mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 10:57:03 +08:00
Merge pull request #822 from stabey/upstream-pr/xai-media
feat(providers): 新增 xAI Provider(设备码 OAuth + 原生图像/视频)
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -17,8 +17,8 @@ use crate::ai_serving::transport::{
|
||||
ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput, GROK_CHAT_PATH,
|
||||
};
|
||||
use crate::ai_serving::{
|
||||
apply_codex_openai_special_headers, build_chatgpt_web_image_request_body,
|
||||
build_codex_openai_image_api_provider_request_body,
|
||||
apply_codex_openai_special_headers, apply_xai_upstream_payload_edits,
|
||||
build_chatgpt_web_image_request_body, build_codex_openai_image_api_provider_request_body,
|
||||
build_gemini_image_request_body_from_openai_image_request,
|
||||
build_openai_image_api_provider_request_body, build_openai_image_provider_request_body,
|
||||
default_model_for_openai_image_operation, normalize_openai_image_request,
|
||||
@@ -211,7 +211,7 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
|
||||
upstream_is_stream,
|
||||
)
|
||||
};
|
||||
let Some(provider_request_body) = provider_request_body else {
|
||||
let Some(mut provider_request_body) = provider_request_body else {
|
||||
mark_skipped_local_openai_image_candidate_with_failure_diagnostic(
|
||||
state,
|
||||
input,
|
||||
@@ -229,6 +229,11 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
|
||||
.await;
|
||||
return None;
|
||||
};
|
||||
apply_xai_upstream_payload_edits(
|
||||
&mut provider_request_body,
|
||||
transport.provider.provider_type.as_str(),
|
||||
provider_api_format,
|
||||
);
|
||||
let Some(mut provider_request_headers) = (if is_grok {
|
||||
build_grok_browser_headers(GrokHeaderInput {
|
||||
transport,
|
||||
|
||||
@@ -8,6 +8,7 @@ use crate::ai_serving::planner::{
|
||||
build_ai_execution_decision_response, resolve_transport_request_encoding_policy,
|
||||
AiExecutionDecisionResponseParts,
|
||||
};
|
||||
use crate::ai_serving::transport::xai::video::is_native_video_request;
|
||||
use crate::ai_serving::transport::{
|
||||
resolve_transport_execution_timeouts, resolve_transport_profile,
|
||||
};
|
||||
@@ -33,7 +34,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
|
||||
let Some(resolved) = resolve_local_video_create_candidate_payload_parts(
|
||||
state, parts, body_json, trace_id, input, &attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
@@ -52,9 +53,32 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
|
||||
.await;
|
||||
let transport_profile = resolve_transport_profile(&transport);
|
||||
let mut extra_fields = serde_json::Map::new();
|
||||
if is_native_video_request(&transport.provider.provider_type, parts.uri.path()) {
|
||||
extra_fields.insert(
|
||||
"video_client_protocol".to_string(),
|
||||
serde_json::json!("xai"),
|
||||
);
|
||||
}
|
||||
|
||||
if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) {
|
||||
extra_fields.insert("proxy".to_string(), proxy_value);
|
||||
}
|
||||
if transport.provider.provider_type.eq_ignore_ascii_case("xai") {
|
||||
extra_fields.insert("video_provider_xai".into(), serde_json::json!(true));
|
||||
if let Some(duration) = resolved.provider_request_body.get("duration") {
|
||||
extra_fields.insert("video_duration".into(), duration.clone());
|
||||
}
|
||||
if parts.uri.path() == "/openai/v1/videos" {
|
||||
extra_fields.insert(
|
||||
"video_size".into(),
|
||||
body_json
|
||||
.get("size")
|
||||
.filter(|v| v.as_str().is_some_and(|s| !s.trim().is_empty()))
|
||||
.cloned()
|
||||
.unwrap_or_else(|| serde_json::json!("720x1280")),
|
||||
);
|
||||
}
|
||||
}
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let report_context = build_local_execution_report_context(LocalExecutionReportContextParts {
|
||||
auth_context: &input.auth_context,
|
||||
|
||||
@@ -3,15 +3,23 @@ use std::sync::Arc;
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::ai_serving::planner::candidate_preparation::resolve_candidate_mapped_model;
|
||||
use crate::ai_serving::planner::candidate_preparation::{
|
||||
prepare_header_authenticated_candidate, resolve_candidate_mapped_model, OauthPreparationContext,
|
||||
};
|
||||
use crate::ai_serving::planner::spec_metadata::local_video_create_spec_metadata;
|
||||
use crate::ai_serving::transport::xai::video::{
|
||||
convert_openai_video_request, is_explicit_native_video_path, is_native_video_request,
|
||||
};
|
||||
use crate::ai_serving::transport::{
|
||||
build_video_create_headers, build_video_create_request_body, build_video_create_upstream_url,
|
||||
resolve_video_create_auth, video_create_transport_unsupported_reason,
|
||||
ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput,
|
||||
};
|
||||
use crate::ai_serving::{CandidateFailureDiagnostic, GatewayProviderTransportSnapshot};
|
||||
use crate::AppState;
|
||||
use crate::ai_serving::{
|
||||
apply_xai_upstream_payload_edits, CandidateFailureDiagnostic, GatewayProviderTransportSnapshot,
|
||||
PlannerAppState,
|
||||
};
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
use super::support::{
|
||||
mark_skipped_local_video_candidate, mark_skipped_local_video_candidate_with_failure_diagnostic,
|
||||
@@ -37,11 +45,16 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
||||
input: &LocalVideoCreateDecisionInput,
|
||||
attempt: &LocalVideoCreateCandidateAttempt,
|
||||
spec: LocalVideoCreateSpec,
|
||||
) -> Option<LocalVideoCreateCandidatePayloadParts> {
|
||||
) -> Result<Option<LocalVideoCreateCandidatePayloadParts>, GatewayError> {
|
||||
let spec_metadata = local_video_create_spec_metadata(spec);
|
||||
let candidate = &attempt.eligible.candidate;
|
||||
let transport = &attempt.eligible.transport;
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
if is_explicit_native_video_path(parts.uri.path())
|
||||
&& !transport.provider.provider_type.eq_ignore_ascii_case("xai")
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let provider_family = provider_video_create_family(spec.family);
|
||||
let transport_unsupported_reason = video_create_transport_unsupported_reason(
|
||||
@@ -60,23 +73,39 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
||||
skip_reason,
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let auth = resolve_video_create_auth(transport, provider_family);
|
||||
let Some((auth_header, auth_value)) = auth else {
|
||||
mark_skipped_local_video_candidate(
|
||||
state,
|
||||
input,
|
||||
let prepared_candidate = match prepare_header_authenticated_candidate(
|
||||
PlannerAppState::new(state),
|
||||
transport,
|
||||
candidate,
|
||||
resolve_video_create_auth(transport, provider_family),
|
||||
OauthPreparationContext {
|
||||
trace_id,
|
||||
candidate,
|
||||
attempt.candidate_index,
|
||||
&attempt.candidate_id,
|
||||
"transport_auth_unavailable",
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
api_format: spec_metadata.api_format,
|
||||
operation: "video_create_candidate_request",
|
||||
},
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(prepared) => prepared,
|
||||
Err(skip_reason) => {
|
||||
mark_skipped_local_video_candidate(
|
||||
state,
|
||||
input,
|
||||
trace_id,
|
||||
candidate,
|
||||
attempt.candidate_index,
|
||||
&attempt.candidate_id,
|
||||
skip_reason,
|
||||
)
|
||||
.await;
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
let auth_header = prepared_candidate.auth_header;
|
||||
let auth_value = prepared_candidate.auth_value;
|
||||
|
||||
let mapped_model = match resolve_candidate_mapped_model(candidate) {
|
||||
Ok(mapped_model) => mapped_model,
|
||||
@@ -91,7 +120,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
||||
skip_reason,
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
|
||||
@@ -117,10 +146,10 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let Some(provider_request_body) = build_video_create_request_body(
|
||||
let Some(mut provider_request_body) = build_video_create_request_body(
|
||||
body_json,
|
||||
provider_family,
|
||||
&mapped_model,
|
||||
@@ -142,11 +171,28 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
if transport.provider.provider_type.eq_ignore_ascii_case("xai")
|
||||
&& !is_native_video_request(&transport.provider.provider_type, parts.uri.path())
|
||||
{
|
||||
provider_request_body =
|
||||
convert_openai_video_request(&provider_request_body).map_err(|message| {
|
||||
GatewayError::Client {
|
||||
status: http::StatusCode::BAD_REQUEST,
|
||||
message: message.to_string(),
|
||||
}
|
||||
})?;
|
||||
}
|
||||
apply_xai_upstream_payload_edits(
|
||||
&mut provider_request_body,
|
||||
transport.provider.provider_type.as_str(),
|
||||
spec_metadata.api_format,
|
||||
);
|
||||
|
||||
let Some(provider_request_headers) =
|
||||
build_video_create_headers(ProviderVideoCreateHeadersInput {
|
||||
transport,
|
||||
headers: effective_headers,
|
||||
auth_header: &auth_header,
|
||||
auth_value: &auth_value,
|
||||
@@ -170,10 +216,10 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
Some(LocalVideoCreateCandidatePayloadParts {
|
||||
Ok(Some(LocalVideoCreateCandidatePayloadParts {
|
||||
transport: Arc::clone(transport),
|
||||
auth_header,
|
||||
auth_value,
|
||||
@@ -181,7 +227,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
||||
provider_request_headers,
|
||||
provider_request_body,
|
||||
upstream_url,
|
||||
})
|
||||
}))
|
||||
}
|
||||
|
||||
fn provider_video_create_family(family: LocalVideoCreateFamily) -> ProviderVideoCreateFamily {
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -488,6 +488,7 @@ impl ResponsesWebSocketBodyNormalization {
|
||||
digest.update([match self.reasoning_replay_policy {
|
||||
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds => 0,
|
||||
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque => 1,
|
||||
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted => 2,
|
||||
}]);
|
||||
update_normalization_optional_json_digest(&mut digest, self.model_directive_patch.as_ref());
|
||||
digest.finalize().into()
|
||||
|
||||
@@ -12,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,
|
||||
|
||||
@@ -58,6 +58,10 @@ pub(crate) mod windsurf {
|
||||
pub(crate) use aether_provider_transport::windsurf::*;
|
||||
}
|
||||
|
||||
pub(crate) mod xai {
|
||||
pub(crate) use aether_provider_transport::xai::*;
|
||||
}
|
||||
|
||||
pub(crate) use aether_provider_transport::{
|
||||
append_transport_diagnostics_to_value, apply_codex_fingerprint_convergence,
|
||||
apply_codex_fingerprint_convergence_with_context, apply_local_auth_config_header_overrides,
|
||||
|
||||
@@ -53,6 +53,8 @@ const AI_ANY_ROUTE_PATTERNS: &[&str] = &[
|
||||
"/v1beta/operations/{*operation_path}",
|
||||
"/v1/videos",
|
||||
"/v1/videos/{*video_path}",
|
||||
"/openai/v1/videos",
|
||||
"/openai/v1/videos/{*video_path}",
|
||||
"/upload/v1beta/files",
|
||||
"/v1beta/files",
|
||||
"/v1beta/files/{*file_path}",
|
||||
|
||||
@@ -536,6 +536,9 @@ mod tests {
|
||||
|
||||
fn sample_sparse_stored_task() -> StoredVideoTask {
|
||||
let snapshot = LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
||||
local_short_id: None,
|
||||
native_response: None,
|
||||
xai_provider: false,
|
||||
local_task_id: "task-1".to_string(),
|
||||
upstream_task_id: "ext-1".to_string(),
|
||||
created_at_unix_ms: 1,
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -123,6 +123,15 @@ impl GatewayDataState {
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn attach_video_task_repository_for_tests<T>(mut self, repository: Arc<T>) -> Self
|
||||
where
|
||||
T: VideoTaskRepository + 'static,
|
||||
{
|
||||
self.video_task_reader = Some(repository.clone());
|
||||
self.video_task_writer = Some(repository);
|
||||
self
|
||||
}
|
||||
|
||||
pub(crate) fn with_video_task_repository_for_tests<T>(repository: Arc<T>) -> Self
|
||||
where
|
||||
T: VideoTaskRepository + 'static,
|
||||
|
||||
@@ -1464,6 +1464,22 @@ pub(crate) async fn maybe_execute_sync_via_local_video_decision(
|
||||
.await
|
||||
}
|
||||
|
||||
fn supports_local_video_get(
|
||||
parts: &http::request::Parts,
|
||||
decision: &GatewayControlDecision,
|
||||
) -> bool {
|
||||
parts.method == http::Method::GET
|
||||
&& decision.route_kind.as_deref() == Some("video")
|
||||
&& (crate::video_tasks::resolve_video_task_read_lookup_key(
|
||||
decision.route_family.as_deref(),
|
||||
parts.uri.path(),
|
||||
)
|
||||
.is_some()
|
||||
|| (decision.route_family.as_deref() == Some("openai")
|
||||
&& crate::video_tasks::extract_openai_task_id_from_content_path(parts.uri.path())
|
||||
.is_some()))
|
||||
}
|
||||
|
||||
pub(crate) fn maybe_execute_sync_request<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
@@ -1477,7 +1493,7 @@ pub(crate) fn maybe_execute_sync_request<'a>(
|
||||
};
|
||||
#[cfg(not(test))]
|
||||
{
|
||||
if parts.method != http::Method::POST {
|
||||
if parts.method != http::Method::POST && !supports_local_video_get(parts, decision) {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
return maybe_execute_sync_local_path(state, parts, body_bytes, trace_id, decision)
|
||||
@@ -1490,6 +1506,7 @@ pub(crate) fn maybe_execute_sync_request<'a>(
|
||||
.unwrap_or_default()
|
||||
.is_empty()
|
||||
&& parts.method != http::Method::POST
|
||||
&& !supports_local_video_get(parts, decision)
|
||||
{
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
@@ -1511,7 +1528,7 @@ pub(crate) fn maybe_execute_stream_request<'a>(
|
||||
};
|
||||
#[cfg(not(test))]
|
||||
{
|
||||
if parts.method != http::Method::POST {
|
||||
if parts.method != http::Method::POST && !supports_local_video_get(parts, decision) {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
return maybe_execute_stream_local_path(state, parts, body_bytes, trace_id, decision)
|
||||
@@ -1524,6 +1541,7 @@ pub(crate) fn maybe_execute_stream_request<'a>(
|
||||
.unwrap_or_default()
|
||||
.is_empty()
|
||||
&& parts.method != http::Method::POST
|
||||
&& !supports_local_video_get(parts, decision)
|
||||
{
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
|
||||
@@ -32,6 +32,10 @@ fn request_has_execution_runtime_via_guard(headers: &HeaderMap) -> bool {
|
||||
}
|
||||
|
||||
pub(crate) fn frontdoor_self_loop_public_ai_path(path: &str) -> bool {
|
||||
let path = path
|
||||
.strip_prefix("/openai")
|
||||
.filter(|p| *p == "/v1/videos" || p.starts_with("/v1/videos/"))
|
||||
.unwrap_or(path);
|
||||
matches!(
|
||||
path,
|
||||
"/v1/messages"
|
||||
|
||||
@@ -74,7 +74,8 @@ fn validate_batch_access_token_import(
|
||||
) -> Result<(), String> {
|
||||
if !provider_type_supports_access_token_import(provider_type) {
|
||||
return Err(
|
||||
"Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok Provider".to_string(),
|
||||
"Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok / xAI Provider"
|
||||
.to_string(),
|
||||
);
|
||||
}
|
||||
if provider_type.eq_ignore_ascii_case("claude_code") {
|
||||
|
||||
@@ -214,7 +214,12 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
} else {
|
||||
let sso_from_cookie = grok_cookie_session_token(provider_type, raw_token);
|
||||
let token_input = sso_from_cookie.as_deref().unwrap_or(raw_token);
|
||||
let (refresh_token, access_token) = import_tokens_from_raw_token(token_input);
|
||||
let (refresh_token, access_token) =
|
||||
if provider_type.trim().eq_ignore_ascii_case("xai") {
|
||||
(None, Some(token_input.to_string()))
|
||||
} else {
|
||||
import_tokens_from_raw_token(token_input)
|
||||
};
|
||||
let (refresh_token, access_token) = normalize_provider_import_tokens(
|
||||
provider_type,
|
||||
refresh_token.as_deref(),
|
||||
@@ -262,6 +267,7 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
let object = normalized_claude_object.as_ref().unwrap_or(object);
|
||||
let is_grok = provider_type.trim().eq_ignore_ascii_case("grok");
|
||||
let is_windsurf = provider_type.trim().eq_ignore_ascii_case("windsurf");
|
||||
let is_xai = provider_type.trim().eq_ignore_ascii_case("xai");
|
||||
let is_codex_agent_identity = provider_type.trim().eq_ignore_ascii_case("codex")
|
||||
&& aether_provider_transport::is_codex_agent_identity_auth_config_value(item);
|
||||
if is_codex_agent_identity {
|
||||
@@ -336,14 +342,6 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let (refresh_token, access_token) = normalize_provider_import_tokens(
|
||||
provider_type,
|
||||
refresh_token.as_deref(),
|
||||
access_token
|
||||
.as_deref()
|
||||
.or(session_token.as_deref())
|
||||
.or(header_bearer_token.as_deref()),
|
||||
);
|
||||
let windsurf_api_key = is_windsurf
|
||||
.then(|| {
|
||||
coerce_admin_provider_oauth_import_str(
|
||||
@@ -351,6 +349,22 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
)
|
||||
})
|
||||
.flatten();
|
||||
let xai_api_key = is_xai
|
||||
.then(|| {
|
||||
coerce_admin_provider_oauth_import_str(
|
||||
object.get("api_key").or_else(|| object.get("apiKey")),
|
||||
)
|
||||
})
|
||||
.flatten();
|
||||
let (refresh_token, access_token) = normalize_provider_import_tokens(
|
||||
provider_type,
|
||||
refresh_token.as_deref(),
|
||||
access_token
|
||||
.as_deref()
|
||||
.or(session_token.as_deref())
|
||||
.or(header_bearer_token.as_deref())
|
||||
.or(xai_api_key.as_deref()),
|
||||
);
|
||||
let windsurf_token = is_windsurf
|
||||
.then(|| {
|
||||
coerce_admin_provider_oauth_import_str(
|
||||
@@ -1577,4 +1591,23 @@ mod tests {
|
||||
assert!(entries[1].access_token.is_none());
|
||||
assert!(entries[1].raw_credentials.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_xai_api_key_json_and_raw_lines_as_access_token() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"xai",
|
||||
r#"{"api_key":"xai-api-key","email":"a@x.ai"}
|
||||
{"refresh_token":"xai-refresh"}
|
||||
xai-raw-api-key"#,
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 3);
|
||||
assert!(entries[0].refresh_token.is_none());
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some("xai-api-key"));
|
||||
assert_eq!(entries[0].email.as_deref(), Some("[email protected]"));
|
||||
assert_eq!(entries[1].refresh_token.as_deref(), Some("xai-refresh"));
|
||||
assert!(entries[1].access_token.is_none());
|
||||
assert!(entries[2].refresh_token.is_none());
|
||||
assert_eq!(entries[2].access_token.as_deref(), Some("xai-raw-api-key"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -186,10 +186,10 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
|
||||
));
|
||||
};
|
||||
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
|
||||
if provider_type != "kiro" && provider_type != "windsurf" {
|
||||
if provider_type != "kiro" && provider_type != "windsurf" && provider_type != "xai" {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"设备授权仅支持 Kiro / Windsurf provider",
|
||||
"设备授权仅支持 Kiro / Windsurf / xAI provider",
|
||||
));
|
||||
}
|
||||
let Some(principal) = request_context
|
||||
@@ -219,6 +219,19 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
|
||||
)
|
||||
.await;
|
||||
|
||||
if provider_type == "xai" {
|
||||
return super::xai::handle_admin_provider_oauth_xai_device_authorize(
|
||||
state,
|
||||
&provider_id,
|
||||
&provider,
|
||||
principal,
|
||||
runtime_endpoint.as_ref(),
|
||||
request_proxy,
|
||||
payload.proxy_node_id.as_deref(),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
if provider_type == "windsurf" {
|
||||
let session_id = generate_provider_oauth_nonce();
|
||||
let login_option = payload
|
||||
|
||||
@@ -2,6 +2,7 @@ mod authorize;
|
||||
mod lease;
|
||||
mod poll;
|
||||
mod session;
|
||||
mod xai;
|
||||
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
|
||||
@@ -479,6 +479,18 @@ pub(super) async fn handle_admin_provider_oauth_device_poll(
|
||||
)
|
||||
.await;
|
||||
|
||||
if provider_type == "xai" {
|
||||
return super::xai::handle_admin_provider_oauth_xai_device_poll(
|
||||
state,
|
||||
&provider,
|
||||
&endpoints,
|
||||
request_proxy,
|
||||
session_id,
|
||||
session,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
if provider_type == "windsurf" {
|
||||
return handle_admin_provider_oauth_windsurf_browser_device_poll(
|
||||
state,
|
||||
|
||||
@@ -0,0 +1,368 @@
|
||||
use super::session::attach_admin_provider_oauth_device_poll_terminal_response;
|
||||
use crate::control::GatewayAdminPrincipalContext;
|
||||
use crate::handlers::admin::provider::oauth::dispatch::helpers::admin_provider_oauth_key_name_from_auth_config;
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::oauth::provisioning::{
|
||||
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::runtime::spawn_provider_oauth_account_state_refresh_after_update;
|
||||
use crate::handlers::admin::provider::oauth::state::{
|
||||
current_unix_secs, generate_provider_oauth_nonce,
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::GatewayError;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use aether_data::repository::provider_oauth::{
|
||||
StoredAdminProviderOAuthDeviceSession, KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
|
||||
};
|
||||
use aether_oauth::core::OAuthError;
|
||||
use aether_oauth::provider::providers::{
|
||||
XaiDevicePollOutcome, XaiProviderOAuthAdapter, XAI_CLIENT_ID, XAI_DEVICE_CODE_URL,
|
||||
XAI_TOKEN_URL,
|
||||
};
|
||||
use aether_oauth::provider::ProviderOAuthTransportContext;
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde_json::{json, Value};
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_xai_device_authorize(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
principal: &GatewayAdminPrincipalContext,
|
||||
runtime_endpoint: Option<&StoredProviderCatalogEndpoint>,
|
||||
request_proxy: Option<ProxySnapshot>,
|
||||
proxy_node_id: Option<&str>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let device_url = state.provider_oauth_token_url("xai_device", XAI_DEVICE_CODE_URL);
|
||||
let token_url = state.provider_oauth_token_url("xai", XAI_TOKEN_URL);
|
||||
let adapter =
|
||||
XaiProviderOAuthAdapter::default().with_endpoint_overrides(&device_url, &token_url);
|
||||
let ctx = ProviderOAuthTransportContext {
|
||||
provider_id: provider_id.to_string(),
|
||||
provider_type: provider.provider_type.clone(),
|
||||
endpoint_id: runtime_endpoint.map(|endpoint| endpoint.id.clone()),
|
||||
key_id: None,
|
||||
auth_type: Some("oauth".to_string()),
|
||||
decrypted_api_key: None,
|
||||
decrypted_auth_config: None,
|
||||
provider_config: provider.config.clone(),
|
||||
endpoint_config: runtime_endpoint.and_then(|endpoint| endpoint.config.clone()),
|
||||
key_config: None,
|
||||
network: aether_oauth::network::OAuthNetworkContext::provider_operation(
|
||||
request_proxy.clone(),
|
||||
),
|
||||
};
|
||||
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
|
||||
let authorization = match adapter.start_device_flow(&executor, &ctx).await {
|
||||
Ok(authorization) => authorization,
|
||||
Err(error) => {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
sanitize_xai_oauth_error(&error),
|
||||
));
|
||||
}
|
||||
};
|
||||
|
||||
let now_unix_secs = current_unix_secs();
|
||||
let session_id = generate_provider_oauth_nonce();
|
||||
let session = StoredAdminProviderOAuthDeviceSession {
|
||||
session_id: session_id.clone(),
|
||||
provider_id: provider_id.to_string(),
|
||||
initiated_by_user_id: principal.user_id.clone(),
|
||||
initiated_by_session_id: principal.session_id.clone(),
|
||||
initiated_by_management_token_id: principal.management_token_id.clone(),
|
||||
region: String::new(),
|
||||
client_id: XAI_CLIENT_ID.to_string(),
|
||||
client_secret: String::new(),
|
||||
device_code: authorization.device_code.clone(),
|
||||
auth_type: Some("device".to_string()),
|
||||
social_provider: None,
|
||||
code_verifier: None,
|
||||
redirect_uri: Some(token_url),
|
||||
machine_id: None,
|
||||
interval: authorization.interval,
|
||||
expires_at_unix_secs: now_unix_secs.saturating_add(authorization.expires_in),
|
||||
status: "pending".to_string(),
|
||||
proxy_node_id: proxy_node_id
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned),
|
||||
created_at_unix_ms: now_unix_secs,
|
||||
key_id: None,
|
||||
email: None,
|
||||
replaced: false,
|
||||
error_msg: None,
|
||||
};
|
||||
if let Err(response) = state
|
||||
.save_provider_oauth_device_session(
|
||||
&session_id,
|
||||
&session,
|
||||
authorization
|
||||
.expires_in
|
||||
.saturating_add(KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS),
|
||||
)
|
||||
.await
|
||||
{
|
||||
return Ok(response);
|
||||
}
|
||||
|
||||
Ok(Json(json!({
|
||||
"session_id": session_id,
|
||||
"user_code": authorization.user_code,
|
||||
"verification_uri": authorization.verification_uri,
|
||||
"verification_uri_complete": authorization.verification_uri_complete,
|
||||
"expires_in": authorization.expires_in,
|
||||
"interval": authorization.interval,
|
||||
"auth_type": "device",
|
||||
}))
|
||||
.into_response())
|
||||
}
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_xai_device_poll(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoints: &[StoredProviderCatalogEndpoint],
|
||||
request_proxy: Option<ProxySnapshot>,
|
||||
session_id: &str,
|
||||
mut session: StoredAdminProviderOAuthDeviceSession,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let token_url = session
|
||||
.redirect_uri
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.unwrap_or_else(|| state.provider_oauth_token_url("xai", XAI_TOKEN_URL));
|
||||
let adapter =
|
||||
XaiProviderOAuthAdapter::default().with_endpoint_overrides(XAI_DEVICE_CODE_URL, token_url);
|
||||
let ctx = ProviderOAuthTransportContext {
|
||||
provider_id: provider.id.clone(),
|
||||
provider_type: provider.provider_type.clone(),
|
||||
endpoint_id: None,
|
||||
key_id: None,
|
||||
auth_type: Some("oauth".to_string()),
|
||||
decrypted_api_key: None,
|
||||
decrypted_auth_config: None,
|
||||
provider_config: provider.config.clone(),
|
||||
endpoint_config: None,
|
||||
key_config: None,
|
||||
network: aether_oauth::network::OAuthNetworkContext::provider_operation(
|
||||
request_proxy.clone(),
|
||||
),
|
||||
};
|
||||
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
|
||||
let outcome = match adapter
|
||||
.poll_device_token(&executor, &ctx, &session.device_code)
|
||||
.await
|
||||
{
|
||||
Ok(outcome) => outcome,
|
||||
Err(error) => {
|
||||
return Ok(xai_device_poll_terminal_from_error(
|
||||
state,
|
||||
session_id,
|
||||
&mut session,
|
||||
&error,
|
||||
)
|
||||
.await);
|
||||
}
|
||||
};
|
||||
|
||||
match outcome {
|
||||
XaiDevicePollOutcome::Pending => {
|
||||
Ok(Json(json!({"status": "pending", "replaced": false})).into_response())
|
||||
}
|
||||
XaiDevicePollOutcome::SlowDown => {
|
||||
Ok(Json(json!({"status": "slow_down", "replaced": false})).into_response())
|
||||
}
|
||||
XaiDevicePollOutcome::Authorized(result) => {
|
||||
persist_xai_device_authorization(
|
||||
state,
|
||||
provider,
|
||||
endpoints,
|
||||
request_proxy,
|
||||
session_id,
|
||||
session,
|
||||
*result,
|
||||
)
|
||||
.await
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn persist_xai_device_authorization(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoints: &[StoredProviderCatalogEndpoint],
|
||||
request_proxy: Option<ProxySnapshot>,
|
||||
session_id: &str,
|
||||
mut session: StoredAdminProviderOAuthDeviceSession,
|
||||
result: aether_oauth::provider::ProviderOAuthTokenSet,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let access_token = result.token_set.access_token.trim().to_string();
|
||||
if access_token.is_empty() {
|
||||
return Ok(Json(json!({
|
||||
"status": "error",
|
||||
"error": "xAI token 响应缺少 access_token",
|
||||
"replaced": false,
|
||||
}))
|
||||
.into_response());
|
||||
}
|
||||
let mut auth_config = result.auth_config.as_object().cloned().unwrap_or_default();
|
||||
auth_config.insert("provider_type".to_string(), json!("xai"));
|
||||
auth_config.insert("auth_method".to_string(), json!("oauth"));
|
||||
auth_config.insert("using_api".to_string(), json!(false));
|
||||
|
||||
let duplicate = match state
|
||||
.find_duplicate_provider_oauth_key(&provider.id, &auth_config, None)
|
||||
.await
|
||||
{
|
||||
Ok(duplicate) => duplicate,
|
||||
Err(detail) => {
|
||||
return Ok(Json(json!({
|
||||
"status": "error",
|
||||
"error": detail,
|
||||
"replaced": false,
|
||||
}))
|
||||
.into_response());
|
||||
}
|
||||
};
|
||||
|
||||
let api_formats = provider_oauth_active_api_formats(endpoints);
|
||||
let key_proxy = provider_oauth_key_proxy_value(session.proxy_node_id.as_deref());
|
||||
let expires_at = result.token_set.expires_at_unix_secs;
|
||||
let email = auth_config
|
||||
.get("email")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
let mut replaced = false;
|
||||
let persisted_key = if let Some(existing_key) = duplicate {
|
||||
replaced = true;
|
||||
match state
|
||||
.update_existing_provider_oauth_catalog_key(
|
||||
&existing_key,
|
||||
&provider.provider_type,
|
||||
&access_token,
|
||||
&auth_config,
|
||||
&api_formats,
|
||||
key_proxy.clone(),
|
||||
expires_at,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(key) => key,
|
||||
None => {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"provider oauth write unavailable",
|
||||
));
|
||||
}
|
||||
}
|
||||
} else {
|
||||
let key_name = admin_provider_oauth_key_name_from_auth_config(
|
||||
&provider.provider_type,
|
||||
&auth_config,
|
||||
None,
|
||||
);
|
||||
match state
|
||||
.create_provider_oauth_catalog_key(
|
||||
&provider.id,
|
||||
&provider.provider_type,
|
||||
&key_name,
|
||||
&access_token,
|
||||
&auth_config,
|
||||
&api_formats,
|
||||
key_proxy,
|
||||
expires_at,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(key) => key,
|
||||
None => {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"provider oauth write unavailable",
|
||||
));
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
spawn_provider_oauth_account_state_refresh_after_update(
|
||||
state.cloned_app(),
|
||||
provider.clone(),
|
||||
persisted_key.id.clone(),
|
||||
request_proxy.clone(),
|
||||
);
|
||||
|
||||
session.status = "authorized".to_string();
|
||||
session.key_id = Some(persisted_key.id.clone());
|
||||
session.email = email.clone();
|
||||
session.replaced = replaced;
|
||||
session.error_msg = None;
|
||||
let _ = state
|
||||
.save_provider_oauth_device_session(session_id, &session, 60)
|
||||
.await;
|
||||
|
||||
Ok(attach_admin_provider_oauth_device_poll_terminal_response(
|
||||
session_id,
|
||||
"authorized",
|
||||
Json(json!({
|
||||
"status": "authorized",
|
||||
"key_id": persisted_key.id,
|
||||
"email": email,
|
||||
"replaced": replaced,
|
||||
}))
|
||||
.into_response(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn xai_device_poll_terminal_from_error(
|
||||
state: &AdminAppState<'_>,
|
||||
session_id: &str,
|
||||
session: &mut StoredAdminProviderOAuthDeviceSession,
|
||||
error: &OAuthError,
|
||||
) -> Response<Body> {
|
||||
let (status, message) = match error {
|
||||
OAuthError::InvalidRequest(detail) if detail.to_ascii_lowercase().contains("expired") => {
|
||||
("expired", "设备码已过期".to_string())
|
||||
}
|
||||
OAuthError::InvalidRequest(detail) if detail.to_ascii_lowercase().contains("denied") => {
|
||||
("error", "用户拒绝授权".to_string())
|
||||
}
|
||||
_ => ("error", sanitize_xai_oauth_error(error)),
|
||||
};
|
||||
session.status = status.to_string();
|
||||
session.error_msg = Some(message.clone());
|
||||
let _ = state
|
||||
.save_provider_oauth_device_session(session_id, session, 30)
|
||||
.await;
|
||||
attach_admin_provider_oauth_device_poll_terminal_response(
|
||||
session_id,
|
||||
status,
|
||||
Json(json!({
|
||||
"status": status,
|
||||
"error": message,
|
||||
"replaced": false,
|
||||
}))
|
||||
.into_response(),
|
||||
)
|
||||
}
|
||||
|
||||
fn sanitize_xai_oauth_error(error: &OAuthError) -> String {
|
||||
match error {
|
||||
OAuthError::InvalidRequest(_) => "xAI 设备授权失败: 请求参数无效".to_string(),
|
||||
OAuthError::HttpStatus { status_code, .. } => {
|
||||
format!("xAI 设备授权失败: HTTP {status_code}")
|
||||
}
|
||||
_ => "xAI 设备授权失败".to_string(),
|
||||
}
|
||||
}
|
||||
@@ -715,7 +715,7 @@ async fn resolve_admin_provider_oauth_single_import_tokens(
|
||||
if !provider_type_supports_access_token_import(provider_type) {
|
||||
return Err(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok Provider",
|
||||
"Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok / xAI Provider",
|
||||
));
|
||||
}
|
||||
|
||||
@@ -867,7 +867,7 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
flatten_claude_code_credentials_payload(&mut raw_payload);
|
||||
}
|
||||
let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken");
|
||||
let access_token_input = import_payload_string_any(
|
||||
let mut access_token_input = import_payload_string_any(
|
||||
&raw_payload,
|
||||
&[
|
||||
"access_token",
|
||||
@@ -879,6 +879,9 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
],
|
||||
)
|
||||
.or_else(|| provider_oauth_import_authorization_bearer_token_from_object(&raw_payload));
|
||||
if provider_type == "xai" && access_token_input.is_none() {
|
||||
access_token_input = import_payload_string(&raw_payload, "api_key", "apiKey");
|
||||
}
|
||||
let imported_expires_at =
|
||||
import_payload_u64_any(&raw_payload, &["expires_at", "expiresAt", "expired"]);
|
||||
let (refresh_token_input, access_token_input) = normalize_provider_import_tokens(
|
||||
@@ -901,7 +904,11 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
if !create_agent_identity && refresh_token_input.is_none() && access_token_input.is_none() {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Refresh Token、Access Token 或 sso_token 不能为空",
|
||||
if provider_type == "xai" {
|
||||
"Refresh Token、Access Token 或 api_key 不能为空"
|
||||
} else {
|
||||
"Refresh Token、Access Token 或 sso_token 不能为空"
|
||||
},
|
||||
));
|
||||
}
|
||||
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
|
||||
|
||||
@@ -70,6 +70,12 @@ pub(super) async fn handle_admin_provider_oauth_start_key(
|
||||
"Windsurf 请使用浏览器登录或导入凭据。",
|
||||
));
|
||||
}
|
||||
if provider_type == "xai" {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"xAI 请使用设备授权或导入凭据。",
|
||||
));
|
||||
}
|
||||
let Some(template) = admin_provider_oauth_template(&provider_type) else {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
@@ -167,6 +173,12 @@ pub(super) async fn handle_admin_provider_oauth_start_provider(
|
||||
"Windsurf 请使用浏览器登录或导入凭据。",
|
||||
));
|
||||
}
|
||||
if provider_type == "xai" {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"xAI 请使用设备授权或导入凭据。",
|
||||
));
|
||||
}
|
||||
let Some(template) = admin_provider_oauth_template(&provider_type) else {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
|
||||
@@ -121,6 +121,9 @@ pub(super) fn normalize_provider_import_tokens(
|
||||
if provider_type == "grok" {
|
||||
return (None, access_token.or(refresh_token));
|
||||
}
|
||||
if provider_type == "xai" {
|
||||
return (refresh_token, access_token);
|
||||
}
|
||||
if provider_type == "claude_code" {
|
||||
if access_token.is_none() && refresh_token.as_deref().is_some_and(is_claude_access_token) {
|
||||
return (None, refresh_token);
|
||||
@@ -237,7 +240,7 @@ pub(super) fn provider_oauth_import_authorization_bearer_token_from_object(
|
||||
pub(super) fn provider_type_supports_access_token_import(provider_type: &str) -> bool {
|
||||
matches!(
|
||||
provider_type.trim().to_ascii_lowercase().as_str(),
|
||||
"claude_code" | "codex" | "chatgpt_web" | "grok"
|
||||
"claude_code" | "codex" | "chatgpt_web" | "grok" | "xai"
|
||||
)
|
||||
}
|
||||
|
||||
@@ -331,6 +334,15 @@ pub(super) fn build_provider_access_token_import_auth_config(
|
||||
auth_config.insert("sso_token".to_string(), json!(access_token));
|
||||
auth_config.insert("auth_method".to_string(), json!("sso_token"));
|
||||
}
|
||||
if provider_type.trim().eq_ignore_ascii_case("xai") {
|
||||
if refresh_token.is_some() {
|
||||
auth_config.insert("auth_method".to_string(), json!("oauth"));
|
||||
auth_config.insert("using_api".to_string(), json!(false));
|
||||
} else {
|
||||
auth_config.insert("auth_method".to_string(), json!("api_key"));
|
||||
auth_config.insert("using_api".to_string(), json!(true));
|
||||
}
|
||||
}
|
||||
|
||||
auth_config.insert(
|
||||
"access_token_import_temporary".to_string(),
|
||||
@@ -532,6 +544,41 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalize_xai_import_keeps_refresh_token_separate_from_api_key() {
|
||||
let (refresh_token, access_token) =
|
||||
normalize_provider_import_tokens("xai", Some("xai-refresh-token"), None);
|
||||
assert_eq!(refresh_token.as_deref(), Some("xai-refresh-token"));
|
||||
assert!(access_token.is_none());
|
||||
|
||||
let (refresh_token, access_token) =
|
||||
normalize_provider_import_tokens("xai", None, Some("xai-api-key"));
|
||||
assert!(refresh_token.is_none());
|
||||
assert_eq!(access_token.as_deref(), Some("xai-api-key"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_xai_auth_config_from_api_key_and_oauth_tokens() {
|
||||
let (api_key_config, _) =
|
||||
build_provider_access_token_import_auth_config("xai", "xai-api-key", None, None, None);
|
||||
assert_eq!(api_key_config.get("auth_method"), Some(&json!("api_key")));
|
||||
assert_eq!(api_key_config.get("using_api"), Some(&json!(true)));
|
||||
|
||||
let (oauth_config, _) = build_provider_access_token_import_auth_config(
|
||||
"xai",
|
||||
"xai-access-token",
|
||||
Some("xai-refresh-token"),
|
||||
None,
|
||||
None,
|
||||
);
|
||||
assert_eq!(oauth_config.get("auth_method"), Some(&json!("oauth")));
|
||||
assert_eq!(oauth_config.get("using_api"), Some(&json!(false)));
|
||||
assert_eq!(
|
||||
oauth_config.get("refresh_token"),
|
||||
Some(&json!("xai-refresh-token"))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn flattens_only_claude_ai_oauth_credentials_and_converts_expiry_ms() {
|
||||
let mut payload = json!({
|
||||
|
||||
@@ -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<StoredProviderCatalogKey>,
|
||||
proxy_override: Option<ProxySnapshot>,
|
||||
) -> ProviderQuotaRefreshFuture<'a> {
|
||||
Box::pin(refresh_xai_provider_quota_locally(
|
||||
state,
|
||||
provider,
|
||||
endpoint,
|
||||
keys,
|
||||
proxy_override,
|
||||
))
|
||||
}
|
||||
|
||||
@@ -7,3 +7,4 @@ pub(crate) mod grok;
|
||||
pub(crate) mod kiro;
|
||||
pub(crate) mod shared;
|
||||
pub(crate) mod windsurf;
|
||||
pub(crate) mod xai;
|
||||
|
||||
@@ -1715,6 +1715,7 @@ fn provider_quota_url_has_allowed_origin(provider_name: &str, value: &str) -> bo
|
||||
"gemini_cli" => host == "cloudcode-pa.googleapis.com",
|
||||
"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",
|
||||
|
||||
@@ -0,0 +1,313 @@
|
||||
use super::shared::{
|
||||
build_provider_quota_execution_plan, build_quota_snapshot_payload,
|
||||
default_provider_quota_execution_timeouts, execute_provider_quota_plan,
|
||||
extract_execution_error_message, oauth_refresh_auto_removed_result,
|
||||
persist_provider_quota_refresh_state, quota_key_auto_removed,
|
||||
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::quota::parse_xai_billing_response;
|
||||
use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use aether_provider_pool::{build_xai_pool_billing_request, build_xai_pool_user_request};
|
||||
use aether_provider_transport::xai::{
|
||||
extract_xai_user_id_from_auth_config, extract_xai_user_id_from_value, xai_auth_uses_api,
|
||||
};
|
||||
use serde_json::{json, Value};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
async fn execute_xai_quota_plan(
|
||||
state: &AdminAppState<'_>,
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
spec: aether_provider_pool::ProviderPoolQuotaRequestSpec,
|
||||
proxy_override: Option<&ProxySnapshot>,
|
||||
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
|
||||
let proxy = match proxy_override {
|
||||
Some(proxy) => Some(proxy.clone()),
|
||||
None => {
|
||||
state
|
||||
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let plan = build_provider_quota_execution_plan(
|
||||
transport,
|
||||
spec,
|
||||
proxy,
|
||||
state.resolve_transport_profile(transport),
|
||||
timeouts,
|
||||
);
|
||||
|
||||
execute_provider_quota_plan(state, transport, plan, "xai").await
|
||||
}
|
||||
|
||||
fn xai_authorization_from_header(authorization: &(String, String)) -> (String, String) {
|
||||
authorization.clone()
|
||||
}
|
||||
|
||||
fn enrich_xai_subscription_title(mut metadata: Value, auth_config: Option<&str>) -> Value {
|
||||
if metadata
|
||||
.get("subscription_title")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty())
|
||||
{
|
||||
return metadata;
|
||||
}
|
||||
let Some(config) = auth_config
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(|value| serde_json::from_str::<Value>(value).ok())
|
||||
else {
|
||||
return metadata;
|
||||
};
|
||||
let title = ["subscription_tier", "subscriptionTier", "tier", "plan"]
|
||||
.iter()
|
||||
.find_map(|field| {
|
||||
config
|
||||
.get(*field)
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
});
|
||||
if let Some(title) = title {
|
||||
if let Some(object) = metadata.as_object_mut() {
|
||||
object.insert("subscription_title".to_string(), json!(title));
|
||||
}
|
||||
}
|
||||
metadata
|
||||
}
|
||||
|
||||
pub(crate) async fn refresh_xai_provider_quota_locally(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
keys: Vec<StoredProviderCatalogKey>,
|
||||
proxy_override: Option<ProxySnapshot>,
|
||||
) -> Result<Option<serde_json::Value>, GatewayError> {
|
||||
let mut results = Vec::new();
|
||||
let mut success_count = 0usize;
|
||||
let mut failed_count = 0usize;
|
||||
let mut auto_removed_count = 0usize;
|
||||
|
||||
for key in keys {
|
||||
let transport = match state
|
||||
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
|
||||
.await?
|
||||
{
|
||||
Some(transport) => transport,
|
||||
None => {
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"status": "error",
|
||||
"message": "Provider transport snapshot unavailable",
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
if xai_auth_uses_api(
|
||||
transport.key.auth_type.as_str(),
|
||||
transport.key.decrypted_auth_config.as_deref(),
|
||||
) {
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"status": "skipped",
|
||||
"message": "xAI API Key 账号没有 Grok Build 订阅额度接口,请使用设备授权账号查询额度。",
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
|
||||
let authorization = match state.resolve_local_oauth_header_auth(&transport).await? {
|
||||
Some(auth) => auth,
|
||||
_ => {
|
||||
if quota_key_auto_removed(state, &key.id).await? {
|
||||
auto_removed_count += 1;
|
||||
results.push(oauth_refresh_auto_removed_result(&key));
|
||||
continue;
|
||||
}
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"status": "error",
|
||||
"message": "缺少 OAuth 认证信息,请先授权/刷新 Token",
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
let fallback_user_id =
|
||||
extract_xai_user_id_from_auth_config(transport.key.decrypted_auth_config.as_deref());
|
||||
let user_id = match execute_xai_quota_plan(
|
||||
state,
|
||||
&transport,
|
||||
build_xai_pool_user_request(
|
||||
&transport.key.id,
|
||||
xai_authorization_from_header(&authorization),
|
||||
),
|
||||
proxy_override.as_ref(),
|
||||
)
|
||||
.await?
|
||||
{
|
||||
ProviderQuotaExecutionOutcome::Response(result) if result.status_code == 200 => result
|
||||
.body
|
||||
.as_ref()
|
||||
.and_then(|body| body.json_body.as_ref())
|
||||
.and_then(extract_xai_user_id_from_value)
|
||||
.or(fallback_user_id),
|
||||
_ => fallback_user_id,
|
||||
};
|
||||
|
||||
let result = match execute_xai_quota_plan(
|
||||
state,
|
||||
&transport,
|
||||
build_xai_pool_billing_request(
|
||||
&transport.key.id,
|
||||
xai_authorization_from_header(&authorization),
|
||||
user_id.as_deref(),
|
||||
),
|
||||
proxy_override.as_ref(),
|
||||
)
|
||||
.await?
|
||||
{
|
||||
ProviderQuotaExecutionOutcome::Response(result) => result,
|
||||
ProviderQuotaExecutionOutcome::Failure(_) => {
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"status": "error",
|
||||
"message": "xAI billing 请求执行失败",
|
||||
"status_code": 502,
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
let mut metadata_update = None::<serde_json::Value>;
|
||||
let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) =
|
||||
quota_refresh_success_invalid_state(&key);
|
||||
let mut status = "error".to_string();
|
||||
let mut message = None::<String>;
|
||||
|
||||
if result.status_code == 200 {
|
||||
if let Some(body_json) = result
|
||||
.body
|
||||
.as_ref()
|
||||
.and_then(|body| body.json_body.as_ref())
|
||||
{
|
||||
metadata_update =
|
||||
parse_xai_billing_response(body_json, now_unix_secs).map(|metadata| {
|
||||
json!({
|
||||
"xai": enrich_xai_subscription_title(
|
||||
metadata,
|
||||
transport.key.decrypted_auth_config.as_deref(),
|
||||
)
|
||||
})
|
||||
});
|
||||
if metadata_update.is_some() {
|
||||
status = "success".to_string();
|
||||
} else {
|
||||
status = "no_metadata".to_string();
|
||||
message = Some("响应中未包含可用的 Grok Build 额度信息".to_string());
|
||||
}
|
||||
} else {
|
||||
status = "no_metadata".to_string();
|
||||
message = Some("响应中未包含配额信息".to_string());
|
||||
}
|
||||
} else {
|
||||
message = Some(
|
||||
extract_execution_error_message(&result)
|
||||
.unwrap_or_else(|| format!("xAI billing 返回状态码 {}", result.status_code)),
|
||||
);
|
||||
if result.status_code == 401 || result.status_code == 403 {
|
||||
let reason = message
|
||||
.clone()
|
||||
.unwrap_or_else(|| "账户访问被禁止".to_string());
|
||||
oauth_invalid_at_unix_secs = Some(now_unix_secs);
|
||||
oauth_invalid_reason = Some(format!("账户访问被禁止: {reason}"));
|
||||
status = if result.status_code == 401 {
|
||||
"unauthorized".to_string()
|
||||
} else {
|
||||
"forbidden".to_string()
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
if !persist_provider_quota_refresh_state(
|
||||
state,
|
||||
&key.id,
|
||||
metadata_update.as_ref(),
|
||||
oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason,
|
||||
None,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"status": "error",
|
||||
"message": "Key 状态写入失败",
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
|
||||
if status == "success" {
|
||||
success_count += 1;
|
||||
} else {
|
||||
failed_count += 1;
|
||||
}
|
||||
|
||||
let mut payload = serde_json::Map::new();
|
||||
payload.insert("key_id".to_string(), json!(key.id));
|
||||
payload.insert("key_name".to_string(), json!(key.name));
|
||||
payload.insert("status".to_string(), json!(status));
|
||||
if let Some(message) = message {
|
||||
payload.insert("message".to_string(), json!(message));
|
||||
}
|
||||
if let Some(metadata) = metadata_update.as_ref().and_then(|value| value.get("xai")) {
|
||||
payload.insert(
|
||||
"metadata".to_string(),
|
||||
admin_provider_metadata_bucket_safe_json("xai", Some(metadata)),
|
||||
);
|
||||
}
|
||||
if let Some(quota_snapshot) = build_quota_snapshot_payload(
|
||||
"xai",
|
||||
key.status_snapshot.as_ref(),
|
||||
metadata_update.as_ref(),
|
||||
) {
|
||||
payload.insert("quota_snapshot".to_string(), quota_snapshot);
|
||||
}
|
||||
results.push(serde_json::Value::Object(payload));
|
||||
}
|
||||
|
||||
Ok(Some(json!({
|
||||
"success": success_count,
|
||||
"failed": failed_count,
|
||||
"total": results.len(),
|
||||
"results": results,
|
||||
"message": format!("已处理 {} 个 Key", results.len()),
|
||||
"auto_removed": auto_removed_count,
|
||||
})))
|
||||
}
|
||||
@@ -932,6 +932,13 @@ fn admin_pool_build_account_quota(
|
||||
return Some(account_quota);
|
||||
}
|
||||
}
|
||||
"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())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -4,9 +4,9 @@ pub(crate) fn normalize_provider_type_input(value: &str) -> Result<String, Strin
|
||||
let normalized = value.trim().to_ascii_lowercase();
|
||||
match normalized.as_str() {
|
||||
"custom" | "claude_code" | "kiro" | "codex" | "chatgpt_web" | "gemini_cli"
|
||||
| "antigravity" | "vertex_ai" | "grok" | "windsurf" => Ok(normalized),
|
||||
| "antigravity" | "vertex_ai" | "grok" | "windsurf" | "xai" => Ok(normalized),
|
||||
_ => Err(
|
||||
"provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai / grok / windsurf"
|
||||
"provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai / grok / windsurf / xai"
|
||||
.to_string(),
|
||||
),
|
||||
}
|
||||
@@ -405,6 +405,14 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalize_provider_type_supports_xai() {
|
||||
assert_eq!(
|
||||
normalize_provider_type_input(" xAI ").expect("type should normalize"),
|
||||
"xai"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalize_api_format_list_dedupes_canonical_formats() {
|
||||
assert_eq!(
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -1388,6 +1388,168 @@ fn build_kiro_quota_status_snapshot(
|
||||
}))
|
||||
}
|
||||
|
||||
fn build_xai_quota_status_snapshot(
|
||||
upstream_metadata: Option<&Value>,
|
||||
source: &str,
|
||||
) -> Option<Value> {
|
||||
let metadata = provider_quota_metadata_bucket(upstream_metadata, "xai")?;
|
||||
let observed_at_unix_secs = provider_quota_timestamp_unix_secs(metadata.get("updated_at"));
|
||||
let usage_limit = metadata
|
||||
.get("usage_limit")
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64);
|
||||
let current_usage = metadata
|
||||
.get("current_usage")
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64);
|
||||
let remaining = metadata
|
||||
.get("remaining")
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64);
|
||||
let usage_ratio = metadata
|
||||
.get("usage_percentage")
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64)
|
||||
.map(|value| (value / 100.0).clamp(0.0, 1.0))
|
||||
.or_else(|| {
|
||||
current_usage
|
||||
.zip(usage_limit)
|
||||
.and_then(|(current_usage, usage_limit)| {
|
||||
(usage_limit > 0.0).then_some((current_usage / usage_limit).clamp(0.0, 1.0))
|
||||
})
|
||||
});
|
||||
let remaining_ratio = usage_ratio.map(|value| (1.0 - value).max(0.0));
|
||||
let next_reset_at = provider_quota_timestamp_unix_secs(metadata.get("next_reset_at"));
|
||||
let reset_seconds = quota_window_reset_seconds(observed_at_unix_secs, next_reset_at);
|
||||
let plan_type = metadata
|
||||
.get("subscription_title")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
let period_type = metadata
|
||||
.get("period_type")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
let usage_label = match period_type.as_deref() {
|
||||
Some("monthly") => "月额度",
|
||||
Some("weekly") => "周额度",
|
||||
_ => "额度",
|
||||
};
|
||||
|
||||
let mut windows = Vec::new();
|
||||
if usage_ratio.is_some()
|
||||
|| remaining.is_some()
|
||||
|| usage_limit.is_some()
|
||||
|| current_usage.is_some()
|
||||
|| next_reset_at.is_some()
|
||||
{
|
||||
windows.push(json!({
|
||||
"code": "usage",
|
||||
"label": usage_label,
|
||||
"scope": "account",
|
||||
"unit": if usage_limit.is_some() { "usd" } else { "percent" },
|
||||
"used_ratio": usage_ratio,
|
||||
"remaining_ratio": remaining_ratio,
|
||||
"used_value": current_usage,
|
||||
"remaining_value": remaining,
|
||||
"limit_value": usage_limit,
|
||||
"reset_at": next_reset_at,
|
||||
"reset_seconds": reset_seconds,
|
||||
}));
|
||||
}
|
||||
|
||||
let prepaid_balance = metadata
|
||||
.get("prepaid_balance")
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64);
|
||||
if prepaid_balance.is_some_and(|value| value > 0.0) {
|
||||
windows.push(json!({
|
||||
"code": "prepaid",
|
||||
"label": "预付额度",
|
||||
"scope": "account",
|
||||
"unit": "usd",
|
||||
"used_ratio": serde_json::Value::Null,
|
||||
"remaining_ratio": serde_json::Value::Null,
|
||||
"remaining_value": prepaid_balance,
|
||||
"reset_at": serde_json::Value::Null,
|
||||
"reset_seconds": serde_json::Value::Null,
|
||||
}));
|
||||
}
|
||||
|
||||
let on_demand_cap = metadata
|
||||
.get("on_demand_cap")
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64);
|
||||
let on_demand_used = metadata
|
||||
.get("on_demand_used")
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64);
|
||||
let on_demand_enabled = metadata
|
||||
.get("on_demand_enabled")
|
||||
.and_then(admin_provider_quota_pure::coerce_json_bool)
|
||||
!= Some(false);
|
||||
if on_demand_enabled && on_demand_cap.is_some_and(|value| value > 0.0) {
|
||||
let on_demand_remaining = on_demand_cap
|
||||
.zip(on_demand_used)
|
||||
.map(|(cap, used)| (cap - used).max(0.0));
|
||||
let on_demand_ratio = on_demand_cap
|
||||
.zip(on_demand_used)
|
||||
.and_then(|(cap, used)| (cap > 0.0).then_some((used / cap).clamp(0.0, 1.0)));
|
||||
windows.push(json!({
|
||||
"code": "on_demand",
|
||||
"label": "按需额度",
|
||||
"scope": "account",
|
||||
"unit": "usd",
|
||||
"used_ratio": on_demand_ratio,
|
||||
"remaining_ratio": on_demand_ratio.map(|value| (1.0 - value).max(0.0)),
|
||||
"used_value": on_demand_used,
|
||||
"remaining_value": on_demand_remaining,
|
||||
"limit_value": on_demand_cap,
|
||||
"reset_at": serde_json::Value::Null,
|
||||
"reset_seconds": serde_json::Value::Null,
|
||||
}));
|
||||
}
|
||||
|
||||
if windows.is_empty() && plan_type.is_none() && observed_at_unix_secs.is_none() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let prepaid_available = prepaid_balance.is_some_and(|value| value > 0.0);
|
||||
let on_demand_available = on_demand_enabled
|
||||
&& on_demand_cap.is_some_and(|value| value > 0.0)
|
||||
&& on_demand_used
|
||||
.zip(on_demand_cap)
|
||||
.is_some_and(|(used, cap)| used < cap);
|
||||
let usage_exhausted = remaining.is_some_and(|value| value <= 0.0)
|
||||
|| usage_ratio.is_some_and(|value| value >= 1.0 - 1e-6);
|
||||
let exhausted = usage_exhausted && !prepaid_available && !on_demand_available;
|
||||
let reason = if exhausted {
|
||||
Some("额度已耗尽".to_string())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let label = if exhausted {
|
||||
Some("额度耗尽")
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let code = if exhausted { "exhausted" } else { "ok" };
|
||||
|
||||
Some(json!({
|
||||
"version": 2,
|
||||
"provider_type": "xai",
|
||||
"code": code,
|
||||
"label": label,
|
||||
"reason": reason,
|
||||
"freshness": "fresh",
|
||||
"source": source,
|
||||
"observed_at": observed_at_unix_secs,
|
||||
"exhausted": exhausted,
|
||||
"usage_ratio": usage_ratio,
|
||||
"updated_at": observed_at_unix_secs,
|
||||
"reset_at": next_reset_at,
|
||||
"reset_seconds": reset_seconds,
|
||||
"plan_type": plan_type,
|
||||
"windows": windows,
|
||||
}))
|
||||
}
|
||||
|
||||
fn build_chatgpt_web_quota_status_snapshot(
|
||||
upstream_metadata: Option<&Value>,
|
||||
source: &str,
|
||||
@@ -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();
|
||||
|
||||
@@ -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")),
|
||||
|
||||
@@ -172,6 +172,7 @@ fn provider_uses_bearer_oauth_runtime(provider_type: &str) -> bool {
|
||||
| "antigravity"
|
||||
| "kiro"
|
||||
| "windsurf"
|
||||
| "xai"
|
||||
)
|
||||
}
|
||||
|
||||
@@ -406,6 +407,22 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn recognizes_xai_oauth_as_bearer_runtime() {
|
||||
let semantics = provider_key_auth_semantics(&sample_key("oauth"), "xai");
|
||||
|
||||
assert!(semantics.oauth_managed());
|
||||
assert!(semantics.can_refresh_oauth());
|
||||
assert_eq!(
|
||||
semantics.credential_kind(),
|
||||
ProviderKeyCredentialKind::OAuthSession
|
||||
);
|
||||
assert_eq!(
|
||||
semantics.runtime_auth_kind(),
|
||||
ProviderKeyRuntimeAuthKind::Bearer
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn refresh_capability_requires_stored_refresh_token() {
|
||||
let semantics = provider_key_auth_semantics(&sample_key("oauth"), "codex");
|
||||
|
||||
@@ -186,6 +186,8 @@ fn frontend_path_bypasses_static(path: &str) -> bool {
|
||||
"/health" | "/test-connection" | crate::constants::READYZ_PATH
|
||||
) || path.starts_with("/api/")
|
||||
|| path.starts_with("/v1/")
|
||||
|| path == "/openai/v1/videos"
|
||||
|| path.starts_with("/openai/v1/videos/")
|
||||
|| path.starts_with("/v1beta/")
|
||||
|| path.starts_with("/upload/")
|
||||
|| path.starts_with("/_gateway/")
|
||||
|
||||
@@ -290,6 +290,14 @@ impl provider_transport::VideoTaskTransportSnapshotLookup for AppState {
|
||||
.await
|
||||
.map_err(GatewayError::into_message)
|
||||
}
|
||||
|
||||
async fn resolve_video_task_proxy(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<ProxySnapshot> {
|
||||
self.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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![
|
||||
|
||||
@@ -983,6 +983,297 @@ async fn gateway_rejects_generic_oauth_start_for_windsurf_provider_impl() {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn gateway_rejects_generic_oauth_start_for_xai_provider() {
|
||||
run_admin_oauth_test(
|
||||
"gateway_rejects_generic_oauth_start_for_xai_provider",
|
||||
gateway_rejects_generic_oauth_start_for_xai_provider_impl,
|
||||
);
|
||||
}
|
||||
|
||||
async fn gateway_rejects_generic_oauth_start_for_xai_provider_impl() {
|
||||
let mut provider = sample_provider("provider-xai", "xai", 10);
|
||||
provider.provider_type = "xai".to_string();
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider],
|
||||
vec![],
|
||||
vec![],
|
||||
));
|
||||
let state = AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(GatewayDataState::with_provider_catalog_reader_for_tests(
|
||||
provider_catalog_repository,
|
||||
));
|
||||
|
||||
let response = local_admin_provider_oauth_response(
|
||||
&state,
|
||||
http::Method::POST,
|
||||
"/api/admin/provider-oauth/providers/provider-xai/start",
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
|
||||
let body = to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
.expect("body should read");
|
||||
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
|
||||
assert!(
|
||||
payload["detail"].as_str().is_some_and(|detail| {
|
||||
detail.contains("设备授权") || detail.contains("导入凭据")
|
||||
}),
|
||||
"payload={payload}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn gateway_handles_admin_provider_oauth_device_authorize_for_xai() {
|
||||
run_admin_oauth_test(
|
||||
"gateway_handles_admin_provider_oauth_device_authorize_for_xai",
|
||||
gateway_handles_admin_provider_oauth_device_authorize_for_xai_impl,
|
||||
);
|
||||
}
|
||||
|
||||
async fn gateway_handles_admin_provider_oauth_device_authorize_for_xai_impl() {
|
||||
let authorize_hits = Arc::new(Mutex::new(0usize));
|
||||
let authorize_hits_clone = Arc::clone(&authorize_hits);
|
||||
let oidc_server = Router::new().fallback(any(move |_request: Request| {
|
||||
let authorize_hits_inner = Arc::clone(&authorize_hits_clone);
|
||||
async move {
|
||||
*authorize_hits_inner.lock().expect("mutex should lock") += 1;
|
||||
Json(json!({
|
||||
"device_code": "xai-device-code",
|
||||
"user_code": "XAI-CODE",
|
||||
"verification_uri": "https://auth.x.ai/activate",
|
||||
"verification_uri_complete": "https://auth.x.ai/activate?user_code=XAI-CODE",
|
||||
"expires_in": 600,
|
||||
"interval": 5,
|
||||
}))
|
||||
}
|
||||
}));
|
||||
|
||||
let mut provider = sample_provider("provider-xai", "xai", 10);
|
||||
provider.provider_type = "xai".to_string();
|
||||
let endpoint = sample_endpoint(
|
||||
"endpoint-xai-responses",
|
||||
"provider-xai",
|
||||
"openai:responses",
|
||||
"https://cli-chat-proxy.grok.com/v1",
|
||||
);
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider],
|
||||
vec![endpoint],
|
||||
vec![],
|
||||
));
|
||||
let (oidc_url, oidc_handle) = start_server(oidc_server).await;
|
||||
let state = AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(GatewayDataState::with_provider_catalog_reader_for_tests(
|
||||
provider_catalog_repository,
|
||||
))
|
||||
.with_provider_oauth_token_url_for_tests(
|
||||
"xai_device",
|
||||
format!("{oidc_url}/oauth2/device/code"),
|
||||
)
|
||||
.with_provider_oauth_token_url_for_tests("xai", format!("{oidc_url}/oauth2/token"));
|
||||
|
||||
let response = local_admin_provider_oauth_response(
|
||||
&state,
|
||||
http::Method::POST,
|
||||
"/api/admin/provider-oauth/providers/provider-xai/device-authorize",
|
||||
Some(json!({ "proxy_node_id": "proxy-node-xai" })),
|
||||
)
|
||||
.await;
|
||||
let status = response.status();
|
||||
let body = to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
.expect("body should read");
|
||||
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
|
||||
assert_eq!(status, StatusCode::OK, "payload={payload}");
|
||||
let session_id = payload["session_id"]
|
||||
.as_str()
|
||||
.expect("session_id should exist")
|
||||
.to_string();
|
||||
assert_eq!(payload["user_code"], "XAI-CODE");
|
||||
assert_eq!(payload["verification_uri"], "https://auth.x.ai/activate");
|
||||
assert_eq!(
|
||||
payload["verification_uri_complete"],
|
||||
"https://auth.x.ai/activate?user_code=XAI-CODE"
|
||||
);
|
||||
assert_eq!(payload["auth_type"], "device");
|
||||
assert!(payload.get("callback_required").is_none() || payload["callback_required"] == false);
|
||||
assert_eq!(*authorize_hits.lock().expect("mutex should lock"), 1);
|
||||
|
||||
let stored = state
|
||||
.load_provider_oauth_device_session_for_tests(&format!("device_auth_session:{session_id}"))
|
||||
.expect("device session should be stored");
|
||||
let stored: serde_json::Value =
|
||||
serde_json::from_str(&stored).expect("device session json should parse");
|
||||
assert_eq!(stored["provider_id"], "provider-xai");
|
||||
assert_eq!(stored["device_code"], "xai-device-code");
|
||||
assert_eq!(stored["auth_type"], "device");
|
||||
assert_eq!(stored["redirect_uri"], format!("{oidc_url}/oauth2/token"));
|
||||
assert_eq!(stored["proxy_node_id"], "proxy-node-xai");
|
||||
assert_eq!(stored["status"], "pending");
|
||||
|
||||
oidc_handle.abort();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn gateway_handles_admin_provider_oauth_device_poll_for_xai() {
|
||||
run_admin_oauth_test(
|
||||
"gateway_handles_admin_provider_oauth_device_poll_for_xai",
|
||||
gateway_handles_admin_provider_oauth_device_poll_for_xai_impl,
|
||||
);
|
||||
}
|
||||
|
||||
async fn gateway_handles_admin_provider_oauth_device_poll_for_xai_impl() {
|
||||
let token_hits = Arc::new(Mutex::new(0usize));
|
||||
let token_hits_clone = Arc::clone(&token_hits);
|
||||
let access_token = sample_kiro_device_access_token("[email protected]");
|
||||
let id_token = access_token.clone();
|
||||
let token_server = Router::new().fallback(any(move |_request: Request| {
|
||||
let token_hits_inner = Arc::clone(&token_hits_clone);
|
||||
let access_token = access_token.clone();
|
||||
let id_token = id_token.clone();
|
||||
async move {
|
||||
let hit = {
|
||||
let mut hits = token_hits_inner.lock().expect("mutex should lock");
|
||||
*hits += 1;
|
||||
*hits
|
||||
};
|
||||
if hit == 1 {
|
||||
return (
|
||||
StatusCode::BAD_REQUEST,
|
||||
Json(json!({ "error": "authorization_pending" })),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
Json(json!({
|
||||
"access_token": access_token,
|
||||
"refresh_token": "xai-refresh-token",
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3600,
|
||||
"id_token": id_token,
|
||||
}))
|
||||
.into_response()
|
||||
}
|
||||
}));
|
||||
|
||||
let mut provider = sample_provider("provider-xai", "xai", 10);
|
||||
provider.provider_type = "xai".to_string();
|
||||
let endpoint = sample_endpoint(
|
||||
"endpoint-xai-responses",
|
||||
"provider-xai",
|
||||
"openai:responses",
|
||||
"https://cli-chat-proxy.grok.com/v1",
|
||||
);
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider],
|
||||
vec![endpoint],
|
||||
vec![],
|
||||
));
|
||||
let (token_url, token_handle) = start_server(token_server).await;
|
||||
let resolved_token_url = format!("{token_url}/oauth2/token");
|
||||
let state = AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_provider_catalog_repository_for_tests(
|
||||
provider_catalog_repository.clone(),
|
||||
)
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||
)
|
||||
.with_provider_oauth_device_session_entry_for_tests(
|
||||
"session-xai",
|
||||
json!({
|
||||
"provider_id": "provider-xai",
|
||||
"region": "",
|
||||
"client_id": "b1a00492-073a-47ea-816f-4c329264a828",
|
||||
"client_secret": "",
|
||||
"device_code": "xai-device-code",
|
||||
"auth_type": "device",
|
||||
"social_provider": null,
|
||||
"code_verifier": null,
|
||||
"redirect_uri": resolved_token_url,
|
||||
"machine_id": null,
|
||||
"interval": 5,
|
||||
"expires_at_unix_secs": 4_102_444_800u64,
|
||||
"status": "pending",
|
||||
"proxy_node_id": null,
|
||||
"created_at_unix_ms": 1_711_000_000u64,
|
||||
"key_id": null,
|
||||
"email": null,
|
||||
"replaced": false,
|
||||
"error_msg": null,
|
||||
}),
|
||||
)
|
||||
.with_provider_oauth_token_url_for_tests("xai", resolved_token_url.clone());
|
||||
|
||||
let pending = local_admin_provider_oauth_response(
|
||||
&state,
|
||||
http::Method::POST,
|
||||
"/api/admin/provider-oauth/providers/provider-xai/device-poll",
|
||||
Some(json!({ "session_id": "session-xai" })),
|
||||
)
|
||||
.await;
|
||||
let pending_body = to_bytes(pending.into_body(), usize::MAX)
|
||||
.await
|
||||
.expect("pending body should read");
|
||||
let pending_payload: serde_json::Value =
|
||||
serde_json::from_slice(&pending_body).expect("pending json should parse");
|
||||
assert_eq!(
|
||||
pending_payload["status"], "pending",
|
||||
"payload={pending_payload}"
|
||||
);
|
||||
|
||||
let authorized = local_admin_provider_oauth_response(
|
||||
&state,
|
||||
http::Method::POST,
|
||||
"/api/admin/provider-oauth/providers/provider-xai/device-poll",
|
||||
Some(json!({ "session_id": "session-xai" })),
|
||||
)
|
||||
.await;
|
||||
let status = authorized.status();
|
||||
let body = to_bytes(authorized.into_body(), usize::MAX)
|
||||
.await
|
||||
.expect("authorized body should read");
|
||||
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
|
||||
assert_eq!(status, StatusCode::OK, "payload={payload}");
|
||||
assert_eq!(payload["status"], "authorized");
|
||||
assert_eq!(payload["email"], "[email protected]");
|
||||
assert_eq!(payload["replaced"], false);
|
||||
assert_eq!(*token_hits.lock().expect("mutex should lock"), 2);
|
||||
|
||||
let stored = state
|
||||
.load_provider_oauth_device_session_for_tests("device_auth_session:session-xai")
|
||||
.expect("device session should persist");
|
||||
let stored: serde_json::Value =
|
||||
serde_json::from_str(&stored).expect("device session json should parse");
|
||||
assert_eq!(stored["status"], "authorized");
|
||||
let key_id = stored["key_id"]
|
||||
.as_str()
|
||||
.expect("key_id should be stored")
|
||||
.to_string();
|
||||
assert_eq!(payload["key_id"], key_id);
|
||||
|
||||
let persisted = provider_catalog_repository
|
||||
.list_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await
|
||||
.expect("keys should load")
|
||||
.into_iter()
|
||||
.next()
|
||||
.expect("persisted key should exist");
|
||||
assert_eq!(persisted.auth_type, "oauth");
|
||||
let decrypted_auth_config = decrypt_persisted_provider_auth_config(&persisted);
|
||||
let auth_config: serde_json::Value =
|
||||
serde_json::from_str(&decrypted_auth_config).expect("auth config should parse");
|
||||
assert_eq!(auth_config["provider_type"], "xai");
|
||||
assert_eq!(auth_config["auth_method"], "oauth");
|
||||
assert_eq!(auth_config["using_api"], false);
|
||||
assert_eq!(auth_config["email"], "[email protected]");
|
||||
|
||||
token_handle.abort();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn gateway_handles_admin_provider_oauth_device_poll_for_windsurf_one_time_token() {
|
||||
run_admin_oauth_test(
|
||||
|
||||
@@ -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<I, S>(node_ids: I) -> Arc<InMemoryProxyNodeRepository>
|
||||
where
|
||||
I: IntoIterator<Item = S>,
|
||||
S: AsRef<str>,
|
||||
{
|
||||
video_proxy_node_repository_at_url(node_ids, "http://127.0.0.1:1")
|
||||
}
|
||||
|
||||
pub(super) fn video_proxy_node_repository_at_url<I, S>(
|
||||
node_ids: I,
|
||||
proxy_url: &str,
|
||||
) -> Arc<InMemoryProxyNodeRepository>
|
||||
where
|
||||
I: IntoIterator<Item = S>,
|
||||
S: AsRef<str>,
|
||||
@@ -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<InMemoryProviderCatalogReadRepository> {
|
||||
video_provider_catalog_repository_with_proxy(
|
||||
provider_id,
|
||||
provider_type,
|
||||
endpoint_id,
|
||||
api_format,
|
||||
endpoint_base_url,
|
||||
key_id,
|
||||
upstream_api_key,
|
||||
None,
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) fn video_provider_catalog_repository_with_proxy(
|
||||
provider_id: &str,
|
||||
provider_type: &str,
|
||||
endpoint_id: &str,
|
||||
api_format: &str,
|
||||
endpoint_base_url: &str,
|
||||
key_id: &str,
|
||||
upstream_api_key: &str,
|
||||
proxy: Option<serde_json::Value>,
|
||||
) -> Arc<InMemoryProviderCatalogReadRepository> {
|
||||
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,
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -0,0 +1,427 @@
|
||||
use super::*;
|
||||
use aether_data::repository::auth::{
|
||||
InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot,
|
||||
};
|
||||
use aether_data::repository::candidate_selection::InMemoryMinimalCandidateSelectionReadRepository;
|
||||
use aether_data::repository::candidates::InMemoryRequestCandidateRepository;
|
||||
use aether_data_contracts::repository::candidate_selection::{
|
||||
StoredMinimalCandidateSelectionRow, StoredProviderModelMapping,
|
||||
};
|
||||
use sha2::{Digest, Sha256};
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
fn sample_auth_snapshot(api_key_id: &str, user_id: &str) -> StoredAuthApiKeySnapshot {
|
||||
StoredAuthApiKeySnapshot::new(
|
||||
user_id.to_string(),
|
||||
"video-user".to_string(),
|
||||
Some("[email protected]".to_string()),
|
||||
"user".to_string(),
|
||||
"local".to_string(),
|
||||
true,
|
||||
false,
|
||||
Some(json!(["openai"])),
|
||||
Some(json!(["openai:video"])),
|
||||
Some(json!(["video-model"])),
|
||||
api_key_id.to_string(),
|
||||
Some("default".to_string()),
|
||||
true,
|
||||
false,
|
||||
false,
|
||||
Some(60),
|
||||
Some(5),
|
||||
Some(4_102_444_800),
|
||||
Some(json!(["openai"])),
|
||||
Some(json!(["openai:video"])),
|
||||
Some(json!(["video-model"])),
|
||||
)
|
||||
.expect("auth snapshot should build")
|
||||
}
|
||||
|
||||
fn sample_candidate_row() -> StoredMinimalCandidateSelectionRow {
|
||||
StoredMinimalCandidateSelectionRow {
|
||||
provider_id: "provider-openai-video-local-1".to_string(),
|
||||
provider_name: "openai".to_string(),
|
||||
provider_type: "xai".to_string(),
|
||||
provider_priority: 10,
|
||||
provider_is_active: true,
|
||||
endpoint_id: "endpoint-openai-video-local-1".to_string(),
|
||||
endpoint_api_format: "openai:video".to_string(),
|
||||
endpoint_api_family: Some("openai".to_string()),
|
||||
endpoint_kind: Some("video".to_string()),
|
||||
endpoint_is_active: true,
|
||||
key_id: "key-openai-video-local-1".to_string(),
|
||||
key_name: "prod".to_string(),
|
||||
key_auth_type: "api_key".to_string(),
|
||||
key_is_active: true,
|
||||
key_api_formats: Some(vec!["openai:video".to_string()]),
|
||||
key_allowed_models: None,
|
||||
key_capabilities: None,
|
||||
key_internal_priority: 5,
|
||||
key_global_priority_by_format: Some(json!({"openai:video": 1})),
|
||||
model_id: "model-openai-video-local-1".to_string(),
|
||||
global_model_id: "global-model-openai-video-local-1".to_string(),
|
||||
global_model_name: "video-model".to_string(),
|
||||
global_model_mappings: None,
|
||||
global_model_supports_streaming: Some(false),
|
||||
model_provider_model_name: "grok-imagine-video".to_string(),
|
||||
model_provider_model_mappings: Some(vec![StoredProviderModelMapping {
|
||||
name: "grok-imagine-video".to_string(),
|
||||
priority: 1,
|
||||
api_formats: Some(vec!["openai:video".to_string()]),
|
||||
endpoint_ids: None,
|
||||
operations: None,
|
||||
}]),
|
||||
model_supports_streaming: Some(false),
|
||||
model_is_active: true,
|
||||
model_is_available: true,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn xai_video_native_and_compatibility_http_lifecycle() {
|
||||
Box::pin(assert_xai_video_http_lifecycle(Arc::new(
|
||||
InMemoryVideoTaskRepository::default(),
|
||||
)))
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn xai_video_native_and_compatibility_http_lifecycle_postgres() {
|
||||
let configured_database_url = std::env::var("AETHER_TEST_DATABASE_URL").ok();
|
||||
let managed_database = if configured_database_url.is_none() {
|
||||
Some(
|
||||
aether_testkit::ManagedPostgresServer::start()
|
||||
.await
|
||||
.expect("temporary PostgreSQL should start"),
|
||||
)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let database_url = configured_database_url.unwrap_or_else(|| {
|
||||
managed_database
|
||||
.as_ref()
|
||||
.expect("managed test database should exist")
|
||||
.database_url()
|
||||
.to_string()
|
||||
});
|
||||
let pool = sqlx::postgres::PgPoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect(&database_url)
|
||||
.await
|
||||
.expect("test database should connect");
|
||||
aether_data::driver::postgres::run_migrations(&pool)
|
||||
.await
|
||||
.expect("test database should migrate");
|
||||
// Preserve the production column constraints and unique indexes while isolating test rows.
|
||||
sqlx::query("CREATE TEMP TABLE video_tasks (LIKE public.video_tasks INCLUDING ALL)")
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("isolated video task table should be created");
|
||||
let repository =
|
||||
Arc::new(aether_data::repository::video_tasks::SqlxVideoTaskRepository::new(pool.clone()));
|
||||
Box::pin(assert_xai_video_http_lifecycle(repository)).await;
|
||||
pool.close().await;
|
||||
}
|
||||
|
||||
async fn assert_xai_video_http_lifecycle<T>(repository: Arc<T>)
|
||||
where
|
||||
T: aether_data_contracts::repository::video_tasks::VideoTaskRepository + 'static,
|
||||
{
|
||||
let static_dir = std::env::temp_dir().join(format!(
|
||||
"aether-xai-video-static-{}",
|
||||
std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_nanos()
|
||||
));
|
||||
std::fs::create_dir_all(&static_dir).unwrap();
|
||||
std::fs::write(
|
||||
static_dir.join("index.html"),
|
||||
"<html>Aether test frontend</html>",
|
||||
)
|
||||
.unwrap();
|
||||
let seen = Arc::new(Mutex::new(Vec::<serde_json::Value>::new()));
|
||||
let calls = Arc::new(AtomicUsize::new(0));
|
||||
// Exercise the real HTTP executor, including production method gates, instead of
|
||||
// the test execution-runtime override that used to hide rejected GET requests.
|
||||
let video_url = Arc::new(Mutex::new(String::new()));
|
||||
let runtime = Router::new()
|
||||
.route("/v1/videos/{operation}", any({
|
||||
let seen = seen.clone();
|
||||
let calls = calls.clone();
|
||||
let video_url = video_url.clone();
|
||||
move |request: Request| {
|
||||
let seen = seen.clone();
|
||||
let calls = calls.clone();
|
||||
let video_url = video_url.clone();
|
||||
async move {
|
||||
let (parts, body) = request.into_parts();
|
||||
assert_eq!(parts.headers["authorization"], "Bearer upstream-video-key");
|
||||
let bytes = to_bytes(body, usize::MAX).await.unwrap();
|
||||
let body: serde_json::Value = serde_json::from_slice(&bytes).unwrap_or(json!(null));
|
||||
seen.lock().unwrap().push(json!({
|
||||
"method": parts.method.as_str(),
|
||||
"url": parts.uri.path(),
|
||||
"body": {"json_body": body}
|
||||
}));
|
||||
let response = if parts.method == http::Method::POST {
|
||||
json!({"request_id":"upstream-video-id", "provider_extension":{"accepted":true}})
|
||||
} else {
|
||||
assert_eq!(parts.uri.path(), "/v1/videos/upstream-video-id");
|
||||
if calls.fetch_add(1, Ordering::SeqCst) == 0 {
|
||||
json!({"status":"pending"})
|
||||
} else {
|
||||
json!({"status":"done", "model":"grok-imagine-video", "video":{"url":video_url.lock().unwrap().clone(), "duration":6, "respect_moderation":true}, "provider_extension":"preserved"})
|
||||
}
|
||||
};
|
||||
Json(response)
|
||||
}
|
||||
}
|
||||
}))
|
||||
.route("/test.mp4", any(|request: Request| async move {
|
||||
assert!(request.headers().get("authorization").is_none());
|
||||
assert!(request.headers().get("x-xai-token-auth").is_none());
|
||||
([("content-type", "video/mp4")], "test-video-bytes")
|
||||
}));
|
||||
let (runtime_url, runtime_handle) = start_server(runtime).await;
|
||||
let expected_video_url = format!("{runtime_url}/test.mp4");
|
||||
*video_url.lock().unwrap() = expected_video_url.clone();
|
||||
let state_factory = || {
|
||||
let auth = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![
|
||||
(
|
||||
Some(format!("{:x}", Sha256::digest(b"owner-key"))),
|
||||
sample_auth_snapshot("owner-api-key", "owner"),
|
||||
),
|
||||
(
|
||||
Some(format!("{:x}", Sha256::digest(b"foreign-key"))),
|
||||
sample_auth_snapshot("foreign-api-key", "foreign"),
|
||||
),
|
||||
]));
|
||||
let candidates = Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
|
||||
sample_candidate_row(),
|
||||
]));
|
||||
let catalog = video_provider_catalog_repository_with_proxy(
|
||||
"provider-openai-video-local-1",
|
||||
"xai",
|
||||
"endpoint-openai-video-local-1",
|
||||
"openai:video",
|
||||
"http://video-provider.invalid/v1",
|
||||
"key-openai-video-local-1",
|
||||
"upstream-video-key",
|
||||
Some(json!({"enabled":true,"node_id":"video-proxy"})),
|
||||
);
|
||||
AppState::new().expect("gateway should build").with_video_task_truth_source_mode(VideoTaskTruthSourceMode::RustAuthoritative).with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
|
||||
auth, candidates, catalog, Arc::new(InMemoryRequestCandidateRepository::default()), DEVELOPMENT_ENCRYPTION_KEY
|
||||
).attach_video_task_repository_for_tests(repository.clone())
|
||||
.attach_proxy_node_repository_for_tests(video_proxy_node_repository_at_url(["video-proxy"], &runtime_url))
|
||||
)
|
||||
};
|
||||
let router_factory =
|
||||
|| crate::attach_static_frontend(build_router_with_state(state_factory()), &static_dir);
|
||||
let (gateway_url, gateway_handle) = start_server(router_factory()).await;
|
||||
let client = reqwest::Client::new();
|
||||
assert_eq!(
|
||||
client
|
||||
.get(&gateway_url)
|
||||
.send()
|
||||
.await
|
||||
.unwrap()
|
||||
.text()
|
||||
.await
|
||||
.unwrap(),
|
||||
"<html>Aether test frontend</html>"
|
||||
);
|
||||
for (path, native) in [
|
||||
("/v1/videos/generations", true),
|
||||
("/v1/videos", true),
|
||||
("/v1/videos/edits", true),
|
||||
("/v1/videos/extensions", true),
|
||||
("/openai/v1/videos", false),
|
||||
] {
|
||||
calls.store(0, Ordering::SeqCst);
|
||||
let body = if native {
|
||||
json!({"model":"video-model","prompt":"A cat","duration":6,"aspect_ratio":"1:1","video":{"url":"https://example.com/input.mp4"},"future_option":true})
|
||||
} else {
|
||||
json!({"model":"video-model","prompt":"A cat","seconds":"6","size":"1280x720"})
|
||||
};
|
||||
let response = client
|
||||
.post(format!("{gateway_url}{path}"))
|
||||
.bearer_auth("owner-key")
|
||||
.json(&body)
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
let status = response.status();
|
||||
let result: serde_json::Value = response.json().await.unwrap();
|
||||
assert_eq!(status, StatusCode::OK, "{path}: {result}");
|
||||
let id = result[if native { "request_id" } else { "id" }]
|
||||
.as_str()
|
||||
.unwrap();
|
||||
assert_ne!(id, "upstream-video-id");
|
||||
if native {
|
||||
assert!(result.get("id").is_none());
|
||||
assert_eq!(result["provider_extension"]["accepted"], true);
|
||||
} else {
|
||||
assert_eq!(result["status"], "queued");
|
||||
}
|
||||
let request = seen.lock().unwrap().last().unwrap().clone();
|
||||
let suffix = if path.ends_with("/edits") {
|
||||
"edits"
|
||||
} else if path.ends_with("/extensions") {
|
||||
"extensions"
|
||||
} else {
|
||||
"generations"
|
||||
};
|
||||
assert_eq!(request["url"], format!("/v1/videos/{suffix}"));
|
||||
assert_eq!(request["body"]["json_body"]["model"], "grok-imagine-video");
|
||||
assert_eq!(request["body"]["json_body"]["duration"], 6);
|
||||
if native {
|
||||
assert_eq!(request["body"]["json_body"]["future_option"], true);
|
||||
} else {
|
||||
assert_eq!(request["body"]["json_body"]["aspect_ratio"], "16:9");
|
||||
assert_eq!(request["body"]["json_body"]["resolution"], "720p");
|
||||
assert!(request["body"]["json_body"].get("seconds").is_none());
|
||||
assert!(request["body"]["json_body"].get("size").is_none());
|
||||
}
|
||||
let query = format!(
|
||||
"{gateway_url}{}/{id}",
|
||||
if native {
|
||||
"/v1/videos"
|
||||
} else {
|
||||
"/openai/v1/videos"
|
||||
}
|
||||
);
|
||||
let before = seen.lock().unwrap().len();
|
||||
let denied = client
|
||||
.get(&query)
|
||||
.bearer_auth("foreign-key")
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(denied.status(), StatusCode::NOT_FOUND);
|
||||
assert_eq!(seen.lock().unwrap().len(), before);
|
||||
let denied_content = client
|
||||
.get(format!("{gateway_url}/openai/v1/videos/{id}/content"))
|
||||
.bearer_auth("foreign-key")
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(denied_content.status(), StatusCode::NOT_FOUND);
|
||||
assert_eq!(seen.lock().unwrap().len(), before);
|
||||
let pending: serde_json::Value = client
|
||||
.get(&query)
|
||||
.bearer_auth("owner-key")
|
||||
.send()
|
||||
.await
|
||||
.unwrap()
|
||||
.json()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
pending["status"],
|
||||
if native { "pending" } else { "queued" },
|
||||
"{path}: {pending}"
|
||||
);
|
||||
let done: serde_json::Value = client
|
||||
.get(&query)
|
||||
.bearer_auth("owner-key")
|
||||
.send()
|
||||
.await
|
||||
.unwrap()
|
||||
.json()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(done["status"], if native { "done" } else { "completed" });
|
||||
if native {
|
||||
assert_eq!(done["video"]["respect_moderation"], true);
|
||||
assert_eq!(done["provider_extension"], "preserved");
|
||||
} else {
|
||||
assert_eq!(done["video_url"], expected_video_url);
|
||||
}
|
||||
let stored = repository
|
||||
.find(VideoTaskLookupKey::Id(id))
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
stored.client_api_format.as_deref(),
|
||||
Some(if native { "xai:video" } else { "openai:video" })
|
||||
);
|
||||
assert_eq!(
|
||||
stored.external_task_id.as_deref(),
|
||||
Some("upstream-video-id")
|
||||
);
|
||||
assert!(stored.request_metadata.is_none());
|
||||
assert!(stored.original_request_body.is_none());
|
||||
// A new gateway instance must reconstruct the pinned provider/credential and protocol.
|
||||
let (restart_url, restart_handle) = start_server(router_factory()).await;
|
||||
let restored: serde_json::Value = client
|
||||
.get(format!(
|
||||
"{restart_url}{}/{id}",
|
||||
if native {
|
||||
"/v1/videos"
|
||||
} else {
|
||||
"/openai/v1/videos"
|
||||
}
|
||||
))
|
||||
.bearer_auth("owner-key")
|
||||
.send()
|
||||
.await
|
||||
.unwrap()
|
||||
.json()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(restored["status"], done["status"]);
|
||||
if native {
|
||||
assert_eq!(restored["video"]["respect_moderation"], true);
|
||||
}
|
||||
let compat: serde_json::Value = client
|
||||
.get(format!("{restart_url}/openai/v1/videos/{id}"))
|
||||
.bearer_auth("owner-key")
|
||||
.send()
|
||||
.await
|
||||
.unwrap()
|
||||
.json()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(compat["status"], "completed");
|
||||
assert_eq!(compat["video_url"], expected_video_url);
|
||||
let native_view: serde_json::Value = client
|
||||
.get(format!("{restart_url}/v1/videos/{id}"))
|
||||
.bearer_auth("owner-key")
|
||||
.send()
|
||||
.await
|
||||
.unwrap()
|
||||
.json()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(native_view["status"], "done");
|
||||
assert_eq!(native_view["video"]["respect_moderation"], true);
|
||||
for prefix in ["/v1/videos", "/openai/v1/videos"] {
|
||||
let content = client
|
||||
.get(format!("{restart_url}{prefix}/{id}/content"))
|
||||
.bearer_auth("owner-key")
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(content.status(), StatusCode::OK);
|
||||
assert_eq!(content.headers()["content-type"], "video/mp4");
|
||||
assert_eq!(content.bytes().await.unwrap(), "test-video-bytes");
|
||||
}
|
||||
restart_handle.abort();
|
||||
}
|
||||
let before = seen.lock().unwrap().len();
|
||||
let bad = client
|
||||
.post(format!("{gateway_url}/openai/v1/videos"))
|
||||
.bearer_auth("owner-key")
|
||||
.json(&json!({"model":"video-model","prompt":"cat","seconds":"wrong"}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(bad.status(), StatusCode::BAD_REQUEST);
|
||||
assert_eq!(seen.lock().unwrap().len(), before);
|
||||
gateway_handle.abort();
|
||||
runtime_handle.abort();
|
||||
std::fs::remove_dir_all(&static_dir).unwrap();
|
||||
}
|
||||
@@ -11,6 +11,9 @@ use super::{
|
||||
fn rust_authoritative_service_builds_openai_cancel_follow_up_plan() {
|
||||
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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user