mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-07 18:07:47 +08:00
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:
+1
@@ -0,0 +1 @@
|
||||
test binary
|
||||
@@ -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,
|
||||
|
||||
+13
-3
@@ -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
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+7
@@ -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,
|
||||
|
||||
@@ -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![
|
||||
|
||||
@@ -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(¤t)? {
|
||||
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");
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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(&[]);
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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")),
|
||||
|
||||
@@ -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,
|
||||
};
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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");
|
||||
|
||||
@@ -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))
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user