mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 10:57:03 +08:00
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:
@@ -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());
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user