feat(gateway): Codex/OpenAI Responses WebSocket 代理模式

在 /v1/responses 上支持 WebSocket 升级,把客户端帧中继到上游 Codex /
OpenAI Responses WebSocket 端点,同时保持既有的路由、鉴权、配额与用量
语义:

- 路由与准入:control/route/ai.rs 识别 WebSocket 升级请求;
  websocket/ingress.rs 复用 API Key 鉴权、IP 规则与并发许可,并引入
  独立的 WebSocket 连接许可
- 中继:websocket/responses/* 按 connection / session / turn 分层,
  帧解析归一化、socket 写入有界、continuation 保持调度亲和性
- 配额:orchestration/codex_quota_breaker.rs 在账号配额耗尽时熔断并
  自动恢复,不再直接断开客户端连接
- 用量:每个 turn 的终态用量落库,request_metadata 记录
  websocket_mode / websocket_transport,管理端与 usage 视图暴露
  is_websocket
- 管理端:provider 可配置 Responses WebSocket 开关
This commit is contained in:
AAEE86
2026-08-17 14:50:33 +08:00
committed by ZheFox
parent 9a0d346ff3
commit 71b54070e8
72 changed files with 10441 additions and 108 deletions
@@ -69,6 +69,9 @@ pub(crate) use aether_ai_formats::api::{
};
pub(crate) use aether_ai_formats::protocol::stream::CanonicalUsage as StreamingCanonicalUsage;
pub(crate) use aether_ai_formats::CODEX_RESPONSES_LITE_HEADER;
/// Codex client identity headers re-exported for out-of-crate probe binaries,
/// which must reach `aether_ai_formats` through this seam.
pub use aether_ai_formats::{CODEX_CLIENT_ORIGINATOR, CODEX_CLIENT_USER_AGENT};
pub(crate) fn parse_direct_request_body(
parts: &http::request::Parts,
+7 -6
View File
@@ -52,16 +52,17 @@ pub(crate) use self::planner::{
build_standard_family_sync_plan_and_reports, build_standard_stream_plan_from_decision,
build_standard_sync_plan_from_decision, candidate_auth_channel_skip_reason,
codex_model_capabilities_for_transport, extract_pool_sticky_session_token,
maybe_build_stream_decision_payload, maybe_build_stream_plan_payload,
maybe_build_sync_decision_payload, maybe_build_sync_plan_payload,
planner_is_matching_stream_request, provider_key_pool_score_id, provider_key_pool_score_scope,
read_candidate_transport_snapshot, record_local_runtime_candidate_skip_reason,
resolve_tunnel_scheduler_affinity_context, resolve_upstream_is_stream_for_provider,
set_local_openai_chat_execution_exhausted_diagnostic,
maybe_build_responses_websocket_decision, maybe_build_stream_decision_payload,
maybe_build_stream_plan_payload, maybe_build_sync_decision_payload,
maybe_build_sync_plan_payload, planner_is_matching_stream_request, provider_key_pool_score_id,
provider_key_pool_score_scope, read_candidate_transport_snapshot,
record_local_runtime_candidate_skip_reason, resolve_tunnel_scheduler_affinity_context,
resolve_upstream_is_stream_for_provider, set_local_openai_chat_execution_exhausted_diagnostic,
set_local_openai_image_execution_exhausted_diagnostic, validate_final_openai_provider_request,
CandidateFailureDiagnostic, CandidateFailureDiagnosticKind, EligibleLocalExecutionCandidate,
GatewayAuthApiKeySnapshot, GatewayProviderTransportSnapshot, LocalExecutionAttemptSource,
LocalExecutionCandidateKind, LocalResolvedOAuthRequestAuth, PlannerAppState,
ResponsesWebSocketBodyNormalization, ResponsesWebSocketDecision,
SkippedLocalExecutionCandidate,
};
pub(crate) use self::pure::*;
@@ -80,8 +80,9 @@ pub(crate) use self::standard::{
build_local_stream_plan_and_reports as build_standard_family_stream_plan_and_reports,
build_local_sync_attempt_source as build_standard_family_sync_attempt_source,
build_local_sync_plan_and_reports as build_standard_family_sync_plan_and_reports,
codex_model_capabilities_for_transport, set_local_openai_chat_execution_exhausted_diagnostic,
validate_final_openai_provider_request,
codex_model_capabilities_for_transport, maybe_build_responses_websocket_decision,
set_local_openai_chat_execution_exhausted_diagnostic, validate_final_openai_provider_request,
ResponsesWebSocketBodyNormalization, ResponsesWebSocketDecision,
};
pub(crate) use self::state::{
GatewayAuthApiKeySnapshot, GatewayProviderTransportSnapshot, LocalResolvedOAuthRequestAuth,
@@ -42,13 +42,14 @@ pub(crate) use self::openai::{
build_local_openai_responses_sync_attempt_source_for_kind,
build_local_openai_responses_sync_plan_and_reports_for_kind, copy_request_number_field,
copy_request_number_field_as, map_openai_reasoning_effort_to_claude_output,
map_openai_reasoning_effort_to_gemini_budget, maybe_build_stream_local_decision_payload,
map_openai_reasoning_effort_to_gemini_budget, maybe_build_responses_websocket_decision,
maybe_build_stream_local_decision_payload,
maybe_build_stream_local_openai_responses_decision_payload,
maybe_build_sync_local_decision_payload,
maybe_build_sync_local_openai_embedding_decision_payload,
maybe_build_sync_local_openai_responses_decision_payload, parse_openai_stop_sequences,
resolve_openai_chat_max_tokens, set_local_openai_chat_execution_exhausted_diagnostic,
value_as_u64,
value_as_u64, ResponsesWebSocketBodyNormalization, ResponsesWebSocketDecision,
};
pub(crate) use crate::ai_serving::normalize_standard_request_to_openai_chat_request;
pub(crate) use crate::ai_serving::{
@@ -23,6 +23,8 @@ pub(crate) use responses::{
build_local_openai_responses_stream_plan_and_reports_for_kind,
build_local_openai_responses_sync_attempt_source_for_kind,
build_local_openai_responses_sync_plan_and_reports_for_kind,
maybe_build_responses_websocket_decision,
maybe_build_stream_local_openai_responses_decision_payload,
maybe_build_sync_local_openai_responses_decision_payload,
maybe_build_sync_local_openai_responses_decision_payload, ResponsesWebSocketBodyNormalization,
ResponsesWebSocketDecision,
};
@@ -1,6 +1,16 @@
use crate::ai_serving::planner::common::endpoint_config_forces_body_stream_field;
use crate::ai_serving::planner::plan_builders::{AiStreamAttempt, AiSyncAttempt};
use crate::ai_serving::planner::spec_metadata::local_openai_responses_spec_metadata;
use crate::ai_serving::planner::standard::codex::codex_model_capabilities_for_transport;
use crate::ai_serving::planner::standard::normalize::build_local_openai_responses_request_body_with_codex_model_capabilities;
use crate::ai_serving::GatewayControlDecision;
use crate::orchestration::{
codex_quota_breaker_blocks_candidate, log_codex_quota_breaker_check_failure,
responses_websocket_adapter, ResponsesWebSocketAdapter,
};
use crate::{AiExecutionDecision, AppState, GatewayError};
use aether_runtime_state::RuntimeLockLease;
use std::collections::BTreeSet;
mod decision;
mod plans;
@@ -165,3 +175,314 @@ pub(crate) async fn maybe_build_stream_local_openai_responses_decision_payload(
Ok(None)
}
/// One eligible upstream plus the adapter that is allowed to speak to it.
///
/// The adapter is selected from the provider-scoped capability before the
/// decision leaves the planner. This prevents a public Responses socket from
/// choosing an arbitrary provider protocol after scheduling has completed.
pub(crate) struct ResponsesWebSocketDecision {
pub(crate) execution: AiExecutionDecision,
pub(crate) adapter: ResponsesWebSocketAdapter,
pub(crate) normalization: ResponsesWebSocketBodyNormalization,
}
/// Everything needed to re-run provider-body normalization for the candidate a
/// socket is already bound to.
///
/// A continuation turn (`previous_response_id` on the bound upstream) cannot
/// re-enter the planner, because planning selects a candidate and a different
/// key would break the response chain. Without this, such turns reached the
/// provider with only their `model` rewritten — skipping model directives,
/// endpoint body rules, and the Codex body contract that turn 1 received.
///
/// This value holds cloned scalars and JSON only: no candidate, no pool key
/// lease, no `AppState`. It cannot influence selection.
#[derive(Debug, Clone)]
pub(crate) struct ResponsesWebSocketBodyNormalization {
provider_type: String,
provider_api_format: String,
client_api_format: String,
mapped_model: String,
requested_model: String,
upstream_is_stream: bool,
force_body_stream_field: bool,
body_rules: Option<serde_json::Value>,
request_headers: http::HeaderMap,
codex_model_capabilities: Option<crate::ai_serving::CodexResponsesModelCapabilities>,
model_directive_patch: Option<serde_json::Value>,
}
impl ResponsesWebSocketBodyNormalization {
/// Builds a normalizer for a plain `openai:responses` upstream with no
/// endpoint body rules, directives or Codex capabilities, so relay tests can
/// construct a bound connection without standing up a provider snapshot.
#[cfg(test)]
pub(crate) fn for_tests(mapped_model: &str) -> Self {
Self {
provider_type: "openai".to_string(),
provider_api_format: "openai:responses".to_string(),
client_api_format: "openai:responses".to_string(),
mapped_model: mapped_model.to_string(),
requested_model: mapped_model.to_string(),
upstream_is_stream: true,
force_body_stream_field: false,
body_rules: None,
request_headers: http::HeaderMap::new(),
codex_model_capabilities: None,
model_directive_patch: None,
}
}
#[cfg(test)]
pub(crate) fn with_provider_type_for_tests(mut self, provider_type: &str) -> Self {
self.provider_type = provider_type.to_string();
self
}
#[cfg(test)]
pub(crate) fn with_model_directive_patch_for_tests(mut self, patch: serde_json::Value) -> Self {
self.model_directive_patch = Some(patch);
self
}
/// Applies the same body transformations the planner applied on the turn
/// that bound this upstream.
///
/// Mirrors the same-format branch of
/// `resolve_local_openai_responses_candidate_payload_parts`. The
/// cross-format, Kiro, Windsurf and Antigravity branches are unreachable
/// here: the WebSocket planner only returns candidates whose provider API
/// format is `openai:responses`.
///
/// Returns `None` when normalization fails, leaving the caller to fall back
/// to the unnormalized event — a continuation cannot re-select a candidate,
/// so failing the turn outright would be worse than sending it as-is.
pub(crate) fn normalize_response_create(
&self,
client_event: &serde_json::Value,
) -> Option<serde_json::Value> {
use crate::ai_serving::planner::common::{
enforce_provider_body_stream_policy, request_requires_body_stream_field,
};
let source_model = client_event
.get("model")
.and_then(serde_json::Value::as_str)
.unwrap_or(self.requested_model.as_str());
let require_body_stream_field =
request_requires_body_stream_field(client_event, self.force_body_stream_field);
let mut body = build_local_openai_responses_request_body_with_codex_model_capabilities(
client_event,
&self.mapped_model,
self.upstream_is_stream,
self.force_body_stream_field,
self.provider_type.as_str(),
self.provider_api_format.as_str(),
self.body_rules.as_ref(),
&self.request_headers,
self.codex_model_capabilities.as_ref(),
false,
)?;
if let Some(patch) = self.model_directive_patch.as_ref() {
crate::ai_serving::apply_model_directive_mapping_patch(&mut body, patch);
// The patch is a deep merge and may reintroduce `stream`.
enforce_provider_body_stream_policy(
&mut body,
self.provider_api_format.as_str(),
self.upstream_is_stream,
require_body_stream_field,
);
}
crate::ai_serving::finalize_openai_provider_request_with_codex_model_capabilities(
&mut body,
crate::ai_serving::OpenAiProviderRequestFinalization {
source_api_format: self.client_api_format.as_str(),
provider_api_format: self.provider_api_format.as_str(),
provider_type: self.provider_type.as_str(),
provider_model: self.mapped_model.as_str(),
source_model,
body_rules: self.body_rules.as_ref(),
upstream_is_stream: self.upstream_is_stream,
require_body_stream_field,
},
self.codex_model_capabilities.as_ref(),
)
.ok()?;
Some(body)
}
}
/// Builds one upstream decision for a Responses WebSocket turn. The session
/// reuses this decision for same-model turns and invokes the planner again when
/// a later `response.create` changes the public model.
pub(crate) async fn maybe_build_responses_websocket_decision(
state: &AppState,
parts: &http::request::Parts,
trace_id: &str,
decision: &GatewayControlDecision,
body_json: &serde_json::Value,
excluded_key_ids: Option<&BTreeSet<String>>,
excluded_codex_account_ids: Option<&BTreeSet<String>>,
) -> Result<Option<ResponsesWebSocketDecision>, GatewayError> {
let Some(spec) = resolve_stream_spec(crate::ai_serving::OPENAI_RESPONSES_STREAM_PLAN_KIND)
else {
return Ok(None);
};
let Some(input) = resolve_local_openai_responses_decision_input(
state,
parts,
trace_id,
decision,
body_json,
spec.decision_kind,
)
.await?
else {
return Ok(None);
};
let body_json = input.effective_body_json(body_json);
let (mut source, _) = build_local_openai_responses_candidate_attempt_source(
state, trace_id, &input, body_json, spec,
)
.await?;
while let Some(attempt) = source.next_attempt().await? {
let pool_key_lease = attempt.eligible.orchestration.pool_key_lease.clone();
if excluded_key_ids
.is_some_and(|key_ids| key_ids.contains(attempt.eligible.candidate.key_id.as_str()))
{
release_responses_websocket_planning_lease(state, pool_key_lease.as_ref()).await;
continue;
}
let Some(adapter) = responses_websocket_adapter(
&attempt.eligible.transport.provider.provider_type,
attempt.eligible.transport.provider.config.as_ref(),
) else {
release_responses_websocket_planning_lease(state, pool_key_lease.as_ref()).await;
continue;
};
// Captured before `attempt` is consumed so a later continuation turn can
// reproduce this candidate's body normalization without re-planning.
let transport = std::sync::Arc::clone(&attempt.eligible.transport);
let candidate_provider_api_format = attempt.eligible.provider_api_format.clone();
let payload = match maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
.await
{
Ok(Some(payload)) => payload,
Ok(None) => {
release_responses_websocket_planning_lease(state, pool_key_lease.as_ref()).await;
continue;
}
Err(error) => {
release_responses_websocket_planning_lease(state, pool_key_lease.as_ref()).await;
return Err(error);
}
};
if payload
.provider_type
.as_deref()
.is_some_and(|value| value.trim().eq_ignore_ascii_case("codex"))
&& crate::orchestration::codex_account_id_from_headers(
&payload.provider_request_headers,
)
.is_some_and(|account_id| {
excluded_codex_account_ids
.is_some_and(|account_ids| account_ids.contains(account_id))
})
{
release_responses_websocket_planning_lease(state, pool_key_lease.as_ref()).await;
continue;
}
match codex_quota_breaker_blocks_candidate(
state,
payload.provider_type.as_deref(),
payload.key_id.as_deref(),
&payload.provider_request_headers,
)
.await
{
Ok(true) => {
release_responses_websocket_planning_lease(state, pool_key_lease.as_ref()).await;
continue;
}
Ok(false) => {}
Err(error) => log_codex_quota_breaker_check_failure(&error),
}
if payload
.provider_type
.as_deref()
.is_some_and(|value| adapter.supports_provider_type(value))
&& payload.provider_api_format.as_deref().is_some_and(|value| {
crate::ai_serving::normalize_api_format_alias(value) == "openai:responses"
})
{
let mapped_model = payload.mapped_model.clone().unwrap_or_default();
let source_model = body_json
.get("model")
.and_then(serde_json::Value::as_str)
.unwrap_or(input.requested_model.as_str());
let normalization = ResponsesWebSocketBodyNormalization {
provider_type: transport.provider.provider_type.clone(),
provider_api_format: candidate_provider_api_format.clone(),
client_api_format: local_openai_responses_spec_metadata(spec)
.api_format
.to_string(),
requested_model: input.requested_model.clone(),
upstream_is_stream: payload.upstream_is_stream,
force_body_stream_field: endpoint_config_forces_body_stream_field(
transport.endpoint.config.as_ref(),
),
body_rules: transport.endpoint.body_rules.clone(),
request_headers: input.effective_headers(&parts.headers).clone(),
codex_model_capabilities: codex_model_capabilities_for_transport(
&transport,
candidate_provider_api_format.as_str(),
mapped_model.as_str(),
source_model,
),
model_directive_patch: input
.model_directive_policy
.resolve_reasoning(
candidate_provider_api_format.as_str(),
Some(&input.requested_model),
)
.mapping_patch_for_mapped_model(mapped_model.as_str())
.ok()
.flatten(),
mapped_model,
};
return Ok(Some(ResponsesWebSocketDecision {
execution: payload,
adapter,
normalization,
}));
}
release_responses_websocket_planning_lease(state, pool_key_lease.as_ref()).await;
}
Ok(None)
}
async fn release_responses_websocket_planning_lease(
state: &AppState,
lease: Option<&RuntimeLockLease>,
) {
let Some(lease) = lease else {
return;
};
if let Err(error) =
crate::handlers::shared::provider_pool::release_admin_provider_pool_key_lease(
state.runtime_state.as_ref(),
lease,
)
.await
{
tracing::warn!(
error = ?error,
"gateway Responses WebSocket planner failed to release an unused pool key lease"
);
}
}
+11 -3
View File
@@ -1,13 +1,17 @@
use axum::body::Body;
use axum::extract::Request;
use axum::http::{header, HeaderValue, Response, StatusCode};
use axum::routing::{any, post};
use axum::routing::{any, get, post};
use axum::Router;
use super::{aliyun, claude, doubao, gemini, jina, openai};
use crate::api::response::build_local_http_error_response_with_request_path;
use crate::headers::extract_or_generate_trace_id;
use crate::{handlers::proxy::proxy_request, state::AppState, GatewayError};
use crate::{
handlers::proxy::{proxy_request, responses_websocket},
state::AppState,
GatewayError,
};
// Router registration patterns live here so AI public ingress has a single mount registry.
// They intentionally stay separate from manifest-facing route inventories in constants.rs,
@@ -51,7 +55,11 @@ const AI_ANY_ROUTE_PATTERNS: &[&str] = &[
pub(crate) fn mount_ai_routes(mut router: Router<AppState>) -> Router<AppState> {
for path in AI_POST_ROUTE_PATTERNS {
router = router.route(path, post(proxy_request));
router = if *path == "/v1/responses" {
router.route(path, get(responses_websocket).post(proxy_request))
} else {
router.route(path, post(proxy_request))
};
}
for path in CLAUDE_POST_ROUTE_PATTERNS {
router = router.route(
+34
View File
@@ -50,12 +50,40 @@ pub(crate) async fn health(State(state): State<AppState>) -> impl IntoResponse {
"rejected": snapshot.rejected,
})
});
let websocket_connection_concurrency =
state
.websocket_connection_concurrency_snapshot()
.map(|snapshot| {
json!({
"limit": snapshot.limit,
"in_flight": snapshot.in_flight,
"available_permits": snapshot.available_permits,
"high_watermark": snapshot.high_watermark,
"rejected": snapshot.rejected,
})
});
let distributed_websocket_connection_concurrency = state
.distributed_websocket_connection_concurrency_snapshot()
.await
.ok()
.flatten()
.map(|snapshot| {
json!({
"limit": snapshot.limit,
"in_flight": snapshot.in_flight,
"available_permits": snapshot.available_permits,
"high_watermark": snapshot.high_watermark,
"rejected": snapshot.rejected,
})
});
Json(json!({
"status": "ok",
"component": "aether-gateway",
"control_api_enabled": true,
"request_concurrency": request_concurrency,
"distributed_request_concurrency": distributed_request_concurrency,
"websocket_connection_concurrency": websocket_connection_concurrency,
"distributed_websocket_connection_concurrency": distributed_websocket_connection_concurrency,
}))
}
@@ -113,6 +141,12 @@ pub(crate) async fn frontdoor_manifest(State(state): State<AppState>) -> impl In
"execution_runtime_configured": state.execution_runtime_configured(),
"request_concurrency_enabled": state.request_concurrency_snapshot().is_some(),
"distributed_request_concurrency_enabled": state.distributed_request_gate.is_some(),
"websocket_connection_concurrency_enabled": state
.websocket_connection_concurrency_snapshot()
.is_some(),
"distributed_websocket_connection_concurrency_enabled": state
.distributed_websocket_connection_gate
.is_some(),
"frontdoor_cors_enabled": cors_enabled,
"frontdoor_cors_allow_credentials": cors_allow_credentials,
"frontdoor_cors_allowed_origins": cors_allowed_origins,
+51 -1
View File
@@ -35,7 +35,10 @@ pub(super) fn classify_ai_public_route(
"openai:rerank",
true,
))
} else if method == http::Method::POST
} else if (method == http::Method::POST
|| (method == http::Method::GET
&& normalized_path == "/v1/responses"
&& is_websocket_upgrade_request(headers)))
&& matches!(normalized_path, "/v1/responses" | "/v1/responses/compact")
{
if normalized_path.ends_with("/compact") {
@@ -199,6 +202,24 @@ fn claude_request_auth_channel(headers: &http::HeaderMap) -> &'static str {
}
}
fn is_websocket_upgrade_request(headers: &http::HeaderMap) -> bool {
let has_upgrade_connection = headers
.get(http::header::CONNECTION)
.and_then(|value| value.to_str().ok())
.is_some_and(|value| {
value
.split(',')
.map(str::trim)
.any(|value| value.eq_ignore_ascii_case("upgrade"))
});
let has_websocket_upgrade = headers
.get(http::header::UPGRADE)
.and_then(|value| value.to_str().ok())
.is_some_and(|value| value.eq_ignore_ascii_case("websocket"));
has_upgrade_connection && has_websocket_upgrade
}
fn is_gemini_operation_method(method: &http::Method, normalized_path: &str) -> bool {
method == http::Method::GET
|| (method == http::Method::POST && normalized_path.ends_with(":cancel"))
@@ -242,3 +263,32 @@ fn classify_antigravity_v1internal_route(
execution_runtime_candidate,
))
}
#[cfg(test)]
mod tests {
use axum::http::header::{CONNECTION, UPGRADE};
use axum::http::{HeaderMap, HeaderValue, Method};
use super::classify_ai_public_route;
#[test]
fn classifies_websocket_upgrade_on_responses_route() {
let mut headers = HeaderMap::new();
headers.insert(CONNECTION, HeaderValue::from_static("keep-alive, Upgrade"));
headers.insert(UPGRADE, HeaderValue::from_static("websocket"));
let route = classify_ai_public_route(&Method::GET, "/v1/responses", &headers)
.expect("Responses WebSocket should be an AI public route");
assert_eq!(route.route_class, "ai_public");
assert_eq!(route.route_family, "openai");
assert_eq!(route.route_kind, "responses");
assert_eq!(route.auth_endpoint_signature, "openai:responses");
}
#[test]
fn does_not_classify_plain_get_as_responses_websocket() {
assert!(
classify_ai_public_route(&Method::GET, "/v1/responses", &HeaderMap::new()).is_none()
);
}
}
@@ -0,0 +1,62 @@
//! Shared admission helpers for local upstream execution.
//!
//! The stream candidate loop and long-lived WebSocket turns both need to
//! participate in the same gateway-wide upstream execution gate. Keep the
//! provider abstraction here so tests can supply an isolated gate while
//! production callers use `AppState` directly.
use std::time::Duration;
use aether_runtime::{ConcurrencyGate, ConcurrencyPermit};
use tokio::time::timeout;
use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{AppState, GatewayError};
pub(crate) const UPSTREAM_EXECUTION_GATE_NAME: &str = "gateway_upstream_execution";
pub(crate) trait UpstreamExecutionGateProvider {
fn upstream_execution_gate(&self) -> Option<&ConcurrencyGate>;
fn upstream_execution_gate_queue_budget(&self) -> Duration;
}
impl UpstreamExecutionGateProvider for AppState {
fn upstream_execution_gate(&self) -> Option<&ConcurrencyGate> {
self.upstream_execution_gate.as_deref()
}
fn upstream_execution_gate_queue_budget(&self) -> Duration {
self.frontdoor_runtime_guards.internal_gate_queue_budget
}
}
/// Acquires the shared gateway-wide upstream execution permit.
///
/// A missing gate is an intentional configuration (unlimited), so callers
/// receive `Ok(None)`. Saturation keeps the existing candidate-level
/// `AdmissionTimeout` contract used by the HTTP stream path.
pub(crate) async fn acquire_upstream_execution_gate(
state: &(impl UpstreamExecutionGateProvider + ?Sized),
trace_id: &str,
) -> Result<Option<ConcurrencyPermit>, GatewayError> {
let Some(gate) = state.upstream_execution_gate() else {
return Ok(None);
};
let budget = state.upstream_execution_gate_queue_budget();
let gate_wait_started_at = std::time::Instant::now();
match timeout(budget, gate.acquire()).await {
Ok(Ok(permit)) => {
observe_gateway_stage_ms(
"upstream_execution_gate_wait",
gate_wait_started_at.elapsed().as_millis() as u64,
);
Ok(Some(permit))
}
Ok(Err(err)) => Err(GatewayError::Internal(err.to_string())),
Err(_) => Err(GatewayError::AdmissionTimeout {
trace_id: trace_id.to_string(),
gate: UPSTREAM_EXECUTION_GATE_NAME,
queue_budget_ms: budget.as_millis() as u64,
}),
}
}
@@ -3,6 +3,7 @@ use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
pub(crate) mod admission;
mod chatgpt_web_image;
mod constants;
mod fallback;
@@ -23,6 +24,9 @@ pub(crate) mod transport;
mod transport_failure;
mod windsurf;
pub(crate) use self::admission::{
acquire_upstream_execution_gate, UpstreamExecutionGateProvider, UPSTREAM_EXECUTION_GATE_NAME,
};
pub(crate) use self::chatgpt_web_image::maybe_execute_chatgpt_web_image_sync;
pub(crate) use self::constants::{
MAX_ERROR_BODY_BYTES, MAX_STREAM_PREFETCH_BYTES, MAX_STREAM_PREFETCH_FRAMES,
@@ -20,9 +20,11 @@ use crate::ai_serving::LocalExecutionAttemptSource;
use crate::clock::current_unix_ms;
use crate::control::GatewayControlDecision;
use crate::execution_runtime::{
build_transport_error_stop_response, execute_execution_runtime_stream_with_retry_scope,
acquire_upstream_execution_gate, build_transport_error_stop_response,
execute_execution_runtime_stream_with_retry_scope,
execute_execution_runtime_sync_with_retry_scope,
mark_stream_candidate_watchdog_terminal_started, StreamCandidateWatchdogProgress,
UpstreamExecutionGateProvider, UPSTREAM_EXECUTION_GATE_NAME,
};
use crate::executor::{
build_local_execution_exhaustion, mark_deferred_upstream_response, LocalExecutionRequestOutcome,
@@ -43,7 +45,6 @@ use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{AppState, GatewayError};
const DEFAULT_STREAM_FIRST_BYTE_WATCHDOG_TIMEOUT_MS: u64 = 30_000;
const UPSTREAM_EXECUTION_GATE_NAME: &str = "gateway_upstream_execution";
const UPSTREAM_TARGET_GATE_NAME: &str = "gateway_upstream_target";
const UPSTREAM_EXECUTION_GATE_HOLD_STREAM_RESPONSE_ENV: &str =
"AETHER_GATEWAY_UPSTREAM_EXECUTION_GATE_HOLD_STREAM_RESPONSE";
@@ -1612,47 +1613,6 @@ fn hold_response_upstream_execution_permit(
Response::from_parts(parts, Body::from_stream(stream))
}
trait UpstreamExecutionGateProvider {
fn upstream_execution_gate(&self) -> Option<&aether_runtime::ConcurrencyGate>;
fn upstream_execution_gate_queue_budget(&self) -> Duration;
}
impl UpstreamExecutionGateProvider for AppState {
fn upstream_execution_gate(&self) -> Option<&aether_runtime::ConcurrencyGate> {
self.upstream_execution_gate.as_deref()
}
fn upstream_execution_gate_queue_budget(&self) -> Duration {
self.frontdoor_runtime_guards.internal_gate_queue_budget
}
}
async fn acquire_upstream_execution_gate(
state: &(impl UpstreamExecutionGateProvider + ?Sized),
trace_id: &str,
) -> Result<Option<ConcurrencyPermit>, GatewayError> {
let Some(gate) = state.upstream_execution_gate() else {
return Ok(None);
};
let budget = state.upstream_execution_gate_queue_budget();
let gate_wait_started_at = std::time::Instant::now();
match timeout(budget, gate.acquire()).await {
Ok(Ok(permit)) => {
observe_gateway_stage_ms(
"upstream_execution_gate_wait",
gate_wait_started_at.elapsed().as_millis() as u64,
);
Ok(Some(permit))
}
Ok(Err(err)) => Err(GatewayError::Internal(err.to_string())),
Err(_) => Err(GatewayError::AdmissionTimeout {
trace_id: trace_id.to_string(),
gate: UPSTREAM_EXECUTION_GATE_NAME,
queue_budget_ms: budget.as_millis() as u64,
}),
}
}
pub(crate) async fn mark_unused_local_candidate_items<T, FPlan, FContext>(
state: &AppState,
remaining: Vec<T>,
@@ -1300,6 +1300,10 @@ fn provider_query_pool_catalog_key_context(
provider_type,
quota_snapshot,
),
quota_hard_blocked: admin_provider_pool_pure::admin_pool_key_quota_hard_blocked(
key,
provider_type,
),
health_score,
latency_avg_ms,
catalog_lru_score: Some(key.last_used_at_unix_secs.unwrap_or(0) as f64),
@@ -154,6 +154,8 @@ pub(crate) struct AdminProviderCreateRequest {
#[serde(default)]
pub(crate) codex_fingerprint_convergence_enabled: Option<bool>,
#[serde(default)]
pub(crate) responses_websocket_enabled: Option<bool>,
#[serde(default)]
pub(crate) is_active: Option<bool>,
#[serde(default)]
pub(crate) concurrent_limit: Option<i32>,
@@ -215,6 +217,8 @@ pub(crate) struct AdminProviderUpdateRequest {
#[serde(default)]
pub(crate) codex_fingerprint_convergence_enabled: Option<bool>,
#[serde(default)]
pub(crate) responses_websocket_enabled: Option<bool>,
#[serde(default)]
pub(crate) is_active: Option<bool>,
#[serde(default)]
pub(crate) concurrent_limit: Option<i32>,
@@ -4,7 +4,7 @@ use crate::handlers::admin::provider::shared::support::{
};
use crate::handlers::admin::shared::unix_secs_to_rfc3339;
use crate::handlers::public::{request_candidate_event_unix_ms, request_candidate_status_label};
use crate::orchestration::codex_cyber_flag_passthrough_enabled;
use crate::orchestration::{codex_cyber_flag_passthrough_enabled, responses_websocket_adapter};
use crate::provider_key_auth::provider_key_effective_api_formats;
use aether_data_contracts::repository::candidates::{
RequestCandidateStatus, StoredRequestCandidate,
@@ -222,6 +222,7 @@ pub(crate) fn build_admin_provider_summary_value(
&provider.provider_type,
provider.config.as_ref(),
),
"responses_websocket_enabled": responses_websocket_adapter(&provider.provider_type, provider.config.as_ref()).is_some(),
"ops_quota_alert_enabled": ops_quota_alert_enabled,
"created_at": endpoint_timestamp_or_now(provider.created_at_unix_ms, now_unix_secs),
"updated_at": endpoint_timestamp_or_now(provider.updated_at_unix_secs, now_unix_secs),
@@ -208,6 +208,53 @@ pub(crate) fn normalize_chat_pii_redaction_config(
}
}
pub(crate) fn set_responses_websocket_enabled(
config: &mut serde_json::Map<String, serde_json::Value>,
enabled: bool,
) -> Result<(), String> {
let mut responses = match config.remove("responses_websocket") {
None => serde_json::Map::new(),
Some(serde_json::Value::Object(config)) => config,
Some(_) => return Err("config.responses_websocket 必须是 JSON 对象".to_string()),
};
responses.insert("enabled".to_string(), serde_json::Value::Bool(enabled));
config.insert(
"responses_websocket".to_string(),
serde_json::Value::Object(responses),
);
Ok(())
}
pub(crate) fn remove_responses_websocket_enabled(
config: &mut serde_json::Map<String, serde_json::Value>,
) {
let Some(serde_json::Value::Object(responses)) = config.get_mut("responses_websocket") else {
return;
};
responses.remove("enabled");
if responses.is_empty() {
config.remove("responses_websocket");
}
}
pub(crate) fn validate_responses_websocket_config(
config: &serde_json::Map<String, serde_json::Value>,
) -> Result<(), String> {
if let Some(value) = config.get("responses_websocket") {
let responses = value
.as_object()
.ok_or_else(|| "config.responses_websocket 必须是 JSON 对象".to_string())?;
let enabled = responses
.get("enabled")
.ok_or_else(|| "config.responses_websocket.enabled 为必填布尔值".to_string())?;
if !enabled.is_boolean() {
return Err("config.responses_websocket.enabled 必须是布尔值".to_string());
}
}
Ok(())
}
pub(crate) fn validate_vertex_api_formats(
provider_type: &str,
auth_type: &str,
@@ -258,7 +305,9 @@ mod tests {
normalize_api_format_list, normalize_auth_type, normalize_auth_type_by_format,
normalize_chat_pii_redaction_config, normalize_pool_advanced_config,
normalize_provider_type_input, normalize_rate_multipliers,
reconcile_allow_auth_channel_mismatch_formats, validate_vertex_api_formats,
reconcile_allow_auth_channel_mismatch_formats, remove_responses_websocket_enabled,
set_responses_websocket_enabled, validate_responses_websocket_config,
validate_vertex_api_formats,
};
use serde_json::json;
@@ -317,6 +366,21 @@ mod tests {
);
}
#[test]
fn responses_websocket_setting_is_available_to_explicitly_enabled_providers() {
let mut config = serde_json::Map::new();
set_responses_websocket_enabled(&mut config, true)
.expect("Responses setting should be accepted");
assert_eq!(
config.get("responses_websocket"),
Some(&json!({"enabled": true}))
);
validate_responses_websocket_config(&config).expect("Responses setting should validate");
remove_responses_websocket_enabled(&mut config);
assert!(config.get("responses_websocket").is_none());
}
#[test]
fn normalize_auth_type_supports_bearer() {
assert_eq!(
@@ -7,6 +7,8 @@ use crate::handlers::admin::provider::shared::support::{
use crate::handlers::admin::provider::write::normalize::normalize_chat_pii_redaction_config;
use crate::handlers::admin::provider::write::normalize::normalize_pool_advanced_config;
use crate::handlers::admin::provider::write::normalize::normalize_provider_type_input;
use crate::handlers::admin::provider::write::normalize::set_responses_websocket_enabled;
use crate::handlers::admin::provider::write::normalize::validate_responses_websocket_config;
use crate::handlers::admin::request::AdminAppState;
use crate::handlers::admin::shared::normalize_json_object;
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider;
@@ -179,6 +181,10 @@ pub(crate) async fn build_admin_create_provider_record(
config_map.insert("chat_pii_redaction".to_string(), value);
}
}
if let Some(enabled) = payload.responses_websocket_enabled {
set_responses_websocket_enabled(&mut config_map, enabled)?;
}
validate_responses_websocket_config(&config_map)?;
let config = (!config_map.is_empty()).then_some(serde_json::Value::Object(config_map));
crate::provider_transport::validate_anthropic_compatibility_profile_config(config.as_ref())
.map_err(|_| "无效的 Anthropic compatibility profile".to_string())?;
@@ -7,6 +7,8 @@ use crate::handlers::admin::provider::shared::support::{
use crate::handlers::admin::provider::write::normalize::normalize_chat_pii_redaction_config;
use crate::handlers::admin::provider::write::normalize::normalize_pool_advanced_config;
use crate::handlers::admin::provider::write::normalize::normalize_provider_type_input;
use crate::handlers::admin::provider::write::normalize::set_responses_websocket_enabled;
use crate::handlers::admin::provider::write::normalize::validate_responses_websocket_config;
use crate::handlers::admin::request::AdminAppState;
use crate::handlers::admin::shared::normalize_json_object;
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider;
@@ -337,6 +339,14 @@ pub(crate) async fn build_admin_update_provider_record(
}
}
if fields.contains("responses_websocket_enabled") {
let enabled = payload
.responses_websocket_enabled
.ok_or_else(|| "responses_websocket_enabled 必须是布尔值".to_string())?;
set_responses_websocket_enabled(&mut config_map, enabled)?;
}
validate_responses_websocket_config(&config_map)?;
updated.config = (!config_map.is_empty()).then_some(serde_json::Value::Object(config_map));
crate::provider_transport::validate_anthropic_compatibility_profile_config(
updated.config.as_ref(),
@@ -1,5 +1,6 @@
mod body_buffer;
mod local;
mod websocket;
use self::body_buffer::{
buffer_and_normalize_request_body, build_request_body_buffer_error_response,
@@ -8,6 +9,7 @@ use self::body_buffer::{
use self::local::{
maybe_build_local_admin_proxy_response, maybe_build_local_internal_proxy_response,
};
pub(crate) use self::websocket::responses::responses_websocket;
use super::internal::resolve_local_proxy_execution_path;
pub(crate) use super::public::matches_model_mapping_for_models;
use crate::ai_serving::api::{
@@ -0,0 +1,293 @@
//! Authenticated public WebSocket upgrade admission shared by AI adapters.
use std::future::Future;
use std::net::SocketAddr;
use axum::body::Body;
use axum::extract::ws::{WebSocket, WebSocketUpgrade};
use axum::http::{HeaderMap, Method, Response, StatusCode, Uri};
use tracing::{info, warn};
use crate::api::response::{
build_local_auth_rejection_response, build_local_http_error_response,
build_local_overloaded_response,
};
use crate::control::{
trusted_auth_local_rejection, GatewayControlDecision, GatewayLocalAuthRejection,
};
use crate::handlers::proxy::websocket::session::{WebSocketSessionLimits, WEBSOCKET_LOG_TRANSPORT};
use crate::handlers::shared::ip_rules_allow;
use crate::headers::{effective_client_ip, extract_or_generate_trace_id};
use crate::router::RequestAdmissionError;
use crate::{AppState, GatewayError};
/// Request facts that survive the HTTP Upgrade and are needed by a protocol
/// adapter for planning, rate limiting, and connection-scoped audit logs.
pub(crate) struct WebSocketRequestContext {
pub(crate) trace_id: String,
pub(crate) headers: HeaderMap,
pub(crate) uri: Uri,
pub(crate) remote_addr: SocketAddr,
pub(crate) decision: GatewayControlDecision,
pub(crate) rpm_bypassed: bool,
/// Held for the lifetime of the upgraded socket. The Responses session
/// polls its health and closes the client when a distributed lease is
/// revoked or expires.
pub(crate) websocket_connection_permit: Option<aether_runtime::AdmissionPermit>,
}
/// Adapter-specific wording and event identifiers for generic upgrade checks.
#[derive(Clone, Copy)]
pub(crate) struct WebSocketIngressSpec {
pub(crate) route_unavailable_message: &'static str,
pub(crate) ip_whitelist_failure_event_name: &'static str,
}
/// Performs the HTTP-only part of an AI WebSocket request.
///
/// The ordinary request permit covers only the HTTP Upgrade window. A
/// dedicated WebSocket connection permit is held for the socket lifetime so
/// idle clients cannot consume capacity reserved for normal HTTP requests.
pub(crate) async fn upgrade_authenticated_ai_websocket<F, Fut>(
state: AppState,
remote_addr: SocketAddr,
ws: WebSocketUpgrade,
headers: HeaderMap,
uri: Uri,
limits: WebSocketSessionLimits,
spec: WebSocketIngressSpec,
run_session: F,
) -> Result<Response<Body>, GatewayError>
where
F: FnOnce(WebSocket, AppState, WebSocketRequestContext) -> Fut + Send + 'static,
Fut: Future<Output = ()> + Send + 'static,
{
let trace_id = extract_or_generate_trace_id(&headers);
let client_ip = effective_client_ip(&headers, &remote_addr);
if state.admin_security_ip_blacklisted(client_ip).await? {
return build_local_http_error_response(
&trace_id,
None,
StatusCode::FORBIDDEN,
"当前 IP 已被禁止访问",
);
}
let request_context = crate::control::resolve_public_request_context(
&state,
&Method::GET,
&uri,
&headers,
&trace_id,
)
.await?;
let Some(decision) = request_context.control_decision else {
return build_local_http_error_response(
&trace_id,
None,
StatusCode::NOT_FOUND,
spec.route_unavailable_message,
);
};
if let Some(rejection) = trusted_auth_local_rejection(Some(&decision), &headers) {
return build_local_auth_rejection_response(&trace_id, Some(&decision), &rejection);
}
let Some(auth_context) = decision.auth_context.as_ref() else {
return build_local_auth_rejection_response(
&trace_id,
Some(&decision),
&GatewayLocalAuthRejection::InvalidApiKey,
);
};
if !auth_context.access_allowed {
return build_local_auth_rejection_response(
&trace_id,
Some(&decision),
&GatewayLocalAuthRejection::InvalidApiKey,
);
}
if !ip_rules_allow(auth_context.ip_rules.as_deref(), client_ip) {
return build_local_auth_rejection_response(
&trace_id,
Some(&decision),
&GatewayLocalAuthRejection::IpNotAllowed {
remote_ip: client_ip.to_string(),
},
);
}
let ip_whitelisted = match state.admin_security_ip_whitelisted(client_ip).await {
Ok(value) => value,
Err(error) => {
warn!(
event_name = spec.ip_whitelist_failure_event_name,
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %trace_id,
client_ip = %client_ip,
error = ?error,
"gateway continued with WebSocket rate limiting after IP whitelist check error"
);
false
}
};
let request_permit = match state.try_acquire_request_permit().await {
Ok(permit) => permit,
Err(error) => {
return websocket_admission_error_response(
&trace_id,
&decision,
Some(uri.path()),
error,
)
}
};
let websocket_connection_permit = match state.try_acquire_websocket_connection_permit().await {
Ok(permit) => permit,
Err(error) => {
return websocket_admission_error_response(
&trace_id,
&decision,
Some(uri.path()),
error,
)
}
};
let context = WebSocketRequestContext {
trace_id,
headers,
uri,
remote_addr,
decision,
rpm_bypassed: ip_whitelisted,
websocket_connection_permit,
};
Ok(ws
.max_frame_size(limits.max_frame_size)
.max_message_size(limits.max_message_size)
.on_upgrade(move |socket| async move {
drop(request_permit);
run_session(socket, state, context).await;
}))
}
fn websocket_admission_error_response(
trace_id: &str,
decision: &GatewayControlDecision,
request_path: Option<&str>,
error: RequestAdmissionError,
) -> Result<Response<Body>, GatewayError> {
match error {
RequestAdmissionError::Local(aether_runtime::ConcurrencyError::Saturated {
gate,
limit,
})
| RequestAdmissionError::Distributed(
aether_runtime_state::RuntimeSemaphoreError::Saturated { gate, limit },
)
| RequestAdmissionError::Distributed(
aether_runtime_state::RuntimeSemaphoreError::Unavailable { gate, limit, .. },
) => build_local_overloaded_response(trace_id, Some(decision), request_path, gate, limit),
RequestAdmissionError::Local(aether_runtime::ConcurrencyError::Closed { gate }) => Err(
GatewayError::Internal(format!("gateway concurrency gate {gate} is closed")),
),
RequestAdmissionError::Distributed(
aether_runtime_state::RuntimeSemaphoreError::InvalidConfiguration(message),
) => Err(GatewayError::Internal(message)),
}
}
/// Connection-level access log fields which are independent of a protocol's
/// per-turn usage lifecycle.
#[derive(Clone, Copy)]
pub(crate) struct WebSocketConnectionLogSpec {
pub(crate) opened_event_name: &'static str,
pub(crate) closed_event_name: &'static str,
pub(crate) opened_message: &'static str,
pub(crate) closed_message: &'static str,
pub(crate) execution_path: &'static str,
pub(crate) provider_type: &'static str,
}
pub(crate) struct WebSocketConnectionLog {
spec: WebSocketConnectionLogSpec,
trace_id: String,
remote_addr: SocketAddr,
path: String,
route_class: String,
user_id: String,
api_key_id: String,
started_at: std::time::Instant,
}
impl WebSocketConnectionLog {
pub(crate) fn new(context: &WebSocketRequestContext, spec: WebSocketConnectionLogSpec) -> Self {
let auth_context = context.decision.auth_context.as_ref();
Self {
spec,
trace_id: context.trace_id.clone(),
remote_addr: context.remote_addr,
path: context.uri.path().to_string(),
route_class: context
.decision
.route_class
.as_deref()
.unwrap_or("ai_public")
.to_string(),
user_id: auth_context
.map(|auth_context| auth_context.user_id.clone())
.unwrap_or_else(|| "-".to_string()),
api_key_id: auth_context
.map(|auth_context| auth_context.api_key_id.clone())
.unwrap_or_else(|| "-".to_string()),
started_at: std::time::Instant::now(),
}
}
pub(crate) fn log_opened(&self) {
info!(
event_name = self.spec.opened_event_name,
log_type = "access",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
status = "upgraded",
status_code = 101u16,
trace_id = %self.trace_id,
remote_addr = %self.remote_addr,
method = "GET",
path = %self.path,
user_id = %self.user_id,
api_key_id = %self.api_key_id,
route_class = %self.route_class,
execution_path = self.spec.execution_path,
provider_type = self.spec.provider_type,
message = self.spec.opened_message,
);
}
}
impl Drop for WebSocketConnectionLog {
fn drop(&mut self) {
info!(
event_name = self.spec.closed_event_name,
log_type = "access",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
status = "closed",
status_code = 101u16,
trace_id = %self.trace_id,
remote_addr = %self.remote_addr,
method = "GET",
path = %self.path,
user_id = %self.user_id,
api_key_id = %self.api_key_id,
route_class = %self.route_class,
execution_path = self.spec.execution_path,
provider_type = self.spec.provider_type,
elapsed_ms = self.started_at.elapsed().as_millis() as u64,
message = self.spec.closed_message,
);
}
}
@@ -0,0 +1,12 @@
//! Shared infrastructure for public AI WebSocket bridges.
//!
//! Protocol adapters live below [`responses`]. This layer deliberately owns
//! only transport concerns that are common to future adapters: authenticated
//! upgrade admission, connection limits, upstream handshakes, and frame
//! conversion. It does not interpret provider events or make routing
//! decisions.
pub(crate) mod ingress;
pub(crate) mod responses;
pub(crate) mod session;
pub(crate) mod transport;
@@ -0,0 +1,186 @@
//! Provider-specific hooks for the standard Responses WebSocket session.
use async_trait::async_trait;
use serde_json::Value;
use super::adapters::CODEX_RESPONSES_WEBSOCKET_ADAPTER;
use crate::ai_serving::AiExecutionDecision;
use crate::handlers::proxy::websocket::transport::UpstreamWebSocketErrorCodes;
use crate::orchestration::ResponsesWebSocketAdapter;
use crate::AppState;
#[derive(Debug, Clone, Copy)]
pub(super) struct ResponsesWebSocketDrainDirective {
pub(super) error_code: &'static str,
/// The terminal upstream event may be replayed only when the session has
/// not exposed any standard Responses event to the client.
pub(super) retry_current_turn: bool,
/// When present, the exhausted provider key remains excluded from later
/// turns on this client socket until the upstream's reported reset time.
pub(super) retry_exclusion_until_unix_secs: Option<u64>,
}
/// Provider-specific observation produced while relaying an upstream frame.
/// The session can make the retry/drain decision synchronously, while the
/// optional persistence sink runs outside the frame-forwarding path.
#[derive(Debug, Clone)]
pub(super) struct ResponsesWebSocketAdapterObservation {
pub(super) drain: Option<ResponsesWebSocketDrainDirective>,
pub(super) quota_metadata: Option<Value>,
}
/// Provider identity used by the shared session's temporary exclusion table.
/// The session does not need to know how a provider derives its account id.
#[derive(Debug, Clone, Default)]
pub(super) struct ResponsesWebSocketExclusionIdentity {
pub(super) account_id: Option<String>,
}
/// Whether receiving an upstream event still leaves the active client turn
/// safe to replay on a freshly bound upstream. The shared session keeps the
/// conservative default; provider adapters may explicitly whitelist their
/// documented, pre-response advisory events.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum ResponsesWebSocketRebindSafety {
Safe,
Unsafe { reason: &'static str },
}
/// Boundary between the standard Responses protocol engine and provider
/// behavior. Adapters receive already-planned provider requests; they never
/// own public WebSocket parsing, turn accounting, or model scheduling.
#[async_trait]
pub(super) trait ResponsesWebSocketProtocolAdapter: Send + Sync {
fn kind(&self) -> ResponsesWebSocketAdapter;
fn upstream_errors(&self) -> UpstreamWebSocketErrorCodes;
/// Adds provider-specific metadata to an otherwise standard Responses
/// stream report. The event payload is never rewritten for the client.
fn decorate_turn_report_context(&self, report_context: &mut Option<Value>, event: &Value);
/// Whether this adapter needs the shared session to parse each upstream
/// text event before normal turn accounting runs.
fn observes_upstream_events(&self) -> bool;
/// Classifies whether a received upstream event can be followed by a
/// transparent quota-driven rebind. An adapter must return `Safe` only
/// for events that neither create public Responses state nor make a replay
/// observably ambiguous to the client.
fn rebind_safety_for_upstream_event(&self, event: &Value) -> ResponsesWebSocketRebindSafety;
/// Lets an adapter classify provider-only events. Returning a directive
/// asks the shared session to drain after the active standard response.
fn observe_upstream_event(&self, event: &Value)
-> Option<ResponsesWebSocketAdapterObservation>;
fn exhaustion_exclusion_identity(
&self,
_decision: &AiExecutionDecision,
) -> Option<ResponsesWebSocketExclusionIdentity> {
None
}
/// Persists an adapter observation outside the frame-forwarding path.
async fn persist_upstream_observation(
&self,
state: &AppState,
trace_id: &str,
report_context: Option<&Value>,
observation: ResponsesWebSocketAdapterObservation,
);
}
pub(super) fn resolve_responses_websocket_adapter(
kind: ResponsesWebSocketAdapter,
) -> &'static dyn ResponsesWebSocketProtocolAdapter {
match kind {
ResponsesWebSocketAdapter::Standard => &STANDARD_RESPONSES_WEBSOCKET_ADAPTER,
ResponsesWebSocketAdapter::Codex => &CODEX_RESPONSES_WEBSOCKET_ADAPTER,
}
}
struct StandardResponsesWebSocketAdapter;
const STANDARD_UPSTREAM_WEBSOCKET_ERRORS: UpstreamWebSocketErrorCodes =
UpstreamWebSocketErrorCodes {
upstream_url_missing: "responses_upstream_url_missing",
upstream_url_invalid: "responses_upstream_url_invalid",
headers_invalid: "responses_websocket_headers_invalid",
client_build_failed: "responses_websocket_client_build_failed",
proxy_invalid: "responses_websocket_proxy_invalid",
tunnel_proxy_unsupported: "responses_websocket_tunnel_proxy_unsupported",
handshake_failed: "responses_websocket_handshake_failed",
upgrade_rejected: "responses_websocket_upgrade_rejected",
upgrade_failed: "responses_websocket_upgrade_failed",
};
static STANDARD_RESPONSES_WEBSOCKET_ADAPTER: StandardResponsesWebSocketAdapter =
StandardResponsesWebSocketAdapter;
#[async_trait]
impl ResponsesWebSocketProtocolAdapter for StandardResponsesWebSocketAdapter {
fn kind(&self) -> ResponsesWebSocketAdapter {
ResponsesWebSocketAdapter::Standard
}
fn upstream_errors(&self) -> UpstreamWebSocketErrorCodes {
STANDARD_UPSTREAM_WEBSOCKET_ERRORS
}
fn decorate_turn_report_context(&self, _report_context: &mut Option<Value>, _event: &Value) {}
fn observes_upstream_events(&self) -> bool {
false
}
fn rebind_safety_for_upstream_event(&self, event: &Value) -> ResponsesWebSocketRebindSafety {
let reason = if is_standard_responses_event(event) {
"standard_response_event"
} else {
"unrecognized_upstream_event"
};
ResponsesWebSocketRebindSafety::Unsafe { reason }
}
fn observe_upstream_event(
&self,
_event: &Value,
) -> Option<ResponsesWebSocketAdapterObservation> {
None
}
async fn persist_upstream_observation(
&self,
_state: &AppState,
_trace_id: &str,
_report_context: Option<&Value>,
_observation: ResponsesWebSocketAdapterObservation,
) {
}
}
pub(super) fn is_standard_responses_event(event: &Value) -> bool {
event
.get("type")
.and_then(Value::as_str)
.is_some_and(|event_type| event_type.starts_with("response."))
}
#[cfg(test)]
mod tests {
use super::{resolve_responses_websocket_adapter, ResponsesWebSocketProtocolAdapter};
use crate::orchestration::ResponsesWebSocketAdapter;
#[test]
fn standard_adapter_has_no_codex_extensions() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
assert_eq!(adapter.kind(), ResponsesWebSocketAdapter::Standard);
assert!(!adapter.observes_upstream_events());
assert_eq!(
adapter.upstream_errors().handshake_failed,
"responses_websocket_handshake_failed"
);
}
}
@@ -0,0 +1,290 @@
//! Codex-specific extensions for the standard Responses WebSocket session.
use async_trait::async_trait;
use serde_json::{Map, Value};
use super::super::adapter::{
is_standard_responses_event, ResponsesWebSocketAdapterObservation,
ResponsesWebSocketDrainDirective, ResponsesWebSocketExclusionIdentity,
ResponsesWebSocketProtocolAdapter, ResponsesWebSocketRebindSafety,
};
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,
};
use crate::AppState;
const CODEX_WEBSOCKET_LOG_TARGET: &str = "aether_gateway::handlers::proxy::codex_ws";
const CODEX_WEBSOCKET_RATE_LIMITS_REPORT_CONTEXT_FIELD: &str = "codex_websocket_rate_limits";
const CODEX_UPSTREAM_WEBSOCKET_ERRORS: UpstreamWebSocketErrorCodes = UpstreamWebSocketErrorCodes {
upstream_url_missing: "codex_upstream_url_missing",
upstream_url_invalid: "codex_upstream_url_invalid",
headers_invalid: "codex_websocket_headers_invalid",
client_build_failed: "codex_websocket_client_build_failed",
proxy_invalid: "codex_websocket_proxy_invalid",
tunnel_proxy_unsupported: "codex_websocket_tunnel_proxy_unsupported",
handshake_failed: "codex_websocket_handshake_failed",
upgrade_rejected: "codex_websocket_upgrade_rejected",
upgrade_failed: "codex_websocket_upgrade_failed",
};
pub(crate) static CODEX_RESPONSES_WEBSOCKET_ADAPTER: CodexResponsesWebSocketAdapter =
CodexResponsesWebSocketAdapter;
pub(crate) struct CodexResponsesWebSocketAdapter;
#[async_trait]
impl ResponsesWebSocketProtocolAdapter for CodexResponsesWebSocketAdapter {
fn kind(&self) -> ResponsesWebSocketAdapter {
ResponsesWebSocketAdapter::Codex
}
fn upstream_errors(&self) -> UpstreamWebSocketErrorCodes {
CODEX_UPSTREAM_WEBSOCKET_ERRORS
}
fn decorate_turn_report_context(&self, report_context: &mut Option<Value>, event: &Value) {
let Some(rate_limits) = parse_codex_rate_limits(event) else {
return;
};
let context = report_context.get_or_insert_with(|| Value::Object(Map::new()));
let Some(context) = context.as_object_mut() else {
return;
};
context.insert(
CODEX_WEBSOCKET_RATE_LIMITS_REPORT_CONTEXT_FIELD.to_string(),
rate_limits,
);
}
fn observes_upstream_events(&self) -> bool {
true
}
fn rebind_safety_for_upstream_event(&self, event: &Value) -> ResponsesWebSocketRebindSafety {
if let Some(chunks) = event.get("chunks").and_then(Value::as_array) {
if chunks.is_empty() {
return ResponsesWebSocketRebindSafety::Unsafe {
reason: "unrecognized_upstream_event",
};
}
return chunks
.iter()
.map(codex_direct_rebind_safety)
.find(|safety| matches!(safety, ResponsesWebSocketRebindSafety::Unsafe { .. }))
.unwrap_or(ResponsesWebSocketRebindSafety::Safe);
}
codex_direct_rebind_safety(event)
}
fn observe_upstream_event(
&self,
event: &Value,
) -> Option<ResponsesWebSocketAdapterObservation> {
let rate_limits = parse_codex_rate_limits(event)?;
let 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());
Some(ResponsesWebSocketAdapterObservation {
drain: exhausted.then_some(ResponsesWebSocketDrainDirective {
error_code: "codex_account_quota_exhausted",
retry_current_turn: true,
retry_exclusion_until_unix_secs,
}),
quota_metadata: Some(rate_limits),
})
}
fn exhaustion_exclusion_identity(
&self,
decision: &AiExecutionDecision,
) -> Option<ResponsesWebSocketExclusionIdentity> {
Some(ResponsesWebSocketExclusionIdentity {
account_id: codex_account_id_from_headers(&decision.provider_request_headers)
.map(str::to_string),
})
}
async fn persist_upstream_observation(
&self,
state: &AppState,
trace_id: &str,
report_context: Option<&Value>,
observation: ResponsesWebSocketAdapterObservation,
) {
let Some(rate_limits) = observation.quota_metadata else {
return;
};
if let Err(error) =
sync_codex_websocket_quota_metadata(state, report_context, rate_limits).await
{
tracing::warn!(
target: CODEX_WEBSOCKET_LOG_TARGET,
event_name = "codex_websocket_quota_sync_failed",
log_type = "ops",
transport = "websocket",
websocket = true,
trace_id = %trace_id,
error = ?error,
"gateway failed to persist Codex WebSocket quota metadata"
);
}
}
}
fn codex_direct_rebind_safety(event: &Value) -> ResponsesWebSocketRebindSafety {
let event_type = event
.get("type")
.and_then(Value::as_str)
.unwrap_or_default();
if matches!(event_type, "codex.rate_limits" | "codex.response.metadata") {
// Codex emits these as pre-response advisory metadata. They do
// not create a public `response.*` object, so a replacement
// upstream can safely emit its own current snapshot.
return ResponsesWebSocketRebindSafety::Safe;
}
if event_type == "error" && parse_codex_rate_limits(event).is_some() {
// The quota error is withheld from the client when the shared
// session successfully rebinds, therefore it remains replay-safe.
return ResponsesWebSocketRebindSafety::Safe;
}
let reason = if is_standard_responses_event(event) {
"standard_response_event"
} else {
"unrecognized_upstream_event"
};
ResponsesWebSocketRebindSafety::Unsafe { reason }
}
fn parse_codex_rate_limits(event: &Value) -> Option<Value> {
aether_admin::provider::quota::parse_codex_websocket_rate_limits_response(
event,
current_unix_secs(),
)
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::{
CodexResponsesWebSocketAdapter, ResponsesWebSocketProtocolAdapter,
ResponsesWebSocketRebindSafety,
};
#[test]
fn codex_rate_limit_chunk_is_kept_for_the_terminal_report() {
let adapter = CodexResponsesWebSocketAdapter;
assert!(adapter.observes_upstream_events());
let mut context = Some(json!({"key_id": "codex-key"}));
adapter.decorate_turn_report_context(
&mut context,
&json!({
"chunks": [{
"type": "codex.rate_limits",
"plan_type": "free",
"rate_limits": {
"allowed": true,
"limit_reached": false,
"primary": {
"used_percent": 91,
"window_minutes": 43200,
"reset_after_seconds": 2590791
}
}
}]
}),
);
assert_eq!(
context.as_ref().and_then(
|context| context.pointer("/codex_websocket_rate_limits/primary_used_percent")
),
Some(&json!(91.0))
);
}
#[test]
fn usage_limit_error_is_kept_for_the_terminal_report() {
let adapter = CodexResponsesWebSocketAdapter;
let mut context = Some(json!({"key_id": "codex-key"}));
adapter.decorate_turn_report_context(
&mut context,
&json!({
"type": "error",
"error": {
"type": "usage_limit_reached",
"plan_type": "free",
"resets_at": 1_787_274_385u64,
},
"status_code": 429,
"headers": {
"X-Codex-Primary-Used-Percent": "100",
"X-Codex-Primary-Reset-At": "1787274385",
},
}),
);
assert_eq!(
context
.as_ref()
.and_then(|context| context.pointer("/codex_websocket_rate_limits/allowed")),
Some(&json!(false))
);
assert_eq!(
context.as_ref().and_then(|context| {
context.pointer("/codex_websocket_rate_limits/primary_used_percent")
}),
Some(&json!(100.0))
);
}
#[test]
fn only_known_codex_pre_response_metadata_is_safe_to_rebind() {
let adapter = CodexResponsesWebSocketAdapter;
assert_eq!(
adapter.rebind_safety_for_upstream_event(&json!({
"type": "codex.rate_limits",
"rate_limits": {"allowed": true}
})),
ResponsesWebSocketRebindSafety::Safe
);
assert_eq!(
adapter.rebind_safety_for_upstream_event(&json!({
"type": "codex.response.metadata"
})),
ResponsesWebSocketRebindSafety::Safe
);
assert_eq!(
adapter.rebind_safety_for_upstream_event(&json!({
"chunks": [
{"type": "codex.rate_limits", "rate_limits": {"allowed": true}},
{"type": "codex.response.metadata"}
]
})),
ResponsesWebSocketRebindSafety::Safe
);
assert_eq!(
adapter.rebind_safety_for_upstream_event(&json!({
"type": "response.created"
})),
ResponsesWebSocketRebindSafety::Unsafe {
reason: "standard_response_event"
}
);
assert_eq!(
adapter.rebind_safety_for_upstream_event(&json!({
"type": "codex.unknown"
})),
ResponsesWebSocketRebindSafety::Unsafe {
reason: "unrecognized_upstream_event"
}
);
}
}
@@ -0,0 +1,5 @@
//! Provider-specific Responses WebSocket adapters.
mod codex;
pub(super) use codex::CODEX_RESPONSES_WEBSOCKET_ADAPTER;
@@ -0,0 +1,80 @@
//! Per-turn resource admission for the Responses WebSocket bridge.
//!
//! A WebSocket connection may live for a long time, but each `response.create`
//! is still one active upstream execution. Keep the resource leases attached
//! to the turn instead of the socket so idle connections do not consume
//! upstream capacity.
use std::time::Instant;
use aether_contracts::ExecutionPlan;
use crate::execution_runtime::acquire_upstream_execution_gate;
use crate::provider_pool_demand::{
acquire_provider_pool_in_flight_guard, ProviderPoolInFlightGuard,
};
use crate::upstream_admission::UpstreamTargetAdmissionPermit;
use crate::{AppState, GatewayError};
pub(super) struct ResponsesWebSocketTurnAdmission {
upstream_execution: Option<aether_runtime::ConcurrencyPermit>,
upstream_target: Option<UpstreamTargetAdmissionPermit>,
provider_pool: Option<ProviderPoolInFlightGuard>,
acquired_at: Instant,
}
impl ResponsesWebSocketTurnAdmission {
pub(super) async fn acquire(
state: &AppState,
plan: &ExecutionPlan,
trace_id: &str,
) -> Result<Self, GatewayError> {
let upstream_execution = acquire_upstream_execution_gate(state, trace_id).await?;
let upstream_target = match state
.upstream_target_admission
.acquire(plan, trace_id)
.await
{
Ok(permit) => permit,
Err(error) => {
drop(upstream_execution);
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;
Ok(Self {
upstream_execution,
upstream_target,
provider_pool,
acquired_at: Instant::now(),
})
}
/// Release the distributed provider token before the turn's persistence
/// work. The remaining permits are local RAII guards and are dropped with
/// this value.
pub(super) async fn release(mut self) {
if let Some(provider_pool) = self.provider_pool.take() {
provider_pool.release().await;
}
drop(self.upstream_target.take());
drop(self.upstream_execution.take());
}
}
impl Drop for ResponsesWebSocketTurnAdmission {
fn drop(&mut self) {
crate::stage_metrics::observe_gateway_stage_ms(
"websocket_turn_admission_held",
self.acquired_at.elapsed().as_millis() as u64,
);
}
}
@@ -0,0 +1,413 @@
//! Identity of the physical upstream connection backing a Responses session.
//!
//! A Responses continuation carries state that lives on one provider socket.
//! Comparing only the selected key is therefore not sufficient: transport
//! settings, stable account headers, and the protocol adapter can all change
//! the connection that would receive the next event. Rotating bearer values
//! are intentionally excluded because they do not change an already-upgraded
//! socket's physical binding.
use std::collections::{BTreeMap, BTreeSet};
use std::fmt;
use aether_contracts::{ProxySnapshot, ResolvedTransportProfile};
use sha2::{Digest, Sha256};
use super::adapter::ResponsesWebSocketProtocolAdapter;
use crate::ai_serving::AiExecutionDecision;
use crate::handlers::proxy::websocket::transport::{
websocket_handshake_headers, websocket_upstream_url,
};
use crate::orchestration::ResponsesWebSocketAdapter;
/// Stable, comparable identity for the actual WebSocket connection target.
///
/// The identity deliberately owns the normalized handshake values rather than
/// retaining a reference to the planner decision. A later re-plan can then
/// be compared without accidentally ignoring a field that changes the
/// physical connection.
#[derive(Clone, PartialEq)]
pub(super) struct UpstreamBindingIdentity {
adapter_kind: ResponsesWebSocketAdapter,
provider_id: Option<String>,
endpoint_id: Option<String>,
key_id: Option<String>,
upstream_url: String,
handshake_headers: BTreeMap<String, String>,
/// Authentication values are not part of a stable key binding when the
/// planner has already supplied a key identity. If that identity is
/// unavailable, retain only a one-way fingerprint so two accounts cannot
/// accidentally share a continuation socket.
auth_fingerprint: Option<[u8; 32]>,
proxy: Option<ProxySnapshot>,
transport_profile: Option<ResolvedTransportProfile>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum UpstreamBindingIdentityError {
MissingUpstreamUrl,
InvalidUpstreamUrl,
InvalidHandshakeHeaders,
}
impl UpstreamBindingIdentity {
/// Builds an identity from the same normalized URL and headers used by
/// the WebSocket transport client.
pub(super) fn from_decision(
adapter: &'static dyn ResponsesWebSocketProtocolAdapter,
decision: &AiExecutionDecision,
) -> Result<Self, UpstreamBindingIdentityError> {
let raw_url = decision
.upstream_url
.as_deref()
.filter(|value| !value.trim().is_empty())
.ok_or(UpstreamBindingIdentityError::MissingUpstreamUrl)?;
let upstream_url = websocket_upstream_url(raw_url, "invalid")
.map_err(|_| UpstreamBindingIdentityError::InvalidUpstreamUrl)?
.to_string();
let headers = websocket_handshake_headers(&decision.provider_request_headers, "invalid")
.map_err(|_| UpstreamBindingIdentityError::InvalidHandshakeHeaders)?;
let authentication_header_names = authentication_header_names(decision);
let mut handshake_headers = BTreeMap::new();
let mut authentication_headers = BTreeMap::new();
for (name, value) in &headers {
let name = name.as_str().to_ascii_lowercase();
let value = value
.to_str()
.map_err(|_| UpstreamBindingIdentityError::InvalidHandshakeHeaders)?;
if authentication_header_names.contains(name.as_str()) {
authentication_headers.insert(name, value.to_string());
} else {
handshake_headers.insert(name, value.to_string());
}
}
let auth_fingerprint = decision
.key_id
.is_none()
.then(|| fingerprint_headers(&authentication_headers))
.filter(|_| !authentication_headers.is_empty());
Ok(Self {
adapter_kind: adapter.kind(),
provider_id: decision.provider_id.clone(),
endpoint_id: decision.endpoint_id.clone(),
key_id: decision.key_id.clone(),
upstream_url,
handshake_headers,
auth_fingerprint,
proxy: effective_proxy_snapshot(decision.proxy.as_ref()),
transport_profile: decision.transport_profile.clone(),
})
}
}
/// Header names that carry credentials in the provider handshake. The
/// planner's explicit `auth_header` extends this list for provider-specific
/// schemes; unknown headers remain part of the stable handshake identity.
fn authentication_header_names(decision: &AiExecutionDecision) -> BTreeSet<String> {
let mut names = BTreeSet::from([
"authorization".to_string(),
"proxy-authorization".to_string(),
"x-api-key".to_string(),
"api-key".to_string(),
"x-goog-api-key".to_string(),
"x-azure-api-key".to_string(),
]);
if let Some(name) = decision
.auth_header
.as_deref()
.map(str::trim)
.filter(|name| !name.is_empty())
{
names.insert(name.to_ascii_lowercase());
}
names
}
fn fingerprint_headers(headers: &BTreeMap<String, String>) -> [u8; 32] {
let mut hasher = Sha256::new();
for (name, value) in headers {
hasher.update((name.len() as u64).to_be_bytes());
hasher.update(name.as_bytes());
hasher.update((value.len() as u64).to_be_bytes());
hasher.update(value.as_bytes());
}
hasher.finalize().into()
}
/// Normalize only values that are provably direct transport. Keep node/tunnel
/// fields even though the current WebSocket builder rejects those proxies: a
/// re-plan must not accidentally reuse an already-bound direct socket for a
/// decision that selected a different proxy topology.
fn effective_proxy_snapshot(proxy: Option<&ProxySnapshot>) -> Option<ProxySnapshot> {
let proxy = proxy?;
if proxy.enabled == Some(false) {
return None;
}
let mut normalized = proxy.clone();
normalized.url = normalized
.url
.take()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty());
normalized.mode = normalized
.mode
.take()
.map(|value| value.trim().to_ascii_lowercase())
.filter(|value| !value.is_empty());
normalized.node_id = normalized
.node_id
.take()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty());
normalized.label = normalized
.label
.take()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty());
let has_effective_proxy = normalized.url.is_some()
|| normalized.node_id.is_some()
|| normalized.mode.is_some()
|| normalized.extra.is_some();
has_effective_proxy.then_some(normalized)
}
impl fmt::Debug for UpstreamBindingIdentity {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("UpstreamBindingIdentity")
.field("adapter_kind", &self.adapter_kind)
.field("provider_id", &self.provider_id)
.field("endpoint_id", &self.endpoint_id)
.field("key_id", &self.key_id)
.field("upstream_url", &self.upstream_url)
.field(
"handshake_header_names",
&self.handshake_headers.keys().collect::<Vec<_>>(),
)
.field("proxy_configured", &self.proxy.is_some())
.field(
"transport_profile_id",
&self
.transport_profile
.as_ref()
.map(|profile| profile.profile_id.as_str()),
)
.finish()
}
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use serde_json::json;
use super::{UpstreamBindingIdentity, UpstreamBindingIdentityError};
use crate::ai_serving::AiExecutionDecision;
use crate::handlers::proxy::websocket::responses::adapter::resolve_responses_websocket_adapter;
use crate::orchestration::ResponsesWebSocketAdapter;
fn decision() -> AiExecutionDecision {
AiExecutionDecision {
action: "execute".to_string(),
decision_kind: None,
execution_strategy: None,
conversion_mode: None,
request_id: Some("request-1".to_string()),
candidate_id: Some("candidate-1".to_string()),
provider_name: Some("provider".to_string()),
provider_type: Some("openai".to_string()),
provider_id: Some("provider-1".to_string()),
endpoint_id: Some("endpoint-1".to_string()),
key_id: Some("key-1".to_string()),
upstream_base_url: Some("https://api.example.test".to_string()),
upstream_url: Some("https://api.example.test/v1/responses".to_string()),
provider_request_method: Some("POST".to_string()),
auth_header: Some("authorization".to_string()),
auth_value: Some("Bearer secret".to_string()),
provider_api_format: Some("openai:responses".to_string()),
client_api_format: Some("openai:responses".to_string()),
provider_contract: None,
client_contract: None,
model_name: Some("gpt-5.6-sol".to_string()),
mapped_model: None,
prompt_cache_key: None,
extra_headers: BTreeMap::new(),
provider_request_headers: BTreeMap::from([
("Authorization".to_string(), "Bearer secret".to_string()),
("X-Client".to_string(), "aether".to_string()),
("Connection".to_string(), "keep-alive".to_string()),
]),
provider_request_body: Some(json!({"model": "gpt-5.6-sol"})),
provider_request_body_base64: None,
content_type: Some("application/json".to_string()),
content_encoding: None,
request_gzip: None,
proxy: None,
transport_profile: None,
timeouts: None,
upstream_is_stream: true,
report_kind: None,
report_context: None,
auth_context: None,
}
}
#[test]
fn identity_normalizes_url_and_hop_by_hop_headers() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
let identity = UpstreamBindingIdentity::from_decision(adapter, &decision()).unwrap();
assert_eq!(identity.upstream_url, "wss://api.example.test/v1/responses");
assert_eq!(
identity.handshake_headers,
BTreeMap::from([("x-client".to_string(), "aether".to_string())])
);
}
#[test]
fn identity_changes_when_physical_binding_changes() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
let base = decision();
let identity = UpstreamBindingIdentity::from_decision(adapter, &base).unwrap();
let codex_adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
assert_ne!(
identity,
UpstreamBindingIdentity::from_decision(codex_adapter, &base).unwrap()
);
for mutate in [
|decision: &mut AiExecutionDecision| {
decision.key_id = Some("key-2".to_string());
},
|decision: &mut AiExecutionDecision| {
decision.upstream_url = Some("https://other.example.test/v1/responses".to_string());
},
|decision: &mut AiExecutionDecision| {
decision
.provider_request_headers
.insert("X-Client".to_string(), "other".to_string());
},
|decision: &mut AiExecutionDecision| {
decision.proxy = Some(aether_contracts::ProxySnapshot {
enabled: Some(true),
url: Some("http://proxy.example.test:8080".to_string()),
..Default::default()
});
},
|decision: &mut AiExecutionDecision| {
decision.transport_profile = Some(aether_contracts::ResolvedTransportProfile {
profile_id: "chrome136".to_string(),
..Default::default()
});
},
] {
let mut changed = base.clone();
mutate(&mut changed);
let changed_identity =
UpstreamBindingIdentity::from_decision(adapter, &changed).unwrap();
assert_ne!(identity, changed_identity);
}
let mut rotated = base.clone();
rotated
.provider_request_headers
.insert("Authorization".to_string(), "Bearer rotated".to_string());
assert_eq!(
identity,
UpstreamBindingIdentity::from_decision(adapter, &rotated).unwrap()
);
}
#[test]
fn stable_key_identity_ignores_custom_auth_value_rotation() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
let mut base = decision();
base.auth_header = Some("X-Provider-Token".to_string());
base.provider_request_headers.remove("Authorization");
base.provider_request_headers.insert(
"X-Provider-Token".to_string(),
"provider-token-1".to_string(),
);
let identity = UpstreamBindingIdentity::from_decision(adapter, &base).unwrap();
assert!(!identity.handshake_headers.contains_key("x-provider-token"));
assert!(identity.auth_fingerprint.is_none());
let mut rotated = base;
rotated.provider_request_headers.insert(
"X-Provider-Token".to_string(),
"provider-token-2".to_string(),
);
assert_eq!(
identity,
UpstreamBindingIdentity::from_decision(adapter, &rotated).unwrap()
);
}
#[test]
fn missing_key_identity_fingerprints_authentication_values() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
let mut first = decision();
first.key_id = None;
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
assert!(first_identity.auth_fingerprint.is_some());
let mut same_account_rotation = first.clone();
same_account_rotation.provider_request_headers.insert(
"Authorization".to_string(),
"Bearer different-account-or-token".to_string(),
);
let changed_identity =
UpstreamBindingIdentity::from_decision(adapter, &same_account_rotation).unwrap();
assert_ne!(first_identity, changed_identity);
let mut non_auth_change = first;
non_auth_change
.provider_request_headers
.insert("X-Client".to_string(), "other-client".to_string());
assert_ne!(
first_identity,
UpstreamBindingIdentity::from_decision(adapter, &non_auth_change).unwrap()
);
}
#[test]
fn disabled_proxy_is_equivalent_to_direct_transport() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
let direct = decision();
let direct_identity = UpstreamBindingIdentity::from_decision(adapter, &direct).unwrap();
let mut explicitly_disabled = direct;
explicitly_disabled.proxy = Some(aether_contracts::ProxySnapshot {
enabled: Some(false),
url: Some("http://ignored.example.test:8080".to_string()),
..Default::default()
});
assert_eq!(
direct_identity,
UpstreamBindingIdentity::from_decision(adapter, &explicitly_disabled).unwrap()
);
}
#[test]
fn identity_rejects_missing_or_invalid_connection_fields() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
let mut missing = decision();
missing.upstream_url = None;
assert_eq!(
UpstreamBindingIdentity::from_decision(adapter, &missing),
Err(UpstreamBindingIdentityError::MissingUpstreamUrl)
);
let mut invalid = decision();
invalid.upstream_url = Some("file:///tmp/responses".to_string());
assert_eq!(
UpstreamBindingIdentity::from_decision(adapter, &invalid),
Err(UpstreamBindingIdentityError::InvalidUpstreamUrl)
);
}
}
@@ -0,0 +1,770 @@
//! Client-side Responses WebSocket event forwarding and follow-up planning.
use axum::body::Bytes;
use axum::extract::ws::{Message as AxumWsMessage, WebSocket};
use futures_util::SinkExt;
use serde_json::Value;
use uuid::Uuid;
use wreq::ws::message::Message as WreqWsMessage;
use super::adapter::{resolve_responses_websocket_adapter, ResponsesWebSocketDrainDirective};
use super::lifecycle::{
await_pending_turn_finalization, queue_turn_finalization,
send_responses_websocket_turn_start_error, ActiveResponsesWebSocketTurn,
};
use super::quota::{mark_active_response_retry_unsafe, send_previous_response_not_found};
use super::request::{
build_planning_parts, changed_followup_response_create_model,
continuation_requires_same_upstream, normalize_followup_response_create,
planned_response_create_event, provider_model_from_decision,
response_create_has_previous_response_id, response_create_model_or_current,
};
use super::state::{ActiveResponsesWebSocketRequest, BoundResponsesConnection};
use super::turn::{
begin_responses_websocket_turn, prepare_responses_websocket_turn_decision,
ResponsesWebSocketTurnObservation, ResponsesWebSocketTurnOutcome,
};
use super::upstream::{bind_responses_upstream, decision_reuses_bound_upstream};
use crate::ai_serving::maybe_build_responses_websocket_decision;
use crate::clock::current_unix_secs;
use crate::control::{request_model_local_rejection, GatewayControlDecision};
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
use crate::handlers::proxy::websocket::session::{CLOSE_INTERNAL_ERROR, WEBSOCKET_LOG_TRANSPORT};
use crate::handlers::proxy::websocket::transport::{
client_close_to_upstream, close_client_socket, close_upstream_socket, send_client_message,
send_gateway_error, send_gateway_error_with_status, send_upstream_message,
};
use crate::orchestration::release_pool_key_lease_from_report_context;
use crate::rate_limit::FrontdoorUserRpmOutcome;
use crate::AppState;
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
macro_rules! debug {
($($arg:tt)*) => {
tracing::debug!(target: LOG_TARGET, $($arg)*)
};
}
macro_rules! warn {
($($arg:tt)*) => {
tracing::warn!(target: LOG_TARGET, $($arg)*)
};
}
pub(super) enum RelayDisposition {
Continue,
Close,
UpstreamError(&'static str),
}
pub(super) fn adapter_drain_ready(
pending_adapter_drain: Option<ResponsesWebSocketDrainDirective>,
response_in_flight: bool,
observation: Option<ResponsesWebSocketTurnObservation>,
upstream_closed: bool,
) -> bool {
pending_adapter_drain.is_some()
&& (upstream_closed
|| !response_in_flight
|| matches!(
observation,
Some(ResponsesWebSocketTurnObservation::Terminal(_))
))
}
pub(super) async fn forward_client_message(
client_message: AxumWsMessage,
bound: &mut BoundResponsesConnection,
client_socket: &mut WebSocket,
state: &AppState,
context: &WebSocketRequestContext,
) -> RelayDisposition {
match client_message {
AxumWsMessage::Text(text) => {
let text = text.to_string();
let client_event = serde_json::from_str::<Value>(&text).ok();
let is_response_create = client_event
.as_ref()
.and_then(|event| event.get("type"))
.and_then(Value::as_str)
== Some("response.create");
if !is_response_create {
if bound.upstream.is_none() {
send_gateway_error(
client_socket,
"responses_websocket_upstream_rebind_required",
"Send a new response.create to select another Provider connection",
)
.await;
return RelayDisposition::Continue;
}
// We cannot reconstruct arbitrary Responses control events on
// a replacement socket. A concurrent quota error must be
// surfaced rather than replaying only the response.create.
mark_active_response_retry_unsafe(bound, "client_control_event");
return send_upstream_message(
bound
.upstream
.as_mut()
.expect("upstream presence was checked above"),
WreqWsMessage::text(text),
)
.await
.map(|()| RelayDisposition::Continue)
.unwrap_or(RelayDisposition::UpstreamError(
"responses_websocket_send_failed",
));
}
if bound.response_in_flight {
send_gateway_error(
client_socket,
"response_already_in_progress",
"This connection runs one response at a time",
)
.await;
return RelayDisposition::Continue;
}
// A prior terminal turn may still be writing usage/audit and
// projecting provider effects. Do not let a new independent turn
// plan against stale health, adaptive, or pool state.
await_pending_turn_finalization(bound).await;
match consume_response_create_rate_limit(state, &context.decision, context.rpm_bypassed)
.await
{
Ok(true) => {}
Ok(false) => {
send_gateway_error_with_status(
client_socket,
429,
"rate_limit_exceeded",
"Too many response.create events; retry later",
)
.await;
return RelayDisposition::Continue;
}
Err(()) => {
send_gateway_error_with_status(
client_socket,
503,
"gateway_rate_limit_unavailable",
"Gateway could not evaluate the response rate limit",
)
.await;
close_client_socket(
client_socket,
CLOSE_INTERNAL_ERROR,
"rate_limit_unavailable",
)
.await;
return RelayDisposition::Close;
}
}
let Some(client_event) = client_event else {
send_gateway_error(
client_socket,
"invalid_response_create",
"response.create must be valid JSON",
)
.await;
return RelayDisposition::Continue;
};
if bound.upstream.is_none() {
if response_create_has_previous_response_id(&client_event) {
send_previous_response_not_found(client_socket).await;
return RelayDisposition::Continue;
}
let mut client_event = client_event;
let requested_model = match response_create_model_or_current(
&mut client_event,
&bound.client_model,
) {
Ok(model) => model,
Err(code) => {
send_gateway_error(
client_socket,
code,
"response.create.model must be a non-empty string",
)
.await;
return RelayDisposition::Continue;
}
};
return forward_replanned_response_create(
bound,
client_socket,
state,
context,
client_event,
requested_model,
)
.await;
}
let changed_model =
match changed_followup_response_create_model(&client_event, &bound.client_model) {
Ok(model) => model,
Err(code) => {
send_gateway_error(
client_socket,
code,
"response.create.model must be a non-empty string",
)
.await;
return RelayDisposition::Continue;
}
};
if let Some(requested_model) = changed_model {
return forward_replanned_response_create(
bound,
client_socket,
state,
context,
client_event,
requested_model,
)
.await;
}
if !response_create_has_previous_response_id(&client_event) {
return forward_replanned_response_create(
bound,
client_socket,
state,
context,
client_event,
bound.client_model.clone(),
)
.await;
}
let outbound = match normalize_followup_response_create(
&client_event,
&bound.provider_model,
&bound.body_normalization,
) {
Ok(value) => value,
Err(code) => {
send_gateway_error(
client_socket,
code,
"Gateway could not prepare the response.create event",
)
.await;
return RelayDisposition::Continue;
}
};
let provider_event = match serde_json::from_str::<Value>(&outbound) {
Ok(event) => event,
Err(_) => {
send_gateway_error(
client_socket,
"response_create_serialization_failed",
"Gateway could not prepare the response.create event",
)
.await;
return RelayDisposition::Continue;
}
};
let turn_index = bound.next_turn_index;
let turn_request_id = Uuid::new_v4().to_string();
let logical_turn_id = Uuid::new_v4().to_string();
debug!(
event_name = "responses_websocket_response_create_forwarding",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
turn_index,
client_model = %bound.client_model,
provider_model = %bound.provider_model,
model_replanned = false,
has_previous_response_id = response_create_has_previous_response_id(&client_event),
"gateway is forwarding a Responses response.create"
);
let turn_decision = prepare_responses_websocket_turn_decision(
&bound.decision_template,
turn_request_id,
false,
&client_event,
&provider_event,
&context.trace_id,
turn_index,
&logical_turn_id,
1,
);
let planning_parts = build_planning_parts(context);
let mut turn = match begin_responses_websocket_turn(
state,
&planning_parts,
&context.decision,
turn_decision,
&client_event,
)
.await
{
Ok(turn) => turn,
Err(error) => {
warn!(
event_name = "responses_websocket_followup_turn_lifecycle_start_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error = ?error,
"gateway could not start Responses WebSocket follow-up usage/audit lifecycle"
);
send_responses_websocket_turn_start_error(client_socket, &error).await;
return RelayDisposition::Continue;
}
};
turn.set_provider_response_headers(bound.upstream_response_headers.clone());
bound.active_turn = Some(ActiveResponsesWebSocketTurn::new(state, turn));
bound.active_response_create = Some(ActiveResponsesWebSocketRequest::new(
client_event.clone(),
turn_index,
logical_turn_id,
));
bound.next_turn_index = bound.next_turn_index.saturating_add(1);
bound.response_in_flight = true;
let Some(upstream) = bound.upstream.as_mut() else {
return RelayDisposition::UpstreamError("responses_websocket_send_failed");
};
match send_upstream_message(upstream, WreqWsMessage::text(outbound)).await {
Ok(()) => {
if let Some(turn) = bound.active_turn.as_mut() {
turn.mark_upstream_request_sent();
}
RelayDisposition::Continue
}
Err(_) => RelayDisposition::UpstreamError("responses_websocket_send_failed"),
}
}
AxumWsMessage::Binary(data) => {
if bound.upstream.is_some() {
mark_active_response_retry_unsafe(bound, "client_binary_frame");
send_upstream_message(
bound
.upstream
.as_mut()
.expect("upstream presence was checked above"),
WreqWsMessage::Binary(data),
)
.await
.map(|()| RelayDisposition::Continue)
.unwrap_or(RelayDisposition::UpstreamError(
"responses_websocket_send_failed",
))
} else {
send_gateway_error(
client_socket,
"responses_websocket_upstream_rebind_required",
"Send a new response.create to select another Provider connection",
)
.await;
RelayDisposition::Continue
}
}
AxumWsMessage::Ping(data) => match bound.upstream.as_mut() {
Some(upstream) => send_upstream_message(upstream, WreqWsMessage::Ping(data))
.await
.map(|()| RelayDisposition::Continue)
.unwrap_or(RelayDisposition::UpstreamError(
"responses_websocket_send_failed",
)),
None => send_client_message(client_socket, AxumWsMessage::Pong(data))
.await
.map(|()| RelayDisposition::Continue)
.unwrap_or(RelayDisposition::Close),
},
AxumWsMessage::Pong(data) => match bound.upstream.as_mut() {
Some(upstream) => send_upstream_message(upstream, WreqWsMessage::Pong(data))
.await
.map(|()| RelayDisposition::Continue)
.unwrap_or(RelayDisposition::UpstreamError(
"responses_websocket_send_failed",
)),
None => RelayDisposition::Continue,
},
AxumWsMessage::Close(frame) => {
if let Some(upstream) = bound.upstream.as_mut() {
close_upstream_socket(upstream, client_close_to_upstream(frame)).await;
}
RelayDisposition::Close
}
}
}
async fn forward_replanned_response_create(
bound: &mut BoundResponsesConnection,
client_socket: &mut WebSocket,
state: &AppState,
context: &WebSocketRequestContext,
client_event: Value,
requested_model: String,
) -> RelayDisposition {
let planning_parts = build_planning_parts(context);
let client_event_text = match serde_json::to_vec(&client_event) {
Ok(value) => Bytes::from(value),
Err(_) => {
send_gateway_error(
client_socket,
"invalid_response_create",
"response.create must be valid JSON",
)
.await;
return RelayDisposition::Continue;
}
};
match request_model_local_rejection(
state,
Some(&context.decision),
&planning_parts.uri,
&planning_parts.headers,
&client_event_text,
)
.await
{
Ok(Some(_)) => {
send_gateway_error(
client_socket,
"model_not_allowed",
"The requested model is not available to this API key",
)
.await;
return RelayDisposition::Continue;
}
Ok(None) => {}
Err(error) => {
warn!(
event_name = "responses_websocket_followup_model_access_check_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
requested_model = %requested_model,
error = ?error,
"gateway failed to evaluate follow-up WebSocket model access policy"
);
send_gateway_error(
client_socket,
"gateway_auth_unavailable",
"Gateway could not evaluate request access",
)
.await;
close_client_socket(
client_socket,
CLOSE_INTERNAL_ERROR,
"gateway_auth_unavailable",
)
.await;
return RelayDisposition::Close;
}
}
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);
let excluded_key_ids = (!excluded_key_ids.is_empty()).then_some(&excluded_key_ids);
let excluded_codex_account_ids =
(!excluded_codex_account_ids.is_empty()).then_some(&excluded_codex_account_ids);
let planned = match maybe_build_responses_websocket_decision(
state,
&planning_parts,
&turn_request_id,
&context.decision,
&client_event,
excluded_key_ids,
excluded_codex_account_ids,
)
.await
{
Ok(Some(decision)) => decision,
Ok(None) => {
send_gateway_error_with_status(
client_socket,
503,
"responses_provider_unavailable",
"No eligible WebSocket-enabled Responses provider is available for the requested model",
)
.await;
return RelayDisposition::Continue;
}
Err(error) => {
warn!(
event_name = "responses_websocket_followup_model_planning_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
requested_model = %requested_model,
error = ?error,
"gateway failed to re-plan Responses WebSocket follow-up model"
);
send_gateway_error_with_status(
client_socket,
503,
"responses_provider_unavailable",
"Gateway could not prepare the requested model",
)
.await;
return RelayDisposition::Continue;
}
};
let adapter = resolve_responses_websocket_adapter(planned.adapter);
let normalization = planned.normalization;
let decision = planned.execution;
let reuses_bound_upstream = decision_reuses_bound_upstream(bound, adapter, &decision);
if continuation_requires_same_upstream(&client_event, reuses_bound_upstream) {
release_pool_key_lease_from_report_context(state, decision.report_context.as_ref()).await;
debug!(
event_name = "responses_websocket_continuation_rebind_rejected",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
requested_model = %requested_model,
previous_key_id = ?bound.decision_template.key_id,
planned_key_id = ?decision.key_id,
error_code = "previous_response_not_found",
"gateway refused to move a Responses continuation to a different upstream account or connection"
);
send_previous_response_not_found(client_socket).await;
return RelayDisposition::Continue;
}
let provider_event =
match planned_response_create_event(&decision, &client_event).and_then(|event| {
serde_json::from_str::<Value>(&event)
.map_err(|_| "response_create_serialization_failed")
}) {
Ok(event) => event,
Err(code) => {
release_pool_key_lease_from_report_context(state, decision.report_context.as_ref())
.await;
send_gateway_error(
client_socket,
code,
"Gateway could not prepare the requested model",
)
.await;
return RelayDisposition::Continue;
}
};
let turn_index = bound.next_turn_index;
let turn_decision = prepare_responses_websocket_turn_decision(
&decision,
turn_request_id,
true,
&client_event,
&provider_event,
&context.trace_id,
turn_index,
&logical_turn_id,
1,
);
let mut turn = match begin_responses_websocket_turn(
state,
&planning_parts,
&context.decision,
turn_decision,
&client_event,
)
.await
{
Ok(turn) => turn,
Err(error) => {
warn!(
event_name = "responses_websocket_replanned_turn_lifecycle_start_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
requested_model = %requested_model,
error = ?error,
"gateway could not start re-planned WebSocket usage/audit lifecycle"
);
send_responses_websocket_turn_start_error(client_socket, &error).await;
return RelayDisposition::Continue;
}
};
if reuses_bound_upstream {
let outbound = match serde_json::to_string(&provider_event) {
Ok(outbound) => outbound,
Err(_) => {
queue_turn_finalization(
bound,
state,
turn,
ResponsesWebSocketTurnOutcome::upstream_send_failed(),
)
.await;
send_gateway_error(
client_socket,
"response_create_serialization_failed",
"Gateway could not prepare the requested model",
)
.await;
return RelayDisposition::Continue;
}
};
let Some(upstream) = bound.upstream.as_mut() else {
queue_turn_finalization(
bound,
state,
turn,
ResponsesWebSocketTurnOutcome::upstream_send_failed(),
)
.await;
return RelayDisposition::UpstreamError("responses_websocket_send_failed");
};
if send_upstream_message(upstream, WreqWsMessage::text(outbound))
.await
.is_err()
{
queue_turn_finalization(
bound,
state,
turn,
ResponsesWebSocketTurnOutcome::upstream_send_failed(),
)
.await;
return RelayDisposition::UpstreamError("responses_websocket_send_failed");
}
turn.mark_upstream_request_sent();
turn.set_provider_response_headers(bound.upstream_response_headers.clone());
let provider_model =
provider_model_from_decision(&decision).unwrap_or_else(|| bound.provider_model.clone());
let previous_client_model = std::mem::replace(&mut bound.client_model, requested_model);
let previous_provider_model = std::mem::replace(&mut bound.provider_model, provider_model);
bound.decision_template = decision;
// The re-plan keeps this upstream but resolved a new model, so later
// continuations must normalize against the new plan, not the old one.
bound.body_normalization = normalization;
bound.active_turn = Some(ActiveResponsesWebSocketTurn::new(state, turn));
bound.active_response_create = Some(ActiveResponsesWebSocketRequest::new(
client_event.clone(),
turn_index,
logical_turn_id.clone(),
));
bound.next_turn_index = bound.next_turn_index.saturating_add(1);
bound.response_in_flight = true;
debug!(
event_name = "responses_websocket_followup_model_replanned",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
turn_index,
previous_client_model = %previous_client_model,
client_model = %bound.client_model,
previous_provider_model = %previous_provider_model,
provider_model = %bound.provider_model,
upstream_rebound = false,
model_replanned = true,
"gateway re-planned a Responses WebSocket model on the existing upstream"
);
return RelayDisposition::Continue;
}
let mut replacement =
match bind_responses_upstream(&decision, normalization, &client_event, adapter).await {
Ok(connection) => connection,
Err(code) => {
queue_turn_finalization(
bound,
state,
turn,
ResponsesWebSocketTurnOutcome::upstream_connect_failed(code),
)
.await;
warn!(
event_name = "responses_websocket_followup_model_rebind_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
requested_model = %requested_model,
error_code = code,
"gateway failed to rebind Responses WebSocket follow-up model"
);
send_gateway_error_with_status(
client_socket,
502,
code,
"Gateway could not establish the requested model",
)
.await;
return RelayDisposition::Continue;
}
};
turn.mark_upstream_request_sent();
turn.set_provider_response_headers(replacement.upstream_response_headers.clone());
let previous_client_model = bound.client_model.clone();
let previous_provider_model = bound.provider_model.clone();
let replacement_upstream = replacement
.upstream
.take()
.expect("newly bound Responses upstream should be present");
if let Some(mut previous_upstream) = bound.upstream.replace(replacement_upstream) {
close_upstream_socket(&mut previous_upstream, None).await;
}
bound.adapter = replacement.adapter;
bound.client_model = replacement.client_model;
bound.provider_model = replacement.provider_model;
bound.response_in_flight = true;
bound.decision_template = replacement.decision_template;
bound.body_normalization = replacement.body_normalization;
bound.binding_identity = replacement.binding_identity;
bound.active_turn = Some(ActiveResponsesWebSocketTurn::new(state, turn));
bound.active_response_create = Some(ActiveResponsesWebSocketRequest::new(
client_event,
turn_index,
logical_turn_id,
));
bound.next_turn_index = bound.next_turn_index.saturating_add(1);
bound.upstream_response_headers = replacement.upstream_response_headers;
bound.pending_adapter_drain = replacement.pending_adapter_drain;
debug!(
event_name = "responses_websocket_followup_model_rebound",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
turn_index,
previous_client_model = %previous_client_model,
requested_model = %requested_model,
previous_provider_model = %previous_provider_model,
provider_model = %bound.provider_model,
upstream_rebound = true,
model_replanned = true,
"gateway rebound Responses WebSocket for a follow-up model"
);
RelayDisposition::Continue
}
pub(super) async fn consume_response_create_rate_limit(
state: &AppState,
decision: &GatewayControlDecision,
rpm_bypassed: bool,
) -> Result<bool, ()> {
if rpm_bypassed {
return Ok(true);
}
match state
.frontdoor_user_rpm()
.check_and_consume(state, Some(decision))
.await
.map_err(|_| ())?
{
FrontdoorUserRpmOutcome::Rejected(_) => Ok(false),
FrontdoorUserRpmOutcome::Allowed | FrontdoorUserRpmOutcome::NotApplicable => Ok(true),
}
}
@@ -0,0 +1,575 @@
//! Connection-level Responses WebSocket FSM.
use std::time::Duration;
use axum::extract::ws::WebSocket;
use futures_util::{SinkExt, StreamExt};
use serde_json::Value;
use tokio::time::sleep;
use wreq::ws::message::Message as WreqWsMessage;
use super::client::{adapter_drain_ready, forward_client_message, RelayDisposition};
use super::frame::ParsedResponsesWebSocketFrame;
use super::lifecycle::{
await_pending_adapter_observation, finalize_active_turn, queue_turn_finalization,
ActiveResponsesWebSocketTurn,
};
use super::quota::{
active_continuation_can_retry_from_full_input, detach_exhausted_upstream,
is_usage_limit_error_event, mark_active_response_retry_unsafe,
observe_active_response_rebind_safety, retry_active_turn_after_quota_exhaustion,
send_previous_response_not_found, should_request_full_continuation_retry,
};
use super::relay_policy::{
classify_quota_relay, classify_upstream_frame, fatal_relay_policy, FatalRelaySignal,
QuotaRelayAction, QuotaRelayFacts, UpstreamFrameAction, UpstreamFrameKind,
};
use super::state::BoundResponsesConnection;
use super::turn::{
ResponsesWebSocketTurn, ResponsesWebSocketTurnObservation, ResponsesWebSocketTurnOutcome,
};
use super::upstream::{close_bound_upstream, receive_optional_upstream};
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
use crate::handlers::proxy::websocket::session::{
wait_for_optional_deadline, CLOSE_INTERNAL_ERROR, CLOSE_TRY_AGAIN,
RESPONSES_WEBSOCKET_SESSION_LIMITS, WEBSOCKET_LOG_TRANSPORT,
};
use crate::handlers::proxy::websocket::transport::{
close_client_socket, send_client_message, send_gateway_error_with_status,
send_responses_websocket_error, upstream_message_to_client,
};
use crate::AppState;
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
macro_rules! debug {
($($arg:tt)*) => {
tracing::debug!(target: LOG_TARGET, $($arg)*)
};
}
macro_rules! warn {
($($arg:tt)*) => {
tracing::warn!(target: LOG_TARGET, $($arg)*)
};
}
pub(super) async fn relay_bound_connection(
client_socket: &mut WebSocket,
bound: &mut BoundResponsesConnection,
state: &AppState,
context: &WebSocketRequestContext,
connection_permit: Option<aether_runtime::AdmissionPermit>,
) {
let connection_deadline = sleep(RESPONSES_WEBSOCKET_SESSION_LIMITS.max_connection_duration);
tokio::pin!(connection_deadline);
loop {
let active_turn_deadline = bound.active_turn.as_ref().map(|turn| turn.deadline());
tokio::select! {
_ = &mut connection_deadline => {
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::connection_limit_reached(),
).await;
send_gateway_error_with_status(
client_socket,
503,
"websocket_connection_limit_reached",
"WebSocket connection duration limit reached; reconnect to continue",
).await;
bound.active_response_create = None;
close_bound_upstream(bound).await;
close_client_socket(client_socket, CLOSE_TRY_AGAIN, "connection_limit_reached").await;
break;
}
_ = wait_for_optional_deadline(active_turn_deadline.map(|deadline| deadline.deadline)) => {
let Some(turn_deadline) = active_turn_deadline else {
continue;
};
warn!(
event_name = "responses_websocket_turn_timeout",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
timeout_phase = ?turn_deadline.phase,
timeout_ms = turn_deadline.timeout.as_millis() as u64,
"Responses WebSocket response did not reach its configured deadline"
);
finalize_active_turn(bound, state, turn_deadline.phase.outcome()).await;
send_gateway_error_with_status(
client_socket,
504,
turn_deadline.phase.error_code(),
turn_deadline.phase.client_message(),
).await;
bound.active_response_create = None;
close_bound_upstream(bound).await;
close_client_socket(
client_socket,
CLOSE_TRY_AGAIN,
turn_deadline.phase.error_code(),
).await;
break;
}
_ = wait_for_connection_permit_loss(connection_permit.as_ref()) => {
let policy = fatal_relay_policy(FatalRelaySignal::ConnectionAdmissionLost);
warn!(
event_name = "responses_websocket_connection_admission_lost",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
"gateway closed Responses WebSocket after its connection admission became unhealthy"
);
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::connection_admission_lost(),
).await;
bound.active_response_create = None;
close_bound_upstream(bound).await;
send_gateway_error_with_status(
client_socket,
policy.status_code,
policy.error_code,
policy.client_message,
).await;
close_client_socket(client_socket, policy.close_code, policy.close_reason).await;
break;
}
client_message = client_socket.next() => {
let Some(client_message) = client_message else {
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::client_disconnected(),
).await;
bound.active_response_create = None;
close_bound_upstream(bound).await;
break;
};
let Ok(client_message) = client_message else {
warn!(
event_name = "responses_websocket_client_receive_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
"client WebSocket receive failed"
);
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::client_disconnected(),
).await;
bound.active_response_create = None;
close_bound_upstream(bound).await;
break;
};
match forward_client_message(client_message, bound, client_socket, state, context).await {
RelayDisposition::Continue => {}
RelayDisposition::Close => {
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::client_disconnected(),
).await;
break;
}
RelayDisposition::UpstreamError(code) => {
warn!(
event_name = "responses_websocket_upstream_send_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error_code = code,
"Upstream WebSocket send failed"
);
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::upstream_send_failed(),
).await;
send_gateway_error_with_status(
client_socket,
502,
code,
"Gateway could not forward the WebSocket event upstream",
).await;
bound.active_response_create = None;
close_bound_upstream(bound).await;
close_client_socket(client_socket, CLOSE_INTERNAL_ERROR, code).await;
break;
}
}
}
upstream_message = receive_optional_upstream(&mut bound.upstream) => {
let Some(upstream_message) = upstream_message else {
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::upstream_closed(),
).await;
bound.active_response_create = None;
bound.upstream = None;
close_client_socket(client_socket, 1000, "upstream_closed").await;
break;
};
let Ok(upstream_message) = upstream_message else {
warn!(
event_name = "responses_websocket_upstream_receive_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
"Upstream WebSocket receive failed"
);
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::upstream_receive_failed(),
).await;
send_gateway_error_with_status(
client_socket,
502,
"responses_websocket_receive_failed",
"Provider connection closed unexpectedly",
).await;
bound.active_response_create = None;
bound.upstream = None;
close_client_socket(client_socket, CLOSE_INTERNAL_ERROR, "upstream_receive_failed").await;
break;
};
let parsed_upstream_frame = match &upstream_message {
WreqWsMessage::Text(text) => {
ParsedResponsesWebSocketFrame::parse(text.as_str()).ok()
}
_ => None,
};
let parsed_upstream_event = parsed_upstream_frame
.as_ref()
.map(ParsedResponsesWebSocketFrame::event);
if let WreqWsMessage::Text(text) = &upstream_message {
debug!(
event_name = "responses_websocket_upstream_event",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
event_type = %parsed_upstream_frame
.as_ref()
.map(ParsedResponsesWebSocketFrame::event_type_for_log)
.unwrap_or_else(|| "invalid_json".to_string()),
frame_bytes = text.len(),
chunked = parsed_upstream_frame
.as_ref()
.is_some_and(ParsedResponsesWebSocketFrame::is_chunked),
active_turn = bound.active_turn.is_some(),
"gateway received Responses WebSocket event"
);
}
if matches!(&upstream_message, WreqWsMessage::Binary(_)) {
mark_active_response_retry_unsafe(bound, "upstream_binary_frame");
} else if matches!(&upstream_message, WreqWsMessage::Text(_))
&& parsed_upstream_event.is_none()
{
mark_active_response_retry_unsafe(bound, "invalid_upstream_event");
}
if let Some(event) = parsed_upstream_event {
observe_active_response_rebind_safety(bound, event);
if bound.pending_adapter_drain.is_none()
&& bound.adapter.observes_upstream_events()
{
let adapter = bound.adapter;
if let Some(observation) = adapter.observe_upstream_event(event) {
let directive = observation.drain;
await_pending_adapter_observation(bound).await;
let state_for_observation = state.clone();
let trace_id = context.trace_id.clone();
let report_context = bound.decision_template.report_context.clone();
bound.pending_adapter_observation = Some(tokio::spawn(async move {
adapter
.persist_upstream_observation(
&state_for_observation,
&trace_id,
report_context.as_ref(),
observation,
)
.await;
}));
if let Some(directive) = directive {
bound.pending_adapter_drain = Some(directive);
// A definitive quota signal must be visible to
// the next planner before a transparent retry.
await_pending_adapter_observation(bound).await;
}
}
}
}
let observation = match &upstream_message {
WreqWsMessage::Text(text) => {
let adapter = bound.adapter;
match parsed_upstream_frame.as_ref() {
Some(frame) => bound
.active_turn
.as_mut()
.and_then(|turn| turn.observe_upstream_frame(frame, adapter)),
None => {
if let Some(turn) = bound.active_turn.as_mut() {
turn.observe_invalid_upstream_text(text.as_str())
}
else {
None
}
}
}
}
_ => None,
};
update_response_in_flight(bound, parsed_upstream_frame.as_ref());
if matches!(
observation,
Some(ResponsesWebSocketTurnObservation::Started)
| Some(ResponsesWebSocketTurnObservation::Terminal(_))
) {
if let Some(turn) = bound.active_turn.as_mut() {
turn.mark_stream_started(state).await;
}
}
let terminal_outcome = match observation {
Some(ResponsesWebSocketTurnObservation::Terminal(outcome)) => Some(outcome),
_ => None,
};
if terminal_outcome.is_some() {
bound.response_in_flight = false;
}
if matches!(&upstream_message, WreqWsMessage::Text(_))
&& parsed_upstream_frame.is_none()
{
let policy = fatal_relay_policy(FatalRelaySignal::InvalidUpstreamText);
finalize_active_turn(
bound,
state,
terminal_outcome.unwrap_or_else(
ResponsesWebSocketTurnOutcome::upstream_receive_failed,
),
)
.await;
bound.active_response_create = None;
send_responses_websocket_error(
client_socket,
policy.status_code,
"server_error",
policy.error_code,
policy.client_message,
)
.await;
close_bound_upstream(bound).await;
close_client_socket(
client_socket,
policy.close_code,
policy.close_reason,
)
.await;
break;
}
let is_close = matches!(upstream_message, WreqWsMessage::Close(_));
let drain_for_adapter = adapter_drain_ready(
bound.pending_adapter_drain,
bound.response_in_flight,
observation,
is_close,
);
let quota_facts = QuotaRelayFacts {
drain_ready: drain_for_adapter,
retry_current_turn: bound
.pending_adapter_drain
.is_some_and(|directive| directive.retry_current_turn),
transparent_retry_failed: false,
usage_limit_error: parsed_upstream_event.is_some_and(is_usage_limit_error_event),
continuation_retry_eligible: active_continuation_can_retry_from_full_input(bound),
upstream_closed: is_close,
};
let mut quota_relay_action = classify_quota_relay(quota_facts);
if matches!(quota_relay_action, QuotaRelayAction::AttemptTransparentRetry) {
let mut retry_turn = bound.active_turn.take().map(ActiveResponsesWebSocketTurn::disarm);
if let Some(turn) = retry_turn.as_mut() {
turn.release_admission().await;
}
if retry_active_turn_after_quota_exhaustion(bound, state, context).await {
if let Some(turn) = retry_turn {
queue_turn_finalization(
bound,
state,
turn,
terminal_outcome.unwrap_or_else(
ResponsesWebSocketTurnOutcome::upstream_closed,
),
)
.await;
}
continue;
}
bound.active_turn = retry_turn.map(|turn| ActiveResponsesWebSocketTurn::new(state, turn));
quota_relay_action = classify_quota_relay(QuotaRelayFacts {
retry_current_turn: false,
transparent_retry_failed: true,
..quota_facts
});
}
if matches!(
quota_relay_action,
QuotaRelayAction::RequestFullContinuationRetry
) {
let directive = bound
.pending_adapter_drain
.expect("adapter drain state should be present");
debug!(
event_name = "responses_websocket_continuation_retry_required",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error_code = "previous_response_not_found",
"gateway will ask the client to retry the continuation with complete input"
);
let mut turn = bound.active_turn.take().map(ActiveResponsesWebSocketTurn::disarm);
if let Some(active_turn) = turn.as_mut() {
active_turn.release_admission().await;
}
send_previous_response_not_found(client_socket).await;
if let Some(turn) = turn {
queue_turn_finalization(
bound,
state,
turn,
terminal_outcome.unwrap_or_else(
ResponsesWebSocketTurnOutcome::upstream_closed,
),
)
.await;
}
bound.active_response_create = None;
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
continue;
}
if matches!(quota_relay_action, QuotaRelayAction::ForwardQuotaAndDetach) {
let directive = bound
.pending_adapter_drain
.expect("adapter drain state should be present");
finalize_active_turn(
bound,
state,
terminal_outcome
.unwrap_or_else(ResponsesWebSocketTurnOutcome::provider_quota_exhausted),
)
.await;
send_gateway_error_with_status(
client_socket,
429,
directive.error_code,
"Provider connection closed after reporting exhausted quota; send a new response.create to select another Provider connection",
)
.await;
bound.active_response_create = None;
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
continue;
}
if let Err(error) = send_client_message(
client_socket,
upstream_message_to_client(upstream_message.clone()),
).await {
warn!(
event_name = "responses_websocket_client_send_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error_code = error.as_str(),
"gateway could not relay a provider event to the client"
);
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::client_disconnected(),
).await;
bound.active_response_create = None;
close_bound_upstream(bound).await;
break;
}
if let (Some(turn), Some(frame)) =
(bound.active_turn.as_mut(), parsed_upstream_frame.as_ref())
{
turn.capture_client_frame(frame.event());
}
if let Some(outcome) = terminal_outcome {
finalize_active_turn(bound, state, outcome).await;
bound.active_response_create = None;
} else if is_close {
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::upstream_closed(),
)
.await;
}
if drain_for_adapter {
let directive = bound
.pending_adapter_drain
.expect("adapter drain state should be present");
bound.active_response_create = None;
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
continue;
}
if is_close {
bound.upstream = None;
break;
}
}
}
}
}
async fn wait_for_connection_permit_loss(permit: Option<&aether_runtime::AdmissionPermit>) {
let Some(permit) = permit else {
std::future::pending::<()>().await;
return;
};
let mut health = tokio::time::interval(Duration::from_secs(1));
health.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
loop {
health.tick().await;
if !permit.is_healthy() {
return;
}
}
}
fn update_response_in_flight(
bound: &mut BoundResponsesConnection,
frame: Option<&ParsedResponsesWebSocketFrame<'_>>,
) {
let Some(frame) = frame else {
return;
};
let frame_kind = if frame.is_terminal() {
UpstreamFrameKind::Terminal
} else if frame.is_started() {
UpstreamFrameKind::Started
} else {
UpstreamFrameKind::Other
};
match classify_upstream_frame(frame_kind) {
UpstreamFrameAction::Continue if frame.is_started() => {
bound.response_in_flight = true;
}
UpstreamFrameAction::FinalizeTurn => {
bound.response_in_flight = false;
}
UpstreamFrameAction::Continue | UpstreamFrameAction::FinalizeAndClose => {}
}
}
@@ -0,0 +1,371 @@
//! Parsed OpenAI Responses WebSocket text frames.
//!
//! A relay frame is parsed once and then shared by the protocol adapter, turn
//! accounting, retry safety, and connection lifecycle code. Keeping the raw
//! text as a borrow avoids copying the websocket payload while the relay is
//! processing it.
use serde_json::Value;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) struct ResponsesWebSocketFrameTerminal {
pub(super) status_code: u16,
pub(super) cancelled: bool,
}
#[derive(Debug)]
pub(super) struct ParsedResponsesWebSocketFrame<'a> {
raw_text: &'a str,
event: Value,
event_type: Option<String>,
status: Option<u16>,
started: bool,
terminal: Option<ResponsesWebSocketFrameTerminal>,
terminal_event: Option<Value>,
chunked: bool,
}
impl<'a> ParsedResponsesWebSocketFrame<'a> {
pub(super) fn parse(raw_text: &'a str) -> serde_json::Result<Self> {
let event = serde_json::from_str::<Value>(raw_text)?;
let events = protocol_events_of(&event);
let started = events.iter().copied().any(event_is_started);
// A batch carries at most one terminal in practice. Taking the first
// in document order keeps the outcome deterministic if that ever
// stops being true.
let terminal_entry = events
.iter()
.copied()
.find_map(|candidate| terminal_for_event(candidate).map(|term| (candidate, term)));
let terminal = terminal_entry.map(|(_, terminal)| terminal);
// The terminal event describes the turn's outcome, so it is the one
// worth naming in logs and recording as the terminal error body.
let event_type = terminal_entry
.map(|(candidate, _)| candidate)
.or_else(|| events.last().copied())
.and_then(event_type_of)
.map(str::to_string);
let terminal_event = terminal_entry.map(|(candidate, _)| candidate.clone());
let chunked = event.get("chunks").and_then(Value::as_array).is_some();
let status = terminal.map(|terminal| terminal.status_code);
Ok(Self {
raw_text,
event,
event_type,
status,
started,
terminal,
terminal_event,
chunked,
})
}
/// The protocol events this frame carries.
///
/// Codex batches standard `response.*` events into a `{"chunks":[...]}`
/// envelope, so one frame can carry several events — and the terminal one
/// may be buried inside the batch. Every consumer that interprets event
/// semantics must walk this rather than the envelope, or a batched
/// `response.completed` goes unnoticed and wedges the turn.
pub(super) fn protocol_events(&self) -> Vec<&Value> {
protocol_events_of(&self.event)
}
/// The individual event that ended the turn, unwrapped from its batch.
pub(super) fn terminal_event(&self) -> Option<&Value> {
self.terminal_event.as_ref()
}
pub(super) fn is_chunked(&self) -> bool {
self.chunked
}
pub(super) fn raw_text(&self) -> &'a str {
self.raw_text
}
pub(super) fn event(&self) -> &Value {
&self.event
}
pub(super) fn event_type(&self) -> Option<&str> {
self.event_type.as_deref()
}
pub(super) fn status(&self) -> Option<u16> {
self.status
}
pub(super) fn is_started(&self) -> bool {
self.started
}
pub(super) fn is_terminal(&self) -> bool {
self.terminal.is_some()
}
pub(super) fn terminal(&self) -> Option<ResponsesWebSocketFrameTerminal> {
self.terminal
}
/// Return a bounded label suitable for structured logs. Event payloads
/// are never inserted directly into a log field.
pub(super) fn event_type_for_log(&self) -> String {
self.event_type
.as_deref()
.map(safe_websocket_event_label)
.unwrap_or_else(|| "invalid_json".to_string())
}
}
/// Flattens a frame into the events it carries. An envelope may name its own
/// `type` *and* batch further events under `chunks`; both are protocol events.
fn protocol_events_of(event: &Value) -> Vec<&Value> {
let mut events = Vec::new();
if event_type_of(event).is_some() {
events.push(event);
}
if let Some(chunks) = event.get("chunks").and_then(Value::as_array) {
events.extend(chunks.iter().filter(|chunk| event_type_of(chunk).is_some()));
}
// An unrecognized shape is still relayed and still accounted for, so it
// must not vanish from the observer's view of the stream.
if events.is_empty() {
events.push(event);
}
events
}
fn event_type_of(event: &Value) -> Option<&str> {
event.get("type").and_then(Value::as_str)
}
fn event_is_started(event: &Value) -> bool {
matches!(
event_type_of(event).unwrap_or_default(),
"response.created" | "response.in_progress" | "response.queued"
)
}
fn terminal_for_event(event: &Value) -> Option<ResponsesWebSocketFrameTerminal> {
match event_type_of(event).unwrap_or_default() {
"response.completed" => Some(ResponsesWebSocketFrameTerminal {
status_code: websocket_event_status_code(event, 200),
cancelled: false,
}),
"response.incomplete" => Some(ResponsesWebSocketFrameTerminal {
status_code: websocket_event_status_code(event, 502),
cancelled: false,
}),
"response.cancelled" => Some(ResponsesWebSocketFrameTerminal {
status_code: 499,
cancelled: true,
}),
"response.failed" => Some(ResponsesWebSocketFrameTerminal {
status_code: websocket_event_status_code(event, 502),
cancelled: false,
}),
"error" => Some(ResponsesWebSocketFrameTerminal {
status_code: websocket_event_status_code(event, 502),
cancelled: false,
}),
_ => None,
}
}
fn websocket_event_status_code(event: &Value, default: u16) -> u16 {
if let Some(status_code) = event
.get("status_code")
.or_else(|| event.get("status"))
.or_else(|| {
event
.get("response")
.and_then(|response| response.get("status_code"))
})
.and_then(Value::as_u64)
.and_then(|value| u16::try_from(value).ok())
.filter(|value| *value > 0)
{
return status_code;
}
let error_code = [
event.pointer("/error/type"),
event.pointer("/error/code"),
event.pointer("/response/error/type"),
event.pointer("/response/error/code"),
]
.into_iter()
.flatten()
.filter_map(Value::as_str)
.map(str::to_ascii_lowercase)
.find(|value| !value.trim().is_empty());
match error_code.as_deref() {
Some(
"usage_limit_reached" | "insufficient_quota" | "rate_limit_exceeded" | "quota_exceeded",
) => 429,
Some("invalid_api_key" | "authentication_error") => 401,
Some("invalid_request_error" | "invalid_request" | "model_not_found") => 400,
Some("overloaded" | "server_error" | "service_unavailable") => 503,
_ => default,
}
}
fn safe_websocket_event_label(value: &str) -> String {
let value = value.trim();
if value.is_empty()
|| value.len() > 80
|| !value
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'_' | b'-'))
{
return "unknown".to_string();
}
value.to_string()
}
#[cfg(test)]
mod tests {
use super::ParsedResponsesWebSocketFrame;
#[test]
fn parses_started_frame_once_with_raw_text_and_event_metadata() {
let raw = r#"{"type":"response.in_progress","response":{"status":200}}"#;
let frame = ParsedResponsesWebSocketFrame::parse(raw).expect("valid frame");
assert_eq!(frame.raw_text(), raw);
assert_eq!(frame.event_type(), Some("response.in_progress"));
assert_eq!(frame.status(), None);
assert!(frame.is_started());
assert!(!frame.is_terminal());
assert_eq!(frame.event()["response"]["status"], 200);
assert_eq!(frame.event_type_for_log(), "response.in_progress");
}
#[test]
fn classifies_terminal_status_and_cancellation() {
let completed = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"response.completed","status_code":201}"#,
)
.expect("valid frame");
assert_eq!(completed.status(), Some(201));
assert_eq!(
completed
.terminal()
.map(|terminal| (terminal.status_code, terminal.cancelled)),
Some((201, false))
);
let cancelled = ParsedResponsesWebSocketFrame::parse(r#"{"type":"response.cancelled"}"#)
.expect("valid frame");
assert_eq!(cancelled.status(), Some(499));
assert_eq!(
cancelled
.terminal()
.map(|terminal| (terminal.status_code, terminal.cancelled)),
Some((499, true))
);
let error = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"error","status_code":429,"error":{"type":"usage_limit_reached"}}"#,
)
.expect("valid frame");
assert_eq!(error.status(), Some(429));
assert!(error.is_terminal());
let failed = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"response.failed","response":{"error":{"code":"rate_limit_exceeded"}}}"#,
)
.expect("valid frame");
assert_eq!(failed.status(), Some(429));
}
#[test]
fn detects_a_terminal_batched_inside_a_chunks_envelope() {
let frame = ParsedResponsesWebSocketFrame::parse(
r#"{"chunks":[{"type":"response.output_text.delta","delta":"hi"},{"type":"response.completed","response":{"usage":{"total_tokens":8}}}]}"#,
)
.expect("valid frame");
assert!(frame.is_chunked());
assert!(frame.is_terminal());
assert_eq!(frame.status(), Some(200));
// The label and the recorded error body must name the event that ended
// the turn, not the envelope.
assert_eq!(frame.event_type(), Some("response.completed"));
assert_eq!(
frame.terminal_event().and_then(|event| event
.pointer("/response/usage/total_tokens")
.and_then(serde_json::Value::as_u64)),
Some(8)
);
assert_eq!(frame.protocol_events().len(), 2);
}
#[test]
fn detects_a_start_event_batched_inside_a_chunks_envelope() {
let frame = ParsedResponsesWebSocketFrame::parse(
r#"{"chunks":[{"type":"codex.rate_limits"},{"type":"response.created"}]}"#,
)
.expect("valid frame");
assert!(frame.is_started());
assert!(!frame.is_terminal());
assert_eq!(frame.protocol_events().len(), 2);
}
#[test]
fn an_envelope_may_carry_its_own_type_alongside_batched_events() {
let frame = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"codex.response.metadata","chunks":[{"type":"response.failed","response":{"error":{"code":"rate_limit_exceeded"}}}]}"#,
)
.expect("valid frame");
assert_eq!(frame.protocol_events().len(), 2);
assert!(frame.is_terminal());
assert_eq!(frame.status(), Some(429));
assert_eq!(frame.event_type(), Some("response.failed"));
}
#[test]
fn a_batch_without_a_terminal_does_not_end_the_turn() {
let frame = ParsedResponsesWebSocketFrame::parse(
r#"{"chunks":[{"type":"response.output_text.delta","delta":"a"},{"type":"response.output_text.delta","delta":"b"}]}"#,
)
.expect("valid frame");
assert!(!frame.is_terminal());
assert!(!frame.is_started());
assert!(frame.terminal_event().is_none());
}
#[test]
fn an_unrecognized_shape_is_still_surfaced_as_one_event() {
let frame =
ParsedResponsesWebSocketFrame::parse(r#"{"unexpected":true}"#).expect("valid frame");
assert_eq!(frame.protocol_events().len(), 1);
assert!(!frame.is_chunked());
assert!(!frame.is_terminal());
assert_eq!(frame.event_type(), None);
assert_eq!(frame.event_type_for_log(), "invalid_json");
}
#[test]
fn preserves_safe_log_label_boundaries() {
let unsafe_label =
ParsedResponsesWebSocketFrame::parse(r#"{"type":"not safe / contains spaces"}"#)
.expect("valid frame");
assert_eq!(unsafe_label.event_type_for_log(), "unknown");
let missing_label =
ParsedResponsesWebSocketFrame::parse(r#"{"message":"ok"}"#).expect("valid frame");
assert_eq!(missing_label.event_type_for_log(), "invalid_json");
}
#[test]
fn rejects_invalid_json() {
assert!(ParsedResponsesWebSocketFrame::parse("not-json").is_err());
}
}
@@ -0,0 +1,251 @@
//! Turn finalization and terminal error mapping for a Responses WebSocket.
//!
//! A connection can outlive a turn, so persistence and adapter observation
//! handles are joined in order before the next turn is planned.
use std::time::Duration;
use axum::extract::ws::WebSocket;
use tokio::task::JoinHandle;
use tokio::time::timeout;
use super::state::BoundResponsesConnection;
use super::turn::{
spawn_responses_websocket_turn_finalization, ResponsesWebSocketTurn,
ResponsesWebSocketTurnOutcome,
};
use crate::handlers::proxy::websocket::session::{
CLOSE_INTERNAL_ERROR, CLOSE_POLICY_VIOLATION, CLOSE_TRY_AGAIN, WEBSOCKET_LOG_TRANSPORT,
};
use crate::handlers::proxy::websocket::transport::send_responses_websocket_error;
use crate::{AppState, GatewayError};
const RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT: Duration = Duration::from_secs(5);
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
macro_rules! warn {
($($arg:tt)*) => {
tracing::warn!(target: LOG_TARGET, $($arg)*)
};
}
/// Owns the in-flight turn so that losing the relay task still finalizes it.
///
/// Every ordinary exit path takes the turn out of here and finalizes it
/// explicitly. This guard only covers the paths that are not exit paths at all
/// — a panic in the relay loop, or the task being dropped — where the turn
/// would otherwise be discarded with its usage row left `Pending`, its
/// candidate row left `Streaming`, and its distributed pool key lease leaked
/// until the lease expires. Mirrors the HTTP path's `DirectPassthroughFinalizer`.
pub(super) struct ActiveResponsesWebSocketTurn {
turn: Option<ResponsesWebSocketTurn>,
state: AppState,
}
impl ActiveResponsesWebSocketTurn {
pub(super) fn new(state: &AppState, turn: ResponsesWebSocketTurn) -> Self {
Self {
turn: Some(turn),
state: state.clone(),
}
}
/// Hands the turn back to a caller that will finalize it explicitly.
pub(super) fn disarm(mut self) -> ResponsesWebSocketTurn {
self.turn
.take()
.expect("an armed active turn always holds its turn")
}
}
impl std::ops::Deref for ActiveResponsesWebSocketTurn {
type Target = ResponsesWebSocketTurn;
fn deref(&self) -> &Self::Target {
self.turn
.as_ref()
.expect("an armed active turn always holds its turn")
}
}
impl std::ops::DerefMut for ActiveResponsesWebSocketTurn {
fn deref_mut(&mut self) -> &mut Self::Target {
self.turn
.as_mut()
.expect("an armed active turn always holds its turn")
}
}
impl Drop for ActiveResponsesWebSocketTurn {
fn drop(&mut self) {
let Some(turn) = self.turn.take() else {
return;
};
let state = self.state.clone();
// No runtime means the process is going down; the spawn could not
// complete anyway.
if let Ok(handle) = tokio::runtime::Handle::try_current() {
warn!(
event_name = "responses_websocket_turn_abandoned",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
"gateway finalized a Responses WebSocket turn whose relay task went away"
);
handle.spawn(async move {
turn.finalize_detached(
&state,
ResponsesWebSocketTurnOutcome::relay_task_abandoned(),
)
.await;
});
}
}
}
pub(super) async fn finalize_active_turn(
bound: &mut BoundResponsesConnection,
state: &AppState,
outcome: ResponsesWebSocketTurnOutcome,
) {
if let Some(turn) = bound.active_turn.take() {
queue_turn_finalization(bound, state, turn.disarm(), outcome).await;
}
}
pub(super) async fn queue_turn_finalization(
bound: &mut BoundResponsesConnection,
state: &AppState,
turn: ResponsesWebSocketTurn,
outcome: ResponsesWebSocketTurnOutcome,
) {
await_pending_adapter_observation(bound).await;
await_pending_turn_finalization(bound).await;
bound.pending_turn_finalization =
Some(spawn_responses_websocket_turn_finalization(state.clone(), turn, outcome).await);
}
pub(super) async fn await_pending_adapter_observation(bound: &mut BoundResponsesConnection) {
if let Some(mut handle) = bound.pending_adapter_observation.take() {
match timeout(RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT, &mut handle).await {
Ok(Err(error)) => {
warn!(
event_name = "responses_websocket_adapter_observation_join_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
error = ?error,
"gateway Responses WebSocket adapter observation task failed"
);
}
Ok(Ok(())) => {}
Err(_) => {
handle.abort();
let _ = handle.await;
warn!(
event_name = "responses_websocket_adapter_observation_timeout",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
timeout_ms = RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT.as_millis() as u64,
"gateway stopped waiting for a Responses WebSocket adapter observation"
);
}
}
}
}
pub(super) async fn finalize_unbound_turn(
state: AppState,
turn: ResponsesWebSocketTurn,
outcome: ResponsesWebSocketTurnOutcome,
) -> JoinHandle<()> {
spawn_responses_websocket_turn_finalization(state, turn, outcome).await
}
pub(super) async fn await_turn_finalization_handle(handle: JoinHandle<()>) {
// Do not abort terminal persistence here. Each I/O stage inside the turn
// finalizer is independently bounded, and aborting the owner would skip
// pool-lease cleanup and leave usage/candidate state non-terminal.
match handle.await {
Ok(()) => {}
Err(error) => {
warn!(
event_name = "responses_websocket_turn_finalization_join_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
error = ?error,
"gateway Responses WebSocket turn finalizer task failed"
);
}
}
}
pub(super) async fn await_pending_turn_finalization(bound: &mut BoundResponsesConnection) {
if let Some(handle) = bound.pending_turn_finalization.take() {
await_turn_finalization_handle(handle).await;
}
}
pub(super) async fn send_responses_websocket_turn_start_error(
client_socket: &mut WebSocket,
error: &GatewayError,
) {
match error {
GatewayError::Client { status, message } => {
let (error_type, code) = if status.as_u16() == 429 {
("rate_limit_error", "gateway_request_capacity_exceeded")
} else {
("invalid_request_error", "gateway_request_not_allowed")
};
send_responses_websocket_error(
client_socket,
status.as_u16(),
error_type,
code,
message,
)
.await;
}
GatewayError::AdmissionTimeout { .. } => {
send_responses_websocket_error(
client_socket,
503,
"server_error",
"gateway_admission_timeout",
"Gateway capacity is busy; retry this response",
)
.await;
}
GatewayError::LocalExecutionPlanningTimeout { .. } => {
send_responses_websocket_error(
client_socket,
504,
"server_error",
"gateway_planning_timeout",
"Gateway planning timed out; retry this response",
)
.await;
}
_ => {
send_responses_websocket_error(
client_socket,
500,
"server_error",
"responses_websocket_turn_start_failed",
"Gateway could not start this response",
)
.await;
}
}
}
pub(super) fn responses_websocket_turn_start_close(error: &GatewayError) -> (u16, &'static str) {
match error {
GatewayError::Client { .. } => (CLOSE_POLICY_VIOLATION, "request_not_allowed"),
GatewayError::AdmissionTimeout { .. }
| GatewayError::LocalExecutionPlanningTimeout { .. } => (CLOSE_TRY_AGAIN, "gateway_busy"),
_ => (CLOSE_INTERNAL_ERROR, "turn_start_failed"),
}
}
@@ -0,0 +1,61 @@
//! OpenAI Responses WebSocket protocol entry point, session engine, and adapters.
//!
//! The route is protocol-oriented. `session` bootstraps the authenticated
//! connection, `connection` owns the socket FSM, `client` and `quota` own
//! protocol/retry policy, and `lifecycle`/`turn` bridge each turn into the
//! existing usage and audit runtime. Adapters contain only provider-specific
//! connection and metadata behavior.
mod adapter;
mod adapters;
mod admission;
mod binding;
mod client;
mod connection;
mod frame;
mod lifecycle;
mod quota;
mod relay_policy;
mod request;
mod session;
mod state;
mod turn;
mod upstream;
use std::net::SocketAddr;
use axum::body::Body;
use axum::extract::ws::WebSocketUpgrade;
use axum::extract::{ConnectInfo, State};
use axum::http::{HeaderMap, Response, Uri};
use crate::handlers::proxy::websocket::ingress::{
upgrade_authenticated_ai_websocket, WebSocketIngressSpec,
};
use crate::handlers::proxy::websocket::session::RESPONSES_WEBSOCKET_SESSION_LIMITS;
use crate::{AppState, GatewayError};
pub(crate) async fn responses_websocket(
State(state): State<AppState>,
ConnectInfo(remote_addr): ConnectInfo<SocketAddr>,
ws: WebSocketUpgrade,
headers: HeaderMap,
uri: Uri,
) -> Result<Response<Body>, GatewayError> {
upgrade_authenticated_ai_websocket(
state,
remote_addr,
ws,
headers,
uri,
RESPONSES_WEBSOCKET_SESSION_LIMITS,
RESPONSES_WEBSOCKET_INGRESS_SPEC,
session::run_responses_websocket,
)
.await
}
const RESPONSES_WEBSOCKET_INGRESS_SPEC: WebSocketIngressSpec = WebSocketIngressSpec {
route_unavailable_message: "WebSocket route is unavailable",
ip_whitelist_failure_event_name: "responses_websocket_ip_whitelist_check_failed",
};
@@ -0,0 +1,380 @@
//! Quota exhaustion, replay safety, and upstream replacement policy.
use axum::extract::ws::WebSocket;
use futures_util::SinkExt;
use serde_json::Value;
use uuid::Uuid;
use wreq::ws::message::Message as WreqWsMessage;
use super::adapter::{
resolve_responses_websocket_adapter, ResponsesWebSocketDrainDirective,
ResponsesWebSocketRebindSafety,
};
use super::lifecycle::{queue_turn_finalization, ActiveResponsesWebSocketTurn};
use super::request::{
build_planning_parts, planned_response_create_event, response_create_has_previous_response_id,
};
use super::state::BoundResponsesConnection;
use super::turn::{
begin_responses_websocket_turn, prepare_responses_websocket_turn_decision,
ResponsesWebSocketTurnOutcome,
};
use super::upstream::{bind_responses_upstream, close_bound_upstream};
use crate::ai_serving::maybe_build_responses_websocket_decision;
use crate::clock::current_unix_secs;
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
use crate::handlers::proxy::websocket::session::WEBSOCKET_LOG_TRANSPORT;
use crate::handlers::proxy::websocket::transport::{
close_upstream_socket, send_responses_websocket_error,
};
use crate::orchestration::release_pool_key_lease_from_report_context;
use crate::AppState;
const PREVIOUS_RESPONSE_NOT_FOUND_MESSAGE: &str =
"Previous response was not found. Retrying the full request.";
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
macro_rules! debug {
($($arg:tt)*) => {
tracing::debug!(target: LOG_TARGET, $($arg)*)
};
}
macro_rules! warn {
($($arg:tt)*) => {
tracing::warn!(target: LOG_TARGET, $($arg)*)
};
}
pub(super) async fn detach_exhausted_upstream(
bound: &mut BoundResponsesConnection,
directive: ResponsesWebSocketDrainDirective,
trace_id: &str,
) {
let exclusion = record_exhausted_bound_key(bound, directive.retry_exclusion_until_unix_secs);
close_bound_upstream(bound).await;
bound.response_in_flight = false;
bound.pending_adapter_drain = None;
let now_unix_secs = current_unix_secs();
debug!(
event_name = "responses_websocket_upstream_detached",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %trace_id,
reason = directive.error_code,
exhausted_key_id = ?exclusion.as_ref().map(|(key_id, _)| key_id),
retry_exclusion_until_unix_secs = ?exclusion.as_ref().map(|(_, until)| until),
exhausted_exclusion_count = bound.exhausted_exclusions.len(now_unix_secs),
"gateway detached an exhausted Responses WebSocket upstream while preserving the client socket"
);
}
pub(super) fn record_exhausted_bound_key(
bound: &mut BoundResponsesConnection,
reset_at_unix_secs: Option<u64>,
) -> Option<(String, u64)> {
let key_id = bound
.decision_template
.key_id
.as_deref()
.map(str::trim)
.filter(|key_id| !key_id.is_empty())?
.to_string();
let provider_account_id = bound
.adapter
.exhaustion_exclusion_identity(&bound.decision_template)
.and_then(|identity| identity.account_id);
let exclusion_until = bound.exhausted_exclusions.exclude(
key_id.clone(),
provider_account_id,
reset_at_unix_secs,
current_unix_secs(),
);
Some((key_id, exclusion_until))
}
pub(super) async fn retry_active_turn_after_quota_exhaustion(
bound: &mut BoundResponsesConnection,
state: &AppState,
context: &WebSocketRequestContext,
) -> bool {
let Some(active) = bound.active_response_create.as_mut() else {
return false;
};
if let Some(reason) = active.quota_retry_block_reason() {
debug!(
event_name = "responses_websocket_quota_retry_skipped",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
turn_index = active.turn_index,
logical_turn_id = %active.logical_turn_id,
turn_attempt = active.turn_attempt,
reason,
"gateway will not transparently replay an unsafe Responses WebSocket turn"
);
return false;
}
active.retry_attempted = true;
active.turn_attempt = active.turn_attempt.saturating_add(1);
let client_event = active.client_event.clone();
let turn_index = active.turn_index;
let logical_turn_id = active.logical_turn_id.clone();
let turn_attempt = active.turn_attempt;
let retry_exclusion_until_unix_secs = bound
.pending_adapter_drain
.and_then(|directive| directive.retry_exclusion_until_unix_secs);
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 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);
let excluded_codex_account_ids = bound.exhausted_exclusions.codex_account_ids(now_unix_secs);
let excluded_key_ids = (!excluded_key_ids.is_empty()).then_some(&excluded_key_ids);
let excluded_codex_account_ids =
(!excluded_codex_account_ids.is_empty()).then_some(&excluded_codex_account_ids);
let planned = match maybe_build_responses_websocket_decision(
state,
&planning_parts,
&turn_request_id,
&context.decision,
&client_event,
excluded_key_ids,
excluded_codex_account_ids,
)
.await
{
Ok(Some(decision)) => decision,
Ok(None) => {
warn!(
event_name = "responses_websocket_quota_retry_provider_unavailable",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
exhausted_key_id = ?exhausted_key_id,
"gateway could not find an alternate Responses WebSocket provider after quota exhaustion"
);
return false;
}
Err(error) => {
warn!(
event_name = "responses_websocket_quota_retry_planning_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
exhausted_key_id = ?exhausted_key_id,
error = ?error,
"gateway could not plan an alternate Responses WebSocket provider after quota exhaustion"
);
return false;
}
};
let adapter = resolve_responses_websocket_adapter(planned.adapter);
let normalization = planned.normalization;
let decision = planned.execution;
if exhausted_key_id.as_deref() == decision.key_id.as_deref() {
release_pool_key_lease_from_report_context(state, decision.report_context.as_ref()).await;
warn!(
event_name = "responses_websocket_quota_retry_selected_exhausted_key",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
key_id = ?decision.key_id,
"gateway rejected an alternate Responses WebSocket plan that reused the exhausted key"
);
return false;
}
let provider_event = match planned_response_create_event(&decision, &client_event).and_then(
|event| {
serde_json::from_str::<Value>(&event)
.map_err(|_| "response_create_serialization_failed")
},
) {
Ok(event) => event,
Err(code) => {
release_pool_key_lease_from_report_context(state, decision.report_context.as_ref())
.await;
warn!(
event_name = "responses_websocket_quota_retry_normalization_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error_code = code,
"gateway could not rebuild a Responses response.create for transparent quota retry"
);
return false;
}
};
let turn_decision = prepare_responses_websocket_turn_decision(
&decision,
turn_request_id,
true,
&client_event,
&provider_event,
&context.trace_id,
turn_index,
&logical_turn_id,
turn_attempt,
);
let mut turn = match begin_responses_websocket_turn(
state,
&planning_parts,
&context.decision,
turn_decision,
&client_event,
)
.await
{
Ok(turn) => turn,
Err(error) => {
warn!(
event_name = "responses_websocket_quota_retry_reporting_unavailable",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error = ?error,
"gateway could not start usage and audit tracking for transparent quota retry"
);
return false;
}
};
let mut replacement = match bind_responses_upstream(
&decision,
normalization,
&client_event,
adapter,
)
.await
{
Ok(connection) => connection,
Err(code) => {
queue_turn_finalization(
bound,
state,
turn,
ResponsesWebSocketTurnOutcome::upstream_connect_failed(code),
)
.await;
warn!(
event_name = "responses_websocket_quota_retry_rebind_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error_code = code,
"gateway could not bind an alternate Responses WebSocket provider after quota exhaustion"
);
return false;
}
};
turn.mark_upstream_request_sent();
turn.set_provider_response_headers(replacement.upstream_response_headers.clone());
let replacement_upstream = replacement
.upstream
.take()
.expect("newly bound Responses upstream should be present");
if let Some(mut previous_upstream) = bound.upstream.replace(replacement_upstream) {
close_upstream_socket(&mut previous_upstream, None).await;
}
let previous_key_id = bound.decision_template.key_id.clone();
bound.adapter = replacement.adapter;
bound.client_model = replacement.client_model;
bound.provider_model = replacement.provider_model;
bound.response_in_flight = true;
bound.decision_template = replacement.decision_template;
bound.body_normalization = replacement.body_normalization;
bound.binding_identity = replacement.binding_identity;
bound.active_turn = Some(ActiveResponsesWebSocketTurn::new(state, turn));
bound.upstream_response_headers = replacement.upstream_response_headers;
bound.pending_adapter_drain = None;
debug!(
event_name = "responses_websocket_quota_retry_rebound",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
turn_index,
logical_turn_id = %logical_turn_id,
turn_attempt,
previous_key_id = ?previous_key_id,
key_id = ?bound.decision_template.key_id,
"gateway transparently rebound a Responses WebSocket turn after quota exhaustion"
);
true
}
pub(super) fn active_continuation_can_retry_from_full_input(
bound: &BoundResponsesConnection,
) -> bool {
bound.active_response_create.as_ref().is_some_and(|active| {
response_create_has_previous_response_id(&active.client_event)
&& active.retry_unsafe_reason.is_none()
})
}
pub(super) fn is_usage_limit_error_event(event: &Value) -> bool {
let is_error = |value: &Value| {
value.get("type").and_then(Value::as_str) == Some("error")
&& value.pointer("/error/type").and_then(Value::as_str) == Some("usage_limit_reached")
};
is_error(event)
|| event
.get("chunks")
.and_then(Value::as_array)
.is_some_and(|chunks| chunks.iter().any(is_error))
}
pub(super) fn should_request_full_continuation_retry(
bound: &BoundResponsesConnection,
retry_current_turn: bool,
upstream_event: Option<&Value>,
) -> bool {
retry_current_turn
&& active_continuation_can_retry_from_full_input(bound)
&& upstream_event.is_some_and(is_usage_limit_error_event)
}
pub(super) async fn send_previous_response_not_found(client_socket: &mut WebSocket) {
send_responses_websocket_error(
client_socket,
400,
"invalid_request_error",
"previous_response_not_found",
PREVIOUS_RESPONSE_NOT_FOUND_MESSAGE,
)
.await;
}
pub(super) fn observe_active_response_rebind_safety(
bound: &mut BoundResponsesConnection,
event: &Value,
) {
let ResponsesWebSocketRebindSafety::Unsafe { reason } =
bound.adapter.rebind_safety_for_upstream_event(event)
else {
return;
};
if let Some(active) = bound.active_response_create.as_mut() {
active.mark_retry_unsafe(reason);
}
}
pub(super) fn mark_active_response_retry_unsafe(
bound: &mut BoundResponsesConnection,
reason: &'static str,
) {
if let Some(active) = bound.active_response_create.as_mut() {
active.mark_retry_unsafe(reason);
}
}
@@ -0,0 +1,298 @@
//! Pure relay policy decisions for the Responses WebSocket session.
//!
//! The session owns sockets, provider planning, and usage persistence. This
//! module deliberately owns none of those resources: it only turns observed
//! protocol facts into a bounded action. Keeping this layer dependency-free
//! makes the failure paths executable with `rustc --test` without linking the
//! full gateway (which is useful on constrained CI/diagnostic hosts).
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FatalRelaySignal {
ConnectionAdmissionLost,
InvalidUpstreamText,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct FatalRelayPolicy {
pub status_code: u16,
pub close_code: u16,
pub error_code: &'static str,
pub client_message: &'static str,
pub close_reason: &'static str,
}
/// Map a local relay failure to the status/event/close tuple sent after the
/// HTTP upgrade. In particular, capacity loss is retryable (1013), while a
/// malformed provider frame is an internal relay error (1011).
pub const fn fatal_relay_policy(signal: FatalRelaySignal) -> FatalRelayPolicy {
match signal {
FatalRelaySignal::ConnectionAdmissionLost => FatalRelayPolicy {
status_code: 503,
close_code: 1013,
error_code: "gateway_connection_admission_lost",
client_message: "Gateway capacity lease was lost; reconnect to continue",
close_reason: "connection_admission_lost",
},
FatalRelaySignal::InvalidUpstreamText => FatalRelayPolicy {
status_code: 502,
close_code: 1011,
error_code: "responses_websocket_invalid_upstream_event",
client_message: "Provider returned an invalid WebSocket event",
close_reason: "invalid_upstream_event",
},
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum UpstreamFrameKind {
Other,
Started,
Terminal,
Close,
InvalidText,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum UpstreamFrameAction {
Continue,
FinalizeTurn,
FinalizeAndClose,
}
/// Classify the lifecycle effect of one upstream frame. A malformed text
/// frame and a non-terminal close both finalize the active turn before the
/// client socket is closed; a valid terminal event finalizes the turn but is
/// still eligible for the normal downstream forwarding path.
pub const fn classify_upstream_frame(kind: UpstreamFrameKind) -> UpstreamFrameAction {
match kind {
UpstreamFrameKind::Other | UpstreamFrameKind::Started => UpstreamFrameAction::Continue,
UpstreamFrameKind::Terminal => UpstreamFrameAction::FinalizeTurn,
UpstreamFrameKind::Close | UpstreamFrameKind::InvalidText => {
UpstreamFrameAction::FinalizeAndClose
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct QuotaRelayFacts {
/// The adapter has observed a definitive quota signal and is ready to
/// drain/rebind the current upstream.
pub drain_ready: bool,
/// The adapter allows a transparent replay of this turn.
pub retry_current_turn: bool,
/// The session already attempted the adapter-approved transparent replay
/// and could not bind an alternate upstream. A continuation may request
/// complete input only after that first recovery path was exhausted.
pub transparent_retry_failed: bool,
/// The event contains the definitive `usage_limit_reached` error. A
/// merely exhausted-looking rate-limit snapshot must not trigger retry.
pub usage_limit_error: bool,
/// The active request is a continuation that can be retried from complete
/// input after the old account is detached.
pub continuation_retry_eligible: bool,
pub upstream_closed: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum QuotaRelayAction {
None,
AttemptTransparentRetry,
RequestFullContinuationRetry,
ForwardQuotaAndDetach,
}
/// Decide the quota branch before any response is forwarded to the client.
/// `retry_current_turn` intentionally wins over the continuation branch: the
/// session attempts the normal transparent retry first, then calls this again
/// with `transparent_retry_failed` after that attempt fails. This preserves
/// the Codex recovery order while making each fallback explicit.
pub const fn classify_quota_relay(facts: QuotaRelayFacts) -> QuotaRelayAction {
if !facts.drain_ready {
return QuotaRelayAction::None;
}
if facts.usage_limit_error && facts.retry_current_turn && !facts.transparent_retry_failed {
return QuotaRelayAction::AttemptTransparentRetry;
}
if facts.usage_limit_error
&& facts.continuation_retry_eligible
&& (facts.transparent_retry_failed || !facts.retry_current_turn)
{
return QuotaRelayAction::RequestFullContinuationRetry;
}
if facts.upstream_closed {
return QuotaRelayAction::ForwardQuotaAndDetach;
}
QuotaRelayAction::None
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Debug, Clone, Copy)]
struct MockUpstream {
frames: &'static [UpstreamFrameKind],
cursor: usize,
}
impl MockUpstream {
const fn new(frames: &'static [UpstreamFrameKind]) -> Self {
Self { frames, cursor: 0 }
}
fn next(&mut self) -> Option<UpstreamFrameKind> {
let frame = self.frames.get(self.cursor).copied()?;
self.cursor += 1;
Some(frame)
}
}
#[test]
fn mock_upstream_terminal_event_finalizes_without_waiting_for_an_extra_frame() {
let mut upstream = MockUpstream::new(&[
UpstreamFrameKind::Started,
UpstreamFrameKind::Other,
UpstreamFrameKind::Terminal,
]);
assert_eq!(
classify_upstream_frame(upstream.next().unwrap()),
UpstreamFrameAction::Continue
);
assert_eq!(
classify_upstream_frame(upstream.next().unwrap()),
UpstreamFrameAction::Continue
);
assert_eq!(
classify_upstream_frame(upstream.next().unwrap()),
UpstreamFrameAction::FinalizeTurn
);
assert_eq!(upstream.next(), None);
}
#[test]
fn mock_upstream_quota_429_attempts_one_transparent_retry_then_can_close() {
let first = classify_quota_relay(QuotaRelayFacts {
drain_ready: true,
retry_current_turn: true,
transparent_retry_failed: false,
usage_limit_error: true,
continuation_retry_eligible: false,
upstream_closed: false,
});
assert_eq!(first, QuotaRelayAction::AttemptTransparentRetry);
// A failed transparent retry must not loop forever. Once the adapter
// no longer permits replay, the terminal quota event is forwarded and
// the exhausted upstream is detached.
let after_retry_failure = classify_quota_relay(QuotaRelayFacts {
drain_ready: true,
retry_current_turn: false,
transparent_retry_failed: true,
usage_limit_error: true,
continuation_retry_eligible: false,
upstream_closed: true,
});
assert_eq!(after_retry_failure, QuotaRelayAction::ForwardQuotaAndDetach);
}
#[test]
fn continuation_quota_can_request_full_input_retry_without_replaying_partial_state() {
assert_eq!(
classify_quota_relay(QuotaRelayFacts {
drain_ready: true,
retry_current_turn: false,
transparent_retry_failed: false,
usage_limit_error: true,
continuation_retry_eligible: true,
upstream_closed: true,
}),
QuotaRelayAction::RequestFullContinuationRetry
);
assert_eq!(
classify_quota_relay(QuotaRelayFacts {
drain_ready: true,
retry_current_turn: false,
transparent_retry_failed: true,
usage_limit_error: true,
continuation_retry_eligible: true,
upstream_closed: true,
}),
QuotaRelayAction::RequestFullContinuationRetry
);
}
#[test]
fn continuation_quota_without_transparent_retry_support_uses_full_input_retry() {
assert_eq!(
classify_quota_relay(QuotaRelayFacts {
drain_ready: true,
retry_current_turn: false,
transparent_retry_failed: false,
usage_limit_error: true,
continuation_retry_eligible: true,
upstream_closed: false,
}),
QuotaRelayAction::RequestFullContinuationRetry
);
}
#[test]
fn connection_admission_loss_is_retryable_and_invalid_json_is_terminal() {
assert_eq!(
fatal_relay_policy(FatalRelaySignal::ConnectionAdmissionLost),
FatalRelayPolicy {
status_code: 503,
close_code: 1013,
error_code: "gateway_connection_admission_lost",
client_message: "Gateway capacity lease was lost; reconnect to continue",
close_reason: "connection_admission_lost",
}
);
assert_eq!(
fatal_relay_policy(FatalRelaySignal::InvalidUpstreamText),
FatalRelayPolicy {
status_code: 502,
close_code: 1011,
error_code: "responses_websocket_invalid_upstream_event",
client_message: "Provider returned an invalid WebSocket event",
close_reason: "invalid_upstream_event",
}
);
}
#[test]
fn invalid_json_never_maps_to_a_waiting_state() {
let mut upstream = MockUpstream::new(&[UpstreamFrameKind::InvalidText]);
let action = classify_upstream_frame(upstream.next().unwrap());
assert_eq!(action, UpstreamFrameAction::FinalizeAndClose);
assert_eq!(upstream.next(), None);
}
#[test]
fn quota_snapshot_without_definitive_error_does_not_trigger_retry() {
assert_eq!(
classify_quota_relay(QuotaRelayFacts {
drain_ready: true,
retry_current_turn: true,
transparent_retry_failed: false,
usage_limit_error: false,
continuation_retry_eligible: false,
upstream_closed: false,
}),
QuotaRelayAction::None
);
assert_eq!(
classify_quota_relay(QuotaRelayFacts {
drain_ready: true,
retry_current_turn: false,
transparent_retry_failed: true,
usage_limit_error: false,
continuation_retry_eligible: false,
upstream_closed: true,
}),
QuotaRelayAction::ForwardQuotaAndDetach
);
}
}
@@ -0,0 +1,352 @@
//! Responses WebSocket request normalization and model-selection helpers.
//!
//! These functions translate client protocol events into the HTTP-shaped
//! planning input and provider `response.create` events. They deliberately do
//! not depend on connection state or perform I/O.
use axum::http::header::{AUTHORIZATION, CONNECTION, CONTENT_TYPE, UPGRADE};
use axum::http::Method;
use serde_json::Value;
use crate::ai_serving::{AiExecutionDecision, ResponsesWebSocketBodyNormalization};
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
use crate::headers::request_origin_from_headers_and_remote_addr;
pub(super) fn build_planning_parts(context: &WebSocketRequestContext) -> http::request::Parts {
let mut request = http::Request::builder()
.method(Method::POST)
.uri(context.uri.clone())
.body(())
.expect("a validated request URI should build planning request parts");
let headers = request.headers_mut();
*headers = context.headers.clone();
headers.remove(AUTHORIZATION);
headers.remove("x-api-key");
headers.remove("api-key");
headers.remove("x-goog-api-key");
headers.remove(CONNECTION);
headers.remove(UPGRADE);
headers.remove("sec-websocket-key");
headers.remove("sec-websocket-version");
headers.remove("sec-websocket-protocol");
headers.remove("sec-websocket-extensions");
headers.insert(
CONTENT_TYPE,
http::HeaderValue::from_static("application/json"),
);
request
.extensions_mut()
.insert(request_origin_from_headers_and_remote_addr(
&context.headers,
&context.remote_addr,
));
request.into_parts().0
}
pub(super) fn planned_response_create_event(
decision: &AiExecutionDecision,
fallback: &Value,
) -> Result<String, &'static str> {
let event = decision
.provider_request_body
.clone()
.unwrap_or_else(|| fallback.clone());
finish_response_create_event(event, fallback)
}
/// Restores the WebSocket protocol framing that provider-body normalization is
/// not aware of.
///
/// `previous_response_id` is on the Codex unsupported-field list and `generate`
/// is not an HTTP body option at all, so normalization strips both — yet they
/// are the entire point of WebSocket mode. They must be re-grafted from the
/// client event afterwards. `stream`/`background` go the other way: the
/// normalizer inserts `stream`, and the WebSocket protocol has no use for it.
fn finish_response_create_event(
mut event: Value,
client_event: &Value,
) -> Result<String, &'static str> {
let object = event
.as_object_mut()
.ok_or("responses_websocket_request_invalid")?;
object.insert(
"type".to_string(),
Value::String("response.create".to_string()),
);
for field in ["previous_response_id", "generate"] {
if let Some(value) = client_event.get(field) {
if value.is_null() {
object.remove(field);
} else {
object.insert(field.to_string(), value.clone());
}
}
}
object.remove("stream");
object.remove("background");
serde_json::to_string(&event).map_err(|_| "responses_websocket_request_invalid")
}
pub(super) fn response_create_has_previous_response_id(event: &Value) -> bool {
event
.get("previous_response_id")
.is_some_and(|value| !value.is_null())
}
pub(super) fn continuation_requires_same_upstream(
event: &Value,
reuses_bound_upstream: bool,
) -> bool {
response_create_has_previous_response_id(event) && !reuses_bound_upstream
}
pub(super) fn changed_followup_response_create_model(
event: &Value,
current_client_model: &str,
) -> Result<Option<String>, &'static str> {
let Some(object) = event.as_object() else {
return Err("invalid_response_create");
};
let Some(model) = object.get("model") else {
return Ok(None);
};
let Some(model) = model
.as_str()
.map(str::trim)
.filter(|model| !model.is_empty())
else {
return Err("invalid_response_create_model");
};
if model.eq_ignore_ascii_case(current_client_model) {
Ok(None)
} else {
Ok(Some(model.to_string()))
}
}
pub(super) fn response_create_model_or_current(
event: &mut Value,
current_client_model: &str,
) -> Result<String, &'static str> {
let Some(object) = event.as_object_mut() else {
return Err("invalid_response_create");
};
let Some(model) = object.get("model") else {
object.insert(
"model".to_string(),
Value::String(current_client_model.to_string()),
);
return Ok(current_client_model.to_string());
};
let Some(model) = model
.as_str()
.map(str::trim)
.filter(|model| !model.is_empty())
else {
return Err("invalid_response_create_model");
};
Ok(model.to_string())
}
pub(super) fn provider_model_from_decision(decision: &AiExecutionDecision) -> Option<String> {
decision
.provider_request_body
.as_ref()
.and_then(|body| body.get("model"))
.and_then(Value::as_str)
.or(decision.mapped_model.as_deref())
.map(str::trim)
.filter(|model| !model.is_empty())
.map(str::to_string)
}
/// Prepares a continuation `response.create` for the already-bound upstream.
///
/// The turn cannot be re-planned without risking a different provider key, so
/// the binding's retained normalizer is replayed instead. That keeps model
/// directives, endpoint body rules and the Codex body contract applied on every
/// turn rather than only on the one that bound the socket.
pub(super) fn normalize_followup_response_create(
event: &Value,
provider_model: &str,
normalization: &ResponsesWebSocketBodyNormalization,
) -> Result<String, &'static str> {
if event.as_object().is_none() {
return Err("invalid_response_create");
}
if event.get("type").and_then(Value::as_str) != Some("response.create") {
return Err("invalid_response_create");
}
// Normalization is best-effort here: a continuation cannot fall back to
// another candidate, so a body the contract rejects is still better sent
// than dropped.
let mut normalized = normalization
.normalize_response_create(event)
.unwrap_or_else(|| event.clone());
let Some(object) = normalized.as_object_mut() else {
return Err("invalid_response_create");
};
// A continuation must never switch models mid-socket, and normalization is
// allowed to rewrite `model` (the Codex image-tool path does).
object.insert(
"model".to_string(),
Value::String(provider_model.to_string()),
);
finish_response_create_event(normalized, event)
.map_err(|_| "response_create_serialization_failed")
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::{normalize_followup_response_create, response_create_has_previous_response_id};
use crate::ai_serving::ResponsesWebSocketBodyNormalization;
fn normalized_continuation(
event: &serde_json::Value,
normalization: &ResponsesWebSocketBodyNormalization,
) -> serde_json::Value {
let outbound = normalize_followup_response_create(event, "provider-model", normalization)
.expect("continuation should normalize");
serde_json::from_str(&outbound).expect("normalized event should be JSON")
}
#[test]
fn continuation_keeps_protocol_state_that_provider_normalization_strips() {
// `previous_response_id` is on the Codex unsupported-field list, so
// normalization removes it — yet it is what continues the chain. If
// this regresses, every continuation turn silently starts a new one.
let event = json!({
"type": "response.create",
"model": "public-model",
"previous_response_id": "resp_123",
"input": [],
"stream": true,
"background": true,
});
let normalized = normalized_continuation(
&event,
&ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_provider_type_for_tests("codex"),
);
assert_eq!(normalized["type"], "response.create");
assert_eq!(normalized["previous_response_id"], "resp_123");
assert_eq!(normalized["model"], "provider-model");
assert!(normalized.get("stream").is_none());
assert!(normalized.get("background").is_none());
}
#[test]
fn continuation_strips_fields_the_codex_backend_rejects() {
// The point of the fix: before it, turns 2..N reached Codex with the
// client's raw body, so a `temperature` that turn 1 had stripped would
// be rejected upstream. This also proves normalization really runs
// rather than silently falling back to the unmodified event.
let event = json!({
"type": "response.create",
"model": "public-model",
"previous_response_id": "resp_123",
"temperature": 0.7,
"top_p": 0.9,
"input": [],
});
let normalized = normalized_continuation(
&event,
&ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_provider_type_for_tests("codex"),
);
assert!(normalized.get("temperature").is_none());
assert!(normalized.get("top_p").is_none());
assert_eq!(normalized["store"], false);
// ...and the protocol state survives the same pass.
assert_eq!(normalized["previous_response_id"], "resp_123");
}
#[test]
fn continuation_keeps_a_warmup_generate_flag() {
let event = json!({
"type": "response.create",
"model": "public-model",
"previous_response_id": "resp_123",
"generate": false,
"input": [],
});
let normalized = normalized_continuation(
&event,
&ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_provider_type_for_tests("codex"),
);
assert_eq!(normalized["generate"], false);
}
#[test]
fn continuation_applies_the_model_directive_patch_the_binding_turn_received() {
let event = json!({
"type": "response.create",
"model": "public-model",
"previous_response_id": "resp_123",
"input": [],
});
let normalized = normalized_continuation(
&event,
&ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_model_directive_patch_for_tests(json!({"reasoning": {"effort": "high"}})),
);
assert_eq!(normalized["reasoning"]["effort"], "high");
}
#[test]
fn continuation_still_forces_the_bound_provider_model() {
let event = json!({
"type": "response.create",
"model": "some-other-model",
"previous_response_id": "resp_123",
"input": [],
});
let normalized = normalized_continuation(
&event,
&ResponsesWebSocketBodyNormalization::for_tests("provider-model"),
);
assert_eq!(normalized["model"], "provider-model");
}
#[test]
fn a_continuation_that_is_not_a_response_create_is_rejected() {
let normalization = ResponsesWebSocketBodyNormalization::for_tests("provider-model");
assert!(normalize_followup_response_create(
&json!({"type": "response.cancel"}),
"provider-model",
&normalization,
)
.is_err());
assert!(normalize_followup_response_create(
&json!("not an object"),
"provider-model",
&normalization,
)
.is_err());
}
#[test]
fn previous_response_id_is_protocol_state_even_when_not_a_string() {
assert!(response_create_has_previous_response_id(
&json!({"previous_response_id": 42})
));
assert!(!response_create_has_previous_response_id(
&json!({"previous_response_id": null})
));
assert!(!response_create_has_previous_response_id(&json!({})));
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,141 @@
//! Mutable state owned by one Responses WebSocket connection.
//!
//! The session loop is intentionally kept separate from these containers. A
//! connection may survive many `response.create` turns, while the turn
//! lifecycle and upstream binding are replaced independently.
use serde_json::Value;
use std::collections::{BTreeMap, BTreeSet};
use tokio::task::JoinHandle;
use super::adapter::{ResponsesWebSocketDrainDirective, ResponsesWebSocketProtocolAdapter};
use super::binding::UpstreamBindingIdentity;
use super::lifecycle::ActiveResponsesWebSocketTurn;
use super::request::response_create_has_previous_response_id;
use crate::ai_serving::{AiExecutionDecision, ResponsesWebSocketBodyNormalization};
const EXHAUSTED_KEY_EXCLUSION_FALLBACK_SECONDS: u64 = 300;
/// All mutable state associated with the physical upstream connection.
pub(super) struct BoundResponsesConnection {
pub(super) upstream: Option<wreq::ws::WebSocket>,
pub(super) adapter: &'static dyn ResponsesWebSocketProtocolAdapter,
pub(super) client_model: String,
pub(super) provider_model: String,
pub(super) response_in_flight: bool,
pub(super) decision_template: AiExecutionDecision,
/// Reproduces this binding's provider-body normalization for continuation
/// turns, which must not re-enter the planner. Replaced whenever the
/// binding or its decision is replaced.
pub(super) body_normalization: ResponsesWebSocketBodyNormalization,
pub(super) binding_identity: UpstreamBindingIdentity,
pub(super) active_turn: Option<ActiveResponsesWebSocketTurn>,
pub(super) active_response_create: Option<ActiveResponsesWebSocketRequest>,
pub(super) next_turn_index: u64,
pub(super) upstream_response_headers: BTreeMap<String, String>,
pub(super) pending_adapter_drain: Option<ResponsesWebSocketDrainDirective>,
pub(super) pending_adapter_observation: Option<JoinHandle<()>>,
pub(super) exhausted_exclusions: ExhaustedResponsesWebSocketExclusions,
pub(super) pending_turn_finalization: Option<JoinHandle<()>>,
}
/// Connection-local fallback in addition to the distributed account breaker.
/// A key and its provider account are excluded until the upstream's reset
/// deadline (or a short fallback when the terminal payload lacks one), so an
/// unusually long-lived client socket does not keep it unavailable after the
/// quota has recovered.
#[derive(Debug, Default)]
pub(super) struct ExhaustedResponsesWebSocketExclusions {
expires_at_by_key: BTreeMap<String, u64>,
expires_at_by_codex_account: BTreeMap<String, u64>,
}
impl ExhaustedResponsesWebSocketExclusions {
pub(super) fn exclude(
&mut self,
key_id: String,
codex_account_id: Option<String>,
reset_at_unix_secs: Option<u64>,
now_unix_secs: u64,
) -> u64 {
self.prune(now_unix_secs);
let requested_expiry = reset_at_unix_secs
.filter(|reset_at| *reset_at > now_unix_secs)
.unwrap_or_else(|| {
now_unix_secs.saturating_add(EXHAUSTED_KEY_EXCLUSION_FALLBACK_SECONDS)
});
let expiry = self
.expires_at_by_key
.entry(key_id)
.and_modify(|existing| *existing = (*existing).max(requested_expiry))
.or_insert(requested_expiry);
if let Some(account_id) = codex_account_id {
self.expires_at_by_codex_account
.entry(account_id)
.and_modify(|existing| *existing = (*existing).max(requested_expiry))
.or_insert(requested_expiry);
}
*expiry
}
pub(super) fn codex_account_ids(&mut self, now_unix_secs: u64) -> BTreeSet<String> {
self.prune(now_unix_secs);
self.expires_at_by_codex_account.keys().cloned().collect()
}
pub(super) fn key_ids(&mut self, now_unix_secs: u64) -> BTreeSet<String> {
self.prune(now_unix_secs);
self.expires_at_by_key.keys().cloned().collect()
}
pub(super) fn len(&mut self, now_unix_secs: u64) -> usize {
self.prune(now_unix_secs);
self.expires_at_by_key.len() + self.expires_at_by_codex_account.len()
}
fn prune(&mut self, now_unix_secs: u64) {
self.expires_at_by_key
.retain(|_, expires_at| *expires_at > now_unix_secs);
self.expires_at_by_codex_account
.retain(|_, expires_at| *expires_at > now_unix_secs);
}
}
#[derive(Debug, Clone)]
pub(super) struct ActiveResponsesWebSocketRequest {
pub(super) client_event: Value,
pub(super) turn_index: u64,
pub(super) logical_turn_id: String,
pub(super) turn_attempt: u32,
pub(super) retry_attempted: bool,
pub(super) retry_unsafe_reason: Option<&'static str>,
}
impl ActiveResponsesWebSocketRequest {
pub(super) fn new(client_event: Value, turn_index: u64, logical_turn_id: String) -> Self {
Self {
client_event,
turn_index,
logical_turn_id,
turn_attempt: 1,
retry_attempted: false,
retry_unsafe_reason: None,
}
}
pub(super) fn quota_retry_block_reason(&self) -> Option<&'static str> {
if self.retry_attempted {
Some("quota_retry_already_attempted")
} else if let Some(reason) = self.retry_unsafe_reason {
Some(reason)
} else if response_create_has_previous_response_id(&self.client_event) {
Some("previous_response_id")
} else {
None
}
}
pub(super) fn mark_retry_unsafe(&mut self, reason: &'static str) {
self.retry_unsafe_reason.get_or_insert(reason);
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,113 @@
//! Physical upstream WebSocket binding and transport helpers.
use serde_json::Value;
use wreq::ws::message::Message as WreqWsMessage;
use super::adapter::ResponsesWebSocketProtocolAdapter;
use super::binding::{UpstreamBindingIdentity, UpstreamBindingIdentityError};
use super::request::planned_response_create_event;
use super::state::{BoundResponsesConnection, ExhaustedResponsesWebSocketExclusions};
use crate::ai_serving::{AiExecutionDecision, ResponsesWebSocketBodyNormalization};
use crate::handlers::proxy::websocket::session::RESPONSES_WEBSOCKET_SESSION_LIMITS;
use crate::handlers::proxy::websocket::transport::{
close_upstream_socket, connect_upstream_websocket, send_upstream_message,
};
pub(super) async fn bind_responses_upstream(
decision: &AiExecutionDecision,
normalization: ResponsesWebSocketBodyNormalization,
initial_event: &Value,
adapter: &'static dyn ResponsesWebSocketProtocolAdapter,
) -> Result<BoundResponsesConnection, &'static str> {
let binding_identity =
UpstreamBindingIdentity::from_decision(adapter, decision).map_err(|error| match error {
UpstreamBindingIdentityError::MissingUpstreamUrl => {
adapter.upstream_errors().upstream_url_missing
}
UpstreamBindingIdentityError::InvalidUpstreamUrl => {
adapter.upstream_errors().upstream_url_invalid
}
UpstreamBindingIdentityError::InvalidHandshakeHeaders => {
adapter.upstream_errors().headers_invalid
}
})?;
let mut upstream = connect_upstream_websocket(
decision,
RESPONSES_WEBSOCKET_SESSION_LIMITS,
adapter.upstream_errors(),
)
.await?;
let first_event = planned_response_create_event(decision, initial_event)?;
send_upstream_message(&mut upstream.socket, WreqWsMessage::text(first_event))
.await
.map_err(|_| "responses_websocket_initial_send_failed")?;
let client_model = initial_event
.get("model")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.ok_or("responses_websocket_model_missing")?
.to_string();
let provider_model = decision
.provider_request_body
.as_ref()
.and_then(|body| body.get("model"))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.or_else(|| {
decision
.mapped_model
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
})
.ok_or("responses_websocket_mapped_model_missing")?
.to_string();
Ok(BoundResponsesConnection {
upstream: Some(upstream.socket),
adapter,
client_model,
provider_model,
response_in_flight: true,
decision_template: decision.clone(),
body_normalization: normalization,
binding_identity,
active_turn: None,
active_response_create: None,
next_turn_index: 2,
upstream_response_headers: upstream.response_headers,
pending_adapter_drain: None,
pending_adapter_observation: None,
exhausted_exclusions: ExhaustedResponsesWebSocketExclusions::default(),
pending_turn_finalization: None,
})
}
pub(super) async fn receive_optional_upstream(
upstream: &mut Option<wreq::ws::WebSocket>,
) -> Option<Result<WreqWsMessage, ()>> {
match upstream.as_mut() {
Some(upstream) => upstream.recv().await.map(|message| message.map_err(|_| ())),
None => std::future::pending().await,
}
}
pub(super) async fn close_bound_upstream(bound: &mut BoundResponsesConnection) {
if let Some(mut upstream) = bound.upstream.take() {
close_upstream_socket(&mut upstream, None).await;
}
}
pub(super) fn decision_reuses_bound_upstream(
bound: &BoundResponsesConnection,
adapter: &'static dyn ResponsesWebSocketProtocolAdapter,
decision: &AiExecutionDecision,
) -> bool {
bound.upstream.is_some()
&& UpstreamBindingIdentity::from_decision(adapter, decision)
.map(|identity| bound.binding_identity == identity)
.unwrap_or(false)
}
@@ -0,0 +1,48 @@
//! Connection-scoped limits and primitives shared by AI WebSocket sessions.
use std::time::{Duration, Instant};
/// The public Responses WebSocket contract is intentionally bounded so a
/// single active socket cannot retain gateway resources indefinitely.
#[derive(Debug, Clone, Copy)]
pub(crate) struct WebSocketSessionLimits {
pub(crate) max_frame_size: usize,
pub(crate) max_message_size: usize,
pub(crate) initial_message_timeout: Duration,
pub(crate) max_connection_duration: Duration,
}
pub(crate) const RESPONSES_WEBSOCKET_SESSION_LIMITS: WebSocketSessionLimits =
WebSocketSessionLimits {
max_frame_size: 16 << 20,
max_message_size: 16 << 20,
initial_message_timeout: Duration::from_secs(60),
max_connection_duration: Duration::from_secs(60 * 60),
};
/// A peer that stops draining its receive window must not be able to pin the
/// relay loop. Session loops await socket writes inside a `tokio::select!`,
/// so an unbounded write also suspends the connection and per-turn deadlines
/// that would otherwise reclaim the upstream socket and the shared upstream
/// admission permits.
pub(crate) const RELAY_WRITE_TIMEOUT: Duration = Duration::from_secs(30);
/// Frames the gateway emits while tearing a session down are best-effort: the
/// session is ending either way, so an unresponsive peer must not delay
/// releasing the upstream.
pub(crate) const TEARDOWN_WRITE_TIMEOUT: Duration = Duration::from_secs(5);
pub(crate) const CLOSE_POLICY_VIOLATION: u16 = 1008;
pub(crate) const CLOSE_INTERNAL_ERROR: u16 = 1011;
pub(crate) const CLOSE_TRY_AGAIN: u16 = 1013;
pub(crate) const WEBSOCKET_LOG_TRANSPORT: &str = "websocket";
/// Waits for an optional per-turn deadline without allocating a timer when no
/// turn is active. The protocol adapter retains ownership of the deadline's
/// meaning and terminal outcome.
pub(crate) async fn wait_for_optional_deadline(deadline: Option<Instant>) {
match deadline {
Some(deadline) => tokio::time::sleep_until(tokio::time::Instant::from_std(deadline)).await,
None => std::future::pending::<()>().await,
}
}
@@ -0,0 +1,406 @@
//! Upstream WebSocket handshake and frame conversion utilities.
//!
//! These helpers intentionally do not parse messages. A protocol adapter is
//! responsible for deciding when and what to send, while this module owns the
//! HTTP-to-WebSocket transport conversion and provider transport profile.
use std::collections::BTreeMap;
use std::time::Duration;
use axum::extract::ws::{CloseFrame as AxumCloseFrame, Message as AxumWsMessage, WebSocket};
use axum::http::header::{
ACCEPT, ACCEPT_ENCODING, CONNECTION, CONTENT_ENCODING, CONTENT_LENGTH, CONTENT_TYPE, HOST,
TRANSFER_ENCODING, UPGRADE,
};
use axum::http::HeaderMap;
use futures_util::{SinkExt, TryFutureExt};
use serde_json::json;
use url::Url;
use wreq::ws::message::{CloseFrame as WreqCloseFrame, Message as WreqWsMessage};
use crate::ai_serving::AiExecutionDecision;
use crate::execution_runtime::transport::{
build_browser_wreq_client, build_request_headers, ExecutionTransportControls,
};
use crate::handlers::proxy::websocket::session::{
WebSocketSessionLimits, RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT,
};
#[derive(Clone, Copy)]
pub(crate) struct UpstreamWebSocketErrorCodes {
pub(crate) upstream_url_missing: &'static str,
pub(crate) upstream_url_invalid: &'static str,
pub(crate) headers_invalid: &'static str,
pub(crate) client_build_failed: &'static str,
pub(crate) proxy_invalid: &'static str,
pub(crate) tunnel_proxy_unsupported: &'static str,
pub(crate) handshake_failed: &'static str,
pub(crate) upgrade_rejected: &'static str,
pub(crate) upgrade_failed: &'static str,
}
pub(crate) struct UpstreamWebSocketConnection {
pub(crate) socket: wreq::ws::WebSocket,
pub(crate) response_headers: BTreeMap<String, String>,
}
pub(crate) async fn connect_upstream_websocket(
decision: &AiExecutionDecision,
limits: WebSocketSessionLimits,
errors: UpstreamWebSocketErrorCodes,
) -> Result<UpstreamWebSocketConnection, &'static str> {
let upstream_url = decision
.upstream_url
.as_deref()
.ok_or(errors.upstream_url_missing)?;
let upstream_url = websocket_upstream_url(upstream_url, errors.upstream_url_invalid)?;
let headers =
websocket_handshake_headers(&decision.provider_request_headers, errors.headers_invalid)?;
let client = build_websocket_client(decision, errors)?;
let response = client
.websocket(upstream_url.as_str())
.headers(headers)
.max_frame_size(limits.max_frame_size)
.max_message_size(limits.max_message_size)
.send()
.await
.map_err(|_| errors.handshake_failed)?;
if response.status().as_u16() != 101 {
return Err(errors.upgrade_rejected);
}
let response_headers = websocket_response_headers(response.headers());
let socket = response
.into_websocket()
.await
.map_err(|_| errors.upgrade_failed)?;
Ok(UpstreamWebSocketConnection {
socket,
response_headers,
})
}
fn websocket_response_headers(headers: &HeaderMap) -> BTreeMap<String, String> {
headers
.iter()
.filter_map(|(name, value)| {
value
.to_str()
.ok()
.map(|value| (name.as_str().to_string(), value.to_string()))
})
.collect()
}
pub(crate) fn websocket_upstream_url(
raw: &str,
invalid_code: &'static str,
) -> Result<Url, &'static str> {
let mut url = Url::parse(raw).map_err(|_| invalid_code)?;
if url.host_str().is_none() || !url.username().is_empty() || url.password().is_some() {
return Err(invalid_code);
}
let websocket_scheme = match url.scheme() {
"https" => "wss",
"http" => "ws",
"wss" | "ws" => return Ok(url),
_ => return Err(invalid_code),
};
url.set_scheme(websocket_scheme).map_err(|_| invalid_code)?;
Ok(url)
}
pub(crate) fn websocket_handshake_headers(
provider_headers: &BTreeMap<String, String>,
invalid_code: &'static str,
) -> Result<HeaderMap, &'static str> {
let mut headers =
build_request_headers(provider_headers, None, false).map_err(|_| invalid_code)?;
for header in [
ACCEPT,
ACCEPT_ENCODING,
CONNECTION,
CONTENT_ENCODING,
CONTENT_LENGTH,
CONTENT_TYPE,
HOST,
TRANSFER_ENCODING,
UPGRADE,
] {
headers.remove(header);
}
Ok(headers)
}
fn build_websocket_client(
decision: &AiExecutionDecision,
errors: UpstreamWebSocketErrorCodes,
) -> Result<wreq::Client, &'static str> {
let timeouts = websocket_timeouts(decision);
if let Some(profile) = decision.transport_profile.as_ref() {
return build_browser_wreq_client(
timeouts.as_ref(),
decision.proxy.as_ref(),
profile,
ExecutionTransportControls::default(),
false,
)
.map_err(|_| errors.client_build_failed);
}
let mut builder = wreq::Client::builder();
if let Some(connect_ms) = timeouts.as_ref().and_then(|timeouts| timeouts.connect_ms) {
builder = builder.connect_timeout(Duration::from_millis(connect_ms));
}
if let Some(proxy) = decision
.proxy
.as_ref()
.filter(|proxy| proxy.enabled != Some(false))
{
if let Some(proxy_url) = proxy
.url
.as_deref()
.map(str::trim)
.filter(|url| !url.is_empty())
{
let proxy = wreq::Proxy::all(proxy_url).map_err(|_| errors.proxy_invalid)?;
builder = builder.proxy(proxy);
} else if proxy.node_id.is_some() || proxy.mode.as_deref() == Some("tunnel") {
return Err(errors.tunnel_proxy_unsupported);
}
}
builder.build().map_err(|_| errors.client_build_failed)
}
pub(crate) fn websocket_timeouts(
decision: &AiExecutionDecision,
) -> Option<aether_contracts::ExecutionTimeouts> {
let mut timeouts = decision.timeouts.clone()?;
timeouts.read_ms = None;
timeouts.first_byte_ms = None;
timeouts.total_ms = None;
Some(timeouts)
}
/// Why a frame did not reach its peer. A timeout is reported separately from
/// a socket error because the two describe different peers: one has gone away,
/// the other is still connected but has stopped reading.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum WebSocketWriteError {
Failed,
TimedOut,
}
impl WebSocketWriteError {
pub(crate) const fn as_str(self) -> &'static str {
match self {
Self::Failed => "write_failed",
Self::TimedOut => "write_timeout",
}
}
}
/// Relays one frame to the client under [`RELAY_WRITE_TIMEOUT`].
pub(crate) async fn send_client_message(
client_socket: &mut WebSocket,
message: AxumWsMessage,
) -> Result<(), WebSocketWriteError> {
bounded_send(
RELAY_WRITE_TIMEOUT,
client_socket.send(message).map_err(|_| ()),
)
.await
}
/// Sends one frame to the upstream under [`RELAY_WRITE_TIMEOUT`].
pub(crate) async fn send_upstream_message(
upstream: &mut wreq::ws::WebSocket,
message: WreqWsMessage,
) -> Result<(), WebSocketWriteError> {
bounded_send(RELAY_WRITE_TIMEOUT, upstream.send(message).map_err(|_| ())).await
}
/// Best-effort teardown write. The caller is already ending the session, so
/// the outcome only matters for keeping the wait bounded.
async fn send_teardown_message<F>(write: F)
where
F: std::future::Future<Output = Result<(), ()>>,
{
let _ = bounded_send(TEARDOWN_WRITE_TIMEOUT, write).await;
}
async fn bounded_send<F>(budget: Duration, write: F) -> Result<(), WebSocketWriteError>
where
F: std::future::Future<Output = Result<(), ()>>,
{
match tokio::time::timeout(budget, write).await {
Ok(Ok(())) => Ok(()),
Ok(Err(())) => Err(WebSocketWriteError::Failed),
Err(_) => Err(WebSocketWriteError::TimedOut),
}
}
/// Sends a WebSocket Close frame upstream without waiting on an unresponsive
/// provider. The socket is dropped by the caller either way.
pub(crate) async fn close_upstream_socket(
upstream: &mut wreq::ws::WebSocket,
frame: Option<WreqCloseFrame>,
) {
send_teardown_message(upstream.send(WreqWsMessage::Close(frame)).map_err(|_| ())).await;
}
pub(crate) fn upstream_message_to_client(message: WreqWsMessage) -> AxumWsMessage {
match message {
WreqWsMessage::Text(text) => AxumWsMessage::Text(text.to_string().into()),
WreqWsMessage::Binary(data) => AxumWsMessage::Binary(data),
WreqWsMessage::Ping(data) => AxumWsMessage::Ping(data),
WreqWsMessage::Pong(data) => AxumWsMessage::Pong(data),
WreqWsMessage::Close(frame) => AxumWsMessage::Close(frame.map(|frame| AxumCloseFrame {
code: frame.code.into(),
reason: frame.reason.to_string().into(),
})),
}
}
pub(crate) fn client_close_to_upstream(frame: Option<AxumCloseFrame>) -> Option<WreqCloseFrame> {
frame.map(|frame| WreqCloseFrame {
code: frame.code.into(),
reason: frame.reason.to_string().into(),
})
}
/// Builds a Responses WebSocket error event in the shape understood by the
/// official client implementations. The status is part of the event body,
/// not the WebSocket handshake, because the connection is already upgraded.
pub(crate) fn responses_websocket_error_event(
status: u16,
error_type: &str,
code: &str,
message: &str,
) -> serde_json::Value {
json!({
"type": "error",
"status": status,
"error": {
"type": error_type,
"code": code,
"message": message,
},
})
}
pub(crate) async fn send_responses_websocket_error(
client_socket: &mut WebSocket,
status: u16,
error_type: &str,
code: &str,
message: &str,
) {
let event = responses_websocket_error_event(status, error_type, code, message);
send_teardown_message(
client_socket
.send(AxumWsMessage::Text(event.to_string().into()))
.map_err(|_| ()),
)
.await;
}
pub(crate) async fn send_gateway_error(client_socket: &mut WebSocket, code: &str, message: &str) {
send_gateway_error_with_status(client_socket, 400, code, message).await;
}
pub(crate) async fn send_gateway_error_with_status(
client_socket: &mut WebSocket,
status: u16,
code: &str,
message: &str,
) {
send_responses_websocket_error(client_socket, status, "gateway_error", code, message).await;
}
pub(crate) async fn close_client_socket(client_socket: &mut WebSocket, code: u16, reason: &str) {
send_teardown_message(
client_socket
.send(AxumWsMessage::Close(Some(AxumCloseFrame {
code,
reason: reason.to_string().into(),
})))
.map_err(|_| ()),
)
.await;
}
#[cfg(test)]
mod tests {
use super::{
bounded_send, responses_websocket_error_event, websocket_upstream_url, WebSocketWriteError,
RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT,
};
use std::time::Duration;
#[tokio::test]
async fn a_peer_that_never_drains_its_window_times_out_instead_of_pinning_the_relay() {
let stalled = std::future::pending::<Result<(), ()>>();
let outcome = bounded_send(Duration::from_millis(1), stalled).await;
assert_eq!(outcome, Err(WebSocketWriteError::TimedOut));
}
#[tokio::test]
async fn a_socket_error_is_reported_separately_from_a_stalled_peer() {
let outcome = bounded_send(RELAY_WRITE_TIMEOUT, std::future::ready(Err(()))).await;
assert_eq!(outcome, Err(WebSocketWriteError::Failed));
assert_eq!(WebSocketWriteError::Failed.as_str(), "write_failed");
assert_eq!(WebSocketWriteError::TimedOut.as_str(), "write_timeout");
}
#[tokio::test]
async fn a_write_that_completes_within_its_budget_succeeds() {
let outcome = bounded_send(RELAY_WRITE_TIMEOUT, std::future::ready(Ok::<(), ()>(()))).await;
assert_eq!(outcome, Ok(()));
}
#[test]
fn teardown_writes_are_given_a_shorter_budget_than_relayed_frames() {
assert!(TEARDOWN_WRITE_TIMEOUT < RELAY_WRITE_TIMEOUT);
}
#[test]
fn builds_a_client_compatible_responses_error_event() {
let event = responses_websocket_error_event(
400,
"invalid_request_error",
"previous_response_not_found",
"Previous response was not found.",
);
assert_eq!(event["type"], "error");
assert_eq!(event["status"], 400);
assert_eq!(event["error"]["type"], "invalid_request_error");
assert_eq!(event["error"]["code"], "previous_response_not_found");
assert_eq!(
event["error"]["message"],
"Previous response was not found."
);
}
#[test]
fn maps_http_url_to_websocket_url_without_losing_path_or_query() {
let url = websocket_upstream_url(
"https://example.test/backend-api/codex/responses?x=1",
"invalid",
)
.expect("URL should be converted");
assert_eq!(
url.as_str(),
"wss://example.test/backend-api/codex/responses?x=1"
);
}
#[test]
fn rejects_upstream_url_with_credentials() {
assert!(websocket_upstream_url("https://token@example.test/responses", "invalid").is_err());
}
}
@@ -479,6 +479,7 @@ fn build_users_me_usage_record_payload(
"response_time_ms": item.response_time_ms,
"first_byte_time_ms": item.first_byte_time_ms,
"is_stream": item.is_stream,
"is_websocket": item.is_websocket(),
"upstream_is_stream": upstream_is_stream,
"client_requested_stream": client_is_stream,
"client_is_stream": client_is_stream,
@@ -562,6 +563,7 @@ fn build_users_me_usage_active_payload(item: &StoredRequestUsageAudit) -> serde_
"api_format": item.api_format,
"endpoint_api_format": item.endpoint_api_format,
"is_stream": item.is_stream,
"is_websocket": item.is_websocket(),
"upstream_is_stream": upstream_is_stream,
"client_requested_stream": client_is_stream,
"client_is_stream": client_is_stream,
@@ -1723,6 +1725,23 @@ mod tests {
assert_eq!(active["reasoning_effort"], "max");
}
#[test]
fn user_usage_payloads_expose_websocket_transport() {
let item = StoredRequestUsageAudit {
request_metadata: Some(json!({
"websocket_mode": true,
"websocket_transport": "responses",
})),
..sample_usage("completed")
};
let record = build_users_me_usage_record_payload(&item, false, &BTreeMap::new(), false);
let active = build_users_me_usage_active_payload(&item);
assert_eq!(record["is_websocket"], true);
assert_eq!(active["is_websocket"], true);
}
#[test]
fn user_usage_active_override_uses_terminal_candidate_latency() {
let candidate = sample_candidate(
@@ -967,6 +967,12 @@ fn build_codex_quota_status_snapshot(
let credits_unlimited = metadata
.get("credits_unlimited")
.and_then(admin_provider_quota_pure::coerce_json_bool);
let allowed = metadata
.get("allowed")
.and_then(admin_provider_quota_pure::coerce_json_bool);
let limit_reached = metadata
.get("limit_reached")
.and_then(admin_provider_quota_pure::coerce_json_bool);
let reset_credits = build_codex_reset_credits_status_snapshot(metadata, observed_at_unix_secs);
let windows = [
@@ -996,6 +1002,8 @@ fn build_codex_quota_status_snapshot(
&& credits_has_credits.is_none()
&& credits_balance.is_none()
&& credits_unlimited.is_none()
&& allowed.is_none()
&& limit_reached.is_none()
&& reset_credits.is_none()
&& observed_at_unix_secs.is_none()
{
@@ -1031,11 +1039,19 @@ fn build_codex_quota_status_snapshot(
.filter_map(admin_provider_quota_pure::coerce_json_u64)
.min();
let reset_at = quota_windows_min_reset_at(&primary_windows);
let exhausted_by_credits = primary_windows.is_empty()
let explicitly_blocked = allowed == Some(false) || limit_reached == Some(true);
let explicitly_available =
!explicitly_blocked && (allowed == Some(true) || limit_reached == Some(false));
let exhausted_by_credits = !explicitly_available
&& primary_windows.is_empty()
&& credits_unlimited != Some(true)
&& credits_has_credits == Some(false);
let exhausted_by_window = usage_ratio.is_some_and(|value| value >= 1.0 - 1e-6);
let exhausted = exhausted_by_credits || exhausted_by_window;
let exhausted_by_window =
!explicitly_available && usage_ratio.is_some_and(|value| value >= 1.0 - 1e-6);
let exhausted_by_signal = admin_provider_quota_pure::codex_rate_limit_metadata_exhausted(
&Value::Object(metadata.clone()),
);
let exhausted = exhausted_by_signal || exhausted_by_credits || exhausted_by_window;
let mut credits = Map::new();
if let Some(value) = credits_has_credits {
@@ -1048,7 +1064,9 @@ fn build_codex_quota_status_snapshot(
credits.insert("unlimited".to_string(), json!(value));
}
let reason = if exhausted_by_credits {
let reason = if exhausted_by_signal {
Some("上游已拒绝继续使用该账号")
} else if exhausted_by_credits {
Some("无可用积分")
} else if exhausted_by_window {
Some("额度窗口已耗尽")
@@ -1071,6 +1089,8 @@ fn build_codex_quota_status_snapshot(
"reset_at": reset_at,
"reset_seconds": reset_seconds,
"plan_type": plan_type,
"allowed": allowed,
"limit_reached": limit_reached,
"credits": if credits.is_empty() {
Value::Null
} else {
@@ -3833,6 +3853,28 @@ mod tests {
);
}
#[test]
fn sync_provider_key_quota_status_snapshot_honors_codex_limit_signal() {
let payload = sync_provider_key_quota_status_snapshot(
None,
"codex",
Some(&json!({
"codex": {
"updated_at": 1_775_800_000u64,
"allowed": false,
"limit_reached": true
}
})),
"websocket_response_body",
)
.expect("explicit Codex limit signal should build a quota snapshot");
assert_eq!(payload.pointer("/quota/code"), Some(&json!("exhausted")));
assert_eq!(payload.pointer("/quota/exhausted"), Some(&json!(true)));
assert_eq!(payload.pointer("/quota/allowed"), Some(&json!(false)));
assert_eq!(payload.pointer("/quota/limit_reached"), Some(&json!(true)));
}
#[test]
fn sync_provider_key_quota_status_snapshot_drops_codex_usage_state_when_window_resets() {
let current_status_snapshot = json!({
+1
View File
@@ -90,6 +90,7 @@ pub(crate) use self::ai_serving::api::{
EXECUTION_RUNTIME_SYNC_DECISION_ACTION, GEMINI_FILES_DOWNLOAD_PLAN_KIND,
OPENAI_VIDEO_CONTENT_PLAN_KIND,
};
pub use self::ai_serving::api::{CODEX_CLIENT_ORIGINATOR, CODEX_CLIENT_USER_AGENT};
pub(crate) use self::ai_serving::{
AiExecutionDecision, AiExecutionPlanPayload, AiStreamAttempt, AiSyncAttempt,
};
+51 -1
View File
@@ -1330,9 +1330,21 @@ struct Args {
#[arg(long, env = "AETHER_GATEWAY_MAX_IN_FLIGHT_REQUESTS")]
max_in_flight_requests: Option<usize>,
/// Maximum number of long-lived public WebSocket connections. When unset,
/// this follows `max_in_flight_requests` while remaining an independent
/// gate. Set `AETHER_GATEWAY_MAX_WEBSOCKET_CONNECTIONS` to override it.
#[arg(long, env = "AETHER_GATEWAY_MAX_WEBSOCKET_CONNECTIONS")]
max_websocket_connections: Option<usize>,
#[arg(long, env = "AETHER_GATEWAY_DISTRIBUTED_REQUEST_LIMIT")]
distributed_request_limit: Option<usize>,
/// Optional distributed limit for long-lived WebSocket connections. When
/// omitted, the distributed request limit is reused; set it to 0 to keep
/// WebSocket admission local-only.
#[arg(long, env = "AETHER_GATEWAY_DISTRIBUTED_WEBSOCKET_CONNECTION_LIMIT")]
distributed_websocket_connection_limit: Option<usize>,
#[arg(long, env = "AETHER_GATEWAY_DISTRIBUTED_REQUEST_REDIS_URL")]
distributed_request_redis_url: Option<String>,
@@ -1820,6 +1832,15 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
.max_in_flight_requests
.filter(|limit| *limit > 0)
.unwrap_or_else(automatic_gateway_request_concurrency);
let websocket_connection_limit = args
.max_websocket_connections
.filter(|limit| *limit > 0)
.unwrap_or(request_concurrency_limit);
let distributed_websocket_connection_limit = match args.distributed_websocket_connection_limit {
Some(limit) if limit > 0 => Some(limit),
Some(_) => None,
None => args.distributed_request_limit.filter(|limit| *limit > 0),
};
let usage_queue_request_concurrency_hint = usage_queue_request_concurrency_hint(
Some(request_concurrency_limit),
args.distributed_request_limit,
@@ -1941,6 +1962,14 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
"auto"
},
distributed_request_limit = args.distributed_request_limit.unwrap_or_default(),
max_websocket_connections = websocket_connection_limit,
max_websocket_connections_source = if args.max_websocket_connections.is_some() {
"explicit"
} else {
"request_concurrency_fallback"
},
distributed_websocket_connection_limit =
distributed_websocket_connection_limit.unwrap_or_default(),
distributed_request_redis_configured = args
.distributed_request_redis_url
.as_deref()
@@ -1995,7 +2024,9 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
{
state = state.with_video_task_store_path(path)?;
}
state = state.with_request_concurrency_limit(request_concurrency_limit);
state = state
.with_request_concurrency_limit(request_concurrency_limit)
.with_websocket_connection_limit(websocket_connection_limit);
if let Some(limit) = args.distributed_request_limit.filter(|limit| *limit > 0) {
let distributed_gate = state
.runtime_state()
@@ -2013,6 +2044,23 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
})?;
state = state.with_distributed_request_concurrency_gate(distributed_gate);
}
if let Some(limit) = distributed_websocket_connection_limit {
let distributed_gate = state
.runtime_state()
.semaphore(
"gateway_websocket_connections_distributed",
limit,
RuntimeSemaphoreConfig {
lease_ttl_ms: args.distributed_request_lease_ttl_ms.max(1),
renew_interval_ms: args.distributed_request_renew_interval_ms.max(1),
command_timeout_ms: Some(args.distributed_request_command_timeout_ms.max(1)),
},
)
.map_err(|err| {
std::io::Error::new(std::io::ErrorKind::InvalidInput, err.to_string())
})?;
state = state.with_distributed_websocket_connection_gate(distributed_gate);
}
if matches!(args.deployment_topology, DeploymentTopologyArg::MultiNode)
&& !state.has_usage_data_writer()
{
@@ -2531,7 +2579,9 @@ mod tests {
video_task_poller_batch_size: 32,
video_task_store_path: None,
max_in_flight_requests: None,
max_websocket_connections: None,
distributed_request_limit: None,
distributed_websocket_connection_limit: None,
distributed_request_redis_url: None,
distributed_request_redis_key_prefix: None,
distributed_request_lease_ttl_ms: 30_000,
@@ -0,0 +1,374 @@
//! Short-lived runtime circuit breaker for exhausted Codex accounts.
//!
//! Persisted provider-key quota snapshots remain the durable source of truth.
//! This module closes the interval between receiving a definitive WebSocket
//! `usage_limit_reached` event and every scheduler/cache replica observing the
//! persisted snapshot. The account-scoped entry also protects pools that
//! contain more than one catalog key for the same ChatGPT account.
use std::collections::BTreeMap;
use serde_json::{json, Map, Value};
use sha2::{Digest, Sha256};
use tracing::{info, warn};
use crate::clock::current_unix_secs;
use crate::{AppState, GatewayError};
const CODEX_QUOTA_BREAKER_KEY_PREFIX: &str = "aether:codex:quota-breaker:v1";
const CODEX_QUOTA_BREAKER_FALLBACK_TTL_SECONDS: u64 = 300;
const CODEX_QUOTA_BREAKER_MAX_TTL_SECONDS: u64 = 31 * 24 * 60 * 60;
/// Installs immediate runtime exclusions for a definitive Codex quota
/// exhaustion signal. When RuntimeState is backed by Redis the exclusions are
/// shared by every gateway node; the in-memory backend still protects the
/// current node.
pub(crate) async fn install_codex_quota_exhaustion_breaker(
state: &AppState,
report_context: Option<&Value>,
quota_metadata: &Value,
source: &str,
) -> Result<bool, GatewayError> {
if !aether_admin::provider::quota::codex_rate_limit_metadata_exhausted(quota_metadata) {
return Ok(false);
}
let keys = codex_quota_breaker_keys_from_report_context(report_context);
if keys.is_empty() {
return Ok(false);
}
let now_unix_secs = current_unix_secs();
let (ttl_seconds, reset_at_unix_secs) = codex_quota_breaker_ttl(quota_metadata, now_unix_secs);
let value = json!({
"version": 1,
"observed_at": now_unix_secs,
"reset_at": reset_at_unix_secs,
"source": source,
})
.to_string();
for key in &keys {
state.runtime_kv_setex(key, &value, ttl_seconds).await?;
}
info!(
event_name = "codex_account_quota_breaker_installed",
log_type = "event",
scope_count = keys.len(),
ttl_seconds,
reset_at_unix_secs = ?reset_at_unix_secs,
source,
"gateway installed immediate Codex quota exhaustion exclusions"
);
Ok(true)
}
/// Returns whether a planned Codex request is temporarily blocked by a
/// definitive account quota signal that has not yet been observed in the
/// durable provider catalog.
pub(crate) async fn codex_quota_breaker_blocks_candidate(
state: &AppState,
provider_type: Option<&str>,
key_id: Option<&str>,
provider_request_headers: &BTreeMap<String, String>,
) -> Result<bool, GatewayError> {
if !provider_type.is_some_and(|value| value.trim().eq_ignore_ascii_case("codex")) {
return Ok(false);
}
for key in codex_quota_breaker_keys(
key_id,
codex_account_id_from_headers(provider_request_headers),
) {
if state.runtime_kv_exists(&key).await? {
return Ok(true);
}
}
Ok(false)
}
fn codex_quota_breaker_keys_from_report_context(report_context: Option<&Value>) -> Vec<String> {
let key_id = report_context
.and_then(|context| context.get("key_id"))
.and_then(Value::as_str);
let account_id = report_context
.and_then(|context| context.get("provider_request_headers"))
.and_then(Value::as_object)
.and_then(account_id_from_header_object);
codex_quota_breaker_keys(key_id, account_id)
}
fn codex_quota_breaker_keys(key_id: Option<&str>, account_id: Option<&str>) -> Vec<String> {
let mut keys = Vec::with_capacity(2);
if let Some(account_id) = normalize_identifier(account_id) {
keys.push(codex_quota_breaker_runtime_key("account", account_id));
}
if let Some(key_id) = normalize_identifier(key_id) {
keys.push(codex_quota_breaker_runtime_key("key", key_id));
}
keys
}
fn codex_quota_breaker_runtime_key(scope: &str, identifier: &str) -> String {
format!(
"{CODEX_QUOTA_BREAKER_KEY_PREFIX}:{scope}:{}",
opaque_identifier(identifier)
)
}
fn opaque_identifier(identifier: &str) -> String {
let digest = Sha256::digest(identifier.as_bytes());
let mut encoded = String::with_capacity(digest.len().saturating_mul(2));
for byte in digest {
use std::fmt::Write as _;
let _ = write!(encoded, "{byte:02x}");
}
encoded
}
fn normalize_identifier(value: Option<&str>) -> Option<&str> {
value.map(str::trim).filter(|value| !value.is_empty())
}
pub(crate) fn codex_account_id_from_headers(headers: &BTreeMap<String, String>) -> Option<&str> {
headers.iter().find_map(|(name, value)| {
name.trim()
.eq_ignore_ascii_case("chatgpt-account-id")
.then_some(value.as_str())
.and_then(|value| normalize_identifier(Some(value)))
})
}
fn account_id_from_header_object(headers: &Map<String, Value>) -> Option<&str> {
headers.iter().find_map(|(name, value)| {
name.trim()
.eq_ignore_ascii_case("chatgpt-account-id")
.then(|| value.as_str())
.flatten()
.and_then(|value| normalize_identifier(Some(value)))
})
}
fn codex_quota_breaker_ttl(quota_metadata: &Value, now_unix_secs: u64) -> (u64, Option<u64>) {
let reset_at = codex_quota_exhaustion_reset_at(quota_metadata, now_unix_secs);
let ttl_seconds = reset_at
.and_then(|reset_at| reset_at.checked_sub(now_unix_secs))
.filter(|ttl| *ttl > 0)
.unwrap_or(CODEX_QUOTA_BREAKER_FALLBACK_TTL_SECONDS)
.clamp(1, CODEX_QUOTA_BREAKER_MAX_TTL_SECONDS);
(ttl_seconds, reset_at)
}
/// Returns the latest reset deadline required for a currently exhausted Codex
/// quota window. It is shared by the distributed breaker and the per-socket
/// retry exclusion so both stop excluding the account at the same time.
pub(crate) fn codex_quota_exhaustion_reset_at(
quota_metadata: &Value,
now_unix_secs: u64,
) -> Option<u64> {
let Some(metadata) = quota_metadata.as_object() else {
return None;
};
let exhausted_windows = ["primary", "secondary"]
.into_iter()
.filter(|prefix| codex_window_is_exhausted(metadata, prefix))
.collect::<Vec<_>>();
let prefixes = if exhausted_windows.is_empty() {
vec!["primary", "secondary"]
} else {
exhausted_windows
};
prefixes
.iter()
.filter_map(|prefix| codex_window_reset_at(metadata, prefix, now_unix_secs))
.filter(|reset_at| *reset_at > now_unix_secs)
.max()
}
fn codex_window_is_exhausted(metadata: &Map<String, Value>, prefix: &str) -> bool {
let used_percent = metadata
.get(&format!("{prefix}_used_percent"))
.and_then(aether_admin::provider::quota::coerce_json_f64);
used_percent.is_some_and(|used_percent| used_percent >= 100.0 - 1e-6)
}
fn codex_window_reset_at(
metadata: &Map<String, Value>,
prefix: &str,
observed_at_unix_secs: u64,
) -> Option<u64> {
metadata
.get(&format!("{prefix}_reset_at"))
.and_then(aether_admin::provider::quota::coerce_json_u64)
.filter(|reset_at| *reset_at > observed_at_unix_secs)
.or_else(|| {
metadata
.get(&format!("{prefix}_reset_after_seconds"))
.and_then(aether_admin::provider::quota::coerce_json_u64)
.and_then(|seconds| observed_at_unix_secs.checked_add(seconds))
})
}
pub(crate) fn log_codex_quota_breaker_install_failure(error: &GatewayError) {
warn!(
event_name = "codex_account_quota_breaker_install_failed",
log_type = "ops",
error = ?error,
"gateway could not install the immediate Codex quota exhaustion breaker"
);
}
pub(crate) fn log_codex_quota_breaker_check_failure(error: &GatewayError) {
warn!(
event_name = "codex_account_quota_breaker_check_failed",
log_type = "ops",
transport = "websocket",
websocket = true,
error = ?error,
"gateway could not check the immediate Codex quota exhaustion breaker; allowing candidate selection"
);
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use serde_json::json;
use super::{
codex_quota_breaker_blocks_candidate, codex_quota_breaker_keys,
codex_quota_breaker_keys_from_report_context, codex_quota_breaker_ttl,
codex_quota_exhaustion_reset_at, install_codex_quota_exhaustion_breaker,
};
use crate::AppState;
#[test]
fn account_scope_is_shared_across_catalog_keys() {
let first = codex_quota_breaker_keys(Some("key-first"), Some("account-123"));
let second = codex_quota_breaker_keys(Some("key-second"), Some("account-123"));
assert_eq!(first.first(), second.first());
assert_ne!(first.last(), second.last());
assert!(!first.first().is_some_and(|key| key.contains("account-123")));
}
#[test]
fn report_context_and_planned_headers_use_the_same_account_scope() {
let report_context = json!({
"key_id": "key-first",
"provider_request_headers": {
"ChatGPT-Account-ID": "account-123"
}
});
let context_keys = codex_quota_breaker_keys_from_report_context(Some(&report_context));
let planned_keys = codex_quota_breaker_keys(Some("key-second"), Some("account-123"));
assert_eq!(context_keys.first(), planned_keys.first());
}
#[test]
fn ttl_uses_the_exhausted_window_reset_deadline() {
let (ttl, reset_at) = codex_quota_breaker_ttl(
&json!({
"primary_used_percent": 100,
"primary_reset_at": 1_000,
"secondary_used_percent": 10,
"secondary_reset_at": 2_000,
}),
500,
);
assert_eq!(ttl, 500);
assert_eq!(reset_at, Some(1_000));
}
#[test]
fn ttl_falls_back_when_an_error_has_no_reset_metadata() {
let (ttl, reset_at) = codex_quota_breaker_ttl(&json!({"allowed": false}), 500);
assert_eq!(ttl, 300);
assert_eq!(reset_at, None);
}
#[test]
fn reset_deadline_uses_relative_metadata_when_absolute_reset_is_stale() {
assert_eq!(
codex_quota_exhaustion_reset_at(
&json!({
"primary_used_percent": 100,
"primary_reset_at": 999,
"primary_reset_after_seconds": 120,
}),
1_000,
),
Some(1_120)
);
}
#[test]
fn account_header_matching_is_case_insensitive() {
let headers =
BTreeMap::from([("CHATGPT-ACCOUNT-ID".to_string(), "account-123".to_string())]);
let keys = codex_quota_breaker_keys(
Some("key-first"),
headers.iter().find_map(|(name, value)| {
name.eq_ignore_ascii_case("chatgpt-account-id")
.then_some(value.as_str())
}),
);
assert_eq!(keys.len(), 2);
}
#[tokio::test]
async fn exhausted_account_immediately_blocks_a_different_catalog_key() {
let state = AppState::new().expect("gateway state should build");
let report_context = json!({
"key_id": "key-first",
"provider_request_headers": {
"ChatGPT-Account-ID": "account-123"
}
});
let quota_metadata = json!({
"allowed": false,
"limit_reached": true,
"primary_used_percent": 100,
"primary_reset_after_seconds": 60,
});
assert!(install_codex_quota_exhaustion_breaker(
&state,
Some(&report_context),
&quota_metadata,
"test",
)
.await
.expect("breaker installation should succeed"));
let same_account_other_key =
BTreeMap::from([("chatgpt-account-id".to_string(), "account-123".to_string())]);
assert!(codex_quota_breaker_blocks_candidate(
&state,
Some("codex"),
Some("key-second"),
&same_account_other_key,
)
.await
.expect("breaker lookup should succeed"));
let other_account =
BTreeMap::from([("chatgpt-account-id".to_string(), "account-456".to_string())]);
assert!(!codex_quota_breaker_blocks_candidate(
&state,
Some("codex"),
Some("key-second"),
&other_account,
)
.await
.expect("breaker lookup should succeed"));
}
}
+260 -12
View File
@@ -29,7 +29,8 @@ use tracing::warn;
use super::{
classify_failure_disposition, local_failover_error_message, project_local_adaptive_rate_limit,
project_local_adaptive_success, project_local_failure_health, project_local_key_circuit_closed,
project_local_key_circuit_failure, project_local_success_health, FailureScope,
project_local_key_circuit_failure, project_local_success_health,
resolve_local_failover_analysis_for_attempt, FailureScope, LocalFailoverAnalysis,
LocalFailoverClassification,
};
use crate::ai_serving::extract_pool_sticky_session_token;
@@ -325,6 +326,41 @@ pub(crate) fn spawn_local_oauth_success_effect(
});
}
/// Inputs for the terminal effects of a failed streaming attempt.
///
/// The status/body are deliberately supplied by the transport-specific caller:
/// a WebSocket terminal event may carry its own status and error body, while a
/// normal stream failure gets them from the HTTP response. Keeping this type at
/// the orchestration boundary prevents each transport from rebuilding the
/// health, adaptive, OAuth, and pool effect sequence independently.
#[derive(Debug, Clone, Copy)]
pub(crate) struct LocalStreamFailureEffect<'a> {
pub(crate) status_code: u16,
pub(crate) headers: &'a BTreeMap<String, String>,
pub(crate) response_text: Option<&'a str>,
pub(crate) stream_timeout: bool,
}
impl<'a> LocalStreamFailureEffect<'a> {
pub(crate) const fn new(
status_code: u16,
headers: &'a BTreeMap<String, String>,
response_text: Option<&'a str>,
) -> Self {
Self {
status_code,
headers,
response_text,
stream_timeout: false,
}
}
pub(crate) const fn with_stream_timeout(mut self) -> Self {
self.stream_timeout = true;
self
}
}
struct PoolFeedbackContext {
pool_config: AdminProviderPoolConfig,
sticky_session_token: Option<String>,
@@ -377,36 +413,157 @@ pub(crate) async fn apply_local_execution_effect(
}
LocalExecutionEffect::PoolSuccessSync { payload } => {
record_sync_pool_success_effect(state, context, payload).await;
release_pool_key_lease_effect(state, context).await;
release_local_pool_key_lease(state, context).await;
}
LocalExecutionEffect::PoolSuccessStream { payload } => {
record_stream_pool_success_effect(state, context, payload).await;
release_pool_key_lease_effect(state, context).await;
release_local_pool_key_lease(state, context).await;
}
LocalExecutionEffect::PoolError(effect) => {
record_pool_error_effect(state, context, effect).await;
release_pool_key_lease_effect(state, context).await;
release_local_pool_key_lease(state, context).await;
}
LocalExecutionEffect::PoolStreamTimeout => {
record_pool_stream_timeout_effect(state, context).await;
release_pool_key_lease_effect(state, context).await;
release_local_pool_key_lease(state, context).await;
}
}
}
async fn release_pool_key_lease_effect(state: &AppState, context: LocalExecutionEffectContext<'_>) {
/// Apply the provider/key effects shared by every successful streaming
/// transport. Usage persistence and request-candidate terminal status remain
/// owned by the report layer; this helper only projects execution health,
/// adaptive state, and pool feedback.
pub(crate) async fn apply_local_stream_success_effects(
state: &AppState,
context: LocalExecutionEffectContext<'_>,
payload: &GatewayStreamReportRequest,
) {
apply_local_execution_effect(
state,
context,
LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect),
)
.await;
apply_local_execution_effect(
state,
context,
LocalExecutionEffect::AdaptiveSuccess(LocalAdaptiveSuccessEffect),
)
.await;
apply_local_execution_effect(
state,
context,
LocalExecutionEffect::PoolSuccessStream { payload },
)
.await;
}
/// Apply the provider/key effects shared by every failed streaming attempt.
/// The returned analysis is the same failover classification used by the
/// normal stream runtime, allowing the caller to make a transport-specific
/// retry/close decision without re-running policy evaluation.
pub(crate) async fn apply_local_stream_failure_effects(
state: &AppState,
context: LocalExecutionEffectContext<'_>,
effect: LocalStreamFailureEffect<'_>,
) -> LocalFailoverAnalysis {
let analysis = resolve_local_failover_analysis_for_attempt(
state,
context.plan,
context.report_context,
effect.status_code,
effect.response_text,
)
.await;
if effect.stream_timeout {
apply_local_execution_effect(state, context, LocalExecutionEffect::PoolStreamTimeout).await;
}
apply_local_execution_effect(
state,
context,
LocalExecutionEffect::AttemptFailure(LocalAttemptFailureEffect {
status_code: effect.status_code,
classification: analysis.classification,
}),
)
.await;
apply_local_execution_effect(
state,
context,
LocalExecutionEffect::AdaptiveRateLimit(LocalAdaptiveRateLimitEffect {
status_code: effect.status_code,
classification: analysis.classification,
headers: Some(effect.headers),
}),
)
.await;
apply_local_execution_effect(
state,
context,
LocalExecutionEffect::HealthFailure(LocalHealthFailureEffect {
status_code: effect.status_code,
classification: analysis.classification,
}),
)
.await;
apply_local_execution_effect(
state,
context,
LocalExecutionEffect::OauthInvalidation(LocalOAuthInvalidationEffect {
status_code: effect.status_code,
response_text: effect.response_text,
}),
)
.await;
apply_local_execution_effect(
state,
context,
LocalExecutionEffect::PoolError(LocalPoolErrorEffect {
status_code: effect.status_code,
classification: analysis.classification,
headers: effect.headers,
error_body: effect.response_text,
}),
)
.await;
analysis
}
pub(crate) async fn release_local_pool_key_lease(
state: &AppState,
context: LocalExecutionEffectContext<'_>,
) {
let metadata = local_execution_candidate_metadata_from_report_context(context.report_context);
let Some(lease) = metadata.pool_key_lease else {
return;
};
release_pool_key_lease(state, &lease).await;
}
/// Releases a lease carried by a planned-but-not-started report context. This
/// path has no execution plan yet, so it intentionally omits candidate health
/// logging and only performs the distributed lock cleanup.
pub(crate) async fn release_pool_key_lease_from_report_context(
state: &AppState,
report_context: Option<&Value>,
) {
let metadata = local_execution_candidate_metadata_from_report_context(report_context);
let Some(lease) = metadata.pool_key_lease else {
return;
};
release_pool_key_lease(state, &lease).await;
}
async fn release_pool_key_lease(state: &AppState, lease: &aether_runtime_state::RuntimeLockLease) {
if let Err(err) =
release_admin_provider_pool_key_lease(state.runtime_state.as_ref(), &lease).await
release_admin_provider_pool_key_lease(state.runtime_state.as_ref(), lease).await
{
warn!(
error = ?err,
provider_id = %context.plan.provider_id,
key_id = %context.plan.key_id,
"gateway orchestration effects: failed to release pool key lease"
"gateway orchestration effects: failed to release a planned pool key lease"
);
}
}
@@ -1967,20 +2124,22 @@ mod tests {
use serde_json::{json, Value};
use super::{
apply_local_execution_effect, execution_plan_bearer_matches_transport,
apply_local_execution_effect, apply_local_stream_failure_effects,
apply_local_stream_success_effects, execution_plan_bearer_matches_transport,
local_candidate_failure_should_apply_key_effects,
local_candidate_failure_should_record_pool_error, pool_score_feedback_gate_allows,
pool_score_hard_state_for_status, resolve_pool_feedback_context,
LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect,
LocalExecutionEffect, LocalExecutionEffectContext, LocalHealthFailureEffect,
LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, LocalOAuthSuccessEffect,
LocalPoolErrorEffect, ProviderKeyEffectLockPool,
LocalPoolErrorEffect, LocalStreamFailureEffect, ProviderKeyEffectLockPool,
};
use crate::data::{GatewayDataConfig, GatewayDataState};
use crate::orchestration::{
apply_local_report_effect, LocalFailoverClassification, LocalReportEffect,
};
use crate::scheduler::affinity::SCHEDULER_AFFINITY_TTL;
use crate::usage::GatewayStreamReportRequest;
use crate::AppState;
use aether_scheduler_core::{
build_scheduler_affinity_cache_key_for_api_key_id,
@@ -2031,6 +2190,22 @@ mod tests {
plan
}
fn sample_stream_report() -> GatewayStreamReportRequest {
GatewayStreamReportRequest {
trace_id: "trace-stream-effects".to_string(),
report_kind: "openai_chat_stream_success".to_string(),
report_context: None,
status_code: 200,
headers: BTreeMap::new(),
provider_body_base64: None,
provider_body_state: None,
client_body_base64: None,
client_body_state: None,
terminal_summary: None,
telemetry: None,
}
}
#[test]
fn pool_score_feedback_gate_suppresses_repeated_success_writes() {
super::POOL_SCORE_FEEDBACK_GATE.clear();
@@ -2828,6 +3003,79 @@ mod tests {
.is_some());
}
#[tokio::test]
async fn stream_success_effect_helper_projects_health_and_scheduler_affinity() {
let state = AppState::new().expect("gateway state should build");
let plan = sample_plan();
let report_context = json!({
"api_key_id": "api-key-1",
"client_api_format": "openai:chat",
"model": "gpt-5",
});
let cache_key =
build_scheduler_affinity_cache_key_for_api_key_id("api-key-1", "openai:chat", "gpt-5")
.expect("scheduler affinity cache key should build");
let payload = sample_stream_report();
apply_local_stream_success_effects(
&state,
LocalExecutionEffectContext {
plan: &plan,
report_context: Some(&report_context),
},
&payload,
)
.await;
assert_eq!(
state.read_scheduler_affinity_target(cache_key.as_str(), SCHEDULER_AFFINITY_TTL),
Some(SchedulerAffinityTarget {
provider_id: "prov-1".to_string(),
endpoint_id: "ep-1".to_string(),
key_id: "key-1".to_string(),
})
);
}
#[tokio::test]
async fn stream_failure_effect_helper_returns_analysis_and_projects_health() {
let state = health_state();
let plan = sample_plan();
let headers = BTreeMap::new();
let analysis = apply_local_stream_failure_effects(
&state,
LocalExecutionEffectContext {
plan: &plan,
report_context: None,
},
LocalStreamFailureEffect::new(503, &headers, Some("upstream unavailable"))
.with_stream_timeout(),
)
.await;
assert_eq!(
analysis.classification,
LocalFailoverClassification::UseDefault
);
assert_eq!(analysis.decision.as_str(), "use_default");
let stored_key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
.await
.expect("provider catalog keys should load")
.into_iter()
.next()
.expect("stored key should exist");
assert_eq!(
stored_key
.health_by_format
.as_ref()
.and_then(|value| value.get("openai:chat"))
.and_then(|value| value.get("consecutive_failures"))
.and_then(Value::as_u64),
Some(1)
);
}
#[tokio::test]
async fn configured_stop_pattern_keeps_scheduler_affinity_cache() {
let state = AppState::new().expect("gateway state should build");
+15 -3
View File
@@ -7,6 +7,7 @@ use crate::AppState;
mod adaptive;
mod attempt;
mod classifier;
mod codex_quota_breaker;
mod effects;
mod health;
mod oauth_error;
@@ -32,11 +33,20 @@ pub(crate) use self::classifier::{
FailureTokenAction, LocalFailoverClassification, LocalFailoverInput,
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,
};
pub(crate) use self::effects::{
apply_local_execution_effect, spawn_local_oauth_success_effect, LocalAdaptiveRateLimitEffect,
LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect, LocalExecutionEffect,
LocalExecutionEffectContext, LocalHealthFailureEffect, LocalHealthSuccessEffect,
LocalOAuthInvalidationEffect, LocalOAuthSuccessEffect, LocalPoolErrorEffect,
apply_local_stream_failure_effects, apply_local_stream_success_effects,
release_local_pool_key_lease,
release_pool_key_lease_from_report_context, LocalAdaptiveRateLimitEffect,
LocalStreamFailureEffect,
};
pub(crate) use self::health::{
project_local_failure_health, project_local_key_circuit_closed,
@@ -48,8 +58,9 @@ pub(crate) use self::oauth_error::{
pub(crate) use self::policy::{
append_local_failover_policy_to_value, codex_cyber_flag_passthrough_enabled,
cyber_continue_failover_enabled, local_failover_policy_from_report_context,
local_failover_policy_from_transport, resolve_local_failover_policy, LocalFailoverPolicy,
LocalFailoverRegexRule, CYBER_CONTINUE_FAILOVER_CONFIG_KEY,
local_failover_policy_from_transport, resolve_local_failover_policy,
responses_websocket_adapter, LocalFailoverPolicy, LocalFailoverRegexRule,
ResponsesWebSocketAdapter, CYBER_CONTINUE_FAILOVER_CONFIG_KEY, RESPONSES_WEBSOCKET_CONFIG_KEY,
};
pub(crate) use self::recovery::{
analyze_local_failover, analyze_local_transport_error, apply_provider_failure_disposition,
@@ -59,7 +70,8 @@ pub(crate) use self::recovery::{
#[cfg(test)]
pub(crate) use self::report_effects::clear_local_report_effect_caches_for_tests;
pub(crate) use self::report_effects::{
apply_local_report_effect, store_local_gemini_file_mapping, LocalReportEffect,
apply_local_report_effect, store_local_gemini_file_mapping,
sync_codex_websocket_quota_metadata, LocalReportEffect,
};
pub(crate) async fn resolve_local_failover_analysis_for_attempt(
@@ -8,6 +8,7 @@ use crate::provider_transport::GatewayProviderTransportSnapshot;
use crate::AppState;
pub(crate) const CYBER_CONTINUE_FAILOVER_CONFIG_KEY: &str = "cyber_continue_failover";
pub(crate) const RESPONSES_WEBSOCKET_CONFIG_KEY: &str = "responses_websocket";
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct LocalFailoverPolicy {
@@ -297,6 +298,63 @@ pub(crate) fn codex_cyber_flag_passthrough_enabled(
.unwrap_or(true)
}
/// Selects the protocol adapter responsible for one eligible Responses
/// WebSocket upstream. Provider-scoped feature switches remain the source of
/// truth; this enum only identifies provider-specific extensions around the
/// otherwise standard Responses WebSocket protocol.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum ResponsesWebSocketAdapter {
/// A provider that speaks the standard OpenAI Responses WebSocket protocol.
Standard,
/// Standard protocol plus Codex account and quota extensions.
Codex,
}
impl ResponsesWebSocketAdapter {
pub(crate) fn supports_provider_type(self, provider_type: &str) -> bool {
match self {
Self::Standard => {
!provider_type.trim().is_empty()
&& !provider_type.trim().eq_ignore_ascii_case("codex")
}
Self::Codex => provider_type.trim().eq_ignore_ascii_case("codex"),
}
}
}
/// Whether a provider explicitly enables the standard Responses WebSocket
/// bridge. The setting is provider-scoped so rollout remains opt-in per
/// verified upstream.
pub(crate) fn responses_websocket_enabled(provider_config: Option<&Value>) -> bool {
provider_config
.and_then(|config| config.get(RESPONSES_WEBSOCKET_CONFIG_KEY))
.and_then(Value::as_object)
.and_then(|responses| responses.get("enabled"))
.and_then(Value::as_bool)
.unwrap_or(false)
}
/// Returns the enabled Responses WebSocket adapter for a provider. The shared
/// protocol bridge remains opt-in, while this resolver isolates provider-only
/// extensions from candidate planning and the session engine.
pub(crate) fn responses_websocket_adapter(
provider_type: &str,
provider_config: Option<&Value>,
) -> Option<ResponsesWebSocketAdapter> {
let provider_type = provider_type.trim();
if provider_type.is_empty() {
return None;
}
if !responses_websocket_enabled(provider_config) {
return None;
}
Some(if provider_type.eq_ignore_ascii_case("codex") {
ResponsesWebSocketAdapter::Codex
} else {
ResponsesWebSocketAdapter::Standard
})
}
fn local_failover_regex_rule_to_value(rule: &LocalFailoverRegexRule) -> Value {
json!({
"pattern": rule.pattern,
@@ -368,7 +426,9 @@ mod tests {
use super::{
append_local_failover_policy_to_value, local_failover_policy_from_report_context,
local_failover_policy_from_transport, LocalFailoverPolicy, LocalFailoverRegexRule,
local_failover_policy_from_transport, responses_websocket_adapter,
responses_websocket_enabled, LocalFailoverPolicy, LocalFailoverRegexRule,
ResponsesWebSocketAdapter,
};
use crate::provider_transport::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
@@ -593,4 +653,41 @@ mod tests {
transport.provider.provider_type = "llm".to_string();
assert!(local_failover_policy_from_transport(&transport).stop_cyber_policy_errors);
}
#[test]
fn responses_websocket_requires_an_explicit_provider_switch() {
assert!(!responses_websocket_enabled(None));
assert!(!responses_websocket_enabled(Some(&json!({
"responses_websocket": {"enabled": false}
}))));
assert!(responses_websocket_enabled(Some(&json!({
"responses_websocket": {"enabled": true}
}))));
assert_eq!(
responses_websocket_adapter(
"custom",
Some(&json!({"responses_websocket": {"enabled": false}})),
),
None
);
assert_eq!(
responses_websocket_adapter(
"custom",
Some(&json!({"responses_websocket": {"enabled": true}})),
),
Some(ResponsesWebSocketAdapter::Standard)
);
assert_eq!(
responses_websocket_adapter(
"codex",
Some(&json!({"responses_websocket": {"enabled": true}})),
),
Some(ResponsesWebSocketAdapter::Codex)
);
assert!(ResponsesWebSocketAdapter::Codex.supports_provider_type("CODEX"));
assert!(!ResponsesWebSocketAdapter::Codex.supports_provider_type("openai"));
assert!(ResponsesWebSocketAdapter::Standard.supports_provider_type("custom"));
assert!(!ResponsesWebSocketAdapter::Standard.supports_provider_type("codex"));
}
}
@@ -16,12 +16,16 @@ use serde_json::{json, Value};
use tracing::warn;
use uuid::Uuid;
use super::codex_quota_breaker::{
install_codex_quota_exhaustion_breaker, log_codex_quota_breaker_install_failure,
};
use crate::clock::current_unix_secs;
use crate::handlers::shared::sync_provider_key_quota_status_snapshot;
use crate::log_ids::short_request_id;
use crate::{AppState, GatewayError};
const RUNTIME_METADATA_CAS_MAX_ATTEMPTS: usize = 16;
const CODEX_WEBSOCKET_RATE_LIMITS_REPORT_CONTEXT_FIELD: &str = "codex_websocket_rate_limits";
static GROK_CHINESE_WAIT_DURATION_RE: OnceLock<Regex> = OnceLock::new();
static GROK_ENGLISH_WAIT_DURATION_RE: OnceLock<Regex> = OnceLock::new();
@@ -95,6 +99,12 @@ fn report_context_provider_response_headers(
(!out.is_empty()).then_some(out)
}
fn codex_websocket_quota_from_report_context(report_context: Option<&Value>) -> Option<Value> {
report_context
.and_then(|context| context.get(CODEX_WEBSOCKET_RATE_LIMITS_REPORT_CONTEXT_FIELD))
.filter(|value| value.as_object().is_some_and(|object| !object.is_empty()))
.cloned()
}
fn merge_metadata_object(
current: Option<&Value>,
section_key: &str,
@@ -257,6 +267,35 @@ fn gemini_cli_credits_from_stream_payload(
latest
}
fn codex_websocket_quota_from_stream_payload(
payload: &GatewayStreamReportRequest,
now_unix_secs: u64,
) -> Option<Value> {
let body_base64 = payload.provider_body_base64.as_deref()?;
let body = base64::engine::general_purpose::STANDARD
.decode(body_base64)
.ok()?;
let text = std::str::from_utf8(&body).ok()?;
let mut latest = None::<Value>;
for raw_line in text.lines() {
let line = raw_line.trim_matches('\r').trim();
let data = line.strip_prefix("data:").map(str::trim).unwrap_or(line);
if data.is_empty() || data == "[DONE]" || data.starts_with(':') {
continue;
}
let Ok(value) = serde_json::from_str::<Value>(data) else {
continue;
};
if let Some(quota) = admin_provider_quota_pure::parse_codex_websocket_rate_limits_response(
&value,
now_unix_secs,
) {
latest = Some(quota);
}
}
latest
}
async fn sync_gemini_cli_credits_from_report(
state: &AppState,
report_context: Option<&Value>,
@@ -572,7 +611,24 @@ async fn apply_local_sync_report_effect(state: &AppState, payload: &GatewaySyncR
}
async fn apply_local_stream_report_effect(state: &AppState, payload: &GatewayStreamReportRequest) {
if (200..300).contains(&payload.status_code) {
let websocket_quota_seen = match sync_codex_websocket_quota_from_stream_payload(state, payload)
.await
{
Ok(Some(_)) => true,
Ok(None) => false,
Err(err) => {
warn!(
event_name = "codex_realtime_quota_sync_failed",
log_type = "ops",
report_kind = %payload.report_kind,
report_request_id = %short_request_id(report_request_id(payload.report_context.as_ref())),
error = ?err,
"gateway failed to persist Codex realtime quota from WebSocket response body"
);
false
}
};
if !websocket_quota_seen && (200..300).contains(&payload.status_code) {
if let Err(err) = sync_codex_quota_from_response_headers(
state,
payload.report_context.as_ref(),
@@ -762,11 +818,6 @@ async fn sync_codex_quota_from_response_headers(
report_context: Option<&Value>,
headers: &BTreeMap<String, String>,
) -> Result<bool, GatewayError> {
let key_id = match report_context_key_id(report_context) {
Some(value) => value,
None => return Ok(false),
};
let now_unix_secs = current_unix_secs();
let observed_at_unix_secs = report_context_u64(
report_context,
@@ -791,6 +842,97 @@ async fn sync_codex_quota_from_response_headers(
}) else {
return Ok(false);
};
sync_codex_quota_metadata_with_observation(
state,
report_context,
parsed,
"response_headers",
observed_at_unix_secs,
request_started_at_unix_ms,
request_order_id,
observed_reset_generation,
observed_credential_generation,
)
.await
}
async fn sync_codex_websocket_quota_from_stream_payload(
state: &AppState,
payload: &GatewayStreamReportRequest,
) -> Result<Option<bool>, GatewayError> {
let now_unix_secs = current_unix_secs();
let parsed = codex_websocket_quota_from_report_context(payload.report_context.as_ref())
.or_else(|| codex_websocket_quota_from_stream_payload(payload, now_unix_secs));
let Some(parsed) = parsed else {
return Ok(None);
};
Ok(Some(
sync_codex_quota_metadata(
state,
payload.report_context.as_ref(),
parsed,
"websocket_response_body",
)
.await?,
))
}
async fn sync_codex_quota_metadata(
state: &AppState,
report_context: Option<&Value>,
parsed: Value,
source: &'static str,
) -> Result<bool, GatewayError> {
let observed_at_unix_secs = parsed
.get("updated_at")
.and_then(admin_provider_quota_pure::coerce_json_u64)
.filter(|value| *value > 0)
.unwrap_or_else(current_unix_secs);
let request_started_at_unix_ms =
report_context_u64(report_context, "provider_request_started_at_unix_ms");
let request_order_id = report_context_string(report_context, "provider_request_order_id");
let observed_reset_generation =
report_context_u64(report_context, "codex_quota_reset_generation");
let observed_credential_generation =
report_context_string(report_context, "codex_credential_generation");
sync_codex_quota_metadata_with_observation(
state,
report_context,
parsed,
source,
observed_at_unix_secs,
request_started_at_unix_ms,
request_order_id,
observed_reset_generation,
observed_credential_generation,
)
.await
}
#[allow(clippy::too_many_arguments)]
async fn sync_codex_quota_metadata_with_observation(
state: &AppState,
report_context: Option<&Value>,
parsed: Value,
source: &'static str,
observed_at_unix_secs: u64,
request_started_at_unix_ms: Option<u64>,
request_order_id: Option<&str>,
observed_reset_generation: Option<u64>,
observed_credential_generation: Option<&str>,
) -> Result<bool, GatewayError> {
if admin_provider_quota_pure::codex_rate_limit_metadata_exhausted(&parsed) {
if let Err(error) =
install_codex_quota_exhaustion_breaker(state, report_context, &parsed, source).await
{
log_codex_quota_breaker_install_failure(&error);
}
}
let key_id = match report_context_key_id(report_context) {
Some(value) => value,
None => return Ok(false),
};
// Runtime headers can be partial (for example only the primary window),
// so absence never authoritatively removes another stored window.
let coverage = admin_provider_quota_pure::CodexQuotaWindowCoverage::Patch;
@@ -845,7 +987,7 @@ async fn sync_codex_quota_from_response_headers(
key.status_snapshot.as_ref(),
provider.provider_type.as_str(),
updated_upstream_metadata.as_ref(),
"response_headers",
source,
);
let updated = state
.update_provider_catalog_key_runtime_metadata(
@@ -872,6 +1014,14 @@ async fn sync_codex_quota_from_response_headers(
Ok(false)
}
pub(crate) async fn sync_codex_websocket_quota_metadata(
state: &AppState,
report_context: Option<&Value>,
parsed: Value,
) -> Result<bool, GatewayError> {
sync_codex_quota_metadata(state, report_context, parsed, "websocket_response_body").await
}
#[cfg(test)]
pub(crate) fn clear_local_report_effect_caches_for_tests() {}
@@ -1104,6 +1254,33 @@ mod tests {
assert_eq!(status["quota"]["provider_type"], json!("gemini_cli"));
}
#[test]
fn codex_quota_snapshot_match_requires_explicit_signal_projection() {
let parsed = json!({
"allowed": false,
"limit_reached": true,
});
assert!(!codex_quota_snapshot_matches_metadata(
Some(&json!({
"quota": {
"exhausted": true
}
})),
&parsed,
));
assert!(codex_quota_snapshot_matches_metadata(
Some(&json!({
"quota": {
"exhausted": true,
"allowed": false,
"limit_reached": true
}
})),
&parsed,
));
}
#[test]
fn grok_quota_feedback_decrements_the_matching_window() {
let mut bucket = json!({
@@ -1290,4 +1467,42 @@ mod tests {
None
);
}
#[test]
fn codex_quota_snapshot_does_not_roll_back_exhaustion() {
let current = json!({
"allowed": false,
"limit_reached": true,
"primary_used_percent": 100.0,
"primary_reset_at": 2_000,
"updated_at": 100
});
let delayed = json!({
"allowed": true,
"limit_reached": false,
"primary_used_percent": 99.0,
"primary_reset_at": 2_000,
"updated_at": 101
});
assert!(codex_snapshot_regresses(&current, &delayed));
}
#[test]
fn codex_quota_snapshot_allows_a_new_reset_window() {
let current = json!({
"allowed": false,
"limit_reached": true,
"primary_used_percent": 100.0,
"primary_reset_at": 2_000,
"updated_at": 100
});
let refreshed = json!({
"allowed": true,
"limit_reached": false,
"primary_used_percent": 1.0,
"primary_reset_at": 1_000,
"updated_at": 101
});
assert!(!codex_snapshot_regresses(&current, &refreshed));
}
}
@@ -68,12 +68,14 @@ impl ProviderPoolInFlightGuard {
if self.released {
return;
}
self.released = true;
match &mut self.kind {
ProviderPoolInFlightGuardKind::Local {
provider_id,
counter,
} => decrement_local_provider_in_flight(provider_id, counter),
} => {
self.released = true;
decrement_local_provider_in_flight(provider_id, counter);
}
ProviderPoolInFlightGuardKind::Runtime {
runtime,
tokens_key,
@@ -85,11 +87,17 @@ impl ProviderPoolInFlightGuard {
if let Some(handle) = renew_handle.take() {
handle.abort();
}
if let Err(err) = runtime.score_remove(tokens_key, token).await {
debug!(
error = ?err,
"gateway provider pool demand: failed to release in-flight token"
);
// Mark the guard released only after Redis confirms removal.
// If this future is cancelled, Drop still schedules the same
// idempotent cleanup instead of leaving the token until TTL.
match runtime.score_remove(tokens_key, token).await {
Ok(_) => self.released = true,
Err(err) => {
debug!(
error = ?err,
"gateway provider pool demand: failed to release in-flight token; scheduling drop fallback"
);
}
}
}
}
+1
View File
@@ -94,6 +94,7 @@ const STAGES: &[&str] = &[
"stream_usage_pending",
"stream_provider_in_flight",
"stream_upstream_target_admission",
"websocket_turn_admission_held",
"stream_upstream_headers",
"stream_first_frame",
"stream_first_data",
+2
View File
@@ -377,11 +377,13 @@ pub struct AppState {
pub(crate) frontdoor_runtime_guards: Arc<FrontdoorRuntimeGuardConfig>,
pub(crate) request_body_buffer_budget: Arc<Semaphore>,
pub(crate) request_gate: Option<Arc<ConcurrencyGate>>,
pub(crate) websocket_connection_gate: Option<Arc<ConcurrencyGate>>,
pub(crate) auth_snapshot_load_gate: Option<Arc<ConcurrencyGate>>,
pub(crate) candidate_planning_gate: Option<Arc<ConcurrencyGate>>,
pub(crate) upstream_execution_gate: Option<Arc<ConcurrencyGate>>,
pub(crate) upstream_target_admission: Arc<crate::upstream_admission::UpstreamTargetAdmission>,
pub(crate) distributed_request_gate: Option<Arc<RuntimeSemaphore>>,
pub(crate) distributed_websocket_connection_gate: Option<Arc<RuntimeSemaphore>>,
pub(crate) client: reqwest::Client,
pub(crate) owner_forward_client: reqwest::Client,
pub(crate) auth_context_cache: Arc<AuthContextCache>,
+84 -2
View File
@@ -325,6 +325,7 @@ impl AppState {
frontdoor_runtime_guards.request_body_buffer_budget_permits,
)),
request_gate: None,
websocket_connection_gate: None,
auth_snapshot_load_gate: frontdoor_runtime_guards
.auth_snapshot_load_gate_limit
.map(|limit| Arc::new(ConcurrencyGate::new("gateway_auth_snapshot_load", limit))),
@@ -341,6 +342,7 @@ impl AppState {
),
),
distributed_request_gate: None,
distributed_websocket_connection_gate: None,
client,
owner_forward_client,
auth_context_cache: Arc::new(AuthContextCache::default()),
@@ -574,8 +576,20 @@ impl AppState {
}
pub fn with_request_concurrency_limit(mut self, limit: usize) -> Self {
self.request_gate = Some(Arc::new(ConcurrencyGate::new(
"gateway_requests",
let limit = limit.max(1);
self.request_gate = Some(Arc::new(ConcurrencyGate::new("gateway_requests", limit)));
if self.websocket_connection_gate.is_none() {
self.websocket_connection_gate = Some(Arc::new(ConcurrencyGate::new(
"gateway_websocket_connections",
limit,
)));
}
self
}
pub fn with_websocket_connection_limit(mut self, limit: usize) -> Self {
self.websocket_connection_gate = Some(Arc::new(ConcurrencyGate::new(
"gateway_websocket_connections",
limit.max(1),
)));
self
@@ -646,6 +660,11 @@ impl AppState {
self
}
pub fn with_distributed_websocket_connection_gate(mut self, gate: RuntimeSemaphore) -> Self {
self.distributed_websocket_connection_gate = Some(Arc::new(gate));
self
}
pub fn with_frontdoor_cors_config(mut self, config: FrontdoorCorsConfig) -> Self {
self.frontdoor_cors = Some(Arc::new(config));
self
@@ -1224,6 +1243,12 @@ impl AppState {
self.request_gate.as_ref().map(|gate| gate.snapshot())
}
pub(crate) fn websocket_connection_concurrency_snapshot(&self) -> Option<ConcurrencySnapshot> {
self.websocket_connection_gate
.as_ref()
.map(|gate| gate.snapshot())
}
pub(crate) fn auth_snapshot_load_concurrency_snapshot(&self) -> Option<ConcurrencySnapshot> {
self.auth_snapshot_load_gate
.as_ref()
@@ -1251,6 +1276,15 @@ impl AppState {
}
}
pub(crate) async fn distributed_websocket_connection_concurrency_snapshot(
&self,
) -> Result<Option<RuntimeSemaphoreSnapshot>, RuntimeSemaphoreError> {
match self.distributed_websocket_connection_gate.as_ref() {
Some(gate) => gate.snapshot().await.map(Some),
None => Ok(None),
}
}
pub(crate) async fn metric_samples(&self) -> Vec<MetricSample> {
let now = std::time::Instant::now();
let snapshot = self.metric_snapshot.read().await.clone();
@@ -1557,6 +1591,9 @@ impl AppState {
if let Some(snapshot) = self.request_concurrency_snapshot() {
samples.extend(snapshot.to_metric_samples("gateway_requests"));
}
if let Some(snapshot) = self.websocket_connection_concurrency_snapshot() {
samples.extend(snapshot.to_metric_samples("gateway_websocket_connections"));
}
if let Some(snapshot) = self.auth_snapshot_load_concurrency_snapshot() {
samples.extend(snapshot.to_metric_samples("gateway_auth_snapshot_load"));
}
@@ -1600,6 +1637,28 @@ impl AppState {
)])],
}
};
let distributed_websocket_connection_metrics = async {
let Some(gate) = self.distributed_websocket_connection_gate.as_ref() else {
return Vec::new();
};
match tokio::time::timeout(DISTRIBUTED_CONCURRENCY_METRICS_TIMEOUT, gate.snapshot())
.await
{
Ok(Ok(snapshot)) => {
snapshot.to_metric_samples("gateway_websocket_connections_distributed")
}
Ok(Err(_)) | Err(_) => vec![MetricSample::new(
"concurrency_unavailable",
"Whether the distributed concurrency gate is currently unavailable.",
MetricKind::Gauge,
1,
)
.with_labels(vec![MetricLabel::new(
"gate",
"gateway_websocket_connections_distributed",
)])],
}
};
let postgres_observability_metrics = async {
match tokio::time::timeout(
POSTGRES_OBSERVABILITY_METRICS_TIMEOUT,
@@ -1649,6 +1708,7 @@ impl AppState {
);
let (
distributed_request_metrics,
distributed_websocket_connection_metrics,
postgres_observability_metrics,
postgres_activity_group_metrics,
redis_runtime_metrics,
@@ -1656,6 +1716,7 @@ impl AppState {
usage_counter_pending_health_metrics,
) = tokio::join!(
distributed_request_metrics,
distributed_websocket_connection_metrics,
postgres_observability_metrics,
postgres_activity_group_metrics,
redis_runtime_metrics,
@@ -1663,6 +1724,7 @@ impl AppState {
usage_counter_pending_health_metrics,
);
samples.extend(distributed_request_metrics);
samples.extend(distributed_websocket_connection_metrics);
samples.extend(postgres_observability_metrics);
samples.extend(postgres_activity_group_metrics);
samples.extend(redis_runtime_metrics);
@@ -1783,6 +1845,26 @@ impl AppState {
Ok(AdmissionPermit::from_parts(local, distributed))
}
pub(crate) async fn try_acquire_websocket_connection_permit(
&self,
) -> Result<Option<AdmissionPermit>, RequestAdmissionError> {
let local = self
.websocket_connection_gate
.as_ref()
.map(|gate| gate.try_acquire())
.transpose()
.map_err(RequestAdmissionError::Local)?;
let distributed = match self.distributed_websocket_connection_gate.as_ref() {
Some(gate) => Some(
gate.try_acquire()
.await
.map_err(RequestAdmissionError::Distributed)?,
),
None => None,
};
Ok(AdmissionPermit::from_parts(local, distributed))
}
pub fn has_auth_api_key_data_reader(&self) -> bool {
self.data.has_auth_api_key_reader()
}
@@ -53,6 +53,83 @@ fn memory_runtime_semaphore(gate: &'static str, limit: usize) -> RuntimeSemaphor
.expect("memory runtime semaphore should build")
}
#[test]
fn gateway_websocket_connections_use_independent_admission() {
run_concurrency_test(
"gateway_websocket_connections_use_independent_admission",
gateway_websocket_connections_use_independent_admission_impl,
);
}
async fn gateway_websocket_connections_use_independent_admission_impl() {
let state = AppState::new()
.expect("gateway state should build")
.with_request_concurrency_limit(1)
.with_websocket_connection_limit(1);
let request_permit = state
.try_acquire_request_permit()
.await
.expect("request admission should succeed")
.expect("request gate should return a permit");
let websocket_permit = state
.try_acquire_websocket_connection_permit()
.await
.expect("WebSocket admission should be independent from request admission")
.expect("WebSocket connection gate should return a permit");
assert_eq!(
state
.request_concurrency_snapshot()
.expect("request gate should be configured")
.in_flight,
1
);
assert_eq!(
state
.websocket_connection_concurrency_snapshot()
.expect("WebSocket connection gate should be configured")
.in_flight,
1
);
let websocket_error = state
.try_acquire_websocket_connection_permit()
.await
.expect_err("second WebSocket connection should be rejected");
assert!(matches!(
websocket_error,
crate::router::RequestAdmissionError::Local(aether_runtime::ConcurrencyError::Saturated {
gate: "gateway_websocket_connections",
limit: 1,
})
));
drop(request_permit);
let replacement_request_permit = state
.try_acquire_request_permit()
.await
.expect("request admission should remain available independently")
.expect("request gate should return a replacement permit");
assert_eq!(
state
.websocket_connection_concurrency_snapshot()
.expect("WebSocket connection gate should be configured")
.in_flight,
1
);
drop(websocket_permit);
let replacement_websocket_permit = state
.try_acquire_websocket_connection_permit()
.await
.expect("WebSocket admission should recover after release")
.expect("WebSocket connection gate should return a replacement permit");
drop(replacement_request_permit);
drop(replacement_websocket_permit);
}
fn sample_decision() -> crate::control::GatewayControlDecision {
crate::control::GatewayControlDecision {
public_path: "/v1/chat/completions".to_string(),
@@ -316,9 +393,14 @@ async fn gateway_exposes_request_concurrency_metrics_impl() {
let state = AppState::new()
.expect("gateway state should build")
.with_request_concurrency_limit(3)
.with_websocket_connection_limit(7)
.with_distributed_request_concurrency_gate(memory_runtime_semaphore(
"gateway_requests_distributed",
5,
))
.with_distributed_websocket_connection_gate(memory_runtime_semaphore(
"gateway_websocket_connections_distributed",
9,
));
assert!(state.prewarm_metric_snapshot().await);
let gateway = build_router_with_state(state);
@@ -344,6 +426,15 @@ async fn gateway_exposes_request_concurrency_metrics_impl() {
assert!(body.contains("concurrency_available_permits{gate=\"gateway_requests\"} 3"));
assert!(body.contains("concurrency_in_flight{gate=\"gateway_requests_distributed\"} 0"));
assert!(body.contains("concurrency_available_permits{gate=\"gateway_requests_distributed\"} 5"));
assert!(body.contains("concurrency_in_flight{gate=\"gateway_websocket_connections\"} 0"));
assert!(
body.contains("concurrency_available_permits{gate=\"gateway_websocket_connections\"} 7")
);
assert!(body
.contains("concurrency_in_flight{gate=\"gateway_websocket_connections_distributed\"} 0"));
assert!(body.contains(
"concurrency_available_permits{gate=\"gateway_websocket_connections_distributed\"} 9"
));
assert!(body.contains("tunnel_proxy_connections 0"));
assert!(body.contains("tunnel_nodes 0"));
assert!(body.contains("tunnel_active_streams 0"));
@@ -1008,6 +1008,185 @@ mod tests {
assert_eq!(quota.get("updated_at"), quota.get("observed_at"));
}
#[tokio::test]
async fn submit_stream_report_updates_codex_quota_from_websocket_response_body() {
crate::orchestration::clear_local_report_effect_caches_for_tests();
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider_catalog_provider(
"provider-codex-websocket",
"codex",
)],
Vec::new(),
vec![sample_provider_catalog_key(
"key-codex-websocket",
"provider-codex-websocket",
)],
));
let state = build_provider_catalog_test_state(Arc::clone(&provider_catalog_repository));
let websocket_event = json!({
"chunks": [{
"type": "codex.rate_limits",
"plan_type": "free",
"rate_limits": {
"allowed": true,
"limit_reached": false,
"primary": {
"used_percent": 91,
"window_minutes": 43200,
"reset_after_seconds": 2590791,
"reset_at": 1787154563u64
}
}
}]
});
let body = format!("data: {websocket_event}\n\n");
submit_stream_report(
&state,
GatewayStreamReportRequest {
trace_id: "trace-codex-reporting-websocket".to_string(),
report_kind: "openai_responses_stream_success".to_string(),
report_context: Some(json!({
"request_id": "req-codex-reporting-websocket",
"key_id": "key-codex-websocket",
"websocket_mode": true
})),
status_code: 200,
headers: sample_codex_paid_headers(),
provider_body_base64: Some(
base64::engine::general_purpose::STANDARD.encode(body.as_bytes()),
),
provider_body_state: Some(UsageBodyCaptureState::Inline),
client_body_base64: None,
client_body_state: None,
terminal_summary: None,
telemetry: None,
},
)
.await
.expect("stream report should stay local");
let reloaded = provider_catalog_repository
.list_keys_by_ids(&["key-codex-websocket".to_string()])
.await
.expect("keys should list");
let codex = reloaded[0]
.upstream_metadata
.as_ref()
.and_then(serde_json::Value::as_object)
.and_then(|metadata| metadata.get("codex"))
.and_then(serde_json::Value::as_object)
.expect("codex metadata should exist");
assert_eq!(codex.get("plan_type"), Some(&json!("free")));
assert_eq!(codex.get("allowed"), Some(&json!(true)));
assert_eq!(codex.get("limit_reached"), Some(&json!(false)));
assert_eq!(codex.get("primary_used_percent"), Some(&json!(91.0)));
assert_eq!(codex.get("primary_window_minutes"), Some(&json!(43_200u64)));
let quota = reloaded[0]
.status_snapshot
.as_ref()
.and_then(serde_json::Value::as_object)
.and_then(|snapshot| snapshot.get("quota"))
.and_then(serde_json::Value::as_object)
.expect("quota snapshot should exist");
assert_eq!(quota.get("source"), Some(&json!("websocket_response_body")));
assert_eq!(quota.get("code"), Some(&json!("ok")));
assert_eq!(quota.get("usage_ratio"), Some(&json!(0.91)));
}
#[tokio::test]
async fn submit_stream_report_marks_codex_websocket_usage_limit_error_exhausted() {
crate::orchestration::clear_local_report_effect_caches_for_tests();
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider_catalog_provider(
"provider-codex-websocket-limit",
"codex",
)],
Vec::new(),
vec![sample_provider_catalog_key(
"key-codex-websocket-limit",
"provider-codex-websocket-limit",
)],
));
let state = build_provider_catalog_test_state(Arc::clone(&provider_catalog_repository));
let websocket_event = json!({
"type": "error",
"error": {
"type": "usage_limit_reached",
"plan_type": "free",
"resets_at": 1_787_274_385u64,
"resets_in_seconds": 2_590_077u64,
},
"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-After-Seconds": "2590078",
"X-Codex-Primary-Reset-At": "1787274385",
"X-Codex-Credits-Has-Credits": "False",
},
});
let body = format!("data: {websocket_event}\n\n");
submit_stream_report(
&state,
GatewayStreamReportRequest {
trace_id: "trace-codex-reporting-websocket-limit".to_string(),
report_kind: "openai_responses_stream_success".to_string(),
report_context: Some(json!({
"request_id": "req-codex-reporting-websocket-limit",
"key_id": "key-codex-websocket-limit",
"websocket_mode": true,
})),
status_code: 429,
headers: BTreeMap::new(),
provider_body_base64: Some(
base64::engine::general_purpose::STANDARD.encode(body.as_bytes()),
),
provider_body_state: Some(UsageBodyCaptureState::Inline),
client_body_base64: None,
client_body_state: None,
terminal_summary: None,
telemetry: None,
},
)
.await
.expect("stream report should stay local");
let reloaded = provider_catalog_repository
.list_keys_by_ids(&["key-codex-websocket-limit".to_string()])
.await
.expect("keys should list");
let codex = reloaded[0]
.upstream_metadata
.as_ref()
.and_then(serde_json::Value::as_object)
.and_then(|metadata| metadata.get("codex"))
.and_then(serde_json::Value::as_object)
.expect("codex metadata should exist");
assert_eq!(codex.get("allowed"), Some(&json!(false)));
assert_eq!(codex.get("limit_reached"), Some(&json!(true)));
assert_eq!(codex.get("primary_used_percent"), Some(&json!(100.0)));
assert_eq!(
codex.get("primary_reset_at"),
Some(&json!(1_787_274_385u64))
);
let quota = reloaded[0]
.status_snapshot
.as_ref()
.and_then(serde_json::Value::as_object)
.and_then(|snapshot| snapshot.get("quota"))
.and_then(serde_json::Value::as_object)
.expect("quota snapshot should exist");
assert_eq!(quota.get("source"), Some(&json!("websocket_response_body")));
assert_eq!(quota.get("code"), Some(&json!("exhausted")));
}
#[tokio::test]
async fn submit_stream_report_updates_codex_quota_from_provider_response_headers() {
crate::orchestration::clear_local_report_effect_caches_for_tests();
@@ -1216,6 +1216,7 @@ fn admin_usage_active_request_json(
"api_key_name": api_key_name,
"provider_key_name": provider_key_name,
"is_stream": item.is_stream,
"is_websocket": item.is_websocket(),
"upstream_is_stream": upstream_is_stream,
"client_requested_stream": client_is_stream,
"client_is_stream": client_is_stream,
@@ -1348,6 +1349,7 @@ pub fn admin_usage_record_json(
"end_to_end_first_byte_time_ms"
)),
);
object.insert("is_websocket".to_string(), json!(item.is_websocket()));
object.insert("is_stream".to_string(), json!(item.is_stream));
object.insert(
UPSTREAM_IS_STREAM_KEY.to_string(),
@@ -2697,6 +2699,30 @@ mod tests {
}
}
#[test]
fn admin_usage_payloads_expose_websocket_transport() {
let item = StoredRequestUsageAudit {
request_metadata: Some(json!({
"websocket_mode": true,
"websocket_transport": "responses",
})),
..sample_usage("completed", Some(200), None)
};
let record = admin_usage_record_json(
&item,
&BTreeMap::new(),
&BTreeMap::new(),
false,
false,
None,
);
let active = admin_usage_active_request_json(&item, None, None, None);
assert_eq!(record["is_websocket"], true);
assert_eq!(active["is_websocket"], true);
}
#[test]
fn admin_usage_record_infers_client_family_from_user_agent() {
let item = StoredRequestUsageAudit {
+7
View File
@@ -111,6 +111,13 @@ pub fn admin_pool_key_account_quota_exhausted(
aether_provider_pool::provider_pool_key_account_quota_exhausted(key, provider_type)
}
pub fn admin_pool_key_quota_hard_blocked(
key: &StoredProviderCatalogKey,
provider_type: &str,
) -> bool {
aether_provider_pool::provider_pool_key_quota_hard_blocked(key, provider_type)
}
fn admin_pool_has_proxy(key: &StoredProviderCatalogKey) -> bool {
match key.proxy.as_ref() {
Some(Value::Object(values)) => !values.is_empty(),
+268 -1
View File
@@ -2019,6 +2019,174 @@ pub fn parse_codex_wham_usage_response(
Some(serde_json::Value::Object(result))
}
/// Normalizes quota metadata emitted by the Codex Responses WebSocket.
///
/// The upstream normally sends a `codex.rate_limits` item inside a `chunks`
/// envelope. When the account is already exhausted it can instead send a
/// terminal `usage_limit_reached` error whose embedded `X-Codex-*` headers
/// contain the authoritative final quota snapshot.
pub fn parse_codex_websocket_rate_limits_response(
value: &serde_json::Value,
updated_at_unix_secs: u64,
) -> Option<serde_json::Value> {
let mut latest = parse_codex_websocket_quota_event(value, updated_at_unix_secs);
for chunk in value
.get("chunks")
.and_then(serde_json::Value::as_array)
.into_iter()
.flatten()
{
if let Some(parsed) = parse_codex_websocket_quota_event(chunk, updated_at_unix_secs) {
latest = Some(parsed);
}
}
latest
}
fn parse_codex_websocket_quota_event(
value: &serde_json::Value,
updated_at_unix_secs: u64,
) -> Option<serde_json::Value> {
parse_codex_websocket_rate_limits_chunk(value, updated_at_unix_secs)
.or_else(|| parse_codex_websocket_usage_limit_error(value, updated_at_unix_secs))
}
/// Returns whether normalized Codex rate-limit metadata says that the account
/// cannot accept another request. Explicit upstream flags take precedence, and
/// the percentage fallback keeps older payloads working when those flags are
/// absent.
pub fn codex_rate_limit_metadata_exhausted(value: &serde_json::Value) -> bool {
let allowed = value.get("allowed").and_then(coerce_json_bool);
let limit_reached = value.get("limit_reached").and_then(coerce_json_bool);
if allowed == Some(false) || limit_reached == Some(true) {
return true;
}
if allowed == Some(true) || limit_reached == Some(false) {
return false;
}
["primary_used_percent", "secondary_used_percent"]
.into_iter()
.filter_map(|key| value.get(key))
.filter_map(coerce_json_f64)
.any(|used_percent| used_percent >= 100.0 - 1e-6)
}
fn parse_codex_websocket_rate_limits_chunk(
value: &serde_json::Value,
updated_at_unix_secs: u64,
) -> Option<serde_json::Value> {
let root = value.as_object()?;
if root.get("type").and_then(serde_json::Value::as_str) != Some("codex.rate_limits") {
return None;
}
let rate_limits = root
.get("rate_limits")
.and_then(serde_json::Value::as_object)?;
let mut result = serde_json::Map::new();
let plan_type = root
.get("plan_type")
.or_else(|| rate_limits.get("plan_type"))
.and_then(serde_json::Value::as_str)
.and_then(|value| normalize_codex_plan_type(Some(value)));
if let Some(plan_type) = plan_type {
result.insert("plan_type".to_string(), json!(plan_type));
}
if let Some(allowed) = rate_limits.get("allowed").and_then(coerce_json_bool) {
result.insert("allowed".to_string(), json!(allowed));
}
if let Some(limit_reached) = rate_limits.get("limit_reached").and_then(coerce_json_bool) {
result.insert("limit_reached".to_string(), json!(limit_reached));
}
if let Some(primary) = rate_limits
.get("primary")
.and_then(serde_json::Value::as_object)
{
codex_write_window(&mut result, primary, "primary");
}
if let Some(secondary) = rate_limits
.get("secondary")
.and_then(serde_json::Value::as_object)
{
codex_write_window(&mut result, secondary, "secondary");
}
if result.is_empty() {
return None;
}
result.insert("updated_at".to_string(), json!(updated_at_unix_secs));
Some(serde_json::Value::Object(result))
}
fn parse_codex_websocket_usage_limit_error(
value: &serde_json::Value,
updated_at_unix_secs: u64,
) -> Option<serde_json::Value> {
let root = value.as_object()?;
if root.get("type").and_then(serde_json::Value::as_str) != Some("error") {
return None;
}
let status_code = root
.get("status_code")
.or_else(|| root.get("status"))
.and_then(coerce_json_u64);
if status_code != Some(429) {
return None;
}
let error = root.get("error").and_then(serde_json::Value::as_object)?;
if error.get("type").and_then(serde_json::Value::as_str) != Some("usage_limit_reached") {
return None;
}
let headers = root
.get("headers")
.and_then(serde_json::Value::as_object)
.map(|headers| {
headers
.iter()
.filter_map(|(name, value)| {
value
.as_str()
.map(|value| (name.clone(), value.to_string()))
})
.collect::<BTreeMap<_, _>>()
})
.unwrap_or_default();
let mut result = parse_codex_usage_headers(&headers, updated_at_unix_secs)
.and_then(|value| value.as_object().cloned())
.unwrap_or_default();
if !result.contains_key("plan_type") {
if let Some(plan_type) = error
.get("plan_type")
.and_then(serde_json::Value::as_str)
.and_then(|value| normalize_codex_plan_type(Some(value)))
{
result.insert("plan_type".to_string(), json!(plan_type));
}
}
if !result.contains_key("primary_reset_at") {
if let Some(reset_at) = error.get("resets_at").and_then(coerce_json_u64) {
result.insert("primary_reset_at".to_string(), json!(reset_at));
}
}
if !result.contains_key("primary_reset_after_seconds") {
if let Some(reset_after_seconds) = error.get("resets_in_seconds").and_then(coerce_json_u64)
{
result.insert(
"primary_reset_after_seconds".to_string(),
json!(reset_after_seconds),
);
}
}
// `usage_limit_reached` is a definitive, account-wide terminal signal.
// Preserve that fact even if an intermediary strips some Codex headers.
result.insert("allowed".to_string(), json!(false));
result.insert("limit_reached".to_string(), json!(true));
result.insert("updated_at".to_string(), json!(updated_at_unix_secs));
Some(serde_json::Value::Object(result))
}
fn parse_codex_reset_credit_timestamp(value: Option<&serde_json::Value>) -> Option<u64> {
let value = value?;
if let Some(timestamp) = coerce_json_u64(value) {
@@ -3374,10 +3542,11 @@ pub fn parse_chatgpt_web_conversation_init_response(
mod tests {
use super::{
codex_build_invalid_state, codex_oauth_success_request_order_is_stale,
codex_rate_limit_metadata_exhausted,
codex_runtime_invalid_reason, extract_execution_error_detail,
merge_codex_quota_metadata_snapshot, normalize_codex_reset_credit_consume_outcome,
parse_antigravity_usage_response, parse_chatgpt_web_conversation_init_response,
parse_codex_backend_me_response, parse_codex_usage_headers,
parse_codex_usage_headers, parse_codex_websocket_rate_limits_response,
parse_codex_wham_reset_credits_detail_response, parse_codex_wham_usage_response,
parse_gemini_cli_retrieve_user_quota_response,
parse_gemini_cli_v1internal_credits_response, parse_windsurf_model_configs_response,
@@ -5395,6 +5564,104 @@ mod tests {
assert!(parsed.get("secondary_window_minutes").is_none());
}
#[test]
fn parses_codex_websocket_rate_limits_from_chunk_envelope() {
let parsed = parse_codex_websocket_rate_limits_response(
&json!({
"chunks": [
{"type": "response.output_text.delta", "delta": "ignored"},
{
"type": "codex.rate_limits",
"plan_type": "free",
"rate_limits": {
"allowed": true,
"limit_reached": false,
"primary": {
"used_percent": 91,
"window_minutes": 43200,
"reset_after_seconds": 2590791,
"reset_at": 1787154563u64
}
}
}
]
}),
1_787_000_000,
)
.expect("Codex WebSocket quota chunk should parse");
assert_eq!(parsed.get("plan_type"), Some(&json!("free")));
assert_eq!(parsed.get("allowed"), Some(&json!(true)));
assert_eq!(parsed.get("limit_reached"), Some(&json!(false)));
assert_eq!(parsed.get("primary_used_percent"), Some(&json!(91.0)));
assert_eq!(
parsed.get("primary_window_minutes"),
Some(&json!(43_200u64))
);
assert_eq!(
parsed.get("primary_reset_after_seconds"),
Some(&json!(2_590_791u64))
);
assert!(parsed.get("secondary_used_percent").is_none());
}
#[test]
fn parses_codex_websocket_usage_limit_error_headers() {
let parsed = parse_codex_websocket_rate_limits_response(
&json!({
"type": "error",
"error": {
"type": "usage_limit_reached",
"plan_type": "free",
"resets_at": 1_787_274_385u64,
"resets_in_seconds": 2_590_077u64,
},
"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-After-Seconds": "2590078",
"X-Codex-Primary-Reset-At": "1787274385",
"X-Codex-Credits-Has-Credits": "False",
},
}),
1_787_000_000,
)
.expect("Codex usage-limit error should parse as quota metadata");
assert_eq!(parsed.get("allowed"), Some(&json!(false)));
assert_eq!(parsed.get("limit_reached"), Some(&json!(true)));
assert_eq!(parsed.get("plan_type"), Some(&json!("free")));
assert_eq!(parsed.get("primary_used_percent"), Some(&json!(100.0)));
assert_eq!(
parsed.get("primary_reset_at"),
Some(&json!(1_787_274_385u64))
);
assert!(codex_rate_limit_metadata_exhausted(&parsed));
}
#[test]
fn codex_rate_limit_metadata_detects_explicit_and_window_exhaustion() {
assert!(codex_rate_limit_metadata_exhausted(&json!({
"allowed": false
})));
assert!(codex_rate_limit_metadata_exhausted(&json!({
"limit_reached": true
})));
assert!(codex_rate_limit_metadata_exhausted(&json!({
"primary_used_percent": 100
})));
assert!(!codex_rate_limit_metadata_exhausted(&json!({
"allowed": true,
"primary_used_percent": 100
})));
assert!(!codex_rate_limit_metadata_exhausted(&json!({
"limit_reached": false,
"secondary_used_percent": 100
})));
}
#[test]
fn parses_codex_reset_credit_count_from_wham_usage() {
let parsed = parse_codex_wham_usage_response(
+1 -1
View File
@@ -89,7 +89,7 @@ pub struct AiExecutionPlanPayload {
pub auth_context: Option<ExecutionRuntimeAuthContext>,
}
#[derive(Debug, Deserialize, Serialize)]
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct AiExecutionDecision {
pub action: String,
#[serde(default)]
@@ -185,6 +185,7 @@ SELECT
OR NULLIF(BTRIM("usage".request_metadata->>'provider_actual_service_tier'), '') IS NOT NULL
OR ("usage".request_metadata->>'client_requested_stream') IN ('true', 'false')
OR ("usage".request_metadata->>'upstream_is_stream') IN ('true', 'false')
OR ("usage".request_metadata->>'websocket_mode') IN ('true', 'false')
THEN jsonb_strip_nulls(jsonb_build_object(
'client_ip',
NULLIF(BTRIM("usage".request_metadata->>'client_ip'), ''),
@@ -213,6 +214,12 @@ SELECT
WHEN ("usage".request_metadata->>'upstream_is_stream') IN ('true', 'false')
THEN ("usage".request_metadata->>'upstream_is_stream')::boolean
ELSE NULL
END,
'websocket_mode',
CASE
WHEN ("usage".request_metadata->>'websocket_mode') IN ('true', 'false')
THEN ("usage".request_metadata->>'websocket_mode')::boolean
ELSE NULL
END
))::json
ELSE NULL::json
@@ -185,6 +185,7 @@ SELECT
OR NULLIF(BTRIM("usage".request_metadata->>'provider_actual_service_tier'), '') IS NOT NULL
OR ("usage".request_metadata->>'client_requested_stream') IN ('true', 'false')
OR ("usage".request_metadata->>'upstream_is_stream') IN ('true', 'false')
OR ("usage".request_metadata->>'websocket_mode') IN ('true', 'false')
THEN jsonb_strip_nulls(jsonb_build_object(
'client_ip',
NULLIF(BTRIM("usage".request_metadata->>'client_ip'), ''),
@@ -213,6 +214,12 @@ SELECT
WHEN ("usage".request_metadata->>'upstream_is_stream') IN ('true', 'false')
THEN ("usage".request_metadata->>'upstream_is_stream')::boolean
ELSE NULL
END,
'websocket_mode',
CASE
WHEN ("usage".request_metadata->>'websocket_mode') IN ('true', 'false')
THEN ("usage".request_metadata->>'websocket_mode')::boolean
ELSE NULL
END
))::json
ELSE NULL::json
@@ -3221,6 +3221,8 @@ fn usage_sql_uses_json_null_placeholders_for_usage_payload_columns() {
assert!(sql.contains("request_metadata->>'provider_reasoning_effort'"));
assert!(sql.contains("request_metadata->>'provider_service_tier'"));
assert!(sql.contains("request_metadata->>'provider_actual_service_tier'"));
assert!(sql.contains("request_metadata->>'websocket_mode'"));
assert!(sql.contains("'websocket_mode'"));
assert!(sql.contains("AS client_family"));
assert!(sql.contains("request_metadata->'client_session_affinity'->>'client_family'"));
assert!(sql.contains("request_metadata->>'client_family'"));
@@ -37,4 +37,5 @@ pub use types::{
PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, PROVIDER_REASONING_EFFORT_METADATA_KEY,
PROVIDER_SERVICE_TIER_METADATA_KEY, REQUESTED_REASONING_EFFORT_METADATA_KEY,
ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY, ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY,
WEBSOCKET_MODE_METADATA_KEY, WEBSOCKET_TRANSPORT_METADATA_KEY,
};
@@ -9,6 +9,8 @@ pub const PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY: &str = "provider_actual_ser
pub const PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY: &str = "provider_cache_ttl_minutes";
pub const ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY: &str = "routing_candidate_skip_reason";
pub const ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY: &str = "routing_failure_diagnostic";
pub const WEBSOCKET_MODE_METADATA_KEY: &str = "websocket_mode";
pub const WEBSOCKET_TRANSPORT_METADATA_KEY: &str = "websocket_transport";
pub fn extract_provider_reasoning_effort_from_body(value: Option<&Value>) -> Option<String> {
let object = value.and_then(Value::as_object)?;
@@ -539,6 +541,11 @@ impl StoredRequestUsageAudit {
usage_request_metadata_client_family(self.request_metadata.as_ref())
}
pub fn is_websocket(&self) -> bool {
self.request_metadata_bool(WEBSOCKET_MODE_METADATA_KEY)
.unwrap_or(false)
}
fn billing_snapshot_resolved_number(&self, key: &str) -> Option<f64> {
self.request_metadata_object()
.and_then(|metadata| metadata.get("billing_snapshot"))
@@ -2394,7 +2401,8 @@ mod tests {
extract_provider_actual_service_tier_from_response,
extract_provider_service_tier_from_body, resolve_provider_cache_ttl_minutes,
StoredRequestUsageAudit, UpsertUsageRecord, UsageBodyCaptureState, UsageBodyCaptureStorage,
UsageBodyField, UsageProviderPerformanceQuery,
UsageBodyField, UsageProviderPerformanceQuery, WEBSOCKET_MODE_METADATA_KEY,
WEBSOCKET_TRANSPORT_METADATA_KEY,
};
use serde_json::{json, Value};
@@ -2663,6 +2671,19 @@ mod tests {
assert_eq!(usage.settlement_price_per_request(), Some(0.02));
}
#[test]
fn websocket_transport_uses_typed_request_metadata() {
let mut usage = sample_usage();
assert!(!usage.is_websocket());
usage.request_metadata = Some(json!({
WEBSOCKET_MODE_METADATA_KEY: true,
WEBSOCKET_TRANSPORT_METADATA_KEY: "responses",
}));
assert!(usage.is_websocket());
}
#[test]
fn settlement_accessors_fall_back_to_billing_snapshot_and_legacy_output_price() {
let mut usage = sample_usage();
+28 -1
View File
@@ -38,6 +38,7 @@ pub struct PoolMemberSignals {
pub quota_reset_seconds: Option<f64>,
pub account_blocked: bool,
pub quota_exhausted: bool,
pub quota_hard_blocked: bool,
pub health_score: Option<f64>,
pub latency_avg_ms: Option<f64>,
pub catalog_lru_score: Option<f64>,
@@ -224,7 +225,9 @@ fn schedule_pool_group<Candidate>(
continue;
}
if pool_config.skip_exhausted_accounts && item.key_context.quota_exhausted {
if item.key_context.quota_hard_blocked
|| (pool_config.skip_exhausted_accounts && item.key_context.quota_exhausted)
{
skipped.push(PoolSkippedCandidate {
candidate: item.candidate,
skip_reason: POOL_ACCOUNT_EXHAUSTED_SKIP_REASON,
@@ -901,6 +904,30 @@ mod tests {
);
}
#[test]
fn pool_scheduler_always_skips_hard_quota_blocks() {
let ready = sample_candidate("provider-pool", "endpoint-1", "key-ready", 10, true);
let mut hard_blocked =
sample_candidate("provider-pool", "endpoint-1", "key-blocked", 10, true);
hard_blocked.key_context.quota_exhausted = true;
hard_blocked.key_context.quota_hard_blocked = true;
let outcome = run_pool_scheduler(vec![ready, hard_blocked], &BTreeMap::new(), "seed");
assert_eq!(
outcome
.candidates
.iter()
.map(|item| item.candidate.as_str())
.collect::<Vec<_>>(),
vec!["key-ready"]
);
assert_eq!(
outcome.skipped_candidates[0].skip_reason,
POOL_ACCOUNT_EXHAUSTED_SKIP_REASON
);
}
#[test]
fn pool_scheduler_promotes_sticky_hit_before_other_sorted_keys() {
let key_a = sample_candidate("provider-pool", "endpoint-1", "key-a", 10, true)
+117 -3
View File
@@ -34,9 +34,10 @@ pub use providers::{
WINDSURF_MODEL_CONFIGS_PATH, WINDSURF_RATE_LIMIT_PATH, WINDSURF_USER_STATUS_PATH,
};
pub use quota::{
provider_pool_key_account_quota_exhausted, provider_pool_key_scheduling_label,
provider_pool_member_quota_snapshot, provider_pool_quota_metadata_provider_type,
provider_pool_quota_metadata_updated_at, provider_pool_quota_snapshot_updated_at,
provider_pool_key_account_quota_exhausted, provider_pool_key_quota_hard_blocked,
provider_pool_key_scheduling_label, provider_pool_member_quota_snapshot,
provider_pool_quota_metadata_provider_type, provider_pool_quota_metadata_updated_at,
provider_pool_quota_snapshot_updated_at,
};
pub use quota_refresh::ProviderPoolQuotaRequestSpec;
pub use service::ProviderPoolService;
@@ -693,6 +694,15 @@ mod tests {
#[test]
fn provider_quota_exhaustion_is_adapter_owned() {
assert!(provider_pool_key_account_quota_exhausted(
&sample_key(Some(json!({
"codex": {
"allowed": false,
"limit_reached": true
}
}))),
"codex",
));
assert!(provider_pool_key_account_quota_exhausted(
&sample_key(Some(json!({
"codex": {
@@ -758,6 +768,35 @@ mod tests {
}))),
"codex",
));
assert!(!provider_pool_key_account_quota_exhausted(
&sample_key(Some(json!({
"codex": {
"allowed": true,
"primary_used_percent": 100.0
}
}))),
"codex",
));
let mut explicit_codex_limit = sample_key(None);
explicit_codex_limit.status_snapshot = Some(json!({
"quota": {
"version": 2,
"provider_type": "codex",
"exhausted": true,
"allowed": false,
"limit_reached": true,
"usage_ratio": 0.91,
"windows": [{
"code": "weekly",
"used_ratio": 0.91
}]
}
}));
assert!(provider_pool_key_account_quota_exhausted(
&explicit_codex_limit,
"codex",
));
}
#[test]
@@ -817,6 +856,8 @@ mod tests {
&sample_key(Some(json!({
"codex": {
"updated_at": now.saturating_sub(600),
"allowed": false,
"limit_reached": true,
"primary_used_percent": 100.0,
"primary_reset_at": now.saturating_sub(60)
}
@@ -835,6 +876,79 @@ mod tests {
));
}
#[test]
fn codex_explicit_quota_block_is_hard_until_reset() {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.expect("system time should be after unix epoch")
.as_secs();
assert!(provider_pool_key_quota_hard_blocked(
&sample_key(Some(json!({
"codex": {
"updated_at": now,
"allowed": false,
"limit_reached": true,
"primary_reset_at": now.saturating_add(3600)
}
}))),
"codex",
));
assert!(!provider_pool_key_quota_hard_blocked(
&sample_key(Some(json!({
"codex": {
"updated_at": now,
"primary_used_percent": 100.0,
"primary_reset_at": now.saturating_add(3600)
}
}))),
"codex",
));
assert!(!provider_pool_key_quota_hard_blocked(
&sample_key(Some(json!({
"codex": {
"updated_at": now.saturating_sub(600),
"allowed": false,
"limit_reached": true,
"primary_reset_at": now.saturating_sub(60)
}
}))),
"codex",
));
let mut snapshot_blocked = sample_key(None);
snapshot_blocked.status_snapshot = Some(json!({
"quota": {
"version": 2,
"provider_type": "codex",
"exhausted": true,
"allowed": false,
"limit_reached": true,
"usage_ratio": 0.91,
"reset_at": now.saturating_add(3600)
}
}));
assert!(provider_pool_key_quota_hard_blocked(
&snapshot_blocked,
"codex",
));
snapshot_blocked.status_snapshot = Some(json!({
"quota": {
"version": 2,
"provider_type": "codex",
"exhausted": true,
"allowed": false,
"limit_reached": true,
"usage_ratio": 0.91,
"reset_at": now.saturating_sub(60)
}
}));
assert!(!provider_pool_key_quota_hard_blocked(
&snapshot_blocked,
"codex",
));
}
#[test]
fn grok_quota_tier_boundaries_match_pool_modes() {
assert_eq!(
@@ -64,6 +64,7 @@ pub trait ProviderPoolAdapter: Send + Sync {
quota_reset_seconds: provider_pool_quota_reset_seconds(input.key),
account_blocked: provider_pool_account_blocked(input.key),
quota_exhausted: self.quota_exhausted(input),
quota_hard_blocked: self.quota_hard_blocked(input),
..PoolMemberSignals::default()
}
}
@@ -72,6 +73,10 @@ pub trait ProviderPoolAdapter: Send + Sync {
provider_pool_quota_snapshot_exhausted_decision(input.key, input.provider_type)
.unwrap_or(false)
}
fn quota_hard_blocked(&self, _input: &ProviderPoolMemberInput<'_>) -> bool {
false
}
}
pub(crate) fn provider_pool_matching_endpoint<F>(
@@ -11,8 +11,9 @@ use crate::provider::{
};
use crate::quota::{
provider_pool_current_unix_secs, provider_pool_json_bool, provider_pool_json_f64,
provider_pool_metadata_bucket, provider_pool_quota_snapshot_exhausted_decision,
provider_pool_reset_deadline_elapsed, provider_pool_timestamp_unix_secs,
provider_pool_member_quota_snapshot, provider_pool_metadata_bucket,
provider_pool_quota_snapshot_exhausted_decision, provider_pool_reset_deadline_elapsed,
provider_pool_timestamp_unix_secs,
};
use crate::quota_refresh::ProviderPoolQuotaRequestSpec;
@@ -48,6 +49,23 @@ impl ProviderPoolAdapter for CodexProviderPoolAdapter {
}
fn quota_exhausted(&self, input: &ProviderPoolMemberInput<'_>) -> bool {
if let Some(quota_snapshot) =
provider_pool_member_quota_snapshot(input.key, input.provider_type)
{
let explicitly_exhausted = provider_pool_json_bool(quota_snapshot.get("allowed"))
== Some(false)
|| provider_pool_json_bool(quota_snapshot.get("limit_reached")) == Some(true);
if explicitly_exhausted {
let observed_at = provider_pool_timestamp_unix_secs(
quota_snapshot
.get("observed_at")
.or_else(|| quota_snapshot.get("updated_at")),
);
return !provider_pool_current_unix_secs().is_some_and(|now_unix_secs| {
provider_pool_reset_deadline_elapsed(quota_snapshot, observed_at, now_unix_secs)
});
}
}
if let Some(exhausted) =
provider_pool_quota_snapshot_exhausted_decision(input.key, input.provider_type)
{
@@ -57,6 +75,10 @@ impl ProviderPoolAdapter for CodexProviderPoolAdapter {
.is_some_and(quota_exhausted_from_bucket)
}
fn quota_hard_blocked(&self, input: &ProviderPoolMemberInput<'_>) -> bool {
codex_explicit_quota_block_active(input.key, input.provider_type)
}
fn quota_refresh_endpoint(
&self,
endpoints: &[StoredProviderCatalogEndpoint],
@@ -72,6 +94,35 @@ impl ProviderPoolAdapter for CodexProviderPoolAdapter {
}
}
fn codex_explicit_quota_block_active(
key: &aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey,
provider_type: &str,
) -> bool {
let Some(quota_snapshot) = provider_pool_member_quota_snapshot(key, provider_type) else {
return provider_pool_metadata_bucket(key.upstream_metadata.as_ref(), provider_type)
.is_some_and(|bucket| {
(provider_pool_json_bool(bucket.get("allowed")) == Some(false)
|| provider_pool_json_bool(bucket.get("limit_reached")) == Some(true))
&& !["primary", "secondary"]
.into_iter()
.any(|prefix| codex_window_reset_elapsed(bucket, prefix))
});
};
let explicitly_blocked = provider_pool_json_bool(quota_snapshot.get("allowed")) == Some(false)
|| provider_pool_json_bool(quota_snapshot.get("limit_reached")) == Some(true);
if !explicitly_blocked {
return false;
}
let observed_at = provider_pool_timestamp_unix_secs(
quota_snapshot
.get("observed_at")
.or_else(|| quota_snapshot.get("updated_at")),
);
!provider_pool_current_unix_secs().is_some_and(|now_unix_secs| {
provider_pool_reset_deadline_elapsed(quota_snapshot, observed_at, now_unix_secs)
})
}
fn build_codex_wham_headers(
resolved_oauth_auth: Option<(String, String)>,
decrypted_api_key: Option<&str>,
@@ -248,6 +299,19 @@ fn codex_window_used_percent_exhausted(bucket: &Map<String, Value>, prefix: &str
}
pub(crate) fn quota_exhausted_from_bucket(bucket: &Map<String, Value>) -> bool {
let allowed = provider_pool_json_bool(bucket.get("allowed"));
let limit_reached = provider_pool_json_bool(bucket.get("limit_reached"));
if allowed == Some(false) || limit_reached == Some(true) {
let reset_elapsed = ["primary", "secondary"]
.into_iter()
.any(|prefix| codex_window_reset_elapsed(bucket, prefix));
if !reset_elapsed {
return true;
}
}
if allowed == Some(true) || limit_reached == Some(false) {
return false;
}
if provider_pool_json_bool(bucket.get("credits_unlimited")) == Some(true) {
return false;
}
+12
View File
@@ -18,6 +18,18 @@ pub fn provider_pool_key_account_quota_exhausted(
})
}
pub fn provider_pool_key_quota_hard_blocked(
key: &StoredProviderCatalogKey,
provider_type: &str,
) -> bool {
let adapter = ProviderPoolService::with_builtin_adapters().adapter(provider_type);
adapter.quota_hard_blocked(&ProviderPoolMemberInput {
provider_type,
key,
auth_config: None,
})
}
pub fn provider_pool_member_quota_snapshot<'a>(
key: &'a StoredProviderCatalogKey,
provider_type: &str,
@@ -10,7 +10,8 @@ use aether_data_contracts::repository::usage::{
PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY, PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY,
PROVIDER_REASONING_EFFORT_METADATA_KEY, PROVIDER_SERVICE_TIER_METADATA_KEY,
REQUESTED_REASONING_EFFORT_METADATA_KEY, ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY,
ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY,
ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY, WEBSOCKET_MODE_METADATA_KEY,
WEBSOCKET_TRANSPORT_METADATA_KEY,
};
use serde_json::{json, Map, Value};
@@ -348,6 +349,8 @@ fn copy_allowed_metadata_fields(source: &Map<String, Value>, target: &mut Map<St
copy_bool(source, target, UPSTREAM_IS_STREAM_KEY);
copy_non_null_value(source, target, "client_session_affinity");
copy_bool(source, target, "api_key_is_standalone");
copy_bool(source, target, WEBSOCKET_MODE_METADATA_KEY);
copy_non_empty_string(source, target, WEBSOCKET_TRANSPORT_METADATA_KEY);
copy_non_empty_string(source, target, "request_path");
copy_non_empty_string(source, target, "request_query_string");
copy_non_empty_string(source, target, "request_path_and_query");
@@ -402,6 +405,8 @@ fn move_allowed_metadata_fields(mut source: Map<String, Value>, target: &mut Map
remove_bool(&mut source, target, UPSTREAM_IS_STREAM_KEY);
remove_non_null_value(&mut source, target, "client_session_affinity");
remove_bool(&mut source, target, "api_key_is_standalone");
remove_bool(&mut source, target, WEBSOCKET_MODE_METADATA_KEY);
remove_non_empty_string(&mut source, target, WEBSOCKET_TRANSPORT_METADATA_KEY);
remove_non_empty_string(&mut source, target, "request_path");
remove_non_empty_string(&mut source, target, "request_query_string");
remove_non_empty_string(&mut source, target, "request_path_and_query");
@@ -907,6 +912,24 @@ mod tests {
);
}
#[test]
fn sanitizes_websocket_transport_metadata() {
let metadata = sanitize_usage_request_metadata(Some(json!({
"websocket_mode": true,
"websocket_transport": "responses",
"untrusted_field": "drop-me",
})))
.expect("WebSocket metadata should remain");
assert_eq!(
metadata,
json!({
"websocket_mode": true,
"websocket_transport": "responses",
})
);
}
#[test]
fn sanitizes_request_path_query_metadata() {
let metadata = sanitize_usage_request_metadata(Some(json!({
+34 -1
View File
@@ -2,7 +2,10 @@ use std::collections::BTreeMap;
use aether_ai_formats::UPSTREAM_IS_STREAM_KEY;
use aether_contracts::{ExecutionPlan, ExecutionTelemetry};
use aether_data_contracts::repository::usage::{UpsertUsageRecord, UsageBodyCaptureState};
use aether_data_contracts::repository::usage::{
UpsertUsageRecord, UsageBodyCaptureState, WEBSOCKET_MODE_METADATA_KEY,
WEBSOCKET_TRANSPORT_METADATA_KEY,
};
use aether_data_contracts::DataLayerError;
use base64::Engine as _;
use serde_json::{json, Map, Value};
@@ -2114,6 +2117,18 @@ fn build_runtime_request_metadata_seed_from_parts(
Value::Bool(api_key_is_standalone),
);
}
if let Some(websocket_mode) = context_bool(context, WEBSOCKET_MODE_METADATA_KEY) {
metadata.insert(
WEBSOCKET_MODE_METADATA_KEY.to_string(),
Value::Bool(websocket_mode),
);
}
if let Some(websocket_transport) = context_string(context, WEBSOCKET_TRANSPORT_METADATA_KEY) {
metadata.insert(
WEBSOCKET_TRANSPORT_METADATA_KEY.to_string(),
Value::String(websocket_transport),
);
}
let provider_source_bytes = provider_request_body_base64.and_then(decoded_base64_len_hint);
append_runtime_body_capture_metadata(
&mut metadata,
@@ -3763,6 +3778,8 @@ mod tests {
Some(&json!({
"candidate_id": "cand-pending-event-1",
"candidate_index": 3,
"websocket_mode": true,
"websocket_transport": "responses",
"original_request_body": {"messages": [{"content": "omit me"}]},
"provider_request_body": {
"input": [{"type": "compaction_trigger"}]
@@ -3785,6 +3802,22 @@ mod tests {
assert!(record.provider_request_body.is_none());
assert_eq!(record.candidate_id.as_deref(), Some("cand-pending-event-1"));
assert_eq!(record.candidate_index, Some(3));
assert_eq!(
record
.request_metadata
.as_ref()
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("websocket_mode")),
Some(&json!(true))
);
assert_eq!(
record
.request_metadata
.as_ref()
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("websocket_transport")),
Some(&json!("responses"))
);
}
#[test]
+194
View File
@@ -0,0 +1,194 @@
# WebSocket Mode
The Responses API supports a WebSocket mode for long-running, tool-call-heavy workflows. In this mode, you keep a persistent connection to `/v1/responses` and continue each turn by sending only new input items plus `previous_response_id`.
WebSocket mode is compatible with both Zero Data Retention (ZDR) and `store=false`.
## Why use WebSocket mode
WebSocket mode is most useful when a workflow involves many model-tool round trips (for example, agentic coding or orchestration loops with repeated tool calls).
Because the connection stays open and each turn sends only incremental input, WebSocket mode reduces per-turn continuation overhead and improves end-to-end latency across long chains. For rollouts with 20+ tool calls, we have seen up to roughly 40% faster end-to-end execution.
## Connect and create responses
In WebSocket mode, start each turn by sending a `response.create` event from the client. The payload mirrors the normal [Responses create body](https://developers.openai.com/api/reference/resources/responses/methods/create), except that transport-specific fields like `stream` and `background` are not used.
```python
from websocket import create_connection
import json
import os
ws = create_connection(
"wss://api.openai.com/v1/responses",
header=[
f"Authorization: Bearer {os.environ['OPENAI_API_KEY']}",
],
)
ws.send(
json.dumps(
{
"type": "response.create",
"model": "gpt-5.6",
"store": False,
"input": [
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "Find fizz_buzz()"}],
}
],
"tools": [],
}
)
)
```
Clients can optionally warm up request state by sending `response.create` with `generate: false`. This is useful when you already know the tools, instructions, and/or custom messages you plan to send with an upcoming turn. `generate: false` does not return a model output, but prepares request state so the next generated turn can start faster. The warmup request returns a response ID that you can chain from with `previous_response_id`, including on later turns in a response chain. The next section explains how to continue a session using `previous_response_id` and incremental inputs.
## Continue with incremental inputs
To continue a run, send another `response.create` with:
- `previous_response_id` set to the prior response ID.
- `input` containing only new items (for example, tool outputs and the next user message).
```python
ws.send(
json.dumps(
{
"type": "response.create",
"model": "gpt-5.6",
"store": False,
"previous_response_id": "resp_123",
"input": [
{
"type": "function_call_output",
"call_id": "call_123",
"output": "tool result",
},
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "Now optimize it."}],
},
],
"tools": [],
}
)
)
```
## How continuation works
WebSocket mode uses the same `previous_response_id` chaining semantics as HTTP mode, but it adds a lower-latency continuation path on the active socket.
On an active WebSocket connection, the service keeps one previous-response state in a connection-local in-memory cache (the most recent response). Continuing from that most recent response is fast because the service can reuse connection-local state. Because the previous-response state is retained only in memory and is not written to disk, you can use WebSocket mode in a way that is compatible with `store=false` and Zero Data Retention (ZDR).
If a `previous_response_id` is not in the in-memory cache, behavior depends on whether you store responses:
- With `store=true`, the service may hydrate older response IDs from persisted state when available. Continuation can still work, but it usually loses the in-memory latency benefit.
- With `store=false` (including ZDR), there is no persisted fallback. If the ID is uncached, the request returns `previous_response_not_found`.
If a turn fails (`4xx` or `5xx`), the service evicts the referenced `previous_response_id` from the connection-local cache. This prevents reusing stale cached state for that failed continuation.
## Compaction and creating new responses
If you are using compaction, there are two different continuation patterns:
### Server-side compaction (`context_management`)
When you enable server-side compaction (`context_management` with `compact_threshold`), compaction happens during normal `/responses` generation. In WebSocket mode, you continue the same way you normally do: send the next `response.create` with the latest `previous_response_id` and only new input items.
### Standalone `/responses/compact`
The standalone [`/responses/compact` endpoint](https://developers.openai.com/api/docs/api-reference/responses/compact) returns a new compacted input window, not a response ID. After compaction, create a new response on your WebSocket connection using the compacted window as `input` (plus the next user/tool items).
Start a new chain by omitting `previous_response_id` or setting it to `null`. Pass the compacted output as-is; do not prune the returned window.
```python
# Compact your current window (HTTP call)
compacted = client.responses.compact(
model="gpt-5.6",
input=long_input_items_array,
)
# Start a new response on the WebSocket using the compacted window
ws.send(
json.dumps(
{
"type": "response.create",
"model": "gpt-5.6",
"store": False,
"input": [
*compacted.output,
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "Continue from here."}],
},
],
"tools": [],
}
)
)
```
## Connection behavior and limits
- Server events and ordering match the existing Responses streaming event model.
- A single WebSocket connection can receive multiple `response.create` messages, but it runs them sequentially (one in-flight response at a time).
- No multiplexing support today. Use multiple connections if you need parallel runs.
- Connection duration is limited to 60 minutes. Reconnect when the limit is reached.
- Aether binds each upstream WebSocket to one selected provider key. A provider must explicitly enable the standard Responses WebSocket capability and expose an `openai:responses` endpoint before it is eligible for this bridge.
- The Codex adapter additionally watches Codex quota events. A `usage_limit_reached` terminal error immediately marks the bound account unavailable. If the client has not received a standard `response.*` event and the request has no `previous_response_id`, Aether retries that one turn once on another eligible key without closing the public socket.
- After a standard response event has reached the client, after a retry has already been attempted, or for a request using `previous_response_id`, Aether forwards the provider terminal error and detaches only the exhausted upstream. If the upstream closes immediately after the quota signal, Aether emits a recoverable gateway error instead. The public WebSocket stays open so a later independent `response.create` can select another key.
- Aether does not transparently move an existing response chain to another provider key. Connection-local `previous_response_id` state cannot be transferred safely, especially with `store=false`/ZDR; send a new request with complete input after an exhausted continuation.
## Reconnect and recover
When a connection closes (or hits the 60-minute limit), open a new WebSocket connection and continue with one of these patterns:
1. If your prior response is persisted (`store=true`) and you have a valid response ID, continue with `previous_response_id` and new input items.
2. If you cannot continue the chain (for example, `store=false`/ZDR or `previous_response_not_found`), start a new response by setting `previous_response_id` to `null` (or omitting it) and send the full input context for the next turn.
3. If you compacted context with `/responses/compact`, use the returned compacted window as the base `input` for that new response, then append the latest user/tool items.
## Errors to handle
`previous_response_not_found`
```json
{
"type": "error",
"status": 400,
"error": {
"code": "previous_response_not_found",
"message": "Previous response with id 'resp_abc' not found.",
"param": "previous_response_id"
}
}
```
`websocket_connection_limit_reached`
```json
{
"type": "error",
"error": {
"type": "invalid_request_error",
"code": "websocket_connection_limit_reached",
"message": "Responses websocket connection limit reached (60 minutes). Create a new websocket connection to continue."
},
"status": 400
}
```
## Related guides
- [Conversation state](https://developers.openai.com/api/docs/guides/conversation-state)
- [Streaming API responses](https://developers.openai.com/api/docs/guides/streaming-responses)
- [Responses streaming events reference](https://developers.openai.com/api/docs/api-reference/responses-streaming)