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();
}
@@ -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,
@@ -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;
@@ -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);
@@ -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")