mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-05 00:47:48 +08:00
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:
@@ -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,
|
||||
|
||||
+1
@@ -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(
|
||||
|
||||
@@ -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?,
|
||||
)),
|
||||
|
||||
@@ -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("a));
|
||||
}
|
||||
|
||||
#[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("a));
|
||||
}
|
||||
|
||||
#[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
|
||||
};
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user