Merge commit 'refs/pull/481/head' of github-fawney19:fawney19/Aether

# Conflicts:
#	apps/aether-gateway/src/ai_serving/api.rs
#	apps/aether-gateway/src/ai_serving/planner/standard/family/request.rs
#	apps/aether-gateway/src/ai_serving/planner/standard/openai/chat/decision/request.rs
#	apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/request.rs
This commit is contained in:
fawney19
2026-05-19 01:46:41 +08:00
124 changed files with 14402 additions and 702 deletions

View File

@@ -65,6 +65,8 @@ tracing.workspace = true
url.workspace = true
uuid.workspace = true
webpki-roots.workspace = true
wreq.workspace = true
wreq-util.workspace = true
[target.'cfg(not(target_env = "msvc"))'.dependencies]
tikv-jemallocator = "0.6"

View File

@@ -41,18 +41,21 @@ pub(crate) use crate::ai_serving::{
AiExecutionDecision, AiExecutionPlanPayload, AiStreamAttempt, AiSyncAttempt,
};
pub(crate) use aether_ai_formats::api::{
build_core_error_body_for_client_format, core_error_background_report_kind,
core_error_default_client_api_format, core_success_background_report_kind,
encode_kiro_sse_events, implicit_sync_finalize_report_kind, is_core_error_finalize_kind,
build_core_error_body_for_client_format, convert_standard_chat_response,
core_error_background_report_kind, core_error_default_client_api_format,
core_success_background_report_kind, encode_kiro_sse_events,
implicit_sync_finalize_report_kind, is_core_error_finalize_kind,
normalize_provider_private_report_context, normalize_provider_private_response_value,
provider_private_response_allows_sync_finalize, resolve_claude_stream_spec,
resolve_claude_sync_spec, resolve_gemini_stream_spec, resolve_gemini_sync_spec,
resolve_local_image_stream_spec, resolve_local_image_sync_spec,
resolve_local_same_format_stream_spec, resolve_local_same_format_sync_spec,
sanitize_request_path_and_query, AiControlPlanRequest, ExecutionRuntimeAuthContext,
sanitize_request_path_and_query, AiControlPlanRequest, CanonicalContentPart,
CanonicalStreamEvent, CanonicalStreamFrame, ClaudeClientEmitter, ExecutionRuntimeAuthContext,
LocalCoreSyncErrorKind, LocalOpenAiImageSpec, LocalSameFormatProviderFamily,
LocalSameFormatProviderSpec, LocalStandardSourceFamily, LocalStandardSourceMode,
LocalStandardSpec, StreamingStandardTerminalObserver, EXECUTION_RUNTIME_STREAM_DECISION_ACTION,
LocalStandardSpec, OpenAIChatClientEmitter, OpenAIResponsesClientEmitter,
StreamingStandardTerminalObserver, EXECUTION_RUNTIME_STREAM_DECISION_ACTION,
EXECUTION_RUNTIME_SYNC_DECISION_ACTION, GEMINI_FILES_DOWNLOAD_PLAN_KIND,
GEMINI_VIDEO_CANCEL_SYNC_PLAN_KIND, OPENAI_EMBEDDING_SYNC_PLAN_KIND,
OPENAI_IMAGE_STREAM_PLAN_KIND, OPENAI_IMAGE_SYNC_FINALIZE_REPORT_KIND,
@@ -60,6 +63,7 @@ pub(crate) use aether_ai_formats::api::{
OPENAI_VIDEO_CONTENT_PLAN_KIND, OPENAI_VIDEO_DELETE_SYNC_PLAN_KIND,
OPENAI_VIDEO_REMIX_SYNC_PLAN_KIND,
};
pub(crate) use aether_ai_formats::protocol::stream::CanonicalUsage as StreamingCanonicalUsage;
pub(crate) fn parse_direct_request_body(
parts: &http::request::Parts,

View File

@@ -70,7 +70,10 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
let proxy = state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport)
.await;
let transport_profile = resolve_transport_profile(&resolved.transport);
let transport_profile = resolved
.transport_profile
.clone()
.or_else(|| resolve_transport_profile(&resolved.transport));
let mut extra_fields = serde_json::Map::new();
if let Some(proxy_value) =
build_request_trace_proxy_value(Some(&resolved.transport), proxy.as_ref())
@@ -154,6 +157,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
upstream_url,
provider_request_headers,
provider_request_body,
transport_profile: _,
} = resolved;
let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {

View File

@@ -1,6 +1,7 @@
use std::collections::BTreeMap;
use std::sync::Arc;
use aether_contracts::ResolvedTransportProfile;
use serde_json::Value;
use crate::ai_serving::planner::common::{
@@ -12,7 +13,8 @@ use crate::ai_serving::transport::antigravity::{
AntigravityRequestEnvelopeSupport, AntigravityRequestSideSupport,
};
use crate::ai_serving::transport::{
build_same_format_provider_headers, SameFormatProviderHeadersInput,
build_grok_browser_headers, build_grok_upstream_url, build_same_format_provider_headers,
GrokHeaderInput, SameFormatProviderHeadersInput, GROK_CHAT_PATH,
};
use crate::ai_serving::{CandidateFailureDiagnostic, GatewayProviderTransportSnapshot};
use crate::AppState;
@@ -96,6 +98,7 @@ pub(crate) struct LocalSameFormatProviderCandidatePayloadParts {
pub(super) upstream_url: String,
pub(super) provider_request_headers: BTreeMap<String, String>,
pub(super) provider_request_body: Value,
pub(super) transport_profile: Option<ResolvedTransportProfile>,
}
pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
@@ -247,15 +250,27 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
base_provider_request_body
};
let Some(upstream_url) = super::super::request::build_same_format_upstream_url(
parts,
&prepared.transport,
&prepared.mapped_model,
prepared.provider_api_format.as_str(),
spec,
prepared.upstream_is_stream,
prepared.kiro_auth.as_ref(),
) else {
let is_grok = prepared
.transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("grok");
let transport_profile =
crate::ai_serving::transport::resolve_transport_profile(&prepared.transport);
let Some(upstream_url) = (if is_grok {
Some(build_grok_upstream_url(&prepared.transport, GROK_CHAT_PATH))
} else {
super::super::request::build_same_format_upstream_url(
parts,
&prepared.transport,
&prepared.mapped_model,
prepared.provider_api_format.as_str(),
spec,
prepared.upstream_is_stream,
prepared.kiro_auth.as_ref(),
)
}) else {
mark_skipped_local_same_format_provider_candidate_with_failure_diagnostic(
state,
input,
@@ -278,7 +293,18 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
.as_ref()
.map(build_antigravity_static_identity_headers)
.unwrap_or_default();
let Some(provider_request_headers) =
let Some(provider_request_headers) = (if is_grok {
build_grok_browser_headers(GrokHeaderInput {
transport: &prepared.transport,
transport_profile: transport_profile.as_ref(),
request_headers: Some(effective_headers),
content_type: "application/json",
accept: "text/event-stream",
header_rules: prepared.transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body,
original_request_body: body_json,
})
} else {
build_same_format_provider_headers(SameFormatProviderHeadersInput {
headers: effective_headers,
provider_request_body: &provider_request_body,
@@ -295,7 +321,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
.as_ref()
.map(|auth| auth.machine_id.as_str()),
})
else {
}) else {
mark_skipped_local_same_format_provider_candidate_with_failure_diagnostic(
state,
input,
@@ -327,5 +353,6 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
upstream_url,
provider_request_headers,
provider_request_body,
transport_profile,
})
}

View File

@@ -63,7 +63,10 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
.app()
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&transport)
.await;
let transport_profile = resolve_transport_profile(&transport);
let transport_profile = resolved
.transport_profile
.clone()
.or_else(|| resolve_transport_profile(&transport));
let mut extra_fields = serde_json::Map::new();
if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) {
extra_fields.insert("proxy".to_string(), proxy_value);

View File

@@ -1,17 +1,19 @@
use std::collections::BTreeMap;
use std::sync::Arc;
use aether_contracts::ResolvedTransportProfile;
use serde_json::Value;
use crate::ai_serving::planner::candidate_preparation::{
prepare_header_authenticated_candidate, OauthPreparationContext,
};
use crate::ai_serving::planner::spec_metadata::local_openai_image_spec_metadata;
use crate::ai_serving::pure::normalize_openai_image_request_with_options;
use crate::ai_serving::transport::{
build_openai_image_headers, build_openai_image_upstream_url,
build_standard_provider_request_headers, openai_image_transport_unsupported_reason,
resolve_openai_image_auth, ProviderOpenAiImageHeadersInput,
StandardProviderRequestHeadersInput,
build_grok_browser_headers, build_grok_upstream_url, build_openai_image_headers,
build_openai_image_upstream_url, build_standard_provider_request_headers,
openai_image_transport_unsupported_reason, resolve_openai_image_auth, GrokHeaderInput,
ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput, GROK_CHAT_PATH,
};
use crate::ai_serving::{
apply_codex_openai_responses_special_body_edits, apply_codex_openai_responses_special_headers,
@@ -21,6 +23,7 @@ use crate::ai_serving::{
normalize_openai_image_request, request_conversion_direct_auth, CandidateFailureDiagnostic,
GatewayProviderTransportSnapshot, PlannerAppState, RequestConversionKind,
};
use crate::image_capabilities::openai_image_normalize_options_for_provider;
use crate::AppState;
use super::support::{
@@ -43,6 +46,7 @@ pub(super) struct LocalOpenAiImageCandidatePayloadParts {
pub(super) provider_request_body: Value,
pub(super) upstream_url: String,
pub(super) input_summary: Value,
pub(super) transport_profile: Option<ResolvedTransportProfile>,
}
pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
@@ -121,8 +125,13 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
let auth_header = prepared_candidate.auth_header;
let auth_value = prepared_candidate.auth_value;
let Some(normalized_request) = normalize_openai_image_request(parts, body_json, body_base64)
else {
let normalized_request = normalize_openai_image_request_with_options(
parts,
body_json,
body_base64,
openai_image_normalize_options_for_provider(&transport.provider.provider_type),
);
let Some(normalized_request) = normalized_request else {
mark_skipped_local_openai_image_candidate_with_failure_diagnostic(
state,
input,
@@ -146,8 +155,16 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
.provider_type
.trim()
.eq_ignore_ascii_case("chatgpt_web");
let is_grok = transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("grok");
let transport_profile = crate::ai_serving::transport::resolve_transport_profile(transport);
let upstream_url = if is_chatgpt_web {
chatgpt_web_image_internal_url(&transport.endpoint.base_url)
} else if is_grok {
build_grok_upstream_url(transport, GROK_CHAT_PATH)
} else {
build_openai_image_upstream_url(transport, parts.uri.query())
};
@@ -169,7 +186,18 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
);
}
let Some(mut provider_request_headers) =
let Some(mut provider_request_headers) = (if is_grok {
build_grok_browser_headers(GrokHeaderInput {
transport,
transport_profile: transport_profile.as_ref(),
request_headers: Some(effective_headers),
content_type: "application/json",
accept: "*/*",
header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body,
original_request_body: body_json,
})
} else {
build_openai_image_headers(ProviderOpenAiImageHeadersInput {
headers: effective_headers,
auth_header: &auth_header,
@@ -178,7 +206,7 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
provider_request_body: &provider_request_body,
original_request_body: body_json,
})
else {
}) else {
mark_skipped_local_openai_image_candidate_with_failure_diagnostic(
state,
input,
@@ -198,6 +226,7 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
};
if is_chatgpt_web {
provider_request_headers.insert("x-aether-chatgpt-web-image".to_string(), "1".to_string());
} else if is_grok {
} else {
apply_codex_openai_responses_special_headers(
&mut provider_request_headers,
@@ -223,7 +252,7 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
.unwrap_or_default()
.to_string();
let input_summary = if is_chatgpt_web {
let input_summary = if is_chatgpt_web || is_grok {
provider_request_body.clone()
} else {
normalized_request.summary_json
@@ -240,6 +269,7 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
provider_request_body,
upstream_url,
input_summary,
transport_profile,
})
}
@@ -425,6 +455,7 @@ async fn resolve_local_openai_image_to_gemini_candidate_payload_parts(
provider_request_body: converted.body_json,
upstream_url,
input_summary: converted.summary_json,
transport_profile: None,
})
}

View File

@@ -149,7 +149,10 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
),
&resolved.transport,
);
let transport_profile = resolve_transport_profile(&resolved.transport);
let transport_profile = resolved
.transport_profile
.clone()
.or_else(|| resolve_transport_profile(&resolved.transport));
let timeouts = resolve_transport_execution_timeouts(&resolved.transport);
let super::request::LocalStandardCandidatePayloadParts {
auth_header,
@@ -162,6 +165,7 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
upstream_is_stream,
envelope_name: _,
transport,
transport_profile: _,
} = resolved;
let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {

View File

@@ -1,6 +1,7 @@
use std::collections::BTreeMap;
use std::sync::Arc;
use aether_contracts::ResolvedTransportProfile;
use serde_json::Value;
use crate::ai_serving::planner::candidate_preparation::{
@@ -21,10 +22,11 @@ use crate::ai_serving::transport::kiro::{
KIRO_ENVELOPE_NAME,
};
use crate::ai_serving::transport::{
build_kiro_cross_format_upstream_url, build_openai_image_headers,
build_openai_image_upstream_url, build_standard_provider_request_headers,
openai_image_transport_unsupported_reason, resolve_openai_image_auth,
ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput,
build_grok_browser_headers, build_grok_upstream_url, build_kiro_cross_format_upstream_url,
build_openai_image_headers, build_openai_image_upstream_url,
build_standard_provider_request_headers, openai_image_transport_unsupported_reason,
resolve_grok_session_auth, resolve_openai_image_auth, GrokHeaderInput,
ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput, GROK_CHAT_PATH,
};
use crate::ai_serving::{
build_openai_image_request_body_from_gemini_image_request, gemini_request_is_image_generation,
@@ -49,6 +51,14 @@ pub(crate) struct LocalStandardCandidatePayloadParts {
pub(super) upstream_is_stream: bool,
pub(super) envelope_name: Option<&'static str>,
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
pub(super) transport_profile: Option<ResolvedTransportProfile>,
}
fn is_grok_text_provider_api_format(provider_api_format: &str) -> bool {
matches!(
crate::ai_serving::normalize_api_format_alias(provider_api_format).as_str(),
"openai:chat" | "openai:responses" | "openai:responses:compact" | "claude:messages"
)
}
pub(crate) async fn resolve_local_standard_candidate_payload_parts(
@@ -64,8 +74,14 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
let planner_state = crate::ai_serving::PlannerAppState::new(state);
let candidate = &attempt.eligible.candidate;
let transport = &attempt.eligible.transport;
let transport_profile = crate::ai_serving::transport::resolve_transport_profile(transport);
let provider_api_format = attempt.eligible.provider_api_format.as_str();
let effective_headers = input.effective_headers(&parts.headers);
let is_grok = transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("grok");
if spec_metadata.api_format == "gemini:generate_content"
&& provider_api_format == "openai:image"
&& gemini_request_is_image_generation(body_json)
@@ -76,6 +92,104 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
.await;
}
let is_kiro_claude_cli = is_kiro_claude_messages_transport(transport, provider_api_format);
if is_grok && is_grok_text_provider_api_format(provider_api_format) {
let prepared_candidate = match prepare_header_authenticated_candidate(
planner_state,
transport,
candidate,
resolve_grok_session_auth(transport),
OauthPreparationContext {
trace_id,
api_format: provider_api_format,
operation: "standard_family_grok_text_request",
},
)
.await
{
Ok(prepared) => prepared,
Err(skip_reason) => {
mark_skipped_local_standard_candidate(
state,
input,
trace_id,
candidate,
attempt.candidate_index,
&attempt.candidate_id,
skip_reason,
)
.await;
return None;
}
};
let mut provider_request_body = body_json.clone();
if let Some(object) = provider_request_body.as_object_mut() {
object.insert(
"model".to_string(),
serde_json::Value::String(prepared_candidate.mapped_model.clone()),
);
}
let upstream_is_stream = resolve_upstream_is_stream_for_provider(
transport.endpoint.config.as_ref(),
transport.provider.provider_type.as_str(),
provider_api_format,
spec_metadata.require_streaming,
false,
);
let force_body_stream_field =
endpoint_config_forces_body_stream_field(transport.endpoint.config.as_ref());
enforce_provider_body_stream_policy(
&mut provider_request_body,
provider_api_format,
upstream_is_stream,
request_requires_body_stream_field(body_json, force_body_stream_field),
);
let upstream_url = build_grok_upstream_url(transport, GROK_CHAT_PATH);
let Some(provider_request_headers) = build_grok_browser_headers(GrokHeaderInput {
transport,
transport_profile: transport_profile.as_ref(),
request_headers: Some(effective_headers),
content_type: "application/json",
accept: "text/event-stream",
header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body,
original_request_body: body_json,
}) else {
mark_skipped_local_standard_candidate_with_failure_diagnostic(
state,
input,
trace_id,
candidate,
attempt.candidate_index,
&attempt.candidate_id,
"transport_header_rules_apply_failed",
CandidateFailureDiagnostic::header_rules_apply_failed(
spec_metadata.api_format,
provider_api_format,
"grok_standard_family_headers",
),
)
.await;
return None;
};
return Some(LocalStandardCandidatePayloadParts {
auth_header: prepared_candidate.auth_header,
auth_value: prepared_candidate.auth_value,
mapped_model: prepared_candidate.mapped_model,
provider_api_format: provider_api_format.to_string(),
provider_request_body,
provider_request_headers,
upstream_url,
upstream_is_stream,
envelope_name: None,
transport: Arc::clone(transport),
transport_profile,
});
}
let Some(conversion_kind) =
crate::ai_serving::request_conversion_kind(spec_metadata.api_format, provider_api_format)
else {
@@ -362,6 +476,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
upstream_is_stream,
envelope_name: None,
transport: Arc::clone(transport),
transport_profile: None,
})
}
@@ -498,6 +613,7 @@ async fn resolve_local_gemini_image_to_openai_image_candidate_payload_parts(
upstream_is_stream,
envelope_name: None,
transport: Arc::clone(transport),
transport_profile: None,
})
}
@@ -617,5 +733,6 @@ async fn build_kiro_cross_format_payload_parts(
upstream_is_stream,
envelope_name: Some(KIRO_ENVELOPE_NAME),
transport: Arc::clone(transport),
transport_profile: None,
})
}

View File

@@ -68,7 +68,10 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
let proxy = state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport)
.await;
let transport_profile = resolve_transport_profile(&resolved.transport);
let transport_profile = resolved
.transport_profile
.clone()
.or_else(|| resolve_transport_profile(&resolved.transport));
let timeouts = resolve_transport_execution_timeouts(&resolved.transport);
let mut extra_fields = serde_json::Map::new();
if let Some(proxy_value) =
@@ -100,6 +103,7 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
envelope_name,
transport,
request_redacted,
transport_profile: _,
} = resolved;
let original_request_body_json = if request_redacted {
Some(&provider_request_body)

View File

@@ -3,6 +3,7 @@ use std::collections::BTreeMap;
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
use aether_contracts::ResolvedTransportProfile;
use serde_json::Value;
use crate::ai_serving::planner::candidate_preparation::{
@@ -27,8 +28,9 @@ use crate::ai_serving::transport::kiro::{
};
use crate::ai_serving::transport::local_openai_chat_transport_unsupported_reason;
use crate::ai_serving::transport::{
build_kiro_cross_format_upstream_url, build_standard_provider_request_headers,
StandardProviderRequestHeadersInput,
build_grok_browser_headers, build_grok_upstream_url, build_kiro_cross_format_upstream_url,
build_standard_provider_request_headers, GrokHeaderInput, StandardProviderRequestHeadersInput,
GROK_CHAT_PATH,
};
use crate::ai_serving::{
ai_local_execution_contract_for_formats, request_conversion_direct_auth,
@@ -64,6 +66,14 @@ pub(crate) struct LocalOpenAiChatCandidatePayloadParts {
pub(super) envelope_name: Option<&'static str>,
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
pub(super) request_redacted: bool,
pub(super) transport_profile: Option<ResolvedTransportProfile>,
}
fn is_grok_text_provider_api_format(provider_api_format: &str) -> bool {
matches!(
crate::ai_serving::normalize_api_format_alias(provider_api_format).as_str(),
"openai:chat" | "openai:responses" | "openai:responses:compact" | "claude:messages"
)
}
fn request_identity_response_encoding_when_redacted(
@@ -177,6 +187,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
let candidate = &eligible.candidate;
let provider_api_format = eligible.provider_api_format.as_str();
let transport = &eligible.transport;
let transport_profile = crate::ai_serving::transport::resolve_transport_profile(transport);
let force_body_stream_field =
endpoint_config_forces_body_stream_field(transport.endpoint.config.as_ref());
let enable_model_directives =
@@ -191,6 +202,129 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
.await?;
let body_json = redaction.body_json.as_ref();
let effective_headers = input.effective_headers(&parts.headers);
let is_grok = transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("grok");
if is_grok && is_grok_text_provider_api_format(provider_api_format) {
let prepared_candidate = match prepare_header_authenticated_candidate(
planner_state,
transport,
candidate,
crate::ai_serving::transport::resolve_grok_session_auth(transport),
OauthPreparationContext {
trace_id,
api_format: provider_api_format,
operation: "openai_chat_same_format",
},
)
.await
{
Ok(prepared) => prepared,
Err(skip_reason) => {
mark_skipped_local_openai_chat_candidate(
state,
input,
trace_id,
candidate,
candidate_index,
candidate_id,
skip_reason,
)
.await;
return Ok(None);
}
};
let Some(provider_request_body) = build_local_openai_chat_request_body(
body_json,
&prepared_candidate.mapped_model,
upstream_is_stream,
force_body_stream_field,
transport.endpoint.body_rules.as_ref(),
effective_headers,
enable_model_directives,
) else {
mark_skipped_local_openai_chat_candidate_with_extra_data(
state,
input,
trace_id,
candidate,
candidate_index,
candidate_id,
"provider_request_body_build_failed",
request_body_build_failure_extra_data(
body_json,
"openai:chat",
provider_api_format,
),
)
.await;
return Ok(None);
};
let upstream_url = build_grok_upstream_url(transport, GROK_CHAT_PATH);
let Some(mut provider_request_headers) = build_grok_browser_headers(GrokHeaderInput {
transport,
transport_profile: transport_profile.as_ref(),
request_headers: Some(effective_headers),
content_type: "application/json",
accept: "text/event-stream",
header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body,
original_request_body: body_json,
}) else {
mark_skipped_local_openai_chat_candidate_with_failure_diagnostic(
state,
input,
trace_id,
candidate,
candidate_index,
candidate_id,
"transport_header_rules_apply_failed",
CandidateFailureDiagnostic::header_rules_apply_failed(
"openai:chat",
provider_api_format,
"grok_openai_chat_headers",
),
)
.await;
return Ok(None);
};
let (execution_strategy, conversion_mode) =
ai_local_execution_contract_for_formats("openai:chat", provider_api_format);
let resolved_report_kind =
if decision_kind == OPENAI_CHAT_STREAM_PLAN_KIND || !upstream_is_stream {
report_kind.to_string()
} else {
"openai_chat_sync_finalize".to_string()
};
request_identity_response_encoding_when_redacted(
&mut provider_request_headers,
redaction.redacted,
);
return Ok(Some(LocalOpenAiChatCandidatePayloadParts {
auth_header: prepared_candidate.auth_header,
auth_value: prepared_candidate.auth_value,
mapped_model: prepared_candidate.mapped_model,
provider_api_format: provider_api_format.to_string(),
provider_request_body,
provider_request_headers,
upstream_url,
execution_strategy,
conversion_mode,
report_kind: resolved_report_kind,
envelope_name: None,
transport: Arc::clone(transport),
request_redacted: redaction.redacted,
transport_profile,
}));
}
if provider_api_format == "openai:chat" {
if let Some(skip_reason) = local_openai_chat_transport_unsupported_reason(transport) {
@@ -352,6 +486,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
envelope_name: None,
transport: Arc::clone(transport),
request_redacted: redaction.redacted,
transport_profile,
}));
};
@@ -640,6 +775,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
envelope_name: None,
transport: Arc::clone(transport),
request_redacted: redaction.redacted,
transport_profile: None,
}))
}
@@ -777,6 +913,7 @@ async fn build_kiro_openai_chat_cross_format_payload_parts(
envelope_name: Some(KIRO_ENVELOPE_NAME),
transport: Arc::clone(transport),
request_redacted,
transport_profile: None,
})
}

View File

@@ -67,7 +67,10 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
let proxy = state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport)
.await;
let transport_profile = resolve_transport_profile(&resolved.transport);
let transport_profile = resolved
.transport_profile
.clone()
.or_else(|| resolve_transport_profile(&resolved.transport));
let timeouts = resolve_transport_execution_timeouts(&resolved.transport);
let mut extra_fields = serde_json::Map::new();
if let Some(proxy_value) =
@@ -175,6 +178,7 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
envelope_name: _,
upstream_is_stream,
transport,
transport_profile: _,
} = resolved;
let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {

View File

@@ -1,6 +1,7 @@
use std::collections::BTreeMap;
use std::sync::Arc;
use aether_contracts::ResolvedTransportProfile;
use serde_json::Value;
use tracing::debug;
@@ -35,8 +36,10 @@ use crate::ai_serving::transport::kiro::{
KiroRequestAuth, KIRO_ENVELOPE_NAME,
};
use crate::ai_serving::transport::{
build_kiro_cross_format_upstream_url, build_standard_provider_request_headers,
local_standard_transport_unsupported_reason_with_network, StandardProviderRequestHeadersInput,
build_grok_browser_headers, build_grok_upstream_url, build_kiro_cross_format_upstream_url,
build_standard_provider_request_headers,
local_standard_transport_unsupported_reason_with_network, GrokHeaderInput,
StandardProviderRequestHeadersInput, GROK_CHAT_PATH,
};
use crate::ai_serving::{
ai_local_execution_contract_for_formats, request_conversion_direct_auth,
@@ -56,6 +59,13 @@ use super::LocalOpenAiResponsesSpec;
const ANTIGRAVITY_ENVELOPE_NAME: &str = "antigravity:v1internal";
fn is_grok_text_provider_api_format(provider_api_format: &str) -> bool {
matches!(
crate::ai_serving::normalize_api_format_alias(provider_api_format).as_str(),
"openai:chat" | "openai:responses" | "openai:responses:compact" | "claude:messages"
)
}
pub(crate) struct LocalOpenAiResponsesCandidatePayloadParts {
pub(super) auth_header: String,
pub(super) auth_value: String,
@@ -70,6 +80,7 @@ pub(crate) struct LocalOpenAiResponsesCandidatePayloadParts {
pub(super) envelope_name: Option<&'static str>,
pub(super) upstream_is_stream: bool,
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
pub(super) transport_profile: Option<ResolvedTransportProfile>,
}
#[allow(clippy::too_many_arguments)]
@@ -90,12 +101,22 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
let candidate = &eligible.candidate;
let provider_api_format = eligible.provider_api_format.as_str();
let transport = &eligible.transport;
let transport_profile = crate::ai_serving::transport::resolve_transport_profile(transport);
let is_antigravity = is_antigravity_provider_transport(transport);
let is_kiro_claude_cli = is_kiro_claude_messages_transport(transport, provider_api_format);
let is_grok = transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("grok");
let same_format = api_format_alias_matches(provider_api_format, &client_api_format);
let conversion_kind = request_conversion_kind(spec_metadata.api_format, provider_api_format);
let transport_unsupported_reason = if same_format && is_kiro_claude_cli {
let transport_unsupported_reason = if is_grok
&& is_grok_text_provider_api_format(provider_api_format)
{
None
} else if same_format && is_kiro_claude_cli {
local_kiro_request_transport_unsupported_reason_with_network(transport)
} else if same_format {
local_standard_transport_unsupported_reason_with_network(transport, provider_api_format)
@@ -154,7 +175,9 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
None
};
let direct_auth = if kiro_auth.is_some() {
let direct_auth = if is_grok && is_grok_text_provider_api_format(provider_api_format) {
crate::ai_serving::transport::resolve_grok_session_auth(transport)
} else if kiro_auth.is_some() {
None
} else if same_format {
match crate::ai_serving::normalize_api_format_alias(provider_api_format).as_str() {
@@ -237,42 +260,57 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
let force_body_stream_field =
endpoint_config_forces_body_stream_field(transport.endpoint.config.as_ref());
let effective_headers = input.effective_headers(&parts.headers);
let Some(mut base_provider_request_body) = (if needs_bidirectional_conversion {
build_cross_format_openai_responses_request_body(
body_json,
&mapped_model,
spec_metadata.api_format,
provider_api_format,
upstream_is_stream,
force_body_stream_field,
transport.provider.provider_type.as_str(),
if is_kiro_claude_cli {
None
} else {
transport.endpoint.body_rules.as_ref()
},
Some(input.auth_context.api_key_id.as_str()),
effective_headers,
enable_model_directives,
)
} else {
build_local_openai_responses_request_body(
body_json,
&mapped_model,
upstream_is_stream,
force_body_stream_field,
transport.provider.provider_type.as_str(),
provider_api_format,
if is_kiro_claude_cli {
None
} else {
transport.endpoint.body_rules.as_ref()
},
Some(input.auth_context.api_key_id.as_str()),
effective_headers,
enable_model_directives,
)
}) else {
let Some(mut base_provider_request_body) =
(if is_grok && is_grok_text_provider_api_format(provider_api_format) {
build_local_openai_responses_request_body(
body_json,
&mapped_model,
upstream_is_stream,
force_body_stream_field,
transport.provider.provider_type.as_str(),
spec_metadata.api_format,
transport.endpoint.body_rules.as_ref(),
Some(input.auth_context.api_key_id.as_str()),
effective_headers,
enable_model_directives,
)
} else if needs_bidirectional_conversion {
build_cross_format_openai_responses_request_body(
body_json,
&mapped_model,
spec_metadata.api_format,
provider_api_format,
upstream_is_stream,
force_body_stream_field,
transport.provider.provider_type.as_str(),
if is_kiro_claude_cli {
None
} else {
transport.endpoint.body_rules.as_ref()
},
Some(input.auth_context.api_key_id.as_str()),
effective_headers,
enable_model_directives,
)
} else {
build_local_openai_responses_request_body(
body_json,
&mapped_model,
upstream_is_stream,
force_body_stream_field,
transport.provider.provider_type.as_str(),
provider_api_format,
if is_kiro_claude_cli {
None
} else {
transport.endpoint.body_rules.as_ref()
},
Some(input.auth_context.api_key_id.as_str()),
effective_headers,
enable_model_directives,
)
})
else {
mark_skipped_local_openai_responses_candidate_with_extra_data(
state,
input,
@@ -391,7 +429,9 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
.await;
}
let Some(upstream_url) = (if needs_bidirectional_conversion {
let Some(upstream_url) = (if is_grok && is_grok_text_provider_api_format(provider_api_format) {
Some(build_grok_upstream_url(transport, GROK_CHAT_PATH))
} else if needs_bidirectional_conversion {
build_cross_format_openai_responses_upstream_url(
parts,
transport,
@@ -428,48 +468,86 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
.as_ref()
.map(build_antigravity_static_identity_headers)
.unwrap_or_default();
let Some(resolved_headers) =
build_standard_provider_request_headers(StandardProviderRequestHeadersInput {
let resolved_headers = if is_grok && is_grok_text_provider_api_format(provider_api_format) {
let Some(headers) = build_grok_browser_headers(GrokHeaderInput {
transport,
provider_api_format,
same_format,
headers: effective_headers,
auth_header: &auth_header,
auth_value: &auth_value,
extra_headers: &extra_headers,
transport_profile: transport_profile.as_ref(),
request_headers: Some(effective_headers),
content_type: "application/json",
accept: "text/event-stream",
header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body,
original_request_body: body_json,
upstream_is_stream,
})
else {
mark_skipped_local_openai_responses_candidate_with_failure_diagnostic(
state,
input,
trace_id,
candidate,
candidate_index,
candidate_id,
"transport_header_rules_apply_failed",
CandidateFailureDiagnostic::header_rules_apply_failed(
spec_metadata.api_format,
}) else {
mark_skipped_local_openai_responses_candidate_with_failure_diagnostic(
state,
input,
trace_id,
candidate,
candidate_index,
candidate_id,
"transport_header_rules_apply_failed",
CandidateFailureDiagnostic::header_rules_apply_failed(
spec_metadata.api_format,
provider_api_format,
"grok_openai_responses_headers",
),
)
.await;
return None;
};
crate::ai_serving::transport::StandardProviderRequestHeaders {
headers,
auth_header: auth_header.clone(),
auth_value: auth_value.clone(),
}
} else {
let Some(resolved_headers) =
build_standard_provider_request_headers(StandardProviderRequestHeadersInput {
transport,
provider_api_format,
"openai_responses_headers",
),
)
.await;
return None;
same_format,
headers: effective_headers,
auth_header: &auth_header,
auth_value: &auth_value,
extra_headers: &extra_headers,
header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body,
original_request_body: body_json,
upstream_is_stream,
})
else {
mark_skipped_local_openai_responses_candidate_with_failure_diagnostic(
state,
input,
trace_id,
candidate,
candidate_index,
candidate_id,
"transport_header_rules_apply_failed",
CandidateFailureDiagnostic::header_rules_apply_failed(
spec_metadata.api_format,
provider_api_format,
"openai_responses_headers",
),
)
.await;
return None;
};
resolved_headers
};
let mut provider_request_headers = resolved_headers.headers;
apply_codex_openai_responses_special_headers(
&mut provider_request_headers,
&provider_request_body,
effective_headers,
transport.provider.provider_type.as_str(),
provider_api_format,
Some(trace_id),
transport.key.decrypted_auth_config.as_deref(),
);
if !is_grok {
apply_codex_openai_responses_special_headers(
&mut provider_request_headers,
&provider_request_body,
effective_headers,
transport.provider.provider_type.as_str(),
provider_api_format,
Some(trace_id),
transport.key.decrypted_auth_config.as_deref(),
);
}
let (execution_strategy, conversion_mode) =
ai_local_execution_contract_for_formats(spec_metadata.api_format, provider_api_format);
@@ -517,6 +595,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
},
upstream_is_stream,
transport: Arc::clone(transport),
transport_profile,
})
}
@@ -668,5 +747,6 @@ async fn build_kiro_openai_responses_payload_parts(
envelope_name: Some(KIRO_ENVELOPE_NAME),
upstream_is_stream,
transport: Arc::clone(transport),
transport_profile: None,
})
}

View File

@@ -61,6 +61,7 @@ pub(crate) use aether_ai_formats::api::{
maybe_build_standard_sync_finalize_product_from_normalized_payload, model_directive_base_model,
normalize_api_format_alias, normalize_claude_request_to_openai_chat_request,
normalize_gemini_request_to_openai_chat_request, normalize_openai_image_request,
normalize_openai_image_request_with_options,
normalize_openai_responses_request_to_openai_chat_request,
normalize_provider_private_report_context, normalize_provider_private_response_value,
normalize_standard_request_to_openai_chat_request, openai_image_operation_from_path,
@@ -97,10 +98,10 @@ pub(crate) use aether_ai_formats::api::{
LocalStandardSourceMode, LocalStandardSpec, LocalSyncReportParts, LocalVideoCreateFamily,
LocalVideoCreateSpec, NormalizedOpenAiImageRequest, OpenAIChatClientEmitter,
OpenAIChatProviderState, OpenAIResponsesClientEmitter, OpenAIResponsesProviderState,
OpenAiImageOperation, OpenAiImageRequestForGemini, OpenAiImageResponseFormat,
OpenAiImageStreamState, OpenAiImageSyncFinalizeProduct, ProviderAdaptationDescriptor,
ProviderAdaptationSurface, ProviderPrivateStreamNormalizer, RequestConversionKind,
StandardCrossFormatSyncProduct, StandardSyncFinalizeNormalizedProduct,
OpenAiImageNormalizeOptions, OpenAiImageOperation, OpenAiImageRequestForGemini,
OpenAiImageResponseFormat, OpenAiImageStreamState, OpenAiImageSyncFinalizeProduct,
ProviderAdaptationDescriptor, ProviderAdaptationSurface, ProviderPrivateStreamNormalizer,
RequestConversionKind, StandardCrossFormatSyncProduct, StandardSyncFinalizeNormalizedProduct,
StreamingStandardFormatMatrix, SyncChatResponseConversionKind, SyncCliResponseConversionKind,
SyncToStreamBridgeOutcome, ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME, CLAUDE_CHAT_STREAM_PLAN_KIND,
CLAUDE_CHAT_STREAM_SUCCESS_REPORT_KIND, CLAUDE_CHAT_SYNC_ERROR_REPORT_KIND,

View File

@@ -14,6 +14,10 @@ pub(crate) mod kiro {
pub(crate) use aether_provider_transport::kiro::*;
}
pub(crate) mod grok {
pub(crate) use aether_provider_transport::grok::*;
}
pub(crate) mod oauth_refresh {
pub(crate) use aether_provider_transport::oauth_refresh::*;
}
@@ -54,6 +58,7 @@ pub(crate) use aether_provider_transport::{
body_rules_are_locally_supported, body_rules_handle_path, body_rules_have_enabled_rules,
build_cross_format_openai_chat_upstream_url, build_cross_format_openai_responses_upstream_url,
build_gemini_files_headers, build_gemini_files_request_body, build_gemini_files_upstream_url,
build_grok_app_chat_body, build_grok_browser_headers, build_grok_upstream_url,
build_kiro_cross_format_upstream_url, build_local_openai_chat_upstream_url,
build_local_openai_responses_upstream_url, build_openai_image_headers,
build_openai_image_upstream_url, build_passthrough_headers, build_request_trace_proxy_value,
@@ -72,21 +77,23 @@ pub(crate) use aether_provider_transport::{
openai_image_transport_unsupported_reason, request_conversion_direct_auth,
request_conversion_enabled_for_transport, request_conversion_transport_supported,
request_conversion_transport_unsupported_reason, request_pair_allowed_for_transport,
resolve_gemini_files_auth, resolve_openai_image_auth, resolve_same_format_provider_direct_auth,
resolve_transport_execution_timeouts, resolve_transport_profile,
resolve_transport_proxy_snapshot, resolve_transport_proxy_snapshot_with_tunnel_affinity,
resolve_video_create_auth, same_format_provider_transport_supported,
same_format_provider_transport_unsupported_reason, should_skip_upstream_passthrough_header,
should_try_same_format_provider_oauth_auth, supports_local_gemini_transport_with_network,
resolve_gemini_files_auth, resolve_grok_session_auth, resolve_openai_image_auth,
resolve_same_format_provider_direct_auth, resolve_transport_execution_timeouts,
resolve_transport_profile, resolve_transport_proxy_snapshot,
resolve_transport_proxy_snapshot_with_tunnel_affinity, resolve_video_create_auth,
same_format_provider_transport_supported, same_format_provider_transport_unsupported_reason,
should_skip_upstream_passthrough_header, should_try_same_format_provider_oauth_auth,
supports_local_gemini_transport_with_network,
supports_local_generic_oauth_request_auth_resolution,
supports_local_oauth_request_auth_resolution, transport_proxy_is_locally_supported,
video_create_transport_unsupported_reason, CandidateTransportPolicyFacts,
GatewayProviderTransportSnapshot, GeminiFilesHeadersInput, GeminiFilesRequestBodyError,
GeminiFilesRequestBodyParts, LocalResolvedOAuthRequestAuth, ProviderOpenAiImageHeadersInput,
ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput, SameFormatProviderFamily,
SameFormatProviderHeadersInput, SameFormatProviderRequestBehavior,
GeminiFilesRequestBodyParts, GrokHeaderInput, LocalResolvedOAuthRequestAuth,
ProviderOpenAiImageHeadersInput, ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput,
SameFormatProviderFamily, SameFormatProviderHeadersInput, SameFormatProviderRequestBehavior,
SameFormatProviderRequestBehaviorParams, SameFormatProviderRequestBodyInput,
SameFormatProviderUpstreamUrlParams, StandardPlanFallbackAcceptPolicy,
StandardPlanFallbackHeadersInput, StandardProviderRequestHeaders,
StandardProviderRequestHeadersInput, TransportRequestUrlParams,
StandardProviderRequestHeadersInput, TransportRequestUrlParams, GROK_CHAT_PATH,
GROK_INTERNAL_HEADER, GROK_RATE_LIMITS_PATH,
};

View File

@@ -17,7 +17,6 @@ const AI_POST_ROUTE_PATTERNS: &[&str] = &[
"/v1/responses/compact",
"/v1/images/generations",
"/v1/images/edits",
"/v1/images/variations",
];
const AI_ANY_ROUTE_PATTERNS: &[&str] = &[

View File

@@ -115,7 +115,6 @@ pub(crate) const RUST_FRONTDOOR_OWNED_ROUTE_PATTERNS: &[&str] = &[
"/v1/rerank",
"/v1/images/generations",
"/v1/images/edits",
"/v1/images/variations",
"/v1/messages",
"/v1/messages/count_tokens",
"/v1/responses",

View File

@@ -55,7 +55,7 @@ pub(super) fn classify_ai_public_route(
} else if method == http::Method::POST
&& matches!(
normalized_path,
"/v1/images/generations" | "/v1/images/edits" | "/v1/images/variations"
"/v1/images/generations" | "/v1/images/edits"
)
{
Some(classified(

View File

@@ -80,6 +80,28 @@ fn classifies_openai_chat_and_responses_separately_from_embedding() {
assert_ne!(responses.route_kind.as_deref(), Some("embedding"));
}
#[test]
fn classifies_openai_image_generation_and_edit_but_not_variation() {
let headers = headers(&[("authorization", "Bearer sk-test")]);
for path in ["/v1/images/generations", "/v1/images/edits"] {
let uri: Uri = path.parse().expect("uri should parse");
let decision = classify_control_route(&http::Method::POST, &uri, &headers)
.expect("image route should classify");
assert_eq!(decision.route_family.as_deref(), Some("openai"));
assert_eq!(decision.route_kind.as_deref(), Some("image"));
assert_eq!(
decision.auth_endpoint_signature.as_deref(),
Some("openai:image")
);
assert!(decision.is_execution_runtime_candidate());
}
let variation_uri: Uri = "/v1/images/variations".parse().expect("uri should parse");
assert!(classify_control_route(&http::Method::POST, &variation_uri, &headers).is_none());
}
#[test]
fn classifies_models_list_as_claude_when_headers_match() {
let headers = headers(&[

View File

@@ -3,9 +3,10 @@ use std::io::Error as IoError;
use std::time::Instant;
use aether_contracts::{
ExecutionPlan, ExecutionResult, ExecutionTelemetry, RequestBody, ResponseBody, StreamFrame,
StreamFramePayload, StreamFrameType, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER,
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER,
ExecutionPlan, ExecutionResult, ExecutionTelemetry, RequestBody, ResolvedTransportProfile,
ResponseBody, StreamFrame, StreamFramePayload, StreamFrameType,
EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER,
TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_POOL_SCOPE_KEY,
};
use axum::body::Bytes;
use base64::Engine as _;
@@ -30,6 +31,7 @@ const CHATGPT_WEB_CLIENT_VERSION: &str = "prod-be885abbfcfe7b1f511e88b3003d9ee44
const CHATGPT_WEB_BUILD_NUMBER: &str = "5955942";
const CHATGPT_WEB_SEC_CH_UA: &str =
r#""Microsoft Edge";v="143", "Chromium";v="143", "Not A(Brand";v="24""#;
const CHATGPT_WEB_BROWSER_PROFILE: &str = "chrome143";
pub(crate) struct ChatGptWebImageStream {
pub(crate) frame_stream: BoxStream<'static, Result<Bytes, IoError>>,
@@ -921,7 +923,7 @@ async fn execute_subrequest(
provider_api_format: plan.provider_api_format.clone(),
model_name: plan.model_name.clone(),
proxy: plan.proxy.clone(),
transport_profile: plan.transport_profile.clone(),
transport_profile: chatgpt_web_image_transport_profile(plan),
timeouts: plan.timeouts.clone(),
};
DirectSyncExecutionRuntime::new()
@@ -929,6 +931,34 @@ async fn execute_subrequest(
.await
}
fn chatgpt_web_image_transport_profile(plan: &ExecutionPlan) -> Option<ResolvedTransportProfile> {
match plan.transport_profile.as_ref() {
Some(profile)
if profile
.backend
.trim()
.eq_ignore_ascii_case(TRANSPORT_BACKEND_BROWSER_WREQ) =>
{
Some(profile.clone())
}
_ => Some(default_chatgpt_web_image_transport_profile()),
}
}
fn default_chatgpt_web_image_transport_profile() -> ResolvedTransportProfile {
ResolvedTransportProfile {
profile_id: CHATGPT_WEB_BROWSER_PROFILE.to_string(),
backend: TRANSPORT_BACKEND_BROWSER_WREQ.to_string(),
http_mode: TRANSPORT_HTTP_MODE_AUTO.to_string(),
pool_scope: TRANSPORT_POOL_SCOPE_KEY.to_string(),
header_fingerprint: None,
extra: Some(json!({
"browser_profile": CHATGPT_WEB_BROWSER_PROFILE,
"source": "chatgpt_web_image_default",
})),
}
}
fn web_base_headers(fp: &WebFingerprint, token: &str, path: &str) -> BTreeMap<String, String> {
let mut headers = BTreeMap::from([
("user-agent".to_string(), fp.user_agent.to_string()),
@@ -1962,6 +1992,30 @@ mod tests {
}
}
#[test]
fn chatgpt_web_image_subrequests_default_to_browser_wreq_transport() {
let plan = sample_plan(
CHATGPT_WEB_DEFAULT_BASE_URL,
json!({"prompt": "draw a small test image"}),
false,
);
let profile = chatgpt_web_image_transport_profile(&plan).expect("transport profile");
assert_eq!(profile.backend, TRANSPORT_BACKEND_BROWSER_WREQ);
assert_eq!(profile.profile_id, CHATGPT_WEB_BROWSER_PROFILE);
assert_eq!(profile.http_mode, TRANSPORT_HTTP_MODE_AUTO);
assert_eq!(profile.pool_scope, TRANSPORT_POOL_SCOPE_KEY);
assert_eq!(
profile
.extra
.as_ref()
.and_then(|value| value.get("source"))
.and_then(Value::as_str),
Some("chatgpt_web_image_default")
);
}
async fn start_mock_chatgpt_web() -> (String, tokio::task::JoinHandle<()>) {
let app = Router::new().fallback(any(|request: Request| async move {
let path = request.uri().path().to_string();

File diff suppressed because it is too large Load Diff

View File

@@ -6,6 +6,7 @@ use serde_json::{Map, Value};
mod chatgpt_web_image;
mod constants;
mod fallback;
mod grok;
mod kiro_web_search;
pub(crate) mod ndjson;
mod oauth_retry;
@@ -54,6 +55,7 @@ pub(crate) use sync::{
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessBuild,
LocalVideoSyncSuccessOutcome,
};
pub(crate) use transport::execute_sync_plan_with_report_context as execute_execution_runtime_sync_plan_with_report_context;
pub(crate) use transport::{
execute_sync_plan as execute_execution_runtime_sync_plan, DirectSyncExecutionRuntime,
DirectUpstreamStreamExecution, ExecutionRuntimeTransportError,

View File

@@ -370,6 +370,8 @@ impl IntoResponse for ExecutionRuntimeAppError {
) => StatusCode::BAD_REQUEST,
ExecutionRuntimeServerError::Transport(
ExecutionRuntimeTransportError::ClientBuild(_)
| ExecutionRuntimeTransportError::BrowserClientBuild(_)
| ExecutionRuntimeTransportError::BrowserBody(_)
| ExecutionRuntimeTransportError::UpstreamRequest(_)
| ExecutionRuntimeTransportError::RelayError(_)
| ExecutionRuntimeTransportError::InvalidJson(_),

View File

@@ -60,6 +60,7 @@ use crate::constants::{CONTROL_CANDIDATE_ID_HEADER, CONTROL_REQUEST_ID_HEADER};
use crate::control::GatewayControlDecision;
use crate::execution_runtime::build_direct_execution_frame_stream;
use crate::execution_runtime::chatgpt_web_image::maybe_execute_chatgpt_web_image_stream;
use crate::execution_runtime::grok::maybe_execute_grok_stream;
use crate::execution_runtime::kiro_web_search::maybe_execute_kiro_web_search_stream;
use crate::execution_runtime::oauth_retry::refresh_oauth_plan_auth_for_retry;
#[cfg(test)]
@@ -525,6 +526,58 @@ pub(crate) async fn execute_execution_runtime_stream(
key_id.as_str(),
)
.await;
match maybe_execute_grok_stream(&plan, report_context.as_ref()).await {
Ok(Some(grok_stream)) => {
return execute_stream_from_frame_stream(
state,
plan,
trace_id,
decision,
plan_kind,
report_kind,
grok_stream.report_context.or(report_context),
candidate_started_unix_secs,
stream_started_at,
grok_stream.frame_stream,
provider_pool_in_flight_guard.take(),
)
.await;
}
Ok(None) => {}
Err(err) => {
info!(
event_name = "grok_execution_unavailable",
log_type = "ops",
trace_id = %trace_id,
request_id = %plan_request_id_for_log,
candidate_id = ?plan.candidate_id,
provider_name = provider_name.as_str(),
endpoint_id = %endpoint_id,
key_id = %key_id,
model_name = model_name.as_str(),
candidate_index = candidate_index.as_str(),
error = %err,
"gateway Grok stream execution unavailable"
);
let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status(
state,
&plan,
report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed,
status_code: None,
error_type: Some("grok_execution_unavailable".to_string()),
error_message: Some(format!("{err:?}")),
latency_ms: None,
started_at_unix_ms: Some(candidate_started_unix_secs),
finished_at_unix_ms: Some(terminal_unix_secs),
},
)
.await;
return Ok(None);
}
}
match maybe_execute_kiro_web_search_stream(state, &plan, report_context.as_ref()).await {
Ok(Some(kiro_web_search)) => {
return execute_stream_from_frame_stream(

View File

@@ -18,7 +18,9 @@ use crate::ai_serving::api::{
normalize_provider_private_report_context, StreamingStandardTerminalObserver,
};
use crate::execution_runtime::ndjson::encode_stream_frame_ndjson;
use crate::execution_runtime::transport::DirectUpstreamResponse;
use crate::execution_runtime::transport::{
format_wreq_upstream_request_error, DirectUpstreamResponse,
};
use crate::execution_runtime::DirectUpstreamStreamExecution;
use crate::GatewayError;
@@ -235,6 +237,62 @@ pub(crate) fn build_direct_execution_frame_stream(
}
}
}
DirectUpstreamResponse::BrowserWreq(response) => {
let mut bytes_stream = response.bytes_stream();
while let Some(item) = bytes_stream.next().await {
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
if !first_chunk_telemetry_emitted {
match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
first_chunk_telemetry_emitted = true;
}
upstream_bytes += chunk.len() as u64;
observe_stream_chunk(
&mut stream_terminal_observer,
&normalized_observer_context,
private_stream_normalizer.as_mut(),
&mut observer_buffered,
chunk.as_ref(),
);
match encode_data_frame(&chunk) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
}
Err(err) => {
let message = format_wreq_upstream_request_error(&err);
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
status_code,
upstream_bytes,
error = %message,
"upstream body stream read error"
);
match encode_error_frame(status_code, message) {
Ok(frame) => yield Ok(frame),
Err(encode_err) => {
yield Err(encode_err);
return;
}
}
break;
}
}
}
}
DirectUpstreamResponse::LocalTunnel(mut response) => loop {
match response.next_chunk().await {
Ok(Some(chunk)) => {
@@ -454,6 +512,35 @@ async fn buffer_non_sse_upstream_body(
}
}
}
DirectUpstreamResponse::BrowserWreq(response) => {
let mut bytes_stream = response.bytes_stream();
while let Some(item) = bytes_stream.next().await {
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
upstream_bytes += chunk.len() as u64;
body_bytes.extend_from_slice(&chunk);
}
Err(err) => {
let message = format_wreq_upstream_request_error(&err);
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
upstream_bytes,
error = %message,
"upstream body stream read error"
);
return Err(BufferedUpstreamBodyError {
message,
ttfb_ms,
upstream_bytes,
});
}
}
}
}
DirectUpstreamResponse::LocalTunnel(mut response) => loop {
match response.next_chunk().await {
Ok(Some(chunk)) => {

View File

@@ -36,14 +36,15 @@ use crate::api::response::{
use crate::clock::current_unix_ms as current_request_candidate_unix_ms;
use crate::control::GatewayControlDecision;
use crate::execution_runtime::chatgpt_web_image::maybe_execute_chatgpt_web_image_sync;
use crate::execution_runtime::grok::maybe_execute_grok_sync;
use crate::execution_runtime::oauth_retry::refresh_oauth_plan_auth_for_retry;
#[cfg(test)]
use crate::execution_runtime::remote_compat::post_sync_plan_to_remote_execution_runtime;
use crate::execution_runtime::submission::submit_local_core_error_or_sync_finalize;
use crate::execution_runtime::transport::{
build_request_body, collect_response_headers, decode_response_body_bytes,
response_body_is_json, send_request, DirectSyncExecutionRuntime,
ExecutionRuntimeTransportError,
format_upstream_request_error, format_wreq_upstream_request_error, response_body_is_json,
send_request, DirectHttpResponse, DirectSyncExecutionRuntime, ExecutionRuntimeTransportError,
};
use crate::execution_runtime::{
analyze_local_candidate_failover_sync, apply_endpoint_response_header_rules,
@@ -682,23 +683,46 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
.await
.map_err(SyncExecutionFailure::from_transport)?;
let ttfb_ms = started_at.elapsed().as_millis() as u64;
let status_code = response.status().as_u16();
let headers = collect_response_headers(response.headers());
let status_code = response.status_code();
let headers = response.headers();
progress.record_response_started(status_code, ttfb_ms).await;
let mut upstream_stream = response.bytes_stream();
let mut body_bytes = Vec::new();
while let Some(chunk) = upstream_stream.next().await {
let chunk = chunk.map_err(|err| {
SyncExecutionFailure::from_transport(ExecutionRuntimeTransportError::UpstreamRequest(
crate::execution_runtime::transport::format_upstream_request_error(&err),
))
})?;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
progress
.observe_chunk(&chunk, status_code, elapsed_ms)
.await;
body_bytes.extend_from_slice(&chunk);
match response {
DirectHttpResponse::Reqwest(response) => {
let mut upstream_stream = response.bytes_stream();
while let Some(chunk) = upstream_stream.next().await {
let chunk = chunk.map_err(|err| {
SyncExecutionFailure::from_transport(
ExecutionRuntimeTransportError::UpstreamRequest(
format_upstream_request_error(&err),
),
)
})?;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
progress
.observe_chunk(&chunk, status_code, elapsed_ms)
.await;
body_bytes.extend_from_slice(&chunk);
}
}
DirectHttpResponse::BrowserWreq(response) => {
let mut upstream_stream = response.bytes_stream();
while let Some(chunk) = upstream_stream.next().await {
let chunk = chunk.map_err(|err| {
SyncExecutionFailure::from_transport(
ExecutionRuntimeTransportError::UpstreamRequest(
format_wreq_upstream_request_error(&err),
),
)
})?;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
progress
.observe_chunk(&chunk, status_code, elapsed_ms)
.await;
body_bytes.extend_from_slice(&chunk);
}
}
}
let decoded_body_bytes =
@@ -1077,64 +1101,106 @@ async fn execute_execution_runtime_sync_impl(
.await;
#[cfg(not(test))]
let mut result = {
match maybe_execute_chatgpt_web_image_sync(state, &plan, report_context.as_ref()).await {
match maybe_execute_grok_sync(&plan, report_context.as_ref()).await {
Ok(Some(result)) => result,
Ok(None) => match execute_direct_sync_runtime_candidate(
state,
&plan,
report_context.as_ref(),
trace_id,
plan_kind,
plan_request_id_for_log.as_str(),
plan_candidate_id.as_deref(),
provider_name.as_str(),
endpoint_id.as_str(),
key_id.as_str(),
model_name.as_str(),
candidate_index.as_str(),
progress_snapshot.clone(),
)
.await
{
Ok(result) => result,
Err(err) => {
warn!(
event_name = "sync_execution_runtime_unavailable",
log_type = "ops",
trace_id = %trace_id,
request_id = %plan_request_id_for_log,
candidate_id = ?plan_candidate_id,
provider_name,
endpoint_id,
key_id,
model_name,
candidate_index = candidate_index.as_str(),
error_type = err.error_type,
error = %err.message,
"gateway in-process sync execution unavailable"
);
let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status(
Ok(None) => {
match maybe_execute_chatgpt_web_image_sync(state, &plan, report_context.as_ref())
.await
{
Ok(Some(result)) => result,
Ok(None) => match execute_direct_sync_runtime_candidate(
state,
&plan,
report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed,
status_code: err.status_code,
error_type: Some(err.error_type.to_string()),
error_message: Some(err.message),
latency_ms: err.latency_ms,
started_at_unix_ms: Some(candidate_started_unix_secs),
finished_at_unix_ms: Some(terminal_unix_secs),
},
trace_id,
plan_kind,
plan_request_id_for_log.as_str(),
plan_candidate_id.as_deref(),
provider_name.as_str(),
endpoint_id.as_str(),
key_id.as_str(),
model_name.as_str(),
candidate_index.as_str(),
progress_snapshot.clone(),
)
.await;
return Ok(None);
.await
{
Ok(result) => result,
Err(err) => {
warn!(
event_name = "sync_execution_runtime_unavailable",
log_type = "ops",
trace_id = %trace_id,
request_id = %plan_request_id_for_log,
candidate_id = ?plan_candidate_id,
provider_name,
endpoint_id,
key_id,
model_name,
candidate_index = candidate_index.as_str(),
error_type = err.error_type,
error = %err.message,
"gateway in-process sync execution unavailable"
);
let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status(
state,
&plan,
report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed,
status_code: err.status_code,
error_type: Some(err.error_type.to_string()),
error_message: Some(err.message),
latency_ms: err.latency_ms,
started_at_unix_ms: Some(candidate_started_unix_secs),
finished_at_unix_ms: Some(terminal_unix_secs),
},
)
.await;
return Ok(None);
}
},
Err(err) => {
warn!(
event_name = "chatgpt_web_image_execution_unavailable",
log_type = "ops",
trace_id = %trace_id,
request_id = %plan_request_id_for_log,
candidate_id = ?plan_candidate_id,
provider_name,
endpoint_id,
key_id,
model_name,
candidate_index = candidate_index.as_str(),
error = %err,
"gateway ChatGPT-Web image execution unavailable"
);
let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status(
state,
&plan,
report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed,
status_code: None,
error_type: Some(
"chatgpt_web_image_execution_unavailable".to_string(),
),
error_message: Some(err.to_string()),
latency_ms: None,
started_at_unix_ms: Some(candidate_started_unix_secs),
finished_at_unix_ms: Some(terminal_unix_secs),
},
)
.await;
return Ok(None);
}
}
},
}
Err(err) => {
warn!(
event_name = "chatgpt_web_image_execution_unavailable",
event_name = "grok_execution_unavailable",
log_type = "ops",
trace_id = %trace_id,
request_id = %plan_request_id_for_log,
@@ -1145,7 +1211,7 @@ async fn execute_execution_runtime_sync_impl(
model_name,
candidate_index = candidate_index.as_str(),
error = %err,
"gateway ChatGPT-Web image execution unavailable"
"gateway Grok execution unavailable"
);
let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status(
@@ -1155,7 +1221,7 @@ async fn execute_execution_runtime_sync_impl(
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed,
status_code: None,
error_type: Some("chatgpt_web_image_execution_unavailable".to_string()),
error_type: Some("grok_execution_unavailable".to_string()),
error_message: Some(err.to_string()),
latency_ms: None,
started_at_unix_ms: Some(candidate_started_unix_secs),
@@ -1212,30 +1278,72 @@ async fn execute_execution_runtime_sync_impl(
.trim()
.is_empty()
{
match maybe_execute_chatgpt_web_image_sync(state, &plan, report_context.as_ref()).await
{
match maybe_execute_grok_sync(&plan, report_context.as_ref()).await {
Ok(Some(result)) => result,
Ok(None) => match execute_direct_sync_runtime_candidate(
Ok(None) => match maybe_execute_chatgpt_web_image_sync(
state,
&plan,
report_context.as_ref(),
trace_id,
plan_kind,
plan_request_id_for_log.as_str(),
plan_candidate_id.as_deref(),
provider_name.as_str(),
endpoint_id.as_str(),
key_id.as_str(),
model_name.as_str(),
candidate_index.as_str(),
progress_snapshot.clone(),
)
.await
{
Ok(result) => result,
Ok(Some(result)) => result,
Ok(None) => match execute_direct_sync_runtime_candidate(
state,
&plan,
report_context.as_ref(),
trace_id,
plan_kind,
plan_request_id_for_log.as_str(),
plan_candidate_id.as_deref(),
provider_name.as_str(),
endpoint_id.as_str(),
key_id.as_str(),
model_name.as_str(),
candidate_index.as_str(),
progress_snapshot.clone(),
)
.await
{
Ok(result) => result,
Err(err) => {
warn!(
event_name = "sync_execution_runtime_unavailable",
log_type = "ops",
trace_id = %trace_id,
request_id = %plan_request_id_for_log,
candidate_id = ?plan_candidate_id,
provider_name,
endpoint_id,
key_id,
model_name,
candidate_index = candidate_index.as_str(),
error_type = err.error_type,
error = %err.message,
"gateway in-process sync execution unavailable"
);
let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status(
state,
&plan,
report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed,
status_code: err.status_code,
error_type: Some(err.error_type.to_string()),
error_message: Some(err.message),
latency_ms: err.latency_ms,
started_at_unix_ms: Some(candidate_started_unix_secs),
finished_at_unix_ms: Some(terminal_unix_secs),
},
)
.await;
return Ok(None);
}
},
Err(err) => {
warn!(
event_name = "sync_execution_runtime_unavailable",
event_name = "chatgpt_web_image_execution_unavailable",
log_type = "ops",
trace_id = %trace_id,
request_id = %plan_request_id_for_log,
@@ -1245,9 +1353,8 @@ async fn execute_execution_runtime_sync_impl(
key_id,
model_name,
candidate_index = candidate_index.as_str(),
error_type = err.error_type,
error = %err.message,
"gateway in-process sync execution unavailable"
error = %err,
"gateway ChatGPT-Web image execution unavailable"
);
let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status(
@@ -1256,10 +1363,12 @@ async fn execute_execution_runtime_sync_impl(
report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed,
status_code: err.status_code,
error_type: Some(err.error_type.to_string()),
error_message: Some(err.message),
latency_ms: err.latency_ms,
status_code: None,
error_type: Some(
"chatgpt_web_image_execution_unavailable".to_string(),
),
error_message: Some(err.to_string()),
latency_ms: None,
started_at_unix_ms: Some(candidate_started_unix_secs),
finished_at_unix_ms: Some(terminal_unix_secs),
},
@@ -1270,7 +1379,7 @@ async fn execute_execution_runtime_sync_impl(
},
Err(err) => {
warn!(
event_name = "chatgpt_web_image_execution_unavailable",
event_name = "grok_execution_unavailable",
log_type = "ops",
trace_id = %trace_id,
request_id = %plan_request_id_for_log,
@@ -1281,7 +1390,7 @@ async fn execute_execution_runtime_sync_impl(
model_name,
candidate_index = candidate_index.as_str(),
error = %err,
"gateway ChatGPT-Web image execution unavailable"
"gateway Grok execution unavailable"
);
let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status(
@@ -1291,7 +1400,7 @@ async fn execute_execution_runtime_sync_impl(
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed,
status_code: None,
error_type: Some("chatgpt_web_image_execution_unavailable".to_string()),
error_type: Some("grok_execution_unavailable".to_string()),
error_message: Some(err.to_string()),
latency_ms: None,
started_at_unix_ms: Some(candidate_started_unix_secs),

View File

@@ -8,7 +8,8 @@ use aether_contracts::{
ExecutionPlan, ExecutionResult, ExecutionTelemetry, ProxySnapshot, ResolvedTransportProfile,
ResponseBody, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER,
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER,
TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_HTTP_MODE_HTTP1_ONLY,
TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_BACKEND_REQWEST_RUSTLS,
TRANSPORT_HTTP_MODE_HTTP1_ONLY,
};
use aether_data::repository::proxy_nodes::ProxyNodeTrafficMutation;
use aether_http::{apply_http_client_config, HttpClientConfig};
@@ -81,6 +82,52 @@ pub(crate) fn format_upstream_request_error(err: &reqwest::Error) -> String {
detail
}
pub(crate) fn format_wreq_upstream_request_error(err: &wreq::Error) -> String {
let mut kinds = Vec::new();
if err.is_connect() {
kinds.push("connect");
}
if err.is_timeout() {
kinds.push("timeout");
}
if err.is_redirect() {
kinds.push("redirect");
}
if err.is_body() {
kinds.push("body");
}
if err.is_decode() {
kinds.push("decode");
}
if err.is_request() {
kinds.push("request");
}
let mut detail = err.to_string();
let mut source = err.source();
while let Some(cause) = source {
let cause_text = cause.to_string();
if !cause_text.is_empty() && !detail.contains(&cause_text) {
detail.push_str(": ");
detail.push_str(&cause_text);
}
source = cause.source();
}
if let Some(uri) = err.uri() {
detail.push_str(" [uri=");
detail.push_str(&uri.to_string());
detail.push(']');
}
if !kinds.is_empty() {
detail.push_str(" [kind=");
detail.push_str(&kinds.join(","));
detail.push(']');
}
detail
}
#[derive(Debug, Error)]
pub(crate) enum ExecutionRuntimeTransportError {
#[error("stream execution is not supported for this plan")]
@@ -107,6 +154,10 @@ pub(crate) enum ExecutionRuntimeTransportError {
BodyEncode(serde_json::Error),
#[error("failed to build HTTP client: {0}")]
ClientBuild(reqwest::Error),
#[error("failed to build browser impersonation HTTP client: {0}")]
BrowserClientBuild(wreq::Error),
#[error("browser impersonation response body failed: {0}")]
BrowserBody(String),
#[error("failed to execute upstream request: {0}")]
UpstreamRequest(String),
#[error("hub relay request failed: {0}")]
@@ -136,7 +187,7 @@ struct RelayRequestMeta {
pub(crate) struct DirectSyncExecutionRuntime;
#[derive(Debug, Clone, Copy, Default)]
struct ExecutionTransportControls {
pub(crate) struct ExecutionTransportControls {
follow_redirects: Option<bool>,
http1_only: bool,
accept_invalid_certs: bool,
@@ -144,6 +195,7 @@ struct ExecutionTransportControls {
pub(crate) enum DirectUpstreamResponse {
Reqwest(reqwest::Response),
BrowserWreq(wreq::Response),
LocalTunnel(tunnel::DirectRelayResponse),
}
@@ -172,11 +224,9 @@ impl DirectSyncExecutionRuntime {
let started_at = Instant::now();
let response = send_request(plan, body_bytes).await?;
let ttfb_ms = started_at.elapsed().as_millis() as u64;
let status_code = response.status().as_u16();
let headers = collect_response_headers(response.headers());
let body_bytes = response.bytes().await.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err))
})?;
let status_code = response.status_code();
let headers = response.headers();
let body_bytes = response.bytes().await?;
let decoded_body_bytes = decode_response_body_bytes(&headers, &body_bytes)
.unwrap_or_else(|| body_bytes.to_vec());
let elapsed_ms = started_at.elapsed().as_millis() as u64;
@@ -230,8 +280,8 @@ impl DirectSyncExecutionRuntime {
let started_at = Instant::now();
let response = send_request(plan, body_bytes).await?;
let status_code = response.status().as_u16();
let headers = collect_response_headers(response.headers());
let status_code = response.status_code();
let headers = response.headers();
let stream_summary_report_context = build_stream_summary_report_context(plan);
@@ -242,7 +292,7 @@ impl DirectSyncExecutionRuntime {
headers,
provider_api_format: plan.provider_api_format.clone(),
stream_summary_report_context,
response: DirectUpstreamResponse::Reqwest(response),
response: response.into_direct_upstream_response(),
started_at,
})
}
@@ -252,6 +302,15 @@ pub(crate) async fn execute_sync_plan(
state: &AppState,
trace_id: Option<&str>,
plan: &ExecutionPlan,
) -> Result<ExecutionResult, GatewayError> {
execute_sync_plan_with_report_context(state, trace_id, plan, None).await
}
pub(crate) async fn execute_sync_plan_with_report_context(
state: &AppState,
trace_id: Option<&str>,
plan: &ExecutionPlan,
report_context: Option<&serde_json::Value>,
) -> Result<ExecutionResult, GatewayError> {
#[cfg(test)]
{
@@ -275,6 +334,18 @@ pub(crate) async fn execute_sync_plan(
.map_err(|err| GatewayError::Internal(err.to_string()));
}
match super::grok::maybe_execute_grok_sync(plan, report_context).await {
Ok(Some(result)) => {
record_manual_proxy_request_outcome(state, plan, result.status_code).await;
return Ok(result);
}
Ok(None) => {}
Err(err) => {
record_manual_proxy_request_failure(state, plan).await;
return Err(GatewayError::Internal(err.to_string()));
}
}
let _ = trace_id;
match DirectSyncExecutionRuntime::new().execute_sync(plan).await {
Ok(result) => {
@@ -554,7 +625,7 @@ fn build_direct_tunnel_request_meta(
pub(crate) async fn send_request(
plan: &ExecutionPlan,
body_bytes: Vec<u8>,
) -> Result<reqwest::Response, ExecutionRuntimeTransportError> {
) -> Result<DirectHttpResponse, ExecutionRuntimeTransportError> {
if let Some(detail) = gateway_frontdoor_self_loop_guard_error(plan.url.as_str()) {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(detail));
}
@@ -572,6 +643,18 @@ pub(crate) async fn send_request(
.and_then(|timeouts| timeouts.total_ms)
.map(Duration::from_millis);
if transport_profile_uses_browser_wreq(plan.transport_profile.as_ref()) {
return send_via_browser_wreq_transport(
plan,
method,
headers,
body_bytes,
total_timeout,
transport_controls,
)
.await;
}
if let Some(node_id) = resolve_tunnel_node_id(plan.proxy.as_ref()) {
return send_via_tunnel_relay(
plan,
@@ -582,7 +665,8 @@ pub(crate) async fn send_request(
total_timeout,
transport_controls,
)
.await;
.await
.map(DirectHttpResponse::Reqwest);
}
let client = build_client(
@@ -596,9 +680,95 @@ pub(crate) async fn send_request(
if let Some(timeout) = total_timeout {
request = request.timeout(timeout);
}
request.send().await.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err))
})
request
.send()
.await
.map(DirectHttpResponse::Reqwest)
.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err))
})
}
pub(crate) enum DirectHttpResponse {
Reqwest(reqwest::Response),
BrowserWreq(wreq::Response),
}
impl DirectHttpResponse {
pub(crate) fn status_code(&self) -> u16 {
match self {
DirectHttpResponse::Reqwest(response) => response.status().as_u16(),
DirectHttpResponse::BrowserWreq(response) => response.status().as_u16(),
}
}
pub(crate) fn headers(&self) -> BTreeMap<String, String> {
match self {
DirectHttpResponse::Reqwest(response) => collect_response_headers(response.headers()),
DirectHttpResponse::BrowserWreq(response) => {
collect_response_headers(response.headers())
}
}
}
pub(crate) async fn bytes(self) -> Result<Bytes, ExecutionRuntimeTransportError> {
match self {
DirectHttpResponse::Reqwest(response) => response.bytes().await.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err))
}),
DirectHttpResponse::BrowserWreq(response) => response.bytes().await.map_err(|err| {
ExecutionRuntimeTransportError::BrowserBody(format_wreq_upstream_request_error(
&err,
))
}),
}
}
fn into_direct_upstream_response(self) -> DirectUpstreamResponse {
match self {
DirectHttpResponse::Reqwest(response) => DirectUpstreamResponse::Reqwest(response),
DirectHttpResponse::BrowserWreq(response) => {
DirectUpstreamResponse::BrowserWreq(response)
}
}
}
}
async fn send_via_browser_wreq_transport(
plan: &ExecutionPlan,
method: reqwest::Method,
headers: HeaderMap,
body_bytes: Vec<u8>,
total_timeout: Option<Duration>,
transport_controls: ExecutionTransportControls,
) -> Result<DirectHttpResponse, ExecutionRuntimeTransportError> {
let profile = plan.transport_profile.as_ref().ok_or_else(|| {
ExecutionRuntimeTransportError::UnsupportedTransportProfile(String::new())
})?;
let client = build_browser_wreq_client(
plan.timeouts.as_ref(),
plan.proxy.as_ref(),
profile,
transport_controls,
)?;
let method = wreq::Method::from_bytes(method.as_str().as_bytes())
.map_err(ExecutionRuntimeTransportError::InvalidMethod)?;
let mut request = client
.request(method, plan.url.as_str())
.headers(headers)
.body(body_bytes);
if let Some(timeout) = total_timeout {
request = request.timeout(timeout);
}
request
.send()
.await
.map(DirectHttpResponse::BrowserWreq)
.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format_wreq_upstream_request_error(
&err,
))
})
}
async fn send_via_tunnel_relay(
@@ -905,6 +1075,96 @@ fn build_client(
.map_err(ExecutionRuntimeTransportError::ClientBuild)
}
pub(crate) fn build_browser_wreq_client(
timeouts: Option<&aether_contracts::ExecutionTimeouts>,
proxy: Option<&ProxySnapshot>,
transport_profile: &ResolvedTransportProfile,
transport_controls: ExecutionTransportControls,
) -> Result<wreq::Client, ExecutionRuntimeTransportError> {
let emulation = browser_wreq_emulation_from_profile(transport_profile)?;
let mut builder = wreq::Client::builder().emulation(emulation);
if transport_controls.follow_redirects == Some(true) {
builder = builder.redirect(wreq::redirect::Policy::limited(10));
}
if transport_controls.http1_only || transport_profile_http1_only(Some(transport_profile)) {
builder = builder.http1_only();
}
if transport_controls.accept_invalid_certs {
builder = builder.cert_verification(false).verify_hostname(false);
}
if let Some(connect_ms) = timeouts.and_then(|timeouts| timeouts.connect_ms) {
builder = builder.connect_timeout(Duration::from_millis(connect_ms));
}
if let Some(total_ms) = timeouts.and_then(|timeouts| timeouts.total_ms) {
builder = builder.timeout(Duration::from_millis(total_ms));
}
if let Some(read_ms) = timeouts.and_then(|timeouts| timeouts.read_ms) {
builder = builder.read_timeout(Duration::from_millis(read_ms));
}
if let Some(proxy_url) = resolve_proxy_url(proxy)? {
let proxy = wreq::Proxy::all(proxy_url.as_str())
.map_err(ExecutionRuntimeTransportError::BrowserClientBuild)?;
builder = builder.proxy(proxy);
}
builder
.build()
.map_err(ExecutionRuntimeTransportError::BrowserClientBuild)
}
fn browser_wreq_emulation_from_profile(
profile: &ResolvedTransportProfile,
) -> Result<wreq_util::Emulation, ExecutionRuntimeTransportError> {
match normalize_browser_profile_name(browser_transport_profile_name(profile)).as_str() {
"chrome100" => Ok(wreq_util::Emulation::Chrome100),
"chrome101" => Ok(wreq_util::Emulation::Chrome101),
"chrome104" => Ok(wreq_util::Emulation::Chrome104),
"chrome105" => Ok(wreq_util::Emulation::Chrome105),
"chrome106" => Ok(wreq_util::Emulation::Chrome106),
"chrome107" => Ok(wreq_util::Emulation::Chrome107),
"chrome108" => Ok(wreq_util::Emulation::Chrome108),
"chrome109" => Ok(wreq_util::Emulation::Chrome109),
"chrome110" => Ok(wreq_util::Emulation::Chrome110),
"chrome114" => Ok(wreq_util::Emulation::Chrome114),
"chrome116" => Ok(wreq_util::Emulation::Chrome116),
"chrome117" => Ok(wreq_util::Emulation::Chrome117),
"chrome118" => Ok(wreq_util::Emulation::Chrome118),
"chrome119" => Ok(wreq_util::Emulation::Chrome119),
"chrome120" => Ok(wreq_util::Emulation::Chrome120),
"chrome123" => Ok(wreq_util::Emulation::Chrome123),
"chrome124" => Ok(wreq_util::Emulation::Chrome124),
"chrome126" => Ok(wreq_util::Emulation::Chrome126),
"chrome127" => Ok(wreq_util::Emulation::Chrome127),
"chrome128" => Ok(wreq_util::Emulation::Chrome128),
"chrome129" => Ok(wreq_util::Emulation::Chrome129),
"chrome130" => Ok(wreq_util::Emulation::Chrome130),
"chrome131" => Ok(wreq_util::Emulation::Chrome131),
"chrome132" => Ok(wreq_util::Emulation::Chrome132),
"chrome133" => Ok(wreq_util::Emulation::Chrome133),
"chrome134" => Ok(wreq_util::Emulation::Chrome134),
"chrome135" => Ok(wreq_util::Emulation::Chrome135),
"chrome136" => Ok(wreq_util::Emulation::Chrome136),
"chrome137" => Ok(wreq_util::Emulation::Chrome137),
"chrome138" => Ok(wreq_util::Emulation::Chrome138),
"chrome139" => Ok(wreq_util::Emulation::Chrome139),
"chrome140" => Ok(wreq_util::Emulation::Chrome140),
"chrome141" => Ok(wreq_util::Emulation::Chrome141),
"chrome142" => Ok(wreq_util::Emulation::Chrome142),
"chrome143" => Ok(wreq_util::Emulation::Chrome143),
"chrome144" => Ok(wreq_util::Emulation::Chrome144),
"chrome145" => Ok(wreq_util::Emulation::Chrome145),
other => Err(ExecutionRuntimeTransportError::UnsupportedTransportProfile(
format!("browser_wreq:{other}"),
)),
}
}
fn normalize_browser_profile_name(value: String) -> String {
value
.trim()
.to_ascii_lowercase()
.replace(['_', '-', ' '], "")
}
fn validate_reqwest_transport_profile(
transport_profile: Option<&ResolvedTransportProfile>,
) -> Result<(), ExecutionRuntimeTransportError> {
@@ -923,6 +1183,56 @@ fn validate_reqwest_transport_profile(
))
}
fn transport_profile_uses_browser_wreq(
transport_profile: Option<&ResolvedTransportProfile>,
) -> bool {
transport_profile
.map(|profile| {
profile
.backend
.trim()
.eq_ignore_ascii_case(TRANSPORT_BACKEND_BROWSER_WREQ)
})
.unwrap_or(false)
}
fn browser_transport_profile_name(profile: &ResolvedTransportProfile) -> String {
profile
.extra
.as_ref()
.and_then(|value| {
value
.get("browser_profile")
.or_else(|| value.get("impersonate"))
.and_then(Value::as_str)
})
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.or_else(|| {
profile
.profile_id
.trim()
.is_empty()
.then_some("chrome136".to_string())
.or_else(|| Some(profile.profile_id.trim().to_string()))
})
.unwrap_or_else(|| "chrome136".to_string())
}
fn insert_browser_control_header(
headers: &mut HeaderMap,
name: &'static str,
value: &str,
) -> Result<(), ExecutionRuntimeTransportError> {
headers.insert(
HeaderName::from_static(name),
HeaderValue::from_str(value)
.map_err(|_| ExecutionRuntimeTransportError::InvalidHeaderValue(name.to_string()))?,
);
Ok(())
}
fn transport_profile_http1_only(transport_profile: Option<&ResolvedTransportProfile>) -> bool {
transport_profile
.map(|profile| {
@@ -991,7 +1301,7 @@ fn resolve_proxy_url(
Ok(None)
}
fn build_request_headers(
pub(crate) fn build_request_headers(
headers: &BTreeMap<String, String>,
content_encoding: Option<&str>,
allow_passthrough_content_encoding: bool,
@@ -1180,22 +1490,22 @@ mod tests {
use aether_contracts::{
ExecutionPlan, ExecutionTimeouts, ProxySnapshot, RequestBody, ResolvedTransportProfile,
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER,
TRANSPORT_BACKEND_REQWEST_RUSTLS,
TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_BACKEND_REQWEST_RUSTLS,
};
use aether_data::repository::proxy_nodes::{
InMemoryProxyNodeRepository, ProxyNodeReadRepository, StoredProxyNode,
};
use axum::body::Bytes;
use axum::body::{Body, Bytes};
use axum::extract::ws::Message;
use axum::extract::Path;
use axum::http::HeaderMap as AxumHeaderMap;
use axum::routing::post;
use axum::routing::{any, post};
use axum::{Json, Router};
use serde_json::json;
use tokio::sync::watch;
use super::{
build_client, build_request_headers, execute_sync_plan,
build_browser_wreq_client, build_client, build_request_headers, execute_sync_plan,
record_manual_proxy_request_failure, record_manual_proxy_request_outcome,
record_manual_proxy_request_success, record_manual_proxy_stream_error,
resolve_execution_transport_controls, DirectSyncExecutionRuntime,
@@ -1431,6 +1741,228 @@ mod tests {
);
}
#[tokio::test]
async fn direct_sync_execution_runtime_routes_browser_wreq_transport_in_process() {
async fn browser_upstream(headers: AxumHeaderMap, body: Bytes) -> axum::response::Response {
assert_eq!(
headers
.get("content-type")
.and_then(|value| value.to_str().ok()),
Some("application/json")
);
assert!(
headers
.get(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER)
.is_none(),
"internal execution control headers must not leak upstream"
);
assert_eq!(body.as_ref(), br#"{"modelName":"auto"}"#);
axum::response::Response::builder()
.status(http::StatusCode::ACCEPTED)
.header("content-type", "application/json")
.body(Body::from(
json!({
"ok": true,
"via": "browser_wreq"
})
.to_string(),
))
.expect("response should build")
}
let listener = crate::test_support::bind_loopback_listener()
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should resolve");
let app = Router::new().route("/request", any(browser_upstream));
let server = tokio::spawn(async move {
axum::serve(listener, app)
.await
.expect("test server should run");
});
let plan = ExecutionPlan {
request_id: "req-browser-wreq".into(),
candidate_id: None,
provider_name: Some("grok".into()),
provider_id: "provider-1".into(),
endpoint_id: "endpoint-1".into(),
key_id: "key-1".into(),
method: "POST".into(),
url: format!("http://{addr}/request"),
headers: BTreeMap::from([
("content-type".into(), "application/json".into()),
(
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.into(),
"true".into(),
),
]),
content_type: Some("application/json".into()),
content_encoding: None,
body: RequestBody::from_json(json!({"modelName":"auto"})),
stream: false,
client_api_format: "openai:responses".into(),
provider_api_format: "grok:rate_limits".into(),
model_name: Some("grok-quota".into()),
proxy: None,
transport_profile: Some(ResolvedTransportProfile {
profile_id: "chrome136".into(),
backend: TRANSPORT_BACKEND_BROWSER_WREQ.into(),
http_mode: "auto".into(),
pool_scope: "key".into(),
header_fingerprint: None,
extra: Some(json!({
"browser_profile": "chrome136"
})),
}),
timeouts: Some(ExecutionTimeouts {
total_ms: Some(5_000),
..ExecutionTimeouts::default()
}),
};
let result = DirectSyncExecutionRuntime::new()
.execute_sync(&plan)
.await
.expect("browser wreq transport plan should execute in-process");
server.abort();
assert_eq!(result.status_code, http::StatusCode::ACCEPTED.as_u16());
assert_eq!(
result
.body
.and_then(|body| body.json_body)
.and_then(|body| body.get("via").cloned()),
Some(json!("browser_wreq"))
);
}
#[test]
fn browser_wreq_transport_rejects_unknown_profile() {
let profile = ResolvedTransportProfile {
profile_id: "firefox999".into(),
backend: TRANSPORT_BACKEND_BROWSER_WREQ.into(),
http_mode: "auto".into(),
pool_scope: "key".into(),
header_fingerprint: None,
extra: None,
};
let error = match build_browser_wreq_client(
None,
None,
&profile,
ExecutionTransportControls::default(),
) {
Ok(_) => panic!("unknown browser profile should fail loudly"),
Err(error) => error,
};
assert!(matches!(
error,
ExecutionRuntimeTransportError::UnsupportedTransportProfile(backend)
if backend == "browser_wreq:firefox999"
));
}
#[tokio::test]
async fn execute_sync_plan_routes_grok_marker_through_grok_runtime() {
let listener = crate::test_support::bind_loopback_listener()
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should resolve");
let app = Router::new().route(
"/rest/app-chat/conversations/new",
post(|body: Bytes| async move {
let body_json: serde_json::Value =
serde_json::from_slice(&body).expect("request body should be json");
if body_json.get("message").and_then(serde_json::Value::as_str)
!= Some("[user]: hello")
{
return (
axum::http::StatusCode::BAD_REQUEST,
Json(json!({
"error": {
"message": "expected grok app-chat message",
"body": body_json,
}
})),
);
}
(
axum::http::StatusCode::OK,
Json(json!({
"result": {
"response": {
"token": "pong",
"messageTag": "final"
}
}
})),
)
}),
);
let server = tokio::spawn(async move {
axum::serve(listener, app)
.await
.expect("test server should run");
});
let plan = ExecutionPlan {
request_id: "req-grok-runtime".into(),
candidate_id: Some("cand-grok".into()),
provider_name: Some("grok".into()),
provider_id: "provider-grok".into(),
endpoint_id: "endpoint-grok".into(),
key_id: "key-grok".into(),
method: "POST".into(),
url: format!("http://{addr}/rest/app-chat/conversations/new"),
headers: BTreeMap::from([
("content-type".into(), "application/json".into()),
(
aether_provider_transport::GROK_INTERNAL_HEADER.into(),
"1".into(),
),
]),
content_type: Some("application/json".into()),
content_encoding: None,
body: RequestBody::from_json(json!({
"model": "grok-4.20-0309-non-reasoning",
"messages": [{"role": "user", "content": "hello"}],
})),
stream: true,
client_api_format: "openai:chat".into(),
provider_api_format: "openai:chat".into(),
model_name: Some("grok-4.20-0309-non-reasoning".into()),
proxy: None,
transport_profile: None,
timeouts: Some(ExecutionTimeouts {
connect_ms: Some(5_000),
total_ms: Some(5_000),
..ExecutionTimeouts::default()
}),
};
let report_context = json!({"mapped_model": "grok-4.20-fast"});
let result = super::super::grok::maybe_execute_grok_sync(&plan, Some(&report_context))
.await
.expect("grok runtime plan should execute")
.expect("grok runtime should handle marked plan");
server.abort();
assert_eq!(result.status_code, http::StatusCode::OK.as_u16());
assert_eq!(
result
.body
.and_then(|body| body.json_body)
.and_then(|body| body["choices"][0]["message"]["content"]
.as_str()
.map(str::to_string)),
Some("pong".to_string())
);
}
#[tokio::test]
async fn execute_sync_plan_records_manual_proxy_success() {
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![

View File

@@ -38,7 +38,7 @@ pub(super) async fn build_admin_create_api_key_install_session_response(
Err(_) => {
return Ok(build_admin_api_keys_bad_request_response(
"请求数据验证失败",
))
));
}
};

View File

@@ -249,9 +249,10 @@ async fn admin_monitoring_resilience_status_returns_local_payload() {
let recommendations = payload["recommendations"]
.as_array()
.expect("recommendations should be array");
assert!(recommendations.iter().any(|item| item
.as_str()
.is_some_and(|value| value.contains("prod-key"))));
assert!(recommendations.iter().any(|item| {
item.as_str()
.is_some_and(|value| value.contains("prod-key"))
}));
assert!(payload["timestamp"].as_str().is_some());
}

View File

@@ -32,7 +32,7 @@ pub(super) async fn maybe_handle(
.into_response(),
));
};
let Some(_provider) = state
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
@@ -129,7 +129,9 @@ pub(super) async fn maybe_handle(
Json(serde_json::Value::Array(
created
.iter()
.map(|model| build_admin_provider_model_response(model, now_unix_secs))
.map(|model| {
build_admin_provider_model_response(&provider, model, now_unix_secs)
})
.collect(),
))
.into_response(),

View File

@@ -31,7 +31,7 @@ pub(super) async fn maybe_handle(
.into_response(),
));
};
let Some(_provider) = state
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
@@ -90,8 +90,12 @@ pub(super) async fn maybe_handle(
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
Json(build_admin_provider_model_response(&created, now_unix_secs))
.into_response()
Json(build_admin_provider_model_response(
&provider,
&created,
now_unix_secs,
))
.into_response()
}
None => (
http::StatusCode::INTERNAL_SERVER_ERROR,

View File

@@ -1,9 +1,13 @@
use crate::handlers::admin::provider::shared::model_test_capabilities::{
admin_provider_model_supports_image_generation, admin_provider_model_test_capabilities_payload,
};
use crate::handlers::admin::request::AdminAppState;
use crate::GatewayError;
use aether_admin::provider::models as admin_provider_models_pure;
use aether_data_contracts::repository::global_models::{
AdminProviderModelListQuery, StoredAdminProviderModel,
};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider;
use std::time::{SystemTime, UNIX_EPOCH};
pub(super) fn admin_provider_model_effective_input_price(
@@ -26,10 +30,32 @@ pub(super) fn admin_provider_model_effective_capability(
}
pub(super) fn build_admin_provider_model_response(
provider: &StoredProviderCatalogProvider,
model: &StoredAdminProviderModel,
now_unix_secs: u64,
) -> serde_json::Value {
admin_provider_models_pure::build_admin_provider_model_response(model, now_unix_secs)
let mut payload =
admin_provider_models_pure::build_admin_provider_model_response(model, now_unix_secs);
let fallback_supports_image_generation = payload
.get("effective_supports_image_generation")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false);
let supports_image_generation = admin_provider_model_supports_image_generation(
&provider.provider_type,
&model.provider_model_name,
fallback_supports_image_generation,
);
if let Some(object) = payload.as_object_mut() {
object.insert(
"model_test_capabilities".to_string(),
admin_provider_model_test_capabilities_payload(
&provider.provider_type,
&model.provider_model_name,
supports_image_generation,
),
);
}
payload
}
pub(super) async fn build_admin_provider_models_payload(
@@ -48,9 +74,10 @@ pub(super) async fn build_admin_provider_models_payload(
.ok()?
.into_iter()
.next()?;
let provider_id = provider.id.clone();
let mut models = state
.list_admin_provider_models(&AdminProviderModelListQuery {
provider_id: provider.id,
provider_id,
is_active,
offset: skip,
limit,
@@ -70,7 +97,7 @@ pub(super) async fn build_admin_provider_models_payload(
Some(serde_json::Value::Array(
models
.iter()
.map(|model| build_admin_provider_model_response(model, now_unix_secs))
.map(|model| build_admin_provider_model_response(&provider, model, now_unix_secs))
.collect(),
))
}
@@ -80,9 +107,15 @@ pub(super) async fn build_admin_provider_model_payload(
provider_id: &str,
model_id: &str,
) -> Option<serde_json::Value> {
if !state.has_global_model_data_reader() {
if !state.has_provider_catalog_data_reader() || !state.has_global_model_data_reader() {
return None;
}
let provider = state
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
.await
.ok()?
.into_iter()
.next()?;
let model = state
.get_admin_provider_model(provider_id, model_id)
.await
@@ -92,7 +125,11 @@ pub(super) async fn build_admin_provider_model_payload(
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
Some(build_admin_provider_model_response(&model, now_unix_secs))
Some(build_admin_provider_model_response(
&provider,
&model,
now_unix_secs,
))
}
pub(super) async fn admin_provider_model_name_exists(

View File

@@ -33,6 +33,20 @@ pub(super) async fn maybe_handle(
.into_response(),
));
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
));
};
let Some(existing) = state
.get_admin_provider_model(&provider_id, &model_id)
.await?
@@ -110,8 +124,12 @@ pub(super) async fn maybe_handle(
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
Json(build_admin_provider_model_response(&updated, now_unix_secs))
.into_response()
Json(build_admin_provider_model_response(
&provider,
&updated,
now_unix_secs,
))
.into_response()
}
None => (
http::StatusCode::NOT_FOUND,

View File

@@ -1,3 +1,4 @@
use super::super::helpers::admin_provider_oauth_key_name_from_auth_config;
use super::super::token_import::{
build_provider_access_token_import_auth_config, provider_type_supports_access_token_import,
};
@@ -24,13 +25,11 @@ use crate::handlers::admin::provider::oauth::runtime::{
use crate::handlers::admin::provider::oauth::state::{
admin_provider_oauth_template, exchange_admin_provider_oauth_refresh_token,
};
use crate::handlers::admin::provider::shared::support::ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL;
use crate::handlers::admin::request::{AdminAppState, AdminProviderOAuthTemplate};
use crate::GatewayError;
use aether_admin::provider::oauth::parse_admin_provider_oauth_kiro_batch_import_entries;
use aether_contracts::ProxySnapshot;
use serde_json::{json, Map, Value};
use std::time::{SystemTime, UNIX_EPOCH};
struct AdminProviderOAuthResolvedBatchImport {
access_token: String,
@@ -45,7 +44,7 @@ pub(super) fn estimate_admin_provider_oauth_batch_import_total(
if provider_type.eq_ignore_ascii_case("kiro") {
parse_admin_provider_oauth_kiro_batch_import_entries(raw_credentials).len()
} else {
parse_admin_provider_oauth_batch_import_entries(raw_credentials).len()
parse_admin_provider_oauth_batch_import_entries(provider_type, raw_credentials).len()
}
}
@@ -67,7 +66,8 @@ pub(super) async fn execute_admin_provider_oauth_batch_import_for_provider_type(
)
.await
} else {
let entries = parse_admin_provider_oauth_batch_import_entries(raw_credentials);
let entries =
parse_admin_provider_oauth_batch_import_entries(provider_type, raw_credentials);
execute_admin_provider_oauth_batch_import(
state,
provider_id,
@@ -82,7 +82,7 @@ pub(super) async fn execute_admin_provider_oauth_batch_import_for_provider_type(
async fn resolve_admin_provider_oauth_batch_import_tokens(
state: &AdminAppState<'_>,
template: AdminProviderOAuthTemplate,
template: Option<AdminProviderOAuthTemplate>,
provider_type: &str,
entry: &AdminProviderOAuthBatchImportEntry,
request_proxy: Option<ProxySnapshot>,
@@ -99,6 +99,29 @@ async fn resolve_admin_provider_oauth_batch_import_tokens(
.filter(|value| !value.is_empty());
if let Some(refresh_token) = refresh_token {
let Some(template) = template else {
if provider_type_supports_access_token_import(provider_type) {
if let Some(access_token) = access_token {
let (auth_config, expires_at) = build_provider_access_token_import_auth_config(
provider_type,
access_token,
Some(refresh_token),
entry.expires_at,
Some("Provider 不支持 Refresh Token 交换,已回退为 Session Token 导入"),
);
return Ok(AdminProviderOAuthResolvedBatchImport {
access_token: access_token.to_string(),
auth_config,
expires_at,
});
}
}
return Err(
"该 Provider 不支持 Refresh Token 导入,请提供 sso_token 或 access_token"
.to_string(),
);
};
let token_payload = match exchange_admin_provider_oauth_refresh_token(
state,
template,
@@ -152,7 +175,7 @@ async fn resolve_admin_provider_oauth_batch_import_tokens(
if let Some(access_token) = access_token {
if !provider_type_supports_access_token_import(provider_type) {
return Err("Access Token 导入仅支持 Codex / ChatGPT Web Provider".to_string());
return Err("Access Token 导入仅支持 Codex / ChatGPT Web / Grok Provider".to_string());
}
let (auth_config, expires_at) = build_provider_access_token_import_auth_config(
provider_type,
@@ -204,25 +227,7 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
});
};
let Some(template) = admin_provider_oauth_template(provider_type) else {
return Ok(AdminProviderOAuthBatchImportOutcome {
total: entries.len(),
success: 0,
failed: entries.len(),
results: entries
.iter()
.enumerate()
.map(|(index, _)| {
json!({
"index": index,
"status": "error",
"error": ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL,
"replaced": false,
})
})
.collect(),
});
};
let template = admin_provider_oauth_template(provider_type);
let endpoint_resolution =
resolve_provider_oauth_runtime_endpoints(state, &provider, provider_type).await?;
@@ -340,24 +345,11 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
}
}
} else {
let key_name = auth_config
.get("email")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|email| format!("{provider_type}_{email}"))
.unwrap_or_else(|| {
format!(
"{}_{}_{}",
provider_type,
SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0),
index
)
});
let key_name = admin_provider_oauth_key_name_from_auth_config(
provider_type,
&auth_config,
Some(index),
);
match create_provider_oauth_catalog_key(
state,
provider_id,

View File

@@ -8,7 +8,7 @@ use super::parse::{
};
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
use crate::handlers::admin::provider::oauth::state::{
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
build_admin_provider_oauth_backend_unavailable_response,
is_fixed_provider_type_for_provider_oauth,
};
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_batch_import_provider_id;
@@ -60,10 +60,6 @@ pub(in super::super) async fn handle_admin_provider_oauth_batch_import(
"该 Provider 不是固定类型,无法使用 provider-oauth",
));
}
if provider_type != "kiro" && admin_provider_oauth_template(&provider_type).is_none() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let total = estimate_admin_provider_oauth_batch_import_total(
&provider_type,
payload.credentials.as_str(),

View File

@@ -1,4 +1,4 @@
use super::super::token_import::{import_tokens_from_raw_token, normalize_single_import_tokens};
use super::super::token_import::{import_tokens_from_raw_token, normalize_provider_import_tokens};
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
use crate::handlers::admin::provider::oauth::state::{current_unix_secs, json_u64_value};
use axum::{
@@ -25,8 +25,15 @@ pub(super) struct AdminProviderOAuthBatchImportEntry {
pub account_id: Option<String>,
pub account_user_id: Option<String>,
pub plan_type: Option<String>,
pub pool_tier: Option<String>,
pub user_id: Option<String>,
pub email: Option<String>,
pub account_name: Option<String>,
pub sso_rw_token: Option<String>,
pub cf_cookies: Option<String>,
pub cf_clearance: Option<String>,
pub user_agent: Option<String>,
pub browser_profile: Option<String>,
}
#[derive(Debug, Clone)]
@@ -67,16 +74,72 @@ fn coerce_admin_provider_oauth_import_str(value: Option<&serde_json::Value>) ->
.map(ToOwned::to_owned)
}
fn grok_cookie_value(raw: &str, name: &str) -> Option<String> {
raw.trim()
.strip_prefix("Cookie:")
.unwrap_or_else(|| raw.trim())
.split(';')
.filter_map(|segment| segment.trim().split_once('='))
.find_map(|(cookie_name, cookie_value)| {
cookie_name
.trim()
.eq_ignore_ascii_case(name)
.then(|| cookie_value.trim())
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
}
fn grok_cookie_profile(raw: &str) -> Option<String> {
let raw = raw
.trim()
.strip_prefix("Cookie:")
.unwrap_or_else(|| raw.trim());
let parts = raw
.split(';')
.filter_map(|segment| {
let (cookie_name, cookie_value) = segment.trim().split_once('=')?;
let cookie_name = cookie_name.trim();
let cookie_value = cookie_value.trim();
if cookie_name.is_empty()
|| cookie_value.is_empty()
|| cookie_name.eq_ignore_ascii_case("sso")
|| cookie_name.eq_ignore_ascii_case("sso-rw")
{
return None;
}
Some(format!("{cookie_name}={cookie_value}"))
})
.collect::<Vec<_>>();
(!parts.is_empty()).then(|| parts.join("; "))
}
fn grok_cookie_session_token(provider_type: &str, raw: &str) -> Option<String> {
provider_type
.trim()
.eq_ignore_ascii_case("grok")
.then(|| grok_cookie_value(raw, "sso"))
.flatten()
}
fn extract_admin_provider_oauth_batch_import_entry(
provider_type: &str,
item: &serde_json::Value,
) -> Option<AdminProviderOAuthBatchImportEntry> {
match item {
serde_json::Value::String(value) => {
let refresh_token = value.trim();
if refresh_token.is_empty() {
let raw_token = value.trim();
if raw_token.is_empty() {
None
} else {
let (refresh_token, access_token) = import_tokens_from_raw_token(refresh_token);
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) = normalize_provider_import_tokens(
provider_type,
refresh_token.as_deref(),
access_token.as_deref(),
);
Some(AdminProviderOAuthBatchImportEntry {
refresh_token,
access_token,
@@ -84,8 +147,15 @@ fn extract_admin_provider_oauth_batch_import_entry(
account_id: None,
account_user_id: None,
plan_type: None,
user_id: None,
pool_tier: None,
user_id: grok_cookie_value(raw_token, "x-userid"),
email: None,
account_name: None,
sso_rw_token: grok_cookie_value(raw_token, "sso-rw"),
cf_cookies: grok_cookie_profile(raw_token),
cf_clearance: grok_cookie_value(raw_token, "cf_clearance"),
user_agent: None,
browser_profile: None,
})
}
}
@@ -100,8 +170,34 @@ fn extract_admin_provider_oauth_batch_import_entry(
.get("access_token")
.or_else(|| object.get("accessToken")),
);
let (refresh_token, access_token) =
normalize_single_import_tokens(refresh_token.as_deref(), access_token.as_deref());
let grok_token_alias = if provider_type.trim().eq_ignore_ascii_case("grok") {
object.get("token")
} else {
None
};
let grok_cookie = if provider_type.trim().eq_ignore_ascii_case("grok") {
coerce_admin_provider_oauth_import_str(
object.get("cookie").or_else(|| object.get("cookieHeader")),
)
} else {
None
};
let session_token = coerce_admin_provider_oauth_import_str(
object
.get("sso_token")
.or_else(|| object.get("ssoToken"))
.or(grok_token_alias),
)
.or_else(|| {
grok_cookie
.as_deref()
.and_then(|cookie| grok_cookie_value(cookie, "sso"))
});
let (refresh_token, access_token) = normalize_provider_import_tokens(
provider_type,
refresh_token.as_deref(),
access_token.as_deref().or(session_token.as_deref()),
);
if refresh_token.is_none() && access_token.is_none() {
return None;
}
@@ -129,14 +225,65 @@ fn extract_admin_provider_oauth_batch_import_entry(
.or_else(|| object.get("chatgptPlanType")),
)
.map(|value| value.to_ascii_lowercase());
let pool_tier = coerce_admin_provider_oauth_import_str(
object
.get("pool_tier")
.or_else(|| object.get("poolTier"))
.or_else(|| object.get("tier")),
)
.map(|value| value.to_ascii_lowercase());
let user_id = coerce_admin_provider_oauth_import_str(
object
.get("user_id")
.or_else(|| object.get("userId"))
.or_else(|| object.get("chatgpt_user_id"))
.or_else(|| object.get("chatgptUserId")),
);
)
.or_else(|| {
grok_cookie
.as_deref()
.and_then(|cookie| grok_cookie_value(cookie, "x-userid"))
});
let email = coerce_admin_provider_oauth_import_str(object.get("email"));
let account_name = coerce_admin_provider_oauth_import_str(
object
.get("account_name")
.or_else(|| object.get("accountName")),
);
let sso_rw_token = coerce_admin_provider_oauth_import_str(
object
.get("sso_rw_token")
.or_else(|| object.get("ssoRwToken")),
)
.or_else(|| {
grok_cookie
.as_deref()
.and_then(|cookie| grok_cookie_value(cookie, "sso-rw"))
});
let cf_clearance = coerce_admin_provider_oauth_import_str(
object
.get("cf_clearance")
.or_else(|| object.get("cfClearance")),
)
.or_else(|| {
grok_cookie
.as_deref()
.and_then(|cookie| grok_cookie_value(cookie, "cf_clearance"))
});
let cf_cookies = coerce_admin_provider_oauth_import_str(
object.get("cf_cookies").or_else(|| object.get("cfCookies")),
)
.or_else(|| grok_cookie.as_deref().and_then(grok_cookie_profile));
let user_agent = coerce_admin_provider_oauth_import_str(
object.get("user_agent").or_else(|| object.get("userAgent")),
);
let browser_profile = coerce_admin_provider_oauth_import_str(
object
.get("browser_profile")
.or_else(|| object.get("browserProfile"))
.or_else(|| object.get("browser"))
.or_else(|| object.get("impersonate")),
);
Some(AdminProviderOAuthBatchImportEntry {
refresh_token,
access_token,
@@ -144,8 +291,15 @@ fn extract_admin_provider_oauth_batch_import_entry(
account_id,
account_user_id,
plan_type,
pool_tier,
user_id,
email,
account_name,
sso_rw_token,
cf_cookies,
cf_clearance,
user_agent,
browser_profile,
})
}
_ => None,
@@ -153,6 +307,7 @@ fn extract_admin_provider_oauth_batch_import_entry(
}
pub(super) fn parse_admin_provider_oauth_batch_import_entries(
provider_type: &str,
raw_credentials: &str,
) -> Vec<AdminProviderOAuthBatchImportEntry> {
let raw = raw_credentials.trim();
@@ -165,7 +320,9 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries(
{
return items
.iter()
.filter_map(extract_admin_provider_oauth_batch_import_entry)
.filter_map(|item| {
extract_admin_provider_oauth_batch_import_entry(provider_type, item)
})
.collect();
}
}
@@ -174,7 +331,7 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries(
if let Ok(value @ serde_json::Value::Object(_)) =
serde_json::from_str::<serde_json::Value>(raw)
{
return extract_admin_provider_oauth_batch_import_entry(&value)
return extract_admin_provider_oauth_batch_import_entry(provider_type, &value)
.into_iter()
.collect();
}
@@ -183,18 +340,19 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries(
raw.lines()
.map(str::trim)
.filter(|line| !line.is_empty() && !line.starts_with('#'))
.map(|token| {
let (refresh_token, access_token) = import_tokens_from_raw_token(token);
AdminProviderOAuthBatchImportEntry {
refresh_token,
access_token,
expires_at: None,
account_id: None,
account_user_id: None,
plan_type: None,
user_id: None,
email: None,
.filter_map(|line| {
if line.starts_with('{') {
return serde_json::from_str::<serde_json::Value>(line)
.ok()
.and_then(|value| {
extract_admin_provider_oauth_batch_import_entry(provider_type, &value)
});
}
extract_admin_provider_oauth_batch_import_entry(
provider_type,
&serde_json::Value::String(line.to_string()),
)
})
.collect()
}
@@ -204,10 +362,8 @@ pub(super) fn apply_admin_provider_oauth_batch_import_hints(
entry: &AdminProviderOAuthBatchImportEntry,
auth_config: &mut serde_json::Map<String, serde_json::Value>,
) {
if !matches!(
provider_type.trim().to_ascii_lowercase().as_str(),
"codex" | "chatgpt_web"
) {
let provider_type = provider_type.trim().to_ascii_lowercase();
if !matches!(provider_type.as_str(), "codex" | "chatgpt_web" | "grok") {
return;
}
if let Some(account_id) = entry.account_id.as_ref() {
@@ -225,6 +381,11 @@ pub(super) fn apply_admin_provider_oauth_batch_import_hints(
.entry("plan_type".to_string())
.or_insert_with(|| json!(plan_type));
}
if let Some(pool_tier) = entry.pool_tier.as_ref() {
auth_config
.entry("pool_tier".to_string())
.or_insert_with(|| json!(pool_tier));
}
if let Some(user_id) = entry.user_id.as_ref() {
auth_config
.entry("user_id".to_string())
@@ -235,6 +396,36 @@ pub(super) fn apply_admin_provider_oauth_batch_import_hints(
.entry("email".to_string())
.or_insert_with(|| json!(email));
}
if let Some(account_name) = entry.account_name.as_ref() {
auth_config
.entry("account_name".to_string())
.or_insert_with(|| json!(account_name));
}
if let Some(sso_rw_token) = entry.sso_rw_token.as_ref() {
auth_config
.entry("sso_rw_token".to_string())
.or_insert_with(|| json!(sso_rw_token));
}
if let Some(cf_cookies) = entry.cf_cookies.as_ref() {
auth_config
.entry("cf_cookies".to_string())
.or_insert_with(|| json!(cf_cookies));
}
if let Some(cf_clearance) = entry.cf_clearance.as_ref() {
auth_config
.entry("cf_clearance".to_string())
.or_insert_with(|| json!(cf_clearance));
}
if let Some(user_agent) = entry.user_agent.as_ref() {
auth_config
.entry("user_agent".to_string())
.or_insert_with(|| json!(user_agent));
}
if let Some(browser_profile) = entry.browser_profile.as_ref() {
auth_config
.entry("browser_profile".to_string())
.or_insert_with(|| json!(browser_profile));
}
}
pub(super) async fn extract_admin_provider_oauth_batch_error_detail(
@@ -337,6 +528,7 @@ mod tests {
#[test]
fn parses_access_token_only_entry() {
let entries = parse_admin_provider_oauth_batch_import_entries(
"codex",
r#"[{"accessToken":"at_1","expiresAt":2100000000,"accountId":"acc-1","email":"u@example.com"}]"#,
);
@@ -356,10 +548,89 @@ mod tests {
"exp": 2_000_000_000u64,
}));
let entries = parse_admin_provider_oauth_batch_import_entries(&token);
let entries = parse_admin_provider_oauth_batch_import_entries("codex", &token);
assert_eq!(entries.len(), 1);
assert_eq!(entries[0].refresh_token, None);
assert_eq!(entries[0].access_token.as_deref(), Some(token.as_str()));
}
#[test]
fn parses_grok_jsonl_session_entries() {
let entries = parse_admin_provider_oauth_batch_import_entries(
"grok",
r#"{"sso_token":"sso-1","cf_clearance":"cf-1","pool_tier":"heavy","email":"grok@example.com","browser_profile":"chrome136"}"#,
);
assert_eq!(entries.len(), 1);
assert_eq!(entries[0].refresh_token, None);
assert_eq!(entries[0].access_token.as_deref(), Some("sso-1"));
assert_eq!(entries[0].cf_clearance.as_deref(), Some("cf-1"));
assert_eq!(entries[0].pool_tier.as_deref(), Some("heavy"));
assert_eq!(entries[0].email.as_deref(), Some("grok@example.com"));
assert_eq!(entries[0].browser_profile.as_deref(), Some("chrome136"));
}
#[test]
fn parses_grok_token_alias_with_account_traits() {
let entries = parse_admin_provider_oauth_batch_import_entries(
"grok",
r#"[{"token":"sso-1","planType":"super","tier":"heavy","accountName":"Grok Heavy"}]"#,
);
assert_eq!(entries.len(), 1);
assert_eq!(entries[0].refresh_token, None);
assert_eq!(entries[0].access_token.as_deref(), Some("sso-1"));
assert_eq!(entries[0].plan_type.as_deref(), Some("super"));
assert_eq!(entries[0].pool_tier.as_deref(), Some("heavy"));
assert_eq!(entries[0].account_name.as_deref(), Some("Grok Heavy"));
}
#[test]
fn parses_grok_plain_line_as_session_token() {
let entries = parse_admin_provider_oauth_batch_import_entries("grok", "opaque-sso-token");
assert_eq!(entries.len(), 1);
assert_eq!(entries[0].refresh_token, None);
assert_eq!(entries[0].access_token.as_deref(), Some("opaque-sso-token"));
}
#[test]
fn parses_grok_cookie_line_as_session_metadata() {
let entries = parse_admin_provider_oauth_batch_import_entries(
"grok",
"i18nextLng=zh; cf_clearance=cf-1; sso-rw=rw-1; sso=sso-1; x-userid=user-1",
);
assert_eq!(entries.len(), 1);
assert_eq!(entries[0].refresh_token, None);
assert_eq!(entries[0].access_token.as_deref(), Some("sso-1"));
assert_eq!(entries[0].sso_rw_token.as_deref(), Some("rw-1"));
assert_eq!(
entries[0].cf_cookies.as_deref(),
Some("i18nextLng=zh; cf_clearance=cf-1; x-userid=user-1")
);
assert_eq!(entries[0].cf_clearance.as_deref(), Some("cf-1"));
assert_eq!(entries[0].user_id.as_deref(), Some("user-1"));
}
#[test]
fn parses_grok_cookie_object_as_session_metadata() {
let entries = parse_admin_provider_oauth_batch_import_entries(
"grok",
r#"[{"cookie":"cf_clearance=cf-1; sso-rw=rw-1; sso=sso-1; x-userid=user-1","tier":"heavy"}]"#,
);
assert_eq!(entries.len(), 1);
assert_eq!(entries[0].refresh_token, None);
assert_eq!(entries[0].access_token.as_deref(), Some("sso-1"));
assert_eq!(entries[0].sso_rw_token.as_deref(), Some("rw-1"));
assert_eq!(
entries[0].cf_cookies.as_deref(),
Some("cf_clearance=cf-1; x-userid=user-1")
);
assert_eq!(entries[0].cf_clearance.as_deref(), Some("cf-1"));
assert_eq!(entries[0].user_id.as_deref(), Some("user-1"));
assert_eq!(entries[0].pool_tier.as_deref(), Some("heavy"));
}
}

View File

@@ -10,7 +10,7 @@ use super::progress::{
};
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
use crate::handlers::admin::provider::oauth::state::{
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
build_admin_provider_oauth_backend_unavailable_response,
is_fixed_provider_type_for_provider_oauth,
};
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_batch_import_task_provider_id;
@@ -124,10 +124,6 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
"该 Provider 不是固定类型,无法使用 provider-oauth",
));
}
if provider_type != "kiro" && admin_provider_oauth_template(&provider_type).is_none() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let total = estimate_admin_provider_oauth_batch_import_total(
&provider_type,
payload.credentials.as_str(),

View File

@@ -3,6 +3,8 @@ use axum::{
body::Body,
response::{IntoResponse, Response},
};
use serde_json::{Map, Value};
use std::time::{SystemTime, UNIX_EPOCH};
pub(super) fn attach_admin_provider_oauth_audit_response(
response: Response<Body>,
@@ -19,3 +21,79 @@ pub(super) fn attach_admin_provider_oauth_audit_response(
};
attach_admin_audit_response(response, event_name, action, target_type, &target_id)
}
pub(super) fn admin_provider_oauth_key_name_from_auth_config(
provider_type: &str,
auth_config: &Map<String, Value>,
batch_index: Option<usize>,
) -> String {
let provider_type = provider_type.trim();
if let Some(email) = trimmed_auth_config_string(auth_config, "email") {
return format!("{provider_type}_{email}");
}
if provider_type.eq_ignore_ascii_case("grok") {
if let Some(user_id) = trimmed_auth_config_string(auth_config, "user_id") {
return format!("grok_{user_id}");
}
}
let timestamp = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
match batch_index {
Some(index) => format!("{provider_type}_{timestamp}_{index}"),
None => format!("账号_{timestamp}"),
}
}
fn trimmed_auth_config_string(auth_config: &Map<String, Value>, key: &str) -> Option<String> {
auth_config
.get(key)
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::{json, Map};
#[test]
fn grok_default_key_name_uses_full_user_id() {
let mut auth_config = Map::new();
auth_config.insert(
"user_id".to_string(),
json!("1619039a-0191-4e0a-a490-8f4ad21262c9"),
);
assert_eq!(
admin_provider_oauth_key_name_from_auth_config("grok", &auth_config, None),
"grok_1619039a-0191-4e0a-a490-8f4ad21262c9"
);
}
#[test]
fn default_key_name_prefers_email_over_grok_user_id() {
let mut auth_config = Map::new();
auth_config.insert("email".to_string(), json!("grok@example.com"));
auth_config.insert("user_id".to_string(), json!("user-1"));
assert_eq!(
admin_provider_oauth_key_name_from_auth_config("grok", &auth_config, None),
"grok_grok@example.com"
);
}
#[test]
fn batch_default_key_name_keeps_existing_timestamp_shape() {
let auth_config = Map::new();
let name = admin_provider_oauth_key_name_from_auth_config("codex", &auth_config, Some(3));
assert!(name.starts_with("codex_"));
assert!(name.ends_with("_3"));
}
}

View File

@@ -14,8 +14,9 @@ use super::super::state::{
exchange_admin_provider_oauth_refresh_token, is_fixed_provider_type_for_provider_oauth,
json_u64_value,
};
use super::helpers::admin_provider_oauth_key_name_from_auth_config;
use super::token_import::{
build_provider_access_token_import_auth_config, normalize_single_import_tokens,
build_provider_access_token_import_auth_config, normalize_provider_import_tokens,
provider_type_supports_access_token_import,
};
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_import_provider_id;
@@ -31,7 +32,6 @@ use axum::{
Json,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
struct AdminProviderOAuthSingleImportTokens {
access_token: String,
@@ -72,7 +72,8 @@ fn apply_single_import_hints(
payload: &serde_json::Map<String, serde_json::Value>,
auth_config: &mut serde_json::Map<String, serde_json::Value>,
) {
if !provider_type_supports_access_token_import(provider_type) {
let provider_type = provider_type.trim().to_ascii_lowercase();
if !matches!(provider_type.as_str(), "codex" | "chatgpt_web" | "grok") {
return;
}
@@ -110,19 +111,40 @@ fn apply_single_import_hints(
&["user_id", "userId", "chatgpt_user_id", "chatgptUserId"][..],
),
("account_name", &["account_name", "accountName"][..]),
("sso_rw_token", &["sso_rw_token", "ssoRwToken"][..]),
(
"cf_cookies",
&["cf_cookies", "cfCookies", "cookie", "cookieHeader"][..],
),
("cf_clearance", &["cf_clearance", "cfClearance"][..]),
("user_agent", &["user_agent", "userAgent"][..]),
(
"browser_profile",
&[
"browser_profile",
"browserProfile",
"browser",
"impersonate",
][..],
),
("pool_tier", &["pool_tier", "poolTier", "tier"][..]),
] {
let Some(value) = import_payload_string_any(payload, keys) else {
continue;
};
auth_config
.entry(target.to_string())
.or_insert_with(|| json!(value));
auth_config.entry(target.to_string()).or_insert_with(|| {
if target == "plan_type" || target == "pool_tier" {
json!(value.to_ascii_lowercase())
} else {
json!(value)
}
});
}
}
async fn resolve_admin_provider_oauth_single_import_tokens(
state: &AdminAppState<'_>,
template: AdminProviderOAuthTemplate,
template: Option<AdminProviderOAuthTemplate>,
provider_type: &str,
refresh_token: Option<&str>,
access_token: Option<&str>,
@@ -133,6 +155,32 @@ async fn resolve_admin_provider_oauth_single_import_tokens(
.map(str::trim)
.filter(|value| !value.is_empty())
{
let Some(template) = template else {
if provider_type_supports_access_token_import(provider_type) {
if let Some(access_token) = access_token
.map(str::trim)
.filter(|value| !value.is_empty())
{
let (auth_config, expires_at) = build_provider_access_token_import_auth_config(
provider_type,
access_token,
Some(refresh_token),
imported_expires_at,
Some("Provider 不支持 Refresh Token 交换,已回退为 Session Token 导入"),
);
return Ok(AdminProviderOAuthSingleImportTokens {
access_token: access_token.to_string(),
auth_config,
expires_at,
});
}
}
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不支持 Refresh Token 导入,请提供 sso_token 或 access_token",
));
};
let token_payload = match state
.exchange_admin_provider_oauth_refresh_token(
template,
@@ -200,7 +248,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 导入仅支持 Codex / ChatGPT Web Provider",
"Access Token 导入仅支持 Codex / ChatGPT Web / Grok Provider",
));
}
@@ -248,18 +296,11 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
}
};
let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken");
let access_token_input = import_payload_string(&raw_payload, "access_token", "accessToken");
let imported_expires_at = import_payload_u64(&raw_payload, "expires_at", "expiresAt");
let (refresh_token_input, access_token_input) = normalize_single_import_tokens(
refresh_token_input.as_deref(),
access_token_input.as_deref(),
let access_token_input = import_payload_string_any(
&raw_payload,
&["access_token", "accessToken", "sso_token", "ssoToken"],
);
if 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 不能为空",
));
}
let imported_expires_at = import_payload_u64(&raw_payload, "expires_at", "expiresAt");
let name = raw_payload
.get("name")
.and_then(serde_json::Value::as_str)
@@ -285,6 +326,17 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
let (refresh_token_input, access_token_input) = normalize_provider_import_tokens(
&provider_type,
refresh_token_input.as_deref(),
access_token_input.as_deref(),
);
if 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 !is_fixed_provider_type_for_provider_oauth(&provider_type) {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
@@ -297,9 +349,10 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
"Kiro 不支持单条 Refresh Token 导入,请使用批量导入或设备授权。",
));
}
let Some(template) = admin_provider_oauth_template(&provider_type) else {
let template = admin_provider_oauth_template(&provider_type);
if template.is_none() && !provider_type_supports_access_token_import(&provider_type) {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
};
}
let endpoint_resolution =
resolve_provider_oauth_runtime_endpoints(state, &provider, &provider_type).await?;
let endpoints = endpoint_resolution.endpoints;
@@ -380,25 +433,9 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
}
}
} else {
let name = name
.or_else(|| {
auth_config
.get("email")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
.unwrap_or_else(|| {
format!(
"账号_{}",
SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0)
)
});
let name = name.unwrap_or_else(|| {
admin_provider_oauth_key_name_from_auth_config(&provider_type, &auth_config, None)
});
match state
.create_provider_oauth_catalog_key(
&provider_id,

View File

@@ -81,6 +81,28 @@ pub(super) fn normalize_single_import_tokens(
(refresh_token, access_token)
}
pub(super) fn normalize_provider_import_tokens(
provider_type: &str,
refresh_token: Option<&str>,
access_token: Option<&str>,
) -> (Option<String>, Option<String>) {
let provider_type = provider_type.trim().to_ascii_lowercase();
let refresh_token = refresh_token
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let access_token = access_token
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
if provider_type == "grok" {
return (None, access_token.or(refresh_token));
}
normalize_single_import_tokens(refresh_token.as_deref(), access_token.as_deref())
}
pub(super) fn import_tokens_from_raw_token(token: &str) -> (Option<String>, Option<String>) {
if looks_like_access_token(token) {
(None, Some(token.trim().to_string()))
@@ -98,7 +120,7 @@ pub(super) fn decode_access_token_expires_at(access_token: &str) -> Option<u64>
pub(super) fn provider_type_supports_access_token_import(provider_type: &str) -> bool {
matches!(
provider_type.trim().to_ascii_lowercase().as_str(),
"codex" | "chatgpt_web"
"codex" | "chatgpt_web" | "grok"
)
}
@@ -123,6 +145,11 @@ pub(super) fn build_provider_access_token_import_auth_config(
auth_config.insert("refresh_token".to_string(), json!(refresh_token));
}
if provider_type.trim().eq_ignore_ascii_case("grok") {
auth_config.insert("sso_token".to_string(), json!(access_token));
auth_config.insert("auth_method".to_string(), json!("sso_token"));
}
auth_config.insert(
"access_token_import_temporary".to_string(),
json!(refresh_token.is_none()),
@@ -149,7 +176,7 @@ pub(super) fn build_provider_access_token_import_auth_config(
mod tests {
use super::{
build_provider_access_token_import_auth_config, decode_access_token_expires_at,
looks_like_access_token, normalize_single_import_tokens,
looks_like_access_token, normalize_provider_import_tokens, normalize_single_import_tokens,
};
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use serde_json::json;
@@ -250,4 +277,34 @@ mod tests {
Some(&json!(true))
);
}
#[test]
fn normalize_grok_import_treats_opaque_session_as_access_token() {
let (refresh_token, access_token) =
normalize_provider_import_tokens("grok", Some("sso_session_token"), None);
assert!(refresh_token.is_none());
assert_eq!(access_token.as_deref(), Some("sso_session_token"));
}
#[test]
fn builds_grok_auth_config_from_session_token() {
let (auth_config, expires_at) = build_provider_access_token_import_auth_config(
"grok",
"sso_session_token",
None,
Some(2_200_000_000),
None,
);
assert_eq!(expires_at, Some(2_200_000_000));
assert_eq!(
auth_config.get("sso_token"),
Some(&json!("sso_session_token"))
);
assert_eq!(auth_config.get("auth_method"), Some(&json!("sso_token")));
assert_eq!(
auth_config.get("expires_at"),
Some(&json!(2_200_000_000u64))
);
}
}

View File

@@ -8,8 +8,10 @@ use crate::GatewayError;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
};
use aether_provider_transport::provider_types::provider_type_is_fixed;
use serde_json::json;
use aether_provider_transport::{
grok_browser_transport_fingerprint_from_auth_config, provider_types::provider_type_is_fixed,
};
use serde_json::{json, Map, Value};
use std::time::{SystemTime, UNIX_EPOCH};
use uuid::Uuid;
@@ -92,6 +94,16 @@ pub(crate) fn build_provider_oauth_auth_config_from_token_payload(
(auth_config, access_token, refresh_token, expires_at)
}
fn grok_oauth_catalog_key_fingerprint(
provider_type: &str,
auth_config: &Map<String, Value>,
) -> Option<Value> {
if !provider_type.trim().eq_ignore_ascii_case("grok") {
return None;
}
grok_browser_transport_fingerprint_from_auth_config(auth_config)
}
pub(crate) async fn create_provider_oauth_catalog_key(
state: &AdminAppState<'_>,
provider_id: &str,
@@ -136,7 +148,7 @@ pub(crate) async fn create_provider_oauth_catalog_key(
None,
expires_at_unix_secs,
proxy,
None,
grok_oauth_catalog_key_fingerprint(provider_type, auth_config),
)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
record.internal_priority = 50;
@@ -193,6 +205,9 @@ pub(crate) async fn update_existing_provider_oauth_catalog_key(
updated.expires_at_unix_secs = expires_at_unix_secs;
updated.oauth_invalid_at_unix_secs = None;
updated.oauth_invalid_reason = None;
if updated.fingerprint.is_none() {
updated.fingerprint = grok_oauth_catalog_key_fingerprint(provider_type, auth_config);
}
updated.health_by_format = Some(json!({}));
updated.circuit_breaker_by_format = Some(json!({}));
updated.error_count = Some(0);
@@ -223,7 +238,9 @@ fn provider_oauth_catalog_key_api_formats(
#[cfg(test)]
mod tests {
use super::provider_oauth_token_payload_expires_at_unix_secs;
use super::{
grok_oauth_catalog_key_fingerprint, provider_oauth_token_payload_expires_at_unix_secs,
};
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use serde_json::json;
@@ -273,4 +290,60 @@ mod tests {
Some(2_000_000_000)
);
}
#[test]
fn grok_oauth_catalog_key_fingerprint_uses_browser_wreq_profile() {
let auth_config = json!({
"sso_token": "abc",
"browser_profile": "chrome-137",
});
let auth_config = auth_config.as_object().expect("object");
let fingerprint = grok_oauth_catalog_key_fingerprint("grok", auth_config)
.expect("fingerprint should resolve");
assert_eq!(
fingerprint["transport_profile"]["profile_id"],
json!("chrome137")
);
assert_eq!(
fingerprint["transport_profile"]["backend"],
json!("browser_wreq")
);
assert_eq!(
fingerprint["transport_profile"]["extra"]["browser_profile"],
json!("chrome137")
);
}
#[test]
fn grok_oauth_catalog_key_fingerprint_infers_profile_from_user_agent() {
let auth_config = json!({
"sso_token": "abc",
"user_agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/137.0.0.0 Safari/537.36",
});
let auth_config = auth_config.as_object().expect("object");
let fingerprint = grok_oauth_catalog_key_fingerprint("grok", auth_config)
.expect("fingerprint should resolve");
assert_eq!(
fingerprint["transport_profile"]["profile_id"],
json!("chrome137")
);
assert_eq!(
fingerprint["transport_profile"]["extra"]["browser_profile"],
json!("chrome137")
);
}
#[test]
fn grok_oauth_catalog_key_fingerprint_ignores_non_grok_providers() {
let auth_config = json!({
"browser_profile": "chrome136",
});
let auth_config = auth_config.as_object().expect("object");
assert!(grok_oauth_catalog_key_fingerprint("openai", auth_config).is_none());
}
}

View File

@@ -4,6 +4,7 @@ use std::pin::Pin;
use super::antigravity::refresh_antigravity_provider_quota_locally;
use super::chatgpt_web::refresh_chatgpt_web_provider_quota_locally;
use super::codex::refresh_codex_provider_quota_locally;
use super::grok::refresh_grok_provider_quota_locally;
use super::kiro::refresh_kiro_provider_quota_locally;
use crate::handlers::admin::request::AdminAppState;
use crate::GatewayError;
@@ -33,6 +34,7 @@ const PROVIDER_QUOTA_REFRESH_HANDLERS: &[(&str, ProviderQuotaRefreshHandler)] =
refresh_chatgpt_web_provider_quota_locally_boxed,
),
("codex", refresh_codex_provider_quota_locally_boxed),
("grok", refresh_grok_provider_quota_locally_boxed),
("kiro", refresh_kiro_provider_quota_locally_boxed),
];
@@ -117,3 +119,19 @@ fn refresh_kiro_provider_quota_locally_boxed<'a>(
proxy_override,
))
}
fn refresh_grok_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_grok_provider_quota_locally(
state,
provider,
endpoint,
keys,
proxy_override,
))
}

View File

@@ -0,0 +1,826 @@
use super::shared::{
build_quota_snapshot_payload, default_provider_quota_execution_timeouts,
execute_provider_quota_plan, extract_execution_error_message,
persist_provider_quota_refresh_state, quota_refresh_success_invalid_state,
ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::provider::shared::payloads::{
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
};
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::GatewayError;
use aether_contracts::{
ExecutionPlan, ExecutionResult, ProxySnapshot, RequestBody, ResolvedTransportProfile,
};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_provider_pool::{
grok_pool_tier_from_quota_bucket, grok_supported_quota_windows_for_tier,
};
use aether_provider_transport::grok_browser_profile_metadata_from_resolved_transport_profile;
use base64::Engine as _;
use serde_json::json;
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
use uuid::Uuid;
const GROK_DEFAULT_BASE_URL: &str = "https://grok.com";
const GROK_RATE_LIMITS_PATH: &str = "/rest/rate-limits";
const GROK_STATSIG_ID: &str = "ZTpUeXBlRXJyb3I6IENhbm5vdCByZWFkIHByb3BlcnRpZXMgb2YgdW5kZWZpbmVkIChyZWFkaW5nICdjaGlsZE5vZGVzJyk=";
fn grok_base_url(endpoint: &StoredProviderCatalogEndpoint) -> String {
let base_url = endpoint.base_url.trim().trim_end_matches('/');
if base_url.is_empty() {
GROK_DEFAULT_BASE_URL.to_string()
} else {
base_url.to_string()
}
}
fn grok_auth_config(
transport: &AdminGatewayProviderTransportSnapshot,
) -> Option<serde_json::Value> {
transport
.key
.decrypted_auth_config
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.and_then(|value| serde_json::from_str::<serde_json::Value>(value).ok())
}
fn grok_auth_string(auth_config: Option<&serde_json::Value>, fields: &[&str]) -> Option<String> {
let object = auth_config.and_then(serde_json::Value::as_object)?;
fields.iter().find_map(|field| {
object
.get(*field)
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
}
fn build_grok_quota_headers(
auth_config: Option<&serde_json::Value>,
transport_profile: Option<&ResolvedTransportProfile>,
base_url: &str,
) -> Option<BTreeMap<String, String>> {
let cookie = build_grok_quota_cookie(auth_config).unwrap_or_default();
let browser_profile =
grok_browser_profile_metadata_from_resolved_transport_profile(transport_profile?)?;
Some(BTreeMap::from([
("accept".to_string(), "*/*".to_string()),
(
"accept-language".to_string(),
"zh-CN,zh;q=0.9,en;q=0.8".to_string(),
),
(
"baggage".to_string(),
"sentry-environment=production,sentry-release=d6add6fb0460641fd482d767a335ef72b9b6abb8,sentry-public_key=b311e0f2690c81f25e2c4cf6d4f7ce1c".to_string(),
),
("content-type".to_string(), "application/json".to_string()),
("origin".to_string(), base_url.to_string()),
("priority".to_string(), "u=1, i".to_string()),
("referer".to_string(), format!("{base_url}/")),
("sec-ch-ua".to_string(), browser_profile.sec_ch_ua),
("sec-ch-ua-mobile".to_string(), "?0".to_string()),
("sec-ch-ua-model".to_string(), String::new()),
(
"sec-ch-ua-platform".to_string(),
browser_profile.sec_ch_ua_platform,
),
("sec-fetch-dest".to_string(), "empty".to_string()),
("sec-fetch-mode".to_string(), "cors".to_string()),
("sec-fetch-site".to_string(), "same-origin".to_string()),
("user-agent".to_string(), browser_profile.user_agent),
("cookie".to_string(), cookie),
("x-statsig-id".to_string(), GROK_STATSIG_ID.to_string()),
("x-xai-request-id".to_string(), Uuid::new_v4().to_string()),
]))
}
fn build_grok_quota_cookie(auth_config: Option<&serde_json::Value>) -> Option<String> {
let token = grok_auth_string(auth_config, &["sso_token", "access_token", "token"])?;
let token = strip_cookie_prefix(token.trim(), "sso=");
if token.is_empty() {
return None;
}
let sso_rw = grok_auth_string(auth_config, &["sso_rw_token", "ssoRwToken"])
.map(|value| strip_cookie_prefix(value.trim(), "sso-rw="))
.filter(|value| !value.is_empty())
.unwrap_or_else(|| token.clone());
let mut parts = vec![format!("sso={token}"), format!("sso-rw={sso_rw}")];
if let Some(extra_cookies) =
grok_auth_string(auth_config, &["cf_cookies", "cfCookies", "cookie"])
.and_then(|value| normalize_grok_extra_cookies(value.as_str()))
{
parts.push(extra_cookies);
}
let cf_clearance = grok_auth_string(auth_config, &["cf_clearance", "cfClearance"])
.map(|value| strip_cookie_prefix(value.trim(), "cf_clearance="))
.filter(|value| !value.is_empty());
if let Some(cf_clearance) = cf_clearance {
if !parts.iter().any(|part| part.contains("cf_clearance=")) {
parts.push(format!("cf_clearance={cf_clearance}"));
}
}
Some(parts.join("; "))
}
fn strip_cookie_prefix(value: &str, prefix: &str) -> String {
value
.strip_prefix(prefix)
.map(str::trim)
.unwrap_or(value)
.to_string()
}
fn normalize_grok_extra_cookies(value: &str) -> Option<String> {
let parts = value
.trim()
.trim_matches(';')
.split(';')
.filter_map(|segment| {
let (name, value) = segment.trim().split_once('=')?;
let name = name.trim();
let value = value.trim();
if name.is_empty()
|| value.is_empty()
|| name.eq_ignore_ascii_case("sso")
|| name.eq_ignore_ascii_case("sso-rw")
{
return None;
}
Some(format!("{name}={value}"))
})
.collect::<Vec<_>>();
(!parts.is_empty()).then(|| parts.join("; "))
}
#[derive(Debug, Clone, Copy, PartialEq)]
struct GrokRateLimitSnapshot {
remaining: f64,
total: f64,
window_seconds: u64,
wait_time_seconds: Option<u64>,
}
impl GrokRateLimitSnapshot {
fn reset_after_seconds(self) -> u64 {
self.wait_time_seconds.unwrap_or(self.window_seconds)
}
fn reset_at_source(self) -> &'static str {
if self.wait_time_seconds.is_some() {
"grok_rate_limits_wait_time"
} else {
"grok_rate_limits_window"
}
}
}
fn parse_grok_rate_limits(body: &serde_json::Value) -> Option<GrokRateLimitSnapshot> {
let remaining = body
.get("remainingQueries")
.and_then(serde_json::Value::as_f64)?;
let total = body
.get("totalQueries")
.and_then(serde_json::Value::as_f64)
.unwrap_or(remaining.max(0.0));
let window_seconds = body
.get("windowSizeSeconds")
.and_then(serde_json::Value::as_u64)
.unwrap_or(72_000);
let wait_time_seconds = body
.get("waitTimeSeconds")
.and_then(serde_json::Value::as_u64);
Some(GrokRateLimitSnapshot {
remaining,
total,
window_seconds,
wait_time_seconds,
})
}
fn grok_pool_tier_hint_for_refresh(
key: &StoredProviderCatalogKey,
auth_config: Option<&serde_json::Value>,
) -> Option<&'static str> {
key.status_snapshot
.as_ref()
.and_then(serde_json::Value::as_object)
.and_then(|snapshot| snapshot.get("quota"))
.and_then(serde_json::Value::as_object)
.and_then(grok_pool_tier_from_quota_bucket)
.or_else(|| {
key.upstream_metadata
.as_ref()
.and_then(serde_json::Value::as_object)
.and_then(|metadata| metadata.get("grok"))
.and_then(serde_json::Value::as_object)
.and_then(grok_pool_tier_from_quota_bucket)
})
.or_else(|| {
auth_config
.and_then(serde_json::Value::as_object)
.and_then(grok_pool_tier_from_quota_bucket)
})
}
async fn execute_grok_quota_plan(
state: &AdminAppState<'_>,
transport: &AdminGatewayProviderTransportSnapshot,
endpoint: &StoredProviderCatalogEndpoint,
body: serde_json::Value,
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 transport_profile = state.resolve_transport_profile(transport);
let base_url = grok_base_url(endpoint);
let headers = build_grok_quota_headers(
grok_auth_config(transport).as_ref(),
transport_profile.as_ref(),
&base_url,
)
.ok_or_else(|| {
GatewayError::Internal("unsupported Grok browser transport profile".to_string())
})?;
let plan = ExecutionPlan {
request_id: format!("grok-quota:{}", transport.key.id),
candidate_id: None,
provider_name: Some("grok".to_string()),
provider_id: transport.provider.id.clone(),
endpoint_id: transport.endpoint.id.clone(),
key_id: transport.key.id.clone(),
method: "POST".to_string(),
url: format!(
"{}/{}",
base_url,
GROK_RATE_LIMITS_PATH.trim_start_matches('/')
),
headers,
content_type: Some("application/json".to_string()),
content_encoding: None,
body: RequestBody::from_json(body),
stream: false,
client_api_format: "openai:responses".to_string(),
provider_api_format: "grok:rate_limits".to_string(),
model_name: Some("grok-quota".to_string()),
proxy,
transport_profile,
timeouts,
};
execute_provider_quota_plan(state, transport, plan, "grok").await
}
fn grok_quota_error_detail(result: &ExecutionResult) -> Option<String> {
extract_execution_error_message(result).or_else(|| {
let body = result.body.as_ref()?.body_bytes_b64.as_deref()?;
let decoded = base64::engine::general_purpose::STANDARD
.decode(body)
.ok()?;
let text = String::from_utf8_lossy(&decoded).trim().to_string();
(!text.is_empty()).then_some(text)
})
}
fn grok_is_cloudflare_challenge(message: &str) -> bool {
let lowered = message.to_ascii_lowercase();
lowered.contains("cloudflare")
|| lowered.contains("just a moment")
|| lowered.contains("__cf_chl")
|| lowered.contains("cf-ray")
}
fn grok_quota_invalid_reason(status_code: u16, upstream_message: Option<&str>) -> String {
let message = upstream_message.unwrap_or_default().trim();
if status_code == 403 && grok_is_cloudflare_challenge(message) {
return format!(
"{OAUTH_REFRESH_FAILED_PREFIX}Grok Cloudflare 验证失败,请重新从同一浏览器复制最新 Cookie 和 User-Agent或配置可通过 Cloudflare 的代理运行时"
);
}
let detail = if message.is_empty() {
match status_code {
401 => "Grok Token 无效或已过期",
403 => "Grok 账户访问受限",
_ => "Grok 请求失败",
}
} else {
message
};
match status_code {
401 => format!("{OAUTH_EXPIRED_PREFIX}{detail}"),
403 => format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}{detail}"),
_ => detail.to_string(),
}
}
fn grok_quota_result_message(reason: &str) -> String {
for prefix in [
OAUTH_REFRESH_FAILED_PREFIX,
OAUTH_EXPIRED_PREFIX,
OAUTH_ACCOUNT_BLOCK_PREFIX,
] {
if let Some(message) = reason.strip_prefix(prefix) {
return message.trim().to_string();
}
}
reason.trim().to_string()
}
pub(crate) async fn refresh_grok_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;
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 grok_auth_config(&transport).is_none() {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 Grok 账号会话信息,请先导入 Token",
}));
continue;
}
let auth_config = grok_auth_config(&transport);
let mut quota_by_model = serde_json::Map::new();
let mut refreshed = false;
let mut invalid_reason = None::<String>;
let mut invalid_at = key.oauth_invalid_at_unix_secs;
let mut last_status_code = None::<u16>;
let mut last_error_message = None::<String>;
let mut metadata_update = serde_json::Map::new();
let base_url = grok_base_url(endpoint);
let supported_windows = grok_supported_quota_windows_for_tier(
grok_pool_tier_hint_for_refresh(&key, auth_config.as_ref()),
);
for (quota_key, mode_name) in supported_windows.iter().copied() {
let result = match execute_grok_quota_plan(
state,
&transport,
endpoint,
json!({ "modelName": mode_name }),
proxy_override.as_ref(),
)
.await?
{
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
last_error_message = Some(format!("rate-limits 请求执行失败: {detail}"));
continue;
}
};
last_status_code = Some(result.status_code);
if result.status_code == 200 {
if let Some(body_json) = result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
{
if let Some(rate_limit) = parse_grok_rate_limits(body_json) {
refreshed = true;
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let reset_after_seconds = rate_limit.reset_after_seconds();
let reset_at = now_unix_secs.saturating_add(reset_after_seconds);
quota_by_model.insert(
(*quota_key).to_string(),
json!({
"display_name": *mode_name,
"remaining_fraction": if rate_limit.total > 0.0 { Some((rate_limit.remaining / rate_limit.total).clamp(0.0, 1.0)) } else { None::<f64> },
"used_percent": if rate_limit.total > 0.0 { Some(((rate_limit.total - rate_limit.remaining).max(0.0) / rate_limit.total * 100.0).clamp(0.0, 100.0)) } else { None::<f64> },
"remaining": rate_limit.remaining,
"total": rate_limit.total,
"window_seconds": rate_limit.window_seconds,
"wait_time_seconds": rate_limit.wait_time_seconds,
"reset_after_seconds": reset_after_seconds,
"reset_at": reset_at,
"next_reset_at": reset_at,
"reset_at_source": rate_limit.reset_at_source(),
"is_exhausted": rate_limit.remaining <= 0.0,
}),
);
} else {
last_error_message = Some(
"Grok rate-limits 未返回 remainingQueries/totalQueries".to_string(),
);
}
} else {
last_error_message = Some("Grok rate-limits 未返回 JSON 数据".to_string());
}
} else if matches!(result.status_code, 401 | 403) {
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
invalid_at = Some(now_unix_secs);
let error_detail = grok_quota_error_detail(&result);
invalid_reason = Some(grok_quota_invalid_reason(
result.status_code,
error_detail.as_deref(),
));
last_error_message = invalid_reason.as_deref().map(grok_quota_result_message);
} else {
let error_detail =
grok_quota_error_detail(&result).unwrap_or_else(|| "Grok 请求失败".to_string());
last_error_message = Some(format!(
"Grok rate-limits 请求失败({}): {error_detail}",
result.status_code
));
}
}
if refreshed {
if let Some(pool_tier) = grok_pool_tier_from_quota_bucket(&quota_by_model)
.or_else(|| grok_pool_tier_hint_for_refresh(&key, auth_config.as_ref()))
{
let pool_tier_value = json!(pool_tier);
metadata_update.insert("pool_tier".to_string(), pool_tier_value.clone());
metadata_update
.entry("plan_type".to_string())
.or_insert(pool_tier_value);
}
metadata_update.insert(
"updated_at".to_string(),
json!(SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0)),
);
metadata_update.insert("base_url".to_string(), json!(base_url));
metadata_update.insert("quota_by_model".to_string(), json!(quota_by_model));
}
let metadata_update_value = if metadata_update.is_empty() {
None
} else {
Some(serde_json::Value::Object({
let mut map = serde_json::Map::new();
map.insert(
"grok".to_string(),
serde_json::Value::Object(metadata_update.clone()),
);
map
}))
};
if !persist_provider_quota_refresh_state(
state,
&key.id,
metadata_update_value.as_ref(),
invalid_at,
invalid_reason,
None,
)
.await?
{
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Key 状态写入失败",
}));
continue;
}
if refreshed {
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!(if refreshed { "success" } else { "error" }),
);
if let Some(metadata) = metadata_update.get("quota_by_model").cloned() {
payload.insert("metadata".to_string(), metadata);
}
if let Some(quota_snapshot) = build_quota_snapshot_payload(
"grok",
key.status_snapshot.as_ref(),
metadata_update_value.as_ref(),
) {
payload.insert("quota_snapshot".to_string(), quota_snapshot);
}
if !refreshed {
payload.insert(
"message".to_string(),
json!(last_error_message.unwrap_or_else(|| {
"Grok rate-limits 未返回可用配额数据".to_string()
})),
);
if let Some(status_code) = last_status_code {
payload.insert("status_code".to_string(), json!(status_code));
}
}
results.push(serde_json::Value::Object(payload));
}
Ok(Some(json!({
"success": success_count,
"failed": failed_count,
"total": success_count + failed_count,
"results": results,
"message": format!("已处理 {} 个 Key", success_count + failed_count),
"auto_removed": 0,
})))
}
#[cfg(test)]
mod tests {
use super::{
build_grok_quota_cookie, build_grok_quota_headers, grok_pool_tier_hint_for_refresh,
grok_quota_error_detail, grok_quota_invalid_reason, grok_quota_result_message,
parse_grok_rate_limits,
};
use crate::handlers::admin::provider::shared::payloads::OAUTH_REFRESH_FAILED_PREFIX;
use aether_contracts::{ExecutionResult, ResponseBody};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use base64::Engine as _;
use serde_json::json;
use std::collections::BTreeMap;
fn sample_key(
status_snapshot: Option<serde_json::Value>,
upstream_metadata: Option<serde_json::Value>,
) -> StoredProviderCatalogKey {
let mut key = StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"key-1".to_string(),
"oauth".to_string(),
None,
true,
)
.expect("key should build");
key.status_snapshot = status_snapshot;
key.upstream_metadata = upstream_metadata;
key
}
#[test]
fn quota_cookie_preserves_grok_session_and_clearance() {
let auth_config = json!({
"sso_token": "sso=abc",
"sso_rw_token": "sso-rw=rw",
"cf_clearance": "cf"
});
let cookie = build_grok_quota_cookie(Some(&auth_config)).expect("cookie should build");
assert_eq!(cookie, "sso=abc; sso-rw=rw; cf_clearance=cf");
}
#[test]
fn quota_cookie_removes_duplicate_session_cookies_from_cf_profile() {
let auth_config = json!({
"sso_token": "abc",
"sso_rw_token": "rw",
"cf_cookies": "i18nextLng=zh; sso=ignored; sso-rw=ignored-rw; cf_clearance=cf"
});
let cookie = build_grok_quota_cookie(Some(&auth_config)).expect("cookie should build");
assert_eq!(cookie, "sso=abc; sso-rw=rw; i18nextLng=zh; cf_clearance=cf");
}
#[test]
fn quota_headers_use_resolved_transport_profile_user_agent() {
let auth_config = json!({
"sso_token": "abc",
"user_agent": "Mozilla/5.0 custom"
});
let transport_profile = aether_provider_transport::grok_browser_resolved_transport_profile(
Some("chrome137"),
"test",
)
.expect("profile should resolve");
let headers = build_grok_quota_headers(
Some(&auth_config),
Some(&transport_profile),
"https://grok.com",
)
.expect("headers should build");
assert!(headers
.get("user-agent")
.is_some_and(|value| value.contains("Chrome/137.0.0.0")));
assert_eq!(
headers.get("sec-ch-ua"),
Some(
&r#""Google Chrome";v="137", "Chromium";v="137", "Not(A:Brand";v="24""#.to_string()
)
);
}
#[test]
fn quota_headers_default_to_chrome136_clearance_profile() {
let auth_config = json!({
"sso_token": "abc"
});
let transport_profile =
aether_provider_transport::grok_browser_resolved_transport_profile(None, "test")
.expect("profile should resolve");
let headers = build_grok_quota_headers(
Some(&auth_config),
Some(&transport_profile),
"https://grok.com",
)
.expect("headers should build");
assert!(headers
.get("user-agent")
.is_some_and(|value| value.contains("Chrome/136.0.0.0")));
assert_eq!(
headers.get("sec-ch-ua"),
Some(
&r#""Google Chrome";v="136", "Chromium";v="136", "Not(A:Brand";v="24""#.to_string()
)
);
assert_eq!(
headers.get("sec-ch-ua-platform"),
Some(&r#""macOS""#.to_string())
);
assert!(headers.contains_key("x-statsig-id"));
assert!(headers.contains_key("x-xai-request-id"));
}
#[test]
fn quota_headers_do_not_mark_rate_limits_as_grok_app_chat_runtime() {
let auth_config = json!({
"sso_token": "abc"
});
let transport_profile =
aether_provider_transport::grok_browser_resolved_transport_profile(None, "test")
.expect("profile should resolve");
let headers = build_grok_quota_headers(
Some(&auth_config),
Some(&transport_profile),
"https://grok.com",
)
.expect("headers should build");
assert!(!headers.contains_key(aether_provider_transport::GROK_INTERNAL_HEADER));
}
#[test]
fn parses_grok_wait_time_seconds_as_authoritative_reset_delay() {
let body = json!({
"windowSizeSeconds": 86_400,
"remainingQueries": 0,
"waitTimeSeconds": 12_648,
"totalQueries": 30,
"lowEffortRateLimits": null,
"highEffortRateLimits": null
});
let rate_limits = parse_grok_rate_limits(&body).expect("rate limits should parse");
assert_eq!(rate_limits.remaining, 0.0);
assert_eq!(rate_limits.total, 30.0);
assert_eq!(rate_limits.window_seconds, 86_400);
assert_eq!(rate_limits.wait_time_seconds, Some(12_648));
assert_eq!(rate_limits.reset_after_seconds(), 12_648);
assert_eq!(rate_limits.reset_at_source(), "grok_rate_limits_wait_time");
}
#[test]
fn parses_grok_rate_limits_falls_back_to_window_when_wait_time_is_absent() {
let body = json!({
"windowSizeSeconds": 86_400,
"remainingQueries": 12,
"totalQueries": 30
});
let rate_limits = parse_grok_rate_limits(&body).expect("rate limits should parse");
assert_eq!(rate_limits.remaining, 12.0);
assert_eq!(rate_limits.total, 30.0);
assert_eq!(rate_limits.window_seconds, 86_400);
assert_eq!(rate_limits.wait_time_seconds, None);
assert_eq!(rate_limits.reset_after_seconds(), 86_400);
assert_eq!(rate_limits.reset_at_source(), "grok_rate_limits_window");
}
#[test]
fn infers_grok_pool_tier_from_live_quota_totals() {
let key = sample_key(
Some(json!({
"quota": {
"pool_tier": "heavy"
}
})),
None,
);
assert_eq!(
grok_pool_tier_hint_for_refresh(&key, Some(&json!({}))),
Some("heavy")
);
}
#[test]
fn infers_basic_grok_pool_tier_from_fast_quota_when_auto_is_absent() {
let key = sample_key(
None,
Some(json!({
"grok": {
"plan_type": "basic"
}
})),
);
assert_eq!(grok_pool_tier_hint_for_refresh(&key, None), Some("basic"));
}
#[test]
fn cloudflare_challenge_403_is_not_account_block() {
let body = "<!DOCTYPE html><html><head><title>Just a moment...</title></head><body>Cloudflare</body></html>";
let result = ExecutionResult {
request_id: "grok-quota:test".to_string(),
candidate_id: None,
status_code: 403,
headers: BTreeMap::new(),
body: Some(ResponseBody {
json_body: None,
body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)),
}),
telemetry: None,
error: None,
};
let detail = grok_quota_error_detail(&result).expect("html body should be decoded");
let reason = grok_quota_invalid_reason(result.status_code, Some(&detail));
assert!(reason.starts_with("[REFRESH_FAILED] "));
assert!(!reason.starts_with("[ACCOUNT_BLOCK] "));
assert!(reason.contains("Cloudflare"));
}
#[test]
fn quota_result_message_removes_status_prefix() {
let reason = format!("{OAUTH_REFRESH_FAILED_PREFIX}Grok Cloudflare 验证失败");
assert_eq!(
grok_quota_result_message(&reason),
"Grok Cloudflare 验证失败"
);
}
}

View File

@@ -2,5 +2,6 @@ pub(crate) mod antigravity;
pub(crate) mod chatgpt_web;
pub(crate) mod codex;
pub(crate) mod dispatch;
pub(crate) mod grok;
pub(crate) mod kiro;
pub(crate) mod shared;

View File

@@ -63,6 +63,13 @@ fn select_provider_oauth_runtime_endpoint(
.trim()
.eq_ignore_ascii_case("openai:image")
}),
"grok" => matching_endpoint(endpoints, include_inactive, |endpoint| {
endpoint
.api_format
.trim()
.eq_ignore_ascii_case("openai:chat")
})
.or_else(|| matching_endpoint(endpoints, include_inactive, |_| true)),
"antigravity" => matching_endpoint(endpoints, include_inactive, |endpoint| {
endpoint
.api_format

View File

@@ -40,7 +40,7 @@ pub(super) async fn admin_provider_ops_sub2api_balance_payload(
"query_balance",
message,
None,
)
);
}
};

View File

@@ -57,7 +57,7 @@ pub(super) async fn handle_admin_provider_ops_action(
Err(_) => {
return Ok(Some(bad_request_detail_response(
"请求体必须是合法的 JSON 对象",
)))
)));
}
};
let payload =
@@ -67,7 +67,7 @@ pub(super) async fn handle_admin_provider_ops_action(
Err(_) => {
return Ok(Some(bad_request_detail_response(
"请求体必须是合法的 JSON 对象",
)))
)));
}
};
payload.config

View File

@@ -3,7 +3,9 @@ use crate::handlers::admin::provider::shared::support::{
};
use crate::handlers::admin::request::AdminAppState;
use crate::handlers::admin::shared::{provider_key_status_snapshot_payload, unix_secs_to_rfc3339};
use crate::provider_key_auth::{provider_key_auth_semantics, provider_key_effective_api_formats};
use crate::provider_key_auth::{
provider_key_auth_semantics, provider_key_can_refresh_oauth, provider_key_effective_api_formats,
};
use aether_admin::provider::pool as admin_provider_pool_pure;
use aether_admin::provider::quota as admin_provider_quota_pure;
use aether_data_contracts::repository::pool_scores::StoredPoolMemberScore;
@@ -597,6 +599,185 @@ fn admin_pool_build_antigravity_account_quota_from_snapshot(
))
}
fn admin_pool_grok_quota_window_label(
window: &serde_json::Map<String, serde_json::Value>,
) -> String {
let raw_code = window
.get("code")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.unwrap_or_default()
.trim_start_matches("model:")
.to_ascii_lowercase();
let raw_label = window
.get("label")
.or_else(|| window.get("model"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or(raw_code.as_str())
.to_ascii_lowercase();
match raw_label.as_str() {
"quota_auto" | "auto" => "Auto".to_string(),
"quota_fast" | "fast" => "Fast".to_string(),
"quota_expert" | "expert" => "Expert".to_string(),
"quota_heavy" | "heavy" => "Heavy".to_string(),
"quota_grok_4_3" | "grok-420-computer-use-sa" => "Grok 4.3".to_string(),
_ => match raw_code.as_str() {
"quota_auto" | "auto" => "Auto".to_string(),
"quota_fast" | "fast" => "Fast".to_string(),
"quota_expert" | "expert" => "Expert".to_string(),
"quota_heavy" | "heavy" => "Heavy".to_string(),
"quota_grok_4_3" | "grok-420-computer-use-sa" => "Grok 4.3".to_string(),
_ => window
.get("label")
.or_else(|| window.get("model"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or("模式")
.to_string(),
},
}
}
fn admin_pool_quota_window_remaining_percent(
window: &serde_json::Map<String, serde_json::Value>,
) -> Option<f64> {
admin_pool_json_to_f64(window.get("remaining_ratio"))
.map(|value| (value * 100.0).clamp(0.0, 100.0))
.or_else(|| {
admin_pool_json_to_f64(window.get("used_ratio"))
.map(|value| ((1.0 - value) * 100.0).clamp(0.0, 100.0))
})
.or_else(|| {
admin_pool_json_to_f64(window.get("remaining_value"))
.zip(admin_pool_json_to_f64(window.get("limit_value")))
.and_then(|(remaining, limit)| {
(limit > 0.0).then_some((remaining / limit * 100.0).clamp(0.0, 100.0))
})
})
.or_else(|| {
admin_pool_json_to_f64(window.get("used_value"))
.zip(admin_pool_json_to_f64(window.get("limit_value")))
.and_then(|(used, limit)| {
(limit > 0.0).then_some(((1.0 - used / limit) * 100.0).clamp(0.0, 100.0))
})
})
}
fn admin_pool_quota_window_value_text(
window: &serde_json::Map<String, serde_json::Value>,
) -> Option<String> {
let limit_value =
admin_pool_json_to_f64(window.get("limit_value")).filter(|value| *value > 0.0)?;
if let Some(remaining_value) = admin_pool_json_to_f64(window.get("remaining_value")) {
return Some(format!(
"{}/{}",
admin_pool_format_quota_value(remaining_value),
admin_pool_format_quota_value(limit_value),
));
}
admin_pool_json_to_f64(window.get("used_value")).map(|used_value| {
format!(
"{}/{}",
admin_pool_format_quota_value((limit_value - used_value).max(0.0)),
admin_pool_format_quota_value(limit_value),
)
})
}
fn admin_pool_build_grok_account_quota_from_snapshot(
quota_snapshot: &serde_json::Map<String, serde_json::Value>,
) -> Option<String> {
let code = quota_snapshot
.get("code")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.unwrap_or_default();
if code.eq_ignore_ascii_case("banned") {
return quota_snapshot
.get("label")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned)
.or_else(|| Some("账号已封禁".to_string()));
}
if code.eq_ignore_ascii_case("forbidden") {
return quota_snapshot
.get("label")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned)
.or_else(|| Some("访问受限".to_string()));
}
let model_parts = admin_pool_quota_windows(quota_snapshot)
.into_iter()
.filter(|window| {
window
.get("scope")
.and_then(serde_json::Value::as_str)
.is_some_and(|scope| scope.eq_ignore_ascii_case("model"))
})
.filter_map(|window| {
let remaining_percent = admin_pool_quota_window_remaining_percent(window)?;
let mut part = format!(
"{}剩余 {}",
admin_pool_grok_quota_window_label(window),
admin_pool_format_percent(remaining_percent),
);
if let Some(value_text) = admin_pool_quota_window_value_text(window) {
part.push_str(&format!(" ({value_text})"));
}
Some(part)
})
.collect::<Vec<_>>();
if !model_parts.is_empty() {
return Some(model_parts.join(" | "));
}
let window = admin_pool_quota_window(quota_snapshot, "usage")
.or_else(|| admin_pool_quota_windows(quota_snapshot).into_iter().next())?;
let remaining_value = admin_pool_json_to_f64(window.get("remaining_value"));
let limit_value = admin_pool_json_to_f64(window.get("limit_value"));
if let (Some(remaining_value), Some(limit_value)) = (remaining_value, limit_value) {
if limit_value > 0.0 && remaining_value <= 0.0 {
return Some(format!(
"剩余 {}/{}",
admin_pool_format_quota_value(remaining_value),
admin_pool_format_quota_value(limit_value),
));
}
}
if let Some(remaining_percent) = admin_pool_quota_window_remaining_percent(window) {
if let Some(value_text) = admin_pool_quota_window_value_text(window) {
return Some(format!(
"剩余 {} ({value_text})",
admin_pool_format_percent(remaining_percent),
));
}
return Some(format!(
"剩余 {}",
admin_pool_format_percent(remaining_percent),
));
}
match (remaining_value, limit_value) {
(Some(remaining_value), Some(limit_value)) if limit_value > 0.0 => Some(format!(
"剩余 {}/{}",
admin_pool_format_quota_value(remaining_value),
admin_pool_format_quota_value(limit_value),
)),
_ => quota_snapshot
.get("label")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned),
}
}
fn admin_pool_build_gemini_cli_account_quota_from_snapshot(
quota_snapshot: &serde_json::Map<String, serde_json::Value>,
) -> Option<String> {
@@ -699,6 +880,13 @@ fn admin_pool_build_account_quota(
return Some(account_quota);
}
}
"grok" => {
if let Some(account_quota) =
admin_pool_build_grok_account_quota_from_snapshot(quota_snapshot)
{
return Some(account_quota);
}
}
"gemini_cli" => {
if let Some(account_quota) =
admin_pool_build_gemini_cli_account_quota_from_snapshot(quota_snapshot)
@@ -962,7 +1150,10 @@ pub(super) fn build_admin_pool_key_payload(
);
payload.insert(
"can_refresh_oauth".to_string(),
json!(auth_semantics.can_refresh_oauth()),
json!(provider_key_can_refresh_oauth(
auth_semantics,
auth_config.as_ref()
)),
);
payload.insert(
"can_export_oauth".to_string(),
@@ -1225,4 +1416,39 @@ mod tests {
assert_eq!(usage["total_tokens"], json!(375));
assert_eq!(usage["total_cost_usd"], json!("0.60000000"));
}
#[test]
fn grok_model_quota_is_rendered_for_pool_rows() {
let quota_snapshot = json!({
"provider_type": "grok",
"code": "ok",
"exhausted": false,
"plan_type": "heavy",
"pool_tier": "heavy",
"windows": [
{
"code": "model:quota_auto",
"label": "auto",
"scope": "model",
"remaining_ratio": 0.4,
"used_value": 90,
"limit_value": 150
},
{
"code": "model:quota_heavy",
"label": "heavy",
"scope": "model",
"remaining_ratio": 0.0,
"used_value": 20,
"limit_value": 20
}
]
});
let quota_snapshot = quota_snapshot.as_object().unwrap();
assert_eq!(
admin_pool_build_account_quota("grok", Some(quota_snapshot)),
Some("Auto剩余 40.0% (60/150) | Heavy剩余 0.0% (0/20)".to_string())
);
}
}

View File

@@ -3,7 +3,7 @@ use super::{
AdminPoolResolveSelectionRequest, ADMIN_POOL_PROVIDER_CATALOG_READER_UNAVAILABLE_DETAIL,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::provider_key_auth::provider_key_auth_semantics;
use crate::provider_key_auth::{provider_key_auth_semantics, provider_key_can_refresh_oauth};
use crate::GatewayError;
use aether_admin::provider::pool as admin_provider_pool_pure;
use axum::{
@@ -94,6 +94,7 @@ pub(super) async fn build_admin_pool_resolve_selection_response(
.iter()
.map(|key| {
let auth_semantics = provider_key_auth_semantics(key, &provider_type);
let auth_config = state.parse_catalog_auth_config_json(key);
json!({
"key_id": key.id,
"key_name": key.name,
@@ -102,7 +103,7 @@ pub(super) async fn build_admin_pool_resolve_selection_response(
"credential_kind": auth_semantics.credential_kind().as_str(),
"runtime_auth_kind": auth_semantics.runtime_auth_kind().as_str(),
"oauth_managed": auth_semantics.oauth_managed(),
"can_refresh_oauth": auth_semantics.can_refresh_oauth(),
"can_refresh_oauth": provider_key_can_refresh_oauth(auth_semantics, auth_config.as_ref()),
"can_export_oauth": auth_semantics.can_export_oauth(),
"can_edit_oauth": auth_semantics.can_edit_oauth(),
})

View File

@@ -17,6 +17,9 @@ use crate::ai_serving::{
};
use crate::clock::current_unix_ms;
use crate::execution_runtime;
use crate::handlers::admin::provider::shared::model_test_capabilities::{
admin_provider_model_supports_image_generation, admin_provider_model_test_capabilities_payload,
};
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::handlers::shared::provider_pool::{
admin_provider_pool_config_from_config_value, read_admin_provider_pool_runtime_state,
@@ -100,6 +103,181 @@ struct ProviderQueryKeyFetchResult {
has_success: bool,
}
fn provider_query_model_id(model: &Value) -> Option<&str> {
model
.get("id")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
}
fn provider_query_grok_required_tier_rank(model_id: &str) -> Option<u8> {
match model_id.trim() {
"grok-4.20-0309-non-reasoning" | "grok-4.20-fast" | "grok-imagine-image-lite" => Some(0),
"grok-4.20-0309"
| "grok-4.20-0309-reasoning"
| "grok-4.20-0309-non-reasoning-super"
| "grok-4.20-0309-super"
| "grok-4.20-0309-reasoning-super"
| "grok-4.20-auto"
| "grok-4.20-expert"
| "grok-4.3-beta"
| "grok-imagine-image"
| "grok-imagine-image-pro"
| "grok-imagine-image-edit" => Some(1),
"grok-4.20-0309-non-reasoning-heavy"
| "grok-4.20-0309-heavy"
| "grok-4.20-0309-reasoning-heavy"
| "grok-4.20-multi-agent-0309"
| "grok-4.20-heavy" => Some(2),
_ => None,
}
}
fn provider_query_normalize_grok_pool_tier(value: Option<&str>) -> Option<&'static str> {
match value?.trim().to_ascii_lowercase().as_str() {
"basic" => Some("basic"),
"super" => Some("super"),
"heavy" => Some("heavy"),
_ => None,
}
}
fn provider_query_grok_pool_tier_rank(value: Option<&str>) -> u8 {
match provider_query_normalize_grok_pool_tier(value).unwrap_or("basic") {
"heavy" => 2,
"super" => 1,
_ => 0,
}
}
fn provider_query_grok_quota_string(quota: &Map<String, Value>, fields: &[&str]) -> Option<String> {
fields.iter().find_map(|field| {
quota
.get(*field)
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
}
fn provider_query_grok_window_limit(quota: &Map<String, Value>, model_name: &str) -> Option<f64> {
quota
.get("windows")
.and_then(Value::as_array)?
.iter()
.filter_map(Value::as_object)
.find(|window| {
window
.get("model")
.and_then(Value::as_str)
.is_some_and(|value| value.trim() == model_name)
})
.and_then(|window| window.get("limit_value"))
.and_then(Value::as_f64)
.filter(|value| value.is_finite() && *value > 0.0)
}
fn provider_query_grok_pool_tier_from_quota(quota: &Map<String, Value>) -> Option<&'static str> {
if let Some(tier) =
provider_query_grok_quota_string(quota, &["pool_tier", "tier", "plan_type", "plan"])
.and_then(|value| provider_query_normalize_grok_pool_tier(Some(&value)))
{
return Some(tier);
}
if let Some(auto_total) = provider_query_grok_window_limit(quota, "quota_auto") {
if (auto_total - 150.0).abs() < f64::EPSILON {
return Some("heavy");
}
if (auto_total - 50.0).abs() < f64::EPSILON {
return Some("super");
}
}
if let Some(fast_total) = provider_query_grok_window_limit(quota, "quota_fast") {
if (fast_total - 400.0).abs() < f64::EPSILON {
return Some("heavy");
}
if (fast_total - 140.0).abs() < f64::EPSILON {
return Some("super");
}
if (fast_total - 30.0).abs() < f64::EPSILON {
return Some("basic");
}
}
None
}
fn provider_query_grok_key_pool_tier(key: &StoredProviderCatalogKey) -> Option<&'static str> {
key.status_snapshot
.as_ref()
.and_then(Value::as_object)
.and_then(|snapshot| snapshot.get("quota"))
.and_then(Value::as_object)
.and_then(provider_query_grok_pool_tier_from_quota)
}
fn provider_query_filter_models_for_key(
provider: &StoredProviderCatalogProvider,
key: &StoredProviderCatalogKey,
models: Vec<Value>,
) -> Vec<Value> {
if !provider.provider_type.trim().eq_ignore_ascii_case("grok") {
return models;
}
let allowed_rank = provider_query_grok_pool_tier_rank(provider_query_grok_key_pool_tier(key));
models
.into_iter()
.filter(|model| {
provider_query_model_id(model)
.and_then(provider_query_grok_required_tier_rank)
.is_some_and(|required_rank| required_rank <= allowed_rank)
})
.collect()
}
fn provider_query_attach_model_test_capabilities(
provider: &StoredProviderCatalogProvider,
models: Vec<Value>,
) -> Vec<Value> {
models
.into_iter()
.map(|mut model| {
let Some(object) = model.as_object_mut() else {
return model;
};
let model_id = object
.get("id")
.and_then(Value::as_str)
.map(str::trim)
.unwrap_or_default()
.to_string();
let supports_image_generation = admin_provider_model_supports_image_generation(
&provider.provider_type,
&model_id,
object
.get("supports_image_generation")
.or_else(|| object.get("effective_supports_image_generation"))
.and_then(Value::as_bool)
.unwrap_or(false),
);
object.insert(
"model_test_capabilities".to_string(),
admin_provider_model_test_capabilities_payload(
&provider.provider_type,
&model_id,
supports_image_generation,
),
);
model
})
.collect()
}
fn provider_query_codex_preset_fallback(
provider: &StoredProviderCatalogProvider,
) -> Option<ProviderQueryKeyFetchResult> {
@@ -245,8 +423,9 @@ async fn provider_query_fetch_models_for_key(
if let Some(cached_models) =
provider_query_read_cached_models(state, &provider.id, &key.id).await
{
let models = provider_query_filter_models_for_key(provider, key, cached_models);
return Ok(ProviderQueryKeyFetchResult {
models: cached_models,
models,
error: None,
from_cache: true,
has_success: true,
@@ -257,8 +436,13 @@ async fn provider_query_fetch_models_for_key(
let selected_endpoints = selected_models_fetch_endpoints(endpoints, key);
if selected_endpoints.is_empty() {
if let Some(models) = preset_models_for_provider(&provider.provider_type) {
let models = provider_query_filter_models_for_key(
provider,
key,
aggregate_models_for_cache(&models),
);
return Ok(ProviderQueryKeyFetchResult {
models: aggregate_models_for_cache(&models),
models,
error: None,
from_cache: false,
has_success: true,
@@ -342,7 +526,7 @@ async fn provider_query_fetch_models_for_key(
}
Ok(ProviderQueryKeyFetchResult {
models: unique_models,
models: provider_query_filter_models_for_key(provider, key, unique_models),
error,
from_cache: false,
has_success: outcome.has_success,
@@ -397,11 +581,12 @@ pub(crate) async fn build_admin_provider_query_models_response(
force_refresh,
)
.await?;
let success = !result.models.is_empty();
let models = provider_query_attach_model_test_capabilities(&provider, result.models);
let success = !models.is_empty();
return Ok(Json(json!({
"success": success,
"data": {
"models": result.models,
"models": models,
"error": result.error,
"from_cache": result.from_cache,
},
@@ -429,6 +614,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
{
if let Some(models) = provider_query_read_provider_cached_models(state, &provider.id).await
{
let models = provider_query_attach_model_test_capabilities(&provider, models);
return Ok(Json(json!({
"success": !models.is_empty(),
"data": {
@@ -504,6 +690,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
if !success && error.is_none() {
error = Some(ADMIN_PROVIDER_QUERY_NO_MODELS_FROM_KEY_DETAIL.to_string());
}
let models = provider_query_attach_model_test_capabilities(&provider, models);
Ok(Json(json!({
"success": success,
@@ -519,3 +706,152 @@ pub(crate) async fn build_admin_provider_query_models_response(
}))
.into_response())
}
#[cfg(test)]
mod tests {
use super::*;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
fn grok_provider() -> StoredProviderCatalogProvider {
let mut provider = StoredProviderCatalogProvider::new(
"provider-1".to_string(),
"Grok".to_string(),
None,
"grok".to_string(),
)
.expect("provider should build");
provider.provider_type = "grok".to_string();
provider
}
fn grok_key_with_quota(quota: Value) -> StoredProviderCatalogKey {
let mut key = StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"key-1".to_string(),
"oauth".to_string(),
None,
true,
)
.expect("key should build");
key.status_snapshot = Some(json!({ "quota": quota }));
key
}
fn model(id: &str) -> Value {
json!({ "id": id })
}
fn filtered_ids(key: &StoredProviderCatalogKey) -> Vec<String> {
provider_query_filter_models_for_key(
&grok_provider(),
key,
vec![
model("grok-4.20-0309-non-reasoning"),
model("grok-4.20-auto"),
model("grok-4.20-heavy"),
model("grok-imagine-image-lite"),
model("grok-imagine-image"),
model("grok-imagine-image-edit"),
],
)
.into_iter()
.filter_map(|item| item.get("id").and_then(Value::as_str).map(str::to_string))
.collect()
}
#[test]
fn provider_query_grok_basic_tier_hides_super_and_heavy_models() {
let key = grok_key_with_quota(json!({ "pool_tier": "basic" }));
assert_eq!(
filtered_ids(&key),
["grok-4.20-0309-non-reasoning", "grok-imagine-image-lite"]
);
}
#[test]
fn provider_query_grok_super_tier_hides_heavy_models() {
let key = grok_key_with_quota(json!({ "plan_type": "super" }));
assert_eq!(
filtered_ids(&key),
[
"grok-4.20-0309-non-reasoning",
"grok-4.20-auto",
"grok-imagine-image-lite",
"grok-imagine-image",
"grok-imagine-image-edit"
]
);
}
#[test]
fn provider_query_grok_heavy_tier_keeps_full_non_video_catalog() {
let key = grok_key_with_quota(json!({ "pool_tier": "heavy" }));
assert_eq!(
filtered_ids(&key),
[
"grok-4.20-0309-non-reasoning",
"grok-4.20-auto",
"grok-4.20-heavy",
"grok-imagine-image-lite",
"grok-imagine-image",
"grok-imagine-image-edit"
]
);
}
#[test]
fn provider_query_grok_tier_falls_back_to_live_quota_windows() {
let key = grok_key_with_quota(json!({
"windows": [
{ "model": "quota_fast", "limit_value": 140.0 }
]
}));
assert_eq!(
filtered_ids(&key),
[
"grok-4.20-0309-non-reasoning",
"grok-4.20-auto",
"grok-imagine-image-lite",
"grok-imagine-image",
"grok-imagine-image-edit"
]
);
}
#[test]
fn provider_query_attaches_model_test_capabilities_to_models() {
let models = provider_query_attach_model_test_capabilities(
&grok_provider(),
vec![
model("grok-4.20-fast"),
model("grok-imagine-image"),
model("grok-imagine-image-edit"),
],
);
assert!(models[0]["model_test_capabilities"]["openai:image"].is_null());
assert_eq!(
models[1]["model_test_capabilities"]["openai:image"]["max_generation_count"],
json!(4)
);
assert_eq!(
models[1]["model_test_capabilities"]["openai:image"]["supports_generation"],
json!(true)
);
assert_eq!(
models[2]["model_test_capabilities"]["openai:image"]["supports_generation"],
json!(false)
);
assert_eq!(
models[2]["model_test_capabilities"]["openai:image"]["supports_edit"],
json!(true)
);
}
}

View File

@@ -81,6 +81,7 @@ use tracing::{debug, warn};
use uuid::Uuid;
mod adapter;
mod capabilities;
mod model_mapping;
mod summary;
@@ -88,13 +89,17 @@ use self::adapter::{
provider_query_antigravity_test_unsupported_reason,
provider_query_antigravity_unsupported_reason,
provider_query_default_antigravity_endpoint_test_body,
provider_query_model_test_endpoint_priority, provider_query_normalize_api_format_alias,
provider_query_standard_test_client_api_format,
provider_query_grok_test_unsupported_reason, provider_query_model_test_endpoint_priority,
provider_query_normalize_api_format_alias, provider_query_standard_test_client_api_format,
provider_query_standard_test_unsupported_reason,
provider_query_test_adapter_for_provider_api_format,
provider_query_transport_supports_model_test_execution,
provider_query_unsupported_test_api_format_message, ProviderQueryTestAdapter,
};
use self::capabilities::{
provider_query_openai_image_normalize_failure_message,
provider_query_openai_image_normalize_options,
};
use self::model_mapping::{
provider_query_resolve_explicit_mapped_effective_model,
provider_query_resolve_global_effective_model,
@@ -530,12 +535,20 @@ fn provider_query_build_test_request_body_for_route(
provider_query_build_test_request_body_with_model_policy(payload, model, override_custom_model)
}
fn provider_query_build_test_request_body_with_model_policy(
fn provider_query_build_test_request_body_for_api_format(
payload: &Value,
model: &str,
override_custom_model: bool,
route_path: &str,
client_api_format: &str,
) -> Value {
let client_api_format = provider_query_normalize_api_format_alias(client_api_format);
let override_custom_model = route_path.ends_with("/test-model-failover")
|| provider_query_extract_mapped_model_name(payload).is_some();
if let Some(mut body) = provider_query_extract_request_body(payload) {
let has_conversation = provider_query_request_body_has_conversation_for_api_format(
&body,
client_api_format.as_str(),
);
if let Some(object) = body.as_object_mut() {
if override_custom_model {
object.insert("model".to_string(), Value::String(model.to_string()));
@@ -544,6 +557,123 @@ fn provider_query_build_test_request_body_with_model_policy(
.entry("model".to_string())
.or_insert_with(|| Value::String(model.to_string()));
}
if !has_conversation {
provider_query_insert_default_test_conversation(
object,
client_api_format.as_str(),
payload,
);
}
}
return body;
}
let message = provider_query_extract_message(payload)
.unwrap_or_else(|| DEFAULT_PROVIDER_QUERY_TEST_MESSAGE.to_string());
match client_api_format.as_str() {
"openai:responses" | "openai:responses:compact" => json!({
"model": model,
"input": message,
"max_output_tokens": 30,
"temperature": 0.7,
"stream": true,
}),
"claude:messages" => json!({
"model": model,
"messages": [{
"role": "user",
"content": message
}],
"max_tokens": 30,
"temperature": 0.7,
"stream": true,
}),
_ => json!({
"model": model,
"messages": [{
"role": "user",
"content": message
}],
"max_tokens": 30,
"temperature": 0.7,
"stream": true,
}),
}
}
fn provider_query_build_grok_test_request_body_for_api_format(
payload: &Value,
model: &str,
route_path: &str,
client_api_format: &str,
) -> Value {
provider_query_build_test_request_body_for_api_format(
payload,
model,
route_path,
client_api_format,
)
}
fn provider_query_insert_default_test_conversation(
object: &mut Map<String, Value>,
client_api_format: &str,
payload: &Value,
) {
let message = provider_query_extract_message(payload)
.unwrap_or_else(|| DEFAULT_PROVIDER_QUERY_TEST_MESSAGE.to_string());
match client_api_format {
"openai:responses" | "openai:responses:compact" => {
object.insert("input".to_string(), Value::String(message));
}
"claude:messages" => {
object.insert(
"messages".to_string(),
json!([{ "role": "user", "content": message }]),
);
}
_ => {
object.insert(
"messages".to_string(),
json!([{ "role": "user", "content": message }]),
);
}
}
}
fn provider_query_grok_test_client_api_format(provider_api_format: &str) -> &'static str {
match provider_query_normalize_api_format_alias(provider_api_format).as_str() {
"openai:responses" | "openai:responses:compact" => "openai:responses",
"claude:messages" => "claude:messages",
_ => "openai:chat",
}
}
fn provider_query_build_test_request_body_with_model_policy(
payload: &Value,
model: &str,
override_custom_model: bool,
) -> Value {
if let Some(mut body) = provider_query_extract_request_body(payload) {
let has_conversation = provider_query_request_body_has_conversation(&body);
if let Some(object) = body.as_object_mut() {
if override_custom_model {
object.insert("model".to_string(), Value::String(model.to_string()));
} else {
object
.entry("model".to_string())
.or_insert_with(|| Value::String(model.to_string()));
}
if !has_conversation {
object.insert(
"messages".to_string(),
json!([{
"role": "user",
"content": provider_query_extract_message(payload)
.unwrap_or_else(|| DEFAULT_PROVIDER_QUERY_TEST_MESSAGE.to_string())
}]),
);
}
}
return body;
}
@@ -561,6 +691,73 @@ fn provider_query_build_test_request_body_with_model_policy(
})
}
fn provider_query_request_body_has_conversation(body: &Value) -> bool {
body.get("messages")
.and_then(Value::as_array)
.map(|messages| {
messages
.iter()
.any(|message| value_has_non_empty_text(message.get("content")))
})
.unwrap_or(false)
|| value_has_non_empty_text(body.get("input"))
|| value_has_non_empty_text(body.get("prompt"))
|| value_has_non_empty_text(body.get("query"))
|| value_has_non_empty_text(body.get("system"))
}
fn provider_query_request_body_has_conversation_for_api_format(
body: &Value,
client_api_format: &str,
) -> bool {
match provider_query_normalize_api_format_alias(client_api_format).as_str() {
"openai:responses" | "openai:responses:compact" => {
value_has_non_empty_text(body.get("input"))
|| value_has_non_empty_text(body.get("prompt"))
}
"claude:messages" => {
body.get("messages")
.and_then(Value::as_array)
.map(|messages| {
messages
.iter()
.any(|message| value_has_non_empty_text(message.get("content")))
})
.unwrap_or(false)
|| value_has_non_empty_text(body.get("system"))
}
_ => provider_query_request_body_has_conversation(body),
}
}
fn provider_query_request_body_is_openai_responses_shape(body: &Value) -> bool {
let Some(object) = body.as_object() else {
return false;
};
[
"input",
"tools",
"tool_choice",
"instructions",
"previous_response_id",
]
.iter()
.any(|key| object.contains_key(*key))
}
fn value_has_non_empty_text(value: Option<&Value>) -> bool {
match value {
Some(Value::String(value)) => !value.trim().is_empty(),
Some(Value::Array(values)) => values
.iter()
.any(|value| value_has_non_empty_text(Some(value))),
Some(Value::Object(values)) => values
.values()
.any(|value| value_has_non_empty_text(Some(value))),
_ => false,
}
}
fn provider_query_request_body_model<'a>(request_body: &'a Value, fallback: &'a str) -> &'a str {
request_body
.get("model")
@@ -1562,6 +1759,32 @@ fn provider_query_chatgpt_web_image_internal_url(base_url: &str) -> String {
format!("{base_url}/__aether/chatgpt-web-image")
}
fn provider_query_openai_image_test_upstream_url(
transport: &AdminGatewayProviderTransportSnapshot,
request_query: Option<&str>,
) -> String {
if transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("chatgpt_web")
{
provider_query_chatgpt_web_image_internal_url(&transport.endpoint.base_url)
} else if transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("grok")
{
crate::provider_transport::build_grok_upstream_url(
transport,
crate::provider_transport::GROK_CHAT_PATH,
)
} else {
crate::provider_transport::build_openai_image_upstream_url(transport, request_query)
}
}
async fn provider_query_finalize_openai_image_result(
route_path: &str,
trace_id: &str,
@@ -1661,12 +1884,16 @@ async fn provider_query_execute_openai_image_test_candidate(
*synthetic_request.headers_mut() = incoming_request_headers;
let (parts, _) = synthetic_request.into_parts();
let Some(normalized_request) =
crate::ai_serving::normalize_openai_image_request(&parts, &request_body, None)
else {
let provider_type = transport.provider.provider_type.as_str();
let Some(normalized_request) = crate::ai_serving::normalize_openai_image_request_with_options(
&parts,
&request_body,
None,
provider_query_openai_image_normalize_options(provider_type),
) else {
return Ok(provider_query_skipped_execution_outcome(
request_body.clone(),
"Provider request body could not be normalized for openai:image",
provider_query_openai_image_normalize_failure_message(provider_type, &request_body),
));
};
@@ -1675,6 +1902,11 @@ async fn provider_query_execute_openai_image_test_candidate(
.provider_type
.trim()
.eq_ignore_ascii_case("chatgpt_web");
let is_grok = transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("grok");
let mut provider_request_body = if is_chatgpt_web {
match crate::ai_serving::build_chatgpt_web_image_request_body(&parts, &request_body, None) {
Ok(body) => body,
@@ -1702,17 +1934,33 @@ async fn provider_query_execute_openai_image_test_candidate(
"Provider auth is unavailable for openai:image",
));
};
let transport_profile = state.resolve_transport_profile(&transport);
let Some(mut request_headers) = crate::provider_transport::build_openai_image_headers(
crate::provider_transport::ProviderOpenAiImageHeadersInput {
headers: &parts.headers,
auth_header: &auth_header,
auth_value: &auth_value,
header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body,
original_request_body: &request_body,
},
) else {
let Some(mut request_headers) = (if is_grok {
crate::provider_transport::build_grok_browser_headers(
crate::provider_transport::GrokHeaderInput {
transport: &transport,
transport_profile: transport_profile.as_ref(),
request_headers: Some(&parts.headers),
content_type: "application/json",
accept: "*/*",
header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body,
original_request_body: &request_body,
},
)
} else {
crate::provider_transport::build_openai_image_headers(
crate::provider_transport::ProviderOpenAiImageHeadersInput {
headers: &parts.headers,
auth_header: &auth_header,
auth_value: &auth_value,
header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body,
original_request_body: &request_body,
},
)
}) else {
return Ok(ProviderQueryExecutionOutcome {
status: "failed",
skip_reason: None,
@@ -1728,6 +1976,7 @@ async fn provider_query_execute_openai_image_test_candidate(
};
if is_chatgpt_web {
request_headers.insert("x-aether-chatgpt-web-image".to_string(), "1".to_string());
} else if is_grok {
} else {
crate::ai_serving::apply_codex_openai_responses_special_headers(
&mut request_headers,
@@ -1761,16 +2010,12 @@ async fn provider_query_execute_openai_image_test_candidate(
.filter(|value| !value.is_empty())
.unwrap_or(request_model.as_str())
.to_string();
let image_request = if is_chatgpt_web {
let image_request = if is_chatgpt_web || is_grok {
provider_request_body.clone()
} else {
normalized_request.summary_json.clone()
};
let request_url = if is_chatgpt_web {
provider_query_chatgpt_web_image_internal_url(&transport.endpoint.base_url)
} else {
crate::provider_transport::build_openai_image_upstream_url(&transport, parts.uri.query())
};
let request_url = provider_query_openai_image_test_upstream_url(&transport, parts.uri.query());
let upstream_is_stream = provider_request_body
.get("stream")
.and_then(Value::as_bool)
@@ -1796,11 +2041,27 @@ async fn provider_query_execute_openai_image_test_candidate(
proxy: state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&transport)
.await,
transport_profile: state.resolve_transport_profile(&transport),
transport_profile: transport_profile.clone(),
timeouts: state.resolve_transport_execution_timeouts(&transport),
};
let result = if is_chatgpt_web {
let result = if is_grok {
let report_context = json!({
"client_api_format": "openai:image",
"provider_api_format": "openai:image",
"provider_type": "grok",
"model": request_model,
"mapped_model": mapped_model,
"image_request": image_request.clone(),
});
state
.execute_execution_runtime_sync_plan_with_report_context(
Some(trace_id),
&plan,
Some(&report_context),
)
.await?
} else if is_chatgpt_web {
let report_context = json!({
"client_api_format": "openai:image",
"provider_api_format": "openai:image",
@@ -2115,6 +2376,177 @@ async fn provider_query_execute_antigravity_test_candidate(
})
}
async fn provider_query_execute_grok_test_candidate(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
candidate: &ProviderQueryTestCandidate,
payload: &Value,
route_path: &str,
trace_id: &str,
) -> Result<ProviderQueryExecutionOutcome, GatewayError> {
let Some(transport) = state
.read_provider_transport_snapshot(&provider.id, &candidate.endpoint.id, &candidate.key.id)
.await?
else {
return Ok(provider_query_skipped_execution_outcome(
Value::Null,
"Provider transport snapshot is unavailable",
));
};
let provider_api_format =
provider_query_normalize_api_format_alias(&candidate.endpoint.api_format);
let client_api_format = provider_query_grok_test_client_api_format(&provider_api_format);
let request_body = provider_query_build_grok_test_request_body_for_api_format(
payload,
&candidate.effective_model,
route_path,
client_api_format,
);
if let Some(reason) =
provider_query_grok_test_unsupported_reason(&transport, &provider_api_format)
{
return Ok(provider_query_skipped_execution_outcome(
request_body,
format!(
"{} ({reason})",
provider_query_unsupported_test_api_format_message(&candidate.endpoint.api_format)
),
));
}
let incoming_request_headers = provider_query_extract_request_headers(payload);
let mut synthetic_request = http::Request::builder()
.uri(route_path)
.body(())
.map_err(|err| GatewayError::Internal(err.to_string()))?;
*synthetic_request.headers_mut() = incoming_request_headers;
let (parts, _) = synthetic_request.into_parts();
let request_model =
provider_query_request_body_model(&request_body, &candidate.effective_model);
let request_url = crate::provider_transport::build_grok_upstream_url(
&transport,
crate::provider_transport::GROK_CHAT_PATH,
);
let provider_request_body = crate::provider_transport::build_grok_app_chat_body(
client_api_format,
Some(request_model),
&request_body,
);
let report_context = json!({
"provider_type": provider.provider_type,
"provider_api_format": provider_api_format,
"client_api_format": client_api_format,
"model": request_model,
"mapped_model": candidate.effective_model,
"request_path": route_path,
"request_body": request_body,
});
let transport_profile = state.resolve_transport_profile(&transport);
let Some(request_headers) = crate::provider_transport::build_grok_browser_headers(
crate::provider_transport::GrokHeaderInput {
transport: &transport,
transport_profile: transport_profile.as_ref(),
request_headers: Some(&parts.headers),
content_type: "application/json",
accept: "text/event-stream",
header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body,
original_request_body: &request_body,
},
) else {
return Ok(ProviderQueryExecutionOutcome {
status: "failed",
skip_reason: None,
error_message: Some("provider request headers build failed".to_string()),
status_code: None,
latency_ms: None,
request_url,
request_headers: BTreeMap::new(),
request_body: provider_request_body,
response_headers: BTreeMap::new(),
response_body: None,
});
};
let plan = ExecutionPlan {
request_id: trace_id.to_string(),
candidate_id: Some(format!("provider-query-{}", candidate.key.id)),
provider_name: Some(provider.name.clone()),
provider_id: provider.id.clone(),
endpoint_id: candidate.endpoint.id.clone(),
key_id: candidate.key.id.clone(),
method: "POST".to_string(),
url: request_url.clone(),
headers: request_headers.clone(),
content_type: Some("application/json".to_string()),
content_encoding: None,
body: RequestBody::from_json(request_body.clone()),
stream: true,
client_api_format: client_api_format.to_string(),
provider_api_format: provider_api_format.clone(),
model_name: Some(request_model.to_string()),
proxy: state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&transport)
.await,
transport_profile,
timeouts: state.resolve_transport_execution_timeouts(&transport),
};
let result = match state
.execute_execution_runtime_sync_plan_with_report_context(
Some(trace_id),
&plan,
Some(&report_context),
)
.await
{
Ok(result) => result,
Err(err) => {
return Ok(ProviderQueryExecutionOutcome {
status: "failed",
skip_reason: None,
error_message: Some(format!("model test execution failed: {err:?}")),
status_code: None,
latency_ms: None,
request_url,
request_headers,
request_body: provider_request_body,
response_headers: BTreeMap::new(),
response_body: None,
});
}
};
let response_body = result.body.as_ref().and_then(|body| body.json_body.clone());
let did_fail = result.status_code >= 400 || response_body.is_none();
let error_message = if did_fail {
provider_query_extract_error_message(&result).or_else(|| {
response_body.is_none().then(|| {
format!(
"Provider returned HTTP {} without a model-test response body",
result.status_code
)
})
})
} else {
None
};
Ok(ProviderQueryExecutionOutcome {
status: if did_fail { "failed" } else { "success" },
skip_reason: None,
error_message,
status_code: Some(result.status_code),
latency_ms: result.telemetry.as_ref().and_then(|value| value.elapsed_ms),
request_url,
request_headers,
request_body: provider_request_body,
response_headers: result.headers,
response_body,
})
}
async fn provider_query_execute_standard_test_candidate(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
@@ -2132,22 +2564,25 @@ async fn provider_query_execute_standard_test_candidate(
"Provider transport snapshot is unavailable",
));
};
let original_request_body = provider_query_build_test_request_body_for_route(
let provider_api_format = candidate.endpoint.api_format.as_str();
let normalized_provider_api_format =
crate::ai_serving::normalize_api_format_alias(provider_api_format);
let client_api_format =
provider_query_standard_test_client_api_format(normalized_provider_api_format.as_str());
let original_request_body = provider_query_build_test_request_body_for_api_format(
payload,
&candidate.effective_model,
route_path,
client_api_format,
);
if !provider_query_transport_supports_model_test_execution(
state,
&transport,
candidate.endpoint.api_format.as_str(),
provider_api_format,
) {
return Ok(provider_query_skipped_execution_outcome(
original_request_body,
provider_query_standard_test_unsupported_reason(
&transport,
candidate.endpoint.api_format.as_str(),
),
provider_query_standard_test_unsupported_reason(&transport, provider_api_format),
));
}
@@ -2159,11 +2594,6 @@ async fn provider_query_execute_standard_test_candidate(
let request_model =
provider_query_request_body_model(&request_body, &candidate.effective_model);
let provider_api_format = candidate.endpoint.api_format.as_str();
let normalized_provider_api_format =
crate::ai_serving::normalize_api_format_alias(provider_api_format);
let client_api_format =
provider_query_standard_test_client_api_format(normalized_provider_api_format.as_str());
let upstream_is_stream = provider_query_resolve_standard_test_upstream_is_stream(
transport.endpoint.config.as_ref(),
transport.provider.provider_type.as_str(),
@@ -2229,12 +2659,20 @@ async fn provider_query_execute_standard_test_candidate(
}
"openai:responses" | "openai:responses:compact" => {
let Some(mut provider_request_body) =
crate::ai_serving::build_cross_format_openai_chat_request_body(
&request_body,
request_model,
normalized_provider_api_format.as_str(),
upstream_is_stream,
)
(if provider_query_request_body_is_openai_responses_shape(&request_body) {
crate::ai_serving::build_local_openai_responses_request_body(
&request_body,
request_model,
upstream_is_stream,
)
} else {
crate::ai_serving::build_cross_format_openai_chat_request_body(
&request_body,
request_model,
normalized_provider_api_format.as_str(),
upstream_is_stream,
)
})
else {
return Ok(provider_query_skipped_execution_outcome(
request_body.clone(),
@@ -2659,6 +3097,12 @@ async fn build_admin_provider_query_kiro_failover_response(
)
.await
}
Some(ProviderQueryTestAdapter::Grok) => {
provider_query_execute_grok_test_candidate(
state, &provider, candidate, payload, route_path, &trace_id,
)
.await
}
Some(ProviderQueryTestAdapter::Standard) => {
provider_query_execute_standard_test_candidate(
state, &provider, candidate, payload, route_path, &trace_id,
@@ -2718,7 +3162,10 @@ async fn build_admin_provider_query_kiro_failover_response(
));
if is_success {
success_body = response_body;
success_stream = matches!(adapter, Some(ProviderQueryTestAdapter::Kiro));
success_stream = matches!(
adapter,
Some(ProviderQueryTestAdapter::Kiro | ProviderQueryTestAdapter::Grok)
);
winning_candidate_index = Some(candidate_index);
break;
}

View File

@@ -10,6 +10,7 @@ use serde_json::{json, Value};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum ProviderQueryTestAdapter {
Standard,
Grok,
Kiro,
OpenAiImage,
Antigravity,
@@ -134,6 +135,65 @@ pub(super) fn provider_query_antigravity_test_unsupported_reason(
}
}
}
pub(super) fn provider_query_grok_test_unsupported_reason(
transport: &AdminGatewayProviderTransportSnapshot,
api_format: &str,
) -> Option<&'static str> {
if !transport.provider.is_active {
return Some("provider_inactive");
}
if !transport.endpoint.is_active {
return Some("endpoint_inactive");
}
if !transport.key.is_active {
return Some("key_inactive");
}
if !transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("grok")
{
return Some("transport_provider_type_unsupported");
}
let normalized_api_format = provider_query_normalize_api_format_alias(api_format);
if !matches!(
normalized_api_format.as_str(),
"openai:chat" | "openai:responses" | "openai:responses:compact" | "claude:messages"
) {
return Some("transport_api_format_mismatch");
}
if provider_query_normalize_api_format_alias(&transport.endpoint.api_format)
!= normalized_api_format
{
return Some("transport_api_format_mismatch");
}
if crate::provider_transport::resolve_grok_session_auth(transport).is_none() {
return Some("transport_oauth_resolution_unsupported");
}
if !crate::provider_transport::header_rules_are_locally_supported(
transport.endpoint.header_rules.as_ref(),
) {
return Some("transport_header_rules_unsupported");
}
if !crate::provider_transport::body_rules_are_locally_supported(
transport.endpoint.body_rules.as_ref(),
) {
return Some("transport_body_rules_unsupported");
}
if !crate::provider_transport::transport_proxy_is_locally_supported(transport) {
return Some("transport_proxy_unsupported");
}
if crate::provider_transport::transport_profile_is_configured(transport)
&& crate::provider_transport::resolve_transport_profile(transport).is_none()
{
return Some("transport_profile_unsupported");
}
None
}
pub(super) fn provider_query_normalize_api_format_alias(value: &str) -> String {
crate::ai_serving::normalize_api_format_alias(value)
}
@@ -147,6 +207,15 @@ pub(super) fn provider_query_test_adapter_for_provider_api_format(
}
let normalized_api_format = provider_query_normalize_api_format_alias(api_format);
if provider_type.trim().eq_ignore_ascii_case("grok") {
return match normalized_api_format.as_str() {
"openai:chat" | "openai:responses" | "openai:responses:compact" | "claude:messages" => {
Some(ProviderQueryTestAdapter::Grok)
}
"openai:image" => Some(ProviderQueryTestAdapter::OpenAiImage),
_ => None,
};
}
if normalized_api_format == "openai:image" {
return Some(ProviderQueryTestAdapter::OpenAiImage);
}
@@ -182,6 +251,16 @@ pub(super) fn provider_query_model_test_endpoint_priority(
let normalized_api_format = provider_query_normalize_api_format_alias(api_format);
match provider_query_test_adapter_for_provider_api_format(provider_type, api_format)? {
ProviderQueryTestAdapter::Kiro => Some(0),
ProviderQueryTestAdapter::Grok => {
if matches!(
normalized_api_format.as_str(),
"openai:chat" | "openai:responses" | "openai:responses:compact" | "claude:messages"
) {
Some(0)
} else {
Some(2)
}
}
ProviderQueryTestAdapter::Antigravity => Some(1),
ProviderQueryTestAdapter::OpenAiImage => Some(2),
ProviderQueryTestAdapter::Standard => {
@@ -234,6 +313,9 @@ pub(super) fn provider_query_transport_supports_model_test_execution(
)
.is_none()
}
Some(ProviderQueryTestAdapter::Grok) => {
provider_query_grok_test_unsupported_reason(transport, api_format).is_none()
}
Some(ProviderQueryTestAdapter::Standard) => match crate::ai_serving::normalize_api_format_alias(api_format).as_str() {
"openai:chat" => {
crate::provider_transport::policy::supports_local_openai_chat_transport(transport)

View File

@@ -0,0 +1,50 @@
use crate::handlers::admin::provider::shared::model_test_capabilities::{
admin_provider_openai_image_normalize_options, admin_provider_openai_image_test_capability,
AdminProviderOpenAiImageTestCapability,
};
use serde_json::Value;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) struct ProviderQueryOpenAiImageTestCapability(AdminProviderOpenAiImageTestCapability);
pub(super) fn provider_query_openai_image_test_capability(
provider_type: &str,
) -> ProviderQueryOpenAiImageTestCapability {
ProviderQueryOpenAiImageTestCapability(admin_provider_openai_image_test_capability(
provider_type,
))
}
pub(super) fn provider_query_openai_image_normalize_options(
provider_type: &str,
) -> crate::ai_serving::OpenAiImageNormalizeOptions {
admin_provider_openai_image_normalize_options(provider_type)
}
pub(super) fn provider_query_openai_image_requested_count(request_body: &Value) -> Option<u64> {
request_body.get("n").and_then(|value| {
value.as_u64().or_else(|| {
value
.as_str()
.map(str::trim)
.filter(|value| !value.is_empty())
.and_then(|value| value.parse::<u64>().ok())
})
})
}
pub(super) fn provider_query_openai_image_normalize_failure_message(
provider_type: &str,
request_body: &Value,
) -> String {
let capability = provider_query_openai_image_test_capability(provider_type);
if provider_query_openai_image_requested_count(request_body)
.is_some_and(|value| !capability.0.supports_generation_count(value))
{
return format!(
"Provider request body could not be normalized for openai:image: selected provider supports n=1..{} for generation",
capability.0.max_generation_count
);
}
"Provider request body could not be normalized for openai:image".to_string()
}

View File

@@ -1,3 +1,5 @@
use std::collections::BTreeMap;
use super::super::provider_query_key_display_name;
use super::{ProviderQueryExecutionOutcome, ProviderQueryTestCandidate};
use serde_json::{json, Value};
@@ -22,13 +24,43 @@ pub(super) fn provider_query_test_attempt_payload(
"status_code": execution.status_code,
"latency_ms": execution.latency_ms,
"request_url": execution.request_url,
"request_headers": execution.request_headers,
"request_headers": provider_query_redact_diagnostic_headers(&execution.request_headers),
"request_body": execution.request_body,
"response_headers": execution.response_headers,
"response_headers": provider_query_redact_diagnostic_headers(&execution.response_headers),
"response_body": execution.response_body,
})
}
fn provider_query_redact_diagnostic_headers(
headers: &BTreeMap<String, String>,
) -> BTreeMap<String, String> {
headers
.iter()
.map(|(name, value)| {
if provider_query_header_is_sensitive(name) {
(name.clone(), "<redacted>".to_string())
} else {
(name.clone(), value.clone())
}
})
.collect()
}
fn provider_query_header_is_sensitive(name: &str) -> bool {
matches!(
name.trim().to_ascii_lowercase().as_str(),
"authorization"
| "proxy-authorization"
| "cookie"
| "set-cookie"
| "x-api-key"
| "api-key"
| "x-goog-api-key"
| "anthropic-api-key"
| "openai-api-key"
)
}
pub(super) fn provider_query_candidate_summary_payload(
total_candidates: usize,
total_attempts: usize,
@@ -133,3 +165,37 @@ pub(super) fn provider_query_candidate_summary_payload(
.unwrap_or(Value::Null),
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn provider_query_diagnostic_headers_redact_credentials() {
let headers = BTreeMap::from([
("cookie".to_string(), "sso=secret".to_string()),
("authorization".to_string(), "Bearer secret".to_string()),
("x-goog-api-key".to_string(), "secret".to_string()),
("content-type".to_string(), "application/json".to_string()),
]);
let redacted = provider_query_redact_diagnostic_headers(&headers);
assert_eq!(
redacted.get("cookie").map(String::as_str),
Some("<redacted>")
);
assert_eq!(
redacted.get("authorization").map(String::as_str),
Some("<redacted>")
);
assert_eq!(
redacted.get("x-goog-api-key").map(String::as_str),
Some("<redacted>")
);
assert_eq!(
redacted.get("content-type").map(String::as_str),
Some("application/json")
);
}
}

View File

@@ -1,6 +1,68 @@
use super::*;
use crate::handlers::admin::request::AdminGatewayProviderTransportSnapshot;
use serde_json::json;
fn sample_openai_image_transport(provider_type: &str) -> AdminGatewayProviderTransportSnapshot {
AdminGatewayProviderTransportSnapshot {
provider: crate::provider_transport::snapshot::GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Provider".to_string(),
provider_type: provider_type.to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: crate::provider_transport::snapshot::GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "openai:image".to_string(),
api_family: None,
endpoint_kind: None,
is_active: true,
base_url: "https://grok.com/".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: crate::provider_transport::snapshot::GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "oauth".to_string(),
is_active: true,
api_formats: None,
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
decrypted_api_key: String::new(),
decrypted_auth_config: Some(
json!({
"sso_token": "abc",
"sso_rw_token": "rw"
})
.to_string(),
),
},
}
}
#[test]
fn provider_query_test_request_body_preserves_custom_model() {
let payload = json!({
@@ -28,6 +90,41 @@ fn provider_query_test_request_body_defaults_missing_model() {
assert_eq!(body["model"], json!("fallback-model"));
}
#[test]
fn provider_query_test_request_body_fills_empty_conversation() {
let payload = json!({
"request_body": {
"model": "custom-upstream-model",
"messages": []
}
});
let body = provider_query_build_test_request_body(&payload, "fallback-model");
assert_eq!(body["model"], json!("custom-upstream-model"));
assert_eq!(
body["messages"],
json!([{ "role": "user", "content": DEFAULT_PROVIDER_QUERY_TEST_MESSAGE }])
);
}
#[test]
fn provider_query_test_request_body_keeps_non_empty_conversation() {
let payload = json!({
"request_body": {
"model": "custom-upstream-model",
"messages": [{ "role": "user", "content": "custom prompt" }]
}
});
let body = provider_query_build_test_request_body(&payload, "fallback-model");
assert_eq!(
body["messages"],
json!([{ "role": "user", "content": "custom prompt" }])
);
}
#[test]
fn provider_query_failover_request_body_overrides_custom_model() {
let payload = json!({
@@ -168,6 +265,54 @@ fn provider_query_standard_test_aggregates_responses_stream_body() {
assert_eq!(body["output"][0]["content"][0]["text"], json!("Hello"));
}
#[test]
fn provider_query_standard_test_aggregates_responses_image_generation_call() {
let stream_body = concat!(
"event: response.created\n",
"data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_img_123\",\"object\":\"response\",\"model\":\"gpt-5.4-mini\",\"status\":\"in_progress\",\"output\":[]}}\n\n",
"event: response.output_item.done\n",
"data: {\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"id\":\"ig_123\",\"type\":\"image_generation_call\",\"status\":\"completed\",\"output_format\":\"png\",\"result\":\"aGVsbG8=\"}}\n\n",
"event: response.completed\n",
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_img_123\",\"object\":\"response\",\"model\":\"gpt-5.4-mini\",\"status\":\"completed\",\"output\":[]}}\n\n",
);
let result = aether_contracts::ExecutionResult {
request_id: "provider-test".to_string(),
candidate_id: Some("candidate-0".to_string()),
status_code: 200,
headers: BTreeMap::new(),
body: Some(aether_contracts::ResponseBody {
json_body: None,
body_bytes_b64: Some(
base64::engine::general_purpose::STANDARD.encode(stream_body.as_bytes()),
),
}),
telemetry: None,
error: None,
};
let body = provider_query_standard_execution_response_body("openai:responses", &result)
.expect("responses image stream body should aggregate");
assert_eq!(body["output"][0]["type"], json!("image_generation_call"));
assert_eq!(body["output"][0]["result"], json!("aGVsbG8="));
}
#[test]
fn provider_query_responses_test_request_body_defaults_to_responses_input() {
let payload = json!({"message": "hello from responses"});
let body = provider_query_build_test_request_body_for_api_format(
&payload,
"gpt-5.4-mini",
"/api/admin/provider-query/test-model",
"openai:responses",
);
assert_eq!(body["model"], json!("gpt-5.4-mini"));
assert_eq!(body["input"], json!("hello from responses"));
assert!(body.get("messages").is_none());
}
#[test]
fn provider_query_test_adapter_routes_fixed_provider_endpoint_types() {
assert_eq!(
@@ -204,6 +349,22 @@ fn provider_query_test_adapter_routes_fixed_provider_endpoint_types() {
),
Some(ProviderQueryTestAdapter::Antigravity)
);
assert_eq!(
provider_query_test_adapter_for_provider_api_format("grok", "openai:chat"),
Some(ProviderQueryTestAdapter::Grok)
);
assert_eq!(
provider_query_test_adapter_for_provider_api_format("grok", "openai:responses"),
Some(ProviderQueryTestAdapter::Grok)
);
assert_eq!(
provider_query_test_adapter_for_provider_api_format("grok", "claude:messages"),
Some(ProviderQueryTestAdapter::Grok)
);
assert_eq!(
provider_query_test_adapter_for_provider_api_format("grok", "openai:image"),
Some(ProviderQueryTestAdapter::OpenAiImage)
);
assert_eq!(
provider_query_test_adapter_for_provider_api_format("custom", "openai:embedding"),
Some(ProviderQueryTestAdapter::Standard)
@@ -244,12 +405,154 @@ fn provider_query_endpoint_priority_prefers_text_before_cli_and_image() {
provider_query_model_test_endpoint_priority("chatgpt_web", "openai:image"),
Some(2)
);
assert_eq!(
provider_query_model_test_endpoint_priority("grok", "openai:chat"),
Some(0)
);
assert_eq!(
provider_query_model_test_endpoint_priority("grok", "openai:responses"),
Some(0)
);
assert_eq!(
provider_query_model_test_endpoint_priority("antigravity", "gemini:generate_content"),
Some(1)
);
}
#[test]
fn provider_query_grok_model_test_body_maps_non_reasoning_model_to_fast_mode() {
let payload = json!({
"request_body": {
"model": "grok-4.20-0309-non-reasoning",
"messages": [
{"role": "system", "content": "be concise"},
{"role": "user", "content": "hello"}
]
}
});
let request_body = provider_query_build_test_request_body_for_route(
&payload,
"grok-4.20-0309-non-reasoning",
"/api/admin/provider-query/test-model",
);
let upstream_body = crate::provider_transport::build_grok_app_chat_body(
"openai:chat",
Some(provider_query_request_body_model(
&request_body,
"grok-4.20-0309-non-reasoning",
)),
&request_body,
);
assert_eq!(upstream_body["modeId"], json!("fast"));
assert_eq!(
upstream_body["message"],
json!("[system]: be concise\n\n[user]: hello")
);
}
#[test]
fn provider_query_grok_model_test_uses_responses_client_body_for_responses_endpoint() {
let payload = json!({
"request_body": {
"model": "grok-4.20-0309-non-reasoning",
"input": "hello from responses body"
}
});
let request_body = provider_query_build_grok_test_request_body_for_api_format(
&payload,
"grok-4.20-0309-non-reasoning",
"/api/admin/provider-query/test-model",
"openai:responses",
);
let upstream_body = crate::provider_transport::build_grok_app_chat_body(
provider_query_grok_test_client_api_format("openai:responses"),
Some(provider_query_request_body_model(
&request_body,
"grok-4.20-0309-non-reasoning",
)),
&request_body,
);
assert_eq!(upstream_body["modeId"], json!("fast"));
assert_eq!(upstream_body["message"], json!("hello from responses body"));
}
#[test]
fn provider_query_grok_model_test_uses_responses_input_when_existing_body_has_messages() {
let payload = json!({
"request_body": {
"model": "grok-4.20-0309-non-reasoning",
"messages": [{
"role": "user",
"content": "hello from stale chat body"
}]
}
});
let request_body = provider_query_build_grok_test_request_body_for_api_format(
&payload,
"grok-4.20-0309-non-reasoning",
"/api/admin/provider-query/test-model",
"openai:responses",
);
assert_eq!(request_body["model"], json!("grok-4.20-0309-non-reasoning"));
assert_eq!(
request_body["input"],
json!("Hello! This is a test message.")
);
assert!(request_body.get("messages").is_some());
let upstream_body = crate::provider_transport::build_grok_app_chat_body(
provider_query_grok_test_client_api_format("openai:responses"),
Some(provider_query_request_body_model(
&request_body,
"grok-4.20-0309-non-reasoning",
)),
&request_body,
);
assert_eq!(upstream_body["modeId"], json!("fast"));
assert_eq!(
upstream_body["message"],
json!("Hello! This is a test message.")
);
}
#[test]
fn provider_query_grok_model_test_defaults_claude_messages_body_for_claude_endpoint() {
let payload = json!({});
let request_body = provider_query_build_grok_test_request_body_for_api_format(
&payload,
"grok-4.20-0309-non-reasoning",
"/api/admin/provider-query/test-model",
"claude:messages",
);
assert_eq!(request_body["model"], json!("grok-4.20-0309-non-reasoning"));
assert_eq!(
request_body["messages"],
json!([{ "role": "user", "content": DEFAULT_PROVIDER_QUERY_TEST_MESSAGE }])
);
let upstream_body = crate::provider_transport::build_grok_app_chat_body(
provider_query_grok_test_client_api_format("claude:messages"),
Some(provider_query_request_body_model(
&request_body,
"grok-4.20-0309-non-reasoning",
)),
&request_body,
);
assert_eq!(upstream_body["modeId"], json!("fast"));
assert_eq!(
upstream_body["message"],
json!("[user]: Hello! This is a test message.")
);
}
#[test]
fn provider_query_candidate_summary_marks_unused_after_first_success() {
let attempts = vec![json!({
@@ -360,3 +663,76 @@ fn provider_query_failover_image_test_request_body_overrides_model() {
assert_eq!(body["model"], json!("new-image-model"));
}
#[test]
fn provider_query_grok_image_test_allows_multi_generation_count() {
let request = http::Request::builder()
.uri("/v1/images/generations")
.body(())
.expect("request should build");
let (parts, _) = request.into_parts();
let body = json!({
"model": "grok-imagine-image",
"prompt": "draw",
"n": 2
});
let normalized = crate::ai_serving::normalize_openai_image_request_with_options(
&parts,
&body,
None,
provider_query_openai_image_normalize_options("grok"),
)
.expect("grok image model tests should allow multi-image generation");
let provider_body = crate::ai_serving::build_openai_image_provider_request_body(&normalized);
assert_eq!(provider_body["n"], json!(2));
}
#[test]
fn provider_query_grok_image_test_uses_grok_app_chat_upstream_url() {
let transport = sample_openai_image_transport("grok");
assert_eq!(
provider_query_openai_image_test_upstream_url(&transport, Some("trace=1")),
"https://grok.com/rest/app-chat/conversations/new"
);
}
#[test]
fn provider_query_chatgpt_web_image_test_uses_internal_upstream_url() {
let transport = sample_openai_image_transport("chatgpt_web");
assert_eq!(
provider_query_openai_image_test_upstream_url(&transport, Some("trace=1")),
"https://grok.com/__aether/chatgpt-web-image"
);
}
#[test]
fn provider_query_non_grok_image_test_keeps_single_generation_boundary() {
let request = http::Request::builder()
.uri("/v1/images/generations")
.body(())
.expect("request should build");
let (parts, _) = request.into_parts();
let body = json!({
"model": "gpt-image-2",
"prompt": "draw",
"n": 2
});
assert!(
crate::ai_serving::normalize_openai_image_request_with_options(
&parts,
&body,
None,
provider_query_openai_image_normalize_options("chatgpt_web"),
)
.is_none()
);
assert_eq!(
provider_query_openai_image_normalize_failure_message("chatgpt_web", &body),
"Provider request body could not be normalized for openai:image: selected provider supports n=1..1 for generation"
);
}

View File

@@ -1,3 +1,4 @@
pub(crate) mod model_test_capabilities;
pub(crate) mod paths;
pub(crate) mod payloads;
pub(crate) mod support;

View File

@@ -0,0 +1,121 @@
use crate::image_capabilities::{
openai_image_normalize_options_for_provider, openai_image_provider_max_generation_count,
};
use serde_json::{json, Value};
const GROK_IMAGE_MODEL_IDS: &[&str] = &[
"grok-imagine-image-lite",
"grok-imagine-image",
"grok-imagine-image-pro",
"grok-imagine-image-edit",
];
const GROK_IMAGE_EDIT_MODEL_ID: &str = "grok-imagine-image-edit";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct AdminProviderOpenAiImageTestCapability {
pub(crate) max_generation_count: u64,
}
impl AdminProviderOpenAiImageTestCapability {
pub(crate) fn supports_generation_count(self, count: u64) -> bool {
count >= 1 && count <= self.max_generation_count
}
}
pub(crate) fn admin_provider_openai_image_test_capability(
provider_type: &str,
) -> AdminProviderOpenAiImageTestCapability {
AdminProviderOpenAiImageTestCapability {
max_generation_count: openai_image_provider_max_generation_count(provider_type),
}
}
pub(crate) fn admin_provider_openai_image_normalize_options(
provider_type: &str,
) -> crate::ai_serving::OpenAiImageNormalizeOptions {
openai_image_normalize_options_for_provider(provider_type)
}
pub(crate) fn admin_provider_model_test_capabilities_payload(
provider_type: &str,
model_id: &str,
supports_image_generation: bool,
) -> Value {
let provider_type = provider_type.trim();
let model_id = model_id.trim();
let is_grok_image_edit =
provider_type.eq_ignore_ascii_case("grok") && model_id == GROK_IMAGE_EDIT_MODEL_ID;
let openai_image = if supports_image_generation {
Some(json!({
"max_generation_count": admin_provider_openai_image_test_capability(provider_type).max_generation_count,
"supports_generation": !is_grok_image_edit,
"supports_edit": is_grok_image_edit,
}))
} else {
None
};
json!({
"openai:image": openai_image,
})
}
pub(crate) fn admin_provider_model_supports_image_generation(
provider_type: &str,
model_id: &str,
fallback_supports_image_generation: bool,
) -> bool {
if provider_type.trim().eq_ignore_ascii_case("grok") {
let model_id = model_id.trim();
return GROK_IMAGE_MODEL_IDS
.iter()
.any(|candidate| model_id.eq_ignore_ascii_case(candidate));
}
fallback_supports_image_generation
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn grok_image_generation_models_expose_multi_image_capability() {
let payload =
admin_provider_model_test_capabilities_payload("grok", "grok-imagine-image", true);
assert_eq!(payload["openai:image"]["max_generation_count"], 4);
assert_eq!(payload["openai:image"]["supports_generation"], true);
assert_eq!(payload["openai:image"]["supports_edit"], false);
}
#[test]
fn grok_image_edit_model_is_edit_only_for_generation_tests() {
let payload =
admin_provider_model_test_capabilities_payload("grok", "grok-imagine-image-edit", true);
assert_eq!(payload["openai:image"]["max_generation_count"], 4);
assert_eq!(payload["openai:image"]["supports_generation"], false);
assert_eq!(payload["openai:image"]["supports_edit"], true);
}
#[test]
fn non_image_models_report_null_image_test_capability() {
let payload = admin_provider_model_test_capabilities_payload("openai", "gpt-5.5", false);
assert!(payload["openai:image"].is_null());
}
#[test]
fn grok_image_support_uses_catalog_model_ids_not_global_fallback() {
assert!(admin_provider_model_supports_image_generation(
"grok",
"grok-imagine-image-pro",
false,
));
assert!(!admin_provider_model_supports_image_generation(
"grok",
"grok-4.20-fast",
true,
));
}
}

View File

@@ -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" => Ok(normalized),
| "antigravity" | "vertex_ai" | "grok" => Ok(normalized),
_ => Err(
"provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai"
"provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai / grok"
.to_string(),
),
}
@@ -260,6 +260,14 @@ mod tests {
);
}
#[test]
fn normalize_provider_type_supports_grok() {
assert_eq!(
normalize_provider_type_input(" Grok ").expect("type should normalize"),
"grok"
);
}
#[test]
fn normalize_api_format_list_dedupes_canonical_formats() {
assert_eq!(

View File

@@ -205,7 +205,7 @@ pub(crate) async fn build_admin_update_provider_record(
updated.stream_first_byte_timeout_secs = match payload.stream_first_byte_timeout {
Some(value) if (1.0..=300.0).contains(&value) => Some(value),
Some(_) => {
return Err("stream_first_byte_timeout 必须是 1 到 300 之间的数字".to_string())
return Err("stream_first_byte_timeout 必须是 1 到 300 之间的数字".to_string());
}
None => None,
};

View File

@@ -250,4 +250,19 @@ impl<'a> AdminAppState<'a> {
crate::execution_runtime::execute_execution_runtime_sync_plan(self.app, trace_id, plan)
.await
}
pub(crate) async fn execute_execution_runtime_sync_plan_with_report_context(
&self,
trace_id: Option<&str>,
plan: &aether_contracts::ExecutionPlan,
report_context: Option<&serde_json::Value>,
) -> Result<aether_contracts::ExecutionResult, GatewayError> {
crate::execution_runtime::execute_execution_runtime_sync_plan_with_report_context(
self.app,
trace_id,
plan,
report_context,
)
.await
}
}

View File

@@ -18,8 +18,7 @@ use std::collections::BTreeMap;
use std::io::Read;
use url::Url;
const KIRO_IDC_AMZ_USER_AGENT: &str =
"aws-sdk-js/3.738.0 ua/2.1 os/other lang/js md/browser#unknown_unknown api/sso-oidc#3.738.0 m/E KiroIDE";
const KIRO_IDC_AMZ_USER_AGENT: &str = "aws-sdk-js/3.738.0 ua/2.1 os/other lang/js md/browser#unknown_unknown api/sso-oidc#3.738.0 m/E KiroIDE";
const ADMIN_PROVIDER_OAUTH_TIMEOUT_MS: u64 = 30_000;
const ADMIN_PROVIDER_OAUTH_PROXY_TIMEOUT_MS: u64 = 60_000;

View File

@@ -84,7 +84,7 @@ impl<'a> AdminAppState<'a> {
Json(json!({ "detail": "请求数据验证失败" })),
)
.into_response(),
))
));
}
};

View File

@@ -1197,7 +1197,7 @@ impl<'a> AdminAppState<'a> {
Err(_) => {
return Ok(Err(invalid_request(format!(
"Provider '{provider_name}' 配置格式无效"
))))
))));
}
};
let mut updated = invalid!(
@@ -1230,7 +1230,7 @@ impl<'a> AdminAppState<'a> {
Err(_) => {
return Ok(Err(invalid_request(format!(
"Provider '{provider_name}' 配置格式无效"
))))
))));
}
};
let (mut record, shift_existing_priorities_from) =
@@ -1301,7 +1301,7 @@ impl<'a> AdminAppState<'a> {
Err(_) => {
return Ok(Err(invalid_request(
"Provider Endpoint 配置格式无效",
)))
)));
}
};
let (fields, payload) = patch.into_parts();
@@ -1504,7 +1504,7 @@ impl<'a> AdminAppState<'a> {
) {
Ok(patch) => patch,
Err(_) => {
return Ok(Err(invalid_request("Provider Key 配置格式无效")))
return Ok(Err(invalid_request("Provider Key 配置格式无效")));
}
};
let mut updated = invalid!(
@@ -1985,7 +1985,7 @@ impl<'a> AdminAppState<'a> {
Err(_) => {
return Ok(Err(invalid_request(
"merge_mode 仅支持 skip / overwrite / error",
)))
)));
}
};
let empty = Vec::new();

View File

@@ -1,6 +1,9 @@
use crate::async_task::CancelVideoTaskError;
use crate::control::GatewayControlDecision;
use crate::control::GatewayPublicRequestContext;
use crate::image_capabilities::{
openai_image_gateway_max_generation_count, openai_image_gateway_max_generation_count_for_model,
};
use crate::{AppState, GatewayError};
use aether_data_contracts::repository::video_tasks::{
StoredVideoTask, VideoTaskQueryFilter, VideoTaskStatus,
@@ -18,9 +21,6 @@ const AI_PUBLIC_METHOD_NOT_ALLOWED_DETAIL: &str = "Method not allowed";
const AI_PUBLIC_UNAUTHORIZED_DETAIL: &str = "Unauthorized";
const OPENAI_IMAGE_PROMPT_DETAIL: &str = "图片生成/编辑请求缺少 prompt";
const OPENAI_IMAGE_EDIT_INPUT_DETAIL: &str = "图片编辑请求至少需要 1 张输入图片";
const OPENAI_IMAGE_VARIATION_INPUT_DETAIL: &str = "图片变体请求需要 image 文件";
const OPENAI_IMAGE_N_DETAIL: &str = "当前 Codex 图片反代仅支持 n=1";
const OPENAI_IMAGE_STREAM_VARIATION_DETAIL: &str = "图片变体接口当前仅支持同步响应";
const OPENAI_IMAGE_PARTIAL_IMAGES_DETAIL: &str =
"partial_images 仅支持 0-3且必须配合 stream=true";
const OPENAI_IMAGE_STYLE_DETAIL: &str = "当前 Codex 图片反代暂不支持 style 参数";
@@ -57,7 +57,6 @@ const OPENAI_RERANK_STREAM_UNSUPPORTED_DETAIL: &str = "Rerank requests do not su
enum OpenAiImageOperation {
Generate,
Edit,
Variation,
}
impl OpenAiImageOperation {
@@ -65,7 +64,6 @@ impl OpenAiImageOperation {
match path {
"/v1/images/generations" => Some(Self::Generate),
"/v1/images/edits" => Some(Self::Edit),
"/v1/images/variations" => Some(Self::Variation),
_ => None,
}
}
@@ -202,7 +200,7 @@ fn maybe_build_local_openai_request_validation_response(
if decision.route_kind.as_deref() != Some("image")
|| !matches!(
request_context.request_path.as_str(),
"/v1/images/generations" | "/v1/images/edits" | "/v1/images/variations"
"/v1/images/generations" | "/v1/images/edits"
)
{
return None;
@@ -240,19 +238,13 @@ fn maybe_build_local_openai_request_validation_response(
OPENAI_IMAGE_EDIT_INPUT_DETAIL,
));
}
OpenAiImageOperation::Variation if validation.image_count == 0 => {
return Some(build_ai_public_error_response(
http::StatusCode::BAD_REQUEST,
OPENAI_IMAGE_VARIATION_INPUT_DETAIL,
));
}
_ => {}
}
if validation.n.is_some_and(|value| value != 1) {
if let Some(detail) = validate_openai_image_n(&validation) {
return Some(build_ai_public_error_response(
http::StatusCode::BAD_REQUEST,
OPENAI_IMAGE_N_DETAIL,
detail,
));
}
@@ -272,15 +264,6 @@ fn maybe_build_local_openai_request_validation_response(
));
}
if validation.stream {
if operation == OpenAiImageOperation::Variation {
return Some(build_ai_public_error_response(
http::StatusCode::BAD_REQUEST,
OPENAI_IMAGE_STREAM_VARIATION_DETAIL,
));
}
}
if validation
.response_format
.as_deref()
@@ -360,6 +343,23 @@ fn maybe_build_local_openai_request_validation_response(
None
}
fn openai_image_n_detail(max_generation_count: u64) -> String {
if max_generation_count >= openai_image_gateway_max_generation_count() {
format!("当前图片反代仅支持 n=1..{max_generation_count}")
} else {
format!("当前图片模型仅支持 n=1..{max_generation_count}")
}
}
fn validate_openai_image_n(validation: &OpenAiImageValidationInput) -> Option<String> {
let max_generation_count =
openai_image_gateway_max_generation_count_for_model(validation.model.as_deref());
validation
.n
.is_some_and(|value| value == 0 || value > max_generation_count)
.then(|| openai_image_n_detail(max_generation_count))
}
fn validate_openai_embedding_request(
content_type: Option<&str>,
request_body: &Bytes,
@@ -535,7 +535,6 @@ fn parse_openai_image_validation_input(
OpenAiImageOperation::Generate | OpenAiImageOperation::Edit => {
OPENAI_IMAGE_PROMPT_DETAIL
}
OpenAiImageOperation::Variation => OPENAI_IMAGE_VARIATION_INPUT_DETAIL,
});
}
@@ -1260,7 +1259,8 @@ fn estimate_text_tokens(text: &str) -> u64 {
#[cfg(test)]
mod tests {
use super::{
estimate_claude_count_tokens, parse_openai_image_validation_input, OpenAiImageOperation,
estimate_claude_count_tokens, parse_openai_image_validation_input, validate_openai_image_n,
OpenAiImageOperation,
};
use axum::body::Bytes;
use serde_json::json;
@@ -1344,4 +1344,31 @@ mod tests {
assert_eq!(validation.prompt.as_deref(), Some("edit this image"));
assert_eq!(validation.image_count, 1);
}
#[test]
fn image_validation_restricts_multi_image_count_to_grok_models() {
let openai_body = Bytes::from_static(br#"{"model":"gpt-image-2","prompt":"draw","n":2}"#);
let openai_validation = parse_openai_image_validation_input(
OpenAiImageOperation::Generate,
Some("application/json"),
&openai_body,
)
.expect("valid image payload should parse");
assert_eq!(
validate_openai_image_n(&openai_validation).as_deref(),
Some("当前图片模型仅支持 n=1..1")
);
let grok_body =
Bytes::from_static(br#"{"model":"grok-imagine-image-lite","prompt":"draw","n":4}"#);
let grok_validation = parse_openai_image_validation_input(
OpenAiImageOperation::Generate,
Some("application/json"),
&grok_body,
)
.expect("valid grok image payload should parse");
assert!(validate_openai_image_n(&grok_validation).is_none());
}
}

View File

@@ -1,7 +1,7 @@
use crate::handlers::shared::{json_string_list, unix_secs_to_rfc3339};
use crate::provider_key_auth::{
provider_key_auth_semantics, provider_key_configured_api_formats,
provider_key_inherits_provider_api_formats,
provider_key_auth_semantics, provider_key_can_refresh_oauth,
provider_key_configured_api_formats, provider_key_inherits_provider_api_formats,
};
use crate::AppState;
use aether_admin::provider::quota as admin_provider_quota_pure;
@@ -10,6 +10,9 @@ use aether_admin::provider::status as admin_provider_status_pure;
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
use aether_crypto::{decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use aether_provider_pool::{
grok_pool_tier_from_quota_bucket, grok_supported_quota_windows_for_tier,
};
use serde_json::{json, Map, Value};
use std::borrow::Cow;
use std::time::{SystemTime, UNIX_EPOCH};
@@ -461,6 +464,24 @@ fn model_quota_window_snapshot(
item: &Map<String, Value>,
observed_at_unix_secs: Option<u64>,
) -> Option<Value> {
let remaining_value = item
.get("remaining")
.or_else(|| item.get("remaining_value"))
.and_then(admin_provider_quota_pure::coerce_json_f64);
let limit_value = item
.get("total")
.or_else(|| item.get("limit_value"))
.and_then(admin_provider_quota_pure::coerce_json_f64)
.filter(|value| *value > 0.0);
let used_value = item
.get("used")
.or_else(|| item.get("used_value"))
.and_then(admin_provider_quota_pure::coerce_json_f64)
.or_else(|| {
remaining_value
.zip(limit_value)
.map(|(remaining, limit)| (limit - remaining).max(0.0))
});
let used_ratio = item
.get("used_percent")
.and_then(admin_provider_quota_pure::coerce_json_f64)
@@ -489,6 +510,8 @@ fn model_quota_window_snapshot(
&& reset_at.is_none()
&& reset_seconds.is_none()
&& is_exhausted.is_none()
&& remaining_value.is_none()
&& limit_value.is_none()
{
return None;
}
@@ -507,12 +530,29 @@ fn model_quota_window_snapshot(
window.insert("model".to_string(), json!(model_name));
window.insert("used_ratio".to_string(), json!(used_ratio));
window.insert("remaining_ratio".to_string(), json!(remaining_ratio));
window.insert("used_value".to_string(), json!(used_value));
window.insert("remaining_value".to_string(), json!(remaining_value));
window.insert("limit_value".to_string(), json!(limit_value));
window.insert("reset_at".to_string(), json!(reset_at));
window.insert("reset_seconds".to_string(), json!(reset_seconds));
window.insert("is_exhausted".to_string(), json!(is_exhausted));
Some(Value::Object(window))
}
fn provider_quota_metadata_string(
metadata: &Map<String, Value>,
fields: &[&str],
) -> Option<String> {
fields.iter().find_map(|field| {
metadata
.get(*field)
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
}
fn quota_windows_usage_ratio(windows: &[Value]) -> Option<f64> {
windows
.iter()
@@ -1126,6 +1166,72 @@ fn build_antigravity_quota_status_snapshot(
}))
}
fn build_grok_quota_status_snapshot(
upstream_metadata: Option<&Value>,
source: &str,
) -> Option<Value> {
let metadata = provider_quota_metadata_bucket(upstream_metadata, "grok")?;
let observed_at_unix_secs = provider_quota_timestamp_unix_secs(metadata.get("updated_at"));
let inferred_pool_tier = grok_pool_tier_from_quota_bucket(metadata);
let pool_tier = provider_quota_metadata_string(metadata, &["pool_tier", "tier"])
.or_else(|| inferred_pool_tier.map(ToOwned::to_owned));
let plan_type = provider_quota_metadata_string(metadata, &["plan_type", "plan"])
.or_else(|| pool_tier.clone());
let supported_windows = grok_supported_quota_windows_for_tier(pool_tier.as_deref());
let windows = provider_quota_model_bucket(metadata)
.map(|models| {
models
.iter()
.filter_map(|(model_name, item)| {
if !supported_windows
.iter()
.any(|(quota_key, _)| *quota_key == model_name.as_str())
{
return None;
}
model_quota_window_snapshot(
model_name,
item.as_object()?,
observed_at_unix_secs,
)
})
.collect::<Vec<_>>()
})
.unwrap_or_default();
if windows.is_empty() && observed_at_unix_secs.is_none() {
return None;
}
let usage_ratio = quota_windows_usage_ratio(&windows);
let reset_seconds = quota_windows_min_reset_seconds(&windows);
let reset_at = quota_windows_min_reset_at(&windows);
let exhausted = quota_windows_all_exhausted(&windows);
Some(json!({
"version": 2,
"provider_type": "grok",
"code": if exhausted { "exhausted" } else { "ok" },
"label": if exhausted { Some("额度耗尽") } else { None::<&str> },
"reason": if exhausted {
Some("所有 Grok 模式额度已耗尽")
} else {
None::<&str>
},
"freshness": "fresh",
"source": source,
"observed_at": observed_at_unix_secs,
"exhausted": exhausted,
"usage_ratio": usage_ratio,
"updated_at": observed_at_unix_secs,
"reset_at": reset_at,
"reset_seconds": reset_seconds,
"plan_type": plan_type,
"pool_tier": pool_tier,
"windows": windows,
}))
}
fn build_gemini_cli_quota_status_snapshot(
upstream_metadata: Option<&Value>,
source: &str,
@@ -1228,6 +1334,7 @@ pub(crate) fn sync_provider_key_quota_status_snapshot(
"kiro" => build_kiro_quota_status_snapshot(upstream_metadata, source),
"chatgpt_web" => build_chatgpt_web_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),
_ => None,
}?;
@@ -1590,7 +1697,10 @@ pub(crate) fn build_admin_provider_key_response(
);
payload.insert(
"can_refresh_oauth".to_string(),
json!(auth_semantics.can_refresh_oauth()),
json!(provider_key_can_refresh_oauth(
auth_semantics,
auth_config.as_ref()
)),
);
payload.insert(
"can_export_oauth".to_string(),
@@ -2099,6 +2209,69 @@ mod tests {
assert_eq!(window.get("remaining_ratio"), Some(&json!(0.96)));
}
#[test]
fn provider_key_status_snapshot_payload_backfills_grok_model_quota() {
let mut key = sample_catalog_key();
key.upstream_metadata = Some(json!({
"grok": {
"updated_at": 1_778_067_246u64,
"pool_tier": "heavy",
"plan_type": "heavy",
"quota_by_model": {
"quota_auto": {
"display_name": "auto",
"remaining_fraction": 0.4,
"used_percent": 60.0,
"remaining": 60.0,
"total": 150.0,
"reset_at": 1_778_157_172u64,
"is_exhausted": false
},
"quota_heavy": {
"display_name": "heavy",
"remaining_fraction": 0.0,
"used_percent": 100.0,
"reset_at": 1_778_157_172u64,
"is_exhausted": true
}
}
}
}));
let payload = provider_key_status_snapshot_payload(&key, "grok");
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("grok quota windows should exist");
assert_eq!(quota.get("provider_type"), Some(&json!("grok")));
assert_eq!(quota.get("code"), Some(&json!("ok")));
assert_eq!(quota.get("plan_type"), Some(&json!("heavy")));
assert_eq!(quota.get("pool_tier"), Some(&json!("heavy")));
assert_eq!(quota.get("exhausted"), Some(&json!(false)));
assert_eq!(quota.get("usage_ratio"), Some(&json!(1.0)));
assert_eq!(quota.get("reset_at"), Some(&json!(1_778_157_172u64)));
assert_eq!(windows.len(), 2);
assert!(windows.iter().any(|window| {
window
.get("code")
.and_then(Value::as_str)
.is_some_and(|code| code == "model:quota_auto")
}));
let auto = windows
.iter()
.filter_map(Value::as_object)
.find(|window| window.get("code") == Some(&json!("model:quota_auto")))
.expect("auto quota window should exist");
assert_eq!(auto.get("remaining_value"), Some(&json!(60.0)));
assert_eq!(auto.get("limit_value"), Some(&json!(150.0)));
assert_eq!(auto.get("used_value"), Some(&json!(90.0)));
}
#[test]
fn provider_key_status_snapshot_payload_preserves_existing_materialized_quota_snapshot() {
let mut key = sample_catalog_key();

View File

@@ -0,0 +1,60 @@
use crate::ai_serving::OpenAiImageNormalizeOptions;
const DEFAULT_OPENAI_IMAGE_MAX_GENERATION_COUNT: u64 = 1;
const GROK_OPENAI_IMAGE_MAX_GENERATION_COUNT: u64 = 4;
pub(crate) fn openai_image_gateway_max_generation_count() -> u64 {
GROK_OPENAI_IMAGE_MAX_GENERATION_COUNT
}
pub(crate) fn openai_image_gateway_max_generation_count_for_model(model: Option<&str>) -> u64 {
if model.is_some_and(is_grok_openai_image_model) {
GROK_OPENAI_IMAGE_MAX_GENERATION_COUNT
} else {
DEFAULT_OPENAI_IMAGE_MAX_GENERATION_COUNT
}
}
pub(crate) fn openai_image_provider_max_generation_count(provider_type: &str) -> u64 {
if provider_type.trim().eq_ignore_ascii_case("grok") {
GROK_OPENAI_IMAGE_MAX_GENERATION_COUNT
} else {
DEFAULT_OPENAI_IMAGE_MAX_GENERATION_COUNT
}
}
pub(crate) fn openai_image_normalize_options_for_provider(
provider_type: &str,
) -> OpenAiImageNormalizeOptions {
OpenAiImageNormalizeOptions::with_max_generation_count(
openai_image_provider_max_generation_count(provider_type),
)
}
fn is_grok_openai_image_model(model: &str) -> bool {
model
.trim()
.to_ascii_lowercase()
.contains("grok-imagine-image")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn grok_owns_gateway_wide_image_generation_count_ceiling() {
assert_eq!(openai_image_gateway_max_generation_count(), 4);
assert_eq!(
openai_image_gateway_max_generation_count_for_model(Some("grok-imagine-image-lite")),
4
);
assert_eq!(
openai_image_gateway_max_generation_count_for_model(Some("gpt-image-2")),
1
);
assert_eq!(openai_image_gateway_max_generation_count_for_model(None), 1);
assert_eq!(openai_image_provider_max_generation_count("grok"), 4);
assert_eq!(openai_image_provider_max_generation_count("openai"), 1);
}
}

View File

@@ -45,6 +45,7 @@ mod frontdoor_loop_guard;
mod handlers;
mod headers;
mod hooks;
mod image_capabilities;
mod log_ids;
mod maintenance;
pub(crate) mod middleware;

View File

@@ -3,12 +3,14 @@ use std::sync::{Mutex, OnceLock};
use std::time::{Duration, Instant};
use aether_admin::provider::quota as admin_provider_quota_pure;
use aether_provider_pool::grok_quota_window_key_for_model;
use aether_usage_runtime::{
extract_gemini_file_mapping_entries, gemini_file_mapping_cache_key, normalize_gemini_file_name,
report_request_id, GatewayStreamReportRequest, GatewaySyncReportRequest,
GEMINI_FILE_MAPPING_TTL_SECONDS,
};
use serde_json::Value;
use regex::Regex;
use serde_json::{json, Value};
use tracing::warn;
use uuid::Uuid;
@@ -23,6 +25,8 @@ const CODEX_QUOTA_CACHE_MAX_ENTRIES: usize = 4096;
type HeaderFingerprintCache = Mutex<HashMap<String, (String, Instant)>>;
static CODEX_QUOTA_HEADER_FINGERPRINT_CACHE: OnceLock<HeaderFingerprintCache> = OnceLock::new();
static GROK_CHINESE_WAIT_DURATION_RE: OnceLock<Regex> = OnceLock::new();
static GROK_ENGLISH_WAIT_DURATION_RE: OnceLock<Regex> = OnceLock::new();
#[derive(Debug, Clone, Copy)]
pub(crate) enum LocalReportEffect<'a> {
@@ -167,6 +171,264 @@ fn merge_metadata_object(
Some(Value::Object(merged))
}
fn grok_report_context_model(report_context: Option<&Value>) -> Option<String> {
report_context
.and_then(|context| context.get("mapped_model"))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
fn grok_chinese_wait_duration_regex() -> &'static Regex {
GROK_CHINESE_WAIT_DURATION_RE.get_or_init(|| {
Regex::new(r"(?i)(?:(\d+)\s*天)?\s*(?:(\d+)\s*(?:小时|小時))?\s*(?:(\d+)\s*分钟)?\s*(?:(\d+)\s*秒)?")
.expect("grok Chinese wait duration regex should compile")
})
}
fn grok_english_wait_duration_regex() -> &'static Regex {
GROK_ENGLISH_WAIT_DURATION_RE.get_or_init(|| {
Regex::new(
r"(?i)(?:(\d+)\s*(?:d|day|days))?\s*(?:(\d+)\s*(?:h|hour|hours))?\s*(?:(\d+)\s*(?:m|min|mins|minute|minutes))?\s*(?:(\d+)\s*(?:s|sec|secs|second|seconds))?",
)
.expect("grok English wait duration regex should compile")
})
}
fn grok_duration_capture_seconds(captures: regex::Captures<'_>) -> Option<u64> {
let values = [1usize, 2, 3, 4]
.into_iter()
.map(|index| {
captures
.get(index)
.and_then(|item| item.as_str().parse::<u64>().ok())
.unwrap_or(0)
})
.collect::<Vec<_>>();
let seconds = values[0]
.saturating_mul(86_400)
.saturating_add(values[1].saturating_mul(3_600))
.saturating_add(values[2].saturating_mul(60))
.saturating_add(values[3]);
(seconds > 0).then_some(seconds)
}
fn grok_wait_duration_seconds_from_text(text: &str) -> Option<u64> {
let text = text.trim();
if text.is_empty() {
return None;
}
for captures in grok_chinese_wait_duration_regex().captures_iter(text) {
if let Some(seconds) = grok_duration_capture_seconds(captures) {
return Some(seconds);
}
}
for captures in grok_english_wait_duration_regex().captures_iter(text) {
if let Some(seconds) = grok_duration_capture_seconds(captures) {
return Some(seconds);
}
}
None
}
fn grok_response_error_text(value: &Value) -> Option<String> {
match value {
Value::String(text) => Some(text.trim().to_string()).filter(|text| !text.is_empty()),
Value::Object(object) => {
if let Some(error) = object.get("error") {
if let Some(text) = grok_response_error_text(error) {
return Some(text);
}
}
for key in ["message", "detail", "reason", "error"] {
if let Some(text) = object
.get(key)
.and_then(Value::as_str)
.map(str::trim)
.filter(|text| !text.is_empty())
{
return Some(text.to_string());
}
}
None
}
_ => None,
}
}
fn grok_upstream_response_body(report_context: Option<&Value>) -> Option<&Value> {
report_context
.and_then(|context| context.get("upstream_response"))
.and_then(|response| response.get("body"))
}
fn grok_quota_reset_after_seconds(
body_json: Option<&Value>,
report_context: Option<&Value>,
) -> Option<u64> {
body_json
.and_then(grok_response_error_text)
.and_then(|text| grok_wait_duration_seconds_from_text(text.as_str()))
.or_else(|| {
grok_upstream_response_body(report_context)
.and_then(grok_response_error_text)
.and_then(|text| grok_wait_duration_seconds_from_text(text.as_str()))
})
}
fn grok_apply_quota_feedback(
bucket: &mut serde_json::Map<String, Value>,
model: &str,
status_code: u16,
reset_after_seconds: Option<u64>,
now_unix_secs: u64,
) -> bool {
let Some(quota_key) = grok_quota_window_key_for_model(Some(model)) else {
return false;
};
let quota_by_model = if bucket.contains_key("quota_by_model") {
bucket.get_mut("quota_by_model")
} else {
bucket.get_mut("models")
};
let Some(window) = quota_by_model
.and_then(Value::as_object_mut)
.and_then(|models| models.get_mut(quota_key))
.and_then(Value::as_object_mut)
else {
return false;
};
let total = window
.get("total")
.and_then(admin_provider_quota_pure::coerce_json_f64)
.filter(|value| *value > 0.0);
let current_remaining = window
.get("remaining")
.and_then(admin_provider_quota_pure::coerce_json_f64)
.or(total)
.unwrap_or(0.0);
let next_remaining = match status_code {
429 => 0.0,
code if code >= 400 && reset_after_seconds.is_some() => 0.0,
code if code < 300 => (current_remaining - 1.0).max(0.0),
_ => return false,
};
window.insert("remaining".to_string(), json!(next_remaining));
if let Some(total) = total {
window.insert("total".to_string(), json!(total));
window.insert(
"remaining_fraction".to_string(),
json!((next_remaining / total).clamp(0.0, 1.0)),
);
window.insert(
"used_percent".to_string(),
json!(((total - next_remaining).max(0.0) / total * 100.0).clamp(0.0, 100.0)),
);
} else if status_code == 429 {
window.insert("remaining_fraction".to_string(), json!(0.0));
window.insert("used_percent".to_string(), json!(100.0));
}
if let Some(reset_after_seconds) = reset_after_seconds.filter(|seconds| *seconds > 0) {
let reset_at = now_unix_secs.saturating_add(reset_after_seconds);
window.insert("reset_at".to_string(), json!(reset_at));
window.insert("next_reset_at".to_string(), json!(reset_at));
window.insert(
"reset_after_seconds".to_string(),
json!(reset_after_seconds),
);
window.insert("reset_at_source".to_string(), json!("grok_upstream_error"));
}
window.insert("is_exhausted".to_string(), json!(next_remaining <= 0.0));
true
}
fn grok_mark_quota_bucket_updated(bucket: &mut serde_json::Map<String, Value>, now_unix_secs: u64) {
bucket.insert("updated_at".to_string(), json!(now_unix_secs));
}
async fn sync_grok_quota_from_report_context(
state: &AppState,
report_context: Option<&Value>,
status_code: u16,
body_json: Option<&Value>,
) -> Result<bool, GatewayError> {
let key_id = match report_context_key_id(report_context) {
Some(value) => value,
None => return Ok(false),
};
let Some(model) = grok_report_context_model(report_context) else {
return Ok(false);
};
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("grok") {
return Ok(false);
}
let Some(grok_bucket) = key
.upstream_metadata
.as_ref()
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("grok"))
.and_then(Value::as_object)
.cloned()
else {
return Ok(false);
};
let mut updated_grok_bucket = grok_bucket;
let now_unix_secs = current_unix_secs();
if !grok_apply_quota_feedback(
&mut updated_grok_bucket,
model.as_str(),
status_code,
grok_quota_reset_after_seconds(body_json, report_context),
now_unix_secs,
) {
return Ok(false);
}
grok_mark_quota_bucket_updated(&mut updated_grok_bucket, now_unix_secs);
let updated_upstream_metadata = merge_metadata_object(
key.upstream_metadata.as_ref(),
"grok",
Value::Object(updated_grok_bucket),
);
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(),
"report_effect",
);
let mut updated_key = key;
updated_key.upstream_metadata = updated_upstream_metadata;
updated_key.status_snapshot = updated_status_snapshot;
updated_key.updated_at_unix_secs = Some(now_unix_secs);
Ok(state
.update_provider_catalog_key(&updated_key)
.await?
.is_some())
}
async fn apply_local_sync_report_effect(state: &AppState, payload: &GatewaySyncReportRequest) {
apply_local_gemini_file_mapping_report_effect(state, payload).await;
if let Err(err) = sync_codex_quota_from_response_headers(
@@ -185,6 +447,23 @@ async fn apply_local_sync_report_effect(state: &AppState, payload: &GatewaySyncR
"gateway failed to persist codex realtime quota from sync response headers"
);
}
if let Err(err) = sync_grok_quota_from_report_context(
state,
payload.report_context.as_ref(),
payload.status_code,
payload.body_json.as_ref(),
)
.await
{
warn!(
event_name = "grok_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 grok realtime quota from sync response"
);
}
}
async fn apply_local_stream_report_effect(state: &AppState, payload: &GatewayStreamReportRequest) {
@@ -204,6 +483,23 @@ async fn apply_local_stream_report_effect(state: &AppState, payload: &GatewayStr
"gateway failed to persist codex realtime quota from stream response headers"
);
}
if let Err(err) = sync_grok_quota_from_report_context(
state,
payload.report_context.as_ref(),
payload.status_code,
grok_upstream_response_body(payload.report_context.as_ref()),
)
.await
{
warn!(
event_name = "grok_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 grok realtime quota from stream response"
);
}
}
async fn apply_local_gemini_file_mapping_report_effect(
@@ -443,3 +739,196 @@ pub(crate) fn clear_local_report_effect_caches_for_tests() {
.clear();
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn grok_quota_feedback_decrements_the_matching_window() {
let mut bucket = json!({
"quota_by_model": {
"quota_fast": {
"display_name": "fast",
"remaining": 30.0,
"total": 30.0,
"remaining_fraction": 1.0,
"used_percent": 0.0,
"is_exhausted": false
}
}
})
.as_object()
.cloned()
.expect("bucket should be object");
assert!(grok_apply_quota_feedback(
&mut bucket,
"grok-4.20-fast",
200,
None,
1_700_000_000
));
let fast = bucket
.get("quota_by_model")
.and_then(Value::as_object)
.and_then(|models| models.get("quota_fast"))
.and_then(Value::as_object)
.expect("fast window should exist");
assert_eq!(fast.get("remaining"), Some(&json!(29.0)));
assert_eq!(fast.get("remaining_fraction"), Some(&json!(29.0 / 30.0)));
assert_eq!(fast.get("used_percent"), Some(&json!(100.0 / 30.0)));
assert_eq!(fast.get("is_exhausted"), Some(&json!(false)));
}
#[test]
fn grok_report_context_model_requires_mapped_model() {
assert_eq!(
grok_report_context_model(Some(&json!({
"model": "grok-4.20-0309-reasoning"
}))),
None
);
assert_eq!(
grok_report_context_model(Some(&json!({
"mapped_model": "grok-4.20-fast"
}))),
Some("grok-4.20-fast".to_string())
);
}
#[test]
fn grok_quota_feedback_zeros_rate_limited_window() {
let mut bucket = json!({
"quota_by_model": {
"quota_fast": {
"display_name": "fast",
"remaining": 1.0,
"total": 30.0,
"remaining_fraction": 1.0 / 30.0,
"used_percent": 29.0 / 30.0 * 100.0,
"is_exhausted": false
}
}
})
.as_object()
.cloned()
.expect("bucket should be object");
assert!(grok_apply_quota_feedback(
&mut bucket,
"grok-4.20-fast",
429,
None,
1_700_000_000
));
let fast = bucket
.get("quota_by_model")
.and_then(Value::as_object)
.and_then(|models| models.get("quota_fast"))
.and_then(Value::as_object)
.expect("fast window should exist");
assert_eq!(fast.get("remaining"), Some(&json!(0.0)));
assert_eq!(fast.get("remaining_fraction"), Some(&json!(0.0)));
assert_eq!(fast.get("used_percent"), Some(&json!(100.0)));
assert_eq!(fast.get("is_exhausted"), Some(&json!(true)));
}
#[test]
fn grok_quota_feedback_records_reset_after_when_upstream_mentions_wait_time() {
let mut bucket = json!({
"quota_by_model": {
"quota_fast": {
"display_name": "fast",
"remaining": 1.0,
"total": 30.0,
"remaining_fraction": 1.0 / 30.0,
"used_percent": 29.0 / 30.0 * 100.0,
"is_exhausted": false,
"reset_at": 10,
"next_reset_at": 10
}
}
})
.as_object()
.cloned()
.expect("bucket should be object");
let parsed = grok_quota_reset_after_seconds(
Some(&json!({
"error": {
"message": "升级到 SuperGrok 获得更高使用上限,或等待 6小时 13分钟。"
}
})),
None,
);
assert_eq!(parsed, Some(22_380));
assert!(grok_apply_quota_feedback(
&mut bucket,
"grok-4.20-fast",
503,
parsed,
1_700_000_000
));
let fast = bucket
.get("quota_by_model")
.and_then(Value::as_object)
.and_then(|models| models.get("quota_fast"))
.and_then(Value::as_object)
.expect("fast window should exist");
assert_eq!(fast.get("remaining"), Some(&json!(0.0)));
assert_eq!(fast.get("is_exhausted"), Some(&json!(true)));
assert_eq!(fast.get("reset_after_seconds"), Some(&json!(22_380u64)));
assert_eq!(fast.get("reset_at"), Some(&json!(1_700_022_380u64)));
assert_eq!(fast.get("next_reset_at"), Some(&json!(1_700_022_380u64)));
assert_eq!(
fast.get("reset_at_source"),
Some(&json!("grok_upstream_error"))
);
}
#[test]
fn grok_realtime_quota_bucket_updates_observed_timestamp() {
let mut bucket = json!({
"updated_at": 1_600_000_000u64,
"quota_by_model": {
"quota_fast": {
"display_name": "fast",
"remaining": 1.0,
"total": 30.0,
"remaining_fraction": 1.0 / 30.0,
"used_percent": 29.0 / 30.0 * 100.0,
"is_exhausted": false
}
}
})
.as_object()
.cloned()
.expect("bucket should be object");
grok_mark_quota_bucket_updated(&mut bucket, 1_700_000_000);
assert_eq!(bucket.get("updated_at"), Some(&json!(1_700_000_000u64)));
}
#[test]
fn grok_wait_duration_parser_handles_english_and_chinese_messages() {
assert_eq!(
grok_wait_duration_seconds_from_text("wait 6h 13m"),
Some(22_380)
);
assert_eq!(
grok_wait_duration_seconds_from_text("等待 6小时13分钟"),
Some(22_380)
);
assert_eq!(
grok_wait_duration_seconds_from_text("no duration here"),
None
);
}
}

View File

@@ -4,6 +4,7 @@ use aether_data_contracts::repository::provider_catalog::{
use aether_provider_transport::provider_types::{
fixed_provider_key_inherits_api_formats, provider_type_is_fixed,
};
use serde_json::{Map, Value};
use std::collections::BTreeSet;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -81,6 +82,18 @@ impl ProviderKeyAuthSemantics {
}
}
pub(crate) fn provider_key_can_refresh_oauth(
auth_semantics: ProviderKeyAuthSemantics,
auth_config: Option<&Map<String, Value>>,
) -> bool {
auth_semantics.can_refresh_oauth()
&& auth_config
.and_then(|config| config.get("refresh_token"))
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty())
}
fn normalized_auth_type(key: &StoredProviderCatalogKey) -> String {
key.auth_type.trim().to_ascii_lowercase()
}
@@ -106,6 +119,10 @@ fn provider_uses_bearer_oauth_runtime(provider_type: &str) -> bool {
)
}
fn provider_uses_grok_session_runtime(provider_type: &str) -> bool {
provider_type.trim().eq_ignore_ascii_case("grok")
}
fn provider_key_is_legacy_kiro_oauth_session(
key: &StoredProviderCatalogKey,
provider_type: &str,
@@ -122,7 +139,8 @@ pub(crate) fn provider_key_auth_semantics(
) -> ProviderKeyAuthSemantics {
let auth_type = normalized_auth_type(key);
let oauth_managed = auth_type == "oauth"
|| provider_key_is_legacy_kiro_oauth_session(key, provider_type, &auth_type);
|| provider_key_is_legacy_kiro_oauth_session(key, provider_type, &auth_type)
|| (provider_uses_grok_session_runtime(provider_type) && key_has_auth_config(key));
let credential_kind = if oauth_managed {
ProviderKeyCredentialKind::OAuthSession
} else if matches!(auth_type.as_str(), "service_account" | "vertex_ai") {
@@ -135,6 +153,8 @@ pub(crate) fn provider_key_auth_semantics(
ProviderKeyCredentialKind::OAuthSession => {
if provider_uses_bearer_oauth_runtime(provider_type) {
ProviderKeyRuntimeAuthKind::Bearer
} else if provider_uses_grok_session_runtime(provider_type) {
ProviderKeyRuntimeAuthKind::Unknown
} else {
ProviderKeyRuntimeAuthKind::Unknown
}
@@ -226,7 +246,7 @@ pub(crate) fn provider_key_effective_api_formats(
#[cfg(test)]
mod tests {
use super::{
provider_active_api_formats, provider_key_auth_semantics,
provider_active_api_formats, provider_key_auth_semantics, provider_key_can_refresh_oauth,
provider_key_configured_api_formats, provider_key_effective_api_formats,
provider_key_inherits_provider_api_formats, ProviderKeyCredentialKind,
ProviderKeyRuntimeAuthKind,
@@ -293,6 +313,46 @@ mod tests {
);
}
#[test]
fn recognizes_grok_oauth_session_as_managed_without_bearer_runtime() {
let mut key = sample_key("oauth");
key.encrypted_auth_config = Some(r#"{"sso_token":"abc"}"#.to_string());
let semantics = provider_key_auth_semantics(&key, "grok");
assert!(semantics.oauth_managed());
assert_eq!(
semantics.credential_kind(),
ProviderKeyCredentialKind::OAuthSession
);
assert_eq!(
semantics.runtime_auth_kind(),
ProviderKeyRuntimeAuthKind::Unknown
);
}
#[test]
fn refresh_capability_requires_stored_refresh_token() {
let semantics = provider_key_auth_semantics(&sample_key("oauth"), "codex");
assert!(!provider_key_can_refresh_oauth(
semantics,
json!({
"access_token": "access-token",
"access_token_import_temporary": true
})
.as_object()
));
assert!(!provider_key_can_refresh_oauth(
semantics,
json!({ "refresh_token": " " }).as_object()
));
assert!(provider_key_can_refresh_oauth(
semantics,
json!({ "refresh_token": "refresh-token" }).as_object()
));
}
#[test]
fn recognizes_legacy_kiro_bearer_key_with_auth_config_as_oauth_managed() {
let mut key = sample_key("bearer");

View File

@@ -94,6 +94,7 @@ async fn gateway_handles_admin_provider_models_locally_with_trusted_admin_princi
assert_eq!(items[0]["effective_input_price"], 3.0);
assert_eq!(items[0]["effective_output_price"], 15.0);
assert_eq!(items[0]["effective_supports_streaming"], true);
assert!(items[0]["model_test_capabilities"]["openai:image"].is_null());
assert_eq!(items[0]["created_at"], "2024-03-21T05:46:40Z");
assert_eq!(items[0]["updated_at"], "2024-03-21T05:48:20Z");
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);

View File

@@ -1967,6 +1967,144 @@ async fn gateway_streams_codex_openai_responses_upstream_for_admin_pool_model_te
execution_runtime_handle.abort();
}
#[tokio::test]
async fn gateway_routes_grok_responses_admin_pool_model_test_through_grok_runtime() {
let execution_runtime = Router::new().route(
"/v1/execute/sync",
any(move |Json(plan): Json<ExecutionPlan>| async move {
assert_eq!(plan.provider_id, "provider-grok");
assert_eq!(plan.endpoint_id, "endpoint-grok-responses");
assert_eq!(plan.key_id, "key-grok-oauth");
assert_eq!(plan.client_api_format, "openai:responses");
assert_eq!(plan.provider_api_format, "openai:responses");
assert_eq!(plan.url, "https://grok.com/rest/app-chat/conversations/new");
assert_eq!(plan.model_name.as_deref(), Some("grok-4.20-fast"));
assert!(plan.stream, "Grok model test should request a stream");
assert_eq!(
plan.headers
.get(aether_provider_transport::GROK_INTERNAL_HEADER)
.map(String::as_str),
Some("1")
);
assert_eq!(
plan.headers.get("cookie").map(String::as_str),
Some("sso=grok-sso; sso-rw=grok-rw")
);
let body = plan.body.json_body.as_ref().expect("json body");
assert_eq!(body["model"], json!("grok-4.20-fast"));
assert_eq!(body["input"], json!("Hello! This is a test message."));
assert_eq!(
body["messages"][0]["content"],
json!("stale chat-shaped frontend body")
);
Json(json!({
"request_id": plan.request_id,
"candidate_id": plan.candidate_id,
"status_code": 200,
"headers": {
"content-type": "application/json"
},
"body": {
"json_body": {
"id": "resp-grok-model-test",
"model": "grok-4.20-fast",
"output_text": "ok"
}
},
"telemetry": {
"elapsed_ms": 18
}
}))
}),
);
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
let mut provider = sample_provider("provider-grok", "Grok", 10);
provider.provider_type = "grok".to_string();
provider.config = Some(json!({"pool_advanced": {}}));
let mut key = sample_key(
"key-grok-oauth",
"provider-grok",
"openai:responses",
"__placeholder__",
);
key.auth_type = "oauth".to_string();
key.encrypted_auth_config = Some(
aether_crypto::encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
r#"{
"provider_type":"grok",
"sso_token":"grok-sso",
"sso_rw_token":"grok-rw"
}"#,
)
.expect("auth config should encrypt"),
);
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![sample_endpoint(
"endpoint-grok-responses",
"provider-grok",
"openai:responses",
"https://grok.com",
)],
vec![key],
));
let gateway = build_router_with_state(
build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(GatewayDataState::with_provider_transport_reader_for_tests(
provider_catalog_repository,
DEVELOPMENT_ENCRYPTION_KEY.to_string(),
)),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!(
"{gateway_url}/api/admin/provider-query/test-model-failover"
))
.header(GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({
"provider_id": "provider-grok",
"mode": "pool",
"model": "grok-4.20-fast",
"failover_models": ["grok-4.20-fast"],
"api_format": "openai:responses",
"endpoint_id": "endpoint-grok-responses",
"request_body": {
"model": "grok-4.20-fast",
"messages": [{
"role": "user",
"content": "stale chat-shaped frontend body"
}],
"stream": true
}
}))
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["success"], json!(true));
assert_eq!(payload["attempts"][0]["status"], json!("success"));
assert_eq!(
payload["attempts"][0]["request_body"]["message"],
json!("Hello! This is a test message.")
);
assert_eq!(
payload["attempts"][0]["request_headers"][aether_provider_transport::GROK_INTERNAL_HEADER],
json!("1")
);
gateway_handle.abort();
execution_runtime_handle.abort();
}
#[tokio::test]
async fn gateway_uses_pool_scheduler_order_for_admin_pool_model_test() {
let execution_runtime = Router::new().route(

View File

@@ -4,6 +4,7 @@ use super::{
InMemoryVideoTaskRepository, UpsertVideoTask, VideoTaskLookupKey, VideoTaskReadRepository,
VideoTaskStatus, VideoTaskWriteRepository, DEVELOPMENT_ENCRYPTION_KEY,
};
use crate::image_capabilities::openai_image_gateway_max_generation_count;
use crate::tests::{
any, build_router_with_state, build_state_with_execution_runtime_override, json, start_server,
to_bytes, AppState, Arc, Body, Json, Mutex, Request, Router, StatusCode, EXECUTION_PATH_HEADER,
@@ -582,7 +583,7 @@ async fn gateway_does_not_locally_reject_image_model_name_on_chat_completions()
}
#[tokio::test]
async fn gateway_rejects_image_request_with_n_greater_than_one_without_hitting_fallback_probe() {
async fn gateway_rejects_image_request_with_n_greater_than_four_without_hitting_fallback_probe() {
let fallback_probe_hits = Arc::new(Mutex::new(0usize));
let fallback_probe_hits_clone = Arc::clone(&fallback_probe_hits);
let fallback_probe = Router::new().route(
@@ -615,9 +616,9 @@ async fn gateway_rejects_image_request_with_n_greater_than_one_without_hitting_f
.header(http::header::CONTENT_TYPE, "application/json")
.body(
serde_json::to_vec(&json!({
"model": "gpt-image-2",
"model": "grok-imagine-image-lite",
"prompt": "draw",
"n": 2,
"n": 5,
"response_format": "b64_json"
}))
.expect("request body should encode"),
@@ -635,7 +636,13 @@ async fn gateway_rejects_image_request_with_n_greater_than_one_without_hitting_f
Some(EXECUTION_PATH_LOCAL_AI_PUBLIC)
);
let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["detail"], "当前 Codex 图片反代仅支持 n=1");
assert_eq!(
payload["detail"],
format!(
"当前图片反代仅支持 n=1..{}",
openai_image_gateway_max_generation_count()
)
);
assert_eq!(*fallback_probe_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort();
@@ -643,7 +650,7 @@ async fn gateway_rejects_image_request_with_n_greater_than_one_without_hitting_f
}
#[tokio::test]
async fn gateway_rejects_variation_request_without_image_without_hitting_fallback_probe() {
async fn gateway_does_not_mount_image_variation_route_without_hitting_fallback_probe() {
let fallback_probe_hits = Arc::new(Mutex::new(0usize));
let fallback_probe_hits_clone = Arc::clone(&fallback_probe_hits);
let fallback_probe = Router::new().route(
@@ -685,16 +692,7 @@ async fn gateway_rejects_variation_request_without_image_without_hitting_fallbac
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
assert_eq!(
response
.headers()
.get(EXECUTION_PATH_HEADER)
.and_then(|value| value.to_str().ok()),
Some(EXECUTION_PATH_LOCAL_AI_PUBLIC)
);
let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["detail"], "图片变体请求需要 image 文件");
assert_eq!(response.status(), StatusCode::NOT_FOUND);
assert_eq!(*fallback_probe_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort();