mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-12 04:09:48 +08:00
Merge pull request #822 from stabey/upstream-pr/xai-media
feat(providers): 新增 xAI Provider(设备码 OAuth + 原生图像/视频)
This commit is contained in:
Generated
+1
@@ -748,6 +748,7 @@ dependencies = [
|
|||||||
"async-trait",
|
"async-trait",
|
||||||
"serde",
|
"serde",
|
||||||
"serde_json",
|
"serde_json",
|
||||||
|
"sha2",
|
||||||
"url",
|
"url",
|
||||||
"uuid",
|
"uuid",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -583,6 +583,11 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
|||||||
source_model,
|
source_model,
|
||||||
codex_model_capabilities.as_ref(),
|
codex_model_capabilities.as_ref(),
|
||||||
);
|
);
|
||||||
|
crate::ai_serving::transport::xai::insert_cli_identity_headers_if_needed(
|
||||||
|
transport.as_ref(),
|
||||||
|
prepared.provider_api_format.as_str(),
|
||||||
|
&mut provider_request_headers,
|
||||||
|
);
|
||||||
request_identity_response_encoding_when_redacted(
|
request_identity_response_encoding_when_redacted(
|
||||||
&mut provider_request_headers,
|
&mut provider_request_headers,
|
||||||
redaction.redacted,
|
redaction.redacted,
|
||||||
|
|||||||
@@ -17,8 +17,8 @@ use crate::ai_serving::transport::{
|
|||||||
ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput, GROK_CHAT_PATH,
|
ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput, GROK_CHAT_PATH,
|
||||||
};
|
};
|
||||||
use crate::ai_serving::{
|
use crate::ai_serving::{
|
||||||
apply_codex_openai_special_headers, build_chatgpt_web_image_request_body,
|
apply_codex_openai_special_headers, apply_xai_upstream_payload_edits,
|
||||||
build_codex_openai_image_api_provider_request_body,
|
build_chatgpt_web_image_request_body, build_codex_openai_image_api_provider_request_body,
|
||||||
build_gemini_image_request_body_from_openai_image_request,
|
build_gemini_image_request_body_from_openai_image_request,
|
||||||
build_openai_image_api_provider_request_body, build_openai_image_provider_request_body,
|
build_openai_image_api_provider_request_body, build_openai_image_provider_request_body,
|
||||||
default_model_for_openai_image_operation, normalize_openai_image_request,
|
default_model_for_openai_image_operation, normalize_openai_image_request,
|
||||||
@@ -211,7 +211,7 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
|
|||||||
upstream_is_stream,
|
upstream_is_stream,
|
||||||
)
|
)
|
||||||
};
|
};
|
||||||
let Some(provider_request_body) = provider_request_body else {
|
let Some(mut provider_request_body) = provider_request_body else {
|
||||||
mark_skipped_local_openai_image_candidate_with_failure_diagnostic(
|
mark_skipped_local_openai_image_candidate_with_failure_diagnostic(
|
||||||
state,
|
state,
|
||||||
input,
|
input,
|
||||||
@@ -229,6 +229,11 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
|
|||||||
.await;
|
.await;
|
||||||
return None;
|
return None;
|
||||||
};
|
};
|
||||||
|
apply_xai_upstream_payload_edits(
|
||||||
|
&mut provider_request_body,
|
||||||
|
transport.provider.provider_type.as_str(),
|
||||||
|
provider_api_format,
|
||||||
|
);
|
||||||
let Some(mut provider_request_headers) = (if is_grok {
|
let Some(mut provider_request_headers) = (if is_grok {
|
||||||
build_grok_browser_headers(GrokHeaderInput {
|
build_grok_browser_headers(GrokHeaderInput {
|
||||||
transport,
|
transport,
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ use crate::ai_serving::planner::{
|
|||||||
build_ai_execution_decision_response, resolve_transport_request_encoding_policy,
|
build_ai_execution_decision_response, resolve_transport_request_encoding_policy,
|
||||||
AiExecutionDecisionResponseParts,
|
AiExecutionDecisionResponseParts,
|
||||||
};
|
};
|
||||||
|
use crate::ai_serving::transport::xai::video::is_native_video_request;
|
||||||
use crate::ai_serving::transport::{
|
use crate::ai_serving::transport::{
|
||||||
resolve_transport_execution_timeouts, resolve_transport_profile,
|
resolve_transport_execution_timeouts, resolve_transport_profile,
|
||||||
};
|
};
|
||||||
@@ -33,7 +34,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
|
|||||||
let Some(resolved) = resolve_local_video_create_candidate_payload_parts(
|
let Some(resolved) = resolve_local_video_create_candidate_payload_parts(
|
||||||
state, parts, body_json, trace_id, input, &attempt, spec,
|
state, parts, body_json, trace_id, input, &attempt, spec,
|
||||||
)
|
)
|
||||||
.await
|
.await?
|
||||||
else {
|
else {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
@@ -52,9 +53,32 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
|
|||||||
.await;
|
.await;
|
||||||
let transport_profile = resolve_transport_profile(&transport);
|
let transport_profile = resolve_transport_profile(&transport);
|
||||||
let mut extra_fields = serde_json::Map::new();
|
let mut extra_fields = serde_json::Map::new();
|
||||||
|
if is_native_video_request(&transport.provider.provider_type, parts.uri.path()) {
|
||||||
|
extra_fields.insert(
|
||||||
|
"video_client_protocol".to_string(),
|
||||||
|
serde_json::json!("xai"),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) {
|
if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) {
|
||||||
extra_fields.insert("proxy".to_string(), proxy_value);
|
extra_fields.insert("proxy".to_string(), proxy_value);
|
||||||
}
|
}
|
||||||
|
if transport.provider.provider_type.eq_ignore_ascii_case("xai") {
|
||||||
|
extra_fields.insert("video_provider_xai".into(), serde_json::json!(true));
|
||||||
|
if let Some(duration) = resolved.provider_request_body.get("duration") {
|
||||||
|
extra_fields.insert("video_duration".into(), duration.clone());
|
||||||
|
}
|
||||||
|
if parts.uri.path() == "/openai/v1/videos" {
|
||||||
|
extra_fields.insert(
|
||||||
|
"video_size".into(),
|
||||||
|
body_json
|
||||||
|
.get("size")
|
||||||
|
.filter(|v| v.as_str().is_some_and(|s| !s.trim().is_empty()))
|
||||||
|
.cloned()
|
||||||
|
.unwrap_or_else(|| serde_json::json!("720x1280")),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
let effective_headers = input.effective_headers(&parts.headers);
|
let effective_headers = input.effective_headers(&parts.headers);
|
||||||
let report_context = build_local_execution_report_context(LocalExecutionReportContextParts {
|
let report_context = build_local_execution_report_context(LocalExecutionReportContextParts {
|
||||||
auth_context: &input.auth_context,
|
auth_context: &input.auth_context,
|
||||||
|
|||||||
@@ -3,15 +3,23 @@ use std::sync::Arc;
|
|||||||
|
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
|
|
||||||
use crate::ai_serving::planner::candidate_preparation::resolve_candidate_mapped_model;
|
use crate::ai_serving::planner::candidate_preparation::{
|
||||||
|
prepare_header_authenticated_candidate, resolve_candidate_mapped_model, OauthPreparationContext,
|
||||||
|
};
|
||||||
use crate::ai_serving::planner::spec_metadata::local_video_create_spec_metadata;
|
use crate::ai_serving::planner::spec_metadata::local_video_create_spec_metadata;
|
||||||
|
use crate::ai_serving::transport::xai::video::{
|
||||||
|
convert_openai_video_request, is_explicit_native_video_path, is_native_video_request,
|
||||||
|
};
|
||||||
use crate::ai_serving::transport::{
|
use crate::ai_serving::transport::{
|
||||||
build_video_create_headers, build_video_create_request_body, build_video_create_upstream_url,
|
build_video_create_headers, build_video_create_request_body, build_video_create_upstream_url,
|
||||||
resolve_video_create_auth, video_create_transport_unsupported_reason,
|
resolve_video_create_auth, video_create_transport_unsupported_reason,
|
||||||
ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput,
|
ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput,
|
||||||
};
|
};
|
||||||
use crate::ai_serving::{CandidateFailureDiagnostic, GatewayProviderTransportSnapshot};
|
use crate::ai_serving::{
|
||||||
use crate::AppState;
|
apply_xai_upstream_payload_edits, CandidateFailureDiagnostic, GatewayProviderTransportSnapshot,
|
||||||
|
PlannerAppState,
|
||||||
|
};
|
||||||
|
use crate::{AppState, GatewayError};
|
||||||
|
|
||||||
use super::support::{
|
use super::support::{
|
||||||
mark_skipped_local_video_candidate, mark_skipped_local_video_candidate_with_failure_diagnostic,
|
mark_skipped_local_video_candidate, mark_skipped_local_video_candidate_with_failure_diagnostic,
|
||||||
@@ -37,11 +45,16 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
|||||||
input: &LocalVideoCreateDecisionInput,
|
input: &LocalVideoCreateDecisionInput,
|
||||||
attempt: &LocalVideoCreateCandidateAttempt,
|
attempt: &LocalVideoCreateCandidateAttempt,
|
||||||
spec: LocalVideoCreateSpec,
|
spec: LocalVideoCreateSpec,
|
||||||
) -> Option<LocalVideoCreateCandidatePayloadParts> {
|
) -> Result<Option<LocalVideoCreateCandidatePayloadParts>, GatewayError> {
|
||||||
let spec_metadata = local_video_create_spec_metadata(spec);
|
let spec_metadata = local_video_create_spec_metadata(spec);
|
||||||
let candidate = &attempt.eligible.candidate;
|
let candidate = &attempt.eligible.candidate;
|
||||||
let transport = &attempt.eligible.transport;
|
let transport = &attempt.eligible.transport;
|
||||||
let effective_headers = input.effective_headers(&parts.headers);
|
let effective_headers = input.effective_headers(&parts.headers);
|
||||||
|
if is_explicit_native_video_path(parts.uri.path())
|
||||||
|
&& !transport.provider.provider_type.eq_ignore_ascii_case("xai")
|
||||||
|
{
|
||||||
|
return Ok(None);
|
||||||
|
}
|
||||||
|
|
||||||
let provider_family = provider_video_create_family(spec.family);
|
let provider_family = provider_video_create_family(spec.family);
|
||||||
let transport_unsupported_reason = video_create_transport_unsupported_reason(
|
let transport_unsupported_reason = video_create_transport_unsupported_reason(
|
||||||
@@ -60,23 +73,39 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
|||||||
skip_reason,
|
skip_reason,
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
return None;
|
return Ok(None);
|
||||||
}
|
}
|
||||||
|
|
||||||
let auth = resolve_video_create_auth(transport, provider_family);
|
let prepared_candidate = match prepare_header_authenticated_candidate(
|
||||||
let Some((auth_header, auth_value)) = auth else {
|
PlannerAppState::new(state),
|
||||||
mark_skipped_local_video_candidate(
|
transport,
|
||||||
state,
|
candidate,
|
||||||
input,
|
resolve_video_create_auth(transport, provider_family),
|
||||||
|
OauthPreparationContext {
|
||||||
trace_id,
|
trace_id,
|
||||||
candidate,
|
api_format: spec_metadata.api_format,
|
||||||
attempt.candidate_index,
|
operation: "video_create_candidate_request",
|
||||||
&attempt.candidate_id,
|
},
|
||||||
"transport_auth_unavailable",
|
)
|
||||||
)
|
.await
|
||||||
.await;
|
{
|
||||||
return None;
|
Ok(prepared) => prepared,
|
||||||
|
Err(skip_reason) => {
|
||||||
|
mark_skipped_local_video_candidate(
|
||||||
|
state,
|
||||||
|
input,
|
||||||
|
trace_id,
|
||||||
|
candidate,
|
||||||
|
attempt.candidate_index,
|
||||||
|
&attempt.candidate_id,
|
||||||
|
skip_reason,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
return Ok(None);
|
||||||
|
}
|
||||||
};
|
};
|
||||||
|
let auth_header = prepared_candidate.auth_header;
|
||||||
|
let auth_value = prepared_candidate.auth_value;
|
||||||
|
|
||||||
let mapped_model = match resolve_candidate_mapped_model(candidate) {
|
let mapped_model = match resolve_candidate_mapped_model(candidate) {
|
||||||
Ok(mapped_model) => mapped_model,
|
Ok(mapped_model) => mapped_model,
|
||||||
@@ -91,7 +120,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
|||||||
skip_reason,
|
skip_reason,
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
return None;
|
return Ok(None);
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -117,10 +146,10 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
return None;
|
return Ok(None);
|
||||||
};
|
};
|
||||||
|
|
||||||
let Some(provider_request_body) = build_video_create_request_body(
|
let Some(mut provider_request_body) = build_video_create_request_body(
|
||||||
body_json,
|
body_json,
|
||||||
provider_family,
|
provider_family,
|
||||||
&mapped_model,
|
&mapped_model,
|
||||||
@@ -142,11 +171,28 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
return None;
|
return Ok(None);
|
||||||
};
|
};
|
||||||
|
if transport.provider.provider_type.eq_ignore_ascii_case("xai")
|
||||||
|
&& !is_native_video_request(&transport.provider.provider_type, parts.uri.path())
|
||||||
|
{
|
||||||
|
provider_request_body =
|
||||||
|
convert_openai_video_request(&provider_request_body).map_err(|message| {
|
||||||
|
GatewayError::Client {
|
||||||
|
status: http::StatusCode::BAD_REQUEST,
|
||||||
|
message: message.to_string(),
|
||||||
|
}
|
||||||
|
})?;
|
||||||
|
}
|
||||||
|
apply_xai_upstream_payload_edits(
|
||||||
|
&mut provider_request_body,
|
||||||
|
transport.provider.provider_type.as_str(),
|
||||||
|
spec_metadata.api_format,
|
||||||
|
);
|
||||||
|
|
||||||
let Some(provider_request_headers) =
|
let Some(provider_request_headers) =
|
||||||
build_video_create_headers(ProviderVideoCreateHeadersInput {
|
build_video_create_headers(ProviderVideoCreateHeadersInput {
|
||||||
|
transport,
|
||||||
headers: effective_headers,
|
headers: effective_headers,
|
||||||
auth_header: &auth_header,
|
auth_header: &auth_header,
|
||||||
auth_value: &auth_value,
|
auth_value: &auth_value,
|
||||||
@@ -170,10 +216,10 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
return None;
|
return Ok(None);
|
||||||
};
|
};
|
||||||
|
|
||||||
Some(LocalVideoCreateCandidatePayloadParts {
|
Ok(Some(LocalVideoCreateCandidatePayloadParts {
|
||||||
transport: Arc::clone(transport),
|
transport: Arc::clone(transport),
|
||||||
auth_header,
|
auth_header,
|
||||||
auth_value,
|
auth_value,
|
||||||
@@ -181,7 +227,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
|||||||
provider_request_headers,
|
provider_request_headers,
|
||||||
provider_request_body,
|
provider_request_body,
|
||||||
upstream_url,
|
upstream_url,
|
||||||
})
|
}))
|
||||||
}
|
}
|
||||||
|
|
||||||
fn provider_video_create_family(family: LocalVideoCreateFamily) -> ProviderVideoCreateFamily {
|
fn provider_video_create_family(family: LocalVideoCreateFamily) -> ProviderVideoCreateFamily {
|
||||||
|
|||||||
@@ -13,7 +13,9 @@ pub(crate) fn openai_responses_reasoning_replay_policy(
|
|||||||
base_url: &str,
|
base_url: &str,
|
||||||
_provider_model: &str,
|
_provider_model: &str,
|
||||||
) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy {
|
) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy {
|
||||||
if is_deepseek_provider(provider_type, base_url) {
|
if provider_type.trim().eq_ignore_ascii_case("xai") {
|
||||||
|
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
|
||||||
|
} else if is_deepseek_provider(provider_type, base_url) {
|
||||||
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
|
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
|
||||||
} else {
|
} else {
|
||||||
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
|
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
|
||||||
@@ -238,6 +240,27 @@ mod tests {
|
|||||||
openai_responses_reasoning_replay_policy,
|
openai_responses_reasoning_replay_policy,
|
||||||
};
|
};
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_reasoning_policy_comes_from_provider_type() {
|
||||||
|
use crate::ai_serving::OpenAiResponsesReasoningReplayPolicy;
|
||||||
|
assert_eq!(
|
||||||
|
openai_responses_reasoning_replay_policy(
|
||||||
|
"xai",
|
||||||
|
"https://custom.example/v1",
|
||||||
|
"grok-4.6"
|
||||||
|
),
|
||||||
|
OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
openai_responses_reasoning_replay_policy(
|
||||||
|
"openai",
|
||||||
|
"https://custom.example/v1",
|
||||||
|
"grok-4.6"
|
||||||
|
),
|
||||||
|
OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn detects_deepseek_provider_only_by_official_host() {
|
fn detects_deepseek_provider_only_by_official_host() {
|
||||||
assert!(!is_deepseek_provider(
|
assert!(!is_deepseek_provider(
|
||||||
|
|||||||
@@ -488,6 +488,7 @@ impl ResponsesWebSocketBodyNormalization {
|
|||||||
digest.update([match self.reasoning_replay_policy {
|
digest.update([match self.reasoning_replay_policy {
|
||||||
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds => 0,
|
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds => 0,
|
||||||
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque => 1,
|
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque => 1,
|
||||||
|
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted => 2,
|
||||||
}]);
|
}]);
|
||||||
update_normalization_optional_json_digest(&mut digest, self.model_directive_patch.as_ref());
|
update_normalization_optional_json_digest(&mut digest, self.model_directive_patch.as_ref());
|
||||||
digest.finalize().into()
|
digest.finalize().into()
|
||||||
|
|||||||
@@ -12,7 +12,8 @@ pub(crate) use aether_ai_formats::api::{
|
|||||||
apply_codex_openai_responses_websocket_continuation_body_edits_with_source_model_and_capabilities,
|
apply_codex_openai_responses_websocket_continuation_body_edits_with_source_model_and_capabilities,
|
||||||
apply_codex_openai_special_headers, apply_model_directive_mapping_patch,
|
apply_codex_openai_special_headers, apply_model_directive_mapping_patch,
|
||||||
apply_model_directive_overrides_from_model, apply_model_directive_overrides_from_request,
|
apply_model_directive_overrides_from_model, apply_model_directive_overrides_from_request,
|
||||||
apply_openai_responses_compact_special_body_edits, build_chatgpt_web_image_request_body,
|
apply_openai_responses_compact_special_body_edits, apply_xai_upstream_payload_edits,
|
||||||
|
apply_xai_upstream_payload_edits_with_client, build_chatgpt_web_image_request_body,
|
||||||
build_codex_model_catalog_metadata, build_codex_openai_image_api_provider_request_body,
|
build_codex_model_catalog_metadata, build_codex_openai_image_api_provider_request_body,
|
||||||
build_core_error_body_for_client_format, build_cross_format_openai_chat_request_body,
|
build_core_error_body_for_client_format, build_cross_format_openai_chat_request_body,
|
||||||
build_cross_format_openai_chat_request_body_with_model_directives,
|
build_cross_format_openai_chat_request_body_with_model_directives,
|
||||||
|
|||||||
@@ -58,6 +58,10 @@ pub(crate) mod windsurf {
|
|||||||
pub(crate) use aether_provider_transport::windsurf::*;
|
pub(crate) use aether_provider_transport::windsurf::*;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub(crate) mod xai {
|
||||||
|
pub(crate) use aether_provider_transport::xai::*;
|
||||||
|
}
|
||||||
|
|
||||||
pub(crate) use aether_provider_transport::{
|
pub(crate) use aether_provider_transport::{
|
||||||
append_transport_diagnostics_to_value, apply_codex_fingerprint_convergence,
|
append_transport_diagnostics_to_value, apply_codex_fingerprint_convergence,
|
||||||
apply_codex_fingerprint_convergence_with_context, apply_local_auth_config_header_overrides,
|
apply_codex_fingerprint_convergence_with_context, apply_local_auth_config_header_overrides,
|
||||||
|
|||||||
@@ -53,6 +53,8 @@ const AI_ANY_ROUTE_PATTERNS: &[&str] = &[
|
|||||||
"/v1beta/operations/{*operation_path}",
|
"/v1beta/operations/{*operation_path}",
|
||||||
"/v1/videos",
|
"/v1/videos",
|
||||||
"/v1/videos/{*video_path}",
|
"/v1/videos/{*video_path}",
|
||||||
|
"/openai/v1/videos",
|
||||||
|
"/openai/v1/videos/{*video_path}",
|
||||||
"/upload/v1beta/files",
|
"/upload/v1beta/files",
|
||||||
"/v1beta/files",
|
"/v1beta/files",
|
||||||
"/v1beta/files/{*file_path}",
|
"/v1beta/files/{*file_path}",
|
||||||
|
|||||||
@@ -536,6 +536,9 @@ mod tests {
|
|||||||
|
|
||||||
fn sample_sparse_stored_task() -> StoredVideoTask {
|
fn sample_sparse_stored_task() -> StoredVideoTask {
|
||||||
let snapshot = LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
let snapshot = LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: None,
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: false,
|
||||||
local_task_id: "task-1".to_string(),
|
local_task_id: "task-1".to_string(),
|
||||||
upstream_task_id: "ext-1".to_string(),
|
upstream_task_id: "ext-1".to_string(),
|
||||||
created_at_unix_ms: 1,
|
created_at_unix_ms: 1,
|
||||||
|
|||||||
@@ -140,6 +140,8 @@ pub(crate) const RUST_FRONTDOOR_OWNED_ROUTE_PATTERNS: &[&str] = &[
|
|||||||
"/v1beta/models/{model}/operations/{id}",
|
"/v1beta/models/{model}/operations/{id}",
|
||||||
"/v1beta/operations",
|
"/v1beta/operations",
|
||||||
"/v1beta/operations/{id}",
|
"/v1beta/operations/{id}",
|
||||||
|
"/openai/v1/videos",
|
||||||
|
"/openai/v1/videos/{path...}",
|
||||||
"/v1/videos",
|
"/v1/videos",
|
||||||
"/v1/videos/{path...}",
|
"/v1/videos/{path...}",
|
||||||
"/upload/v1beta/files",
|
"/upload/v1beta/files",
|
||||||
|
|||||||
@@ -137,7 +137,11 @@ pub(super) fn classify_ai_public_route(
|
|||||||
.with_client_surface(detect_claude_client_surface(headers))
|
.with_client_surface(detect_claude_client_surface(headers))
|
||||||
.with_api_operation(ApiOperation::ClaudeMessagesCreate),
|
.with_api_operation(ApiOperation::ClaudeMessagesCreate),
|
||||||
)
|
)
|
||||||
} else if normalized_path.starts_with("/v1/videos") {
|
} else if normalized_path == "/v1/videos"
|
||||||
|
|| normalized_path.starts_with("/v1/videos/")
|
||||||
|
|| normalized_path == "/openai/v1/videos"
|
||||||
|
|| normalized_path.starts_with("/openai/v1/videos/")
|
||||||
|
{
|
||||||
Some(classified(
|
Some(classified(
|
||||||
"ai_public",
|
"ai_public",
|
||||||
"openai",
|
"openai",
|
||||||
|
|||||||
@@ -123,6 +123,15 @@ impl GatewayDataState {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
|
pub(crate) fn attach_video_task_repository_for_tests<T>(mut self, repository: Arc<T>) -> Self
|
||||||
|
where
|
||||||
|
T: VideoTaskRepository + 'static,
|
||||||
|
{
|
||||||
|
self.video_task_reader = Some(repository.clone());
|
||||||
|
self.video_task_writer = Some(repository);
|
||||||
|
self
|
||||||
|
}
|
||||||
|
|
||||||
pub(crate) fn with_video_task_repository_for_tests<T>(repository: Arc<T>) -> Self
|
pub(crate) fn with_video_task_repository_for_tests<T>(repository: Arc<T>) -> Self
|
||||||
where
|
where
|
||||||
T: VideoTaskRepository + 'static,
|
T: VideoTaskRepository + 'static,
|
||||||
|
|||||||
@@ -1464,6 +1464,22 @@ pub(crate) async fn maybe_execute_sync_via_local_video_decision(
|
|||||||
.await
|
.await
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn supports_local_video_get(
|
||||||
|
parts: &http::request::Parts,
|
||||||
|
decision: &GatewayControlDecision,
|
||||||
|
) -> bool {
|
||||||
|
parts.method == http::Method::GET
|
||||||
|
&& decision.route_kind.as_deref() == Some("video")
|
||||||
|
&& (crate::video_tasks::resolve_video_task_read_lookup_key(
|
||||||
|
decision.route_family.as_deref(),
|
||||||
|
parts.uri.path(),
|
||||||
|
)
|
||||||
|
.is_some()
|
||||||
|
|| (decision.route_family.as_deref() == Some("openai")
|
||||||
|
&& crate::video_tasks::extract_openai_task_id_from_content_path(parts.uri.path())
|
||||||
|
.is_some()))
|
||||||
|
}
|
||||||
|
|
||||||
pub(crate) fn maybe_execute_sync_request<'a>(
|
pub(crate) fn maybe_execute_sync_request<'a>(
|
||||||
state: &'a AppState,
|
state: &'a AppState,
|
||||||
parts: &'a http::request::Parts,
|
parts: &'a http::request::Parts,
|
||||||
@@ -1477,7 +1493,7 @@ pub(crate) fn maybe_execute_sync_request<'a>(
|
|||||||
};
|
};
|
||||||
#[cfg(not(test))]
|
#[cfg(not(test))]
|
||||||
{
|
{
|
||||||
if parts.method != http::Method::POST {
|
if parts.method != http::Method::POST && !supports_local_video_get(parts, decision) {
|
||||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||||
}
|
}
|
||||||
return maybe_execute_sync_local_path(state, parts, body_bytes, trace_id, decision)
|
return maybe_execute_sync_local_path(state, parts, body_bytes, trace_id, decision)
|
||||||
@@ -1490,6 +1506,7 @@ pub(crate) fn maybe_execute_sync_request<'a>(
|
|||||||
.unwrap_or_default()
|
.unwrap_or_default()
|
||||||
.is_empty()
|
.is_empty()
|
||||||
&& parts.method != http::Method::POST
|
&& parts.method != http::Method::POST
|
||||||
|
&& !supports_local_video_get(parts, decision)
|
||||||
{
|
{
|
||||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||||
}
|
}
|
||||||
@@ -1511,7 +1528,7 @@ pub(crate) fn maybe_execute_stream_request<'a>(
|
|||||||
};
|
};
|
||||||
#[cfg(not(test))]
|
#[cfg(not(test))]
|
||||||
{
|
{
|
||||||
if parts.method != http::Method::POST {
|
if parts.method != http::Method::POST && !supports_local_video_get(parts, decision) {
|
||||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||||
}
|
}
|
||||||
return maybe_execute_stream_local_path(state, parts, body_bytes, trace_id, decision)
|
return maybe_execute_stream_local_path(state, parts, body_bytes, trace_id, decision)
|
||||||
@@ -1524,6 +1541,7 @@ pub(crate) fn maybe_execute_stream_request<'a>(
|
|||||||
.unwrap_or_default()
|
.unwrap_or_default()
|
||||||
.is_empty()
|
.is_empty()
|
||||||
&& parts.method != http::Method::POST
|
&& parts.method != http::Method::POST
|
||||||
|
&& !supports_local_video_get(parts, decision)
|
||||||
{
|
{
|
||||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -32,6 +32,10 @@ fn request_has_execution_runtime_via_guard(headers: &HeaderMap) -> bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn frontdoor_self_loop_public_ai_path(path: &str) -> bool {
|
pub(crate) fn frontdoor_self_loop_public_ai_path(path: &str) -> bool {
|
||||||
|
let path = path
|
||||||
|
.strip_prefix("/openai")
|
||||||
|
.filter(|p| *p == "/v1/videos" || p.starts_with("/v1/videos/"))
|
||||||
|
.unwrap_or(path);
|
||||||
matches!(
|
matches!(
|
||||||
path,
|
path,
|
||||||
"/v1/messages"
|
"/v1/messages"
|
||||||
|
|||||||
@@ -74,7 +74,8 @@ fn validate_batch_access_token_import(
|
|||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
if !provider_type_supports_access_token_import(provider_type) {
|
if !provider_type_supports_access_token_import(provider_type) {
|
||||||
return Err(
|
return Err(
|
||||||
"Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok Provider".to_string(),
|
"Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok / xAI Provider"
|
||||||
|
.to_string(),
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
if provider_type.eq_ignore_ascii_case("claude_code") {
|
if provider_type.eq_ignore_ascii_case("claude_code") {
|
||||||
|
|||||||
@@ -214,7 +214,12 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
|||||||
} else {
|
} else {
|
||||||
let sso_from_cookie = grok_cookie_session_token(provider_type, raw_token);
|
let sso_from_cookie = grok_cookie_session_token(provider_type, raw_token);
|
||||||
let token_input = sso_from_cookie.as_deref().unwrap_or(raw_token);
|
let token_input = sso_from_cookie.as_deref().unwrap_or(raw_token);
|
||||||
let (refresh_token, access_token) = import_tokens_from_raw_token(token_input);
|
let (refresh_token, access_token) =
|
||||||
|
if provider_type.trim().eq_ignore_ascii_case("xai") {
|
||||||
|
(None, Some(token_input.to_string()))
|
||||||
|
} else {
|
||||||
|
import_tokens_from_raw_token(token_input)
|
||||||
|
};
|
||||||
let (refresh_token, access_token) = normalize_provider_import_tokens(
|
let (refresh_token, access_token) = normalize_provider_import_tokens(
|
||||||
provider_type,
|
provider_type,
|
||||||
refresh_token.as_deref(),
|
refresh_token.as_deref(),
|
||||||
@@ -262,6 +267,7 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
|||||||
let object = normalized_claude_object.as_ref().unwrap_or(object);
|
let object = normalized_claude_object.as_ref().unwrap_or(object);
|
||||||
let is_grok = provider_type.trim().eq_ignore_ascii_case("grok");
|
let is_grok = provider_type.trim().eq_ignore_ascii_case("grok");
|
||||||
let is_windsurf = provider_type.trim().eq_ignore_ascii_case("windsurf");
|
let is_windsurf = provider_type.trim().eq_ignore_ascii_case("windsurf");
|
||||||
|
let is_xai = provider_type.trim().eq_ignore_ascii_case("xai");
|
||||||
let is_codex_agent_identity = provider_type.trim().eq_ignore_ascii_case("codex")
|
let is_codex_agent_identity = provider_type.trim().eq_ignore_ascii_case("codex")
|
||||||
&& aether_provider_transport::is_codex_agent_identity_auth_config_value(item);
|
&& aether_provider_transport::is_codex_agent_identity_auth_config_value(item);
|
||||||
if is_codex_agent_identity {
|
if is_codex_agent_identity {
|
||||||
@@ -336,14 +342,6 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
|||||||
} else {
|
} else {
|
||||||
None
|
None
|
||||||
};
|
};
|
||||||
let (refresh_token, access_token) = normalize_provider_import_tokens(
|
|
||||||
provider_type,
|
|
||||||
refresh_token.as_deref(),
|
|
||||||
access_token
|
|
||||||
.as_deref()
|
|
||||||
.or(session_token.as_deref())
|
|
||||||
.or(header_bearer_token.as_deref()),
|
|
||||||
);
|
|
||||||
let windsurf_api_key = is_windsurf
|
let windsurf_api_key = is_windsurf
|
||||||
.then(|| {
|
.then(|| {
|
||||||
coerce_admin_provider_oauth_import_str(
|
coerce_admin_provider_oauth_import_str(
|
||||||
@@ -351,6 +349,22 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
|||||||
)
|
)
|
||||||
})
|
})
|
||||||
.flatten();
|
.flatten();
|
||||||
|
let xai_api_key = is_xai
|
||||||
|
.then(|| {
|
||||||
|
coerce_admin_provider_oauth_import_str(
|
||||||
|
object.get("api_key").or_else(|| object.get("apiKey")),
|
||||||
|
)
|
||||||
|
})
|
||||||
|
.flatten();
|
||||||
|
let (refresh_token, access_token) = normalize_provider_import_tokens(
|
||||||
|
provider_type,
|
||||||
|
refresh_token.as_deref(),
|
||||||
|
access_token
|
||||||
|
.as_deref()
|
||||||
|
.or(session_token.as_deref())
|
||||||
|
.or(header_bearer_token.as_deref())
|
||||||
|
.or(xai_api_key.as_deref()),
|
||||||
|
);
|
||||||
let windsurf_token = is_windsurf
|
let windsurf_token = is_windsurf
|
||||||
.then(|| {
|
.then(|| {
|
||||||
coerce_admin_provider_oauth_import_str(
|
coerce_admin_provider_oauth_import_str(
|
||||||
@@ -1577,4 +1591,23 @@ mod tests {
|
|||||||
assert!(entries[1].access_token.is_none());
|
assert!(entries[1].access_token.is_none());
|
||||||
assert!(entries[1].raw_credentials.is_none());
|
assert!(entries[1].raw_credentials.is_none());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn parses_xai_api_key_json_and_raw_lines_as_access_token() {
|
||||||
|
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||||
|
"xai",
|
||||||
|
r#"{"api_key":"xai-api-key","email":"a@x.ai"}
|
||||||
|
{"refresh_token":"xai-refresh"}
|
||||||
|
xai-raw-api-key"#,
|
||||||
|
);
|
||||||
|
|
||||||
|
assert_eq!(entries.len(), 3);
|
||||||
|
assert!(entries[0].refresh_token.is_none());
|
||||||
|
assert_eq!(entries[0].access_token.as_deref(), Some("xai-api-key"));
|
||||||
|
assert_eq!(entries[0].email.as_deref(), Some("[email protected]"));
|
||||||
|
assert_eq!(entries[1].refresh_token.as_deref(), Some("xai-refresh"));
|
||||||
|
assert!(entries[1].access_token.is_none());
|
||||||
|
assert!(entries[2].refresh_token.is_none());
|
||||||
|
assert_eq!(entries[2].access_token.as_deref(), Some("xai-raw-api-key"));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -186,10 +186,10 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
|
|||||||
));
|
));
|
||||||
};
|
};
|
||||||
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
|
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
|
||||||
if provider_type != "kiro" && provider_type != "windsurf" {
|
if provider_type != "kiro" && provider_type != "windsurf" && provider_type != "xai" {
|
||||||
return Ok(build_internal_control_error_response(
|
return Ok(build_internal_control_error_response(
|
||||||
http::StatusCode::BAD_REQUEST,
|
http::StatusCode::BAD_REQUEST,
|
||||||
"设备授权仅支持 Kiro / Windsurf provider",
|
"设备授权仅支持 Kiro / Windsurf / xAI provider",
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
let Some(principal) = request_context
|
let Some(principal) = request_context
|
||||||
@@ -219,6 +219,19 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
|
|||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
|
if provider_type == "xai" {
|
||||||
|
return super::xai::handle_admin_provider_oauth_xai_device_authorize(
|
||||||
|
state,
|
||||||
|
&provider_id,
|
||||||
|
&provider,
|
||||||
|
principal,
|
||||||
|
runtime_endpoint.as_ref(),
|
||||||
|
request_proxy,
|
||||||
|
payload.proxy_node_id.as_deref(),
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
}
|
||||||
|
|
||||||
if provider_type == "windsurf" {
|
if provider_type == "windsurf" {
|
||||||
let session_id = generate_provider_oauth_nonce();
|
let session_id = generate_provider_oauth_nonce();
|
||||||
let login_option = payload
|
let login_option = payload
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ mod authorize;
|
|||||||
mod lease;
|
mod lease;
|
||||||
mod poll;
|
mod poll;
|
||||||
mod session;
|
mod session;
|
||||||
|
mod xai;
|
||||||
|
|
||||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||||
use crate::GatewayError;
|
use crate::GatewayError;
|
||||||
|
|||||||
@@ -479,6 +479,18 @@ pub(super) async fn handle_admin_provider_oauth_device_poll(
|
|||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
|
if provider_type == "xai" {
|
||||||
|
return super::xai::handle_admin_provider_oauth_xai_device_poll(
|
||||||
|
state,
|
||||||
|
&provider,
|
||||||
|
&endpoints,
|
||||||
|
request_proxy,
|
||||||
|
session_id,
|
||||||
|
session,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
}
|
||||||
|
|
||||||
if provider_type == "windsurf" {
|
if provider_type == "windsurf" {
|
||||||
return handle_admin_provider_oauth_windsurf_browser_device_poll(
|
return handle_admin_provider_oauth_windsurf_browser_device_poll(
|
||||||
state,
|
state,
|
||||||
|
|||||||
@@ -0,0 +1,368 @@
|
|||||||
|
use super::session::attach_admin_provider_oauth_device_poll_terminal_response;
|
||||||
|
use crate::control::GatewayAdminPrincipalContext;
|
||||||
|
use crate::handlers::admin::provider::oauth::dispatch::helpers::admin_provider_oauth_key_name_from_auth_config;
|
||||||
|
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||||
|
use crate::handlers::admin::provider::oauth::provisioning::{
|
||||||
|
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
|
||||||
|
};
|
||||||
|
use crate::handlers::admin::provider::oauth::runtime::spawn_provider_oauth_account_state_refresh_after_update;
|
||||||
|
use crate::handlers::admin::provider::oauth::state::{
|
||||||
|
current_unix_secs, generate_provider_oauth_nonce,
|
||||||
|
};
|
||||||
|
use crate::handlers::admin::request::AdminAppState;
|
||||||
|
use crate::GatewayError;
|
||||||
|
use aether_contracts::ProxySnapshot;
|
||||||
|
use aether_data::repository::provider_oauth::{
|
||||||
|
StoredAdminProviderOAuthDeviceSession, KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS,
|
||||||
|
};
|
||||||
|
use aether_data_contracts::repository::provider_catalog::{
|
||||||
|
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
|
||||||
|
};
|
||||||
|
use aether_oauth::core::OAuthError;
|
||||||
|
use aether_oauth::provider::providers::{
|
||||||
|
XaiDevicePollOutcome, XaiProviderOAuthAdapter, XAI_CLIENT_ID, XAI_DEVICE_CODE_URL,
|
||||||
|
XAI_TOKEN_URL,
|
||||||
|
};
|
||||||
|
use aether_oauth::provider::ProviderOAuthTransportContext;
|
||||||
|
use axum::{
|
||||||
|
body::Body,
|
||||||
|
http,
|
||||||
|
response::{IntoResponse, Response},
|
||||||
|
Json,
|
||||||
|
};
|
||||||
|
use serde_json::{json, Value};
|
||||||
|
|
||||||
|
pub(super) async fn handle_admin_provider_oauth_xai_device_authorize(
|
||||||
|
state: &AdminAppState<'_>,
|
||||||
|
provider_id: &str,
|
||||||
|
provider: &StoredProviderCatalogProvider,
|
||||||
|
principal: &GatewayAdminPrincipalContext,
|
||||||
|
runtime_endpoint: Option<&StoredProviderCatalogEndpoint>,
|
||||||
|
request_proxy: Option<ProxySnapshot>,
|
||||||
|
proxy_node_id: Option<&str>,
|
||||||
|
) -> Result<Response<Body>, GatewayError> {
|
||||||
|
let device_url = state.provider_oauth_token_url("xai_device", XAI_DEVICE_CODE_URL);
|
||||||
|
let token_url = state.provider_oauth_token_url("xai", XAI_TOKEN_URL);
|
||||||
|
let adapter =
|
||||||
|
XaiProviderOAuthAdapter::default().with_endpoint_overrides(&device_url, &token_url);
|
||||||
|
let ctx = ProviderOAuthTransportContext {
|
||||||
|
provider_id: provider_id.to_string(),
|
||||||
|
provider_type: provider.provider_type.clone(),
|
||||||
|
endpoint_id: runtime_endpoint.map(|endpoint| endpoint.id.clone()),
|
||||||
|
key_id: None,
|
||||||
|
auth_type: Some("oauth".to_string()),
|
||||||
|
decrypted_api_key: None,
|
||||||
|
decrypted_auth_config: None,
|
||||||
|
provider_config: provider.config.clone(),
|
||||||
|
endpoint_config: runtime_endpoint.and_then(|endpoint| endpoint.config.clone()),
|
||||||
|
key_config: None,
|
||||||
|
network: aether_oauth::network::OAuthNetworkContext::provider_operation(
|
||||||
|
request_proxy.clone(),
|
||||||
|
),
|
||||||
|
};
|
||||||
|
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
|
||||||
|
let authorization = match adapter.start_device_flow(&executor, &ctx).await {
|
||||||
|
Ok(authorization) => authorization,
|
||||||
|
Err(error) => {
|
||||||
|
return Ok(build_internal_control_error_response(
|
||||||
|
http::StatusCode::BAD_REQUEST,
|
||||||
|
sanitize_xai_oauth_error(&error),
|
||||||
|
));
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
let now_unix_secs = current_unix_secs();
|
||||||
|
let session_id = generate_provider_oauth_nonce();
|
||||||
|
let session = StoredAdminProviderOAuthDeviceSession {
|
||||||
|
session_id: session_id.clone(),
|
||||||
|
provider_id: provider_id.to_string(),
|
||||||
|
initiated_by_user_id: principal.user_id.clone(),
|
||||||
|
initiated_by_session_id: principal.session_id.clone(),
|
||||||
|
initiated_by_management_token_id: principal.management_token_id.clone(),
|
||||||
|
region: String::new(),
|
||||||
|
client_id: XAI_CLIENT_ID.to_string(),
|
||||||
|
client_secret: String::new(),
|
||||||
|
device_code: authorization.device_code.clone(),
|
||||||
|
auth_type: Some("device".to_string()),
|
||||||
|
social_provider: None,
|
||||||
|
code_verifier: None,
|
||||||
|
redirect_uri: Some(token_url),
|
||||||
|
machine_id: None,
|
||||||
|
interval: authorization.interval,
|
||||||
|
expires_at_unix_secs: now_unix_secs.saturating_add(authorization.expires_in),
|
||||||
|
status: "pending".to_string(),
|
||||||
|
proxy_node_id: proxy_node_id
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.map(ToOwned::to_owned),
|
||||||
|
created_at_unix_ms: now_unix_secs,
|
||||||
|
key_id: None,
|
||||||
|
email: None,
|
||||||
|
replaced: false,
|
||||||
|
error_msg: None,
|
||||||
|
};
|
||||||
|
if let Err(response) = state
|
||||||
|
.save_provider_oauth_device_session(
|
||||||
|
&session_id,
|
||||||
|
&session,
|
||||||
|
authorization
|
||||||
|
.expires_in
|
||||||
|
.saturating_add(KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS),
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
{
|
||||||
|
return Ok(response);
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(Json(json!({
|
||||||
|
"session_id": session_id,
|
||||||
|
"user_code": authorization.user_code,
|
||||||
|
"verification_uri": authorization.verification_uri,
|
||||||
|
"verification_uri_complete": authorization.verification_uri_complete,
|
||||||
|
"expires_in": authorization.expires_in,
|
||||||
|
"interval": authorization.interval,
|
||||||
|
"auth_type": "device",
|
||||||
|
}))
|
||||||
|
.into_response())
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) async fn handle_admin_provider_oauth_xai_device_poll(
|
||||||
|
state: &AdminAppState<'_>,
|
||||||
|
provider: &StoredProviderCatalogProvider,
|
||||||
|
endpoints: &[StoredProviderCatalogEndpoint],
|
||||||
|
request_proxy: Option<ProxySnapshot>,
|
||||||
|
session_id: &str,
|
||||||
|
mut session: StoredAdminProviderOAuthDeviceSession,
|
||||||
|
) -> Result<Response<Body>, GatewayError> {
|
||||||
|
let token_url = session
|
||||||
|
.redirect_uri
|
||||||
|
.as_deref()
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.map(ToOwned::to_owned)
|
||||||
|
.unwrap_or_else(|| state.provider_oauth_token_url("xai", XAI_TOKEN_URL));
|
||||||
|
let adapter =
|
||||||
|
XaiProviderOAuthAdapter::default().with_endpoint_overrides(XAI_DEVICE_CODE_URL, token_url);
|
||||||
|
let ctx = ProviderOAuthTransportContext {
|
||||||
|
provider_id: provider.id.clone(),
|
||||||
|
provider_type: provider.provider_type.clone(),
|
||||||
|
endpoint_id: None,
|
||||||
|
key_id: None,
|
||||||
|
auth_type: Some("oauth".to_string()),
|
||||||
|
decrypted_api_key: None,
|
||||||
|
decrypted_auth_config: None,
|
||||||
|
provider_config: provider.config.clone(),
|
||||||
|
endpoint_config: None,
|
||||||
|
key_config: None,
|
||||||
|
network: aether_oauth::network::OAuthNetworkContext::provider_operation(
|
||||||
|
request_proxy.clone(),
|
||||||
|
),
|
||||||
|
};
|
||||||
|
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
|
||||||
|
let outcome = match adapter
|
||||||
|
.poll_device_token(&executor, &ctx, &session.device_code)
|
||||||
|
.await
|
||||||
|
{
|
||||||
|
Ok(outcome) => outcome,
|
||||||
|
Err(error) => {
|
||||||
|
return Ok(xai_device_poll_terminal_from_error(
|
||||||
|
state,
|
||||||
|
session_id,
|
||||||
|
&mut session,
|
||||||
|
&error,
|
||||||
|
)
|
||||||
|
.await);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
match outcome {
|
||||||
|
XaiDevicePollOutcome::Pending => {
|
||||||
|
Ok(Json(json!({"status": "pending", "replaced": false})).into_response())
|
||||||
|
}
|
||||||
|
XaiDevicePollOutcome::SlowDown => {
|
||||||
|
Ok(Json(json!({"status": "slow_down", "replaced": false})).into_response())
|
||||||
|
}
|
||||||
|
XaiDevicePollOutcome::Authorized(result) => {
|
||||||
|
persist_xai_device_authorization(
|
||||||
|
state,
|
||||||
|
provider,
|
||||||
|
endpoints,
|
||||||
|
request_proxy,
|
||||||
|
session_id,
|
||||||
|
session,
|
||||||
|
*result,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn persist_xai_device_authorization(
|
||||||
|
state: &AdminAppState<'_>,
|
||||||
|
provider: &StoredProviderCatalogProvider,
|
||||||
|
endpoints: &[StoredProviderCatalogEndpoint],
|
||||||
|
request_proxy: Option<ProxySnapshot>,
|
||||||
|
session_id: &str,
|
||||||
|
mut session: StoredAdminProviderOAuthDeviceSession,
|
||||||
|
result: aether_oauth::provider::ProviderOAuthTokenSet,
|
||||||
|
) -> Result<Response<Body>, GatewayError> {
|
||||||
|
let access_token = result.token_set.access_token.trim().to_string();
|
||||||
|
if access_token.is_empty() {
|
||||||
|
return Ok(Json(json!({
|
||||||
|
"status": "error",
|
||||||
|
"error": "xAI token 响应缺少 access_token",
|
||||||
|
"replaced": false,
|
||||||
|
}))
|
||||||
|
.into_response());
|
||||||
|
}
|
||||||
|
let mut auth_config = result.auth_config.as_object().cloned().unwrap_or_default();
|
||||||
|
auth_config.insert("provider_type".to_string(), json!("xai"));
|
||||||
|
auth_config.insert("auth_method".to_string(), json!("oauth"));
|
||||||
|
auth_config.insert("using_api".to_string(), json!(false));
|
||||||
|
|
||||||
|
let duplicate = match state
|
||||||
|
.find_duplicate_provider_oauth_key(&provider.id, &auth_config, None)
|
||||||
|
.await
|
||||||
|
{
|
||||||
|
Ok(duplicate) => duplicate,
|
||||||
|
Err(detail) => {
|
||||||
|
return Ok(Json(json!({
|
||||||
|
"status": "error",
|
||||||
|
"error": detail,
|
||||||
|
"replaced": false,
|
||||||
|
}))
|
||||||
|
.into_response());
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
let api_formats = provider_oauth_active_api_formats(endpoints);
|
||||||
|
let key_proxy = provider_oauth_key_proxy_value(session.proxy_node_id.as_deref());
|
||||||
|
let expires_at = result.token_set.expires_at_unix_secs;
|
||||||
|
let email = auth_config
|
||||||
|
.get("email")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.map(ToOwned::to_owned);
|
||||||
|
let mut replaced = false;
|
||||||
|
let persisted_key = if let Some(existing_key) = duplicate {
|
||||||
|
replaced = true;
|
||||||
|
match state
|
||||||
|
.update_existing_provider_oauth_catalog_key(
|
||||||
|
&existing_key,
|
||||||
|
&provider.provider_type,
|
||||||
|
&access_token,
|
||||||
|
&auth_config,
|
||||||
|
&api_formats,
|
||||||
|
key_proxy.clone(),
|
||||||
|
expires_at,
|
||||||
|
)
|
||||||
|
.await?
|
||||||
|
{
|
||||||
|
Some(key) => key,
|
||||||
|
None => {
|
||||||
|
return Ok(build_internal_control_error_response(
|
||||||
|
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||||
|
"provider oauth write unavailable",
|
||||||
|
));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
let key_name = admin_provider_oauth_key_name_from_auth_config(
|
||||||
|
&provider.provider_type,
|
||||||
|
&auth_config,
|
||||||
|
None,
|
||||||
|
);
|
||||||
|
match state
|
||||||
|
.create_provider_oauth_catalog_key(
|
||||||
|
&provider.id,
|
||||||
|
&provider.provider_type,
|
||||||
|
&key_name,
|
||||||
|
&access_token,
|
||||||
|
&auth_config,
|
||||||
|
&api_formats,
|
||||||
|
key_proxy,
|
||||||
|
expires_at,
|
||||||
|
)
|
||||||
|
.await?
|
||||||
|
{
|
||||||
|
Some(key) => key,
|
||||||
|
None => {
|
||||||
|
return Ok(build_internal_control_error_response(
|
||||||
|
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||||
|
"provider oauth write unavailable",
|
||||||
|
));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
spawn_provider_oauth_account_state_refresh_after_update(
|
||||||
|
state.cloned_app(),
|
||||||
|
provider.clone(),
|
||||||
|
persisted_key.id.clone(),
|
||||||
|
request_proxy.clone(),
|
||||||
|
);
|
||||||
|
|
||||||
|
session.status = "authorized".to_string();
|
||||||
|
session.key_id = Some(persisted_key.id.clone());
|
||||||
|
session.email = email.clone();
|
||||||
|
session.replaced = replaced;
|
||||||
|
session.error_msg = None;
|
||||||
|
let _ = state
|
||||||
|
.save_provider_oauth_device_session(session_id, &session, 60)
|
||||||
|
.await;
|
||||||
|
|
||||||
|
Ok(attach_admin_provider_oauth_device_poll_terminal_response(
|
||||||
|
session_id,
|
||||||
|
"authorized",
|
||||||
|
Json(json!({
|
||||||
|
"status": "authorized",
|
||||||
|
"key_id": persisted_key.id,
|
||||||
|
"email": email,
|
||||||
|
"replaced": replaced,
|
||||||
|
}))
|
||||||
|
.into_response(),
|
||||||
|
))
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn xai_device_poll_terminal_from_error(
|
||||||
|
state: &AdminAppState<'_>,
|
||||||
|
session_id: &str,
|
||||||
|
session: &mut StoredAdminProviderOAuthDeviceSession,
|
||||||
|
error: &OAuthError,
|
||||||
|
) -> Response<Body> {
|
||||||
|
let (status, message) = match error {
|
||||||
|
OAuthError::InvalidRequest(detail) if detail.to_ascii_lowercase().contains("expired") => {
|
||||||
|
("expired", "设备码已过期".to_string())
|
||||||
|
}
|
||||||
|
OAuthError::InvalidRequest(detail) if detail.to_ascii_lowercase().contains("denied") => {
|
||||||
|
("error", "用户拒绝授权".to_string())
|
||||||
|
}
|
||||||
|
_ => ("error", sanitize_xai_oauth_error(error)),
|
||||||
|
};
|
||||||
|
session.status = status.to_string();
|
||||||
|
session.error_msg = Some(message.clone());
|
||||||
|
let _ = state
|
||||||
|
.save_provider_oauth_device_session(session_id, session, 30)
|
||||||
|
.await;
|
||||||
|
attach_admin_provider_oauth_device_poll_terminal_response(
|
||||||
|
session_id,
|
||||||
|
status,
|
||||||
|
Json(json!({
|
||||||
|
"status": status,
|
||||||
|
"error": message,
|
||||||
|
"replaced": false,
|
||||||
|
}))
|
||||||
|
.into_response(),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn sanitize_xai_oauth_error(error: &OAuthError) -> String {
|
||||||
|
match error {
|
||||||
|
OAuthError::InvalidRequest(_) => "xAI 设备授权失败: 请求参数无效".to_string(),
|
||||||
|
OAuthError::HttpStatus { status_code, .. } => {
|
||||||
|
format!("xAI 设备授权失败: HTTP {status_code}")
|
||||||
|
}
|
||||||
|
_ => "xAI 设备授权失败".to_string(),
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -715,7 +715,7 @@ async fn resolve_admin_provider_oauth_single_import_tokens(
|
|||||||
if !provider_type_supports_access_token_import(provider_type) {
|
if !provider_type_supports_access_token_import(provider_type) {
|
||||||
return Err(build_internal_control_error_response(
|
return Err(build_internal_control_error_response(
|
||||||
http::StatusCode::BAD_REQUEST,
|
http::StatusCode::BAD_REQUEST,
|
||||||
"Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok Provider",
|
"Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok / xAI Provider",
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -867,7 +867,7 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
|||||||
flatten_claude_code_credentials_payload(&mut raw_payload);
|
flatten_claude_code_credentials_payload(&mut raw_payload);
|
||||||
}
|
}
|
||||||
let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken");
|
let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken");
|
||||||
let access_token_input = import_payload_string_any(
|
let mut access_token_input = import_payload_string_any(
|
||||||
&raw_payload,
|
&raw_payload,
|
||||||
&[
|
&[
|
||||||
"access_token",
|
"access_token",
|
||||||
@@ -879,6 +879,9 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
|||||||
],
|
],
|
||||||
)
|
)
|
||||||
.or_else(|| provider_oauth_import_authorization_bearer_token_from_object(&raw_payload));
|
.or_else(|| provider_oauth_import_authorization_bearer_token_from_object(&raw_payload));
|
||||||
|
if provider_type == "xai" && access_token_input.is_none() {
|
||||||
|
access_token_input = import_payload_string(&raw_payload, "api_key", "apiKey");
|
||||||
|
}
|
||||||
let imported_expires_at =
|
let imported_expires_at =
|
||||||
import_payload_u64_any(&raw_payload, &["expires_at", "expiresAt", "expired"]);
|
import_payload_u64_any(&raw_payload, &["expires_at", "expiresAt", "expired"]);
|
||||||
let (refresh_token_input, access_token_input) = normalize_provider_import_tokens(
|
let (refresh_token_input, access_token_input) = normalize_provider_import_tokens(
|
||||||
@@ -901,7 +904,11 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
|||||||
if !create_agent_identity && refresh_token_input.is_none() && access_token_input.is_none() {
|
if !create_agent_identity && refresh_token_input.is_none() && access_token_input.is_none() {
|
||||||
return Ok(build_internal_control_error_response(
|
return Ok(build_internal_control_error_response(
|
||||||
http::StatusCode::BAD_REQUEST,
|
http::StatusCode::BAD_REQUEST,
|
||||||
"Refresh Token、Access Token 或 sso_token 不能为空",
|
if provider_type == "xai" {
|
||||||
|
"Refresh Token、Access Token 或 api_key 不能为空"
|
||||||
|
} else {
|
||||||
|
"Refresh Token、Access Token 或 sso_token 不能为空"
|
||||||
|
},
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
|
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
|
||||||
|
|||||||
@@ -70,6 +70,12 @@ pub(super) async fn handle_admin_provider_oauth_start_key(
|
|||||||
"Windsurf 请使用浏览器登录或导入凭据。",
|
"Windsurf 请使用浏览器登录或导入凭据。",
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
if provider_type == "xai" {
|
||||||
|
return Ok(build_internal_control_error_response(
|
||||||
|
http::StatusCode::BAD_REQUEST,
|
||||||
|
"xAI 请使用设备授权或导入凭据。",
|
||||||
|
));
|
||||||
|
}
|
||||||
let Some(template) = admin_provider_oauth_template(&provider_type) else {
|
let Some(template) = admin_provider_oauth_template(&provider_type) else {
|
||||||
return Ok(build_internal_control_error_response(
|
return Ok(build_internal_control_error_response(
|
||||||
http::StatusCode::BAD_REQUEST,
|
http::StatusCode::BAD_REQUEST,
|
||||||
@@ -167,6 +173,12 @@ pub(super) async fn handle_admin_provider_oauth_start_provider(
|
|||||||
"Windsurf 请使用浏览器登录或导入凭据。",
|
"Windsurf 请使用浏览器登录或导入凭据。",
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
if provider_type == "xai" {
|
||||||
|
return Ok(build_internal_control_error_response(
|
||||||
|
http::StatusCode::BAD_REQUEST,
|
||||||
|
"xAI 请使用设备授权或导入凭据。",
|
||||||
|
));
|
||||||
|
}
|
||||||
let Some(template) = admin_provider_oauth_template(&provider_type) else {
|
let Some(template) = admin_provider_oauth_template(&provider_type) else {
|
||||||
return Ok(build_internal_control_error_response(
|
return Ok(build_internal_control_error_response(
|
||||||
http::StatusCode::BAD_REQUEST,
|
http::StatusCode::BAD_REQUEST,
|
||||||
|
|||||||
@@ -121,6 +121,9 @@ pub(super) fn normalize_provider_import_tokens(
|
|||||||
if provider_type == "grok" {
|
if provider_type == "grok" {
|
||||||
return (None, access_token.or(refresh_token));
|
return (None, access_token.or(refresh_token));
|
||||||
}
|
}
|
||||||
|
if provider_type == "xai" {
|
||||||
|
return (refresh_token, access_token);
|
||||||
|
}
|
||||||
if provider_type == "claude_code" {
|
if provider_type == "claude_code" {
|
||||||
if access_token.is_none() && refresh_token.as_deref().is_some_and(is_claude_access_token) {
|
if access_token.is_none() && refresh_token.as_deref().is_some_and(is_claude_access_token) {
|
||||||
return (None, refresh_token);
|
return (None, refresh_token);
|
||||||
@@ -237,7 +240,7 @@ pub(super) fn provider_oauth_import_authorization_bearer_token_from_object(
|
|||||||
pub(super) fn provider_type_supports_access_token_import(provider_type: &str) -> bool {
|
pub(super) fn provider_type_supports_access_token_import(provider_type: &str) -> bool {
|
||||||
matches!(
|
matches!(
|
||||||
provider_type.trim().to_ascii_lowercase().as_str(),
|
provider_type.trim().to_ascii_lowercase().as_str(),
|
||||||
"claude_code" | "codex" | "chatgpt_web" | "grok"
|
"claude_code" | "codex" | "chatgpt_web" | "grok" | "xai"
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -331,6 +334,15 @@ pub(super) fn build_provider_access_token_import_auth_config(
|
|||||||
auth_config.insert("sso_token".to_string(), json!(access_token));
|
auth_config.insert("sso_token".to_string(), json!(access_token));
|
||||||
auth_config.insert("auth_method".to_string(), json!("sso_token"));
|
auth_config.insert("auth_method".to_string(), json!("sso_token"));
|
||||||
}
|
}
|
||||||
|
if provider_type.trim().eq_ignore_ascii_case("xai") {
|
||||||
|
if refresh_token.is_some() {
|
||||||
|
auth_config.insert("auth_method".to_string(), json!("oauth"));
|
||||||
|
auth_config.insert("using_api".to_string(), json!(false));
|
||||||
|
} else {
|
||||||
|
auth_config.insert("auth_method".to_string(), json!("api_key"));
|
||||||
|
auth_config.insert("using_api".to_string(), json!(true));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
auth_config.insert(
|
auth_config.insert(
|
||||||
"access_token_import_temporary".to_string(),
|
"access_token_import_temporary".to_string(),
|
||||||
@@ -532,6 +544,41 @@ mod tests {
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn normalize_xai_import_keeps_refresh_token_separate_from_api_key() {
|
||||||
|
let (refresh_token, access_token) =
|
||||||
|
normalize_provider_import_tokens("xai", Some("xai-refresh-token"), None);
|
||||||
|
assert_eq!(refresh_token.as_deref(), Some("xai-refresh-token"));
|
||||||
|
assert!(access_token.is_none());
|
||||||
|
|
||||||
|
let (refresh_token, access_token) =
|
||||||
|
normalize_provider_import_tokens("xai", None, Some("xai-api-key"));
|
||||||
|
assert!(refresh_token.is_none());
|
||||||
|
assert_eq!(access_token.as_deref(), Some("xai-api-key"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn builds_xai_auth_config_from_api_key_and_oauth_tokens() {
|
||||||
|
let (api_key_config, _) =
|
||||||
|
build_provider_access_token_import_auth_config("xai", "xai-api-key", None, None, None);
|
||||||
|
assert_eq!(api_key_config.get("auth_method"), Some(&json!("api_key")));
|
||||||
|
assert_eq!(api_key_config.get("using_api"), Some(&json!(true)));
|
||||||
|
|
||||||
|
let (oauth_config, _) = build_provider_access_token_import_auth_config(
|
||||||
|
"xai",
|
||||||
|
"xai-access-token",
|
||||||
|
Some("xai-refresh-token"),
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
);
|
||||||
|
assert_eq!(oauth_config.get("auth_method"), Some(&json!("oauth")));
|
||||||
|
assert_eq!(oauth_config.get("using_api"), Some(&json!(false)));
|
||||||
|
assert_eq!(
|
||||||
|
oauth_config.get("refresh_token"),
|
||||||
|
Some(&json!("xai-refresh-token"))
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn flattens_only_claude_ai_oauth_credentials_and_converts_expiry_ms() {
|
fn flattens_only_claude_ai_oauth_credentials_and_converts_expiry_ms() {
|
||||||
let mut payload = json!({
|
let mut payload = json!({
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ use super::gemini_cli::refresh_gemini_cli_provider_quota_locally;
|
|||||||
use super::grok::refresh_grok_provider_quota_locally;
|
use super::grok::refresh_grok_provider_quota_locally;
|
||||||
use super::kiro::refresh_kiro_provider_quota_locally;
|
use super::kiro::refresh_kiro_provider_quota_locally;
|
||||||
use super::windsurf::refresh_windsurf_provider_quota_locally;
|
use super::windsurf::refresh_windsurf_provider_quota_locally;
|
||||||
|
use super::xai::refresh_xai_provider_quota_locally;
|
||||||
use crate::handlers::admin::request::AdminAppState;
|
use crate::handlers::admin::request::AdminAppState;
|
||||||
use crate::GatewayError;
|
use crate::GatewayError;
|
||||||
use aether_contracts::ProxySnapshot;
|
use aether_contracts::ProxySnapshot;
|
||||||
@@ -43,6 +44,7 @@ const PROVIDER_QUOTA_REFRESH_HANDLERS: &[(&str, ProviderQuotaRefreshHandler)] =
|
|||||||
("grok", refresh_grok_provider_quota_locally_boxed),
|
("grok", refresh_grok_provider_quota_locally_boxed),
|
||||||
("kiro", refresh_kiro_provider_quota_locally_boxed),
|
("kiro", refresh_kiro_provider_quota_locally_boxed),
|
||||||
("windsurf", refresh_windsurf_provider_quota_locally_boxed),
|
("windsurf", refresh_windsurf_provider_quota_locally_boxed),
|
||||||
|
("xai", refresh_xai_provider_quota_locally_boxed),
|
||||||
];
|
];
|
||||||
|
|
||||||
pub(crate) async fn refresh_provider_pool_quota_locally(
|
pub(crate) async fn refresh_provider_pool_quota_locally(
|
||||||
@@ -174,3 +176,19 @@ fn refresh_windsurf_provider_quota_locally_boxed<'a>(
|
|||||||
proxy_override,
|
proxy_override,
|
||||||
))
|
))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn refresh_xai_provider_quota_locally_boxed<'a>(
|
||||||
|
state: &'a AdminAppState<'a>,
|
||||||
|
provider: &'a StoredProviderCatalogProvider,
|
||||||
|
endpoint: &'a StoredProviderCatalogEndpoint,
|
||||||
|
keys: Vec<StoredProviderCatalogKey>,
|
||||||
|
proxy_override: Option<ProxySnapshot>,
|
||||||
|
) -> ProviderQuotaRefreshFuture<'a> {
|
||||||
|
Box::pin(refresh_xai_provider_quota_locally(
|
||||||
|
state,
|
||||||
|
provider,
|
||||||
|
endpoint,
|
||||||
|
keys,
|
||||||
|
proxy_override,
|
||||||
|
))
|
||||||
|
}
|
||||||
|
|||||||
@@ -7,3 +7,4 @@ pub(crate) mod grok;
|
|||||||
pub(crate) mod kiro;
|
pub(crate) mod kiro;
|
||||||
pub(crate) mod shared;
|
pub(crate) mod shared;
|
||||||
pub(crate) mod windsurf;
|
pub(crate) mod windsurf;
|
||||||
|
pub(crate) mod xai;
|
||||||
|
|||||||
@@ -1715,6 +1715,7 @@ fn provider_quota_url_has_allowed_origin(provider_name: &str, value: &str) -> bo
|
|||||||
"gemini_cli" => host == "cloudcode-pa.googleapis.com",
|
"gemini_cli" => host == "cloudcode-pa.googleapis.com",
|
||||||
"chatgpt_web" | "codex" => host == "chatgpt.com",
|
"chatgpt_web" | "codex" => host == "chatgpt.com",
|
||||||
"grok" => host == "grok.com",
|
"grok" => host == "grok.com",
|
||||||
|
"xai" => host == "cli-chat-proxy.grok.com",
|
||||||
"windsurf" => host == "server.codeium.com",
|
"windsurf" => host == "server.codeium.com",
|
||||||
"kiro" => kiro_quota_host_is_allowed(host),
|
"kiro" => kiro_quota_host_is_allowed(host),
|
||||||
_ => false,
|
_ => false,
|
||||||
@@ -1814,6 +1815,14 @@ mod tests {
|
|||||||
),
|
),
|
||||||
("codex", "https://chatgpt.com/backend-api/wham/usage"),
|
("codex", "https://chatgpt.com/backend-api/wham/usage"),
|
||||||
("grok", "https://grok.com/rest/rate-limits"),
|
("grok", "https://grok.com/rest/rate-limits"),
|
||||||
|
(
|
||||||
|
"xai",
|
||||||
|
"https://cli-chat-proxy.grok.com/v1/billing?format=credits",
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"xai",
|
||||||
|
"https://cli-chat-proxy.grok.com/v1/user",
|
||||||
|
),
|
||||||
(
|
(
|
||||||
"windsurf",
|
"windsurf",
|
||||||
"https://server.codeium.com/exa.seat_management_pb.SeatManagementService/GetUserStatus",
|
"https://server.codeium.com/exa.seat_management_pb.SeatManagementService/GetUserStatus",
|
||||||
@@ -1847,6 +1856,11 @@ mod tests {
|
|||||||
"https://chatgpt.com.attacker.test/backend-api/wham/usage",
|
"https://chatgpt.com.attacker.test/backend-api/wham/usage",
|
||||||
),
|
),
|
||||||
("grok", "https://grok.com.attacker.test/rest/rate-limits"),
|
("grok", "https://grok.com.attacker.test/rest/rate-limits"),
|
||||||
|
(
|
||||||
|
"xai",
|
||||||
|
"https://cli-chat-proxy.grok.com.attacker.test/v1/billing",
|
||||||
|
),
|
||||||
|
("xai", "https://api.x.ai/v1/billing?format=credits"),
|
||||||
("windsurf", "https://server.codeium.com.attacker.test/quota"),
|
("windsurf", "https://server.codeium.com.attacker.test/quota"),
|
||||||
(
|
(
|
||||||
"gemini_cli",
|
"gemini_cli",
|
||||||
|
|||||||
@@ -0,0 +1,313 @@
|
|||||||
|
use super::shared::{
|
||||||
|
build_provider_quota_execution_plan, build_quota_snapshot_payload,
|
||||||
|
default_provider_quota_execution_timeouts, execute_provider_quota_plan,
|
||||||
|
extract_execution_error_message, oauth_refresh_auto_removed_result,
|
||||||
|
persist_provider_quota_refresh_state, quota_key_auto_removed,
|
||||||
|
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
|
||||||
|
};
|
||||||
|
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||||
|
use crate::GatewayError;
|
||||||
|
use aether_admin::provider::quota::parse_xai_billing_response;
|
||||||
|
use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json;
|
||||||
|
use aether_contracts::ProxySnapshot;
|
||||||
|
use aether_data_contracts::repository::provider_catalog::{
|
||||||
|
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||||
|
};
|
||||||
|
use aether_provider_pool::{build_xai_pool_billing_request, build_xai_pool_user_request};
|
||||||
|
use aether_provider_transport::xai::{
|
||||||
|
extract_xai_user_id_from_auth_config, extract_xai_user_id_from_value, xai_auth_uses_api,
|
||||||
|
};
|
||||||
|
use serde_json::{json, Value};
|
||||||
|
use std::time::{SystemTime, UNIX_EPOCH};
|
||||||
|
|
||||||
|
async fn execute_xai_quota_plan(
|
||||||
|
state: &AdminAppState<'_>,
|
||||||
|
transport: &AdminGatewayProviderTransportSnapshot,
|
||||||
|
spec: aether_provider_pool::ProviderPoolQuotaRequestSpec,
|
||||||
|
proxy_override: Option<&ProxySnapshot>,
|
||||||
|
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
|
||||||
|
let proxy = match proxy_override {
|
||||||
|
Some(proxy) => Some(proxy.clone()),
|
||||||
|
None => {
|
||||||
|
state
|
||||||
|
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
|
||||||
|
.await
|
||||||
|
}
|
||||||
|
};
|
||||||
|
let timeouts = state
|
||||||
|
.resolve_transport_execution_timeouts(transport)
|
||||||
|
.or(Some(default_provider_quota_execution_timeouts(
|
||||||
|
proxy.as_ref(),
|
||||||
|
)));
|
||||||
|
let plan = build_provider_quota_execution_plan(
|
||||||
|
transport,
|
||||||
|
spec,
|
||||||
|
proxy,
|
||||||
|
state.resolve_transport_profile(transport),
|
||||||
|
timeouts,
|
||||||
|
);
|
||||||
|
|
||||||
|
execute_provider_quota_plan(state, transport, plan, "xai").await
|
||||||
|
}
|
||||||
|
|
||||||
|
fn xai_authorization_from_header(authorization: &(String, String)) -> (String, String) {
|
||||||
|
authorization.clone()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn enrich_xai_subscription_title(mut metadata: Value, auth_config: Option<&str>) -> Value {
|
||||||
|
if metadata
|
||||||
|
.get("subscription_title")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.map(str::trim)
|
||||||
|
.is_some_and(|value| !value.is_empty())
|
||||||
|
{
|
||||||
|
return metadata;
|
||||||
|
}
|
||||||
|
let Some(config) = auth_config
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.and_then(|value| serde_json::from_str::<Value>(value).ok())
|
||||||
|
else {
|
||||||
|
return metadata;
|
||||||
|
};
|
||||||
|
let title = ["subscription_tier", "subscriptionTier", "tier", "plan"]
|
||||||
|
.iter()
|
||||||
|
.find_map(|field| {
|
||||||
|
config
|
||||||
|
.get(*field)
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.map(ToOwned::to_owned)
|
||||||
|
});
|
||||||
|
if let Some(title) = title {
|
||||||
|
if let Some(object) = metadata.as_object_mut() {
|
||||||
|
object.insert("subscription_title".to_string(), json!(title));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
metadata
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(crate) async fn refresh_xai_provider_quota_locally(
|
||||||
|
state: &AdminAppState<'_>,
|
||||||
|
provider: &StoredProviderCatalogProvider,
|
||||||
|
endpoint: &StoredProviderCatalogEndpoint,
|
||||||
|
keys: Vec<StoredProviderCatalogKey>,
|
||||||
|
proxy_override: Option<ProxySnapshot>,
|
||||||
|
) -> Result<Option<serde_json::Value>, GatewayError> {
|
||||||
|
let mut results = Vec::new();
|
||||||
|
let mut success_count = 0usize;
|
||||||
|
let mut failed_count = 0usize;
|
||||||
|
let mut auto_removed_count = 0usize;
|
||||||
|
|
||||||
|
for key in keys {
|
||||||
|
let transport = match state
|
||||||
|
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
|
||||||
|
.await?
|
||||||
|
{
|
||||||
|
Some(transport) => transport,
|
||||||
|
None => {
|
||||||
|
failed_count += 1;
|
||||||
|
results.push(json!({
|
||||||
|
"key_id": key.id,
|
||||||
|
"key_name": key.name,
|
||||||
|
"status": "error",
|
||||||
|
"message": "Provider transport snapshot unavailable",
|
||||||
|
}));
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
if xai_auth_uses_api(
|
||||||
|
transport.key.auth_type.as_str(),
|
||||||
|
transport.key.decrypted_auth_config.as_deref(),
|
||||||
|
) {
|
||||||
|
results.push(json!({
|
||||||
|
"key_id": key.id,
|
||||||
|
"key_name": key.name,
|
||||||
|
"status": "skipped",
|
||||||
|
"message": "xAI API Key 账号没有 Grok Build 订阅额度接口,请使用设备授权账号查询额度。",
|
||||||
|
}));
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
let authorization = match state.resolve_local_oauth_header_auth(&transport).await? {
|
||||||
|
Some(auth) => auth,
|
||||||
|
_ => {
|
||||||
|
if quota_key_auto_removed(state, &key.id).await? {
|
||||||
|
auto_removed_count += 1;
|
||||||
|
results.push(oauth_refresh_auto_removed_result(&key));
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
failed_count += 1;
|
||||||
|
results.push(json!({
|
||||||
|
"key_id": key.id,
|
||||||
|
"key_name": key.name,
|
||||||
|
"status": "error",
|
||||||
|
"message": "缺少 OAuth 认证信息,请先授权/刷新 Token",
|
||||||
|
}));
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
let fallback_user_id =
|
||||||
|
extract_xai_user_id_from_auth_config(transport.key.decrypted_auth_config.as_deref());
|
||||||
|
let user_id = match execute_xai_quota_plan(
|
||||||
|
state,
|
||||||
|
&transport,
|
||||||
|
build_xai_pool_user_request(
|
||||||
|
&transport.key.id,
|
||||||
|
xai_authorization_from_header(&authorization),
|
||||||
|
),
|
||||||
|
proxy_override.as_ref(),
|
||||||
|
)
|
||||||
|
.await?
|
||||||
|
{
|
||||||
|
ProviderQuotaExecutionOutcome::Response(result) if result.status_code == 200 => result
|
||||||
|
.body
|
||||||
|
.as_ref()
|
||||||
|
.and_then(|body| body.json_body.as_ref())
|
||||||
|
.and_then(extract_xai_user_id_from_value)
|
||||||
|
.or(fallback_user_id),
|
||||||
|
_ => fallback_user_id,
|
||||||
|
};
|
||||||
|
|
||||||
|
let result = match execute_xai_quota_plan(
|
||||||
|
state,
|
||||||
|
&transport,
|
||||||
|
build_xai_pool_billing_request(
|
||||||
|
&transport.key.id,
|
||||||
|
xai_authorization_from_header(&authorization),
|
||||||
|
user_id.as_deref(),
|
||||||
|
),
|
||||||
|
proxy_override.as_ref(),
|
||||||
|
)
|
||||||
|
.await?
|
||||||
|
{
|
||||||
|
ProviderQuotaExecutionOutcome::Response(result) => result,
|
||||||
|
ProviderQuotaExecutionOutcome::Failure(_) => {
|
||||||
|
failed_count += 1;
|
||||||
|
results.push(json!({
|
||||||
|
"key_id": key.id,
|
||||||
|
"key_name": key.name,
|
||||||
|
"status": "error",
|
||||||
|
"message": "xAI billing 请求执行失败",
|
||||||
|
"status_code": 502,
|
||||||
|
}));
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
let now_unix_secs = SystemTime::now()
|
||||||
|
.duration_since(UNIX_EPOCH)
|
||||||
|
.ok()
|
||||||
|
.map(|duration| duration.as_secs())
|
||||||
|
.unwrap_or(0);
|
||||||
|
let mut metadata_update = None::<serde_json::Value>;
|
||||||
|
let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) =
|
||||||
|
quota_refresh_success_invalid_state(&key);
|
||||||
|
let mut status = "error".to_string();
|
||||||
|
let mut message = None::<String>;
|
||||||
|
|
||||||
|
if result.status_code == 200 {
|
||||||
|
if let Some(body_json) = result
|
||||||
|
.body
|
||||||
|
.as_ref()
|
||||||
|
.and_then(|body| body.json_body.as_ref())
|
||||||
|
{
|
||||||
|
metadata_update =
|
||||||
|
parse_xai_billing_response(body_json, now_unix_secs).map(|metadata| {
|
||||||
|
json!({
|
||||||
|
"xai": enrich_xai_subscription_title(
|
||||||
|
metadata,
|
||||||
|
transport.key.decrypted_auth_config.as_deref(),
|
||||||
|
)
|
||||||
|
})
|
||||||
|
});
|
||||||
|
if metadata_update.is_some() {
|
||||||
|
status = "success".to_string();
|
||||||
|
} else {
|
||||||
|
status = "no_metadata".to_string();
|
||||||
|
message = Some("响应中未包含可用的 Grok Build 额度信息".to_string());
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
status = "no_metadata".to_string();
|
||||||
|
message = Some("响应中未包含配额信息".to_string());
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
message = Some(
|
||||||
|
extract_execution_error_message(&result)
|
||||||
|
.unwrap_or_else(|| format!("xAI billing 返回状态码 {}", result.status_code)),
|
||||||
|
);
|
||||||
|
if result.status_code == 401 || result.status_code == 403 {
|
||||||
|
let reason = message
|
||||||
|
.clone()
|
||||||
|
.unwrap_or_else(|| "账户访问被禁止".to_string());
|
||||||
|
oauth_invalid_at_unix_secs = Some(now_unix_secs);
|
||||||
|
oauth_invalid_reason = Some(format!("账户访问被禁止: {reason}"));
|
||||||
|
status = if result.status_code == 401 {
|
||||||
|
"unauthorized".to_string()
|
||||||
|
} else {
|
||||||
|
"forbidden".to_string()
|
||||||
|
};
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !persist_provider_quota_refresh_state(
|
||||||
|
state,
|
||||||
|
&key.id,
|
||||||
|
metadata_update.as_ref(),
|
||||||
|
oauth_invalid_at_unix_secs,
|
||||||
|
oauth_invalid_reason,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.await?
|
||||||
|
{
|
||||||
|
failed_count += 1;
|
||||||
|
results.push(json!({
|
||||||
|
"key_id": key.id,
|
||||||
|
"key_name": key.name,
|
||||||
|
"status": "error",
|
||||||
|
"message": "Key 状态写入失败",
|
||||||
|
}));
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
if status == "success" {
|
||||||
|
success_count += 1;
|
||||||
|
} else {
|
||||||
|
failed_count += 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut payload = serde_json::Map::new();
|
||||||
|
payload.insert("key_id".to_string(), json!(key.id));
|
||||||
|
payload.insert("key_name".to_string(), json!(key.name));
|
||||||
|
payload.insert("status".to_string(), json!(status));
|
||||||
|
if let Some(message) = message {
|
||||||
|
payload.insert("message".to_string(), json!(message));
|
||||||
|
}
|
||||||
|
if let Some(metadata) = metadata_update.as_ref().and_then(|value| value.get("xai")) {
|
||||||
|
payload.insert(
|
||||||
|
"metadata".to_string(),
|
||||||
|
admin_provider_metadata_bucket_safe_json("xai", Some(metadata)),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
if let Some(quota_snapshot) = build_quota_snapshot_payload(
|
||||||
|
"xai",
|
||||||
|
key.status_snapshot.as_ref(),
|
||||||
|
metadata_update.as_ref(),
|
||||||
|
) {
|
||||||
|
payload.insert("quota_snapshot".to_string(), quota_snapshot);
|
||||||
|
}
|
||||||
|
results.push(serde_json::Value::Object(payload));
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(Some(json!({
|
||||||
|
"success": success_count,
|
||||||
|
"failed": failed_count,
|
||||||
|
"total": results.len(),
|
||||||
|
"results": results,
|
||||||
|
"message": format!("已处理 {} 个 Key", results.len()),
|
||||||
|
"auto_removed": auto_removed_count,
|
||||||
|
})))
|
||||||
|
}
|
||||||
@@ -932,6 +932,13 @@ fn admin_pool_build_account_quota(
|
|||||||
return Some(account_quota);
|
return Some(account_quota);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
"xai" => {
|
||||||
|
if let Some(account_quota) =
|
||||||
|
admin_pool_build_kiro_account_quota_from_snapshot(quota_snapshot)
|
||||||
|
{
|
||||||
|
return Some(account_quota);
|
||||||
|
}
|
||||||
|
}
|
||||||
"chatgpt_web" => {
|
"chatgpt_web" => {
|
||||||
if let Some(account_quota) =
|
if let Some(account_quota) =
|
||||||
admin_pool_build_chatgpt_web_account_quota_from_snapshot(quota_snapshot)
|
admin_pool_build_chatgpt_web_account_quota_from_snapshot(quota_snapshot)
|
||||||
@@ -1591,4 +1598,29 @@ mod tests {
|
|||||||
Some("Auto剩余 40.0% (60/150) | Heavy剩余 0.0% (0/20)".to_string())
|
Some("Auto剩余 40.0% (60/150) | Heavy剩余 0.0% (0/20)".to_string())
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_account_quota_is_rendered_as_remaining_percent() {
|
||||||
|
let quota_snapshot = json!({
|
||||||
|
"provider_type": "xai",
|
||||||
|
"code": "ok",
|
||||||
|
"exhausted": false,
|
||||||
|
"plan_type": "SuperGrok",
|
||||||
|
"windows": [
|
||||||
|
{
|
||||||
|
"code": "usage",
|
||||||
|
"label": "周额度",
|
||||||
|
"scope": "account",
|
||||||
|
"used_ratio": 0.46,
|
||||||
|
"remaining_ratio": 0.54
|
||||||
|
}
|
||||||
|
]
|
||||||
|
});
|
||||||
|
let quota_snapshot = quota_snapshot.as_object().unwrap();
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
admin_pool_build_account_quota("xai", Some(quota_snapshot)),
|
||||||
|
Some("剩余 54.0%".to_string())
|
||||||
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3511,6 +3511,11 @@ async fn provider_query_execute_standard_test_candidate(
|
|||||||
codex_model_capabilities.as_ref(),
|
codex_model_capabilities.as_ref(),
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
crate::provider_transport::insert_cli_identity_headers_if_needed(
|
||||||
|
&transport,
|
||||||
|
provider_api_format,
|
||||||
|
&mut request_headers,
|
||||||
|
);
|
||||||
if !uses_vertex_query_auth {
|
if !uses_vertex_query_auth {
|
||||||
if let (Some(auth_header), Some(auth_value)) =
|
if let (Some(auth_header), Some(auth_value)) =
|
||||||
(auth_header.as_deref(), auth_value.as_deref())
|
(auth_header.as_deref(), auth_value.as_deref())
|
||||||
|
|||||||
@@ -4,9 +4,9 @@ pub(crate) fn normalize_provider_type_input(value: &str) -> Result<String, Strin
|
|||||||
let normalized = value.trim().to_ascii_lowercase();
|
let normalized = value.trim().to_ascii_lowercase();
|
||||||
match normalized.as_str() {
|
match normalized.as_str() {
|
||||||
"custom" | "claude_code" | "kiro" | "codex" | "chatgpt_web" | "gemini_cli"
|
"custom" | "claude_code" | "kiro" | "codex" | "chatgpt_web" | "gemini_cli"
|
||||||
| "antigravity" | "vertex_ai" | "grok" | "windsurf" => Ok(normalized),
|
| "antigravity" | "vertex_ai" | "grok" | "windsurf" | "xai" => Ok(normalized),
|
||||||
_ => Err(
|
_ => Err(
|
||||||
"provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai / grok / windsurf"
|
"provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai / grok / windsurf / xai"
|
||||||
.to_string(),
|
.to_string(),
|
||||||
),
|
),
|
||||||
}
|
}
|
||||||
@@ -405,6 +405,14 @@ mod tests {
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn normalize_provider_type_supports_xai() {
|
||||||
|
assert_eq!(
|
||||||
|
normalize_provider_type_input(" xAI ").expect("type should normalize"),
|
||||||
|
"xai"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn normalize_api_format_list_dedupes_canonical_formats() {
|
fn normalize_api_format_list_dedupes_canonical_formats() {
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
|
|||||||
@@ -90,6 +90,8 @@ pub(super) struct ResponsesWebSocketContinuationRecord {
|
|||||||
/// request JSON can never set it.
|
/// request JSON can never set it.
|
||||||
#[serde(default)]
|
#[serde(default)]
|
||||||
deepseek_opaque_reasoning_replay: bool,
|
deepseek_opaque_reasoning_replay: bool,
|
||||||
|
#[serde(default)]
|
||||||
|
xai_encrypted_reasoning_replay: bool,
|
||||||
/// A prior turn stored PII sentinels whose restore mapping exists only on
|
/// A prior turn stored PII sentinels whose restore mapping exists only on
|
||||||
/// the original downstream socket. Such a chain cannot safely resume on a
|
/// the original downstream socket. Such a chain cannot safely resume on a
|
||||||
/// new socket without leaking sentinels, so lookup succeeds but bootstrap
|
/// new socket without leaking sentinels, so lookup succeeds but bootstrap
|
||||||
@@ -122,6 +124,10 @@ impl ResponsesWebSocketContinuationRecord {
|
|||||||
normalization.reasoning_replay_policy(),
|
normalization.reasoning_replay_policy(),
|
||||||
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
|
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
|
||||||
),
|
),
|
||||||
|
xai_encrypted_reasoning_replay: matches!(
|
||||||
|
normalization.reasoning_replay_policy(),
|
||||||
|
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
|
||||||
|
),
|
||||||
has_connection_local_redaction,
|
has_connection_local_redaction,
|
||||||
responses_lite_static_config,
|
responses_lite_static_config,
|
||||||
};
|
};
|
||||||
@@ -156,7 +162,9 @@ impl ResponsesWebSocketContinuationRecord {
|
|||||||
pub(super) fn reasoning_replay_policy(
|
pub(super) fn reasoning_replay_policy(
|
||||||
&self,
|
&self,
|
||||||
) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy {
|
) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy {
|
||||||
if self.deepseek_opaque_reasoning_replay {
|
if self.xai_encrypted_reasoning_replay {
|
||||||
|
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
|
||||||
|
} else if self.deepseek_opaque_reasoning_replay {
|
||||||
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
|
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
|
||||||
} else {
|
} else {
|
||||||
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
|
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
|
||||||
@@ -476,6 +484,7 @@ mod tests {
|
|||||||
binding_fingerprint: [7; 32],
|
binding_fingerprint: [7; 32],
|
||||||
normalization_fingerprint: [9; 32],
|
normalization_fingerprint: [9; 32],
|
||||||
deepseek_opaque_reasoning_replay: false,
|
deepseek_opaque_reasoning_replay: false,
|
||||||
|
xai_encrypted_reasoning_replay: false,
|
||||||
has_connection_local_redaction: false,
|
has_connection_local_redaction: false,
|
||||||
responses_lite_static_config: Some(ResponsesLiteStaticConfig::from_response_create(
|
responses_lite_static_config: Some(ResponsesLiteStaticConfig::from_response_create(
|
||||||
&json!({
|
&json!({
|
||||||
@@ -714,6 +723,29 @@ mod tests {
|
|||||||
assert_eq!(decoded, record());
|
assert_eq!(decoded, record());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn serialized_record_preserves_xai_replay_policy_and_reads_legacy_records() {
|
||||||
|
let mut expected = record();
|
||||||
|
expected.xai_encrypted_reasoning_replay = true;
|
||||||
|
let mut serialized = serde_json::to_value(&expected).unwrap();
|
||||||
|
let decoded: ResponsesWebSocketContinuationRecord =
|
||||||
|
serde_json::from_value(serialized.clone()).unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
decoded.reasoning_replay_policy(),
|
||||||
|
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
|
||||||
|
);
|
||||||
|
serialized
|
||||||
|
.as_object_mut()
|
||||||
|
.unwrap()
|
||||||
|
.remove("xai_encrypted_reasoning_replay");
|
||||||
|
let legacy: ResponsesWebSocketContinuationRecord =
|
||||||
|
serde_json::from_value(serialized).unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
legacy.reasoning_replay_policy(),
|
||||||
|
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn serialized_record_preserves_only_the_server_derived_reasoning_replay_policy_bit() {
|
fn serialized_record_preserves_only_the_server_derived_reasoning_replay_policy_bit() {
|
||||||
let mut expected = record();
|
let mut expected = record();
|
||||||
|
|||||||
@@ -1388,6 +1388,168 @@ fn build_kiro_quota_status_snapshot(
|
|||||||
}))
|
}))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn build_xai_quota_status_snapshot(
|
||||||
|
upstream_metadata: Option<&Value>,
|
||||||
|
source: &str,
|
||||||
|
) -> Option<Value> {
|
||||||
|
let metadata = provider_quota_metadata_bucket(upstream_metadata, "xai")?;
|
||||||
|
let observed_at_unix_secs = provider_quota_timestamp_unix_secs(metadata.get("updated_at"));
|
||||||
|
let usage_limit = metadata
|
||||||
|
.get("usage_limit")
|
||||||
|
.and_then(admin_provider_quota_pure::coerce_json_f64);
|
||||||
|
let current_usage = metadata
|
||||||
|
.get("current_usage")
|
||||||
|
.and_then(admin_provider_quota_pure::coerce_json_f64);
|
||||||
|
let remaining = metadata
|
||||||
|
.get("remaining")
|
||||||
|
.and_then(admin_provider_quota_pure::coerce_json_f64);
|
||||||
|
let usage_ratio = metadata
|
||||||
|
.get("usage_percentage")
|
||||||
|
.and_then(admin_provider_quota_pure::coerce_json_f64)
|
||||||
|
.map(|value| (value / 100.0).clamp(0.0, 1.0))
|
||||||
|
.or_else(|| {
|
||||||
|
current_usage
|
||||||
|
.zip(usage_limit)
|
||||||
|
.and_then(|(current_usage, usage_limit)| {
|
||||||
|
(usage_limit > 0.0).then_some((current_usage / usage_limit).clamp(0.0, 1.0))
|
||||||
|
})
|
||||||
|
});
|
||||||
|
let remaining_ratio = usage_ratio.map(|value| (1.0 - value).max(0.0));
|
||||||
|
let next_reset_at = provider_quota_timestamp_unix_secs(metadata.get("next_reset_at"));
|
||||||
|
let reset_seconds = quota_window_reset_seconds(observed_at_unix_secs, next_reset_at);
|
||||||
|
let plan_type = metadata
|
||||||
|
.get("subscription_title")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.map(ToOwned::to_owned);
|
||||||
|
let period_type = metadata
|
||||||
|
.get("period_type")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.map(ToOwned::to_owned);
|
||||||
|
let usage_label = match period_type.as_deref() {
|
||||||
|
Some("monthly") => "月额度",
|
||||||
|
Some("weekly") => "周额度",
|
||||||
|
_ => "额度",
|
||||||
|
};
|
||||||
|
|
||||||
|
let mut windows = Vec::new();
|
||||||
|
if usage_ratio.is_some()
|
||||||
|
|| remaining.is_some()
|
||||||
|
|| usage_limit.is_some()
|
||||||
|
|| current_usage.is_some()
|
||||||
|
|| next_reset_at.is_some()
|
||||||
|
{
|
||||||
|
windows.push(json!({
|
||||||
|
"code": "usage",
|
||||||
|
"label": usage_label,
|
||||||
|
"scope": "account",
|
||||||
|
"unit": if usage_limit.is_some() { "usd" } else { "percent" },
|
||||||
|
"used_ratio": usage_ratio,
|
||||||
|
"remaining_ratio": remaining_ratio,
|
||||||
|
"used_value": current_usage,
|
||||||
|
"remaining_value": remaining,
|
||||||
|
"limit_value": usage_limit,
|
||||||
|
"reset_at": next_reset_at,
|
||||||
|
"reset_seconds": reset_seconds,
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
|
||||||
|
let prepaid_balance = metadata
|
||||||
|
.get("prepaid_balance")
|
||||||
|
.and_then(admin_provider_quota_pure::coerce_json_f64);
|
||||||
|
if prepaid_balance.is_some_and(|value| value > 0.0) {
|
||||||
|
windows.push(json!({
|
||||||
|
"code": "prepaid",
|
||||||
|
"label": "预付额度",
|
||||||
|
"scope": "account",
|
||||||
|
"unit": "usd",
|
||||||
|
"used_ratio": serde_json::Value::Null,
|
||||||
|
"remaining_ratio": serde_json::Value::Null,
|
||||||
|
"remaining_value": prepaid_balance,
|
||||||
|
"reset_at": serde_json::Value::Null,
|
||||||
|
"reset_seconds": serde_json::Value::Null,
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
|
||||||
|
let on_demand_cap = metadata
|
||||||
|
.get("on_demand_cap")
|
||||||
|
.and_then(admin_provider_quota_pure::coerce_json_f64);
|
||||||
|
let on_demand_used = metadata
|
||||||
|
.get("on_demand_used")
|
||||||
|
.and_then(admin_provider_quota_pure::coerce_json_f64);
|
||||||
|
let on_demand_enabled = metadata
|
||||||
|
.get("on_demand_enabled")
|
||||||
|
.and_then(admin_provider_quota_pure::coerce_json_bool)
|
||||||
|
!= Some(false);
|
||||||
|
if on_demand_enabled && on_demand_cap.is_some_and(|value| value > 0.0) {
|
||||||
|
let on_demand_remaining = on_demand_cap
|
||||||
|
.zip(on_demand_used)
|
||||||
|
.map(|(cap, used)| (cap - used).max(0.0));
|
||||||
|
let on_demand_ratio = on_demand_cap
|
||||||
|
.zip(on_demand_used)
|
||||||
|
.and_then(|(cap, used)| (cap > 0.0).then_some((used / cap).clamp(0.0, 1.0)));
|
||||||
|
windows.push(json!({
|
||||||
|
"code": "on_demand",
|
||||||
|
"label": "按需额度",
|
||||||
|
"scope": "account",
|
||||||
|
"unit": "usd",
|
||||||
|
"used_ratio": on_demand_ratio,
|
||||||
|
"remaining_ratio": on_demand_ratio.map(|value| (1.0 - value).max(0.0)),
|
||||||
|
"used_value": on_demand_used,
|
||||||
|
"remaining_value": on_demand_remaining,
|
||||||
|
"limit_value": on_demand_cap,
|
||||||
|
"reset_at": serde_json::Value::Null,
|
||||||
|
"reset_seconds": serde_json::Value::Null,
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
|
||||||
|
if windows.is_empty() && plan_type.is_none() && observed_at_unix_secs.is_none() {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
|
||||||
|
let prepaid_available = prepaid_balance.is_some_and(|value| value > 0.0);
|
||||||
|
let on_demand_available = on_demand_enabled
|
||||||
|
&& on_demand_cap.is_some_and(|value| value > 0.0)
|
||||||
|
&& on_demand_used
|
||||||
|
.zip(on_demand_cap)
|
||||||
|
.is_some_and(|(used, cap)| used < cap);
|
||||||
|
let usage_exhausted = remaining.is_some_and(|value| value <= 0.0)
|
||||||
|
|| usage_ratio.is_some_and(|value| value >= 1.0 - 1e-6);
|
||||||
|
let exhausted = usage_exhausted && !prepaid_available && !on_demand_available;
|
||||||
|
let reason = if exhausted {
|
||||||
|
Some("额度已耗尽".to_string())
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
let label = if exhausted {
|
||||||
|
Some("额度耗尽")
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
let code = if exhausted { "exhausted" } else { "ok" };
|
||||||
|
|
||||||
|
Some(json!({
|
||||||
|
"version": 2,
|
||||||
|
"provider_type": "xai",
|
||||||
|
"code": code,
|
||||||
|
"label": label,
|
||||||
|
"reason": reason,
|
||||||
|
"freshness": "fresh",
|
||||||
|
"source": source,
|
||||||
|
"observed_at": observed_at_unix_secs,
|
||||||
|
"exhausted": exhausted,
|
||||||
|
"usage_ratio": usage_ratio,
|
||||||
|
"updated_at": observed_at_unix_secs,
|
||||||
|
"reset_at": next_reset_at,
|
||||||
|
"reset_seconds": reset_seconds,
|
||||||
|
"plan_type": plan_type,
|
||||||
|
"windows": windows,
|
||||||
|
}))
|
||||||
|
}
|
||||||
|
|
||||||
fn build_chatgpt_web_quota_status_snapshot(
|
fn build_chatgpt_web_quota_status_snapshot(
|
||||||
upstream_metadata: Option<&Value>,
|
upstream_metadata: Option<&Value>,
|
||||||
source: &str,
|
source: &str,
|
||||||
@@ -2255,6 +2417,7 @@ pub(crate) fn sync_provider_key_quota_status_snapshot(
|
|||||||
let mut quota = match normalized_provider_type.as_str() {
|
let mut quota = match normalized_provider_type.as_str() {
|
||||||
"codex" => build_codex_quota_status_snapshot(upstream_metadata, source),
|
"codex" => build_codex_quota_status_snapshot(upstream_metadata, source),
|
||||||
"kiro" => build_kiro_quota_status_snapshot(upstream_metadata, source),
|
"kiro" => build_kiro_quota_status_snapshot(upstream_metadata, source),
|
||||||
|
"xai" => build_xai_quota_status_snapshot(upstream_metadata, source),
|
||||||
"chatgpt_web" => build_chatgpt_web_quota_status_snapshot(upstream_metadata, source),
|
"chatgpt_web" => build_chatgpt_web_quota_status_snapshot(upstream_metadata, source),
|
||||||
"windsurf" => build_windsurf_quota_status_snapshot(upstream_metadata, source),
|
"windsurf" => build_windsurf_quota_status_snapshot(upstream_metadata, source),
|
||||||
"antigravity" => build_antigravity_quota_status_snapshot(upstream_metadata, source),
|
"antigravity" => build_antigravity_quota_status_snapshot(upstream_metadata, source),
|
||||||
@@ -3622,6 +3785,43 @@ mod tests {
|
|||||||
assert_eq!(auto.get("used_value"), Some(&json!(90.0)));
|
assert_eq!(auto.get("used_value"), Some(&json!(90.0)));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn provider_key_status_snapshot_payload_backfills_xai_weekly_credits() {
|
||||||
|
let mut key = sample_catalog_key();
|
||||||
|
key.upstream_metadata = Some(json!({
|
||||||
|
"xai": {
|
||||||
|
"updated_at": 1_778_067_246u64,
|
||||||
|
"usage_percentage": 46.0,
|
||||||
|
"period_type": "weekly",
|
||||||
|
"next_reset_at": 1_778_157_172u64,
|
||||||
|
"subscription_title": "SuperGrok",
|
||||||
|
"prepaid_balance": 0.0,
|
||||||
|
"on_demand_cap": 0.0,
|
||||||
|
"on_demand_used": 0.0
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
|
||||||
|
let payload = provider_key_status_snapshot_payload(&key, "xai");
|
||||||
|
let quota = payload
|
||||||
|
.get("quota")
|
||||||
|
.and_then(Value::as_object)
|
||||||
|
.expect("quota snapshot should be object");
|
||||||
|
let windows = quota
|
||||||
|
.get("windows")
|
||||||
|
.and_then(Value::as_array)
|
||||||
|
.expect("xai quota windows should exist");
|
||||||
|
|
||||||
|
assert_eq!(quota.get("provider_type"), Some(&json!("xai")));
|
||||||
|
assert_eq!(quota.get("code"), Some(&json!("ok")));
|
||||||
|
assert_eq!(quota.get("exhausted"), Some(&json!(false)));
|
||||||
|
assert_eq!(quota.get("plan_type"), Some(&json!("SuperGrok")));
|
||||||
|
assert_eq!(quota.get("usage_ratio"), Some(&json!(0.46)));
|
||||||
|
assert_eq!(quota.get("reset_at"), Some(&json!(1_778_157_172u64)));
|
||||||
|
assert_eq!(windows.len(), 1);
|
||||||
|
assert_eq!(windows[0].get("code"), Some(&json!("usage")));
|
||||||
|
assert_eq!(windows[0].get("label"), Some(&json!("周额度")));
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn provider_key_status_snapshot_payload_backfills_gemini_cli_account_credits() {
|
fn provider_key_status_snapshot_payload_backfills_gemini_cli_account_credits() {
|
||||||
let mut key = sample_catalog_key();
|
let mut key = sample_catalog_key();
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ pub(crate) fn openai_image_provider_max_generation_count(provider_type: &str) ->
|
|||||||
GROK_OPENAI_IMAGE_MAX_GENERATION_COUNT
|
GROK_OPENAI_IMAGE_MAX_GENERATION_COUNT
|
||||||
} else if matches!(
|
} else if matches!(
|
||||||
provider_type.trim().to_ascii_lowercase().as_str(),
|
provider_type.trim().to_ascii_lowercase().as_str(),
|
||||||
"openai" | "codex"
|
"openai" | "codex" | "xai"
|
||||||
) {
|
) {
|
||||||
OPENAI_IMAGE_MAX_GENERATION_COUNT
|
OPENAI_IMAGE_MAX_GENERATION_COUNT
|
||||||
} else {
|
} else {
|
||||||
@@ -58,6 +58,7 @@ mod tests {
|
|||||||
assert_eq!(openai_image_provider_max_generation_count("grok"), 4);
|
assert_eq!(openai_image_provider_max_generation_count("grok"), 4);
|
||||||
assert_eq!(openai_image_provider_max_generation_count("openai"), 10);
|
assert_eq!(openai_image_provider_max_generation_count("openai"), 10);
|
||||||
assert_eq!(openai_image_provider_max_generation_count("codex"), 10);
|
assert_eq!(openai_image_provider_max_generation_count("codex"), 10);
|
||||||
|
assert_eq!(openai_image_provider_max_generation_count("xai"), 10);
|
||||||
assert_eq!(openai_image_provider_max_generation_count("custom"), 1);
|
assert_eq!(openai_image_provider_max_generation_count("custom"), 1);
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
openai_image_provider_max_generation_count_for_model("openai", Some("dall-e-3")),
|
openai_image_provider_max_generation_count_for_model("openai", Some("dall-e-3")),
|
||||||
|
|||||||
@@ -172,6 +172,7 @@ fn provider_uses_bearer_oauth_runtime(provider_type: &str) -> bool {
|
|||||||
| "antigravity"
|
| "antigravity"
|
||||||
| "kiro"
|
| "kiro"
|
||||||
| "windsurf"
|
| "windsurf"
|
||||||
|
| "xai"
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -406,6 +407,22 @@ mod tests {
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn recognizes_xai_oauth_as_bearer_runtime() {
|
||||||
|
let semantics = provider_key_auth_semantics(&sample_key("oauth"), "xai");
|
||||||
|
|
||||||
|
assert!(semantics.oauth_managed());
|
||||||
|
assert!(semantics.can_refresh_oauth());
|
||||||
|
assert_eq!(
|
||||||
|
semantics.credential_kind(),
|
||||||
|
ProviderKeyCredentialKind::OAuthSession
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
semantics.runtime_auth_kind(),
|
||||||
|
ProviderKeyRuntimeAuthKind::Bearer
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn refresh_capability_requires_stored_refresh_token() {
|
fn refresh_capability_requires_stored_refresh_token() {
|
||||||
let semantics = provider_key_auth_semantics(&sample_key("oauth"), "codex");
|
let semantics = provider_key_auth_semantics(&sample_key("oauth"), "codex");
|
||||||
|
|||||||
@@ -186,6 +186,8 @@ fn frontend_path_bypasses_static(path: &str) -> bool {
|
|||||||
"/health" | "/test-connection" | crate::constants::READYZ_PATH
|
"/health" | "/test-connection" | crate::constants::READYZ_PATH
|
||||||
) || path.starts_with("/api/")
|
) || path.starts_with("/api/")
|
||||||
|| path.starts_with("/v1/")
|
|| path.starts_with("/v1/")
|
||||||
|
|| path == "/openai/v1/videos"
|
||||||
|
|| path.starts_with("/openai/v1/videos/")
|
||||||
|| path.starts_with("/v1beta/")
|
|| path.starts_with("/v1beta/")
|
||||||
|| path.starts_with("/upload/")
|
|| path.starts_with("/upload/")
|
||||||
|| path.starts_with("/_gateway/")
|
|| path.starts_with("/_gateway/")
|
||||||
|
|||||||
@@ -290,6 +290,14 @@ impl provider_transport::VideoTaskTransportSnapshotLookup for AppState {
|
|||||||
.await
|
.await
|
||||||
.map_err(GatewayError::into_message)
|
.map_err(GatewayError::into_message)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn resolve_video_task_proxy(
|
||||||
|
&self,
|
||||||
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
|
) -> Option<ProxySnapshot> {
|
||||||
|
self.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
|
||||||
|
.await
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[async_trait]
|
#[async_trait]
|
||||||
|
|||||||
@@ -1789,6 +1789,7 @@ fn admin_provider_oauth_quota_mod_stays_thin() {
|
|||||||
"pub(crate) mod dispatch;",
|
"pub(crate) mod dispatch;",
|
||||||
"pub(crate) mod kiro;",
|
"pub(crate) mod kiro;",
|
||||||
"pub(crate) mod shared;",
|
"pub(crate) mod shared;",
|
||||||
|
"pub(crate) mod xai;",
|
||||||
] {
|
] {
|
||||||
assert!(
|
assert!(
|
||||||
quota_mod.contains(pattern),
|
quota_mod.contains(pattern),
|
||||||
@@ -1861,6 +1862,7 @@ fn admin_provider_oauth_quota_mod_stays_thin() {
|
|||||||
"refresh_antigravity_provider_quota_locally",
|
"refresh_antigravity_provider_quota_locally",
|
||||||
"refresh_gemini_cli_provider_quota_locally",
|
"refresh_gemini_cli_provider_quota_locally",
|
||||||
"refresh_chatgpt_web_provider_quota_locally",
|
"refresh_chatgpt_web_provider_quota_locally",
|
||||||
|
"refresh_xai_provider_quota_locally",
|
||||||
] {
|
] {
|
||||||
assert!(
|
assert!(
|
||||||
quota_dispatch.contains(pattern),
|
quota_dispatch.contains(pattern),
|
||||||
|
|||||||
@@ -1452,6 +1452,7 @@ fn ai_serving_planner_separates_local_candidate_resolution_from_ranking() {
|
|||||||
"GeminiCliProviderPoolAdapter",
|
"GeminiCliProviderPoolAdapter",
|
||||||
"KiroProviderPoolAdapter",
|
"KiroProviderPoolAdapter",
|
||||||
"ChatGptWebProviderPoolAdapter",
|
"ChatGptWebProviderPoolAdapter",
|
||||||
|
"XaiProviderPoolAdapter",
|
||||||
"CLAUDE_CODE_PROVIDER_POOL_ADAPTER",
|
"CLAUDE_CODE_PROVIDER_POOL_ADAPTER",
|
||||||
"VERTEX_AI_PROVIDER_POOL_ADAPTER",
|
"VERTEX_AI_PROVIDER_POOL_ADAPTER",
|
||||||
"provider_types_for_capability",
|
"provider_types_for_capability",
|
||||||
@@ -1478,6 +1479,7 @@ fn ai_serving_planner_separates_local_candidate_resolution_from_ranking() {
|
|||||||
"pub mod gemini_cli;",
|
"pub mod gemini_cli;",
|
||||||
"pub mod kiro;",
|
"pub mod kiro;",
|
||||||
"pub mod chatgpt_web;",
|
"pub mod chatgpt_web;",
|
||||||
|
"pub mod xai;",
|
||||||
] {
|
] {
|
||||||
assert!(
|
assert!(
|
||||||
provider_pool_providers.contains(pattern),
|
provider_pool_providers.contains(pattern),
|
||||||
@@ -1513,6 +1515,14 @@ fn ai_serving_planner_separates_local_candidate_resolution_from_ranking() {
|
|||||||
"crates/aether-provider/pool/src/providers/kiro.rs",
|
"crates/aether-provider/pool/src/providers/kiro.rs",
|
||||||
vec!["KiroProviderPoolAdapter", "quota_exhausted_from_bucket"],
|
vec!["KiroProviderPoolAdapter", "quota_exhausted_from_bucket"],
|
||||||
),
|
),
|
||||||
|
(
|
||||||
|
"crates/aether-provider/pool/src/providers/xai.rs",
|
||||||
|
vec![
|
||||||
|
"XaiProviderPoolAdapter",
|
||||||
|
"build_xai_pool_billing_request",
|
||||||
|
"quota_exhausted_from_bucket",
|
||||||
|
],
|
||||||
|
),
|
||||||
(
|
(
|
||||||
"crates/aether-provider/pool/src/providers/chatgpt_web.rs",
|
"crates/aether-provider/pool/src/providers/chatgpt_web.rs",
|
||||||
vec![
|
vec![
|
||||||
|
|||||||
@@ -983,6 +983,297 @@ async fn gateway_rejects_generic_oauth_start_for_windsurf_provider_impl() {
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn gateway_rejects_generic_oauth_start_for_xai_provider() {
|
||||||
|
run_admin_oauth_test(
|
||||||
|
"gateway_rejects_generic_oauth_start_for_xai_provider",
|
||||||
|
gateway_rejects_generic_oauth_start_for_xai_provider_impl,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn gateway_rejects_generic_oauth_start_for_xai_provider_impl() {
|
||||||
|
let mut provider = sample_provider("provider-xai", "xai", 10);
|
||||||
|
provider.provider_type = "xai".to_string();
|
||||||
|
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||||
|
vec![provider],
|
||||||
|
vec![],
|
||||||
|
vec![],
|
||||||
|
));
|
||||||
|
let state = AppState::new()
|
||||||
|
.expect("gateway should build")
|
||||||
|
.with_data_state_for_tests(GatewayDataState::with_provider_catalog_reader_for_tests(
|
||||||
|
provider_catalog_repository,
|
||||||
|
));
|
||||||
|
|
||||||
|
let response = local_admin_provider_oauth_response(
|
||||||
|
&state,
|
||||||
|
http::Method::POST,
|
||||||
|
"/api/admin/provider-oauth/providers/provider-xai/start",
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
|
||||||
|
let body = to_bytes(response.into_body(), usize::MAX)
|
||||||
|
.await
|
||||||
|
.expect("body should read");
|
||||||
|
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
|
||||||
|
assert!(
|
||||||
|
payload["detail"].as_str().is_some_and(|detail| {
|
||||||
|
detail.contains("设备授权") || detail.contains("导入凭据")
|
||||||
|
}),
|
||||||
|
"payload={payload}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn gateway_handles_admin_provider_oauth_device_authorize_for_xai() {
|
||||||
|
run_admin_oauth_test(
|
||||||
|
"gateway_handles_admin_provider_oauth_device_authorize_for_xai",
|
||||||
|
gateway_handles_admin_provider_oauth_device_authorize_for_xai_impl,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn gateway_handles_admin_provider_oauth_device_authorize_for_xai_impl() {
|
||||||
|
let authorize_hits = Arc::new(Mutex::new(0usize));
|
||||||
|
let authorize_hits_clone = Arc::clone(&authorize_hits);
|
||||||
|
let oidc_server = Router::new().fallback(any(move |_request: Request| {
|
||||||
|
let authorize_hits_inner = Arc::clone(&authorize_hits_clone);
|
||||||
|
async move {
|
||||||
|
*authorize_hits_inner.lock().expect("mutex should lock") += 1;
|
||||||
|
Json(json!({
|
||||||
|
"device_code": "xai-device-code",
|
||||||
|
"user_code": "XAI-CODE",
|
||||||
|
"verification_uri": "https://auth.x.ai/activate",
|
||||||
|
"verification_uri_complete": "https://auth.x.ai/activate?user_code=XAI-CODE",
|
||||||
|
"expires_in": 600,
|
||||||
|
"interval": 5,
|
||||||
|
}))
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
|
||||||
|
let mut provider = sample_provider("provider-xai", "xai", 10);
|
||||||
|
provider.provider_type = "xai".to_string();
|
||||||
|
let endpoint = sample_endpoint(
|
||||||
|
"endpoint-xai-responses",
|
||||||
|
"provider-xai",
|
||||||
|
"openai:responses",
|
||||||
|
"https://cli-chat-proxy.grok.com/v1",
|
||||||
|
);
|
||||||
|
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||||
|
vec![provider],
|
||||||
|
vec![endpoint],
|
||||||
|
vec![],
|
||||||
|
));
|
||||||
|
let (oidc_url, oidc_handle) = start_server(oidc_server).await;
|
||||||
|
let state = AppState::new()
|
||||||
|
.expect("gateway should build")
|
||||||
|
.with_data_state_for_tests(GatewayDataState::with_provider_catalog_reader_for_tests(
|
||||||
|
provider_catalog_repository,
|
||||||
|
))
|
||||||
|
.with_provider_oauth_token_url_for_tests(
|
||||||
|
"xai_device",
|
||||||
|
format!("{oidc_url}/oauth2/device/code"),
|
||||||
|
)
|
||||||
|
.with_provider_oauth_token_url_for_tests("xai", format!("{oidc_url}/oauth2/token"));
|
||||||
|
|
||||||
|
let response = local_admin_provider_oauth_response(
|
||||||
|
&state,
|
||||||
|
http::Method::POST,
|
||||||
|
"/api/admin/provider-oauth/providers/provider-xai/device-authorize",
|
||||||
|
Some(json!({ "proxy_node_id": "proxy-node-xai" })),
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
let status = response.status();
|
||||||
|
let body = to_bytes(response.into_body(), usize::MAX)
|
||||||
|
.await
|
||||||
|
.expect("body should read");
|
||||||
|
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
|
||||||
|
assert_eq!(status, StatusCode::OK, "payload={payload}");
|
||||||
|
let session_id = payload["session_id"]
|
||||||
|
.as_str()
|
||||||
|
.expect("session_id should exist")
|
||||||
|
.to_string();
|
||||||
|
assert_eq!(payload["user_code"], "XAI-CODE");
|
||||||
|
assert_eq!(payload["verification_uri"], "https://auth.x.ai/activate");
|
||||||
|
assert_eq!(
|
||||||
|
payload["verification_uri_complete"],
|
||||||
|
"https://auth.x.ai/activate?user_code=XAI-CODE"
|
||||||
|
);
|
||||||
|
assert_eq!(payload["auth_type"], "device");
|
||||||
|
assert!(payload.get("callback_required").is_none() || payload["callback_required"] == false);
|
||||||
|
assert_eq!(*authorize_hits.lock().expect("mutex should lock"), 1);
|
||||||
|
|
||||||
|
let stored = state
|
||||||
|
.load_provider_oauth_device_session_for_tests(&format!("device_auth_session:{session_id}"))
|
||||||
|
.expect("device session should be stored");
|
||||||
|
let stored: serde_json::Value =
|
||||||
|
serde_json::from_str(&stored).expect("device session json should parse");
|
||||||
|
assert_eq!(stored["provider_id"], "provider-xai");
|
||||||
|
assert_eq!(stored["device_code"], "xai-device-code");
|
||||||
|
assert_eq!(stored["auth_type"], "device");
|
||||||
|
assert_eq!(stored["redirect_uri"], format!("{oidc_url}/oauth2/token"));
|
||||||
|
assert_eq!(stored["proxy_node_id"], "proxy-node-xai");
|
||||||
|
assert_eq!(stored["status"], "pending");
|
||||||
|
|
||||||
|
oidc_handle.abort();
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn gateway_handles_admin_provider_oauth_device_poll_for_xai() {
|
||||||
|
run_admin_oauth_test(
|
||||||
|
"gateway_handles_admin_provider_oauth_device_poll_for_xai",
|
||||||
|
gateway_handles_admin_provider_oauth_device_poll_for_xai_impl,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn gateway_handles_admin_provider_oauth_device_poll_for_xai_impl() {
|
||||||
|
let token_hits = Arc::new(Mutex::new(0usize));
|
||||||
|
let token_hits_clone = Arc::clone(&token_hits);
|
||||||
|
let access_token = sample_kiro_device_access_token("[email protected]");
|
||||||
|
let id_token = access_token.clone();
|
||||||
|
let token_server = Router::new().fallback(any(move |_request: Request| {
|
||||||
|
let token_hits_inner = Arc::clone(&token_hits_clone);
|
||||||
|
let access_token = access_token.clone();
|
||||||
|
let id_token = id_token.clone();
|
||||||
|
async move {
|
||||||
|
let hit = {
|
||||||
|
let mut hits = token_hits_inner.lock().expect("mutex should lock");
|
||||||
|
*hits += 1;
|
||||||
|
*hits
|
||||||
|
};
|
||||||
|
if hit == 1 {
|
||||||
|
return (
|
||||||
|
StatusCode::BAD_REQUEST,
|
||||||
|
Json(json!({ "error": "authorization_pending" })),
|
||||||
|
)
|
||||||
|
.into_response();
|
||||||
|
}
|
||||||
|
Json(json!({
|
||||||
|
"access_token": access_token,
|
||||||
|
"refresh_token": "xai-refresh-token",
|
||||||
|
"token_type": "Bearer",
|
||||||
|
"expires_in": 3600,
|
||||||
|
"id_token": id_token,
|
||||||
|
}))
|
||||||
|
.into_response()
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
|
||||||
|
let mut provider = sample_provider("provider-xai", "xai", 10);
|
||||||
|
provider.provider_type = "xai".to_string();
|
||||||
|
let endpoint = sample_endpoint(
|
||||||
|
"endpoint-xai-responses",
|
||||||
|
"provider-xai",
|
||||||
|
"openai:responses",
|
||||||
|
"https://cli-chat-proxy.grok.com/v1",
|
||||||
|
);
|
||||||
|
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||||
|
vec![provider],
|
||||||
|
vec![endpoint],
|
||||||
|
vec![],
|
||||||
|
));
|
||||||
|
let (token_url, token_handle) = start_server(token_server).await;
|
||||||
|
let resolved_token_url = format!("{token_url}/oauth2/token");
|
||||||
|
let state = AppState::new()
|
||||||
|
.expect("gateway should build")
|
||||||
|
.with_data_state_for_tests(
|
||||||
|
GatewayDataState::with_provider_catalog_repository_for_tests(
|
||||||
|
provider_catalog_repository.clone(),
|
||||||
|
)
|
||||||
|
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||||
|
)
|
||||||
|
.with_provider_oauth_device_session_entry_for_tests(
|
||||||
|
"session-xai",
|
||||||
|
json!({
|
||||||
|
"provider_id": "provider-xai",
|
||||||
|
"region": "",
|
||||||
|
"client_id": "b1a00492-073a-47ea-816f-4c329264a828",
|
||||||
|
"client_secret": "",
|
||||||
|
"device_code": "xai-device-code",
|
||||||
|
"auth_type": "device",
|
||||||
|
"social_provider": null,
|
||||||
|
"code_verifier": null,
|
||||||
|
"redirect_uri": resolved_token_url,
|
||||||
|
"machine_id": null,
|
||||||
|
"interval": 5,
|
||||||
|
"expires_at_unix_secs": 4_102_444_800u64,
|
||||||
|
"status": "pending",
|
||||||
|
"proxy_node_id": null,
|
||||||
|
"created_at_unix_ms": 1_711_000_000u64,
|
||||||
|
"key_id": null,
|
||||||
|
"email": null,
|
||||||
|
"replaced": false,
|
||||||
|
"error_msg": null,
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
.with_provider_oauth_token_url_for_tests("xai", resolved_token_url.clone());
|
||||||
|
|
||||||
|
let pending = local_admin_provider_oauth_response(
|
||||||
|
&state,
|
||||||
|
http::Method::POST,
|
||||||
|
"/api/admin/provider-oauth/providers/provider-xai/device-poll",
|
||||||
|
Some(json!({ "session_id": "session-xai" })),
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
let pending_body = to_bytes(pending.into_body(), usize::MAX)
|
||||||
|
.await
|
||||||
|
.expect("pending body should read");
|
||||||
|
let pending_payload: serde_json::Value =
|
||||||
|
serde_json::from_slice(&pending_body).expect("pending json should parse");
|
||||||
|
assert_eq!(
|
||||||
|
pending_payload["status"], "pending",
|
||||||
|
"payload={pending_payload}"
|
||||||
|
);
|
||||||
|
|
||||||
|
let authorized = local_admin_provider_oauth_response(
|
||||||
|
&state,
|
||||||
|
http::Method::POST,
|
||||||
|
"/api/admin/provider-oauth/providers/provider-xai/device-poll",
|
||||||
|
Some(json!({ "session_id": "session-xai" })),
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
let status = authorized.status();
|
||||||
|
let body = to_bytes(authorized.into_body(), usize::MAX)
|
||||||
|
.await
|
||||||
|
.expect("authorized body should read");
|
||||||
|
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
|
||||||
|
assert_eq!(status, StatusCode::OK, "payload={payload}");
|
||||||
|
assert_eq!(payload["status"], "authorized");
|
||||||
|
assert_eq!(payload["email"], "[email protected]");
|
||||||
|
assert_eq!(payload["replaced"], false);
|
||||||
|
assert_eq!(*token_hits.lock().expect("mutex should lock"), 2);
|
||||||
|
|
||||||
|
let stored = state
|
||||||
|
.load_provider_oauth_device_session_for_tests("device_auth_session:session-xai")
|
||||||
|
.expect("device session should persist");
|
||||||
|
let stored: serde_json::Value =
|
||||||
|
serde_json::from_str(&stored).expect("device session json should parse");
|
||||||
|
assert_eq!(stored["status"], "authorized");
|
||||||
|
let key_id = stored["key_id"]
|
||||||
|
.as_str()
|
||||||
|
.expect("key_id should be stored")
|
||||||
|
.to_string();
|
||||||
|
assert_eq!(payload["key_id"], key_id);
|
||||||
|
|
||||||
|
let persisted = provider_catalog_repository
|
||||||
|
.list_keys_by_ids(std::slice::from_ref(&key_id))
|
||||||
|
.await
|
||||||
|
.expect("keys should load")
|
||||||
|
.into_iter()
|
||||||
|
.next()
|
||||||
|
.expect("persisted key should exist");
|
||||||
|
assert_eq!(persisted.auth_type, "oauth");
|
||||||
|
let decrypted_auth_config = decrypt_persisted_provider_auth_config(&persisted);
|
||||||
|
let auth_config: serde_json::Value =
|
||||||
|
serde_json::from_str(&decrypted_auth_config).expect("auth config should parse");
|
||||||
|
assert_eq!(auth_config["provider_type"], "xai");
|
||||||
|
assert_eq!(auth_config["auth_method"], "oauth");
|
||||||
|
assert_eq!(auth_config["using_api"], false);
|
||||||
|
assert_eq!(auth_config["email"], "[email protected]");
|
||||||
|
|
||||||
|
token_handle.abort();
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn gateway_handles_admin_provider_oauth_device_poll_for_windsurf_one_time_token() {
|
fn gateway_handles_admin_provider_oauth_device_poll_for_windsurf_one_time_token() {
|
||||||
run_admin_oauth_test(
|
run_admin_oauth_test(
|
||||||
|
|||||||
@@ -36,6 +36,7 @@ mod openai_sync_task;
|
|||||||
mod registry_poller;
|
mod registry_poller;
|
||||||
mod routing;
|
mod routing;
|
||||||
mod stream;
|
mod stream;
|
||||||
|
mod xai;
|
||||||
|
|
||||||
/// Seed online manual proxy nodes for video execution fixtures.
|
/// Seed online manual proxy nodes for video execution fixtures.
|
||||||
///
|
///
|
||||||
@@ -44,6 +45,17 @@ mod stream;
|
|||||||
/// the same deployment-state record; the loopback URL is never contacted when
|
/// the same deployment-state record; the loopback URL is never contacted when
|
||||||
/// the execution-runtime override is active.
|
/// the execution-runtime override is active.
|
||||||
pub(super) fn video_proxy_node_repository<I, S>(node_ids: I) -> Arc<InMemoryProxyNodeRepository>
|
pub(super) fn video_proxy_node_repository<I, S>(node_ids: I) -> Arc<InMemoryProxyNodeRepository>
|
||||||
|
where
|
||||||
|
I: IntoIterator<Item = S>,
|
||||||
|
S: AsRef<str>,
|
||||||
|
{
|
||||||
|
video_proxy_node_repository_at_url(node_ids, "http://127.0.0.1:1")
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) fn video_proxy_node_repository_at_url<I, S>(
|
||||||
|
node_ids: I,
|
||||||
|
proxy_url: &str,
|
||||||
|
) -> Arc<InMemoryProxyNodeRepository>
|
||||||
where
|
where
|
||||||
I: IntoIterator<Item = S>,
|
I: IntoIterator<Item = S>,
|
||||||
S: AsRef<str>,
|
S: AsRef<str>,
|
||||||
@@ -68,7 +80,7 @@ where
|
|||||||
1,
|
1,
|
||||||
)
|
)
|
||||||
.expect("video test proxy node should build")
|
.expect("video test proxy node should build")
|
||||||
.with_manual_proxy_fields(Some("http://127.0.0.1:1".to_string()), None, None)
|
.with_manual_proxy_fields(Some(proxy_url.to_string()), None, None)
|
||||||
.with_tunnel_generation(format!("video-test-generation-{node_id}"))
|
.with_tunnel_generation(format!("video-test-generation-{node_id}"))
|
||||||
});
|
});
|
||||||
Arc::new(InMemoryProxyNodeRepository::seed(nodes))
|
Arc::new(InMemoryProxyNodeRepository::seed(nodes))
|
||||||
@@ -86,6 +98,28 @@ pub(super) fn video_provider_catalog_repository(
|
|||||||
endpoint_base_url: &str,
|
endpoint_base_url: &str,
|
||||||
key_id: &str,
|
key_id: &str,
|
||||||
upstream_api_key: &str,
|
upstream_api_key: &str,
|
||||||
|
) -> Arc<InMemoryProviderCatalogReadRepository> {
|
||||||
|
video_provider_catalog_repository_with_proxy(
|
||||||
|
provider_id,
|
||||||
|
provider_type,
|
||||||
|
endpoint_id,
|
||||||
|
api_format,
|
||||||
|
endpoint_base_url,
|
||||||
|
key_id,
|
||||||
|
upstream_api_key,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) fn video_provider_catalog_repository_with_proxy(
|
||||||
|
provider_id: &str,
|
||||||
|
provider_type: &str,
|
||||||
|
endpoint_id: &str,
|
||||||
|
api_format: &str,
|
||||||
|
endpoint_base_url: &str,
|
||||||
|
key_id: &str,
|
||||||
|
upstream_api_key: &str,
|
||||||
|
proxy: Option<serde_json::Value>,
|
||||||
) -> Arc<InMemoryProviderCatalogReadRepository> {
|
) -> Arc<InMemoryProviderCatalogReadRepository> {
|
||||||
fn seal_bound_credential(
|
fn seal_bound_credential(
|
||||||
provider_id: &str,
|
provider_id: &str,
|
||||||
@@ -117,7 +151,7 @@ pub(super) fn video_provider_catalog_repository(
|
|||||||
false,
|
false,
|
||||||
None,
|
None,
|
||||||
Some(2),
|
Some(2),
|
||||||
None,
|
proxy,
|
||||||
Some(20.0),
|
Some(20.0),
|
||||||
None,
|
None,
|
||||||
None,
|
None,
|
||||||
|
|||||||
@@ -13,7 +13,8 @@ use serde_json::json;
|
|||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
build_state_with_execution_runtime_override, start_server, video_provider_catalog_repository,
|
build_state_with_execution_runtime_override, start_server, video_provider_catalog_repository,
|
||||||
AppState, VideoTaskTruthSourceMode,
|
video_provider_catalog_repository_with_proxy, video_proxy_node_repository_at_url, AppState,
|
||||||
|
VideoTaskTruthSourceMode,
|
||||||
};
|
};
|
||||||
|
|
||||||
fn sample_due_openai_task(upstream_base_url: &str) -> UpsertVideoTask {
|
fn sample_due_openai_task(upstream_base_url: &str) -> UpsertVideoTask {
|
||||||
@@ -279,13 +280,13 @@ async fn gateway_background_video_task_poller_refreshes_due_openai_task_from_rep
|
|||||||
);
|
);
|
||||||
|
|
||||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||||
let upstream_api_root = format!("{upstream_url}/v1");
|
let upstream_api_root = "http://video-provider.invalid/v1".to_string();
|
||||||
let repository = Arc::new(InMemoryVideoTaskRepository::default());
|
let repository = Arc::new(InMemoryVideoTaskRepository::default());
|
||||||
repository
|
repository
|
||||||
.upsert(sample_due_openai_task(&upstream_api_root))
|
.upsert(sample_due_openai_task(&upstream_api_root))
|
||||||
.await
|
.await
|
||||||
.expect("task upsert should succeed");
|
.expect("task upsert should succeed");
|
||||||
let provider_catalog_repository = video_provider_catalog_repository(
|
let provider_catalog_repository = video_provider_catalog_repository_with_proxy(
|
||||||
"provider-openai-video-local-1",
|
"provider-openai-video-local-1",
|
||||||
"openai",
|
"openai",
|
||||||
"endpoint-openai-video-local-1",
|
"endpoint-openai-video-local-1",
|
||||||
@@ -293,6 +294,7 @@ async fn gateway_background_video_task_poller_refreshes_due_openai_task_from_rep
|
|||||||
&upstream_api_root,
|
&upstream_api_root,
|
||||||
"key-openai-video-local-1",
|
"key-openai-video-local-1",
|
||||||
"sk-upstream-openai-video",
|
"sk-upstream-openai-video",
|
||||||
|
Some(json!({"enabled":true,"node_id":"poller-video-proxy"})),
|
||||||
);
|
);
|
||||||
|
|
||||||
let gateway_state = AppState::new()
|
let gateway_state = AppState::new()
|
||||||
@@ -302,7 +304,7 @@ async fn gateway_background_video_task_poller_refreshes_due_openai_task_from_rep
|
|||||||
Arc::clone(&repository),
|
Arc::clone(&repository),
|
||||||
provider_catalog_repository,
|
provider_catalog_repository,
|
||||||
DEVELOPMENT_ENCRYPTION_KEY,
|
DEVELOPMENT_ENCRYPTION_KEY,
|
||||||
),
|
).attach_proxy_node_repository_for_tests(video_proxy_node_repository_at_url(["poller-video-proxy"], &upstream_url)),
|
||||||
)
|
)
|
||||||
.with_video_task_truth_source_mode(VideoTaskTruthSourceMode::RustAuthoritative)
|
.with_video_task_truth_source_mode(VideoTaskTruthSourceMode::RustAuthoritative)
|
||||||
.with_video_task_poller_config(std::time::Duration::from_millis(25), 8);
|
.with_video_task_poller_config(std::time::Duration::from_millis(25), 8);
|
||||||
|
|||||||
@@ -14,8 +14,7 @@ use crate::constants::{
|
|||||||
use super::{build_router, start_server};
|
use super::{build_router, start_server};
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn gateway_locally_denies_video_control_sync_even_with_opt_in_headers_when_execution_runtime_missing(
|
async fn gateway_hides_video_task_from_unauthenticated_caller_with_opt_in_headers() {
|
||||||
) {
|
|
||||||
let execute_hits = Arc::new(Mutex::new(0usize));
|
let execute_hits = Arc::new(Mutex::new(0usize));
|
||||||
let execute_hits_clone = Arc::clone(&execute_hits);
|
let execute_hits_clone = Arc::clone(&execute_hits);
|
||||||
let public_hits = Arc::new(Mutex::new(0usize));
|
let public_hits = Arc::new(Mutex::new(0usize));
|
||||||
@@ -66,13 +65,9 @@ async fn gateway_locally_denies_video_control_sync_even_with_opt_in_headers_when
|
|||||||
.await
|
.await
|
||||||
.expect("request should succeed");
|
.expect("request should succeed");
|
||||||
|
|
||||||
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE);
|
assert_eq!(response.status(), StatusCode::NOT_FOUND);
|
||||||
let payload: serde_json::Value = response.json().await.expect("body should parse");
|
let payload: serde_json::Value = response.json().await.expect("body should parse");
|
||||||
assert_eq!(payload["error"]["type"], "http_error");
|
assert_eq!(payload, crate::video_tasks::not_found_body());
|
||||||
assert_eq!(
|
|
||||||
payload["error"]["message"],
|
|
||||||
"当前 OpenAI Video 请求无法在本地执行:没有匹配到可用的执行路径"
|
|
||||||
);
|
|
||||||
assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0);
|
assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0);
|
||||||
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
|
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
|
||||||
|
|
||||||
@@ -81,8 +76,7 @@ async fn gateway_locally_denies_video_control_sync_even_with_opt_in_headers_when
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn gateway_locally_denies_video_control_sync_without_opt_in_header_when_execution_runtime_missing(
|
async fn gateway_hides_video_task_without_calling_public_or_control_upstream() {
|
||||||
) {
|
|
||||||
let execute_hits = Arc::new(Mutex::new(0usize));
|
let execute_hits = Arc::new(Mutex::new(0usize));
|
||||||
let execute_hits_clone = Arc::clone(&execute_hits);
|
let execute_hits_clone = Arc::clone(&execute_hits);
|
||||||
let public_hits = Arc::new(Mutex::new(0usize));
|
let public_hits = Arc::new(Mutex::new(0usize));
|
||||||
@@ -142,13 +136,9 @@ async fn gateway_locally_denies_video_control_sync_without_opt_in_header_when_ex
|
|||||||
.await
|
.await
|
||||||
.expect("request should succeed");
|
.expect("request should succeed");
|
||||||
|
|
||||||
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE);
|
assert_eq!(response.status(), StatusCode::NOT_FOUND);
|
||||||
let payload: serde_json::Value = response.json().await.expect("body should parse");
|
let payload: serde_json::Value = response.json().await.expect("body should parse");
|
||||||
assert_eq!(payload["error"]["type"], "http_error");
|
assert_eq!(payload, crate::video_tasks::not_found_body());
|
||||||
assert_eq!(
|
|
||||||
payload["error"]["message"],
|
|
||||||
"当前 OpenAI Video 请求无法在本地执行:没有匹配到可用的执行路径"
|
|
||||||
);
|
|
||||||
assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0);
|
assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0);
|
||||||
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
|
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
@@ -165,7 +155,7 @@ async fn gateway_locally_denies_video_control_sync_without_opt_in_header_when_ex
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn gateway_skips_video_get_control_sync_without_opt_in_header() {
|
async fn gateway_hides_video_task_from_unauthenticated_caller_without_opt_in_headers() {
|
||||||
let execute_hits = Arc::new(Mutex::new(0usize));
|
let execute_hits = Arc::new(Mutex::new(0usize));
|
||||||
let execute_hits_clone = Arc::clone(&execute_hits);
|
let execute_hits_clone = Arc::clone(&execute_hits);
|
||||||
let public_hits = Arc::new(Mutex::new(0usize));
|
let public_hits = Arc::new(Mutex::new(0usize));
|
||||||
@@ -211,13 +201,9 @@ async fn gateway_skips_video_get_control_sync_without_opt_in_header() {
|
|||||||
.await
|
.await
|
||||||
.expect("request should succeed");
|
.expect("request should succeed");
|
||||||
|
|
||||||
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE);
|
assert_eq!(response.status(), StatusCode::NOT_FOUND);
|
||||||
let payload: serde_json::Value = response.json().await.expect("body should parse");
|
let payload: serde_json::Value = response.json().await.expect("body should parse");
|
||||||
assert_eq!(payload["error"]["type"], "http_error");
|
assert_eq!(payload, crate::video_tasks::not_found_body());
|
||||||
assert_eq!(
|
|
||||||
payload["error"]["message"],
|
|
||||||
"当前 OpenAI Video 请求无法在本地执行:没有匹配到可用的执行路径"
|
|
||||||
);
|
|
||||||
assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0);
|
assert_eq!(*execute_hits.lock().expect("mutex should lock"), 0);
|
||||||
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
|
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,427 @@
|
|||||||
|
use super::*;
|
||||||
|
use aether_data::repository::auth::{
|
||||||
|
InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot,
|
||||||
|
};
|
||||||
|
use aether_data::repository::candidate_selection::InMemoryMinimalCandidateSelectionReadRepository;
|
||||||
|
use aether_data::repository::candidates::InMemoryRequestCandidateRepository;
|
||||||
|
use aether_data_contracts::repository::candidate_selection::{
|
||||||
|
StoredMinimalCandidateSelectionRow, StoredProviderModelMapping,
|
||||||
|
};
|
||||||
|
use sha2::{Digest, Sha256};
|
||||||
|
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||||
|
|
||||||
|
fn sample_auth_snapshot(api_key_id: &str, user_id: &str) -> StoredAuthApiKeySnapshot {
|
||||||
|
StoredAuthApiKeySnapshot::new(
|
||||||
|
user_id.to_string(),
|
||||||
|
"video-user".to_string(),
|
||||||
|
Some("[email protected]".to_string()),
|
||||||
|
"user".to_string(),
|
||||||
|
"local".to_string(),
|
||||||
|
true,
|
||||||
|
false,
|
||||||
|
Some(json!(["openai"])),
|
||||||
|
Some(json!(["openai:video"])),
|
||||||
|
Some(json!(["video-model"])),
|
||||||
|
api_key_id.to_string(),
|
||||||
|
Some("default".to_string()),
|
||||||
|
true,
|
||||||
|
false,
|
||||||
|
false,
|
||||||
|
Some(60),
|
||||||
|
Some(5),
|
||||||
|
Some(4_102_444_800),
|
||||||
|
Some(json!(["openai"])),
|
||||||
|
Some(json!(["openai:video"])),
|
||||||
|
Some(json!(["video-model"])),
|
||||||
|
)
|
||||||
|
.expect("auth snapshot should build")
|
||||||
|
}
|
||||||
|
|
||||||
|
fn sample_candidate_row() -> StoredMinimalCandidateSelectionRow {
|
||||||
|
StoredMinimalCandidateSelectionRow {
|
||||||
|
provider_id: "provider-openai-video-local-1".to_string(),
|
||||||
|
provider_name: "openai".to_string(),
|
||||||
|
provider_type: "xai".to_string(),
|
||||||
|
provider_priority: 10,
|
||||||
|
provider_is_active: true,
|
||||||
|
endpoint_id: "endpoint-openai-video-local-1".to_string(),
|
||||||
|
endpoint_api_format: "openai:video".to_string(),
|
||||||
|
endpoint_api_family: Some("openai".to_string()),
|
||||||
|
endpoint_kind: Some("video".to_string()),
|
||||||
|
endpoint_is_active: true,
|
||||||
|
key_id: "key-openai-video-local-1".to_string(),
|
||||||
|
key_name: "prod".to_string(),
|
||||||
|
key_auth_type: "api_key".to_string(),
|
||||||
|
key_is_active: true,
|
||||||
|
key_api_formats: Some(vec!["openai:video".to_string()]),
|
||||||
|
key_allowed_models: None,
|
||||||
|
key_capabilities: None,
|
||||||
|
key_internal_priority: 5,
|
||||||
|
key_global_priority_by_format: Some(json!({"openai:video": 1})),
|
||||||
|
model_id: "model-openai-video-local-1".to_string(),
|
||||||
|
global_model_id: "global-model-openai-video-local-1".to_string(),
|
||||||
|
global_model_name: "video-model".to_string(),
|
||||||
|
global_model_mappings: None,
|
||||||
|
global_model_supports_streaming: Some(false),
|
||||||
|
model_provider_model_name: "grok-imagine-video".to_string(),
|
||||||
|
model_provider_model_mappings: Some(vec![StoredProviderModelMapping {
|
||||||
|
name: "grok-imagine-video".to_string(),
|
||||||
|
priority: 1,
|
||||||
|
api_formats: Some(vec!["openai:video".to_string()]),
|
||||||
|
endpoint_ids: None,
|
||||||
|
operations: None,
|
||||||
|
}]),
|
||||||
|
model_supports_streaming: Some(false),
|
||||||
|
model_is_active: true,
|
||||||
|
model_is_available: true,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn xai_video_native_and_compatibility_http_lifecycle() {
|
||||||
|
Box::pin(assert_xai_video_http_lifecycle(Arc::new(
|
||||||
|
InMemoryVideoTaskRepository::default(),
|
||||||
|
)))
|
||||||
|
.await;
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn xai_video_native_and_compatibility_http_lifecycle_postgres() {
|
||||||
|
let configured_database_url = std::env::var("AETHER_TEST_DATABASE_URL").ok();
|
||||||
|
let managed_database = if configured_database_url.is_none() {
|
||||||
|
Some(
|
||||||
|
aether_testkit::ManagedPostgresServer::start()
|
||||||
|
.await
|
||||||
|
.expect("temporary PostgreSQL should start"),
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
let database_url = configured_database_url.unwrap_or_else(|| {
|
||||||
|
managed_database
|
||||||
|
.as_ref()
|
||||||
|
.expect("managed test database should exist")
|
||||||
|
.database_url()
|
||||||
|
.to_string()
|
||||||
|
});
|
||||||
|
let pool = sqlx::postgres::PgPoolOptions::new()
|
||||||
|
.max_connections(1)
|
||||||
|
.connect(&database_url)
|
||||||
|
.await
|
||||||
|
.expect("test database should connect");
|
||||||
|
aether_data::driver::postgres::run_migrations(&pool)
|
||||||
|
.await
|
||||||
|
.expect("test database should migrate");
|
||||||
|
// Preserve the production column constraints and unique indexes while isolating test rows.
|
||||||
|
sqlx::query("CREATE TEMP TABLE video_tasks (LIKE public.video_tasks INCLUDING ALL)")
|
||||||
|
.execute(&pool)
|
||||||
|
.await
|
||||||
|
.expect("isolated video task table should be created");
|
||||||
|
let repository =
|
||||||
|
Arc::new(aether_data::repository::video_tasks::SqlxVideoTaskRepository::new(pool.clone()));
|
||||||
|
Box::pin(assert_xai_video_http_lifecycle(repository)).await;
|
||||||
|
pool.close().await;
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn assert_xai_video_http_lifecycle<T>(repository: Arc<T>)
|
||||||
|
where
|
||||||
|
T: aether_data_contracts::repository::video_tasks::VideoTaskRepository + 'static,
|
||||||
|
{
|
||||||
|
let static_dir = std::env::temp_dir().join(format!(
|
||||||
|
"aether-xai-video-static-{}",
|
||||||
|
std::time::SystemTime::now()
|
||||||
|
.duration_since(std::time::UNIX_EPOCH)
|
||||||
|
.unwrap()
|
||||||
|
.as_nanos()
|
||||||
|
));
|
||||||
|
std::fs::create_dir_all(&static_dir).unwrap();
|
||||||
|
std::fs::write(
|
||||||
|
static_dir.join("index.html"),
|
||||||
|
"<html>Aether test frontend</html>",
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
let seen = Arc::new(Mutex::new(Vec::<serde_json::Value>::new()));
|
||||||
|
let calls = Arc::new(AtomicUsize::new(0));
|
||||||
|
// Exercise the real HTTP executor, including production method gates, instead of
|
||||||
|
// the test execution-runtime override that used to hide rejected GET requests.
|
||||||
|
let video_url = Arc::new(Mutex::new(String::new()));
|
||||||
|
let runtime = Router::new()
|
||||||
|
.route("/v1/videos/{operation}", any({
|
||||||
|
let seen = seen.clone();
|
||||||
|
let calls = calls.clone();
|
||||||
|
let video_url = video_url.clone();
|
||||||
|
move |request: Request| {
|
||||||
|
let seen = seen.clone();
|
||||||
|
let calls = calls.clone();
|
||||||
|
let video_url = video_url.clone();
|
||||||
|
async move {
|
||||||
|
let (parts, body) = request.into_parts();
|
||||||
|
assert_eq!(parts.headers["authorization"], "Bearer upstream-video-key");
|
||||||
|
let bytes = to_bytes(body, usize::MAX).await.unwrap();
|
||||||
|
let body: serde_json::Value = serde_json::from_slice(&bytes).unwrap_or(json!(null));
|
||||||
|
seen.lock().unwrap().push(json!({
|
||||||
|
"method": parts.method.as_str(),
|
||||||
|
"url": parts.uri.path(),
|
||||||
|
"body": {"json_body": body}
|
||||||
|
}));
|
||||||
|
let response = if parts.method == http::Method::POST {
|
||||||
|
json!({"request_id":"upstream-video-id", "provider_extension":{"accepted":true}})
|
||||||
|
} else {
|
||||||
|
assert_eq!(parts.uri.path(), "/v1/videos/upstream-video-id");
|
||||||
|
if calls.fetch_add(1, Ordering::SeqCst) == 0 {
|
||||||
|
json!({"status":"pending"})
|
||||||
|
} else {
|
||||||
|
json!({"status":"done", "model":"grok-imagine-video", "video":{"url":video_url.lock().unwrap().clone(), "duration":6, "respect_moderation":true}, "provider_extension":"preserved"})
|
||||||
|
}
|
||||||
|
};
|
||||||
|
Json(response)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
.route("/test.mp4", any(|request: Request| async move {
|
||||||
|
assert!(request.headers().get("authorization").is_none());
|
||||||
|
assert!(request.headers().get("x-xai-token-auth").is_none());
|
||||||
|
([("content-type", "video/mp4")], "test-video-bytes")
|
||||||
|
}));
|
||||||
|
let (runtime_url, runtime_handle) = start_server(runtime).await;
|
||||||
|
let expected_video_url = format!("{runtime_url}/test.mp4");
|
||||||
|
*video_url.lock().unwrap() = expected_video_url.clone();
|
||||||
|
let state_factory = || {
|
||||||
|
let auth = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![
|
||||||
|
(
|
||||||
|
Some(format!("{:x}", Sha256::digest(b"owner-key"))),
|
||||||
|
sample_auth_snapshot("owner-api-key", "owner"),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
Some(format!("{:x}", Sha256::digest(b"foreign-key"))),
|
||||||
|
sample_auth_snapshot("foreign-api-key", "foreign"),
|
||||||
|
),
|
||||||
|
]));
|
||||||
|
let candidates = Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
|
||||||
|
sample_candidate_row(),
|
||||||
|
]));
|
||||||
|
let catalog = video_provider_catalog_repository_with_proxy(
|
||||||
|
"provider-openai-video-local-1",
|
||||||
|
"xai",
|
||||||
|
"endpoint-openai-video-local-1",
|
||||||
|
"openai:video",
|
||||||
|
"http://video-provider.invalid/v1",
|
||||||
|
"key-openai-video-local-1",
|
||||||
|
"upstream-video-key",
|
||||||
|
Some(json!({"enabled":true,"node_id":"video-proxy"})),
|
||||||
|
);
|
||||||
|
AppState::new().expect("gateway should build").with_video_task_truth_source_mode(VideoTaskTruthSourceMode::RustAuthoritative).with_data_state_for_tests(
|
||||||
|
crate::data::GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
|
||||||
|
auth, candidates, catalog, Arc::new(InMemoryRequestCandidateRepository::default()), DEVELOPMENT_ENCRYPTION_KEY
|
||||||
|
).attach_video_task_repository_for_tests(repository.clone())
|
||||||
|
.attach_proxy_node_repository_for_tests(video_proxy_node_repository_at_url(["video-proxy"], &runtime_url))
|
||||||
|
)
|
||||||
|
};
|
||||||
|
let router_factory =
|
||||||
|
|| crate::attach_static_frontend(build_router_with_state(state_factory()), &static_dir);
|
||||||
|
let (gateway_url, gateway_handle) = start_server(router_factory()).await;
|
||||||
|
let client = reqwest::Client::new();
|
||||||
|
assert_eq!(
|
||||||
|
client
|
||||||
|
.get(&gateway_url)
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.text()
|
||||||
|
.await
|
||||||
|
.unwrap(),
|
||||||
|
"<html>Aether test frontend</html>"
|
||||||
|
);
|
||||||
|
for (path, native) in [
|
||||||
|
("/v1/videos/generations", true),
|
||||||
|
("/v1/videos", true),
|
||||||
|
("/v1/videos/edits", true),
|
||||||
|
("/v1/videos/extensions", true),
|
||||||
|
("/openai/v1/videos", false),
|
||||||
|
] {
|
||||||
|
calls.store(0, Ordering::SeqCst);
|
||||||
|
let body = if native {
|
||||||
|
json!({"model":"video-model","prompt":"A cat","duration":6,"aspect_ratio":"1:1","video":{"url":"https://example.com/input.mp4"},"future_option":true})
|
||||||
|
} else {
|
||||||
|
json!({"model":"video-model","prompt":"A cat","seconds":"6","size":"1280x720"})
|
||||||
|
};
|
||||||
|
let response = client
|
||||||
|
.post(format!("{gateway_url}{path}"))
|
||||||
|
.bearer_auth("owner-key")
|
||||||
|
.json(&body)
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let status = response.status();
|
||||||
|
let result: serde_json::Value = response.json().await.unwrap();
|
||||||
|
assert_eq!(status, StatusCode::OK, "{path}: {result}");
|
||||||
|
let id = result[if native { "request_id" } else { "id" }]
|
||||||
|
.as_str()
|
||||||
|
.unwrap();
|
||||||
|
assert_ne!(id, "upstream-video-id");
|
||||||
|
if native {
|
||||||
|
assert!(result.get("id").is_none());
|
||||||
|
assert_eq!(result["provider_extension"]["accepted"], true);
|
||||||
|
} else {
|
||||||
|
assert_eq!(result["status"], "queued");
|
||||||
|
}
|
||||||
|
let request = seen.lock().unwrap().last().unwrap().clone();
|
||||||
|
let suffix = if path.ends_with("/edits") {
|
||||||
|
"edits"
|
||||||
|
} else if path.ends_with("/extensions") {
|
||||||
|
"extensions"
|
||||||
|
} else {
|
||||||
|
"generations"
|
||||||
|
};
|
||||||
|
assert_eq!(request["url"], format!("/v1/videos/{suffix}"));
|
||||||
|
assert_eq!(request["body"]["json_body"]["model"], "grok-imagine-video");
|
||||||
|
assert_eq!(request["body"]["json_body"]["duration"], 6);
|
||||||
|
if native {
|
||||||
|
assert_eq!(request["body"]["json_body"]["future_option"], true);
|
||||||
|
} else {
|
||||||
|
assert_eq!(request["body"]["json_body"]["aspect_ratio"], "16:9");
|
||||||
|
assert_eq!(request["body"]["json_body"]["resolution"], "720p");
|
||||||
|
assert!(request["body"]["json_body"].get("seconds").is_none());
|
||||||
|
assert!(request["body"]["json_body"].get("size").is_none());
|
||||||
|
}
|
||||||
|
let query = format!(
|
||||||
|
"{gateway_url}{}/{id}",
|
||||||
|
if native {
|
||||||
|
"/v1/videos"
|
||||||
|
} else {
|
||||||
|
"/openai/v1/videos"
|
||||||
|
}
|
||||||
|
);
|
||||||
|
let before = seen.lock().unwrap().len();
|
||||||
|
let denied = client
|
||||||
|
.get(&query)
|
||||||
|
.bearer_auth("foreign-key")
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(denied.status(), StatusCode::NOT_FOUND);
|
||||||
|
assert_eq!(seen.lock().unwrap().len(), before);
|
||||||
|
let denied_content = client
|
||||||
|
.get(format!("{gateway_url}/openai/v1/videos/{id}/content"))
|
||||||
|
.bearer_auth("foreign-key")
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(denied_content.status(), StatusCode::NOT_FOUND);
|
||||||
|
assert_eq!(seen.lock().unwrap().len(), before);
|
||||||
|
let pending: serde_json::Value = client
|
||||||
|
.get(&query)
|
||||||
|
.bearer_auth("owner-key")
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.json()
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
pending["status"],
|
||||||
|
if native { "pending" } else { "queued" },
|
||||||
|
"{path}: {pending}"
|
||||||
|
);
|
||||||
|
let done: serde_json::Value = client
|
||||||
|
.get(&query)
|
||||||
|
.bearer_auth("owner-key")
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.json()
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(done["status"], if native { "done" } else { "completed" });
|
||||||
|
if native {
|
||||||
|
assert_eq!(done["video"]["respect_moderation"], true);
|
||||||
|
assert_eq!(done["provider_extension"], "preserved");
|
||||||
|
} else {
|
||||||
|
assert_eq!(done["video_url"], expected_video_url);
|
||||||
|
}
|
||||||
|
let stored = repository
|
||||||
|
.find(VideoTaskLookupKey::Id(id))
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
stored.client_api_format.as_deref(),
|
||||||
|
Some(if native { "xai:video" } else { "openai:video" })
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
stored.external_task_id.as_deref(),
|
||||||
|
Some("upstream-video-id")
|
||||||
|
);
|
||||||
|
assert!(stored.request_metadata.is_none());
|
||||||
|
assert!(stored.original_request_body.is_none());
|
||||||
|
// A new gateway instance must reconstruct the pinned provider/credential and protocol.
|
||||||
|
let (restart_url, restart_handle) = start_server(router_factory()).await;
|
||||||
|
let restored: serde_json::Value = client
|
||||||
|
.get(format!(
|
||||||
|
"{restart_url}{}/{id}",
|
||||||
|
if native {
|
||||||
|
"/v1/videos"
|
||||||
|
} else {
|
||||||
|
"/openai/v1/videos"
|
||||||
|
}
|
||||||
|
))
|
||||||
|
.bearer_auth("owner-key")
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.json()
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(restored["status"], done["status"]);
|
||||||
|
if native {
|
||||||
|
assert_eq!(restored["video"]["respect_moderation"], true);
|
||||||
|
}
|
||||||
|
let compat: serde_json::Value = client
|
||||||
|
.get(format!("{restart_url}/openai/v1/videos/{id}"))
|
||||||
|
.bearer_auth("owner-key")
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.json()
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(compat["status"], "completed");
|
||||||
|
assert_eq!(compat["video_url"], expected_video_url);
|
||||||
|
let native_view: serde_json::Value = client
|
||||||
|
.get(format!("{restart_url}/v1/videos/{id}"))
|
||||||
|
.bearer_auth("owner-key")
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.json()
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(native_view["status"], "done");
|
||||||
|
assert_eq!(native_view["video"]["respect_moderation"], true);
|
||||||
|
for prefix in ["/v1/videos", "/openai/v1/videos"] {
|
||||||
|
let content = client
|
||||||
|
.get(format!("{restart_url}{prefix}/{id}/content"))
|
||||||
|
.bearer_auth("owner-key")
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(content.status(), StatusCode::OK);
|
||||||
|
assert_eq!(content.headers()["content-type"], "video/mp4");
|
||||||
|
assert_eq!(content.bytes().await.unwrap(), "test-video-bytes");
|
||||||
|
}
|
||||||
|
restart_handle.abort();
|
||||||
|
}
|
||||||
|
let before = seen.lock().unwrap().len();
|
||||||
|
let bad = client
|
||||||
|
.post(format!("{gateway_url}/openai/v1/videos"))
|
||||||
|
.bearer_auth("owner-key")
|
||||||
|
.json(&json!({"model":"video-model","prompt":"cat","seconds":"wrong"}))
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(bad.status(), StatusCode::BAD_REQUEST);
|
||||||
|
assert_eq!(seen.lock().unwrap().len(), before);
|
||||||
|
gateway_handle.abort();
|
||||||
|
runtime_handle.abort();
|
||||||
|
std::fs::remove_dir_all(&static_dir).unwrap();
|
||||||
|
}
|
||||||
@@ -11,6 +11,9 @@ use super::{
|
|||||||
fn rust_authoritative_service_builds_openai_cancel_follow_up_plan() {
|
fn rust_authoritative_service_builds_openai_cancel_follow_up_plan() {
|
||||||
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
||||||
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: None,
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: false,
|
||||||
local_task_id: "task-local-123".to_string(),
|
local_task_id: "task-local-123".to_string(),
|
||||||
upstream_task_id: "ext-video-task-123".to_string(),
|
upstream_task_id: "ext-video-task-123".to_string(),
|
||||||
created_at_unix_ms: 1712345678,
|
created_at_unix_ms: 1712345678,
|
||||||
@@ -92,6 +95,9 @@ fn rust_authoritative_service_builds_openai_cancel_follow_up_plan() {
|
|||||||
fn rust_authoritative_service_builds_openai_remix_follow_up_plan() {
|
fn rust_authoritative_service_builds_openai_remix_follow_up_plan() {
|
||||||
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
||||||
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: None,
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: false,
|
||||||
local_task_id: "task-local-123".to_string(),
|
local_task_id: "task-local-123".to_string(),
|
||||||
upstream_task_id: "ext-video-task-123".to_string(),
|
upstream_task_id: "ext-video-task-123".to_string(),
|
||||||
created_at_unix_ms: 1712345678,
|
created_at_unix_ms: 1712345678,
|
||||||
@@ -177,6 +183,9 @@ fn rust_authoritative_service_builds_openai_remix_follow_up_plan() {
|
|||||||
fn rust_authoritative_service_builds_openai_delete_follow_up_plan() {
|
fn rust_authoritative_service_builds_openai_delete_follow_up_plan() {
|
||||||
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
||||||
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: None,
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: false,
|
||||||
local_task_id: "task-local-123".to_string(),
|
local_task_id: "task-local-123".to_string(),
|
||||||
upstream_task_id: "ext-video-task-123".to_string(),
|
upstream_task_id: "ext-video-task-123".to_string(),
|
||||||
created_at_unix_ms: 1712345678,
|
created_at_unix_ms: 1712345678,
|
||||||
@@ -332,6 +341,9 @@ fn rust_authoritative_service_builds_gemini_cancel_follow_up_plan() {
|
|||||||
fn rust_authoritative_service_builds_openai_read_refresh_plan() {
|
fn rust_authoritative_service_builds_openai_read_refresh_plan() {
|
||||||
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
||||||
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: None,
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: false,
|
||||||
local_task_id: "task-local-123".to_string(),
|
local_task_id: "task-local-123".to_string(),
|
||||||
upstream_task_id: "ext-video-task-123".to_string(),
|
upstream_task_id: "ext-video-task-123".to_string(),
|
||||||
created_at_unix_ms: 1712345678,
|
created_at_unix_ms: 1712345678,
|
||||||
@@ -407,6 +419,9 @@ fn rust_authoritative_service_builds_gemini_read_refresh_plan() {
|
|||||||
fn rust_authoritative_service_builds_poll_refresh_batch_for_active_tasks_only() {
|
fn rust_authoritative_service_builds_poll_refresh_batch_for_active_tasks_only() {
|
||||||
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
||||||
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: None,
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: false,
|
||||||
local_task_id: "task-active-123".to_string(),
|
local_task_id: "task-active-123".to_string(),
|
||||||
upstream_task_id: "ext-video-task-123".to_string(),
|
upstream_task_id: "ext-video-task-123".to_string(),
|
||||||
created_at_unix_ms: 1712345678,
|
created_at_unix_ms: 1712345678,
|
||||||
@@ -428,6 +443,9 @@ fn rust_authoritative_service_builds_poll_refresh_batch_for_active_tasks_only()
|
|||||||
transport: sample_transport("https://api.openai.example", "openai:video"),
|
transport: sample_transport("https://api.openai.example", "openai:video"),
|
||||||
}));
|
}));
|
||||||
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: None,
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: false,
|
||||||
local_task_id: "task-completed-123".to_string(),
|
local_task_id: "task-completed-123".to_string(),
|
||||||
upstream_task_id: "ext-video-task-999".to_string(),
|
upstream_task_id: "ext-video-task-999".to_string(),
|
||||||
created_at_unix_ms: 1712345678,
|
created_at_unix_ms: 1712345678,
|
||||||
@@ -471,6 +489,9 @@ fn file_video_task_store_persists_snapshots_across_service_rebuilds() {
|
|||||||
)
|
)
|
||||||
.expect("file-backed service should build");
|
.expect("file-backed service should build");
|
||||||
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: None,
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: false,
|
||||||
local_task_id: "task-file-123".to_string(),
|
local_task_id: "task-file-123".to_string(),
|
||||||
upstream_task_id: "ext-video-task-123".to_string(),
|
upstream_task_id: "ext-video-task-123".to_string(),
|
||||||
created_at_unix_ms: 1712345678,
|
created_at_unix_ms: 1712345678,
|
||||||
|
|||||||
@@ -10,6 +10,9 @@ use super::{
|
|||||||
fn rust_authoritative_service_projects_openai_status_into_local_read_response() {
|
fn rust_authoritative_service_projects_openai_status_into_local_read_response() {
|
||||||
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
||||||
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: None,
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: false,
|
||||||
local_task_id: "task-local-123".to_string(),
|
local_task_id: "task-local-123".to_string(),
|
||||||
upstream_task_id: "ext-video-task-123".to_string(),
|
upstream_task_id: "ext-video-task-123".to_string(),
|
||||||
created_at_unix_ms: 1712345678,
|
created_at_unix_ms: 1712345678,
|
||||||
@@ -93,6 +96,9 @@ fn rust_authoritative_service_projects_openai_status_into_local_read_response()
|
|||||||
fn rust_authoritative_service_builds_openai_content_stream_plan_from_direct_video_url() {
|
fn rust_authoritative_service_builds_openai_content_stream_plan_from_direct_video_url() {
|
||||||
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
||||||
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: None,
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: false,
|
||||||
local_task_id: "task-local-123".to_string(),
|
local_task_id: "task-local-123".to_string(),
|
||||||
upstream_task_id: "ext-video-task-123".to_string(),
|
upstream_task_id: "ext-video-task-123".to_string(),
|
||||||
created_at_unix_ms: 1712345678,
|
created_at_unix_ms: 1712345678,
|
||||||
@@ -159,6 +165,9 @@ fn rust_authoritative_service_builds_openai_content_stream_plan_from_direct_vide
|
|||||||
fn rust_authoritative_service_returns_processing_content_response_for_pending_openai_task() {
|
fn rust_authoritative_service_returns_processing_content_response_for_pending_openai_task() {
|
||||||
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
||||||
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: None,
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: false,
|
||||||
local_task_id: "task-local-123".to_string(),
|
local_task_id: "task-local-123".to_string(),
|
||||||
upstream_task_id: "ext-video-task-123".to_string(),
|
upstream_task_id: "ext-video-task-123".to_string(),
|
||||||
created_at_unix_ms: 1712345678,
|
created_at_unix_ms: 1712345678,
|
||||||
|
|||||||
@@ -218,6 +218,9 @@ fn rust_authoritative_video_truth_source_can_background_success_report() {
|
|||||||
fn rust_authoritative_service_reads_openai_task_from_local_registry() {
|
fn rust_authoritative_service_reads_openai_task_from_local_registry() {
|
||||||
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
||||||
let snapshot = LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
let snapshot = LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: None,
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: false,
|
||||||
local_task_id: "task-local-123".to_string(),
|
local_task_id: "task-local-123".to_string(),
|
||||||
upstream_task_id: "ext-video-task-123".to_string(),
|
upstream_task_id: "ext-video-task-123".to_string(),
|
||||||
created_at_unix_ms: 1712345678,
|
created_at_unix_ms: 1712345678,
|
||||||
@@ -266,6 +269,9 @@ fn rust_authoritative_service_reads_openai_task_from_local_registry() {
|
|||||||
fn rust_authoritative_service_applies_cancel_and_delete_mutations() {
|
fn rust_authoritative_service_applies_cancel_and_delete_mutations() {
|
||||||
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
||||||
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
service.record_snapshot(LocalVideoTaskSnapshot::OpenAi(OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: None,
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: false,
|
||||||
local_task_id: "task-local-123".to_string(),
|
local_task_id: "task-local-123".to_string(),
|
||||||
upstream_task_id: "ext-video-task-123".to_string(),
|
upstream_task_id: "ext-video-task-123".to_string(),
|
||||||
created_at_unix_ms: 1712345678,
|
created_at_unix_ms: 1712345678,
|
||||||
|
|||||||
@@ -3546,6 +3546,258 @@ pub fn parse_kiro_usage_response(
|
|||||||
Some(serde_json::Value::Object(result))
|
Some(serde_json::Value::Object(result))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn parse_xai_billing_response(
|
||||||
|
value: &serde_json::Value,
|
||||||
|
updated_at_unix_secs: u64,
|
||||||
|
) -> Option<serde_json::Value> {
|
||||||
|
let root = value.as_object()?;
|
||||||
|
let config = root
|
||||||
|
.get("config")
|
||||||
|
.and_then(serde_json::Value::as_object)
|
||||||
|
.unwrap_or(root);
|
||||||
|
|
||||||
|
let usage_percentage = coerce_json_f64_from_map(config, "creditUsagePercent")
|
||||||
|
.or_else(|| extract_xai_product_usage_percent(config));
|
||||||
|
let period = config.get("currentPeriod");
|
||||||
|
let period_type = period
|
||||||
|
.and_then(|value| value.get("type").or_else(|| value.get("periodType")))
|
||||||
|
.and_then(normalize_xai_period_type);
|
||||||
|
let next_reset_at = period
|
||||||
|
.and_then(|value| value.get("end"))
|
||||||
|
.and_then(parse_xai_timestamp)
|
||||||
|
.or_else(|| config.get("billingPeriodEnd").and_then(parse_xai_timestamp));
|
||||||
|
let monthly_limit =
|
||||||
|
coerce_xai_cents_dollars(config.get("monthlyLimit")).filter(|value| *value > 0.0);
|
||||||
|
let current_usage = if monthly_limit.is_some() {
|
||||||
|
coerce_xai_cents_dollars(config.get("used"))
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
let remaining = monthly_limit
|
||||||
|
.zip(current_usage)
|
||||||
|
.map(|(limit, used)| (limit - used).max(0.0));
|
||||||
|
let usage_percentage = usage_percentage.or_else(|| {
|
||||||
|
monthly_limit
|
||||||
|
.zip(current_usage)
|
||||||
|
.map(|(limit, used)| ((used / limit) * 100.0).clamp(0.0, 100.0))
|
||||||
|
});
|
||||||
|
let usage_percentage = match usage_percentage {
|
||||||
|
Some(value) => Some(value.clamp(0.0, 100.0)),
|
||||||
|
None if period_type.is_some() || next_reset_at.is_some() => Some(0.0),
|
||||||
|
None => None,
|
||||||
|
};
|
||||||
|
let prepaid_balance = coerce_xai_cents_dollars(config.get("prepaidBalance"));
|
||||||
|
let on_demand_cap = coerce_xai_cents_dollars(config.get("onDemandCap"));
|
||||||
|
let on_demand_used = coerce_xai_cents_dollars(config.get("onDemandUsed"));
|
||||||
|
let on_demand_enabled = coerce_json_bool_from_map(root, "onDemandEnabled")
|
||||||
|
.or_else(|| coerce_json_bool_from_map(config, "onDemandEnabled"));
|
||||||
|
let subscription_title = first_json_string_by_paths(
|
||||||
|
value,
|
||||||
|
&[
|
||||||
|
&["subscriptionTier"],
|
||||||
|
&["subscription_tier"],
|
||||||
|
&["config", "subscriptionTier"],
|
||||||
|
&["config", "subscription_title"],
|
||||||
|
],
|
||||||
|
);
|
||||||
|
|
||||||
|
if usage_percentage.is_none()
|
||||||
|
&& monthly_limit.is_none()
|
||||||
|
&& current_usage.is_none()
|
||||||
|
&& prepaid_balance.is_none()
|
||||||
|
&& on_demand_cap.is_none()
|
||||||
|
&& next_reset_at.is_none()
|
||||||
|
&& subscription_title.is_none()
|
||||||
|
{
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut result = serde_json::Map::new();
|
||||||
|
result.insert("updated_at".to_string(), json!(updated_at_unix_secs));
|
||||||
|
if let Some(value) = usage_percentage {
|
||||||
|
result.insert("usage_percentage".to_string(), json!(value));
|
||||||
|
}
|
||||||
|
if let Some(value) = monthly_limit {
|
||||||
|
result.insert("usage_limit".to_string(), json!(value));
|
||||||
|
}
|
||||||
|
if let Some(value) = current_usage {
|
||||||
|
result.insert("current_usage".to_string(), json!(value));
|
||||||
|
}
|
||||||
|
if let Some(value) = remaining {
|
||||||
|
result.insert("remaining".to_string(), json!(value));
|
||||||
|
}
|
||||||
|
if let Some(value) = next_reset_at {
|
||||||
|
result.insert("next_reset_at".to_string(), json!(value));
|
||||||
|
}
|
||||||
|
if let Some(value) = period_type {
|
||||||
|
result.insert("period_type".to_string(), json!(value));
|
||||||
|
}
|
||||||
|
if let Some(value) = prepaid_balance {
|
||||||
|
result.insert("prepaid_balance".to_string(), json!(value));
|
||||||
|
}
|
||||||
|
if let Some(value) = on_demand_cap {
|
||||||
|
result.insert("on_demand_cap".to_string(), json!(value));
|
||||||
|
}
|
||||||
|
if let Some(value) = on_demand_used {
|
||||||
|
result.insert("on_demand_used".to_string(), json!(value));
|
||||||
|
}
|
||||||
|
if let Some(value) = on_demand_enabled {
|
||||||
|
result.insert("on_demand_enabled".to_string(), json!(value));
|
||||||
|
}
|
||||||
|
if let Some(value) = subscription_title {
|
||||||
|
result.insert("subscription_title".to_string(), json!(value));
|
||||||
|
}
|
||||||
|
Some(serde_json::Value::Object(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn coerce_json_f64_from_map(
|
||||||
|
object: &serde_json::Map<String, serde_json::Value>,
|
||||||
|
key: &str,
|
||||||
|
) -> Option<f64> {
|
||||||
|
object.get(key).and_then(coerce_json_f64)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn coerce_json_bool_from_map(
|
||||||
|
object: &serde_json::Map<String, serde_json::Value>,
|
||||||
|
key: &str,
|
||||||
|
) -> Option<bool> {
|
||||||
|
object.get(key).and_then(coerce_json_bool)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn extract_xai_product_usage_percent(
|
||||||
|
config: &serde_json::Map<String, serde_json::Value>,
|
||||||
|
) -> Option<f64> {
|
||||||
|
let items = config.get("productUsage")?.as_array()?;
|
||||||
|
let grok_build = items.iter().find(|item| {
|
||||||
|
item.get("product")
|
||||||
|
.and_then(serde_json::Value::as_str)
|
||||||
|
.is_some_and(|product| product.eq_ignore_ascii_case("GrokBuild"))
|
||||||
|
});
|
||||||
|
grok_build
|
||||||
|
.or(items.first())
|
||||||
|
.and_then(|item| item.get("usagePercent").and_then(coerce_json_f64))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn coerce_xai_cents_dollars(value: Option<&serde_json::Value>) -> Option<f64> {
|
||||||
|
let value = value?;
|
||||||
|
let cents = match value {
|
||||||
|
serde_json::Value::Object(object) => object.get("val").and_then(coerce_json_f64)?,
|
||||||
|
other => coerce_json_f64(other)?,
|
||||||
|
};
|
||||||
|
Some(cents / 100.0)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn normalize_xai_period_type(value: &serde_json::Value) -> Option<String> {
|
||||||
|
let raw = value
|
||||||
|
.as_str()
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())?;
|
||||||
|
let lowered = raw.to_ascii_lowercase();
|
||||||
|
if lowered.contains("week") {
|
||||||
|
Some("weekly".to_string())
|
||||||
|
} else if lowered.contains("month") {
|
||||||
|
Some("monthly".to_string())
|
||||||
|
} else {
|
||||||
|
Some(raw.to_string())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn parse_xai_timestamp(value: &serde_json::Value) -> Option<u64> {
|
||||||
|
if let Some(value) = coerce_json_u64(value) {
|
||||||
|
return Some(if value > 1_000_000_000_000 {
|
||||||
|
value / 1000
|
||||||
|
} else {
|
||||||
|
value
|
||||||
|
});
|
||||||
|
}
|
||||||
|
let raw = value.as_str()?.trim();
|
||||||
|
if raw.is_empty() {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
chrono::DateTime::parse_from_rfc3339(raw)
|
||||||
|
.ok()
|
||||||
|
.and_then(|timestamp| u64::try_from(timestamp.timestamp()).ok())
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod xai_quota_tests {
|
||||||
|
use super::parse_xai_billing_response;
|
||||||
|
use serde_json::json;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn parse_xai_credits_percent_and_weekly_period() {
|
||||||
|
let metadata = parse_xai_billing_response(
|
||||||
|
&json!({
|
||||||
|
"config": {
|
||||||
|
"currentPeriod": {
|
||||||
|
"type": "USAGE_PERIOD_TYPE_WEEKLY",
|
||||||
|
"start": "2026-08-08T01:53:09.930537+00:00",
|
||||||
|
"end": "2026-08-15T01:53:09.930537+00:00"
|
||||||
|
},
|
||||||
|
"creditUsagePercent": 46.0,
|
||||||
|
"productUsage": [
|
||||||
|
{"product": "GrokBuild", "usagePercent": 41.0},
|
||||||
|
{"product": "GrokChat"}
|
||||||
|
],
|
||||||
|
"onDemandCap": {"val": 0},
|
||||||
|
"onDemandUsed": {"val": 0},
|
||||||
|
"prepaidBalance": {"val": 0}
|
||||||
|
},
|
||||||
|
"subscriptionTier": "SuperGrok"
|
||||||
|
}),
|
||||||
|
1_775_000_000,
|
||||||
|
)
|
||||||
|
.expect("credits payload should parse");
|
||||||
|
|
||||||
|
assert_eq!(metadata["usage_percentage"], json!(46.0));
|
||||||
|
assert_eq!(metadata["period_type"], json!("weekly"));
|
||||||
|
assert_eq!(metadata["next_reset_at"], json!(1_786_758_789u64));
|
||||||
|
assert_eq!(metadata["prepaid_balance"], json!(0.0));
|
||||||
|
assert_eq!(metadata["on_demand_cap"], json!(0.0));
|
||||||
|
assert_eq!(metadata["subscription_title"], json!("SuperGrok"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn parse_xai_omitted_percent_as_fresh_weekly_zero() {
|
||||||
|
let metadata = parse_xai_billing_response(
|
||||||
|
&json!({
|
||||||
|
"config": {
|
||||||
|
"currentPeriod": {
|
||||||
|
"type": "USAGE_PERIOD_TYPE_WEEKLY",
|
||||||
|
"end": "2026-08-15T01:53:09.930537+00:00"
|
||||||
|
},
|
||||||
|
"isUnifiedBillingUser": true
|
||||||
|
}
|
||||||
|
}),
|
||||||
|
1_775_000_000,
|
||||||
|
)
|
||||||
|
.expect("fresh weekly period should parse");
|
||||||
|
|
||||||
|
assert_eq!(metadata["usage_percentage"], json!(0.0));
|
||||||
|
assert_eq!(metadata["period_type"], json!("weekly"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn parse_xai_legacy_monthly_cents() {
|
||||||
|
let metadata = parse_xai_billing_response(
|
||||||
|
&json!({
|
||||||
|
"config": {
|
||||||
|
"monthlyLimit": {"val": 2500},
|
||||||
|
"used": {"val": 1000},
|
||||||
|
"billingPeriodEnd": "2026-09-01T00:00:00Z"
|
||||||
|
}
|
||||||
|
}),
|
||||||
|
1_775_000_000,
|
||||||
|
)
|
||||||
|
.expect("legacy monthly payload should parse");
|
||||||
|
|
||||||
|
assert_eq!(metadata["usage_limit"], json!(25.0));
|
||||||
|
assert_eq!(metadata["current_usage"], json!(10.0));
|
||||||
|
assert_eq!(metadata["remaining"], json!(15.0));
|
||||||
|
assert_eq!(metadata["usage_percentage"], json!(40.0));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
pub fn parse_windsurf_user_status_response(
|
pub fn parse_windsurf_user_status_response(
|
||||||
value: &serde_json::Value,
|
value: &serde_json::Value,
|
||||||
updated_at_unix_secs: u64,
|
updated_at_unix_secs: u64,
|
||||||
|
|||||||
@@ -285,6 +285,23 @@ pub fn enrich_admin_provider_oauth_auth_config(
|
|||||||
],
|
],
|
||||||
);
|
);
|
||||||
|
|
||||||
|
if provider_type.trim().eq_ignore_ascii_case("xai") {
|
||||||
|
auth_config.insert("auth_method".to_string(), json!("oauth"));
|
||||||
|
auth_config.insert("using_api".to_string(), json!(false));
|
||||||
|
if let Some(id_token) = ["id_token", "idToken"]
|
||||||
|
.iter()
|
||||||
|
.find_map(|field| json_non_empty_string(token_payload.get(field)))
|
||||||
|
{
|
||||||
|
auth_config
|
||||||
|
.entry("id_token".to_string())
|
||||||
|
.or_insert_with(|| json!(id_token.clone()));
|
||||||
|
if let Some(claims) = decode_jwt_claims(&id_token) {
|
||||||
|
merge_missing_auth_config_fields(auth_config, &claims, &["email", "sub"]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
if provider_type.trim().eq_ignore_ascii_case("claude_code") {
|
if provider_type.trim().eq_ignore_ascii_case("claude_code") {
|
||||||
if let Some(organization_uuid) = token_payload_object
|
if let Some(organization_uuid) = token_payload_object
|
||||||
.get("organization")
|
.get("organization")
|
||||||
@@ -554,6 +571,28 @@ mod tests {
|
|||||||
assert_eq!(auth_config.get("is_fedramp"), Some(&json!(true)));
|
assert_eq!(auth_config.get("is_fedramp"), Some(&json!(true)));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_enrichment_marks_oauth_and_extracts_id_token_identity() {
|
||||||
|
let id_token = sample_unsigned_jwt(json!({
|
||||||
|
"email": "[email protected]",
|
||||||
|
"sub": "user-xai-1",
|
||||||
|
}));
|
||||||
|
let token_payload = json!({
|
||||||
|
"access_token": "access-token",
|
||||||
|
"refresh_token": "refresh-token",
|
||||||
|
"id_token": id_token,
|
||||||
|
});
|
||||||
|
let mut auth_config = serde_json::Map::new();
|
||||||
|
|
||||||
|
enrich_admin_provider_oauth_auth_config("xai", &mut auth_config, &token_payload);
|
||||||
|
|
||||||
|
assert_eq!(auth_config.get("auth_method"), Some(&json!("oauth")));
|
||||||
|
assert_eq!(auth_config.get("using_api"), Some(&json!(false)));
|
||||||
|
assert_eq!(auth_config.get("email"), Some(&json!("[email protected]")));
|
||||||
|
assert_eq!(auth_config.get("sub"), Some(&json!("user-xai-1")));
|
||||||
|
assert_eq!(auth_config.get("id_token"), Some(&json!(id_token)));
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn decode_jwt_claims_rejects_oversized_payload_before_decode() {
|
fn decode_jwt_claims_rejects_oversized_payload_before_decode() {
|
||||||
let max_encoded_len = MAX_UNVERIFIED_JWT_CLAIMS_BYTES
|
let max_encoded_len = MAX_UNVERIFIED_JWT_CLAIMS_BYTES
|
||||||
|
|||||||
@@ -208,6 +208,10 @@ pub use crate::formats::{
|
|||||||
resolve_stream_spec as resolve_openai_responses_stream_spec,
|
resolve_stream_spec as resolve_openai_responses_stream_spec,
|
||||||
resolve_sync_spec as resolve_openai_responses_sync_spec, LocalOpenAiResponsesSpec,
|
resolve_sync_spec as resolve_openai_responses_sync_spec, LocalOpenAiResponsesSpec,
|
||||||
},
|
},
|
||||||
|
xai::{
|
||||||
|
apply_xai_upstream_payload_edits, apply_xai_upstream_payload_edits_with_client,
|
||||||
|
xai_supports_native_image_generation,
|
||||||
|
},
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
shared::{
|
shared::{
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ pub mod request;
|
|||||||
pub mod response;
|
pub mod response;
|
||||||
pub mod spec;
|
pub mod spec;
|
||||||
pub mod stream;
|
pub mod stream;
|
||||||
|
pub mod xai;
|
||||||
|
|
||||||
const TOOL_ERROR_PREFIX: &str = "[tool error]";
|
const TOOL_ERROR_PREFIX: &str = "[tool error]";
|
||||||
const AETHER_REASONING_ITEM_ID_PREFIX: &str = "rs_aether_";
|
const AETHER_REASONING_ITEM_ID_PREFIX: &str = "rs_aether_";
|
||||||
@@ -85,6 +86,8 @@ pub enum OpenAiResponsesReasoningReplayPolicy {
|
|||||||
#[default]
|
#[default]
|
||||||
OpenAiItemIds,
|
OpenAiItemIds,
|
||||||
DeepSeekOpaque,
|
DeepSeekOpaque,
|
||||||
|
/// xAI replays encrypted state without requiring OpenAI's item-ID prefix.
|
||||||
|
XaiEncrypted,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Builds a stable, wire-compatible ID for a reasoning item synthesized by Aether.
|
/// Builds a stable, wire-compatible ID for a reasoning item synthesized by Aether.
|
||||||
@@ -234,6 +237,14 @@ fn openai_responses_reasoning_item_is_replayable(
|
|||||||
{
|
{
|
||||||
return true;
|
return true;
|
||||||
}
|
}
|
||||||
|
if policy == OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
|
||||||
|
&& object
|
||||||
|
.get("encrypted_content")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.is_some_and(|value| !value.trim().is_empty())
|
||||||
|
{
|
||||||
|
return true;
|
||||||
|
}
|
||||||
let Some(id) = object
|
let Some(id) = object
|
||||||
.get("id")
|
.get("id")
|
||||||
.and_then(Value::as_str)
|
.and_then(Value::as_str)
|
||||||
@@ -334,6 +345,36 @@ mod tests {
|
|||||||
OPENAI_RESPONSES_OPERATION_COMPACT,
|
OPENAI_RESPONSES_OPERATION_COMPACT,
|
||||||
};
|
};
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_encrypted_replay_accepts_native_ids_but_excludes_foreign_carriers() {
|
||||||
|
let body = serde_json::json!({"input": [
|
||||||
|
{"type": "reasoning", "id": "native-xai-id", "encrypted_content": "opaque-xai-state"},
|
||||||
|
{"type": "reasoning", "encrypted_content": "opaque-idless-state"},
|
||||||
|
{"type": "reasoning", "id": "rs_foreign", "encrypted_content": "cpa-gemini-responses-carrier-v1:foreign"},
|
||||||
|
{"type": "reasoning", "id": "foreign-id", "summary": []}
|
||||||
|
]});
|
||||||
|
let mut xai = body.clone();
|
||||||
|
assert_eq!(
|
||||||
|
super::strip_incompatible_openai_responses_reasoning_items_with_policy(
|
||||||
|
&mut xai,
|
||||||
|
"openai:responses",
|
||||||
|
super::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted,
|
||||||
|
),
|
||||||
|
2
|
||||||
|
);
|
||||||
|
assert_eq!(xai["input"].as_array().unwrap().len(), 2);
|
||||||
|
assert_eq!(xai["input"][0], body["input"][0]);
|
||||||
|
assert_eq!(xai["input"][1], body["input"][1]);
|
||||||
|
let mut openai = body;
|
||||||
|
assert_eq!(
|
||||||
|
super::strip_incompatible_openai_responses_reasoning_items(
|
||||||
|
&mut openai,
|
||||||
|
"openai:responses"
|
||||||
|
),
|
||||||
|
4
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn gemini_tool_signature_carrier_roundtrips_direction_and_exact_value() {
|
fn gemini_tool_signature_carrier_roundtrips_direction_and_exact_value() {
|
||||||
let signature = " opaque-signature-with-padding== ";
|
let signature = " opaque-signature-with-padding== ";
|
||||||
|
|||||||
@@ -0,0 +1,914 @@
|
|||||||
|
use serde_json::{json, Map, Value};
|
||||||
|
|
||||||
|
const XAI_RESPONSES_UNSUPPORTED_BODY_FIELDS: &[&str] = &[
|
||||||
|
"previous_response_id",
|
||||||
|
"prompt_cache_retention",
|
||||||
|
"safety_identifier",
|
||||||
|
"stream_options",
|
||||||
|
"stop",
|
||||||
|
"metadata",
|
||||||
|
];
|
||||||
|
const XAI_WEB_SEARCH_TOOL_TYPE: &str = "web_search";
|
||||||
|
const XAI_IMAGE_GENERATION_TOOL_TYPE: &str = "image_generation";
|
||||||
|
const XAI_TOOL_SEARCH_TOOL_TYPE: &str = "tool_search";
|
||||||
|
const XAI_GROK_IMAGE_GENERATION_MIN: XaiGrokVersion = XaiGrokVersion { major: 4, minor: 6 };
|
||||||
|
|
||||||
|
#[derive(Clone, Copy)]
|
||||||
|
struct XaiGrokVersion {
|
||||||
|
major: i32,
|
||||||
|
minor: i32,
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn apply_xai_upstream_payload_edits(
|
||||||
|
body: &mut Value,
|
||||||
|
provider_type: &str,
|
||||||
|
provider_api_format: &str,
|
||||||
|
) {
|
||||||
|
apply_xai_upstream_payload_edits_with_client(
|
||||||
|
body,
|
||||||
|
provider_type,
|
||||||
|
provider_api_format,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn apply_xai_upstream_payload_edits_with_client(
|
||||||
|
body: &mut Value,
|
||||||
|
provider_type: &str,
|
||||||
|
provider_api_format: &str,
|
||||||
|
client_api_format: Option<&str>,
|
||||||
|
client_body: Option<&Value>,
|
||||||
|
) {
|
||||||
|
if !provider_type.trim().eq_ignore_ascii_case("xai") {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
normalize_xai_image_refs(body);
|
||||||
|
if crate::is_openai_responses_family_format(provider_api_format) {
|
||||||
|
restore_xai_web_search_from_client(body, client_api_format, client_body);
|
||||||
|
sanitize_xai_responses_body(body);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn sanitize_xai_responses_body(body: &mut Value) {
|
||||||
|
let Some(object) = body.as_object_mut() else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
for field in XAI_RESPONSES_UNSUPPORTED_BODY_FIELDS {
|
||||||
|
object.remove(*field);
|
||||||
|
}
|
||||||
|
let keep_image_generation = object
|
||||||
|
.get("model")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.is_some_and(xai_supports_native_image_generation);
|
||||||
|
normalize_xai_tool_arrays(object, keep_image_generation);
|
||||||
|
rewrite_xai_web_search_tool_choice(object);
|
||||||
|
prune_xai_orphaned_tool_choice(object);
|
||||||
|
rewrite_xai_image_generation_tool_choice(object);
|
||||||
|
drop_tool_choice_without_tools(object);
|
||||||
|
strip_unsupported_reasoning_effort(object);
|
||||||
|
sanitize_xai_input_encrypted_content(object);
|
||||||
|
}
|
||||||
|
|
||||||
|
fn restore_xai_web_search_from_client(
|
||||||
|
body: &mut Value,
|
||||||
|
client_api_format: Option<&str>,
|
||||||
|
client_body: Option<&Value>,
|
||||||
|
) {
|
||||||
|
let Some(client_api_format) = client_api_format else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
let Some(client_body) = client_body else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
if !client_requests_web_search(client_api_format, client_body) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
ensure_xai_web_search_tool(body);
|
||||||
|
// Claude names a hosted tool in tool_choice just like a client function.
|
||||||
|
// Resolve that name against the original declaration, never by name alone.
|
||||||
|
if crate::normalize_api_format_alias(client_api_format) == "claude:messages" {
|
||||||
|
let choice = &client_body["tool_choice"];
|
||||||
|
if choice["type"] == "tool"
|
||||||
|
&& choice["name"].as_str().is_some_and(|name| {
|
||||||
|
request_tools(client_body)
|
||||||
|
.iter()
|
||||||
|
.any(|tool| is_web_search_tool(tool) && tool_name(tool) == Some(name))
|
||||||
|
})
|
||||||
|
{
|
||||||
|
body["tool_choice"] = json!({"type": XAI_WEB_SEARCH_TOOL_TYPE});
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn client_requests_web_search(client_api_format: &str, client_body: &Value) -> bool {
|
||||||
|
let format = crate::normalize_api_format_alias(client_api_format);
|
||||||
|
match format.as_str() {
|
||||||
|
"openai:chat" => {
|
||||||
|
object_has_non_null_field(client_body, "web_search_options")
|
||||||
|
|| request_tools(client_body).iter().any(is_web_search_tool)
|
||||||
|
}
|
||||||
|
"claude:messages" => request_tools(client_body).iter().any(is_web_search_tool),
|
||||||
|
"gemini:generate_content" => gemini_request_has_google_search(client_body),
|
||||||
|
_ => false,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn gemini_request_has_google_search(body: &Value) -> bool {
|
||||||
|
request_tools(body).iter().any(|tool| {
|
||||||
|
tool.get("googleSearch").is_some()
|
||||||
|
|| tool.get("google_search").is_some()
|
||||||
|
|| tool
|
||||||
|
.get("googleSearchRetrieval")
|
||||||
|
.is_some_and(|value| !value.is_null())
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn object_has_non_null_field(body: &Value, field: &str) -> bool {
|
||||||
|
body.get(field).is_some_and(|value| !value.is_null())
|
||||||
|
}
|
||||||
|
|
||||||
|
fn ensure_xai_web_search_tool(body: &mut Value) {
|
||||||
|
let Some(object) = body.as_object_mut() else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
if tools_array(object).iter().any(is_web_search_tool) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
let tools = object
|
||||||
|
.entry("tools".to_string())
|
||||||
|
.or_insert_with(|| Value::Array(Vec::new()));
|
||||||
|
if let Some(tools) = tools.as_array_mut() {
|
||||||
|
tools.push(json!({ "type": XAI_WEB_SEARCH_TOOL_TYPE }));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn normalize_xai_tool_arrays(object: &mut Map<String, Value>, keep_image_generation: bool) {
|
||||||
|
if let Some(tools) = object.get_mut("tools").and_then(Value::as_array_mut) {
|
||||||
|
*tools = normalize_xai_tool_list(tools, keep_image_generation);
|
||||||
|
if tools.is_empty() {
|
||||||
|
object.remove("tools");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
let Some(input) = object.get_mut("input").and_then(Value::as_array_mut) else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
for item in input {
|
||||||
|
let Some(item_object) = item.as_object_mut() else {
|
||||||
|
continue;
|
||||||
|
};
|
||||||
|
if item_object.get("type").and_then(Value::as_str) != Some("additional_tools") {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
if let Some(tools) = item_object.get_mut("tools").and_then(Value::as_array_mut) {
|
||||||
|
*tools = normalize_xai_tool_list(tools, keep_image_generation);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn normalize_xai_tool_list(tools: &[Value], keep_image_generation: bool) -> Vec<Value> {
|
||||||
|
tools
|
||||||
|
.iter()
|
||||||
|
.filter_map(|tool| normalize_xai_tool(tool, keep_image_generation))
|
||||||
|
.collect()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn normalize_xai_tool(tool: &Value, keep_image_generation: bool) -> Option<Value> {
|
||||||
|
let Some(object) = tool.as_object() else {
|
||||||
|
return Some(tool.clone());
|
||||||
|
};
|
||||||
|
let tool_type = tool_type(tool).unwrap_or("function");
|
||||||
|
if tool_type == XAI_TOOL_SEARCH_TOOL_TYPE {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
if tool_type == XAI_IMAGE_GENERATION_TOOL_TYPE && !keep_image_generation {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
if tool_type == "custom" && tool_name(tool).is_some_and(|name| name == "apply_patch") {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut next = object.clone();
|
||||||
|
if tool_type.starts_with("web_search") {
|
||||||
|
next.insert(
|
||||||
|
"type".to_string(),
|
||||||
|
Value::String(XAI_WEB_SEARCH_TOOL_TYPE.to_string()),
|
||||||
|
);
|
||||||
|
next.remove("name");
|
||||||
|
next.remove("external_web_access");
|
||||||
|
return Some(Value::Object(next));
|
||||||
|
}
|
||||||
|
if tool_type == "custom" {
|
||||||
|
next.insert("type".to_string(), Value::String("function".to_string()));
|
||||||
|
if let Some(custom) = next.remove("custom") {
|
||||||
|
if let Some(custom_object) = custom.as_object() {
|
||||||
|
for (key, value) in custom_object {
|
||||||
|
next.entry(key.clone()).or_insert_with(|| value.clone());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !next.contains_key("parameters") {
|
||||||
|
next.insert(
|
||||||
|
"parameters".to_string(),
|
||||||
|
json!({"type": "object", "properties": {}}),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
return Some(Value::Object(next));
|
||||||
|
}
|
||||||
|
if tool_type == "function" && !next.contains_key("parameters") {
|
||||||
|
next.insert(
|
||||||
|
"parameters".to_string(),
|
||||||
|
json!({"type": "object", "properties": {}}),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
Some(Value::Object(next))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn rewrite_xai_web_search_tool_choice(object: &mut Map<String, Value>) {
|
||||||
|
let Some(choice) = object.get("tool_choice").cloned() else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
let Some(choice_type) = choice.as_object().and_then(|value| {
|
||||||
|
value
|
||||||
|
.get("type")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.map(str::trim)
|
||||||
|
.map(str::to_ascii_lowercase)
|
||||||
|
}) else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
if is_web_search_choice_type(&choice_type) {
|
||||||
|
object.insert(
|
||||||
|
"tool_choice".to_string(),
|
||||||
|
json!({
|
||||||
|
"type": "allowed_tools",
|
||||||
|
"mode": "required",
|
||||||
|
"tools": [{ "type": XAI_WEB_SEARCH_TOOL_TYPE }]
|
||||||
|
}),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn rewrite_xai_image_generation_tool_choice(object: &mut Map<String, Value>) {
|
||||||
|
let has_image_generation = tools_array(object)
|
||||||
|
.iter()
|
||||||
|
.any(|tool| tool_type(tool).is_some_and(|value| value == XAI_IMAGE_GENERATION_TOOL_TYPE));
|
||||||
|
if !has_image_generation {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
let Some(choice) = object.get("tool_choice").cloned() else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
// xAI's allowed_tools schema cannot contain image_generation. Preserve an
|
||||||
|
// image-only restriction before filtering image entries out of mixed lists.
|
||||||
|
let image_only = is_allowed_tools_image_generation_only(&choice);
|
||||||
|
if choice["type"] == XAI_IMAGE_GENERATION_TOOL_TYPE || image_only {
|
||||||
|
let mode = if image_only && choice["mode"] == "auto" {
|
||||||
|
"auto"
|
||||||
|
} else {
|
||||||
|
"required"
|
||||||
|
};
|
||||||
|
keep_only_image_generation_tools(object);
|
||||||
|
object.insert("tool_choice".to_string(), Value::String(mode.to_string()));
|
||||||
|
} else if choice["type"] == "allowed_tools" {
|
||||||
|
filter_image_generation_from_allowed_tools(object);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn is_allowed_tools_image_generation_only(choice: &Value) -> bool {
|
||||||
|
let Some(object) = choice.as_object() else {
|
||||||
|
return false;
|
||||||
|
};
|
||||||
|
if object.get("type").and_then(Value::as_str) != Some("allowed_tools") {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
let Some(tools) = object.get("tools").and_then(Value::as_array) else {
|
||||||
|
return false;
|
||||||
|
};
|
||||||
|
!tools.is_empty()
|
||||||
|
&& tools.iter().all(|tool| {
|
||||||
|
tool_type(tool).is_some_and(|value| value == XAI_IMAGE_GENERATION_TOOL_TYPE)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn keep_only_image_generation_tools(object: &mut Map<String, Value>) {
|
||||||
|
let Some(tools) = object.get_mut("tools").and_then(Value::as_array_mut) else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
tools.retain(|tool| {
|
||||||
|
tool_type(tool).is_some_and(|value| value == XAI_IMAGE_GENERATION_TOOL_TYPE)
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
fn filter_image_generation_from_allowed_tools(object: &mut Map<String, Value>) {
|
||||||
|
let Some(choice) = object.get_mut("tool_choice").and_then(Value::as_object_mut) else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
let Some(tools) = choice.get_mut("tools").and_then(Value::as_array_mut) else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
tools
|
||||||
|
.retain(|tool| tool_type(tool).is_none_or(|value| value != XAI_IMAGE_GENERATION_TOOL_TYPE));
|
||||||
|
}
|
||||||
|
|
||||||
|
fn is_web_search_choice_type(value: &str) -> bool {
|
||||||
|
value == XAI_WEB_SEARCH_TOOL_TYPE || value.starts_with("web_search")
|
||||||
|
}
|
||||||
|
|
||||||
|
fn prune_xai_orphaned_tool_choice(object: &mut Map<String, Value>) {
|
||||||
|
let available = collect_available_tool_choice_keys(object);
|
||||||
|
let Some(choice) = object.get("tool_choice").cloned() else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
if choice.as_str().is_some() {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
let Some(choice_object) = choice.as_object() else {
|
||||||
|
object.remove("tool_choice");
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
let choice_type = choice_object
|
||||||
|
.get("type")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.unwrap_or_default()
|
||||||
|
.trim()
|
||||||
|
.to_ascii_lowercase();
|
||||||
|
if choice_type == "allowed_tools" {
|
||||||
|
let Some(allowed) = choice_object.get("tools").and_then(Value::as_array) else {
|
||||||
|
object.remove("tool_choice");
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
let kept = allowed
|
||||||
|
.iter()
|
||||||
|
.filter(|tool| tool_matches_available(tool, &available))
|
||||||
|
.cloned()
|
||||||
|
.collect::<Vec<_>>();
|
||||||
|
if kept.is_empty() {
|
||||||
|
object.remove("tool_choice");
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
if let Some(choice) = object.get_mut("tool_choice").and_then(Value::as_object_mut) {
|
||||||
|
choice.insert("tools".to_string(), Value::Array(kept));
|
||||||
|
}
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
if choice_type.is_empty() {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
if !tool_matches_available(&choice, &available) {
|
||||||
|
object.remove("tool_choice");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn collect_available_tool_choice_keys(object: &Map<String, Value>) -> Vec<ToolChoiceKey> {
|
||||||
|
let mut keys = Vec::new();
|
||||||
|
collect_tool_choice_keys(tools_array(object), &mut keys);
|
||||||
|
if let Some(input) = object.get("input").and_then(Value::as_array) {
|
||||||
|
for item in input {
|
||||||
|
if item.get("type").and_then(Value::as_str) == Some("additional_tools") {
|
||||||
|
collect_tool_choice_keys(
|
||||||
|
item.get("tools")
|
||||||
|
.and_then(Value::as_array)
|
||||||
|
.map(Vec::as_slice)
|
||||||
|
.unwrap_or(&[]),
|
||||||
|
&mut keys,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
keys
|
||||||
|
}
|
||||||
|
|
||||||
|
fn collect_tool_choice_keys(tools: &[Value], keys: &mut Vec<ToolChoiceKey>) {
|
||||||
|
for tool in tools {
|
||||||
|
let Some(tool_type) = tool_type(tool) else {
|
||||||
|
continue;
|
||||||
|
};
|
||||||
|
if matches!(tool_type, "function" | "custom") {
|
||||||
|
if let Some(name) = tool_name(tool) {
|
||||||
|
keys.push(ToolChoiceKey::Named {
|
||||||
|
name: name.to_ascii_lowercase(),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
keys.push(ToolChoiceKey::Hosted(tool_type.to_ascii_lowercase()));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn tool_matches_available(choice: &Value, available: &[ToolChoiceKey]) -> bool {
|
||||||
|
let Some(object) = choice.as_object() else {
|
||||||
|
return false;
|
||||||
|
};
|
||||||
|
let choice_type = object
|
||||||
|
.get("type")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.unwrap_or_default()
|
||||||
|
.trim()
|
||||||
|
.to_ascii_lowercase();
|
||||||
|
if matches!(choice_type.as_str(), "function" | "custom" | "tool") {
|
||||||
|
let Some(name) = tool_choice_name(object) else {
|
||||||
|
return false;
|
||||||
|
};
|
||||||
|
return available.iter().any(|key| {
|
||||||
|
matches!(
|
||||||
|
key,
|
||||||
|
ToolChoiceKey::Named { name: available_name, .. }
|
||||||
|
if available_name == &name.to_ascii_lowercase()
|
||||||
|
)
|
||||||
|
});
|
||||||
|
}
|
||||||
|
if is_web_search_choice_type(&choice_type) {
|
||||||
|
return available.iter().any(
|
||||||
|
|key| matches!(key, ToolChoiceKey::Hosted(value) if value == XAI_WEB_SEARCH_TOOL_TYPE),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
available
|
||||||
|
.iter()
|
||||||
|
.any(|key| matches!(key, ToolChoiceKey::Hosted(value) if value == &choice_type))
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Clone, Debug)]
|
||||||
|
enum ToolChoiceKey {
|
||||||
|
Named { name: String },
|
||||||
|
Hosted(String),
|
||||||
|
}
|
||||||
|
|
||||||
|
fn drop_tool_choice_without_tools(object: &mut Map<String, Value>) {
|
||||||
|
if xai_request_has_tools(object) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
object.remove("tools");
|
||||||
|
object.remove("tool_choice");
|
||||||
|
object.remove("parallel_tool_calls");
|
||||||
|
}
|
||||||
|
|
||||||
|
fn xai_request_has_tools(object: &Map<String, Value>) -> bool {
|
||||||
|
if !tools_array(object).is_empty() {
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
object
|
||||||
|
.get("input")
|
||||||
|
.and_then(Value::as_array)
|
||||||
|
.into_iter()
|
||||||
|
.flatten()
|
||||||
|
.any(|item| {
|
||||||
|
item.get("type")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.is_some_and(|value| value == "additional_tools")
|
||||||
|
&& item
|
||||||
|
.get("tools")
|
||||||
|
.and_then(Value::as_array)
|
||||||
|
.is_some_and(|tools| !tools.is_empty())
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn strip_unsupported_reasoning_effort(object: &mut Map<String, Value>) {
|
||||||
|
let model = object
|
||||||
|
.get("model")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.unwrap_or_default();
|
||||||
|
if xai_model_supports_reasoning_effort(model) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
let Some(reasoning) = object.get_mut("reasoning") else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
let Some(reasoning_object) = reasoning.as_object_mut() else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
reasoning_object.remove("effort");
|
||||||
|
if reasoning_object.is_empty() {
|
||||||
|
object.remove("reasoning");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn xai_model_supports_reasoning_effort(model: &str) -> bool {
|
||||||
|
let lowered = model.trim().to_ascii_lowercase();
|
||||||
|
let name = lowered.rsplit('/').next().unwrap_or(lowered.as_str());
|
||||||
|
if name.is_empty() || name.contains("non-reasoning") || name.contains("imagine") {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
name.starts_with("grok-3-mini")
|
||||||
|
|| name.starts_with("grok-4")
|
||||||
|
|| name.starts_with("grok-build")
|
||||||
|
|| name.starts_with("grok-composer")
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn xai_supports_native_image_generation(model: &str) -> bool {
|
||||||
|
let lowered = model.trim().to_ascii_lowercase();
|
||||||
|
let name = lowered.rsplit('/').next().unwrap_or(lowered.as_str());
|
||||||
|
let Some(rest) = name.strip_prefix("grok-") else {
|
||||||
|
return false;
|
||||||
|
};
|
||||||
|
if rest == "4.20" || rest.starts_with("4.20-") {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
parse_grok_version_prefix(rest).is_some_and(grok_version_at_least_image_generation)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn parse_grok_version_prefix(rest: &str) -> Option<XaiGrokVersion> {
|
||||||
|
let major_len = rest
|
||||||
|
.find(|ch: char| !ch.is_ascii_digit())
|
||||||
|
.unwrap_or(rest.len());
|
||||||
|
if major_len == 0 {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
let major = rest[..major_len].parse().ok()?;
|
||||||
|
if major_len == rest.len() || !rest[major_len..].starts_with('.') {
|
||||||
|
return Some(XaiGrokVersion { major, minor: -1 });
|
||||||
|
}
|
||||||
|
let after_dot = &rest[major_len + 1..];
|
||||||
|
let minor_len = after_dot
|
||||||
|
.find(|ch: char| !ch.is_ascii_digit())
|
||||||
|
.unwrap_or(after_dot.len());
|
||||||
|
if minor_len == 0 {
|
||||||
|
return Some(XaiGrokVersion { major, minor: -1 });
|
||||||
|
}
|
||||||
|
let minor = after_dot[..minor_len].parse().ok()?;
|
||||||
|
Some(XaiGrokVersion { major, minor })
|
||||||
|
}
|
||||||
|
|
||||||
|
fn grok_version_at_least_image_generation(version: XaiGrokVersion) -> bool {
|
||||||
|
let minor = if version.minor < 0 { 0 } else { version.minor };
|
||||||
|
(version.major, minor)
|
||||||
|
>= (
|
||||||
|
XAI_GROK_IMAGE_GENERATION_MIN.major,
|
||||||
|
XAI_GROK_IMAGE_GENERATION_MIN.minor,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn sanitize_xai_input_encrypted_content(object: &mut Map<String, Value>) {
|
||||||
|
let Some(input) = object.get_mut("input").and_then(Value::as_array_mut) else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
let mut kept = Vec::new();
|
||||||
|
for item in input.iter() {
|
||||||
|
let Some(item_object) = item.as_object() else {
|
||||||
|
kept.push(item.clone());
|
||||||
|
continue;
|
||||||
|
};
|
||||||
|
let item_type = item_object
|
||||||
|
.get("type")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.unwrap_or_default();
|
||||||
|
if item_type != "reasoning" && item_type != "compaction" {
|
||||||
|
kept.push(item.clone());
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
let Some(encrypted) = item_object.get("encrypted_content") else {
|
||||||
|
kept.push(item.clone());
|
||||||
|
continue;
|
||||||
|
};
|
||||||
|
let valid = encrypted
|
||||||
|
.as_str()
|
||||||
|
.is_some_and(|value| !value.trim().is_empty());
|
||||||
|
if valid {
|
||||||
|
kept.push(item.clone());
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
if item_type == "compaction" {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
let mut next = item_object.clone();
|
||||||
|
next.remove("encrypted_content");
|
||||||
|
kept.push(Value::Object(next));
|
||||||
|
}
|
||||||
|
*input = kept;
|
||||||
|
}
|
||||||
|
|
||||||
|
fn normalize_xai_image_refs(value: &mut Value) {
|
||||||
|
match value {
|
||||||
|
Value::Object(object) => {
|
||||||
|
for key in ["image", "images", "reference_images"] {
|
||||||
|
match object.get_mut(key) {
|
||||||
|
Some(Value::Array(items)) if key != "image" => {
|
||||||
|
for item in items {
|
||||||
|
normalize_xai_image_ref(item);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Some(item) if key == "image" => normalize_xai_image_ref(item),
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for child in object.values_mut() {
|
||||||
|
normalize_xai_image_refs(child);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Value::Array(items) => {
|
||||||
|
for item in items {
|
||||||
|
normalize_xai_image_refs(item);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn normalize_xai_image_ref(value: &mut Value) {
|
||||||
|
let Some(object) = value.as_object_mut() else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
let original_url = object
|
||||||
|
.get("url")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.map(ToOwned::to_owned);
|
||||||
|
let image_url = object.get("image_url").cloned();
|
||||||
|
let resolved_url = original_url.clone().or_else(|| match image_url.as_ref() {
|
||||||
|
Some(Value::String(url)) => {
|
||||||
|
let trimmed = url.trim();
|
||||||
|
(!trimmed.is_empty()).then(|| trimmed.to_string())
|
||||||
|
}
|
||||||
|
Some(Value::Object(inner)) => inner
|
||||||
|
.get("url")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.map(ToOwned::to_owned),
|
||||||
|
_ => None,
|
||||||
|
});
|
||||||
|
let Some(url) = resolved_url else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
if original_url.as_deref() == Some(url.as_str()) && image_url.is_none() {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
object.insert("url".to_string(), Value::String(url));
|
||||||
|
object.remove("image_url");
|
||||||
|
}
|
||||||
|
|
||||||
|
fn request_tools(body: &Value) -> &[Value] {
|
||||||
|
body.get("tools")
|
||||||
|
.and_then(Value::as_array)
|
||||||
|
.map(Vec::as_slice)
|
||||||
|
.unwrap_or(&[])
|
||||||
|
}
|
||||||
|
|
||||||
|
fn tools_array(object: &Map<String, Value>) -> &[Value] {
|
||||||
|
object
|
||||||
|
.get("tools")
|
||||||
|
.and_then(Value::as_array)
|
||||||
|
.map(Vec::as_slice)
|
||||||
|
.unwrap_or(&[])
|
||||||
|
}
|
||||||
|
|
||||||
|
fn tool_type(tool: &Value) -> Option<&str> {
|
||||||
|
tool.get("type").and_then(Value::as_str).map(str::trim)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn tool_name(tool: &Value) -> Option<&str> {
|
||||||
|
tool.get("name")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.or_else(|| {
|
||||||
|
tool.get("function")
|
||||||
|
.and_then(Value::as_object)
|
||||||
|
.and_then(|value| value.get("name"))
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
})
|
||||||
|
.or_else(|| {
|
||||||
|
tool.get("custom")
|
||||||
|
.and_then(Value::as_object)
|
||||||
|
.and_then(|value| value.get("name"))
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
})
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
}
|
||||||
|
|
||||||
|
fn tool_choice_name(choice: &Map<String, Value>) -> Option<&str> {
|
||||||
|
choice
|
||||||
|
.get("name")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.or_else(|| {
|
||||||
|
choice
|
||||||
|
.get("function")
|
||||||
|
.and_then(Value::as_object)
|
||||||
|
.and_then(|value| value.get("name"))
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
})
|
||||||
|
.or_else(|| {
|
||||||
|
choice
|
||||||
|
.get("custom")
|
||||||
|
.and_then(Value::as_object)
|
||||||
|
.and_then(|value| value.get("name"))
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
})
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
}
|
||||||
|
|
||||||
|
fn is_web_search_tool(tool: &Value) -> bool {
|
||||||
|
tool_type(tool).is_some_and(is_web_search_choice_type)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use serde_json::json;
|
||||||
|
|
||||||
|
use super::{
|
||||||
|
apply_xai_upstream_payload_edits, apply_xai_upstream_payload_edits_with_client,
|
||||||
|
xai_model_supports_reasoning_effort, xai_supports_native_image_generation,
|
||||||
|
XAI_RESPONSES_UNSUPPORTED_BODY_FIELDS,
|
||||||
|
};
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_responses_edits_strip_continuation_fields_and_empty_tool_choice() {
|
||||||
|
let mut body = json!({
|
||||||
|
"model": "grok-4.6",
|
||||||
|
"input": "hello",
|
||||||
|
"previous_response_id": "resp_123",
|
||||||
|
"prompt_cache_retention": "24h",
|
||||||
|
"safety_identifier": "user-1",
|
||||||
|
"stream_options": {"include_obfuscation": true},
|
||||||
|
"stop": ["END"],
|
||||||
|
"metadata": {
|
||||||
|
"user_id": "{\"device_id\":\"dev-1\",\"account_uuid\":\"acct-1\",\"session_id\":\"sess-1\"}"
|
||||||
|
},
|
||||||
|
"include": ["reasoning.encrypted_content", "file_search_call.results"],
|
||||||
|
"tool_choice": "auto",
|
||||||
|
"parallel_tool_calls": true,
|
||||||
|
"tools": []
|
||||||
|
});
|
||||||
|
|
||||||
|
apply_xai_upstream_payload_edits(&mut body, "xai", "openai:responses");
|
||||||
|
|
||||||
|
for field in XAI_RESPONSES_UNSUPPORTED_BODY_FIELDS {
|
||||||
|
assert!(body.get(*field).is_none(), "{field} should be stripped");
|
||||||
|
}
|
||||||
|
assert!(body.get("tool_choice").is_none());
|
||||||
|
assert!(body.get("parallel_tool_calls").is_none());
|
||||||
|
assert!(body.get("tools").is_none());
|
||||||
|
assert_eq!(
|
||||||
|
body["include"],
|
||||||
|
json!(["reasoning.encrypted_content", "file_search_call.results"])
|
||||||
|
);
|
||||||
|
assert_eq!(body["model"], "grok-4.6");
|
||||||
|
assert_eq!(body["input"], "hello");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_responses_edits_keep_reasoning_effort_for_thinking_models() {
|
||||||
|
let mut body = json!({
|
||||||
|
"model": "grok-4.6",
|
||||||
|
"reasoning": {"effort": "high", "summary": "auto"}
|
||||||
|
});
|
||||||
|
apply_xai_upstream_payload_edits(&mut body, "xai", "openai:responses");
|
||||||
|
assert_eq!(body["reasoning"]["effort"], "high");
|
||||||
|
assert_eq!(body["reasoning"]["summary"], "auto");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_responses_edits_strip_reasoning_effort_for_non_thinking_models() {
|
||||||
|
let mut body = json!({
|
||||||
|
"model": "grok-4.20-0309-non-reasoning",
|
||||||
|
"reasoning": {"effort": "high"}
|
||||||
|
});
|
||||||
|
apply_xai_upstream_payload_edits(&mut body, "xai", "openai:responses");
|
||||||
|
assert!(body.get("reasoning").is_none());
|
||||||
|
assert!(!xai_model_supports_reasoning_effort(
|
||||||
|
"grok-4.20-0309-non-reasoning"
|
||||||
|
));
|
||||||
|
assert!(xai_model_supports_reasoning_effort("xai/grok-4.5"));
|
||||||
|
assert!(!xai_model_supports_reasoning_effort("grok-imagine-image"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_hosted_tool_choice_rewrites_web_search_and_image_generation() {
|
||||||
|
let mut web_search = json!({
|
||||||
|
"model": "grok-4.6",
|
||||||
|
"tools": [{"type": "web_search_preview", "name": "web_search"}],
|
||||||
|
"tool_choice": {"type": "web_search"}
|
||||||
|
});
|
||||||
|
apply_xai_upstream_payload_edits(&mut web_search, "xai", "openai:responses");
|
||||||
|
assert_eq!(web_search["tools"][0]["type"], "web_search");
|
||||||
|
assert!(web_search["tools"][0].get("name").is_none());
|
||||||
|
assert_eq!(web_search["tool_choice"]["type"], "allowed_tools");
|
||||||
|
assert_eq!(web_search["tool_choice"]["mode"], "required");
|
||||||
|
assert_eq!(web_search["tool_choice"]["tools"][0]["type"], "web_search");
|
||||||
|
|
||||||
|
let mut image = json!({
|
||||||
|
"model": "grok-4.6",
|
||||||
|
"tools": [
|
||||||
|
{"type": "web_search"},
|
||||||
|
{"type": "image_generation", "action": "generate"}
|
||||||
|
],
|
||||||
|
"tool_choice": {"type": "image_generation"}
|
||||||
|
});
|
||||||
|
apply_xai_upstream_payload_edits(&mut image, "xai", "openai:responses");
|
||||||
|
assert_eq!(image["tool_choice"], "required");
|
||||||
|
assert_eq!(image["tools"].as_array().map(Vec::len), Some(1));
|
||||||
|
assert_eq!(image["tools"][0]["type"], "image_generation");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_strips_image_generation_on_older_conversation_models() {
|
||||||
|
let mut body = json!({
|
||||||
|
"model": "grok-4.5",
|
||||||
|
"tools": [
|
||||||
|
{"type": "function", "name": "lookup", "parameters": {"type": "object"}},
|
||||||
|
{"type": "image_generation"}
|
||||||
|
],
|
||||||
|
"tool_choice": {"type": "image_generation"}
|
||||||
|
});
|
||||||
|
apply_xai_upstream_payload_edits(&mut body, "xai", "openai:responses");
|
||||||
|
assert_eq!(body["tools"].as_array().map(Vec::len), Some(1));
|
||||||
|
assert_eq!(body["tools"][0]["name"], "lookup");
|
||||||
|
assert!(body.get("tool_choice").is_none());
|
||||||
|
assert!(xai_supports_native_image_generation("grok-4.6"));
|
||||||
|
assert!(!xai_supports_native_image_generation("grok-4.20-0309"));
|
||||||
|
assert!(!xai_supports_native_image_generation("grok-4.5"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_restores_web_search_from_chat_and_claude_clients() {
|
||||||
|
let mut chat_body = json!({
|
||||||
|
"model": "grok-4.6",
|
||||||
|
"input": "search this"
|
||||||
|
});
|
||||||
|
apply_xai_upstream_payload_edits_with_client(
|
||||||
|
&mut chat_body,
|
||||||
|
"xai",
|
||||||
|
"openai:responses",
|
||||||
|
Some("openai:chat"),
|
||||||
|
Some(&json!({
|
||||||
|
"messages": [{"role": "user", "content": "news"}],
|
||||||
|
"web_search_options": {"search_context_size": "high"}
|
||||||
|
})),
|
||||||
|
);
|
||||||
|
assert_eq!(chat_body["tools"][0]["type"], "web_search");
|
||||||
|
|
||||||
|
let mut claude_body = json!({
|
||||||
|
"model": "grok-4.6",
|
||||||
|
"input": "search this",
|
||||||
|
"tools": [{
|
||||||
|
"type": "function",
|
||||||
|
"name": "lookup",
|
||||||
|
"parameters": {"type": "object", "properties": {}}
|
||||||
|
}],
|
||||||
|
"tool_choice": {"type": "function", "name": "web_search"}
|
||||||
|
});
|
||||||
|
apply_xai_upstream_payload_edits_with_client(
|
||||||
|
&mut claude_body,
|
||||||
|
"xai",
|
||||||
|
"openai:responses",
|
||||||
|
Some("claude:messages"),
|
||||||
|
Some(&json!({
|
||||||
|
"tools": [
|
||||||
|
{"type": "web_search_20250305", "name": "web_search"},
|
||||||
|
{"name": "lookup", "input_schema": {"type": "object"}}
|
||||||
|
],
|
||||||
|
"tool_choice": {"type": "tool", "name": "web_search"}
|
||||||
|
})),
|
||||||
|
);
|
||||||
|
assert!(claude_body["tools"]
|
||||||
|
.as_array()
|
||||||
|
.into_iter()
|
||||||
|
.flatten()
|
||||||
|
.any(|tool| tool["type"] == "web_search"));
|
||||||
|
assert_eq!(claude_body["tool_choice"]["type"], "allowed_tools");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_image_refs_rewrite_openai_aliases_without_touching_chat_parts() {
|
||||||
|
let mut body = json!({
|
||||||
|
"model": "grok-imagine-image",
|
||||||
|
"prompt": "edit this",
|
||||||
|
"image": {"image_url": "https://cdn.example/a.png"},
|
||||||
|
"reference_images": [
|
||||||
|
{"image_url": {"url": "https://cdn.example/b.png"}}
|
||||||
|
],
|
||||||
|
"input": [{
|
||||||
|
"type": "message",
|
||||||
|
"content": [{
|
||||||
|
"type": "image_url",
|
||||||
|
"image_url": {"url": "https://cdn.example/chat.png"}
|
||||||
|
}]
|
||||||
|
}]
|
||||||
|
});
|
||||||
|
|
||||||
|
apply_xai_upstream_payload_edits(&mut body, "xai", "openai:image");
|
||||||
|
|
||||||
|
assert_eq!(body["image"]["url"], "https://cdn.example/a.png");
|
||||||
|
assert!(body["image"].get("image_url").is_none());
|
||||||
|
assert_eq!(
|
||||||
|
body["reference_images"][0]["url"],
|
||||||
|
"https://cdn.example/b.png"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
body["input"][0]["content"][0]["image_url"]["url"],
|
||||||
|
"https://cdn.example/chat.png"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn other_providers_are_left_untouched() {
|
||||||
|
let mut body = json!({
|
||||||
|
"previous_response_id": "resp_123",
|
||||||
|
"image": {"image_url": "https://cdn.example/a.png"}
|
||||||
|
});
|
||||||
|
apply_xai_upstream_payload_edits(&mut body, "codex", "openai:responses");
|
||||||
|
assert_eq!(body["previous_response_id"], "resp_123");
|
||||||
|
assert_eq!(body["image"]["image_url"], "https://cdn.example/a.png");
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -49,6 +49,10 @@ pub fn resolve_execution_runtime_stream_plan_kind_with_client_surface(
|
|||||||
method: &Method,
|
method: &Method,
|
||||||
path: &str,
|
path: &str,
|
||||||
) -> Option<&'static str> {
|
) -> Option<&'static str> {
|
||||||
|
let path = path
|
||||||
|
.strip_prefix("/openai")
|
||||||
|
.filter(|p| *p == "/v1/videos" || p.starts_with("/v1/videos/"))
|
||||||
|
.unwrap_or(path);
|
||||||
if route_class != Some("ai_public") {
|
if route_class != Some("ai_public") {
|
||||||
return None;
|
return None;
|
||||||
}
|
}
|
||||||
@@ -181,6 +185,10 @@ pub fn resolve_execution_runtime_sync_plan_kind_with_client_surface(
|
|||||||
method: &Method,
|
method: &Method,
|
||||||
path: &str,
|
path: &str,
|
||||||
) -> Option<&'static str> {
|
) -> Option<&'static str> {
|
||||||
|
let path = path
|
||||||
|
.strip_prefix("/openai")
|
||||||
|
.filter(|p| *p == "/v1/videos" || p.starts_with("/v1/videos/"))
|
||||||
|
.unwrap_or(path);
|
||||||
if route_class != Some("ai_public") {
|
if route_class != Some("ai_public") {
|
||||||
return None;
|
return None;
|
||||||
}
|
}
|
||||||
@@ -206,7 +214,10 @@ pub fn resolve_execution_runtime_sync_plan_kind_with_client_surface(
|
|||||||
if route_family == Some("openai")
|
if route_family == Some("openai")
|
||||||
&& route_kind == Some("video")
|
&& route_kind == Some("video")
|
||||||
&& *method == Method::POST
|
&& *method == Method::POST
|
||||||
&& path == "/v1/videos"
|
&& matches!(
|
||||||
|
path,
|
||||||
|
"/v1/videos" | "/v1/videos/generations" | "/v1/videos/edits" | "/v1/videos/extensions"
|
||||||
|
)
|
||||||
{
|
{
|
||||||
return Some(OPENAI_VIDEO_CREATE_SYNC_PLAN_KIND);
|
return Some(OPENAI_VIDEO_CREATE_SYNC_PLAN_KIND);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ use crate::formats::openai::responses::codex::{
|
|||||||
apply_codex_openai_responses_chat_body_edits, apply_codex_openai_responses_special_body_edits,
|
apply_codex_openai_responses_chat_body_edits, apply_codex_openai_responses_special_body_edits,
|
||||||
apply_openai_responses_compact_special_body_edits,
|
apply_openai_responses_compact_special_body_edits,
|
||||||
};
|
};
|
||||||
|
use crate::formats::openai::responses::xai::apply_xai_upstream_payload_edits_with_client;
|
||||||
use crate::formats::shared::standard_normalize::{
|
use crate::formats::shared::standard_normalize::{
|
||||||
build_local_openai_chat_request_body_with_model_directives,
|
build_local_openai_chat_request_body_with_model_directives,
|
||||||
is_claude_messages_shaped_body_on_openai_chat_endpoint,
|
is_claude_messages_shaped_body_on_openai_chat_endpoint,
|
||||||
@@ -121,6 +122,11 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and
|
|||||||
enable_model_directives: bool,
|
enable_model_directives: bool,
|
||||||
reasoning_replay_policy: crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy,
|
reasoning_replay_policy: crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy,
|
||||||
) -> Option<Value> {
|
) -> Option<Value> {
|
||||||
|
let reasoning_replay_policy = if provider_type.trim().eq_ignore_ascii_case("xai") {
|
||||||
|
crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
|
||||||
|
} else {
|
||||||
|
reasoning_replay_policy
|
||||||
|
};
|
||||||
let mut format_context = FormatContext::default()
|
let mut format_context = FormatContext::default()
|
||||||
.with_mapped_model(mapped_model)
|
.with_mapped_model(mapped_model)
|
||||||
.with_request_path(request_path)
|
.with_request_path(request_path)
|
||||||
@@ -133,13 +139,10 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and
|
|||||||
client_api_format,
|
client_api_format,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
);
|
);
|
||||||
// DeepSeek's Responses continuation state is opaque. Parsing a same-wire-format
|
// DeepSeek and xAI replay opaque provider state. Preserve their native
|
||||||
// request through the canonical model would discard its id-less `reasoning_text`
|
// Responses input items: canonical conversion can lose reasoning IDs and
|
||||||
// items and future provider-owned fields even though no conversion is required.
|
// encrypted-only items even when source and destination formats are equal.
|
||||||
// Keep that provider-specific route wire-preserving, while retaining canonical
|
let mut provider_request_body = if is_wire_preserving_responses_hop(
|
||||||
// normalization for ordinary OpenAI Responses and for Responses/Compact
|
|
||||||
// cross-format conversions.
|
|
||||||
let mut provider_request_body = if is_wire_preserving_deepseek_responses_hop(
|
|
||||||
source_api_format.as_ref(),
|
source_api_format.as_ref(),
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
reasoning_replay_policy,
|
reasoning_replay_policy,
|
||||||
@@ -200,6 +203,13 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and
|
|||||||
&mut provider_request_body,
|
&mut provider_request_body,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
);
|
);
|
||||||
|
apply_xai_upstream_payload_edits_with_client(
|
||||||
|
&mut provider_request_body,
|
||||||
|
provider_type,
|
||||||
|
provider_api_format,
|
||||||
|
Some(client_api_format),
|
||||||
|
Some(body_json),
|
||||||
|
);
|
||||||
crate::formats::openai::responses::strip_incompatible_openai_responses_reasoning_items_with_policy(
|
crate::formats::openai::responses::strip_incompatible_openai_responses_reasoning_items_with_policy(
|
||||||
&mut provider_request_body,
|
&mut provider_request_body,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
@@ -224,14 +234,16 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and
|
|||||||
Some(provider_request_body)
|
Some(provider_request_body)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn is_wire_preserving_deepseek_responses_hop(
|
fn is_wire_preserving_responses_hop(
|
||||||
source_api_format: &str,
|
source_api_format: &str,
|
||||||
provider_api_format: &str,
|
provider_api_format: &str,
|
||||||
reasoning_replay_policy: crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy,
|
reasoning_replay_policy: crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy,
|
||||||
) -> bool {
|
) -> bool {
|
||||||
if reasoning_replay_policy
|
if !matches!(
|
||||||
!= crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
|
reasoning_replay_policy,
|
||||||
{
|
crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
|
||||||
|
| crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
|
||||||
|
) {
|
||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
let source_api_format = aether_ai_formats::normalize_api_format_alias(source_api_format);
|
let source_api_format = aether_ai_formats::normalize_api_format_alias(source_api_format);
|
||||||
@@ -2077,4 +2089,316 @@ mod tests {
|
|||||||
);
|
);
|
||||||
assert_eq!(gemini["toolConfig"]["functionCallingConfig"]["mode"], "ANY");
|
assert_eq!(gemini["toolConfig"]["functionCallingConfig"]["mode"], "ANY");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_keeps_client_search_functions_distinct_from_hosted_search() {
|
||||||
|
for name in ["web_search", "web_search_internal"] {
|
||||||
|
for hosted in [false, true] {
|
||||||
|
let mut tools = vec![json!({
|
||||||
|
"name": name,
|
||||||
|
"description": "Search internal documents",
|
||||||
|
"input_schema": {"type": "object", "properties": {"query": {"type": "string"}}}
|
||||||
|
})];
|
||||||
|
if hosted {
|
||||||
|
tools.push(json!({"type": "web_search_20260209", "name": "internet_search"}));
|
||||||
|
}
|
||||||
|
let request = json!({
|
||||||
|
"model": "source", "max_tokens": 64,
|
||||||
|
"messages": [{"role": "user", "content": "Search internal documents"}],
|
||||||
|
"tools": tools,
|
||||||
|
"tool_choice": {"type": "tool", "name": name}
|
||||||
|
});
|
||||||
|
let converted = build_standard_request_body(
|
||||||
|
&request,
|
||||||
|
"claude:messages",
|
||||||
|
"grok-4.6",
|
||||||
|
"xai",
|
||||||
|
"openai:responses",
|
||||||
|
"/v1/messages",
|
||||||
|
true,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
converted["tool_choice"],
|
||||||
|
json!({"type": "function", "name": name})
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
converted["tools"]
|
||||||
|
.as_array()
|
||||||
|
.unwrap()
|
||||||
|
.iter()
|
||||||
|
.any(|tool| tool["type"] == "web_search"),
|
||||||
|
hosted
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
let request = json!({
|
||||||
|
"model": "source", "max_tokens": 64,
|
||||||
|
"messages": [{"role": "user", "content": "Search the internet"}],
|
||||||
|
"tools": [{"type": "web_search_20260209", "name": "internet_search"}],
|
||||||
|
"tool_choice": {"type": "tool", "name": "internet_search"}
|
||||||
|
});
|
||||||
|
let converted = build_standard_request_body(
|
||||||
|
&request,
|
||||||
|
"claude:messages",
|
||||||
|
"grok-4.6",
|
||||||
|
"xai",
|
||||||
|
"openai:responses",
|
||||||
|
"/v1/messages",
|
||||||
|
true,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
converted["tool_choice"],
|
||||||
|
json!({
|
||||||
|
"type": "allowed_tools", "mode": "required", "tools": [{"type": "web_search"}]
|
||||||
|
})
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_preserves_function_choices_in_chat_and_responses_requests() {
|
||||||
|
for name in ["web_search", "web_search_internal"] {
|
||||||
|
for (client, request) in [
|
||||||
|
(
|
||||||
|
"openai:chat",
|
||||||
|
json!({
|
||||||
|
"messages": [{"role": "user", "content": "search"}],
|
||||||
|
"tools": [{"type": "function", "function": {"name": name, "parameters": {"type": "object"}}}],
|
||||||
|
"tool_choice": {"type": "function", "function": {"name": name}}
|
||||||
|
}),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"openai:responses",
|
||||||
|
json!({
|
||||||
|
"input": "search",
|
||||||
|
"tools": [{"type": "function", "name": name, "parameters": {"type": "object"}}],
|
||||||
|
"tool_choice": {"type": "function", "name": name}
|
||||||
|
}),
|
||||||
|
),
|
||||||
|
] {
|
||||||
|
let converted = build_standard_request_body(
|
||||||
|
&request,
|
||||||
|
client,
|
||||||
|
"grok-4.6",
|
||||||
|
"xai",
|
||||||
|
"openai:responses",
|
||||||
|
"/v1/responses",
|
||||||
|
true,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
converted["tool_choice"],
|
||||||
|
json!({"type": "function", "name": name})
|
||||||
|
);
|
||||||
|
assert_eq!(converted["tools"].as_array().unwrap().len(), 1);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_image_allowed_tools_preserves_mode_and_restricts_available_tools() {
|
||||||
|
for mode in ["auto", "required"] {
|
||||||
|
for mixed in [false, true] {
|
||||||
|
let mut allowed = vec![json!({"type": "image_generation"})];
|
||||||
|
if mixed {
|
||||||
|
allowed.push(json!({"type": "function", "name": "lookup"}));
|
||||||
|
}
|
||||||
|
let request = json!({
|
||||||
|
"input": "Draw a cat",
|
||||||
|
"tools": [
|
||||||
|
{"type": "web_search"}, {"type": "image_generation"},
|
||||||
|
{"type": "function", "name": "lookup", "parameters": {"type": "object"}}
|
||||||
|
],
|
||||||
|
"tool_choice": {"type": "allowed_tools", "mode": mode, "tools": allowed}
|
||||||
|
});
|
||||||
|
let converted = build_standard_request_body(
|
||||||
|
&request,
|
||||||
|
"openai:responses",
|
||||||
|
"grok-4.6",
|
||||||
|
"xai",
|
||||||
|
"openai:responses",
|
||||||
|
"/v1/responses",
|
||||||
|
true,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
if mixed {
|
||||||
|
assert_eq!(
|
||||||
|
converted["tool_choice"],
|
||||||
|
json!({
|
||||||
|
"type": "allowed_tools", "mode": mode,
|
||||||
|
"tools": [{"type": "function", "name": "lookup"}]
|
||||||
|
})
|
||||||
|
);
|
||||||
|
assert_eq!(converted["tools"].as_array().unwrap().len(), 3);
|
||||||
|
} else {
|
||||||
|
assert_eq!(converted["tool_choice"], mode);
|
||||||
|
assert_eq!(converted["tools"], json!([{"type": "image_generation"}]));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_responses_preserves_requested_encrypted_reasoning_and_replayed_input() {
|
||||||
|
let reasoning = json!({"type": "reasoning", "id": "550e8400-e29b-41d4-a716-446655440000", "summary": [], "encrypted_content": "opaque-xai-state"});
|
||||||
|
let request = json!({
|
||||||
|
"input": [reasoning.clone(), {"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "Previous answer"}]}, {"role": "user", "content": "Continue"}],
|
||||||
|
"include": ["reasoning.encrypted_content"], "store": false
|
||||||
|
});
|
||||||
|
let converted = build_standard_request_body(
|
||||||
|
&request,
|
||||||
|
"openai:responses",
|
||||||
|
"grok-4.6",
|
||||||
|
"xai",
|
||||||
|
"openai:responses",
|
||||||
|
"/v1/responses",
|
||||||
|
true,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(converted["include"], request["include"]);
|
||||||
|
assert_eq!(converted["input"][0], reasoning);
|
||||||
|
assert_eq!(converted["store"], false);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_standard_conversion_strips_unsupported_responses_fields() {
|
||||||
|
let request = json!({
|
||||||
|
"model": "source-model",
|
||||||
|
"messages": [{"role": "user", "content": "Hello xAI"}],
|
||||||
|
"max_tokens": 128,
|
||||||
|
"stop": ["END"],
|
||||||
|
"stream_options": {"include_usage": true},
|
||||||
|
"metadata": {"user_id": "claude-session"},
|
||||||
|
"web_search_options": {"search_context_size": "high"}
|
||||||
|
});
|
||||||
|
let converted = build_standard_request_body(
|
||||||
|
&request,
|
||||||
|
"openai:chat",
|
||||||
|
"grok-4.6",
|
||||||
|
"xai",
|
||||||
|
"openai:responses",
|
||||||
|
"/v1/chat/completions",
|
||||||
|
true,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.expect("chat should convert onto xAI Responses");
|
||||||
|
|
||||||
|
assert_eq!(converted["model"], "grok-4.6");
|
||||||
|
assert!(converted.get("stop").is_none());
|
||||||
|
assert!(converted.get("stream_options").is_none());
|
||||||
|
assert!(converted.get("previous_response_id").is_none());
|
||||||
|
assert!(converted.get("metadata").is_none());
|
||||||
|
assert!(converted.get("input").is_some() || converted.get("messages").is_none());
|
||||||
|
assert_eq!(converted["max_output_tokens"], 128);
|
||||||
|
assert_eq!(converted["tools"][0]["type"], "web_search");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_standard_conversion_covers_claude_and_gemini_clients() {
|
||||||
|
let claude = json!({
|
||||||
|
"model": "claude-sonnet",
|
||||||
|
"max_tokens": 64,
|
||||||
|
"messages": [{"role": "user", "content": "Hello xAI"}],
|
||||||
|
"metadata": {
|
||||||
|
"user_id": "{\"device_id\":\"dev-1\",\"account_uuid\":\"acct-1\",\"session_id\":\"sess-1\"}"
|
||||||
|
},
|
||||||
|
"tools": [
|
||||||
|
{"type": "web_search_20250305", "name": "web_search"},
|
||||||
|
{
|
||||||
|
"name": "lookup",
|
||||||
|
"description": "Look something up",
|
||||||
|
"input_schema": {"type": "object", "properties": {}}
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"tool_choice": {"type": "tool", "name": "web_search"}
|
||||||
|
});
|
||||||
|
let converted = build_standard_request_body(
|
||||||
|
&claude,
|
||||||
|
"claude:messages",
|
||||||
|
"grok-4.6",
|
||||||
|
"xai",
|
||||||
|
"openai:responses",
|
||||||
|
"/v1/messages",
|
||||||
|
true,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.expect("claude should convert onto xAI Responses");
|
||||||
|
assert_eq!(converted["model"], "grok-4.6");
|
||||||
|
assert!(converted.get("metadata").is_none());
|
||||||
|
assert!(converted.get("context_management").is_none());
|
||||||
|
assert!(converted
|
||||||
|
.get("include")
|
||||||
|
.and_then(Value::as_array)
|
||||||
|
.into_iter()
|
||||||
|
.flatten()
|
||||||
|
.any(|item| item == "reasoning.encrypted_content"));
|
||||||
|
assert!(converted["tools"]
|
||||||
|
.as_array()
|
||||||
|
.into_iter()
|
||||||
|
.flatten()
|
||||||
|
.any(|tool| tool["type"] == "web_search"));
|
||||||
|
assert_eq!(converted["tool_choice"]["type"], "allowed_tools");
|
||||||
|
assert!(converted.get("input").is_some());
|
||||||
|
|
||||||
|
let gemini = json!({
|
||||||
|
"model": "gemini-2.5-pro",
|
||||||
|
"contents": [{
|
||||||
|
"role": "user",
|
||||||
|
"parts": [{"text": "Hello xAI"}]
|
||||||
|
}],
|
||||||
|
"tools": [{"googleSearch": {}}]
|
||||||
|
});
|
||||||
|
let converted = build_standard_request_body(
|
||||||
|
&gemini,
|
||||||
|
"gemini:generate_content",
|
||||||
|
"grok-4.6",
|
||||||
|
"xai",
|
||||||
|
"openai:responses",
|
||||||
|
"/v1beta/models/gemini-2.5-pro:generateContent",
|
||||||
|
false,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.expect("gemini should convert onto xAI Responses");
|
||||||
|
assert_eq!(converted["model"], "grok-4.6");
|
||||||
|
assert_eq!(converted["tools"][0]["type"], "web_search");
|
||||||
|
assert!(converted.get("input").is_some());
|
||||||
|
|
||||||
|
let same_format = json!({
|
||||||
|
"model": "grok-4.6",
|
||||||
|
"input": "hello",
|
||||||
|
"previous_response_id": "resp_123",
|
||||||
|
"stop": ["END"],
|
||||||
|
"metadata": {"user_id": "claude-session"}
|
||||||
|
});
|
||||||
|
let converted = build_standard_request_body(
|
||||||
|
&same_format,
|
||||||
|
"openai:responses",
|
||||||
|
"grok-4.6",
|
||||||
|
"xai",
|
||||||
|
"openai:responses",
|
||||||
|
"/v1/responses",
|
||||||
|
true,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.expect("same-format xAI Responses should sanitize in place");
|
||||||
|
assert!(converted.get("previous_response_id").is_none());
|
||||||
|
assert!(converted.get("stop").is_none());
|
||||||
|
assert!(converted.get("metadata").is_none());
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -56,6 +56,10 @@ pub use formats::openai::responses::codex::{
|
|||||||
pub use formats::openai::responses::request::{
|
pub use formats::openai::responses::request::{
|
||||||
validate_openai_responses_request_contract, OpenAiResponsesRequestContractViolation,
|
validate_openai_responses_request_contract, OpenAiResponsesRequestContractViolation,
|
||||||
};
|
};
|
||||||
|
pub use formats::openai::responses::xai::{
|
||||||
|
apply_xai_upstream_payload_edits, apply_xai_upstream_payload_edits_with_client,
|
||||||
|
xai_model_supports_reasoning_effort, xai_supports_native_image_generation,
|
||||||
|
};
|
||||||
pub use formats::openai::responses::{
|
pub use formats::openai::responses::{
|
||||||
normalize_openai_responses_message_item_ids, openai_responses_message_item_id,
|
normalize_openai_responses_message_item_ids, openai_responses_message_item_id,
|
||||||
openai_responses_request_operation, openai_responses_synthetic_reasoning_item_id,
|
openai_responses_request_operation, openai_responses_synthetic_reasoning_item_id,
|
||||||
|
|||||||
@@ -102,6 +102,11 @@ INNER JOIN LATERAL (
|
|||||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||||
AND LOWER($3) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
|
AND LOWER($3) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
|
||||||
)
|
)
|
||||||
|
OR (
|
||||||
|
LOWER(BTRIM(p.provider_type)) = 'xai'
|
||||||
|
AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')
|
||||||
|
AND LOWER($3) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video')
|
||||||
|
)
|
||||||
OR (
|
OR (
|
||||||
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
||||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||||
@@ -127,7 +132,8 @@ INNER JOIN LATERAL (
|
|||||||
'vertex_ai',
|
'vertex_ai',
|
||||||
'antigravity',
|
'antigravity',
|
||||||
'kiro',
|
'kiro',
|
||||||
'windsurf'
|
'windsurf',
|
||||||
|
'xai'
|
||||||
)
|
)
|
||||||
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
||||||
)
|
)
|
||||||
@@ -187,6 +193,11 @@ WHERE p.is_active = TRUE
|
|||||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||||
AND LOWER($3) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
|
AND LOWER($3) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
|
||||||
)
|
)
|
||||||
|
OR (
|
||||||
|
LOWER(BTRIM(p.provider_type)) = 'xai'
|
||||||
|
AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')
|
||||||
|
AND LOWER($3) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video')
|
||||||
|
)
|
||||||
OR (
|
OR (
|
||||||
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
||||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||||
@@ -212,7 +223,8 @@ WHERE p.is_active = TRUE
|
|||||||
'vertex_ai',
|
'vertex_ai',
|
||||||
'antigravity',
|
'antigravity',
|
||||||
'kiro',
|
'kiro',
|
||||||
'windsurf'
|
'windsurf',
|
||||||
|
'xai'
|
||||||
)
|
)
|
||||||
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
||||||
)
|
)
|
||||||
@@ -365,6 +377,11 @@ INNER JOIN LATERAL (
|
|||||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||||
AND LOWER($4) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
|
AND LOWER($4) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
|
||||||
)
|
)
|
||||||
|
OR (
|
||||||
|
LOWER(BTRIM(p.provider_type)) = 'xai'
|
||||||
|
AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')
|
||||||
|
AND LOWER($4) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video')
|
||||||
|
)
|
||||||
OR (
|
OR (
|
||||||
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
||||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||||
@@ -390,7 +407,8 @@ INNER JOIN LATERAL (
|
|||||||
'vertex_ai',
|
'vertex_ai',
|
||||||
'antigravity',
|
'antigravity',
|
||||||
'kiro',
|
'kiro',
|
||||||
'windsurf'
|
'windsurf',
|
||||||
|
'xai'
|
||||||
)
|
)
|
||||||
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
||||||
)
|
)
|
||||||
@@ -451,6 +469,11 @@ WHERE p.is_active = TRUE
|
|||||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||||
AND LOWER($4) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
|
AND LOWER($4) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
|
||||||
)
|
)
|
||||||
|
OR (
|
||||||
|
LOWER(BTRIM(p.provider_type)) = 'xai'
|
||||||
|
AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')
|
||||||
|
AND LOWER($4) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video')
|
||||||
|
)
|
||||||
OR (
|
OR (
|
||||||
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
||||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||||
@@ -476,7 +499,8 @@ WHERE p.is_active = TRUE
|
|||||||
'vertex_ai',
|
'vertex_ai',
|
||||||
'antigravity',
|
'antigravity',
|
||||||
'kiro',
|
'kiro',
|
||||||
'windsurf'
|
'windsurf',
|
||||||
|
'xai'
|
||||||
)
|
)
|
||||||
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
||||||
)
|
)
|
||||||
@@ -632,11 +656,16 @@ WHERE p.is_active = TRUE
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
OR (
|
OR (
|
||||||
LOWER(BTRIM(p.provider_type)) = 'grok'
|
LOWER(BTRIM(p.provider_type)) = 'grok'
|
||||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||||
AND LOWER($6) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
|
AND LOWER($6) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
|
||||||
)
|
)
|
||||||
|
OR (
|
||||||
|
LOWER(BTRIM(p.provider_type)) = 'xai'
|
||||||
|
AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')
|
||||||
|
AND LOWER($6) IN ('openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video')
|
||||||
|
)
|
||||||
OR (
|
OR (
|
||||||
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
||||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||||
@@ -662,7 +691,8 @@ WHERE p.is_active = TRUE
|
|||||||
'vertex_ai',
|
'vertex_ai',
|
||||||
'antigravity',
|
'antigravity',
|
||||||
'kiro',
|
'kiro',
|
||||||
'windsurf'
|
'windsurf',
|
||||||
|
'xai'
|
||||||
)
|
)
|
||||||
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
||||||
)
|
)
|
||||||
@@ -1717,6 +1747,24 @@ mod tests {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn candidate_selection_sql_allows_xai_oauth_responses_auth() {
|
||||||
|
let requested_model_sql = requested_model_selection_sql();
|
||||||
|
for sql in [
|
||||||
|
LIST_FOR_EXACT_API_FORMAT_SQL,
|
||||||
|
LIST_FOR_EXACT_API_FORMAT_AND_GLOBAL_MODEL_SQL,
|
||||||
|
LIST_POOL_KEYS_FOR_GROUP_SQL,
|
||||||
|
requested_model_sql.as_str(),
|
||||||
|
] {
|
||||||
|
assert!(sql.contains("LOWER(BTRIM(p.provider_type)) = 'xai'"));
|
||||||
|
assert!(sql.contains("LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')"));
|
||||||
|
assert!(sql.contains(
|
||||||
|
"'openai:responses', 'openai:responses:compact', 'openai:image', 'openai:video'"
|
||||||
|
));
|
||||||
|
assert!(sql.contains("'xai'"));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn candidate_selection_sql_allows_windsurf_openai_chat_managed_keys() {
|
fn candidate_selection_sql_allows_windsurf_openai_chat_managed_keys() {
|
||||||
let requested_model_sql = requested_model_selection_sql();
|
let requested_model_sql = requested_model_selection_sql();
|
||||||
|
|||||||
@@ -346,6 +346,16 @@ fn key_auth_channel_matches(row: &StoredMinimalCandidateSelectionRow, api_format
|
|||||||
"openai:chat" | "openai:responses" | "claude:messages" | "openai:image"
|
"openai:chat" | "openai:responses" | "claude:messages" | "openai:image"
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
"xai" => {
|
||||||
|
matches!(auth_type.as_str(), "oauth" | "bearer" | "api_key")
|
||||||
|
&& matches!(
|
||||||
|
api_format.as_str(),
|
||||||
|
"openai:responses"
|
||||||
|
| "openai:responses:compact"
|
||||||
|
| "openai:image"
|
||||||
|
| "openai:video"
|
||||||
|
)
|
||||||
|
}
|
||||||
"windsurf" => {
|
"windsurf" => {
|
||||||
matches!(auth_type.as_str(), "oauth" | "api_key" | "bearer")
|
matches!(auth_type.as_str(), "oauth" | "api_key" | "bearer")
|
||||||
&& api_format == "openai:chat"
|
&& api_format == "openai:chat"
|
||||||
@@ -591,6 +601,59 @@ mod tests {
|
|||||||
assert_eq!(rows[0].global_model_name, "grok-4.20-0309-non-reasoning");
|
assert_eq!(rows[0].global_model_name, "grok-4.20-0309-non-reasoning");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn includes_xai_oauth_rows_for_responses_models() {
|
||||||
|
let mut row = sample_row("provider-xai", "openai:responses", "grok-4", 10);
|
||||||
|
row.provider_type = "xai".to_string();
|
||||||
|
row.provider_name = "xai".to_string();
|
||||||
|
row.key_auth_type = "oauth".to_string();
|
||||||
|
row.key_api_formats = Some(vec![
|
||||||
|
"openai:responses".to_string(),
|
||||||
|
"openai:responses:compact".to_string(),
|
||||||
|
]);
|
||||||
|
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![row]);
|
||||||
|
|
||||||
|
let rows = repository
|
||||||
|
.list_for_exact_api_format("openai:responses")
|
||||||
|
.await
|
||||||
|
.expect("list should succeed");
|
||||||
|
|
||||||
|
assert_eq!(rows.len(), 1);
|
||||||
|
assert_eq!(rows[0].provider_type, "xai");
|
||||||
|
assert_eq!(rows[0].global_model_name, "grok-4");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn includes_xai_oauth_rows_for_image_and_video_models() {
|
||||||
|
let mut image = sample_row("provider-xai", "openai:image", "grok-imagine-image", 10);
|
||||||
|
image.provider_type = "xai".to_string();
|
||||||
|
image.provider_name = "xai".to_string();
|
||||||
|
image.key_auth_type = "oauth".to_string();
|
||||||
|
image.key_api_formats = Some(vec!["openai:image".to_string(), "openai:video".to_string()]);
|
||||||
|
|
||||||
|
let mut video = image.clone();
|
||||||
|
video.endpoint_id = "endpoint-video".to_string();
|
||||||
|
video.endpoint_api_format = "openai:video".to_string();
|
||||||
|
video.global_model_name = "grok-imagine-video".to_string();
|
||||||
|
video.model_provider_model_name = "grok-imagine-video".to_string();
|
||||||
|
|
||||||
|
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![image, video]);
|
||||||
|
|
||||||
|
let image_rows = repository
|
||||||
|
.list_for_exact_api_format("openai:image")
|
||||||
|
.await
|
||||||
|
.expect("list should succeed");
|
||||||
|
assert_eq!(image_rows.len(), 1);
|
||||||
|
assert_eq!(image_rows[0].global_model_name, "grok-imagine-image");
|
||||||
|
|
||||||
|
let video_rows = repository
|
||||||
|
.list_for_exact_api_format("openai:video")
|
||||||
|
.await
|
||||||
|
.expect("list should succeed");
|
||||||
|
assert_eq!(video_rows.len(), 1);
|
||||||
|
assert_eq!(video_rows[0].global_model_name, "grok-imagine-video");
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn requested_model_filter_respects_endpoint_scoped_default_mapping() {
|
async fn requested_model_filter_respects_endpoint_scoped_default_mapping() {
|
||||||
let mut selected = sample_row("provider-1", "openai:chat", "deepseek-v4-pro", 10);
|
let mut selected = sample_row("provider-1", "openai:chat", "deepseek-v4-pro", 10);
|
||||||
|
|||||||
@@ -546,7 +546,7 @@ pub fn endpoint_supports_rust_models_fetch(api_format: &str) -> bool {
|
|||||||
pub fn provider_type_uses_preset_models(provider_type: &str) -> bool {
|
pub fn provider_type_uses_preset_models(provider_type: &str) -> bool {
|
||||||
matches!(
|
matches!(
|
||||||
provider_type.trim().to_ascii_lowercase().as_str(),
|
provider_type.trim().to_ascii_lowercase().as_str(),
|
||||||
"claude_code" | "gemini_cli" | "grok"
|
"claude_code" | "gemini_cli" | "grok" | "xai"
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -604,6 +604,22 @@ pub fn preset_models_for_provider(provider_type: &str) -> Option<Vec<Value>> {
|
|||||||
preset_model("grok-imagine-image-pro", "xai", "Grok Imagine Image Pro", "openai:image"),
|
preset_model("grok-imagine-image-pro", "xai", "Grok Imagine Image Pro", "openai:image"),
|
||||||
preset_model("grok-imagine-image-edit", "xai", "Grok Imagine Image Edit", "openai:image"),
|
preset_model("grok-imagine-image-edit", "xai", "Grok Imagine Image Edit", "openai:image"),
|
||||||
],
|
],
|
||||||
|
"xai" => vec![
|
||||||
|
preset_model("grok-4.6", "xai", "Grok 4.6", "openai:responses"),
|
||||||
|
preset_model("grok-build-0.1", "xai", "Grok Build 0.1", "openai:responses"),
|
||||||
|
preset_model("grok-4.5", "xai", "Grok 4.5", "openai:responses"),
|
||||||
|
preset_model("grok-4.3", "xai", "Grok 4.3", "openai:responses"),
|
||||||
|
preset_model("grok-4.20-0309-reasoning", "xai", "Grok 4.20 0309 Reasoning", "openai:responses"),
|
||||||
|
preset_model("grok-4.20-0309-non-reasoning", "xai", "Grok 4.20 0309 Non-Reasoning", "openai:responses"),
|
||||||
|
preset_model("grok-4.20-multi-agent-0309", "xai", "Grok 4.20 Multi-Agent 0309", "openai:responses"),
|
||||||
|
preset_model("grok-3-mini", "xai", "Grok 3 Mini", "openai:responses"),
|
||||||
|
preset_model("grok-3-mini-fast", "xai", "Grok 3 Mini Fast", "openai:responses"),
|
||||||
|
preset_model("grok-composer-2.5-fast", "xai", "Grok Composer 2.5 Fast", "openai:responses"),
|
||||||
|
preset_model("grok-imagine-image", "xai", "Grok Imagine Image", "openai:image"),
|
||||||
|
preset_model("grok-imagine-image-quality", "xai", "Grok Imagine Image Quality", "openai:image"),
|
||||||
|
preset_model("grok-imagine-video", "xai", "Grok Imagine Video", "openai:video"),
|
||||||
|
preset_model("grok-imagine-video-1.5", "xai", "Grok Imagine Video 1.5", "openai:video"),
|
||||||
|
],
|
||||||
_ => return None,
|
_ => return None,
|
||||||
};
|
};
|
||||||
Some(models)
|
Some(models)
|
||||||
@@ -1977,4 +1993,39 @@ mod tests {
|
|||||||
assert_eq!(models[15]["api_formats"], json!(["openai:image"]));
|
assert_eq!(models[15]["api_formats"], json!(["openai:image"]));
|
||||||
assert_eq!(models[18]["api_formats"], json!(["openai:image"]));
|
assert_eq!(models[18]["api_formats"], json!(["openai:image"]));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn preset_models_cover_xai_cli_catalog() {
|
||||||
|
let models = preset_models_for_provider("xai").expect("preset models should exist");
|
||||||
|
let model_ids = models
|
||||||
|
.iter()
|
||||||
|
.map(|model| model["id"].as_str().expect("model id"))
|
||||||
|
.collect::<Vec<_>>();
|
||||||
|
assert_eq!(
|
||||||
|
model_ids,
|
||||||
|
vec![
|
||||||
|
"grok-4.6",
|
||||||
|
"grok-build-0.1",
|
||||||
|
"grok-4.5",
|
||||||
|
"grok-4.3",
|
||||||
|
"grok-4.20-0309-reasoning",
|
||||||
|
"grok-4.20-0309-non-reasoning",
|
||||||
|
"grok-4.20-multi-agent-0309",
|
||||||
|
"grok-3-mini",
|
||||||
|
"grok-3-mini-fast",
|
||||||
|
"grok-composer-2.5-fast",
|
||||||
|
"grok-imagine-image",
|
||||||
|
"grok-imagine-image-quality",
|
||||||
|
"grok-imagine-video",
|
||||||
|
"grok-imagine-video-1.5",
|
||||||
|
]
|
||||||
|
);
|
||||||
|
assert!(models.iter().all(|model| model["owned_by"] == json!("xai")));
|
||||||
|
assert_eq!(models[0]["api_formats"], json!(["openai:responses"]));
|
||||||
|
assert_eq!(models[10]["api_formats"], json!(["openai:image"]));
|
||||||
|
assert_eq!(models[12]["api_formats"], json!(["openai:video"]));
|
||||||
|
assert!(models
|
||||||
|
.iter()
|
||||||
|
.any(|model| model["id"] == "grok-imagine-image"));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -150,6 +150,27 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[
|
|||||||
uses_json_payload: false,
|
uses_json_payload: false,
|
||||||
include_scope_in_token_request: true,
|
include_scope_in_token_request: true,
|
||||||
},
|
},
|
||||||
|
GenericProviderOAuthTemplate {
|
||||||
|
provider_type: "xai",
|
||||||
|
display_name: "xAI",
|
||||||
|
authorize_url: "https://auth.x.ai/oauth2/device/code",
|
||||||
|
token_url: "https://auth.x.ai/oauth2/token",
|
||||||
|
client_id: "b1a00492-073a-47ea-816f-4c329264a828",
|
||||||
|
client_id_env: None,
|
||||||
|
client_secret_env: None,
|
||||||
|
scopes: &[
|
||||||
|
"openid",
|
||||||
|
"profile",
|
||||||
|
"email",
|
||||||
|
"offline_access",
|
||||||
|
"grok-cli:access",
|
||||||
|
"api:access",
|
||||||
|
],
|
||||||
|
redirect_uri: "",
|
||||||
|
use_pkce: false,
|
||||||
|
uses_json_payload: false,
|
||||||
|
include_scope_in_token_request: false,
|
||||||
|
},
|
||||||
];
|
];
|
||||||
|
|
||||||
#[derive(Clone)]
|
#[derive(Clone)]
|
||||||
@@ -212,6 +233,10 @@ impl GenericProviderOAuthAdapter {
|
|||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub(super) fn token_url_for_provider(&self) -> String {
|
||||||
|
self.token_url()
|
||||||
|
}
|
||||||
|
|
||||||
fn token_url(&self) -> String {
|
fn token_url(&self) -> String {
|
||||||
self.token_url_override
|
self.token_url_override
|
||||||
.clone()
|
.clone()
|
||||||
@@ -389,7 +414,10 @@ impl GenericProviderOAuthAdapter {
|
|||||||
self.token_set_from_payload(payload)
|
self.token_set_from_payload(payload)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn token_set_from_payload(&self, payload: Value) -> Result<ProviderOAuthTokenSet, OAuthError> {
|
pub(super) fn token_set_from_payload(
|
||||||
|
&self,
|
||||||
|
payload: Value,
|
||||||
|
) -> Result<ProviderOAuthTokenSet, OAuthError> {
|
||||||
let token_set = OAuthTokenSet::from_token_payload(payload.clone())
|
let token_set = OAuthTokenSet::from_token_payload(payload.clone())
|
||||||
.ok_or_else(|| OAuthError::invalid_response("token response missing access_token"))?;
|
.ok_or_else(|| OAuthError::invalid_response("token response missing access_token"))?;
|
||||||
let mut auth_config = serde_json::Map::new();
|
let mut auth_config = serde_json::Map::new();
|
||||||
@@ -945,6 +973,7 @@ mod tests {
|
|||||||
fn resolves_generic_provider_templates() {
|
fn resolves_generic_provider_templates() {
|
||||||
assert!(template_for_provider_type("codex").is_some());
|
assert!(template_for_provider_type("codex").is_some());
|
||||||
assert!(template_for_provider_type("claude_code").is_some());
|
assert!(template_for_provider_type("claude_code").is_some());
|
||||||
|
assert!(template_for_provider_type("xai").is_some());
|
||||||
assert!(template_for_provider_type("kiro").is_none());
|
assert!(template_for_provider_type("kiro").is_none());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ mod codex;
|
|||||||
mod generic;
|
mod generic;
|
||||||
mod kiro;
|
mod kiro;
|
||||||
mod windsurf;
|
mod windsurf;
|
||||||
|
mod xai;
|
||||||
|
|
||||||
pub use antigravity::{AntigravityProviderOAuthAdapter, ANTIGRAVITY_USER_INFO_URL};
|
pub use antigravity::{AntigravityProviderOAuthAdapter, ANTIGRAVITY_USER_INFO_URL};
|
||||||
pub use claude_code::{
|
pub use claude_code::{
|
||||||
@@ -27,3 +28,7 @@ pub use windsurf::{
|
|||||||
WindsurfProviderOAuthAdapter, WINDSURF_CLIENT_ID, WINDSURF_PROVIDER_TYPE,
|
WindsurfProviderOAuthAdapter, WINDSURF_CLIENT_ID, WINDSURF_PROVIDER_TYPE,
|
||||||
WINDSURF_SHOW_AUTH_TOKEN_REDIRECT, WINDSURF_SIGNIN_URL,
|
WINDSURF_SHOW_AUTH_TOKEN_REDIRECT, WINDSURF_SIGNIN_URL,
|
||||||
};
|
};
|
||||||
|
pub use xai::{
|
||||||
|
XaiDevicePollOutcome, XaiProviderOAuthAdapter, XAI_CLIENT_ID, XAI_DEVICE_CODE_GRANT_TYPE,
|
||||||
|
XAI_DEVICE_CODE_URL, XAI_OAUTH_SCOPES, XAI_PROVIDER_TYPE, XAI_TOKEN_URL,
|
||||||
|
};
|
||||||
|
|||||||
@@ -0,0 +1,668 @@
|
|||||||
|
use super::generic::{template_for_provider_type, GenericProviderOAuthAdapter};
|
||||||
|
use crate::core::{
|
||||||
|
current_unix_secs, redacted_oauth_error_body_excerpt, OAuthDeviceAuthorization, OAuthError,
|
||||||
|
};
|
||||||
|
use crate::network::{OAuthHttpExecutor, OAuthHttpRequest};
|
||||||
|
use crate::provider::{
|
||||||
|
ProviderOAuthAccount, ProviderOAuthAdapter, ProviderOAuthCapabilities,
|
||||||
|
ProviderOAuthImportInput, ProviderOAuthRequestAuth, ProviderOAuthTokenSet,
|
||||||
|
ProviderOAuthTransportContext,
|
||||||
|
};
|
||||||
|
use async_trait::async_trait;
|
||||||
|
use serde_json::{json, Map, Value};
|
||||||
|
use std::collections::BTreeMap;
|
||||||
|
use url::form_urlencoded;
|
||||||
|
|
||||||
|
pub const XAI_PROVIDER_TYPE: &str = "xai";
|
||||||
|
pub const XAI_DEVICE_CODE_URL: &str = "https://auth.x.ai/oauth2/device/code";
|
||||||
|
pub const XAI_TOKEN_URL: &str = "https://auth.x.ai/oauth2/token";
|
||||||
|
pub const XAI_CLIENT_ID: &str = "b1a00492-073a-47ea-816f-4c329264a828";
|
||||||
|
pub const XAI_OAUTH_SCOPES: &[&str] = &[
|
||||||
|
"openid",
|
||||||
|
"profile",
|
||||||
|
"email",
|
||||||
|
"offline_access",
|
||||||
|
"grok-cli:access",
|
||||||
|
"api:access",
|
||||||
|
];
|
||||||
|
pub const XAI_DEVICE_CODE_GRANT_TYPE: &str = "urn:ietf:params:oauth:grant-type:device_code";
|
||||||
|
|
||||||
|
const DEFAULT_DEVICE_EXPIRES_IN_SECS: u64 = 600;
|
||||||
|
const DEFAULT_DEVICE_POLL_INTERVAL_SECS: u64 = 5;
|
||||||
|
|
||||||
|
#[derive(Debug, Clone, PartialEq)]
|
||||||
|
pub enum XaiDevicePollOutcome {
|
||||||
|
Pending,
|
||||||
|
SlowDown,
|
||||||
|
Authorized(Box<ProviderOAuthTokenSet>),
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub struct XaiProviderOAuthAdapter {
|
||||||
|
inner: GenericProviderOAuthAdapter,
|
||||||
|
device_url_override: Option<String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl std::fmt::Debug for XaiProviderOAuthAdapter {
|
||||||
|
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||||
|
formatter
|
||||||
|
.debug_struct("XaiProviderOAuthAdapter")
|
||||||
|
.field(
|
||||||
|
"has_device_url_override",
|
||||||
|
&self.device_url_override.is_some(),
|
||||||
|
)
|
||||||
|
.finish_non_exhaustive()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Default for XaiProviderOAuthAdapter {
|
||||||
|
fn default() -> Self {
|
||||||
|
Self {
|
||||||
|
inner: GenericProviderOAuthAdapter::new(
|
||||||
|
template_for_provider_type(XAI_PROVIDER_TYPE).expect("xai template should exist"),
|
||||||
|
),
|
||||||
|
device_url_override: None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl XaiProviderOAuthAdapter {
|
||||||
|
pub fn with_endpoint_overrides(
|
||||||
|
mut self,
|
||||||
|
device_url: impl Into<String>,
|
||||||
|
token_url: impl Into<String>,
|
||||||
|
) -> Self {
|
||||||
|
self.device_url_override = Some(device_url.into());
|
||||||
|
self.inner = self.inner.with_token_url_override(token_url);
|
||||||
|
self
|
||||||
|
}
|
||||||
|
|
||||||
|
fn device_url(&self) -> String {
|
||||||
|
self.device_url_override
|
||||||
|
.clone()
|
||||||
|
.unwrap_or_else(|| XAI_DEVICE_CODE_URL.to_string())
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn start_device_flow(
|
||||||
|
&self,
|
||||||
|
executor: &dyn OAuthHttpExecutor,
|
||||||
|
ctx: &ProviderOAuthTransportContext,
|
||||||
|
) -> Result<OAuthDeviceAuthorization, OAuthError> {
|
||||||
|
let form = form_urlencoded::Serializer::new(String::new())
|
||||||
|
.append_pair("client_id", XAI_CLIENT_ID)
|
||||||
|
.append_pair("scope", &XAI_OAUTH_SCOPES.join(" "))
|
||||||
|
.finish()
|
||||||
|
.into_bytes();
|
||||||
|
let response = executor
|
||||||
|
.execute(OAuthHttpRequest {
|
||||||
|
request_id: "provider-oauth:xai-device-code".to_string(),
|
||||||
|
method: reqwest::Method::POST,
|
||||||
|
url: self.device_url(),
|
||||||
|
headers: form_headers(),
|
||||||
|
content_type: Some("application/x-www-form-urlencoded".to_string()),
|
||||||
|
json_body: None,
|
||||||
|
body_bytes: Some(form),
|
||||||
|
network: ctx.network.clone(),
|
||||||
|
transport_profile: None,
|
||||||
|
})
|
||||||
|
.await?;
|
||||||
|
if !(200..300).contains(&response.status_code) {
|
||||||
|
return Err(OAuthError::HttpStatus {
|
||||||
|
status_code: response.status_code,
|
||||||
|
body_excerpt: redacted_oauth_error_body_excerpt(&response.body_text),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
let payload = response_json(&response)
|
||||||
|
.ok_or_else(|| OAuthError::invalid_response("xAI device code response is not json"))?;
|
||||||
|
parse_device_authorization(&payload)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn poll_device_token(
|
||||||
|
&self,
|
||||||
|
executor: &dyn OAuthHttpExecutor,
|
||||||
|
ctx: &ProviderOAuthTransportContext,
|
||||||
|
device_code: &str,
|
||||||
|
) -> Result<XaiDevicePollOutcome, OAuthError> {
|
||||||
|
let device_code = device_code.trim();
|
||||||
|
if device_code.is_empty() {
|
||||||
|
return Err(OAuthError::invalid_request("xAI device_code is required"));
|
||||||
|
}
|
||||||
|
let form = form_urlencoded::Serializer::new(String::new())
|
||||||
|
.append_pair("grant_type", XAI_DEVICE_CODE_GRANT_TYPE)
|
||||||
|
.append_pair("device_code", device_code)
|
||||||
|
.append_pair("client_id", XAI_CLIENT_ID)
|
||||||
|
.finish()
|
||||||
|
.into_bytes();
|
||||||
|
let response = executor
|
||||||
|
.execute(OAuthHttpRequest {
|
||||||
|
request_id: "provider-oauth:xai-device-token".to_string(),
|
||||||
|
method: reqwest::Method::POST,
|
||||||
|
url: self.inner.token_url_for_provider(),
|
||||||
|
headers: form_headers(),
|
||||||
|
content_type: Some("application/x-www-form-urlencoded".to_string()),
|
||||||
|
json_body: None,
|
||||||
|
body_bytes: Some(form),
|
||||||
|
network: ctx.network.clone(),
|
||||||
|
transport_profile: None,
|
||||||
|
})
|
||||||
|
.await?;
|
||||||
|
let payload = response_json(&response);
|
||||||
|
if let Some(error_code) = payload.as_ref().and_then(oauth_error_code) {
|
||||||
|
return match error_code.as_str() {
|
||||||
|
"authorization_pending" => Ok(XaiDevicePollOutcome::Pending),
|
||||||
|
"slow_down" => Ok(XaiDevicePollOutcome::SlowDown),
|
||||||
|
"expired_token" => Err(OAuthError::invalid_request("xAI device code expired")),
|
||||||
|
"access_denied" => Err(OAuthError::invalid_request(
|
||||||
|
"xAI device authorization denied",
|
||||||
|
)),
|
||||||
|
other => Err(OAuthError::invalid_response(format!(
|
||||||
|
"xAI device token error: {other}"
|
||||||
|
))),
|
||||||
|
};
|
||||||
|
}
|
||||||
|
if !(200..300).contains(&response.status_code) {
|
||||||
|
return Err(OAuthError::HttpStatus {
|
||||||
|
status_code: response.status_code,
|
||||||
|
body_excerpt: redacted_oauth_error_body_excerpt(&response.body_text),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
let payload = payload
|
||||||
|
.ok_or_else(|| OAuthError::invalid_response("xAI device token response is not json"))?;
|
||||||
|
let mut token_set = self.inner.token_set_from_payload(payload)?;
|
||||||
|
let raw_payload = token_set.token_set.raw_payload.clone();
|
||||||
|
mark_oauth_auth_config(&mut token_set.auth_config);
|
||||||
|
enrich_xai_identity(&mut token_set.auth_config, raw_payload.as_ref());
|
||||||
|
Ok(XaiDevicePollOutcome::Authorized(Box::new(token_set)))
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn import_raw_api_key(
|
||||||
|
&self,
|
||||||
|
input: &ProviderOAuthImportInput,
|
||||||
|
api_key: &str,
|
||||||
|
) -> Result<ProviderOAuthTokenSet, OAuthError> {
|
||||||
|
let api_key = api_key.trim();
|
||||||
|
if api_key.is_empty() {
|
||||||
|
return Err(OAuthError::invalid_request("xAI api_key is required"));
|
||||||
|
}
|
||||||
|
let mut auth_config = Map::new();
|
||||||
|
auth_config.insert("provider_type".to_string(), json!(XAI_PROVIDER_TYPE));
|
||||||
|
auth_config.insert("auth_method".to_string(), json!("api_key"));
|
||||||
|
auth_config.insert("using_api".to_string(), json!(true));
|
||||||
|
auth_config.insert("updated_at".to_string(), json!(current_unix_secs()));
|
||||||
|
if let Some(name) = input
|
||||||
|
.name
|
||||||
|
.as_deref()
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
{
|
||||||
|
auth_config.insert("name".to_string(), json!(name));
|
||||||
|
}
|
||||||
|
Ok(ProviderOAuthTokenSet {
|
||||||
|
token_set: crate::core::OAuthTokenSet {
|
||||||
|
access_token: api_key.to_string(),
|
||||||
|
refresh_token: None,
|
||||||
|
token_type: Some("Bearer".to_string()),
|
||||||
|
scope: None,
|
||||||
|
expires_at_unix_secs: None,
|
||||||
|
raw_payload: None,
|
||||||
|
},
|
||||||
|
auth_config: Value::Object(auth_config),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl ProviderOAuthAdapter for XaiProviderOAuthAdapter {
|
||||||
|
fn provider_type(&self) -> &'static str {
|
||||||
|
XAI_PROVIDER_TYPE
|
||||||
|
}
|
||||||
|
|
||||||
|
fn capabilities(&self) -> ProviderOAuthCapabilities {
|
||||||
|
ProviderOAuthCapabilities {
|
||||||
|
supports_authorization_code: false,
|
||||||
|
supports_cookie_authorization: false,
|
||||||
|
supports_refresh_token_import: true,
|
||||||
|
supports_batch_import: true,
|
||||||
|
supports_device_flow: true,
|
||||||
|
supports_account_probe: false,
|
||||||
|
rotates_refresh_token: true,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn import_credentials(
|
||||||
|
&self,
|
||||||
|
executor: &dyn OAuthHttpExecutor,
|
||||||
|
ctx: &ProviderOAuthTransportContext,
|
||||||
|
input: ProviderOAuthImportInput,
|
||||||
|
) -> Result<ProviderOAuthTokenSet, OAuthError> {
|
||||||
|
if let Some(api_key) =
|
||||||
|
raw_credential_string(input.raw_credentials.as_ref(), &["api_key", "apiKey"])
|
||||||
|
{
|
||||||
|
return self.import_raw_api_key(&input, &api_key).await;
|
||||||
|
}
|
||||||
|
let refresh_token = input
|
||||||
|
.refresh_token
|
||||||
|
.as_deref()
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.map(ToOwned::to_owned)
|
||||||
|
.or_else(|| {
|
||||||
|
raw_credential_string(
|
||||||
|
input.raw_credentials.as_ref(),
|
||||||
|
&["refresh_token", "refreshToken"],
|
||||||
|
)
|
||||||
|
});
|
||||||
|
if let Some(refresh_token) = refresh_token {
|
||||||
|
let mut imported = self
|
||||||
|
.inner
|
||||||
|
.import_credentials(
|
||||||
|
executor,
|
||||||
|
ctx,
|
||||||
|
ProviderOAuthImportInput {
|
||||||
|
refresh_token: Some(refresh_token),
|
||||||
|
..input
|
||||||
|
},
|
||||||
|
)
|
||||||
|
.await?;
|
||||||
|
let raw_payload = imported.token_set.raw_payload.clone();
|
||||||
|
mark_oauth_auth_config(&mut imported.auth_config);
|
||||||
|
enrich_xai_identity(&mut imported.auth_config, raw_payload.as_ref());
|
||||||
|
return Ok(imported);
|
||||||
|
}
|
||||||
|
if let Some(access_token) = raw_credential_string(
|
||||||
|
input.raw_credentials.as_ref(),
|
||||||
|
&["access_token", "accessToken"],
|
||||||
|
) {
|
||||||
|
return self.import_raw_api_key(&input, &access_token).await;
|
||||||
|
}
|
||||||
|
Err(OAuthError::invalid_request(
|
||||||
|
"xAI credentials require api_key, access_token, or refresh_token",
|
||||||
|
))
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn refresh(
|
||||||
|
&self,
|
||||||
|
executor: &dyn OAuthHttpExecutor,
|
||||||
|
ctx: &ProviderOAuthTransportContext,
|
||||||
|
account: &ProviderOAuthAccount,
|
||||||
|
) -> Result<ProviderOAuthTokenSet, OAuthError> {
|
||||||
|
let mut refreshed = self.inner.refresh(executor, ctx, account).await?;
|
||||||
|
let raw_payload = refreshed.token_set.raw_payload.clone();
|
||||||
|
mark_oauth_auth_config(&mut refreshed.auth_config);
|
||||||
|
enrich_xai_identity(&mut refreshed.auth_config, raw_payload.as_ref());
|
||||||
|
Ok(refreshed)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn resolve_request_auth(
|
||||||
|
&self,
|
||||||
|
account: &ProviderOAuthAccount,
|
||||||
|
) -> Result<ProviderOAuthRequestAuth, OAuthError> {
|
||||||
|
self.inner.resolve_request_auth(account)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn account_fingerprint(&self, account: &ProviderOAuthAccount) -> Option<String> {
|
||||||
|
self.inner.account_fingerprint(account)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn mark_oauth_auth_config(auth_config: &mut Value) {
|
||||||
|
let Some(object) = auth_config.as_object_mut() else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
object.insert("provider_type".to_string(), json!(XAI_PROVIDER_TYPE));
|
||||||
|
object.insert("auth_method".to_string(), json!("oauth"));
|
||||||
|
object.insert("using_api".to_string(), json!(false));
|
||||||
|
}
|
||||||
|
|
||||||
|
fn enrich_xai_identity(auth_config: &mut Value, raw_payload: Option<&Value>) {
|
||||||
|
let Some(object) = auth_config.as_object_mut() else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
let id_token = raw_payload
|
||||||
|
.and_then(|payload| payload.get("id_token"))
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty());
|
||||||
|
if let Some(id_token) = id_token {
|
||||||
|
object
|
||||||
|
.entry("id_token".to_string())
|
||||||
|
.or_insert_with(|| json!(id_token));
|
||||||
|
if let Some(claims) = decode_jwt_claims(id_token) {
|
||||||
|
if !object.contains_key("email") {
|
||||||
|
if let Some(email) = claims.get("email").and_then(Value::as_str) {
|
||||||
|
let email = email.trim();
|
||||||
|
if !email.is_empty() {
|
||||||
|
object.insert("email".to_string(), json!(email));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !object.contains_key("sub") {
|
||||||
|
if let Some(sub) = claims.get("sub").and_then(Value::as_str) {
|
||||||
|
let sub = sub.trim();
|
||||||
|
if !sub.is_empty() {
|
||||||
|
object.insert("sub".to_string(), json!(sub));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn parse_device_authorization(payload: &Value) -> Result<OAuthDeviceAuthorization, OAuthError> {
|
||||||
|
let device_code =
|
||||||
|
json_non_empty_string(payload, &["device_code", "deviceCode"]).ok_or_else(|| {
|
||||||
|
OAuthError::invalid_response("xAI device code response missing device_code")
|
||||||
|
})?;
|
||||||
|
let user_code =
|
||||||
|
json_non_empty_string(payload, &["user_code", "userCode"]).ok_or_else(|| {
|
||||||
|
OAuthError::invalid_response("xAI device code response missing user_code")
|
||||||
|
})?;
|
||||||
|
let verification_uri = json_non_empty_string(
|
||||||
|
payload,
|
||||||
|
&["verification_uri", "verificationUri", "verification_url"],
|
||||||
|
)
|
||||||
|
.unwrap_or_default();
|
||||||
|
let verification_uri_complete = json_non_empty_string(
|
||||||
|
payload,
|
||||||
|
&[
|
||||||
|
"verification_uri_complete",
|
||||||
|
"verificationUriComplete",
|
||||||
|
"verification_url_complete",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
.unwrap_or_else(|| verification_uri.clone());
|
||||||
|
if verification_uri.is_empty() && verification_uri_complete.is_empty() {
|
||||||
|
return Err(OAuthError::invalid_response(
|
||||||
|
"xAI device code response missing verification URI",
|
||||||
|
));
|
||||||
|
}
|
||||||
|
Ok(OAuthDeviceAuthorization {
|
||||||
|
device_code,
|
||||||
|
user_code,
|
||||||
|
verification_uri: if verification_uri.is_empty() {
|
||||||
|
verification_uri_complete.clone()
|
||||||
|
} else {
|
||||||
|
verification_uri
|
||||||
|
},
|
||||||
|
verification_uri_complete,
|
||||||
|
expires_in: json_u64(payload, &["expires_in", "expiresIn"])
|
||||||
|
.unwrap_or(DEFAULT_DEVICE_EXPIRES_IN_SECS),
|
||||||
|
interval: json_u64(payload, &["interval"]).unwrap_or(DEFAULT_DEVICE_POLL_INTERVAL_SECS),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn raw_credential_string(raw: Option<&Value>, keys: &[&str]) -> Option<String> {
|
||||||
|
let object = raw?.as_object()?;
|
||||||
|
keys.iter().find_map(|key| {
|
||||||
|
object
|
||||||
|
.get(*key)
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.map(ToOwned::to_owned)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn response_json(response: &crate::network::OAuthHttpResponse) -> Option<Value> {
|
||||||
|
response
|
||||||
|
.json_body
|
||||||
|
.clone()
|
||||||
|
.or_else(|| serde_json::from_str::<Value>(&response.body_text).ok())
|
||||||
|
}
|
||||||
|
|
||||||
|
fn oauth_error_code(payload: &Value) -> Option<String> {
|
||||||
|
payload
|
||||||
|
.get("error")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.map(ToOwned::to_owned)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn json_non_empty_string(payload: &Value, keys: &[&str]) -> Option<String> {
|
||||||
|
keys.iter().find_map(|key| {
|
||||||
|
payload
|
||||||
|
.get(*key)
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.map(ToOwned::to_owned)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn json_u64(payload: &Value, keys: &[&str]) -> Option<u64> {
|
||||||
|
keys.iter().find_map(|key| match payload.get(*key)? {
|
||||||
|
Value::Number(number) => number.as_u64(),
|
||||||
|
Value::String(string) => string.trim().parse::<u64>().ok(),
|
||||||
|
_ => None,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn form_headers() -> BTreeMap<String, String> {
|
||||||
|
BTreeMap::from([
|
||||||
|
(
|
||||||
|
"content-type".to_string(),
|
||||||
|
"application/x-www-form-urlencoded".to_string(),
|
||||||
|
),
|
||||||
|
("accept".to_string(), "application/json".to_string()),
|
||||||
|
])
|
||||||
|
}
|
||||||
|
|
||||||
|
fn decode_jwt_claims(token: &str) -> Option<Map<String, Value>> {
|
||||||
|
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
|
||||||
|
const MAX_UNVERIFIED_JWT_CLAIMS_BYTES: usize = 64 * 1024;
|
||||||
|
|
||||||
|
let payload = token.split('.').nth(1)?;
|
||||||
|
let max_encoded_len = MAX_UNVERIFIED_JWT_CLAIMS_BYTES
|
||||||
|
.saturating_add(2)
|
||||||
|
.checked_div(3)
|
||||||
|
.unwrap_or(usize::MAX)
|
||||||
|
.saturating_mul(4);
|
||||||
|
if payload.len() > max_encoded_len {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
let bytes = URL_SAFE_NO_PAD.decode(payload.as_bytes()).ok()?;
|
||||||
|
if bytes.len() > MAX_UNVERIFIED_JWT_CLAIMS_BYTES {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
serde_json::from_slice::<Value>(&bytes)
|
||||||
|
.ok()?
|
||||||
|
.as_object()
|
||||||
|
.cloned()
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::{
|
||||||
|
XaiDevicePollOutcome, XaiProviderOAuthAdapter, XAI_CLIENT_ID, XAI_DEVICE_CODE_GRANT_TYPE,
|
||||||
|
XAI_OAUTH_SCOPES, XAI_PROVIDER_TYPE,
|
||||||
|
};
|
||||||
|
use crate::network::{OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse};
|
||||||
|
use crate::provider::{
|
||||||
|
ProviderOAuthAccount, ProviderOAuthAdapter, ProviderOAuthImportInput,
|
||||||
|
ProviderOAuthTransportContext,
|
||||||
|
};
|
||||||
|
use async_trait::async_trait;
|
||||||
|
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
|
||||||
|
use serde_json::{json, Value};
|
||||||
|
use std::collections::BTreeMap;
|
||||||
|
use std::sync::{Arc, Mutex};
|
||||||
|
|
||||||
|
#[derive(Clone)]
|
||||||
|
struct ScriptedExecutor {
|
||||||
|
seen_request: Arc<Mutex<Option<OAuthHttpRequest>>>,
|
||||||
|
status_code: u16,
|
||||||
|
payload: Value,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl OAuthHttpExecutor for ScriptedExecutor {
|
||||||
|
async fn execute(
|
||||||
|
&self,
|
||||||
|
request: OAuthHttpRequest,
|
||||||
|
) -> Result<OAuthHttpResponse, crate::core::OAuthError> {
|
||||||
|
*self.seen_request.lock().expect("mutex should lock") = Some(request);
|
||||||
|
Ok(OAuthHttpResponse {
|
||||||
|
status_code: self.status_code,
|
||||||
|
body_text: self.payload.to_string(),
|
||||||
|
json_body: Some(self.payload.clone()),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn transport_context() -> ProviderOAuthTransportContext {
|
||||||
|
ProviderOAuthTransportContext {
|
||||||
|
provider_id: "provider-xai".to_string(),
|
||||||
|
provider_type: XAI_PROVIDER_TYPE.to_string(),
|
||||||
|
endpoint_id: None,
|
||||||
|
key_id: None,
|
||||||
|
auth_type: Some("oauth".to_string()),
|
||||||
|
decrypted_api_key: None,
|
||||||
|
decrypted_auth_config: None,
|
||||||
|
provider_config: None,
|
||||||
|
endpoint_config: None,
|
||||||
|
key_config: None,
|
||||||
|
network: crate::network::OAuthNetworkContext::provider_operation(None),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn encoded_jwt(claims: &Value) -> String {
|
||||||
|
format!(
|
||||||
|
"header.{}.signature",
|
||||||
|
URL_SAFE_NO_PAD.encode(serde_json::to_vec(claims).expect("claims should encode"))
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn imports_api_key_as_official_api_credential() {
|
||||||
|
let adapter = XaiProviderOAuthAdapter::default();
|
||||||
|
let executor = ScriptedExecutor {
|
||||||
|
seen_request: Arc::new(Mutex::new(None)),
|
||||||
|
status_code: 200,
|
||||||
|
payload: json!({}),
|
||||||
|
};
|
||||||
|
let result = adapter
|
||||||
|
.import_credentials(
|
||||||
|
&executor,
|
||||||
|
&transport_context(),
|
||||||
|
ProviderOAuthImportInput {
|
||||||
|
provider_type: XAI_PROVIDER_TYPE.to_string(),
|
||||||
|
name: Some("work".to_string()),
|
||||||
|
refresh_token: None,
|
||||||
|
raw_credentials: Some(json!({"api_key": "xai-key-123"})),
|
||||||
|
network: crate::network::OAuthNetworkContext::provider_operation(None),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.expect("api key import should succeed");
|
||||||
|
|
||||||
|
assert_eq!(result.token_set.access_token, "xai-key-123");
|
||||||
|
assert_eq!(result.auth_config["using_api"], json!(true));
|
||||||
|
assert_eq!(result.auth_config["auth_method"], json!("api_key"));
|
||||||
|
assert!(executor.seen_request.lock().expect("lock").is_none());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn device_poll_treats_authorization_pending_as_pending() {
|
||||||
|
let adapter = XaiProviderOAuthAdapter::default();
|
||||||
|
let seen = Arc::new(Mutex::new(None));
|
||||||
|
let executor = ScriptedExecutor {
|
||||||
|
seen_request: Arc::clone(&seen),
|
||||||
|
status_code: 400,
|
||||||
|
payload: json!({"error": "authorization_pending"}),
|
||||||
|
};
|
||||||
|
let outcome = adapter
|
||||||
|
.poll_device_token(&executor, &transport_context(), "device-code")
|
||||||
|
.await
|
||||||
|
.expect("pending should not be fatal");
|
||||||
|
assert_eq!(outcome, XaiDevicePollOutcome::Pending);
|
||||||
|
|
||||||
|
let request = seen.lock().expect("lock").clone().expect("request");
|
||||||
|
let body = request.body_bytes.expect("body");
|
||||||
|
let fields = url::form_urlencoded::parse(&body)
|
||||||
|
.into_owned()
|
||||||
|
.collect::<BTreeMap<_, _>>();
|
||||||
|
assert_eq!(fields["grant_type"], XAI_DEVICE_CODE_GRANT_TYPE);
|
||||||
|
assert_eq!(fields["device_code"], "device-code");
|
||||||
|
assert_eq!(fields["client_id"], XAI_CLIENT_ID);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn refresh_posts_client_id_and_refresh_token_without_scope() {
|
||||||
|
let adapter = XaiProviderOAuthAdapter::default();
|
||||||
|
let seen = Arc::new(Mutex::new(None));
|
||||||
|
let id_token = encoded_jwt(&json!({"email": "[email protected]", "sub": "subject-1"}));
|
||||||
|
let executor = ScriptedExecutor {
|
||||||
|
seen_request: Arc::clone(&seen),
|
||||||
|
status_code: 200,
|
||||||
|
payload: json!({
|
||||||
|
"access_token": "new-access",
|
||||||
|
"refresh_token": "new-refresh",
|
||||||
|
"id_token": id_token,
|
||||||
|
"expires_in": 3600
|
||||||
|
}),
|
||||||
|
};
|
||||||
|
let account = ProviderOAuthAccount {
|
||||||
|
provider_type: XAI_PROVIDER_TYPE.to_string(),
|
||||||
|
access_token: "old-access".to_string(),
|
||||||
|
auth_config: json!({
|
||||||
|
"provider_type": XAI_PROVIDER_TYPE,
|
||||||
|
"refresh_token": "old-refresh",
|
||||||
|
"using_api": false,
|
||||||
|
}),
|
||||||
|
expires_at_unix_secs: None,
|
||||||
|
identity: BTreeMap::new(),
|
||||||
|
};
|
||||||
|
let result = adapter
|
||||||
|
.refresh(&executor, &transport_context(), &account)
|
||||||
|
.await
|
||||||
|
.expect("refresh should succeed");
|
||||||
|
assert_eq!(result.token_set.access_token, "new-access");
|
||||||
|
assert_eq!(result.auth_config["using_api"], json!(false));
|
||||||
|
assert_eq!(result.auth_config["email"], json!("[email protected]"));
|
||||||
|
assert_eq!(result.auth_config["sub"], json!("subject-1"));
|
||||||
|
|
||||||
|
let request = seen.lock().expect("lock").clone().expect("request");
|
||||||
|
let body = request.body_bytes.expect("body");
|
||||||
|
let fields = url::form_urlencoded::parse(&body)
|
||||||
|
.into_owned()
|
||||||
|
.collect::<BTreeMap<_, _>>();
|
||||||
|
assert_eq!(fields["grant_type"], "refresh_token");
|
||||||
|
assert_eq!(fields["client_id"], XAI_CLIENT_ID);
|
||||||
|
assert_eq!(fields["refresh_token"], "old-refresh");
|
||||||
|
assert!(!fields.contains_key("scope"));
|
||||||
|
assert!(XAI_OAUTH_SCOPES.join(" ").contains("grok-cli:access"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn start_device_flow_posts_client_id_and_scope() {
|
||||||
|
let adapter = XaiProviderOAuthAdapter::default();
|
||||||
|
let seen = Arc::new(Mutex::new(None));
|
||||||
|
let executor = ScriptedExecutor {
|
||||||
|
seen_request: Arc::clone(&seen),
|
||||||
|
status_code: 200,
|
||||||
|
payload: json!({
|
||||||
|
"device_code": "dc-1",
|
||||||
|
"user_code": "ABCD-EFGH",
|
||||||
|
"verification_uri": "https://auth.x.ai/device",
|
||||||
|
"verification_uri_complete": "https://auth.x.ai/device?user_code=ABCD-EFGH",
|
||||||
|
"expires_in": 600,
|
||||||
|
"interval": 5
|
||||||
|
}),
|
||||||
|
};
|
||||||
|
let authorization = adapter
|
||||||
|
.start_device_flow(&executor, &transport_context())
|
||||||
|
.await
|
||||||
|
.expect("device start should succeed");
|
||||||
|
assert_eq!(authorization.user_code, "ABCD-EFGH");
|
||||||
|
assert_eq!(authorization.device_code, "dc-1");
|
||||||
|
|
||||||
|
let request = seen.lock().expect("lock").clone().expect("request");
|
||||||
|
let body = request.body_bytes.expect("body");
|
||||||
|
let fields = url::form_urlencoded::parse(&body)
|
||||||
|
.into_owned()
|
||||||
|
.collect::<BTreeMap<_, _>>();
|
||||||
|
assert_eq!(fields["client_id"], XAI_CLIENT_ID);
|
||||||
|
assert_eq!(fields["scope"], XAI_OAUTH_SCOPES.join(" "));
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -21,7 +21,7 @@ impl ProviderOAuthService {
|
|||||||
use super::providers::{
|
use super::providers::{
|
||||||
AntigravityProviderOAuthAdapter, ClaudeCodeProviderOAuthAdapter,
|
AntigravityProviderOAuthAdapter, ClaudeCodeProviderOAuthAdapter,
|
||||||
CodexProviderOAuthAdapter, GenericProviderOAuthAdapter, KiroProviderOAuthAdapter,
|
CodexProviderOAuthAdapter, GenericProviderOAuthAdapter, KiroProviderOAuthAdapter,
|
||||||
WindsurfProviderOAuthAdapter,
|
WindsurfProviderOAuthAdapter, XaiProviderOAuthAdapter,
|
||||||
};
|
};
|
||||||
|
|
||||||
let mut service = Self::new()
|
let mut service = Self::new()
|
||||||
@@ -29,7 +29,8 @@ impl ProviderOAuthService {
|
|||||||
.with_adapter(Arc::new(ClaudeCodeProviderOAuthAdapter::default()))
|
.with_adapter(Arc::new(ClaudeCodeProviderOAuthAdapter::default()))
|
||||||
.with_adapter(Arc::new(CodexProviderOAuthAdapter::default()))
|
.with_adapter(Arc::new(CodexProviderOAuthAdapter::default()))
|
||||||
.with_adapter(Arc::new(AntigravityProviderOAuthAdapter::default()))
|
.with_adapter(Arc::new(AntigravityProviderOAuthAdapter::default()))
|
||||||
.with_adapter(Arc::new(WindsurfProviderOAuthAdapter));
|
.with_adapter(Arc::new(WindsurfProviderOAuthAdapter))
|
||||||
|
.with_adapter(Arc::new(XaiProviderOAuthAdapter::default()));
|
||||||
for provider_type in ["chatgpt_web", "gemini_cli"] {
|
for provider_type in ["chatgpt_web", "gemini_cli"] {
|
||||||
if let Some(adapter) = GenericProviderOAuthAdapter::for_provider_type(provider_type) {
|
if let Some(adapter) = GenericProviderOAuthAdapter::for_provider_type(provider_type) {
|
||||||
service = service.with_adapter(Arc::new(adapter));
|
service = service.with_adapter(Arc::new(adapter));
|
||||||
@@ -144,6 +145,7 @@ mod tests {
|
|||||||
"antigravity",
|
"antigravity",
|
||||||
"kiro",
|
"kiro",
|
||||||
"windsurf",
|
"windsurf",
|
||||||
|
"xai",
|
||||||
] {
|
] {
|
||||||
assert!(
|
assert!(
|
||||||
service.adapter(provider_type).is_ok(),
|
service.adapter(provider_type).is_ok(),
|
||||||
|
|||||||
@@ -22,18 +22,20 @@ pub use providers::{
|
|||||||
build_windsurf_pool_model_configs_request,
|
build_windsurf_pool_model_configs_request,
|
||||||
build_windsurf_pool_model_configs_request_with_base_url, build_windsurf_pool_quota_request,
|
build_windsurf_pool_model_configs_request_with_base_url, build_windsurf_pool_quota_request,
|
||||||
build_windsurf_pool_quota_request_with_base_url, build_windsurf_pool_rate_limit_request,
|
build_windsurf_pool_quota_request_with_base_url, build_windsurf_pool_rate_limit_request,
|
||||||
build_windsurf_pool_rate_limit_request_with_base_url, enrich_chatgpt_web_quota_metadata,
|
build_windsurf_pool_rate_limit_request_with_base_url, build_xai_pool_billing_request,
|
||||||
grok_mode_id_for_model, grok_pool_tier_from_quota_bucket, grok_quota_window_key_for_model,
|
build_xai_pool_user_request, enrich_chatgpt_web_quota_metadata, grok_mode_id_for_model,
|
||||||
|
grok_pool_tier_from_quota_bucket, grok_quota_window_key_for_model,
|
||||||
grok_supported_quota_windows_for_tier, normalize_chatgpt_web_image_quota_limit,
|
grok_supported_quota_windows_for_tier, normalize_chatgpt_web_image_quota_limit,
|
||||||
AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, CodexProviderPoolAdapter,
|
AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, CodexProviderPoolAdapter,
|
||||||
DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, GrokProviderPoolAdapter,
|
DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, GrokProviderPoolAdapter,
|
||||||
KiroPoolQuotaAuthInput, KiroProviderPoolAdapter, UnsupportedQuotaProviderPoolAdapter,
|
KiroPoolQuotaAuthInput, KiroProviderPoolAdapter, UnsupportedQuotaProviderPoolAdapter,
|
||||||
ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH, ANTIGRAVITY_RETRIEVE_USER_QUOTA_SUMMARY_PATH,
|
XaiProviderPoolAdapter, ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH,
|
||||||
CHATGPT_WEB_CONVERSATION_INIT_PATH, CHATGPT_WEB_DEFAULT_BASE_URL,
|
ANTIGRAVITY_RETRIEVE_USER_QUOTA_SUMMARY_PATH, CHATGPT_WEB_CONVERSATION_INIT_PATH,
|
||||||
CODEX_WHAM_RESET_CREDITS_CONSUME_URL, CODEX_WHAM_RESET_CREDITS_URL, CODEX_WHAM_USAGE_URL,
|
CHATGPT_WEB_DEFAULT_BASE_URL, CODEX_WHAM_RESET_CREDITS_CONSUME_URL,
|
||||||
GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH, GEMINI_CLI_USER_AGENT, KIRO_USAGE_LIMITS_PATH,
|
CODEX_WHAM_RESET_CREDITS_URL, CODEX_WHAM_USAGE_URL, GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH,
|
||||||
KIRO_USAGE_SDK_VERSION, WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_PATH,
|
GEMINI_CLI_USER_AGENT, KIRO_USAGE_LIMITS_PATH, KIRO_USAGE_SDK_VERSION,
|
||||||
WINDSURF_USER_STATUS_PATH,
|
WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_PATH, WINDSURF_USER_STATUS_PATH,
|
||||||
|
XAI_BILLING_PATH, XAI_USER_PATH,
|
||||||
};
|
};
|
||||||
pub use quota::{
|
pub use quota::{
|
||||||
provider_pool_key_account_quota_exhausted, provider_pool_key_model_quota_exhausted,
|
provider_pool_key_account_quota_exhausted, provider_pool_key_model_quota_exhausted,
|
||||||
@@ -81,7 +83,8 @@ mod tests {
|
|||||||
"grok",
|
"grok",
|
||||||
"kiro",
|
"kiro",
|
||||||
"vertex_ai",
|
"vertex_ai",
|
||||||
"windsurf"
|
"windsurf",
|
||||||
|
"xai"
|
||||||
]
|
]
|
||||||
);
|
);
|
||||||
assert!(service
|
assert!(service
|
||||||
@@ -104,7 +107,8 @@ mod tests {
|
|||||||
"gemini_cli",
|
"gemini_cli",
|
||||||
"grok",
|
"grok",
|
||||||
"kiro",
|
"kiro",
|
||||||
"windsurf"
|
"windsurf",
|
||||||
|
"xai"
|
||||||
]
|
]
|
||||||
);
|
);
|
||||||
assert!(service.supports_quota_refresh("codex"));
|
assert!(service.supports_quota_refresh("codex"));
|
||||||
@@ -112,6 +116,7 @@ mod tests {
|
|||||||
assert!(service.supports_quota_refresh("grok"));
|
assert!(service.supports_quota_refresh("grok"));
|
||||||
assert!(service.supports_quota_refresh("gemini_cli"));
|
assert!(service.supports_quota_refresh("gemini_cli"));
|
||||||
assert!(service.supports_quota_refresh("windsurf"));
|
assert!(service.supports_quota_refresh("windsurf"));
|
||||||
|
assert!(service.supports_quota_refresh("xai"));
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
service.quota_refresh_unsupported_message("claude_code"),
|
service.quota_refresh_unsupported_message("claude_code"),
|
||||||
"Claude Code 暂不支持自动刷新额度:上游没有稳定可用的账号额度查询接口"
|
"Claude Code 暂不支持自动刷新额度:上游没有稳定可用的账号额度查询接口"
|
||||||
@@ -642,11 +647,11 @@ mod tests {
|
|||||||
|
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
free_first["providers"],
|
free_first["providers"],
|
||||||
json!(["codex", "grok", "kiro", "windsurf"])
|
json!(["codex", "grok", "kiro", "windsurf", "xai"])
|
||||||
);
|
);
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
recent_refresh["providers"],
|
recent_refresh["providers"],
|
||||||
json!(["codex", "grok", "kiro", "windsurf"])
|
json!(["codex", "grok", "kiro", "windsurf", "xai"])
|
||||||
);
|
);
|
||||||
assert_eq!(free_first["default_enabled"], json!(false));
|
assert_eq!(free_first["default_enabled"], json!(false));
|
||||||
assert_eq!(recent_refresh["default_enabled"], json!(false));
|
assert_eq!(recent_refresh["default_enabled"], json!(false));
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ pub mod grok;
|
|||||||
pub mod kiro;
|
pub mod kiro;
|
||||||
pub mod unsupported;
|
pub mod unsupported;
|
||||||
pub mod windsurf;
|
pub mod windsurf;
|
||||||
|
pub mod xai;
|
||||||
|
|
||||||
pub use antigravity::AntigravityProviderPoolAdapter;
|
pub use antigravity::AntigravityProviderPoolAdapter;
|
||||||
pub use antigravity::{
|
pub use antigravity::{
|
||||||
@@ -51,3 +52,7 @@ pub use windsurf::{
|
|||||||
WINDSURF_DEFAULT_BASE_URL, WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_PATH,
|
WINDSURF_DEFAULT_BASE_URL, WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_PATH,
|
||||||
WINDSURF_USER_STATUS_PATH,
|
WINDSURF_USER_STATUS_PATH,
|
||||||
};
|
};
|
||||||
|
pub use xai::{
|
||||||
|
build_xai_pool_billing_request, build_xai_pool_user_request, XaiProviderPoolAdapter,
|
||||||
|
XAI_BILLING_PATH, XAI_USER_PATH,
|
||||||
|
};
|
||||||
|
|||||||
@@ -0,0 +1,255 @@
|
|||||||
|
use std::collections::BTreeMap;
|
||||||
|
|
||||||
|
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogEndpoint;
|
||||||
|
use aether_provider_transport::xai::{
|
||||||
|
insert_cli_identity_headers, XAI_CHAT_PROXY_BASE_URL, XAI_PROVIDER_TYPE,
|
||||||
|
};
|
||||||
|
use serde_json::{Map, Value};
|
||||||
|
|
||||||
|
use crate::capability::ProviderPoolCapabilities;
|
||||||
|
use crate::provider::{
|
||||||
|
provider_pool_endpoint_format_matches, provider_pool_matching_endpoint, ProviderPoolAdapter,
|
||||||
|
ProviderPoolMemberInput,
|
||||||
|
};
|
||||||
|
use crate::quota::{
|
||||||
|
provider_pool_current_unix_secs, provider_pool_json_bool, provider_pool_json_f64,
|
||||||
|
provider_pool_metadata_bucket, provider_pool_model_quota_exhausted,
|
||||||
|
provider_pool_quota_snapshot_exhausted_decision, provider_pool_reset_deadline_elapsed,
|
||||||
|
provider_pool_timestamp_unix_secs,
|
||||||
|
};
|
||||||
|
use crate::quota_refresh::ProviderPoolQuotaRequestSpec;
|
||||||
|
|
||||||
|
pub const XAI_USER_PATH: &str = "/user";
|
||||||
|
pub const XAI_BILLING_PATH: &str = "/billing?format=credits";
|
||||||
|
|
||||||
|
#[derive(Debug, Clone, Default)]
|
||||||
|
pub struct XaiProviderPoolAdapter;
|
||||||
|
|
||||||
|
impl ProviderPoolAdapter for XaiProviderPoolAdapter {
|
||||||
|
fn provider_type(&self) -> &'static str {
|
||||||
|
XAI_PROVIDER_TYPE
|
||||||
|
}
|
||||||
|
|
||||||
|
fn capabilities(&self) -> ProviderPoolCapabilities {
|
||||||
|
ProviderPoolCapabilities {
|
||||||
|
plan_tier: true,
|
||||||
|
quota_reset: true,
|
||||||
|
quota_refresh: true,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn quota_exhausted(&self, input: &ProviderPoolMemberInput<'_>) -> bool {
|
||||||
|
if let Some(exhausted) = input.provider_model_name.and_then(|model| {
|
||||||
|
provider_pool_model_quota_exhausted(input.key, input.provider_type, model)
|
||||||
|
}) {
|
||||||
|
return exhausted;
|
||||||
|
}
|
||||||
|
if let Some(exhausted) =
|
||||||
|
provider_pool_quota_snapshot_exhausted_decision(input.key, input.provider_type)
|
||||||
|
{
|
||||||
|
return exhausted;
|
||||||
|
}
|
||||||
|
provider_pool_metadata_bucket(input.key.upstream_metadata.as_ref(), input.provider_type)
|
||||||
|
.is_some_and(quota_exhausted_from_bucket)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn quota_refresh_endpoint(
|
||||||
|
&self,
|
||||||
|
endpoints: &[StoredProviderCatalogEndpoint],
|
||||||
|
include_inactive: bool,
|
||||||
|
) -> Option<StoredProviderCatalogEndpoint> {
|
||||||
|
provider_pool_matching_endpoint(endpoints, include_inactive, |endpoint| {
|
||||||
|
provider_pool_endpoint_format_matches(endpoint, "openai:responses")
|
||||||
|
})
|
||||||
|
.or_else(|| provider_pool_matching_endpoint(endpoints, include_inactive, |_| true))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn quota_refresh_missing_endpoint_message(&self) -> String {
|
||||||
|
"找不到有效的 openai:responses 端点".to_string()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn build_xai_pool_user_request(
|
||||||
|
key_id: &str,
|
||||||
|
authorization: (String, String),
|
||||||
|
) -> ProviderPoolQuotaRequestSpec {
|
||||||
|
build_xai_pool_request(
|
||||||
|
format!("xai-user:{key_id}"),
|
||||||
|
"xai:user",
|
||||||
|
"user",
|
||||||
|
XAI_USER_PATH,
|
||||||
|
authorization,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn build_xai_pool_billing_request(
|
||||||
|
key_id: &str,
|
||||||
|
authorization: (String, String),
|
||||||
|
user_id: Option<&str>,
|
||||||
|
) -> ProviderPoolQuotaRequestSpec {
|
||||||
|
build_xai_pool_request(
|
||||||
|
format!("xai-billing:{key_id}"),
|
||||||
|
"xai:billing",
|
||||||
|
"billing",
|
||||||
|
XAI_BILLING_PATH,
|
||||||
|
authorization,
|
||||||
|
user_id,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn build_xai_pool_request(
|
||||||
|
request_id: String,
|
||||||
|
provider_api_format: &str,
|
||||||
|
model_name: &str,
|
||||||
|
path: &str,
|
||||||
|
authorization: (String, String),
|
||||||
|
user_id: Option<&str>,
|
||||||
|
) -> ProviderPoolQuotaRequestSpec {
|
||||||
|
let mut headers = BTreeMap::from([
|
||||||
|
(authorization.0, authorization.1),
|
||||||
|
("accept".to_string(), "application/json".to_string()),
|
||||||
|
]);
|
||||||
|
insert_cli_identity_headers(&mut headers);
|
||||||
|
if let Some(user_id) = user_id.map(str::trim).filter(|value| !value.is_empty()) {
|
||||||
|
headers.insert("x-userid".to_string(), user_id.to_string());
|
||||||
|
}
|
||||||
|
|
||||||
|
ProviderPoolQuotaRequestSpec {
|
||||||
|
request_id,
|
||||||
|
provider_name: XAI_PROVIDER_TYPE.to_string(),
|
||||||
|
quota_kind: XAI_PROVIDER_TYPE.to_string(),
|
||||||
|
method: "GET".to_string(),
|
||||||
|
url: format!("{}{path}", XAI_CHAT_PROXY_BASE_URL.trim_end_matches('/')),
|
||||||
|
headers,
|
||||||
|
content_type: None,
|
||||||
|
json_body: None,
|
||||||
|
client_api_format: "openai:responses".to_string(),
|
||||||
|
provider_api_format: provider_api_format.to_string(),
|
||||||
|
model_name: Some(model_name.to_string()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(crate) fn quota_exhausted_from_bucket(bucket: &Map<String, Value>) -> bool {
|
||||||
|
if provider_pool_current_unix_secs().is_some_and(|now| {
|
||||||
|
provider_pool_reset_deadline_elapsed(
|
||||||
|
bucket,
|
||||||
|
provider_pool_timestamp_unix_secs(bucket.get("updated_at")),
|
||||||
|
now,
|
||||||
|
)
|
||||||
|
}) {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
|
||||||
|
let usage_exhausted = provider_pool_json_f64(bucket.get("remaining"))
|
||||||
|
.is_some_and(|value| value <= 0.0)
|
||||||
|
|| provider_pool_json_f64(bucket.get("usage_percentage"))
|
||||||
|
.is_some_and(|value| value >= 100.0 - 1e-6)
|
||||||
|
|| match (
|
||||||
|
provider_pool_json_f64(bucket.get("usage_limit")),
|
||||||
|
provider_pool_json_f64(bucket.get("current_usage")),
|
||||||
|
) {
|
||||||
|
(Some(limit), Some(current)) if limit > 0.0 => current >= limit,
|
||||||
|
_ => false,
|
||||||
|
};
|
||||||
|
if !usage_exhausted {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
|
||||||
|
let prepaid_available =
|
||||||
|
provider_pool_json_f64(bucket.get("prepaid_balance")).is_some_and(|value| value > 0.0);
|
||||||
|
if prepaid_available {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
|
||||||
|
let on_demand_enabled = provider_pool_json_bool(bucket.get("on_demand_enabled")) != Some(false);
|
||||||
|
let on_demand_cap = provider_pool_json_f64(bucket.get("on_demand_cap")).unwrap_or(0.0);
|
||||||
|
let on_demand_used = provider_pool_json_f64(bucket.get("on_demand_used")).unwrap_or(0.0);
|
||||||
|
if on_demand_enabled && on_demand_cap > 0.0 && on_demand_used < on_demand_cap {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
|
||||||
|
true
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::{
|
||||||
|
build_xai_pool_billing_request, build_xai_pool_user_request, quota_exhausted_from_bucket,
|
||||||
|
};
|
||||||
|
use aether_provider_transport::xai::{
|
||||||
|
XAI_CHAT_PROXY_BASE_URL, XAI_CLIENT_IDENTIFIER_VALUE, XAI_TOKEN_AUTH_VALUE,
|
||||||
|
};
|
||||||
|
use serde_json::{json, Map};
|
||||||
|
|
||||||
|
fn bucket(value: serde_json::Value) -> Map<String, serde_json::Value> {
|
||||||
|
value.as_object().cloned().expect("bucket should be object")
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn user_and_billing_requests_pin_cli_chat_proxy_and_identity_headers() {
|
||||||
|
let authorization = ("authorization".to_string(), "Bearer xai-access".to_string());
|
||||||
|
let user = build_xai_pool_user_request("key-1", authorization.clone());
|
||||||
|
let billing = build_xai_pool_billing_request("key-1", authorization, Some("user-42"));
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
user.url,
|
||||||
|
format!("{}/user", XAI_CHAT_PROXY_BASE_URL.trim_end_matches('/'))
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
billing.url,
|
||||||
|
format!(
|
||||||
|
"{}/billing?format=credits",
|
||||||
|
XAI_CHAT_PROXY_BASE_URL.trim_end_matches('/')
|
||||||
|
)
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
user.headers.get("x-xai-token-auth").map(String::as_str),
|
||||||
|
Some(XAI_TOKEN_AUTH_VALUE)
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
user.headers
|
||||||
|
.get("x-grok-client-identifier")
|
||||||
|
.map(String::as_str),
|
||||||
|
Some(XAI_CLIENT_IDENTIFIER_VALUE)
|
||||||
|
);
|
||||||
|
assert!(!user.headers.contains_key("x-userid"));
|
||||||
|
assert_eq!(
|
||||||
|
billing.headers.get("x-userid").map(String::as_str),
|
||||||
|
Some("user-42")
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
billing.headers.get("authorization").map(String::as_str),
|
||||||
|
Some("Bearer xai-access")
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn percent_exhausted_without_prepaid_or_on_demand_is_exhausted() {
|
||||||
|
assert!(quota_exhausted_from_bucket(&bucket(json!({
|
||||||
|
"usage_percentage": 100.0,
|
||||||
|
"prepaid_balance": 0.0,
|
||||||
|
"on_demand_cap": 0.0,
|
||||||
|
"on_demand_used": 0.0
|
||||||
|
}))));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn unified_billing_zero_on_demand_cap_is_not_exhausted_when_percent_remains() {
|
||||||
|
assert!(!quota_exhausted_from_bucket(&bucket(json!({
|
||||||
|
"usage_percentage": 46.0,
|
||||||
|
"prepaid_balance": 0.0,
|
||||||
|
"on_demand_cap": 0.0,
|
||||||
|
"on_demand_used": 0.0
|
||||||
|
}))));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn prepaid_balance_keeps_account_available_after_weekly_pool_hits_100() {
|
||||||
|
assert!(!quota_exhausted_from_bucket(&bucket(json!({
|
||||||
|
"usage_percentage": 100.0,
|
||||||
|
"prepaid_balance": 12.5,
|
||||||
|
"on_demand_cap": 0.0
|
||||||
|
}))));
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -13,8 +13,8 @@ use crate::provider::{ProviderPoolAdapter, ProviderPoolMemberInput};
|
|||||||
use crate::providers::{
|
use crate::providers::{
|
||||||
AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, CodexProviderPoolAdapter,
|
AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, CodexProviderPoolAdapter,
|
||||||
DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, GrokProviderPoolAdapter,
|
DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, GrokProviderPoolAdapter,
|
||||||
KiroProviderPoolAdapter, WindsurfProviderPoolAdapter, CLAUDE_CODE_PROVIDER_POOL_ADAPTER,
|
KiroProviderPoolAdapter, WindsurfProviderPoolAdapter, XaiProviderPoolAdapter,
|
||||||
VERTEX_AI_PROVIDER_POOL_ADAPTER,
|
CLAUDE_CODE_PROVIDER_POOL_ADAPTER, VERTEX_AI_PROVIDER_POOL_ADAPTER,
|
||||||
};
|
};
|
||||||
|
|
||||||
#[derive(Clone)]
|
#[derive(Clone)]
|
||||||
@@ -55,6 +55,7 @@ impl ProviderPoolService {
|
|||||||
.with_adapter(Arc::new(KiroProviderPoolAdapter))
|
.with_adapter(Arc::new(KiroProviderPoolAdapter))
|
||||||
.with_adapter(Arc::new(ChatGptWebProviderPoolAdapter))
|
.with_adapter(Arc::new(ChatGptWebProviderPoolAdapter))
|
||||||
.with_adapter(Arc::new(WindsurfProviderPoolAdapter))
|
.with_adapter(Arc::new(WindsurfProviderPoolAdapter))
|
||||||
|
.with_adapter(Arc::new(XaiProviderPoolAdapter))
|
||||||
.with_adapter(Arc::new(VERTEX_AI_PROVIDER_POOL_ADAPTER))
|
.with_adapter(Arc::new(VERTEX_AI_PROVIDER_POOL_ADAPTER))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -754,6 +754,90 @@ mod tests {
|
|||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_responses_transport_converts_standard_client_protocols() {
|
||||||
|
let transport = transport_snapshot("xai", "openai:responses", "oauth", true, None);
|
||||||
|
|
||||||
|
for client_api_format in ["openai:chat", "claude:messages", "gemini:generate_content"] {
|
||||||
|
assert!(
|
||||||
|
request_pair_allowed_for_transport(
|
||||||
|
&transport,
|
||||||
|
client_api_format,
|
||||||
|
"openai:responses"
|
||||||
|
),
|
||||||
|
"{client_api_format} should convert onto xAI Responses"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
candidate_transport_pair_skip_reason(&transport, client_api_format),
|
||||||
|
None
|
||||||
|
);
|
||||||
|
}
|
||||||
|
assert!(request_conversion_transport_supported(
|
||||||
|
&transport,
|
||||||
|
RequestConversionKind::ToOpenAiResponses
|
||||||
|
));
|
||||||
|
assert!(!request_pair_allowed_for_transport(
|
||||||
|
&transport,
|
||||||
|
"openai:image",
|
||||||
|
"openai:responses"
|
||||||
|
));
|
||||||
|
assert!(!request_pair_allowed_for_transport(
|
||||||
|
&transport,
|
||||||
|
"openai:video",
|
||||||
|
"openai:responses"
|
||||||
|
));
|
||||||
|
for isolated in ["openai:responses:compact", "openai:image", "openai:video"] {
|
||||||
|
assert!(
|
||||||
|
!request_pair_allowed_for_transport(&transport, isolated, "openai:responses"),
|
||||||
|
"{isolated} must not convert onto xAI Responses"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_compact_and_media_endpoints_are_same_format_only() {
|
||||||
|
let compact = transport_snapshot("xai", "openai:responses:compact", "oauth", true, None);
|
||||||
|
assert!(request_pair_allowed_for_transport(
|
||||||
|
&compact,
|
||||||
|
"openai:responses:compact",
|
||||||
|
"openai:responses:compact"
|
||||||
|
));
|
||||||
|
for client_api_format in [
|
||||||
|
"openai:chat",
|
||||||
|
"openai:responses",
|
||||||
|
"claude:messages",
|
||||||
|
"gemini:generate_content",
|
||||||
|
] {
|
||||||
|
assert!(
|
||||||
|
!request_pair_allowed_for_transport(
|
||||||
|
&compact,
|
||||||
|
client_api_format,
|
||||||
|
"openai:responses:compact"
|
||||||
|
),
|
||||||
|
"{client_api_format} must not convert onto xAI compact"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
for api_format in ["openai:image", "openai:video"] {
|
||||||
|
let transport = transport_snapshot("xai", api_format, "oauth", true, None);
|
||||||
|
assert!(
|
||||||
|
request_pair_allowed_for_transport(&transport, api_format, api_format),
|
||||||
|
"{api_format} same-format transport should be allowed"
|
||||||
|
);
|
||||||
|
for client_api_format in [
|
||||||
|
"openai:chat",
|
||||||
|
"openai:responses",
|
||||||
|
"claude:messages",
|
||||||
|
"gemini:generate_content",
|
||||||
|
] {
|
||||||
|
assert!(
|
||||||
|
!request_pair_allowed_for_transport(&transport, client_api_format, api_format),
|
||||||
|
"{client_api_format} must not convert onto {api_format}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn windsurf_openai_chat_anchor_supports_cross_format_conversion_via_cascade() {
|
fn windsurf_openai_chat_anchor_supports_cross_format_conversion_via_cascade() {
|
||||||
let mut transport = transport_snapshot("windsurf", "openai:chat", "oauth", true, None);
|
let mut transport = transport_snapshot("windsurf", "openai:chat", "oauth", true, None);
|
||||||
|
|||||||
@@ -30,6 +30,7 @@ pub mod url;
|
|||||||
pub mod vertex;
|
pub mod vertex;
|
||||||
mod video;
|
mod video;
|
||||||
pub mod windsurf;
|
pub mod windsurf;
|
||||||
|
pub mod xai;
|
||||||
|
|
||||||
pub use aether_oauth as oauth;
|
pub use aether_oauth as oauth;
|
||||||
pub use agent_identity::{
|
pub use agent_identity::{
|
||||||
@@ -195,3 +196,10 @@ pub use windsurf::{
|
|||||||
local_windsurf_request_transport_unsupported_reason_with_network, GET_CHAT_MESSAGE_PATH,
|
local_windsurf_request_transport_unsupported_reason_with_network, GET_CHAT_MESSAGE_PATH,
|
||||||
WINDSURF_ENVELOPE_NAME,
|
WINDSURF_ENVELOPE_NAME,
|
||||||
};
|
};
|
||||||
|
pub use xai::{
|
||||||
|
extract_xai_user_id_from_auth_config, extract_xai_user_id_from_value,
|
||||||
|
insert_cli_identity_headers, insert_cli_identity_headers_if_needed, is_xai_provider_transport,
|
||||||
|
resolved_xai_request_base_url, resolved_xai_upstream_base_url,
|
||||||
|
should_attach_cli_identity_headers, xai_auth_uses_api, xai_uses_official_api, XAI_API_BASE_URL,
|
||||||
|
XAI_CHAT_PROXY_BASE_URL, XAI_PROVIDER_TYPE,
|
||||||
|
};
|
||||||
|
|||||||
@@ -84,6 +84,11 @@ fn is_dedicated_openai_image_provider(transport: &GatewayProviderTransportSnapsh
|
|||||||
.trim()
|
.trim()
|
||||||
.eq_ignore_ascii_case("codex")
|
.eq_ignore_ascii_case("codex")
|
||||||
|| is_grok_provider_transport(transport)
|
|| is_grok_provider_transport(transport)
|
||||||
|
|| transport
|
||||||
|
.provider
|
||||||
|
.provider_type
|
||||||
|
.trim()
|
||||||
|
.eq_ignore_ascii_case("xai")
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn resolve_openai_image_auth(
|
pub fn resolve_openai_image_auth(
|
||||||
@@ -92,7 +97,10 @@ pub fn resolve_openai_image_auth(
|
|||||||
if is_grok_provider_transport(transport) {
|
if is_grok_provider_transport(transport) {
|
||||||
return resolve_grok_session_auth(transport);
|
return resolve_grok_session_auth(transport);
|
||||||
}
|
}
|
||||||
resolve_local_openai_bearer_auth(transport)
|
resolve_local_openai_bearer_auth(transport).or_else(|| {
|
||||||
|
crate::generic_oauth::resolve_local_generic_oauth_transport_authorization(transport)
|
||||||
|
.map(|value| ("authorization".to_string(), value))
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn build_openai_image_upstream_url(
|
pub fn build_openai_image_upstream_url(
|
||||||
@@ -100,7 +108,11 @@ pub fn build_openai_image_upstream_url(
|
|||||||
request_path: Option<&str>,
|
request_path: Option<&str>,
|
||||||
request_query: Option<&str>,
|
request_query: Option<&str>,
|
||||||
) -> String {
|
) -> String {
|
||||||
build_openai_image_url(&transport.endpoint.base_url, request_path, request_query)
|
build_openai_image_url(
|
||||||
|
&crate::xai::resolved_xai_request_base_url(transport, "openai:image"),
|
||||||
|
request_path,
|
||||||
|
request_query,
|
||||||
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn build_openai_image_headers(
|
pub fn build_openai_image_headers(
|
||||||
@@ -113,6 +125,11 @@ pub fn build_openai_image_headers(
|
|||||||
&BTreeMap::new(),
|
&BTreeMap::new(),
|
||||||
);
|
);
|
||||||
provider_request_headers.insert("content-type".to_string(), "application/json".to_string());
|
provider_request_headers.insert("content-type".to_string(), "application/json".to_string());
|
||||||
|
crate::xai::insert_cli_identity_headers_if_needed(
|
||||||
|
input.transport,
|
||||||
|
"openai:image",
|
||||||
|
&mut provider_request_headers,
|
||||||
|
);
|
||||||
if let Some(accept) = input.accept {
|
if let Some(accept) = input.accept {
|
||||||
provider_request_headers.insert("accept".to_string(), accept.to_string());
|
provider_request_headers.insert("accept".to_string(), accept.to_string());
|
||||||
} else {
|
} else {
|
||||||
@@ -280,6 +297,48 @@ mod tests {
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_oauth_image_uses_cli_proxy() {
|
||||||
|
let mut transport = sample_transport();
|
||||||
|
transport.provider.provider_type = "xai".to_string();
|
||||||
|
transport.endpoint.base_url = "https://cli-chat-proxy.grok.com/v1".to_string();
|
||||||
|
transport.key.auth_type = "oauth".to_string();
|
||||||
|
transport.key.decrypted_auth_config =
|
||||||
|
Some(r#"{"refresh_token":"rt","using_api":false}"#.to_string());
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
openai_image_transport_unsupported_reason(&transport, "openai:image"),
|
||||||
|
None
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
build_openai_image_upstream_url(&transport, Some("/v1/images/generations"), None),
|
||||||
|
"https://cli-chat-proxy.grok.com/v1/images/generations"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
build_openai_image_upstream_url(&transport, Some("/v1/images/edits"), None),
|
||||||
|
"https://cli-chat-proxy.grok.com/v1/images/edits"
|
||||||
|
);
|
||||||
|
let headers = build_openai_image_headers(ProviderOpenAiImageHeadersInput {
|
||||||
|
transport: &transport,
|
||||||
|
headers: &HeaderMap::new(),
|
||||||
|
auth_header: "authorization",
|
||||||
|
auth_value: "Bearer test-token",
|
||||||
|
accept: None,
|
||||||
|
header_rules: None,
|
||||||
|
provider_request_body: &json!({"prompt": "A cat"}),
|
||||||
|
original_request_body: &json!({"prompt": "A cat"}),
|
||||||
|
})
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
headers.get("x-xai-token-auth").map(String::as_str),
|
||||||
|
Some("xai-grok-cli")
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
headers.get("authorization").map(String::as_str),
|
||||||
|
Some("Bearer test-token")
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn codex_is_supported_by_dedicated_openai_image_transport_policy() {
|
fn codex_is_supported_by_dedicated_openai_image_transport_policy() {
|
||||||
let mut transport = sample_transport();
|
let mut transport = sample_transport();
|
||||||
|
|||||||
@@ -275,6 +275,17 @@ const WINDSURF_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
|||||||
..STANDARD_RUNTIME_POLICY
|
..STANDARD_RUNTIME_POLICY
|
||||||
};
|
};
|
||||||
|
|
||||||
|
const XAI_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||||
|
fixed_provider: true,
|
||||||
|
api_format_inheritance: ProviderApiFormatInheritance::OAuthOrBearer,
|
||||||
|
enable_format_conversion_by_default: true,
|
||||||
|
oauth_is_bearer_like: true,
|
||||||
|
supports_model_fetch: false,
|
||||||
|
supports_local_openai_chat_transport: false,
|
||||||
|
supports_local_same_format_transport: true,
|
||||||
|
..STANDARD_RUNTIME_POLICY
|
||||||
|
};
|
||||||
|
|
||||||
const CLAUDE_CODE_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
|
const CLAUDE_CODE_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
|
||||||
provider_type: "claude_code",
|
provider_type: "claude_code",
|
||||||
version: 2,
|
version: 2,
|
||||||
@@ -446,6 +457,39 @@ const WINDSURF_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTem
|
|||||||
runtime_policy: WINDSURF_RUNTIME_POLICY,
|
runtime_policy: WINDSURF_RUNTIME_POLICY,
|
||||||
};
|
};
|
||||||
|
|
||||||
|
const XAI_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
|
||||||
|
provider_type: "xai",
|
||||||
|
version: 2,
|
||||||
|
base_url: crate::xai::XAI_CHAT_PROXY_BASE_URL,
|
||||||
|
endpoints: &[
|
||||||
|
FixedProviderEndpointTemplate {
|
||||||
|
item_key: "openai:responses",
|
||||||
|
api_format: "openai:responses",
|
||||||
|
custom_path: None,
|
||||||
|
config_defaults: FORCE_STREAM_ENDPOINT_CONFIG_DEFAULTS,
|
||||||
|
},
|
||||||
|
FixedProviderEndpointTemplate {
|
||||||
|
item_key: "openai:responses:compact",
|
||||||
|
api_format: "openai:responses:compact",
|
||||||
|
custom_path: None,
|
||||||
|
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||||
|
},
|
||||||
|
FixedProviderEndpointTemplate {
|
||||||
|
item_key: "openai:image",
|
||||||
|
api_format: "openai:image",
|
||||||
|
custom_path: None,
|
||||||
|
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||||
|
},
|
||||||
|
FixedProviderEndpointTemplate {
|
||||||
|
item_key: "openai:video",
|
||||||
|
api_format: "openai:video",
|
||||||
|
custom_path: None,
|
||||||
|
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||||
|
},
|
||||||
|
],
|
||||||
|
runtime_policy: XAI_RUNTIME_POLICY,
|
||||||
|
};
|
||||||
|
|
||||||
pub fn provider_type_is_fixed(provider_type: &str) -> bool {
|
pub fn provider_type_is_fixed(provider_type: &str) -> bool {
|
||||||
provider_runtime_policy(provider_type).fixed_provider
|
provider_runtime_policy(provider_type).fixed_provider
|
||||||
}
|
}
|
||||||
@@ -498,6 +542,7 @@ pub fn fixed_provider_template(provider_type: &str) -> Option<&'static FixedProv
|
|||||||
"vertex_ai" => Some(&VERTEX_AI_FIXED_PROVIDER_TEMPLATE),
|
"vertex_ai" => Some(&VERTEX_AI_FIXED_PROVIDER_TEMPLATE),
|
||||||
"antigravity" => Some(&ANTIGRAVITY_FIXED_PROVIDER_TEMPLATE),
|
"antigravity" => Some(&ANTIGRAVITY_FIXED_PROVIDER_TEMPLATE),
|
||||||
"windsurf" => Some(&WINDSURF_FIXED_PROVIDER_TEMPLATE),
|
"windsurf" => Some(&WINDSURF_FIXED_PROVIDER_TEMPLATE),
|
||||||
|
"xai" => Some(&XAI_FIXED_PROVIDER_TEMPLATE),
|
||||||
_ => None,
|
_ => None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -613,6 +658,16 @@ pub fn provider_type_admin_oauth_template(provider_type: &str) -> Option<Provide
|
|||||||
redirect_uri: "show-auth-token",
|
redirect_uri: "show-auth-token",
|
||||||
use_pkce: false,
|
use_pkce: false,
|
||||||
}),
|
}),
|
||||||
|
"xai" => Some(ProviderOAuthTemplate {
|
||||||
|
provider_type: "xai",
|
||||||
|
display_name: "xAI",
|
||||||
|
authorize_url: aether_oauth::provider::providers::XAI_DEVICE_CODE_URL,
|
||||||
|
token_url: aether_oauth::provider::providers::XAI_TOKEN_URL,
|
||||||
|
client_id: aether_oauth::provider::providers::XAI_CLIENT_ID,
|
||||||
|
scopes: aether_oauth::provider::providers::XAI_OAUTH_SCOPES,
|
||||||
|
redirect_uri: "",
|
||||||
|
use_pkce: false,
|
||||||
|
}),
|
||||||
_ => None,
|
_ => None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -825,6 +880,50 @@ mod tests {
|
|||||||
assert!(ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES.contains(&"windsurf"));
|
assert!(ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES.contains(&"windsurf"));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_fixed_provider_template_exposes_responses_media_endpoints() {
|
||||||
|
let template = fixed_provider_template("xai").expect("xai template should exist");
|
||||||
|
assert_eq!(template.provider_type, "xai");
|
||||||
|
assert_eq!(template.base_url, crate::xai::XAI_CHAT_PROXY_BASE_URL);
|
||||||
|
assert_eq!(template.version, 2);
|
||||||
|
assert_eq!(
|
||||||
|
template
|
||||||
|
.endpoints
|
||||||
|
.iter()
|
||||||
|
.map(|item| item.api_format)
|
||||||
|
.collect::<Vec<_>>(),
|
||||||
|
vec![
|
||||||
|
"openai:responses",
|
||||||
|
"openai:responses:compact",
|
||||||
|
"openai:image",
|
||||||
|
"openai:video"
|
||||||
|
]
|
||||||
|
);
|
||||||
|
|
||||||
|
let policy = provider_runtime_policy("xai");
|
||||||
|
assert!(policy.fixed_provider);
|
||||||
|
assert!(policy.enable_format_conversion_by_default);
|
||||||
|
assert!(policy.oauth_is_bearer_like);
|
||||||
|
assert!(!policy.supports_model_fetch);
|
||||||
|
assert!(policy.supports_local_same_format_transport);
|
||||||
|
assert!(!policy.supports_local_openai_chat_transport);
|
||||||
|
assert!(fixed_provider_key_inherits_api_formats(
|
||||||
|
"xai", "oauth", None
|
||||||
|
));
|
||||||
|
assert!(fixed_provider_key_inherits_api_formats(
|
||||||
|
"xai", "bearer", None
|
||||||
|
));
|
||||||
|
|
||||||
|
let template = provider_type_admin_oauth_template("xai").expect("xai oauth template");
|
||||||
|
assert_eq!(template.provider_type, "xai");
|
||||||
|
assert_eq!(template.display_name, "xAI");
|
||||||
|
assert_eq!(
|
||||||
|
template.token_url,
|
||||||
|
aether_oauth::provider::providers::XAI_TOKEN_URL
|
||||||
|
);
|
||||||
|
assert!(!ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES.contains(&"xai"));
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn fixed_provider_key_inheritance_keeps_oauth_and_kiro_configured_bearer_keys_open() {
|
fn fixed_provider_key_inheritance_keeps_oauth_and_kiro_configured_bearer_keys_open() {
|
||||||
assert!(fixed_provider_key_inherits_api_formats(
|
assert!(fixed_provider_key_inherits_api_formats(
|
||||||
|
|||||||
@@ -42,6 +42,11 @@ pub fn apply_transport_request_body_semantics(
|
|||||||
{
|
{
|
||||||
sanitize_claude_code_request_body(provider_request_body);
|
sanitize_claude_code_request_body(provider_request_body);
|
||||||
}
|
}
|
||||||
|
aether_ai_formats::apply_xai_upstream_payload_edits(
|
||||||
|
provider_request_body,
|
||||||
|
transport.provider.provider_type.as_str(),
|
||||||
|
provider_api_format.as_str(),
|
||||||
|
);
|
||||||
if provider_api_format == "gemini:embedding" && is_vertex_transport_context(transport) {
|
if provider_api_format == "gemini:embedding" && is_vertex_transport_context(transport) {
|
||||||
apply_vertex_gemini_embedding_body_semantics(provider_request_body)?;
|
apply_vertex_gemini_embedding_body_semantics(provider_request_body)?;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -120,6 +120,12 @@ fn build_transport_request_url_inner(
|
|||||||
return Some(url);
|
return Some(url);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
let xai_base =
|
||||||
|
crate::xai::resolved_xai_upstream_base_url(transport, &normalized_provider_api_format);
|
||||||
|
let request_base_url = xai_base
|
||||||
|
.as_deref()
|
||||||
|
.unwrap_or(transport.endpoint.base_url.as_str());
|
||||||
|
|
||||||
let custom_path_template = transport
|
let custom_path_template = transport
|
||||||
.endpoint
|
.endpoint
|
||||||
.custom_path
|
.custom_path
|
||||||
@@ -164,7 +170,7 @@ fn build_transport_request_url_inner(
|
|||||||
path.to_string()
|
path.to_string()
|
||||||
};
|
};
|
||||||
let mut url = build_passthrough_path_url(
|
let mut url = build_passthrough_path_url(
|
||||||
&transport.endpoint.base_url,
|
request_base_url,
|
||||||
normalized_path.as_str(),
|
normalized_path.as_str(),
|
||||||
params.request_query,
|
params.request_query,
|
||||||
blocked_keys,
|
blocked_keys,
|
||||||
@@ -190,75 +196,68 @@ fn build_transport_request_url_inner(
|
|||||||
|
|
||||||
let url = match normalized_provider_api_format.as_str() {
|
let url = match normalized_provider_api_format.as_str() {
|
||||||
"openai:chat" => Some(build_openai_chat_url(
|
"openai:chat" => Some(build_openai_chat_url(
|
||||||
&transport.endpoint.base_url,
|
request_base_url,
|
||||||
params.request_query,
|
params.request_query,
|
||||||
)),
|
)),
|
||||||
"openai:responses" => Some(build_openai_responses_url(
|
"openai:responses" => Some(build_openai_responses_url(
|
||||||
&transport.endpoint.base_url,
|
request_base_url,
|
||||||
params.request_query,
|
params.request_query,
|
||||||
false,
|
false,
|
||||||
)),
|
)),
|
||||||
"openai:responses:compact" => Some(build_openai_responses_url(
|
"openai:responses:compact" => Some(build_openai_responses_url(
|
||||||
&transport.endpoint.base_url,
|
request_base_url,
|
||||||
params.request_query,
|
params.request_query,
|
||||||
true,
|
true,
|
||||||
)),
|
)),
|
||||||
"openai:search" => Some(build_openai_search_url(
|
"openai:search" => Some(build_openai_search_url(
|
||||||
&transport.endpoint.base_url,
|
request_base_url,
|
||||||
params.request_query,
|
params.request_query,
|
||||||
)),
|
)),
|
||||||
"openai:realtime" => build_passthrough_path_url(
|
"openai:realtime" => build_passthrough_path_url(
|
||||||
&transport.endpoint.base_url,
|
request_base_url,
|
||||||
"/v1/realtime",
|
"/v1/realtime",
|
||||||
params.request_query,
|
params.request_query,
|
||||||
GATEWAY_CREDENTIAL_QUERY_KEYS,
|
GATEWAY_CREDENTIAL_QUERY_KEYS,
|
||||||
)
|
)
|
||||||
.and_then(|url| replace_realtime_model_query(url, params.mapped_model?)),
|
.and_then(|url| replace_realtime_model_query(url, params.mapped_model?)),
|
||||||
"codex:live" => build_passthrough_path_url(
|
"codex:live" => build_passthrough_path_url(
|
||||||
&transport.endpoint.base_url,
|
request_base_url,
|
||||||
"/live",
|
"/live",
|
||||||
params.request_query,
|
params.request_query,
|
||||||
GATEWAY_CREDENTIAL_QUERY_KEYS,
|
GATEWAY_CREDENTIAL_QUERY_KEYS,
|
||||||
),
|
),
|
||||||
"openai:embedding" | "jina:embedding" => {
|
"openai:embedding" | "jina:embedding" => {
|
||||||
build_provider_embedding_v1_url(&transport.endpoint.base_url, params.request_query)
|
build_provider_embedding_v1_url(request_base_url, params.request_query)
|
||||||
|
}
|
||||||
|
"aliyun:multimodal_embedding" => {
|
||||||
|
build_aliyun_multimodal_embedding_url(request_base_url, params.request_query)
|
||||||
}
|
}
|
||||||
"aliyun:multimodal_embedding" => build_aliyun_multimodal_embedding_url(
|
|
||||||
&transport.endpoint.base_url,
|
|
||||||
params.request_query,
|
|
||||||
),
|
|
||||||
"openai:rerank" | "jina:rerank" => {
|
"openai:rerank" | "jina:rerank" => {
|
||||||
build_provider_rerank_v1_url(&transport.endpoint.base_url, params.request_query)
|
build_provider_rerank_v1_url(request_base_url, params.request_query)
|
||||||
}
|
}
|
||||||
"claude:messages" => Some(if is_claude_count_tokens {
|
"claude:messages" => Some(if is_claude_count_tokens {
|
||||||
build_default_claude_count_tokens_url(
|
build_default_claude_count_tokens_url(request_base_url, params.request_query)
|
||||||
&transport.endpoint.base_url,
|
|
||||||
params.request_query,
|
|
||||||
)
|
|
||||||
} else {
|
} else {
|
||||||
build_claude_messages_url(&transport.endpoint.base_url, params.request_query)
|
build_claude_messages_url(request_base_url, params.request_query)
|
||||||
}),
|
}),
|
||||||
"gemini:generate_content" => build_gemini_content_url(
|
"gemini:generate_content" => build_gemini_content_url(
|
||||||
&transport.endpoint.base_url,
|
request_base_url,
|
||||||
params.mapped_model?,
|
params.mapped_model?,
|
||||||
params.upstream_is_stream,
|
params.upstream_is_stream,
|
||||||
params.request_query,
|
params.request_query,
|
||||||
),
|
),
|
||||||
"gemini:embedding" => build_gemini_embedding_url(
|
"gemini:embedding" => build_gemini_embedding_url(
|
||||||
&transport.endpoint.base_url,
|
request_base_url,
|
||||||
params.mapped_model?,
|
params.mapped_model?,
|
||||||
params.request_query,
|
params.request_query,
|
||||||
gemini_embedding_batch,
|
gemini_embedding_batch,
|
||||||
),
|
),
|
||||||
"gemini:interactions" => {
|
"gemini:interactions" => {
|
||||||
build_gemini_interactions_url(&transport.endpoint.base_url, params.request_query)
|
build_gemini_interactions_url(request_base_url, params.request_query)
|
||||||
|
}
|
||||||
|
"doubao:embedding" => {
|
||||||
|
build_passthrough_path_url(request_base_url, "/embeddings", params.request_query, &[])
|
||||||
}
|
}
|
||||||
"doubao:embedding" => build_passthrough_path_url(
|
|
||||||
&transport.endpoint.base_url,
|
|
||||||
"/embeddings",
|
|
||||||
params.request_query,
|
|
||||||
&[],
|
|
||||||
),
|
|
||||||
_ => None,
|
_ => None,
|
||||||
}?;
|
}?;
|
||||||
|
|
||||||
@@ -2417,4 +2416,82 @@ mod tests {
|
|||||||
"https://api.example.com/v1/messages?model=claude%26admin%3Dtrue%23fragment"
|
"https://api.example.com/v1/messages?model=claude%26admin%3Dtrue%23fragment"
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_oauth_responses_use_cli_chat_proxy() {
|
||||||
|
let mut transport = sample_transport(
|
||||||
|
"xai",
|
||||||
|
"openai:responses",
|
||||||
|
"https://cli-chat-proxy.grok.com/v1",
|
||||||
|
None,
|
||||||
|
);
|
||||||
|
transport.key.auth_type = "oauth".to_string();
|
||||||
|
transport.key.decrypted_auth_config =
|
||||||
|
Some(r#"{"refresh_token":"rt","using_api":false}"#.to_string());
|
||||||
|
|
||||||
|
let url = build_transport_request_url(
|
||||||
|
&transport,
|
||||||
|
TransportRequestUrlParams {
|
||||||
|
provider_api_format: "openai:responses",
|
||||||
|
mapped_model: Some("grok-4"),
|
||||||
|
upstream_is_stream: true,
|
||||||
|
request_query: None,
|
||||||
|
kiro_api_region: None,
|
||||||
|
api_operation: None,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
.expect("xai oauth responses URL");
|
||||||
|
|
||||||
|
assert_eq!(url, "https://cli-chat-proxy.grok.com/v1/responses");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_compact_and_using_api_use_official_api() {
|
||||||
|
let mut oauth = sample_transport(
|
||||||
|
"xai",
|
||||||
|
"openai:responses:compact",
|
||||||
|
"https://cli-chat-proxy.grok.com/v1",
|
||||||
|
None,
|
||||||
|
);
|
||||||
|
oauth.key.auth_type = "oauth".to_string();
|
||||||
|
oauth.key.decrypted_auth_config =
|
||||||
|
Some(r#"{"refresh_token":"rt","using_api":false}"#.to_string());
|
||||||
|
|
||||||
|
let compact = build_transport_request_url(
|
||||||
|
&oauth,
|
||||||
|
TransportRequestUrlParams {
|
||||||
|
provider_api_format: "openai:responses:compact",
|
||||||
|
mapped_model: Some("grok-4"),
|
||||||
|
upstream_is_stream: false,
|
||||||
|
request_query: None,
|
||||||
|
kiro_api_region: None,
|
||||||
|
api_operation: None,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
.expect("xai compact URL");
|
||||||
|
assert_eq!(compact, "https://api.x.ai/v1/responses/compact");
|
||||||
|
|
||||||
|
let mut api_key = sample_transport(
|
||||||
|
"xai",
|
||||||
|
"openai:responses",
|
||||||
|
"https://cli-chat-proxy.grok.com/v1",
|
||||||
|
None,
|
||||||
|
);
|
||||||
|
api_key.key.auth_type = "oauth".to_string();
|
||||||
|
api_key.key.decrypted_auth_config = Some(r#"{"using_api":true}"#.to_string());
|
||||||
|
|
||||||
|
let official = build_transport_request_url(
|
||||||
|
&api_key,
|
||||||
|
TransportRequestUrlParams {
|
||||||
|
provider_api_format: "openai:responses",
|
||||||
|
mapped_model: Some("grok-4"),
|
||||||
|
upstream_is_stream: true,
|
||||||
|
request_query: None,
|
||||||
|
kiro_api_region: None,
|
||||||
|
api_operation: None,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
.expect("xai api key URL");
|
||||||
|
assert_eq!(official, "https://api.x.ai/v1/responses");
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -396,6 +396,12 @@ pub fn build_standard_provider_request_headers(
|
|||||||
force_identity_accept_encoding(&mut headers);
|
force_identity_accept_encoding(&mut headers);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
crate::xai::insert_cli_identity_headers_if_needed(
|
||||||
|
input.transport,
|
||||||
|
input.provider_api_format,
|
||||||
|
&mut headers,
|
||||||
|
);
|
||||||
|
|
||||||
let declared_connection_headers =
|
let declared_connection_headers =
|
||||||
crate::headers::declared_connection_header_names(input.headers, input.extra_headers);
|
crate::headers::declared_connection_header_names(input.headers, input.extra_headers);
|
||||||
crate::headers::remove_declared_connection_headers(&mut headers, &declared_connection_headers);
|
crate::headers::remove_declared_connection_headers(&mut headers, &declared_connection_headers);
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
use std::collections::BTreeMap;
|
use std::collections::BTreeMap;
|
||||||
use std::fmt;
|
use std::fmt;
|
||||||
|
|
||||||
|
use aether_contracts::ProxySnapshot;
|
||||||
use aether_data_contracts::repository::video_tasks::StoredVideoTask;
|
use aether_data_contracts::repository::video_tasks::StoredVideoTask;
|
||||||
use aether_video_tasks_core::{
|
use aether_video_tasks_core::{
|
||||||
LocalVideoTaskSnapshot, LocalVideoTaskTransport, LocalVideoTaskTransportBridgeInput,
|
LocalVideoTaskSnapshot, LocalVideoTaskTransport, LocalVideoTaskTransportBridgeInput,
|
||||||
@@ -12,11 +13,13 @@ use super::auth::{
|
|||||||
build_passthrough_headers_with_auth, resolve_local_gemini_auth,
|
build_passthrough_headers_with_auth, resolve_local_gemini_auth,
|
||||||
resolve_local_openai_bearer_auth,
|
resolve_local_openai_bearer_auth,
|
||||||
};
|
};
|
||||||
use super::network::{resolve_transport_execution_timeouts, resolve_transport_profile};
|
use super::network::{
|
||||||
|
resolve_transport_execution_timeouts, resolve_transport_profile,
|
||||||
|
resolve_transport_proxy_snapshot,
|
||||||
|
};
|
||||||
use super::policy::{
|
use super::policy::{
|
||||||
local_gemini_transport_unsupported_reason_with_network,
|
local_gemini_transport_unsupported_reason_with_network,
|
||||||
local_standard_transport_unsupported_reason_with_network, supports_local_gemini_transport,
|
local_standard_transport_unsupported_reason_with_network,
|
||||||
supports_local_standard_transport,
|
|
||||||
};
|
};
|
||||||
use super::rules::{
|
use super::rules::{
|
||||||
apply_local_body_rules_with_request_headers, apply_local_header_rules_with_request_headers,
|
apply_local_body_rules_with_request_headers, apply_local_header_rules_with_request_headers,
|
||||||
@@ -32,6 +35,7 @@ pub enum ProviderVideoCreateFamily {
|
|||||||
|
|
||||||
#[derive(Clone, Copy)]
|
#[derive(Clone, Copy)]
|
||||||
pub struct ProviderVideoCreateHeadersInput<'a> {
|
pub struct ProviderVideoCreateHeadersInput<'a> {
|
||||||
|
pub transport: &'a GatewayProviderTransportSnapshot,
|
||||||
pub headers: &'a http::HeaderMap,
|
pub headers: &'a http::HeaderMap,
|
||||||
pub auth_header: &'a str,
|
pub auth_header: &'a str,
|
||||||
pub auth_value: &'a str,
|
pub auth_value: &'a str,
|
||||||
@@ -79,6 +83,13 @@ pub trait VideoTaskTransportSnapshotLookup: Send + Sync {
|
|||||||
endpoint_id: &str,
|
endpoint_id: &str,
|
||||||
key_id: &str,
|
key_id: &str,
|
||||||
) -> Result<Option<GatewayProviderTransportSnapshot>, String>;
|
) -> Result<Option<GatewayProviderTransportSnapshot>, String>;
|
||||||
|
|
||||||
|
async fn resolve_video_task_proxy(
|
||||||
|
&self,
|
||||||
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
|
) -> Option<ProxySnapshot> {
|
||||||
|
resolve_transport_proxy_snapshot(transport)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn resolve_local_video_task_transport(
|
pub fn resolve_local_video_task_transport(
|
||||||
@@ -89,13 +100,17 @@ pub fn resolve_local_video_task_transport(
|
|||||||
let api_format = api_format.trim();
|
let api_format = api_format.trim();
|
||||||
let (auth_header, auth_value) = match api_format {
|
let (auth_header, auth_value) = match api_format {
|
||||||
"openai:video" => {
|
"openai:video" => {
|
||||||
if !supports_local_standard_transport(transport, api_format) {
|
if local_standard_transport_unsupported_reason_with_network(transport, api_format)
|
||||||
|
.is_some()
|
||||||
|
{
|
||||||
return None;
|
return None;
|
||||||
}
|
}
|
||||||
resolve_local_openai_bearer_auth(transport)?
|
resolve_openai_compatible_video_auth(transport)?
|
||||||
}
|
}
|
||||||
"gemini:video" => {
|
"gemini:video" => {
|
||||||
if !supports_local_gemini_transport(transport, api_format) {
|
if local_gemini_transport_unsupported_reason_with_network(transport, api_format)
|
||||||
|
.is_some()
|
||||||
|
{
|
||||||
return None;
|
return None;
|
||||||
}
|
}
|
||||||
resolve_local_gemini_auth(transport)?
|
resolve_local_gemini_auth(transport)?
|
||||||
@@ -103,9 +118,9 @@ pub fn resolve_local_video_task_transport(
|
|||||||
_ => return None,
|
_ => return None,
|
||||||
};
|
};
|
||||||
|
|
||||||
Some(LocalVideoTaskTransport::from_bridge_input(
|
let mut resolved =
|
||||||
LocalVideoTaskTransportBridgeInput {
|
LocalVideoTaskTransport::from_bridge_input(LocalVideoTaskTransportBridgeInput {
|
||||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
upstream_base_url: crate::xai::resolved_xai_request_base_url(transport, api_format),
|
||||||
provider_name: Some(transport.provider.name.clone()),
|
provider_name: Some(transport.provider.name.clone()),
|
||||||
provider_id: transport.provider.id.clone(),
|
provider_id: transport.provider.id.clone(),
|
||||||
endpoint_id: transport.endpoint.id.clone(),
|
endpoint_id: transport.endpoint.id.clone(),
|
||||||
@@ -114,11 +129,12 @@ pub fn resolve_local_video_task_transport(
|
|||||||
auth_value,
|
auth_value,
|
||||||
content_type: Some("application/json".to_string()),
|
content_type: Some("application/json".to_string()),
|
||||||
model_name,
|
model_name,
|
||||||
proxy: None,
|
proxy: resolve_transport_proxy_snapshot(transport),
|
||||||
transport_profile: resolve_transport_profile(transport),
|
transport_profile: resolve_transport_profile(transport),
|
||||||
timeouts: resolve_transport_execution_timeouts(transport),
|
timeouts: resolve_transport_execution_timeouts(transport),
|
||||||
},
|
});
|
||||||
))
|
crate::xai::insert_cli_identity_headers_if_needed(transport, api_format, &mut resolved.headers);
|
||||||
|
Some(resolved)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn video_create_transport_unsupported_reason(
|
pub fn video_create_transport_unsupported_reason(
|
||||||
@@ -141,7 +157,7 @@ pub fn resolve_video_create_auth(
|
|||||||
family: ProviderVideoCreateFamily,
|
family: ProviderVideoCreateFamily,
|
||||||
) -> Option<(String, String)> {
|
) -> Option<(String, String)> {
|
||||||
match family {
|
match family {
|
||||||
ProviderVideoCreateFamily::OpenAi => resolve_local_openai_bearer_auth(transport),
|
ProviderVideoCreateFamily::OpenAi => resolve_openai_compatible_video_auth(transport),
|
||||||
ProviderVideoCreateFamily::Gemini => resolve_local_gemini_auth(transport),
|
ProviderVideoCreateFamily::Gemini => resolve_local_gemini_auth(transport),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -173,6 +189,15 @@ pub fn build_video_create_request_body(
|
|||||||
Some(provider_request_body)
|
Some(provider_request_body)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn resolve_openai_compatible_video_auth(
|
||||||
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
|
) -> Option<(String, String)> {
|
||||||
|
resolve_local_openai_bearer_auth(transport).or_else(|| {
|
||||||
|
crate::generic_oauth::resolve_local_generic_oauth_transport_authorization(transport)
|
||||||
|
.map(|value| ("authorization".to_string(), value))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
pub fn build_video_create_upstream_url(
|
pub fn build_video_create_upstream_url(
|
||||||
transport: &GatewayProviderTransportSnapshot,
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
request_path: &str,
|
request_path: &str,
|
||||||
@@ -193,7 +218,13 @@ pub fn build_video_create_upstream_url(
|
|||||||
ProviderVideoCreateFamily::Gemini => &["key"][..],
|
ProviderVideoCreateFamily::Gemini => &["key"][..],
|
||||||
};
|
};
|
||||||
return build_passthrough_path_url(
|
return build_passthrough_path_url(
|
||||||
&transport.endpoint.base_url,
|
&crate::xai::resolved_xai_request_base_url(
|
||||||
|
transport,
|
||||||
|
match family {
|
||||||
|
ProviderVideoCreateFamily::OpenAi => "openai:video",
|
||||||
|
ProviderVideoCreateFamily::Gemini => "gemini:video",
|
||||||
|
},
|
||||||
|
),
|
||||||
path,
|
path,
|
||||||
request_query,
|
request_query,
|
||||||
blocked_keys,
|
blocked_keys,
|
||||||
@@ -202,8 +233,14 @@ pub fn build_video_create_upstream_url(
|
|||||||
|
|
||||||
match family {
|
match family {
|
||||||
ProviderVideoCreateFamily::OpenAi => build_passthrough_path_url(
|
ProviderVideoCreateFamily::OpenAi => build_passthrough_path_url(
|
||||||
&transport.endpoint.base_url,
|
&crate::xai::resolved_xai_request_base_url(transport, "openai:video"),
|
||||||
openai_video_api_root_request_path(request_path),
|
if crate::xai::is_xai_provider_transport(transport)
|
||||||
|
&& matches!(request_path, "/v1/videos" | "/openai/v1/videos")
|
||||||
|
{
|
||||||
|
"/videos/generations"
|
||||||
|
} else {
|
||||||
|
openai_video_api_root_request_path(request_path)
|
||||||
|
},
|
||||||
request_query,
|
request_query,
|
||||||
&[],
|
&[],
|
||||||
),
|
),
|
||||||
@@ -216,6 +253,7 @@ pub fn build_video_create_upstream_url(
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn openai_video_api_root_request_path(request_path: &str) -> &str {
|
fn openai_video_api_root_request_path(request_path: &str) -> &str {
|
||||||
|
let request_path = request_path.strip_prefix("/openai").unwrap_or(request_path);
|
||||||
if request_path.starts_with("/v1/") {
|
if request_path.starts_with("/v1/") {
|
||||||
&request_path[3..]
|
&request_path[3..]
|
||||||
} else {
|
} else {
|
||||||
@@ -232,6 +270,11 @@ pub fn build_video_create_headers(
|
|||||||
input.auth_value,
|
input.auth_value,
|
||||||
&BTreeMap::new(),
|
&BTreeMap::new(),
|
||||||
);
|
);
|
||||||
|
crate::xai::insert_cli_identity_headers_if_needed(
|
||||||
|
input.transport,
|
||||||
|
"openai:video",
|
||||||
|
&mut provider_request_headers,
|
||||||
|
);
|
||||||
if !apply_local_header_rules_with_request_headers(
|
if !apply_local_header_rules_with_request_headers(
|
||||||
&mut provider_request_headers,
|
&mut provider_request_headers,
|
||||||
input.header_rules,
|
input.header_rules,
|
||||||
@@ -281,16 +324,22 @@ pub async fn reconstruct_local_video_task_snapshot(
|
|||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
|
|
||||||
let Some(local_transport) =
|
let Some(mut local_transport) =
|
||||||
resolve_local_video_task_transport(&transport, provider_api_format, task.model.clone())
|
resolve_local_video_task_transport(&transport, provider_api_format, task.model.clone())
|
||||||
else {
|
else {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
|
|
||||||
Ok(LocalVideoTaskSnapshot::from_stored_task_with_transport(
|
// Resolve deployment-managed nodes, system defaults and tunnel affinity just as
|
||||||
task,
|
// creation does; serialized task metadata intentionally contains no credentials.
|
||||||
local_transport,
|
local_transport.proxy = lookup.resolve_video_task_proxy(&transport).await;
|
||||||
))
|
|
||||||
|
let mut snapshot =
|
||||||
|
LocalVideoTaskSnapshot::from_stored_task_with_transport(task, local_transport);
|
||||||
|
if let Some(LocalVideoTaskSnapshot::OpenAi(seed)) = &mut snapshot {
|
||||||
|
seed.xai_provider = crate::xai::is_xai_provider_transport(&transport);
|
||||||
|
}
|
||||||
|
Ok(snapshot)
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
@@ -441,6 +490,46 @@ mod tests {
|
|||||||
assert_eq!(transport.provider_id, "provider-1");
|
assert_eq!(transport.provider_id, "provider-1");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn reconstructs_video_with_configured_proxy_and_profile() {
|
||||||
|
let mut transport = sample_transport("openai:video", "oauth");
|
||||||
|
transport.provider.provider_type = "xai".into();
|
||||||
|
transport.endpoint.base_url = "https://cli-chat-proxy.grok.com/v1".into();
|
||||||
|
transport.provider.proxy = Some(json!({"enabled":true,"url":"http://127.0.0.1:9876"}));
|
||||||
|
transport.provider.config = Some(json!({"fingerprint":{"transport_profile":{
|
||||||
|
"profile_id":"test-video","backend":"reqwest_rustls","http_mode":"auto","pool_scope":"key"
|
||||||
|
}}}));
|
||||||
|
transport.key.decrypted_auth_config = Some(r#"{"using_api":false}"#.into());
|
||||||
|
let lookup = TestLookup(Some(transport));
|
||||||
|
let snapshot = reconstruct_local_video_task_snapshot(&lookup, &sample_stored_video_task())
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.expect("proxied video must resume after restart");
|
||||||
|
let LocalVideoTaskSnapshot::OpenAi(seed) = snapshot else {
|
||||||
|
panic!("expected OpenAI video")
|
||||||
|
};
|
||||||
|
assert!(seed.xai_provider);
|
||||||
|
assert_eq!(
|
||||||
|
seed.transport.proxy.as_ref().unwrap().url.as_deref(),
|
||||||
|
Some("http://127.0.0.1:9876/")
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
seed.transport
|
||||||
|
.transport_profile
|
||||||
|
.as_ref()
|
||||||
|
.unwrap()
|
||||||
|
.profile_id,
|
||||||
|
"test-video"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
seed.transport
|
||||||
|
.headers
|
||||||
|
.get("x-xai-token-auth")
|
||||||
|
.map(String::as_str),
|
||||||
|
Some("xai-grok-cli")
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn resolves_gemini_video_transport() {
|
fn resolves_gemini_video_transport() {
|
||||||
let transport = resolve_local_video_task_transport(
|
let transport = resolve_local_video_task_transport(
|
||||||
@@ -489,6 +578,120 @@ mod tests {
|
|||||||
assert_eq!(url, "https://api.openai.example/v1/videos?trace=1");
|
assert_eq!(url, "https://api.openai.example/v1/videos?trace=1");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_video_create_paths_preserve_auth_hosts_and_custom_endpoints() {
|
||||||
|
for (auth, base) in [
|
||||||
|
("oauth", "https://cli-chat-proxy.grok.com/v1"),
|
||||||
|
("api_key", "https://api.x.ai/v1"),
|
||||||
|
] {
|
||||||
|
let mut transport = sample_transport("openai:video", auth);
|
||||||
|
transport.provider.provider_type = "xai".into();
|
||||||
|
transport.endpoint.base_url = "https://cli-chat-proxy.grok.com/v1".into();
|
||||||
|
transport.key.decrypted_auth_config =
|
||||||
|
(auth == "oauth").then(|| r#"{"using_api":false}"#.into());
|
||||||
|
for path in ["/v1/videos", "/openai/v1/videos", "/v1/videos/generations"] {
|
||||||
|
assert_eq!(
|
||||||
|
build_video_create_upstream_url(
|
||||||
|
&transport,
|
||||||
|
path,
|
||||||
|
Some("trace=1"),
|
||||||
|
"grok-imagine-video",
|
||||||
|
ProviderVideoCreateFamily::OpenAi
|
||||||
|
)
|
||||||
|
.unwrap(),
|
||||||
|
format!("{base}/videos/generations?trace=1")
|
||||||
|
);
|
||||||
|
}
|
||||||
|
transport.endpoint.base_url = "https://gateway.example/prefix/v1".into();
|
||||||
|
assert_eq!(
|
||||||
|
build_video_create_upstream_url(
|
||||||
|
&transport,
|
||||||
|
"/openai/v1/videos",
|
||||||
|
None,
|
||||||
|
"grok-imagine-video",
|
||||||
|
ProviderVideoCreateFamily::OpenAi
|
||||||
|
)
|
||||||
|
.unwrap(),
|
||||||
|
"https://gateway.example/prefix/v1/videos/generations"
|
||||||
|
);
|
||||||
|
transport.endpoint.custom_path = Some("/custom/videos/generations".into());
|
||||||
|
let url = build_video_create_upstream_url(
|
||||||
|
&transport,
|
||||||
|
"/openai/v1/videos",
|
||||||
|
None,
|
||||||
|
"grok-imagine-video",
|
||||||
|
ProviderVideoCreateFamily::OpenAi,
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
assert!(url.ends_with("/custom/videos/generations"), "{url}");
|
||||||
|
}
|
||||||
|
let transport = sample_transport("openai:video", "api_key");
|
||||||
|
assert_eq!(
|
||||||
|
build_video_create_upstream_url(
|
||||||
|
&transport,
|
||||||
|
"/openai/v1/videos",
|
||||||
|
None,
|
||||||
|
"sora",
|
||||||
|
ProviderVideoCreateFamily::OpenAi
|
||||||
|
),
|
||||||
|
build_video_create_upstream_url(
|
||||||
|
&transport,
|
||||||
|
"/v1/videos",
|
||||||
|
None,
|
||||||
|
"sora",
|
||||||
|
ProviderVideoCreateFamily::OpenAi
|
||||||
|
)
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_oauth_video_uses_cli_proxy() {
|
||||||
|
let mut transport = sample_transport("openai:video", "oauth");
|
||||||
|
transport.provider.provider_type = "xai".to_string();
|
||||||
|
transport.endpoint.base_url = "https://cli-chat-proxy.grok.com/v1".to_string();
|
||||||
|
transport.key.decrypted_auth_config =
|
||||||
|
Some(r#"{"refresh_token":"rt","using_api":false}"#.to_string());
|
||||||
|
let url = build_video_create_upstream_url(
|
||||||
|
&transport,
|
||||||
|
"/v1/videos/generations",
|
||||||
|
None,
|
||||||
|
"grok-imagine-video",
|
||||||
|
ProviderVideoCreateFamily::OpenAi,
|
||||||
|
)
|
||||||
|
.expect("url should build");
|
||||||
|
|
||||||
|
assert_eq!(url, "https://cli-chat-proxy.grok.com/v1/videos/generations");
|
||||||
|
let headers = build_video_create_headers(ProviderVideoCreateHeadersInput {
|
||||||
|
transport: &transport,
|
||||||
|
headers: &http::HeaderMap::new(),
|
||||||
|
auth_header: "authorization",
|
||||||
|
auth_value: "Bearer test-token",
|
||||||
|
header_rules: None,
|
||||||
|
provider_request_body: &json!({"prompt": "A cat"}),
|
||||||
|
original_request_body: &json!({"prompt": "A cat"}),
|
||||||
|
})
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
headers.get("x-xai-token-auth").map(String::as_str),
|
||||||
|
Some("xai-grok-cli")
|
||||||
|
);
|
||||||
|
|
||||||
|
let reconstructed = super::resolve_local_video_task_transport(
|
||||||
|
&transport,
|
||||||
|
"openai:video",
|
||||||
|
Some("grok-imagine-video".into()),
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
reconstructed.upstream_base_url,
|
||||||
|
"https://cli-chat-proxy.grok.com/v1"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
reconstructed.headers.get("x-xai-token-auth"),
|
||||||
|
headers.get("x-xai-token-auth")
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn builds_gemini_video_create_url_and_removes_client_key_query() {
|
fn builds_gemini_video_create_url_and_removes_client_key_query() {
|
||||||
let transport = sample_transport("gemini:video", "api_key");
|
let transport = sample_transport("gemini:video", "api_key");
|
||||||
@@ -512,6 +715,7 @@ mod tests {
|
|||||||
let provider_request_body = json!({"prompt": "make a clip"});
|
let provider_request_body = json!({"prompt": "make a clip"});
|
||||||
let original_request_body = provider_request_body.clone();
|
let original_request_body = provider_request_body.clone();
|
||||||
let headers = build_video_create_headers(ProviderVideoCreateHeadersInput {
|
let headers = build_video_create_headers(ProviderVideoCreateHeadersInput {
|
||||||
|
transport: &sample_transport("openai:video", "bearer"),
|
||||||
headers: &http::HeaderMap::new(),
|
headers: &http::HeaderMap::new(),
|
||||||
auth_header: "authorization",
|
auth_header: "authorization",
|
||||||
auth_value: "Bearer secret",
|
auth_value: "Bearer secret",
|
||||||
|
|||||||
@@ -0,0 +1,444 @@
|
|||||||
|
pub mod video;
|
||||||
|
|
||||||
|
use std::collections::BTreeMap;
|
||||||
|
|
||||||
|
use aether_ai_formats::normalize_api_format_alias;
|
||||||
|
use serde_json::Value;
|
||||||
|
|
||||||
|
use crate::snapshot::GatewayProviderTransportSnapshot;
|
||||||
|
|
||||||
|
pub const XAI_PROVIDER_TYPE: &str = "xai";
|
||||||
|
pub const XAI_CHAT_PROXY_BASE_URL: &str = "https://cli-chat-proxy.grok.com/v1";
|
||||||
|
pub const XAI_API_BASE_URL: &str = "https://api.x.ai/v1";
|
||||||
|
pub const XAI_CLIENT_VERSION: &str = "0.2.120";
|
||||||
|
pub const XAI_TOKEN_AUTH_HEADER: &str = "x-xai-token-auth";
|
||||||
|
pub const XAI_TOKEN_AUTH_VALUE: &str = "xai-grok-cli";
|
||||||
|
pub const XAI_CLIENT_VERSION_HEADER: &str = "x-grok-client-version";
|
||||||
|
pub const XAI_CLIENT_IDENTIFIER_HEADER: &str = "x-grok-client-identifier";
|
||||||
|
pub const XAI_CLIENT_IDENTIFIER_VALUE: &str = "grok-shell";
|
||||||
|
pub const XAI_AUTHENTICATE_RESPONSE_HEADER: &str = "x-authenticateresponse";
|
||||||
|
pub const XAI_AUTHENTICATE_RESPONSE_VALUE: &str = "authenticate-response";
|
||||||
|
|
||||||
|
pub fn xai_cli_user_agent() -> String {
|
||||||
|
format!("xai-grok-workspace/{XAI_CLIENT_VERSION}")
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn is_xai_provider_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||||
|
transport
|
||||||
|
.provider
|
||||||
|
.provider_type
|
||||||
|
.trim()
|
||||||
|
.eq_ignore_ascii_case(XAI_PROVIDER_TYPE)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn xai_uses_official_api(api_format: &str) -> bool {
|
||||||
|
matches!(
|
||||||
|
normalize_api_format_alias(api_format).as_str(),
|
||||||
|
"openai:responses:compact"
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn resolved_xai_upstream_base_url(
|
||||||
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
|
api_format: &str,
|
||||||
|
) -> Option<String> {
|
||||||
|
if !is_xai_provider_transport(transport) {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
let stored = transport.endpoint.base_url.trim();
|
||||||
|
if xai_uses_official_api(api_format) {
|
||||||
|
if stored.is_empty()
|
||||||
|
|| is_cli_chat_proxy_base_url(stored)
|
||||||
|
|| is_official_api_base_url(stored)
|
||||||
|
{
|
||||||
|
return Some(XAI_API_BASE_URL.to_string());
|
||||||
|
}
|
||||||
|
return Some(trim_base_url(stored));
|
||||||
|
}
|
||||||
|
if xai_using_api(transport) {
|
||||||
|
if stored.is_empty() || is_cli_chat_proxy_base_url(stored) {
|
||||||
|
return Some(XAI_API_BASE_URL.to_string());
|
||||||
|
}
|
||||||
|
return Some(trim_base_url(stored));
|
||||||
|
}
|
||||||
|
if stored.is_empty() || is_official_api_base_url(stored) {
|
||||||
|
return Some(XAI_CHAT_PROXY_BASE_URL.to_string());
|
||||||
|
}
|
||||||
|
Some(trim_base_url(stored))
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn resolved_xai_request_base_url(
|
||||||
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
|
api_format: &str,
|
||||||
|
) -> String {
|
||||||
|
resolved_xai_upstream_base_url(transport, api_format)
|
||||||
|
.unwrap_or_else(|| trim_base_url(&transport.endpoint.base_url))
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn should_attach_cli_identity_headers(
|
||||||
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
|
api_format: &str,
|
||||||
|
) -> bool {
|
||||||
|
if !is_xai_provider_transport(transport) {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
if xai_uses_official_api(api_format) {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
resolved_xai_upstream_base_url(transport, api_format)
|
||||||
|
.as_deref()
|
||||||
|
.is_some_and(is_cli_chat_proxy_base_url)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn insert_cli_identity_headers(headers: &mut BTreeMap<String, String>) {
|
||||||
|
let user_agent = xai_cli_user_agent();
|
||||||
|
for (name, value) in [
|
||||||
|
(XAI_TOKEN_AUTH_HEADER, XAI_TOKEN_AUTH_VALUE),
|
||||||
|
(XAI_CLIENT_VERSION_HEADER, XAI_CLIENT_VERSION),
|
||||||
|
("user-agent", user_agent.as_str()),
|
||||||
|
(XAI_CLIENT_IDENTIFIER_HEADER, XAI_CLIENT_IDENTIFIER_VALUE),
|
||||||
|
(
|
||||||
|
XAI_AUTHENTICATE_RESPONSE_HEADER,
|
||||||
|
XAI_AUTHENTICATE_RESPONSE_VALUE,
|
||||||
|
),
|
||||||
|
] {
|
||||||
|
if !headers
|
||||||
|
.keys()
|
||||||
|
.any(|existing| existing.eq_ignore_ascii_case(name))
|
||||||
|
{
|
||||||
|
headers.insert(name.to_string(), value.to_string());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn insert_cli_identity_headers_if_needed(
|
||||||
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
|
api_format: &str,
|
||||||
|
headers: &mut BTreeMap<String, String>,
|
||||||
|
) {
|
||||||
|
if should_attach_cli_identity_headers(transport, api_format) {
|
||||||
|
insert_cli_identity_headers(headers);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn xai_auth_uses_api(auth_type: &str, decrypted_auth_config: Option<&str>) -> bool {
|
||||||
|
if let Some(value) = auth_config_using_api(decrypted_auth_config) {
|
||||||
|
return value;
|
||||||
|
}
|
||||||
|
let auth_type = auth_type.trim().to_ascii_lowercase();
|
||||||
|
if auth_type == "oauth" || auth_config_has_refresh_token(decrypted_auth_config) {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
matches!(auth_type.as_str(), "api_key" | "bearer" | "apikey")
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn extract_xai_user_id_from_auth_config(raw_auth_config: Option<&str>) -> Option<String> {
|
||||||
|
let value = parse_auth_config(raw_auth_config)?;
|
||||||
|
extract_xai_user_id_from_value(&value)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn extract_xai_user_id_from_value(value: &Value) -> Option<String> {
|
||||||
|
const PATHS: &[&[&str]] = &[
|
||||||
|
&["userId"],
|
||||||
|
&["user_id"],
|
||||||
|
&["id"],
|
||||||
|
&["sub"],
|
||||||
|
&["user", "userId"],
|
||||||
|
&["user", "id"],
|
||||||
|
&["user", "user_id"],
|
||||||
|
&["user", "sub"],
|
||||||
|
];
|
||||||
|
PATHS.iter().find_map(|path| {
|
||||||
|
let mut current = value;
|
||||||
|
for key in *path {
|
||||||
|
current = current.get(*key)?;
|
||||||
|
}
|
||||||
|
coerce_xai_id(current)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn xai_using_api(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||||
|
xai_auth_uses_api(
|
||||||
|
transport.key.auth_type.as_str(),
|
||||||
|
transport.key.decrypted_auth_config.as_deref(),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn coerce_xai_id(value: &Value) -> Option<String> {
|
||||||
|
match value {
|
||||||
|
Value::String(text) => {
|
||||||
|
let trimmed = text.trim();
|
||||||
|
(!trimmed.is_empty()).then(|| trimmed.to_string())
|
||||||
|
}
|
||||||
|
Value::Number(number) => {
|
||||||
|
let rendered = number.to_string();
|
||||||
|
(!rendered.is_empty()).then_some(rendered)
|
||||||
|
}
|
||||||
|
_ => None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn auth_config_using_api(raw_auth_config: Option<&str>) -> Option<bool> {
|
||||||
|
let value = parse_auth_config(raw_auth_config)?;
|
||||||
|
let using_api = value.get("using_api")?;
|
||||||
|
match using_api {
|
||||||
|
Value::Bool(value) => Some(*value),
|
||||||
|
Value::String(value) => value.trim().parse::<bool>().ok(),
|
||||||
|
_ => None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn auth_config_has_refresh_token(raw_auth_config: Option<&str>) -> bool {
|
||||||
|
let value = match parse_auth_config(raw_auth_config) {
|
||||||
|
Some(value) => value,
|
||||||
|
None => return false,
|
||||||
|
};
|
||||||
|
["refresh_token", "refreshToken"]
|
||||||
|
.iter()
|
||||||
|
.find_map(|field| value.get(*field).and_then(Value::as_str))
|
||||||
|
.map(str::trim)
|
||||||
|
.is_some_and(|value| !value.is_empty())
|
||||||
|
}
|
||||||
|
|
||||||
|
fn parse_auth_config(raw_auth_config: Option<&str>) -> Option<Value> {
|
||||||
|
raw_auth_config
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.and_then(|value| serde_json::from_str::<Value>(value).ok())
|
||||||
|
}
|
||||||
|
|
||||||
|
fn trim_base_url(url: &str) -> String {
|
||||||
|
url.trim().trim_end_matches('/').to_string()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn normalize_base_url(url: &str) -> String {
|
||||||
|
trim_base_url(url).to_ascii_lowercase()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn is_official_api_base_url(url: &str) -> bool {
|
||||||
|
normalize_base_url(url) == normalize_base_url(XAI_API_BASE_URL)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn is_cli_chat_proxy_base_url(url: &str) -> bool {
|
||||||
|
normalize_base_url(url) == normalize_base_url(XAI_CHAT_PROXY_BASE_URL)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::{
|
||||||
|
insert_cli_identity_headers_if_needed, is_xai_provider_transport,
|
||||||
|
resolved_xai_upstream_base_url, should_attach_cli_identity_headers, XAI_API_BASE_URL,
|
||||||
|
XAI_CHAT_PROXY_BASE_URL, XAI_CLIENT_IDENTIFIER_VALUE, XAI_TOKEN_AUTH_VALUE,
|
||||||
|
};
|
||||||
|
use crate::snapshot::{
|
||||||
|
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||||
|
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
|
||||||
|
};
|
||||||
|
use std::collections::BTreeMap;
|
||||||
|
|
||||||
|
fn sample_transport(
|
||||||
|
auth_type: &str,
|
||||||
|
auth_config: Option<&str>,
|
||||||
|
base_url: &str,
|
||||||
|
) -> GatewayProviderTransportSnapshot {
|
||||||
|
GatewayProviderTransportSnapshot {
|
||||||
|
provider: GatewayProviderTransportProvider {
|
||||||
|
id: "provider-xai".to_string(),
|
||||||
|
name: "xAI".to_string(),
|
||||||
|
provider_type: "xai".to_string(),
|
||||||
|
website: None,
|
||||||
|
is_active: true,
|
||||||
|
keep_priority_on_conversion: false,
|
||||||
|
enable_format_conversion: true,
|
||||||
|
concurrent_limit: None,
|
||||||
|
max_retries: None,
|
||||||
|
proxy: None,
|
||||||
|
request_timeout_secs: None,
|
||||||
|
stream_first_byte_timeout_secs: None,
|
||||||
|
config: None,
|
||||||
|
},
|
||||||
|
endpoint: GatewayProviderTransportEndpoint {
|
||||||
|
id: "endpoint-xai".to_string(),
|
||||||
|
provider_id: "provider-xai".to_string(),
|
||||||
|
api_format: "openai:responses".to_string(),
|
||||||
|
api_family: None,
|
||||||
|
endpoint_kind: None,
|
||||||
|
is_active: true,
|
||||||
|
base_url: base_url.to_string(),
|
||||||
|
header_rules: None,
|
||||||
|
body_rules: None,
|
||||||
|
max_retries: None,
|
||||||
|
custom_path: None,
|
||||||
|
config: None,
|
||||||
|
format_acceptance_config: None,
|
||||||
|
proxy: None,
|
||||||
|
},
|
||||||
|
key: GatewayProviderTransportKey {
|
||||||
|
id: "key-xai".to_string(),
|
||||||
|
provider_id: "provider-xai".to_string(),
|
||||||
|
name: "key".to_string(),
|
||||||
|
auth_type: auth_type.to_string(),
|
||||||
|
is_active: true,
|
||||||
|
api_formats: None,
|
||||||
|
auth_type_by_format: None,
|
||||||
|
allow_auth_channel_mismatch_formats: None,
|
||||||
|
allowed_models: None,
|
||||||
|
capabilities: None,
|
||||||
|
rate_multipliers: None,
|
||||||
|
global_priority_by_format: None,
|
||||||
|
expires_at_unix_secs: None,
|
||||||
|
proxy: None,
|
||||||
|
fingerprint: None,
|
||||||
|
upstream_metadata: None,
|
||||||
|
decrypted_api_key: "access-token".to_string(),
|
||||||
|
decrypted_auth_config: auth_config.map(ToOwned::to_owned),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn oauth_defaults_to_cli_chat_proxy_for_responses() {
|
||||||
|
let transport = sample_transport(
|
||||||
|
"oauth",
|
||||||
|
Some(r#"{"refresh_token":"rt","using_api":false}"#),
|
||||||
|
XAI_CHAT_PROXY_BASE_URL,
|
||||||
|
);
|
||||||
|
assert!(is_xai_provider_transport(&transport));
|
||||||
|
assert_eq!(
|
||||||
|
resolved_xai_upstream_base_url(&transport, "openai:responses").as_deref(),
|
||||||
|
Some(XAI_CHAT_PROXY_BASE_URL)
|
||||||
|
);
|
||||||
|
assert!(should_attach_cli_identity_headers(
|
||||||
|
&transport,
|
||||||
|
"openai:responses"
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn compact_and_using_api_stay_on_official_api() {
|
||||||
|
let oauth = sample_transport(
|
||||||
|
"oauth",
|
||||||
|
Some(r#"{"refresh_token":"rt","using_api":false}"#),
|
||||||
|
XAI_CHAT_PROXY_BASE_URL,
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
resolved_xai_upstream_base_url(&oauth, "openai:responses:compact").as_deref(),
|
||||||
|
Some(XAI_API_BASE_URL)
|
||||||
|
);
|
||||||
|
assert!(!should_attach_cli_identity_headers(
|
||||||
|
&oauth,
|
||||||
|
"openai:responses:compact"
|
||||||
|
));
|
||||||
|
|
||||||
|
let api_key = sample_transport(
|
||||||
|
"oauth",
|
||||||
|
Some(r#"{"using_api":true}"#),
|
||||||
|
XAI_CHAT_PROXY_BASE_URL,
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
resolved_xai_upstream_base_url(&api_key, "openai:responses").as_deref(),
|
||||||
|
Some(XAI_API_BASE_URL)
|
||||||
|
);
|
||||||
|
assert!(!should_attach_cli_identity_headers(
|
||||||
|
&api_key,
|
||||||
|
"openai:responses"
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn media_routing_and_cli_headers_follow_auth_and_base_url() {
|
||||||
|
for api_format in ["openai:image", "openai:video"] {
|
||||||
|
for stored in ["", XAI_API_BASE_URL, XAI_CHAT_PROXY_BASE_URL] {
|
||||||
|
for (auth_type, config, expected) in [
|
||||||
|
(
|
||||||
|
"oauth",
|
||||||
|
Some(r#"{"refresh_token":"rt","using_api":false}"#),
|
||||||
|
XAI_CHAT_PROXY_BASE_URL,
|
||||||
|
),
|
||||||
|
("oauth", Some(r#"{"using_api":true}"#), XAI_API_BASE_URL),
|
||||||
|
("bearer", None, XAI_API_BASE_URL),
|
||||||
|
] {
|
||||||
|
let transport = sample_transport(auth_type, config, stored);
|
||||||
|
assert_eq!(
|
||||||
|
resolved_xai_upstream_base_url(&transport, api_format).as_deref(),
|
||||||
|
Some(expected)
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
should_attach_cli_identity_headers(&transport, api_format),
|
||||||
|
expected == XAI_CHAT_PROXY_BASE_URL
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
let custom = sample_transport("oauth", None, "https://custom.example/v1");
|
||||||
|
assert_eq!(
|
||||||
|
resolved_xai_upstream_base_url(&custom, api_format).as_deref(),
|
||||||
|
Some("https://custom.example/v1")
|
||||||
|
);
|
||||||
|
assert!(!should_attach_cli_identity_headers(&custom, api_format));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn bearer_without_refresh_uses_official_api() {
|
||||||
|
let transport = sample_transport("bearer", None, XAI_CHAT_PROXY_BASE_URL);
|
||||||
|
assert_eq!(
|
||||||
|
resolved_xai_upstream_base_url(&transport, "openai:responses").as_deref(),
|
||||||
|
Some(XAI_API_BASE_URL)
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn cli_headers_do_not_override_existing_values() {
|
||||||
|
let transport = sample_transport(
|
||||||
|
"oauth",
|
||||||
|
Some(r#"{"refresh_token":"rt"}"#),
|
||||||
|
XAI_CHAT_PROXY_BASE_URL,
|
||||||
|
);
|
||||||
|
let mut headers = BTreeMap::from([(
|
||||||
|
"x-grok-client-identifier".to_string(),
|
||||||
|
"custom-client".to_string(),
|
||||||
|
)]);
|
||||||
|
insert_cli_identity_headers_if_needed(&transport, "openai:responses", &mut headers);
|
||||||
|
assert_eq!(
|
||||||
|
headers.get("x-grok-client-identifier").map(String::as_str),
|
||||||
|
Some("custom-client")
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
headers.get("x-xai-token-auth").map(String::as_str),
|
||||||
|
Some(XAI_TOKEN_AUTH_VALUE)
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
headers.get("x-authenticateresponse").map(String::as_str),
|
||||||
|
Some("authenticate-response")
|
||||||
|
);
|
||||||
|
assert_ne!(
|
||||||
|
headers.get("x-grok-client-identifier").map(String::as_str),
|
||||||
|
Some(XAI_CLIENT_IDENTIFIER_VALUE)
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn extracts_user_id_from_user_payload_and_auth_config_sub() {
|
||||||
|
use super::{
|
||||||
|
extract_xai_user_id_from_auth_config, extract_xai_user_id_from_value, xai_auth_uses_api,
|
||||||
|
};
|
||||||
|
use serde_json::json;
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
extract_xai_user_id_from_value(&json!({"userId": "user-42"})).as_deref(),
|
||||||
|
Some("user-42")
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
extract_xai_user_id_from_auth_config(Some(r#"{"sub":"subject-1"}"#)).as_deref(),
|
||||||
|
Some("subject-1")
|
||||||
|
);
|
||||||
|
assert!(!xai_auth_uses_api(
|
||||||
|
"oauth",
|
||||||
|
Some(r#"{"refresh_token":"rt","using_api":false}"#)
|
||||||
|
));
|
||||||
|
assert!(xai_auth_uses_api(
|
||||||
|
"bearer",
|
||||||
|
Some(r#"{"api_key":"xai-key","using_api":true}"#)
|
||||||
|
));
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,147 @@
|
|||||||
|
use serde_json::{json, Value};
|
||||||
|
|
||||||
|
/// Native xAI video requests live under /v1; the OpenAI-compatible adapter under /openai/v1.
|
||||||
|
pub fn is_native_video_request(provider_type: &str, path: &str) -> bool {
|
||||||
|
provider_type.trim().eq_ignore_ascii_case("xai")
|
||||||
|
&& matches!(
|
||||||
|
path,
|
||||||
|
"/v1/videos" | "/v1/videos/generations" | "/v1/videos/edits" | "/v1/videos/extensions"
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn is_explicit_native_video_path(path: &str) -> bool {
|
||||||
|
matches!(
|
||||||
|
path,
|
||||||
|
"/v1/videos/generations" | "/v1/videos/edits" | "/v1/videos/extensions"
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Convert the OpenAI video request contract to xAI's native contract.
|
||||||
|
/// Native requests bypass this adapter so provider-specific fields remain intact.
|
||||||
|
pub fn convert_openai_video_request(body: &Value) -> Result<Value, &'static str> {
|
||||||
|
let prompt = text(&body["prompt"]).ok_or("prompt is required")?;
|
||||||
|
let seconds = match &body["seconds"] {
|
||||||
|
Value::Null => 4,
|
||||||
|
Value::String(value) if value.trim().is_empty() => 4,
|
||||||
|
Value::String(value) => value
|
||||||
|
.trim()
|
||||||
|
.parse::<i64>()
|
||||||
|
.map_err(|_| "seconds must be an integer")?,
|
||||||
|
value => value.as_i64().ok_or("seconds must be an integer")?,
|
||||||
|
}
|
||||||
|
.clamp(1, 15);
|
||||||
|
let size = text(&body["size"]).unwrap_or("720x1280");
|
||||||
|
let default_ratio = match size {
|
||||||
|
"720x1280" | "1024x1792" => "9:16",
|
||||||
|
"1280x720" | "1792x1024" => "16:9",
|
||||||
|
_ => return Err("size must be one of 720x1280, 1280x720, 1024x1792, or 1792x1024"),
|
||||||
|
};
|
||||||
|
let ratio = match text(&body["aspect_ratio"])
|
||||||
|
.unwrap_or("")
|
||||||
|
.to_ascii_lowercase()
|
||||||
|
.as_str()
|
||||||
|
{
|
||||||
|
"square" | "1:1" => "1:1",
|
||||||
|
"landscape" | "16:9" => "16:9",
|
||||||
|
"portrait" | "9:16" => "9:16",
|
||||||
|
"4:3" => "4:3",
|
||||||
|
"3:4" => "3:4",
|
||||||
|
"3:2" => "3:2",
|
||||||
|
"2:3" => "2:3",
|
||||||
|
_ => default_ratio,
|
||||||
|
};
|
||||||
|
let resolution = if text(&body["resolution"]).is_some_and(|v| v.eq_ignore_ascii_case("480p")) {
|
||||||
|
"480p"
|
||||||
|
} else {
|
||||||
|
"720p"
|
||||||
|
};
|
||||||
|
if text(&body["input_reference"]["file_id"]).is_some() {
|
||||||
|
return Err("input_reference.file_id is not supported for xAI video generation; use input_reference.image_url");
|
||||||
|
}
|
||||||
|
let image = text(&body["input_reference"]["image_url"])
|
||||||
|
.or_else(|| image_url(&body["image"]))
|
||||||
|
.or_else(|| text(&body["image_url"]));
|
||||||
|
let references: Vec<_> = ["reference_images", "reference_image_urls"]
|
||||||
|
.into_iter()
|
||||||
|
.filter_map(|key| body[key].as_array())
|
||||||
|
.flatten()
|
||||||
|
.filter_map(image_url)
|
||||||
|
.map(|url| json!({"url":url}))
|
||||||
|
.collect();
|
||||||
|
if references.len() > 7 {
|
||||||
|
return Err("reference_images supports at most 7 images on xAI");
|
||||||
|
}
|
||||||
|
if image.is_some() && !references.is_empty() {
|
||||||
|
return Err("image and reference_images cannot be combined on xAI");
|
||||||
|
}
|
||||||
|
let mut result = json!({"model":body["model"], "prompt":prompt, "duration":seconds, "aspect_ratio":ratio, "resolution":resolution});
|
||||||
|
if let Some(url) = image {
|
||||||
|
result["image"] = json!({"url":url});
|
||||||
|
}
|
||||||
|
if !references.is_empty() {
|
||||||
|
result["reference_images"] = json!(references);
|
||||||
|
}
|
||||||
|
Ok(result)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn text(value: &Value) -> Option<&str> {
|
||||||
|
value
|
||||||
|
.as_str()
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
}
|
||||||
|
|
||||||
|
fn image_url(value: &Value) -> Option<&str> {
|
||||||
|
text(value)
|
||||||
|
.or_else(|| text(&value["url"]))
|
||||||
|
.or_else(|| text(&value["image_url"]))
|
||||||
|
.or_else(|| text(&value["image_url"]["url"]))
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_video_compatibility_maps_duration_size_and_references() {
|
||||||
|
let converted = convert_openai_video_request(&json!({
|
||||||
|
"model":"grok-imagine-video", "prompt":"A cat", "seconds":"8", "size":"1280x720",
|
||||||
|
"reference_images":[{"image_url":{"url":"https://example.com/a.png"}}],
|
||||||
|
"reference_image_urls":["https://example.com/b.png"]
|
||||||
|
}))
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
converted,
|
||||||
|
json!({"model":"grok-imagine-video", "prompt":"A cat", "duration":8,
|
||||||
|
"aspect_ratio":"16:9", "resolution":"720p", "reference_images":[{"url":"https://example.com/a.png"},{"url":"https://example.com/b.png"}]})
|
||||||
|
);
|
||||||
|
let defaults = convert_openai_video_request(&json!({"prompt":"A cat"})).unwrap();
|
||||||
|
assert_eq!(defaults["duration"], 4);
|
||||||
|
assert_eq!(defaults["aspect_ratio"], "9:16");
|
||||||
|
for (seconds, expected) in [(-1, 1), (30, 15)] {
|
||||||
|
assert_eq!(
|
||||||
|
convert_openai_video_request(&json!({"prompt":"A cat", "seconds":seconds}))
|
||||||
|
.unwrap()["duration"],
|
||||||
|
expected
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_video_compatibility_validates_requests_and_maps_image_input() {
|
||||||
|
for invalid in [
|
||||||
|
json!({}),
|
||||||
|
json!({"prompt":"cat","seconds":"1.5"}),
|
||||||
|
json!({"prompt":"cat","size":"foo"}),
|
||||||
|
json!({"prompt":"cat","input_reference":{"file_id":"file-1"}}),
|
||||||
|
json!({"prompt":"cat","image":"https://example.com/a.png","reference_images":["https://example.com/b.png"]}),
|
||||||
|
json!({"prompt":"cat","reference_images":vec!["https://example.com/a.png";8]}),
|
||||||
|
] {
|
||||||
|
assert!(convert_openai_video_request(&invalid).is_err(), "{invalid}");
|
||||||
|
}
|
||||||
|
let body = convert_openai_video_request(&json!({"prompt":"cat","input_reference":{"image_url":"https://example.com/a.png"},"aspect_ratio":"square","resolution":"480p"})).unwrap();
|
||||||
|
assert_eq!(body["image"]["url"], "https://example.com/a.png");
|
||||||
|
assert_eq!(body["aspect_ratio"], "1:1");
|
||||||
|
assert_eq!(body["resolution"], "480p");
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,5 +1,8 @@
|
|||||||
use std::path::PathBuf;
|
use std::path::PathBuf;
|
||||||
use std::process::{Child, Command, Stdio};
|
use std::process::{Child, Command, Stdio};
|
||||||
|
use std::sync::atomic::{AtomicU64, Ordering};
|
||||||
|
|
||||||
|
static POSTGRES_WORKDIR_SEQ: AtomicU64 = AtomicU64::new(0);
|
||||||
|
|
||||||
use aether_data::driver::postgres::PostgresPoolConfig;
|
use aether_data::driver::postgres::PostgresPoolConfig;
|
||||||
use aether_data::{DataBackends, DataLayerConfig};
|
use aether_data::{DataBackends, DataLayerConfig};
|
||||||
@@ -21,10 +24,20 @@ pub struct ManagedPostgresServer {
|
|||||||
impl ManagedPostgresServer {
|
impl ManagedPostgresServer {
|
||||||
pub async fn start() -> Result<Self, Box<dyn std::error::Error>> {
|
pub async fn start() -> Result<Self, Box<dyn std::error::Error>> {
|
||||||
let port = reserve_local_port()?;
|
let port = reserve_local_port()?;
|
||||||
|
// pid+port is not unique: cargo test shares one PID, and ephemeral ports
|
||||||
|
// are reused after the listener is dropped. Parallel e2e tests then hit
|
||||||
|
// create_dir AlreadyExists.
|
||||||
|
let seq = POSTGRES_WORKDIR_SEQ.fetch_add(1, Ordering::Relaxed);
|
||||||
|
let nanos = std::time::SystemTime::now()
|
||||||
|
.duration_since(std::time::UNIX_EPOCH)
|
||||||
|
.map(|duration| duration.as_nanos())
|
||||||
|
.unwrap_or(0);
|
||||||
let workdir = std::env::temp_dir().join(format!(
|
let workdir = std::env::temp_dir().join(format!(
|
||||||
"aether-postgres-baseline-{}-{}",
|
"aether-postgres-baseline-{}-{}-{}-{}",
|
||||||
std::process::id(),
|
std::process::id(),
|
||||||
port
|
port,
|
||||||
|
seq,
|
||||||
|
nanos
|
||||||
));
|
));
|
||||||
let data_dir = workdir.join("data");
|
let data_dir = workdir.join("data");
|
||||||
std::fs::create_dir(&workdir)?;
|
std::fs::create_dir(&workdir)?;
|
||||||
|
|||||||
@@ -6649,7 +6649,7 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
use std::collections::BTreeMap;
|
use std::collections::BTreeMap;
|
||||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
|
||||||
use std::sync::{Arc, Mutex};
|
use std::sync::{Arc, Mutex};
|
||||||
use std::time::Instant;
|
use std::time::Instant;
|
||||||
|
|
||||||
@@ -7375,9 +7375,27 @@ mod tests {
|
|||||||
queue: Arc<dyn RuntimeQueueStore>,
|
queue: Arc<dyn RuntimeQueueStore>,
|
||||||
policy_started: Arc<tokio::sync::Notify>,
|
policy_started: Arc<tokio::sync::Notify>,
|
||||||
release_policy: Arc<tokio::sync::Notify>,
|
release_policy: Arc<tokio::sync::Notify>,
|
||||||
|
policy_released: Arc<AtomicBool>,
|
||||||
policy_reads: Arc<AtomicUsize>,
|
policy_reads: Arc<AtomicUsize>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl BlockingPolicyQueueConfiguredUsageStore {
|
||||||
|
fn new(queue: Arc<dyn RuntimeQueueStore>) -> Self {
|
||||||
|
Self {
|
||||||
|
queue,
|
||||||
|
policy_started: Arc::new(tokio::sync::Notify::new()),
|
||||||
|
release_policy: Arc::new(tokio::sync::Notify::new()),
|
||||||
|
policy_released: Arc::new(AtomicBool::new(false)),
|
||||||
|
policy_reads: Arc::new(AtomicUsize::new(0)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn release_blocked_policy(&self) {
|
||||||
|
self.policy_released.store(true, Ordering::Release);
|
||||||
|
self.release_policy.notify_waiters();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Default)]
|
#[derive(Default)]
|
||||||
struct FailingPolicyUsageStore {
|
struct FailingPolicyUsageStore {
|
||||||
inner: NoRedisUsageStore,
|
inner: NoRedisUsageStore,
|
||||||
@@ -8409,7 +8427,18 @@ mod tests {
|
|||||||
async fn body_capture_policy(&self) -> Result<UsageBodyCapturePolicy, DataLayerError> {
|
async fn body_capture_policy(&self) -> Result<UsageBodyCapturePolicy, DataLayerError> {
|
||||||
self.policy_reads.fetch_add(1, Ordering::AcqRel);
|
self.policy_reads.fetch_add(1, Ordering::AcqRel);
|
||||||
self.policy_started.notify_one();
|
self.policy_started.notify_one();
|
||||||
self.release_policy.notified().await;
|
// Latch the gate: Notify is edge-triggered, and later policy reads
|
||||||
|
// (or a waiter that subscribed after a single notify) must not hang.
|
||||||
|
loop {
|
||||||
|
if self.policy_released.load(Ordering::Acquire) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
let notified = self.release_policy.notified();
|
||||||
|
if self.policy_released.load(Ordering::Acquire) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
notified.await;
|
||||||
|
}
|
||||||
Ok(UsageBodyCapturePolicy::default())
|
Ok(UsageBodyCapturePolicy::default())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -9043,15 +9072,39 @@ mod tests {
|
|||||||
.await
|
.await
|
||||||
.expect("a duplicate first-byte marker must release the terminal barrier");
|
.expect("a duplicate first-byte marker must release the terminal barrier");
|
||||||
|
|
||||||
let records = store.records.lock().expect("records lock");
|
{
|
||||||
assert_eq!(
|
let records = store.records.lock().expect("records lock");
|
||||||
records.len(),
|
assert_eq!(
|
||||||
2,
|
records.len(),
|
||||||
"the duplicate first byte must be coalesced"
|
2,
|
||||||
);
|
"the duplicate first byte must be coalesced"
|
||||||
assert_eq!(records[0].status, "streaming");
|
);
|
||||||
assert_eq!(records[1].status, "completed");
|
assert_eq!(records[0].status, "streaming");
|
||||||
drop(records);
|
assert_eq!(records[1].status, "completed");
|
||||||
|
}
|
||||||
|
|
||||||
|
// The terminal persistence notification can arrive before the submission
|
||||||
|
// dispatcher accounts for its completed task and releases admission.
|
||||||
|
timeout(Duration::from_secs(1), async {
|
||||||
|
loop {
|
||||||
|
let snapshot = runtime.metrics_snapshot();
|
||||||
|
if snapshot.lifecycle_submission_pending == 0
|
||||||
|
&& snapshot.first_byte_persistence_pending == 0
|
||||||
|
&& snapshot.ordered_lifecycle_pending == 0
|
||||||
|
&& runtime
|
||||||
|
.lifecycle_submission
|
||||||
|
.state
|
||||||
|
.admission
|
||||||
|
.available_permits()
|
||||||
|
== CAPACITY
|
||||||
|
{
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
sleep(Duration::from_millis(1)).await;
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.expect("duplicate first-byte submission accounting should drain");
|
||||||
|
|
||||||
let snapshot = runtime.metrics_snapshot();
|
let snapshot = runtime.metrics_snapshot();
|
||||||
assert_eq!(snapshot.lifecycle_submission_pending, 0);
|
assert_eq!(snapshot.lifecycle_submission_pending, 0);
|
||||||
@@ -12325,12 +12378,9 @@ mod tests {
|
|||||||
async fn event_capture_budget_bounds_blocked_policy_waiters_and_releases_on_cancel_or_basic() {
|
async fn event_capture_budget_bounds_blocked_policy_waiters_and_releases_on_cancel_or_basic() {
|
||||||
for limit in [0, 64 * 1024] {
|
for limit in [0, 64 * 1024] {
|
||||||
let runtime = UsageRuntime::new(UsageRuntimeConfig::default()).expect("runtime");
|
let runtime = UsageRuntime::new(UsageRuntimeConfig::default()).expect("runtime");
|
||||||
let store = BlockingPolicyQueueConfiguredUsageStore {
|
let store = BlockingPolicyQueueConfiguredUsageStore::new(Arc::new(
|
||||||
queue: Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())),
|
RuntimeState::memory(MemoryRuntimeStateConfig::default()),
|
||||||
policy_started: Arc::new(tokio::sync::Notify::new()),
|
));
|
||||||
release_policy: Arc::new(tokio::sync::Notify::new()),
|
|
||||||
policy_reads: Arc::new(AtomicUsize::new(0)),
|
|
||||||
};
|
|
||||||
let budget = Arc::new(crate::event_capture_budget::EventCaptureMemoryBudget::new(
|
let budget = Arc::new(crate::event_capture_budget::EventCaptureMemoryBudget::new(
|
||||||
limit,
|
limit,
|
||||||
));
|
));
|
||||||
@@ -12391,7 +12441,7 @@ mod tests {
|
|||||||
.await
|
.await
|
||||||
.expect("replacement policy read starts");
|
.expect("replacement policy read starts");
|
||||||
assert_eq!(budget.retained_bytes(), retained);
|
assert_eq!(budget.retained_bytes(), retained);
|
||||||
store.release_policy.notify_one();
|
store.release_blocked_policy();
|
||||||
let event = timeout(Duration::from_secs(2), completing)
|
let event = timeout(Duration::from_secs(2), completing)
|
||||||
.await
|
.await
|
||||||
.expect("Basic policy completes")
|
.expect("Basic policy completes")
|
||||||
@@ -13398,12 +13448,7 @@ mod tests {
|
|||||||
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
|
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
|
||||||
let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0));
|
let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0));
|
||||||
let queue: Arc<dyn RuntimeQueueStore> = tracked_queue.clone();
|
let queue: Arc<dyn RuntimeQueueStore> = tracked_queue.clone();
|
||||||
let store = BlockingPolicyQueueConfiguredUsageStore {
|
let store = BlockingPolicyQueueConfiguredUsageStore::new(queue);
|
||||||
queue,
|
|
||||||
policy_started: Arc::new(tokio::sync::Notify::new()),
|
|
||||||
release_policy: Arc::new(tokio::sync::Notify::new()),
|
|
||||||
policy_reads: Arc::new(AtomicUsize::new(0)),
|
|
||||||
};
|
|
||||||
let runtime = UsageRuntime::new(config).expect("usage runtime should build");
|
let runtime = UsageRuntime::new(config).expect("usage runtime should build");
|
||||||
let request_id = "req-terminal-seed-waits-for-turn";
|
let request_id = "req-terminal-seed-waits-for-turn";
|
||||||
let plan = terminal_test_plan(request_id);
|
let plan = terminal_test_plan(request_id);
|
||||||
@@ -13425,9 +13470,10 @@ mod tests {
|
|||||||
assert_eq!(blocked_snapshot.terminal_submission_in_flight, 0);
|
assert_eq!(blocked_snapshot.terminal_submission_in_flight, 0);
|
||||||
assert!(blocked_snapshot.lifecycle_submission_pending >= 2);
|
assert!(blocked_snapshot.lifecycle_submission_pending >= 2);
|
||||||
|
|
||||||
store.release_policy.notify_waiters();
|
store.release_blocked_policy();
|
||||||
timeout(Duration::from_secs(2), async {
|
timeout(Duration::from_secs(2), async {
|
||||||
loop {
|
loop {
|
||||||
|
store.release_blocked_policy();
|
||||||
let snapshot = runtime.metrics_snapshot();
|
let snapshot = runtime.metrics_snapshot();
|
||||||
if tracked_queue.successful_appends.load(Ordering::Acquire) == 1
|
if tracked_queue.successful_appends.load(Ordering::Acquire) == 1
|
||||||
&& snapshot.lifecycle_submission_pending == 0
|
&& snapshot.lifecycle_submission_pending == 0
|
||||||
@@ -13435,7 +13481,7 @@ mod tests {
|
|||||||
{
|
{
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
tokio::task::yield_now().await;
|
sleep(Duration::from_millis(1)).await;
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
.await
|
.await
|
||||||
@@ -13468,12 +13514,7 @@ mod tests {
|
|||||||
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
|
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
|
||||||
let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0));
|
let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0));
|
||||||
let queue: Arc<dyn RuntimeQueueStore> = tracked_queue.clone();
|
let queue: Arc<dyn RuntimeQueueStore> = tracked_queue.clone();
|
||||||
let store = BlockingPolicyQueueConfiguredUsageStore {
|
let store = BlockingPolicyQueueConfiguredUsageStore::new(queue);
|
||||||
queue,
|
|
||||||
policy_started: Arc::new(tokio::sync::Notify::new()),
|
|
||||||
release_policy: Arc::new(tokio::sync::Notify::new()),
|
|
||||||
policy_reads: Arc::new(AtomicUsize::new(0)),
|
|
||||||
};
|
|
||||||
let runtime = UsageRuntime::new(config).expect("usage runtime should build");
|
let runtime = UsageRuntime::new(config).expect("usage runtime should build");
|
||||||
let policy_started = store.policy_started.notified();
|
let policy_started = store.policy_started.notified();
|
||||||
|
|
||||||
@@ -13520,9 +13561,10 @@ mod tests {
|
|||||||
assert_eq!(blocked_snapshot.terminal_submission_in_flight, 1);
|
assert_eq!(blocked_snapshot.terminal_submission_in_flight, 1);
|
||||||
assert!(blocked_snapshot.lifecycle_submission_pending <= BACKLOG + 1);
|
assert!(blocked_snapshot.lifecycle_submission_pending <= BACKLOG + 1);
|
||||||
|
|
||||||
store.release_policy.notify_waiters();
|
store.release_blocked_policy();
|
||||||
timeout(Duration::from_secs(5), async {
|
timeout(Duration::from_secs(5), async {
|
||||||
loop {
|
loop {
|
||||||
|
store.release_blocked_policy();
|
||||||
let snapshot = runtime.metrics_snapshot();
|
let snapshot = runtime.metrics_snapshot();
|
||||||
if tracked_queue.successful_appends.load(Ordering::Acquire) == BACKLOG + 1
|
if tracked_queue.successful_appends.load(Ordering::Acquire) == BACKLOG + 1
|
||||||
&& snapshot.lifecycle_submission_pending == 0
|
&& snapshot.lifecycle_submission_pending == 0
|
||||||
@@ -13530,7 +13572,7 @@ mod tests {
|
|||||||
{
|
{
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
tokio::task::yield_now().await;
|
sleep(Duration::from_millis(1)).await;
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
.await
|
.await
|
||||||
@@ -13828,12 +13870,7 @@ mod tests {
|
|||||||
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
|
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
|
||||||
let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0));
|
let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0));
|
||||||
let queue: Arc<dyn RuntimeQueueStore> = tracked_queue.clone();
|
let queue: Arc<dyn RuntimeQueueStore> = tracked_queue.clone();
|
||||||
let store = BlockingPolicyQueueConfiguredUsageStore {
|
let store = BlockingPolicyQueueConfiguredUsageStore::new(queue);
|
||||||
queue,
|
|
||||||
policy_started: Arc::new(tokio::sync::Notify::new()),
|
|
||||||
release_policy: Arc::new(tokio::sync::Notify::new()),
|
|
||||||
policy_reads: Arc::new(AtomicUsize::new(0)),
|
|
||||||
};
|
|
||||||
let runtime = UsageRuntime::new(config).expect("usage runtime should build");
|
let runtime = UsageRuntime::new(config).expect("usage runtime should build");
|
||||||
let policy_started = store.policy_started.notified();
|
let policy_started = store.policy_started.notified();
|
||||||
runtime
|
runtime
|
||||||
@@ -13892,10 +13929,10 @@ mod tests {
|
|||||||
.expect("terminal submissions should reach the execution backlog");
|
.expect("terminal submissions should reach the execution backlog");
|
||||||
let saturated_snapshot = runtime.metrics_snapshot();
|
let saturated_snapshot = runtime.metrics_snapshot();
|
||||||
|
|
||||||
store.release_policy.notify_waiters();
|
store.release_blocked_policy();
|
||||||
let all_completed = timeout(Duration::from_secs(2), async {
|
let all_completed = timeout(Duration::from_secs(2), async {
|
||||||
loop {
|
loop {
|
||||||
store.release_policy.notify_waiters();
|
store.release_blocked_policy();
|
||||||
if tracked_queue.successful_appends.load(Ordering::Acquire)
|
if tracked_queue.successful_appends.load(Ordering::Acquire)
|
||||||
== EXCESS_SUBMISSIONS + 1
|
== EXCESS_SUBMISSIONS + 1
|
||||||
&& runtime.metrics_snapshot().terminal_submission_in_flight == 0
|
&& runtime.metrics_snapshot().terminal_submission_in_flight == 0
|
||||||
|
|||||||
@@ -12,5 +12,6 @@ aether-data-contracts.workspace = true
|
|||||||
async-trait.workspace = true
|
async-trait.workspace = true
|
||||||
serde.workspace = true
|
serde.workspace = true
|
||||||
serde_json.workspace = true
|
serde_json.workspace = true
|
||||||
|
sha2.workspace = true
|
||||||
url.workspace = true
|
url.workspace = true
|
||||||
uuid.workspace = true
|
uuid.workspace = true
|
||||||
|
|||||||
@@ -43,6 +43,27 @@ pub fn map_openai_stored_task_to_read_response(
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn build_openai_stored_task_body(task: StoredVideoTask, status: VideoTaskStatus) -> Value {
|
fn build_openai_stored_task_body(task: StoredVideoTask, status: VideoTaskStatus) -> Value {
|
||||||
|
if task.client_api_format.as_deref() == Some("xai:video") {
|
||||||
|
let mut body = json!({"status":match status {
|
||||||
|
VideoTaskStatus::Completed => "done",
|
||||||
|
VideoTaskStatus::Expired => "expired",
|
||||||
|
VideoTaskStatus::Failed | VideoTaskStatus::Cancelled | VideoTaskStatus::Deleted => "failed",
|
||||||
|
_ => "pending",
|
||||||
|
}});
|
||||||
|
if let Some(model) = task.model {
|
||||||
|
body["model"] = json!(model);
|
||||||
|
}
|
||||||
|
if let Some(url) = task.video_url {
|
||||||
|
body["video"] = json!({"url":url});
|
||||||
|
if let Some(duration) = task.duration_seconds {
|
||||||
|
body["video"]["duration"] = json!(duration);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if status == VideoTaskStatus::Failed {
|
||||||
|
body["error"] = json!({"code":sanitize_video_task_error_code(task.error_code).unwrap_or_else(|| "unknown".into()),"message":"Video generation failed"});
|
||||||
|
}
|
||||||
|
return body;
|
||||||
|
}
|
||||||
let mut body = json!({
|
let mut body = json!({
|
||||||
"id": task.id,
|
"id": task.id,
|
||||||
"object": "video",
|
"object": "video",
|
||||||
@@ -57,6 +78,9 @@ fn build_openai_stored_task_body(task: StoredVideoTask, status: VideoTaskStatus)
|
|||||||
if let Some(prompt) = task.prompt {
|
if let Some(prompt) = task.prompt {
|
||||||
body["prompt"] = Value::String(prompt);
|
body["prompt"] = Value::String(prompt);
|
||||||
}
|
}
|
||||||
|
if let Some(seconds) = task.duration_seconds {
|
||||||
|
body["seconds"] = json!(seconds.to_string());
|
||||||
|
}
|
||||||
if let Some(size) = task.size {
|
if let Some(size) = task.size {
|
||||||
body["size"] = Value::String(size);
|
body["size"] = Value::String(size);
|
||||||
}
|
}
|
||||||
@@ -91,21 +115,97 @@ fn map_openai_stored_task_status(status: VideoTaskStatus) -> &'static str {
|
|||||||
}
|
}
|
||||||
|
|
||||||
impl OpenAiVideoTaskSeed {
|
impl OpenAiVideoTaskSeed {
|
||||||
|
pub fn uses_xai_provider(&self) -> bool {
|
||||||
|
self.xai_provider || self.is_xai_native()
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn is_xai_native(&self) -> bool {
|
||||||
|
self.persistence.client_api_format == "xai:video"
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn native_create_body_json(&self) -> Value {
|
||||||
|
let mut body = self.native_response.clone().unwrap_or_else(|| json!({}));
|
||||||
|
body["request_id"] = json!(self.local_task_id);
|
||||||
|
if body.get("id").is_some() {
|
||||||
|
body["id"] = json!(self.local_task_id);
|
||||||
|
}
|
||||||
|
body
|
||||||
|
}
|
||||||
|
|
||||||
|
fn native_read_body_json(&self) -> Value {
|
||||||
|
if let Some(mut body) = self.native_response.clone().filter(|body| {
|
||||||
|
body.get("status").is_some()
|
||||||
|
|| body.get("error").is_some()
|
||||||
|
|| body.get("code").is_some()
|
||||||
|
}) {
|
||||||
|
if body.get("request_id").is_some() {
|
||||||
|
body["request_id"] = json!(self.local_task_id);
|
||||||
|
}
|
||||||
|
if body.get("id").is_some() {
|
||||||
|
body["id"] = json!(self.local_task_id);
|
||||||
|
}
|
||||||
|
return body;
|
||||||
|
}
|
||||||
|
let mut body = json!({"status":match self.status {
|
||||||
|
LocalVideoTaskStatus::Completed => "done",
|
||||||
|
LocalVideoTaskStatus::Expired => "expired",
|
||||||
|
LocalVideoTaskStatus::Failed | LocalVideoTaskStatus::Cancelled | LocalVideoTaskStatus::Deleted => "failed",
|
||||||
|
_ => "pending",
|
||||||
|
}});
|
||||||
|
if let Some(model) = &self.model {
|
||||||
|
body["model"] = json!(model);
|
||||||
|
}
|
||||||
|
if let Some(url) = &self.video_url {
|
||||||
|
body["video"] = json!({"url":url});
|
||||||
|
if let Some(duration) = self.seconds.as_deref().and_then(|v| v.parse::<u64>().ok()) {
|
||||||
|
body["video"]["duration"] = json!(duration);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if self.error_code.is_some() {
|
||||||
|
body["error"] = json!({"code":self.error_code,"message":"Video generation failed"});
|
||||||
|
}
|
||||||
|
body
|
||||||
|
}
|
||||||
|
|
||||||
pub fn apply_provider_body(&mut self, provider_body: &Map<String, Value>) {
|
pub fn apply_provider_body(&mut self, provider_body: &Map<String, Value>) {
|
||||||
|
if self.uses_xai_provider() {
|
||||||
|
self.native_response = Some(Value::Object(provider_body.clone()));
|
||||||
|
}
|
||||||
|
|
||||||
let raw_status = provider_body
|
let raw_status = provider_body
|
||||||
.get("status")
|
.get("status")
|
||||||
.and_then(Value::as_str)
|
.and_then(Value::as_str)
|
||||||
.map(str::trim)
|
.map(str::trim)
|
||||||
.unwrap_or_default();
|
.unwrap_or_default();
|
||||||
self.status = match raw_status {
|
// Accept xAI's native lifecycle vocabulary alongside OpenAI's fields.
|
||||||
"queued" => LocalVideoTaskStatus::Queued,
|
self.status = match raw_status.to_ascii_lowercase().as_str() {
|
||||||
"processing" => LocalVideoTaskStatus::Processing,
|
"queued" | "pending" => LocalVideoTaskStatus::Queued,
|
||||||
"completed" => LocalVideoTaskStatus::Completed,
|
"processing" | "in_progress" | "running" => LocalVideoTaskStatus::Processing,
|
||||||
"failed" => LocalVideoTaskStatus::Failed,
|
"completed" | "done" | "succeeded" | "success" => LocalVideoTaskStatus::Completed,
|
||||||
"cancelled" => LocalVideoTaskStatus::Cancelled,
|
"failed" | "error" => LocalVideoTaskStatus::Failed,
|
||||||
|
"cancelled" | "canceled" => LocalVideoTaskStatus::Cancelled,
|
||||||
"expired" => LocalVideoTaskStatus::Expired,
|
"expired" => LocalVideoTaskStatus::Expired,
|
||||||
_ => LocalVideoTaskStatus::Submitted,
|
_ => LocalVideoTaskStatus::Submitted,
|
||||||
};
|
};
|
||||||
|
let error = provider_body.get("error").filter(|value| !value.is_null());
|
||||||
|
let error_code = provider_body
|
||||||
|
.get("code")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.filter(|value| !value.trim().is_empty())
|
||||||
|
.or_else(|| {
|
||||||
|
error
|
||||||
|
.and_then(|value| value.get("code"))
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
});
|
||||||
|
// xAI may report a failed job as a 200 response with code/error only.
|
||||||
|
if (error.is_some() || error_code.is_some())
|
||||||
|
&& !matches!(
|
||||||
|
self.status,
|
||||||
|
LocalVideoTaskStatus::Cancelled | LocalVideoTaskStatus::Expired
|
||||||
|
)
|
||||||
|
{
|
||||||
|
self.status = LocalVideoTaskStatus::Failed;
|
||||||
|
}
|
||||||
self.progress_percent = provider_body
|
self.progress_percent = provider_body
|
||||||
.get("progress")
|
.get("progress")
|
||||||
.and_then(Value::as_u64)
|
.and_then(Value::as_u64)
|
||||||
@@ -117,20 +217,35 @@ impl OpenAiVideoTaskSeed {
|
|||||||
});
|
});
|
||||||
self.completed_at_unix_secs = provider_body.get("completed_at").and_then(Value::as_u64);
|
self.completed_at_unix_secs = provider_body.get("completed_at").and_then(Value::as_u64);
|
||||||
self.expires_at_unix_secs = provider_body.get("expires_at").and_then(Value::as_u64);
|
self.expires_at_unix_secs = provider_body.get("expires_at").and_then(Value::as_u64);
|
||||||
let error = provider_body.get("error").and_then(Value::as_object);
|
self.error_code = sanitize_video_task_error_code(error_code.map(str::to_string));
|
||||||
self.error_code = sanitize_video_task_error_code(
|
|
||||||
error
|
|
||||||
.and_then(|value| value.get("code"))
|
|
||||||
.and_then(Value::as_str)
|
|
||||||
.map(str::to_string),
|
|
||||||
);
|
|
||||||
self.error_message = None;
|
self.error_message = None;
|
||||||
self.video_url = provider_body
|
self.video_url = provider_body
|
||||||
.get("video_url")
|
.get("video_url")
|
||||||
.or_else(|| provider_body.get("url"))
|
.or_else(|| provider_body.get("url"))
|
||||||
.or_else(|| provider_body.get("result_url"))
|
.or_else(|| provider_body.get("result_url"))
|
||||||
|
.or_else(|| {
|
||||||
|
provider_body
|
||||||
|
.get("video")
|
||||||
|
.and_then(|video| video.get("url"))
|
||||||
|
})
|
||||||
.and_then(Value::as_str)
|
.and_then(Value::as_str)
|
||||||
.map(str::to_string);
|
.map(str::to_string);
|
||||||
|
if let Some(seconds) = provider_body
|
||||||
|
.get("seconds")
|
||||||
|
.or_else(|| {
|
||||||
|
provider_body
|
||||||
|
.get("video")
|
||||||
|
.and_then(|video| video.get("duration"))
|
||||||
|
})
|
||||||
|
.filter(|value| value.is_string() || value.is_number())
|
||||||
|
{
|
||||||
|
self.seconds = Some(
|
||||||
|
seconds
|
||||||
|
.as_str()
|
||||||
|
.map(str::to_string)
|
||||||
|
.unwrap_or_else(|| seconds.to_string()),
|
||||||
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn build_content_stream_action(
|
pub fn build_content_stream_action(
|
||||||
@@ -239,6 +354,9 @@ impl OpenAiVideoTaskSeed {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn client_body_json(&self) -> Value {
|
pub fn client_body_json(&self) -> Value {
|
||||||
|
if self.is_xai_native() {
|
||||||
|
return self.native_read_body_json();
|
||||||
|
}
|
||||||
let mut body = json!({
|
let mut body = json!({
|
||||||
"id": self.local_task_id,
|
"id": self.local_task_id,
|
||||||
"object": "video",
|
"object": "video",
|
||||||
@@ -259,6 +377,9 @@ impl OpenAiVideoTaskSeed {
|
|||||||
if let Some(seconds) = &self.seconds {
|
if let Some(seconds) = &self.seconds {
|
||||||
body["seconds"] = Value::String(seconds.clone());
|
body["seconds"] = Value::String(seconds.clone());
|
||||||
}
|
}
|
||||||
|
if let Some(video_url) = &self.video_url {
|
||||||
|
body["video_url"] = Value::String(video_url.clone());
|
||||||
|
}
|
||||||
if let Some(remixed_from_video_id) = &self.remixed_from_video_id {
|
if let Some(remixed_from_video_id) = &self.remixed_from_video_id {
|
||||||
body["remixed_from_video_id"] = Value::String(remixed_from_video_id.clone());
|
body["remixed_from_video_id"] = Value::String(remixed_from_video_id.clone());
|
||||||
}
|
}
|
||||||
@@ -357,12 +478,20 @@ impl OpenAiVideoTaskSeed {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn build_get_follow_up_plan(&self, trace_id: &str) -> Option<ExecutionPlan> {
|
pub fn build_get_follow_up_plan(&self, trace_id: &str) -> Option<ExecutionPlan> {
|
||||||
if !matches!(
|
let refreshable = matches!(
|
||||||
self.status,
|
self.status,
|
||||||
LocalVideoTaskStatus::Submitted
|
LocalVideoTaskStatus::Submitted
|
||||||
| LocalVideoTaskStatus::Queued
|
| LocalVideoTaskStatus::Queued
|
||||||
| LocalVideoTaskStatus::Processing
|
| LocalVideoTaskStatus::Processing
|
||||||
) {
|
) || (self.uses_xai_provider()
|
||||||
|
&& self.native_response.is_none()
|
||||||
|
&& matches!(
|
||||||
|
self.status,
|
||||||
|
LocalVideoTaskStatus::Completed
|
||||||
|
| LocalVideoTaskStatus::Failed
|
||||||
|
| LocalVideoTaskStatus::Expired
|
||||||
|
));
|
||||||
|
if !refreshable {
|
||||||
return None;
|
return None;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -573,7 +702,12 @@ impl OpenAiVideoTaskSeed {
|
|||||||
};
|
};
|
||||||
let mut record = UpsertVideoTask {
|
let mut record = UpsertVideoTask {
|
||||||
id: self.local_task_id.clone(),
|
id: self.local_task_id.clone(),
|
||||||
short_id: None,
|
// The production schema requires a unique, non-null short_id (at most 16 chars).
|
||||||
|
// Derive it deterministically so repeated capture and legacy snapshot reloads agree.
|
||||||
|
short_id: Some(self.local_short_id.clone().unwrap_or_else(|| {
|
||||||
|
use sha2::{Digest, Sha256};
|
||||||
|
format!("{:x}", Sha256::digest(self.local_task_id.as_bytes()))[..16].to_string()
|
||||||
|
})),
|
||||||
request_id: self.persistence.request_id.clone(),
|
request_id: self.persistence.request_id.clone(),
|
||||||
user_id: self.user_id.clone(),
|
user_id: self.user_id.clone(),
|
||||||
api_key_id: self.api_key_id.clone(),
|
api_key_id: self.api_key_id.clone(),
|
||||||
@@ -589,7 +723,11 @@ impl OpenAiVideoTaskSeed {
|
|||||||
model: self.model.clone().or_else(|| Some(String::new())),
|
model: self.model.clone().or_else(|| Some(String::new())),
|
||||||
prompt: self.prompt.clone().or_else(|| Some(String::new())),
|
prompt: self.prompt.clone().or_else(|| Some(String::new())),
|
||||||
original_request_body: None,
|
original_request_body: None,
|
||||||
duration_seconds: request_body_u32(&self.persistence.original_request_body, "seconds"),
|
duration_seconds: self
|
||||||
|
.seconds
|
||||||
|
.as_deref()
|
||||||
|
.and_then(|value| value.parse().ok())
|
||||||
|
.or_else(|| request_body_u32(&self.persistence.original_request_body, "seconds")),
|
||||||
resolution: request_body_string(&self.persistence.original_request_body, "resolution"),
|
resolution: request_body_string(&self.persistence.original_request_body, "resolution"),
|
||||||
aspect_ratio: request_body_string(
|
aspect_ratio: request_body_string(
|
||||||
&self.persistence.original_request_body,
|
&self.persistence.original_request_body,
|
||||||
@@ -697,6 +835,9 @@ mod tests {
|
|||||||
#[test]
|
#[test]
|
||||||
fn builds_minimal_openai_persistence_record_without_sensitive_snapshot() {
|
fn builds_minimal_openai_persistence_record_without_sensitive_snapshot() {
|
||||||
let seed = OpenAiVideoTaskSeed {
|
let seed = OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: None,
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: false,
|
||||||
local_task_id: "task-openai-sensitive".to_string(),
|
local_task_id: "task-openai-sensitive".to_string(),
|
||||||
upstream_task_id: "upstream-openai-sensitive".to_string(),
|
upstream_task_id: "upstream-openai-sensitive".to_string(),
|
||||||
created_at_unix_ms: 1_712_345_678,
|
created_at_unix_ms: 1_712_345_678,
|
||||||
@@ -747,6 +888,12 @@ mod tests {
|
|||||||
|
|
||||||
let record = seed.to_upsert_record();
|
let record = seed.to_upsert_record();
|
||||||
|
|
||||||
|
let short_id = record
|
||||||
|
.short_id
|
||||||
|
.as_deref()
|
||||||
|
.expect("database short_id is required");
|
||||||
|
assert_eq!(short_id.len(), 16);
|
||||||
|
assert_eq!(seed.to_upsert_record().short_id, record.short_id);
|
||||||
assert_eq!(record.error_code.as_deref(), Some("provider_error"));
|
assert_eq!(record.error_code.as_deref(), Some("provider_error"));
|
||||||
assert!(record.original_request_body.is_none());
|
assert!(record.original_request_body.is_none());
|
||||||
assert!(record.progress_message.is_none());
|
assert!(record.progress_message.is_none());
|
||||||
@@ -759,6 +906,8 @@ mod tests {
|
|||||||
|
|
||||||
let mut stored = record.into_stored();
|
let mut stored = record.into_stored();
|
||||||
stored.status = VideoTaskStatus::Completed;
|
stored.status = VideoTaskStatus::Completed;
|
||||||
|
// Migrated tasks can already have a short ID unrelated to the derived ID.
|
||||||
|
stored.short_id = Some("legacy-short-id".to_string());
|
||||||
let snapshot =
|
let snapshot =
|
||||||
LocalVideoTaskSnapshot::from_stored_task_with_transport(&stored, seed.transport)
|
LocalVideoTaskSnapshot::from_stored_task_with_transport(&stored, seed.transport)
|
||||||
.expect("stored task should reconstruct with current transport");
|
.expect("stored task should reconstruct with current transport");
|
||||||
@@ -766,6 +915,21 @@ mod tests {
|
|||||||
panic!("expected OpenAI snapshot");
|
panic!("expected OpenAI snapshot");
|
||||||
};
|
};
|
||||||
assert_eq!(restored.prompt, stored.prompt);
|
assert_eq!(restored.prompt, stored.prompt);
|
||||||
|
assert_eq!(restored.to_upsert_record().short_id, stored.short_id);
|
||||||
|
let mut embedded = stored.clone();
|
||||||
|
let mut legacy_snapshot =
|
||||||
|
serde_json::to_value(LocalVideoTaskSnapshot::OpenAi(restored.clone())).unwrap();
|
||||||
|
legacy_snapshot["OpenAi"]
|
||||||
|
.as_object_mut()
|
||||||
|
.unwrap()
|
||||||
|
.remove("local_short_id");
|
||||||
|
embedded.request_metadata = Some(json!({"rust_local_snapshot": legacy_snapshot}));
|
||||||
|
let embedded_snapshot = LocalVideoTaskSnapshot::from_stored_task(&embedded)
|
||||||
|
.expect("legacy embedded snapshot should hydrate");
|
||||||
|
assert_eq!(
|
||||||
|
embedded_snapshot.to_upsert_record().short_id,
|
||||||
|
stored.short_id
|
||||||
|
);
|
||||||
assert_eq!(restored.to_upsert_record().video_url, stored.video_url);
|
assert_eq!(restored.to_upsert_record().video_url, stored.video_url);
|
||||||
let Some(LocalVideoTaskContentAction::StreamPlan(plan)) =
|
let Some(LocalVideoTaskContentAction::StreamPlan(plan)) =
|
||||||
restored.build_content_stream_action(None, "trace-download")
|
restored.build_content_stream_action(None, "trace-download")
|
||||||
|
|||||||
@@ -8,7 +8,9 @@ use uuid::Uuid;
|
|||||||
use crate::{LocalVideoTaskRegistryMutation, LocalVideoTaskStatus, VideoTaskTruthSourceMode};
|
use crate::{LocalVideoTaskRegistryMutation, LocalVideoTaskStatus, VideoTaskTruthSourceMode};
|
||||||
|
|
||||||
pub fn extract_openai_task_id_from_path(path: &str) -> Option<&str> {
|
pub fn extract_openai_task_id_from_path(path: &str) -> Option<&str> {
|
||||||
let suffix = path.strip_prefix("/v1/videos/")?;
|
let suffix = path
|
||||||
|
.strip_prefix("/v1/videos/")
|
||||||
|
.or_else(|| path.strip_prefix("/openai/v1/videos/"))?;
|
||||||
if suffix.is_empty()
|
if suffix.is_empty()
|
||||||
|| suffix.contains('/')
|
|| suffix.contains('/')
|
||||||
|| suffix.ends_with(":cancel")
|
|| suffix.ends_with(":cancel")
|
||||||
@@ -29,21 +31,27 @@ pub fn extract_gemini_short_id_from_path(path: &str) -> Option<&str> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn extract_openai_task_id_from_cancel_path(path: &str) -> Option<&str> {
|
pub fn extract_openai_task_id_from_cancel_path(path: &str) -> Option<&str> {
|
||||||
let suffix = path.strip_prefix("/v1/videos/")?;
|
let suffix = path
|
||||||
|
.strip_prefix("/v1/videos/")
|
||||||
|
.or_else(|| path.strip_prefix("/openai/v1/videos/"))?;
|
||||||
suffix
|
suffix
|
||||||
.strip_suffix("/cancel")
|
.strip_suffix("/cancel")
|
||||||
.filter(|value| !value.is_empty())
|
.filter(|value| !value.is_empty())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn extract_openai_task_id_from_remix_path(path: &str) -> Option<&str> {
|
pub fn extract_openai_task_id_from_remix_path(path: &str) -> Option<&str> {
|
||||||
let suffix = path.strip_prefix("/v1/videos/")?;
|
let suffix = path
|
||||||
|
.strip_prefix("/v1/videos/")
|
||||||
|
.or_else(|| path.strip_prefix("/openai/v1/videos/"))?;
|
||||||
suffix
|
suffix
|
||||||
.strip_suffix("/remix")
|
.strip_suffix("/remix")
|
||||||
.filter(|value| !value.is_empty())
|
.filter(|value| !value.is_empty())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn extract_openai_task_id_from_content_path(path: &str) -> Option<&str> {
|
pub fn extract_openai_task_id_from_content_path(path: &str) -> Option<&str> {
|
||||||
let suffix = path.strip_prefix("/v1/videos/")?;
|
let suffix = path
|
||||||
|
.strip_prefix("/v1/videos/")
|
||||||
|
.or_else(|| path.strip_prefix("/openai/v1/videos/"))?;
|
||||||
suffix
|
suffix
|
||||||
.strip_suffix("/content")
|
.strip_suffix("/content")
|
||||||
.filter(|value| !value.is_empty())
|
.filter(|value| !value.is_empty())
|
||||||
|
|||||||
@@ -73,7 +73,7 @@ async fn read_openai_video_task_response(
|
|||||||
}
|
}
|
||||||
None => state.find_stored_video_task(lookup).await?,
|
None => state.find_stored_video_task(lookup).await?,
|
||||||
};
|
};
|
||||||
let Some(task) = task else {
|
let Some(mut task) = task else {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -81,6 +81,9 @@ async fn read_openai_video_task_response(
|
|||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if request_path.starts_with("/openai/v1/videos/") {
|
||||||
|
task.client_api_format = Some("openai:video".into());
|
||||||
|
}
|
||||||
Ok(Some(map_openai_stored_task_to_read_response(task)))
|
Ok(Some(map_openai_stored_task_to_read_response(task)))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -105,13 +105,8 @@ impl VideoTaskService {
|
|||||||
if self.truth_source_mode != VideoTaskTruthSourceMode::RustAuthoritative {
|
if self.truth_source_mode != VideoTaskTruthSourceMode::RustAuthoritative {
|
||||||
return None;
|
return None;
|
||||||
}
|
}
|
||||||
match route_family {
|
self.snapshot_for_route(route_family, request_path)
|
||||||
Some("openai") => extract_openai_task_id_from_path(request_path)
|
.map(|snapshot| snapshot.read_response_for_path(request_path))
|
||||||
.and_then(|task_id| self.store.read_openai(task_id)),
|
|
||||||
Some("gemini") => extract_gemini_short_id_from_path(request_path)
|
|
||||||
.and_then(|short_id| self.store.read_gemini(short_id)),
|
|
||||||
_ => None,
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn read_response_for_user(
|
pub fn read_response_for_user(
|
||||||
@@ -126,7 +121,7 @@ impl VideoTaskService {
|
|||||||
let snapshot = self.snapshot_for_route(route_family, request_path)?;
|
let snapshot = self.snapshot_for_route(route_family, request_path)?;
|
||||||
snapshot
|
snapshot
|
||||||
.belongs_to_user(user_id)
|
.belongs_to_user(user_id)
|
||||||
.then(|| snapshot.read_response())
|
.then(|| snapshot.read_response_for_path(request_path))
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn snapshot_for_route(
|
pub fn snapshot_for_route(
|
||||||
|
|||||||
@@ -29,6 +29,7 @@ impl LocalVideoTaskSnapshot {
|
|||||||
// contain stale identity fields after a task import or repair.
|
// contain stale identity fields after a task import or repair.
|
||||||
match &mut snapshot {
|
match &mut snapshot {
|
||||||
Self::OpenAi(seed) => {
|
Self::OpenAi(seed) => {
|
||||||
|
seed.local_short_id = task.short_id.clone();
|
||||||
seed.user_id = task.user_id.clone();
|
seed.user_id = task.user_id.clone();
|
||||||
seed.api_key_id = task.api_key_id.clone();
|
seed.api_key_id = task.api_key_id.clone();
|
||||||
}
|
}
|
||||||
@@ -51,6 +52,9 @@ impl LocalVideoTaskSnapshot {
|
|||||||
"openai:video" => {
|
"openai:video" => {
|
||||||
let upstream_task_id = non_empty_owned(task.external_task_id.as_ref())?;
|
let upstream_task_id = non_empty_owned(task.external_task_id.as_ref())?;
|
||||||
Some(Self::OpenAi(OpenAiVideoTaskSeed {
|
Some(Self::OpenAi(OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: task.short_id.clone(),
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: persistence.client_api_format == "xai:video",
|
||||||
local_task_id: task.id.clone(),
|
local_task_id: task.id.clone(),
|
||||||
upstream_task_id,
|
upstream_task_id,
|
||||||
created_at_unix_ms: task.created_at_unix_ms,
|
created_at_unix_ms: task.created_at_unix_ms,
|
||||||
@@ -142,6 +146,19 @@ impl LocalVideoTaskSnapshot {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn read_response_for_path(&self, path: &str) -> LocalVideoTaskReadResponse {
|
||||||
|
if let Self::OpenAi(seed) = self {
|
||||||
|
let mut seed = seed.clone();
|
||||||
|
if path.starts_with("/openai/v1/videos/") {
|
||||||
|
seed.persistence.client_api_format = "openai:video".to_string();
|
||||||
|
} else if path.starts_with("/v1/videos/") && seed.uses_xai_provider() {
|
||||||
|
seed.persistence.client_api_format = "xai:video".to_string();
|
||||||
|
}
|
||||||
|
return Self::OpenAi(seed).read_response();
|
||||||
|
}
|
||||||
|
self.read_response()
|
||||||
|
}
|
||||||
|
|
||||||
pub fn read_response(&self) -> LocalVideoTaskReadResponse {
|
pub fn read_response(&self) -> LocalVideoTaskReadResponse {
|
||||||
match self {
|
match self {
|
||||||
Self::OpenAi(seed) => match seed.status {
|
Self::OpenAi(seed) => match seed.status {
|
||||||
|
|||||||
@@ -19,14 +19,17 @@ impl LocalVideoTaskSeed {
|
|||||||
) -> Option<Self> {
|
) -> Option<Self> {
|
||||||
let transport = LocalVideoTaskTransport::from_plan(plan)?;
|
let transport = LocalVideoTaskTransport::from_plan(plan)?;
|
||||||
let persistence = LocalVideoTaskPersistence::from_report_context(report_context, plan);
|
let persistence = LocalVideoTaskPersistence::from_report_context(report_context, plan);
|
||||||
match report_kind {
|
let mut seed = match report_kind {
|
||||||
"openai_video_create_sync_finalize" => {
|
"openai_video_create_sync_finalize" => {
|
||||||
let upstream_id = provider_body.get("id").and_then(Value::as_str)?.trim();
|
let upstream_id = openai_video_provider_task_id(provider_body)?;
|
||||||
if upstream_id.is_empty() {
|
|
||||||
return None;
|
|
||||||
}
|
|
||||||
|
|
||||||
Some(Self::OpenAiCreate(OpenAiVideoTaskSeed {
|
Some(Self::OpenAiCreate(OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: None,
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: report_context
|
||||||
|
.get("video_provider_xai")
|
||||||
|
.and_then(Value::as_bool)
|
||||||
|
.unwrap_or(false),
|
||||||
local_task_id: context_text(report_context, "local_task_id")
|
local_task_id: context_text(report_context, "local_task_id")
|
||||||
.unwrap_or_else(|| Uuid::new_v4().to_string()),
|
.unwrap_or_else(|| Uuid::new_v4().to_string()),
|
||||||
upstream_task_id: upstream_id.to_string(),
|
upstream_task_id: upstream_id.to_string(),
|
||||||
@@ -37,8 +40,12 @@ impl LocalVideoTaskSeed {
|
|||||||
model: context_text(report_context, "model")
|
model: context_text(report_context, "model")
|
||||||
.or_else(|| request_body_text(report_context, "model")),
|
.or_else(|| request_body_text(report_context, "model")),
|
||||||
prompt: request_body_text(report_context, "prompt"),
|
prompt: request_body_text(report_context, "prompt"),
|
||||||
size: request_body_text(report_context, "size"),
|
size: context_text(report_context, "video_size")
|
||||||
seconds: request_body_text(report_context, "seconds"),
|
.or_else(|| request_body_text(report_context, "size")),
|
||||||
|
seconds: context_u64(report_context, "video_duration")
|
||||||
|
.map(|v| v.to_string())
|
||||||
|
.or_else(|| request_body_text(report_context, "seconds"))
|
||||||
|
.or_else(|| request_body_text(report_context, "duration")),
|
||||||
remixed_from_video_id: None,
|
remixed_from_video_id: None,
|
||||||
status: LocalVideoTaskStatus::Submitted,
|
status: LocalVideoTaskStatus::Submitted,
|
||||||
progress_percent: 0,
|
progress_percent: 0,
|
||||||
@@ -52,12 +59,15 @@ impl LocalVideoTaskSeed {
|
|||||||
}))
|
}))
|
||||||
}
|
}
|
||||||
"openai_video_remix_sync_finalize" => {
|
"openai_video_remix_sync_finalize" => {
|
||||||
let upstream_id = provider_body.get("id").and_then(Value::as_str)?.trim();
|
let upstream_id = openai_video_provider_task_id(provider_body)?;
|
||||||
if upstream_id.is_empty() {
|
|
||||||
return None;
|
|
||||||
}
|
|
||||||
|
|
||||||
Some(Self::OpenAiRemix(OpenAiVideoTaskSeed {
|
Some(Self::OpenAiRemix(OpenAiVideoTaskSeed {
|
||||||
|
local_short_id: None,
|
||||||
|
native_response: None,
|
||||||
|
xai_provider: report_context
|
||||||
|
.get("video_provider_xai")
|
||||||
|
.and_then(Value::as_bool)
|
||||||
|
.unwrap_or(false),
|
||||||
local_task_id: context_text(report_context, "local_task_id")
|
local_task_id: context_text(report_context, "local_task_id")
|
||||||
.unwrap_or_else(|| Uuid::new_v4().to_string()),
|
.unwrap_or_else(|| Uuid::new_v4().to_string()),
|
||||||
upstream_task_id: upstream_id.to_string(),
|
upstream_task_id: upstream_id.to_string(),
|
||||||
@@ -68,8 +78,12 @@ impl LocalVideoTaskSeed {
|
|||||||
model: context_text(report_context, "model")
|
model: context_text(report_context, "model")
|
||||||
.or_else(|| request_body_text(report_context, "model")),
|
.or_else(|| request_body_text(report_context, "model")),
|
||||||
prompt: request_body_text(report_context, "prompt"),
|
prompt: request_body_text(report_context, "prompt"),
|
||||||
size: request_body_text(report_context, "size"),
|
size: context_text(report_context, "video_size")
|
||||||
seconds: request_body_text(report_context, "seconds"),
|
.or_else(|| request_body_text(report_context, "size")),
|
||||||
|
seconds: context_u64(report_context, "video_duration")
|
||||||
|
.map(|v| v.to_string())
|
||||||
|
.or_else(|| request_body_text(report_context, "seconds"))
|
||||||
|
.or_else(|| request_body_text(report_context, "duration")),
|
||||||
remixed_from_video_id: context_text(report_context, "task_id")
|
remixed_from_video_id: context_text(report_context, "task_id")
|
||||||
.or_else(|| request_body_text(report_context, "remix_video_id")),
|
.or_else(|| request_body_text(report_context, "remix_video_id")),
|
||||||
status: LocalVideoTaskStatus::Submitted,
|
status: LocalVideoTaskStatus::Submitted,
|
||||||
@@ -110,7 +124,11 @@ impl LocalVideoTaskSeed {
|
|||||||
}))
|
}))
|
||||||
}
|
}
|
||||||
_ => None,
|
_ => None,
|
||||||
|
}?;
|
||||||
|
if let Self::OpenAiCreate(task) | Self::OpenAiRemix(task) = &mut seed {
|
||||||
|
task.apply_provider_body(provider_body);
|
||||||
}
|
}
|
||||||
|
Some(seed)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn success_report_kind(&self) -> &'static str {
|
pub fn success_report_kind(&self) -> &'static str {
|
||||||
@@ -144,12 +162,28 @@ impl LocalVideoTaskSeed {
|
|||||||
|
|
||||||
pub fn client_body_json(&self) -> Value {
|
pub fn client_body_json(&self) -> Value {
|
||||||
match self {
|
match self {
|
||||||
Self::OpenAiCreate(seed) | Self::OpenAiRemix(seed) => seed.client_body_json(),
|
Self::OpenAiCreate(seed) | Self::OpenAiRemix(seed) => {
|
||||||
|
if seed.is_xai_native() {
|
||||||
|
seed.native_create_body_json()
|
||||||
|
} else {
|
||||||
|
seed.client_body_json()
|
||||||
|
}
|
||||||
|
}
|
||||||
Self::GeminiCreate(seed) => seed.client_body_json(),
|
Self::GeminiCreate(seed) => seed.client_body_json(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn openai_video_provider_task_id(body: &Map<String, Value>) -> Option<&str> {
|
||||||
|
// xAI's OpenAI-compatible video creation returns request_id instead of id.
|
||||||
|
["id", "request_id"].into_iter().find_map(|field| {
|
||||||
|
body.get(field)
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
impl VideoTaskTruthSourceMode {
|
impl VideoTaskTruthSourceMode {
|
||||||
pub fn prepare_sync_success(
|
pub fn prepare_sync_success(
|
||||||
self,
|
self,
|
||||||
@@ -353,6 +387,234 @@ mod tests {
|
|||||||
resolve_local_sync_success_background_report_kind,
|
resolve_local_sync_success_background_report_kind,
|
||||||
};
|
};
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_native_video_protocol_survives_persistence_and_preserves_provider_fields() {
|
||||||
|
use crate::{
|
||||||
|
LocalVideoTaskContentAction, LocalVideoTaskSnapshot, VideoTaskService,
|
||||||
|
VideoTaskTruthSourceMode,
|
||||||
|
};
|
||||||
|
let mut plan =
|
||||||
|
build_internal_finalize_video_plan("native-create", "openai:video", None).unwrap();
|
||||||
|
plan.url = "https://api.x.ai/v1/videos/generations".into();
|
||||||
|
plan.headers
|
||||||
|
.insert("authorization".into(), "Bearer test-key".into());
|
||||||
|
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
||||||
|
let context = json!({"local_task_id":"native-local", "user_id":"owner", "model":"grok-imagine-video", "video_client_protocol":"xai", "video_duration":6});
|
||||||
|
let success = service
|
||||||
|
.prepare_sync_success(
|
||||||
|
"openai_video_create_sync_finalize",
|
||||||
|
json!({"request_id":"native-upstream", "future_field":true})
|
||||||
|
.as_object()
|
||||||
|
.unwrap(),
|
||||||
|
context.as_object().unwrap(),
|
||||||
|
&plan,
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
success.client_body_json(),
|
||||||
|
json!({"request_id":"native-local","future_field":true})
|
||||||
|
);
|
||||||
|
let mut snapshot = success.to_snapshot();
|
||||||
|
let body = json!({"status":"done","video":{"url":"https://vidgen.x.ai/video.mp4","duration":6,"respect_moderation":true},"future_field":[1,2]});
|
||||||
|
snapshot.apply_provider_body(body.as_object().unwrap());
|
||||||
|
assert_eq!(snapshot.read_response().body_json, body);
|
||||||
|
assert_eq!(
|
||||||
|
snapshot
|
||||||
|
.read_response_for_path("/openai/v1/videos/native-local")
|
||||||
|
.body_json["status"],
|
||||||
|
"completed"
|
||||||
|
);
|
||||||
|
let LocalVideoTaskSnapshot::OpenAi(seed) = &snapshot else {
|
||||||
|
panic!("openai task expected")
|
||||||
|
};
|
||||||
|
let Some(LocalVideoTaskContentAction::StreamPlan(download)) =
|
||||||
|
seed.build_content_stream_action(None, "download")
|
||||||
|
else {
|
||||||
|
panic!("download expected")
|
||||||
|
};
|
||||||
|
assert_eq!(download.url, "https://vidgen.x.ai/video.mp4");
|
||||||
|
assert!(download.headers.is_empty());
|
||||||
|
let stored = snapshot.to_upsert_record().into_stored();
|
||||||
|
assert!(stored.request_metadata.is_none());
|
||||||
|
assert_eq!(stored.client_api_format.as_deref(), Some("xai:video"));
|
||||||
|
let restored = LocalVideoTaskSnapshot::from_stored_task_with_transport(
|
||||||
|
&stored,
|
||||||
|
seed.transport.clone(),
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(restored.read_response().body_json["status"], "done");
|
||||||
|
service.record_snapshot(restored);
|
||||||
|
let poll = service
|
||||||
|
.prepare_read_refresh_sync_plan_for_user(
|
||||||
|
Some("openai"),
|
||||||
|
"/v1/videos/native-local",
|
||||||
|
"owner",
|
||||||
|
"poll",
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(poll.plan.url, "https://api.x.ai/v1/videos/native-upstream");
|
||||||
|
assert!(service
|
||||||
|
.prepare_read_refresh_sync_plan_for_user(
|
||||||
|
Some("openai"),
|
||||||
|
"/v1/videos/native-local",
|
||||||
|
"foreign",
|
||||||
|
"poll"
|
||||||
|
)
|
||||||
|
.is_none());
|
||||||
|
assert!(service.apply_read_refresh_projection(&poll, body.as_object().unwrap()));
|
||||||
|
assert_eq!(
|
||||||
|
service
|
||||||
|
.read_response_for_user(Some("openai"), "/v1/videos/native-local", "owner")
|
||||||
|
.unwrap()
|
||||||
|
.body_json,
|
||||||
|
body
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_video_lifecycle_creates_polls_persists_and_downloads() {
|
||||||
|
use crate::{
|
||||||
|
LocalVideoTaskContentAction, LocalVideoTaskSnapshot, VideoTaskService,
|
||||||
|
VideoTaskTruthSourceMode,
|
||||||
|
};
|
||||||
|
for api_root in ["https://cli-chat-proxy.grok.com/v1", "https://api.x.ai/v1"] {
|
||||||
|
let mut plan =
|
||||||
|
build_internal_finalize_video_plan("xai-create", "openai:video", None).unwrap();
|
||||||
|
plan.url = format!("{api_root}/videos/generations");
|
||||||
|
plan.headers
|
||||||
|
.insert("authorization".into(), "Bearer test-token".into());
|
||||||
|
let service = VideoTaskService::new(VideoTaskTruthSourceMode::RustAuthoritative);
|
||||||
|
let context = json!({"local_task_id": "local-video", "model": "grok-imagine-video", "original_request_body": {"prompt": "A cat", "seconds": "6"}});
|
||||||
|
let success = service
|
||||||
|
.prepare_sync_success(
|
||||||
|
"openai_video_create_sync_finalize",
|
||||||
|
json!({"request_id": "xai-request"}).as_object().unwrap(),
|
||||||
|
context.as_object().unwrap(),
|
||||||
|
&plan,
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(success.client_body_json()["id"], "local-video");
|
||||||
|
assert_eq!(success.client_body_json()["status"], "queued");
|
||||||
|
let snapshot = success.to_snapshot();
|
||||||
|
assert_eq!(
|
||||||
|
snapshot.to_upsert_record().external_task_id.as_deref(),
|
||||||
|
Some("xai-request")
|
||||||
|
);
|
||||||
|
service.record_snapshot(snapshot.clone());
|
||||||
|
let poll = service
|
||||||
|
.prepare_poll_refresh_plan_for_snapshot(snapshot, "xai-poll")
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(poll.plan.method, "GET");
|
||||||
|
assert_eq!(poll.plan.url, format!("{api_root}/videos/xai-request"));
|
||||||
|
assert_eq!(
|
||||||
|
poll.plan.headers.get("authorization"),
|
||||||
|
plan.headers.get("authorization")
|
||||||
|
);
|
||||||
|
assert!(service.apply_read_refresh_projection(
|
||||||
|
&poll,
|
||||||
|
json!({"status": "pending"}).as_object().unwrap()
|
||||||
|
));
|
||||||
|
assert_eq!(
|
||||||
|
service
|
||||||
|
.read_response(Some("openai"), "/v1/videos/local-video")
|
||||||
|
.unwrap()
|
||||||
|
.body_json["status"],
|
||||||
|
"queued"
|
||||||
|
);
|
||||||
|
assert!(service.apply_read_refresh_projection(&poll, json!({
|
||||||
|
"status": "done", "video": {"url": "https://vidgen.x.ai/result.mp4", "duration": 6}
|
||||||
|
}).as_object().unwrap()));
|
||||||
|
let snapshot = service
|
||||||
|
.snapshot_for_route(Some("openai"), "/v1/videos/local-video")
|
||||||
|
.unwrap();
|
||||||
|
assert!(!snapshot.is_active_for_refresh());
|
||||||
|
let record = snapshot.to_upsert_record();
|
||||||
|
assert_eq!(
|
||||||
|
record.status,
|
||||||
|
aether_data_contracts::repository::video_tasks::VideoTaskStatus::Completed
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
record.video_url.as_deref(),
|
||||||
|
Some("https://vidgen.x.ai/result.mp4")
|
||||||
|
);
|
||||||
|
assert_eq!(record.duration_seconds, Some(6));
|
||||||
|
let response = snapshot.read_response();
|
||||||
|
assert_eq!(response.body_json["status"], "completed");
|
||||||
|
assert_eq!(response.body_json["progress"], 100);
|
||||||
|
assert_eq!(
|
||||||
|
response.body_json["video_url"],
|
||||||
|
"https://vidgen.x.ai/result.mp4"
|
||||||
|
);
|
||||||
|
let LocalVideoTaskSnapshot::OpenAi(seed) = snapshot else {
|
||||||
|
panic!("OpenAI video expected")
|
||||||
|
};
|
||||||
|
let Some(LocalVideoTaskContentAction::StreamPlan(download)) =
|
||||||
|
seed.build_content_stream_action(None, "download")
|
||||||
|
else {
|
||||||
|
panic!("download expected")
|
||||||
|
};
|
||||||
|
assert_eq!(download.url, "https://vidgen.x.ai/result.mp4");
|
||||||
|
assert!(
|
||||||
|
download.headers.is_empty(),
|
||||||
|
"provider credentials must not be sent to the media CDN"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn xai_video_errors_are_terminal_even_without_a_status() {
|
||||||
|
use crate::{LocalVideoTaskSnapshot, VideoTaskTruthSourceMode};
|
||||||
|
let mut plan =
|
||||||
|
build_internal_finalize_video_plan("xai-create", "openai:video", None).unwrap();
|
||||||
|
plan.url = "https://cli-chat-proxy.grok.com/v1/videos/generations".into();
|
||||||
|
for body in [
|
||||||
|
json!({"code": "content_policy_violation", "error": "Rejected"}),
|
||||||
|
json!({"error": {"code": "content_policy_violation", "message": "Rejected"}}),
|
||||||
|
json!({"status": "failed", "error": "Rejected"}),
|
||||||
|
] {
|
||||||
|
let mut snapshot = VideoTaskTruthSourceMode::RustAuthoritative
|
||||||
|
.prepare_sync_success(
|
||||||
|
"openai_video_create_sync_finalize",
|
||||||
|
json!({"request_id": "xai-request"}).as_object().unwrap(),
|
||||||
|
&Default::default(),
|
||||||
|
&plan,
|
||||||
|
)
|
||||||
|
.unwrap()
|
||||||
|
.to_snapshot();
|
||||||
|
snapshot.apply_provider_body(body.as_object().unwrap());
|
||||||
|
assert!(!snapshot.is_active_for_refresh());
|
||||||
|
assert_eq!(snapshot.read_response().body_json["status"], "failed");
|
||||||
|
let LocalVideoTaskSnapshot::OpenAi(seed) = snapshot else {
|
||||||
|
panic!("OpenAI video expected")
|
||||||
|
};
|
||||||
|
assert!(seed.error_message.is_none());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn openai_video_id_takes_precedence_over_xai_alias() {
|
||||||
|
assert_eq!(
|
||||||
|
super::openai_video_provider_task_id(
|
||||||
|
json!({"id": "openai-id", "request_id": "trace-id"})
|
||||||
|
.as_object()
|
||||||
|
.unwrap()
|
||||||
|
),
|
||||||
|
Some("openai-id")
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
super::openai_video_provider_task_id(
|
||||||
|
json!({"id": " ", "request_id": "xai-id"})
|
||||||
|
.as_object()
|
||||||
|
.unwrap()
|
||||||
|
),
|
||||||
|
Some("xai-id")
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
super::openai_video_provider_task_id(json!({"request_id": " "}).as_object().unwrap()),
|
||||||
|
None
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn builds_local_sync_finalize_read_response_for_supported_video_finalize_kinds() {
|
fn builds_local_sync_finalize_read_response_for_supported_video_finalize_kinds() {
|
||||||
let delete_response = build_local_sync_finalize_read_response(
|
let delete_response = build_local_sync_finalize_read_response(
|
||||||
|
|||||||
@@ -71,8 +71,16 @@ impl LocalVideoTaskPersistence {
|
|||||||
.unwrap_or_else(|| plan.request_id.clone()),
|
.unwrap_or_else(|| plan.request_id.clone()),
|
||||||
username: context_text(report_context, "username"),
|
username: context_text(report_context, "username"),
|
||||||
api_key_name: context_text(report_context, "api_key_name"),
|
api_key_name: context_text(report_context, "api_key_name"),
|
||||||
client_api_format: context_text(report_context, "client_api_format")
|
client_api_format: if report_context
|
||||||
.unwrap_or_else(|| plan.client_api_format.clone()),
|
.get("video_client_protocol")
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
== Some("xai")
|
||||||
|
{
|
||||||
|
"xai:video".to_string()
|
||||||
|
} else {
|
||||||
|
context_text(report_context, "client_api_format")
|
||||||
|
.unwrap_or_else(|| plan.client_api_format.clone())
|
||||||
|
},
|
||||||
provider_api_format: context_text(report_context, "provider_api_format")
|
provider_api_format: context_text(report_context, "provider_api_format")
|
||||||
.unwrap_or_else(|| plan.provider_api_format.clone()),
|
.unwrap_or_else(|| plan.provider_api_format.clone()),
|
||||||
original_request_body: report_context
|
original_request_body: report_context
|
||||||
|
|||||||
@@ -201,6 +201,13 @@ pub struct LocalVideoTaskPersistence {
|
|||||||
|
|
||||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||||
pub struct OpenAiVideoTaskSeed {
|
pub struct OpenAiVideoTaskSeed {
|
||||||
|
/// Preserve existing database identity; older snapshots derive it from the local task ID.
|
||||||
|
#[serde(default)]
|
||||||
|
pub local_short_id: Option<String>,
|
||||||
|
#[serde(default)]
|
||||||
|
pub native_response: Option<Value>,
|
||||||
|
#[serde(default)]
|
||||||
|
pub xai_provider: bool,
|
||||||
pub local_task_id: String,
|
pub local_task_id: String,
|
||||||
pub upstream_task_id: String,
|
pub upstream_task_id: String,
|
||||||
pub created_at_unix_ms: u64,
|
pub created_at_unix_ms: u64,
|
||||||
|
|||||||
@@ -0,0 +1,147 @@
|
|||||||
|
# xAI provider behavior
|
||||||
|
|
||||||
|
The following rules preserve the provider-specific behavior of the `xai` provider
|
||||||
|
across Aether's request and transport layers.
|
||||||
|
|
||||||
|
## Responses and tools
|
||||||
|
|
||||||
|
- HTTP requests drop `previous_response_id`. Clients must supply conversation
|
||||||
|
history; this provider does not add an HTTP response-ID history store.
|
||||||
|
- `metadata.user_id` is removed. Claude clients copy it onto converted Responses
|
||||||
|
bodies and xAI rejects the field.
|
||||||
|
- Preserve requested `reasoning.encrypted_content`. On a native Responses-to-Responses
|
||||||
|
hop, keep provider-owned input items instead of rebuilding them through the canonical
|
||||||
|
format. xAI encrypted reasoning may have IDs that do not use OpenAI's `rs` prefix.
|
||||||
|
Aether's Gemini signature carriers remain excluded from xAI replay.
|
||||||
|
- The replay policy is selected from the configured provider type. A model called
|
||||||
|
`grok-*` on another provider does not opt into that policy. WebSocket continuation
|
||||||
|
metadata retains the selected policy across reconnects.
|
||||||
|
- A regular client function called `web_search` remains a function. Claude hosted
|
||||||
|
search choices are resolved against the original typed tool declaration, including
|
||||||
|
declarations with a different name.
|
||||||
|
- When only `image_generation` is allowed, keep only that tool and retain the requested
|
||||||
|
`auto` or `required` mode. For mixed allowed-tool lists, remove the image choice while
|
||||||
|
preserving the other allowed entries, as required by xAI's tool-choice schema.
|
||||||
|
- Reasoning effort is stripped for models that do not accept it.
|
||||||
|
- OpenAI-style image reference aliases in a request body are rewritten to xAI's
|
||||||
|
shape without touching chat message parts.
|
||||||
|
|
||||||
|
## Routing and credentials
|
||||||
|
|
||||||
|
OAuth requests default to `https://cli-chat-proxy.grok.com/v1`; API-key or
|
||||||
|
`using_api=true` requests default to `https://api.x.ai/v1`. Explicit custom gateways
|
||||||
|
are preserved. Compact remains on the official endpoint. CLI identity headers are
|
||||||
|
applied where the selected upstream requires them.
|
||||||
|
|
||||||
|
Account binding uses the xAI device code flow: the gateway requests a device code,
|
||||||
|
the operator authorizes it out of band, and the gateway polls for the token set.
|
||||||
|
There is no local callback listener, so headless deployments can bind accounts.
|
||||||
|
Refresh tokens can also be imported individually or in batches, and are rotated
|
||||||
|
on refresh.
|
||||||
|
|
||||||
|
Quota refresh reads `/user` and `/billing?format=credits` and stores a structured
|
||||||
|
usage snapshot. A prepaid balance keeps an account selectable after the weekly
|
||||||
|
allowance is exhausted. API-key accounts skip the subscription billing surface.
|
||||||
|
|
||||||
|
## Images and videos
|
||||||
|
|
||||||
|
OAuth media requests default to `https://cli-chat-proxy.grok.com/v1`; API-key or
|
||||||
|
`using_api=true` requests default to `https://api.x.ai/v1`. Explicit custom gateways
|
||||||
|
are preserved. Compact remains on the official endpoint. CLI identity headers are
|
||||||
|
applied to media requests and restored when a persisted video task's polling transport
|
||||||
|
is reconstructed.
|
||||||
|
|
||||||
|
Aether's OpenAI-compatible task parser accepts xAI's `request_id` creation field,
|
||||||
|
status aliases such as `pending` and `done`, nested `video.url` and `video.duration`,
|
||||||
|
and failure payloads containing `code` / `error` without a status. Existing OpenAI
|
||||||
|
`id` takes precedence. The client receives Aether's local task ID; polling uses the
|
||||||
|
upstream task ID and selected credential. Completed video downloads use the returned
|
||||||
|
media URL without forwarding provider authentication headers to the media host.
|
||||||
|
|
||||||
|
### Public video protocols
|
||||||
|
|
||||||
|
The xAI provider supports two video surfaces:
|
||||||
|
|
||||||
|
| Operation | xAI native | OpenAI compatible |
|
||||||
|
| --- | --- | --- |
|
||||||
|
| Create | `POST /v1/videos/generations` | `POST /openai/v1/videos` |
|
||||||
|
| Edit / extend | `POST /v1/videos/edits`, `POST /v1/videos/extensions` | — |
|
||||||
|
| Retrieve | `GET /v1/videos/{request_id}` | `GET /openai/v1/videos/{id}` |
|
||||||
|
| Download | use the returned `video.url` | `GET /openai/v1/videos/{id}/content` |
|
||||||
|
|
||||||
|
For xAI, `POST /v1/videos` is a native creation alias. Other providers retain
|
||||||
|
Aether's existing OpenAI-compatible `/v1/videos` behavior. xAI callers using
|
||||||
|
OpenAI `seconds` / `size` parameters must use `/openai/v1/videos`. The adapter
|
||||||
|
maps these to numeric `duration`, `aspect_ratio`, and `resolution`; it also adapts
|
||||||
|
image references. This implementation defaults to 4 seconds, portrait, and 720p,
|
||||||
|
clamps `duration` to 1-15, and validates inputs. Explicit native requests retain
|
||||||
|
native parameters and additional provider fields.
|
||||||
|
|
||||||
|
Default xAI creation targets `/videos/generations` on the selected upstream host.
|
||||||
|
Explicit custom endpoint paths still take precedence. Native generation, editing,
|
||||||
|
and extension paths only select xAI provider candidates.
|
||||||
|
|
||||||
|
Native creation returns `request_id`; native retrieval preserves `done`, nested
|
||||||
|
`video.url`, and provider fields such as `respect_moderation`. The identifier is
|
||||||
|
an opaque Aether task ID so queries remain scoped to the owning user and pinned
|
||||||
|
to the original upstream task and credential. The explicit `/openai/v1/videos`
|
||||||
|
surface projects `id`, `completed`, and `video_url`.
|
||||||
|
|
||||||
|
The task row records the native client protocol as `xai:video`, while its provider
|
||||||
|
transport remains `openai:video`. This survives restart without storing request
|
||||||
|
bodies or credentials. Raw native responses are cached only in memory; after
|
||||||
|
reconstruction the gateway refreshes from the original provider to recover its
|
||||||
|
response fields, including for completed tasks. If refreshing is unavailable,
|
||||||
|
the stored task still provides the native status and media URL projection.
|
||||||
|
|
||||||
|
OpenAI/xAI task persistence supplies a stable 16-character `short_id`, as required
|
||||||
|
by the PostgreSQL schema. Existing rows retain their original short ID across
|
||||||
|
reconstruction, including legacy embedded snapshots. This internal identifier is
|
||||||
|
separate from the opaque local task ID returned to clients; no schema change or
|
||||||
|
historical row rewrite is needed.
|
||||||
|
|
||||||
|
Task retrieval and content downloads are admitted by the production GET execution
|
||||||
|
gate. Reconstructed tasks resolve proxy nodes, system proxy defaults, tunnel affinity,
|
||||||
|
and transport profiles through the same deployment resolver used for creation;
|
||||||
|
configured proxy routes must not silently turn into direct requests after restart.
|
||||||
|
|
||||||
|
### Runtime configuration
|
||||||
|
|
||||||
|
Standalone Rust deployments must set
|
||||||
|
`AETHER_GATEWAY_VIDEO_TASK_TRUTH_SOURCE_MODE=rust-authoritative` and restart the
|
||||||
|
gateway to enable video task retrieval, polling, and content downloads. The CLI's
|
||||||
|
legacy default is `python-sync-report`: creation can return a task ID in that mode,
|
||||||
|
but the local task read/refresh paths are disabled and may return HTTP 503.
|
||||||
|
|
||||||
|
When the gateway also serves the frontend, `/openai/v1/videos` and its subpaths
|
||||||
|
must bypass the static SPA handler and be mounted as API routes. Otherwise a
|
||||||
|
successful-looking HTTP 200 response to a video query may contain `text/html`
|
||||||
|
instead of the task's JSON response. The lifecycle regression includes the static
|
||||||
|
frontend to cover this production configuration.
|
||||||
|
|
||||||
|
## Regression coverage
|
||||||
|
|
||||||
|
The format tests cover client and hosted search choices, image-only and mixed tool
|
||||||
|
restrictions, encrypted reasoning replay, image reference rewriting, and unchanged
|
||||||
|
OpenAI replay restrictions. Transport tests cover OAuth/API-key/custom routing and
|
||||||
|
media identity headers. Video-task tests exercise creation, polling, terminal
|
||||||
|
projection, persistence fields, content-download planning, and status-less errors
|
||||||
|
using local fixtures. They do not make paid generation requests.
|
||||||
|
|
||||||
|
The HTTP regression exercises all native creation paths and the compatibility
|
||||||
|
prefix through the public router and candidate planner, then checks polling,
|
||||||
|
cross-user denial, persistence, retrieval from a fresh gateway instance, and downloads
|
||||||
|
through both prefixes without leaking authorization to the media host. It uses the
|
||||||
|
real HTTP executor and a managed proxy node backed by a local test server, with no
|
||||||
|
execution-runtime override. The background poller also has a real HTTP proxy-node
|
||||||
|
regression, so production method guards and transport reconstruction are exercised.
|
||||||
|
CI also runs the same HTTP lifecycle with the PostgreSQL repository and the
|
||||||
|
production column constraints/indexes in an isolated temporary table. This catches
|
||||||
|
persistence failures that the in-memory repository cannot expose. The test uses
|
||||||
|
local `initdb`, `postgres`, and `pg_ctl` (already provided by the gateway CI job),
|
||||||
|
or an explicit `AETHER_TEST_DATABASE_URL` pointing to an isolated test database.
|
||||||
|
|
||||||
|
```sh
|
||||||
|
cargo test -p aether-ai-formats -p aether-provider-transport -p aether-video-tasks-core --lib
|
||||||
|
cargo test -p aether-gateway --lib xai
|
||||||
|
```
|
||||||
@@ -376,7 +376,7 @@ function jsonValueContainsAgentIdentity(value: unknown): boolean {
|
|||||||
export interface DeviceAuthorizeRequest {
|
export interface DeviceAuthorizeRequest {
|
||||||
start_url?: string
|
start_url?: string
|
||||||
region?: string
|
region?: string
|
||||||
auth_type?: 'builder_id' | 'identity_center' | 'google' | 'github' | 'browser'
|
auth_type?: 'builder_id' | 'identity_center' | 'google' | 'github' | 'browser' | 'device'
|
||||||
login_option?: 'google' | 'github' | 'default'
|
login_option?: 'google' | 'github' | 'default'
|
||||||
redirect_uri?: string
|
redirect_uri?: string
|
||||||
proxy_node_id?: string
|
proxy_node_id?: string
|
||||||
|
|||||||
@@ -462,6 +462,23 @@ export interface GrokUpstreamMetadata {
|
|||||||
account_user_id?: string | null
|
account_user_id?: string | null
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export interface XaiUpstreamMetadata {
|
||||||
|
updated_at?: number
|
||||||
|
subscription_title?: string
|
||||||
|
usage_percentage?: number
|
||||||
|
remaining_percentage?: number
|
||||||
|
usage_label?: string
|
||||||
|
usage_limit?: number
|
||||||
|
current_usage?: number
|
||||||
|
remaining?: number
|
||||||
|
next_reset_at?: number
|
||||||
|
prepaid_balance?: number
|
||||||
|
on_demand_cap?: number
|
||||||
|
on_demand_used?: number
|
||||||
|
on_demand_remaining?: number
|
||||||
|
period_type?: string
|
||||||
|
}
|
||||||
|
|
||||||
export interface GeminiCliTierMetadata {
|
export interface GeminiCliTierMetadata {
|
||||||
id?: string | null
|
id?: string | null
|
||||||
tierType?: string | null
|
tierType?: string | null
|
||||||
@@ -520,6 +537,7 @@ export interface UpstreamMetadata {
|
|||||||
chatgpt_web?: ChatGPTWebUpstreamMetadata
|
chatgpt_web?: ChatGPTWebUpstreamMetadata
|
||||||
grok?: GrokUpstreamMetadata
|
grok?: GrokUpstreamMetadata
|
||||||
gemini_cli?: GeminiCliUpstreamMetadata
|
gemini_cli?: GeminiCliUpstreamMetadata
|
||||||
|
xai?: XaiUpstreamMetadata
|
||||||
}
|
}
|
||||||
|
|
||||||
// 按格式的健康度数据
|
// 按格式的健康度数据
|
||||||
@@ -758,7 +776,7 @@ export interface HealthRelatedMonitorResponse {
|
|||||||
related_providers: HealthRelatedMonitor[]
|
related_providers: HealthRelatedMonitor[]
|
||||||
}
|
}
|
||||||
|
|
||||||
export type ProviderType = 'custom' | 'claude_code' | 'codex' | 'chatgpt_web' | 'gemini_cli' | 'antigravity' | 'kiro' | 'grok' | 'windsurf' | 'vertex_ai'
|
export type ProviderType = 'custom' | 'claude_code' | 'codex' | 'chatgpt_web' | 'gemini_cli' | 'antigravity' | 'kiro' | 'grok' | 'xai' | 'windsurf' | 'vertex_ai'
|
||||||
|
|
||||||
export interface ClaudeCodeAdvancedConfig {
|
export interface ClaudeCodeAdvancedConfig {
|
||||||
// 会话数量控制:null/undefined 表示不限制
|
// 会话数量控制:null/undefined 表示不限制
|
||||||
|
|||||||
@@ -334,7 +334,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
|
|||||||
label: 'Free/Team 优先',
|
label: 'Free/Team 优先',
|
||||||
description: '兼容旧配置:优先消耗 Free、Team 或两者',
|
description: '兼容旧配置:优先消耗 Free、Team 或两者',
|
||||||
evidence_hint: '依据 plan_type,保留旧 free_only/team_only/both 语义',
|
evidence_hint: '依据 plan_type,保留旧 free_only/team_only/both 语义',
|
||||||
providers: ['codex', 'grok', 'kiro', 'windsurf'],
|
providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'],
|
||||||
modes: [
|
modes: [
|
||||||
{ value: 'free_only', label: 'Free' },
|
{ value: 'free_only', label: 'Free' },
|
||||||
{ value: 'team_only', label: 'Team' },
|
{ value: 'team_only', label: 'Team' },
|
||||||
@@ -347,7 +347,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
|
|||||||
label: 'Free 优先',
|
label: 'Free 优先',
|
||||||
description: '优先消耗 Free 账号(依赖 plan_type)',
|
description: '优先消耗 Free 账号(依赖 plan_type)',
|
||||||
evidence_hint: '依据 plan_type(Free 账号优先调度)',
|
evidence_hint: '依据 plan_type(Free 账号优先调度)',
|
||||||
providers: ['codex', 'grok', 'kiro', 'windsurf'],
|
providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'],
|
||||||
modes: null,
|
modes: null,
|
||||||
default_mode: null,
|
default_mode: null,
|
||||||
},
|
},
|
||||||
@@ -356,7 +356,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
|
|||||||
label: 'Team 优先',
|
label: 'Team 优先',
|
||||||
description: '优先消耗 Team 账号(依赖 plan_type)',
|
description: '优先消耗 Team 账号(依赖 plan_type)',
|
||||||
evidence_hint: '依据 plan_type(Team 账号优先调度)',
|
evidence_hint: '依据 plan_type(Team 账号优先调度)',
|
||||||
providers: ['codex', 'grok', 'kiro', 'windsurf'],
|
providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'],
|
||||||
modes: null,
|
modes: null,
|
||||||
default_mode: null,
|
default_mode: null,
|
||||||
},
|
},
|
||||||
@@ -365,7 +365,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
|
|||||||
label: 'Plus 优先',
|
label: 'Plus 优先',
|
||||||
description: '优先消耗 Plus 账号(依赖 plan_type)',
|
description: '优先消耗 Plus 账号(依赖 plan_type)',
|
||||||
evidence_hint: '依据 plan_type(Plus 账号优先调度)',
|
evidence_hint: '依据 plan_type(Plus 账号优先调度)',
|
||||||
providers: ['codex', 'grok', 'kiro', 'windsurf'],
|
providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'],
|
||||||
modes: null,
|
modes: null,
|
||||||
default_mode: null,
|
default_mode: null,
|
||||||
},
|
},
|
||||||
@@ -374,7 +374,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
|
|||||||
label: 'Pro 优先',
|
label: 'Pro 优先',
|
||||||
description: '优先消耗 Pro 账号(依赖 plan_type)',
|
description: '优先消耗 Pro 账号(依赖 plan_type)',
|
||||||
evidence_hint: '依据 plan_type(Pro 账号优先调度)',
|
evidence_hint: '依据 plan_type(Pro 账号优先调度)',
|
||||||
providers: ['codex', 'grok', 'kiro', 'windsurf'],
|
providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'],
|
||||||
modes: null,
|
modes: null,
|
||||||
default_mode: null,
|
default_mode: null,
|
||||||
},
|
},
|
||||||
@@ -392,7 +392,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
|
|||||||
label: '额度刷新优先',
|
label: '额度刷新优先',
|
||||||
description: '优先选即将刷新额度的账号',
|
description: '优先选即将刷新额度的账号',
|
||||||
evidence_hint: '依据账号额度重置倒计时(next_reset / reset_seconds)',
|
evidence_hint: '依据账号额度重置倒计时(next_reset / reset_seconds)',
|
||||||
providers: ['codex', 'grok', 'kiro', 'windsurf'],
|
providers: ['codex', 'grok', 'kiro', 'windsurf', 'xai'],
|
||||||
default_enabled_providers: ['codex', 'windsurf'],
|
default_enabled_providers: ['codex', 'windsurf'],
|
||||||
modes: null,
|
modes: null,
|
||||||
default_mode: null,
|
default_mode: null,
|
||||||
|
|||||||
@@ -244,6 +244,129 @@
|
|||||||
</div>
|
</div>
|
||||||
</template>
|
</template>
|
||||||
|
|
||||||
|
<!-- xAI: 设备授权 -->
|
||||||
|
<template v-else-if="isXaiProvider">
|
||||||
|
<div class="space-y-3">
|
||||||
|
<div class="h-[265px]">
|
||||||
|
<div
|
||||||
|
v-if="device.status === 'error' || device.status === 'expired'"
|
||||||
|
class="rounded-xl border border-destructive/20 bg-destructive/5 p-5"
|
||||||
|
>
|
||||||
|
<div class="flex flex-col items-center text-center space-y-3">
|
||||||
|
<div class="w-10 h-10 rounded-full bg-destructive/10 flex items-center justify-center">
|
||||||
|
<AlertCircle class="w-5 h-5 text-destructive" />
|
||||||
|
</div>
|
||||||
|
<div class="space-y-1">
|
||||||
|
<p class="text-sm font-medium text-destructive">
|
||||||
|
{{ legacyT(device.status === 'expired' ? '授权已过期' : '授权失败') }}
|
||||||
|
</p>
|
||||||
|
<p class="text-xs text-muted-foreground">
|
||||||
|
{{ legacyT(device.error || '请重试') }}
|
||||||
|
</p>
|
||||||
|
</div>
|
||||||
|
<Button
|
||||||
|
size="sm"
|
||||||
|
variant="outline"
|
||||||
|
@click="resetDevice"
|
||||||
|
>
|
||||||
|
{{ legacyT('重新开始') }}
|
||||||
|
</Button>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<div
|
||||||
|
v-else-if="device.starting && !device.session_id"
|
||||||
|
class="flex items-center justify-center py-12"
|
||||||
|
>
|
||||||
|
<div class="text-center">
|
||||||
|
<div class="animate-spin rounded-full h-6 w-6 border-b-2 border-primary mx-auto mb-3" />
|
||||||
|
<p class="text-xs text-muted-foreground">
|
||||||
|
{{ legacyT('正在准备设备授权...') }}
|
||||||
|
</p>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<div
|
||||||
|
v-else-if="device.session_id && device.status === 'pending'"
|
||||||
|
class="rounded-xl border border-border bg-muted/20 p-5"
|
||||||
|
>
|
||||||
|
<div class="flex flex-col items-center text-center space-y-4">
|
||||||
|
<div class="relative">
|
||||||
|
<div class="absolute inset-0 rounded-full bg-primary/20 animate-ping" />
|
||||||
|
<div class="relative w-10 h-10 rounded-full bg-primary/10 flex items-center justify-center">
|
||||||
|
<ExternalLink class="w-5 h-5 text-primary" />
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<div class="space-y-1">
|
||||||
|
<p class="text-sm font-medium">
|
||||||
|
{{ legacyT('在浏览器中输入设备码完成授权') }}
|
||||||
|
</p>
|
||||||
|
<p class="text-xs text-muted-foreground">
|
||||||
|
{{ legacyT('授权完成后此页面将自动更新') }}
|
||||||
|
</p>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<div
|
||||||
|
v-if="device.user_code"
|
||||||
|
class="flex items-center gap-2 rounded-lg border border-border bg-background px-3 py-2"
|
||||||
|
>
|
||||||
|
<span class="text-lg font-mono font-bold tracking-[0.2em]">{{ device.user_code }}</span>
|
||||||
|
<button
|
||||||
|
class="p-1 rounded hover:bg-muted transition-colors"
|
||||||
|
:title="legacyT('复制设备码')"
|
||||||
|
@click="copyToClipboard(device.user_code)"
|
||||||
|
>
|
||||||
|
<Copy class="w-3.5 h-3.5 text-muted-foreground" />
|
||||||
|
</button>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<div class="flex items-center gap-1.5 text-xs text-muted-foreground">
|
||||||
|
<div class="animate-spin rounded-full h-3 w-3 border-[1.5px] border-primary/30 border-t-primary" />
|
||||||
|
<span>{{ remainingText }}</span>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<div class="flex gap-2 w-full">
|
||||||
|
<Button
|
||||||
|
class="flex-1"
|
||||||
|
size="sm"
|
||||||
|
:disabled="!device.verification_uri_complete && !device.verification_uri"
|
||||||
|
@click="openDeviceVerificationUrl"
|
||||||
|
>
|
||||||
|
<ExternalLink class="w-3.5 h-3.5 mr-1.5" />
|
||||||
|
{{ legacyT('打开授权页面') }}
|
||||||
|
</Button>
|
||||||
|
<Button
|
||||||
|
size="sm"
|
||||||
|
variant="outline"
|
||||||
|
:disabled="!device.verification_uri_complete && !device.verification_uri"
|
||||||
|
@click="copyToClipboard(device.verification_uri_complete || device.verification_uri)"
|
||||||
|
>
|
||||||
|
<Copy class="w-3.5 h-3.5" />
|
||||||
|
</Button>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<div
|
||||||
|
v-else
|
||||||
|
class="flex h-full flex-col items-center justify-center gap-3"
|
||||||
|
>
|
||||||
|
<p class="text-xs text-muted-foreground text-center">
|
||||||
|
{{ legacyT('使用 xAI 设备授权登录 Grok CLI,或改为导入 API Key / Refresh Token。') }}
|
||||||
|
</p>
|
||||||
|
<Button
|
||||||
|
class="w-full"
|
||||||
|
:disabled="device.starting"
|
||||||
|
@click="startDeviceAuth"
|
||||||
|
>
|
||||||
|
{{ device.starting ? legacyT('正在准备授权...') : legacyT('开始授权') }}
|
||||||
|
</Button>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</template>
|
||||||
|
|
||||||
<!-- Kiro: 设备授权模式 -->
|
<!-- Kiro: 设备授权模式 -->
|
||||||
<template v-else-if="isKiroProvider">
|
<template v-else-if="isKiroProvider">
|
||||||
<div class="space-y-3">
|
<div class="space-y-3">
|
||||||
@@ -994,7 +1117,7 @@ let oauthInitRequestId = 0
|
|||||||
let oauthCompleteRequestId = 0
|
let oauthCompleteRequestId = 0
|
||||||
|
|
||||||
// 设备授权状态
|
// 设备授权状态
|
||||||
type DeviceAuthType = 'default' | 'google' | 'github' | 'builder_id' | 'identity_center'
|
type DeviceAuthType = 'default' | 'google' | 'github' | 'builder_id' | 'identity_center' | 'device'
|
||||||
type WindsurfLoginOption = 'default' | 'google' | 'github'
|
type WindsurfLoginOption = 'default' | 'google' | 'github'
|
||||||
|
|
||||||
interface DeviceAuthState {
|
interface DeviceAuthState {
|
||||||
@@ -1075,9 +1198,10 @@ const isOpen = computed(() => props.open)
|
|||||||
const isKiroProvider = computed(() => (props.providerType || '').toLowerCase() === 'kiro')
|
const isKiroProvider = computed(() => (props.providerType || '').toLowerCase() === 'kiro')
|
||||||
const isGrokProvider = computed(() => (props.providerType || '').toLowerCase() === 'grok')
|
const isGrokProvider = computed(() => (props.providerType || '').toLowerCase() === 'grok')
|
||||||
const isWindsurfProvider = computed(() => (props.providerType || '').toLowerCase() === 'windsurf')
|
const isWindsurfProvider = computed(() => (props.providerType || '').toLowerCase() === 'windsurf')
|
||||||
|
const isXaiProvider = computed(() => (props.providerType || '').toLowerCase() === 'xai')
|
||||||
const isCodexProvider = computed(() => (props.providerType || '').toLowerCase() === 'codex')
|
const isCodexProvider = computed(() => (props.providerType || '').toLowerCase() === 'codex')
|
||||||
const isClaudeCodeProvider = computed(() => (props.providerType || '').toLowerCase() === 'claude_code')
|
const isClaudeCodeProvider = computed(() => (props.providerType || '').toLowerCase() === 'claude_code')
|
||||||
const isDeviceBrowserProvider = computed(() => isKiroProvider.value || isWindsurfProvider.value)
|
const isDeviceBrowserProvider = computed(() => isKiroProvider.value || isWindsurfProvider.value || isXaiProvider.value)
|
||||||
const showAuthorizationMode = computed(() => !isGrokProvider.value)
|
const showAuthorizationMode = computed(() => !isGrokProvider.value)
|
||||||
const defaultMode = computed<DialogMode>(() => (isGrokProvider.value ? 'import' : 'oauth'))
|
const defaultMode = computed<DialogMode>(() => (isGrokProvider.value ? 'import' : 'oauth'))
|
||||||
|
|
||||||
@@ -1101,7 +1225,7 @@ const isManualDeviceCallbackPending = computed(() =>
|
|||||||
|
|
||||||
const authorizationModeLabel = computed(() => {
|
const authorizationModeLabel = computed(() => {
|
||||||
if (isWindsurfProvider.value) return legacyT('浏览器登录')
|
if (isWindsurfProvider.value) return legacyT('浏览器登录')
|
||||||
if (isDeviceBrowserProvider.value) return legacyT('设备授权')
|
if (isXaiProvider.value || isDeviceBrowserProvider.value) return legacyT('设备授权')
|
||||||
return legacyT('获取授权')
|
return legacyT('获取授权')
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -1215,6 +1339,9 @@ const importManualPlaceholder = computed(() => {
|
|||||||
if (isClaudeCodeProvider.value) {
|
if (isClaudeCodeProvider.value) {
|
||||||
return legacyT('粘贴 Claude Refresh Token 或 Claude Code .credentials.json 内容')
|
return legacyT('粘贴 Claude Refresh Token 或 Claude Code .credentials.json 内容')
|
||||||
}
|
}
|
||||||
|
if (isXaiProvider.value) {
|
||||||
|
return legacyT('粘贴 xAI API Key、Access Token,或包含 refresh_token / api_key 的 JSON')
|
||||||
|
}
|
||||||
if (isWindsurfProvider.value) {
|
if (isWindsurfProvider.value) {
|
||||||
return legacyT('粘贴 show-auth-token Token、API key 或 JSON 内容')
|
return legacyT('粘贴 show-auth-token Token、API key 或 JSON 内容')
|
||||||
}
|
}
|
||||||
@@ -1477,11 +1604,17 @@ function resetDevice() {
|
|||||||
totp.stop()
|
totp.stop()
|
||||||
const { auth_type, start_url, region, totp_secret } = device.value
|
const { auth_type, start_url, region, totp_secret } = device.value
|
||||||
device.value = createInitialDeviceState()
|
device.value = createInitialDeviceState()
|
||||||
device.value.auth_type = isWindsurfProvider.value ? (auth_type === 'google' || auth_type === 'github' ? auth_type : 'default') : auth_type
|
device.value.auth_type = isXaiProvider.value
|
||||||
|
? 'device'
|
||||||
|
: isWindsurfProvider.value
|
||||||
|
? (auth_type === 'google' || auth_type === 'github' ? auth_type : 'default')
|
||||||
|
: auth_type
|
||||||
device.value.start_url = start_url
|
device.value.start_url = start_url
|
||||||
device.value.region = region
|
device.value.region = region
|
||||||
device.value.totp_secret = totp_secret
|
device.value.totp_secret = totp_secret
|
||||||
if (!isWindsurfProvider.value && (device.value.auth_type === 'google' || device.value.auth_type === 'github')) {
|
if (isXaiProvider.value) {
|
||||||
|
void ensureXaiDeviceAuth()
|
||||||
|
} else if (!isWindsurfProvider.value && (device.value.auth_type === 'google' || device.value.auth_type === 'github')) {
|
||||||
void ensureKiroSocialDeviceAuth()
|
void ensureKiroSocialDeviceAuth()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1500,6 +1633,8 @@ function resetForm() {
|
|||||||
device.value = createInitialDeviceState()
|
device.value = createInitialDeviceState()
|
||||||
if (isWindsurfProvider.value) {
|
if (isWindsurfProvider.value) {
|
||||||
device.value.auth_type = 'default'
|
device.value.auth_type = 'default'
|
||||||
|
} else if (isXaiProvider.value) {
|
||||||
|
device.value.auth_type = 'device'
|
||||||
}
|
}
|
||||||
importText.value = ''
|
importText.value = ''
|
||||||
importing.value = false
|
importing.value = false
|
||||||
@@ -1531,6 +1666,8 @@ function switchMode(newMode: DialogMode) {
|
|||||||
if (newMode === 'oauth') {
|
if (newMode === 'oauth') {
|
||||||
if (isKiroProvider.value) {
|
if (isKiroProvider.value) {
|
||||||
void ensureKiroSocialDeviceAuth()
|
void ensureKiroSocialDeviceAuth()
|
||||||
|
} else if (isXaiProvider.value) {
|
||||||
|
void ensureXaiDeviceAuth()
|
||||||
} else if (!oauth.value.authorization_url && !oauth.value.starting) {
|
} else if (!oauth.value.authorization_url && !oauth.value.starting) {
|
||||||
initOAuth()
|
initOAuth()
|
||||||
}
|
}
|
||||||
@@ -1833,6 +1970,29 @@ function parseImportText(text: string): {
|
|||||||
return { refresh_token: trimmed }
|
return { refresh_token: trimmed }
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if (isXaiProvider.value) {
|
||||||
|
try {
|
||||||
|
const parsed: unknown = JSON.parse(trimmed)
|
||||||
|
if (typeof parsed === 'object' && parsed !== null) {
|
||||||
|
const obj = parsed as Record<string, unknown>
|
||||||
|
const apiKey = normalizeStringField(obj.api_key) ?? normalizeStringField(obj.apiKey)
|
||||||
|
const refreshToken = normalizeStringField(obj.refresh_token) ?? normalizeStringField(obj.refreshToken)
|
||||||
|
const accessToken = normalizeStringField(obj.access_token) ?? normalizeStringField(obj.accessToken) ?? apiKey
|
||||||
|
if (refreshToken || accessToken) {
|
||||||
|
return {
|
||||||
|
refresh_token: refreshToken,
|
||||||
|
access_token: accessToken,
|
||||||
|
name: normalizeStringField(obj.name) ?? normalizeStringField(obj.email),
|
||||||
|
email: normalizeStringField(obj.email),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} catch {
|
||||||
|
// Raw xAI API keys / access tokens are imported as access_token.
|
||||||
|
}
|
||||||
|
return { access_token: trimmed }
|
||||||
|
}
|
||||||
|
|
||||||
if (isGrokProvider.value) {
|
if (isGrokProvider.value) {
|
||||||
const cookieImport = parseGrokCookieImport(trimmed)
|
const cookieImport = parseGrokCookieImport(trimmed)
|
||||||
if (cookieImport) {
|
if (cookieImport) {
|
||||||
@@ -2367,17 +2527,20 @@ async function startDeviceAuth() {
|
|||||||
device.value.error = ''
|
device.value.error = ''
|
||||||
try {
|
try {
|
||||||
const isWindsurf = isWindsurfProvider.value
|
const isWindsurf = isWindsurfProvider.value
|
||||||
|
const isXai = isXaiProvider.value
|
||||||
const isBuilderID = requestedAuthType === 'builder_id'
|
const isBuilderID = requestedAuthType === 'builder_id'
|
||||||
const isSocial = requestedAuthType === 'google' || requestedAuthType === 'github'
|
const isSocial = !isXai && (requestedAuthType === 'google' || requestedAuthType === 'github')
|
||||||
const windsurfLoginOption: WindsurfLoginOption = isSocial ? requestedAuthType : 'default'
|
const windsurfLoginOption: WindsurfLoginOption = isSocial ? requestedAuthType : 'default'
|
||||||
const authTypeForRequest = isWindsurf
|
const authTypeForRequest = isWindsurf
|
||||||
? 'browser'
|
? 'browser'
|
||||||
: (requestedAuthType === 'default' ? 'google' : requestedAuthType)
|
: isXai
|
||||||
|
? 'device'
|
||||||
|
: (requestedAuthType === 'default' ? 'google' : requestedAuthType)
|
||||||
const resp = await startDeviceAuthorize(props.providerId, {
|
const resp = await startDeviceAuthorize(props.providerId, {
|
||||||
auth_type: authTypeForRequest,
|
auth_type: authTypeForRequest,
|
||||||
login_option: isWindsurf ? windsurfLoginOption : undefined,
|
login_option: isWindsurf ? windsurfLoginOption : undefined,
|
||||||
start_url: isWindsurf ? undefined : (isBuilderID ? BUILDER_ID_START_URL : (isSocial ? undefined : (device.value.start_url.trim() || undefined))),
|
start_url: (isWindsurf || isXai) ? undefined : (isBuilderID ? BUILDER_ID_START_URL : (isSocial ? undefined : (device.value.start_url.trim() || undefined))),
|
||||||
region: isWindsurf ? undefined : (isBuilderID || isSocial ? BUILDER_ID_REGION : (device.value.region.trim() || undefined)),
|
region: (isWindsurf || isXai) ? undefined : (isBuilderID || isSocial ? BUILDER_ID_REGION : (device.value.region.trim() || undefined)),
|
||||||
proxy_node_id: selectedProxyNodeId.value || undefined,
|
proxy_node_id: selectedProxyNodeId.value || undefined,
|
||||||
})
|
})
|
||||||
if (requestId !== deviceAuthRequestId || device.value.auth_type !== requestedAuthType) return
|
if (requestId !== deviceAuthRequestId || device.value.auth_type !== requestedAuthType) return
|
||||||
@@ -2417,6 +2580,13 @@ async function ensureKiroSocialDeviceAuth() {
|
|||||||
await startDeviceAuth()
|
await startDeviceAuth()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async function ensureXaiDeviceAuth() {
|
||||||
|
if (!props.open || !props.providerId || !isXaiProvider.value) return
|
||||||
|
if (device.value.starting) return
|
||||||
|
if (device.value.session_id && (device.value.status === 'pending' || device.value.status === 'authorized')) return
|
||||||
|
await startDeviceAuth()
|
||||||
|
}
|
||||||
|
|
||||||
function scheduleDevicePoll() {
|
function scheduleDevicePoll() {
|
||||||
if (devicePollTimer) clearTimeout(devicePollTimer)
|
if (devicePollTimer) clearTimeout(devicePollTimer)
|
||||||
devicePollTimer = setTimeout(() => pollDevice(), device.value.interval * 1000)
|
devicePollTimer = setTimeout(() => pollDevice(), device.value.interval * 1000)
|
||||||
@@ -2525,6 +2695,9 @@ watch(
|
|||||||
}
|
}
|
||||||
if (isWindsurfProvider.value) {
|
if (isWindsurfProvider.value) {
|
||||||
device.value.auth_type = 'default'
|
device.value.auth_type = 'default'
|
||||||
|
} else if (isXaiProvider.value) {
|
||||||
|
device.value.auth_type = 'device'
|
||||||
|
void ensureXaiDeviceAuth()
|
||||||
} else if (isKiroProvider.value) {
|
} else if (isKiroProvider.value) {
|
||||||
void ensureKiroSocialDeviceAuth()
|
void ensureKiroSocialDeviceAuth()
|
||||||
} else {
|
} else {
|
||||||
@@ -2554,6 +2727,9 @@ watch(
|
|||||||
device.value.auth_type = ['default', 'google', 'github'].includes(device.value.auth_type)
|
device.value.auth_type = ['default', 'google', 'github'].includes(device.value.auth_type)
|
||||||
? device.value.auth_type
|
? device.value.auth_type
|
||||||
: 'default'
|
: 'default'
|
||||||
|
} else if (props.open && isXaiProvider.value && mode.value === 'oauth') {
|
||||||
|
device.value.auth_type = 'device'
|
||||||
|
void ensureXaiDeviceAuth()
|
||||||
} else if (props.open && isKiroProvider.value && mode.value === 'oauth') {
|
} else if (props.open && isKiroProvider.value && mode.value === 'oauth') {
|
||||||
void ensureKiroSocialDeviceAuth()
|
void ensureKiroSocialDeviceAuth()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -514,6 +514,65 @@
|
|||||||
</div>
|
</div>
|
||||||
</template>
|
</template>
|
||||||
</div>
|
</div>
|
||||||
|
<!-- xAI / Grok Build 订阅额度 -->
|
||||||
|
<div
|
||||||
|
v-if="provider.provider_type === 'xai' && hasXaiQuotaDisplayData(key)"
|
||||||
|
class="mt-2 p-2 rounded-md bg-muted/30"
|
||||||
|
>
|
||||||
|
<ProviderQuotaSectionHeader
|
||||||
|
:title="legacyT('账号配额')"
|
||||||
|
:loading="refreshingQuota"
|
||||||
|
:updated-text="getXaiQuotaDisplay(key)?.updated_at ? formatKiroUpdatedAt(getXaiQuotaDisplay(key)?.updated_at || 0) : null"
|
||||||
|
/>
|
||||||
|
<div class="space-y-2">
|
||||||
|
<ProviderQuotaProgressRow
|
||||||
|
v-if="getXaiQuotaDisplay(key)?.usage_percentage !== undefined || getXaiQuotaDisplay(key)?.remaining_percentage !== undefined"
|
||||||
|
:label="legacyT(getXaiUsageLabel(key))"
|
||||||
|
:used-percent="getXaiUsedPercent(key)"
|
||||||
|
:remaining-percent="getXaiRemainingPercent(key)"
|
||||||
|
:meter-class="getQuotaRemainingClass(getXaiUsedPercent(key))"
|
||||||
|
:bar-class="getQuotaRemainingBarColor(getXaiUsedPercent(key))"
|
||||||
|
:reset-text="getXaiQuotaDisplay(key)?.next_reset_at
|
||||||
|
? `${formatKiroResetTime(getXaiQuotaDisplay(key)?.next_reset_at)}${legacyT('重置')}`
|
||||||
|
: null"
|
||||||
|
>
|
||||||
|
<template
|
||||||
|
v-if="getXaiQuotaDisplay(key)?.usage_limit != null"
|
||||||
|
#footer
|
||||||
|
>
|
||||||
|
<div class="flex items-center justify-between text-[9px] text-muted-foreground/70 mt-0.5">
|
||||||
|
<span>
|
||||||
|
{{ formatKiroUsage(getXaiQuotaDisplay(key)?.current_usage) }} /
|
||||||
|
{{ formatKiroUsage(getXaiQuotaDisplay(key)?.usage_limit) }}
|
||||||
|
</span>
|
||||||
|
<span v-if="getXaiQuotaDisplay(key)?.next_reset_at">
|
||||||
|
{{ formatKiroResetTime(getXaiQuotaDisplay(key)?.next_reset_at) }}{{ legacyT('重置') }}
|
||||||
|
</span>
|
||||||
|
</div>
|
||||||
|
</template>
|
||||||
|
</ProviderQuotaProgressRow>
|
||||||
|
<div
|
||||||
|
v-if="getXaiQuotaDisplay(key)?.prepaid_balance != null"
|
||||||
|
class="text-[10px] text-muted-foreground"
|
||||||
|
>
|
||||||
|
{{ legacyT('预付额度') }}: {{ formatKiroUsage(getXaiQuotaDisplay(key)?.prepaid_balance) }}
|
||||||
|
</div>
|
||||||
|
<ProviderQuotaProgressRow
|
||||||
|
v-if="getXaiQuotaDisplay(key)?.on_demand_cap"
|
||||||
|
:label="legacyT('按需额度')"
|
||||||
|
:used-percent="getXaiOnDemandUsedPercent(key)"
|
||||||
|
:meter-class="getQuotaRemainingClass(getXaiOnDemandUsedPercent(key))"
|
||||||
|
:bar-class="getQuotaRemainingBarColor(getXaiOnDemandUsedPercent(key))"
|
||||||
|
>
|
||||||
|
<template #footer>
|
||||||
|
<div class="text-[9px] text-muted-foreground/70 mt-0.5">
|
||||||
|
{{ formatKiroUsage(getXaiQuotaDisplay(key)?.on_demand_used) }} /
|
||||||
|
{{ formatKiroUsage(getXaiQuotaDisplay(key)?.on_demand_cap) }}
|
||||||
|
</div>
|
||||||
|
</template>
|
||||||
|
</ProviderQuotaProgressRow>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
<!-- Windsurf 上游额度信息 -->
|
<!-- Windsurf 上游额度信息 -->
|
||||||
<div
|
<div
|
||||||
v-if="provider.provider_type === 'windsurf' && (hasWindsurfQuotaDisplayData(key) || isWindsurfUnavailableKey(key) || isWindsurfExhaustedKey(key))"
|
v-if="provider.provider_type === 'windsurf' && (hasWindsurfQuotaDisplayData(key) || isWindsurfUnavailableKey(key) || isWindsurfExhaustedKey(key))"
|
||||||
@@ -1005,6 +1064,7 @@ import type {
|
|||||||
GrokUpstreamMetadata,
|
GrokUpstreamMetadata,
|
||||||
KiroUpstreamMetadata,
|
KiroUpstreamMetadata,
|
||||||
WindsurfUpstreamMetadata,
|
WindsurfUpstreamMetadata,
|
||||||
|
XaiUpstreamMetadata,
|
||||||
QuotaResetCreditsSnapshot,
|
QuotaResetCreditsSnapshot,
|
||||||
QuotaStatusSnapshot,
|
QuotaStatusSnapshot,
|
||||||
QuotaWindowSnapshot,
|
QuotaWindowSnapshot,
|
||||||
@@ -1858,7 +1918,7 @@ function quotaSnapshotHasDisplayData(quota: QuotaStatusSnapshot | null | undefin
|
|||||||
|
|
||||||
function getQuotaSnapshotForProvider(
|
function getQuotaSnapshotForProvider(
|
||||||
key: EndpointAPIKey,
|
key: EndpointAPIKey,
|
||||||
providerType: 'codex' | 'kiro' | 'windsurf' | 'antigravity' | 'chatgpt_web' | 'gemini_cli' | 'grok',
|
providerType: 'codex' | 'kiro' | 'windsurf' | 'antigravity' | 'chatgpt_web' | 'gemini_cli' | 'grok' | 'xai',
|
||||||
): QuotaStatusSnapshot | null {
|
): QuotaStatusSnapshot | null {
|
||||||
const quota = key.status_snapshot?.quota
|
const quota = key.status_snapshot?.quota
|
||||||
if (!quota) return null
|
if (!quota) return null
|
||||||
@@ -2179,6 +2239,90 @@ function hasKiroQuotaDisplayData(key: EndpointAPIKey): boolean {
|
|||||||
return !!kiro && (kiro.usage_percentage !== undefined || kiro.usage_limit !== undefined)
|
return !!kiro && (kiro.usage_percentage !== undefined || kiro.usage_limit !== undefined)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function getXaiQuotaDisplay(key: EndpointAPIKey): XaiUpstreamMetadata | null {
|
||||||
|
const quota = getQuotaSnapshotForProvider(key, 'xai')
|
||||||
|
if (!quota) return null
|
||||||
|
|
||||||
|
const display: XaiUpstreamMetadata = {}
|
||||||
|
const updatedAt = getQuotaSnapshotUpdatedAt(quota)
|
||||||
|
if (updatedAt !== undefined) display.updated_at = updatedAt
|
||||||
|
if (quota.plan_type) display.subscription_title = quota.plan_type
|
||||||
|
|
||||||
|
const usageWindow =
|
||||||
|
getQuotaWindow(quota, 'usage')
|
||||||
|
?? getQuotaWindowByScope(quota, 'account')[0]
|
||||||
|
?? null
|
||||||
|
if (usageWindow) {
|
||||||
|
const usedPercent = getQuotaWindowUsedPercent(usageWindow)
|
||||||
|
const remainingPercent = getQuotaWindowRemainingPercent(usageWindow)
|
||||||
|
if (usedPercent !== undefined) display.usage_percentage = usedPercent
|
||||||
|
if (remainingPercent !== undefined) display.remaining_percentage = remainingPercent
|
||||||
|
const usageLabel = String(usageWindow.label || '').trim()
|
||||||
|
if (usageLabel) display.usage_label = usageLabel
|
||||||
|
if (typeof usageWindow.used_value === 'number') display.current_usage = usageWindow.used_value
|
||||||
|
if (typeof usageWindow.limit_value === 'number') display.usage_limit = usageWindow.limit_value
|
||||||
|
if (typeof usageWindow.remaining_value === 'number') display.remaining = usageWindow.remaining_value
|
||||||
|
const nextResetAt =
|
||||||
|
getQuotaWindowResetAt(usageWindow)
|
||||||
|
?? (() => {
|
||||||
|
const resetSeconds = getQuotaWindowResetSeconds(usageWindow)
|
||||||
|
if (updatedAt === undefined || resetSeconds === undefined) return undefined
|
||||||
|
return updatedAt + resetSeconds
|
||||||
|
})()
|
||||||
|
if (nextResetAt !== undefined) display.next_reset_at = nextResetAt
|
||||||
|
}
|
||||||
|
|
||||||
|
const prepaidWindow = getQuotaWindow(quota, 'prepaid')
|
||||||
|
if (typeof prepaidWindow?.remaining_value === 'number') {
|
||||||
|
display.prepaid_balance = prepaidWindow.remaining_value
|
||||||
|
}
|
||||||
|
|
||||||
|
const onDemandWindow = getQuotaWindow(quota, 'on_demand')
|
||||||
|
if (typeof onDemandWindow?.limit_value === 'number') display.on_demand_cap = onDemandWindow.limit_value
|
||||||
|
if (typeof onDemandWindow?.used_value === 'number') display.on_demand_used = onDemandWindow.used_value
|
||||||
|
if (typeof onDemandWindow?.remaining_value === 'number') display.on_demand_remaining = onDemandWindow.remaining_value
|
||||||
|
|
||||||
|
return Object.keys(display).length > 0 ? display : null
|
||||||
|
}
|
||||||
|
|
||||||
|
function hasXaiQuotaDisplayData(key: EndpointAPIKey): boolean {
|
||||||
|
const xai = getXaiQuotaDisplay(key)
|
||||||
|
return !!xai && (
|
||||||
|
xai.usage_percentage !== undefined
|
||||||
|
|| xai.remaining_percentage !== undefined
|
||||||
|
|| xai.prepaid_balance !== undefined
|
||||||
|
|| xai.on_demand_cap !== undefined
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
function getXaiUsageLabel(key: EndpointAPIKey): string {
|
||||||
|
const display = getXaiQuotaDisplay(key)
|
||||||
|
if (display?.usage_label) return display.usage_label
|
||||||
|
const title = display?.subscription_title
|
||||||
|
return title ? `使用额度 (${title})` : '使用额度'
|
||||||
|
}
|
||||||
|
|
||||||
|
function getXaiUsedPercent(key: EndpointAPIKey): number {
|
||||||
|
return Math.min(Math.max(100 - getXaiRemainingPercent(key), 0), 100)
|
||||||
|
}
|
||||||
|
|
||||||
|
function getXaiRemainingPercent(key: EndpointAPIKey): number {
|
||||||
|
const xai = getXaiQuotaDisplay(key)
|
||||||
|
if (xai?.remaining_percentage != null && Number.isFinite(xai.remaining_percentage)) {
|
||||||
|
return Math.min(Math.max(xai.remaining_percentage, 0), 100)
|
||||||
|
}
|
||||||
|
if (xai?.usage_percentage != null && Number.isFinite(xai.usage_percentage)) {
|
||||||
|
return Math.min(Math.max(100 - xai.usage_percentage, 0), 100)
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
function getXaiOnDemandUsedPercent(key: EndpointAPIKey): number {
|
||||||
|
const xai = getXaiQuotaDisplay(key)
|
||||||
|
if (!xai?.on_demand_cap || xai.on_demand_cap <= 0) return 0
|
||||||
|
return Math.max(Math.min(((xai.on_demand_used || 0) / xai.on_demand_cap) * 100, 100), 0)
|
||||||
|
}
|
||||||
|
|
||||||
type GrokQuotaDisplay = GrokUpstreamMetadata & {
|
type GrokQuotaDisplay = GrokUpstreamMetadata & {
|
||||||
usage_percentage?: number
|
usage_percentage?: number
|
||||||
usage_limit?: number
|
usage_limit?: number
|
||||||
@@ -2696,6 +2840,28 @@ function shouldAutoRefreshGrokQuota(): boolean {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function shouldAutoRefreshXaiQuota(): boolean {
|
||||||
|
if (provider.value?.provider_type !== 'xai') return false
|
||||||
|
const now = Math.floor(Date.now() / 1000)
|
||||||
|
|
||||||
|
for (const { key } of allKeys.value) {
|
||||||
|
if (!key.is_active) continue
|
||||||
|
|
||||||
|
if (isTokenExpiringSoon(key, now)) return true
|
||||||
|
|
||||||
|
if (!hasXaiQuotaDisplayData(key)) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
const updatedAt = getXaiQuotaDisplay(key)?.updated_at
|
||||||
|
if (typeof updatedAt !== 'number' || (now - updatedAt) > AUTO_QUOTA_REFRESH_STALE_SECONDS) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
function shouldAutoRefreshWindsurfQuota(): boolean {
|
function shouldAutoRefreshWindsurfQuota(): boolean {
|
||||||
if (provider.value?.provider_type !== 'windsurf') return false
|
if (provider.value?.provider_type !== 'windsurf') return false
|
||||||
const now = Math.floor(Date.now() / 1000)
|
const now = Math.floor(Date.now() / 1000)
|
||||||
@@ -2824,7 +2990,7 @@ async function autoRefreshQuotaInBackground(): Promise<boolean> {
|
|||||||
if (refreshingQuota.value) return false
|
if (refreshingQuota.value) return false
|
||||||
|
|
||||||
const providerType = provider.value?.provider_type
|
const providerType = provider.value?.provider_type
|
||||||
if (providerType !== 'codex' && providerType !== 'gemini_cli' && providerType !== 'antigravity' && providerType !== 'kiro' && providerType !== 'windsurf' && providerType !== 'chatgpt_web' && providerType !== 'grok') return false
|
if (providerType !== 'codex' && providerType !== 'gemini_cli' && providerType !== 'antigravity' && providerType !== 'kiro' && providerType !== 'windsurf' && providerType !== 'chatgpt_web' && providerType !== 'grok' && providerType !== 'xai') return false
|
||||||
|
|
||||||
// 检查是否需要刷新
|
// 检查是否需要刷新
|
||||||
let shouldRefresh = false
|
let shouldRefresh = false
|
||||||
@@ -2838,6 +3004,8 @@ async function autoRefreshQuotaInBackground(): Promise<boolean> {
|
|||||||
shouldRefresh = shouldAutoRefreshKiroQuota()
|
shouldRefresh = shouldAutoRefreshKiroQuota()
|
||||||
} else if (providerType === 'grok') {
|
} else if (providerType === 'grok') {
|
||||||
shouldRefresh = shouldAutoRefreshGrokQuota()
|
shouldRefresh = shouldAutoRefreshGrokQuota()
|
||||||
|
} else if (providerType === 'xai') {
|
||||||
|
shouldRefresh = shouldAutoRefreshXaiQuota()
|
||||||
} else if (providerType === 'windsurf') {
|
} else if (providerType === 'windsurf') {
|
||||||
shouldRefresh = shouldAutoRefreshWindsurfQuota()
|
shouldRefresh = shouldAutoRefreshWindsurfQuota()
|
||||||
} else if (providerType === 'chatgpt_web') {
|
} else if (providerType === 'chatgpt_web') {
|
||||||
@@ -2856,6 +3024,8 @@ async function autoRefreshQuotaInBackground(): Promise<boolean> {
|
|||||||
hadCachedQuota = allKeys.value.some(({ key }) => key.is_active && hasKiroQuotaDisplayData(key))
|
hadCachedQuota = allKeys.value.some(({ key }) => key.is_active && hasKiroQuotaDisplayData(key))
|
||||||
} else if (providerType === 'grok') {
|
} else if (providerType === 'grok') {
|
||||||
hadCachedQuota = allKeys.value.some(({ key }) => key.is_active && hasGrokQuotaDisplayData(key))
|
hadCachedQuota = allKeys.value.some(({ key }) => key.is_active && hasGrokQuotaDisplayData(key))
|
||||||
|
} else if (providerType === 'xai') {
|
||||||
|
hadCachedQuota = allKeys.value.some(({ key }) => key.is_active && hasXaiQuotaDisplayData(key))
|
||||||
} else if (providerType === 'windsurf') {
|
} else if (providerType === 'windsurf') {
|
||||||
hadCachedQuota = allKeys.value.some(({ key }) => key.is_active && hasWindsurfQuotaDisplayData(key))
|
hadCachedQuota = allKeys.value.some(({ key }) => key.is_active && hasWindsurfQuotaDisplayData(key))
|
||||||
} else if (providerType === 'chatgpt_web') {
|
} else if (providerType === 'chatgpt_web') {
|
||||||
|
|||||||
@@ -60,6 +60,9 @@
|
|||||||
<SelectItem value="grok">
|
<SelectItem value="grok">
|
||||||
Grok
|
Grok
|
||||||
</SelectItem>
|
</SelectItem>
|
||||||
|
<SelectItem value="xai">
|
||||||
|
xAI
|
||||||
|
</SelectItem>
|
||||||
<SelectItem value="kiro">
|
<SelectItem value="kiro">
|
||||||
Kiro
|
Kiro
|
||||||
</SelectItem>
|
</SelectItem>
|
||||||
@@ -93,6 +96,9 @@
|
|||||||
<SelectItem value="grok">
|
<SelectItem value="grok">
|
||||||
Grok
|
Grok
|
||||||
</SelectItem>
|
</SelectItem>
|
||||||
|
<SelectItem value="xai">
|
||||||
|
xAI
|
||||||
|
</SelectItem>
|
||||||
<SelectItem value="kiro">
|
<SelectItem value="kiro">
|
||||||
Kiro
|
Kiro
|
||||||
</SelectItem>
|
</SelectItem>
|
||||||
|
|||||||
@@ -54,6 +54,24 @@ describe('provider quota display components', () => {
|
|||||||
unmount()
|
unmount()
|
||||||
})
|
})
|
||||||
|
|
||||||
|
it('fills the remaining bar even when used percent is zero', () => {
|
||||||
|
const { root, unmount } = mount(ProviderQuotaProgressRow, {
|
||||||
|
label: '周额度',
|
||||||
|
usedPercent: 0,
|
||||||
|
remainingPercent: 86,
|
||||||
|
meterClass: 'text-green-600',
|
||||||
|
barClass: 'bg-green-500',
|
||||||
|
resetText: '5天0小时后重置',
|
||||||
|
})
|
||||||
|
|
||||||
|
expect(root.querySelector('[data-testid="provider-quota-progress-meter"]')?.textContent?.trim()).toBe('86.0%')
|
||||||
|
expect((root.querySelector('[data-testid="provider-quota-progress-bar"]') as HTMLElement).style.width).toBe('86%')
|
||||||
|
expect(root.textContent).toContain('周额度')
|
||||||
|
expect(root.querySelector('[data-testid="provider-quota-progress-reset"]')?.textContent).toBe('5天0小时后重置')
|
||||||
|
|
||||||
|
unmount()
|
||||||
|
})
|
||||||
|
|
||||||
it('renders section loading and updated state', () => {
|
it('renders section loading and updated state', () => {
|
||||||
const Probe = defineComponent({
|
const Probe = defineComponent({
|
||||||
setup() {
|
setup() {
|
||||||
|
|||||||
@@ -1386,6 +1386,7 @@ function formatAuthType(authType: string): string {
|
|||||||
if (lowered === 'antigravity') return 'Antigravity OAuth'
|
if (lowered === 'antigravity') return 'Antigravity OAuth'
|
||||||
if (lowered === 'kiro') return 'Kiro OAuth'
|
if (lowered === 'kiro') return 'Kiro OAuth'
|
||||||
if (lowered === 'grok') return 'Grok OAuth'
|
if (lowered === 'grok') return 'Grok OAuth'
|
||||||
|
if (lowered === 'xai') return 'xAI OAuth'
|
||||||
return authType
|
return authType
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -39,10 +39,12 @@ const MODEL_TEST_OAUTH_INHERITS_PROVIDER_FORMATS = new Set([
|
|||||||
'vertex_ai',
|
'vertex_ai',
|
||||||
'antigravity',
|
'antigravity',
|
||||||
'kiro',
|
'kiro',
|
||||||
|
'xai',
|
||||||
])
|
])
|
||||||
|
|
||||||
const MODEL_TEST_BEARER_INHERITS_PROVIDER_FORMATS = new Set([
|
const MODEL_TEST_BEARER_INHERITS_PROVIDER_FORMATS = new Set([
|
||||||
'chatgpt_web',
|
'chatgpt_web',
|
||||||
|
'xai',
|
||||||
])
|
])
|
||||||
|
|
||||||
const MODEL_TEST_DIAGNOSTIC_LABELS: Record<string, string> = {
|
const MODEL_TEST_DIAGNOSTIC_LABELS: Record<string, string> = {
|
||||||
|
|||||||
@@ -16,6 +16,12 @@ describe('providerTypeUtils', () => {
|
|||||||
expect(isKeyManagedProviderType('grok')).toBe(false)
|
expect(isKeyManagedProviderType('grok')).toBe(false)
|
||||||
})
|
})
|
||||||
|
|
||||||
|
it('treats xAI as an OAuth account provider', () => {
|
||||||
|
expect(isOAuthAccountProviderType('xai')).toBe(true)
|
||||||
|
expect(isOAuthAccountProviderType('xAI')).toBe(true)
|
||||||
|
expect(isKeyManagedProviderType('xai')).toBe(false)
|
||||||
|
})
|
||||||
|
|
||||||
it('treats Windsurf as an OAuth account provider', () => {
|
it('treats Windsurf as an OAuth account provider', () => {
|
||||||
expect(isOAuthAccountProviderType('windsurf')).toBe(true)
|
expect(isOAuthAccountProviderType('windsurf')).toBe(true)
|
||||||
expect(isOAuthAccountProviderType('Windsurf')).toBe(true)
|
expect(isOAuthAccountProviderType('Windsurf')).toBe(true)
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ const oauthAccountProviderTypes = new Set([
|
|||||||
'antigravity',
|
'antigravity',
|
||||||
'kiro',
|
'kiro',
|
||||||
'grok',
|
'grok',
|
||||||
|
'xai',
|
||||||
'windsurf',
|
'windsurf',
|
||||||
])
|
])
|
||||||
|
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user