mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 10:57:03 +08:00
feat(providers): add xAI provider with device code OAuth
Add a separate `xai` provider type for xAI Grok CLI subscription accounts. It is independent of the existing `grok` provider, which reverse-proxies grok.com with browser cookies; behavior of `grok` is unchanged. Account binding uses the xAI device code flow, so no local callback listener is needed and headless deployments can bind accounts. Refresh tokens can also be imported individually or in batches, and are rotated on refresh. OAuth requests default to the cli-chat-proxy Responses API; API keys and compact stay on api.x.ai. Explicit custom gateways are preserved. Only `openai:responses` and `openai:responses:compact` are exposed; Chat, Claude and Gemini clients reach the provider through Aether's existing cross-format conversion rather than new native endpoints. Upstream Responses payloads are sanitized for what xAI actually rejects: `previous_response_id` and `metadata.user_id` are dropped, hosted `tool_choice` is rewritten, `web_search` is restored for converted clients, `image_generation` is stripped on older Grok conversation models, unsupported reasoning effort is removed, and requested `reasoning.encrypted_content` is preserved with a replay policy keyed on the configured provider type rather than the model name. Quota refresh reads /user and /billing?format=credits and stores a structured usage snapshot; a prepaid balance keeps an account selectable after the weekly allowance is exhausted. API-key accounts skip the subscription billing surface. The admin UI shows remaining weekly quota as a labeled bar in the provider drawer and the pool list. Co-Authored-By: Claude Opus 5 <[email protected]>
This commit is contained in:
@@ -3546,6 +3546,258 @@ pub fn parse_kiro_usage_response(
|
||||
Some(serde_json::Value::Object(result))
|
||||
}
|
||||
|
||||
pub fn parse_xai_billing_response(
|
||||
value: &serde_json::Value,
|
||||
updated_at_unix_secs: u64,
|
||||
) -> Option<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(
|
||||
value: &serde_json::Value,
|
||||
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 let Some(organization_uuid) = token_payload_object
|
||||
.get("organization")
|
||||
@@ -554,6 +571,28 @@ mod tests {
|
||||
assert_eq!(auth_config.get("is_fedramp"), Some(&json!(true)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn xai_enrichment_marks_oauth_and_extracts_id_token_identity() {
|
||||
let id_token = sample_unsigned_jwt(json!({
|
||||
"email": "[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]
|
||||
fn decode_jwt_claims_rejects_oversized_payload_before_decode() {
|
||||
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_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::{
|
||||
|
||||
@@ -7,6 +7,7 @@ pub mod request;
|
||||
pub mod response;
|
||||
pub mod spec;
|
||||
pub mod stream;
|
||||
pub mod xai;
|
||||
|
||||
const TOOL_ERROR_PREFIX: &str = "[tool error]";
|
||||
const AETHER_REASONING_ITEM_ID_PREFIX: &str = "rs_aether_";
|
||||
@@ -85,6 +86,8 @@ pub enum OpenAiResponsesReasoningReplayPolicy {
|
||||
#[default]
|
||||
OpenAiItemIds,
|
||||
DeepSeekOpaque,
|
||||
/// xAI replays encrypted state without requiring OpenAI's item-ID prefix.
|
||||
XaiEncrypted,
|
||||
}
|
||||
|
||||
/// Builds a stable, wire-compatible ID for a reasoning item synthesized by Aether.
|
||||
@@ -234,6 +237,14 @@ fn openai_responses_reasoning_item_is_replayable(
|
||||
{
|
||||
return true;
|
||||
}
|
||||
if policy == OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
|
||||
&& object
|
||||
.get("encrypted_content")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|value| !value.trim().is_empty())
|
||||
{
|
||||
return true;
|
||||
}
|
||||
let Some(id) = object
|
||||
.get("id")
|
||||
.and_then(Value::as_str)
|
||||
@@ -334,6 +345,36 @@ mod tests {
|
||||
OPENAI_RESPONSES_OPERATION_COMPACT,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn xai_encrypted_replay_accepts_native_ids_but_excludes_foreign_carriers() {
|
||||
let body = serde_json::json!({"input": [
|
||||
{"type": "reasoning", "id": "native-xai-id", "encrypted_content": "opaque-xai-state"},
|
||||
{"type": "reasoning", "encrypted_content": "opaque-idless-state"},
|
||||
{"type": "reasoning", "id": "rs_foreign", "encrypted_content": "cpa-gemini-responses-carrier-v1:foreign"},
|
||||
{"type": "reasoning", "id": "foreign-id", "summary": []}
|
||||
]});
|
||||
let mut xai = body.clone();
|
||||
assert_eq!(
|
||||
super::strip_incompatible_openai_responses_reasoning_items_with_policy(
|
||||
&mut xai,
|
||||
"openai:responses",
|
||||
super::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted,
|
||||
),
|
||||
2
|
||||
);
|
||||
assert_eq!(xai["input"].as_array().unwrap().len(), 2);
|
||||
assert_eq!(xai["input"][0], body["input"][0]);
|
||||
assert_eq!(xai["input"][1], body["input"][1]);
|
||||
let mut openai = body;
|
||||
assert_eq!(
|
||||
super::strip_incompatible_openai_responses_reasoning_items(
|
||||
&mut openai,
|
||||
"openai:responses"
|
||||
),
|
||||
4
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn gemini_tool_signature_carrier_roundtrips_direction_and_exact_value() {
|
||||
let signature = " opaque-signature-with-padding== ";
|
||||
|
||||
@@ -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-4.6",
|
||||
"prompt": "edit this",
|
||||
"image": {"image_url": "https://cdn.example/a.png"},
|
||||
"reference_images": [
|
||||
{"image_url": {"url": "https://cdn.example/b.png"}}
|
||||
],
|
||||
"input": [{
|
||||
"type": "message",
|
||||
"content": [{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "https://cdn.example/chat.png"}
|
||||
}]
|
||||
}]
|
||||
});
|
||||
|
||||
apply_xai_upstream_payload_edits(&mut body, "xai", "openai:responses");
|
||||
|
||||
assert_eq!(body["image"]["url"], "https://cdn.example/a.png");
|
||||
assert!(body["image"].get("image_url").is_none());
|
||||
assert_eq!(
|
||||
body["reference_images"][0]["url"],
|
||||
"https://cdn.example/b.png"
|
||||
);
|
||||
assert_eq!(
|
||||
body["input"][0]["content"][0]["image_url"]["url"],
|
||||
"https://cdn.example/chat.png"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn other_providers_are_left_untouched() {
|
||||
let mut body = json!({
|
||||
"previous_response_id": "resp_123",
|
||||
"image": {"image_url": "https://cdn.example/a.png"}
|
||||
});
|
||||
apply_xai_upstream_payload_edits(&mut body, "codex", "openai:responses");
|
||||
assert_eq!(body["previous_response_id"], "resp_123");
|
||||
assert_eq!(body["image"]["image_url"], "https://cdn.example/a.png");
|
||||
}
|
||||
}
|
||||
@@ -17,6 +17,7 @@ use crate::formats::openai::responses::codex::{
|
||||
apply_codex_openai_responses_chat_body_edits, apply_codex_openai_responses_special_body_edits,
|
||||
apply_openai_responses_compact_special_body_edits,
|
||||
};
|
||||
use crate::formats::openai::responses::xai::apply_xai_upstream_payload_edits_with_client;
|
||||
use crate::formats::shared::standard_normalize::{
|
||||
build_local_openai_chat_request_body_with_model_directives,
|
||||
is_claude_messages_shaped_body_on_openai_chat_endpoint,
|
||||
@@ -121,6 +122,11 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and
|
||||
enable_model_directives: bool,
|
||||
reasoning_replay_policy: crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy,
|
||||
) -> Option<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()
|
||||
.with_mapped_model(mapped_model)
|
||||
.with_request_path(request_path)
|
||||
@@ -133,13 +139,10 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and
|
||||
client_api_format,
|
||||
provider_api_format,
|
||||
);
|
||||
// DeepSeek's Responses continuation state is opaque. Parsing a same-wire-format
|
||||
// request through the canonical model would discard its id-less `reasoning_text`
|
||||
// items and future provider-owned fields even though no conversion is required.
|
||||
// Keep that provider-specific route wire-preserving, while retaining canonical
|
||||
// normalization for ordinary OpenAI Responses and for Responses/Compact
|
||||
// cross-format conversions.
|
||||
let mut provider_request_body = if is_wire_preserving_deepseek_responses_hop(
|
||||
// DeepSeek and xAI replay opaque provider state. Preserve their native
|
||||
// Responses input items: canonical conversion can lose reasoning IDs and
|
||||
// encrypted-only items even when source and destination formats are equal.
|
||||
let mut provider_request_body = if is_wire_preserving_responses_hop(
|
||||
source_api_format.as_ref(),
|
||||
provider_api_format,
|
||||
reasoning_replay_policy,
|
||||
@@ -200,6 +203,13 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and
|
||||
&mut provider_request_body,
|
||||
provider_api_format,
|
||||
);
|
||||
apply_xai_upstream_payload_edits_with_client(
|
||||
&mut provider_request_body,
|
||||
provider_type,
|
||||
provider_api_format,
|
||||
Some(client_api_format),
|
||||
Some(body_json),
|
||||
);
|
||||
crate::formats::openai::responses::strip_incompatible_openai_responses_reasoning_items_with_policy(
|
||||
&mut provider_request_body,
|
||||
provider_api_format,
|
||||
@@ -224,14 +234,16 @@ pub fn build_standard_request_body_with_model_directives_and_request_headers_and
|
||||
Some(provider_request_body)
|
||||
}
|
||||
|
||||
fn is_wire_preserving_deepseek_responses_hop(
|
||||
fn is_wire_preserving_responses_hop(
|
||||
source_api_format: &str,
|
||||
provider_api_format: &str,
|
||||
reasoning_replay_policy: crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy,
|
||||
) -> bool {
|
||||
if reasoning_replay_policy
|
||||
!= crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
|
||||
{
|
||||
if !matches!(
|
||||
reasoning_replay_policy,
|
||||
crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
|
||||
| crate::formats::openai::responses::OpenAiResponsesReasoningReplayPolicy::XaiEncrypted
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
let source_api_format = aether_ai_formats::normalize_api_format_alias(source_api_format);
|
||||
@@ -2077,4 +2089,316 @@ mod tests {
|
||||
);
|
||||
assert_eq!(gemini["toolConfig"]["functionCallingConfig"]["mode"], "ANY");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn xai_keeps_client_search_functions_distinct_from_hosted_search() {
|
||||
for name in ["web_search", "web_search_internal"] {
|
||||
for hosted in [false, true] {
|
||||
let mut tools = vec![json!({
|
||||
"name": name,
|
||||
"description": "Search internal documents",
|
||||
"input_schema": {"type": "object", "properties": {"query": {"type": "string"}}}
|
||||
})];
|
||||
if hosted {
|
||||
tools.push(json!({"type": "web_search_20260209", "name": "internet_search"}));
|
||||
}
|
||||
let request = json!({
|
||||
"model": "source", "max_tokens": 64,
|
||||
"messages": [{"role": "user", "content": "Search internal documents"}],
|
||||
"tools": tools,
|
||||
"tool_choice": {"type": "tool", "name": name}
|
||||
});
|
||||
let converted = build_standard_request_body(
|
||||
&request,
|
||||
"claude:messages",
|
||||
"grok-4.6",
|
||||
"xai",
|
||||
"openai:responses",
|
||||
"/v1/messages",
|
||||
true,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
converted["tool_choice"],
|
||||
json!({"type": "function", "name": name})
|
||||
);
|
||||
assert_eq!(
|
||||
converted["tools"]
|
||||
.as_array()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.any(|tool| tool["type"] == "web_search"),
|
||||
hosted
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let request = json!({
|
||||
"model": "source", "max_tokens": 64,
|
||||
"messages": [{"role": "user", "content": "Search the internet"}],
|
||||
"tools": [{"type": "web_search_20260209", "name": "internet_search"}],
|
||||
"tool_choice": {"type": "tool", "name": "internet_search"}
|
||||
});
|
||||
let converted = build_standard_request_body(
|
||||
&request,
|
||||
"claude:messages",
|
||||
"grok-4.6",
|
||||
"xai",
|
||||
"openai:responses",
|
||||
"/v1/messages",
|
||||
true,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
converted["tool_choice"],
|
||||
json!({
|
||||
"type": "allowed_tools", "mode": "required", "tools": [{"type": "web_search"}]
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn xai_preserves_function_choices_in_chat_and_responses_requests() {
|
||||
for name in ["web_search", "web_search_internal"] {
|
||||
for (client, request) in [
|
||||
(
|
||||
"openai:chat",
|
||||
json!({
|
||||
"messages": [{"role": "user", "content": "search"}],
|
||||
"tools": [{"type": "function", "function": {"name": name, "parameters": {"type": "object"}}}],
|
||||
"tool_choice": {"type": "function", "function": {"name": name}}
|
||||
}),
|
||||
),
|
||||
(
|
||||
"openai:responses",
|
||||
json!({
|
||||
"input": "search",
|
||||
"tools": [{"type": "function", "name": name, "parameters": {"type": "object"}}],
|
||||
"tool_choice": {"type": "function", "name": name}
|
||||
}),
|
||||
),
|
||||
] {
|
||||
let converted = build_standard_request_body(
|
||||
&request,
|
||||
client,
|
||||
"grok-4.6",
|
||||
"xai",
|
||||
"openai:responses",
|
||||
"/v1/responses",
|
||||
true,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
converted["tool_choice"],
|
||||
json!({"type": "function", "name": name})
|
||||
);
|
||||
assert_eq!(converted["tools"].as_array().unwrap().len(), 1);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn xai_image_allowed_tools_preserves_mode_and_restricts_available_tools() {
|
||||
for mode in ["auto", "required"] {
|
||||
for mixed in [false, true] {
|
||||
let mut allowed = vec![json!({"type": "image_generation"})];
|
||||
if mixed {
|
||||
allowed.push(json!({"type": "function", "name": "lookup"}));
|
||||
}
|
||||
let request = json!({
|
||||
"input": "Draw a cat",
|
||||
"tools": [
|
||||
{"type": "web_search"}, {"type": "image_generation"},
|
||||
{"type": "function", "name": "lookup", "parameters": {"type": "object"}}
|
||||
],
|
||||
"tool_choice": {"type": "allowed_tools", "mode": mode, "tools": allowed}
|
||||
});
|
||||
let converted = build_standard_request_body(
|
||||
&request,
|
||||
"openai:responses",
|
||||
"grok-4.6",
|
||||
"xai",
|
||||
"openai:responses",
|
||||
"/v1/responses",
|
||||
true,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.unwrap();
|
||||
if mixed {
|
||||
assert_eq!(
|
||||
converted["tool_choice"],
|
||||
json!({
|
||||
"type": "allowed_tools", "mode": mode,
|
||||
"tools": [{"type": "function", "name": "lookup"}]
|
||||
})
|
||||
);
|
||||
assert_eq!(converted["tools"].as_array().unwrap().len(), 3);
|
||||
} else {
|
||||
assert_eq!(converted["tool_choice"], mode);
|
||||
assert_eq!(converted["tools"], json!([{"type": "image_generation"}]));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn xai_responses_preserves_requested_encrypted_reasoning_and_replayed_input() {
|
||||
let reasoning = json!({"type": "reasoning", "id": "550e8400-e29b-41d4-a716-446655440000", "summary": [], "encrypted_content": "opaque-xai-state"});
|
||||
let request = json!({
|
||||
"input": [reasoning.clone(), {"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "Previous answer"}]}, {"role": "user", "content": "Continue"}],
|
||||
"include": ["reasoning.encrypted_content"], "store": false
|
||||
});
|
||||
let converted = build_standard_request_body(
|
||||
&request,
|
||||
"openai:responses",
|
||||
"grok-4.6",
|
||||
"xai",
|
||||
"openai:responses",
|
||||
"/v1/responses",
|
||||
true,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(converted["include"], request["include"]);
|
||||
assert_eq!(converted["input"][0], reasoning);
|
||||
assert_eq!(converted["store"], false);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn xai_standard_conversion_strips_unsupported_responses_fields() {
|
||||
let request = json!({
|
||||
"model": "source-model",
|
||||
"messages": [{"role": "user", "content": "Hello xAI"}],
|
||||
"max_tokens": 128,
|
||||
"stop": ["END"],
|
||||
"stream_options": {"include_usage": true},
|
||||
"metadata": {"user_id": "claude-session"},
|
||||
"web_search_options": {"search_context_size": "high"}
|
||||
});
|
||||
let converted = build_standard_request_body(
|
||||
&request,
|
||||
"openai:chat",
|
||||
"grok-4.6",
|
||||
"xai",
|
||||
"openai:responses",
|
||||
"/v1/chat/completions",
|
||||
true,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("chat should convert onto xAI Responses");
|
||||
|
||||
assert_eq!(converted["model"], "grok-4.6");
|
||||
assert!(converted.get("stop").is_none());
|
||||
assert!(converted.get("stream_options").is_none());
|
||||
assert!(converted.get("previous_response_id").is_none());
|
||||
assert!(converted.get("metadata").is_none());
|
||||
assert!(converted.get("input").is_some() || converted.get("messages").is_none());
|
||||
assert_eq!(converted["max_output_tokens"], 128);
|
||||
assert_eq!(converted["tools"][0]["type"], "web_search");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn xai_standard_conversion_covers_claude_and_gemini_clients() {
|
||||
let claude = json!({
|
||||
"model": "claude-sonnet",
|
||||
"max_tokens": 64,
|
||||
"messages": [{"role": "user", "content": "Hello xAI"}],
|
||||
"metadata": {
|
||||
"user_id": "{\"device_id\":\"dev-1\",\"account_uuid\":\"acct-1\",\"session_id\":\"sess-1\"}"
|
||||
},
|
||||
"tools": [
|
||||
{"type": "web_search_20250305", "name": "web_search"},
|
||||
{
|
||||
"name": "lookup",
|
||||
"description": "Look something up",
|
||||
"input_schema": {"type": "object", "properties": {}}
|
||||
}
|
||||
],
|
||||
"tool_choice": {"type": "tool", "name": "web_search"}
|
||||
});
|
||||
let converted = build_standard_request_body(
|
||||
&claude,
|
||||
"claude:messages",
|
||||
"grok-4.6",
|
||||
"xai",
|
||||
"openai:responses",
|
||||
"/v1/messages",
|
||||
true,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("claude should convert onto xAI Responses");
|
||||
assert_eq!(converted["model"], "grok-4.6");
|
||||
assert!(converted.get("metadata").is_none());
|
||||
assert!(converted.get("context_management").is_none());
|
||||
assert!(converted
|
||||
.get("include")
|
||||
.and_then(Value::as_array)
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.any(|item| item == "reasoning.encrypted_content"));
|
||||
assert!(converted["tools"]
|
||||
.as_array()
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.any(|tool| tool["type"] == "web_search"));
|
||||
assert_eq!(converted["tool_choice"]["type"], "allowed_tools");
|
||||
assert!(converted.get("input").is_some());
|
||||
|
||||
let gemini = json!({
|
||||
"model": "gemini-2.5-pro",
|
||||
"contents": [{
|
||||
"role": "user",
|
||||
"parts": [{"text": "Hello xAI"}]
|
||||
}],
|
||||
"tools": [{"googleSearch": {}}]
|
||||
});
|
||||
let converted = build_standard_request_body(
|
||||
&gemini,
|
||||
"gemini:generate_content",
|
||||
"grok-4.6",
|
||||
"xai",
|
||||
"openai:responses",
|
||||
"/v1beta/models/gemini-2.5-pro:generateContent",
|
||||
false,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("gemini should convert onto xAI Responses");
|
||||
assert_eq!(converted["model"], "grok-4.6");
|
||||
assert_eq!(converted["tools"][0]["type"], "web_search");
|
||||
assert!(converted.get("input").is_some());
|
||||
|
||||
let same_format = json!({
|
||||
"model": "grok-4.6",
|
||||
"input": "hello",
|
||||
"previous_response_id": "resp_123",
|
||||
"stop": ["END"],
|
||||
"metadata": {"user_id": "claude-session"}
|
||||
});
|
||||
let converted = build_standard_request_body(
|
||||
&same_format,
|
||||
"openai:responses",
|
||||
"grok-4.6",
|
||||
"xai",
|
||||
"openai:responses",
|
||||
"/v1/responses",
|
||||
true,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("same-format xAI Responses should sanitize in place");
|
||||
assert!(converted.get("previous_response_id").is_none());
|
||||
assert!(converted.get("stop").is_none());
|
||||
assert!(converted.get("metadata").is_none());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -56,6 +56,10 @@ pub use formats::openai::responses::codex::{
|
||||
pub use formats::openai::responses::request::{
|
||||
validate_openai_responses_request_contract, OpenAiResponsesRequestContractViolation,
|
||||
};
|
||||
pub use formats::openai::responses::xai::{
|
||||
apply_xai_upstream_payload_edits, apply_xai_upstream_payload_edits_with_client,
|
||||
xai_model_supports_reasoning_effort, xai_supports_native_image_generation,
|
||||
};
|
||||
pub use formats::openai::responses::{
|
||||
normalize_openai_responses_message_item_ids, openai_responses_message_item_id,
|
||||
openai_responses_request_operation, openai_responses_synthetic_reasoning_item_id,
|
||||
|
||||
@@ -102,6 +102,11 @@ INNER JOIN LATERAL (
|
||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||
AND LOWER($3) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
|
||||
)
|
||||
OR (
|
||||
LOWER(BTRIM(p.provider_type)) = 'xai'
|
||||
AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')
|
||||
AND LOWER($3) IN ('openai:responses', 'openai:responses:compact')
|
||||
)
|
||||
OR (
|
||||
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||
@@ -127,7 +132,8 @@ INNER JOIN LATERAL (
|
||||
'vertex_ai',
|
||||
'antigravity',
|
||||
'kiro',
|
||||
'windsurf'
|
||||
'windsurf',
|
||||
'xai'
|
||||
)
|
||||
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
||||
)
|
||||
@@ -187,6 +193,11 @@ WHERE p.is_active = TRUE
|
||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||
AND LOWER($3) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
|
||||
)
|
||||
OR (
|
||||
LOWER(BTRIM(p.provider_type)) = 'xai'
|
||||
AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')
|
||||
AND LOWER($3) IN ('openai:responses', 'openai:responses:compact')
|
||||
)
|
||||
OR (
|
||||
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||
@@ -212,7 +223,8 @@ WHERE p.is_active = TRUE
|
||||
'vertex_ai',
|
||||
'antigravity',
|
||||
'kiro',
|
||||
'windsurf'
|
||||
'windsurf',
|
||||
'xai'
|
||||
)
|
||||
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
||||
)
|
||||
@@ -365,6 +377,11 @@ INNER JOIN LATERAL (
|
||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||
AND LOWER($4) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
|
||||
)
|
||||
OR (
|
||||
LOWER(BTRIM(p.provider_type)) = 'xai'
|
||||
AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')
|
||||
AND LOWER($4) IN ('openai:responses', 'openai:responses:compact')
|
||||
)
|
||||
OR (
|
||||
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||
@@ -390,7 +407,8 @@ INNER JOIN LATERAL (
|
||||
'vertex_ai',
|
||||
'antigravity',
|
||||
'kiro',
|
||||
'windsurf'
|
||||
'windsurf',
|
||||
'xai'
|
||||
)
|
||||
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
||||
)
|
||||
@@ -451,6 +469,11 @@ WHERE p.is_active = TRUE
|
||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||
AND LOWER($4) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
|
||||
)
|
||||
OR (
|
||||
LOWER(BTRIM(p.provider_type)) = 'xai'
|
||||
AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')
|
||||
AND LOWER($4) IN ('openai:responses', 'openai:responses:compact')
|
||||
)
|
||||
OR (
|
||||
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||
@@ -476,7 +499,8 @@ WHERE p.is_active = TRUE
|
||||
'vertex_ai',
|
||||
'antigravity',
|
||||
'kiro',
|
||||
'windsurf'
|
||||
'windsurf',
|
||||
'xai'
|
||||
)
|
||||
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
||||
)
|
||||
@@ -637,6 +661,11 @@ WHERE p.is_active = TRUE
|
||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||
AND LOWER($6) IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image')
|
||||
)
|
||||
OR (
|
||||
LOWER(BTRIM(p.provider_type)) = 'xai'
|
||||
AND LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')
|
||||
AND LOWER($6) IN ('openai:responses', 'openai:responses:compact')
|
||||
)
|
||||
OR (
|
||||
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||
@@ -662,7 +691,8 @@ WHERE p.is_active = TRUE
|
||||
'vertex_ai',
|
||||
'antigravity',
|
||||
'kiro',
|
||||
'windsurf'
|
||||
'windsurf',
|
||||
'xai'
|
||||
)
|
||||
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
||||
)
|
||||
@@ -1717,6 +1747,22 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn candidate_selection_sql_allows_xai_oauth_responses_auth() {
|
||||
let requested_model_sql = requested_model_selection_sql();
|
||||
for sql in [
|
||||
LIST_FOR_EXACT_API_FORMAT_SQL,
|
||||
LIST_FOR_EXACT_API_FORMAT_AND_GLOBAL_MODEL_SQL,
|
||||
LIST_POOL_KEYS_FOR_GROUP_SQL,
|
||||
requested_model_sql.as_str(),
|
||||
] {
|
||||
assert!(sql.contains("LOWER(BTRIM(p.provider_type)) = 'xai'"));
|
||||
assert!(sql.contains("LOWER(BTRIM(pak.auth_type)) IN ('oauth', 'bearer', 'api_key')"));
|
||||
assert!(sql.contains("'openai:responses', 'openai:responses:compact'"));
|
||||
assert!(sql.contains("'xai'"));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn candidate_selection_sql_allows_windsurf_openai_chat_managed_keys() {
|
||||
let requested_model_sql = requested_model_selection_sql();
|
||||
|
||||
@@ -346,6 +346,13 @@ fn key_auth_channel_matches(row: &StoredMinimalCandidateSelectionRow, api_format
|
||||
"openai:chat" | "openai:responses" | "claude:messages" | "openai:image"
|
||||
)
|
||||
}
|
||||
"xai" => {
|
||||
matches!(auth_type.as_str(), "oauth" | "bearer" | "api_key")
|
||||
&& matches!(
|
||||
api_format.as_str(),
|
||||
"openai:responses" | "openai:responses:compact"
|
||||
)
|
||||
}
|
||||
"windsurf" => {
|
||||
matches!(auth_type.as_str(), "oauth" | "api_key" | "bearer")
|
||||
&& api_format == "openai:chat"
|
||||
@@ -591,6 +598,28 @@ mod tests {
|
||||
assert_eq!(rows[0].global_model_name, "grok-4.20-0309-non-reasoning");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn includes_xai_oauth_rows_for_responses_models() {
|
||||
let mut row = sample_row("provider-xai", "openai:responses", "grok-4", 10);
|
||||
row.provider_type = "xai".to_string();
|
||||
row.provider_name = "xai".to_string();
|
||||
row.key_auth_type = "oauth".to_string();
|
||||
row.key_api_formats = Some(vec![
|
||||
"openai:responses".to_string(),
|
||||
"openai:responses:compact".to_string(),
|
||||
]);
|
||||
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![row]);
|
||||
|
||||
let rows = repository
|
||||
.list_for_exact_api_format("openai:responses")
|
||||
.await
|
||||
.expect("list should succeed");
|
||||
|
||||
assert_eq!(rows.len(), 1);
|
||||
assert_eq!(rows[0].provider_type, "xai");
|
||||
assert_eq!(rows[0].global_model_name, "grok-4");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn requested_model_filter_respects_endpoint_scoped_default_mapping() {
|
||||
let mut selected = sample_row("provider-1", "openai:chat", "deepseek-v4-pro", 10);
|
||||
|
||||
@@ -546,7 +546,7 @@ pub fn endpoint_supports_rust_models_fetch(api_format: &str) -> bool {
|
||||
pub fn provider_type_uses_preset_models(provider_type: &str) -> bool {
|
||||
matches!(
|
||||
provider_type.trim().to_ascii_lowercase().as_str(),
|
||||
"claude_code" | "gemini_cli" | "grok"
|
||||
"claude_code" | "gemini_cli" | "grok" | "xai"
|
||||
)
|
||||
}
|
||||
|
||||
@@ -604,6 +604,18 @@ pub fn preset_models_for_provider(provider_type: &str) -> Option<Vec<Value>> {
|
||||
preset_model("grok-imagine-image-pro", "xai", "Grok Imagine Image Pro", "openai:image"),
|
||||
preset_model("grok-imagine-image-edit", "xai", "Grok Imagine Image Edit", "openai:image"),
|
||||
],
|
||||
"xai" => vec![
|
||||
preset_model("grok-4.6", "xai", "Grok 4.6", "openai:responses"),
|
||||
preset_model("grok-build-0.1", "xai", "Grok Build 0.1", "openai:responses"),
|
||||
preset_model("grok-4.5", "xai", "Grok 4.5", "openai:responses"),
|
||||
preset_model("grok-4.3", "xai", "Grok 4.3", "openai:responses"),
|
||||
preset_model("grok-4.20-0309-reasoning", "xai", "Grok 4.20 0309 Reasoning", "openai:responses"),
|
||||
preset_model("grok-4.20-0309-non-reasoning", "xai", "Grok 4.20 0309 Non-Reasoning", "openai:responses"),
|
||||
preset_model("grok-4.20-multi-agent-0309", "xai", "Grok 4.20 Multi-Agent 0309", "openai:responses"),
|
||||
preset_model("grok-3-mini", "xai", "Grok 3 Mini", "openai:responses"),
|
||||
preset_model("grok-3-mini-fast", "xai", "Grok 3 Mini Fast", "openai:responses"),
|
||||
preset_model("grok-composer-2.5-fast", "xai", "Grok Composer 2.5 Fast", "openai:responses"),
|
||||
],
|
||||
_ => return None,
|
||||
};
|
||||
Some(models)
|
||||
@@ -1977,4 +1989,32 @@ mod tests {
|
||||
assert_eq!(models[15]["api_formats"], json!(["openai:image"]));
|
||||
assert_eq!(models[18]["api_formats"], json!(["openai:image"]));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn preset_models_cover_xai_cli_catalog() {
|
||||
let models = preset_models_for_provider("xai").expect("preset models should exist");
|
||||
let model_ids = models
|
||||
.iter()
|
||||
.map(|model| model["id"].as_str().expect("model id"))
|
||||
.collect::<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",
|
||||
]
|
||||
);
|
||||
assert!(models.iter().all(|model| model["owned_by"] == json!("xai")));
|
||||
assert!(models
|
||||
.iter()
|
||||
.all(|model| model["api_formats"] == json!(["openai:responses"])));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -150,6 +150,27 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[
|
||||
uses_json_payload: false,
|
||||
include_scope_in_token_request: true,
|
||||
},
|
||||
GenericProviderOAuthTemplate {
|
||||
provider_type: "xai",
|
||||
display_name: "xAI",
|
||||
authorize_url: "https://auth.x.ai/oauth2/device/code",
|
||||
token_url: "https://auth.x.ai/oauth2/token",
|
||||
client_id: "b1a00492-073a-47ea-816f-4c329264a828",
|
||||
client_id_env: None,
|
||||
client_secret_env: None,
|
||||
scopes: &[
|
||||
"openid",
|
||||
"profile",
|
||||
"email",
|
||||
"offline_access",
|
||||
"grok-cli:access",
|
||||
"api:access",
|
||||
],
|
||||
redirect_uri: "",
|
||||
use_pkce: false,
|
||||
uses_json_payload: false,
|
||||
include_scope_in_token_request: false,
|
||||
},
|
||||
];
|
||||
|
||||
#[derive(Clone)]
|
||||
@@ -212,6 +233,10 @@ impl GenericProviderOAuthAdapter {
|
||||
self
|
||||
}
|
||||
|
||||
pub(super) fn token_url_for_provider(&self) -> String {
|
||||
self.token_url()
|
||||
}
|
||||
|
||||
fn token_url(&self) -> String {
|
||||
self.token_url_override
|
||||
.clone()
|
||||
@@ -389,7 +414,10 @@ impl GenericProviderOAuthAdapter {
|
||||
self.token_set_from_payload(payload)
|
||||
}
|
||||
|
||||
fn token_set_from_payload(&self, payload: Value) -> Result<ProviderOAuthTokenSet, OAuthError> {
|
||||
pub(super) fn token_set_from_payload(
|
||||
&self,
|
||||
payload: Value,
|
||||
) -> Result<ProviderOAuthTokenSet, OAuthError> {
|
||||
let token_set = OAuthTokenSet::from_token_payload(payload.clone())
|
||||
.ok_or_else(|| OAuthError::invalid_response("token response missing access_token"))?;
|
||||
let mut auth_config = serde_json::Map::new();
|
||||
@@ -945,6 +973,7 @@ mod tests {
|
||||
fn resolves_generic_provider_templates() {
|
||||
assert!(template_for_provider_type("codex").is_some());
|
||||
assert!(template_for_provider_type("claude_code").is_some());
|
||||
assert!(template_for_provider_type("xai").is_some());
|
||||
assert!(template_for_provider_type("kiro").is_none());
|
||||
}
|
||||
|
||||
|
||||
@@ -4,6 +4,7 @@ mod codex;
|
||||
mod generic;
|
||||
mod kiro;
|
||||
mod windsurf;
|
||||
mod xai;
|
||||
|
||||
pub use antigravity::{AntigravityProviderOAuthAdapter, ANTIGRAVITY_USER_INFO_URL};
|
||||
pub use claude_code::{
|
||||
@@ -27,3 +28,7 @@ pub use windsurf::{
|
||||
WindsurfProviderOAuthAdapter, WINDSURF_CLIENT_ID, WINDSURF_PROVIDER_TYPE,
|
||||
WINDSURF_SHOW_AUTH_TOKEN_REDIRECT, WINDSURF_SIGNIN_URL,
|
||||
};
|
||||
pub use xai::{
|
||||
XaiDevicePollOutcome, XaiProviderOAuthAdapter, XAI_CLIENT_ID, XAI_DEVICE_CODE_GRANT_TYPE,
|
||||
XAI_DEVICE_CODE_URL, XAI_OAUTH_SCOPES, XAI_PROVIDER_TYPE, XAI_TOKEN_URL,
|
||||
};
|
||||
|
||||
@@ -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::{
|
||||
AntigravityProviderOAuthAdapter, ClaudeCodeProviderOAuthAdapter,
|
||||
CodexProviderOAuthAdapter, GenericProviderOAuthAdapter, KiroProviderOAuthAdapter,
|
||||
WindsurfProviderOAuthAdapter,
|
||||
WindsurfProviderOAuthAdapter, XaiProviderOAuthAdapter,
|
||||
};
|
||||
|
||||
let mut service = Self::new()
|
||||
@@ -29,7 +29,8 @@ impl ProviderOAuthService {
|
||||
.with_adapter(Arc::new(ClaudeCodeProviderOAuthAdapter::default()))
|
||||
.with_adapter(Arc::new(CodexProviderOAuthAdapter::default()))
|
||||
.with_adapter(Arc::new(AntigravityProviderOAuthAdapter::default()))
|
||||
.with_adapter(Arc::new(WindsurfProviderOAuthAdapter));
|
||||
.with_adapter(Arc::new(WindsurfProviderOAuthAdapter))
|
||||
.with_adapter(Arc::new(XaiProviderOAuthAdapter::default()));
|
||||
for provider_type in ["chatgpt_web", "gemini_cli"] {
|
||||
if let Some(adapter) = GenericProviderOAuthAdapter::for_provider_type(provider_type) {
|
||||
service = service.with_adapter(Arc::new(adapter));
|
||||
@@ -144,6 +145,7 @@ mod tests {
|
||||
"antigravity",
|
||||
"kiro",
|
||||
"windsurf",
|
||||
"xai",
|
||||
] {
|
||||
assert!(
|
||||
service.adapter(provider_type).is_ok(),
|
||||
|
||||
@@ -22,18 +22,20 @@ pub use providers::{
|
||||
build_windsurf_pool_model_configs_request,
|
||||
build_windsurf_pool_model_configs_request_with_base_url, build_windsurf_pool_quota_request,
|
||||
build_windsurf_pool_quota_request_with_base_url, build_windsurf_pool_rate_limit_request,
|
||||
build_windsurf_pool_rate_limit_request_with_base_url, enrich_chatgpt_web_quota_metadata,
|
||||
grok_mode_id_for_model, grok_pool_tier_from_quota_bucket, grok_quota_window_key_for_model,
|
||||
build_windsurf_pool_rate_limit_request_with_base_url, build_xai_pool_billing_request,
|
||||
build_xai_pool_user_request, enrich_chatgpt_web_quota_metadata, grok_mode_id_for_model,
|
||||
grok_pool_tier_from_quota_bucket, grok_quota_window_key_for_model,
|
||||
grok_supported_quota_windows_for_tier, normalize_chatgpt_web_image_quota_limit,
|
||||
AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, CodexProviderPoolAdapter,
|
||||
DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, GrokProviderPoolAdapter,
|
||||
KiroPoolQuotaAuthInput, KiroProviderPoolAdapter, UnsupportedQuotaProviderPoolAdapter,
|
||||
ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH, ANTIGRAVITY_RETRIEVE_USER_QUOTA_SUMMARY_PATH,
|
||||
CHATGPT_WEB_CONVERSATION_INIT_PATH, CHATGPT_WEB_DEFAULT_BASE_URL,
|
||||
CODEX_WHAM_RESET_CREDITS_CONSUME_URL, CODEX_WHAM_RESET_CREDITS_URL, CODEX_WHAM_USAGE_URL,
|
||||
GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH, GEMINI_CLI_USER_AGENT, KIRO_USAGE_LIMITS_PATH,
|
||||
KIRO_USAGE_SDK_VERSION, WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_PATH,
|
||||
WINDSURF_USER_STATUS_PATH,
|
||||
XaiProviderPoolAdapter, ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH,
|
||||
ANTIGRAVITY_RETRIEVE_USER_QUOTA_SUMMARY_PATH, CHATGPT_WEB_CONVERSATION_INIT_PATH,
|
||||
CHATGPT_WEB_DEFAULT_BASE_URL, CODEX_WHAM_RESET_CREDITS_CONSUME_URL,
|
||||
CODEX_WHAM_RESET_CREDITS_URL, CODEX_WHAM_USAGE_URL, GEMINI_CLI_RETRIEVE_USER_QUOTA_PATH,
|
||||
GEMINI_CLI_USER_AGENT, KIRO_USAGE_LIMITS_PATH, KIRO_USAGE_SDK_VERSION,
|
||||
WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_PATH, WINDSURF_USER_STATUS_PATH,
|
||||
XAI_BILLING_PATH, XAI_USER_PATH,
|
||||
};
|
||||
pub use quota::{
|
||||
provider_pool_key_account_quota_exhausted, provider_pool_key_model_quota_exhausted,
|
||||
@@ -81,7 +83,8 @@ mod tests {
|
||||
"grok",
|
||||
"kiro",
|
||||
"vertex_ai",
|
||||
"windsurf"
|
||||
"windsurf",
|
||||
"xai"
|
||||
]
|
||||
);
|
||||
assert!(service
|
||||
@@ -104,7 +107,8 @@ mod tests {
|
||||
"gemini_cli",
|
||||
"grok",
|
||||
"kiro",
|
||||
"windsurf"
|
||||
"windsurf",
|
||||
"xai"
|
||||
]
|
||||
);
|
||||
assert!(service.supports_quota_refresh("codex"));
|
||||
@@ -112,6 +116,7 @@ mod tests {
|
||||
assert!(service.supports_quota_refresh("grok"));
|
||||
assert!(service.supports_quota_refresh("gemini_cli"));
|
||||
assert!(service.supports_quota_refresh("windsurf"));
|
||||
assert!(service.supports_quota_refresh("xai"));
|
||||
assert_eq!(
|
||||
service.quota_refresh_unsupported_message("claude_code"),
|
||||
"Claude Code 暂不支持自动刷新额度:上游没有稳定可用的账号额度查询接口"
|
||||
@@ -642,11 +647,11 @@ mod tests {
|
||||
|
||||
assert_eq!(
|
||||
free_first["providers"],
|
||||
json!(["codex", "grok", "kiro", "windsurf"])
|
||||
json!(["codex", "grok", "kiro", "windsurf", "xai"])
|
||||
);
|
||||
assert_eq!(
|
||||
recent_refresh["providers"],
|
||||
json!(["codex", "grok", "kiro", "windsurf"])
|
||||
json!(["codex", "grok", "kiro", "windsurf", "xai"])
|
||||
);
|
||||
assert_eq!(free_first["default_enabled"], json!(false));
|
||||
assert_eq!(recent_refresh["default_enabled"], json!(false));
|
||||
|
||||
@@ -7,6 +7,7 @@ pub mod grok;
|
||||
pub mod kiro;
|
||||
pub mod unsupported;
|
||||
pub mod windsurf;
|
||||
pub mod xai;
|
||||
|
||||
pub use antigravity::AntigravityProviderPoolAdapter;
|
||||
pub use antigravity::{
|
||||
@@ -51,3 +52,7 @@ pub use windsurf::{
|
||||
WINDSURF_DEFAULT_BASE_URL, WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_PATH,
|
||||
WINDSURF_USER_STATUS_PATH,
|
||||
};
|
||||
pub use xai::{
|
||||
build_xai_pool_billing_request, build_xai_pool_user_request, XaiProviderPoolAdapter,
|
||||
XAI_BILLING_PATH, XAI_USER_PATH,
|
||||
};
|
||||
|
||||
@@ -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::{
|
||||
AntigravityProviderPoolAdapter, ChatGptWebProviderPoolAdapter, CodexProviderPoolAdapter,
|
||||
DefaultProviderPoolAdapter, GeminiCliProviderPoolAdapter, GrokProviderPoolAdapter,
|
||||
KiroProviderPoolAdapter, WindsurfProviderPoolAdapter, CLAUDE_CODE_PROVIDER_POOL_ADAPTER,
|
||||
VERTEX_AI_PROVIDER_POOL_ADAPTER,
|
||||
KiroProviderPoolAdapter, WindsurfProviderPoolAdapter, XaiProviderPoolAdapter,
|
||||
CLAUDE_CODE_PROVIDER_POOL_ADAPTER, VERTEX_AI_PROVIDER_POOL_ADAPTER,
|
||||
};
|
||||
|
||||
#[derive(Clone)]
|
||||
@@ -55,6 +55,7 @@ impl ProviderPoolService {
|
||||
.with_adapter(Arc::new(KiroProviderPoolAdapter))
|
||||
.with_adapter(Arc::new(ChatGptWebProviderPoolAdapter))
|
||||
.with_adapter(Arc::new(WindsurfProviderPoolAdapter))
|
||||
.with_adapter(Arc::new(XaiProviderPoolAdapter))
|
||||
.with_adapter(Arc::new(VERTEX_AI_PROVIDER_POOL_ADAPTER))
|
||||
}
|
||||
|
||||
|
||||
@@ -754,6 +754,63 @@ mod tests {
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn xai_responses_transport_converts_standard_client_protocols() {
|
||||
let transport = transport_snapshot("xai", "openai:responses", "oauth", true, None);
|
||||
|
||||
for client_api_format in ["openai:chat", "claude:messages", "gemini:generate_content"] {
|
||||
assert!(
|
||||
request_pair_allowed_for_transport(
|
||||
&transport,
|
||||
client_api_format,
|
||||
"openai:responses"
|
||||
),
|
||||
"{client_api_format} should convert onto xAI Responses"
|
||||
);
|
||||
assert_eq!(
|
||||
candidate_transport_pair_skip_reason(&transport, client_api_format),
|
||||
None
|
||||
);
|
||||
}
|
||||
assert!(request_conversion_transport_supported(
|
||||
&transport,
|
||||
RequestConversionKind::ToOpenAiResponses
|
||||
));
|
||||
assert!(
|
||||
!request_pair_allowed_for_transport(
|
||||
&transport,
|
||||
"openai:responses:compact",
|
||||
"openai:responses"
|
||||
),
|
||||
"compact must not convert onto xAI Responses"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn xai_compact_endpoint_is_same_format_only() {
|
||||
let compact = transport_snapshot("xai", "openai:responses:compact", "oauth", true, None);
|
||||
assert!(request_pair_allowed_for_transport(
|
||||
&compact,
|
||||
"openai:responses:compact",
|
||||
"openai:responses:compact"
|
||||
));
|
||||
for client_api_format in [
|
||||
"openai:chat",
|
||||
"openai:responses",
|
||||
"claude:messages",
|
||||
"gemini:generate_content",
|
||||
] {
|
||||
assert!(
|
||||
!request_pair_allowed_for_transport(
|
||||
&compact,
|
||||
client_api_format,
|
||||
"openai:responses:compact"
|
||||
),
|
||||
"{client_api_format} must not convert onto xAI compact"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_openai_chat_anchor_supports_cross_format_conversion_via_cascade() {
|
||||
let mut transport = transport_snapshot("windsurf", "openai:chat", "oauth", true, None);
|
||||
|
||||
@@ -30,6 +30,7 @@ pub mod url;
|
||||
pub mod vertex;
|
||||
mod video;
|
||||
pub mod windsurf;
|
||||
pub mod xai;
|
||||
|
||||
pub use aether_oauth as oauth;
|
||||
pub use agent_identity::{
|
||||
@@ -195,3 +196,10 @@ pub use windsurf::{
|
||||
local_windsurf_request_transport_unsupported_reason_with_network, GET_CHAT_MESSAGE_PATH,
|
||||
WINDSURF_ENVELOPE_NAME,
|
||||
};
|
||||
pub use xai::{
|
||||
extract_xai_user_id_from_auth_config, extract_xai_user_id_from_value,
|
||||
insert_cli_identity_headers, insert_cli_identity_headers_if_needed, is_xai_provider_transport,
|
||||
resolved_xai_request_base_url, resolved_xai_upstream_base_url,
|
||||
should_attach_cli_identity_headers, xai_auth_uses_api, xai_uses_official_api, XAI_API_BASE_URL,
|
||||
XAI_CHAT_PROXY_BASE_URL, XAI_PROVIDER_TYPE,
|
||||
};
|
||||
|
||||
@@ -275,6 +275,17 @@ const WINDSURF_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
..STANDARD_RUNTIME_POLICY
|
||||
};
|
||||
|
||||
const XAI_RUNTIME_POLICY: ProviderRuntimePolicy = ProviderRuntimePolicy {
|
||||
fixed_provider: true,
|
||||
api_format_inheritance: ProviderApiFormatInheritance::OAuthOrBearer,
|
||||
enable_format_conversion_by_default: true,
|
||||
oauth_is_bearer_like: true,
|
||||
supports_model_fetch: false,
|
||||
supports_local_openai_chat_transport: false,
|
||||
supports_local_same_format_transport: true,
|
||||
..STANDARD_RUNTIME_POLICY
|
||||
};
|
||||
|
||||
const CLAUDE_CODE_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
|
||||
provider_type: "claude_code",
|
||||
version: 2,
|
||||
@@ -446,6 +457,27 @@ const WINDSURF_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTem
|
||||
runtime_policy: WINDSURF_RUNTIME_POLICY,
|
||||
};
|
||||
|
||||
const XAI_FIXED_PROVIDER_TEMPLATE: FixedProviderTemplate = FixedProviderTemplate {
|
||||
provider_type: "xai",
|
||||
version: 2,
|
||||
base_url: crate::xai::XAI_CHAT_PROXY_BASE_URL,
|
||||
endpoints: &[
|
||||
FixedProviderEndpointTemplate {
|
||||
item_key: "openai:responses",
|
||||
api_format: "openai:responses",
|
||||
custom_path: None,
|
||||
config_defaults: FORCE_STREAM_ENDPOINT_CONFIG_DEFAULTS,
|
||||
},
|
||||
FixedProviderEndpointTemplate {
|
||||
item_key: "openai:responses:compact",
|
||||
api_format: "openai:responses:compact",
|
||||
custom_path: None,
|
||||
config_defaults: EMPTY_ENDPOINT_CONFIG_DEFAULTS,
|
||||
},
|
||||
],
|
||||
runtime_policy: XAI_RUNTIME_POLICY,
|
||||
};
|
||||
|
||||
pub fn provider_type_is_fixed(provider_type: &str) -> bool {
|
||||
provider_runtime_policy(provider_type).fixed_provider
|
||||
}
|
||||
@@ -498,6 +530,7 @@ pub fn fixed_provider_template(provider_type: &str) -> Option<&'static FixedProv
|
||||
"vertex_ai" => Some(&VERTEX_AI_FIXED_PROVIDER_TEMPLATE),
|
||||
"antigravity" => Some(&ANTIGRAVITY_FIXED_PROVIDER_TEMPLATE),
|
||||
"windsurf" => Some(&WINDSURF_FIXED_PROVIDER_TEMPLATE),
|
||||
"xai" => Some(&XAI_FIXED_PROVIDER_TEMPLATE),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
@@ -613,6 +646,16 @@ pub fn provider_type_admin_oauth_template(provider_type: &str) -> Option<Provide
|
||||
redirect_uri: "show-auth-token",
|
||||
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,
|
||||
}
|
||||
}
|
||||
@@ -825,6 +868,45 @@ mod tests {
|
||||
assert!(ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES.contains(&"windsurf"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn xai_fixed_provider_template_exposes_responses_endpoints() {
|
||||
let template = fixed_provider_template("xai").expect("xai template should exist");
|
||||
assert_eq!(template.provider_type, "xai");
|
||||
assert_eq!(template.base_url, crate::xai::XAI_CHAT_PROXY_BASE_URL);
|
||||
assert_eq!(template.version, 2);
|
||||
assert_eq!(
|
||||
template
|
||||
.endpoints
|
||||
.iter()
|
||||
.map(|item| item.api_format)
|
||||
.collect::<Vec<_>>(),
|
||||
vec!["openai:responses", "openai:responses:compact"]
|
||||
);
|
||||
|
||||
let policy = provider_runtime_policy("xai");
|
||||
assert!(policy.fixed_provider);
|
||||
assert!(policy.enable_format_conversion_by_default);
|
||||
assert!(policy.oauth_is_bearer_like);
|
||||
assert!(!policy.supports_model_fetch);
|
||||
assert!(policy.supports_local_same_format_transport);
|
||||
assert!(!policy.supports_local_openai_chat_transport);
|
||||
assert!(fixed_provider_key_inherits_api_formats(
|
||||
"xai", "oauth", None
|
||||
));
|
||||
assert!(fixed_provider_key_inherits_api_formats(
|
||||
"xai", "bearer", None
|
||||
));
|
||||
|
||||
let template = provider_type_admin_oauth_template("xai").expect("xai oauth template");
|
||||
assert_eq!(template.provider_type, "xai");
|
||||
assert_eq!(template.display_name, "xAI");
|
||||
assert_eq!(
|
||||
template.token_url,
|
||||
aether_oauth::provider::providers::XAI_TOKEN_URL
|
||||
);
|
||||
assert!(!ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES.contains(&"xai"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn fixed_provider_key_inheritance_keeps_oauth_and_kiro_configured_bearer_keys_open() {
|
||||
assert!(fixed_provider_key_inherits_api_formats(
|
||||
|
||||
@@ -42,6 +42,11 @@ pub fn apply_transport_request_body_semantics(
|
||||
{
|
||||
sanitize_claude_code_request_body(provider_request_body);
|
||||
}
|
||||
aether_ai_formats::apply_xai_upstream_payload_edits(
|
||||
provider_request_body,
|
||||
transport.provider.provider_type.as_str(),
|
||||
provider_api_format.as_str(),
|
||||
);
|
||||
if provider_api_format == "gemini:embedding" && is_vertex_transport_context(transport) {
|
||||
apply_vertex_gemini_embedding_body_semantics(provider_request_body)?;
|
||||
}
|
||||
|
||||
@@ -120,6 +120,12 @@ fn build_transport_request_url_inner(
|
||||
return Some(url);
|
||||
}
|
||||
|
||||
let xai_base =
|
||||
crate::xai::resolved_xai_upstream_base_url(transport, &normalized_provider_api_format);
|
||||
let request_base_url = xai_base
|
||||
.as_deref()
|
||||
.unwrap_or(transport.endpoint.base_url.as_str());
|
||||
|
||||
let custom_path_template = transport
|
||||
.endpoint
|
||||
.custom_path
|
||||
@@ -164,7 +170,7 @@ fn build_transport_request_url_inner(
|
||||
path.to_string()
|
||||
};
|
||||
let mut url = build_passthrough_path_url(
|
||||
&transport.endpoint.base_url,
|
||||
request_base_url,
|
||||
normalized_path.as_str(),
|
||||
params.request_query,
|
||||
blocked_keys,
|
||||
@@ -190,75 +196,68 @@ fn build_transport_request_url_inner(
|
||||
|
||||
let url = match normalized_provider_api_format.as_str() {
|
||||
"openai:chat" => Some(build_openai_chat_url(
|
||||
&transport.endpoint.base_url,
|
||||
request_base_url,
|
||||
params.request_query,
|
||||
)),
|
||||
"openai:responses" => Some(build_openai_responses_url(
|
||||
&transport.endpoint.base_url,
|
||||
request_base_url,
|
||||
params.request_query,
|
||||
false,
|
||||
)),
|
||||
"openai:responses:compact" => Some(build_openai_responses_url(
|
||||
&transport.endpoint.base_url,
|
||||
request_base_url,
|
||||
params.request_query,
|
||||
true,
|
||||
)),
|
||||
"openai:search" => Some(build_openai_search_url(
|
||||
&transport.endpoint.base_url,
|
||||
request_base_url,
|
||||
params.request_query,
|
||||
)),
|
||||
"openai:realtime" => build_passthrough_path_url(
|
||||
&transport.endpoint.base_url,
|
||||
request_base_url,
|
||||
"/v1/realtime",
|
||||
params.request_query,
|
||||
GATEWAY_CREDENTIAL_QUERY_KEYS,
|
||||
)
|
||||
.and_then(|url| replace_realtime_model_query(url, params.mapped_model?)),
|
||||
"codex:live" => build_passthrough_path_url(
|
||||
&transport.endpoint.base_url,
|
||||
request_base_url,
|
||||
"/live",
|
||||
params.request_query,
|
||||
GATEWAY_CREDENTIAL_QUERY_KEYS,
|
||||
),
|
||||
"openai:embedding" | "jina:embedding" => {
|
||||
build_provider_embedding_v1_url(&transport.endpoint.base_url, params.request_query)
|
||||
build_provider_embedding_v1_url(request_base_url, params.request_query)
|
||||
}
|
||||
"aliyun:multimodal_embedding" => {
|
||||
build_aliyun_multimodal_embedding_url(request_base_url, params.request_query)
|
||||
}
|
||||
"aliyun:multimodal_embedding" => build_aliyun_multimodal_embedding_url(
|
||||
&transport.endpoint.base_url,
|
||||
params.request_query,
|
||||
),
|
||||
"openai:rerank" | "jina:rerank" => {
|
||||
build_provider_rerank_v1_url(&transport.endpoint.base_url, params.request_query)
|
||||
build_provider_rerank_v1_url(request_base_url, params.request_query)
|
||||
}
|
||||
"claude:messages" => Some(if is_claude_count_tokens {
|
||||
build_default_claude_count_tokens_url(
|
||||
&transport.endpoint.base_url,
|
||||
params.request_query,
|
||||
)
|
||||
build_default_claude_count_tokens_url(request_base_url, params.request_query)
|
||||
} else {
|
||||
build_claude_messages_url(&transport.endpoint.base_url, params.request_query)
|
||||
build_claude_messages_url(request_base_url, params.request_query)
|
||||
}),
|
||||
"gemini:generate_content" => build_gemini_content_url(
|
||||
&transport.endpoint.base_url,
|
||||
request_base_url,
|
||||
params.mapped_model?,
|
||||
params.upstream_is_stream,
|
||||
params.request_query,
|
||||
),
|
||||
"gemini:embedding" => build_gemini_embedding_url(
|
||||
&transport.endpoint.base_url,
|
||||
request_base_url,
|
||||
params.mapped_model?,
|
||||
params.request_query,
|
||||
gemini_embedding_batch,
|
||||
),
|
||||
"gemini:interactions" => {
|
||||
build_gemini_interactions_url(&transport.endpoint.base_url, params.request_query)
|
||||
build_gemini_interactions_url(request_base_url, params.request_query)
|
||||
}
|
||||
"doubao:embedding" => {
|
||||
build_passthrough_path_url(request_base_url, "/embeddings", params.request_query, &[])
|
||||
}
|
||||
"doubao:embedding" => build_passthrough_path_url(
|
||||
&transport.endpoint.base_url,
|
||||
"/embeddings",
|
||||
params.request_query,
|
||||
&[],
|
||||
),
|
||||
_ => None,
|
||||
}?;
|
||||
|
||||
@@ -2417,4 +2416,82 @@ mod tests {
|
||||
"https://api.example.com/v1/messages?model=claude%26admin%3Dtrue%23fragment"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn xai_oauth_responses_use_cli_chat_proxy() {
|
||||
let mut transport = sample_transport(
|
||||
"xai",
|
||||
"openai:responses",
|
||||
"https://cli-chat-proxy.grok.com/v1",
|
||||
None,
|
||||
);
|
||||
transport.key.auth_type = "oauth".to_string();
|
||||
transport.key.decrypted_auth_config =
|
||||
Some(r#"{"refresh_token":"rt","using_api":false}"#.to_string());
|
||||
|
||||
let url = build_transport_request_url(
|
||||
&transport,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "openai:responses",
|
||||
mapped_model: Some("grok-4"),
|
||||
upstream_is_stream: true,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.expect("xai oauth responses URL");
|
||||
|
||||
assert_eq!(url, "https://cli-chat-proxy.grok.com/v1/responses");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn xai_compact_and_using_api_use_official_api() {
|
||||
let mut oauth = sample_transport(
|
||||
"xai",
|
||||
"openai:responses:compact",
|
||||
"https://cli-chat-proxy.grok.com/v1",
|
||||
None,
|
||||
);
|
||||
oauth.key.auth_type = "oauth".to_string();
|
||||
oauth.key.decrypted_auth_config =
|
||||
Some(r#"{"refresh_token":"rt","using_api":false}"#.to_string());
|
||||
|
||||
let compact = build_transport_request_url(
|
||||
&oauth,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "openai:responses:compact",
|
||||
mapped_model: Some("grok-4"),
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.expect("xai compact URL");
|
||||
assert_eq!(compact, "https://api.x.ai/v1/responses/compact");
|
||||
|
||||
let mut api_key = sample_transport(
|
||||
"xai",
|
||||
"openai:responses",
|
||||
"https://cli-chat-proxy.grok.com/v1",
|
||||
None,
|
||||
);
|
||||
api_key.key.auth_type = "oauth".to_string();
|
||||
api_key.key.decrypted_auth_config = Some(r#"{"using_api":true}"#.to_string());
|
||||
|
||||
let official = build_transport_request_url(
|
||||
&api_key,
|
||||
TransportRequestUrlParams {
|
||||
provider_api_format: "openai:responses",
|
||||
mapped_model: Some("grok-4"),
|
||||
upstream_is_stream: true,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
)
|
||||
.expect("xai api key URL");
|
||||
assert_eq!(official, "https://api.x.ai/v1/responses");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -396,6 +396,12 @@ pub fn build_standard_provider_request_headers(
|
||||
force_identity_accept_encoding(&mut headers);
|
||||
}
|
||||
|
||||
crate::xai::insert_cli_identity_headers_if_needed(
|
||||
input.transport,
|
||||
input.provider_api_format,
|
||||
&mut headers,
|
||||
);
|
||||
|
||||
let declared_connection_headers =
|
||||
crate::headers::declared_connection_header_names(input.headers, input.extra_headers);
|
||||
crate::headers::remove_declared_connection_headers(&mut headers, &declared_connection_headers);
|
||||
|
||||
@@ -0,0 +1,409 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use aether_ai_formats::normalize_api_format_alias;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
pub const XAI_PROVIDER_TYPE: &str = "xai";
|
||||
pub const XAI_CHAT_PROXY_BASE_URL: &str = "https://cli-chat-proxy.grok.com/v1";
|
||||
pub const XAI_API_BASE_URL: &str = "https://api.x.ai/v1";
|
||||
pub const XAI_CLIENT_VERSION: &str = "0.2.120";
|
||||
pub const XAI_TOKEN_AUTH_HEADER: &str = "x-xai-token-auth";
|
||||
pub const XAI_TOKEN_AUTH_VALUE: &str = "xai-grok-cli";
|
||||
pub const XAI_CLIENT_VERSION_HEADER: &str = "x-grok-client-version";
|
||||
pub const XAI_CLIENT_IDENTIFIER_HEADER: &str = "x-grok-client-identifier";
|
||||
pub const XAI_CLIENT_IDENTIFIER_VALUE: &str = "grok-shell";
|
||||
pub const XAI_AUTHENTICATE_RESPONSE_HEADER: &str = "x-authenticateresponse";
|
||||
pub const XAI_AUTHENTICATE_RESPONSE_VALUE: &str = "authenticate-response";
|
||||
|
||||
pub fn xai_cli_user_agent() -> String {
|
||||
format!("xai-grok-workspace/{XAI_CLIENT_VERSION}")
|
||||
}
|
||||
|
||||
pub fn is_xai_provider_transport(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(XAI_PROVIDER_TYPE)
|
||||
}
|
||||
|
||||
pub fn xai_uses_official_api(api_format: &str) -> bool {
|
||||
matches!(
|
||||
normalize_api_format_alias(api_format).as_str(),
|
||||
"openai:responses:compact"
|
||||
)
|
||||
}
|
||||
|
||||
pub fn resolved_xai_upstream_base_url(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> Option<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 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}"#)
|
||||
));
|
||||
}
|
||||
}
|
||||
@@ -1,5 +1,8 @@
|
||||
use std::path::PathBuf;
|
||||
use std::process::{Child, Command, Stdio};
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
|
||||
static POSTGRES_WORKDIR_SEQ: AtomicU64 = AtomicU64::new(0);
|
||||
|
||||
use aether_data::driver::postgres::PostgresPoolConfig;
|
||||
use aether_data::{DataBackends, DataLayerConfig};
|
||||
@@ -21,10 +24,20 @@ pub struct ManagedPostgresServer {
|
||||
impl ManagedPostgresServer {
|
||||
pub async fn start() -> Result<Self, Box<dyn std::error::Error>> {
|
||||
let port = reserve_local_port()?;
|
||||
// pid+port is not unique: cargo test shares one PID, and ephemeral ports
|
||||
// are reused after the listener is dropped. Parallel e2e tests then hit
|
||||
// create_dir AlreadyExists.
|
||||
let seq = POSTGRES_WORKDIR_SEQ.fetch_add(1, Ordering::Relaxed);
|
||||
let nanos = std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.map(|duration| duration.as_nanos())
|
||||
.unwrap_or(0);
|
||||
let workdir = std::env::temp_dir().join(format!(
|
||||
"aether-postgres-baseline-{}-{}",
|
||||
"aether-postgres-baseline-{}-{}-{}-{}",
|
||||
std::process::id(),
|
||||
port
|
||||
port,
|
||||
seq,
|
||||
nanos
|
||||
));
|
||||
let data_dir = workdir.join("data");
|
||||
std::fs::create_dir(&workdir)?;
|
||||
|
||||
@@ -6649,7 +6649,7 @@ mod tests {
|
||||
}
|
||||
|
||||
use std::collections::BTreeMap;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::Instant;
|
||||
|
||||
@@ -7375,9 +7375,27 @@ mod tests {
|
||||
queue: Arc<dyn RuntimeQueueStore>,
|
||||
policy_started: Arc<tokio::sync::Notify>,
|
||||
release_policy: Arc<tokio::sync::Notify>,
|
||||
policy_released: Arc<AtomicBool>,
|
||||
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)]
|
||||
struct FailingPolicyUsageStore {
|
||||
inner: NoRedisUsageStore,
|
||||
@@ -8409,7 +8427,18 @@ mod tests {
|
||||
async fn body_capture_policy(&self) -> Result<UsageBodyCapturePolicy, DataLayerError> {
|
||||
self.policy_reads.fetch_add(1, Ordering::AcqRel);
|
||||
self.policy_started.notify_one();
|
||||
self.release_policy.notified().await;
|
||||
// Latch the gate: Notify is edge-triggered, and later policy reads
|
||||
// (or a waiter that subscribed after a single notify) must not hang.
|
||||
loop {
|
||||
if self.policy_released.load(Ordering::Acquire) {
|
||||
break;
|
||||
}
|
||||
let notified = self.release_policy.notified();
|
||||
if self.policy_released.load(Ordering::Acquire) {
|
||||
break;
|
||||
}
|
||||
notified.await;
|
||||
}
|
||||
Ok(UsageBodyCapturePolicy::default())
|
||||
}
|
||||
}
|
||||
@@ -9043,15 +9072,39 @@ mod tests {
|
||||
.await
|
||||
.expect("a duplicate first-byte marker must release the terminal barrier");
|
||||
|
||||
let records = store.records.lock().expect("records lock");
|
||||
assert_eq!(
|
||||
records.len(),
|
||||
2,
|
||||
"the duplicate first byte must be coalesced"
|
||||
);
|
||||
assert_eq!(records[0].status, "streaming");
|
||||
assert_eq!(records[1].status, "completed");
|
||||
drop(records);
|
||||
{
|
||||
let records = store.records.lock().expect("records lock");
|
||||
assert_eq!(
|
||||
records.len(),
|
||||
2,
|
||||
"the duplicate first byte must be coalesced"
|
||||
);
|
||||
assert_eq!(records[0].status, "streaming");
|
||||
assert_eq!(records[1].status, "completed");
|
||||
}
|
||||
|
||||
// The terminal persistence notification can arrive before the submission
|
||||
// dispatcher accounts for its completed task and releases admission.
|
||||
timeout(Duration::from_secs(1), async {
|
||||
loop {
|
||||
let snapshot = runtime.metrics_snapshot();
|
||||
if snapshot.lifecycle_submission_pending == 0
|
||||
&& snapshot.first_byte_persistence_pending == 0
|
||||
&& snapshot.ordered_lifecycle_pending == 0
|
||||
&& runtime
|
||||
.lifecycle_submission
|
||||
.state
|
||||
.admission
|
||||
.available_permits()
|
||||
== CAPACITY
|
||||
{
|
||||
break;
|
||||
}
|
||||
sleep(Duration::from_millis(1)).await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("duplicate first-byte submission accounting should drain");
|
||||
|
||||
let snapshot = runtime.metrics_snapshot();
|
||||
assert_eq!(snapshot.lifecycle_submission_pending, 0);
|
||||
@@ -12325,12 +12378,9 @@ mod tests {
|
||||
async fn event_capture_budget_bounds_blocked_policy_waiters_and_releases_on_cancel_or_basic() {
|
||||
for limit in [0, 64 * 1024] {
|
||||
let runtime = UsageRuntime::new(UsageRuntimeConfig::default()).expect("runtime");
|
||||
let store = BlockingPolicyQueueConfiguredUsageStore {
|
||||
queue: Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())),
|
||||
policy_started: Arc::new(tokio::sync::Notify::new()),
|
||||
release_policy: Arc::new(tokio::sync::Notify::new()),
|
||||
policy_reads: Arc::new(AtomicUsize::new(0)),
|
||||
};
|
||||
let store = BlockingPolicyQueueConfiguredUsageStore::new(Arc::new(
|
||||
RuntimeState::memory(MemoryRuntimeStateConfig::default()),
|
||||
));
|
||||
let budget = Arc::new(crate::event_capture_budget::EventCaptureMemoryBudget::new(
|
||||
limit,
|
||||
));
|
||||
@@ -12391,7 +12441,7 @@ mod tests {
|
||||
.await
|
||||
.expect("replacement policy read starts");
|
||||
assert_eq!(budget.retained_bytes(), retained);
|
||||
store.release_policy.notify_one();
|
||||
store.release_blocked_policy();
|
||||
let event = timeout(Duration::from_secs(2), completing)
|
||||
.await
|
||||
.expect("Basic policy completes")
|
||||
@@ -13398,12 +13448,7 @@ mod tests {
|
||||
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
|
||||
let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0));
|
||||
let queue: Arc<dyn RuntimeQueueStore> = tracked_queue.clone();
|
||||
let store = BlockingPolicyQueueConfiguredUsageStore {
|
||||
queue,
|
||||
policy_started: Arc::new(tokio::sync::Notify::new()),
|
||||
release_policy: Arc::new(tokio::sync::Notify::new()),
|
||||
policy_reads: Arc::new(AtomicUsize::new(0)),
|
||||
};
|
||||
let store = BlockingPolicyQueueConfiguredUsageStore::new(queue);
|
||||
let runtime = UsageRuntime::new(config).expect("usage runtime should build");
|
||||
let request_id = "req-terminal-seed-waits-for-turn";
|
||||
let plan = terminal_test_plan(request_id);
|
||||
@@ -13425,9 +13470,10 @@ mod tests {
|
||||
assert_eq!(blocked_snapshot.terminal_submission_in_flight, 0);
|
||||
assert!(blocked_snapshot.lifecycle_submission_pending >= 2);
|
||||
|
||||
store.release_policy.notify_waiters();
|
||||
store.release_blocked_policy();
|
||||
timeout(Duration::from_secs(2), async {
|
||||
loop {
|
||||
store.release_blocked_policy();
|
||||
let snapshot = runtime.metrics_snapshot();
|
||||
if tracked_queue.successful_appends.load(Ordering::Acquire) == 1
|
||||
&& snapshot.lifecycle_submission_pending == 0
|
||||
@@ -13435,7 +13481,7 @@ mod tests {
|
||||
{
|
||||
break;
|
||||
}
|
||||
tokio::task::yield_now().await;
|
||||
sleep(Duration::from_millis(1)).await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
@@ -13468,12 +13514,7 @@ mod tests {
|
||||
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
|
||||
let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0));
|
||||
let queue: Arc<dyn RuntimeQueueStore> = tracked_queue.clone();
|
||||
let store = BlockingPolicyQueueConfiguredUsageStore {
|
||||
queue,
|
||||
policy_started: Arc::new(tokio::sync::Notify::new()),
|
||||
release_policy: Arc::new(tokio::sync::Notify::new()),
|
||||
policy_reads: Arc::new(AtomicUsize::new(0)),
|
||||
};
|
||||
let store = BlockingPolicyQueueConfiguredUsageStore::new(queue);
|
||||
let runtime = UsageRuntime::new(config).expect("usage runtime should build");
|
||||
let policy_started = store.policy_started.notified();
|
||||
|
||||
@@ -13520,9 +13561,10 @@ mod tests {
|
||||
assert_eq!(blocked_snapshot.terminal_submission_in_flight, 1);
|
||||
assert!(blocked_snapshot.lifecycle_submission_pending <= BACKLOG + 1);
|
||||
|
||||
store.release_policy.notify_waiters();
|
||||
store.release_blocked_policy();
|
||||
timeout(Duration::from_secs(5), async {
|
||||
loop {
|
||||
store.release_blocked_policy();
|
||||
let snapshot = runtime.metrics_snapshot();
|
||||
if tracked_queue.successful_appends.load(Ordering::Acquire) == BACKLOG + 1
|
||||
&& snapshot.lifecycle_submission_pending == 0
|
||||
@@ -13530,7 +13572,7 @@ mod tests {
|
||||
{
|
||||
break;
|
||||
}
|
||||
tokio::task::yield_now().await;
|
||||
sleep(Duration::from_millis(1)).await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
@@ -13828,12 +13870,7 @@ mod tests {
|
||||
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
|
||||
let tracked_queue = Arc::new(FlakyAppendQueueStore::new(inner_queue, 0));
|
||||
let queue: Arc<dyn RuntimeQueueStore> = tracked_queue.clone();
|
||||
let store = BlockingPolicyQueueConfiguredUsageStore {
|
||||
queue,
|
||||
policy_started: Arc::new(tokio::sync::Notify::new()),
|
||||
release_policy: Arc::new(tokio::sync::Notify::new()),
|
||||
policy_reads: Arc::new(AtomicUsize::new(0)),
|
||||
};
|
||||
let store = BlockingPolicyQueueConfiguredUsageStore::new(queue);
|
||||
let runtime = UsageRuntime::new(config).expect("usage runtime should build");
|
||||
let policy_started = store.policy_started.notified();
|
||||
runtime
|
||||
@@ -13892,10 +13929,10 @@ mod tests {
|
||||
.expect("terminal submissions should reach the execution backlog");
|
||||
let saturated_snapshot = runtime.metrics_snapshot();
|
||||
|
||||
store.release_policy.notify_waiters();
|
||||
store.release_blocked_policy();
|
||||
let all_completed = timeout(Duration::from_secs(2), async {
|
||||
loop {
|
||||
store.release_policy.notify_waiters();
|
||||
store.release_blocked_policy();
|
||||
if tracked_queue.successful_appends.load(Ordering::Acquire)
|
||||
== EXCESS_SUBMISSIONS + 1
|
||||
&& runtime.metrics_snapshot().terminal_submission_in_flight == 0
|
||||
|
||||
Reference in New Issue
Block a user