perf(gateway): scale request hot paths for 20k streams

Shard and singleflight hot-path caches, batch and prioritize candidate and usage lifecycle persistence, and extend database and pressure-test instrumentation for 20k concurrent streams.
This commit is contained in:
elky
2026-07-22 02:11:08 +08:00
parent 7756c0913f
commit fc92c4f431
124 changed files with 36325 additions and 3217 deletions
@@ -1,3 +1,5 @@
use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
@@ -5,14 +7,26 @@ use aether_cache::ExpiringMap;
use aether_data::repository::auth::*;
use aether_data::DataLayerError;
use async_trait::async_trait;
use tokio::sync::futures::OwnedNotified;
use tokio::sync::Notify;
// Security revalidation bypasses this throughput-oriented read cache.
const AUTH_API_KEY_SNAPSHOT_CACHE_TTL: Duration = Duration::from_secs(30);
const AUTH_API_KEY_SNAPSHOT_CACHE_MAX_ENTRIES: usize = 16_384;
tokio::task_local! {
static AUTH_API_KEY_READ_CACHE_BYPASS: ();
}
pub(super) struct CachedAuthApiKeyReadRepository {
inner: Arc<dyn AuthApiKeyReadRepository>,
snapshots: ExpiringMap<AuthApiKeySnapshotCacheKey, Option<StoredAuthApiKeySnapshot>>,
load_guard: tokio::sync::Mutex<()>,
// Loads for unrelated identities must not queue behind one global mutex. A
// per-key notification keeps same-key loads singleflight without retaining
// an async mutex (or a cancelled waiter) in the map.
inflight: std::sync::Mutex<HashMap<AuthApiKeySnapshotCacheKey, Arc<AuthApiKeyInflightState>>>,
generation: AtomicU64,
mutation: std::sync::Mutex<()>,
}
impl CachedAuthApiKeyReadRepository {
@@ -20,7 +34,68 @@ impl CachedAuthApiKeyReadRepository {
Self {
inner,
snapshots: ExpiringMap::new(),
load_guard: tokio::sync::Mutex::new(()),
inflight: std::sync::Mutex::new(HashMap::new()),
generation: AtomicU64::new(0),
mutation: std::sync::Mutex::new(()),
}
}
fn insert_if_generation(
&self,
cache_key: AuthApiKeySnapshotCacheKey,
value: Option<StoredAuthApiKeySnapshot>,
generation: u64,
) {
let Ok(_mutation) = self.mutation.lock() else {
return;
};
if self.generation.load(Ordering::Acquire) != generation {
return;
}
self.snapshots.insert(
cache_key,
value,
AUTH_API_KEY_SNAPSHOT_CACHE_TTL,
AUTH_API_KEY_SNAPSHOT_CACHE_MAX_ENTRIES,
);
}
pub(super) fn clear_cache(&self) {
let Ok(_mutation) = self.mutation.lock() else {
return;
};
self.generation.fetch_add(1, Ordering::AcqRel);
self.snapshots.clear();
let states = self
.inflight
.lock()
.map(|mut inflight| inflight.drain().map(|(_, state)| state).collect::<Vec<_>>())
.unwrap_or_default();
for state in states {
state.complete();
}
}
fn register_inflight(
&self,
cache_key: &AuthApiKeySnapshotCacheKey,
) -> AuthApiKeyInflightRegistration<'_> {
match self.inflight.lock() {
Ok(mut inflight) => {
if let Some(state) = inflight.get(cache_key) {
AuthApiKeyInflightRegistration::Follower(state.waiter())
} else {
let state = Arc::new(AuthApiKeyInflightState::new());
inflight.insert(cache_key.clone(), Arc::clone(&state));
AuthApiKeyInflightRegistration::Leader(AuthApiKeyInflightGuard {
cache: self,
cache_key: Some(cache_key.clone()),
state,
generation: self.generation.load(Ordering::Acquire),
})
}
}
Err(_) => AuthApiKeyInflightRegistration::Bypass,
}
}
@@ -43,6 +118,54 @@ impl CachedAuthApiKeyReadRepository {
}
}
impl super::GatewayDataState {
pub(crate) fn clear_auth_api_key_read_cache(&self) {
if let Some(repository) = self.auth_api_key_reader.as_ref() {
repository.clear_cache();
}
}
#[cfg(test)]
pub(crate) fn with_cached_auth_api_key_repository_for_tests<T>(repository: Arc<T>) -> Self
where
T: AuthRepository + 'static,
{
let inner: Arc<dyn AuthApiKeyReadRepository> = repository.clone();
let cached: Arc<dyn AuthApiKeyReadRepository> =
Arc::new(CachedAuthApiKeyReadRepository::new(inner));
let mut state = Self::with_auth_api_key_repository_for_tests(repository);
state.auth_api_key_reader = Some(cached);
state
}
pub(crate) async fn read_auth_api_key_snapshot_strong(
&self,
user_id: &str,
api_key_id: &str,
now_unix_secs: u64,
) -> Result<Option<crate::data::auth::GatewayAuthApiKeySnapshot>, DataLayerError> {
AUTH_API_KEY_READ_CACHE_BYPASS
.scope(
(),
self.read_auth_api_key_snapshot(user_id, api_key_id, now_unix_secs),
)
.await
}
pub(crate) async fn read_auth_api_key_snapshot_by_key_hash_strong(
&self,
key_hash: &str,
now_unix_secs: u64,
) -> Result<Option<crate::data::auth::GatewayAuthApiKeySnapshot>, DataLayerError> {
AUTH_API_KEY_READ_CACHE_BYPASS
.scope(
(),
self.read_auth_api_key_snapshot_by_key_hash(key_hash, now_unix_secs),
)
.await
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
enum AuthApiKeySnapshotCacheKey {
KeyHash(String),
@@ -50,12 +173,142 @@ enum AuthApiKeySnapshotCacheKey {
UserApiKeyIds { user_id: String, api_key_id: String },
}
enum AuthApiKeyInflightRegistration<'a> {
Leader(AuthApiKeyInflightGuard<'a>),
Follower(AuthApiKeyInflightWaiter),
Bypass,
}
struct AuthApiKeyInflightState {
completed: AtomicBool,
error: std::sync::Mutex<Option<DataLayerError>>,
notify: Arc<Notify>,
}
impl AuthApiKeyInflightState {
fn new() -> Self {
Self {
completed: AtomicBool::new(false),
error: std::sync::Mutex::new(None),
notify: Arc::new(Notify::new()),
}
}
fn waiter(self: &Arc<Self>) -> AuthApiKeyInflightWaiter {
AuthApiKeyInflightWaiter {
state: Arc::clone(self),
notified: Arc::clone(&self.notify).notified_owned(),
}
}
fn complete(&self) {
if !self.completed.swap(true, Ordering::AcqRel) {
self.notify.notify_waiters();
}
}
fn fail(&self, error: DataLayerError) {
if let Ok(mut current) = self.error.lock() {
*current = Some(error);
}
self.complete();
}
fn error(&self) -> Option<DataLayerError> {
self.error.lock().ok().and_then(|error| error.clone())
}
}
struct AuthApiKeyInflightWaiter {
state: Arc<AuthApiKeyInflightState>,
notified: OwnedNotified,
}
impl AuthApiKeyInflightWaiter {
async fn wait(self) -> Result<(), DataLayerError> {
let Self { state, notified } = self;
if state.completed.load(Ordering::Acquire) {
return state.error().map_or(Ok(()), Err);
}
tokio::pin!(notified);
if notified.as_mut().enable() || state.completed.load(Ordering::Acquire) {
return state.error().map_or(Ok(()), Err);
}
notified.await;
state.error().map_or(Ok(()), Err)
}
}
struct AuthApiKeyInflightGuard<'a> {
cache: &'a CachedAuthApiKeyReadRepository,
cache_key: Option<AuthApiKeySnapshotCacheKey>,
state: Arc<AuthApiKeyInflightState>,
generation: u64,
}
impl AuthApiKeyInflightGuard<'_> {
fn fail(&self, error: DataLayerError) {
let Some(cache_key) = self.cache_key.as_ref() else {
return;
};
let removed_current = self
.cache
.inflight
.lock()
.map(|mut inflight| {
if inflight
.get(cache_key)
.is_some_and(|current| Arc::ptr_eq(current, &self.state))
{
inflight.remove(cache_key);
true
} else {
false
}
})
.unwrap_or(false);
if removed_current {
self.state.fail(error);
}
}
}
impl Drop for AuthApiKeyInflightGuard<'_> {
fn drop(&mut self) {
let Some(cache_key) = self.cache_key.take() else {
return;
};
let removed = self
.cache
.inflight
.lock()
.map(|mut inflight| {
inflight
.get(&cache_key)
.is_some_and(|current| Arc::ptr_eq(current, &self.state))
&& inflight.remove(&cache_key).is_some()
})
.unwrap_or(false);
if removed {
self.state.complete();
}
}
}
#[async_trait]
impl AuthApiKeyReadRepository for CachedAuthApiKeyReadRepository {
async fn find_api_key_snapshot(
&self,
key: AuthApiKeyLookupKey<'_>,
) -> Result<Option<StoredAuthApiKeySnapshot>, DataLayerError> {
if AUTH_API_KEY_READ_CACHE_BYPASS
.try_with(|_| true)
.unwrap_or(false)
{
return self.inner.find_api_key_snapshot(key).await;
}
let cache_key = Self::cache_key(key);
if let Some(value) = self
.snapshots
@@ -64,22 +317,43 @@ impl AuthApiKeyReadRepository for CachedAuthApiKeyReadRepository {
return Ok(value);
}
let _guard = self.load_guard.lock().await;
if let Some(value) = self
.snapshots
.get_fresh(&cache_key, AUTH_API_KEY_SNAPSHOT_CACHE_TTL)
{
return Ok(value);
}
loop {
match self.register_inflight(&cache_key) {
AuthApiKeyInflightRegistration::Leader(guard) => {
if let Some(value) = self
.snapshots
.get_fresh(&cache_key, AUTH_API_KEY_SNAPSHOT_CACHE_TTL)
{
return Ok(value);
}
let value = self.inner.find_api_key_snapshot(key).await?;
self.snapshots.insert(
cache_key,
value.clone(),
AUTH_API_KEY_SNAPSHOT_CACHE_TTL,
AUTH_API_KEY_SNAPSHOT_CACHE_MAX_ENTRIES,
);
Ok(value)
let value = match self.inner.find_api_key_snapshot(key).await {
Ok(value) => value,
Err(error) => {
guard.fail(error.clone());
return Err(error);
}
};
self.insert_if_generation(cache_key.clone(), value.clone(), guard.generation);
return Ok(value);
}
AuthApiKeyInflightRegistration::Follower(waiter) => {
waiter.wait().await?;
if let Some(value) = self
.snapshots
.get_fresh(&cache_key, AUTH_API_KEY_SNAPSHOT_CACHE_TTL)
{
return Ok(value);
}
}
AuthApiKeyInflightRegistration::Bypass => {
let generation = self.generation.load(Ordering::Acquire);
let value = self.inner.find_api_key_snapshot(key).await?;
self.insert_if_generation(cache_key.clone(), value.clone(), generation);
return Ok(value);
}
}
}
}
async fn list_api_key_snapshots_by_ids(
@@ -98,6 +372,10 @@ impl AuthApiKeyReadRepository for CachedAuthApiKeyReadRepository {
Ok(snapshots)
}
fn clear_cache(&self) {
CachedAuthApiKeyReadRepository::clear_cache(self);
}
async fn list_export_api_keys_by_user_ids(
&self,
user_ids: &[String],
@@ -178,3 +456,262 @@ impl AuthApiKeyReadRepository for CachedAuthApiKeyReadRepository {
self.inner.list_export_standalone_api_keys().await
}
}
#[cfg(test)]
mod tests {
use super::*;
use aether_data::repository::auth::InMemoryAuthApiKeySnapshotRepository;
fn sample_snapshot(api_key_id: &str, user_id: &str) -> StoredAuthApiKeySnapshot {
StoredAuthApiKeySnapshot::new(
user_id.to_string(),
"alice".to_string(),
Some("[email protected]".to_string()),
"user".to_string(),
"local".to_string(),
true,
false,
Some(serde_json::json!(["openai"])),
Some(serde_json::json!(["openai:chat"])),
Some(serde_json::json!(["gpt-4.1"])),
api_key_id.to_string(),
Some("default".to_string()),
true,
false,
false,
Some(60),
Some(5),
Some(200),
Some(serde_json::json!(["openai"])),
Some(serde_json::json!(["openai:chat"])),
Some(serde_json::json!(["gpt-4.1"])),
)
.expect("snapshot should build")
}
#[tokio::test]
async fn concurrent_same_key_loads_once_and_reuses_cached_snapshot() {
let inner = Arc::new(
InMemoryAuthApiKeySnapshotRepository::seed([(
None,
sample_snapshot("key-a", "user-a"),
)])
.with_lookup_delay_for_tests(Duration::from_millis(25)),
);
let repository = Arc::new(CachedAuthApiKeyReadRepository::new(inner.clone()));
let mut tasks = Vec::new();
for _ in 0..32 {
let repository = Arc::clone(&repository);
tasks.push(tokio::spawn(async move {
repository
.find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId("key-a"))
.await
.expect("cached lookup should succeed")
.expect("snapshot should exist")
}));
}
for task in tasks {
assert_eq!(
task.await.expect("lookup task should join").api_key_id,
"key-a"
);
}
assert_eq!(inner.snapshot_lookup_count("key-a"), 1);
assert!(repository
.inflight
.lock()
.expect("inflight lock should not be poisoned")
.is_empty());
}
#[tokio::test]
async fn different_keys_use_independent_inflight_notifications() {
let inner = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed([]));
let repository = CachedAuthApiKeyReadRepository::new(inner);
let key_a = AuthApiKeySnapshotCacheKey::ApiKeyId("key-a".to_string());
let key_b = AuthApiKeySnapshotCacheKey::ApiKeyId("key-b".to_string());
let leader_a = repository.register_inflight(&key_a);
let follower_a = repository.register_inflight(&key_a);
let leader_b = repository.register_inflight(&key_b);
assert!(matches!(
leader_a,
AuthApiKeyInflightRegistration::Leader(_)
));
assert!(matches!(
follower_a,
AuthApiKeyInflightRegistration::Follower(_)
));
assert!(matches!(
leader_b,
AuthApiKeyInflightRegistration::Leader(_)
));
}
#[tokio::test]
async fn cancelled_leader_releases_inflight_key() {
let inner = Arc::new(
InMemoryAuthApiKeySnapshotRepository::seed([])
.with_lookup_delay_for_tests(Duration::from_secs(30)),
);
let repository = Arc::new(CachedAuthApiKeyReadRepository::new(inner));
let lookup_repository = Arc::clone(&repository);
let task = tokio::spawn(async move {
lookup_repository
.find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId("cancelled-key"))
.await
});
tokio::time::timeout(Duration::from_secs(1), async {
loop {
if repository
.inflight
.lock()
.expect("inflight lock should not be poisoned")
.contains_key(&AuthApiKeySnapshotCacheKey::ApiKeyId(
"cancelled-key".to_string(),
))
{
break;
}
tokio::task::yield_now().await;
}
})
.await
.expect("lookup should register its inflight key");
task.abort();
let _ = task.await;
assert!(repository
.inflight
.lock()
.expect("inflight lock should not be poisoned")
.is_empty());
}
#[tokio::test]
async fn follower_observes_completion_before_first_poll() {
let inner = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed([]));
let repository = CachedAuthApiKeyReadRepository::new(inner);
let key = AuthApiKeySnapshotCacheKey::ApiKeyId("key-a".to_string());
let leader = match repository.register_inflight(&key) {
AuthApiKeyInflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
let follower = match repository.register_inflight(&key) {
AuthApiKeyInflightRegistration::Follower(waiter) => waiter,
_ => panic!("second registration should follow"),
};
// Complete before the OwnedNotified is polled. A bare notify_waiters()
// broadcast would be lost in this ordering.
drop(leader);
tokio::time::timeout(Duration::from_millis(100), follower.wait())
.await
.expect("completed follower must not miss the broadcast")
.expect("successful flight should not publish an error");
}
#[tokio::test]
async fn failed_flight_shares_error_and_preserves_replacement() {
let inner = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed([]));
let repository = CachedAuthApiKeyReadRepository::new(inner);
let key = AuthApiKeySnapshotCacheKey::ApiKeyId("key-a".to_string());
let old_leader = match repository.register_inflight(&key) {
AuthApiKeyInflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
let follower = match repository.register_inflight(&key) {
AuthApiKeyInflightRegistration::Follower(waiter) => waiter,
_ => panic!("second registration should follow"),
};
old_leader.fail(DataLayerError::Sql(
"forced auth snapshot load failure".to_string(),
));
let error = tokio::time::timeout(Duration::from_millis(100), follower.wait())
.await
.expect("failed flight should release its follower")
.expect_err("follower should observe the leader error");
assert_eq!(
error.to_string(),
"sql error: forced auth snapshot load failure"
);
let replacement = match repository.register_inflight(&key) {
AuthApiKeyInflightRegistration::Leader(guard) => guard,
_ => panic!("failed flight should allow a replacement"),
};
drop(old_leader);
assert!(matches!(
repository.register_inflight(&key),
AuthApiKeyInflightRegistration::Follower(_)
));
drop(replacement);
}
#[tokio::test]
async fn strong_read_bypasses_a_fresh_cached_allow_snapshot() {
let inner = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed([(
None,
sample_snapshot("key-a", "user-a"),
)]));
let repository = CachedAuthApiKeyReadRepository::new(inner.clone());
let lookup = AuthApiKeyLookupKey::UserApiKeyIds {
user_id: "user-a",
api_key_id: "key-a",
};
let cached = repository
.find_api_key_snapshot(lookup)
.await
.expect("initial lookup should succeed")
.expect("snapshot should exist");
assert!(!cached.api_key_is_locked);
assert!(inner
.set_user_api_key_locked("user-a", "key-a", true)
.await
.expect("cross-node lock should succeed"));
let still_cached = repository
.find_api_key_snapshot(lookup)
.await
.expect("cached lookup should succeed")
.expect("snapshot should exist");
assert!(!still_cached.api_key_is_locked);
let strong = AUTH_API_KEY_READ_CACHE_BYPASS
.scope((), repository.find_api_key_snapshot(lookup))
.await
.expect("strong lookup should succeed")
.expect("snapshot should exist");
assert!(strong.api_key_is_locked);
assert_eq!(inner.snapshot_lookup_count("key-a"), 2);
}
#[test]
fn clear_cache_rejects_old_leader_publication() {
let inner = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed([]));
let repository = CachedAuthApiKeyReadRepository::new(inner);
let key = AuthApiKeySnapshotCacheKey::ApiKeyId("key-a".to_string());
let leader = match repository.register_inflight(&key) {
AuthApiKeyInflightRegistration::Leader(guard) => guard,
_ => panic!("key-a should register a leader"),
};
let old_generation = leader.generation;
repository.clear_cache();
repository.insert_if_generation(
key.clone(),
Some(sample_snapshot("stale-key", "user-a")),
old_generation,
);
assert!(repository
.snapshots
.get_fresh(&key, AUTH_API_KEY_SNAPSHOT_CACHE_TTL)
.is_none());
}
}
@@ -1,6 +1,6 @@
use std::collections::HashMap;
use std::future::Future;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::time::Duration;
@@ -12,12 +12,13 @@ use aether_data_contracts::repository::candidate_selection::{
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
};
use async_trait::async_trait;
use tokio::sync::Notify;
use tokio::sync::{Notify, OwnedSemaphorePermit, Semaphore};
use tokio::time::timeout;
use tracing::warn;
const CANDIDATE_SELECTION_CACHE_TTL: Duration = Duration::from_secs(5);
const CANDIDATE_SELECTION_CACHE_MAX_ENTRIES: usize = 4096;
const CANDIDATE_SELECTION_CACHE_MAX_INFLIGHT: usize = 4096;
#[cfg(not(test))]
const CANDIDATE_SELECTION_CACHE_LOAD_TIMEOUT: Duration = Duration::from_secs(10);
#[cfg(test)]
@@ -30,10 +31,10 @@ const CANDIDATE_SELECTION_CACHE_INFLIGHT_WAIT_TIMEOUT: Duration = Duration::from
pub(super) struct CachedMinimalCandidateSelectionReadRepository {
inner: Arc<dyn MinimalCandidateSelectionReadRepository>,
entries: ExpiringMap<CandidateSelectionCacheKey, Vec<StoredMinimalCandidateSelectionRow>>,
inflight: Mutex<HashMap<CandidateSelectionCacheKey, u64>>,
inflight_notify: Notify,
next_inflight_token: AtomicU64,
inflight: Mutex<HashMap<CandidateSelectionCacheKey, Arc<InflightState>>>,
epoch: AtomicU64,
mutation: Mutex<()>,
admission: Arc<Semaphore>,
}
impl CachedMinimalCandidateSelectionReadRepository {
@@ -42,9 +43,9 @@ impl CachedMinimalCandidateSelectionReadRepository {
inner,
entries: ExpiringMap::new(),
inflight: Mutex::new(HashMap::new()),
inflight_notify: Notify::new(),
next_inflight_token: AtomicU64::new(1),
epoch: AtomicU64::new(0),
mutation: Mutex::new(()),
admission: Arc::new(Semaphore::new(CANDIDATE_SELECTION_CACHE_MAX_INFLIGHT)),
}
}
@@ -61,18 +62,35 @@ impl CachedMinimalCandidateSelectionReadRepository {
return Ok(rows);
}
self.load_after_cache_miss(key, load).await
}
async fn load_after_cache_miss<F, Fut>(
&self,
key: CandidateSelectionCacheKey,
load: F,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>
where
F: Fn() -> Fut,
Fut: Future<Output = Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>>,
{
loop {
let notified = self.inflight_notify.notified();
match self.register_inflight(&key) {
InflightRegistration::Bypass => {
return load_candidate_selection_rows_with_timeout(&key, load()).await;
InflightRegistration::Saturated => {
return Err(DataLayerError::TimedOut(format!(
"candidate selection cache admission saturated for {key:?}"
)));
}
InflightRegistration::Follower => {
if timeout(CANDIDATE_SELECTION_CACHE_INFLIGHT_WAIT_TIMEOUT, notified)
.await
.is_err()
InflightRegistration::Follower(state) => {
match timeout(
CANDIDATE_SELECTION_CACHE_INFLIGHT_WAIT_TIMEOUT,
state.wait(),
)
.await
{
self.expire_inflight(&key);
Ok(Ok(())) => {}
Ok(Err(error)) => return Err(error),
Err(_) => self.expire_inflight(&key, &state),
}
if let Some(rows) = self.entries.get_fresh(&key, CANDIDATE_SELECTION_CACHE_TTL)
{
@@ -80,60 +98,107 @@ impl CachedMinimalCandidateSelectionReadRepository {
}
continue;
}
InflightRegistration::Leader(token) => {
let mut guard = InflightGuard::new(self, key.clone(), token);
let load_epoch = self.epoch.load(Ordering::Acquire);
InflightRegistration::Leader(mut guard) => {
// A writer may have populated the cache after the first
// miss but before this flight was registered.
if let Some(rows) = self.entries.get_fresh(&key, CANDIDATE_SELECTION_CACHE_TTL)
{
return Ok(rows);
}
let result = load_candidate_selection_rows_with_timeout(&key, load()).await;
if let Ok(rows) = &result {
if load_epoch == self.epoch.load(Ordering::Acquire) {
self.entries.insert(
key.clone(),
rows.clone(),
CANDIDATE_SELECTION_CACHE_TTL,
CANDIDATE_SELECTION_CACHE_MAX_ENTRIES,
);
}
match &result {
Ok(rows) => guard.finish_loaded(rows.clone()),
Err(error) => guard.finish(Some(error.clone())),
}
guard.finish();
return result;
}
}
}
}
fn register_inflight(&self, key: &CandidateSelectionCacheKey) -> 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 register_inflight(&self, key: &CandidateSelectionCacheKey) -> InflightRegistration<'_> {
let mut inflight = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(state) = inflight.get(key) {
return InflightRegistration::Follower(Arc::clone(state));
}
if inflight.len() >= CANDIDATE_SELECTION_CACHE_MAX_INFLIGHT {
return InflightRegistration::Saturated;
}
let Ok(admission) = Arc::clone(&self.admission).try_acquire_owned() else {
return InflightRegistration::Saturated;
};
let state = Arc::new(InflightState {
epoch: self.epoch.load(Ordering::Acquire),
notify: Notify::new(),
completed: AtomicBool::new(false),
error: Mutex::new(None),
});
inflight.insert(key.clone(), Arc::clone(&state));
InflightRegistration::Leader(InflightGuard::new(self, key.clone(), state, admission))
}
fn finish_inflight(&self, key: &CandidateSelectionCacheKey, 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;
}
fn finish_inflight(
&self,
key: &CandidateSelectionCacheKey,
state: &Arc<InflightState>,
rows: Option<Vec<StoredMinimalCandidateSelectionRow>>,
admission: Option<OwnedSemaphorePermit>,
) -> Option<Arc<InflightState>> {
let _mutation = self
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let mut inflight = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
drop(admission);
if !inflight
.get(key)
.is_some_and(|current| Arc::ptr_eq(current, state))
|| self.epoch.load(Ordering::Acquire) != state.epoch
{
return None;
}
if removed {
self.inflight_notify.notify_waiters();
if let Some(rows) = rows {
self.entries.insert(
key.clone(),
rows,
CANDIDATE_SELECTION_CACHE_TTL,
CANDIDATE_SELECTION_CACHE_MAX_ENTRIES,
);
}
let removed = inflight.remove(key);
drop(inflight);
drop(_mutation);
removed
}
fn expire_inflight(&self, key: &CandidateSelectionCacheKey) {
let mut removed = false;
if let Ok(mut inflight) = self.inflight.lock() {
removed = inflight.remove(key).is_some();
}
if removed {
fn expire_inflight(&self, key: &CandidateSelectionCacheKey, state: &Arc<InflightState>) {
let _mutation = self
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let removed = {
let mut inflight = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if inflight
.get(key)
.is_some_and(|current| Arc::ptr_eq(current, state))
{
inflight.remove(key)
} else {
None
}
};
drop(_mutation);
if let Some(state) = removed {
warn!(
event_name = "candidate_selection_cache_inflight_expired",
log_type = "ops",
@@ -141,64 +206,143 @@ impl CachedMinimalCandidateSelectionReadRepository {
wait_timeout_ms = CANDIDATE_SELECTION_CACHE_INFLIGHT_WAIT_TIMEOUT.as_millis() as u64,
"gateway candidate selection cache expired stale inflight load"
);
self.inflight_notify.notify_waiters();
state.complete(None);
}
}
fn clear(&self) {
let _mutation = self
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
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 {
let states = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.drain()
.map(|(_, state)| state)
.collect::<Vec<_>>();
if !states.is_empty() {
warn!(
event_name = "candidate_selection_cache_inflight_cleared",
log_type = "ops",
"gateway candidate selection cache cleared in-flight loads"
);
self.inflight_notify.notify_waiters();
for state in states {
state.complete(None);
}
}
}
}
enum InflightRegistration {
Leader(u64),
Follower,
Bypass,
struct InflightState {
epoch: u64,
notify: Notify,
completed: AtomicBool,
error: Mutex<Option<DataLayerError>>,
}
impl InflightState {
fn complete(&self, error: Option<DataLayerError>) {
if let Some(error) = error {
if let Ok(mut current) = self.error.lock() {
*current = Some(error);
}
}
if !self.completed.swap(true, Ordering::AcqRel) {
self.notify.notify_waiters();
}
}
async fn wait(&self) -> Result<(), DataLayerError> {
loop {
if self.completed.load(Ordering::Acquire) {
return self
.error
.lock()
.ok()
.and_then(|error| error.clone())
.map_or(Ok(()), Err);
}
// Register before checking completion a second time. Creating a
// Notified future alone is insufficient because notify_waiters()
// can otherwise run before the future's first poll.
let mut notified = Box::pin(self.notify.notified());
notified.as_mut().enable();
if self.completed.load(Ordering::Acquire) {
return self
.error
.lock()
.ok()
.and_then(|error| error.clone())
.map_or(Ok(()), Err);
}
notified.await;
}
}
}
enum InflightRegistration<'a> {
Leader(InflightGuard<'a>),
Follower(Arc<InflightState>),
Saturated,
}
struct InflightGuard<'a> {
cache: &'a CachedMinimalCandidateSelectionReadRepository,
key: Option<CandidateSelectionCacheKey>,
token: u64,
state: Arc<InflightState>,
admission: Option<OwnedSemaphorePermit>,
}
impl<'a> InflightGuard<'a> {
fn new(
cache: &'a CachedMinimalCandidateSelectionReadRepository,
key: CandidateSelectionCacheKey,
token: u64,
state: Arc<InflightState>,
admission: OwnedSemaphorePermit,
) -> Self {
Self {
cache,
key: Some(key),
token,
state,
admission: Some(admission),
}
}
fn finish(&mut self) {
if let Some(key) = self.key.take() {
self.cache.finish_inflight(&key, self.token);
fn epoch(&self) -> u64 {
self.state.epoch
}
fn finish_loaded(&mut self, rows: Vec<StoredMinimalCandidateSelectionRow>) {
let removed = self.key.take().and_then(|key| {
self.cache
.finish_inflight(&key, &self.state, Some(rows), self.admission.take())
});
self.admission.take();
if let Some(removed) = removed {
removed.complete(None);
}
}
fn finish(&mut self, error: Option<DataLayerError>) {
let removed = self.key.take().and_then(|key| {
self.cache
.finish_inflight(&key, &self.state, None, self.admission.take())
});
self.admission.take();
if let Some(removed) = removed {
removed.complete(error);
}
}
}
impl Drop for InflightGuard<'_> {
fn drop(&mut self) {
self.finish();
self.finish(None);
}
}
@@ -573,6 +717,202 @@ mod tests {
assert_eq!(inner.calls(), 1);
}
#[tokio::test]
async fn candidate_selection_cache_only_notifies_matching_inflight_key() {
let inner = Arc::new(StubCandidateSelectionRepository::new(Duration::ZERO));
let cache = CachedMinimalCandidateSelectionReadRepository::new(inner);
let key_a = CandidateSelectionCacheKey::ApiFormat {
api_format: "openai:chat".to_string(),
};
let key_b = CandidateSelectionCacheKey::ApiFormat {
api_format: "anthropic:messages".to_string(),
};
let mut leader_a = match cache.register_inflight(&key_a) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("first key registration should lead"),
};
let mut leader_b = match cache.register_inflight(&key_b) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("second key registration should lead independently"),
};
let state_a = match cache.register_inflight(&key_a) {
InflightRegistration::Follower(state) => state,
_ => panic!("duplicate key registration should follow"),
};
leader_b.finish(None);
assert!(
tokio::time::timeout(Duration::from_millis(10), state_a.wait())
.await
.is_err(),
"completing another key must not wake this follower"
);
leader_a.finish(None);
tokio::time::timeout(Duration::from_millis(100), state_a.wait())
.await
.expect("completing the matching key must wake its follower")
.expect("successful completion should not publish an error");
assert!(cache.inflight.lock().unwrap().is_empty());
}
#[tokio::test]
async fn candidate_selection_cache_follower_observes_completion_before_first_poll() {
let inner = Arc::new(StubCandidateSelectionRepository::new(Duration::ZERO));
let cache = CachedMinimalCandidateSelectionReadRepository::new(inner);
let key = CandidateSelectionCacheKey::ApiFormat {
api_format: "openai:chat".to_string(),
};
let mut leader = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
let state = match cache.register_inflight(&key) {
InflightRegistration::Follower(state) => state,
_ => panic!("second registration should follow"),
};
// Complete before wait() is constructed or polled. notify_waiters()
// alone would lose this notification and wait for the full timeout.
leader.finish(None);
tokio::time::timeout(Duration::from_millis(100), state.wait())
.await
.expect("completed follower must not miss the broadcast")
.expect("successful completion should not publish an error");
}
#[tokio::test]
async fn candidate_selection_cache_shares_leader_failure() {
let inner = Arc::new(StubCandidateSelectionRepository::new(Duration::ZERO));
let cache = CachedMinimalCandidateSelectionReadRepository::new(inner);
let key = CandidateSelectionCacheKey::ApiFormat {
api_format: "openai:chat".to_string(),
};
let mut leader = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
let state = match cache.register_inflight(&key) {
InflightRegistration::Follower(state) => state,
_ => panic!("second registration should follow"),
};
leader.finish(Some(DataLayerError::Sql(
"forced candidate cache load failure".to_string(),
)));
let error = tokio::time::timeout(Duration::from_millis(100), state.wait())
.await
.expect("failed load should release its follower")
.expect_err("follower should observe the leader failure");
assert_eq!(
error.to_string(),
"sql error: forced candidate cache load failure"
);
assert!(cache.inflight.lock().unwrap().is_empty());
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn candidate_selection_cache_broadcasts_to_all_same_key_followers() {
const FOLLOWERS: usize = 64;
let inner = Arc::new(StubCandidateSelectionRepository::new(Duration::ZERO));
let cache = CachedMinimalCandidateSelectionReadRepository::new(inner);
let key = CandidateSelectionCacheKey::ApiFormat {
api_format: "openai:chat".to_string(),
};
let mut leader = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
let mut tasks = Vec::with_capacity(FOLLOWERS);
for _ in 0..FOLLOWERS {
let state = match cache.register_inflight(&key) {
InflightRegistration::Follower(state) => state,
_ => panic!("same-key registration should follow"),
};
tasks.push(tokio::spawn(async move { state.wait().await }));
}
tokio::task::yield_now().await;
leader.finish(None);
for task in tasks {
tokio::time::timeout(Duration::from_millis(250), task)
.await
.expect("all same-key followers should receive completion")
.expect("follower task should finish")
.expect("successful completion should not publish an error");
}
}
#[tokio::test]
async fn candidate_selection_cache_cancelled_guard_wakes_registered_follower() {
let inner = Arc::new(StubCandidateSelectionRepository::new(Duration::ZERO));
let cache = CachedMinimalCandidateSelectionReadRepository::new(inner);
let key = CandidateSelectionCacheKey::ApiFormat {
api_format: "openai:chat".to_string(),
};
let guard = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
let state = match cache.register_inflight(&key) {
InflightRegistration::Follower(state) => state,
_ => panic!("second registration should follow"),
};
drop(guard);
tokio::time::timeout(Duration::from_millis(100), state.wait())
.await
.expect("leader cancellation must wake an existing follower")
.expect("leader cancellation should allow a retry");
assert!(cache.inflight.lock().unwrap().is_empty());
}
#[tokio::test]
async fn candidate_selection_cache_clear_wakes_follower_and_old_guard_keeps_new_flight() {
let inner = Arc::new(StubCandidateSelectionRepository::new(Duration::ZERO));
let cache = CachedMinimalCandidateSelectionReadRepository::new(inner);
let key = CandidateSelectionCacheKey::ApiFormat {
api_format: "openai:chat".to_string(),
};
let old_guard = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
let old_state = match cache.register_inflight(&key) {
InflightRegistration::Follower(state) => state,
_ => panic!("second registration should follow"),
};
let old_epoch = old_guard.epoch();
cache.clear();
assert!(cache.epoch.load(Ordering::Acquire) > old_epoch);
let mut new_guard = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("registration after clear should lead"),
};
assert_ne!(old_guard.epoch(), new_guard.epoch());
// Dropping the invalidated leader's RAII guard must not remove the
// new flight created for the same key after clear().
drop(old_guard);
assert!(cache
.inflight
.lock()
.unwrap()
.get(&key)
.is_some_and(|current| Arc::ptr_eq(current, &new_guard.state)));
tokio::time::timeout(Duration::from_millis(100), old_state.wait())
.await
.expect("clear must wake followers of the invalidated flight")
.expect("clear should allow a retry");
new_guard.finish(None);
assert!(cache.inflight.lock().unwrap().is_empty());
}
#[tokio::test]
async fn candidate_selection_cache_clear_invalidates_entries() {
let inner = Arc::new(StubCandidateSelectionRepository::new(Duration::ZERO));
@@ -587,6 +927,123 @@ mod tests {
assert_eq!(inner.calls(), 2);
}
#[test]
fn candidate_selection_cache_rejects_publication_from_pre_clear_flight() {
let inner = Arc::new(StubCandidateSelectionRepository::new(Duration::ZERO));
let cache = CachedMinimalCandidateSelectionReadRepository::new(inner);
let key = CandidateSelectionCacheKey::ApiFormat {
api_format: "openai:chat".to_string(),
};
let mut stale_leader = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
cache.clear();
stale_leader.finish_loaded(Vec::new());
assert!(cache
.entries
.get_fresh(&key, CANDIDATE_SELECTION_CACHE_TTL)
.is_none());
}
#[test]
fn candidate_selection_cache_clear_keeps_active_load_admission_bounded() {
let inner = Arc::new(StubCandidateSelectionRepository::new(Duration::ZERO));
let mut cache = CachedMinimalCandidateSelectionReadRepository::new(inner);
cache.admission = Arc::new(Semaphore::new(2));
let key = CandidateSelectionCacheKey::ApiFormat {
api_format: "openai:chat".to_string(),
};
let mut detached_leaders = Vec::new();
for _ in 0..2 {
let leader = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("load below the hard active limit should lead"),
};
cache.clear();
detached_leaders.push(leader);
}
assert_eq!(cache.admission.available_permits(), 0);
assert!(matches!(
cache.register_inflight(&key),
InflightRegistration::Saturated
));
drop(detached_leaders);
assert_eq!(cache.admission.available_permits(), 2);
assert!(matches!(
cache.register_inflight(&key),
InflightRegistration::Leader(_)
));
}
#[test]
fn candidate_selection_cache_expired_leader_cannot_publish_over_replacement() {
let inner = Arc::new(StubCandidateSelectionRepository::new(Duration::ZERO));
let cache = CachedMinimalCandidateSelectionReadRepository::new(inner);
let key = CandidateSelectionCacheKey::ApiFormat {
api_format: "openai:chat".to_string(),
};
let mut old_leader = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
let old_state = Arc::clone(&old_leader.state);
cache.expire_inflight(&key, &old_state);
let mut replacement = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("expiration should allow a replacement leader"),
};
assert_eq!(old_leader.epoch(), replacement.epoch());
old_leader.finish_loaded(vec![sample_row("stale-key", 1)]);
assert!(cache
.entries
.get_fresh(&key, CANDIDATE_SELECTION_CACHE_TTL)
.is_none());
replacement.finish_loaded(vec![sample_row("fresh-key", 2)]);
let cached = cache
.entries
.get_fresh(&key, CANDIDATE_SELECTION_CACHE_TTL)
.expect("replacement should publish");
assert_eq!(cached[0].key_id, "fresh-key");
}
#[tokio::test]
async fn candidate_selection_cache_rechecks_fresh_entry_after_leader_registration() {
let inner = Arc::new(StubCandidateSelectionRepository::new(Duration::ZERO));
let cache = CachedMinimalCandidateSelectionReadRepository::new(inner);
let key = CandidateSelectionCacheKey::ApiFormat {
api_format: "openai:chat".to_string(),
};
let expected = vec![sample_row("cached-key", 1)];
cache.entries.insert(
key.clone(),
expected.clone(),
CANDIDATE_SELECTION_CACHE_TTL,
CANDIDATE_SELECTION_CACHE_MAX_ENTRIES,
);
let loads = AtomicUsize::new(0);
// Exercise the post-initial-miss path directly to model a concurrent
// writer filling the cache immediately before flight registration.
let rows = cache
.load_after_cache_miss(key, || async {
loads.fetch_add(1, Ordering::SeqCst);
Ok(Vec::new())
})
.await
.expect("fresh entry should satisfy the lookup");
assert_eq!(rows, expected);
assert_eq!(loads.load(Ordering::SeqCst), 0);
assert!(cache.inflight.lock().unwrap().is_empty());
}
#[tokio::test]
async fn candidate_selection_cache_releases_inflight_when_leader_is_cancelled() {
let inner = Arc::new(FirstLoadPendingThenFastRepository::new());
+105 -3
View File
@@ -1,8 +1,10 @@
use super::{
ApiKeyLastUsedDelta, DataLayerError, GatewayDataState, GeminiFileMappingListQuery,
GeminiFileMappingStats, ProviderCatalogKeyListQuery, PublicHealthStatusCount,
PublicHealthTimelineBucket, StoredGeminiFileMapping, StoredGeminiFileMappingListPage,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
GeminiFileMappingStats, ProviderCatalogKeyAdaptiveStateUpdate,
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListQuery,
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
PublicHealthStatusCount, PublicHealthTimelineBucket, StoredGeminiFileMapping,
StoredGeminiFileMappingListPage, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider, StoredRequestCandidate,
UpsertGeminiFileMappingRecord, UpsertRequestCandidateRecord,
@@ -370,6 +372,34 @@ impl GatewayDataState {
Ok(updated)
}
pub(crate) async fn update_provider_catalog_key_oauth_runtime_state(
&self,
key_id: &str,
oauth_invalid_at_unix_secs: Option<u64>,
oauth_invalid_reason: Option<&str>,
encrypted_auth_config_update: Option<&str>,
updated_at_unix_secs: Option<u64>,
) -> Result<bool, DataLayerError> {
let updated = match &self.provider_catalog_writer {
Some(repository) => {
repository
.update_key_oauth_runtime_state(
key_id,
oauth_invalid_at_unix_secs,
oauth_invalid_reason,
encrypted_auth_config_update,
updated_at_unix_secs,
)
.await
}
None => Ok(false),
}?;
if updated {
self.clear_provider_catalog_cache();
}
Ok(updated)
}
pub(crate) async fn create_provider_catalog_key(
&self,
key: &StoredProviderCatalogKey,
@@ -681,4 +711,76 @@ impl GatewayDataState {
}
Ok(updated)
}
pub(crate) async fn reset_provider_catalog_key_error_count(
&self,
key_id: &str,
) -> Result<bool, DataLayerError> {
let updated = match &self.provider_catalog_writer {
Some(repository) => repository.reset_key_error_count(key_id).await,
None => Ok(false),
}?;
if updated {
self.clear_provider_catalog_cache();
}
Ok(updated)
}
pub(crate) async fn compare_and_update_provider_catalog_key_adaptive_state(
&self,
update: &ProviderCatalogKeyAdaptiveStateUpdate,
) -> Result<bool, DataLayerError> {
let Some(repository) = &self.provider_catalog_writer else {
return Ok(false);
};
let updated = repository
.compare_and_update_key_adaptive_state(update)
.await?;
// A false CAS result normally means another instance won the write. Drop the
// five-second read cache before the caller reloads and retries.
self.clear_provider_catalog_cache();
Ok(updated)
}
pub(crate) async fn update_provider_catalog_key_runtime_metadata(
&self,
update: &ProviderCatalogKeyRuntimeMetadataUpdate,
) -> Result<bool, DataLayerError> {
let Some(repository) = &self.provider_catalog_writer else {
return Ok(false);
};
let updated = repository.update_key_runtime_metadata(update).await?;
// A false result is a namespace CAS conflict. Drop the read cache so
// the caller's retry observes the writer that won the race.
self.clear_provider_catalog_cache();
Ok(updated)
}
pub(crate) async fn update_provider_catalog_key_status_snapshot(
&self,
update: &ProviderCatalogKeyStatusSnapshotUpdate,
) -> Result<bool, DataLayerError> {
let Some(repository) = &self.provider_catalog_writer else {
return Ok(false);
};
let updated = repository.update_key_status_snapshot(update).await?;
if updated {
self.clear_provider_catalog_cache();
}
Ok(updated)
}
pub(crate) async fn compare_and_update_provider_catalog_key_health_state(
&self,
update: &ProviderCatalogKeyHealthStateUpdate,
) -> Result<bool, DataLayerError> {
let Some(repository) = &self.provider_catalog_writer else {
return Ok(false);
};
let updated = repository
.compare_and_update_key_health_state(update)
.await?;
self.clear_provider_catalog_cache();
Ok(updated)
}
}
+444 -125
View File
@@ -3,76 +3,248 @@ use aether_data_contracts::repository::candidate_selection::MinimalCandidateSele
use aether_data_contracts::repository::candidates::RequestCandidateReadRepository;
use aether_data_contracts::repository::provider_catalog::ProviderCatalogReadRepository;
use aether_runtime_state::RuntimeQueueStore;
use std::collections::HashSet;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::Notify;
use std::time::Duration;
use super::{GatewayDataConfig, GatewayDataState, StoredSystemConfigEntry};
use super::{
GatewayDataConfig, GatewayDataState, StoredSystemConfigEntry, SystemConfigValueCacheState,
SystemConfigValueInflightCompletion, SystemConfigValueInflightState,
};
const SYSTEM_CONFIG_VALUE_CACHE_TTL: Duration = Duration::from_secs(30);
fn system_config_value_load_state() -> &'static SystemConfigValueLoadState {
static STATE: std::sync::OnceLock<SystemConfigValueLoadState> = std::sync::OnceLock::new();
STATE.get_or_init(SystemConfigValueLoadState::default)
}
#[derive(Debug, Default)]
struct SystemConfigValueLoadState {
inflight: std::sync::Mutex<HashSet<String>>,
notify: Notify,
}
const SYSTEM_CONFIG_VALUE_CACHE_MAX_ENTRIES: usize = 512;
const SYSTEM_CONFIG_VALUE_CACHE_MAX_INFLIGHT: usize = 512;
enum SystemConfigValueLoadRegistration<'a> {
Leader(SystemConfigValueLoadGuard<'a>),
Follower,
Bypass,
Follower(Arc<SystemConfigValueInflightState>),
Saturated,
}
struct SystemConfigValueLoadGuard<'a> {
state: &'a SystemConfigValueLoadState,
cache: &'a SystemConfigValueCacheState,
key: Option<String>,
state: Arc<SystemConfigValueInflightState>,
admission: Option<tokio::sync::OwnedSemaphorePermit>,
}
impl SystemConfigValueInflightState {
async fn wait(&self) -> SystemConfigValueInflightCompletion {
loop {
if let Some(completion) = self.completion.get() {
return completion.clone();
}
let mut notified = Box::pin(self.notify.notified());
notified.as_mut().enable();
if let Some(completion) = self.completion.get() {
return completion.clone();
}
notified.await;
}
}
}
impl SystemConfigValueCacheState {
fn get(&self, key: &str) -> Option<Option<serde_json::Value>> {
self.entries.get_fresh(key, SYSTEM_CONFIG_VALUE_CACHE_TTL)
}
fn register(&self, key: &str) -> SystemConfigValueLoadRegistration<'_> {
{
let inflight = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(state) = inflight.get(key) {
return SystemConfigValueLoadRegistration::Follower(Arc::clone(state));
}
}
let _mutation = self
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let mut inflight = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(state) = inflight.get(key) {
return SystemConfigValueLoadRegistration::Follower(Arc::clone(state));
}
if inflight.len() >= SYSTEM_CONFIG_VALUE_CACHE_MAX_INFLIGHT {
return SystemConfigValueLoadRegistration::Saturated;
}
let Ok(admission) = Arc::clone(&self.admission).try_acquire_owned() else {
return SystemConfigValueLoadRegistration::Saturated;
};
let state = Arc::new(SystemConfigValueInflightState {
notify: Arc::new(tokio::sync::Notify::new()),
completion: std::sync::OnceLock::new(),
});
inflight.insert(key.to_string(), Arc::clone(&state));
SystemConfigValueLoadRegistration::Leader(SystemConfigValueLoadGuard {
cache: self,
key: Some(key.to_string()),
state,
admission: Some(admission),
})
}
fn finish_loaded(
&self,
key: &str,
state: &Arc<SystemConfigValueInflightState>,
admission: tokio::sync::OwnedSemaphorePermit,
value: Option<serde_json::Value>,
) {
self.finish_current(
key,
state,
admission,
SystemConfigValueInflightCompletion::Loaded,
Some(value),
);
}
fn finish_current(
&self,
key: &str,
state: &Arc<SystemConfigValueInflightState>,
admission: tokio::sync::OwnedSemaphorePermit,
completion: SystemConfigValueInflightCompletion,
cache_value: Option<Option<serde_json::Value>>,
) {
let _mutation = self
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let mut inflight = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
drop(admission);
debug_assert!(self.admission.available_permits() > 0);
if !inflight
.get(key)
.is_some_and(|current| Arc::ptr_eq(current, state))
{
return;
}
if let Some(value) = cache_value {
self.entries.insert(
key.to_string(),
value,
SYSTEM_CONFIG_VALUE_CACHE_TTL,
SYSTEM_CONFIG_VALUE_CACHE_MAX_ENTRIES,
);
}
let completed = state.completion.set(completion).is_ok();
inflight.remove(key);
drop(inflight);
drop(_mutation);
if completed {
state.notify.notify_waiters();
}
}
fn invalidate(&self, key: &str) {
let _mutation = self
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
self.entries.remove(key);
let state = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.remove(key);
let completed = state.as_ref().is_some_and(|state| {
state
.completion
.set(SystemConfigValueInflightCompletion::Invalidated)
.is_ok()
});
drop(_mutation);
if completed {
state
.expect("completed state should exist")
.notify
.notify_waiters();
}
}
fn clear(&self) {
let _mutation = self
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
self.entries.clear();
let states = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.drain()
.map(|(_, state)| state)
.collect::<Vec<_>>();
let completed = states
.iter()
.filter(|state| {
state
.completion
.set(SystemConfigValueInflightCompletion::Invalidated)
.is_ok()
})
.cloned()
.collect::<Vec<_>>();
drop(_mutation);
for state in completed {
state.notify.notify_waiters();
}
}
}
impl Default for SystemConfigValueCacheState {
fn default() -> Self {
Self {
entries: aether_cache::ExpiringMap::default(),
inflight: std::sync::Mutex::new(std::collections::HashMap::new()),
mutation: std::sync::Mutex::new(()),
admission: Arc::new(tokio::sync::Semaphore::new(
SYSTEM_CONFIG_VALUE_CACHE_MAX_INFLIGHT,
)),
}
}
}
impl SystemConfigValueLoadGuard<'_> {
fn finish_loaded(&mut self, value: Option<serde_json::Value>) {
if let Some(key) = self.key.take() {
let admission = self
.admission
.take()
.expect("active system config leader must own admission");
self.cache
.finish_loaded(&key, &self.state, admission, value);
}
}
fn finish(&mut self, completion: SystemConfigValueInflightCompletion) {
if let Some(key) = self.key.take() {
let admission = self
.admission
.take()
.expect("active system config leader must own admission");
self.cache
.finish_current(&key, &self.state, admission, completion, None);
}
}
}
impl Drop for SystemConfigValueLoadGuard<'_> {
fn drop(&mut self) {
if let Some(key) = self.key.take() {
self.state.finish(&key);
}
}
}
impl SystemConfigValueLoadState {
fn register(&self, key: &str) -> SystemConfigValueLoadRegistration<'_> {
match self.inflight.lock() {
Ok(mut inflight) => {
if inflight.contains(key) {
SystemConfigValueLoadRegistration::Follower
} else {
inflight.insert(key.to_string());
SystemConfigValueLoadRegistration::Leader(SystemConfigValueLoadGuard {
state: self,
key: Some(key.to_string()),
})
}
}
Err(_) => SystemConfigValueLoadRegistration::Bypass,
}
}
fn notified(&self) -> tokio::sync::futures::Notified<'_> {
self.notify.notified()
}
fn finish(&self, key: &str) {
let removed = self
.inflight
.lock()
.map(|mut inflight| inflight.remove(key))
.unwrap_or(false);
if removed {
self.notify.notify_waiters();
}
self.finish(SystemConfigValueInflightCompletion::Cancelled);
}
}
@@ -83,6 +255,139 @@ fn current_system_config_updated_at_unix_secs() -> u64 {
.as_secs()
}
#[cfg(test)]
mod system_config_value_cache_tests {
use super::*;
fn leader<'a>(
cache: &'a SystemConfigValueCacheState,
key: &str,
) -> SystemConfigValueLoadGuard<'a> {
match cache.register(key) {
SystemConfigValueLoadRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
}
}
#[tokio::test]
async fn completion_before_first_poll_releases_system_config_follower() {
let cache = SystemConfigValueCacheState::default();
let mut leader = leader(&cache, "key-a");
let follower = match cache.register("key-a") {
SystemConfigValueLoadRegistration::Follower(state) => state,
_ => panic!("second registration should follow"),
};
leader.finish_loaded(Some(serde_json::json!({"version": 1})));
assert!(matches!(
tokio::time::timeout(Duration::from_millis(100), follower.wait())
.await
.expect("completed follower must not miss the notification"),
SystemConfigValueInflightCompletion::Loaded
));
assert_eq!(
cache.get("key-a"),
Some(Some(serde_json::json!({"version": 1})))
);
}
#[tokio::test]
async fn system_config_follower_receives_leader_failure() {
let cache = SystemConfigValueCacheState::default();
let mut leader = leader(&cache, "key-a");
let follower = match cache.register("key-a") {
SystemConfigValueLoadRegistration::Follower(state) => state,
_ => panic!("second registration should follow"),
};
leader.finish(SystemConfigValueInflightCompletion::Failed(
DataLayerError::Sql("forced system config failure".to_string()),
));
let completion = tokio::time::timeout(Duration::from_millis(100), follower.wait())
.await
.expect("failed load should release its follower");
let SystemConfigValueInflightCompletion::Failed(error) = completion else {
panic!("follower should observe the leader failure");
};
assert_eq!(error.to_string(), "sql error: forced system config failure");
}
#[tokio::test]
async fn invalidation_wins_and_old_guard_preserves_replacement() {
let cache = SystemConfigValueCacheState::default();
let mut old_leader = leader(&cache, "key-a");
let old_follower = match cache.register("key-a") {
SystemConfigValueLoadRegistration::Follower(state) => state,
_ => panic!("second registration should follow"),
};
cache.invalidate("key-a");
let replacement = leader(&cache, "key-a");
old_leader.finish_loaded(Some(serde_json::json!({"stale": true})));
assert!(matches!(
old_follower.wait().await,
SystemConfigValueInflightCompletion::Invalidated
));
assert_eq!(cache.get("key-a"), None);
assert!(matches!(
cache.register("key-a"),
SystemConfigValueLoadRegistration::Follower(_)
));
drop(replacement);
}
#[test]
fn system_config_states_are_independent_and_inflight_is_bounded() {
let first = SystemConfigValueCacheState::default();
let second = SystemConfigValueCacheState::default();
let first_guard = leader(&first, "shared-key");
let second_guard = leader(&second, "shared-key");
drop(first_guard);
drop(second_guard);
let mut guards = Vec::with_capacity(SYSTEM_CONFIG_VALUE_CACHE_MAX_INFLIGHT);
for index in 0..SYSTEM_CONFIG_VALUE_CACHE_MAX_INFLIGHT {
let key = format!("bounded-{index}");
guards.push(leader(&first, &key));
}
assert!(matches!(
first.register("over-capacity"),
SystemConfigValueLoadRegistration::Saturated
));
drop(guards);
}
#[test]
fn capacity_full_cancelled_system_config_follower_can_retry() {
let cache = SystemConfigValueCacheState::default();
let mut active = Vec::with_capacity(SYSTEM_CONFIG_VALUE_CACHE_MAX_INFLIGHT - 1);
for index in 0..SYSTEM_CONFIG_VALUE_CACHE_MAX_INFLIGHT - 1 {
active.push(leader(&cache, &format!("active-{index}")));
}
let current = leader(&cache, "retry-key");
let follower = match cache.register("retry-key") {
SystemConfigValueLoadRegistration::Follower(state) => state,
_ => panic!("same-key request should follow at full capacity"),
};
assert_eq!(cache.admission.available_permits(), 0);
assert!(matches!(
cache.register("over-capacity"),
SystemConfigValueLoadRegistration::Saturated
));
drop(current);
assert!(matches!(
follower.completion.get(),
Some(SystemConfigValueInflightCompletion::Cancelled)
));
let mut replacement = leader(&cache, "retry-key");
replacement.finish(SystemConfigValueInflightCompletion::Cancelled);
assert_eq!(cache.admission.available_permits(), 1);
drop(active);
}
}
impl GatewayDataState {
pub(crate) fn disabled() -> Self {
Self::default()
@@ -478,6 +783,12 @@ impl GatewayDataState {
self.wallet_reader.is_some()
}
#[cfg(test)]
pub(crate) fn without_wallet_reader_for_tests(mut self) -> Self {
self.wallet_reader = None;
self
}
pub(crate) fn has_wallet_writer(&self) -> bool {
self.wallet_writer.is_some()
}
@@ -502,82 +813,87 @@ impl GatewayDataState {
.get(key)
.map(|entry| entry.value.clone()));
}
let cached_value = self
.system_config_value_cache
.read()
.expect("system config value cache lock")
.get(key)
.cloned();
if let Some((cached_at, value)) = cached_value {
if cached_at.elapsed() <= SYSTEM_CONFIG_VALUE_CACHE_TTL {
return Ok(value);
}
if let Some(value) = self.system_config_value_cache.get(key) {
return Ok(value);
}
let load_state = system_config_value_load_state();
loop {
let notified = load_state.notified();
match load_state.register(key) {
SystemConfigValueLoadRegistration::Bypass => {
let Some(backends) = self.backends.as_ref() else {
return Ok(None);
};
let value = crate::request_diagnostics::observe_db_operation(
"system_config_value",
self.database_pool_summary(),
backends.find_system_config_value(key),
)
.await?;
self.system_config_value_cache
.write()
.expect("system config value cache lock")
.insert(key.to_string(), (Instant::now(), value.clone()));
return Ok(value);
match self.system_config_value_cache.register(key) {
SystemConfigValueLoadRegistration::Saturated => {
return Err(DataLayerError::TimedOut(format!(
"system config cache admission saturated for key '{key}'"
)));
}
SystemConfigValueLoadRegistration::Follower => {
notified.await;
let cached_value = self
.system_config_value_cache
.read()
.expect("system config value cache lock")
.get(key)
.cloned();
if let Some((cached_at, value)) = cached_value {
if cached_at.elapsed() <= SYSTEM_CONFIG_VALUE_CACHE_TTL {
SystemConfigValueLoadRegistration::Follower(state) => match state.wait().await {
SystemConfigValueInflightCompletion::Failed(error) => {
return Err(error);
}
SystemConfigValueInflightCompletion::Loaded => {
if let Some(value) = self.system_config_value_cache.get(key) {
return Ok(value);
}
}
}
SystemConfigValueLoadRegistration::Leader(_guard) => {
let cached_value = self
.system_config_value_cache
.read()
.expect("system config value cache lock")
.get(key)
.cloned();
if let Some((cached_at, value)) = cached_value {
if cached_at.elapsed() <= SYSTEM_CONFIG_VALUE_CACHE_TTL {
SystemConfigValueInflightCompletion::Cancelled
| SystemConfigValueInflightCompletion::Invalidated => {}
},
SystemConfigValueLoadRegistration::Leader(mut guard) => {
if let Some(value) = self.system_config_value_cache.get(key) {
guard.finish(SystemConfigValueInflightCompletion::Loaded);
return Ok(value);
}
match self.load_system_config_value_uncached(key).await {
Ok(value) => {
guard.finish_loaded(value.clone());
return Ok(value);
}
Err(error) => {
guard
.finish(SystemConfigValueInflightCompletion::Failed(error.clone()));
return Err(error);
}
}
let Some(backends) = self.backends.as_ref() else {
return Ok(None);
};
let value = crate::request_diagnostics::observe_db_operation(
"system_config_value",
self.database_pool_summary(),
backends.find_system_config_value(key),
)
.await?;
self.system_config_value_cache
.write()
.expect("system config value cache lock")
.insert(key.to_string(), (Instant::now(), value.clone()));
return Ok(value);
}
}
}
}
async fn load_system_config_value_uncached(
&self,
key: &str,
) -> Result<Option<serde_json::Value>, DataLayerError> {
let Some(backends) = self.backends.as_ref() else {
return Ok(None);
};
crate::request_diagnostics::observe_db_operation(
"system_config_value",
self.database_pool_summary(),
backends.find_system_config_value(key),
)
.await
}
pub(crate) async fn find_system_config_value_strong(
&self,
key: &str,
) -> Result<Option<serde_json::Value>, DataLayerError> {
if let Some(values) = &self.system_config_values {
return Ok(values
.read()
.expect("system config values lock")
.get(key)
.map(|entry| entry.value.clone()));
}
let Some(backends) = self.backends.as_ref() else {
return Ok(None);
};
crate::request_diagnostics::observe_db_operation(
"system_config_value_strong",
self.database_pool_summary(),
backends.find_system_config_value(key),
)
.await
}
pub(crate) async fn upsert_system_config_value(
&self,
key: &str,
@@ -668,10 +984,7 @@ impl GatewayDataState {
}
fn clear_cached_system_config_value(&self, key: &str) {
self.system_config_value_cache
.write()
.expect("system config value cache lock")
.remove(key);
self.system_config_value_cache.invalidate(key);
}
pub(crate) async fn read_admin_system_stats(
@@ -687,24 +1000,30 @@ impl GatewayDataState {
&self,
target: aether_data::repository::system::AdminSystemPurgeTarget,
) -> Result<aether_data::repository::system::AdminSystemPurgeSummary, DataLayerError> {
if matches!(
target,
let purges_config = matches!(
&target,
aether_data::repository::system::AdminSystemPurgeTarget::Config
) {
);
if purges_config {
if let Some(values) = &self.system_config_values {
let mut values = values.write().expect("system config values lock");
let deleted = values.len() as u64;
values.clear();
self.system_config_value_cache.clear();
let mut summary =
aether_data::repository::system::AdminSystemPurgeSummary::default();
summary.add("system_configs", deleted);
return Ok(summary);
}
}
match self.backends.as_ref() {
let result = match self.backends.as_ref() {
Some(backends) => backends.purge_admin_system_data(target).await,
None => Ok(aether_data::repository::system::AdminSystemPurgeSummary::default()),
};
if purges_config && result.is_ok() {
self.system_config_value_cache.clear();
}
result
}
pub(crate) async fn export_admin_system_usage_aggregates(
@@ -13,7 +13,7 @@ use aether_data_contracts::repository::provider_catalog::{
};
use aether_data_contracts::repository::settlement::{StoredUsageSettlement, UsageSettlementInput};
use aether_data_contracts::repository::usage::{
ProxyNodeCounterDelta, StoredRequestUsageAudit, UpsertUsageRecord,
ProxyNodeCounterDelta, StoredRequestUsageAudit, UpsertUsageRecord, UsageWriteRepository,
};
use aether_data_contracts::repository::video_tasks::{StoredVideoTask, VideoTaskLookupKey};
use aether_runtime_state::RuntimeQueueStore;
@@ -252,6 +252,12 @@ impl UsageRuntimeAccess for GatewayDataState {
GatewayDataState::usage_worker_queue(self)
}
fn supports_first_byte_usage_fast_path(&self) -> bool {
self.usage_writer
.as_ref()
.is_some_and(|repository| repository.supports_first_byte_usage_fast_path())
}
fn usage_worker_should_defer_for_database_pressure(&self) -> bool {
self.database_pool_summary()
.as_ref()
@@ -323,12 +329,57 @@ impl aether_usage_runtime::ManualProxyNodeCounter for GatewayDataState {
#[async_trait]
impl UsageRecordWriter for GatewayDataState {
fn supports_first_byte_usage_batch(&self) -> bool {
self.usage_writer
.as_ref()
.is_some_and(|repository| repository.supports_first_byte_usage_batch())
}
fn first_byte_usage_writer_identity(&self) -> Option<usize> {
self.usage_writer
.as_ref()
.map(|repository| std::sync::Arc::as_ptr(repository) as *const () as usize)
}
fn supports_pending_usage_batch(&self) -> bool {
self.usage_writer
.as_ref()
.is_some_and(|repository| repository.supports_pending_usage_batch())
}
fn pending_usage_writer_identity(&self) -> Option<usize> {
self.usage_writer
.as_ref()
.map(|repository| std::sync::Arc::as_ptr(repository) as *const () as usize)
}
async fn upsert_usage_record(
&self,
record: UpsertUsageRecord,
) -> Result<Option<StoredRequestUsageAudit>, DataLayerError> {
GatewayDataState::upsert_usage(self, record).await
}
async fn upsert_first_byte_usage_record(
&self,
record: UpsertUsageRecord,
) -> Result<(), DataLayerError> {
GatewayDataState::upsert_first_byte_usage(self, record).await
}
async fn upsert_first_byte_usage_records(
&self,
records: Vec<UpsertUsageRecord>,
) -> Result<(), DataLayerError> {
GatewayDataState::upsert_first_byte_usage_many(self, records).await
}
async fn upsert_pending_usage_records(
&self,
records: Vec<UpsertUsageRecord>,
) -> Result<(), DataLayerError> {
GatewayDataState::upsert_pending_usage_many(self, records).await
}
}
#[cfg(test)]
+39 -10
View File
@@ -2,8 +2,8 @@ use std::collections::BTreeMap;
use std::collections::HashMap;
use std::fmt;
use std::sync::Arc;
use std::sync::OnceLock;
use std::sync::RwLock;
use std::time::Instant;
use super::auth::GatewayAuthApiKeySnapshot;
use super::candidates::{read_request_candidate_trace, RequestCandidateTrace};
@@ -13,6 +13,7 @@ use crate::provider_transport::{
read_provider_transport_snapshot, GatewayProviderTransportSnapshot,
};
use crate::video_tasks::LocalVideoTaskReadResponse;
use aether_cache::ExpiringMap;
use aether_data::repository::announcements::{
AnnouncementListQuery, AnnouncementReadRepository, AnnouncementWriteRepository,
CreateAnnouncementRecord, StoredAnnouncement, StoredAnnouncementPage, UpdateAnnouncementRecord,
@@ -120,8 +121,10 @@ use aether_data_contracts::repository::pool_scores::{
UpsertPoolMemberScore,
};
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyListQuery, ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyHealthStateUpdate,
ProviderCatalogKeyListQuery, ProviderCatalogKeyRuntimeMetadataUpdate,
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository,
ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
};
@@ -196,10 +199,30 @@ pub(crate) struct GatewayDataState {
wallet_writer: Option<Arc<dyn WalletWriteRepository>>,
settlement_writer: Option<Arc<dyn SettlementWriteRepository>>,
system_config_values: Option<Arc<RwLock<BTreeMap<String, StoredSystemConfigEntry>>>>,
system_config_value_cache: Arc<RwLock<BTreeMap<String, (Instant, Option<serde_json::Value>)>>>,
system_config_value_cache: Arc<SystemConfigValueCacheState>,
billing_model_context_cache: Arc<BillingModelContextCacheState>,
}
pub(super) struct SystemConfigValueCacheState {
pub(super) entries: ExpiringMap<String, Option<serde_json::Value>>,
pub(super) inflight: std::sync::Mutex<HashMap<String, Arc<SystemConfigValueInflightState>>>,
pub(super) mutation: std::sync::Mutex<()>,
pub(super) admission: Arc<tokio::sync::Semaphore>,
}
pub(super) struct SystemConfigValueInflightState {
pub(super) notify: Arc<tokio::sync::Notify>,
pub(super) completion: OnceLock<SystemConfigValueInflightCompletion>,
}
#[derive(Clone)]
pub(super) enum SystemConfigValueInflightCompletion {
Loaded,
Failed(DataLayerError),
Cancelled,
Invalidated,
}
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
pub(super) enum BillingModelContextCacheKey {
ByModelId {
@@ -214,14 +237,20 @@ pub(super) enum BillingModelContextCacheKey {
},
}
#[derive(Default)]
pub(super) struct BillingModelContextCacheState {
pub(super) entries:
RwLock<HashMap<BillingModelContextCacheKey, (Instant, Option<StoredBillingModelContext>)>>,
pub(super) inflight: std::sync::Mutex<HashMap<BillingModelContextCacheKey, u64>>,
pub(super) inflight_notify: tokio::sync::Notify,
pub(super) next_inflight_token: std::sync::atomic::AtomicU64,
pub(super) entries: ExpiringMap<BillingModelContextCacheKey, Option<StoredBillingModelContext>>,
pub(super) inflight: std::sync::Mutex<
HashMap<BillingModelContextCacheKey, Arc<BillingModelContextInflightState>>,
>,
pub(super) epoch: std::sync::atomic::AtomicU64,
pub(super) mutation: std::sync::Mutex<()>,
pub(super) admission: Arc<tokio::sync::Semaphore>,
}
pub(super) struct BillingModelContextInflightState {
pub(super) epoch: u64,
pub(super) completion: std::sync::OnceLock<Result<(), DataLayerError>>,
pub(super) notify: tokio::sync::Notify,
}
impl fmt::Debug for GatewayDataState {
@@ -4,6 +4,7 @@ use super::{
PoolMemberHardState, PoolMemberIdentity, PoolMemberProbeAttempt, PoolMemberProbeResult,
PoolMemberScheduleFeedback, PoolScoreScope, StoredPoolMemberScore, UpsertPoolMemberScore,
};
use aether_data_contracts::repository::pool_scores::PoolMemberScoreUpsertMode;
impl GatewayDataState {
pub(crate) async fn list_ranked_pool_members(
@@ -56,6 +57,20 @@ impl GatewayDataState {
}
}
pub(crate) async fn upsert_pool_member_score_with_mode(
&self,
score: UpsertPoolMemberScore,
mode: PoolMemberScoreUpsertMode,
) -> Result<Option<StoredPoolMemberScore>, DataLayerError> {
match &self.pool_score_writer {
Some(repository) => repository
.upsert_pool_member_score_with_mode(score, mode)
.await
.map(Some),
None => Ok(None),
}
}
pub(crate) async fn record_pool_member_probe_result(
&self,
result: PoolMemberProbeResult,
@@ -1,7 +1,7 @@
use std::collections::HashMap;
use std::future::Future;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::sync::{Arc, Mutex, OnceLock};
use std::time::Duration;
use aether_cache::ExpiringMap;
@@ -16,14 +16,16 @@ use tokio::sync::Notify;
const PROVIDER_CATALOG_CACHE_TTL: Duration = Duration::from_secs(5);
const PROVIDER_CATALOG_CACHE_MAX_ENTRIES: usize = 1024;
const PROVIDER_CATALOG_CACHE_MAX_INFLIGHT: usize = 1024;
const PROVIDER_CATALOG_CACHE_LOAD_TIMEOUT: Duration = Duration::from_secs(10);
pub(super) struct CachedProviderCatalogReadRepository {
inner: Arc<dyn ProviderCatalogReadRepository>,
entries: ExpiringMap<ProviderCatalogCacheKey, ProviderCatalogCacheValue>,
inflight: Mutex<HashMap<ProviderCatalogCacheKey, u64>>,
inflight_notify: Notify,
next_inflight_token: AtomicU64,
inflight: Mutex<HashMap<ProviderCatalogCacheKey, Arc<ProviderCatalogInflightState>>>,
admission: Arc<tokio::sync::Semaphore>,
epoch: AtomicU64,
mutation: Mutex<()>,
}
impl CachedProviderCatalogReadRepository {
@@ -32,9 +34,11 @@ impl CachedProviderCatalogReadRepository {
inner,
entries: ExpiringMap::new(),
inflight: Mutex::new(HashMap::new()),
inflight_notify: Notify::new(),
next_inflight_token: AtomicU64::new(1),
admission: Arc::new(tokio::sync::Semaphore::new(
PROVIDER_CATALOG_CACHE_MAX_INFLIGHT,
)),
epoch: AtomicU64::new(0),
mutation: Mutex::new(()),
}
}
@@ -47,78 +51,181 @@ impl CachedProviderCatalogReadRepository {
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();
if let Some(value) = self.entries.get_fresh(&key, PROVIDER_CATALOG_CACHE_TTL) {
return Ok(value);
}
match self.register_inflight(&key) {
InflightRegistration::Bypass => return load().await,
InflightRegistration::Follower => {
notified.await;
InflightRegistration::Saturated => {
return Err(DataLayerError::TimedOut(format!(
"provider catalog cache admission saturated for {key:?}"
)));
}
InflightRegistration::Follower(state) => {
state.wait().await;
match self.follower_completion(&state) {
Some(ProviderCatalogInflightCompletion::Loaded(value)) => return Ok(value),
Some(ProviderCatalogInflightCompletion::Failed(error)) => {
return Err(error);
}
Some(
ProviderCatalogInflightCompletion::Cancelled
| ProviderCatalogInflightCompletion::Invalidated,
)
| None => continue,
}
}
InflightRegistration::Leader(mut guard) => {
if let Some(value) = self.entries.get_fresh(&key, PROVIDER_CATALOG_CACHE_TTL) {
guard.finish(ProviderCatalogInflightCompletion::Loaded(value.clone()));
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,
);
let result =
match tokio::time::timeout(PROVIDER_CATALOG_CACHE_LOAD_TIMEOUT, load())
.await
{
Ok(result) => result,
Err(_) => Err(DataLayerError::TimedOut(format!(
"provider catalog cache load exceeded {}ms for {key:?}",
PROVIDER_CATALOG_CACHE_LOAD_TIMEOUT.as_millis()
))),
};
match result {
Ok(value) => {
guard.finish_loaded(value.clone());
return Ok(value);
}
Err(error) => {
guard.finish(ProviderCatalogInflightCompletion::Failed(error.clone()));
return Err(error);
}
}
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)
fn register_inflight(&self, key: &ProviderCatalogCacheKey) -> InflightRegistration<'_> {
{
let inflight = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(state) = inflight.get(key) {
return InflightRegistration::Follower(Arc::clone(state));
}
Err(_) => InflightRegistration::Bypass,
}
let _mutation = self
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let mut inflight = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(state) = inflight.get(key) {
return InflightRegistration::Follower(Arc::clone(state));
}
if inflight.len() >= PROVIDER_CATALOG_CACHE_MAX_INFLIGHT {
return InflightRegistration::Saturated;
}
let Ok(admission) = Arc::clone(&self.admission).try_acquire_owned() else {
return InflightRegistration::Saturated;
};
let state = Arc::new(ProviderCatalogInflightState {
notify: Arc::new(Notify::new()),
completion: OnceLock::new(),
epoch: self.epoch.load(Ordering::Acquire),
});
inflight.insert(key.clone(), Arc::clone(&state));
InflightRegistration::Leader(InflightGuard {
cache: self,
key: Some(key.clone()),
state,
admission: Some(admission),
})
}
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) {
fn finish_inflight(
&self,
key: &ProviderCatalogCacheKey,
state: &Arc<ProviderCatalogInflightState>,
admission: tokio::sync::OwnedSemaphorePermit,
completion: ProviderCatalogInflightCompletion,
cache_value: Option<ProviderCatalogCacheValue>,
) {
let _mutation = self
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let removed = {
let mut inflight = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
drop(admission);
debug_assert!(self.admission.available_permits() > 0);
if inflight
.get(key)
.is_some_and(|current| Arc::ptr_eq(current, state))
&& self.epoch.load(Ordering::Acquire) == state.epoch
{
if let Some(value) = cache_value {
self.entries.insert(
key.clone(),
value,
PROVIDER_CATALOG_CACHE_TTL,
PROVIDER_CATALOG_CACHE_MAX_ENTRIES,
);
}
state.complete(completion);
inflight.remove(key);
removed = true;
true
} else {
false
}
}
};
drop(_mutation);
if removed {
self.inflight_notify.notify_waiters();
state.notify.notify_waiters();
}
}
fn follower_completion(
&self,
state: &ProviderCatalogInflightState,
) -> Option<ProviderCatalogInflightCompletion> {
let before = self.epoch.load(Ordering::Acquire);
let completion = state.completion.get().cloned();
let after = self.epoch.load(Ordering::Acquire);
(before == state.epoch && before == after)
.then_some(completion)
.flatten()
}
fn clear(&self) {
let _mutation = self
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
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();
let states = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.drain()
.map(|(_, state)| state)
.collect::<Vec<_>>();
for state in &states {
state.complete(ProviderCatalogInflightCompletion::Invalidated);
}
if cleared_inflight {
self.inflight_notify.notify_waiters();
drop(_mutation);
for state in states {
state.notify.notify_waiters();
}
}
}
@@ -335,41 +442,83 @@ enum ProviderCatalogCacheValue {
KeyStats(Vec<StoredProviderCatalogKeyStats>),
}
enum InflightRegistration {
Leader(u64),
Follower,
Bypass,
struct ProviderCatalogInflightState {
notify: Arc<Notify>,
completion: OnceLock<ProviderCatalogInflightCompletion>,
epoch: u64,
}
impl ProviderCatalogInflightState {
fn complete(&self, completion: ProviderCatalogInflightCompletion) {
let _ = self.completion.set(completion);
}
async fn wait(&self) {
loop {
if self.completion.get().is_some() {
return;
}
let mut notified = Box::pin(self.notify.notified());
notified.as_mut().enable();
if self.completion.get().is_some() {
return;
}
notified.await;
}
}
}
#[derive(Clone)]
enum ProviderCatalogInflightCompletion {
Loaded(ProviderCatalogCacheValue),
Failed(DataLayerError),
Cancelled,
Invalidated,
}
enum InflightRegistration<'a> {
Leader(InflightGuard<'a>),
Follower(Arc<ProviderCatalogInflightState>),
Saturated,
}
struct InflightGuard<'a> {
cache: &'a CachedProviderCatalogReadRepository,
key: Option<ProviderCatalogCacheKey>,
token: u64,
state: Arc<ProviderCatalogInflightState>,
admission: Option<tokio::sync::OwnedSemaphorePermit>,
}
impl<'a> InflightGuard<'a> {
fn new(
cache: &'a CachedProviderCatalogReadRepository,
key: ProviderCatalogCacheKey,
token: u64,
) -> Self {
Self {
cache,
key: Some(key),
token,
}
impl InflightGuard<'_> {
fn finish_loaded(&mut self, value: ProviderCatalogCacheValue) {
let completion = ProviderCatalogInflightCompletion::Loaded(value.clone());
self.finish_with_cache(completion, Some(value));
}
fn finish(&mut self) {
fn finish(&mut self, completion: ProviderCatalogInflightCompletion) {
self.finish_with_cache(completion, None);
}
fn finish_with_cache(
&mut self,
completion: ProviderCatalogInflightCompletion,
cache_value: Option<ProviderCatalogCacheValue>,
) {
if let Some(key) = self.key.take() {
self.cache.finish_inflight(&key, self.token);
let admission = self
.admission
.take()
.expect("active provider catalog leader must own admission");
self.cache
.finish_inflight(&key, &self.state, admission, completion, cache_value);
}
}
}
impl Drop for InflightGuard<'_> {
fn drop(&mut self) {
self.finish();
self.finish(ProviderCatalogInflightCompletion::Cancelled);
}
}
@@ -384,3 +533,173 @@ fn normalize_ids(ids: &[String]) -> Vec<String> {
normalized.dedup();
normalized
}
#[cfg(test)]
mod tests {
use super::*;
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
fn cache() -> CachedProviderCatalogReadRepository {
CachedProviderCatalogReadRepository::new(Arc::new(
InMemoryProviderCatalogReadRepository::seed(Vec::new(), Vec::new(), Vec::new()),
))
}
fn provider(id: &str) -> StoredProviderCatalogProvider {
StoredProviderCatalogProvider::new(
id.to_string(),
id.to_string(),
None,
"openai".to_string(),
)
.expect("provider should be valid")
}
#[tokio::test]
async fn provider_catalog_follower_observes_completion_before_first_poll() {
let cache = cache();
let key = ProviderCatalogCacheKey::Providers { active_only: false };
let mut leader = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
let follower = match cache.register_inflight(&key) {
InflightRegistration::Follower(state) => state,
_ => panic!("second registration should follow"),
};
leader.finish(ProviderCatalogInflightCompletion::Loaded(
ProviderCatalogCacheValue::Providers(Vec::new()),
));
tokio::time::timeout(Duration::from_millis(100), follower.wait())
.await
.expect("completion before the first poll must release the follower");
assert!(matches!(
cache.follower_completion(&follower),
Some(ProviderCatalogInflightCompletion::Loaded(_))
));
}
#[tokio::test]
async fn provider_catalog_follower_receives_leader_failure() {
let cache = cache();
let key = ProviderCatalogCacheKey::Providers { active_only: false };
let mut leader = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
let follower = match cache.register_inflight(&key) {
InflightRegistration::Follower(state) => state,
_ => panic!("second registration should follow"),
};
leader.finish(ProviderCatalogInflightCompletion::Failed(
DataLayerError::Sql("forced provider catalog failure".to_string()),
));
tokio::time::timeout(Duration::from_millis(100), follower.wait())
.await
.expect("failed load must release the follower");
let Some(ProviderCatalogInflightCompletion::Failed(error)) =
cache.follower_completion(&follower)
else {
panic!("follower should observe the failed completion");
};
assert_eq!(
error.to_string(),
"sql error: forced provider catalog failure"
);
}
#[test]
fn provider_catalog_old_guard_cannot_remove_replacement_after_clear() {
let cache = cache();
let key = ProviderCatalogCacheKey::Providers { active_only: false };
let old_leader = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
cache.clear();
let replacement = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("clear should admit a replacement"),
};
drop(old_leader);
assert!(matches!(
cache.register_inflight(&key),
InflightRegistration::Follower(_)
));
drop(replacement);
}
#[test]
fn provider_catalog_old_flight_after_clear_cannot_overwrite_replacement() {
let cache = cache();
let key = ProviderCatalogCacheKey::Providers { active_only: false };
let mut old_leader = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
cache.clear();
let mut replacement = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("clear should admit a replacement"),
};
replacement.finish_loaded(ProviderCatalogCacheValue::Providers(vec![provider(
"fresh",
)]));
old_leader.finish_loaded(ProviderCatalogCacheValue::Providers(vec![provider(
"stale",
)]));
let Some(ProviderCatalogCacheValue::Providers(cached)) =
cache.entries.get_fresh(&key, PROVIDER_CATALOG_CACHE_TTL)
else {
panic!("replacement value should remain cached");
};
assert_eq!(cached[0].id, "fresh");
}
#[test]
fn provider_catalog_capacity_full_cancelled_follower_can_retry_after_repeated_clear() {
let cache = cache();
let key = ProviderCatalogCacheKey::Providers { active_only: false };
let mut active = Vec::with_capacity(PROVIDER_CATALOG_CACHE_MAX_INFLIGHT);
for _ in 0..PROVIDER_CATALOG_CACHE_MAX_INFLIGHT - 1 {
let leader = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("each available permit should admit one leader"),
};
active.push(leader);
cache.clear();
}
let current = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("the final available permit should admit a leader"),
};
let follower = match cache.register_inflight(&key) {
InflightRegistration::Follower(state) => state,
_ => panic!("the same-key request should follow at full capacity"),
};
assert_eq!(cache.admission.available_permits(), 0);
assert!(matches!(
cache.register_inflight(&ProviderCatalogCacheKey::Providers { active_only: true }),
InflightRegistration::Saturated
));
drop(current);
assert!(matches!(
cache.follower_completion(&follower),
Some(ProviderCatalogInflightCompletion::Cancelled)
));
let mut replacement = match cache.register_inflight(&key) {
InflightRegistration::Leader(guard) => guard,
_ => panic!("cancelled follower retry should use the released permit"),
};
replacement.finish(ProviderCatalogInflightCompletion::Cancelled);
assert_eq!(cache.admission.available_permits(), 1);
}
}
@@ -1,4 +1,7 @@
use std::sync::Arc;
use std::collections::HashMap;
use std::future::Future;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex, OnceLock};
use std::time::Duration;
use aether_cache::ExpiringMap;
@@ -9,16 +12,22 @@ use aether_data_contracts::repository::routing_profiles::{
StoredRoutingGroupVersion,
};
use async_trait::async_trait;
use dashmap::DashMap;
use tokio::sync::Notify;
// Routing profile writes clear this local cache. Keep the fallback TTL bounded
// so a missed cross-node invalidation cannot route traffic stale for minutes.
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;
const ROUTING_GROUP_CACHE_MAX_INFLIGHT: usize = 4_096;
const ROUTING_GROUP_CACHE_LOAD_TIMEOUT: Duration = Duration::from_secs(10);
pub(super) struct CachedRoutingGroupReadRepository {
inner: Arc<dyn RoutingGroupReadRepository>,
entries: ExpiringMap<RoutingGroupCacheKey, RoutingGroupCacheValue>,
load_guards: DashMap<RoutingGroupCacheKey, Arc<tokio::sync::Mutex<()>>>,
inflight: Mutex<HashMap<RoutingGroupCacheKey, Arc<RoutingGroupInflightState>>>,
admission: Arc<tokio::sync::Semaphore>,
generation: AtomicU64,
mutation: Mutex<()>,
}
impl CachedRoutingGroupReadRepository {
@@ -26,58 +35,245 @@ impl CachedRoutingGroupReadRepository {
Self {
inner,
entries: ExpiringMap::new(),
load_guards: DashMap::new(),
inflight: Mutex::new(HashMap::new()),
admission: Arc::new(tokio::sync::Semaphore::new(
ROUTING_GROUP_CACHE_MAX_INFLIGHT,
)),
generation: AtomicU64::new(0),
mutation: Mutex::new(()),
}
}
fn clear(&self) {
let _mutation = self
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
self.generation.fetch_add(1, Ordering::AcqRel);
self.entries.clear();
self.load_guards.clear();
let states = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.drain()
.map(|(_, state)| {
let _ = state
.completion
.set(RoutingGroupInflightCompletion::Invalidated);
state
})
.collect::<Vec<_>>();
drop(_mutation);
for state in states {
state.notify.notify_waiters();
}
}
async fn get_or_load(
async fn get_or_load<F, Fut>(
&self,
key: RoutingGroupCacheKey,
load: impl std::future::Future<Output = Result<RoutingGroupCacheValue, DataLayerError>>,
) -> Result<RoutingGroupCacheValue, DataLayerError> {
if let Some((value, _age)) = self
.entries
.get_with_age(&key, ROUTING_GROUP_CACHE_STALE_TTL)
{
return Ok(value);
mut load: F,
) -> Result<RoutingGroupCacheValue, DataLayerError>
where
F: FnMut() -> Fut,
Fut: Future<Output = Result<RoutingGroupCacheValue, DataLayerError>>,
{
loop {
if let Some(value) = self.cached_value(&key) {
return Ok(value);
}
match self.register_inflight(&key) {
RoutingGroupInflightRegistration::Saturated => {
return Err(DataLayerError::TimedOut(format!(
"routing group cache admission saturated for {key:?}"
)));
}
RoutingGroupInflightRegistration::Leader(mut guard) => {
// A previous leader can publish between the optimistic
// cache read and this registration.
if let Some(value) = self.cached_value(&key) {
guard.finish(RoutingGroupInflightCompletion::Loaded(value.clone()));
return Ok(value);
}
let result = match tokio::time::timeout(
ROUTING_GROUP_CACHE_LOAD_TIMEOUT,
load(),
)
.await
{
Ok(result) => result,
Err(_) => Err(DataLayerError::TimedOut(format!(
"routing group cache load exceeded {}ms for {key:?}",
ROUTING_GROUP_CACHE_LOAD_TIMEOUT.as_millis()
))),
};
match result {
Ok(value) => {
self.insert_if_generation(key.clone(), value.clone(), guard.generation);
guard.finish(RoutingGroupInflightCompletion::Loaded(value.clone()));
return Ok(value);
}
Err(error) => {
guard.finish(RoutingGroupInflightCompletion::Failed(
SharedDataLayerError::from(&error),
));
return Err(error);
}
}
}
RoutingGroupInflightRegistration::Follower(state) => {
state.wait().await;
match self.follower_completion(&state) {
Some(RoutingGroupInflightCompletion::Loaded(value)) => return Ok(value),
Some(RoutingGroupInflightCompletion::Failed(error)) => {
return Err(error.into_data_layer_error());
}
Some(
RoutingGroupInflightCompletion::Cancelled
| RoutingGroupInflightCompletion::Invalidated,
)
| None => continue,
}
}
}
}
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);
}
fn cached_value(&self, key: &RoutingGroupCacheKey) -> Option<RoutingGroupCacheValue> {
self.entries
.get_with_age(key, ROUTING_GROUP_CACHE_STALE_TTL)
.map(|(value, _age)| value)
}
fn follower_completion(
&self,
state: &RoutingGroupInflightState,
) -> Option<RoutingGroupInflightCompletion> {
// Double-check the generation around the completion read. This keeps
// the hot follower path lock-free while ensuring clear() cannot race
// between an old-generation check and returning a loaded value.
let before = self.generation.load(Ordering::Acquire);
let completion = state.completion.get().cloned();
let after = self.generation.load(Ordering::Acquire);
(before == state.generation && before == after)
.then_some(completion)
.flatten()
}
fn insert_if_generation(
&self,
key: RoutingGroupCacheKey,
value: RoutingGroupCacheValue,
generation: u64,
) {
let _mutation = self
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if self.generation.load(Ordering::Acquire) != generation {
return;
}
let value = load.await?;
self.entries.insert(
key,
value.clone(),
value,
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();
fn register_inflight(
&self,
key: &RoutingGroupCacheKey,
) -> RoutingGroupInflightRegistration<'_> {
// Existing followers only touch the per-key map. Avoid taking the
// mutation lock on the 20k-request hot path; that lock is reserved
// for leader insertion and invalidation ordering.
{
let inflight = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(state) = inflight.get(key) {
return RoutingGroupInflightRegistration::Follower(Arc::clone(state));
}
}
let _mutation = self
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let mut inflight = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(state) = inflight.get(key) {
return RoutingGroupInflightRegistration::Follower(Arc::clone(state));
}
if inflight.len() >= ROUTING_GROUP_CACHE_MAX_INFLIGHT {
return RoutingGroupInflightRegistration::Saturated;
}
let Ok(admission) = Arc::clone(&self.admission).try_acquire_owned() else {
return RoutingGroupInflightRegistration::Saturated;
};
let state = Arc::new(RoutingGroupInflightState {
notify: Arc::new(Notify::new()),
completion: OnceLock::new(),
generation: self.generation.load(Ordering::Acquire),
});
let generation = state.generation;
inflight.insert(key.clone(), Arc::clone(&state));
RoutingGroupInflightRegistration::Leader(RoutingGroupInflightGuard {
cache: self,
key: Some(key.clone()),
state,
generation,
admission: Some(admission),
})
}
fn finish_inflight(
&self,
key: &RoutingGroupCacheKey,
state: &Arc<RoutingGroupInflightState>,
admission: tokio::sync::OwnedSemaphorePermit,
completion: RoutingGroupInflightCompletion,
) {
let _mutation = self
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let removed = {
let mut inflight = self
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
drop(admission);
debug_assert!(self.admission.available_permits() > 0);
if inflight
.get(key)
.is_some_and(|current| Arc::ptr_eq(current, state))
{
let _ = state.completion.set(completion);
inflight.remove(key);
true
} else {
false
}
};
drop(_mutation);
if removed {
state.notify.notify_waiters();
}
self.load_guards
.entry(key.clone())
.or_insert_with(|| Arc::new(tokio::sync::Mutex::new(())))
.clone()
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
enum RoutingGroupCacheKey {
ListGroups,
HasAnyBinding,
FindById(String),
FindByName(String),
FindSystemDefault,
@@ -92,11 +288,120 @@ enum RoutingGroupCacheKey {
#[derive(Debug, Clone)]
enum RoutingGroupCacheValue {
Groups(Vec<StoredRoutingGroup>),
Bool(bool),
Group(Option<StoredRoutingGroup>),
Bindings(Vec<StoredRoutingGroupBinding>),
Versions(Vec<StoredRoutingGroupVersion>),
}
struct RoutingGroupInflightState {
notify: Arc<Notify>,
completion: OnceLock<RoutingGroupInflightCompletion>,
generation: u64,
}
impl RoutingGroupInflightState {
async fn wait(&self) {
loop {
if self.completion.get().is_some() {
return;
}
// Register before checking completion a second time. This closes
// the completion/notify race even if the follower is not polled
// until after the leader has broadcast with notify_waiters().
let mut notified = Box::pin(self.notify.notified());
notified.as_mut().enable();
if self.completion.get().is_some() {
return;
}
notified.await;
}
}
}
#[derive(Clone)]
enum RoutingGroupInflightCompletion {
Loaded(RoutingGroupCacheValue),
Failed(SharedDataLayerError),
Cancelled,
Invalidated,
}
enum RoutingGroupInflightRegistration<'a> {
Leader(RoutingGroupInflightGuard<'a>),
Follower(Arc<RoutingGroupInflightState>),
Saturated,
}
struct RoutingGroupInflightGuard<'a> {
cache: &'a CachedRoutingGroupReadRepository,
key: Option<RoutingGroupCacheKey>,
state: Arc<RoutingGroupInflightState>,
generation: u64,
admission: Option<tokio::sync::OwnedSemaphorePermit>,
}
impl RoutingGroupInflightGuard<'_> {
fn finish(&mut self, completion: RoutingGroupInflightCompletion) {
if let Some(key) = self.key.take() {
let admission = self
.admission
.take()
.expect("active routing group leader must own admission");
self.cache
.finish_inflight(&key, &self.state, admission, completion);
}
}
}
impl Drop for RoutingGroupInflightGuard<'_> {
fn drop(&mut self) {
self.finish(RoutingGroupInflightCompletion::Cancelled);
}
}
#[derive(Clone)]
enum SharedDataLayerError {
InvalidConfiguration(String),
InvalidInput(String),
Postgres(String),
Redis(String),
Sql(String),
TimedOut(String),
UnexpectedValue(String),
}
impl From<&DataLayerError> for SharedDataLayerError {
fn from(error: &DataLayerError) -> Self {
match error {
DataLayerError::InvalidConfiguration(message) => {
Self::InvalidConfiguration(message.clone())
}
DataLayerError::InvalidInput(message) => Self::InvalidInput(message.clone()),
DataLayerError::Postgres(message) => Self::Postgres(message.clone()),
DataLayerError::Redis(message) => Self::Redis(message.clone()),
DataLayerError::Sql(message) => Self::Sql(message.clone()),
DataLayerError::TimedOut(message) => Self::TimedOut(message.clone()),
DataLayerError::UnexpectedValue(message) => Self::UnexpectedValue(message.clone()),
}
}
}
impl SharedDataLayerError {
fn into_data_layer_error(self) -> DataLayerError {
match self {
Self::InvalidConfiguration(message) => DataLayerError::InvalidConfiguration(message),
Self::InvalidInput(message) => DataLayerError::InvalidInput(message),
Self::Postgres(message) => DataLayerError::Postgres(message),
Self::Redis(message) => DataLayerError::Redis(message),
Self::Sql(message) => DataLayerError::Sql(message),
Self::TimedOut(message) => DataLayerError::TimedOut(message),
Self::UnexpectedValue(message) => DataLayerError::UnexpectedValue(message),
}
}
}
fn lookup_cache_key(lookup: &RoutingGroupLookupKey<'_>) -> RoutingGroupCacheKey {
match lookup {
RoutingGroupLookupKey::Id(id) => RoutingGroupCacheKey::FindById((*id).to_string()),
@@ -122,7 +427,7 @@ impl RoutingGroupReadRepository for CachedRoutingGroupReadRepository {
async fn list_routing_groups(&self) -> Result<Vec<StoredRoutingGroup>, DataLayerError> {
match self
.get_or_load(RoutingGroupCacheKey::ListGroups, async {
.get_or_load(RoutingGroupCacheKey::ListGroups, || async {
self.inner
.list_routing_groups()
.await
@@ -140,12 +445,16 @@ impl RoutingGroupReadRepository for CachedRoutingGroupReadRepository {
lookup: RoutingGroupLookupKey<'_>,
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
let key = lookup_cache_key(&lookup);
let lookup_for_load = lookup.clone();
match self
.get_or_load(key, async {
self.inner
.find_routing_group(lookup)
.await
.map(RoutingGroupCacheValue::Group)
.get_or_load(key, move || {
let lookup = lookup_for_load.clone();
async move {
self.inner
.find_routing_group(lookup)
.await
.map(RoutingGroupCacheValue::Group)
}
})
.await?
{
@@ -164,7 +473,7 @@ impl RoutingGroupReadRepository for CachedRoutingGroupReadRepository {
subject_id: query.subject_id.clone(),
};
match self
.get_or_load(key, async {
.get_or_load(key, || async {
self.inner
.list_routing_group_bindings(query)
.await
@@ -177,13 +486,28 @@ impl RoutingGroupReadRepository for CachedRoutingGroupReadRepository {
}
}
async fn has_any_routing_group_binding(&self) -> Result<bool, DataLayerError> {
match self
.get_or_load(RoutingGroupCacheKey::HasAnyBinding, || async {
self.inner
.has_any_routing_group_binding()
.await
.map(RoutingGroupCacheValue::Bool)
})
.await?
{
RoutingGroupCacheValue::Bool(value) => Ok(value),
_ => Ok(false),
}
}
async fn list_routing_group_versions(
&self,
group_id: &str,
) -> Result<Vec<StoredRoutingGroupVersion>, DataLayerError> {
let key = RoutingGroupCacheKey::Versions(group_id.to_string());
match self
.get_or_load(key, async {
.get_or_load(key, || async {
self.inner
.list_routing_group_versions(group_id)
.await
@@ -201,11 +525,14 @@ impl RoutingGroupReadRepository for CachedRoutingGroupReadRepository {
mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::sync::{Barrier, Semaphore};
use super::*;
#[derive(Default)]
struct CountingRoutingGroupReadRepository {
list_calls: AtomicUsize,
has_any_binding_calls: AtomicUsize,
}
#[async_trait]
@@ -229,6 +556,11 @@ mod tests {
Ok(Vec::new())
}
async fn has_any_routing_group_binding(&self) -> Result<bool, DataLayerError> {
self.has_any_binding_calls.fetch_add(1, Ordering::AcqRel);
Ok(false)
}
async fn list_routing_group_versions(
&self,
_group_id: &str,
@@ -237,6 +569,131 @@ mod tests {
}
}
#[derive(Clone, Copy)]
enum ControlledHasAnyBindingBehavior {
WaitThenSuccess(bool),
WaitThenError,
FirstWaitFalseThenTrue,
FirstPendingThenTrue,
}
struct ControlledRoutingGroupReadRepository {
behavior: ControlledHasAnyBindingBehavior,
has_any_binding_calls: AtomicUsize,
started: Semaphore,
release: Semaphore,
}
impl ControlledRoutingGroupReadRepository {
fn new(behavior: ControlledHasAnyBindingBehavior) -> Self {
Self {
behavior,
has_any_binding_calls: AtomicUsize::new(0),
started: Semaphore::new(0),
release: Semaphore::new(0),
}
}
fn calls(&self) -> usize {
self.has_any_binding_calls.load(Ordering::Acquire)
}
async fn wait_until_started(&self) {
self.started
.acquire()
.await
.expect("started semaphore should remain open")
.forget();
}
async fn wait_for_release(&self) {
self.release
.acquire()
.await
.expect("release semaphore should remain open")
.forget();
}
}
#[async_trait]
impl RoutingGroupReadRepository for ControlledRoutingGroupReadRepository {
async fn list_routing_groups(&self) -> Result<Vec<StoredRoutingGroup>, DataLayerError> {
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 has_any_routing_group_binding(&self) -> Result<bool, DataLayerError> {
let call = self.has_any_binding_calls.fetch_add(1, Ordering::AcqRel) + 1;
self.started.add_permits(1);
match self.behavior {
ControlledHasAnyBindingBehavior::WaitThenSuccess(value) => {
self.wait_for_release().await;
Ok(value)
}
ControlledHasAnyBindingBehavior::WaitThenError => {
self.wait_for_release().await;
Err(DataLayerError::TimedOut(
"shared routing load error".to_string(),
))
}
ControlledHasAnyBindingBehavior::FirstWaitFalseThenTrue if call == 1 => {
self.wait_for_release().await;
Ok(false)
}
ControlledHasAnyBindingBehavior::FirstPendingThenTrue if call == 1 => {
std::future::pending().await
}
ControlledHasAnyBindingBehavior::FirstWaitFalseThenTrue
| ControlledHasAnyBindingBehavior::FirstPendingThenTrue => Ok(true),
}
}
async fn list_routing_group_versions(
&self,
_group_id: &str,
) -> Result<Vec<StoredRoutingGroupVersion>, DataLayerError> {
Ok(Vec::new())
}
}
async fn wait_for_same_key_inflight_participants(
repository: &CachedRoutingGroupReadRepository,
participants: usize,
) {
tokio::time::timeout(Duration::from_secs(1), async {
loop {
let participant_count = repository
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.get(&RoutingGroupCacheKey::HasAnyBinding)
.map(Arc::strong_count)
.unwrap_or_default();
// One Arc is retained by the map and every active request
// owns one through its leader guard or follower state.
if participant_count >= participants + 1 {
break;
}
tokio::task::yield_now().await;
}
})
.await
.expect("all same-key requests should join one inflight load");
}
#[tokio::test]
async fn clear_local_cache_forces_next_load() {
let inner = Arc::new(CountingRoutingGroupReadRepository::default());
@@ -252,11 +709,236 @@ mod tests {
.expect("cached list should load");
assert_eq!(inner.list_calls.load(Ordering::Acquire), 1);
assert!(!repository
.has_any_routing_group_binding()
.await
.expect("initial binding existence should load"));
assert!(!repository
.has_any_routing_group_binding()
.await
.expect("cached binding existence should load"));
assert_eq!(inner.has_any_binding_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);
assert!(!repository
.has_any_routing_group_binding()
.await
.expect("cleared binding existence should reload"));
assert_eq!(inner.has_any_binding_calls.load(Ordering::Acquire), 2);
}
#[tokio::test]
async fn follower_does_not_miss_completion_before_first_poll() {
let inner = Arc::new(CountingRoutingGroupReadRepository::default());
let repository = CachedRoutingGroupReadRepository::new(inner);
let key = RoutingGroupCacheKey::HasAnyBinding;
let mut leader = match repository.register_inflight(&key) {
RoutingGroupInflightRegistration::Leader(leader) => leader,
_ => panic!("first registration should lead"),
};
let follower = match repository.register_inflight(&key) {
RoutingGroupInflightRegistration::Follower(state) => state,
_ => panic!("second registration should follow"),
};
// Complete before the follower wait future is created or polled. A
// bare notify_waiters().await sequence would sleep forever here.
leader.finish(RoutingGroupInflightCompletion::Loaded(
RoutingGroupCacheValue::Bool(true),
));
tokio::time::timeout(Duration::from_millis(100), follower.wait())
.await
.expect("completed follower should not miss the broadcast");
assert!(matches!(
repository.follower_completion(&follower),
Some(RoutingGroupInflightCompletion::Loaded(
RoutingGroupCacheValue::Bool(true)
))
));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_same_key_requests_share_one_completed_load() {
const TASKS: usize = 64;
let inner = Arc::new(ControlledRoutingGroupReadRepository::new(
ControlledHasAnyBindingBehavior::WaitThenSuccess(true),
));
let repository = Arc::new(CachedRoutingGroupReadRepository::new(inner.clone()));
let barrier = Arc::new(Barrier::new(TASKS + 1));
let mut tasks = Vec::with_capacity(TASKS);
for _ in 0..TASKS {
let repository = Arc::clone(&repository);
let barrier = Arc::clone(&barrier);
tasks.push(tokio::spawn(async move {
barrier.wait().await;
repository.has_any_routing_group_binding().await
}));
}
barrier.wait().await;
inner.wait_until_started().await;
wait_for_same_key_inflight_participants(&repository, TASKS).await;
inner.release.add_permits(TASKS);
for task in tasks {
assert!(task
.await
.expect("request task should finish")
.expect("shared load should succeed"));
}
assert_eq!(inner.calls(), 1);
assert!(repository
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.is_empty());
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_same_key_followers_share_leader_error() {
const TASKS: usize = 32;
let inner = Arc::new(ControlledRoutingGroupReadRepository::new(
ControlledHasAnyBindingBehavior::WaitThenError,
));
let repository = Arc::new(CachedRoutingGroupReadRepository::new(inner.clone()));
let barrier = Arc::new(Barrier::new(TASKS + 1));
let mut tasks = Vec::with_capacity(TASKS);
for _ in 0..TASKS {
let repository = Arc::clone(&repository);
let barrier = Arc::clone(&barrier);
tasks.push(tokio::spawn(async move {
barrier.wait().await;
repository.has_any_routing_group_binding().await
}));
}
barrier.wait().await;
inner.wait_until_started().await;
wait_for_same_key_inflight_participants(&repository, TASKS).await;
inner.release.add_permits(TASKS);
for task in tasks {
let error = task
.await
.expect("request task should finish")
.expect_err("shared load should fail");
assert!(matches!(
error,
DataLayerError::TimedOut(message) if message == "shared routing load error"
));
}
assert_eq!(inner.calls(), 1);
}
#[tokio::test]
async fn cancelled_leader_wakes_follower_for_retry() {
let inner = Arc::new(ControlledRoutingGroupReadRepository::new(
ControlledHasAnyBindingBehavior::FirstPendingThenTrue,
));
let repository = Arc::new(CachedRoutingGroupReadRepository::new(inner.clone()));
let leader_repository = Arc::clone(&repository);
let leader =
tokio::spawn(async move { leader_repository.has_any_routing_group_binding().await });
inner.wait_until_started().await;
let follower_repository = Arc::clone(&repository);
let follower =
tokio::spawn(async move { follower_repository.has_any_routing_group_binding().await });
wait_for_same_key_inflight_participants(&repository, 2).await;
leader.abort();
let _ = leader.await;
assert!(tokio::time::timeout(Duration::from_secs(1), follower)
.await
.expect("follower should not remain stuck after leader cancellation")
.expect("follower task should finish")
.expect("retried load should succeed"));
assert_eq!(inner.calls(), 2);
}
#[tokio::test]
async fn clear_wakes_followers_and_rejects_old_generation_publication() {
let inner = Arc::new(ControlledRoutingGroupReadRepository::new(
ControlledHasAnyBindingBehavior::FirstWaitFalseThenTrue,
));
let repository = Arc::new(CachedRoutingGroupReadRepository::new(inner.clone()));
let leader_repository = Arc::clone(&repository);
let leader =
tokio::spawn(async move { leader_repository.has_any_routing_group_binding().await });
inner.wait_until_started().await;
let follower_repository = Arc::clone(&repository);
let follower =
tokio::spawn(async move { follower_repository.has_any_routing_group_binding().await });
wait_for_same_key_inflight_participants(&repository, 2).await;
repository.clear_local_cache();
assert!(tokio::time::timeout(Duration::from_secs(1), follower)
.await
.expect("clear should wake the follower")
.expect("follower task should finish")
.expect("new-generation load should succeed"));
assert_eq!(inner.calls(), 2);
inner.release.add_permits(1);
assert!(!leader
.await
.expect("old leader task should finish")
.expect("old leader load should succeed"));
assert!(repository
.has_any_routing_group_binding()
.await
.expect("new-generation value should remain cached"));
assert_eq!(inner.calls(), 2);
}
#[test]
fn capacity_full_cancelled_follower_can_retry_after_repeated_clear() {
let inner = Arc::new(CountingRoutingGroupReadRepository::default());
let repository = CachedRoutingGroupReadRepository::new(inner);
let key = RoutingGroupCacheKey::HasAnyBinding;
let mut active = Vec::with_capacity(ROUTING_GROUP_CACHE_MAX_INFLIGHT);
for _ in 0..ROUTING_GROUP_CACHE_MAX_INFLIGHT - 1 {
let leader = match repository.register_inflight(&key) {
RoutingGroupInflightRegistration::Leader(guard) => guard,
_ => panic!("each available permit should admit one leader"),
};
active.push(leader);
repository.clear();
}
let current = match repository.register_inflight(&key) {
RoutingGroupInflightRegistration::Leader(guard) => guard,
_ => panic!("the final available permit should admit a leader"),
};
let follower = match repository.register_inflight(&key) {
RoutingGroupInflightRegistration::Follower(state) => state,
_ => panic!("the same-key request should follow at full capacity"),
};
assert_eq!(repository.admission.available_permits(), 0);
assert!(matches!(
repository.register_inflight(&RoutingGroupCacheKey::ListGroups),
RoutingGroupInflightRegistration::Saturated
));
drop(current);
assert!(matches!(
repository.follower_completion(&follower),
Some(RoutingGroupInflightCompletion::Cancelled)
));
let mut replacement = match repository.register_inflight(&key) {
RoutingGroupInflightRegistration::Leader(guard) => guard,
_ => panic!("cancelled follower retry should use the released permit"),
};
replacement.finish(RoutingGroupInflightCompletion::Cancelled);
assert_eq!(repository.admission.available_permits(), 1);
}
}
+627 -154
View File
@@ -5,7 +5,8 @@ use super::{
AdminBillingRuleWriteInput, AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery,
AdminRedeemCodeListQuery, AdminWalletLedgerQuery, AdminWalletListQuery,
AdminWalletRefundRequestListQuery, AnnouncementListQuery, AuditLogListQuery,
BackgroundTaskListQuery, BackgroundTaskSummary, BillingModelContextCacheKey, BillingPlanRecord,
BackgroundTaskListQuery, BackgroundTaskSummary, BillingModelContextCacheKey,
BillingModelContextCacheState, BillingModelContextInflightState, BillingPlanRecord,
BillingPlanWriteInput, CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput,
CreateAdminRedeemCodeBatchResult, CreateAnnouncementRecord, CreateManualWalletRechargeInput,
CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput,
@@ -57,38 +58,93 @@ fn normalize_optional_billing_context_cache_part(value: Option<&str>) -> Option<
.map(ToOwned::to_owned)
}
enum BillingModelContextInflightRegistration {
Leader(u64),
Follower,
Bypass,
enum BillingModelContextInflightRegistration<'a> {
Leader(BillingModelContextInflightGuard<'a>),
Follower(std::sync::Arc<BillingModelContextInflightState>),
Saturated,
}
struct BillingModelContextInflightGuard<'a> {
state: &'a GatewayDataState,
key: Option<BillingModelContextCacheKey>,
token: u64,
inflight_state: std::sync::Arc<BillingModelContextInflightState>,
admission: Option<tokio::sync::OwnedSemaphorePermit>,
}
impl<'a> BillingModelContextInflightGuard<'a> {
fn new(state: &'a GatewayDataState, key: BillingModelContextCacheKey, token: u64) -> Self {
fn new(
state: &'a GatewayDataState,
key: BillingModelContextCacheKey,
inflight_state: std::sync::Arc<BillingModelContextInflightState>,
admission: tokio::sync::OwnedSemaphorePermit,
) -> Self {
Self {
state,
key: Some(key),
token,
inflight_state,
admission: Some(admission),
}
}
fn finish(&mut self) {
if let Some(key) = self.key.take() {
self.state
.finish_billing_model_context_inflight(&key, self.token);
fn epoch(&self) -> u64 {
self.inflight_state.epoch
}
fn finish(&mut self, error: Option<DataLayerError>) {
let removed = self.key.take().and_then(|key| {
self.state.finish_billing_model_context_inflight(
&key,
&self.inflight_state,
self.admission.take(),
)
});
self.admission.take();
if let Some(removed) = removed {
removed.complete(error.map_or(Ok(()), Err));
}
}
}
impl Drop for BillingModelContextInflightGuard<'_> {
fn drop(&mut self) {
self.finish();
self.finish(None);
}
}
impl BillingModelContextInflightState {
fn complete(&self, result: Result<(), DataLayerError>) {
if self.completion.set(result).is_ok() {
self.notify.notify_waiters();
}
}
async fn wait(&self) -> Result<(), DataLayerError> {
loop {
if let Some(result) = self.completion.get() {
return result.clone();
}
let mut notified = Box::pin(self.notify.notified());
notified.as_mut().enable();
if let Some(result) = self.completion.get() {
return result.clone();
}
notified.await;
}
}
}
impl Default for BillingModelContextCacheState {
fn default() -> Self {
Self {
entries: aether_cache::ExpiringMap::default(),
inflight: std::sync::Mutex::new(std::collections::HashMap::new()),
epoch: std::sync::atomic::AtomicU64::new(0),
mutation: std::sync::Mutex::new(()),
admission: std::sync::Arc::new(tokio::sync::Semaphore::new(
GatewayDataState::BILLING_MODEL_CONTEXT_CACHE_MAX_INFLIGHT,
)),
}
}
}
@@ -98,6 +154,7 @@ impl GatewayDataState {
const MAINTENANCE_POOL_PRESSURE_MAX_DEFER: Duration = Duration::from_secs(30);
const BILLING_MODEL_CONTEXT_CACHE_TTL: Duration = Duration::from_secs(30);
const BILLING_MODEL_CONTEXT_CACHE_MAX_ENTRIES: usize = 4096;
const BILLING_MODEL_CONTEXT_CACHE_MAX_INFLIGHT: usize = 4096;
#[cfg(not(test))]
const BILLING_MODEL_CONTEXT_CACHE_INFLIGHT_WAIT_TIMEOUT: Duration = Duration::from_secs(10);
#[cfg(test)]
@@ -1143,6 +1200,63 @@ impl GatewayDataState {
.await
}
pub(crate) async fn upsert_first_byte_usage(
&self,
usage: UpsertUsageRecord,
) -> Result<(), DataLayerError> {
crate::request_diagnostics::observe_db_operation(
"usage_first_byte_upsert",
self.database_pool_summary(),
async {
match &self.usage_writer {
Some(repository) => repository.upsert_first_byte(usage).await,
None => Ok(()),
}
},
)
.await
}
pub(crate) async fn upsert_first_byte_usage_many(
&self,
usages: Vec<UpsertUsageRecord>,
) -> Result<(), DataLayerError> {
if usages.is_empty() {
return Ok(());
}
crate::request_diagnostics::observe_db_operation(
"usage_first_byte_upsert_batch",
self.database_pool_summary(),
async {
match &self.usage_writer {
Some(repository) => repository.upsert_first_byte_many(usages).await,
None => Ok(()),
}
},
)
.await
}
pub(crate) async fn upsert_pending_usage_many(
&self,
usages: Vec<UpsertUsageRecord>,
) -> Result<(), DataLayerError> {
if usages.is_empty() {
return Ok(());
}
crate::request_diagnostics::observe_db_operation(
"usage_pending_upsert_batch",
self.database_pool_summary(),
async {
match &self.usage_writer {
Some(repository) => repository.upsert_pending_many(usages).await,
None => Ok(()),
}
},
)
.await
}
#[allow(dead_code)]
pub(crate) async fn rebuild_api_key_usage_stats(&self) -> Result<u64, DataLayerError> {
match &self.usage_writer {
@@ -1832,54 +1946,52 @@ impl GatewayDataState {
return Ok(value);
}
loop {
let notified = self.billing_model_context_cache.inflight_notify.notified();
match self.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Bypass => {
let load_epoch = self
.billing_model_context_cache
.epoch
.load(std::sync::atomic::Ordering::Acquire);
return self
.load_billing_model_context_by_name(
key,
provider_id,
provider_api_key_id,
global_model_name,
load_epoch,
)
.await;
BillingModelContextInflightRegistration::Saturated => {
return Err(DataLayerError::TimedOut(format!(
"billing model context cache admission saturated for {key:?}"
)));
}
BillingModelContextInflightRegistration::Follower => {
if timeout(
BillingModelContextInflightRegistration::Follower(inflight_state) => {
match timeout(
Self::BILLING_MODEL_CONTEXT_CACHE_INFLIGHT_WAIT_TIMEOUT,
notified,
inflight_state.wait(),
)
.await
.is_err()
{
self.expire_billing_model_context_inflight(&key);
Ok(Ok(())) => {}
Ok(Err(error)) => return Err(error),
Err(_) => self.expire_billing_model_context_inflight(&key, &inflight_state),
}
if let Some(value) = self.cached_billing_model_context(&key) {
return Ok(value);
}
continue;
}
BillingModelContextInflightRegistration::Leader(token) => {
let mut guard = BillingModelContextInflightGuard::new(self, key.clone(), token);
let load_epoch = self
.billing_model_context_cache
.epoch
.load(std::sync::atomic::Ordering::Acquire);
let result = self
.load_billing_model_context_by_name(
BillingModelContextInflightRegistration::Leader(mut guard) => {
if let Some(value) = self.cached_billing_model_context(&key) {
return Ok(value);
}
let load_epoch = guard.epoch();
let result = match timeout(
Self::BILLING_MODEL_CONTEXT_CACHE_INFLIGHT_WAIT_TIMEOUT,
self.load_billing_model_context_by_name(
key,
provider_id,
provider_api_key_id,
global_model_name,
load_epoch,
)
.await;
guard.finish();
&guard.inflight_state,
),
)
.await
{
Ok(result) => result,
Err(_) => Err(DataLayerError::TimedOut(
"billing model context load timed out".to_string(),
)),
};
guard.finish(result.as_ref().err().cloned());
return result;
}
}
@@ -1901,54 +2013,52 @@ impl GatewayDataState {
return Ok(value);
}
loop {
let notified = self.billing_model_context_cache.inflight_notify.notified();
match self.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Bypass => {
let load_epoch = self
.billing_model_context_cache
.epoch
.load(std::sync::atomic::Ordering::Acquire);
return self
.load_billing_model_context_by_model_id(
key,
provider_id,
provider_api_key_id,
model_id,
load_epoch,
)
.await;
BillingModelContextInflightRegistration::Saturated => {
return Err(DataLayerError::TimedOut(format!(
"billing model context cache admission saturated for {key:?}"
)));
}
BillingModelContextInflightRegistration::Follower => {
if timeout(
BillingModelContextInflightRegistration::Follower(inflight_state) => {
match timeout(
Self::BILLING_MODEL_CONTEXT_CACHE_INFLIGHT_WAIT_TIMEOUT,
notified,
inflight_state.wait(),
)
.await
.is_err()
{
self.expire_billing_model_context_inflight(&key);
Ok(Ok(())) => {}
Ok(Err(error)) => return Err(error),
Err(_) => self.expire_billing_model_context_inflight(&key, &inflight_state),
}
if let Some(value) = self.cached_billing_model_context(&key) {
return Ok(value);
}
continue;
}
BillingModelContextInflightRegistration::Leader(token) => {
let mut guard = BillingModelContextInflightGuard::new(self, key.clone(), token);
let load_epoch = self
.billing_model_context_cache
.epoch
.load(std::sync::atomic::Ordering::Acquire);
let result = self
.load_billing_model_context_by_model_id(
BillingModelContextInflightRegistration::Leader(mut guard) => {
if let Some(value) = self.cached_billing_model_context(&key) {
return Ok(value);
}
let load_epoch = guard.epoch();
let result = match timeout(
Self::BILLING_MODEL_CONTEXT_CACHE_INFLIGHT_WAIT_TIMEOUT,
self.load_billing_model_context_by_model_id(
key,
provider_id,
provider_api_key_id,
model_id,
load_epoch,
)
.await;
guard.finish();
&guard.inflight_state,
),
)
.await
{
Ok(result) => result,
Err(_) => Err(DataLayerError::TimedOut(
"billing model context load timed out".to_string(),
)),
};
guard.finish(result.as_ref().err().cloned());
return result;
}
}
@@ -1962,6 +2072,7 @@ impl GatewayDataState {
provider_api_key_id: Option<&str>,
global_model_name: &str,
load_epoch: u64,
load_flight: &std::sync::Arc<BillingModelContextInflightState>,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
crate::request_diagnostics::observe_db_operation(
"billing_model_context",
@@ -1972,11 +2083,16 @@ impl GatewayDataState {
let value = repository
.find_model_context(provider_id, provider_api_key_id, global_model_name)
.await?;
self.remember_billing_model_context(key, value.clone(), load_epoch);
self.remember_billing_model_context(
key,
value.clone(),
load_epoch,
load_flight,
);
Ok(value)
}
None => {
self.remember_billing_model_context(key, None, load_epoch);
self.remember_billing_model_context(key, None, load_epoch, load_flight);
Ok(None)
}
}
@@ -1992,6 +2108,7 @@ impl GatewayDataState {
provider_api_key_id: Option<&str>,
model_id: &str,
load_epoch: u64,
load_flight: &std::sync::Arc<BillingModelContextInflightState>,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
crate::request_diagnostics::observe_db_operation(
"billing_model_context",
@@ -2006,11 +2123,16 @@ impl GatewayDataState {
model_id,
)
.await?;
self.remember_billing_model_context(key, value.clone(), load_epoch);
self.remember_billing_model_context(
key,
value.clone(),
load_epoch,
load_flight,
);
Ok(value)
}
None => {
self.remember_billing_model_context(key, None, load_epoch);
self.remember_billing_model_context(key, None, load_epoch, load_flight);
Ok(None)
}
}
@@ -2022,44 +2144,91 @@ impl GatewayDataState {
fn register_billing_model_context_inflight(
&self,
key: &BillingModelContextCacheKey,
) -> BillingModelContextInflightRegistration {
match self.billing_model_context_cache.inflight.lock() {
Ok(mut inflight) => {
if inflight.contains_key(key) {
return BillingModelContextInflightRegistration::Follower;
}
let token = self
.billing_model_context_cache
.next_inflight_token
.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
inflight.insert(key.clone(), token);
BillingModelContextInflightRegistration::Leader(token)
}
Err(_) => BillingModelContextInflightRegistration::Bypass,
) -> BillingModelContextInflightRegistration<'_> {
let mut inflight = self
.billing_model_context_cache
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(inflight_state) = inflight.get(key) {
return BillingModelContextInflightRegistration::Follower(std::sync::Arc::clone(
inflight_state,
));
}
if inflight.len() >= Self::BILLING_MODEL_CONTEXT_CACHE_MAX_INFLIGHT {
return BillingModelContextInflightRegistration::Saturated;
}
let Ok(admission) =
std::sync::Arc::clone(&self.billing_model_context_cache.admission).try_acquire_owned()
else {
return BillingModelContextInflightRegistration::Saturated;
};
let inflight_state = std::sync::Arc::new(BillingModelContextInflightState {
epoch: self
.billing_model_context_cache
.epoch
.load(std::sync::atomic::Ordering::Acquire),
completion: std::sync::OnceLock::new(),
notify: tokio::sync::Notify::new(),
});
inflight.insert(key.clone(), std::sync::Arc::clone(&inflight_state));
BillingModelContextInflightRegistration::Leader(BillingModelContextInflightGuard::new(
self,
key.clone(),
inflight_state,
admission,
))
}
fn finish_billing_model_context_inflight(
&self,
key: &BillingModelContextCacheKey,
inflight_state: &std::sync::Arc<BillingModelContextInflightState>,
admission: Option<tokio::sync::OwnedSemaphorePermit>,
) -> Option<std::sync::Arc<BillingModelContextInflightState>> {
let mut inflight = self
.billing_model_context_cache
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
drop(admission);
if inflight
.get(key)
.is_some_and(|current| std::sync::Arc::ptr_eq(current, inflight_state))
{
inflight.remove(key)
} else {
None
}
}
fn finish_billing_model_context_inflight(&self, key: &BillingModelContextCacheKey, token: u64) {
let mut removed = false;
if let Ok(mut inflight) = self.billing_model_context_cache.inflight.lock() {
if inflight.get(key).copied() == Some(token) {
inflight.remove(key);
removed = true;
fn expire_billing_model_context_inflight(
&self,
key: &BillingModelContextCacheKey,
inflight_state: &std::sync::Arc<BillingModelContextInflightState>,
) {
let _mutation = self
.billing_model_context_cache
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let removed = {
let mut inflight = self
.billing_model_context_cache
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if inflight
.get(key)
.is_some_and(|current| std::sync::Arc::ptr_eq(current, inflight_state))
{
inflight.remove(key)
} else {
None
}
}
if removed {
self.billing_model_context_cache
.inflight_notify
.notify_waiters();
}
}
fn expire_billing_model_context_inflight(&self, key: &BillingModelContextCacheKey) {
let mut removed = false;
if let Ok(mut inflight) = self.billing_model_context_cache.inflight.lock() {
removed = inflight.remove(key).is_some();
}
if removed {
};
drop(_mutation);
if let Some(removed) = removed {
tracing::warn!(
event_name = "billing_model_context_cache_inflight_expired",
log_type = "ops",
@@ -2067,9 +2236,7 @@ impl GatewayDataState {
wait_timeout_ms = Self::BILLING_MODEL_CONTEXT_CACHE_INFLIGHT_WAIT_TIMEOUT.as_millis() as u64,
"gateway billing model context cache expired stale inflight load"
);
self.billing_model_context_cache
.inflight_notify
.notify_waiters();
removed.complete(Ok(()));
}
}
@@ -2079,13 +2246,7 @@ impl GatewayDataState {
) -> Option<Option<StoredBillingModelContext>> {
self.billing_model_context_cache
.entries
.read()
.expect("billing model context cache lock")
.get(key)
.and_then(|(cached_at, value)| {
(cached_at.elapsed() <= Self::BILLING_MODEL_CONTEXT_CACHE_TTL)
.then(|| value.clone())
})
.get_fresh(key, Self::BILLING_MODEL_CONTEXT_CACHE_TTL)
}
fn remember_billing_model_context(
@@ -2093,12 +2254,13 @@ impl GatewayDataState {
key: BillingModelContextCacheKey,
value: Option<StoredBillingModelContext>,
load_epoch: u64,
load_flight: &std::sync::Arc<BillingModelContextInflightState>,
) {
let mut cache = self
let _mutation = self
.billing_model_context_cache
.entries
.write()
.expect("billing model context cache lock");
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if load_epoch
!= self
.billing_model_context_cache
@@ -2107,44 +2269,53 @@ impl GatewayDataState {
{
return;
}
cache.retain(|_, (cached_at, _)| {
cached_at.elapsed() <= Self::BILLING_MODEL_CONTEXT_CACHE_TTL
});
if cache.len() >= Self::BILLING_MODEL_CONTEXT_CACHE_MAX_ENTRIES {
if let Some(oldest_key) = cache
.iter()
.min_by_key(|(_, (cached_at, _))| *cached_at)
.map(|(key, _)| key.clone())
{
cache.remove(&oldest_key);
}
let inflight = self
.billing_model_context_cache
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if !inflight
.get(&key)
.is_some_and(|current| std::sync::Arc::ptr_eq(current, load_flight))
{
return;
}
cache.insert(key, (Instant::now(), value));
self.billing_model_context_cache.entries.insert(
key,
value,
Self::BILLING_MODEL_CONTEXT_CACHE_TTL,
Self::BILLING_MODEL_CONTEXT_CACHE_MAX_ENTRIES,
);
}
pub(super) fn clear_billing_model_context_cache(&self) {
let _mutation = self
.billing_model_context_cache
.mutation
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
self.billing_model_context_cache
.epoch
.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
self.billing_model_context_cache
.entries
.write()
.expect("billing model context cache lock")
.clear();
let mut cleared_inflight = false;
if let Ok(mut inflight) = self.billing_model_context_cache.inflight.lock() {
cleared_inflight = !inflight.is_empty();
inflight.clear();
}
if cleared_inflight {
self.billing_model_context_cache.entries.clear();
let inflight_states = self
.billing_model_context_cache
.inflight
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.drain()
.map(|(_, state)| state)
.collect::<Vec<_>>();
drop(_mutation);
if !inflight_states.is_empty() {
tracing::warn!(
event_name = "billing_model_context_cache_inflight_cleared",
log_type = "ops",
"gateway billing model context cache cleared in-flight loads"
);
self.billing_model_context_cache
.inflight_notify
.notify_waiters();
for inflight_state in inflight_states {
inflight_state.complete(Ok(()));
}
}
}
@@ -2572,11 +2743,14 @@ mod tests {
use aether_data_contracts::repository::global_models::{
StoredAdminGlobalModel, StoredPublicGlobalModel, UpdateAdminGlobalModelRecord,
};
use aether_data_contracts::DataLayerError;
use async_trait::async_trait;
use serde_json::json;
use tokio::sync::Barrier;
use super::GatewayDataState;
use super::{
BillingModelContextCacheKey, BillingModelContextInflightRegistration, GatewayDataState,
};
struct SlowBillingContextRepository {
calls: AtomicUsize,
@@ -2590,6 +2764,13 @@ mod tests {
release_first_read: Barrier,
}
struct ConcurrentBillingContextRepository {
calls: AtomicUsize,
context: StoredBillingModelContext,
entered: Barrier,
release: Barrier,
}
#[async_trait]
impl BillingReadRepository for SlowBillingContextRepository {
async fn find_model_context(
@@ -2628,6 +2809,21 @@ mod tests {
}
}
#[async_trait]
impl BillingReadRepository for ConcurrentBillingContextRepository {
async fn find_model_context(
&self,
_provider_id: &str,
_provider_api_key_id: Option<&str>,
_global_model_name: &str,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
self.calls.fetch_add(1, Ordering::AcqRel);
self.entered.wait().await;
self.release.wait().await;
Ok(Some(self.context.clone()))
}
}
fn billing_context() -> StoredBillingModelContext {
StoredBillingModelContext::new(
"provider-1".to_string(),
@@ -2649,6 +2845,283 @@ mod tests {
.expect("billing context should build")
}
fn billing_cache_key(global_model_name: impl Into<String>) -> BillingModelContextCacheKey {
BillingModelContextCacheKey::ByGlobalModelName {
provider_id: "provider-1".to_string(),
provider_api_key_id: Some("key-1".to_string()),
global_model_name: global_model_name.into(),
}
}
#[tokio::test]
async fn billing_model_context_cancelled_leader_cannot_lose_follower_wakeup() {
let state = GatewayDataState::default();
let key = billing_cache_key("lost-wakeup");
let leader = match state.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
let follower = match state.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Follower(inflight_state) => inflight_state,
_ => panic!("second registration should follow"),
};
// Complete before wait() is constructed or polled. A bare global
// notify_waiters() broadcast loses this ordering.
drop(leader);
tokio::time::timeout(Duration::from_millis(100), follower.wait())
.await
.expect("cancelled flight must release an unpolled follower")
.expect("leader cancellation should allow a retry");
assert!(matches!(
state.register_billing_model_context_inflight(&key),
BillingModelContextInflightRegistration::Leader(_)
));
}
#[tokio::test]
async fn billing_model_context_failed_flight_fans_out_error() {
let state = GatewayDataState::default();
let key = billing_cache_key("failed-flight");
let mut leader = match state.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
let followers = (0..2)
.map(
|_| match state.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Follower(inflight_state) => {
inflight_state
}
_ => panic!("same-key registration should follow"),
},
)
.collect::<Vec<_>>();
leader.finish(Some(DataLayerError::Sql(
"forced billing context load failure".to_string(),
)));
for follower in followers {
let error = tokio::time::timeout(Duration::from_millis(100), follower.wait())
.await
.expect("failed flight should release every follower")
.expect_err("follower should receive the leader failure");
assert_eq!(
error.to_string(),
"sql error: forced billing context load failure"
);
}
assert!(state
.billing_model_context_cache
.inflight
.lock()
.unwrap()
.is_empty());
}
#[tokio::test]
async fn billing_model_context_clear_wakes_old_follower_without_removing_replacement() {
let state = GatewayDataState::default();
let key = billing_cache_key("clear-replacement");
let old_leader = match state.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
let old_follower = match state.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Follower(inflight_state) => inflight_state,
_ => panic!("second registration should follow"),
};
let old_epoch = old_leader.epoch();
state.clear_billing_model_context_cache();
let replacement_leader = match state.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Leader(guard) => guard,
_ => panic!("clear should allow a replacement leader"),
};
let replacement_follower = match state.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Follower(inflight_state) => inflight_state,
_ => panic!("registration behind replacement should follow"),
};
assert_ne!(replacement_leader.epoch(), old_epoch);
drop(old_leader);
assert!(state
.billing_model_context_cache
.inflight
.lock()
.unwrap()
.get(&key)
.is_some_and(|current| std::sync::Arc::ptr_eq(
current,
&replacement_leader.inflight_state
)));
tokio::time::timeout(Duration::from_millis(100), old_follower.wait())
.await
.expect("clear should wake the invalidated flight")
.expect("clear should allow an immediate retry");
drop(replacement_leader);
tokio::time::timeout(Duration::from_millis(100), replacement_follower.wait())
.await
.expect("old guard must not strand the replacement follower")
.expect("replacement completion should succeed");
}
#[tokio::test]
async fn billing_model_context_timeout_expiration_allows_replacement() {
let state = GatewayDataState::default();
let key = billing_cache_key("timeout-replacement");
let old_leader = match state.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
let old_follower = match state.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Follower(inflight_state) => inflight_state,
_ => panic!("second registration should follow"),
};
state.expire_billing_model_context_inflight(&key, &old_follower);
old_follower
.wait()
.await
.expect("expired flight should permit a retry");
let replacement_leader = match state.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Leader(guard) => guard,
_ => panic!("timeout should allow a replacement leader"),
};
drop(old_leader);
assert!(state
.billing_model_context_cache
.inflight
.lock()
.unwrap()
.get(&key)
.is_some_and(|current| std::sync::Arc::ptr_eq(
current,
&replacement_leader.inflight_state
)));
}
#[test]
fn billing_model_context_expired_leader_cannot_publish_over_replacement() {
let state = GatewayDataState::default();
let key = billing_cache_key("timeout-publication-replacement");
let old_leader = match state.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Leader(guard) => guard,
_ => panic!("first registration should lead"),
};
let old_flight = std::sync::Arc::clone(&old_leader.inflight_state);
let load_epoch = old_leader.epoch();
state.expire_billing_model_context_inflight(&key, &old_flight);
let replacement = match state.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Leader(guard) => guard,
_ => panic!("expiration should allow a replacement leader"),
};
assert_eq!(replacement.epoch(), load_epoch);
let mut fresh = billing_context();
fresh.default_price_per_request = Some(2.0);
state.remember_billing_model_context(
key.clone(),
Some(fresh),
replacement.epoch(),
&replacement.inflight_state,
);
let mut stale = billing_context();
stale.default_price_per_request = Some(1.0);
state.remember_billing_model_context(key.clone(), Some(stale), load_epoch, &old_flight);
let cached = state
.cached_billing_model_context(&key)
.expect("replacement should publish")
.expect("billing context should exist");
assert_eq!(cached.default_price_per_request, Some(2.0));
}
#[test]
fn billing_model_context_inflight_limit_rejects_only_new_keys() {
let state = GatewayDataState::default();
let mut leaders =
Vec::with_capacity(GatewayDataState::BILLING_MODEL_CONTEXT_CACHE_MAX_INFLIGHT);
for index in 0..GatewayDataState::BILLING_MODEL_CONTEXT_CACHE_MAX_INFLIGHT {
let key = billing_cache_key(format!("model-{index}"));
match state.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Leader(guard) => leaders.push(guard),
_ => panic!("unique key below the hard limit should lead"),
}
}
assert!(matches!(
state.register_billing_model_context_inflight(&billing_cache_key("overflow")),
BillingModelContextInflightRegistration::Saturated
));
assert!(matches!(
state.register_billing_model_context_inflight(&billing_cache_key("model-0")),
BillingModelContextInflightRegistration::Follower(_)
));
assert_eq!(
state
.billing_model_context_cache
.inflight
.lock()
.unwrap()
.len(),
GatewayDataState::BILLING_MODEL_CONTEXT_CACHE_MAX_INFLIGHT
);
drop(leaders);
assert!(state
.billing_model_context_cache
.inflight
.lock()
.unwrap()
.is_empty());
}
#[tokio::test]
async fn billing_model_context_different_keys_load_concurrently() {
let repository = Arc::new(ConcurrentBillingContextRepository {
calls: AtomicUsize::new(0),
context: billing_context(),
entered: Barrier::new(3),
release: Barrier::new(3),
});
let state = Arc::new(GatewayDataState::with_billing_reader_for_tests(
repository.clone(),
));
let task_a = {
let state = Arc::clone(&state);
tokio::spawn(async move {
state
.find_billing_model_context("provider-1", Some("key-1"), "model-a")
.await
})
};
let task_b = {
let state = Arc::clone(&state);
tokio::spawn(async move {
state
.find_billing_model_context("provider-1", Some("key-1"), "model-b")
.await
})
};
tokio::time::timeout(Duration::from_secs(1), repository.entered.wait())
.await
.expect("different cache keys should enter the repository concurrently");
assert_eq!(repository.calls.load(Ordering::Acquire), 2);
repository.release.wait().await;
task_a
.await
.expect("first lookup should join")
.expect("first lookup should succeed");
task_b
.await
.expect("second lookup should join")
.expect("second lookup should succeed");
}
#[tokio::test]
async fn billing_model_context_cache_coalesces_concurrent_loads() {
let repository = Arc::new(SlowBillingContextRepository {
@@ -17,10 +17,10 @@ use super::{
OAuthProviderWriteRepository, PoolMemberScoreWriteRepository, PoolScoreReadRepository,
ProviderCatalogReadRepository, ProviderCatalogWriteRepository, ProviderQuotaReadRepository,
ProviderQuotaWriteRepository, ProxyNodeReadRepository, ProxyNodeWriteRepository,
RequestCandidateReadRepository, RequestCandidateWriteRepository, SettlementWriteRepository,
StoredSystemConfigEntry, StoredUserPreferenceRecord, UsageReadRepository, UsageWriteRepository,
UserReadRepository, VideoTaskReadRepository, VideoTaskWriteRepository, WalletReadRepository,
WalletWriteRepository,
RequestCandidateReadRepository, RequestCandidateWriteRepository, RoutingGroupReadRepository,
RoutingGroupWriteRepository, SettlementWriteRepository, StoredSystemConfigEntry,
StoredUserPreferenceRecord, UsageReadRepository, UsageWriteRepository, UserReadRepository,
VideoTaskReadRepository, VideoTaskWriteRepository, WalletReadRepository, WalletWriteRepository,
};
mod announcements;
@@ -867,6 +867,16 @@ impl GatewayDataState {
self
}
#[cfg(test)]
pub(crate) fn with_routing_group_repository_for_tests<T>(mut self, repository: Arc<T>) -> Self
where
T: RoutingGroupReadRepository + RoutingGroupWriteRepository + 'static,
{
self.routing_group_reader = Some(repository.clone());
self.routing_group_writer = Some(repository);
self
}
#[cfg(test)]
pub(crate) fn with_auth_api_key_reader(
mut self,