mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-05 00:47:48 +08:00
Merge branch 'fawney19:main' into main
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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<SchedulerAffinityTarget> {
|
||||
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<u64>,
|
||||
) {
|
||||
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<u64>,
|
||||
) {
|
||||
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<u64>,
|
||||
) {
|
||||
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<SchedulerAffinityScope> {
|
||||
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))
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
);
|
||||
}
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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<String, u32>,
|
||||
scanned_rows_by_format: BTreeMap<String, u32>,
|
||||
resolved_global_model_names: BTreeMap<String, String>,
|
||||
fallback_scanned_api_formats: BTreeSet<String>,
|
||||
fallback_offsets: BTreeMap<String, u32>,
|
||||
fallback_scan_epoch: u32,
|
||||
exhausted_api_formats: BTreeSet<String>,
|
||||
seen_candidate_keys: BTreeSet<String>,
|
||||
}
|
||||
@@ -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<RequestedModelCandidateRowsPage, GatewayError> {
|
||||
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::<Vec<_>>();
|
||||
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::<Vec<_>>();
|
||||
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<Vec<StoredApiFormatCandidateRowsQuery>>,
|
||||
}
|
||||
|
||||
impl PagedFallbackRepository {
|
||||
fn new(total_rows: u32) -> Self {
|
||||
Self {
|
||||
total_rows,
|
||||
page_queries: Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
fn page_queries(&self) -> Vec<StoredApiFormatCandidateRowsQuery> {
|
||||
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<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
panic!("routing fallback must not use the unbounded API-format query")
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_page(
|
||||
&self,
|
||||
query: &StoredApiFormatCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, 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<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model(
|
||||
&self,
|
||||
_api_format: &str,
|
||||
_requested_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model_page(
|
||||
&self,
|
||||
_query: &StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group(
|
||||
&self,
|
||||
_query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group_key_ids(
|
||||
&self,
|
||||
_query: &StoredPoolKeyCandidateRowsByKeyIdsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
}
|
||||
|
||||
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::<Vec<_>>();
|
||||
let repository: Arc<dyn MinimalCandidateSelectionReadRepository> =
|
||||
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::<Vec<_>>();
|
||||
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<dyn MinimalCandidateSelectionReadRepository> =
|
||||
|
||||
@@ -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::<Vec<_>>(),
|
||||
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::<Vec<_>>()
|
||||
} 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());
|
||||
|
||||
@@ -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<aether_scheduler_core::ClientSessionAffinity>,
|
||||
pub(crate) policy_context: Option<crate::scheduler::affinity::SchedulerAffinityPolicyContext>,
|
||||
pub(crate) routing_overlay: Option<aether_routing_core::RankingOverlay>,
|
||||
}
|
||||
|
||||
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<Option<ResolvedTunnelSchedulerAffinityContext>, 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,
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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<u64>,
|
||||
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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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")
|
||||
|
||||
+1
@@ -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")
|
||||
|
||||
+32
-6
@@ -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,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>;
|
||||
|
||||
async fn read_minimal_candidate_selection_rows_for_api_format_page(
|
||||
&self,
|
||||
query: &StoredApiFormatCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, 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<RequestedModelCandidateRowsPage, DataLayerError> {
|
||||
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,
|
||||
|
||||
@@ -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<Vec<StoredMinimalCandidateSelectionRow>, 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());
|
||||
|
||||
@@ -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<Vec<StoredMinimalCandidateSelectionRow>, 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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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<Vec<StoredMinimalCandidateSelectionRow>, 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,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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<u64>,
|
||||
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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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<BTreeMap<String, String>>,
|
||||
pub user_agent: Option<String>,
|
||||
pub browser_profile: Option<String>,
|
||||
pub organization_uuid: Option<String>,
|
||||
pub scopes: Option<serde_json::Value>,
|
||||
pub subscription_type: Option<String>,
|
||||
pub rate_limit_tier: Option<String>,
|
||||
}
|
||||
|
||||
#[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("[email protected]"));
|
||||
}
|
||||
|
||||
#[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(
|
||||
|
||||
+14
-156
@@ -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
|
||||
}
|
||||
|
||||
@@ -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<String>,
|
||||
proxy_node_id: Option<String>,
|
||||
}
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_cookie_authorize(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&axum::body::Bytes>,
|
||||
) -> Result<Response<Body>, 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<ClaudeCookieAuthorizeRequest, Response<Body>> {
|
||||
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::<serde_json::Value>(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<String> {
|
||||
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<String, serde_json::Value>,
|
||||
key: &str,
|
||||
) -> Option<String> {
|
||||
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<Body> {
|
||||
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:?}"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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<String, String>;
|
||||
|
||||
struct ClaudeCookieTaskRequest {
|
||||
entries: Vec<ClaudeCookieTaskEntry>,
|
||||
proxy_node_id: Option<String>,
|
||||
}
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_start_cookie_task(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Response<Body>, 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<Value>,
|
||||
created_at: u64,
|
||||
started_at: u64,
|
||||
finished_at: Option<u64>,
|
||||
) {
|
||||
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<Body>) -> 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::<Value>(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<ClaudeCookieTaskRequest, Response<Body>> {
|
||||
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::<Value>(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::<Result<Vec<_>, _>>()?
|
||||
} 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::<Vec<_>>();
|
||||
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<String, Value>, key: &str) -> Option<String> {
|
||||
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<Body> {
|
||||
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::<Vec<_>>();
|
||||
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());
|
||||
}
|
||||
}
|
||||
@@ -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::<serde_json::Value>(request_body) {
|
||||
let mut raw_payload = match serde_json::from_slice::<serde_json::Value>(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,
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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<Response<Body>, 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<Response<Body>, 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<Response<Body>, 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<Response<Body>, 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,
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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<String, Value>) {
|
||||
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<u64>,
|
||||
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());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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<String> {
|
||||
@@ -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<RuntimeLockLease>,
|
||||
) {
|
||||
release_provider_oauth_account_locks(state, leases).await;
|
||||
}
|
||||
|
||||
pub(crate) async fn release_provider_oauth_account_locks(
|
||||
state: &AdminAppState<'_>,
|
||||
leases: Vec<RuntimeLockLease>,
|
||||
) {
|
||||
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<String, serde_json::Value>,
|
||||
) -> Option<String> {
|
||||
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<String, serde_json::Value>,
|
||||
operation: &str,
|
||||
) -> Result<Vec<RuntimeLockLease>, 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<String, serde_json::Value>,
|
||||
existing_auth_config: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> Option<bool> {
|
||||
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<String, serde_json::Value>,
|
||||
existing_auth_config: &serde_json::Map<String, serde_json::Value>,
|
||||
@@ -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": "[email protected]",
|
||||
"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": "[email protected]"
|
||||
}));
|
||||
let same_uuid_other_email = auth_config(json!({
|
||||
"accountUuid": "account-1",
|
||||
"email": "[email protected]"
|
||||
}));
|
||||
let email_only_uppercase = auth_config(json!({"email": "[email protected]"}));
|
||||
let email_only_lowercase = auth_config(json!({"email": "[email protected]"}));
|
||||
|
||||
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": "[email protected]"
|
||||
}));
|
||||
|
||||
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!({
|
||||
|
||||
@@ -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<String>,
|
||||
key_proxy: Option<Value>,
|
||||
request_proxy: Option<ProxySnapshot>,
|
||||
lock_operation: &'static str,
|
||||
) -> Result<Response<Body>, 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<String, Value>,
|
||||
|
||||
@@ -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<ProxySnapshot>,
|
||||
) -> Result<serde_json::Value, Response<Body>> {
|
||||
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",
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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::{
|
||||
|
||||
@@ -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",
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -20,9 +20,14 @@ pub(crate) fn admin_provider_oauth_template(
|
||||
}
|
||||
|
||||
pub(crate) fn build_admin_provider_oauth_supported_types_payload() -> Vec<serde_json::Value> {
|
||||
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<serde_
|
||||
"authorize_url": template.authorize_url,
|
||||
"token_url": template.token_url,
|
||||
"use_pkce": template.use_pkce,
|
||||
"supports_authorization_code": capabilities.as_ref().is_some_and(|value| value.supports_authorization_code),
|
||||
"supports_cookie_authorization": capabilities.as_ref().is_some_and(|value| value.supports_cookie_authorization),
|
||||
"supports_refresh_token_import": capabilities.as_ref().is_some_and(|value| value.supports_refresh_token_import),
|
||||
"supports_batch_import": capabilities.as_ref().is_some_and(|value| value.supports_batch_import),
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
|
||||
@@ -10,6 +10,7 @@ const POOL_ALLOWED_SCHEDULING_PRESETS: &[&str] = &[
|
||||
"load_balance",
|
||||
"single_account",
|
||||
"priority_first",
|
||||
"free_team_first",
|
||||
"free_first",
|
||||
"team_first",
|
||||
"plus_first",
|
||||
@@ -166,8 +167,9 @@ fn parse_pool_score_rules(pool_advanced: &Map<String, Value>) -> PoolMemberScore
|
||||
|
||||
fn normalize_pool_preset_mode(preset: &str, raw_mode: Option<&Value>) -> Option<String> {
|
||||
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]
|
||||
|
||||
@@ -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::<Vec<_>>();
|
||||
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| {
|
||||
|
||||
@@ -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::<Vec<_>>(),
|
||||
["cache_affinity", "recent_refresh"]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_standard_test_resolves_codex_responses_upstream_streaming() {
|
||||
assert!(provider_query_resolve_standard_test_upstream_is_stream(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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<String> {
|
||||
provider_oauth_provider_id_for_suffix(request_path, "/cookie-authorize")
|
||||
}
|
||||
|
||||
pub(crate) fn admin_provider_oauth_cookie_task_provider_id(request_path: &str) -> Option<String> {
|
||||
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<String> {
|
||||
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());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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::<serde_json::Value>(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()
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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";
|
||||
|
||||
@@ -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<String> {
|
||||
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::<RoutingPoolPolicyOverride>(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::<Vec<_>>();
|
||||
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<ManagedRedisServer> {
|
||||
@@ -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();
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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<StoredRoutingGroup> = 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<bool, GatewayRoutingSelectionError> {
|
||||
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<T>() -> Result<T, DataLayerError> {
|
||||
Err(DataLayerError::Sql(
|
||||
"routing repository unavailable".to_string(),
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl RoutingGroupReadRepository for FailingRoutingGroupRepository {
|
||||
async fn list_routing_groups(&self) -> Result<Vec<StoredRoutingGroup>, DataLayerError> {
|
||||
Self::failure()
|
||||
}
|
||||
|
||||
async fn find_routing_group(
|
||||
&self,
|
||||
lookup: RoutingGroupLookupKey<'_>,
|
||||
) -> Result<Option<StoredRoutingGroup>, 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<Vec<StoredRoutingGroupBinding>, DataLayerError> {
|
||||
Self::failure()
|
||||
}
|
||||
|
||||
async fn list_routing_group_versions(
|
||||
&self,
|
||||
_group_id: &str,
|
||||
) -> Result<Vec<StoredRoutingGroupVersion>, 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();
|
||||
|
||||
@@ -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<SchedulerAffinityScope>,
|
||||
}
|
||||
|
||||
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<SchedulerAffinityPolicyContext> {
|
||||
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<String, Value>,
|
||||
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<SchedulerAffinityTarget> {
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -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::<ExecutionPlan>::new()));
|
||||
let execution_plans_clone = Arc::clone(&execution_plans);
|
||||
let execution_runtime = Router::new().route(
|
||||
"/v1/execute/sync",
|
||||
any(move |Json(plan): Json<ExecutionPlan>| {
|
||||
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": "[email protected]"
|
||||
}
|
||||
}),
|
||||
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"], "[email protected]");
|
||||
|
||||
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"], "[email protected]");
|
||||
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::<ExecutionPlan>::new()));
|
||||
let execution_plans_clone = Arc::clone(&execution_plans);
|
||||
let execution_runtime = Router::new().route(
|
||||
"/v1/execute/sync",
|
||||
any(move |Json(plan): Json<ExecutionPlan>| {
|
||||
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",
|
||||
|
||||
@@ -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::<Vec<_>>();
|
||||
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();
|
||||
|
||||
Reference in New Issue
Block a user