diff --git a/Cargo.lock b/Cargo.lock index d83fc53c1..e475509d2 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -748,6 +748,7 @@ dependencies = [ "async-trait", "serde", "serde_json", + "sha2", "url", "uuid", ] diff --git a/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/request.rs b/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/request.rs index 41008896e..582c23342 100644 --- a/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/request.rs +++ b/apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/request.rs @@ -583,6 +583,11 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts( source_model, codex_model_capabilities.as_ref(), ); + crate::ai_serving::transport::xai::insert_cli_identity_headers_if_needed( + transport.as_ref(), + prepared.provider_api_format.as_str(), + &mut provider_request_headers, + ); request_identity_response_encoding_when_redacted( &mut provider_request_headers, redaction.redacted, diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/image/request.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/image/request.rs index 3ca11f0d2..be3ff2fc0 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/image/request.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/image/request.rs @@ -17,8 +17,8 @@ use crate::ai_serving::transport::{ ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput, GROK_CHAT_PATH, }; use crate::ai_serving::{ - apply_codex_openai_special_headers, build_chatgpt_web_image_request_body, - build_codex_openai_image_api_provider_request_body, + apply_codex_openai_special_headers, apply_xai_upstream_payload_edits, + build_chatgpt_web_image_request_body, build_codex_openai_image_api_provider_request_body, build_gemini_image_request_body_from_openai_image_request, build_openai_image_api_provider_request_body, build_openai_image_provider_request_body, default_model_for_openai_image_operation, normalize_openai_image_request, @@ -211,7 +211,7 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts( upstream_is_stream, ) }; - let Some(provider_request_body) = provider_request_body else { + let Some(mut provider_request_body) = provider_request_body else { mark_skipped_local_openai_image_candidate_with_failure_diagnostic( state, input, @@ -229,6 +229,11 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts( .await; return None; }; + apply_xai_upstream_payload_edits( + &mut provider_request_body, + transport.provider.provider_type.as_str(), + provider_api_format, + ); let Some(mut provider_request_headers) = (if is_grok { build_grok_browser_headers(GrokHeaderInput { transport, diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/video/decision.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/video/decision.rs index c1207443e..8120ca3d7 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/video/decision.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/video/decision.rs @@ -8,6 +8,7 @@ use crate::ai_serving::planner::{ build_ai_execution_decision_response, resolve_transport_request_encoding_policy, AiExecutionDecisionResponseParts, }; +use crate::ai_serving::transport::xai::video::is_native_video_request; use crate::ai_serving::transport::{ resolve_transport_execution_timeouts, resolve_transport_profile, }; @@ -33,7 +34,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat let Some(resolved) = resolve_local_video_create_candidate_payload_parts( state, parts, body_json, trace_id, input, &attempt, spec, ) - .await + .await? else { return Ok(None); }; @@ -52,9 +53,32 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat .await; let transport_profile = resolve_transport_profile(&transport); let mut extra_fields = serde_json::Map::new(); + if is_native_video_request(&transport.provider.provider_type, parts.uri.path()) { + extra_fields.insert( + "video_client_protocol".to_string(), + serde_json::json!("xai"), + ); + } + if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) { extra_fields.insert("proxy".to_string(), proxy_value); } + if transport.provider.provider_type.eq_ignore_ascii_case("xai") { + extra_fields.insert("video_provider_xai".into(), serde_json::json!(true)); + if let Some(duration) = resolved.provider_request_body.get("duration") { + extra_fields.insert("video_duration".into(), duration.clone()); + } + if parts.uri.path() == "/openai/v1/videos" { + extra_fields.insert( + "video_size".into(), + body_json + .get("size") + .filter(|v| v.as_str().is_some_and(|s| !s.trim().is_empty())) + .cloned() + .unwrap_or_else(|| serde_json::json!("720x1280")), + ); + } + } let effective_headers = input.effective_headers(&parts.headers); let report_context = build_local_execution_report_context(LocalExecutionReportContextParts { auth_context: &input.auth_context, diff --git a/apps/aether-gateway/src/ai_serving/planner/specialized/video/request.rs b/apps/aether-gateway/src/ai_serving/planner/specialized/video/request.rs index 795384f94..5910ba566 100644 --- a/apps/aether-gateway/src/ai_serving/planner/specialized/video/request.rs +++ b/apps/aether-gateway/src/ai_serving/planner/specialized/video/request.rs @@ -3,15 +3,23 @@ use std::sync::Arc; use serde_json::Value; -use crate::ai_serving::planner::candidate_preparation::resolve_candidate_mapped_model; +use crate::ai_serving::planner::candidate_preparation::{ + prepare_header_authenticated_candidate, resolve_candidate_mapped_model, OauthPreparationContext, +}; use crate::ai_serving::planner::spec_metadata::local_video_create_spec_metadata; +use crate::ai_serving::transport::xai::video::{ + convert_openai_video_request, is_explicit_native_video_path, is_native_video_request, +}; use crate::ai_serving::transport::{ build_video_create_headers, build_video_create_request_body, build_video_create_upstream_url, resolve_video_create_auth, video_create_transport_unsupported_reason, ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput, }; -use crate::ai_serving::{CandidateFailureDiagnostic, GatewayProviderTransportSnapshot}; -use crate::AppState; +use crate::ai_serving::{ + apply_xai_upstream_payload_edits, CandidateFailureDiagnostic, GatewayProviderTransportSnapshot, + PlannerAppState, +}; +use crate::{AppState, GatewayError}; use super::support::{ mark_skipped_local_video_candidate, mark_skipped_local_video_candidate_with_failure_diagnostic, @@ -37,11 +45,16 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts( input: &LocalVideoCreateDecisionInput, attempt: &LocalVideoCreateCandidateAttempt, spec: LocalVideoCreateSpec, -) -> Option { +) -> Result, GatewayError> { let spec_metadata = local_video_create_spec_metadata(spec); let candidate = &attempt.eligible.candidate; let transport = &attempt.eligible.transport; let effective_headers = input.effective_headers(&parts.headers); + if is_explicit_native_video_path(parts.uri.path()) + && !transport.provider.provider_type.eq_ignore_ascii_case("xai") + { + return Ok(None); + } let provider_family = provider_video_create_family(spec.family); let transport_unsupported_reason = video_create_transport_unsupported_reason( @@ -60,23 +73,39 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts( skip_reason, ) .await; - return None; + return Ok(None); } - let auth = resolve_video_create_auth(transport, provider_family); - let Some((auth_header, auth_value)) = auth else { - mark_skipped_local_video_candidate( - state, - input, + let prepared_candidate = match prepare_header_authenticated_candidate( + PlannerAppState::new(state), + transport, + candidate, + resolve_video_create_auth(transport, provider_family), + OauthPreparationContext { trace_id, - candidate, - attempt.candidate_index, - &attempt.candidate_id, - "transport_auth_unavailable", - ) - .await; - return None; + api_format: spec_metadata.api_format, + operation: "video_create_candidate_request", + }, + ) + .await + { + Ok(prepared) => prepared, + Err(skip_reason) => { + mark_skipped_local_video_candidate( + state, + input, + trace_id, + candidate, + attempt.candidate_index, + &attempt.candidate_id, + skip_reason, + ) + .await; + return Ok(None); + } }; + let auth_header = prepared_candidate.auth_header; + let auth_value = prepared_candidate.auth_value; let mapped_model = match resolve_candidate_mapped_model(candidate) { Ok(mapped_model) => mapped_model, @@ -91,7 +120,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts( skip_reason, ) .await; - return None; + return Ok(None); } }; @@ -117,10 +146,10 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts( ), ) .await; - return None; + return Ok(None); }; - let Some(provider_request_body) = build_video_create_request_body( + let Some(mut provider_request_body) = build_video_create_request_body( body_json, provider_family, &mapped_model, @@ -142,11 +171,28 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts( ), ) .await; - return None; + return Ok(None); }; + if transport.provider.provider_type.eq_ignore_ascii_case("xai") + && !is_native_video_request(&transport.provider.provider_type, parts.uri.path()) + { + provider_request_body = + convert_openai_video_request(&provider_request_body).map_err(|message| { + GatewayError::Client { + status: http::StatusCode::BAD_REQUEST, + message: message.to_string(), + } + })?; + } + apply_xai_upstream_payload_edits( + &mut provider_request_body, + transport.provider.provider_type.as_str(), + spec_metadata.api_format, + ); let Some(provider_request_headers) = build_video_create_headers(ProviderVideoCreateHeadersInput { + transport, headers: effective_headers, auth_header: &auth_header, auth_value: &auth_value, @@ -170,10 +216,10 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts( ), ) .await; - return None; + return Ok(None); }; - Some(LocalVideoCreateCandidatePayloadParts { + Ok(Some(LocalVideoCreateCandidatePayloadParts { transport: Arc::clone(transport), auth_header, auth_value, @@ -181,7 +227,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts( provider_request_headers, provider_request_body, upstream_url, - }) + })) } fn provider_video_create_family(family: LocalVideoCreateFamily) -> ProviderVideoCreateFamily { diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/deepseek.rs b/apps/aether-gateway/src/ai_serving/planner/standard/deepseek.rs index 6bb5ac92a..bc0eb87d6 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/deepseek.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/deepseek.rs @@ -13,7 +13,9 @@ pub(crate) fn openai_responses_reasoning_replay_policy( base_url: &str, _provider_model: &str, ) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy { - if is_deepseek_provider(provider_type, base_url) { + if provider_type.trim().eq_ignore_ascii_case("xai") { + crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted + } else if is_deepseek_provider(provider_type, base_url) { crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque } else { crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds @@ -238,6 +240,27 @@ mod tests { openai_responses_reasoning_replay_policy, }; + #[test] + fn xai_reasoning_policy_comes_from_provider_type() { + use crate::ai_serving::OpenAiResponsesReasoningReplayPolicy; + assert_eq!( + openai_responses_reasoning_replay_policy( + "xai", + "https://custom.example/v1", + "grok-4.6" + ), + OpenAiResponsesReasoningReplayPolicy::XaiEncrypted + ); + assert_eq!( + openai_responses_reasoning_replay_policy( + "openai", + "https://custom.example/v1", + "grok-4.6" + ), + OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds + ); + } + #[test] fn detects_deepseek_provider_only_by_official_host() { assert!(!is_deepseek_provider( diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/mod.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/mod.rs index 7ea4f4bdf..c6fcc2779 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/mod.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/mod.rs @@ -488,6 +488,7 @@ impl ResponsesWebSocketBodyNormalization { digest.update([match self.reasoning_replay_policy { crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds => 0, crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque => 1, + crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted => 2, }]); update_normalization_optional_json_digest(&mut digest, self.model_directive_patch.as_ref()); digest.finalize().into() diff --git a/apps/aether-gateway/src/ai_serving/pure/mod.rs b/apps/aether-gateway/src/ai_serving/pure/mod.rs index 91115566b..42c45f186 100644 --- a/apps/aether-gateway/src/ai_serving/pure/mod.rs +++ b/apps/aether-gateway/src/ai_serving/pure/mod.rs @@ -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_special_headers, apply_model_directive_mapping_patch, apply_model_directive_overrides_from_model, apply_model_directive_overrides_from_request, - apply_openai_responses_compact_special_body_edits, build_chatgpt_web_image_request_body, + apply_openai_responses_compact_special_body_edits, apply_xai_upstream_payload_edits, + apply_xai_upstream_payload_edits_with_client, build_chatgpt_web_image_request_body, build_codex_model_catalog_metadata, build_codex_openai_image_api_provider_request_body, build_core_error_body_for_client_format, build_cross_format_openai_chat_request_body, build_cross_format_openai_chat_request_body_with_model_directives, diff --git a/apps/aether-gateway/src/ai_serving/transport.rs b/apps/aether-gateway/src/ai_serving/transport.rs index 97fe5434a..8c81cb17d 100644 --- a/apps/aether-gateway/src/ai_serving/transport.rs +++ b/apps/aether-gateway/src/ai_serving/transport.rs @@ -58,6 +58,10 @@ pub(crate) mod windsurf { pub(crate) use aether_provider_transport::windsurf::*; } +pub(crate) mod xai { + pub(crate) use aether_provider_transport::xai::*; +} + pub(crate) use aether_provider_transport::{ append_transport_diagnostics_to_value, apply_codex_fingerprint_convergence, apply_codex_fingerprint_convergence_with_context, apply_local_auth_config_header_overrides, diff --git a/apps/aether-gateway/src/api/ai/registry.rs b/apps/aether-gateway/src/api/ai/registry.rs index 4fe6cddbf..685e0d956 100644 --- a/apps/aether-gateway/src/api/ai/registry.rs +++ b/apps/aether-gateway/src/api/ai/registry.rs @@ -53,6 +53,8 @@ const AI_ANY_ROUTE_PATTERNS: &[&str] = &[ "/v1beta/operations/{*operation_path}", "/v1/videos", "/v1/videos/{*video_path}", + "/openai/v1/videos", + "/openai/v1/videos/{*video_path}", "/upload/v1beta/files", "/v1beta/files", "/v1beta/files/{*file_path}", diff --git a/apps/aether-gateway/src/async_task/runtime.rs b/apps/aether-gateway/src/async_task/runtime.rs index d1e4a64b5..f5d301ee2 100644 --- a/apps/aether-gateway/src/async_task/runtime.rs +++ b/apps/aether-gateway/src/async_task/runtime.rs @@ -536,6 +536,9 @@ mod tests { fn sample_sparse_stored_task() -> StoredVideoTask { let snapshot = LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-1".to_string(), upstream_task_id: "ext-1".to_string(), created_at_unix_ms: 1, diff --git a/apps/aether-gateway/src/constants.rs b/apps/aether-gateway/src/constants.rs index e1edeb5d9..9a2bcbe14 100644 --- a/apps/aether-gateway/src/constants.rs +++ b/apps/aether-gateway/src/constants.rs @@ -140,6 +140,8 @@ pub(crate) const RUST_FRONTDOOR_OWNED_ROUTE_PATTERNS: &[&str] = &[ "/v1beta/models/{model}/operations/{id}", "/v1beta/operations", "/v1beta/operations/{id}", + "/openai/v1/videos", + "/openai/v1/videos/{path...}", "/v1/videos", "/v1/videos/{path...}", "/upload/v1beta/files", diff --git a/apps/aether-gateway/src/control/route/ai.rs b/apps/aether-gateway/src/control/route/ai.rs index f6ac5c83b..856f62543 100644 --- a/apps/aether-gateway/src/control/route/ai.rs +++ b/apps/aether-gateway/src/control/route/ai.rs @@ -137,7 +137,11 @@ pub(super) fn classify_ai_public_route( .with_client_surface(detect_claude_client_surface(headers)) .with_api_operation(ApiOperation::ClaudeMessagesCreate), ) - } else if normalized_path.starts_with("/v1/videos") { + } else if normalized_path == "/v1/videos" + || normalized_path.starts_with("/v1/videos/") + || normalized_path == "/openai/v1/videos" + || normalized_path.starts_with("/openai/v1/videos/") + { Some(classified( "ai_public", "openai", diff --git a/apps/aether-gateway/src/data/state/testing/video_tasks.rs b/apps/aether-gateway/src/data/state/testing/video_tasks.rs index 9db104d98..24952c220 100644 --- a/apps/aether-gateway/src/data/state/testing/video_tasks.rs +++ b/apps/aether-gateway/src/data/state/testing/video_tasks.rs @@ -123,6 +123,15 @@ impl GatewayDataState { } #[cfg(test)] + pub(crate) fn attach_video_task_repository_for_tests(mut self, repository: Arc) -> 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(repository: Arc) -> Self where T: VideoTaskRepository + 'static, diff --git a/apps/aether-gateway/src/executor/orchestration.rs b/apps/aether-gateway/src/executor/orchestration.rs index e79f96038..0473bfa3f 100644 --- a/apps/aether-gateway/src/executor/orchestration.rs +++ b/apps/aether-gateway/src/executor/orchestration.rs @@ -1464,6 +1464,22 @@ pub(crate) async fn maybe_execute_sync_via_local_video_decision( .await } +fn supports_local_video_get( + parts: &http::request::Parts, + decision: &GatewayControlDecision, +) -> bool { + parts.method == http::Method::GET + && decision.route_kind.as_deref() == Some("video") + && (crate::video_tasks::resolve_video_task_read_lookup_key( + decision.route_family.as_deref(), + parts.uri.path(), + ) + .is_some() + || (decision.route_family.as_deref() == Some("openai") + && crate::video_tasks::extract_openai_task_id_from_content_path(parts.uri.path()) + .is_some())) +} + pub(crate) fn maybe_execute_sync_request<'a>( state: &'a AppState, parts: &'a http::request::Parts, @@ -1477,7 +1493,7 @@ pub(crate) fn maybe_execute_sync_request<'a>( }; #[cfg(not(test))] { - if parts.method != http::Method::POST { + if parts.method != http::Method::POST && !supports_local_video_get(parts, decision) { return Ok(LocalExecutionRequestOutcome::NoPath); } return maybe_execute_sync_local_path(state, parts, body_bytes, trace_id, decision) @@ -1490,6 +1506,7 @@ pub(crate) fn maybe_execute_sync_request<'a>( .unwrap_or_default() .is_empty() && parts.method != http::Method::POST + && !supports_local_video_get(parts, decision) { return Ok(LocalExecutionRequestOutcome::NoPath); } @@ -1511,7 +1528,7 @@ pub(crate) fn maybe_execute_stream_request<'a>( }; #[cfg(not(test))] { - if parts.method != http::Method::POST { + if parts.method != http::Method::POST && !supports_local_video_get(parts, decision) { return Ok(LocalExecutionRequestOutcome::NoPath); } return maybe_execute_stream_local_path(state, parts, body_bytes, trace_id, decision) @@ -1524,6 +1541,7 @@ pub(crate) fn maybe_execute_stream_request<'a>( .unwrap_or_default() .is_empty() && parts.method != http::Method::POST + && !supports_local_video_get(parts, decision) { return Ok(LocalExecutionRequestOutcome::NoPath); } diff --git a/apps/aether-gateway/src/frontdoor_loop_guard.rs b/apps/aether-gateway/src/frontdoor_loop_guard.rs index 52c50feff..b3824be3d 100644 --- a/apps/aether-gateway/src/frontdoor_loop_guard.rs +++ b/apps/aether-gateway/src/frontdoor_loop_guard.rs @@ -32,6 +32,10 @@ fn request_has_execution_runtime_via_guard(headers: &HeaderMap) -> bool { } pub(crate) fn frontdoor_self_loop_public_ai_path(path: &str) -> bool { + let path = path + .strip_prefix("/openai") + .filter(|p| *p == "/v1/videos" || p.starts_with("/v1/videos/")) + .unwrap_or(path); matches!( path, "/v1/messages" diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/execution.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/execution.rs index 42d9c0c15..f2a355f3d 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/execution.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/execution.rs @@ -74,7 +74,8 @@ fn validate_batch_access_token_import( ) -> Result<(), String> { if !provider_type_supports_access_token_import(provider_type) { return Err( - "Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok Provider".to_string(), + "Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok / xAI Provider" + .to_string(), ); } if provider_type.eq_ignore_ascii_case("claude_code") { diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/parse.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/parse.rs index 60cff9220..0ad07cd61 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/parse.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/parse.rs @@ -214,7 +214,12 @@ fn extract_admin_provider_oauth_batch_import_entry( } else { let sso_from_cookie = grok_cookie_session_token(provider_type, raw_token); let token_input = sso_from_cookie.as_deref().unwrap_or(raw_token); - let (refresh_token, access_token) = import_tokens_from_raw_token(token_input); + let (refresh_token, access_token) = + if provider_type.trim().eq_ignore_ascii_case("xai") { + (None, Some(token_input.to_string())) + } else { + import_tokens_from_raw_token(token_input) + }; let (refresh_token, access_token) = normalize_provider_import_tokens( provider_type, refresh_token.as_deref(), @@ -262,6 +267,7 @@ fn extract_admin_provider_oauth_batch_import_entry( let object = normalized_claude_object.as_ref().unwrap_or(object); let is_grok = provider_type.trim().eq_ignore_ascii_case("grok"); let is_windsurf = provider_type.trim().eq_ignore_ascii_case("windsurf"); + let is_xai = provider_type.trim().eq_ignore_ascii_case("xai"); let is_codex_agent_identity = provider_type.trim().eq_ignore_ascii_case("codex") && aether_provider_transport::is_codex_agent_identity_auth_config_value(item); if is_codex_agent_identity { @@ -336,14 +342,6 @@ fn extract_admin_provider_oauth_batch_import_entry( } else { None }; - let (refresh_token, access_token) = normalize_provider_import_tokens( - provider_type, - refresh_token.as_deref(), - access_token - .as_deref() - .or(session_token.as_deref()) - .or(header_bearer_token.as_deref()), - ); let windsurf_api_key = is_windsurf .then(|| { coerce_admin_provider_oauth_import_str( @@ -351,6 +349,22 @@ fn extract_admin_provider_oauth_batch_import_entry( ) }) .flatten(); + let xai_api_key = is_xai + .then(|| { + coerce_admin_provider_oauth_import_str( + object.get("api_key").or_else(|| object.get("apiKey")), + ) + }) + .flatten(); + let (refresh_token, access_token) = normalize_provider_import_tokens( + provider_type, + refresh_token.as_deref(), + access_token + .as_deref() + .or(session_token.as_deref()) + .or(header_bearer_token.as_deref()) + .or(xai_api_key.as_deref()), + ); let windsurf_token = is_windsurf .then(|| { coerce_admin_provider_oauth_import_str( @@ -1577,4 +1591,23 @@ mod tests { assert!(entries[1].access_token.is_none()); assert!(entries[1].raw_credentials.is_none()); } + + #[test] + fn parses_xai_api_key_json_and_raw_lines_as_access_token() { + let entries = parse_admin_provider_oauth_batch_import_entries( + "xai", + r#"{"api_key":"xai-api-key","email":"a@x.ai"} +{"refresh_token":"xai-refresh"} +xai-raw-api-key"#, + ); + + assert_eq!(entries.len(), 3); + assert!(entries[0].refresh_token.is_none()); + assert_eq!(entries[0].access_token.as_deref(), Some("xai-api-key")); + assert_eq!(entries[0].email.as_deref(), Some("a@x.ai")); + 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")); + } } diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/authorize.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/authorize.rs index b4ce3c5e5..fc92acb64 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/authorize.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/authorize.rs @@ -186,10 +186,10 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize( )); }; let provider_type = provider.provider_type.trim().to_ascii_lowercase(); - if provider_type != "kiro" && provider_type != "windsurf" { + if provider_type != "kiro" && provider_type != "windsurf" && provider_type != "xai" { return Ok(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, - "设备授权仅支持 Kiro / Windsurf provider", + "设备授权仅支持 Kiro / Windsurf / xAI provider", )); } let Some(principal) = request_context @@ -219,6 +219,19 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize( ) .await; + if provider_type == "xai" { + return super::xai::handle_admin_provider_oauth_xai_device_authorize( + state, + &provider_id, + &provider, + principal, + runtime_endpoint.as_ref(), + request_proxy, + payload.proxy_node_id.as_deref(), + ) + .await; + } + if provider_type == "windsurf" { let session_id = generate_provider_oauth_nonce(); let login_option = payload diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/mod.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/mod.rs index 295a93336..172f99cde 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/mod.rs @@ -2,6 +2,7 @@ mod authorize; mod lease; mod poll; mod session; +mod xai; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::GatewayError; diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/poll.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/poll.rs index 8074a0024..649623f51 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/poll.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/poll.rs @@ -479,6 +479,18 @@ pub(super) async fn handle_admin_provider_oauth_device_poll( ) .await; + if provider_type == "xai" { + return super::xai::handle_admin_provider_oauth_xai_device_poll( + state, + &provider, + &endpoints, + request_proxy, + session_id, + session, + ) + .await; + } + if provider_type == "windsurf" { return handle_admin_provider_oauth_windsurf_browser_device_poll( state, diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/xai.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/xai.rs new file mode 100644 index 000000000..bf42d5db6 --- /dev/null +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/xai.rs @@ -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, + proxy_node_id: Option<&str>, +) -> Result, 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, + session_id: &str, + mut session: StoredAdminProviderOAuthDeviceSession, +) -> Result, 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, + session_id: &str, + mut session: StoredAdminProviderOAuthDeviceSession, + result: aether_oauth::provider::ProviderOAuthTokenSet, +) -> Result, 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 { + 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(), + } +} diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/import.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/import.rs index 5ef0ea8b9..c7707e786 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/import.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/import.rs @@ -715,7 +715,7 @@ async fn resolve_admin_provider_oauth_single_import_tokens( if !provider_type_supports_access_token_import(provider_type) { return Err(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, - "Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok Provider", + "Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok / xAI Provider", )); } @@ -867,7 +867,7 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token( flatten_claude_code_credentials_payload(&mut raw_payload); } let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken"); - let access_token_input = import_payload_string_any( + let mut access_token_input = import_payload_string_any( &raw_payload, &[ "access_token", @@ -879,6 +879,9 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token( ], ) .or_else(|| provider_oauth_import_authorization_bearer_token_from_object(&raw_payload)); + if provider_type == "xai" && access_token_input.is_none() { + access_token_input = import_payload_string(&raw_payload, "api_key", "apiKey"); + } let imported_expires_at = import_payload_u64_any(&raw_payload, &["expires_at", "expiresAt", "expired"]); let (refresh_token_input, access_token_input) = normalize_provider_import_tokens( @@ -901,7 +904,11 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token( if !create_agent_identity && refresh_token_input.is_none() && access_token_input.is_none() { return Ok(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, - "Refresh Token、Access Token 或 sso_token 不能为空", + if provider_type == "xai" { + "Refresh Token、Access Token 或 api_key 不能为空" + } else { + "Refresh Token、Access Token 或 sso_token 不能为空" + }, )); } if !is_fixed_provider_type_for_provider_oauth(&provider_type) { diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/start.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/start.rs index d909b0320..184ec1c1b 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/start.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/start.rs @@ -70,6 +70,12 @@ pub(super) async fn handle_admin_provider_oauth_start_key( "Windsurf 请使用浏览器登录或导入凭据。", )); } + if provider_type == "xai" { + return Ok(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "xAI 请使用设备授权或导入凭据。", + )); + } let Some(template) = admin_provider_oauth_template(&provider_type) else { return Ok(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, @@ -167,6 +173,12 @@ pub(super) async fn handle_admin_provider_oauth_start_provider( "Windsurf 请使用浏览器登录或导入凭据。", )); } + if provider_type == "xai" { + return Ok(build_internal_control_error_response( + http::StatusCode::BAD_REQUEST, + "xAI 请使用设备授权或导入凭据。", + )); + } let Some(template) = admin_provider_oauth_template(&provider_type) else { return Ok(build_internal_control_error_response( http::StatusCode::BAD_REQUEST, diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/token_import.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/token_import.rs index f446d4104..b6566f052 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/token_import.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/token_import.rs @@ -121,6 +121,9 @@ pub(super) fn normalize_provider_import_tokens( if provider_type == "grok" { return (None, access_token.or(refresh_token)); } + if provider_type == "xai" { + return (refresh_token, access_token); + } if provider_type == "claude_code" { if access_token.is_none() && refresh_token.as_deref().is_some_and(is_claude_access_token) { return (None, refresh_token); @@ -237,7 +240,7 @@ pub(super) fn provider_oauth_import_authorization_bearer_token_from_object( pub(super) fn provider_type_supports_access_token_import(provider_type: &str) -> bool { matches!( provider_type.trim().to_ascii_lowercase().as_str(), - "claude_code" | "codex" | "chatgpt_web" | "grok" + "claude_code" | "codex" | "chatgpt_web" | "grok" | "xai" ) } @@ -331,6 +334,15 @@ pub(super) fn build_provider_access_token_import_auth_config( auth_config.insert("sso_token".to_string(), json!(access_token)); auth_config.insert("auth_method".to_string(), json!("sso_token")); } + if provider_type.trim().eq_ignore_ascii_case("xai") { + if refresh_token.is_some() { + auth_config.insert("auth_method".to_string(), json!("oauth")); + auth_config.insert("using_api".to_string(), json!(false)); + } else { + auth_config.insert("auth_method".to_string(), json!("api_key")); + auth_config.insert("using_api".to_string(), json!(true)); + } + } auth_config.insert( "access_token_import_temporary".to_string(), @@ -532,6 +544,41 @@ mod tests { ); } + #[test] + fn normalize_xai_import_keeps_refresh_token_separate_from_api_key() { + let (refresh_token, access_token) = + normalize_provider_import_tokens("xai", Some("xai-refresh-token"), None); + assert_eq!(refresh_token.as_deref(), Some("xai-refresh-token")); + assert!(access_token.is_none()); + + let (refresh_token, access_token) = + normalize_provider_import_tokens("xai", None, Some("xai-api-key")); + assert!(refresh_token.is_none()); + assert_eq!(access_token.as_deref(), Some("xai-api-key")); + } + + #[test] + fn builds_xai_auth_config_from_api_key_and_oauth_tokens() { + let (api_key_config, _) = + build_provider_access_token_import_auth_config("xai", "xai-api-key", None, None, None); + assert_eq!(api_key_config.get("auth_method"), Some(&json!("api_key"))); + assert_eq!(api_key_config.get("using_api"), Some(&json!(true))); + + let (oauth_config, _) = build_provider_access_token_import_auth_config( + "xai", + "xai-access-token", + Some("xai-refresh-token"), + None, + None, + ); + assert_eq!(oauth_config.get("auth_method"), Some(&json!("oauth"))); + assert_eq!(oauth_config.get("using_api"), Some(&json!(false))); + assert_eq!( + oauth_config.get("refresh_token"), + Some(&json!("xai-refresh-token")) + ); + } + #[test] fn flattens_only_claude_ai_oauth_credentials_and_converts_expiry_ms() { let mut payload = json!({ diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/dispatch.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/dispatch.rs index 2963ba5e7..268b261df 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/dispatch.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/dispatch.rs @@ -8,6 +8,7 @@ use super::gemini_cli::refresh_gemini_cli_provider_quota_locally; use super::grok::refresh_grok_provider_quota_locally; use super::kiro::refresh_kiro_provider_quota_locally; use super::windsurf::refresh_windsurf_provider_quota_locally; +use super::xai::refresh_xai_provider_quota_locally; use crate::handlers::admin::request::AdminAppState; use crate::GatewayError; use aether_contracts::ProxySnapshot; @@ -43,6 +44,7 @@ const PROVIDER_QUOTA_REFRESH_HANDLERS: &[(&str, ProviderQuotaRefreshHandler)] = ("grok", refresh_grok_provider_quota_locally_boxed), ("kiro", refresh_kiro_provider_quota_locally_boxed), ("windsurf", refresh_windsurf_provider_quota_locally_boxed), + ("xai", refresh_xai_provider_quota_locally_boxed), ]; pub(crate) async fn refresh_provider_pool_quota_locally( @@ -174,3 +176,19 @@ fn refresh_windsurf_provider_quota_locally_boxed<'a>( proxy_override, )) } + +fn refresh_xai_provider_quota_locally_boxed<'a>( + state: &'a AdminAppState<'a>, + provider: &'a StoredProviderCatalogProvider, + endpoint: &'a StoredProviderCatalogEndpoint, + keys: Vec, + proxy_override: Option, +) -> ProviderQuotaRefreshFuture<'a> { + Box::pin(refresh_xai_provider_quota_locally( + state, + provider, + endpoint, + keys, + proxy_override, + )) +} diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/mod.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/mod.rs index 8629fff76..3f794abae 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/mod.rs @@ -7,3 +7,4 @@ pub(crate) mod grok; pub(crate) mod kiro; pub(crate) mod shared; pub(crate) mod windsurf; +pub(crate) mod xai; diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/shared.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/shared.rs index 7f1129faf..e9c086b7e 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/shared.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/shared.rs @@ -1715,6 +1715,7 @@ fn provider_quota_url_has_allowed_origin(provider_name: &str, value: &str) -> bo "gemini_cli" => host == "cloudcode-pa.googleapis.com", "chatgpt_web" | "codex" => host == "chatgpt.com", "grok" => host == "grok.com", + "xai" => host == "cli-chat-proxy.grok.com", "windsurf" => host == "server.codeium.com", "kiro" => kiro_quota_host_is_allowed(host), _ => false, @@ -1814,6 +1815,14 @@ mod tests { ), ("codex", "https://chatgpt.com/backend-api/wham/usage"), ("grok", "https://grok.com/rest/rate-limits"), + ( + "xai", + "https://cli-chat-proxy.grok.com/v1/billing?format=credits", + ), + ( + "xai", + "https://cli-chat-proxy.grok.com/v1/user", + ), ( "windsurf", "https://server.codeium.com/exa.seat_management_pb.SeatManagementService/GetUserStatus", @@ -1847,6 +1856,11 @@ mod tests { "https://chatgpt.com.attacker.test/backend-api/wham/usage", ), ("grok", "https://grok.com.attacker.test/rest/rate-limits"), + ( + "xai", + "https://cli-chat-proxy.grok.com.attacker.test/v1/billing", + ), + ("xai", "https://api.x.ai/v1/billing?format=credits"), ("windsurf", "https://server.codeium.com.attacker.test/quota"), ( "gemini_cli", diff --git a/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/xai.rs b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/xai.rs new file mode 100644 index 000000000..32cae49fd --- /dev/null +++ b/apps/aether-gateway/src/handlers/admin/provider/oauth/quota/xai.rs @@ -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 { + 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).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, + proxy_override: Option, +) -> Result, 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::; + 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::; + + 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, + }))) +} diff --git a/apps/aether-gateway/src/handlers/admin/provider/pool_admin/payloads.rs b/apps/aether-gateway/src/handlers/admin/provider/pool_admin/payloads.rs index 5dec0f9d5..da5c80b6b 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/pool_admin/payloads.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/pool_admin/payloads.rs @@ -932,6 +932,13 @@ fn admin_pool_build_account_quota( return Some(account_quota); } } + "xai" => { + if let Some(account_quota) = + admin_pool_build_kiro_account_quota_from_snapshot(quota_snapshot) + { + return Some(account_quota); + } + } "chatgpt_web" => { if let Some(account_quota) = admin_pool_build_chatgpt_web_account_quota_from_snapshot(quota_snapshot) @@ -1591,4 +1598,29 @@ mod tests { Some("Auto剩余 40.0% (60/150) | Heavy剩余 0.0% (0/20)".to_string()) ); } + + #[test] + fn xai_account_quota_is_rendered_as_remaining_percent() { + let quota_snapshot = json!({ + "provider_type": "xai", + "code": "ok", + "exhausted": false, + "plan_type": "SuperGrok", + "windows": [ + { + "code": "usage", + "label": "周额度", + "scope": "account", + "used_ratio": 0.46, + "remaining_ratio": 0.54 + } + ] + }); + let quota_snapshot = quota_snapshot.as_object().unwrap(); + + assert_eq!( + admin_pool_build_account_quota("xai", Some(quota_snapshot)), + Some("剩余 54.0%".to_string()) + ); + } } diff --git a/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test.rs b/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test.rs index 8319b1409..01ba34459 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/query/models/model_test.rs @@ -3511,6 +3511,11 @@ async fn provider_query_execute_standard_test_candidate( codex_model_capabilities.as_ref(), ); } + crate::provider_transport::insert_cli_identity_headers_if_needed( + &transport, + provider_api_format, + &mut request_headers, + ); if !uses_vertex_query_auth { if let (Some(auth_header), Some(auth_value)) = (auth_header.as_deref(), auth_value.as_deref()) diff --git a/apps/aether-gateway/src/handlers/admin/provider/write/normalize.rs b/apps/aether-gateway/src/handlers/admin/provider/write/normalize.rs index 65a6d364e..2649209f8 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/write/normalize.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/write/normalize.rs @@ -4,9 +4,9 @@ pub(crate) fn normalize_provider_type_input(value: &str) -> Result Ok(normalized), + | "antigravity" | "vertex_ai" | "grok" | "windsurf" | "xai" => Ok(normalized), _ => Err( - "provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai / grok / windsurf" + "provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai / grok / windsurf / xai" .to_string(), ), } @@ -405,6 +405,14 @@ mod tests { ); } + #[test] + fn normalize_provider_type_supports_xai() { + assert_eq!( + normalize_provider_type_input(" xAI ").expect("type should normalize"), + "xai" + ); + } + #[test] fn normalize_api_format_list_dedupes_canonical_formats() { assert_eq!( diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/continuation.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/continuation.rs index d8de44134..7aaa58ba6 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/continuation.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/continuation.rs @@ -90,6 +90,8 @@ pub(super) struct ResponsesWebSocketContinuationRecord { /// request JSON can never set it. #[serde(default)] deepseek_opaque_reasoning_replay: bool, + #[serde(default)] + xai_encrypted_reasoning_replay: bool, /// A prior turn stored PII sentinels whose restore mapping exists only on /// the original downstream socket. Such a chain cannot safely resume on a /// new socket without leaking sentinels, so lookup succeeds but bootstrap @@ -122,6 +124,10 @@ impl ResponsesWebSocketContinuationRecord { normalization.reasoning_replay_policy(), crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque ), + xai_encrypted_reasoning_replay: matches!( + normalization.reasoning_replay_policy(), + crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted + ), has_connection_local_redaction, responses_lite_static_config, }; @@ -156,7 +162,9 @@ impl ResponsesWebSocketContinuationRecord { pub(super) fn reasoning_replay_policy( &self, ) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy { - if self.deepseek_opaque_reasoning_replay { + if self.xai_encrypted_reasoning_replay { + crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted + } else if self.deepseek_opaque_reasoning_replay { crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque } else { crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds @@ -476,6 +484,7 @@ mod tests { binding_fingerprint: [7; 32], normalization_fingerprint: [9; 32], deepseek_opaque_reasoning_replay: false, + xai_encrypted_reasoning_replay: false, has_connection_local_redaction: false, responses_lite_static_config: Some(ResponsesLiteStaticConfig::from_response_create( &json!({ @@ -714,6 +723,29 @@ mod tests { assert_eq!(decoded, record()); } + #[test] + fn serialized_record_preserves_xai_replay_policy_and_reads_legacy_records() { + let mut expected = record(); + expected.xai_encrypted_reasoning_replay = true; + let mut serialized = serde_json::to_value(&expected).unwrap(); + let decoded: ResponsesWebSocketContinuationRecord = + serde_json::from_value(serialized.clone()).unwrap(); + assert_eq!( + decoded.reasoning_replay_policy(), + crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted + ); + serialized + .as_object_mut() + .unwrap() + .remove("xai_encrypted_reasoning_replay"); + let legacy: ResponsesWebSocketContinuationRecord = + serde_json::from_value(serialized).unwrap(); + assert_eq!( + legacy.reasoning_replay_policy(), + crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds + ); + } + #[test] fn serialized_record_preserves_only_the_server_derived_reasoning_replay_policy_bit() { let mut expected = record(); diff --git a/apps/aether-gateway/src/handlers/shared/catalog.rs b/apps/aether-gateway/src/handlers/shared/catalog.rs index 736640b3e..b8c585c4c 100644 --- a/apps/aether-gateway/src/handlers/shared/catalog.rs +++ b/apps/aether-gateway/src/handlers/shared/catalog.rs @@ -1388,6 +1388,168 @@ fn build_kiro_quota_status_snapshot( })) } +fn build_xai_quota_status_snapshot( + upstream_metadata: Option<&Value>, + source: &str, +) -> Option { + let metadata = provider_quota_metadata_bucket(upstream_metadata, "xai")?; + let observed_at_unix_secs = provider_quota_timestamp_unix_secs(metadata.get("updated_at")); + let usage_limit = metadata + .get("usage_limit") + .and_then(admin_provider_quota_pure::coerce_json_f64); + let current_usage = metadata + .get("current_usage") + .and_then(admin_provider_quota_pure::coerce_json_f64); + let remaining = metadata + .get("remaining") + .and_then(admin_provider_quota_pure::coerce_json_f64); + let usage_ratio = metadata + .get("usage_percentage") + .and_then(admin_provider_quota_pure::coerce_json_f64) + .map(|value| (value / 100.0).clamp(0.0, 1.0)) + .or_else(|| { + current_usage + .zip(usage_limit) + .and_then(|(current_usage, usage_limit)| { + (usage_limit > 0.0).then_some((current_usage / usage_limit).clamp(0.0, 1.0)) + }) + }); + let remaining_ratio = usage_ratio.map(|value| (1.0 - value).max(0.0)); + let next_reset_at = provider_quota_timestamp_unix_secs(metadata.get("next_reset_at")); + let reset_seconds = quota_window_reset_seconds(observed_at_unix_secs, next_reset_at); + let plan_type = metadata + .get("subscription_title") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned); + let period_type = metadata + .get("period_type") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned); + let usage_label = match period_type.as_deref() { + Some("monthly") => "月额度", + Some("weekly") => "周额度", + _ => "额度", + }; + + let mut windows = Vec::new(); + if usage_ratio.is_some() + || remaining.is_some() + || usage_limit.is_some() + || current_usage.is_some() + || next_reset_at.is_some() + { + windows.push(json!({ + "code": "usage", + "label": usage_label, + "scope": "account", + "unit": if usage_limit.is_some() { "usd" } else { "percent" }, + "used_ratio": usage_ratio, + "remaining_ratio": remaining_ratio, + "used_value": current_usage, + "remaining_value": remaining, + "limit_value": usage_limit, + "reset_at": next_reset_at, + "reset_seconds": reset_seconds, + })); + } + + let prepaid_balance = metadata + .get("prepaid_balance") + .and_then(admin_provider_quota_pure::coerce_json_f64); + if prepaid_balance.is_some_and(|value| value > 0.0) { + windows.push(json!({ + "code": "prepaid", + "label": "预付额度", + "scope": "account", + "unit": "usd", + "used_ratio": serde_json::Value::Null, + "remaining_ratio": serde_json::Value::Null, + "remaining_value": prepaid_balance, + "reset_at": serde_json::Value::Null, + "reset_seconds": serde_json::Value::Null, + })); + } + + let on_demand_cap = metadata + .get("on_demand_cap") + .and_then(admin_provider_quota_pure::coerce_json_f64); + let on_demand_used = metadata + .get("on_demand_used") + .and_then(admin_provider_quota_pure::coerce_json_f64); + let on_demand_enabled = metadata + .get("on_demand_enabled") + .and_then(admin_provider_quota_pure::coerce_json_bool) + != Some(false); + if on_demand_enabled && on_demand_cap.is_some_and(|value| value > 0.0) { + let on_demand_remaining = on_demand_cap + .zip(on_demand_used) + .map(|(cap, used)| (cap - used).max(0.0)); + let on_demand_ratio = on_demand_cap + .zip(on_demand_used) + .and_then(|(cap, used)| (cap > 0.0).then_some((used / cap).clamp(0.0, 1.0))); + windows.push(json!({ + "code": "on_demand", + "label": "按需额度", + "scope": "account", + "unit": "usd", + "used_ratio": on_demand_ratio, + "remaining_ratio": on_demand_ratio.map(|value| (1.0 - value).max(0.0)), + "used_value": on_demand_used, + "remaining_value": on_demand_remaining, + "limit_value": on_demand_cap, + "reset_at": serde_json::Value::Null, + "reset_seconds": serde_json::Value::Null, + })); + } + + if windows.is_empty() && plan_type.is_none() && observed_at_unix_secs.is_none() { + return None; + } + + let prepaid_available = prepaid_balance.is_some_and(|value| value > 0.0); + let on_demand_available = on_demand_enabled + && on_demand_cap.is_some_and(|value| value > 0.0) + && on_demand_used + .zip(on_demand_cap) + .is_some_and(|(used, cap)| used < cap); + let usage_exhausted = remaining.is_some_and(|value| value <= 0.0) + || usage_ratio.is_some_and(|value| value >= 1.0 - 1e-6); + let exhausted = usage_exhausted && !prepaid_available && !on_demand_available; + let reason = if exhausted { + Some("额度已耗尽".to_string()) + } else { + None + }; + let label = if exhausted { + Some("额度耗尽") + } else { + None + }; + let code = if exhausted { "exhausted" } else { "ok" }; + + Some(json!({ + "version": 2, + "provider_type": "xai", + "code": code, + "label": label, + "reason": reason, + "freshness": "fresh", + "source": source, + "observed_at": observed_at_unix_secs, + "exhausted": exhausted, + "usage_ratio": usage_ratio, + "updated_at": observed_at_unix_secs, + "reset_at": next_reset_at, + "reset_seconds": reset_seconds, + "plan_type": plan_type, + "windows": windows, + })) +} + fn build_chatgpt_web_quota_status_snapshot( upstream_metadata: Option<&Value>, source: &str, @@ -2255,6 +2417,7 @@ pub(crate) fn sync_provider_key_quota_status_snapshot( let mut quota = match normalized_provider_type.as_str() { "codex" => build_codex_quota_status_snapshot(upstream_metadata, source), "kiro" => build_kiro_quota_status_snapshot(upstream_metadata, source), + "xai" => build_xai_quota_status_snapshot(upstream_metadata, source), "chatgpt_web" => build_chatgpt_web_quota_status_snapshot(upstream_metadata, source), "windsurf" => build_windsurf_quota_status_snapshot(upstream_metadata, source), "antigravity" => build_antigravity_quota_status_snapshot(upstream_metadata, source), @@ -3622,6 +3785,43 @@ mod tests { assert_eq!(auto.get("used_value"), Some(&json!(90.0))); } + #[test] + fn provider_key_status_snapshot_payload_backfills_xai_weekly_credits() { + let mut key = sample_catalog_key(); + key.upstream_metadata = Some(json!({ + "xai": { + "updated_at": 1_778_067_246u64, + "usage_percentage": 46.0, + "period_type": "weekly", + "next_reset_at": 1_778_157_172u64, + "subscription_title": "SuperGrok", + "prepaid_balance": 0.0, + "on_demand_cap": 0.0, + "on_demand_used": 0.0 + } + })); + + let payload = provider_key_status_snapshot_payload(&key, "xai"); + let quota = payload + .get("quota") + .and_then(Value::as_object) + .expect("quota snapshot should be object"); + let windows = quota + .get("windows") + .and_then(Value::as_array) + .expect("xai quota windows should exist"); + + assert_eq!(quota.get("provider_type"), Some(&json!("xai"))); + assert_eq!(quota.get("code"), Some(&json!("ok"))); + assert_eq!(quota.get("exhausted"), Some(&json!(false))); + assert_eq!(quota.get("plan_type"), Some(&json!("SuperGrok"))); + assert_eq!(quota.get("usage_ratio"), Some(&json!(0.46))); + assert_eq!(quota.get("reset_at"), Some(&json!(1_778_157_172u64))); + assert_eq!(windows.len(), 1); + assert_eq!(windows[0].get("code"), Some(&json!("usage"))); + assert_eq!(windows[0].get("label"), Some(&json!("周额度"))); + } + #[test] fn provider_key_status_snapshot_payload_backfills_gemini_cli_account_credits() { let mut key = sample_catalog_key(); diff --git a/apps/aether-gateway/src/image_capabilities.rs b/apps/aether-gateway/src/image_capabilities.rs index 21dc5f183..0c5c10c1d 100644 --- a/apps/aether-gateway/src/image_capabilities.rs +++ b/apps/aether-gateway/src/image_capabilities.rs @@ -13,7 +13,7 @@ pub(crate) fn openai_image_provider_max_generation_count(provider_type: &str) -> GROK_OPENAI_IMAGE_MAX_GENERATION_COUNT } else if matches!( provider_type.trim().to_ascii_lowercase().as_str(), - "openai" | "codex" + "openai" | "codex" | "xai" ) { OPENAI_IMAGE_MAX_GENERATION_COUNT } else { @@ -58,6 +58,7 @@ mod tests { assert_eq!(openai_image_provider_max_generation_count("grok"), 4); assert_eq!(openai_image_provider_max_generation_count("openai"), 10); assert_eq!(openai_image_provider_max_generation_count("codex"), 10); + assert_eq!(openai_image_provider_max_generation_count("xai"), 10); assert_eq!(openai_image_provider_max_generation_count("custom"), 1); assert_eq!( openai_image_provider_max_generation_count_for_model("openai", Some("dall-e-3")), diff --git a/apps/aether-gateway/src/provider_key_auth.rs b/apps/aether-gateway/src/provider_key_auth.rs index 62379e62f..1d8021e7c 100644 --- a/apps/aether-gateway/src/provider_key_auth.rs +++ b/apps/aether-gateway/src/provider_key_auth.rs @@ -172,6 +172,7 @@ fn provider_uses_bearer_oauth_runtime(provider_type: &str) -> bool { | "antigravity" | "kiro" | "windsurf" + | "xai" ) } @@ -406,6 +407,22 @@ mod tests { ); } + #[test] + fn recognizes_xai_oauth_as_bearer_runtime() { + let semantics = provider_key_auth_semantics(&sample_key("oauth"), "xai"); + + assert!(semantics.oauth_managed()); + assert!(semantics.can_refresh_oauth()); + assert_eq!( + semantics.credential_kind(), + ProviderKeyCredentialKind::OAuthSession + ); + assert_eq!( + semantics.runtime_auth_kind(), + ProviderKeyRuntimeAuthKind::Bearer + ); + } + #[test] fn refresh_capability_requires_stored_refresh_token() { let semantics = provider_key_auth_semantics(&sample_key("oauth"), "codex"); diff --git a/apps/aether-gateway/src/router.rs b/apps/aether-gateway/src/router.rs index e7d48f818..157c46653 100644 --- a/apps/aether-gateway/src/router.rs +++ b/apps/aether-gateway/src/router.rs @@ -186,6 +186,8 @@ fn frontend_path_bypasses_static(path: &str) -> bool { "/health" | "/test-connection" | crate::constants::READYZ_PATH ) || path.starts_with("/api/") || path.starts_with("/v1/") + || path == "/openai/v1/videos" + || path.starts_with("/openai/v1/videos/") || path.starts_with("/v1beta/") || path.starts_with("/upload/") || path.starts_with("/_gateway/") diff --git a/apps/aether-gateway/src/state/integrations.rs b/apps/aether-gateway/src/state/integrations.rs index db94e5786..caf97843a 100644 --- a/apps/aether-gateway/src/state/integrations.rs +++ b/apps/aether-gateway/src/state/integrations.rs @@ -290,6 +290,14 @@ impl provider_transport::VideoTaskTransportSnapshotLookup for AppState { .await .map_err(GatewayError::into_message) } + + async fn resolve_video_task_proxy( + &self, + transport: &GatewayProviderTransportSnapshot, + ) -> Option { + self.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport) + .await + } } #[async_trait] diff --git a/apps/aether-gateway/src/tests/architecture/admin_provider.rs b/apps/aether-gateway/src/tests/architecture/admin_provider.rs index 9c27d423d..5e95f5c57 100644 --- a/apps/aether-gateway/src/tests/architecture/admin_provider.rs +++ b/apps/aether-gateway/src/tests/architecture/admin_provider.rs @@ -1789,6 +1789,7 @@ fn admin_provider_oauth_quota_mod_stays_thin() { "pub(crate) mod dispatch;", "pub(crate) mod kiro;", "pub(crate) mod shared;", + "pub(crate) mod xai;", ] { assert!( quota_mod.contains(pattern), @@ -1861,6 +1862,7 @@ fn admin_provider_oauth_quota_mod_stays_thin() { "refresh_antigravity_provider_quota_locally", "refresh_gemini_cli_provider_quota_locally", "refresh_chatgpt_web_provider_quota_locally", + "refresh_xai_provider_quota_locally", ] { assert!( quota_dispatch.contains(pattern), diff --git a/apps/aether-gateway/src/tests/architecture/ai_serving.rs b/apps/aether-gateway/src/tests/architecture/ai_serving.rs index b5ef1144f..fa71ab5c7 100644 --- a/apps/aether-gateway/src/tests/architecture/ai_serving.rs +++ b/apps/aether-gateway/src/tests/architecture/ai_serving.rs @@ -1452,6 +1452,7 @@ fn ai_serving_planner_separates_local_candidate_resolution_from_ranking() { "GeminiCliProviderPoolAdapter", "KiroProviderPoolAdapter", "ChatGptWebProviderPoolAdapter", + "XaiProviderPoolAdapter", "CLAUDE_CODE_PROVIDER_POOL_ADAPTER", "VERTEX_AI_PROVIDER_POOL_ADAPTER", "provider_types_for_capability", @@ -1478,6 +1479,7 @@ fn ai_serving_planner_separates_local_candidate_resolution_from_ranking() { "pub mod gemini_cli;", "pub mod kiro;", "pub mod chatgpt_web;", + "pub mod xai;", ] { assert!( 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", 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", vec![ diff --git a/apps/aether-gateway/src/tests/control/admin/oauth.rs b/apps/aether-gateway/src/tests/control/admin/oauth.rs index d75debf5b..50e9cfc95 100644 --- a/apps/aether-gateway/src/tests/control/admin/oauth.rs +++ b/apps/aether-gateway/src/tests/control/admin/oauth.rs @@ -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("user@x.ai"); + 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"], "user@x.ai"); + 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"], "user@x.ai"); + + token_handle.abort(); +} + #[test] fn gateway_handles_admin_provider_oauth_device_poll_for_windsurf_one_time_token() { run_admin_oauth_test( diff --git a/apps/aether-gateway/src/tests/video/mod.rs b/apps/aether-gateway/src/tests/video/mod.rs index a93ac3b2d..61b94cd08 100644 --- a/apps/aether-gateway/src/tests/video/mod.rs +++ b/apps/aether-gateway/src/tests/video/mod.rs @@ -36,6 +36,7 @@ mod openai_sync_task; mod registry_poller; mod routing; mod stream; +mod xai; /// 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 execution-runtime override is active. pub(super) fn video_proxy_node_repository(node_ids: I) -> Arc +where + I: IntoIterator, + S: AsRef, +{ + video_proxy_node_repository_at_url(node_ids, "http://127.0.0.1:1") +} + +pub(super) fn video_proxy_node_repository_at_url( + node_ids: I, + proxy_url: &str, +) -> Arc where I: IntoIterator, S: AsRef, @@ -68,7 +80,7 @@ where 1, ) .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}")) }); Arc::new(InMemoryProxyNodeRepository::seed(nodes)) @@ -86,6 +98,28 @@ pub(super) fn video_provider_catalog_repository( endpoint_base_url: &str, key_id: &str, upstream_api_key: &str, +) -> Arc { + 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, ) -> Arc { fn seal_bound_credential( provider_id: &str, @@ -117,7 +151,7 @@ pub(super) fn video_provider_catalog_repository( false, None, Some(2), - None, + proxy, Some(20.0), None, None, diff --git a/apps/aether-gateway/src/tests/video/registry_poller.rs b/apps/aether-gateway/src/tests/video/registry_poller.rs index b47b5938a..35bf4a0c3 100644 --- a/apps/aether-gateway/src/tests/video/registry_poller.rs +++ b/apps/aether-gateway/src/tests/video/registry_poller.rs @@ -13,7 +13,8 @@ use serde_json::json; use super::{ 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 { @@ -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_api_root = format!("{upstream_url}/v1"); + let upstream_api_root = "http://video-provider.invalid/v1".to_string(); let repository = Arc::new(InMemoryVideoTaskRepository::default()); repository .upsert(sample_due_openai_task(&upstream_api_root)) .await .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", "openai", "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, "key-openai-video-local-1", "sk-upstream-openai-video", + Some(json!({"enabled":true,"node_id":"poller-video-proxy"})), ); 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), provider_catalog_repository, 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_poller_config(std::time::Duration::from_millis(25), 8); diff --git a/apps/aether-gateway/src/tests/video/routing.rs b/apps/aether-gateway/src/tests/video/routing.rs index 07f551281..ba787c274 100644 --- a/apps/aether-gateway/src/tests/video/routing.rs +++ b/apps/aether-gateway/src/tests/video/routing.rs @@ -14,8 +14,7 @@ use crate::constants::{ use super::{build_router, start_server}; #[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_clone = Arc::clone(&execute_hits); 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 .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"); - assert_eq!(payload["error"]["type"], "http_error"); - assert_eq!( - payload["error"]["message"], - "当前 OpenAI Video 请求无法在本地执行:没有匹配到可用的执行路径" - ); + assert_eq!(payload, crate::video_tasks::not_found_body()); assert_eq!(*execute_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] -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_clone = Arc::clone(&execute_hits); 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 .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"); - assert_eq!(payload["error"]["type"], "http_error"); - assert_eq!( - payload["error"]["message"], - "当前 OpenAI Video 请求无法在本地执行:没有匹配到可用的执行路径" - ); + assert_eq!(payload, crate::video_tasks::not_found_body()); assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0); assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); assert_eq!( @@ -165,7 +155,7 @@ async fn gateway_locally_denies_video_control_sync_without_opt_in_header_when_ex } #[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_clone = Arc::clone(&execute_hits); 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 .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"); - assert_eq!(payload["error"]["type"], "http_error"); - assert_eq!( - payload["error"]["message"], - "当前 OpenAI Video 请求无法在本地执行:没有匹配到可用的执行路径" - ); + assert_eq!(payload, crate::video_tasks::not_found_body()); assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0); assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); diff --git a/apps/aether-gateway/src/tests/video/xai.rs b/apps/aether-gateway/src/tests/video/xai.rs new file mode 100644 index 000000000..059b4b844 --- /dev/null +++ b/apps/aether-gateway/src/tests/video/xai.rs @@ -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("video@example.com".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(repository: Arc) +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"), + "Aether test frontend", + ) + .unwrap(); + let seen = Arc::new(Mutex::new(Vec::::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(), + "Aether test frontend" + ); + 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(); +} diff --git a/apps/aether-gateway/src/video_tasks/tests/plans.rs b/apps/aether-gateway/src/video_tasks/tests/plans.rs index e40ec578a..154910309 100644 --- a/apps/aether-gateway/src/video_tasks/tests/plans.rs +++ b/apps/aether-gateway/src/video_tasks/tests/plans.rs @@ -11,6 +11,9 @@ use super::{ fn rust_authoritative_service_builds_openai_cancel_follow_up_plan() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), 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() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), 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() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), 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() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), 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() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-active-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), 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"), })); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-completed-123".to_string(), upstream_task_id: "ext-video-task-999".to_string(), 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"); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-file-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), created_at_unix_ms: 1712345678, diff --git a/apps/aether-gateway/src/video_tasks/tests/projection.rs b/apps/aether-gateway/src/video_tasks/tests/projection.rs index fa7978bfd..fed35f5d1 100644 --- a/apps/aether-gateway/src/video_tasks/tests/projection.rs +++ b/apps/aether-gateway/src/video_tasks/tests/projection.rs @@ -10,6 +10,9 @@ use super::{ fn rust_authoritative_service_projects_openai_status_into_local_read_response() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), 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() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), 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() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), created_at_unix_ms: 1712345678, diff --git a/apps/aether-gateway/src/video_tasks/tests/sync.rs b/apps/aether-gateway/src/video_tasks/tests/sync.rs index 9634d322a..7d85ef8d8 100644 --- a/apps/aether-gateway/src/video_tasks/tests/sync.rs +++ b/apps/aether-gateway/src/video_tasks/tests/sync.rs @@ -218,6 +218,9 @@ fn rust_authoritative_video_truth_source_can_background_success_report() { fn rust_authoritative_service_reads_openai_task_from_local_registry() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); let snapshot = LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), 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() { let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative); service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-local-123".to_string(), upstream_task_id: "ext-video-task-123".to_string(), created_at_unix_ms: 1712345678, diff --git a/crates/aether-admin/src/provider/quota.rs b/crates/aether-admin/src/provider/quota.rs index 2a4647159..1ed101ad8 100644 --- a/crates/aether-admin/src/provider/quota.rs +++ b/crates/aether-admin/src/provider/quota.rs @@ -3546,6 +3546,258 @@ pub fn parse_kiro_usage_response( Some(serde_json::Value::Object(result)) } +pub fn parse_xai_billing_response( + value: &serde_json::Value, + updated_at_unix_secs: u64, +) -> Option { + 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, + key: &str, +) -> Option { + object.get(key).and_then(coerce_json_f64) +} + +fn coerce_json_bool_from_map( + object: &serde_json::Map, + key: &str, +) -> Option { + object.get(key).and_then(coerce_json_bool) +} + +fn extract_xai_product_usage_percent( + config: &serde_json::Map, +) -> Option { + 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 { + 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 { + 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 { + 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( value: &serde_json::Value, updated_at_unix_secs: u64, diff --git a/crates/aether-admin/src/provider/state.rs b/crates/aether-admin/src/provider/state.rs index 1d4a40bf4..b7cbc2ade 100644 --- a/crates/aether-admin/src/provider/state.rs +++ b/crates/aether-admin/src/provider/state.rs @@ -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 let Some(organization_uuid) = token_payload_object .get("organization") @@ -554,6 +571,28 @@ mod tests { 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": "grok@x.ai", + "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!("grok@x.ai"))); + assert_eq!(auth_config.get("sub"), Some(&json!("user-xai-1"))); + assert_eq!(auth_config.get("id_token"), Some(&json!(id_token))); + } + #[test] fn decode_jwt_claims_rejects_oversized_payload_before_decode() { let max_encoded_len = MAX_UNVERIFIED_JWT_CLAIMS_BYTES diff --git a/crates/aether-ai/formats/src/api.rs b/crates/aether-ai/formats/src/api.rs index de48589f7..1f0a310bd 100644 --- a/crates/aether-ai/formats/src/api.rs +++ b/crates/aether-ai/formats/src/api.rs @@ -208,6 +208,10 @@ pub use crate::formats::{ resolve_stream_spec as resolve_openai_responses_stream_spec, 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::{ diff --git a/crates/aether-ai/formats/src/formats/openai/responses/mod.rs b/crates/aether-ai/formats/src/formats/openai/responses/mod.rs index 36053f67e..1d22c1cd1 100644 --- a/crates/aether-ai/formats/src/formats/openai/responses/mod.rs +++ b/crates/aether-ai/formats/src/formats/openai/responses/mod.rs @@ -7,6 +7,7 @@ pub mod request; pub mod response; pub mod spec; pub mod stream; +pub mod xai; const TOOL_ERROR_PREFIX: &str = "[tool error]"; const AETHER_REASONING_ITEM_ID_PREFIX: &str = "rs_aether_"; @@ -85,6 +86,8 @@ pub enum OpenAiResponsesReasoningReplayPolicy { #[default] OpenAiItemIds, 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. @@ -234,6 +237,14 @@ fn openai_responses_reasoning_item_is_replayable( { 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 .get("id") .and_then(Value::as_str) @@ -334,6 +345,36 @@ mod tests { 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] fn gemini_tool_signature_carrier_roundtrips_direction_and_exact_value() { let signature = " opaque-signature-with-padding== "; diff --git a/crates/aether-ai/formats/src/formats/openai/responses/xai.rs b/crates/aether-ai/formats/src/formats/openai/responses/xai.rs new file mode 100644 index 000000000..212f5212d --- /dev/null +++ b/crates/aether-ai/formats/src/formats/openai/responses/xai.rs @@ -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, 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 { + tools + .iter() + .filter_map(|tool| normalize_xai_tool(tool, keep_image_generation)) + .collect() +} + +fn normalize_xai_tool(tool: &Value, keep_image_generation: bool) -> Option { + 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) { + 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) { + 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) { + 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) { + 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) { + 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::>(); + 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) -> Vec { + 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) { + 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) { + 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) -> 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) { + 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 { + 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) { + 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) -> &[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) -> 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"); + } +} diff --git a/crates/aether-ai/formats/src/formats/shared/routing.rs b/crates/aether-ai/formats/src/formats/shared/routing.rs index 3dec3d5b7..c85aec418 100644 --- a/crates/aether-ai/formats/src/formats/shared/routing.rs +++ b/crates/aether-ai/formats/src/formats/shared/routing.rs @@ -49,6 +49,10 @@ pub fn resolve_execution_runtime_stream_plan_kind_with_client_surface( method: &Method, path: &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") { return None; } @@ -181,6 +185,10 @@ pub fn resolve_execution_runtime_sync_plan_kind_with_client_surface( method: &Method, path: &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") { return None; } @@ -206,7 +214,10 @@ pub fn resolve_execution_runtime_sync_plan_kind_with_client_surface( if route_family == Some("openai") && route_kind == Some("video") && *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); } diff --git a/crates/aether-ai/formats/src/formats/shared/standard_matrix.rs b/crates/aether-ai/formats/src/formats/shared/standard_matrix.rs index e46b2d808..c85760d3b 100644 --- a/crates/aether-ai/formats/src/formats/shared/standard_matrix.rs +++ b/crates/aether-ai/formats/src/formats/shared/standard_matrix.rs @@ -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_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::{ build_local_openai_chat_request_body_with_model_directives, 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, reasoning_replay_policy: crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy, ) -> Option { + 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() .with_mapped_model(mapped_model) .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, provider_api_format, ); - // DeepSeek's Responses continuation state is opaque. Parsing a same-wire-format - // request through the canonical model would discard its id-less `reasoning_text` - // items and future provider-owned fields even though no conversion is required. - // Keep that provider-specific route wire-preserving, while retaining canonical - // normalization for ordinary OpenAI Responses and for Responses/Compact - // cross-format conversions. - let mut provider_request_body = if is_wire_preserving_deepseek_responses_hop( + // DeepSeek and xAI replay opaque provider state. Preserve their native + // Responses input items: canonical conversion can lose reasoning IDs and + // encrypted-only items even when source and destination formats are equal. + let mut provider_request_body = if is_wire_preserving_responses_hop( source_api_format.as_ref(), provider_api_format, reasoning_replay_policy, @@ -200,6 +203,13 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and &mut provider_request_body, 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( &mut provider_request_body, provider_api_format, @@ -224,14 +234,16 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and Some(provider_request_body) } -fn is_wire_preserving_deepseek_responses_hop( +fn is_wire_preserving_responses_hop( source_api_format: &str, provider_api_format: &str, reasoning_replay_policy: crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy, ) -> bool { - if reasoning_replay_policy - != crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque - { + if !matches!( + reasoning_replay_policy, + crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque + | crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted + ) { return false; } 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"); } + + #[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()); + } } diff --git a/crates/aether-ai/formats/src/lib.rs b/crates/aether-ai/formats/src/lib.rs index be40c6d6f..4496e1900 100644 --- a/crates/aether-ai/formats/src/lib.rs +++ b/crates/aether-ai/formats/src/lib.rs @@ -56,6 +56,10 @@ pub use formats::openai::responses::codex::{ pub use formats::openai::responses::request::{ 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::{ normalize_openai_responses_message_item_ids, openai_responses_message_item_id, openai_responses_request_operation, openai_responses_synthetic_reasoning_item_id, diff --git a/crates/aether-data/adapters/postgres/src/candidate_selection.rs b/crates/aether-data/adapters/postgres/src/candidate_selection.rs index 0b1650263..dddec6d1b 100644 --- a/crates/aether-data/adapters/postgres/src/candidate_selection.rs +++ b/crates/aether-data/adapters/postgres/src/candidate_selection.rs @@ -102,6 +102,11 @@ INNER JOIN LATERAL ( AND LOWER(BTRIM(pak.auth_type)) = 'oauth' 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 ( LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') AND LOWER(BTRIM(pak.auth_type)) = 'oauth' @@ -127,7 +132,8 @@ INNER JOIN LATERAL ( 'vertex_ai', 'antigravity', 'kiro', - 'windsurf' + 'windsurf', + 'xai' ) 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($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 ( LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') AND LOWER(BTRIM(pak.auth_type)) = 'oauth' @@ -212,7 +223,8 @@ WHERE p.is_active = TRUE 'vertex_ai', 'antigravity', 'kiro', - 'windsurf' + 'windsurf', + 'xai' ) AND LOWER(BTRIM(pak.auth_type)) <> 'oauth' ) @@ -365,6 +377,11 @@ INNER JOIN LATERAL ( AND LOWER(BTRIM(pak.auth_type)) = 'oauth' 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 ( LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') AND LOWER(BTRIM(pak.auth_type)) = 'oauth' @@ -390,7 +407,8 @@ INNER JOIN LATERAL ( 'vertex_ai', 'antigravity', 'kiro', - 'windsurf' + 'windsurf', + 'xai' ) 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($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 ( LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') AND LOWER(BTRIM(pak.auth_type)) = 'oauth' @@ -476,7 +499,8 @@ WHERE p.is_active = TRUE 'vertex_ai', 'antigravity', 'kiro', - 'windsurf' + 'windsurf', + 'xai' ) AND LOWER(BTRIM(pak.auth_type)) <> 'oauth' ) @@ -632,11 +656,16 @@ WHERE p.is_active = TRUE ) ) ) - OR ( - LOWER(BTRIM(p.provider_type)) = 'grok' - AND LOWER(BTRIM(pak.auth_type)) = 'oauth' - AND LOWER($6) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image') - ) + OR ( + LOWER(BTRIM(p.provider_type)) = 'grok' + AND LOWER(BTRIM(pak.auth_type)) = 'oauth' + 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 ( LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') AND LOWER(BTRIM(pak.auth_type)) = 'oauth' @@ -662,7 +691,8 @@ WHERE p.is_active = TRUE 'vertex_ai', 'antigravity', 'kiro', - 'windsurf' + 'windsurf', + 'xai' ) 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] fn candidate_selection_sql_allows_windsurf_openai_chat_managed_keys() { let requested_model_sql = requested_model_selection_sql(); diff --git a/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs b/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs index 665803064..cf70cc1bb 100644 --- a/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs +++ b/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs @@ -346,6 +346,16 @@ fn key_auth_channel_matches(row: &StoredMinimalCandidateSelectionRow, api_format "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" => { matches!(auth_type.as_str(), "oauth" | "api_key" | "bearer") && api_format == "openai:chat" @@ -591,6 +601,59 @@ mod tests { 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] async fn requested_model_filter_respects_endpoint_scoped_default_mapping() { let mut selected = sample_row("provider-1", "openai:chat", "deepseek-v4-pro", 10); diff --git a/crates/aether-model-fetch/src/logic.rs b/crates/aether-model-fetch/src/logic.rs index 9b8523921..270b13ba9 100644 --- a/crates/aether-model-fetch/src/logic.rs +++ b/crates/aether-model-fetch/src/logic.rs @@ -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 { matches!( 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> { 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"), ], + "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, }; Some(models) @@ -1977,4 +1993,39 @@ mod tests { assert_eq!(models[15]["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::>(); + 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")); + } } diff --git a/crates/aether-oauth/src/provider/providers/generic.rs b/crates/aether-oauth/src/provider/providers/generic.rs index d3f761573..cfe85b19d 100644 --- a/crates/aether-oauth/src/provider/providers/generic.rs +++ b/crates/aether-oauth/src/provider/providers/generic.rs @@ -150,6 +150,27 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[ uses_json_payload: false, 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)] @@ -212,6 +233,10 @@ impl GenericProviderOAuthAdapter { self } + pub(super) fn token_url_for_provider(&self) -> String { + self.token_url() + } + fn token_url(&self) -> String { self.token_url_override .clone() @@ -389,7 +414,10 @@ impl GenericProviderOAuthAdapter { self.token_set_from_payload(payload) } - fn token_set_from_payload(&self, payload: Value) -> Result { + pub(super) fn token_set_from_payload( + &self, + payload: Value, + ) -> Result { let token_set = OAuthTokenSet::from_token_payload(payload.clone()) .ok_or_else(|| OAuthError::invalid_response("token response missing access_token"))?; let mut auth_config = serde_json::Map::new(); @@ -945,6 +973,7 @@ mod tests { fn resolves_generic_provider_templates() { assert!(template_for_provider_type("codex").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()); } diff --git a/crates/aether-oauth/src/provider/providers/mod.rs b/crates/aether-oauth/src/provider/providers/mod.rs index f6fe878a8..f0cfeea68 100644 --- a/crates/aether-oauth/src/provider/providers/mod.rs +++ b/crates/aether-oauth/src/provider/providers/mod.rs @@ -4,6 +4,7 @@ mod codex; mod generic; mod kiro; mod windsurf; +mod xai; pub use antigravity::{AntigravityProviderOAuthAdapter, ANTIGRAVITY_USER_INFO_URL}; pub use claude_code::{ @@ -27,3 +28,7 @@ pub use windsurf::{ WindsurfProviderOAuthAdapter, WINDSURF_CLIENT_ID, WINDSURF_PROVIDER_TYPE, 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, +}; diff --git a/crates/aether-oauth/src/provider/providers/xai.rs b/crates/aether-oauth/src/provider/providers/xai.rs new file mode 100644 index 000000000..f936bdd04 --- /dev/null +++ b/crates/aether-oauth/src/provider/providers/xai.rs @@ -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), +} + +#[derive(Clone)] +pub struct XaiProviderOAuthAdapter { + inner: GenericProviderOAuthAdapter, + device_url_override: Option, +} + +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, + token_url: impl Into, + ) -> 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 { + 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 { + 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 { + 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 { + 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 { + 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 { + self.inner.resolve_request_auth(account) + } + + fn account_fingerprint(&self, account: &ProviderOAuthAccount) -> Option { + 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 { + 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 { + 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 { + response + .json_body + .clone() + .or_else(|| serde_json::from_str::(&response.body_text).ok()) +} + +fn oauth_error_code(payload: &Value) -> Option { + 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 { + 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 { + keys.iter().find_map(|key| match payload.get(*key)? { + Value::Number(number) => number.as_u64(), + Value::String(string) => string.trim().parse::().ok(), + _ => None, + }) +} + +fn form_headers() -> BTreeMap { + 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> { + 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::(&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>>, + status_code: u16, + payload: Value, + } + + #[async_trait] + impl OAuthHttpExecutor for ScriptedExecutor { + async fn execute( + &self, + request: OAuthHttpRequest, + ) -> Result { + *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::>(); + 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": "user@x.ai", "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!("user@x.ai")); + 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::>(); + 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::>(); + assert_eq!(fields["client_id"], XAI_CLIENT_ID); + assert_eq!(fields["scope"], XAI_OAUTH_SCOPES.join(" ")); + } +} diff --git a/crates/aether-oauth/src/provider/service.rs b/crates/aether-oauth/src/provider/service.rs index 8b121f5de..1c5edbf4d 100644 --- a/crates/aether-oauth/src/provider/service.rs +++ b/crates/aether-oauth/src/provider/service.rs @@ -21,7 +21,7 @@ impl ProviderOAuthService { use super::providers::{ AntigravityProviderOAuthAdapter, ClaudeCodeProviderOAuthAdapter, CodexProviderOAuthAdapter, GenericProviderOAuthAdapter, KiroProviderOAuthAdapter, - WindsurfProviderOAuthAdapter, + WindsurfProviderOAuthAdapter, XaiProviderOAuthAdapter, }; let mut service = Self::new() @@ -29,7 +29,8 @@ impl ProviderOAuthService { .with_adapter(Arc::new(ClaudeCodeProviderOAuthAdapter::default())) .with_adapter(Arc::new(CodexProviderOAuthAdapter::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"] { if let Some(adapter) = GenericProviderOAuthAdapter::for_provider_type(provider_type) { service = service.with_adapter(Arc::new(adapter)); @@ -144,6 +145,7 @@ mod tests { "antigravity", "kiro", "windsurf", + "xai", ] { assert!( service.adapter(provider_type).is_ok(), diff --git a/crates/aether-provider/pool/src/lib.rs b/crates/aether-provider/pool/src/lib.rs index 2eb0f0529..06fd0ea75 100644 --- a/crates/aether-provider/pool/src/lib.rs +++ b/crates/aether-provider/pool/src/lib.rs @@ -22,18 +22,20 @@ pub use providers::{ build_windsurf_pool_model_configs_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_rate_limit_request_with_base_url, enrich_chatgpt_web_quota_metadata, - grok_mode_id_for_model, grok_pool_tier_from_quota_bucket, grok_quota_window_key_for_model, + build_windsurf_pool_rate_limit_request_with_base_url, build_xai_pool_billing_request, + 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, AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, CodexProviderPoolAdapter, DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, GrokProviderPoolAdapter, KiroPoolQuotaAuthInput, KiroProviderPoolAdapter, UnsupportedQuotaProviderPoolAdapter, - ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH, ANTIGRAVITY_RETRIEVE_USER_QUOTA_SUMMARY_PATH, - CHATGPT_WEB_CONVERSATION_INIT_PATH, CHATGPT_WEB_DEFAULT_BASE_URL, - CODEX_WHAM_RESET_CREDITS_CONSUME_URL, CODEX_WHAM_RESET_CREDITS_URL, CODEX_WHAM_USAGE_URL, - GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH, GEMINI_CLI_USER_AGENT, KIRO_USAGE_LIMITS_PATH, - KIRO_USAGE_SDK_VERSION, WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_PATH, - WINDSURF_USER_STATUS_PATH, + XaiProviderPoolAdapter, ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH, + ANTIGRAVITY_RETRIEVE_USER_QUOTA_SUMMARY_PATH, CHATGPT_WEB_CONVERSATION_INIT_PATH, + CHATGPT_WEB_DEFAULT_BASE_URL, CODEX_WHAM_RESET_CREDITS_CONSUME_URL, + CODEX_WHAM_RESET_CREDITS_URL, CODEX_WHAM_USAGE_URL, GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH, + GEMINI_CLI_USER_AGENT, KIRO_USAGE_LIMITS_PATH, KIRO_USAGE_SDK_VERSION, + WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_PATH, WINDSURF_USER_STATUS_PATH, + XAI_BILLING_PATH, XAI_USER_PATH, }; pub use quota::{ provider_pool_key_account_quota_exhausted, provider_pool_key_model_quota_exhausted, @@ -81,7 +83,8 @@ mod tests { "grok", "kiro", "vertex_ai", - "windsurf" + "windsurf", + "xai" ] ); assert!(service @@ -104,7 +107,8 @@ mod tests { "gemini_cli", "grok", "kiro", - "windsurf" + "windsurf", + "xai" ] ); assert!(service.supports_quota_refresh("codex")); @@ -112,6 +116,7 @@ mod tests { assert!(service.supports_quota_refresh("grok")); assert!(service.supports_quota_refresh("gemini_cli")); assert!(service.supports_quota_refresh("windsurf")); + assert!(service.supports_quota_refresh("xai")); assert_eq!( service.quota_refresh_unsupported_message("claude_code"), "Claude Code 暂不支持自动刷新额度:上游没有稳定可用的账号额度查询接口" @@ -642,11 +647,11 @@ mod tests { assert_eq!( free_first["providers"], - json!(["codex", "grok", "kiro", "windsurf"]) + json!(["codex", "grok", "kiro", "windsurf", "xai"]) ); assert_eq!( recent_refresh["providers"], - json!(["codex", "grok", "kiro", "windsurf"]) + json!(["codex", "grok", "kiro", "windsurf", "xai"]) ); assert_eq!(free_first["default_enabled"], json!(false)); assert_eq!(recent_refresh["default_enabled"], json!(false)); diff --git a/crates/aether-provider/pool/src/providers/mod.rs b/crates/aether-provider/pool/src/providers/mod.rs index 72734c81e..9260caaf9 100644 --- a/crates/aether-provider/pool/src/providers/mod.rs +++ b/crates/aether-provider/pool/src/providers/mod.rs @@ -7,6 +7,7 @@ pub mod grok; pub mod kiro; pub mod unsupported; pub mod windsurf; +pub mod xai; pub use antigravity::AntigravityProviderPoolAdapter; pub use antigravity::{ @@ -51,3 +52,7 @@ pub use windsurf::{ WINDSURF_DEFAULT_BASE_URL, WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_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, +}; diff --git a/crates/aether-provider/pool/src/providers/xai.rs b/crates/aether-provider/pool/src/providers/xai.rs new file mode 100644 index 000000000..4817c7d16 --- /dev/null +++ b/crates/aether-provider/pool/src/providers/xai.rs @@ -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 { + 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) -> 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 { + 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 + })))); + } +} diff --git a/crates/aether-provider/pool/src/service.rs b/crates/aether-provider/pool/src/service.rs index f61b07a7f..a0668c817 100644 --- a/crates/aether-provider/pool/src/service.rs +++ b/crates/aether-provider/pool/src/service.rs @@ -13,8 +13,8 @@ use crate::provider::{ProviderPoolAdapter, ProviderPoolMemberInput}; use crate::providers::{ AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, CodexProviderPoolAdapter, DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, GrokProviderPoolAdapter, - KiroProviderPoolAdapter, WindsurfProviderPoolAdapter, CLAUDE_CODE_PROVIDER_POOL_ADAPTER, - VERTEX_AI_PROVIDER_POOL_ADAPTER, + KiroProviderPoolAdapter, WindsurfProviderPoolAdapter, XaiProviderPoolAdapter, + CLAUDE_CODE_PROVIDER_POOL_ADAPTER, VERTEX_AI_PROVIDER_POOL_ADAPTER, }; #[derive(Clone)] @@ -55,6 +55,7 @@ impl ProviderPoolService { .with_adapter(Arc::new(KiroProviderPoolAdapter)) .with_adapter(Arc::new(ChatGptWebProviderPoolAdapter)) .with_adapter(Arc::new(WindsurfProviderPoolAdapter)) + .with_adapter(Arc::new(XaiProviderPoolAdapter)) .with_adapter(Arc::new(VERTEX_AI_PROVIDER_POOL_ADAPTER)) } diff --git a/crates/aether-provider/transport/src/conversion.rs b/crates/aether-provider/transport/src/conversion.rs index 90d02cd75..421f82557 100644 --- a/crates/aether-provider/transport/src/conversion.rs +++ b/crates/aether-provider/transport/src/conversion.rs @@ -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] fn windsurf_openai_chat_anchor_supports_cross_format_conversion_via_cascade() { let mut transport = transport_snapshot("windsurf", "openai:chat", "oauth", true, None); diff --git a/crates/aether-provider/transport/src/lib.rs b/crates/aether-provider/transport/src/lib.rs index 9fa3ab2ee..21411f369 100644 --- a/crates/aether-provider/transport/src/lib.rs +++ b/crates/aether-provider/transport/src/lib.rs @@ -30,6 +30,7 @@ pub mod url; pub mod vertex; mod video; pub mod windsurf; +pub mod xai; pub use aether_oauth as oauth; pub use agent_identity::{ @@ -195,3 +196,10 @@ pub use windsurf::{ local_windsurf_request_transport_unsupported_reason_with_network, GET_CHAT_MESSAGE_PATH, 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, +}; diff --git a/crates/aether-provider/transport/src/openai_image/mod.rs b/crates/aether-provider/transport/src/openai_image/mod.rs index 2170611c6..2390d274d 100644 --- a/crates/aether-provider/transport/src/openai_image/mod.rs +++ b/crates/aether-provider/transport/src/openai_image/mod.rs @@ -84,6 +84,11 @@ fn is_dedicated_openai_image_provider(transport: &GatewayProviderTransportSnapsh .trim() .eq_ignore_ascii_case("codex") || is_grok_provider_transport(transport) + || transport + .provider + .provider_type + .trim() + .eq_ignore_ascii_case("xai") } pub fn resolve_openai_image_auth( @@ -92,7 +97,10 @@ pub fn resolve_openai_image_auth( if is_grok_provider_transport(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( @@ -100,7 +108,11 @@ pub fn build_openai_image_upstream_url( request_path: Option<&str>, request_query: Option<&str>, ) -> 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( @@ -113,6 +125,11 @@ pub fn build_openai_image_headers( &BTreeMap::new(), ); 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 { provider_request_headers.insert("accept".to_string(), accept.to_string()); } 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] fn codex_is_supported_by_dedicated_openai_image_transport_policy() { let mut transport = sample_transport(); diff --git a/crates/aether-provider/transport/src/provider_types.rs b/crates/aether-provider/transport/src/provider_types.rs index 4e9045c8d..cac278079 100644 --- a/crates/aether-provider/transport/src/provider_types.rs +++ b/crates/aether-provider/transport/src/provider_types.rs @@ -275,6 +275,17 @@ const WINDSURF_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy { ..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 { provider_type: "claude_code", version: 2, @@ -446,6 +457,39 @@ const WINDSURF_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTem 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 { 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), "antigravity" => Some(&ANTIGRAVITY_FIXED_PROVIDER_TEMPLATE), "windsurf" => Some(&WINDSURF_FIXED_PROVIDER_TEMPLATE), + "xai" => Some(&XAI_FIXED_PROVIDER_TEMPLATE), _ => None, } } @@ -613,6 +658,16 @@ pub fn provider_type_admin_oauth_template(provider_type: &str) -> Option 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, } } @@ -825,6 +880,50 @@ mod tests { 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![ + "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] fn fixed_provider_key_inheritance_keeps_oauth_and_kiro_configured_bearer_keys_open() { assert!(fixed_provider_key_inherits_api_formats( diff --git a/crates/aether-provider/transport/src/request_body.rs b/crates/aether-provider/transport/src/request_body.rs index 3fda480cb..185180985 100644 --- a/crates/aether-provider/transport/src/request_body.rs +++ b/crates/aether-provider/transport/src/request_body.rs @@ -42,6 +42,11 @@ pub fn apply_transport_request_body_semantics( { 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) { apply_vertex_gemini_embedding_body_semantics(provider_request_body)?; } diff --git a/crates/aether-provider/transport/src/request_url/mod.rs b/crates/aether-provider/transport/src/request_url/mod.rs index 22dc905c6..43168999c 100644 --- a/crates/aether-provider/transport/src/request_url/mod.rs +++ b/crates/aether-provider/transport/src/request_url/mod.rs @@ -120,6 +120,12 @@ fn build_transport_request_url_inner( 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 .endpoint .custom_path @@ -164,7 +170,7 @@ fn build_transport_request_url_inner( path.to_string() }; let mut url = build_passthrough_path_url( - &transport.endpoint.base_url, + request_base_url, normalized_path.as_str(), params.request_query, blocked_keys, @@ -190,75 +196,68 @@ fn build_transport_request_url_inner( let url = match normalized_provider_api_format.as_str() { "openai:chat" => Some(build_openai_chat_url( - &transport.endpoint.base_url, + request_base_url, params.request_query, )), "openai:responses" => Some(build_openai_responses_url( - &transport.endpoint.base_url, + request_base_url, params.request_query, false, )), "openai:responses:compact" => Some(build_openai_responses_url( - &transport.endpoint.base_url, + request_base_url, params.request_query, true, )), "openai:search" => Some(build_openai_search_url( - &transport.endpoint.base_url, + request_base_url, params.request_query, )), "openai:realtime" => build_passthrough_path_url( - &transport.endpoint.base_url, + request_base_url, "/v1/realtime", params.request_query, GATEWAY_CREDENTIAL_QUERY_KEYS, ) .and_then(|url| replace_realtime_model_query(url, params.mapped_model?)), "codex:live" => build_passthrough_path_url( - &transport.endpoint.base_url, + request_base_url, "/live", params.request_query, GATEWAY_CREDENTIAL_QUERY_KEYS, ), "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" => { - 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 { - build_default_claude_count_tokens_url( - &transport.endpoint.base_url, - params.request_query, - ) + build_default_claude_count_tokens_url(request_base_url, params.request_query) } 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( - &transport.endpoint.base_url, + request_base_url, params.mapped_model?, params.upstream_is_stream, params.request_query, ), "gemini:embedding" => build_gemini_embedding_url( - &transport.endpoint.base_url, + request_base_url, params.mapped_model?, params.request_query, gemini_embedding_batch, ), "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, }?; @@ -2417,4 +2416,82 @@ mod tests { "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"); + } } diff --git a/crates/aether-provider/transport/src/standard/mod.rs b/crates/aether-provider/transport/src/standard/mod.rs index 56b5c0402..c5ed29970 100644 --- a/crates/aether-provider/transport/src/standard/mod.rs +++ b/crates/aether-provider/transport/src/standard/mod.rs @@ -396,6 +396,12 @@ pub fn build_standard_provider_request_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 = crate::headers::declared_connection_header_names(input.headers, input.extra_headers); crate::headers::remove_declared_connection_headers(&mut headers, &declared_connection_headers); diff --git a/crates/aether-provider/transport/src/video/mod.rs b/crates/aether-provider/transport/src/video/mod.rs index 85dfa0cd9..69af710bd 100644 --- a/crates/aether-provider/transport/src/video/mod.rs +++ b/crates/aether-provider/transport/src/video/mod.rs @@ -1,6 +1,7 @@ use std::collections::BTreeMap; use std::fmt; +use aether_contracts::ProxySnapshot; use aether_data_contracts::repository::video_tasks::StoredVideoTask; use aether_video_tasks_core::{ LocalVideoTaskSnapshot, LocalVideoTaskTransport, LocalVideoTaskTransportBridgeInput, @@ -12,11 +13,13 @@ use super::auth::{ build_passthrough_headers_with_auth, resolve_local_gemini_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::{ local_gemini_transport_unsupported_reason_with_network, - local_standard_transport_unsupported_reason_with_network, supports_local_gemini_transport, - supports_local_standard_transport, + local_standard_transport_unsupported_reason_with_network, }; use super::rules::{ apply_local_body_rules_with_request_headers, apply_local_header_rules_with_request_headers, @@ -32,6 +35,7 @@ pub enum ProviderVideoCreateFamily { #[derive(Clone, Copy)] pub struct ProviderVideoCreateHeadersInput<'a> { + pub transport: &'a GatewayProviderTransportSnapshot, pub headers: &'a http::HeaderMap, pub auth_header: &'a str, pub auth_value: &'a str, @@ -79,6 +83,13 @@ pub trait VideoTaskTransportSnapshotLookup: Send + Sync { endpoint_id: &str, key_id: &str, ) -> Result, String>; + + async fn resolve_video_task_proxy( + &self, + transport: &GatewayProviderTransportSnapshot, + ) -> Option { + resolve_transport_proxy_snapshot(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 (auth_header, auth_value) = match api_format { "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; } - resolve_local_openai_bearer_auth(transport)? + resolve_openai_compatible_video_auth(transport)? } "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; } resolve_local_gemini_auth(transport)? @@ -103,9 +118,9 @@ pub fn resolve_local_video_task_transport( _ => return None, }; - Some(LocalVideoTaskTransport::from_bridge_input( - LocalVideoTaskTransportBridgeInput { - upstream_base_url: transport.endpoint.base_url.clone(), + let mut resolved = + LocalVideoTaskTransport::from_bridge_input(LocalVideoTaskTransportBridgeInput { + upstream_base_url: crate::xai::resolved_xai_request_base_url(transport, api_format), provider_name: Some(transport.provider.name.clone()), provider_id: transport.provider.id.clone(), endpoint_id: transport.endpoint.id.clone(), @@ -114,11 +129,12 @@ pub fn resolve_local_video_task_transport( auth_value, content_type: Some("application/json".to_string()), model_name, - proxy: None, + proxy: resolve_transport_proxy_snapshot(transport), transport_profile: resolve_transport_profile(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( @@ -141,7 +157,7 @@ pub fn resolve_video_create_auth( family: ProviderVideoCreateFamily, ) -> Option<(String, String)> { 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), } } @@ -173,6 +189,15 @@ pub fn build_video_create_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( transport: &GatewayProviderTransportSnapshot, request_path: &str, @@ -193,7 +218,13 @@ pub fn build_video_create_upstream_url( ProviderVideoCreateFamily::Gemini => &["key"][..], }; 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, request_query, blocked_keys, @@ -202,8 +233,14 @@ pub fn build_video_create_upstream_url( match family { ProviderVideoCreateFamily::OpenAi => build_passthrough_path_url( - &transport.endpoint.base_url, - openai_video_api_root_request_path(request_path), + &crate::xai::resolved_xai_request_base_url(transport, "openai:video"), + 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, &[], ), @@ -216,6 +253,7 @@ pub fn build_video_create_upstream_url( } 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/") { &request_path[3..] } else { @@ -232,6 +270,11 @@ pub fn build_video_create_headers( input.auth_value, &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( &mut provider_request_headers, input.header_rules, @@ -281,16 +324,22 @@ pub async fn reconstruct_local_video_task_snapshot( return Ok(None); }; - let Some(local_transport) = + let Some(mut local_transport) = resolve_local_video_task_transport(&transport, provider_api_format, task.model.clone()) else { return Ok(None); }; - Ok(LocalVideoTaskSnapshot::from_stored_task_with_transport( - task, - local_transport, - )) + // Resolve deployment-managed nodes, system defaults and tunnel affinity just as + // creation does; serialized task metadata intentionally contains no credentials. + 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)] @@ -441,6 +490,46 @@ mod tests { 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] fn resolves_gemini_video_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"); } + #[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] fn builds_gemini_video_create_url_and_removes_client_key_query() { let transport = sample_transport("gemini:video", "api_key"); @@ -512,6 +715,7 @@ mod tests { let provider_request_body = json!({"prompt": "make a clip"}); let original_request_body = provider_request_body.clone(); let headers = build_video_create_headers(ProviderVideoCreateHeadersInput { + transport: &sample_transport("openai:video", "bearer"), headers: &http::HeaderMap::new(), auth_header: "authorization", auth_value: "Bearer secret", diff --git a/crates/aether-provider/transport/src/xai.rs b/crates/aether-provider/transport/src/xai.rs new file mode 100644 index 000000000..5dbeae020 --- /dev/null +++ b/crates/aether-provider/transport/src/xai.rs @@ -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 { + 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) { + 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, +) { + 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 { + 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 { + 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 { + 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 { + 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::().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 { + raw_auth_config + .map(str::trim) + .filter(|value| !value.is_empty()) + .and_then(|value| serde_json::from_str::(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}"#) + )); + } +} diff --git a/crates/aether-provider/transport/src/xai/video.rs b/crates/aether-provider/transport/src/xai/video.rs new file mode 100644 index 000000000..b715952d2 --- /dev/null +++ b/crates/aether-provider/transport/src/xai/video.rs @@ -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 { + 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::() + .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"); + } +} diff --git a/crates/aether-testing/testkit/src/postgres.rs b/crates/aether-testing/testkit/src/postgres.rs index 957301f6f..76876b5b4 100644 --- a/crates/aether-testing/testkit/src/postgres.rs +++ b/crates/aether-testing/testkit/src/postgres.rs @@ -1,5 +1,8 @@ use std::path::PathBuf; 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::{DataBackends, DataLayerConfig}; @@ -21,10 +24,20 @@ pub struct ManagedPostgresServer { impl ManagedPostgresServer { pub async fn start() -> Result> { 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!( - "aether-postgres-baseline-{}-{}", + "aether-postgres-baseline-{}-{}-{}-{}", std::process::id(), - port + port, + seq, + nanos )); let data_dir = workdir.join("data"); std::fs::create_dir(&workdir)?; diff --git a/crates/aether-usage/runtime/src/runtime.rs b/crates/aether-usage/runtime/src/runtime.rs index cb02c868f..2eec0db56 100644 --- a/crates/aether-usage/runtime/src/runtime.rs +++ b/crates/aether-usage/runtime/src/runtime.rs @@ -6649,7 +6649,7 @@ mod tests { } use std::collections::BTreeMap; - use std::sync::atomic::{AtomicUsize, Ordering}; + use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; use std::time::Instant; @@ -7375,9 +7375,27 @@ mod tests { queue: Arc, policy_started: Arc, release_policy: Arc, + policy_released: Arc, policy_reads: Arc, } + impl BlockingPolicyQueueConfiguredUsageStore { + fn new(queue: Arc) -> 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)] struct FailingPolicyUsageStore { inner: NoRedisUsageStore, @@ -8409,7 +8427,18 @@ mod tests { async fn body_capture_policy(&self) -> Result { self.policy_reads.fetch_add(1, Ordering::AcqRel); 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()) } } @@ -9043,15 +9072,39 @@ mod tests { .await .expect("a duplicate first-byte marker must release the terminal barrier"); - let records = store.records.lock().expect("records lock"); - assert_eq!( - records.len(), - 2, - "the duplicate first byte must be coalesced" - ); - assert_eq!(records[0].status, "streaming"); - assert_eq!(records[1].status, "completed"); - drop(records); + { + let records = store.records.lock().expect("records lock"); + assert_eq!( + records.len(), + 2, + "the duplicate first byte must be coalesced" + ); + assert_eq!(records[0].status, "streaming"); + 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(); 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() { for limit in [0, 64 * 1024] { let runtime = UsageRuntime::new(UsageRuntimeConfig::default()).expect("runtime"); - let store = BlockingPolicyQueueConfiguredUsageStore { - queue: Arc::new(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 store = BlockingPolicyQueueConfiguredUsageStore::new(Arc::new( + RuntimeState::memory(MemoryRuntimeStateConfig::default()), + )); let budget = Arc::new(crate::event_capture_budget::EventCaptureMemoryBudget::new( limit, )); @@ -12391,7 +12441,7 @@ mod tests { .await .expect("replacement policy read starts"); assert_eq!(budget.retained_bytes(), retained); - store.release_policy.notify_one(); + store.release_blocked_policy(); let event = timeout(Duration::from_secs(2), completing) .await .expect("Basic policy completes") @@ -13398,12 +13448,7 @@ mod tests { Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0)); let queue: Arc = tracked_queue.clone(); - let store = BlockingPolicyQueueConfiguredUsageStore { - 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 store = BlockingPolicyQueueConfiguredUsageStore::new(queue); let runtime = UsageRuntime::new(config).expect("usage runtime should build"); let request_id = "req-terminal-seed-waits-for-turn"; let plan = terminal_test_plan(request_id); @@ -13425,9 +13470,10 @@ mod tests { assert_eq!(blocked_snapshot.terminal_submission_in_flight, 0); assert!(blocked_snapshot.lifecycle_submission_pending >= 2); - store.release_policy.notify_waiters(); + store.release_blocked_policy(); timeout(Duration::from_secs(2), async { loop { + store.release_blocked_policy(); let snapshot = runtime.metrics_snapshot(); if tracked_queue.successful_appends.load(Ordering::Acquire) == 1 && snapshot.lifecycle_submission_pending == 0 @@ -13435,7 +13481,7 @@ mod tests { { break; } - tokio::task::yield_now().await; + sleep(Duration::from_millis(1)).await; } }) .await @@ -13468,12 +13514,7 @@ mod tests { Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0)); let queue: Arc = tracked_queue.clone(); - let store = BlockingPolicyQueueConfiguredUsageStore { - 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 store = BlockingPolicyQueueConfiguredUsageStore::new(queue); let runtime = UsageRuntime::new(config).expect("usage runtime should build"); let policy_started = store.policy_started.notified(); @@ -13520,9 +13561,10 @@ mod tests { assert_eq!(blocked_snapshot.terminal_submission_in_flight, 1); assert!(blocked_snapshot.lifecycle_submission_pending <= BACKLOG + 1); - store.release_policy.notify_waiters(); + store.release_blocked_policy(); timeout(Duration::from_secs(5), async { loop { + store.release_blocked_policy(); let snapshot = runtime.metrics_snapshot(); if tracked_queue.successful_appends.load(Ordering::Acquire) == BACKLOG + 1 && snapshot.lifecycle_submission_pending == 0 @@ -13530,7 +13572,7 @@ mod tests { { break; } - tokio::task::yield_now().await; + sleep(Duration::from_millis(1)).await; } }) .await @@ -13828,12 +13870,7 @@ mod tests { Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0)); let queue: Arc = tracked_queue.clone(); - let store = BlockingPolicyQueueConfiguredUsageStore { - 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 store = BlockingPolicyQueueConfiguredUsageStore::new(queue); let runtime = UsageRuntime::new(config).expect("usage runtime should build"); let policy_started = store.policy_started.notified(); runtime @@ -13892,10 +13929,10 @@ mod tests { .expect("terminal submissions should reach the execution backlog"); let saturated_snapshot = runtime.metrics_snapshot(); - store.release_policy.notify_waiters(); + store.release_blocked_policy(); let all_completed = timeout(Duration::from_secs(2), async { loop { - store.release_policy.notify_waiters(); + store.release_blocked_policy(); if tracked_queue.successful_appends.load(Ordering::Acquire) == EXCESS_SUBMISSIONS + 1 && runtime.metrics_snapshot().terminal_submission_in_flight == 0 diff --git a/crates/aether-video-tasks-core/Cargo.toml b/crates/aether-video-tasks-core/Cargo.toml index 25f284b52..f50899619 100644 --- a/crates/aether-video-tasks-core/Cargo.toml +++ b/crates/aether-video-tasks-core/Cargo.toml @@ -12,5 +12,6 @@ aether-data-contracts.workspace = true async-trait.workspace = true serde.workspace = true serde_json.workspace = true +sha2.workspace = true url.workspace = true uuid.workspace = true diff --git a/crates/aether-video-tasks-core/src/openai.rs b/crates/aether-video-tasks-core/src/openai.rs index 8b0d657cb..5738f07d1 100644 --- a/crates/aether-video-tasks-core/src/openai.rs +++ b/crates/aether-video-tasks-core/src/openai.rs @@ -43,6 +43,27 @@ pub fn map_openai_stored_task_to_read_response( } 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!({ "id": task.id, "object": "video", @@ -57,6 +78,9 @@ fn build_openai_stored_task_body(task: StoredVideoTask, status: VideoTaskStatus) if let Some(prompt) = task.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 { body["size"] = Value::String(size); } @@ -91,21 +115,97 @@ fn map_openai_stored_task_status(status: VideoTaskStatus) -> &'static str { } 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::().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) { + if self.uses_xai_provider() { + self.native_response = Some(Value::Object(provider_body.clone())); + } + let raw_status = provider_body .get("status") .and_then(Value::as_str) .map(str::trim) .unwrap_or_default(); - self.status = match raw_status { - "queued" => LocalVideoTaskStatus::Queued, - "processing" => LocalVideoTaskStatus::Processing, - "completed" => LocalVideoTaskStatus::Completed, - "failed" => LocalVideoTaskStatus::Failed, - "cancelled" => LocalVideoTaskStatus::Cancelled, + // Accept xAI's native lifecycle vocabulary alongside OpenAI's fields. + self.status = match raw_status.to_ascii_lowercase().as_str() { + "queued" | "pending" => LocalVideoTaskStatus::Queued, + "processing" | "in_progress" | "running" => LocalVideoTaskStatus::Processing, + "completed" | "done" | "succeeded" | "success" => LocalVideoTaskStatus::Completed, + "failed" | "error" => LocalVideoTaskStatus::Failed, + "cancelled" | "canceled" => LocalVideoTaskStatus::Cancelled, "expired" => LocalVideoTaskStatus::Expired, _ => 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 .get("progress") .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.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 - .and_then(|value| value.get("code")) - .and_then(Value::as_str) - .map(str::to_string), - ); + self.error_code = sanitize_video_task_error_code(error_code.map(str::to_string)); self.error_message = None; self.video_url = provider_body .get("video_url") .or_else(|| provider_body.get("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) .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( @@ -239,6 +354,9 @@ impl OpenAiVideoTaskSeed { } pub fn client_body_json(&self) -> Value { + if self.is_xai_native() { + return self.native_read_body_json(); + } let mut body = json!({ "id": self.local_task_id, "object": "video", @@ -259,6 +377,9 @@ impl OpenAiVideoTaskSeed { if let Some(seconds) = &self.seconds { 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 { 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 { - if !matches!( + let refreshable = matches!( self.status, LocalVideoTaskStatus::Submitted | LocalVideoTaskStatus::Queued | LocalVideoTaskStatus::Processing - ) { + ) || (self.uses_xai_provider() + && self.native_response.is_none() + && matches!( + self.status, + LocalVideoTaskStatus::Completed + | LocalVideoTaskStatus::Failed + | LocalVideoTaskStatus::Expired + )); + if !refreshable { return None; } @@ -573,7 +702,12 @@ impl OpenAiVideoTaskSeed { }; let mut record = UpsertVideoTask { 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(), user_id: self.user_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())), prompt: self.prompt.clone().or_else(|| Some(String::new())), 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"), aspect_ratio: request_body_string( &self.persistence.original_request_body, @@ -697,6 +835,9 @@ mod tests { #[test] fn builds_minimal_openai_persistence_record_without_sensitive_snapshot() { let seed = OpenAiVideoTaskSeed { + local_short_id: None, + native_response: None, + xai_provider: false, local_task_id: "task-openai-sensitive".to_string(), upstream_task_id: "upstream-openai-sensitive".to_string(), created_at_unix_ms: 1_712_345_678, @@ -747,6 +888,12 @@ mod tests { 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!(record.original_request_body.is_none()); assert!(record.progress_message.is_none()); @@ -759,6 +906,8 @@ mod tests { let mut stored = record.into_stored(); 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 = LocalVideoTaskSnapshot::from_stored_task_with_transport(&stored, seed.transport) .expect("stored task should reconstruct with current transport"); @@ -766,6 +915,21 @@ mod tests { panic!("expected OpenAI snapshot"); }; 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); let Some(LocalVideoTaskContentAction::StreamPlan(plan)) = restored.build_content_stream_action(None, "trace-download") diff --git a/crates/aether-video-tasks-core/src/path.rs b/crates/aether-video-tasks-core/src/path.rs index c266fd4e0..f04251770 100644 --- a/crates/aether-video-tasks-core/src/path.rs +++ b/crates/aether-video-tasks-core/src/path.rs @@ -8,7 +8,9 @@ use uuid::Uuid; use crate::{LocalVideoTaskRegistryMutation, LocalVideoTaskStatus, VideoTaskTruthSourceMode}; 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() || suffix.contains('/') || 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> { - let suffix = path.strip_prefix("/v1/videos/")?; + let suffix = path + .strip_prefix("/v1/videos/") + .or_else(|| path.strip_prefix("/openai/v1/videos/"))?; suffix .strip_suffix("/cancel") .filter(|value| !value.is_empty()) } 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 .strip_suffix("/remix") .filter(|value| !value.is_empty()) } 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 .strip_suffix("/content") .filter(|value| !value.is_empty()) diff --git a/crates/aether-video-tasks-core/src/read_side.rs b/crates/aether-video-tasks-core/src/read_side.rs index fd8749ce1..003dfcdc1 100644 --- a/crates/aether-video-tasks-core/src/read_side.rs +++ b/crates/aether-video-tasks-core/src/read_side.rs @@ -73,7 +73,7 @@ async fn read_openai_video_task_response( } None => state.find_stored_video_task(lookup).await?, }; - let Some(task) = task else { + let Some(mut task) = task else { return Ok(None); }; @@ -81,6 +81,9 @@ async fn read_openai_video_task_response( 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))) } diff --git a/crates/aether-video-tasks-core/src/service.rs b/crates/aether-video-tasks-core/src/service.rs index c94fc00db..a4ba50545 100644 --- a/crates/aether-video-tasks-core/src/service.rs +++ b/crates/aether-video-tasks-core/src/service.rs @@ -105,13 +105,8 @@ impl VideoTaskService { if self.truth_source_mode != VideoTaskTruthSourceMode::RustAuthoritative { return None; } - match route_family { - Some("openai") => extract_openai_task_id_from_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, - } + self.snapshot_for_route(route_family, request_path) + .map(|snapshot| snapshot.read_response_for_path(request_path)) } pub fn read_response_for_user( @@ -126,7 +121,7 @@ impl VideoTaskService { let snapshot = self.snapshot_for_route(route_family, request_path)?; snapshot .belongs_to_user(user_id) - .then(|| snapshot.read_response()) + .then(|| snapshot.read_response_for_path(request_path)) } pub fn snapshot_for_route( diff --git a/crates/aether-video-tasks-core/src/snapshot.rs b/crates/aether-video-tasks-core/src/snapshot.rs index 6aea1cb21..7979dc051 100644 --- a/crates/aether-video-tasks-core/src/snapshot.rs +++ b/crates/aether-video-tasks-core/src/snapshot.rs @@ -29,6 +29,7 @@ impl LocalVideoTaskSnapshot { // contain stale identity fields after a task import or repair. match &mut snapshot { Self::OpenAi(seed) => { + seed.local_short_id = task.short_id.clone(); seed.user_id = task.user_id.clone(); seed.api_key_id = task.api_key_id.clone(); } @@ -51,6 +52,9 @@ impl LocalVideoTaskSnapshot { "openai:video" => { let upstream_task_id = non_empty_owned(task.external_task_id.as_ref())?; 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(), upstream_task_id, 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 { match self { Self::OpenAi(seed) => match seed.status { diff --git a/crates/aether-video-tasks-core/src/sync.rs b/crates/aether-video-tasks-core/src/sync.rs index 2bf49bd57..11ef4b8ce 100644 --- a/crates/aether-video-tasks-core/src/sync.rs +++ b/crates/aether-video-tasks-core/src/sync.rs @@ -19,14 +19,17 @@ impl LocalVideoTaskSeed { ) -> Option { let transport = LocalVideoTaskTransport::from_plan(plan)?; let persistence = LocalVideoTaskPersistence::from_report_context(report_context, plan); - match report_kind { + let mut seed = match report_kind { "openai_video_create_sync_finalize" => { - let upstream_id = provider_body.get("id").and_then(Value::as_str)?.trim(); - if upstream_id.is_empty() { - return None; - } + let upstream_id = openai_video_provider_task_id(provider_body)?; 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") .unwrap_or_else(|| Uuid::new_v4().to_string()), upstream_task_id: upstream_id.to_string(), @@ -37,8 +40,12 @@ impl LocalVideoTaskSeed { model: context_text(report_context, "model") .or_else(|| request_body_text(report_context, "model")), prompt: request_body_text(report_context, "prompt"), - size: request_body_text(report_context, "size"), - seconds: request_body_text(report_context, "seconds"), + size: context_text(report_context, "video_size") + .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, status: LocalVideoTaskStatus::Submitted, progress_percent: 0, @@ -52,12 +59,15 @@ impl LocalVideoTaskSeed { })) } "openai_video_remix_sync_finalize" => { - let upstream_id = provider_body.get("id").and_then(Value::as_str)?.trim(); - if upstream_id.is_empty() { - return None; - } + let upstream_id = openai_video_provider_task_id(provider_body)?; 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") .unwrap_or_else(|| Uuid::new_v4().to_string()), upstream_task_id: upstream_id.to_string(), @@ -68,8 +78,12 @@ impl LocalVideoTaskSeed { model: context_text(report_context, "model") .or_else(|| request_body_text(report_context, "model")), prompt: request_body_text(report_context, "prompt"), - size: request_body_text(report_context, "size"), - seconds: request_body_text(report_context, "seconds"), + size: context_text(report_context, "video_size") + .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") .or_else(|| request_body_text(report_context, "remix_video_id")), status: LocalVideoTaskStatus::Submitted, @@ -110,7 +124,11 @@ impl LocalVideoTaskSeed { })) } _ => 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 { @@ -144,12 +162,28 @@ impl LocalVideoTaskSeed { pub fn client_body_json(&self) -> Value { 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(), } } } +fn openai_video_provider_task_id(body: &Map) -> 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 { pub fn prepare_sync_success( self, @@ -353,6 +387,234 @@ mod tests { 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] fn builds_local_sync_finalize_read_response_for_supported_video_finalize_kinds() { let delete_response = build_local_sync_finalize_read_response( diff --git a/crates/aether-video-tasks-core/src/transport_domain.rs b/crates/aether-video-tasks-core/src/transport_domain.rs index 9503e0087..9f3b267ae 100644 --- a/crates/aether-video-tasks-core/src/transport_domain.rs +++ b/crates/aether-video-tasks-core/src/transport_domain.rs @@ -71,8 +71,16 @@ impl LocalVideoTaskPersistence { .unwrap_or_else(|| plan.request_id.clone()), username: context_text(report_context, "username"), api_key_name: context_text(report_context, "api_key_name"), - client_api_format: context_text(report_context, "client_api_format") - .unwrap_or_else(|| plan.client_api_format.clone()), + client_api_format: if report_context + .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") .unwrap_or_else(|| plan.provider_api_format.clone()), original_request_body: report_context diff --git a/crates/aether-video-tasks-core/src/types.rs b/crates/aether-video-tasks-core/src/types.rs index 81d8a29f1..69e6b255a 100644 --- a/crates/aether-video-tasks-core/src/types.rs +++ b/crates/aether-video-tasks-core/src/types.rs @@ -201,6 +201,13 @@ pub struct LocalVideoTaskPersistence { #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct OpenAiVideoTaskSeed { + /// Preserve existing database identity; older snapshots derive it from the local task ID. + #[serde(default)] + pub local_short_id: Option, + #[serde(default)] + pub native_response: Option, + #[serde(default)] + pub xai_provider: bool, pub local_task_id: String, pub upstream_task_id: String, pub created_at_unix_ms: u64, diff --git a/docs/operations/xai-provider.md b/docs/operations/xai-provider.md new file mode 100644 index 000000000..71aee7f72 --- /dev/null +++ b/docs/operations/xai-provider.md @@ -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 +``` diff --git a/frontend/src/api/endpoints/provider_oauth.ts b/frontend/src/api/endpoints/provider_oauth.ts index 5b0c1c68e..048113212 100644 --- a/frontend/src/api/endpoints/provider_oauth.ts +++ b/frontend/src/api/endpoints/provider_oauth.ts @@ -376,7 +376,7 @@ function jsonValueContainsAgentIdentity(value: unknown): boolean { export interface DeviceAuthorizeRequest { start_url?: 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' redirect_uri?: string proxy_node_id?: string diff --git a/frontend/src/api/endpoints/types/provider.ts b/frontend/src/api/endpoints/types/provider.ts index 0d55c80c1..9aa2bd0e6 100644 --- a/frontend/src/api/endpoints/types/provider.ts +++ b/frontend/src/api/endpoints/types/provider.ts @@ -462,6 +462,23 @@ export interface GrokUpstreamMetadata { 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 { id?: string | null tierType?: string | null @@ -520,6 +537,7 @@ export interface UpstreamMetadata { chatgpt_web?: ChatGPTWebUpstreamMetadata grok?: GrokUpstreamMetadata gemini_cli?: GeminiCliUpstreamMetadata + xai?: XaiUpstreamMetadata } // 按格式的健康度数据 @@ -758,7 +776,7 @@ export interface HealthRelatedMonitorResponse { 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 { // 会话数量控制:null/undefined 表示不限制 diff --git a/frontend/src/features/pool/components/PoolSchedulingDialog.vue b/frontend/src/features/pool/components/PoolSchedulingDialog.vue index f9bd2ea49..2008c7226 100644 --- a/frontend/src/features/pool/components/PoolSchedulingDialog.vue +++ b/frontend/src/features/pool/components/PoolSchedulingDialog.vue @@ -334,7 +334,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [ label: 'Free/Team 优先', description: '兼容旧配置:优先消耗 Free、Team 或两者', evidence_hint: '依据 plan_type,保留旧 free_only/team_only/both 语义', - providers: ['codex', 'grok', 'kiro', 'windsurf'], + providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'], modes: [ { value: 'free_only', label: 'Free' }, { value: 'team_only', label: 'Team' }, @@ -347,7 +347,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [ label: 'Free 优先', description: '优先消耗 Free 账号(依赖 plan_type)', evidence_hint: '依据 plan_type(Free 账号优先调度)', - providers: ['codex', 'grok', 'kiro', 'windsurf'], + providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'], modes: null, default_mode: null, }, @@ -356,7 +356,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [ label: 'Team 优先', description: '优先消耗 Team 账号(依赖 plan_type)', evidence_hint: '依据 plan_type(Team 账号优先调度)', - providers: ['codex', 'grok', 'kiro', 'windsurf'], + providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'], modes: null, default_mode: null, }, @@ -365,7 +365,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [ label: 'Plus 优先', description: '优先消耗 Plus 账号(依赖 plan_type)', evidence_hint: '依据 plan_type(Plus 账号优先调度)', - providers: ['codex', 'grok', 'kiro', 'windsurf'], + providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'], modes: null, default_mode: null, }, @@ -374,7 +374,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [ label: 'Pro 优先', description: '优先消耗 Pro 账号(依赖 plan_type)', evidence_hint: '依据 plan_type(Pro 账号优先调度)', - providers: ['codex', 'grok', 'kiro', 'windsurf'], + providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'], modes: null, default_mode: null, }, @@ -392,7 +392,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [ label: '额度刷新优先', description: '优先选即将刷新额度的账号', evidence_hint: '依据账号额度重置倒计时(next_reset / reset_seconds)', - providers: ['codex', 'grok', 'kiro', 'windsurf'], + providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'], default_enabled_providers: ['codex', 'windsurf'], modes: null, default_mode: null, diff --git a/frontend/src/features/providers/components/OAuthAccountDialog.vue b/frontend/src/features/providers/components/OAuthAccountDialog.vue index 9686c8c9c..38877ae81 100644 --- a/frontend/src/features/providers/components/OAuthAccountDialog.vue +++ b/frontend/src/features/providers/components/OAuthAccountDialog.vue @@ -244,6 +244,129 @@ + + + + +
+ +
+ + + +
+ {{ legacyT('预付额度') }}: {{ formatKiroUsage(getXaiQuotaDisplay(key)?.prepaid_balance) }} +
+ + + +
+
{ + 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 & { usage_percentage?: number usage_limit?: number @@ -2696,6 +2840,28 @@ function shouldAutoRefreshGrokQuota(): boolean { 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 { if (provider.value?.provider_type !== 'windsurf') return false const now = Math.floor(Date.now() / 1000) @@ -2824,7 +2990,7 @@ async function autoRefreshQuotaInBackground(): Promise { if (refreshingQuota.value) return false 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 @@ -2838,6 +3004,8 @@ async function autoRefreshQuotaInBackground(): Promise { shouldRefresh = shouldAutoRefreshKiroQuota() } else if (providerType === 'grok') { shouldRefresh = shouldAutoRefreshGrokQuota() + } else if (providerType === 'xai') { + shouldRefresh = shouldAutoRefreshXaiQuota() } else if (providerType === 'windsurf') { shouldRefresh = shouldAutoRefreshWindsurfQuota() } else if (providerType === 'chatgpt_web') { @@ -2856,6 +3024,8 @@ async function autoRefreshQuotaInBackground(): Promise { hadCachedQuota = allKeys.value.some(({ key }) => key.is_active && hasKiroQuotaDisplayData(key)) } else if (providerType === 'grok') { 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') { hadCachedQuota = allKeys.value.some(({ key }) => key.is_active && hasWindsurfQuotaDisplayData(key)) } else if (providerType === 'chatgpt_web') { diff --git a/frontend/src/features/providers/components/ProviderFormDialog.vue b/frontend/src/features/providers/components/ProviderFormDialog.vue index 727528acf..c8720e027 100644 --- a/frontend/src/features/providers/components/ProviderFormDialog.vue +++ b/frontend/src/features/providers/components/ProviderFormDialog.vue @@ -60,6 +60,9 @@ Grok + + xAI + Kiro @@ -93,6 +96,9 @@ Grok + + xAI + Kiro diff --git a/frontend/src/features/providers/components/__tests__/provider-quota-display.spec.ts b/frontend/src/features/providers/components/__tests__/provider-quota-display.spec.ts index a7b9ae8d4..8eee317a3 100644 --- a/frontend/src/features/providers/components/__tests__/provider-quota-display.spec.ts +++ b/frontend/src/features/providers/components/__tests__/provider-quota-display.spec.ts @@ -54,6 +54,24 @@ describe('provider quota display components', () => { 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', () => { const Probe = defineComponent({ setup() { diff --git a/frontend/src/features/providers/components/provider-tabs/ModelTestDialog.vue b/frontend/src/features/providers/components/provider-tabs/ModelTestDialog.vue index ca9de71af..10ee913cf 100644 --- a/frontend/src/features/providers/components/provider-tabs/ModelTestDialog.vue +++ b/frontend/src/features/providers/components/provider-tabs/ModelTestDialog.vue @@ -1386,6 +1386,7 @@ function formatAuthType(authType: string): string { if (lowered === 'antigravity') return 'Antigravity OAuth' if (lowered === 'kiro') return 'Kiro OAuth' if (lowered === 'grok') return 'Grok OAuth' + if (lowered === 'xai') return 'xAI OAuth' return authType } diff --git a/frontend/src/features/providers/components/provider-tabs/model-test-capabilities.ts b/frontend/src/features/providers/components/provider-tabs/model-test-capabilities.ts index 56148567b..790c65915 100644 --- a/frontend/src/features/providers/components/provider-tabs/model-test-capabilities.ts +++ b/frontend/src/features/providers/components/provider-tabs/model-test-capabilities.ts @@ -39,10 +39,12 @@ const MODEL_TEST_OAUTH_INHERITS_PROVIDER_FORMATS = new Set([ 'vertex_ai', 'antigravity', 'kiro', + 'xai', ]) const MODEL_TEST_BEARER_INHERITS_PROVIDER_FORMATS = new Set([ 'chatgpt_web', + 'xai', ]) const MODEL_TEST_DIAGNOSTIC_LABELS: Record = { diff --git a/frontend/src/features/providers/utils/__tests__/providerTypeUtils.spec.ts b/frontend/src/features/providers/utils/__tests__/providerTypeUtils.spec.ts index 6ef1ece4b..75ef811a5 100644 --- a/frontend/src/features/providers/utils/__tests__/providerTypeUtils.spec.ts +++ b/frontend/src/features/providers/utils/__tests__/providerTypeUtils.spec.ts @@ -16,6 +16,12 @@ describe('providerTypeUtils', () => { 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', () => { expect(isOAuthAccountProviderType('windsurf')).toBe(true) expect(isOAuthAccountProviderType('Windsurf')).toBe(true) diff --git a/frontend/src/features/providers/utils/providerTypeUtils.ts b/frontend/src/features/providers/utils/providerTypeUtils.ts index 2ea917f52..33e475bd1 100644 --- a/frontend/src/features/providers/utils/providerTypeUtils.ts +++ b/frontend/src/features/providers/utils/providerTypeUtils.ts @@ -12,6 +12,7 @@ const oauthAccountProviderTypes = new Set([ 'antigravity', 'kiro', 'grok', + 'xai', 'windsurf', ]) diff --git a/frontend/src/i18n/messages.ts b/frontend/src/i18n/messages.ts index e53bdf21e..dd58b6316 100644 --- a/frontend/src/i18n/messages.ts +++ b/frontend/src/i18n/messages.ts @@ -2146,7 +2146,10 @@ const legacyExactEnglishMessages: Record = { '账号不可用': 'Account unavailable', '日额度': 'Daily quota', '周额度': 'Weekly quota', + '月额度': 'Monthly quota', '剩余额度': 'Remaining quota', + '预付额度': 'Prepaid credits', + '按需额度': 'On-demand credits', '点击编辑优先级': 'Edit priority', '点击编辑倍率': 'Edit multiplier', '同步失败': 'Sync failed', diff --git a/frontend/src/utils/__tests__/providerKeyQuota.spec.ts b/frontend/src/utils/__tests__/providerKeyQuota.spec.ts index e1737d9df..ccfda2f14 100644 --- a/frontend/src/utils/__tests__/providerKeyQuota.spec.ts +++ b/frontend/src/utils/__tests__/providerKeyQuota.spec.ts @@ -359,4 +359,30 @@ describe('providerKeyQuota', () => { }, }, 'windsurf')).toBe('可用模型 3 个') }) + + it('formats xAI weekly credits as remaining percent', () => { + expect(getQuotaDisplayText({ + status_snapshot: { + oauth: { code: 'valid' }, + account: { code: 'ok', blocked: false }, + quota: { + provider_type: 'xai', + code: 'ok', + exhausted: false, + windows: [ + { + code: 'usage', + scope: 'account', + used_ratio: 0.46, + remaining_ratio: 0.54, + }, + { + code: 'prepaid', + remaining_value: 12.5, + }, + ], + }, + }, + }, 'xai')).toBe('剩余 54.0% | 预付剩余 12.5') + }) }) diff --git a/frontend/src/utils/oauth-icons.ts b/frontend/src/utils/oauth-icons.ts index c6c555e4e..5d9d14285 100644 --- a/frontend/src/utils/oauth-icons.ts +++ b/frontend/src/utils/oauth-icons.ts @@ -5,6 +5,7 @@ export const OAUTH_ICONS: Record = { google: ``, gemini_cli: ``, grok: ``, + xai: ``, } // Default icon when provider type is not found diff --git a/frontend/src/utils/providerKeyQuota.ts b/frontend/src/utils/providerKeyQuota.ts index 22f681922..61ba49070 100644 --- a/frontend/src/utils/providerKeyQuota.ts +++ b/frontend/src/utils/providerKeyQuota.ts @@ -265,6 +265,29 @@ function getKiroQuotaText(quota: QuotaStatusSnapshot): string | null { return normalizeText(quota.label) } +function getXaiQuotaText(quota: QuotaStatusSnapshot): string | null { + const parts: string[] = [] + const usageText = getKiroQuotaText(quota) + if (usageText) parts.push(usageText) + + const prepaid = getQuotaWindow(quota, 'prepaid') + if (typeof prepaid?.remaining_value === 'number') { + parts.push(`预付剩余 ${formatQuotaValue(prepaid.remaining_value)}`) + } + + const onDemand = getQuotaWindow(quota, 'on_demand') + const onDemandRemaining = getQuotaWindowRemainingPercent(onDemand) + if (onDemandRemaining != null) { + const valueText = getQuotaWindowValueText(onDemand) + parts.push(`按需剩余 ${formatPercent(onDemandRemaining)}${valueText ? ` (${valueText})` : ''}`) + } else if (typeof onDemand?.remaining_value === 'number') { + parts.push(`按需剩余 ${formatQuotaValue(onDemand.remaining_value)}`) + } + + if (parts.length > 0) return parts.join(' | ') + return normalizeText(quota.label) +} + function getGrokQuotaText(quota: QuotaStatusSnapshot): string | null { const code = normalizeText(quota.code)?.toLowerCase() if (code === 'banned') { @@ -452,6 +475,8 @@ export function getQuotaSnapshotFallbackText( return getCodexQuotaText(quota) case 'kiro': return getKiroQuotaText(quota) + case 'xai': + return getXaiQuotaText(quota) case 'grok': return getGrokQuotaText(quota) case 'windsurf': diff --git a/frontend/src/views/admin/PoolManagement.vue b/frontend/src/views/admin/PoolManagement.vue index a853829b6..1e05aac93 100644 --- a/frontend/src/views/admin/PoolManagement.vue +++ b/frontend/src/views/admin/PoolManagement.vue @@ -1700,6 +1700,7 @@ const showAccountQuotaColumn = computed(() => { || selectedProviderType.value === 'antigravity' || selectedProviderType.value === 'grok' || selectedProviderType.value === 'chatgpt_web' + || selectedProviderType.value === 'xai' }) const desktopColumnWidths = computed(() => { @@ -2149,6 +2150,7 @@ const quotaRefreshSupported = computed(() => { || selectedProviderType.value === 'antigravity' || selectedProviderType.value === 'grok' || selectedProviderType.value === 'chatgpt_web' + || selectedProviderType.value === 'xai' }) function canResetCycleStats(_key: PoolKeyDetail): boolean { @@ -3480,8 +3482,10 @@ function normalizeQuotaLabel(label: string): string { if (/spark/i.test(normalized) && normalized.includes('周')) return 'Spark周' if (normalized.includes('5H')) return '5H' if (normalized.includes('周')) return '周' + if (normalized.includes('月')) return '月' if (normalized.includes('最低剩余')) return '最低' if (normalized === '剩余' || normalized.includes('剩余')) return '剩余' + if (normalized === '额度') return '额度' return normalized } @@ -3490,6 +3494,9 @@ function getQuotaProgressLabel(label: string): string { if (label === '5H') return '5H' if (label === '周') return '周' if (label === '月') return '月' + if (label === '周额度') return '周' + if (label === '月额度') return '月' + if (label === '额度') return '额度' if (label === 'Spark5H') return 'Spark5H' if (label === 'Spark周') return 'Spark周' if (label === '最低') return '最低' @@ -3498,7 +3505,7 @@ function getQuotaProgressLabel(label: string): string { } function getQuotaProgressCountdown(item: QuotaProgressItem) { - const staticResetLabels = ['日', '5H', '周', '月', 'Spark5H', 'Spark周', 'Spark月', 'Auto', 'Fast', 'Expert', 'Heavy', 'Grok 4.3', '生图'] + const staticResetLabels = ['日', '5H', '周', '月', '周额度', '月额度', '额度', 'Spark5H', 'Spark周', 'Spark月', 'Auto', 'Fast', 'Expert', 'Heavy', 'Grok 4.3', '生图'] if (!item.allowDynamicReset && !staticResetLabels.includes(item.label)) return null if (item.resetAtSeconds == null && item.resetSeconds == null) return null return getCodexResetCountdown( @@ -3567,6 +3574,7 @@ function getQuotaLabelOrder(label: string): number { if (label === 'Prompt') return 12 if (label === 'Flex') return 13 if (label === '剩余') return 14 + if (label === '额度') return 14 if (label === '最低') return 15 if (label === '生图') return 16 if (label === '速率') return 17 @@ -3758,7 +3766,7 @@ function buildQuotaProgressItemsFromSnapshot(key: PoolKeyDetail): QuotaProgressI .filter((item): item is QuotaProgressItem => item != null) } - if (providerType === 'kiro') { + if (providerType === 'kiro' || providerType === 'xai') { const quotaResetAtSeconds = getQuotaSnapshotResetAtSeconds(quota) const quotaResetSeconds = getQuotaSnapshotResetSeconds(quota) const window = getQuotaSnapshotWindow(quota, 'usage') @@ -3772,12 +3780,13 @@ function buildQuotaProgressItemsFromSnapshot(key: PoolKeyDetail): QuotaProgressI : undefined return [{ - label: '剩余', + label: normalizeQuotaLabel(String(window?.label || '').trim() || '剩余'), remainingPercent, detail, resetAtSeconds: normalizeUnixSeconds(window?.reset_at ?? quotaResetAtSeconds ?? null), resetSeconds: normalizeRemainingSeconds(window?.reset_seconds ?? quotaResetSeconds ?? null), updatedAtSeconds: getQuotaSnapshotUpdatedAtSeconds(quota), + allowDynamicReset: true, }] }