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

Shard and singleflight hot-path caches, batch and prioritize candidate and usage lifecycle persistence, and extend database and pressure-test instrumentation for 20k concurrent streams.
This commit is contained in:
elky
2026-07-22 02:11:08 +08:00
parent 7756c0913f
commit fc92c4f431
124 changed files with 36325 additions and 3217 deletions
+22 -3
View File
@@ -31,6 +31,7 @@ use super::{
AdminWalletPaymentOrderRecord, AdminWalletRefundRecord, AdminWalletTransactionRecord,
CachedProviderTransportSnapshot, FrontdoorCorsConfig, LocalExecutionRuntimeMissDiagnostic,
LocalProviderDeleteTaskState, ProviderTransportSnapshotCacheKey,
ProviderTransportSnapshotFlight,
};
const DEFAULT_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 120_000;
@@ -54,8 +55,8 @@ const DEFAULT_UPSTREAM_EXECUTION_GATE_LIMIT: usize = 10_000;
const DEFAULT_UPSTREAM_TARGET_GATE_LIMIT: usize = 10_000;
const MAX_AUTH_SNAPSHOT_LOAD_GATE_LIMIT: usize = 1024;
const MAX_CANDIDATE_PLANNING_GATE_LIMIT: usize = 8192;
const MAX_UPSTREAM_EXECUTION_GATE_LIMIT: usize = 16_384;
const MAX_UPSTREAM_TARGET_GATE_LIMIT: usize = 16_384;
const MAX_UPSTREAM_EXECUTION_GATE_LIMIT: usize = 32_768;
const MAX_UPSTREAM_TARGET_GATE_LIMIT: usize = 32_768;
const AUTH_SNAPSHOT_LOAD_GATE_LIMIT_PER_CPU: usize = 16;
const CANDIDATE_PLANNING_GATE_LIMIT_PER_CPU: usize = 256;
const UPSTREAM_EXECUTION_GATE_LIMIT_PER_CPU: usize = 1024;
@@ -408,6 +409,7 @@ pub struct AppState {
pub(crate) scheduler_affinity_epoch: Arc<AtomicU64>,
pub(crate) dashboard_response_cache: Arc<DashboardResponseCache>,
pub(crate) system_config_cache: Arc<SystemConfigCache>,
pub(crate) endpoint_response_header_rules_cache: Arc<JsonValueCache<String>>,
pub(crate) candidate_row_page_cache: Arc<super::super::cache::CandidateRowPageCache>,
pub(crate) candidate_page_cache: Arc<super::super::cache::CandidatePageCache>,
pub(crate) candidate_resolved_page_cache: Arc<super::super::cache::CandidateResolvedPageCache>,
@@ -430,8 +432,9 @@ pub struct AppState {
pub(crate) tunnel: crate::tunnel::EmbeddedTunnelState,
pub(crate) provider_transport_snapshot_cache:
Arc<DashMap<ProviderTransportSnapshotCacheKey, CachedProviderTransportSnapshot>>,
pub(crate) provider_transport_snapshot_cache_generation: Arc<AtomicU64>,
pub(crate) provider_transport_snapshot_inflight:
Arc<DashMap<ProviderTransportSnapshotCacheKey, Arc<TokioMutex<()>>>>,
Arc<DashMap<ProviderTransportSnapshotCacheKey, Arc<ProviderTransportSnapshotFlight>>>,
pub(crate) provider_key_rpm_resets: Arc<StdMutex<HashMap<String, u64>>>,
pub(crate) local_execution_runtime_miss_diagnostics:
Arc<DashMap<String, LocalExecutionRuntimeMissDiagnostic>>,
@@ -547,6 +550,22 @@ mod tests {
parse_gate_limit_value(Some("4096"), TEST_PROFILE, TEST_CAPACITY),
Some(4096)
);
assert_eq!(
parse_gate_limit_value(
Some("20000"),
UPSTREAM_TARGET_GATE_AUTO_PROFILE,
TEST_CAPACITY
),
Some(20_000)
);
assert_eq!(
parse_gate_limit_value(
Some("20000"),
UPSTREAM_EXECUTION_GATE_AUTO_PROFILE,
TEST_CAPACITY
),
Some(20_000)
);
}
#[test]
+98 -1
View File
@@ -1,16 +1,113 @@
use std::sync::Arc;
use std::sync::Mutex as StdMutex;
use std::time::Duration;
use super::super::error::GatewayError;
use super::super::provider_transport;
use tokio::sync::Notify;
#[derive(Debug, Clone)]
pub(crate) enum ProviderTransportSnapshotFlightResult {
Published(Arc<provider_transport::GatewayProviderTransportSnapshot>),
Missing,
Invalidated,
Retry,
Error(GatewayError),
}
/// One transport snapshot load per cache key. The completion value is kept
/// alongside the notification so a waiter that is scheduled after the
/// broadcast can still observe the result without a lost wakeup.
#[derive(Debug)]
pub(crate) struct ProviderTransportSnapshotFlight {
generation: u64,
notify: Arc<Notify>,
result: StdMutex<Option<ProviderTransportSnapshotFlightResult>>,
}
impl ProviderTransportSnapshotFlight {
pub(crate) fn new(generation: u64) -> Self {
Self {
generation,
notify: Arc::new(Notify::new()),
result: StdMutex::new(None),
}
}
pub(crate) fn generation(&self) -> u64 {
self.generation
}
fn result(&self) -> Option<ProviderTransportSnapshotFlightResult> {
self.result
.lock()
.map(|result| result.clone())
.unwrap_or_else(|poisoned| poisoned.into_inner().clone())
}
/// Completes the flight once. A clear/cancellation may win the race with
/// the leader; in that case the old result must not overwrite Invalidated.
pub(crate) fn complete(&self, result: ProviderTransportSnapshotFlightResult) -> bool {
let completed = match self.result.lock() {
Ok(mut current) => {
if current.is_some() {
false
} else {
*current = Some(result);
true
}
}
Err(poisoned) => {
let mut current = poisoned.into_inner();
if current.is_some() {
false
} else {
*current = Some(result);
true
}
}
};
if completed {
self.notify.notify_waiters();
}
completed
}
pub(crate) async fn wait(&self) -> ProviderTransportSnapshotFlightResult {
loop {
if let Some(result) = self.result() {
return result;
}
// Register before checking the completion state a second time.
// This closes the result-check/notify race even when the task has
// not been polled yet when the leader broadcasts completion.
let mut notified = Box::pin(self.notify.notified());
notified.as_mut().enable();
if let Some(result) = self.result() {
return result;
}
notified.await;
}
}
}
pub(crate) const AUTH_API_KEY_LAST_USED_TTL: Duration = Duration::from_secs(60);
pub(crate) const AUTH_API_KEY_LAST_USED_MAX_ENTRIES: usize = 10_000;
// Keep normal freshness short for cross-node configuration propagation. A
// stale entry is served while one background refresh runs, so an idle burst
// does not turn expiry into a synchronous database wait.
pub(crate) const PROVIDER_TRANSPORT_SNAPSHOT_CACHE_TTL: Duration = Duration::from_secs(1);
pub(crate) const PROVIDER_TRANSPORT_SNAPSHOT_CACHE_STALE_TTL: Duration = Duration::from_secs(30);
// Keep the last known transport usable across normal idle periods. The first
// stale request starts a single background refresh, and every local catalog or
// credential mutation advances the generation and clears the entry.
pub(crate) const PROVIDER_TRANSPORT_SNAPSHOT_CACHE_STALE_TTL: Duration =
Duration::from_secs(5 * 60);
pub(crate) const PROVIDER_TRANSPORT_SNAPSHOT_CACHE_MAX_ENTRIES: usize = 1_024;
#[derive(Debug, Clone)]
pub(crate) struct CachedProviderTransportSnapshot {
pub(crate) loaded_at: std::time::Instant,
pub(crate) generation: u64,
pub(crate) snapshot: Arc<provider_transport::GatewayProviderTransportSnapshot>,
}
+200 -26
View File
@@ -677,40 +677,78 @@ impl AppState {
Ok(updated)
}
pub(crate) async fn update_provider_catalog_key_runtime_state(
pub(crate) async fn compare_and_update_provider_catalog_key_adaptive_state(
&self,
key: &provider_catalog::StoredProviderCatalogKey,
) -> Result<Option<provider_catalog::StoredProviderCatalogKey>, GatewayError> {
update: &provider_catalog::ProviderCatalogKeyAdaptiveStateUpdate,
) -> Result<bool, GatewayError> {
let updated = self
.data
.update_provider_catalog_key(key)
.compare_and_update_provider_catalog_key_adaptive_state(update)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated.is_some() {
// A CAS conflict means a remote writer changed state, so local runtime reads
// must be refreshed even though this instance did not update the row.
self.invalidate_provider_runtime_state_caches();
Ok(updated)
}
pub(crate) async fn update_provider_catalog_key_runtime_metadata(
&self,
update: &provider_catalog::ProviderCatalogKeyRuntimeMetadataUpdate,
) -> Result<bool, GatewayError> {
let updated = self
.data
.update_provider_catalog_key_runtime_metadata(update)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
// A false result is a namespace CAS conflict. Invalidate the runtime
// snapshots before the caller reloads and retries. Upstream metadata is
// part of the transport snapshot, unlike health/adaptive state.
self.invalidate_provider_transport_runtime_state_caches();
Ok(updated)
}
pub(crate) async fn update_provider_catalog_key_status_snapshot(
&self,
update: &provider_catalog::ProviderCatalogKeyStatusSnapshotUpdate,
) -> Result<bool, GatewayError> {
let updated = self
.data
.update_provider_catalog_key_status_snapshot(update)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated {
self.invalidate_provider_runtime_state_caches();
}
Ok(updated)
}
pub(crate) async fn update_provider_catalog_key_success_health_state(
pub(crate) async fn compare_and_update_provider_catalog_key_health_state(
&self,
key_id: &str,
is_active: bool,
health_by_format: Option<&serde_json::Value>,
circuit_breaker_by_format: Option<&serde_json::Value>,
update: &provider_catalog::ProviderCatalogKeyHealthStateUpdate,
) -> Result<bool, GatewayError> {
let updated = self
.data
.update_provider_catalog_key_health_state(
key_id,
is_active,
health_by_format,
circuit_breaker_by_format,
)
.compare_and_update_provider_catalog_key_health_state(update)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
// On conflict another Gateway changed the health snapshot; invalidate all
// health-sensitive caches before the retry reads it back.
self.invalidate_provider_health_routing_caches();
Ok(updated)
}
pub(crate) async fn reset_provider_catalog_key_error_count(
&self,
key_id: &str,
) -> Result<bool, GatewayError> {
let updated = self
.data
.reset_provider_catalog_key_error_count(key_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated {
self.invalidate_provider_runtime_state_caches();
self.invalidate_provider_health_routing_caches();
}
Ok(updated)
}
@@ -874,9 +912,34 @@ impl AppState {
.ok()
.map(|duration| duration.as_secs());
self.update_provider_catalog_key(&key)
let cleared = self
.data
.clear_provider_catalog_key_oauth_invalid_marker(key_id)
.await
.map(|updated| updated.is_some())
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if !cleared {
return Ok(false);
}
// The marker write already committed. Invalidate before any follow-up
// status patch so error/false paths cannot retain an invalid transport.
self.invalidate_provider_transport_runtime_state_caches();
let oauth = key
.status_snapshot
.as_ref()
.and_then(serde_json::Value::as_object)
.and_then(|snapshot| snapshot.get("oauth"))
.cloned()
.unwrap_or(serde_json::Value::Null);
let updated = self
.update_provider_catalog_key_status_snapshot(
&provider_catalog::ProviderCatalogKeyStatusSnapshotUpdate {
key_id: key_id.to_string(),
status_snapshot_patch: serde_json::json!({"oauth":oauth}),
updated_at_unix_secs: key.updated_at_unix_secs,
},
)
.await?;
Ok(updated)
}
pub(crate) fn put_provider_delete_task(&self, task: LocalProviderDeleteTaskState) {
@@ -998,7 +1061,10 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated {
self.invalidate_provider_health_routing_caches();
// This administrator-facing API also writes `is_active`, which is
// part of the transport snapshot. Runtime health CAS updates use the
// separate compare-and-update API above and keep transport cached.
self.invalidate_provider_routing_caches();
}
Ok(updated)
}
@@ -1360,7 +1426,7 @@ mod tests {
}
#[tokio::test]
async fn provider_catalog_health_update_keeps_scheduler_affinity_cache() {
async fn provider_catalog_runtime_health_update_keeps_scheduler_affinity_and_transport_cache() {
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider()],
vec![sample_endpoint()],
@@ -1373,6 +1439,12 @@ mod tests {
.with_encryption_key_for_tests("test-encryption-key"),
);
let transport_before = state
.read_provider_transport_snapshot_arc("provider-1", "endpoint-1", "key-1")
.await
.expect("provider transport should read")
.expect("provider transport should exist");
let cache_key = "scheduler_affinity:api-key-1:openai:chat:gpt-5";
let ttl = Duration::from_secs(300);
let target = SchedulerAffinityTarget {
@@ -1390,7 +1462,15 @@ mod tests {
}
});
let updated = state
.update_provider_catalog_key_health_state("key-1", true, Some(&health_by_format), None)
.compare_and_update_provider_catalog_key_health_state(
&aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyHealthStateUpdate {
key_id: "key-1".to_string(),
expected_health_by_format: None,
expected_circuit_breaker_by_format: None,
health_by_format: Some(health_by_format),
circuit_breaker_by_format: None,
},
)
.await
.expect("key health update should succeed");
@@ -1400,6 +1480,96 @@ mod tests {
state.read_scheduler_affinity_target(cache_key, ttl),
Some(target)
);
let transport_after = state
.read_provider_transport_snapshot_arc("provider-1", "endpoint-1", "key-1")
.await
.expect("provider transport should read after health update")
.expect("provider transport should still exist");
assert!(
Arc::ptr_eq(&transport_before, &transport_after),
"health-only writes must not invalidate transport configuration"
);
}
#[tokio::test]
async fn provider_catalog_admin_health_update_invalidates_transport_when_active_changes() {
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider()],
vec![sample_endpoint()],
vec![sample_key()],
));
let state = AppState::new()
.expect("app state should build")
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(repository)
.with_encryption_key_for_tests("test-encryption-key"),
);
let transport_before = state
.read_provider_transport_snapshot_arc("provider-1", "endpoint-1", "key-1")
.await
.expect("provider transport should read")
.expect("provider transport should exist");
assert!(transport_before.key.is_active);
assert!(state
.update_provider_catalog_key_health_state("key-1", false, None, None)
.await
.expect("administrator health update should succeed"));
let transport_after = state
.read_provider_transport_snapshot_arc("provider-1", "endpoint-1", "key-1")
.await
.expect("provider transport should read after active update")
.expect("provider transport should still exist");
assert!(!transport_after.key.is_active);
assert!(!Arc::ptr_eq(&transport_before, &transport_after));
}
#[tokio::test]
async fn clearing_oauth_invalid_marker_invalidates_transport_snapshot() {
let mut key = sample_key();
key.auth_type = "oauth".to_string();
key.oauth_invalid_at_unix_secs = Some(1_700_000_000);
key.oauth_invalid_reason = Some("[REFRESH_FAILED] stale token".to_string());
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider()],
vec![sample_endpoint()],
vec![key],
));
let state = AppState::new()
.expect("app state should build")
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(repository)
.with_encryption_key_for_tests("test-encryption-key"),
);
let transport_before = state
.read_provider_transport_snapshot_arc("provider-1", "endpoint-1", "key-1")
.await
.expect("provider transport should read")
.expect("provider transport should exist");
assert!(state
.clear_provider_catalog_key_oauth_invalid_marker("key-1")
.await
.expect("OAuth invalid marker should clear"));
let transport_after = state
.read_provider_transport_snapshot_arc("provider-1", "endpoint-1", "key-1")
.await
.expect("provider transport should reload")
.expect("provider transport should exist");
let persisted = state
.read_provider_catalog_keys_by_ids(&["key-1".to_string()])
.await
.expect("provider key should reload")
.into_iter()
.next()
.expect("provider key should exist");
assert!(persisted.oauth_invalid_at_unix_secs.is_none());
assert!(persisted.oauth_invalid_reason.is_none());
assert!(!Arc::ptr_eq(&transport_before, &transport_after));
}
#[tokio::test]
@@ -1442,14 +1612,18 @@ mod tests {
);
assert!(state.candidate_page_cache.get(&cache_key, ttl).is_some());
let mut updated_key = sample_key();
updated_key.status_snapshot = Some(serde_json::json!({"source": "runtime"}));
let updated = state
.update_provider_catalog_key_runtime_state(&updated_key)
.update_provider_catalog_key_status_snapshot(
&aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyStatusSnapshotUpdate {
key_id: "key-1".to_string(),
status_snapshot_patch: serde_json::json!({"source": "runtime"}),
updated_at_unix_secs: None,
},
)
.await
.expect("runtime state update should succeed");
assert!(updated.is_some());
assert!(updated);
assert!(state.candidate_page_cache.get(&cache_key, ttl).is_some());
}
}
File diff suppressed because it is too large Load Diff
@@ -485,6 +485,20 @@ impl RequestCandidateRuntimeWriter for AppState {
) -> Result<Option<StoredRequestCandidate>, GatewayError> {
AppState::upsert_request_candidate(self, candidate).await
}
async fn enqueue_request_candidate_status(
&self,
candidate: UpsertRequestCandidateRecord,
) -> Result<Option<()>, GatewayError> {
AppState::enqueue_request_candidate_status(self, candidate).await
}
fn try_enqueue_request_candidate_status(
&self,
candidate: UpsertRequestCandidateRecord,
) -> Result<(), UpsertRequestCandidateRecord> {
AppState::try_enqueue_request_candidate_status(self, candidate)
}
}
#[async_trait]
+2 -1
View File
@@ -32,7 +32,8 @@ pub(crate) use self::app::{
FrontdoorRuntimeGuardConfig, REQUEST_BODY_BUFFER_PERMIT_BYTES,
};
pub(crate) use self::cache::{
CachedProviderTransportSnapshot, AUTH_API_KEY_LAST_USED_MAX_ENTRIES,
CachedProviderTransportSnapshot, ProviderTransportSnapshotFlight,
ProviderTransportSnapshotFlightResult, AUTH_API_KEY_LAST_USED_MAX_ENTRIES,
AUTH_API_KEY_LAST_USED_TTL, PROVIDER_TRANSPORT_SNAPSHOT_CACHE_MAX_ENTRIES,
PROVIDER_TRANSPORT_SNAPSHOT_CACHE_STALE_TTL, PROVIDER_TRANSPORT_SNAPSHOT_CACHE_TTL,
};
File diff suppressed because it is too large Load Diff
@@ -59,6 +59,7 @@ impl AppState {
if cache_key.is_empty() {
return Ok(None);
}
let cache_generation = self.auth_snapshot_cache.generation();
let snapshot = self
.auth_snapshot_cache
.get_or_load(
@@ -74,10 +75,11 @@ impl AppState {
)
.await?;
if let Some(snapshot) = snapshot.as_ref() {
self.auth_snapshot_cache.insert(
self.auth_snapshot_cache.insert_if_generation(
AuthSnapshotCacheKey::user_api_key_ids(&snapshot.user_id, &snapshot.api_key_id),
Some(snapshot.clone()),
AUTH_API_KEY_SNAPSHOT_RUNTIME_CACHE_TTL,
cache_generation,
);
}
Ok(snapshot)
@@ -4,6 +4,8 @@ use crate::constants::{BUILTIN_DEFAULT_USER_GROUP_ID, DEFAULT_USER_GROUP_CONFIG_
use crate::{AppState, GatewayError};
use std::time::Duration;
// Membership changes should propagate promptly across gateway instances. The
// routing fast path skips this lookup entirely when no bindings exist.
const USER_GROUPS_FOR_USER_CACHE_TTL: Duration = Duration::from_secs(30);
impl AppState {
@@ -20,6 +20,16 @@ fn local_mutation_outcome<T>(outcome: AdminBillingMutationOutcome<T>) -> LocalMu
}
impl AppState {
fn finish_billing_model_context_mutation<T>(
&self,
outcome: LocalMutationOutcome<T>,
) -> LocalMutationOutcome<T> {
if matches!(&outcome, LocalMutationOutcome::Applied(_)) {
self.auth_request_cost_upper_bound_cache.clear();
}
outcome
}
pub(crate) async fn admin_billing_enabled_default_value_exists(
&self,
api_format: &str,
@@ -81,14 +91,18 @@ impl AppState {
.lock()
.expect("admin billing rule store should lock")
.insert(record.id.clone(), record.clone());
return Ok(LocalMutationOutcome::Applied(record));
return Ok(
self.finish_billing_model_context_mutation(LocalMutationOutcome::Applied(record))
);
}
self.data
let outcome = self
.data
.create_admin_billing_rule(input)
.await
.map(local_mutation_outcome)
.map_err(data_error)
.map_err(data_error)?;
Ok(self.finish_billing_model_context_mutation(outcome))
}
pub(crate) async fn list_admin_billing_rules(
@@ -171,14 +185,20 @@ impl AppState {
record.dimension_mappings = input.dimension_mappings.clone();
record.is_enabled = input.is_enabled;
record.updated_at_unix_secs = chrono::Utc::now().timestamp().max(0) as u64;
return Ok(LocalMutationOutcome::Applied(record.clone()));
return Ok(
self.finish_billing_model_context_mutation(LocalMutationOutcome::Applied(
record.clone(),
)),
);
}
self.data
let outcome = self
.data
.update_admin_billing_rule(rule_id, input)
.await
.map(local_mutation_outcome)
.map_err(data_error)
.map_err(data_error)?;
Ok(self.finish_billing_model_context_mutation(outcome))
}
pub(crate) async fn create_admin_billing_collector(
@@ -207,14 +227,18 @@ impl AppState {
.lock()
.expect("admin billing collector store should lock")
.insert(record.id.clone(), record.clone());
return Ok(LocalMutationOutcome::Applied(record));
return Ok(
self.finish_billing_model_context_mutation(LocalMutationOutcome::Applied(record))
);
}
self.data
let outcome = self
.data
.create_admin_billing_collector(input)
.await
.map(local_mutation_outcome)
.map_err(data_error)
.map_err(data_error)?;
Ok(self.finish_billing_model_context_mutation(outcome))
}
pub(crate) async fn list_admin_billing_collectors(
@@ -313,14 +337,20 @@ impl AppState {
record.priority = input.priority;
record.is_enabled = input.is_enabled;
record.updated_at_unix_secs = chrono::Utc::now().timestamp().max(0) as u64;
return Ok(LocalMutationOutcome::Applied(record.clone()));
return Ok(
self.finish_billing_model_context_mutation(LocalMutationOutcome::Applied(
record.clone(),
)),
);
}
self.data
let outcome = self
.data
.update_admin_billing_collector(collector_id, input)
.await
.map(local_mutation_outcome)
.map_err(data_error)
.map_err(data_error)?;
Ok(self.finish_billing_model_context_mutation(outcome))
}
pub(crate) async fn apply_admin_billing_preset(
@@ -389,23 +419,27 @@ impl AppState {
}
}
}
return Ok(LocalMutationOutcome::Applied(
AdminBillingPresetApplyResult {
preset: preset.to_string(),
mode: mode.to_string(),
created,
updated,
skipped,
errors: Vec::new(),
},
));
return Ok(
self.finish_billing_model_context_mutation(LocalMutationOutcome::Applied(
AdminBillingPresetApplyResult {
preset: preset.to_string(),
mode: mode.to_string(),
created,
updated,
skipped,
errors: Vec::new(),
},
)),
);
}
self.data
let outcome = self
.data
.apply_admin_billing_preset(preset, mode, collectors)
.await
.map(local_mutation_outcome)
.map_err(data_error)
.map_err(data_error)?;
Ok(self.finish_billing_model_context_mutation(outcome))
}
pub(crate) async fn find_payment_gateway_config(
@@ -549,3 +583,135 @@ impl AppState {
self.find_user_daily_quota_availability(user_id).await
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use serde_json::json;
use super::{
AdminBillingCollectorWriteInput, AdminBillingRuleWriteInput, AppState, LocalMutationOutcome,
};
const CACHE_KEY: &str = "billing-mutation-test";
fn rule_input() -> AdminBillingRuleWriteInput {
AdminBillingRuleWriteInput {
name: "chat-input".to_string(),
task_type: "chat".to_string(),
global_model_id: Some("global-model".to_string()),
model_id: Some("model-1".to_string()),
expression: "input_tokens * 0.01".to_string(),
variables: json!({"base": 1}),
dimension_mappings: json!({"input_tokens": "input_tokens"}),
is_enabled: true,
}
}
fn collector_input(dimension_name: &str) -> AdminBillingCollectorWriteInput {
AdminBillingCollectorWriteInput {
api_format: "openai".to_string(),
task_type: "chat".to_string(),
dimension_name: dimension_name.to_string(),
source_type: "request".to_string(),
source_path: Some("usage.input_tokens".to_string()),
value_type: "float".to_string(),
transform_expression: None,
default_value: None,
priority: 10,
is_enabled: true,
}
}
fn prime_cache(state: &AppState) {
state.auth_request_cost_upper_bound_cache.insert(
CACHE_KEY.to_string(),
Some(1.0),
Duration::from_secs(60),
);
assert!(state
.auth_request_cost_upper_bound_cache
.get(&CACHE_KEY.to_string(), Duration::from_secs(60),)
.is_some());
}
fn assert_cache_cleared(state: &AppState) {
assert_eq!(
state
.auth_request_cost_upper_bound_cache
.get(&CACHE_KEY.to_string(), Duration::from_secs(60),),
None
);
}
#[tokio::test]
async fn applied_admin_billing_mutations_clear_auth_cost_cache() {
let state = AppState::new().expect("app state should build");
prime_cache(&state);
let created_rule = state
.create_admin_billing_rule(&rule_input())
.await
.expect("rule create should succeed");
let rule_id = match created_rule {
LocalMutationOutcome::Applied(record) => record.id,
other => panic!("expected applied rule create, got {other:?}"),
};
assert_cache_cleared(&state);
prime_cache(&state);
let updated_rule = state
.update_admin_billing_rule(&rule_id, &rule_input())
.await
.expect("rule update should succeed");
assert!(matches!(updated_rule, LocalMutationOutcome::Applied(_)));
assert_cache_cleared(&state);
prime_cache(&state);
let created_collector = state
.create_admin_billing_collector(&collector_input("latency"))
.await
.expect("collector create should succeed");
let collector_id = match created_collector {
LocalMutationOutcome::Applied(record) => record.id,
other => panic!("expected applied collector create, got {other:?}"),
};
assert_cache_cleared(&state);
prime_cache(&state);
let updated_collector = state
.update_admin_billing_collector(&collector_id, &collector_input("latency"))
.await
.expect("collector update should succeed");
assert!(matches!(
updated_collector,
LocalMutationOutcome::Applied(_)
));
assert_cache_cleared(&state);
prime_cache(&state);
let applied_preset = state
.apply_admin_billing_preset("test", "merge", &[collector_input("images")])
.await
.expect("preset apply should succeed");
assert!(matches!(applied_preset, LocalMutationOutcome::Applied(_)));
assert_cache_cleared(&state);
}
#[tokio::test]
async fn non_applied_admin_billing_mutation_keeps_auth_cost_cache() {
let state = AppState::new().expect("app state should build");
prime_cache(&state);
let outcome = state
.update_admin_billing_rule("missing", &rule_input())
.await
.expect("missing rule update should return a local outcome");
assert_eq!(outcome, LocalMutationOutcome::NotFound);
assert!(state
.auth_request_cost_upper_bound_cache
.get(&CACHE_KEY.to_string(), Duration::from_secs(60),)
.is_some());
}
}
@@ -127,6 +127,44 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
/// Persist a candidate status when the caller does not need the materialized row.
///
/// Lifecycle updates are emitted on the hot path (in particular the first-byte
/// `pending -> streaming` transition). Rebuilding `StoredRequestCandidate` here
/// only to discard it adds validation and clones for every update, especially when
/// the async queue is enabled.
pub(crate) async fn enqueue_request_candidate_status(
&self,
candidate: candidates::UpsertRequestCandidateRecord,
) -> Result<Option<()>, GatewayError> {
if let Some(queue) = self.request_candidate_queue.as_ref() {
queue
.enqueue_or_fallback(candidate)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
return Ok(Some(()));
}
self.data
.upsert_request_candidate(candidate)
.await
.map(|stored| stored.map(|_| ()))
.map_err(|err| GatewayError::Internal(err.to_string()))
}
/// Try the in-memory lifecycle lane without awaiting or touching the repository.
/// The returned record must be persisted through `enqueue_request_candidate_status`
/// when the queue is disabled or closed.
pub(crate) fn try_enqueue_request_candidate_status(
&self,
candidate: candidates::UpsertRequestCandidateRecord,
) -> Result<(), candidates::UpsertRequestCandidateRecord> {
let Some(queue) = self.request_candidate_queue.as_ref() else {
return Err(candidate);
};
queue.try_enqueue_priority_status(candidate)
}
}
fn stored_request_candidate_from_upsert(