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
@@ -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(),