mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-10 11:19:50 +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:
@@ -1161,36 +1161,54 @@ async fn resolve_priority_candidate_page_with_cache(
|
||||
.resolved_page_cache_model_directive_policy_hash(),
|
||||
cursor.resolution_mode,
|
||||
);
|
||||
let page_candidates_for_fallback = page_candidates.clone();
|
||||
let page_candidates_for_load = page_candidates;
|
||||
let cache = cursor.state.app().candidate_resolved_page_cache.clone();
|
||||
let page_candidates_for_fallback = page_candidates;
|
||||
let app = cursor.state.app();
|
||||
let cache = app.candidate_resolved_page_cache.clone();
|
||||
let ttl = candidate_page_cache_ttl_from_env();
|
||||
let stale_ttl = candidate_page_cache_stale_ttl(ttl);
|
||||
// Keep request-owned planner inputs borrowed until the cache tells us a
|
||||
// cold load or stale refresh is actually needed. Fresh hits must not pay
|
||||
// for deep copies of candidate pages and auth/routing snapshots.
|
||||
let client_api_format = cursor.client_api_format.as_str();
|
||||
let requested_model = cursor.requested_model.as_str();
|
||||
let auth_snapshot = &cursor.auth_snapshot;
|
||||
let client_session_affinity = cursor.client_session_affinity.as_ref();
|
||||
let required_capabilities = cursor.required_capabilities.as_ref();
|
||||
let routing_policy = cursor.routing_policy.as_ref();
|
||||
let request_auth_channel = cursor.request_auth_channel.as_deref();
|
||||
let resolution_mode = cursor.resolution_mode;
|
||||
let cached = cache
|
||||
.get_or_load_once_stale_while_refreshing(
|
||||
.get_or_load_once_stale_while_revalidating(
|
||||
key,
|
||||
ttl,
|
||||
stale_ttl,
|
||||
|| async move {
|
||||
let (candidates, resolved_skipped) =
|
||||
resolve_and_rank_logical_local_execution_candidates(
|
||||
cursor.state,
|
||||
page_candidates_for_load,
|
||||
&cursor.client_api_format,
|
||||
Some(&cursor.requested_model),
|
||||
Some(&cursor.auth_snapshot),
|
||||
cursor.client_session_affinity.as_ref(),
|
||||
cursor.required_capabilities.as_ref(),
|
||||
cursor.routing_policy.as_ref(),
|
||||
cursor.sticky_session_token.as_deref(),
|
||||
cursor.request_auth_channel.as_deref(),
|
||||
cursor.resolution_mode,
|
||||
)
|
||||
.await;
|
||||
Ok::<_, GatewayError>(Some(Arc::new(CandidateResolvedPageSnapshot {
|
||||
candidates,
|
||||
resolved_skipped,
|
||||
})))
|
||||
|| {
|
||||
resolve_candidate_page_snapshot(
|
||||
(*app).clone(),
|
||||
page_candidates_for_fallback.clone(),
|
||||
client_api_format.to_owned(),
|
||||
requested_model.to_owned(),
|
||||
auth_snapshot.clone(),
|
||||
client_session_affinity.cloned(),
|
||||
required_capabilities.cloned(),
|
||||
routing_policy.cloned(),
|
||||
request_auth_channel.map(ToOwned::to_owned),
|
||||
resolution_mode,
|
||||
)
|
||||
},
|
||||
|| {
|
||||
resolve_candidate_page_snapshot(
|
||||
(*app).clone(),
|
||||
page_candidates_for_fallback.clone(),
|
||||
client_api_format.to_owned(),
|
||||
requested_model.to_owned(),
|
||||
auth_snapshot.clone(),
|
||||
client_session_affinity.cloned(),
|
||||
required_capabilities.cloned(),
|
||||
routing_policy.cloned(),
|
||||
request_auth_channel.map(ToOwned::to_owned),
|
||||
resolution_mode,
|
||||
)
|
||||
},
|
||||
CacheLoadObserver::new()
|
||||
.on_hit(record_candidate_page_resolve_cache_hit)
|
||||
@@ -1228,6 +1246,39 @@ async fn resolve_priority_candidate_page_with_cache(
|
||||
}
|
||||
}
|
||||
|
||||
async fn resolve_candidate_page_snapshot(
|
||||
app: AppState,
|
||||
page_candidates: Vec<SchedulerMinimalCandidateSelectionCandidate>,
|
||||
client_api_format: String,
|
||||
requested_model: String,
|
||||
auth_snapshot: GatewayAuthApiKeySnapshot,
|
||||
client_session_affinity: Option<ClientSessionAffinity>,
|
||||
required_capabilities: Option<Value>,
|
||||
routing_policy: Option<ResolvedRoutingPolicy>,
|
||||
request_auth_channel: Option<String>,
|
||||
resolution_mode: LocalCandidateResolutionMode,
|
||||
) -> Result<Option<Arc<CandidateResolvedPageSnapshot>>, GatewayError> {
|
||||
let state = PlannerAppState::new(&app);
|
||||
let (candidates, resolved_skipped) = resolve_and_rank_logical_local_execution_candidates(
|
||||
state,
|
||||
page_candidates,
|
||||
&client_api_format,
|
||||
Some(&requested_model),
|
||||
Some(&auth_snapshot),
|
||||
client_session_affinity.as_ref(),
|
||||
required_capabilities.as_ref(),
|
||||
routing_policy.as_ref(),
|
||||
None,
|
||||
request_auth_channel.as_deref(),
|
||||
resolution_mode,
|
||||
)
|
||||
.await;
|
||||
Ok(Some(Arc::new(CandidateResolvedPageSnapshot {
|
||||
candidates,
|
||||
resolved_skipped,
|
||||
})))
|
||||
}
|
||||
|
||||
fn should_cache_resolved_candidate_page(cursor: &RequestedModelAttemptPageCursor<'_>) -> bool {
|
||||
cursor.sticky_session_token.is_none()
|
||||
&& cursor
|
||||
|
||||
@@ -321,6 +321,7 @@ pub(crate) struct LocalCandidatePreselectionPageCursor<'a> {
|
||||
scanned_rows_by_format: BTreeMap<String, u32>,
|
||||
resolved_global_model_names: BTreeMap<String, String>,
|
||||
fallback_scanned_api_formats: BTreeSet<String>,
|
||||
exhausted_api_formats: BTreeSet<String>,
|
||||
seen_candidate_keys: BTreeSet<String>,
|
||||
}
|
||||
|
||||
@@ -402,6 +403,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
scanned_rows_by_format: BTreeMap::new(),
|
||||
resolved_global_model_names: BTreeMap::new(),
|
||||
fallback_scanned_api_formats: BTreeSet::new(),
|
||||
exhausted_api_formats: BTreeSet::new(),
|
||||
seen_candidate_keys: BTreeSet::new(),
|
||||
}
|
||||
}
|
||||
@@ -426,15 +428,38 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
}
|
||||
}
|
||||
|
||||
// Deferred pages and formats already proven exhausted require no planning
|
||||
// permit. This is the common second-target path for a single-candidate
|
||||
// model, so keep it entirely in memory before joining the shared gate.
|
||||
while self.format_index < self.candidate_api_formats.len() {
|
||||
let candidate_api_format = self.candidate_api_formats[self.format_index].clone();
|
||||
if let Some(outcome) = self.pop_deferred_page(&candidate_api_format) {
|
||||
return Ok(Some(outcome));
|
||||
}
|
||||
let Some(outcome) = self
|
||||
.next_page_for_api_format_with_planning_gate(&candidate_api_format)
|
||||
.await?
|
||||
else {
|
||||
if self.api_format_is_exhausted(&candidate_api_format) {
|
||||
self.format_index += 1;
|
||||
continue;
|
||||
}
|
||||
break;
|
||||
}
|
||||
if self.format_index >= self.candidate_api_formats.len() {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
// One next_page call may need to confirm exhaustion across several API
|
||||
// formats. Hold one permit for that scan instead of rejoining the gate
|
||||
// once per format.
|
||||
let _permit = acquire_candidate_planning_gate(self.state, &self.trace_id).await?;
|
||||
while self.format_index < self.candidate_api_formats.len() {
|
||||
let candidate_api_format = self.candidate_api_formats[self.format_index].clone();
|
||||
if let Some(outcome) = self.pop_deferred_page(&candidate_api_format) {
|
||||
return Ok(Some(outcome));
|
||||
}
|
||||
if self.api_format_is_exhausted(&candidate_api_format) {
|
||||
self.format_index += 1;
|
||||
continue;
|
||||
}
|
||||
let Some(outcome) = self.next_page_for_api_format(&candidate_api_format).await? else {
|
||||
self.format_index += 1;
|
||||
continue;
|
||||
};
|
||||
@@ -453,6 +478,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
self.scanned_rows_by_format.clear();
|
||||
self.resolved_global_model_names.clear();
|
||||
self.fallback_scanned_api_formats.clear();
|
||||
self.exhausted_api_formats.clear();
|
||||
self.seen_candidate_keys.clear();
|
||||
self.priority_page_emitted = false;
|
||||
self.deferred_pages_by_format.clear();
|
||||
@@ -656,22 +682,6 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
self.next_priority_page().await
|
||||
}
|
||||
|
||||
async fn next_page_for_api_format_with_planning_gate(
|
||||
&mut self,
|
||||
candidate_api_format: &str,
|
||||
) -> Result<
|
||||
Option<
|
||||
AiCandidatePreselectionOutcome<
|
||||
SchedulerMinimalCandidateSelectionCandidate,
|
||||
SkippedLocalExecutionCandidate,
|
||||
>,
|
||||
>,
|
||||
GatewayError,
|
||||
> {
|
||||
let _permit = acquire_candidate_planning_gate(self.state, &self.trace_id).await?;
|
||||
self.next_page_for_api_format(candidate_api_format).await
|
||||
}
|
||||
|
||||
async fn split_priority_conversion_page(
|
||||
&self,
|
||||
candidate_api_format: &str,
|
||||
@@ -802,6 +812,9 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
if normalized_api_format.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
if self.exhausted_api_formats.contains(&normalized_api_format) {
|
||||
return Ok(None);
|
||||
}
|
||||
let routing_model = self.routing_model(candidate_api_format).to_string();
|
||||
let requested_names = requested_model_candidate_names(&routing_model, false);
|
||||
let scanned = *self
|
||||
@@ -809,6 +822,8 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
.get(&normalized_api_format)
|
||||
.unwrap_or(&0);
|
||||
if scanned >= REQUESTED_MODEL_MAX_SCANNED_ROWS {
|
||||
self.exhausted_api_formats
|
||||
.insert(normalized_api_format.clone());
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
@@ -839,6 +854,8 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
.unwrap_or(&0);
|
||||
let remaining = REQUESTED_MODEL_MAX_SCANNED_ROWS.saturating_sub(scanned);
|
||||
if remaining == 0 {
|
||||
self.exhausted_api_formats
|
||||
.insert(normalized_api_format.clone());
|
||||
return Ok(None);
|
||||
}
|
||||
let limit = REQUESTED_MODEL_CANDIDATE_PAGE_SIZE.min(remaining);
|
||||
@@ -963,10 +980,12 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
>,
|
||||
GatewayError,
|
||||
> {
|
||||
if !self
|
||||
if self
|
||||
.fallback_scanned_api_formats
|
||||
.insert(normalized_api_format.to_string())
|
||||
.contains(normalized_api_format)
|
||||
{
|
||||
self.exhausted_api_formats
|
||||
.insert(normalized_api_format.to_string());
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
@@ -990,8 +1009,20 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
self.build_page_outcome_from_rows(candidate_api_format, normalized_api_format, rows)
|
||||
.await
|
||||
let outcome = self
|
||||
.build_page_outcome_from_rows(candidate_api_format, normalized_api_format, rows)
|
||||
.await?;
|
||||
self.fallback_scanned_api_formats
|
||||
.insert(normalized_api_format.to_string());
|
||||
self.exhausted_api_formats
|
||||
.insert(normalized_api_format.to_string());
|
||||
Ok(outcome)
|
||||
}
|
||||
|
||||
fn api_format_is_exhausted(&self, candidate_api_format: &str) -> bool {
|
||||
let normalized_api_format = normalize_api_format(candidate_api_format);
|
||||
normalized_api_format.is_empty()
|
||||
|| self.exhausted_api_formats.contains(&normalized_api_format)
|
||||
}
|
||||
|
||||
async fn build_page_outcome_from_rows(
|
||||
@@ -1280,14 +1311,78 @@ mod tests {
|
||||
use crate::AppState;
|
||||
use aether_data::repository::candidate_selection::InMemoryMinimalCandidateSelectionReadRepository;
|
||||
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
|
||||
use aether_data::DataLayerError;
|
||||
use aether_data_contracts::repository::candidate_selection::{
|
||||
MinimalCandidateSelectionReadRepository, StoredProviderModelMapping,
|
||||
MinimalCandidateSelectionReadRepository, StoredPoolKeyCandidateRowsByKeyIdsQuery,
|
||||
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
|
||||
StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::sync::Arc;
|
||||
|
||||
#[derive(Default)]
|
||||
struct EmptyFallbackCountingRepository {
|
||||
fallback_reads: AtomicUsize,
|
||||
}
|
||||
|
||||
impl EmptyFallbackCountingRepository {
|
||||
fn fallback_reads(&self) -> usize {
|
||||
self.fallback_reads.load(Ordering::Acquire)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl MinimalCandidateSelectionReadRepository for EmptyFallbackCountingRepository {
|
||||
async fn list_for_exact_api_format(
|
||||
&self,
|
||||
_api_format: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.fallback_reads.fetch_add(1, Ordering::AcqRel);
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_global_model(
|
||||
&self,
|
||||
_api_format: &str,
|
||||
_global_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model(
|
||||
&self,
|
||||
_api_format: &str,
|
||||
_requested_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model_page(
|
||||
&self,
|
||||
_query: &StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group(
|
||||
&self,
|
||||
_query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group_key_ids(
|
||||
&self,
|
||||
_query: &StoredPoolKeyCandidateRowsByKeyIdsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
}
|
||||
|
||||
fn unrestricted_auth_snapshot() -> GatewayAuthApiKeySnapshot {
|
||||
GatewayAuthApiKeySnapshot {
|
||||
user_id: "user-1".to_string(),
|
||||
@@ -1317,6 +1412,165 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn empty_fallback_is_scanned_once_and_then_skipped_in_memory() {
|
||||
let repository = Arc::new(EmptyFallbackCountingRepository::default());
|
||||
let data_state =
|
||||
GatewayDataState::with_minimal_candidate_selection_reader_for_tests(repository.clone());
|
||||
let app = AppState::new()
|
||||
.expect("gateway state should build")
|
||||
.with_data_state_for_tests(data_state);
|
||||
let auth_snapshot = unrestricted_auth_snapshot();
|
||||
let model_directive_policy =
|
||||
crate::system_features::ModelDirectivePolicySnapshot::load(&app).await;
|
||||
let mut cursor = LocalCandidatePreselectionPageCursor::new(
|
||||
PlannerAppState::new(&app),
|
||||
&model_directive_policy,
|
||||
"openai:search",
|
||||
"missing-model",
|
||||
None,
|
||||
false,
|
||||
None,
|
||||
&auth_snapshot,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
|
||||
false,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(cursor
|
||||
.next_page()
|
||||
.await
|
||||
.expect("empty preselection should succeed")
|
||||
.is_none());
|
||||
assert_eq!(repository.fallback_reads(), 1);
|
||||
assert!(cursor.api_format_is_exhausted("openai:search"));
|
||||
|
||||
// The second target must not reacquire the gate or rescan the empty
|
||||
// fallback after the format was proven exhausted.
|
||||
assert!(cursor
|
||||
.next_page()
|
||||
.await
|
||||
.expect("exhausted preselection should succeed")
|
||||
.is_none());
|
||||
assert_eq!(repository.fallback_reads(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn restart_scan_clears_exhaustion_and_allows_fallback_to_be_read_again() {
|
||||
let repository = Arc::new(EmptyFallbackCountingRepository::default());
|
||||
let data_state =
|
||||
GatewayDataState::with_minimal_candidate_selection_reader_for_tests(repository.clone());
|
||||
let app = AppState::new()
|
||||
.expect("gateway state should build")
|
||||
.with_data_state_for_tests(data_state);
|
||||
let auth_snapshot = unrestricted_auth_snapshot();
|
||||
let model_directive_policy =
|
||||
crate::system_features::ModelDirectivePolicySnapshot::load(&app).await;
|
||||
let mut cursor = LocalCandidatePreselectionPageCursor::new(
|
||||
PlannerAppState::new(&app),
|
||||
&model_directive_policy,
|
||||
"openai:search",
|
||||
"missing-model",
|
||||
None,
|
||||
false,
|
||||
None,
|
||||
&auth_snapshot,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
|
||||
false,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(cursor
|
||||
.next_page()
|
||||
.await
|
||||
.expect("initial scan should succeed")
|
||||
.is_none());
|
||||
assert_eq!(repository.fallback_reads(), 1);
|
||||
cursor.restart_scan();
|
||||
assert!(!cursor.api_format_is_exhausted("openai:search"));
|
||||
assert!(cursor
|
||||
.next_page()
|
||||
.await
|
||||
.expect("restarted scan should succeed")
|
||||
.is_none());
|
||||
assert_eq!(repository.fallback_reads(), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn fallback_can_supply_a_real_second_candidate_after_fast_path_page() {
|
||||
let mut first = standard_candidate_row("provider-first", "openai:chat", 0);
|
||||
first.global_model_name = "gpt-5".to_string();
|
||||
first.global_model_mappings = Some(vec!["gpt-5(?:\\.\\d+)?".to_string()]);
|
||||
first.model_provider_model_name = "gpt-5.1".to_string();
|
||||
|
||||
let mut second = standard_candidate_row("provider-second", "openai:chat", 1);
|
||||
second.global_model_name = "gpt-5".to_string();
|
||||
second.global_model_mappings = Some(vec!["gpt-5(?:\\.\\d+)?".to_string()]);
|
||||
second.model_provider_model_name = "gpt-5-secondary".to_string();
|
||||
|
||||
let repository: Arc<dyn MinimalCandidateSelectionReadRepository> =
|
||||
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed([
|
||||
first, second,
|
||||
]));
|
||||
let data_state =
|
||||
GatewayDataState::with_minimal_candidate_selection_reader_for_tests(repository);
|
||||
let app = AppState::new()
|
||||
.expect("gateway state should build")
|
||||
.with_data_state_for_tests(data_state);
|
||||
let auth_snapshot = unrestricted_auth_snapshot();
|
||||
let model_directive_policy =
|
||||
crate::system_features::ModelDirectivePolicySnapshot::load(&app).await;
|
||||
let mut cursor = LocalCandidatePreselectionPageCursor::new(
|
||||
PlannerAppState::new(&app),
|
||||
&model_directive_policy,
|
||||
"openai:chat",
|
||||
"gpt-5.1",
|
||||
None,
|
||||
false,
|
||||
None,
|
||||
&auth_snapshot,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
|
||||
false,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
let first_page = cursor
|
||||
.next_page()
|
||||
.await
|
||||
.expect("fast-path candidate should load")
|
||||
.expect("first candidate should be present");
|
||||
assert_eq!(first_page.candidates.len(), 1);
|
||||
assert_eq!(first_page.candidates[0].provider_id, "provider-first");
|
||||
|
||||
let second_page = cursor
|
||||
.next_page()
|
||||
.await
|
||||
.expect("fallback candidate should load")
|
||||
.expect("second candidate should not be skipped");
|
||||
assert_eq!(second_page.candidates.len(), 1);
|
||||
assert_eq!(second_page.candidates[0].provider_id, "provider-second");
|
||||
assert!(cursor
|
||||
.next_page()
|
||||
.await
|
||||
.expect("exhausted formats should finish in memory")
|
||||
.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn priority_page_cache_requires_fixed_order_or_explicit_affinity() {
|
||||
let repository: Arc<dyn MinimalCandidateSelectionReadRepository> =
|
||||
|
||||
@@ -18,6 +18,7 @@ use crate::ai_serving::{
|
||||
ExecutionRuntimeAuthContext, GatewayAuthApiKeySnapshot, GatewayProviderTransportSnapshot,
|
||||
PlannerAppState, CODEX_RESPONSES_LITE_HEADER,
|
||||
};
|
||||
use crate::cache::CacheLoadObserver;
|
||||
use crate::client_session_affinity::client_session_affinity_from_api_request;
|
||||
use crate::clock::current_unix_secs;
|
||||
use crate::routing::{
|
||||
@@ -29,7 +30,10 @@ use crate::routing::{
|
||||
use crate::stage_metrics::observe_gateway_stage_ms;
|
||||
use crate::{AiExecutionDecision, AppState, GatewayError};
|
||||
|
||||
// Keep normal freshness bounded for cross-node routing changes. Stale values
|
||||
// are served while a single background refresh updates the cache.
|
||||
const ROUTING_GROUP_SELECTION_CACHE_TTL: Duration = Duration::from_secs(30);
|
||||
const ROUTING_GROUP_SELECTION_CACHE_STALE_TTL: Duration = Duration::from_secs(120);
|
||||
const CODEX_ACCOUNT_ID_HEADER: &str = "chatgpt-account-id";
|
||||
const CODEX_FEDRAMP_HEADER: &str = "x-openai-fedramp";
|
||||
|
||||
@@ -392,47 +396,64 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
|
||||
let explicit_group = routing_header_value_str(&parts.headers, ROUTING_GROUP_HEADER);
|
||||
let selected_group = match state.routing_group_read_repository() {
|
||||
Some(repository) => {
|
||||
let user_groups_lookup_started_at = std::time::Instant::now();
|
||||
let user_group_ids = match state
|
||||
.list_user_groups_for_user(&input.auth_context.user_id)
|
||||
.await
|
||||
{
|
||||
Ok(groups) => groups.into_iter().map(|group| group.id).collect::<Vec<_>>(),
|
||||
Err(error) => {
|
||||
warn!(
|
||||
user_id = %input.auth_context.user_id,
|
||||
error = ?error,
|
||||
"gateway routing profile user group lookup failed"
|
||||
);
|
||||
Vec::new()
|
||||
}
|
||||
// Explicit non-default groups are authorized against principal
|
||||
// bindings, so both selection and its cache key must retain the
|
||||
// caller context. Only the implicit no-binding system-default
|
||||
// path is global and can skip the membership lookup.
|
||||
let principal_context_required = if explicit_group.is_some() {
|
||||
true
|
||||
} else {
|
||||
!matches!(repository.has_any_routing_group_binding().await, Ok(false))
|
||||
};
|
||||
observe_gateway_stage_ms(
|
||||
"routing_user_groups_lookup",
|
||||
user_groups_lookup_started_at.elapsed().as_millis() as u64,
|
||||
);
|
||||
let user_group_ids = if principal_context_required {
|
||||
let user_groups_lookup_started_at = std::time::Instant::now();
|
||||
let user_group_ids = match state
|
||||
.list_user_groups_for_user(&input.auth_context.user_id)
|
||||
.await
|
||||
{
|
||||
Ok(groups) => groups.into_iter().map(|group| group.id).collect::<Vec<_>>(),
|
||||
Err(error) => {
|
||||
warn!(
|
||||
user_id = %input.auth_context.user_id,
|
||||
error = ?error,
|
||||
"gateway routing profile user group lookup failed"
|
||||
);
|
||||
Vec::new()
|
||||
}
|
||||
};
|
||||
observe_gateway_stage_ms(
|
||||
"routing_user_groups_lookup",
|
||||
user_groups_lookup_started_at.elapsed().as_millis() as u64,
|
||||
);
|
||||
user_group_ids
|
||||
} else {
|
||||
Vec::new()
|
||||
};
|
||||
let selection_user_id =
|
||||
principal_context_required.then(|| input.auth_context.user_id.clone());
|
||||
let selection_api_key_id =
|
||||
principal_context_required.then(|| input.auth_context.api_key_id.clone());
|
||||
let selection_cache_key = routing_group_selection_cache_key(
|
||||
explicit_group.as_deref(),
|
||||
Some(input.auth_context.user_id.as_str()),
|
||||
Some(input.auth_context.api_key_id.as_str()),
|
||||
selection_user_id.as_deref(),
|
||||
selection_api_key_id.as_deref(),
|
||||
&user_group_ids,
|
||||
);
|
||||
let user_id = input.auth_context.user_id.clone();
|
||||
let api_key_id = input.auth_context.api_key_id.clone();
|
||||
let group_selection_started_at = std::time::Instant::now();
|
||||
let selection = state
|
||||
.routing_group_selection_cache
|
||||
.get_or_load_once(
|
||||
.get_or_load_once_stale_while_revalidating(
|
||||
selection_cache_key,
|
||||
ROUTING_GROUP_SELECTION_CACHE_TTL,
|
||||
|| async move {
|
||||
ROUTING_GROUP_SELECTION_CACHE_STALE_TTL,
|
||||
|| async {
|
||||
let selection_load_started_at = std::time::Instant::now();
|
||||
let selection = select_gateway_routing_group(
|
||||
repository.as_ref(),
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: explicit_group.as_deref(),
|
||||
user_id: Some(user_id.as_str()),
|
||||
api_key_id: Some(api_key_id.as_str()),
|
||||
user_id: selection_user_id.as_deref(),
|
||||
api_key_id: selection_api_key_id.as_deref(),
|
||||
user_group_ids: &user_group_ids,
|
||||
},
|
||||
)
|
||||
@@ -444,6 +465,33 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
|
||||
);
|
||||
Ok::<_, GatewayError>(Some(selection))
|
||||
},
|
||||
|| {
|
||||
let repository = repository.clone();
|
||||
let explicit_group = explicit_group.clone();
|
||||
let user_id = selection_user_id.clone();
|
||||
let api_key_id = selection_api_key_id.clone();
|
||||
let user_group_ids = user_group_ids.clone();
|
||||
async move {
|
||||
let selection_load_started_at = std::time::Instant::now();
|
||||
let selection = select_gateway_routing_group(
|
||||
repository.as_ref(),
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: explicit_group.as_deref(),
|
||||
user_id: user_id.as_deref(),
|
||||
api_key_id: api_key_id.as_deref(),
|
||||
user_group_ids: &user_group_ids,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.map_err(routing_selection_error)?;
|
||||
observe_gateway_stage_ms(
|
||||
"routing_group_selection_load",
|
||||
selection_load_started_at.elapsed().as_millis() as u64,
|
||||
);
|
||||
Ok::<_, GatewayError>(Some(selection))
|
||||
}
|
||||
},
|
||||
CacheLoadObserver::default(),
|
||||
)
|
||||
.await?
|
||||
.unwrap_or_default();
|
||||
@@ -919,12 +967,119 @@ fn ensure_report_context_routing_trace(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use super::*;
|
||||
use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository;
|
||||
use aether_data_contracts::repository::routing_profiles::{
|
||||
CreateRoutingGroupBindingRecord, CreateRoutingGroupRecord, RoutingGroupBindingSubject,
|
||||
RoutingGroupWriteRepository,
|
||||
};
|
||||
use aether_provider_transport::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn explicit_routing_selection_cache_key_is_principal_specific() {
|
||||
let first = routing_group_selection_cache_key(
|
||||
Some("private"),
|
||||
Some("user-1"),
|
||||
Some("key-1"),
|
||||
&["team-1".to_string()],
|
||||
);
|
||||
let second = routing_group_selection_cache_key(
|
||||
Some("private"),
|
||||
Some("user-2"),
|
||||
Some("key-2"),
|
||||
&["team-2".to_string()],
|
||||
);
|
||||
|
||||
assert_ne!(first, second);
|
||||
assert!(first.contains("user=user-1"));
|
||||
assert!(first.contains("api_key=key-1"));
|
||||
assert!(first.contains("groups=team-1"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn explicit_routing_attachment_authorizes_and_caches_per_principal() {
|
||||
let repository = Arc::new(InMemoryRoutingGroupRepository::default());
|
||||
repository
|
||||
.create_routing_group(CreateRoutingGroupRecord {
|
||||
id: "private-group".to_string(),
|
||||
name: "private".to_string(),
|
||||
description: None,
|
||||
enabled: true,
|
||||
is_system_default: false,
|
||||
config_json: json!({}),
|
||||
version: 1,
|
||||
created_at: 1,
|
||||
updated_at: 1,
|
||||
published_at: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
repository
|
||||
.create_routing_group_binding(CreateRoutingGroupBindingRecord {
|
||||
id: "binding-user-1".to_string(),
|
||||
group_id: "private-group".to_string(),
|
||||
subject_type: RoutingGroupBindingSubject::User,
|
||||
subject_id: "user-1".to_string(),
|
||||
is_default: false,
|
||||
allow_explicit_select: true,
|
||||
created_at: 1,
|
||||
updated_at: 1,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let state = AppState::new().unwrap().with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::disabled()
|
||||
.with_routing_group_repository_for_tests(repository),
|
||||
);
|
||||
let (parts, _) = http::Request::builder()
|
||||
.header(ROUTING_GROUP_HEADER, "private-group")
|
||||
.body(())
|
||||
.unwrap()
|
||||
.into_parts();
|
||||
|
||||
let mut allowed = sample_decision_input();
|
||||
attach_routing_policy_to_local_requested_model_input(
|
||||
&state,
|
||||
&parts,
|
||||
&mut allowed,
|
||||
&json!({"model": "gpt-5"}),
|
||||
"openai:chat",
|
||||
)
|
||||
.await
|
||||
.expect("bound principal should explicitly select the private group");
|
||||
let policy = allowed
|
||||
.routing_policy
|
||||
.as_ref()
|
||||
.expect("explicit selection should attach routing policy");
|
||||
assert_eq!(policy.group_id.as_deref(), Some("private-group"));
|
||||
assert_eq!(policy.selection_source, "explicit_header");
|
||||
|
||||
let mut denied = sample_decision_input();
|
||||
denied.auth_context.user_id = "user-2".to_string();
|
||||
denied.auth_context.api_key_id = "api-key-2".to_string();
|
||||
let error = attach_routing_policy_to_local_requested_model_input(
|
||||
&state,
|
||||
&parts,
|
||||
&mut denied,
|
||||
&json!({"model": "gpt-5"}),
|
||||
"openai:chat",
|
||||
)
|
||||
.await
|
||||
.expect_err("another principal must not reuse the authorized cache entry");
|
||||
match error {
|
||||
GatewayError::Client { status, message } => {
|
||||
assert_eq!(status, StatusCode::FORBIDDEN);
|
||||
assert!(message.contains("not allowed for this principal"));
|
||||
}
|
||||
other => panic!("unexpected explicit selection error: {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_auth_context() -> ExecutionRuntimeAuthContext {
|
||||
ExecutionRuntimeAuthContext {
|
||||
user_id: "user-1".to_string(),
|
||||
|
||||
Reference in New Issue
Block a user