Cache provider catalog lookups

This commit is contained in:
fawney19
2026-05-27 13:56:39 +08:00
parent 674cc85005
commit ccfc4cbddc
16 changed files with 582 additions and 75 deletions
+70 -14
View File
@@ -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)
}
}
+12 -1
View File
@@ -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());
+2
View File
@@ -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();
}