diff --git a/apps/aether-gateway/src/ai_serving/mod.rs b/apps/aether-gateway/src/ai_serving/mod.rs index fb696fd3f..9038c71ae 100644 --- a/apps/aether-gateway/src/ai_serving/mod.rs +++ b/apps/aether-gateway/src/ai_serving/mod.rs @@ -55,7 +55,8 @@ pub(crate) use self::planner::{ maybe_build_sync_decision_payload, maybe_build_sync_plan_payload, planner_is_matching_stream_request, provider_key_pool_score_id, provider_key_pool_score_scope, read_candidate_transport_snapshot, record_local_runtime_candidate_skip_reason, - resolve_upstream_is_stream_for_provider, set_local_openai_chat_execution_exhausted_diagnostic, + resolve_tunnel_scheduler_affinity_context, resolve_upstream_is_stream_for_provider, + set_local_openai_chat_execution_exhausted_diagnostic, set_local_openai_image_execution_exhausted_diagnostic, validate_final_openai_provider_request, CandidateFailureDiagnostic, CandidateFailureDiagnosticKind, EligibleLocalExecutionCandidate, GatewayAuthApiKeySnapshot, GatewayProviderTransportSnapshot, LocalExecutionAttemptSource, diff --git a/apps/aether-gateway/src/ai_serving/planner/candidate_affinity_cache.rs b/apps/aether-gateway/src/ai_serving/planner/candidate_affinity_cache.rs index 3a77c8b95..9a00a457d 100644 --- a/apps/aether-gateway/src/ai_serving/planner/candidate_affinity_cache.rs +++ b/apps/aether-gateway/src/ai_serving/planner/candidate_affinity_cache.rs @@ -1,6 +1,8 @@ +use aether_routing_core::ResolvedRoutingPolicy; use aether_scheduler_core::{ - build_scheduler_affinity_cache_key_for_api_key_id_with_client_session, ClientSessionAffinity, - SchedulerAffinityTarget, SchedulerMinimalCandidateSelectionCandidate, + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope, + ClientSessionAffinity, SchedulerAffinityScope, SchedulerAffinityTarget, + SchedulerMinimalCandidateSelectionCandidate, }; use crate::ai_serving::{GatewayAuthApiKeySnapshot, PlannerAppState}; @@ -20,6 +22,7 @@ pub(crate) fn read_cached_scheduler_affinity_target( client_session_affinity: Option<&ClientSessionAffinity>, client_api_format: &str, requested_model: Option<&str>, + routing_policy: Option<&ResolvedRoutingPolicy>, ) -> Option { if !has_explicit_session_affinity(client_session_affinity) { return None; @@ -30,12 +33,15 @@ pub(crate) fn read_cached_scheduler_affinity_target( let api_key_id = auth_snapshot .map(|snapshot| snapshot.api_key_id.trim()) .filter(|value| !value.is_empty())?; - let cache_key = build_scheduler_affinity_cache_key_for_api_key_id_with_client_session( - api_key_id, - client_api_format, - requested_model, - client_session_affinity, - )?; + let affinity_scope = scheduler_affinity_scope_for_routing_policy(routing_policy); + let cache_key = + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope( + api_key_id, + client_api_format, + requested_model, + client_session_affinity, + affinity_scope.as_ref(), + )?; state .app() @@ -72,6 +78,53 @@ pub(crate) fn remember_scheduler_affinity_for_candidate_at_epoch( requested_model: &str, candidate: &SchedulerMinimalCandidateSelectionCandidate, expected_epoch: Option, +) { + remember_scheduler_affinity_for_candidate_with_scope_at_epoch( + state, + auth_snapshot, + client_session_affinity, + client_api_format, + requested_model, + candidate, + None, + expected_epoch, + ); +} + +#[allow(clippy::too_many_arguments)] +pub(crate) fn remember_scheduler_affinity_for_candidate_with_routing_policy_at_epoch( + state: PlannerAppState<'_>, + auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, + client_session_affinity: Option<&ClientSessionAffinity>, + client_api_format: &str, + requested_model: &str, + candidate: &SchedulerMinimalCandidateSelectionCandidate, + routing_policy: Option<&ResolvedRoutingPolicy>, + expected_epoch: Option, +) { + let affinity_scope = scheduler_affinity_scope_for_routing_policy(routing_policy); + remember_scheduler_affinity_for_candidate_with_scope_at_epoch( + state, + auth_snapshot, + client_session_affinity, + client_api_format, + requested_model, + candidate, + affinity_scope.as_ref(), + expected_epoch, + ); +} + +#[allow(clippy::too_many_arguments)] +fn remember_scheduler_affinity_for_candidate_with_scope_at_epoch( + state: PlannerAppState<'_>, + auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, + client_session_affinity: Option<&ClientSessionAffinity>, + client_api_format: &str, + requested_model: &str, + candidate: &SchedulerMinimalCandidateSelectionCandidate, + affinity_scope: Option<&SchedulerAffinityScope>, + expected_epoch: Option, ) { if !has_explicit_session_affinity(client_session_affinity) { return; @@ -82,12 +135,15 @@ pub(crate) fn remember_scheduler_affinity_for_candidate_at_epoch( else { return; }; - let Some(cache_key) = build_scheduler_affinity_cache_key_for_api_key_id_with_client_session( - api_key_id, - client_api_format, - requested_model, - client_session_affinity, - ) else { + let Some(cache_key) = + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope( + api_key_id, + client_api_format, + requested_model, + client_session_affinity, + affinity_scope, + ) + else { return; }; @@ -103,3 +159,15 @@ pub(crate) fn remember_scheduler_affinity_for_candidate_at_epoch( expected_epoch, ); } + +fn scheduler_affinity_scope_for_routing_policy( + routing_policy: Option<&ResolvedRoutingPolicy>, +) -> Option { + let policy = routing_policy?; + let group_id = policy + .group_id + .as_deref() + .map(str::trim) + .filter(|group_id| !group_id.is_empty())?; + Some(SchedulerAffinityScope::new(group_id, policy.group_version)) +} diff --git a/apps/aether-gateway/src/ai_serving/planner/candidate_materialization.rs b/apps/aether-gateway/src/ai_serving/planner/candidate_materialization.rs index 384b5644c..3aff505cf 100644 --- a/apps/aether-gateway/src/ai_serving/planner/candidate_materialization.rs +++ b/apps/aether-gateway/src/ai_serving/planner/candidate_materialization.rs @@ -24,7 +24,7 @@ use tokio::time::Instant; use tracing::warn; use uuid::Uuid; -use crate::ai_serving::planner::candidate_affinity_cache::remember_scheduler_affinity_for_candidate_at_epoch; +use crate::ai_serving::planner::candidate_affinity_cache::remember_scheduler_affinity_for_candidate_with_routing_policy_at_epoch; use crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy; use crate::ai_serving::planner::candidate_resolution::{ resolve_and_rank_logical_local_execution_candidates, EligibleLocalExecutionCandidate, @@ -421,6 +421,7 @@ where self.client_session_affinity, self.client_api_format, self.requested_model, + self.routing_policy, candidates, ); } @@ -691,6 +692,7 @@ where client_session_affinity, client_api_format, requested_model, + routing_policy, &candidates, ); } @@ -983,15 +985,7 @@ impl<'a> RequestedModelAttemptPageCursor<'a> { "candidate_page_load", page_started_at.elapsed().as_millis() as u64, ); - if matches!(error, GatewayError::AdmissionTimeout { .. }) { - return Err(error); - } - warn!( - trace_id = %self.trace_id, - error = ?error, - "gateway lazy requested-model candidate page read failed" - ); - return Ok(false); + return Err(error); } }; observe_gateway_stage_ms( @@ -1035,6 +1029,7 @@ impl<'a> RequestedModelAttemptPageCursor<'a> { self.client_session_affinity.as_ref(), &self.client_api_format, Some(&self.requested_model), + self.routing_policy.as_ref(), &candidates, ); self.remembered_affinity = true; @@ -1259,6 +1254,7 @@ pub(crate) fn remember_first_local_candidate_affinity( client_session_affinity: Option<&ClientSessionAffinity>, client_api_format: &str, requested_model: Option<&str>, + routing_policy: Option<&ResolvedRoutingPolicy>, candidates: &[EligibleLocalExecutionCandidate], ) { let Some(first_candidate) = candidates.first() else { @@ -1268,13 +1264,14 @@ pub(crate) fn remember_first_local_candidate_affinity( .map(str::trim) .filter(|value| !value.is_empty()) .unwrap_or(first_candidate.candidate.global_model_name.as_str()); - remember_scheduler_affinity_for_candidate_at_epoch( + remember_scheduler_affinity_for_candidate_with_routing_policy_at_epoch( state, auth_snapshot, client_session_affinity, client_api_format, affinity_requested_model, &first_candidate.candidate, + routing_policy, first_candidate.orchestration.scheduler_affinity_epoch, ); } diff --git a/apps/aether-gateway/src/ai_serving/planner/candidate_ranking.rs b/apps/aether-gateway/src/ai_serving/planner/candidate_ranking.rs index daaf1a7c1..9c492583f 100644 --- a/apps/aether-gateway/src/ai_serving/planner/candidate_ranking.rs +++ b/apps/aether-gateway/src/ai_serving/planner/candidate_ranking.rs @@ -66,6 +66,7 @@ impl AiCandidateRankingPort for GatewayLocalCandidateRankingPort<'_> { self.client_session_affinity, normalized_client_api_format, affinity_requested_model, + self.routing_policy, )) } @@ -221,17 +222,21 @@ fn routing_overlaid_candidate( let mut overlaid = candidate.clone(); overlaid.provider_priority = policy .ranking_overlay - .provider_priority_or_unspecified(candidate.provider_id.as_str()); + .provider_priority(candidate.provider_id.as_str(), candidate.provider_priority); let overlaid_key_priority = match kind { LocalExecutionCandidateKind::SingleKey => policy .ranking_overlay - .key_priority_or_unspecified(candidate.key_id.as_str()), + .key_priority_overrides + .get(candidate.key_id.as_str()), LocalExecutionCandidateKind::PoolGroup => policy .ranking_overlay - .pool_priority_or_unspecified(candidate.provider_id.as_str()), + .pool_priority_overrides + .get(candidate.provider_id.as_str()), }; - overlaid.key_internal_priority = overlaid_key_priority; - overlaid.key_global_priority_for_format = Some(overlaid_key_priority); + if let Some(overlaid_key_priority) = overlaid_key_priority.copied() { + overlaid.key_internal_priority = overlaid_key_priority; + overlaid.key_global_priority_for_format = Some(overlaid_key_priority); + } overlaid } @@ -354,7 +359,7 @@ mod tests { } #[test] - fn routing_policy_priorities_do_not_fall_back_to_candidate_priorities() { + fn routing_policy_priorities_fall_back_to_candidate_priorities() { let mut candidate = sample_candidate("endpoint-1", "key-1"); candidate.provider_priority = 7; candidate.key_internal_priority = 3; @@ -380,18 +385,9 @@ mod tests { &candidate, ); - assert_eq!( - overlaid.provider_priority, - aether_routing_core::ROUTING_PRIORITY_UNSPECIFIED - ); - assert_eq!( - overlaid.key_internal_priority, - aether_routing_core::ROUTING_PRIORITY_UNSPECIFIED - ); - assert_eq!( - overlaid.key_global_priority_for_format, - Some(aether_routing_core::ROUTING_PRIORITY_UNSPECIFIED) - ); + assert_eq!(overlaid.provider_priority, 7); + assert_eq!(overlaid.key_internal_priority, 3); + assert_eq!(overlaid.key_global_priority_for_format, Some(2)); } #[test] diff --git a/apps/aether-gateway/src/ai_serving/planner/candidate_source.rs b/apps/aether-gateway/src/ai_serving/planner/candidate_source.rs index d5d938c9a..d5ad3c92c 100644 --- a/apps/aether-gateway/src/ai_serving/planner/candidate_source.rs +++ b/apps/aether-gateway/src/ai_serving/planner/candidate_source.rs @@ -29,9 +29,10 @@ use crate::cache::{ }; use crate::clock::request_distribution_seed; use crate::data::candidate_selection::{ - read_requested_model_rows_fast_path_page, requested_model_candidate_names, - MinimalCandidateSelectionRowSource, RequestedModelCandidateRowsPage, - REQUESTED_MODEL_CANDIDATE_PAGE_SIZE, REQUESTED_MODEL_MAX_SCANNED_ROWS, + read_api_format_rows_fallback_page, read_requested_model_rows_fast_path_page, + requested_model_candidate_names, MinimalCandidateSelectionRowSource, + RequestedModelCandidateRowsPage, REQUESTED_MODEL_CANDIDATE_PAGE_SIZE, + REQUESTED_MODEL_MAX_SCANNED_ROWS, }; use crate::scheduler::candidate::SchedulerSkippedCandidate; use crate::scheduler::config::{SchedulerOrderingConfig, SchedulerSchedulingMode}; @@ -120,7 +121,10 @@ impl AiCandidatePreselectionPort for GatewayLocalCandidatePreselectionPort<'_> { self.require_streaming, self.required_capabilities, auth_snapshot, - self.client_session_affinity, + self.routing_policy + .is_none() + .then_some(self.client_session_affinity) + .flatten(), self.ranking_seed, false, self.request_operation, @@ -320,7 +324,8 @@ pub(crate) struct LocalCandidatePreselectionPageCursor<'a> { requested_name_offsets: BTreeMap, scanned_rows_by_format: BTreeMap, resolved_global_model_names: BTreeMap, - fallback_scanned_api_formats: BTreeSet, + fallback_offsets: BTreeMap, + fallback_scan_epoch: u32, exhausted_api_formats: BTreeSet, seen_candidate_keys: BTreeSet, } @@ -402,7 +407,8 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> { requested_name_offsets: BTreeMap::new(), scanned_rows_by_format: BTreeMap::new(), resolved_global_model_names: BTreeMap::new(), - fallback_scanned_api_formats: BTreeSet::new(), + fallback_offsets: BTreeMap::new(), + fallback_scan_epoch: 0, exhausted_api_formats: BTreeSet::new(), seen_candidate_keys: BTreeSet::new(), } @@ -421,13 +427,35 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> { > { if !self.priority_page_emitted { self.priority_page_emitted = true; - let priority_page = self.cached_next_priority_page().await?; + let mut priority_page = self.cached_next_priority_page().await?; + if self.routing_policy.is_some() { + while let Some(mut page) = self.next_page_after_priority().await? { + priority_page.candidates.append(&mut page.candidates); + priority_page + .skipped_candidates + .append(&mut page.skipped_candidates); + } + } if !priority_page.candidates.is_empty() || !priority_page.skipped_candidates.is_empty() { return Ok(Some(priority_page)); } } + self.next_page_after_priority().await + } + + async fn next_page_after_priority( + &mut self, + ) -> Result< + Option< + AiCandidatePreselectionOutcome< + SchedulerMinimalCandidateSelectionCandidate, + SkippedLocalExecutionCandidate, + >, + >, + GatewayError, + > { // 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. @@ -477,7 +505,8 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> { self.requested_name_offsets.clear(); self.scanned_rows_by_format.clear(); self.resolved_global_model_names.clear(); - self.fallback_scanned_api_formats.clear(); + self.fallback_offsets.clear(); + self.fallback_scan_epoch = self.fallback_scan_epoch.wrapping_add(1); self.exhausted_api_formats.clear(); self.seen_candidate_keys.clear(); self.priority_page_emitted = false; @@ -967,6 +996,63 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> { } } + async fn read_api_format_rows_fallback_page_cached( + &self, + normalized_api_format: &str, + offset: u32, + limit: u32, + ) -> Result { + let key = CandidateRowPageCacheKey::for_api_format_fallback( + normalized_api_format, + offset, + limit, + self.fallback_scan_epoch, + ); + let cache = self.state.app().candidate_row_page_cache.clone(); + let ttl = candidate_page_cache_ttl_from_env(); + let stale_ttl = candidate_page_cache_stale_ttl(ttl); + let cached = cache + .get_or_load_once_stale_while_refreshing( + key, + ttl, + stale_ttl, + || async { + let page = read_api_format_rows_fallback_page( + self.state.app().data.as_ref(), + normalized_api_format, + offset, + limit, + ) + .await + .map_err(|err| GatewayError::Internal(err.to_string()))?; + Ok::<_, GatewayError>(Some(Arc::new(page))) + }, + CacheLoadObserver::new() + .on_hit(record_candidate_row_page_cache_hit) + .on_miss(record_candidate_row_page_cache_miss) + .on_load(record_candidate_row_page_cache_load) + .on_follower_wait(record_candidate_row_page_cache_follower_wait), + ) + .await?; + + match cached { + Some(page) => { + if page.rows.is_empty() { + record_candidate_row_page_cache_none(); + } + Ok(page.as_ref().clone()) + } + None => { + record_candidate_row_page_cache_none(); + Ok(RequestedModelCandidateRowsPage { + rows: Vec::new(), + scanned_rows: 0, + end_of_requested_name: true, + }) + } + } + } + async fn next_fallback_page_for_api_format( &mut self, candidate_api_format: &str, @@ -980,43 +1066,67 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> { >, GatewayError, > { - if self - .fallback_scanned_api_formats - .contains(normalized_api_format) - { - self.exhausted_api_formats - .insert(normalized_api_format.to_string()); - return Ok(None); - } - let routing_model = self.routing_model(candidate_api_format).to_string(); - let rows = self - .state - .app() - .data - .read_minimal_candidate_selection_rows_for_api_format(normalized_api_format) - .await - .map_err(|err| GatewayError::Internal(err.to_string()))? - .into_iter() - .filter(|row| { - row_supports_requested_model_with_model_directives_and_request_operation( - row, - &routing_model, - normalized_api_format, - false, - self.request_operation.as_deref(), - ) - }) - .collect::>(); + loop { + let scanned = *self + .scanned_rows_by_format + .get(normalized_api_format) + .unwrap_or(&0); + let remaining = REQUESTED_MODEL_MAX_SCANNED_ROWS.saturating_sub(scanned); + if remaining == 0 { + self.exhausted_api_formats + .insert(normalized_api_format.to_string()); + return Ok(None); + } + let limit = REQUESTED_MODEL_CANDIDATE_PAGE_SIZE.min(remaining); + let offset = *self + .fallback_offsets + .get(normalized_api_format) + .unwrap_or(&0); + let page = self + .read_api_format_rows_fallback_page_cached(normalized_api_format, offset, limit) + .await?; + let page_scanned = page.scanned_rows.min(limit); + let end_of_format = page.end_of_requested_name || page_scanned < limit; + self.fallback_offsets.insert( + normalized_api_format.to_string(), + offset.saturating_add(page_scanned), + ); + let total_scanned = scanned.saturating_add(page_scanned); + self.scanned_rows_by_format + .insert(normalized_api_format.to_string(), total_scanned); + if end_of_format || total_scanned >= REQUESTED_MODEL_MAX_SCANNED_ROWS { + self.exhausted_api_formats + .insert(normalized_api_format.to_string()); + } + if page_scanned == 0 { + return Ok(None); + } - 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) + let rows = page + .rows + .into_iter() + .take(page_scanned as usize) + .filter(|row| { + row_supports_requested_model_with_model_directives_and_request_operation( + row, + &routing_model, + normalized_api_format, + false, + self.request_operation.as_deref(), + ) + }) + .collect::>(); + if let Some(outcome) = self + .build_page_outcome_from_rows(candidate_api_format, normalized_api_format, rows) + .await? + { + return Ok(Some(outcome)); + } + if self.exhausted_api_formats.contains(normalized_api_format) { + return Ok(None); + } + } } fn api_format_is_exhausted(&self, candidate_api_format: &str) -> bool { @@ -1125,7 +1235,10 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> { candidates, self.required_capabilities.as_ref(), auth_snapshot, - self.client_session_affinity.as_ref(), + self.routing_policy + .is_none() + .then_some(self.client_session_affinity.as_ref()) + .flatten(), self.ranking_seed, ) .await?; @@ -1313,22 +1426,116 @@ mod tests { use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::DataLayerError; use aether_data_contracts::repository::candidate_selection::{ - MinimalCandidateSelectionReadRepository, StoredPoolKeyCandidateRowsByKeyIdsQuery, - StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping, - StoredRequestedModelCandidateRowsQuery, + MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery, + 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; + use std::sync::{Arc, Mutex}; #[derive(Default)] struct EmptyFallbackCountingRepository { fallback_reads: AtomicUsize, } + struct PagedFallbackRepository { + total_rows: u32, + page_queries: Mutex>, + } + + impl PagedFallbackRepository { + fn new(total_rows: u32) -> Self { + Self { + total_rows, + page_queries: Mutex::new(Vec::new()), + } + } + + fn page_queries(&self) -> Vec { + self.page_queries + .lock() + .expect("fallback query lock") + .clone() + } + } + + #[async_trait] + impl MinimalCandidateSelectionReadRepository for PagedFallbackRepository { + async fn list_for_exact_api_format( + &self, + _api_format: &str, + ) -> Result, DataLayerError> { + panic!("routing fallback must not use the unbounded API-format query") + } + + async fn list_for_exact_api_format_page( + &self, + query: &StoredApiFormatCandidateRowsQuery, + ) -> Result, DataLayerError> { + self.page_queries + .lock() + .expect("fallback query lock") + .push(query.clone()); + if normalize_api_format(&query.api_format) != "openai:chat" { + return Ok(Vec::new()); + } + let end = query + .offset + .saturating_add(query.limit) + .min(self.total_rows); + Ok((query.offset..end) + .map(|index| { + standard_candidate_row( + format!("fallback-provider-{index:04}").as_str(), + "openai:chat", + i32::try_from(index).expect("test provider priority should fit"), + ) + }) + .collect()) + } + + async fn list_for_exact_api_format_and_global_model( + &self, + _api_format: &str, + _global_model_name: &str, + ) -> Result, DataLayerError> { + Ok(Vec::new()) + } + + async fn list_for_exact_api_format_and_requested_model( + &self, + _api_format: &str, + _requested_model_name: &str, + ) -> Result, DataLayerError> { + Ok(Vec::new()) + } + + async fn list_for_exact_api_format_and_requested_model_page( + &self, + _query: &StoredRequestedModelCandidateRowsQuery, + ) -> Result, DataLayerError> { + Ok(Vec::new()) + } + + async fn list_pool_key_rows_for_group( + &self, + _query: &StoredPoolKeyCandidateRowsQuery, + ) -> Result, DataLayerError> { + Ok(Vec::new()) + } + + async fn list_pool_key_rows_for_group_key_ids( + &self, + _query: &StoredPoolKeyCandidateRowsByKeyIdsQuery, + ) -> Result, DataLayerError> { + Ok(Vec::new()) + } + } + impl EmptyFallbackCountingRepository { fn fallback_reads(&self) -> usize { self.fallback_reads.load(Ordering::Acquire) @@ -1571,6 +1778,155 @@ mod tests { .is_none()); } + #[tokio::test] + async fn routing_policy_collects_candidate_pages_before_final_ranking() { + let rows = (0..300) + .map(|index| { + standard_candidate_row( + format!("provider-{index:03}").as_str(), + "openai:chat", + index, + ) + }) + .collect::>(); + let repository: Arc = + Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)); + let app = AppState::new() + .expect("gateway state should build") + .with_data_state_for_tests( + GatewayDataState::with_minimal_candidate_selection_reader_for_tests(repository), + ); + let auth_snapshot = unrestricted_auth_snapshot(); + let model_directive_policy = + crate::system_features::ModelDirectivePolicySnapshot::load(&app).await; + let routing_policy = ResolvedRoutingPolicy { + group_id: Some("routing-group-1".to_string()), + group_version: Some(1), + selection_source: "test".to_string(), + requested_model: "gpt-5".to_string(), + resolved_model: "gpt-5".to_string(), + priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider, + scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder, + keep_priority_on_conversion: false, + ranking_overlay: Default::default(), + mutation_plan: Default::default(), + pool_policy_overrides: Default::default(), + matched_rules: Vec::new(), + }; + let mut cursor = LocalCandidatePreselectionPageCursor::new( + PlannerAppState::new(&app), + &model_directive_policy, + "openai:chat", + "gpt-5", + None, + false, + None, + &auth_snapshot, + Some(&routing_policy), + None, + None, + true, + LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat, + false, + None, + ) + .await; + + let candidates = cursor + .next_page() + .await + .expect("routing candidate scan should succeed") + .expect("routing candidates should be present") + .candidates; + + assert_eq!(candidates.len(), 300); + assert!(cursor + .next_page() + .await + .expect("routing scan should be exhausted") + .is_none()); + } + + #[tokio::test] + async fn routing_fallback_uses_bounded_api_format_pages() { + let repository = Arc::new(PagedFallbackRepository::new( + REQUESTED_MODEL_MAX_SCANNED_ROWS + REQUESTED_MODEL_CANDIDATE_PAGE_SIZE, + )); + let app = AppState::new() + .expect("gateway state should build") + .with_data_state_for_tests( + GatewayDataState::with_minimal_candidate_selection_reader_for_tests( + repository.clone(), + ), + ); + let auth_snapshot = unrestricted_auth_snapshot(); + let model_directive_policy = + crate::system_features::ModelDirectivePolicySnapshot::load(&app).await; + let routing_policy = ResolvedRoutingPolicy { + group_id: Some("routing-group-fallback".to_string()), + group_version: Some(1), + selection_source: "test".to_string(), + requested_model: "gpt-5".to_string(), + resolved_model: "gpt-5".to_string(), + priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider, + scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder, + keep_priority_on_conversion: false, + ranking_overlay: Default::default(), + mutation_plan: Default::default(), + pool_policy_overrides: Default::default(), + matched_rules: Vec::new(), + }; + let mut cursor = LocalCandidatePreselectionPageCursor::new( + PlannerAppState::new(&app), + &model_directive_policy, + "openai:chat", + "gpt-5", + None, + false, + None, + &auth_snapshot, + Some(&routing_policy), + None, + None, + true, + LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat, + false, + None, + ) + .await; + + let candidates = cursor + .next_page() + .await + .expect("routing fallback scan should succeed") + .expect("routing fallback candidates should be present") + .candidates; + + assert_eq!(candidates.len(), REQUESTED_MODEL_MAX_SCANNED_ROWS as usize); + assert!(cursor + .next_page() + .await + .expect("bounded routing fallback should be exhausted") + .is_none()); + let page_queries = repository + .page_queries() + .into_iter() + .filter(|query| normalize_api_format(&query.api_format) == "openai:chat") + .collect::>(); + assert_eq!( + page_queries.len(), + (REQUESTED_MODEL_MAX_SCANNED_ROWS / REQUESTED_MODEL_CANDIDATE_PAGE_SIZE) as usize + ); + for (index, query) in page_queries.iter().enumerate() { + assert_eq!(query.limit, REQUESTED_MODEL_CANDIDATE_PAGE_SIZE); + assert_eq!( + query.offset, + u32::try_from(index).expect("page index should fit") + * REQUESTED_MODEL_CANDIDATE_PAGE_SIZE + ); + } + } + #[tokio::test] async fn priority_page_cache_requires_fixed_order_or_explicit_affinity() { let repository: Arc = diff --git a/apps/aether-gateway/src/ai_serving/planner/decision_input.rs b/apps/aether-gateway/src/ai_serving/planner/decision_input.rs index aff69b55f..784106113 100644 --- a/apps/aether-gateway/src/ai_serving/planner/decision_input.rs +++ b/apps/aether-gateway/src/ai_serving/planner/decision_input.rs @@ -11,7 +11,6 @@ use async_trait::async_trait; use http::StatusCode; use http::{HeaderMap, HeaderName, HeaderValue}; use serde_json::{json, Value}; -use tracing::warn; use crate::ai_serving::planner::common::extract_standard_requested_model; use crate::ai_serving::{ @@ -408,29 +407,28 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input( let principal_context_required = if explicit_group.is_some() { true } else { - !matches!(repository.has_any_routing_group_binding().await, Ok(false)) + repository + .has_any_routing_group_binding() + .await + .map_err(|error| { + routing_selection_error(GatewayRoutingSelectionError::Repository( + error.to_string(), + )) + })? }; let user_group_ids = if principal_context_required { let user_groups_lookup_started_at = std::time::Instant::now(); - let user_group_ids = match state + let user_groups = state .list_user_groups_for_user(&input.auth_context.user_id) - .await - { - Ok(groups) => groups.into_iter().map(|group| group.id).collect::>(), - Err(error) => { - warn!( - user_id = %input.auth_context.user_id, - error = ?error, - "gateway routing profile user group lookup failed" - ); - Vec::new() - } - }; + .await; observe_gateway_stage_ms( "routing_user_groups_lookup", user_groups_lookup_started_at.elapsed().as_millis() as u64, ); - user_group_ids + user_groups? + .into_iter() + .map(|group| group.id) + .collect::>() } else { Vec::new() }; @@ -735,9 +733,14 @@ pub(crate) async fn resolve_local_authenticated_decision_input( } fn routing_selection_error(error: GatewayRoutingSelectionError) -> GatewayError { - GatewayError::Client { - status: StatusCode::FORBIDDEN, - message: error.to_string(), + match error { + GatewayRoutingSelectionError::Repository(message) => { + GatewayError::Internal(format!("routing group repository lookup failed: {message}")) + } + error => GatewayError::Client { + status: StatusCode::FORBIDDEN, + message: error.to_string(), + }, } } @@ -1006,6 +1009,21 @@ mod tests { assert!(first.contains("groups=team-1")); } + #[test] + fn routing_repository_failure_maps_to_internal_gateway_error() { + let error = routing_selection_error(GatewayRoutingSelectionError::Repository( + "sql error: database unavailable".to_string(), + )); + + match error { + GatewayError::Internal(message) => { + assert!(message.contains("routing group repository lookup failed")); + assert!(message.contains("database unavailable")); + } + other => panic!("unexpected routing repository error mapping: {other:?}"), + } + } + #[tokio::test] async fn explicit_routing_attachment_authorizes_and_caches_per_principal() { let repository = Arc::new(InMemoryRoutingGroupRepository::default()); diff --git a/apps/aether-gateway/src/ai_serving/planner/mod.rs b/apps/aether-gateway/src/ai_serving/planner/mod.rs index d3161464c..9f8eb3a97 100644 --- a/apps/aether-gateway/src/ai_serving/planner/mod.rs +++ b/apps/aether-gateway/src/ai_serving/planner/mod.rs @@ -93,6 +93,71 @@ pub(crate) use aether_ai_serving::{ CandidateFailureDiagnostic, CandidateFailureDiagnosticKind, }; +pub(crate) struct ResolvedTunnelSchedulerAffinityContext { + pub(crate) requested_model: String, + pub(crate) client_session_affinity: Option, + pub(crate) policy_context: Option, + pub(crate) routing_overlay: Option, +} + +pub(crate) async fn resolve_tunnel_scheduler_affinity_context( + state: &AppState, + parts: &http::request::Parts, + decision: &GatewayControlDecision, + requested_model: String, + body_json: &serde_json::Value, + client_api_format: &str, +) -> Result, GatewayError> { + let Some(auth_context) = decision.auth_context.as_ref() else { + return Ok(None); + }; + let execution_auth_context = + crate::ai_serving::build_execution_runtime_auth_context(auth_context); + let Some(auth_snapshot) = state + .read_cached_auth_api_key_snapshot( + &execution_auth_context.user_id, + &execution_auth_context.api_key_id, + crate::clock::current_unix_secs(), + ) + .await? + else { + return Ok(None); + }; + let resolved_auth_input = decision_input::ResolvedLocalDecisionAuthInput { + auth_context: execution_auth_context, + auth_snapshot, + required_capabilities: None, + model_directive_policy: decision.model_directive_policy.clone(), + }; + let mut input = decision_input::build_local_requested_model_decision_input( + resolved_auth_input, + requested_model, + ); + decision_input::attach_routing_policy_to_local_requested_model_input( + state, + parts, + &mut input, + body_json, + client_api_format, + ) + .await?; + let policy_context = input + .routing_policy + .as_ref() + .map(crate::scheduler::affinity::SchedulerAffinityPolicyContext::from_routing_policy); + let routing_overlay = input + .routing_policy + .as_ref() + .map(|policy| policy.ranking_overlay.clone()); + + Ok(Some(ResolvedTunnelSchedulerAffinityContext { + requested_model: input.requested_model, + client_session_affinity: input.client_session_affinity, + policy_context, + routing_overlay, + })) +} + pub(crate) async fn maybe_build_sync_decision_payload( state: &AppState, parts: &http::request::Parts, diff --git a/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/payload.rs b/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/payload.rs index 60cffac9f..0d3c578b8 100644 --- a/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/payload.rs +++ b/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/payload.rs @@ -199,6 +199,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_ original_request_body_json, original_request_body_base64: None, client_session_affinity: input.client_session_affinity.as_ref(), + routing_policy: input.routing_policy.as_ref(), scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch, client_requested_stream: body_json .get("stream") diff --git a/apps/aether-gateway/src/ai_serving/planner/report_context.rs b/apps/aether-gateway/src/ai_serving/planner/report_context.rs index 61fa55688..7e3a29343 100644 --- a/apps/aether-gateway/src/ai_serving/planner/report_context.rs +++ b/apps/aether-gateway/src/ai_serving/planner/report_context.rs @@ -6,6 +6,7 @@ use aether_ai_serving::{ provider_stream_event_api_format_for_provider_type as ai_provider_stream_event_api_format_for_provider_type, AiExecutionReportContextParts, AiRequestOrigin, }; +use aether_routing_core::ResolvedRoutingPolicy; use aether_runtime_state::RuntimeLockLease; use aether_scheduler_core::{ClientSessionAffinity, SchedulerRankingOutcome}; use serde_json::{Map, Value}; @@ -20,8 +21,9 @@ use crate::client_session_affinity::{ }; use crate::orchestration::{ insert_pool_key_lease_report_context_fields, ExecutionAttemptIdentity, - SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD, + ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD, SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD, }; +use crate::scheduler::affinity::insert_scheduler_affinity_policy_report_context_field; pub(crate) struct LocalExecutionReportContextParts<'a> { pub(crate) auth_context: &'a ExecutionRuntimeAuthContext, @@ -55,6 +57,7 @@ pub(crate) struct LocalExecutionReportContextParts<'a> { pub(crate) original_request_body_json: Option<&'a Value>, pub(crate) original_request_body_base64: Option<&'a str>, pub(crate) client_session_affinity: Option<&'a ClientSessionAffinity>, + pub(crate) routing_policy: Option<&'a ResolvedRoutingPolicy>, pub(crate) scheduler_affinity_epoch: Option, pub(crate) client_requested_stream: bool, pub(crate) upstream_is_stream: bool, @@ -105,6 +108,16 @@ pub(crate) fn build_local_execution_report_context( merge_incoming_tls_fingerprint(&mut extra_fields, incoming_tls); } insert_pool_key_lease_report_context_fields(&mut extra_fields, parts.pool_key_lease); + insert_scheduler_affinity_policy_report_context_field(&mut extra_fields, parts.routing_policy); + if let Some(override_policy) = parts + .routing_policy + .and_then(|policy| policy.pool_policy_overrides.get(parts.provider_id)) + .filter(|override_policy| !override_policy.scheduling_presets.is_empty()) + { + if let Ok(value) = serde_json::to_value(override_policy) { + extra_fields.insert(ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD.to_string(), value); + } + } if let Some(epoch) = parts.scheduler_affinity_epoch { extra_fields.insert( SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD.to_string(), @@ -315,6 +328,7 @@ mod tests { original_request_body_json: Some(&json!({"model": "gpt-5"})), original_request_body_base64: None, client_session_affinity: Some(&client_session_affinity), + routing_policy: None, scheduler_affinity_epoch: None, client_requested_stream: false, upstream_is_stream: false, @@ -397,6 +411,7 @@ mod tests { })), original_request_body_base64: None, client_session_affinity: None, + routing_policy: None, scheduler_affinity_epoch: None, client_requested_stream: false, upstream_is_stream: true, @@ -463,6 +478,7 @@ mod tests { original_request_body_json: Some(&json!({"model": "gpt-5"})), original_request_body_base64: None, client_session_affinity: None, + routing_policy: None, scheduler_affinity_epoch: None, client_requested_stream: false, upstream_is_stream: false, diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/files/decision.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/files/decision.rs index 53b07a984..8f329cf8b 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/files/decision.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/files/decision.rs @@ -107,6 +107,7 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat original_request_body_json: Some(body_json), original_request_body_base64: resolved.provider_request_body_base64.as_deref(), client_session_affinity: input.client_session_affinity.as_ref(), + routing_policy: input.routing_policy.as_ref(), scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch, client_requested_stream: spec_metadata.require_streaming, upstream_is_stream: spec_metadata.require_streaming, diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/image/decision.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/image/decision.rs index 7c8ba9543..ea13663db 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/image/decision.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/image/decision.rs @@ -120,6 +120,7 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat original_request_body_json: Some(body_json), original_request_body_base64: body_base64, client_session_affinity: input.client_session_affinity.as_ref(), + routing_policy: input.routing_policy.as_ref(), scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch, client_requested_stream: spec_metadata.require_streaming, upstream_is_stream, diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/video/decision.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/video/decision.rs index 39682407a..f252354db 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/video/decision.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/video/decision.rs @@ -88,6 +88,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat original_request_body_json: Some(body_json), original_request_body_base64: None, client_session_affinity: input.client_session_affinity.as_ref(), + routing_policy: input.routing_policy.as_ref(), scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch, client_requested_stream: false, upstream_is_stream: false, diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/family/payload.rs b/apps/aether-gateway/src/ai_serving/planner/standard/family/payload.rs index 6e5ec6ef5..cbb556095 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/family/payload.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/family/payload.rs @@ -140,6 +140,7 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate( original_request_body_json, original_request_body_base64: None, client_session_affinity: input.client_session_affinity.as_ref(), + routing_policy: input.routing_policy.as_ref(), scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch, client_requested_stream: body_json .get("stream") diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/decision/payload.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/decision/payload.rs index 113829caf..4b6644845 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/decision/payload.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/decision/payload.rs @@ -193,6 +193,7 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate original_request_body_json, original_request_body_base64: None, client_session_affinity: input.client_session_affinity.as_ref(), + routing_policy: input.routing_policy.as_ref(), scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch, client_requested_stream: body_json .get("stream") diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/payload.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/payload.rs index 5b318b047..ae0f4f4c3 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/payload.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/payload.rs @@ -142,6 +142,7 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand original_request_body_json, original_request_body_base64: None, client_session_affinity: input.client_session_affinity.as_ref(), + routing_policy: input.routing_policy.as_ref(), scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch, client_requested_stream: body_json .get("stream") diff --git a/apps/aether-gateway/src/cache/candidate_page.rs b/apps/aether-gateway/src/cache/candidate_page.rs index a4ec7e821..6512fae8d 100644 --- a/apps/aether-gateway/src/cache/candidate_page.rs +++ b/apps/aether-gateway/src/cache/candidate_page.rs @@ -72,11 +72,21 @@ struct CandidatePageCacheMetrics { #[derive(Debug, Clone, PartialEq, Eq, Hash)] pub(crate) struct CandidateRowPageCacheKey { api_format: String, - requested_model_name: String, - requested_name: String, + kind: CandidateRowPageCacheKind, offset: u32, limit: u32, - enable_model_directives: bool, +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +enum CandidateRowPageCacheKind { + RequestedModel { + requested_model_name: String, + requested_name: String, + enable_model_directives: bool, + }, + ApiFormatFallback { + scan_epoch: u32, + }, } #[derive(Debug, Clone, PartialEq, Eq, Hash)] @@ -119,11 +129,27 @@ impl CandidateRowPageCacheKey { ) -> Self { Self { api_format: normalize_api_format(api_format), - requested_model_name: normalize_text_key(requested_model_name), - requested_name: normalize_text_key(requested_name), + kind: CandidateRowPageCacheKind::RequestedModel { + requested_model_name: normalize_text_key(requested_model_name), + requested_name: normalize_text_key(requested_name), + enable_model_directives, + }, + offset, + limit, + } + } + + pub(crate) fn for_api_format_fallback( + api_format: &str, + offset: u32, + limit: u32, + scan_epoch: u32, + ) -> Self { + Self { + api_format: normalize_api_format(api_format), + kind: CandidateRowPageCacheKind::ApiFormatFallback { scan_epoch }, offset, limit, - enable_model_directives, } } } diff --git a/apps/aether-gateway/src/control/route/oauth.rs b/apps/aether-gateway/src/control/route/oauth.rs index f2a76717d..ab6279f74 100644 --- a/apps/aether-gateway/src/control/route/oauth.rs +++ b/apps/aether-gateway/src/control/route/oauth.rs @@ -222,6 +222,39 @@ pub(super) fn classify_oauth_route( "admin:provider_oauth", false, )) + } else if method == http::Method::POST + && normalized_path.starts_with("/api/admin/provider-oauth/providers/") + && normalized_path.ends_with("/cookie-authorize/tasks") + { + Some(classified( + "admin_proxy", + "provider_oauth_manage", + "start_cookie_authorize_task", + "admin:provider_oauth", + false, + )) + } else if method == http::Method::GET + && normalized_path.starts_with("/api/admin/provider-oauth/providers/") + && normalized_path.contains("/cookie-authorize/tasks/") + { + Some(classified( + "admin_proxy", + "provider_oauth_manage", + "get_cookie_authorize_task_status", + "admin:provider_oauth", + false, + )) + } else if method == http::Method::POST + && normalized_path.starts_with("/api/admin/provider-oauth/providers/") + && normalized_path.ends_with("/cookie-authorize") + { + Some(classified( + "admin_proxy", + "provider_oauth_manage", + "cookie_authorize", + "admin:provider_oauth", + false, + )) } else if method == http::Method::POST && normalized_path.starts_with("/api/admin/provider-oauth/providers/") && normalized_path.ends_with("/agent-identity-import/tasks") diff --git a/apps/aether-gateway/src/control/tests/admin_oauth.rs b/apps/aether-gateway/src/control/tests/admin_oauth.rs index 9326e5c1e..53d406a80 100644 --- a/apps/aether-gateway/src/control/tests/admin_oauth.rs +++ b/apps/aether-gateway/src/control/tests/admin_oauth.rs @@ -109,6 +109,27 @@ fn classifies_admin_provider_oauth_maintenance_routes_as_admin_proxy_route() { "admin:provider_oauth", "admin:provider_oauth:write", ), + ( + http::Method::POST, + "/api/admin/provider-oauth/providers/provider-123/cookie-authorize", + "cookie_authorize", + "admin:provider_oauth", + "admin:provider_oauth:write", + ), + ( + http::Method::POST, + "/api/admin/provider-oauth/providers/provider-123/cookie-authorize/tasks", + "start_cookie_authorize_task", + "admin:provider_oauth", + "admin:provider_oauth:write", + ), + ( + http::Method::GET, + "/api/admin/provider-oauth/providers/provider-123/cookie-authorize/tasks/claude-cookie-task-123", + "get_cookie_authorize_task_status", + "admin:provider_oauth", + "admin:provider_oauth:read", + ), ( http::Method::POST, "/api/admin/provider-oauth/providers/provider-123/agent-identity-import/tasks", diff --git a/apps/aether-gateway/src/data/candidate_selection.rs b/apps/aether-gateway/src/data/candidate_selection.rs index 003a61661..eb453e67b 100644 --- a/apps/aether-gateway/src/data/candidate_selection.rs +++ b/apps/aether-gateway/src/data/candidate_selection.rs @@ -1,7 +1,7 @@ use aether_data::DataLayerError; use aether_data_contracts::repository::candidate_selection::{ - StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery, - StoredRequestedModelCandidateRowsQuery, + StoredApiFormatCandidateRowsQuery, StoredMinimalCandidateSelectionRow, + StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery, }; use aether_scheduler_core::{ auth_constraints_allow_api_format, collect_global_model_names_for_required_capability, @@ -39,6 +39,19 @@ pub(crate) trait MinimalCandidateSelectionRowSource { api_format: &str, ) -> Result, DataLayerError>; + async fn read_minimal_candidate_selection_rows_for_api_format_page( + &self, + query: &StoredApiFormatCandidateRowsQuery, + ) -> Result, DataLayerError> { + Ok(self + .read_minimal_candidate_selection_rows_for_api_format(&query.api_format) + .await? + .into_iter() + .skip(query.offset as usize) + .take(query.limit as usize) + .collect()) + } + async fn read_pool_key_candidate_rows_for_group( &self, query: &StoredPoolKeyCandidateRowsQuery, @@ -231,6 +244,30 @@ pub(crate) async fn read_requested_model_rows_fast_path_page( }) } +pub(crate) async fn read_api_format_rows_fallback_page( + state: &(impl MinimalCandidateSelectionRowSource + Sync), + api_format: &str, + offset: u32, + limit: u32, +) -> Result { + let limit = limit.max(1); + let rows = state + .read_minimal_candidate_selection_rows_for_api_format_page( + &StoredApiFormatCandidateRowsQuery { + api_format: api_format.to_string(), + offset, + limit, + }, + ) + .await?; + let scanned_rows = rows.len() as u32; + Ok(RequestedModelCandidateRowsPage { + rows, + scanned_rows, + end_of_requested_name: scanned_rows < limit, + }) +} + pub(crate) async fn enumerate_minimal_candidate_selection_with_required_capabilities( state: &(impl MinimalCandidateSelectionRowSource + Sync), api_format: &str, diff --git a/apps/aether-gateway/src/data/state/candidate_cache.rs b/apps/aether-gateway/src/data/state/candidate_cache.rs index 10e8e40e3..cd907cf4b 100644 --- a/apps/aether-gateway/src/data/state/candidate_cache.rs +++ b/apps/aether-gateway/src/data/state/candidate_cache.rs @@ -7,9 +7,10 @@ use std::time::Duration; use aether_cache::ExpiringMap; use aether_data::DataLayerError; use aether_data_contracts::repository::candidate_selection::{ - MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow, - StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery, - StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery, + MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery, + StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder, + StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery, + StoredRequestedModelCandidateRowsQuery, }; use async_trait::async_trait; use tokio::sync::{Notify, OwnedSemaphorePermit, Semaphore}; @@ -389,6 +390,19 @@ impl MinimalCandidateSelectionReadRepository for CachedMinimalCandidateSelection .await } + async fn list_for_exact_api_format_page( + &self, + query: &StoredApiFormatCandidateRowsQuery, + ) -> Result, DataLayerError> { + let key = CandidateSelectionCacheKey::ApiFormatPage { + api_format: normalize_api_format_key(&query.api_format), + offset: query.offset, + limit: query.limit, + }; + self.get_or_load(key, || self.inner.list_for_exact_api_format_page(query)) + .await + } + async fn list_for_exact_api_format_and_global_model( &self, api_format: &str, @@ -480,6 +494,11 @@ enum CandidateSelectionCacheKey { ApiFormat { api_format: String, }, + ApiFormatPage { + api_format: String, + offset: u32, + limit: u32, + }, ApiFormatAndGlobalModel { api_format: String, global_model_name: String, @@ -1044,6 +1063,36 @@ mod tests { assert!(cache.inflight.lock().unwrap().is_empty()); } + #[tokio::test] + async fn candidate_selection_cache_keys_api_format_pages_by_offset() { + let inner = Arc::new(StubCandidateSelectionRepository::new(Duration::ZERO)); + let cache = CachedMinimalCandidateSelectionReadRepository::new(inner.clone()); + let first = StoredApiFormatCandidateRowsQuery { + api_format: "openai:chat".to_string(), + offset: 0, + limit: 256, + }; + let second = StoredApiFormatCandidateRowsQuery { + offset: 256, + ..first.clone() + }; + + cache + .list_for_exact_api_format_page(&first) + .await + .expect("first page should load"); + cache + .list_for_exact_api_format_page(&first) + .await + .expect("first page should be cached"); + cache + .list_for_exact_api_format_page(&second) + .await + .expect("second page should load independently"); + + assert_eq!(inner.calls(), 2); + } + #[tokio::test] async fn candidate_selection_cache_releases_inflight_when_leader_is_cancelled() { let inner = Arc::new(FirstLoadPendingThenFastRepository::new()); diff --git a/apps/aether-gateway/src/data/state/integrations.rs b/apps/aether-gateway/src/data/state/integrations.rs index 6283bf1d1..f0361bad9 100644 --- a/apps/aether-gateway/src/data/state/integrations.rs +++ b/apps/aether-gateway/src/data/state/integrations.rs @@ -176,6 +176,14 @@ impl MinimalCandidateSelectionRowSource for GatewayDataState { .await } + async fn read_minimal_candidate_selection_rows_for_api_format_page( + &self, + query: &aether_data_contracts::repository::candidate_selection::StoredApiFormatCandidateRowsQuery, + ) -> Result, DataLayerError> { + self.list_minimal_candidate_selection_rows_for_api_format_page(query) + .await + } + async fn read_pool_key_candidate_rows_for_group( &self, query: &aether_data_contracts::repository::candidate_selection::StoredPoolKeyCandidateRowsQuery, diff --git a/apps/aether-gateway/src/data/state/mod.rs b/apps/aether-gateway/src/data/state/mod.rs index 39a2c46e9..685639d4f 100644 --- a/apps/aether-gateway/src/data/state/mod.rs +++ b/apps/aether-gateway/src/data/state/mod.rs @@ -97,9 +97,9 @@ use aether_data_contracts::repository::billing::{ UserPlanEntitlementRecord, }; use aether_data_contracts::repository::candidate_selection::{ - MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow, - StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery, - StoredRequestedModelCandidateRowsQuery, + MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery, + StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsByKeyIdsQuery, + StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery, }; use aether_data_contracts::repository::candidates::{ PublicHealthStatusCount, PublicHealthTimelineBucket, RequestCandidateReadRepository, diff --git a/apps/aether-gateway/src/data/state/models.rs b/apps/aether-gateway/src/data/state/models.rs index 7bc99e32c..f579d4a90 100644 --- a/apps/aether-gateway/src/data/state/models.rs +++ b/apps/aether-gateway/src/data/state/models.rs @@ -2,11 +2,12 @@ use super::{ AdminGlobalModelListQuery, AdminProviderModelListQuery, CreateAdminGlobalModelRecord, DataLayerError, GatewayDataState, PublicCatalogModelListQuery, PublicCatalogModelSearchQuery, PublicGlobalModelQuery, StoredAdminGlobalModel, StoredAdminGlobalModelPage, - StoredAdminProviderModel, StoredMinimalCandidateSelectionRow, - StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery, - StoredProviderActiveGlobalModel, StoredProviderModelStats, StoredPublicCatalogModel, - StoredPublicGlobalModel, StoredPublicGlobalModelPage, StoredRequestedModelCandidateRowsQuery, - UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord, + StoredAdminProviderModel, StoredApiFormatCandidateRowsQuery, + StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsByKeyIdsQuery, + StoredPoolKeyCandidateRowsQuery, StoredProviderActiveGlobalModel, StoredProviderModelStats, + StoredPublicCatalogModel, StoredPublicGlobalModel, StoredPublicGlobalModelPage, + StoredRequestedModelCandidateRowsQuery, UpdateAdminGlobalModelRecord, + UpsertAdminProviderModelRecord, }; impl GatewayDataState { @@ -98,6 +99,23 @@ impl GatewayDataState { .await } + pub(crate) async fn list_minimal_candidate_selection_rows_for_api_format_page( + &self, + query: &StoredApiFormatCandidateRowsQuery, + ) -> Result, DataLayerError> { + crate::request_diagnostics::observe_db_operation( + "candidate_selection", + self.database_pool_summary(), + async { + match &self.minimal_candidate_selection_reader { + Some(repository) => repository.list_for_exact_api_format_page(query).await, + None => Ok(Vec::new()), + } + }, + ) + .await + } + pub(crate) async fn list_pool_key_candidate_rows_for_group( &self, query: &StoredPoolKeyCandidateRowsQuery, diff --git a/apps/aether-gateway/src/dispatch/pool_scheduler.rs b/apps/aether-gateway/src/dispatch/pool_scheduler.rs index f4eecd995..8fe81fce5 100644 --- a/apps/aether-gateway/src/dispatch/pool_scheduler.rs +++ b/apps/aether-gateway/src/dispatch/pool_scheduler.rs @@ -42,7 +42,7 @@ use crate::handlers::shared::provider_pool::{ use crate::handlers::shared::provider_pool::{ admin_provider_pool_quota_probe_active_members_key, read_admin_provider_pool_key_cooldown_reason, AdminProviderPoolConfig, - AdminProviderPoolRuntimeState, + AdminProviderPoolRuntimeState, AdminProviderPoolSchedulingPreset, }; use crate::handlers::shared::{parse_catalog_auth_config_json, provider_key_health_summary}; use crate::maintenance::spawn_pool_quota_probe_replenish_for_request; @@ -104,6 +104,7 @@ async fn schedule_pool_page_candidates( state: PlannerAppState<'_>, candidates: Vec, sticky_session_token: Option<&str>, + effective_pool_config: Option<&AdminProviderPoolConfig>, ) -> ( Vec, Vec, @@ -115,7 +116,10 @@ async fn schedule_pool_page_candidates( let mut provider_runtime_requirements = BTreeMap::)>::new(); for candidate in &candidates { - let Some(pool_config) = pool_config_for_candidate(candidate) else { + let Some(pool_config) = effective_pool_config + .cloned() + .or_else(|| pool_config_for_candidate(candidate)) + else { continue; }; let entry = provider_runtime_requirements @@ -170,10 +174,16 @@ async fn schedule_pool_page_candidates( } } - let outcome = apply_local_execution_pool_scheduler_with_runtime_map_outcome( + let effective_pool_config_by_provider = effective_pool_config + .map(|config| { + BTreeMap::from([(candidates[0].candidate.provider_id.clone(), config.clone())]) + }) + .unwrap_or_default(); + let outcome = apply_local_execution_pool_scheduler_with_runtime_map_outcome_and_configs( candidates, &runtime_by_provider, &key_context_by_id, + &effective_pool_config_by_provider, ); let scheduled = outcome.candidates; let skipped = outcome.skipped; @@ -347,6 +357,8 @@ pub(crate) struct PoolKeyCursor<'a> { requested_model: Option, request_auth_channel: Option, routing_overlay: Option, + routing_allowed_key_ids: Option>, + effective_pool_config: Option, runtime_miss_trace_id: Option, record_runtime_miss_diagnostic: bool, pool_key_order: StoredPoolKeyCandidateOrder, @@ -358,7 +370,11 @@ pub(crate) struct PoolKeyCursor<'a> { max_scanned_keys: u32, absolute_max_scanned_keys: u32, score_top_n: u32, - score_phase_loaded: bool, + score_next_offset: u32, + score_phase_exhausted: bool, + score_schedule_interest_count: usize, + routing_allowed_key_offset: usize, + routing_allowed_rows: Option>, skip_reason_counts: BTreeMap<&'static str, u32>, next_pool_key_index: u32, sticky_candidate_loaded: bool, @@ -404,15 +420,27 @@ impl<'a> PoolKeyCursor<'a> { request_auth_channel: Option<&str>, routing_policy: Option<&ResolvedRoutingPolicy>, ) -> Self { - let pool_key_order = pool_key_candidate_order_for_group(&group, routing_policy); + let effective_pool_config = effective_pool_config_for_group(&group, routing_policy); + let pool_key_order = + pool_key_candidate_order_for_group(&group, effective_pool_config.as_ref()); let routing_overlay = routing_policy.map(|policy| policy.ranking_overlay.clone()); - let pool_config = pool_config_for_candidate(&group); - let score_top_n = pool_config + let routing_allowed_key_ids = routing_policy + .map(|policy| &policy.ranking_overlay.allowed_keys) + .filter(|key_ids| !key_ids.is_empty()) + .map(|key_ids| { + let mut seen = BTreeSet::new(); + key_ids + .iter() + .filter(|key_id| seen.insert((*key_id).clone())) + .cloned() + .collect::>() + }); + let score_top_n = effective_pool_config .as_ref() .map(|config| config.score_top_n) .unwrap_or(u64::from(aether_dispatch_core::DEFAULT_POOL_PAGE_SIZE)) .clamp(1, u64::from(u32::MAX)) as u32; - let configured_max_scanned_keys = pool_config + let configured_max_scanned_keys = effective_pool_config .as_ref() .map(|config| config.score_fallback_scan_limit) .unwrap_or(u64::from(aether_dispatch_core::DEFAULT_POOL_MAX_SCAN)) @@ -427,6 +455,8 @@ impl<'a> PoolKeyCursor<'a> { requested_model: requested_model.map(str::to_string), request_auth_channel: request_auth_channel.map(str::to_string), routing_overlay, + routing_allowed_key_ids, + effective_pool_config, runtime_miss_trace_id: None, record_runtime_miss_diagnostic: false, pool_key_order, @@ -438,7 +468,11 @@ impl<'a> PoolKeyCursor<'a> { max_scanned_keys: max_scanned_keys.max(window_config.window_size), absolute_max_scanned_keys: absolute_max_scanned_keys.max(window_config.window_size), score_top_n, - score_phase_loaded: false, + score_next_offset: 0, + score_phase_exhausted: false, + score_schedule_interest_count: 0, + routing_allowed_key_offset: 0, + routing_allowed_rows: None, skip_reason_counts: BTreeMap::new(), next_pool_key_index: 0, sticky_candidate_loaded: false, @@ -582,8 +616,11 @@ impl<'a> PoolKeyCursor<'a> { } async fn next_page_candidates(&mut self) -> Option> { - if !self.score_phase_loaded { - self.score_phase_loaded = true; + if self.routing_allowed_key_ids.is_some() { + return self.next_routing_allowed_candidates().await; + } + + if !self.score_phase_exhausted { if let Some(score_candidates) = self.next_score_candidates().await { return Some(score_candidates); } @@ -636,20 +673,177 @@ impl<'a> PoolKeyCursor<'a> { return None; } - self.scanned_keys += rows.len() as u32; - self.budget_scanned_keys += rows.len() as u32; - self.next_offset = self.next_offset.saturating_add(rows.len() as u32); - Some(self.build_page_eligible_candidates(rows).await) + let row_count = u32::try_from(rows.len()).unwrap_or(u32::MAX); + self.next_offset = self.next_offset.saturating_add(row_count); + let seen_count_before = self.seen_key_ids.len(); + let candidates = self.build_page_eligible_candidates(rows).await; + let distinct_row_count = self.seen_key_ids.len().saturating_sub(seen_count_before); + self.scanned_keys = self.scanned_keys.saturating_add(row_count); + self.budget_scanned_keys = self + .budget_scanned_keys + .saturating_add(u32::try_from(distinct_row_count).unwrap_or(u32::MAX)); + Some(candidates) + } + + async fn next_routing_allowed_candidates( + &mut self, + ) -> Option> { + loop { + if self.budget_scanned_keys >= self.max_scanned_keys + || self.scanned_keys >= self.absolute_max_scanned_keys + { + return None; + } + + let key_ids = self.routing_allowed_key_ids.as_ref()?; + let has_buffered_rows = self + .routing_allowed_rows + .as_ref() + .is_some_and(|rows| !rows.is_empty()); + if !has_buffered_rows && self.routing_allowed_key_offset >= key_ids.len() { + return None; + } + let limit = self + .page_size + .min(self.max_scanned_keys - self.budget_scanned_keys) + .min(self.absolute_max_scanned_keys - self.scanned_keys) + as usize; + if matches!( + self.pool_key_order, + StoredPoolKeyCandidateOrder::InternalPriority + ) && self.routing_allowed_rows.is_none() + { + let query = StoredPoolKeyCandidateRowsByKeyIdsQuery { + api_format: self.group.candidate.endpoint_api_format.clone(), + provider_id: self.group.candidate.provider_id.clone(), + endpoint_id: self.group.candidate.endpoint_id.clone(), + model_id: self.group.candidate.model_id.clone(), + selected_provider_model_name: self + .group + .candidate + .selected_provider_model_name + .clone(), + key_ids: key_ids.clone(), + }; + let mut rows = match self + .state + .app() + .list_pool_key_candidate_rows_for_group_key_ids(&query) + .await + { + Ok(rows) => rows, + Err(err) => { + warn!( + event_name = "pool_group_routing_key_load_failed", + log_type = "event", + provider_id = %self.group.candidate.provider_id, + endpoint_id = %self.group.candidate.endpoint_id, + model_id = %self.group.candidate.model_id, + allowed_key_count = query.key_ids.len(), + error = ?err, + "gateway pool scheduler failed to materialize routing-allowed pool keys" + ); + return None; + } + }; + rows.sort_by(|left, right| { + let left_priority = self + .routing_overlay + .as_ref() + .map_or(left.key_internal_priority, |overlay| { + overlay.key_priority(&left.key_id, left.key_internal_priority) + }); + let right_priority = self + .routing_overlay + .as_ref() + .map_or(right.key_internal_priority, |overlay| { + overlay.key_priority(&right.key_id, right.key_internal_priority) + }); + left_priority + .cmp(&right_priority) + .then(left.key_id.cmp(&right.key_id)) + }); + self.routing_allowed_key_offset = key_ids.len(); + self.routing_allowed_rows = Some(rows.into()); + } + + let rows = if let Some(rows) = self.routing_allowed_rows.as_mut() { + let page_len = limit.min(rows.len()); + rows.drain(..page_len).collect::>() + } else { + let end = self + .routing_allowed_key_offset + .saturating_add(limit) + .min(key_ids.len()); + let page_key_ids = key_ids[self.routing_allowed_key_offset..end].to_vec(); + self.routing_allowed_key_offset = end; + let query = StoredPoolKeyCandidateRowsByKeyIdsQuery { + api_format: self.group.candidate.endpoint_api_format.clone(), + provider_id: self.group.candidate.provider_id.clone(), + endpoint_id: self.group.candidate.endpoint_id.clone(), + model_id: self.group.candidate.model_id.clone(), + selected_provider_model_name: self + .group + .candidate + .selected_provider_model_name + .clone(), + key_ids: page_key_ids, + }; + match self + .state + .app() + .list_pool_key_candidate_rows_for_group_key_ids(&query) + .await + { + Ok(rows) => rows, + Err(err) => { + warn!( + event_name = "pool_group_routing_key_load_failed", + log_type = "event", + provider_id = %self.group.candidate.provider_id, + endpoint_id = %self.group.candidate.endpoint_id, + model_id = %self.group.candidate.model_id, + allowed_key_count = query.key_ids.len(), + error = ?err, + "gateway pool scheduler failed to materialize routing-allowed pool keys" + ); + return None; + } + } + }; + if rows.is_empty() { + if self.routing_allowed_rows.is_some() { + return None; + } + continue; + } + + let row_count = u32::try_from(rows.len()).unwrap_or(u32::MAX); + let seen_count_before = self.seen_key_ids.len(); + let candidates = self.build_page_eligible_candidates(rows).await; + let distinct_row_count = self.seen_key_ids.len().saturating_sub(seen_count_before); + self.scanned_keys = self.scanned_keys.saturating_add(row_count); + self.budget_scanned_keys = self + .budget_scanned_keys + .saturating_add(u32::try_from(distinct_row_count).unwrap_or(u32::MAX)); + return Some(candidates); + } } async fn next_score_candidates(&mut self) -> Option> { - if self.scanned_keys >= self.absolute_max_scanned_keys { + if self.score_phase_exhausted + || self.score_next_offset >= self.score_top_n + || self.scanned_keys >= self.absolute_max_scanned_keys + { + self.score_phase_exhausted = true; return None; } let limit = self - .score_top_n + .page_size + .min(self.score_top_n - self.score_next_offset) .min(self.absolute_max_scanned_keys - self.scanned_keys); if limit == 0 { + self.score_phase_exhausted = true; return None; } let scope = provider_key_pool_score_scope(); @@ -661,7 +855,7 @@ impl<'a> PoolKeyCursor<'a> { scope_id: scope.scope_id.clone(), hard_states: vec![PoolMemberHardState::Available, PoolMemberHardState::Unknown], probe_statuses: None, - offset: 0, + offset: self.score_next_offset as usize, limit: limit as usize, }; let score_started_at = std::time::Instant::now(); @@ -682,6 +876,7 @@ impl<'a> PoolKeyCursor<'a> { error = ?err, "gateway pool scheduler failed to read ranked pool member scores" ); + self.score_phase_exhausted = true; return None; } }; @@ -690,8 +885,14 @@ impl<'a> PoolKeyCursor<'a> { score_started_at.elapsed().as_millis() as u64, ); if scores.is_empty() { + self.score_phase_exhausted = true; return None; } + let score_count = u32::try_from(scores.len()).unwrap_or(u32::MAX); + self.score_next_offset = self.score_next_offset.saturating_add(score_count); + if score_count < limit || self.score_next_offset >= self.score_top_n { + self.score_phase_exhausted = true; + } self.spawn_score_schedule_interest_recording(&scores); @@ -731,6 +932,7 @@ impl<'a> PoolKeyCursor<'a> { error = ?err, "gateway pool scheduler failed to materialize ranked pool keys" ); + self.score_phase_exhausted = true; return None; } }; @@ -738,7 +940,7 @@ impl<'a> PoolKeyCursor<'a> { "pool_score_key_rows", rows_started_at.elapsed().as_millis() as u64, ); - let materialized_row_count = rows.len() as u32; + let materialized_row_count = u32::try_from(rows.len()).unwrap_or(u32::MAX); let missing_score_count = scores.len().saturating_sub(rows.len()); if missing_score_count > 0 { *self @@ -746,15 +948,18 @@ impl<'a> PoolKeyCursor<'a> { .entry("pool_score_member_missing") .or_insert(0) += u32::try_from(missing_score_count).unwrap_or(u32::MAX); } + let seen_count_before = self.seen_key_ids.len(); + let candidates = self.build_page_eligible_candidates(rows).await; + let distinct_row_count = self.seen_key_ids.len().saturating_sub(seen_count_before); self.scanned_keys = self.scanned_keys.saturating_add(materialized_row_count); self.budget_scanned_keys = self .budget_scanned_keys - .saturating_add(materialized_row_count); - Some(self.build_page_eligible_candidates(rows).await) + .saturating_add(u32::try_from(distinct_row_count).unwrap_or(u32::MAX)); + Some(candidates) } async fn sticky_candidate(&mut self) -> Option { - let pool_config = pool_config_for_candidate(&self.group)?; + let pool_config = self.effective_pool_config.as_ref()?.clone(); if !admin_provider_pool_cache_affinity_enabled(&pool_config) { return None; } @@ -767,6 +972,13 @@ impl<'a> PoolKeyCursor<'a> { ) .await; let sticky_key_id = runtime.sticky_bound_key_id?; + if self + .routing_overlay + .as_ref() + .is_some_and(|overlay| !overlay.key_allowed(sticky_key_id.as_str())) + { + return None; + } if self.seen_key_ids.contains(&sticky_key_id) { return None; } @@ -821,6 +1033,7 @@ impl<'a> PoolKeyCursor<'a> { self.state, candidates, self.sticky_session_token.as_deref(), + self.effective_pool_config.as_ref(), ) .await; self.record_skipped_candidates(&skipped); @@ -945,11 +1158,16 @@ impl<'a> PoolKeyCursor<'a> { ); } - fn spawn_score_schedule_interest_recording(&self, scores: &[StoredPoolMemberScore]) { - if scores.is_empty() || !self.state.app().data.has_pool_score_writer() { + fn spawn_score_schedule_interest_recording(&mut self, scores: &[StoredPoolMemberScore]) { + if scores.is_empty() + || self.score_schedule_interest_count >= POOL_SCORE_SCHEDULE_INTEREST_MAX_PER_BATCH + || !self.state.app().data.has_pool_score_writer() + { return; } + let remaining_interest_budget = POOL_SCORE_SCHEDULE_INTEREST_MAX_PER_BATCH + .saturating_sub(self.score_schedule_interest_count); let scheduled_at = current_unix_ms() / 1000; let provider_id = self.group.candidate.provider_id.clone(); let endpoint_id = self.group.candidate.endpoint_id.clone(); @@ -962,7 +1180,7 @@ impl<'a> PoolKeyCursor<'a> { >= POOL_SCORE_SCHEDULE_INTEREST_MIN_INTERVAL_SECS }) }) - .take(POOL_SCORE_SCHEDULE_INTEREST_MAX_PER_BATCH) + .take(remaining_interest_budget) .map(|score| PoolMemberScheduleFeedback { identity: PoolMemberIdentity { pool_kind: score.pool_kind.clone(), @@ -1008,6 +1226,9 @@ impl<'a> PoolKeyCursor<'a> { ); return; }; + self.score_schedule_interest_count = self + .score_schedule_interest_count + .saturating_add(score_count); let app = self.state.app().clone(); @@ -1056,6 +1277,14 @@ impl<'a> PoolKeyCursor<'a> { &mut self, candidate: aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate, ) -> Option { + if self + .routing_overlay + .as_ref() + .is_some_and(|overlay| !overlay.key_allowed(candidate.key_id.as_str())) + { + self.record_skip_reason(ROUTING_PROFILE_DISALLOWED_KEY_SKIP_REASON); + return None; + } if !self.seen_key_ids.insert(candidate.key_id.clone()) { return None; } @@ -1338,11 +1567,26 @@ fn apply_local_execution_pool_scheduler_with_runtime_map_outcome( candidates: Vec, runtime_by_provider: &BTreeMap, key_context_by_id: &BTreeMap, +) -> PoolSchedulerApplyOutcome { + apply_local_execution_pool_scheduler_with_runtime_map_outcome_and_configs( + candidates, + runtime_by_provider, + key_context_by_id, + &BTreeMap::new(), + ) +} + +fn apply_local_execution_pool_scheduler_with_runtime_map_outcome_and_configs( + candidates: Vec, + runtime_by_provider: &BTreeMap, + key_context_by_id: &BTreeMap, + effective_pool_config_by_provider: &BTreeMap, ) -> PoolSchedulerApplyOutcome { let (scheduled, skipped) = run_local_execution_pool_scheduler_with_runtime_map( candidates.clone(), runtime_by_provider, key_context_by_id, + effective_pool_config_by_provider, true, ); let mut active_probe_evicted_members_by_provider = @@ -1370,6 +1614,7 @@ fn apply_local_execution_pool_scheduler_with_runtime_map_outcome( candidates, runtime_by_provider, key_context_by_id, + effective_pool_config_by_provider, false, ); merge_active_probe_evictions( @@ -1433,6 +1678,7 @@ fn run_local_execution_pool_scheduler_with_runtime_map( candidates: Vec, runtime_by_provider: &BTreeMap, key_context_by_id: &BTreeMap, + effective_pool_config_by_provider: &BTreeMap, enforce_active_probe_seal: bool, ) -> ( Vec, @@ -1449,7 +1695,10 @@ fn run_local_execution_pool_scheduler_with_runtime_map( .get(&candidate.candidate.key_id) .cloned() .unwrap_or_default(); - let admin_pool_config = pool_config_for_candidate(&candidate); + let admin_pool_config = effective_pool_config_by_provider + .get(&candidate.candidate.provider_id) + .cloned() + .or_else(|| pool_config_for_candidate(&candidate)); if let Some(config) = admin_pool_config.as_ref() { if enforce_active_probe_seal && should_enforce_active_probe_sealed_pool(config) { @@ -1512,6 +1761,36 @@ fn pool_config_for_candidate( admin_provider_pool_config_from_config_value(candidate.transport.provider.config.as_ref()) } +fn effective_pool_config_for_group( + group: &EligibleLocalExecutionCandidate, + routing_policy: Option<&ResolvedRoutingPolicy>, +) -> Option { + let mut pool_config = pool_config_for_candidate(group)?; + let override_policy = routing_policy + .and_then(|policy| { + policy + .pool_policy_overrides + .get(group.candidate.provider_id.as_str()) + }) + .filter(|override_policy| !override_policy.scheduling_presets.is_empty()); + if let Some(override_policy) = override_policy { + let scheduling_presets = override_policy + .scheduling_presets + .iter() + .map(|preset| AdminProviderPoolSchedulingPreset { + preset: preset.preset.clone(), + enabled: preset.enabled, + mode: preset.mode.clone(), + }) + .collect::>(); + pool_config.lru_enabled = scheduling_presets + .iter() + .any(|preset| preset.enabled && preset.preset.eq_ignore_ascii_case("lru")); + pool_config.scheduling_presets = scheduling_presets; + } + Some(pool_config) +} + fn should_enforce_active_probe_sealed_pool(pool_config: &AdminProviderPoolConfig) -> bool { pool_config.probing_enabled } @@ -1532,38 +1811,20 @@ fn should_trigger_active_probe_burst_for_request( fn pool_key_candidate_order_for_group( group: &EligibleLocalExecutionCandidate, - routing_policy: Option<&ResolvedRoutingPolicy>, + pool_config: Option<&AdminProviderPoolConfig>, ) -> StoredPoolKeyCandidateOrder { - let Some(pool_config) = pool_config_for_candidate(group) else { + let Some(pool_config) = pool_config else { return StoredPoolKeyCandidateOrder::InternalPriority; }; - let override_presets = routing_policy - .and_then(|policy| { - policy - .pool_policy_overrides - .get(group.candidate.provider_id.as_str()) + let presets = pool_config + .scheduling_presets + .iter() + .map(|preset| PoolSchedulingPreset { + preset: preset.preset.clone(), + enabled: preset.enabled, + mode: preset.mode.clone(), }) - .filter(|override_policy| !override_policy.scheduling_presets.is_empty()); - let presets = match override_presets { - Some(override_policy) => override_policy - .scheduling_presets - .iter() - .map(|preset| PoolSchedulingPreset { - preset: preset.preset.clone(), - enabled: preset.enabled, - mode: preset.mode.clone(), - }) - .collect::>(), - None => pool_config - .scheduling_presets - .iter() - .map(|preset| PoolSchedulingPreset { - preset: preset.preset.clone(), - enabled: preset.enabled, - mode: preset.mode.clone(), - }) - .collect::>(), - }; + .collect::>(); let active_presets = ProviderPoolService::with_builtin_adapters() .normalize_scheduling_presets(group.transport.provider.provider_type.as_str(), &presets) .into_iter() @@ -1675,8 +1936,9 @@ mod tests { admin_provider_pool_quota_probe_active_members_key, apply_local_execution_pool_scheduler, apply_local_execution_pool_scheduler_with_runtime_map, apply_local_execution_pool_scheduler_with_runtime_map_outcome, - build_pool_catalog_key_context, pool_config_for_candidate, - pool_key_requires_reauth_for_scheduling, + apply_local_execution_pool_scheduler_with_runtime_map_outcome_and_configs, + build_pool_catalog_key_context, effective_pool_config_for_group, pool_config_for_candidate, + pool_key_candidate_order_for_group, pool_key_requires_reauth_for_scheduling, prune_unschedulable_active_probe_members_for_request, remove_active_probe_members_for_request, should_trigger_active_probe_burst_for_request, PoolCatalogKeyContext, PoolKeyCursor, POOL_ACTIVE_PROBE_SEALED_SKIP_REASON, @@ -1689,7 +1951,8 @@ mod tests { }; use crate::data::GatewayDataState; use crate::handlers::shared::provider_pool::{ - record_admin_provider_pool_error, AdminProviderPoolRuntimeState, + admin_provider_pool_cache_affinity_enabled, record_admin_provider_pool_error, + AdminProviderPoolRuntimeState, }; use crate::orchestration::LocalExecutionCandidateMetadata; use crate::{AppState, LocalExecutionRuntimeMissDiagnostic}; @@ -1712,7 +1975,8 @@ mod tests { GatewayProviderTransportProvider, }; use aether_routing_core::{ - RankingOverlay, ResolvedRoutingPolicy, RoutingSchedulingMode, RoutingSetPriorityMode, + RankingOverlay, ResolvedRoutingPolicy, RoutingPoolPolicyOverride, RoutingSchedulingMode, + RoutingSchedulingPreset, RoutingSetPriorityMode, }; use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate; use serde_json::json; @@ -2555,7 +2819,7 @@ mod tests { .iter() .map(|item| item.candidate.key_id.as_str()) .collect::>(), - vec!["key-plus", "key-pro", "key-team"] + vec!["key-pro", "key-plus", "key-team"] ); } @@ -2755,6 +3019,145 @@ mod tests { assert_eq!(lru_cursor.pool_key_order, StoredPoolKeyCandidateOrder::Lru); } + #[test] + fn routing_pool_override_is_effective_for_page_scheduling_and_sticky_mode() { + let base_config = Some(json!({ + "pool_advanced": { + "scheduling_presets": [{"preset": "plus_first", "enabled": true}] + } + })); + let key_plus = sample_eligible_candidate( + "provider-pool", + "endpoint-1", + "key-plus", + 10, + base_config.clone(), + ); + let key_pro = + sample_eligible_candidate("provider-pool", "endpoint-1", "key-pro", 10, base_config); + let mut routing_policy = routing_policy_with_allowed_keys([]); + routing_policy.pool_policy_overrides.insert( + "provider-pool".to_string(), + RoutingPoolPolicyOverride { + scheduling_presets: vec![RoutingSchedulingPreset { + preset: "pro_first".to_string(), + enabled: true, + mode: Some("pro_only".to_string()), + }], + }, + ); + let effective_config = effective_pool_config_for_group(&key_plus, Some(&routing_policy)) + .expect("effective pool config should parse"); + let effective_configs = BTreeMap::from([("provider-pool".to_string(), effective_config)]); + let key_context_by_id = BTreeMap::from([ + ( + "key-plus".to_string(), + PoolCatalogKeyContext { + plan_tier: Some("plus".to_string()), + ..PoolCatalogKeyContext::default() + }, + ), + ( + "key-pro".to_string(), + PoolCatalogKeyContext { + plan_tier: Some("pro".to_string()), + ..PoolCatalogKeyContext::default() + }, + ), + ]); + + let outcome = apply_local_execution_pool_scheduler_with_runtime_map_outcome_and_configs( + vec![key_plus, key_pro], + &BTreeMap::new(), + &key_context_by_id, + &effective_configs, + ); + + assert!(outcome.skipped.is_empty()); + assert_eq!( + outcome + .candidates + .iter() + .map(|candidate| candidate.candidate.key_id.as_str()) + .collect::>(), + ["key-pro", "key-plus"] + ); + + let app = AppState::new().expect("state should build"); + let cache_affinity_group = sample_eligible_candidate( + "provider-pool", + "endpoint-1", + "pool-group", + 10, + Some(json!({ + "pool_advanced": { + "scheduling_presets": [{"preset": "load_balance", "enabled": true}] + } + })), + ); + routing_policy.pool_policy_overrides.insert( + "provider-pool".to_string(), + RoutingPoolPolicyOverride { + scheduling_presets: vec![RoutingSchedulingPreset { + preset: "cache_affinity".to_string(), + enabled: true, + mode: None, + }], + }, + ); + let cursor = PoolKeyCursor::new_with_routing_policy( + PlannerAppState::new(&app), + cache_affinity_group, + None, + None, + None, + Some(&routing_policy), + ); + assert!(admin_provider_pool_cache_affinity_enabled( + cursor + .effective_pool_config + .as_ref() + .expect("effective config should exist") + )); + } + + #[test] + fn routing_pool_lru_override_recomputes_derived_lru_state() { + let group = sample_eligible_candidate( + "provider-pool", + "endpoint-1", + "pool-group", + 10, + Some(json!({ + "pool_advanced": { + "scheduling_presets": [ + {"preset": "cache_affinity", "enabled": true} + ] + } + })), + ); + let mut routing_policy = routing_policy_with_allowed_keys([]); + routing_policy.pool_policy_overrides.insert( + "provider-pool".to_string(), + RoutingPoolPolicyOverride { + scheduling_presets: vec![RoutingSchedulingPreset { + preset: "lru".to_string(), + enabled: true, + mode: None, + }], + }, + ); + + let effective_config = effective_pool_config_for_group(&group, Some(&routing_policy)) + .expect("effective pool config should parse"); + + assert!(effective_config.lru_enabled); + assert_eq!( + pool_key_candidate_order_for_group(&group, Some(&effective_config)), + StoredPoolKeyCandidateOrder::Lru + ); + } + #[test] fn pool_key_cursor_records_runtime_miss_when_exhausted_without_returning_key() { let app = AppState::new().expect("state should build"); @@ -2918,6 +3321,179 @@ mod tests { ); } + #[tokio::test] + async fn routing_allowed_key_outside_score_top_n_is_materialized_directly() { + let provider_config = Some(json!({ + "pool_advanced": { + "score_top_n": 16, + "score_fallback_scan_limit": 16 + } + })); + let (provider, endpoint, keys, rows) = large_pool_fixture(32, provider_config.clone()); + let scores = (0..16) + .map(|index| { + sample_provider_key_pool_score( + "provider-pool", + &format!("key-{index:05}"), + 1_000.0 - index as f64, + ) + }) + .collect::>(); + let data_state = + GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests( + Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![endpoint], + keys, + )), + Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)), + ) + .with_pool_score_repository_for_tests(Arc::new( + InMemoryPoolMemberScoreRepository::seed(scores), + )) + .with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY); + let app = AppState::new() + .expect("state should build") + .with_data_state_for_tests(data_state); + let group = sample_eligible_candidate( + "provider-pool", + "endpoint-1", + "pool-group", + 10, + provider_config, + ); + let routing_policy = routing_policy_with_allowed_keys(["key-00031"]); + let mut cursor = PoolKeyCursor::new_with_routing_policy( + PlannerAppState::new(&app), + group, + None, + None, + None, + Some(&routing_policy), + ); + + let candidate = cursor + .next_key() + .await + .expect("the only routing-allowed key should not be lost to score preselection"); + + assert_eq!(candidate.candidate.key_id, "key-00031"); + assert_eq!(cursor.scanned_keys, 1); + assert_eq!(cursor.budget_scanned_keys, 1); + } + + #[tokio::test] + async fn routing_allowed_keys_continue_across_pool_pages() { + let provider_config = Some(json!({ + "pool_advanced": { + "scheduling_presets": [ + {"preset": "lru", "enabled": false} + ], + "score_top_n": 16, + "score_fallback_scan_limit": 128 + } + })); + let (provider, endpoint, keys, rows) = large_pool_fixture(80, provider_config.clone()); + let data_state = + GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests( + Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![endpoint], + keys, + )), + Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)), + ) + .with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY); + let app = AppState::new() + .expect("state should build") + .with_data_state_for_tests(data_state); + let group = sample_eligible_candidate( + "provider-pool", + "endpoint-1", + "pool-group", + 10, + provider_config, + ); + let mut routing_policy = routing_policy_with_allowed_keys([]); + routing_policy.ranking_overlay.allowed_keys = + (0..80).map(|index| format!("key-{index:05}")).collect(); + let mut cursor = PoolKeyCursor::new_with_routing_policy( + PlannerAppState::new(&app), + group, + None, + None, + None, + Some(&routing_policy), + ); + + let mut returned_key_ids = Vec::new(); + while let Some(candidate) = cursor.next_key().await { + returned_key_ids.push(candidate.candidate.key_id); + } + + assert_eq!(returned_key_ids.len(), 32); + assert!(returned_key_ids.iter().any(|key_id| key_id == "key-00064")); + assert!(returned_key_ids.iter().any(|key_id| key_id == "key-00079")); + assert_eq!(cursor.scanned_keys, 80); + assert_eq!(cursor.budget_scanned_keys, 80); + assert_eq!(cursor.routing_allowed_key_offset, 80); + } + + #[tokio::test] + async fn routing_allowed_keys_preserve_internal_priority_without_active_presets() { + let provider_config = Some(json!({ + "pool_advanced": { + "scheduling_presets": [ + {"preset": "lru", "enabled": false} + ], + "score_top_n": 2, + "score_fallback_scan_limit": 2 + } + })); + let (provider, endpoint, mut keys, mut rows) = + large_pool_fixture(2, provider_config.clone()); + keys[0].internal_priority = 100; + keys[1].internal_priority = 1; + rows[0].key_internal_priority = 100; + rows[1].key_internal_priority = 1; + let data_state = + GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests( + Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![endpoint], + keys, + )), + Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)), + ) + .with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY); + let app = AppState::new() + .expect("state should build") + .with_data_state_for_tests(data_state); + let group = sample_eligible_candidate( + "provider-pool", + "endpoint-1", + "pool-group", + 10, + provider_config, + ); + let routing_policy = routing_policy_with_allowed_keys(["key-00000", "key-00001"]); + let mut cursor = PoolKeyCursor::new_with_routing_policy( + PlannerAppState::new(&app), + group, + None, + None, + None, + Some(&routing_policy), + ); + + let candidate = cursor + .next_key() + .await + .expect("an allowed key should be schedulable"); + + assert_eq!(candidate.candidate.key_id, "key-00001"); + } + #[tokio::test] async fn pool_key_cursor_allows_parallel_requests_to_use_same_healthy_key() { let app = AppState::new().expect("state should build"); @@ -3314,6 +3890,123 @@ mod tests { ); } + #[tokio::test] + async fn score_candidates_continue_across_pool_windows() { + let provider_config = Some(json!({ + "pool_advanced": { + "score_top_n": 128, + "score_fallback_scan_limit": 128 + } + })); + let (provider, endpoint, keys, rows) = large_pool_fixture(128, provider_config.clone()); + let scores = (0..128) + .map(|index| { + sample_provider_key_pool_score( + "provider-pool", + &format!("key-{index:05}"), + 1_000.0 - index as f64, + ) + }) + .collect::>(); + let data_state = + GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests( + Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![endpoint], + keys, + )), + Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)), + ) + .with_pool_score_repository_for_tests(Arc::new( + InMemoryPoolMemberScoreRepository::seed(scores), + )) + .with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY); + let app = AppState::new() + .expect("state should build") + .with_data_state_for_tests(data_state); + let group = sample_eligible_candidate( + "provider-pool", + "endpoint-1", + "pool-group", + 10, + provider_config, + ); + let mut cursor = PoolKeyCursor::new(PlannerAppState::new(&app), group, None, None, None); + + let mut returned_key_ids = Vec::new(); + while let Some(candidate) = cursor.next_key().await { + returned_key_ids.push(candidate.candidate.key_id); + } + + assert_eq!(returned_key_ids.len(), 32); + assert!(returned_key_ids.iter().any(|key_id| key_id == "key-00112")); + assert!(returned_key_ids.iter().any(|key_id| key_id == "key-00127")); + assert_eq!(cursor.score_next_offset, 128); + assert!(cursor.score_phase_exhausted); + assert_eq!(cursor.score_schedule_interest_count, 16); + assert_eq!(cursor.scanned_keys, 128); + assert_eq!(cursor.budget_scanned_keys, 128); + } + + #[tokio::test] + async fn score_and_fallback_duplicate_does_not_consume_scan_budget_twice() { + let provider_config = Some(json!({ + "pool_advanced": { + "score_top_n": 1, + "score_fallback_scan_limit": 2 + } + })); + let (provider, endpoint, keys, rows) = large_pool_fixture(2, provider_config.clone()); + let data_state = + GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests( + Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![endpoint], + keys, + )), + Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)), + ) + .with_pool_score_repository_for_tests(Arc::new( + InMemoryPoolMemberScoreRepository::seed(vec![sample_provider_key_pool_score( + "provider-pool", + "key-00000", + 1_000.0, + )]), + )) + .with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY); + let app = AppState::new() + .expect("state should build") + .with_data_state_for_tests(data_state); + let group = sample_eligible_candidate( + "provider-pool", + "endpoint-1", + "pool-group", + 10, + provider_config, + ); + let mut cursor = PoolKeyCursor::new(PlannerAppState::new(&app), group, None, None, None); + + let mut returned_key_ids = vec![ + cursor + .next_key() + .await + .expect("one candidate should schedule") + .candidate + .key_id, + cursor + .next_key() + .await + .expect("the other candidate should schedule") + .candidate + .key_id, + ]; + returned_key_ids.sort(); + + assert_eq!(returned_key_ids, ["key-00000", "key-00001"]); + assert_eq!(cursor.scanned_keys, 3); + assert_eq!(cursor.budget_scanned_keys, 2); + } + #[tokio::test] async fn pool_scheduler_skips_invalid_and_exhausted_high_priority_hot_pool_before_fallback_provider( ) { diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/execution.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/execution.rs index c4a6f5042..9aefb9a54 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/execution.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/execution.rs @@ -1,7 +1,8 @@ use super::super::helpers::admin_provider_oauth_key_name_from_auth_config; use super::super::token_import::{ build_provider_access_token_import_auth_config, decode_access_token_expires_at, - provider_oauth_import_authorization_bearer_token, provider_type_supports_access_token_import, + is_claude_session_key, provider_oauth_import_authorization_bearer_token, + provider_type_supports_access_token_import, validate_claude_access_token_import, }; use super::kiro_import::execute_admin_provider_oauth_kiro_batch_import; use super::parse::{ @@ -56,6 +57,33 @@ fn sanitize_windsurf_batch_import_error(error: &OAuthError) -> String { } } +fn can_fallback_batch_refresh_to_access_token( + provider_type: &str, + access_token: Option<&str>, +) -> bool { + !provider_type.eq_ignore_ascii_case("claude_code") + && access_token.is_some() + && provider_type_supports_access_token_import(provider_type) +} + +fn validate_batch_access_token_import( + provider_type: &str, + access_token: &str, + imported_expires_at: Option, + now_unix_secs: u64, +) -> Result<(), String> { + if !provider_type_supports_access_token_import(provider_type) { + return Err( + "Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok Provider".to_string(), + ); + } + if provider_type.eq_ignore_ascii_case("claude_code") { + validate_claude_access_token_import(access_token, imported_expires_at, now_unix_secs) + .map_err(str::to_string)?; + } + Ok(()) +} + const CODEX_AGENT_IDENTITY_SAFE_FIELDS: &[(&str, &[&str])] = &[ ("agent_runtime_id", &["agent_runtime_id", "agentRuntimeId"]), ( @@ -249,6 +277,15 @@ async fn resolve_admin_provider_oauth_batch_import_tokens( .as_deref() .map(str::trim) .filter(|value| !value.is_empty()); + let is_claude = provider_type.eq_ignore_ascii_case("claude_code"); + if is_claude + && refresh_token + .into_iter() + .chain(access_token) + .any(is_claude_session_key) + { + return Err("Claude sessionKey 请使用 Cookie 授权,不能作为导入凭据".to_string()); + } if provider_type.eq_ignore_ascii_case("windsurf") { let token_for_import = refresh_token.or(access_token); @@ -307,7 +344,7 @@ async fn resolve_admin_provider_oauth_batch_import_tokens( if let Some(refresh_token) = refresh_token { let Some(template) = template else { - if provider_type_supports_access_token_import(provider_type) { + if can_fallback_batch_refresh_to_access_token(provider_type, access_token) { if let Some(access_token) = access_token { let (auth_config, expires_at) = build_provider_access_token_import_auth_config( provider_type, @@ -340,7 +377,7 @@ async fn resolve_admin_provider_oauth_batch_import_tokens( Ok(payload) => payload, Err(response) => { let detail = extract_admin_provider_oauth_batch_error_detail(response).await; - if provider_type_supports_access_token_import(provider_type) { + if can_fallback_batch_refresh_to_access_token(provider_type, access_token) { if let Some(access_token) = access_token { let (auth_config, expires_at) = build_provider_access_token_import_auth_config( @@ -381,9 +418,17 @@ async fn resolve_admin_provider_oauth_batch_import_tokens( } if let Some(access_token) = access_token { - if !provider_type_supports_access_token_import(provider_type) { - return Err("Access Token 导入仅支持 Codex / ChatGPT Web / Grok Provider".to_string()); - } + let now_unix_secs = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .ok() + .map(|duration| duration.as_secs()) + .unwrap_or(0); + validate_batch_access_token_import( + provider_type, + access_token, + entry.expires_at, + now_unix_secs, + )?; let (auth_config, expires_at) = build_provider_access_token_import_auth_config( provider_type, access_token, @@ -734,7 +779,8 @@ pub(super) async fn execute_admin_provider_oauth_batch_import( mod tests { use super::super::parse::parse_admin_provider_oauth_batch_import_entries; use super::{ - codex_agent_identity_auth_config_from_import, sanitize_windsurf_batch_import_error, + can_fallback_batch_refresh_to_access_token, codex_agent_identity_auth_config_from_import, + sanitize_windsurf_batch_import_error, validate_batch_access_token_import, }; use aether_oauth::core::OAuthError; use serde_json::json; @@ -763,6 +809,51 @@ mod tests { assert!(!detail.contains("secret-token")); } + #[test] + fn claude_batch_refresh_failure_never_falls_back_to_imported_access_token() { + assert!(!can_fallback_batch_refresh_to_access_token( + "claude_code", + Some("sk-ant-oat01-stale") + )); + assert!(can_fallback_batch_refresh_to_access_token( + "codex", + Some("fallback-access-token") + )); + } + + #[test] + fn claude_batch_access_only_requires_oat_prefix_and_future_expiry() { + let now = 2_000_000_000; + assert!(validate_batch_access_token_import( + "claude_code", + "sk-ant-oat01-valid", + Some(now + 3600), + now, + ) + .is_ok()); + assert!(validate_batch_access_token_import( + "claude_code", + "not-an-oat", + Some(now + 3600), + now, + ) + .is_err()); + assert!(validate_batch_access_token_import( + "claude_code", + "sk-ant-oat01-missing-expiry", + None, + now, + ) + .is_err()); + assert!(validate_batch_access_token_import( + "claude_code", + "sk-ant-oat01-expired", + Some(now), + now, + ) + .is_err()); + } + #[test] fn normalizes_codex_agent_identity_import_without_access_token() { let entries = parse_admin_provider_oauth_batch_import_entries( diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/mod.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/mod.rs index d0d429db0..87d66a654 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/mod.rs @@ -6,6 +6,7 @@ mod progress; mod task; pub(super) use orchestration::handle_admin_provider_oauth_batch_import; +pub(super) use parse::build_admin_provider_oauth_batch_task_state; pub(super) use task::{ handle_admin_provider_oauth_start_agent_identity_import_task, handle_admin_provider_oauth_start_batch_import_task, diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/parse.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/parse.rs index 5560c75b2..a8cb9de1f 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/parse.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/parse.rs @@ -1,6 +1,6 @@ use super::super::token_import::{ - import_tokens_from_raw_token, normalize_provider_import_tokens, - normalize_provider_oauth_import_headers_from_object, + flatten_claude_code_credentials_payload, import_tokens_from_raw_token, + normalize_provider_import_tokens, normalize_provider_oauth_import_headers_from_object, provider_oauth_import_authorization_bearer_token, }; use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response; @@ -46,6 +46,10 @@ pub(super) struct AdminProviderOAuthBatchImportEntry { pub request_headers: Option>, pub user_agent: Option, pub browser_profile: Option, + pub organization_uuid: Option, + pub scopes: Option, + pub subscription_type: Option, + pub rate_limit_tier: Option, } #[derive(Debug, Clone)] @@ -239,10 +243,23 @@ fn extract_admin_provider_oauth_batch_import_entry( request_headers: None, user_agent: None, browser_profile: None, + organization_uuid: None, + scopes: None, + subscription_type: None, + rate_limit_tier: None, }) } } serde_json::Value::Object(object) => { + let is_claude = provider_type.trim().eq_ignore_ascii_case("claude_code"); + let normalized_claude_object = if is_claude { + let mut normalized = object.clone(); + flatten_claude_code_credentials_payload(&mut normalized); + Some(normalized) + } else { + None + }; + let object = normalized_claude_object.as_ref().unwrap_or(object); let is_grok = provider_type.trim().eq_ignore_ascii_case("grok"); let is_windsurf = provider_type.trim().eq_ignore_ascii_case("windsurf"); let is_codex_agent_identity = provider_type.trim().eq_ignore_ascii_case("codex") @@ -271,6 +288,10 @@ fn extract_admin_provider_oauth_batch_import_entry( request_headers: None, user_agent: None, browser_profile: None, + organization_uuid: None, + scopes: None, + subscription_type: None, + rate_limit_tier: None, }); } let refresh_token = coerce_admin_provider_oauth_import_str( @@ -463,6 +484,39 @@ fn extract_admin_provider_oauth_batch_import_entry( .or_else(|| object.get("browser")) .or_else(|| object.get("impersonate")), ); + let organization_uuid = is_claude + .then(|| { + coerce_admin_provider_oauth_import_str( + object + .get("organization_uuid") + .or_else(|| object.get("organizationUuid")) + .or_else(|| object.get("org_uuid")), + ) + }) + .flatten(); + let scopes = is_claude + .then(|| object.get("scopes")) + .flatten() + .filter(|value| value.is_array() || value.is_string()) + .cloned(); + let subscription_type = is_claude + .then(|| { + coerce_admin_provider_oauth_import_str( + object + .get("subscription_type") + .or_else(|| object.get("subscriptionType")), + ) + }) + .flatten(); + let rate_limit_tier = is_claude + .then(|| { + coerce_admin_provider_oauth_import_str( + object + .get("rate_limit_tier") + .or_else(|| object.get("rateLimitTier")), + ) + }) + .flatten(); Some(AdminProviderOAuthBatchImportEntry { parse_error: None, refresh_token, @@ -486,6 +540,10 @@ fn extract_admin_provider_oauth_batch_import_entry( request_headers, user_agent, browser_profile, + organization_uuid, + scopes, + subscription_type, + rate_limit_tier, }) } _ => None, @@ -718,6 +776,10 @@ fn parse_error_entry(error: String) -> AdminProviderOAuthBatchImportEntry { request_headers: None, user_agent: None, browser_profile: None, + organization_uuid: None, + scopes: None, + subscription_type: None, + rate_limit_tier: None, } } @@ -768,6 +830,29 @@ pub(super) fn apply_admin_provider_oauth_batch_import_hints( } return; } + if provider_type == "claude_code" { + if let Some(organization_uuid) = entry.organization_uuid.as_ref() { + auth_config + .entry("org_uuid".to_string()) + .or_insert_with(|| json!(organization_uuid)); + } + if let Some(scopes) = entry.scopes.as_ref() { + auth_config + .entry("scopes".to_string()) + .or_insert_with(|| scopes.clone()); + } + if let Some(subscription_type) = entry.subscription_type.as_ref() { + auth_config + .entry("subscription_type".to_string()) + .or_insert_with(|| json!(subscription_type)); + } + if let Some(rate_limit_tier) = entry.rate_limit_tier.as_ref() { + auth_config + .entry("rate_limit_tier".to_string()) + .or_insert_with(|| json!(rate_limit_tier)); + } + return; + } if !matches!(provider_type.as_str(), "codex" | "chatgpt_web" | "grok") { return; } @@ -878,7 +963,7 @@ pub(super) fn build_admin_provider_oauth_batch_import_response( })) } -pub(super) fn build_admin_provider_oauth_batch_task_state( +pub(in super::super) fn build_admin_provider_oauth_batch_task_state( task_id: &str, provider_id: &str, provider_type: &str, @@ -995,6 +1080,46 @@ mod tests { assert_eq!(entries[0].email.as_deref(), Some("u@example.com")); } + #[test] + fn parses_claude_credentials_json_and_ignores_mcp_oauth() { + let entries = parse_admin_provider_oauth_batch_import_entries( + "claude_code", + r#"{ + "claudeAiOauth": { + "accessToken": "sk-ant-oat01-access", + "refreshToken": "sk-ant-ort01-refresh", + "expiresAt": 2100000000123, + "scopes": ["user:profile"], + "subscriptionType": "pro", + "rateLimitTier": "tier_1", + "organizationUuid": "org-123" + }, + "mcpOAuth": {"accessToken": "must-not-be-imported"} + }"#, + ); + + assert_eq!(entries.len(), 1); + assert_eq!( + entries[0].access_token.as_deref(), + Some("sk-ant-oat01-access") + ); + assert_eq!( + entries[0].refresh_token.as_deref(), + Some("sk-ant-ort01-refresh") + ); + assert_eq!(entries[0].expires_at, Some(2_100_000_000)); + assert_eq!(entries[0].organization_uuid.as_deref(), Some("org-123")); + assert_eq!(entries[0].scopes, Some(json!(["user:profile"]))); + assert_eq!(entries[0].subscription_type.as_deref(), Some("pro")); + assert_eq!(entries[0].rate_limit_tier.as_deref(), Some("tier_1")); + + let ignored = parse_admin_provider_oauth_batch_import_entries( + "claude_code", + r#"{"mcpOAuth":{"accessToken":"must-not-be-imported"}}"#, + ); + assert!(ignored.is_empty()); + } + #[test] fn preserves_codex_agent_identity_entry_without_access_token() { let entries = parse_admin_provider_oauth_batch_import_entries( diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/complete/provider.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/complete/provider.rs index f805012ce..7b7760e38 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/complete/provider.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/complete/provider.rs @@ -1,17 +1,8 @@ -use super::super::super::duplicates::{ - acquire_codex_oauth_account_locks, find_duplicate_provider_oauth_key, - release_codex_oauth_account_locks, -}; use super::super::super::errors::build_internal_control_error_response; use super::super::super::provisioning::{ - build_provider_oauth_auth_config_from_token_payload, create_provider_oauth_catalog_key, - provider_oauth_active_api_formats, provider_oauth_key_proxy_value, - update_existing_provider_oauth_catalog_key, -}; -use super::super::super::runtime::{ - resolve_provider_oauth_runtime_endpoints, - spawn_provider_oauth_account_state_refresh_after_update, + provider_oauth_key_proxy_value, provision_provider_oauth_token_payload_for_provider, }; +use super::super::super::runtime::resolve_provider_oauth_runtime_endpoints; use super::super::super::state::{ admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response, is_fixed_provider_type_for_provider_oauth, @@ -25,11 +16,8 @@ use crate::GatewayError; use axum::{ body::{Body, Bytes}, http, - response::{IntoResponse, Response}, - Json, + response::Response, }; -use serde_json::json; -use std::time::{SystemTime, UNIX_EPOCH}; pub(super) async fn handle_admin_provider_oauth_complete_provider( state: &AdminAppState<'_>, @@ -151,145 +139,15 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider( Err(response) => return Ok(response), }; - let (auth_config, access_token, refresh_token, expires_at) = - build_provider_oauth_auth_config_from_token_payload(&provider_type, &token_payload); - let Some(access_token) = access_token else { - return Ok(build_internal_control_error_response( - http::StatusCode::BAD_REQUEST, - "token exchange 返回缺少 access_token", - )); - }; - - let api_formats = provider_oauth_active_api_formats(&endpoints); - let codex_oauth_account_leases = if provider_type == "codex" { - match acquire_codex_oauth_account_locks( - state, - &provider_id, - &auth_config, - "provider-complete", - ) - .await - { - Ok(leases) => leases, - Err(error) => { - return Ok(build_internal_control_error_response( - error.status_code(), - error.detail(), - )); - } - } - } else { - Vec::new() - }; - let duplicate = match state - .find_duplicate_provider_oauth_key(&provider_id, &auth_config, None) - .await - { - Ok(duplicate) => duplicate, - Err(detail) => { - release_codex_oauth_account_locks(state, codex_oauth_account_leases).await; - return Ok(build_internal_control_error_response( - if provider_type == "codex" { - http::StatusCode::CONFLICT - } else { - http::StatusCode::BAD_REQUEST - }, - detail, - )); - } - }; - - let replaced = duplicate.is_some(); - let persisted_key = if let Some(existing_key) = duplicate { - let update_result = state - .update_existing_provider_oauth_catalog_key( - &existing_key, - &provider_type, - &access_token, - &auth_config, - &api_formats, - key_proxy.clone(), - expires_at, - ) - .await; - match update_result { - Err(error) => { - release_codex_oauth_account_locks(state, codex_oauth_account_leases).await; - return Err(error); - } - Ok(Some(key)) => key, - Ok(None) => { - release_codex_oauth_account_locks(state, codex_oauth_account_leases).await; - return Ok(build_internal_control_error_response( - http::StatusCode::SERVICE_UNAVAILABLE, - "provider oauth write unavailable", - )); - } - } - } else { - let name = payload - .name - .or_else(|| { - auth_config - .get("email") - .and_then(serde_json::Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned) - }) - .unwrap_or_else(|| { - format!( - "账号_{}", - SystemTime::now() - .duration_since(UNIX_EPOCH) - .ok() - .map(|duration| duration.as_secs()) - .unwrap_or(0) - ) - }); - let create_result = state - .create_provider_oauth_catalog_key( - &provider_id, - &provider_type, - &name, - &access_token, - &auth_config, - &api_formats, - key_proxy.clone(), - expires_at, - ) - .await; - match create_result { - Err(error) => { - release_codex_oauth_account_locks(state, codex_oauth_account_leases).await; - return Err(error); - } - Ok(Some(key)) => key, - Ok(None) => { - release_codex_oauth_account_locks(state, codex_oauth_account_leases).await; - return Ok(build_internal_control_error_response( - http::StatusCode::SERVICE_UNAVAILABLE, - "provider oauth write unavailable", - )); - } - } - }; - release_codex_oauth_account_locks(state, codex_oauth_account_leases).await; - - spawn_provider_oauth_account_state_refresh_after_update( - state.cloned_app(), - provider.clone(), - persisted_key.id.clone(), - request_proxy.clone(), - ); - - Ok(Json(json!({ - "key_id": persisted_key.id, - "provider_type": provider_type, - "expires_at": expires_at, - "has_refresh_token": refresh_token.is_some(), - "email": auth_config.get("email").cloned().unwrap_or(serde_json::Value::Null), - "replaced": replaced, - })) - .into_response()) + provision_provider_oauth_token_payload_for_provider( + state, + &provider, + &endpoints, + &token_payload, + payload.name, + key_proxy, + request_proxy, + "provider-complete", + ) + .await } diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/cookie.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/cookie.rs new file mode 100644 index 000000000..45d799808 --- /dev/null +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/cookie.rs @@ -0,0 +1,214 @@ +use super::super::errors::build_internal_control_error_response; +use super::super::provisioning::{ + provider_oauth_key_proxy_value, provision_provider_oauth_token_payload_for_provider, +}; +use super::super::runtime::resolve_provider_oauth_runtime_endpoints; +use super::super::state::authorize_admin_provider_oauth_with_cookie; +use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_cookie_provider_id; +use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; +use crate::GatewayError; +use axum::{body::Body, http, response::Response}; + +pub(super) const MAX_CLAUDE_COOKIE_AUTHORIZE_BODY_BYTES: usize = 32 * 1024; +pub(super) const MAX_CLAUDE_SESSION_KEY_BYTES: usize = 16 * 1024; + +struct ClaudeCookieAuthorizeRequest { + session_key: String, + name: Option, + proxy_node_id: Option, +} + +pub(super) async fn handle_admin_provider_oauth_cookie_authorize( + state: &AdminAppState<'_>, + request_context: &AdminRequestContext<'_>, + request_body: Option<&axum::body::Bytes>, +) -> Result, GatewayError> { + if !state.has_provider_catalog_data_reader() { + return Ok(super::super::state::build_admin_provider_oauth_backend_unavailable_response()); + } + let Some(provider_id) = admin_provider_oauth_cookie_provider_id(request_context.path()) else { + return Ok(build_internal_control_error_response( + http::StatusCode::NOT_FOUND, + "Provider 不存在", + )); + }; + let payload = match parse_claude_cookie_authorize_request(request_body) { + Ok(payload) => payload, + Err(response) => return Ok(response), + }; + + let Some(provider) = state + .read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id)) + .await? + .into_iter() + .next() + else { + return Ok(build_internal_control_error_response( + http::StatusCode::NOT_FOUND, + "Provider 不存在", + )); + }; + let provider_type = provider.provider_type.trim().to_ascii_lowercase(); + if provider_type != "claude_code" { + return Ok(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "Cookie 授权仅支持 Claude Code Provider", + )); + } + + let endpoint_resolution = + resolve_provider_oauth_runtime_endpoints(state, &provider, &provider_type).await?; + let endpoints = endpoint_resolution.endpoints; + let request_proxy = state + .resolve_admin_provider_oauth_operation_proxy_snapshot( + payload.proxy_node_id.as_deref(), + &[ + endpoint_resolution + .runtime_endpoint + .as_ref() + .and_then(|endpoint| endpoint.proxy.as_ref()), + provider.proxy.as_ref(), + ], + ) + .await; + let key_proxy = provider_oauth_key_proxy_value(payload.proxy_node_id.as_deref()); + let token_payload = match authorize_admin_provider_oauth_with_cookie( + state, + payload.session_key, + request_proxy.clone(), + ) + .await + { + Ok(payload) => payload, + Err(response) => return Ok(response), + }; + + provision_provider_oauth_token_payload_for_provider( + state, + &provider, + &endpoints, + &token_payload, + payload.name, + key_proxy, + request_proxy, + "cookie-authorize", + ) + .await +} + +fn parse_claude_cookie_authorize_request( + request_body: Option<&axum::body::Bytes>, +) -> Result> { + let Some(request_body) = request_body else { + return Err(bad_cookie_request("请求体必须是合法的 JSON 对象")); + }; + if request_body.len() > MAX_CLAUDE_COOKIE_AUTHORIZE_BODY_BYTES { + return Err(bad_cookie_request("Cookie 授权请求体过大")); + } + let payload = serde_json::from_slice::(request_body) + .ok() + .and_then(|value| value.as_object().cloned()) + .ok_or_else(|| bad_cookie_request("请求体必须是合法的 JSON 对象"))?; + let cookie = ["cookie", "session_key", "sessionKey"] + .into_iter() + .find_map(|key| payload.get(key).and_then(serde_json::Value::as_str)) + .ok_or_else(|| bad_cookie_request("Cookie 不能为空"))?; + let session_key = normalize_claude_session_key(cookie) + .ok_or_else(|| bad_cookie_request("Cookie 中缺少有效的 sessionKey"))?; + + Ok(ClaudeCookieAuthorizeRequest { + session_key, + name: optional_trimmed_string(&payload, "name"), + proxy_node_id: optional_trimmed_string(&payload, "proxy_node_id") + .or_else(|| optional_trimmed_string(&payload, "proxyNodeId")), + }) +} + +pub(super) fn normalize_claude_session_key(raw: &str) -> Option { + let raw = raw.trim(); + if raw.is_empty() || raw.len() > MAX_CLAUDE_SESSION_KEY_BYTES || raw.contains(['\r', '\n']) { + return None; + } + let cookie = raw + .split_once(':') + .filter(|(name, _)| name.trim().eq_ignore_ascii_case("cookie")) + .map(|(_, value)| value.trim()) + .unwrap_or(raw); + + if !cookie.contains('=') { + return valid_session_key_value(cookie).then(|| cookie.to_string()); + } + + let mut session_key = None; + for segment in cookie.split(';') { + let (name, value) = segment.trim().split_once('=')?; + if !name.trim().eq_ignore_ascii_case("sessionKey") { + continue; + } + if session_key.is_some() || !valid_session_key_value(value.trim()) { + return None; + } + session_key = Some(value.trim().to_string()); + } + session_key +} + +fn valid_session_key_value(value: &str) -> bool { + !value.is_empty() + && value.len() <= MAX_CLAUDE_SESSION_KEY_BYTES + && !value.contains(['\r', '\n', ';']) + && http::HeaderValue::from_str(value).is_ok() +} + +fn optional_trimmed_string( + payload: &serde_json::Map, + key: &str, +) -> Option { + payload + .get(key) + .and_then(serde_json::Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) +} + +fn bad_cookie_request(detail: &'static str) -> Response { + build_internal_control_error_response(http::StatusCode::BAD_REQUEST, detail) +} + +#[cfg(test)] +mod tests { + use super::normalize_claude_session_key; + + #[test] + fn normalizes_supported_claude_cookie_inputs() { + for (input, expected) in [ + ("sk-ant-sid01-raw", "sk-ant-sid01-raw"), + ("sessionKey=sk-ant-sid01-pair", "sk-ant-sid01-pair"), + ( + "Cookie: other=value; sessionKey=sk-ant-sid01-header; theme=dark", + "sk-ant-sid01-header", + ), + ] { + assert_eq!( + normalize_claude_session_key(input).as_deref(), + Some(expected) + ); + } + } + + #[test] + fn rejects_ambiguous_or_unsafe_claude_cookie_inputs() { + for input in [ + "", + "foo=bar", + "sessionKey=one; sessionKey=two", + "sessionKey=value\r\nx-leak: yes", + ] { + assert!( + normalize_claude_session_key(input).is_none(), + "input={input:?}" + ); + } + } +} diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/cookie_task.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/cookie_task.rs new file mode 100644 index 000000000..e5a266c74 --- /dev/null +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/cookie_task.rs @@ -0,0 +1,682 @@ +use super::super::errors::build_internal_control_error_response; +use super::super::provisioning::{ + provider_oauth_key_proxy_value, provision_provider_oauth_token_payload_for_provider, +}; +use super::super::runtime::resolve_provider_oauth_runtime_endpoints; +use super::super::state::{ + authorize_admin_provider_oauth_with_cookie, + build_admin_provider_oauth_backend_unavailable_response, +}; +use super::batch::build_admin_provider_oauth_batch_task_state; +use super::cookie::{normalize_claude_session_key, MAX_CLAUDE_SESSION_KEY_BYTES}; +use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_cookie_task_provider_id; +use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; +use crate::task_runtime::{ + append_event_with_logging, now_unix_secs, task_definition, update_run_status, + upsert_run_with_logging, TASK_KEY_PROVIDER_OAUTH_BATCH_IMPORT, +}; +use crate::GatewayError; +use aether_data_contracts::repository::background_tasks::{ + BackgroundTaskKind, BackgroundTaskStatus, UpsertBackgroundTaskRun, +}; +use axum::{ + body::{to_bytes, Body, Bytes}, + http, + response::{IntoResponse, Response}, + Json, +}; +use futures_util::{stream, StreamExt}; +use serde_json::{json, Value}; +use std::collections::HashSet; +use std::time::{SystemTime, UNIX_EPOCH}; +use tokio::task; +use uuid::Uuid; + +const CLAUDE_COOKIE_TASK_IMPORT_KIND: &str = "cookie_authorize"; +const CLAUDE_COOKIE_TASK_ID_PREFIX: &str = "claude-cookie-"; +const MAX_CLAUDE_COOKIE_TASK_ENTRIES: usize = 20; +const MAX_CLAUDE_COOKIE_TASK_BODY_BYTES: usize = 768 * 1024; +const CLAUDE_COOKIE_AUTHORIZATION_CONCURRENCY: usize = 3; +const MAX_SAFE_ERROR_DETAIL_BYTES: usize = 512; + +type ClaudeCookieTaskEntry = Result; + +struct ClaudeCookieTaskRequest { + entries: Vec, + proxy_node_id: Option, +} + +pub(super) async fn handle_admin_provider_oauth_start_cookie_task( + state: &AdminAppState<'_>, + request_context: &AdminRequestContext<'_>, + request_body: Option<&Bytes>, +) -> Result, GatewayError> { + if !state.has_provider_catalog_data_reader() { + return Ok(build_admin_provider_oauth_backend_unavailable_response()); + } + let Some(provider_id) = admin_provider_oauth_cookie_task_provider_id(request_context.path()) + else { + return Ok(build_internal_control_error_response( + http::StatusCode::NOT_FOUND, + "Provider 不存在", + )); + }; + let payload = match parse_claude_cookie_task_request(request_body) { + Ok(payload) => payload, + Err(response) => return Ok(response), + }; + + let Some(provider) = state + .read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id)) + .await? + .into_iter() + .next() + else { + return Ok(build_internal_control_error_response( + http::StatusCode::NOT_FOUND, + "Provider 不存在", + )); + }; + let provider_type = provider.provider_type.trim().to_ascii_lowercase(); + if provider_type != "claude_code" { + return Ok(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "Cookie 授权仅支持 Claude Code Provider", + )); + } + + let endpoint_resolution = + resolve_provider_oauth_runtime_endpoints(state, &provider, &provider_type).await?; + let endpoints = endpoint_resolution.endpoints; + let request_proxy = state + .resolve_admin_provider_oauth_operation_proxy_snapshot( + payload.proxy_node_id.as_deref(), + &[ + endpoint_resolution + .runtime_endpoint + .as_ref() + .and_then(|endpoint| endpoint.proxy.as_ref()), + provider.proxy.as_ref(), + ], + ) + .await; + let key_proxy = provider_oauth_key_proxy_value(payload.proxy_node_id.as_deref()); + + let task_id = format!("{CLAUDE_COOKIE_TASK_ID_PREFIX}{}", Uuid::new_v4()); + let total = payload.entries.len(); + let created_at = now_unix_secs(); + let submitted_state = build_admin_provider_oauth_batch_task_state( + &task_id, + &provider_id, + &provider_type, + CLAUDE_COOKIE_TASK_IMPORT_KIND, + "submitted", + total, + 0, + 0, + 0, + 0, + 0, + Some("任务已提交,等待执行"), + None, + Vec::new(), + created_at, + None, + None, + ); + if state + .save_provider_oauth_batch_task_payload(&task_id, &submitted_state) + .await + .is_err() + { + return Ok(build_internal_control_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "provider oauth batch task redis unavailable", + )); + } + + if state.has_background_task_data_writer() { + let max_attempts = task_definition(TASK_KEY_PROVIDER_OAUTH_BATCH_IMPORT) + .map(|item| item.retry_policy.max_attempts) + .unwrap_or(1); + let run = UpsertBackgroundTaskRun { + id: task_id.clone(), + task_key: TASK_KEY_PROVIDER_OAUTH_BATCH_IMPORT.to_string(), + kind: BackgroundTaskKind::OnDemand, + trigger: "manual".to_string(), + status: BackgroundTaskStatus::Queued, + attempt: 1, + max_attempts, + owner_instance: Some(state.app().tunnel.local_instance_id().to_string()), + progress_percent: 0, + progress_message: Some("Claude Cookie authorization queued".to_string()), + payload_json: Some(json!({ + "provider_id": provider_id.clone(), + "provider_type": provider_type.clone(), + "import_kind": CLAUDE_COOKIE_TASK_IMPORT_KIND, + "total": total, + })), + result_json: None, + error_message: None, + cancel_requested: false, + created_by: Some("admin".to_string()), + created_at_unix_secs: created_at, + started_at_unix_secs: None, + finished_at_unix_secs: None, + updated_at_unix_secs: created_at, + }; + let _ = upsert_run_with_logging(state.app(), run).await; + append_event_with_logging( + state.app(), + &task_id, + "queued", + "Claude Cookie authorization queued", + Some(json!({ + "provider_id": provider_id.clone(), + "provider_type": provider_type.clone(), + "import_kind": CLAUDE_COOKIE_TASK_IMPORT_KIND, + "total": total, + })), + ) + .await; + } + + let task_state = state.cloned_app(); + let task_id_for_worker = task_id.clone(); + let provider_id_for_worker = provider_id.clone(); + let provider_type_for_worker = provider_type.clone(); + task::spawn(async move { + let started_at = current_unix_secs_or(created_at); + let task_admin_state = AdminAppState::new(&task_state); + save_cookie_task_state( + &task_admin_state, + &task_id_for_worker, + &provider_id_for_worker, + &provider_type_for_worker, + "processing", + total, + 0, + 0, + 0, + 0, + 0, + Some("正在获取 Claude 授权"), + Vec::new(), + created_at, + started_at, + None, + ) + .await; + let _ = update_run_status( + &task_state, + &task_id_for_worker, + BackgroundTaskStatus::Running, + Some(1), + Some("Claude Cookie authorization started".to_string()), + None, + None, + Some(started_at), + None, + ) + .await; + append_event_with_logging( + &task_state, + &task_id_for_worker, + "running", + "Claude Cookie authorization started", + None, + ) + .await; + + let mut pending = stream::iter(payload.entries.into_iter().enumerate().map( + |(index, entry)| { + let proxy = request_proxy.clone(); + let task_admin_state = &task_admin_state; + async move { + let result = match entry { + Ok(session_key) => authorize_admin_provider_oauth_with_cookie( + task_admin_state, + session_key, + proxy, + ) + .await + .map_err(|_| "Claude Cookie 授权失败".to_string()), + Err(detail) => Err(detail), + }; + (index, result) + } + }, + )) + .buffer_unordered(CLAUDE_COOKIE_AUTHORIZATION_CONCURRENCY); + let mut authorization_results = Vec::with_capacity(total); + while let Some(result) = pending.next().await { + authorization_results.push(result); + } + authorization_results.sort_by_key(|(index, _)| *index); + + let mut success = 0usize; + let mut failed = 0usize; + let mut created_count = 0usize; + let mut replaced_count = 0usize; + let mut error_samples = Vec::new(); + + for (index, authorization_result) in authorization_results { + let result = match authorization_result { + Ok(token_payload) => match provision_provider_oauth_token_payload_for_provider( + &task_admin_state, + &provider, + &endpoints, + &token_payload, + None, + key_proxy.clone(), + request_proxy.clone(), + "cookie-authorize-batch", + ) + .await + { + Ok(response) => cookie_task_item_from_response(index, response).await, + Err(_) => cookie_task_error(index, "provider oauth write unavailable"), + }, + Err(detail) => cookie_task_error(index, detail.as_str()), + }; + + if result.get("status").and_then(Value::as_str) == Some("success") { + success += 1; + if result.get("replaced").and_then(Value::as_bool) == Some(true) { + replaced_count += 1; + } else { + created_count += 1; + } + } else { + failed += 1; + error_samples.push(result); + } + let processed = success.saturating_add(failed); + let message = format!("处理中 {processed}/{total}"); + save_cookie_task_state( + &task_admin_state, + &task_id_for_worker, + &provider_id_for_worker, + &provider_type_for_worker, + "processing", + total, + processed, + success, + failed, + created_count, + replaced_count, + Some(message.as_str()), + error_samples.clone(), + created_at, + started_at, + None, + ) + .await; + } + + let finished_at = current_unix_secs_or(started_at); + let message = format!("授权完成:成功 {success},失败 {failed}"); + save_cookie_task_state( + &task_admin_state, + &task_id_for_worker, + &provider_id_for_worker, + &provider_type_for_worker, + "completed", + total, + total, + success, + failed, + created_count, + replaced_count, + Some(message.as_str()), + error_samples, + created_at, + started_at, + Some(finished_at), + ) + .await; + let _ = update_run_status( + &task_state, + &task_id_for_worker, + BackgroundTaskStatus::Succeeded, + Some(100), + Some(message), + Some(json!({ + "provider_id": provider_id_for_worker, + "provider_type": provider_type_for_worker, + "import_kind": CLAUDE_COOKIE_TASK_IMPORT_KIND, + "total": total, + "success": success, + "failed": failed, + "created_count": created_count, + "replaced_count": replaced_count, + })), + None, + None, + Some(finished_at), + ) + .await; + append_event_with_logging( + &task_state, + &task_id_for_worker, + "succeeded", + "Claude Cookie authorization completed", + None, + ) + .await; + }); + + Ok(Json(submitted_state).into_response()) +} + +#[allow(clippy::too_many_arguments)] +async fn save_cookie_task_state( + state: &AdminAppState<'_>, + task_id: &str, + provider_id: &str, + provider_type: &str, + status: &str, + total: usize, + processed: usize, + success: usize, + failed: usize, + created_count: usize, + replaced_count: usize, + message: Option<&str>, + error_samples: Vec, + created_at: u64, + started_at: u64, + finished_at: Option, +) { + let task_state = build_admin_provider_oauth_batch_task_state( + task_id, + provider_id, + provider_type, + CLAUDE_COOKIE_TASK_IMPORT_KIND, + status, + total, + processed, + success, + failed, + created_count, + replaced_count, + message, + None, + error_samples, + created_at, + Some(started_at), + finished_at, + ); + let _ = state + .save_provider_oauth_batch_task_payload(task_id, &task_state) + .await; +} + +async fn cookie_task_item_from_response(index: usize, response: Response) -> Value { + let status = response.status(); + let body = to_bytes(response.into_body(), crate::MAX_ERROR_BODY_BYTES) + .await + .ok(); + let payload = body + .as_deref() + .and_then(|body| serde_json::from_slice::(body).ok()); + if status.is_success() { + let Some(payload) = payload else { + return cookie_task_error(index, "provider oauth write unavailable"); + }; + let Some(key_id) = payload.get("key_id").and_then(Value::as_str) else { + return cookie_task_error(index, "provider oauth write unavailable"); + }; + return json!({ + "index": index, + "status": "success", + "key_id": key_id, + "email": payload.get("email").cloned().unwrap_or(Value::Null), + "replaced": payload.get("replaced").and_then(Value::as_bool).unwrap_or(false), + "error": Value::Null, + }); + } + + let detail = payload + .as_ref() + .and_then(|payload| payload.get("detail")) + .and_then(Value::as_str) + .and_then(safe_error_detail) + .unwrap_or("Claude 账号创建或更新失败"); + cookie_task_error(index, detail) +} + +fn safe_error_detail(detail: &str) -> Option<&str> { + let detail = detail.trim(); + if detail.is_empty() || detail.len() > MAX_SAFE_ERROR_DETAIL_BYTES { + return None; + } + let normalized = detail.to_ascii_lowercase(); + if normalized.contains("sessionkey") + || normalized.contains("sk-ant-") + || normalized.contains("cookie:") + { + return None; + } + Some(detail) +} + +fn cookie_task_error(index: usize, detail: &str) -> Value { + json!({ + "index": index, + "status": "error", + "error": detail, + "replaced": false, + }) +} + +fn parse_claude_cookie_task_request( + request_body: Option<&Bytes>, +) -> Result> { + let Some(request_body) = request_body else { + return Err(bad_cookie_task_request("请求体必须是合法的 JSON 对象")); + }; + if request_body.len() > MAX_CLAUDE_COOKIE_TASK_BODY_BYTES { + return Err(bad_cookie_task_request("Cookie 授权请求体过大")); + } + let payload = serde_json::from_slice::(request_body) + .ok() + .and_then(|value| value.as_object().cloned()) + .ok_or_else(|| bad_cookie_task_request("请求体必须是合法的 JSON 对象"))?; + + let legacy_keys = ["cookie", "session_key", "sessionKey"]; + let has_legacy_cookie = legacy_keys.iter().any(|key| payload.contains_key(*key)); + let raw_entries = if let Some(cookies) = payload.get("cookies") { + if has_legacy_cookie { + return Err(bad_cookie_task_request("cookie 与 cookies 不能同时提供")); + } + let cookies = cookies + .as_array() + .ok_or_else(|| bad_cookie_task_request("cookies 必须是字符串数组"))?; + if cookies.is_empty() { + return Err(bad_cookie_task_request("Cookie 不能为空")); + } + if cookies.len() > MAX_CLAUDE_COOKIE_TASK_ENTRIES { + return Err(bad_cookie_task_request("Cookie 批量授权最多支持 20 条")); + } + cookies + .iter() + .map(|value| { + value + .as_str() + .map(ToOwned::to_owned) + .ok_or_else(|| bad_cookie_task_request("cookies 必须是字符串数组")) + }) + .collect::, _>>()? + } else { + let raw = legacy_keys + .into_iter() + .find_map(|key| payload.get(key).and_then(Value::as_str)) + .ok_or_else(|| bad_cookie_task_request("Cookie 不能为空"))?; + let entries = raw + .lines() + .map(str::trim) + .filter(|line| !line.is_empty()) + .map(ToOwned::to_owned) + .collect::>(); + if entries.is_empty() { + return Err(bad_cookie_task_request("Cookie 不能为空")); + } + if entries.len() > MAX_CLAUDE_COOKIE_TASK_ENTRIES { + return Err(bad_cookie_task_request("Cookie 批量授权最多支持 20 条")); + } + entries + }; + + let mut seen_session_keys = HashSet::new(); + let entries = raw_entries + .into_iter() + .map(|raw| { + let session_key = + normalize_claude_session_key(&raw).ok_or_else(|| "Cookie 格式无效".to_string())?; + if !seen_session_keys.insert(session_key.clone()) { + return Err("Cookie 重复".to_string()); + } + Ok(session_key) + }) + .collect(); + Ok(ClaudeCookieTaskRequest { + entries, + proxy_node_id: optional_trimmed_string(&payload, "proxy_node_id") + .or_else(|| optional_trimmed_string(&payload, "proxyNodeId")), + }) +} + +fn optional_trimmed_string(payload: &serde_json::Map, key: &str) -> Option { + payload + .get(key) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) +} + +fn bad_cookie_task_request(detail: &'static str) -> Response { + build_internal_control_error_response(http::StatusCode::BAD_REQUEST, detail) +} + +fn current_unix_secs_or(fallback: u64) -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .ok() + .map(|duration| duration.as_secs()) + .unwrap_or(fallback) +} + +#[cfg(test)] +mod tests { + use super::{ + parse_claude_cookie_task_request, safe_error_detail, MAX_CLAUDE_COOKIE_TASK_BODY_BYTES, + MAX_CLAUDE_COOKIE_TASK_ENTRIES, MAX_CLAUDE_SESSION_KEY_BYTES, + }; + use axum::body::{to_bytes, Bytes}; + use serde_json::json; + + #[test] + fn parses_canonical_and_multiline_cookie_batches_without_retaining_raw_headers() { + for payload in [ + json!({ + "cookies": [ + "sessionKey=sk-ant-sid01-one", + "Cookie: theme=dark; sessionKey=sk-ant-sid01-two" + ], + "proxy_node_id": "proxy-1" + }), + json!({ + "cookie": "sessionKey=sk-ant-sid01-one\n\nCookie: sessionKey=sk-ant-sid01-two", + "proxyNodeId": "proxy-1" + }), + ] { + let body = Bytes::from(payload.to_string()); + let parsed = + parse_claude_cookie_task_request(Some(&body)).expect("cookie batch should parse"); + assert_eq!(parsed.entries.len(), 2); + assert_eq!(parsed.entries[0].as_deref(), Ok("sk-ant-sid01-one")); + assert_eq!(parsed.entries[1].as_deref(), Ok("sk-ant-sid01-two")); + assert_eq!(parsed.proxy_node_id.as_deref(), Some("proxy-1")); + } + } + + #[test] + fn keeps_invalid_cookie_lines_as_independent_sanitized_results() { + let body = Bytes::from( + json!({ + "cookies": ["foo=bar", "sessionKey=valid", "Cookie: sessionKey=valid"] + }) + .to_string(), + ); + let parsed = + parse_claude_cookie_task_request(Some(&body)).expect("request should be accepted"); + assert_eq!(parsed.entries.len(), 3); + assert_eq!( + parsed.entries[0].as_ref().expect_err("entry should fail"), + "Cookie 格式无效" + ); + assert_eq!(parsed.entries[1].as_deref(), Ok("valid")); + assert_eq!( + parsed.entries[2].as_ref().expect_err("entry should fail"), + "Cookie 重复" + ); + } + + #[test] + fn accepts_twenty_maximum_length_session_keys_within_batch_body_limit() { + let cookies = (0..MAX_CLAUDE_COOKIE_TASK_ENTRIES) + .map(|index| { + let prefix = format!("{index:02}-"); + format!( + "{prefix}{}", + "x".repeat(MAX_CLAUDE_SESSION_KEY_BYTES - prefix.len()) + ) + }) + .collect::>(); + let body = Bytes::from(json!({"cookies": cookies}).to_string()); + assert!(body.len() < MAX_CLAUDE_COOKIE_TASK_BODY_BYTES); + let parsed = parse_claude_cookie_task_request(Some(&body)) + .expect("maximum valid batch should parse"); + assert_eq!(parsed.entries.len(), MAX_CLAUDE_COOKIE_TASK_ENTRIES); + assert!(parsed.entries.iter().all(Result::is_ok)); + } + + #[tokio::test] + async fn rejects_ambiguous_or_oversized_cookie_batches_without_echoing_secrets() { + let too_many = vec!["sessionKey=value"; MAX_CLAUDE_COOKIE_TASK_ENTRIES + 1]; + for payload in [ + json!({"cookie": "sessionKey=secret", "cookies": ["sessionKey=other"]}), + json!({"cookies": too_many}), + json!({"cookies": []}), + ] { + let body = Bytes::from(payload.to_string()); + let response = match parse_claude_cookie_task_request(Some(&body)) { + Ok(_) => panic!("request should fail"), + Err(response) => response, + }; + let response_body = to_bytes(response.into_body(), usize::MAX) + .await + .expect("body should read"); + let text = String::from_utf8_lossy(&response_body); + assert!(!text.contains("secret")); + assert!(!text.contains("other")); + } + + let oversized_body = Bytes::from(vec![b'x'; MAX_CLAUDE_COOKIE_TASK_BODY_BYTES + 1]); + let response = match parse_claude_cookie_task_request(Some(&oversized_body)) { + Ok(_) => panic!("oversized request should fail"), + Err(response) => response, + }; + assert_eq!(response.status(), http::StatusCode::BAD_REQUEST); + } + + #[test] + fn error_detail_filter_rejects_possible_cookie_or_token_leaks() { + assert_eq!(safe_error_detail("账号重复"), Some("账号重复")); + assert!(safe_error_detail("sessionKey=secret").is_none()); + assert!(safe_error_detail("upstream sk-ant-oat01-secret").is_none()); + assert!(safe_error_detail("Cookie: secret").is_none()); + } +} diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/import.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/import.rs index 72e2e739c..ad48fce19 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/import.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/import.rs @@ -20,9 +20,10 @@ use super::super::state::{ use super::helpers::admin_provider_oauth_key_name_from_auth_config; use super::token_import::{ build_provider_access_token_import_auth_config, decode_access_token_expires_at, + flatten_claude_code_credentials_payload, is_claude_session_key, normalize_provider_import_tokens, normalize_provider_oauth_import_headers_from_object, provider_oauth_import_authorization_bearer_token_from_object, - provider_type_supports_access_token_import, + provider_type_supports_access_token_import, validate_claude_access_token_import, }; use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_import_provider_id; use crate::handlers::admin::request::{ @@ -487,6 +488,29 @@ fn apply_single_import_hints( } return; } + if provider_type == "claude_code" { + for (target, keys) in [ + ( + "org_uuid", + &["org_uuid", "organization_uuid", "organizationUuid"][..], + ), + ( + "subscription_type", + &["subscription_type", "subscriptionType"][..], + ), + ("rate_limit_tier", &["rate_limit_tier", "rateLimitTier"][..]), + ] { + if let Some(value) = import_payload_string_any(payload, keys) { + auth_config + .entry(target.to_string()) + .or_insert_with(|| json!(value)); + } + } + if let Some(scopes) = payload.get("scopes").cloned() { + auth_config.entry("scopes".to_string()).or_insert(scopes); + } + return; + } if !matches!(provider_type.as_str(), "codex" | "chatgpt_web" | "grok") { return; } @@ -617,7 +641,9 @@ async fn resolve_admin_provider_oauth_single_import_tokens( { Ok(payload) => payload, Err(response) => { - if provider_type_supports_access_token_import(provider_type) { + if !provider_type.eq_ignore_ascii_case("claude_code") + && provider_type_supports_access_token_import(provider_type) + { if let Some(access_token) = access_token .map(str::trim) .filter(|value| !value.is_empty()) @@ -671,10 +697,25 @@ async fn resolve_admin_provider_oauth_single_import_tokens( "Refresh Token 或 Access Token 不能为空", )); }; + if provider_type.eq_ignore_ascii_case("claude_code") { + let now_unix_secs = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .ok() + .map(|duration| duration.as_secs()) + .unwrap_or(0); + if let Err(detail) = + validate_claude_access_token_import(access_token, imported_expires_at, now_unix_secs) + { + return Err(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + detail, + )); + } + } if !provider_type_supports_access_token_import(provider_type) { return Err(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, - "Access Token 导入仅支持 Codex / ChatGPT Web / Grok Provider", + "Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok Provider", )); } @@ -776,7 +817,7 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token( "请求体必须是合法的 JSON 对象", )); }; - let raw_payload = match serde_json::from_slice::(request_body) { + let mut raw_payload = match serde_json::from_slice::(request_body) { Ok(serde_json::Value::Object(map)) => map, _ => { return Ok(build_internal_control_error_response( @@ -797,21 +838,6 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token( } else { None }; - let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken"); - let access_token_input = import_payload_string_any( - &raw_payload, - &[ - "access_token", - "accessToken", - "sso_token", - "ssoToken", - "session_token", - "sessionToken", - ], - ) - .or_else(|| provider_oauth_import_authorization_bearer_token_from_object(&raw_payload)); - let imported_expires_at = - import_payload_u64_any(&raw_payload, &["expires_at", "expiresAt", "expired"]); let name = raw_payload .get("name") .and_then(serde_json::Value::as_str) @@ -837,11 +863,41 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token( )); }; let provider_type = provider.provider_type.trim().to_ascii_lowercase(); + if provider_type == "claude_code" { + flatten_claude_code_credentials_payload(&mut raw_payload); + } + let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken"); + let access_token_input = import_payload_string_any( + &raw_payload, + &[ + "access_token", + "accessToken", + "sso_token", + "ssoToken", + "session_token", + "sessionToken", + ], + ) + .or_else(|| provider_oauth_import_authorization_bearer_token_from_object(&raw_payload)); + let imported_expires_at = + import_payload_u64_any(&raw_payload, &["expires_at", "expiresAt", "expired"]); let (refresh_token_input, access_token_input) = normalize_provider_import_tokens( &provider_type, refresh_token_input.as_deref(), access_token_input.as_deref(), ); + if provider_type == "claude_code" + && refresh_token_input + .as_deref() + .into_iter() + .chain(access_token_input.as_deref()) + .any(is_claude_session_key) + { + return Ok(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "Claude sessionKey 请使用 Cookie 授权,不能作为导入凭据", + )); + } if !create_agent_identity && refresh_token_input.is_none() && access_token_input.is_none() { return Ok(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/mod.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/mod.rs index b25aca6fb..630b9e48a 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/mod.rs @@ -6,9 +6,11 @@ use crate::handlers::admin::provider::shared::paths::{ admin_provider_oauth_agent_identity_import_task_provider_id, admin_provider_oauth_batch_import_provider_id, admin_provider_oauth_batch_import_task_provider_id, admin_provider_oauth_complete_key_id, - admin_provider_oauth_complete_provider_id, admin_provider_oauth_device_authorize_provider_id, - admin_provider_oauth_import_provider_id, admin_provider_oauth_refresh_key_id, - admin_provider_oauth_start_key_id, admin_provider_oauth_start_provider_id, + admin_provider_oauth_complete_provider_id, admin_provider_oauth_cookie_provider_id, + admin_provider_oauth_cookie_task_provider_id, + admin_provider_oauth_device_authorize_provider_id, admin_provider_oauth_import_provider_id, + admin_provider_oauth_refresh_key_id, admin_provider_oauth_start_key_id, + admin_provider_oauth_start_provider_id, }; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::GatewayError; @@ -21,6 +23,8 @@ use axum::{ mod batch; mod complete; +mod cookie; +mod cookie_task; mod device; mod helpers; mod import; @@ -94,6 +98,12 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response( )); } + if route_kind == Some("get_cookie_authorize_task_status") && *method == http::Method::GET { + return Ok(Some( + tasks::handle_admin_provider_oauth_cookie_task_status(state, request_context).await?, + )); + } + if route_kind == Some("complete_key_oauth") && *method == http::Method::POST { let response = complete::handle_admin_provider_oauth_complete_key( state, @@ -156,6 +166,38 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response( ))); } + if route_kind == Some("cookie_authorize") && *method == http::Method::POST { + let response = cookie::handle_admin_provider_oauth_cookie_authorize( + state, + request_context, + request_body, + ) + .await?; + return Ok(Some(helpers::attach_admin_provider_oauth_audit_response( + response, + "admin_provider_oauth_cookie_authorized", + "authorize_provider_oauth_with_cookie", + "provider", + admin_provider_oauth_cookie_provider_id(request_context.path()), + ))); + } + + if route_kind == Some("start_cookie_authorize_task") && *method == http::Method::POST { + let response = cookie_task::handle_admin_provider_oauth_start_cookie_task( + state, + request_context, + request_body, + ) + .await?; + return Ok(Some(helpers::attach_admin_provider_oauth_audit_response( + response, + "admin_provider_oauth_cookie_task_started", + "start_provider_oauth_cookie_task", + "provider", + admin_provider_oauth_cookie_task_provider_id(request_context.path()), + ))); + } + if route_kind == Some("batch_import_oauth") && *method == http::Method::POST { let response = batch::handle_admin_provider_oauth_batch_import(state, request_context, request_body) @@ -226,7 +268,12 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response( if matches!( route_kind, - Some("refresh_key_oauth" | "import_refresh_token") + Some( + "refresh_key_oauth" + | "import_refresh_token" + | "cookie_authorize" + | "start_cookie_authorize_task", + ) ) { return Ok(Some( build_admin_provider_oauth_backend_unavailable_response(), diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/tasks.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/tasks.rs index eece2024f..922197d8f 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/tasks.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/tasks.rs @@ -1,7 +1,7 @@ use super::super::errors::build_internal_control_error_response; use crate::handlers::admin::provider::shared::paths::{ admin_provider_oauth_agent_identity_import_task_path, - admin_provider_oauth_batch_import_task_path, + admin_provider_oauth_batch_import_task_path, admin_provider_oauth_cookie_task_path, }; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::shared::attach_admin_audit_response; @@ -14,20 +14,41 @@ use axum::{ }; const PROVIDER_AGENT_IDENTITY_IMPORT_KIND: &str = "agent_identity"; +const PROVIDER_OAUTH_BATCH_IMPORT_KIND: &str = "oauth_batch"; +const PROVIDER_COOKIE_AUTHORIZE_IMPORT_KIND: &str = "cookie_authorize"; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum ProviderOAuthTaskRouteKind { + BatchImport, + AgentIdentity, + CookieAuthorize, +} fn provider_oauth_import_task_matches_route( task_id: &str, payload: &serde_json::Value, - agent_identity_only: bool, + route_kind: ProviderOAuthTaskRouteKind, ) -> bool { let has_agent_prefix = task_id.starts_with("agent-identity-"); + let has_cookie_prefix = task_id.starts_with("claude-cookie-"); let import_kind = payload .get("import_kind") .and_then(serde_json::Value::as_str); - if agent_identity_only { - has_agent_prefix && import_kind == Some(PROVIDER_AGENT_IDENTITY_IMPORT_KIND) - } else { - !has_agent_prefix && import_kind != Some(PROVIDER_AGENT_IDENTITY_IMPORT_KIND) + match route_kind { + ProviderOAuthTaskRouteKind::BatchImport => { + !has_agent_prefix + && !has_cookie_prefix + && matches!( + import_kind, + None | Some("") | Some(PROVIDER_OAUTH_BATCH_IMPORT_KIND) + ) + } + ProviderOAuthTaskRouteKind::AgentIdentity => { + has_agent_prefix && import_kind == Some(PROVIDER_AGENT_IDENTITY_IMPORT_KIND) + } + ProviderOAuthTaskRouteKind::CookieAuthorize => { + has_cookie_prefix && import_kind == Some(PROVIDER_COOKIE_AUTHORIZE_IMPORT_KIND) + } } } @@ -35,30 +56,63 @@ pub(super) async fn handle_admin_provider_oauth_batch_import_task_status( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, ) -> Result, GatewayError> { - handle_admin_provider_oauth_import_task_status(state, request_context, false).await + handle_admin_provider_oauth_import_task_status( + state, + request_context, + ProviderOAuthTaskRouteKind::BatchImport, + ) + .await } pub(super) async fn handle_admin_provider_oauth_agent_identity_import_task_status( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, ) -> Result, GatewayError> { - handle_admin_provider_oauth_import_task_status(state, request_context, true).await + handle_admin_provider_oauth_import_task_status( + state, + request_context, + ProviderOAuthTaskRouteKind::AgentIdentity, + ) + .await +} + +pub(super) async fn handle_admin_provider_oauth_cookie_task_status( + state: &AdminAppState<'_>, + request_context: &AdminRequestContext<'_>, +) -> Result, GatewayError> { + handle_admin_provider_oauth_import_task_status( + state, + request_context, + ProviderOAuthTaskRouteKind::CookieAuthorize, + ) + .await } async fn handle_admin_provider_oauth_import_task_status( state: &AdminAppState<'_>, request_context: &AdminRequestContext<'_>, - agent_identity_only: bool, + route_kind: ProviderOAuthTaskRouteKind, ) -> Result, GatewayError> { - let task_path = if agent_identity_only { - admin_provider_oauth_agent_identity_import_task_path(request_context.path()) - } else { - admin_provider_oauth_batch_import_task_path(request_context.path()) + let not_found_detail = match route_kind { + ProviderOAuthTaskRouteKind::BatchImport => "批量导入任务不存在或已过期", + ProviderOAuthTaskRouteKind::AgentIdentity => "Agent Identity 导入任务不存在或已过期", + ProviderOAuthTaskRouteKind::CookieAuthorize => "Cookie 授权任务不存在或已过期", + }; + let task_path = match route_kind { + ProviderOAuthTaskRouteKind::BatchImport => { + admin_provider_oauth_batch_import_task_path(request_context.path()) + } + ProviderOAuthTaskRouteKind::AgentIdentity => { + admin_provider_oauth_agent_identity_import_task_path(request_context.path()) + } + ProviderOAuthTaskRouteKind::CookieAuthorize => { + admin_provider_oauth_cookie_task_path(request_context.path()) + } }; let Some((provider_id, task_id)) = task_path else { return Ok(build_internal_control_error_response( http::StatusCode::NOT_FOUND, - "批量导入任务不存在", + not_found_detail, )); }; let payload = match state @@ -69,7 +123,7 @@ async fn handle_admin_provider_oauth_import_task_status( Ok(None) => { return Ok(build_internal_control_error_response( http::StatusCode::NOT_FOUND, - "批量导入任务不存在或已过期", + not_found_detail, )); } Err(_) => { @@ -79,10 +133,10 @@ async fn handle_admin_provider_oauth_import_task_status( )); } }; - if !provider_oauth_import_task_matches_route(&task_id, &payload, agent_identity_only) { + if !provider_oauth_import_task_matches_route(&task_id, &payload, route_kind) { return Ok(build_internal_control_error_response( http::StatusCode::NOT_FOUND, - "导入任务不存在或已过期", + not_found_detail, )); } let status = payload @@ -91,20 +145,25 @@ async fn handle_admin_provider_oauth_import_task_status( .map(ToOwned::to_owned) .unwrap_or_default(); let response = Json(payload).into_response(); - let (completed_event, failed_event, action, target_type) = if agent_identity_only { - ( + let (completed_event, failed_event, action, target_type) = match route_kind { + ProviderOAuthTaskRouteKind::AgentIdentity => ( "admin_provider_oauth_agent_identity_import_completed_viewed", "admin_provider_oauth_agent_identity_import_failed_viewed", "view_provider_agent_identity_import_terminal_state", "provider_agent_identity_import_task", - ) - } else { - ( + ), + ProviderOAuthTaskRouteKind::CookieAuthorize => ( + "admin_provider_oauth_cookie_task_completed_viewed", + "admin_provider_oauth_cookie_task_failed_viewed", + "view_provider_oauth_cookie_task_terminal_state", + "provider_oauth_cookie_task", + ), + ProviderOAuthTaskRouteKind::BatchImport => ( "admin_provider_oauth_batch_task_completed_viewed", "admin_provider_oauth_batch_task_failed_viewed", "view_provider_oauth_batch_task_terminal_state", "provider_oauth_batch_task", - ) + ), }; Ok(match status.as_str() { "completed" => attach_admin_audit_response( @@ -127,38 +186,69 @@ async fn handle_admin_provider_oauth_import_task_status( #[cfg(test)] mod tests { - use super::provider_oauth_import_task_matches_route; + use super::{provider_oauth_import_task_matches_route, ProviderOAuthTaskRouteKind}; use serde_json::json; #[test] fn import_task_status_routes_are_bidirectionally_isolated() { let agent_payload = json!({ "import_kind": "agent_identity" }); let batch_payload = json!({ "import_kind": "oauth_batch" }); + let cookie_payload = json!({ "import_kind": "cookie_authorize" }); assert!(provider_oauth_import_task_matches_route( "agent-identity-task-1", &agent_payload, - true, + ProviderOAuthTaskRouteKind::AgentIdentity, )); assert!(!provider_oauth_import_task_matches_route( "agent-identity-task-1", &agent_payload, - false, + ProviderOAuthTaskRouteKind::BatchImport, )); assert!(provider_oauth_import_task_matches_route( "batch-task-1", &batch_payload, - false, + ProviderOAuthTaskRouteKind::BatchImport, )); assert!(!provider_oauth_import_task_matches_route( "batch-task-1", &batch_payload, - true, + ProviderOAuthTaskRouteKind::AgentIdentity, + )); + assert!(provider_oauth_import_task_matches_route( + "claude-cookie-task-1", + &cookie_payload, + ProviderOAuthTaskRouteKind::CookieAuthorize, + )); + for route_kind in [ + ProviderOAuthTaskRouteKind::BatchImport, + ProviderOAuthTaskRouteKind::AgentIdentity, + ] { + assert!(!provider_oauth_import_task_matches_route( + "claude-cookie-task-1", + &cookie_payload, + route_kind, + )); + } + for (task_id, payload) in [ + ("batch-task-1", &batch_payload), + ("agent-identity-task-1", &agent_payload), + ] { + assert!(!provider_oauth_import_task_matches_route( + task_id, + payload, + ProviderOAuthTaskRouteKind::CookieAuthorize, + )); + } + assert!(!provider_oauth_import_task_matches_route( + "claude-cookie-task-1", + &batch_payload, + ProviderOAuthTaskRouteKind::CookieAuthorize, )); assert!(provider_oauth_import_task_matches_route( "legacy-batch-task", &json!({}), - false, + ProviderOAuthTaskRouteKind::BatchImport, )); } } diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/token_import.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/token_import.rs index eb26af782..0cc4794f2 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/token_import.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/token_import.rs @@ -100,6 +100,12 @@ pub(super) fn normalize_provider_import_tokens( if provider_type == "grok" { return (None, access_token.or(refresh_token)); } + if provider_type == "claude_code" { + if access_token.is_none() && refresh_token.as_deref().is_some_and(is_claude_access_token) { + return (None, refresh_token); + } + return (refresh_token, access_token); + } normalize_single_import_tokens(refresh_token.as_deref(), access_token.as_deref()) } @@ -210,10 +216,75 @@ pub(super) fn provider_oauth_import_authorization_bearer_token_from_object( pub(super) fn provider_type_supports_access_token_import(provider_type: &str) -> bool { matches!( provider_type.trim().to_ascii_lowercase().as_str(), - "codex" | "chatgpt_web" | "grok" + "claude_code" | "codex" | "chatgpt_web" | "grok" ) } +pub(super) fn is_claude_access_token(value: &str) -> bool { + value.trim().starts_with("sk-ant-oat") +} + +pub(super) fn is_claude_session_key(value: &str) -> bool { + value.trim().starts_with("sk-ant-sid") +} + +pub(super) fn flatten_claude_code_credentials_payload(payload: &mut Map) { + let nested = payload + .get("claudeAiOauth") + .or_else(|| payload.get("claude_ai_oauth")) + .and_then(Value::as_object) + .cloned(); + let Some(nested) = nested else { + return; + }; + + for (target, aliases) in [ + ("access_token", &["access_token", "accessToken"][..]), + ("refresh_token", &["refresh_token", "refreshToken"][..]), + ("scopes", &["scopes"][..]), + ( + "subscription_type", + &["subscription_type", "subscriptionType"][..], + ), + ("rate_limit_tier", &["rate_limit_tier", "rateLimitTier"][..]), + ( + "organization_uuid", + &["organization_uuid", "organizationUuid", "org_uuid"][..], + ), + ] { + if payload.contains_key(target) { + continue; + } + if let Some(value) = aliases.iter().find_map(|key| nested.get(*key)).cloned() { + payload.insert(target.to_string(), value); + } + } + + if !payload.contains_key("expires_at") { + if let Some(expires_at) = json_u64_value(nested.get("expires_at")) { + payload.insert("expires_at".to_string(), json!(expires_at)); + } else if let Some(expires_at_ms) = json_u64_value(nested.get("expiresAt")) { + payload.insert("expires_at".to_string(), json!(expires_at_ms / 1_000)); + } + } +} + +pub(super) fn validate_claude_access_token_import( + access_token: &str, + imported_expires_at: Option, + now_unix_secs: u64, +) -> Result<(), &'static str> { + if !is_claude_access_token(access_token) { + return Err("Claude Access Token 格式无效,请导入 sk-ant-oat 凭据"); + } + if imported_expires_at.is_none_or(|expires_at| expires_at <= now_unix_secs) { + return Err( + "Claude Access Token 单独导入必须提供有效的未来 expires_at;建议导入完整 Claude credentials 或 Refresh Token", + ); + } + Ok(()) +} + pub(super) fn build_provider_access_token_import_auth_config( provider_type: &str, access_token: &str, @@ -266,9 +337,10 @@ pub(super) fn build_provider_access_token_import_auth_config( mod tests { use super::{ build_provider_access_token_import_auth_config, decode_access_token_expires_at, - looks_like_access_token, normalize_provider_import_tokens, - normalize_provider_oauth_import_headers, normalize_single_import_tokens, - provider_oauth_import_authorization_bearer_token, + flatten_claude_code_credentials_payload, looks_like_access_token, + normalize_provider_import_tokens, normalize_provider_oauth_import_headers, + normalize_single_import_tokens, provider_oauth_import_authorization_bearer_token, + validate_claude_access_token_import, }; use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _}; use serde_json::json; @@ -431,4 +503,64 @@ mod tests { Some(&json!(2_200_000_000u64)) ); } + + #[test] + fn flattens_only_claude_ai_oauth_credentials_and_converts_expiry_ms() { + let mut payload = json!({ + "claudeAiOauth": { + "accessToken": "sk-ant-oat01-access", + "refreshToken": "sk-ant-ort01-refresh", + "expiresAt": 2_100_000_000_123u64, + "scopes": ["user:profile"], + "subscriptionType": "pro", + "rateLimitTier": "tier_1", + "organizationUuid": "org-123" + }, + "mcpOAuth": { + "accessToken": "must-not-be-imported" + } + }) + .as_object() + .cloned() + .expect("payload should be an object"); + + flatten_claude_code_credentials_payload(&mut payload); + + assert_eq!( + payload.get("access_token"), + Some(&json!("sk-ant-oat01-access")) + ); + assert_eq!( + payload.get("refresh_token"), + Some(&json!("sk-ant-ort01-refresh")) + ); + assert_eq!(payload.get("expires_at"), Some(&json!(2_100_000_000u64))); + assert_eq!(payload.get("organization_uuid"), Some(&json!("org-123"))); + assert_ne!( + payload.get("access_token"), + Some(&json!("must-not-be-imported")) + ); + } + + #[test] + fn validates_claude_access_token_prefix_and_future_expiry() { + assert!(validate_claude_access_token_import( + "sk-ant-oat01-access", + Some(2_100_000_000), + 2_000_000_000, + ) + .is_ok()); + assert!(validate_claude_access_token_import( + "arbitrary-token", + Some(2_100_000_000), + 2_000_000_000, + ) + .is_err()); + assert!(validate_claude_access_token_import( + "sk-ant-oat01-expired", + Some(1_900_000_000), + 2_000_000_000, + ) + .is_err()); + } } diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/duplicates.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/duplicates.rs index 8bf4b85a5..40cd83e7b 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/duplicates.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/duplicates.rs @@ -8,6 +8,7 @@ use std::time::{Duration, SystemTime, UNIX_EPOCH}; use uuid::Uuid; const CODEX_OAUTH_ACCOUNT_LOCK_TTL: Duration = Duration::from_secs(180); +const CLAUDE_OAUTH_ACCOUNT_LOCK_TTL: Duration = Duration::from_secs(180); #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub(crate) enum CodexOAuthAccountLockError { @@ -34,6 +35,31 @@ impl CodexOAuthAccountLockError { } } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum ClaudeOAuthAccountLockError { + MissingIdentity, + Contended, + Unavailable, +} + +impl ClaudeOAuthAccountLockError { + pub(crate) const fn status_code(self) -> http::StatusCode { + match self { + Self::MissingIdentity => http::StatusCode::BAD_REQUEST, + Self::Contended => http::StatusCode::CONFLICT, + Self::Unavailable => http::StatusCode::SERVICE_UNAVAILABLE, + } + } + + pub(crate) const fn detail(self) -> &'static str { + match self { + Self::MissingIdentity => "Claude 账号身份字段缺失,无法安全写入授权", + Self::Contended => "该 Claude 账号正在更新授权,请稍后重试", + Self::Unavailable => "Claude 账号授权锁暂不可用,请稍后重试", + } + } +} + fn normalize_codex_plan_group_for_provider_oauth( plan_type: Option<&serde_json::Value>, ) -> Option { @@ -208,23 +234,96 @@ pub(crate) async fn acquire_codex_oauth_account_locks( pub(crate) async fn release_codex_oauth_account_locks( state: &AdminAppState<'_>, leases: Vec, +) { + release_provider_oauth_account_locks(state, leases).await; +} + +pub(crate) async fn release_provider_oauth_account_locks( + state: &AdminAppState<'_>, + leases: Vec, ) { for lease in leases.into_iter().rev() { match state.runtime_state().lock_release(&lease).await { Ok(true) => {} Ok(false) => tracing::warn!( lock_key = %lease.key, - "gateway Codex OAuth account lock was not owned during release" + "gateway provider OAuth account lock was not owned during release" ), Err(error) => tracing::warn!( lock_key = %lease.key, error = ?error, - "gateway Codex OAuth account lock release failed" + "gateway provider OAuth account lock release failed" ), } } } +fn claude_oauth_account_lock_key( + provider_id: &str, + auth_config: &serde_json::Map, +) -> Option { + let (identity_kind, identity) = normalize_provider_oauth_identity_value_from_keys( + auth_config, + &["account_uuid", "accountUuid"], + ) + .map(|value| ("account_uuid", value)) + .or_else(|| { + normalize_provider_oauth_identity_value_from_keys(auth_config, &["email"]) + .map(|value| ("email", value.to_ascii_lowercase())) + })?; + let mut digest = Sha256::new(); + digest.update(provider_id.trim().as_bytes()); + digest.update([0]); + digest.update(identity_kind.as_bytes()); + digest.update([0]); + digest.update(identity.as_bytes()); + Some(format!( + "provider_oauth_claude_account:{:x}", + digest.finalize() + )) +} + +pub(crate) async fn acquire_claude_oauth_account_lock( + state: &AdminAppState<'_>, + provider_id: &str, + auth_config: &serde_json::Map, + operation: &str, +) -> Result, ClaudeOAuthAccountLockError> { + let Some(lock_key) = claude_oauth_account_lock_key(provider_id, auth_config) else { + return Err(ClaudeOAuthAccountLockError::MissingIdentity); + }; + let owner = format!( + "aether-gateway-claude-oauth-{}-{}", + operation.trim(), + Uuid::new_v4() + ); + let lease = match state + .runtime_state() + .lock_try_acquire( + lock_key.as_str(), + owner.as_str(), + CLAUDE_OAUTH_ACCOUNT_LOCK_TTL, + ) + .await + { + Ok(Some(lease)) => lease, + Ok(None) => return Err(ClaudeOAuthAccountLockError::Contended), + Err(error) => { + tracing::warn!( + provider_id = %provider_id, + lock_key = %lock_key, + operation, + error = ?error, + "gateway Claude OAuth account lock unavailable" + ); + return Err(ClaudeOAuthAccountLockError::Unavailable); + } + }; + + state.app().data.clear_provider_catalog_cache(); + Ok(vec![lease]) +} + fn is_openai_provider_oauth_provider_type(value: Option<&serde_json::Value>) -> bool { value .and_then(serde_json::Value::as_str) @@ -242,6 +341,37 @@ fn is_windsurf_provider_oauth_provider_type(value: Option<&serde_json::Value>) - .is_some_and(|provider_type| provider_type.eq_ignore_ascii_case("windsurf")) } +fn is_claude_provider_oauth_provider_type(value: Option<&serde_json::Value>) -> bool { + value + .and_then(serde_json::Value::as_str) + .map(str::trim) + .is_some_and(|provider_type| provider_type.eq_ignore_ascii_case("claude_code")) +} + +fn match_claude_provider_oauth_identity( + new_auth_config: &serde_json::Map, + existing_auth_config: &serde_json::Map, +) -> Option { + if !is_claude_provider_oauth_provider_type(new_auth_config.get("provider_type")) + && !is_claude_provider_oauth_provider_type(existing_auth_config.get("provider_type")) + { + return None; + } + + let new_account_uuid = normalize_provider_oauth_identity_value_from_keys( + new_auth_config, + &["account_uuid", "accountUuid"], + ); + let existing_account_uuid = normalize_provider_oauth_identity_value_from_keys( + existing_auth_config, + &["account_uuid", "accountUuid"], + ); + match (new_account_uuid, existing_account_uuid) { + (Some(left), Some(right)) => Some(left == right), + _ => None, + } +} + fn match_codex_provider_oauth_identity( new_auth_config: &serde_json::Map, existing_auth_config: &serde_json::Map, @@ -429,6 +559,10 @@ pub(crate) async fn find_duplicate_provider_oauth_key( let new_email = normalize_provider_oauth_identity_value(auth_config.get("email")); let new_user_id = normalize_provider_oauth_identity_value(auth_config.get("user_id")); let new_account_id = normalize_provider_oauth_identity_value(auth_config.get("account_id")); + let new_account_uuid = normalize_provider_oauth_identity_value_from_keys( + auth_config, + &["account_uuid", "accountUuid"], + ); let new_agent_runtime_id = normalize_provider_oauth_identity_value( auth_config .get("agent_runtime_id") @@ -442,6 +576,7 @@ pub(crate) async fn find_duplicate_provider_oauth_key( if new_email.is_none() && new_user_id.is_none() && new_account_id.is_none() + && new_account_uuid.is_none() && new_agent_runtime_id.is_none() && new_credential_fingerprint.is_none() { @@ -486,17 +621,22 @@ pub(crate) async fn find_duplicate_provider_oauth_key( .is_some_and(|value| value.eq_ignore_ascii_case("windsurf")); let mut is_duplicate = false; + let claude_identity_match = + match_claude_provider_oauth_identity(auth_config, &existing_auth_config); let codex_identity_match = match_codex_provider_oauth_identity(auth_config, &existing_auth_config); let windsurf_identity_match = match_windsurf_provider_oauth_identity(auth_config, &existing_auth_config); - if let Some(codex_identity_match) = codex_identity_match { + if let Some(claude_identity_match) = claude_identity_match { + is_duplicate = claude_identity_match; + } else if let Some(codex_identity_match) = codex_identity_match { is_duplicate = codex_identity_match; } else if let Some(windsurf_identity_match) = windsurf_identity_match { is_duplicate = windsurf_identity_match; } - if codex_identity_match.is_none() + if claude_identity_match.is_none() + && codex_identity_match.is_none() && windsurf_identity_match.is_none() && !is_duplicate && new_user_id.is_some() @@ -507,7 +647,8 @@ pub(crate) async fn find_duplicate_provider_oauth_key( is_duplicate = true; } - if codex_identity_match.is_none() + if claude_identity_match.is_none() + && codex_identity_match.is_none() && windsurf_identity_match.is_none() && !is_duplicate && !is_windsurf @@ -548,19 +689,20 @@ pub(crate) async fn find_duplicate_provider_oauth_key( if existing_provider_oauth_key_is_replaceable(&existing_key) { return Ok(Some(existing_key)); } - let identifier = - normalize_provider_oauth_identity_value(auth_config.get("account_user_id")) - .or_else(|| normalize_provider_oauth_identity_value(auth_config.get("account_id"))) - .or_else(|| new_agent_runtime_id.clone()) - .or_else(|| { - normalize_provider_oauth_identity_value( - auth_config.get("credential_fingerprint"), - ) - .map(|value| format!("fingerprint:{value}")) - }) - .or_else(|| new_email.clone()) - .or_else(|| new_user_id.clone()) - .unwrap_or_default(); + let identifier = normalize_provider_oauth_identity_value_from_keys( + auth_config, + &["account_uuid", "accountUuid"], + ) + .or_else(|| normalize_provider_oauth_identity_value(auth_config.get("account_user_id"))) + .or_else(|| normalize_provider_oauth_identity_value(auth_config.get("account_id"))) + .or_else(|| new_agent_runtime_id.clone()) + .or_else(|| { + normalize_provider_oauth_identity_value(auth_config.get("credential_fingerprint")) + .map(|value| format!("fingerprint:{value}")) + }) + .or_else(|| new_email.clone()) + .or_else(|| new_user_id.clone()) + .unwrap_or_default(); return Err(format!( "该 OAuth 账号 ({identifier}) 已存在于当前 Provider 中(名称: {})", existing_key.name @@ -573,9 +715,12 @@ pub(crate) async fn find_duplicate_provider_oauth_key( #[cfg(test)] mod tests { use super::{ - acquire_codex_oauth_account_locks, codex_agent_identity_account_lock_keys, - match_codex_provider_oauth_identity, match_windsurf_provider_oauth_identity, - release_codex_oauth_account_locks, CodexOAuthAccountLockError, + acquire_claude_oauth_account_lock, acquire_codex_oauth_account_locks, + claude_oauth_account_lock_key, codex_agent_identity_account_lock_keys, + match_claude_provider_oauth_identity, match_codex_provider_oauth_identity, + match_windsurf_provider_oauth_identity, release_codex_oauth_account_locks, + release_provider_oauth_account_locks, ClaudeOAuthAccountLockError, + CodexOAuthAccountLockError, }; use crate::handlers::admin::request::AdminAppState; use crate::AppState; @@ -604,6 +749,99 @@ mod tests { ); } + #[test] + fn claude_identity_prefers_account_uuid_without_using_organization_uuid() { + let new_auth_config = auth_config(json!({ + "provider_type": "claude_code", + "account_uuid": "account-1", + "org_uuid": "shared-org" + })); + let same_account = auth_config(json!({ + "provider_type": "claude_code", + "accountUuid": "account-1", + "org_uuid": "other-org" + })); + let different_account = auth_config(json!({ + "provider_type": "claude_code", + "account_uuid": "account-2", + "email": "same@example.com", + "org_uuid": "shared-org" + })); + let organization_only = auth_config(json!({ + "provider_type": "claude_code", + "org_uuid": "shared-org" + })); + + assert_eq!( + match_claude_provider_oauth_identity(&new_auth_config, &same_account), + Some(true) + ); + assert_eq!( + match_claude_provider_oauth_identity(&new_auth_config, &different_account), + Some(false) + ); + assert_eq!( + match_claude_provider_oauth_identity(&new_auth_config, &organization_only), + None + ); + } + + #[test] + fn claude_account_lock_prefers_uuid_and_falls_back_to_normalized_email() { + let with_uuid = auth_config(json!({ + "account_uuid": "account-1", + "email": "first@example.com" + })); + let same_uuid_other_email = auth_config(json!({ + "accountUuid": "account-1", + "email": "other@example.com" + })); + let email_only_uppercase = auth_config(json!({"email": "User@Example.COM"})); + let email_only_lowercase = auth_config(json!({"email": "user@example.com"})); + + let uuid_key = claude_oauth_account_lock_key("provider-claude", &with_uuid) + .expect("uuid lock key should build"); + assert_eq!( + Some(uuid_key.as_str()), + claude_oauth_account_lock_key("provider-claude", &same_uuid_other_email).as_deref() + ); + assert_eq!( + claude_oauth_account_lock_key("provider-claude", &email_only_uppercase), + claude_oauth_account_lock_key("provider-claude", &email_only_lowercase) + ); + assert!(uuid_key.starts_with("provider_oauth_claude_account:")); + assert!(!uuid_key.contains("account-1")); + assert!(claude_oauth_account_lock_key("provider-claude", &Map::new()).is_none()); + } + + #[tokio::test] + async fn concurrent_claude_writes_contend_on_the_same_account_lock() { + let app = AppState::new().expect("app state should build"); + let state = AdminAppState::new(&app); + let config = auth_config(json!({ + "provider_type": "claude_code", + "account_uuid": "account-shared", + "email": "claude@example.com" + })); + + let first = + acquire_claude_oauth_account_lock(&state, "provider-claude", &config, "first-test") + .await + .expect("first Claude lock should acquire"); + let second = + acquire_claude_oauth_account_lock(&state, "provider-claude", &config, "second-test") + .await + .expect_err("second Claude lock should contend"); + assert_eq!(second, ClaudeOAuthAccountLockError::Contended); + + release_provider_oauth_account_locks(&state, first).await; + let retry = + acquire_claude_oauth_account_lock(&state, "provider-claude", &config, "retry-test") + .await + .expect("Claude lock should be reusable after release"); + release_provider_oauth_account_locks(&state, retry).await; + } + #[test] fn codex_agent_identity_matches_runtime_without_account_metadata() { let new_auth_config = auth_config(json!({ diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/provisioning.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/provisioning.rs index efd12abe5..5c3eaafa2 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/provisioning.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/provisioning.rs @@ -1,3 +1,9 @@ +use super::duplicates::{ + acquire_claude_oauth_account_lock, acquire_codex_oauth_account_locks, + release_provider_oauth_account_locks, +}; +use super::errors::build_internal_control_error_response; +use super::runtime::spawn_provider_oauth_account_state_refresh_after_update; use super::state::{ decode_jwt_claims, enrich_admin_provider_oauth_auth_config, json_non_empty_string, json_u64_value, @@ -9,15 +15,22 @@ use crate::handlers::admin::admin_provider_pool_config; use crate::handlers::admin::request::AdminAppState; use crate::provider_key_auth::provider_active_api_formats; use crate::GatewayError; +use aether_contracts::ProxySnapshot; use aether_data_contracts::repository::pool_scores::{ GetPoolMemberScoresByIdsQuery, PoolMemberIdentity, }; use aether_data_contracts::repository::provider_catalog::{ - StoredProviderCatalogEndpoint, StoredProviderCatalogKey, + StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, }; use aether_provider_transport::{ grok_browser_transport_fingerprint_from_auth_config, provider_types::provider_type_is_fixed, }; +use axum::{ + body::Body, + http, + response::{IntoResponse, Response}, + Json, +}; use serde_json::{json, Map, Value}; use std::time::{SystemTime, UNIX_EPOCH}; use uuid::Uuid; @@ -101,6 +114,168 @@ pub(crate) fn build_provider_oauth_auth_config_from_token_payload( (auth_config, access_token, refresh_token, expires_at) } +pub(crate) async fn provision_provider_oauth_token_payload_for_provider( + state: &AdminAppState<'_>, + provider: &StoredProviderCatalogProvider, + endpoints: &[StoredProviderCatalogEndpoint], + token_payload: &Value, + requested_name: Option, + key_proxy: Option, + request_proxy: Option, + lock_operation: &'static str, +) -> Result, GatewayError> { + let provider_id = provider.id.clone(); + let provider_type = provider.provider_type.trim().to_ascii_lowercase(); + let (auth_config, access_token, refresh_token, expires_at) = + build_provider_oauth_auth_config_from_token_payload(&provider_type, token_payload); + let Some(access_token) = access_token else { + return Ok(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "token exchange 返回缺少 access_token", + )); + }; + + let api_formats = provider_oauth_active_api_formats(endpoints); + let oauth_account_leases = if provider_type == "codex" { + match acquire_codex_oauth_account_locks(state, &provider_id, &auth_config, lock_operation) + .await + { + Ok(leases) => leases, + Err(error) => { + return Ok(build_internal_control_error_response( + error.status_code(), + error.detail(), + )); + } + } + } else if provider_type == "claude_code" { + match acquire_claude_oauth_account_lock(state, &provider_id, &auth_config, lock_operation) + .await + { + Ok(leases) => leases, + Err(error) => { + return Ok(build_internal_control_error_response( + error.status_code(), + error.detail(), + )); + } + } + } else { + Vec::new() + }; + let duplicate = match state + .find_duplicate_provider_oauth_key(&provider_id, &auth_config, None) + .await + { + Ok(duplicate) => duplicate, + Err(detail) => { + release_provider_oauth_account_locks(state, oauth_account_leases).await; + return Ok(build_internal_control_error_response( + if provider_type == "codex" { + http::StatusCode::CONFLICT + } else { + http::StatusCode::BAD_REQUEST + }, + detail, + )); + } + }; + + let replaced = duplicate.is_some(); + let persisted_key = if let Some(existing_key) = duplicate { + match state + .update_existing_provider_oauth_catalog_key( + &existing_key, + &provider_type, + &access_token, + &auth_config, + &api_formats, + key_proxy.clone(), + expires_at, + ) + .await + { + Err(error) => { + release_provider_oauth_account_locks(state, oauth_account_leases).await; + return Err(error); + } + Ok(Some(key)) => key, + Ok(None) => { + release_provider_oauth_account_locks(state, oauth_account_leases).await; + return Ok(build_internal_control_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "provider oauth write unavailable", + )); + } + } + } else { + let name = requested_name + .or_else(|| { + auth_config + .get("email") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + }) + .unwrap_or_else(|| { + format!( + "账号_{}", + SystemTime::now() + .duration_since(UNIX_EPOCH) + .ok() + .map(|duration| duration.as_secs()) + .unwrap_or(0) + ) + }); + match state + .create_provider_oauth_catalog_key( + &provider_id, + &provider_type, + &name, + &access_token, + &auth_config, + &api_formats, + key_proxy, + expires_at, + ) + .await + { + Err(error) => { + release_provider_oauth_account_locks(state, oauth_account_leases).await; + return Err(error); + } + Ok(Some(key)) => key, + Ok(None) => { + release_provider_oauth_account_locks(state, oauth_account_leases).await; + return Ok(build_internal_control_error_response( + http::StatusCode::SERVICE_UNAVAILABLE, + "provider oauth write unavailable", + )); + } + } + }; + release_provider_oauth_account_locks(state, oauth_account_leases).await; + + spawn_provider_oauth_account_state_refresh_after_update( + state.cloned_app(), + provider.clone(), + persisted_key.id.clone(), + request_proxy, + ); + + Ok(Json(json!({ + "key_id": persisted_key.id, + "provider_type": provider_type, + "expires_at": expires_at, + "has_refresh_token": refresh_token.is_some(), + "temporary": refresh_token.is_none(), + "email": auth_config.get("email").cloned().unwrap_or(Value::Null), + "replaced": replaced, + })) + .into_response()) +} + fn grok_oauth_catalog_key_fingerprint( provider_type: &str, auth_config: &Map, diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/state/exchange.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/state/exchange.rs index 5dce3cbaa..bfda465fb 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/state/exchange.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/state/exchange.rs @@ -3,8 +3,13 @@ use super::super::errors::{ }; use crate::handlers::admin::request::{AdminAppState, AdminProviderOAuthTemplate}; use aether_contracts::ProxySnapshot; -use aether_oauth::provider::providers::GenericProviderOAuthAdapter; -use aether_oauth::provider::{ProviderOAuthService, ProviderOAuthTransportContext}; +use aether_oauth::provider::providers::{ + ClaudeCodeProviderOAuthAdapter, GenericProviderOAuthAdapter, CLAUDE_CODE_PROVIDER_TYPE, + CLAUDE_CODE_TOKEN_URL, CLAUDE_CODE_WEB_BASE_URL, +}; +use aether_oauth::provider::{ + ProviderOAuthCookieAuthorizationInput, ProviderOAuthService, ProviderOAuthTransportContext, +}; use axum::{body::Body, http, response::Response}; use std::sync::Arc; @@ -140,3 +145,40 @@ pub(crate) async fn exchange_admin_provider_oauth_refresh_token( ) }) } + +pub(crate) async fn authorize_admin_provider_oauth_with_cookie( + state: &AdminAppState<'_>, + session_key: String, + proxy: Option, +) -> Result> { + let web_base_url = + state.provider_oauth_token_url("claude_code_cookie_base_url", CLAUDE_CODE_WEB_BASE_URL); + let token_url = + state.provider_oauth_token_url(CLAUDE_CODE_PROVIDER_TYPE, CLAUDE_CODE_TOKEN_URL); + let service = ProviderOAuthService::new().with_adapter(Arc::new( + ClaudeCodeProviderOAuthAdapter::default().with_endpoint_overrides(web_base_url, token_url), + )); + let ctx = provider_oauth_exchange_context(CLAUDE_CODE_PROVIDER_TYPE, proxy); + let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state); + let result = service + .authorize_with_cookie( + &executor, + &ctx, + ProviderOAuthCookieAuthorizationInput { session_key }, + ) + .await + .map_err(|error| { + let detail = if matches!(error, aether_oauth::core::OAuthError::InvalidRequest(_)) { + "Claude Cookie 格式无效" + } else { + "Claude Cookie 授权失败" + }; + build_internal_control_error_response(http::StatusCode::BAD_REQUEST, detail) + })?; + token_payload_from_provider_oauth_result(result).map_err(|_| { + build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "Claude Cookie 授权返回缺少 access_token", + ) + }) +} diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/state/mod.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/state/mod.rs index ccca0fff2..7fb1936e7 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/state/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/state/mod.rs @@ -5,7 +5,8 @@ mod template; pub(crate) use self::auth_config::enrich_admin_provider_oauth_auth_config; pub(crate) use self::exchange::{ - exchange_admin_provider_oauth_code, exchange_admin_provider_oauth_refresh_token, + authorize_admin_provider_oauth_with_cookie, exchange_admin_provider_oauth_code, + exchange_admin_provider_oauth_refresh_token, }; pub(crate) use self::storage::build_provider_oauth_start_response; pub(crate) use self::template::{ diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/state/storage.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/state/storage.rs index 2a13a76bf..aa2c8f384 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/state/storage.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/state/storage.rs @@ -17,7 +17,7 @@ pub(crate) fn build_provider_oauth_start_response( "authorization_url": authorization_url, "redirect_uri": template.redirect_uri, "provider_type": template.provider_type, - "instructions": "1) 打开 authorization_url 完成授权\n2) 授权后会跳转到 redirect_uri(localhost)\n3) 复制浏览器地址栏完整 URL,调用 complete 接口粘贴 callback_url", + "instructions": "1) 打开 authorization_url 完成授权\n2) 复制授权页面显示的授权码或浏览器中的完整回调 URL\n3) 调用 complete 接口粘贴 callback_url", }) } diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/state/template.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/state/template.rs index b1a36b6eb..e332080b7 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/state/template.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/state/template.rs @@ -20,9 +20,14 @@ pub(crate) fn admin_provider_oauth_template( } pub(crate) fn build_admin_provider_oauth_supported_types_payload() -> Vec { + let service = aether_oauth::provider::ProviderOAuthService::with_builtin_adapters(); admin_provider_oauth_template_types() .filter_map(|provider_type| admin_provider_oauth_template(provider_type)) .map(|template| { + let capabilities = service + .adapter(template.provider_type) + .ok() + .map(|adapter| adapter.capabilities()); json!({ "provider_type": template.provider_type, "display_name": template.display_name, @@ -31,6 +36,10 @@ pub(crate) fn build_admin_provider_oauth_supported_types_payload() -> Vec) -> PoolMemberScore fn normalize_pool_preset_mode(preset: &str, raw_mode: Option<&Value>) -> Option { match preset { - "free_first" | "team_first" | "plus_first" | "pro_first" => { + "free_team_first" | "free_first" | "team_first" | "plus_first" | "pro_first" => { let default_mode = match preset { + "free_team_first" => "both", "free_first" => "free_only", "team_first" => "team_only", "plus_first" => "plus_only", @@ -180,6 +182,9 @@ fn normalize_pool_preset_mode(preset: &str, raw_mode: Option<&Value>) -> Option< .filter(|value| !value.is_empty()) .map(|value| value.to_ascii_lowercase()) .filter(|value| match preset { + "free_team_first" => { + matches!(value.as_str(), "free_only" | "team_only" | "both") + } "free_first" => value == "free_only", "team_first" => value == "team_only", "plus_first" => value == "plus_only", @@ -786,7 +791,7 @@ mod tests { } #[test] - fn retired_free_team_first_preset_is_rejected() { + fn legacy_free_team_first_preset_is_preserved() { let config = admin_provider_pool_config_from_config_value(Some(&json!({ "pool_advanced": { "scheduling_presets": [ @@ -797,8 +802,11 @@ mod tests { .expect("pool config should parse"); assert_eq!(config.scheduling_presets.len(), 1); - assert_eq!(config.scheduling_presets[0].preset, "lru"); - assert_eq!(config.scheduling_presets[0].mode, None); + assert_eq!(config.scheduling_presets[0].preset, "free_team_first"); + assert_eq!( + config.scheduling_presets[0].mode.as_deref(), + Some("team_only") + ); } #[test] diff --git a/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test.rs b/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test.rs index e96973d01..86cdf14be 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test.rs @@ -65,6 +65,8 @@ use aether_model_fetch::{ aggregate_models_for_cache, fetch_models_from_transports, json_string_list, preset_models_for_provider, selected_models_fetch_endpoints, }; +use aether_pool_core::PoolSchedulingPreset; +use aether_provider_pool::ProviderPoolService; use aether_scheduler_core::provider_key_circuit_payload_is_active_open_at; use axum::{ body::{to_bytes, Body}, @@ -1106,15 +1108,26 @@ fn provider_query_pool_sort_seed() -> String { fn provider_query_ai_pool_scheduling_config( config: &AdminProviderPoolConfig, + provider_type: &str, ) -> AiPoolSchedulingConfig { + let presets = config + .scheduling_presets + .iter() + .map(|preset| PoolSchedulingPreset { + preset: preset.preset.clone(), + enabled: preset.enabled, + mode: preset.mode.clone(), + }) + .collect::>(); + let normalized_presets = ProviderPoolService::with_builtin_adapters() + .normalize_scheduling_presets(provider_type, &presets); AiPoolSchedulingConfig { - scheduling_presets: config - .scheduling_presets - .iter() + scheduling_presets: normalized_presets + .into_iter() .map(|preset| AiPoolSchedulingPreset { - preset: preset.preset.clone(), + preset: preset.preset, enabled: preset.enabled, - mode: preset.mode.clone(), + mode: preset.mode, }) .collect(), lru_enabled: config.lru_enabled, @@ -1324,7 +1337,8 @@ async fn provider_query_apply_pool_scheduler_to_test_candidates( provider.id.clone(), provider_query_ai_pool_runtime_state(&runtime), ); - let pool_config = provider_query_ai_pool_scheduling_config(pool_config); + let pool_config = + provider_query_ai_pool_scheduling_config(pool_config, provider.provider_type.as_str()); let inputs = keys .into_iter() .map(|key| { diff --git a/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test/tests.rs b/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test/tests.rs index dd7285a5f..4c050c452 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test/tests.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test/tests.rs @@ -320,6 +320,31 @@ fn provider_query_model_test_empty_selected_key_ids_keep_default_selection() { assert!(provider_query_extract_api_key_ids(&json!({ "api_key_ids": [] })).is_none()); } +#[test] +fn provider_query_codex_pool_config_injects_recent_refresh() { + let raw_config = json!({ + "pool_advanced": { + "scheduling_presets": [{ + "preset": "cache_affinity", + "enabled": true + }] + } + }); + let config = admin_provider_pool_config_from_config_value(Some(&raw_config)) + .expect("pool config should parse"); + + let normalized = provider_query_ai_pool_scheduling_config(&config, "codex"); + + assert_eq!( + normalized + .scheduling_presets + .iter() + .map(|preset| preset.preset.as_str()) + .collect::>(), + ["cache_affinity", "recent_refresh"] + ); +} + #[test] fn provider_query_standard_test_resolves_codex_responses_upstream_streaming() { assert!(provider_query_resolve_standard_test_upstream_is_stream( diff --git a/apps/aether-gateway/src/handlers/admin/provider/shared/paths/mod.rs b/apps/aether-gateway/src/handlers/admin/provider/shared/paths/mod.rs index af35e1780..d8e49f74e 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/shared/paths/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/shared/paths/mod.rs @@ -24,7 +24,9 @@ pub(crate) use self::oauth::{ admin_provider_oauth_agent_identity_import_task_provider_id, admin_provider_oauth_batch_import_provider_id, admin_provider_oauth_batch_import_task_path, admin_provider_oauth_batch_import_task_provider_id, admin_provider_oauth_complete_key_id, - admin_provider_oauth_complete_provider_id, admin_provider_oauth_device_authorize_provider_id, + admin_provider_oauth_complete_provider_id, admin_provider_oauth_cookie_provider_id, + admin_provider_oauth_cookie_task_path, admin_provider_oauth_cookie_task_provider_id, + admin_provider_oauth_device_authorize_provider_id, admin_provider_oauth_device_poll_provider_id, admin_provider_oauth_import_provider_id, admin_provider_oauth_refresh_key_id, admin_provider_oauth_start_key_id, admin_provider_oauth_start_provider_id, diff --git a/apps/aether-gateway/src/handlers/admin/provider/shared/paths/oauth.rs b/apps/aether-gateway/src/handlers/admin/provider/shared/paths/oauth.rs index 4b6eae70f..8ac2a969a 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/shared/paths/oauth.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/shared/paths/oauth.rs @@ -34,6 +34,33 @@ pub(crate) fn admin_provider_oauth_import_provider_id(request_path: &str) -> Opt provider_oauth_provider_id_for_suffix(request_path, "/import-refresh-token") } +pub(crate) fn admin_provider_oauth_cookie_provider_id(request_path: &str) -> Option { + provider_oauth_provider_id_for_suffix(request_path, "/cookie-authorize") +} + +pub(crate) fn admin_provider_oauth_cookie_task_provider_id(request_path: &str) -> Option { + provider_oauth_provider_id_for_suffix(request_path, "/cookie-authorize/tasks") +} + +pub(crate) fn admin_provider_oauth_cookie_task_path( + request_path: &str, +) -> Option<(String, String)> { + let suffix = request_path + .strip_prefix("/api/admin/provider-oauth/providers/")? + .strip_suffix("/") + .unwrap_or(request_path.strip_prefix("/api/admin/provider-oauth/providers/")?); + let (provider_id, task_id) = suffix.split_once("/cookie-authorize/tasks/")?; + if provider_id.is_empty() + || provider_id.contains('/') + || task_id.is_empty() + || task_id.contains('/') + || !task_id.starts_with("claude-cookie-") + { + return None; + } + Some((provider_id.to_string(), task_id.to_string())) +} + pub(crate) fn admin_provider_oauth_batch_import_provider_id(request_path: &str) -> Option { provider_oauth_provider_id_for_suffix(request_path, "/batch-import") } @@ -110,6 +137,7 @@ mod tests { use super::{ admin_provider_oauth_agent_identity_import_task_path, admin_provider_oauth_agent_identity_import_task_provider_id, + admin_provider_oauth_cookie_task_path, admin_provider_oauth_cookie_task_provider_id, }; #[test] @@ -139,4 +167,28 @@ mod tests { ) .is_none()); } + + #[test] + fn parses_dedicated_claude_cookie_task_paths() { + assert_eq!( + admin_provider_oauth_cookie_task_provider_id( + "/api/admin/provider-oauth/providers/provider-claude/cookie-authorize/tasks", + ) + .as_deref(), + Some("provider-claude") + ); + assert_eq!( + admin_provider_oauth_cookie_task_path( + "/api/admin/provider-oauth/providers/provider-claude/cookie-authorize/tasks/claude-cookie-task-1", + ), + Some(( + "provider-claude".to_string(), + "claude-cookie-task-1".to_string(), + )) + ); + assert!(admin_provider_oauth_cookie_task_path( + "/api/admin/provider-oauth/providers/provider-claude/cookie-authorize/tasks/task-1", + ) + .is_none()); + } } diff --git a/apps/aether-gateway/src/handlers/admin/request/provider/oauth.rs b/apps/aether-gateway/src/handlers/admin/request/provider/oauth.rs index 7b93dfe2c..3306f8764 100644 --- a/apps/aether-gateway/src/handlers/admin/request/provider/oauth.rs +++ b/apps/aether-gateway/src/handlers/admin/request/provider/oauth.rs @@ -555,6 +555,7 @@ impl<'a> AdminAppState<'a> { json_body, body_bytes, network, + transport_profile: None, }; let response = aether_oauth::network::OAuthHttpExecutor::execute( &crate::oauth::GatewayOAuthHttpExecutor::new(*self), diff --git a/apps/aether-gateway/src/handlers/proxy/mod.rs b/apps/aether-gateway/src/handlers/proxy/mod.rs index c9fd5dcdc..8808c2077 100644 --- a/apps/aether-gateway/src/handlers/proxy/mod.rs +++ b/apps/aether-gateway/src/handlers/proxy/mod.rs @@ -346,20 +346,6 @@ async fn maybe_forward_public_request_to_tunnel_owner( }) else { return Ok(None); }; - let cache_affinity_enabled = match read_scheduler_ordering_config(state).await { - Ok(config) => config.scheduling_mode == SchedulerSchedulingMode::CacheAffinity, - Err(err) => { - warn!( - trace_id = %request_context.trace_id, - error = ?err, - "gateway failed to load scheduler config while checking tunnel affinity forwarding mode" - ); - SchedulerSchedulingMode::default() == SchedulerSchedulingMode::CacheAffinity - } - }; - if !cache_affinity_enabled { - return Ok(None); - } let Some(api_format) = decision .auth_endpoint_signature .as_deref() @@ -382,21 +368,66 @@ async fn maybe_forward_public_request_to_tunnel_owner( crate::headers::decoded_request_body_bytes(&parts.headers, body.as_ref()).ok()?; serde_json::from_slice::(body.as_ref()).ok() }); - let client_session_affinity = - crate::client_session_affinity::client_session_affinity_from_api_request( - api_format, - &parts.headers, - body_json.as_ref(), - ); - let Some(target) = crate::scheduler::affinity::read_cached_scheduler_affinity_target( + let empty_body_json = serde_json::Value::Null; + let affinity_context = match crate::ai_serving::resolve_tunnel_scheduler_affinity_context( state, - &auth_context.api_key_id, - client_session_affinity.as_ref(), + parts, + decision, + requested_model, + body_json.as_ref().unwrap_or(&empty_body_json), api_format, - &requested_model, - ) else { + ) + .await + { + Ok(Some(context)) => context, + Ok(None) => return Ok(None), + Err(err) => { + warn!( + trace_id = %request_context.trace_id, + error = ?err, + "gateway failed to resolve routing policy while checking tunnel affinity forwarding" + ); + return Ok(None); + } + }; + let target = if let Some(policy_context) = affinity_context.policy_context.as_ref() { + crate::scheduler::affinity::read_cached_scheduler_affinity_target_with_policy_context( + state, + &auth_context.api_key_id, + affinity_context.client_session_affinity.as_ref(), + api_format, + &affinity_context.requested_model, + policy_context, + ) + } else { + let cache_affinity_enabled = match read_scheduler_ordering_config(state).await { + Ok(config) => config.scheduling_mode == SchedulerSchedulingMode::CacheAffinity, + Err(err) => { + warn!( + trace_id = %request_context.trace_id, + error = ?err, + "gateway failed to load scheduler config while checking tunnel affinity forwarding mode" + ); + SchedulerSchedulingMode::default() == SchedulerSchedulingMode::CacheAffinity + } + }; + if !cache_affinity_enabled { + return Ok(None); + } + crate::scheduler::affinity::read_cached_scheduler_affinity_target( + state, + &auth_context.api_key_id, + affinity_context.client_session_affinity.as_ref(), + api_format, + &affinity_context.requested_model, + ) + }; + let Some(target) = target else { return Ok(None); }; + if !routing_overlay_allows_affinity_target(affinity_context.routing_overlay.as_ref(), &target) { + return Ok(None); + } let transport = match state .read_provider_transport_snapshot(&target.provider_id, &target.endpoint_id, &target.key_id) @@ -547,6 +578,16 @@ async fn maybe_forward_public_request_to_tunnel_owner( Ok(Some(response)) } +fn routing_overlay_allows_affinity_target( + routing_overlay: Option<&aether_routing_core::RankingOverlay>, + target: &aether_scheduler_core::SchedulerAffinityTarget, +) -> bool { + routing_overlay.is_none_or(|overlay| { + overlay.provider_allowed(target.provider_id.as_str()) + && overlay.key_allowed(target.key_id.as_str()) + }) +} + fn owner_forward_request_is_stream( parts: &http::request::Parts, decision: &GatewayControlDecision, @@ -2327,14 +2368,51 @@ mod tests { api_key_remote_ip_allowed, buffer_and_normalize_request_body, diagnostic_is_auth_api_key_concurrency_limited, local_execution_runtime_miss_detail, owner_forward_request_is_stream, restore_redacted_stream_execution_response, - restore_redacted_sync_execution_response, GatewayControlDecision, - LocalExecutionRuntimeMissDiagnostic, RequestBodyBufferError, RequestBodyBufferPolicy, + restore_redacted_sync_execution_response, routing_overlay_allows_affinity_target, + GatewayControlDecision, LocalExecutionRuntimeMissDiagnostic, RequestBodyBufferError, + RequestBodyBufferPolicy, }; use axum::body::{to_bytes, Body, Bytes}; use axum::http::{header, HeaderMap, HeaderValue, Method, Response}; use serde_json::json; use tokio::sync::Semaphore; + #[test] + fn routing_overlay_blocks_disallowed_tunnel_affinity_target() { + let target = aether_scheduler_core::SchedulerAffinityTarget { + provider_id: "provider-allowed".to_string(), + endpoint_id: "endpoint-1".to_string(), + key_id: "key-allowed".to_string(), + }; + let matching = aether_routing_core::RankingOverlay { + allowed_providers: vec!["provider-allowed".to_string()], + allowed_keys: vec!["key-allowed".to_string()], + ..aether_routing_core::RankingOverlay::default() + }; + let wrong_provider = aether_routing_core::RankingOverlay { + allowed_providers: vec!["provider-other".to_string()], + ..matching.clone() + }; + let wrong_key = aether_routing_core::RankingOverlay { + allowed_keys: vec!["key-other".to_string()], + ..matching.clone() + }; + + assert!(routing_overlay_allows_affinity_target(None, &target)); + assert!(routing_overlay_allows_affinity_target( + Some(&matching), + &target + )); + assert!(!routing_overlay_allows_affinity_target( + Some(&wrong_provider), + &target + )); + assert!(!routing_overlay_allows_affinity_target( + Some(&wrong_key), + &target + )); + } + #[test] fn owner_forward_uses_search_protocol_timeout_semantics() { let request = http::Request::builder() diff --git a/apps/aether-gateway/src/handlers/shared/request_utils.rs b/apps/aether-gateway/src/handlers/shared/request_utils.rs index 9cf01b71b..8267b1382 100644 --- a/apps/aether-gateway/src/handlers/shared/request_utils.rs +++ b/apps/aether-gateway/src/handlers/shared/request_utils.rs @@ -244,6 +244,12 @@ pub(crate) fn admin_proxy_local_requires_buffered_body( http::Method::POST, Some("import_refresh_token"), ) + | (Some("provider_oauth_manage"), http::Method::POST, Some("cookie_authorize")) + | ( + Some("provider_oauth_manage"), + http::Method::POST, + Some("start_cookie_authorize_task"), + ) | (Some("provider_oauth_manage"), http::Method::POST, Some("batch_import_oauth")) | ( Some("provider_oauth_manage"), diff --git a/apps/aether-gateway/src/oauth/http_executor.rs b/apps/aether-gateway/src/oauth/http_executor.rs index f4481d9ac..bfe43bf0d 100644 --- a/apps/aether-gateway/src/oauth/http_executor.rs +++ b/apps/aether-gateway/src/oauth/http_executor.rs @@ -69,7 +69,7 @@ impl<'a> OAuthHttpExecutor for GatewayOAuthHttpExecutor<'a> { provider_api_format: "oauth:exchange".to_string(), model_name: Some("oauth-exchange".to_string()), proxy: request.network.proxy, - transport_profile: None, + transport_profile: request.transport_profile, timeouts: Some(ExecutionTimeouts { connect_ms: Some(timeouts.connect_ms), read_ms: Some(timeouts.read_ms), diff --git a/apps/aether-gateway/src/orchestration/attempt.rs b/apps/aether-gateway/src/orchestration/attempt.rs index 987eb9d05..1ae62662f 100644 --- a/apps/aether-gateway/src/orchestration/attempt.rs +++ b/apps/aether-gateway/src/orchestration/attempt.rs @@ -35,6 +35,7 @@ pub(crate) struct LocalExecutionCandidateMetadata { } pub(crate) const SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD: &str = "scheduler_affinity_epoch"; +pub(crate) const ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD: &str = "routing_pool_policy_override"; pub(crate) const POOL_KEY_LEASE_KEY_REPORT_FIELD: &str = "pool_key_lease_key"; pub(crate) const POOL_KEY_LEASE_OWNER_REPORT_FIELD: &str = "pool_key_lease_owner"; pub(crate) const POOL_KEY_LEASE_TOKEN_REPORT_FIELD: &str = "pool_key_lease_token"; diff --git a/apps/aether-gateway/src/orchestration/effects.rs b/apps/aether-gateway/src/orchestration/effects.rs index ebe093c38..8e4f99d27 100644 --- a/apps/aether-gateway/src/orchestration/effects.rs +++ b/apps/aether-gateway/src/orchestration/effects.rs @@ -13,8 +13,9 @@ use aether_data_contracts::repository::provider_catalog::{ ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyHealthStateUpdate, }; +use aether_routing_core::RoutingPoolPolicyOverride; use aether_scheduler_core::{ - build_scheduler_affinity_cache_key_for_api_key_id_with_client_session, + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope, count_recent_rpm_requests_for_provider_key, ClientSessionAffinity, SchedulerAffinityTarget, }; use aether_usage_runtime::{ @@ -41,9 +42,16 @@ use crate::handlers::shared::provider_pool::{ admin_provider_pool_key_terminal_error_reason, record_admin_provider_pool_error, record_admin_provider_pool_stream_timeout, record_admin_provider_pool_success, release_admin_provider_pool_key_lease, AdminProviderPoolConfig, + AdminProviderPoolSchedulingPreset, +}; +use crate::orchestration::{ + local_execution_candidate_metadata_from_report_context, + ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD, +}; +use crate::scheduler::affinity::{ + scheduler_affinity_policy_context_from_report_context, SCHEDULER_AFFINITY_POLICY_REPORT_FIELD, + SCHEDULER_AFFINITY_TTL, }; -use crate::orchestration::local_execution_candidate_metadata_from_report_context; -use crate::scheduler::affinity::SCHEDULER_AFFINITY_TTL; use crate::scheduler::config::{read_scheduler_ordering_config, SchedulerSchedulingMode}; use crate::AppState; @@ -343,11 +351,22 @@ fn report_context_string_field<'a>( fn local_scheduler_affinity_cache_key(report_context: Option<&Value>) -> Option { let client_session_affinity = local_client_session_affinity(report_context); - build_scheduler_affinity_cache_key_for_api_key_id_with_client_session( + let policy_context = scheduler_affinity_policy_context_from_report_context(report_context); + if report_context + .and_then(|context| context.get(SCHEDULER_AFFINITY_POLICY_REPORT_FIELD)) + .is_some() + && policy_context.is_none() + { + return None; + } + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope( report_context_string_field(report_context, "api_key_id")?, report_context_string_field(report_context, "client_api_format")?, report_context_string_field(report_context, "model")?, client_session_affinity.as_ref(), + policy_context + .as_ref() + .and_then(|context| context.scope.as_ref()), ) } @@ -504,7 +523,17 @@ async fn local_scheduler_affinity_matches_failed_target( local_execution_plan_uses_pool(state, plan).await } -async fn scheduler_cache_affinity_enabled(state: &AppState) -> bool { +async fn scheduler_cache_affinity_enabled( + state: &AppState, + report_context: Option<&Value>, +) -> bool { + if report_context + .and_then(|context| context.get(SCHEDULER_AFFINITY_POLICY_REPORT_FIELD)) + .is_some() + { + return scheduler_affinity_policy_context_from_report_context(report_context) + .is_some_and(|context| context.cache_affinity_enabled()); + } match read_scheduler_ordering_config(state).await { Ok(config) => config.scheduling_mode == SchedulerSchedulingMode::CacheAffinity, Err(error) => { @@ -523,7 +552,7 @@ async fn remember_successful_local_scheduler_affinity( state: &AppState, context: LocalExecutionEffectContext<'_>, ) { - if !scheduler_cache_affinity_enabled(state).await { + if !scheduler_cache_affinity_enabled(state, context.report_context).await { return; } let Some(cache_key) = local_scheduler_affinity_cache_key(context.report_context) else { @@ -577,12 +606,33 @@ async fn resolve_pool_feedback_context( } }; - let Some(pool_config) = + let Some(mut pool_config) = admin_provider_pool_config_from_config_value(transport.provider.config.as_ref()) else { return None; }; + if let Some(override_policy) = context + .report_context + .and_then(|report_context| report_context.get(ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD)) + .and_then(|value| serde_json::from_value::(value.clone()).ok()) + .filter(|override_policy| !override_policy.scheduling_presets.is_empty()) + { + let scheduling_presets = override_policy + .scheduling_presets + .into_iter() + .map(|preset| AdminProviderPoolSchedulingPreset { + preset: preset.preset, + enabled: preset.enabled, + mode: preset.mode, + }) + .collect::>(); + pool_config.lru_enabled = scheduling_presets + .iter() + .any(|preset| preset.enabled && preset.preset.eq_ignore_ascii_case("lru")); + pool_config.scheduling_presets = scheduling_presets; + } + let sticky_session_token = pool_feedback_request_body(plan, context.report_context) .and_then(extract_pool_sticky_session_token); @@ -1758,10 +1808,11 @@ mod tests { apply_local_execution_effect, execution_plan_bearer_matches_transport, local_candidate_failure_should_apply_key_effects, local_candidate_failure_should_record_pool_error, pool_score_feedback_gate_allows, - pool_score_hard_state_for_status, LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, - LocalAttemptFailureEffect, LocalExecutionEffect, LocalExecutionEffectContext, - LocalHealthFailureEffect, LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, - LocalPoolErrorEffect, ProviderKeyEffectLockPool, + pool_score_hard_state_for_status, resolve_pool_feedback_context, + LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect, + LocalExecutionEffect, LocalExecutionEffectContext, LocalHealthFailureEffect, + LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, LocalPoolErrorEffect, + ProviderKeyEffectLockPool, }; use crate::data::{GatewayDataConfig, GatewayDataState}; use crate::orchestration::LocalFailoverClassification; @@ -1770,7 +1821,8 @@ mod tests { use aether_scheduler_core::{ build_scheduler_affinity_cache_key_for_api_key_id, build_scheduler_affinity_cache_key_for_api_key_id_with_client_session, - ClientSessionAffinity, SchedulerAffinityTarget, + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope, + ClientSessionAffinity, SchedulerAffinityScope, SchedulerAffinityTarget, }; async fn start_managed_redis_or_skip() -> Option { @@ -2221,6 +2273,61 @@ mod tests { ) } + #[tokio::test] + async fn pool_feedback_uses_routing_profile_scheduling_override() { + let mut provider = sample_pool_health_provider(); + provider.config = Some(json!({ + "pool_advanced": { + "scheduling_presets": [{ + "preset": "lru", + "enabled": true + }] + } + })); + let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![sample_health_endpoint()], + vec![sample_health_key()], + )); + let state = AppState::new() + .expect("gateway state should build") + .with_data_state_for_tests( + GatewayDataState::with_provider_catalog_repository_for_tests(repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ); + let plan = sample_plan(); + let report_context = json!({ + "routing_pool_policy_override": { + "scheduling_presets": [{ + "preset": "cache_affinity", + "enabled": true + }] + } + }); + + let feedback = resolve_pool_feedback_context( + &state, + LocalExecutionEffectContext { + plan: &plan, + report_context: Some(&report_context), + }, + ) + .await + .expect("pool feedback context should resolve"); + + assert!( + crate::handlers::shared::provider_pool::admin_provider_pool_cache_affinity_enabled( + &feedback.pool_config + ) + ); + assert!(!feedback.pool_config.lru_enabled); + assert_eq!(feedback.pool_config.scheduling_presets.len(), 1); + assert_eq!( + feedback.pool_config.scheduling_presets[0].preset, + "cache_affinity" + ); + } + fn health_state_with_key(key: StoredProviderCatalogKey) -> AppState { let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( vec![sample_health_provider()], @@ -2632,6 +2739,144 @@ mod tests { ); } + #[tokio::test] + async fn routing_profile_cache_affinity_overrides_legacy_fixed_mode_on_success() { + let state = AppState::new() + .expect("gateway state should build") + .with_data_state_for_tests( + GatewayDataState::disabled().with_system_config_values_for_tests(vec![( + "scheduling_mode".to_string(), + json!("fixed_order"), + )]), + ); + let plan = sample_plan(); + let affinity = session_affinity(); + let scope = SchedulerAffinityScope::new("routing-group-1", Some(7)); + let report_context = json!({ + "api_key_id": "api-key-1", + "client_api_format": "openai:chat", + "model": "gpt-5", + "client_session_affinity": { + "client_family": "generic", + "session_key": "session=session-1;agent=coder" + }, + "scheduler_affinity_policy": { + "scheduling_mode": "cache_affinity", + "scope": { + "routing_group_id": "routing-group-1", + "routing_group_version": 7 + } + } + }); + let scoped_cache_key = + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope( + "api-key-1", + "openai:chat", + "gpt-5", + Some(&affinity), + Some(&scope), + ) + .expect("scoped scheduler affinity cache key should build"); + + apply_local_execution_effect( + &state, + LocalExecutionEffectContext { + plan: &plan, + report_context: Some(&report_context), + }, + LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect), + ) + .await; + + assert_eq!( + state.read_scheduler_affinity_target(scoped_cache_key.as_str(), SCHEDULER_AFFINITY_TTL), + Some(SchedulerAffinityTarget { + provider_id: "prov-1".to_string(), + endpoint_id: "ep-1".to_string(), + key_id: "key-1".to_string(), + }) + ); + assert!(state + .read_scheduler_affinity_target( + session_scheduler_affinity_cache_key().as_str(), + SCHEDULER_AFFINITY_TTL + ) + .is_none()); + } + + #[tokio::test] + async fn routing_profile_fixed_mode_overrides_legacy_cache_affinity_on_success() { + let state = AppState::new().expect("gateway state should build"); + let plan = sample_plan(); + let report_context = json!({ + "api_key_id": "api-key-1", + "client_api_format": "openai:chat", + "model": "gpt-5", + "scheduler_affinity_policy": { + "scheduling_mode": "fixed_order", + "scope": { + "routing_group_id": "routing-group-1", + "routing_group_version": 7 + } + } + }); + let scope = SchedulerAffinityScope::new("routing-group-1", Some(7)); + let scoped_cache_key = + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope( + "api-key-1", + "openai:chat", + "gpt-5", + None, + Some(&scope), + ) + .expect("scoped scheduler affinity cache key should build"); + + apply_local_execution_effect( + &state, + LocalExecutionEffectContext { + plan: &plan, + report_context: Some(&report_context), + }, + LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect), + ) + .await; + + assert!(state + .read_scheduler_affinity_target(scoped_cache_key.as_str(), SCHEDULER_AFFINITY_TTL) + .is_none()); + } + + #[tokio::test] + async fn malformed_routing_affinity_context_does_not_fall_back_to_legacy_mode() { + let state = AppState::new().expect("gateway state should build"); + let plan = sample_plan(); + let report_context = json!({ + "api_key_id": "api-key-1", + "client_api_format": "openai:chat", + "model": "gpt-5", + "scheduler_affinity_policy": { + "scheduling_mode": "unknown" + } + }); + let legacy_cache_key = + build_scheduler_affinity_cache_key_for_api_key_id("api-key-1", "openai:chat", "gpt-5") + .expect("legacy scheduler affinity cache key should build"); + + apply_local_execution_effect( + &state, + LocalExecutionEffectContext { + plan: &plan, + report_context: Some(&report_context), + }, + LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect), + ) + .await; + + assert!(state + .read_scheduler_affinity_target(legacy_cache_key.as_str(), SCHEDULER_AFFINITY_TTL) + .is_none()); + } + #[tokio::test] async fn health_success_keeps_scheduler_affinity_after_health_state_update() { let state = health_state(); diff --git a/apps/aether-gateway/src/orchestration/mod.rs b/apps/aether-gateway/src/orchestration/mod.rs index b9782b319..369a2d9a5 100644 --- a/apps/aether-gateway/src/orchestration/mod.rs +++ b/apps/aether-gateway/src/orchestration/mod.rs @@ -22,7 +22,8 @@ pub(crate) use self::attempt::{ attempt_identity_from_report_context, build_local_attempt_identities, insert_pool_key_lease_report_context_fields, local_attempt_slot_count, local_execution_candidate_metadata_from_report_context, ExecutionAttemptIdentity, - LocalExecutionCandidateMetadata, SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD, + LocalExecutionCandidateMetadata, ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD, + SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD, }; pub(crate) use self::classifier::{ classify_anthropic_failure_disposition, classify_failure_disposition, classify_local_failover, diff --git a/apps/aether-gateway/src/routing/selection.rs b/apps/aether-gateway/src/routing/selection.rs index f4750e5e7..bef1ccfcd 100644 --- a/apps/aether-gateway/src/routing/selection.rs +++ b/apps/aether-gateway/src/routing/selection.rs @@ -14,6 +14,8 @@ pub(crate) enum GatewayRoutingSelectionError { Disabled(String), #[error("routing group was explicitly requested but is not allowed for this principal: {0}")] Forbidden(String), + #[error("routing group repository lookup failed: {0}")] + Repository(String), } #[derive(Debug, Clone, Default)] @@ -42,22 +44,21 @@ pub(crate) async fn select_gateway_routing_group( let group = repository .find_routing_group(RoutingGroupLookupKey::Id(explicit)) .await - .ok() - .flatten() - .or({ - let group: Option = repository - .find_routing_group(RoutingGroupLookupKey::Name(explicit)) - .await - .unwrap_or_default(); - group - }); + .map_err(repository_selection_error)?; + let group = match group { + Some(group) => Some(group), + None => repository + .find_routing_group(RoutingGroupLookupKey::Name(explicit)) + .await + .map_err(repository_selection_error)?, + }; let Some(group) = group else { return Err(GatewayRoutingSelectionError::NotFound(explicit.to_string())); }; if !group.enabled { return Err(GatewayRoutingSelectionError::Disabled(group.id)); } - if !explicit_group_allowed(repository, &group.id, &input).await { + if !explicit_group_allowed(repository, &group, &input).await? { return Err(GatewayRoutingSelectionError::Forbidden(group.id)); } return Ok(GatewayRoutingGroupSelection { @@ -70,13 +71,15 @@ pub(crate) async fn select_gateway_routing_group( // produce a group. The data-state repository answers this with a cached // existence query, so the common "routing configured but unused" case // does not materialize the binding table per API key/user. - let has_bindings = repository.has_any_routing_group_binding().await; - if matches!(has_bindings, Ok(false)) { + let has_bindings = repository + .has_any_routing_group_binding() + .await + .map_err(repository_selection_error)?; + if !has_bindings { let system_default = repository .find_routing_group(RoutingGroupLookupKey::SystemDefault) .await - .ok() - .flatten() + .map_err(repository_selection_error)? .filter(|group| group.enabled); return Ok(GatewayRoutingGroupSelection { group: system_default, @@ -92,13 +95,12 @@ pub(crate) async fn select_gateway_routing_group( subject_id: Some(subject_id.to_string()), }) .await - .unwrap_or_default(); + .map_err(repository_selection_error)?; for binding in bindings.into_iter().filter(|binding| binding.is_default) { let group = repository .find_routing_group(RoutingGroupLookupKey::Id(&binding.group_id)) .await - .ok() - .flatten(); + .map_err(repository_selection_error)?; if let Some(group) = group.filter(|group| group.enabled) { return Ok(GatewayRoutingGroupSelection { group: Some(group), @@ -111,8 +113,7 @@ pub(crate) async fn select_gateway_routing_group( let system_default = repository .find_routing_group(RoutingGroupLookupKey::SystemDefault) .await - .ok() - .flatten() + .map_err(repository_selection_error)? .filter(|group| group.enabled); Ok(GatewayRoutingGroupSelection { group: system_default, @@ -122,31 +123,30 @@ pub(crate) async fn select_gateway_routing_group( async fn explicit_group_allowed( repository: &(impl RoutingGroupReadRepository + ?Sized), - group_id: &str, + group: &StoredRoutingGroup, input: &GatewayRoutingSelectionInput<'_>, -) -> bool { - if let Ok(Some(group)) = repository - .find_routing_group(RoutingGroupLookupKey::Id(group_id)) - .await - { - if group.is_system_default { - return true; - } +) -> Result { + if group.is_system_default { + return Ok(true); } for (subject_type, subject_id, _) in default_binding_candidates(input) { let bindings = repository .list_routing_group_bindings(&RoutingGroupBindingQuery { - group_id: Some(group_id.to_string()), + group_id: Some(group.id.clone()), subject_type: Some(subject_type), subject_id: Some(subject_id.to_string()), }) .await - .unwrap_or_default(); + .map_err(repository_selection_error)?; if bindings.iter().any(|binding| binding.allow_explicit_select) { - return true; + return Ok(true); } } - false + Ok(false) +} + +fn repository_selection_error(error: impl std::fmt::Display) -> GatewayRoutingSelectionError { + GatewayRoutingSelectionError::Repository(error.to_string()) } fn default_binding_candidates<'a>( @@ -178,11 +178,57 @@ mod tests { use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository; use aether_data_contracts::repository::routing_profiles::{ CreateRoutingGroupBindingRecord, CreateRoutingGroupRecord, RoutingGroupWriteRepository, + StoredRoutingGroupBinding, StoredRoutingGroupVersion, }; + use aether_data_contracts::DataLayerError; + use async_trait::async_trait; use serde_json::json; use super::*; + struct FailingRoutingGroupRepository { + id_lookup_is_missing: bool, + } + + impl FailingRoutingGroupRepository { + fn failure() -> Result { + Err(DataLayerError::Sql( + "routing repository unavailable".to_string(), + )) + } + } + + #[async_trait] + impl RoutingGroupReadRepository for FailingRoutingGroupRepository { + async fn list_routing_groups(&self) -> Result, DataLayerError> { + Self::failure() + } + + async fn find_routing_group( + &self, + lookup: RoutingGroupLookupKey<'_>, + ) -> Result, DataLayerError> { + if self.id_lookup_is_missing && matches!(lookup, RoutingGroupLookupKey::Id(_)) { + return Ok(None); + } + Self::failure() + } + + async fn list_routing_group_bindings( + &self, + _query: &RoutingGroupBindingQuery, + ) -> Result, DataLayerError> { + Self::failure() + } + + async fn list_routing_group_versions( + &self, + _group_id: &str, + ) -> Result, DataLayerError> { + Self::failure() + } + } + #[tokio::test] async fn selects_api_key_default_binding() { let repository = InMemoryRoutingGroupRepository::default(); @@ -336,6 +382,58 @@ mod tests { ); } + #[tokio::test] + async fn propagates_explicit_name_lookup_failure_after_missing_id() { + let repository = FailingRoutingGroupRepository { + id_lookup_is_missing: true, + }; + + let error = select_gateway_routing_group( + &repository, + GatewayRoutingSelectionInput { + explicit_group: Some("group-name"), + user_id: Some("user-1"), + api_key_id: Some("api-key-1"), + user_group_ids: &[], + }, + ) + .await + .unwrap_err(); + + assert_eq!( + error, + GatewayRoutingSelectionError::Repository( + "sql error: routing repository unavailable".to_string() + ) + ); + } + + #[tokio::test] + async fn propagates_implicit_binding_lookup_failure() { + let repository = FailingRoutingGroupRepository { + id_lookup_is_missing: false, + }; + + let error = select_gateway_routing_group( + &repository, + GatewayRoutingSelectionInput { + explicit_group: None, + user_id: Some("user-1"), + api_key_id: Some("api-key-1"), + user_group_ids: &[], + }, + ) + .await + .unwrap_err(); + + assert_eq!( + error, + GatewayRoutingSelectionError::Repository( + "sql error: routing repository unavailable".to_string() + ) + ); + } + #[tokio::test] async fn rejects_explicit_disabled_group() { let repository = InMemoryRoutingGroupRepository::default(); diff --git a/apps/aether-gateway/src/scheduler/affinity.rs b/apps/aether-gateway/src/scheduler/affinity.rs index 1c3fd7525..401b1b7f7 100644 --- a/apps/aether-gateway/src/scheduler/affinity.rs +++ b/apps/aether-gateway/src/scheduler/affinity.rs @@ -1,13 +1,65 @@ use std::time::Duration; +use aether_routing_core::{ResolvedRoutingPolicy, RoutingSchedulingMode}; use aether_scheduler_core::{ - build_scheduler_affinity_cache_key_for_api_key_id_with_client_session, ClientSessionAffinity, - SchedulerAffinityTarget, + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session, + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope, + ClientSessionAffinity, SchedulerAffinityScope, SchedulerAffinityTarget, }; +use serde::{Deserialize, Serialize}; +use serde_json::Value; use super::state::SchedulerRuntimeState; pub(crate) const SCHEDULER_AFFINITY_TTL: Duration = Duration::from_secs(300); +pub(crate) const SCHEDULER_AFFINITY_POLICY_REPORT_FIELD: &str = "scheduler_affinity_policy"; + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub(crate) struct SchedulerAffinityPolicyContext { + pub(crate) scheduling_mode: RoutingSchedulingMode, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub(crate) scope: Option, +} + +impl SchedulerAffinityPolicyContext { + pub(crate) fn from_routing_policy(policy: &ResolvedRoutingPolicy) -> Self { + let scope = policy + .group_id + .as_deref() + .map(str::trim) + .filter(|group_id| !group_id.is_empty()) + .map(|group_id| SchedulerAffinityScope::new(group_id, policy.group_version)); + Self { + scheduling_mode: policy.scheduling_mode, + scope, + } + } + + pub(crate) fn cache_affinity_enabled(&self) -> bool { + self.scheduling_mode == RoutingSchedulingMode::CacheAffinity + } +} + +pub(crate) fn scheduler_affinity_policy_context_from_report_context( + report_context: Option<&Value>, +) -> Option { + report_context + .and_then(|context| context.get(SCHEDULER_AFFINITY_POLICY_REPORT_FIELD)) + .and_then(|value| serde_json::from_value(value.clone()).ok()) +} + +pub(crate) fn insert_scheduler_affinity_policy_report_context_field( + extra_fields: &mut serde_json::Map, + routing_policy: Option<&ResolvedRoutingPolicy>, +) { + let Some(routing_policy) = routing_policy else { + return; + }; + let context = SchedulerAffinityPolicyContext::from_routing_policy(routing_policy); + if let Ok(value) = serde_json::to_value(context) { + extra_fields.insert(SCHEDULER_AFFINITY_POLICY_REPORT_FIELD.to_string(), value); + } +} pub(crate) fn read_cached_scheduler_affinity_target( state: &(impl SchedulerRuntimeState + ?Sized), @@ -24,3 +76,25 @@ pub(crate) fn read_cached_scheduler_affinity_target( )?; state.read_cached_scheduler_affinity_target(&cache_key, SCHEDULER_AFFINITY_TTL) } + +pub(crate) fn read_cached_scheduler_affinity_target_with_policy_context( + state: &(impl SchedulerRuntimeState + ?Sized), + api_key_id: &str, + client_session_affinity: Option<&ClientSessionAffinity>, + api_format: &str, + global_model_name: &str, + policy_context: &SchedulerAffinityPolicyContext, +) -> Option { + if !policy_context.cache_affinity_enabled() { + return None; + } + let cache_key = + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope( + api_key_id, + api_format, + global_model_name, + client_session_affinity, + policy_context.scope.as_ref(), + )?; + state.read_cached_scheduler_affinity_target(&cache_key, SCHEDULER_AFFINITY_TTL) +} diff --git a/apps/aether-gateway/src/tests/control/admin/oauth.rs b/apps/aether-gateway/src/tests/control/admin/oauth.rs index dc1457d55..dd2164256 100644 --- a/apps/aether-gateway/src/tests/control/admin/oauth.rs +++ b/apps/aether-gateway/src/tests/control/admin/oauth.rs @@ -6,6 +6,7 @@ use aether_contracts::{ use aether_crypto::{ decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY, }; +use aether_data::repository::background_tasks::InMemoryBackgroundTaskRepository; use aether_data::repository::management_tokens::{ InMemoryManagementTokenRepository, ManagementTokenReadRepository, }; @@ -15,6 +16,7 @@ use aether_data::repository::oauth_providers::{ use aether_data::repository::pool_scores::InMemoryPoolMemberScoreRepository; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::repository::proxy_nodes::InMemoryProxyNodeRepository; +use aether_data_contracts::repository::background_tasks::BackgroundTaskReadRepository; use aether_data_contracts::repository::pool_scores::{ GetPoolMemberScoresByIdsQuery, PoolMemberHardState, PoolMemberIdentity, PoolScoreReadRepository, }; @@ -26,7 +28,7 @@ use axum::response::{IntoResponse, Response}; use axum::routing::{any, delete, get, patch, post, put}; use axum::{extract::Request, Json, Router}; use http::{HeaderMap, HeaderValue, StatusCode}; -use serde_json::json; +use serde_json::{json, Value}; use super::super::{ build_router_with_state, build_state_with_execution_runtime_override, hash_management_token, @@ -309,17 +311,489 @@ async fn gateway_handles_admin_provider_oauth_supported_types_locally_with_trust let items = payload.as_array().expect("items should be array"); assert_eq!(items.len(), 6); assert_eq!(items[0]["provider_type"], "claude_code"); + assert_eq!( + items[0]["authorize_url"], + "https://claude.ai/oauth/authorize" + ); + assert_eq!( + items[0]["token_url"], + "https://platform.claude.com/v1/oauth/token" + ); + assert_eq!( + items[0]["redirect_uri"], + "https://platform.claude.com/oauth/code/callback" + ); + assert_eq!(items[0]["supports_cookie_authorization"], true); assert_eq!(items[1]["provider_type"], "codex"); assert_eq!(items[2]["provider_type"], "chatgpt_web"); assert_eq!(items[3]["provider_type"], "gemini_cli"); assert_eq!(items[4]["provider_type"], "antigravity"); assert_eq!(items[5]["provider_type"], "windsurf"); + assert!(items[1..] + .iter() + .all(|item| item["supports_cookie_authorization"] == false)); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); upstream_handle.abort(); } +#[test] +fn gateway_authorizes_claude_cookie_without_persisting_cookie() { + run_admin_oauth_test( + "gateway_authorizes_claude_cookie_without_persisting_cookie", + gateway_authorizes_claude_cookie_without_persisting_cookie_impl, + ); +} + +async fn gateway_authorizes_claude_cookie_without_persisting_cookie_impl() { + let execution_plans = Arc::new(Mutex::new(Vec::::new())); + let execution_plans_clone = Arc::clone(&execution_plans); + let execution_runtime = Router::new().route( + "/v1/execute/sync", + any(move |Json(plan): Json| { + let execution_plans_inner = Arc::clone(&execution_plans_clone); + async move { + execution_plans_inner + .lock() + .expect("mutex should lock") + .push(plan.clone()); + let json_body = match plan.request_id.as_str() { + "provider-oauth:claude-cookie-organizations" => json!([ + {"uuid": "org-personal", "raven_type": "personal"}, + {"uuid": "org-team", "raven_type": "team"} + ]), + "provider-oauth:claude-cookie-authorize" => { + let state = plan + .body + .json_body + .as_ref() + .and_then(|body| body.get("state")) + .and_then(serde_json::Value::as_str) + .expect("authorize plan should contain state"); + let mut redirect = + url::Url::parse("https://platform.claude.com/oauth/code/callback") + .expect("redirect URL should parse"); + redirect + .query_pairs_mut() + .append_pair("code", "claude-authorization-code") + .append_pair("state", state); + json!({"redirect_uri": redirect.to_string()}) + } + "provider-oauth:exchange-code" => json!({ + "access_token": "sk-ant-oat01-created", + "refresh_token": "sk-ant-ort01-created", + "expires_in": 3600, + "organization": {"uuid": "org-team"}, + "account": { + "uuid": "account-claude-123", + "email_address": "claude@example.com" + } + }), + unexpected => panic!("unexpected execution plan: {unexpected}"), + }; + Json(json!({ + "request_id": plan.request_id, + "status_code": 200, + "headers": {"content-type": "application/json"}, + "body": {"json_body": json_body} + })) + } + }), + ); + + let mut provider = sample_provider("provider-claude", "claude", 10); + provider.provider_type = "claude_code".to_string(); + provider.proxy = Some(json!({ + "mode": "tunnel", + "node_id": "proxy-node-claude", + "url": "http://proxy.example:8080", + "enabled": true + })); + let endpoint = sample_endpoint( + "endpoint-claude-messages", + "provider-claude", + "claude:messages", + "https://api.anthropic.com", + ); + let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![endpoint], + vec![], + )); + + let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; + let state = build_state_with_execution_runtime_override(execution_runtime_url) + .with_data_state_for_tests( + GatewayDataState::with_provider_catalog_repository_for_tests( + provider_catalog_repository.clone(), + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY), + ) + .with_provider_oauth_token_url_for_tests( + "claude_code_cookie_base_url", + "https://claude.example", + ) + .with_provider_oauth_token_url_for_tests( + "claude_code", + "https://platform.example/v1/oauth/token", + ); + + let response = local_admin_provider_oauth_response( + &state, + http::Method::POST, + "/api/admin/provider-oauth/providers/provider-claude/cookie-authorize", + Some(json!({ + "cookie": "Cookie: other=value; sessionKey=sk-ant-sid01-secret; theme=dark", + "name": "claude-cookie-account" + })), + ) + .await; + let status = response.status(); + let payload: serde_json::Value = serde_json::from_slice( + &to_bytes(response.into_body(), usize::MAX) + .await + .expect("body should read"), + ) + .expect("response should be JSON"); + assert_eq!(status, StatusCode::OK, "payload={payload}"); + assert_eq!(payload["provider_type"], "claude_code"); + assert_eq!(payload["has_refresh_token"], true); + assert_eq!(payload["temporary"], false); + assert_eq!(payload["email"], "claude@example.com"); + + let keys = provider_catalog_repository + .list_keys_by_provider_ids(&["provider-claude".to_string()]) + .await + .expect("keys should load"); + assert_eq!(keys.len(), 1); + let decrypted_api_key = decrypt_python_fernet_ciphertext( + DEVELOPMENT_ENCRYPTION_KEY, + keys[0] + .encrypted_api_key + .as_deref() + .expect("api key should be encrypted"), + ) + .expect("api key should decrypt"); + assert_eq!(decrypted_api_key, "sk-ant-oat01-created"); + let decrypted_auth_config = decrypt_python_fernet_ciphertext( + DEVELOPMENT_ENCRYPTION_KEY, + keys[0] + .encrypted_auth_config + .as_deref() + .expect("auth config should be encrypted"), + ) + .expect("auth config should decrypt"); + let auth_config: serde_json::Value = + serde_json::from_str(&decrypted_auth_config).expect("auth config should parse"); + assert_eq!(auth_config["refresh_token"], "sk-ant-ort01-created"); + assert_eq!(auth_config["org_uuid"], "org-team"); + assert_eq!(auth_config["account_uuid"], "account-claude-123"); + assert_eq!(auth_config["email"], "claude@example.com"); + assert!(!decrypted_auth_config.contains("sk-ant-sid01-secret")); + assert!(!decrypted_auth_config.contains("sessionKey")); + assert!(auth_config.get("cookie").is_none()); + + let plans = execution_plans.lock().expect("mutex should lock"); + assert_eq!(plans.len(), 3); + for plan in &plans[..2] { + assert_eq!( + plan.headers.get("cookie").map(String::as_str), + Some("sessionKey=sk-ant-sid01-secret") + ); + assert_eq!( + plan.headers + .get(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER) + .map(String::as_str), + Some("false") + ); + assert_eq!( + plan.proxy.as_ref().and_then(|proxy| proxy.mode.as_deref()), + Some("tunnel") + ); + assert!(plan.transport_profile.is_none()); + } + let token_plan = &plans[2]; + assert_eq!(token_plan.request_id, "provider-oauth:exchange-code"); + assert!(!token_plan.headers.contains_key("cookie")); + assert!(token_plan + .body + .json_body + .as_ref() + .is_some_and(|body| body.get("scope").is_none())); + + execution_runtime_handle.abort(); +} + +#[test] +fn gateway_batch_authorizes_claude_cookies_as_redacted_task() { + run_admin_oauth_test( + "gateway_batch_authorizes_claude_cookies_as_redacted_task", + gateway_batch_authorizes_claude_cookies_as_redacted_task_impl, + ); +} + +async fn gateway_batch_authorizes_claude_cookies_as_redacted_task_impl() { + let execution_plans = Arc::new(Mutex::new(Vec::::new())); + let execution_plans_clone = Arc::clone(&execution_plans); + let execution_runtime = Router::new().route( + "/v1/execute/sync", + any(move |Json(plan): Json| { + let execution_plans_inner = Arc::clone(&execution_plans_clone); + async move { + execution_plans_inner + .lock() + .expect("mutex should lock") + .push(plan.clone()); + let json_body = match plan.request_id.as_str() { + "provider-oauth:claude-cookie-organizations" => { + json!([{"uuid": "org-team", "raven_type": "team"}]) + } + "provider-oauth:claude-cookie-authorize" => { + let cookie = plan + .headers + .get("cookie") + .map(String::as_str) + .expect("authorize plan should contain cookie"); + let label = if cookie.contains("batch-sid-one") { + "one" + } else if cookie.contains("batch-sid-two") { + "two" + } else { + panic!("unexpected Cookie authorization input") + }; + let state = plan + .body + .json_body + .as_ref() + .and_then(|body| body.get("state")) + .and_then(serde_json::Value::as_str) + .expect("authorize plan should contain state"); + let mut redirect = + url::Url::parse("https://platform.claude.com/oauth/code/callback") + .expect("redirect URL should parse"); + redirect + .query_pairs_mut() + .append_pair("code", format!("code-{label}").as_str()) + .append_pair("state", state); + json!({"redirect_uri": redirect.to_string()}) + } + "provider-oauth:exchange-code" => { + let code = plan + .body + .json_body + .as_ref() + .and_then(|body| body.get("code")) + .and_then(serde_json::Value::as_str) + .expect("token plan should contain code"); + let label = code.strip_prefix("code-").expect("code should be tagged"); + json!({ + "access_token": format!("sk-ant-oat01-{label}"), + "refresh_token": format!("sk-ant-ort01-{label}"), + "expires_in": 3600, + "organization": {"uuid": "org-team"}, + "account": { + "uuid": format!("account-{label}"), + "email_address": format!("{label}@example.com") + } + }) + } + unexpected => panic!("unexpected execution plan: {unexpected}"), + }; + Json(json!({ + "request_id": plan.request_id, + "status_code": 200, + "headers": {"content-type": "application/json"}, + "body": {"json_body": json_body} + })) + } + }), + ); + + let mut provider = sample_provider("provider-claude", "claude", 10); + provider.provider_type = "claude_code".to_string(); + provider.proxy = Some(json!({ + "mode": "tunnel", + "node_id": "proxy-node-claude", + "url": "http://proxy.example:8080", + "enabled": true + })); + let endpoint = sample_endpoint( + "endpoint-claude-messages", + "provider-claude", + "claude:messages", + "https://api.anthropic.com", + ); + let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( + vec![provider], + vec![endpoint], + vec![], + )); + let background_task_repository = Arc::new(InMemoryBackgroundTaskRepository::default()); + let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; + let state = build_state_with_execution_runtime_override(execution_runtime_url) + .with_data_state_for_tests( + GatewayDataState::with_provider_catalog_repository_for_tests( + provider_catalog_repository.clone(), + ) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) + .with_background_task_repository_for_tests(background_task_repository.clone()), + ) + .with_provider_oauth_token_url_for_tests( + "claude_code_cookie_base_url", + "https://claude.example", + ) + .with_provider_oauth_token_url_for_tests( + "claude_code", + "https://platform.example/v1/oauth/token", + ); + + let response = local_admin_provider_oauth_response( + &state, + http::Method::POST, + "/api/admin/provider-oauth/providers/provider-claude/cookie-authorize/tasks", + Some(json!({ + "cookies": [ + "sessionKey=batch-sid-one", + "foo=bar", + "Cookie: sessionKey=batch-sid-one", + "sessionKey=batch-sid-two" + ] + })), + ) + .await; + assert_eq!(response.status(), StatusCode::OK); + let submitted: serde_json::Value = serde_json::from_slice( + &to_bytes(response.into_body(), usize::MAX) + .await + .expect("submitted body should read"), + ) + .expect("submitted body should parse"); + assert_eq!(submitted["status"], "submitted"); + assert_eq!(submitted["total"], 4); + assert_eq!(submitted["import_kind"], "cookie_authorize"); + let task_id = submitted["task_id"] + .as_str() + .expect("task id should exist") + .to_string(); + assert!(task_id.starts_with("claude-cookie-")); + + let mut status_payload = Value::Null; + for _ in 0..80 { + let response = local_admin_provider_oauth_response( + &state, + http::Method::GET, + format!( + "/api/admin/provider-oauth/providers/provider-claude/cookie-authorize/tasks/{task_id}" + ) + .as_str(), + None, + ) + .await; + assert_eq!(response.status(), StatusCode::OK); + status_payload = serde_json::from_slice( + &to_bytes(response.into_body(), usize::MAX) + .await + .expect("status body should read"), + ) + .expect("status body should parse"); + if status_payload["status"] == "completed" { + break; + } + tokio::time::sleep(std::time::Duration::from_millis(25)).await; + } + + assert_eq!(status_payload["status"], "completed"); + assert_eq!(status_payload["import_kind"], "cookie_authorize"); + assert_eq!(status_payload["total"], 4); + assert_eq!(status_payload["processed"], 4); + assert_eq!(status_payload["success"], 2); + assert_eq!(status_payload["failed"], 2); + assert_eq!(status_payload["created_count"], 2); + assert_eq!(status_payload["replaced_count"], 0); + let error_samples = status_payload["error_samples"] + .as_array() + .expect("error samples should be an array"); + assert_eq!(error_samples.len(), 2); + assert_eq!(error_samples[0]["index"], 1); + assert_eq!(error_samples[1]["index"], 2); + + let status_text = status_payload.to_string(); + for forbidden in ["sessionKey", "batch-sid-one", "batch-sid-two"] { + assert!(!status_text.contains(forbidden), "forbidden={forbidden}"); + } + let raw_task_state = state + .load_provider_oauth_batch_task_for_tests( + format!("provider_oauth_batch_task:{task_id}").as_str(), + ) + .expect("raw task state should exist"); + for forbidden in ["sessionKey", "batch-sid-one", "batch-sid-two"] { + assert!( + !raw_task_state.contains(forbidden), + "raw task state contains {forbidden}" + ); + } + let background_run = background_task_repository + .find_run(&task_id) + .await + .expect("background run should load") + .expect("background run should exist"); + let background_text = serde_json::to_string(&background_run) + .expect("background run should serialize for assertion"); + for forbidden in ["sessionKey", "batch-sid-one", "batch-sid-two"] { + assert!( + !background_text.contains(forbidden), + "forbidden={forbidden}" + ); + } + + let keys = provider_catalog_repository + .list_keys_by_provider_ids(&["provider-claude".to_string()]) + .await + .expect("keys should load"); + assert_eq!(keys.len(), 2); + for key in &keys { + let auth_config = decrypt_python_fernet_ciphertext( + DEVELOPMENT_ENCRYPTION_KEY, + key.encrypted_auth_config + .as_deref() + .expect("auth config should be encrypted"), + ) + .expect("auth config should decrypt"); + assert!(!auth_config.contains("batch-sid")); + assert!(!auth_config.contains("sessionKey")); + } + + let plans = execution_plans.lock().expect("mutex should lock"); + assert_eq!(plans.len(), 6); + assert_eq!( + plans + .iter() + .filter(|plan| plan.request_id == "provider-oauth:claude-cookie-organizations") + .count(), + 2 + ); + assert_eq!( + plans + .iter() + .filter(|plan| plan.request_id == "provider-oauth:claude-cookie-authorize") + .count(), + 2 + ); + assert_eq!( + plans + .iter() + .filter(|plan| plan.request_id == "provider-oauth:exchange-code") + .count(), + 2 + ); + assert!(plans.iter().all(|plan| { + plan.proxy.as_ref().and_then(|proxy| proxy.mode.as_deref()) == Some("tunnel") + })); + + execution_runtime_handle.abort(); +} + #[test] fn gateway_handles_admin_provider_oauth_device_authorize_for_windsurf_browser() { run_admin_oauth_test( @@ -7955,6 +8429,7 @@ async fn gateway_handles_admin_provider_oauth_unavailable_routes_locally_with_tr let client = reqwest::Client::new(); for path in [ "/api/admin/provider-oauth/providers/provider-123/import-refresh-token", + "/api/admin/provider-oauth/providers/provider-123/cookie-authorize/tasks", "/api/admin/provider-oauth/providers/provider-123/batch-import", "/api/admin/provider-oauth/providers/provider-123/batch-import/tasks", "/api/admin/provider-oauth/providers/provider-123/device-authorize", diff --git a/apps/aether-gateway/src/tests/control/admin/pool.rs b/apps/aether-gateway/src/tests/control/admin/pool.rs index 3713b81d9..a2c98945a 100644 --- a/apps/aether-gateway/src/tests/control/admin/pool.rs +++ b/apps/aether-gateway/src/tests/control/admin/pool.rs @@ -279,11 +279,24 @@ async fn gateway_handles_admin_pool_scheduling_presets_locally_with_trusted_admi assert_eq!(response.status(), StatusCode::OK); let payload: serde_json::Value = response.json().await.expect("json body should parse"); let items = payload.as_array().expect("payload should be an array"); - assert_eq!(items.len(), 14); - assert_eq!(items[0]["name"], "lru"); - assert_eq!(items[1]["name"], "cache_affinity"); - assert_eq!(items[8]["name"], "pro_first"); - assert_eq!(items[13]["name"], "team_first"); + assert_eq!(items.len(), 15); + let preset_names = items + .iter() + .filter_map(|item| item["name"].as_str()) + .collect::>(); + assert_eq!(preset_names.first(), Some(&"lru")); + assert_eq!(preset_names.get(1), Some(&"cache_affinity")); + assert_eq!(preset_names.last(), Some(&"team_first")); + + let preset_index = |name| { + preset_names + .iter() + .position(|preset_name| *preset_name == name) + .unwrap_or_else(|| panic!("missing scheduling preset {name}")) + }; + assert!(preset_index("cost_first") < preset_index("free_team_first")); + assert!(preset_index("free_team_first") < preset_index("free_first")); + assert!(preset_index("pro_first") < preset_index("team_first")); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); gateway_handle.abort(); diff --git a/crates/aether-admin/src/provider/state.rs b/crates/aether-admin/src/provider/state.rs index 381ccf9cb..e3450287e 100644 --- a/crates/aether-admin/src/provider/state.rs +++ b/crates/aether-admin/src/provider/state.rs @@ -38,6 +38,17 @@ pub fn provider_oauth_pkce_s256(verifier: &str) -> String { pub fn parse_provider_oauth_callback_params(callback_url: &str) -> BTreeMap { let mut merged = BTreeMap::new(); let raw_callback_url = callback_url.trim(); + if !raw_callback_url.contains("://") { + if let Some((code, state)) = raw_callback_url.split_once('#') { + let code = code.trim(); + let state = state.strip_prefix("state=").unwrap_or(state).trim(); + if !code.is_empty() && !code.contains('=') && !state.is_empty() { + merged.insert("code".to_string(), code.to_string()); + merged.insert("state".to_string(), state.to_string()); + return merged; + } + } + } let parsed_url = Url::parse(raw_callback_url).or_else(|_| { Url::parse(&format!( "https://aether.local/{}", @@ -262,6 +273,36 @@ pub fn enrich_admin_provider_oauth_auth_config( ], ); + if provider_type.trim().eq_ignore_ascii_case("claude_code") { + if let Some(organization_uuid) = token_payload_object + .get("organization") + .and_then(Value::as_object) + .and_then(|value| value.get("uuid")) + .cloned() + { + auth_config + .entry("org_uuid".to_string()) + .or_insert(organization_uuid); + } + if let Some(account) = token_payload_object + .get("account") + .and_then(Value::as_object) + { + if let Some(account_uuid) = account.get("uuid").cloned() { + auth_config + .entry("account_uuid".to_string()) + .or_insert(account_uuid); + } + if let Some(email) = account.get("email_address").cloned() { + auth_config + .entry("email_address".to_string()) + .or_insert_with(|| email.clone()); + auth_config.entry("email".to_string()).or_insert(email); + } + } + return; + } + if !provider_type_uses_openai_chatgpt_identity(provider_type) { return; } @@ -407,6 +448,21 @@ mod tests { assert_eq!(params.get("state").map(String::as_str), Some("nonce-value")); } + #[test] + fn parse_provider_oauth_callback_params_reads_raw_claude_code_and_state() { + for (input, expected_state) in [ + ("claude-code#nonce-value", "nonce-value"), + ("claude-code#state=nonce-value", "nonce-value"), + ] { + let params = parse_provider_oauth_callback_params(input); + assert_eq!(params.get("code").map(String::as_str), Some("claude-code")); + assert_eq!( + params.get("state").map(String::as_str), + Some(expected_state) + ); + } + } + #[test] fn parse_provider_oauth_callback_params_reads_relative_show_auth_token_url() { let params = parse_provider_oauth_callback_params( diff --git a/crates/aether-data/adapters/mysql/migrations/20260727000000_repair_routing_default_uniqueness.sql b/crates/aether-data/adapters/mysql/migrations/20260727000000_repair_routing_default_uniqueness.sql new file mode 100644 index 000000000..7bc76ba5f --- /dev/null +++ b/crates/aether-data/adapters/mysql/migrations/20260727000000_repair_routing_default_uniqueness.sql @@ -0,0 +1,115 @@ +UPDATE routing_groups +SET is_system_default = 0 +WHERE id IN ( + SELECT id + FROM ( + SELECT + id, + ROW_NUMBER() OVER (ORDER BY enabled DESC, updated_at DESC, id ASC) AS default_rank + FROM routing_groups + WHERE is_system_default = 1 + ) AS ranked_defaults + WHERE default_rank > 1 +); + +UPDATE routing_group_bindings +SET is_default = 0 +WHERE id IN ( + SELECT id + FROM ( + SELECT + id, + ROW_NUMBER() OVER ( + PARTITION BY subject_type, subject_id + ORDER BY created_at ASC, id ASC + ) AS default_rank + FROM routing_group_bindings + WHERE is_default = 1 + ) AS ranked_defaults + WHERE default_rank > 1 +); + +SET @aether_routing_groups_default_guard_column_sql := IF( + ( + SELECT COUNT(*) + FROM information_schema.columns + WHERE table_schema = DATABASE() + AND table_name = 'routing_groups' + AND column_name = 'system_default_unique_guard' + ) = 0, + 'ALTER TABLE routing_groups ADD COLUMN system_default_unique_guard TINYINT GENERATED ALWAYS AS (CASE WHEN is_system_default = 1 THEN 1 ELSE NULL END) VIRTUAL INVISIBLE', + 'DO 0' +); + +PREPARE aether_routing_groups_default_guard_column_stmt + FROM @aether_routing_groups_default_guard_column_sql; +EXECUTE aether_routing_groups_default_guard_column_stmt; +DEALLOCATE PREPARE aether_routing_groups_default_guard_column_stmt; + +SET @aether_routing_bindings_default_type_guard_column_sql := IF( + ( + SELECT COUNT(*) + FROM information_schema.columns + WHERE table_schema = DATABASE() + AND table_name = 'routing_group_bindings' + AND column_name = 'default_subject_type_guard' + ) = 0, + 'ALTER TABLE routing_group_bindings ADD COLUMN default_subject_type_guard VARCHAR(32) GENERATED ALWAYS AS (CASE WHEN is_default = 1 THEN subject_type ELSE NULL END) VIRTUAL INVISIBLE', + 'DO 0' +); + +PREPARE aether_routing_bindings_default_type_guard_column_stmt + FROM @aether_routing_bindings_default_type_guard_column_sql; +EXECUTE aether_routing_bindings_default_type_guard_column_stmt; +DEALLOCATE PREPARE aether_routing_bindings_default_type_guard_column_stmt; + +SET @aether_routing_bindings_default_id_guard_column_sql := IF( + ( + SELECT COUNT(*) + FROM information_schema.columns + WHERE table_schema = DATABASE() + AND table_name = 'routing_group_bindings' + AND column_name = 'default_subject_id_guard' + ) = 0, + 'ALTER TABLE routing_group_bindings ADD COLUMN default_subject_id_guard VARCHAR(64) GENERATED ALWAYS AS (CASE WHEN is_default = 1 THEN subject_id ELSE NULL END) VIRTUAL INVISIBLE', + 'DO 0' +); + +PREPARE aether_routing_bindings_default_id_guard_column_stmt + FROM @aether_routing_bindings_default_id_guard_column_sql; +EXECUTE aether_routing_bindings_default_id_guard_column_stmt; +DEALLOCATE PREPARE aether_routing_bindings_default_id_guard_column_stmt; + +SET @aether_routing_groups_default_unique_index_sql := IF( + ( + SELECT COUNT(*) + FROM information_schema.statistics + WHERE table_schema = DATABASE() + AND table_name = 'routing_groups' + AND index_name = 'routing_groups_one_system_default_key' + ) = 0, + 'CREATE UNIQUE INDEX routing_groups_one_system_default_key ON routing_groups (system_default_unique_guard)', + 'DO 0' +); + +PREPARE aether_routing_groups_default_unique_index_stmt + FROM @aether_routing_groups_default_unique_index_sql; +EXECUTE aether_routing_groups_default_unique_index_stmt; +DEALLOCATE PREPARE aether_routing_groups_default_unique_index_stmt; + +SET @aether_routing_bindings_default_unique_index_sql := IF( + ( + SELECT COUNT(*) + FROM information_schema.statistics + WHERE table_schema = DATABASE() + AND table_name = 'routing_group_bindings' + AND index_name = 'routing_group_bindings_subject_default_key' + ) = 0, + 'CREATE UNIQUE INDEX routing_group_bindings_subject_default_key ON routing_group_bindings (default_subject_type_guard, default_subject_id_guard)', + 'DO 0' +); + +PREPARE aether_routing_bindings_default_unique_index_stmt + FROM @aether_routing_bindings_default_unique_index_sql; +EXECUTE aether_routing_bindings_default_unique_index_stmt; +DEALLOCATE PREPARE aether_routing_bindings_default_unique_index_stmt; diff --git a/crates/aether-data/adapters/mysql/src/candidate_selection.rs b/crates/aether-data/adapters/mysql/src/candidate_selection.rs index d3d6879f4..1967af316 100644 --- a/crates/aether-data/adapters/mysql/src/candidate_selection.rs +++ b/crates/aether-data/adapters/mysql/src/candidate_selection.rs @@ -4,10 +4,10 @@ use async_trait::async_trait; use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row}; use aether_data_contracts::repository::candidate_selection::{ - MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow, - StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery, - StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping, - StoredRequestedModelCandidateRowsQuery, + MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery, + StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder, + StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery, + StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery, }; use aether_data_contracts::DataLayerError; @@ -37,7 +37,6 @@ SELECT pak.capabilities AS key_capabilities, pak.internal_priority AS key_internal_priority, pak.global_priority_by_format AS key_global_priority_by_format, - pak.last_used_at AS key_last_used_at_unix_secs, m.id AS model_id, m.global_model_id AS global_model_id, gm.name AS global_model_name, @@ -60,6 +59,9 @@ WHERE p.is_active = 1 AND gm.is_active = 1 "#; +const REQUESTED_MODEL_RAW_PAGE_SIZE: u32 = 256; +const REQUESTED_MODEL_RAW_SCAN_LIMIT: u32 = 2048; + #[derive(Debug, Clone)] pub struct MysqlMinimalCandidateSelectionReadRepository { pool: MysqlPool, @@ -70,7 +72,52 @@ struct CandidateSelectionRow { row: StoredMinimalCandidateSelectionRow, provider_pool_enabled: bool, key_auth_config: Option, - key_last_used_at_unix_secs: Option, +} + +#[derive(Debug)] +struct ExactPageAccumulator { + rows: Vec, + offset: usize, + limit: usize, + target_len: usize, +} + +impl ExactPageAccumulator { + fn new(offset: u32, limit: u32) -> Self { + let offset = usize::try_from(offset).unwrap_or(usize::MAX); + let limit = usize::try_from(limit).unwrap_or(usize::MAX); + Self { + rows: Vec::new(), + offset, + limit, + target_len: offset.saturating_add(limit), + } + } + + fn is_full(&self) -> bool { + self.rows.len() >= self.target_len + } + + fn push_matching(&mut self, rows: I, mut predicate: F) + where + I: IntoIterator, + F: FnMut(&T) -> bool, + { + let remaining = self.target_len.saturating_sub(self.rows.len()); + self.rows.extend( + rows.into_iter() + .filter(|row| predicate(row)) + .take(remaining), + ); + } + + fn into_page(self) -> Vec { + self.rows + .into_iter() + .skip(self.offset) + .take(self.limit) + .collect() + } } impl MysqlMinimalCandidateSelectionReadRepository { @@ -115,6 +162,263 @@ impl MysqlMinimalCandidateSelectionReadRepository { let rows = self.load_rows_for_api_format(api_format).await?; Ok(sort_rows(select_pool_rows(rows), true)) } + + async fn load_selected_rows_for_api_format_page( + &self, + api_format: &str, + limit: u32, + offset: u32, + ) -> Result, DataLayerError> { + let mut builder = api_format_page_query(api_format, limit, offset); + let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; + rows.iter().map(map_candidate_selection_row).collect() + } + + async fn selected_rows_for_api_format_page( + &self, + query: &StoredApiFormatCandidateRowsQuery, + ) -> Result, DataLayerError> { + if query.limit == 0 { + return Ok(Vec::new()); + } + let target_len = query.offset.saturating_add(query.limit); + let mut raw_offset = 0_u32; + let mut selected = Vec::new(); + while selected.len() < target_len as usize && raw_offset < REQUESTED_MODEL_RAW_SCAN_LIMIT { + let raw_limit = + target_len.min(REQUESTED_MODEL_RAW_SCAN_LIMIT.saturating_sub(raw_offset)); + let rows = self + .load_selected_rows_for_api_format_page(&query.api_format, raw_limit, raw_offset) + .await?; + let raw_len = rows.len() as u32; + selected.extend(rows.into_iter().filter(|item| { + api_format_matches(&item.row.endpoint_api_format, &query.api_format) + && item.row.key_supports_api_format(&query.api_format) + && key_auth_channel_matches(item, &query.api_format) + })); + if raw_len < raw_limit || raw_len == 0 { + break; + } + let next_offset = raw_offset.saturating_add(raw_len); + if next_offset == raw_offset { + break; + } + raw_offset = next_offset; + } + let rows = selected.into_iter().map(|item| item.row).collect(); + Ok(sort_rows(dedupe_candidate_selection_rows(rows), true) + .into_iter() + .skip(query.offset as usize) + .take(query.limit as usize) + .collect()) + } +} + +fn api_format_page_query( + api_format: &str, + limit: u32, + offset: u32, +) -> QueryBuilder<'static, MySql> { + let mut builder = QueryBuilder::::new("WITH candidate_rows AS ("); + builder.push(CANDIDATE_SELECTION_COLUMNS); + push_candidate_api_format_filters(&mut builder, api_format); + push_selected_pool_rows(&mut builder); + builder.push( + r#" +ORDER BY + global_model_name ASC, + provider_priority ASC, + key_internal_priority ASC, + provider_id ASC, + endpoint_id ASC, + key_id ASC, + model_id ASC +LIMIT "#, + ); + builder.push_bind(i64::from(limit)); + builder.push(" OFFSET "); + builder.push_bind(i64::from(offset)); + builder +} + +fn requested_model_page_query( + api_format: &str, + requested_model_name: &str, + limit: u32, + offset: u32, +) -> QueryBuilder<'static, MySql> { + let mut builder = QueryBuilder::::new("WITH candidate_rows AS ("); + builder.push(CANDIDATE_SELECTION_COLUMNS); + push_candidate_api_format_filters(&mut builder, api_format); + push_requested_model_sql_filter(&mut builder, requested_model_name); + push_selected_pool_rows(&mut builder); + builder.push( + r#" +ORDER BY + global_model_name ASC, + provider_priority ASC, + key_internal_priority ASC, + provider_id ASC, + endpoint_id ASC, + key_id ASC, + model_id ASC +LIMIT "#, + ); + builder.push_bind(i64::from(limit)); + builder.push(" OFFSET "); + builder.push_bind(i64::from(offset)); + builder +} + +fn pool_key_group_query(query: &StoredPoolKeyCandidateRowsQuery) -> QueryBuilder<'static, MySql> { + let mut builder = QueryBuilder::::new(CANDIDATE_SELECTION_COLUMNS); + push_candidate_api_format_filters(&mut builder, &query.api_format); + push_pool_key_group_filters( + &mut builder, + &query.provider_id, + &query.endpoint_id, + &query.model_id, + ); + push_pool_key_order(&mut builder, &query.order); + builder.push(" LIMIT "); + builder.push_bind(i64::from(query.limit)); + builder.push(" OFFSET "); + builder.push_bind(i64::from(query.offset)); + builder +} + +fn pool_key_group_by_key_ids_query( + query: &StoredPoolKeyCandidateRowsByKeyIdsQuery, +) -> QueryBuilder<'static, MySql> { + let mut builder = QueryBuilder::::new(CANDIDATE_SELECTION_COLUMNS); + push_candidate_api_format_filters(&mut builder, &query.api_format); + push_pool_key_group_filters( + &mut builder, + &query.provider_id, + &query.endpoint_id, + &query.model_id, + ); + builder.push(" AND pak.id IN ("); + { + let mut separated = builder.separated(", "); + for key_id in &query.key_ids { + separated.push_bind(key_id.clone()); + } + } + builder.push(") ORDER BY FIELD(pak.id, "); + { + let mut separated = builder.separated(", "); + for key_id in &query.key_ids { + separated.push_bind(key_id.clone()); + } + } + builder.push(") ASC, pak.id ASC"); + builder +} + +fn push_candidate_api_format_filters(builder: &mut QueryBuilder<'_, MySql>, api_format: &str) { + let canonical_api_format = normalize_api_format(api_format); + let storage_aliases = sql_match_aliases(&api_format_aliases(&canonical_api_format)); + let permission_aliases = + sql_match_aliases(&api_format_permission_aliases(&canonical_api_format)); + builder.push(" AND LOWER(pe.api_format) IN ("); + { + let mut separated = builder.separated(", "); + for alias in storage_aliases { + separated.push_bind(alias); + } + } + builder.push(") AND (pak.api_formats IS NULL OR TRIM(pak.api_formats) = ''"); + for alias in permission_aliases { + builder.push(" OR JSON_SEARCH(LOWER(pak.api_formats), 'one', "); + builder.push_bind(alias); + builder.push(") IS NOT NULL"); + } + builder.push(")"); + push_key_auth_channel_filter(builder, &canonical_api_format); +} + +fn push_requested_model_sql_filter( + builder: &mut QueryBuilder<'_, MySql>, + requested_model_name: &str, +) { + builder.push(" AND (gm.name = "); + builder.push_bind(requested_model_name.to_string()); + builder.push(" OR m.provider_model_name = "); + builder.push_bind(requested_model_name.to_string()); + builder.push(" OR (m.provider_model_mappings IS NOT NULL AND LOCATE("); + builder.push_bind(requested_model_name.to_string()); + builder.push(", m.provider_model_mappings) > 0))"); +} + +fn push_selected_pool_rows(builder: &mut QueryBuilder<'_, MySql>) { + builder.push( + r#" +), +ranked_rows AS ( + SELECT + candidate_rows.*, + CASE + WHEN JSON_EXTRACT(provider_config, '$.pool_advanced') IS NOT NULL + AND JSON_TYPE(JSON_EXTRACT(provider_config, '$.pool_advanced')) <> 'NULL' + THEN 1 ELSE 0 + END AS provider_pool_enabled, + ROW_NUMBER() OVER ( + PARTITION BY provider_id, endpoint_id, model_id + ORDER BY key_internal_priority ASC, key_id ASC + ) AS pool_rank + FROM candidate_rows +), +selected_rows AS ( + SELECT * + FROM ranked_rows + WHERE provider_pool_enabled = 0 OR pool_rank = 1 +) +SELECT * +FROM selected_rows"#, + ); +} + +fn push_pool_key_group_filters( + builder: &mut QueryBuilder<'_, MySql>, + provider_id: &str, + endpoint_id: &str, + model_id: &str, +) { + builder.push(" AND p.id = "); + builder.push_bind(provider_id.to_string()); + builder.push(" AND pe.id = "); + builder.push_bind(endpoint_id.to_string()); + builder.push(" AND m.id = "); + builder.push_bind(model_id.to_string()); +} + +fn push_pool_key_order(builder: &mut QueryBuilder<'_, MySql>, order: &StoredPoolKeyCandidateOrder) { + match order { + StoredPoolKeyCandidateOrder::InternalPriority => { + builder.push(" ORDER BY pak.internal_priority ASC, pak.id ASC"); + } + StoredPoolKeyCandidateOrder::Lru => { + builder.push( + " ORDER BY pak.last_used_at IS NOT NULL ASC, pak.last_used_at ASC, pak.internal_priority ASC, pak.id ASC", + ); + } + StoredPoolKeyCandidateOrder::CacheAffinity => { + builder.push( + " ORDER BY pak.last_used_at IS NULL ASC, pak.last_used_at DESC, pak.internal_priority ASC, pak.id ASC", + ); + } + StoredPoolKeyCandidateOrder::SingleAccount => { + builder.push( + " ORDER BY pak.internal_priority ASC, pak.last_used_at IS NULL ASC, pak.last_used_at DESC, pak.id ASC", + ); + } + StoredPoolKeyCandidateOrder::LoadBalance { seed } => { + builder.push(" ORDER BY MD5(CONCAT("); + builder.push_bind(seed.clone()); + builder.push(", ':', pak.id)) ASC, pak.id ASC"); + } + } } #[async_trait] @@ -126,6 +430,13 @@ impl MinimalCandidateSelectionReadRepository for MysqlMinimalCandidateSelectionR self.selected_rows_for_api_format(api_format).await } + async fn list_for_exact_api_format_page( + &self, + query: &StoredApiFormatCandidateRowsQuery, + ) -> Result, DataLayerError> { + self.selected_rows_for_api_format_page(query).await + } + async fn list_for_exact_api_format_and_global_model( &self, api_format: &str, @@ -146,57 +457,76 @@ impl MinimalCandidateSelectionReadRepository for MysqlMinimalCandidateSelectionR api_format: &str, requested_model_name: &str, ) -> Result, DataLayerError> { - self.list_for_exact_api_format_and_requested_model_page( - &StoredRequestedModelCandidateRowsQuery { - api_format: api_format.to_string(), - requested_model_name: requested_model_name.to_string(), - offset: 0, - limit: u32::MAX, - }, - ) - .await + let rows = self + .selected_rows_for_api_format(api_format) + .await? + .into_iter() + .filter(|row| row_matches_requested_model(row, requested_model_name, api_format)) + .collect::>(); + Ok(sort_rows(rows, true)) } async fn list_for_exact_api_format_and_requested_model_page( &self, query: &StoredRequestedModelCandidateRowsQuery, ) -> Result, DataLayerError> { - let rows = self - .selected_rows_for_api_format(&query.api_format) - .await? + if query.limit == 0 { + return Ok(Vec::new()); + } + // The SQL model predicate is a coarse superset, so fill the exact page across raw windows. + let mut exact_page = ExactPageAccumulator::new(query.offset, query.limit); + let mut raw_offset = 0_u32; + while !exact_page.is_full() && raw_offset < REQUESTED_MODEL_RAW_SCAN_LIMIT { + let raw_limit = REQUESTED_MODEL_RAW_PAGE_SIZE + .min(REQUESTED_MODEL_RAW_SCAN_LIMIT.saturating_sub(raw_offset)); + let mut builder = requested_model_page_query( + &query.api_format, + &query.requested_model_name, + raw_limit, + raw_offset, + ); + let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; + let raw_len = u32::try_from(rows.len()).unwrap_or(u32::MAX); + let items = rows + .iter() + .map(map_candidate_selection_row) + .collect::, _>>()?; + exact_page.push_matching(items, |item| { + row_matches_requested_model( + &item.row, + &query.requested_model_name, + &query.api_format, + ) + }); + raw_offset = raw_offset.saturating_add(raw_len); + if raw_len < raw_limit || raw_len == 0 { + break; + } + } + let rows = exact_page + .into_page() .into_iter() - .filter(|row| { - row_matches_requested_model(row, &query.requested_model_name, &query.api_format) - }) - .collect::>(); - Ok(sort_rows(rows, true) - .into_iter() - .skip(query.offset as usize) - .take(query.limit as usize) - .collect()) + .map(|item| item.row) + .collect(); + Ok(sort_rows(dedupe_candidate_selection_rows(rows), true)) } async fn list_pool_key_rows_for_group( &self, query: &StoredPoolKeyCandidateRowsQuery, ) -> Result, DataLayerError> { - let rows = self - .load_rows_for_api_format(&query.api_format) - .await? - .into_iter() - .filter(|row| { - row.row.provider_id == query.provider_id - && row.row.endpoint_id == query.endpoint_id - && row.row.model_id == query.model_id - }) - .collect::>(); - let mut rows = sort_pool_key_rows(rows, &query.order); - Ok(rows - .drain(..) - .skip(query.offset as usize) - .take(query.limit as usize) - .map(|item| item.row) - .collect()) + if query.limit == 0 { + return Ok(Vec::new()); + } + let mut builder = pool_key_group_query(query); + let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; + let rows = rows + .iter() + .map(map_candidate_selection_row) + .collect::, _>>()?; + Ok(dedupe_candidate_selection_rows( + rows.into_iter().map(|item| item.row).collect(), + )) } async fn list_pool_key_rows_for_group_key_ids( @@ -212,25 +542,23 @@ impl MinimalCandidateSelectionReadRepository for MysqlMinimalCandidateSelectionR .enumerate() .map(|(index, key_id)| (key_id.as_str(), index)) .collect::>(); - let mut rows = self - .load_rows_for_api_format(&query.api_format) - .await? + let mut builder = pool_key_group_by_key_ids_query(query); + let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; + let mut rows = rows + .iter() + .map(map_candidate_selection_row) + .collect::, _>>()? .into_iter() - .filter(|row| { - row.row.provider_id == query.provider_id - && row.row.endpoint_id == query.endpoint_id - && row.row.model_id == query.model_id - && key_order.contains_key(row.row.key_id.as_str()) - }) .map(|item| item.row) .collect::>(); + rows = dedupe_candidate_selection_rows(rows); rows.sort_by(|left, right| { key_order .get(left.key_id.as_str()) .cmp(&key_order.get(right.key_id.as_str())) .then(left.key_id.cmp(&right.key_id)) }); - Ok(dedupe_candidate_selection_rows(rows)) + Ok(rows) } } @@ -283,67 +611,6 @@ fn sort_rows( rows } -fn sort_pool_key_rows( - mut rows: Vec, - order: &StoredPoolKeyCandidateOrder, -) -> Vec { - rows.sort_by(|left, right| match order { - StoredPoolKeyCandidateOrder::InternalPriority => compare_pool_key_internal(left, right), - StoredPoolKeyCandidateOrder::Lru => left - .key_last_used_at_unix_secs - .cmp(&right.key_last_used_at_unix_secs) - .then_with(|| compare_pool_key_internal(left, right)), - StoredPoolKeyCandidateOrder::CacheAffinity => right - .key_last_used_at_unix_secs - .cmp(&left.key_last_used_at_unix_secs) - .then_with(|| compare_pool_key_internal(left, right)), - StoredPoolKeyCandidateOrder::SingleAccount => left - .row - .key_internal_priority - .cmp(&right.row.key_internal_priority) - .then_with(|| { - right - .key_last_used_at_unix_secs - .cmp(&left.key_last_used_at_unix_secs) - }) - .then(left.row.key_id.cmp(&right.row.key_id)), - StoredPoolKeyCandidateOrder::LoadBalance { seed } => { - stable_pool_key_hash(seed.as_str(), left.row.key_id.as_str()) - .cmp(&stable_pool_key_hash( - seed.as_str(), - right.row.key_id.as_str(), - )) - .then(left.row.key_id.cmp(&right.row.key_id)) - } - }); - rows -} - -fn compare_pool_key_internal( - left: &CandidateSelectionRow, - right: &CandidateSelectionRow, -) -> std::cmp::Ordering { - left.row - .key_internal_priority - .cmp(&right.row.key_internal_priority) - .then(left.row.key_id.cmp(&right.row.key_id)) -} - -fn stable_pool_key_hash(seed: &str, key_id: &str) -> u64 { - let mut hash = 0xcbf29ce484222325u64; - for byte in seed - .as_bytes() - .iter() - .copied() - .chain(std::iter::once(b':')) - .chain(key_id.as_bytes().iter().copied()) - { - hash ^= u64::from(byte); - hash = hash.wrapping_mul(0x100000001b3); - } - hash -} - fn row_matches_requested_model( row: &StoredMinimalCandidateSelectionRow, requested_model_name: &str, @@ -418,6 +685,57 @@ fn mapping_scope_matches( }) } +fn push_key_auth_channel_filter(builder: &mut QueryBuilder<'_, MySql>, api_format: &str) { + builder.push(" AND ((LOWER(TRIM(p.provider_type)) = 'codex'"); + builder.push(" AND LOWER(TRIM(pak.auth_type)) = 'oauth' AND "); + builder.push_bind(api_format.to_string()); + builder.push( + " IN ('openai:responses', 'openai:responses:compact', 'openai:search', 'openai:image'))", + ); + + builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'chatgpt_web'"); + builder.push(" AND LOWER(TRIM(pak.auth_type)) IN ('oauth', 'bearer') AND "); + builder.push_bind(api_format.to_string()); + builder.push(" = 'openai:image')"); + + builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'claude_code'"); + builder.push(" AND LOWER(TRIM(pak.auth_type)) = 'oauth' AND "); + builder.push_bind(api_format.to_string()); + builder.push(" = 'claude:messages')"); + + builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'kiro' AND "); + builder.push_bind(api_format.to_string()); + builder.push(" = 'claude:messages' AND (LOWER(TRIM(pak.auth_type)) = 'oauth'"); + builder.push(" OR (LOWER(TRIM(pak.auth_type)) = 'bearer'"); + builder.push(" AND pak.auth_config IS NOT NULL AND TRIM(pak.auth_config) <> '')))"); + + builder.push(" OR (LOWER(TRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')"); + builder.push(" AND LOWER(TRIM(pak.auth_type)) = 'oauth' AND "); + builder.push_bind(api_format.to_string()); + builder.push(" = 'gemini:generate_content')"); + + builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'grok'"); + builder.push(" AND LOWER(TRIM(pak.auth_type)) = 'oauth' AND "); + builder.push_bind(api_format.to_string()); + builder.push(" IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image'))"); + + builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'windsurf'"); + builder.push(" AND LOWER(TRIM(pak.auth_type)) IN ('oauth', 'api_key', 'bearer') AND "); + builder.push_bind(api_format.to_string()); + builder.push(" = 'openai:chat')"); + + builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'vertex_ai'"); + builder.push( + " AND LOWER(TRIM(pak.auth_type)) IN ('api_key', 'service_account', 'vertex_ai') AND ", + ); + builder.push_bind(api_format.to_string()); + builder.push(" IN ('gemini:generate_content', 'gemini:embedding'))"); + + builder.push( + " OR (LOWER(TRIM(p.provider_type)) NOT IN ('chatgpt_web', 'claude_code', 'codex', 'gemini_cli', 'grok', 'vertex_ai', 'antigravity', 'kiro', 'windsurf') AND LOWER(TRIM(pak.auth_type)) <> 'oauth'))", + ); +} + fn key_auth_channel_matches(row: &CandidateSelectionRow, api_format: &str) -> bool { let provider_type = row.row.provider_type.trim().to_ascii_lowercase(); let auth_type = row.row.key_auth_type.trim().to_ascii_lowercase(); @@ -543,10 +861,6 @@ fn map_candidate_selection_row(row: &MySqlRow) -> Result, _>("key_last_used_at_unix_secs") - .map_sql_err()? - .and_then(|value| u64::try_from(value).ok()), }) } @@ -770,6 +1084,10 @@ fn api_format_aliases(api_format: &str) -> Vec { aether_ai_formats::api_format_storage_aliases(api_format) } +fn api_format_permission_aliases(api_format: &str) -> Vec { + aether_ai_formats::api_format_permission_storage_aliases(api_format) +} + fn normalize_api_format(api_format: &str) -> String { aether_ai_formats::normalize_api_format_alias(api_format) } @@ -791,7 +1109,134 @@ fn sql_match_aliases(api_formats: &[String]) -> Vec { #[cfg(test)] mod tests { - use super::{vertex_key_auth_channel_matches, MysqlMinimalCandidateSelectionReadRepository}; + use super::{ + api_format_page_query, pool_key_group_by_key_ids_query, pool_key_group_query, + requested_model_page_query, vertex_key_auth_channel_matches, ExactPageAccumulator, + MysqlMinimalCandidateSelectionReadRepository, REQUESTED_MODEL_RAW_SCAN_LIMIT, + }; + use aether_data_contracts::repository::candidate_selection::{ + StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery, + StoredPoolKeyCandidateRowsQuery, + }; + + #[test] + fn api_format_page_query_uses_portable_sql_pagination_and_stable_order() { + let query = api_format_page_query("openai:chat", 256, 512); + let sql = query.sql(); + + assert!(sql.contains("ROW_NUMBER() OVER (")); + assert!(sql.contains("JSON_SEARCH(LOWER(pak.api_formats), 'one', ?")); + assert!(sql.contains("WHERE provider_pool_enabled = 0 OR pool_rank = 1")); + assert!(sql.contains( + "ORDER BY\n global_model_name ASC,\n provider_priority ASC,\n key_internal_priority ASC," + )); + assert!(sql.contains("LIMIT ? OFFSET ?")); + } + + #[test] + fn requested_model_page_query_filters_and_pages_before_fetch() { + let query = requested_model_page_query("openai:chat", "gpt-5", 256, 256); + let sql = query.sql(); + + assert!(sql.contains("LOWER(pe.api_format) IN (")); + assert!(sql.contains("JSON_SEARCH(LOWER(pak.api_formats), 'one', ?")); + assert!(sql.contains("LOWER(TRIM(p.provider_type)) = 'codex'")); + assert!(sql.contains("AND (gm.name = ? OR m.provider_model_name = ?")); + assert!(sql.contains("LOCATE(?, m.provider_model_mappings) > 0")); + assert!(sql.contains("ROW_NUMBER() OVER (")); + assert!(sql.contains("LIMIT ? OFFSET ?")); + } + + #[test] + fn exact_page_accumulator_continues_after_coarse_false_positives() { + let mut accumulator = ExactPageAccumulator::new(1, 2); + accumulator.push_matching(vec![("coarse-1", false), ("coarse-2", false)], |row| row.1); + assert!(!accumulator.is_full()); + + accumulator.push_matching( + vec![ + ("exact-1", true), + ("coarse-3", false), + ("exact-2", true), + ("exact-3", true), + ], + |row| row.1, + ); + + assert!(accumulator.is_full()); + assert_eq!( + accumulator.into_page(), + vec![("exact-2", true), ("exact-3", true)] + ); + assert_eq!(REQUESTED_MODEL_RAW_SCAN_LIMIT, 2048); + } + + #[test] + fn pool_group_query_pushes_group_filters_order_and_page_into_sql() { + let query = pool_key_group_query(&StoredPoolKeyCandidateRowsQuery { + api_format: "openai:chat".to_string(), + provider_id: "provider-1".to_string(), + endpoint_id: "endpoint-1".to_string(), + model_id: "model-1".to_string(), + selected_provider_model_name: "gpt-5".to_string(), + order: StoredPoolKeyCandidateOrder::Lru, + offset: 64, + limit: 64, + }); + let sql = query.sql(); + + assert!(sql.contains("LOWER(pe.api_format) IN (")); + assert!(sql.contains("JSON_SEARCH(LOWER(pak.api_formats), 'one', ?")); + assert!(sql.contains("LOWER(TRIM(p.provider_type)) = 'codex'")); + assert!(sql.contains("AND p.id = ? AND pe.id = ? AND m.id = ?")); + assert!(sql.contains( + "ORDER BY pak.last_used_at IS NOT NULL ASC, pak.last_used_at ASC, pak.internal_priority ASC, pak.id ASC" + )); + assert!(sql.contains("LIMIT ? OFFSET ?")); + } + + #[test] + fn load_balance_pool_group_query_uses_seeded_sql_order_before_paging() { + let query = pool_key_group_query(&StoredPoolKeyCandidateRowsQuery { + api_format: "openai:chat".to_string(), + provider_id: "provider-1".to_string(), + endpoint_id: "endpoint-1".to_string(), + model_id: "model-1".to_string(), + selected_provider_model_name: "gpt-5".to_string(), + order: StoredPoolKeyCandidateOrder::LoadBalance { + seed: "request-1".to_string(), + }, + offset: 128, + limit: 64, + }); + let sql = query.sql(); + + assert!(sql.contains("AND p.id = ? AND pe.id = ? AND m.id = ?")); + assert!( + sql.contains("ORDER BY MD5(CONCAT(?, ':', pak.id)) ASC, pak.id ASC LIMIT ? OFFSET ?") + ); + assert!(!sql.contains("ROW_NUMBER() OVER (")); + } + + #[test] + fn pool_group_by_key_ids_query_filters_ids_and_preserves_requested_order() { + let query = pool_key_group_by_key_ids_query(&StoredPoolKeyCandidateRowsByKeyIdsQuery { + api_format: "openai:chat".to_string(), + provider_id: "provider-1".to_string(), + endpoint_id: "endpoint-1".to_string(), + model_id: "model-1".to_string(), + selected_provider_model_name: "gpt-5".to_string(), + key_ids: vec!["key-2".to_string(), "key-1".to_string()], + }); + let sql = query.sql(); + + assert!(sql.contains("LOWER(pe.api_format) IN (")); + assert!(sql.contains("JSON_SEARCH(LOWER(pak.api_formats), 'one', ?")); + assert!(sql.contains("LOWER(TRIM(p.provider_type)) = 'codex'")); + assert!(sql.contains("AND p.id = ? AND pe.id = ? AND m.id = ?")); + assert!(sql.contains("AND pak.id IN (?, ?)")); + assert!(sql.contains("ORDER BY FIELD(pak.id, ?, ?) ASC, pak.id ASC")); + } #[tokio::test] async fn repository_builds_from_lazy_pool() { diff --git a/crates/aether-data/adapters/mysql/src/routing_profiles.rs b/crates/aether-data/adapters/mysql/src/routing_profiles.rs index 06ec187ab..0d5b76e17 100644 --- a/crates/aether-data/adapters/mysql/src/routing_profiles.rs +++ b/crates/aether-data/adapters/mysql/src/routing_profiles.rs @@ -1,6 +1,6 @@ use async_trait::async_trait; use serde_json::Value; -use sqlx::{mysql::MySqlRow, Row}; +use sqlx::{mysql::MySqlRow, Acquire, Row}; use aether_data_contracts::repository::routing_profiles::*; use aether_data_contracts::DataLayerError; @@ -56,24 +56,6 @@ impl MysqlRoutingGroupRepository { pub fn new(pool: MysqlPool) -> Self { Self { pool } } - - async fn reload_group(&self, id: &str) -> Result, DataLayerError> { - self.find_routing_group(RoutingGroupLookupKey::Id(id)).await - } - - async fn find_binding_by_id( - &self, - id: &str, - ) -> Result, DataLayerError> { - let row = sqlx::query(&format!( - "{ROUTING_GROUP_BINDING_SELECT} WHERE id = ? LIMIT 1" - )) - .bind(id) - .fetch_optional(&self.pool) - .await - .map_sql_err()?; - row.as_ref().map(map_binding_row).transpose() - } } #[async_trait] @@ -170,6 +152,24 @@ impl RoutingGroupWriteRepository for MysqlRoutingGroupRepository { record: CreateRoutingGroupRecord, ) -> Result { let group = StoredRoutingGroup::new(record)?; + let mut connection = self.pool.acquire().await.map_sql_err()?; + sqlx::query("SET TRANSACTION ISOLATION LEVEL SERIALIZABLE") + .execute(&mut *connection) + .await + .map_sql_err()?; + let mut tx = connection.begin().await.map_sql_err()?; + sqlx::query("SELECT id FROM routing_groups ORDER BY id FOR UPDATE") + .fetch_all(&mut *tx) + .await + .map_sql_err()?; + if group.is_system_default { + sqlx::query( + "UPDATE routing_groups SET is_system_default = 0 WHERE is_system_default = 1", + ) + .execute(&mut *tx) + .await + .map_sql_err()?; + } sqlx::query( r#" INSERT INTO routing_groups ( @@ -192,9 +192,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) .bind(group.created_at) .bind(group.updated_at) .bind(group.published_at) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; + tx.commit().await.map_sql_err()?; Ok(group) } @@ -203,10 +204,36 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) id: &str, patch: UpdateRoutingGroupRecord, ) -> Result, DataLayerError> { - let Some(mut group) = self.reload_group(id).await? else { + let mut connection = self.pool.acquire().await.map_sql_err()?; + sqlx::query("SET TRANSACTION ISOLATION LEVEL SERIALIZABLE") + .execute(&mut *connection) + .await + .map_sql_err()?; + let mut tx = connection.begin().await.map_sql_err()?; + sqlx::query("SELECT id FROM routing_groups ORDER BY id FOR UPDATE") + .fetch_all(&mut *tx) + .await + .map_sql_err()?; + let row = sqlx::query(&format!( + "{ROUTING_GROUP_SELECT} WHERE id = ? LIMIT 1 FOR UPDATE" + )) + .bind(id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(mut group) = row.as_ref().map(map_group_row).transpose()? else { return Ok(None); }; apply_group_patch(&mut group, patch)?; + if group.is_system_default { + sqlx::query( + "UPDATE routing_groups SET is_system_default = 0 WHERE is_system_default = 1 AND id <> ?", + ) + .bind(id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } sqlx::query( r#" UPDATE routing_groups @@ -233,9 +260,10 @@ WHERE id = ? .bind(group.updated_at) .bind(group.published_at) .bind(id) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; + tx.commit().await.map_sql_err()?; Ok(Some(group)) } @@ -266,6 +294,30 @@ WHERE id = ? record: CreateRoutingGroupBindingRecord, ) -> Result { let binding = StoredRoutingGroupBinding::new(record)?; + let mut connection = self.pool.acquire().await.map_sql_err()?; + sqlx::query("SET TRANSACTION ISOLATION LEVEL SERIALIZABLE") + .execute(&mut *connection) + .await + .map_sql_err()?; + let mut tx = connection.begin().await.map_sql_err()?; + sqlx::query("SELECT id FROM routing_group_bindings ORDER BY id FOR UPDATE") + .fetch_all(&mut *tx) + .await + .map_sql_err()?; + if binding.is_default { + sqlx::query( + r#" +UPDATE routing_group_bindings +SET is_default = 0 +WHERE is_default = 1 AND subject_type = ? AND subject_id = ? +"#, + ) + .bind(binding_subject_to_database(binding.subject_type)) + .bind(&binding.subject_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } sqlx::query( r#" INSERT INTO routing_group_bindings ( @@ -283,9 +335,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?) .bind(binding.allow_explicit_select) .bind(binding.created_at) .bind(binding.updated_at) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; + tx.commit().await.map_sql_err()?; Ok(binding) } @@ -306,10 +359,45 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?) id: &str, patch: UpdateRoutingGroupBindingRecord, ) -> Result, DataLayerError> { - let Some(mut binding) = self.find_binding_by_id(id).await? else { + let mut connection = self.pool.acquire().await.map_sql_err()?; + sqlx::query("SET TRANSACTION ISOLATION LEVEL SERIALIZABLE") + .execute(&mut *connection) + .await + .map_sql_err()?; + let mut tx = connection.begin().await.map_sql_err()?; + sqlx::query("SELECT id FROM routing_group_bindings ORDER BY id FOR UPDATE") + .fetch_all(&mut *tx) + .await + .map_sql_err()?; + let row = sqlx::query(&format!( + "{ROUTING_GROUP_BINDING_SELECT} WHERE id = ? LIMIT 1 FOR UPDATE" + )) + .bind(id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(mut binding) = row.as_ref().map(map_binding_row).transpose()? else { return Ok(None); }; apply_binding_patch(&mut binding, patch)?; + if binding.is_default { + sqlx::query( + r#" +UPDATE routing_group_bindings +SET is_default = 0 +WHERE is_default = 1 + AND subject_type = ? + AND subject_id = ? + AND id <> ? +"#, + ) + .bind(binding_subject_to_database(binding.subject_type)) + .bind(&binding.subject_id) + .bind(id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } sqlx::query( r#" UPDATE routing_group_bindings @@ -329,9 +417,10 @@ WHERE id = ? .bind(binding.allow_explicit_select) .bind(binding.updated_at) .bind(id) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; + tx.commit().await.map_sql_err()?; Ok(Some(binding)) } diff --git a/crates/aether-data/adapters/postgres/migrations/20260727000000_repair_routing_default_uniqueness.sql b/crates/aether-data/adapters/postgres/migrations/20260727000000_repair_routing_default_uniqueness.sql new file mode 100644 index 000000000..16a9713ef --- /dev/null +++ b/crates/aether-data/adapters/postgres/migrations/20260727000000_repair_routing_default_uniqueness.sql @@ -0,0 +1,36 @@ +WITH ranked_defaults AS ( + SELECT + id, + ROW_NUMBER() OVER (ORDER BY enabled DESC, updated_at DESC, id ASC) AS default_rank + FROM public.routing_groups + WHERE is_system_default = TRUE +) +UPDATE public.routing_groups AS routing_group +SET is_system_default = FALSE +FROM ranked_defaults +WHERE routing_group.id = ranked_defaults.id + AND ranked_defaults.default_rank > 1; + +WITH ranked_defaults AS ( + SELECT + id, + ROW_NUMBER() OVER ( + PARTITION BY subject_type, subject_id + ORDER BY created_at ASC, id ASC + ) AS default_rank + FROM public.routing_group_bindings + WHERE is_default = TRUE +) +UPDATE public.routing_group_bindings AS binding +SET is_default = FALSE +FROM ranked_defaults +WHERE binding.id = ranked_defaults.id + AND ranked_defaults.default_rank > 1; + +CREATE UNIQUE INDEX IF NOT EXISTS routing_groups_one_system_default_key + ON public.routing_groups (is_system_default) + WHERE is_system_default = TRUE; + +CREATE UNIQUE INDEX IF NOT EXISTS routing_group_bindings_subject_default_key + ON public.routing_group_bindings (subject_type, subject_id) + WHERE is_default = TRUE; diff --git a/crates/aether-data/adapters/postgres/src/candidate_selection.rs b/crates/aether-data/adapters/postgres/src/candidate_selection.rs index b34179137..47a912158 100644 --- a/crates/aether-data/adapters/postgres/src/candidate_selection.rs +++ b/crates/aether-data/adapters/postgres/src/candidate_selection.rs @@ -4,10 +4,10 @@ use sqlx::{PgPool, Row}; use std::collections::BTreeSet; use aether_data_contracts::repository::candidate_selection::{ - MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow, - StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery, - StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping, - StoredRequestedModelCandidateRowsQuery, + MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery, + StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder, + StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery, + StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery, }; use aether_data_contracts::DataLayerError; @@ -764,6 +764,44 @@ impl SqlxMinimalCandidateSelectionReadRepository { Ok(dedupe_candidate_selection_rows(rows)) } + pub async fn list_for_exact_api_format_page( + &self, + query: &StoredApiFormatCandidateRowsQuery, + ) -> Result, DataLayerError> { + if query.limit == 0 { + return Ok(Vec::new()); + } + let canonical_api_format = normalize_api_format(&query.api_format); + let storage_aliases = api_format_aliases(&canonical_api_format); + let sql_match_aliases = + sql_match_aliases(&api_format_permission_aliases(&canonical_api_format)); + let fetch_limit = i64::from(query.offset.saturating_add(query.limit)); + let sql = format!("{LIST_FOR_EXACT_API_FORMAT_SQL}\nLIMIT $4\nOFFSET $5"); + let mut rows = Vec::new(); + for api_format in storage_aliases { + rows.extend( + Self::collect_query_rows( + sqlx::query(sql.as_str()) + .bind(api_format) + .bind(sql_match_aliases.clone()) + .bind(canonical_api_format.clone()) + .bind(fetch_limit) + .bind(0_i64) + .fetch(&self.pool), + map_candidate_selection_row, + ) + .await?, + ); + } + let mut rows = dedupe_candidate_selection_rows(rows); + sort_candidate_selection_rows(&mut rows, true); + Ok(rows + .into_iter() + .skip(query.offset as usize) + .take(query.limit as usize) + .collect()) + } + pub async fn list_for_exact_api_format_and_global_model( &self, api_format: &str, @@ -1109,6 +1147,24 @@ fn dedupe_candidate_selection_rows( .collect() } +fn sort_candidate_selection_rows( + rows: &mut [StoredMinimalCandidateSelectionRow], + include_global_model: bool, +) { + rows.sort_by(|left, right| { + let global_model_order = include_global_model + .then(|| left.global_model_name.cmp(&right.global_model_name)) + .unwrap_or(std::cmp::Ordering::Equal); + global_model_order + .then(left.provider_priority.cmp(&right.provider_priority)) + .then(left.key_internal_priority.cmp(&right.key_internal_priority)) + .then(left.provider_id.cmp(&right.provider_id)) + .then(left.endpoint_id.cmp(&right.endpoint_id)) + .then(left.key_id.cmp(&right.key_id)) + .then(left.model_id.cmp(&right.model_id)) + }); +} + #[async_trait] impl MinimalCandidateSelectionReadRepository for SqlxMinimalCandidateSelectionReadRepository { async fn list_for_exact_api_format( @@ -1118,6 +1174,13 @@ impl MinimalCandidateSelectionReadRepository for SqlxMinimalCandidateSelectionRe Self::list_for_exact_api_format(self, api_format).await } + async fn list_for_exact_api_format_page( + &self, + query: &StoredApiFormatCandidateRowsQuery, + ) -> Result, DataLayerError> { + Self::list_for_exact_api_format_page(self, query).await + } + async fn list_for_exact_api_format_and_global_model( &self, api_format: &str, diff --git a/crates/aether-data/adapters/postgres/src/routing_profiles.rs b/crates/aether-data/adapters/postgres/src/routing_profiles.rs index a5f35626c..e34003efc 100644 --- a/crates/aether-data/adapters/postgres/src/routing_profiles.rs +++ b/crates/aether-data/adapters/postgres/src/routing_profiles.rs @@ -63,24 +63,6 @@ impl PostgresRoutingGroupRepository { pub fn new(pool: PgPool) -> Self { Self { pool } } - - async fn reload_group(&self, id: &str) -> Result, DataLayerError> { - self.find_routing_group(RoutingGroupLookupKey::Id(id)).await - } - - async fn find_binding_by_id( - &self, - id: &str, - ) -> Result, DataLayerError> { - let row = sqlx::query(&format!( - "{ROUTING_GROUP_BINDING_SELECT} WHERE id = $1 LIMIT 1" - )) - .bind(id) - .fetch_optional(&self.pool) - .await - .map_postgres_err()?; - row.as_ref().map(map_binding_row).transpose() - } } #[async_trait] @@ -180,6 +162,19 @@ impl RoutingGroupWriteRepository for PostgresRoutingGroupRepository { record: CreateRoutingGroupRecord, ) -> Result { let group = StoredRoutingGroup::new(record)?; + let mut tx = self.pool.begin().await.map_postgres_err()?; + if group.is_system_default { + sqlx::query("LOCK TABLE routing_groups IN SHARE ROW EXCLUSIVE MODE") + .execute(&mut *tx) + .await + .map_postgres_err()?; + sqlx::query( + "UPDATE routing_groups SET is_system_default = FALSE WHERE is_system_default = TRUE", + ) + .execute(&mut *tx) + .await + .map_postgres_err()?; + } sqlx::query( r#" INSERT INTO routing_groups ( @@ -199,9 +194,10 @@ VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10) .bind(group.created_at) .bind(group.updated_at) .bind(group.published_at) - .execute(&self.pool) + .execute(&mut *tx) .await .map_postgres_err()?; + tx.commit().await.map_postgres_err()?; Ok(group) } @@ -210,10 +206,31 @@ VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10) id: &str, patch: UpdateRoutingGroupRecord, ) -> Result, DataLayerError> { - let Some(mut group) = self.reload_group(id).await? else { + let mut tx = self.pool.begin().await.map_postgres_err()?; + sqlx::query("LOCK TABLE routing_groups IN SHARE ROW EXCLUSIVE MODE") + .execute(&mut *tx) + .await + .map_postgres_err()?; + let row = sqlx::query(&format!( + "{ROUTING_GROUP_SELECT} WHERE id = $1 LIMIT 1 FOR UPDATE" + )) + .bind(id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let Some(mut group) = row.as_ref().map(map_group_row).transpose()? else { return Ok(None); }; apply_group_patch(&mut group, patch)?; + if group.is_system_default { + sqlx::query( + "UPDATE routing_groups SET is_system_default = FALSE WHERE is_system_default = TRUE AND id <> $1", + ) + .bind(id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + } sqlx::query( r#" UPDATE routing_groups @@ -237,9 +254,10 @@ WHERE id = $1 .bind(group.version) .bind(group.updated_at) .bind(group.published_at) - .execute(&self.pool) + .execute(&mut *tx) .await .map_postgres_err()?; + tx.commit().await.map_postgres_err()?; Ok(Some(group)) } @@ -270,6 +288,25 @@ WHERE id = $1 record: CreateRoutingGroupBindingRecord, ) -> Result { let binding = StoredRoutingGroupBinding::new(record)?; + let mut tx = self.pool.begin().await.map_postgres_err()?; + if binding.is_default { + sqlx::query("LOCK TABLE routing_group_bindings IN SHARE ROW EXCLUSIVE MODE") + .execute(&mut *tx) + .await + .map_postgres_err()?; + sqlx::query( + r#" +UPDATE routing_group_bindings +SET is_default = FALSE +WHERE is_default = TRUE AND subject_type = $1 AND subject_id = $2 +"#, + ) + .bind(binding_subject_to_database(binding.subject_type)) + .bind(&binding.subject_id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + } sqlx::query( r#" INSERT INTO routing_group_bindings ( @@ -287,9 +324,10 @@ VALUES ($1, $2, $3, $4, $5, $6, $7, $8) .bind(binding.allow_explicit_select) .bind(binding.created_at) .bind(binding.updated_at) - .execute(&self.pool) + .execute(&mut *tx) .await .map_postgres_err()?; + tx.commit().await.map_postgres_err()?; Ok(binding) } @@ -310,10 +348,40 @@ VALUES ($1, $2, $3, $4, $5, $6, $7, $8) id: &str, patch: UpdateRoutingGroupBindingRecord, ) -> Result, DataLayerError> { - let Some(mut binding) = self.find_binding_by_id(id).await? else { + let mut tx = self.pool.begin().await.map_postgres_err()?; + sqlx::query("LOCK TABLE routing_group_bindings IN SHARE ROW EXCLUSIVE MODE") + .execute(&mut *tx) + .await + .map_postgres_err()?; + let row = sqlx::query(&format!( + "{ROUTING_GROUP_BINDING_SELECT} WHERE id = $1 LIMIT 1 FOR UPDATE" + )) + .bind(id) + .fetch_optional(&mut *tx) + .await + .map_postgres_err()?; + let Some(mut binding) = row.as_ref().map(map_binding_row).transpose()? else { return Ok(None); }; apply_binding_patch(&mut binding, patch)?; + if binding.is_default { + sqlx::query( + r#" +UPDATE routing_group_bindings +SET is_default = FALSE +WHERE is_default = TRUE + AND subject_type = $1 + AND subject_id = $2 + AND id <> $3 +"#, + ) + .bind(binding_subject_to_database(binding.subject_type)) + .bind(&binding.subject_id) + .bind(id) + .execute(&mut *tx) + .await + .map_postgres_err()?; + } sqlx::query( r#" UPDATE routing_group_bindings @@ -333,9 +401,10 @@ WHERE id = $1 .bind(binding.is_default) .bind(binding.allow_explicit_select) .bind(binding.updated_at) - .execute(&self.pool) + .execute(&mut *tx) .await .map_postgres_err()?; + tx.commit().await.map_postgres_err()?; Ok(Some(binding)) } diff --git a/crates/aether-data/adapters/sqlite/migrations/20260727000000_repair_routing_default_uniqueness.sql b/crates/aether-data/adapters/sqlite/migrations/20260727000000_repair_routing_default_uniqueness.sql new file mode 100644 index 000000000..805705e63 --- /dev/null +++ b/crates/aether-data/adapters/sqlite/migrations/20260727000000_repair_routing_default_uniqueness.sql @@ -0,0 +1,38 @@ +UPDATE routing_groups +SET is_system_default = 0 +WHERE id IN ( + SELECT id + FROM ( + SELECT + id, + ROW_NUMBER() OVER (ORDER BY enabled DESC, updated_at DESC, id ASC) AS default_rank + FROM routing_groups + WHERE is_system_default = 1 + ) AS ranked_defaults + WHERE default_rank > 1 +); + +UPDATE routing_group_bindings +SET is_default = 0 +WHERE id IN ( + SELECT id + FROM ( + SELECT + id, + ROW_NUMBER() OVER ( + PARTITION BY subject_type, subject_id + ORDER BY created_at ASC, id ASC + ) AS default_rank + FROM routing_group_bindings + WHERE is_default = 1 + ) AS ranked_defaults + WHERE default_rank > 1 +); + +CREATE UNIQUE INDEX IF NOT EXISTS routing_groups_one_system_default_key + ON routing_groups (is_system_default) + WHERE is_system_default = 1; + +CREATE UNIQUE INDEX IF NOT EXISTS routing_group_bindings_subject_default_key + ON routing_group_bindings (subject_type, subject_id) + WHERE is_default = 1; diff --git a/crates/aether-data/adapters/sqlite/src/candidate_selection.rs b/crates/aether-data/adapters/sqlite/src/candidate_selection.rs index d09d78197..b8d4cf38f 100644 --- a/crates/aether-data/adapters/sqlite/src/candidate_selection.rs +++ b/crates/aether-data/adapters/sqlite/src/candidate_selection.rs @@ -4,10 +4,10 @@ use async_trait::async_trait; use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite}; use aether_data_contracts::repository::candidate_selection::{ - MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow, - StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery, - StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping, - StoredRequestedModelCandidateRowsQuery, + MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery, + StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder, + StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery, + StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery, }; use aether_data_contracts::DataLayerError; @@ -68,6 +68,9 @@ WHERE p.is_active = 1 AND gm.is_active = 1 "#; +const REQUESTED_MODEL_RAW_PAGE_SIZE: u32 = 256; +const REQUESTED_MODEL_RAW_SCAN_LIMIT: u32 = 2048; + #[derive(Debug, Clone)] pub struct SqliteMinimalCandidateSelectionReadRepository { pool: SqlitePool, @@ -77,7 +80,6 @@ pub struct SqliteMinimalCandidateSelectionReadRepository { struct CandidateSelectionRow { row: StoredMinimalCandidateSelectionRow, key_auth_config: Option, - key_last_used_at_unix_secs: Option, } #[derive(Debug, Clone, Copy)] @@ -99,6 +101,58 @@ struct SqlPage { offset: i64, } +#[derive(Debug)] +struct ExactPageAccumulator { + rows: Vec, + offset: usize, + limit: usize, + target_len: usize, +} + +impl ExactPageAccumulator { + fn new(offset: u32, limit: u32) -> Self { + let offset = usize::try_from(offset).unwrap_or(usize::MAX); + let limit = usize::try_from(limit).unwrap_or(usize::MAX); + Self { + rows: Vec::new(), + offset, + limit, + target_len: offset.saturating_add(limit), + } + } + + fn is_full(&self) -> bool { + self.rows.len() >= self.target_len + } + + fn push_matching(&mut self, rows: I, mut predicate: F) + where + I: IntoIterator, + F: FnMut(&T) -> bool, + { + let remaining = self.target_len.saturating_sub(self.rows.len()); + self.rows.extend( + rows.into_iter() + .filter(|row| predicate(row)) + .take(remaining), + ); + } + + fn into_page(self) -> Vec { + self.rows + .into_iter() + .skip(self.offset) + .take(self.limit) + .collect() + } +} + +#[derive(Debug)] +struct RequestedModelRawPage { + rows: Vec, + raw_len: u32, +} + impl SqliteMinimalCandidateSelectionReadRepository { pub fn new(pool: SqlitePool) -> Self { Self { pool } @@ -148,44 +202,7 @@ impl SqliteMinimalCandidateSelectionReadRepository { ); } } - builder.push( - r#" -), -pool_rows AS ( - SELECT candidate.* - FROM candidate_rows candidate - WHERE candidate.provider_pool_enabled = 1 - AND NOT EXISTS ( - SELECT 1 - FROM candidate_rows other - WHERE other.provider_pool_enabled = 1 - AND other.provider_id = candidate.provider_id - AND other.endpoint_id = candidate.endpoint_id - AND other.model_id = candidate.model_id - AND ( - other.key_internal_priority < candidate.key_internal_priority - OR ( - other.key_internal_priority = candidate.key_internal_priority - AND other.key_id < candidate.key_id - ) - ) - ) -), -selected_rows AS ( - SELECT * FROM candidate_rows WHERE provider_pool_enabled = 0 - UNION ALL - SELECT * FROM pool_rows -) -SELECT * FROM selected_rows -"#, - ); - push_selected_rows_order(&mut builder, order); - if let Some(page) = page { - builder.push(" LIMIT "); - builder.push_bind(page.limit); - builder.push(" OFFSET "); - builder.push_bind(page.offset); - } + push_selected_rows_query_tail(&mut builder, order, page); let query_rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; let mut items = query_rows @@ -211,6 +228,42 @@ SELECT * FROM selected_rows }; Ok(dedupe_candidate_selection_rows(rows)) } + + async fn load_requested_model_raw_page( + &self, + api_format: &str, + requested_model_name: &str, + page: SqlPage, + ) -> Result { + let canonical_api_format = normalize_api_format(api_format); + let storage_aliases = api_format_aliases(&canonical_api_format); + let match_aliases = + sql_match_aliases(&api_format_permission_aliases(&canonical_api_format)); + let mut builder = QueryBuilder::::new("WITH candidate_rows AS ("); + builder.push(CANDIDATE_SELECTION_COLUMNS); + push_candidate_sql_filters_for_aliases( + &mut builder, + &storage_aliases, + &match_aliases, + &canonical_api_format, + ); + push_requested_model_sql_filter(&mut builder, requested_model_name, &match_aliases); + push_selected_rows_query_tail(&mut builder, SelectedRowsOrder::WithGlobalModel, Some(page)); + + let query_rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; + let raw_len = u32::try_from(query_rows.len()).unwrap_or(u32::MAX); + let mut rows = query_rows + .iter() + .map(map_candidate_selection_row) + .collect::, _>>()?; + rows.retain(|item| { + api_format_matches(&item.row.endpoint_api_format, &canonical_api_format) + && item.row.key_supports_api_format(&canonical_api_format) + && key_auth_channel_matches(item, &canonical_api_format) + }); + + Ok(RequestedModelRawPage { rows, raw_len }) + } } #[async_trait] @@ -222,6 +275,33 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection self.selected_rows_for_api_format(api_format).await } + async fn list_for_exact_api_format_page( + &self, + query: &StoredApiFormatCandidateRowsQuery, + ) -> Result, DataLayerError> { + if query.limit == 0 { + return Ok(Vec::new()); + } + let fetch_limit = query.offset.saturating_add(query.limit); + let mut rows = self + .load_selected_rows_for_api_format( + &query.api_format, + SelectedRowsFilter::None, + SelectedRowsOrder::WithGlobalModel, + Some(SqlPage { + limit: i64::from(fetch_limit), + offset: 0, + }), + ) + .await?; + sort_candidate_selection_rows(&mut rows, true); + Ok(rows + .into_iter() + .skip(query.offset as usize) + .take(query.limit as usize) + .collect()) + } + async fn list_for_exact_api_format_and_global_model( &self, api_format: &str, @@ -254,28 +334,58 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection &self, query: &StoredRequestedModelCandidateRowsQuery, ) -> Result, DataLayerError> { - self.load_selected_rows_for_api_format( - &query.api_format, - SelectedRowsFilter::RequestedModel(&query.requested_model_name), - SelectedRowsOrder::WithGlobalModel, - Some(SqlPage { - limit: i64::from(query.limit.max(1)), - offset: i64::from(query.offset), - }), - ) - .await + if query.limit == 0 { + return Ok(Vec::new()); + } + + let mut exact_page = ExactPageAccumulator::new(query.offset, query.limit); + let mut raw_offset = 0_u32; + while !exact_page.is_full() && raw_offset < REQUESTED_MODEL_RAW_SCAN_LIMIT { + let raw_limit = REQUESTED_MODEL_RAW_PAGE_SIZE + .min(REQUESTED_MODEL_RAW_SCAN_LIMIT.saturating_sub(raw_offset)); + let raw_page = self + .load_requested_model_raw_page( + &query.api_format, + &query.requested_model_name, + SqlPage { + limit: i64::from(raw_limit), + offset: i64::from(raw_offset), + }, + ) + .await?; + exact_page.push_matching(raw_page.rows, |item| { + row_matches_requested_model( + &item.row, + &query.requested_model_name, + &query.api_format, + ) + }); + raw_offset = raw_offset.saturating_add(raw_page.raw_len); + if raw_page.raw_len < raw_limit || raw_page.raw_len == 0 { + break; + } + } + + let rows = exact_page + .into_page() + .into_iter() + .map(|item| item.row) + .collect(); + Ok(dedupe_candidate_selection_rows(rows)) } async fn list_pool_key_rows_for_group( &self, query: &StoredPoolKeyCandidateRowsQuery, ) -> Result, DataLayerError> { + if query.limit == 0 { + return Ok(Vec::new()); + } let canonical_api_format = normalize_api_format(&query.api_format); let storage_aliases = api_format_aliases(&canonical_api_format); let match_aliases = sql_match_aliases(&api_format_permission_aliases(&canonical_api_format)); let mut rows = Vec::::new(); - let page_in_sql = !matches!(query.order, StoredPoolKeyCandidateOrder::LoadBalance { .. }); for storage_api_format in storage_aliases { let mut builder = QueryBuilder::::new(CANDIDATE_SELECTION_COLUMNS); @@ -286,15 +396,11 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection builder.push_bind(&query.endpoint_id); builder.push(" AND m.id = "); builder.push_bind(&query.model_id); - if page_in_sql { - push_pool_key_order(&mut builder, &query.order); - builder.push(" LIMIT "); - builder.push_bind(i64::from(query.limit.max(1))); - builder.push(" OFFSET "); - builder.push_bind(i64::from(query.offset)); - } else { - builder.push(" ORDER BY pak.id ASC"); - } + push_pool_key_order(&mut builder, &query.order); + builder.push(" LIMIT "); + builder.push_bind(i64::from(query.limit)); + builder.push(" OFFSET "); + builder.push_bind(i64::from(query.offset)); let query_rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; let mut items = query_rows @@ -309,20 +415,9 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection rows.extend(items); } - if page_in_sql { - Ok(dedupe_candidate_selection_rows( - rows.into_iter().map(|item| item.row).collect(), - )) - } else { - Ok(dedupe_candidate_selection_rows( - sort_pool_key_rows(rows, &query.order) - .into_iter() - .skip(query.offset as usize) - .take(query.limit as usize) - .map(|item| item.row) - .collect(), - )) - } + Ok(dedupe_candidate_selection_rows( + rows.into_iter().map(|item| item.row).collect(), + )) } async fn list_pool_key_rows_for_group_key_ids( @@ -411,6 +506,19 @@ fn push_candidate_sql_filters( push_key_auth_channel_sql_filter(builder, storage_api_format); } +fn push_candidate_sql_filters_for_aliases( + builder: &mut QueryBuilder<'_, Sqlite>, + storage_api_formats: &[String], + match_aliases: &[String], + requested_api_format: &str, +) { + builder.push(" AND LOWER(COALESCE(pe.api_format, '')) IN ("); + push_bind_list(builder, storage_api_formats); + builder.push(")"); + push_key_api_format_sql_filter(builder, match_aliases); + push_key_auth_channel_sql_filter(builder, requested_api_format); +} + fn push_key_api_format_sql_filter( builder: &mut QueryBuilder<'_, Sqlite>, match_aliases: &[String], @@ -626,6 +734,51 @@ fn push_requested_model_sql_filter( ); } +fn push_selected_rows_query_tail( + builder: &mut QueryBuilder<'_, Sqlite>, + order: SelectedRowsOrder, + page: Option, +) { + builder.push( + r#" +), +pool_rows AS ( + SELECT candidate.* + FROM candidate_rows candidate + WHERE candidate.provider_pool_enabled = 1 + AND NOT EXISTS ( + SELECT 1 + FROM candidate_rows other + WHERE other.provider_pool_enabled = 1 + AND other.provider_id = candidate.provider_id + AND other.endpoint_id = candidate.endpoint_id + AND other.model_id = candidate.model_id + AND ( + other.key_internal_priority < candidate.key_internal_priority + OR ( + other.key_internal_priority = candidate.key_internal_priority + AND other.key_id < candidate.key_id + ) + ) + ) +), +selected_rows AS ( + SELECT * FROM candidate_rows WHERE provider_pool_enabled = 0 + UNION ALL + SELECT * FROM pool_rows +) +SELECT * FROM selected_rows +"#, + ); + push_selected_rows_order(builder, order); + if let Some(page) = page { + builder.push(" LIMIT "); + builder.push_bind(page.limit); + builder.push(" OFFSET "); + builder.push_bind(page.offset); + } +} + fn push_selected_rows_order(builder: &mut QueryBuilder<'_, Sqlite>, order: SelectedRowsOrder) { builder.push(" ORDER BY "); if matches!(order, SelectedRowsOrder::WithGlobalModel) { @@ -660,12 +813,52 @@ fn push_pool_key_order( ); } StoredPoolKeyCandidateOrder::LoadBalance { seed } => { - let _ = seed; - builder.push(" ORDER BY pak.id ASC"); + push_seeded_pool_key_order(builder, seed); } } } +fn push_seeded_pool_key_order(builder: &mut QueryBuilder<'_, Sqlite>, seed: &str) { + const HEX_DIGITS: &[u8; 16] = b"0123456789abcdef"; + const PLACEHOLDERS: &[u8; 16] = b"ghijklmnopqrstuv"; + + let mut digits_by_rank = HEX_DIGITS.map(char::from); + digits_by_rank.sort_by(|left, right| { + stable_pool_key_hash(seed, &left.to_string()) + .cmp(&stable_pool_key_hash(seed, &right.to_string())) + .then(left.cmp(right)) + }); + let mut rank_by_digit = ['0'; 16]; + for (rank, digit) in digits_by_rank.into_iter().enumerate() { + let digit_index = digit + .to_digit(16) + .expect("seeded pool-key rank input must be hexadecimal") + as usize; + rank_by_digit[digit_index] = char::from(HEX_DIGITS[rank]); + } + + // SQLite has no built-in hash; placeholders avoid cascading replacements + // while remapping every key-id nibble to a seed-derived rank. + builder.push(" ORDER BY "); + for _ in 0..(HEX_DIGITS.len() + PLACEHOLDERS.len()) { + builder.push("replace("); + } + builder.push("lower(hex(pak.id))"); + for (digit, placeholder) in HEX_DIGITS.iter().zip(PLACEHOLDERS) { + builder.push(format!( + ", '{}', '{}')", + char::from(*digit), + char::from(*placeholder) + )); + } + for (placeholder, rank) in PLACEHOLDERS.iter().zip(rank_by_digit) { + builder.push(format!(", '{}', ", char::from(*placeholder))); + builder.push_bind(rank.to_string()); + builder.push(")"); + } + builder.push(" ASC, pak.id ASC"); +} + fn push_bind_list(builder: &mut QueryBuilder<'_, Sqlite>, values: &[String]) { let mut separated = builder.separated(", "); for value in values { @@ -673,52 +866,6 @@ fn push_bind_list(builder: &mut QueryBuilder<'_, Sqlite>, values: &[String]) { } } -fn sort_pool_key_rows( - mut rows: Vec, - order: &StoredPoolKeyCandidateOrder, -) -> Vec { - rows.sort_by(|left, right| match order { - StoredPoolKeyCandidateOrder::InternalPriority => compare_pool_key_internal(left, right), - StoredPoolKeyCandidateOrder::Lru => left - .key_last_used_at_unix_secs - .cmp(&right.key_last_used_at_unix_secs) - .then_with(|| compare_pool_key_internal(left, right)), - StoredPoolKeyCandidateOrder::CacheAffinity => right - .key_last_used_at_unix_secs - .cmp(&left.key_last_used_at_unix_secs) - .then_with(|| compare_pool_key_internal(left, right)), - StoredPoolKeyCandidateOrder::SingleAccount => left - .row - .key_internal_priority - .cmp(&right.row.key_internal_priority) - .then_with(|| { - right - .key_last_used_at_unix_secs - .cmp(&left.key_last_used_at_unix_secs) - }) - .then(left.row.key_id.cmp(&right.row.key_id)), - StoredPoolKeyCandidateOrder::LoadBalance { seed } => { - stable_pool_key_hash(seed.as_str(), left.row.key_id.as_str()) - .cmp(&stable_pool_key_hash( - seed.as_str(), - right.row.key_id.as_str(), - )) - .then(left.row.key_id.cmp(&right.row.key_id)) - } - }); - rows -} - -fn compare_pool_key_internal( - left: &CandidateSelectionRow, - right: &CandidateSelectionRow, -) -> std::cmp::Ordering { - left.row - .key_internal_priority - .cmp(&right.row.key_internal_priority) - .then(left.row.key_id.cmp(&right.row.key_id)) -} - fn stable_pool_key_hash(seed: &str, key_id: &str) -> u64 { let mut hash = 0xcbf29ce484222325u64; for byte in seed @@ -875,6 +1022,24 @@ fn dedupe_candidate_selection_rows( .collect() } +fn sort_candidate_selection_rows( + rows: &mut [StoredMinimalCandidateSelectionRow], + include_global_model: bool, +) { + rows.sort_by(|left, right| { + let global_model_order = include_global_model + .then(|| left.global_model_name.cmp(&right.global_model_name)) + .unwrap_or(std::cmp::Ordering::Equal); + global_model_order + .then(left.provider_priority.cmp(&right.provider_priority)) + .then(left.key_internal_priority.cmp(&right.key_internal_priority)) + .then(left.provider_id.cmp(&right.provider_id)) + .then(left.endpoint_id.cmp(&right.endpoint_id)) + .then(left.key_id.cmp(&right.key_id)) + .then(left.model_id.cmp(&right.model_id)) + }); +} + fn map_candidate_selection_row(row: &SqliteRow) -> Result { let _provider_config = parse_json(row.try_get("provider_config").ok().flatten())?; let global_model_config = parse_json(row.try_get("global_model_config").ok().flatten())?; @@ -931,10 +1096,6 @@ fn map_candidate_selection_row(row: &SqliteRow) -> Result, _>("key_last_used_at_unix_secs") - .map_sql_err()? - .and_then(|value| u64::try_from(value).ok()), }) } @@ -1177,8 +1338,9 @@ fn sql_match_aliases(api_formats: &[String]) -> Vec { #[cfg(test)] mod tests { use super::{ - push_key_auth_channel_sql_filter, vertex_key_auth_channel_matches, - SqliteMinimalCandidateSelectionReadRepository, + push_key_auth_channel_sql_filter, push_pool_key_order, vertex_key_auth_channel_matches, + ExactPageAccumulator, SqliteMinimalCandidateSelectionReadRepository, + REQUESTED_MODEL_RAW_SCAN_LIMIT, }; use crate::run_migrations; use aether_data_contracts::repository::candidate_selection::{ @@ -1216,6 +1378,57 @@ mod tests { assert!(vertex_clause.contains("gemini:embedding")); } + #[test] + fn exact_page_accumulator_continues_after_coarse_false_positives() { + let mut accumulator = ExactPageAccumulator::new(1, 2); + accumulator.push_matching(vec![("coarse-1", false), ("coarse-2", false)], |row| row.1); + assert!(!accumulator.is_full()); + + accumulator.push_matching( + vec![ + ("exact-1", true), + ("coarse-3", false), + ("exact-2", true), + ("exact-3", true), + ], + |row| row.1, + ); + + assert!(accumulator.is_full()); + assert_eq!( + accumulator.into_page(), + vec![("exact-2", true), ("exact-3", true)] + ); + assert_eq!(REQUESTED_MODEL_RAW_SCAN_LIMIT, 2048); + } + + #[test] + fn load_balance_pool_key_order_is_seeded_and_pageable_in_sql() { + let sql_for_seed = |seed: &str| { + let mut builder = + sqlx::QueryBuilder::::new("SELECT pak.id FROM provider_api_keys pak"); + push_pool_key_order( + &mut builder, + &StoredPoolKeyCandidateOrder::LoadBalance { + seed: seed.to_string(), + }, + ); + builder.push(" LIMIT "); + builder.push_bind(64_i64); + builder.push(" OFFSET "); + builder.push_bind(128_i64); + builder.sql().to_string() + }; + + let first_seed_sql = sql_for_seed("seed-a"); + let second_seed_sql = sql_for_seed("seed-b"); + + assert!(first_seed_sql.contains("lower(hex(pak.id))")); + assert!(first_seed_sql.contains("ASC, pak.id ASC LIMIT ? OFFSET ?")); + assert_eq!(first_seed_sql, second_seed_sql); + assert_eq!(first_seed_sql.matches('?').count(), 18); + } + #[tokio::test] async fn sqlite_repository_reads_candidate_selection_rows() { let pool = sqlx::sqlite::SqlitePoolOptions::new() @@ -1336,6 +1549,97 @@ mod tests { assert_eq!(search_rows[0].endpoint_api_format, "openai:search"); } + #[tokio::test] + async fn sqlite_requested_model_page_crosses_coarse_false_positive_windows() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + seed_requested_model_pagination(&pool).await; + + let repository = SqliteMinimalCandidateSelectionReadRepository::new(pool); + let rows = repository + .list_for_exact_api_format_and_requested_model_page( + &StoredRequestedModelCandidateRowsQuery { + api_format: "openai:chat".to_string(), + requested_model_name: "sqlite-page-target".to_string(), + offset: 1, + limit: 1, + }, + ) + .await + .expect("requested model page should cross the coarse-only window"); + + assert_eq!(rows.len(), 1); + assert_eq!(rows[0].model_id, "model-pagination-exact-1"); + } + + #[tokio::test] + async fn sqlite_load_balance_pool_key_pages_use_stable_seeded_order() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_migrations(&pool) + .await + .expect("sqlite migrations should run"); + seed_candidate_selection(&pool).await; + + let repository = SqliteMinimalCandidateSelectionReadRepository::new(pool); + let load_page = |seed: &str, offset, limit| StoredPoolKeyCandidateRowsQuery { + api_format: "openai:chat".to_string(), + provider_id: "provider-1".to_string(), + endpoint_id: "endpoint-1".to_string(), + model_id: "model-1".to_string(), + selected_provider_model_name: "provider-model".to_string(), + order: StoredPoolKeyCandidateOrder::LoadBalance { + seed: seed.to_string(), + }, + offset, + limit, + }; + + let seed_a_first = repository + .list_pool_key_rows_for_group(&load_page("seed-a", 0, 1)) + .await + .expect("first load-balance page should load"); + let seed_a_second = repository + .list_pool_key_rows_for_group(&load_page("seed-a", 1, 1)) + .await + .expect("second load-balance page should load"); + let seed_a_replay = repository + .list_pool_key_rows_for_group(&load_page("seed-a", 0, 2)) + .await + .expect("replayed load-balance window should load"); + let seed_b = repository + .list_pool_key_rows_for_group(&load_page("seed-b", 0, 2)) + .await + .expect("alternate load-balance seed should load"); + + let seed_a_pages = seed_a_first + .iter() + .chain(&seed_a_second) + .map(|row| row.key_id.as_str()) + .collect::>(); + let seed_a_replay = seed_a_replay + .iter() + .map(|row| row.key_id.as_str()) + .collect::>(); + let seed_b = seed_b + .iter() + .map(|row| row.key_id.as_str()) + .collect::>(); + + assert_eq!(seed_a_pages, vec!["key-1", "key-2"]); + assert_eq!(seed_a_pages, seed_a_replay); + assert_eq!(seed_b, vec!["key-2", "key-1"]); + } + async fn seed_candidate_selection(pool: &sqlx::SqlitePool) { sqlx::query( r#" @@ -1455,4 +1759,94 @@ VALUES ( .await .expect("candidate selection rows should seed"); } + + async fn seed_requested_model_pagination(pool: &sqlx::SqlitePool) { + sqlx::query( + r#" +INSERT INTO providers ( + id, name, provider_type, provider_priority, is_active, created_at, updated_at +) +VALUES ( + 'provider-pagination', 'Pagination Provider', 'custom', 10, 1, 1, 1 +); + +INSERT INTO provider_endpoints ( + id, provider_id, name, base_url, api_format, is_active, created_at, updated_at +) +VALUES ( + 'endpoint-pagination', 'provider-pagination', 'Pagination Endpoint', + 'https://example.test', 'openai:chat', 1, 1, 1 +); + +INSERT INTO provider_api_keys ( + id, provider_id, name, auth_type, api_formats, internal_priority, + is_active, created_at, updated_at +) +VALUES ( + 'key-pagination', 'provider-pagination', 'Pagination Key', 'api_key', + '["openai:chat"]', 10, 1, 1, 1 +); + +WITH RECURSIVE sequence(value) AS ( + SELECT 0 + UNION ALL + SELECT value + 1 FROM sequence WHERE value < 255 +) +INSERT INTO global_models ( + id, name, is_active, created_at, updated_at +) +SELECT + printf('global-pagination-false-%03d', value), + printf('a-pagination-false-%03d', value), + 1, 1, 1 +FROM sequence; + +INSERT INTO global_models ( + id, name, is_active, created_at, updated_at +) +VALUES + ('global-pagination-exact-0', 'z-pagination-exact-0', 1, 1, 1), + ('global-pagination-exact-1', 'z-pagination-exact-1', 1, 1, 1); + +WITH RECURSIVE sequence(value) AS ( + SELECT 0 + UNION ALL + SELECT value + 1 FROM sequence WHERE value < 255 +) +INSERT INTO models ( + id, provider_id, global_model_id, provider_model_name, provider_model_mappings, + is_active, is_available, created_at, updated_at +) +SELECT + printf('model-pagination-false-%03d', value), + 'provider-pagination', + printf('global-pagination-false-%03d', value), + 'upstream-false', + '[{"name":"sqlite-page-target-noise","api_formats":["openai:chat"],"priority":1}]', + 1, 1, 1, 1 +FROM sequence; + +INSERT INTO models ( + id, provider_id, global_model_id, provider_model_name, provider_model_mappings, + is_active, is_available, created_at, updated_at +) +VALUES + ( + 'model-pagination-exact-0', 'provider-pagination', 'global-pagination-exact-0', + 'upstream-exact-0', + '[{"name":"sqlite-page-target","api_formats":["openai:chat"],"priority":1}]', + 1, 1, 1, 1 + ), + ( + 'model-pagination-exact-1', 'provider-pagination', 'global-pagination-exact-1', + 'upstream-exact-1', + '[{"name":"sqlite-page-target","api_formats":["openai:chat"],"priority":1}]', + 1, 1, 1, 1 + ); +"#, + ) + .execute(pool) + .await + .expect("requested model pagination rows should seed"); + } } diff --git a/crates/aether-data/adapters/sqlite/src/routing_profiles.rs b/crates/aether-data/adapters/sqlite/src/routing_profiles.rs index 3c44f612d..22bcdf8ac 100644 --- a/crates/aether-data/adapters/sqlite/src/routing_profiles.rs +++ b/crates/aether-data/adapters/sqlite/src/routing_profiles.rs @@ -56,24 +56,6 @@ impl SqliteRoutingGroupRepository { pub fn new(pool: SqlitePool) -> Self { Self { pool } } - - async fn reload_group(&self, id: &str) -> Result, DataLayerError> { - self.find_routing_group(RoutingGroupLookupKey::Id(id)).await - } - - async fn find_binding_by_id( - &self, - id: &str, - ) -> Result, DataLayerError> { - let row = sqlx::query(&format!( - "{ROUTING_GROUP_BINDING_SELECT} WHERE id = ? LIMIT 1" - )) - .bind(id) - .fetch_optional(&self.pool) - .await - .map_sql_err()?; - row.as_ref().map(map_binding_row).transpose() - } } #[async_trait] @@ -170,6 +152,19 @@ impl RoutingGroupWriteRepository for SqliteRoutingGroupRepository { record: CreateRoutingGroupRecord, ) -> Result { let group = StoredRoutingGroup::new(record)?; + let mut tx = self.pool.begin().await.map_sql_err()?; + sqlx::query("UPDATE routing_groups SET is_system_default = is_system_default WHERE 0") + .execute(&mut *tx) + .await + .map_sql_err()?; + if group.is_system_default { + sqlx::query( + "UPDATE routing_groups SET is_system_default = 0 WHERE is_system_default = 1", + ) + .execute(&mut *tx) + .await + .map_sql_err()?; + } sqlx::query( r#" INSERT INTO routing_groups ( @@ -192,9 +187,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) .bind(group.created_at) .bind(group.updated_at) .bind(group.published_at) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; + tx.commit().await.map_sql_err()?; Ok(group) } @@ -203,10 +199,29 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) id: &str, patch: UpdateRoutingGroupRecord, ) -> Result, DataLayerError> { - let Some(mut group) = self.reload_group(id).await? else { + let mut tx = self.pool.begin().await.map_sql_err()?; + sqlx::query("UPDATE routing_groups SET is_system_default = is_system_default WHERE 0") + .execute(&mut *tx) + .await + .map_sql_err()?; + let row = sqlx::query(&format!("{ROUTING_GROUP_SELECT} WHERE id = ? LIMIT 1")) + .bind(id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(mut group) = row.as_ref().map(map_group_row).transpose()? else { return Ok(None); }; apply_group_patch(&mut group, patch)?; + if group.is_system_default { + sqlx::query( + "UPDATE routing_groups SET is_system_default = 0 WHERE is_system_default = 1 AND id <> ?", + ) + .bind(id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } sqlx::query( r#" UPDATE routing_groups @@ -233,9 +248,10 @@ WHERE id = ? .bind(group.updated_at) .bind(group.published_at) .bind(id) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; + tx.commit().await.map_sql_err()?; Ok(Some(group)) } @@ -266,6 +282,25 @@ WHERE id = ? record: CreateRoutingGroupBindingRecord, ) -> Result { let binding = StoredRoutingGroupBinding::new(record)?; + let mut tx = self.pool.begin().await.map_sql_err()?; + sqlx::query("UPDATE routing_group_bindings SET is_default = is_default WHERE 0") + .execute(&mut *tx) + .await + .map_sql_err()?; + if binding.is_default { + sqlx::query( + r#" +UPDATE routing_group_bindings +SET is_default = 0 +WHERE is_default = 1 AND subject_type = ? AND subject_id = ? +"#, + ) + .bind(binding_subject_to_database(binding.subject_type)) + .bind(&binding.subject_id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } sqlx::query( r#" INSERT INTO routing_group_bindings ( @@ -283,9 +318,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?) .bind(binding.allow_explicit_select) .bind(binding.created_at) .bind(binding.updated_at) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; + tx.commit().await.map_sql_err()?; Ok(binding) } @@ -306,10 +342,40 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?) id: &str, patch: UpdateRoutingGroupBindingRecord, ) -> Result, DataLayerError> { - let Some(mut binding) = self.find_binding_by_id(id).await? else { + let mut tx = self.pool.begin().await.map_sql_err()?; + sqlx::query("UPDATE routing_group_bindings SET is_default = is_default WHERE 0") + .execute(&mut *tx) + .await + .map_sql_err()?; + let row = sqlx::query(&format!( + "{ROUTING_GROUP_BINDING_SELECT} WHERE id = ? LIMIT 1" + )) + .bind(id) + .fetch_optional(&mut *tx) + .await + .map_sql_err()?; + let Some(mut binding) = row.as_ref().map(map_binding_row).transpose()? else { return Ok(None); }; apply_binding_patch(&mut binding, patch)?; + if binding.is_default { + sqlx::query( + r#" +UPDATE routing_group_bindings +SET is_default = 0 +WHERE is_default = 1 + AND subject_type = ? + AND subject_id = ? + AND id <> ? +"#, + ) + .bind(binding_subject_to_database(binding.subject_type)) + .bind(&binding.subject_id) + .bind(id) + .execute(&mut *tx) + .await + .map_sql_err()?; + } sqlx::query( r#" UPDATE routing_group_bindings @@ -329,9 +395,10 @@ WHERE id = ? .bind(binding.allow_explicit_select) .bind(binding.updated_at) .bind(id) - .execute(&self.pool) + .execute(&mut *tx) .await .map_sql_err()?; + tx.commit().await.map_sql_err()?; Ok(Some(binding)) } @@ -526,4 +593,255 @@ mod tests { 1 ); } + + #[tokio::test] + async fn sqlite_keeps_system_and_subject_defaults_unique() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_sqlite_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteRoutingGroupRepository::new(pool); + + for (id, is_system_default) in [("group-1", true), ("group-2", true), ("group-3", false)] { + repository + .create_routing_group(group_record(id, is_system_default)) + .await + .expect("group should create"); + } + assert_eq!(system_default_ids(&repository).await, vec!["group-2"]); + + repository + .update_routing_group( + "group-1", + UpdateRoutingGroupRecord { + is_system_default: Some(true), + updated_at: 2, + ..UpdateRoutingGroupRecord::default() + }, + ) + .await + .expect("group should update"); + assert_eq!(system_default_ids(&repository).await, vec!["group-1"]); + + repository + .create_routing_group_binding(binding_record("binding-1", "group-1", "subject-1", true)) + .await + .expect("binding should create"); + repository + .create_routing_group_binding(binding_record("binding-2", "group-2", "subject-1", true)) + .await + .expect("binding should create"); + repository + .create_routing_group_binding(binding_record("binding-3", "group-3", "subject-2", true)) + .await + .expect("binding should create"); + + assert_eq!( + default_binding_ids(&repository, "subject-1").await, + vec!["binding-2"] + ); + assert_eq!( + default_binding_ids(&repository, "subject-2").await, + vec!["binding-3"] + ); + + repository + .update_routing_group_binding( + "binding-1", + UpdateRoutingGroupBindingRecord { + is_default: Some(true), + updated_at: 2, + ..UpdateRoutingGroupBindingRecord::default() + }, + ) + .await + .expect("binding should update"); + assert_eq!( + default_binding_ids(&repository, "subject-1").await, + vec!["binding-1"] + ); + assert_eq!( + default_binding_ids(&repository, "subject-2").await, + vec!["binding-3"] + ); + + repository + .update_routing_group_binding( + "binding-3", + UpdateRoutingGroupBindingRecord { + subject_id: Some("subject-1".to_string()), + updated_at: 3, + ..UpdateRoutingGroupBindingRecord::default() + }, + ) + .await + .expect("binding should move"); + assert_eq!( + default_binding_ids(&repository, "subject-1").await, + vec!["binding-3"] + ); + assert!(default_binding_ids(&repository, "subject-2") + .await + .is_empty()); + } + + #[tokio::test] + async fn sqlite_repair_migration_resolves_existing_duplicate_defaults() { + let pool = sqlx::sqlite::SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("sqlite pool should connect"); + run_sqlite_migrations(&pool) + .await + .expect("sqlite migrations should run"); + let repository = SqliteRoutingGroupRepository::new(pool.clone()); + sqlx::raw_sql( + r#" +DROP INDEX routing_groups_one_system_default_key; +DROP INDEX routing_group_bindings_subject_default_key; +"#, + ) + .execute(&pool) + .await + .expect("unique indexes should be removable to simulate a pre-repair database"); + + for id in ["group-1", "group-2", "group-3"] { + repository + .create_routing_group(group_record(id, false)) + .await + .expect("group should create"); + } + sqlx::query( + r#" +UPDATE routing_groups +SET is_system_default = 1, + enabled = CASE id WHEN 'group-3' THEN 0 ELSE 1 END, + updated_at = CASE id + WHEN 'group-1' THEN 1 + WHEN 'group-2' THEN 2 + ELSE 3 + END +"#, + ) + .execute(&pool) + .await + .expect("duplicate system defaults should seed"); + + for (id, subject_id) in [ + ("binding-3", "subject-1"), + ("binding-2", "subject-1"), + ("binding-1", "subject-1"), + ("binding-4", "subject-2"), + ] { + repository + .create_routing_group_binding(binding_record(id, "group-1", subject_id, false)) + .await + .expect("binding should create"); + } + sqlx::query( + r#" +UPDATE routing_group_bindings +SET is_default = 1, + created_at = CASE id WHEN 'binding-3' THEN 2 ELSE 1 END +"#, + ) + .execute(&pool) + .await + .expect("duplicate binding defaults should seed"); + + let repair_migration = + include_str!("../migrations/20260727000000_repair_routing_default_uniqueness.sql"); + for _ in 0..2 { + sqlx::raw_sql(repair_migration) + .execute(&pool) + .await + .expect("repair migration should be idempotent"); + } + + assert_eq!(system_default_ids(&repository).await, vec!["group-2"]); + assert_eq!( + default_binding_ids(&repository, "subject-1").await, + vec!["binding-1"] + ); + assert_eq!( + default_binding_ids(&repository, "subject-2").await, + vec!["binding-4"] + ); + + sqlx::query("UPDATE routing_groups SET is_system_default = 1 WHERE id = 'group-3'") + .execute(&pool) + .await + .expect_err("database should reject a second system default"); + sqlx::query("UPDATE routing_group_bindings SET is_default = 1 WHERE id = 'binding-2'") + .execute(&pool) + .await + .expect_err("database should reject a second default for the same subject"); + } + + fn group_record(id: &str, is_system_default: bool) -> CreateRoutingGroupRecord { + CreateRoutingGroupRecord { + id: id.to_string(), + name: id.to_string(), + description: None, + enabled: true, + is_system_default, + config_json: json!({}), + version: 1, + created_at: 1, + updated_at: 1, + published_at: None, + } + } + + fn binding_record( + id: &str, + group_id: &str, + subject_id: &str, + is_default: bool, + ) -> CreateRoutingGroupBindingRecord { + CreateRoutingGroupBindingRecord { + id: id.to_string(), + group_id: group_id.to_string(), + subject_type: RoutingGroupBindingSubject::ApiKey, + subject_id: subject_id.to_string(), + is_default, + allow_explicit_select: true, + created_at: 1, + updated_at: 1, + } + } + + async fn system_default_ids(repository: &SqliteRoutingGroupRepository) -> Vec { + repository + .list_routing_groups() + .await + .expect("groups should list") + .into_iter() + .filter(|group| group.is_system_default) + .map(|group| group.id) + .collect() + } + + async fn default_binding_ids( + repository: &SqliteRoutingGroupRepository, + subject_id: &str, + ) -> Vec { + repository + .list_routing_group_bindings(&RoutingGroupBindingQuery { + group_id: None, + subject_type: Some(RoutingGroupBindingSubject::ApiKey), + subject_id: Some(subject_id.to_string()), + }) + .await + .expect("bindings should list") + .into_iter() + .filter(|binding| binding.is_default) + .map(|binding| binding.id) + .collect() + } } diff --git a/crates/aether-data/contracts/src/repository/candidate_selection/mod.rs b/crates/aether-data/contracts/src/repository/candidate_selection/mod.rs index a1a1619a7..bc46fb39d 100644 --- a/crates/aether-data/contracts/src/repository/candidate_selection/mod.rs +++ b/crates/aether-data/contracts/src/repository/candidate_selection/mod.rs @@ -2,7 +2,8 @@ mod types; pub use types::{ MinimalCandidateSelectionReadRepository, MinimalCandidateSelectionRepository, - StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder, - StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery, - StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery, + StoredApiFormatCandidateRowsQuery, StoredMinimalCandidateSelectionRow, + StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery, + StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping, + StoredRequestedModelCandidateRowsQuery, }; diff --git a/crates/aether-data/contracts/src/repository/candidate_selection/types.rs b/crates/aether-data/contracts/src/repository/candidate_selection/types.rs index 46f160b6c..48e74b3db 100644 --- a/crates/aether-data/contracts/src/repository/candidate_selection/types.rs +++ b/crates/aether-data/contracts/src/repository/candidate_selection/types.rs @@ -89,6 +89,13 @@ pub struct StoredRequestedModelCandidateRowsQuery { pub limit: u32, } +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +pub struct StoredApiFormatCandidateRowsQuery { + pub api_format: String, + pub offset: u32, + pub limit: u32, +} + impl StoredMinimalCandidateSelectionRow { pub fn supports_streaming(&self) -> bool { self.model_supports_streaming @@ -119,6 +126,19 @@ pub trait MinimalCandidateSelectionReadRepository: Send + Sync { api_format: &str, ) -> Result, crate::DataLayerError>; + async fn list_for_exact_api_format_page( + &self, + query: &StoredApiFormatCandidateRowsQuery, + ) -> Result, crate::DataLayerError> { + Ok(self + .list_for_exact_api_format(&query.api_format) + .await? + .into_iter() + .skip(query.offset as usize) + .take(query.limit as usize) + .collect()) + } + async fn list_for_exact_api_format_and_global_model( &self, api_format: &str, diff --git a/crates/aether-data/contracts/src/repository/pool_scores/helpers.rs b/crates/aether-data/contracts/src/repository/pool_scores/helpers.rs index 2a3d9a0be..c7579ce42 100644 --- a/crates/aether-data/contracts/src/repository/pool_scores/helpers.rs +++ b/crates/aether-data/contracts/src/repository/pool_scores/helpers.rs @@ -17,7 +17,10 @@ pub fn merge_score_reason_patch( } pub fn score_with_delta(score: f64, delta_basis_points: Option) -> f64 { - let delta = delta_basis_points.unwrap_or_default() as f64 / 10_000.0; + let Some(delta_basis_points) = delta_basis_points else { + return score; + }; + let delta = delta_basis_points as f64 / 10_000.0; (score + delta).clamp(0.0, 1.0) } @@ -46,3 +49,19 @@ pub fn u64_opt_from_i64( ) -> Result, crate::DataLayerError> { value.map(|value| u64_from_i64(value, field)).transpose() } + +#[cfg(test)] +mod tests { + use super::score_with_delta; + + #[test] + fn score_without_delta_is_unchanged() { + assert_eq!(score_with_delta(1_000.0, None), 1_000.0); + } + + #[test] + fn score_with_delta_remains_bounded() { + assert_eq!(score_with_delta(0.95, Some(1_000)), 1.0); + assert_eq!(score_with_delta(0.05, Some(-1_000)), 0.0); + } +} diff --git a/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs b/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs index 2906bf693..94f6262c1 100644 --- a/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs +++ b/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs @@ -3,9 +3,10 @@ use std::sync::RwLock; use async_trait::async_trait; use super::{ - MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow, - StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery, - StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery, + MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery, + StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder, + StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery, + StoredRequestedModelCandidateRowsQuery, }; use crate::DataLayerError; @@ -61,6 +62,19 @@ impl MinimalCandidateSelectionReadRepository for InMemoryMinimalCandidateSelecti Ok(rows) } + async fn list_for_exact_api_format_page( + &self, + query: &StoredApiFormatCandidateRowsQuery, + ) -> Result, DataLayerError> { + Ok(self + .list_for_exact_api_format(&query.api_format) + .await? + .into_iter() + .skip(query.offset as usize) + .take(query.limit as usize) + .collect()) + } + async fn list_for_exact_api_format_and_global_model( &self, api_format: &str, @@ -351,8 +365,9 @@ fn key_auth_channel_matches(row: &StoredMinimalCandidateSelectionRow, api_format mod tests { use super::InMemoryMinimalCandidateSelectionReadRepository; use crate::repository::candidate_selection::{ - MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow, - StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping, + MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery, + StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder, + StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery, }; @@ -697,6 +712,27 @@ mod tests { assert_eq!(rows[1].provider_id, "provider-2"); } + #[tokio::test] + async fn lists_exact_api_format_in_stable_pages() { + let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![ + sample_row("provider-3", "openai:chat", "gpt-4.1", 30), + sample_row("provider-1", "openai:chat", "gpt-4.1", 10), + sample_row("provider-2", "openai:chat", "gpt-4.1", 20), + ]); + + let page = repository + .list_for_exact_api_format_page(&StoredApiFormatCandidateRowsQuery { + api_format: "openai:chat".to_string(), + offset: 1, + limit: 1, + }) + .await + .expect("API-format page should load"); + + assert_eq!(page.len(), 1); + assert_eq!(page[0].provider_id, "provider-2"); + } + #[tokio::test] async fn list_pool_key_rows_for_group_returns_requested_page_only() { let mut rows = Vec::new(); diff --git a/crates/aether-data/runtime/src/repository/candidate_selection/mod.rs b/crates/aether-data/runtime/src/repository/candidate_selection/mod.rs index 3671ffeff..68428a5c9 100644 --- a/crates/aether-data/runtime/src/repository/candidate_selection/mod.rs +++ b/crates/aether-data/runtime/src/repository/candidate_selection/mod.rs @@ -3,9 +3,10 @@ mod memory; #[allow(unused_imports)] pub(crate) use aether_data_contracts::repository::candidate_selection::{ MinimalCandidateSelectionReadRepository, MinimalCandidateSelectionRepository, - StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder, - StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery, - StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery, + StoredApiFormatCandidateRowsQuery, StoredMinimalCandidateSelectionRow, + StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery, + StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping, + StoredRequestedModelCandidateRowsQuery, }; #[cfg(feature = "mysql")] pub use aether_data_mysql::MysqlMinimalCandidateSelectionReadRepository; diff --git a/crates/aether-data/runtime/src/repository/provider_oauth.rs b/crates/aether-data/runtime/src/repository/provider_oauth.rs index 456a4a0fa..a7a5abf6b 100644 --- a/crates/aether-data/runtime/src/repository/provider_oauth.rs +++ b/crates/aether-data/runtime/src/repository/provider_oauth.rs @@ -72,6 +72,13 @@ pub fn build_provider_oauth_batch_task_status_payload( "submitted" | "processing" | "completed" | "failed" => raw_status, _ => "failed", }; + let import_kind = match state.get("import_kind").and_then(serde_json::Value::as_str) { + Some("oauth_batch" | "agent_identity" | "cookie_authorize") => state + .get("import_kind") + .and_then(serde_json::Value::as_str) + .unwrap_or_default(), + _ => "", + }; let error_samples = state .get("error_samples") .and_then(serde_json::Value::as_array) @@ -94,6 +101,7 @@ pub fn build_provider_oauth_batch_task_status_payload( .get("provider_type") .and_then(serde_json::Value::as_str) .unwrap_or_default(), + "import_kind": import_kind, "status": normalized_status, "total": state.get("total").and_then(serde_json::Value::as_i64).unwrap_or(0), "processed": state.get("processed").and_then(serde_json::Value::as_i64).unwrap_or(0), @@ -169,6 +177,7 @@ mod tests { let input = json!({ "task_id": "task-123", "provider_type": "codex", + "import_kind": "oauth_batch", "status": "weird", "total": 4, "processed": 2, @@ -194,6 +203,10 @@ mod tests { payload.get("provider_id").and_then(|v| v.as_str()), Some("provider-123") ); + assert_eq!( + payload.get("import_kind").and_then(|v| v.as_str()), + Some("oauth_batch") + ); assert_eq!( payload.get("status").and_then(|v| v.as_str()), Some("failed") diff --git a/crates/aether-data/runtime/src/repository/routing_profiles/memory.rs b/crates/aether-data/runtime/src/repository/routing_profiles/memory.rs index 29a52f9ba..b6f11e675 100644 --- a/crates/aether-data/runtime/src/repository/routing_profiles/memory.rs +++ b/crates/aether-data/runtime/src/repository/routing_profiles/memory.rs @@ -1,6 +1,7 @@ use std::collections::BTreeMap; use std::sync::RwLock; +use aether_data_contracts::repository::routing_profiles::{apply_binding_patch, apply_group_patch}; use async_trait::async_trait; use super::{ @@ -149,10 +150,13 @@ impl RoutingGroupWriteRepository for InMemoryRoutingGroupRepository { record: CreateRoutingGroupRecord, ) -> Result { let group = StoredRoutingGroup::new(record)?; - self.groups - .write() - .expect("routing group repository lock") - .insert(group.id.clone(), group.clone()); + let mut groups = self.groups.write().expect("routing group repository lock"); + if group.is_system_default { + for existing in groups.values_mut() { + existing.is_system_default = false; + } + } + groups.insert(group.id.clone(), group.clone()); Ok(group) } @@ -162,42 +166,17 @@ impl RoutingGroupWriteRepository for InMemoryRoutingGroupRepository { patch: UpdateRoutingGroupRecord, ) -> Result, DataLayerError> { let mut groups = self.groups.write().expect("routing group repository lock"); - let Some(group) = groups.get_mut(id) else { + let Some(mut group) = groups.get(id).cloned() else { return Ok(None); }; - if let Some(name) = patch.name { - if name.trim().is_empty() { - return Err(DataLayerError::InvalidInput( - "routing_groups.name is empty".to_string(), - )); + apply_group_patch(&mut group, patch)?; + if group.is_system_default { + for existing in groups.values_mut() { + existing.is_system_default = false; } - group.name = name; } - if let Some(description) = patch.description { - group.description = description; - } - if let Some(enabled) = patch.enabled { - group.enabled = enabled; - } - if let Some(is_system_default) = patch.is_system_default { - group.is_system_default = is_system_default; - } - if let Some(config_json) = patch.config_json { - if !config_json.is_object() { - return Err(DataLayerError::InvalidInput( - "routing_groups.config_json must be a JSON object".to_string(), - )); - } - group.config_json = config_json; - } - if let Some(version) = patch.version { - group.version = version.max(1); - } - if let Some(published_at) = patch.published_at { - group.published_at = published_at; - } - group.updated_at = patch.updated_at; - Ok(Some(group.clone())) + groups.insert(id.to_string(), group.clone()); + Ok(Some(group)) } async fn delete_routing_group(&self, id: &str) -> Result { @@ -214,10 +193,19 @@ impl RoutingGroupWriteRepository for InMemoryRoutingGroupRepository { record: CreateRoutingGroupBindingRecord, ) -> Result { let binding = StoredRoutingGroupBinding::new(record)?; - self.bindings + let mut bindings = self + .bindings .write() - .expect("routing group binding repository lock") - .insert(binding.id.clone(), binding.clone()); + .expect("routing group binding repository lock"); + if binding.is_default { + for existing in bindings.values_mut().filter(|existing| { + existing.subject_type == binding.subject_type + && existing.subject_id == binding.subject_id + }) { + existing.is_default = false; + } + } + bindings.insert(binding.id.clone(), binding.clone()); Ok(binding) } @@ -239,36 +227,20 @@ impl RoutingGroupWriteRepository for InMemoryRoutingGroupRepository { .bindings .write() .expect("routing group binding repository lock"); - let Some(binding) = bindings.get_mut(id) else { + let Some(mut binding) = bindings.get(id).cloned() else { return Ok(None); }; - if let Some(group_id) = patch.group_id { - if group_id.trim().is_empty() { - return Err(DataLayerError::InvalidInput( - "routing_group_bindings.group_id is empty".to_string(), - )); + apply_binding_patch(&mut binding, patch)?; + if binding.is_default { + for existing in bindings.values_mut().filter(|existing| { + existing.subject_type == binding.subject_type + && existing.subject_id == binding.subject_id + }) { + existing.is_default = false; } - binding.group_id = group_id; } - if let Some(subject_type) = patch.subject_type { - binding.subject_type = subject_type; - } - if let Some(subject_id) = patch.subject_id { - if subject_id.trim().is_empty() { - return Err(DataLayerError::InvalidInput( - "routing_group_bindings.subject_id is empty".to_string(), - )); - } - binding.subject_id = subject_id; - } - if let Some(is_default) = patch.is_default { - binding.is_default = is_default; - } - if let Some(allow_explicit_select) = patch.allow_explicit_select { - binding.allow_explicit_select = allow_explicit_select; - } - binding.updated_at = patch.updated_at; - Ok(Some(binding.clone())) + bindings.insert(id.to_string(), binding.clone()); + Ok(Some(binding)) } async fn create_routing_group_version( @@ -347,4 +319,153 @@ mod tests { 1 ); } + + #[tokio::test] + async fn keeps_system_and_subject_defaults_unique() { + let repository = InMemoryRoutingGroupRepository::default(); + for (id, is_system_default) in [("group-1", true), ("group-2", true), ("group-3", false)] { + repository + .create_routing_group(group_record(id, is_system_default)) + .await + .expect("group should store"); + } + + assert_eq!(system_default_ids(&repository).await, vec!["group-2"]); + + repository + .update_routing_group( + "group-1", + UpdateRoutingGroupRecord { + is_system_default: Some(true), + updated_at: 2, + ..UpdateRoutingGroupRecord::default() + }, + ) + .await + .expect("group should update"); + assert_eq!(system_default_ids(&repository).await, vec!["group-1"]); + + repository + .create_routing_group_binding(binding_record("binding-1", "group-1", "subject-1", true)) + .await + .expect("binding should store"); + repository + .create_routing_group_binding(binding_record("binding-2", "group-2", "subject-1", true)) + .await + .expect("binding should store"); + repository + .create_routing_group_binding(binding_record("binding-3", "group-3", "subject-2", true)) + .await + .expect("binding should store"); + + assert_eq!( + default_binding_ids(&repository, "subject-1").await, + vec!["binding-2"] + ); + assert_eq!( + default_binding_ids(&repository, "subject-2").await, + vec!["binding-3"] + ); + + repository + .update_routing_group_binding( + "binding-1", + UpdateRoutingGroupBindingRecord { + is_default: Some(true), + updated_at: 2, + ..UpdateRoutingGroupBindingRecord::default() + }, + ) + .await + .expect("binding should update"); + assert_eq!( + default_binding_ids(&repository, "subject-1").await, + vec!["binding-1"] + ); + assert_eq!( + default_binding_ids(&repository, "subject-2").await, + vec!["binding-3"] + ); + + repository + .update_routing_group_binding( + "binding-3", + UpdateRoutingGroupBindingRecord { + subject_id: Some("subject-1".to_string()), + updated_at: 3, + ..UpdateRoutingGroupBindingRecord::default() + }, + ) + .await + .expect("binding should move"); + assert_eq!( + default_binding_ids(&repository, "subject-1").await, + vec!["binding-3"] + ); + assert!(default_binding_ids(&repository, "subject-2") + .await + .is_empty()); + } + + fn group_record(id: &str, is_system_default: bool) -> CreateRoutingGroupRecord { + CreateRoutingGroupRecord { + id: id.to_string(), + name: id.to_string(), + description: None, + enabled: true, + is_system_default, + config_json: json!({}), + version: 1, + created_at: 1, + updated_at: 1, + published_at: None, + } + } + + fn binding_record( + id: &str, + group_id: &str, + subject_id: &str, + is_default: bool, + ) -> CreateRoutingGroupBindingRecord { + CreateRoutingGroupBindingRecord { + id: id.to_string(), + group_id: group_id.to_string(), + subject_type: RoutingGroupBindingSubject::ApiKey, + subject_id: subject_id.to_string(), + is_default, + allow_explicit_select: true, + created_at: 1, + updated_at: 1, + } + } + + async fn system_default_ids(repository: &InMemoryRoutingGroupRepository) -> Vec { + repository + .list_routing_groups() + .await + .expect("groups should list") + .into_iter() + .filter(|group| group.is_system_default) + .map(|group| group.id) + .collect() + } + + async fn default_binding_ids( + repository: &InMemoryRoutingGroupRepository, + subject_id: &str, + ) -> Vec { + repository + .list_routing_group_bindings(&RoutingGroupBindingQuery { + group_id: None, + subject_type: Some(RoutingGroupBindingSubject::ApiKey), + subject_id: Some(subject_id.to_string()), + }) + .await + .expect("bindings should list") + .into_iter() + .filter(|binding| binding.is_default) + .map(|binding| binding.id) + .collect() + } } diff --git a/crates/aether-oauth/src/identity/providers/custom_oidc.rs b/crates/aether-oauth/src/identity/providers/custom_oidc.rs index 95a9722ff..d267b3f58 100644 --- a/crates/aether-oauth/src/identity/providers/custom_oidc.rs +++ b/crates/aether-oauth/src/identity/providers/custom_oidc.rs @@ -81,6 +81,7 @@ impl IdentityOAuthProvider for CustomOidcIdentityOAuthProvider { json_body: None, body_bytes: Some(body_bytes), network: ctx.network.clone(), + transport_profile: None, }) .await?; if !(200..300).contains(&response.status_code) { @@ -123,6 +124,7 @@ impl IdentityOAuthProvider for CustomOidcIdentityOAuthProvider { json_body: None, body_bytes: None, network, + transport_profile: None, }) .await?; if !(200..300).contains(&response.status_code) { diff --git a/crates/aether-oauth/src/network/executor.rs b/crates/aether-oauth/src/network/executor.rs index 86e703214..82831a6b3 100644 --- a/crates/aether-oauth/src/network/executor.rs +++ b/crates/aether-oauth/src/network/executor.rs @@ -1,11 +1,12 @@ use crate::core::OAuthError; +use aether_contracts::ResolvedTransportProfile; use async_trait::async_trait; use serde_json::Value; use std::collections::BTreeMap; use super::OAuthNetworkContext; -#[derive(Debug, Clone, PartialEq)] +#[derive(Clone, PartialEq)] pub struct OAuthHttpRequest { pub request_id: String, pub method: reqwest::Method, @@ -15,6 +16,31 @@ pub struct OAuthHttpRequest { pub json_body: Option, pub body_bytes: Option>, pub network: OAuthNetworkContext, + pub transport_profile: Option, +} + +impl std::fmt::Debug for OAuthHttpRequest { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("OAuthHttpRequest") + .field("request_id", &self.request_id) + .field("method", &self.method) + .field("url", &self.url) + .field("header_names", &self.headers.keys().collect::>()) + .field("content_type", &self.content_type) + .field("has_json_body", &self.json_body.is_some()) + .field("body_bytes_len", &self.body_bytes.as_ref().map(Vec::len)) + .field("network_policy", &self.network.policy) + .field("has_proxy", &self.network.proxy.is_some()) + .field( + "transport_profile_id", + &self + .transport_profile + .as_ref() + .map(|profile| profile.profile_id.as_str()), + ) + .finish() + } } #[derive(Debug, Clone, PartialEq)] diff --git a/crates/aether-oauth/src/provider/account.rs b/crates/aether-oauth/src/provider/account.rs index b783837f8..007a914f1 100644 --- a/crates/aether-oauth/src/provider/account.rs +++ b/crates/aether-oauth/src/provider/account.rs @@ -6,6 +6,7 @@ use std::collections::BTreeMap; #[derive(Debug, Clone, PartialEq, Eq)] pub struct ProviderOAuthCapabilities { pub supports_authorization_code: bool, + pub supports_cookie_authorization: bool, pub supports_refresh_token_import: bool, pub supports_batch_import: bool, pub supports_device_flow: bool, @@ -16,6 +17,7 @@ pub struct ProviderOAuthCapabilities { impl ProviderOAuthCapabilities { pub const GENERIC_AUTH_CODE: Self = Self { supports_authorization_code: true, + supports_cookie_authorization: false, supports_refresh_token_import: true, supports_batch_import: true, supports_device_flow: false, @@ -24,6 +26,20 @@ impl ProviderOAuthCapabilities { }; } +#[derive(Clone, PartialEq, Eq)] +pub struct ProviderOAuthCookieAuthorizationInput { + pub session_key: String, +} + +impl std::fmt::Debug for ProviderOAuthCookieAuthorizationInput { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProviderOAuthCookieAuthorizationInput") + .field("session_key", &"") + .finish() + } +} + #[derive(Debug, Clone, PartialEq)] pub struct ProviderOAuthTransportContext { pub provider_id: String, diff --git a/crates/aether-oauth/src/provider/adapter.rs b/crates/aether-oauth/src/provider/adapter.rs index 92f94eebc..2ad630e09 100644 --- a/crates/aether-oauth/src/provider/adapter.rs +++ b/crates/aether-oauth/src/provider/adapter.rs @@ -1,7 +1,7 @@ use super::{ ProviderOAuthAccount, ProviderOAuthAccountState, ProviderOAuthCapabilities, - ProviderOAuthImportInput, ProviderOAuthRequestAuth, ProviderOAuthTokenSet, - ProviderOAuthTransportContext, + ProviderOAuthCookieAuthorizationInput, ProviderOAuthImportInput, ProviderOAuthRequestAuth, + ProviderOAuthTokenSet, ProviderOAuthTransportContext, }; use crate::core::{OAuthAuthorizeResponse, OAuthError}; use crate::network::OAuthHttpExecutor; @@ -42,6 +42,17 @@ pub trait ProviderOAuthAdapter: Send + Sync { )) } + async fn authorize_with_cookie( + &self, + _executor: &dyn OAuthHttpExecutor, + _ctx: &ProviderOAuthTransportContext, + _input: ProviderOAuthCookieAuthorizationInput, + ) -> Result { + Err(OAuthError::UnsupportedProvider( + self.provider_type().to_string(), + )) + } + async fn import_credentials( &self, executor: &dyn OAuthHttpExecutor, diff --git a/crates/aether-oauth/src/provider/mod.rs b/crates/aether-oauth/src/provider/mod.rs index 8303814f7..0113c6d3d 100644 --- a/crates/aether-oauth/src/provider/mod.rs +++ b/crates/aether-oauth/src/provider/mod.rs @@ -5,8 +5,8 @@ mod service; pub use account::{ ProviderOAuthAccount, ProviderOAuthAccountState, ProviderOAuthCapabilities, - ProviderOAuthImportInput, ProviderOAuthRequestAuth, ProviderOAuthTokenSet, - ProviderOAuthTransportContext, + ProviderOAuthCookieAuthorizationInput, ProviderOAuthImportInput, ProviderOAuthRequestAuth, + ProviderOAuthTokenSet, ProviderOAuthTransportContext, }; pub use adapter::{ProviderOAuthAdapter, ProviderOAuthProbeResult}; pub use service::ProviderOAuthService; diff --git a/crates/aether-oauth/src/provider/providers/claude_code.rs b/crates/aether-oauth/src/provider/providers/claude_code.rs new file mode 100644 index 000000000..584b54a3a --- /dev/null +++ b/crates/aether-oauth/src/provider/providers/claude_code.rs @@ -0,0 +1,779 @@ +use super::generic::{template_for_provider_type, GenericProviderOAuthAdapter}; +use crate::core::{ + generate_oauth_nonce, generate_pkce_verifier, pkce_s256, OAuthAuthorizeResponse, +}; +use crate::network::{OAuthHttpExecutor, OAuthHttpRequest}; +use crate::provider::{ + ProviderOAuthAccount, ProviderOAuthAdapter, ProviderOAuthCapabilities, + ProviderOAuthCookieAuthorizationInput, ProviderOAuthImportInput, ProviderOAuthProbeResult, + ProviderOAuthRequestAuth, ProviderOAuthTokenSet, ProviderOAuthTransportContext, +}; +use crate::OAuthError; +use aether_contracts::{ + ResolvedTransportProfile, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, + TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_POOL_SCOPE_KEY, +}; +use async_trait::async_trait; +use serde_json::{json, Value}; +use std::collections::BTreeMap; +use url::Url; + +pub const CLAUDE_CODE_PROVIDER_TYPE: &str = "claude_code"; +pub const CLAUDE_CODE_CLIENT_ID: &str = "9d1c250a-e61b-44d9-88ed-5944d1962f5e"; +pub const CLAUDE_CODE_WEB_BASE_URL: &str = "https://claude.ai"; +pub const CLAUDE_CODE_AUTHORIZE_URL: &str = "https://claude.ai/oauth/authorize"; +pub const CLAUDE_CODE_TOKEN_URL: &str = "https://platform.claude.com/v1/oauth/token"; +pub const CLAUDE_CODE_REDIRECT_URI: &str = "https://platform.claude.com/oauth/code/callback"; +pub const CLAUDE_CODE_OAUTH_SCOPES: &[&str] = &[ + "org:create_api_key", + "user:profile", + "user:inference", + "user:sessions:claude_code", + "user:mcp_servers", + "user:file_upload", +]; +pub const CLAUDE_CODE_COOKIE_SCOPE: &str = + "user:profile user:inference user:sessions:claude_code user:mcp_servers user:file_upload"; + +const MAX_CLAUDE_SESSION_KEY_BYTES: usize = 16 * 1024; +const CLAUDE_CODE_BROWSER_USER_AGENT: &str = "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/136.0.0.0 Safari/537.36"; + +#[derive(Debug, Clone)] +pub struct ClaudeCodeProviderOAuthAdapter { + inner: GenericProviderOAuthAdapter, + web_base_url: String, +} + +impl Default for ClaudeCodeProviderOAuthAdapter { + fn default() -> Self { + Self { + inner: GenericProviderOAuthAdapter::new( + template_for_provider_type(CLAUDE_CODE_PROVIDER_TYPE) + .expect("claude code oauth template should exist"), + ), + web_base_url: CLAUDE_CODE_WEB_BASE_URL.to_string(), + } + } +} + +impl ClaudeCodeProviderOAuthAdapter { + pub fn with_endpoint_overrides( + mut self, + web_base_url: impl Into, + token_url: impl Into, + ) -> Self { + self.web_base_url = web_base_url.into(); + self.inner = self.inner.with_token_url_override(token_url); + self + } + + fn web_url(&self, path_segments: &[&str]) -> Result { + let mut url = Url::parse(self.web_base_url.trim()) + .map_err(|_| OAuthError::invalid_request("claude web base url must be absolute"))?; + url.set_query(None); + url.set_fragment(None); + { + let mut segments = url + .path_segments_mut() + .map_err(|_| OAuthError::invalid_request("claude web base url is invalid"))?; + segments.clear(); + segments.extend(path_segments.iter().copied()); + } + Ok(url.to_string()) + } + + fn session_cookie(session_key: &str) -> Result { + let session_key = session_key.trim(); + if session_key.is_empty() + || session_key.len() > MAX_CLAUDE_SESSION_KEY_BYTES + || session_key.contains(['\r', '\n', ';']) + || http::HeaderValue::from_str(session_key).is_err() + { + return Err(OAuthError::invalid_request("invalid Claude sessionKey")); + } + Ok(format!("sessionKey={session_key}")) + } + + async fn organization_uuid( + &self, + executor: &dyn OAuthHttpExecutor, + ctx: &ProviderOAuthTransportContext, + cookie: &str, + ) -> Result { + let response = executor + .execute(OAuthHttpRequest { + request_id: "provider-oauth:claude-cookie-organizations".to_string(), + method: reqwest::Method::GET, + url: self.web_url(&["api", "organizations"])?, + headers: cookie_headers(cookie, false), + content_type: None, + json_body: None, + body_bytes: None, + network: ctx.network.clone(), + transport_profile: claude_code_oauth_transport_profile_for_context(ctx), + }) + .await?; + ensure_success(&response)?; + let organizations = response + .json_body + .or_else(|| serde_json::from_str::(&response.body_text).ok()) + .and_then(|value| value.as_array().cloned()) + .ok_or_else(|| { + OAuthError::invalid_response("Claude organizations response is invalid") + })?; + + let organization = if organizations.len() == 1 { + organizations.first() + } else { + organizations + .iter() + .find(|organization| { + organization + .get("raven_type") + .and_then(Value::as_str) + .is_some_and(|value| value.eq_ignore_ascii_case("team")) + }) + .or_else(|| organizations.first()) + } + .ok_or_else(|| OAuthError::invalid_response("Claude account has no organizations"))?; + + organization + .get("uuid") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .ok_or_else(|| OAuthError::invalid_response("Claude organization is missing uuid")) + } + + async fn authorization_code( + &self, + executor: &dyn OAuthHttpExecutor, + ctx: &ProviderOAuthTransportContext, + cookie: &str, + organization_uuid: &str, + state: &str, + code_challenge: &str, + ) -> Result { + let response = executor + .execute(OAuthHttpRequest { + request_id: "provider-oauth:claude-cookie-authorize".to_string(), + method: reqwest::Method::POST, + url: self.web_url(&["v1", "oauth", organization_uuid, "authorize"])?, + headers: cookie_headers(cookie, true), + content_type: Some("application/json".to_string()), + json_body: Some(json!({ + "response_type": "code", + "client_id": CLAUDE_CODE_CLIENT_ID, + "organization_uuid": organization_uuid, + "redirect_uri": CLAUDE_CODE_REDIRECT_URI, + "scope": CLAUDE_CODE_COOKIE_SCOPE, + "state": state, + "code_challenge": code_challenge, + "code_challenge_method": "S256", + })), + body_bytes: None, + network: ctx.network.clone(), + transport_profile: claude_code_oauth_transport_profile_for_context(ctx), + }) + .await?; + ensure_success(&response)?; + let payload = response + .json_body + .or_else(|| serde_json::from_str::(&response.body_text).ok()) + .ok_or_else(|| OAuthError::invalid_response("Claude authorize response is invalid"))?; + let redirect_uri = payload + .get("redirect_uri") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + .ok_or_else(|| { + OAuthError::invalid_response("Claude authorize response is missing redirect_uri") + })?; + + validate_authorization_redirect(&redirect_uri, state) + } +} + +#[async_trait] +impl ProviderOAuthAdapter for ClaudeCodeProviderOAuthAdapter { + fn provider_type(&self) -> &'static str { + CLAUDE_CODE_PROVIDER_TYPE + } + + fn capabilities(&self) -> ProviderOAuthCapabilities { + ProviderOAuthCapabilities { + supports_cookie_authorization: true, + ..self.inner.capabilities() + } + } + + fn build_authorize_url( + &self, + ctx: &ProviderOAuthTransportContext, + state: &str, + code_challenge: Option<&str>, + ) -> Result { + let mut response = self.inner.build_authorize_url(ctx, state, code_challenge)?; + let mut url = Url::parse(&response.authorize_url) + .map_err(|_| OAuthError::invalid_response("invalid Claude authorize_url"))?; + url.query_pairs_mut().append_pair("code", "true"); + response.authorize_url = url.to_string(); + Ok(response) + } + + async fn exchange_code( + &self, + executor: &dyn OAuthHttpExecutor, + ctx: &ProviderOAuthTransportContext, + code: &str, + state: &str, + pkce_verifier: Option<&str>, + ) -> Result { + self.inner + .exchange_code(executor, ctx, code, state, pkce_verifier) + .await + } + + async fn authorize_with_cookie( + &self, + executor: &dyn OAuthHttpExecutor, + ctx: &ProviderOAuthTransportContext, + input: ProviderOAuthCookieAuthorizationInput, + ) -> Result { + let cookie = Self::session_cookie(&input.session_key)?; + let organization_uuid = self.organization_uuid(executor, ctx, &cookie).await?; + let state = generate_oauth_nonce(); + let verifier = generate_pkce_verifier(); + let challenge = pkce_s256(&verifier); + let code = self + .authorization_code( + executor, + ctx, + &cookie, + &organization_uuid, + &state, + &challenge, + ) + .await?; + self.inner + .exchange_code(executor, ctx, &code, &state, Some(&verifier)) + .await + } + + async fn import_credentials( + &self, + executor: &dyn OAuthHttpExecutor, + ctx: &ProviderOAuthTransportContext, + input: ProviderOAuthImportInput, + ) -> Result { + self.inner.import_credentials(executor, ctx, input).await + } + + async fn refresh( + &self, + executor: &dyn OAuthHttpExecutor, + ctx: &ProviderOAuthTransportContext, + account: &ProviderOAuthAccount, + ) -> Result { + self.inner.refresh(executor, ctx, account).await + } + + fn resolve_request_auth( + &self, + account: &ProviderOAuthAccount, + ) -> Result { + self.inner.resolve_request_auth(account) + } + + fn account_fingerprint(&self, account: &ProviderOAuthAccount) -> Option { + self.inner.account_fingerprint(account) + } + + async fn probe_account_state( + &self, + executor: &dyn OAuthHttpExecutor, + ctx: &ProviderOAuthTransportContext, + account: &ProviderOAuthAccount, + ) -> Result, OAuthError> { + self.inner.probe_account_state(executor, ctx, account).await + } +} + +pub(super) fn claude_code_oauth_transport_profile() -> ResolvedTransportProfile { + ResolvedTransportProfile { + profile_id: "claude_oauth_chrome136".to_string(), + backend: TRANSPORT_BACKEND_BROWSER_WREQ.to_string(), + http_mode: TRANSPORT_HTTP_MODE_AUTO.to_string(), + pool_scope: TRANSPORT_POOL_SCOPE_KEY.to_string(), + header_fingerprint: None, + extra: Some(json!({ "browser_profile": "chrome136" })), + } +} + +fn claude_code_oauth_transport_profile_for_context( + ctx: &ProviderOAuthTransportContext, +) -> Option { + // Node-only proxies must execute through the tunnel runtime, which cannot use browser_wreq. + // The explicit browser headers still keep that fallback compatible with Claude's web flow. + let tunnel_only_proxy = ctx.network.proxy.as_ref().is_some_and(|proxy| { + if proxy.enabled == Some(false) { + return false; + } + let has_proxy_url = proxy + .url + .as_deref() + .map(str::trim) + .is_some_and(|value| !value.is_empty()); + let has_node_id = proxy + .node_id + .as_deref() + .map(str::trim) + .is_some_and(|value| !value.is_empty()); + let tunnel_mode = proxy + .mode + .as_deref() + .map(str::trim) + .is_some_and(|value| value.eq_ignore_ascii_case("tunnel")); + has_node_id && (tunnel_mode || !has_proxy_url) + }); + (!tunnel_only_proxy).then(claude_code_oauth_transport_profile) +} + +fn cookie_headers(cookie: &str, json_request: bool) -> BTreeMap { + let mut headers = BTreeMap::from([ + ("accept".to_string(), "application/json".to_string()), + ("accept-language".to_string(), "en-US,en;q=0.9".to_string()), + ("cache-control".to_string(), "no-cache".to_string()), + ("cookie".to_string(), cookie.to_string()), + ( + "user-agent".to_string(), + CLAUDE_CODE_BROWSER_USER_AGENT.to_string(), + ), + ( + EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.to_string(), + "false".to_string(), + ), + ]); + if json_request { + headers.insert("content-type".to_string(), "application/json".to_string()); + headers.insert("origin".to_string(), CLAUDE_CODE_WEB_BASE_URL.to_string()); + headers.insert( + "referer".to_string(), + format!("{CLAUDE_CODE_WEB_BASE_URL}/new"), + ); + } + headers +} + +fn ensure_success(response: &crate::network::OAuthHttpResponse) -> Result<(), OAuthError> { + if (200..300).contains(&response.status_code) { + return Ok(()); + } + Err(OAuthError::HttpStatus { + status_code: response.status_code, + body_excerpt: "Claude Cookie authorization request failed".to_string(), + }) +} + +fn validate_authorization_redirect( + redirect_uri: &str, + expected_state: &str, +) -> Result { + let redirect = Url::parse(redirect_uri) + .map_err(|_| OAuthError::invalid_response("Claude authorize redirect_uri is invalid"))?; + let expected = Url::parse(CLAUDE_CODE_REDIRECT_URI).map_err(|_| { + OAuthError::invalid_response("Claude redirect URI configuration is invalid") + })?; + if redirect.scheme() != expected.scheme() + || redirect.host_str() != expected.host_str() + || redirect.port_or_known_default() != expected.port_or_known_default() + || redirect.path() != expected.path() + || !redirect.username().is_empty() + || redirect.password().is_some() + || redirect.fragment().is_some() + { + return Err(OAuthError::invalid_response( + "Claude authorize redirect_uri target is invalid", + )); + } + + let mut code = None; + let mut state = None; + for (key, value) in redirect.query_pairs() { + match key.as_ref() { + "code" if code.is_none() => code = Some(value.into_owned()), + "state" if state.is_none() => state = Some(value.into_owned()), + "code" | "state" => { + return Err(OAuthError::invalid_response( + "Claude authorize redirect_uri has duplicate parameters", + )); + } + _ => {} + } + } + let code = code + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) + .ok_or_else(|| { + OAuthError::invalid_response("Claude authorize redirect_uri is missing code") + })?; + let state = state + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) + .ok_or_else(|| { + OAuthError::invalid_response("Claude authorize redirect_uri is missing state") + })?; + if state != expected_state { + return Err(OAuthError::InvalidState); + } + Ok(code) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::network::{OAuthHttpResponse, OAuthNetworkContext}; + use aether_contracts::ProxySnapshot; + use async_trait::async_trait; + use std::sync::{Arc, Mutex}; + + #[derive(Debug, Clone, Copy, Default)] + enum RedirectMode { + #[default] + Matching, + WrongState, + HostileHost, + } + + #[derive(Clone)] + struct RecordingExecutor { + requests: Arc>>, + organizations: Value, + token_payload: Value, + redirect_mode: RedirectMode, + } + + impl Default for RecordingExecutor { + fn default() -> Self { + Self { + requests: Arc::new(Mutex::new(Vec::new())), + organizations: json!([ + {"uuid": "org-personal", "raven_type": "personal"}, + {"uuid": "org-team", "raven_type": "team"} + ]), + token_payload: json!({ + "access_token": "sk-ant-oat01-new", + "refresh_token": "sk-ant-ort01-new", + "expires_in": 3600, + "organization": {"uuid": "org-team"}, + "account": { + "uuid": "account-123", + "email_address": "alice@example.com" + } + }), + redirect_mode: RedirectMode::Matching, + } + } + } + + #[async_trait] + impl OAuthHttpExecutor for RecordingExecutor { + async fn execute( + &self, + request: OAuthHttpRequest, + ) -> Result { + self.requests + .lock() + .expect("requests lock") + .push(request.clone()); + + let payload = if request.url.ends_with("/api/organizations") { + self.organizations.clone() + } else if request.url.contains("/authorize") { + let requested_state = request + .json_body + .as_ref() + .and_then(|body| body.get("state")) + .and_then(Value::as_str) + .expect("authorize request should contain state"); + let state = match self.redirect_mode { + RedirectMode::Matching | RedirectMode::HostileHost => requested_state, + RedirectMode::WrongState => "wrong-state", + }; + let redirect_base = match self.redirect_mode { + RedirectMode::HostileHost => { + "https://platform.claude.com.evil/oauth/code/callback" + } + _ => CLAUDE_CODE_REDIRECT_URI, + }; + let mut redirect = Url::parse(redirect_base).expect("redirect URL should parse"); + redirect + .query_pairs_mut() + .append_pair("code", "authorization-code") + .append_pair("state", state); + json!({"redirect_uri": redirect.to_string()}) + } else { + self.token_payload.clone() + }; + + Ok(OAuthHttpResponse { + status_code: 200, + body_text: payload.to_string(), + json_body: Some(payload), + }) + } + } + + fn context(proxy: Option) -> ProviderOAuthTransportContext { + ProviderOAuthTransportContext { + provider_id: "provider-claude".to_string(), + provider_type: CLAUDE_CODE_PROVIDER_TYPE.to_string(), + endpoint_id: None, + key_id: None, + auth_type: Some("oauth".to_string()), + decrypted_api_key: None, + decrypted_auth_config: None, + provider_config: None, + endpoint_config: None, + key_config: None, + network: OAuthNetworkContext::provider_operation(proxy), + } + } + + #[test] + fn builds_current_manual_authorize_url() { + let adapter = ClaudeCodeProviderOAuthAdapter::default(); + let response = adapter + .build_authorize_url(&context(None), "state-123", Some("challenge-123")) + .expect("authorize URL should build"); + let url = Url::parse(&response.authorize_url).expect("authorize URL should parse"); + let query = url + .query_pairs() + .map(|(key, value)| (key.into_owned(), value.into_owned())) + .collect::>(); + + assert_eq!( + format!( + "{}://{}{}", + url.scheme(), + url.host_str().unwrap_or_default(), + url.path() + ), + CLAUDE_CODE_AUTHORIZE_URL + ); + assert_eq!( + query.get("client_id").map(String::as_str), + Some(CLAUDE_CODE_CLIENT_ID) + ); + assert_eq!( + query.get("redirect_uri").map(String::as_str), + Some(CLAUDE_CODE_REDIRECT_URI) + ); + assert_eq!( + query.get("scope").map(String::as_str), + Some(CLAUDE_CODE_OAUTH_SCOPES.join(" ").as_str()) + ); + assert_eq!(query.get("code").map(String::as_str), Some("true")); + assert_eq!( + query.get("code_challenge").map(String::as_str), + Some("challenge-123") + ); + assert!(adapter.capabilities().supports_cookie_authorization); + } + + #[tokio::test] + async fn cookie_authorization_uses_team_org_safe_headers_and_current_token_contract() { + let executor = RecordingExecutor::default(); + let adapter = ClaudeCodeProviderOAuthAdapter::default().with_endpoint_overrides( + "https://claude.test", + "https://platform.test/v1/oauth/token", + ); + let session_key = "sk-ant-sid01-secret"; + + let result = adapter + .authorize_with_cookie( + &executor, + &context(None), + ProviderOAuthCookieAuthorizationInput { + session_key: session_key.to_string(), + }, + ) + .await + .expect("cookie authorization should succeed"); + + assert_eq!(result.token_set.access_token, "sk-ant-oat01-new"); + assert_eq!( + result.token_set.refresh_token.as_deref(), + Some("sk-ant-ort01-new") + ); + assert_eq!(result.auth_config["org_uuid"], "org-team"); + assert_eq!(result.auth_config["account_uuid"], "account-123"); + assert_eq!(result.auth_config["email"], "alice@example.com"); + + let requests = executor.requests.lock().expect("requests lock").clone(); + assert_eq!(requests.len(), 3); + assert_eq!(requests[0].method, reqwest::Method::GET); + assert!(requests[0].url.ends_with("/api/organizations")); + assert!(requests[1].url.ends_with("/v1/oauth/org-team/authorize")); + for request in &requests[..2] { + assert_eq!( + request.headers.get("cookie").map(String::as_str), + Some("sessionKey=sk-ant-sid01-secret") + ); + assert_eq!( + request + .headers + .get(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER) + .map(String::as_str), + Some("false") + ); + assert_eq!( + request.headers.get("user-agent").map(String::as_str), + Some(CLAUDE_CODE_BROWSER_USER_AGENT) + ); + assert_eq!( + request + .transport_profile + .as_ref() + .map(|profile| profile.backend.as_str()), + Some(TRANSPORT_BACKEND_BROWSER_WREQ) + ); + assert!(!format!("{request:?}").contains(session_key)); + } + assert_eq!( + requests[1] + .json_body + .as_ref() + .and_then(|body| body.get("scope")) + .and_then(Value::as_str), + Some(CLAUDE_CODE_COOKIE_SCOPE) + ); + + let token_request = &requests[2]; + assert_eq!(token_request.url, "https://platform.test/v1/oauth/token"); + assert!(!token_request.headers.contains_key("cookie")); + assert!(token_request.transport_profile.is_none()); + assert_eq!( + token_request.headers.get("user-agent").map(String::as_str), + Some("axios/1.13.6") + ); + let token_body = token_request + .json_body + .as_ref() + .expect("token request should be JSON"); + assert!(token_body.get("scope").is_none()); + assert_eq!(token_body["redirect_uri"], CLAUDE_CODE_REDIRECT_URI); + assert_eq!(token_body["code"], "authorization-code"); + assert!(token_body.get("code_verifier").is_some()); + } + + #[tokio::test] + async fn rejects_wrong_state_and_hostile_authorize_redirects() { + for redirect_mode in [RedirectMode::WrongState, RedirectMode::HostileHost] { + let executor = RecordingExecutor { + redirect_mode, + ..RecordingExecutor::default() + }; + let adapter = ClaudeCodeProviderOAuthAdapter::default().with_endpoint_overrides( + "https://claude.test", + "https://platform.test/v1/oauth/token", + ); + let error = adapter + .authorize_with_cookie( + &executor, + &context(None), + ProviderOAuthCookieAuthorizationInput { + session_key: "sk-ant-sid01-secret".to_string(), + }, + ) + .await + .expect_err("unsafe redirect should be rejected"); + match redirect_mode { + RedirectMode::WrongState => assert!(matches!(error, OAuthError::InvalidState)), + RedirectMode::HostileHost => { + assert!(matches!(error, OAuthError::InvalidResponse(_))) + } + RedirectMode::Matching => unreachable!(), + } + } + } + + #[test] + fn tunnel_proxy_falls_back_from_browser_transport_even_when_url_is_present() { + for proxy in [ + ProxySnapshot { + node_id: Some("node-only".to_string()), + ..ProxySnapshot::default() + }, + ProxySnapshot { + mode: Some("tunnel".to_string()), + node_id: Some("node-with-url".to_string()), + url: Some("http://127.0.0.1:9999".to_string()), + ..ProxySnapshot::default() + }, + ] { + assert!( + claude_code_oauth_transport_profile_for_context(&context(Some(proxy))).is_none() + ); + } + + let url_proxy = ProxySnapshot { + mode: Some("url".to_string()), + node_id: Some("metadata-node".to_string()), + url: Some("http://127.0.0.1:9999".to_string()), + ..ProxySnapshot::default() + }; + assert!( + claude_code_oauth_transport_profile_for_context(&context(Some(url_proxy))).is_some() + ); + assert_eq!( + cookie_headers("sessionKey=test", false) + .get("user-agent") + .map(String::as_str), + Some(CLAUDE_CODE_BROWSER_USER_AGENT) + ); + } + + #[tokio::test] + async fn refresh_rotates_claude_refresh_token_without_scope() { + let executor = RecordingExecutor::default(); + let adapter = ClaudeCodeProviderOAuthAdapter::default().with_endpoint_overrides( + "https://claude.test", + "https://platform.test/v1/oauth/token", + ); + let account = ProviderOAuthAccount { + provider_type: CLAUDE_CODE_PROVIDER_TYPE.to_string(), + access_token: "sk-ant-oat01-old".to_string(), + auth_config: json!({ + "provider_type": CLAUDE_CODE_PROVIDER_TYPE, + "refresh_token": "sk-ant-ort01-old", + "email": "old@example.com" + }), + expires_at_unix_secs: Some(1), + identity: BTreeMap::new(), + }; + + let refreshed = adapter + .refresh(&executor, &context(None), &account) + .await + .expect("refresh should succeed"); + + assert_eq!(refreshed.token_set.access_token, "sk-ant-oat01-new"); + assert_eq!( + refreshed.token_set.refresh_token.as_deref(), + Some("sk-ant-ort01-new") + ); + assert_eq!(refreshed.auth_config["refresh_token"], "sk-ant-ort01-new"); + let requests = executor.requests.lock().expect("requests lock"); + assert_eq!(requests.len(), 1); + let body = requests[0] + .json_body + .as_ref() + .expect("refresh request should be JSON"); + assert_eq!(body["grant_type"], "refresh_token"); + assert_eq!(body["refresh_token"], "sk-ant-ort01-old"); + assert!(body.get("scope").is_none()); + } +} diff --git a/crates/aether-oauth/src/provider/providers/generic.rs b/crates/aether-oauth/src/provider/providers/generic.rs index b68769cf4..9e27dca46 100644 --- a/crates/aether-oauth/src/provider/providers/generic.rs +++ b/crates/aether-oauth/src/provider/providers/generic.rs @@ -12,6 +12,11 @@ use sha2::{Digest, Sha256}; use std::collections::BTreeMap; use url::form_urlencoded; +use super::claude_code::{ + CLAUDE_CODE_AUTHORIZE_URL, CLAUDE_CODE_CLIENT_ID, CLAUDE_CODE_OAUTH_SCOPES, + CLAUDE_CODE_PROVIDER_TYPE, CLAUDE_CODE_REDIRECT_URI, CLAUDE_CODE_TOKEN_URL, +}; + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct GenericProviderOAuthTemplate { pub provider_type: &'static str, @@ -24,20 +29,22 @@ pub struct GenericProviderOAuthTemplate { pub redirect_uri: &'static str, pub use_pkce: bool, pub uses_json_payload: bool, + pub include_scope_in_token_request: bool, } pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[ GenericProviderOAuthTemplate { - provider_type: "claude_code", + provider_type: CLAUDE_CODE_PROVIDER_TYPE, display_name: "ClaudeCode", - authorize_url: "https://claude.ai/oauth/authorize", - token_url: "https://console.anthropic.com/v1/oauth/token", - client_id: "9d1c250a-e61b-44d9-88ed-5944d1962f5e", + authorize_url: CLAUDE_CODE_AUTHORIZE_URL, + token_url: CLAUDE_CODE_TOKEN_URL, + client_id: CLAUDE_CODE_CLIENT_ID, client_secret: "", - scopes: &["org:create_api_key", "user:profile", "user:inference"], - redirect_uri: "http://localhost:54545/callback", + scopes: CLAUDE_CODE_OAUTH_SCOPES, + redirect_uri: CLAUDE_CODE_REDIRECT_URI, use_pkce: true, uses_json_payload: true, + include_scope_in_token_request: false, }, GenericProviderOAuthTemplate { provider_type: "codex", @@ -50,6 +57,7 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[ redirect_uri: "http://localhost:1455/auth/callback", use_pkce: true, uses_json_payload: false, + include_scope_in_token_request: true, }, GenericProviderOAuthTemplate { provider_type: "chatgpt_web", @@ -62,6 +70,7 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[ redirect_uri: "http://localhost:1455/auth/callback", use_pkce: true, uses_json_payload: false, + include_scope_in_token_request: true, }, GenericProviderOAuthTemplate { provider_type: "gemini_cli", @@ -78,6 +87,7 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[ redirect_uri: "http://localhost:8085/oauth2callback", use_pkce: false, uses_json_payload: false, + include_scope_in_token_request: true, }, GenericProviderOAuthTemplate { provider_type: "antigravity", @@ -96,6 +106,7 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[ redirect_uri: "http://localhost:51121/oauth2callback", use_pkce: true, uses_json_payload: false, + include_scope_in_token_request: true, }, ]; @@ -185,19 +196,22 @@ impl GenericProviderOAuthAdapter { Value::String(code_or_refresh_token.to_string()), ); } - if let Some(scope) = scope.as_ref() { - body.insert("scope".to_string(), Value::String(scope.clone())); + if self.template.include_scope_in_token_request { + if let Some(scope) = scope.as_ref() { + body.insert("scope".to_string(), Value::String(scope.clone())); + } } executor .execute(OAuthHttpRequest { request_id: request_id.clone(), method: reqwest::Method::POST, url: self.token_url(), - headers: json_headers(), + headers: json_headers(self.template.provider_type), content_type: Some("application/json".to_string()), json_body: Some(Value::Object(body)), body_bytes: None, network: ctx.network.clone(), + transport_profile: None, }) .await? } else { @@ -214,8 +228,10 @@ impl GenericProviderOAuthAdapter { } else { form.append_pair("refresh_token", code_or_refresh_token); } - if let Some(scope) = scope.as_ref() { - form.append_pair("scope", scope); + if self.template.include_scope_in_token_request { + if let Some(scope) = scope.as_ref() { + form.append_pair("scope", scope); + } } if !self.template.client_secret.trim().is_empty() { form.append_pair("client_secret", self.template.client_secret); @@ -232,6 +248,7 @@ impl GenericProviderOAuthAdapter { json_body: None, body_bytes: Some(form_body), network: ctx.network.clone(), + transport_profile: None, }) .await? }; @@ -422,11 +439,19 @@ fn form_headers() -> BTreeMap { ]) } -fn json_headers() -> BTreeMap { - BTreeMap::from([ +fn json_headers(provider_type: &str) -> BTreeMap { + let mut headers = BTreeMap::from([ ("content-type".to_string(), "application/json".to_string()), ("accept".to_string(), "application/json".to_string()), - ]) + ]); + if provider_type.eq_ignore_ascii_case(CLAUDE_CODE_PROVIDER_TYPE) { + headers.insert( + "accept".to_string(), + "application/json, text/plain, */*".to_string(), + ); + headers.insert("user-agent".to_string(), "axios/1.13.6".to_string()); + } + headers } fn truncate_body(body: &str) -> String { @@ -470,6 +495,32 @@ fn enrich_generic_identity( } } } + if provider_type.eq_ignore_ascii_case(CLAUDE_CODE_PROVIDER_TYPE) { + if let Some(organization_uuid) = token_payload + .get("organization") + .and_then(Value::as_object) + .and_then(|value| value.get("uuid")) + .cloned() + { + auth_config + .entry("org_uuid".to_string()) + .or_insert(organization_uuid); + } + if let Some(account) = token_payload.get("account").and_then(Value::as_object) { + if let Some(account_uuid) = account.get("uuid").cloned() { + auth_config + .entry("account_uuid".to_string()) + .or_insert(account_uuid); + } + if let Some(email) = account.get("email_address").cloned() { + auth_config + .entry("email_address".to_string()) + .or_insert_with(|| email.clone()); + auth_config.entry("email".to_string()).or_insert(email); + } + } + return; + } if !matches!( provider_type.trim().to_ascii_lowercase().as_str(), "codex" | "chatgpt_web" diff --git a/crates/aether-oauth/src/provider/providers/kiro.rs b/crates/aether-oauth/src/provider/providers/kiro.rs index c3d02bcf3..921a87e69 100644 --- a/crates/aether-oauth/src/provider/providers/kiro.rs +++ b/crates/aether-oauth/src/provider/providers/kiro.rs @@ -298,6 +298,7 @@ impl KiroProviderOAuthAdapter { })), body_bytes: None, network: ctx.network.clone(), + transport_profile: None, }) .await?; if !(200..300).contains(&response.status_code) { @@ -376,6 +377,7 @@ impl KiroProviderOAuthAdapter { })), body_bytes: None, network: ctx.network.clone(), + transport_profile: None, }) .await?; if !(200..300).contains(&response.status_code) { @@ -420,6 +422,7 @@ impl ProviderOAuthAdapter for KiroProviderOAuthAdapter { fn capabilities(&self) -> ProviderOAuthCapabilities { ProviderOAuthCapabilities { supports_authorization_code: false, + supports_cookie_authorization: false, supports_refresh_token_import: true, supports_batch_import: true, supports_device_flow: true, diff --git a/crates/aether-oauth/src/provider/providers/mod.rs b/crates/aether-oauth/src/provider/providers/mod.rs index 447108648..a9e72c092 100644 --- a/crates/aether-oauth/src/provider/providers/mod.rs +++ b/crates/aether-oauth/src/provider/providers/mod.rs @@ -1,10 +1,16 @@ mod antigravity; +mod claude_code; mod codex; mod generic; mod kiro; mod windsurf; pub use antigravity::AntigravityProviderOAuthAdapter; +pub use claude_code::{ + ClaudeCodeProviderOAuthAdapter, CLAUDE_CODE_AUTHORIZE_URL, CLAUDE_CODE_CLIENT_ID, + CLAUDE_CODE_COOKIE_SCOPE, CLAUDE_CODE_OAUTH_SCOPES, CLAUDE_CODE_PROVIDER_TYPE, + CLAUDE_CODE_REDIRECT_URI, CLAUDE_CODE_TOKEN_URL, CLAUDE_CODE_WEB_BASE_URL, +}; pub use codex::CodexProviderOAuthAdapter; pub use generic::{ GenericProviderOAuthAdapter, GenericProviderOAuthTemplate, GENERIC_PROVIDER_OAUTH_TEMPLATES, diff --git a/crates/aether-oauth/src/provider/providers/windsurf.rs b/crates/aether-oauth/src/provider/providers/windsurf.rs index bcb54f656..d437b0d27 100644 --- a/crates/aether-oauth/src/provider/providers/windsurf.rs +++ b/crates/aether-oauth/src/provider/providers/windsurf.rs @@ -104,6 +104,7 @@ impl WindsurfProviderOAuthAdapter { json_body, body_bytes, network: ctx.network.clone(), + transport_profile: None, }) .await; match response { @@ -234,6 +235,7 @@ impl WindsurfProviderOAuthAdapter { json_body: Some(json!({ "email": email, "password": password })), body_bytes: None, network: ctx.network.clone(), + transport_profile: None, }) .await?; if !(200..300).contains(&login_response.status_code) { @@ -266,6 +268,7 @@ impl WindsurfProviderOAuthAdapter { json_body: None, body_bytes: Some(Vec::new()), network: ctx.network.clone(), + transport_profile: None, }) .await; match response { @@ -388,6 +391,7 @@ impl ProviderOAuthAdapter for WindsurfProviderOAuthAdapter { fn capabilities(&self) -> ProviderOAuthCapabilities { ProviderOAuthCapabilities { supports_authorization_code: false, + supports_cookie_authorization: false, supports_refresh_token_import: true, supports_batch_import: true, supports_device_flow: true, diff --git a/crates/aether-oauth/src/provider/service.rs b/crates/aether-oauth/src/provider/service.rs index 4cc7b74c6..8b121f5de 100644 --- a/crates/aether-oauth/src/provider/service.rs +++ b/crates/aether-oauth/src/provider/service.rs @@ -1,6 +1,7 @@ use super::{ - ProviderOAuthAdapter, ProviderOAuthImportInput, ProviderOAuthProbeResult, - ProviderOAuthRequestAuth, ProviderOAuthTokenSet, ProviderOAuthTransportContext, + ProviderOAuthAdapter, ProviderOAuthCookieAuthorizationInput, ProviderOAuthImportInput, + ProviderOAuthProbeResult, ProviderOAuthRequestAuth, ProviderOAuthTokenSet, + ProviderOAuthTransportContext, }; use crate::core::{OAuthAdapterRegistry, OAuthAuthorizeResponse, OAuthError}; use crate::network::OAuthHttpExecutor; @@ -18,16 +19,18 @@ impl ProviderOAuthService { pub fn with_builtin_adapters() -> Self { use super::providers::{ - AntigravityProviderOAuthAdapter, CodexProviderOAuthAdapter, - GenericProviderOAuthAdapter, KiroProviderOAuthAdapter, WindsurfProviderOAuthAdapter, + AntigravityProviderOAuthAdapter, ClaudeCodeProviderOAuthAdapter, + CodexProviderOAuthAdapter, GenericProviderOAuthAdapter, KiroProviderOAuthAdapter, + WindsurfProviderOAuthAdapter, }; let mut service = Self::new() .with_adapter(Arc::new(KiroProviderOAuthAdapter::default())) + .with_adapter(Arc::new(ClaudeCodeProviderOAuthAdapter::default())) .with_adapter(Arc::new(CodexProviderOAuthAdapter::default())) .with_adapter(Arc::new(AntigravityProviderOAuthAdapter::default())) .with_adapter(Arc::new(WindsurfProviderOAuthAdapter)); - for provider_type in ["claude_code", "chatgpt_web", "gemini_cli"] { + for provider_type in ["chatgpt_web", "gemini_cli"] { if let Some(adapter) = GenericProviderOAuthAdapter::for_provider_type(provider_type) { service = service.with_adapter(Arc::new(adapter)); } @@ -83,6 +86,17 @@ impl ProviderOAuthService { .await } + pub async fn authorize_with_cookie( + &self, + executor: &dyn OAuthHttpExecutor, + ctx: &ProviderOAuthTransportContext, + input: ProviderOAuthCookieAuthorizationInput, + ) -> Result { + self.adapter(&ctx.provider_type)? + .authorize_with_cookie(executor, ctx, input) + .await + } + pub async fn refresh( &self, executor: &dyn OAuthHttpExecutor, diff --git a/crates/aether-pool-core/src/scheduler.rs b/crates/aether-pool-core/src/scheduler.rs index 61342b3c1..c0d5f1c79 100644 --- a/crates/aether-pool-core/src/scheduler.rs +++ b/crates/aether-pool-core/src/scheduler.rs @@ -360,21 +360,12 @@ fn build_pool_sort_vectors( let lru_ranks = lru_rank_indices(items, false); let cache_affinity_ranks = lru_rank_indices(items, true); - if lru_enabled { - for item in items { - let key_id = item.item.facts.key_id.clone(); - vectors - .entry(key_id.clone()) - .or_default() - .push(*lru_ranks.get(&key_id).unwrap_or(&0)); - } - } - for preset in presets { let ranks = match preset.preset.as_str() { "cache_affinity" => cache_affinity_ranks.clone(), "priority_first" => priority_first_ranks(items, &lru_ranks), "single_account" => single_account_ranks(items), + "free_team_first" => plan_ranks(items, &lru_ranks, preset.mode.as_deref()), "plus_first" => plan_ranks(items, &lru_ranks, Some("plus_only")), "pro_first" => plan_ranks(items, &lru_ranks, Some("pro_only")), "free_first" => plan_ranks(items, &lru_ranks, Some("free_only")), @@ -396,6 +387,16 @@ fn build_pool_sort_vectors( } } + if lru_enabled { + for item in items { + let key_id = item.item.facts.key_id.clone(); + vectors + .entry(key_id.clone()) + .or_default() + .push(*lru_ranks.get(&key_id).unwrap_or(&0)); + } + } + vectors } @@ -415,13 +416,13 @@ fn lru_rank_indices( fn priority_first_ranks( items: &[PoolGroupCandidateOrdering], - _lru_ranks: &BTreeMap, + lru_ranks: &BTreeMap, ) -> BTreeMap { let scores = collect_metric_scores(items, |item| { Some(f64::from(item.item.facts.key_internal_priority)) }); if !score_map_has_variation(&scores) { - return neutral_rank_indices(items); + return lru_ranks.clone(); } rank_indices_from_score_map(items, &scores, false) } @@ -457,7 +458,7 @@ fn single_account_ranks( fn plan_ranks( items: &[PoolGroupCandidateOrdering], - _lru_ranks: &BTreeMap, + lru_ranks: &BTreeMap, mode: Option<&str>, ) -> BTreeMap { let scores = items @@ -473,14 +474,14 @@ fn plan_ranks( }) .collect::>(); if !score_map_has_variation(&scores) { - return neutral_rank_indices(items); + return lru_ranks.clone(); } rank_indices_from_score_map(items, &scores, false) } fn health_first_ranks( items: &[PoolGroupCandidateOrdering], - _lru_ranks: &BTreeMap, + lru_ranks: &BTreeMap, ) -> BTreeMap { let scores = collect_metric_scores(items, |item| { item.item @@ -489,60 +490,63 @@ fn health_first_ranks( .map(|score| 1.0 - score.clamp(0.0, 1.0)) }); if !score_map_has_signal(&scores) { - return neutral_rank_indices(items); + return lru_ranks.clone(); } rank_indices_from_score_map(items, &scores, false) } fn latency_first_ranks( items: &[PoolGroupCandidateOrdering], - _lru_ranks: &BTreeMap, + lru_ranks: &BTreeMap, ) -> BTreeMap { let scores = collect_metric_scores(items, |item| item.item.key_context.latency_avg_ms); if !score_map_has_signal(&scores) { - return neutral_rank_indices(items); + return lru_ranks.clone(); } rank_indices_from_score_map(items, &scores, false) } fn cost_first_ranks( items: &[PoolGroupCandidateOrdering], - _lru_ranks: &BTreeMap, + lru_ranks: &BTreeMap, cost_limit_per_key_tokens: Option, ) -> BTreeMap { + let has_cost_signal = items.iter().any(|item| item.cost_usage > 0); let scores = collect_metric_scores(items, |item| { - cost_penalty(item, cost_limit_per_key_tokens).or(item.item.key_context.quota_usage_ratio) + cost_penalty(item, cost_limit_per_key_tokens, has_cost_signal) + .or(item.item.key_context.quota_usage_ratio) }); if !score_map_has_signal(&scores) { - return neutral_rank_indices(items); + return lru_ranks.clone(); } rank_indices_from_score_map(items, &scores, false) } fn quota_balanced_ranks( items: &[PoolGroupCandidateOrdering], - _lru_ranks: &BTreeMap, + lru_ranks: &BTreeMap, cost_limit_per_key_tokens: Option, ) -> BTreeMap { + let has_cost_signal = items.iter().any(|item| item.cost_usage > 0); let scores = collect_metric_scores(items, |item| { item.item .key_context .quota_usage_ratio - .or_else(|| cost_penalty(item, cost_limit_per_key_tokens)) + .or_else(|| cost_penalty(item, cost_limit_per_key_tokens, has_cost_signal)) }); if !score_map_has_signal(&scores) { - return neutral_rank_indices(items); + return lru_ranks.clone(); } rank_indices_from_score_map(items, &scores, false) } fn recent_refresh_ranks( items: &[PoolGroupCandidateOrdering], - _lru_ranks: &BTreeMap, + lru_ranks: &BTreeMap, ) -> BTreeMap { let scores = collect_metric_scores(items, |item| item.item.key_context.quota_reset_seconds); if !score_map_has_signal(&scores) { - return neutral_rank_indices(items); + return lru_ranks.clone(); } rank_indices_from_score_map(items, &scores, false) } @@ -653,29 +657,33 @@ fn rank_indices_from_score_map( .then(left.2.cmp(&right.2)) }); - decorated - .into_iter() - .enumerate() - .map(|(rank, (_, _, _, key_id))| (key_id, rank)) - .collect() -} - -fn neutral_rank_indices( - items: &[PoolGroupCandidateOrdering], -) -> BTreeMap { - items - .iter() - .map(|item| (item.item.facts.key_id.clone(), 0)) - .collect() + let mut ranks = BTreeMap::new(); + let mut previous_score = None::<(bool, f64)>; + let mut rank = 0; + for (index, (missing, score, _, key_id)) in decorated.into_iter().enumerate() { + if previous_score + .as_ref() + .is_some_and(|previous| previous.0 != missing || previous.1 != score) + { + rank = index; + } + previous_score = Some((missing, score)); + ranks.insert(key_id, rank); + } + ranks } fn cost_penalty( item: &PoolGroupCandidateOrdering, cost_limit_per_key_tokens: Option, + has_cost_signal: bool, ) -> Option { - if item.cost_usage == 0 { + if !has_cost_signal { return None; } + if item.cost_usage == 0 { + return Some(0.0); + } if let Some(limit) = cost_limit_per_key_tokens.filter(|limit| *limit > 0) { return Some((item.cost_usage as f64 / limit as f64).clamp(0.0, 1.0)); @@ -1006,6 +1014,223 @@ mod tests { ); } + #[test] + fn cost_first_treats_zero_usage_as_the_lowest_cost() { + let unused = sample_candidate("provider-pool", "endpoint-1", "key-unused", 10, true) + .with_presets(vec![PoolSchedulingPreset { + preset: "cost_first".to_string(), + enabled: true, + mode: None, + }]); + let used = sample_candidate("provider-pool", "endpoint-1", "key-used", 10, true) + .with_presets(vec![PoolSchedulingPreset { + preset: "cost_first".to_string(), + enabled: true, + mode: None, + }]); + let runtime_by_provider = BTreeMap::from([( + "provider-pool".to_string(), + PoolRuntimeState { + cost_window_usage_by_key: BTreeMap::from([("key-used".to_string(), 100)]), + lru_score_by_key: BTreeMap::from([ + ("key-used".to_string(), 0.0), + ("key-unused".to_string(), 100.0), + ]), + ..PoolRuntimeState::default() + }, + )]); + + let outcome = run_pool_scheduler(vec![used, unused], &runtime_by_provider, "seed"); + + assert!(outcome.skipped_candidates.is_empty()); + assert_eq!( + outcome + .candidates + .iter() + .map(|item| item.candidate.as_str()) + .collect::>(), + vec!["key-unused", "key-used"] + ); + } + + #[test] + fn cost_first_falls_back_to_quota_without_cost_signal() { + let high_quota_usage = + sample_candidate("provider-pool", "endpoint-1", "key-high-quota", 10, true) + .with_presets(vec![PoolSchedulingPreset { + preset: "cost_first".to_string(), + enabled: true, + mode: None, + }]) + .with_quota_usage(0.8); + let low_quota_usage = + sample_candidate("provider-pool", "endpoint-1", "key-low-quota", 10, true) + .with_presets(vec![PoolSchedulingPreset { + preset: "cost_first".to_string(), + enabled: true, + mode: None, + }]) + .with_quota_usage(0.2); + let runtime_by_provider = BTreeMap::from([( + "provider-pool".to_string(), + PoolRuntimeState { + lru_score_by_key: BTreeMap::from([ + ("key-high-quota".to_string(), 0.0), + ("key-low-quota".to_string(), 100.0), + ]), + ..PoolRuntimeState::default() + }, + )]); + + let outcome = run_pool_scheduler( + vec![high_quota_usage, low_quota_usage], + &runtime_by_provider, + "seed", + ); + + assert_eq!( + outcome + .candidates + .iter() + .map(|item| item.candidate.as_str()) + .collect::>(), + vec!["key-low-quota", "key-high-quota"] + ); + } + + #[test] + fn cost_first_falls_back_to_lru_without_cost_or_quota_signal() { + let key_recent = sample_candidate("provider-pool", "endpoint-1", "key-recent", 10, true) + .with_presets(vec![PoolSchedulingPreset { + preset: "cost_first".to_string(), + enabled: true, + mode: None, + }]); + let key_old = sample_candidate("provider-pool", "endpoint-1", "key-old", 10, true) + .with_presets(vec![PoolSchedulingPreset { + preset: "cost_first".to_string(), + enabled: true, + mode: None, + }]); + let runtime_by_provider = BTreeMap::from([( + "provider-pool".to_string(), + PoolRuntimeState { + lru_score_by_key: BTreeMap::from([ + ("key-recent".to_string(), 100.0), + ("key-old".to_string(), 0.0), + ]), + ..PoolRuntimeState::default() + }, + )]); + + let outcome = run_pool_scheduler(vec![key_recent, key_old], &runtime_by_provider, "seed"); + + assert_eq!( + outcome + .candidates + .iter() + .map(|item| item.candidate.as_str()) + .collect::>(), + vec!["key-old", "key-recent"] + ); + } + + #[test] + fn legacy_free_team_first_preserves_all_modes() { + let scheduled = |mode: &str, key_plans: &[(&str, &str)]| { + let candidates = key_plans + .iter() + .map(|(key_id, plan)| { + sample_candidate("provider-pool", "endpoint-1", key_id, 10, true) + .with_presets(vec![PoolSchedulingPreset { + preset: "free_team_first".to_string(), + enabled: true, + mode: Some(mode.to_string()), + }]) + .with_plan(plan) + }) + .collect::>(); + run_pool_scheduler(candidates, &BTreeMap::new(), "seed") + .candidates + .into_iter() + .map(|item| item.candidate) + .collect::>() + }; + + assert_eq!( + scheduled("free_only", &[("key-team", "team"), ("key-free", "free")]), + ["key-free", "key-team"] + ); + assert_eq!( + scheduled("team_only", &[("key-free", "free"), ("key-team", "team")]), + ["key-team", "key-free"] + ); + assert_eq!( + scheduled( + "both", + &[ + ("key-plus", "plus"), + ("key-team", "team"), + ("key-free", "free") + ], + ), + ["key-team", "key-free", "key-plus"] + ); + } + + #[test] + fn strategy_ties_use_lru_as_the_final_tiebreaker() { + let recent_free = + sample_candidate("provider-pool", "endpoint-1", "key-recent-free", 10, true) + .with_presets(vec![PoolSchedulingPreset { + preset: "free_first".to_string(), + enabled: true, + mode: None, + }]) + .with_plan("free"); + let plus = sample_candidate("provider-pool", "endpoint-1", "key-plus", 10, true) + .with_presets(vec![PoolSchedulingPreset { + preset: "free_first".to_string(), + enabled: true, + mode: None, + }]) + .with_plan("plus"); + let old_free = sample_candidate("provider-pool", "endpoint-1", "key-old-free", 10, true) + .with_presets(vec![PoolSchedulingPreset { + preset: "free_first".to_string(), + enabled: true, + mode: None, + }]) + .with_plan("free"); + let runtime_by_provider = BTreeMap::from([( + "provider-pool".to_string(), + PoolRuntimeState { + lru_score_by_key: BTreeMap::from([ + ("key-recent-free".to_string(), 200.0), + ("key-old-free".to_string(), 10.0), + ("key-plus".to_string(), 0.0), + ]), + ..PoolRuntimeState::default() + }, + )]); + + let outcome = run_pool_scheduler( + vec![recent_free, plus, old_free], + &runtime_by_provider, + "seed", + ); + + assert!(outcome.skipped_candidates.is_empty()); + assert_eq!( + outcome + .candidates + .iter() + .map(|item| item.candidate.as_str()) + .collect::>(), + ["key-old-free", "key-recent-free", "key-plus"] + ); + } + #[test] fn pool_scheduler_applies_distribution_mode_before_strategy_presets() { let key_cache_hit = @@ -1276,6 +1501,7 @@ mod tests { fn with_cost_limit(self, limit: u64) -> Self; fn with_presets(self, presets: Vec) -> Self; fn with_plan(self, plan: &str) -> Self; + fn with_quota_usage(self, ratio: f64) -> Self; } impl TestCandidateExt for PoolCandidateInput { @@ -1297,5 +1523,10 @@ mod tests { self.key_context.plan_tier = Some(plan.to_string()); self } + + fn with_quota_usage(mut self, ratio: f64) -> Self { + self.key_context.quota_usage_ratio = Some(ratio); + self + } } } diff --git a/crates/aether-provider/pool/src/lib.rs b/crates/aether-provider/pool/src/lib.rs index 6ac735d5a..46d9dae20 100644 --- a/crates/aether-provider/pool/src/lib.rs +++ b/crates/aether-provider/pool/src/lib.rs @@ -604,6 +604,10 @@ mod tests { .iter() .find(|item| item["name"] == "recent_refresh") .expect("recent_refresh should exist"); + let legacy_free_team = items + .iter() + .find(|item| item["name"] == "free_team_first") + .expect("legacy free_team_first should remain configurable"); assert_eq!( free_first["providers"], @@ -613,6 +617,14 @@ mod tests { recent_refresh["providers"], json!(["codex", "grok", "kiro", "windsurf"]) ); + assert_eq!(free_first["default_enabled"], json!(false)); + assert_eq!(recent_refresh["default_enabled"], json!(false)); + assert_eq!( + recent_refresh["default_enabled_providers"], + json!(["codex", "windsurf"]) + ); + assert_eq!(legacy_free_team["default_mode"], json!("both")); + assert_eq!(legacy_free_team["modes"].as_array().map(Vec::len), Some(3)); } #[test] @@ -657,6 +669,17 @@ mod tests { ["cache_affinity", "recent_refresh"] ); + let legacy_free_team = service.normalize_scheduling_presets( + "codex", + &[PoolSchedulingPreset { + preset: "free_team_first".to_string(), + enabled: true, + mode: Some("team_only".to_string()), + }], + ); + assert_eq!(legacy_free_team[0].preset, "free_team_first"); + assert_eq!(legacy_free_team[0].mode.as_deref(), Some("team_only")); + let unsupported = service.normalize_scheduling_presets( "chatgpt_web", &[PoolSchedulingPreset { diff --git a/crates/aether-provider/pool/src/presets.rs b/crates/aether-provider/pool/src/presets.rs index 19fc5ce42..a3302d428 100644 --- a/crates/aether-provider/pool/src/presets.rs +++ b/crates/aether-provider/pool/src/presets.rs @@ -110,6 +110,7 @@ pub fn build_admin_pool_scheduling_presets_payload() -> Value { "依据窗口成本/Token 用量,缺失时回退配额使用率", &service, ), + legacy_free_team_first_preset_payload(&service), provider_pool_preset_payload( "free_first", "Free 优先", @@ -201,6 +202,24 @@ pub fn build_admin_pool_scheduling_presets_payload() -> Value { ]) } +fn legacy_free_team_first_preset_payload(service: &ProviderPoolService) -> Value { + let mut payload = provider_pool_preset_payload( + "free_team_first", + "Free/Team 优先", + "兼容旧配置:优先消耗 Free、Team 或两者", + Some(ProviderPoolCapability::PlanTier), + "依据 plan_type,保留旧 free_only/team_only/both 语义", + service, + ); + payload["modes"] = json!([ + {"value": "free_only", "label": "Free"}, + {"value": "team_only", "label": "Team"}, + {"value": "both", "label": "Free + Team"} + ]); + payload["default_mode"] = json!("both"); + payload +} + fn provider_pool_preset_payload( name: &'static str, label: &'static str, @@ -212,11 +231,24 @@ fn provider_pool_preset_payload( let providers = capability .map(|capability| service.provider_types_for_capability(capability)) .unwrap_or_default(); + let default_enabled_providers = service + .provider_types() + .filter(|provider_type| { + service + .adapter(provider_type) + .default_scheduling_presets() + .iter() + .any(|preset| preset.enabled && preset.preset.eq_ignore_ascii_case(name)) + }) + .map(str::to_string) + .collect::>(); json!({ "name": name, "label": label, "description": description, "providers": providers, + "default_enabled": name == "cache_affinity", + "default_enabled_providers": default_enabled_providers, "modes": Value::Null, "default_mode": Value::Null, "mutex_group": provider_pool_preset_mutex_group(name), @@ -226,7 +258,7 @@ fn provider_pool_preset_payload( fn provider_pool_supports_preset(adapter: &dyn ProviderPoolAdapter, preset: &str) -> bool { match preset { - "free_first" | "plus_first" | "pro_first" | "team_first" => adapter + "free_first" | "free_team_first" | "plus_first" | "pro_first" | "team_first" => adapter .capabilities() .supports(ProviderPoolCapability::PlanTier), "recent_refresh" => adapter diff --git a/crates/aether-provider/transport/src/agent_identity.rs b/crates/aether-provider/transport/src/agent_identity.rs index 28f5ca3ce..82313e77a 100644 --- a/crates/aether-provider/transport/src/agent_identity.rs +++ b/crates/aether-provider/transport/src/agent_identity.rs @@ -695,6 +695,7 @@ async fn create_codex_agent_identity_from_session_token_with_auth_api_base_url( })), body_bytes: None, network, + transport_profile: None, }) .await .map_err(|_| CodexAgentIdentityEnrollmentError::TaskRegistrationRequestFailed)?; @@ -772,6 +773,7 @@ async fn register_codex_agent_identity_from_access_token_with_auth_api_base_url( })), body_bytes: None, network, + transport_profile: None, }) .await .map_err(|_| CodexAgentIdentityEnrollmentError::RegistrationRequestFailed)?; diff --git a/crates/aether-provider/transport/src/provider_types.rs b/crates/aether-provider/transport/src/provider_types.rs index 0a85b477d..5f7935c1d 100644 --- a/crates/aether-provider/transport/src/provider_types.rs +++ b/crates/aether-provider/transport/src/provider_types.rs @@ -540,13 +540,13 @@ pub fn provider_type_admin_oauth_template(provider_type: &str) -> Option Some(ProviderOAuthTemplate { provider_type: "claude_code", - display_name: "ClaudeCode", - authorize_url: "https://claude.ai/oauth/authorize", - token_url: "https://console.anthropic.com/v1/oauth/token", - client_id: "9d1c250a-e61b-44d9-88ed-5944d1962f5e", + display_name: "Claude Code", + authorize_url: aether_oauth::provider::providers::CLAUDE_CODE_AUTHORIZE_URL, + token_url: aether_oauth::provider::providers::CLAUDE_CODE_TOKEN_URL, + client_id: aether_oauth::provider::providers::CLAUDE_CODE_CLIENT_ID, client_secret: "", - scopes: &["org:create_api_key", "user:profile", "user:inference"], - redirect_uri: "http://localhost:54545/callback", + scopes: aether_oauth::provider::providers::CLAUDE_CODE_OAUTH_SCOPES, + redirect_uri: aether_oauth::provider::providers::CLAUDE_CODE_REDIRECT_URI, use_pkce: true, }), "codex" => Some(ProviderOAuthTemplate { diff --git a/crates/aether-routing-core/src/lib.rs b/crates/aether-routing-core/src/lib.rs index f2d151355..740a4dac9 100644 --- a/crates/aether-routing-core/src/lib.rs +++ b/crates/aether-routing-core/src/lib.rs @@ -33,4 +33,6 @@ pub use trace::{ RoutingCandidateTrace, RoutingDecisionTrace, RoutingPatchSummary, RoutingPoolExpansionTrace, RoutingRuntimeFacts, }; -pub use validation::{validate_routing_group_config, RoutingValidationError}; +pub use validation::{ + validate_routing_group_config, RoutingValidationError, MAX_ROUTING_ALLOWED_KEYS, +}; diff --git a/crates/aether-routing-core/src/policy.rs b/crates/aether-routing-core/src/policy.rs index dbcbb34ec..1a784cc6e 100644 --- a/crates/aether-routing-core/src/policy.rs +++ b/crates/aether-routing-core/src/policy.rs @@ -11,6 +11,7 @@ use crate::conditions::RoutingConditionContext; use crate::model::{RoutingGroupConfig, RoutingModelPolicy, RoutingPoolPolicyOverride}; use crate::mutations::{validate_header_patch, validate_json_patch_operations, MutationPlan}; use crate::ranking::RankingOverlay; +use crate::validation::validate_routing_group_config; #[derive(Debug, Error, Clone, PartialEq, Eq)] pub enum RoutingPolicyError { @@ -68,6 +69,9 @@ pub fn resolve_routing_policy( config: &RoutingGroupConfig, input: RoutingPolicyInput<'_>, ) -> Result { + validate_routing_group_config(config) + .map_err(|error| RoutingPolicyError::InvalidConfig(error.to_string()))?; + if !model_allowed(&config.allowed_models, input.requested_model) && !model_allowed(&config.allowed_models, input.resolved_model) { diff --git a/crates/aether-routing-core/src/ranking.rs b/crates/aether-routing-core/src/ranking.rs index 05086c150..a8bc5b867 100644 --- a/crates/aether-routing-core/src/ranking.rs +++ b/crates/aether-routing-core/src/ranking.rs @@ -40,6 +40,13 @@ impl RankingOverlay { .unwrap_or(fallback) } + pub fn pool_priority(&self, provider_id: &str, fallback: i32) -> i32 { + self.pool_priority_overrides + .get(provider_id) + .copied() + .unwrap_or(fallback) + } + pub fn provider_priority_or_unspecified(&self, provider_id: &str) -> i32 { self.provider_priority_overrides .get(provider_id) @@ -100,15 +107,18 @@ pub fn rank_vector_for_candidate( ) -> RoutingCandidateRankVector { RoutingCandidateRankVector { provider_priority_before: facts.provider_priority, - provider_priority_after: overlay.provider_priority_or_unspecified(&facts.provider_id), + provider_priority_after: overlay + .provider_priority(&facts.provider_id, facts.provider_priority), key_priority_before: facts.key_priority, key_priority_after: match facts.candidate_kind { CandidateKind::Provider => facts .key_id .as_deref() - .map(|key_id| overlay.key_priority_or_unspecified(key_id)) - .unwrap_or(ROUTING_PRIORITY_UNSPECIFIED), - CandidateKind::PoolGroup => overlay.pool_priority_or_unspecified(&facts.provider_id), + .map(|key_id| overlay.key_priority(key_id, facts.key_priority)) + .unwrap_or(facts.key_priority), + CandidateKind::PoolGroup => { + overlay.pool_priority(&facts.provider_id, facts.key_priority) + } }, } } @@ -142,7 +152,7 @@ mod tests { } #[test] - fn rank_vector_marks_missing_routing_priorities_unspecified() { + fn rank_vector_falls_back_to_existing_priorities() { let facts = RoutingCandidateFacts { candidate_kind: CandidateKind::Provider, provider_id: "provider-a".to_string(), @@ -154,8 +164,8 @@ mod tests { }; let vector = rank_vector_for_candidate(&RankingOverlay::default(), &facts); - assert_eq!(vector.provider_priority_after, ROUTING_PRIORITY_UNSPECIFIED); - assert_eq!(vector.key_priority_after, ROUTING_PRIORITY_UNSPECIFIED); + assert_eq!(vector.provider_priority_after, 10); + assert_eq!(vector.key_priority_after, 20); } #[test] diff --git a/crates/aether-routing-core/src/validation.rs b/crates/aether-routing-core/src/validation.rs index 9ce88d85c..0ac3123be 100644 --- a/crates/aether-routing-core/src/validation.rs +++ b/crates/aether-routing-core/src/validation.rs @@ -2,9 +2,29 @@ use std::collections::BTreeSet; use thiserror::Error; -use crate::model::RoutingGroupConfig; +use crate::model::{RoutingGroupConfig, RoutingPoolPolicyOverride}; use crate::mutations::{validate_header_patch, validate_json_patch_operations}; -use crate::RoutingAction; +use crate::{RoutingAction, RoutingRulePhase}; + +pub const MAX_ROUTING_ALLOWED_KEYS: usize = 512; + +const ROUTING_POOL_PRESETS: &[&str] = &[ + "lru", + "cache_affinity", + "load_balance", + "single_account", + "priority_first", + "free_team_first", + "free_first", + "team_first", + "plus_first", + "pro_first", + "health_first", + "latency_first", + "cost_first", + "quota_balanced", + "recent_refresh", +]; #[derive(Debug, Error, Clone, PartialEq, Eq)] pub enum RoutingValidationError { @@ -16,6 +36,34 @@ pub enum RoutingValidationError { EmptyModelSelector, #[error("invalid mutation action: {0}")] InvalidMutation(String), + #[error("routing rule {rule_id} uses unsupported {action} action in provider_request phase")] + ProviderRequestActionNotAllowed { + rule_id: String, + action: &'static str, + }, + #[error("routing key selector {selector} contains {count} entries; maximum is {max}")] + TooManyAllowedKeys { + selector: String, + count: usize, + max: usize, + }, + #[error("routing pool policy {selector} has an empty provider id")] + EmptyPoolProviderId { selector: String }, + #[error("routing pool policy {selector} uses unsupported preset: {preset}")] + UnsupportedPoolPreset { selector: String, preset: String }, + #[error("routing pool policy {selector} contains duplicate preset: {preset}")] + DuplicatePoolPreset { selector: String, preset: String }, + #[error("routing pool policy {selector} preset {preset} has invalid mode: {mode}")] + InvalidPoolPresetMode { + selector: String, + preset: String, + mode: String, + }, + #[error("routing pool policy {selector} enables mutually exclusive distribution presets: {presets:?}")] + ConflictingPoolDistributionPresets { + selector: String, + presets: Vec, + }, } pub fn validate_routing_group_config( @@ -26,6 +74,22 @@ pub fn validate_routing_group_config( if model_policy.model.trim().is_empty() { return Err(RoutingValidationError::EmptyModelSelector); } + validate_allowed_key_count( + format!("model:{}", model_policy.model.trim()), + model_policy.allowed_keys.len(), + )?; + for (provider_id, override_policy) in &model_policy.pool_policy_overrides { + let model_selector = format!("model:{}", model_policy.model.trim()); + if provider_id.trim().is_empty() { + return Err(RoutingValidationError::EmptyPoolProviderId { + selector: model_selector, + }); + } + validate_pool_policy_override( + format!("{model_selector}:provider:{}", provider_id.trim()), + override_policy, + )?; + } } for rule in &config.rules { if rule.id.trim().is_empty() { @@ -35,6 +99,17 @@ pub fn validate_routing_group_config( return Err(RoutingValidationError::DuplicateRuleId(rule.id.clone())); } for action in &rule.actions { + if rule.phase == RoutingRulePhase::ProviderRequest + && !matches!( + action, + RoutingAction::JsonPatchBody { .. } | RoutingAction::PatchHeaders { .. } + ) + { + return Err(RoutingValidationError::ProviderRequestActionNotAllowed { + rule_id: rule.id.clone(), + action: routing_action_name(action), + }); + } match action { RoutingAction::JsonPatchBody { patch } => { validate_json_patch_operations(patch).map_err(|error| { @@ -46,9 +121,372 @@ pub fn validate_routing_group_config( RoutingValidationError::InvalidMutation(error.to_string()) })?; } + RoutingAction::RestrictKeys { key_ids } => { + validate_allowed_key_count(format!("rule:{}", rule.id), key_ids.len())?; + } _ => {} } } } Ok(()) } + +fn validate_pool_policy_override( + selector: String, + override_policy: &RoutingPoolPolicyOverride, +) -> Result<(), RoutingValidationError> { + let mut seen = BTreeSet::new(); + let mut enabled_distribution_presets = Vec::new(); + for preset in &override_policy.scheduling_presets { + let normalized = preset.preset.trim().to_ascii_lowercase(); + if !ROUTING_POOL_PRESETS.contains(&normalized.as_str()) { + return Err(RoutingValidationError::UnsupportedPoolPreset { + selector, + preset: normalized, + }); + } + if !seen.insert(normalized.clone()) { + return Err(RoutingValidationError::DuplicatePoolPreset { + selector, + preset: normalized, + }); + } + if let Some(mode) = preset.mode.as_deref() { + let mode = mode.trim().to_ascii_lowercase(); + if !routing_pool_preset_mode_valid(&normalized, &mode) { + return Err(RoutingValidationError::InvalidPoolPresetMode { + selector, + preset: normalized, + mode, + }); + } + } + if preset.enabled && routing_pool_distribution_preset(&normalized) { + enabled_distribution_presets.push(normalized); + } + } + if enabled_distribution_presets.len() > 1 { + return Err(RoutingValidationError::ConflictingPoolDistributionPresets { + selector, + presets: enabled_distribution_presets, + }); + } + Ok(()) +} + +fn routing_pool_preset_mode_valid(preset: &str, mode: &str) -> bool { + match preset { + "free_team_first" => matches!(mode, "free_only" | "team_only" | "both"), + "free_first" => mode == "free_only", + "team_first" => mode == "team_only", + "plus_first" => mode == "plus_only", + "pro_first" => mode == "pro_only", + _ => false, + } +} + +fn routing_pool_distribution_preset(preset: &str) -> bool { + matches!( + preset, + "lru" | "cache_affinity" | "load_balance" | "single_account" + ) +} + +fn validate_allowed_key_count( + selector: String, + count: usize, +) -> Result<(), RoutingValidationError> { + if count > MAX_ROUTING_ALLOWED_KEYS { + return Err(RoutingValidationError::TooManyAllowedKeys { + selector, + count, + max: MAX_ROUTING_ALLOWED_KEYS, + }); + } + Ok(()) +} + +fn routing_action_name(action: &RoutingAction) -> &'static str { + match action { + RoutingAction::RestrictModels { .. } => "restrict_models", + RoutingAction::RestrictProviders { .. } => "restrict_providers", + RoutingAction::RestrictKeys { .. } => "restrict_keys", + RoutingAction::SetScheduling { .. } => "set_scheduling", + RoutingAction::SetProviderPriority { .. } => "set_provider_priority", + RoutingAction::SetKeyPriority { .. } => "set_key_priority", + RoutingAction::JsonPatchBody { .. } => "json_patch_body", + RoutingAction::PatchHeaders { .. } => "patch_headers", + } +} + +#[cfg(test)] +mod tests { + use serde_json::json; + + use crate::{ + RoutingCondition, RoutingGroupConfig, RoutingHeaderPatch, RoutingJsonPatchOperation, + RoutingModelPolicy, RoutingRule, + }; + + use super::*; + + fn provider_request_config(actions: Vec) -> RoutingGroupConfig { + RoutingGroupConfig { + rules: vec![RoutingRule { + id: "provider-rule".to_string(), + priority: 0, + enabled: true, + phase: RoutingRulePhase::ProviderRequest, + conditions: RoutingCondition::default(), + actions, + stop_processing: false, + }], + ..RoutingGroupConfig::default() + } + } + + #[test] + fn rejects_non_mutation_actions_in_provider_request_phase() { + let actions = [ + ( + RoutingAction::RestrictModels { + models: vec!["gpt-5".to_string()], + }, + "restrict_models", + ), + ( + RoutingAction::RestrictProviders { + provider_ids: vec!["provider-1".to_string()], + }, + "restrict_providers", + ), + ( + RoutingAction::RestrictKeys { + key_ids: vec!["key-1".to_string()], + }, + "restrict_keys", + ), + ( + RoutingAction::SetScheduling { + priority_mode: None, + scheduling_mode: None, + keep_priority_on_conversion: Some(true), + }, + "set_scheduling", + ), + ( + RoutingAction::SetProviderPriority { + provider_id: "provider-1".to_string(), + priority: 1, + }, + "set_provider_priority", + ), + ( + RoutingAction::SetKeyPriority { + key_id: "key-1".to_string(), + priority: 1, + }, + "set_key_priority", + ), + ]; + + for (action, action_name) in actions { + let error = validate_routing_group_config(&provider_request_config(vec![action])) + .expect_err("provider_request must only accept mutations"); + assert_eq!( + error, + RoutingValidationError::ProviderRequestActionNotAllowed { + rule_id: "provider-rule".to_string(), + action: action_name, + } + ); + } + } + + #[test] + fn accepts_mutation_actions_in_provider_request_phase() { + let config = provider_request_config(vec![ + RoutingAction::JsonPatchBody { + patch: vec![RoutingJsonPatchOperation::Add { + path: "/metadata/routed".to_string(), + value: json!(true), + }], + }, + RoutingAction::PatchHeaders { + patch: vec![RoutingHeaderPatch::Set { + name: "x-routing-profile".to_string(), + value: "provider-rule".to_string(), + }], + }, + ]); + + validate_routing_group_config(&config) + .expect("provider_request mutation actions should remain valid"); + } + + #[test] + fn rejects_oversized_model_allowed_key_selector() { + let config = RoutingGroupConfig { + model_policies: vec![RoutingModelPolicy { + model: "gpt-5".to_string(), + allowed_keys: key_ids(MAX_ROUTING_ALLOWED_KEYS + 1), + ..RoutingModelPolicy::default() + }], + ..RoutingGroupConfig::default() + }; + + assert_eq!( + validate_routing_group_config(&config), + Err(RoutingValidationError::TooManyAllowedKeys { + selector: "model:gpt-5".to_string(), + count: MAX_ROUTING_ALLOWED_KEYS + 1, + max: MAX_ROUTING_ALLOWED_KEYS, + }) + ); + } + + #[test] + fn rejects_oversized_rule_allowed_key_selector() { + let config = RoutingGroupConfig { + rules: vec![RoutingRule { + id: "restrict-keys".to_string(), + priority: 0, + enabled: true, + phase: RoutingRulePhase::ClientRequest, + conditions: RoutingCondition::default(), + actions: vec![RoutingAction::RestrictKeys { + key_ids: key_ids(MAX_ROUTING_ALLOWED_KEYS + 1), + }], + stop_processing: false, + }], + ..RoutingGroupConfig::default() + }; + + assert_eq!( + validate_routing_group_config(&config), + Err(RoutingValidationError::TooManyAllowedKeys { + selector: "rule:restrict-keys".to_string(), + count: MAX_ROUTING_ALLOWED_KEYS + 1, + max: MAX_ROUTING_ALLOWED_KEYS, + }) + ); + } + + #[test] + fn accepts_allowed_key_selectors_at_the_limit() { + let config = RoutingGroupConfig { + model_policies: vec![RoutingModelPolicy { + model: "gpt-5".to_string(), + allowed_keys: key_ids(MAX_ROUTING_ALLOWED_KEYS), + ..RoutingModelPolicy::default() + }], + ..RoutingGroupConfig::default() + }; + + validate_routing_group_config(&config) + .expect("allowed key selectors at the scan limit should remain valid"); + } + + #[test] + fn rejects_unknown_pool_override_preset() { + let config = pool_override_config(vec![pool_preset("typo_priority", true, None)]); + + assert_eq!( + validate_routing_group_config(&config), + Err(RoutingValidationError::UnsupportedPoolPreset { + selector: "model:gpt-5:provider:provider-1".to_string(), + preset: "typo_priority".to_string(), + }) + ); + } + + #[test] + fn rejects_duplicate_pool_override_presets() { + let config = pool_override_config(vec![ + pool_preset("health_first", true, None), + pool_preset(" HEALTH_FIRST ", false, None), + ]); + + assert_eq!( + validate_routing_group_config(&config), + Err(RoutingValidationError::DuplicatePoolPreset { + selector: "model:gpt-5:provider:provider-1".to_string(), + preset: "health_first".to_string(), + }) + ); + } + + #[test] + fn rejects_invalid_pool_override_mode() { + let config = + pool_override_config(vec![pool_preset("free_team_first", true, Some("pro_only"))]); + + assert_eq!( + validate_routing_group_config(&config), + Err(RoutingValidationError::InvalidPoolPresetMode { + selector: "model:gpt-5:provider:provider-1".to_string(), + preset: "free_team_first".to_string(), + mode: "pro_only".to_string(), + }) + ); + } + + #[test] + fn rejects_conflicting_pool_override_distribution_modes() { + let config = pool_override_config(vec![ + pool_preset("lru", true, None), + pool_preset("cache_affinity", true, None), + ]); + + assert_eq!( + validate_routing_group_config(&config), + Err(RoutingValidationError::ConflictingPoolDistributionPresets { + selector: "model:gpt-5:provider:provider-1".to_string(), + presets: vec!["lru".to_string(), "cache_affinity".to_string()], + }) + ); + } + + #[test] + fn accepts_valid_pool_override_modes_and_disabled_alternatives() { + let config = pool_override_config(vec![ + pool_preset("cache_affinity", true, None), + pool_preset("lru", false, None), + pool_preset("free_team_first", true, Some("team_only")), + ]); + + validate_routing_group_config(&config).expect("valid pool override should pass"); + } + + fn pool_override_config( + scheduling_presets: Vec, + ) -> RoutingGroupConfig { + RoutingGroupConfig { + model_policies: vec![RoutingModelPolicy { + model: "gpt-5".to_string(), + pool_policy_overrides: std::collections::BTreeMap::from([( + "provider-1".to_string(), + RoutingPoolPolicyOverride { scheduling_presets }, + )]), + ..RoutingModelPolicy::default() + }], + ..RoutingGroupConfig::default() + } + } + + fn pool_preset( + preset: &str, + enabled: bool, + mode: Option<&str>, + ) -> crate::RoutingSchedulingPreset { + crate::RoutingSchedulingPreset { + preset: preset.to_string(), + enabled, + mode: mode.map(str::to_string), + } + } + + fn key_ids(count: usize) -> Vec { + (0..count).map(|index| format!("key-{index}")).collect() + } +} diff --git a/crates/aether-scheduler-core/src/affinity.rs b/crates/aether-scheduler-core/src/affinity.rs index 6b49b5914..1b8a25f69 100644 --- a/crates/aether-scheduler-core/src/affinity.rs +++ b/crates/aether-scheduler-core/src/affinity.rs @@ -1,3 +1,4 @@ +use serde::{Deserialize, Serialize}; use sha2::{Digest, Sha256}; use crate::SchedulerMinimalCandidateSelectionCandidate; @@ -9,6 +10,26 @@ pub struct SchedulerAffinityTarget { pub key_id: String, } +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct SchedulerAffinityScope { + pub routing_group_id: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub routing_group_version: Option, +} + +impl SchedulerAffinityScope { + pub fn new(routing_group_id: impl Into, routing_group_version: Option) -> Self { + Self { + routing_group_id: routing_group_id.into(), + routing_group_version, + } + } + + pub fn is_valid(&self) -> bool { + !self.routing_group_id.trim().is_empty() + } +} + #[derive(Debug, Clone, Default, PartialEq, Eq)] pub struct ClientSessionAffinity { pub client_family: Option, @@ -56,6 +77,22 @@ pub fn build_scheduler_affinity_cache_key_for_api_key_id_with_client_session( api_format: &str, global_model_name: &str, client_session_affinity: Option<&ClientSessionAffinity>, +) -> Option { + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope( + api_key_id, + api_format, + global_model_name, + client_session_affinity, + None, + ) +} + +pub fn build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope( + api_key_id: &str, + api_format: &str, + global_model_name: &str, + client_session_affinity: Option<&ClientSessionAffinity>, + affinity_scope: Option<&SchedulerAffinityScope>, ) -> Option { let api_key_id = api_key_id.trim(); if api_key_id.is_empty() { @@ -67,36 +104,63 @@ pub fn build_scheduler_affinity_cache_key_for_api_key_id_with_client_session( return None; } - let legacy_key = format!("scheduler_affinity:{api_key_id}:{api_format}:{global_model_name}"); - let Some(client_session_affinity) = client_session_affinity else { - return Some(legacy_key); - }; - let Some(session_key) = client_session_affinity - .session_key - .as_deref() + let affinity_scope = affinity_scope.filter(|scope| scope.is_valid()); + let session_key = client_session_affinity + .and_then(|affinity| affinity.session_key.as_deref()) .map(str::trim) - .filter(|value| !value.is_empty()) - else { - return Some(legacy_key); - }; + .filter(|value| !value.is_empty()); + if session_key.is_none() && affinity_scope.is_none() { + return Some(format!( + "scheduler_affinity:{api_key_id}:{api_format}:{global_model_name}" + )); + } let client_family = client_session_affinity - .client_family - .as_deref() + .and_then(|affinity| affinity.client_family.as_deref()) .map(str::trim) .filter(|value| !value.is_empty()) .map(str::to_ascii_lowercase) .unwrap_or_else(|| "generic".to_string()); - let session_hash = hash_session_key(session_key); + let session_hash = match affinity_scope { + Some(scope) => hash_scoped_session_key(session_key, scope), + None => hash_session_key(session_key.expect("session key should exist without scope")), + }; Some(format!( "scheduler_affinity:v2:{api_key_id}:{api_format}:{global_model_name}:{client_family}:{session_hash}" )) } +fn hash_scoped_session_key( + session_key: Option<&str>, + affinity_scope: &SchedulerAffinityScope, +) -> String { + let mut hasher = Sha256::new(); + hasher.update(b"scheduler-affinity-scope-v1\0"); + hasher.update(affinity_scope.routing_group_id.trim().as_bytes()); + hasher.update(b"\0"); + match affinity_scope.routing_group_version { + Some(version) => hasher.update(version.to_string().as_bytes()), + None => hasher.update(b"unversioned"), + } + hasher.update(b"\0"); + match session_key { + Some(session_key) => { + hasher.update(b"session\0"); + hasher.update(session_key.as_bytes()); + } + None => hasher.update(b"api-key-scope"), + } + hex_digest(hasher.finalize()) +} + fn hash_session_key(session_key: &str) -> String { let mut hasher = Sha256::new(); hasher.update(session_key.as_bytes()); - let digest = hasher.finalize(); + hex_digest(hasher.finalize()) +} + +fn hex_digest(digest: impl AsRef<[u8]>) -> String { + let digest = digest.as_ref(); digest.iter().map(|byte| format!("{byte:02x}")).collect() } @@ -142,8 +206,9 @@ mod tests { use super::{ build_scheduler_affinity_cache_key_for_api_key_id, build_scheduler_affinity_cache_key_for_api_key_id_with_client_session, + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope, candidate_affinity_hash, candidate_key, matches_affinity_target, ClientSessionAffinity, - SchedulerAffinityTarget, + SchedulerAffinityScope, SchedulerAffinityTarget, }; use crate::SchedulerMinimalCandidateSelectionCandidate; @@ -262,6 +327,47 @@ mod tests { assert_ne!(left_key, other_client_key); } + #[test] + fn routing_scoped_affinity_keys_split_groups_and_versions() { + let affinity = + ClientSessionAffinity::new(Some("generic".to_string()), Some("session-a".to_string())); + let group_one_v1 = SchedulerAffinityScope::new("group-1", Some(1)); + let group_one_v2 = SchedulerAffinityScope::new("group-1", Some(2)); + let group_two_v1 = SchedulerAffinityScope::new("group-2", Some(1)); + let key_for = |scope: &SchedulerAffinityScope| { + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope( + "api-key-1", + "openai:chat", + "gpt-5", + Some(&affinity), + Some(scope), + ) + .expect("scoped affinity key should build") + }; + + assert_ne!(key_for(&group_one_v1), key_for(&group_one_v2)); + assert_ne!(key_for(&group_one_v1), key_for(&group_two_v1)); + } + + #[test] + fn routing_scoped_affinity_key_isolated_without_client_session() { + let group_one = SchedulerAffinityScope::new("group-1", Some(1)); + let group_two = SchedulerAffinityScope::new("group-2", Some(1)); + let key_for = |scope: &SchedulerAffinityScope| { + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope( + "api-key-1", + "openai:chat", + "gpt-5", + None, + Some(scope), + ) + .expect("scoped affinity key should build") + }; + + assert!(key_for(&group_one).starts_with("scheduler_affinity:v2:")); + assert_ne!(key_for(&group_one), key_for(&group_two)); + } + #[test] fn affinity_hash_is_candidate_specific() { let left = sample_candidate("1"); diff --git a/crates/aether-scheduler-core/src/lib.rs b/crates/aether-scheduler-core/src/lib.rs index 3fbfe0498..55198d60f 100644 --- a/crates/aether-scheduler-core/src/lib.rs +++ b/crates/aether-scheduler-core/src/lib.rs @@ -9,8 +9,10 @@ mod request_candidate; pub use affinity::{ build_scheduler_affinity_cache_key_for_api_key_id, - build_scheduler_affinity_cache_key_for_api_key_id_with_client_session, candidate_affinity_hash, - candidate_key, matches_affinity_target, ClientSessionAffinity, SchedulerAffinityTarget, + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session, + build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope, + candidate_affinity_hash, candidate_key, matches_affinity_target, ClientSessionAffinity, + SchedulerAffinityScope, SchedulerAffinityTarget, }; pub use auth::{ api_format_matches_allowed_value, auth_constraints_allow_api_format, diff --git a/frontend/src/api/__tests__/provider-oauth-agent-identity-routes.spec.ts b/frontend/src/api/__tests__/provider-oauth-agent-identity-routes.spec.ts index c3ac3cc4c..48a2e4646 100644 --- a/frontend/src/api/__tests__/provider-oauth-agent-identity-routes.spec.ts +++ b/frontend/src/api/__tests__/provider-oauth-agent-identity-routes.spec.ts @@ -13,12 +13,17 @@ vi.mock('@/api/client', () => ({ })) import { + authorizeProviderWithCookie, + getProviderCookieAuthorizeTaskStatus, getBatchImportOAuthTaskStatus, importProviderRefreshToken, + startProviderCookieAuthorizeTask, startBatchImportOAuthTask, } from '@/api/endpoints/provider_oauth' -describe('Agent Identity OAuth management routes', () => { +const CLAUDE_COOKIE_AUTHORIZE_TIMEOUT_MS = 4 * 60 * 1000 + +describe('Provider OAuth management routes', () => { beforeEach(() => { getMock.mockReset() postMock.mockReset() @@ -137,4 +142,37 @@ describe('Agent Identity OAuth management routes', () => { }, ) }) + + it('allows the sequential Claude Cookie exchange to outlive the global timeout', async () => { + const payload = { + cookie: 'sessionKey=claude-session-key', + proxy_node_id: 'proxy-1', + } + + await authorizeProviderWithCookie('provider-claude', payload) + + expect(postMock).toHaveBeenCalledWith( + '/api/admin/provider-oauth/providers/provider-claude/cookie-authorize', + payload, + { timeout: CLAUDE_COOKIE_AUTHORIZE_TIMEOUT_MS }, + ) + }) + + it('starts and polls Claude Cookie batches through the dedicated task routes', async () => { + const payload = { + cookies: ['sessionKey=claude-1', 'sessionKey=claude-2'], + proxy_node_id: 'proxy-1', + } + + await startProviderCookieAuthorizeTask('provider-claude', payload) + await getProviderCookieAuthorizeTaskStatus('provider-claude', 'claude-cookie-task-1') + + expect(postMock).toHaveBeenCalledWith( + '/api/admin/provider-oauth/providers/provider-claude/cookie-authorize/tasks', + payload, + ) + expect(getMock).toHaveBeenCalledWith( + '/api/admin/provider-oauth/providers/provider-claude/cookie-authorize/tasks/claude-cookie-task-1', + ) + }) }) diff --git a/frontend/src/api/endpoints/pool.ts b/frontend/src/api/endpoints/pool.ts index 2926164f2..82621013e 100644 --- a/frontend/src/api/endpoints/pool.ts +++ b/frontend/src/api/endpoints/pool.ts @@ -103,6 +103,8 @@ export interface PoolPresetMeta { label: string description: string providers: string[] + default_enabled?: boolean + default_enabled_providers?: string[] modes?: PoolPresetModeMeta[] | null default_mode?: string | null mutex_group?: string | null diff --git a/frontend/src/api/endpoints/provider_oauth.ts b/frontend/src/api/endpoints/provider_oauth.ts index bbd9bab2f..cce12ee9c 100644 --- a/frontend/src/api/endpoints/provider_oauth.ts +++ b/frontend/src/api/endpoints/provider_oauth.ts @@ -1,5 +1,7 @@ import client from '../client' +const CLAUDE_COOKIE_AUTHORIZE_TIMEOUT_MS = 4 * 60 * 1000 + export interface ProviderOAuthStartResponse { authorization_url: string redirect_uri: string @@ -36,6 +38,17 @@ export interface ProviderOAuthCompleteResponseWithKey { detail?: string } +export interface ProviderCookieAuthorizeRequest { + cookie: string + name?: string + proxy_node_id?: string +} + +export interface ProviderCookieAuthorizeBatchTaskRequest { + cookies: string[] + proxy_node_id?: string +} + export interface OAuthBatchImportResultItem { index: number status: 'success' | 'error' @@ -47,9 +60,11 @@ export interface OAuthBatchImportResultItem { } export type OAuthBatchImportTaskStatus = 'submitted' | 'processing' | 'completed' | 'failed' +export type OAuthBatchImportKind = 'oauth_batch' | 'agent_identity' | 'cookie_authorize' export interface OAuthBatchImportTaskStartResponse { task_id: string + import_kind?: OAuthBatchImportKind status: OAuthBatchImportTaskStatus total: number processed: number @@ -63,6 +78,7 @@ export interface OAuthBatchImportTaskStartResponse { export interface OAuthBatchImportTaskStatusResponse { task_id: string + import_kind?: OAuthBatchImportKind provider_id: string provider_type: string status: OAuthBatchImportTaskStatus @@ -233,6 +249,39 @@ export async function completeProviderLevelOAuth( return resp.data } +export async function authorizeProviderWithCookie( + providerId: string, + data: ProviderCookieAuthorizeRequest +): Promise { + const resp = await client.post( + `/api/admin/provider-oauth/providers/${providerId}/cookie-authorize`, + data, + { timeout: CLAUDE_COOKIE_AUTHORIZE_TIMEOUT_MS }, + ) + return resp.data +} + +export async function startProviderCookieAuthorizeTask( + providerId: string, + data: ProviderCookieAuthorizeBatchTaskRequest, +): Promise { + const resp = await client.post( + `/api/admin/provider-oauth/providers/${providerId}/cookie-authorize/tasks`, + data, + ) + return resp.data +} + +export async function getProviderCookieAuthorizeTaskStatus( + providerId: string, + taskId: string, +): Promise { + const resp = await client.get( + `/api/admin/provider-oauth/providers/${providerId}/cookie-authorize/tasks/${taskId}`, + ) + return resp.data +} + export async function importProviderRefreshToken( providerId: string, data: { diff --git a/frontend/src/components/common/JsonImportInput.vue b/frontend/src/components/common/JsonImportInput.vue index 9b1b3dfb7..10ecc508d 100644 --- a/frontend/src/components/common/JsonImportInput.vue +++ b/frontend/src/components/common/JsonImportInput.vue @@ -10,52 +10,65 @@ >
-
-
- -
-
-

- {{ localizedDropTitle }} -

-

- {{ localizedDropHint }} -

+
+
+
+ +
+
+

+ {{ localizedDropTitle }} +

+

+ {{ localizedDropHint }} +

+
-
-
- -