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
@@ -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://[email protected]/responses", "invalid").is_err());
}
}