From e83399db2fd1d16f07d5a6e8c563c9a199c190c1 Mon Sep 17 00:00:00 2001 From: stabey <36232531+stabey@users.noreply.github.com> Date: Mon, 14 Sep 2026 21:09:03 +0800 Subject: [PATCH 1/2] feat(providers): add xAI provider with device code OAuth Add a separate `xai` provider type for xAI Grok CLI subscription accounts. It is independent of the existing `grok` provider, which reverse-proxies grok.com with browser cookies; behavior of `grok` is unchanged. Account binding uses the xAI device code flow, so no local callback listener is needed and headless deployments can bind accounts. Refresh tokens can also be imported individually or in batches, and are rotated on refresh. OAuth requests default to the cli-chat-proxy Responses API; API keys and compact stay on api.x.ai. Explicit custom gateways are preserved. Only `openai:responses` and `openai:responses:compact` are exposed; Chat, Claude and Gemini clients reach the provider through Aether's existing cross-format conversion rather than new native endpoints. Upstream Responses payloads are sanitized for what xAI actually rejects: `previous_response_id` and `metadata.user_id` are dropped, hosted `tool_choice` is rewritten, `web_search` is restored for converted clients, `image_generation` is stripped on older Grok conversation models, unsupported reasoning effort is removed, and requested `reasoning.encrypted_content` is preserved with a replay policy keyed on the configured provider type rather than the model name. 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. The admin UI shows remaining weekly quota as a labeled bar in the provider drawer and the pool list. Co-Authored-By: Claude Opus 5 --- .../passthrough/provider/family/request.rs | 5 + .../ai_serving/planner/standard/deepseek.rs | 25 +- .../planner/standard/openai/responses/mod.rs | 1 + .../aether-gateway/src/ai_serving/pure/mod.rs | 3 +- .../src/ai_serving/transport.rs | 4 + .../oauth/dispatch/batch/execution.rs | 3 +- .../provider/oauth/dispatch/batch/parse.rs | 51 +- .../oauth/dispatch/device/authorize.rs | 17 +- .../provider/oauth/dispatch/device/mod.rs | 1 + .../provider/oauth/dispatch/device/poll.rs | 12 + .../provider/oauth/dispatch/device/xai.rs | 368 +++++++ .../admin/provider/oauth/dispatch/import.rs | 13 +- .../admin/provider/oauth/dispatch/start.rs | 12 + .../provider/oauth/dispatch/token_import.rs | 49 +- .../admin/provider/oauth/quota/dispatch.rs | 18 + .../admin/provider/oauth/quota/mod.rs | 1 + .../admin/provider/oauth/quota/shared.rs | 14 + .../admin/provider/oauth/quota/xai.rs | 313 ++++++ .../admin/provider/pool_admin/payloads.rs | 32 + .../admin/provider/query/models/model_test.rs | 5 + .../admin/provider/write/normalize.rs | 12 +- .../proxy/websocket/responses/continuation.rs | 34 +- .../src/handlers/shared/catalog.rs | 200 ++++ apps/aether-gateway/src/provider_key_auth.rs | 17 + .../src/tests/architecture/admin_provider.rs | 2 + .../src/tests/architecture/ai_serving.rs | 10 + .../src/tests/control/admin/oauth.rs | 291 ++++++ crates/aether-admin/src/provider/quota.rs | 252 +++++ crates/aether-admin/src/provider/state.rs | 39 + crates/aether-ai/formats/src/api.rs | 4 + .../src/formats/openai/responses/mod.rs | 41 + .../src/formats/openai/responses/xai.rs | 914 ++++++++++++++++++ .../src/formats/shared/standard_matrix.rs | 346 ++++++- crates/aether-ai/formats/src/lib.rs | 4 + .../postgres/src/candidate_selection.rs | 56 +- .../repository/candidate_selection/memory.rs | 29 + crates/aether-model-fetch/src/logic.rs | 42 +- .../src/provider/providers/generic.rs | 31 +- .../src/provider/providers/mod.rs | 5 + .../src/provider/providers/xai.rs | 668 +++++++++++++ crates/aether-oauth/src/provider/service.rs | 6 +- crates/aether-provider/pool/src/lib.rs | 29 +- .../aether-provider/pool/src/providers/mod.rs | 5 + .../aether-provider/pool/src/providers/xai.rs | 255 +++++ crates/aether-provider/pool/src/service.rs | 5 +- .../transport/src/conversion.rs | 57 ++ crates/aether-provider/transport/src/lib.rs | 8 + .../transport/src/provider_types.rs | 82 ++ .../transport/src/request_body.rs | 5 + .../transport/src/request_url/mod.rs | 131 ++- .../transport/src/standard/mod.rs | 6 + crates/aether-provider/transport/src/xai.rs | 409 ++++++++ crates/aether-testing/testkit/src/postgres.rs | 17 +- crates/aether-usage/runtime/src/runtime.rs | 121 ++- docs/operations/xai-provider.md | 57 ++ frontend/src/api/endpoints/provider_oauth.ts | 2 +- frontend/src/api/endpoints/types/provider.ts | 20 +- .../pool/components/PoolSchedulingDialog.vue | 12 +- .../components/OAuthAccountDialog.vue | 194 +++- .../components/ProviderDetailDrawer.vue | 174 +++- .../components/ProviderFormDialog.vue | 6 + .../__tests__/provider-quota-display.spec.ts | 18 + .../provider-tabs/ModelTestDialog.vue | 1 + .../provider-tabs/model-test-capabilities.ts | 2 + .../utils/__tests__/providerTypeUtils.spec.ts | 6 + .../providers/utils/providerTypeUtils.ts | 1 + frontend/src/i18n/messages.ts | 3 + .../utils/__tests__/providerKeyQuota.spec.ts | 26 + frontend/src/utils/oauth-icons.ts | 1 + frontend/src/utils/providerKeyQuota.ts | 25 + frontend/src/views/admin/PoolManagement.vue | 15 +- 71 files changed, 5495 insertions(+), 148 deletions(-) create mode 100644 apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/device/xai.rs create mode 100644 apps/aether-gateway/src/handlers/admin/provider/oauth/quota/xai.rs create mode 100644 crates/aether-ai/formats/src/formats/openai/responses/xai.rs create mode 100644 crates/aether-oauth/src/provider/providers/xai.rs create mode 100644 crates/aether-provider/pool/src/providers/xai.rs create mode 100644 crates/aether-provider/transport/src/xai.rs create mode 100644 docs/operations/xai-provider.md 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/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/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/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/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/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..c12b105f4 --- /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-4.6", + "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:responses"); + + 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/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..192026a50 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') + ) 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') + ) 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') + ) 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') + ) 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' ) @@ -637,6 +661,11 @@ WHERE p.is_active = TRUE 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') + ) 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,22 @@ 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'")); + 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..3719865c2 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,13 @@ 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" + ) + } "windsurf" => { matches!(auth_type.as_str(), "oauth" | "api_key" | "bearer") && api_format == "openai:chat" @@ -591,6 +598,28 @@ 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 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..ebd84e031 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,18 @@ 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"), + ], _ => return None, }; Some(models) @@ -1977,4 +1989,32 @@ 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", + ] + ); + assert!(models.iter().all(|model| model["owned_by"] == json!("xai"))); + assert!(models + .iter() + .all(|model| model["api_formats"] == json!(["openai:responses"]))); + } } 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..109e7b47f 100644 --- a/crates/aether-provider/transport/src/conversion.rs +++ b/crates/aether-provider/transport/src/conversion.rs @@ -754,6 +754,63 @@ 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:responses:compact", + "openai:responses" + ), + "compact must not convert onto xAI Responses" + ); + } + + #[test] + fn xai_compact_endpoint_is_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" + ); + } + } + #[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/provider_types.rs b/crates/aether-provider/transport/src/provider_types.rs index 4e9045c8d..42e520c39 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,27 @@ 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, + }, + ], + runtime_policy: XAI_RUNTIME_POLICY, +}; + pub fn provider_type_is_fixed(provider_type: &str) -> bool { provider_runtime_policy(provider_type).fixed_provider } @@ -498,6 +530,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 +646,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 +868,45 @@ mod tests { assert!(ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES.contains(&"windsurf")); } + #[test] + fn xai_fixed_provider_template_exposes_responses_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"] + ); + + 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/xai.rs b/crates/aether-provider/transport/src/xai.rs new file mode 100644 index 000000000..169420d75 --- /dev/null +++ b/crates/aether-provider/transport/src/xai.rs @@ -0,0 +1,409 @@ +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 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-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/docs/operations/xai-provider.md b/docs/operations/xai-provider.md new file mode 100644 index 000000000..1bc7a1f35 --- /dev/null +++ b/docs/operations/xai-provider.md @@ -0,0 +1,57 @@ +# 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. + +## 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 +the fixed-provider endpoint template. OAuth tests cover the device code lifecycle, +token import, and batch import. + +```sh +cargo test -p aether-ai-formats -p aether-provider-transport -p aether-oauth --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, }] } From 04c4a977665599809baf82397090b3838b4f0d1f Mon Sep 17 00:00:00 2001 From: stabey <36232531+stabey@users.noreply.github.com> Date: Mon, 14 Sep 2026 21:17:21 +0800 Subject: [PATCH 2/2] feat(xai): add native image and video endpoints Expose the xAI Imagine image and video surfaces on top of the `xai` provider, and make the shared OpenAI video-task layer survive the production configuration they need. Native video requests live under /v1 (generations, edits, extensions, with /v1/videos as a creation alias that only selects xAI candidates); the OpenAI-compatible adapter stays under /openai/v1/videos and maps `seconds` / `size` onto numeric duration, aspect ratio and resolution. Clients receive an opaque Aether task ID scoped to the owning user; polling uses the upstream task ID and the original credential, and completed downloads fetch the returned media URL without forwarding provider authorization to the media host. Three fixes to the shared video layer are required for this to work outside tests: - OpenAI/xAI task persistence now supplies a stable 16-character short_id, which the PostgreSQL schema requires. Existing rows keep their original value across reconstruction, so no schema change or historical rewrite is needed. - Task retrieval and content downloads are admitted by the production GET execution gate, and reconstructed tasks resolve proxy nodes, system proxy defaults, tunnel affinity and transport profiles through the same deployment resolver used for creation. A configured proxy route no longer silently becomes a direct request after restart. - When the gateway also serves the frontend, /openai/v1/videos and its subpaths bypass the static SPA handler. Otherwise a video query returns HTTP 200 with text/html instead of the task JSON. Co-Authored-By: Claude Opus 5 --- Cargo.lock | 1 + .../planner/specialized/image/request.rs | 11 +- .../planner/specialized/video/decision.rs | 26 +- .../planner/specialized/video/request.rs | 94 +++- apps/aether-gateway/src/api/ai/registry.rs | 2 + apps/aether-gateway/src/async_task/runtime.rs | 3 + apps/aether-gateway/src/constants.rs | 2 + apps/aether-gateway/src/control/route/ai.rs | 6 +- .../src/data/state/testing/video_tasks.rs | 9 + .../src/executor/orchestration.rs | 22 +- .../src/frontdoor_loop_guard.rs | 4 + apps/aether-gateway/src/image_capabilities.rs | 3 +- apps/aether-gateway/src/router.rs | 2 + apps/aether-gateway/src/state/integrations.rs | 8 + apps/aether-gateway/src/tests/video/mod.rs | 38 +- .../src/tests/video/registry_poller.rs | 10 +- .../aether-gateway/src/tests/video/routing.rs | 32 +- apps/aether-gateway/src/tests/video/xai.rs | 427 ++++++++++++++++++ .../src/video_tasks/tests/plans.rs | 21 + .../src/video_tasks/tests/projection.rs | 9 + .../src/video_tasks/tests/sync.rs | 6 + .../src/formats/openai/responses/xai.rs | 4 +- .../formats/src/formats/shared/routing.rs | 13 +- .../postgres/src/candidate_selection.rs | 32 +- .../repository/candidate_selection/memory.rs | 36 +- crates/aether-model-fetch/src/logic.rs | 13 +- .../transport/src/conversion.rs | 45 +- .../transport/src/openai_image/mod.rs | 63 ++- .../transport/src/provider_types.rs | 21 +- .../transport/src/video/mod.rs | 246 +++++++++- crates/aether-provider/transport/src/xai.rs | 35 ++ .../transport/src/xai/video.rs | 147 ++++++ crates/aether-video-tasks-core/Cargo.toml | 1 + crates/aether-video-tasks-core/src/openai.rs | 198 +++++++- crates/aether-video-tasks-core/src/path.rs | 16 +- .../aether-video-tasks-core/src/read_side.rs | 5 +- crates/aether-video-tasks-core/src/service.rs | 11 +- .../aether-video-tasks-core/src/snapshot.rs | 17 + crates/aether-video-tasks-core/src/sync.rs | 290 +++++++++++- .../src/transport_domain.rs | 12 +- crates/aether-video-tasks-core/src/types.rs | 7 + docs/operations/xai-provider.md | 96 +++- 42 files changed, 1880 insertions(+), 164 deletions(-) create mode 100644 apps/aether-gateway/src/tests/video/xai.rs create mode 100644 crates/aether-provider/transport/src/xai/video.rs 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/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/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/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/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/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-ai/formats/src/formats/openai/responses/xai.rs b/crates/aether-ai/formats/src/formats/openai/responses/xai.rs index c12b105f4..212f5212d 100644 --- a/crates/aether-ai/formats/src/formats/openai/responses/xai.rs +++ b/crates/aether-ai/formats/src/formats/openai/responses/xai.rs @@ -872,7 +872,7 @@ mod tests { #[test] fn xai_image_refs_rewrite_openai_aliases_without_touching_chat_parts() { let mut body = json!({ - "model": "grok-4.6", + "model": "grok-imagine-image", "prompt": "edit this", "image": {"image_url": "https://cdn.example/a.png"}, "reference_images": [ @@ -887,7 +887,7 @@ mod tests { }] }); - apply_xai_upstream_payload_edits(&mut body, "xai", "openai:responses"); + 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()); 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-data/adapters/postgres/src/candidate_selection.rs b/crates/aether-data/adapters/postgres/src/candidate_selection.rs index 192026a50..dddec6d1b 100644 --- a/crates/aether-data/adapters/postgres/src/candidate_selection.rs +++ b/crates/aether-data/adapters/postgres/src/candidate_selection.rs @@ -105,7 +105,7 @@ INNER JOIN LATERAL ( 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') + AND LOWER($3) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video') ) OR ( LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') @@ -196,7 +196,7 @@ WHERE p.is_active = TRUE 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') + AND LOWER($3) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video') ) OR ( LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') @@ -380,7 +380,7 @@ INNER JOIN LATERAL ( 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') + AND LOWER($4) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video') ) OR ( LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') @@ -472,7 +472,7 @@ WHERE p.is_active = TRUE 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') + AND LOWER($4) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video') ) OR ( LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity') @@ -656,16 +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)) = 'xai' - AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key') - AND LOWER($6) IN ('openai:responses', 'openai:responses:compact') - ) + 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' @@ -1758,7 +1758,9 @@ mod tests { ] { 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'")); + assert!(sql.contains( + "'openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video'" + )); assert!(sql.contains("'xai'")); } } 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 3719865c2..cf70cc1bb 100644 --- a/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs +++ b/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs @@ -350,7 +350,10 @@ fn key_auth_channel_matches(row: &StoredMinimalCandidateSelectionRow, api_format matches!(auth_type.as_str(), "oauth" | "bearer" | "api_key") && matches!( api_format.as_str(), - "openai:responses" | "openai:responses:compact" + "openai:responses" + | "openai:responses:compact" + | "openai:image" + | "openai:video" ) } "windsurf" => { @@ -620,6 +623,37 @@ mod tests { 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 ebd84e031..270b13ba9 100644 --- a/crates/aether-model-fetch/src/logic.rs +++ b/crates/aether-model-fetch/src/logic.rs @@ -615,6 +615,10 @@ pub fn preset_models_for_provider(provider_type: &str) -> Option> { 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, }; @@ -2010,11 +2014,18 @@ mod tests { "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() - .all(|model| model["api_formats"] == json!(["openai:responses"]))); + .any(|model| model["id"] == "grok-imagine-image")); } } diff --git a/crates/aether-provider/transport/src/conversion.rs b/crates/aether-provider/transport/src/conversion.rs index 109e7b47f..421f82557 100644 --- a/crates/aether-provider/transport/src/conversion.rs +++ b/crates/aether-provider/transport/src/conversion.rs @@ -776,18 +776,26 @@ mod tests { &transport, RequestConversionKind::ToOpenAiResponses )); - assert!( - !request_pair_allowed_for_transport( - &transport, - "openai:responses:compact", - "openai:responses" - ), - "compact must not convert onto xAI Responses" - ); + 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_endpoint_is_same_format_only() { + 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, @@ -809,6 +817,25 @@ mod tests { "{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] 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 42e520c39..cac278079 100644 --- a/crates/aether-provider/transport/src/provider_types.rs +++ b/crates/aether-provider/transport/src/provider_types.rs @@ -474,6 +474,18 @@ const XAI_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate 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, }; @@ -869,7 +881,7 @@ mod tests { } #[test] - fn xai_fixed_provider_template_exposes_responses_endpoints() { + 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); @@ -880,7 +892,12 @@ mod tests { .iter() .map(|item| item.api_format) .collect::>(), - vec!["openai:responses", "openai:responses:compact"] + vec![ + "openai:responses", + "openai:responses:compact", + "openai:image", + "openai:video" + ] ); let policy = provider_runtime_policy("xai"); 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 index 169420d75..5dbeae020 100644 --- a/crates/aether-provider/transport/src/xai.rs +++ b/crates/aether-provider/transport/src/xai.rs @@ -1,3 +1,5 @@ +pub mod video; + use std::collections::BTreeMap; use aether_ai_formats::normalize_api_format_alias; @@ -343,6 +345,39 @@ mod tests { )); } + #[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); 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-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 index 1bc7a1f35..71aee7f72 100644 --- a/docs/operations/xai-provider.md +++ b/docs/operations/xai-provider.md @@ -43,15 +43,105 @@ Quota refresh reads `/user` and `/billing?format=credits` and stores a structure 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 -the fixed-provider endpoint template. OAuth tests cover the device code lifecycle, -token import, and batch import. +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-oauth --lib +cargo test -p aether-ai-formats -p aether-provider-transport -p aether-video-tasks-core --lib cargo test -p aether-gateway --lib xai ```