Merge origin/main into main

Integrate upstream updates while preserving the local analytics dashboards and schema-only migration changes.

Combine user account analysis with upstream user/group usage statistics in separate tabs, retain all migration versions, and keep the deleted audit document removed.

Validation: gateway all-target cargo check, frontend type check and 57 focused tests, 48 migration tests, schema composition checks, and diff whitespace checks.
This commit is contained in:
elky
2026-10-02 11:57:18 +08:00
343 changed files with 27929 additions and 2549 deletions
+8 -3
View File
@@ -69,9 +69,14 @@ pub(crate) use aether_ai_formats::api::{
OPENAI_VIDEO_REMIX_SYNC_PLAN_KIND,
};
pub(crate) use aether_ai_formats::protocol::stream::CanonicalUsage as StreamingCanonicalUsage;
/// Codex client identity headers re-exported for out-of-crate probe binaries,
/// which must reach `aether_ai_formats` through this seam.
pub use aether_ai_formats::{CODEX_CLIENT_ORIGINATOR, CODEX_CLIENT_USER_AGENT};
/// Codex client identity accessors re-exported for out-of-crate probe binaries,
/// which must reach the runtime profile through this seam.
pub use aether_ai_formats::{codex_client_originator, codex_client_user_agent};
/// Codex 动态客户端画像 API 只允许经此根缝进入 gateway,避免其它模块直接依赖 formats crate。
pub(crate) use aether_ai_formats::{
codex_client_profile, codex_client_version, set_codex_cli_version, set_codex_client_profile,
CodexClientProfile,
};
pub(crate) use aether_ai_formats::{CODEX_RESPONSES_LITE_HEADER, UPSTREAM_IS_STREAM_KEY};
pub(crate) fn parse_direct_request_body(
@@ -18,6 +18,12 @@ pub(crate) fn maybe_build_local_stream_rewriter<'a>(
}
impl LocalStreamRewriter<'_> {
pub(crate) fn into_owned(self) -> LocalStreamRewriter<'static> {
LocalStreamRewriter {
inner: self.inner.into_owned(),
}
}
pub(crate) fn push_chunk(&mut self, chunk: &[u8]) -> Result<Vec<u8>, GatewayError> {
self.inner.push_chunk(chunk).map_err(map_surface_error)
}
@@ -93,6 +93,16 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
&self,
candidate: Self::Candidate,
) -> Self::Skipped {
warn!(
event_name = "local_candidate_skipped",
log_type = "event",
provider_id = %candidate.provider_id,
endpoint_id = %candidate.endpoint_id,
key_id = %candidate.key_id,
api_format = %candidate.endpoint_api_format,
skip_reason = "transport_snapshot_missing",
"local execution candidate skipped during planning"
);
SkippedLocalExecutionCandidate {
candidate,
skip_reason: "transport_snapshot_missing",
@@ -145,6 +155,16 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
transport: Self::Transport,
skip_reason: &'static str,
) -> Self::Skipped {
warn!(
event_name = "local_candidate_skipped",
log_type = "event",
provider_id = %candidate.provider_id,
endpoint_id = %candidate.endpoint_id,
key_id = %candidate.key_id,
api_format = %candidate.endpoint_api_format,
skip_reason,
"local execution candidate skipped during planning"
);
SkippedLocalExecutionCandidate {
candidate,
skip_reason,
@@ -314,6 +314,20 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
return Ok(None);
}
// Same-format requests skip `apply_transport_request_body_semantics`, so the opt-in
// Claude Code body mimicry has to be applied here as well.
if crate::ai_serving::transport::claude_code::apply_claude_code_body_mimicry_for_transport(
&mut base_provider_request_body,
&transport,
prepared.provider_api_format.as_str(),
) {
compatibility_edits.push(SameFormatProviderCompatibilityEdit {
field: "body".to_string(),
action: SameFormatProviderCompatibilityEditAction::ProviderCompatibilityRewrite,
detail: "applied Claude Code body mimicry for provider compatibility".to_string(),
});
}
let antigravity_auth = if prepared.is_antigravity {
let mut antigravity_support = classify_local_antigravity_request_support(
&transport,
@@ -583,6 +597,11 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
source_model,
codex_model_capabilities.as_ref(),
);
crate::ai_serving::transport::xai::insert_cli_identity_headers_if_needed(
transport.as_ref(),
prepared.provider_api_format.as_str(),
&mut provider_request_headers,
);
request_identity_response_encoding_when_redacted(
&mut provider_request_headers,
redaction.redacted,
@@ -17,8 +17,8 @@ use crate::ai_serving::transport::{
ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput, GROK_CHAT_PATH,
};
use crate::ai_serving::{
apply_codex_openai_special_headers, build_chatgpt_web_image_request_body,
build_codex_openai_image_api_provider_request_body,
apply_codex_openai_special_headers, apply_xai_upstream_payload_edits,
build_chatgpt_web_image_request_body, build_codex_openai_image_api_provider_request_body,
build_gemini_image_request_body_from_openai_image_request,
build_openai_image_api_provider_request_body, build_openai_image_provider_request_body,
default_model_for_openai_image_operation, normalize_openai_image_request,
@@ -211,7 +211,7 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
upstream_is_stream,
)
};
let Some(provider_request_body) = provider_request_body else {
let Some(mut provider_request_body) = provider_request_body else {
mark_skipped_local_openai_image_candidate_with_failure_diagnostic(
state,
input,
@@ -229,6 +229,11 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
.await;
return None;
};
apply_xai_upstream_payload_edits(
&mut provider_request_body,
transport.provider.provider_type.as_str(),
provider_api_format,
);
let Some(mut provider_request_headers) = (if is_grok {
build_grok_browser_headers(GrokHeaderInput {
transport,
@@ -8,6 +8,7 @@ use crate::ai_serving::planner::{
build_ai_execution_decision_response, resolve_transport_request_encoding_policy,
AiExecutionDecisionResponseParts,
};
use crate::ai_serving::transport::xai::video::is_native_video_request;
use crate::ai_serving::transport::{
resolve_transport_execution_timeouts, resolve_transport_profile,
};
@@ -33,7 +34,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
let Some(resolved) = resolve_local_video_create_candidate_payload_parts(
state, parts, body_json, trace_id, input, &attempt, spec,
)
.await
.await?
else {
return Ok(None);
};
@@ -52,9 +53,32 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
.await;
let transport_profile = resolve_transport_profile(&transport);
let mut extra_fields = serde_json::Map::new();
if is_native_video_request(&transport.provider.provider_type, parts.uri.path()) {
extra_fields.insert(
"video_client_protocol".to_string(),
serde_json::json!("xai"),
);
}
if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) {
extra_fields.insert("proxy".to_string(), proxy_value);
}
if transport.provider.provider_type.eq_ignore_ascii_case("xai") {
extra_fields.insert("video_provider_xai".into(), serde_json::json!(true));
if let Some(duration) = resolved.provider_request_body.get("duration") {
extra_fields.insert("video_duration".into(), duration.clone());
}
if parts.uri.path() == "/openai/v1/videos" {
extra_fields.insert(
"video_size".into(),
body_json
.get("size")
.filter(|v| v.as_str().is_some_and(|s| !s.trim().is_empty()))
.cloned()
.unwrap_or_else(|| serde_json::json!("720x1280")),
);
}
}
let effective_headers = input.effective_headers(&parts.headers);
let report_context = build_local_execution_report_context(LocalExecutionReportContextParts {
auth_context: &input.auth_context,
@@ -3,15 +3,23 @@ use std::sync::Arc;
use serde_json::Value;
use crate::ai_serving::planner::candidate_preparation::resolve_candidate_mapped_model;
use crate::ai_serving::planner::candidate_preparation::{
prepare_header_authenticated_candidate, resolve_candidate_mapped_model, OauthPreparationContext,
};
use crate::ai_serving::planner::spec_metadata::local_video_create_spec_metadata;
use crate::ai_serving::transport::xai::video::{
convert_openai_video_request, is_explicit_native_video_path, is_native_video_request,
};
use crate::ai_serving::transport::{
build_video_create_headers, build_video_create_request_body, build_video_create_upstream_url,
resolve_video_create_auth, video_create_transport_unsupported_reason,
ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput,
};
use crate::ai_serving::{CandidateFailureDiagnostic, GatewayProviderTransportSnapshot};
use crate::AppState;
use crate::ai_serving::{
apply_xai_upstream_payload_edits, CandidateFailureDiagnostic, GatewayProviderTransportSnapshot,
PlannerAppState,
};
use crate::{AppState, GatewayError};
use super::support::{
mark_skipped_local_video_candidate, mark_skipped_local_video_candidate_with_failure_diagnostic,
@@ -37,11 +45,16 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
input: &LocalVideoCreateDecisionInput,
attempt: &LocalVideoCreateCandidateAttempt,
spec: LocalVideoCreateSpec,
) -> Option<LocalVideoCreateCandidatePayloadParts> {
) -> Result<Option<LocalVideoCreateCandidatePayloadParts>, GatewayError> {
let spec_metadata = local_video_create_spec_metadata(spec);
let candidate = &attempt.eligible.candidate;
let transport = &attempt.eligible.transport;
let effective_headers = input.effective_headers(&parts.headers);
if is_explicit_native_video_path(parts.uri.path())
&& !transport.provider.provider_type.eq_ignore_ascii_case("xai")
{
return Ok(None);
}
let provider_family = provider_video_create_family(spec.family);
let transport_unsupported_reason = video_create_transport_unsupported_reason(
@@ -60,23 +73,39 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
skip_reason,
)
.await;
return None;
return Ok(None);
}
let auth = resolve_video_create_auth(transport, provider_family);
let Some((auth_header, auth_value)) = auth else {
mark_skipped_local_video_candidate(
state,
input,
let prepared_candidate = match prepare_header_authenticated_candidate(
PlannerAppState::new(state),
transport,
candidate,
resolve_video_create_auth(transport, provider_family),
OauthPreparationContext {
trace_id,
candidate,
attempt.candidate_index,
&attempt.candidate_id,
"transport_auth_unavailable",
)
.await;
return None;
api_format: spec_metadata.api_format,
operation: "video_create_candidate_request",
},
)
.await
{
Ok(prepared) => prepared,
Err(skip_reason) => {
mark_skipped_local_video_candidate(
state,
input,
trace_id,
candidate,
attempt.candidate_index,
&attempt.candidate_id,
skip_reason,
)
.await;
return Ok(None);
}
};
let auth_header = prepared_candidate.auth_header;
let auth_value = prepared_candidate.auth_value;
let mapped_model = match resolve_candidate_mapped_model(candidate) {
Ok(mapped_model) => mapped_model,
@@ -91,7 +120,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
skip_reason,
)
.await;
return None;
return Ok(None);
}
};
@@ -117,10 +146,10 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
),
)
.await;
return None;
return Ok(None);
};
let Some(provider_request_body) = build_video_create_request_body(
let Some(mut provider_request_body) = build_video_create_request_body(
body_json,
provider_family,
&mapped_model,
@@ -142,11 +171,28 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
),
)
.await;
return None;
return Ok(None);
};
if transport.provider.provider_type.eq_ignore_ascii_case("xai")
&& !is_native_video_request(&transport.provider.provider_type, parts.uri.path())
{
provider_request_body =
convert_openai_video_request(&provider_request_body).map_err(|message| {
GatewayError::Client {
status: http::StatusCode::BAD_REQUEST,
message: message.to_string(),
}
})?;
}
apply_xai_upstream_payload_edits(
&mut provider_request_body,
transport.provider.provider_type.as_str(),
spec_metadata.api_format,
);
let Some(provider_request_headers) =
build_video_create_headers(ProviderVideoCreateHeadersInput {
transport,
headers: effective_headers,
auth_header: &auth_header,
auth_value: &auth_value,
@@ -170,10 +216,10 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
),
)
.await;
return None;
return Ok(None);
};
Some(LocalVideoCreateCandidatePayloadParts {
Ok(Some(LocalVideoCreateCandidatePayloadParts {
transport: Arc::clone(transport),
auth_header,
auth_value,
@@ -181,7 +227,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
provider_request_headers,
provider_request_body,
upstream_url,
})
}))
}
fn provider_video_create_family(family: LocalVideoCreateFamily) -> ProviderVideoCreateFamily {
@@ -505,9 +505,12 @@ fn projects_uuid_prompt_cache_identity_into_missing_session_headers() {
assert_eq!(headers.get("x-client-request-id"), None);
assert_eq!(
headers.get("user-agent").map(String::as_str),
Some(aether_ai_formats::CODEX_CLIENT_USER_AGENT)
Some(aether_ai_formats::codex_client_user_agent().as_str())
);
assert_eq!(
headers.get("originator"),
Some(&aether_ai_formats::codex_client_originator())
);
assert_eq!(headers.get("originator"), Some(&"codex_cli_rs".to_string()));
assert!(!headers.contains_key("version"));
assert_eq!(headers.get("x-openai-fedramp"), Some(&"true".to_string()));
assert_eq!(
@@ -615,9 +618,12 @@ fn injects_only_codex_client_headers_for_images_requests() {
);
assert_eq!(
headers.get("user-agent").map(String::as_str),
Some(aether_ai_formats::CODEX_CLIENT_USER_AGENT)
Some(aether_ai_formats::codex_client_user_agent().as_str())
);
assert_eq!(
headers.get("originator"),
Some(&aether_ai_formats::codex_client_originator())
);
assert_eq!(headers.get("originator"), Some(&"codex_cli_rs".to_string()));
assert!(!headers.contains_key("version"));
assert_eq!(headers.get("x-openai-fedramp"), Some(&"true".to_string()));
for name in ["x-client-request-id", "session-id", "thread-id"] {
@@ -699,9 +705,12 @@ fn preserves_client_context_headers_and_enforces_codex_provider_identity() {
);
assert_eq!(
headers.get("user-agent").map(String::as_str),
Some(aether_ai_formats::CODEX_CLIENT_USER_AGENT)
Some(aether_ai_formats::codex_client_user_agent().as_str())
);
assert_eq!(
headers.get("originator"),
Some(&aether_ai_formats::codex_client_originator())
);
assert_eq!(headers.get("originator"), Some(&"codex_cli_rs".to_string()));
assert_eq!(
headers
.keys()
@@ -763,9 +772,12 @@ fn compact_projects_uuid_prompt_cache_identity_into_session_headers() {
assert_eq!(headers.get("x-client-request-id"), None);
assert_eq!(
headers.get("user-agent").map(String::as_str),
Some(aether_ai_formats::CODEX_CLIENT_USER_AGENT)
Some(aether_ai_formats::codex_client_user_agent().as_str())
);
assert_eq!(
headers.get("originator"),
Some(&aether_ai_formats::codex_client_originator())
);
assert_eq!(headers.get("originator"), Some(&"codex_cli_rs".to_string()));
assert!(!headers.contains_key("version"));
assert_eq!(headers.get("x-openai-fedramp"), Some(&"true".to_string()));
assert_eq!(
@@ -13,7 +13,9 @@ pub(crate) fn openai_responses_reasoning_replay_policy(
base_url: &str,
_provider_model: &str,
) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy {
if is_deepseek_provider(provider_type, base_url) {
if provider_type.trim().eq_ignore_ascii_case("xai") {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
} else if is_deepseek_provider(provider_type, base_url) {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
} else {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
@@ -238,6 +240,27 @@ mod tests {
openai_responses_reasoning_replay_policy,
};
#[test]
fn xai_reasoning_policy_comes_from_provider_type() {
use crate::ai_serving::OpenAiResponsesReasoningReplayPolicy;
assert_eq!(
openai_responses_reasoning_replay_policy(
"xai",
"https://custom.example/v1",
"grok-4.6"
),
OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
);
assert_eq!(
openai_responses_reasoning_replay_policy(
"openai",
"https://custom.example/v1",
"grok-4.6"
),
OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
);
}
#[test]
fn detects_deepseek_provider_only_by_official_host() {
assert!(!is_deepseek_provider(
@@ -4,7 +4,7 @@ use crate::ai_serving::transport::apply_standard_provider_request_body_rules_wit
use crate::ai_serving::{
apply_codex_openai_responses_chat_body_edits,
apply_openai_responses_compact_special_body_edits,
build_cross_format_openai_chat_request_body_with_model_directives as surface_build_cross_format_openai_chat_request_body,
build_cross_format_openai_chat_request_body_with_provider_context as surface_build_cross_format_openai_chat_request_body,
build_local_openai_chat_request_body_with_model_directives as surface_build_local_openai_chat_request_body,
GatewayProviderTransportSnapshot,
};
@@ -73,9 +73,11 @@ pub(crate) fn build_cross_format_openai_chat_request_body(
let provider_request_body = surface_build_cross_format_openai_chat_request_body(
body_json,
mapped_model,
provider_type,
provider_api_format,
upstream_is_stream,
enable_model_directives,
user_api_key_id,
)?;
let mut provider_request_body =
apply_standard_provider_request_body_rules_with_request_headers(
@@ -110,6 +112,42 @@ pub(crate) fn build_cross_format_openai_chat_request_body(
Some(provider_request_body)
}
#[cfg(test)]
mod antigravity_schema_tests {
use super::*;
use serde_json::json;
#[test]
fn antigravity_chat_route_preserves_tool_schema_and_alternate_responses_shape() {
let schema = json!({"type": "object", "properties": {"mode": {"const": "fast"}}});
let body = json!({"model": "client", "messages": [{"role": "user", "content": "hi"}],
"tools": [{"type": "function", "function": {"name": "probe", "parameters": schema}}]});
let responses_body = json!({"model": "client", "input": "hi",
"tools": [{"type": "function", "name": "probe", "parameters": schema}]});
for input in [body, responses_body] {
for provider in ["antigravity", "gemini"] {
let output = build_cross_format_openai_chat_request_body(
&input,
"claude-test",
provider,
"gemini:generate_content",
true,
false,
None,
None,
&http::HeaderMap::new(),
false,
)
.unwrap();
let parameters = &output["tools"][0]["functionDeclarations"][0]["parameters"];
assert_eq!(parameters == &schema, provider == "antigravity");
assert!(output.get("stream").is_none());
assert_eq!(output["contents"][0]["parts"][0]["text"], "hi");
}
}
}
}
pub(crate) fn build_cross_format_openai_chat_upstream_url(
parts: &http::request::Parts,
transport: &GatewayProviderTransportSnapshot,
@@ -3,7 +3,7 @@ use serde_json::Value;
use crate::ai_serving::transport::apply_standard_provider_request_body_rules_with_request_headers;
use crate::ai_serving::{
apply_openai_responses_compact_special_body_edits,
build_cross_format_openai_responses_request_body_with_model_directives_and_history_scope as surface_build_cross_format_openai_responses_request_body,
build_cross_format_openai_responses_request_body_with_provider_context as surface_build_cross_format_openai_responses_request_body,
build_local_openai_responses_request_body_with_model_directives as surface_build_local_openai_responses_request_body,
GatewayProviderTransportSnapshot,
};
@@ -218,6 +218,7 @@ pub(crate) fn build_cross_format_openai_responses_request_body_with_codex_model_
body_json,
mapped_model,
client_api_format,
provider_type,
provider_api_format,
upstream_is_stream,
enable_model_directives,
@@ -274,6 +275,41 @@ pub(crate) fn build_local_openai_responses_upstream_url(
)
}
#[cfg(test)]
mod antigravity_schema_tests {
use super::*;
use serde_json::json;
#[test]
fn antigravity_responses_route_preserves_tool_schema_without_changing_public_gemini() {
let schema = json!({"type": "object", "properties": {"mode": {"const": "fast"}}});
let input = json!({"model": "client", "input": "hi",
"tools": [{"type": "function", "name": "probe", "parameters": schema}]});
for provider in ["antigravity", "gemini"] {
let output =
build_cross_format_openai_responses_request_body_with_codex_model_capabilities(
&input,
"claude-test",
"openai:responses",
"gemini:generate_content",
true,
false,
provider,
None,
&http::HeaderMap::new(),
Some("antigravity-schema-test"),
None,
false,
)
.unwrap();
let parameters = &output["tools"][0]["functionDeclarations"][0]["parameters"];
assert_eq!(parameters == &schema, provider == "antigravity");
assert!(output.get("stream").is_none());
assert_eq!(output["contents"][0]["parts"][0]["text"], "hi");
}
}
}
pub(crate) fn build_cross_format_openai_responses_upstream_url(
parts: &http::request::Parts,
transport: &GatewayProviderTransportSnapshot,
@@ -140,7 +140,7 @@ fn finalize_openai_chat_provider_request_body(
mapped_model,
source_model,
);
crate::ai_serving::finalize_openai_provider_request_with_codex_model_capabilities_and_reasoning_replay_policy(
let finalization_failure = crate::ai_serving::finalize_openai_provider_request_with_codex_model_capabilities_and_reasoning_replay_policy(
provider_request_body,
crate::ai_serving::OpenAiProviderRequestFinalization {
source_api_format: "openai:chat",
@@ -170,7 +170,17 @@ fn finalize_openai_chat_provider_request_body(
provider_api_format,
"openai_chat_request_finalization",
)
})
});
if finalization_failure.is_none() {
// This builder does not go through `apply_transport_request_body_semantics`, so the
// Claude Code body mimicry must be applied here for Chat -> claude_code requests.
crate::ai_serving::transport::claude_code::apply_claude_code_body_mimicry_for_transport(
provider_request_body,
transport,
provider_api_format,
);
}
finalization_failure
}
#[allow(clippy::too_many_arguments)]
@@ -2763,7 +2773,7 @@ mod tests {
payload.provider_request_body["userAgent"],
"vscode/1.X.X (Antigravity/4.3.0)"
);
assert_eq!(payload.provider_request_body["requestType"], "agent");
assert!(payload.provider_request_body.get("requestType").is_none());
assert!(payload.provider_request_body.get("contents").is_none());
assert!(payload.provider_request_body["request"]
.get("contents")
@@ -1,4 +1,4 @@
use aether_routing_core::RoutingExecutionPolicy;
use aether_routing_core::{RoutingExecutionPolicy, RoutingSchedulingMode};
use async_trait::async_trait;
use std::collections::VecDeque;
use tracing::warn;
@@ -207,7 +207,12 @@ impl LocalOpenAiChatStreamAttemptSource<'_> {
async fn next_raw_attempt_with_target_select(
&mut self,
) -> Result<Option<LocalOpenAiChatCandidateAttempt>, GatewayError> {
let select_window = openai_chat_stream_target_select_window();
let select_window = openai_chat_stream_target_select_window_for_mode(
self.input
.routing_policy
.as_ref()
.map(|policy| policy.scheduling_mode),
);
if select_window <= 1 {
return self.next_raw_attempt_linear().await;
}
@@ -365,6 +370,15 @@ fn openai_chat_stream_target_select_window() -> usize {
.clamp(1, MAX_OPENAI_CHAT_STREAM_TARGET_SELECT_WINDOW)
}
fn openai_chat_stream_target_select_window_for_mode(
scheduling_mode: Option<RoutingSchedulingMode>,
) -> usize {
if scheduling_mode == Some(RoutingSchedulingMode::FixedOrder) {
return 1;
}
openai_chat_stream_target_select_window()
}
#[derive(Clone, Copy)]
struct TargetSelectCandidateIdentity<'a> {
provider_id: &'a str,
@@ -574,4 +588,14 @@ mod tests {
assert_eq!(select_target_index(19, &choices), 1);
}
#[test]
fn fixed_order_disables_stream_target_selection() {
assert_eq!(
openai_chat_stream_target_select_window_for_mode(Some(
RoutingSchedulingMode::FixedOrder,
)),
1
);
}
}
@@ -635,6 +635,13 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts_with_
{
log_responses_to_chat_tool_conversion(trace_id, body_json, &base_provider_request_body);
}
// This builder does not go through `apply_transport_request_body_semantics`, so the
// Claude Code body mimicry must be applied here for Responses -> claude_code requests.
crate::ai_serving::transport::claude_code::apply_claude_code_body_mimicry_for_transport(
&mut base_provider_request_body,
&transport,
provider_api_format,
);
let provider_request_body = base_provider_request_body;
if let Some(kiro_auth) = kiro_auth.as_ref() {
@@ -488,6 +488,7 @@ impl ResponsesWebSocketBodyNormalization {
digest.update([match self.reasoning_replay_policy {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds => 0,
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque => 1,
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted => 2,
}]);
update_normalization_optional_json_digest(&mut digest, self.model_directive_patch.as_ref());
digest.finalize().into()
@@ -12,13 +12,16 @@ pub(crate) use aether_ai_formats::api::{
apply_codex_openai_responses_websocket_continuation_body_edits_with_source_model_and_capabilities,
apply_codex_openai_special_headers, apply_model_directive_mapping_patch,
apply_model_directive_overrides_from_model, apply_model_directive_overrides_from_request,
apply_openai_responses_compact_special_body_edits, build_chatgpt_web_image_request_body,
apply_openai_responses_compact_special_body_edits, apply_xai_upstream_payload_edits,
apply_xai_upstream_payload_edits_with_client, build_chatgpt_web_image_request_body,
build_codex_model_catalog_metadata, build_codex_openai_image_api_provider_request_body,
build_core_error_body_for_client_format, build_cross_format_openai_chat_request_body,
build_cross_format_openai_chat_request_body_with_model_directives,
build_cross_format_openai_chat_request_body_with_provider_context,
build_cross_format_openai_responses_request_body,
build_cross_format_openai_responses_request_body_with_model_directives,
build_cross_format_openai_responses_request_body_with_model_directives_and_history_scope,
build_cross_format_openai_responses_request_body_with_provider_context,
build_gemini_image_request_body_from_openai_image_request,
build_gemini_image_response_from_openai_image_response,
build_gemini_image_response_from_openai_responses_image_response, build_generated_tool_call_id,
@@ -181,7 +184,7 @@ pub(crate) use aether_ai_formats::{
openai_responses_request_operation, openai_responses_synthetic_reasoning_item_id,
strip_incompatible_openai_responses_reasoning_items,
strip_incompatible_openai_responses_reasoning_items_with_policy, ApiOperation, ClientSurface,
CODEX_CLIENT_VERSION, OPENAI_RESPONSES_OPERATION_COMPACT,
OPENAI_RESPONSES_OPERATION_COMPACT,
};
pub(crate) fn plan_kind_matches_api_operation(
@@ -58,6 +58,10 @@ pub(crate) mod windsurf {
pub(crate) use aether_provider_transport::windsurf::*;
}
pub(crate) mod xai {
pub(crate) use aether_provider_transport::xai::*;
}
pub(crate) use aether_provider_transport::{
append_transport_diagnostics_to_value, apply_codex_fingerprint_convergence,
apply_codex_fingerprint_convergence_with_context, apply_local_auth_config_header_overrides,
@@ -53,6 +53,8 @@ const AI_ANY_ROUTE_PATTERNS: &[&str] = &[
"/v1beta/operations/{*operation_path}",
"/v1/videos",
"/v1/videos/{*video_path}",
"/openai/v1/videos",
"/openai/v1/videos/{*video_path}",
"/upload/v1beta/files",
"/v1beta/files",
"/v1beta/files/{*file_path}",
@@ -536,6 +536,9 @@ mod tests {
fn sample_sparse_stored_task() -> StoredVideoTask {
let snapshot = LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: false,
local_task_id: "task-1".to_string(),
upstream_task_id: "ext-1".to_string(),
created_at_unix_ms: 1,
+12 -1
View File
@@ -799,6 +799,7 @@ mod tests {
use aes_gcm::aead::{Aead, AeadCore, KeyInit, OsRng, Payload};
use aes_gcm::Aes256Gcm;
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
use base64::Engine as _;
use bytes::Bytes;
use chrono::{DateTime, Utc};
use serde_json::json;
@@ -1243,8 +1244,18 @@ mod tests {
assert_eq!(restored.key_id, None);
assert_eq!(restored.export_version.as_deref(), Some("2.3"));
// 17 个互不相同的合法 base64-32 字节直接密钥:本段只验证“legacy 候选 >16 → TooManyLegacyKeys”,
// 不测口令强度、不解密。直接密钥走 decode_direct_fernet_key(生产已支持路径),跳过 PBKDF2,
// 避免本用例为计数语义再付 17×10 万次迭代;上半段 DEVELOPMENT_ENCRYPTION_KEY 真实 v1 兼容
// 与 wrong-legacy-secret 派生路径保持不变。
let too_many: Vec<_> = (0..17)
.map(|index| BackupDecryptionKey::historical(format!("legacy-{index}")).unwrap())
.map(|index| {
let mut material = [0u8; 32];
material[0] = index as u8 + 1;
material[31] = index as u8 + 1;
let secret = base64::engine::general_purpose::STANDARD.encode(material);
BackupDecryptionKey::historical(secret).unwrap()
})
.collect();
assert!(matches!(
restore_backup_json(
@@ -7,7 +7,7 @@
#[path = "support/responses_ws_probe.rs"]
mod responses_ws_probe;
use aether_gateway::{CODEX_CLIENT_ORIGINATOR, CODEX_CLIENT_USER_AGENT};
use aether_gateway::{codex_client_originator, codex_client_user_agent};
use clap::Parser;
use http::header::{AUTHORIZATION, USER_AGENT};
use http::{HeaderMap, HeaderName, HeaderValue};
@@ -78,14 +78,12 @@ fn handshake_headers(access_token: &str, account_id: &str) -> Result<HeaderMap,
let mut headers = HeaderMap::new();
headers.insert(AUTHORIZATION, bearer_authorization_value(access_token)?);
headers.insert(HeaderName::from_static("chatgpt-account-id"), account_id);
headers.insert(
USER_AGENT,
HeaderValue::from_static(CODEX_CLIENT_USER_AGENT),
);
headers.insert(
HeaderName::from_static("originator"),
HeaderValue::from_static(CODEX_CLIENT_ORIGINATOR),
);
let user_agent = HeaderValue::from_str(&codex_client_user_agent())
.map_err(|_| ProbeFailure::MissingConfiguration)?;
headers.insert(USER_AGENT, user_agent);
let originator = HeaderValue::from_str(&codex_client_originator())
.map_err(|_| ProbeFailure::MissingConfiguration)?;
headers.insert(HeaderName::from_static("originator"), originator);
Ok(headers)
}
@@ -111,6 +109,18 @@ mod tests {
assert!(headers.contains_key("chatgpt-account-id"));
assert!(headers.contains_key(USER_AGENT));
assert!(headers.contains_key("originator"));
assert_eq!(
headers
.get(USER_AGENT)
.and_then(|value| value.to_str().ok()),
Some(aether_gateway::codex_client_user_agent().as_str())
);
assert_eq!(
headers
.get("originator")
.and_then(|value| value.to_str().ok()),
Some(aether_gateway::codex_client_originator().as_str())
);
assert_eq!(
CodexResponsesProbeProfile::sent_header_names(),
vec![
+468
View File
@@ -0,0 +1,468 @@
//! Codex 客户端画像的运行时发布与官方 CLI 版本刷新。
use std::collections::BTreeMap;
use std::future::Future;
use std::time::Duration;
use aether_runtime_state::RuntimeState;
use futures_util::StreamExt as _;
use reqwest::{redirect::Policy, Client};
use semver::Version;
use serde::{Deserialize, Serialize};
use tracing::{info, warn};
use crate::ai_serving::api::{codex_client_version, set_codex_cli_version};
use crate::AppState;
const CLI_RELEASE_ENDPOINT: &str = "https://registry.npmjs.org/@openai%2Fcodex/latest";
const PROFILE_CACHE_KEY: &str = "aether:codex:client-profile:v1";
const PROFILE_CACHE_TTL: Duration = Duration::from_secs(30 * 24 * 60 * 60);
const PROFILE_REFRESH_INTERVAL: Duration = Duration::from_secs(24 * 60 * 60);
const RELEASE_CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
const RELEASE_REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
const MAX_RELEASE_BYTES: usize = 256 * 1024;
const CLI_TARGETS: [&str; 6] = [
"darwin-arm64",
"darwin-x64",
"linux-arm64",
"linux-x64",
"win32-arm64",
"win32-x64",
];
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
struct NpmRelease {
name: String,
version: String,
optional_dependencies: BTreeMap<String, String>,
}
#[derive(Debug, Deserialize, Serialize)]
struct CachedProfile {
version: String,
verified_at_unix_secs: u64,
}
#[derive(Debug, thiserror::Error)]
enum ProfileRefreshError {
#[error("Codex CLI release client initialization failed: {0}")]
Client(#[from] reqwest::Error),
#[error("Codex CLI release request returned HTTP {0}")]
HttpStatus(u16),
#[error("Codex CLI release response exceeded {MAX_RELEASE_BYTES} bytes")]
ResponseTooLarge,
#[error("Codex CLI release metadata is invalid")]
InvalidMetadata,
#[error("Codex CLI release version is older than the active profile")]
Rollback,
#[error("Codex CLI profile cache operation failed: {0}")]
Cache(String),
}
fn version_sequence(version: &str) -> Result<u64, ProfileRefreshError> {
let parsed = Version::parse(version).map_err(|_| ProfileRefreshError::InvalidMetadata)?;
if !parsed.pre.is_empty()
|| !parsed.build.is_empty()
|| parsed.major > 999
|| parsed.minor > 999
|| parsed.patch > 999
{
return Err(ProfileRefreshError::InvalidMetadata);
}
Ok(1 + parsed.major * 1_000_000 + parsed.minor * 1_000 + parsed.patch)
}
/// 校验官方 npm stable 标签及六个平台依赖来自同一版本发布。
fn parse_cli_release(bytes: &[u8]) -> Result<String, ProfileRefreshError> {
if bytes.len() > MAX_RELEASE_BYTES {
return Err(ProfileRefreshError::ResponseTooLarge);
}
let release = serde_json::from_slice::<NpmRelease>(bytes)
.map_err(|_| ProfileRefreshError::InvalidMetadata)?;
let sequence = version_sequence(&release.version)?;
if sequence == 0
|| release.name != "@openai/codex"
|| CLI_TARGETS.iter().any(|target| {
release
.optional_dependencies
.get(&format!("@openai/codex-{target}"))
!= Some(&format!("npm:@openai/codex@{}-{target}", release.version))
})
{
return Err(ProfileRefreshError::InvalidMetadata);
}
Ok(release.version)
}
fn refresh_enabled_from(value: Option<&str>) -> bool {
!value.is_some_and(|value| {
matches!(
value.trim().to_ascii_lowercase().as_str(),
"0" | "false" | "off"
)
})
}
fn refresh_enabled() -> bool {
refresh_enabled_from(
std::env::var("AETHER_CODEX_CLIENT_PROFILE_REFRESH")
.ok()
.as_deref(),
)
}
fn fixed_version_from(value: Option<&str>) -> Option<String> {
let value = value?.trim();
if value.is_empty() || version_sequence(value).is_err() {
None
} else {
Some(value.to_owned())
}
}
fn fixed_version_override() -> Option<String> {
let value = std::env::var("AETHER_CODEX_CLIENT_VERSION").ok()?;
let version = fixed_version_from(Some(&value));
if version.is_none() {
warn!(
event_name = "codex_client_profile_fixed_version_invalid",
"AETHER_CODEX_CLIENT_VERSION is invalid; using cached or built-in profile"
);
}
version
}
fn build_release_client() -> Result<Client, ProfileRefreshError> {
Client::builder()
.https_only(true)
.no_proxy()
.redirect(Policy::none())
.connect_timeout(RELEASE_CONNECT_TIMEOUT)
.timeout(RELEASE_REQUEST_TIMEOUT)
.build()
.map_err(ProfileRefreshError::Client)
}
async fn fetch_latest_cli_version(client: &Client) -> Result<String, ProfileRefreshError> {
let response = client
.get(CLI_RELEASE_ENDPOINT)
.send()
.await
.map_err(ProfileRefreshError::Client)?;
if !response.status().is_success() {
return Err(ProfileRefreshError::HttpStatus(response.status().as_u16()));
}
if response
.content_length()
.is_some_and(|length| length > MAX_RELEASE_BYTES as u64)
{
return Err(ProfileRefreshError::ResponseTooLarge);
}
let mut bytes = Vec::new();
let mut stream = response.bytes_stream();
while let Some(chunk) = stream.next().await {
let chunk = chunk.map_err(ProfileRefreshError::Client)?;
if bytes.len().saturating_add(chunk.len()) > MAX_RELEASE_BYTES {
return Err(ProfileRefreshError::ResponseTooLarge);
}
bytes.extend_from_slice(&chunk);
}
parse_cli_release(&bytes)
}
async fn restore_cached_profile(runtime: &RuntimeState) -> Result<(), ProfileRefreshError> {
let Some(raw) = runtime
.kv_get(PROFILE_CACHE_KEY)
.await
.map_err(|err| ProfileRefreshError::Cache(err.to_string()))?
else {
return Ok(());
};
let cached = serde_json::from_str::<CachedProfile>(&raw)
.map_err(|_| ProfileRefreshError::InvalidMetadata)?;
if let Some(version) = cached_version_to_restore(&cached, &codex_client_version())? {
set_codex_cli_version(&version).map_err(|_| ProfileRefreshError::InvalidMetadata)?;
info!(
event_name = "codex_client_profile_restored",
version = %version,
verified_at_unix_secs = cached.verified_at_unix_secs,
"restored cached Codex CLI profile"
);
}
Ok(())
}
fn cached_version_to_restore(
cached: &CachedProfile,
active_version: &str,
) -> Result<Option<String>, ProfileRefreshError> {
let cached_sequence = version_sequence(&cached.version)?;
let active_sequence = version_sequence(active_version)?;
Ok((cached_sequence >= active_sequence).then(|| cached.version.clone()))
}
async fn refresh_once_with_fetch<F, Fut>(
runtime: &RuntimeState,
fixed_version: Option<&str>,
refresh_is_enabled: bool,
fetch_latest: F,
) -> Result<String, ProfileRefreshError>
where
F: FnOnce() -> Fut,
Fut: Future<Output = Result<String, ProfileRefreshError>>,
{
if let Some(version) = fixed_version {
set_codex_cli_version(version).map_err(|_| ProfileRefreshError::InvalidMetadata)?;
return Ok(version.to_owned());
}
if let Err(error) = restore_cached_profile(runtime).await {
// 缓存损坏或暂时不可用不应阻断官方版本检查;当前进程继续使用旧画像。
warn!(
event_name = "codex_client_profile_cache_restore_failed",
error = %error,
"could not restore cached Codex CLI profile"
);
}
if !refresh_is_enabled {
return Ok(codex_client_version());
}
let version = fetch_latest().await?;
let current = codex_client_version();
if version_sequence(&version)? < version_sequence(&current)? {
return Err(ProfileRefreshError::Rollback);
}
let cached = CachedProfile {
version: version.clone(),
verified_at_unix_secs: chrono::Utc::now().timestamp().max(0) as u64,
};
let serialized =
serde_json::to_string(&cached).map_err(|_| ProfileRefreshError::InvalidMetadata)?;
set_codex_cli_version(&version).map_err(|_| ProfileRefreshError::InvalidMetadata)?;
if let Err(error) = runtime
.kv_set(PROFILE_CACHE_KEY, serialized, Some(PROFILE_CACHE_TTL))
.await
{
// 本地画像已经完成原子替换;缓存写失败只影响下次进程启动的恢复。
warn!(
event_name = "codex_client_profile_cache_write_failed",
error = %error,
"published Codex CLI profile locally but could not persist the cache"
);
}
Ok(version)
}
async fn refresh_once(runtime: &RuntimeState) -> Result<String, ProfileRefreshError> {
let fixed_version = fixed_version_override();
refresh_once_with_fetch(
runtime,
fixed_version.as_deref(),
refresh_enabled(),
|| async {
let client = build_release_client()?;
fetch_latest_cli_version(&client).await
},
)
.await
}
pub(crate) async fn prewarm(runtime: &RuntimeState) -> Result<String, String> {
refresh_once(runtime).await.map_err(|err| err.to_string())
}
pub(crate) fn spawn_worker(app: AppState) -> tokio::task::JoinHandle<()> {
crate::task_runtime::spawn_singleton_worker(
app,
crate::task_runtime::TASK_KEY_CODEX_CLIENT_PROFILE,
|app| async move {
let mut interval = tokio::time::interval(PROFILE_REFRESH_INTERVAL);
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
// 启动阶段由 prewarm 完成一次检查;后台任务只负责后续每日刷新,避免重复建连。
interval.tick().await;
loop {
interval.tick().await;
match refresh_once(app.runtime_state()).await {
Ok(version) => info!(
event_name = "codex_client_profile_refreshed",
version = %version,
"refreshed Codex CLI profile"
),
Err(error) => warn!(
event_name = "codex_client_profile_refresh_failed",
error = %error,
"keeping the previous Codex CLI profile after refresh failure"
),
}
}
},
)
}
#[cfg(test)]
mod tests {
use std::sync::{
atomic::{AtomicBool, Ordering},
Mutex, OnceLock,
};
use std::time::Duration;
use aether_runtime_state::{MemoryRuntimeStateConfig, RuntimeState};
use super::{
cached_version_to_restore, fixed_version_from, parse_cli_release, refresh_enabled_from,
refresh_once_with_fetch, CachedProfile, ProfileRefreshError, PROFILE_CACHE_KEY,
};
use crate::ai_serving::api::{
codex_client_profile, codex_client_version, set_codex_cli_version,
set_codex_client_profile, CodexClientProfile,
};
static PROFILE_TEST_LOCK: OnceLock<Mutex<()>> = OnceLock::new();
struct ProfileRestore(CodexClientProfile);
impl Drop for ProfileRestore {
fn drop(&mut self) {
set_codex_client_profile(self.0.clone());
}
}
fn profile_restore_guard() -> (std::sync::MutexGuard<'static, ()>, ProfileRestore) {
let lock = PROFILE_TEST_LOCK.get_or_init(|| Mutex::new(()));
let guard = lock.lock().expect("profile test lock");
let restore = ProfileRestore(codex_client_profile());
(guard, restore)
}
#[test]
fn accepts_only_one_verified_cli_release_for_all_targets() {
let body = serde_json::json!({
"name": "@openai/codex",
"version": "0.200.1",
"optionalDependencies": {
"@openai/codex-darwin-arm64": "npm:@openai/[email protected]",
"@openai/codex-darwin-x64": "npm:@openai/[email protected]",
"@openai/codex-linux-arm64": "npm:@openai/[email protected]",
"@openai/codex-linux-x64": "npm:@openai/[email protected]",
"@openai/codex-win32-arm64": "npm:@openai/[email protected]",
"@openai/codex-win32-x64": "npm:@openai/[email protected]"
}
});
assert_eq!(
parse_cli_release(&serde_json::to_vec(&body).unwrap()).unwrap(),
"0.200.1"
);
}
#[test]
fn rejects_incomplete_platform_release() {
let body = serde_json::json!({
"name": "@openai/codex",
"version": "0.200.1",
"optionalDependencies": {}
});
assert!(parse_cli_release(&serde_json::to_vec(&body).unwrap()).is_err());
}
#[test]
fn refresh_and_fixed_version_environment_policies_are_strict() {
assert!(!refresh_enabled_from(Some("off")));
assert!(!refresh_enabled_from(Some(" FALSE ")));
assert!(refresh_enabled_from(None));
assert_eq!(
fixed_version_from(Some(" 0.200.1 ")).as_deref(),
Some("0.200.1")
);
assert!(fixed_version_from(Some("0.200.1-beta.1")).is_none());
assert!(fixed_version_from(Some("1.2")).is_none());
}
#[test]
fn cached_profile_never_rewinds_active_profile() {
let cached = CachedProfile {
version: "0.200.1".to_string(),
verified_at_unix_secs: 1,
};
assert_eq!(
cached_version_to_restore(&cached, "0.200.0").unwrap(),
Some("0.200.1".to_string())
);
assert_eq!(cached_version_to_restore(&cached, "0.201.0").unwrap(), None);
}
#[tokio::test]
async fn cache_hit_is_restored_without_network_when_refresh_is_disabled() {
let (_lock, _restore) = profile_restore_guard();
let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default());
runtime
.kv_set(
PROFILE_CACHE_KEY,
serde_json::to_string(&CachedProfile {
version: "0.200.1".to_string(),
verified_at_unix_secs: 1,
})
.unwrap(),
Some(Duration::from_secs(60)),
)
.await
.unwrap();
let result = refresh_once_with_fetch(&runtime, None, false, || async {
Err(ProfileRefreshError::HttpStatus(599))
})
.await
.unwrap();
assert_eq!(result, "0.200.1");
assert_eq!(codex_client_version(), "0.200.1");
}
#[tokio::test]
async fn refresh_failure_keeps_previous_profile() {
let (_lock, _restore) = profile_restore_guard();
let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default());
let before = codex_client_profile();
let result = refresh_once_with_fetch(&runtime, None, true, || async {
Err(ProfileRefreshError::HttpStatus(503))
})
.await;
assert!(matches!(result, Err(ProfileRefreshError::HttpStatus(503))));
assert_eq!(codex_client_profile(), before);
}
#[tokio::test]
async fn fixed_version_override_skips_network_and_publishes_profile() {
let (_lock, _restore) = profile_restore_guard();
let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default());
let fetch_called = AtomicBool::new(false);
let result = refresh_once_with_fetch(&runtime, Some("0.220.0"), true, || async {
fetch_called.store(true, Ordering::SeqCst);
Ok("0.221.0".to_string())
})
.await
.unwrap();
assert_eq!(result, "0.220.0");
assert!(!fetch_called.load(Ordering::SeqCst));
assert_eq!(codex_client_version(), "0.220.0");
}
#[tokio::test]
async fn rollback_is_rejected_without_replacing_profile() {
let (_lock, _restore) = profile_restore_guard();
set_codex_cli_version("0.220.0").unwrap();
let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default());
let result =
refresh_once_with_fetch(&runtime, None, true, || async { Ok("0.219.9".to_string()) })
.await;
assert!(matches!(result, Err(ProfileRefreshError::Rollback)));
assert_eq!(codex_client_version(), "0.220.0");
}
}
+2
View File
@@ -140,6 +140,8 @@ pub(crate) const RUST_FRONTDOOR_OWNED_ROUTE_PATTERNS: &[&str] = &[
"/v1beta/models/{model}/operations/{id}",
"/v1beta/operations",
"/v1beta/operations/{id}",
"/openai/v1/videos",
"/openai/v1/videos/{path...}",
"/v1/videos",
"/v1/videos/{path...}",
"/upload/v1beta/files",
@@ -605,6 +605,20 @@ pub(super) fn classify_admin_observability_family_route(
"admin:stats",
false,
))
} else if method == http::Method::GET
&& matches!(
normalized_path,
"/api/admin/stats/leaderboard/user-groups"
| "/api/admin/stats/leaderboard/user-groups/"
)
{
Some(classified(
"admin_proxy",
"stats_manage",
"leaderboard_user_groups",
"admin:stats",
false,
))
} else if method == http::Method::GET
&& matches!(
normalized_path,
+5 -1
View File
@@ -137,7 +137,11 @@ pub(super) fn classify_ai_public_route(
.with_client_surface(detect_claude_client_surface(headers))
.with_api_operation(ApiOperation::ClaudeMessagesCreate),
)
} else if normalized_path.starts_with("/v1/videos") {
} else if normalized_path == "/v1/videos"
|| normalized_path.starts_with("/v1/videos/")
|| normalized_path == "/openai/v1/videos"
|| normalized_path.starts_with("/openai/v1/videos/")
{
Some(classified(
"ai_public",
"openai",
@@ -1,6 +1,7 @@
use http::Uri;
use crate::control::management_token_required_permission;
use crate::control::{management_token_required_permission, GatewayPublicRequestContext};
use crate::handlers::shared::local_proxy_route_requires_buffered_body;
use super::{classify_control_route, headers};
@@ -206,6 +207,10 @@ fn classifies_admin_system_maintenance_write_routes_as_admin_proxy_route() {
"/api/admin/system/important-notification/test",
"important_notification_test",
),
(
"/api/admin/system/cleanup/usage/manual",
"cleanup_usage_manual",
),
("/api/admin/system/cleanup", "cleanup"),
("/api/admin/system/purge/config", "purge_config"),
("/api/admin/system/purge/users", "purge_users"),
@@ -235,6 +240,28 @@ fn classifies_admin_system_maintenance_write_routes_as_admin_proxy_route() {
Some("admin:system")
);
assert!(!decision.is_execution_runtime_candidate());
if matches!(
expected_kind,
"config_import"
| "users_import"
| "data_import"
| "smtp_test"
| "important_notification_test"
| "cleanup_usage_manual"
) {
let context = GatewayPublicRequestContext::from_request_parts(
"trace-system-maintenance-write",
&http::Method::POST,
&uri,
&headers,
Some(decision),
);
assert!(
local_proxy_route_requires_buffered_body(&context),
"POST {path} should buffer request body"
);
}
}
}
@@ -303,6 +330,20 @@ fn classifies_admin_system_update_routes_as_admin_proxy_routes() {
Some("admin:system")
);
assert!(!decision.is_execution_runtime_candidate());
if matches!(expected_kind, "prepare_update" | "apply_update") {
let context = GatewayPublicRequestContext::from_request_parts(
"trace-system-update-write",
&method,
&uri,
&headers,
Some(decision),
);
assert!(
local_proxy_route_requires_buffered_body(&context),
"{method} {path} should buffer request body"
);
}
}
}
@@ -197,6 +197,28 @@ fn classifies_admin_stats_leaderboard_models_as_admin_proxy_route() {
assert!(!decision.is_execution_runtime_candidate());
}
#[test]
fn classifies_admin_stats_leaderboard_user_groups_as_admin_proxy_route() {
let headers = headers(&[]);
let uri: Uri = "/api/admin/stats/leaderboard/user-groups"
.parse()
.expect("uri should parse");
let decision =
classify_control_route(&http::Method::GET, &uri, &headers).expect("route should classify");
assert_eq!(decision.route_class.as_deref(), Some("admin_proxy"));
assert_eq!(decision.route_family.as_deref(), Some("stats_manage"));
assert_eq!(
decision.route_kind.as_deref(),
Some("leaderboard_user_groups")
);
assert_eq!(
decision.auth_endpoint_signature.as_deref(),
Some("admin:stats")
);
assert!(!decision.is_execution_runtime_candidate());
}
#[test]
fn classifies_admin_stats_leaderboard_users_as_admin_proxy_route() {
let headers = headers(&[]);
+7 -5
View File
@@ -62,8 +62,9 @@ pub(crate) use aether_data::repository::users::{
StoredUserPreferenceRecord, StoredUserSessionRecord,
};
use aether_data::repository::wallet::{
AdjustWalletBalanceInput, AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery,
AdminRedeemCodeListQuery, AdminWalletLedgerQuery, AdminWalletListQuery,
AdjustWalletBalanceInBatchInput, AdjustWalletBalanceInput, AdminPaymentOrderListQuery,
AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery,
AdminUserWalletBalanceBatchUserOutcome, AdminWalletLedgerQuery, AdminWalletListQuery,
AdminWalletRefundRequestListQuery, CompareAndSwapPaymentOrderStripeClientSecretInput,
CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput,
CreateAdminRedeemCodeBatchResult, CreateManualWalletRechargeInput,
@@ -72,13 +73,14 @@ use aether_data::repository::wallet::{
CreateWalletRefundRequestOutcome, CreditAdminPaymentOrderInput,
DeleteAdminRedeemCodeBatchInput, DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput,
FailAdminWalletRefundInput, FailWalletRechargeCheckoutInput, InitializeAuthWalletOutcome,
PrepareAdminUserWalletBalanceBatchInput, PrepareAdminUserWalletBalanceBatchOutcome,
ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome,
ReclaimWalletRechargeCheckoutInput, RedeemWalletCodeInput, RedeemWalletCodeOutcome,
StoredAdminPaymentCallback, StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder,
StoredAdminPaymentOrderPage, StoredAdminRedeemCode, StoredAdminRedeemCodeBatch,
StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage, StoredAdminWalletLedgerPage,
StoredAdminWalletListPage, StoredAdminWalletRefund, StoredAdminWalletRefundPage,
StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction,
StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage, StoredAdminUserWalletBalanceBatch,
StoredAdminWalletLedgerPage, StoredAdminWalletListPage, StoredAdminWalletRefund,
StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction,
StoredAdminWalletTransactionPage, StoredWalletDailyUsageLedger,
StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, UpdateAdminWalletRefundGatewayInput,
UpdateWalletRechargeCheckoutInput, WalletLookupKey, WalletMutationOutcome,
+93 -23
View File
@@ -1,9 +1,10 @@
use super::{
read_decision_trace, read_provider_transport_snapshot, read_request_candidate_trace,
AdjustWalletBalanceInput, AdminBillingCollectorRecord, AdminBillingCollectorWriteInput,
AdminBillingMutationOutcome, AdminBillingPresetApplyResult, AdminBillingRuleRecord,
AdminBillingRuleWriteInput, AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery,
AdminRedeemCodeListQuery, AdminWalletLedgerQuery, AdminWalletListQuery,
AdjustWalletBalanceInBatchInput, AdjustWalletBalanceInput, AdminBillingCollectorRecord,
AdminBillingCollectorWriteInput, AdminBillingMutationOutcome, AdminBillingPresetApplyResult,
AdminBillingRuleRecord, AdminBillingRuleWriteInput, AdminPaymentOrderListQuery,
AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery,
AdminUserWalletBalanceBatchUserOutcome, AdminWalletLedgerQuery, AdminWalletListQuery,
AdminWalletRefundRequestListQuery, AnnouncementListQuery, AuditLogListQuery,
BackgroundTaskListQuery, BackgroundTaskSummary, BillingModelContextCacheKey,
BillingModelContextCacheState, BillingModelContextInflightState, BillingPlanRecord,
@@ -17,24 +18,25 @@ use super::{
DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput,
FailWalletRechargeCheckoutInput, GatewayDataState, GatewayProviderTransportSnapshot,
LocalVideoTaskReadResponse, PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord,
PaymentGatewayConfigWriteInput, PaymentGatewaySecretCasUpdate, ProcessAdminWalletRefundInput,
ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, ReclaimWalletRechargeCheckoutInput,
ReconcileUsagePolicyCostInput, RedeemWalletCodeInput, RedeemWalletCodeOutcome,
ReleaseUsagePolicyRequestAdmissionInput, RequestAuditBundle, RequestCandidateTrace,
ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, ReserveUsagePolicyRequestInput,
ReserveUsagePolicyRequestOutcome, StoredAdminAuditLogPage, StoredAdminPaymentCallbackPage,
StoredAdminPaymentOrder, StoredAdminPaymentOrderPage, StoredAdminRedeemCodeBatch,
StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage, StoredAdminWalletLedgerPage,
StoredAdminWalletListPage, StoredAdminWalletRefund, StoredAdminWalletRefundPage,
StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction,
StoredAdminWalletTransactionPage, StoredAnnouncement, StoredAnnouncementPage,
StoredBackgroundTaskEvent, StoredBackgroundTaskRun, StoredBackgroundTaskRunPage,
StoredBillingModelContext, StoredProviderQuotaSnapshot, StoredProviderUsageSummary,
StoredRequestUsageAudit, StoredSuspiciousActivity, StoredUsagePolicyCostReservation,
StoredUsagePolicyRequestAdmission, StoredUsageSettlement, StoredUserAuditLogPage,
StoredUserAuthRecord, StoredUserExportRow, StoredUserSummary, StoredVideoTask,
StoredWalletDailyUsageLedger, StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot,
UpdateAdminWalletRefundGatewayInput, UpdateAnnouncementRecord,
PaymentGatewayConfigWriteInput, PaymentGatewaySecretCasUpdate,
PrepareAdminUserWalletBalanceBatchInput, PrepareAdminUserWalletBalanceBatchOutcome,
ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome,
ReclaimWalletRechargeCheckoutInput, ReconcileUsagePolicyCostInput, RedeemWalletCodeInput,
RedeemWalletCodeOutcome, ReleaseUsagePolicyRequestAdmissionInput, RequestAuditBundle,
RequestCandidateTrace, ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome,
ReserveUsagePolicyRequestInput, ReserveUsagePolicyRequestOutcome, StoredAdminAuditLogPage,
StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, StoredAdminPaymentOrderPage,
StoredAdminRedeemCodeBatch, StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage,
StoredAdminUserWalletBalanceBatch, StoredAdminWalletLedgerPage, StoredAdminWalletListPage,
StoredAdminWalletRefund, StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestPage,
StoredAdminWalletTransaction, StoredAdminWalletTransactionPage, StoredAnnouncement,
StoredAnnouncementPage, StoredBackgroundTaskEvent, StoredBackgroundTaskRun,
StoredBackgroundTaskRunPage, StoredBillingModelContext, StoredProviderQuotaSnapshot,
StoredProviderUsageSummary, StoredRequestUsageAudit, StoredSuspiciousActivity,
StoredUsagePolicyCostReservation, StoredUsagePolicyRequestAdmission, StoredUsageSettlement,
StoredUserAuditLogPage, StoredUserAuthRecord, StoredUserExportRow, StoredUserSummary,
StoredVideoTask, StoredWalletDailyUsageLedger, StoredWalletDailyUsageLedgerPage,
StoredWalletSnapshot, UpdateAdminWalletRefundGatewayInput, UpdateAnnouncementRecord,
UpdateWalletRechargeCheckoutInput, UpsertBackgroundTaskEvent, UpsertBackgroundTaskRun,
UpsertUsageRecord, UpsertVideoTask, UsageSettlementInput, UserDailyQuotaAvailabilityRecord,
UserPlanEntitlementRecord, VideoTaskLookupKey, VideoTaskModelCount, VideoTaskQueryFilter,
@@ -1101,13 +1103,81 @@ impl GatewayDataState {
pub(crate) async fn adjust_wallet_balance(
&self,
input: AdjustWalletBalanceInput,
) -> Result<Option<(StoredWalletSnapshot, StoredAdminWalletTransaction)>, DataLayerError> {
) -> Result<Option<(StoredWalletSnapshot, Option<StoredAdminWalletTransaction>)>, DataLayerError>
{
match &self.wallet_writer {
Some(repository) => repository.adjust_wallet_balance(input).await,
None => Ok(None),
}
}
pub(crate) async fn prepare_admin_user_wallet_balance_batch(
&self,
input: PrepareAdminUserWalletBalanceBatchInput,
) -> Result<Option<PrepareAdminUserWalletBalanceBatchOutcome>, DataLayerError> {
match &self.wallet_writer {
Some(repository) => repository
.prepare_admin_user_wallet_balance_batch(input)
.await
.map(Some),
None => Ok(None),
}
}
pub(crate) async fn get_admin_user_wallet_balance_batch(
&self,
admin_user_id: &str,
idempotency_key: &str,
request_fingerprint: &str,
) -> Result<Option<PrepareAdminUserWalletBalanceBatchOutcome>, DataLayerError> {
match &self.wallet_writer {
Some(repository) => {
repository
.get_admin_user_wallet_balance_batch(
admin_user_id,
idempotency_key,
request_fingerprint,
)
.await
}
None => Ok(None),
}
}
pub(crate) async fn adjust_admin_user_wallet_balance_batch_user(
&self,
input: AdjustWalletBalanceInBatchInput,
) -> Result<Option<AdminUserWalletBalanceBatchUserOutcome>, DataLayerError> {
match &self.wallet_writer {
Some(repository) => repository
.adjust_admin_user_wallet_balance_batch_user(input)
.await
.map(Some),
None => Ok(None),
}
}
pub(crate) async fn record_admin_user_wallet_balance_batch_failure(
&self,
admin_user_id: &str,
idempotency_key: &str,
user_id: &str,
reason: &str,
) -> Result<Option<AdminUserWalletBalanceBatchUserOutcome>, DataLayerError> {
match &self.wallet_writer {
Some(repository) => repository
.record_admin_user_wallet_balance_batch_failure(
admin_user_id,
idempotency_key,
user_id,
reason,
)
.await
.map(Some),
None => Ok(None),
}
}
pub(crate) async fn create_manual_wallet_recharge(
&self,
input: CreateManualWalletRechargeInput,
@@ -123,6 +123,15 @@ impl GatewayDataState {
}
#[cfg(test)]
pub(crate) fn attach_video_task_repository_for_tests<T>(mut self, repository: Arc<T>) -> Self
where
T: VideoTaskRepository + 'static,
{
self.video_task_reader = Some(repository.clone());
self.video_task_writer = Some(repository);
self
}
pub(crate) fn with_video_task_repository_for_tests<T>(repository: Arc<T>) -> Self
where
T: VideoTaskRepository + 'static,
@@ -131,8 +131,13 @@ async fn schedule_pool_page_candidates(
entry.1.insert(candidate.candidate.key_id.clone());
}
let key_context_by_id =
read_pool_catalog_key_contexts_by_id(state, &candidates, provider_model_name).await;
let key_context_by_id = read_pool_catalog_key_contexts_by_id(
state,
&candidates,
provider_model_name,
effective_pool_config,
)
.await;
let mut runtime_by_provider = BTreeMap::new();
let mut pool_config_by_provider = BTreeMap::new();
@@ -634,7 +639,9 @@ impl<'a> PoolKeyCursor<'a> {
if !self.score_phase_exhausted {
if let Some(score_candidates) = self.next_score_candidates().await {
return Some(score_candidates);
if !score_candidates.is_empty() {
return Some(score_candidates);
}
}
}
@@ -1027,6 +1034,26 @@ impl<'a> PoolKeyCursor<'a> {
return None;
}
if pool_config.reserve_minimum_quota
&& admin_provider_pool_pure::admin_pool_key_minimum_quota_reached(
&key,
self.group.candidate.provider_type.as_str(),
Some(self.group.candidate.selected_provider_model_name.as_str()),
)
{
self.seen_key_ids.insert(key.id.clone());
self.record_skip_reason(POOL_ACCOUNT_EXHAUSTED_SKIP_REASON);
self.skipped_candidates
.push(SkippedLocalExecutionCandidate {
candidate: pool_candidate_from_catalog_key(&self.group, key),
skip_reason: POOL_ACCOUNT_EXHAUSTED_SKIP_REASON,
transport: None,
ranking: self.group.ranking.clone(),
extra_data: None,
});
return None;
}
let candidate = pool_candidate_from_catalog_key(&self.group, key);
self.build_eligible_candidate(candidate).await
}
@@ -1427,15 +1454,23 @@ async fn read_pool_catalog_key_contexts_by_id(
state: PlannerAppState<'_>,
candidates: &[EligibleLocalExecutionCandidate],
provider_model_name: Option<&str>,
effective_pool_config: Option<&AdminProviderPoolConfig>,
) -> BTreeMap<String, PoolCatalogKeyContext> {
let mut key_ids = Vec::new();
let mut provider_type_by_key_id = BTreeMap::<String, String>::new();
let mut reserve_minimum_quota_key_ids = BTreeSet::new();
for candidate in candidates {
if pool_config_for_candidate(candidate).is_none() {
let Some(pool_config) = effective_pool_config
.cloned()
.or_else(|| pool_config_for_candidate(candidate))
else {
continue;
}
};
let key_id = candidate.candidate.key_id.clone();
if pool_config.reserve_minimum_quota {
reserve_minimum_quota_key_ids.insert(key_id.clone());
}
if let Entry::Vacant(entry) = provider_type_by_key_id.entry(key_id.clone()) {
entry.insert(candidate.transport.provider.provider_type.clone());
key_ids.push(key_id);
@@ -1487,16 +1522,20 @@ async fn read_pool_catalog_key_contexts_by_id(
.get(&key.id)
.map(String::as_str)
.unwrap_or_default();
(
key.id.clone(),
build_pool_catalog_key_context(
state,
&provider_pool_service,
let mut context = build_pool_catalog_key_context(
state,
&provider_pool_service,
&key,
provider_type,
provider_model_name,
);
context.quota_exhausted |= reserve_minimum_quota_key_ids.contains(&key.id)
&& admin_provider_pool_pure::admin_pool_key_minimum_quota_reached(
&key,
provider_type,
provider_model_name,
),
)
);
(key.id.clone(), context)
})
.collect::<BTreeMap<_, _>>();
// A key can disappear between the candidate-row and catalog reads. Keep
@@ -3962,6 +4001,110 @@ mod tests {
}));
}
#[tokio::test]
async fn pool_key_cursor_reserve_minimum_quota_filters_pages_and_sticky_hits() {
for reserve_enabled in [false, true] {
for sticky in [false, true] {
for used_percent in [99.0, 98.0, 83.0] {
let provider_config = Some(json!({
"pool_advanced": {
"reserve_minimum_quota": reserve_enabled,
"skip_exhausted_accounts": false
}
}));
let provider =
sample_codex_pool_provider("provider-pool", 0, provider_config.clone());
let endpoint = sample_codex_pool_endpoint("provider-pool", "endpoint-1");
let mut reserved = sample_codex_pool_key("provider-pool", "key-low");
reserved.status_snapshot = Some(json!({
"quota": {
"provider_type": "codex",
"updated_at": 100,
"allowed": false,
"exhausted": true,
"code": "exhausted",
"windows": [{
"code": "weekly",
"scope": "account",
"used_ratio": 1.0,
"reset_at": 4_102_444_800u64
}]
}
}));
reserved.upstream_metadata = Some(json!({
"codex": {
"updated_at": 200,
"primary_used_percent": used_percent,
"primary_reset_at": 4_102_444_800u64
}
}));
let ready = sample_codex_pool_key("provider-pool", "key-ready");
let rows = vec![
sample_codex_pool_row("provider-pool", "endpoint-1", "key-low", 0),
sample_codex_pool_row("provider-pool", "endpoint-1", "key-ready", 0),
];
let data_state = GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests(
Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider], vec![endpoint], vec![reserved, ready],
)),
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)),
)
.with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY);
let app = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data_state);
let group =
sample_codex_pool_group("provider-pool", "endpoint-1", 0, provider_config);
let pool_config =
pool_config_for_candidate(&group).expect("pool config should parse");
let sticky_token = sticky.then_some("reserve-session");
if sticky {
record_admin_provider_pool_success(
app.runtime_state.as_ref(),
"provider-pool",
"key-low",
&pool_config,
sticky_token,
0,
None,
)
.await;
}
let mut cursor = PoolKeyCursor::new(
PlannerAppState::new(&app),
group,
sticky_token,
None,
None,
);
cursor.window_size = 1;
cursor.page_size = 1;
let mut returned = Vec::new();
while let Some(candidate) = cursor.next_key().await {
returned.push(candidate.candidate.key_id);
}
let reserve_reached = reserve_enabled && used_percent >= 99.0;
assert_eq!(
returned.contains(&"key-low".to_string()),
!reserve_reached,
"reserve={reserve_enabled}, sticky={sticky}, used={used_percent}"
);
assert!(returned.contains(&"key-ready".to_string()));
if reserve_reached {
assert_eq!(
cursor
.skip_reason_counts
.get(POOL_ACCOUNT_EXHAUSTED_SKIP_REASON),
Some(&1)
);
} else if sticky {
assert_eq!(returned.first().map(String::as_str), Some("key-low"));
}
}
}
}
}
#[tokio::test]
async fn pool_key_cursor_does_not_spend_effective_scan_budget_on_exhausted_accounts() {
let provider_config = Some(json!({
@@ -4168,6 +4311,115 @@ mod tests {
);
}
#[tokio::test]
async fn inactive_pool_key_with_stale_score_does_not_exhaust_pool() {
let provider_config = Some(json!({
"pool_advanced": {
"score_top_n": 128,
"scheduling_presets": [
{"preset": "single_account", "enabled": true},
{"preset": "priority_first", "enabled": true}
]
}
}));
let (provider, endpoint, mut keys, mut rows) =
large_pool_fixture(2, provider_config.clone());
keys[1].is_active = false;
rows.retain(|row| row.key_id != "key-00001");
let scores = vec![
sample_provider_key_pool_score("provider-pool", "key-00000", 5.0),
sample_provider_key_pool_score("provider-pool", "key-00001", 20.0),
];
let data_state =
GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests(
Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
keys,
)),
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)),
)
.with_pool_score_repository_for_tests(Arc::new(
InMemoryPoolMemberScoreRepository::seed(scores),
))
.with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY);
let app = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data_state);
let group = sample_eligible_candidate(
"provider-pool",
"endpoint-1",
"pool-group",
10,
provider_config,
);
let mut cursor = PoolKeyCursor::new(PlannerAppState::new(&app), group, None, None, None);
let candidate = cursor
.next_key()
.await
.expect("active key must stay schedulable beside a stale inactive score");
assert_eq!(candidate.candidate.key_id, "key-00000");
assert_eq!(
cursor.skip_reason_counts.get("pool_score_member_missing"),
Some(&1)
);
}
#[tokio::test]
async fn stale_inactive_score_only_does_not_exhaust_pool() {
let provider_config = Some(json!({
"pool_advanced": {
"score_top_n": 128,
"scheduling_presets": [
{"preset": "single_account", "enabled": true},
{"preset": "priority_first", "enabled": true}
]
}
}));
let (provider, endpoint, mut keys, mut rows) =
large_pool_fixture(2, provider_config.clone());
keys[1].is_active = false;
rows.retain(|row| row.key_id != "key-00001");
let scores = vec![sample_provider_key_pool_score(
"provider-pool",
"key-00001",
20.0,
)];
let data_state =
GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests(
Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
keys,
)),
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)),
)
.with_pool_score_repository_for_tests(Arc::new(
InMemoryPoolMemberScoreRepository::seed(scores),
))
.with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY);
let app = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data_state);
let group = sample_eligible_candidate(
"provider-pool",
"endpoint-1",
"pool-group",
10,
provider_config,
);
let mut cursor = PoolKeyCursor::new(PlannerAppState::new(&app), group, None, None, None);
let candidate = cursor
.next_key()
.await
.expect("catalog rows must remain schedulable when the only score is stale");
assert_eq!(candidate.candidate.key_id, "key-00000");
}
#[tokio::test]
async fn score_candidates_continue_across_pool_windows() {
let provider_config = Some(json!({
@@ -4913,15 +5165,6 @@ mod tests {
))
}
fn provider_catalog_credential_state() -> AppState {
AppState::new()
.expect("credential state should build")
.with_data_state_for_tests(
GatewayDataState::disabled()
.with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY),
)
}
fn large_pool_fixture(
key_count: usize,
provider_config: Option<serde_json::Value>,
@@ -4972,18 +5215,12 @@ mod tests {
)
.expect("endpoint transport should build");
let credential_state = provider_catalog_credential_state();
// 这些用例只验证池扫描、跳过计数和游标预算,不会发起请求或读取凭据。
// 留空凭据可跳过无关的 Fernet 加解密,同时避免复用绑定密文破坏 key_id AAD。
let mut keys = Vec::with_capacity(key_count);
let mut rows = Vec::with_capacity(key_count);
for index in 0..key_count {
let key_id = format!("key-{index:05}");
let encrypted_api_key = credential_state
.seal_provider_catalog_key_api_key(
"provider-pool",
&key_id,
&format!("secret-{index}"),
)
.expect("api key should encrypt");
let mut key = StoredProviderCatalogKey::new(
key_id.clone(),
"provider-pool".to_string(),
@@ -4995,7 +5232,7 @@ mod tests {
.expect("key should build")
.with_transport_fields(
Some(json!(["openai:chat"])),
encrypted_api_key,
None,
None,
None,
None,
@@ -5130,10 +5367,8 @@ mod tests {
.expect("endpoint transport should build")
}
/// 这些测试只检查池调度状态,不涉及凭据解密,因此不构造无关的密文。
fn sample_codex_pool_key(provider_id: &str, key_id: &str) -> StoredProviderCatalogKey {
let encrypted_api_key = provider_catalog_credential_state()
.seal_provider_catalog_key_api_key(provider_id, key_id, &format!("secret-{key_id}"))
.expect("api key should encrypt");
let mut key = StoredProviderCatalogKey::new(
key_id.to_string(),
provider_id.to_string(),
@@ -5145,7 +5380,7 @@ mod tests {
.expect("key should build")
.with_transport_fields(
Some(json!(["openai:responses"])),
encrypted_api_key,
None,
None,
None,
Some(json!({"openai:responses": 1})),
@@ -3199,13 +3199,15 @@ fn openai_responses_body(
let response_id = format!("resp_{}", Uuid::new_v4());
let mut output = Vec::new();
if !collected.thinking.trim().is_empty() {
let thinking = collected.thinking.trim();
output.push(json!({
"id": openai_responses_synthetic_reasoning_item_id(&response_id, 0),
"type": "reasoning",
"status": "completed",
"summary": [{
"type": "summary_text",
"text": collected.thinking.trim(),
"summary": [],
"content": [{
"type": "reasoning_text",
"text": thinking,
}],
}));
}
@@ -4723,6 +4725,15 @@ mod tests {
serde_json::json!(usage.reasoning_tokens)
);
assert_eq!(body["output"][0]["type"], serde_json::json!("reasoning"));
assert_eq!(
body["output"][0]["content"][0]["type"],
serde_json::json!("reasoning_text")
);
assert_eq!(
body["output"][0]["content"][0]["text"],
serde_json::json!("short reasoning")
);
assert_eq!(body["output"][0]["summary"], serde_json::json!([]));
assert_eq!(body["output"][1]["type"], serde_json::json!("message"));
assert!(body["output"][1]["id"]
.as_str()
@@ -4906,7 +4917,12 @@ mod tests {
assert!(body.contains("event: response.created"));
assert!(body.contains("event: response.in_progress"));
assert!(body.contains("event: response.reasoning_summary_part.added"));
// Thinking must stay off the summary channel or clients that render
// both (Codex) print the raw chain-of-thought twice.
assert!(!body.contains("event: response.reasoning_summary_part.added"));
assert!(!body.contains("event: response.reasoning_summary_text.delta"));
assert!(!body.contains("event: response.reasoning_summary_text.done"));
assert!(body.contains("\"type\":\"reasoning_text\""));
assert!(body.contains("event: response.content_part.added"));
assert!(body.contains("event: response.output_text.done"));
assert!(body.contains("event: response.completed"));
@@ -6513,6 +6513,21 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
let normalized_stream_report_context =
normalize_provider_private_report_context(report_context.as_ref());
// Observers follow the live protocol stream across prefetch and transfer.
// Diagnostic capture limits must never determine parser state.
let stream_usage_report_context = normalized_stream_report_context.clone().or_else(|| {
Some(json!({
"provider_api_format": plan.provider_api_format.as_str(),
"client_api_format": plan.client_api_format.as_str(),
}))
});
let mut stream_usage_observer = stream_usage_report_context
.as_ref()
.map(|_| StreamingStandardTerminalObserver::default());
let mut stream_usage_observer_buffered =
StreamUsageObservationBuffer::new(max_stream_body_buffer_bytes);
let mut provider_error_inspection = ProviderStreamErrorInspection::default();
let mut prefetched_provider_error = None;
let upstream_headers = headers.clone();
let mut private_stream_normalizer =
maybe_build_provider_private_stream_normalizer(report_context.as_ref());
@@ -6613,7 +6628,8 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
stream_commit_gate.commit();
}
let mut prefetched_chunks: Vec<Bytes> = Vec::new();
let mut provider_prefetched_body = Vec::new();
let mut provider_prefetched_body = StreamBodyCapture::default();
let mut provider_prefetched_bytes = 0_u64;
let mut provider_prefetched_body_truncated = false;
let mut prefetched_body = Vec::new();
let mut prefetched_inspection_body = Vec::new();
@@ -6866,10 +6882,12 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
}
}
append_stream_capture_bytes(
provider_prefetched_bytes =
provider_prefetched_bytes.saturating_add(chunk.len() as u64);
append_budgeted_stream_capture_bytes(
&mut provider_prefetched_body,
&chunk,
MAX_STREAM_PREFETCH_BYTES,
max_stream_body_buffer_bytes,
&mut provider_prefetched_body_truncated,
);
append_stream_capture_bytes(
@@ -7114,6 +7132,22 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
} else {
chunk
};
if let Some(error) = provider_error_inspection
.observe(stream_usage_report_context.as_ref(), &normalized_chunk)
{
prefetched_provider_error.get_or_insert(error);
}
if let (Some(observer), Some(context)) = (
stream_usage_observer.as_mut(),
stream_usage_report_context.as_ref(),
) {
observe_stream_usage_bytes(
observer,
context,
&mut stream_usage_observer_buffered,
&normalized_chunk,
);
}
let rewritten_chunk = if let Some(rewriter) = local_stream_rewriter.as_mut() {
match rewriter.push_chunk(&normalized_chunk) {
Ok(rewritten_chunk) => rewritten_chunk,
@@ -7244,17 +7278,21 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
if stream_commit_gate.is_uncommitted() {
stream_commit_gate.commit();
}
let prefetched_response_history_persisted = if let Some(record) = local_stream_rewriter
if let Some(record) = local_stream_rewriter
.as_mut()
.and_then(|rewriter| rewriter.take_response_history_record())
{
crate::ai_serving::persist_response_history_record(state, record).await;
true
} else {
false
};
drop(private_stream_normalizer);
drop(local_stream_rewriter);
}
// Keep partial records and conversion state; replaying the bounded
// inspection/capture prefix loses any bytes consumed beyond that prefix.
let mut private_stream_normalizer = private_stream_normalizer.map(|parser| parser.into_owned());
let mut local_stream_rewriter = local_stream_rewriter.map(|parser| parser.into_owned());
if sync_json_stream_bridge_active {
private_stream_normalizer = None;
local_stream_rewriter = None;
stream_usage_observer = None;
}
let initial_usage_telemetry = prefetched_usage_telemetry.clone().or_else(|| {
prefetched_telemetry
@@ -7297,7 +7335,6 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
let headers_for_report = headers.clone();
let report_kind_owned = report_kind;
let report_context_owned = report_context;
let normalized_stream_report_context_owned = normalized_stream_report_context;
let lifecycle_seed_for_report = lifecycle_seed;
let provider_prefetched_body_for_report = provider_prefetched_body;
let prefetched_body_for_report = prefetched_body;
@@ -7339,40 +7376,10 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
let _stream_total_guard =
StageElapsedGuard::from_started_at("stream_total", stream_started_at_for_report);
let _provider_pool_in_flight_guard = provider_pool_in_flight_guard_for_report;
let mut provider_buffered_body = StreamBodyCapture::default();
let mut provider_buffered_body = provider_prefetched_body_for_report;
let mut buffered_body = StreamBodyCapture::default();
let mut provider_body_truncated = false;
let mut provider_body_truncated = provider_prefetched_body_truncated;
let mut client_body_truncated = false;
let mut private_stream_normalizer = if sync_json_stream_bridge_active_for_report {
None
} else {
maybe_build_provider_private_stream_normalizer(report_context_owned.as_ref())
};
let mut local_stream_rewriter = if sync_json_stream_bridge_active_for_report {
None
} else {
maybe_build_stream_response_rewriter(normalized_stream_report_context_owned.as_ref())
};
let stream_usage_report_context =
normalized_stream_report_context_owned.clone().or_else(|| {
Some(serde_json::json!({
"provider_api_format": plan_for_report.provider_api_format.as_str(),
"client_api_format": plan_for_report.client_api_format.as_str(),
}))
});
let mut stream_usage_observer = stream_usage_report_context
.as_ref()
.filter(|_| !sync_json_stream_bridge_active_for_report)
.map(|_| StreamingStandardTerminalObserver::default());
let mut stream_usage_observer_buffered =
StreamUsageObservationBuffer::new(max_stream_body_buffer_bytes);
let mut provider_error_inspection = ProviderStreamErrorInspection::default();
append_budgeted_stream_capture_bytes(
&mut provider_buffered_body,
&provider_prefetched_body_for_report,
max_stream_body_buffer_bytes,
&mut provider_body_truncated,
);
append_budgeted_stream_capture_bytes(
&mut buffered_body,
&prefetched_body_for_report,
@@ -7416,9 +7423,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
} else {
initial_elapsed_ms
}));
let provider_stream_bytes = Arc::new(AtomicU64::new(
u64::try_from(provider_prefetched_body_for_report.len()).unwrap_or(u64::MAX),
));
let provider_stream_bytes = Arc::new(AtomicU64::new(provider_prefetched_bytes));
let client_stream_bytes = Arc::new(AtomicU64::new(
u64::try_from(prefetched_body_for_report.len()).unwrap_or(u64::MAX),
));
@@ -7514,96 +7519,20 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
}
})
};
if !provider_prefetched_body_for_report.is_empty() {
let normalized_prefetched_chunk = if let Some(normalizer) =
private_stream_normalizer.as_mut()
{
match normalizer.push_chunk(&provider_prefetched_body_for_report) {
Ok(normalized_chunk) => Some(normalized_chunk),
Err(err) => {
warn!(
event_name = "stream_execution_prefetch_normalize_restore_failed",
log_type = "ops",
trace_id = %trace_id_owned,
request_id = %request_id_for_report_log,
candidate_id = ?candidate_id_for_report.as_deref(),
error_category = "stream_normalization_restore_failed",
"gateway failed to restore private stream normalization state after prefetch"
);
terminal_failure = Some(build_stream_failure_report(
"execution_runtime_stream_rewrite_error",
format!(
"failed to restore private stream normalization state after prefetch: {err:?}"
),
502,
));
None
}
}
} else {
None
};
let replay_chunk = normalized_prefetched_chunk
.as_deref()
.unwrap_or(provider_prefetched_body_for_report.as_slice());
if let Some(error_body_json) = provider_error_inspection
.observe(stream_usage_report_context.as_ref(), replay_chunk)
{
provider_error_forwarded_to_client = !prefetched_body_for_report.is_empty();
let error_status_code = resolve_provider_stream_error_status_code(
plan_for_report.provider_api_format.as_str(),
status_code,
&error_body_json,
);
terminal_failure = Some(build_stream_failure_from_provider_error_body(
error_status_code,
&error_body_json,
));
}
if let (Some(observer), Some(report_context)) = (
stream_usage_observer.as_mut(),
stream_usage_report_context.as_ref(),
) {
observe_stream_usage_bytes(
observer,
report_context,
&mut stream_usage_observer_buffered,
replay_chunk,
);
}
if terminal_failure.is_none() {
if let Some(rewriter) = local_stream_rewriter.as_mut() {
if let Err(err) = rewriter.push_chunk(replay_chunk) {
warn!(
event_name = "stream_execution_prefetch_rewrite_restore_failed",
log_type = "ops",
trace_id = %trace_id_owned,
request_id = %request_id_for_report_log,
candidate_id = ?candidate_id_for_report.as_deref(),
error_category = "stream_rewrite_restore_failed",
"gateway failed to restore local stream rewrite state after prefetch"
);
terminal_failure = Some(build_stream_failure_report(
"execution_runtime_stream_rewrite_error",
format!(
"failed to restore local stream rewrite state after prefetch: {err:?}"
),
502,
));
}
}
}
if prefetched_response_history_persisted {
if let Some(rewriter) = local_stream_rewriter.as_mut() {
let _ = rewriter.take_response_history_record();
}
}
if let Some(error_body_json) = prefetched_provider_error {
provider_error_forwarded_to_client = !prefetched_body_for_report.is_empty();
let error_status_code = resolve_provider_stream_error_status_code(
plan_for_report.provider_api_format.as_str(),
status_code,
&error_body_json,
);
terminal_failure = Some(build_stream_failure_from_provider_error_body(
error_status_code,
&error_body_json,
));
}
// These buffers restore parser/rewriter state above. Audit capture owns
// its budgeted copies; retaining semantic prefetch duplicates for the
// rest of the stream would bypass the capture memory limit.
drop(provider_prefetched_body_for_report);
// Parser state is already current and capture owns its budgeted bytes.
// This output prefix is needed only to initialize client-side trackers.
drop(prefetched_body_for_report);
if terminal_failure.is_none() && !reached_eof {
@@ -9363,6 +9292,188 @@ mod tests {
.unwrap()
}
#[tokio::test]
async fn prefetch_handoff_preserves_large_responses_setup_event() {
let event = format!(
"event: response.created\ndata: {}\n\n",
json!({"type":"response.created", "response": {
"id":"resp-large-setup", "status":"in_progress", "output":[],
"tools":[{"name":"write", "description":"x".repeat(64 * 1024)}]
}})
);
let done = "event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp-large-setup\",\"status\":\"completed\",\"output\":[],\"usage\":{\"input_tokens\":7,\"output_tokens\":2}}}\n\n";
// Include the two observed transport boundaries, exact/near budget
// boundaries, and multiple prefetch chunks crossing the budget.
for cuts in [
vec![16_383],
vec![16_384],
vec![17_735],
vec![17_741],
vec![8_192, 17_735],
] {
let mut chunks = Vec::new();
let mut start = 0;
for end in cuts {
chunks.push(&event[start..end]);
start = end;
}
chunks.push(&event[start..]);
chunks.push(done);
let response = execute_generic_sse_precommit(chunks, json!({}), None, false)
.await
.expect("large setup should commit at the bounded prefetch limit");
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
let body = String::from_utf8(body.to_vec()).unwrap();
assert!(
body.starts_with(&event),
"setup bytes lost or duplicated at split {start}"
);
let events: Vec<Value> = body
.lines()
.filter_map(|line| line.strip_prefix("data: "))
.filter(|payload| *payload != "[DONE]")
.map(|payload| {
serde_json::from_str(payload).expect("every SSE payload must be valid JSON")
})
.collect();
assert_eq!(events.len(), 2, "events must be forwarded exactly once");
assert_eq!(events[1]["type"], "response.completed");
}
}
#[tokio::test]
async fn prefetch_handoff_keeps_audit_usage_and_private_conversion() {
for private in [false, true] {
let request_id = format!("handoff-audit-{}", uuid::Uuid::new_v4());
let mut plan = if private {
antigravity_gemini_stream_plan(&request_id)
} else {
native_anthropic_stream_plan(&request_id)
};
if !private {
plan.provider_api_format = "openai:responses".into();
plan.client_api_format = "openai:responses".into();
}
let context = json!({
"request_id": request_id, "candidate_id": plan.candidate_id,
"candidate_index":0, "retry_index":0,
"provider_api_format": plan.provider_api_format,
"client_api_format": plan.client_api_format,
"needs_conversion": private, "has_envelope": private,
"envelope_name": if private { "antigravity:v1internal" } else { "" },
});
let repository = Arc::new(InMemoryUsageReadRepository::default());
let catalog = provider_catalog_for_plan(&plan, None);
let state = AppState::new()
.unwrap()
.with_data_state_for_tests(
crate::data::GatewayDataState::with_usage_repository_for_tests(Arc::clone(
&repository,
))
.with_provider_catalog_reader(Arc::new(catalog))
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
.with_system_config_values_for_tests([(
"request_record_level".into(),
json!("full"),
)]),
)
.with_usage_runtime_for_tests(UsageRuntimeConfig {
enabled: true,
..Default::default()
});
let text = "hello".repeat(12_000);
let payload = if private {
json!({"response":{"candidates":[{"content":{"role":"model","parts":[{"text":text}]},
"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":1234,"candidatesTokenCount":567},
"modelVersion":"gemini-3.7-flash-tiered"}})
} else {
json!({"type":"response.completed","response":{"id":"resp-handoff-usage","status":"completed",
"output":[{"type":"message","id":"msg-handoff","role":"assistant","status":"completed",
"content":[{"type":"output_text","text":text,"annotations":[]}]}],
"usage":{"input_tokens":1234,"output_tokens":567,"total_tokens":1801}}})
};
let input = format!("data: {payload}\n\n");
// One complete large chunk exercises an already-emitted prefetch
// result; the private path exercises incomplete normalization too.
let chunks = if private {
vec![input[..17_735].to_string(), input[17_735..].to_string()]
} else {
vec![input.clone()]
};
let frames = stream! {
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
frame_type:StreamFrameType::Headers,
payload:StreamFramePayload::Headers { status_code:200,
headers:BTreeMap::from([("content-type".into(),"text/event-stream".into())]),
response_observation:None },
}));
for chunk in chunks {
yield Ok(ndjson_frame(StreamFrame { frame_type:StreamFrameType::Data,
payload:StreamFramePayload::Data { text:Some(chunk),chunk_b64:None } }));
}
yield Ok(ndjson_frame(StreamFrame::eof()));
}.boxed();
let response = execute_stream_from_frame_stream(
&state,
plan,
"trace-handoff-audit",
&test_decision(),
OPENAI_RESPONSES_STREAM_PLAN_KIND,
Some("openai_responses_stream_success".into()),
Some(context),
crate::clock::current_unix_ms(),
Instant::now(),
RequestStageTrace::from_env(),
false,
frames,
None,
)
.await
.unwrap()
.unwrap();
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
let body = String::from_utf8(body.to_vec()).unwrap();
let events: Vec<Value> = body
.lines()
.filter_map(|l| l.strip_prefix("data: "))
.filter(|p| *p != "[DONE]")
.map(|p| serde_json::from_str(p).unwrap())
.collect();
assert_eq!(
events
.iter()
.filter(|e| e["type"] == "response.completed")
.count(),
1
);
assert!(body.contains(&text));
let usage = tokio::time::timeout(Duration::from_secs(3), async {
loop {
if let Some(u) = repository
.find_by_request_id(&request_id)
.await
.unwrap()
.filter(|u| u.status == "completed" || u.status == "failed")
{
break u;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("usage should finalize");
assert_eq!(usage.status, "completed", "{:?}", usage.error_message);
assert_eq!(usage.input_tokens, 1234);
assert_eq!(usage.output_tokens, 567);
let captured = usage.response_body.as_ref().expect("provider capture");
assert!(
captured["metadata"].get("dropped_chunks").is_none(),
"{captured}"
);
assert_eq!(captured["chunks"].as_array().unwrap(), &vec![payload]);
}
}
#[tokio::test]
async fn generic_stream_success_regex_matches_fragmented_plain_body() {
for chunks in [
@@ -9901,7 +10012,7 @@ mod tests {
let mut buffer = super::StreamUsageObservationBuffer::new(32 * 1024);
let mut rewriter = super::maybe_build_stream_response_rewriter(Some(&context)).unwrap();
let mut delivered = Vec::new();
for chunk in chunks {
for (index, chunk) in chunks.into_iter().enumerate() {
provider.append(chunk, 32 * 1024, &mut provider_truncated);
super::observe_stream_usage_bytes(
observer.as_mut().unwrap(),
@@ -9912,6 +10023,10 @@ mod tests {
let output = rewriter.push_chunk(chunk).unwrap();
client.append(&output, 32 * 1024, &mut client_truncated);
delivered.extend(output);
if index == 0 {
// Task handoff must also work when audit admits no bytes.
rewriter = rewriter.into_owned();
}
}
let tail = rewriter.finish().unwrap();
client.append(&tail, 32 * 1024, &mut client_truncated);
@@ -11931,7 +12046,11 @@ mod tests {
.expect("response body should read");
let body = String::from_utf8(body.to_vec()).expect("response body should be utf8");
assert!(
body.contains("event: response.reasoning_summary_text.delta\n"),
body.contains("event: response.reasoning_text.delta\n"),
"{body}"
);
assert!(
!body.contains("event: response.reasoning_summary_text.delta\n"),
"{body}"
);
assert!(
@@ -196,6 +196,78 @@ fn maybe_build_invalid_provider_success_finalize_response(
)?))
}
fn local_sync_needs_conversion(payload: &GatewaySyncReportRequest) -> bool {
payload
.report_context
.as_ref()
.and_then(|value| value.get("needs_conversion"))
.and_then(|value| value.as_bool())
.unwrap_or(false)
}
/// A successful upstream response that needed conversion but could not be
/// converted must not reach the client in the provider's own format.
fn maybe_build_unconverted_cross_format_success_response(
trace_id: &str,
decision: &GatewayControlDecision,
payload: &GatewaySyncReportRequest,
) -> Result<Option<Response<Body>>, GatewayError> {
if payload.status_code >= 400
|| !local_sync_needs_conversion(payload)
|| !is_core_error_finalize_kind(payload.report_kind.as_str())
{
return Ok(None);
}
let client_api_format = resolve_local_sync_client_api_format(payload);
let provider_api_format = resolve_local_sync_provider_api_format(payload);
warn!(
event_name = "local_core_finalize_cross_format_success_unconverted",
log_type = "event",
trace_id = %trace_id,
report_kind = %payload.report_kind,
status_code = payload.status_code,
client_api_format = %client_api_format,
provider_api_format = %provider_api_format,
"gateway could not convert a successful provider response to the client format"
);
let message = format!(
"Provider returned HTTP {} but its {provider_api_format} response could not be converted to {client_api_format}.",
payload.status_code
);
let body_json = build_core_error_body_for_client_format(
&client_api_format,
&message,
Some("response_conversion_failed"),
LocalCoreSyncErrorKind::ServerError,
)
.unwrap_or_else(|| {
serde_json::json!({
"error": {
"message": message,
"type": "server_error",
"code": "response_conversion_failed"
}
})
});
let mut response_headers = payload.headers.clone();
response_headers.remove("content-encoding");
response_headers.remove("content-length");
response_headers.insert("content-type".to_string(), "application/json".to_string());
let body_bytes =
serde_json::to_vec(&body_json).map_err(|err| GatewayError::Internal(err.to_string()))?;
response_headers.insert("content-length".to_string(), body_bytes.len().to_string());
Ok(Some(build_client_response_from_parts(
StatusCode::BAD_GATEWAY.as_u16(),
&response_headers,
Body::from(body_bytes),
trace_id,
Some(decision),
)?))
}
fn local_core_sync_finalize_has_invalid_provider_success(
payload: &GatewaySyncReportRequest,
) -> Result<bool, GatewayError> {
@@ -274,6 +346,12 @@ pub(crate) fn resolve_local_core_error_response_body_json(
return Ok(Some(body_json));
}
// A 2xx cross-format body that is not JSON (e.g. an aggregated SSE capture)
// carries no upstream error; wrapping it as one would ship raw provider
// bytes to the client under the success status.
if payload.status_code < 400 && local_sync_needs_conversion(payload) {
return Ok(None);
}
let Some(body_text) = decode_local_sync_body_text(payload)? else {
return Ok(None);
};
@@ -626,6 +704,10 @@ pub(crate) async fn submit_local_core_error_or_sync_finalize(
maybe_build_local_core_error_response(trace_id, decision, &payload)?
{
response
} else if let Some(response) =
maybe_build_unconverted_cross_format_success_response(trace_id, decision, &payload)?
{
response
} else {
warn!(
event_name = "local_core_finalize_fallback_raw_response_body",
@@ -937,6 +1019,128 @@ mod tests {
);
}
#[tokio::test]
async fn local_core_sync_finalize_converts_forced_responses_stream_for_gemini_client() {
use base64::Engine as _;
// Forced-stream xAI shape: the terminal response echoes request
// metadata and encrypted reasoning next to the real answer.
let raw_sse = concat!(
"event: response.created\n",
"data: {\"type\":\"response.created\",\"sequence_number\":0,\"response\":{\"id\":\"resp_xai_123\",\"object\":\"response\",\"status\":\"in_progress\",\"model\":\"grok-4.7-build\",\"output\":[],\"parallel_tool_calls\":true,\"tools\":[]}}\n\n",
"event: response.output_item.done\n",
"data: {\"type\":\"response.output_item.done\",\"sequence_number\":1,\"output_index\":0,\"item\":{\"id\":\"rs_xai_123\",\"type\":\"reasoning\",\"status\":\"completed\",\"summary\":[],\"encrypted_content\":\"opaque-xai-reasoning\"}}\n\n",
"event: response.output_text.delta\n",
"data: {\"type\":\"response.output_text.delta\",\"sequence_number\":2,\"item_id\":\"msg_xai_123\",\"output_index\":1,\"content_index\":0,\"delta\":\"Hi there, friend\"}\n\n",
"event: response.output_item.done\n",
"data: {\"type\":\"response.output_item.done\",\"sequence_number\":3,\"output_index\":1,\"item\":{\"id\":\"msg_xai_123\",\"type\":\"message\",\"status\":\"completed\",\"role\":\"assistant\",\"content\":[{\"type\":\"output_text\",\"text\":\"Hi there, friend\",\"annotations\":[]}]}}\n\n",
"event: response.completed\n",
"data: {\"type\":\"response.completed\",\"sequence_number\":4,\"response\":{\"id\":\"resp_xai_123\",\"object\":\"response\",\"status\":\"completed\",\"model\":\"grok-4.7-build\",\"output\":[],\"parallel_tool_calls\":true,\"tool_choice\":\"auto\",\"tools\":[],\"text\":{\"format\":{\"type\":\"text\"}},\"temperature\":0.7,\"store\":false,\"usage\":{\"input_tokens\":1249,\"output_tokens\":12,\"total_tokens\":1261}}}\n\n",
);
let mut payload = core_finalize_payload(
"gemini_chat_sync_finalize",
"gemini:generate_content",
"openai:responses",
200,
json!(null),
);
payload.body_json = None;
payload.body_base64 = Some(base64::engine::general_purpose::STANDARD.encode(raw_sse));
payload.report_context = Some(json!({
"client_api_format": "gemini:generate_content",
"provider_api_format": "openai:responses",
"provider_stream_event_api_format": "openai:responses",
"model": "grok-4.7",
"mapped_model": "grok-4.7",
"needs_conversion": true,
}));
let state = AppState::new().expect("state should build");
let response = submit_local_core_error_or_sync_finalize(
&state,
"trace-forced-responses-gemini",
&test_decision(),
payload,
)
.await
.expect("finalize should build a response");
assert_eq!(response.status(), http::StatusCode::OK);
let body_bytes = to_bytes(response.into_body(), usize::MAX)
.await
.expect("body should read");
let body =
serde_json::from_slice::<serde_json::Value>(&body_bytes).expect("body should decode");
assert!(body.get("error").is_none(), "unexpected error body: {body}");
let parts = body["candidates"][0]["content"]["parts"]
.as_array()
.expect("gemini parts");
assert!(parts.iter().any(|part| part["text"] == "Hi there, friend"));
let text = String::from_utf8_lossy(&body_bytes);
assert!(!text.contains("opaque-xai-reasoning") && !text.contains("response.created"));
}
#[tokio::test]
async fn local_core_sync_finalize_never_wraps_unconvertible_success_sse_as_client_error() {
use base64::Engine as _;
// A complete stream whose output the Gemini client cannot represent.
let raw_sse = concat!(
"event: response.output_item.done\n",
"data: {\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"id\":\"future_item_123\",\"type\":\"future_output\",\"payload\":\"must-not-drop\"}}\n\n",
"event: response.completed\n",
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_raw_123\",\"object\":\"response\",\"status\":\"completed\",\"model\":\"grok-4.7\",\"output\":[]}}\n\n",
);
let mut payload = core_finalize_payload(
"gemini_chat_sync_finalize",
"gemini:generate_content",
"openai:responses",
200,
json!(null),
);
payload.body_json = None;
payload.body_base64 = Some(base64::engine::general_purpose::STANDARD.encode(raw_sse));
payload.report_context = Some(json!({
"client_api_format": "gemini:generate_content",
"provider_api_format": "openai:responses",
"provider_stream_event_api_format": "openai:responses",
"needs_conversion": true,
}));
assert!(maybe_build_local_core_error_response(
"trace-raw-success-sse",
&test_decision(),
&payload,
)
.expect("response build should not error")
.is_none());
let state = AppState::new().expect("state should build");
let response = submit_local_core_error_or_sync_finalize(
&state,
"trace-raw-success-sse",
&test_decision(),
payload,
)
.await
.expect("finalize should build a response");
assert_eq!(response.status(), http::StatusCode::BAD_GATEWAY);
let body = serde_json::from_slice::<serde_json::Value>(
&to_bytes(response.into_body(), usize::MAX)
.await
.expect("body should read"),
)
.expect("body should decode");
let message = body["error"]["message"]
.as_str()
.expect("error message should exist");
assert!(
message.contains("could not be converted") && !message.contains("must-not-drop"),
"unexpected message: {message}"
);
}
#[tokio::test]
async fn submit_local_core_finalize_keeps_http_200_for_success_image_body() {
let payload = core_finalize_payload(
@@ -9737,7 +9737,10 @@ mod tests {
headers: BTreeMap::from([("content-type".into(), "application/json".into())]),
content_type: Some("application/json".into()),
content_encoding: Some(encoding.into()),
body: RequestBody::from_json(json!({"model": "gpt-4.1"})),
body: RequestBody::from_json(json!({
"model": "gpt-4.1",
"service_tier": "ultrafast"
})),
stream: false,
client_api_format: "openai:chat".into(),
provider_api_format: "openai:chat".into(),
@@ -9758,7 +9761,7 @@ mod tests {
result.body.and_then(|body| body.json_body),
Some(json!({
"content_encoding": encoding,
"body": {"model": "gpt-4.1"},
"body": {"model": "gpt-4.1", "service_tier": "ultrafast"},
}))
);
}
@@ -1464,6 +1464,22 @@ pub(crate) async fn maybe_execute_sync_via_local_video_decision(
.await
}
fn supports_local_video_get(
parts: &http::request::Parts,
decision: &GatewayControlDecision,
) -> bool {
parts.method == http::Method::GET
&& decision.route_kind.as_deref() == Some("video")
&& (crate::video_tasks::resolve_video_task_read_lookup_key(
decision.route_family.as_deref(),
parts.uri.path(),
)
.is_some()
|| (decision.route_family.as_deref() == Some("openai")
&& crate::video_tasks::extract_openai_task_id_from_content_path(parts.uri.path())
.is_some()))
}
pub(crate) fn maybe_execute_sync_request<'a>(
state: &'a AppState,
parts: &'a http::request::Parts,
@@ -1477,7 +1493,7 @@ pub(crate) fn maybe_execute_sync_request<'a>(
};
#[cfg(not(test))]
{
if parts.method != http::Method::POST {
if parts.method != http::Method::POST && !supports_local_video_get(parts, decision) {
return Ok(LocalExecutionRequestOutcome::NoPath);
}
return maybe_execute_sync_local_path(state, parts, body_bytes, trace_id, decision)
@@ -1490,6 +1506,7 @@ pub(crate) fn maybe_execute_sync_request<'a>(
.unwrap_or_default()
.is_empty()
&& parts.method != http::Method::POST
&& !supports_local_video_get(parts, decision)
{
return Ok(LocalExecutionRequestOutcome::NoPath);
}
@@ -1511,7 +1528,7 @@ pub(crate) fn maybe_execute_stream_request<'a>(
};
#[cfg(not(test))]
{
if parts.method != http::Method::POST {
if parts.method != http::Method::POST && !supports_local_video_get(parts, decision) {
return Ok(LocalExecutionRequestOutcome::NoPath);
}
return maybe_execute_stream_local_path(state, parts, body_bytes, trace_id, decision)
@@ -1524,6 +1541,7 @@ pub(crate) fn maybe_execute_stream_request<'a>(
.unwrap_or_default()
.is_empty()
&& parts.method != http::Method::POST
&& !supports_local_video_get(parts, decision)
{
return Ok(LocalExecutionRequestOutcome::NoPath);
}
@@ -32,6 +32,10 @@ fn request_has_execution_runtime_via_guard(headers: &HeaderMap) -> bool {
}
pub(crate) fn frontdoor_self_loop_public_ai_path(path: &str) -> bool {
let path = path
.strip_prefix("/openai")
.filter(|p| *p == "/v1/videos" || p.starts_with("/v1/videos/"))
.unwrap_or(path);
matches!(
path,
"/v1/messages"
@@ -63,13 +63,14 @@ pub(in super::super) async fn build_admin_wallet_adjust_response(
}
let operator_id = admin_wallet_operator_id(request_context);
let has_wallet_writer = state.has_wallet_data_writer();
let Some((wallet, transaction)) = state
let Some((wallet, Some(transaction))) = state
.admin_adjust_wallet_balance(
&wallet_id,
amount_usd,
&balance_type,
operator_id.as_deref(),
description.as_deref(),
false,
)
.await?
else {
@@ -12,3 +12,138 @@ pub(crate) use self::stats::{
};
pub(crate) use self::stats::{AdminStatsTimeRange, AdminStatsUsageFilter};
pub(crate) use self::usage::maybe_build_local_admin_usage_response;
pub(crate) async fn resolve_usage_user_group_scope(
state: &crate::handlers::admin::request::AdminAppState<'_>,
query: Option<&str>,
include_inactive: bool,
exclude_admin: bool,
) -> Result<Result<Option<Vec<String>>, String>, crate::GatewayError> {
let group_id = crate::handlers::admin::shared::query_param_value(query, "user_group_id");
let Some(group_id) = group_id else {
return Ok(Ok(None));
};
if crate::handlers::admin::shared::query_param_value(query, "user_id").is_some() {
return Ok(Err(
"user_id and user_group_id cannot be used together".to_string()
));
}
if !state.has_user_data_reader() {
return Ok(Err("user group data is unavailable".to_string()));
}
if group_id == UNGROUPED_USAGE_ID {
let ids = ungrouped_usage_users(state)
.await?
.into_iter()
.filter(|user| include_inactive || user.is_active)
.filter(|user| !exclude_admin || !user.role.eq_ignore_ascii_case("admin"))
.map(|user| user.id)
.collect();
return Ok(Ok(Some(ids)));
}
match state
.resolve_usage_user_group_member_ids(&group_id, include_inactive, exclude_admin)
.await?
{
Some(user_ids) => Ok(Ok(Some(user_ids))),
None => Ok(Err("user_group_id does not exist".to_string())),
}
}
/// Reserved statistics-only scope; never a permission group.
pub(crate) const UNGROUPED_USAGE_ID: &str = "__ungrouped__";
pub(crate) async fn ungrouped_usage_users(
state: &crate::handlers::admin::request::AdminAppState<'_>,
) -> Result<Vec<aether_data::repository::users::StoredUserSummary>, crate::GatewayError> {
use aether_data::repository::users::UserExportListQuery;
let mut users = Vec::new();
let mut skip = 0;
loop {
let page = state
.list_export_users_page(&UserExportListQuery {
skip,
limit: 500,
..Default::default()
})
.await?;
let count = page.len();
if count == 0 {
break;
}
let ids = page.into_iter().map(|user| user.id).collect::<Vec<_>>();
let grouped = state
.list_user_group_memberships_by_user_ids(&ids)
.await?
.into_iter()
.map(|membership| membership.user_id)
.collect::<std::collections::BTreeSet<_>>();
let ids = ids
.into_iter()
.filter(|id| !grouped.contains(id))
.collect::<Vec<_>>();
users.extend(
state
.list_users_by_ids(&ids)
.await?
.into_iter()
.filter(|user| !user.is_deleted),
);
skip += count;
if count < 500 {
break;
}
}
Ok(users)
}
/// Current group provider policy, resolved to the provider-name dimension used by usage rollups.
/// None is unrestricted; Some(empty) deliberately matches no usage.
pub(crate) async fn usage_group_provider_names(
state: &crate::handlers::admin::request::AdminAppState<'_>,
group: &aether_data::repository::users::StoredUserGroup,
) -> Result<Option<Vec<String>>, crate::GatewayError> {
if matches!(
group.allowed_providers_mode.as_str(),
"unrestricted" | "inherit"
) {
return Ok(None);
}
if group.allowed_providers_mode != "specific" {
return Ok(Some(Vec::new()));
}
let allowed = group.allowed_providers.as_deref().unwrap_or_default();
let providers = state.list_provider_catalog_providers(false).await?;
let mut names = providers
.into_iter()
.filter(|provider| {
allowed.iter().any(|value| {
let value = value.trim();
value.eq_ignore_ascii_case(&provider.id)
|| value.eq_ignore_ascii_case(&provider.name)
|| value.eq_ignore_ascii_case(&provider.provider_type)
})
})
.map(|provider| provider.name)
.collect::<Vec<_>>();
names.sort();
names.dedup();
Ok(Some(names))
}
pub(crate) async fn resolve_usage_group_provider_names(
state: &crate::handlers::admin::request::AdminAppState<'_>,
query: Option<&str>,
) -> Result<Option<Vec<String>>, crate::GatewayError> {
let Some(id) = crate::handlers::admin::shared::query_param_value(query, "user_group_id") else {
return Ok(None);
};
if id == UNGROUPED_USAGE_ID {
return Ok(None);
}
let Some(group) = state.find_user_group_by_id(&id).await? else {
return Ok(Some(Vec::new()));
};
usage_group_provider_names(state, &group).await
}
@@ -169,9 +169,11 @@ pub(super) async fn build_admin_monitoring_system_status_response(
let today_usage = state
.summarize_usage_audits(&UsageAuditSummaryQuery {
provider_names: None,
created_from_unix_secs: today_start.timestamp().max(0) as u64,
created_until_unix_secs: now_unix_secs.saturating_add(1),
user_id: None,
user_ids: None,
provider_name: None,
model: None,
})
@@ -1,3 +1,4 @@
use super::super::resolve_usage_user_group_scope;
use super::range::{build_comparison_range, parse_bounded_u32};
use super::resolve_admin_usage_time_range;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
@@ -109,6 +110,7 @@ pub(super) async fn maybe_build_local_admin_stats_analytics_response(
};
let current_summary = state
.summarize_usage_audits(&UsageAuditSummaryQuery {
provider_names: None,
created_from_unix_secs: current_from_unix_secs,
created_until_unix_secs: current_until_unix_secs,
..Default::default()
@@ -116,6 +118,7 @@ pub(super) async fn maybe_build_local_admin_stats_analytics_response(
.await?;
let comparison_summary = state
.summarize_usage_audits(&UsageAuditSummaryQuery {
provider_names: None,
created_from_unix_secs: comparison_from_unix_secs,
created_until_unix_secs: comparison_until_unix_secs,
..Default::default()
@@ -294,6 +297,17 @@ pub(super) async fn maybe_build_local_admin_stats_analytics_response(
}
let filters = AdminStatsUsageFilter::from_query(request_context.query_string());
let user_ids = match resolve_usage_user_group_scope(
state,
request_context.query_string(),
false,
false,
)
.await?
{
Ok(value) => value,
Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))),
};
let query_granularity = match granularity {
AdminStatsGranularity::Hour => UsageTimeSeriesGranularity::Hour,
AdminStatsGranularity::Day
@@ -306,11 +320,17 @@ pub(super) async fn maybe_build_local_admin_stats_analytics_response(
};
let buckets = state
.summarize_usage_time_series(&UsageTimeSeriesQuery {
provider_names: super::super::resolve_usage_group_provider_names(
state,
request_context.query_string(),
)
.await?,
created_from_unix_secs,
created_until_unix_secs,
granularity: query_granularity,
tz_offset_minutes: time_range.tz_offset_minutes,
user_id: filters.user_id,
user_ids,
provider_name: filters.provider_name,
model: filters.model,
})
@@ -72,11 +72,13 @@ pub(super) async fn maybe_build_local_admin_stats_cost_response(
};
let buckets = state
.summarize_usage_time_series(&UsageTimeSeriesQuery {
provider_names: None,
created_from_unix_secs,
created_until_unix_secs,
granularity: UsageTimeSeriesGranularity::Day,
tz_offset_minutes: time_range.tz_offset_minutes,
user_id: None,
user_ids: None,
provider_name: None,
model: None,
})
@@ -3,12 +3,13 @@ use crate::GatewayError;
use aether_data_contracts::repository::usage::StoredRequestUsageAudit;
pub(super) use aether_admin::observability::stats::{
build_admin_stats_leaderboard_response, build_api_key_leaderboard_items,
build_api_key_leaderboard_items_from_summaries, build_model_leaderboard_items,
build_model_leaderboard_items_from_summaries, build_user_leaderboard_items,
build_user_leaderboard_items_from_summaries, compare_leaderboard_items, compute_dense_rank,
AdminStatsLeaderboardItem, AdminStatsLeaderboardMetric, AdminStatsLeaderboardNameMode,
AdminStatsSortOrder, AdminStatsUserMetadata,
build_admin_stats_leaderboard_response, build_admin_stats_user_group_leaderboard_response,
build_api_key_leaderboard_items, build_api_key_leaderboard_items_from_summaries,
build_model_leaderboard_items, build_model_leaderboard_items_from_summaries,
build_user_leaderboard_items, build_user_leaderboard_items_from_summaries,
compare_leaderboard_items, compute_dense_rank, AdminStatsLeaderboardItem,
AdminStatsLeaderboardMetric, AdminStatsLeaderboardNameMode, AdminStatsSortOrder,
AdminStatsUserMetadata,
};
pub(super) async fn load_user_leaderboard_metadata(
@@ -1,7 +1,9 @@
use super::super::resolve_usage_user_group_scope;
use super::leaderboard::{
build_admin_stats_leaderboard_response, build_api_key_leaderboard_items_from_summaries,
build_model_leaderboard_items_from_summaries, build_user_leaderboard_items_from_summaries,
compare_leaderboard_items, load_user_leaderboard_metadata, AdminStatsLeaderboardNameMode,
build_admin_stats_leaderboard_response, build_admin_stats_user_group_leaderboard_response,
build_api_key_leaderboard_items_from_summaries, build_model_leaderboard_items_from_summaries,
build_user_leaderboard_items_from_summaries, compare_leaderboard_items,
load_user_leaderboard_metadata, AdminStatsLeaderboardItem, AdminStatsLeaderboardNameMode,
};
use super::range::{parse_bounded_u32, parse_nonnegative_usize};
use super::resolve_admin_usage_time_range;
@@ -14,6 +16,7 @@ use aether_admin::observability::stats::{
};
use aether_data_contracts::repository::usage::{UsageLeaderboardGroupBy, UsageLeaderboardQuery};
use axum::{body::Body, http, response::Response};
use std::collections::{BTreeMap, BTreeSet};
pub(super) async fn maybe_build_local_admin_stats_leaderboard_response(
state: &AdminAppState<'_>,
@@ -75,10 +78,12 @@ pub(super) async fn maybe_build_local_admin_stats_leaderboard_response(
};
let summaries = state
.summarize_usage_leaderboard(&UsageLeaderboardQuery {
provider_names: None,
created_from_unix_secs,
created_until_unix_secs,
group_by: UsageLeaderboardGroupBy::Model,
user_id: filters.user_id,
user_ids: None,
provider_name: filters.provider_name,
model: filters.model,
})
@@ -152,10 +157,12 @@ pub(super) async fn maybe_build_local_admin_stats_leaderboard_response(
};
let summaries = state
.summarize_usage_leaderboard(&UsageLeaderboardQuery {
provider_names: None,
created_from_unix_secs,
created_until_unix_secs,
group_by: UsageLeaderboardGroupBy::ApiKey,
user_id: filters.user_id,
user_ids: None,
provider_name: filters.provider_name,
model: filters.model,
})
@@ -206,6 +213,181 @@ pub(super) async fn maybe_build_local_admin_stats_leaderboard_response(
)));
}
if request_context
.decision()
.and_then(|decision| decision.route_kind.as_deref())
== Some("leaderboard_user_groups")
&& request_context.method() == http::Method::GET
&& matches!(
request_context.path(),
"/api/admin/stats/leaderboard/user-groups"
| "/api/admin/stats/leaderboard/user-groups/"
)
{
let time_range = match resolve_admin_usage_time_range(query) {
Ok(value) => value,
Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))),
};
let metric = match AdminStatsLeaderboardMetric::parse(query) {
Ok(value) => value,
Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))),
};
let order = match AdminStatsSortOrder::parse(query) {
Ok(value) => value,
Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))),
};
let limit = match query_param_value(query, "limit")
.map(|value| parse_bounded_u32("limit", &value, 1, 100))
.transpose()
{
Ok(Some(value)) => value as usize,
Ok(None) => 10,
Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))),
};
let offset = match query_param_value(query, "offset")
.map(|value| parse_nonnegative_usize("offset", &value))
.transpose()
{
Ok(Some(value)) => value,
Ok(None) => 0,
Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))),
};
let empty_counts = BTreeMap::new();
if !state.has_usage_data_reader() || !state.has_user_data_reader() {
return Ok(Some(build_admin_stats_user_group_leaderboard_response(
metric,
Some(&time_range),
&[],
&empty_counts,
&empty_counts,
offset,
limit,
)));
}
let include_inactive = query_param_bool(query, "include_inactive", false);
let exclude_admin = query_param_bool(query, "exclude_admin", false);
let filters = AdminStatsUsageFilter::from_query(query);
if filters.user_id.is_some() {
return Ok(Some(admin_stats_bad_request_response(
"user_id is not supported for the user group leaderboard".to_string(),
)));
}
let Some((created_from_unix_secs, created_until_unix_secs)) = time_range.to_unix_bounds()
else {
return Ok(Some(build_admin_stats_user_group_leaderboard_response(
metric,
Some(&time_range),
&[],
&empty_counts,
&empty_counts,
offset,
limit,
)));
};
let mut leaderboard = Vec::new();
let mut member_counts = BTreeMap::new();
let mut active_member_counts = BTreeMap::new();
for group in state.list_user_groups().await? {
let members = state.list_user_group_members(&group.id).await?;
let member_count = members.iter().filter(|member| !member.is_deleted).count();
let active_member_count = members
.iter()
.filter(|member| !member.is_deleted && member.is_active)
.count();
let user_ids = members
.iter()
.filter(|member| !member.is_deleted)
.filter(|member| include_inactive || member.is_active)
.filter(|member| !exclude_admin || !member.role.eq_ignore_ascii_case("admin"))
.map(|member| member.user_id.clone())
.collect::<Vec<_>>();
let summaries = state
.summarize_usage_leaderboard(&UsageLeaderboardQuery {
created_from_unix_secs,
created_until_unix_secs,
group_by: UsageLeaderboardGroupBy::User,
user_id: None,
user_ids: Some(user_ids),
provider_names: super::super::usage_group_provider_names(state, &group).await?,
provider_name: filters.provider_name.clone(),
model: filters.model.clone(),
})
.await?;
let user_ids = summaries
.iter()
.map(|row| row.group_key.clone())
.collect::<Vec<_>>();
let metadata = load_user_leaderboard_metadata(state, &user_ids).await?;
let users = build_user_leaderboard_items_from_summaries(
&summaries,
&metadata,
state.has_auth_user_data_reader(),
state.has_user_data_reader(),
include_inactive,
exclude_admin,
);
let mut item = AdminStatsLeaderboardItem {
id: group.id.clone(),
name: group.name,
requests: 0,
tokens: 0,
cost: 0.0,
};
for user in users {
item.requests = item.requests.saturating_add(user.requests);
item.tokens = item.tokens.saturating_add(user.tokens);
item.cost += user.cost;
}
member_counts.insert(group.id.clone(), member_count);
active_member_counts.insert(group.id, active_member_count);
leaderboard.push(item);
}
let ungrouped = super::super::ungrouped_usage_users(state).await?;
let id = super::super::UNGROUPED_USAGE_ID.to_string();
member_counts.insert(id.clone(), ungrouped.len());
active_member_counts.insert(
id.clone(),
ungrouped.iter().filter(|user| user.is_active).count(),
);
let user_ids = ungrouped
.into_iter()
.filter(|user| include_inactive || user.is_active)
.filter(|user| !exclude_admin || !user.role.eq_ignore_ascii_case("admin"))
.map(|user| user.id)
.collect();
let rows = state
.summarize_usage_leaderboard(&UsageLeaderboardQuery {
created_from_unix_secs,
created_until_unix_secs,
group_by: UsageLeaderboardGroupBy::User,
user_id: None,
user_ids: Some(user_ids),
provider_names: None,
provider_name: filters.provider_name.clone(),
model: filters.model.clone(),
})
.await?;
leaderboard.push(AdminStatsLeaderboardItem {
id,
name: "Ungrouped".to_string(),
requests: rows.iter().map(|row| row.request_count).sum(),
tokens: rows.iter().map(|row| row.total_tokens).sum(),
cost: rows.iter().map(|row| row.total_cost_usd).sum(),
});
leaderboard.sort_by(|left, right| compare_leaderboard_items(metric, order, left, right));
return Ok(Some(build_admin_stats_user_group_leaderboard_response(
metric,
Some(&time_range),
&leaderboard,
&member_counts,
&active_member_counts,
offset,
limit,
)));
}
if request_context
.decision()
.and_then(|decision| decision.route_kind.as_deref())
@@ -253,6 +435,13 @@ pub(super) async fn maybe_build_local_admin_stats_leaderboard_response(
let include_inactive = query_param_bool(query, "include_inactive", false);
let exclude_admin = query_param_bool(query, "exclude_admin", false);
let filters = AdminStatsUsageFilter::from_query(query);
let scoped_user_ids =
match resolve_usage_user_group_scope(state, query, include_inactive, exclude_admin)
.await?
{
Ok(value) => value,
Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))),
};
let Some((created_from_unix_secs, created_until_unix_secs)) = time_range.to_unix_bounds()
else {
return Ok(Some(admin_stats_leaderboard_empty_response(
@@ -262,10 +451,13 @@ pub(super) async fn maybe_build_local_admin_stats_leaderboard_response(
};
let summaries = state
.summarize_usage_leaderboard(&UsageLeaderboardQuery {
provider_names: super::super::resolve_usage_group_provider_names(state, query)
.await?,
created_from_unix_secs,
created_until_unix_secs,
group_by: UsageLeaderboardGroupBy::User,
user_id: filters.user_id,
user_ids: scoped_user_ids,
provider_name: filters.provider_name,
model: filters.model,
})
@@ -1,3 +1,4 @@
use super::super::resolve_usage_user_group_scope;
use super::super::stats::resolve_admin_usage_time_range;
use super::analytics::admin_usage_api_key_names;
use super::analytics::admin_usage_provider_key_names;
@@ -121,21 +122,30 @@ fn apply_admin_usage_status_filter(query: &mut UsageAuditListQuery, status: Opti
"pending" | "streaming" | "completed" | "cancelled" => {
query.statuses = Some(vec![status]);
}
"has_fallback" | "has_retry" => {}
"has_fallback" | "has_retry" | "has_skipped_candidate" => {}
_ => {}
}
}
#[derive(Clone, Copy, Debug, Default)]
#[derive(Clone, Debug, Default)]
struct AdminUsageAttemptFlags {
has_fallback: bool,
has_retry: bool,
/// 是否存在"被调度跳过"的候选(调度阶段判定本次不可用,从未向上游发起请求)。
///
/// 这是与 has_fallback 正交的信号:has_fallback 表示"更靠前的候选真的失败并被换掉",
/// 而本字段表示"更靠前的候选压根没被发出去"。两者在日志列表里观感都是"换了提供商",
/// 但用户拿不到 has_fallback 小图标时容易误判为调度错误,故单独暴露。
has_skipped_candidate: bool,
/// 跳过原因(去重、保持出现顺序),用于前端 tooltip 直接说明"为什么没用它"。
skipped_candidate_reasons: Vec<String>,
}
fn admin_usage_attempt_status_filter(status: Option<&str>) -> Option<&'static str> {
match status?.trim().to_ascii_lowercase().as_str() {
"has_fallback" => Some("has_fallback"),
"has_retry" => Some("has_retry"),
"has_skipped_candidate" => Some("has_skipped_candidate"),
_ => None,
}
}
@@ -190,25 +200,52 @@ fn admin_usage_attempt_flags_from_candidates(
})
});
let has_retry = candidates.iter().any(admin_usage_candidate_was_retried);
let skipped_candidate_reasons = admin_usage_skipped_candidate_reasons(candidates);
AdminUsageAttemptFlags {
has_fallback,
has_retry,
has_skipped_candidate: !skipped_candidate_reasons.is_empty(),
skipped_candidate_reasons,
}
}
/// 收集被跳过候选的原因,去重并保持候选顺序(决定性的在前,便于阅读)。
fn admin_usage_skipped_candidate_reasons(candidates: &[StoredRequestCandidate]) -> Vec<String> {
let mut reasons = Vec::new();
for candidate in candidates
.iter()
.filter(|candidate| candidate.status == RequestCandidateStatus::Skipped)
{
let Some(reason) = candidate
.skip_reason
.as_deref()
.map(str::trim)
.filter(|reason| !reason.is_empty())
else {
continue;
};
if !reasons.iter().any(|existing| existing == reason) {
reasons.push(reason.to_string());
}
}
reasons
}
fn admin_usage_attempt_flags_for_item(
item: &StoredRequestUsageAudit,
flags_by_usage_id: &BTreeMap<String, AdminUsageAttemptFlags>,
request_candidate_reader_available: bool,
) -> AdminUsageAttemptFlags {
flags_by_usage_id.get(&item.id).copied().unwrap_or_else(|| {
flags_by_usage_id.get(&item.id).cloned().unwrap_or_else(|| {
if request_candidate_reader_available {
AdminUsageAttemptFlags::default()
} else {
AdminUsageAttemptFlags {
has_fallback: admin_usage_has_fallback(item),
has_retry: false,
has_skipped_candidate: false,
skipped_candidate_reasons: Vec::new(),
}
}
})
@@ -477,6 +514,8 @@ fn admin_usage_matches_attempt_status(
match status {
"has_fallback" => flags.has_fallback,
"has_retry" => flags.has_retry,
// 与 has_fallback 区分:这里是"更靠前的候选被调度跳过、根本没发出去"
"has_skipped_candidate" => flags.has_skipped_candidate,
_ => true,
}
}
@@ -548,6 +587,9 @@ fn build_admin_usage_records_response_with_attempt_flags(
);
record["has_fallback"] = json!(flags.has_fallback);
record["has_retry"] = json!(flags.has_retry);
// 被跳过的候选:前端据此提示"这次没用某个提供商,是因为它在调度阶段就被排除了"。
record["has_skipped_candidate"] = json!(flags.has_skipped_candidate);
record["skipped_candidate_reasons"] = json!(flags.skipped_candidate_reasons);
record
})
.collect();
@@ -799,11 +841,18 @@ pub(super) async fn maybe_build_local_admin_usage_summary_response(
&Default::default(),
)));
};
let user_ids = match resolve_usage_user_group_scope(state, query, false, false).await? {
Ok(value) => value,
Err(detail) => return Ok(Some(admin_usage_bad_request_response(detail))),
};
let summary = state
.summarize_usage_audits(&UsageAuditSummaryQuery {
provider_names: super::super::resolve_usage_group_provider_names(state, query)
.await?,
created_from_unix_secs,
created_until_unix_secs,
user_id: query_param_value(query, "user_id"),
user_ids,
provider_name: query_param_value(query, "provider"),
model: query_param_value(query, "model"),
})
@@ -1202,9 +1251,11 @@ mod tests {
use aether_data_contracts::repository::candidates::{
RequestCandidateStatus, StoredRequestCandidate,
};
use aether_data_contracts::repository::usage::StoredRequestUsageAudit;
use serde_json::json;
use super::{
admin_usage_attempt_flags_from_candidates, admin_usage_skipped_candidate_reasons,
admin_usage_terminal_candidate_state_override, build_admin_usage_keyword_search_query,
build_admin_usage_records_query, latest_admin_usage_image_progress,
AdminUsageSearchContext,
@@ -1246,6 +1297,144 @@ mod tests {
.expect("candidate should build")
}
/// 构造一条"被调度跳过"的候选(从未向上游发起请求)。
fn skipped_candidate(candidate_index: i32, reason: &str) -> StoredRequestCandidate {
let mut candidate = sample_candidate(
candidate_index,
RequestCandidateStatus::Skipped,
None,
None,
None,
);
candidate.skip_reason = Some(reason.to_string());
// 跳过候选没有开始时间,is_attempted 因此为 false
candidate.started_at_unix_ms = None;
candidate
}
#[test]
fn skipped_candidate_reasons_are_deduplicated_in_candidate_order() {
let reasons = admin_usage_skipped_candidate_reasons(&[
skipped_candidate(0, "key_rpm_exhausted"),
skipped_candidate(1, "provider_inactive"),
skipped_candidate(2, "key_rpm_exhausted"),
]);
assert_eq!(
reasons,
vec![
"key_rpm_exhausted".to_string(),
"provider_inactive".to_string()
]
);
}
#[test]
fn skipped_candidate_reasons_ignore_attempted_candidates() {
// 真正发起过请求的失败候选不属于"被跳过",避免与 has_fallback 语义混淆
let failed = sample_candidate(
0,
RequestCandidateStatus::Failed,
Some(503),
Some(1_000),
Some("upstream exploded"),
);
assert!(admin_usage_skipped_candidate_reasons(&[failed]).is_empty());
}
#[test]
fn attempt_flags_report_skipped_candidates_without_fallback() {
let candidates = vec![
skipped_candidate(0, "key_rpm_exhausted"),
sample_candidate(
1,
RequestCandidateStatus::Success,
Some(200),
Some(900),
None,
),
];
let flags = admin_usage_attempt_flags_from_candidates(&sample_usage_audit(), &candidates);
// 这正是用户遇到的场景:换了提供商,但没有任何候选失败过
assert!(flags.has_skipped_candidate);
assert!(!flags.has_fallback);
assert_eq!(
flags.skipped_candidate_reasons,
vec!["key_rpm_exhausted".to_string()]
);
}
#[test]
fn attempt_flags_keep_fallback_and_skipped_candidate_independent() {
let candidates = vec![
skipped_candidate(0, "provider_inactive"),
sample_candidate(
1,
RequestCandidateStatus::Failed,
Some(503),
Some(500),
None,
),
sample_candidate(
2,
RequestCandidateStatus::Success,
Some(200),
Some(700),
None,
),
];
let flags = admin_usage_attempt_flags_from_candidates(&sample_usage_audit(), &candidates);
assert!(flags.has_skipped_candidate);
assert!(flags.has_fallback);
}
/// 最小可用的用量审计行,仅用于驱动 flags 计算(其中候选 id 为空即可)。
fn sample_usage_audit() -> StoredRequestUsageAudit {
StoredRequestUsageAudit::new(
"usage-1".to_string(),
"req-1".to_string(),
Some("user-1".to_string()),
Some("api-key-1".to_string()),
Some("alice".to_string()),
Some("default".to_string()),
"OpenAI".to_string(),
"gpt-4.1".to_string(),
None,
None,
None,
None,
None,
Some("openai:chat".to_string()),
Some("openai".to_string()),
Some("chat".to_string()),
Some("openai:chat".to_string()),
Some("openai".to_string()),
Some("chat".to_string()),
false,
false,
10,
20,
30,
0.0,
0.0,
Some(200),
None,
None,
None,
None,
"completed".to_string(),
"settled".to_string(),
1_000,
1_001,
None,
)
.expect("usage should build")
}
#[test]
fn admin_usage_active_override_uses_current_terminal_candidate_latency() {
let candidate = sample_candidate(
@@ -74,7 +74,8 @@ fn validate_batch_access_token_import(
) -> Result<(), String> {
if !provider_type_supports_access_token_import(provider_type) {
return Err(
"Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok Provider".to_string(),
"Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok / xAI Provider"
.to_string(),
);
}
if provider_type.eq_ignore_ascii_case("claude_code") {
@@ -214,7 +214,12 @@ fn extract_admin_provider_oauth_batch_import_entry(
} else {
let sso_from_cookie = grok_cookie_session_token(provider_type, raw_token);
let token_input = sso_from_cookie.as_deref().unwrap_or(raw_token);
let (refresh_token, access_token) = import_tokens_from_raw_token(token_input);
let (refresh_token, access_token) =
if provider_type.trim().eq_ignore_ascii_case("xai") {
(None, Some(token_input.to_string()))
} else {
import_tokens_from_raw_token(token_input)
};
let (refresh_token, access_token) = normalize_provider_import_tokens(
provider_type,
refresh_token.as_deref(),
@@ -262,6 +267,7 @@ fn extract_admin_provider_oauth_batch_import_entry(
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_xai = provider_type.trim().eq_ignore_ascii_case("xai");
let is_codex_agent_identity = provider_type.trim().eq_ignore_ascii_case("codex")
&& aether_provider_transport::is_codex_agent_identity_auth_config_value(item);
if is_codex_agent_identity {
@@ -336,14 +342,6 @@ fn extract_admin_provider_oauth_batch_import_entry(
} else {
None
};
let (refresh_token, access_token) = normalize_provider_import_tokens(
provider_type,
refresh_token.as_deref(),
access_token
.as_deref()
.or(session_token.as_deref())
.or(header_bearer_token.as_deref()),
);
let windsurf_api_key = is_windsurf
.then(|| {
coerce_admin_provider_oauth_import_str(
@@ -351,6 +349,22 @@ fn extract_admin_provider_oauth_batch_import_entry(
)
})
.flatten();
let xai_api_key = is_xai
.then(|| {
coerce_admin_provider_oauth_import_str(
object.get("api_key").or_else(|| object.get("apiKey")),
)
})
.flatten();
let (refresh_token, access_token) = normalize_provider_import_tokens(
provider_type,
refresh_token.as_deref(),
access_token
.as_deref()
.or(session_token.as_deref())
.or(header_bearer_token.as_deref())
.or(xai_api_key.as_deref()),
);
let windsurf_token = is_windsurf
.then(|| {
coerce_admin_provider_oauth_import_str(
@@ -1577,4 +1591,23 @@ mod tests {
assert!(entries[1].access_token.is_none());
assert!(entries[1].raw_credentials.is_none());
}
#[test]
fn parses_xai_api_key_json_and_raw_lines_as_access_token() {
let entries = parse_admin_provider_oauth_batch_import_entries(
"xai",
r#"{"api_key":"xai-api-key","email":"a@x.ai"}
{"refresh_token":"xai-refresh"}
xai-raw-api-key"#,
);
assert_eq!(entries.len(), 3);
assert!(entries[0].refresh_token.is_none());
assert_eq!(entries[0].access_token.as_deref(), Some("xai-api-key"));
assert_eq!(entries[0].email.as_deref(), Some("[email protected]"));
assert_eq!(entries[1].refresh_token.as_deref(), Some("xai-refresh"));
assert!(entries[1].access_token.is_none());
assert!(entries[2].refresh_token.is_none());
assert_eq!(entries[2].access_token.as_deref(), Some("xai-raw-api-key"));
}
}
@@ -186,10 +186,10 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if provider_type != "kiro" && provider_type != "windsurf" {
if provider_type != "kiro" && provider_type != "windsurf" && provider_type != "xai" {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"设备授权仅支持 Kiro / Windsurf provider",
"设备授权仅支持 Kiro / Windsurf / xAI provider",
));
}
let Some(principal) = request_context
@@ -219,6 +219,19 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
)
.await;
if provider_type == "xai" {
return super::xai::handle_admin_provider_oauth_xai_device_authorize(
state,
&provider_id,
&provider,
principal,
runtime_endpoint.as_ref(),
request_proxy,
payload.proxy_node_id.as_deref(),
)
.await;
}
if provider_type == "windsurf" {
let session_id = generate_provider_oauth_nonce();
let login_option = payload
@@ -2,6 +2,7 @@ mod authorize;
mod lease;
mod poll;
mod session;
mod xai;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
@@ -479,6 +479,18 @@ pub(super) async fn handle_admin_provider_oauth_device_poll(
)
.await;
if provider_type == "xai" {
return super::xai::handle_admin_provider_oauth_xai_device_poll(
state,
&provider,
&endpoints,
request_proxy,
session_id,
session,
)
.await;
}
if provider_type == "windsurf" {
return handle_admin_provider_oauth_windsurf_browser_device_poll(
state,
@@ -0,0 +1,368 @@
use super::session::attach_admin_provider_oauth_device_poll_terminal_response;
use crate::control::GatewayAdminPrincipalContext;
use crate::handlers::admin::provider::oauth::dispatch::helpers::admin_provider_oauth_key_name_from_auth_config;
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
use crate::handlers::admin::provider::oauth::provisioning::{
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
};
use crate::handlers::admin::provider::oauth::runtime::spawn_provider_oauth_account_state_refresh_after_update;
use crate::handlers::admin::provider::oauth::state::{
current_unix_secs, generate_provider_oauth_nonce,
};
use crate::handlers::admin::request::AdminAppState;
use crate::GatewayError;
use aether_contracts::ProxySnapshot;
use aether_data::repository::provider_oauth::{
StoredAdminProviderOAuthDeviceSession, KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS,
};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
};
use aether_oauth::core::OAuthError;
use aether_oauth::provider::providers::{
XaiDevicePollOutcome, XaiProviderOAuthAdapter, XAI_CLIENT_ID, XAI_DEVICE_CODE_URL,
XAI_TOKEN_URL,
};
use aether_oauth::provider::ProviderOAuthTransportContext;
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::{json, Value};
pub(super) async fn handle_admin_provider_oauth_xai_device_authorize(
state: &AdminAppState<'_>,
provider_id: &str,
provider: &StoredProviderCatalogProvider,
principal: &GatewayAdminPrincipalContext,
runtime_endpoint: Option<&StoredProviderCatalogEndpoint>,
request_proxy: Option<ProxySnapshot>,
proxy_node_id: Option<&str>,
) -> Result<Response<Body>, GatewayError> {
let device_url = state.provider_oauth_token_url("xai_device", XAI_DEVICE_CODE_URL);
let token_url = state.provider_oauth_token_url("xai", XAI_TOKEN_URL);
let adapter =
XaiProviderOAuthAdapter::default().with_endpoint_overrides(&device_url, &token_url);
let ctx = ProviderOAuthTransportContext {
provider_id: provider_id.to_string(),
provider_type: provider.provider_type.clone(),
endpoint_id: runtime_endpoint.map(|endpoint| endpoint.id.clone()),
key_id: None,
auth_type: Some("oauth".to_string()),
decrypted_api_key: None,
decrypted_auth_config: None,
provider_config: provider.config.clone(),
endpoint_config: runtime_endpoint.and_then(|endpoint| endpoint.config.clone()),
key_config: None,
network: aether_oauth::network::OAuthNetworkContext::provider_operation(
request_proxy.clone(),
),
};
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
let authorization = match adapter.start_device_flow(&executor, &ctx).await {
Ok(authorization) => authorization,
Err(error) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
sanitize_xai_oauth_error(&error),
));
}
};
let now_unix_secs = current_unix_secs();
let session_id = generate_provider_oauth_nonce();
let session = StoredAdminProviderOAuthDeviceSession {
session_id: session_id.clone(),
provider_id: provider_id.to_string(),
initiated_by_user_id: principal.user_id.clone(),
initiated_by_session_id: principal.session_id.clone(),
initiated_by_management_token_id: principal.management_token_id.clone(),
region: String::new(),
client_id: XAI_CLIENT_ID.to_string(),
client_secret: String::new(),
device_code: authorization.device_code.clone(),
auth_type: Some("device".to_string()),
social_provider: None,
code_verifier: None,
redirect_uri: Some(token_url),
machine_id: None,
interval: authorization.interval,
expires_at_unix_secs: now_unix_secs.saturating_add(authorization.expires_in),
status: "pending".to_string(),
proxy_node_id: proxy_node_id
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned),
created_at_unix_ms: now_unix_secs,
key_id: None,
email: None,
replaced: false,
error_msg: None,
};
if let Err(response) = state
.save_provider_oauth_device_session(
&session_id,
&session,
authorization
.expires_in
.saturating_add(KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS),
)
.await
{
return Ok(response);
}
Ok(Json(json!({
"session_id": session_id,
"user_code": authorization.user_code,
"verification_uri": authorization.verification_uri,
"verification_uri_complete": authorization.verification_uri_complete,
"expires_in": authorization.expires_in,
"interval": authorization.interval,
"auth_type": "device",
}))
.into_response())
}
pub(super) async fn handle_admin_provider_oauth_xai_device_poll(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoints: &[StoredProviderCatalogEndpoint],
request_proxy: Option<ProxySnapshot>,
session_id: &str,
mut session: StoredAdminProviderOAuthDeviceSession,
) -> Result<Response<Body>, GatewayError> {
let token_url = session
.redirect_uri
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.unwrap_or_else(|| state.provider_oauth_token_url("xai", XAI_TOKEN_URL));
let adapter =
XaiProviderOAuthAdapter::default().with_endpoint_overrides(XAI_DEVICE_CODE_URL, token_url);
let ctx = ProviderOAuthTransportContext {
provider_id: provider.id.clone(),
provider_type: provider.provider_type.clone(),
endpoint_id: None,
key_id: None,
auth_type: Some("oauth".to_string()),
decrypted_api_key: None,
decrypted_auth_config: None,
provider_config: provider.config.clone(),
endpoint_config: None,
key_config: None,
network: aether_oauth::network::OAuthNetworkContext::provider_operation(
request_proxy.clone(),
),
};
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
let outcome = match adapter
.poll_device_token(&executor, &ctx, &session.device_code)
.await
{
Ok(outcome) => outcome,
Err(error) => {
return Ok(xai_device_poll_terminal_from_error(
state,
session_id,
&mut session,
&error,
)
.await);
}
};
match outcome {
XaiDevicePollOutcome::Pending => {
Ok(Json(json!({"status": "pending", "replaced": false})).into_response())
}
XaiDevicePollOutcome::SlowDown => {
Ok(Json(json!({"status": "slow_down", "replaced": false})).into_response())
}
XaiDevicePollOutcome::Authorized(result) => {
persist_xai_device_authorization(
state,
provider,
endpoints,
request_proxy,
session_id,
session,
*result,
)
.await
}
}
}
async fn persist_xai_device_authorization(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoints: &[StoredProviderCatalogEndpoint],
request_proxy: Option<ProxySnapshot>,
session_id: &str,
mut session: StoredAdminProviderOAuthDeviceSession,
result: aether_oauth::provider::ProviderOAuthTokenSet,
) -> Result<Response<Body>, GatewayError> {
let access_token = result.token_set.access_token.trim().to_string();
if access_token.is_empty() {
return Ok(Json(json!({
"status": "error",
"error": "xAI token 响应缺少 access_token",
"replaced": false,
}))
.into_response());
}
let mut auth_config = result.auth_config.as_object().cloned().unwrap_or_default();
auth_config.insert("provider_type".to_string(), json!("xai"));
auth_config.insert("auth_method".to_string(), json!("oauth"));
auth_config.insert("using_api".to_string(), json!(false));
let duplicate = match state
.find_duplicate_provider_oauth_key(&provider.id, &auth_config, None)
.await
{
Ok(duplicate) => duplicate,
Err(detail) => {
return Ok(Json(json!({
"status": "error",
"error": detail,
"replaced": false,
}))
.into_response());
}
};
let api_formats = provider_oauth_active_api_formats(endpoints);
let key_proxy = provider_oauth_key_proxy_value(session.proxy_node_id.as_deref());
let expires_at = result.token_set.expires_at_unix_secs;
let email = auth_config
.get("email")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let mut replaced = false;
let persisted_key = if let Some(existing_key) = duplicate {
replaced = true;
match state
.update_existing_provider_oauth_catalog_key(
&existing_key,
&provider.provider_type,
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
expires_at,
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
} else {
let key_name = admin_provider_oauth_key_name_from_auth_config(
&provider.provider_type,
&auth_config,
None,
);
match state
.create_provider_oauth_catalog_key(
&provider.id,
&provider.provider_type,
&key_name,
&access_token,
&auth_config,
&api_formats,
key_proxy,
expires_at,
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
};
spawn_provider_oauth_account_state_refresh_after_update(
state.cloned_app(),
provider.clone(),
persisted_key.id.clone(),
request_proxy.clone(),
);
session.status = "authorized".to_string();
session.key_id = Some(persisted_key.id.clone());
session.email = email.clone();
session.replaced = replaced;
session.error_msg = None;
let _ = state
.save_provider_oauth_device_session(session_id, &session, 60)
.await;
Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"authorized",
Json(json!({
"status": "authorized",
"key_id": persisted_key.id,
"email": email,
"replaced": replaced,
}))
.into_response(),
))
}
async fn xai_device_poll_terminal_from_error(
state: &AdminAppState<'_>,
session_id: &str,
session: &mut StoredAdminProviderOAuthDeviceSession,
error: &OAuthError,
) -> Response<Body> {
let (status, message) = match error {
OAuthError::InvalidRequest(detail) if detail.to_ascii_lowercase().contains("expired") => {
("expired", "设备码已过期".to_string())
}
OAuthError::InvalidRequest(detail) if detail.to_ascii_lowercase().contains("denied") => {
("error", "用户拒绝授权".to_string())
}
_ => ("error", sanitize_xai_oauth_error(error)),
};
session.status = status.to_string();
session.error_msg = Some(message.clone());
let _ = state
.save_provider_oauth_device_session(session_id, session, 30)
.await;
attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
status,
Json(json!({
"status": status,
"error": message,
"replaced": false,
}))
.into_response(),
)
}
fn sanitize_xai_oauth_error(error: &OAuthError) -> String {
match error {
OAuthError::InvalidRequest(_) => "xAI 设备授权失败: 请求参数无效".to_string(),
OAuthError::HttpStatus { status_code, .. } => {
format!("xAI 设备授权失败: HTTP {status_code}")
}
_ => "xAI 设备授权失败".to_string(),
}
}
@@ -715,7 +715,7 @@ async fn resolve_admin_provider_oauth_single_import_tokens(
if !provider_type_supports_access_token_import(provider_type) {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok Provider",
"Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok / xAI Provider",
));
}
@@ -867,7 +867,7 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
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(
let mut access_token_input = import_payload_string_any(
&raw_payload,
&[
"access_token",
@@ -879,6 +879,9 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
],
)
.or_else(|| provider_oauth_import_authorization_bearer_token_from_object(&raw_payload));
if provider_type == "xai" && access_token_input.is_none() {
access_token_input = import_payload_string(&raw_payload, "api_key", "apiKey");
}
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(
@@ -901,7 +904,11 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
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,
"Refresh Token、Access Token 或 sso_token 不能为空",
if provider_type == "xai" {
"Refresh Token、Access Token 或 api_key 不能为空"
} else {
"Refresh Token、Access Token 或 sso_token 不能为空"
},
));
}
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
@@ -70,6 +70,12 @@ pub(super) async fn handle_admin_provider_oauth_start_key(
"Windsurf 请使用浏览器登录或导入凭据。",
));
}
if provider_type == "xai" {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"xAI 请使用设备授权或导入凭据。",
));
}
let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
@@ -167,6 +173,12 @@ pub(super) async fn handle_admin_provider_oauth_start_provider(
"Windsurf 请使用浏览器登录或导入凭据。",
));
}
if provider_type == "xai" {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"xAI 请使用设备授权或导入凭据。",
));
}
let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
@@ -121,6 +121,9 @@ pub(super) fn normalize_provider_import_tokens(
if provider_type == "grok" {
return (None, access_token.or(refresh_token));
}
if provider_type == "xai" {
return (refresh_token, access_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);
@@ -237,7 +240,7 @@ 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(),
"claude_code" | "codex" | "chatgpt_web" | "grok"
"claude_code" | "codex" | "chatgpt_web" | "grok" | "xai"
)
}
@@ -331,6 +334,15 @@ pub(super) fn build_provider_access_token_import_auth_config(
auth_config.insert("sso_token".to_string(), json!(access_token));
auth_config.insert("auth_method".to_string(), json!("sso_token"));
}
if provider_type.trim().eq_ignore_ascii_case("xai") {
if refresh_token.is_some() {
auth_config.insert("auth_method".to_string(), json!("oauth"));
auth_config.insert("using_api".to_string(), json!(false));
} else {
auth_config.insert("auth_method".to_string(), json!("api_key"));
auth_config.insert("using_api".to_string(), json!(true));
}
}
auth_config.insert(
"access_token_import_temporary".to_string(),
@@ -532,6 +544,41 @@ mod tests {
);
}
#[test]
fn normalize_xai_import_keeps_refresh_token_separate_from_api_key() {
let (refresh_token, access_token) =
normalize_provider_import_tokens("xai", Some("xai-refresh-token"), None);
assert_eq!(refresh_token.as_deref(), Some("xai-refresh-token"));
assert!(access_token.is_none());
let (refresh_token, access_token) =
normalize_provider_import_tokens("xai", None, Some("xai-api-key"));
assert!(refresh_token.is_none());
assert_eq!(access_token.as_deref(), Some("xai-api-key"));
}
#[test]
fn builds_xai_auth_config_from_api_key_and_oauth_tokens() {
let (api_key_config, _) =
build_provider_access_token_import_auth_config("xai", "xai-api-key", None, None, None);
assert_eq!(api_key_config.get("auth_method"), Some(&json!("api_key")));
assert_eq!(api_key_config.get("using_api"), Some(&json!(true)));
let (oauth_config, _) = build_provider_access_token_import_auth_config(
"xai",
"xai-access-token",
Some("xai-refresh-token"),
None,
None,
);
assert_eq!(oauth_config.get("auth_method"), Some(&json!("oauth")));
assert_eq!(oauth_config.get("using_api"), Some(&json!(false)));
assert_eq!(
oauth_config.get("refresh_token"),
Some(&json!("xai-refresh-token"))
);
}
#[test]
fn flattens_only_claude_ai_oauth_credentials_and_converts_expiry_ms() {
let mut payload = json!({
@@ -0,0 +1,221 @@
use super::shared::{
build_provider_quota_execution_plan, build_quota_snapshot_payload,
default_provider_quota_execution_timeouts, execute_provider_quota_plan,
extract_execution_error_message, oauth_refresh_auto_removed_result,
persist_provider_quota_refresh_state, quota_key_auto_removed,
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::GatewayError;
use aether_admin::provider::quota::parse_claude_code_oauth_usage_response;
use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json;
use aether_contracts::ProxySnapshot;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_provider_pool::build_claude_code_pool_quota_request;
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
async fn execute_claude_code_quota_plan(
state: &AdminAppState<'_>,
transport: &AdminGatewayProviderTransportSnapshot,
authorization: (String, String),
proxy_override: Option<&ProxySnapshot>,
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
let proxy = match proxy_override {
Some(proxy) => Some(proxy.clone()),
None => {
state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await
}
};
let timeouts = state
.resolve_transport_execution_timeouts(transport)
.or(Some(default_provider_quota_execution_timeouts(
proxy.as_ref(),
)));
let spec = build_claude_code_pool_quota_request(&transport.key.id, authorization);
let plan = build_provider_quota_execution_plan(
transport,
spec,
proxy,
state.resolve_transport_profile(transport),
timeouts,
);
execute_provider_quota_plan(state, transport, plan, "claude_code").await
}
pub(crate) async fn refresh_claude_code_provider_quota_locally(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
proxy_override: Option<ProxySnapshot>,
) -> Result<Option<serde_json::Value>, GatewayError> {
let mut results = Vec::new();
let mut success_count = 0usize;
let mut failed_count = 0usize;
let mut auto_removed_count = 0usize;
for key in keys {
let transport = match state
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
.await?
{
Some(transport) => transport,
None => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Provider transport snapshot unavailable",
}));
continue;
}
};
let authorization = match state.resolve_local_oauth_header_auth(&transport).await? {
Some(auth) => auth,
_ => {
if quota_key_auto_removed(state, &key.id).await? {
auto_removed_count += 1;
results.push(oauth_refresh_auto_removed_result(&key));
continue;
}
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 OAuth 认证信息,请先授权/刷新 Token",
}));
continue;
}
};
let result = match execute_claude_code_quota_plan(
state,
&transport,
authorization,
proxy_override.as_ref(),
)
.await?
{
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(_) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "oauth/usage 请求执行失败",
"status_code": 502,
}));
continue;
}
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut metadata_update = None::<serde_json::Value>;
let (oauth_invalid_at_unix_secs, oauth_invalid_reason) =
quota_refresh_success_invalid_state(&key);
let mut status = "error".to_string();
let mut message = None::<String>;
if result.status_code == 200 {
if let Some(body_json) = result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
{
metadata_update = parse_claude_code_oauth_usage_response(body_json, now_unix_secs)
.map(|metadata| json!({ "claude_code": metadata }));
if metadata_update.is_some() {
status = "success".to_string();
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含额度窗口".to_string());
}
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含配额信息".to_string());
}
} else {
message = Some(match result.status_code {
401 => "oauth/usage 返回 401,Token 可能已失效,请刷新 Token".to_string(),
403 => "oauth/usage 返回 403,该账号缺少 user:profile 权限(如 Setup Token),无法查询额度"
.to_string(),
429 => "oauth/usage 被限流,请稍后重试".to_string(),
code => format!("oauth/usage 返回状态码 {code}"),
});
}
if !persist_provider_quota_refresh_state(
state,
&key.id,
metadata_update.as_ref(),
oauth_invalid_at_unix_secs,
oauth_invalid_reason,
None,
)
.await?
{
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Key 状态写入失败",
}));
continue;
}
if status == "success" {
success_count += 1;
} else {
failed_count += 1;
}
let mut payload = serde_json::Map::new();
payload.insert("key_id".to_string(), json!(key.id));
payload.insert("key_name".to_string(), json!(key.name));
payload.insert("status".to_string(), json!(status));
if let Some(message) = message {
payload.insert("message".to_string(), json!(message));
}
if let Some(metadata) = metadata_update
.as_ref()
.and_then(|value| value.get("claude_code"))
{
payload.insert(
"metadata".to_string(),
admin_provider_metadata_bucket_safe_json("claude_code", Some(metadata)),
);
}
if let Some(quota_snapshot) = build_quota_snapshot_payload(
"claude_code",
key.status_snapshot.as_ref(),
metadata_update.as_ref(),
) {
payload.insert("quota_snapshot".to_string(), quota_snapshot);
}
results.push(serde_json::Value::Object(payload));
}
Ok(Some(json!({
"success": success_count,
"failed": failed_count,
"total": results.len(),
"results": results,
"message": format!("已处理 {} 个 Key", results.len()),
"auto_removed": auto_removed_count,
})))
}
@@ -3,11 +3,13 @@ use std::pin::Pin;
use super::antigravity::refresh_antigravity_provider_quota_locally;
use super::chatgpt_web::refresh_chatgpt_web_provider_quota_locally;
use super::claude_code::refresh_claude_code_provider_quota_locally;
use super::codex::refresh_codex_provider_quota_locally;
use super::gemini_cli::refresh_gemini_cli_provider_quota_locally;
use super::grok::refresh_grok_provider_quota_locally;
use super::kiro::refresh_kiro_provider_quota_locally;
use super::windsurf::refresh_windsurf_provider_quota_locally;
use super::xai::refresh_xai_provider_quota_locally;
use crate::handlers::admin::request::AdminAppState;
use crate::GatewayError;
use aether_contracts::ProxySnapshot;
@@ -35,6 +37,10 @@ const PROVIDER_QUOTA_REFRESH_HANDLERS: &[(&str, ProviderQuotaRefreshHandler)] =
"chatgpt_web",
refresh_chatgpt_web_provider_quota_locally_boxed,
),
(
"claude_code",
refresh_claude_code_provider_quota_locally_boxed,
),
("codex", refresh_codex_provider_quota_locally_boxed),
(
"gemini_cli",
@@ -43,6 +49,7 @@ const PROVIDER_QUOTA_REFRESH_HANDLERS: &[(&str, ProviderQuotaRefreshHandler)] =
("grok", refresh_grok_provider_quota_locally_boxed),
("kiro", refresh_kiro_provider_quota_locally_boxed),
("windsurf", refresh_windsurf_provider_quota_locally_boxed),
("xai", refresh_xai_provider_quota_locally_boxed),
];
pub(crate) async fn refresh_provider_pool_quota_locally(
@@ -111,6 +118,22 @@ fn refresh_codex_provider_quota_locally_boxed<'a>(
))
}
fn refresh_claude_code_provider_quota_locally_boxed<'a>(
state: &'a AdminAppState<'a>,
provider: &'a StoredProviderCatalogProvider,
endpoint: &'a StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
proxy_override: Option<ProxySnapshot>,
) -> ProviderQuotaRefreshFuture<'a> {
Box::pin(refresh_claude_code_provider_quota_locally(
state,
provider,
endpoint,
keys,
proxy_override,
))
}
fn refresh_gemini_cli_provider_quota_locally_boxed<'a>(
state: &'a AdminAppState<'a>,
provider: &'a StoredProviderCatalogProvider,
@@ -174,3 +197,19 @@ fn refresh_windsurf_provider_quota_locally_boxed<'a>(
proxy_override,
))
}
fn refresh_xai_provider_quota_locally_boxed<'a>(
state: &'a AdminAppState<'a>,
provider: &'a StoredProviderCatalogProvider,
endpoint: &'a StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
proxy_override: Option<ProxySnapshot>,
) -> ProviderQuotaRefreshFuture<'a> {
Box::pin(refresh_xai_provider_quota_locally(
state,
provider,
endpoint,
keys,
proxy_override,
))
}
@@ -1,5 +1,6 @@
pub(crate) mod antigravity;
pub(crate) mod chatgpt_web;
pub(crate) mod claude_code;
pub(crate) mod codex;
pub(crate) mod dispatch;
pub(crate) mod gemini_cli;
@@ -7,3 +8,4 @@ pub(crate) mod grok;
pub(crate) mod kiro;
pub(crate) mod shared;
pub(crate) mod windsurf;
pub(crate) mod xai;
@@ -1713,8 +1713,10 @@ fn provider_quota_url_has_allowed_origin(provider_name: &str, value: &str) -> bo
| "daily-cloudcode-pa.sandbox.googleapis.com"
),
"gemini_cli" => host == "cloudcode-pa.googleapis.com",
"claude_code" => host == "api.anthropic.com",
"chatgpt_web" | "codex" => host == "chatgpt.com",
"grok" => host == "grok.com",
"xai" => host == "cli-chat-proxy.grok.com",
"windsurf" => host == "server.codeium.com",
"kiro" => kiro_quota_host_is_allowed(host),
_ => false,
@@ -1814,6 +1816,14 @@ mod tests {
),
("codex", "https://chatgpt.com/backend-api/wham/usage"),
("grok", "https://grok.com/rest/rate-limits"),
(
"xai",
"https://cli-chat-proxy.grok.com/v1/billing?format=credits",
),
(
"xai",
"https://cli-chat-proxy.grok.com/v1/user",
),
(
"windsurf",
"https://server.codeium.com/exa.seat_management_pb.SeatManagementService/GetUserStatus",
@@ -1847,6 +1857,11 @@ mod tests {
"https://chatgpt.com.attacker.test/backend-api/wham/usage",
),
("grok", "https://grok.com.attacker.test/rest/rate-limits"),
(
"xai",
"https://cli-chat-proxy.grok.com.attacker.test/v1/billing",
),
("xai", "https://api.x.ai/v1/billing?format=credits"),
("windsurf", "https://server.codeium.com.attacker.test/quota"),
(
"gemini_cli",
@@ -0,0 +1,313 @@
use super::shared::{
build_provider_quota_execution_plan, build_quota_snapshot_payload,
default_provider_quota_execution_timeouts, execute_provider_quota_plan,
extract_execution_error_message, oauth_refresh_auto_removed_result,
persist_provider_quota_refresh_state, quota_key_auto_removed,
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::GatewayError;
use aether_admin::provider::quota::parse_xai_billing_response;
use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json;
use aether_contracts::ProxySnapshot;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_provider_pool::{build_xai_pool_billing_request, build_xai_pool_user_request};
use aether_provider_transport::xai::{
extract_xai_user_id_from_auth_config, extract_xai_user_id_from_value, xai_auth_uses_api,
};
use serde_json::{json, Value};
use std::time::{SystemTime, UNIX_EPOCH};
async fn execute_xai_quota_plan(
state: &AdminAppState<'_>,
transport: &AdminGatewayProviderTransportSnapshot,
spec: aether_provider_pool::ProviderPoolQuotaRequestSpec,
proxy_override: Option<&ProxySnapshot>,
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
let proxy = match proxy_override {
Some(proxy) => Some(proxy.clone()),
None => {
state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await
}
};
let timeouts = state
.resolve_transport_execution_timeouts(transport)
.or(Some(default_provider_quota_execution_timeouts(
proxy.as_ref(),
)));
let plan = build_provider_quota_execution_plan(
transport,
spec,
proxy,
state.resolve_transport_profile(transport),
timeouts,
);
execute_provider_quota_plan(state, transport, plan, "xai").await
}
fn xai_authorization_from_header(authorization: &(String, String)) -> (String, String) {
authorization.clone()
}
fn enrich_xai_subscription_title(mut metadata: Value, auth_config: Option<&str>) -> Value {
if metadata
.get("subscription_title")
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty())
{
return metadata;
}
let Some(config) = auth_config
.map(str::trim)
.filter(|value| !value.is_empty())
.and_then(|value| serde_json::from_str::<Value>(value).ok())
else {
return metadata;
};
let title = ["subscription_tier", "subscriptionTier", "tier", "plan"]
.iter()
.find_map(|field| {
config
.get(*field)
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
});
if let Some(title) = title {
if let Some(object) = metadata.as_object_mut() {
object.insert("subscription_title".to_string(), json!(title));
}
}
metadata
}
pub(crate) async fn refresh_xai_provider_quota_locally(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
proxy_override: Option<ProxySnapshot>,
) -> Result<Option<serde_json::Value>, GatewayError> {
let mut results = Vec::new();
let mut success_count = 0usize;
let mut failed_count = 0usize;
let mut auto_removed_count = 0usize;
for key in keys {
let transport = match state
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
.await?
{
Some(transport) => transport,
None => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Provider transport snapshot unavailable",
}));
continue;
}
};
if xai_auth_uses_api(
transport.key.auth_type.as_str(),
transport.key.decrypted_auth_config.as_deref(),
) {
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "skipped",
"message": "xAI API Key 账号没有 Grok Build 订阅额度接口,请使用设备授权账号查询额度。",
}));
continue;
}
let authorization = match state.resolve_local_oauth_header_auth(&transport).await? {
Some(auth) => auth,
_ => {
if quota_key_auto_removed(state, &key.id).await? {
auto_removed_count += 1;
results.push(oauth_refresh_auto_removed_result(&key));
continue;
}
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 OAuth 认证信息,请先授权/刷新 Token",
}));
continue;
}
};
let fallback_user_id =
extract_xai_user_id_from_auth_config(transport.key.decrypted_auth_config.as_deref());
let user_id = match execute_xai_quota_plan(
state,
&transport,
build_xai_pool_user_request(
&transport.key.id,
xai_authorization_from_header(&authorization),
),
proxy_override.as_ref(),
)
.await?
{
ProviderQuotaExecutionOutcome::Response(result) if result.status_code == 200 => result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
.and_then(extract_xai_user_id_from_value)
.or(fallback_user_id),
_ => fallback_user_id,
};
let result = match execute_xai_quota_plan(
state,
&transport,
build_xai_pool_billing_request(
&transport.key.id,
xai_authorization_from_header(&authorization),
user_id.as_deref(),
),
proxy_override.as_ref(),
)
.await?
{
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(_) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "xAI billing 请求执行失败",
"status_code": 502,
}));
continue;
}
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut metadata_update = None::<serde_json::Value>;
let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) =
quota_refresh_success_invalid_state(&key);
let mut status = "error".to_string();
let mut message = None::<String>;
if result.status_code == 200 {
if let Some(body_json) = result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
{
metadata_update =
parse_xai_billing_response(body_json, now_unix_secs).map(|metadata| {
json!({
"xai": enrich_xai_subscription_title(
metadata,
transport.key.decrypted_auth_config.as_deref(),
)
})
});
if metadata_update.is_some() {
status = "success".to_string();
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含可用的 Grok Build 额度信息".to_string());
}
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含配额信息".to_string());
}
} else {
message = Some(
extract_execution_error_message(&result)
.unwrap_or_else(|| format!("xAI billing 返回状态码 {}", result.status_code)),
);
if result.status_code == 401 || result.status_code == 403 {
let reason = message
.clone()
.unwrap_or_else(|| "账户访问被禁止".to_string());
oauth_invalid_at_unix_secs = Some(now_unix_secs);
oauth_invalid_reason = Some(format!("账户访问被禁止: {reason}"));
status = if result.status_code == 401 {
"unauthorized".to_string()
} else {
"forbidden".to_string()
};
}
}
if !persist_provider_quota_refresh_state(
state,
&key.id,
metadata_update.as_ref(),
oauth_invalid_at_unix_secs,
oauth_invalid_reason,
None,
)
.await?
{
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Key 状态写入失败",
}));
continue;
}
if status == "success" {
success_count += 1;
} else {
failed_count += 1;
}
let mut payload = serde_json::Map::new();
payload.insert("key_id".to_string(), json!(key.id));
payload.insert("key_name".to_string(), json!(key.name));
payload.insert("status".to_string(), json!(status));
if let Some(message) = message {
payload.insert("message".to_string(), json!(message));
}
if let Some(metadata) = metadata_update.as_ref().and_then(|value| value.get("xai")) {
payload.insert(
"metadata".to_string(),
admin_provider_metadata_bucket_safe_json("xai", Some(metadata)),
);
}
if let Some(quota_snapshot) = build_quota_snapshot_payload(
"xai",
key.status_snapshot.as_ref(),
metadata_update.as_ref(),
) {
payload.insert("quota_snapshot".to_string(), quota_snapshot);
}
results.push(serde_json::Value::Object(payload));
}
Ok(Some(json!({
"success": success_count,
"failed": failed_count,
"total": results.len(),
"results": results,
"message": format!("已处理 {} 个 Key", results.len()),
"auto_removed": auto_removed_count,
})))
}
@@ -411,6 +411,7 @@ pub(crate) fn admin_provider_pool_config_from_config_value(
unschedulable_rules: Vec::new(),
lru_enabled: false,
skip_exhausted_accounts: false,
reserve_minimum_quota: false,
sticky_session_ttl_seconds: 3600,
latency_window_seconds: 3600,
latency_sample_limit: 50,
@@ -446,6 +447,10 @@ pub(crate) fn admin_provider_pool_config_from_config_value(
.get("skip_exhausted_accounts")
.and_then(Value::as_bool)
.unwrap_or(false),
reserve_minimum_quota: pool_advanced
.get("reserve_minimum_quota")
.and_then(Value::as_bool)
.unwrap_or(false),
sticky_session_ttl_seconds: pool_advanced
.get("sticky_session_ttl_seconds")
.and_then(json_u64)
@@ -574,6 +579,22 @@ mod tests {
let config = admin_provider_pool_config(&provider).expect("pool config should exist");
assert!(!config.skip_exhausted_accounts);
assert!(!config.reserve_minimum_quota);
}
#[test]
fn parses_reserve_minimum_quota_independently_of_skip_exhausted_accounts() {
for enabled in [false, true] {
let provider = sample_provider(json!({
"pool_advanced": {
"reserve_minimum_quota": enabled,
"skip_exhausted_accounts": false
}
}));
let config = admin_provider_pool_config(&provider).expect("pool config should exist");
assert_eq!(config.reserve_minimum_quota, enabled);
assert!(!config.skip_exhausted_accounts);
}
}
#[test]
@@ -651,6 +651,7 @@ mod tests {
unschedulable_rules: Vec::new(),
lru_enabled: true,
skip_exhausted_accounts: false,
reserve_minimum_quota: false,
sticky_session_ttl_seconds: 120,
latency_window_seconds: 600,
latency_sample_limit: 10,
@@ -932,6 +932,13 @@ fn admin_pool_build_account_quota(
return Some(account_quota);
}
}
"xai" => {
if let Some(account_quota) =
admin_pool_build_kiro_account_quota_from_snapshot(quota_snapshot)
{
return Some(account_quota);
}
}
"chatgpt_web" => {
if let Some(account_quota) =
admin_pool_build_chatgpt_web_account_quota_from_snapshot(quota_snapshot)
@@ -1110,10 +1117,16 @@ pub(super) fn build_admin_pool_key_payload(
let health_score = admin_pool_health_score(key);
let circuit_breaker_open = false;
let auth_semantics = provider_key_auth_semantics(key, provider_type);
let account_quota_exhausted = pool_config
.as_ref()
.is_some_and(|config| config.skip_exhausted_accounts)
&& admin_provider_pool_pure::admin_pool_key_account_quota_exhausted(key, provider_type);
let account_quota_exhausted = pool_config.as_ref().is_some_and(|config| {
(config.skip_exhausted_accounts
&& admin_provider_pool_pure::admin_pool_key_account_quota_exhausted(key, provider_type))
|| (config.reserve_minimum_quota
&& admin_provider_pool_pure::admin_pool_key_minimum_quota_reached(
key,
provider_type,
None,
))
});
let auth_config = state.parse_catalog_auth_config_json(key);
let oauth_expires_at =
admin_pool_derive_oauth_expires_at(provider_type, key, auth_config.as_ref());
@@ -1591,4 +1604,29 @@ mod tests {
Some("Auto剩余 40.0% (60/150) | Heavy剩余 0.0% (0/20)".to_string())
);
}
#[test]
fn xai_account_quota_is_rendered_as_remaining_percent() {
let quota_snapshot = json!({
"provider_type": "xai",
"code": "ok",
"exhausted": false,
"plan_type": "SuperGrok",
"windows": [
{
"code": "usage",
"label": "周额度",
"scope": "account",
"used_ratio": 0.46,
"remaining_ratio": 0.54
}
]
});
let quota_snapshot = quota_snapshot.as_object().unwrap();
assert_eq!(
admin_pool_build_account_quota("xai", Some(quota_snapshot)),
Some("剩余 54.0%".to_string())
);
}
}
@@ -488,9 +488,16 @@ pub(super) fn admin_pool_key_visible_status_filter(
) {
return status;
}
if pool_config.is_some_and(|config| config.skip_exhausted_accounts)
&& admin_provider_pool_pure::admin_pool_key_account_quota_exhausted(key, provider_type)
{
if pool_config.is_some_and(|config| {
(config.skip_exhausted_accounts
&& admin_provider_pool_pure::admin_pool_key_account_quota_exhausted(key, provider_type))
|| (config.reserve_minimum_quota
&& admin_provider_pool_pure::admin_pool_key_minimum_quota_reached(
key,
provider_type,
None,
))
}) {
return "quota_exhausted";
}
if !key.is_active {
@@ -594,11 +594,12 @@ async fn provider_query_fetch_models_for_key(
});
}
let dynamic_client_version = crate::ai_serving::api::codex_client_version();
let client_version = is_codex.then(|| {
codex_catalog
.as_ref()
.map(|catalog| catalog.client_version.as_str())
.unwrap_or(crate::ai_serving::CODEX_CLIENT_VERSION)
.unwrap_or(dynamic_client_version.as_str())
});
let outcome =
match fetch_models_from_transports_for_management(state.app(), &transports, client_version)
@@ -1340,6 +1340,7 @@ async fn provider_query_apply_pool_scheduler_to_test_candidates(
provider.id.clone(),
provider_query_ai_pool_runtime_state(&runtime),
);
let reserve_minimum_quota = pool_config.reserve_minimum_quota;
let pool_config =
provider_query_ai_pool_scheduling_config(pool_config, provider.provider_type.as_str());
let inputs = keys
@@ -1351,6 +1352,14 @@ async fn provider_query_apply_pool_scheduler_to_test_candidates(
effective_model: effective_model.to_string(),
scheduler_skip_reason: None,
};
let mut key_context =
provider_query_pool_catalog_key_context(state, &key, &provider.provider_type);
key_context.quota_exhausted |= reserve_minimum_quota
&& admin_provider_pool_pure::admin_pool_key_minimum_quota_reached(
&key,
&provider.provider_type,
Some(effective_model),
);
AiPoolCandidateInput {
facts: AiPoolCandidateFacts {
provider_id: provider.id.clone(),
@@ -1362,11 +1371,7 @@ async fn provider_query_apply_pool_scheduler_to_test_candidates(
key_internal_priority: key.internal_priority,
},
pool_config: Some(pool_config.clone()),
key_context: provider_query_pool_catalog_key_context(
state,
&key,
&provider.provider_type,
),
key_context,
candidate,
}
})
@@ -3511,6 +3516,11 @@ async fn provider_query_execute_standard_test_candidate(
codex_model_capabilities.as_ref(),
);
}
crate::provider_transport::insert_cli_identity_headers_if_needed(
&transport,
provider_api_format,
&mut request_headers,
);
if !uses_vertex_query_auth {
if let (Some(auth_header), Some(auth_value)) =
(auth_header.as_deref(), auth_value.as_deref())
@@ -65,6 +65,7 @@ pub(crate) struct AdminProviderPoolConfig {
pub(crate) unschedulable_rules: Vec<AdminProviderPoolUnschedulableRule>,
pub(crate) lru_enabled: bool,
pub(crate) skip_exhausted_accounts: bool,
pub(crate) reserve_minimum_quota: bool,
pub(crate) sticky_session_ttl_seconds: u64,
pub(crate) latency_window_seconds: u64,
pub(crate) latency_sample_limit: u64,
@@ -4,9 +4,9 @@ pub(crate) fn normalize_provider_type_input(value: &str) -> Result<String, Strin
let normalized = value.trim().to_ascii_lowercase();
match normalized.as_str() {
"custom" | "claude_code" | "kiro" | "codex" | "chatgpt_web" | "gemini_cli"
| "antigravity" | "vertex_ai" | "grok" | "windsurf" => Ok(normalized),
| "antigravity" | "vertex_ai" | "grok" | "windsurf" | "xai" => Ok(normalized),
_ => Err(
"provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai / grok / windsurf"
"provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai / grok / windsurf / xai"
.to_string(),
),
}
@@ -405,6 +405,14 @@ mod tests {
);
}
#[test]
fn normalize_provider_type_supports_xai() {
assert_eq!(
normalize_provider_type_input(" xAI ").expect("type should normalize"),
"xai"
);
}
#[test]
fn normalize_api_format_list_dedupes_canonical_formats() {
assert_eq!(
@@ -388,10 +388,11 @@ impl<'a> AdminAppState<'a> {
balance_type: &str,
operator_id: Option<&str>,
description: Option<&str>,
clamp_deduction_to_available_balance: bool,
) -> Result<
Option<(
aether_data::repository::wallet::StoredWalletSnapshot,
crate::AdminWalletTransactionRecord,
Option<crate::AdminWalletTransactionRecord>,
)>,
GatewayError,
> {
@@ -402,6 +403,7 @@ impl<'a> AdminAppState<'a> {
balance_type,
operator_id,
description,
clamp_deduction_to_available_balance,
)
.await
}
@@ -17,6 +17,64 @@ impl<'a> AdminAppState<'a> {
pub(crate) fn cloned_app(&self) -> AppState {
self.app.clone()
}
pub(crate) async fn get_admin_user_wallet_balance_batch(
&self,
admin_user_id: &str,
idempotency_key: &str,
request_fingerprint: &str,
) -> Result<
Option<aether_data::repository::wallet::PrepareAdminUserWalletBalanceBatchOutcome>,
GatewayError,
> {
self.app
.get_admin_user_wallet_balance_batch(
admin_user_id,
idempotency_key,
request_fingerprint,
)
.await
}
pub(crate) async fn prepare_admin_user_wallet_balance_batch(
&self,
input: aether_data::repository::wallet::PrepareAdminUserWalletBalanceBatchInput,
) -> Result<
aether_data::repository::wallet::PrepareAdminUserWalletBalanceBatchOutcome,
GatewayError,
> {
self.app
.prepare_admin_user_wallet_balance_batch(input)
.await
}
pub(crate) async fn record_admin_user_wallet_balance_batch_failure(
&self,
admin_user_id: &str,
idempotency_key: &str,
user_id: &str,
reason: &str,
) -> Result<aether_data::repository::wallet::AdminUserWalletBalanceBatchUserOutcome, GatewayError>
{
self.app
.record_admin_user_wallet_balance_batch_failure(
admin_user_id,
idempotency_key,
user_id,
reason,
)
.await
}
pub(crate) async fn adjust_admin_user_wallet_balance_batch_user(
&self,
input: aether_data::repository::wallet::AdjustWalletBalanceInBatchInput,
) -> Result<aether_data::repository::wallet::AdminUserWalletBalanceBatchUserOutcome, GatewayError>
{
self.app
.adjust_admin_user_wallet_balance_batch_user(input)
.await
}
}
impl<'a> AsRef<AppState> for AdminAppState<'a> {
@@ -126,6 +126,30 @@ impl<'a> AdminAppState<'a> {
self.app.list_user_group_members(group_id).await
}
pub(crate) async fn resolve_usage_user_group_member_ids(
&self,
group_id: &str,
include_inactive: bool,
exclude_admin: bool,
) -> Result<Option<Vec<String>>, GatewayError> {
if self.find_user_group_by_id(group_id).await?.is_none() {
return Ok(None);
}
let mut user_ids = self
.list_user_group_members(group_id)
.await?
.into_iter()
.filter(|member| !member.is_deleted)
.filter(|member| include_inactive || member.is_active)
.filter(|member| !exclude_admin || !member.role.eq_ignore_ascii_case("admin"))
.map(|member| member.user_id)
.collect::<Vec<_>>();
user_ids.sort();
user_ids.dedup();
Ok(Some(user_ids))
}
pub(crate) async fn replace_user_group_members(
&self,
group_id: &str,
@@ -64,6 +64,7 @@ pub(crate) async fn build_admin_list_user_api_keys_response(
"rate_limit": record.rate_limit,
"concurrent_limit": record.concurrent_limit,
"feature_settings": record.feature_settings,
"ip_rules": record.ip_rules,
"expires_at": format_optional_unix_secs_iso8601(record.expires_at_unix_secs),
"last_used_at": format_optional_unix_secs_iso8601(record.last_used_at_unix_secs),
"created_at": format_optional_unix_secs_iso8601(record.created_at_unix_secs),
@@ -1,11 +1,16 @@
use super::{
build_admin_users_bad_request_response, build_admin_users_permission_denied_response,
build_admin_users_read_only_response, disabled_user_policy_detail, disabled_user_policy_field,
build_admin_users_read_only_response, build_admin_users_wallet_permission_denied_response,
disabled_user_policy_detail, disabled_user_policy_field,
management_token_may_adjust_admin_wallet_balance,
management_token_may_administer_user_accounts, normalize_admin_user_role,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::attach_admin_audit_response;
use crate::GatewayError;
use aether_data::repository::wallet::{
AdminUserWalletBalanceBatchUserOutcome, PrepareAdminUserWalletBalanceBatchOutcome,
};
use axum::{
body::{Body, Bytes},
http,
@@ -13,9 +18,10 @@ use axum::{
Json,
};
use serde_json::{json, Value};
use sha2::{Digest as _, Sha256};
use std::collections::{BTreeMap, BTreeSet};
#[derive(Debug, Clone, Default, serde::Deserialize)]
#[derive(Debug, Clone, Default, serde::Deserialize, serde::Serialize)]
struct AdminUserSelectionFilters {
#[serde(default)]
search: Option<String>,
@@ -27,7 +33,7 @@ struct AdminUserSelectionFilters {
group_id: Option<String>,
}
#[derive(Debug, Clone, Default)]
#[derive(Debug, Clone, Default, serde::Serialize)]
struct AdminUserSelectionRequest {
user_ids: Vec<String>,
group_ids: Vec<String>,
@@ -40,6 +46,7 @@ struct AdminUserBatchActionRequest {
selection: AdminUserSelectionRequest,
action: String,
payload: Option<Value>,
idempotency_key: Option<String>,
}
#[derive(Debug, serde::Deserialize)]
@@ -48,6 +55,8 @@ struct RawAdminUserBatchActionRequest {
action: String,
#[serde(default)]
payload: Option<Value>,
#[serde(default)]
idempotency_key: Option<String>,
}
#[derive(Debug, Clone, Default)]
@@ -68,7 +77,7 @@ struct AdminUserSelectionItem {
matched_by: Vec<String>,
}
#[derive(Debug, Clone, serde::Serialize)]
#[derive(Debug, Clone, serde::Deserialize, serde::Serialize)]
struct AdminUserSelectionWarning {
#[serde(rename = "type")]
warning_type: String,
@@ -88,6 +97,7 @@ struct AdminUserBatchMutation {
role: Option<String>,
is_active: Option<bool>,
unlimited: Option<bool>,
wallet_balance_adjustment: Option<AdminUserWalletBalanceAdjustment>,
modified_fields: Vec<&'static str>,
}
@@ -97,6 +107,28 @@ impl AdminUserBatchMutation {
}
}
#[derive(Debug, Clone, Copy)]
struct AdminUserWalletBalanceAdjustment {
operation: AdminUserWalletBalanceOperation,
amount: f64,
}
#[derive(Debug, Clone, Copy)]
enum AdminUserWalletBalanceOperation {
Add,
Deduct,
}
enum AdminBatchWalletBalanceAdjustmentError {
WalletLookup,
BalanceAdjustment,
}
enum AdminBatchWalletLimitModeError {
WalletLookup,
Mutation,
}
pub(in super::super) async fn build_admin_resolve_user_selection_response(
state: &AdminAppState<'_>,
_request_context: &AdminRequestContext<'_>,
@@ -128,10 +160,29 @@ pub(in super::super) async fn build_admin_user_batch_action_response(
Ok(value) => value,
Err(detail) => return Ok(build_admin_user_batch_bad_request_response(detail)),
};
let mutation = match parse_batch_mutation(&request.action, request.payload) {
let mutation = match parse_batch_mutation(&request.action, request.payload.clone()) {
Ok(value) => value,
Err(detail) => return Ok(build_admin_user_batch_bad_request_response(detail)),
};
if mutation.wallet_balance_adjustment.is_some() {
if !management_token_may_adjust_admin_wallet_balance(request_context) {
return Ok(build_admin_users_wallet_permission_denied_response(
request_context,
));
}
if !state.has_auth_wallet_write_capability() {
return Ok(build_admin_users_read_only_response(
"当前为只读模式,无法批量调整用户钱包余额",
));
}
return build_admin_user_wallet_balance_batch_response(
state,
request_context,
request,
mutation,
)
.await;
}
let resolved = match resolve_admin_user_selection(state, request.selection).await {
Ok(value) => value,
Err(detail) => return Ok(build_admin_user_batch_bad_request_response(detail)),
@@ -177,9 +228,29 @@ pub(in super::super) async fn build_admin_user_batch_action_response(
.iter()
.map(|user_id| json!({ "user_id": user_id, "reason": "用户不存在或已删除" }))
.collect::<Vec<_>>();
let mut completed_user_ids = Vec::new();
let mut uncertain_user_ids = Vec::new();
let mut unprocessed_user_ids = Vec::new();
let mut interrupted = false;
for item in &resolved.items {
if state.find_user_auth_by_id(&item.user_id).await?.is_none() {
for (item_index, item) in resolved.items.iter().enumerate() {
let user = match state.find_user_auth_by_id(&item.user_id).await {
Ok(user) => user,
Err(_) => {
record_batch_action_interruption(
&resolved.items,
item_index,
false,
"读取用户状态失败,批次已中止,该用户未执行",
&mut failures,
&mut uncertain_user_ids,
&mut unprocessed_user_ids,
);
interrupted = true;
break;
}
};
if user.is_none() {
failures.push(json!({
"user_id": item.user_id,
"reason": "用户不存在或已删除",
@@ -202,17 +273,92 @@ pub(in super::super) async fn build_admin_user_batch_action_response(
}
if let Some(unlimited) = mutation.unlimited {
if !apply_batch_user_wallet_limit_mode(state, &item.user_id, unlimited).await? {
failures.push(json!({
"user_id": item.user_id,
"reason": "用户钱包不可用",
}));
continue;
match apply_batch_user_wallet_limit_mode(state, &item.user_id, unlimited).await {
Ok(true) => {}
Ok(false) => {
failures.push(json!({
"user_id": item.user_id,
"reason": "用户钱包不可用",
}));
continue;
}
Err(AdminBatchWalletLimitModeError::WalletLookup) => {
record_batch_action_interruption(
&resolved.items,
item_index,
false,
"读取用户钱包失败,批次已中止,该用户未执行",
&mut failures,
&mut uncertain_user_ids,
&mut unprocessed_user_ids,
);
interrupted = true;
break;
}
Err(AdminBatchWalletLimitModeError::Mutation) => {
record_batch_action_interruption(
&resolved.items,
item_index,
true,
"用户钱包更新结果未确认,批次已中止,请核对钱包后再重试",
&mut failures,
&mut uncertain_user_ids,
&mut unprocessed_user_ids,
);
interrupted = true;
break;
}
}
}
if mutation.has_auth_user_fields()
&& state
if let Some(adjustment) = mutation.wallet_balance_adjustment {
match apply_batch_user_wallet_balance_adjustment(
state,
&item.user_id,
adjustment,
current_admin_user_id,
)
.await
{
Ok(true) => {}
Ok(false) => {
failures.push(json!({
"user_id": item.user_id,
"reason": "用户钱包不可用",
}));
continue;
}
Err(AdminBatchWalletBalanceAdjustmentError::WalletLookup) => {
record_batch_action_interruption(
&resolved.items,
item_index,
false,
"读取用户钱包失败,批次已中止,该用户未执行",
&mut failures,
&mut uncertain_user_ids,
&mut unprocessed_user_ids,
);
interrupted = true;
break;
}
Err(AdminBatchWalletBalanceAdjustmentError::BalanceAdjustment) => {
record_batch_action_interruption(
&resolved.items,
item_index,
true,
"余额调整结果未确认,批次已中止,请核对钱包后再重试",
&mut failures,
&mut uncertain_user_ids,
&mut unprocessed_user_ids,
);
interrupted = true;
break;
}
}
}
if mutation.has_auth_user_fields() {
let updated_user = match state
.update_local_auth_user_admin_fields(
&item.user_id,
mutation.role.clone(),
@@ -226,22 +372,39 @@ pub(in super::super) async fn build_admin_user_batch_action_response(
None,
mutation.is_active,
)
.await?
.is_none()
{
failures.push(json!({
"user_id": item.user_id,
"reason": "用户不存在或已删除",
}));
continue;
.await
{
Ok(user) => user,
Err(_) => {
record_batch_action_interruption(
&resolved.items,
item_index,
true,
"用户更新结果未确认,批次已中止,请核对后再重试",
&mut failures,
&mut uncertain_user_ids,
&mut unprocessed_user_ids,
);
interrupted = true;
break;
}
};
if updated_user.is_none() {
failures.push(json!({
"user_id": item.user_id,
"reason": "用户不存在或已删除",
}));
continue;
}
}
success += 1;
completed_user_ids.push(item.user_id.clone());
}
let failed = failures.len();
let total = success + failed;
let response = Json(json!({
let mut response_payload = json!({
"total": total,
"success": success,
"failed": failed,
@@ -249,8 +412,14 @@ pub(in super::super) async fn build_admin_user_batch_action_response(
"warnings": resolved.warnings,
"action": request.action.trim().to_ascii_lowercase(),
"modified_fields": mutation.modified_fields,
}))
.into_response();
"interrupted": interrupted,
});
if interrupted {
response_payload["completed_user_ids"] = json!(completed_user_ids);
response_payload["uncertain_user_ids"] = json!(uncertain_user_ids);
response_payload["unprocessed_user_ids"] = json!(unprocessed_user_ids);
}
let response = Json(response_payload).into_response();
Ok(attach_admin_audit_response(
response,
@@ -261,6 +430,408 @@ pub(in super::super) async fn build_admin_user_batch_action_response(
))
}
fn record_batch_action_interruption(
items: &[AdminUserSelectionItem],
item_index: usize,
current_result_uncertain: bool,
reason: &str,
failures: &mut Vec<Value>,
uncertain_user_ids: &mut Vec<String>,
unprocessed_user_ids: &mut Vec<String>,
) {
let current_item = &items[item_index];
failures.push(json!({
"user_id": current_item.user_id,
"reason": reason,
}));
if current_result_uncertain {
uncertain_user_ids.push(current_item.user_id.clone());
} else {
unprocessed_user_ids.push(current_item.user_id.clone());
}
for item in items.iter().skip(item_index + 1) {
failures.push(json!({
"user_id": item.user_id,
"reason": "因前序错误未执行",
}));
unprocessed_user_ids.push(item.user_id.clone());
}
}
async fn build_admin_user_wallet_balance_batch_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request: AdminUserBatchActionRequest,
mutation: AdminUserBatchMutation,
) -> Result<Response<Body>, GatewayError> {
let Some(idempotency_key) = request
.idempotency_key
.as_deref()
.map(str::trim)
.filter(|value| {
!value.is_empty()
&& value.len() <= 128
&& value.bytes().all(|byte| (0x21..=0x7e).contains(&byte))
})
else {
return Ok(build_admin_user_batch_bad_request_response(
"余额批量操作必须提供有效的 idempotency_key".to_string(),
));
};
let Some(admin_user_id) = request_context
.decision()
.and_then(|decision| decision.admin_principal.as_ref())
.map(|principal| principal.user_id.clone())
else {
return Ok(build_admin_users_permission_denied_response(
request_context,
));
};
let action = request.action.trim().to_ascii_lowercase();
let fingerprint_payload = json!({
"selection": &request.selection,
"action": &action,
"payload": &request.payload,
});
let encoded = serde_json::to_vec(&fingerprint_payload)
.map_err(|error| GatewayError::Internal(error.to_string()))?;
let request_fingerprint = format!("{:x}", Sha256::digest(encoded));
let existing = state
.get_admin_user_wallet_balance_batch(&admin_user_id, idempotency_key, &request_fingerprint)
.await?;
let batch = match existing {
Some(PrepareAdminUserWalletBalanceBatchOutcome::Conflict) => {
return Ok(build_admin_user_batch_idempotency_conflict_response());
}
Some(PrepareAdminUserWalletBalanceBatchOutcome::Ready(batch)) => batch,
None => {
let resolved =
match resolve_admin_user_selection(state, request.selection.clone()).await {
Ok(value) => value,
Err(detail) => return Ok(build_admin_user_batch_bad_request_response(detail)),
};
let warnings = serde_json::to_value(&resolved.warnings)
.ok()
.and_then(|value| value.as_array().cloned())
.unwrap_or_default();
let prepared = state
.prepare_admin_user_wallet_balance_batch(
aether_data::repository::wallet::PrepareAdminUserWalletBalanceBatchInput {
admin_user_id: admin_user_id.clone(),
idempotency_key: idempotency_key.to_string(),
request_fingerprint: request_fingerprint.clone(),
target_user_ids: resolved
.items
.iter()
.map(|item| item.user_id.clone())
.collect(),
missing_user_ids: resolved.missing_user_ids,
warnings,
},
)
.await?;
match prepared {
PrepareAdminUserWalletBalanceBatchOutcome::Conflict => {
return Ok(build_admin_user_batch_idempotency_conflict_response());
}
PrepareAdminUserWalletBalanceBatchOutcome::Ready(batch) => batch,
}
}
};
let warnings: Vec<AdminUserSelectionWarning> =
serde_json::from_value(Value::Array(batch.warnings.clone()))
.map_err(|error| GatewayError::Internal(error.to_string()))?;
let resolved = ResolvedAdminUserSelection {
items: batch
.target_user_ids
.iter()
.map(|user_id| AdminUserSelectionItem {
user_id: user_id.clone(),
username: String::new(),
email: None,
role: "user".to_string(),
is_active: true,
matched_by: Vec::new(),
})
.collect(),
missing_user_ids: batch.missing_user_ids.clone(),
warnings,
};
let adjustment = mutation
.wallet_balance_adjustment
.expect("wallet balance action should have an adjustment");
let signed_amount = match adjustment.operation {
AdminUserWalletBalanceOperation::Add => adjustment.amount,
AdminUserWalletBalanceOperation::Deduct => -adjustment.amount,
};
let mut outcomes = batch.user_outcomes;
let mut completed_user_ids = outcomes
.iter()
.filter_map(|(user_id, outcome)| {
matches!(outcome, AdminUserWalletBalanceBatchUserOutcome::Succeeded)
.then_some(user_id.clone())
})
.collect::<Vec<_>>();
let mut success = completed_user_ids.len();
let mut failures = resolved
.missing_user_ids
.iter()
.map(|user_id| json!({ "user_id": user_id, "reason": "用户不存在或已删除" }))
.collect::<Vec<_>>();
for (user_id, outcome) in &outcomes {
if let AdminUserWalletBalanceBatchUserOutcome::Failed(reason) = outcome {
failures.push(json!({ "user_id": user_id, "reason": reason }));
}
}
let mut uncertain_user_ids = Vec::new();
let mut unprocessed_user_ids = Vec::new();
let mut interrupted = false;
for (item_index, item) in resolved.items.iter().enumerate() {
if outcomes.contains_key(&item.user_id) {
continue;
}
match state.find_user_auth_by_id(&item.user_id).await {
Err(_) => {
failures.push(json!({
"user_id": item.user_id,
"reason": "读取用户状态失败,批次已中止,该用户未执行",
}));
unprocessed_user_ids.push(item.user_id.clone());
interrupted = true;
}
Ok(None) => {
let outcome = match state
.record_admin_user_wallet_balance_batch_failure(
&admin_user_id,
idempotency_key,
&item.user_id,
"用户不存在或已删除",
)
.await
{
Ok(outcome) => outcome,
Err(_) => {
failures.push(json!({
"user_id": item.user_id,
"reason": "记录用户状态失败,批次已中止,该用户未执行",
}));
unprocessed_user_ids.push(item.user_id.clone());
interrupted = true;
append_wallet_batch_unprocessed_suffix(
&resolved.items,
item_index + 1,
&outcomes,
&mut failures,
&mut unprocessed_user_ids,
);
break;
}
};
outcomes.insert(item.user_id.clone(), outcome.clone());
match outcome {
AdminUserWalletBalanceBatchUserOutcome::Succeeded => {
completed_user_ids.push(item.user_id.clone());
success += 1;
}
AdminUserWalletBalanceBatchUserOutcome::Failed(reason) => {
failures.push(json!({ "user_id": item.user_id, "reason": reason }));
}
}
}
Ok(Some(_)) => {
let wallet = match state
.find_wallet(aether_data::repository::wallet::WalletLookupKey::UserId(
&item.user_id,
))
.await
{
Ok(Some(wallet)) => wallet,
Ok(None) => {
let outcome = match state
.record_admin_user_wallet_balance_batch_failure(
&admin_user_id,
idempotency_key,
&item.user_id,
"用户钱包不可用",
)
.await
{
Ok(outcome) => outcome,
Err(_) => {
failures.push(json!({
"user_id": item.user_id,
"reason": "记录用户钱包状态失败,批次已中止,该用户未执行",
}));
unprocessed_user_ids.push(item.user_id.clone());
interrupted = true;
append_wallet_batch_unprocessed_suffix(
&resolved.items,
item_index + 1,
&outcomes,
&mut failures,
&mut unprocessed_user_ids,
);
break;
}
};
outcomes.insert(item.user_id.clone(), outcome.clone());
match outcome {
AdminUserWalletBalanceBatchUserOutcome::Succeeded => {
completed_user_ids.push(item.user_id.clone());
success += 1;
}
AdminUserWalletBalanceBatchUserOutcome::Failed(reason) => {
failures.push(json!({ "user_id": item.user_id, "reason": reason }));
}
}
continue;
}
Err(_) => {
failures.push(json!({
"user_id": item.user_id,
"reason": "读取用户钱包失败,批次已中止,该用户未执行",
}));
unprocessed_user_ids.push(item.user_id.clone());
interrupted = true;
append_wallet_batch_unprocessed_suffix(
&resolved.items,
item_index + 1,
&outcomes,
&mut failures,
&mut unprocessed_user_ids,
);
break;
}
};
let result = state
.adjust_admin_user_wallet_balance_batch_user(
aether_data::repository::wallet::AdjustWalletBalanceInBatchInput {
admin_user_id: admin_user_id.clone(),
idempotency_key: idempotency_key.to_string(),
user_id: item.user_id.clone(),
adjustment: aether_data::repository::wallet::AdjustWalletBalanceInput {
wallet_id: wallet.id,
amount_usd: signed_amount,
balance_type: "recharge".to_string(),
operator_id: Some(admin_user_id.clone()),
description: Some("管理员批量调整用户余额".to_string()),
clamp_deduction_to_available_balance: true,
batch_context: None,
},
},
)
.await;
match result {
Ok(AdminUserWalletBalanceBatchUserOutcome::Succeeded) => {
outcomes.insert(
item.user_id.clone(),
AdminUserWalletBalanceBatchUserOutcome::Succeeded,
);
completed_user_ids.push(item.user_id.clone());
success += 1;
}
Ok(AdminUserWalletBalanceBatchUserOutcome::Failed(reason)) => {
outcomes.insert(
item.user_id.clone(),
AdminUserWalletBalanceBatchUserOutcome::Failed(reason.clone()),
);
failures.push(json!({ "user_id": item.user_id, "reason": reason }));
}
Err(_) => {
failures.push(json!({
"user_id": item.user_id,
"reason": "余额调整结果未确认,批次已中止,请使用同一批次重试以核对结果",
}));
uncertain_user_ids.push(item.user_id.clone());
interrupted = true;
append_wallet_batch_unprocessed_suffix(
&resolved.items,
item_index + 1,
&outcomes,
&mut failures,
&mut unprocessed_user_ids,
);
break;
}
}
}
}
if interrupted {
for pending in resolved.items.iter().skip(item_index + 1) {
if outcomes.contains_key(&pending.user_id) {
continue;
}
failures.push(json!({
"user_id": pending.user_id,
"reason": "因前序错误未执行",
}));
unprocessed_user_ids.push(pending.user_id.clone());
}
break;
}
}
let failed = failures.len();
let total = success + failed;
let mut response_payload = json!({
"total": total,
"success": success,
"failed": failed,
"failures": failures,
"warnings": resolved.warnings,
"action": action,
"modified_fields": mutation.modified_fields,
"interrupted": interrupted,
});
if interrupted {
response_payload["completed_user_ids"] = json!(completed_user_ids);
response_payload["uncertain_user_ids"] = json!(uncertain_user_ids);
response_payload["unprocessed_user_ids"] = json!(unprocessed_user_ids);
}
let response = Json(response_payload).into_response();
Ok(attach_admin_audit_response(
response,
"admin_users_batch_action_executed",
"batch_update_users",
"user_batch",
"users",
))
}
fn append_wallet_batch_unprocessed_suffix(
items: &[AdminUserSelectionItem],
start_index: usize,
outcomes: &BTreeMap<String, AdminUserWalletBalanceBatchUserOutcome>,
failures: &mut Vec<Value>,
unprocessed_user_ids: &mut Vec<String>,
) {
for pending in items.iter().skip(start_index) {
if outcomes.contains_key(&pending.user_id) {
continue;
}
failures.push(json!({
"user_id": pending.user_id,
"reason": "因前序错误未执行",
}));
unprocessed_user_ids.push(pending.user_id.clone());
}
}
fn build_admin_user_batch_idempotency_conflict_response() -> Response<Body> {
(
http::StatusCode::CONFLICT,
Json(json!({
"detail": "idempotency_key was already used with a different request",
"error_code": "idempotency_key_conflict",
})),
)
.into_response()
}
fn parse_resolve_selection_request(
request_body: Option<&Bytes>,
) -> Result<AdminUserSelectionRequest, String> {
@@ -286,6 +857,7 @@ fn parse_batch_action_request(
selection: parse_selection_request_value(raw.selection)?,
action: raw.action,
payload: raw.payload,
idempotency_key: raw.idempotency_key,
})
}
_ => Err("Invalid JSON request body".to_string()),
@@ -604,10 +1176,37 @@ fn parse_batch_mutation(
}),
"update_access_control" => parse_access_control_mutation(payload),
"update_role" => parse_role_mutation(payload),
"adjust_wallet_balance" => parse_wallet_balance_adjustment_mutation(payload),
_ => Err("不支持的批量操作".to_string()),
}
}
fn parse_wallet_balance_adjustment_mutation(
payload: Option<Value>,
) -> Result<AdminUserBatchMutation, String> {
let Some(Value::Object(payload)) = payload else {
return Err("payload 必须是对象".to_string());
};
let operation = match payload.get("operation").and_then(Value::as_str) {
Some("add") => AdminUserWalletBalanceOperation::Add,
Some("deduct") => AdminUserWalletBalanceOperation::Deduct,
_ => return Err("operation 必须为 add 或 deduct".to_string()),
};
let amount = payload
.get("amount")
.and_then(Value::as_f64)
.ok_or_else(|| "amount 必须为大于 0 的有限数字".to_string())?;
if !amount.is_finite() || amount <= 0.0 {
return Err("amount 必须为大于 0 的有限数字".to_string());
}
Ok(AdminUserBatchMutation {
wallet_balance_adjustment: Some(AdminUserWalletBalanceAdjustment { operation, amount }),
modified_fields: vec!["wallet_balance"],
..AdminUserBatchMutation::default()
})
}
fn parse_role_mutation(payload: Option<Value>) -> Result<AdminUserBatchMutation, String> {
let Some(Value::Object(payload)) = payload else {
return Err("payload 必须是对象".to_string());
@@ -700,30 +1299,68 @@ async fn apply_batch_user_wallet_limit_mode(
state: &AdminAppState<'_>,
user_id: &str,
unlimited: bool,
) -> Result<bool, GatewayError> {
) -> Result<bool, AdminBatchWalletLimitModeError> {
let desired_limit_mode = if unlimited { "unlimited" } else { "finite" };
match state
let wallet = state
.find_wallet(aether_data::repository::wallet::WalletLookupKey::UserId(
user_id,
))
.await?
{
.await
.map_err(|_| AdminBatchWalletLimitModeError::WalletLookup)?;
match wallet {
Some(wallet) => {
if wallet.limit_mode.eq_ignore_ascii_case(desired_limit_mode) {
return Ok(true);
}
Ok(state
.update_auth_user_wallet_limit_mode(user_id, desired_limit_mode)
.await?
.await
.map_err(|_| AdminBatchWalletLimitModeError::Mutation)?
.is_some())
}
None => Ok(state
.initialize_auth_user_wallet(user_id, 0.0, unlimited)
.await?
.await
.map_err(|_| AdminBatchWalletLimitModeError::Mutation)?
.is_some()),
}
}
async fn apply_batch_user_wallet_balance_adjustment(
state: &AdminAppState<'_>,
user_id: &str,
adjustment: AdminUserWalletBalanceAdjustment,
operator_id: Option<&str>,
) -> Result<bool, AdminBatchWalletBalanceAdjustmentError> {
// Resolve only the wallet ID; the repository clamps the deduction under its row lock.
let Some(wallet) = state
.find_wallet(aether_data::repository::wallet::WalletLookupKey::UserId(
user_id,
))
.await
.map_err(|_| AdminBatchWalletBalanceAdjustmentError::WalletLookup)?
else {
return Ok(false);
};
let amount = match adjustment.operation {
AdminUserWalletBalanceOperation::Add => adjustment.amount,
AdminUserWalletBalanceOperation::Deduct => -adjustment.amount,
};
state
.admin_adjust_wallet_balance(
&wallet.id,
amount,
"recharge",
operator_id,
Some("管理员批量调整用户余额"),
true,
)
.await
.map_err(|_| AdminBatchWalletBalanceAdjustmentError::BalanceAdjustment)
.map(|result| result.is_some())
}
fn build_admin_user_batch_bad_request_response(detail: String) -> Response<Body> {
if detail.as_str() == "缺少 user_id" {
return build_admin_users_bad_request_response("缺少 user_id");
@@ -49,12 +49,14 @@ use self::shared::AdminUpdateUserPatch;
use self::shared::{
admin_default_user_initial_gift, build_admin_users_bad_request_response,
build_admin_users_data_unavailable_response, build_admin_users_permission_denied_response,
build_admin_users_read_only_response, disabled_user_policy_detail, disabled_user_policy_field,
format_optional_datetime_iso8601, legacy_admin_list_policy_mode,
legacy_admin_rate_limit_policy_mode, management_token_may_administer_user_accounts,
normalize_admin_optional_user_email, normalize_admin_user_group_ids, normalize_admin_user_role,
normalize_admin_username, validate_admin_user_password, AdminCreateUserApiKeyRequest,
AdminCreateUserRequest, AdminToggleUserApiKeyLockRequest, AdminUpdateUserApiKeyRequest,
build_admin_users_read_only_response, build_admin_users_wallet_permission_denied_response,
disabled_user_policy_detail, disabled_user_policy_field, format_optional_datetime_iso8601,
legacy_admin_list_policy_mode, legacy_admin_rate_limit_policy_mode,
management_token_may_adjust_admin_wallet_balance,
management_token_may_administer_user_accounts, normalize_admin_optional_user_email,
normalize_admin_user_group_ids, normalize_admin_user_role, normalize_admin_username,
validate_admin_user_password, AdminCreateUserApiKeyRequest, AdminCreateUserRequest,
AdminToggleUserApiKeyLockRequest, AdminUpdateUserApiKeyRequest,
};
pub(crate) use self::shared::{
normalize_admin_list_policy_mode, normalize_admin_rate_limit_policy_mode,
@@ -165,6 +165,18 @@ pub(super) fn management_token_may_administer_user_accounts(
})
}
pub(super) fn management_token_may_adjust_admin_wallet_balance(
request_context: &crate::handlers::admin::request::AdminRequestContext<'_>,
) -> bool {
request_context.decision().is_some_and(|decision| {
crate::control::management_token_principal_has_permission(decision, "admin:wallets:write")
|| crate::control::management_token_principal_has_permission(
decision,
"admin:wallets:admin",
)
})
}
pub(super) fn build_admin_users_permission_denied_response(
request_context: &crate::handlers::admin::request::AdminRequestContext<'_>,
) -> Response<Body> {
@@ -192,6 +204,34 @@ pub(super) fn build_admin_users_permission_denied_response(
)
}
pub(super) fn build_admin_users_wallet_permission_denied_response(
request_context: &crate::handlers::admin::request::AdminRequestContext<'_>,
) -> Response<Body> {
let actor_id = request_context
.decision()
.and_then(|decision| decision.admin_principal.as_ref())
.and_then(|principal| principal.management_token_id.as_deref())
.unwrap_or("unknown");
crate::handlers::admin::shared::attach_admin_audit_response(
(
http::StatusCode::FORBIDDEN,
Json(json!({
"detail": "management token permission denied",
"required_permissions": ["admin:wallets:write", "admin:wallets:admin"],
"permission_mode": "any_of",
"route_family": request_context.route_family(),
"route_kind": request_context.route_kind(),
"request_path": request_context.path(),
})),
)
.into_response(),
"admin_user_wallet_balance_permission_denied",
"permission_denied",
"admin_user_wallet_balance",
actor_id,
)
}
pub(super) fn normalize_admin_optional_user_email(
value: Option<&str>,
) -> Result<Option<String>, String> {
@@ -397,9 +437,52 @@ pub(super) fn format_optional_datetime_iso8601(
#[cfg(test)]
mod tests {
use super::{normalize_admin_user_api_formats, AdminUpdateUserApiKeyRequest};
use super::{
build_admin_users_wallet_permission_denied_response, normalize_admin_user_api_formats,
AdminUpdateUserApiKeyRequest,
};
use crate::control::{GatewayControlDecision, GatewayPublicRequestContext};
use crate::handlers::admin::request::AdminRequestContext;
use axum::http::{HeaderMap, Method, Uri};
use serde_json::json;
#[test]
fn wallet_permission_denial_uses_wallet_audit_category() {
let uri: Uri = "/api/admin/users/batch-action"
.parse()
.expect("uri should parse");
let method = Method::POST;
let headers = HeaderMap::new();
let decision = GatewayControlDecision::synthetic(
uri.path(),
Some("admin_proxy".to_string()),
Some("users_manage".to_string()),
Some("batch_user_action".to_string()),
Some("admin:users".to_string()),
);
let context = GatewayPublicRequestContext::from_request_parts(
"trace-wallet-permission-denied",
&method,
&uri,
&headers,
Some(decision),
);
let request_context = AdminRequestContext::new(&context);
let response = build_admin_users_wallet_permission_denied_response(&request_context);
let event = response
.extensions()
.get::<crate::audit::AdminAuditEvent>()
.expect("wallet denial should attach an audit event");
assert_eq!(
event.event_name,
"admin_user_wallet_balance_permission_denied"
);
assert_eq!(event.action, "permission_denied");
assert_eq!(event.target_type, "admin_user_wallet_balance");
}
#[test]
fn admin_user_api_formats_accept_current_canonical_signatures() {
assert_eq!(
@@ -90,6 +90,8 @@ pub(super) struct ResponsesWebSocketContinuationRecord {
/// request JSON can never set it.
#[serde(default)]
deepseek_opaque_reasoning_replay: bool,
#[serde(default)]
xai_encrypted_reasoning_replay: bool,
/// A prior turn stored PII sentinels whose restore mapping exists only on
/// the original downstream socket. Such a chain cannot safely resume on a
/// new socket without leaking sentinels, so lookup succeeds but bootstrap
@@ -122,6 +124,10 @@ impl ResponsesWebSocketContinuationRecord {
normalization.reasoning_replay_policy(),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
),
xai_encrypted_reasoning_replay: matches!(
normalization.reasoning_replay_policy(),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
),
has_connection_local_redaction,
responses_lite_static_config,
};
@@ -156,7 +162,9 @@ impl ResponsesWebSocketContinuationRecord {
pub(super) fn reasoning_replay_policy(
&self,
) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy {
if self.deepseek_opaque_reasoning_replay {
if self.xai_encrypted_reasoning_replay {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
} else if self.deepseek_opaque_reasoning_replay {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
} else {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
@@ -476,6 +484,7 @@ mod tests {
binding_fingerprint: [7; 32],
normalization_fingerprint: [9; 32],
deepseek_opaque_reasoning_replay: false,
xai_encrypted_reasoning_replay: false,
has_connection_local_redaction: false,
responses_lite_static_config: Some(ResponsesLiteStaticConfig::from_response_create(
&json!({
@@ -714,6 +723,29 @@ mod tests {
assert_eq!(decoded, record());
}
#[test]
fn serialized_record_preserves_xai_replay_policy_and_reads_legacy_records() {
let mut expected = record();
expected.xai_encrypted_reasoning_replay = true;
let mut serialized = serde_json::to_value(&expected).unwrap();
let decoded: ResponsesWebSocketContinuationRecord =
serde_json::from_value(serialized.clone()).unwrap();
assert_eq!(
decoded.reasoning_replay_policy(),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
);
serialized
.as_object_mut()
.unwrap()
.remove("xai_encrypted_reasoning_replay");
let legacy: ResponsesWebSocketContinuationRecord =
serde_json::from_value(serialized).unwrap();
assert_eq!(
legacy.reasoning_replay_policy(),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
);
}
#[test]
fn serialized_record_preserves_only_the_server_derived_reasoning_replay_policy_bit() {
let mut expected = record();
@@ -567,6 +567,7 @@ fn build_users_me_usage_record_payload(
"id": item.id,
"model": item.model,
"target_model": serde_json::Value::Null,
"response_model": item.provider_response_model(),
"api_format": item.api_format,
"endpoint_api_format": item.endpoint_api_format,
"has_format_conversion": item.has_format_conversion,
@@ -681,6 +682,7 @@ fn build_users_me_usage_active_payload(item: &StoredRequestUsageAudit) -> serde_
"client_ip": users_me_usage_metadata_string(item, "client_ip"),
"user_agent": users_me_usage_metadata_string(item, "user_agent"),
"target_model": item.target_model,
"response_model": item.provider_response_model(),
"has_fallback": item.has_fallback(),
});
payload["end_to_end_time_ms"] = json!(users_me_usage_metadata_u64(item, "end_to_end_time_ms"));
@@ -1865,6 +1867,25 @@ mod tests {
assert_eq!(payload["cache_creation_ephemeral_1h_input_tokens"], 6);
}
#[test]
fn user_usage_payloads_expose_response_model_separately_from_mapping() {
let item = StoredRequestUsageAudit {
target_model: Some("provider-mapped-model".to_string()),
request_metadata: Some(json!({
"provider_response_model": "gpt-5.1"
})),
..sample_usage("completed")
};
let record = build_users_me_usage_record_payload(&item, false, &BTreeMap::new(), false);
let active = build_users_me_usage_active_payload(&item);
for payload in [&record, &active] {
assert_eq!(payload["target_model"], "provider-mapped-model");
assert_eq!(payload["response_model"], "gpt-5.1");
}
}
#[test]
fn user_usage_payloads_project_end_to_end_timings_from_metadata() {
let item = StoredRequestUsageAudit {
@@ -21,6 +21,7 @@ use aether_crypto::{
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use aether_provider_pool::{
grok_pool_tier_from_quota_bucket, grok_supported_quota_windows_for_tier,
provider_pool_codex_metadata_has_account_quota,
};
use aether_scheduler_core::provider_key_circuit_payload_is_active_open_at;
use serde_json::{json, Map, Value};
@@ -1111,9 +1112,8 @@ fn build_codex_quota_status_snapshot(
source: &str,
) -> Option<Value> {
let metadata = provider_quota_metadata_bucket(upstream_metadata, "codex")?;
let observed_at_unix_secs = metadata
.get("updated_at")
.and_then(admin_provider_quota_pure::coerce_json_u64);
let observed_at_unix_secs = provider_quota_timestamp_unix_secs(metadata.get("observed_at"))
.or_else(|| provider_quota_timestamp_unix_secs(metadata.get("updated_at")));
let plan_type = metadata
.get("plan_type")
.and_then(Value::as_str)
@@ -1388,6 +1388,168 @@ fn build_kiro_quota_status_snapshot(
}))
}
fn build_xai_quota_status_snapshot(
upstream_metadata: Option<&Value>,
source: &str,
) -> Option<Value> {
let metadata = provider_quota_metadata_bucket(upstream_metadata, "xai")?;
let observed_at_unix_secs = provider_quota_timestamp_unix_secs(metadata.get("updated_at"));
let usage_limit = metadata
.get("usage_limit")
.and_then(admin_provider_quota_pure::coerce_json_f64);
let current_usage = metadata
.get("current_usage")
.and_then(admin_provider_quota_pure::coerce_json_f64);
let remaining = metadata
.get("remaining")
.and_then(admin_provider_quota_pure::coerce_json_f64);
let usage_ratio = metadata
.get("usage_percentage")
.and_then(admin_provider_quota_pure::coerce_json_f64)
.map(|value| (value / 100.0).clamp(0.0, 1.0))
.or_else(|| {
current_usage
.zip(usage_limit)
.and_then(|(current_usage, usage_limit)| {
(usage_limit > 0.0).then_some((current_usage / usage_limit).clamp(0.0, 1.0))
})
});
let remaining_ratio = usage_ratio.map(|value| (1.0 - value).max(0.0));
let next_reset_at = provider_quota_timestamp_unix_secs(metadata.get("next_reset_at"));
let reset_seconds = quota_window_reset_seconds(observed_at_unix_secs, next_reset_at);
let plan_type = metadata
.get("subscription_title")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let period_type = metadata
.get("period_type")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let usage_label = match period_type.as_deref() {
Some("monthly") => "月额度",
Some("weekly") => "周额度",
_ => "额度",
};
let mut windows = Vec::new();
if usage_ratio.is_some()
|| remaining.is_some()
|| usage_limit.is_some()
|| current_usage.is_some()
|| next_reset_at.is_some()
{
windows.push(json!({
"code": "usage",
"label": usage_label,
"scope": "account",
"unit": if usage_limit.is_some() { "usd" } else { "percent" },
"used_ratio": usage_ratio,
"remaining_ratio": remaining_ratio,
"used_value": current_usage,
"remaining_value": remaining,
"limit_value": usage_limit,
"reset_at": next_reset_at,
"reset_seconds": reset_seconds,
}));
}
let prepaid_balance = metadata
.get("prepaid_balance")
.and_then(admin_provider_quota_pure::coerce_json_f64);
if prepaid_balance.is_some_and(|value| value > 0.0) {
windows.push(json!({
"code": "prepaid",
"label": "预付额度",
"scope": "account",
"unit": "usd",
"used_ratio": serde_json::Value::Null,
"remaining_ratio": serde_json::Value::Null,
"remaining_value": prepaid_balance,
"reset_at": serde_json::Value::Null,
"reset_seconds": serde_json::Value::Null,
}));
}
let on_demand_cap = metadata
.get("on_demand_cap")
.and_then(admin_provider_quota_pure::coerce_json_f64);
let on_demand_used = metadata
.get("on_demand_used")
.and_then(admin_provider_quota_pure::coerce_json_f64);
let on_demand_enabled = metadata
.get("on_demand_enabled")
.and_then(admin_provider_quota_pure::coerce_json_bool)
!= Some(false);
if on_demand_enabled && on_demand_cap.is_some_and(|value| value > 0.0) {
let on_demand_remaining = on_demand_cap
.zip(on_demand_used)
.map(|(cap, used)| (cap - used).max(0.0));
let on_demand_ratio = on_demand_cap
.zip(on_demand_used)
.and_then(|(cap, used)| (cap > 0.0).then_some((used / cap).clamp(0.0, 1.0)));
windows.push(json!({
"code": "on_demand",
"label": "按需额度",
"scope": "account",
"unit": "usd",
"used_ratio": on_demand_ratio,
"remaining_ratio": on_demand_ratio.map(|value| (1.0 - value).max(0.0)),
"used_value": on_demand_used,
"remaining_value": on_demand_remaining,
"limit_value": on_demand_cap,
"reset_at": serde_json::Value::Null,
"reset_seconds": serde_json::Value::Null,
}));
}
if windows.is_empty() && plan_type.is_none() && observed_at_unix_secs.is_none() {
return None;
}
let prepaid_available = prepaid_balance.is_some_and(|value| value > 0.0);
let on_demand_available = on_demand_enabled
&& on_demand_cap.is_some_and(|value| value > 0.0)
&& on_demand_used
.zip(on_demand_cap)
.is_some_and(|(used, cap)| used < cap);
let usage_exhausted = remaining.is_some_and(|value| value <= 0.0)
|| usage_ratio.is_some_and(|value| value >= 1.0 - 1e-6);
let exhausted = usage_exhausted && !prepaid_available && !on_demand_available;
let reason = if exhausted {
Some("额度已耗尽".to_string())
} else {
None
};
let label = if exhausted {
Some("额度耗尽")
} else {
None
};
let code = if exhausted { "exhausted" } else { "ok" };
Some(json!({
"version": 2,
"provider_type": "xai",
"code": code,
"label": label,
"reason": reason,
"freshness": "fresh",
"source": source,
"observed_at": observed_at_unix_secs,
"exhausted": exhausted,
"usage_ratio": usage_ratio,
"updated_at": observed_at_unix_secs,
"reset_at": next_reset_at,
"reset_seconds": reset_seconds,
"plan_type": plan_type,
"windows": windows,
}))
}
fn build_chatgpt_web_quota_status_snapshot(
upstream_metadata: Option<&Value>,
source: &str,
@@ -2132,6 +2294,111 @@ fn build_gemini_cli_quota_status_snapshot(
}))
}
fn build_claude_code_quota_status_snapshot(
upstream_metadata: Option<&Value>,
source: &str,
) -> Option<Value> {
let metadata = provider_quota_metadata_bucket(upstream_metadata, "claude_code")?;
let observed_at_unix_secs = provider_quota_timestamp_unix_secs(metadata.get("updated_at"));
// (metadata prefix, window code, window minutes, account-wide?). Display labels are
// resolved by the frontend from `code` so they follow the UI locale.
let definitions: [(&str, &str, u64, bool); 4] = [
("five_hour", "5h", 300, true),
("seven_day", "weekly", 10_080, true),
("seven_day_sonnet", "weekly_sonnet", 10_080, false),
("seven_day_fable", "weekly_fable", 10_080, false),
];
let mut windows = Vec::new();
for (prefix, code, window_minutes, account_wide) in definitions {
let used_percent = metadata
.get(&format!("{prefix}_used_percent"))
.and_then(Value::as_f64);
let reset_at =
provider_quota_timestamp_unix_secs(metadata.get(&format!("{prefix}_reset_at")));
// A window whose reset time already passed no longer describes current usage.
let expired = reset_at
.zip(observed_at_unix_secs)
.is_some_and(|(reset_at, observed_at)| reset_at <= observed_at);
let Some(used_percent) = used_percent else {
continue;
};
let used_ratio = if expired {
0.0
} else {
(used_percent / 100.0).clamp(0.0, 1.0)
};
let reset_seconds = reset_at
.zip(observed_at_unix_secs)
.map(|(reset_at, observed_at)| reset_at.saturating_sub(observed_at));
let mut window = json!({
"code": code,
"scope": if account_wide { "account" } else { "model" },
"unit": "percent",
"used_ratio": used_ratio,
"remaining_ratio": 1.0 - used_ratio,
"reset_at": reset_at,
"reset_seconds": reset_seconds,
"window_minutes": window_minutes,
"is_exhausted": used_ratio >= 1.0 - 1e-6,
});
if !account_wide {
window["quota_group"] = json!(code);
}
windows.push(window);
}
if windows.is_empty() {
return None;
}
let account_windows = windows
.iter()
.filter(|window| window.get("scope").and_then(Value::as_str) == Some("account"))
.cloned()
.collect::<Vec<_>>();
let blocking_windows = account_windows
.iter()
.filter(|window| window.get("is_exhausted").and_then(Value::as_bool) == Some(true))
.cloned()
.collect::<Vec<_>>();
let exhausted = !blocking_windows.is_empty();
// The account is usable again only once every exhausted window resets.
let reset_at = if exhausted {
blocking_windows
.iter()
.filter_map(|window| provider_quota_timestamp_unix_secs(window.get("reset_at")))
.max()
} else {
None
};
let reset_seconds = if exhausted {
blocking_windows
.iter()
.filter_map(|window| window.get("reset_seconds").and_then(Value::as_u64))
.max()
} else {
None
};
Some(json!({
"version": 2,
"provider_type": "claude_code",
"code": if exhausted { "exhausted" } else { "ok" },
"freshness": "fresh",
"source": source,
"observed_at": observed_at_unix_secs,
"exhausted": exhausted,
"usage_ratio": quota_windows_usage_ratio(&account_windows),
"updated_at": observed_at_unix_secs,
"reset_at": reset_at,
"reset_seconds": reset_seconds,
"reset_credits": build_codex_reset_credits_status_snapshot(
metadata,
observed_at_unix_secs,
),
"windows": windows,
}))
}
fn build_codex_reset_credits_status_snapshot(
metadata: &Map<String, Value>,
observed_at_unix_secs: Option<u64>,
@@ -2255,11 +2522,13 @@ pub(crate) fn sync_provider_key_quota_status_snapshot(
let mut quota = match normalized_provider_type.as_str() {
"codex" => build_codex_quota_status_snapshot(upstream_metadata, source),
"kiro" => build_kiro_quota_status_snapshot(upstream_metadata, source),
"xai" => build_xai_quota_status_snapshot(upstream_metadata, source),
"chatgpt_web" => build_chatgpt_web_quota_status_snapshot(upstream_metadata, source),
"windsurf" => build_windsurf_quota_status_snapshot(upstream_metadata, source),
"antigravity" => build_antigravity_quota_status_snapshot(upstream_metadata, source),
"grok" => build_grok_quota_status_snapshot(upstream_metadata, source),
"gemini_cli" => build_gemini_cli_quota_status_snapshot(upstream_metadata, source),
"claude_code" => build_claude_code_quota_status_snapshot(upstream_metadata, source),
_ => None,
}?;
if normalized_provider_type == "codex" {
@@ -2339,17 +2608,19 @@ fn codex_upstream_metadata_is_at_least_as_fresh(
let Some(metadata) = provider_quota_metadata_bucket(upstream_metadata, "codex") else {
return false;
};
let Some(metadata_updated_at) = metadata
.get("updated_at")
.and_then(admin_provider_quota_pure::coerce_json_u64)
// Identity, reset-credit, and model-only updates do not replace the
// account's quota observation, even when their timestamp is newer.
if !provider_pool_codex_metadata_has_account_quota(metadata) {
return false;
}
let Some(metadata_updated_at) = provider_quota_timestamp_unix_secs(metadata.get("observed_at"))
.or_else(|| provider_quota_timestamp_unix_secs(metadata.get("updated_at")))
else {
return false;
};
let snapshot_updated_at = quota_snapshot.and_then(|quota| {
quota
.get("updated_at")
.or_else(|| quota.get("observed_at"))
.and_then(admin_provider_quota_pure::coerce_json_u64)
provider_quota_timestamp_unix_secs(quota.get("observed_at"))
.or_else(|| provider_quota_timestamp_unix_secs(quota.get("updated_at")))
});
snapshot_updated_at.is_none_or(|updated_at| metadata_updated_at >= updated_at)
@@ -2448,6 +2719,19 @@ pub(crate) fn provider_key_status_snapshot_payload(
let mut snapshot = provider_key_status_snapshot_object(Some(&payload))
.or_else(|| default_provider_key_status_snapshot().as_object().cloned())
.unwrap_or_default();
// Legacy snapshots can retain an exhausted summary after a window reset or
// newer quota observation. Use the same decision as scheduling so the
// account list and its status filter do not keep displaying that stale block.
if provider_type.trim().eq_ignore_ascii_case("codex")
&& !aether_provider_pool::provider_pool_key_account_quota_exhausted(key, provider_type)
{
if let Some(quota) = snapshot.get_mut("quota").and_then(Value::as_object_mut) {
quota.insert("exhausted".to_string(), json!(false));
if quota.get("code").and_then(Value::as_str) == Some("exhausted") {
quota.insert("code".to_string(), json!("ok"));
}
}
}
snapshot.insert(
"oauth".to_string(),
build_provider_key_oauth_status_snapshot(key),
@@ -3559,6 +3843,59 @@ mod tests {
assert_eq!(window.get("used_value"), Some(&json!(0.0)));
}
#[test]
fn provider_key_status_snapshot_payload_backfills_claude_code_usage_windows() {
let mut key = sample_catalog_key();
key.upstream_metadata = Some(json!({
"claude_code": {
"updated_at": 1_800_000_000u64,
"five_hour_used_percent": 100.0,
"five_hour_reset_at": 1_800_003_600u64,
"seven_day_used_percent": 40.0,
"seven_day_reset_at": 1_800_400_000u64,
"seven_day_sonnet_used_percent": 10.0,
"seven_day_sonnet_reset_at": 1_800_400_000u64,
"reset_credits": {
"available_count": 2,
"updated_at": 1_800_000_000u64,
"detail_source": "claude_oauth_usage",
"credits": [{
"display_key": "Key-1",
"status": "available",
"expires_at": 1_800_144_000u64
}]
}
}
}));
let payload = provider_key_status_snapshot_payload(&key, "claude_code");
let quota = payload
.get("quota")
.and_then(Value::as_object)
.expect("quota snapshot should be object");
assert_eq!(quota.get("provider_type"), Some(&json!("claude_code")));
// An exhausted 5h window blocks the whole account until it resets.
assert_eq!(quota.get("exhausted"), Some(&json!(true)));
assert_eq!(quota.get("reset_at"), Some(&json!(1_800_003_600u64)));
let windows = quota
.get("windows")
.and_then(Value::as_array)
.expect("windows should exist");
assert_eq!(windows.len(), 3);
assert_eq!(windows[0]["code"], json!("5h"));
assert_eq!(windows[0]["scope"], json!("account"));
assert_eq!(windows[0]["window_minutes"], json!(300));
assert_eq!(windows[1]["code"], json!("weekly"));
assert_eq!(windows[1]["used_ratio"], json!(0.4));
assert_eq!(windows[2]["code"], json!("weekly_sonnet"));
assert_eq!(windows[2]["scope"], json!("model"));
assert_eq!(quota["reset_credits"]["available_count"], json!(2));
assert_eq!(
quota["reset_credits"]["credits"][0]["remaining_seconds"],
json!(144_000u64)
);
}
#[test]
fn provider_key_status_snapshot_payload_backfills_grok_model_quota() {
let mut key = sample_catalog_key();
@@ -3622,6 +3959,43 @@ mod tests {
assert_eq!(auto.get("used_value"), Some(&json!(90.0)));
}
#[test]
fn provider_key_status_snapshot_payload_backfills_xai_weekly_credits() {
let mut key = sample_catalog_key();
key.upstream_metadata = Some(json!({
"xai": {
"updated_at": 1_778_067_246u64,
"usage_percentage": 46.0,
"period_type": "weekly",
"next_reset_at": 1_778_157_172u64,
"subscription_title": "SuperGrok",
"prepaid_balance": 0.0,
"on_demand_cap": 0.0,
"on_demand_used": 0.0
}
}));
let payload = provider_key_status_snapshot_payload(&key, "xai");
let quota = payload
.get("quota")
.and_then(Value::as_object)
.expect("quota snapshot should be object");
let windows = quota
.get("windows")
.and_then(Value::as_array)
.expect("xai quota windows should exist");
assert_eq!(quota.get("provider_type"), Some(&json!("xai")));
assert_eq!(quota.get("code"), Some(&json!("ok")));
assert_eq!(quota.get("exhausted"), Some(&json!(false)));
assert_eq!(quota.get("plan_type"), Some(&json!("SuperGrok")));
assert_eq!(quota.get("usage_ratio"), Some(&json!(0.46)));
assert_eq!(quota.get("reset_at"), Some(&json!(1_778_157_172u64)));
assert_eq!(windows.len(), 1);
assert_eq!(windows[0].get("code"), Some(&json!("usage")));
assert_eq!(windows[0].get("label"), Some(&json!("周额度")));
}
#[test]
fn provider_key_status_snapshot_payload_backfills_gemini_cli_account_credits() {
let mut key = sample_catalog_key();
@@ -4013,6 +4387,116 @@ mod tests {
);
}
#[test]
fn provider_key_status_snapshot_payload_clears_stale_codex_exhaustion_summary() {
let mut key = sample_catalog_key();
key.status_snapshot = Some(json!({
"quota": {
"provider_type": "codex",
"code": "exhausted",
"exhausted": true,
"windows": [{
"code": "weekly",
"scope": "account",
"used_ratio": 0.83,
"remaining_ratio": 0.17,
"reset_at": 4_102_444_800u64
}]
}
}));
let payload = provider_key_status_snapshot_payload(&key, "codex");
assert_eq!(payload["quota"]["code"], "ok");
assert_eq!(payload["quota"]["exhausted"], false);
assert!(payload["quota"]["label"].is_null());
// An explicit current upstream refusal is not a stale percentage summary.
key.status_snapshot.as_mut().unwrap()["quota"]["allowed"] = json!(false);
let payload = provider_key_status_snapshot_payload(&key, "codex");
assert_eq!(payload["quota"]["code"], "exhausted");
assert_eq!(payload["quota"]["exhausted"], true);
// Missing capacity evidence must not clear an exhausted summary either.
key.status_snapshot = Some(json!({"quota": {
"provider_type": "codex", "code": "exhausted", "exhausted": true
}}));
let payload = provider_key_status_snapshot_payload(&key, "codex");
assert_eq!(payload["quota"]["code"], "exhausted");
assert_eq!(payload["quota"]["exhausted"], true);
}
#[test]
fn provider_key_status_snapshot_payload_refreshes_codex_timestamp_formats() {
for updated_at in [
json!(1_900_000_000u64),
json!(1_900_000_000_000u64),
json!("2030-03-17T17:46:40Z"),
] {
let mut key = sample_catalog_key();
key.upstream_metadata = Some(json!({
"codex": {
"updated_at": updated_at,
"primary_used_percent": 83.0,
"primary_reset_at": 4_102_444_800u64
}
}));
key.status_snapshot = Some(json!({
"quota": {
"provider_type": "codex",
"updated_at": 1_899_999_000u64,
"code": "exhausted",
"exhausted": true,
"allowed": false,
"windows": [{
"code": "weekly",
"scope": "account",
"used_ratio": 1.0,
"remaining_ratio": 0.0,
"reset_at": 4_102_444_800u64
}]
}
}));
let payload = provider_key_status_snapshot_payload(&key, "codex");
assert_eq!(payload["quota"]["code"], "ok");
assert_eq!(payload["quota"]["updated_at"], 1_900_000_000u64);
assert_eq!(payload["quota"]["windows"][0]["used_ratio"], 0.83);
assert!(payload["quota"]["allowed"].is_null());
}
}
#[test]
fn provider_key_status_snapshot_payload_preserves_codex_account_quota_on_unrelated_updates() {
for patch in [
json!({"plan_type": "pro"}),
json!({"spark_primary_used_percent": 83.0}),
json!({"credits_unlimited": false}),
json!({"windows": [{"code": "weekly", "reset_at": 4_102_444_800u64}]}),
] {
let mut metadata = patch;
metadata["updated_at"] = json!("2030-03-17T17:46:40Z");
let mut key = sample_catalog_key();
key.upstream_metadata = Some(json!({"codex": metadata}));
key.status_snapshot = Some(json!({"quota": {
"provider_type": "codex", "updated_at": 200,
"code": "exhausted", "exhausted": true, "allowed": false,
"windows": [{"code": "weekly", "scope": "account", "used_ratio": 1.0,
"remaining_ratio": 0.0, "reset_at": 4_102_444_800u64}]
}}));
let payload = provider_key_status_snapshot_payload(&key, "codex");
assert_eq!(payload["quota"]["code"], "exhausted", "{metadata}");
assert_eq!(payload["quota"]["exhausted"], true, "{metadata}");
assert_eq!(payload["quota"]["allowed"], false, "{metadata}");
assert_eq!(payload["quota"]["updated_at"], 200, "{metadata}");
assert_eq!(
payload["quota"]["windows"][0]["code"], "weekly",
"{metadata}"
);
assert_eq!(
payload["quota"]["windows"][0]["used_ratio"], 1.0,
"{metadata}"
);
}
}
#[test]
fn provider_key_status_snapshot_payload_restores_complete_codex_cache() {
let mut key = sample_catalog_key();
@@ -276,6 +276,10 @@ pub(crate) fn admin_proxy_local_requires_buffered_body(
| (Some("system_manage"), http::Method::POST, Some("config_import"))
| (Some("system_manage"), http::Method::POST, Some("users_import"))
| (Some("system_manage"), http::Method::POST, Some("data_import"))
| (Some("system_manage"), http::Method::POST, Some("cleanup_usage_manual"))
| (Some("system_manage"), http::Method::POST, Some("smtp_test"))
| (Some("system_manage"), http::Method::POST, Some("prepare_update"))
| (Some("system_manage"), http::Method::POST, Some("apply_update"))
| (Some("system_manage"), http::Method::PUT, Some("settings_set"))
| (Some("system_manage"), http::Method::PUT, Some("config_set"))
| (Some("system_manage"), http::Method::PUT, Some("email_template_set"))
@@ -611,4 +615,27 @@ mod tests {
"/v1/chat/completions?key=passthrough"
);
}
#[test]
fn manual_cleanup_route_requires_buffered_body() {
use crate::control::GatewayPublicRequestContext;
let mut decision = GatewayControlDecision::synthetic(
"/api/admin/system/cleanup/usage/manual",
Some("admin_proxy".to_string()),
Some("system_manage".to_string()),
Some("cleanup_usage_manual".to_string()),
Some("system_manage:cleanup_usage_manual".to_string()),
);
decision.route_class = Some("admin_proxy".to_string());
let uri: http::Uri = "/api/admin/system/cleanup/usage/manual".parse().unwrap();
let headers = http::HeaderMap::new();
let context = GatewayPublicRequestContext::from_request_parts(
"trace-manual-cleanup",
&http::Method::POST,
&uri,
&headers,
Some(decision),
);
assert!(super::admin_proxy_local_requires_buffered_body(&context));
}
}
@@ -13,7 +13,7 @@ pub(crate) fn openai_image_provider_max_generation_count(provider_type: &str) ->
GROK_OPENAI_IMAGE_MAX_GENERATION_COUNT
} else if matches!(
provider_type.trim().to_ascii_lowercase().as_str(),
"openai" | "codex"
"openai" | "codex" | "xai"
) {
OPENAI_IMAGE_MAX_GENERATION_COUNT
} else {
@@ -58,6 +58,7 @@ mod tests {
assert_eq!(openai_image_provider_max_generation_count("grok"), 4);
assert_eq!(openai_image_provider_max_generation_count("openai"), 10);
assert_eq!(openai_image_provider_max_generation_count("codex"), 10);
assert_eq!(openai_image_provider_max_generation_count("xai"), 10);
assert_eq!(openai_image_provider_max_generation_count("custom"), 1);
assert_eq!(
openai_image_provider_max_generation_count_for_model("openai", Some("dall-e-3")),
+2 -1
View File
@@ -37,6 +37,7 @@ mod bark_push;
mod cache;
mod client_session_affinity;
mod clock;
mod codex_profile;
mod constants;
mod control;
mod data;
@@ -92,12 +93,12 @@ mod usage;
mod video_tasks;
mod wallet_runtime;
pub use self::ai_serving::api::{codex_client_originator, codex_client_user_agent};
pub(crate) use self::ai_serving::api::{
AiControlPlanRequest, EXECUTION_RUNTIME_STREAM_DECISION_ACTION,
EXECUTION_RUNTIME_SYNC_DECISION_ACTION, GEMINI_FILES_DOWNLOAD_PLAN_KIND,
OPENAI_VIDEO_CONTENT_PLAN_KIND,
};
pub use self::ai_serving::api::{CODEX_CLIENT_ORIGINATOR, CODEX_CLIENT_USER_AGENT};
pub(crate) use self::ai_serving::{
AiExecutionDecision, AiExecutionPlanPayload, AiStreamAttempt, AiSyncAttempt,
};
+19 -1
View File
@@ -2513,6 +2513,20 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
);
}
}
match state.prewarm_codex_client_profile().await {
Ok(version) => {
info!(
codex_client_version = %version,
"prewarmed Codex client profile"
);
}
Err(err) => {
warn!(
error = %err,
"failed to refresh Codex client profile; built-in or cached profile remains active"
);
}
}
match prewarm_direct_h2c_sender_cache_from_env_for_startup().await {
Ok(Some(report)) => {
if report.failed_targets > 0 {
@@ -4963,7 +4977,11 @@ mod tests {
builder
.http1()
.timer(TokioTimer::new())
.header_read_timeout(std::time::Duration::from_millis(10))
// Keep hyper's own header timeout far from the 5ms first-request
// deadline: when a slow runner lets both expire before the next
// poll, `select!` may pick the connection branch and surface
// hyper's header-timeout error instead of the clean deadline close.
.header_read_timeout(std::time::Duration::from_secs(30))
.max_buf_size(super::MIN_GATEWAY_HTTP_HEADER_MAX_BYTES)
.max_headers(super::MIN_GATEWAY_HTTP_MAX_HEADERS);
builder
+8 -12
View File
@@ -19,6 +19,8 @@ use sha2::{Digest, Sha256};
use tokio::sync::{Mutex, Semaphore};
use tracing::{debug, info, warn};
use crate::ai_serving::api::codex_client_version;
const CODEX_CATALOG_SCHEMA_VERSION: u32 = 2;
const CODEX_CATALOG_CREDENTIAL_SCOPE_DOMAIN: &str = "aether-codex-catalog-credential-v2";
const CODEX_CLIENT_VERSION_MAX_LEN: usize = 64;
@@ -99,7 +101,7 @@ pub(crate) fn normalize_codex_client_version(raw: Option<&str>) -> NormalizedCod
used_fallback: false,
},
None => NormalizedCodexClientVersion {
value: crate::ai_serving::CODEX_CLIENT_VERSION.to_string(),
value: codex_client_version(),
used_fallback: true,
},
}
@@ -1521,7 +1523,7 @@ where
.await?;
let scope = target.credential_scope()?;
let state = runtime.codex_catalog_runtime_state();
let mut version = Version::parse(crate::ai_serving::CODEX_CLIENT_VERSION).ok()?;
let mut version = Version::parse(&codex_client_version()).ok()?;
if let Some(recent) =
read_recent_codex_catalog_client_version(state, provider_id, key_id, scope).await
{
@@ -2461,10 +2463,7 @@ mod tests {
let initial = read_codex_management_catalog(&runtime, TEST_PROVIDER_ID, TEST_KEY_ID)
.await
.expect("management context");
assert_eq!(
initial.client_version,
crate::ai_serving::CODEX_CLIENT_VERSION
);
assert_eq!(initial.client_version, codex_client_version());
assert!(initial.models.is_none());
seed_catalog(&runtime, &version("0.200.0")).await;
@@ -2498,10 +2497,7 @@ mod tests {
let rebound = read_codex_management_catalog(&runtime, TEST_PROVIDER_ID, TEST_KEY_ID)
.await
.unwrap();
assert_eq!(
rebound.client_version,
crate::ai_serving::CODEX_CLIENT_VERSION
);
assert_eq!(rebound.client_version, codex_client_version());
assert!(rebound.models.is_none());
}
@@ -2557,7 +2553,7 @@ mod tests {
format!("1.2.3-{}", "x".repeat(CODEX_CLIENT_VERSION_MAX_LEN)),
] {
let normalized = normalize_codex_client_version(Some(&raw));
assert_eq!(normalized.as_str(), crate::ai_serving::CODEX_CLIENT_VERSION);
assert_eq!(normalized.as_str(), codex_client_version());
assert!(normalized.used_fallback());
assert!(!catalog_lkg_key(&target(), normalized.as_str()).contains(&raw));
}
@@ -3753,7 +3749,7 @@ mod tests {
.await
.expect("seed legacy cache");
let load = load_one(&runtime, &version(crate::ai_serving::CODEX_CLIENT_VERSION)).await;
let load = load_one(&runtime, &version(&codex_client_version())).await;
assert!(load.snapshot(TEST_PROVIDER_ID, TEST_KEY_ID).is_none());
assert_eq!(runtime.execution_count(), 1);
}
@@ -551,6 +551,24 @@ async fn sync_grok_quota_from_report_context(
async fn apply_local_sync_report_effect(state: &AppState, payload: &GatewaySyncReportRequest) {
apply_local_gemini_file_mapping_report_effect(state, payload).await;
if claude_code_quota_headers_reportable(payload.status_code) {
if let Err(err) = sync_claude_code_quota_from_response_headers(
state,
payload.report_context.as_ref(),
&payload.headers,
)
.await
{
warn!(
event_name = "claude_code_realtime_quota_sync_failed",
log_type = "ops",
report_kind = %payload.report_kind,
report_request_id = %short_request_id(report_request_id(payload.report_context.as_ref())),
error = ?err,
"gateway failed to persist claude_code realtime quota from sync response headers"
);
}
}
if (200..300).contains(&payload.status_code) {
if let Err(err) = sync_codex_quota_from_response_headers(
state,
@@ -641,6 +659,24 @@ async fn apply_local_stream_report_effect(state: &AppState, payload: &GatewayStr
);
}
}
if claude_code_quota_headers_reportable(payload.status_code) {
if let Err(err) = sync_claude_code_quota_from_response_headers(
state,
payload.report_context.as_ref(),
&payload.headers,
)
.await
{
warn!(
event_name = "claude_code_realtime_quota_sync_failed",
log_type = "ops",
report_kind = %payload.report_kind,
report_request_id = %short_request_id(report_request_id(payload.report_context.as_ref())),
error = ?err,
"gateway failed to persist claude_code realtime quota from stream response headers"
);
}
}
if let Err(err) = sync_grok_quota_from_report_context(
state,
payload.report_context.as_ref(),
@@ -896,6 +932,134 @@ async fn sync_codex_quota_from_response_headers(
.await
}
fn claude_code_quota_headers_reportable(status_code: u16) -> bool {
// Real limit 429s carry the unified headers (fingerprint-rejection 429s do not, and then
// the parser finds nothing to record).
(200..300).contains(&status_code) || status_code == 429
}
/// Passive sampling of Anthropic's `anthropic-ratelimit-unified-*` response headers into the
/// `claude_code` quota metadata, so the 5H / weekly windows stay fresh between active refreshes.
async fn sync_claude_code_quota_from_response_headers(
state: &AppState,
report_context: Option<&Value>,
headers: &BTreeMap<String, String>,
) -> Result<bool, GatewayError> {
let observed_at_unix_secs = report_context_u64(
report_context,
"provider_response_headers_observed_at_unix_ms",
)
.map(|value| value / 1_000)
.filter(|value| *value > 0)
.unwrap_or_else(current_unix_secs);
let parsed = report_context_provider_response_headers(report_context)
.and_then(|headers| {
admin_provider_quota_pure::parse_claude_code_usage_headers(
&headers,
observed_at_unix_secs,
)
})
.or_else(|| {
admin_provider_quota_pure::parse_claude_code_usage_headers(
headers,
observed_at_unix_secs,
)
});
let Some(parsed) = parsed else {
return Ok(false);
};
let Some(key_id) = report_context_key_id(report_context) else {
return Ok(false);
};
for attempt in 0..RUNTIME_METADATA_CAS_MAX_ATTEMPTS {
let Some(key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
else {
return Ok(false);
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&key.provider_id))
.await?
.into_iter()
.next()
else {
return Ok(false);
};
if !provider
.provider_type
.trim()
.eq_ignore_ascii_case("claude_code")
{
return Ok(false);
}
let expected_namespace_value =
upstream_metadata_namespace_value(key.upstream_metadata.as_ref(), "claude_code");
let mut bucket = expected_namespace_value
.as_ref()
.and_then(Value::as_object)
.cloned()
.unwrap_or_default();
// Never let an older observation overwrite a newer active refresh.
if bucket
.get("updated_at")
.and_then(admin_provider_quota_pure::coerce_json_u64)
.is_some_and(|stored| stored > observed_at_unix_secs)
{
return Ok(false);
}
let Some(patch) = parsed.as_object() else {
return Ok(false);
};
// Headers can describe only some windows; absent windows keep their stored value.
for (field, value) in patch {
bucket.insert(field.clone(), value.clone());
}
let next_bucket = Value::Object(bucket);
if expected_namespace_value.as_ref() == Some(&next_bucket) {
return Ok(false);
}
let updated_upstream_metadata = merge_metadata_object(
key.upstream_metadata.as_ref(),
"claude_code",
next_bucket.clone(),
);
let updated_status_snapshot = sync_provider_key_quota_status_snapshot(
key.status_snapshot.as_ref(),
provider.provider_type.as_str(),
updated_upstream_metadata.as_ref(),
"response_headers",
);
let updated = state
.update_provider_catalog_key_runtime_metadata(
&ProviderCatalogKeyRuntimeMetadataUpdate {
key_id: key_id.clone(),
namespace: "claude_code".to_string(),
expected_upstream_metadata_value: expected_namespace_value,
upstream_metadata_value: next_bucket,
status_snapshot_patch: quota_status_snapshot_patch(
updated_status_snapshot.as_ref(),
),
updated_at_unix_secs: Some(observed_at_unix_secs),
},
)
.await?;
if updated {
return Ok(true);
}
if attempt + 1 < RUNTIME_METADATA_CAS_MAX_ATTEMPTS {
let backoff_us = 50_u64.saturating_mul((attempt + 1) as u64).min(1_000);
tokio::time::sleep(Duration::from_micros(backoff_us)).await;
}
}
Ok(false)
}
async fn sync_codex_websocket_quota_from_stream_payload(
state: &AppState,
payload: &GatewayStreamReportRequest,
@@ -172,6 +172,7 @@ fn provider_uses_bearer_oauth_runtime(provider_type: &str) -> bool {
| "antigravity"
| "kiro"
| "windsurf"
| "xai"
)
}
@@ -406,6 +407,22 @@ mod tests {
);
}
#[test]
fn recognizes_xai_oauth_as_bearer_runtime() {
let semantics = provider_key_auth_semantics(&sample_key("oauth"), "xai");
assert!(semantics.oauth_managed());
assert!(semantics.can_refresh_oauth());
assert_eq!(
semantics.credential_kind(),
ProviderKeyCredentialKind::OAuthSession
);
assert_eq!(
semantics.runtime_auth_kind(),
ProviderKeyRuntimeAuthKind::Bearer
);
}
#[test]
fn refresh_capability_requires_stored_refresh_token() {
let semantics = provider_key_auth_semantics(&sample_key("oauth"), "codex");
+2
View File
@@ -186,6 +186,8 @@ fn frontend_path_bypasses_static(path: &str) -> bool {
"/health" | "/test-connection" | crate::constants::READYZ_PATH
) || path.starts_with("/api/")
|| path.starts_with("/v1/")
|| path == "/openai/v1/videos"
|| path.starts_with("/openai/v1/videos/")
|| path.starts_with("/v1beta/")
|| path.starts_with("/upload/")
|| path.starts_with("/_gateway/")
@@ -23,6 +23,16 @@ pub(super) fn resolve_scheduler_candidate_selectability(
if let Some(skip_reason) =
current_candidate_runtime_skip_reason(&candidate, runtime_snapshot, now_unix_secs)
{
tracing::debug!(
event_name = "scheduler_candidate_skipped",
log_type = "event",
provider_id = %candidate.provider_id,
endpoint_id = %candidate.endpoint_id,
key_id = %candidate.key_id,
api_format = %candidate.endpoint_api_format,
skip_reason,
"scheduler candidate skipped during runtime selectability resolution"
);
if emitted_skipped_keys.insert(key) {
skipped.push(SchedulerSkippedCandidate {
candidate,
@@ -40,10 +40,6 @@ pub(super) async fn read_candidate_runtime_selection_snapshot(
) -> Result<CandidateRuntimeSelectionSnapshot, GatewayError> {
let provider_concurrent_limits = read_provider_concurrent_limits(state, candidates).await?;
let provider_pool_state = read_provider_pool_state_map(state, candidates).await?;
let provider_skip_exhausted_accounts = provider_pool_state
.iter()
.map(|(provider_id, state)| (provider_id.clone(), state.skip_exhausted_accounts))
.collect::<BTreeMap<_, _>>();
let pool_provider_ids = provider_pool_state
.iter()
.filter_map(|(provider_id, state)| state.pool_enabled.then_some(provider_id.clone()))
@@ -62,7 +58,7 @@ pub(super) async fn read_candidate_runtime_selection_snapshot(
let key_account_quota_exhausted = read_key_account_quota_exhaustion_map(
candidates,
&provider_key_rpm_states,
&provider_skip_exhausted_accounts,
&provider_pool_state,
);
let key_oauth_invalid =
read_key_oauth_invalid_map(candidates, &provider_key_rpm_states, now_unix_secs);
@@ -360,6 +356,7 @@ async fn read_provider_quota_block_map(
struct ProviderPoolState {
pool_enabled: bool,
skip_exhausted_accounts: bool,
reserve_minimum_quota: bool,
}
async fn read_provider_pool_state_map(
@@ -391,11 +388,17 @@ async fn read_provider_pool_state_map(
.and_then(|value| value.get("skip_exhausted_accounts"))
.and_then(serde_json::Value::as_bool)
.unwrap_or(false);
let reserve_minimum_quota = pool_advanced
.and_then(serde_json::Value::as_object)
.and_then(|value| value.get("reserve_minimum_quota"))
.and_then(serde_json::Value::as_bool)
.unwrap_or(false);
(
provider.id,
ProviderPoolState {
pool_enabled: pool_advanced.is_some(),
skip_exhausted_accounts,
reserve_minimum_quota,
},
)
})
@@ -405,7 +408,7 @@ async fn read_provider_pool_state_map(
fn read_key_account_quota_exhaustion_map(
candidates: &[SchedulerMinimalCandidateSelectionCandidate],
provider_key_rpm_states: &BTreeMap<String, StoredProviderCatalogKey>,
provider_skip_exhausted_accounts: &BTreeMap<String, bool>,
provider_pool_state: &BTreeMap<String, ProviderPoolState>,
) -> BTreeMap<String, bool> {
candidates
.iter()
@@ -431,11 +434,19 @@ fn read_key_account_quota_exhaustion_map(
candidate.provider_type.as_str(),
candidate.selected_provider_model_name.as_str(),
);
let skip_configured = provider_skip_exhausted_accounts
let pool_state = provider_pool_state
.get(candidate.provider_id.as_str())
.copied()
.unwrap_or(false);
hard_blocked || (skip_configured && account_exhausted)
.unwrap_or_default();
let reserve_reached = pool_state.reserve_minimum_quota
&& admin_provider_pool_pure::admin_pool_key_minimum_quota_reached(
key,
candidate.provider_type.as_str(),
Some(candidate.selected_provider_model_name.as_str()),
);
hard_blocked
|| reserve_reached
|| (pool_state.skip_exhausted_accounts && account_exhausted)
});
(candidate.key_id.clone(), exhausted)
})
@@ -566,3 +577,64 @@ fn read_provider_key_rpm_reset_at_map(
})
.collect::<BTreeMap<_, _>>()
}
#[cfg(test)]
mod reserve_minimum_quota_tests {
use super::*;
use serde_json::json;
#[test]
fn reserve_minimum_quota_is_independent_of_skip_exhausted_accounts() {
let candidate = SchedulerMinimalCandidateSelectionCandidate {
provider_id: "provider-codex".to_string(),
provider_name: "codex".to_string(),
provider_type: "codex".to_string(),
provider_priority: 0,
endpoint_id: "endpoint-codex".to_string(),
endpoint_api_format: "openai:responses".to_string(),
key_id: "key-codex".to_string(),
key_name: "codex".to_string(),
key_auth_type: "oauth".to_string(),
key_internal_priority: 0,
key_global_priority_for_format: None,
key_capabilities: None,
model_id: "model-codex".to_string(),
global_model_id: "global-model-codex".to_string(),
global_model_name: "gpt-5".to_string(),
selected_provider_model_name: "gpt-5".to_string(),
supports_streaming: true,
mapping_matched_model: None,
};
let mut key = StoredProviderCatalogKey::new(
candidate.key_id.clone(),
candidate.provider_id.clone(),
"codex".to_string(),
"oauth".to_string(),
None,
true,
)
.expect("key should build");
for reserve_enabled in [false, true] {
for used_percent in [99.0, 98.0] {
key.upstream_metadata =
Some(json!({"codex": {"primary_used_percent": used_percent}}));
let exhausted = read_key_account_quota_exhaustion_map(
std::slice::from_ref(&candidate),
&BTreeMap::from([(key.id.clone(), key.clone())]),
&BTreeMap::from([(
candidate.provider_id.clone(),
ProviderPoolState {
pool_enabled: true,
reserve_minimum_quota: reserve_enabled,
skip_exhausted_accounts: false,
},
)]),
);
assert_eq!(
exhausted.get(&key.id),
Some(&(reserve_enabled && used_percent >= 99.0))
);
}
}
}
}
+19
View File
@@ -495,6 +495,25 @@ pub struct AppState {
Arc<StdMutex<HashMap<String, aether_data::repository::wallet::StoredWalletSnapshot>>>,
>,
#[cfg(test)]
pub(crate) auth_wallet_adjustment_error_for_tests: Option<String>,
#[cfg(test)]
pub(crate) auth_wallet_lookup_error_for_tests: Option<String>,
#[cfg(test)]
pub(crate) auth_wallet_batch_store_for_tests: Option<
Arc<
StdMutex<
HashMap<
(String, String),
aether_data::repository::wallet::StoredAdminUserWalletBalanceBatch,
>,
>,
>,
>,
#[cfg(test)]
pub(crate) auth_wallet_batch_operation_lock_for_tests: Arc<TokioMutex<()>>,
#[cfg(test)]
pub(crate) auth_wallet_batch_failure_record_error_for_tests: Option<String>,
#[cfg(test)]
pub(crate) admin_wallet_payment_order_store:
Option<Arc<StdMutex<HashMap<String, AdminWalletPaymentOrderRecord>>>>,
#[cfg(test)]
+40
View File
@@ -804,12 +804,43 @@ impl AppState {
if updated.is_some() {
self.invalidate_provider_routing_caches();
}
if let Some(key) = updated.as_ref().filter(|key| !key.is_active) {
self.delete_inactive_provider_catalog_key_pool_scores(
key.provider_id.as_str(),
key.id.as_str(),
)
.await;
}
match updated {
Some(key) => self.open_provider_catalog_key(key).await.map(Some),
None => Ok(None),
}
}
async fn delete_inactive_provider_catalog_key_pool_scores(
&self,
provider_id: &str,
key_id: &str,
) {
if let Err(err) = self
.data
.delete_pool_member_scores_for_member(
&pool_scores::PoolMemberIdentity::provider_api_key(
provider_id.to_string(),
key_id.to_string(),
),
)
.await
{
warn!(
provider_id,
key_id,
error = ?err,
"gateway provider catalog key deactivate: failed to delete pool member scores"
);
}
}
pub(crate) async fn compare_and_update_provider_catalog_key_admin_state(
&self,
update: &provider_catalog::ProviderCatalogKeyAdminCasUpdate,
@@ -843,6 +874,15 @@ impl AppState {
if updated.as_ref().is_some_and(|keys| !keys.is_empty()) {
self.invalidate_provider_routing_caches();
}
if let Some(keys) = updated.as_ref() {
for key in keys.iter().filter(|key| !key.is_active) {
self.delete_inactive_provider_catalog_key_pool_scores(
key.provider_id.as_str(),
key.id.as_str(),
)
.await;
}
}
match updated {
Some(keys) => self.open_provider_catalog_keys(keys).await.map(Some),
None => Ok(None),
+19
View File
@@ -54,6 +54,7 @@ use super::super::router::RequestAdmissionError;
use super::super::{control::GatewayControlDecision, error::GatewayError};
use super::super::{provider_transport, usage};
use crate::codex_profile::spawn_worker as spawn_codex_client_profile_worker;
use crate::maintenance::spawn_account_self_check_worker;
use crate::maintenance::spawn_audit_cleanup_worker;
use crate::maintenance::spawn_db_maintenance_worker;
@@ -149,6 +150,10 @@ fn system_config_key_affects_provider_transport_snapshot(key: &str) -> bool {
}
impl AppState {
pub async fn prewarm_codex_client_profile(&self) -> Result<String, String> {
crate::codex_profile::prewarm(self.runtime_state()).await
}
pub async fn prewarm_chat_pii_redaction_runtime_config(&self) -> Result<bool, String> {
crate::privacy::read_chat_pii_redaction_runtime_config(self)
.await
@@ -471,6 +476,16 @@ impl AppState {
#[cfg(test)]
auth_wallet_store: Some(Arc::new(StdMutex::new(HashMap::new()))),
#[cfg(test)]
auth_wallet_adjustment_error_for_tests: None,
#[cfg(test)]
auth_wallet_lookup_error_for_tests: None,
#[cfg(test)]
auth_wallet_batch_store_for_tests: Some(Arc::new(StdMutex::new(HashMap::new()))),
#[cfg(test)]
auth_wallet_batch_operation_lock_for_tests: Arc::new(TokioMutex::new(())),
#[cfg(test)]
auth_wallet_batch_failure_record_error_for_tests: None,
#[cfg(test)]
admin_wallet_payment_order_store: Some(Arc::new(StdMutex::new(HashMap::new()))),
#[cfg(test)]
admin_payment_callback_store: Some(Arc::new(StdMutex::new(HashMap::new()))),
@@ -2337,6 +2352,10 @@ impl AppState {
crate::task_runtime::TASK_KEY_MODEL_FETCH_WORKER,
spawn_model_fetch_worker(background_state.clone()),
);
supervise_worker(
crate::task_runtime::TASK_KEY_CODEX_CLIENT_PROFILE,
Some(spawn_codex_client_profile_worker(background_state.clone())),
);
supervise_worker(
crate::task_runtime::TASK_KEY_VIDEO_TASK_POLLER,
spawn_video_task_poller(background_state.clone()),
@@ -290,6 +290,14 @@ impl provider_transport::VideoTaskTransportSnapshotLookup for AppState {
.await
.map_err(GatewayError::into_message)
}
async fn resolve_video_task_proxy(
&self,
transport: &GatewayProviderTransportSnapshot,
) -> Option<ProxySnapshot> {
self.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await
}
}
#[async_trait]
@@ -1,8 +1,198 @@
use crate::{AdminWalletPaymentOrderRecord, AdminWalletTransactionRecord, AppState, GatewayError};
use aether_data::repository::wallet::{
AdjustWalletBalanceInBatchInput, AdminUserWalletBalanceBatchUserOutcome,
PrepareAdminUserWalletBalanceBatchInput, PrepareAdminUserWalletBalanceBatchOutcome,
StoredAdminUserWalletBalanceBatch,
};
use std::collections::BTreeMap;
use super::admin_wallet_build_order_no;
impl AppState {
pub(crate) async fn prepare_admin_user_wallet_balance_batch(
&self,
input: PrepareAdminUserWalletBalanceBatchInput,
) -> Result<PrepareAdminUserWalletBalanceBatchOutcome, GatewayError> {
#[cfg(test)]
if let Some(store) = self.auth_wallet_batch_store_for_tests.as_ref() {
let mut batches = store.lock().expect("auth wallet batch store should lock");
let key = (input.admin_user_id.clone(), input.idempotency_key.clone());
if let Some(existing) = batches.get(&key) {
if existing.request_fingerprint != input.request_fingerprint {
return Ok(PrepareAdminUserWalletBalanceBatchOutcome::Conflict);
}
return Ok(PrepareAdminUserWalletBalanceBatchOutcome::Ready(
existing.clone(),
));
}
let batch = StoredAdminUserWalletBalanceBatch {
admin_user_id: input.admin_user_id,
idempotency_key: input.idempotency_key,
request_fingerprint: input.request_fingerprint,
target_user_ids: input.target_user_ids,
missing_user_ids: input.missing_user_ids,
warnings: input.warnings,
user_outcomes: BTreeMap::new(),
};
batches.insert(key, batch.clone());
return Ok(PrepareAdminUserWalletBalanceBatchOutcome::Ready(batch));
}
self.data
.prepare_admin_user_wallet_balance_batch(input)
.await
.map_err(|error| GatewayError::Internal(error.to_string()))?
.ok_or_else(|| {
GatewayError::Internal("admin wallet batch storage is unavailable".to_string())
})
}
pub(crate) async fn get_admin_user_wallet_balance_batch(
&self,
admin_user_id: &str,
idempotency_key: &str,
request_fingerprint: &str,
) -> Result<Option<PrepareAdminUserWalletBalanceBatchOutcome>, GatewayError> {
#[cfg(test)]
if let Some(store) = self.auth_wallet_batch_store_for_tests.as_ref() {
let batches = store.lock().expect("auth wallet batch store should lock");
return Ok(batches
.get(&(admin_user_id.to_string(), idempotency_key.to_string()))
.map(|existing| {
if existing.request_fingerprint != request_fingerprint {
PrepareAdminUserWalletBalanceBatchOutcome::Conflict
} else {
PrepareAdminUserWalletBalanceBatchOutcome::Ready(existing.clone())
}
}));
}
self.data
.get_admin_user_wallet_balance_batch(
admin_user_id,
idempotency_key,
request_fingerprint,
)
.await
.map_err(|error| GatewayError::Internal(error.to_string()))
}
pub(crate) async fn adjust_admin_user_wallet_balance_batch_user(
&self,
input: AdjustWalletBalanceInBatchInput,
) -> Result<AdminUserWalletBalanceBatchUserOutcome, GatewayError> {
#[cfg(test)]
if let Some(store) = self.auth_wallet_batch_store_for_tests.as_ref() {
let _operation = self.auth_wallet_batch_operation_lock_for_tests.lock().await;
let key = (input.admin_user_id.clone(), input.idempotency_key.clone());
if let Some(existing) = store
.lock()
.expect("auth wallet batch store should lock")
.get(&key)
.and_then(|batch| batch.user_outcomes.get(&input.user_id))
.cloned()
{
return Ok(existing);
}
let target_exists = store
.lock()
.expect("auth wallet batch store should lock")
.get(&key)
.is_some_and(|batch| batch.target_user_ids.contains(&input.user_id));
if !target_exists {
return Err(GatewayError::Internal(
"user is outside the prepared admin wallet batch".to_string(),
));
}
let adjustment = input.adjustment;
let result = self
.admin_adjust_wallet_balance(
&adjustment.wallet_id,
adjustment.amount_usd,
&adjustment.balance_type,
adjustment.operator_id.as_deref(),
adjustment.description.as_deref(),
adjustment.clamp_deduction_to_available_balance,
)
.await?;
let outcome = if result.is_some() {
AdminUserWalletBalanceBatchUserOutcome::Succeeded
} else {
AdminUserWalletBalanceBatchUserOutcome::Failed("用户钱包不可用".to_string())
};
if let Some(batch) = store
.lock()
.expect("auth wallet batch store should lock")
.get_mut(&key)
{
batch.user_outcomes.insert(input.user_id, outcome.clone());
}
return Ok(outcome);
}
self.data
.adjust_admin_user_wallet_balance_batch_user(input)
.await
.map_err(|error| GatewayError::Internal(error.to_string()))?
.ok_or_else(|| {
GatewayError::Internal("admin wallet batch storage is unavailable".to_string())
})
}
pub(crate) async fn record_admin_user_wallet_balance_batch_failure(
&self,
admin_user_id: &str,
idempotency_key: &str,
user_id: &str,
reason: &str,
) -> Result<AdminUserWalletBalanceBatchUserOutcome, GatewayError> {
#[cfg(test)]
if self
.auth_wallet_batch_failure_record_error_for_tests
.as_deref()
== Some(user_id)
{
return Err(GatewayError::Internal(
"injected wallet batch failure-record error".to_string(),
));
}
#[cfg(test)]
if let Some(store) = self.auth_wallet_batch_store_for_tests.as_ref() {
let _operation = self.auth_wallet_batch_operation_lock_for_tests.lock().await;
let key = (admin_user_id.to_string(), idempotency_key.to_string());
let mut batches = store.lock().expect("auth wallet batch store should lock");
let batch = batches.get_mut(&key).ok_or_else(|| {
GatewayError::Internal("admin wallet batch was not prepared".to_string())
})?;
if !batch.target_user_ids.iter().any(|target| target == user_id) {
return Err(GatewayError::Internal(
"user is outside the prepared admin wallet batch".to_string(),
));
}
return Ok(batch
.user_outcomes
.entry(user_id.to_string())
.or_insert_with(|| {
AdminUserWalletBalanceBatchUserOutcome::Failed(reason.to_string())
})
.clone());
}
self.data
.record_admin_user_wallet_balance_batch_failure(
admin_user_id,
idempotency_key,
user_id,
reason,
)
.await
.map_err(|error| GatewayError::Internal(error.to_string()))?
.ok_or_else(|| {
GatewayError::Internal("admin wallet batch storage is unavailable".to_string())
})
}
pub(crate) async fn admin_adjust_wallet_balance(
&self,
wallet_id: &str,
@@ -10,13 +200,21 @@ impl AppState {
balance_type: &str,
operator_id: Option<&str>,
description: Option<&str>,
clamp_deduction_to_available_balance: bool,
) -> Result<
Option<(
aether_data::repository::wallet::StoredWalletSnapshot,
AdminWalletTransactionRecord,
Option<AdminWalletTransactionRecord>,
)>,
GatewayError,
> {
#[cfg(test)]
if self.auth_wallet_adjustment_error_for_tests.as_deref() == Some(wallet_id) {
return Err(GatewayError::Internal(
"injected test wallet adjustment failure".to_string(),
));
}
#[cfg(test)]
if let Some(store) = self.auth_wallet_store.as_ref() {
let mut guard = store.lock().expect("auth wallet store should lock");
@@ -27,6 +225,18 @@ impl AppState {
let before_recharge = wallet.balance;
let before_gift = wallet.gift_balance;
let before_total = before_recharge + before_gift;
let amount_usd = if clamp_deduction_to_available_balance && amount_usd < 0.0 {
if before_total < 0.0 {
-before_total
} else {
-(-amount_usd).min(before_total)
}
} else {
amount_usd
};
if amount_usd == 0.0 {
return Ok(Some((wallet.clone(), None)));
}
let mut after_recharge = before_recharge;
let mut after_gift = before_gift;
@@ -90,7 +300,7 @@ impl AppState {
let updated_wallet = wallet.clone();
drop(guard);
self.invalidate_auth_context_cache();
return Ok(Some((updated_wallet, transaction)));
return Ok(Some((updated_wallet, Some(transaction))));
}
Ok(self
@@ -100,10 +310,15 @@ impl AppState {
balance_type: balance_type.to_string(),
operator_id: operator_id.map(ToOwned::to_owned),
description: description.map(ToOwned::to_owned),
clamp_deduction_to_available_balance,
batch_context: None,
})
.await?
.map(|(wallet, transaction)| {
(wallet, stored_wallet_transaction_to_gateway(transaction))
(
wallet,
transaction.map(stored_wallet_transaction_to_gateway),
)
}))
}
@@ -112,7 +112,7 @@ impl AppState {
) -> Result<
Option<(
aether_data::repository::wallet::StoredWalletSnapshot,
aether_data::repository::wallet::StoredAdminWalletTransaction,
Option<aether_data::repository::wallet::StoredAdminWalletTransaction>,
)>,
GatewayError,
> {
@@ -5,6 +5,19 @@ impl AppState {
&self,
lookup: aether_data::repository::wallet::WalletLookupKey<'_>,
) -> Result<Option<aether_data::repository::wallet::StoredWalletSnapshot>, GatewayError> {
#[cfg(test)]
if let Some(failed_user_id) = self.auth_wallet_lookup_error_for_tests.as_deref() {
let lookup_user_id = match &lookup {
aether_data::repository::wallet::WalletLookupKey::UserId(user_id) => Some(*user_id),
_ => None,
};
if lookup_user_id == Some(failed_user_id) {
return Err(GatewayError::Internal(
"injected test wallet lookup failure".to_string(),
));
}
}
#[cfg(test)]
if let Some(store) = self.auth_wallet_store.as_ref() {
let wallet = {
+21
View File
@@ -472,6 +472,27 @@ impl AppState {
self
}
pub(crate) fn fail_auth_wallet_adjustment_for_tests(
mut self,
wallet_id: impl Into<String>,
) -> Self {
self.auth_wallet_adjustment_error_for_tests = Some(wallet_id.into());
self
}
pub(crate) fn fail_auth_wallet_lookup_for_tests(mut self, user_id: impl Into<String>) -> Self {
self.auth_wallet_lookup_error_for_tests = Some(user_id.into());
self
}
pub(crate) fn fail_auth_wallet_batch_failure_record_for_tests(
mut self,
user_id: impl Into<String>,
) -> Self {
self.auth_wallet_batch_failure_record_error_for_tests = Some(user_id.into());
self
}
pub(crate) fn with_admin_wallet_payment_orders_for_tests<I>(mut self, orders: I) -> Self
where
I: IntoIterator<Item = crate::AdminWalletPaymentOrderRecord>,
@@ -24,6 +24,7 @@ pub(crate) const TASK_KEY_USAGE_QUEUE_WORKER: &str = "usage.queue.worker";
pub(crate) const TASK_KEY_USAGE_COUNTER_FLUSH: &str = "usage.counter.flush.worker";
pub(crate) const TASK_KEY_VIDEO_TASK_POLLER: &str = "video.task.poller";
pub(crate) const TASK_KEY_MODEL_FETCH_WORKER: &str = "model.fetch.worker";
pub(crate) const TASK_KEY_CODEX_CLIENT_PROFILE: &str = "maintenance.codex.client.profile";
pub(crate) const TASK_KEY_PROVIDER_QUOTA_RESET: &str = "provider.quota.reset.worker";
pub(crate) const TASK_KEY_ACCOUNT_SELF_CHECK: &str = "account.self_check.worker";
pub(crate) const TASK_KEY_POOL_SCORE_REBUILD: &str = "pool.score.rebuild.worker";
@@ -202,6 +203,14 @@ const TASK_DEFINITIONS: &[TaskDefinition] = &[
true,
RETRY_ONCE,
),
TaskDefinition::new(
TASK_KEY_CODEX_CLIENT_PROFILE,
TaskKind::Scheduled,
"daily",
true,
true,
RETRY_ONCE,
),
TaskDefinition::new(
TASK_KEY_PROVIDER_QUOTA_RESET,
TaskKind::Scheduled,
@@ -1972,7 +1972,7 @@ async fn gateway_executes_openai_chat_antigravity_cross_format_sync_via_local_fi
seen_execution_runtime_request.user_agent,
aether_provider_transport::antigravity::ANTIGRAVITY_REQUEST_USER_AGENT
);
assert_eq!(seen_execution_runtime_request.request_type, "agent");
assert_eq!(seen_execution_runtime_request.request_type, "");
assert_eq!(seen_execution_runtime_request.contents_len, 1);
assert!(!seen_execution_runtime_request.request_has_model);
@@ -893,8 +893,8 @@ async fn gateway_executes_openai_responses_cross_format_function_call_upstream_s
},
{
"type": "function_call",
"id": "call_auto_1",
"call_id": "call_auto_1",
"id": "call_auto_0",
"call_id": "call_auto_0",
"name": "get_weather",
"arguments": "{\"location\":\"Tokyo\"}"
}
@@ -1570,7 +1570,7 @@ async fn gateway_executes_openai_responses_antigravity_cross_format_upstream_str
seen_remote_execution_runtime_request.user_agent,
aether_provider_transport::antigravity::ANTIGRAVITY_REQUEST_USER_AGENT
);
assert_eq!(seen_remote_execution_runtime_request.request_type, "agent");
assert_eq!(seen_remote_execution_runtime_request.request_type, "");
assert_eq!(seen_remote_execution_runtime_request.contents_len, 1);
assert!(!seen_remote_execution_runtime_request.request_has_model);
@@ -2155,7 +2155,7 @@ async fn gateway_executes_antigravity_gemini_cli_sync_upstream_stream_via_local_
seen_remote_execution_runtime_request.user_agent,
aether_provider_transport::antigravity::ANTIGRAVITY_REQUEST_USER_AGENT
);
assert_eq!(seen_remote_execution_runtime_request.request_type, "agent");
assert_eq!(seen_remote_execution_runtime_request.request_type, "");
assert_eq!(seen_remote_execution_runtime_request.contents_len, 0);
assert!((seen_remote_execution_runtime_request.exact_temperature - 0.2).abs() < f64::EPSILON);
assert!(!seen_remote_execution_runtime_request.request_has_model);

Some files were not shown because too many files have changed in this diff Show More