mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-10 05:00:19 +08:00
Cache provider catalog lookups
This commit is contained in:
@@ -334,7 +334,7 @@ impl GatewayDataState {
|
||||
encrypted_auth_config: Option<&str>,
|
||||
expires_at_unix_secs: Option<u64>,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.update_key_oauth_credentials(
|
||||
@@ -346,17 +346,25 @@ impl GatewayDataState {
|
||||
.await
|
||||
}
|
||||
None => Ok(false),
|
||||
}?;
|
||||
if updated {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn create_provider_catalog_key(
|
||||
&self,
|
||||
key: &StoredProviderCatalogKey,
|
||||
) -> Result<Option<StoredProviderCatalogKey>, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let created = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.create_key(key).await.map(Some),
|
||||
None => Ok(None),
|
||||
}?;
|
||||
if created.is_some() {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(created)
|
||||
}
|
||||
|
||||
pub(crate) async fn create_provider_catalog_provider(
|
||||
@@ -364,33 +372,45 @@ impl GatewayDataState {
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
shift_existing_priorities_from: Option<i32>,
|
||||
) -> Result<Option<StoredProviderCatalogProvider>, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let created = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository
|
||||
.create_provider(provider, shift_existing_priorities_from)
|
||||
.await
|
||||
.map(Some),
|
||||
None => Ok(None),
|
||||
}?;
|
||||
if created.is_some() {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(created)
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_provider(
|
||||
&self,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
) -> Result<Option<StoredProviderCatalogProvider>, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.update_provider(provider).await.map(Some),
|
||||
None => Ok(None),
|
||||
}?;
|
||||
if updated.is_some() {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn delete_provider_catalog_provider(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let deleted = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.delete_provider(provider_id).await,
|
||||
None => Ok(false),
|
||||
}?;
|
||||
if deleted {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(deleted)
|
||||
}
|
||||
|
||||
pub(crate) async fn cleanup_deleted_provider_catalog_refs(
|
||||
@@ -399,54 +419,74 @@ impl GatewayDataState {
|
||||
endpoint_ids: &[String],
|
||||
key_ids: &[String],
|
||||
) -> Result<(), DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let cleaned = match &self.provider_catalog_writer {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.cleanup_deleted_provider_refs(provider_id, endpoint_ids, key_ids)
|
||||
.await
|
||||
}
|
||||
None => Ok(()),
|
||||
};
|
||||
if !endpoint_ids.is_empty() || !key_ids.is_empty() {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
cleaned
|
||||
}
|
||||
|
||||
pub(crate) async fn create_provider_catalog_endpoint(
|
||||
&self,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
) -> Result<Option<StoredProviderCatalogEndpoint>, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let created = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.create_endpoint(endpoint).await.map(Some),
|
||||
None => Ok(None),
|
||||
}?;
|
||||
if created.is_some() {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(created)
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_endpoint(
|
||||
&self,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
) -> Result<Option<StoredProviderCatalogEndpoint>, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.update_endpoint(endpoint).await.map(Some),
|
||||
None => Ok(None),
|
||||
}?;
|
||||
if updated.is_some() {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn delete_provider_catalog_endpoint(
|
||||
&self,
|
||||
endpoint_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let deleted = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.delete_endpoint(endpoint_id).await,
|
||||
None => Ok(false),
|
||||
}?;
|
||||
if deleted {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(deleted)
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_key(
|
||||
&self,
|
||||
key: &StoredProviderCatalogKey,
|
||||
) -> Result<Option<StoredProviderCatalogKey>, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.update_key(key).await.map(Some),
|
||||
None => Ok(None),
|
||||
}?;
|
||||
if updated.is_some() {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_key_upstream_metadata(
|
||||
@@ -455,34 +495,46 @@ impl GatewayDataState {
|
||||
upstream_metadata: Option<&serde_json::Value>,
|
||||
updated_at_unix_secs: Option<u64>,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.update_key_upstream_metadata(key_id, upstream_metadata, updated_at_unix_secs)
|
||||
.await
|
||||
}
|
||||
None => Ok(false),
|
||||
}?;
|
||||
if updated {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn delete_provider_catalog_key(
|
||||
&self,
|
||||
key_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let deleted = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.delete_key(key_id).await,
|
||||
None => Ok(false),
|
||||
}?;
|
||||
if deleted {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(deleted)
|
||||
}
|
||||
|
||||
pub(crate) async fn clear_provider_catalog_key_oauth_invalid_marker(
|
||||
&self,
|
||||
key_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.clear_key_oauth_invalid_marker(key_id).await,
|
||||
None => Ok(false),
|
||||
}?;
|
||||
if updated {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_key_health_state(
|
||||
@@ -492,7 +544,7 @@ impl GatewayDataState {
|
||||
health_by_format: Option<&serde_json::Value>,
|
||||
circuit_breaker_by_format: Option<&serde_json::Value>,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
match &self.provider_catalog_writer {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.update_key_health_state(
|
||||
@@ -504,6 +556,10 @@ impl GatewayDataState {
|
||||
.await
|
||||
}
|
||||
None => Ok(false),
|
||||
}?;
|
||||
if updated {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
use aether_data::{DataBackends, DataLayerError, DatabaseDriver};
|
||||
use aether_data_contracts::repository::candidate_selection::MinimalCandidateSelectionReadRepository;
|
||||
use aether_data_contracts::repository::provider_catalog::ProviderCatalogReadRepository;
|
||||
use aether_runtime_state::RuntimeQueueStore;
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -99,7 +100,11 @@ impl GatewayDataState {
|
||||
let request_candidate_reader = backends.read().request_candidates();
|
||||
let request_candidate_writer = backends.write().request_candidates();
|
||||
let gemini_file_mapping_writer = backends.write().gemini_file_mappings();
|
||||
let provider_catalog_reader = backends.read().provider_catalog();
|
||||
let provider_catalog_reader = backends.read().provider_catalog().map(|repository| {
|
||||
Arc::new(
|
||||
super::provider_catalog_cache::CachedProviderCatalogReadRepository::new(repository),
|
||||
) as Arc<dyn ProviderCatalogReadRepository>
|
||||
});
|
||||
let provider_catalog_writer = backends.write().provider_catalog();
|
||||
let pool_score_reader = backends.read().pool_scores();
|
||||
let pool_score_writer = backends.write().pool_scores();
|
||||
@@ -276,6 +281,12 @@ impl GatewayDataState {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn clear_provider_catalog_cache(&self) {
|
||||
if let Some(repository) = &self.provider_catalog_reader {
|
||||
repository.clear_local_cache();
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn has_request_candidate_reader(&self) -> bool {
|
||||
self.request_candidate_reader.is_some()
|
||||
}
|
||||
|
||||
@@ -324,6 +324,7 @@ mod core;
|
||||
mod integrations;
|
||||
mod models;
|
||||
mod pool_scores;
|
||||
mod provider_catalog_cache;
|
||||
mod referrals;
|
||||
mod routing_profiles;
|
||||
mod runtime;
|
||||
|
||||
@@ -0,0 +1,360 @@
|
||||
use std::collections::HashMap;
|
||||
use std::future::Future;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::Duration;
|
||||
|
||||
use aether_cache::ExpiringMap;
|
||||
use aether_data::DataLayerError;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyListQuery, ProviderCatalogReadRepository, StoredProviderCatalogEndpoint,
|
||||
StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary,
|
||||
StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use tokio::sync::Notify;
|
||||
|
||||
const PROVIDER_CATALOG_CACHE_TTL: Duration = Duration::from_secs(5);
|
||||
const PROVIDER_CATALOG_CACHE_MAX_ENTRIES: usize = 1024;
|
||||
|
||||
pub(super) struct CachedProviderCatalogReadRepository {
|
||||
inner: Arc<dyn ProviderCatalogReadRepository>,
|
||||
entries: ExpiringMap<ProviderCatalogCacheKey, ProviderCatalogCacheValue>,
|
||||
inflight: Mutex<HashMap<ProviderCatalogCacheKey, u64>>,
|
||||
inflight_notify: Notify,
|
||||
next_inflight_token: AtomicU64,
|
||||
epoch: AtomicU64,
|
||||
}
|
||||
|
||||
impl CachedProviderCatalogReadRepository {
|
||||
pub(super) fn new(inner: Arc<dyn ProviderCatalogReadRepository>) -> Self {
|
||||
Self {
|
||||
inner,
|
||||
entries: ExpiringMap::new(),
|
||||
inflight: Mutex::new(HashMap::new()),
|
||||
inflight_notify: Notify::new(),
|
||||
next_inflight_token: AtomicU64::new(1),
|
||||
epoch: AtomicU64::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_or_load<F, Fut>(
|
||||
&self,
|
||||
key: ProviderCatalogCacheKey,
|
||||
load: F,
|
||||
) -> Result<ProviderCatalogCacheValue, DataLayerError>
|
||||
where
|
||||
F: Fn() -> Fut,
|
||||
Fut: Future<Output = Result<ProviderCatalogCacheValue, DataLayerError>>,
|
||||
{
|
||||
if let Some(value) = self.entries.get_fresh(&key, PROVIDER_CATALOG_CACHE_TTL) {
|
||||
return Ok(value);
|
||||
}
|
||||
|
||||
loop {
|
||||
let notified = self.inflight_notify.notified();
|
||||
match self.register_inflight(&key) {
|
||||
InflightRegistration::Bypass => return load().await,
|
||||
InflightRegistration::Follower => {
|
||||
notified.await;
|
||||
if let Some(value) = self.entries.get_fresh(&key, PROVIDER_CATALOG_CACHE_TTL) {
|
||||
return Ok(value);
|
||||
}
|
||||
}
|
||||
InflightRegistration::Leader(token) => {
|
||||
let mut guard = InflightGuard::new(self, key.clone(), token);
|
||||
let load_epoch = self.epoch.load(Ordering::Acquire);
|
||||
let result = load().await;
|
||||
if let Ok(value) = &result {
|
||||
if load_epoch == self.epoch.load(Ordering::Acquire) {
|
||||
self.entries.insert(
|
||||
key.clone(),
|
||||
value.clone(),
|
||||
PROVIDER_CATALOG_CACHE_TTL,
|
||||
PROVIDER_CATALOG_CACHE_MAX_ENTRIES,
|
||||
);
|
||||
}
|
||||
}
|
||||
guard.finish();
|
||||
return result;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn register_inflight(&self, key: &ProviderCatalogCacheKey) -> InflightRegistration {
|
||||
match self.inflight.lock() {
|
||||
Ok(mut inflight) => {
|
||||
if inflight.contains_key(key) {
|
||||
return InflightRegistration::Follower;
|
||||
}
|
||||
let token = self.next_inflight_token.fetch_add(1, Ordering::AcqRel);
|
||||
inflight.insert(key.clone(), token);
|
||||
InflightRegistration::Leader(token)
|
||||
}
|
||||
Err(_) => InflightRegistration::Bypass,
|
||||
}
|
||||
}
|
||||
|
||||
fn finish_inflight(&self, key: &ProviderCatalogCacheKey, token: u64) {
|
||||
let mut removed = false;
|
||||
if let Ok(mut inflight) = self.inflight.lock() {
|
||||
if inflight.get(key).copied() == Some(token) {
|
||||
inflight.remove(key);
|
||||
removed = true;
|
||||
}
|
||||
}
|
||||
if removed {
|
||||
self.inflight_notify.notify_waiters();
|
||||
}
|
||||
}
|
||||
|
||||
fn clear(&self) {
|
||||
self.epoch.fetch_add(1, Ordering::AcqRel);
|
||||
self.entries.clear();
|
||||
let mut cleared_inflight = false;
|
||||
if let Ok(mut inflight) = self.inflight.lock() {
|
||||
cleared_inflight = !inflight.is_empty();
|
||||
inflight.clear();
|
||||
}
|
||||
if cleared_inflight {
|
||||
self.inflight_notify.notify_waiters();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ProviderCatalogReadRepository for CachedProviderCatalogReadRepository {
|
||||
fn clear_local_cache(&self) {
|
||||
self.clear();
|
||||
self.inner.clear_local_cache();
|
||||
}
|
||||
|
||||
async fn list_providers(
|
||||
&self,
|
||||
active_only: bool,
|
||||
) -> Result<Vec<StoredProviderCatalogProvider>, DataLayerError> {
|
||||
match self
|
||||
.get_or_load(
|
||||
ProviderCatalogCacheKey::Providers { active_only },
|
||||
|| async move {
|
||||
self.inner
|
||||
.list_providers(active_only)
|
||||
.await
|
||||
.map(ProviderCatalogCacheValue::Providers)
|
||||
},
|
||||
)
|
||||
.await?
|
||||
{
|
||||
ProviderCatalogCacheValue::Providers(items) => Ok(items),
|
||||
_ => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_providers_by_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogProvider>, DataLayerError> {
|
||||
let key = ProviderCatalogCacheKey::ProvidersByIds(normalize_ids(provider_ids));
|
||||
match self
|
||||
.get_or_load(key, || async move {
|
||||
self.inner
|
||||
.list_providers_by_ids(provider_ids)
|
||||
.await
|
||||
.map(ProviderCatalogCacheValue::Providers)
|
||||
})
|
||||
.await?
|
||||
{
|
||||
ProviderCatalogCacheValue::Providers(items) => Ok(items),
|
||||
_ => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_endpoints_by_ids(
|
||||
&self,
|
||||
endpoint_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogEndpoint>, DataLayerError> {
|
||||
self.inner.list_endpoints_by_ids(endpoint_ids).await
|
||||
}
|
||||
|
||||
async fn list_endpoints_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogEndpoint>, DataLayerError> {
|
||||
let key = ProviderCatalogCacheKey::EndpointsByProviderIds(normalize_ids(provider_ids));
|
||||
match self
|
||||
.get_or_load(key, || async move {
|
||||
self.inner
|
||||
.list_endpoints_by_provider_ids(provider_ids)
|
||||
.await
|
||||
.map(ProviderCatalogCacheValue::Endpoints)
|
||||
})
|
||||
.await?
|
||||
{
|
||||
ProviderCatalogCacheValue::Endpoints(items) => Ok(items),
|
||||
_ => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_keys_by_ids(
|
||||
&self,
|
||||
key_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
self.inner.list_keys_by_ids(key_ids).await
|
||||
}
|
||||
|
||||
async fn list_keys_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
let key = ProviderCatalogCacheKey::KeysByProviderIds(normalize_ids(provider_ids));
|
||||
match self
|
||||
.get_or_load(key, || async move {
|
||||
self.inner
|
||||
.list_keys_by_provider_ids(provider_ids)
|
||||
.await
|
||||
.map(ProviderCatalogCacheValue::Keys)
|
||||
})
|
||||
.await?
|
||||
{
|
||||
ProviderCatalogCacheValue::Keys(items) => Ok(items),
|
||||
_ => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_key_summaries_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
let key = ProviderCatalogCacheKey::KeySummariesByProviderIds(normalize_ids(provider_ids));
|
||||
match self
|
||||
.get_or_load(key, || async move {
|
||||
self.inner
|
||||
.list_key_summaries_by_provider_ids(provider_ids)
|
||||
.await
|
||||
.map(ProviderCatalogCacheValue::Keys)
|
||||
})
|
||||
.await?
|
||||
{
|
||||
ProviderCatalogCacheValue::Keys(items) => Ok(items),
|
||||
_ => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_key_maintenance_summaries_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKeyMaintenanceSummary>, DataLayerError> {
|
||||
let key = ProviderCatalogCacheKey::KeyMaintenanceSummariesByProviderIds(normalize_ids(
|
||||
provider_ids,
|
||||
));
|
||||
match self
|
||||
.get_or_load(key, || async move {
|
||||
self.inner
|
||||
.list_key_maintenance_summaries_by_provider_ids(provider_ids)
|
||||
.await
|
||||
.map(ProviderCatalogCacheValue::KeyMaintenanceSummaries)
|
||||
})
|
||||
.await?
|
||||
{
|
||||
ProviderCatalogCacheValue::KeyMaintenanceSummaries(items) => Ok(items),
|
||||
_ => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_keys_page(
|
||||
&self,
|
||||
query: &ProviderCatalogKeyListQuery,
|
||||
) -> Result<StoredProviderCatalogKeyPage, DataLayerError> {
|
||||
self.inner.list_keys_page(query).await
|
||||
}
|
||||
|
||||
async fn list_key_stats_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKeyStats>, DataLayerError> {
|
||||
let key = ProviderCatalogCacheKey::KeyStatsByProviderIds(normalize_ids(provider_ids));
|
||||
match self
|
||||
.get_or_load(key, || async move {
|
||||
self.inner
|
||||
.list_key_stats_by_provider_ids(provider_ids)
|
||||
.await
|
||||
.map(ProviderCatalogCacheValue::KeyStats)
|
||||
})
|
||||
.await?
|
||||
{
|
||||
ProviderCatalogCacheValue::KeyStats(items) => Ok(items),
|
||||
_ => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
|
||||
enum ProviderCatalogCacheKey {
|
||||
Providers { active_only: bool },
|
||||
ProvidersByIds(Vec<String>),
|
||||
EndpointsByProviderIds(Vec<String>),
|
||||
KeysByProviderIds(Vec<String>),
|
||||
KeySummariesByProviderIds(Vec<String>),
|
||||
KeyMaintenanceSummariesByProviderIds(Vec<String>),
|
||||
KeyStatsByProviderIds(Vec<String>),
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
enum ProviderCatalogCacheValue {
|
||||
Providers(Vec<StoredProviderCatalogProvider>),
|
||||
Endpoints(Vec<StoredProviderCatalogEndpoint>),
|
||||
Keys(Vec<StoredProviderCatalogKey>),
|
||||
KeyMaintenanceSummaries(Vec<StoredProviderCatalogKeyMaintenanceSummary>),
|
||||
KeyStats(Vec<StoredProviderCatalogKeyStats>),
|
||||
}
|
||||
|
||||
enum InflightRegistration {
|
||||
Leader(u64),
|
||||
Follower,
|
||||
Bypass,
|
||||
}
|
||||
|
||||
struct InflightGuard<'a> {
|
||||
cache: &'a CachedProviderCatalogReadRepository,
|
||||
key: Option<ProviderCatalogCacheKey>,
|
||||
token: u64,
|
||||
}
|
||||
|
||||
impl<'a> InflightGuard<'a> {
|
||||
fn new(
|
||||
cache: &'a CachedProviderCatalogReadRepository,
|
||||
key: ProviderCatalogCacheKey,
|
||||
token: u64,
|
||||
) -> Self {
|
||||
Self {
|
||||
cache,
|
||||
key: Some(key),
|
||||
token,
|
||||
}
|
||||
}
|
||||
|
||||
fn finish(&mut self) {
|
||||
if let Some(key) = self.key.take() {
|
||||
self.cache.finish_inflight(&key, self.token);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for InflightGuard<'_> {
|
||||
fn drop(&mut self) {
|
||||
self.finish();
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_ids(ids: &[String]) -> Vec<String> {
|
||||
let mut normalized = ids
|
||||
.iter()
|
||||
.map(|id| id.trim())
|
||||
.filter(|id| !id.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.collect::<Vec<_>>();
|
||||
normalized.sort();
|
||||
normalized.dedup();
|
||||
normalized
|
||||
}
|
||||
@@ -2,7 +2,7 @@ use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
|
||||
use aether_data_contracts::repository::pool_scores::{
|
||||
GetPoolMemberScoresByIdsQuery, ListPoolMemberProbeCandidatesQuery, PoolMemberHardState,
|
||||
ListPoolMemberProbeCandidatesQuery, ListPoolMemberScoresQuery, PoolMemberHardState,
|
||||
PoolMemberIdentity, PoolMemberProbeAttempt, PoolMemberProbeResult, PoolMemberProbeStatus,
|
||||
StoredPoolMemberScore, POOL_KIND_PROVIDER_KEY_POOL,
|
||||
};
|
||||
@@ -22,7 +22,6 @@ use crate::admin_api::{
|
||||
};
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
use crate::ai_serving::provider_key_pool_score_id;
|
||||
use crate::ai_serving::provider_key_pool_score_scope;
|
||||
use crate::handlers::shared::provider_pool::{
|
||||
admin_provider_pool_quota_probe_active_members_key, AdminProviderPoolConfig,
|
||||
@@ -38,6 +37,7 @@ const POOL_QUOTA_PROBE_REDIS_PREFIX: &str = "ap:quota_probe:last";
|
||||
const POOL_QUOTA_PROBE_DEFAULT_SCAN_INTERVAL_SECONDS: u64 = 60;
|
||||
const POOL_QUOTA_PROBE_MIN_SCAN_INTERVAL_SECONDS: u64 = 15;
|
||||
const POOL_QUOTA_PROBE_DEFAULT_MAX_KEYS_PER_PROVIDER: usize = 50;
|
||||
const POOL_QUOTA_PROBE_PROVIDER_SCORE_READ_LIMIT: usize = 100_000;
|
||||
const POOL_QUOTA_PROBE_DEFAULT_GLOBAL_CONCURRENCY: usize = 16;
|
||||
const POOL_QUOTA_PROBE_PROVIDER_LOCK_TTL_MS: u64 = 30_000;
|
||||
const POOL_QUOTA_PROBE_BURST_TRIGGER_LOCK_TTL_MS: u64 = 30_000;
|
||||
@@ -386,21 +386,25 @@ async fn load_provider_key_account_scores(
|
||||
return BTreeMap::new();
|
||||
}
|
||||
let scope = provider_key_pool_score_scope();
|
||||
let score_ids = key_ids
|
||||
.iter()
|
||||
.map(|key_id| {
|
||||
let identity =
|
||||
PoolMemberIdentity::provider_api_key(provider_id.to_string(), key_id.clone());
|
||||
provider_key_pool_score_id(&identity, &scope)
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
let key_ids = key_ids.iter().map(String::as_str).collect::<BTreeSet<_>>();
|
||||
match state
|
||||
.data
|
||||
.get_pool_member_scores_by_ids(&GetPoolMemberScoresByIdsQuery { ids: score_ids })
|
||||
.list_pool_member_scores(&ListPoolMemberScoresQuery {
|
||||
pool_kind: POOL_KIND_PROVIDER_KEY_POOL.to_string(),
|
||||
pool_id: provider_id.to_string(),
|
||||
capability: Some(scope.capability),
|
||||
scope_kind: Some(scope.scope_kind),
|
||||
scope_id: scope.scope_id,
|
||||
hard_states: Vec::new(),
|
||||
probe_statuses: None,
|
||||
offset: 0,
|
||||
limit: POOL_QUOTA_PROBE_PROVIDER_SCORE_READ_LIMIT,
|
||||
})
|
||||
.await
|
||||
{
|
||||
Ok(scores) => scores
|
||||
.into_iter()
|
||||
.filter(|score| key_ids.contains(score.member_id.as_str()))
|
||||
.map(|score| (score.member_id.clone(), score))
|
||||
.collect(),
|
||||
Err(err) => {
|
||||
|
||||
@@ -645,6 +645,11 @@ async fn record_health_success_effect(
|
||||
.circuit_breaker_by_format
|
||||
.as_ref()
|
||||
.and_then(|current| project_local_key_circuit_closed(Some(current), api_format));
|
||||
if current_key.health_by_format.as_ref() == Some(&health_by_format)
|
||||
&& circuit_breaker_by_format.as_ref() == current_key.circuit_breaker_by_format.as_ref()
|
||||
{
|
||||
return;
|
||||
}
|
||||
let circuit_breaker_update = circuit_breaker_by_format
|
||||
.as_ref()
|
||||
.or(current_key.circuit_breaker_by_format.as_ref());
|
||||
|
||||
@@ -573,12 +573,14 @@ impl AppState {
|
||||
|
||||
pub(crate) fn invalidate_provider_routing_caches(&self) {
|
||||
self.data.clear_minimal_candidate_selection_cache();
|
||||
self.data.clear_provider_catalog_cache();
|
||||
self.clear_provider_transport_snapshot_cache();
|
||||
self.invalidate_scheduler_affinity_cache();
|
||||
}
|
||||
|
||||
pub(crate) fn invalidate_provider_health_routing_caches(&self) {
|
||||
self.data.clear_minimal_candidate_selection_cache();
|
||||
self.data.clear_provider_catalog_cache();
|
||||
self.clear_provider_transport_snapshot_cache();
|
||||
}
|
||||
|
||||
|
||||
@@ -580,6 +580,8 @@ impl StoredProviderCatalogKeyStats {
|
||||
|
||||
#[async_trait]
|
||||
pub trait ProviderCatalogReadRepository: Send + Sync {
|
||||
fn clear_local_cache(&self) {}
|
||||
|
||||
async fn list_providers(
|
||||
&self,
|
||||
active_only: bool,
|
||||
|
||||
+15
@@ -0,0 +1,15 @@
|
||||
SET @aether_provider_key_name_index_sql := IF(
|
||||
(
|
||||
SELECT COUNT(*)
|
||||
FROM information_schema.statistics
|
||||
WHERE table_schema = DATABASE()
|
||||
AND table_name = 'provider_api_keys'
|
||||
AND index_name = 'idx_provider_api_keys_provider_name_id'
|
||||
) = 0,
|
||||
'CREATE INDEX idx_provider_api_keys_provider_name_id ON provider_api_keys (provider_id, name, id)',
|
||||
'DO 0'
|
||||
);
|
||||
|
||||
PREPARE aether_provider_key_name_index_stmt FROM @aether_provider_key_name_index_sql;
|
||||
EXECUTE aether_provider_key_name_index_stmt;
|
||||
DEALLOCATE PREPARE aether_provider_key_name_index_stmt;
|
||||
+2
@@ -0,0 +1,2 @@
|
||||
CREATE INDEX IF NOT EXISTS idx_provider_api_keys_provider_name_id
|
||||
ON public.provider_api_keys USING btree (provider_id, name, id);
|
||||
+2
@@ -0,0 +1,2 @@
|
||||
CREATE INDEX IF NOT EXISTS idx_provider_api_keys_provider_name_id
|
||||
ON provider_api_keys (provider_id, name, id);
|
||||
@@ -197,6 +197,14 @@ CREATE INDEX IF NOT EXISTS idx_provider_api_keys_provider_default_sort ON public
|
||||
|
||||
|
||||
|
||||
--
|
||||
-- Name: idx_provider_api_keys_provider_name_id; Type: INDEX; Schema: public; Owner: -
|
||||
--
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_provider_api_keys_provider_name_id ON public.provider_api_keys USING btree (provider_id, name, id);
|
||||
|
||||
|
||||
|
||||
--
|
||||
-- Name: idx_provider_api_keys_provider_active_priority_id; Type: INDEX; Schema: public; Owner: -
|
||||
--
|
||||
|
||||
@@ -7,7 +7,7 @@ use tracing::info;
|
||||
// Generated by build.rs from schema/bootstrap/postgres.
|
||||
pub(crate) static EMPTY_DATABASE_SNAPSHOT_SQL: &str =
|
||||
include_str!(concat!(env!("OUT_DIR"), "/empty_database_snapshot.sql"));
|
||||
pub(crate) const EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION: i64 = 20260524000000;
|
||||
pub(crate) const EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION: i64 = 20260527000000;
|
||||
|
||||
const PUBLIC_BASE_TABLE_COUNT_SQL: &str = r#"
|
||||
SELECT COUNT(*)::BIGINT
|
||||
|
||||
@@ -314,6 +314,7 @@ fn empty_database_snapshot_covers_current_cutoff_versions() {
|
||||
20260520010000,
|
||||
20260522000000,
|
||||
20260524000000,
|
||||
20260527000000,
|
||||
]
|
||||
);
|
||||
}
|
||||
@@ -390,6 +391,7 @@ fn empty_database_snapshot_sql_includes_usage_body_blobs_and_audit_admin_role()
|
||||
assert!(EMPTY_DATABASE_SNAPSHOT_SQL.contains("ix_usage_counter_deltas_unprocessed"));
|
||||
assert!(EMPTY_DATABASE_SNAPSHOT_SQL.contains("idx_entitlement_usage_entitlement_date"));
|
||||
assert!(EMPTY_DATABASE_SNAPSHOT_SQL.contains("idx_provider_api_keys_provider_default_sort"));
|
||||
assert!(EMPTY_DATABASE_SNAPSHOT_SQL.contains("idx_provider_api_keys_provider_name_id"));
|
||||
assert!(
|
||||
EMPTY_DATABASE_SNAPSHOT_SQL.contains("idx_provider_api_keys_provider_active_priority_id")
|
||||
);
|
||||
@@ -668,6 +670,7 @@ fn mysql_and_sqlite_migrations_include_enabled_incrementals() {
|
||||
20260520000000,
|
||||
20260520010000,
|
||||
20260524000000,
|
||||
20260527000000,
|
||||
]
|
||||
);
|
||||
assert_eq!(
|
||||
@@ -692,6 +695,7 @@ fn mysql_and_sqlite_migrations_include_enabled_incrementals() {
|
||||
20260520000000,
|
||||
20260520010000,
|
||||
20260524000000,
|
||||
20260527000000,
|
||||
]
|
||||
);
|
||||
}
|
||||
@@ -1217,6 +1221,7 @@ fn pending_migrations_from_applied_skips_versions_already_applied() {
|
||||
20260520010000,
|
||||
20260522000000,
|
||||
20260524000000,
|
||||
20260527000000,
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
@@ -341,14 +341,13 @@ impl PoolScoreReadRepository for PostgresPoolMemberScoreRepository {
|
||||
if query.ids.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
let mut builder = QueryBuilder::<Postgres>::new(SCORE_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_in(&mut builder, &mut where_clause, "id", &query.ids);
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let rows = sqlx::query(&format!(
|
||||
"{SCORE_COLUMNS} WHERE id = ANY($1) ORDER BY id ASC"
|
||||
))
|
||||
.bind(&query.ids)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_score_row).collect()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, Row};
|
||||
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
|
||||
|
||||
use super::{
|
||||
InMemoryProviderCatalogReadRepository, ProviderCatalogKeyListQuery,
|
||||
@@ -16,6 +16,33 @@ pub struct MysqlProviderCatalogReadRepository {
|
||||
pool: MysqlPool,
|
||||
}
|
||||
|
||||
const KEY_SELECT_SQL: &str = r#"
|
||||
SELECT
|
||||
id, provider_id, name, auth_type, capabilities, is_active, api_formats,
|
||||
auth_type_by_format, allow_auth_channel_mismatch_formats,
|
||||
COALESCE(api_key, encrypted_key) AS api_key,
|
||||
auth_config, note, internal_priority, rate_multipliers,
|
||||
global_priority_by_format, allowed_models,
|
||||
expires_at AS expires_at_unix_secs,
|
||||
cache_ttl_minutes, max_probe_interval_minutes, proxy, fingerprint,
|
||||
rpm_limit, concurrent_limit, learned_rpm_limit, concurrent_429_count,
|
||||
rpm_429_count, last_429_at AS last_429_at_unix_secs, last_429_type,
|
||||
adjustment_history, utilization_samples,
|
||||
last_probe_increase_at AS last_probe_increase_at_unix_secs,
|
||||
last_rpm_peak, request_count, total_tokens, total_cost_usd,
|
||||
success_count, error_count, total_response_time_ms,
|
||||
last_used_at AS last_used_at_unix_secs, auto_fetch_models,
|
||||
last_models_fetch_at AS last_models_fetch_at_unix_secs,
|
||||
last_models_fetch_error, locked_models, model_include_patterns,
|
||||
model_exclude_patterns, upstream_metadata,
|
||||
oauth_invalid_at AS oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason, status_snapshot,
|
||||
created_at AS created_at_unix_ms,
|
||||
updated_at AS updated_at_unix_secs,
|
||||
health_by_format, circuit_breaker_by_format
|
||||
FROM provider_api_keys
|
||||
"#;
|
||||
|
||||
impl MysqlProviderCatalogReadRepository {
|
||||
pub fn new(pool: MysqlPool) -> Self {
|
||||
Self { pool }
|
||||
@@ -71,37 +98,26 @@ WHERE api_format IS NOT NULL
|
||||
}
|
||||
|
||||
async fn load_keys(&self) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
id, provider_id, name, auth_type, capabilities, is_active, api_formats,
|
||||
auth_type_by_format, allow_auth_channel_mismatch_formats,
|
||||
COALESCE(api_key, encrypted_key) AS api_key,
|
||||
auth_config, note, internal_priority, rate_multipliers,
|
||||
global_priority_by_format, allowed_models,
|
||||
expires_at AS expires_at_unix_secs,
|
||||
cache_ttl_minutes, max_probe_interval_minutes, proxy, fingerprint,
|
||||
rpm_limit, concurrent_limit, learned_rpm_limit, concurrent_429_count,
|
||||
rpm_429_count, last_429_at AS last_429_at_unix_secs, last_429_type,
|
||||
adjustment_history, utilization_samples,
|
||||
last_probe_increase_at AS last_probe_increase_at_unix_secs,
|
||||
last_rpm_peak, request_count, total_tokens, total_cost_usd,
|
||||
success_count, error_count, total_response_time_ms,
|
||||
last_used_at AS last_used_at_unix_secs, auto_fetch_models,
|
||||
last_models_fetch_at AS last_models_fetch_at_unix_secs,
|
||||
last_models_fetch_error, locked_models, model_include_patterns,
|
||||
model_exclude_patterns, upstream_metadata,
|
||||
oauth_invalid_at AS oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason, status_snapshot,
|
||||
created_at AS created_at_unix_ms,
|
||||
updated_at AS updated_at_unix_secs,
|
||||
health_by_format, circuit_breaker_by_format
|
||||
FROM provider_api_keys
|
||||
"#,
|
||||
)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let rows = sqlx::query(KEY_SELECT_SQL)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_key_row).collect()
|
||||
}
|
||||
|
||||
async fn list_keys_by_provider_ids_direct(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
if provider_ids.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let rows = build_list_keys_by_provider_ids_query(provider_ids)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_key_row).collect()
|
||||
}
|
||||
|
||||
@@ -912,20 +928,14 @@ impl ProviderCatalogReadRepository for MysqlProviderCatalogReadRepository {
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_keys_by_provider_ids(provider_ids)
|
||||
.await
|
||||
self.list_keys_by_provider_ids_direct(provider_ids).await
|
||||
}
|
||||
|
||||
async fn list_key_summaries_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_key_summaries_by_provider_ids(provider_ids)
|
||||
.await
|
||||
self.list_keys_by_provider_ids_direct(provider_ids).await
|
||||
}
|
||||
|
||||
async fn list_key_maintenance_summaries_by_provider_ids(
|
||||
@@ -1162,6 +1172,19 @@ fn optional_json_to_string(
|
||||
optional_json_ref_to_string(value.as_ref(), field_name)
|
||||
}
|
||||
|
||||
fn build_list_keys_by_provider_ids_query(provider_ids: &[String]) -> QueryBuilder<'_, MySql> {
|
||||
let mut builder = QueryBuilder::<MySql>::new(KEY_SELECT_SQL);
|
||||
builder.push("WHERE provider_id IN (");
|
||||
{
|
||||
let mut separated = builder.separated(", ");
|
||||
for provider_id in provider_ids {
|
||||
separated.push_bind(provider_id.clone());
|
||||
}
|
||||
}
|
||||
builder.push(") ORDER BY provider_id ASC, name ASC, id ASC");
|
||||
builder
|
||||
}
|
||||
|
||||
fn key_insert_sql() -> &'static str {
|
||||
r#"
|
||||
INSERT INTO provider_api_keys (
|
||||
@@ -1568,6 +1591,7 @@ mod tests {
|
||||
StoredProviderCatalogProvider,
|
||||
};
|
||||
use serde_json::json;
|
||||
use sqlx::Execute;
|
||||
|
||||
#[tokio::test]
|
||||
async fn repository_builds_from_lazy_pool() {
|
||||
@@ -1580,6 +1604,17 @@ mod tests {
|
||||
let _repository = MysqlProviderCatalogReadRepository::new(pool);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn list_keys_by_provider_ids_query_targets_index_aligned_ordering() {
|
||||
let provider_ids = vec!["provider-a".to_string(), "provider-b".to_string()];
|
||||
let mut builder = super::build_list_keys_by_provider_ids_query(&provider_ids);
|
||||
let query = builder.build();
|
||||
let sql = query.sql();
|
||||
|
||||
assert!(sql.contains("WHERE provider_id IN ("));
|
||||
assert!(sql.contains("ORDER BY provider_id ASC, name ASC, id ASC"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mysql_provider_catalog_repository_round_trips_when_url_is_set() {
|
||||
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
|
||||
|
||||
Reference in New Issue
Block a user