mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-07 01:47:47 +08:00
Improve gateway transport and usage runtime
This commit is contained in:
@@ -368,6 +368,12 @@ impl GatewayDataState {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn clear_routing_group_cache(&self) {
|
||||
if let Some(repository) = &self.routing_group_reader {
|
||||
repository.clear_local_cache();
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn clear_provider_catalog_cache(&self) {
|
||||
if let Some(repository) = &self.provider_catalog_reader {
|
||||
repository.clear_local_cache();
|
||||
|
||||
@@ -360,3 +360,5 @@ mod routing_profiles;
|
||||
mod runtime;
|
||||
#[cfg(test)]
|
||||
mod testing;
|
||||
#[cfg(feature = "testkit")]
|
||||
pub(crate) mod testkit;
|
||||
|
||||
@@ -9,14 +9,16 @@ use aether_data_contracts::repository::routing_profiles::{
|
||||
StoredRoutingGroupVersion,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use dashmap::DashMap;
|
||||
|
||||
const ROUTING_GROUP_CACHE_TTL: Duration = Duration::from_secs(5);
|
||||
const ROUTING_GROUP_CACHE_STALE_TTL: Duration = Duration::from_secs(60);
|
||||
const ROUTING_GROUP_CACHE_MAX_ENTRIES: usize = 4_096;
|
||||
const ROUTING_GROUP_CACHE_MAX_LOAD_GUARDS: usize = 4_096;
|
||||
|
||||
pub(super) struct CachedRoutingGroupReadRepository {
|
||||
inner: Arc<dyn RoutingGroupReadRepository>,
|
||||
entries: ExpiringMap<RoutingGroupCacheKey, RoutingGroupCacheValue>,
|
||||
load_guard: tokio::sync::Mutex<()>,
|
||||
load_guards: DashMap<RoutingGroupCacheKey, Arc<tokio::sync::Mutex<()>>>,
|
||||
}
|
||||
|
||||
impl CachedRoutingGroupReadRepository {
|
||||
@@ -24,31 +26,53 @@ impl CachedRoutingGroupReadRepository {
|
||||
Self {
|
||||
inner,
|
||||
entries: ExpiringMap::new(),
|
||||
load_guard: tokio::sync::Mutex::new(()),
|
||||
load_guards: DashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn clear(&self) {
|
||||
self.entries.clear();
|
||||
self.load_guards.clear();
|
||||
}
|
||||
|
||||
async fn get_or_load(
|
||||
&self,
|
||||
key: RoutingGroupCacheKey,
|
||||
load: impl std::future::Future<Output = Result<RoutingGroupCacheValue, DataLayerError>>,
|
||||
) -> Result<RoutingGroupCacheValue, DataLayerError> {
|
||||
if let Some(value) = self.entries.get_fresh(&key, ROUTING_GROUP_CACHE_TTL) {
|
||||
if let Some((value, _age)) = self
|
||||
.entries
|
||||
.get_with_age(&key, ROUTING_GROUP_CACHE_STALE_TTL)
|
||||
{
|
||||
return Ok(value);
|
||||
}
|
||||
let _guard = self.load_guard.lock().await;
|
||||
if let Some(value) = self.entries.get_fresh(&key, ROUTING_GROUP_CACHE_TTL) {
|
||||
let load_guard = self.load_guard_for(&key);
|
||||
let _guard = load_guard.lock().await;
|
||||
if let Some((value, _age)) = self
|
||||
.entries
|
||||
.get_with_age(&key, ROUTING_GROUP_CACHE_STALE_TTL)
|
||||
{
|
||||
return Ok(value);
|
||||
}
|
||||
let value = load.await?;
|
||||
self.entries.insert(
|
||||
key,
|
||||
value.clone(),
|
||||
ROUTING_GROUP_CACHE_TTL,
|
||||
ROUTING_GROUP_CACHE_STALE_TTL,
|
||||
ROUTING_GROUP_CACHE_MAX_ENTRIES,
|
||||
);
|
||||
Ok(value)
|
||||
}
|
||||
|
||||
fn load_guard_for(&self, key: &RoutingGroupCacheKey) -> Arc<tokio::sync::Mutex<()>> {
|
||||
if self.load_guards.len() > ROUTING_GROUP_CACHE_MAX_LOAD_GUARDS {
|
||||
self.load_guards.clear();
|
||||
}
|
||||
self.load_guards
|
||||
.entry(key.clone())
|
||||
.or_insert_with(|| Arc::new(tokio::sync::Mutex::new(())))
|
||||
.clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
|
||||
@@ -92,6 +116,10 @@ fn subject_cache_key(subject: Option<RoutingGroupBindingSubject>) -> Option<&'st
|
||||
|
||||
#[async_trait]
|
||||
impl RoutingGroupReadRepository for CachedRoutingGroupReadRepository {
|
||||
fn clear_local_cache(&self) {
|
||||
self.clear();
|
||||
}
|
||||
|
||||
async fn list_routing_groups(&self) -> Result<Vec<StoredRoutingGroup>, DataLayerError> {
|
||||
match self
|
||||
.get_or_load(RoutingGroupCacheKey::ListGroups, async {
|
||||
@@ -168,3 +196,67 @@ impl RoutingGroupReadRepository for CachedRoutingGroupReadRepository {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
use super::*;
|
||||
|
||||
#[derive(Default)]
|
||||
struct CountingRoutingGroupReadRepository {
|
||||
list_calls: AtomicUsize,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl RoutingGroupReadRepository for CountingRoutingGroupReadRepository {
|
||||
async fn list_routing_groups(&self) -> Result<Vec<StoredRoutingGroup>, DataLayerError> {
|
||||
self.list_calls.fetch_add(1, Ordering::AcqRel);
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn find_routing_group(
|
||||
&self,
|
||||
_lookup: RoutingGroupLookupKey<'_>,
|
||||
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn list_routing_group_bindings(
|
||||
&self,
|
||||
_query: &RoutingGroupBindingQuery,
|
||||
) -> Result<Vec<StoredRoutingGroupBinding>, DataLayerError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn list_routing_group_versions(
|
||||
&self,
|
||||
_group_id: &str,
|
||||
) -> Result<Vec<StoredRoutingGroupVersion>, DataLayerError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn clear_local_cache_forces_next_load() {
|
||||
let inner = Arc::new(CountingRoutingGroupReadRepository::default());
|
||||
let repository = CachedRoutingGroupReadRepository::new(inner.clone());
|
||||
|
||||
repository
|
||||
.list_routing_groups()
|
||||
.await
|
||||
.expect("initial list should load");
|
||||
repository
|
||||
.list_routing_groups()
|
||||
.await
|
||||
.expect("cached list should load");
|
||||
assert_eq!(inner.list_calls.load(Ordering::Acquire), 1);
|
||||
|
||||
repository.clear_local_cache();
|
||||
repository
|
||||
.list_routing_groups()
|
||||
.await
|
||||
.expect("cleared list should reload");
|
||||
assert_eq!(inner.list_calls.load(Ordering::Acquire), 2);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,91 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use aether_data_contracts::repository::candidate_selection::MinimalCandidateSelectionReadRepository;
|
||||
use aether_data_contracts::repository::candidates::{
|
||||
RequestCandidateReadRepository, RequestCandidateRepository, RequestCandidateWriteRepository,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
|
||||
};
|
||||
use aether_data_contracts::repository::usage::{
|
||||
UsageReadRepository, UsageRepository, UsageWriteRepository,
|
||||
};
|
||||
|
||||
use aether_data::repository::auth::AuthApiKeyReadRepository;
|
||||
|
||||
use super::{GatewayDataConfig, GatewayDataState};
|
||||
|
||||
impl GatewayDataState {
|
||||
pub(crate) fn with_openai_chat_pressure_repositories_for_testkit<T, U, V>(
|
||||
auth_api_key_repository: Arc<dyn AuthApiKeyReadRepository>,
|
||||
candidate_selection_repository: Arc<dyn MinimalCandidateSelectionReadRepository>,
|
||||
provider_catalog_repository: Arc<U>,
|
||||
request_candidate_repository: Arc<T>,
|
||||
usage_repository: Arc<V>,
|
||||
encryption_key: impl Into<String>,
|
||||
) -> Self
|
||||
where
|
||||
T: RequestCandidateRepository + 'static,
|
||||
U: ProviderCatalogReadRepository + ProviderCatalogWriteRepository + 'static,
|
||||
V: UsageRepository + 'static,
|
||||
{
|
||||
let request_candidate_reader: Arc<dyn RequestCandidateReadRepository> =
|
||||
request_candidate_repository.clone();
|
||||
let request_candidate_writer: Arc<dyn RequestCandidateWriteRepository> =
|
||||
request_candidate_repository;
|
||||
let provider_catalog_reader: Arc<dyn ProviderCatalogReadRepository> =
|
||||
provider_catalog_repository.clone();
|
||||
let provider_catalog_writer: Arc<dyn ProviderCatalogWriteRepository> =
|
||||
provider_catalog_repository;
|
||||
let usage_reader: Arc<dyn UsageReadRepository> = usage_repository.clone();
|
||||
let usage_writer: Arc<dyn UsageWriteRepository> = usage_repository;
|
||||
|
||||
Self {
|
||||
config: GatewayDataConfig::disabled().with_encryption_key(encryption_key),
|
||||
backends: None,
|
||||
auth_api_key_reader: Some(auth_api_key_repository),
|
||||
auth_api_key_writer: None,
|
||||
auth_module_reader: None,
|
||||
auth_module_writer: None,
|
||||
announcement_reader: None,
|
||||
announcement_writer: None,
|
||||
management_token_reader: None,
|
||||
management_token_writer: None,
|
||||
oauth_provider_reader: None,
|
||||
oauth_provider_writer: None,
|
||||
proxy_node_reader: None,
|
||||
proxy_node_writer: None,
|
||||
billing_reader: None,
|
||||
background_task_reader: None,
|
||||
background_task_writer: None,
|
||||
gemini_file_mapping_reader: None,
|
||||
gemini_file_mapping_writer: None,
|
||||
global_model_reader: None,
|
||||
global_model_writer: None,
|
||||
minimal_candidate_selection_reader: Some(candidate_selection_repository),
|
||||
request_candidate_reader: Some(request_candidate_reader),
|
||||
request_candidate_writer: Some(request_candidate_writer),
|
||||
provider_catalog_reader: Some(provider_catalog_reader),
|
||||
provider_catalog_writer: Some(provider_catalog_writer),
|
||||
pool_score_reader: None,
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: Some(usage_reader),
|
||||
usage_writer: Some(usage_writer),
|
||||
user_reader: None,
|
||||
user_preferences: None,
|
||||
usage_worker_queue: None,
|
||||
video_task_reader: None,
|
||||
video_task_writer: None,
|
||||
wallet_reader: None,
|
||||
wallet_writer: None,
|
||||
settlement_writer: None,
|
||||
system_config_values: None,
|
||||
system_config_value_cache: Default::default(),
|
||||
billing_model_context_cache: Default::default(),
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user