Merge remote-tracking branch 'zhefox/main' into zhefox-main

# Conflicts:
#	crates/aether-admin/src/provider/quota.rs
#	crates/aether-ai/formats/src/formats/openai/chat/stream.rs
#	crates/aether-ai/formats/src/formats/openai/responses/mod.rs
#	crates/aether-provider/pool/src/provider.rs
#	crates/aether-provider/pool/src/quota.rs
This commit is contained in:
zhefox
2026-09-02 15:25:27 +08:00
229 changed files with 51071 additions and 896 deletions
@@ -0,0 +1,249 @@
use std::sync::{Arc, OnceLock};
use std::time::{SystemTime, UNIX_EPOCH};
use aether_provider_transport::CodexFingerprintConvergenceContext;
use http::{request::Parts, HeaderMap};
use serde_json::Value;
use uuid::Uuid;
use crate::client_session_affinity::codex_request_signals_from_request;
#[derive(Debug, Clone)]
pub(crate) struct CodexFingerprintContextSlot(Arc<OnceLock<CodexFingerprintConvergenceContext>>);
impl Default for CodexFingerprintContextSlot {
fn default() -> Self {
Self(Arc::new(OnceLock::new()))
}
}
impl CodexFingerprintContextSlot {
fn resolve(
&self,
headers: &HeaderMap,
body_json: &Value,
) -> CodexFingerprintConvergenceContext {
self.0
.get_or_init(|| {
build_codex_fingerprint_context(headers, body_json, Uuid::now_v7().to_string())
})
.clone()
}
}
pub(crate) fn resolve_codex_fingerprint_context(
parts: &Parts,
body_json: &Value,
) -> CodexFingerprintConvergenceContext {
if let Some(context) = parts
.extensions
.get::<CodexFingerprintConvergenceContext>()
.cloned()
{
return context;
}
if let Some(slot) = parts.extensions.get::<CodexFingerprintContextSlot>() {
return slot.resolve(&parts.headers, body_json);
}
build_codex_fingerprint_context(&parts.headers, body_json, Uuid::now_v7().to_string())
}
pub(crate) fn install_codex_fingerprint_context_slot(parts: &mut Parts) {
if parts
.extensions
.get::<CodexFingerprintConvergenceContext>()
.is_none()
&& parts
.extensions
.get::<CodexFingerprintContextSlot>()
.is_none()
{
parts
.extensions
.insert(CodexFingerprintContextSlot::default());
}
}
pub(crate) fn ensure_codex_fingerprint_context(
parts: &mut Parts,
body_json: &Value,
) -> CodexFingerprintConvergenceContext {
let context = resolve_codex_fingerprint_context(parts, body_json);
if parts
.extensions
.get::<CodexFingerprintConvergenceContext>()
.is_none()
{
parts.extensions.remove::<CodexFingerprintContextSlot>();
parts.extensions.insert(context.clone());
}
context
}
pub(crate) fn attach_codex_logical_turn_context(
parts: &mut Parts,
body_json: &Value,
logical_turn_id: &str,
) -> CodexFingerprintConvergenceContext {
let context =
build_codex_fingerprint_context(&parts.headers, body_json, logical_turn_id.to_string());
parts.extensions.remove::<CodexFingerprintContextSlot>();
parts.extensions.insert(context.clone());
context
}
pub(crate) fn restore_codex_logical_turn_context(
parts: &mut Parts,
context: &CodexFingerprintConvergenceContext,
) {
parts.extensions.remove::<CodexFingerprintContextSlot>();
parts.extensions.insert(context.clone());
}
fn build_codex_fingerprint_context(
headers: &HeaderMap,
body_json: &Value,
logical_turn_id: String,
) -> CodexFingerprintConvergenceContext {
let signals = codex_request_signals_from_request(headers, Some(body_json));
let mut context =
CodexFingerprintConvergenceContext::new(logical_turn_id, current_unix_millis());
if let Some(turn_id) = signals.turn_id {
context = context.with_original_turn_id(turn_id);
}
if let Some(session_id) = signals.thread_id.or(signals.session_id) {
context = context.with_original_client_session_id(session_id);
}
if let Some(prompt_cache_key) = signals.prompt_cache_key {
context = context.with_original_prompt_cache_key(prompt_cache_key);
}
context
}
fn current_unix_millis() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_millis()
.try_into()
.unwrap_or(u64::MAX)
}
#[cfg(test)]
mod tests {
use http::HeaderValue;
use serde_json::json;
use super::*;
#[test]
fn request_signals_are_captured_once_for_the_logical_turn() {
let request = http::Request::builder()
.header("thread-id", "header-thread")
.body(())
.expect("request should build");
let (mut parts, _) = request.into_parts();
let body = json!({
"prompt_cache_key": "client-cache",
"client_metadata": {
"turn_id": "client-turn",
"thread_id": "body-thread"
}
});
let context = attach_codex_logical_turn_context(&mut parts, &body, "logical-turn");
assert_eq!(context.logical_turn_id(), "logical-turn");
assert_eq!(context.original_turn_id(), Some("client-turn"));
assert_eq!(context.original_client_session_id(), Some("header-thread"));
assert_eq!(context.original_prompt_cache_key(), Some("client-cache"));
assert_eq!(
parts.extensions.get::<CodexFingerprintConvergenceContext>(),
Some(&context)
);
}
#[test]
fn restored_context_wins_over_retry_request_signals() {
let original = CodexFingerprintConvergenceContext::new("logical-turn", 1234)
.with_original_turn_id("original-turn")
.with_original_client_session_id("original-thread")
.with_original_prompt_cache_key("original-cache");
let request = http::Request::builder()
.body(())
.expect("request should build");
let (mut parts, _) = request.into_parts();
parts
.headers
.insert("thread-id", HeaderValue::from_static("retry-thread"));
restore_codex_logical_turn_context(&mut parts, &original);
let resolved = resolve_codex_fingerprint_context(
&parts,
&json!({
"prompt_cache_key": "retry-cache",
"client_metadata": {"turn_id": "retry-turn"}
}),
);
assert_eq!(resolved, original);
assert_eq!(resolved.turn_started_at_unix_ms(), 1234);
}
#[test]
fn generated_context_is_persisted_for_http_replanning() {
let request = http::Request::builder()
.header("session-id", "client-session")
.body(())
.expect("request should build");
let (mut parts, _) = request.into_parts();
let body = json!({
"prompt_cache_key": "client-cache",
"client_metadata": {"turn_id": "client-turn"}
});
let first = ensure_codex_fingerprint_context(&mut parts, &body);
let second = resolve_codex_fingerprint_context(
&parts,
&json!({
"prompt_cache_key": "retry-cache",
"client_metadata": {"turn_id": "retry-turn"}
}),
);
assert_eq!(second, first);
assert_eq!(second.original_turn_id(), Some("client-turn"));
assert_eq!(second.original_prompt_cache_key(), Some("client-cache"));
}
#[test]
fn installed_slot_reuses_context_across_cloned_parts() {
let request = http::Request::builder()
.body(())
.expect("request should build");
let (mut parts, _) = request.into_parts();
install_codex_fingerprint_context_slot(&mut parts);
let cloned_parts = parts.clone();
let first = resolve_codex_fingerprint_context(
&parts,
&json!({
"prompt_cache_key": "first-cache",
"client_metadata": {"turn_id": "first-turn"}
}),
);
let second = resolve_codex_fingerprint_context(
&cloned_parts,
&json!({
"prompt_cache_key": "second-cache",
"client_metadata": {"turn_id": "second-turn"}
}),
);
assert_eq!(second, first);
assert_eq!(second.original_turn_id(), Some("first-turn"));
assert_eq!(second.original_prompt_cache_key(), Some("first-cache"));
}
}
@@ -1,5 +1,6 @@
mod adaptation;
pub(crate) mod api;
pub(crate) mod codex_context;
mod finalize;
mod planner;
mod pure;
@@ -86,6 +86,52 @@ impl GatewayLocalCandidatePreselectionPort<'_> {
}
}
/// A Responses compaction request carries the OpenAI-only `compaction_trigger`
/// control item. It must stay on an OpenAI Responses endpoint: treating it as
/// an ordinary cross-format request would make Gemini/Claude candidates look
/// eligible and defer the inevitable lossy-conversion failure until payload
/// construction.
fn request_candidate_api_formats_for_operation(
client_api_format: &str,
require_streaming: bool,
request_operation: Option<&str>,
) -> Vec<String> {
let candidate_api_formats =
crate::ai_serving::request_candidate_api_formats(client_api_format, require_streaming)
.into_iter()
.map(str::to_string)
.collect::<Vec<_>>();
restrict_candidate_api_formats_for_operation(
client_api_format,
request_operation,
candidate_api_formats,
)
}
fn restrict_candidate_api_formats_for_operation(
client_api_format: &str,
request_operation: Option<&str>,
candidate_api_formats: Vec<String>,
) -> Vec<String> {
let is_responses_compaction = request_operation.is_some_and(|operation| {
operation.eq_ignore_ascii_case(crate::ai_serving::OPENAI_RESPONSES_OPERATION_COMPACT)
});
let is_standard_responses_client =
crate::ai_serving::normalize_api_format_alias(client_api_format) == "openai:responses";
if !(is_responses_compaction && is_standard_responses_client) {
return candidate_api_formats;
}
candidate_api_formats
.into_iter()
.filter(|candidate_api_format| {
crate::ai_serving::normalize_api_format_alias(candidate_api_format)
== "openai:responses"
})
.collect()
}
#[async_trait]
impl AiCandidatePreselectionPort for GatewayLocalCandidatePreselectionPort<'_> {
type Candidate = SchedulerMinimalCandidateSelectionCandidate;
@@ -219,11 +265,11 @@ pub(crate) async fn preselect_local_execution_candidates_with_serving(
>,
GatewayError,
> {
let candidate_api_formats =
crate::ai_serving::request_candidate_api_formats(client_api_format, require_streaming)
.into_iter()
.map(str::to_string)
.collect::<Vec<_>>();
let candidate_api_formats = request_candidate_api_formats_for_operation(
client_api_format,
require_streaming,
request_operation,
);
preselect_local_execution_candidates_for_api_formats_with_serving(
state,
model_directive_policy,
@@ -264,6 +310,11 @@ pub(crate) async fn preselect_local_execution_candidates_for_api_formats_with_se
>,
GatewayError,
> {
let candidate_api_formats = restrict_candidate_api_formats_for_operation(
client_api_format,
request_operation,
candidate_api_formats,
);
let model_directive_routing_models = resolve_model_directive_routing_models(
model_directive_policy,
&candidate_api_formats,
@@ -362,11 +413,11 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
allow_priority_page_cache: bool,
trace_id: Option<&str>,
) -> Self {
let candidate_api_formats =
crate::ai_serving::request_candidate_api_formats(client_api_format, require_streaming)
.into_iter()
.map(str::to_string)
.collect::<Vec<_>>();
let candidate_api_formats = request_candidate_api_formats_for_operation(
client_api_format,
require_streaming,
request_operation,
);
let model_directive_routing_models = resolve_model_directive_routing_models(
model_directive_policy,
&candidate_api_formats,
@@ -1437,6 +1488,31 @@ mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
#[test]
fn compaction_operation_excludes_non_responses_provider_formats() {
assert_eq!(
request_candidate_api_formats_for_operation("openai:responses", true, Some("compact"),),
vec!["openai:responses"]
);
assert_eq!(
request_candidate_api_formats_for_operation("openai:responses", true, None),
vec![
"openai:responses",
"openai:chat",
"claude:messages",
"gemini:generate_content"
]
);
assert_eq!(
request_candidate_api_formats_for_operation(
"openai:responses:compact",
false,
Some("compact"),
),
vec!["openai:responses:compact"]
);
}
#[derive(Default)]
struct EmptyFallbackCountingRepository {
fallback_reads: AtomicUsize,
@@ -13,6 +13,7 @@ use http::{HeaderMap, HeaderName, HeaderValue};
use serde_json::{json, Value};
use crate::ai_serving::planner::common::extract_standard_requested_model;
use crate::ai_serving::transport::CodexFingerprintConvergenceContext;
use crate::ai_serving::{
ClientSurface, ExecutionRuntimeAuthContext, GatewayAuthApiKeySnapshot,
GatewayCredentialCarrier, GatewayProviderTransportSnapshot, PlannerAppState,
@@ -55,7 +56,7 @@ pub(crate) struct LocalRequestedModelDecisionInput {
pub(crate) client_surface: Option<ClientSurface>,
pub(crate) gateway_credential_carrier: Option<GatewayCredentialCarrier>,
pub(crate) client_session_affinity: Option<ClientSessionAffinity>,
pub(crate) original_client_session_id: Option<String>,
pub(crate) codex_fingerprint_context: Option<CodexFingerprintConvergenceContext>,
pub(crate) routing_policy: Option<ResolvedRoutingPolicy>,
pub(crate) routing_trace_seed: Option<RoutingDecisionTrace>,
pub(crate) routing_context: Option<LocalRoutingRequestContext>,
@@ -167,7 +168,7 @@ pub(crate) fn apply_provider_request_routing_policy_to_decision_with_websocket_m
provider_api_format.as_str(),
);
}
apply_codex_oauth_fingerprint_convergence_to_decision(
apply_codex_fingerprint_convergence_to_decision(
input,
decision,
transport,
@@ -230,7 +231,7 @@ pub(crate) fn apply_provider_request_routing_policy_to_decision_with_websocket_m
provider_api_format.as_str(),
);
}
apply_codex_oauth_fingerprint_convergence_to_decision(
apply_codex_fingerprint_convergence_to_decision(
input,
decision,
transport,
@@ -291,6 +292,7 @@ pub(crate) fn apply_provider_request_routing_policy_to_decision_with_websocket_m
crate::ai_serving::openai_responses_reasoning_replay_policy(
transport.provider.provider_type.as_str(),
transport.endpoint.base_url.as_str(),
provider_model.as_str(),
)
})
.unwrap_or_default();
@@ -355,7 +357,7 @@ pub(crate) fn apply_provider_request_routing_policy_to_decision_with_websocket_m
if original_provider_request_body.is_some() {
decision.provider_request_body = Some(provider_request_body);
}
apply_codex_oauth_fingerprint_convergence_to_decision(
apply_codex_fingerprint_convergence_to_decision(
input,
decision,
transport,
@@ -365,7 +367,7 @@ pub(crate) fn apply_provider_request_routing_policy_to_decision_with_websocket_m
Ok(())
}
fn apply_codex_oauth_fingerprint_convergence_to_decision(
fn apply_codex_fingerprint_convergence_to_decision(
input: &LocalRequestedModelDecisionInput,
decision: &mut AiExecutionDecision,
transport: Option<&GatewayProviderTransportSnapshot>,
@@ -376,13 +378,24 @@ fn apply_codex_oauth_fingerprint_convergence_to_decision(
else {
return;
};
crate::ai_serving::transport::apply_codex_oauth_fingerprint_convergence(
let Some(context) = input.codex_fingerprint_context.as_ref() else {
return;
};
let applied = crate::ai_serving::transport::apply_codex_fingerprint_convergence_with_context(
transport,
provider_api_format,
input.original_client_session_id.as_deref(),
context,
&mut decision.provider_request_headers,
provider_request_body,
);
if applied {
decision.prompt_cache_key = provider_request_body
.get("prompt_cache_key")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
}
}
struct GatewayAuthenticatedDecisionInputPort<'a> {
@@ -471,7 +484,7 @@ pub(crate) fn build_local_requested_model_decision_input(
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
codex_fingerprint_context: None,
routing_policy: None,
routing_trace_seed: None,
routing_context: None,
@@ -486,7 +499,8 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
body_json: &Value,
client_api_format: &str,
) -> Result<(), GatewayError> {
input.original_client_session_id = original_client_session_id_from_headers(&parts.headers);
input.codex_fingerprint_context =
Some(crate::ai_serving::codex_context::resolve_codex_fingerprint_context(parts, body_json));
let explicit_group = routing_header_value_str(&parts.headers, ROUTING_GROUP_HEADER);
let selected_group = match state.routing_group_read_repository() {
Some(repository) => {
@@ -737,12 +751,6 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
Ok(())
}
fn original_client_session_id_from_headers(headers: &HeaderMap) -> Option<String> {
routing_header_value_str(headers, "session-id")
.or_else(|| routing_header_value_str(headers, "session_id"))
.or_else(|| routing_header_value_str(headers, "x-session-id"))
}
fn try_attach_static_default_routing_policy_to_input(
input: &mut LocalRequestedModelDecisionInput,
parts: &http::request::Parts,
@@ -1106,38 +1114,6 @@ mod tests {
GatewayProviderTransportProvider,
};
#[test]
fn original_client_session_id_accepts_live_header_as_fallback() {
let headers = HeaderMap::from_iter([(
HeaderName::from_static("x-session-id"),
HeaderValue::from_static("live-thread-1"),
)]);
assert_eq!(
original_client_session_id_from_headers(&headers).as_deref(),
Some("live-thread-1")
);
}
#[test]
fn original_client_session_id_prefers_responses_headers_over_live_fallback() {
let headers = HeaderMap::from_iter([
(
HeaderName::from_static("session-id"),
HeaderValue::from_static("responses-session"),
),
(
HeaderName::from_static("x-session-id"),
HeaderValue::from_static("live-thread"),
),
]);
assert_eq!(
original_client_session_id_from_headers(&headers).as_deref(),
Some("responses-session")
);
}
#[test]
fn explicit_routing_selection_cache_key_is_principal_specific() {
let first = routing_group_selection_cache_key(
@@ -1350,7 +1326,7 @@ mod tests {
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
codex_fingerprint_context: None,
routing_policy: None,
routing_trace_seed: None,
model_directive_policy: Default::default(),
@@ -1599,7 +1575,7 @@ mod tests {
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
codex_fingerprint_context: None,
routing_policy: None,
routing_trace_seed: None,
model_directive_policy: Default::default(),
@@ -1669,7 +1645,7 @@ mod tests {
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
codex_fingerprint_context: None,
routing_policy: None,
routing_trace_seed: None,
routing_context: None,
@@ -1754,7 +1730,13 @@ mod tests {
});
let mut with_mutation = sample_decision_input();
for input in [&mut no_context, &mut empty_mutation, &mut with_mutation] {
input.original_client_session_id = Some("client-session-1".to_string());
input.codex_fingerprint_context = Some(
CodexFingerprintConvergenceContext::new(
uuid::Uuid::new_v4().to_string(),
1_756_668_000_000,
)
.with_original_client_session_id("client-session-1".to_string()),
);
}
let mut stable_identity = None;
@@ -1801,6 +1783,10 @@ mod tests {
.provider_request_body
.as_ref()
.expect("request body");
assert_eq!(
decision.prompt_cache_key.as_deref(),
body.get("prompt_cache_key").and_then(Value::as_str)
);
assert_eq!(
body["prompt_cache_key"],
"172c39e6-c0a0-5a70-8b63-e0f8e0d185a3"
@@ -2005,6 +1991,7 @@ mod tests {
let body = decision.provider_request_body.as_ref().expect("body");
assert!(body.get("prompt_cache_key").is_none());
assert!(decision.prompt_cache_key.is_none());
assert!(body.get("client_metadata").is_none());
assert!(!decision.provider_request_headers.contains_key("session-id"));
assert!(!decision.provider_request_headers.contains_key("thread-id"));
@@ -183,6 +183,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
let reasoning_replay_policy = openai_responses_reasoning_replay_policy(
prepared.transport.provider.provider_type.as_str(),
prepared.transport.endpoint.base_url.as_str(),
prepared.mapped_model.as_str(),
);
let redaction = resolve_provider_chat_pii_redaction(
state,
@@ -15,11 +15,25 @@ pub(crate) fn is_deepseek_provider(provider_type: &str, base_url: &str) -> bool
host == "deepseek.com" || host.ends_with(".deepseek.com")
}
fn is_deepseek_model(provider_model: &str) -> bool {
let provider_model = provider_model.trim().to_ascii_lowercase();
let leaf = provider_model
.rsplit(['/', ':'])
.next()
.unwrap_or(provider_model.as_str());
leaf == "deepseek" || leaf.starts_with("deepseek-") || leaf.starts_with("deepseek_")
}
fn is_deepseek_upstream(provider_type: &str, base_url: &str, provider_model: &str) -> bool {
is_deepseek_provider(provider_type, base_url) || is_deepseek_model(provider_model)
}
pub(crate) fn openai_responses_reasoning_replay_policy(
provider_type: &str,
base_url: &str,
provider_model: &str,
) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy {
if is_deepseek_provider(provider_type, base_url) {
if is_deepseek_upstream(provider_type, base_url, provider_model) {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
} else {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
@@ -33,7 +47,11 @@ pub(crate) fn apply_deepseek_tool_call_thinking_compat(
provider_api_format: &str,
original_request_body: Option<&Value>,
) {
if !is_deepseek_provider(provider_type, base_url) {
let provider_model = provider_request_body
.get("model")
.and_then(Value::as_str)
.unwrap_or_default();
if !is_deepseek_upstream(provider_type, base_url, provider_model) {
return;
}
@@ -302,11 +320,35 @@ mod tests {
));
assert!(!is_deepseek_provider("custom", "ftp://api.deepseek.com/v1"));
assert_eq!(
openai_responses_reasoning_replay_policy("custom", "https://api.deepseek.com/v1"),
openai_responses_reasoning_replay_policy(
"custom",
"https://api.deepseek.com/v1",
"deepseek-v4-flash",
),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
);
assert_eq!(
openai_responses_reasoning_replay_policy("openai", "https://api.openai.com/v1"),
openai_responses_reasoning_replay_policy(
"openai",
"https://api.openai.com/v1",
"gpt-5.6-sol",
),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
);
assert_eq!(
openai_responses_reasoning_replay_policy(
"custom",
"https://api.b.ai/v1",
"deepseek-v4-flash",
),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
);
assert_eq!(
openai_responses_reasoning_replay_policy(
"custom",
"https://api.b.ai/v1",
"not-deepseek-compatible",
),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
);
}
@@ -330,8 +372,11 @@ mod tests {
"input": reasoning_items.clone(),
"future_request_field": {"preserve": true}
});
let replay_policy =
openai_responses_reasoning_replay_policy("custom", "https://api.deepseek.com/v1");
let replay_policy = openai_responses_reasoning_replay_policy(
"custom",
"https://api.deepseek.com/v1",
"deepseek-v4-flash",
);
let mut provider_body = crate::ai_serving::build_standard_request_body_with_model_directives_and_request_headers_and_reasoning_replay_policy(
&request,
"openai:responses",
@@ -373,7 +418,11 @@ mod tests {
crate::ai_serving::strip_incompatible_openai_responses_reasoning_items_with_policy(
&mut deepseek,
"openai:responses",
openai_responses_reasoning_replay_policy("custom", "https://api.deepseek.com/v1"),
openai_responses_reasoning_replay_policy(
"custom",
"https://api.deepseek.com/v1",
"deepseek-v4-flash",
),
),
0
);
@@ -383,7 +432,11 @@ mod tests {
crate::ai_serving::strip_incompatible_openai_responses_reasoning_items_with_policy(
&mut openai,
"openai:responses",
openai_responses_reasoning_replay_policy("openai", "https://api.openai.com/v1"),
openai_responses_reasoning_replay_policy(
"openai",
"https://api.openai.com/v1",
"gpt-5.6-sol",
),
),
66
);
@@ -417,6 +470,52 @@ mod tests {
assert_eq!(body["messages"][1]["reasoning_content"], "");
}
#[test]
fn custom_relay_deepseek_model_adds_chat_thinking_compat() {
let mut body = json!({
"model": "deepseek-v4-flash",
"messages": [
{"role": "user", "content": "inspect the repository"},
{"role": "assistant", "content": null, "tool_calls": [{
"id": "call_1",
"type": "function",
"function": {"name": "inspect", "arguments": "{}"}
}]},
{"role": "tool", "tool_call_id": "call_1", "content": "done"}
]
});
apply_deepseek_tool_call_thinking_compat(
&mut body,
"custom",
"https://api.b.ai/v1",
"openai:chat",
None,
);
assert_eq!(body["thinking"]["type"], "enabled");
assert_eq!(body["messages"][1]["reasoning_content"], "");
}
#[test]
fn custom_relay_non_deepseek_model_is_not_rewritten() {
let original = json!({
"model": "not-deepseek-compatible",
"messages": [{"role": "assistant", "content": "done"}]
});
let mut body = original.clone();
apply_deepseek_tool_call_thinking_compat(
&mut body,
"custom",
"https://api.b.ai/v1",
"openai:chat",
None,
);
assert_eq!(body, original);
}
#[test]
fn openai_chat_deepseek_honors_disabled_thinking() {
let original = json!({"reasoning_effort": "none"});
@@ -377,7 +377,7 @@ mod tests {
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
codex_fingerprint_context: None,
routing_policy: None,
routing_trace_seed: None,
routing_context: None,
@@ -597,6 +597,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
let reasoning_replay_policy = openai_responses_reasoning_replay_policy(
transport.provider.provider_type.as_str(),
transport.endpoint.base_url.as_str(),
prepared_candidate.mapped_model.as_str(),
);
let redaction = resolve_provider_chat_pii_redaction(
state,
@@ -102,6 +102,59 @@ fn builds_openai_chat_cross_format_request_body_from_openai_responses_source() {
assert_eq!(provider_request_body["messages"][0]["content"], "hello");
}
#[test]
fn maps_openai_responses_additional_tools_without_message_name() {
let body_json = json!({
"model": "gpt-5",
"input": [
{
"type": "additional_tools",
"role": "developer",
"tools": [{
"type": "function",
"name": "get_weather",
"description": "Get the weather",
"parameters": {
"type": "object",
"properties": {}
}
}]
},
{
"role": "user",
"content": "What is the weather?"
}
]
});
let provider_request_body = build_cross_format_openai_responses_request_body(
&body_json,
"gpt-5-upstream",
"openai:responses",
"openai:chat",
false,
false,
"openai",
None,
None,
&http::HeaderMap::new(),
false,
)
.expect("Responses additional tools should map to a Chat request body");
assert_eq!(
provider_request_body["messages"].as_array().map(Vec::len),
Some(1)
);
assert_eq!(provider_request_body["messages"][0]["role"], "user");
assert!(provider_request_body["messages"][0].get("name").is_none());
assert_eq!(provider_request_body["tools"][0]["type"], "function");
assert_eq!(
provider_request_body["tools"][0]["function"]["name"],
"get_weather"
);
}
#[test]
fn local_openai_responses_wrapper_preserves_body_order_after_edits() {
let body_json: Value = serde_json::from_str(
@@ -159,6 +159,7 @@ fn finalize_openai_chat_provider_request_body(
openai_responses_reasoning_replay_policy(
transport.provider.provider_type.as_str(),
transport.endpoint.base_url.as_str(),
mapped_model,
),
)
.err()
@@ -2182,7 +2183,7 @@ mod tests {
client_surface: None,
gateway_credential_carrier: None,
client_session_affinity: None,
original_client_session_id: None,
codex_fingerprint_context: None,
routing_policy: None,
routing_trace_seed: None,
routing_context: None,
@@ -438,6 +438,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts_with_
let reasoning_replay_policy = openai_responses_reasoning_replay_policy(
transport.provider.provider_type.as_str(),
transport.endpoint.base_url.as_str(),
mapped_model.as_str(),
);
let redaction = resolve_provider_chat_pii_redaction(
state,
@@ -1005,6 +1005,7 @@ pub(crate) async fn maybe_build_responses_websocket_decision(
reasoning_replay_policy: openai_responses_reasoning_replay_policy(
transport.provider.provider_type.as_str(),
transport.endpoint.base_url.as_str(),
mapped_model.as_str(),
),
model_directive_patch: input
.model_directive_policy
@@ -181,7 +181,7 @@ pub(crate) use aether_ai_formats::{
openai_responses_request_operation, openai_responses_synthetic_reasoning_item_id,
strip_incompatible_openai_responses_reasoning_items,
strip_incompatible_openai_responses_reasoning_items_with_policy, ApiOperation, ClientSurface,
CODEX_CLIENT_VERSION,
CODEX_CLIENT_VERSION, OPENAI_RESPONSES_OPERATION_COMPACT,
};
pub(crate) fn plan_kind_matches_api_operation(
@@ -59,9 +59,9 @@ pub(crate) mod windsurf {
}
pub(crate) use aether_provider_transport::{
append_transport_diagnostics_to_value, apply_codex_oauth_fingerprint_convergence,
apply_local_auth_config_header_overrides, apply_local_body_rules,
apply_local_body_rules_with_request_headers, apply_local_header_rules,
append_transport_diagnostics_to_value, apply_codex_fingerprint_convergence,
apply_codex_fingerprint_convergence_with_context, apply_local_auth_config_header_overrides,
apply_local_body_rules, apply_local_body_rules_with_request_headers, apply_local_header_rules,
apply_local_header_rules_with_request_headers, apply_standard_provider_request_body_rules,
apply_standard_provider_request_body_rules_with_request_headers,
apply_transport_request_body_semantics, body_rules_are_locally_supported,
@@ -107,8 +107,9 @@ pub(crate) use aether_provider_transport::{
supports_local_generic_oauth_request_auth_resolution,
supports_local_oauth_request_auth_resolution, transport_proxy_is_locally_supported,
transport_supports_api_operation, video_create_transport_unsupported_reason,
AnthropicCompatibilityProfile, CandidateTransportPolicyFacts, GatewayProviderTransportSnapshot,
GeminiCliRequestAuth, GeminiCliRequestAuthSupport, GeminiCliRequestAuthUnsupportedReason,
AnthropicCompatibilityProfile, CandidateTransportPolicyFacts,
CodexFingerprintConvergenceContext, GatewayProviderTransportSnapshot, GeminiCliRequestAuth,
GeminiCliRequestAuthSupport, GeminiCliRequestAuthUnsupportedReason,
GeminiCliRequestEnvelopeSupport, GeminiFilesHeadersInput, GeminiFilesRequestBodyError,
GeminiFilesRequestBodyParts, GrokHeaderInput, LocalResolvedOAuthRequestAuth,
ProviderOpenAiImageHeadersInput, ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput,
@@ -1,7 +1,10 @@
use axum::routing::get;
use axum::Router;
use crate::{handlers::proxy::proxy_request, state::AppState};
use crate::{
handlers::{proxy::proxy_request, public::vscodex_ws_proxy},
state::AppState,
};
pub(crate) fn mount_public_support_routes(router: Router<AppState>) -> Router<AppState> {
router
@@ -26,6 +29,7 @@ pub(crate) fn mount_public_support_routes(router: Router<AppState>) -> Router<Ap
.route("/api/capabilities", get(proxy_request))
.route("/api/capabilities/user-configurable", get(proxy_request))
.route("/api/capabilities/model/{*model_path}", get(proxy_request))
.route("/api/vscodex/ws", get(vscodex_ws_proxy))
.route("/install/{*install_path}", get(proxy_request))
.route("/install-tunnel/{*install_path}", get(proxy_request))
.route("/i/{*install_path}", get(proxy_request))
@@ -23,6 +23,21 @@ pub(crate) struct ClientSessionScope {
pub(crate) source: ClientSessionSignalSource,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub(crate) struct CodexRequestSignals {
pub(crate) session_id: Option<String>,
pub(crate) thread_id: Option<String>,
pub(crate) turn_id: Option<String>,
pub(crate) prompt_cache_key: Option<String>,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
struct CodexTurnMetadataSignals {
session_id: Option<String>,
thread_id: Option<String>,
turn_id: Option<String>,
}
impl ClientSessionScope {
fn new(
client_family: impl Into<String>,
@@ -136,6 +151,13 @@ pub(crate) fn client_session_scope_from_request(
.or_else(|| extract_scope_from_other_specific_adapters(&request, client_family.as_str()))
}
pub(crate) fn codex_request_signals_from_request(
headers: &http::HeaderMap,
body_json: Option<&Value>,
) -> CodexRequestSignals {
extract_codex_request_signals(&ClientSessionRequest { headers, body_json })
}
fn codex_search_session_scope(request: &ClientSessionRequest<'_>) -> Option<ClientSessionScope> {
let session_id = request
.body_json?
@@ -408,29 +430,7 @@ impl ClientSessionScopeAdapter for CodexSessionScopeAdapter {
}
fn extract_scope(&self, request: &ClientSessionRequest<'_>) -> Option<ClientSessionScope> {
header_value_str(request.headers, "session-id")
.or_else(|| header_value_str(request.headers, "thread-id"))
.or_else(|| header_value_str(request.headers, "session_id"))
.or_else(|| header_value_str(request.headers, "conversation_id"))
.map(|root_session| {
ClientSessionScope::new(
self.family(),
root_session,
None,
header_value_str(request.headers, "chatgpt-account-id"),
ClientSessionSignalSource::Header,
)
})
.or_else(|| {
let body_session = GenericSessionScopeAdapter.extract_scope(request)?;
Some(ClientSessionScope::new(
self.family(),
body_session.session_id,
body_session.agent_id,
header_value_str(request.headers, "chatgpt-account-id"),
body_session.source,
))
})
codex_request_session_scope_from_request(request)
}
}
@@ -778,6 +778,162 @@ fn explicit_aether_session_scope(
))
}
fn extract_codex_request_signals(request: &ClientSessionRequest<'_>) -> CodexRequestSignals {
let body_client_metadata = request
.body_json
.and_then(|body| body.get("client_metadata"))
.and_then(Value::as_object);
let body_turn_metadata = codex_turn_metadata_signals(
body_client_metadata.and_then(|metadata| metadata.get("x-codex-turn-metadata")),
);
let header_turn_metadata = header_value_str(request.headers, "x-codex-turn-metadata")
.map(|raw| parse_codex_turn_metadata(&raw))
.unwrap_or_default();
let native_thread_id = header_value_str(request.headers, "thread-id")
.or_else(|| {
body_client_metadata
.and_then(|metadata| value_at_map_path(metadata, "thread_id"))
.map(ToOwned::to_owned)
})
.or_else(|| body_turn_metadata.thread_id.clone());
let turn_id = body_client_metadata
.and_then(|metadata| value_at_map_path(metadata, "turn_id"))
.map(ToOwned::to_owned)
.or_else(|| body_turn_metadata.turn_id.clone())
.or_else(|| {
request
.body_json
.and_then(|body| value_at_path(body, &["turn_id"]))
.map(ToOwned::to_owned)
})
.or(header_turn_metadata.turn_id);
let prompt_cache_key = request
.body_json
.and_then(|body| value_at_path(body, &["prompt_cache_key"]))
.map(ToOwned::to_owned);
let session_id =
codex_request_session_scope(request, &body_turn_metadata).map(|scope| scope.session_id);
let thread_id = native_thread_id.or_else(|| session_id.clone());
CodexRequestSignals {
session_id,
thread_id,
turn_id,
prompt_cache_key,
}
}
fn codex_request_session_scope_from_request(
request: &ClientSessionRequest<'_>,
) -> Option<ClientSessionScope> {
let body_turn_metadata = codex_turn_metadata_signals(
request
.body_json
.and_then(|body| body.get("client_metadata"))
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("x-codex-turn-metadata")),
);
codex_request_session_scope(request, &body_turn_metadata)
}
fn codex_request_session_scope(
request: &ClientSessionRequest<'_>,
body_turn_metadata: &CodexTurnMetadataSignals,
) -> Option<ClientSessionScope> {
if let Some(scope) = explicit_aether_session_scope(request, CodexSessionScopeAdapter.family()) {
return Some(scope);
}
if let Some(root_session) = header_value_str(request.headers, "session-id")
.or_else(|| header_value_str(request.headers, "thread-id"))
.or_else(|| header_value_str(request.headers, "session_id"))
.or_else(|| header_value_str(request.headers, "conversation_id"))
.or_else(|| header_value_str(request.headers, "x-session-id"))
{
return Some(codex_session_scope(
request,
root_session,
None,
ClientSessionSignalSource::Header,
));
}
let body_client_metadata = request
.body_json
.and_then(|body| body.get("client_metadata"))
.and_then(Value::as_object);
if let Some(root_session) = body_client_metadata
.and_then(|metadata| value_at_map_path(metadata, "session_id"))
.or_else(|| {
body_client_metadata.and_then(|metadata| value_at_map_path(metadata, "thread_id"))
})
.map(ToOwned::to_owned)
.or_else(|| body_turn_metadata.session_id.clone())
.or_else(|| body_turn_metadata.thread_id.clone())
{
return Some(codex_session_scope(
request,
root_session,
None,
ClientSessionSignalSource::Body,
));
}
let generic = GenericSessionScopeAdapter.extract_scope(request)?;
Some(codex_session_scope(
request,
generic.session_id,
generic.agent_id,
generic.source,
))
}
fn codex_session_scope(
request: &ClientSessionRequest<'_>,
session_id: String,
agent_id: Option<String>,
source: ClientSessionSignalSource,
) -> ClientSessionScope {
ClientSessionScope::new(
CodexSessionScopeAdapter.family(),
session_id,
agent_id,
header_value_str(request.headers, "chatgpt-account-id"),
source,
)
}
fn codex_turn_metadata_signals(value: Option<&Value>) -> CodexTurnMetadataSignals {
match value {
Some(Value::Object(metadata)) => codex_turn_metadata_signals_from_map(metadata),
Some(Value::String(raw)) => parse_codex_turn_metadata(raw),
_ => CodexTurnMetadataSignals::default(),
}
}
fn parse_codex_turn_metadata(raw: &str) -> CodexTurnMetadataSignals {
serde_json::from_str::<Map<String, Value>>(raw)
.map(|metadata| codex_turn_metadata_signals_from_map(&metadata))
.unwrap_or_default()
}
fn codex_turn_metadata_signals_from_map(metadata: &Map<String, Value>) -> CodexTurnMetadataSignals {
CodexTurnMetadataSignals {
session_id: value_at_map_path(metadata, "session_id").map(ToOwned::to_owned),
thread_id: value_at_map_path(metadata, "thread_id").map(ToOwned::to_owned),
turn_id: value_at_map_path(metadata, "turn_id").map(ToOwned::to_owned),
}
}
fn value_at_map_path<'a>(object: &'a Map<String, Value>, key: &str) -> Option<&'a str> {
object
.get(key)
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
}
fn normalize_session_key(
account_hint: Option<&str>,
root_session: &str,
@@ -831,12 +987,25 @@ mod tests {
client_session_affinity_from_api_request,
client_session_affinity_from_report_context_value, client_session_affinity_from_request,
client_session_affinity_report_context_value, client_session_scope_from_request,
ClientSessionSignalSource, AETHER_AGENT_ID_HEADER, AETHER_SESSION_ID_HEADER,
codex_request_signals_from_request, ClientSessionSignalSource, AETHER_AGENT_ID_HEADER,
AETHER_SESSION_ID_HEADER,
};
use aether_scheduler_core::ClientSessionAffinity;
use http::{HeaderMap, HeaderValue};
use http::{HeaderMap, HeaderName, HeaderValue};
use serde_json::json;
fn request_headers(values: &[(&str, &str)]) -> HeaderMap {
values
.iter()
.map(|(name, value)| {
(
HeaderName::from_bytes(name.as_bytes()).expect("valid test header name"),
HeaderValue::from_bytes(value.as_bytes()).expect("valid test header value"),
)
})
.collect()
}
#[test]
fn unknown_adapter_extracts_body_session_and_agent() {
let body = json!({
@@ -933,6 +1102,276 @@ mod tests {
);
}
#[test]
fn codex_request_signals_apply_session_precedence() {
let cases = vec![
(
request_headers(&[
(AETHER_SESSION_ID_HEADER, "aether-session"),
("session-id", "header-session"),
]),
json!({"client_metadata": {"session_id": "body-session"}}),
"aether-session",
ClientSessionSignalSource::ExplicitAetherHeader,
),
(
request_headers(&[
("session-id", "header-session"),
("thread-id", "header-thread"),
("session_id", "header-session-underscore"),
("conversation_id", "header-conversation"),
]),
json!({"client_metadata": {"session_id": "body-session"}}),
"header-session",
ClientSessionSignalSource::Header,
),
(
request_headers(&[
("thread-id", "header-thread"),
("session_id", "header-session-underscore"),
("conversation_id", "header-conversation"),
]),
json!({"client_metadata": {"session_id": "body-session"}}),
"header-thread",
ClientSessionSignalSource::Header,
),
(
request_headers(&[
("session_id", "header-session-underscore"),
("conversation_id", "header-conversation"),
]),
json!({"client_metadata": {"session_id": "body-session"}}),
"header-session-underscore",
ClientSessionSignalSource::Header,
),
(
request_headers(&[("conversation_id", "header-conversation")]),
json!({"client_metadata": {"session_id": "body-session"}}),
"header-conversation",
ClientSessionSignalSource::Header,
),
(
HeaderMap::new(),
json!({
"prompt_cache_key": "prompt-cache",
"client_metadata": {
"session_id": "body-session",
"thread_id": "body-thread",
"x-codex-turn-metadata": {
"session_id": "nested-session",
"thread_id": "nested-thread"
}
}
}),
"body-session",
ClientSessionSignalSource::Body,
),
(
HeaderMap::new(),
json!({
"prompt_cache_key": "prompt-cache",
"client_metadata": {
"thread_id": "body-thread",
"x-codex-turn-metadata": {"session_id": "nested-session"}
}
}),
"body-thread",
ClientSessionSignalSource::Body,
),
(
HeaderMap::new(),
json!({
"prompt_cache_key": "prompt-cache",
"client_metadata": {
"x-codex-turn-metadata": {
"session_id": "nested-session",
"thread_id": "nested-thread"
}
}
}),
"nested-session",
ClientSessionSignalSource::Body,
),
(
HeaderMap::new(),
json!({
"prompt_cache_key": "prompt-cache",
"client_metadata": {
"x-codex-turn-metadata": json!({
"thread_id": "nested-thread"
}).to_string()
}
}),
"nested-thread",
ClientSessionSignalSource::Body,
),
(
HeaderMap::new(),
json!({
"prompt_cache_key": "prompt-cache",
"conversation_id": "generic-conversation"
}),
"prompt-cache",
ClientSessionSignalSource::Body,
),
(
HeaderMap::new(),
json!({"metadata": {"session_id": "generic-session"}}),
"generic-session",
ClientSessionSignalSource::Body,
),
];
for (headers, body, expected_session_id, expected_source) in cases {
let signals = codex_request_signals_from_request(&headers, Some(&body));
assert_eq!(signals.session_id.as_deref(), Some(expected_session_id));
let mut codex_headers = headers;
codex_headers.insert(
http::header::USER_AGENT,
HeaderValue::from_static("codex_cli_rs/0.144.1"),
);
let scope = client_session_scope_from_request(&codex_headers, Some(&body))
.expect("Codex scope should reuse the native signal precedence");
assert_eq!(scope.client_family, "codex");
assert_eq!(scope.session_id, expected_session_id);
assert_eq!(scope.source, expected_source);
}
}
#[test]
fn codex_request_signals_extract_thread_and_prompt_cache_independently() {
let body = json!({
"prompt_cache_key": "prompt-cache",
"client_metadata": {
"thread_id": "body-thread",
"x-codex-turn-metadata": {"thread_id": "nested-thread"}
}
});
let headers = request_headers(&[("thread-id", "header-thread")]);
let header_signals = codex_request_signals_from_request(&headers, Some(&body));
assert_eq!(header_signals.thread_id.as_deref(), Some("header-thread"));
assert_eq!(
header_signals.prompt_cache_key.as_deref(),
Some("prompt-cache")
);
let body_signals = codex_request_signals_from_request(&HeaderMap::new(), Some(&body));
assert_eq!(body_signals.thread_id.as_deref(), Some("body-thread"));
let nested_body = json!({
"client_metadata": {
"x-codex-turn-metadata": json!({
"thread_id": "nested-thread"
}).to_string()
}
});
let nested_signals =
codex_request_signals_from_request(&HeaderMap::new(), Some(&nested_body));
assert_eq!(nested_signals.thread_id.as_deref(), Some("nested-thread"));
let session_only_body = json!({"client_metadata": {"session_id": "body-session"}});
let session_only_signals =
codex_request_signals_from_request(&HeaderMap::new(), Some(&session_only_body));
assert_eq!(
session_only_signals.thread_id.as_deref(),
Some("body-session")
);
}
#[test]
fn codex_request_signals_use_live_session_header() {
let headers = request_headers(&[("x-session-id", "live-session")]);
let signals = codex_request_signals_from_request(&headers, None);
assert_eq!(signals.session_id.as_deref(), Some("live-session"));
assert_eq!(signals.thread_id.as_deref(), Some("live-session"));
}
#[test]
fn codex_request_signals_prefer_responses_session_header_over_live_session_header() {
let headers = request_headers(&[
("session-id", "responses-session"),
("x-session-id", "live-session"),
]);
let signals = codex_request_signals_from_request(&headers, None);
assert_eq!(signals.session_id.as_deref(), Some("responses-session"));
assert_eq!(signals.thread_id.as_deref(), Some("responses-session"));
}
#[test]
fn codex_request_signals_apply_turn_precedence() {
let headers = request_headers(&[("x-codex-turn-metadata", r#"{"turn_id":"header-turn"}"#)]);
let direct_body = json!({
"turn_id": "top-level-turn",
"client_metadata": {
"turn_id": "body-turn",
"x-codex-turn-metadata": {"turn_id": "nested-turn"}
}
});
assert_eq!(
codex_request_signals_from_request(&headers, Some(&direct_body))
.turn_id
.as_deref(),
Some("body-turn")
);
let nested_object_body = json!({
"turn_id": "top-level-turn",
"client_metadata": {
"x-codex-turn-metadata": {"turn_id": "nested-object-turn"}
}
});
assert_eq!(
codex_request_signals_from_request(&headers, Some(&nested_object_body))
.turn_id
.as_deref(),
Some("nested-object-turn")
);
let nested_string_body = json!({
"turn_id": "top-level-turn",
"client_metadata": {
"x-codex-turn-metadata": json!({
"turn_id": "nested-string-turn"
}).to_string()
}
});
assert_eq!(
codex_request_signals_from_request(&headers, Some(&nested_string_body))
.turn_id
.as_deref(),
Some("nested-string-turn")
);
let top_level_body = json!({
"turn_id": "top-level-turn",
"client_metadata": {"x-codex-turn-metadata": "not-json"}
});
assert_eq!(
codex_request_signals_from_request(&headers, Some(&top_level_body))
.turn_id
.as_deref(),
Some("top-level-turn")
);
assert_eq!(
codex_request_signals_from_request(&headers, None)
.turn_id
.as_deref(),
Some("header-turn")
);
}
#[test]
fn codex_request_signals_ignore_client_request_id() {
let headers = request_headers(&[("x-client-request-id", "request-only-id")]);
let signals =
codex_request_signals_from_request(&headers, Some(&json!({"model": "gpt-5"})));
assert_eq!(signals, super::CodexRequestSignals::default());
}
#[test]
fn report_context_round_trips_normalized_session_affinity() {
let affinity = ClientSessionAffinity::new(
+34 -5
View File
@@ -205,8 +205,8 @@ async fn available_balance_capacity_usd(
.as_ref()
.is_some_and(|wallet| wallet.limit_mode.eq_ignore_ascii_case("unlimited"));
Ok(match quota.as_ref() {
Some(quota) if !quota.allow_wallet_overage => Some(quota.remaining_usd.max(0.0)),
Some(_) if wallet_is_unlimited => None,
Some(quota) if !quota.allow_wallet_overage => Some(quota.remaining_usd.max(0.0)),
Some(quota) => Some(quota.remaining_usd.max(0.0) + wallet_available_usd.unwrap_or(0.0)),
None if wallet_is_unlimited => None,
None => wallet_available_usd,
@@ -832,9 +832,10 @@ mod tests {
use serde_json::json;
use super::{
execution_plan_balance_capacity_rejection, execution_plan_cost_upper_bound_cache_key,
max_output_tokens_from_request, openai_request_input_is_self_contained,
output_choice_count_upper_bound, request_model_local_rejection, GatewayLocalAuthRejection,
available_balance_capacity_usd, execution_plan_balance_capacity_rejection,
execution_plan_cost_upper_bound_cache_key, max_output_tokens_from_request,
openai_request_input_is_self_contained, output_choice_count_upper_bound,
request_model_local_rejection, GatewayLocalAuthRejection,
};
use crate::control::{GatewayControlAuthContext, GatewayControlDecision};
use crate::data::GatewayDataState;
@@ -939,6 +940,14 @@ mod tests {
fn state_with_quota_and_wallet(
quota: UserDailyQuotaAvailabilityRecord,
context: StoredBillingModelContext,
) -> AppState {
state_with_quota_context_and_wallet(quota, context, sample_wallet("user-1", 30.0))
}
fn state_with_quota_context_and_wallet(
quota: UserDailyQuotaAvailabilityRecord,
context: StoredBillingModelContext,
wallet: StoredWalletSnapshot,
) -> AppState {
let candidate_repository =
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
@@ -952,7 +961,7 @@ mod tests {
AppState::new()
.expect("state should build")
.with_data_state_for_tests(data)
.with_auth_wallets_for_tests(vec![sample_wallet("user-1", 30.0)])
.with_auth_wallets_for_tests(vec![wallet])
}
fn state_with_model_mapping() -> AppState {
@@ -1350,6 +1359,26 @@ mod tests {
}
}
#[tokio::test]
async fn unlimited_wallet_capacity_ignores_exhausted_non_overage_quota() {
let context = billing_context_with_pricing(None, None, None, None);
let mut wallet = sample_wallet("user-1", 0.0);
wallet.limit_mode = "unlimited".to_string();
let state =
state_with_quota_context_and_wallet(quota_availability(0.0, false), context, wallet);
let decision = decision_with_allowed_models(vec!["gpt-5".to_string()]);
let auth_context = decision
.auth_context
.as_ref()
.expect("decision should include auth context");
let capacity = available_balance_capacity_usd(&state, auth_context)
.await
.expect("capacity should resolve");
assert_eq!(capacity, None);
}
#[tokio::test]
async fn positive_balance_does_not_allow_historical_invalid_processing_pricing() {
let context = billing_context_with_pricing(
@@ -797,6 +797,18 @@ pub(super) fn classify_admin_operations_family_route(
"admin:users",
false,
))
} else if method == http::Method::DELETE
&& normalized_path_no_trailing.starts_with("/api/admin/users/")
&& normalized_path_no_trailing.contains("/billing/entitlements/")
&& normalized_path_no_trailing.matches('/').count() == 7
{
Some(classified(
"admin_proxy",
"users_manage",
"revoke_user_billing_entitlement",
"admin:users",
false,
))
} else if method == http::Method::GET
&& normalized_path.starts_with("/api/admin/users/")
&& normalized_path.ends_with("/sessions")
@@ -520,6 +520,65 @@ pub(super) fn classify_public_support_route(
"aether:ccswitch_usage",
false,
))
} else if method == http::Method::POST
&& matches!(normalized_path, "/api/vscodex/pair" | "/api/vscodex/pair/")
{
Some(classified(
"public_support",
"vscodex",
"pairing_exchange",
"public:vscodex",
false,
))
} else if method == http::Method::GET
&& matches!(
normalized_path,
"/api/users/me/vscodex/devices" | "/api/users/me/vscodex/devices/"
)
{
Some(classified(
"public_support",
"users_me",
"vscodex_devices_list",
"user:self",
false,
))
} else if method == http::Method::POST
&& matches!(
normalized_path,
"/api/users/me/vscodex/pairings" | "/api/users/me/vscodex/pairings/"
)
{
Some(classified(
"public_support",
"users_me",
"vscodex_pairing_create",
"user:self",
false,
))
} else if method == http::Method::POST
&& matches!(
normalized_path,
"/api/users/me/vscodex/ws-tickets" | "/api/users/me/vscodex/ws-tickets/"
)
{
Some(classified(
"public_support",
"users_me",
"vscodex_ws_ticket_create",
"user:self",
false,
))
} else if method == http::Method::DELETE
&& has_single_segment_after_prefix(normalized_path, "/api/users/me/vscodex/devices/")
{
Some(classified(
"public_support",
"users_me",
"vscodex_device_delete",
"user:self",
false,
))
} else if method == http::Method::GET
&& matches!(
normalized_path,
@@ -103,6 +103,22 @@ fn classifies_admin_user_billing_routes_as_admin_proxy_route() {
Some("admin:users")
);
let revoke_uri: Uri = "/api/admin/users/user-1/billing/entitlements/entitlement-1"
.parse()
.expect("uri should parse");
let revoke = classify_control_route(&http::Method::DELETE, &revoke_uri, &headers)
.expect("route should classify");
assert_eq!(revoke.route_class.as_deref(), Some("admin_proxy"));
assert_eq!(revoke.route_family.as_deref(), Some("users_manage"));
assert_eq!(
revoke.route_kind.as_deref(),
Some("revoke_user_billing_entitlement")
);
assert_eq!(
revoke.auth_endpoint_signature.as_deref(),
Some("admin:users")
);
let context = GatewayPublicRequestContext::from_request_parts(
"trace-user-billing-grant",
&http::Method::POST,
@@ -440,6 +440,26 @@ fn classifies_users_me_routes_as_public_support_route() {
"/api/users/me/available-models",
"available_models",
),
(
http::Method::GET,
"/api/users/me/vscodex/devices",
"vscodex_devices_list",
),
(
http::Method::POST,
"/api/users/me/vscodex/pairings",
"vscodex_pairing_create",
),
(
http::Method::DELETE,
"/api/users/me/vscodex/devices/device-1",
"vscodex_device_delete",
),
(
http::Method::POST,
"/api/users/me/vscodex/ws-tickets",
"vscodex_ws_ticket_create",
),
(
http::Method::PUT,
"/api/users/me/model-capabilities",
@@ -496,6 +516,49 @@ fn classifies_users_me_routes_as_public_support_route() {
}
}
#[test]
fn vscodex_post_routes_buffer_request_body() {
let headers = headers(&[]);
for path in [
"/api/vscodex/pair",
"/api/users/me/vscodex/pairings",
"/api/users/me/vscodex/ws-tickets",
] {
let uri: Uri = path.parse().expect("uri should parse");
let decision = classify_control_route(&http::Method::POST, &uri, &headers)
.expect("route should classify");
let context = GatewayPublicRequestContext::from_request_parts(
"trace-vscodex",
&http::Method::POST,
&uri,
&headers,
Some(decision),
);
assert!(
local_proxy_route_requires_buffered_body(&context),
"{path} should buffer its JSON body"
);
}
}
#[test]
fn classifies_public_vscodex_pairing_exchange() {
let headers = headers(&[]);
let uri: Uri = "/api/vscodex/pair".parse().expect("uri should parse");
let decision =
classify_control_route(&http::Method::POST, &uri, &headers).expect("route should classify");
assert_eq!(decision.route_class.as_deref(), Some("public_support"));
assert_eq!(decision.route_family.as_deref(), Some("vscodex"));
assert_eq!(decision.route_kind.as_deref(), Some("pairing_exchange"));
assert_eq!(
decision.auth_endpoint_signature.as_deref(),
Some("public:vscodex")
);
assert!(!decision.is_execution_runtime_candidate());
}
#[test]
fn classifies_ccswitch_usage_as_api_key_public_support_route() {
let headers = headers(&[]);
@@ -2578,6 +2578,21 @@ impl GatewayDataState {
}
}
pub(crate) async fn revoke_user_plan_entitlement(
&self,
user_id: &str,
entitlement_id: &str,
) -> Result<AdminBillingMutationOutcome<()>, DataLayerError> {
match &self.billing_reader {
Some(repository) => {
repository
.revoke_user_plan_entitlement(user_id, entitlement_id)
.await
}
None => Ok(AdminBillingMutationOutcome::Unavailable),
}
}
pub(crate) async fn find_user_daily_quota_availability(
&self,
user_id: &str,
@@ -600,11 +600,17 @@ impl<'a> PoolKeyCursor<'a> {
return;
};
self.exhaustion_skip_recorded = true;
record_local_runtime_candidate_skip_reason(
self.state.app(),
trace_id,
self.runtime_miss_pool_exhaustion_skip_reason(),
);
if self.skip_reason_counts.is_empty() {
record_local_runtime_candidate_skip_reason(
self.state.app(),
trace_id,
"pool_group_exhausted",
);
return;
}
for reason in self.skip_reason_counts.keys() {
record_local_runtime_candidate_skip_reason(self.state.app(), trace_id, reason);
}
}
fn runtime_miss_pool_exhaustion_skip_reason(&self) -> &'static str {
@@ -1886,17 +1892,17 @@ fn pool_key_candidate_order_for_group(
})
.collect::<Vec<_>>();
let active_presets = ProviderPoolService::with_builtin_adapters()
.normalize_scheduling_presets(group.transport.provider.provider_type.as_str(), &presets)
.into_iter()
.map(|preset| preset.preset)
.collect::<Vec<_>>();
.normalize_scheduling_presets(group.transport.provider.provider_type.as_str(), &presets);
if let Some(distribution_mode) = active_presets
.iter()
.find(|preset| pool_distribution_mode_preset(preset.as_str()))
.map(String::as_str)
.find(|preset| pool_distribution_mode_preset(preset.preset.as_str()))
{
return match distribution_mode {
"cache_affinity" => StoredPoolKeyCandidateOrder::CacheAffinity,
return match distribution_mode.preset.as_str() {
"cache_affinity" => match distribution_mode.mode.as_deref() {
Some("lru") => StoredPoolKeyCandidateOrder::Lru,
Some("single_account") => StoredPoolKeyCandidateOrder::SingleAccount,
_ => StoredPoolKeyCandidateOrder::CacheAffinity,
},
"load_balance" => StoredPoolKeyCandidateOrder::LoadBalance {
seed: pool_sort_seed(),
},
@@ -2260,7 +2266,7 @@ mod tests {
}
#[test]
fn pool_scheduler_promotes_sticky_hit_before_other_sorted_keys() {
fn pool_scheduler_promotes_sticky_hit_before_lru_secondary_order() {
let key_a = sample_eligible_candidate(
"provider-pool",
"endpoint-1",
@@ -2268,7 +2274,11 @@ mod tests {
10,
Some(json!({
"pool_advanced": {
"scheduling_presets": [{"preset": "cache_affinity", "enabled": true}]
"scheduling_presets": [{
"preset": "cache_affinity",
"enabled": true,
"mode": "lru"
}]
}
})),
);
@@ -2279,7 +2289,11 @@ mod tests {
10,
Some(json!({
"pool_advanced": {
"scheduling_presets": [{"preset": "cache_affinity", "enabled": true}]
"scheduling_presets": [{
"preset": "cache_affinity",
"enabled": true,
"mode": "lru"
}]
}
})),
);
@@ -2313,6 +2327,37 @@ mod tests {
);
}
#[test]
fn cache_affinity_secondary_modes_select_distinct_candidate_orders() {
for (mode, expected) in [
("single_account", StoredPoolKeyCandidateOrder::SingleAccount),
("lru", StoredPoolKeyCandidateOrder::Lru),
] {
let group = sample_eligible_candidate(
"provider-pool",
"endpoint-1",
"key-a",
10,
Some(json!({
"pool_advanced": {
"scheduling_presets": [{
"preset": "cache_affinity",
"enabled": true,
"mode": mode
}]
}
})),
);
let config = pool_config_for_candidate(&group).expect("pool config should parse");
assert!(admin_provider_pool_cache_affinity_enabled(&config));
assert_eq!(
pool_key_candidate_order_for_group(&group, Some(&config)),
expected
);
}
}
#[test]
fn pool_scheduler_ignores_sticky_hit_without_cache_affinity() {
let key_a = sample_eligible_candidate(
@@ -3300,8 +3345,12 @@ mod tests {
.take_local_execution_runtime_miss_diagnostic(trace_id)
.expect("runtime miss diagnostic should exist");
assert_eq!(diagnostic.reason, "all_candidates_skipped");
assert_eq!(diagnostic.skipped_candidate_count, Some(1));
assert_eq!(diagnostic.skipped_candidate_count, Some(2));
assert_eq!(diagnostic.skip_reasons.get("pool_cooldown"), Some(&1));
assert_eq!(
diagnostic.skip_reasons.get("transport_snapshot_missing"),
Some(&1)
);
}
#[tokio::test]
@@ -4639,6 +4688,53 @@ mod tests {
assert!(context.quota_exhausted);
}
#[test]
fn pool_catalog_context_scopes_antigravity_exhaustion_to_requested_model() {
let mut key = sample_catalog_oauth_key("key-antigravity-model-quota");
key.status_snapshot = Some(json!({
"quota": {
"version": 2,
"provider_type": "antigravity",
"exhausted": false,
"windows": [
{
"code": "model:gemini-3.1-pro-high",
"scope": "model",
"model": "gemini-3.1-pro-high",
"used_ratio": 1.0,
"is_exhausted": true
},
{
"code": "model:gemini-3-flash-agent",
"scope": "model",
"model": "gemini-3-flash-agent",
"used_ratio": 0.1,
"is_exhausted": false
}
]
}
}));
let app = app_state_with_catalog_key(key.clone());
let exhausted = build_pool_catalog_key_context(
PlannerAppState::new(&app),
&ProviderPoolService::with_builtin_adapters(),
&key,
"antigravity",
Some("gemini-3.1-pro-high"),
);
let available = build_pool_catalog_key_context(
PlannerAppState::new(&app),
&ProviderPoolService::with_builtin_adapters(),
&key,
"antigravity",
Some("gemini-3-flash-agent"),
);
assert!(exhausted.quota_exhausted);
assert!(!available.quota_exhausted);
}
#[test]
fn pool_catalog_context_marks_known_banned_account_from_metadata() {
let mut key = sample_catalog_oauth_key("key-account-banned");
@@ -5,6 +5,7 @@ use serde_json::Value;
use crate::execution_runtime::MAX_STREAM_PREFETCH_BYTES;
const ANTHROPIC_PRECOMMIT_MAX_WAIT: Duration = Duration::from_millis(750);
const GEMINI_PRECOMMIT_MAX_WAIT: Duration = Duration::from_millis(750);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum StreamCommitPolicy {
@@ -14,6 +15,10 @@ pub(super) enum StreamCommitPolicy {
max_bytes: usize,
max_wait: Duration,
},
FirstGeminiSemanticEvent {
max_bytes: usize,
max_wait: Duration,
},
}
impl StreamCommitPolicy {
@@ -51,6 +56,12 @@ impl StreamCommitPolicy {
max_wait: ANTHROPIC_PRECOMMIT_MAX_WAIT,
};
}
if provider_api_format.eq_ignore_ascii_case("gemini:generate_content") {
return Self::FirstGeminiSemanticEvent {
max_bytes: MAX_STREAM_PREFETCH_BYTES,
max_wait: GEMINI_PRECOMMIT_MAX_WAIT,
};
}
return Self::ResponseHeaders;
}
@@ -78,12 +89,16 @@ impl StreamCommitPolicy {
}
pub(super) const fn requires_bounded_frame_wait(self) -> bool {
matches!(self, Self::FirstAnthropicSemanticEvent { .. })
matches!(
self,
Self::FirstAnthropicSemanticEvent { .. } | Self::FirstGeminiSemanticEvent { .. }
)
}
pub(super) const fn max_precommit_wait(self) -> Option<Duration> {
match self {
Self::FirstAnthropicSemanticEvent { max_wait, .. } => Some(max_wait),
Self::FirstAnthropicSemanticEvent { max_wait, .. }
| Self::FirstGeminiSemanticEvent { max_wait, .. } => Some(max_wait),
Self::ResponseHeaders | Self::FirstClassifiedBody => None,
}
}
@@ -91,6 +106,10 @@ impl StreamCommitPolicy {
pub(super) const fn is_native_anthropic(self) -> bool {
matches!(self, Self::FirstAnthropicSemanticEvent { .. })
}
pub(super) const fn is_gemini(self) -> bool {
matches!(self, Self::FirstGeminiSemanticEvent { .. })
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -113,6 +132,7 @@ pub(super) struct StreamCommitGate {
state: StreamCommitState,
observed_bytes: usize,
anthropic: AnthropicSsePrecommitInspector,
gemini: GeminiSsePrecommitInspector,
}
impl StreamCommitGate {
@@ -127,6 +147,7 @@ impl StreamCommitGate {
state,
observed_bytes: 0,
anthropic: AnthropicSsePrecommitInspector::default(),
gemini: GeminiSsePrecommitInspector::default(),
}
}
@@ -143,21 +164,32 @@ impl StreamCommitGate {
return StreamPrecommitObservation::Commit;
}
let StreamCommitPolicy::FirstAnthropicSemanticEvent { max_bytes, .. } = self.policy else {
return StreamPrecommitObservation::Pending;
let (max_bytes, observation) = match self.policy {
StreamCommitPolicy::FirstAnthropicSemanticEvent { max_bytes, .. } => {
(max_bytes, self.anthropic.observe(chunk, max_bytes))
}
StreamCommitPolicy::FirstGeminiSemanticEvent { max_bytes, .. } => {
(max_bytes, self.gemini.observe(chunk, max_bytes))
}
StreamCommitPolicy::ResponseHeaders | StreamCommitPolicy::FirstClassifiedBody => {
return StreamPrecommitObservation::Pending;
}
};
self.observed_bytes = self.observed_bytes.saturating_add(chunk.len());
match self.anthropic.observe(chunk, max_bytes) {
AnthropicSseObservation::Pending => {}
AnthropicSseObservation::SemanticEvent => {
match observation {
SemanticSseObservation::Pending => {}
SemanticSseObservation::SemanticEvent => {
self.state = StreamCommitState::Committed;
return StreamPrecommitObservation::Commit;
}
AnthropicSseObservation::Error(body_json) => {
SemanticSseObservation::Error {
status_code,
body_json,
} => {
self.state = StreamCommitState::Terminal;
return StreamPrecommitObservation::UpstreamError {
status_code: anthropic_error_status_code(&body_json),
status_code,
body_json,
};
}
@@ -179,10 +211,10 @@ impl StreamCommitGate {
}
#[derive(Debug)]
enum AnthropicSseObservation {
enum SemanticSseObservation {
Pending,
SemanticEvent,
Error(Value),
Error { status_code: u16, body_json: Value },
}
#[derive(Debug, Default)]
@@ -191,7 +223,7 @@ struct AnthropicSsePrecommitInspector {
}
impl AnthropicSsePrecommitInspector {
fn observe(&mut self, chunk: &[u8], max_bytes: usize) -> AnthropicSseObservation {
fn observe(&mut self, chunk: &[u8], max_bytes: usize) -> SemanticSseObservation {
let remaining = max_bytes.saturating_sub(self.buffered.len());
let truncated = chunk.len() > remaining;
self.buffered
@@ -201,15 +233,44 @@ impl AnthropicSsePrecommitInspector {
let record = self.buffered[..record_end].to_vec();
self.buffered.drain(..record_end + separator_len);
match classify_anthropic_sse_record(&record) {
AnthropicSseObservation::Pending => {}
SemanticSseObservation::Pending => {}
decision => return decision,
}
}
if truncated {
AnthropicSseObservation::SemanticEvent
SemanticSseObservation::SemanticEvent
} else {
AnthropicSseObservation::Pending
SemanticSseObservation::Pending
}
}
}
#[derive(Debug, Default)]
struct GeminiSsePrecommitInspector {
buffered: Vec<u8>,
}
impl GeminiSsePrecommitInspector {
fn observe(&mut self, chunk: &[u8], max_bytes: usize) -> SemanticSseObservation {
let remaining = max_bytes.saturating_sub(self.buffered.len());
let truncated = chunk.len() > remaining;
self.buffered
.extend_from_slice(&chunk[..chunk.len().min(remaining)]);
while let Some((record_end, separator_len)) = find_sse_record_boundary(&self.buffered) {
let record = self.buffered[..record_end].to_vec();
self.buffered.drain(..record_end + separator_len);
match classify_gemini_sse_record(&record) {
SemanticSseObservation::Pending => {}
decision => return decision,
}
}
if truncated {
SemanticSseObservation::SemanticEvent
} else {
SemanticSseObservation::Pending
}
}
}
@@ -249,9 +310,9 @@ fn next_sse_line_ending(buffer: &[u8], start: usize) -> Option<(usize, usize)> {
Some((index, ending_len))
}
fn classify_anthropic_sse_record(record: &[u8]) -> AnthropicSseObservation {
fn classify_anthropic_sse_record(record: &[u8]) -> SemanticSseObservation {
let Ok(record) = std::str::from_utf8(record) else {
return AnthropicSseObservation::Pending;
return SemanticSseObservation::Pending;
};
let normalized_record = record.replace("\r\n", "\n").replace('\r', "\n");
let mut event_type = None;
@@ -275,15 +336,18 @@ fn classify_anthropic_sse_record(record: &[u8]) -> AnthropicSseObservation {
}
}
if data.trim().is_empty() {
return AnthropicSseObservation::Pending;
return SemanticSseObservation::Pending;
}
let Ok(body_json) = serde_json::from_str::<Value>(data.trim()) else {
return AnthropicSseObservation::Pending;
return SemanticSseObservation::Pending;
};
let payload_type = body_json.get("type").and_then(Value::as_str).map(str::trim);
if event_type == Some("error") || payload_type == Some("error") {
return AnthropicSseObservation::Error(body_json);
return SemanticSseObservation::Error {
status_code: anthropic_error_status_code(&body_json),
body_json,
};
}
let semantic_type = match (event_type, payload_type) {
@@ -292,12 +356,120 @@ fn classify_anthropic_sse_record(record: &[u8]) -> AnthropicSseObservation {
_ => None,
};
if semantic_type.is_some_and(is_anthropic_semantic_event_type) {
AnthropicSseObservation::SemanticEvent
SemanticSseObservation::SemanticEvent
} else {
AnthropicSseObservation::Pending
SemanticSseObservation::Pending
}
}
fn classify_gemini_sse_record(record: &[u8]) -> SemanticSseObservation {
let Ok(record) = std::str::from_utf8(record) else {
return SemanticSseObservation::Pending;
};
let data = record
.replace("\r\n", "\n")
.replace('\r', "\n")
.lines()
.filter_map(|line| line.strip_prefix("data:").map(str::trim_start))
.collect::<Vec<_>>()
.join("\n");
if data.trim().is_empty() {
return SemanticSseObservation::Pending;
}
if data.trim() == "[DONE]" {
return SemanticSseObservation::SemanticEvent;
}
let Ok(body_json) = serde_json::from_str::<Value>(data.trim()) else {
return SemanticSseObservation::Pending;
};
let response = body_json.get("response").unwrap_or(&body_json);
let Some(candidates) = response.get("candidates").and_then(Value::as_array) else {
return SemanticSseObservation::Pending;
};
for candidate in candidates {
let finish_reason = candidate
.get("finishReason")
.or_else(|| candidate.get("finish_reason"))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty());
if let Some(finish_reason) = finish_reason.filter(|reason| {
matches!(
*reason,
"MALFORMED_FUNCTION_CALL"
| "UNEXPECTED_TOOL_CALL"
| "TOO_MANY_TOOL_CALLS"
| "MISSING_THOUGHT_SIGNATURE"
| "MALFORMED_RESPONSE"
)
}) {
let message = candidate
.get("finishMessage")
.or_else(|| candidate.get("finish_message"))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.unwrap_or_else(|| format!("Gemini stream ended with {finish_reason}"));
return SemanticSseObservation::Error {
status_code: 502,
body_json: serde_json::json!({
"error": {
"type": "upstream_gemini_finish_error",
"code": finish_reason,
"message": message,
"upstream_status": 200
}
}),
};
}
if finish_reason.is_some() {
return SemanticSseObservation::SemanticEvent;
}
let Some(parts) = candidate
.get("content")
.and_then(|content| content.get("parts"))
.and_then(Value::as_array)
else {
continue;
};
if parts.iter().any(gemini_part_is_client_semantic) {
return SemanticSseObservation::SemanticEvent;
}
}
SemanticSseObservation::Pending
}
fn gemini_part_is_client_semantic(part: &Value) -> bool {
let Some(part) = part.as_object() else {
return true;
};
if part
.keys()
.any(|key| !matches!(key.as_str(), "text" | "thought" | "thoughtSignature"))
{
return true;
}
if part.get("thought").and_then(Value::as_bool) == Some(true) {
return false;
}
if part.keys().all(|key| key == "thoughtSignature") {
return false;
}
if part
.get("text")
.and_then(Value::as_str)
.is_some_and(|text| !text.is_empty())
{
return true;
}
false
}
fn is_anthropic_semantic_event_type(event_type: &str) -> bool {
matches!(
event_type,
@@ -346,6 +518,13 @@ mod tests {
}
}
fn gemini_policy() -> StreamCommitPolicy {
StreamCommitPolicy::FirstGeminiSemanticEvent {
max_bytes: 16_384,
max_wait: Duration::from_millis(750),
}
}
#[test]
fn policy_selects_bounded_anthropic_gate_only_for_native_same_format_sse() {
let native = StreamCommitPolicy::for_response(
@@ -384,6 +563,109 @@ mod tests {
.commits_on_response_headers());
}
#[test]
fn policy_selects_bounded_gemini_gate_for_event_streams() {
let policy = StreamCommitPolicy::for_response(
true,
Some("text/event-stream"),
"gemini:generate_content",
"openai:responses",
false,
true,
false,
);
assert!(policy.is_gemini());
assert!(policy.requires_bounded_frame_wait());
assert_eq!(
policy.max_precommit_wait(),
Some(Duration::from_millis(750))
);
}
#[test]
fn gemini_gate_waits_through_thought_and_commits_on_text() {
let mut gate = StreamCommitGate::new(gemini_policy());
let thought = b"data: {\"response\":{\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"thought\":true,\"text\":\"checking\"}]}}]}}\n\n";
let text = b"data: {\"response\":{\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"text\":\"answer\"}]}}]}}\n\n";
assert_eq!(
gate.observe_provider_bytes(thought),
StreamPrecommitObservation::Pending
);
assert_eq!(
gate.observe_provider_bytes(text),
StreamPrecommitObservation::Commit
);
assert_eq!(gate.state(), StreamCommitState::Committed);
}
#[test]
fn gemini_gate_commits_on_function_call_even_with_thought_marker() {
let mut gate = StreamCommitGate::new(gemini_policy());
let tool_call = b"data: {\"response\":{\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"thought\":true,\"functionCall\":{\"name\":\"validate\",\"args\":{}}}]}}]}}\n\n";
assert_eq!(
gate.observe_provider_bytes(tool_call),
StreamPrecommitObservation::Commit
);
assert_eq!(gate.state(), StreamCommitState::Committed);
}
#[test]
fn gemini_gate_rejects_malformed_function_call_before_commit() {
let mut gate = StreamCommitGate::new(gemini_policy());
let thought = b"data: {\"response\":{\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"thought\":true,\"text\":\"calling\"}]}}]}}\n\n";
let malformed = b"data: {\"response\":{\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"thoughtSignature\":\"signature\",\"text\":\"\"}]},\"finishReason\":\"MALFORMED_FUNCTION_CALL\",\"finishMessage\":\"Malformed function call: Function call is empty - no input to parse.\"}]}}\n\n";
assert_eq!(
gate.observe_provider_bytes(thought),
StreamPrecommitObservation::Pending
);
let StreamPrecommitObservation::UpstreamError {
status_code,
body_json,
} = gate.observe_provider_bytes(malformed)
else {
panic!("malformed Gemini function call should fail before stream commit");
};
assert_eq!(status_code, 502);
assert_eq!(body_json["error"]["code"], "MALFORMED_FUNCTION_CALL");
assert_eq!(
body_json["error"]["message"],
"Malformed function call: Function call is empty - no input to parse."
);
assert_eq!(gate.state(), StreamCommitState::Terminal);
}
#[test]
fn gemini_gate_detects_malformed_function_call_across_chunk_boundaries() {
let malformed = b"data: {\"response\":{\"candidates\":[{\"content\":{\"role\":\"model\",\"parts\":[{\"thoughtSignature\":\"signature\",\"text\":\"\"}]},\"finishReason\":\"MALFORMED_FUNCTION_CALL\",\"finishMessage\":\"empty call\"}]}}\r\n\r\n";
for split in 1..malformed.len() {
let mut gate = StreamCommitGate::new(gemini_policy());
let first_observation = gate.observe_provider_bytes(&malformed[..split]);
if !matches!(
first_observation,
StreamPrecommitObservation::UpstreamError {
status_code: 502,
..
}
) {
assert_eq!(first_observation, StreamPrecommitObservation::Pending);
assert!(matches!(
gate.observe_provider_bytes(&malformed[split..]),
StreamPrecommitObservation::UpstreamError {
status_code: 502,
..
}
));
}
assert_eq!(gate.state(), StreamCommitState::Terminal);
}
}
#[test]
fn gate_detects_anthropic_error_across_every_chunk_boundary() {
let event = b"event: error\r\ndata: {\"type\":\"error\",\"error\":{\"type\":\"overloaded_error\",\"message\":\"busy\"}}\r\n\r\n";
@@ -71,6 +71,7 @@ use crate::ai_serving::api::{
UPSTREAM_IS_STREAM_KEY,
};
use crate::ai_serving::is_openai_responses_family_format;
use crate::ai_serving::record_local_runtime_candidate_skip_reason;
use crate::api::response::{
attach_control_metadata_headers, build_client_response, build_client_response_from_parts,
};
@@ -126,7 +127,7 @@ use crate::orchestration::{
LocalOAuthSuccessEffect, LocalPoolErrorEffect,
};
use crate::provider_pool_demand::{
acquire_provider_pool_in_flight_guard, ProviderPoolInFlightGuard,
acquire_provider_pool_execution_guard, ProviderPoolInFlightAdmission, ProviderPoolInFlightGuard,
};
use crate::request_candidate_runtime::{
ensure_execution_request_candidate_slot, persist_local_request_candidate_status_record,
@@ -3772,6 +3773,46 @@ async fn execute_execution_runtime_stream_inner(
plan_kind,
report_context.as_ref(),
);
let candidate_started_unix_secs = current_request_candidate_unix_ms();
let provider_in_flight_started_at = Instant::now();
let mut provider_pool_in_flight_guard =
match acquire_provider_pool_execution_guard(state, &plan).await? {
ProviderPoolInFlightAdmission::Acquired(guard) => guard,
ProviderPoolInFlightAdmission::Saturated { limit } => {
record_local_runtime_candidate_skip_reason(
state,
trace_id,
"provider_key_concurrency_limit_reached",
);
if let Some(retry_scope) = retry_scope_out.as_deref_mut() {
*retry_scope = AiAttemptRetryScope::Candidate;
}
if let Some(snapshot) = request_candidate_status_snapshot.as_ref() {
record_local_request_candidate_status_snapshot(
state,
snapshot,
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Skipped,
status_code: Some(http::StatusCode::TOO_MANY_REQUESTS.as_u16()),
error_type: Some("provider_key_concurrency_limit_reached".to_string()),
error_message: Some(format!(
"provider key concurrency limit reached: {limit}"
)),
latency_ms: Some(0),
started_at_unix_ms: Some(candidate_started_unix_secs),
finished_at_unix_ms: Some(candidate_started_unix_secs),
},
)
.await;
}
return Ok(None);
}
};
observe_gateway_stage_trace_ms(
&mut stage_trace,
"stream_provider_in_flight",
provider_in_flight_started_at.elapsed().as_millis() as u64,
);
// Inline passthrough records its lifecycle seed after upstream headers are
// available. Avoid constructing a throwaway seed on the common path.
let mut lifecycle_seed = (!defer_stream_pending_for_direct_inline)
@@ -3781,7 +3822,6 @@ async fn execute_execution_runtime_stream_inner(
record_stream_pending_lifecycle(state, seed, &mut stage_trace).await;
lifecycle_pending_recorded = true;
}
let candidate_started_unix_secs = current_request_candidate_unix_ms();
if let Some(snapshot) = request_candidate_status_snapshot.clone() {
record_local_request_candidate_status_snapshot(
state,
@@ -3810,20 +3850,6 @@ async fn execute_execution_runtime_stream_inner(
.and_then(|context| context.candidate_index)
.map(|value| value.to_string())
.unwrap_or_else(|| "-".to_string());
let provider_in_flight_started_at = Instant::now();
let mut provider_pool_in_flight_guard = acquire_provider_pool_in_flight_guard(
state.runtime_state.clone(),
&plan.provider_id,
plan.request_id.as_str(),
plan.candidate_id.as_deref(),
key_id.as_str(),
)
.await;
observe_gateway_stage_trace_ms(
&mut stage_trace,
"stream_provider_in_flight",
provider_in_flight_started_at.elapsed().as_millis() as u64,
);
match maybe_execute_grok_stream(&plan, report_context.as_ref()).await {
Ok(Some(grok_stream)) => {
return execute_stream_from_frame_stream_with_retry_scope(
@@ -6470,7 +6496,9 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
}
}
let inspection = if stream_commit_policy.is_native_anthropic() {
let inspection = if stream_commit_policy.is_native_anthropic()
|| stream_commit_policy.is_gemini()
{
StreamPrefetchInspection::NeedMore
} else {
inspect_prefetched_stream_body(
@@ -8211,7 +8239,8 @@ mod tests {
DirectPassthroughFinalizerCore, DirectPassthroughInlineBodyState, DirectPassthroughMode,
PostStopFrameReadBudget, PostStopLimitedStreamReader, ProviderStreamErrorInspection,
ANTHROPIC_POST_STOP_DRAIN_MAX_BYTES, GEMINI_FILES_DOWNLOAD_PLAN_KIND,
OPENAI_CHAT_STREAM_PLAN_KIND, POST_STOP_MAX_EMPTY_CHUNKS_PER_POLL,
OPENAI_CHAT_STREAM_PLAN_KIND, OPENAI_RESPONSES_STREAM_PLAN_KIND,
POST_STOP_MAX_EMPTY_CHUNKS_PER_POLL,
};
use crate::control::GatewayControlDecision;
use crate::stage_metrics::RequestStageTrace;
@@ -8756,6 +8785,36 @@ mod tests {
}
}
fn antigravity_gemini_stream_plan(request_id: &str) -> ExecutionPlan {
ExecutionPlan {
request_id: request_id.to_string(),
candidate_id: Some(format!("candidate-{request_id}")),
provider_name: Some("antigravity".to_string()),
provider_id: format!("provider-{request_id}"),
endpoint_id: format!("endpoint-{request_id}"),
key_id: format!("key-{request_id}"),
method: "POST".to_string(),
url: "https://cloudcode-pa.googleapis.com/v1internal:streamGenerateContent".to_string(),
headers: BTreeMap::from([
("content-type".to_string(), "application/json".to_string()),
("accept".to_string(), "text/event-stream".to_string()),
]),
content_type: Some("application/json".to_string()),
content_encoding: None,
body: RequestBody::from_json(json!({
"model": "gemini-3.7-flash-tiered",
"contents": [{"role": "user", "parts": [{"text": "validate"}]}]
})),
stream: true,
client_api_format: "openai:responses".to_string(),
provider_api_format: "gemini:generate_content".to_string(),
model_name: Some("gemini-3.7-flash-tiered".to_string()),
proxy: None,
transport_profile: None,
timeouts: None,
}
}
struct StreamDropFlag(Arc<AtomicBool>);
impl Drop for StreamDropFlag {
@@ -10504,6 +10563,92 @@ mod tests {
assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR);
}
#[tokio::test]
async fn malformed_antigravity_function_call_retries_before_stream_commit() {
let request_id = "req-antigravity-malformed-function-call";
let plan = antigravity_gemini_stream_plan(request_id);
let provider_catalog = provider_catalog_for_plan(
&plan,
Some(json!({
"failover_rules": {
"continue_status_codes": [502]
}
})),
);
let data_state = crate::data::GatewayDataState::with_provider_transport_reader_for_tests(
Arc::new(provider_catalog),
"development-key",
);
let state = AppState::new()
.expect("app state should build")
.with_data_state_for_tests(data_state);
let frame_stream = stream! {
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
frame_type: StreamFrameType::Headers,
payload: StreamFramePayload::Headers {
status_code: 200,
headers: BTreeMap::from([(
"content-type".to_string(),
"text/event-stream".to_string(),
)]),
response_observation: None,
},
}));
for chunk in [
r#"data: {"response":{"candidates":[{"content":{"role":"model","parts":[{"thought":true,"text":"Validating the document."}]} }],"modelVersion":"gemini-3.7-flash-tiered"}}
"#,
r#"data: {"response":{"candidates":[{"content":{"role":"model","parts":[{"thoughtSignature":"signature","text":""}]},"finishReason":"MALFORMED_FUNCTION_CALL","finishMessage":"Malformed function call: Function call is empty - no input to parse."}],"modelVersion":"gemini-3.7-flash-tiered"}}
"#,
] {
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
frame_type: StreamFrameType::Data,
payload: StreamFramePayload::Data {
chunk_b64: None,
text: Some(chunk.to_string()),
},
}));
}
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame::eof()));
}
.boxed();
let mut retry_scope = AiAttemptRetryScope::Provider;
let response = execute_stream_from_frame_stream_with_retry_scope(
&state,
plan,
"trace-antigravity-malformed-function-call",
&test_decision(),
OPENAI_RESPONSES_STREAM_PLAN_KIND,
Some("openai_responses_stream_success".to_string()),
Some(json!({
"request_id": request_id,
"candidate_id": format!("candidate-{request_id}"),
"candidate_index": 0,
"retry_index": 0,
"provider_api_format": "gemini:generate_content",
"client_api_format": "openai:responses",
"needs_conversion": true
})),
crate::clock::current_unix_ms(),
Instant::now(),
RequestStageTrace::from_env(),
true,
frame_stream,
false,
None,
Some(&mut retry_scope),
None,
None,
)
.await
.expect("malformed Antigravity stream should resolve through failover");
assert!(response.is_none());
assert_eq!(retry_scope, AiAttemptRetryScope::Candidate);
}
fn tunnel_proxy_snapshot(base_url: String) -> aether_contracts::ProxySnapshot {
aether_contracts::ProxySnapshot {
enabled: Some(true),
@@ -34,6 +34,7 @@ use crate::ai_serving::api::{
implicit_sync_finalize_report_kind, maybe_build_sync_finalize_outcome, LocalCoreSyncErrorKind,
LocalCoreSyncFinalizeOutcome,
};
use crate::ai_serving::record_local_runtime_candidate_skip_reason;
use crate::api::response::{
attach_control_metadata_headers, build_client_response, build_client_response_from_parts,
build_client_response_from_parts_with_mutator,
@@ -78,7 +79,9 @@ use crate::orchestration::{
LocalExecutionEffectContext, LocalHealthFailureEffect, LocalHealthSuccessEffect,
LocalOAuthInvalidationEffect, LocalOAuthSuccessEffect, LocalPoolErrorEffect,
};
use crate::provider_pool_demand::acquire_provider_pool_in_flight_guard;
use crate::provider_pool_demand::{
acquire_provider_pool_execution_guard, ProviderPoolInFlightAdmission,
};
use crate::request_candidate_runtime::{
ensure_execution_request_candidate_slot, record_local_request_candidate_extra_data,
record_local_request_candidate_status, record_local_request_candidate_status_snapshot,
@@ -1995,6 +1998,37 @@ async fn execute_execution_runtime_sync_impl(
.unwrap_or_else(|| "-".to_string());
let candidate_started_at = Instant::now();
let candidate_started_unix_secs = current_request_candidate_unix_ms();
let _provider_pool_in_flight_guard = match acquire_provider_pool_execution_guard(state, &plan)
.await?
{
ProviderPoolInFlightAdmission::Acquired(guard) => guard,
ProviderPoolInFlightAdmission::Saturated { limit } => {
record_local_runtime_candidate_skip_reason(
state,
trace_id,
"provider_key_concurrency_limit_reached",
);
if let Some(retry_scope) = retry_scope_out.as_deref_mut() {
*retry_scope = AiAttemptRetryScope::Candidate;
}
record_local_request_candidate_status(
state,
&plan,
report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Skipped,
status_code: Some(StatusCode::TOO_MANY_REQUESTS.as_u16()),
error_type: Some("provider_key_concurrency_limit_reached".to_string()),
error_message: Some(format!("provider key concurrency limit reached: {limit}")),
latency_ms: Some(0),
started_at_unix_ms: Some(candidate_started_unix_secs),
finished_at_unix_ms: Some(candidate_started_unix_secs),
},
)
.await;
return Ok(None);
}
};
let lifecycle_seed = build_lifecycle_usage_seed(&plan, report_context.as_ref());
let usage_data = state.usage_lifecycle_data_state().as_ref().clone();
state
@@ -2024,14 +2058,6 @@ async fn execute_execution_runtime_sync_impl(
candidate_started_at,
);
let result = (async {
let _provider_pool_in_flight_guard = acquire_provider_pool_in_flight_guard(
state.runtime_state.clone(),
&plan.provider_id,
plan_request_id.as_str(),
plan_candidate_id.as_deref(),
key_id.as_str(),
)
.await;
record_sync_execution_active(
state,
&plan,
@@ -3279,7 +3305,14 @@ fn maybe_build_implicit_sync_finalize_outcome(
body_base64: &Option<String>,
telemetry: &Option<ExecutionTelemetry>,
) -> Result<Option<ImplicitSyncFinalizeOutcome>, GatewayError> {
if status_code >= 400 || body_json.is_some() || body_base64.is_none() {
let needs_conversion = report_context
.as_ref()
.and_then(|value| value.get("needs_conversion"))
.and_then(serde_json::Value::as_bool)
.unwrap_or(false);
let has_captured_stream_body = body_json.is_none() && body_base64.is_some();
let has_cross_format_sync_body = needs_conversion && body_json.is_some();
if status_code >= 400 || (!has_captured_stream_body && !has_cross_format_sync_body) {
return Ok(None);
}
@@ -3461,6 +3494,137 @@ mod tests {
.with_execution_runtime_candidate(true)
}
#[tokio::test]
async fn implicit_sync_finalize_converts_chat_json_to_namespaced_responses() {
let decision = GatewayControlDecision::synthetic(
"/v1/responses",
Some("ai_public".to_string()),
Some("openai".to_string()),
Some("responses".to_string()),
Some("openai:responses".to_string()),
)
.with_execution_runtime_candidate(true);
let report_context = Some(json!({
"provider_api_format": "openai:chat",
"client_api_format": "openai:responses",
"needs_conversion": true,
"mapped_model": "qwen-upstream",
"original_request_body": {
"model": "qwen",
"tools": [{
"type": "namespace",
"name": "mcp__vulnerability_report",
"description": "reporting tools",
"tools": [{
"type": "function",
"name": "vulnerability_report",
"description": "write the confirmed report",
"parameters": {
"type": "object",
"properties": {
"report_path": {"type": "string"}
},
"required": ["report_path"]
},
"strict": true
}]
}]
}
}));
let provider_body = Some(json!({
"id": "chatcmpl_namespace_sync",
"object": "chat.completion",
"created": 1_777_777_777,
"model": "qwen-upstream",
"choices": [{
"index": 0,
"message": {
"role": "assistant",
"content": null,
"tool_calls": [{
"id": "call_report_1",
"type": "function",
"function": {
"name": "vulnerability_report",
"arguments": "{\"report_path\":\"reports/sql-001-c1.md\"}"
}
}]
},
"finish_reason": "tool_calls"
}],
"usage": {
"prompt_tokens": 10,
"completion_tokens": 4,
"total_tokens": 14
}
}));
let implicit = maybe_build_implicit_sync_finalize_outcome(
"trace-namespace-sync",
&decision,
"openai_responses_sync",
&report_context,
StatusCode::OK.as_u16(),
&BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
&provider_body,
&None,
&None,
)
.expect("cross-format sync JSON finalize should not error")
.expect("cross-format sync JSON should be finalized");
let response_body = axum::body::to_bytes(implicit.outcome.response.into_body(), usize::MAX)
.await
.expect("response body should read");
let response_json: Value =
serde_json::from_slice(&response_body).expect("response body should be JSON");
assert_eq!(response_json["object"], "response");
assert!(response_json.get("choices").is_none());
assert_eq!(response_json["output"][0]["type"], "function_call");
assert_eq!(response_json["output"][0]["name"], "vulnerability_report");
assert_eq!(
response_json["output"][0]["namespace"],
"mcp__vulnerability_report"
);
assert_eq!(response_json["output"][0]["call_id"], "call_report_1");
}
#[test]
fn implicit_sync_finalize_leaves_same_format_json_on_passthrough_path() {
let report_context = Some(json!({
"provider_api_format": "openai:responses",
"client_api_format": "openai:responses",
"needs_conversion": false
}));
let body_json = Some(json!({
"id": "resp_same_format",
"object": "response",
"status": "completed",
"output": []
}));
let outcome = maybe_build_implicit_sync_finalize_outcome(
"trace-same-format-sync",
&GatewayControlDecision::synthetic(
"/v1/responses",
Some("ai_public".to_string()),
Some("openai".to_string()),
Some("responses".to_string()),
Some("openai:responses".to_string()),
),
"openai_responses_sync",
&report_context,
StatusCode::OK.as_u16(),
&BTreeMap::new(),
&body_json,
&None,
&None,
)
.expect("same-format sync JSON guard should not error");
assert!(outcome.is_none());
}
#[tokio::test]
async fn oversized_upstream_response_builds_claude_502_retry_fallback() {
let mut plan = test_openai_image_plan(false);
@@ -5855,6 +5855,100 @@ mod tests {
);
}
#[tokio::test]
async fn direct_sync_execution_runtime_preserves_gemini_tool_config_on_wire() {
let listener = crate::test_support::bind_loopback_listener()
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should resolve");
let captured_body = Arc::new(Mutex::new(None));
let captured_body_for_handler = Arc::clone(&captured_body);
let app = Router::new().route(
"/generate",
post(move |body: Bytes| {
let captured_body = Arc::clone(&captured_body_for_handler);
async move {
*captured_body
.lock()
.expect("capture lock should not be poisoned") = Some(body.to_vec());
Json(json!({"ok": true}))
}
}),
);
let server = tokio::spawn(async move {
axum::serve(listener, app)
.await
.expect("test server should run");
});
let result = DirectSyncExecutionRuntime::new()
.execute_sync(&ExecutionPlan {
request_id: "req-gemini-tool-config-wire".into(),
candidate_id: Some("cand-gemini-tool-config-wire".into()),
provider_name: Some("google".into()),
provider_id: "prov-gemini-tool-config-wire".into(),
endpoint_id: "ep-gemini-tool-config-wire".into(),
key_id: "key-gemini-tool-config-wire".into(),
method: "POST".into(),
url: format!("http://{addr}/generate"),
headers: BTreeMap::from([("content-type".into(), "application/json".into())]),
content_type: Some("application/json".into()),
content_encoding: None,
body: RequestBody::from_json(json!({
"model": "gemini-3-flash-preview",
"contents": [{
"role": "user",
"parts": [{"text": "Search, then save the result."}]
}],
"tools": [
{"googleSearch": {}},
{"functionDeclarations": [{
"name": "save_result",
"parameters": {
"type": "OBJECT",
"properties": {"result": {"type": "STRING"}}
}
}]}
],
"toolConfig": {
"includeServerSideToolInvocations": true,
"functionCallingConfig": {"mode": "ANY"}
}
})),
stream: false,
client_api_format: "openai:responses".into(),
provider_api_format: "gemini:generate_content".into(),
model_name: Some("gemini-3-flash-preview".into()),
proxy: None,
transport_profile: None,
timeouts: Some(ExecutionTimeouts {
connect_ms: Some(5_000),
total_ms: Some(LOCAL_HTTP_SUCCESS_TIMEOUT_MS),
..ExecutionTimeouts::default()
}),
})
.await
.expect("sync execution should succeed");
server.abort();
assert_eq!(result.status_code, 200);
let body = captured_body
.lock()
.expect("capture lock should not be poisoned")
.take()
.and_then(|body| serde_json::from_slice::<serde_json::Value>(&body).ok())
.expect("upstream should receive a JSON body");
assert_eq!(
body["toolConfig"]["includeServerSideToolInvocations"],
json!(true)
);
assert_eq!(body["toolConfig"]["functionCallingConfig"]["mode"], "ANY");
assert!(body["toolConfig"]
.get("include_server_side_tool_invocations")
.is_none());
}
#[tokio::test]
async fn direct_sync_execution_runtime_applies_non_stream_total_timeout_to_body() {
let listener = crate::test_support::bind_loopback_listener()
@@ -111,6 +111,22 @@ impl LocalExecutionRuntimeMissContext {
})
}
pub(crate) fn all_candidates_skipped_for_reasons(&self, reasons: &[&str]) -> bool {
if reasons.is_empty() || self.candidate_contexts.is_empty() {
return false;
}
self.candidate_contexts.iter().all(|candidate| {
candidate.candidate.status == RequestCandidateStatus::Skipped
&& candidate
.candidate
.skip_reason
.as_deref()
.map(str::trim)
.is_some_and(|value| reasons.contains(&value))
})
}
pub(crate) fn candidate_summary(&self) -> Option<String> {
const MAX_ITEMS: usize = 5;
@@ -167,6 +167,14 @@ fn parse_pool_score_rules(pool_advanced: &Map<String, Value>) -> PoolMemberScore
fn normalize_pool_preset_mode(preset: &str, raw_mode: Option<&Value>) -> Option<String> {
match preset {
"cache_affinity" => Some(
raw_mode
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| matches!(*value, "single_account" | "lru"))
.unwrap_or("single_account")
.to_string(),
),
"free_team_first" | "free_first" | "team_first" | "plus_first" | "pro_first" => {
let default_mode = match preset {
"free_team_first" => "both",
@@ -1342,6 +1342,7 @@ pub(super) fn build_admin_pool_key_payload(
json!(key.internal_priority),
);
payload.insert("rpm_limit".to_string(), json!(key.rpm_limit));
payload.insert("concurrent_limit".to_string(), json!(key.concurrent_limit));
payload.insert(
"cache_ttl_minutes".to_string(),
json!(key.cache_ttl_minutes),
@@ -3177,6 +3177,7 @@ async fn provider_query_execute_standard_test_candidate(
crate::ai_serving::openai_responses_reasoning_replay_policy(
transport.provider.provider_type.as_str(),
transport.endpoint.base_url.as_str(),
request_model,
),
)
else {
@@ -3238,6 +3239,7 @@ async fn provider_query_execute_standard_test_candidate(
crate::ai_serving::openai_responses_reasoning_replay_policy(
transport.provider.provider_type.as_str(),
transport.endpoint.base_url.as_str(),
request_model,
),
)
.is_err()
@@ -34,6 +34,22 @@ fn admin_user_id_from_billing_path(request_path: &str, suffix: &str) -> Option<S
}
}
fn admin_user_entitlement_ids_from_path(request_path: &str) -> Option<(String, String)> {
let rest = request_path
.trim_end_matches('/')
.strip_prefix("/api/admin/users/")?;
let mut parts = rest.split('/');
let user_id = parts.next()?.trim();
if parts.next()? != "billing" || parts.next()? != "entitlements" {
return None;
}
let entitlement_id = parts.next()?.trim();
if user_id.is_empty() || entitlement_id.is_empty() || parts.next().is_some() {
return None;
}
Some((user_id.to_string(), entitlement_id.to_string()))
}
fn admin_user_billing_operator_id(request_context: &AdminRequestContext<'_>) -> Option<String> {
request_context
.decision()
@@ -202,6 +218,60 @@ pub(in super::super) async fn build_admin_list_user_billing_entitlements_respons
}
}
pub(in super::super) async fn build_admin_revoke_user_billing_entitlement_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
) -> Result<Response<Body>, GatewayError> {
let Some((user_id, entitlement_id)) =
admin_user_entitlement_ids_from_path(request_context.path())
else {
return Ok(build_admin_users_bad_request_response("缺少套餐权益 ID"));
};
if state.find_user_auth_by_id(&user_id).await?.is_none() {
return Ok((
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "用户不存在" })),
)
.into_response());
}
match state
.app()
.revoke_user_plan_entitlement(&user_id, &entitlement_id)
.await?
{
crate::LocalMutationOutcome::Applied(()) => {}
crate::LocalMutationOutcome::NotFound => {
return Ok((
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "套餐权益不存在或已失效" })),
)
.into_response());
}
crate::LocalMutationOutcome::Invalid(detail) => {
return Ok(build_admin_users_bad_request_response(detail));
}
crate::LocalMutationOutcome::Unavailable => {
return Ok(build_admin_users_data_unavailable_response());
}
}
let entitlements = match load_admin_user_entitlements_payload(state, &user_id).await? {
Some(value) => value,
None => return Ok(build_admin_users_data_unavailable_response()),
};
Ok(attach_admin_audit_response(
Json(json!({
"items": entitlements["items"].clone(),
"entitlements": entitlements["items"].clone(),
"total": entitlements["total"].clone(),
}))
.into_response(),
"admin_user_plan_revoked",
"revoke_user_billing_entitlement",
"user_plan_entitlement",
&entitlement_id,
))
}
pub(in super::super) async fn build_admin_grant_user_billing_plan_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
@@ -28,6 +28,7 @@ use self::batch::{
use self::billing::{
build_admin_grant_user_billing_plan_response,
build_admin_list_user_billing_entitlements_response,
build_admin_revoke_user_billing_entitlement_response,
};
use self::groups::{
build_admin_create_user_group_response, build_admin_delete_user_group_response,
@@ -8,10 +8,11 @@ use super::{
build_admin_list_user_group_members_response, build_admin_list_user_groups_response,
build_admin_list_user_sessions_response, build_admin_list_users_response,
build_admin_replace_user_group_members_response, build_admin_resolve_user_selection_response,
build_admin_reveal_user_api_key_response, build_admin_set_default_user_group_response,
build_admin_toggle_user_api_key_lock_response, build_admin_update_user_api_key_response,
build_admin_update_user_group_response, build_admin_update_user_response,
build_admin_user_batch_action_response, build_admin_users_data_unavailable_response,
build_admin_reveal_user_api_key_response, build_admin_revoke_user_billing_entitlement_response,
build_admin_set_default_user_group_response, build_admin_toggle_user_api_key_lock_response,
build_admin_update_user_api_key_response, build_admin_update_user_group_response,
build_admin_update_user_response, build_admin_user_batch_action_response,
build_admin_users_data_unavailable_response,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
@@ -58,6 +59,10 @@ fn is_admin_users_route(request_context: &AdminRequestContext<'_>) -> bool {
&& path.starts_with("/api/admin/users/")
&& path.ends_with("/billing/grant-plan")
&& path.matches('/').count() == 6)
|| (request_context.method() == http::Method::DELETE
&& path.starts_with("/api/admin/users/")
&& path.contains("/billing/entitlements/")
&& path.matches('/').count() == 7)
|| ((request_context.method() == http::Method::GET
|| request_context.method() == http::Method::PUT
|| request_context.method() == http::Method::DELETE)
@@ -155,6 +160,9 @@ pub(super) async fn maybe_build_local_admin_users_routes_response(
build_admin_grant_user_billing_plan_response(state, request_context, request_body)
.await?,
)),
Some("revoke_user_billing_entitlement") => Ok(Some(
build_admin_revoke_user_billing_entitlement_response(state, request_context).await?,
)),
Some("get_user") => Ok(Some(
build_admin_get_user_response(state, request_context).await?,
)),
+94 -11
View File
@@ -105,6 +105,12 @@ const LOCAL_EXECUTION_LOOP_DETECTED_DETAIL: &str =
"Gateway detected an execution runtime request loop back into the local frontdoor";
const AUTH_API_KEY_CONCURRENCY_LIMIT_REACHED_DETAIL: &str =
"当前调用方 API Key 并发请求数已达上限,请稍后重试";
const PROVIDER_KEY_CAPACITY_LIMIT_REACHED_DETAIL: &str =
"所有可用上游账号当前均已达到并发或 RPM 上限,请稍后重试";
const PROVIDER_KEY_CAPACITY_LIMIT_SKIP_REASONS: &[&str] = &[
"provider_key_concurrency_limit_reached",
"key_rpm_exhausted",
];
const LOCAL_EXECUTION_PLANNING_TIMEOUT_DETAIL: &str =
"当前 AI 请求在本地执行规划阶段超时,请稍后重试";
const EXECUTION_PATH_TUNNEL_AFFINITY_FORWARD: &str = "tunnel_affinity_forward";
@@ -1095,6 +1101,7 @@ async fn proxy_request_inner(
),
}
let (mut parts, body) = request.into_parts();
crate::ai_serving::codex_context::install_codex_fingerprint_context_slot(&mut parts);
let redaction_slot = crate::privacy::RedactionSessionSlot::default();
parts.extensions.insert(redaction_slot.clone());
parts
@@ -1350,6 +1357,7 @@ async fn proxy_request_inner(
.extensions
.get::<crate::middleware::CfConnectingIp>()
.map(|value| value.0.as_str()),
client_ip,
local_proxy_body.as_ref(),
)
.await
@@ -1913,12 +1921,23 @@ async fn proxy_request_inner(
.all_candidates_skipped_for_reason(AUTH_API_KEY_CONCURRENCY_LIMIT_SKIP_REASON)
|| local_execution_runtime_miss_context
.all_candidates_skipped_for_reason(LEGACY_API_KEY_CONCURRENCY_LIMIT_SKIP_REASON);
let local_execution_runtime_miss_detail = (!auth_api_key_concurrency_limited)
.then(|| {
let provider_key_capacity_limited = local_execution_runtime_miss_diagnostic
.as_ref()
.map(|diagnostic| diagnostic_is_provider_key_capacity_limited(Some(diagnostic)))
.unwrap_or_else(|| {
local_execution_runtime_miss_context
.all_provider_request_body_build_failures_detail()
.all_candidates_skipped_for_reasons(PROVIDER_KEY_CAPACITY_LIMIT_SKIP_REASONS)
});
let local_execution_runtime_miss_detail = provider_key_capacity_limited
.then_some(PROVIDER_KEY_CAPACITY_LIMIT_REACHED_DETAIL.to_string())
.or_else(|| {
(!auth_api_key_concurrency_limited)
.then(|| {
local_execution_runtime_miss_context
.all_provider_request_body_build_failures_detail()
})
.flatten()
})
.flatten()
.or_else(|| {
local_execution_runtime_miss_detail(
control_decision,
@@ -2034,7 +2053,7 @@ async fn proxy_request_inner(
let mut response = build_local_http_error_response(
&trace_id,
control_decision,
http::StatusCode::SERVICE_UNAVAILABLE,
local_execution_runtime_miss_status(provider_key_capacity_limited),
local_execution_runtime_miss_client_message(
local_execution_runtime_miss_detail.as_str(),
)
@@ -2360,6 +2379,30 @@ fn diagnostic_is_auth_api_key_concurrency_limited(
}))
}
fn diagnostic_is_provider_key_capacity_limited(
diagnostic: Option<&LocalExecutionRuntimeMissDiagnostic>,
) -> bool {
let Some(diagnostic) = diagnostic else {
return false;
};
PROVIDER_KEY_CAPACITY_LIMIT_SKIP_REASONS.contains(&diagnostic.reason.as_str())
|| (diagnostic.candidate_count.is_some_and(|candidate_count| {
candidate_count > 0
&& diagnostic.skipped_candidate_count.unwrap_or(0) >= candidate_count
}) && !diagnostic.skip_reasons.is_empty()
&& diagnostic.skip_reasons.iter().all(|(reason, count)| {
PROVIDER_KEY_CAPACITY_LIMIT_SKIP_REASONS.contains(&reason.as_str()) && *count > 0
}))
}
fn local_execution_runtime_miss_status(provider_key_capacity_limited: bool) -> http::StatusCode {
if provider_key_capacity_limited {
http::StatusCode::TOO_MANY_REQUESTS
} else {
http::StatusCode::SERVICE_UNAVAILABLE
}
}
fn local_execution_runtime_miss_route_detail(
decision: Option<&GatewayControlDecision>,
) -> Option<&'static str> {
@@ -2398,14 +2441,15 @@ mod tests {
use super::{
api_key_remote_ip_allowed, buffer_and_normalize_request_body,
diagnostic_is_auth_api_key_concurrency_limited, local_execution_runtime_miss_detail,
owner_forward_request_is_stream, restore_redacted_stream_execution_response,
restore_redacted_sync_execution_response, routing_overlay_allows_affinity_target,
GatewayControlDecision, LocalExecutionRuntimeMissDiagnostic, RequestBodyBufferError,
RequestBodyBufferPolicy,
diagnostic_is_auth_api_key_concurrency_limited,
diagnostic_is_provider_key_capacity_limited, local_execution_runtime_miss_detail,
local_execution_runtime_miss_status, owner_forward_request_is_stream,
restore_redacted_stream_execution_response, restore_redacted_sync_execution_response,
routing_overlay_allows_affinity_target, GatewayControlDecision,
LocalExecutionRuntimeMissDiagnostic, RequestBodyBufferError, RequestBodyBufferPolicy,
};
use axum::body::{to_bytes, Body, Bytes};
use axum::http::{header, HeaderMap, HeaderValue, Method, Response};
use axum::http::{header, HeaderMap, HeaderValue, Method, Response, StatusCode};
use serde_json::json;
use tokio::sync::Semaphore;
@@ -2885,6 +2929,45 @@ mod tests {
Some("当前调用方 API Key 并发请求数已达上限,请稍后重试")
);
}
#[test]
fn provider_key_capacity_requires_every_skip_reason_to_be_capacity_related() {
let capacity_limited = LocalExecutionRuntimeMissDiagnostic {
reason: "candidate_evaluation_incomplete".to_string(),
candidate_count: Some(2),
skipped_candidate_count: Some(2),
skip_reasons: std::collections::BTreeMap::from([
("provider_key_concurrency_limit_reached".to_string(), 1),
("key_rpm_exhausted".to_string(), 1),
]),
..LocalExecutionRuntimeMissDiagnostic::default()
};
let mixed_failure = LocalExecutionRuntimeMissDiagnostic {
reason: "all_candidates_skipped".to_string(),
candidate_count: Some(2),
skipped_candidate_count: Some(2),
skip_reasons: std::collections::BTreeMap::from([
("provider_key_concurrency_limit_reached".to_string(), 1),
("account_quota_exhausted".to_string(), 1),
]),
..LocalExecutionRuntimeMissDiagnostic::default()
};
assert!(diagnostic_is_provider_key_capacity_limited(Some(
&capacity_limited
)));
assert!(!diagnostic_is_provider_key_capacity_limited(Some(
&mixed_failure
)));
assert_eq!(
local_execution_runtime_miss_status(true),
StatusCode::TOO_MANY_REQUESTS
);
assert_eq!(
local_execution_runtime_miss_status(false),
StatusCode::SERVICE_UNAVAILABLE
);
}
}
#[path = "finalize.rs"]
@@ -1062,13 +1062,28 @@ mod tests {
#[tokio::test]
async fn codex_realtime_call_remains_on_the_live_handler() {
let decision = GatewayControlDecision::synthetic(
let mut decision = GatewayControlDecision::synthetic(
"/v1/realtime/calls",
Some("ai_public".to_string()),
Some("codex".to_string()),
Some("live".to_string()),
Some("codex:live".to_string()),
);
decision.auth_context = Some(GatewayControlAuthContext {
user_id: "user-codex-realtime".to_string(),
api_key_id: "key-codex-realtime".to_string(),
username: Some("codex-realtime".to_string()),
api_key_name: Some("codex-realtime".to_string()),
balance_remaining: None,
access_allowed: true,
user_rate_limit: None,
api_key_rate_limit: None,
api_key_is_standalone: true,
admin_bypass_limits: false,
local_rejection: None,
allowed_models: None,
ip_rules: None,
});
let request_context = GatewayPublicRequestContext {
trace_id: "trace-codex-realtime-call".to_string(),
request_method: http::Method::POST,
@@ -49,6 +49,8 @@ pub(super) enum LiveAuthMode {
pub(super) struct PlannedLiveCandidate {
pub(super) execution: AiExecutionDecision,
pub(super) pinned_candidate: ResponsesWebSocketPinnedCandidate,
pub(super) codex_fingerprint_context:
aether_provider_transport::CodexFingerprintConvergenceContext,
pub(super) client_model: String,
pub(super) provider_model: String,
pub(super) auth_mode: LiveAuthMode,
@@ -228,8 +230,9 @@ async fn plan_live_candidate_inner(
if validate_model(client_model).is_err() || client_model.len() > MAX_LIVE_MODEL_BYTES {
return Ok(None);
}
let parts = build_live_planning_parts(headers, remote_addr);
let mut parts = build_live_planning_parts(headers, remote_addr);
let body = json!({"model": client_model, "input": []});
crate::ai_serving::codex_context::install_codex_fingerprint_context_slot(&mut parts);
let execution = maybe_build_pinned_stream_local_same_format_provider_decision_payload(
state,
&parts,
@@ -338,6 +341,8 @@ async fn plan_live_candidate_inner(
Ok(Some(PlannedLiveCandidate {
execution,
pinned_candidate,
codex_fingerprint_context:
crate::ai_serving::codex_context::resolve_codex_fingerprint_context(&parts, &body),
client_model: client_model.to_string(),
provider_model,
auth_mode,
@@ -560,7 +565,11 @@ pub(super) fn build_live_stream_admission_attempt(
remote_addr: &SocketAddr,
upstream_url: String,
) -> Result<Option<AiStreamAttempt>, GatewayError> {
let parts = build_live_planning_parts(headers, remote_addr);
let mut parts = build_live_planning_parts(headers, remote_addr);
crate::ai_serving::codex_context::restore_codex_logical_turn_context(
&mut parts,
&candidate.codex_fingerprint_context,
);
let body = json!({"model": candidate.client_model.as_str(), "input": []});
let mut execution = candidate.execution.clone();
execution.upstream_url = Some(upstream_url);
@@ -922,6 +931,11 @@ mod tests {
"key-1",
)
.unwrap(),
codex_fingerprint_context:
aether_provider_transport::CodexFingerprintConvergenceContext::new(
"test-live-turn",
1,
),
client_model: "global-model".to_string(),
provider_model: "provider-model".to_string(),
auth_mode,
@@ -718,6 +718,11 @@ mod tests {
PlannedLiveCandidate {
execution,
pinned_candidate: binding.pinned_candidate.clone(),
codex_fingerprint_context:
aether_provider_transport::CodexFingerprintConvergenceContext::new(
"test-live-turn",
1,
),
client_model: binding.client_model.clone(),
provider_model: binding.provider_model.clone(),
auth_mode: binding.auth_mode,
@@ -13,8 +13,9 @@ use crate::ai_serving::AiExecutionDecision;
use crate::clock::current_unix_secs;
use crate::handlers::proxy::websocket::transport::UpstreamWebSocketErrorCodes;
use crate::orchestration::{
codex_account_id_from_headers, codex_quota_exhaustion_reset_at,
sync_codex_websocket_quota_metadata, ResponsesWebSocketAdapter,
codex_account_id_from_headers, codex_model_quota_exhaustion_reset_at,
codex_quota_exhaustion_reset_at, sync_codex_websocket_quota_metadata,
ResponsesWebSocketAdapter,
};
use crate::AppState;
@@ -114,16 +115,37 @@ impl ResponsesWebSocketProtocolAdapter for CodexResponsesWebSocketAdapter {
event: &Value,
) -> Option<ResponsesWebSocketAdapterObservation> {
let rate_limits = parse_codex_rate_limits(event)?;
let exhausted =
let account_exhausted =
aether_admin::provider::quota::codex_rate_limit_metadata_exhausted(&rate_limits);
let retry_exclusion_until_unix_secs =
codex_quota_exhaustion_reset_at(&rate_limits, current_unix_secs());
let active_limit_exhausted =
aether_admin::provider::quota::codex_websocket_response_has_usage_limit_error(event);
let now_unix_secs = current_unix_secs();
let scoped_reset_at = if account_exhausted {
codex_quota_exhaustion_reset_at(&rate_limits, now_unix_secs)
} else {
codex_model_quota_exhaustion_reset_at(&rate_limits, now_unix_secs)
};
let retry_exclusion_until_unix_secs = active_limit_exhausted
.then(|| {
aether_admin::provider::quota::codex_websocket_usage_limit_reset_at(
event,
now_unix_secs,
)
})
.flatten()
.or(scoped_reset_at);
Some(ResponsesWebSocketAdapterObservation {
drain: exhausted.then_some(ResponsesWebSocketDrainDirective {
error_code: "codex_account_quota_exhausted",
retry_current_turn: true,
retry_exclusion_until_unix_secs,
}),
drain: (account_exhausted || active_limit_exhausted).then_some(
ResponsesWebSocketDrainDirective {
error_code: if account_exhausted {
"codex_account_quota_exhausted"
} else {
"codex_active_limit_exhausted"
},
retry_current_turn: true,
retry_exclusion_until_unix_secs,
},
),
quota_metadata: Some(rate_limits),
})
}
@@ -337,6 +359,91 @@ mod tests {
);
}
#[test]
fn model_scoped_usage_limit_error_drains_without_account_exhaustion() {
let adapter = CodexResponsesWebSocketAdapter;
let event = json!({
"type": "error",
"error": {
"type": "usage_limit_reached",
"plan_type": "pro",
},
"status_code": 429,
"headers": {
"X-Codex-Plan-Type": "pro",
"X-Codex-Active-Limit": "codex_bengalfox",
"X-Codex-Primary-Used-Percent": "100",
"X-Codex-Primary-Window-Minutes": "300",
"X-Codex-Primary-Reset-At": "4000000000",
"X-Codex-Bengalfox-Limit-Name": "GPT-5.3-Codex-Spark",
"X-Codex-Bengalfox-Primary-Used-Percent": "100",
"X-Codex-Bengalfox-Primary-Window-Minutes": "300",
"X-Codex-Bengalfox-Primary-Reset-At": "4000000000",
},
});
let observation = adapter
.observe_upstream_event(&event)
.expect("model-scoped quota error should be observed");
let drain = observation
.drain
.expect("model-scoped quota error should retry the current turn");
let quota = observation
.quota_metadata
.expect("model-scoped quota metadata should be retained");
assert_eq!(drain.error_code, "codex_active_limit_exhausted");
assert!(drain.retry_current_turn);
assert_eq!(
drain.retry_exclusion_until_unix_secs,
Some(4_000_000_000u64)
);
assert_eq!(quota["spark_primary_used_percent"], json!(100.0));
assert!(quota.get("allowed").is_none());
assert!(quota.get("limit_reached").is_none());
assert!(!aether_admin::provider::quota::codex_rate_limit_metadata_exhausted(&quota));
}
#[test]
fn account_scoped_usage_limit_error_keeps_account_drain_semantics() {
let adapter = CodexResponsesWebSocketAdapter;
let event = json!({
"type": "error",
"error": {
"type": "usage_limit_reached",
"plan_type": "free",
"resets_at": 4_000_000_000u64,
},
"status_code": 429,
"headers": {
"X-Codex-Plan-Type": "free",
"X-Codex-Primary-Used-Percent": "100",
"X-Codex-Primary-Window-Minutes": "43200",
"X-Codex-Primary-Reset-At": "4000000000",
},
});
let observation = adapter
.observe_upstream_event(&event)
.expect("account quota error should be observed");
let drain = observation
.drain
.expect("account quota error should retry the current turn");
let quota = observation
.quota_metadata
.expect("account quota metadata should be retained");
assert_eq!(drain.error_code, "codex_account_quota_exhausted");
assert!(drain.retry_current_turn);
assert_eq!(
drain.retry_exclusion_until_unix_secs,
Some(4_000_000_000u64)
);
assert_eq!(quota["allowed"], json!(false));
assert_eq!(quota["limit_reached"], json!(true));
assert!(aether_admin::provider::quota::codex_rate_limit_metadata_exhausted(&quota));
}
#[test]
fn only_known_codex_pre_response_signals_are_safe_to_rebind() {
let adapter = CodexResponsesWebSocketAdapter;
@@ -11,7 +11,7 @@ use aether_contracts::ExecutionPlan;
use crate::execution_runtime::acquire_upstream_execution_gate;
use crate::provider_pool_demand::{
acquire_provider_pool_in_flight_guard, ProviderPoolInFlightGuard,
acquire_provider_pool_execution_guard, ProviderPoolInFlightAdmission, ProviderPoolInFlightGuard,
};
use crate::upstream_admission::UpstreamTargetAdmissionPermit;
use crate::{AppState, GatewayError};
@@ -41,14 +41,17 @@ impl ResponsesWebSocketTurnAdmission {
return Err(error);
}
};
let provider_pool = acquire_provider_pool_in_flight_guard(
state.runtime_state.clone(),
&plan.provider_id,
&plan.request_id,
plan.candidate_id.as_deref(),
&plan.key_id,
)
.await;
let provider_pool = match acquire_provider_pool_execution_guard(state, plan).await? {
ProviderPoolInFlightAdmission::Acquired(guard) => guard,
ProviderPoolInFlightAdmission::Saturated { limit } => {
drop(upstream_target);
drop(upstream_execution);
return Err(GatewayError::Client {
status: http::StatusCode::TOO_MANY_REQUESTS,
message: format!("上游账号并发已达上限 ({limit})"),
});
}
};
Ok(Self {
upstream_execution,
@@ -1,5 +1,6 @@
//! Client-side Responses WebSocket event forwarding and follow-up planning.
use aether_provider_transport::CodexFingerprintConvergenceContext;
use axum::extract::ws::{Message as AxumWsMessage, WebSocket};
use futures_util::SinkExt;
use serde_json::Value;
@@ -248,7 +249,14 @@ pub(super) async fn forward_client_message(
// derive one strong live control snapshot that every stage below
// shares. The connection's Upgrade-time decision is only the
// immutable identity seed.
let planning_parts = build_planning_parts(context);
let logical_turn_id = Uuid::now_v7().to_string();
let mut planning_parts = build_planning_parts(context);
let codex_fingerprint_context =
crate::ai_serving::codex_context::attach_codex_logical_turn_context(
&mut planning_parts,
&client_event,
&logical_turn_id,
);
let turn_control = match resolve_responses_websocket_turn_control(
state,
context,
@@ -453,6 +461,8 @@ pub(super) async fn forward_client_message(
context,
planning_parts,
client_event,
logical_turn_id,
codex_fingerprint_context,
turn_control,
turn_redaction_session,
)
@@ -466,6 +476,8 @@ pub(super) async fn forward_client_message(
planning_parts,
client_event,
requested_model,
logical_turn_id,
codex_fingerprint_context,
turn_control,
raw_responses_lite_static_config
.expect("independent turns always retain their raw static config"),
@@ -513,6 +525,8 @@ async fn forward_pinned_continuation(
context: &WebSocketRequestContext,
planning_parts: http::request::Parts,
client_event: Value,
logical_turn_id: String,
codex_fingerprint_context: CodexFingerprintConvergenceContext,
turn_control: ResponsesWebSocketTurnControl,
turn_redaction_session: Option<RedactionSession>,
) -> RelayDisposition {
@@ -556,7 +570,6 @@ async fn forward_pinned_continuation(
};
let turn_request_id = Uuid::new_v4().to_string();
let logical_turn_id = Uuid::new_v4().to_string();
let planned = match await_owned_responses_websocket_plan(spawn_owned_responses_websocket_plan(
state.clone(),
planning_parts,
@@ -766,6 +779,7 @@ async fn forward_pinned_continuation(
bound.body_normalization = normalization;
bound.turn_state.begin(
LogicalTurn::new(client_event, turn_index, logical_turn_id)
.with_codex_fingerprint_context(codex_fingerprint_context)
.with_provider_store(provider_event.get("store") == Some(&Value::Bool(true)))
.with_turn_control(turn_control),
turn,
@@ -800,12 +814,13 @@ async fn forward_replanned_response_create(
planning_parts: http::request::Parts,
client_event: Value,
requested_model: String,
logical_turn_id: String,
codex_fingerprint_context: CodexFingerprintConvergenceContext,
turn_control: ResponsesWebSocketTurnControl,
raw_responses_lite_static_config: ResponsesLiteStaticConfig,
turn_redaction_session: Option<RedactionSession>,
) -> RelayDisposition {
let turn_request_id = Uuid::new_v4().to_string();
let logical_turn_id = Uuid::new_v4().to_string();
let now_unix_secs = current_unix_secs();
let excluded_key_ids = bound.exhausted_exclusions.key_ids(now_unix_secs);
let excluded_codex_account_ids = bound.exhausted_exclusions.codex_account_ids(now_unix_secs);
@@ -992,6 +1007,7 @@ async fn forward_replanned_response_create(
bound.body_normalization = normalization;
bound.turn_state.begin(
LogicalTurn::new(client_event.clone(), turn_index, logical_turn_id.clone())
.with_codex_fingerprint_context(codex_fingerprint_context.clone())
.with_provider_store(provider_event.get("store") == Some(&Value::Bool(true)))
.with_turn_control(turn_control),
turn,
@@ -1076,6 +1092,7 @@ async fn forward_replanned_response_create(
bound.binding_identity = replacement.binding_identity;
bound.turn_state.begin(
LogicalTurn::new(client_event, turn_index, logical_turn_id)
.with_codex_fingerprint_context(codex_fingerprint_context)
.with_provider_store(provider_event.get("store") == Some(&Value::Bool(true)))
.with_turn_control(turn_control),
turn,
@@ -146,6 +146,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
};
let turn_index = active.turn_index;
let logical_turn_id = active.logical_turn_id.clone();
let codex_fingerprint_context = active.codex_fingerprint_context.clone();
let turn_attempt = active.turn_attempt;
let retry_exclusion_until_unix_secs = bound
@@ -154,7 +155,13 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
let exhausted_key = record_exhausted_bound_key(bound, retry_exclusion_until_unix_secs);
let exhausted_key_id = exhausted_key.as_ref().map(|(key_id, _)| key_id.clone());
let planning_parts = build_planning_parts(context);
let mut planning_parts = build_planning_parts(context);
if let Some(codex_fingerprint_context) = codex_fingerprint_context.as_ref() {
crate::ai_serving::codex_context::restore_codex_logical_turn_context(
&mut planning_parts,
codex_fingerprint_context,
);
}
let turn_request_id = Uuid::new_v4().to_string();
let now_unix_secs = current_unix_secs();
let excluded_key_ids = bound.exhausted_exclusions.key_ids(now_unix_secs);
@@ -444,7 +444,14 @@ async fn bootstrap_responses_websocket(
let raw_responses_lite_static_config =
ResponsesLiteStaticConfig::from_response_create(&first_event);
let planning_parts = build_planning_parts(context);
let first_logical_turn_id = Uuid::now_v7().to_string();
let mut planning_parts = build_planning_parts(context);
let first_codex_fingerprint_context =
crate::ai_serving::codex_context::attach_codex_logical_turn_context(
&mut planning_parts,
&first_event,
&first_logical_turn_id,
);
let turn_control = match resolve_responses_websocket_turn_control(
&state,
context,
@@ -860,7 +867,6 @@ async fn bootstrap_responses_websocket(
return None;
}
};
let first_logical_turn_id = Uuid::new_v4().to_string();
let first_turn_decision = prepare_responses_websocket_turn_decision(
&decision,
context.trace_id.clone(),
@@ -954,6 +960,7 @@ async fn bootstrap_responses_websocket(
}
bound.turn_state.begin(
LogicalTurn::new(first_event, 1, first_logical_turn_id)
.with_codex_fingerprint_context(first_codex_fingerprint_context)
.with_provider_store(first_provider_event.get("store") == Some(&Value::Bool(true)))
.with_turn_control(turn_control),
first_turn,
@@ -5,6 +5,7 @@
//! 非法组合只能靠调用点的 if 和「记得同时改另外两个字段」来避免。这里把它收敛成
//! 一个枚举:合法组合由类型保证,转换只能走受控 API。
use aether_provider_transport::CodexFingerprintConvergenceContext;
use serde_json::Value;
use super::control::ResponsesWebSocketTurnControl;
@@ -26,6 +27,9 @@ pub(super) struct LogicalTurn {
pub(super) provider_store: bool,
pub(super) turn_index: u64,
pub(super) logical_turn_id: String,
/// Immutable Codex client identity for every provider attempt belonging to
/// this logical turn. A transparent re-plan must never mint a new turn.
pub(super) codex_fingerprint_context: Option<CodexFingerprintConvergenceContext>,
pub(super) turn_attempt: u32,
pub(super) retry_attempted: bool,
pub(super) retry_unsafe_reason: Option<&'static str>,
@@ -42,6 +46,7 @@ impl LogicalTurn {
provider_store: false,
turn_index,
logical_turn_id,
codex_fingerprint_context: None,
turn_attempt: 1,
retry_attempted: false,
retry_unsafe_reason: None,
@@ -54,6 +59,14 @@ impl LogicalTurn {
self
}
pub(super) fn with_codex_fingerprint_context(
mut self,
context: CodexFingerprintConvergenceContext,
) -> Self {
self.codex_fingerprint_context = Some(context);
self
}
pub(super) fn with_provider_store(mut self, provider_store: bool) -> Self {
self.provider_store = provider_store;
self
@@ -28,5 +28,5 @@ pub(crate) use self::support::{
build_api_key_install_session_response, build_proxy_node_install_session_response,
build_unhandled_public_support_response, matches_model_mapping_for_models,
maybe_build_local_admin_announcements_response, maybe_build_local_public_support_response,
CreateApiKeyInstallSessionRequest,
vscodex_ws_proxy, CreateApiKeyInstallSessionRequest,
};
@@ -48,6 +48,8 @@ mod support_payment;
mod support_test_connection;
#[path = "support/user_me.rs"]
mod support_user_me;
#[path = "support/user_me_vscodex.rs"]
mod support_vscodex;
#[path = "support/wallet.rs"]
mod support_wallet;
@@ -89,6 +91,8 @@ use self::support_oauth::maybe_build_local_oauth_response;
use self::support_payment::maybe_build_local_payment_callback_response;
use self::support_test_connection::maybe_build_local_test_connection_response;
use self::support_user_me::maybe_build_local_users_me_response;
pub(crate) use self::support_vscodex::vscodex_ws_proxy;
use self::support_vscodex::{handle_users_me_vscodex_request, maybe_build_local_vscodex_response};
use self::support_wallet::{
build_wallet_balance_payload_for_auth_scope, build_wallet_balance_payload_for_user,
build_wallet_live_today_usage_payload_for_api_key,
@@ -121,6 +125,7 @@ pub(crate) async fn maybe_build_local_public_support_response(
request_context: &GatewayPublicRequestContext,
headers: &http::HeaderMap,
cf_connecting_ip: Option<&str>,
client_ip: std::net::IpAddr,
request_body: Option<&Bytes>,
) -> Option<Response<Body>> {
let decision = request_context.control_decision.as_ref()?;
@@ -192,6 +197,11 @@ pub(crate) async fn maybe_build_local_public_support_response(
.await;
}
if decision.route_family.as_deref() == Some("vscodex") {
return maybe_build_local_vscodex_response(state, request_context, client_ip, request_body)
.await;
}
if decision.route_family.as_deref() == Some("install") {
return maybe_build_local_install_response(state, request_context).await;
}
@@ -2,8 +2,9 @@ use super::{
auth_password_policy_level, base_url_from_request, build_auth_error_response,
build_auth_wallet_summary_payload, decrypt_catalog_secret_with_fallbacks,
encrypt_catalog_secret_with_fallbacks, handle_auth_me,
handle_users_me_api_key_install_session_create, query_param_optional_bool, query_param_value,
resolve_authenticated_local_user, sanitize_public_model_config_for_user, unix_secs_to_rfc3339,
handle_users_me_api_key_install_session_create, handle_users_me_vscodex_request,
query_param_optional_bool, query_param_value, resolve_authenticated_local_user,
sanitize_public_model_config_for_user, unix_secs_to_rfc3339,
users_me_api_key_install_sessions_path_matches, validate_auth_register_password, AppState,
AuthenticatedLocalUserContext, GatewayPublicRequestContext, PUBLIC_CAPABILITY_DEFINITIONS,
};
@@ -18,9 +18,10 @@ use super::{
handle_users_me_preferences_put, handle_users_me_providers_get, handle_users_me_referral_get,
handle_users_me_sessions_get, handle_users_me_update_session, handle_users_me_usage_active_get,
handle_users_me_usage_get, handle_users_me_usage_heatmap_get,
handle_users_me_usage_interval_timeline_get, users_me_api_key_capabilities_path_matches,
users_me_api_key_detail_path_matches, users_me_api_key_install_sessions_path_matches,
users_me_api_key_providers_path_matches, users_me_management_token_detail_path_matches,
handle_users_me_usage_interval_timeline_get, handle_users_me_vscodex_request,
users_me_api_key_capabilities_path_matches, users_me_api_key_detail_path_matches,
users_me_api_key_install_sessions_path_matches, users_me_api_key_providers_path_matches,
users_me_management_token_detail_path_matches,
users_me_management_token_regenerate_path_matches,
users_me_management_token_toggle_path_matches, users_me_management_tokens_root,
users_me_session_detail_path_matches, AppState, GatewayPublicRequestContext,
@@ -55,6 +56,14 @@ pub(crate) async fn maybe_build_local_users_me_response(
{
Some(handle_users_me_delete_other_sessions(state, request_context, headers).await)
}
Some(
"vscodex_devices_list"
| "vscodex_pairing_create"
| "vscodex_device_delete"
| "vscodex_ws_ticket_create",
) => Some(
handle_users_me_vscodex_request(state, request_context, headers, request_body).await,
),
Some("session_delete")
if users_me_session_detail_path_matches(&request_context.request_path) =>
{
@@ -0,0 +1,931 @@
use std::collections::HashMap;
use std::net::IpAddr;
use std::sync::{Arc, LazyLock, Mutex};
use std::time::Duration;
use axum::body::{Body, Bytes};
use axum::extract::{
ws::{CloseFrame as AxumCloseFrame, Message as AxumMessage, WebSocket, WebSocketUpgrade},
ConnectInfo, State,
};
use axum::http::{self, header};
use axum::response::{IntoResponse, Response};
use futures_util::{SinkExt, StreamExt};
use serde::Deserialize;
use serde_json::{json, Map, Value};
use tokio::sync::Semaphore;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::protocol::{
CloseFrame as TungsteniteCloseFrame, WebSocketConfig,
};
use tokio_tungstenite::tungstenite::Message as TungsteniteMessage;
use tracing::warn;
use super::{
build_auth_error_response, build_auth_json_response, module_available_from_env,
resolve_authenticated_local_user, AppState, GatewayPublicRequestContext,
};
const VSCODEX_ENABLED_ENV: &str = "AETHER_VSCODEX_ENABLED";
const VSCODEX_INTERNAL_URL_ENV: &str = "AETHER_VSCODEX_INTERNAL_URL";
const VSCODEX_INTERNAL_TOKEN_ENV: &str = "AETHER_VSCODEX_INTERNAL_TOKEN";
const VSCODEX_REQUEST_TIMEOUT: Duration = Duration::from_secs(15);
const VSCODEX_MAX_RESPONSE_BYTES: usize = 1024 * 1024;
const VSCODEX_WS_MAX_MESSAGE_BYTES: usize = 16 * 1024 * 1024;
const VSCODEX_WS_MAX_CONNECTIONS: usize = 256;
const VSCODEX_WS_MAX_CONNECTIONS_PER_IP: usize = 16;
const VSCODEX_DEVICE_PATH_PREFIX: &str = "/api/users/me/vscodex/devices/";
const VSCODEX_CLIENT_IP_HEADER: &str = "x-aether-client-ip";
static VSCODEX_HTTP_CLIENT: LazyLock<Result<reqwest::Client, reqwest::Error>> =
LazyLock::new(|| {
reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
});
static VSCODEX_WS_CONNECTIONS: LazyLock<Arc<Semaphore>> =
LazyLock::new(|| Arc::new(Semaphore::new(VSCODEX_WS_MAX_CONNECTIONS)));
static VSCODEX_WS_CONNECTIONS_BY_IP: LazyLock<Arc<VscodexWsIpConnectionLimiter>> =
LazyLock::new(|| {
Arc::new(VscodexWsIpConnectionLimiter::new(
VSCODEX_WS_MAX_CONNECTIONS_PER_IP,
))
});
#[derive(Debug)]
struct VscodexWsIpConnectionLimiter {
max_connections: usize,
active: Mutex<HashMap<IpAddr, usize>>,
}
impl VscodexWsIpConnectionLimiter {
fn new(max_connections: usize) -> Self {
Self {
max_connections: max_connections.max(1),
active: Mutex::new(HashMap::new()),
}
}
fn try_acquire(self: &Arc<Self>, client_ip: IpAddr) -> Option<VscodexWsIpConnectionPermit> {
let mut active = self
.active
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let current = active.get(&client_ip).copied().unwrap_or_default();
if current >= self.max_connections {
return None;
}
active.insert(client_ip, current.saturating_add(1));
Some(VscodexWsIpConnectionPermit {
limiter: Arc::clone(self),
client_ip,
})
}
fn release(&self, client_ip: IpAddr) {
let mut active = self
.active
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let Some(current) = active.get_mut(&client_ip) else {
return;
};
if *current <= 1 {
active.remove(&client_ip);
} else {
*current -= 1;
}
}
#[cfg(test)]
fn active_ip_count(&self) -> usize {
self.active
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.len()
}
}
#[derive(Debug)]
struct VscodexWsIpConnectionPermit {
limiter: Arc<VscodexWsIpConnectionLimiter>,
client_ip: IpAddr,
}
impl Drop for VscodexWsIpConnectionPermit {
fn drop(&mut self) {
self.limiter.release(self.client_ip);
}
}
#[derive(Debug)]
struct VscodexSidecarConfig {
base_url: reqwest::Url,
authorization: reqwest::header::HeaderValue,
http_client: reqwest::Client,
}
#[derive(Debug, Default, Deserialize)]
struct CreatePairingRequest {
name: Option<String>,
}
#[derive(Debug, Deserialize)]
struct CreateWsTicketRequest {
device_id: String,
}
#[derive(Debug, Deserialize)]
struct ExchangePairingRequest {
code: String,
name: Option<String>,
}
pub(crate) async fn vscodex_ws_proxy(
State(state): State<AppState>,
ConnectInfo(remote_addr): ConnectInfo<std::net::SocketAddr>,
ws: WebSocketUpgrade,
headers: http::HeaderMap,
) -> Response<Body> {
let request_permit = match state.try_acquire_request_permit().await {
Ok(value) => value,
Err(err) => {
warn!(error = ?err, "VS Codex WebSocket request admission rejected");
return build_auth_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"服务繁忙,请稍后重试",
false,
);
}
};
let client_ip = crate::headers::effective_client_ip(&headers, &remote_addr);
match state.admin_security_ip_blacklisted(client_ip).await {
Ok(true) => {
return build_auth_error_response(
http::StatusCode::FORBIDDEN,
"当前 IP 已被禁止访问",
false,
)
}
Ok(false) => {}
Err(err) => warn!(
client_ip = %client_ip,
error = ?err,
"VS Codex WebSocket IP blacklist check failed open"
),
}
let connection_permit = match Arc::clone(&VSCODEX_WS_CONNECTIONS).try_acquire_owned() {
Ok(value) => value,
Err(_) => {
return build_auth_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"VS Codex 连接数已达上限",
false,
)
}
};
// Only active connections have entries, and each already owns one of the 256 global slots.
let ip_connection_permit = match VSCODEX_WS_CONNECTIONS_BY_IP.try_acquire(client_ip) {
Some(value) => value,
None => {
warn!(
client_ip = %client_ip,
limit = VSCODEX_WS_MAX_CONNECTIONS_PER_IP,
"VS Codex per-IP WebSocket connection limit reached"
);
let mut response = build_auth_error_response(
http::StatusCode::TOO_MANY_REQUESTS,
"当前 IP 的 VS Codex 连接数已达上限",
false,
);
response
.headers_mut()
.insert(header::RETRY_AFTER, http::HeaderValue::from_static("1"));
return response;
}
};
let config = match load_vscodex_sidecar_config() {
Ok(Some(value)) => value,
Ok(None) => {
return build_auth_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"VS Codex 服务未启用",
false,
)
}
Err(detail) => {
warn!(error = %detail, "VS Codex WebSocket sidecar configuration is invalid");
return build_auth_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"VS Codex 服务配置不完整",
false,
);
}
};
let sidecar_url = match build_vscodex_websocket_url(&config.base_url) {
Ok(value) => value,
Err(detail) => {
warn!(error = %detail, "could not build VS Codex sidecar WebSocket URL");
return build_auth_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"VS Codex 服务配置不完整",
false,
);
}
};
let mut sidecar_request = match sidecar_url.as_str().into_client_request() {
Ok(value) => value,
Err(err) => {
warn!(error = %err, "could not build VS Codex sidecar WebSocket request");
return build_auth_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"VS Codex 服务配置不完整",
false,
);
}
};
for header_name in [header::ORIGIN, header::SEC_WEBSOCKET_PROTOCOL] {
if let Some(value) = headers.get(&header_name) {
sidecar_request
.headers_mut()
.insert(header_name, value.clone());
}
}
let mut sidecar_config = WebSocketConfig::default();
sidecar_config.max_message_size = Some(VSCODEX_WS_MAX_MESSAGE_BYTES);
sidecar_config.max_frame_size = Some(VSCODEX_WS_MAX_MESSAGE_BYTES);
let (sidecar_socket, sidecar_response) = match tokio::time::timeout(
VSCODEX_REQUEST_TIMEOUT,
tokio_tungstenite::connect_async_with_config(sidecar_request, Some(sidecar_config), true),
)
.await
{
Ok(Ok(value)) => value,
Ok(Err(err)) => {
warn!(error = %err, "VS Codex sidecar WebSocket handshake failed");
return build_auth_error_response(
http::StatusCode::BAD_GATEWAY,
"VS Codex 服务暂时不可用",
false,
);
}
Err(_) => {
warn!("VS Codex sidecar WebSocket handshake timed out");
return build_auth_error_response(
http::StatusCode::GATEWAY_TIMEOUT,
"VS Codex 服务请求超时",
false,
);
}
};
let selected_protocol = sidecar_response
.headers()
.get(header::SEC_WEBSOCKET_PROTOCOL)
.and_then(|value| value.to_str().ok())
.map(str::to_string);
let ws = ws
.max_message_size(VSCODEX_WS_MAX_MESSAGE_BYTES)
.max_frame_size(VSCODEX_WS_MAX_MESSAGE_BYTES);
let ws = match selected_protocol {
Some(protocol) => ws.protocols([protocol]),
None => ws,
};
drop(request_permit);
ws.on_upgrade(move |browser_socket| async move {
let _connection_permit = connection_permit;
bridge_vscodex_websockets(browser_socket, sidecar_socket, ip_connection_permit).await;
})
}
pub(super) async fn maybe_build_local_vscodex_response(
_state: &AppState,
request_context: &GatewayPublicRequestContext,
client_ip: std::net::IpAddr,
request_body: Option<&Bytes>,
) -> Option<Response<Body>> {
let decision = request_context.control_decision.as_ref()?;
if decision.route_family.as_deref() != Some("vscodex") {
return None;
}
if decision.route_kind.as_deref() != Some("pairing_exchange")
|| !matches!(
request_context.request_path.as_str(),
"/api/vscodex/pair" | "/api/vscodex/pair/"
)
{
return Some(build_auth_error_response(
http::StatusCode::NOT_FOUND,
"VS Codex 接口不存在",
false,
));
}
let config = match load_vscodex_sidecar_config() {
Ok(Some(value)) => value,
Ok(None) => {
return Some(build_auth_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"VS Codex 服务未启用",
false,
))
}
Err(detail) => {
warn!(
error = %detail,
"VS Codex sidecar configuration is invalid"
);
return Some(build_auth_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"VS Codex 服务配置不完整",
false,
));
}
};
let payload = match parse_pairing_exchange_request(request_body) {
Ok(value) => value,
Err(response) => return Some(response),
};
let url = match append_vscodex_sidecar_path(&config.base_url, &["v1", "pairings", "exchange"]) {
Ok(value) => value,
Err(detail) => {
warn!(error = %detail, "could not build VS Codex pairing exchange URL");
return Some(build_auth_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"VS Codex 服务配置不完整",
false,
));
}
};
let request =
build_authenticated_sidecar_request(&config, reqwest::Method::POST, url, Some(payload))
.header(VSCODEX_CLIENT_IP_HEADER, client_ip.to_string());
Some(send_vscodex_sidecar_request(request, "public", "pairing_exchange").await)
}
pub(super) async fn handle_users_me_vscodex_request(
state: &AppState,
request_context: &GatewayPublicRequestContext,
headers: &http::HeaderMap,
request_body: Option<&Bytes>,
) -> Response<Body> {
let auth = match resolve_authenticated_local_user(state, request_context, headers).await {
Ok(value) => value,
Err(response) => return response,
};
let config = match load_vscodex_sidecar_config() {
Ok(Some(value)) => value,
Ok(None) => {
return build_auth_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"VS Codex 服务未启用",
false,
)
}
Err(detail) => {
warn!(
user_id = %auth.user.id,
error = %detail,
"VS Codex sidecar configuration is invalid"
);
return build_auth_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"VS Codex 服务配置不完整",
false,
);
}
};
let Some(route_kind) = request_context
.control_decision
.as_ref()
.and_then(|decision| decision.route_kind.as_deref())
else {
return build_auth_error_response(
http::StatusCode::NOT_FOUND,
"VS Codex 接口不存在",
false,
);
};
let request = match build_vscodex_sidecar_request(
&config,
&auth.user.id,
route_kind,
&request_context.request_path,
request_body,
) {
Ok(value) => value,
Err(response) => return response,
};
send_vscodex_sidecar_request(request, &auth.user.id, route_kind).await
}
fn load_vscodex_sidecar_config() -> Result<Option<VscodexSidecarConfig>, String> {
if !module_available_from_env(VSCODEX_ENABLED_ENV, false) {
return Ok(None);
}
let raw_url = required_env(VSCODEX_INTERNAL_URL_ENV)?;
let base_url = reqwest::Url::parse(&raw_url)
.map_err(|err| format!("{VSCODEX_INTERNAL_URL_ENV} is invalid: {err}"))?;
if !matches!(base_url.scheme(), "http" | "https")
|| !base_url.has_host()
|| !base_url.username().is_empty()
|| base_url.password().is_some()
|| base_url.query().is_some()
|| base_url.fragment().is_some()
|| base_url.cannot_be_a_base()
{
return Err(format!(
"{VSCODEX_INTERNAL_URL_ENV} must be an HTTP(S) base URL without credentials, query, or fragment"
));
}
let token = required_env(VSCODEX_INTERNAL_TOKEN_ENV)?;
let authorization = reqwest::header::HeaderValue::from_str(&format!("Bearer {token}"))
.map_err(|_| format!("{VSCODEX_INTERNAL_TOKEN_ENV} is not a valid HTTP credential"))?;
let http_client = VSCODEX_HTTP_CLIENT
.as_ref()
.map_err(|err| format!("could not initialize VS Codex HTTP client: {err}"))?
.clone();
Ok(Some(VscodexSidecarConfig {
base_url,
authorization,
http_client,
}))
}
fn required_env(key: &str) -> Result<String, String> {
std::env::var(key)
.ok()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.ok_or_else(|| format!("{key} is required"))
}
fn build_vscodex_sidecar_request(
config: &VscodexSidecarConfig,
user_id: &str,
route_kind: &str,
request_path: &str,
request_body: Option<&Bytes>,
) -> Result<reqwest::RequestBuilder, Response<Body>> {
let (method, suffix, payload) = match route_kind {
"vscodex_devices_list" => (reqwest::Method::GET, vec!["devices"], None),
"vscodex_pairing_create" => (
reqwest::Method::POST,
vec!["pairings"],
Some(parse_pairing_request(request_body)?),
),
"vscodex_device_delete" => {
let Some(device_id) = vscodex_device_id_from_path(request_path) else {
return Err(build_auth_error_response(
http::StatusCode::BAD_REQUEST,
"设备标识无效",
false,
));
};
(reqwest::Method::DELETE, vec!["devices", device_id], None)
}
"vscodex_ws_ticket_create" => (
reqwest::Method::POST,
vec!["ws-tickets"],
Some(parse_ws_ticket_request(request_body)?),
),
_ => {
return Err(build_auth_error_response(
http::StatusCode::NOT_FOUND,
"VS Codex 接口不存在",
false,
))
}
};
let url = build_vscodex_sidecar_url(&config.base_url, user_id, &suffix).map_err(|detail| {
warn!(user_id = %user_id, error = %detail, "could not build VS Codex sidecar URL");
build_auth_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"VS Codex 服务配置不完整",
false,
)
})?;
Ok(build_authenticated_sidecar_request(
config, method, url, payload,
))
}
fn build_authenticated_sidecar_request(
config: &VscodexSidecarConfig,
method: reqwest::Method,
url: reqwest::Url,
payload: Option<Value>,
) -> reqwest::RequestBuilder {
let mut request = config
.http_client
.request(method, url)
.header(header::AUTHORIZATION, config.authorization.clone())
.header(header::ACCEPT, "application/json")
.timeout(VSCODEX_REQUEST_TIMEOUT);
if let Some(payload) = payload {
request = request.json(&payload);
}
request
}
fn build_vscodex_sidecar_url(
base_url: &reqwest::Url,
user_id: &str,
suffix: &[&str],
) -> Result<reqwest::Url, String> {
let mut segments = vec!["internal", "v1", "users", user_id];
segments.extend(suffix.iter().copied());
append_vscodex_sidecar_path(base_url, &segments)
}
fn append_vscodex_sidecar_path(
base_url: &reqwest::Url,
suffix: &[&str],
) -> Result<reqwest::Url, String> {
let mut url = base_url.clone();
let mut path_segments = url
.path_segments_mut()
.map_err(|_| "VS Codex sidecar URL cannot contain path segments".to_string())?;
path_segments.pop_if_empty();
path_segments.extend(suffix.iter().copied());
drop(path_segments);
Ok(url)
}
fn build_vscodex_websocket_url(base_url: &reqwest::Url) -> Result<reqwest::Url, String> {
let mut url = append_vscodex_sidecar_path(base_url, &["api", "vscodex", "ws"])?;
let scheme = match url.scheme() {
"http" => "ws",
"https" => "wss",
_ => return Err("VS Codex sidecar URL must use HTTP(S)".to_string()),
};
url.set_scheme(scheme)
.map_err(|_| "could not convert VS Codex sidecar URL to WebSocket".to_string())?;
Ok(url)
}
fn parse_pairing_request(request_body: Option<&Bytes>) -> Result<Value, Response<Body>> {
let payload = parse_json_request::<CreatePairingRequest>(request_body, true)?;
let mut object = Map::new();
if let Some(name) = payload.name {
object.insert("name".to_string(), Value::String(name));
}
Ok(Value::Object(object))
}
fn parse_ws_ticket_request(request_body: Option<&Bytes>) -> Result<Value, Response<Body>> {
let payload = parse_json_request::<CreateWsTicketRequest>(request_body, false)?;
let device_id = payload.device_id.trim();
if !valid_vscodex_device_id(device_id) {
return Err(build_auth_error_response(
http::StatusCode::BAD_REQUEST,
"设备标识无效",
false,
));
}
Ok(json!({ "device_id": device_id }))
}
fn parse_pairing_exchange_request(request_body: Option<&Bytes>) -> Result<Value, Response<Body>> {
let payload = parse_json_request::<ExchangePairingRequest>(request_body, false)?;
let code = payload.code.trim();
if code.is_empty() || code.len() > 256 || code.chars().any(char::is_control) {
return Err(build_auth_error_response(
http::StatusCode::BAD_REQUEST,
"配对码无效",
false,
));
}
let mut object = Map::from_iter([("code".to_string(), Value::String(code.to_string()))]);
if let Some(name) = payload.name {
object.insert("name".to_string(), Value::String(name));
}
Ok(Value::Object(object))
}
fn parse_json_request<T>(
request_body: Option<&Bytes>,
empty_object_allowed: bool,
) -> Result<T, Response<Body>>
where
T: serde::de::DeserializeOwned,
{
let body = request_body.filter(|body| !body.is_empty());
let result = match body {
Some(body) => serde_json::from_slice(body),
None if empty_object_allowed => serde_json::from_slice(b"{}"),
None => {
return Err(build_auth_error_response(
http::StatusCode::BAD_REQUEST,
"缺少请求体",
false,
))
}
};
result.map_err(|_| {
build_auth_error_response(http::StatusCode::BAD_REQUEST, "请求数据验证失败", false)
})
}
fn vscodex_device_id_from_path(path: &str) -> Option<&str> {
let trimmed = path.trim_end_matches('/');
let device_id = trimmed.strip_prefix(VSCODEX_DEVICE_PATH_PREFIX)?;
if device_id.contains('/') || !valid_vscodex_device_id(device_id) {
return None;
}
Some(device_id)
}
fn valid_vscodex_device_id(value: &str) -> bool {
!value.is_empty() && value.len() <= 128 && !value.chars().any(char::is_control)
}
async fn send_vscodex_sidecar_request(
request: reqwest::RequestBuilder,
request_scope: &str,
operation: &str,
) -> Response<Body> {
let mut upstream = match request.send().await {
Ok(value) => value,
Err(err) => {
warn!(
request_scope = %request_scope,
operation = %operation,
error = %err,
"VS Codex sidecar request failed"
);
let (status, detail) = if err.is_timeout() {
(http::StatusCode::GATEWAY_TIMEOUT, "VS Codex 服务请求超时")
} else {
(http::StatusCode::BAD_GATEWAY, "VS Codex 服务暂时不可用")
};
return build_auth_error_response(status, detail, false);
}
};
let status = http::StatusCode::from_u16(upstream.status().as_u16())
.unwrap_or(http::StatusCode::BAD_GATEWAY);
if matches!(
status,
http::StatusCode::UNAUTHORIZED | http::StatusCode::FORBIDDEN
) {
warn!(
request_scope = %request_scope,
operation = %operation,
upstream_status = status.as_u16(),
"VS Codex sidecar rejected gateway credentials"
);
return build_auth_error_response(
http::StatusCode::BAD_GATEWAY,
"VS Codex 服务鉴权失败",
false,
);
}
if status.is_redirection() {
warn!(
request_scope = %request_scope,
operation = %operation,
upstream_status = status.as_u16(),
"VS Codex sidecar returned an unexpected redirect"
);
return build_auth_error_response(
http::StatusCode::BAD_GATEWAY,
"VS Codex 服务返回无效响应",
false,
);
}
if status == http::StatusCode::NO_CONTENT {
return vscodex_no_store_response(status.into_response(), None);
}
let mut response_body = Vec::new();
while let Some(chunk) = match upstream.chunk().await {
Ok(value) => value,
Err(err) => {
warn!(
request_scope = %request_scope,
operation = %operation,
error = %err,
"could not read VS Codex sidecar response"
);
return build_auth_error_response(
http::StatusCode::BAD_GATEWAY,
"VS Codex 服务返回无效响应",
false,
);
}
} {
if response_body.len().saturating_add(chunk.len()) > VSCODEX_MAX_RESPONSE_BYTES {
warn!(
request_scope = %request_scope,
operation = %operation,
"VS Codex sidecar response exceeded the size limit"
);
return build_auth_error_response(
http::StatusCode::BAD_GATEWAY,
"VS Codex 服务返回无效响应",
false,
);
}
response_body.extend_from_slice(&chunk);
}
if response_body.is_empty() {
return build_auth_error_response(
http::StatusCode::BAD_GATEWAY,
"VS Codex 服务返回无效响应",
false,
);
}
let payload = match serde_json::from_slice(&response_body) {
Ok(value) => value,
Err(err) => {
warn!(
request_scope = %request_scope,
operation = %operation,
upstream_status = status.as_u16(),
error = %err,
"VS Codex sidecar returned non-JSON data"
);
return build_auth_error_response(
http::StatusCode::BAD_GATEWAY,
"VS Codex 服务返回无效响应",
false,
);
}
};
let retry_after = upstream.headers().get(header::RETRY_AFTER).cloned();
vscodex_no_store_response(build_auth_json_response(status, payload, None), retry_after)
}
fn vscodex_no_store_response(
mut response: Response<Body>,
retry_after: Option<http::HeaderValue>,
) -> Response<Body> {
response.headers_mut().insert(
header::CACHE_CONTROL,
http::HeaderValue::from_static("no-store"),
);
if let Some(retry_after) = retry_after {
response
.headers_mut()
.insert(header::RETRY_AFTER, retry_after);
}
response
}
async fn bridge_vscodex_websockets<S>(
browser_socket: WebSocket,
sidecar_socket: S,
ip_connection_permit: VscodexWsIpConnectionPermit,
) where
S: futures_util::Stream<
Item = Result<TungsteniteMessage, tokio_tungstenite::tungstenite::Error>,
> + futures_util::Sink<TungsteniteMessage, Error = tokio_tungstenite::tungstenite::Error>
+ Unpin
+ Send
+ 'static,
{
let (mut browser_tx, mut browser_rx) = browser_socket.split();
let (mut sidecar_tx, mut sidecar_rx) = sidecar_socket.split();
let mut ip_connection_permit = Some(ip_connection_permit);
loop {
tokio::select! {
browser_message = browser_rx.next() => {
match browser_message {
Some(Ok(message)) => {
let close = matches!(message, AxumMessage::Close(_));
if let Err(err) = sidecar_tx.send(axum_to_tungstenite_message(message)).await {
warn!(error = %err, "could not forward VS Codex browser WebSocket frame");
break;
}
if close {
break;
}
}
Some(Err(err)) => {
warn!(error = %err, "VS Codex browser WebSocket read failed");
break;
}
None => break,
}
}
sidecar_message = sidecar_rx.next() => {
match sidecar_message {
Some(Ok(TungsteniteMessage::Frame(_))) => continue,
Some(Ok(message)) => {
if ip_connection_permit.is_some() && vscodex_ws_authentication_succeeded(&message) {
ip_connection_permit.take();
}
let close = matches!(message, TungsteniteMessage::Close(_));
if let Err(err) = browser_tx.send(tungstenite_to_axum_message(message)).await {
warn!(error = %err, "could not forward VS Codex sidecar WebSocket frame");
break;
}
if close {
break;
}
}
Some(Err(err)) => {
warn!(error = %err, "VS Codex sidecar WebSocket read failed");
break;
}
None => break,
}
}
}
}
let _ = sidecar_tx.close().await;
let _ = browser_tx.close().await;
}
fn vscodex_ws_authentication_succeeded(message: &TungsteniteMessage) -> bool {
let TungsteniteMessage::Text(text) = message else {
return false;
};
serde_json::from_str::<Value>(text.as_ref())
.ok()
.and_then(|payload| {
payload
.get("type")
.and_then(Value::as_str)
.map(str::to_string)
})
.as_deref()
== Some("auth.ok")
}
fn axum_to_tungstenite_message(message: AxumMessage) -> TungsteniteMessage {
match message {
AxumMessage::Text(text) => TungsteniteMessage::Text(text.to_string().into()),
AxumMessage::Binary(bytes) => TungsteniteMessage::Binary(bytes),
AxumMessage::Ping(bytes) => TungsteniteMessage::Ping(bytes),
AxumMessage::Pong(bytes) => TungsteniteMessage::Pong(bytes),
AxumMessage::Close(frame) => {
TungsteniteMessage::Close(frame.map(|frame| TungsteniteCloseFrame {
code: frame.code.into(),
reason: frame.reason.to_string().into(),
}))
}
}
}
fn tungstenite_to_axum_message(message: TungsteniteMessage) -> AxumMessage {
match message {
TungsteniteMessage::Text(text) => AxumMessage::Text(text.to_string().into()),
TungsteniteMessage::Binary(bytes) => AxumMessage::Binary(bytes),
TungsteniteMessage::Ping(bytes) => AxumMessage::Ping(bytes),
TungsteniteMessage::Pong(bytes) => AxumMessage::Pong(bytes),
TungsteniteMessage::Close(frame) => AxumMessage::Close(frame.map(|frame| AxumCloseFrame {
code: frame.code.into(),
reason: frame.reason.to_string().into(),
})),
TungsteniteMessage::Frame(_) => AxumMessage::Close(None),
}
}
#[cfg(test)]
mod tests {
use super::{vscodex_ws_authentication_succeeded, VscodexWsIpConnectionLimiter};
use std::sync::Arc;
use tokio_tungstenite::tungstenite::Message;
#[test]
fn vscodex_ws_ip_limiter_releases_and_removes_inactive_ips() {
let limiter = Arc::new(VscodexWsIpConnectionLimiter::new(1));
let client_ip = "198.51.100.10".parse().expect("IP should parse");
let permit = limiter
.try_acquire(client_ip)
.expect("first connection should acquire");
assert_eq!(limiter.active_ip_count(), 1);
assert!(limiter.try_acquire(client_ip).is_none());
drop(permit);
assert_eq!(limiter.active_ip_count(), 0);
assert!(limiter.try_acquire(client_ip).is_some());
}
#[test]
fn vscodex_ws_ip_limiter_releases_only_after_sidecar_auth_success() {
assert!(vscodex_ws_authentication_succeeded(&Message::Text(
r#"{"type":"auth.ok","role":"operator"}"#.into()
)));
assert!(!vscodex_ws_authentication_succeeded(&Message::Text(
r#"{"type":"auth","token":"client-controlled"}"#.into()
)));
assert!(!vscodex_ws_authentication_succeeded(&Message::Binary(
br#"{"type":"auth.ok"}"#.to_vec().into()
)));
}
}
@@ -443,7 +443,18 @@ fn provider_quota_metadata_bucket<'a>(
fn provider_quota_timestamp_unix_secs(value: Option<&Value>) -> Option<u64> {
let mut parsed = match value {
Some(Value::Number(number)) => number.as_f64(),
Some(Value::String(text)) => text.trim().parse::<f64>().ok(),
Some(Value::String(text)) => {
let text = text.trim();
if let Ok(timestamp) = text.parse::<f64>() {
Some(timestamp)
} else {
return chrono::DateTime::parse_from_rfc3339(text)
.ok()?
.timestamp()
.try_into()
.ok();
}
}
_ => None,
}?;
if !parsed.is_finite() || parsed <= 0.0 {
@@ -560,7 +571,10 @@ fn model_quota_window_snapshot(
.map(|value| value.clamp(0.0, 1.0))
.or_else(|| used_ratio.map(|value| (1.0 - value).max(0.0)));
let reset_at = provider_quota_timestamp_unix_secs(
item.get("reset_at").or_else(|| item.get("next_reset_at")),
item.get("reset_at")
.or_else(|| item.get("next_reset_at"))
.or_else(|| item.get("reset_time"))
.or_else(|| item.get("next_reset_time")),
);
let reset_seconds = quota_window_reset_seconds(observed_at_unix_secs, reset_at);
let is_exhausted = item
@@ -4094,7 +4108,8 @@ mod tests {
},
"claude-sonnet-4-6": {
"display_name": "Claude Sonnet 4.6 (Thinking)",
"remaining_fraction": 0.3
"remaining_fraction": 0.3,
"reset_time": "2026-04-07T12:34:56Z"
}
}
}
@@ -4139,6 +4154,16 @@ mod tests {
label_for_model("claude-sonnet-4-6"),
Some(json!("Claude Sonnet 4.6 (Thinking)"))
);
let claude_window = windows
.iter()
.filter_map(Value::as_object)
.find(|window| window.get("model") == Some(&json!("claude-sonnet-4-6")))
.expect("Claude quota window should exist");
assert_eq!(
claude_window.get("reset_at"),
Some(&json!(1_775_565_296u64))
);
assert_eq!(claude_window.get("reset_seconds"), Some(&json!(12_011u64)));
}
#[test]
@@ -495,8 +495,14 @@ pub(crate) fn public_support_local_requires_buffered_body(
Some(
"api_keys_create"
| "api_key_install_session_create"
| "management_tokens_create",
| "management_tokens_create"
| "vscodex_pairing_create"
| "vscodex_ws_ticket_create",
),
) | (
Some("vscodex"),
http::Method::POST,
Some("pairing_exchange"),
) | (
Some("wallet"),
http::Method::POST,
@@ -256,12 +256,16 @@ fn auth_config_has_refresh_token(auth_config: Option<&str>) -> bool {
let Ok(value) = serde_json::from_str::<Value>(auth_config) else {
return false;
};
value
.as_object()
.and_then(|object| object.get("refresh_token"))
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty())
let Some(object) = value.as_object() else {
return false;
};
["refresh_token", "refreshToken"].iter().any(|field| {
object
.get(*field)
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty())
})
}
fn now_unix_secs() -> u64 {
@@ -273,7 +277,44 @@ fn now_unix_secs() -> u64 {
#[cfg(test)]
mod tests {
use super::agent_identity_needs_task_recovery;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use super::{
agent_identity_needs_task_recovery, auth_config_has_refresh_token, oauth_refresh_candidate,
};
#[test]
fn legacy_antigravity_refresh_token_is_refreshable() {
assert!(auth_config_has_refresh_token(Some(
r#"{"refreshToken":"legacy-refresh-token"}"#,
)));
}
#[test]
fn expiring_antigravity_oauth_key_is_refresh_candidate() {
let provider = StoredProviderCatalogProvider::new(
"provider-antigravity".to_string(),
"Antigravity".to_string(),
None,
"antigravity".to_string(),
)
.expect("provider should build");
let mut key = StoredProviderCatalogKey::new(
"key-antigravity".to_string(),
provider.id.clone(),
"Antigravity OAuth".to_string(),
"oauth".to_string(),
None,
true,
)
.expect("key should build");
key.encrypted_auth_config = Some("encrypted-auth-config".to_string());
key.expires_at_unix_secs = Some(120);
assert!(oauth_refresh_candidate(&provider, &key, 120));
}
#[test]
fn pending_agent_identity_without_task_is_recoverable() {
@@ -167,17 +167,43 @@ fn codex_quota_breaker_ttl(quota_metadata: &Value, now_unix_secs: u64) -> (u64,
pub(crate) fn codex_quota_exhaustion_reset_at(
quota_metadata: &Value,
now_unix_secs: u64,
) -> Option<u64> {
codex_quota_exhaustion_reset_at_for_prefixes(
quota_metadata,
now_unix_secs,
&["primary", "secondary"],
)
}
/// Returns the reset deadline for an exhausted model-scoped Codex window without promoting that
/// window to account-wide exhaustion.
pub(crate) fn codex_model_quota_exhaustion_reset_at(
quota_metadata: &Value,
now_unix_secs: u64,
) -> Option<u64> {
codex_quota_exhaustion_reset_at_for_prefixes(
quota_metadata,
now_unix_secs,
&["spark_primary", "spark_secondary"],
)
}
fn codex_quota_exhaustion_reset_at_for_prefixes(
quota_metadata: &Value,
now_unix_secs: u64,
window_prefixes: &[&str],
) -> Option<u64> {
let Some(metadata) = quota_metadata.as_object() else {
return None;
};
let exhausted_windows = ["primary", "secondary"]
.into_iter()
let exhausted_windows = window_prefixes
.iter()
.copied()
.filter(|prefix| codex_window_is_exhausted(metadata, prefix))
.collect::<Vec<_>>();
let prefixes = if exhausted_windows.is_empty() {
vec!["primary", "secondary"]
window_prefixes.to_vec()
} else {
exhausted_windows
};
+4 -3
View File
@@ -34,9 +34,10 @@ pub(crate) use self::classifier::{
LocalTransportFailoverClassification,
};
pub(crate) use self::codex_quota_breaker::{
codex_account_id_from_headers, codex_quota_breaker_blocks_candidate,
codex_quota_exhaustion_reset_at, install_codex_quota_exhaustion_breaker,
log_codex_quota_breaker_check_failure, log_codex_quota_breaker_install_failure,
codex_account_id_from_headers, codex_model_quota_exhaustion_reset_at,
codex_quota_breaker_blocks_candidate, codex_quota_exhaustion_reset_at,
install_codex_quota_exhaustion_breaker, log_codex_quota_breaker_check_failure,
log_codex_quota_breaker_install_failure,
};
pub(crate) use self::effects::{
apply_local_execution_effect, apply_local_stream_failure_effects,
+14 -5
View File
@@ -90,11 +90,15 @@ pub(crate) fn provider_key_can_refresh_oauth(
) -> bool {
auth_semantics.can_refresh_oauth()
&& (provider_key_auth_config_is_agent_identity(provider_type, auth_config)
|| auth_config
.and_then(|config| config.get("refresh_token"))
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty()))
|| auth_config.is_some_and(|config| {
["refresh_token", "refreshToken"].iter().any(|field| {
config
.get(*field)
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty())
})
}))
}
pub(crate) fn provider_key_can_export_oauth(
@@ -425,6 +429,11 @@ mod tests {
"codex",
json!({ "refresh_token": "refresh-token" }).as_object()
));
assert!(provider_key_can_refresh_oauth(
provider_key_auth_semantics(&sample_key("oauth"), "antigravity"),
"antigravity",
json!({ "refreshToken": "legacy-refresh-token" }).as_object()
));
assert!(provider_key_can_refresh_oauth(
semantics,
"codex",
+176 -11
View File
@@ -4,14 +4,20 @@ use std::sync::{
};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use aether_runtime_state::RuntimeState;
use aether_contracts::ExecutionPlan;
use aether_runtime_state::{
RuntimeSemaphoreConfig, RuntimeSemaphoreError, RuntimeSemaphorePermit, RuntimeState,
};
use dashmap::DashMap;
use serde::{Deserialize, Serialize};
use tokio::task::JoinHandle;
use tracing::debug;
use uuid::Uuid;
use crate::{AppState, GatewayError};
const PROVIDER_POOL_IN_FLIGHT_TOKENS_PREFIX: &str = "ap:provider_pool:in_flight";
const PROVIDER_KEY_CONCURRENCY_GATE: &str = "provider_key";
const PROVIDER_POOL_DEMAND_SNAPSHOT_PREFIX: &str = "ap:provider_pool:demand";
const PROVIDER_POOL_BURST_PENDING_PREFIX: &str = "ap:quota_probe:burst_pending";
const PROVIDER_POOL_IN_FLIGHT_TOKEN_TTL_MS: u64 = 120_000;
@@ -42,10 +48,17 @@ pub(crate) struct ProviderPoolDemandSnapshot {
pub(crate) struct ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind,
provider_key_permit: Option<RuntimeSemaphorePermit>,
released: bool,
}
pub(crate) enum ProviderPoolInFlightAdmission {
Acquired(Option<ProviderPoolInFlightGuard>),
Saturated { limit: usize },
}
enum ProviderPoolInFlightGuardKind {
Disabled,
Local {
provider_id: String,
counter: Arc<AtomicUsize>,
@@ -69,6 +82,9 @@ impl ProviderPoolInFlightGuard {
return;
}
match &mut self.kind {
ProviderPoolInFlightGuardKind::Disabled => {
self.released = true;
}
ProviderPoolInFlightGuardKind::Local {
provider_id,
counter,
@@ -101,6 +117,16 @@ impl ProviderPoolInFlightGuard {
}
}
}
if self.released {
if let Some(provider_key_permit) = self.provider_key_permit.take() {
if let Err(err) = provider_key_permit.release().await {
debug!(
error = ?err,
"gateway provider pool demand: failed to release provider key permit; scheduling drop fallback"
);
}
}
}
}
}
@@ -111,6 +137,7 @@ impl Drop for ProviderPoolInFlightGuard {
}
self.released = true;
match &mut self.kind {
ProviderPoolInFlightGuardKind::Disabled => {}
ProviderPoolInFlightGuardKind::Local {
provider_id,
counter,
@@ -290,32 +317,83 @@ pub(crate) async fn acquire_provider_pool_in_flight_guard(
candidate_id: Option<&str>,
key_id: &str,
) -> Option<ProviderPoolInFlightGuard> {
acquire_provider_pool_in_flight_guard_with_key_limit(
runtime,
provider_id,
request_id,
candidate_id,
key_id,
None,
)
.await
.ok()
.flatten()
}
pub(crate) async fn acquire_provider_pool_in_flight_guard_with_key_limit(
runtime: Arc<RuntimeState>,
provider_id: &str,
request_id: &str,
candidate_id: Option<&str>,
key_id: &str,
concurrent_limit: Option<usize>,
) -> Result<Option<ProviderPoolInFlightGuard>, RuntimeSemaphoreError> {
let provider_key_permit = match concurrent_limit.filter(|limit| *limit > 0) {
Some(limit) => Some(
runtime
.keyed_semaphore(
PROVIDER_KEY_CONCURRENCY_GATE,
key_id,
limit,
RuntimeSemaphoreConfig::default(),
)?
.try_acquire()
.await?,
),
None => None,
};
let provider_id = provider_id.trim();
if provider_id.is_empty() {
return None;
return Ok(
provider_key_permit.map(|provider_key_permit| ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind::Disabled,
provider_key_permit: Some(provider_key_permit),
released: false,
}),
);
}
match provider_pool_in_flight_mode() {
ProviderPoolInFlightMode::Off => return None,
ProviderPoolInFlightMode::Off => {
return Ok(
provider_key_permit.map(|provider_key_permit| ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind::Disabled,
provider_key_permit: Some(provider_key_permit),
released: false,
}),
);
}
ProviderPoolInFlightMode::Local => {
let counter = increment_local_provider_in_flight(provider_id);
return Some(ProviderPoolInFlightGuard {
return Ok(Some(ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind::Local {
provider_id: provider_id.to_string(),
counter,
},
provider_key_permit,
released: false,
});
}));
}
ProviderPoolInFlightMode::Runtime if runtime.is_memory() => {
let counter = increment_local_provider_in_flight(provider_id);
return Some(ProviderPoolInFlightGuard {
return Ok(Some(ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind::Local {
provider_id: provider_id.to_string(),
counter,
},
provider_key_permit,
released: false,
});
}));
}
ProviderPoolInFlightMode::Runtime => {}
}
@@ -335,7 +413,13 @@ pub(crate) async fn acquire_provider_pool_in_flight_guard(
error = ?err,
"gateway provider pool demand: failed to acquire in-flight token"
);
return None;
return Ok(
provider_key_permit.map(|provider_key_permit| ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind::Disabled,
provider_key_permit: Some(provider_key_permit),
released: false,
}),
);
}
Err(_) => {
debug!(
@@ -343,7 +427,13 @@ pub(crate) async fn acquire_provider_pool_in_flight_guard(
timeout_ms = provider_pool_in_flight_acquire_timeout().as_millis() as u64,
"gateway provider pool demand: skipped in-flight token after acquire timeout"
);
return None;
return Ok(
provider_key_permit.map(|provider_key_permit| ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind::Disabled,
provider_key_permit: Some(provider_key_permit),
released: false,
}),
);
}
}
@@ -355,7 +445,7 @@ pub(crate) async fn acquire_provider_pool_in_flight_guard(
stop_renewal.clone(),
);
Some(ProviderPoolInFlightGuard {
Ok(Some(ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind::Runtime {
runtime,
tokens_key,
@@ -363,8 +453,39 @@ pub(crate) async fn acquire_provider_pool_in_flight_guard(
stop_renewal,
renew_handle: Some(renew_handle),
},
provider_key_permit,
released: false,
})
}))
}
pub(crate) async fn acquire_provider_pool_execution_guard(
state: &AppState,
plan: &ExecutionPlan,
) -> Result<ProviderPoolInFlightAdmission, GatewayError> {
let concurrent_limit = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
.await?
.into_iter()
.find(|key| key.id == plan.key_id)
.and_then(|key| key.concurrent_limit)
.filter(|limit| *limit > 0)
.and_then(|limit| usize::try_from(limit).ok());
match acquire_provider_pool_in_flight_guard_with_key_limit(
state.runtime_state.clone(),
&plan.provider_id,
&plan.request_id,
plan.candidate_id.as_deref(),
&plan.key_id,
concurrent_limit,
)
.await
{
Ok(guard) => Ok(ProviderPoolInFlightAdmission::Acquired(guard)),
Err(RuntimeSemaphoreError::Saturated { limit, .. }) => {
Ok(ProviderPoolInFlightAdmission::Saturated { limit })
}
Err(error) => Err(GatewayError::Internal(error.to_string())),
}
}
pub(crate) async fn provider_pool_live_in_flight_count(
@@ -586,6 +707,50 @@ mod tests {
);
}
#[tokio::test]
async fn provider_key_limit_rejects_concurrent_guard_until_release() {
let runtime = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let first = acquire_provider_pool_in_flight_guard_with_key_limit(
runtime.clone(),
"provider-limit",
"request-1",
Some("candidate-1"),
"key-limit",
Some(1),
)
.await
.expect("first admission should resolve")
.expect("first guard should be acquired");
let second = acquire_provider_pool_in_flight_guard_with_key_limit(
runtime.clone(),
"provider-limit",
"request-2",
Some("candidate-2"),
"key-limit",
Some(1),
)
.await;
assert!(matches!(
second,
Err(RuntimeSemaphoreError::Saturated { limit: 1, .. })
));
first.release().await;
let replacement = acquire_provider_pool_in_flight_guard_with_key_limit(
runtime,
"provider-limit",
"request-3",
Some("candidate-3"),
"key-limit",
Some(1),
)
.await
.expect("replacement admission should resolve")
.expect("replacement guard should acquire after release");
drop(replacement);
}
#[tokio::test]
async fn demand_snapshot_uses_instant_in_flight_for_fast_rise_and_ema_for_fall() {
let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default());
@@ -363,7 +363,7 @@ pub(crate) async fn record_local_request_candidate_status(
report_context: Option<&Value>,
status_update: SchedulerRequestCandidateStatusUpdate,
) {
let Some(record) =
let Some(mut record) =
build_local_request_candidate_status_record(LocalRequestCandidateStatusRecordInput {
plan,
report_context,
@@ -372,6 +372,8 @@ pub(crate) async fn record_local_request_candidate_status(
else {
return;
};
record.skip_reason =
local_request_candidate_skip_reason(record.status, record.error_type.as_deref());
persist_local_request_candidate_status_record(state, record).await;
}
@@ -429,6 +431,7 @@ fn build_local_request_candidate_status_snapshot_record(
started_at_unix_ms,
finished_at_unix_ms,
} = status_update;
let skip_reason = local_request_candidate_skip_reason(status, error_type.as_deref());
UpsertRequestCandidateRecord {
id: snapshot.candidate_id.clone(),
request_id: snapshot.request_id.clone(),
@@ -442,7 +445,7 @@ fn build_local_request_candidate_status_snapshot_record(
endpoint_id: Some(snapshot.endpoint_id.clone()),
key_id: Some(snapshot.key_id.clone()),
status,
skip_reason: None,
skip_reason,
is_cached: None,
status_code,
error_type,
@@ -457,6 +460,17 @@ fn build_local_request_candidate_status_snapshot_record(
}
}
fn local_request_candidate_skip_reason(
status: RequestCandidateStatus,
error_type: Option<&str>,
) -> Option<String> {
(status == RequestCandidateStatus::Skipped)
.then_some(error_type)
.flatten()
.filter(|reason| *reason == "provider_key_concurrency_limit_reached")
.map(ToOwned::to_owned)
}
pub(crate) fn try_enqueue_local_request_candidate_status_snapshot(
state: &(impl RequestCandidateRuntimeWriter + ?Sized),
snapshot: &LocalRequestCandidateStatusSnapshot,
@@ -1078,6 +1092,40 @@ mod tests {
assert_eq!(records[0].status_code, Some(200));
}
#[test]
fn saturated_provider_key_snapshot_persists_capacity_skip_reason() {
let mut plan = sample_plan();
plan.candidate_id = Some("candidate-provider-key-saturated".to_string());
let snapshot = snapshot_local_request_candidate_status(&plan, None)
.expect("candidate snapshot should build");
let writer = SynchronousStatusWriter::default();
try_enqueue_local_request_candidate_status_snapshot(
&writer,
&snapshot,
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Skipped,
status_code: Some(429),
error_type: Some("provider_key_concurrency_limit_reached".to_string()),
error_message: Some("provider key concurrency limit reached: 1".to_string()),
latency_ms: Some(0),
started_at_unix_ms: Some(123),
finished_at_unix_ms: Some(123),
},
)
.expect("saturated status should use the synchronous enqueue path");
let records = writer
.records
.lock()
.expect("synchronous status records lock");
assert_eq!(records.len(), 1);
assert_eq!(
records[0].skip_reason.as_deref(),
Some("provider_key_concurrency_limit_reached")
);
}
fn sample_minimal_candidate() -> SchedulerMinimalCandidateSelectionCandidate {
SchedulerMinimalCandidateSelectionCandidate {
provider_id: "provider-1".to_string(),
+89 -4
View File
@@ -33,7 +33,9 @@ use std::sync::atomic::Ordering;
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use aether_crypto::encrypt_python_fernet_plaintext;
use aether_crypto::{
decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY,
};
const LOCAL_OAUTH_HTTP_TIMEOUT_MS: u64 = 30_000;
const REMOTE_OAUTH_REFRESH_WAIT_TIMEOUT: Duration = Duration::from_secs(35);
@@ -3395,7 +3397,10 @@ mod tests {
use std::sync::Arc;
use std::time::{Duration, Instant};
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
use aether_crypto::{
decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext,
DEVELOPMENT_ENCRYPTION_KEY,
};
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyListQuery,
@@ -3478,9 +3483,17 @@ mod tests {
fn codex_oauth_state(
auth_config: &serde_json::Value,
access_token: &str,
) -> (AppState, Arc<InMemoryProviderCatalogReadRepository>, String) {
provider_oauth_state("codex", auth_config, access_token)
}
fn provider_oauth_state(
provider_type: &str,
auth_config: &serde_json::Value,
access_token: &str,
) -> (AppState, Arc<InMemoryProviderCatalogReadRepository>, String) {
let mut provider = sample_provider();
provider.provider_type = "codex".to_string();
provider.provider_type = provider_type.to_string();
let encrypted_auth_config =
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, &auth_config.to_string())
.expect("auth config should encrypt");
@@ -3490,7 +3503,7 @@ mod tests {
let key = StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"Codex OAuth".to_string(),
format!("{provider_type} OAuth"),
"oauth".to_string(),
None,
true,
@@ -4564,6 +4577,78 @@ mod tests {
);
}
#[tokio::test]
async fn antigravity_refresh_entry_persists_tokens_and_expiry() {
let initial_config = json!({
"provider_type": "antigravity",
"refreshToken": "legacy-refresh-token",
"expires_at": 1,
});
let (state, repository, _) =
provider_oauth_state("antigravity", &initial_config, "stale-access-token");
let transport = state
.read_provider_transport_snapshot("provider-1", "endpoint-1", "key-1")
.await
.expect("transport should load")
.expect("transport should exist");
let expected_credential_fence = state
.capture_provider_transport_credential_fence(&transport)
.await
.expect("credential fence should load")
.expect("credential fence should match");
let expires_at = 4_102_555_900;
let refreshed_entry = crate::provider_transport::CachedOAuthEntry {
provider_type: "antigravity".to_string(),
auth_header_name: "authorization".to_string(),
auth_header_value: "Bearer fresh-access-token".to_string(),
expires_at_unix_secs: Some(expires_at),
metadata: Some(json!({
"provider_type": "antigravity",
"refresh_token": "legacy-refresh-token",
"expires_at": expires_at,
})),
source_fingerprint: None,
};
state
.persist_local_oauth_refresh_entry(
&transport,
&refreshed_entry,
Some(&expected_credential_fence),
)
.await
.expect("Antigravity refresh should persist");
let stored = repository
.list_keys_by_ids(&["key-1".to_string()])
.await
.expect("key should reload")
.pop()
.expect("key should remain");
let access_token = decrypt_python_fernet_ciphertext(
DEVELOPMENT_ENCRYPTION_KEY,
stored
.encrypted_api_key
.as_deref()
.expect("access token should persist"),
)
.expect("access token should decrypt");
let auth_config = decrypt_python_fernet_ciphertext(
DEVELOPMENT_ENCRYPTION_KEY,
stored
.encrypted_auth_config
.as_deref()
.expect("auth config should persist"),
)
.expect("auth config should decrypt");
let auth_config: serde_json::Value =
serde_json::from_str(&auth_config).expect("auth config should parse");
assert_eq!(access_token, "fresh-access-token");
assert_eq!(auth_config["refresh_token"], json!("legacy-refresh-token"));
assert_eq!(stored.expires_at_unix_secs, Some(expires_at));
}
#[tokio::test]
async fn agent_auth_config_fence_rejects_metadata_only_rewrite() {
let initial_config = json!({
@@ -499,11 +499,16 @@ impl AppState {
plan_id: &str,
input: &BillingPlanWriteInput,
) -> Result<LocalMutationOutcome<BillingPlanRecord>, GatewayError> {
self.data
let outcome = self
.data
.update_billing_plan(plan_id, input)
.await
.map(local_mutation_outcome)
.map_err(data_error)
.map_err(data_error)?;
if matches!(&outcome, LocalMutationOutcome::Applied(_)) {
self.invalidate_auth_context_cache();
}
Ok(outcome)
}
pub(crate) async fn set_billing_plan_enabled(
@@ -539,6 +544,23 @@ impl AppState {
.map_err(data_error)
}
pub(crate) async fn revoke_user_plan_entitlement(
&self,
user_id: &str,
entitlement_id: &str,
) -> Result<LocalMutationOutcome<()>, GatewayError> {
let outcome = self
.data
.revoke_user_plan_entitlement(user_id, entitlement_id)
.await
.map(local_mutation_outcome)
.map_err(data_error)?;
if matches!(&outcome, LocalMutationOutcome::Applied(_)) {
self.invalidate_auth_context_cache();
}
Ok(outcome)
}
pub(crate) async fn find_user_daily_quota_availability(
&self,
user_id: &str,
@@ -629,6 +629,7 @@ async fn gateway_pool_list_includes_usage_totals_and_nullable_lru_score() {
"sk-usage",
);
key.name = "usage key".to_string();
key.concurrent_limit = Some(5);
key.request_count = Some(1566);
key.total_tokens = 187_327_321;
key.total_cost_usd = 93.1319297;
@@ -661,6 +662,7 @@ async fn gateway_pool_list_includes_usage_totals_and_nullable_lru_score() {
.expect("json body should parse");
let keys = payload["keys"].as_array().expect("keys should be array");
assert_eq!(keys.len(), 1);
assert_eq!(keys[0]["concurrent_limit"], json!(5));
assert_eq!(keys[0]["request_count"], json!(1566));
assert_eq!(keys[0]["total_tokens"], json!(187_327_321u64));
assert_eq!(keys[0]["total_cost_usd"], json!("93.13192970"));
@@ -3643,6 +3645,7 @@ async fn gateway_batch_updates_shared_pool_key_configuration() {
"api_formats": ["openai:responses"],
"internal_priority": 7,
"rpm_limit": null,
"concurrent_limit": 6,
"auto_fetch_models": false,
"allowed_models": ["gpt-5.6-sol", "gpt-5.6-luna"],
"locked_models": [],
@@ -3668,6 +3671,7 @@ async fn gateway_batch_updates_shared_pool_key_configuration() {
assert_eq!(key.allow_auth_channel_mismatch_formats, Some(json!([])));
assert_eq!(key.internal_priority, 7);
assert_eq!(key.rpm_limit, None);
assert_eq!(key.concurrent_limit, Some(6));
assert_eq!(key.learned_rpm_limit, None);
assert!(!key.auto_fetch_models);
assert_eq!(
@@ -4670,8 +4670,7 @@ async fn gateway_hydrates_antigravity_project_id_from_load_code_assist_for_test_
.lock()
.expect("mutex should lock")
.push(plan.url.clone());
if plan.url == "https://daily-cloudcode-pa.googleapis.com/v1internal:loadCodeAssist"
{
if plan.url == "https://cloudcode-pa.googleapis.com/v1internal:loadCodeAssist" {
assert_eq!(plan.model_name.as_deref(), Some("loadCodeAssist"));
assert_eq!(
plan.headers.get("authorization").map(String::as_str),
@@ -4827,7 +4826,7 @@ async fn gateway_hydrates_antigravity_project_id_from_load_code_assist_for_test_
assert_eq!(
*seen_urls.lock().expect("mutex should lock"),
vec![
"https://daily-cloudcode-pa.googleapis.com/v1internal:loadCodeAssist".to_string(),
"https://cloudcode-pa.googleapis.com/v1internal:loadCodeAssist".to_string(),
"https://daily-cloudcode-pa.googleapis.com/v1internal:generateContent".to_string(),
]
);
@@ -49,6 +49,8 @@ use chrono::{TimeZone, Utc};
#[path = "public_support/dashboard.rs"]
mod dashboard;
#[path = "public_support/vscodex.rs"]
mod vscodex;
#[tokio::test]
async fn gateway_handles_public_announcements_list_without_proxying_upstream() {
@@ -0,0 +1,578 @@
use super::{
any, build_router_with_state, build_test_auth_token, json, sample_auth_session,
sample_auth_user, sample_auth_wallet, set_test_env_var, start_auth_gateway_with_state,
start_server, AppState, Arc, Json, Mutex, Request, Router, StatusCode, Utc,
};
use axum::extract::ws::{Message as AxumWsMessage, WebSocketUpgrade};
use axum::response::IntoResponse;
use futures_util::{SinkExt, StreamExt};
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::Message as TungsteniteMessage;
#[derive(Debug, Clone, PartialEq)]
struct CapturedSidecarRequest {
method: http::Method,
path: String,
authorization: Option<String>,
client_ip: Option<String>,
body: Option<serde_json::Value>,
}
#[test]
fn gateway_authenticates_and_proxies_vscodex_bff_routes() {
std::thread::Builder::new()
.name("vscodex-gateway-test".to_string())
.stack_size(32 * 1024 * 1024)
.spawn(|| {
tokio::runtime::Builder::new_multi_thread()
.enable_all()
.thread_stack_size(32 * 1024 * 1024)
.build()
.expect("test runtime should build")
.block_on(run_vscodex_gateway_integration());
})
.expect("test thread should spawn")
.join()
.expect("test thread should complete");
}
async fn run_vscodex_gateway_integration() {
let captured_requests = Arc::new(Mutex::new(Vec::<CapturedSidecarRequest>::new()));
let captured_requests_for_handler = Arc::clone(&captured_requests);
let captured_ws_handshake = Arc::new(Mutex::new(None::<(Option<String>, Option<String>)>));
let captured_ws_handshake_for_handler = Arc::clone(&captured_ws_handshake);
let sidecar = Router::new()
.route(
"/api/vscodex/ws",
any(move |ws: WebSocketUpgrade, headers: http::HeaderMap| {
let captured_ws_handshake = Arc::clone(&captured_ws_handshake_for_handler);
async move {
let origin = headers
.get(http::header::ORIGIN)
.and_then(|value| value.to_str().ok())
.map(ToOwned::to_owned);
let authorization = headers
.get(http::header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.map(ToOwned::to_owned);
*captured_ws_handshake
.lock()
.expect("WebSocket handshake store should lock") =
Some((origin, authorization));
ws.protocols(["vscodex.v1"])
.on_upgrade(|mut socket| async move {
while let Some(Ok(message)) = socket.next().await {
match message {
AxumWsMessage::Text(text) => {
let response = if text
== r#"{"type":"auth","token":"test-auth-ok"}"# {
r#"{"type":"auth.ok","role":"operator"}"#.to_string()
} else {
format!("echo:{text}")
};
if socket
.send(AxumWsMessage::Text(response.into()))
.await
.is_err()
{
break;
}
}
AxumWsMessage::Binary(bytes) => {
if socket.send(AxumWsMessage::Binary(bytes)).await.is_err()
{
break;
}
}
AxumWsMessage::Close(frame) => {
let _ = socket.send(AxumWsMessage::Close(frame)).await;
break;
}
AxumWsMessage::Ping(_) | AxumWsMessage::Pong(_) => {}
}
}
})
}
}),
)
.route(
"/{*path}",
any(move |request: Request| {
let captured_requests = Arc::clone(&captured_requests_for_handler);
async move {
let method = request.method().clone();
let path = request.uri().path().to_string();
let authorization = request
.headers()
.get(http::header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.map(ToOwned::to_owned);
let client_ip = request
.headers()
.get("x-aether-client-ip")
.and_then(|value| value.to_str().ok())
.map(ToOwned::to_owned);
let body = axum::body::to_bytes(request.into_body(), 1024 * 1024)
.await
.expect("sidecar request body should be readable");
let body = (!body.is_empty()).then(|| {
serde_json::from_slice(&body).expect("sidecar request body should be JSON")
});
captured_requests
.lock()
.expect("captured request store should lock")
.push(CapturedSidecarRequest {
method: method.clone(),
path: path.clone(),
authorization,
client_ip,
body,
});
let (status, payload) = match (method, path.as_str()) {
(http::Method::GET, "/internal/v1/users/user-auth-1/devices") => (
StatusCode::OK,
json!({ "devices": [{ "id": "host-1", "name": "My Mac" }] }),
),
(http::Method::POST, "/internal/v1/users/user-auth-1/pairings") => {
(StatusCode::CREATED, json!({ "code": "PAIR-123" }))
}
(http::Method::POST, "/internal/v1/users/user-auth-1/ws-tickets") => (
StatusCode::CREATED,
json!({
"ticket": "ticket-123",
"ws_url": "wss://aether.example/vscodex/ws"
}),
),
(http::Method::DELETE, "/internal/v1/users/user-auth-1/devices/host-1") => {
(StatusCode::NO_CONTENT, json!({}))
}
(
http::Method::DELETE,
"/internal/v1/users/user-auth-1/devices/missing",
) => (StatusCode::NOT_FOUND, json!({ "detail": "设备不存在" })),
(
http::Method::DELETE,
"/internal/v1/users/user-auth-1/devices/internal-denied",
) => (
StatusCode::UNAUTHORIZED,
json!({ "detail": "internal token invalid" }),
),
(
http::Method::DELETE,
"/internal/v1/users/user-auth-1/devices/redirect",
) => (StatusCode::TEMPORARY_REDIRECT, json!({ "redirect": true })),
(
http::Method::DELETE,
"/internal/v1/users/user-auth-1/devices/empty-ok",
) => return StatusCode::OK.into_response(),
(http::Method::POST, "/v1/pairings/exchange") => (
StatusCode::CREATED,
json!({ "device_id": "host-2", "device_token": "host-secret" }),
),
_ => (
StatusCode::NOT_FOUND,
json!({ "detail": "unexpected path" }),
),
};
let mut response = (status, Json(payload)).into_response();
if status == StatusCode::TEMPORARY_REDIRECT {
response.headers_mut().insert(
http::header::LOCATION,
"/redirect-must-not-be-followed".parse().unwrap(),
);
}
response
}
}),
);
let (sidecar_url, sidecar_handle) = start_server(sidecar).await;
let _enabled = set_test_env_var("AETHER_VSCODEX_ENABLED", "true");
let _internal_url = set_test_env_var("AETHER_VSCODEX_INTERNAL_URL", &sidecar_url);
let _internal_token = set_test_env_var("AETHER_VSCODEX_INTERNAL_TOKEN", "sidecar-secret");
let now = Utc::now();
let user = sample_auth_user(now);
let access_token = build_test_auth_token(
"access",
serde_json::Map::from_iter([
("user_id".to_string(), json!(user.id)),
("role".to_string(), json!(user.role)),
(
"created_at".to_string(),
json!(user.created_at.map(|value| value.to_rfc3339())),
),
("session_id".to_string(), json!("session-vscodex")),
]),
now + chrono::Duration::hours(1),
);
let (gateway_url, upstream_hits, gateway_handle, upstream_handle) =
start_auth_gateway_with_state(
user,
sample_auth_wallet("user-auth-1", now),
[sample_auth_session(
"user-auth-1",
"session-vscodex",
"browser-device-vscodex",
"refresh-vscodex",
now,
)],
)
.await;
let client = reqwest::Client::new();
let unauthenticated = client
.get(format!("{gateway_url}/api/users/me/vscodex/devices"))
.send()
.await
.expect("unauthenticated request should complete");
assert_eq!(unauthenticated.status(), StatusCode::UNAUTHORIZED);
assert!(captured_requests
.lock()
.expect("captured request store should lock")
.is_empty());
let devices = client
.get(format!("{gateway_url}/api/users/me/vscodex/devices"))
.bearer_auth(&access_token)
.header("x-client-device-id", "browser-device-vscodex")
.send()
.await
.expect("devices request should complete");
assert_eq!(devices.status(), StatusCode::OK);
let devices_payload: serde_json::Value =
devices.json().await.expect("devices body should be JSON");
assert_eq!(devices_payload["devices"][0]["id"], "host-1");
let pairing = client
.post(format!("{gateway_url}/api/users/me/vscodex/pairings"))
.bearer_auth(&access_token)
.header("x-client-device-id", "browser-device-vscodex")
.json(&json!({ "name": "My Mac", "user_id": "attacker" }))
.send()
.await
.expect("pairing request should complete");
assert_eq!(pairing.status(), StatusCode::CREATED);
assert_eq!(
pairing
.headers()
.get(http::header::CACHE_CONTROL)
.and_then(|value| value.to_str().ok()),
Some("no-store")
);
let pairing_payload: serde_json::Value =
pairing.json().await.expect("pairing body should be JSON");
assert_eq!(pairing_payload["code"], "PAIR-123");
let ticket = client
.post(format!("{gateway_url}/api/users/me/vscodex/ws-tickets"))
.bearer_auth(&access_token)
.header("x-client-device-id", "browser-device-vscodex")
.json(&json!({ "device_id": "host-1", "user_id": "attacker" }))
.send()
.await
.expect("ticket request should complete");
assert_eq!(ticket.status(), StatusCode::CREATED);
assert_eq!(
ticket
.headers()
.get(http::header::CACHE_CONTROL)
.and_then(|value| value.to_str().ok()),
Some("no-store")
);
let ticket_payload: serde_json::Value =
ticket.json().await.expect("ticket body should be JSON");
assert_eq!(ticket_payload["ticket"], "ticket-123");
assert_eq!(ticket_payload["ws_url"], "wss://aether.example/vscodex/ws");
let deleted = client
.delete(format!("{gateway_url}/api/users/me/vscodex/devices/host-1"))
.bearer_auth(&access_token)
.header("x-client-device-id", "browser-device-vscodex")
.send()
.await
.expect("delete request should complete");
assert_eq!(deleted.status(), StatusCode::NO_CONTENT);
assert_eq!(
deleted
.headers()
.get(http::header::CACHE_CONTROL)
.and_then(|value| value.to_str().ok()),
Some("no-store")
);
let missing = client
.delete(format!(
"{gateway_url}/api/users/me/vscodex/devices/missing"
))
.bearer_auth(&access_token)
.header("x-client-device-id", "browser-device-vscodex")
.send()
.await
.expect("missing device request should complete");
assert_eq!(missing.status(), StatusCode::NOT_FOUND);
let missing_payload: serde_json::Value =
missing.json().await.expect("missing body should be JSON");
assert_eq!(missing_payload["detail"], "设备不存在");
let internal_denied = client
.delete(format!(
"{gateway_url}/api/users/me/vscodex/devices/internal-denied"
))
.bearer_auth(&access_token)
.header("x-client-device-id", "browser-device-vscodex")
.send()
.await
.expect("internal auth failure request should complete");
assert_eq!(internal_denied.status(), StatusCode::BAD_GATEWAY);
let internal_denied_payload: serde_json::Value = internal_denied
.json()
.await
.expect("internal auth failure body should be JSON");
assert_eq!(internal_denied_payload["detail"], "VS Codex 服务鉴权失败");
let redirected = client
.delete(format!(
"{gateway_url}/api/users/me/vscodex/devices/redirect"
))
.bearer_auth(&access_token)
.header("x-client-device-id", "browser-device-vscodex")
.send()
.await
.expect("redirecting sidecar request should complete");
assert_eq!(redirected.status(), StatusCode::BAD_GATEWAY);
let empty_success = client
.delete(format!(
"{gateway_url}/api/users/me/vscodex/devices/empty-ok"
))
.bearer_auth(&access_token)
.header("x-client-device-id", "browser-device-vscodex")
.send()
.await
.expect("empty sidecar success should complete");
assert_eq!(empty_success.status(), StatusCode::BAD_GATEWAY);
let pairing_exchange = client
.post(format!("{gateway_url}/api/vscodex/pair"))
.header("x-aether-client-ip", "203.0.113.99")
.json(&json!({
"code": "PAIR-123",
"name": "Office Mac",
"user_id": "attacker",
"device_token": "stolen"
}))
.send()
.await
.expect("public pairing exchange should complete");
assert_eq!(pairing_exchange.status(), StatusCode::CREATED);
assert_eq!(
pairing_exchange
.headers()
.get(http::header::CACHE_CONTROL)
.and_then(|value| value.to_str().ok()),
Some("no-store")
);
let pairing_exchange_payload: serde_json::Value = pairing_exchange
.json()
.await
.expect("pairing exchange body should be JSON");
assert_eq!(pairing_exchange_payload["device_id"], "host-2");
assert_eq!(pairing_exchange_payload["device_token"], "host-secret");
let mut websocket_request = format!("{gateway_url}/api/vscodex/ws")
.replace("http://", "ws://")
.into_client_request()
.expect("WebSocket request should build");
websocket_request.headers_mut().insert(
http::header::ORIGIN,
"https://aether.example".parse().unwrap(),
);
websocket_request.headers_mut().insert(
http::header::AUTHORIZATION,
"Bearer browser-aether-jwt".parse().unwrap(),
);
websocket_request.headers_mut().insert(
http::header::SEC_WEBSOCKET_PROTOCOL,
"vscodex.v1".parse().unwrap(),
);
let (mut websocket, websocket_response) = tokio_tungstenite::connect_async(websocket_request)
.await
.expect("gateway WebSocket should connect");
assert_eq!(
websocket_response
.headers()
.get(http::header::SEC_WEBSOCKET_PROTOCOL)
.and_then(|value| value.to_str().ok()),
Some("vscodex.v1")
);
websocket
.send(TungsteniteMessage::Text(
"{\"type\":\"auth\",\"ticket\":\"one-time-ticket\"}".into(),
))
.await
.expect("ticket frame should send");
let echoed = websocket
.next()
.await
.expect("echoed frame should arrive")
.expect("echoed frame should be valid");
assert_eq!(
echoed,
TungsteniteMessage::Text("echo:{\"type\":\"auth\",\"ticket\":\"one-time-ticket\"}".into())
);
websocket.close(None).await.expect("WebSocket should close");
assert_eq!(
captured_ws_handshake
.lock()
.expect("WebSocket handshake store should lock")
.clone(),
Some((Some("https://aether.example".to_string()), None))
);
let limited_gateway = build_router_with_state(
AppState::new()
.expect("limited gateway state should build")
.with_request_concurrency_limit(1),
);
let (limited_gateway_url, limited_gateway_handle) = start_server(limited_gateway).await;
let limited_ws_url =
format!("{limited_gateway_url}/api/vscodex/ws").replace("http://", "ws://");
let limited_ws_request = || {
let mut request = limited_ws_url
.as_str()
.into_client_request()
.expect("limited WebSocket request should build");
request
.headers_mut()
.insert("x-real-ip", "198.51.100.50".parse().unwrap());
request
};
let mut held_websockets = Vec::new();
for index in 0..16 {
let (websocket, _) = tokio_tungstenite::connect_async(limited_ws_request())
.await
.unwrap_or_else(|err| panic!("limited WebSocket {index} should connect: {err}"));
held_websockets.push(websocket);
}
let per_ip_limit_error = tokio_tungstenite::connect_async(limited_ws_request())
.await
.expect_err("seventeenth WebSocket from one IP should be rejected");
match per_ip_limit_error {
tokio_tungstenite::tungstenite::Error::Http(response) => {
assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS);
assert_eq!(
response
.headers()
.get(http::header::RETRY_AFTER)
.and_then(|value| value.to_str().ok()),
Some("1")
);
}
other => panic!("expected HTTP per-IP limit rejection, got {other:?}"),
}
held_websockets[0]
.send(TungsteniteMessage::Text(
r#"{"type":"auth","token":"test-auth-ok"}"#.into(),
))
.await
.expect("test authentication frame should send");
let auth_ok =
tokio::time::timeout(std::time::Duration::from_secs(1), held_websockets[0].next())
.await
.expect("test authentication response should arrive in time")
.expect("test authentication response should contain a frame")
.expect("test authentication response should be valid");
assert_eq!(
auth_ok,
TungsteniteMessage::Text(r#"{"type":"auth.ok","role":"operator"}"#.into())
);
tokio::time::sleep(std::time::Duration::from_millis(25)).await;
let (mut replacement_websocket, _) = tokio_tungstenite::connect_async(limited_ws_request())
.await
.expect("sidecar auth success should release one pending per-IP slot");
replacement_websocket
.close(None)
.await
.expect("replacement WebSocket should close");
for mut websocket in held_websockets {
websocket
.close(None)
.await
.expect("held WebSocket should close");
}
limited_gateway_handle.abort();
let blacklisted_gateway = build_router_with_state(
AppState::new()
.expect("blacklisted gateway state should build")
.with_admin_security_blacklist_for_tests([(
"127.0.0.1".to_string(),
"blocked".to_string(),
)]),
);
let (blacklisted_gateway_url, blacklisted_gateway_handle) =
start_server(blacklisted_gateway).await;
let blacklisted_ws_url =
format!("{blacklisted_gateway_url}/api/vscodex/ws").replace("http://", "ws://");
let blacklisted_error = tokio_tungstenite::connect_async(&blacklisted_ws_url)
.await
.expect_err("blacklisted WebSocket should be rejected");
match blacklisted_error {
tokio_tungstenite::tungstenite::Error::Http(response) => {
assert_eq!(response.status(), StatusCode::FORBIDDEN)
}
other => panic!("expected HTTP blacklist rejection, got {other:?}"),
}
blacklisted_gateway_handle.abort();
let requests = captured_requests
.lock()
.expect("captured request store should lock")
.clone();
assert_eq!(requests.len(), 9);
assert!(requests
.iter()
.all(|request| request.authorization.as_deref() == Some("Bearer sidecar-secret")));
assert!(requests[..8]
.iter()
.all(|request| request.path.starts_with("/internal/v1/users/user-auth-1/")));
assert_eq!(requests[8].path, "/v1/pairings/exchange");
assert_eq!(requests[1].body, Some(json!({ "name": "My Mac" })));
assert_eq!(requests[2].body, Some(json!({ "device_id": "host-1" })));
assert_eq!(
requests[8].body,
Some(json!({ "code": "PAIR-123", "name": "Office Mac" }))
);
assert_eq!(requests[8].client_ip.as_deref(), Some("127.0.0.1"));
assert!(requests[..8]
.iter()
.all(|request| request.client_ip.is_none()));
let _disabled = set_test_env_var("AETHER_VSCODEX_ENABLED", "false");
let disabled = client
.get(format!("{gateway_url}/api/users/me/vscodex/devices"))
.bearer_auth(&access_token)
.header("x-client-device-id", "browser-device-vscodex")
.send()
.await
.expect("disabled feature request should complete");
assert_eq!(disabled.status(), StatusCode::SERVICE_UNAVAILABLE);
let disabled_payload: serde_json::Value =
disabled.json().await.expect("disabled body should be JSON");
assert_eq!(disabled_payload["detail"], "VS Codex 服务未启用");
assert_eq!(
captured_requests
.lock()
.expect("captured request store should lock")
.len(),
9
);
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort();
upstream_handle.abort();
sidecar_handle.abort();
}
@@ -55,6 +55,9 @@ async fn resolve_wallet_auth_gate_with_cache(
None => WalletAccessDecision::wallet_unavailable(None),
};
if !auth_snapshot.api_key_is_standalone {
let wallet_is_unlimited = wallet
.as_ref()
.is_some_and(|wallet| wallet.limit_mode.eq_ignore_ascii_case("unlimited"));
let quota = if use_cache {
state
.find_user_daily_quota_availability_for_auth(&auth_snapshot.user_id)
@@ -71,7 +74,11 @@ async fn resolve_wallet_auth_gate_with_cache(
quota.remaining_usd,
))));
}
if decision.failure.is_none() && !quota.allow_wallet_overage && !has_remaining_quota {
if !wallet_is_unlimited
&& decision.failure.is_none()
&& !quota.allow_wallet_overage
&& !has_remaining_quota
{
return Ok(Some(WalletAccessDecision::balance_denied(Some(0.0))));
}
}
@@ -264,6 +271,23 @@ mod tests {
assert_eq!(decision.remaining, Some(4.0));
}
#[tokio::test]
async fn unlimited_wallet_ignores_exhausted_non_overage_quota() {
let mut wallet = empty_user_wallet();
wallet.limit_mode = "unlimited".to_string();
let state = state_with_wallet_and_quota(wallet, Some(quota_availability(10.0, 0.0, false)));
let auth_snapshot = ordinary_user_api_key_snapshot();
let decision = resolve_wallet_auth_gate(&state, &auth_snapshot)
.await
.expect("wallet gate should resolve")
.expect("wallet gate should return a decision");
assert!(decision.allowed);
assert_eq!(decision.failure, None);
assert_eq!(decision.remaining, None);
}
#[tokio::test]
async fn disabled_auth_capacity_cache_still_gates_wallet_reads() {
let mut state = state_with_wallet_and_quota(empty_user_wallet(), None);