Merge pull request #822 from stabey/upstream-pr/xai-media

feat(providers): 新增 xAI Provider(设备码 OAuth + 原生图像/视频)
This commit is contained in:
ZheFox
2026-09-15 09:58:55 +08:00
committed by GitHub
105 changed files with 7347 additions and 284 deletions
Generated
+1
View File
@@ -748,6 +748,7 @@ dependencies = [
"async-trait", "async-trait",
"serde", "serde",
"serde_json", "serde_json",
"sha2",
"url", "url",
"uuid", "uuid",
] ]
@@ -583,6 +583,11 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
source_model, source_model,
codex_model_capabilities.as_ref(), codex_model_capabilities.as_ref(),
); );
crate::ai_serving::transport::xai::insert_cli_identity_headers_if_needed(
transport.as_ref(),
prepared.provider_api_format.as_str(),
&mut provider_request_headers,
);
request_identity_response_encoding_when_redacted( request_identity_response_encoding_when_redacted(
&mut provider_request_headers, &mut provider_request_headers,
redaction.redacted, redaction.redacted,
@@ -17,8 +17,8 @@ use crate::ai_serving::transport::{
ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput, GROK_CHAT_PATH, ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput, GROK_CHAT_PATH,
}; };
use crate::ai_serving::{ use crate::ai_serving::{
apply_codex_openai_special_headers, build_chatgpt_web_image_request_body, apply_codex_openai_special_headers, apply_xai_upstream_payload_edits,
build_codex_openai_image_api_provider_request_body, build_chatgpt_web_image_request_body, build_codex_openai_image_api_provider_request_body,
build_gemini_image_request_body_from_openai_image_request, build_gemini_image_request_body_from_openai_image_request,
build_openai_image_api_provider_request_body, build_openai_image_provider_request_body, build_openai_image_api_provider_request_body, build_openai_image_provider_request_body,
default_model_for_openai_image_operation, normalize_openai_image_request, default_model_for_openai_image_operation, normalize_openai_image_request,
@@ -211,7 +211,7 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
upstream_is_stream, upstream_is_stream,
) )
}; };
let Some(provider_request_body) = provider_request_body else { let Some(mut provider_request_body) = provider_request_body else {
mark_skipped_local_openai_image_candidate_with_failure_diagnostic( mark_skipped_local_openai_image_candidate_with_failure_diagnostic(
state, state,
input, input,
@@ -229,6 +229,11 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
.await; .await;
return None; return None;
}; };
apply_xai_upstream_payload_edits(
&mut provider_request_body,
transport.provider.provider_type.as_str(),
provider_api_format,
);
let Some(mut provider_request_headers) = (if is_grok { let Some(mut provider_request_headers) = (if is_grok {
build_grok_browser_headers(GrokHeaderInput { build_grok_browser_headers(GrokHeaderInput {
transport, transport,
@@ -8,6 +8,7 @@ use crate::ai_serving::planner::{
build_ai_execution_decision_response, resolve_transport_request_encoding_policy, build_ai_execution_decision_response, resolve_transport_request_encoding_policy,
AiExecutionDecisionResponseParts, AiExecutionDecisionResponseParts,
}; };
use crate::ai_serving::transport::xai::video::is_native_video_request;
use crate::ai_serving::transport::{ use crate::ai_serving::transport::{
resolve_transport_execution_timeouts, resolve_transport_profile, resolve_transport_execution_timeouts, resolve_transport_profile,
}; };
@@ -33,7 +34,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
let Some(resolved) = resolve_local_video_create_candidate_payload_parts( let Some(resolved) = resolve_local_video_create_candidate_payload_parts(
state, parts, body_json, trace_id, input, &attempt, spec, state, parts, body_json, trace_id, input, &attempt, spec,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
@@ -52,9 +53,32 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
.await; .await;
let transport_profile = resolve_transport_profile(&transport); let transport_profile = resolve_transport_profile(&transport);
let mut extra_fields = serde_json::Map::new(); let mut extra_fields = serde_json::Map::new();
if is_native_video_request(&transport.provider.provider_type, parts.uri.path()) {
extra_fields.insert(
"video_client_protocol".to_string(),
serde_json::json!("xai"),
);
}
if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) { if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) {
extra_fields.insert("proxy".to_string(), proxy_value); extra_fields.insert("proxy".to_string(), proxy_value);
} }
if transport.provider.provider_type.eq_ignore_ascii_case("xai") {
extra_fields.insert("video_provider_xai".into(), serde_json::json!(true));
if let Some(duration) = resolved.provider_request_body.get("duration") {
extra_fields.insert("video_duration".into(), duration.clone());
}
if parts.uri.path() == "/openai/v1/videos" {
extra_fields.insert(
"video_size".into(),
body_json
.get("size")
.filter(|v| v.as_str().is_some_and(|s| !s.trim().is_empty()))
.cloned()
.unwrap_or_else(|| serde_json::json!("720x1280")),
);
}
}
let effective_headers = input.effective_headers(&parts.headers); let effective_headers = input.effective_headers(&parts.headers);
let report_context = build_local_execution_report_context(LocalExecutionReportContextParts { let report_context = build_local_execution_report_context(LocalExecutionReportContextParts {
auth_context: &input.auth_context, auth_context: &input.auth_context,
@@ -3,15 +3,23 @@ use std::sync::Arc;
use serde_json::Value; use serde_json::Value;
use crate::ai_serving::planner::candidate_preparation::resolve_candidate_mapped_model; use crate::ai_serving::planner::candidate_preparation::{
prepare_header_authenticated_candidate, resolve_candidate_mapped_model, OauthPreparationContext,
};
use crate::ai_serving::planner::spec_metadata::local_video_create_spec_metadata; use crate::ai_serving::planner::spec_metadata::local_video_create_spec_metadata;
use crate::ai_serving::transport::xai::video::{
convert_openai_video_request, is_explicit_native_video_path, is_native_video_request,
};
use crate::ai_serving::transport::{ use crate::ai_serving::transport::{
build_video_create_headers, build_video_create_request_body, build_video_create_upstream_url, build_video_create_headers, build_video_create_request_body, build_video_create_upstream_url,
resolve_video_create_auth, video_create_transport_unsupported_reason, resolve_video_create_auth, video_create_transport_unsupported_reason,
ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput, ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput,
}; };
use crate::ai_serving::{CandidateFailureDiagnostic, GatewayProviderTransportSnapshot}; use crate::ai_serving::{
use crate::AppState; apply_xai_upstream_payload_edits, CandidateFailureDiagnostic, GatewayProviderTransportSnapshot,
PlannerAppState,
};
use crate::{AppState, GatewayError};
use super::support::{ use super::support::{
mark_skipped_local_video_candidate, mark_skipped_local_video_candidate_with_failure_diagnostic, mark_skipped_local_video_candidate, mark_skipped_local_video_candidate_with_failure_diagnostic,
@@ -37,11 +45,16 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
input: &LocalVideoCreateDecisionInput, input: &LocalVideoCreateDecisionInput,
attempt: &LocalVideoCreateCandidateAttempt, attempt: &LocalVideoCreateCandidateAttempt,
spec: LocalVideoCreateSpec, spec: LocalVideoCreateSpec,
) -> Option<LocalVideoCreateCandidatePayloadParts> { ) -> Result<Option<LocalVideoCreateCandidatePayloadParts>, GatewayError> {
let spec_metadata = local_video_create_spec_metadata(spec); let spec_metadata = local_video_create_spec_metadata(spec);
let candidate = &attempt.eligible.candidate; let candidate = &attempt.eligible.candidate;
let transport = &attempt.eligible.transport; let transport = &attempt.eligible.transport;
let effective_headers = input.effective_headers(&parts.headers); let effective_headers = input.effective_headers(&parts.headers);
if is_explicit_native_video_path(parts.uri.path())
&& !transport.provider.provider_type.eq_ignore_ascii_case("xai")
{
return Ok(None);
}
let provider_family = provider_video_create_family(spec.family); let provider_family = provider_video_create_family(spec.family);
let transport_unsupported_reason = video_create_transport_unsupported_reason( let transport_unsupported_reason = video_create_transport_unsupported_reason(
@@ -60,23 +73,39 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
skip_reason, skip_reason,
) )
.await; .await;
return None; return Ok(None);
} }
let auth = resolve_video_create_auth(transport, provider_family); let prepared_candidate = match prepare_header_authenticated_candidate(
let Some((auth_header, auth_value)) = auth else { PlannerAppState::new(state),
mark_skipped_local_video_candidate( transport,
state, candidate,
input, resolve_video_create_auth(transport, provider_family),
OauthPreparationContext {
trace_id, trace_id,
candidate, api_format: spec_metadata.api_format,
attempt.candidate_index, operation: "video_create_candidate_request",
&attempt.candidate_id, },
"transport_auth_unavailable", )
) .await
.await; {
return None; Ok(prepared) => prepared,
Err(skip_reason) => {
mark_skipped_local_video_candidate(
state,
input,
trace_id,
candidate,
attempt.candidate_index,
&attempt.candidate_id,
skip_reason,
)
.await;
return Ok(None);
}
}; };
let auth_header = prepared_candidate.auth_header;
let auth_value = prepared_candidate.auth_value;
let mapped_model = match resolve_candidate_mapped_model(candidate) { let mapped_model = match resolve_candidate_mapped_model(candidate) {
Ok(mapped_model) => mapped_model, Ok(mapped_model) => mapped_model,
@@ -91,7 +120,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
skip_reason, skip_reason,
) )
.await; .await;
return None; return Ok(None);
} }
}; };
@@ -117,10 +146,10 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
), ),
) )
.await; .await;
return None; return Ok(None);
}; };
let Some(provider_request_body) = build_video_create_request_body( let Some(mut provider_request_body) = build_video_create_request_body(
body_json, body_json,
provider_family, provider_family,
&mapped_model, &mapped_model,
@@ -142,11 +171,28 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
), ),
) )
.await; .await;
return None; return Ok(None);
}; };
if transport.provider.provider_type.eq_ignore_ascii_case("xai")
&& !is_native_video_request(&transport.provider.provider_type, parts.uri.path())
{
provider_request_body =
convert_openai_video_request(&provider_request_body).map_err(|message| {
GatewayError::Client {
status: http::StatusCode::BAD_REQUEST,
message: message.to_string(),
}
})?;
}
apply_xai_upstream_payload_edits(
&mut provider_request_body,
transport.provider.provider_type.as_str(),
spec_metadata.api_format,
);
let Some(provider_request_headers) = let Some(provider_request_headers) =
build_video_create_headers(ProviderVideoCreateHeadersInput { build_video_create_headers(ProviderVideoCreateHeadersInput {
transport,
headers: effective_headers, headers: effective_headers,
auth_header: &auth_header, auth_header: &auth_header,
auth_value: &auth_value, auth_value: &auth_value,
@@ -170,10 +216,10 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
), ),
) )
.await; .await;
return None; return Ok(None);
}; };
Some(LocalVideoCreateCandidatePayloadParts { Ok(Some(LocalVideoCreateCandidatePayloadParts {
transport: Arc::clone(transport), transport: Arc::clone(transport),
auth_header, auth_header,
auth_value, auth_value,
@@ -181,7 +227,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
provider_request_headers, provider_request_headers,
provider_request_body, provider_request_body,
upstream_url, upstream_url,
}) }))
} }
fn provider_video_create_family(family: LocalVideoCreateFamily) -> ProviderVideoCreateFamily { fn provider_video_create_family(family: LocalVideoCreateFamily) -> ProviderVideoCreateFamily {
@@ -13,7 +13,9 @@ pub(crate) fn openai_responses_reasoning_replay_policy(
base_url: &str, base_url: &str,
_provider_model: &str, _provider_model: &str,
) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy { ) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy {
if is_deepseek_provider(provider_type, base_url) { if provider_type.trim().eq_ignore_ascii_case("xai") {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
} else if is_deepseek_provider(provider_type, base_url) {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
} else { } else {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
@@ -238,6 +240,27 @@ mod tests {
openai_responses_reasoning_replay_policy, openai_responses_reasoning_replay_policy,
}; };
#[test]
fn xai_reasoning_policy_comes_from_provider_type() {
use crate::ai_serving::OpenAiResponsesReasoningReplayPolicy;
assert_eq!(
openai_responses_reasoning_replay_policy(
"xai",
"https://custom.example/v1",
"grok-4.6"
),
OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
);
assert_eq!(
openai_responses_reasoning_replay_policy(
"openai",
"https://custom.example/v1",
"grok-4.6"
),
OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
);
}
#[test] #[test]
fn detects_deepseek_provider_only_by_official_host() { fn detects_deepseek_provider_only_by_official_host() {
assert!(!is_deepseek_provider( assert!(!is_deepseek_provider(
@@ -488,6 +488,7 @@ impl ResponsesWebSocketBodyNormalization {
digest.update([match self.reasoning_replay_policy { digest.update([match self.reasoning_replay_policy {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds => 0, crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds => 0,
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque => 1, crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque => 1,
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted => 2,
}]); }]);
update_normalization_optional_json_digest(&mut digest, self.model_directive_patch.as_ref()); update_normalization_optional_json_digest(&mut digest, self.model_directive_patch.as_ref());
digest.finalize().into() digest.finalize().into()
@@ -12,7 +12,8 @@ pub(crate) use aether_ai_formats::api::{
apply_codex_openai_responses_websocket_continuation_body_edits_with_source_model_and_capabilities, apply_codex_openai_responses_websocket_continuation_body_edits_with_source_model_and_capabilities,
apply_codex_openai_special_headers, apply_model_directive_mapping_patch, apply_codex_openai_special_headers, apply_model_directive_mapping_patch,
apply_model_directive_overrides_from_model, apply_model_directive_overrides_from_request, apply_model_directive_overrides_from_model, apply_model_directive_overrides_from_request,
apply_openai_responses_compact_special_body_edits, build_chatgpt_web_image_request_body, apply_openai_responses_compact_special_body_edits, apply_xai_upstream_payload_edits,
apply_xai_upstream_payload_edits_with_client, build_chatgpt_web_image_request_body,
build_codex_model_catalog_metadata, build_codex_openai_image_api_provider_request_body, build_codex_model_catalog_metadata, build_codex_openai_image_api_provider_request_body,
build_core_error_body_for_client_format, build_cross_format_openai_chat_request_body, build_core_error_body_for_client_format, build_cross_format_openai_chat_request_body,
build_cross_format_openai_chat_request_body_with_model_directives, build_cross_format_openai_chat_request_body_with_model_directives,
@@ -58,6 +58,10 @@ pub(crate) mod windsurf {
pub(crate) use aether_provider_transport::windsurf::*; pub(crate) use aether_provider_transport::windsurf::*;
} }
pub(crate) mod xai {
pub(crate) use aether_provider_transport::xai::*;
}
pub(crate) use aether_provider_transport::{ pub(crate) use aether_provider_transport::{
append_transport_diagnostics_to_value, apply_codex_fingerprint_convergence, append_transport_diagnostics_to_value, apply_codex_fingerprint_convergence,
apply_codex_fingerprint_convergence_with_context, apply_local_auth_config_header_overrides, apply_codex_fingerprint_convergence_with_context, apply_local_auth_config_header_overrides,
@@ -53,6 +53,8 @@ const AI_ANY_ROUTE_PATTERNS: &[&str] = &[
"/v1beta/operations/{*operation_path}", "/v1beta/operations/{*operation_path}",
"/v1/videos", "/v1/videos",
"/v1/videos/{*video_path}", "/v1/videos/{*video_path}",
"/openai/v1/videos",
"/openai/v1/videos/{*video_path}",
"/upload/v1beta/files", "/upload/v1beta/files",
"/v1beta/files", "/v1beta/files",
"/v1beta/files/{*file_path}", "/v1beta/files/{*file_path}",
@@ -536,6 +536,9 @@ mod tests {
fn sample_sparse_stored_task() -> StoredVideoTask { fn sample_sparse_stored_task() -> StoredVideoTask {
let snapshot = LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { let snapshot = LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: false,
local_task_id: "task-1".to_string(), local_task_id: "task-1".to_string(),
upstream_task_id: "ext-1".to_string(), upstream_task_id: "ext-1".to_string(),
created_at_unix_ms: 1, created_at_unix_ms: 1,
+2
View File
@@ -140,6 +140,8 @@ pub(crate) const RUST_FRONTDOOR_OWNED_ROUTE_PATTERNS: &[&str] = &[
"/v1beta/models/{model}/operations/{id}", "/v1beta/models/{model}/operations/{id}",
"/v1beta/operations", "/v1beta/operations",
"/v1beta/operations/{id}", "/v1beta/operations/{id}",
"/openai/v1/videos",
"/openai/v1/videos/{path...}",
"/v1/videos", "/v1/videos",
"/v1/videos/{path...}", "/v1/videos/{path...}",
"/upload/v1beta/files", "/upload/v1beta/files",
+5 -1
View File
@@ -137,7 +137,11 @@ pub(super) fn classify_ai_public_route(
.with_client_surface(detect_claude_client_surface(headers)) .with_client_surface(detect_claude_client_surface(headers))
.with_api_operation(ApiOperation::ClaudeMessagesCreate), .with_api_operation(ApiOperation::ClaudeMessagesCreate),
) )
} else if normalized_path.starts_with("/v1/videos") { } else if normalized_path == "/v1/videos"
|| normalized_path.starts_with("/v1/videos/")
|| normalized_path == "/openai/v1/videos"
|| normalized_path.starts_with("/openai/v1/videos/")
{
Some(classified( Some(classified(
"ai_public", "ai_public",
"openai", "openai",
@@ -123,6 +123,15 @@ impl GatewayDataState {
} }
#[cfg(test)] #[cfg(test)]
pub(crate) fn attach_video_task_repository_for_tests<T>(mut self, repository: Arc<T>) -> Self
where
T: VideoTaskRepository + 'static,
{
self.video_task_reader = Some(repository.clone());
self.video_task_writer = Some(repository);
self
}
pub(crate) fn with_video_task_repository_for_tests<T>(repository: Arc<T>) -> Self pub(crate) fn with_video_task_repository_for_tests<T>(repository: Arc<T>) -> Self
where where
T: VideoTaskRepository + 'static, T: VideoTaskRepository + 'static,
@@ -1464,6 +1464,22 @@ pub(crate) async fn maybe_execute_sync_via_local_video_decision(
.await .await
} }
fn supports_local_video_get(
parts: &http::request::Parts,
decision: &GatewayControlDecision,
) -> bool {
parts.method == http::Method::GET
&& decision.route_kind.as_deref() == Some("video")
&& (crate::video_tasks::resolve_video_task_read_lookup_key(
decision.route_family.as_deref(),
parts.uri.path(),
)
.is_some()
|| (decision.route_family.as_deref() == Some("openai")
&& crate::video_tasks::extract_openai_task_id_from_content_path(parts.uri.path())
.is_some()))
}
pub(crate) fn maybe_execute_sync_request<'a>( pub(crate) fn maybe_execute_sync_request<'a>(
state: &'a AppState, state: &'a AppState,
parts: &'a http::request::Parts, parts: &'a http::request::Parts,
@@ -1477,7 +1493,7 @@ pub(crate) fn maybe_execute_sync_request<'a>(
}; };
#[cfg(not(test))] #[cfg(not(test))]
{ {
if parts.method != http::Method::POST { if parts.method != http::Method::POST && !supports_local_video_get(parts, decision) {
return Ok(LocalExecutionRequestOutcome::NoPath); return Ok(LocalExecutionRequestOutcome::NoPath);
} }
return maybe_execute_sync_local_path(state, parts, body_bytes, trace_id, decision) return maybe_execute_sync_local_path(state, parts, body_bytes, trace_id, decision)
@@ -1490,6 +1506,7 @@ pub(crate) fn maybe_execute_sync_request<'a>(
.unwrap_or_default() .unwrap_or_default()
.is_empty() .is_empty()
&& parts.method != http::Method::POST && parts.method != http::Method::POST
&& !supports_local_video_get(parts, decision)
{ {
return Ok(LocalExecutionRequestOutcome::NoPath); return Ok(LocalExecutionRequestOutcome::NoPath);
} }
@@ -1511,7 +1528,7 @@ pub(crate) fn maybe_execute_stream_request<'a>(
}; };
#[cfg(not(test))] #[cfg(not(test))]
{ {
if parts.method != http::Method::POST { if parts.method != http::Method::POST && !supports_local_video_get(parts, decision) {
return Ok(LocalExecutionRequestOutcome::NoPath); return Ok(LocalExecutionRequestOutcome::NoPath);
} }
return maybe_execute_stream_local_path(state, parts, body_bytes, trace_id, decision) return maybe_execute_stream_local_path(state, parts, body_bytes, trace_id, decision)
@@ -1524,6 +1541,7 @@ pub(crate) fn maybe_execute_stream_request<'a>(
.unwrap_or_default() .unwrap_or_default()
.is_empty() .is_empty()
&& parts.method != http::Method::POST && parts.method != http::Method::POST
&& !supports_local_video_get(parts, decision)
{ {
return Ok(LocalExecutionRequestOutcome::NoPath); return Ok(LocalExecutionRequestOutcome::NoPath);
} }
@@ -32,6 +32,10 @@ fn request_has_execution_runtime_via_guard(headers: &HeaderMap) -> bool {
} }
pub(crate) fn frontdoor_self_loop_public_ai_path(path: &str) -> bool { pub(crate) fn frontdoor_self_loop_public_ai_path(path: &str) -> bool {
let path = path
.strip_prefix("/openai")
.filter(|p| *p == "/v1/videos" || p.starts_with("/v1/videos/"))
.unwrap_or(path);
matches!( matches!(
path, path,
"/v1/messages" "/v1/messages"
@@ -74,7 +74,8 @@ fn validate_batch_access_token_import(
) -> Result<(), String> { ) -> Result<(), String> {
if !provider_type_supports_access_token_import(provider_type) { if !provider_type_supports_access_token_import(provider_type) {
return Err( return Err(
"Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok Provider".to_string(), "Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok / xAI Provider"
.to_string(),
); );
} }
if provider_type.eq_ignore_ascii_case("claude_code") { if provider_type.eq_ignore_ascii_case("claude_code") {
@@ -214,7 +214,12 @@ fn extract_admin_provider_oauth_batch_import_entry(
} else { } else {
let sso_from_cookie = grok_cookie_session_token(provider_type, raw_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 token_input = sso_from_cookie.as_deref().unwrap_or(raw_token);
let (refresh_token, access_token) = import_tokens_from_raw_token(token_input); let (refresh_token, access_token) =
if provider_type.trim().eq_ignore_ascii_case("xai") {
(None, Some(token_input.to_string()))
} else {
import_tokens_from_raw_token(token_input)
};
let (refresh_token, access_token) = normalize_provider_import_tokens( let (refresh_token, access_token) = normalize_provider_import_tokens(
provider_type, provider_type,
refresh_token.as_deref(), refresh_token.as_deref(),
@@ -262,6 +267,7 @@ fn extract_admin_provider_oauth_batch_import_entry(
let object = normalized_claude_object.as_ref().unwrap_or(object); let object = normalized_claude_object.as_ref().unwrap_or(object);
let is_grok = provider_type.trim().eq_ignore_ascii_case("grok"); let is_grok = provider_type.trim().eq_ignore_ascii_case("grok");
let is_windsurf = provider_type.trim().eq_ignore_ascii_case("windsurf"); let is_windsurf = provider_type.trim().eq_ignore_ascii_case("windsurf");
let is_xai = provider_type.trim().eq_ignore_ascii_case("xai");
let is_codex_agent_identity = provider_type.trim().eq_ignore_ascii_case("codex") let is_codex_agent_identity = provider_type.trim().eq_ignore_ascii_case("codex")
&& aether_provider_transport::is_codex_agent_identity_auth_config_value(item); && aether_provider_transport::is_codex_agent_identity_auth_config_value(item);
if is_codex_agent_identity { if is_codex_agent_identity {
@@ -336,14 +342,6 @@ fn extract_admin_provider_oauth_batch_import_entry(
} else { } else {
None None
}; };
let (refresh_token, access_token) = normalize_provider_import_tokens(
provider_type,
refresh_token.as_deref(),
access_token
.as_deref()
.or(session_token.as_deref())
.or(header_bearer_token.as_deref()),
);
let windsurf_api_key = is_windsurf let windsurf_api_key = is_windsurf
.then(|| { .then(|| {
coerce_admin_provider_oauth_import_str( coerce_admin_provider_oauth_import_str(
@@ -351,6 +349,22 @@ fn extract_admin_provider_oauth_batch_import_entry(
) )
}) })
.flatten(); .flatten();
let xai_api_key = is_xai
.then(|| {
coerce_admin_provider_oauth_import_str(
object.get("api_key").or_else(|| object.get("apiKey")),
)
})
.flatten();
let (refresh_token, access_token) = normalize_provider_import_tokens(
provider_type,
refresh_token.as_deref(),
access_token
.as_deref()
.or(session_token.as_deref())
.or(header_bearer_token.as_deref())
.or(xai_api_key.as_deref()),
);
let windsurf_token = is_windsurf let windsurf_token = is_windsurf
.then(|| { .then(|| {
coerce_admin_provider_oauth_import_str( coerce_admin_provider_oauth_import_str(
@@ -1577,4 +1591,23 @@ mod tests {
assert!(entries[1].access_token.is_none()); assert!(entries[1].access_token.is_none());
assert!(entries[1].raw_credentials.is_none()); assert!(entries[1].raw_credentials.is_none());
} }
#[test]
fn parses_xai_api_key_json_and_raw_lines_as_access_token() {
let entries = parse_admin_provider_oauth_batch_import_entries(
"xai",
r#"{"api_key":"xai-api-key","email":"a@x.ai"}
{"refresh_token":"xai-refresh"}
xai-raw-api-key"#,
);
assert_eq!(entries.len(), 3);
assert!(entries[0].refresh_token.is_none());
assert_eq!(entries[0].access_token.as_deref(), Some("xai-api-key"));
assert_eq!(entries[0].email.as_deref(), Some("[email protected]"));
assert_eq!(entries[1].refresh_token.as_deref(), Some("xai-refresh"));
assert!(entries[1].access_token.is_none());
assert!(entries[2].refresh_token.is_none());
assert_eq!(entries[2].access_token.as_deref(), Some("xai-raw-api-key"));
}
} }
@@ -186,10 +186,10 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
)); ));
}; };
let provider_type = provider.provider_type.trim().to_ascii_lowercase(); let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if provider_type != "kiro" && provider_type != "windsurf" { if provider_type != "kiro" && provider_type != "windsurf" && provider_type != "xai" {
return Ok(build_internal_control_error_response( return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST, http::StatusCode::BAD_REQUEST,
"设备授权仅支持 Kiro / Windsurf provider", "设备授权仅支持 Kiro / Windsurf / xAI provider",
)); ));
} }
let Some(principal) = request_context let Some(principal) = request_context
@@ -219,6 +219,19 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
) )
.await; .await;
if provider_type == "xai" {
return super::xai::handle_admin_provider_oauth_xai_device_authorize(
state,
&provider_id,
&provider,
principal,
runtime_endpoint.as_ref(),
request_proxy,
payload.proxy_node_id.as_deref(),
)
.await;
}
if provider_type == "windsurf" { if provider_type == "windsurf" {
let session_id = generate_provider_oauth_nonce(); let session_id = generate_provider_oauth_nonce();
let login_option = payload let login_option = payload
@@ -2,6 +2,7 @@ mod authorize;
mod lease; mod lease;
mod poll; mod poll;
mod session; mod session;
mod xai;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError; use crate::GatewayError;
@@ -479,6 +479,18 @@ pub(super) async fn handle_admin_provider_oauth_device_poll(
) )
.await; .await;
if provider_type == "xai" {
return super::xai::handle_admin_provider_oauth_xai_device_poll(
state,
&provider,
&endpoints,
request_proxy,
session_id,
session,
)
.await;
}
if provider_type == "windsurf" { if provider_type == "windsurf" {
return handle_admin_provider_oauth_windsurf_browser_device_poll( return handle_admin_provider_oauth_windsurf_browser_device_poll(
state, state,
@@ -0,0 +1,368 @@
use super::session::attach_admin_provider_oauth_device_poll_terminal_response;
use crate::control::GatewayAdminPrincipalContext;
use crate::handlers::admin::provider::oauth::dispatch::helpers::admin_provider_oauth_key_name_from_auth_config;
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
use crate::handlers::admin::provider::oauth::provisioning::{
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
};
use crate::handlers::admin::provider::oauth::runtime::spawn_provider_oauth_account_state_refresh_after_update;
use crate::handlers::admin::provider::oauth::state::{
current_unix_secs, generate_provider_oauth_nonce,
};
use crate::handlers::admin::request::AdminAppState;
use crate::GatewayError;
use aether_contracts::ProxySnapshot;
use aether_data::repository::provider_oauth::{
StoredAdminProviderOAuthDeviceSession, KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS,
};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
};
use aether_oauth::core::OAuthError;
use aether_oauth::provider::providers::{
XaiDevicePollOutcome, XaiProviderOAuthAdapter, XAI_CLIENT_ID, XAI_DEVICE_CODE_URL,
XAI_TOKEN_URL,
};
use aether_oauth::provider::ProviderOAuthTransportContext;
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::{json, Value};
pub(super) async fn handle_admin_provider_oauth_xai_device_authorize(
state: &AdminAppState<'_>,
provider_id: &str,
provider: &StoredProviderCatalogProvider,
principal: &GatewayAdminPrincipalContext,
runtime_endpoint: Option<&StoredProviderCatalogEndpoint>,
request_proxy: Option<ProxySnapshot>,
proxy_node_id: Option<&str>,
) -> Result<Response<Body>, GatewayError> {
let device_url = state.provider_oauth_token_url("xai_device", XAI_DEVICE_CODE_URL);
let token_url = state.provider_oauth_token_url("xai", XAI_TOKEN_URL);
let adapter =
XaiProviderOAuthAdapter::default().with_endpoint_overrides(&device_url, &token_url);
let ctx = ProviderOAuthTransportContext {
provider_id: provider_id.to_string(),
provider_type: provider.provider_type.clone(),
endpoint_id: runtime_endpoint.map(|endpoint| endpoint.id.clone()),
key_id: None,
auth_type: Some("oauth".to_string()),
decrypted_api_key: None,
decrypted_auth_config: None,
provider_config: provider.config.clone(),
endpoint_config: runtime_endpoint.and_then(|endpoint| endpoint.config.clone()),
key_config: None,
network: aether_oauth::network::OAuthNetworkContext::provider_operation(
request_proxy.clone(),
),
};
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
let authorization = match adapter.start_device_flow(&executor, &ctx).await {
Ok(authorization) => authorization,
Err(error) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
sanitize_xai_oauth_error(&error),
));
}
};
let now_unix_secs = current_unix_secs();
let session_id = generate_provider_oauth_nonce();
let session = StoredAdminProviderOAuthDeviceSession {
session_id: session_id.clone(),
provider_id: provider_id.to_string(),
initiated_by_user_id: principal.user_id.clone(),
initiated_by_session_id: principal.session_id.clone(),
initiated_by_management_token_id: principal.management_token_id.clone(),
region: String::new(),
client_id: XAI_CLIENT_ID.to_string(),
client_secret: String::new(),
device_code: authorization.device_code.clone(),
auth_type: Some("device".to_string()),
social_provider: None,
code_verifier: None,
redirect_uri: Some(token_url),
machine_id: None,
interval: authorization.interval,
expires_at_unix_secs: now_unix_secs.saturating_add(authorization.expires_in),
status: "pending".to_string(),
proxy_node_id: proxy_node_id
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned),
created_at_unix_ms: now_unix_secs,
key_id: None,
email: None,
replaced: false,
error_msg: None,
};
if let Err(response) = state
.save_provider_oauth_device_session(
&session_id,
&session,
authorization
.expires_in
.saturating_add(KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS),
)
.await
{
return Ok(response);
}
Ok(Json(json!({
"session_id": session_id,
"user_code": authorization.user_code,
"verification_uri": authorization.verification_uri,
"verification_uri_complete": authorization.verification_uri_complete,
"expires_in": authorization.expires_in,
"interval": authorization.interval,
"auth_type": "device",
}))
.into_response())
}
pub(super) async fn handle_admin_provider_oauth_xai_device_poll(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoints: &[StoredProviderCatalogEndpoint],
request_proxy: Option<ProxySnapshot>,
session_id: &str,
mut session: StoredAdminProviderOAuthDeviceSession,
) -> Result<Response<Body>, GatewayError> {
let token_url = session
.redirect_uri
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.unwrap_or_else(|| state.provider_oauth_token_url("xai", XAI_TOKEN_URL));
let adapter =
XaiProviderOAuthAdapter::default().with_endpoint_overrides(XAI_DEVICE_CODE_URL, token_url);
let ctx = ProviderOAuthTransportContext {
provider_id: provider.id.clone(),
provider_type: provider.provider_type.clone(),
endpoint_id: None,
key_id: None,
auth_type: Some("oauth".to_string()),
decrypted_api_key: None,
decrypted_auth_config: None,
provider_config: provider.config.clone(),
endpoint_config: None,
key_config: None,
network: aether_oauth::network::OAuthNetworkContext::provider_operation(
request_proxy.clone(),
),
};
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
let outcome = match adapter
.poll_device_token(&executor, &ctx, &session.device_code)
.await
{
Ok(outcome) => outcome,
Err(error) => {
return Ok(xai_device_poll_terminal_from_error(
state,
session_id,
&mut session,
&error,
)
.await);
}
};
match outcome {
XaiDevicePollOutcome::Pending => {
Ok(Json(json!({"status": "pending", "replaced": false})).into_response())
}
XaiDevicePollOutcome::SlowDown => {
Ok(Json(json!({"status": "slow_down", "replaced": false})).into_response())
}
XaiDevicePollOutcome::Authorized(result) => {
persist_xai_device_authorization(
state,
provider,
endpoints,
request_proxy,
session_id,
session,
*result,
)
.await
}
}
}
async fn persist_xai_device_authorization(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoints: &[StoredProviderCatalogEndpoint],
request_proxy: Option<ProxySnapshot>,
session_id: &str,
mut session: StoredAdminProviderOAuthDeviceSession,
result: aether_oauth::provider::ProviderOAuthTokenSet,
) -> Result<Response<Body>, GatewayError> {
let access_token = result.token_set.access_token.trim().to_string();
if access_token.is_empty() {
return Ok(Json(json!({
"status": "error",
"error": "xAI token 响应缺少 access_token",
"replaced": false,
}))
.into_response());
}
let mut auth_config = result.auth_config.as_object().cloned().unwrap_or_default();
auth_config.insert("provider_type".to_string(), json!("xai"));
auth_config.insert("auth_method".to_string(), json!("oauth"));
auth_config.insert("using_api".to_string(), json!(false));
let duplicate = match state
.find_duplicate_provider_oauth_key(&provider.id, &auth_config, None)
.await
{
Ok(duplicate) => duplicate,
Err(detail) => {
return Ok(Json(json!({
"status": "error",
"error": detail,
"replaced": false,
}))
.into_response());
}
};
let api_formats = provider_oauth_active_api_formats(endpoints);
let key_proxy = provider_oauth_key_proxy_value(session.proxy_node_id.as_deref());
let expires_at = result.token_set.expires_at_unix_secs;
let email = auth_config
.get("email")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let mut replaced = false;
let persisted_key = if let Some(existing_key) = duplicate {
replaced = true;
match state
.update_existing_provider_oauth_catalog_key(
&existing_key,
&provider.provider_type,
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
expires_at,
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
} else {
let key_name = admin_provider_oauth_key_name_from_auth_config(
&provider.provider_type,
&auth_config,
None,
);
match state
.create_provider_oauth_catalog_key(
&provider.id,
&provider.provider_type,
&key_name,
&access_token,
&auth_config,
&api_formats,
key_proxy,
expires_at,
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
};
spawn_provider_oauth_account_state_refresh_after_update(
state.cloned_app(),
provider.clone(),
persisted_key.id.clone(),
request_proxy.clone(),
);
session.status = "authorized".to_string();
session.key_id = Some(persisted_key.id.clone());
session.email = email.clone();
session.replaced = replaced;
session.error_msg = None;
let _ = state
.save_provider_oauth_device_session(session_id, &session, 60)
.await;
Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"authorized",
Json(json!({
"status": "authorized",
"key_id": persisted_key.id,
"email": email,
"replaced": replaced,
}))
.into_response(),
))
}
async fn xai_device_poll_terminal_from_error(
state: &AdminAppState<'_>,
session_id: &str,
session: &mut StoredAdminProviderOAuthDeviceSession,
error: &OAuthError,
) -> Response<Body> {
let (status, message) = match error {
OAuthError::InvalidRequest(detail) if detail.to_ascii_lowercase().contains("expired") => {
("expired", "设备码已过期".to_string())
}
OAuthError::InvalidRequest(detail) if detail.to_ascii_lowercase().contains("denied") => {
("error", "用户拒绝授权".to_string())
}
_ => ("error", sanitize_xai_oauth_error(error)),
};
session.status = status.to_string();
session.error_msg = Some(message.clone());
let _ = state
.save_provider_oauth_device_session(session_id, session, 30)
.await;
attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
status,
Json(json!({
"status": status,
"error": message,
"replaced": false,
}))
.into_response(),
)
}
fn sanitize_xai_oauth_error(error: &OAuthError) -> String {
match error {
OAuthError::InvalidRequest(_) => "xAI 设备授权失败: 请求参数无效".to_string(),
OAuthError::HttpStatus { status_code, .. } => {
format!("xAI 设备授权失败: HTTP {status_code}")
}
_ => "xAI 设备授权失败".to_string(),
}
}
@@ -715,7 +715,7 @@ async fn resolve_admin_provider_oauth_single_import_tokens(
if !provider_type_supports_access_token_import(provider_type) { if !provider_type_supports_access_token_import(provider_type) {
return Err(build_internal_control_error_response( return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST, http::StatusCode::BAD_REQUEST,
"Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok Provider", "Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok / xAI Provider",
)); ));
} }
@@ -867,7 +867,7 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
flatten_claude_code_credentials_payload(&mut raw_payload); flatten_claude_code_credentials_payload(&mut raw_payload);
} }
let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken"); let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken");
let access_token_input = import_payload_string_any( let mut access_token_input = import_payload_string_any(
&raw_payload, &raw_payload,
&[ &[
"access_token", "access_token",
@@ -879,6 +879,9 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
], ],
) )
.or_else(|| provider_oauth_import_authorization_bearer_token_from_object(&raw_payload)); .or_else(|| provider_oauth_import_authorization_bearer_token_from_object(&raw_payload));
if provider_type == "xai" && access_token_input.is_none() {
access_token_input = import_payload_string(&raw_payload, "api_key", "apiKey");
}
let imported_expires_at = let imported_expires_at =
import_payload_u64_any(&raw_payload, &["expires_at", "expiresAt", "expired"]); import_payload_u64_any(&raw_payload, &["expires_at", "expiresAt", "expired"]);
let (refresh_token_input, access_token_input) = normalize_provider_import_tokens( let (refresh_token_input, access_token_input) = normalize_provider_import_tokens(
@@ -901,7 +904,11 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
if !create_agent_identity && refresh_token_input.is_none() && access_token_input.is_none() { if !create_agent_identity && refresh_token_input.is_none() && access_token_input.is_none() {
return Ok(build_internal_control_error_response( return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST, http::StatusCode::BAD_REQUEST,
"Refresh Token、Access Token 或 sso_token 不能为空", if provider_type == "xai" {
"Refresh Token、Access Token 或 api_key 不能为空"
} else {
"Refresh Token、Access Token 或 sso_token 不能为空"
},
)); ));
} }
if !is_fixed_provider_type_for_provider_oauth(&provider_type) { if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
@@ -70,6 +70,12 @@ pub(super) async fn handle_admin_provider_oauth_start_key(
"Windsurf 请使用浏览器登录或导入凭据。", "Windsurf 请使用浏览器登录或导入凭据。",
)); ));
} }
if provider_type == "xai" {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"xAI 请使用设备授权或导入凭据。",
));
}
let Some(template) = admin_provider_oauth_template(&provider_type) else { let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_internal_control_error_response( return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST, http::StatusCode::BAD_REQUEST,
@@ -167,6 +173,12 @@ pub(super) async fn handle_admin_provider_oauth_start_provider(
"Windsurf 请使用浏览器登录或导入凭据。", "Windsurf 请使用浏览器登录或导入凭据。",
)); ));
} }
if provider_type == "xai" {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"xAI 请使用设备授权或导入凭据。",
));
}
let Some(template) = admin_provider_oauth_template(&provider_type) else { let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_internal_control_error_response( return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST, http::StatusCode::BAD_REQUEST,
@@ -121,6 +121,9 @@ pub(super) fn normalize_provider_import_tokens(
if provider_type == "grok" { if provider_type == "grok" {
return (None, access_token.or(refresh_token)); return (None, access_token.or(refresh_token));
} }
if provider_type == "xai" {
return (refresh_token, access_token);
}
if provider_type == "claude_code" { if provider_type == "claude_code" {
if access_token.is_none() && refresh_token.as_deref().is_some_and(is_claude_access_token) { if access_token.is_none() && refresh_token.as_deref().is_some_and(is_claude_access_token) {
return (None, refresh_token); return (None, refresh_token);
@@ -237,7 +240,7 @@ pub(super) fn provider_oauth_import_authorization_bearer_token_from_object(
pub(super) fn provider_type_supports_access_token_import(provider_type: &str) -> bool { pub(super) fn provider_type_supports_access_token_import(provider_type: &str) -> bool {
matches!( matches!(
provider_type.trim().to_ascii_lowercase().as_str(), provider_type.trim().to_ascii_lowercase().as_str(),
"claude_code" | "codex" | "chatgpt_web" | "grok" "claude_code" | "codex" | "chatgpt_web" | "grok" | "xai"
) )
} }
@@ -331,6 +334,15 @@ pub(super) fn build_provider_access_token_import_auth_config(
auth_config.insert("sso_token".to_string(), json!(access_token)); auth_config.insert("sso_token".to_string(), json!(access_token));
auth_config.insert("auth_method".to_string(), json!("sso_token")); auth_config.insert("auth_method".to_string(), json!("sso_token"));
} }
if provider_type.trim().eq_ignore_ascii_case("xai") {
if refresh_token.is_some() {
auth_config.insert("auth_method".to_string(), json!("oauth"));
auth_config.insert("using_api".to_string(), json!(false));
} else {
auth_config.insert("auth_method".to_string(), json!("api_key"));
auth_config.insert("using_api".to_string(), json!(true));
}
}
auth_config.insert( auth_config.insert(
"access_token_import_temporary".to_string(), "access_token_import_temporary".to_string(),
@@ -532,6 +544,41 @@ mod tests {
); );
} }
#[test]
fn normalize_xai_import_keeps_refresh_token_separate_from_api_key() {
let (refresh_token, access_token) =
normalize_provider_import_tokens("xai", Some("xai-refresh-token"), None);
assert_eq!(refresh_token.as_deref(), Some("xai-refresh-token"));
assert!(access_token.is_none());
let (refresh_token, access_token) =
normalize_provider_import_tokens("xai", None, Some("xai-api-key"));
assert!(refresh_token.is_none());
assert_eq!(access_token.as_deref(), Some("xai-api-key"));
}
#[test]
fn builds_xai_auth_config_from_api_key_and_oauth_tokens() {
let (api_key_config, _) =
build_provider_access_token_import_auth_config("xai", "xai-api-key", None, None, None);
assert_eq!(api_key_config.get("auth_method"), Some(&json!("api_key")));
assert_eq!(api_key_config.get("using_api"), Some(&json!(true)));
let (oauth_config, _) = build_provider_access_token_import_auth_config(
"xai",
"xai-access-token",
Some("xai-refresh-token"),
None,
None,
);
assert_eq!(oauth_config.get("auth_method"), Some(&json!("oauth")));
assert_eq!(oauth_config.get("using_api"), Some(&json!(false)));
assert_eq!(
oauth_config.get("refresh_token"),
Some(&json!("xai-refresh-token"))
);
}
#[test] #[test]
fn flattens_only_claude_ai_oauth_credentials_and_converts_expiry_ms() { fn flattens_only_claude_ai_oauth_credentials_and_converts_expiry_ms() {
let mut payload = json!({ let mut payload = json!({
@@ -8,6 +8,7 @@ use super::gemini_cli::refresh_gemini_cli_provider_quota_locally;
use super::grok::refresh_grok_provider_quota_locally; use super::grok::refresh_grok_provider_quota_locally;
use super::kiro::refresh_kiro_provider_quota_locally; use super::kiro::refresh_kiro_provider_quota_locally;
use super::windsurf::refresh_windsurf_provider_quota_locally; use super::windsurf::refresh_windsurf_provider_quota_locally;
use super::xai::refresh_xai_provider_quota_locally;
use crate::handlers::admin::request::AdminAppState; use crate::handlers::admin::request::AdminAppState;
use crate::GatewayError; use crate::GatewayError;
use aether_contracts::ProxySnapshot; use aether_contracts::ProxySnapshot;
@@ -43,6 +44,7 @@ const PROVIDER_QUOTA_REFRESH_HANDLERS: &[(&str, ProviderQuotaRefreshHandler)] =
("grok", refresh_grok_provider_quota_locally_boxed), ("grok", refresh_grok_provider_quota_locally_boxed),
("kiro", refresh_kiro_provider_quota_locally_boxed), ("kiro", refresh_kiro_provider_quota_locally_boxed),
("windsurf", refresh_windsurf_provider_quota_locally_boxed), ("windsurf", refresh_windsurf_provider_quota_locally_boxed),
("xai", refresh_xai_provider_quota_locally_boxed),
]; ];
pub(crate) async fn refresh_provider_pool_quota_locally( pub(crate) async fn refresh_provider_pool_quota_locally(
@@ -174,3 +176,19 @@ fn refresh_windsurf_provider_quota_locally_boxed<'a>(
proxy_override, proxy_override,
)) ))
} }
fn refresh_xai_provider_quota_locally_boxed<'a>(
state: &'a AdminAppState<'a>,
provider: &'a StoredProviderCatalogProvider,
endpoint: &'a StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
proxy_override: Option<ProxySnapshot>,
) -> ProviderQuotaRefreshFuture<'a> {
Box::pin(refresh_xai_provider_quota_locally(
state,
provider,
endpoint,
keys,
proxy_override,
))
}
@@ -7,3 +7,4 @@ pub(crate) mod grok;
pub(crate) mod kiro; pub(crate) mod kiro;
pub(crate) mod shared; pub(crate) mod shared;
pub(crate) mod windsurf; pub(crate) mod windsurf;
pub(crate) mod xai;
@@ -1715,6 +1715,7 @@ fn provider_quota_url_has_allowed_origin(provider_name: &str, value: &str) -> bo
"gemini_cli" => host == "cloudcode-pa.googleapis.com", "gemini_cli" => host == "cloudcode-pa.googleapis.com",
"chatgpt_web" | "codex" => host == "chatgpt.com", "chatgpt_web" | "codex" => host == "chatgpt.com",
"grok" => host == "grok.com", "grok" => host == "grok.com",
"xai" => host == "cli-chat-proxy.grok.com",
"windsurf" => host == "server.codeium.com", "windsurf" => host == "server.codeium.com",
"kiro" => kiro_quota_host_is_allowed(host), "kiro" => kiro_quota_host_is_allowed(host),
_ => false, _ => false,
@@ -1814,6 +1815,14 @@ mod tests {
), ),
("codex", "https://chatgpt.com/backend-api/wham/usage"), ("codex", "https://chatgpt.com/backend-api/wham/usage"),
("grok", "https://grok.com/rest/rate-limits"), ("grok", "https://grok.com/rest/rate-limits"),
(
"xai",
"https://cli-chat-proxy.grok.com/v1/billing?format=credits",
),
(
"xai",
"https://cli-chat-proxy.grok.com/v1/user",
),
( (
"windsurf", "windsurf",
"https://server.codeium.com/exa.seat_management_pb.SeatManagementService/GetUserStatus", "https://server.codeium.com/exa.seat_management_pb.SeatManagementService/GetUserStatus",
@@ -1847,6 +1856,11 @@ mod tests {
"https://chatgpt.com.attacker.test/backend-api/wham/usage", "https://chatgpt.com.attacker.test/backend-api/wham/usage",
), ),
("grok", "https://grok.com.attacker.test/rest/rate-limits"), ("grok", "https://grok.com.attacker.test/rest/rate-limits"),
(
"xai",
"https://cli-chat-proxy.grok.com.attacker.test/v1/billing",
),
("xai", "https://api.x.ai/v1/billing?format=credits"),
("windsurf", "https://server.codeium.com.attacker.test/quota"), ("windsurf", "https://server.codeium.com.attacker.test/quota"),
( (
"gemini_cli", "gemini_cli",
@@ -0,0 +1,313 @@
use super::shared::{
build_provider_quota_execution_plan, build_quota_snapshot_payload,
default_provider_quota_execution_timeouts, execute_provider_quota_plan,
extract_execution_error_message, oauth_refresh_auto_removed_result,
persist_provider_quota_refresh_state, quota_key_auto_removed,
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::GatewayError;
use aether_admin::provider::quota::parse_xai_billing_response;
use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json;
use aether_contracts::ProxySnapshot;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_provider_pool::{build_xai_pool_billing_request, build_xai_pool_user_request};
use aether_provider_transport::xai::{
extract_xai_user_id_from_auth_config, extract_xai_user_id_from_value, xai_auth_uses_api,
};
use serde_json::{json, Value};
use std::time::{SystemTime, UNIX_EPOCH};
async fn execute_xai_quota_plan(
state: &AdminAppState<'_>,
transport: &AdminGatewayProviderTransportSnapshot,
spec: aether_provider_pool::ProviderPoolQuotaRequestSpec,
proxy_override: Option<&ProxySnapshot>,
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
let proxy = match proxy_override {
Some(proxy) => Some(proxy.clone()),
None => {
state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await
}
};
let timeouts = state
.resolve_transport_execution_timeouts(transport)
.or(Some(default_provider_quota_execution_timeouts(
proxy.as_ref(),
)));
let plan = build_provider_quota_execution_plan(
transport,
spec,
proxy,
state.resolve_transport_profile(transport),
timeouts,
);
execute_provider_quota_plan(state, transport, plan, "xai").await
}
fn xai_authorization_from_header(authorization: &(String, String)) -> (String, String) {
authorization.clone()
}
fn enrich_xai_subscription_title(mut metadata: Value, auth_config: Option<&str>) -> Value {
if metadata
.get("subscription_title")
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty())
{
return metadata;
}
let Some(config) = auth_config
.map(str::trim)
.filter(|value| !value.is_empty())
.and_then(|value| serde_json::from_str::<Value>(value).ok())
else {
return metadata;
};
let title = ["subscription_tier", "subscriptionTier", "tier", "plan"]
.iter()
.find_map(|field| {
config
.get(*field)
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
});
if let Some(title) = title {
if let Some(object) = metadata.as_object_mut() {
object.insert("subscription_title".to_string(), json!(title));
}
}
metadata
}
pub(crate) async fn refresh_xai_provider_quota_locally(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
proxy_override: Option<ProxySnapshot>,
) -> Result<Option<serde_json::Value>, GatewayError> {
let mut results = Vec::new();
let mut success_count = 0usize;
let mut failed_count = 0usize;
let mut auto_removed_count = 0usize;
for key in keys {
let transport = match state
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
.await?
{
Some(transport) => transport,
None => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Provider transport snapshot unavailable",
}));
continue;
}
};
if xai_auth_uses_api(
transport.key.auth_type.as_str(),
transport.key.decrypted_auth_config.as_deref(),
) {
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "skipped",
"message": "xAI API Key 账号没有 Grok Build 订阅额度接口,请使用设备授权账号查询额度。",
}));
continue;
}
let authorization = match state.resolve_local_oauth_header_auth(&transport).await? {
Some(auth) => auth,
_ => {
if quota_key_auto_removed(state, &key.id).await? {
auto_removed_count += 1;
results.push(oauth_refresh_auto_removed_result(&key));
continue;
}
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 OAuth 认证信息,请先授权/刷新 Token",
}));
continue;
}
};
let fallback_user_id =
extract_xai_user_id_from_auth_config(transport.key.decrypted_auth_config.as_deref());
let user_id = match execute_xai_quota_plan(
state,
&transport,
build_xai_pool_user_request(
&transport.key.id,
xai_authorization_from_header(&authorization),
),
proxy_override.as_ref(),
)
.await?
{
ProviderQuotaExecutionOutcome::Response(result) if result.status_code == 200 => result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
.and_then(extract_xai_user_id_from_value)
.or(fallback_user_id),
_ => fallback_user_id,
};
let result = match execute_xai_quota_plan(
state,
&transport,
build_xai_pool_billing_request(
&transport.key.id,
xai_authorization_from_header(&authorization),
user_id.as_deref(),
),
proxy_override.as_ref(),
)
.await?
{
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(_) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "xAI billing 请求执行失败",
"status_code": 502,
}));
continue;
}
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut metadata_update = None::<serde_json::Value>;
let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) =
quota_refresh_success_invalid_state(&key);
let mut status = "error".to_string();
let mut message = None::<String>;
if result.status_code == 200 {
if let Some(body_json) = result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
{
metadata_update =
parse_xai_billing_response(body_json, now_unix_secs).map(|metadata| {
json!({
"xai": enrich_xai_subscription_title(
metadata,
transport.key.decrypted_auth_config.as_deref(),
)
})
});
if metadata_update.is_some() {
status = "success".to_string();
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含可用的 Grok Build 额度信息".to_string());
}
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含配额信息".to_string());
}
} else {
message = Some(
extract_execution_error_message(&result)
.unwrap_or_else(|| format!("xAI billing 返回状态码 {}", result.status_code)),
);
if result.status_code == 401 || result.status_code == 403 {
let reason = message
.clone()
.unwrap_or_else(|| "账户访问被禁止".to_string());
oauth_invalid_at_unix_secs = Some(now_unix_secs);
oauth_invalid_reason = Some(format!("账户访问被禁止: {reason}"));
status = if result.status_code == 401 {
"unauthorized".to_string()
} else {
"forbidden".to_string()
};
}
}
if !persist_provider_quota_refresh_state(
state,
&key.id,
metadata_update.as_ref(),
oauth_invalid_at_unix_secs,
oauth_invalid_reason,
None,
)
.await?
{
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Key 状态写入失败",
}));
continue;
}
if status == "success" {
success_count += 1;
} else {
failed_count += 1;
}
let mut payload = serde_json::Map::new();
payload.insert("key_id".to_string(), json!(key.id));
payload.insert("key_name".to_string(), json!(key.name));
payload.insert("status".to_string(), json!(status));
if let Some(message) = message {
payload.insert("message".to_string(), json!(message));
}
if let Some(metadata) = metadata_update.as_ref().and_then(|value| value.get("xai")) {
payload.insert(
"metadata".to_string(),
admin_provider_metadata_bucket_safe_json("xai", Some(metadata)),
);
}
if let Some(quota_snapshot) = build_quota_snapshot_payload(
"xai",
key.status_snapshot.as_ref(),
metadata_update.as_ref(),
) {
payload.insert("quota_snapshot".to_string(), quota_snapshot);
}
results.push(serde_json::Value::Object(payload));
}
Ok(Some(json!({
"success": success_count,
"failed": failed_count,
"total": results.len(),
"results": results,
"message": format!("已处理 {} 个 Key", results.len()),
"auto_removed": auto_removed_count,
})))
}
@@ -932,6 +932,13 @@ fn admin_pool_build_account_quota(
return Some(account_quota); return Some(account_quota);
} }
} }
"xai" => {
if let Some(account_quota) =
admin_pool_build_kiro_account_quota_from_snapshot(quota_snapshot)
{
return Some(account_quota);
}
}
"chatgpt_web" => { "chatgpt_web" => {
if let Some(account_quota) = if let Some(account_quota) =
admin_pool_build_chatgpt_web_account_quota_from_snapshot(quota_snapshot) admin_pool_build_chatgpt_web_account_quota_from_snapshot(quota_snapshot)
@@ -1591,4 +1598,29 @@ mod tests {
Some("Auto剩余 40.0% (60/150) | Heavy剩余 0.0% (0/20)".to_string()) Some("Auto剩余 40.0% (60/150) | Heavy剩余 0.0% (0/20)".to_string())
); );
} }
#[test]
fn xai_account_quota_is_rendered_as_remaining_percent() {
let quota_snapshot = json!({
"provider_type": "xai",
"code": "ok",
"exhausted": false,
"plan_type": "SuperGrok",
"windows": [
{
"code": "usage",
"label": "周额度",
"scope": "account",
"used_ratio": 0.46,
"remaining_ratio": 0.54
}
]
});
let quota_snapshot = quota_snapshot.as_object().unwrap();
assert_eq!(
admin_pool_build_account_quota("xai", Some(quota_snapshot)),
Some("剩余 54.0%".to_string())
);
}
} }
@@ -3511,6 +3511,11 @@ async fn provider_query_execute_standard_test_candidate(
codex_model_capabilities.as_ref(), codex_model_capabilities.as_ref(),
); );
} }
crate::provider_transport::insert_cli_identity_headers_if_needed(
&transport,
provider_api_format,
&mut request_headers,
);
if !uses_vertex_query_auth { if !uses_vertex_query_auth {
if let (Some(auth_header), Some(auth_value)) = if let (Some(auth_header), Some(auth_value)) =
(auth_header.as_deref(), auth_value.as_deref()) (auth_header.as_deref(), auth_value.as_deref())
@@ -4,9 +4,9 @@ pub(crate) fn normalize_provider_type_input(value: &str) -> Result<String, Strin
let normalized = value.trim().to_ascii_lowercase(); let normalized = value.trim().to_ascii_lowercase();
match normalized.as_str() { match normalized.as_str() {
"custom" | "claude_code" | "kiro" | "codex" | "chatgpt_web" | "gemini_cli" "custom" | "claude_code" | "kiro" | "codex" | "chatgpt_web" | "gemini_cli"
| "antigravity" | "vertex_ai" | "grok" | "windsurf" => Ok(normalized), | "antigravity" | "vertex_ai" | "grok" | "windsurf" | "xai" => Ok(normalized),
_ => Err( _ => Err(
"provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai / grok / windsurf" "provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai / grok / windsurf / xai"
.to_string(), .to_string(),
), ),
} }
@@ -405,6 +405,14 @@ mod tests {
); );
} }
#[test]
fn normalize_provider_type_supports_xai() {
assert_eq!(
normalize_provider_type_input(" xAI ").expect("type should normalize"),
"xai"
);
}
#[test] #[test]
fn normalize_api_format_list_dedupes_canonical_formats() { fn normalize_api_format_list_dedupes_canonical_formats() {
assert_eq!( assert_eq!(
@@ -90,6 +90,8 @@ pub(super) struct ResponsesWebSocketContinuationRecord {
/// request JSON can never set it. /// request JSON can never set it.
#[serde(default)] #[serde(default)]
deepseek_opaque_reasoning_replay: bool, deepseek_opaque_reasoning_replay: bool,
#[serde(default)]
xai_encrypted_reasoning_replay: bool,
/// A prior turn stored PII sentinels whose restore mapping exists only on /// A prior turn stored PII sentinels whose restore mapping exists only on
/// the original downstream socket. Such a chain cannot safely resume on a /// the original downstream socket. Such a chain cannot safely resume on a
/// new socket without leaking sentinels, so lookup succeeds but bootstrap /// new socket without leaking sentinels, so lookup succeeds but bootstrap
@@ -122,6 +124,10 @@ impl ResponsesWebSocketContinuationRecord {
normalization.reasoning_replay_policy(), normalization.reasoning_replay_policy(),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
), ),
xai_encrypted_reasoning_replay: matches!(
normalization.reasoning_replay_policy(),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
),
has_connection_local_redaction, has_connection_local_redaction,
responses_lite_static_config, responses_lite_static_config,
}; };
@@ -156,7 +162,9 @@ impl ResponsesWebSocketContinuationRecord {
pub(super) fn reasoning_replay_policy( pub(super) fn reasoning_replay_policy(
&self, &self,
) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy { ) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy {
if self.deepseek_opaque_reasoning_replay { if self.xai_encrypted_reasoning_replay {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
} else if self.deepseek_opaque_reasoning_replay {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
} else { } else {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
@@ -476,6 +484,7 @@ mod tests {
binding_fingerprint: [7; 32], binding_fingerprint: [7; 32],
normalization_fingerprint: [9; 32], normalization_fingerprint: [9; 32],
deepseek_opaque_reasoning_replay: false, deepseek_opaque_reasoning_replay: false,
xai_encrypted_reasoning_replay: false,
has_connection_local_redaction: false, has_connection_local_redaction: false,
responses_lite_static_config: Some(ResponsesLiteStaticConfig::from_response_create( responses_lite_static_config: Some(ResponsesLiteStaticConfig::from_response_create(
&json!({ &json!({
@@ -714,6 +723,29 @@ mod tests {
assert_eq!(decoded, record()); assert_eq!(decoded, record());
} }
#[test]
fn serialized_record_preserves_xai_replay_policy_and_reads_legacy_records() {
let mut expected = record();
expected.xai_encrypted_reasoning_replay = true;
let mut serialized = serde_json::to_value(&expected).unwrap();
let decoded: ResponsesWebSocketContinuationRecord =
serde_json::from_value(serialized.clone()).unwrap();
assert_eq!(
decoded.reasoning_replay_policy(),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
);
serialized
.as_object_mut()
.unwrap()
.remove("xai_encrypted_reasoning_replay");
let legacy: ResponsesWebSocketContinuationRecord =
serde_json::from_value(serialized).unwrap();
assert_eq!(
legacy.reasoning_replay_policy(),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
);
}
#[test] #[test]
fn serialized_record_preserves_only_the_server_derived_reasoning_replay_policy_bit() { fn serialized_record_preserves_only_the_server_derived_reasoning_replay_policy_bit() {
let mut expected = record(); let mut expected = record();
@@ -1388,6 +1388,168 @@ fn build_kiro_quota_status_snapshot(
})) }))
} }
fn build_xai_quota_status_snapshot(
upstream_metadata: Option<&Value>,
source: &str,
) -> Option<Value> {
let metadata = provider_quota_metadata_bucket(upstream_metadata, "xai")?;
let observed_at_unix_secs = provider_quota_timestamp_unix_secs(metadata.get("updated_at"));
let usage_limit = metadata
.get("usage_limit")
.and_then(admin_provider_quota_pure::coerce_json_f64);
let current_usage = metadata
.get("current_usage")
.and_then(admin_provider_quota_pure::coerce_json_f64);
let remaining = metadata
.get("remaining")
.and_then(admin_provider_quota_pure::coerce_json_f64);
let usage_ratio = metadata
.get("usage_percentage")
.and_then(admin_provider_quota_pure::coerce_json_f64)
.map(|value| (value / 100.0).clamp(0.0, 1.0))
.or_else(|| {
current_usage
.zip(usage_limit)
.and_then(|(current_usage, usage_limit)| {
(usage_limit > 0.0).then_some((current_usage / usage_limit).clamp(0.0, 1.0))
})
});
let remaining_ratio = usage_ratio.map(|value| (1.0 - value).max(0.0));
let next_reset_at = provider_quota_timestamp_unix_secs(metadata.get("next_reset_at"));
let reset_seconds = quota_window_reset_seconds(observed_at_unix_secs, next_reset_at);
let plan_type = metadata
.get("subscription_title")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let period_type = metadata
.get("period_type")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let usage_label = match period_type.as_deref() {
Some("monthly") => "月额度",
Some("weekly") => "周额度",
_ => "额度",
};
let mut windows = Vec::new();
if usage_ratio.is_some()
|| remaining.is_some()
|| usage_limit.is_some()
|| current_usage.is_some()
|| next_reset_at.is_some()
{
windows.push(json!({
"code": "usage",
"label": usage_label,
"scope": "account",
"unit": if usage_limit.is_some() { "usd" } else { "percent" },
"used_ratio": usage_ratio,
"remaining_ratio": remaining_ratio,
"used_value": current_usage,
"remaining_value": remaining,
"limit_value": usage_limit,
"reset_at": next_reset_at,
"reset_seconds": reset_seconds,
}));
}
let prepaid_balance = metadata
.get("prepaid_balance")
.and_then(admin_provider_quota_pure::coerce_json_f64);
if prepaid_balance.is_some_and(|value| value > 0.0) {
windows.push(json!({
"code": "prepaid",
"label": "预付额度",
"scope": "account",
"unit": "usd",
"used_ratio": serde_json::Value::Null,
"remaining_ratio": serde_json::Value::Null,
"remaining_value": prepaid_balance,
"reset_at": serde_json::Value::Null,
"reset_seconds": serde_json::Value::Null,
}));
}
let on_demand_cap = metadata
.get("on_demand_cap")
.and_then(admin_provider_quota_pure::coerce_json_f64);
let on_demand_used = metadata
.get("on_demand_used")
.and_then(admin_provider_quota_pure::coerce_json_f64);
let on_demand_enabled = metadata
.get("on_demand_enabled")
.and_then(admin_provider_quota_pure::coerce_json_bool)
!= Some(false);
if on_demand_enabled && on_demand_cap.is_some_and(|value| value > 0.0) {
let on_demand_remaining = on_demand_cap
.zip(on_demand_used)
.map(|(cap, used)| (cap - used).max(0.0));
let on_demand_ratio = on_demand_cap
.zip(on_demand_used)
.and_then(|(cap, used)| (cap > 0.0).then_some((used / cap).clamp(0.0, 1.0)));
windows.push(json!({
"code": "on_demand",
"label": "按需额度",
"scope": "account",
"unit": "usd",
"used_ratio": on_demand_ratio,
"remaining_ratio": on_demand_ratio.map(|value| (1.0 - value).max(0.0)),
"used_value": on_demand_used,
"remaining_value": on_demand_remaining,
"limit_value": on_demand_cap,
"reset_at": serde_json::Value::Null,
"reset_seconds": serde_json::Value::Null,
}));
}
if windows.is_empty() && plan_type.is_none() && observed_at_unix_secs.is_none() {
return None;
}
let prepaid_available = prepaid_balance.is_some_and(|value| value > 0.0);
let on_demand_available = on_demand_enabled
&& on_demand_cap.is_some_and(|value| value > 0.0)
&& on_demand_used
.zip(on_demand_cap)
.is_some_and(|(used, cap)| used < cap);
let usage_exhausted = remaining.is_some_and(|value| value <= 0.0)
|| usage_ratio.is_some_and(|value| value >= 1.0 - 1e-6);
let exhausted = usage_exhausted && !prepaid_available && !on_demand_available;
let reason = if exhausted {
Some("额度已耗尽".to_string())
} else {
None
};
let label = if exhausted {
Some("额度耗尽")
} else {
None
};
let code = if exhausted { "exhausted" } else { "ok" };
Some(json!({
"version": 2,
"provider_type": "xai",
"code": code,
"label": label,
"reason": reason,
"freshness": "fresh",
"source": source,
"observed_at": observed_at_unix_secs,
"exhausted": exhausted,
"usage_ratio": usage_ratio,
"updated_at": observed_at_unix_secs,
"reset_at": next_reset_at,
"reset_seconds": reset_seconds,
"plan_type": plan_type,
"windows": windows,
}))
}
fn build_chatgpt_web_quota_status_snapshot( fn build_chatgpt_web_quota_status_snapshot(
upstream_metadata: Option<&Value>, upstream_metadata: Option<&Value>,
source: &str, source: &str,
@@ -2255,6 +2417,7 @@ pub(crate) fn sync_provider_key_quota_status_snapshot(
let mut quota = match normalized_provider_type.as_str() { let mut quota = match normalized_provider_type.as_str() {
"codex" => build_codex_quota_status_snapshot(upstream_metadata, source), "codex" => build_codex_quota_status_snapshot(upstream_metadata, source),
"kiro" => build_kiro_quota_status_snapshot(upstream_metadata, source), "kiro" => build_kiro_quota_status_snapshot(upstream_metadata, source),
"xai" => build_xai_quota_status_snapshot(upstream_metadata, source),
"chatgpt_web" => build_chatgpt_web_quota_status_snapshot(upstream_metadata, source), "chatgpt_web" => build_chatgpt_web_quota_status_snapshot(upstream_metadata, source),
"windsurf" => build_windsurf_quota_status_snapshot(upstream_metadata, source), "windsurf" => build_windsurf_quota_status_snapshot(upstream_metadata, source),
"antigravity" => build_antigravity_quota_status_snapshot(upstream_metadata, source), "antigravity" => build_antigravity_quota_status_snapshot(upstream_metadata, source),
@@ -3622,6 +3785,43 @@ mod tests {
assert_eq!(auto.get("used_value"), Some(&json!(90.0))); assert_eq!(auto.get("used_value"), Some(&json!(90.0)));
} }
#[test]
fn provider_key_status_snapshot_payload_backfills_xai_weekly_credits() {
let mut key = sample_catalog_key();
key.upstream_metadata = Some(json!({
"xai": {
"updated_at": 1_778_067_246u64,
"usage_percentage": 46.0,
"period_type": "weekly",
"next_reset_at": 1_778_157_172u64,
"subscription_title": "SuperGrok",
"prepaid_balance": 0.0,
"on_demand_cap": 0.0,
"on_demand_used": 0.0
}
}));
let payload = provider_key_status_snapshot_payload(&key, "xai");
let quota = payload
.get("quota")
.and_then(Value::as_object)
.expect("quota snapshot should be object");
let windows = quota
.get("windows")
.and_then(Value::as_array)
.expect("xai quota windows should exist");
assert_eq!(quota.get("provider_type"), Some(&json!("xai")));
assert_eq!(quota.get("code"), Some(&json!("ok")));
assert_eq!(quota.get("exhausted"), Some(&json!(false)));
assert_eq!(quota.get("plan_type"), Some(&json!("SuperGrok")));
assert_eq!(quota.get("usage_ratio"), Some(&json!(0.46)));
assert_eq!(quota.get("reset_at"), Some(&json!(1_778_157_172u64)));
assert_eq!(windows.len(), 1);
assert_eq!(windows[0].get("code"), Some(&json!("usage")));
assert_eq!(windows[0].get("label"), Some(&json!("周额度")));
}
#[test] #[test]
fn provider_key_status_snapshot_payload_backfills_gemini_cli_account_credits() { fn provider_key_status_snapshot_payload_backfills_gemini_cli_account_credits() {
let mut key = sample_catalog_key(); let mut key = sample_catalog_key();
@@ -13,7 +13,7 @@ pub(crate) fn openai_image_provider_max_generation_count(provider_type: &str) ->
GROK_OPENAI_IMAGE_MAX_GENERATION_COUNT GROK_OPENAI_IMAGE_MAX_GENERATION_COUNT
} else if matches!( } else if matches!(
provider_type.trim().to_ascii_lowercase().as_str(), provider_type.trim().to_ascii_lowercase().as_str(),
"openai" | "codex" "openai" | "codex" | "xai"
) { ) {
OPENAI_IMAGE_MAX_GENERATION_COUNT OPENAI_IMAGE_MAX_GENERATION_COUNT
} else { } else {
@@ -58,6 +58,7 @@ mod tests {
assert_eq!(openai_image_provider_max_generation_count("grok"), 4); assert_eq!(openai_image_provider_max_generation_count("grok"), 4);
assert_eq!(openai_image_provider_max_generation_count("openai"), 10); assert_eq!(openai_image_provider_max_generation_count("openai"), 10);
assert_eq!(openai_image_provider_max_generation_count("codex"), 10); assert_eq!(openai_image_provider_max_generation_count("codex"), 10);
assert_eq!(openai_image_provider_max_generation_count("xai"), 10);
assert_eq!(openai_image_provider_max_generation_count("custom"), 1); assert_eq!(openai_image_provider_max_generation_count("custom"), 1);
assert_eq!( assert_eq!(
openai_image_provider_max_generation_count_for_model("openai", Some("dall-e-3")), openai_image_provider_max_generation_count_for_model("openai", Some("dall-e-3")),
@@ -172,6 +172,7 @@ fn provider_uses_bearer_oauth_runtime(provider_type: &str) -> bool {
| "antigravity" | "antigravity"
| "kiro" | "kiro"
| "windsurf" | "windsurf"
| "xai"
) )
} }
@@ -406,6 +407,22 @@ mod tests {
); );
} }
#[test]
fn recognizes_xai_oauth_as_bearer_runtime() {
let semantics = provider_key_auth_semantics(&sample_key("oauth"), "xai");
assert!(semantics.oauth_managed());
assert!(semantics.can_refresh_oauth());
assert_eq!(
semantics.credential_kind(),
ProviderKeyCredentialKind::OAuthSession
);
assert_eq!(
semantics.runtime_auth_kind(),
ProviderKeyRuntimeAuthKind::Bearer
);
}
#[test] #[test]
fn refresh_capability_requires_stored_refresh_token() { fn refresh_capability_requires_stored_refresh_token() {
let semantics = provider_key_auth_semantics(&sample_key("oauth"), "codex"); let semantics = provider_key_auth_semantics(&sample_key("oauth"), "codex");
+2
View File
@@ -186,6 +186,8 @@ fn frontend_path_bypasses_static(path: &str) -> bool {
"/health" | "/test-connection" | crate::constants::READYZ_PATH "/health" | "/test-connection" | crate::constants::READYZ_PATH
) || path.starts_with("/api/") ) || path.starts_with("/api/")
|| path.starts_with("/v1/") || path.starts_with("/v1/")
|| path == "/openai/v1/videos"
|| path.starts_with("/openai/v1/videos/")
|| path.starts_with("/v1beta/") || path.starts_with("/v1beta/")
|| path.starts_with("/upload/") || path.starts_with("/upload/")
|| path.starts_with("/_gateway/") || path.starts_with("/_gateway/")
@@ -290,6 +290,14 @@ impl provider_transport::VideoTaskTransportSnapshotLookup for AppState {
.await .await
.map_err(GatewayError::into_message) .map_err(GatewayError::into_message)
} }
async fn resolve_video_task_proxy(
&self,
transport: &GatewayProviderTransportSnapshot,
) -> Option<ProxySnapshot> {
self.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await
}
} }
#[async_trait] #[async_trait]
@@ -1789,6 +1789,7 @@ fn admin_provider_oauth_quota_mod_stays_thin() {
"pub(crate) mod dispatch;", "pub(crate) mod dispatch;",
"pub(crate) mod kiro;", "pub(crate) mod kiro;",
"pub(crate) mod shared;", "pub(crate) mod shared;",
"pub(crate) mod xai;",
] { ] {
assert!( assert!(
quota_mod.contains(pattern), quota_mod.contains(pattern),
@@ -1861,6 +1862,7 @@ fn admin_provider_oauth_quota_mod_stays_thin() {
"refresh_antigravity_provider_quota_locally", "refresh_antigravity_provider_quota_locally",
"refresh_gemini_cli_provider_quota_locally", "refresh_gemini_cli_provider_quota_locally",
"refresh_chatgpt_web_provider_quota_locally", "refresh_chatgpt_web_provider_quota_locally",
"refresh_xai_provider_quota_locally",
] { ] {
assert!( assert!(
quota_dispatch.contains(pattern), quota_dispatch.contains(pattern),
@@ -1452,6 +1452,7 @@ fn ai_serving_planner_separates_local_candidate_resolution_from_ranking() {
"GeminiCliProviderPoolAdapter", "GeminiCliProviderPoolAdapter",
"KiroProviderPoolAdapter", "KiroProviderPoolAdapter",
"ChatGptWebProviderPoolAdapter", "ChatGptWebProviderPoolAdapter",
"XaiProviderPoolAdapter",
"CLAUDE_CODE_PROVIDER_POOL_ADAPTER", "CLAUDE_CODE_PROVIDER_POOL_ADAPTER",
"VERTEX_AI_PROVIDER_POOL_ADAPTER", "VERTEX_AI_PROVIDER_POOL_ADAPTER",
"provider_types_for_capability", "provider_types_for_capability",
@@ -1478,6 +1479,7 @@ fn ai_serving_planner_separates_local_candidate_resolution_from_ranking() {
"pub mod gemini_cli;", "pub mod gemini_cli;",
"pub mod kiro;", "pub mod kiro;",
"pub mod chatgpt_web;", "pub mod chatgpt_web;",
"pub mod xai;",
] { ] {
assert!( assert!(
provider_pool_providers.contains(pattern), provider_pool_providers.contains(pattern),
@@ -1513,6 +1515,14 @@ fn ai_serving_planner_separates_local_candidate_resolution_from_ranking() {
"crates/aether-provider/pool/src/providers/kiro.rs", "crates/aether-provider/pool/src/providers/kiro.rs",
vec!["KiroProviderPoolAdapter", "quota_exhausted_from_bucket"], vec!["KiroProviderPoolAdapter", "quota_exhausted_from_bucket"],
), ),
(
"crates/aether-provider/pool/src/providers/xai.rs",
vec![
"XaiProviderPoolAdapter",
"build_xai_pool_billing_request",
"quota_exhausted_from_bucket",
],
),
( (
"crates/aether-provider/pool/src/providers/chatgpt_web.rs", "crates/aether-provider/pool/src/providers/chatgpt_web.rs",
vec![ vec![
@@ -983,6 +983,297 @@ async fn gateway_rejects_generic_oauth_start_for_windsurf_provider_impl() {
); );
} }
#[test]
fn gateway_rejects_generic_oauth_start_for_xai_provider() {
run_admin_oauth_test(
"gateway_rejects_generic_oauth_start_for_xai_provider",
gateway_rejects_generic_oauth_start_for_xai_provider_impl,
);
}
async fn gateway_rejects_generic_oauth_start_for_xai_provider_impl() {
let mut provider = sample_provider("provider-xai", "xai", 10);
provider.provider_type = "xai".to_string();
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![],
vec![],
));
let state = AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(GatewayDataState::with_provider_catalog_reader_for_tests(
provider_catalog_repository,
));
let response = local_admin_provider_oauth_response(
&state,
http::Method::POST,
"/api/admin/provider-oauth/providers/provider-xai/start",
None,
)
.await;
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let body = to_bytes(response.into_body(), usize::MAX)
.await
.expect("body should read");
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
assert!(
payload["detail"].as_str().is_some_and(|detail| {
detail.contains("设备授权") || detail.contains("导入凭据")
}),
"payload={payload}"
);
}
#[test]
fn gateway_handles_admin_provider_oauth_device_authorize_for_xai() {
run_admin_oauth_test(
"gateway_handles_admin_provider_oauth_device_authorize_for_xai",
gateway_handles_admin_provider_oauth_device_authorize_for_xai_impl,
);
}
async fn gateway_handles_admin_provider_oauth_device_authorize_for_xai_impl() {
let authorize_hits = Arc::new(Mutex::new(0usize));
let authorize_hits_clone = Arc::clone(&authorize_hits);
let oidc_server = Router::new().fallback(any(move |_request: Request| {
let authorize_hits_inner = Arc::clone(&authorize_hits_clone);
async move {
*authorize_hits_inner.lock().expect("mutex should lock") += 1;
Json(json!({
"device_code": "xai-device-code",
"user_code": "XAI-CODE",
"verification_uri": "https://auth.x.ai/activate",
"verification_uri_complete": "https://auth.x.ai/activate?user_code=XAI-CODE",
"expires_in": 600,
"interval": 5,
}))
}
}));
let mut provider = sample_provider("provider-xai", "xai", 10);
provider.provider_type = "xai".to_string();
let endpoint = sample_endpoint(
"endpoint-xai-responses",
"provider-xai",
"openai:responses",
"https://cli-chat-proxy.grok.com/v1",
);
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
vec![],
));
let (oidc_url, oidc_handle) = start_server(oidc_server).await;
let state = AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(GatewayDataState::with_provider_catalog_reader_for_tests(
provider_catalog_repository,
))
.with_provider_oauth_token_url_for_tests(
"xai_device",
format!("{oidc_url}/oauth2/device/code"),
)
.with_provider_oauth_token_url_for_tests("xai", format!("{oidc_url}/oauth2/token"));
let response = local_admin_provider_oauth_response(
&state,
http::Method::POST,
"/api/admin/provider-oauth/providers/provider-xai/device-authorize",
Some(json!({ "proxy_node_id": "proxy-node-xai" })),
)
.await;
let status = response.status();
let body = to_bytes(response.into_body(), usize::MAX)
.await
.expect("body should read");
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
assert_eq!(status, StatusCode::OK, "payload={payload}");
let session_id = payload["session_id"]
.as_str()
.expect("session_id should exist")
.to_string();
assert_eq!(payload["user_code"], "XAI-CODE");
assert_eq!(payload["verification_uri"], "https://auth.x.ai/activate");
assert_eq!(
payload["verification_uri_complete"],
"https://auth.x.ai/activate?user_code=XAI-CODE"
);
assert_eq!(payload["auth_type"], "device");
assert!(payload.get("callback_required").is_none() || payload["callback_required"] == false);
assert_eq!(*authorize_hits.lock().expect("mutex should lock"), 1);
let stored = state
.load_provider_oauth_device_session_for_tests(&format!("device_auth_session:{session_id}"))
.expect("device session should be stored");
let stored: serde_json::Value =
serde_json::from_str(&stored).expect("device session json should parse");
assert_eq!(stored["provider_id"], "provider-xai");
assert_eq!(stored["device_code"], "xai-device-code");
assert_eq!(stored["auth_type"], "device");
assert_eq!(stored["redirect_uri"], format!("{oidc_url}/oauth2/token"));
assert_eq!(stored["proxy_node_id"], "proxy-node-xai");
assert_eq!(stored["status"], "pending");
oidc_handle.abort();
}
#[test]
fn gateway_handles_admin_provider_oauth_device_poll_for_xai() {
run_admin_oauth_test(
"gateway_handles_admin_provider_oauth_device_poll_for_xai",
gateway_handles_admin_provider_oauth_device_poll_for_xai_impl,
);
}
async fn gateway_handles_admin_provider_oauth_device_poll_for_xai_impl() {
let token_hits = Arc::new(Mutex::new(0usize));
let token_hits_clone = Arc::clone(&token_hits);
let access_token = sample_kiro_device_access_token("[email protected]");
let id_token = access_token.clone();
let token_server = Router::new().fallback(any(move |_request: Request| {
let token_hits_inner = Arc::clone(&token_hits_clone);
let access_token = access_token.clone();
let id_token = id_token.clone();
async move {
let hit = {
let mut hits = token_hits_inner.lock().expect("mutex should lock");
*hits += 1;
*hits
};
if hit == 1 {
return (
StatusCode::BAD_REQUEST,
Json(json!({ "error": "authorization_pending" })),
)
.into_response();
}
Json(json!({
"access_token": access_token,
"refresh_token": "xai-refresh-token",
"token_type": "Bearer",
"expires_in": 3600,
"id_token": id_token,
}))
.into_response()
}
}));
let mut provider = sample_provider("provider-xai", "xai", 10);
provider.provider_type = "xai".to_string();
let endpoint = sample_endpoint(
"endpoint-xai-responses",
"provider-xai",
"openai:responses",
"https://cli-chat-proxy.grok.com/v1",
);
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
vec![],
));
let (token_url, token_handle) = start_server(token_server).await;
let resolved_token_url = format!("{token_url}/oauth2/token");
let state = AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(
provider_catalog_repository.clone(),
)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
)
.with_provider_oauth_device_session_entry_for_tests(
"session-xai",
json!({
"provider_id": "provider-xai",
"region": "",
"client_id": "b1a00492-073a-47ea-816f-4c329264a828",
"client_secret": "",
"device_code": "xai-device-code",
"auth_type": "device",
"social_provider": null,
"code_verifier": null,
"redirect_uri": resolved_token_url,
"machine_id": null,
"interval": 5,
"expires_at_unix_secs": 4_102_444_800u64,
"status": "pending",
"proxy_node_id": null,
"created_at_unix_ms": 1_711_000_000u64,
"key_id": null,
"email": null,
"replaced": false,
"error_msg": null,
}),
)
.with_provider_oauth_token_url_for_tests("xai", resolved_token_url.clone());
let pending = local_admin_provider_oauth_response(
&state,
http::Method::POST,
"/api/admin/provider-oauth/providers/provider-xai/device-poll",
Some(json!({ "session_id": "session-xai" })),
)
.await;
let pending_body = to_bytes(pending.into_body(), usize::MAX)
.await
.expect("pending body should read");
let pending_payload: serde_json::Value =
serde_json::from_slice(&pending_body).expect("pending json should parse");
assert_eq!(
pending_payload["status"], "pending",
"payload={pending_payload}"
);
let authorized = local_admin_provider_oauth_response(
&state,
http::Method::POST,
"/api/admin/provider-oauth/providers/provider-xai/device-poll",
Some(json!({ "session_id": "session-xai" })),
)
.await;
let status = authorized.status();
let body = to_bytes(authorized.into_body(), usize::MAX)
.await
.expect("authorized body should read");
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
assert_eq!(status, StatusCode::OK, "payload={payload}");
assert_eq!(payload["status"], "authorized");
assert_eq!(payload["email"], "[email protected]");
assert_eq!(payload["replaced"], false);
assert_eq!(*token_hits.lock().expect("mutex should lock"), 2);
let stored = state
.load_provider_oauth_device_session_for_tests("device_auth_session:session-xai")
.expect("device session should persist");
let stored: serde_json::Value =
serde_json::from_str(&stored).expect("device session json should parse");
assert_eq!(stored["status"], "authorized");
let key_id = stored["key_id"]
.as_str()
.expect("key_id should be stored")
.to_string();
assert_eq!(payload["key_id"], key_id);
let persisted = provider_catalog_repository
.list_keys_by_ids(std::slice::from_ref(&key_id))
.await
.expect("keys should load")
.into_iter()
.next()
.expect("persisted key should exist");
assert_eq!(persisted.auth_type, "oauth");
let decrypted_auth_config = decrypt_persisted_provider_auth_config(&persisted);
let auth_config: serde_json::Value =
serde_json::from_str(&decrypted_auth_config).expect("auth config should parse");
assert_eq!(auth_config["provider_type"], "xai");
assert_eq!(auth_config["auth_method"], "oauth");
assert_eq!(auth_config["using_api"], false);
assert_eq!(auth_config["email"], "[email protected]");
token_handle.abort();
}
#[test] #[test]
fn gateway_handles_admin_provider_oauth_device_poll_for_windsurf_one_time_token() { fn gateway_handles_admin_provider_oauth_device_poll_for_windsurf_one_time_token() {
run_admin_oauth_test( run_admin_oauth_test(
+36 -2
View File
@@ -36,6 +36,7 @@ mod openai_sync_task;
mod registry_poller; mod registry_poller;
mod routing; mod routing;
mod stream; mod stream;
mod xai;
/// Seed online manual proxy nodes for video execution fixtures. /// Seed online manual proxy nodes for video execution fixtures.
/// ///
@@ -44,6 +45,17 @@ mod stream;
/// the same deployment-state record; the loopback URL is never contacted when /// the same deployment-state record; the loopback URL is never contacted when
/// the execution-runtime override is active. /// the execution-runtime override is active.
pub(super) fn video_proxy_node_repository<I, S>(node_ids: I) -> Arc<InMemoryProxyNodeRepository> pub(super) fn video_proxy_node_repository<I, S>(node_ids: I) -> Arc<InMemoryProxyNodeRepository>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
video_proxy_node_repository_at_url(node_ids, "http://127.0.0.1:1")
}
pub(super) fn video_proxy_node_repository_at_url<I, S>(
node_ids: I,
proxy_url: &str,
) -> Arc<InMemoryProxyNodeRepository>
where where
I: IntoIterator<Item = S>, I: IntoIterator<Item = S>,
S: AsRef<str>, S: AsRef<str>,
@@ -68,7 +80,7 @@ where
1, 1,
) )
.expect("video test proxy node should build") .expect("video test proxy node should build")
.with_manual_proxy_fields(Some("http://127.0.0.1:1".to_string()), None, None) .with_manual_proxy_fields(Some(proxy_url.to_string()), None, None)
.with_tunnel_generation(format!("video-test-generation-{node_id}")) .with_tunnel_generation(format!("video-test-generation-{node_id}"))
}); });
Arc::new(InMemoryProxyNodeRepository::seed(nodes)) Arc::new(InMemoryProxyNodeRepository::seed(nodes))
@@ -86,6 +98,28 @@ pub(super) fn video_provider_catalog_repository(
endpoint_base_url: &str, endpoint_base_url: &str,
key_id: &str, key_id: &str,
upstream_api_key: &str, upstream_api_key: &str,
) -> Arc<InMemoryProviderCatalogReadRepository> {
video_provider_catalog_repository_with_proxy(
provider_id,
provider_type,
endpoint_id,
api_format,
endpoint_base_url,
key_id,
upstream_api_key,
None,
)
}
pub(super) fn video_provider_catalog_repository_with_proxy(
provider_id: &str,
provider_type: &str,
endpoint_id: &str,
api_format: &str,
endpoint_base_url: &str,
key_id: &str,
upstream_api_key: &str,
proxy: Option<serde_json::Value>,
) -> Arc<InMemoryProviderCatalogReadRepository> { ) -> Arc<InMemoryProviderCatalogReadRepository> {
fn seal_bound_credential( fn seal_bound_credential(
provider_id: &str, provider_id: &str,
@@ -117,7 +151,7 @@ pub(super) fn video_provider_catalog_repository(
false, false,
None, None,
Some(2), Some(2),
None, proxy,
Some(20.0), Some(20.0),
None, None,
None, None,
@@ -13,7 +13,8 @@ use serde_json::json;
use super::{ use super::{
build_state_with_execution_runtime_override, start_server, video_provider_catalog_repository, build_state_with_execution_runtime_override, start_server, video_provider_catalog_repository,
AppState, VideoTaskTruthSourceMode, video_provider_catalog_repository_with_proxy, video_proxy_node_repository_at_url, AppState,
VideoTaskTruthSourceMode,
}; };
fn sample_due_openai_task(upstream_base_url: &str) -> UpsertVideoTask { fn sample_due_openai_task(upstream_base_url: &str) -> UpsertVideoTask {
@@ -279,13 +280,13 @@ async fn gateway_background_video_task_poller_refreshes_due_openai_task_from_rep
); );
let (upstream_url, upstream_handle) = start_server(upstream).await; let (upstream_url, upstream_handle) = start_server(upstream).await;
let upstream_api_root = format!("{upstream_url}/v1"); let upstream_api_root = "http://video-provider.invalid/v1".to_string();
let repository = Arc::new(InMemoryVideoTaskRepository::default()); let repository = Arc::new(InMemoryVideoTaskRepository::default());
repository repository
.upsert(sample_due_openai_task(&upstream_api_root)) .upsert(sample_due_openai_task(&upstream_api_root))
.await .await
.expect("task upsert should succeed"); .expect("task upsert should succeed");
let provider_catalog_repository = video_provider_catalog_repository( let provider_catalog_repository = video_provider_catalog_repository_with_proxy(
"provider-openai-video-local-1", "provider-openai-video-local-1",
"openai", "openai",
"endpoint-openai-video-local-1", "endpoint-openai-video-local-1",
@@ -293,6 +294,7 @@ async fn gateway_background_video_task_poller_refreshes_due_openai_task_from_rep
&upstream_api_root, &upstream_api_root,
"key-openai-video-local-1", "key-openai-video-local-1",
"sk-upstream-openai-video", "sk-upstream-openai-video",
Some(json!({"enabled":true,"node_id":"poller-video-proxy"})),
); );
let gateway_state = AppState::new() let gateway_state = AppState::new()
@@ -302,7 +304,7 @@ async fn gateway_background_video_task_poller_refreshes_due_openai_task_from_rep
Arc::clone(&repository), Arc::clone(&repository),
provider_catalog_repository, provider_catalog_repository,
DEVELOPMENT_ENCRYPTION_KEY, DEVELOPMENT_ENCRYPTION_KEY,
), ).attach_proxy_node_repository_for_tests(video_proxy_node_repository_at_url(["poller-video-proxy"], &upstream_url)),
) )
.with_video_task_truth_source_mode(VideoTaskTruthSourceMode::RustAuthoritative) .with_video_task_truth_source_mode(VideoTaskTruthSourceMode::RustAuthoritative)
.with_video_task_poller_config(std::time::Duration::from_millis(25), 8); .with_video_task_poller_config(std::time::Duration::from_millis(25), 8);
+9 -23
View File
@@ -14,8 +14,7 @@ use crate::constants::{
use super::{build_router, start_server}; use super::{build_router, start_server};
#[tokio::test] #[tokio::test]
async fn gateway_locally_denies_video_control_sync_even_with_opt_in_headers_when_execution_runtime_missing( async fn gateway_hides_video_task_from_unauthenticated_caller_with_opt_in_headers() {
) {
let execute_hits = Arc::new(Mutex::new(0usize)); let execute_hits = Arc::new(Mutex::new(0usize));
let execute_hits_clone = Arc::clone(&execute_hits); let execute_hits_clone = Arc::clone(&execute_hits);
let public_hits = Arc::new(Mutex::new(0usize)); let public_hits = Arc::new(Mutex::new(0usize));
@@ -66,13 +65,9 @@ async fn gateway_locally_denies_video_control_sync_even_with_opt_in_headers_when
.await .await
.expect("request should succeed"); .expect("request should succeed");
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); assert_eq!(response.status(), StatusCode::NOT_FOUND);
let payload: serde_json::Value = response.json().await.expect("body should parse"); let payload: serde_json::Value = response.json().await.expect("body should parse");
assert_eq!(payload["error"]["type"], "http_error"); assert_eq!(payload, crate::video_tasks::not_found_body());
assert_eq!(
payload["error"]["message"],
"当前 OpenAI Video 请求无法在本地执行:没有匹配到可用的执行路径"
);
assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0); assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0);
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
@@ -81,8 +76,7 @@ async fn gateway_locally_denies_video_control_sync_even_with_opt_in_headers_when
} }
#[tokio::test] #[tokio::test]
async fn gateway_locally_denies_video_control_sync_without_opt_in_header_when_execution_runtime_missing( async fn gateway_hides_video_task_without_calling_public_or_control_upstream() {
) {
let execute_hits = Arc::new(Mutex::new(0usize)); let execute_hits = Arc::new(Mutex::new(0usize));
let execute_hits_clone = Arc::clone(&execute_hits); let execute_hits_clone = Arc::clone(&execute_hits);
let public_hits = Arc::new(Mutex::new(0usize)); let public_hits = Arc::new(Mutex::new(0usize));
@@ -142,13 +136,9 @@ async fn gateway_locally_denies_video_control_sync_without_opt_in_header_when_ex
.await .await
.expect("request should succeed"); .expect("request should succeed");
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); assert_eq!(response.status(), StatusCode::NOT_FOUND);
let payload: serde_json::Value = response.json().await.expect("body should parse"); let payload: serde_json::Value = response.json().await.expect("body should parse");
assert_eq!(payload["error"]["type"], "http_error"); assert_eq!(payload, crate::video_tasks::not_found_body());
assert_eq!(
payload["error"]["message"],
"当前 OpenAI Video 请求无法在本地执行:没有匹配到可用的执行路径"
);
assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0); assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0);
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
assert_eq!( assert_eq!(
@@ -165,7 +155,7 @@ async fn gateway_locally_denies_video_control_sync_without_opt_in_header_when_ex
} }
#[tokio::test] #[tokio::test]
async fn gateway_skips_video_get_control_sync_without_opt_in_header() { async fn gateway_hides_video_task_from_unauthenticated_caller_without_opt_in_headers() {
let execute_hits = Arc::new(Mutex::new(0usize)); let execute_hits = Arc::new(Mutex::new(0usize));
let execute_hits_clone = Arc::clone(&execute_hits); let execute_hits_clone = Arc::clone(&execute_hits);
let public_hits = Arc::new(Mutex::new(0usize)); let public_hits = Arc::new(Mutex::new(0usize));
@@ -211,13 +201,9 @@ async fn gateway_skips_video_get_control_sync_without_opt_in_header() {
.await .await
.expect("request should succeed"); .expect("request should succeed");
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); assert_eq!(response.status(), StatusCode::NOT_FOUND);
let payload: serde_json::Value = response.json().await.expect("body should parse"); let payload: serde_json::Value = response.json().await.expect("body should parse");
assert_eq!(payload["error"]["type"], "http_error"); assert_eq!(payload, crate::video_tasks::not_found_body());
assert_eq!(
payload["error"]["message"],
"当前 OpenAI Video 请求无法在本地执行:没有匹配到可用的执行路径"
);
assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0); assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0);
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
+427
View File
@@ -0,0 +1,427 @@
use super::*;
use aether_data::repository::auth::{
InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot,
};
use aether_data::repository::candidate_selection::InMemoryMinimalCandidateSelectionReadRepository;
use aether_data::repository::candidates::InMemoryRequestCandidateRepository;
use aether_data_contracts::repository::candidate_selection::{
StoredMinimalCandidateSelectionRow, StoredProviderModelMapping,
};
use sha2::{Digest, Sha256};
use std::sync::atomic::{AtomicUsize, Ordering};
fn sample_auth_snapshot(api_key_id: &str, user_id: &str) -> StoredAuthApiKeySnapshot {
StoredAuthApiKeySnapshot::new(
user_id.to_string(),
"video-user".to_string(),
Some("[email protected]".to_string()),
"user".to_string(),
"local".to_string(),
true,
false,
Some(json!(["openai"])),
Some(json!(["openai:video"])),
Some(json!(["video-model"])),
api_key_id.to_string(),
Some("default".to_string()),
true,
false,
false,
Some(60),
Some(5),
Some(4_102_444_800),
Some(json!(["openai"])),
Some(json!(["openai:video"])),
Some(json!(["video-model"])),
)
.expect("auth snapshot should build")
}
fn sample_candidate_row() -> StoredMinimalCandidateSelectionRow {
StoredMinimalCandidateSelectionRow {
provider_id: "provider-openai-video-local-1".to_string(),
provider_name: "openai".to_string(),
provider_type: "xai".to_string(),
provider_priority: 10,
provider_is_active: true,
endpoint_id: "endpoint-openai-video-local-1".to_string(),
endpoint_api_format: "openai:video".to_string(),
endpoint_api_family: Some("openai".to_string()),
endpoint_kind: Some("video".to_string()),
endpoint_is_active: true,
key_id: "key-openai-video-local-1".to_string(),
key_name: "prod".to_string(),
key_auth_type: "api_key".to_string(),
key_is_active: true,
key_api_formats: Some(vec!["openai:video".to_string()]),
key_allowed_models: None,
key_capabilities: None,
key_internal_priority: 5,
key_global_priority_by_format: Some(json!({"openai:video": 1})),
model_id: "model-openai-video-local-1".to_string(),
global_model_id: "global-model-openai-video-local-1".to_string(),
global_model_name: "video-model".to_string(),
global_model_mappings: None,
global_model_supports_streaming: Some(false),
model_provider_model_name: "grok-imagine-video".to_string(),
model_provider_model_mappings: Some(vec![StoredProviderModelMapping {
name: "grok-imagine-video".to_string(),
priority: 1,
api_formats: Some(vec!["openai:video".to_string()]),
endpoint_ids: None,
operations: None,
}]),
model_supports_streaming: Some(false),
model_is_active: true,
model_is_available: true,
}
}
#[tokio::test]
async fn xai_video_native_and_compatibility_http_lifecycle() {
Box::pin(assert_xai_video_http_lifecycle(Arc::new(
InMemoryVideoTaskRepository::default(),
)))
.await;
}
#[tokio::test]
async fn xai_video_native_and_compatibility_http_lifecycle_postgres() {
let configured_database_url = std::env::var("AETHER_TEST_DATABASE_URL").ok();
let managed_database = if configured_database_url.is_none() {
Some(
aether_testkit::ManagedPostgresServer::start()
.await
.expect("temporary PostgreSQL should start"),
)
} else {
None
};
let database_url = configured_database_url.unwrap_or_else(|| {
managed_database
.as_ref()
.expect("managed test database should exist")
.database_url()
.to_string()
});
let pool = sqlx::postgres::PgPoolOptions::new()
.max_connections(1)
.connect(&database_url)
.await
.expect("test database should connect");
aether_data::driver::postgres::run_migrations(&pool)
.await
.expect("test database should migrate");
// Preserve the production column constraints and unique indexes while isolating test rows.
sqlx::query("CREATE TEMP TABLE video_tasks (LIKE public.video_tasks INCLUDING ALL)")
.execute(&pool)
.await
.expect("isolated video task table should be created");
let repository =
Arc::new(aether_data::repository::video_tasks::SqlxVideoTaskRepository::new(pool.clone()));
Box::pin(assert_xai_video_http_lifecycle(repository)).await;
pool.close().await;
}
async fn assert_xai_video_http_lifecycle<T>(repository: Arc<T>)
where
T: aether_data_contracts::repository::video_tasks::VideoTaskRepository + 'static,
{
let static_dir = std::env::temp_dir().join(format!(
"aether-xai-video-static-{}",
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
std::fs::create_dir_all(&static_dir).unwrap();
std::fs::write(
static_dir.join("index.html"),
"<html>Aether test frontend</html>",
)
.unwrap();
let seen = Arc::new(Mutex::new(Vec::<serde_json::Value>::new()));
let calls = Arc::new(AtomicUsize::new(0));
// Exercise the real HTTP executor, including production method gates, instead of
// the test execution-runtime override that used to hide rejected GET requests.
let video_url = Arc::new(Mutex::new(String::new()));
let runtime = Router::new()
.route("/v1/videos/{operation}", any({
let seen = seen.clone();
let calls = calls.clone();
let video_url = video_url.clone();
move |request: Request| {
let seen = seen.clone();
let calls = calls.clone();
let video_url = video_url.clone();
async move {
let (parts, body) = request.into_parts();
assert_eq!(parts.headers["authorization"], "Bearer upstream-video-key");
let bytes = to_bytes(body, usize::MAX).await.unwrap();
let body: serde_json::Value = serde_json::from_slice(&bytes).unwrap_or(json!(null));
seen.lock().unwrap().push(json!({
"method": parts.method.as_str(),
"url": parts.uri.path(),
"body": {"json_body": body}
}));
let response = if parts.method == http::Method::POST {
json!({"request_id":"upstream-video-id", "provider_extension":{"accepted":true}})
} else {
assert_eq!(parts.uri.path(), "/v1/videos/upstream-video-id");
if calls.fetch_add(1, Ordering::SeqCst) == 0 {
json!({"status":"pending"})
} else {
json!({"status":"done", "model":"grok-imagine-video", "video":{"url":video_url.lock().unwrap().clone(), "duration":6, "respect_moderation":true}, "provider_extension":"preserved"})
}
};
Json(response)
}
}
}))
.route("/test.mp4", any(|request: Request| async move {
assert!(request.headers().get("authorization").is_none());
assert!(request.headers().get("x-xai-token-auth").is_none());
([("content-type", "video/mp4")], "test-video-bytes")
}));
let (runtime_url, runtime_handle) = start_server(runtime).await;
let expected_video_url = format!("{runtime_url}/test.mp4");
*video_url.lock().unwrap() = expected_video_url.clone();
let state_factory = || {
let auth = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![
(
Some(format!("{:x}", Sha256::digest(b"owner-key"))),
sample_auth_snapshot("owner-api-key", "owner"),
),
(
Some(format!("{:x}", Sha256::digest(b"foreign-key"))),
sample_auth_snapshot("foreign-api-key", "foreign"),
),
]));
let candidates = Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
sample_candidate_row(),
]));
let catalog = video_provider_catalog_repository_with_proxy(
"provider-openai-video-local-1",
"xai",
"endpoint-openai-video-local-1",
"openai:video",
"http://video-provider.invalid/v1",
"key-openai-video-local-1",
"upstream-video-key",
Some(json!({"enabled":true,"node_id":"video-proxy"})),
);
AppState::new().expect("gateway should build").with_video_task_truth_source_mode(VideoTaskTruthSourceMode::RustAuthoritative).with_data_state_for_tests(
crate::data::GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
auth, candidates, catalog, Arc::new(InMemoryRequestCandidateRepository::default()), DEVELOPMENT_ENCRYPTION_KEY
).attach_video_task_repository_for_tests(repository.clone())
.attach_proxy_node_repository_for_tests(video_proxy_node_repository_at_url(["video-proxy"], &runtime_url))
)
};
let router_factory =
|| crate::attach_static_frontend(build_router_with_state(state_factory()), &static_dir);
let (gateway_url, gateway_handle) = start_server(router_factory()).await;
let client = reqwest::Client::new();
assert_eq!(
client
.get(&gateway_url)
.send()
.await
.unwrap()
.text()
.await
.unwrap(),
"<html>Aether test frontend</html>"
);
for (path, native) in [
("/v1/videos/generations", true),
("/v1/videos", true),
("/v1/videos/edits", true),
("/v1/videos/extensions", true),
("/openai/v1/videos", false),
] {
calls.store(0, Ordering::SeqCst);
let body = if native {
json!({"model":"video-model","prompt":"A cat","duration":6,"aspect_ratio":"1:1","video":{"url":"https://example.com/input.mp4"},"future_option":true})
} else {
json!({"model":"video-model","prompt":"A cat","seconds":"6","size":"1280x720"})
};
let response = client
.post(format!("{gateway_url}{path}"))
.bearer_auth("owner-key")
.json(&body)
.send()
.await
.unwrap();
let status = response.status();
let result: serde_json::Value = response.json().await.unwrap();
assert_eq!(status, StatusCode::OK, "{path}: {result}");
let id = result[if native { "request_id" } else { "id" }]
.as_str()
.unwrap();
assert_ne!(id, "upstream-video-id");
if native {
assert!(result.get("id").is_none());
assert_eq!(result["provider_extension"]["accepted"], true);
} else {
assert_eq!(result["status"], "queued");
}
let request = seen.lock().unwrap().last().unwrap().clone();
let suffix = if path.ends_with("/edits") {
"edits"
} else if path.ends_with("/extensions") {
"extensions"
} else {
"generations"
};
assert_eq!(request["url"], format!("/v1/videos/{suffix}"));
assert_eq!(request["body"]["json_body"]["model"], "grok-imagine-video");
assert_eq!(request["body"]["json_body"]["duration"], 6);
if native {
assert_eq!(request["body"]["json_body"]["future_option"], true);
} else {
assert_eq!(request["body"]["json_body"]["aspect_ratio"], "16:9");
assert_eq!(request["body"]["json_body"]["resolution"], "720p");
assert!(request["body"]["json_body"].get("seconds").is_none());
assert!(request["body"]["json_body"].get("size").is_none());
}
let query = format!(
"{gateway_url}{}/{id}",
if native {
"/v1/videos"
} else {
"/openai/v1/videos"
}
);
let before = seen.lock().unwrap().len();
let denied = client
.get(&query)
.bearer_auth("foreign-key")
.send()
.await
.unwrap();
assert_eq!(denied.status(), StatusCode::NOT_FOUND);
assert_eq!(seen.lock().unwrap().len(), before);
let denied_content = client
.get(format!("{gateway_url}/openai/v1/videos/{id}/content"))
.bearer_auth("foreign-key")
.send()
.await
.unwrap();
assert_eq!(denied_content.status(), StatusCode::NOT_FOUND);
assert_eq!(seen.lock().unwrap().len(), before);
let pending: serde_json::Value = client
.get(&query)
.bearer_auth("owner-key")
.send()
.await
.unwrap()
.json()
.await
.unwrap();
assert_eq!(
pending["status"],
if native { "pending" } else { "queued" },
"{path}: {pending}"
);
let done: serde_json::Value = client
.get(&query)
.bearer_auth("owner-key")
.send()
.await
.unwrap()
.json()
.await
.unwrap();
assert_eq!(done["status"], if native { "done" } else { "completed" });
if native {
assert_eq!(done["video"]["respect_moderation"], true);
assert_eq!(done["provider_extension"], "preserved");
} else {
assert_eq!(done["video_url"], expected_video_url);
}
let stored = repository
.find(VideoTaskLookupKey::Id(id))
.await
.unwrap()
.unwrap();
assert_eq!(
stored.client_api_format.as_deref(),
Some(if native { "xai:video" } else { "openai:video" })
);
assert_eq!(
stored.external_task_id.as_deref(),
Some("upstream-video-id")
);
assert!(stored.request_metadata.is_none());
assert!(stored.original_request_body.is_none());
// A new gateway instance must reconstruct the pinned provider/credential and protocol.
let (restart_url, restart_handle) = start_server(router_factory()).await;
let restored: serde_json::Value = client
.get(format!(
"{restart_url}{}/{id}",
if native {
"/v1/videos"
} else {
"/openai/v1/videos"
}
))
.bearer_auth("owner-key")
.send()
.await
.unwrap()
.json()
.await
.unwrap();
assert_eq!(restored["status"], done["status"]);
if native {
assert_eq!(restored["video"]["respect_moderation"], true);
}
let compat: serde_json::Value = client
.get(format!("{restart_url}/openai/v1/videos/{id}"))
.bearer_auth("owner-key")
.send()
.await
.unwrap()
.json()
.await
.unwrap();
assert_eq!(compat["status"], "completed");
assert_eq!(compat["video_url"], expected_video_url);
let native_view: serde_json::Value = client
.get(format!("{restart_url}/v1/videos/{id}"))
.bearer_auth("owner-key")
.send()
.await
.unwrap()
.json()
.await
.unwrap();
assert_eq!(native_view["status"], "done");
assert_eq!(native_view["video"]["respect_moderation"], true);
for prefix in ["/v1/videos", "/openai/v1/videos"] {
let content = client
.get(format!("{restart_url}{prefix}/{id}/content"))
.bearer_auth("owner-key")
.send()
.await
.unwrap();
assert_eq!(content.status(), StatusCode::OK);
assert_eq!(content.headers()["content-type"], "video/mp4");
assert_eq!(content.bytes().await.unwrap(), "test-video-bytes");
}
restart_handle.abort();
}
let before = seen.lock().unwrap().len();
let bad = client
.post(format!("{gateway_url}/openai/v1/videos"))
.bearer_auth("owner-key")
.json(&json!({"model":"video-model","prompt":"cat","seconds":"wrong"}))
.send()
.await
.unwrap();
assert_eq!(bad.status(), StatusCode::BAD_REQUEST);
assert_eq!(seen.lock().unwrap().len(), before);
gateway_handle.abort();
runtime_handle.abort();
std::fs::remove_dir_all(&static_dir).unwrap();
}
@@ -11,6 +11,9 @@ use super::{
fn rust_authoritative_service_builds_openai_cancel_follow_up_plan() { fn rust_authoritative_service_builds_openai_cancel_follow_up_plan() {
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: false,
local_task_id: "task-local-123".to_string(), local_task_id: "task-local-123".to_string(),
upstream_task_id: "ext-video-task-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(),
created_at_unix_ms: 1712345678, created_at_unix_ms: 1712345678,
@@ -92,6 +95,9 @@ fn rust_authoritative_service_builds_openai_cancel_follow_up_plan() {
fn rust_authoritative_service_builds_openai_remix_follow_up_plan() { fn rust_authoritative_service_builds_openai_remix_follow_up_plan() {
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: false,
local_task_id: "task-local-123".to_string(), local_task_id: "task-local-123".to_string(),
upstream_task_id: "ext-video-task-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(),
created_at_unix_ms: 1712345678, created_at_unix_ms: 1712345678,
@@ -177,6 +183,9 @@ fn rust_authoritative_service_builds_openai_remix_follow_up_plan() {
fn rust_authoritative_service_builds_openai_delete_follow_up_plan() { fn rust_authoritative_service_builds_openai_delete_follow_up_plan() {
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: false,
local_task_id: "task-local-123".to_string(), local_task_id: "task-local-123".to_string(),
upstream_task_id: "ext-video-task-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(),
created_at_unix_ms: 1712345678, created_at_unix_ms: 1712345678,
@@ -332,6 +341,9 @@ fn rust_authoritative_service_builds_gemini_cancel_follow_up_plan() {
fn rust_authoritative_service_builds_openai_read_refresh_plan() { fn rust_authoritative_service_builds_openai_read_refresh_plan() {
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: false,
local_task_id: "task-local-123".to_string(), local_task_id: "task-local-123".to_string(),
upstream_task_id: "ext-video-task-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(),
created_at_unix_ms: 1712345678, created_at_unix_ms: 1712345678,
@@ -407,6 +419,9 @@ fn rust_authoritative_service_builds_gemini_read_refresh_plan() {
fn rust_authoritative_service_builds_poll_refresh_batch_for_active_tasks_only() { fn rust_authoritative_service_builds_poll_refresh_batch_for_active_tasks_only() {
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: false,
local_task_id: "task-active-123".to_string(), local_task_id: "task-active-123".to_string(),
upstream_task_id: "ext-video-task-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(),
created_at_unix_ms: 1712345678, created_at_unix_ms: 1712345678,
@@ -428,6 +443,9 @@ fn rust_authoritative_service_builds_poll_refresh_batch_for_active_tasks_only()
transport: sample_transport("https://api.openai.example", "openai:video"), transport: sample_transport("https://api.openai.example", "openai:video"),
})); }));
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: false,
local_task_id: "task-completed-123".to_string(), local_task_id: "task-completed-123".to_string(),
upstream_task_id: "ext-video-task-999".to_string(), upstream_task_id: "ext-video-task-999".to_string(),
created_at_unix_ms: 1712345678, created_at_unix_ms: 1712345678,
@@ -471,6 +489,9 @@ fn file_video_task_store_persists_snapshots_across_service_rebuilds() {
) )
.expect("file-backed service should build"); .expect("file-backed service should build");
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: false,
local_task_id: "task-file-123".to_string(), local_task_id: "task-file-123".to_string(),
upstream_task_id: "ext-video-task-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(),
created_at_unix_ms: 1712345678, created_at_unix_ms: 1712345678,
@@ -10,6 +10,9 @@ use super::{
fn rust_authoritative_service_projects_openai_status_into_local_read_response() { fn rust_authoritative_service_projects_openai_status_into_local_read_response() {
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: false,
local_task_id: "task-local-123".to_string(), local_task_id: "task-local-123".to_string(),
upstream_task_id: "ext-video-task-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(),
created_at_unix_ms: 1712345678, created_at_unix_ms: 1712345678,
@@ -93,6 +96,9 @@ fn rust_authoritative_service_projects_openai_status_into_local_read_response()
fn rust_authoritative_service_builds_openai_content_stream_plan_from_direct_video_url() { fn rust_authoritative_service_builds_openai_content_stream_plan_from_direct_video_url() {
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: false,
local_task_id: "task-local-123".to_string(), local_task_id: "task-local-123".to_string(),
upstream_task_id: "ext-video-task-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(),
created_at_unix_ms: 1712345678, created_at_unix_ms: 1712345678,
@@ -159,6 +165,9 @@ fn rust_authoritative_service_builds_openai_content_stream_plan_from_direct_vide
fn rust_authoritative_service_returns_processing_content_response_for_pending_openai_task() { fn rust_authoritative_service_returns_processing_content_response_for_pending_openai_task() {
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: false,
local_task_id: "task-local-123".to_string(), local_task_id: "task-local-123".to_string(),
upstream_task_id: "ext-video-task-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(),
created_at_unix_ms: 1712345678, created_at_unix_ms: 1712345678,
@@ -218,6 +218,9 @@ fn rust_authoritative_video_truth_source_can_background_success_report() {
fn rust_authoritative_service_reads_openai_task_from_local_registry() { fn rust_authoritative_service_reads_openai_task_from_local_registry() {
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
let snapshot = LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { let snapshot = LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: false,
local_task_id: "task-local-123".to_string(), local_task_id: "task-local-123".to_string(),
upstream_task_id: "ext-video-task-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(),
created_at_unix_ms: 1712345678, created_at_unix_ms: 1712345678,
@@ -266,6 +269,9 @@ fn rust_authoritative_service_reads_openai_task_from_local_registry() {
fn rust_authoritative_service_applies_cancel_and_delete_mutations() { fn rust_authoritative_service_applies_cancel_and_delete_mutations() {
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: false,
local_task_id: "task-local-123".to_string(), local_task_id: "task-local-123".to_string(),
upstream_task_id: "ext-video-task-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(),
created_at_unix_ms: 1712345678, created_at_unix_ms: 1712345678,
+252
View File
@@ -3546,6 +3546,258 @@ pub fn parse_kiro_usage_response(
Some(serde_json::Value::Object(result)) Some(serde_json::Value::Object(result))
} }
pub fn parse_xai_billing_response(
value: &serde_json::Value,
updated_at_unix_secs: u64,
) -> Option<serde_json::Value> {
let root = value.as_object()?;
let config = root
.get("config")
.and_then(serde_json::Value::as_object)
.unwrap_or(root);
let usage_percentage = coerce_json_f64_from_map(config, "creditUsagePercent")
.or_else(|| extract_xai_product_usage_percent(config));
let period = config.get("currentPeriod");
let period_type = period
.and_then(|value| value.get("type").or_else(|| value.get("periodType")))
.and_then(normalize_xai_period_type);
let next_reset_at = period
.and_then(|value| value.get("end"))
.and_then(parse_xai_timestamp)
.or_else(|| config.get("billingPeriodEnd").and_then(parse_xai_timestamp));
let monthly_limit =
coerce_xai_cents_dollars(config.get("monthlyLimit")).filter(|value| *value > 0.0);
let current_usage = if monthly_limit.is_some() {
coerce_xai_cents_dollars(config.get("used"))
} else {
None
};
let remaining = monthly_limit
.zip(current_usage)
.map(|(limit, used)| (limit - used).max(0.0));
let usage_percentage = usage_percentage.or_else(|| {
monthly_limit
.zip(current_usage)
.map(|(limit, used)| ((used / limit) * 100.0).clamp(0.0, 100.0))
});
let usage_percentage = match usage_percentage {
Some(value) => Some(value.clamp(0.0, 100.0)),
None if period_type.is_some() || next_reset_at.is_some() => Some(0.0),
None => None,
};
let prepaid_balance = coerce_xai_cents_dollars(config.get("prepaidBalance"));
let on_demand_cap = coerce_xai_cents_dollars(config.get("onDemandCap"));
let on_demand_used = coerce_xai_cents_dollars(config.get("onDemandUsed"));
let on_demand_enabled = coerce_json_bool_from_map(root, "onDemandEnabled")
.or_else(|| coerce_json_bool_from_map(config, "onDemandEnabled"));
let subscription_title = first_json_string_by_paths(
value,
&[
&["subscriptionTier"],
&["subscription_tier"],
&["config", "subscriptionTier"],
&["config", "subscription_title"],
],
);
if usage_percentage.is_none()
&& monthly_limit.is_none()
&& current_usage.is_none()
&& prepaid_balance.is_none()
&& on_demand_cap.is_none()
&& next_reset_at.is_none()
&& subscription_title.is_none()
{
return None;
}
let mut result = serde_json::Map::new();
result.insert("updated_at".to_string(), json!(updated_at_unix_secs));
if let Some(value) = usage_percentage {
result.insert("usage_percentage".to_string(), json!(value));
}
if let Some(value) = monthly_limit {
result.insert("usage_limit".to_string(), json!(value));
}
if let Some(value) = current_usage {
result.insert("current_usage".to_string(), json!(value));
}
if let Some(value) = remaining {
result.insert("remaining".to_string(), json!(value));
}
if let Some(value) = next_reset_at {
result.insert("next_reset_at".to_string(), json!(value));
}
if let Some(value) = period_type {
result.insert("period_type".to_string(), json!(value));
}
if let Some(value) = prepaid_balance {
result.insert("prepaid_balance".to_string(), json!(value));
}
if let Some(value) = on_demand_cap {
result.insert("on_demand_cap".to_string(), json!(value));
}
if let Some(value) = on_demand_used {
result.insert("on_demand_used".to_string(), json!(value));
}
if let Some(value) = on_demand_enabled {
result.insert("on_demand_enabled".to_string(), json!(value));
}
if let Some(value) = subscription_title {
result.insert("subscription_title".to_string(), json!(value));
}
Some(serde_json::Value::Object(result))
}
fn coerce_json_f64_from_map(
object: &serde_json::Map<String, serde_json::Value>,
key: &str,
) -> Option<f64> {
object.get(key).and_then(coerce_json_f64)
}
fn coerce_json_bool_from_map(
object: &serde_json::Map<String, serde_json::Value>,
key: &str,
) -> Option<bool> {
object.get(key).and_then(coerce_json_bool)
}
fn extract_xai_product_usage_percent(
config: &serde_json::Map<String, serde_json::Value>,
) -> Option<f64> {
let items = config.get("productUsage")?.as_array()?;
let grok_build = items.iter().find(|item| {
item.get("product")
.and_then(serde_json::Value::as_str)
.is_some_and(|product| product.eq_ignore_ascii_case("GrokBuild"))
});
grok_build
.or(items.first())
.and_then(|item| item.get("usagePercent").and_then(coerce_json_f64))
}
fn coerce_xai_cents_dollars(value: Option<&serde_json::Value>) -> Option<f64> {
let value = value?;
let cents = match value {
serde_json::Value::Object(object) => object.get("val").and_then(coerce_json_f64)?,
other => coerce_json_f64(other)?,
};
Some(cents / 100.0)
}
fn normalize_xai_period_type(value: &serde_json::Value) -> Option<String> {
let raw = value
.as_str()
.map(str::trim)
.filter(|value| !value.is_empty())?;
let lowered = raw.to_ascii_lowercase();
if lowered.contains("week") {
Some("weekly".to_string())
} else if lowered.contains("month") {
Some("monthly".to_string())
} else {
Some(raw.to_string())
}
}
fn parse_xai_timestamp(value: &serde_json::Value) -> Option<u64> {
if let Some(value) = coerce_json_u64(value) {
return Some(if value > 1_000_000_000_000 {
value / 1000
} else {
value
});
}
let raw = value.as_str()?.trim();
if raw.is_empty() {
return None;
}
chrono::DateTime::parse_from_rfc3339(raw)
.ok()
.and_then(|timestamp| u64::try_from(timestamp.timestamp()).ok())
}
#[cfg(test)]
mod xai_quota_tests {
use super::parse_xai_billing_response;
use serde_json::json;
#[test]
fn parse_xai_credits_percent_and_weekly_period() {
let metadata = parse_xai_billing_response(
&json!({
"config": {
"currentPeriod": {
"type": "USAGE_PERIOD_TYPE_WEEKLY",
"start": "2026-08-08T01:53:09.930537+00:00",
"end": "2026-08-15T01:53:09.930537+00:00"
},
"creditUsagePercent": 46.0,
"productUsage": [
{"product": "GrokBuild", "usagePercent": 41.0},
{"product": "GrokChat"}
],
"onDemandCap": {"val": 0},
"onDemandUsed": {"val": 0},
"prepaidBalance": {"val": 0}
},
"subscriptionTier": "SuperGrok"
}),
1_775_000_000,
)
.expect("credits payload should parse");
assert_eq!(metadata["usage_percentage"], json!(46.0));
assert_eq!(metadata["period_type"], json!("weekly"));
assert_eq!(metadata["next_reset_at"], json!(1_786_758_789u64));
assert_eq!(metadata["prepaid_balance"], json!(0.0));
assert_eq!(metadata["on_demand_cap"], json!(0.0));
assert_eq!(metadata["subscription_title"], json!("SuperGrok"));
}
#[test]
fn parse_xai_omitted_percent_as_fresh_weekly_zero() {
let metadata = parse_xai_billing_response(
&json!({
"config": {
"currentPeriod": {
"type": "USAGE_PERIOD_TYPE_WEEKLY",
"end": "2026-08-15T01:53:09.930537+00:00"
},
"isUnifiedBillingUser": true
}
}),
1_775_000_000,
)
.expect("fresh weekly period should parse");
assert_eq!(metadata["usage_percentage"], json!(0.0));
assert_eq!(metadata["period_type"], json!("weekly"));
}
#[test]
fn parse_xai_legacy_monthly_cents() {
let metadata = parse_xai_billing_response(
&json!({
"config": {
"monthlyLimit": {"val": 2500},
"used": {"val": 1000},
"billingPeriodEnd": "2026-09-01T00:00:00Z"
}
}),
1_775_000_000,
)
.expect("legacy monthly payload should parse");
assert_eq!(metadata["usage_limit"], json!(25.0));
assert_eq!(metadata["current_usage"], json!(10.0));
assert_eq!(metadata["remaining"], json!(15.0));
assert_eq!(metadata["usage_percentage"], json!(40.0));
}
}
pub fn parse_windsurf_user_status_response( pub fn parse_windsurf_user_status_response(
value: &serde_json::Value, value: &serde_json::Value,
updated_at_unix_secs: u64, updated_at_unix_secs: u64,
+39
View File
@@ -285,6 +285,23 @@ pub fn enrich_admin_provider_oauth_auth_config(
], ],
); );
if provider_type.trim().eq_ignore_ascii_case("xai") {
auth_config.insert("auth_method".to_string(), json!("oauth"));
auth_config.insert("using_api".to_string(), json!(false));
if let Some(id_token) = ["id_token", "idToken"]
.iter()
.find_map(|field| json_non_empty_string(token_payload.get(field)))
{
auth_config
.entry("id_token".to_string())
.or_insert_with(|| json!(id_token.clone()));
if let Some(claims) = decode_jwt_claims(&id_token) {
merge_missing_auth_config_fields(auth_config, &claims, &["email", "sub"]);
}
}
return;
}
if provider_type.trim().eq_ignore_ascii_case("claude_code") { if provider_type.trim().eq_ignore_ascii_case("claude_code") {
if let Some(organization_uuid) = token_payload_object if let Some(organization_uuid) = token_payload_object
.get("organization") .get("organization")
@@ -554,6 +571,28 @@ mod tests {
assert_eq!(auth_config.get("is_fedramp"), Some(&json!(true))); assert_eq!(auth_config.get("is_fedramp"), Some(&json!(true)));
} }
#[test]
fn xai_enrichment_marks_oauth_and_extracts_id_token_identity() {
let id_token = sample_unsigned_jwt(json!({
"email": "[email protected]",
"sub": "user-xai-1",
}));
let token_payload = json!({
"access_token": "access-token",
"refresh_token": "refresh-token",
"id_token": id_token,
});
let mut auth_config = serde_json::Map::new();
enrich_admin_provider_oauth_auth_config("xai", &mut auth_config, &token_payload);
assert_eq!(auth_config.get("auth_method"), Some(&json!("oauth")));
assert_eq!(auth_config.get("using_api"), Some(&json!(false)));
assert_eq!(auth_config.get("email"), Some(&json!("[email protected]")));
assert_eq!(auth_config.get("sub"), Some(&json!("user-xai-1")));
assert_eq!(auth_config.get("id_token"), Some(&json!(id_token)));
}
#[test] #[test]
fn decode_jwt_claims_rejects_oversized_payload_before_decode() { fn decode_jwt_claims_rejects_oversized_payload_before_decode() {
let max_encoded_len = MAX_UNVERIFIED_JWT_CLAIMS_BYTES let max_encoded_len = MAX_UNVERIFIED_JWT_CLAIMS_BYTES
+4
View File
@@ -208,6 +208,10 @@ pub use crate::formats::{
resolve_stream_spec as resolve_openai_responses_stream_spec, resolve_stream_spec as resolve_openai_responses_stream_spec,
resolve_sync_spec as resolve_openai_responses_sync_spec, LocalOpenAiResponsesSpec, resolve_sync_spec as resolve_openai_responses_sync_spec, LocalOpenAiResponsesSpec,
}, },
xai::{
apply_xai_upstream_payload_edits, apply_xai_upstream_payload_edits_with_client,
xai_supports_native_image_generation,
},
}, },
}, },
shared::{ shared::{
@@ -7,6 +7,7 @@ pub mod request;
pub mod response; pub mod response;
pub mod spec; pub mod spec;
pub mod stream; pub mod stream;
pub mod xai;
const TOOL_ERROR_PREFIX: &str = "[tool error]"; const TOOL_ERROR_PREFIX: &str = "[tool error]";
const AETHER_REASONING_ITEM_ID_PREFIX: &str = "rs_aether_"; const AETHER_REASONING_ITEM_ID_PREFIX: &str = "rs_aether_";
@@ -85,6 +86,8 @@ pub enum OpenAiResponsesReasoningReplayPolicy {
#[default] #[default]
OpenAiItemIds, OpenAiItemIds,
DeepSeekOpaque, DeepSeekOpaque,
/// xAI replays encrypted state without requiring OpenAI's item-ID prefix.
XaiEncrypted,
} }
/// Builds a stable, wire-compatible ID for a reasoning item synthesized by Aether. /// Builds a stable, wire-compatible ID for a reasoning item synthesized by Aether.
@@ -234,6 +237,14 @@ fn openai_responses_reasoning_item_is_replayable(
{ {
return true; return true;
} }
if policy == OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
&& object
.get("encrypted_content")
.and_then(Value::as_str)
.is_some_and(|value| !value.trim().is_empty())
{
return true;
}
let Some(id) = object let Some(id) = object
.get("id") .get("id")
.and_then(Value::as_str) .and_then(Value::as_str)
@@ -334,6 +345,36 @@ mod tests {
OPENAI_RESPONSES_OPERATION_COMPACT, OPENAI_RESPONSES_OPERATION_COMPACT,
}; };
#[test]
fn xai_encrypted_replay_accepts_native_ids_but_excludes_foreign_carriers() {
let body = serde_json::json!({"input": [
{"type": "reasoning", "id": "native-xai-id", "encrypted_content": "opaque-xai-state"},
{"type": "reasoning", "encrypted_content": "opaque-idless-state"},
{"type": "reasoning", "id": "rs_foreign", "encrypted_content": "cpa-gemini-responses-carrier-v1:foreign"},
{"type": "reasoning", "id": "foreign-id", "summary": []}
]});
let mut xai = body.clone();
assert_eq!(
super::strip_incompatible_openai_responses_reasoning_items_with_policy(
&mut xai,
"openai:responses",
super::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted,
),
2
);
assert_eq!(xai["input"].as_array().unwrap().len(), 2);
assert_eq!(xai["input"][0], body["input"][0]);
assert_eq!(xai["input"][1], body["input"][1]);
let mut openai = body;
assert_eq!(
super::strip_incompatible_openai_responses_reasoning_items(
&mut openai,
"openai:responses"
),
4
);
}
#[test] #[test]
fn gemini_tool_signature_carrier_roundtrips_direction_and_exact_value() { fn gemini_tool_signature_carrier_roundtrips_direction_and_exact_value() {
let signature = " opaque-signature-with-padding== "; let signature = " opaque-signature-with-padding== ";
@@ -0,0 +1,914 @@
use serde_json::{json, Map, Value};
const XAI_RESPONSES_UNSUPPORTED_BODY_FIELDS: &[&str] = &[
"previous_response_id",
"prompt_cache_retention",
"safety_identifier",
"stream_options",
"stop",
"metadata",
];
const XAI_WEB_SEARCH_TOOL_TYPE: &str = "web_search";
const XAI_IMAGE_GENERATION_TOOL_TYPE: &str = "image_generation";
const XAI_TOOL_SEARCH_TOOL_TYPE: &str = "tool_search";
const XAI_GROK_IMAGE_GENERATION_MIN: XaiGrokVersion = XaiGrokVersion { major: 4, minor: 6 };
#[derive(Clone, Copy)]
struct XaiGrokVersion {
major: i32,
minor: i32,
}
pub fn apply_xai_upstream_payload_edits(
body: &mut Value,
provider_type: &str,
provider_api_format: &str,
) {
apply_xai_upstream_payload_edits_with_client(
body,
provider_type,
provider_api_format,
None,
None,
);
}
pub fn apply_xai_upstream_payload_edits_with_client(
body: &mut Value,
provider_type: &str,
provider_api_format: &str,
client_api_format: Option<&str>,
client_body: Option<&Value>,
) {
if !provider_type.trim().eq_ignore_ascii_case("xai") {
return;
}
normalize_xai_image_refs(body);
if crate::is_openai_responses_family_format(provider_api_format) {
restore_xai_web_search_from_client(body, client_api_format, client_body);
sanitize_xai_responses_body(body);
}
}
fn sanitize_xai_responses_body(body: &mut Value) {
let Some(object) = body.as_object_mut() else {
return;
};
for field in XAI_RESPONSES_UNSUPPORTED_BODY_FIELDS {
object.remove(*field);
}
let keep_image_generation = object
.get("model")
.and_then(Value::as_str)
.is_some_and(xai_supports_native_image_generation);
normalize_xai_tool_arrays(object, keep_image_generation);
rewrite_xai_web_search_tool_choice(object);
prune_xai_orphaned_tool_choice(object);
rewrite_xai_image_generation_tool_choice(object);
drop_tool_choice_without_tools(object);
strip_unsupported_reasoning_effort(object);
sanitize_xai_input_encrypted_content(object);
}
fn restore_xai_web_search_from_client(
body: &mut Value,
client_api_format: Option<&str>,
client_body: Option<&Value>,
) {
let Some(client_api_format) = client_api_format else {
return;
};
let Some(client_body) = client_body else {
return;
};
if !client_requests_web_search(client_api_format, client_body) {
return;
}
ensure_xai_web_search_tool(body);
// Claude names a hosted tool in tool_choice just like a client function.
// Resolve that name against the original declaration, never by name alone.
if crate::normalize_api_format_alias(client_api_format) == "claude:messages" {
let choice = &client_body["tool_choice"];
if choice["type"] == "tool"
&& choice["name"].as_str().is_some_and(|name| {
request_tools(client_body)
.iter()
.any(|tool| is_web_search_tool(tool) && tool_name(tool) == Some(name))
})
{
body["tool_choice"] = json!({"type": XAI_WEB_SEARCH_TOOL_TYPE});
}
}
}
fn client_requests_web_search(client_api_format: &str, client_body: &Value) -> bool {
let format = crate::normalize_api_format_alias(client_api_format);
match format.as_str() {
"openai:chat" => {
object_has_non_null_field(client_body, "web_search_options")
|| request_tools(client_body).iter().any(is_web_search_tool)
}
"claude:messages" => request_tools(client_body).iter().any(is_web_search_tool),
"gemini:generate_content" => gemini_request_has_google_search(client_body),
_ => false,
}
}
fn gemini_request_has_google_search(body: &Value) -> bool {
request_tools(body).iter().any(|tool| {
tool.get("googleSearch").is_some()
|| tool.get("google_search").is_some()
|| tool
.get("googleSearchRetrieval")
.is_some_and(|value| !value.is_null())
})
}
fn object_has_non_null_field(body: &Value, field: &str) -> bool {
body.get(field).is_some_and(|value| !value.is_null())
}
fn ensure_xai_web_search_tool(body: &mut Value) {
let Some(object) = body.as_object_mut() else {
return;
};
if tools_array(object).iter().any(is_web_search_tool) {
return;
}
let tools = object
.entry("tools".to_string())
.or_insert_with(|| Value::Array(Vec::new()));
if let Some(tools) = tools.as_array_mut() {
tools.push(json!({ "type": XAI_WEB_SEARCH_TOOL_TYPE }));
}
}
fn normalize_xai_tool_arrays(object: &mut Map<String, Value>, keep_image_generation: bool) {
if let Some(tools) = object.get_mut("tools").and_then(Value::as_array_mut) {
*tools = normalize_xai_tool_list(tools, keep_image_generation);
if tools.is_empty() {
object.remove("tools");
}
}
let Some(input) = object.get_mut("input").and_then(Value::as_array_mut) else {
return;
};
for item in input {
let Some(item_object) = item.as_object_mut() else {
continue;
};
if item_object.get("type").and_then(Value::as_str) != Some("additional_tools") {
continue;
}
if let Some(tools) = item_object.get_mut("tools").and_then(Value::as_array_mut) {
*tools = normalize_xai_tool_list(tools, keep_image_generation);
}
}
}
fn normalize_xai_tool_list(tools: &[Value], keep_image_generation: bool) -> Vec<Value> {
tools
.iter()
.filter_map(|tool| normalize_xai_tool(tool, keep_image_generation))
.collect()
}
fn normalize_xai_tool(tool: &Value, keep_image_generation: bool) -> Option<Value> {
let Some(object) = tool.as_object() else {
return Some(tool.clone());
};
let tool_type = tool_type(tool).unwrap_or("function");
if tool_type == XAI_TOOL_SEARCH_TOOL_TYPE {
return None;
}
if tool_type == XAI_IMAGE_GENERATION_TOOL_TYPE && !keep_image_generation {
return None;
}
if tool_type == "custom" && tool_name(tool).is_some_and(|name| name == "apply_patch") {
return None;
}
let mut next = object.clone();
if tool_type.starts_with("web_search") {
next.insert(
"type".to_string(),
Value::String(XAI_WEB_SEARCH_TOOL_TYPE.to_string()),
);
next.remove("name");
next.remove("external_web_access");
return Some(Value::Object(next));
}
if tool_type == "custom" {
next.insert("type".to_string(), Value::String("function".to_string()));
if let Some(custom) = next.remove("custom") {
if let Some(custom_object) = custom.as_object() {
for (key, value) in custom_object {
next.entry(key.clone()).or_insert_with(|| value.clone());
}
}
}
if !next.contains_key("parameters") {
next.insert(
"parameters".to_string(),
json!({"type": "object", "properties": {}}),
);
}
return Some(Value::Object(next));
}
if tool_type == "function" && !next.contains_key("parameters") {
next.insert(
"parameters".to_string(),
json!({"type": "object", "properties": {}}),
);
}
Some(Value::Object(next))
}
fn rewrite_xai_web_search_tool_choice(object: &mut Map<String, Value>) {
let Some(choice) = object.get("tool_choice").cloned() else {
return;
};
let Some(choice_type) = choice.as_object().and_then(|value| {
value
.get("type")
.and_then(Value::as_str)
.map(str::trim)
.map(str::to_ascii_lowercase)
}) else {
return;
};
if is_web_search_choice_type(&choice_type) {
object.insert(
"tool_choice".to_string(),
json!({
"type": "allowed_tools",
"mode": "required",
"tools": [{ "type": XAI_WEB_SEARCH_TOOL_TYPE }]
}),
);
}
}
fn rewrite_xai_image_generation_tool_choice(object: &mut Map<String, Value>) {
let has_image_generation = tools_array(object)
.iter()
.any(|tool| tool_type(tool).is_some_and(|value| value == XAI_IMAGE_GENERATION_TOOL_TYPE));
if !has_image_generation {
return;
}
let Some(choice) = object.get("tool_choice").cloned() else {
return;
};
// xAI's allowed_tools schema cannot contain image_generation. Preserve an
// image-only restriction before filtering image entries out of mixed lists.
let image_only = is_allowed_tools_image_generation_only(&choice);
if choice["type"] == XAI_IMAGE_GENERATION_TOOL_TYPE || image_only {
let mode = if image_only && choice["mode"] == "auto" {
"auto"
} else {
"required"
};
keep_only_image_generation_tools(object);
object.insert("tool_choice".to_string(), Value::String(mode.to_string()));
} else if choice["type"] == "allowed_tools" {
filter_image_generation_from_allowed_tools(object);
}
}
fn is_allowed_tools_image_generation_only(choice: &Value) -> bool {
let Some(object) = choice.as_object() else {
return false;
};
if object.get("type").and_then(Value::as_str) != Some("allowed_tools") {
return false;
}
let Some(tools) = object.get("tools").and_then(Value::as_array) else {
return false;
};
!tools.is_empty()
&& tools.iter().all(|tool| {
tool_type(tool).is_some_and(|value| value == XAI_IMAGE_GENERATION_TOOL_TYPE)
})
}
fn keep_only_image_generation_tools(object: &mut Map<String, Value>) {
let Some(tools) = object.get_mut("tools").and_then(Value::as_array_mut) else {
return;
};
tools.retain(|tool| {
tool_type(tool).is_some_and(|value| value == XAI_IMAGE_GENERATION_TOOL_TYPE)
});
}
fn filter_image_generation_from_allowed_tools(object: &mut Map<String, Value>) {
let Some(choice) = object.get_mut("tool_choice").and_then(Value::as_object_mut) else {
return;
};
let Some(tools) = choice.get_mut("tools").and_then(Value::as_array_mut) else {
return;
};
tools
.retain(|tool| tool_type(tool).is_none_or(|value| value != XAI_IMAGE_GENERATION_TOOL_TYPE));
}
fn is_web_search_choice_type(value: &str) -> bool {
value == XAI_WEB_SEARCH_TOOL_TYPE || value.starts_with("web_search")
}
fn prune_xai_orphaned_tool_choice(object: &mut Map<String, Value>) {
let available = collect_available_tool_choice_keys(object);
let Some(choice) = object.get("tool_choice").cloned() else {
return;
};
if choice.as_str().is_some() {
return;
}
let Some(choice_object) = choice.as_object() else {
object.remove("tool_choice");
return;
};
let choice_type = choice_object
.get("type")
.and_then(Value::as_str)
.unwrap_or_default()
.trim()
.to_ascii_lowercase();
if choice_type == "allowed_tools" {
let Some(allowed) = choice_object.get("tools").and_then(Value::as_array) else {
object.remove("tool_choice");
return;
};
let kept = allowed
.iter()
.filter(|tool| tool_matches_available(tool, &available))
.cloned()
.collect::<Vec<_>>();
if kept.is_empty() {
object.remove("tool_choice");
return;
}
if let Some(choice) = object.get_mut("tool_choice").and_then(Value::as_object_mut) {
choice.insert("tools".to_string(), Value::Array(kept));
}
return;
}
if choice_type.is_empty() {
return;
}
if !tool_matches_available(&choice, &available) {
object.remove("tool_choice");
}
}
fn collect_available_tool_choice_keys(object: &Map<String, Value>) -> Vec<ToolChoiceKey> {
let mut keys = Vec::new();
collect_tool_choice_keys(tools_array(object), &mut keys);
if let Some(input) = object.get("input").and_then(Value::as_array) {
for item in input {
if item.get("type").and_then(Value::as_str) == Some("additional_tools") {
collect_tool_choice_keys(
item.get("tools")
.and_then(Value::as_array)
.map(Vec::as_slice)
.unwrap_or(&[]),
&mut keys,
);
}
}
}
keys
}
fn collect_tool_choice_keys(tools: &[Value], keys: &mut Vec<ToolChoiceKey>) {
for tool in tools {
let Some(tool_type) = tool_type(tool) else {
continue;
};
if matches!(tool_type, "function" | "custom") {
if let Some(name) = tool_name(tool) {
keys.push(ToolChoiceKey::Named {
name: name.to_ascii_lowercase(),
});
}
continue;
}
keys.push(ToolChoiceKey::Hosted(tool_type.to_ascii_lowercase()));
}
}
fn tool_matches_available(choice: &Value, available: &[ToolChoiceKey]) -> bool {
let Some(object) = choice.as_object() else {
return false;
};
let choice_type = object
.get("type")
.and_then(Value::as_str)
.unwrap_or_default()
.trim()
.to_ascii_lowercase();
if matches!(choice_type.as_str(), "function" | "custom" | "tool") {
let Some(name) = tool_choice_name(object) else {
return false;
};
return available.iter().any(|key| {
matches!(
key,
ToolChoiceKey::Named { name: available_name, .. }
if available_name == &name.to_ascii_lowercase()
)
});
}
if is_web_search_choice_type(&choice_type) {
return available.iter().any(
|key| matches!(key, ToolChoiceKey::Hosted(value) if value == XAI_WEB_SEARCH_TOOL_TYPE),
);
}
available
.iter()
.any(|key| matches!(key, ToolChoiceKey::Hosted(value) if value == &choice_type))
}
#[derive(Clone, Debug)]
enum ToolChoiceKey {
Named { name: String },
Hosted(String),
}
fn drop_tool_choice_without_tools(object: &mut Map<String, Value>) {
if xai_request_has_tools(object) {
return;
}
object.remove("tools");
object.remove("tool_choice");
object.remove("parallel_tool_calls");
}
fn xai_request_has_tools(object: &Map<String, Value>) -> bool {
if !tools_array(object).is_empty() {
return true;
}
object
.get("input")
.and_then(Value::as_array)
.into_iter()
.flatten()
.any(|item| {
item.get("type")
.and_then(Value::as_str)
.is_some_and(|value| value == "additional_tools")
&& item
.get("tools")
.and_then(Value::as_array)
.is_some_and(|tools| !tools.is_empty())
})
}
fn strip_unsupported_reasoning_effort(object: &mut Map<String, Value>) {
let model = object
.get("model")
.and_then(Value::as_str)
.unwrap_or_default();
if xai_model_supports_reasoning_effort(model) {
return;
}
let Some(reasoning) = object.get_mut("reasoning") else {
return;
};
let Some(reasoning_object) = reasoning.as_object_mut() else {
return;
};
reasoning_object.remove("effort");
if reasoning_object.is_empty() {
object.remove("reasoning");
}
}
pub fn xai_model_supports_reasoning_effort(model: &str) -> bool {
let lowered = model.trim().to_ascii_lowercase();
let name = lowered.rsplit('/').next().unwrap_or(lowered.as_str());
if name.is_empty() || name.contains("non-reasoning") || name.contains("imagine") {
return false;
}
name.starts_with("grok-3-mini")
|| name.starts_with("grok-4")
|| name.starts_with("grok-build")
|| name.starts_with("grok-composer")
}
pub fn xai_supports_native_image_generation(model: &str) -> bool {
let lowered = model.trim().to_ascii_lowercase();
let name = lowered.rsplit('/').next().unwrap_or(lowered.as_str());
let Some(rest) = name.strip_prefix("grok-") else {
return false;
};
if rest == "4.20" || rest.starts_with("4.20-") {
return false;
}
parse_grok_version_prefix(rest).is_some_and(grok_version_at_least_image_generation)
}
fn parse_grok_version_prefix(rest: &str) -> Option<XaiGrokVersion> {
let major_len = rest
.find(|ch: char| !ch.is_ascii_digit())
.unwrap_or(rest.len());
if major_len == 0 {
return None;
}
let major = rest[..major_len].parse().ok()?;
if major_len == rest.len() || !rest[major_len..].starts_with('.') {
return Some(XaiGrokVersion { major, minor: -1 });
}
let after_dot = &rest[major_len + 1..];
let minor_len = after_dot
.find(|ch: char| !ch.is_ascii_digit())
.unwrap_or(after_dot.len());
if minor_len == 0 {
return Some(XaiGrokVersion { major, minor: -1 });
}
let minor = after_dot[..minor_len].parse().ok()?;
Some(XaiGrokVersion { major, minor })
}
fn grok_version_at_least_image_generation(version: XaiGrokVersion) -> bool {
let minor = if version.minor < 0 { 0 } else { version.minor };
(version.major, minor)
>= (
XAI_GROK_IMAGE_GENERATION_MIN.major,
XAI_GROK_IMAGE_GENERATION_MIN.minor,
)
}
fn sanitize_xai_input_encrypted_content(object: &mut Map<String, Value>) {
let Some(input) = object.get_mut("input").and_then(Value::as_array_mut) else {
return;
};
let mut kept = Vec::new();
for item in input.iter() {
let Some(item_object) = item.as_object() else {
kept.push(item.clone());
continue;
};
let item_type = item_object
.get("type")
.and_then(Value::as_str)
.unwrap_or_default();
if item_type != "reasoning" && item_type != "compaction" {
kept.push(item.clone());
continue;
}
let Some(encrypted) = item_object.get("encrypted_content") else {
kept.push(item.clone());
continue;
};
let valid = encrypted
.as_str()
.is_some_and(|value| !value.trim().is_empty());
if valid {
kept.push(item.clone());
continue;
}
if item_type == "compaction" {
continue;
}
let mut next = item_object.clone();
next.remove("encrypted_content");
kept.push(Value::Object(next));
}
*input = kept;
}
fn normalize_xai_image_refs(value: &mut Value) {
match value {
Value::Object(object) => {
for key in ["image", "images", "reference_images"] {
match object.get_mut(key) {
Some(Value::Array(items)) if key != "image" => {
for item in items {
normalize_xai_image_ref(item);
}
}
Some(item) if key == "image" => normalize_xai_image_ref(item),
_ => {}
}
}
for child in object.values_mut() {
normalize_xai_image_refs(child);
}
}
Value::Array(items) => {
for item in items {
normalize_xai_image_refs(item);
}
}
_ => {}
}
}
fn normalize_xai_image_ref(value: &mut Value) {
let Some(object) = value.as_object_mut() else {
return;
};
let original_url = object
.get("url")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let image_url = object.get("image_url").cloned();
let resolved_url = original_url.clone().or_else(|| match image_url.as_ref() {
Some(Value::String(url)) => {
let trimmed = url.trim();
(!trimmed.is_empty()).then(|| trimmed.to_string())
}
Some(Value::Object(inner)) => inner
.get("url")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned),
_ => None,
});
let Some(url) = resolved_url else {
return;
};
if original_url.as_deref() == Some(url.as_str()) && image_url.is_none() {
return;
}
object.insert("url".to_string(), Value::String(url));
object.remove("image_url");
}
fn request_tools(body: &Value) -> &[Value] {
body.get("tools")
.and_then(Value::as_array)
.map(Vec::as_slice)
.unwrap_or(&[])
}
fn tools_array(object: &Map<String, Value>) -> &[Value] {
object
.get("tools")
.and_then(Value::as_array)
.map(Vec::as_slice)
.unwrap_or(&[])
}
fn tool_type(tool: &Value) -> Option<&str> {
tool.get("type").and_then(Value::as_str).map(str::trim)
}
fn tool_name(tool: &Value) -> Option<&str> {
tool.get("name")
.and_then(Value::as_str)
.or_else(|| {
tool.get("function")
.and_then(Value::as_object)
.and_then(|value| value.get("name"))
.and_then(Value::as_str)
})
.or_else(|| {
tool.get("custom")
.and_then(Value::as_object)
.and_then(|value| value.get("name"))
.and_then(Value::as_str)
})
.map(str::trim)
.filter(|value| !value.is_empty())
}
fn tool_choice_name(choice: &Map<String, Value>) -> Option<&str> {
choice
.get("name")
.and_then(Value::as_str)
.or_else(|| {
choice
.get("function")
.and_then(Value::as_object)
.and_then(|value| value.get("name"))
.and_then(Value::as_str)
})
.or_else(|| {
choice
.get("custom")
.and_then(Value::as_object)
.and_then(|value| value.get("name"))
.and_then(Value::as_str)
})
.map(str::trim)
.filter(|value| !value.is_empty())
}
fn is_web_search_tool(tool: &Value) -> bool {
tool_type(tool).is_some_and(is_web_search_choice_type)
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::{
apply_xai_upstream_payload_edits, apply_xai_upstream_payload_edits_with_client,
xai_model_supports_reasoning_effort, xai_supports_native_image_generation,
XAI_RESPONSES_UNSUPPORTED_BODY_FIELDS,
};
#[test]
fn xai_responses_edits_strip_continuation_fields_and_empty_tool_choice() {
let mut body = json!({
"model": "grok-4.6",
"input": "hello",
"previous_response_id": "resp_123",
"prompt_cache_retention": "24h",
"safety_identifier": "user-1",
"stream_options": {"include_obfuscation": true},
"stop": ["END"],
"metadata": {
"user_id": "{\"device_id\":\"dev-1\",\"account_uuid\":\"acct-1\",\"session_id\":\"sess-1\"}"
},
"include": ["reasoning.encrypted_content", "file_search_call.results"],
"tool_choice": "auto",
"parallel_tool_calls": true,
"tools": []
});
apply_xai_upstream_payload_edits(&mut body, "xai", "openai:responses");
for field in XAI_RESPONSES_UNSUPPORTED_BODY_FIELDS {
assert!(body.get(*field).is_none(), "{field} should be stripped");
}
assert!(body.get("tool_choice").is_none());
assert!(body.get("parallel_tool_calls").is_none());
assert!(body.get("tools").is_none());
assert_eq!(
body["include"],
json!(["reasoning.encrypted_content", "file_search_call.results"])
);
assert_eq!(body["model"], "grok-4.6");
assert_eq!(body["input"], "hello");
}
#[test]
fn xai_responses_edits_keep_reasoning_effort_for_thinking_models() {
let mut body = json!({
"model": "grok-4.6",
"reasoning": {"effort": "high", "summary": "auto"}
});
apply_xai_upstream_payload_edits(&mut body, "xai", "openai:responses");
assert_eq!(body["reasoning"]["effort"], "high");
assert_eq!(body["reasoning"]["summary"], "auto");
}
#[test]
fn xai_responses_edits_strip_reasoning_effort_for_non_thinking_models() {
let mut body = json!({
"model": "grok-4.20-0309-non-reasoning",
"reasoning": {"effort": "high"}
});
apply_xai_upstream_payload_edits(&mut body, "xai", "openai:responses");
assert!(body.get("reasoning").is_none());
assert!(!xai_model_supports_reasoning_effort(
"grok-4.20-0309-non-reasoning"
));
assert!(xai_model_supports_reasoning_effort("xai/grok-4.5"));
assert!(!xai_model_supports_reasoning_effort("grok-imagine-image"));
}
#[test]
fn xai_hosted_tool_choice_rewrites_web_search_and_image_generation() {
let mut web_search = json!({
"model": "grok-4.6",
"tools": [{"type": "web_search_preview", "name": "web_search"}],
"tool_choice": {"type": "web_search"}
});
apply_xai_upstream_payload_edits(&mut web_search, "xai", "openai:responses");
assert_eq!(web_search["tools"][0]["type"], "web_search");
assert!(web_search["tools"][0].get("name").is_none());
assert_eq!(web_search["tool_choice"]["type"], "allowed_tools");
assert_eq!(web_search["tool_choice"]["mode"], "required");
assert_eq!(web_search["tool_choice"]["tools"][0]["type"], "web_search");
let mut image = json!({
"model": "grok-4.6",
"tools": [
{"type": "web_search"},
{"type": "image_generation", "action": "generate"}
],
"tool_choice": {"type": "image_generation"}
});
apply_xai_upstream_payload_edits(&mut image, "xai", "openai:responses");
assert_eq!(image["tool_choice"], "required");
assert_eq!(image["tools"].as_array().map(Vec::len), Some(1));
assert_eq!(image["tools"][0]["type"], "image_generation");
}
#[test]
fn xai_strips_image_generation_on_older_conversation_models() {
let mut body = json!({
"model": "grok-4.5",
"tools": [
{"type": "function", "name": "lookup", "parameters": {"type": "object"}},
{"type": "image_generation"}
],
"tool_choice": {"type": "image_generation"}
});
apply_xai_upstream_payload_edits(&mut body, "xai", "openai:responses");
assert_eq!(body["tools"].as_array().map(Vec::len), Some(1));
assert_eq!(body["tools"][0]["name"], "lookup");
assert!(body.get("tool_choice").is_none());
assert!(xai_supports_native_image_generation("grok-4.6"));
assert!(!xai_supports_native_image_generation("grok-4.20-0309"));
assert!(!xai_supports_native_image_generation("grok-4.5"));
}
#[test]
fn xai_restores_web_search_from_chat_and_claude_clients() {
let mut chat_body = json!({
"model": "grok-4.6",
"input": "search this"
});
apply_xai_upstream_payload_edits_with_client(
&mut chat_body,
"xai",
"openai:responses",
Some("openai:chat"),
Some(&json!({
"messages": [{"role": "user", "content": "news"}],
"web_search_options": {"search_context_size": "high"}
})),
);
assert_eq!(chat_body["tools"][0]["type"], "web_search");
let mut claude_body = json!({
"model": "grok-4.6",
"input": "search this",
"tools": [{
"type": "function",
"name": "lookup",
"parameters": {"type": "object", "properties": {}}
}],
"tool_choice": {"type": "function", "name": "web_search"}
});
apply_xai_upstream_payload_edits_with_client(
&mut claude_body,
"xai",
"openai:responses",
Some("claude:messages"),
Some(&json!({
"tools": [
{"type": "web_search_20250305", "name": "web_search"},
{"name": "lookup", "input_schema": {"type": "object"}}
],
"tool_choice": {"type": "tool", "name": "web_search"}
})),
);
assert!(claude_body["tools"]
.as_array()
.into_iter()
.flatten()
.any(|tool| tool["type"] == "web_search"));
assert_eq!(claude_body["tool_choice"]["type"], "allowed_tools");
}
#[test]
fn xai_image_refs_rewrite_openai_aliases_without_touching_chat_parts() {
let mut body = json!({
"model": "grok-imagine-image",
"prompt": "edit this",
"image": {"image_url": "https://cdn.example/a.png"},
"reference_images": [
{"image_url": {"url": "https://cdn.example/b.png"}}
],
"input": [{
"type": "message",
"content": [{
"type": "image_url",
"image_url": {"url": "https://cdn.example/chat.png"}
}]
}]
});
apply_xai_upstream_payload_edits(&mut body, "xai", "openai:image");
assert_eq!(body["image"]["url"], "https://cdn.example/a.png");
assert!(body["image"].get("image_url").is_none());
assert_eq!(
body["reference_images"][0]["url"],
"https://cdn.example/b.png"
);
assert_eq!(
body["input"][0]["content"][0]["image_url"]["url"],
"https://cdn.example/chat.png"
);
}
#[test]
fn other_providers_are_left_untouched() {
let mut body = json!({
"previous_response_id": "resp_123",
"image": {"image_url": "https://cdn.example/a.png"}
});
apply_xai_upstream_payload_edits(&mut body, "codex", "openai:responses");
assert_eq!(body["previous_response_id"], "resp_123");
assert_eq!(body["image"]["image_url"], "https://cdn.example/a.png");
}
}
@@ -49,6 +49,10 @@ pub fn resolve_execution_runtime_stream_plan_kind_with_client_surface(
method: &Method, method: &Method,
path: &str, path: &str,
) -> Option<&'static str> { ) -> Option<&'static str> {
let path = path
.strip_prefix("/openai")
.filter(|p| *p == "/v1/videos" || p.starts_with("/v1/videos/"))
.unwrap_or(path);
if route_class != Some("ai_public") { if route_class != Some("ai_public") {
return None; return None;
} }
@@ -181,6 +185,10 @@ pub fn resolve_execution_runtime_sync_plan_kind_with_client_surface(
method: &Method, method: &Method,
path: &str, path: &str,
) -> Option<&'static str> { ) -> Option<&'static str> {
let path = path
.strip_prefix("/openai")
.filter(|p| *p == "/v1/videos" || p.starts_with("/v1/videos/"))
.unwrap_or(path);
if route_class != Some("ai_public") { if route_class != Some("ai_public") {
return None; return None;
} }
@@ -206,7 +214,10 @@ pub fn resolve_execution_runtime_sync_plan_kind_with_client_surface(
if route_family == Some("openai") if route_family == Some("openai")
&& route_kind == Some("video") && route_kind == Some("video")
&& *method == Method::POST && *method == Method::POST
&& path == "/v1/videos" && matches!(
path,
"/v1/videos" | "/v1/videos/generations" | "/v1/videos/edits" | "/v1/videos/extensions"
)
{ {
return Some(OPENAI_VIDEO_CREATE_SYNC_PLAN_KIND); return Some(OPENAI_VIDEO_CREATE_SYNC_PLAN_KIND);
} }
@@ -17,6 +17,7 @@ use crate::formats::openai::responses::codex::{
apply_codex_openai_responses_chat_body_edits, apply_codex_openai_responses_special_body_edits, apply_codex_openai_responses_chat_body_edits, apply_codex_openai_responses_special_body_edits,
apply_openai_responses_compact_special_body_edits, apply_openai_responses_compact_special_body_edits,
}; };
use crate::formats::openai::responses::xai::apply_xai_upstream_payload_edits_with_client;
use crate::formats::shared::standard_normalize::{ use crate::formats::shared::standard_normalize::{
build_local_openai_chat_request_body_with_model_directives, build_local_openai_chat_request_body_with_model_directives,
is_claude_messages_shaped_body_on_openai_chat_endpoint, is_claude_messages_shaped_body_on_openai_chat_endpoint,
@@ -121,6 +122,11 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and
enable_model_directives: bool, enable_model_directives: bool,
reasoning_replay_policy: crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy, reasoning_replay_policy: crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy,
) -> Option<Value> { ) -> Option<Value> {
let reasoning_replay_policy = if provider_type.trim().eq_ignore_ascii_case("xai") {
crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
} else {
reasoning_replay_policy
};
let mut format_context = FormatContext::default() let mut format_context = FormatContext::default()
.with_mapped_model(mapped_model) .with_mapped_model(mapped_model)
.with_request_path(request_path) .with_request_path(request_path)
@@ -133,13 +139,10 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and
client_api_format, client_api_format,
provider_api_format, provider_api_format,
); );
// DeepSeek's Responses continuation state is opaque. Parsing a same-wire-format // DeepSeek and xAI replay opaque provider state. Preserve their native
// request through the canonical model would discard its id-less `reasoning_text` // Responses input items: canonical conversion can lose reasoning IDs and
// items and future provider-owned fields even though no conversion is required. // encrypted-only items even when source and destination formats are equal.
// Keep that provider-specific route wire-preserving, while retaining canonical let mut provider_request_body = if is_wire_preserving_responses_hop(
// normalization for ordinary OpenAI Responses and for Responses/Compact
// cross-format conversions.
let mut provider_request_body = if is_wire_preserving_deepseek_responses_hop(
source_api_format.as_ref(), source_api_format.as_ref(),
provider_api_format, provider_api_format,
reasoning_replay_policy, reasoning_replay_policy,
@@ -200,6 +203,13 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and
&mut provider_request_body, &mut provider_request_body,
provider_api_format, provider_api_format,
); );
apply_xai_upstream_payload_edits_with_client(
&mut provider_request_body,
provider_type,
provider_api_format,
Some(client_api_format),
Some(body_json),
);
crate::formats::openai::responses::strip_incompatible_openai_responses_reasoning_items_with_policy( crate::formats::openai::responses::strip_incompatible_openai_responses_reasoning_items_with_policy(
&mut provider_request_body, &mut provider_request_body,
provider_api_format, provider_api_format,
@@ -224,14 +234,16 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and
Some(provider_request_body) Some(provider_request_body)
} }
fn is_wire_preserving_deepseek_responses_hop( fn is_wire_preserving_responses_hop(
source_api_format: &str, source_api_format: &str,
provider_api_format: &str, provider_api_format: &str,
reasoning_replay_policy: crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy, reasoning_replay_policy: crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy,
) -> bool { ) -> bool {
if reasoning_replay_policy if !matches!(
!= crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque reasoning_replay_policy,
{ crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
| crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
) {
return false; return false;
} }
let source_api_format = aether_ai_formats::normalize_api_format_alias(source_api_format); let source_api_format = aether_ai_formats::normalize_api_format_alias(source_api_format);
@@ -2077,4 +2089,316 @@ mod tests {
); );
assert_eq!(gemini["toolConfig"]["functionCallingConfig"]["mode"], "ANY"); assert_eq!(gemini["toolConfig"]["functionCallingConfig"]["mode"], "ANY");
} }
#[test]
fn xai_keeps_client_search_functions_distinct_from_hosted_search() {
for name in ["web_search", "web_search_internal"] {
for hosted in [false, true] {
let mut tools = vec![json!({
"name": name,
"description": "Search internal documents",
"input_schema": {"type": "object", "properties": {"query": {"type": "string"}}}
})];
if hosted {
tools.push(json!({"type": "web_search_20260209", "name": "internet_search"}));
}
let request = json!({
"model": "source", "max_tokens": 64,
"messages": [{"role": "user", "content": "Search internal documents"}],
"tools": tools,
"tool_choice": {"type": "tool", "name": name}
});
let converted = build_standard_request_body(
&request,
"claude:messages",
"grok-4.6",
"xai",
"openai:responses",
"/v1/messages",
true,
None,
None,
)
.unwrap();
assert_eq!(
converted["tool_choice"],
json!({"type": "function", "name": name})
);
assert_eq!(
converted["tools"]
.as_array()
.unwrap()
.iter()
.any(|tool| tool["type"] == "web_search"),
hosted
);
}
}
let request = json!({
"model": "source", "max_tokens": 64,
"messages": [{"role": "user", "content": "Search the internet"}],
"tools": [{"type": "web_search_20260209", "name": "internet_search"}],
"tool_choice": {"type": "tool", "name": "internet_search"}
});
let converted = build_standard_request_body(
&request,
"claude:messages",
"grok-4.6",
"xai",
"openai:responses",
"/v1/messages",
true,
None,
None,
)
.unwrap();
assert_eq!(
converted["tool_choice"],
json!({
"type": "allowed_tools", "mode": "required", "tools": [{"type": "web_search"}]
})
);
}
#[test]
fn xai_preserves_function_choices_in_chat_and_responses_requests() {
for name in ["web_search", "web_search_internal"] {
for (client, request) in [
(
"openai:chat",
json!({
"messages": [{"role": "user", "content": "search"}],
"tools": [{"type": "function", "function": {"name": name, "parameters": {"type": "object"}}}],
"tool_choice": {"type": "function", "function": {"name": name}}
}),
),
(
"openai:responses",
json!({
"input": "search",
"tools": [{"type": "function", "name": name, "parameters": {"type": "object"}}],
"tool_choice": {"type": "function", "name": name}
}),
),
] {
let converted = build_standard_request_body(
&request,
client,
"grok-4.6",
"xai",
"openai:responses",
"/v1/responses",
true,
None,
None,
)
.unwrap();
assert_eq!(
converted["tool_choice"],
json!({"type": "function", "name": name})
);
assert_eq!(converted["tools"].as_array().unwrap().len(), 1);
}
}
}
#[test]
fn xai_image_allowed_tools_preserves_mode_and_restricts_available_tools() {
for mode in ["auto", "required"] {
for mixed in [false, true] {
let mut allowed = vec![json!({"type": "image_generation"})];
if mixed {
allowed.push(json!({"type": "function", "name": "lookup"}));
}
let request = json!({
"input": "Draw a cat",
"tools": [
{"type": "web_search"}, {"type": "image_generation"},
{"type": "function", "name": "lookup", "parameters": {"type": "object"}}
],
"tool_choice": {"type": "allowed_tools", "mode": mode, "tools": allowed}
});
let converted = build_standard_request_body(
&request,
"openai:responses",
"grok-4.6",
"xai",
"openai:responses",
"/v1/responses",
true,
None,
None,
)
.unwrap();
if mixed {
assert_eq!(
converted["tool_choice"],
json!({
"type": "allowed_tools", "mode": mode,
"tools": [{"type": "function", "name": "lookup"}]
})
);
assert_eq!(converted["tools"].as_array().unwrap().len(), 3);
} else {
assert_eq!(converted["tool_choice"], mode);
assert_eq!(converted["tools"], json!([{"type": "image_generation"}]));
}
}
}
}
#[test]
fn xai_responses_preserves_requested_encrypted_reasoning_and_replayed_input() {
let reasoning = json!({"type": "reasoning", "id": "550e8400-e29b-41d4-a716-446655440000", "summary": [], "encrypted_content": "opaque-xai-state"});
let request = json!({
"input": [reasoning.clone(), {"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "Previous answer"}]}, {"role": "user", "content": "Continue"}],
"include": ["reasoning.encrypted_content"], "store": false
});
let converted = build_standard_request_body(
&request,
"openai:responses",
"grok-4.6",
"xai",
"openai:responses",
"/v1/responses",
true,
None,
None,
)
.unwrap();
assert_eq!(converted["include"], request["include"]);
assert_eq!(converted["input"][0], reasoning);
assert_eq!(converted["store"], false);
}
#[test]
fn xai_standard_conversion_strips_unsupported_responses_fields() {
let request = json!({
"model": "source-model",
"messages": [{"role": "user", "content": "Hello xAI"}],
"max_tokens": 128,
"stop": ["END"],
"stream_options": {"include_usage": true},
"metadata": {"user_id": "claude-session"},
"web_search_options": {"search_context_size": "high"}
});
let converted = build_standard_request_body(
&request,
"openai:chat",
"grok-4.6",
"xai",
"openai:responses",
"/v1/chat/completions",
true,
None,
None,
)
.expect("chat should convert onto xAI Responses");
assert_eq!(converted["model"], "grok-4.6");
assert!(converted.get("stop").is_none());
assert!(converted.get("stream_options").is_none());
assert!(converted.get("previous_response_id").is_none());
assert!(converted.get("metadata").is_none());
assert!(converted.get("input").is_some() || converted.get("messages").is_none());
assert_eq!(converted["max_output_tokens"], 128);
assert_eq!(converted["tools"][0]["type"], "web_search");
}
#[test]
fn xai_standard_conversion_covers_claude_and_gemini_clients() {
let claude = json!({
"model": "claude-sonnet",
"max_tokens": 64,
"messages": [{"role": "user", "content": "Hello xAI"}],
"metadata": {
"user_id": "{\"device_id\":\"dev-1\",\"account_uuid\":\"acct-1\",\"session_id\":\"sess-1\"}"
},
"tools": [
{"type": "web_search_20250305", "name": "web_search"},
{
"name": "lookup",
"description": "Look something up",
"input_schema": {"type": "object", "properties": {}}
}
],
"tool_choice": {"type": "tool", "name": "web_search"}
});
let converted = build_standard_request_body(
&claude,
"claude:messages",
"grok-4.6",
"xai",
"openai:responses",
"/v1/messages",
true,
None,
None,
)
.expect("claude should convert onto xAI Responses");
assert_eq!(converted["model"], "grok-4.6");
assert!(converted.get("metadata").is_none());
assert!(converted.get("context_management").is_none());
assert!(converted
.get("include")
.and_then(Value::as_array)
.into_iter()
.flatten()
.any(|item| item == "reasoning.encrypted_content"));
assert!(converted["tools"]
.as_array()
.into_iter()
.flatten()
.any(|tool| tool["type"] == "web_search"));
assert_eq!(converted["tool_choice"]["type"], "allowed_tools");
assert!(converted.get("input").is_some());
let gemini = json!({
"model": "gemini-2.5-pro",
"contents": [{
"role": "user",
"parts": [{"text": "Hello xAI"}]
}],
"tools": [{"googleSearch": {}}]
});
let converted = build_standard_request_body(
&gemini,
"gemini:generate_content",
"grok-4.6",
"xai",
"openai:responses",
"/v1beta/models/gemini-2.5-pro:generateContent",
false,
None,
None,
)
.expect("gemini should convert onto xAI Responses");
assert_eq!(converted["model"], "grok-4.6");
assert_eq!(converted["tools"][0]["type"], "web_search");
assert!(converted.get("input").is_some());
let same_format = json!({
"model": "grok-4.6",
"input": "hello",
"previous_response_id": "resp_123",
"stop": ["END"],
"metadata": {"user_id": "claude-session"}
});
let converted = build_standard_request_body(
&same_format,
"openai:responses",
"grok-4.6",
"xai",
"openai:responses",
"/v1/responses",
true,
None,
None,
)
.expect("same-format xAI Responses should sanitize in place");
assert!(converted.get("previous_response_id").is_none());
assert!(converted.get("stop").is_none());
assert!(converted.get("metadata").is_none());
}
} }
+4
View File
@@ -56,6 +56,10 @@ pub use formats::openai::responses::codex::{
pub use formats::openai::responses::request::{ pub use formats::openai::responses::request::{
validate_openai_responses_request_contract, OpenAiResponsesRequestContractViolation, validate_openai_responses_request_contract, OpenAiResponsesRequestContractViolation,
}; };
pub use formats::openai::responses::xai::{
apply_xai_upstream_payload_edits, apply_xai_upstream_payload_edits_with_client,
xai_model_supports_reasoning_effort, xai_supports_native_image_generation,
};
pub use formats::openai::responses::{ pub use formats::openai::responses::{
normalize_openai_responses_message_item_ids, openai_responses_message_item_id, normalize_openai_responses_message_item_ids, openai_responses_message_item_id,
openai_responses_request_operation, openai_responses_synthetic_reasoning_item_id, openai_responses_request_operation, openai_responses_synthetic_reasoning_item_id,
@@ -102,6 +102,11 @@ INNER JOIN LATERAL (
AND LOWER(BTRIM(pak.auth_type)) = 'oauth' AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
AND LOWER($3) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image') AND LOWER($3) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
) )
OR (
LOWER(BTRIM(p.provider_type)) = 'xai'
AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')
AND LOWER($3) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video')
)
OR ( OR (
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
AND LOWER(BTRIM(pak.auth_type)) = 'oauth' AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
@@ -127,7 +132,8 @@ INNER JOIN LATERAL (
'vertex_ai', 'vertex_ai',
'antigravity', 'antigravity',
'kiro', 'kiro',
'windsurf' 'windsurf',
'xai'
) )
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth' AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
) )
@@ -187,6 +193,11 @@ WHERE p.is_active = TRUE
AND LOWER(BTRIM(pak.auth_type)) = 'oauth' AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
AND LOWER($3) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image') AND LOWER($3) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
) )
OR (
LOWER(BTRIM(p.provider_type)) = 'xai'
AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')
AND LOWER($3) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video')
)
OR ( OR (
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
AND LOWER(BTRIM(pak.auth_type)) = 'oauth' AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
@@ -212,7 +223,8 @@ WHERE p.is_active = TRUE
'vertex_ai', 'vertex_ai',
'antigravity', 'antigravity',
'kiro', 'kiro',
'windsurf' 'windsurf',
'xai'
) )
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth' AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
) )
@@ -365,6 +377,11 @@ INNER JOIN LATERAL (
AND LOWER(BTRIM(pak.auth_type)) = 'oauth' AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
AND LOWER($4) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image') AND LOWER($4) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
) )
OR (
LOWER(BTRIM(p.provider_type)) = 'xai'
AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')
AND LOWER($4) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video')
)
OR ( OR (
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
AND LOWER(BTRIM(pak.auth_type)) = 'oauth' AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
@@ -390,7 +407,8 @@ INNER JOIN LATERAL (
'vertex_ai', 'vertex_ai',
'antigravity', 'antigravity',
'kiro', 'kiro',
'windsurf' 'windsurf',
'xai'
) )
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth' AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
) )
@@ -451,6 +469,11 @@ WHERE p.is_active = TRUE
AND LOWER(BTRIM(pak.auth_type)) = 'oauth' AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
AND LOWER($4) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image') AND LOWER($4) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
) )
OR (
LOWER(BTRIM(p.provider_type)) = 'xai'
AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')
AND LOWER($4) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video')
)
OR ( OR (
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
AND LOWER(BTRIM(pak.auth_type)) = 'oauth' AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
@@ -476,7 +499,8 @@ WHERE p.is_active = TRUE
'vertex_ai', 'vertex_ai',
'antigravity', 'antigravity',
'kiro', 'kiro',
'windsurf' 'windsurf',
'xai'
) )
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth' AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
) )
@@ -632,11 +656,16 @@ WHERE p.is_active = TRUE
) )
) )
) )
OR ( OR (
LOWER(BTRIM(p.provider_type)) = 'grok' LOWER(BTRIM(p.provider_type)) = 'grok'
AND LOWER(BTRIM(pak.auth_type)) = 'oauth' AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
AND LOWER($6) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image') AND LOWER($6) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
) )
OR (
LOWER(BTRIM(p.provider_type)) = 'xai'
AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')
AND LOWER($6) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video')
)
OR ( OR (
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
AND LOWER(BTRIM(pak.auth_type)) = 'oauth' AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
@@ -662,7 +691,8 @@ WHERE p.is_active = TRUE
'vertex_ai', 'vertex_ai',
'antigravity', 'antigravity',
'kiro', 'kiro',
'windsurf' 'windsurf',
'xai'
) )
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth' AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
) )
@@ -1717,6 +1747,24 @@ mod tests {
} }
} }
#[test]
fn candidate_selection_sql_allows_xai_oauth_responses_auth() {
let requested_model_sql = requested_model_selection_sql();
for sql in [
LIST_FOR_EXACT_API_FORMAT_SQL,
LIST_FOR_EXACT_API_FORMAT_AND_GLOBAL_MODEL_SQL,
LIST_POOL_KEYS_FOR_GROUP_SQL,
requested_model_sql.as_str(),
] {
assert!(sql.contains("LOWER(BTRIM(p.provider_type)) = 'xai'"));
assert!(sql.contains("LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')"));
assert!(sql.contains(
"'openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video'"
));
assert!(sql.contains("'xai'"));
}
}
#[test] #[test]
fn candidate_selection_sql_allows_windsurf_openai_chat_managed_keys() { fn candidate_selection_sql_allows_windsurf_openai_chat_managed_keys() {
let requested_model_sql = requested_model_selection_sql(); let requested_model_sql = requested_model_selection_sql();
@@ -346,6 +346,16 @@ fn key_auth_channel_matches(row: &StoredMinimalCandidateSelectionRow, api_format
"openai:chat" | "openai:responses" | "claude:messages" | "openai:image" "openai:chat" | "openai:responses" | "claude:messages" | "openai:image"
) )
} }
"xai" => {
matches!(auth_type.as_str(), "oauth" | "bearer" | "api_key")
&& matches!(
api_format.as_str(),
"openai:responses"
| "openai:responses:compact"
| "openai:image"
| "openai:video"
)
}
"windsurf" => { "windsurf" => {
matches!(auth_type.as_str(), "oauth" | "api_key" | "bearer") matches!(auth_type.as_str(), "oauth" | "api_key" | "bearer")
&& api_format == "openai:chat" && api_format == "openai:chat"
@@ -591,6 +601,59 @@ mod tests {
assert_eq!(rows[0].global_model_name, "grok-4.20-0309-non-reasoning"); assert_eq!(rows[0].global_model_name, "grok-4.20-0309-non-reasoning");
} }
#[tokio::test]
async fn includes_xai_oauth_rows_for_responses_models() {
let mut row = sample_row("provider-xai", "openai:responses", "grok-4", 10);
row.provider_type = "xai".to_string();
row.provider_name = "xai".to_string();
row.key_auth_type = "oauth".to_string();
row.key_api_formats = Some(vec![
"openai:responses".to_string(),
"openai:responses:compact".to_string(),
]);
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![row]);
let rows = repository
.list_for_exact_api_format("openai:responses")
.await
.expect("list should succeed");
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].provider_type, "xai");
assert_eq!(rows[0].global_model_name, "grok-4");
}
#[tokio::test]
async fn includes_xai_oauth_rows_for_image_and_video_models() {
let mut image = sample_row("provider-xai", "openai:image", "grok-imagine-image", 10);
image.provider_type = "xai".to_string();
image.provider_name = "xai".to_string();
image.key_auth_type = "oauth".to_string();
image.key_api_formats = Some(vec!["openai:image".to_string(), "openai:video".to_string()]);
let mut video = image.clone();
video.endpoint_id = "endpoint-video".to_string();
video.endpoint_api_format = "openai:video".to_string();
video.global_model_name = "grok-imagine-video".to_string();
video.model_provider_model_name = "grok-imagine-video".to_string();
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![image, video]);
let image_rows = repository
.list_for_exact_api_format("openai:image")
.await
.expect("list should succeed");
assert_eq!(image_rows.len(), 1);
assert_eq!(image_rows[0].global_model_name, "grok-imagine-image");
let video_rows = repository
.list_for_exact_api_format("openai:video")
.await
.expect("list should succeed");
assert_eq!(video_rows.len(), 1);
assert_eq!(video_rows[0].global_model_name, "grok-imagine-video");
}
#[tokio::test] #[tokio::test]
async fn requested_model_filter_respects_endpoint_scoped_default_mapping() { async fn requested_model_filter_respects_endpoint_scoped_default_mapping() {
let mut selected = sample_row("provider-1", "openai:chat", "deepseek-v4-pro", 10); let mut selected = sample_row("provider-1", "openai:chat", "deepseek-v4-pro", 10);
+52 -1
View File
@@ -546,7 +546,7 @@ pub fn endpoint_supports_rust_models_fetch(api_format: &str) -> bool {
pub fn provider_type_uses_preset_models(provider_type: &str) -> bool { pub fn provider_type_uses_preset_models(provider_type: &str) -> bool {
matches!( matches!(
provider_type.trim().to_ascii_lowercase().as_str(), provider_type.trim().to_ascii_lowercase().as_str(),
"claude_code" | "gemini_cli" | "grok" "claude_code" | "gemini_cli" | "grok" | "xai"
) )
} }
@@ -604,6 +604,22 @@ pub fn preset_models_for_provider(provider_type: &str) -> Option<Vec<Value>> {
preset_model("grok-imagine-image-pro", "xai", "Grok Imagine Image Pro", "openai:image"), preset_model("grok-imagine-image-pro", "xai", "Grok Imagine Image Pro", "openai:image"),
preset_model("grok-imagine-image-edit", "xai", "Grok Imagine Image Edit", "openai:image"), preset_model("grok-imagine-image-edit", "xai", "Grok Imagine Image Edit", "openai:image"),
], ],
"xai" => vec![
preset_model("grok-4.6", "xai", "Grok 4.6", "openai:responses"),
preset_model("grok-build-0.1", "xai", "Grok Build 0.1", "openai:responses"),
preset_model("grok-4.5", "xai", "Grok 4.5", "openai:responses"),
preset_model("grok-4.3", "xai", "Grok 4.3", "openai:responses"),
preset_model("grok-4.20-0309-reasoning", "xai", "Grok 4.20 0309 Reasoning", "openai:responses"),
preset_model("grok-4.20-0309-non-reasoning", "xai", "Grok 4.20 0309 Non-Reasoning", "openai:responses"),
preset_model("grok-4.20-multi-agent-0309", "xai", "Grok 4.20 Multi-Agent 0309", "openai:responses"),
preset_model("grok-3-mini", "xai", "Grok 3 Mini", "openai:responses"),
preset_model("grok-3-mini-fast", "xai", "Grok 3 Mini Fast", "openai:responses"),
preset_model("grok-composer-2.5-fast", "xai", "Grok Composer 2.5 Fast", "openai:responses"),
preset_model("grok-imagine-image", "xai", "Grok Imagine Image", "openai:image"),
preset_model("grok-imagine-image-quality", "xai", "Grok Imagine Image Quality", "openai:image"),
preset_model("grok-imagine-video", "xai", "Grok Imagine Video", "openai:video"),
preset_model("grok-imagine-video-1.5", "xai", "Grok Imagine Video 1.5", "openai:video"),
],
_ => return None, _ => return None,
}; };
Some(models) Some(models)
@@ -1977,4 +1993,39 @@ mod tests {
assert_eq!(models[15]["api_formats"], json!(["openai:image"])); assert_eq!(models[15]["api_formats"], json!(["openai:image"]));
assert_eq!(models[18]["api_formats"], json!(["openai:image"])); assert_eq!(models[18]["api_formats"], json!(["openai:image"]));
} }
#[test]
fn preset_models_cover_xai_cli_catalog() {
let models = preset_models_for_provider("xai").expect("preset models should exist");
let model_ids = models
.iter()
.map(|model| model["id"].as_str().expect("model id"))
.collect::<Vec<_>>();
assert_eq!(
model_ids,
vec![
"grok-4.6",
"grok-build-0.1",
"grok-4.5",
"grok-4.3",
"grok-4.20-0309-reasoning",
"grok-4.20-0309-non-reasoning",
"grok-4.20-multi-agent-0309",
"grok-3-mini",
"grok-3-mini-fast",
"grok-composer-2.5-fast",
"grok-imagine-image",
"grok-imagine-image-quality",
"grok-imagine-video",
"grok-imagine-video-1.5",
]
);
assert!(models.iter().all(|model| model["owned_by"] == json!("xai")));
assert_eq!(models[0]["api_formats"], json!(["openai:responses"]));
assert_eq!(models[10]["api_formats"], json!(["openai:image"]));
assert_eq!(models[12]["api_formats"], json!(["openai:video"]));
assert!(models
.iter()
.any(|model| model["id"] == "grok-imagine-image"));
}
} }
@@ -150,6 +150,27 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[
uses_json_payload: false, uses_json_payload: false,
include_scope_in_token_request: true, include_scope_in_token_request: true,
}, },
GenericProviderOAuthTemplate {
provider_type: "xai",
display_name: "xAI",
authorize_url: "https://auth.x.ai/oauth2/device/code",
token_url: "https://auth.x.ai/oauth2/token",
client_id: "b1a00492-073a-47ea-816f-4c329264a828",
client_id_env: None,
client_secret_env: None,
scopes: &[
"openid",
"profile",
"email",
"offline_access",
"grok-cli:access",
"api:access",
],
redirect_uri: "",
use_pkce: false,
uses_json_payload: false,
include_scope_in_token_request: false,
},
]; ];
#[derive(Clone)] #[derive(Clone)]
@@ -212,6 +233,10 @@ impl GenericProviderOAuthAdapter {
self self
} }
pub(super) fn token_url_for_provider(&self) -> String {
self.token_url()
}
fn token_url(&self) -> String { fn token_url(&self) -> String {
self.token_url_override self.token_url_override
.clone() .clone()
@@ -389,7 +414,10 @@ impl GenericProviderOAuthAdapter {
self.token_set_from_payload(payload) self.token_set_from_payload(payload)
} }
fn token_set_from_payload(&self, payload: Value) -> Result<ProviderOAuthTokenSet, OAuthError> { pub(super) fn token_set_from_payload(
&self,
payload: Value,
) -> Result<ProviderOAuthTokenSet, OAuthError> {
let token_set = OAuthTokenSet::from_token_payload(payload.clone()) let token_set = OAuthTokenSet::from_token_payload(payload.clone())
.ok_or_else(|| OAuthError::invalid_response("token response missing access_token"))?; .ok_or_else(|| OAuthError::invalid_response("token response missing access_token"))?;
let mut auth_config = serde_json::Map::new(); let mut auth_config = serde_json::Map::new();
@@ -945,6 +973,7 @@ mod tests {
fn resolves_generic_provider_templates() { fn resolves_generic_provider_templates() {
assert!(template_for_provider_type("codex").is_some()); assert!(template_for_provider_type("codex").is_some());
assert!(template_for_provider_type("claude_code").is_some()); assert!(template_for_provider_type("claude_code").is_some());
assert!(template_for_provider_type("xai").is_some());
assert!(template_for_provider_type("kiro").is_none()); assert!(template_for_provider_type("kiro").is_none());
} }
@@ -4,6 +4,7 @@ mod codex;
mod generic; mod generic;
mod kiro; mod kiro;
mod windsurf; mod windsurf;
mod xai;
pub use antigravity::{AntigravityProviderOAuthAdapter, ANTIGRAVITY_USER_INFO_URL}; pub use antigravity::{AntigravityProviderOAuthAdapter, ANTIGRAVITY_USER_INFO_URL};
pub use claude_code::{ pub use claude_code::{
@@ -27,3 +28,7 @@ pub use windsurf::{
WindsurfProviderOAuthAdapter, WINDSURF_CLIENT_ID, WINDSURF_PROVIDER_TYPE, WindsurfProviderOAuthAdapter, WINDSURF_CLIENT_ID, WINDSURF_PROVIDER_TYPE,
WINDSURF_SHOW_AUTH_TOKEN_REDIRECT, WINDSURF_SIGNIN_URL, WINDSURF_SHOW_AUTH_TOKEN_REDIRECT, WINDSURF_SIGNIN_URL,
}; };
pub use xai::{
XaiDevicePollOutcome, XaiProviderOAuthAdapter, XAI_CLIENT_ID, XAI_DEVICE_CODE_GRANT_TYPE,
XAI_DEVICE_CODE_URL, XAI_OAUTH_SCOPES, XAI_PROVIDER_TYPE, XAI_TOKEN_URL,
};
@@ -0,0 +1,668 @@
use super::generic::{template_for_provider_type, GenericProviderOAuthAdapter};
use crate::core::{
current_unix_secs, redacted_oauth_error_body_excerpt, OAuthDeviceAuthorization, OAuthError,
};
use crate::network::{OAuthHttpExecutor, OAuthHttpRequest};
use crate::provider::{
ProviderOAuthAccount, ProviderOAuthAdapter, ProviderOAuthCapabilities,
ProviderOAuthImportInput, ProviderOAuthRequestAuth, ProviderOAuthTokenSet,
ProviderOAuthTransportContext,
};
use async_trait::async_trait;
use serde_json::{json, Map, Value};
use std::collections::BTreeMap;
use url::form_urlencoded;
pub const XAI_PROVIDER_TYPE: &str = "xai";
pub const XAI_DEVICE_CODE_URL: &str = "https://auth.x.ai/oauth2/device/code";
pub const XAI_TOKEN_URL: &str = "https://auth.x.ai/oauth2/token";
pub const XAI_CLIENT_ID: &str = "b1a00492-073a-47ea-816f-4c329264a828";
pub const XAI_OAUTH_SCOPES: &[&str] = &[
"openid",
"profile",
"email",
"offline_access",
"grok-cli:access",
"api:access",
];
pub const XAI_DEVICE_CODE_GRANT_TYPE: &str = "urn:ietf:params:oauth:grant-type:device_code";
const DEFAULT_DEVICE_EXPIRES_IN_SECS: u64 = 600;
const DEFAULT_DEVICE_POLL_INTERVAL_SECS: u64 = 5;
#[derive(Debug, Clone, PartialEq)]
pub enum XaiDevicePollOutcome {
Pending,
SlowDown,
Authorized(Box<ProviderOAuthTokenSet>),
}
#[derive(Clone)]
pub struct XaiProviderOAuthAdapter {
inner: GenericProviderOAuthAdapter,
device_url_override: Option<String>,
}
impl std::fmt::Debug for XaiProviderOAuthAdapter {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("XaiProviderOAuthAdapter")
.field(
"has_device_url_override",
&self.device_url_override.is_some(),
)
.finish_non_exhaustive()
}
}
impl Default for XaiProviderOAuthAdapter {
fn default() -> Self {
Self {
inner: GenericProviderOAuthAdapter::new(
template_for_provider_type(XAI_PROVIDER_TYPE).expect("xai template should exist"),
),
device_url_override: None,
}
}
}
impl XaiProviderOAuthAdapter {
pub fn with_endpoint_overrides(
mut self,
device_url: impl Into<String>,
token_url: impl Into<String>,
) -> Self {
self.device_url_override = Some(device_url.into());
self.inner = self.inner.with_token_url_override(token_url);
self
}
fn device_url(&self) -> String {
self.device_url_override
.clone()
.unwrap_or_else(|| XAI_DEVICE_CODE_URL.to_string())
}
pub async fn start_device_flow(
&self,
executor: &dyn OAuthHttpExecutor,
ctx: &ProviderOAuthTransportContext,
) -> Result<OAuthDeviceAuthorization, OAuthError> {
let form = form_urlencoded::Serializer::new(String::new())
.append_pair("client_id", XAI_CLIENT_ID)
.append_pair("scope", &XAI_OAUTH_SCOPES.join(" "))
.finish()
.into_bytes();
let response = executor
.execute(OAuthHttpRequest {
request_id: "provider-oauth:xai-device-code".to_string(),
method: reqwest::Method::POST,
url: self.device_url(),
headers: form_headers(),
content_type: Some("application/x-www-form-urlencoded".to_string()),
json_body: None,
body_bytes: Some(form),
network: ctx.network.clone(),
transport_profile: None,
})
.await?;
if !(200..300).contains(&response.status_code) {
return Err(OAuthError::HttpStatus {
status_code: response.status_code,
body_excerpt: redacted_oauth_error_body_excerpt(&response.body_text),
});
}
let payload = response_json(&response)
.ok_or_else(|| OAuthError::invalid_response("xAI device code response is not json"))?;
parse_device_authorization(&payload)
}
pub async fn poll_device_token(
&self,
executor: &dyn OAuthHttpExecutor,
ctx: &ProviderOAuthTransportContext,
device_code: &str,
) -> Result<XaiDevicePollOutcome, OAuthError> {
let device_code = device_code.trim();
if device_code.is_empty() {
return Err(OAuthError::invalid_request("xAI device_code is required"));
}
let form = form_urlencoded::Serializer::new(String::new())
.append_pair("grant_type", XAI_DEVICE_CODE_GRANT_TYPE)
.append_pair("device_code", device_code)
.append_pair("client_id", XAI_CLIENT_ID)
.finish()
.into_bytes();
let response = executor
.execute(OAuthHttpRequest {
request_id: "provider-oauth:xai-device-token".to_string(),
method: reqwest::Method::POST,
url: self.inner.token_url_for_provider(),
headers: form_headers(),
content_type: Some("application/x-www-form-urlencoded".to_string()),
json_body: None,
body_bytes: Some(form),
network: ctx.network.clone(),
transport_profile: None,
})
.await?;
let payload = response_json(&response);
if let Some(error_code) = payload.as_ref().and_then(oauth_error_code) {
return match error_code.as_str() {
"authorization_pending" => Ok(XaiDevicePollOutcome::Pending),
"slow_down" => Ok(XaiDevicePollOutcome::SlowDown),
"expired_token" => Err(OAuthError::invalid_request("xAI device code expired")),
"access_denied" => Err(OAuthError::invalid_request(
"xAI device authorization denied",
)),
other => Err(OAuthError::invalid_response(format!(
"xAI device token error: {other}"
))),
};
}
if !(200..300).contains(&response.status_code) {
return Err(OAuthError::HttpStatus {
status_code: response.status_code,
body_excerpt: redacted_oauth_error_body_excerpt(&response.body_text),
});
}
let payload = payload
.ok_or_else(|| OAuthError::invalid_response("xAI device token response is not json"))?;
let mut token_set = self.inner.token_set_from_payload(payload)?;
let raw_payload = token_set.token_set.raw_payload.clone();
mark_oauth_auth_config(&mut token_set.auth_config);
enrich_xai_identity(&mut token_set.auth_config, raw_payload.as_ref());
Ok(XaiDevicePollOutcome::Authorized(Box::new(token_set)))
}
async fn import_raw_api_key(
&self,
input: &ProviderOAuthImportInput,
api_key: &str,
) -> Result<ProviderOAuthTokenSet, OAuthError> {
let api_key = api_key.trim();
if api_key.is_empty() {
return Err(OAuthError::invalid_request("xAI api_key is required"));
}
let mut auth_config = Map::new();
auth_config.insert("provider_type".to_string(), json!(XAI_PROVIDER_TYPE));
auth_config.insert("auth_method".to_string(), json!("api_key"));
auth_config.insert("using_api".to_string(), json!(true));
auth_config.insert("updated_at".to_string(), json!(current_unix_secs()));
if let Some(name) = input
.name
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
auth_config.insert("name".to_string(), json!(name));
}
Ok(ProviderOAuthTokenSet {
token_set: crate::core::OAuthTokenSet {
access_token: api_key.to_string(),
refresh_token: None,
token_type: Some("Bearer".to_string()),
scope: None,
expires_at_unix_secs: None,
raw_payload: None,
},
auth_config: Value::Object(auth_config),
})
}
}
#[async_trait]
impl ProviderOAuthAdapter for XaiProviderOAuthAdapter {
fn provider_type(&self) -> &'static str {
XAI_PROVIDER_TYPE
}
fn capabilities(&self) -> ProviderOAuthCapabilities {
ProviderOAuthCapabilities {
supports_authorization_code: false,
supports_cookie_authorization: false,
supports_refresh_token_import: true,
supports_batch_import: true,
supports_device_flow: true,
supports_account_probe: false,
rotates_refresh_token: true,
}
}
async fn import_credentials(
&self,
executor: &dyn OAuthHttpExecutor,
ctx: &ProviderOAuthTransportContext,
input: ProviderOAuthImportInput,
) -> Result<ProviderOAuthTokenSet, OAuthError> {
if let Some(api_key) =
raw_credential_string(input.raw_credentials.as_ref(), &["api_key", "apiKey"])
{
return self.import_raw_api_key(&input, &api_key).await;
}
let refresh_token = input
.refresh_token
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.or_else(|| {
raw_credential_string(
input.raw_credentials.as_ref(),
&["refresh_token", "refreshToken"],
)
});
if let Some(refresh_token) = refresh_token {
let mut imported = self
.inner
.import_credentials(
executor,
ctx,
ProviderOAuthImportInput {
refresh_token: Some(refresh_token),
..input
},
)
.await?;
let raw_payload = imported.token_set.raw_payload.clone();
mark_oauth_auth_config(&mut imported.auth_config);
enrich_xai_identity(&mut imported.auth_config, raw_payload.as_ref());
return Ok(imported);
}
if let Some(access_token) = raw_credential_string(
input.raw_credentials.as_ref(),
&["access_token", "accessToken"],
) {
return self.import_raw_api_key(&input, &access_token).await;
}
Err(OAuthError::invalid_request(
"xAI credentials require api_key, access_token, or refresh_token",
))
}
async fn refresh(
&self,
executor: &dyn OAuthHttpExecutor,
ctx: &ProviderOAuthTransportContext,
account: &ProviderOAuthAccount,
) -> Result<ProviderOAuthTokenSet, OAuthError> {
let mut refreshed = self.inner.refresh(executor, ctx, account).await?;
let raw_payload = refreshed.token_set.raw_payload.clone();
mark_oauth_auth_config(&mut refreshed.auth_config);
enrich_xai_identity(&mut refreshed.auth_config, raw_payload.as_ref());
Ok(refreshed)
}
fn resolve_request_auth(
&self,
account: &ProviderOAuthAccount,
) -> Result<ProviderOAuthRequestAuth, OAuthError> {
self.inner.resolve_request_auth(account)
}
fn account_fingerprint(&self, account: &ProviderOAuthAccount) -> Option<String> {
self.inner.account_fingerprint(account)
}
}
fn mark_oauth_auth_config(auth_config: &mut Value) {
let Some(object) = auth_config.as_object_mut() else {
return;
};
object.insert("provider_type".to_string(), json!(XAI_PROVIDER_TYPE));
object.insert("auth_method".to_string(), json!("oauth"));
object.insert("using_api".to_string(), json!(false));
}
fn enrich_xai_identity(auth_config: &mut Value, raw_payload: Option<&Value>) {
let Some(object) = auth_config.as_object_mut() else {
return;
};
let id_token = raw_payload
.and_then(|payload| payload.get("id_token"))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty());
if let Some(id_token) = id_token {
object
.entry("id_token".to_string())
.or_insert_with(|| json!(id_token));
if let Some(claims) = decode_jwt_claims(id_token) {
if !object.contains_key("email") {
if let Some(email) = claims.get("email").and_then(Value::as_str) {
let email = email.trim();
if !email.is_empty() {
object.insert("email".to_string(), json!(email));
}
}
}
if !object.contains_key("sub") {
if let Some(sub) = claims.get("sub").and_then(Value::as_str) {
let sub = sub.trim();
if !sub.is_empty() {
object.insert("sub".to_string(), json!(sub));
}
}
}
}
}
}
fn parse_device_authorization(payload: &Value) -> Result<OAuthDeviceAuthorization, OAuthError> {
let device_code =
json_non_empty_string(payload, &["device_code", "deviceCode"]).ok_or_else(|| {
OAuthError::invalid_response("xAI device code response missing device_code")
})?;
let user_code =
json_non_empty_string(payload, &["user_code", "userCode"]).ok_or_else(|| {
OAuthError::invalid_response("xAI device code response missing user_code")
})?;
let verification_uri = json_non_empty_string(
payload,
&["verification_uri", "verificationUri", "verification_url"],
)
.unwrap_or_default();
let verification_uri_complete = json_non_empty_string(
payload,
&[
"verification_uri_complete",
"verificationUriComplete",
"verification_url_complete",
],
)
.unwrap_or_else(|| verification_uri.clone());
if verification_uri.is_empty() && verification_uri_complete.is_empty() {
return Err(OAuthError::invalid_response(
"xAI device code response missing verification URI",
));
}
Ok(OAuthDeviceAuthorization {
device_code,
user_code,
verification_uri: if verification_uri.is_empty() {
verification_uri_complete.clone()
} else {
verification_uri
},
verification_uri_complete,
expires_in: json_u64(payload, &["expires_in", "expiresIn"])
.unwrap_or(DEFAULT_DEVICE_EXPIRES_IN_SECS),
interval: json_u64(payload, &["interval"]).unwrap_or(DEFAULT_DEVICE_POLL_INTERVAL_SECS),
})
}
fn raw_credential_string(raw: Option<&Value>, keys: &[&str]) -> Option<String> {
let object = raw?.as_object()?;
keys.iter().find_map(|key| {
object
.get(*key)
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
}
fn response_json(response: &crate::network::OAuthHttpResponse) -> Option<Value> {
response
.json_body
.clone()
.or_else(|| serde_json::from_str::<Value>(&response.body_text).ok())
}
fn oauth_error_code(payload: &Value) -> Option<String> {
payload
.get("error")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
fn json_non_empty_string(payload: &Value, keys: &[&str]) -> Option<String> {
keys.iter().find_map(|key| {
payload
.get(*key)
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
}
fn json_u64(payload: &Value, keys: &[&str]) -> Option<u64> {
keys.iter().find_map(|key| match payload.get(*key)? {
Value::Number(number) => number.as_u64(),
Value::String(string) => string.trim().parse::<u64>().ok(),
_ => None,
})
}
fn form_headers() -> BTreeMap<String, String> {
BTreeMap::from([
(
"content-type".to_string(),
"application/x-www-form-urlencoded".to_string(),
),
("accept".to_string(), "application/json".to_string()),
])
}
fn decode_jwt_claims(token: &str) -> Option<Map<String, Value>> {
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
const MAX_UNVERIFIED_JWT_CLAIMS_BYTES: usize = 64 * 1024;
let payload = token.split('.').nth(1)?;
let max_encoded_len = MAX_UNVERIFIED_JWT_CLAIMS_BYTES
.saturating_add(2)
.checked_div(3)
.unwrap_or(usize::MAX)
.saturating_mul(4);
if payload.len() > max_encoded_len {
return None;
}
let bytes = URL_SAFE_NO_PAD.decode(payload.as_bytes()).ok()?;
if bytes.len() > MAX_UNVERIFIED_JWT_CLAIMS_BYTES {
return None;
}
serde_json::from_slice::<Value>(&bytes)
.ok()?
.as_object()
.cloned()
}
#[cfg(test)]
mod tests {
use super::{
XaiDevicePollOutcome, XaiProviderOAuthAdapter, XAI_CLIENT_ID, XAI_DEVICE_CODE_GRANT_TYPE,
XAI_OAUTH_SCOPES, XAI_PROVIDER_TYPE,
};
use crate::network::{OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse};
use crate::provider::{
ProviderOAuthAccount, ProviderOAuthAdapter, ProviderOAuthImportInput,
ProviderOAuthTransportContext,
};
use async_trait::async_trait;
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use serde_json::{json, Value};
use std::collections::BTreeMap;
use std::sync::{Arc, Mutex};
#[derive(Clone)]
struct ScriptedExecutor {
seen_request: Arc<Mutex<Option<OAuthHttpRequest>>>,
status_code: u16,
payload: Value,
}
#[async_trait]
impl OAuthHttpExecutor for ScriptedExecutor {
async fn execute(
&self,
request: OAuthHttpRequest,
) -> Result<OAuthHttpResponse, crate::core::OAuthError> {
*self.seen_request.lock().expect("mutex should lock") = Some(request);
Ok(OAuthHttpResponse {
status_code: self.status_code,
body_text: self.payload.to_string(),
json_body: Some(self.payload.clone()),
})
}
}
fn transport_context() -> ProviderOAuthTransportContext {
ProviderOAuthTransportContext {
provider_id: "provider-xai".to_string(),
provider_type: XAI_PROVIDER_TYPE.to_string(),
endpoint_id: None,
key_id: None,
auth_type: Some("oauth".to_string()),
decrypted_api_key: None,
decrypted_auth_config: None,
provider_config: None,
endpoint_config: None,
key_config: None,
network: crate::network::OAuthNetworkContext::provider_operation(None),
}
}
fn encoded_jwt(claims: &Value) -> String {
format!(
"header.{}.signature",
URL_SAFE_NO_PAD.encode(serde_json::to_vec(claims).expect("claims should encode"))
)
}
#[tokio::test]
async fn imports_api_key_as_official_api_credential() {
let adapter = XaiProviderOAuthAdapter::default();
let executor = ScriptedExecutor {
seen_request: Arc::new(Mutex::new(None)),
status_code: 200,
payload: json!({}),
};
let result = adapter
.import_credentials(
&executor,
&transport_context(),
ProviderOAuthImportInput {
provider_type: XAI_PROVIDER_TYPE.to_string(),
name: Some("work".to_string()),
refresh_token: None,
raw_credentials: Some(json!({"api_key": "xai-key-123"})),
network: crate::network::OAuthNetworkContext::provider_operation(None),
},
)
.await
.expect("api key import should succeed");
assert_eq!(result.token_set.access_token, "xai-key-123");
assert_eq!(result.auth_config["using_api"], json!(true));
assert_eq!(result.auth_config["auth_method"], json!("api_key"));
assert!(executor.seen_request.lock().expect("lock").is_none());
}
#[tokio::test]
async fn device_poll_treats_authorization_pending_as_pending() {
let adapter = XaiProviderOAuthAdapter::default();
let seen = Arc::new(Mutex::new(None));
let executor = ScriptedExecutor {
seen_request: Arc::clone(&seen),
status_code: 400,
payload: json!({"error": "authorization_pending"}),
};
let outcome = adapter
.poll_device_token(&executor, &transport_context(), "device-code")
.await
.expect("pending should not be fatal");
assert_eq!(outcome, XaiDevicePollOutcome::Pending);
let request = seen.lock().expect("lock").clone().expect("request");
let body = request.body_bytes.expect("body");
let fields = url::form_urlencoded::parse(&body)
.into_owned()
.collect::<BTreeMap<_, _>>();
assert_eq!(fields["grant_type"], XAI_DEVICE_CODE_GRANT_TYPE);
assert_eq!(fields["device_code"], "device-code");
assert_eq!(fields["client_id"], XAI_CLIENT_ID);
}
#[tokio::test]
async fn refresh_posts_client_id_and_refresh_token_without_scope() {
let adapter = XaiProviderOAuthAdapter::default();
let seen = Arc::new(Mutex::new(None));
let id_token = encoded_jwt(&json!({"email": "[email protected]", "sub": "subject-1"}));
let executor = ScriptedExecutor {
seen_request: Arc::clone(&seen),
status_code: 200,
payload: json!({
"access_token": "new-access",
"refresh_token": "new-refresh",
"id_token": id_token,
"expires_in": 3600
}),
};
let account = ProviderOAuthAccount {
provider_type: XAI_PROVIDER_TYPE.to_string(),
access_token: "old-access".to_string(),
auth_config: json!({
"provider_type": XAI_PROVIDER_TYPE,
"refresh_token": "old-refresh",
"using_api": false,
}),
expires_at_unix_secs: None,
identity: BTreeMap::new(),
};
let result = adapter
.refresh(&executor, &transport_context(), &account)
.await
.expect("refresh should succeed");
assert_eq!(result.token_set.access_token, "new-access");
assert_eq!(result.auth_config["using_api"], json!(false));
assert_eq!(result.auth_config["email"], json!("[email protected]"));
assert_eq!(result.auth_config["sub"], json!("subject-1"));
let request = seen.lock().expect("lock").clone().expect("request");
let body = request.body_bytes.expect("body");
let fields = url::form_urlencoded::parse(&body)
.into_owned()
.collect::<BTreeMap<_, _>>();
assert_eq!(fields["grant_type"], "refresh_token");
assert_eq!(fields["client_id"], XAI_CLIENT_ID);
assert_eq!(fields["refresh_token"], "old-refresh");
assert!(!fields.contains_key("scope"));
assert!(XAI_OAUTH_SCOPES.join(" ").contains("grok-cli:access"));
}
#[tokio::test]
async fn start_device_flow_posts_client_id_and_scope() {
let adapter = XaiProviderOAuthAdapter::default();
let seen = Arc::new(Mutex::new(None));
let executor = ScriptedExecutor {
seen_request: Arc::clone(&seen),
status_code: 200,
payload: json!({
"device_code": "dc-1",
"user_code": "ABCD-EFGH",
"verification_uri": "https://auth.x.ai/device",
"verification_uri_complete": "https://auth.x.ai/device?user_code=ABCD-EFGH",
"expires_in": 600,
"interval": 5
}),
};
let authorization = adapter
.start_device_flow(&executor, &transport_context())
.await
.expect("device start should succeed");
assert_eq!(authorization.user_code, "ABCD-EFGH");
assert_eq!(authorization.device_code, "dc-1");
let request = seen.lock().expect("lock").clone().expect("request");
let body = request.body_bytes.expect("body");
let fields = url::form_urlencoded::parse(&body)
.into_owned()
.collect::<BTreeMap<_, _>>();
assert_eq!(fields["client_id"], XAI_CLIENT_ID);
assert_eq!(fields["scope"], XAI_OAUTH_SCOPES.join(" "));
}
}
+4 -2
View File
@@ -21,7 +21,7 @@ impl ProviderOAuthService {
use super::providers::{ use super::providers::{
AntigravityProviderOAuthAdapter, ClaudeCodeProviderOAuthAdapter, AntigravityProviderOAuthAdapter, ClaudeCodeProviderOAuthAdapter,
CodexProviderOAuthAdapter, GenericProviderOAuthAdapter, KiroProviderOAuthAdapter, CodexProviderOAuthAdapter, GenericProviderOAuthAdapter, KiroProviderOAuthAdapter,
WindsurfProviderOAuthAdapter, WindsurfProviderOAuthAdapter, XaiProviderOAuthAdapter,
}; };
let mut service = Self::new() let mut service = Self::new()
@@ -29,7 +29,8 @@ impl ProviderOAuthService {
.with_adapter(Arc::new(ClaudeCodeProviderOAuthAdapter::default())) .with_adapter(Arc::new(ClaudeCodeProviderOAuthAdapter::default()))
.with_adapter(Arc::new(CodexProviderOAuthAdapter::default())) .with_adapter(Arc::new(CodexProviderOAuthAdapter::default()))
.with_adapter(Arc::new(AntigravityProviderOAuthAdapter::default())) .with_adapter(Arc::new(AntigravityProviderOAuthAdapter::default()))
.with_adapter(Arc::new(WindsurfProviderOAuthAdapter)); .with_adapter(Arc::new(WindsurfProviderOAuthAdapter))
.with_adapter(Arc::new(XaiProviderOAuthAdapter::default()));
for provider_type in ["chatgpt_web", "gemini_cli"] { for provider_type in ["chatgpt_web", "gemini_cli"] {
if let Some(adapter) = GenericProviderOAuthAdapter::for_provider_type(provider_type) { if let Some(adapter) = GenericProviderOAuthAdapter::for_provider_type(provider_type) {
service = service.with_adapter(Arc::new(adapter)); service = service.with_adapter(Arc::new(adapter));
@@ -144,6 +145,7 @@ mod tests {
"antigravity", "antigravity",
"kiro", "kiro",
"windsurf", "windsurf",
"xai",
] { ] {
assert!( assert!(
service.adapter(provider_type).is_ok(), service.adapter(provider_type).is_ok(),
+17 -12
View File
@@ -22,18 +22,20 @@ pub use providers::{
build_windsurf_pool_model_configs_request, build_windsurf_pool_model_configs_request,
build_windsurf_pool_model_configs_request_with_base_url, build_windsurf_pool_quota_request, build_windsurf_pool_model_configs_request_with_base_url, build_windsurf_pool_quota_request,
build_windsurf_pool_quota_request_with_base_url, build_windsurf_pool_rate_limit_request, build_windsurf_pool_quota_request_with_base_url, build_windsurf_pool_rate_limit_request,
build_windsurf_pool_rate_limit_request_with_base_url, enrich_chatgpt_web_quota_metadata, build_windsurf_pool_rate_limit_request_with_base_url, build_xai_pool_billing_request,
grok_mode_id_for_model, grok_pool_tier_from_quota_bucket, grok_quota_window_key_for_model, build_xai_pool_user_request, enrich_chatgpt_web_quota_metadata, grok_mode_id_for_model,
grok_pool_tier_from_quota_bucket, grok_quota_window_key_for_model,
grok_supported_quota_windows_for_tier, normalize_chatgpt_web_image_quota_limit, grok_supported_quota_windows_for_tier, normalize_chatgpt_web_image_quota_limit,
AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, CodexProviderPoolAdapter, AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, CodexProviderPoolAdapter,
DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, GrokProviderPoolAdapter, DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, GrokProviderPoolAdapter,
KiroPoolQuotaAuthInput, KiroProviderPoolAdapter, UnsupportedQuotaProviderPoolAdapter, KiroPoolQuotaAuthInput, KiroProviderPoolAdapter, UnsupportedQuotaProviderPoolAdapter,
ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH, ANTIGRAVITY_RETRIEVE_USER_QUOTA_SUMMARY_PATH, XaiProviderPoolAdapter, ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH,
CHATGPT_WEB_CONVERSATION_INIT_PATH, CHATGPT_WEB_DEFAULT_BASE_URL, ANTIGRAVITY_RETRIEVE_USER_QUOTA_SUMMARY_PATH, CHATGPT_WEB_CONVERSATION_INIT_PATH,
CODEX_WHAM_RESET_CREDITS_CONSUME_URL, CODEX_WHAM_RESET_CREDITS_URL, CODEX_WHAM_USAGE_URL, CHATGPT_WEB_DEFAULT_BASE_URL, CODEX_WHAM_RESET_CREDITS_CONSUME_URL,
GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH, GEMINI_CLI_USER_AGENT, KIRO_USAGE_LIMITS_PATH, CODEX_WHAM_RESET_CREDITS_URL, CODEX_WHAM_USAGE_URL, GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH,
KIRO_USAGE_SDK_VERSION, WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_PATH, GEMINI_CLI_USER_AGENT, KIRO_USAGE_LIMITS_PATH, KIRO_USAGE_SDK_VERSION,
WINDSURF_USER_STATUS_PATH, WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_PATH, WINDSURF_USER_STATUS_PATH,
XAI_BILLING_PATH, XAI_USER_PATH,
}; };
pub use quota::{ pub use quota::{
provider_pool_key_account_quota_exhausted, provider_pool_key_model_quota_exhausted, provider_pool_key_account_quota_exhausted, provider_pool_key_model_quota_exhausted,
@@ -81,7 +83,8 @@ mod tests {
"grok", "grok",
"kiro", "kiro",
"vertex_ai", "vertex_ai",
"windsurf" "windsurf",
"xai"
] ]
); );
assert!(service assert!(service
@@ -104,7 +107,8 @@ mod tests {
"gemini_cli", "gemini_cli",
"grok", "grok",
"kiro", "kiro",
"windsurf" "windsurf",
"xai"
] ]
); );
assert!(service.supports_quota_refresh("codex")); assert!(service.supports_quota_refresh("codex"));
@@ -112,6 +116,7 @@ mod tests {
assert!(service.supports_quota_refresh("grok")); assert!(service.supports_quota_refresh("grok"));
assert!(service.supports_quota_refresh("gemini_cli")); assert!(service.supports_quota_refresh("gemini_cli"));
assert!(service.supports_quota_refresh("windsurf")); assert!(service.supports_quota_refresh("windsurf"));
assert!(service.supports_quota_refresh("xai"));
assert_eq!( assert_eq!(
service.quota_refresh_unsupported_message("claude_code"), service.quota_refresh_unsupported_message("claude_code"),
"Claude Code 暂不支持自动刷新额度:上游没有稳定可用的账号额度查询接口" "Claude Code 暂不支持自动刷新额度:上游没有稳定可用的账号额度查询接口"
@@ -642,11 +647,11 @@ mod tests {
assert_eq!( assert_eq!(
free_first["providers"], free_first["providers"],
json!(["codex", "grok", "kiro", "windsurf"]) json!(["codex", "grok", "kiro", "windsurf", "xai"])
); );
assert_eq!( assert_eq!(
recent_refresh["providers"], recent_refresh["providers"],
json!(["codex", "grok", "kiro", "windsurf"]) json!(["codex", "grok", "kiro", "windsurf", "xai"])
); );
assert_eq!(free_first["default_enabled"], json!(false)); assert_eq!(free_first["default_enabled"], json!(false));
assert_eq!(recent_refresh["default_enabled"], json!(false)); assert_eq!(recent_refresh["default_enabled"], json!(false));
@@ -7,6 +7,7 @@ pub mod grok;
pub mod kiro; pub mod kiro;
pub mod unsupported; pub mod unsupported;
pub mod windsurf; pub mod windsurf;
pub mod xai;
pub use antigravity::AntigravityProviderPoolAdapter; pub use antigravity::AntigravityProviderPoolAdapter;
pub use antigravity::{ pub use antigravity::{
@@ -51,3 +52,7 @@ pub use windsurf::{
WINDSURF_DEFAULT_BASE_URL, WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_PATH, WINDSURF_DEFAULT_BASE_URL, WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_PATH,
WINDSURF_USER_STATUS_PATH, WINDSURF_USER_STATUS_PATH,
}; };
pub use xai::{
build_xai_pool_billing_request, build_xai_pool_user_request, XaiProviderPoolAdapter,
XAI_BILLING_PATH, XAI_USER_PATH,
};
@@ -0,0 +1,255 @@
use std::collections::BTreeMap;
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogEndpoint;
use aether_provider_transport::xai::{
insert_cli_identity_headers, XAI_CHAT_PROXY_BASE_URL, XAI_PROVIDER_TYPE,
};
use serde_json::{Map, Value};
use crate::capability::ProviderPoolCapabilities;
use crate::provider::{
provider_pool_endpoint_format_matches, provider_pool_matching_endpoint, ProviderPoolAdapter,
ProviderPoolMemberInput,
};
use crate::quota::{
provider_pool_current_unix_secs, provider_pool_json_bool, provider_pool_json_f64,
provider_pool_metadata_bucket, provider_pool_model_quota_exhausted,
provider_pool_quota_snapshot_exhausted_decision, provider_pool_reset_deadline_elapsed,
provider_pool_timestamp_unix_secs,
};
use crate::quota_refresh::ProviderPoolQuotaRequestSpec;
pub const XAI_USER_PATH: &str = "/user";
pub const XAI_BILLING_PATH: &str = "/billing?format=credits";
#[derive(Debug, Clone, Default)]
pub struct XaiProviderPoolAdapter;
impl ProviderPoolAdapter for XaiProviderPoolAdapter {
fn provider_type(&self) -> &'static str {
XAI_PROVIDER_TYPE
}
fn capabilities(&self) -> ProviderPoolCapabilities {
ProviderPoolCapabilities {
plan_tier: true,
quota_reset: true,
quota_refresh: true,
}
}
fn quota_exhausted(&self, input: &ProviderPoolMemberInput<'_>) -> bool {
if let Some(exhausted) = input.provider_model_name.and_then(|model| {
provider_pool_model_quota_exhausted(input.key, input.provider_type, model)
}) {
return exhausted;
}
if let Some(exhausted) =
provider_pool_quota_snapshot_exhausted_decision(input.key, input.provider_type)
{
return exhausted;
}
provider_pool_metadata_bucket(input.key.upstream_metadata.as_ref(), input.provider_type)
.is_some_and(quota_exhausted_from_bucket)
}
fn quota_refresh_endpoint(
&self,
endpoints: &[StoredProviderCatalogEndpoint],
include_inactive: bool,
) -> Option<StoredProviderCatalogEndpoint> {
provider_pool_matching_endpoint(endpoints, include_inactive, |endpoint| {
provider_pool_endpoint_format_matches(endpoint, "openai:responses")
})
.or_else(|| provider_pool_matching_endpoint(endpoints, include_inactive, |_| true))
}
fn quota_refresh_missing_endpoint_message(&self) -> String {
"找不到有效的 openai:responses 端点".to_string()
}
}
pub fn build_xai_pool_user_request(
key_id: &str,
authorization: (String, String),
) -> ProviderPoolQuotaRequestSpec {
build_xai_pool_request(
format!("xai-user:{key_id}"),
"xai:user",
"user",
XAI_USER_PATH,
authorization,
None,
)
}
pub fn build_xai_pool_billing_request(
key_id: &str,
authorization: (String, String),
user_id: Option<&str>,
) -> ProviderPoolQuotaRequestSpec {
build_xai_pool_request(
format!("xai-billing:{key_id}"),
"xai:billing",
"billing",
XAI_BILLING_PATH,
authorization,
user_id,
)
}
fn build_xai_pool_request(
request_id: String,
provider_api_format: &str,
model_name: &str,
path: &str,
authorization: (String, String),
user_id: Option<&str>,
) -> ProviderPoolQuotaRequestSpec {
let mut headers = BTreeMap::from([
(authorization.0, authorization.1),
("accept".to_string(), "application/json".to_string()),
]);
insert_cli_identity_headers(&mut headers);
if let Some(user_id) = user_id.map(str::trim).filter(|value| !value.is_empty()) {
headers.insert("x-userid".to_string(), user_id.to_string());
}
ProviderPoolQuotaRequestSpec {
request_id,
provider_name: XAI_PROVIDER_TYPE.to_string(),
quota_kind: XAI_PROVIDER_TYPE.to_string(),
method: "GET".to_string(),
url: format!("{}{path}", XAI_CHAT_PROXY_BASE_URL.trim_end_matches('/')),
headers,
content_type: None,
json_body: None,
client_api_format: "openai:responses".to_string(),
provider_api_format: provider_api_format.to_string(),
model_name: Some(model_name.to_string()),
}
}
pub(crate) fn quota_exhausted_from_bucket(bucket: &Map<String, Value>) -> bool {
if provider_pool_current_unix_secs().is_some_and(|now| {
provider_pool_reset_deadline_elapsed(
bucket,
provider_pool_timestamp_unix_secs(bucket.get("updated_at")),
now,
)
}) {
return false;
}
let usage_exhausted = provider_pool_json_f64(bucket.get("remaining"))
.is_some_and(|value| value <= 0.0)
|| provider_pool_json_f64(bucket.get("usage_percentage"))
.is_some_and(|value| value >= 100.0 - 1e-6)
|| match (
provider_pool_json_f64(bucket.get("usage_limit")),
provider_pool_json_f64(bucket.get("current_usage")),
) {
(Some(limit), Some(current)) if limit > 0.0 => current >= limit,
_ => false,
};
if !usage_exhausted {
return false;
}
let prepaid_available =
provider_pool_json_f64(bucket.get("prepaid_balance")).is_some_and(|value| value > 0.0);
if prepaid_available {
return false;
}
let on_demand_enabled = provider_pool_json_bool(bucket.get("on_demand_enabled")) != Some(false);
let on_demand_cap = provider_pool_json_f64(bucket.get("on_demand_cap")).unwrap_or(0.0);
let on_demand_used = provider_pool_json_f64(bucket.get("on_demand_used")).unwrap_or(0.0);
if on_demand_enabled && on_demand_cap > 0.0 && on_demand_used < on_demand_cap {
return false;
}
true
}
#[cfg(test)]
mod tests {
use super::{
build_xai_pool_billing_request, build_xai_pool_user_request, quota_exhausted_from_bucket,
};
use aether_provider_transport::xai::{
XAI_CHAT_PROXY_BASE_URL, XAI_CLIENT_IDENTIFIER_VALUE, XAI_TOKEN_AUTH_VALUE,
};
use serde_json::{json, Map};
fn bucket(value: serde_json::Value) -> Map<String, serde_json::Value> {
value.as_object().cloned().expect("bucket should be object")
}
#[test]
fn user_and_billing_requests_pin_cli_chat_proxy_and_identity_headers() {
let authorization = ("authorization".to_string(), "Bearer xai-access".to_string());
let user = build_xai_pool_user_request("key-1", authorization.clone());
let billing = build_xai_pool_billing_request("key-1", authorization, Some("user-42"));
assert_eq!(
user.url,
format!("{}/user", XAI_CHAT_PROXY_BASE_URL.trim_end_matches('/'))
);
assert_eq!(
billing.url,
format!(
"{}/billing?format=credits",
XAI_CHAT_PROXY_BASE_URL.trim_end_matches('/')
)
);
assert_eq!(
user.headers.get("x-xai-token-auth").map(String::as_str),
Some(XAI_TOKEN_AUTH_VALUE)
);
assert_eq!(
user.headers
.get("x-grok-client-identifier")
.map(String::as_str),
Some(XAI_CLIENT_IDENTIFIER_VALUE)
);
assert!(!user.headers.contains_key("x-userid"));
assert_eq!(
billing.headers.get("x-userid").map(String::as_str),
Some("user-42")
);
assert_eq!(
billing.headers.get("authorization").map(String::as_str),
Some("Bearer xai-access")
);
}
#[test]
fn percent_exhausted_without_prepaid_or_on_demand_is_exhausted() {
assert!(quota_exhausted_from_bucket(&bucket(json!({
"usage_percentage": 100.0,
"prepaid_balance": 0.0,
"on_demand_cap": 0.0,
"on_demand_used": 0.0
}))));
}
#[test]
fn unified_billing_zero_on_demand_cap_is_not_exhausted_when_percent_remains() {
assert!(!quota_exhausted_from_bucket(&bucket(json!({
"usage_percentage": 46.0,
"prepaid_balance": 0.0,
"on_demand_cap": 0.0,
"on_demand_used": 0.0
}))));
}
#[test]
fn prepaid_balance_keeps_account_available_after_weekly_pool_hits_100() {
assert!(!quota_exhausted_from_bucket(&bucket(json!({
"usage_percentage": 100.0,
"prepaid_balance": 12.5,
"on_demand_cap": 0.0
}))));
}
}
+3 -2
View File
@@ -13,8 +13,8 @@ use crate::provider::{ProviderPoolAdapter, ProviderPoolMemberInput};
use crate::providers::{ use crate::providers::{
AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, CodexProviderPoolAdapter, AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, CodexProviderPoolAdapter,
DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, GrokProviderPoolAdapter, DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, GrokProviderPoolAdapter,
KiroProviderPoolAdapter, WindsurfProviderPoolAdapter, CLAUDE_CODE_PROVIDER_POOL_ADAPTER, KiroProviderPoolAdapter, WindsurfProviderPoolAdapter, XaiProviderPoolAdapter,
VERTEX_AI_PROVIDER_POOL_ADAPTER, CLAUDE_CODE_PROVIDER_POOL_ADAPTER, VERTEX_AI_PROVIDER_POOL_ADAPTER,
}; };
#[derive(Clone)] #[derive(Clone)]
@@ -55,6 +55,7 @@ impl ProviderPoolService {
.with_adapter(Arc::new(KiroProviderPoolAdapter)) .with_adapter(Arc::new(KiroProviderPoolAdapter))
.with_adapter(Arc::new(ChatGptWebProviderPoolAdapter)) .with_adapter(Arc::new(ChatGptWebProviderPoolAdapter))
.with_adapter(Arc::new(WindsurfProviderPoolAdapter)) .with_adapter(Arc::new(WindsurfProviderPoolAdapter))
.with_adapter(Arc::new(XaiProviderPoolAdapter))
.with_adapter(Arc::new(VERTEX_AI_PROVIDER_POOL_ADAPTER)) .with_adapter(Arc::new(VERTEX_AI_PROVIDER_POOL_ADAPTER))
} }
@@ -754,6 +754,90 @@ mod tests {
)); ));
} }
#[test]
fn xai_responses_transport_converts_standard_client_protocols() {
let transport = transport_snapshot("xai", "openai:responses", "oauth", true, None);
for client_api_format in ["openai:chat", "claude:messages", "gemini:generate_content"] {
assert!(
request_pair_allowed_for_transport(
&transport,
client_api_format,
"openai:responses"
),
"{client_api_format} should convert onto xAI Responses"
);
assert_eq!(
candidate_transport_pair_skip_reason(&transport, client_api_format),
None
);
}
assert!(request_conversion_transport_supported(
&transport,
RequestConversionKind::ToOpenAiResponses
));
assert!(!request_pair_allowed_for_transport(
&transport,
"openai:image",
"openai:responses"
));
assert!(!request_pair_allowed_for_transport(
&transport,
"openai:video",
"openai:responses"
));
for isolated in ["openai:responses:compact", "openai:image", "openai:video"] {
assert!(
!request_pair_allowed_for_transport(&transport, isolated, "openai:responses"),
"{isolated} must not convert onto xAI Responses"
);
}
}
#[test]
fn xai_compact_and_media_endpoints_are_same_format_only() {
let compact = transport_snapshot("xai", "openai:responses:compact", "oauth", true, None);
assert!(request_pair_allowed_for_transport(
&compact,
"openai:responses:compact",
"openai:responses:compact"
));
for client_api_format in [
"openai:chat",
"openai:responses",
"claude:messages",
"gemini:generate_content",
] {
assert!(
!request_pair_allowed_for_transport(
&compact,
client_api_format,
"openai:responses:compact"
),
"{client_api_format} must not convert onto xAI compact"
);
}
for api_format in ["openai:image", "openai:video"] {
let transport = transport_snapshot("xai", api_format, "oauth", true, None);
assert!(
request_pair_allowed_for_transport(&transport, api_format, api_format),
"{api_format} same-format transport should be allowed"
);
for client_api_format in [
"openai:chat",
"openai:responses",
"claude:messages",
"gemini:generate_content",
] {
assert!(
!request_pair_allowed_for_transport(&transport, client_api_format, api_format),
"{client_api_format} must not convert onto {api_format}"
);
}
}
}
#[test] #[test]
fn windsurf_openai_chat_anchor_supports_cross_format_conversion_via_cascade() { fn windsurf_openai_chat_anchor_supports_cross_format_conversion_via_cascade() {
let mut transport = transport_snapshot("windsurf", "openai:chat", "oauth", true, None); let mut transport = transport_snapshot("windsurf", "openai:chat", "oauth", true, None);
@@ -30,6 +30,7 @@ pub mod url;
pub mod vertex; pub mod vertex;
mod video; mod video;
pub mod windsurf; pub mod windsurf;
pub mod xai;
pub use aether_oauth as oauth; pub use aether_oauth as oauth;
pub use agent_identity::{ pub use agent_identity::{
@@ -195,3 +196,10 @@ pub use windsurf::{
local_windsurf_request_transport_unsupported_reason_with_network, GET_CHAT_MESSAGE_PATH, local_windsurf_request_transport_unsupported_reason_with_network, GET_CHAT_MESSAGE_PATH,
WINDSURF_ENVELOPE_NAME, WINDSURF_ENVELOPE_NAME,
}; };
pub use xai::{
extract_xai_user_id_from_auth_config, extract_xai_user_id_from_value,
insert_cli_identity_headers, insert_cli_identity_headers_if_needed, is_xai_provider_transport,
resolved_xai_request_base_url, resolved_xai_upstream_base_url,
should_attach_cli_identity_headers, xai_auth_uses_api, xai_uses_official_api, XAI_API_BASE_URL,
XAI_CHAT_PROXY_BASE_URL, XAI_PROVIDER_TYPE,
};
@@ -84,6 +84,11 @@ fn is_dedicated_openai_image_provider(transport: &GatewayProviderTransportSnapsh
.trim() .trim()
.eq_ignore_ascii_case("codex") .eq_ignore_ascii_case("codex")
|| is_grok_provider_transport(transport) || is_grok_provider_transport(transport)
|| transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("xai")
} }
pub fn resolve_openai_image_auth( pub fn resolve_openai_image_auth(
@@ -92,7 +97,10 @@ pub fn resolve_openai_image_auth(
if is_grok_provider_transport(transport) { if is_grok_provider_transport(transport) {
return resolve_grok_session_auth(transport); return resolve_grok_session_auth(transport);
} }
resolve_local_openai_bearer_auth(transport) resolve_local_openai_bearer_auth(transport).or_else(|| {
crate::generic_oauth::resolve_local_generic_oauth_transport_authorization(transport)
.map(|value| ("authorization".to_string(), value))
})
} }
pub fn build_openai_image_upstream_url( pub fn build_openai_image_upstream_url(
@@ -100,7 +108,11 @@ pub fn build_openai_image_upstream_url(
request_path: Option<&str>, request_path: Option<&str>,
request_query: Option<&str>, request_query: Option<&str>,
) -> String { ) -> String {
build_openai_image_url(&transport.endpoint.base_url, request_path, request_query) build_openai_image_url(
&crate::xai::resolved_xai_request_base_url(transport, "openai:image"),
request_path,
request_query,
)
} }
pub fn build_openai_image_headers( pub fn build_openai_image_headers(
@@ -113,6 +125,11 @@ pub fn build_openai_image_headers(
&BTreeMap::new(), &BTreeMap::new(),
); );
provider_request_headers.insert("content-type".to_string(), "application/json".to_string()); provider_request_headers.insert("content-type".to_string(), "application/json".to_string());
crate::xai::insert_cli_identity_headers_if_needed(
input.transport,
"openai:image",
&mut provider_request_headers,
);
if let Some(accept) = input.accept { if let Some(accept) = input.accept {
provider_request_headers.insert("accept".to_string(), accept.to_string()); provider_request_headers.insert("accept".to_string(), accept.to_string());
} else { } else {
@@ -280,6 +297,48 @@ mod tests {
); );
} }
#[test]
fn xai_oauth_image_uses_cli_proxy() {
let mut transport = sample_transport();
transport.provider.provider_type = "xai".to_string();
transport.endpoint.base_url = "https://cli-chat-proxy.grok.com/v1".to_string();
transport.key.auth_type = "oauth".to_string();
transport.key.decrypted_auth_config =
Some(r#"{"refresh_token":"rt","using_api":false}"#.to_string());
assert_eq!(
openai_image_transport_unsupported_reason(&transport, "openai:image"),
None
);
assert_eq!(
build_openai_image_upstream_url(&transport, Some("/v1/images/generations"), None),
"https://cli-chat-proxy.grok.com/v1/images/generations"
);
assert_eq!(
build_openai_image_upstream_url(&transport, Some("/v1/images/edits"), None),
"https://cli-chat-proxy.grok.com/v1/images/edits"
);
let headers = build_openai_image_headers(ProviderOpenAiImageHeadersInput {
transport: &transport,
headers: &HeaderMap::new(),
auth_header: "authorization",
auth_value: "Bearer test-token",
accept: None,
header_rules: None,
provider_request_body: &json!({"prompt": "A cat"}),
original_request_body: &json!({"prompt": "A cat"}),
})
.unwrap();
assert_eq!(
headers.get("x-xai-token-auth").map(String::as_str),
Some("xai-grok-cli")
);
assert_eq!(
headers.get("authorization").map(String::as_str),
Some("Bearer test-token")
);
}
#[test] #[test]
fn codex_is_supported_by_dedicated_openai_image_transport_policy() { fn codex_is_supported_by_dedicated_openai_image_transport_policy() {
let mut transport = sample_transport(); let mut transport = sample_transport();
@@ -275,6 +275,17 @@ const WINDSURF_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
..STANDARD_RUNTIME_POLICY ..STANDARD_RUNTIME_POLICY
}; };
const XAI_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
fixed_provider: true,
api_format_inheritance: ProviderApiFormatInheritance::OAuthOrBearer,
enable_format_conversion_by_default: true,
oauth_is_bearer_like: true,
supports_model_fetch: false,
supports_local_openai_chat_transport: false,
supports_local_same_format_transport: true,
..STANDARD_RUNTIME_POLICY
};
const CLAUDE_CODE_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate { const CLAUDE_CODE_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
provider_type: "claude_code", provider_type: "claude_code",
version: 2, version: 2,
@@ -446,6 +457,39 @@ const WINDSURF_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTem
runtime_policy: WINDSURF_RUNTIME_POLICY, runtime_policy: WINDSURF_RUNTIME_POLICY,
}; };
const XAI_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
provider_type: "xai",
version: 2,
base_url: crate::xai::XAI_CHAT_PROXY_BASE_URL,
endpoints: &[
FixedProviderEndpointTemplate {
item_key: "openai:responses",
api_format: "openai:responses",
custom_path: None,
config_defaults: FORCE_STREAM_ENDPOINT_CONFIG_DEFAULTS,
},
FixedProviderEndpointTemplate {
item_key: "openai:responses:compact",
api_format: "openai:responses:compact",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
},
FixedProviderEndpointTemplate {
item_key: "openai:image",
api_format: "openai:image",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
},
FixedProviderEndpointTemplate {
item_key: "openai:video",
api_format: "openai:video",
custom_path: None,
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
},
],
runtime_policy: XAI_RUNTIME_POLICY,
};
pub fn provider_type_is_fixed(provider_type: &str) -> bool { pub fn provider_type_is_fixed(provider_type: &str) -> bool {
provider_runtime_policy(provider_type).fixed_provider provider_runtime_policy(provider_type).fixed_provider
} }
@@ -498,6 +542,7 @@ pub fn fixed_provider_template(provider_type: &str) -> Option<&'static FixedProv
"vertex_ai" => Some(&VERTEX_AI_FIXED_PROVIDER_TEMPLATE), "vertex_ai" => Some(&VERTEX_AI_FIXED_PROVIDER_TEMPLATE),
"antigravity" => Some(&ANTIGRAVITY_FIXED_PROVIDER_TEMPLATE), "antigravity" => Some(&ANTIGRAVITY_FIXED_PROVIDER_TEMPLATE),
"windsurf" => Some(&WINDSURF_FIXED_PROVIDER_TEMPLATE), "windsurf" => Some(&WINDSURF_FIXED_PROVIDER_TEMPLATE),
"xai" => Some(&XAI_FIXED_PROVIDER_TEMPLATE),
_ => None, _ => None,
} }
} }
@@ -613,6 +658,16 @@ pub fn provider_type_admin_oauth_template(provider_type: &str) -> Option<Provide
redirect_uri: "show-auth-token", redirect_uri: "show-auth-token",
use_pkce: false, use_pkce: false,
}), }),
"xai" => Some(ProviderOAuthTemplate {
provider_type: "xai",
display_name: "xAI",
authorize_url: aether_oauth::provider::providers::XAI_DEVICE_CODE_URL,
token_url: aether_oauth::provider::providers::XAI_TOKEN_URL,
client_id: aether_oauth::provider::providers::XAI_CLIENT_ID,
scopes: aether_oauth::provider::providers::XAI_OAUTH_SCOPES,
redirect_uri: "",
use_pkce: false,
}),
_ => None, _ => None,
} }
} }
@@ -825,6 +880,50 @@ mod tests {
assert!(ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES.contains(&"windsurf")); assert!(ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES.contains(&"windsurf"));
} }
#[test]
fn xai_fixed_provider_template_exposes_responses_media_endpoints() {
let template = fixed_provider_template("xai").expect("xai template should exist");
assert_eq!(template.provider_type, "xai");
assert_eq!(template.base_url, crate::xai::XAI_CHAT_PROXY_BASE_URL);
assert_eq!(template.version, 2);
assert_eq!(
template
.endpoints
.iter()
.map(|item| item.api_format)
.collect::<Vec<_>>(),
vec![
"openai:responses",
"openai:responses:compact",
"openai:image",
"openai:video"
]
);
let policy = provider_runtime_policy("xai");
assert!(policy.fixed_provider);
assert!(policy.enable_format_conversion_by_default);
assert!(policy.oauth_is_bearer_like);
assert!(!policy.supports_model_fetch);
assert!(policy.supports_local_same_format_transport);
assert!(!policy.supports_local_openai_chat_transport);
assert!(fixed_provider_key_inherits_api_formats(
"xai", "oauth", None
));
assert!(fixed_provider_key_inherits_api_formats(
"xai", "bearer", None
));
let template = provider_type_admin_oauth_template("xai").expect("xai oauth template");
assert_eq!(template.provider_type, "xai");
assert_eq!(template.display_name, "xAI");
assert_eq!(
template.token_url,
aether_oauth::provider::providers::XAI_TOKEN_URL
);
assert!(!ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES.contains(&"xai"));
}
#[test] #[test]
fn fixed_provider_key_inheritance_keeps_oauth_and_kiro_configured_bearer_keys_open() { fn fixed_provider_key_inheritance_keeps_oauth_and_kiro_configured_bearer_keys_open() {
assert!(fixed_provider_key_inherits_api_formats( assert!(fixed_provider_key_inherits_api_formats(
@@ -42,6 +42,11 @@ pub fn apply_transport_request_body_semantics(
{ {
sanitize_claude_code_request_body(provider_request_body); sanitize_claude_code_request_body(provider_request_body);
} }
aether_ai_formats::apply_xai_upstream_payload_edits(
provider_request_body,
transport.provider.provider_type.as_str(),
provider_api_format.as_str(),
);
if provider_api_format == "gemini:embedding" && is_vertex_transport_context(transport) { if provider_api_format == "gemini:embedding" && is_vertex_transport_context(transport) {
apply_vertex_gemini_embedding_body_semantics(provider_request_body)?; apply_vertex_gemini_embedding_body_semantics(provider_request_body)?;
} }
@@ -120,6 +120,12 @@ fn build_transport_request_url_inner(
return Some(url); return Some(url);
} }
let xai_base =
crate::xai::resolved_xai_upstream_base_url(transport, &normalized_provider_api_format);
let request_base_url = xai_base
.as_deref()
.unwrap_or(transport.endpoint.base_url.as_str());
let custom_path_template = transport let custom_path_template = transport
.endpoint .endpoint
.custom_path .custom_path
@@ -164,7 +170,7 @@ fn build_transport_request_url_inner(
path.to_string() path.to_string()
}; };
let mut url = build_passthrough_path_url( let mut url = build_passthrough_path_url(
&transport.endpoint.base_url, request_base_url,
normalized_path.as_str(), normalized_path.as_str(),
params.request_query, params.request_query,
blocked_keys, blocked_keys,
@@ -190,75 +196,68 @@ fn build_transport_request_url_inner(
let url = match normalized_provider_api_format.as_str() { let url = match normalized_provider_api_format.as_str() {
"openai:chat" => Some(build_openai_chat_url( "openai:chat" => Some(build_openai_chat_url(
&transport.endpoint.base_url, request_base_url,
params.request_query, params.request_query,
)), )),
"openai:responses" => Some(build_openai_responses_url( "openai:responses" => Some(build_openai_responses_url(
&transport.endpoint.base_url, request_base_url,
params.request_query, params.request_query,
false, false,
)), )),
"openai:responses:compact" => Some(build_openai_responses_url( "openai:responses:compact" => Some(build_openai_responses_url(
&transport.endpoint.base_url, request_base_url,
params.request_query, params.request_query,
true, true,
)), )),
"openai:search" => Some(build_openai_search_url( "openai:search" => Some(build_openai_search_url(
&transport.endpoint.base_url, request_base_url,
params.request_query, params.request_query,
)), )),
"openai:realtime" => build_passthrough_path_url( "openai:realtime" => build_passthrough_path_url(
&transport.endpoint.base_url, request_base_url,
"/v1/realtime", "/v1/realtime",
params.request_query, params.request_query,
GATEWAY_CREDENTIAL_QUERY_KEYS, GATEWAY_CREDENTIAL_QUERY_KEYS,
) )
.and_then(|url| replace_realtime_model_query(url, params.mapped_model?)), .and_then(|url| replace_realtime_model_query(url, params.mapped_model?)),
"codex:live" => build_passthrough_path_url( "codex:live" => build_passthrough_path_url(
&transport.endpoint.base_url, request_base_url,
"/live", "/live",
params.request_query, params.request_query,
GATEWAY_CREDENTIAL_QUERY_KEYS, GATEWAY_CREDENTIAL_QUERY_KEYS,
), ),
"openai:embedding" | "jina:embedding" => { "openai:embedding" | "jina:embedding" => {
build_provider_embedding_v1_url(&transport.endpoint.base_url, params.request_query) build_provider_embedding_v1_url(request_base_url, params.request_query)
}
"aliyun:multimodal_embedding" => {
build_aliyun_multimodal_embedding_url(request_base_url, params.request_query)
} }
"aliyun:multimodal_embedding" => build_aliyun_multimodal_embedding_url(
&transport.endpoint.base_url,
params.request_query,
),
"openai:rerank" | "jina:rerank" => { "openai:rerank" | "jina:rerank" => {
build_provider_rerank_v1_url(&transport.endpoint.base_url, params.request_query) build_provider_rerank_v1_url(request_base_url, params.request_query)
} }
"claude:messages" => Some(if is_claude_count_tokens { "claude:messages" => Some(if is_claude_count_tokens {
build_default_claude_count_tokens_url( build_default_claude_count_tokens_url(request_base_url, params.request_query)
&transport.endpoint.base_url,
params.request_query,
)
} else { } else {
build_claude_messages_url(&transport.endpoint.base_url, params.request_query) build_claude_messages_url(request_base_url, params.request_query)
}), }),
"gemini:generate_content" => build_gemini_content_url( "gemini:generate_content" => build_gemini_content_url(
&transport.endpoint.base_url, request_base_url,
params.mapped_model?, params.mapped_model?,
params.upstream_is_stream, params.upstream_is_stream,
params.request_query, params.request_query,
), ),
"gemini:embedding" => build_gemini_embedding_url( "gemini:embedding" => build_gemini_embedding_url(
&transport.endpoint.base_url, request_base_url,
params.mapped_model?, params.mapped_model?,
params.request_query, params.request_query,
gemini_embedding_batch, gemini_embedding_batch,
), ),
"gemini:interactions" => { "gemini:interactions" => {
build_gemini_interactions_url(&transport.endpoint.base_url, params.request_query) build_gemini_interactions_url(request_base_url, params.request_query)
}
"doubao:embedding" => {
build_passthrough_path_url(request_base_url, "/embeddings", params.request_query, &[])
} }
"doubao:embedding" => build_passthrough_path_url(
&transport.endpoint.base_url,
"/embeddings",
params.request_query,
&[],
),
_ => None, _ => None,
}?; }?;
@@ -2417,4 +2416,82 @@ mod tests {
"https://api.example.com/v1/messages?model=claude%26admin%3Dtrue%23fragment" "https://api.example.com/v1/messages?model=claude%26admin%3Dtrue%23fragment"
); );
} }
#[test]
fn xai_oauth_responses_use_cli_chat_proxy() {
let mut transport = sample_transport(
"xai",
"openai:responses",
"https://cli-chat-proxy.grok.com/v1",
None,
);
transport.key.auth_type = "oauth".to_string();
transport.key.decrypted_auth_config =
Some(r#"{"refresh_token":"rt","using_api":false}"#.to_string());
let url = build_transport_request_url(
&transport,
TransportRequestUrlParams {
provider_api_format: "openai:responses",
mapped_model: Some("grok-4"),
upstream_is_stream: true,
request_query: None,
kiro_api_region: None,
api_operation: None,
},
)
.expect("xai oauth responses URL");
assert_eq!(url, "https://cli-chat-proxy.grok.com/v1/responses");
}
#[test]
fn xai_compact_and_using_api_use_official_api() {
let mut oauth = sample_transport(
"xai",
"openai:responses:compact",
"https://cli-chat-proxy.grok.com/v1",
None,
);
oauth.key.auth_type = "oauth".to_string();
oauth.key.decrypted_auth_config =
Some(r#"{"refresh_token":"rt","using_api":false}"#.to_string());
let compact = build_transport_request_url(
&oauth,
TransportRequestUrlParams {
provider_api_format: "openai:responses:compact",
mapped_model: Some("grok-4"),
upstream_is_stream: false,
request_query: None,
kiro_api_region: None,
api_operation: None,
},
)
.expect("xai compact URL");
assert_eq!(compact, "https://api.x.ai/v1/responses/compact");
let mut api_key = sample_transport(
"xai",
"openai:responses",
"https://cli-chat-proxy.grok.com/v1",
None,
);
api_key.key.auth_type = "oauth".to_string();
api_key.key.decrypted_auth_config = Some(r#"{"using_api":true}"#.to_string());
let official = build_transport_request_url(
&api_key,
TransportRequestUrlParams {
provider_api_format: "openai:responses",
mapped_model: Some("grok-4"),
upstream_is_stream: true,
request_query: None,
kiro_api_region: None,
api_operation: None,
},
)
.expect("xai api key URL");
assert_eq!(official, "https://api.x.ai/v1/responses");
}
} }
@@ -396,6 +396,12 @@ pub fn build_standard_provider_request_headers(
force_identity_accept_encoding(&mut headers); force_identity_accept_encoding(&mut headers);
} }
crate::xai::insert_cli_identity_headers_if_needed(
input.transport,
input.provider_api_format,
&mut headers,
);
let declared_connection_headers = let declared_connection_headers =
crate::headers::declared_connection_header_names(input.headers, input.extra_headers); crate::headers::declared_connection_header_names(input.headers, input.extra_headers);
crate::headers::remove_declared_connection_headers(&mut headers, &declared_connection_headers); crate::headers::remove_declared_connection_headers(&mut headers, &declared_connection_headers);
+225 -21
View File
@@ -1,6 +1,7 @@
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::fmt; use std::fmt;
use aether_contracts::ProxySnapshot;
use aether_data_contracts::repository::video_tasks::StoredVideoTask; use aether_data_contracts::repository::video_tasks::StoredVideoTask;
use aether_video_tasks_core::{ use aether_video_tasks_core::{
LocalVideoTaskSnapshot, LocalVideoTaskTransport, LocalVideoTaskTransportBridgeInput, LocalVideoTaskSnapshot, LocalVideoTaskTransport, LocalVideoTaskTransportBridgeInput,
@@ -12,11 +13,13 @@ use super::auth::{
build_passthrough_headers_with_auth, resolve_local_gemini_auth, build_passthrough_headers_with_auth, resolve_local_gemini_auth,
resolve_local_openai_bearer_auth, resolve_local_openai_bearer_auth,
}; };
use super::network::{resolve_transport_execution_timeouts, resolve_transport_profile}; use super::network::{
resolve_transport_execution_timeouts, resolve_transport_profile,
resolve_transport_proxy_snapshot,
};
use super::policy::{ use super::policy::{
local_gemini_transport_unsupported_reason_with_network, local_gemini_transport_unsupported_reason_with_network,
local_standard_transport_unsupported_reason_with_network, supports_local_gemini_transport, local_standard_transport_unsupported_reason_with_network,
supports_local_standard_transport,
}; };
use super::rules::{ use super::rules::{
apply_local_body_rules_with_request_headers, apply_local_header_rules_with_request_headers, apply_local_body_rules_with_request_headers, apply_local_header_rules_with_request_headers,
@@ -32,6 +35,7 @@ pub enum ProviderVideoCreateFamily {
#[derive(Clone, Copy)] #[derive(Clone, Copy)]
pub struct ProviderVideoCreateHeadersInput<'a> { pub struct ProviderVideoCreateHeadersInput<'a> {
pub transport: &'a GatewayProviderTransportSnapshot,
pub headers: &'a http::HeaderMap, pub headers: &'a http::HeaderMap,
pub auth_header: &'a str, pub auth_header: &'a str,
pub auth_value: &'a str, pub auth_value: &'a str,
@@ -79,6 +83,13 @@ pub trait VideoTaskTransportSnapshotLookup: Send + Sync {
endpoint_id: &str, endpoint_id: &str,
key_id: &str, key_id: &str,
) -> Result<Option<GatewayProviderTransportSnapshot>, String>; ) -> Result<Option<GatewayProviderTransportSnapshot>, String>;
async fn resolve_video_task_proxy(
&self,
transport: &GatewayProviderTransportSnapshot,
) -> Option<ProxySnapshot> {
resolve_transport_proxy_snapshot(transport)
}
} }
pub fn resolve_local_video_task_transport( pub fn resolve_local_video_task_transport(
@@ -89,13 +100,17 @@ pub fn resolve_local_video_task_transport(
let api_format = api_format.trim(); let api_format = api_format.trim();
let (auth_header, auth_value) = match api_format { let (auth_header, auth_value) = match api_format {
"openai:video" => { "openai:video" => {
if !supports_local_standard_transport(transport, api_format) { if local_standard_transport_unsupported_reason_with_network(transport, api_format)
.is_some()
{
return None; return None;
} }
resolve_local_openai_bearer_auth(transport)? resolve_openai_compatible_video_auth(transport)?
} }
"gemini:video" => { "gemini:video" => {
if !supports_local_gemini_transport(transport, api_format) { if local_gemini_transport_unsupported_reason_with_network(transport, api_format)
.is_some()
{
return None; return None;
} }
resolve_local_gemini_auth(transport)? resolve_local_gemini_auth(transport)?
@@ -103,9 +118,9 @@ pub fn resolve_local_video_task_transport(
_ => return None, _ => return None,
}; };
Some(LocalVideoTaskTransport::from_bridge_input( let mut resolved =
LocalVideoTaskTransportBridgeInput { LocalVideoTaskTransport::from_bridge_input(LocalVideoTaskTransportBridgeInput {
upstream_base_url: transport.endpoint.base_url.clone(), upstream_base_url: crate::xai::resolved_xai_request_base_url(transport, api_format),
provider_name: Some(transport.provider.name.clone()), provider_name: Some(transport.provider.name.clone()),
provider_id: transport.provider.id.clone(), provider_id: transport.provider.id.clone(),
endpoint_id: transport.endpoint.id.clone(), endpoint_id: transport.endpoint.id.clone(),
@@ -114,11 +129,12 @@ pub fn resolve_local_video_task_transport(
auth_value, auth_value,
content_type: Some("application/json".to_string()), content_type: Some("application/json".to_string()),
model_name, model_name,
proxy: None, proxy: resolve_transport_proxy_snapshot(transport),
transport_profile: resolve_transport_profile(transport), transport_profile: resolve_transport_profile(transport),
timeouts: resolve_transport_execution_timeouts(transport), timeouts: resolve_transport_execution_timeouts(transport),
}, });
)) crate::xai::insert_cli_identity_headers_if_needed(transport, api_format, &mut resolved.headers);
Some(resolved)
} }
pub fn video_create_transport_unsupported_reason( pub fn video_create_transport_unsupported_reason(
@@ -141,7 +157,7 @@ pub fn resolve_video_create_auth(
family: ProviderVideoCreateFamily, family: ProviderVideoCreateFamily,
) -> Option<(String, String)> { ) -> Option<(String, String)> {
match family { match family {
ProviderVideoCreateFamily::OpenAi => resolve_local_openai_bearer_auth(transport), ProviderVideoCreateFamily::OpenAi => resolve_openai_compatible_video_auth(transport),
ProviderVideoCreateFamily::Gemini => resolve_local_gemini_auth(transport), ProviderVideoCreateFamily::Gemini => resolve_local_gemini_auth(transport),
} }
} }
@@ -173,6 +189,15 @@ pub fn build_video_create_request_body(
Some(provider_request_body) Some(provider_request_body)
} }
fn resolve_openai_compatible_video_auth(
transport: &GatewayProviderTransportSnapshot,
) -> Option<(String, String)> {
resolve_local_openai_bearer_auth(transport).or_else(|| {
crate::generic_oauth::resolve_local_generic_oauth_transport_authorization(transport)
.map(|value| ("authorization".to_string(), value))
})
}
pub fn build_video_create_upstream_url( pub fn build_video_create_upstream_url(
transport: &GatewayProviderTransportSnapshot, transport: &GatewayProviderTransportSnapshot,
request_path: &str, request_path: &str,
@@ -193,7 +218,13 @@ pub fn build_video_create_upstream_url(
ProviderVideoCreateFamily::Gemini => &["key"][..], ProviderVideoCreateFamily::Gemini => &["key"][..],
}; };
return build_passthrough_path_url( return build_passthrough_path_url(
&transport.endpoint.base_url, &crate::xai::resolved_xai_request_base_url(
transport,
match family {
ProviderVideoCreateFamily::OpenAi => "openai:video",
ProviderVideoCreateFamily::Gemini => "gemini:video",
},
),
path, path,
request_query, request_query,
blocked_keys, blocked_keys,
@@ -202,8 +233,14 @@ pub fn build_video_create_upstream_url(
match family { match family {
ProviderVideoCreateFamily::OpenAi => build_passthrough_path_url( ProviderVideoCreateFamily::OpenAi => build_passthrough_path_url(
&transport.endpoint.base_url, &crate::xai::resolved_xai_request_base_url(transport, "openai:video"),
openai_video_api_root_request_path(request_path), if crate::xai::is_xai_provider_transport(transport)
&& matches!(request_path, "/v1/videos" | "/openai/v1/videos")
{
"/videos/generations"
} else {
openai_video_api_root_request_path(request_path)
},
request_query, request_query,
&[], &[],
), ),
@@ -216,6 +253,7 @@ pub fn build_video_create_upstream_url(
} }
fn openai_video_api_root_request_path(request_path: &str) -> &str { fn openai_video_api_root_request_path(request_path: &str) -> &str {
let request_path = request_path.strip_prefix("/openai").unwrap_or(request_path);
if request_path.starts_with("/v1/") { if request_path.starts_with("/v1/") {
&request_path[3..] &request_path[3..]
} else { } else {
@@ -232,6 +270,11 @@ pub fn build_video_create_headers(
input.auth_value, input.auth_value,
&BTreeMap::new(), &BTreeMap::new(),
); );
crate::xai::insert_cli_identity_headers_if_needed(
input.transport,
"openai:video",
&mut provider_request_headers,
);
if !apply_local_header_rules_with_request_headers( if !apply_local_header_rules_with_request_headers(
&mut provider_request_headers, &mut provider_request_headers,
input.header_rules, input.header_rules,
@@ -281,16 +324,22 @@ pub async fn reconstruct_local_video_task_snapshot(
return Ok(None); return Ok(None);
}; };
let Some(local_transport) = let Some(mut local_transport) =
resolve_local_video_task_transport(&transport, provider_api_format, task.model.clone()) resolve_local_video_task_transport(&transport, provider_api_format, task.model.clone())
else { else {
return Ok(None); return Ok(None);
}; };
Ok(LocalVideoTaskSnapshot::from_stored_task_with_transport( // Resolve deployment-managed nodes, system defaults and tunnel affinity just as
task, // creation does; serialized task metadata intentionally contains no credentials.
local_transport, local_transport.proxy = lookup.resolve_video_task_proxy(&transport).await;
))
let mut snapshot =
LocalVideoTaskSnapshot::from_stored_task_with_transport(task, local_transport);
if let Some(LocalVideoTaskSnapshot::OpenAi(seed)) = &mut snapshot {
seed.xai_provider = crate::xai::is_xai_provider_transport(&transport);
}
Ok(snapshot)
} }
#[cfg(test)] #[cfg(test)]
@@ -441,6 +490,46 @@ mod tests {
assert_eq!(transport.provider_id, "provider-1"); assert_eq!(transport.provider_id, "provider-1");
} }
#[tokio::test]
async fn reconstructs_video_with_configured_proxy_and_profile() {
let mut transport = sample_transport("openai:video", "oauth");
transport.provider.provider_type = "xai".into();
transport.endpoint.base_url = "https://cli-chat-proxy.grok.com/v1".into();
transport.provider.proxy = Some(json!({"enabled":true,"url":"http://127.0.0.1:9876"}));
transport.provider.config = Some(json!({"fingerprint":{"transport_profile":{
"profile_id":"test-video","backend":"reqwest_rustls","http_mode":"auto","pool_scope":"key"
}}}));
transport.key.decrypted_auth_config = Some(r#"{"using_api":false}"#.into());
let lookup = TestLookup(Some(transport));
let snapshot = reconstruct_local_video_task_snapshot(&lookup, &sample_stored_video_task())
.await
.unwrap()
.expect("proxied video must resume after restart");
let LocalVideoTaskSnapshot::OpenAi(seed) = snapshot else {
panic!("expected OpenAI video")
};
assert!(seed.xai_provider);
assert_eq!(
seed.transport.proxy.as_ref().unwrap().url.as_deref(),
Some("http://127.0.0.1:9876/")
);
assert_eq!(
seed.transport
.transport_profile
.as_ref()
.unwrap()
.profile_id,
"test-video"
);
assert_eq!(
seed.transport
.headers
.get("x-xai-token-auth")
.map(String::as_str),
Some("xai-grok-cli")
);
}
#[test] #[test]
fn resolves_gemini_video_transport() { fn resolves_gemini_video_transport() {
let transport = resolve_local_video_task_transport( let transport = resolve_local_video_task_transport(
@@ -489,6 +578,120 @@ mod tests {
assert_eq!(url, "https://api.openai.example/v1/videos?trace=1"); assert_eq!(url, "https://api.openai.example/v1/videos?trace=1");
} }
#[test]
fn xai_video_create_paths_preserve_auth_hosts_and_custom_endpoints() {
for (auth, base) in [
("oauth", "https://cli-chat-proxy.grok.com/v1"),
("api_key", "https://api.x.ai/v1"),
] {
let mut transport = sample_transport("openai:video", auth);
transport.provider.provider_type = "xai".into();
transport.endpoint.base_url = "https://cli-chat-proxy.grok.com/v1".into();
transport.key.decrypted_auth_config =
(auth == "oauth").then(|| r#"{"using_api":false}"#.into());
for path in ["/v1/videos", "/openai/v1/videos", "/v1/videos/generations"] {
assert_eq!(
build_video_create_upstream_url(
&transport,
path,
Some("trace=1"),
"grok-imagine-video",
ProviderVideoCreateFamily::OpenAi
)
.unwrap(),
format!("{base}/videos/generations?trace=1")
);
}
transport.endpoint.base_url = "https://gateway.example/prefix/v1".into();
assert_eq!(
build_video_create_upstream_url(
&transport,
"/openai/v1/videos",
None,
"grok-imagine-video",
ProviderVideoCreateFamily::OpenAi
)
.unwrap(),
"https://gateway.example/prefix/v1/videos/generations"
);
transport.endpoint.custom_path = Some("/custom/videos/generations".into());
let url = build_video_create_upstream_url(
&transport,
"/openai/v1/videos",
None,
"grok-imagine-video",
ProviderVideoCreateFamily::OpenAi,
)
.unwrap();
assert!(url.ends_with("/custom/videos/generations"), "{url}");
}
let transport = sample_transport("openai:video", "api_key");
assert_eq!(
build_video_create_upstream_url(
&transport,
"/openai/v1/videos",
None,
"sora",
ProviderVideoCreateFamily::OpenAi
),
build_video_create_upstream_url(
&transport,
"/v1/videos",
None,
"sora",
ProviderVideoCreateFamily::OpenAi
)
);
}
#[test]
fn xai_oauth_video_uses_cli_proxy() {
let mut transport = sample_transport("openai:video", "oauth");
transport.provider.provider_type = "xai".to_string();
transport.endpoint.base_url = "https://cli-chat-proxy.grok.com/v1".to_string();
transport.key.decrypted_auth_config =
Some(r#"{"refresh_token":"rt","using_api":false}"#.to_string());
let url = build_video_create_upstream_url(
&transport,
"/v1/videos/generations",
None,
"grok-imagine-video",
ProviderVideoCreateFamily::OpenAi,
)
.expect("url should build");
assert_eq!(url, "https://cli-chat-proxy.grok.com/v1/videos/generations");
let headers = build_video_create_headers(ProviderVideoCreateHeadersInput {
transport: &transport,
headers: &http::HeaderMap::new(),
auth_header: "authorization",
auth_value: "Bearer test-token",
header_rules: None,
provider_request_body: &json!({"prompt": "A cat"}),
original_request_body: &json!({"prompt": "A cat"}),
})
.unwrap();
assert_eq!(
headers.get("x-xai-token-auth").map(String::as_str),
Some("xai-grok-cli")
);
let reconstructed = super::resolve_local_video_task_transport(
&transport,
"openai:video",
Some("grok-imagine-video".into()),
)
.unwrap();
assert_eq!(
reconstructed.upstream_base_url,
"https://cli-chat-proxy.grok.com/v1"
);
assert_eq!(
reconstructed.headers.get("x-xai-token-auth"),
headers.get("x-xai-token-auth")
);
}
#[test] #[test]
fn builds_gemini_video_create_url_and_removes_client_key_query() { fn builds_gemini_video_create_url_and_removes_client_key_query() {
let transport = sample_transport("gemini:video", "api_key"); let transport = sample_transport("gemini:video", "api_key");
@@ -512,6 +715,7 @@ mod tests {
let provider_request_body = json!({"prompt": "make a clip"}); let provider_request_body = json!({"prompt": "make a clip"});
let original_request_body = provider_request_body.clone(); let original_request_body = provider_request_body.clone();
let headers = build_video_create_headers(ProviderVideoCreateHeadersInput { let headers = build_video_create_headers(ProviderVideoCreateHeadersInput {
transport: &sample_transport("openai:video", "bearer"),
headers: &http::HeaderMap::new(), headers: &http::HeaderMap::new(),
auth_header: "authorization", auth_header: "authorization",
auth_value: "Bearer secret", auth_value: "Bearer secret",
+444
View File
@@ -0,0 +1,444 @@
pub mod video;
use std::collections::BTreeMap;
use aether_ai_formats::normalize_api_format_alias;
use serde_json::Value;
use crate::snapshot::GatewayProviderTransportSnapshot;
pub const XAI_PROVIDER_TYPE: &str = "xai";
pub const XAI_CHAT_PROXY_BASE_URL: &str = "https://cli-chat-proxy.grok.com/v1";
pub const XAI_API_BASE_URL: &str = "https://api.x.ai/v1";
pub const XAI_CLIENT_VERSION: &str = "0.2.120";
pub const XAI_TOKEN_AUTH_HEADER: &str = "x-xai-token-auth";
pub const XAI_TOKEN_AUTH_VALUE: &str = "xai-grok-cli";
pub const XAI_CLIENT_VERSION_HEADER: &str = "x-grok-client-version";
pub const XAI_CLIENT_IDENTIFIER_HEADER: &str = "x-grok-client-identifier";
pub const XAI_CLIENT_IDENTIFIER_VALUE: &str = "grok-shell";
pub const XAI_AUTHENTICATE_RESPONSE_HEADER: &str = "x-authenticateresponse";
pub const XAI_AUTHENTICATE_RESPONSE_VALUE: &str = "authenticate-response";
pub fn xai_cli_user_agent() -> String {
format!("xai-grok-workspace/{XAI_CLIENT_VERSION}")
}
pub fn is_xai_provider_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case(XAI_PROVIDER_TYPE)
}
pub fn xai_uses_official_api(api_format: &str) -> bool {
matches!(
normalize_api_format_alias(api_format).as_str(),
"openai:responses:compact"
)
}
pub fn resolved_xai_upstream_base_url(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> Option<String> {
if !is_xai_provider_transport(transport) {
return None;
}
let stored = transport.endpoint.base_url.trim();
if xai_uses_official_api(api_format) {
if stored.is_empty()
|| is_cli_chat_proxy_base_url(stored)
|| is_official_api_base_url(stored)
{
return Some(XAI_API_BASE_URL.to_string());
}
return Some(trim_base_url(stored));
}
if xai_using_api(transport) {
if stored.is_empty() || is_cli_chat_proxy_base_url(stored) {
return Some(XAI_API_BASE_URL.to_string());
}
return Some(trim_base_url(stored));
}
if stored.is_empty() || is_official_api_base_url(stored) {
return Some(XAI_CHAT_PROXY_BASE_URL.to_string());
}
Some(trim_base_url(stored))
}
pub fn resolved_xai_request_base_url(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> String {
resolved_xai_upstream_base_url(transport, api_format)
.unwrap_or_else(|| trim_base_url(&transport.endpoint.base_url))
}
pub fn should_attach_cli_identity_headers(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
) -> bool {
if !is_xai_provider_transport(transport) {
return false;
}
if xai_uses_official_api(api_format) {
return false;
}
resolved_xai_upstream_base_url(transport, api_format)
.as_deref()
.is_some_and(is_cli_chat_proxy_base_url)
}
pub fn insert_cli_identity_headers(headers: &mut BTreeMap<String, String>) {
let user_agent = xai_cli_user_agent();
for (name, value) in [
(XAI_TOKEN_AUTH_HEADER, XAI_TOKEN_AUTH_VALUE),
(XAI_CLIENT_VERSION_HEADER, XAI_CLIENT_VERSION),
("user-agent", user_agent.as_str()),
(XAI_CLIENT_IDENTIFIER_HEADER, XAI_CLIENT_IDENTIFIER_VALUE),
(
XAI_AUTHENTICATE_RESPONSE_HEADER,
XAI_AUTHENTICATE_RESPONSE_VALUE,
),
] {
if !headers
.keys()
.any(|existing| existing.eq_ignore_ascii_case(name))
{
headers.insert(name.to_string(), value.to_string());
}
}
}
pub fn insert_cli_identity_headers_if_needed(
transport: &GatewayProviderTransportSnapshot,
api_format: &str,
headers: &mut BTreeMap<String, String>,
) {
if should_attach_cli_identity_headers(transport, api_format) {
insert_cli_identity_headers(headers);
}
}
pub fn xai_auth_uses_api(auth_type: &str, decrypted_auth_config: Option<&str>) -> bool {
if let Some(value) = auth_config_using_api(decrypted_auth_config) {
return value;
}
let auth_type = auth_type.trim().to_ascii_lowercase();
if auth_type == "oauth" || auth_config_has_refresh_token(decrypted_auth_config) {
return false;
}
matches!(auth_type.as_str(), "api_key" | "bearer" | "apikey")
}
pub fn extract_xai_user_id_from_auth_config(raw_auth_config: Option<&str>) -> Option<String> {
let value = parse_auth_config(raw_auth_config)?;
extract_xai_user_id_from_value(&value)
}
pub fn extract_xai_user_id_from_value(value: &Value) -> Option<String> {
const PATHS: &[&[&str]] = &[
&["userId"],
&["user_id"],
&["id"],
&["sub"],
&["user", "userId"],
&["user", "id"],
&["user", "user_id"],
&["user", "sub"],
];
PATHS.iter().find_map(|path| {
let mut current = value;
for key in *path {
current = current.get(*key)?;
}
coerce_xai_id(current)
})
}
fn xai_using_api(transport: &GatewayProviderTransportSnapshot) -> bool {
xai_auth_uses_api(
transport.key.auth_type.as_str(),
transport.key.decrypted_auth_config.as_deref(),
)
}
fn coerce_xai_id(value: &Value) -> Option<String> {
match value {
Value::String(text) => {
let trimmed = text.trim();
(!trimmed.is_empty()).then(|| trimmed.to_string())
}
Value::Number(number) => {
let rendered = number.to_string();
(!rendered.is_empty()).then_some(rendered)
}
_ => None,
}
}
fn auth_config_using_api(raw_auth_config: Option<&str>) -> Option<bool> {
let value = parse_auth_config(raw_auth_config)?;
let using_api = value.get("using_api")?;
match using_api {
Value::Bool(value) => Some(*value),
Value::String(value) => value.trim().parse::<bool>().ok(),
_ => None,
}
}
fn auth_config_has_refresh_token(raw_auth_config: Option<&str>) -> bool {
let value = match parse_auth_config(raw_auth_config) {
Some(value) => value,
None => return false,
};
["refresh_token", "refreshToken"]
.iter()
.find_map(|field| value.get(*field).and_then(Value::as_str))
.map(str::trim)
.is_some_and(|value| !value.is_empty())
}
fn parse_auth_config(raw_auth_config: Option<&str>) -> Option<Value> {
raw_auth_config
.map(str::trim)
.filter(|value| !value.is_empty())
.and_then(|value| serde_json::from_str::<Value>(value).ok())
}
fn trim_base_url(url: &str) -> String {
url.trim().trim_end_matches('/').to_string()
}
fn normalize_base_url(url: &str) -> String {
trim_base_url(url).to_ascii_lowercase()
}
fn is_official_api_base_url(url: &str) -> bool {
normalize_base_url(url) == normalize_base_url(XAI_API_BASE_URL)
}
fn is_cli_chat_proxy_base_url(url: &str) -> bool {
normalize_base_url(url) == normalize_base_url(XAI_CHAT_PROXY_BASE_URL)
}
#[cfg(test)]
mod tests {
use super::{
insert_cli_identity_headers_if_needed, is_xai_provider_transport,
resolved_xai_upstream_base_url, should_attach_cli_identity_headers, XAI_API_BASE_URL,
XAI_CHAT_PROXY_BASE_URL, XAI_CLIENT_IDENTIFIER_VALUE, XAI_TOKEN_AUTH_VALUE,
};
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use std::collections::BTreeMap;
fn sample_transport(
auth_type: &str,
auth_config: Option<&str>,
base_url: &str,
) -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-xai".to_string(),
name: "xAI".to_string(),
provider_type: "xai".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: true,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-xai".to_string(),
provider_id: "provider-xai".to_string(),
api_format: "openai:responses".to_string(),
api_family: None,
endpoint_kind: None,
is_active: true,
base_url: base_url.to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-xai".to_string(),
provider_id: "provider-xai".to_string(),
name: "key".to_string(),
auth_type: auth_type.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,
upstream_metadata: None,
decrypted_api_key: "access-token".to_string(),
decrypted_auth_config: auth_config.map(ToOwned::to_owned),
},
}
}
#[test]
fn oauth_defaults_to_cli_chat_proxy_for_responses() {
let transport = sample_transport(
"oauth",
Some(r#"{"refresh_token":"rt","using_api":false}"#),
XAI_CHAT_PROXY_BASE_URL,
);
assert!(is_xai_provider_transport(&transport));
assert_eq!(
resolved_xai_upstream_base_url(&transport, "openai:responses").as_deref(),
Some(XAI_CHAT_PROXY_BASE_URL)
);
assert!(should_attach_cli_identity_headers(
&transport,
"openai:responses"
));
}
#[test]
fn compact_and_using_api_stay_on_official_api() {
let oauth = sample_transport(
"oauth",
Some(r#"{"refresh_token":"rt","using_api":false}"#),
XAI_CHAT_PROXY_BASE_URL,
);
assert_eq!(
resolved_xai_upstream_base_url(&oauth, "openai:responses:compact").as_deref(),
Some(XAI_API_BASE_URL)
);
assert!(!should_attach_cli_identity_headers(
&oauth,
"openai:responses:compact"
));
let api_key = sample_transport(
"oauth",
Some(r#"{"using_api":true}"#),
XAI_CHAT_PROXY_BASE_URL,
);
assert_eq!(
resolved_xai_upstream_base_url(&api_key, "openai:responses").as_deref(),
Some(XAI_API_BASE_URL)
);
assert!(!should_attach_cli_identity_headers(
&api_key,
"openai:responses"
));
}
#[test]
fn media_routing_and_cli_headers_follow_auth_and_base_url() {
for api_format in ["openai:image", "openai:video"] {
for stored in ["", XAI_API_BASE_URL, XAI_CHAT_PROXY_BASE_URL] {
for (auth_type, config, expected) in [
(
"oauth",
Some(r#"{"refresh_token":"rt","using_api":false}"#),
XAI_CHAT_PROXY_BASE_URL,
),
("oauth", Some(r#"{"using_api":true}"#), XAI_API_BASE_URL),
("bearer", None, XAI_API_BASE_URL),
] {
let transport = sample_transport(auth_type, config, stored);
assert_eq!(
resolved_xai_upstream_base_url(&transport, api_format).as_deref(),
Some(expected)
);
assert_eq!(
should_attach_cli_identity_headers(&transport, api_format),
expected == XAI_CHAT_PROXY_BASE_URL
);
}
}
let custom = sample_transport("oauth", None, "https://custom.example/v1");
assert_eq!(
resolved_xai_upstream_base_url(&custom, api_format).as_deref(),
Some("https://custom.example/v1")
);
assert!(!should_attach_cli_identity_headers(&custom, api_format));
}
}
#[test]
fn bearer_without_refresh_uses_official_api() {
let transport = sample_transport("bearer", None, XAI_CHAT_PROXY_BASE_URL);
assert_eq!(
resolved_xai_upstream_base_url(&transport, "openai:responses").as_deref(),
Some(XAI_API_BASE_URL)
);
}
#[test]
fn cli_headers_do_not_override_existing_values() {
let transport = sample_transport(
"oauth",
Some(r#"{"refresh_token":"rt"}"#),
XAI_CHAT_PROXY_BASE_URL,
);
let mut headers = BTreeMap::from([(
"x-grok-client-identifier".to_string(),
"custom-client".to_string(),
)]);
insert_cli_identity_headers_if_needed(&transport, "openai:responses", &mut headers);
assert_eq!(
headers.get("x-grok-client-identifier").map(String::as_str),
Some("custom-client")
);
assert_eq!(
headers.get("x-xai-token-auth").map(String::as_str),
Some(XAI_TOKEN_AUTH_VALUE)
);
assert_eq!(
headers.get("x-authenticateresponse").map(String::as_str),
Some("authenticate-response")
);
assert_ne!(
headers.get("x-grok-client-identifier").map(String::as_str),
Some(XAI_CLIENT_IDENTIFIER_VALUE)
);
}
#[test]
fn extracts_user_id_from_user_payload_and_auth_config_sub() {
use super::{
extract_xai_user_id_from_auth_config, extract_xai_user_id_from_value, xai_auth_uses_api,
};
use serde_json::json;
assert_eq!(
extract_xai_user_id_from_value(&json!({"userId": "user-42"})).as_deref(),
Some("user-42")
);
assert_eq!(
extract_xai_user_id_from_auth_config(Some(r#"{"sub":"subject-1"}"#)).as_deref(),
Some("subject-1")
);
assert!(!xai_auth_uses_api(
"oauth",
Some(r#"{"refresh_token":"rt","using_api":false}"#)
));
assert!(xai_auth_uses_api(
"bearer",
Some(r#"{"api_key":"xai-key","using_api":true}"#)
));
}
}
@@ -0,0 +1,147 @@
use serde_json::{json, Value};
/// Native xAI video requests live under /v1; the OpenAI-compatible adapter under /openai/v1.
pub fn is_native_video_request(provider_type: &str, path: &str) -> bool {
provider_type.trim().eq_ignore_ascii_case("xai")
&& matches!(
path,
"/v1/videos" | "/v1/videos/generations" | "/v1/videos/edits" | "/v1/videos/extensions"
)
}
pub fn is_explicit_native_video_path(path: &str) -> bool {
matches!(
path,
"/v1/videos/generations" | "/v1/videos/edits" | "/v1/videos/extensions"
)
}
/// Convert the OpenAI video request contract to xAI's native contract.
/// Native requests bypass this adapter so provider-specific fields remain intact.
pub fn convert_openai_video_request(body: &Value) -> Result<Value, &'static str> {
let prompt = text(&body["prompt"]).ok_or("prompt is required")?;
let seconds = match &body["seconds"] {
Value::Null => 4,
Value::String(value) if value.trim().is_empty() => 4,
Value::String(value) => value
.trim()
.parse::<i64>()
.map_err(|_| "seconds must be an integer")?,
value => value.as_i64().ok_or("seconds must be an integer")?,
}
.clamp(1, 15);
let size = text(&body["size"]).unwrap_or("720x1280");
let default_ratio = match size {
"720x1280" | "1024x1792" => "9:16",
"1280x720" | "1792x1024" => "16:9",
_ => return Err("size must be one of 720x1280, 1280x720, 1024x1792, or 1792x1024"),
};
let ratio = match text(&body["aspect_ratio"])
.unwrap_or("")
.to_ascii_lowercase()
.as_str()
{
"square" | "1:1" => "1:1",
"landscape" | "16:9" => "16:9",
"portrait" | "9:16" => "9:16",
"4:3" => "4:3",
"3:4" => "3:4",
"3:2" => "3:2",
"2:3" => "2:3",
_ => default_ratio,
};
let resolution = if text(&body["resolution"]).is_some_and(|v| v.eq_ignore_ascii_case("480p")) {
"480p"
} else {
"720p"
};
if text(&body["input_reference"]["file_id"]).is_some() {
return Err("input_reference.file_id is not supported for xAI video generation; use input_reference.image_url");
}
let image = text(&body["input_reference"]["image_url"])
.or_else(|| image_url(&body["image"]))
.or_else(|| text(&body["image_url"]));
let references: Vec<_> = ["reference_images", "reference_image_urls"]
.into_iter()
.filter_map(|key| body[key].as_array())
.flatten()
.filter_map(image_url)
.map(|url| json!({"url":url}))
.collect();
if references.len() > 7 {
return Err("reference_images supports at most 7 images on xAI");
}
if image.is_some() && !references.is_empty() {
return Err("image and reference_images cannot be combined on xAI");
}
let mut result = json!({"model":body["model"], "prompt":prompt, "duration":seconds, "aspect_ratio":ratio, "resolution":resolution});
if let Some(url) = image {
result["image"] = json!({"url":url});
}
if !references.is_empty() {
result["reference_images"] = json!(references);
}
Ok(result)
}
fn text(value: &Value) -> Option<&str> {
value
.as_str()
.map(str::trim)
.filter(|value| !value.is_empty())
}
fn image_url(value: &Value) -> Option<&str> {
text(value)
.or_else(|| text(&value["url"]))
.or_else(|| text(&value["image_url"]))
.or_else(|| text(&value["image_url"]["url"]))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn xai_video_compatibility_maps_duration_size_and_references() {
let converted = convert_openai_video_request(&json!({
"model":"grok-imagine-video", "prompt":"A cat", "seconds":"8", "size":"1280x720",
"reference_images":[{"image_url":{"url":"https://example.com/a.png"}}],
"reference_image_urls":["https://example.com/b.png"]
}))
.unwrap();
assert_eq!(
converted,
json!({"model":"grok-imagine-video", "prompt":"A cat", "duration":8,
"aspect_ratio":"16:9", "resolution":"720p", "reference_images":[{"url":"https://example.com/a.png"},{"url":"https://example.com/b.png"}]})
);
let defaults = convert_openai_video_request(&json!({"prompt":"A cat"})).unwrap();
assert_eq!(defaults["duration"], 4);
assert_eq!(defaults["aspect_ratio"], "9:16");
for (seconds, expected) in [(-1, 1), (30, 15)] {
assert_eq!(
convert_openai_video_request(&json!({"prompt":"A cat", "seconds":seconds}))
.unwrap()["duration"],
expected
);
}
}
#[test]
fn xai_video_compatibility_validates_requests_and_maps_image_input() {
for invalid in [
json!({}),
json!({"prompt":"cat","seconds":"1.5"}),
json!({"prompt":"cat","size":"foo"}),
json!({"prompt":"cat","input_reference":{"file_id":"file-1"}}),
json!({"prompt":"cat","image":"https://example.com/a.png","reference_images":["https://example.com/b.png"]}),
json!({"prompt":"cat","reference_images":vec!["https://example.com/a.png";8]}),
] {
assert!(convert_openai_video_request(&invalid).is_err(), "{invalid}");
}
let body = convert_openai_video_request(&json!({"prompt":"cat","input_reference":{"image_url":"https://example.com/a.png"},"aspect_ratio":"square","resolution":"480p"})).unwrap();
assert_eq!(body["image"]["url"], "https://example.com/a.png");
assert_eq!(body["aspect_ratio"], "1:1");
assert_eq!(body["resolution"], "480p");
}
}
+15 -2
View File
@@ -1,5 +1,8 @@
use std::path::PathBuf; use std::path::PathBuf;
use std::process::{Child, Command, Stdio}; use std::process::{Child, Command, Stdio};
use std::sync::atomic::{AtomicU64, Ordering};
static POSTGRES_WORKDIR_SEQ: AtomicU64 = AtomicU64::new(0);
use aether_data::driver::postgres::PostgresPoolConfig; use aether_data::driver::postgres::PostgresPoolConfig;
use aether_data::{DataBackends, DataLayerConfig}; use aether_data::{DataBackends, DataLayerConfig};
@@ -21,10 +24,20 @@ pub struct ManagedPostgresServer {
impl ManagedPostgresServer { impl ManagedPostgresServer {
pub async fn start() -> Result<Self, Box<dyn std::error::Error>> { pub async fn start() -> Result<Self, Box<dyn std::error::Error>> {
let port = reserve_local_port()?; let port = reserve_local_port()?;
// pid+port is not unique: cargo test shares one PID, and ephemeral ports
// are reused after the listener is dropped. Parallel e2e tests then hit
// create_dir AlreadyExists.
let seq = POSTGRES_WORKDIR_SEQ.fetch_add(1, Ordering::Relaxed);
let nanos = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|duration| duration.as_nanos())
.unwrap_or(0);
let workdir = std::env::temp_dir().join(format!( let workdir = std::env::temp_dir().join(format!(
"aether-postgres-baseline-{}-{}", "aether-postgres-baseline-{}-{}-{}-{}",
std::process::id(), std::process::id(),
port port,
seq,
nanos
)); ));
let data_dir = workdir.join("data"); let data_dir = workdir.join("data");
std::fs::create_dir(&workdir)?; std::fs::create_dir(&workdir)?;
+79 -42
View File
@@ -6649,7 +6649,7 @@ mod tests {
} }
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex}; use std::sync::{Arc, Mutex};
use std::time::Instant; use std::time::Instant;
@@ -7375,9 +7375,27 @@ mod tests {
queue: Arc<dyn RuntimeQueueStore>, queue: Arc<dyn RuntimeQueueStore>,
policy_started: Arc<tokio::sync::Notify>, policy_started: Arc<tokio::sync::Notify>,
release_policy: Arc<tokio::sync::Notify>, release_policy: Arc<tokio::sync::Notify>,
policy_released: Arc<AtomicBool>,
policy_reads: Arc<AtomicUsize>, policy_reads: Arc<AtomicUsize>,
} }
impl BlockingPolicyQueueConfiguredUsageStore {
fn new(queue: Arc<dyn RuntimeQueueStore>) -> Self {
Self {
queue,
policy_started: Arc::new(tokio::sync::Notify::new()),
release_policy: Arc::new(tokio::sync::Notify::new()),
policy_released: Arc::new(AtomicBool::new(false)),
policy_reads: Arc::new(AtomicUsize::new(0)),
}
}
fn release_blocked_policy(&self) {
self.policy_released.store(true, Ordering::Release);
self.release_policy.notify_waiters();
}
}
#[derive(Default)] #[derive(Default)]
struct FailingPolicyUsageStore { struct FailingPolicyUsageStore {
inner: NoRedisUsageStore, inner: NoRedisUsageStore,
@@ -8409,7 +8427,18 @@ mod tests {
async fn body_capture_policy(&self) -> Result<UsageBodyCapturePolicy, DataLayerError> { async fn body_capture_policy(&self) -> Result<UsageBodyCapturePolicy, DataLayerError> {
self.policy_reads.fetch_add(1, Ordering::AcqRel); self.policy_reads.fetch_add(1, Ordering::AcqRel);
self.policy_started.notify_one(); self.policy_started.notify_one();
self.release_policy.notified().await; // Latch the gate: Notify is edge-triggered, and later policy reads
// (or a waiter that subscribed after a single notify) must not hang.
loop {
if self.policy_released.load(Ordering::Acquire) {
break;
}
let notified = self.release_policy.notified();
if self.policy_released.load(Ordering::Acquire) {
break;
}
notified.await;
}
Ok(UsageBodyCapturePolicy::default()) Ok(UsageBodyCapturePolicy::default())
} }
} }
@@ -9043,15 +9072,39 @@ mod tests {
.await .await
.expect("a duplicate first-byte marker must release the terminal barrier"); .expect("a duplicate first-byte marker must release the terminal barrier");
let records = store.records.lock().expect("records lock"); {
assert_eq!( let records = store.records.lock().expect("records lock");
records.len(), assert_eq!(
2, records.len(),
"the duplicate first byte must be coalesced" 2,
); "the duplicate first byte must be coalesced"
assert_eq!(records[0].status, "streaming"); );
assert_eq!(records[1].status, "completed"); assert_eq!(records[0].status, "streaming");
drop(records); assert_eq!(records[1].status, "completed");
}
// The terminal persistence notification can arrive before the submission
// dispatcher accounts for its completed task and releases admission.
timeout(Duration::from_secs(1), async {
loop {
let snapshot = runtime.metrics_snapshot();
if snapshot.lifecycle_submission_pending == 0
&& snapshot.first_byte_persistence_pending == 0
&& snapshot.ordered_lifecycle_pending == 0
&& runtime
.lifecycle_submission
.state
.admission
.available_permits()
== CAPACITY
{
break;
}
sleep(Duration::from_millis(1)).await;
}
})
.await
.expect("duplicate first-byte submission accounting should drain");
let snapshot = runtime.metrics_snapshot(); let snapshot = runtime.metrics_snapshot();
assert_eq!(snapshot.lifecycle_submission_pending, 0); assert_eq!(snapshot.lifecycle_submission_pending, 0);
@@ -12325,12 +12378,9 @@ mod tests {
async fn event_capture_budget_bounds_blocked_policy_waiters_and_releases_on_cancel_or_basic() { async fn event_capture_budget_bounds_blocked_policy_waiters_and_releases_on_cancel_or_basic() {
for limit in [0, 64 * 1024] { for limit in [0, 64 * 1024] {
let runtime = UsageRuntime::new(UsageRuntimeConfig::default()).expect("runtime"); let runtime = UsageRuntime::new(UsageRuntimeConfig::default()).expect("runtime");
let store = BlockingPolicyQueueConfiguredUsageStore { let store = BlockingPolicyQueueConfiguredUsageStore::new(Arc::new(
queue: Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())), RuntimeState::memory(MemoryRuntimeStateConfig::default()),
policy_started: Arc::new(tokio::sync::Notify::new()), ));
release_policy: Arc::new(tokio::sync::Notify::new()),
policy_reads: Arc::new(AtomicUsize::new(0)),
};
let budget = Arc::new(crate::event_capture_budget::EventCaptureMemoryBudget::new( let budget = Arc::new(crate::event_capture_budget::EventCaptureMemoryBudget::new(
limit, limit,
)); ));
@@ -12391,7 +12441,7 @@ mod tests {
.await .await
.expect("replacement policy read starts"); .expect("replacement policy read starts");
assert_eq!(budget.retained_bytes(), retained); assert_eq!(budget.retained_bytes(), retained);
store.release_policy.notify_one(); store.release_blocked_policy();
let event = timeout(Duration::from_secs(2), completing) let event = timeout(Duration::from_secs(2), completing)
.await .await
.expect("Basic policy completes") .expect("Basic policy completes")
@@ -13398,12 +13448,7 @@ mod tests {
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0)); let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0));
let queue: Arc<dyn RuntimeQueueStore> = tracked_queue.clone(); let queue: Arc<dyn RuntimeQueueStore> = tracked_queue.clone();
let store = BlockingPolicyQueueConfiguredUsageStore { let store = BlockingPolicyQueueConfiguredUsageStore::new(queue);
queue,
policy_started: Arc::new(tokio::sync::Notify::new()),
release_policy: Arc::new(tokio::sync::Notify::new()),
policy_reads: Arc::new(AtomicUsize::new(0)),
};
let runtime = UsageRuntime::new(config).expect("usage runtime should build"); let runtime = UsageRuntime::new(config).expect("usage runtime should build");
let request_id = "req-terminal-seed-waits-for-turn"; let request_id = "req-terminal-seed-waits-for-turn";
let plan = terminal_test_plan(request_id); let plan = terminal_test_plan(request_id);
@@ -13425,9 +13470,10 @@ mod tests {
assert_eq!(blocked_snapshot.terminal_submission_in_flight, 0); assert_eq!(blocked_snapshot.terminal_submission_in_flight, 0);
assert!(blocked_snapshot.lifecycle_submission_pending >= 2); assert!(blocked_snapshot.lifecycle_submission_pending >= 2);
store.release_policy.notify_waiters(); store.release_blocked_policy();
timeout(Duration::from_secs(2), async { timeout(Duration::from_secs(2), async {
loop { loop {
store.release_blocked_policy();
let snapshot = runtime.metrics_snapshot(); let snapshot = runtime.metrics_snapshot();
if tracked_queue.successful_appends.load(Ordering::Acquire) == 1 if tracked_queue.successful_appends.load(Ordering::Acquire) == 1
&& snapshot.lifecycle_submission_pending == 0 && snapshot.lifecycle_submission_pending == 0
@@ -13435,7 +13481,7 @@ mod tests {
{ {
break; break;
} }
tokio::task::yield_now().await; sleep(Duration::from_millis(1)).await;
} }
}) })
.await .await
@@ -13468,12 +13514,7 @@ mod tests {
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0)); let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0));
let queue: Arc<dyn RuntimeQueueStore> = tracked_queue.clone(); let queue: Arc<dyn RuntimeQueueStore> = tracked_queue.clone();
let store = BlockingPolicyQueueConfiguredUsageStore { let store = BlockingPolicyQueueConfiguredUsageStore::new(queue);
queue,
policy_started: Arc::new(tokio::sync::Notify::new()),
release_policy: Arc::new(tokio::sync::Notify::new()),
policy_reads: Arc::new(AtomicUsize::new(0)),
};
let runtime = UsageRuntime::new(config).expect("usage runtime should build"); let runtime = UsageRuntime::new(config).expect("usage runtime should build");
let policy_started = store.policy_started.notified(); let policy_started = store.policy_started.notified();
@@ -13520,9 +13561,10 @@ mod tests {
assert_eq!(blocked_snapshot.terminal_submission_in_flight, 1); assert_eq!(blocked_snapshot.terminal_submission_in_flight, 1);
assert!(blocked_snapshot.lifecycle_submission_pending <= BACKLOG + 1); assert!(blocked_snapshot.lifecycle_submission_pending <= BACKLOG + 1);
store.release_policy.notify_waiters(); store.release_blocked_policy();
timeout(Duration::from_secs(5), async { timeout(Duration::from_secs(5), async {
loop { loop {
store.release_blocked_policy();
let snapshot = runtime.metrics_snapshot(); let snapshot = runtime.metrics_snapshot();
if tracked_queue.successful_appends.load(Ordering::Acquire) == BACKLOG + 1 if tracked_queue.successful_appends.load(Ordering::Acquire) == BACKLOG + 1
&& snapshot.lifecycle_submission_pending == 0 && snapshot.lifecycle_submission_pending == 0
@@ -13530,7 +13572,7 @@ mod tests {
{ {
break; break;
} }
tokio::task::yield_now().await; sleep(Duration::from_millis(1)).await;
} }
}) })
.await .await
@@ -13828,12 +13870,7 @@ mod tests {
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0)); let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0));
let queue: Arc<dyn RuntimeQueueStore> = tracked_queue.clone(); let queue: Arc<dyn RuntimeQueueStore> = tracked_queue.clone();
let store = BlockingPolicyQueueConfiguredUsageStore { let store = BlockingPolicyQueueConfiguredUsageStore::new(queue);
queue,
policy_started: Arc::new(tokio::sync::Notify::new()),
release_policy: Arc::new(tokio::sync::Notify::new()),
policy_reads: Arc::new(AtomicUsize::new(0)),
};
let runtime = UsageRuntime::new(config).expect("usage runtime should build"); let runtime = UsageRuntime::new(config).expect("usage runtime should build");
let policy_started = store.policy_started.notified(); let policy_started = store.policy_started.notified();
runtime runtime
@@ -13892,10 +13929,10 @@ mod tests {
.expect("terminal submissions should reach the execution backlog"); .expect("terminal submissions should reach the execution backlog");
let saturated_snapshot = runtime.metrics_snapshot(); let saturated_snapshot = runtime.metrics_snapshot();
store.release_policy.notify_waiters(); store.release_blocked_policy();
let all_completed = timeout(Duration::from_secs(2), async { let all_completed = timeout(Duration::from_secs(2), async {
loop { loop {
store.release_policy.notify_waiters(); store.release_blocked_policy();
if tracked_queue.successful_appends.load(Ordering::Acquire) if tracked_queue.successful_appends.load(Ordering::Acquire)
== EXCESS_SUBMISSIONS + 1 == EXCESS_SUBMISSIONS + 1
&& runtime.metrics_snapshot().terminal_submission_in_flight == 0 && runtime.metrics_snapshot().terminal_submission_in_flight == 0
@@ -12,5 +12,6 @@ aether-data-contracts.workspace = true
async-trait.workspace = true async-trait.workspace = true
serde.workspace = true serde.workspace = true
serde_json.workspace = true serde_json.workspace = true
sha2.workspace = true
url.workspace = true url.workspace = true
uuid.workspace = true uuid.workspace = true
+181 -17
View File
@@ -43,6 +43,27 @@ pub fn map_openai_stored_task_to_read_response(
} }
fn build_openai_stored_task_body(task: StoredVideoTask, status: VideoTaskStatus) -> Value { fn build_openai_stored_task_body(task: StoredVideoTask, status: VideoTaskStatus) -> Value {
if task.client_api_format.as_deref() == Some("xai:video") {
let mut body = json!({"status":match status {
VideoTaskStatus::Completed => "done",
VideoTaskStatus::Expired => "expired",
VideoTaskStatus::Failed | VideoTaskStatus::Cancelled | VideoTaskStatus::Deleted => "failed",
_ => "pending",
}});
if let Some(model) = task.model {
body["model"] = json!(model);
}
if let Some(url) = task.video_url {
body["video"] = json!({"url":url});
if let Some(duration) = task.duration_seconds {
body["video"]["duration"] = json!(duration);
}
}
if status == VideoTaskStatus::Failed {
body["error"] = json!({"code":sanitize_video_task_error_code(task.error_code).unwrap_or_else(|| "unknown".into()),"message":"Video generation failed"});
}
return body;
}
let mut body = json!({ let mut body = json!({
"id": task.id, "id": task.id,
"object": "video", "object": "video",
@@ -57,6 +78,9 @@ fn build_openai_stored_task_body(task: StoredVideoTask, status: VideoTaskStatus)
if let Some(prompt) = task.prompt { if let Some(prompt) = task.prompt {
body["prompt"] = Value::String(prompt); body["prompt"] = Value::String(prompt);
} }
if let Some(seconds) = task.duration_seconds {
body["seconds"] = json!(seconds.to_string());
}
if let Some(size) = task.size { if let Some(size) = task.size {
body["size"] = Value::String(size); body["size"] = Value::String(size);
} }
@@ -91,21 +115,97 @@ fn map_openai_stored_task_status(status: VideoTaskStatus) -> &'static str {
} }
impl OpenAiVideoTaskSeed { impl OpenAiVideoTaskSeed {
pub fn uses_xai_provider(&self) -> bool {
self.xai_provider || self.is_xai_native()
}
pub fn is_xai_native(&self) -> bool {
self.persistence.client_api_format == "xai:video"
}
pub fn native_create_body_json(&self) -> Value {
let mut body = self.native_response.clone().unwrap_or_else(|| json!({}));
body["request_id"] = json!(self.local_task_id);
if body.get("id").is_some() {
body["id"] = json!(self.local_task_id);
}
body
}
fn native_read_body_json(&self) -> Value {
if let Some(mut body) = self.native_response.clone().filter(|body| {
body.get("status").is_some()
|| body.get("error").is_some()
|| body.get("code").is_some()
}) {
if body.get("request_id").is_some() {
body["request_id"] = json!(self.local_task_id);
}
if body.get("id").is_some() {
body["id"] = json!(self.local_task_id);
}
return body;
}
let mut body = json!({"status":match self.status {
LocalVideoTaskStatus::Completed => "done",
LocalVideoTaskStatus::Expired => "expired",
LocalVideoTaskStatus::Failed | LocalVideoTaskStatus::Cancelled | LocalVideoTaskStatus::Deleted => "failed",
_ => "pending",
}});
if let Some(model) = &self.model {
body["model"] = json!(model);
}
if let Some(url) = &self.video_url {
body["video"] = json!({"url":url});
if let Some(duration) = self.seconds.as_deref().and_then(|v| v.parse::<u64>().ok()) {
body["video"]["duration"] = json!(duration);
}
}
if self.error_code.is_some() {
body["error"] = json!({"code":self.error_code,"message":"Video generation failed"});
}
body
}
pub fn apply_provider_body(&mut self, provider_body: &Map<String, Value>) { pub fn apply_provider_body(&mut self, provider_body: &Map<String, Value>) {
if self.uses_xai_provider() {
self.native_response = Some(Value::Object(provider_body.clone()));
}
let raw_status = provider_body let raw_status = provider_body
.get("status") .get("status")
.and_then(Value::as_str) .and_then(Value::as_str)
.map(str::trim) .map(str::trim)
.unwrap_or_default(); .unwrap_or_default();
self.status = match raw_status { // Accept xAI's native lifecycle vocabulary alongside OpenAI's fields.
"queued" => LocalVideoTaskStatus::Queued, self.status = match raw_status.to_ascii_lowercase().as_str() {
"processing" => LocalVideoTaskStatus::Processing, "queued" | "pending" => LocalVideoTaskStatus::Queued,
"completed" => LocalVideoTaskStatus::Completed, "processing" | "in_progress" | "running" => LocalVideoTaskStatus::Processing,
"failed" => LocalVideoTaskStatus::Failed, "completed" | "done" | "succeeded" | "success" => LocalVideoTaskStatus::Completed,
"cancelled" => LocalVideoTaskStatus::Cancelled, "failed" | "error" => LocalVideoTaskStatus::Failed,
"cancelled" | "canceled" => LocalVideoTaskStatus::Cancelled,
"expired" => LocalVideoTaskStatus::Expired, "expired" => LocalVideoTaskStatus::Expired,
_ => LocalVideoTaskStatus::Submitted, _ => LocalVideoTaskStatus::Submitted,
}; };
let error = provider_body.get("error").filter(|value| !value.is_null());
let error_code = provider_body
.get("code")
.and_then(Value::as_str)
.filter(|value| !value.trim().is_empty())
.or_else(|| {
error
.and_then(|value| value.get("code"))
.and_then(Value::as_str)
});
// xAI may report a failed job as a 200 response with code/error only.
if (error.is_some() || error_code.is_some())
&& !matches!(
self.status,
LocalVideoTaskStatus::Cancelled | LocalVideoTaskStatus::Expired
)
{
self.status = LocalVideoTaskStatus::Failed;
}
self.progress_percent = provider_body self.progress_percent = provider_body
.get("progress") .get("progress")
.and_then(Value::as_u64) .and_then(Value::as_u64)
@@ -117,20 +217,35 @@ impl OpenAiVideoTaskSeed {
}); });
self.completed_at_unix_secs = provider_body.get("completed_at").and_then(Value::as_u64); self.completed_at_unix_secs = provider_body.get("completed_at").and_then(Value::as_u64);
self.expires_at_unix_secs = provider_body.get("expires_at").and_then(Value::as_u64); self.expires_at_unix_secs = provider_body.get("expires_at").and_then(Value::as_u64);
let error = provider_body.get("error").and_then(Value::as_object); self.error_code = sanitize_video_task_error_code(error_code.map(str::to_string));
self.error_code = sanitize_video_task_error_code(
error
.and_then(|value| value.get("code"))
.and_then(Value::as_str)
.map(str::to_string),
);
self.error_message = None; self.error_message = None;
self.video_url = provider_body self.video_url = provider_body
.get("video_url") .get("video_url")
.or_else(|| provider_body.get("url")) .or_else(|| provider_body.get("url"))
.or_else(|| provider_body.get("result_url")) .or_else(|| provider_body.get("result_url"))
.or_else(|| {
provider_body
.get("video")
.and_then(|video| video.get("url"))
})
.and_then(Value::as_str) .and_then(Value::as_str)
.map(str::to_string); .map(str::to_string);
if let Some(seconds) = provider_body
.get("seconds")
.or_else(|| {
provider_body
.get("video")
.and_then(|video| video.get("duration"))
})
.filter(|value| value.is_string() || value.is_number())
{
self.seconds = Some(
seconds
.as_str()
.map(str::to_string)
.unwrap_or_else(|| seconds.to_string()),
);
}
} }
pub fn build_content_stream_action( pub fn build_content_stream_action(
@@ -239,6 +354,9 @@ impl OpenAiVideoTaskSeed {
} }
pub fn client_body_json(&self) -> Value { pub fn client_body_json(&self) -> Value {
if self.is_xai_native() {
return self.native_read_body_json();
}
let mut body = json!({ let mut body = json!({
"id": self.local_task_id, "id": self.local_task_id,
"object": "video", "object": "video",
@@ -259,6 +377,9 @@ impl OpenAiVideoTaskSeed {
if let Some(seconds) = &self.seconds { if let Some(seconds) = &self.seconds {
body["seconds"] = Value::String(seconds.clone()); body["seconds"] = Value::String(seconds.clone());
} }
if let Some(video_url) = &self.video_url {
body["video_url"] = Value::String(video_url.clone());
}
if let Some(remixed_from_video_id) = &self.remixed_from_video_id { if let Some(remixed_from_video_id) = &self.remixed_from_video_id {
body["remixed_from_video_id"] = Value::String(remixed_from_video_id.clone()); body["remixed_from_video_id"] = Value::String(remixed_from_video_id.clone());
} }
@@ -357,12 +478,20 @@ impl OpenAiVideoTaskSeed {
} }
pub fn build_get_follow_up_plan(&self, trace_id: &str) -> Option<ExecutionPlan> { pub fn build_get_follow_up_plan(&self, trace_id: &str) -> Option<ExecutionPlan> {
if !matches!( let refreshable = matches!(
self.status, self.status,
LocalVideoTaskStatus::Submitted LocalVideoTaskStatus::Submitted
| LocalVideoTaskStatus::Queued | LocalVideoTaskStatus::Queued
| LocalVideoTaskStatus::Processing | LocalVideoTaskStatus::Processing
) { ) || (self.uses_xai_provider()
&& self.native_response.is_none()
&& matches!(
self.status,
LocalVideoTaskStatus::Completed
| LocalVideoTaskStatus::Failed
| LocalVideoTaskStatus::Expired
));
if !refreshable {
return None; return None;
} }
@@ -573,7 +702,12 @@ impl OpenAiVideoTaskSeed {
}; };
let mut record = UpsertVideoTask { let mut record = UpsertVideoTask {
id: self.local_task_id.clone(), id: self.local_task_id.clone(),
short_id: None, // The production schema requires a unique, non-null short_id (at most 16 chars).
// Derive it deterministically so repeated capture and legacy snapshot reloads agree.
short_id: Some(self.local_short_id.clone().unwrap_or_else(|| {
use sha2::{Digest, Sha256};
format!("{:x}", Sha256::digest(self.local_task_id.as_bytes()))[..16].to_string()
})),
request_id: self.persistence.request_id.clone(), request_id: self.persistence.request_id.clone(),
user_id: self.user_id.clone(), user_id: self.user_id.clone(),
api_key_id: self.api_key_id.clone(), api_key_id: self.api_key_id.clone(),
@@ -589,7 +723,11 @@ impl OpenAiVideoTaskSeed {
model: self.model.clone().or_else(|| Some(String::new())), model: self.model.clone().or_else(|| Some(String::new())),
prompt: self.prompt.clone().or_else(|| Some(String::new())), prompt: self.prompt.clone().or_else(|| Some(String::new())),
original_request_body: None, original_request_body: None,
duration_seconds: request_body_u32(&self.persistence.original_request_body, "seconds"), duration_seconds: self
.seconds
.as_deref()
.and_then(|value| value.parse().ok())
.or_else(|| request_body_u32(&self.persistence.original_request_body, "seconds")),
resolution: request_body_string(&self.persistence.original_request_body, "resolution"), resolution: request_body_string(&self.persistence.original_request_body, "resolution"),
aspect_ratio: request_body_string( aspect_ratio: request_body_string(
&self.persistence.original_request_body, &self.persistence.original_request_body,
@@ -697,6 +835,9 @@ mod tests {
#[test] #[test]
fn builds_minimal_openai_persistence_record_without_sensitive_snapshot() { fn builds_minimal_openai_persistence_record_without_sensitive_snapshot() {
let seed = OpenAiVideoTaskSeed { let seed = OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: false,
local_task_id: "task-openai-sensitive".to_string(), local_task_id: "task-openai-sensitive".to_string(),
upstream_task_id: "upstream-openai-sensitive".to_string(), upstream_task_id: "upstream-openai-sensitive".to_string(),
created_at_unix_ms: 1_712_345_678, created_at_unix_ms: 1_712_345_678,
@@ -747,6 +888,12 @@ mod tests {
let record = seed.to_upsert_record(); let record = seed.to_upsert_record();
let short_id = record
.short_id
.as_deref()
.expect("database short_id is required");
assert_eq!(short_id.len(), 16);
assert_eq!(seed.to_upsert_record().short_id, record.short_id);
assert_eq!(record.error_code.as_deref(), Some("provider_error")); assert_eq!(record.error_code.as_deref(), Some("provider_error"));
assert!(record.original_request_body.is_none()); assert!(record.original_request_body.is_none());
assert!(record.progress_message.is_none()); assert!(record.progress_message.is_none());
@@ -759,6 +906,8 @@ mod tests {
let mut stored = record.into_stored(); let mut stored = record.into_stored();
stored.status = VideoTaskStatus::Completed; stored.status = VideoTaskStatus::Completed;
// Migrated tasks can already have a short ID unrelated to the derived ID.
stored.short_id = Some("legacy-short-id".to_string());
let snapshot = let snapshot =
LocalVideoTaskSnapshot::from_stored_task_with_transport(&stored, seed.transport) LocalVideoTaskSnapshot::from_stored_task_with_transport(&stored, seed.transport)
.expect("stored task should reconstruct with current transport"); .expect("stored task should reconstruct with current transport");
@@ -766,6 +915,21 @@ mod tests {
panic!("expected OpenAI snapshot"); panic!("expected OpenAI snapshot");
}; };
assert_eq!(restored.prompt, stored.prompt); assert_eq!(restored.prompt, stored.prompt);
assert_eq!(restored.to_upsert_record().short_id, stored.short_id);
let mut embedded = stored.clone();
let mut legacy_snapshot =
serde_json::to_value(LocalVideoTaskSnapshot::OpenAi(restored.clone())).unwrap();
legacy_snapshot["OpenAi"]
.as_object_mut()
.unwrap()
.remove("local_short_id");
embedded.request_metadata = Some(json!({"rust_local_snapshot": legacy_snapshot}));
let embedded_snapshot = LocalVideoTaskSnapshot::from_stored_task(&embedded)
.expect("legacy embedded snapshot should hydrate");
assert_eq!(
embedded_snapshot.to_upsert_record().short_id,
stored.short_id
);
assert_eq!(restored.to_upsert_record().video_url, stored.video_url); assert_eq!(restored.to_upsert_record().video_url, stored.video_url);
let Some(LocalVideoTaskContentAction::StreamPlan(plan)) = let Some(LocalVideoTaskContentAction::StreamPlan(plan)) =
restored.build_content_stream_action(None, "trace-download") restored.build_content_stream_action(None, "trace-download")
+12 -4
View File
@@ -8,7 +8,9 @@ use uuid::Uuid;
use crate::{LocalVideoTaskRegistryMutation, LocalVideoTaskStatus, VideoTaskTruthSourceMode}; use crate::{LocalVideoTaskRegistryMutation, LocalVideoTaskStatus, VideoTaskTruthSourceMode};
pub fn extract_openai_task_id_from_path(path: &str) -> Option<&str> { pub fn extract_openai_task_id_from_path(path: &str) -> Option<&str> {
let suffix = path.strip_prefix("/v1/videos/")?; let suffix = path
.strip_prefix("/v1/videos/")
.or_else(|| path.strip_prefix("/openai/v1/videos/"))?;
if suffix.is_empty() if suffix.is_empty()
|| suffix.contains('/') || suffix.contains('/')
|| suffix.ends_with(":cancel") || suffix.ends_with(":cancel")
@@ -29,21 +31,27 @@ pub fn extract_gemini_short_id_from_path(path: &str) -> Option<&str> {
} }
pub fn extract_openai_task_id_from_cancel_path(path: &str) -> Option<&str> { pub fn extract_openai_task_id_from_cancel_path(path: &str) -> Option<&str> {
let suffix = path.strip_prefix("/v1/videos/")?; let suffix = path
.strip_prefix("/v1/videos/")
.or_else(|| path.strip_prefix("/openai/v1/videos/"))?;
suffix suffix
.strip_suffix("/cancel") .strip_suffix("/cancel")
.filter(|value| !value.is_empty()) .filter(|value| !value.is_empty())
} }
pub fn extract_openai_task_id_from_remix_path(path: &str) -> Option<&str> { pub fn extract_openai_task_id_from_remix_path(path: &str) -> Option<&str> {
let suffix = path.strip_prefix("/v1/videos/")?; let suffix = path
.strip_prefix("/v1/videos/")
.or_else(|| path.strip_prefix("/openai/v1/videos/"))?;
suffix suffix
.strip_suffix("/remix") .strip_suffix("/remix")
.filter(|value| !value.is_empty()) .filter(|value| !value.is_empty())
} }
pub fn extract_openai_task_id_from_content_path(path: &str) -> Option<&str> { pub fn extract_openai_task_id_from_content_path(path: &str) -> Option<&str> {
let suffix = path.strip_prefix("/v1/videos/")?; let suffix = path
.strip_prefix("/v1/videos/")
.or_else(|| path.strip_prefix("/openai/v1/videos/"))?;
suffix suffix
.strip_suffix("/content") .strip_suffix("/content")
.filter(|value| !value.is_empty()) .filter(|value| !value.is_empty())
@@ -73,7 +73,7 @@ async fn read_openai_video_task_response(
} }
None => state.find_stored_video_task(lookup).await?, None => state.find_stored_video_task(lookup).await?,
}; };
let Some(task) = task else { let Some(mut task) = task else {
return Ok(None); return Ok(None);
}; };
@@ -81,6 +81,9 @@ async fn read_openai_video_task_response(
return Ok(None); return Ok(None);
} }
if request_path.starts_with("/openai/v1/videos/") {
task.client_api_format = Some("openai:video".into());
}
Ok(Some(map_openai_stored_task_to_read_response(task))) Ok(Some(map_openai_stored_task_to_read_response(task)))
} }
@@ -105,13 +105,8 @@ impl VideoTaskService {
if self.truth_source_mode != VideoTaskTruthSourceMode::RustAuthoritative { if self.truth_source_mode != VideoTaskTruthSourceMode::RustAuthoritative {
return None; return None;
} }
match route_family { self.snapshot_for_route(route_family, request_path)
Some("openai") => extract_openai_task_id_from_path(request_path) .map(|snapshot| snapshot.read_response_for_path(request_path))
.and_then(|task_id| self.store.read_openai(task_id)),
Some("gemini") => extract_gemini_short_id_from_path(request_path)
.and_then(|short_id| self.store.read_gemini(short_id)),
_ => None,
}
} }
pub fn read_response_for_user( pub fn read_response_for_user(
@@ -126,7 +121,7 @@ impl VideoTaskService {
let snapshot = self.snapshot_for_route(route_family, request_path)?; let snapshot = self.snapshot_for_route(route_family, request_path)?;
snapshot snapshot
.belongs_to_user(user_id) .belongs_to_user(user_id)
.then(|| snapshot.read_response()) .then(|| snapshot.read_response_for_path(request_path))
} }
pub fn snapshot_for_route( pub fn snapshot_for_route(
@@ -29,6 +29,7 @@ impl LocalVideoTaskSnapshot {
// contain stale identity fields after a task import or repair. // contain stale identity fields after a task import or repair.
match &mut snapshot { match &mut snapshot {
Self::OpenAi(seed) => { Self::OpenAi(seed) => {
seed.local_short_id = task.short_id.clone();
seed.user_id = task.user_id.clone(); seed.user_id = task.user_id.clone();
seed.api_key_id = task.api_key_id.clone(); seed.api_key_id = task.api_key_id.clone();
} }
@@ -51,6 +52,9 @@ impl LocalVideoTaskSnapshot {
"openai:video" => { "openai:video" => {
let upstream_task_id = non_empty_owned(task.external_task_id.as_ref())?; let upstream_task_id = non_empty_owned(task.external_task_id.as_ref())?;
Some(Self::OpenAi(OpenAiVideoTaskSeed { Some(Self::OpenAi(OpenAiVideoTaskSeed {
local_short_id: task.short_id.clone(),
native_response: None,
xai_provider: persistence.client_api_format == "xai:video",
local_task_id: task.id.clone(), local_task_id: task.id.clone(),
upstream_task_id, upstream_task_id,
created_at_unix_ms: task.created_at_unix_ms, created_at_unix_ms: task.created_at_unix_ms,
@@ -142,6 +146,19 @@ impl LocalVideoTaskSnapshot {
} }
} }
pub fn read_response_for_path(&self, path: &str) -> LocalVideoTaskReadResponse {
if let Self::OpenAi(seed) = self {
let mut seed = seed.clone();
if path.starts_with("/openai/v1/videos/") {
seed.persistence.client_api_format = "openai:video".to_string();
} else if path.starts_with("/v1/videos/") && seed.uses_xai_provider() {
seed.persistence.client_api_format = "xai:video".to_string();
}
return Self::OpenAi(seed).read_response();
}
self.read_response()
}
pub fn read_response(&self) -> LocalVideoTaskReadResponse { pub fn read_response(&self) -> LocalVideoTaskReadResponse {
match self { match self {
Self::OpenAi(seed) => match seed.status { Self::OpenAi(seed) => match seed.status {
+276 -14
View File
@@ -19,14 +19,17 @@ impl LocalVideoTaskSeed {
) -> Option<Self> { ) -> Option<Self> {
let transport = LocalVideoTaskTransport::from_plan(plan)?; let transport = LocalVideoTaskTransport::from_plan(plan)?;
let persistence = LocalVideoTaskPersistence::from_report_context(report_context, plan); let persistence = LocalVideoTaskPersistence::from_report_context(report_context, plan);
match report_kind { let mut seed = match report_kind {
"openai_video_create_sync_finalize" => { "openai_video_create_sync_finalize" => {
let upstream_id = provider_body.get("id").and_then(Value::as_str)?.trim(); let upstream_id = openai_video_provider_task_id(provider_body)?;
if upstream_id.is_empty() {
return None;
}
Some(Self::OpenAiCreate(OpenAiVideoTaskSeed { Some(Self::OpenAiCreate(OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: report_context
.get("video_provider_xai")
.and_then(Value::as_bool)
.unwrap_or(false),
local_task_id: context_text(report_context, "local_task_id") local_task_id: context_text(report_context, "local_task_id")
.unwrap_or_else(|| Uuid::new_v4().to_string()), .unwrap_or_else(|| Uuid::new_v4().to_string()),
upstream_task_id: upstream_id.to_string(), upstream_task_id: upstream_id.to_string(),
@@ -37,8 +40,12 @@ impl LocalVideoTaskSeed {
model: context_text(report_context, "model") model: context_text(report_context, "model")
.or_else(|| request_body_text(report_context, "model")), .or_else(|| request_body_text(report_context, "model")),
prompt: request_body_text(report_context, "prompt"), prompt: request_body_text(report_context, "prompt"),
size: request_body_text(report_context, "size"), size: context_text(report_context, "video_size")
seconds: request_body_text(report_context, "seconds"), .or_else(|| request_body_text(report_context, "size")),
seconds: context_u64(report_context, "video_duration")
.map(|v| v.to_string())
.or_else(|| request_body_text(report_context, "seconds"))
.or_else(|| request_body_text(report_context, "duration")),
remixed_from_video_id: None, remixed_from_video_id: None,
status: LocalVideoTaskStatus::Submitted, status: LocalVideoTaskStatus::Submitted,
progress_percent: 0, progress_percent: 0,
@@ -52,12 +59,15 @@ impl LocalVideoTaskSeed {
})) }))
} }
"openai_video_remix_sync_finalize" => { "openai_video_remix_sync_finalize" => {
let upstream_id = provider_body.get("id").and_then(Value::as_str)?.trim(); let upstream_id = openai_video_provider_task_id(provider_body)?;
if upstream_id.is_empty() {
return None;
}
Some(Self::OpenAiRemix(OpenAiVideoTaskSeed { Some(Self::OpenAiRemix(OpenAiVideoTaskSeed {
local_short_id: None,
native_response: None,
xai_provider: report_context
.get("video_provider_xai")
.and_then(Value::as_bool)
.unwrap_or(false),
local_task_id: context_text(report_context, "local_task_id") local_task_id: context_text(report_context, "local_task_id")
.unwrap_or_else(|| Uuid::new_v4().to_string()), .unwrap_or_else(|| Uuid::new_v4().to_string()),
upstream_task_id: upstream_id.to_string(), upstream_task_id: upstream_id.to_string(),
@@ -68,8 +78,12 @@ impl LocalVideoTaskSeed {
model: context_text(report_context, "model") model: context_text(report_context, "model")
.or_else(|| request_body_text(report_context, "model")), .or_else(|| request_body_text(report_context, "model")),
prompt: request_body_text(report_context, "prompt"), prompt: request_body_text(report_context, "prompt"),
size: request_body_text(report_context, "size"), size: context_text(report_context, "video_size")
seconds: request_body_text(report_context, "seconds"), .or_else(|| request_body_text(report_context, "size")),
seconds: context_u64(report_context, "video_duration")
.map(|v| v.to_string())
.or_else(|| request_body_text(report_context, "seconds"))
.or_else(|| request_body_text(report_context, "duration")),
remixed_from_video_id: context_text(report_context, "task_id") remixed_from_video_id: context_text(report_context, "task_id")
.or_else(|| request_body_text(report_context, "remix_video_id")), .or_else(|| request_body_text(report_context, "remix_video_id")),
status: LocalVideoTaskStatus::Submitted, status: LocalVideoTaskStatus::Submitted,
@@ -110,7 +124,11 @@ impl LocalVideoTaskSeed {
})) }))
} }
_ => None, _ => None,
}?;
if let Self::OpenAiCreate(task) | Self::OpenAiRemix(task) = &mut seed {
task.apply_provider_body(provider_body);
} }
Some(seed)
} }
pub fn success_report_kind(&self) -> &'static str { pub fn success_report_kind(&self) -> &'static str {
@@ -144,12 +162,28 @@ impl LocalVideoTaskSeed {
pub fn client_body_json(&self) -> Value { pub fn client_body_json(&self) -> Value {
match self { match self {
Self::OpenAiCreate(seed) | Self::OpenAiRemix(seed) => seed.client_body_json(), Self::OpenAiCreate(seed) | Self::OpenAiRemix(seed) => {
if seed.is_xai_native() {
seed.native_create_body_json()
} else {
seed.client_body_json()
}
}
Self::GeminiCreate(seed) => seed.client_body_json(), Self::GeminiCreate(seed) => seed.client_body_json(),
} }
} }
} }
fn openai_video_provider_task_id(body: &Map<String, Value>) -> Option<&str> {
// xAI's OpenAI-compatible video creation returns request_id instead of id.
["id", "request_id"].into_iter().find_map(|field| {
body.get(field)
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
})
}
impl VideoTaskTruthSourceMode { impl VideoTaskTruthSourceMode {
pub fn prepare_sync_success( pub fn prepare_sync_success(
self, self,
@@ -353,6 +387,234 @@ mod tests {
resolve_local_sync_success_background_report_kind, resolve_local_sync_success_background_report_kind,
}; };
#[test]
fn xai_native_video_protocol_survives_persistence_and_preserves_provider_fields() {
use crate::{
LocalVideoTaskContentAction, LocalVideoTaskSnapshot, VideoTaskService,
VideoTaskTruthSourceMode,
};
let mut plan =
build_internal_finalize_video_plan("native-create", "openai:video", None).unwrap();
plan.url = "https://api.x.ai/v1/videos/generations".into();
plan.headers
.insert("authorization".into(), "Bearer test-key".into());
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
let context = json!({"local_task_id":"native-local", "user_id":"owner", "model":"grok-imagine-video", "video_client_protocol":"xai", "video_duration":6});
let success = service
.prepare_sync_success(
"openai_video_create_sync_finalize",
json!({"request_id":"native-upstream", "future_field":true})
.as_object()
.unwrap(),
context.as_object().unwrap(),
&plan,
)
.unwrap();
assert_eq!(
success.client_body_json(),
json!({"request_id":"native-local","future_field":true})
);
let mut snapshot = success.to_snapshot();
let body = json!({"status":"done","video":{"url":"https://vidgen.x.ai/video.mp4","duration":6,"respect_moderation":true},"future_field":[1,2]});
snapshot.apply_provider_body(body.as_object().unwrap());
assert_eq!(snapshot.read_response().body_json, body);
assert_eq!(
snapshot
.read_response_for_path("/openai/v1/videos/native-local")
.body_json["status"],
"completed"
);
let LocalVideoTaskSnapshot::OpenAi(seed) = &snapshot else {
panic!("openai task expected")
};
let Some(LocalVideoTaskContentAction::StreamPlan(download)) =
seed.build_content_stream_action(None, "download")
else {
panic!("download expected")
};
assert_eq!(download.url, "https://vidgen.x.ai/video.mp4");
assert!(download.headers.is_empty());
let stored = snapshot.to_upsert_record().into_stored();
assert!(stored.request_metadata.is_none());
assert_eq!(stored.client_api_format.as_deref(), Some("xai:video"));
let restored = LocalVideoTaskSnapshot::from_stored_task_with_transport(
&stored,
seed.transport.clone(),
)
.unwrap();
assert_eq!(restored.read_response().body_json["status"], "done");
service.record_snapshot(restored);
let poll = service
.prepare_read_refresh_sync_plan_for_user(
Some("openai"),
"/v1/videos/native-local",
"owner",
"poll",
)
.unwrap();
assert_eq!(poll.plan.url, "https://api.x.ai/v1/videos/native-upstream");
assert!(service
.prepare_read_refresh_sync_plan_for_user(
Some("openai"),
"/v1/videos/native-local",
"foreign",
"poll"
)
.is_none());
assert!(service.apply_read_refresh_projection(&poll, body.as_object().unwrap()));
assert_eq!(
service
.read_response_for_user(Some("openai"), "/v1/videos/native-local", "owner")
.unwrap()
.body_json,
body
);
}
#[test]
fn xai_video_lifecycle_creates_polls_persists_and_downloads() {
use crate::{
LocalVideoTaskContentAction, LocalVideoTaskSnapshot, VideoTaskService,
VideoTaskTruthSourceMode,
};
for api_root in ["https://cli-chat-proxy.grok.com/v1", "https://api.x.ai/v1"] {
let mut plan =
build_internal_finalize_video_plan("xai-create", "openai:video", None).unwrap();
plan.url = format!("{api_root}/videos/generations");
plan.headers
.insert("authorization".into(), "Bearer test-token".into());
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
let context = json!({"local_task_id": "local-video", "model": "grok-imagine-video", "original_request_body": {"prompt": "A cat", "seconds": "6"}});
let success = service
.prepare_sync_success(
"openai_video_create_sync_finalize",
json!({"request_id": "xai-request"}).as_object().unwrap(),
context.as_object().unwrap(),
&plan,
)
.unwrap();
assert_eq!(success.client_body_json()["id"], "local-video");
assert_eq!(success.client_body_json()["status"], "queued");
let snapshot = success.to_snapshot();
assert_eq!(
snapshot.to_upsert_record().external_task_id.as_deref(),
Some("xai-request")
);
service.record_snapshot(snapshot.clone());
let poll = service
.prepare_poll_refresh_plan_for_snapshot(snapshot, "xai-poll")
.unwrap();
assert_eq!(poll.plan.method, "GET");
assert_eq!(poll.plan.url, format!("{api_root}/videos/xai-request"));
assert_eq!(
poll.plan.headers.get("authorization"),
plan.headers.get("authorization")
);
assert!(service.apply_read_refresh_projection(
&poll,
json!({"status": "pending"}).as_object().unwrap()
));
assert_eq!(
service
.read_response(Some("openai"), "/v1/videos/local-video")
.unwrap()
.body_json["status"],
"queued"
);
assert!(service.apply_read_refresh_projection(&poll, json!({
"status": "done", "video": {"url": "https://vidgen.x.ai/result.mp4", "duration": 6}
}).as_object().unwrap()));
let snapshot = service
.snapshot_for_route(Some("openai"), "/v1/videos/local-video")
.unwrap();
assert!(!snapshot.is_active_for_refresh());
let record = snapshot.to_upsert_record();
assert_eq!(
record.status,
aether_data_contracts::repository::video_tasks::VideoTaskStatus::Completed
);
assert_eq!(
record.video_url.as_deref(),
Some("https://vidgen.x.ai/result.mp4")
);
assert_eq!(record.duration_seconds, Some(6));
let response = snapshot.read_response();
assert_eq!(response.body_json["status"], "completed");
assert_eq!(response.body_json["progress"], 100);
assert_eq!(
response.body_json["video_url"],
"https://vidgen.x.ai/result.mp4"
);
let LocalVideoTaskSnapshot::OpenAi(seed) = snapshot else {
panic!("OpenAI video expected")
};
let Some(LocalVideoTaskContentAction::StreamPlan(download)) =
seed.build_content_stream_action(None, "download")
else {
panic!("download expected")
};
assert_eq!(download.url, "https://vidgen.x.ai/result.mp4");
assert!(
download.headers.is_empty(),
"provider credentials must not be sent to the media CDN"
);
}
}
#[test]
fn xai_video_errors_are_terminal_even_without_a_status() {
use crate::{LocalVideoTaskSnapshot, VideoTaskTruthSourceMode};
let mut plan =
build_internal_finalize_video_plan("xai-create", "openai:video", None).unwrap();
plan.url = "https://cli-chat-proxy.grok.com/v1/videos/generations".into();
for body in [
json!({"code": "content_policy_violation", "error": "Rejected"}),
json!({"error": {"code": "content_policy_violation", "message": "Rejected"}}),
json!({"status": "failed", "error": "Rejected"}),
] {
let mut snapshot = VideoTaskTruthSourceMode::RustAuthoritative
.prepare_sync_success(
"openai_video_create_sync_finalize",
json!({"request_id": "xai-request"}).as_object().unwrap(),
&Default::default(),
&plan,
)
.unwrap()
.to_snapshot();
snapshot.apply_provider_body(body.as_object().unwrap());
assert!(!snapshot.is_active_for_refresh());
assert_eq!(snapshot.read_response().body_json["status"], "failed");
let LocalVideoTaskSnapshot::OpenAi(seed) = snapshot else {
panic!("OpenAI video expected")
};
assert!(seed.error_message.is_none());
}
}
#[test]
fn openai_video_id_takes_precedence_over_xai_alias() {
assert_eq!(
super::openai_video_provider_task_id(
json!({"id": "openai-id", "request_id": "trace-id"})
.as_object()
.unwrap()
),
Some("openai-id")
);
assert_eq!(
super::openai_video_provider_task_id(
json!({"id": " ", "request_id": "xai-id"})
.as_object()
.unwrap()
),
Some("xai-id")
);
assert_eq!(
super::openai_video_provider_task_id(json!({"request_id": " "}).as_object().unwrap()),
None
);
}
#[test] #[test]
fn builds_local_sync_finalize_read_response_for_supported_video_finalize_kinds() { fn builds_local_sync_finalize_read_response_for_supported_video_finalize_kinds() {
let delete_response = build_local_sync_finalize_read_response( let delete_response = build_local_sync_finalize_read_response(
@@ -71,8 +71,16 @@ impl LocalVideoTaskPersistence {
.unwrap_or_else(|| plan.request_id.clone()), .unwrap_or_else(|| plan.request_id.clone()),
username: context_text(report_context, "username"), username: context_text(report_context, "username"),
api_key_name: context_text(report_context, "api_key_name"), api_key_name: context_text(report_context, "api_key_name"),
client_api_format: context_text(report_context, "client_api_format") client_api_format: if report_context
.unwrap_or_else(|| plan.client_api_format.clone()), .get("video_client_protocol")
.and_then(Value::as_str)
== Some("xai")
{
"xai:video".to_string()
} else {
context_text(report_context, "client_api_format")
.unwrap_or_else(|| plan.client_api_format.clone())
},
provider_api_format: context_text(report_context, "provider_api_format") provider_api_format: context_text(report_context, "provider_api_format")
.unwrap_or_else(|| plan.provider_api_format.clone()), .unwrap_or_else(|| plan.provider_api_format.clone()),
original_request_body: report_context original_request_body: report_context
@@ -201,6 +201,13 @@ pub struct LocalVideoTaskPersistence {
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct OpenAiVideoTaskSeed { pub struct OpenAiVideoTaskSeed {
/// Preserve existing database identity; older snapshots derive it from the local task ID.
#[serde(default)]
pub local_short_id: Option<String>,
#[serde(default)]
pub native_response: Option<Value>,
#[serde(default)]
pub xai_provider: bool,
pub local_task_id: String, pub local_task_id: String,
pub upstream_task_id: String, pub upstream_task_id: String,
pub created_at_unix_ms: u64, pub created_at_unix_ms: u64,
+147
View File
@@ -0,0 +1,147 @@
# xAI provider behavior
The following rules preserve the provider-specific behavior of the `xai` provider
across Aether's request and transport layers.
## Responses and tools
- HTTP requests drop `previous_response_id`. Clients must supply conversation
history; this provider does not add an HTTP response-ID history store.
- `metadata.user_id` is removed. Claude clients copy it onto converted Responses
bodies and xAI rejects the field.
- Preserve requested `reasoning.encrypted_content`. On a native Responses-to-Responses
hop, keep provider-owned input items instead of rebuilding them through the canonical
format. xAI encrypted reasoning may have IDs that do not use OpenAI's `rs` prefix.
Aether's Gemini signature carriers remain excluded from xAI replay.
- The replay policy is selected from the configured provider type. A model called
`grok-*` on another provider does not opt into that policy. WebSocket continuation
metadata retains the selected policy across reconnects.
- A regular client function called `web_search` remains a function. Claude hosted
search choices are resolved against the original typed tool declaration, including
declarations with a different name.
- When only `image_generation` is allowed, keep only that tool and retain the requested
`auto` or `required` mode. For mixed allowed-tool lists, remove the image choice while
preserving the other allowed entries, as required by xAI's tool-choice schema.
- Reasoning effort is stripped for models that do not accept it.
- OpenAI-style image reference aliases in a request body are rewritten to xAI's
shape without touching chat message parts.
## Routing and credentials
OAuth requests default to `https://cli-chat-proxy.grok.com/v1`; API-key or
`using_api=true` requests default to `https://api.x.ai/v1`. Explicit custom gateways
are preserved. Compact remains on the official endpoint. CLI identity headers are
applied where the selected upstream requires them.
Account binding uses the xAI device code flow: the gateway requests a device code,
the operator authorizes it out of band, and the gateway polls for the token set.
There is no local callback listener, so headless deployments can bind accounts.
Refresh tokens can also be imported individually or in batches, and are rotated
on refresh.
Quota refresh reads `/user` and `/billing?format=credits` and stores a structured
usage snapshot. A prepaid balance keeps an account selectable after the weekly
allowance is exhausted. API-key accounts skip the subscription billing surface.
## Images and videos
OAuth media requests default to `https://cli-chat-proxy.grok.com/v1`; API-key or
`using_api=true` requests default to `https://api.x.ai/v1`. Explicit custom gateways
are preserved. Compact remains on the official endpoint. CLI identity headers are
applied to media requests and restored when a persisted video task's polling transport
is reconstructed.
Aether's OpenAI-compatible task parser accepts xAI's `request_id` creation field,
status aliases such as `pending` and `done`, nested `video.url` and `video.duration`,
and failure payloads containing `code` / `error` without a status. Existing OpenAI
`id` takes precedence. The client receives Aether's local task ID; polling uses the
upstream task ID and selected credential. Completed video downloads use the returned
media URL without forwarding provider authentication headers to the media host.
### Public video protocols
The xAI provider supports two video surfaces:
| Operation | xAI native | OpenAI compatible |
| --- | --- | --- |
| Create | `POST /v1/videos/generations` | `POST /openai/v1/videos` |
| Edit / extend | `POST /v1/videos/edits`, `POST /v1/videos/extensions` | — |
| Retrieve | `GET /v1/videos/{request_id}` | `GET /openai/v1/videos/{id}` |
| Download | use the returned `video.url` | `GET /openai/v1/videos/{id}/content` |
For xAI, `POST /v1/videos` is a native creation alias. Other providers retain
Aether's existing OpenAI-compatible `/v1/videos` behavior. xAI callers using
OpenAI `seconds` / `size` parameters must use `/openai/v1/videos`. The adapter
maps these to numeric `duration`, `aspect_ratio`, and `resolution`; it also adapts
image references. This implementation defaults to 4 seconds, portrait, and 720p,
clamps `duration` to 1-15, and validates inputs. Explicit native requests retain
native parameters and additional provider fields.
Default xAI creation targets `/videos/generations` on the selected upstream host.
Explicit custom endpoint paths still take precedence. Native generation, editing,
and extension paths only select xAI provider candidates.
Native creation returns `request_id`; native retrieval preserves `done`, nested
`video.url`, and provider fields such as `respect_moderation`. The identifier is
an opaque Aether task ID so queries remain scoped to the owning user and pinned
to the original upstream task and credential. The explicit `/openai/v1/videos`
surface projects `id`, `completed`, and `video_url`.
The task row records the native client protocol as `xai:video`, while its provider
transport remains `openai:video`. This survives restart without storing request
bodies or credentials. Raw native responses are cached only in memory; after
reconstruction the gateway refreshes from the original provider to recover its
response fields, including for completed tasks. If refreshing is unavailable,
the stored task still provides the native status and media URL projection.
OpenAI/xAI task persistence supplies a stable 16-character `short_id`, as required
by the PostgreSQL schema. Existing rows retain their original short ID across
reconstruction, including legacy embedded snapshots. This internal identifier is
separate from the opaque local task ID returned to clients; no schema change or
historical row rewrite is needed.
Task retrieval and content downloads are admitted by the production GET execution
gate. Reconstructed tasks resolve proxy nodes, system proxy defaults, tunnel affinity,
and transport profiles through the same deployment resolver used for creation;
configured proxy routes must not silently turn into direct requests after restart.
### Runtime configuration
Standalone Rust deployments must set
`AETHER_GATEWAY_VIDEO_TASK_TRUTH_SOURCE_MODE=rust-authoritative` and restart the
gateway to enable video task retrieval, polling, and content downloads. The CLI's
legacy default is `python-sync-report`: creation can return a task ID in that mode,
but the local task read/refresh paths are disabled and may return HTTP 503.
When the gateway also serves the frontend, `/openai/v1/videos` and its subpaths
must bypass the static SPA handler and be mounted as API routes. Otherwise a
successful-looking HTTP 200 response to a video query may contain `text/html`
instead of the task's JSON response. The lifecycle regression includes the static
frontend to cover this production configuration.
## Regression coverage
The format tests cover client and hosted search choices, image-only and mixed tool
restrictions, encrypted reasoning replay, image reference rewriting, and unchanged
OpenAI replay restrictions. Transport tests cover OAuth/API-key/custom routing and
media identity headers. Video-task tests exercise creation, polling, terminal
projection, persistence fields, content-download planning, and status-less errors
using local fixtures. They do not make paid generation requests.
The HTTP regression exercises all native creation paths and the compatibility
prefix through the public router and candidate planner, then checks polling,
cross-user denial, persistence, retrieval from a fresh gateway instance, and downloads
through both prefixes without leaking authorization to the media host. It uses the
real HTTP executor and a managed proxy node backed by a local test server, with no
execution-runtime override. The background poller also has a real HTTP proxy-node
regression, so production method guards and transport reconstruction are exercised.
CI also runs the same HTTP lifecycle with the PostgreSQL repository and the
production column constraints/indexes in an isolated temporary table. This catches
persistence failures that the in-memory repository cannot expose. The test uses
local `initdb`, `postgres`, and `pg_ctl` (already provided by the gateway CI job),
or an explicit `AETHER_TEST_DATABASE_URL` pointing to an isolated test database.
```sh
cargo test -p aether-ai-formats -p aether-provider-transport -p aether-video-tasks-core --lib
cargo test -p aether-gateway --lib xai
```
+1 -1
View File
@@ -376,7 +376,7 @@ function jsonValueContainsAgentIdentity(value: unknown): boolean {
export interface DeviceAuthorizeRequest { export interface DeviceAuthorizeRequest {
start_url?: string start_url?: string
region?: string region?: string
auth_type?: 'builder_id' | 'identity_center' | 'google' | 'github' | 'browser' auth_type?: 'builder_id' | 'identity_center' | 'google' | 'github' | 'browser' | 'device'
login_option?: 'google' | 'github' | 'default' login_option?: 'google' | 'github' | 'default'
redirect_uri?: string redirect_uri?: string
proxy_node_id?: string proxy_node_id?: string
+19 -1
View File
@@ -462,6 +462,23 @@ export interface GrokUpstreamMetadata {
account_user_id?: string | null account_user_id?: string | null
} }
export interface XaiUpstreamMetadata {
updated_at?: number
subscription_title?: string
usage_percentage?: number
remaining_percentage?: number
usage_label?: string
usage_limit?: number
current_usage?: number
remaining?: number
next_reset_at?: number
prepaid_balance?: number
on_demand_cap?: number
on_demand_used?: number
on_demand_remaining?: number
period_type?: string
}
export interface GeminiCliTierMetadata { export interface GeminiCliTierMetadata {
id?: string | null id?: string | null
tierType?: string | null tierType?: string | null
@@ -520,6 +537,7 @@ export interface UpstreamMetadata {
chatgpt_web?: ChatGPTWebUpstreamMetadata chatgpt_web?: ChatGPTWebUpstreamMetadata
grok?: GrokUpstreamMetadata grok?: GrokUpstreamMetadata
gemini_cli?: GeminiCliUpstreamMetadata gemini_cli?: GeminiCliUpstreamMetadata
xai?: XaiUpstreamMetadata
} }
// 按格式的健康度数据 // 按格式的健康度数据
@@ -758,7 +776,7 @@ export interface HealthRelatedMonitorResponse {
related_providers: HealthRelatedMonitor[] related_providers: HealthRelatedMonitor[]
} }
export type ProviderType = 'custom' | 'claude_code' | 'codex' | 'chatgpt_web' | 'gemini_cli' | 'antigravity' | 'kiro' | 'grok' | 'windsurf' | 'vertex_ai' export type ProviderType = 'custom' | 'claude_code' | 'codex' | 'chatgpt_web' | 'gemini_cli' | 'antigravity' | 'kiro' | 'grok' | 'xai' | 'windsurf' | 'vertex_ai'
export interface ClaudeCodeAdvancedConfig { export interface ClaudeCodeAdvancedConfig {
// 会话数量控制:null/undefined 表示不限制 // 会话数量控制:null/undefined 表示不限制
@@ -334,7 +334,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
label: 'Free/Team 优先', label: 'Free/Team 优先',
description: '兼容旧配置:优先消耗 Free、Team 或两者', description: '兼容旧配置:优先消耗 Free、Team 或两者',
evidence_hint: '依据 plan_type,保留旧 free_only/team_only/both 语义', evidence_hint: '依据 plan_type,保留旧 free_only/team_only/both 语义',
providers: ['codex', 'grok', 'kiro', 'windsurf'], providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'],
modes: [ modes: [
{ value: 'free_only', label: 'Free' }, { value: 'free_only', label: 'Free' },
{ value: 'team_only', label: 'Team' }, { value: 'team_only', label: 'Team' },
@@ -347,7 +347,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
label: 'Free 优先', label: 'Free 优先',
description: '优先消耗 Free 账号(依赖 plan_type)', description: '优先消耗 Free 账号(依赖 plan_type)',
evidence_hint: '依据 plan_type(Free 账号优先调度)', evidence_hint: '依据 plan_type(Free 账号优先调度)',
providers: ['codex', 'grok', 'kiro', 'windsurf'], providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'],
modes: null, modes: null,
default_mode: null, default_mode: null,
}, },
@@ -356,7 +356,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
label: 'Team 优先', label: 'Team 优先',
description: '优先消耗 Team 账号(依赖 plan_type)', description: '优先消耗 Team 账号(依赖 plan_type)',
evidence_hint: '依据 plan_type(Team 账号优先调度)', evidence_hint: '依据 plan_type(Team 账号优先调度)',
providers: ['codex', 'grok', 'kiro', 'windsurf'], providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'],
modes: null, modes: null,
default_mode: null, default_mode: null,
}, },
@@ -365,7 +365,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
label: 'Plus 优先', label: 'Plus 优先',
description: '优先消耗 Plus 账号(依赖 plan_type)', description: '优先消耗 Plus 账号(依赖 plan_type)',
evidence_hint: '依据 plan_type(Plus 账号优先调度)', evidence_hint: '依据 plan_type(Plus 账号优先调度)',
providers: ['codex', 'grok', 'kiro', 'windsurf'], providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'],
modes: null, modes: null,
default_mode: null, default_mode: null,
}, },
@@ -374,7 +374,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
label: 'Pro 优先', label: 'Pro 优先',
description: '优先消耗 Pro 账号(依赖 plan_type)', description: '优先消耗 Pro 账号(依赖 plan_type)',
evidence_hint: '依据 plan_type(Pro 账号优先调度)', evidence_hint: '依据 plan_type(Pro 账号优先调度)',
providers: ['codex', 'grok', 'kiro', 'windsurf'], providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'],
modes: null, modes: null,
default_mode: null, default_mode: null,
}, },
@@ -392,7 +392,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
label: '额度刷新优先', label: '额度刷新优先',
description: '优先选即将刷新额度的账号', description: '优先选即将刷新额度的账号',
evidence_hint: '依据账号额度重置倒计时(next_reset / reset_seconds)', evidence_hint: '依据账号额度重置倒计时(next_reset / reset_seconds)',
providers: ['codex', 'grok', 'kiro', 'windsurf'], providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'],
default_enabled_providers: ['codex', 'windsurf'], default_enabled_providers: ['codex', 'windsurf'],
modes: null, modes: null,
default_mode: null, default_mode: null,
@@ -244,6 +244,129 @@
</div> </div>
</template> </template>
<!-- xAI: 设备授权 -->
<template v-else-if="isXaiProvider">
<div class="space-y-3">
<div class="h-[265px]">
<div
v-if="device.status === 'error' || device.status === 'expired'"
class="rounded-xl border border-destructive/20 bg-destructive/5 p-5"
>
<div class="flex flex-col items-center text-center space-y-3">
<div class="w-10 h-10 rounded-full bg-destructive/10 flex items-center justify-center">
<AlertCircle class="w-5 h-5 text-destructive" />
</div>
<div class="space-y-1">
<p class="text-sm font-medium text-destructive">
{{ legacyT(device.status === 'expired' ? '授权已过期' : '授权失败') }}
</p>
<p class="text-xs text-muted-foreground">
{{ legacyT(device.error || '请重试') }}
</p>
</div>
<Button
size="sm"
variant="outline"
@click="resetDevice"
>
{{ legacyT('重新开始') }}
</Button>
</div>
</div>
<div
v-else-if="device.starting && !device.session_id"
class="flex items-center justify-center py-12"
>
<div class="text-center">
<div class="animate-spin rounded-full h-6 w-6 border-b-2 border-primary mx-auto mb-3" />
<p class="text-xs text-muted-foreground">
{{ legacyT('正在准备设备授权...') }}
</p>
</div>
</div>
<div
v-else-if="device.session_id && device.status === 'pending'"
class="rounded-xl border border-border bg-muted/20 p-5"
>
<div class="flex flex-col items-center text-center space-y-4">
<div class="relative">
<div class="absolute inset-0 rounded-full bg-primary/20 animate-ping" />
<div class="relative w-10 h-10 rounded-full bg-primary/10 flex items-center justify-center">
<ExternalLink class="w-5 h-5 text-primary" />
</div>
</div>
<div class="space-y-1">
<p class="text-sm font-medium">
{{ legacyT('在浏览器中输入设备码完成授权') }}
</p>
<p class="text-xs text-muted-foreground">
{{ legacyT('授权完成后此页面将自动更新') }}
</p>
</div>
<div
v-if="device.user_code"
class="flex items-center gap-2 rounded-lg border border-border bg-background px-3 py-2"
>
<span class="text-lg font-mono font-bold tracking-[0.2em]">{{ device.user_code }}</span>
<button
class="p-1 rounded hover:bg-muted transition-colors"
:title="legacyT('复制设备码')"
@click="copyToClipboard(device.user_code)"
>
<Copy class="w-3.5 h-3.5 text-muted-foreground" />
</button>
</div>
<div class="flex items-center gap-1.5 text-xs text-muted-foreground">
<div class="animate-spin rounded-full h-3 w-3 border-[1.5px] border-primary/30 border-t-primary" />
<span>{{ remainingText }}</span>
</div>
<div class="flex gap-2 w-full">
<Button
class="flex-1"
size="sm"
:disabled="!device.verification_uri_complete && !device.verification_uri"
@click="openDeviceVerificationUrl"
>
<ExternalLink class="w-3.5 h-3.5 mr-1.5" />
{{ legacyT('打开授权页面') }}
</Button>
<Button
size="sm"
variant="outline"
:disabled="!device.verification_uri_complete && !device.verification_uri"
@click="copyToClipboard(device.verification_uri_complete || device.verification_uri)"
>
<Copy class="w-3.5 h-3.5" />
</Button>
</div>
</div>
</div>
<div
v-else
class="flex h-full flex-col items-center justify-center gap-3"
>
<p class="text-xs text-muted-foreground text-center">
{{ legacyT('使用 xAI 设备授权登录 Grok CLI,或改为导入 API Key / Refresh Token。') }}
</p>
<Button
class="w-full"
:disabled="device.starting"
@click="startDeviceAuth"
>
{{ device.starting ? legacyT('正在准备授权...') : legacyT('开始授权') }}
</Button>
</div>
</div>
</div>
</template>
<!-- Kiro: 设备授权模式 --> <!-- Kiro: 设备授权模式 -->
<template v-else-if="isKiroProvider"> <template v-else-if="isKiroProvider">
<div class="space-y-3"> <div class="space-y-3">
@@ -994,7 +1117,7 @@ let oauthInitRequestId = 0
let oauthCompleteRequestId = 0 let oauthCompleteRequestId = 0
// 设备授权状态 // 设备授权状态
type DeviceAuthType = 'default' | 'google' | 'github' | 'builder_id' | 'identity_center' type DeviceAuthType = 'default' | 'google' | 'github' | 'builder_id' | 'identity_center' | 'device'
type WindsurfLoginOption = 'default' | 'google' | 'github' type WindsurfLoginOption = 'default' | 'google' | 'github'
interface DeviceAuthState { interface DeviceAuthState {
@@ -1075,9 +1198,10 @@ const isOpen = computed(() => props.open)
const isKiroProvider = computed(() => (props.providerType || '').toLowerCase() === 'kiro') const isKiroProvider = computed(() => (props.providerType || '').toLowerCase() === 'kiro')
const isGrokProvider = computed(() => (props.providerType || '').toLowerCase() === 'grok') const isGrokProvider = computed(() => (props.providerType || '').toLowerCase() === 'grok')
const isWindsurfProvider = computed(() => (props.providerType || '').toLowerCase() === 'windsurf') const isWindsurfProvider = computed(() => (props.providerType || '').toLowerCase() === 'windsurf')
const isXaiProvider = computed(() => (props.providerType || '').toLowerCase() === 'xai')
const isCodexProvider = computed(() => (props.providerType || '').toLowerCase() === 'codex') const isCodexProvider = computed(() => (props.providerType || '').toLowerCase() === 'codex')
const isClaudeCodeProvider = computed(() => (props.providerType || '').toLowerCase() === 'claude_code') const isClaudeCodeProvider = computed(() => (props.providerType || '').toLowerCase() === 'claude_code')
const isDeviceBrowserProvider = computed(() => isKiroProvider.value || isWindsurfProvider.value) const isDeviceBrowserProvider = computed(() => isKiroProvider.value || isWindsurfProvider.value || isXaiProvider.value)
const showAuthorizationMode = computed(() => !isGrokProvider.value) const showAuthorizationMode = computed(() => !isGrokProvider.value)
const defaultMode = computed<DialogMode>(() => (isGrokProvider.value ? 'import' : 'oauth')) const defaultMode = computed<DialogMode>(() => (isGrokProvider.value ? 'import' : 'oauth'))
@@ -1101,7 +1225,7 @@ const isManualDeviceCallbackPending = computed(() =>
const authorizationModeLabel = computed(() => { const authorizationModeLabel = computed(() => {
if (isWindsurfProvider.value) return legacyT('浏览器登录') if (isWindsurfProvider.value) return legacyT('浏览器登录')
if (isDeviceBrowserProvider.value) return legacyT('设备授权') if (isXaiProvider.value || isDeviceBrowserProvider.value) return legacyT('设备授权')
return legacyT('获取授权') return legacyT('获取授权')
}) })
@@ -1215,6 +1339,9 @@ const importManualPlaceholder = computed(() => {
if (isClaudeCodeProvider.value) { if (isClaudeCodeProvider.value) {
return legacyT('粘贴 Claude Refresh Token 或 Claude Code .credentials.json 内容') return legacyT('粘贴 Claude Refresh Token 或 Claude Code .credentials.json 内容')
} }
if (isXaiProvider.value) {
return legacyT('粘贴 xAI API Key、Access Token,或包含 refresh_token / api_key 的 JSON')
}
if (isWindsurfProvider.value) { if (isWindsurfProvider.value) {
return legacyT('粘贴 show-auth-token Token、API key 或 JSON 内容') return legacyT('粘贴 show-auth-token Token、API key 或 JSON 内容')
} }
@@ -1477,11 +1604,17 @@ function resetDevice() {
totp.stop() totp.stop()
const { auth_type, start_url, region, totp_secret } = device.value const { auth_type, start_url, region, totp_secret } = device.value
device.value = createInitialDeviceState() device.value = createInitialDeviceState()
device.value.auth_type = isWindsurfProvider.value ? (auth_type === 'google' || auth_type === 'github' ? auth_type : 'default') : auth_type device.value.auth_type = isXaiProvider.value
? 'device'
: isWindsurfProvider.value
? (auth_type === 'google' || auth_type === 'github' ? auth_type : 'default')
: auth_type
device.value.start_url = start_url device.value.start_url = start_url
device.value.region = region device.value.region = region
device.value.totp_secret = totp_secret device.value.totp_secret = totp_secret
if (!isWindsurfProvider.value && (device.value.auth_type === 'google' || device.value.auth_type === 'github')) { if (isXaiProvider.value) {
void ensureXaiDeviceAuth()
} else if (!isWindsurfProvider.value && (device.value.auth_type === 'google' || device.value.auth_type === 'github')) {
void ensureKiroSocialDeviceAuth() void ensureKiroSocialDeviceAuth()
} }
} }
@@ -1500,6 +1633,8 @@ function resetForm() {
device.value = createInitialDeviceState() device.value = createInitialDeviceState()
if (isWindsurfProvider.value) { if (isWindsurfProvider.value) {
device.value.auth_type = 'default' device.value.auth_type = 'default'
} else if (isXaiProvider.value) {
device.value.auth_type = 'device'
} }
importText.value = '' importText.value = ''
importing.value = false importing.value = false
@@ -1531,6 +1666,8 @@ function switchMode(newMode: DialogMode) {
if (newMode === 'oauth') { if (newMode === 'oauth') {
if (isKiroProvider.value) { if (isKiroProvider.value) {
void ensureKiroSocialDeviceAuth() void ensureKiroSocialDeviceAuth()
} else if (isXaiProvider.value) {
void ensureXaiDeviceAuth()
} else if (!oauth.value.authorization_url && !oauth.value.starting) { } else if (!oauth.value.authorization_url && !oauth.value.starting) {
initOAuth() initOAuth()
} }
@@ -1833,6 +1970,29 @@ function parseImportText(text: string): {
return { refresh_token: trimmed } return { refresh_token: trimmed }
} }
if (isXaiProvider.value) {
try {
const parsed: unknown = JSON.parse(trimmed)
if (typeof parsed === 'object' && parsed !== null) {
const obj = parsed as Record<string, unknown>
const apiKey = normalizeStringField(obj.api_key) ?? normalizeStringField(obj.apiKey)
const refreshToken = normalizeStringField(obj.refresh_token) ?? normalizeStringField(obj.refreshToken)
const accessToken = normalizeStringField(obj.access_token) ?? normalizeStringField(obj.accessToken) ?? apiKey
if (refreshToken || accessToken) {
return {
refresh_token: refreshToken,
access_token: accessToken,
name: normalizeStringField(obj.name) ?? normalizeStringField(obj.email),
email: normalizeStringField(obj.email),
}
}
}
} catch {
// Raw xAI API keys / access tokens are imported as access_token.
}
return { access_token: trimmed }
}
if (isGrokProvider.value) { if (isGrokProvider.value) {
const cookieImport = parseGrokCookieImport(trimmed) const cookieImport = parseGrokCookieImport(trimmed)
if (cookieImport) { if (cookieImport) {
@@ -2367,17 +2527,20 @@ async function startDeviceAuth() {
device.value.error = '' device.value.error = ''
try { try {
const isWindsurf = isWindsurfProvider.value const isWindsurf = isWindsurfProvider.value
const isXai = isXaiProvider.value
const isBuilderID = requestedAuthType === 'builder_id' const isBuilderID = requestedAuthType === 'builder_id'
const isSocial = requestedAuthType === 'google' || requestedAuthType === 'github' const isSocial = !isXai && (requestedAuthType === 'google' || requestedAuthType === 'github')
const windsurfLoginOption: WindsurfLoginOption = isSocial ? requestedAuthType : 'default' const windsurfLoginOption: WindsurfLoginOption = isSocial ? requestedAuthType : 'default'
const authTypeForRequest = isWindsurf const authTypeForRequest = isWindsurf
? 'browser' ? 'browser'
: (requestedAuthType === 'default' ? 'google' : requestedAuthType) : isXai
? 'device'
: (requestedAuthType === 'default' ? 'google' : requestedAuthType)
const resp = await startDeviceAuthorize(props.providerId, { const resp = await startDeviceAuthorize(props.providerId, {
auth_type: authTypeForRequest, auth_type: authTypeForRequest,
login_option: isWindsurf ? windsurfLoginOption : undefined, login_option: isWindsurf ? windsurfLoginOption : undefined,
start_url: isWindsurf ? undefined : (isBuilderID ? BUILDER_ID_START_URL : (isSocial ? undefined : (device.value.start_url.trim() || undefined))), start_url: (isWindsurf || isXai) ? undefined : (isBuilderID ? BUILDER_ID_START_URL : (isSocial ? undefined : (device.value.start_url.trim() || undefined))),
region: isWindsurf ? undefined : (isBuilderID || isSocial ? BUILDER_ID_REGION : (device.value.region.trim() || undefined)), region: (isWindsurf || isXai) ? undefined : (isBuilderID || isSocial ? BUILDER_ID_REGION : (device.value.region.trim() || undefined)),
proxy_node_id: selectedProxyNodeId.value || undefined, proxy_node_id: selectedProxyNodeId.value || undefined,
}) })
if (requestId !== deviceAuthRequestId || device.value.auth_type !== requestedAuthType) return if (requestId !== deviceAuthRequestId || device.value.auth_type !== requestedAuthType) return
@@ -2417,6 +2580,13 @@ async function ensureKiroSocialDeviceAuth() {
await startDeviceAuth() await startDeviceAuth()
} }
async function ensureXaiDeviceAuth() {
if (!props.open || !props.providerId || !isXaiProvider.value) return
if (device.value.starting) return
if (device.value.session_id && (device.value.status === 'pending' || device.value.status === 'authorized')) return
await startDeviceAuth()
}
function scheduleDevicePoll() { function scheduleDevicePoll() {
if (devicePollTimer) clearTimeout(devicePollTimer) if (devicePollTimer) clearTimeout(devicePollTimer)
devicePollTimer = setTimeout(() => pollDevice(), device.value.interval * 1000) devicePollTimer = setTimeout(() => pollDevice(), device.value.interval * 1000)
@@ -2525,6 +2695,9 @@ watch(
} }
if (isWindsurfProvider.value) { if (isWindsurfProvider.value) {
device.value.auth_type = 'default' device.value.auth_type = 'default'
} else if (isXaiProvider.value) {
device.value.auth_type = 'device'
void ensureXaiDeviceAuth()
} else if (isKiroProvider.value) { } else if (isKiroProvider.value) {
void ensureKiroSocialDeviceAuth() void ensureKiroSocialDeviceAuth()
} else { } else {
@@ -2554,6 +2727,9 @@ watch(
device.value.auth_type = ['default', 'google', 'github'].includes(device.value.auth_type) device.value.auth_type = ['default', 'google', 'github'].includes(device.value.auth_type)
? device.value.auth_type ? device.value.auth_type
: 'default' : 'default'
} else if (props.open && isXaiProvider.value && mode.value === 'oauth') {
device.value.auth_type = 'device'
void ensureXaiDeviceAuth()
} else if (props.open && isKiroProvider.value && mode.value === 'oauth') { } else if (props.open && isKiroProvider.value && mode.value === 'oauth') {
void ensureKiroSocialDeviceAuth() void ensureKiroSocialDeviceAuth()
} }
@@ -514,6 +514,65 @@
</div> </div>
</template> </template>
</div> </div>
<!-- xAI / Grok Build 订阅额度 -->
<div
v-if="provider.provider_type === 'xai' && hasXaiQuotaDisplayData(key)"
class="mt-2 p-2 rounded-md bg-muted/30"
>
<ProviderQuotaSectionHeader
:title="legacyT('账号配额')"
:loading="refreshingQuota"
:updated-text="getXaiQuotaDisplay(key)?.updated_at ? formatKiroUpdatedAt(getXaiQuotaDisplay(key)?.updated_at || 0) : null"
/>
<div class="space-y-2">
<ProviderQuotaProgressRow
v-if="getXaiQuotaDisplay(key)?.usage_percentage !== undefined || getXaiQuotaDisplay(key)?.remaining_percentage !== undefined"
:label="legacyT(getXaiUsageLabel(key))"
:used-percent="getXaiUsedPercent(key)"
:remaining-percent="getXaiRemainingPercent(key)"
:meter-class="getQuotaRemainingClass(getXaiUsedPercent(key))"
:bar-class="getQuotaRemainingBarColor(getXaiUsedPercent(key))"
:reset-text="getXaiQuotaDisplay(key)?.next_reset_at
? `${formatKiroResetTime(getXaiQuotaDisplay(key)?.next_reset_at)}${legacyT('重置')}`
: null"
>
<template
v-if="getXaiQuotaDisplay(key)?.usage_limit != null"
#footer
>
<div class="flex items-center justify-between text-[9px] text-muted-foreground/70 mt-0.5">
<span>
{{ formatKiroUsage(getXaiQuotaDisplay(key)?.current_usage) }} /
{{ formatKiroUsage(getXaiQuotaDisplay(key)?.usage_limit) }}
</span>
<span v-if="getXaiQuotaDisplay(key)?.next_reset_at">
{{ formatKiroResetTime(getXaiQuotaDisplay(key)?.next_reset_at) }}{{ legacyT('重置') }}
</span>
</div>
</template>
</ProviderQuotaProgressRow>
<div
v-if="getXaiQuotaDisplay(key)?.prepaid_balance != null"
class="text-[10px] text-muted-foreground"
>
{{ legacyT('预付额度') }}: {{ formatKiroUsage(getXaiQuotaDisplay(key)?.prepaid_balance) }}
</div>
<ProviderQuotaProgressRow
v-if="getXaiQuotaDisplay(key)?.on_demand_cap"
:label="legacyT('按需额度')"
:used-percent="getXaiOnDemandUsedPercent(key)"
:meter-class="getQuotaRemainingClass(getXaiOnDemandUsedPercent(key))"
:bar-class="getQuotaRemainingBarColor(getXaiOnDemandUsedPercent(key))"
>
<template #footer>
<div class="text-[9px] text-muted-foreground/70 mt-0.5">
{{ formatKiroUsage(getXaiQuotaDisplay(key)?.on_demand_used) }} /
{{ formatKiroUsage(getXaiQuotaDisplay(key)?.on_demand_cap) }}
</div>
</template>
</ProviderQuotaProgressRow>
</div>
</div>
<!-- Windsurf 上游额度信息 --> <!-- Windsurf 上游额度信息 -->
<div <div
v-if="provider.provider_type === 'windsurf' && (hasWindsurfQuotaDisplayData(key) || isWindsurfUnavailableKey(key) || isWindsurfExhaustedKey(key))" v-if="provider.provider_type === 'windsurf' && (hasWindsurfQuotaDisplayData(key) || isWindsurfUnavailableKey(key) || isWindsurfExhaustedKey(key))"
@@ -1005,6 +1064,7 @@ import type {
GrokUpstreamMetadata, GrokUpstreamMetadata,
KiroUpstreamMetadata, KiroUpstreamMetadata,
WindsurfUpstreamMetadata, WindsurfUpstreamMetadata,
XaiUpstreamMetadata,
QuotaResetCreditsSnapshot, QuotaResetCreditsSnapshot,
QuotaStatusSnapshot, QuotaStatusSnapshot,
QuotaWindowSnapshot, QuotaWindowSnapshot,
@@ -1858,7 +1918,7 @@ function quotaSnapshotHasDisplayData(quota: QuotaStatusSnapshot | null | undefin
function getQuotaSnapshotForProvider( function getQuotaSnapshotForProvider(
key: EndpointAPIKey, key: EndpointAPIKey,
providerType: 'codex' | 'kiro' | 'windsurf' | 'antigravity' | 'chatgpt_web' | 'gemini_cli' | 'grok', providerType: 'codex' | 'kiro' | 'windsurf' | 'antigravity' | 'chatgpt_web' | 'gemini_cli' | 'grok' | 'xai',
): QuotaStatusSnapshot | null { ): QuotaStatusSnapshot | null {
const quota = key.status_snapshot?.quota const quota = key.status_snapshot?.quota
if (!quota) return null if (!quota) return null
@@ -2179,6 +2239,90 @@ function hasKiroQuotaDisplayData(key: EndpointAPIKey): boolean {
return !!kiro && (kiro.usage_percentage !== undefined || kiro.usage_limit !== undefined) return !!kiro && (kiro.usage_percentage !== undefined || kiro.usage_limit !== undefined)
} }
function getXaiQuotaDisplay(key: EndpointAPIKey): XaiUpstreamMetadata | null {
const quota = getQuotaSnapshotForProvider(key, 'xai')
if (!quota) return null
const display: XaiUpstreamMetadata = {}
const updatedAt = getQuotaSnapshotUpdatedAt(quota)
if (updatedAt !== undefined) display.updated_at = updatedAt
if (quota.plan_type) display.subscription_title = quota.plan_type
const usageWindow =
getQuotaWindow(quota, 'usage')
?? getQuotaWindowByScope(quota, 'account')[0]
?? null
if (usageWindow) {
const usedPercent = getQuotaWindowUsedPercent(usageWindow)
const remainingPercent = getQuotaWindowRemainingPercent(usageWindow)
if (usedPercent !== undefined) display.usage_percentage = usedPercent
if (remainingPercent !== undefined) display.remaining_percentage = remainingPercent
const usageLabel = String(usageWindow.label || '').trim()
if (usageLabel) display.usage_label = usageLabel
if (typeof usageWindow.used_value === 'number') display.current_usage = usageWindow.used_value
if (typeof usageWindow.limit_value === 'number') display.usage_limit = usageWindow.limit_value
if (typeof usageWindow.remaining_value === 'number') display.remaining = usageWindow.remaining_value
const nextResetAt =
getQuotaWindowResetAt(usageWindow)
?? (() => {
const resetSeconds = getQuotaWindowResetSeconds(usageWindow)
if (updatedAt === undefined || resetSeconds === undefined) return undefined
return updatedAt + resetSeconds
})()
if (nextResetAt !== undefined) display.next_reset_at = nextResetAt
}
const prepaidWindow = getQuotaWindow(quota, 'prepaid')
if (typeof prepaidWindow?.remaining_value === 'number') {
display.prepaid_balance = prepaidWindow.remaining_value
}
const onDemandWindow = getQuotaWindow(quota, 'on_demand')
if (typeof onDemandWindow?.limit_value === 'number') display.on_demand_cap = onDemandWindow.limit_value
if (typeof onDemandWindow?.used_value === 'number') display.on_demand_used = onDemandWindow.used_value
if (typeof onDemandWindow?.remaining_value === 'number') display.on_demand_remaining = onDemandWindow.remaining_value
return Object.keys(display).length > 0 ? display : null
}
function hasXaiQuotaDisplayData(key: EndpointAPIKey): boolean {
const xai = getXaiQuotaDisplay(key)
return !!xai && (
xai.usage_percentage !== undefined
|| xai.remaining_percentage !== undefined
|| xai.prepaid_balance !== undefined
|| xai.on_demand_cap !== undefined
)
}
function getXaiUsageLabel(key: EndpointAPIKey): string {
const display = getXaiQuotaDisplay(key)
if (display?.usage_label) return display.usage_label
const title = display?.subscription_title
return title ? `使用额度 (${title})` : '使用额度'
}
function getXaiUsedPercent(key: EndpointAPIKey): number {
return Math.min(Math.max(100 - getXaiRemainingPercent(key), 0), 100)
}
function getXaiRemainingPercent(key: EndpointAPIKey): number {
const xai = getXaiQuotaDisplay(key)
if (xai?.remaining_percentage != null && Number.isFinite(xai.remaining_percentage)) {
return Math.min(Math.max(xai.remaining_percentage, 0), 100)
}
if (xai?.usage_percentage != null && Number.isFinite(xai.usage_percentage)) {
return Math.min(Math.max(100 - xai.usage_percentage, 0), 100)
}
return 0
}
function getXaiOnDemandUsedPercent(key: EndpointAPIKey): number {
const xai = getXaiQuotaDisplay(key)
if (!xai?.on_demand_cap || xai.on_demand_cap <= 0) return 0
return Math.max(Math.min(((xai.on_demand_used || 0) / xai.on_demand_cap) * 100, 100), 0)
}
type GrokQuotaDisplay = GrokUpstreamMetadata & { type GrokQuotaDisplay = GrokUpstreamMetadata & {
usage_percentage?: number usage_percentage?: number
usage_limit?: number usage_limit?: number
@@ -2696,6 +2840,28 @@ function shouldAutoRefreshGrokQuota(): boolean {
return false return false
} }
function shouldAutoRefreshXaiQuota(): boolean {
if (provider.value?.provider_type !== 'xai') return false
const now = Math.floor(Date.now() / 1000)
for (const { key } of allKeys.value) {
if (!key.is_active) continue
if (isTokenExpiringSoon(key, now)) return true
if (!hasXaiQuotaDisplayData(key)) {
return true
}
const updatedAt = getXaiQuotaDisplay(key)?.updated_at
if (typeof updatedAt !== 'number' || (now - updatedAt) > AUTO_QUOTA_REFRESH_STALE_SECONDS) {
return true
}
}
return false
}
function shouldAutoRefreshWindsurfQuota(): boolean { function shouldAutoRefreshWindsurfQuota(): boolean {
if (provider.value?.provider_type !== 'windsurf') return false if (provider.value?.provider_type !== 'windsurf') return false
const now = Math.floor(Date.now() / 1000) const now = Math.floor(Date.now() / 1000)
@@ -2824,7 +2990,7 @@ async function autoRefreshQuotaInBackground(): Promise<boolean> {
if (refreshingQuota.value) return false if (refreshingQuota.value) return false
const providerType = provider.value?.provider_type const providerType = provider.value?.provider_type
if (providerType !== 'codex' && providerType !== 'gemini_cli' && providerType !== 'antigravity' && providerType !== 'kiro' && providerType !== 'windsurf' && providerType !== 'chatgpt_web' && providerType !== 'grok') return false if (providerType !== 'codex' && providerType !== 'gemini_cli' && providerType !== 'antigravity' && providerType !== 'kiro' && providerType !== 'windsurf' && providerType !== 'chatgpt_web' && providerType !== 'grok' && providerType !== 'xai') return false
// 检查是否需要刷新 // 检查是否需要刷新
let shouldRefresh = false let shouldRefresh = false
@@ -2838,6 +3004,8 @@ async function autoRefreshQuotaInBackground(): Promise<boolean> {
shouldRefresh = shouldAutoRefreshKiroQuota() shouldRefresh = shouldAutoRefreshKiroQuota()
} else if (providerType === 'grok') { } else if (providerType === 'grok') {
shouldRefresh = shouldAutoRefreshGrokQuota() shouldRefresh = shouldAutoRefreshGrokQuota()
} else if (providerType === 'xai') {
shouldRefresh = shouldAutoRefreshXaiQuota()
} else if (providerType === 'windsurf') { } else if (providerType === 'windsurf') {
shouldRefresh = shouldAutoRefreshWindsurfQuota() shouldRefresh = shouldAutoRefreshWindsurfQuota()
} else if (providerType === 'chatgpt_web') { } else if (providerType === 'chatgpt_web') {
@@ -2856,6 +3024,8 @@ async function autoRefreshQuotaInBackground(): Promise<boolean> {
hadCachedQuota = allKeys.value.some(({ key }) => key.is_active && hasKiroQuotaDisplayData(key)) hadCachedQuota = allKeys.value.some(({ key }) => key.is_active && hasKiroQuotaDisplayData(key))
} else if (providerType === 'grok') { } else if (providerType === 'grok') {
hadCachedQuota = allKeys.value.some(({ key }) => key.is_active && hasGrokQuotaDisplayData(key)) hadCachedQuota = allKeys.value.some(({ key }) => key.is_active && hasGrokQuotaDisplayData(key))
} else if (providerType === 'xai') {
hadCachedQuota = allKeys.value.some(({ key }) => key.is_active && hasXaiQuotaDisplayData(key))
} else if (providerType === 'windsurf') { } else if (providerType === 'windsurf') {
hadCachedQuota = allKeys.value.some(({ key }) => key.is_active && hasWindsurfQuotaDisplayData(key)) hadCachedQuota = allKeys.value.some(({ key }) => key.is_active && hasWindsurfQuotaDisplayData(key))
} else if (providerType === 'chatgpt_web') { } else if (providerType === 'chatgpt_web') {
@@ -60,6 +60,9 @@
<SelectItem value="grok"> <SelectItem value="grok">
Grok Grok
</SelectItem> </SelectItem>
<SelectItem value="xai">
xAI
</SelectItem>
<SelectItem value="kiro"> <SelectItem value="kiro">
Kiro Kiro
</SelectItem> </SelectItem>
@@ -93,6 +96,9 @@
<SelectItem value="grok"> <SelectItem value="grok">
Grok Grok
</SelectItem> </SelectItem>
<SelectItem value="xai">
xAI
</SelectItem>
<SelectItem value="kiro"> <SelectItem value="kiro">
Kiro Kiro
</SelectItem> </SelectItem>
@@ -54,6 +54,24 @@ describe('provider quota display components', () => {
unmount() unmount()
}) })
it('fills the remaining bar even when used percent is zero', () => {
const { root, unmount } = mount(ProviderQuotaProgressRow, {
label: '周额度',
usedPercent: 0,
remainingPercent: 86,
meterClass: 'text-green-600',
barClass: 'bg-green-500',
resetText: '5天0小时后重置',
})
expect(root.querySelector('[data-testid="provider-quota-progress-meter"]')?.textContent?.trim()).toBe('86.0%')
expect((root.querySelector('[data-testid="provider-quota-progress-bar"]') as HTMLElement).style.width).toBe('86%')
expect(root.textContent).toContain('周额度')
expect(root.querySelector('[data-testid="provider-quota-progress-reset"]')?.textContent).toBe('5天0小时后重置')
unmount()
})
it('renders section loading and updated state', () => { it('renders section loading and updated state', () => {
const Probe = defineComponent({ const Probe = defineComponent({
setup() { setup() {
@@ -1386,6 +1386,7 @@ function formatAuthType(authType: string): string {
if (lowered === 'antigravity') return 'Antigravity OAuth' if (lowered === 'antigravity') return 'Antigravity OAuth'
if (lowered === 'kiro') return 'Kiro OAuth' if (lowered === 'kiro') return 'Kiro OAuth'
if (lowered === 'grok') return 'Grok OAuth' if (lowered === 'grok') return 'Grok OAuth'
if (lowered === 'xai') return 'xAI OAuth'
return authType return authType
} }
@@ -39,10 +39,12 @@ const MODEL_TEST_OAUTH_INHERITS_PROVIDER_FORMATS = new Set([
'vertex_ai', 'vertex_ai',
'antigravity', 'antigravity',
'kiro', 'kiro',
'xai',
]) ])
const MODEL_TEST_BEARER_INHERITS_PROVIDER_FORMATS = new Set([ const MODEL_TEST_BEARER_INHERITS_PROVIDER_FORMATS = new Set([
'chatgpt_web', 'chatgpt_web',
'xai',
]) ])
const MODEL_TEST_DIAGNOSTIC_LABELS: Record<string, string> = { const MODEL_TEST_DIAGNOSTIC_LABELS: Record<string, string> = {
@@ -16,6 +16,12 @@ describe('providerTypeUtils', () => {
expect(isKeyManagedProviderType('grok')).toBe(false) expect(isKeyManagedProviderType('grok')).toBe(false)
}) })
it('treats xAI as an OAuth account provider', () => {
expect(isOAuthAccountProviderType('xai')).toBe(true)
expect(isOAuthAccountProviderType('xAI')).toBe(true)
expect(isKeyManagedProviderType('xai')).toBe(false)
})
it('treats Windsurf as an OAuth account provider', () => { it('treats Windsurf as an OAuth account provider', () => {
expect(isOAuthAccountProviderType('windsurf')).toBe(true) expect(isOAuthAccountProviderType('windsurf')).toBe(true)
expect(isOAuthAccountProviderType('Windsurf')).toBe(true) expect(isOAuthAccountProviderType('Windsurf')).toBe(true)
@@ -12,6 +12,7 @@ const oauthAccountProviderTypes = new Set([
'antigravity', 'antigravity',
'kiro', 'kiro',
'grok', 'grok',
'xai',
'windsurf', 'windsurf',
]) ])

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