fix(ws): harden Responses connection lifecycle

Revalidate control policy per turn, isolate downstream credentials, and make planner/turn ownership cancellation-safe.

Preserve opaque protocol events, align configurable timeout semantics, and extend end-to-end security and settlement coverage.
This commit is contained in:
ZheFox
2026-08-17 18:50:29 +08:00
parent 4a0775c4ea
commit c8118edf36
42 changed files with 4017 additions and 1276 deletions
@@ -1,11 +1,16 @@
//! Authenticated public WebSocket upgrade admission shared by AI adapters.
use std::future::Future;
use std::net::SocketAddr;
use std::net::{IpAddr, SocketAddr};
use axum::body::Body;
use axum::extract::ws::{WebSocket, WebSocketUpgrade};
use axum::http::{HeaderMap, Method, Response, StatusCode, Uri};
use axum::http::header::{
AUTHORIZATION, CONNECTION, COOKIE, HOST, PROXY_AUTHORIZATION, TE, TRAILER, TRANSFER_ENCODING,
UPGRADE,
};
use axum::http::uri::PathAndQuery;
use axum::http::{HeaderMap, HeaderName, Method, Response, StatusCode, Uri};
use tracing::{info, warn};
use crate::api::response::{
@@ -13,7 +18,8 @@ use crate::api::response::{
build_local_overloaded_response,
};
use crate::control::{
trusted_auth_local_rejection, GatewayControlDecision, GatewayLocalAuthRejection,
trusted_auth_local_rejection, GatewayControlDecision, GatewayCredentialCarrier,
GatewayLocalAuthRejection,
};
use crate::handlers::proxy::websocket::session::{WebSocketSessionLimits, WEBSOCKET_LOG_TRANSPORT};
use crate::handlers::shared::ip_rules_allow;
@@ -28,8 +34,10 @@ pub(crate) struct WebSocketRequestContext {
pub(crate) headers: HeaderMap,
pub(crate) uri: Uri,
pub(crate) remote_addr: SocketAddr,
/// Effective client IP resolved once from the authenticated Upgrade. Every
/// turn re-checks live API-key/admin IP policy against this immutable fact.
pub(crate) client_ip: IpAddr,
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.
@@ -40,7 +48,6 @@ pub(crate) struct WebSocketRequestContext {
#[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.
@@ -81,7 +88,7 @@ where
&trace_id,
)
.await?;
let Some(decision) = request_context.control_decision else {
let Some(mut decision) = request_context.control_decision else {
return build_local_http_error_response(
&trace_id,
None,
@@ -92,6 +99,28 @@ where
if let Some(rejection) = trusted_auth_local_rejection(Some(&decision), &headers) {
return build_local_auth_rejection_response(&trace_id, Some(&decision), &rejection);
}
// Browsers attach cookies to WebSocket handshakes automatically and the
// WebSocket API does not let callers add an Authorization header. A
// cookie-only public upgrade would therefore be vulnerable to cross-site
// WebSocket hijacking unless every deployment maintained an Origin
// allowlist. Explicit API-key/bearer credentials (or trusted internal
// auth resolved by the control plane) remain supported.
if !websocket_credential_carrier_is_allowed(decision.gateway_credential_carrier) {
warn!(
event_name = "ai_websocket_cookie_only_auth_rejected",
log_type = "security",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %trace_id,
client_ip = %client_ip,
"gateway rejected cookie-only public WebSocket authentication"
);
return build_local_auth_rejection_response(
&trace_id,
Some(&decision),
&GatewayLocalAuthRejection::InvalidApiKey,
);
}
let Some(auth_context) = decision.auth_context.as_ref() else {
return build_local_auth_rejection_response(
&trace_id,
@@ -99,7 +128,10 @@ where
&GatewayLocalAuthRejection::InvalidApiKey,
);
};
if !auth_context.access_allowed {
if !auth_context.access_allowed
|| auth_context.user_id.trim().is_empty()
|| auth_context.api_key_id.trim().is_empty()
{
return build_local_auth_rejection_response(
&trace_id,
Some(&decision),
@@ -116,22 +148,6 @@ where
);
}
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) => {
@@ -155,13 +171,21 @@ where
}
};
// Authentication has consumed the downstream credentials. From this
// point on the URI and headers become planner input, so retain neither an
// API key from the query string nor client authentication/handshake
// headers. Provider authentication is added independently by the
// planner and is therefore unaffected by this boundary.
let uri = websocket_planning_uri(&uri);
decision.public_query_string = uri.query().map(ToOwned::to_owned);
let headers = websocket_planning_headers(headers);
let context = WebSocketRequestContext {
trace_id,
headers,
uri,
remote_addr,
client_ip,
decision,
rpm_bypassed: ip_whitelisted,
websocket_connection_permit,
};
Ok(ws
@@ -173,6 +197,106 @@ where
}))
}
fn websocket_credential_carrier_is_allowed(carrier: Option<GatewayCredentialCarrier>) -> bool {
carrier != Some(GatewayCredentialCarrier::CookieHeader)
}
fn websocket_planning_uri(uri: &Uri) -> Uri {
let Some(query) = uri.query() else {
return uri.clone();
};
let mut retained = Vec::new();
let mut removed_sensitive_value = false;
for (name, value) in url::form_urlencoded::parse(query.as_bytes()) {
if websocket_query_parameter_is_sensitive(name.as_ref()) {
removed_sensitive_value = true;
} else {
retained.push((name.into_owned(), value.into_owned()));
}
}
if !removed_sensitive_value {
return uri.clone();
}
let mut serializer = url::form_urlencoded::Serializer::new(String::new());
serializer.extend_pairs(retained.iter().map(|(name, value)| (name, value)));
let retained_query = serializer.finish();
let path_and_query = if retained_query.is_empty() {
uri.path().to_string()
} else {
format!("{}?{retained_query}", uri.path())
};
let path_and_query = path_and_query
.parse::<PathAndQuery>()
.expect("a valid URI path plus form-encoded query must remain valid");
let mut parts = uri.clone().into_parts();
parts.path_and_query = Some(path_and_query);
Uri::from_parts(parts).expect("replacing only path-and-query must preserve a valid URI")
}
fn websocket_query_parameter_is_sensitive(name: &str) -> bool {
matches!(
name.to_ascii_lowercase().as_str(),
"key" | "api_key" | "api-key" | "access_token" | "authorization" | "token"
)
}
fn websocket_planning_headers(mut headers: HeaderMap) -> HeaderMap {
// RFC 9110 permits Connection to name additional hop-by-hop fields. Read
// those names before removing Connection itself.
let connection_scoped_names = headers
.get_all(CONNECTION)
.iter()
.filter_map(|value| value.to_str().ok())
.flat_map(|value| value.split(','))
.filter_map(|name| HeaderName::from_bytes(name.trim().as_bytes()).ok())
.collect::<Vec<_>>();
for name in connection_scoped_names {
headers.remove(name);
}
for name in [
AUTHORIZATION,
CONNECTION,
COOKIE,
HOST,
PROXY_AUTHORIZATION,
TE,
TRAILER,
TRANSFER_ENCODING,
UPGRADE,
] {
headers.remove(name);
}
for name in [
"api-key",
"keep-alive",
"proxy-connection",
"x-api-key",
"x-goog-api-key",
crate::constants::GATEWAY_HEADER,
crate::constants::TRUSTED_AUTH_USER_ID_HEADER,
crate::constants::TRUSTED_AUTH_API_KEY_ID_HEADER,
crate::constants::TRUSTED_AUTH_BALANCE_HEADER,
crate::constants::TRUSTED_AUTH_ACCESS_ALLOWED_HEADER,
crate::constants::TRUSTED_ADMIN_USER_ID_HEADER,
crate::constants::TRUSTED_ADMIN_USER_ROLE_HEADER,
crate::constants::TRUSTED_ADMIN_SESSION_ID_HEADER,
crate::constants::TRUSTED_ADMIN_MANAGEMENT_TOKEN_ID_HEADER,
] {
headers.remove(name);
}
let websocket_managed_names = headers
.keys()
.filter(|name| name.as_str().starts_with("sec-websocket-"))
.cloned()
.collect::<Vec<_>>();
for name in websocket_managed_names {
headers.remove(name);
}
headers
}
fn websocket_admission_error_response(
trace_id: &str,
decision: &GatewayControlDecision,
@@ -291,3 +415,113 @@ impl Drop for WebSocketConnectionLog {
);
}
}
#[cfg(test)]
mod tests {
use axum::http::header::{
AUTHORIZATION, CONNECTION, COOKIE, HOST, ORIGIN, SEC_WEBSOCKET_KEY, UPGRADE, USER_AGENT,
};
use axum::http::{HeaderMap, HeaderValue, Uri};
use super::{
websocket_credential_carrier_is_allowed, websocket_planning_headers, websocket_planning_uri,
};
use crate::control::GatewayCredentialCarrier;
#[test]
fn planning_uri_removes_query_credentials_without_losing_safe_parameters() {
let uri: Uri = "/v1/responses?key=downstream-secret&client_hint=a%20b&token=also-secret"
.parse()
.expect("request URI should parse");
let sanitized = websocket_planning_uri(&uri);
assert_eq!(sanitized.path(), "/v1/responses");
assert_eq!(sanitized.query(), Some("client_hint=a+b"));
assert!(!sanitized.to_string().contains("downstream-secret"));
assert!(!sanitized.to_string().contains("also-secret"));
}
#[test]
fn planning_uri_leaves_an_uncredentialed_query_byte_for_byte_unchanged() {
let uri: Uri = "/v1/responses?client_hint=a%20b&empty="
.parse()
.expect("request URI should parse");
assert_eq!(websocket_planning_uri(&uri), uri);
}
#[test]
fn planning_headers_drop_client_auth_cookie_and_websocket_transport_state() {
let mut headers = HeaderMap::new();
headers.insert(
AUTHORIZATION,
HeaderValue::from_static("Bearer client-secret"),
);
headers.insert(COOKIE, HeaderValue::from_static("session=client-secret"));
headers.insert("x-api-key", HeaderValue::from_static("client-secret"));
headers.insert(HOST, HeaderValue::from_static("gateway.example"));
headers.insert(
CONNECTION,
HeaderValue::from_static("keep-alive, Upgrade, x-connection-secret"),
);
headers.insert(UPGRADE, HeaderValue::from_static("websocket"));
headers.insert(SEC_WEBSOCKET_KEY, HeaderValue::from_static("handshake-key"));
headers.insert(
"sec-websocket-future-field",
HeaderValue::from_static("future-handshake-value"),
);
headers.insert(
"x-connection-secret",
HeaderValue::from_static("connection-secret"),
);
headers.insert(ORIGIN, HeaderValue::from_static("https://client.example"));
headers.insert(USER_AGENT, HeaderValue::from_static("codex-cli/test"));
headers.insert("x-client-hint", HeaderValue::from_static("safe"));
let sanitized = websocket_planning_headers(headers);
for name in [
AUTHORIZATION.as_str(),
COOKIE.as_str(),
"x-api-key",
HOST.as_str(),
CONNECTION.as_str(),
UPGRADE.as_str(),
SEC_WEBSOCKET_KEY.as_str(),
"sec-websocket-future-field",
"x-connection-secret",
] {
assert!(sanitized.get(name).is_none(), "{name} must not survive");
}
assert_eq!(
sanitized.get(ORIGIN),
Some(&HeaderValue::from_static("https://client.example"))
);
assert_eq!(
sanitized.get(USER_AGENT),
Some(&HeaderValue::from_static("codex-cli/test"))
);
assert_eq!(
sanitized.get("x-client-hint"),
Some(&HeaderValue::from_static("safe"))
);
}
#[test]
fn websocket_auth_requires_an_explicit_credential_instead_of_cookie_only() {
assert!(!websocket_credential_carrier_is_allowed(Some(
GatewayCredentialCarrier::CookieHeader
)));
for carrier in [
None,
Some(GatewayCredentialCarrier::AuthorizationBearer),
Some(GatewayCredentialCarrier::XApiKey),
Some(GatewayCredentialCarrier::ApiKey),
Some(GatewayCredentialCarrier::XGoogApiKey),
Some(GatewayCredentialCarrier::QueryKey),
] {
assert!(websocket_credential_carrier_is_allowed(carrier));
}
}
}
@@ -46,6 +46,25 @@ pub(super) enum ResponsesWebSocketRebindSafety {
Unsafe { reason: &'static str },
}
/// How an upstream text frame crosses the public Responses WebSocket boundary.
///
/// The normal path is deliberately byte-opaque: callers forward the parsed
/// frame's original text without rebuilding it from a gateway-owned schema.
/// Codex is the only adapter that may peel its documented private batch
/// envelope. Even then, the retained events are borrowed whole so unknown
/// `response.*` event types and unknown fields survive unchanged.
#[derive(Debug, Clone, PartialEq)]
pub(super) enum ResponsesWebSocketRelayDirective<'a> {
/// Forward the provider frame's original text exactly as received.
ForwardOriginal,
/// The provider frame was a private batch envelope. Forward each retained
/// event in document order by serializing the complete borrowed value.
ForwardEvents(Vec<&'a Value>),
/// The entire frame was an explicitly recognized provider-private
/// envelope and therefore has no public event to relay.
SuppressProviderPrivate,
}
/// 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.
@@ -69,6 +88,15 @@ pub(super) trait ResponsesWebSocketProtocolAdapter: Send + Sync {
/// observably ambiguous to the client.
fn rebind_safety_for_upstream_event(&self, event: &Value) -> ResponsesWebSocketRebindSafety;
/// Selects the public relay shape without projecting a provider event
/// through an Aether-owned field or event-type allowlist.
fn relay_directive_for_upstream_event<'a>(
&self,
_event: &'a Value,
) -> ResponsesWebSocketRelayDirective<'a> {
ResponsesWebSocketRelayDirective::ForwardOriginal
}
/// 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)
@@ -169,7 +197,12 @@ pub(super) fn is_standard_responses_event(event: &Value) -> bool {
#[cfg(test)]
mod tests {
use super::{resolve_responses_websocket_adapter, ResponsesWebSocketProtocolAdapter};
use serde_json::json;
use super::{
resolve_responses_websocket_adapter, ResponsesWebSocketProtocolAdapter,
ResponsesWebSocketRelayDirective,
};
use crate::orchestration::ResponsesWebSocketAdapter;
#[test]
@@ -183,4 +216,19 @@ mod tests {
"responses_websocket_handshake_failed"
);
}
#[test]
fn standard_adapter_always_forwards_future_events_opaquely() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
let event = json!({
"type": "response.future_capability.delta",
"delta": {"future_shape": [1, {"nested": true}]},
"unknown_top_level": {"must": "survive"},
});
assert_eq!(
adapter.relay_directive_for_upstream_event(&event),
ResponsesWebSocketRelayDirective::ForwardOriginal
);
}
}
@@ -7,6 +7,7 @@ use super::super::adapter::{
is_standard_responses_event, ResponsesWebSocketAdapterObservation,
ResponsesWebSocketDrainDirective, ResponsesWebSocketExclusionIdentity,
ResponsesWebSocketProtocolAdapter, ResponsesWebSocketRebindSafety,
ResponsesWebSocketRelayDirective,
};
use crate::ai_serving::AiExecutionDecision;
use crate::clock::current_unix_secs;
@@ -66,19 +67,45 @@ impl ResponsesWebSocketProtocolAdapter for CodexResponsesWebSocketAdapter {
}
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() {
let mut saw_event = false;
if event.get("type").and_then(Value::as_str).is_some() {
saw_event = true;
let safety = codex_direct_rebind_safety(event);
if matches!(safety, ResponsesWebSocketRebindSafety::Unsafe { .. }) {
return safety;
}
}
match event.get("chunks") {
Some(Value::Array(chunks)) => {
for chunk in chunks {
saw_event = true;
let safety = codex_direct_rebind_safety(chunk);
if matches!(safety, ResponsesWebSocketRebindSafety::Unsafe { .. }) {
return safety;
}
}
}
Some(_) => {
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);
None => {}
}
codex_direct_rebind_safety(event)
if saw_event {
ResponsesWebSocketRebindSafety::Safe
} else {
ResponsesWebSocketRebindSafety::Unsafe {
reason: "unrecognized_upstream_event",
}
}
}
fn relay_directive_for_upstream_event<'a>(
&self,
event: &'a Value,
) -> ResponsesWebSocketRelayDirective<'a> {
codex_relay_directive(event)
}
fn observe_upstream_event(
@@ -148,9 +175,15 @@ fn codex_direct_rebind_safety(event: &Value) -> ResponsesWebSocketRebindSafety {
// 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.
if event_type == "error"
&& event.pointer("/error/type").and_then(Value::as_str) == Some("usage_limit_reached")
&& parse_codex_rate_limits(event).is_some()
{
// This terminal quota event has not been relayed yet. It can trigger
// one transparent attempt on another key as long as no earlier public
// response event made the logical turn unsafe. If replanning fails,
// the connection layer forwards this exact upstream error instead of
// manufacturing a gateway continuation error.
return ResponsesWebSocketRebindSafety::Safe;
}
let reason = if is_standard_responses_event(event) {
@@ -161,6 +194,55 @@ fn codex_direct_rebind_safety(event: &Value) -> ResponsesWebSocketRebindSafety {
ResponsesWebSocketRebindSafety::Unsafe { reason }
}
fn codex_relay_directive(event: &Value) -> ResponsesWebSocketRelayDirective<'_> {
match event.get("chunks") {
Some(Value::Array(chunks)) if is_explicit_codex_batch_envelope(event) => {
let public_events = chunks
.iter()
.filter(|chunk| !is_codex_private_leaf_event(chunk))
.collect::<Vec<_>>();
if public_events.is_empty() {
ResponsesWebSocketRelayDirective::SuppressProviderPrivate
} else {
ResponsesWebSocketRelayDirective::ForwardEvents(public_events)
}
}
// A malformed or future shape is not proven private. Preserve it
// opaquely rather than guessing at a provider schema.
Some(_) => ResponsesWebSocketRelayDirective::ForwardOriginal,
None if is_codex_private_leaf_event(event) => {
ResponsesWebSocketRelayDirective::SuppressProviderPrivate
}
None => ResponsesWebSocketRelayDirective::ForwardOriginal,
}
}
/// Recognizes only Codex's private batch container. A type-less object must
/// contain exactly `chunks`; unknown siblings could be future public protocol
/// data and therefore force opaque forwarding. A named Codex private root may
/// carry provider metadata alongside its chunks and is safe to peel.
fn is_explicit_codex_batch_envelope(event: &Value) -> bool {
if is_codex_private_event_type(event) {
return true;
}
event.as_object().is_some_and(|object| {
object.len() == 1
&& object.contains_key("chunks")
&& event.get("type").and_then(Value::as_str).is_none()
})
}
fn is_codex_private_leaf_event(event: &Value) -> bool {
is_codex_private_event_type(event) && event.get("chunks").is_none()
}
fn is_codex_private_event_type(event: &Value) -> bool {
matches!(
event.get("type").and_then(Value::as_str),
Some("codex.rate_limits" | "codex.response.metadata")
)
}
fn parse_codex_rate_limits(event: &Value) -> Option<Value> {
aether_admin::provider::quota::parse_codex_websocket_rate_limits_response(
event,
@@ -174,7 +256,7 @@ mod tests {
use super::{
CodexResponsesWebSocketAdapter, ResponsesWebSocketProtocolAdapter,
ResponsesWebSocketRebindSafety,
ResponsesWebSocketRebindSafety, ResponsesWebSocketRelayDirective,
};
#[test]
@@ -245,7 +327,7 @@ mod tests {
}
#[test]
fn only_known_codex_pre_response_metadata_is_safe_to_rebind() {
fn only_known_codex_pre_response_signals_are_safe_to_rebind() {
let adapter = CodexResponsesWebSocketAdapter;
assert_eq!(
@@ -286,5 +368,105 @@ mod tests {
reason: "unrecognized_upstream_event"
}
);
assert_eq!(
adapter.rebind_safety_for_upstream_event(&json!({
"type": "error",
"error": {
"type": "usage_limit_reached",
"plan_type": "plus",
"resets_in_seconds": 3_600
},
"status_code": 429
})),
ResponsesWebSocketRebindSafety::Safe
);
assert_eq!(
adapter.rebind_safety_for_upstream_event(&json!({
"type": "error",
"error": {"type": "usage_limit_reached"}
})),
ResponsesWebSocketRebindSafety::Unsafe {
reason: "unrecognized_upstream_event"
}
);
assert_eq!(
adapter.rebind_safety_for_upstream_event(&json!({
"type": "response.future_capability.delta",
"chunks": [{"type": "codex.rate_limits"}]
})),
ResponsesWebSocketRebindSafety::Unsafe {
reason: "standard_response_event"
}
);
}
#[test]
fn codex_suppresses_only_explicit_private_events_and_envelopes() {
let adapter = CodexResponsesWebSocketAdapter;
for event in [
json!({"type": "codex.rate_limits", "rate_limits": {"allowed": true}}),
json!({"type": "codex.response.metadata", "account_hint": "private"}),
json!({"chunks": [
{"type": "codex.rate_limits"},
{"type": "codex.response.metadata"}
]}),
] {
assert_eq!(
adapter.relay_directive_for_upstream_event(&event),
ResponsesWebSocketRelayDirective::SuppressProviderPrivate
);
}
for event in [
json!({"type": "error", "error": {"type": "usage_limit_reached"}}),
json!({"type": "codex.future_private_maybe", "future": true}),
json!({"chunks": [], "future_envelope_field": {"must": "survive"}}),
json!({"type": "response.future.done", "future_capability": true}),
] {
assert_eq!(
adapter.relay_directive_for_upstream_event(&event),
ResponsesWebSocketRelayDirective::ForwardOriginal
);
}
}
#[test]
fn mixed_codex_batch_forwards_whole_non_private_events_in_order() {
let adapter = CodexResponsesWebSocketAdapter;
let event = json!({
"chunks": [
{
"type": "response.created",
"response": {"id": "resp_future"},
"future_created_field": {"opaque": true}
},
{"type": "codex.rate_limits", "account_hint": "private"},
{
"type": "response.future_capability.delta",
"future_capability": {"nested": [1, 2, 3]},
"sequence_number": 2
},
{"provider_future_event": {"unknown": "must be forwarded"}},
{
"type": "error",
"error": {"type": "future_error", "future_detail": 7}
}
]
});
let ResponsesWebSocketRelayDirective::ForwardEvents(events) =
adapter.relay_directive_for_upstream_event(&event)
else {
panic!("a mixed private envelope must retain all non-private events");
};
assert_eq!(events.len(), 4);
assert_eq!(events[0]["future_created_field"], json!({"opaque": true}));
assert_eq!(events[1]["future_capability"], json!({"nested": [1, 2, 3]}));
assert_eq!(
events[2]["provider_future_event"],
json!({"unknown": "must be forwarded"})
);
assert_eq!(events[3]["error"]["future_detail"], json!(7));
}
}
@@ -2,10 +2,10 @@
//!
//! 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.
//! settings, stable account headers, credentials, and the protocol adapter can
//! all change the connection that would receive the next event. Ordinary Codex
//! OAuth access-token refreshes retain the credential generation and therefore
//! do not unnecessarily replace an already-upgraded socket.
use std::collections::{BTreeMap, BTreeSet};
use std::fmt;
@@ -34,11 +34,15 @@ pub(super) struct UpstreamBindingIdentity {
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]>,
/// One-way identity for the credential generation used by this socket.
///
/// A provider key id identifies a catalog row, not the secret currently
/// stored in that row. Codex decisions carry a server-owned credential
/// generation which is stable across access-token refreshes but rotates
/// when the account/static/refresh credential is replaced. Other
/// decisions conservatively fingerprint the effective authentication
/// handshake values.
credential_fingerprint: [u8; 32],
proxy: Option<ProxySnapshot>,
transport_profile: Option<ResolvedTransportProfile>,
}
@@ -82,11 +86,8 @@ impl UpstreamBindingIdentity {
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());
let credential_fingerprint =
credential_binding_fingerprint(decision, &authentication_headers);
Ok(Self {
adapter_kind: adapter.kind(),
@@ -95,7 +96,7 @@ impl UpstreamBindingIdentity {
key_id: decision.key_id.clone(),
upstream_url,
handshake_headers,
auth_fingerprint,
credential_fingerprint,
proxy: effective_proxy_snapshot(decision.proxy.as_ref()),
transport_profile: decision.transport_profile.clone(),
})
@@ -127,6 +128,7 @@ fn authentication_header_names(decision: &AiExecutionDecision) -> BTreeSet<Strin
fn fingerprint_headers(headers: &BTreeMap<String, String>) -> [u8; 32] {
let mut hasher = Sha256::new();
hasher.update(b"aether-responses-websocket-auth-headers-v1");
for (name, value) in headers {
hasher.update((name.len() as u64).to_be_bytes());
hasher.update(name.as_bytes());
@@ -136,6 +138,69 @@ fn fingerprint_headers(headers: &BTreeMap<String, String>) -> [u8; 32] {
hasher.finalize().into()
}
/// Returns the non-secret credential identity represented by a planner
/// decision. The generation is emitted by Aether's trusted Codex planner from
/// provider-key metadata; it is not sourced from the downstream request.
fn credential_binding_fingerprint(
decision: &AiExecutionDecision,
authentication_headers: &BTreeMap<String, String>,
) -> [u8; 32] {
if decision
.provider_type
.as_deref()
.is_some_and(|provider_type| provider_type.trim().eq_ignore_ascii_case("codex"))
{
if let Some(generation) = decision
.report_context
.as_ref()
.and_then(|context| context.get("codex_credential_generation"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|generation| !generation.is_empty())
{
let mut hasher = Sha256::new();
hasher.update(b"aether-responses-websocket-codex-credential-generation-v1");
hasher.update((generation.len() as u64).to_be_bytes());
hasher.update(generation.as_bytes());
// Only a planner-owned Codex bearer access token is expected to
// rotate without changing credential generation. Compare the
// effective handshake value with the decision's original auth
// value: auth-config/routing/header overrides change only the
// former and therefore must force a rebind.
let stable_authentication_headers = authentication_headers
.iter()
.filter(|(name, value)| {
!is_planner_owned_codex_bearer(decision, name.as_str(), value.as_str())
})
.map(|(name, value)| (name.clone(), value.clone()))
.collect::<BTreeMap<_, _>>();
hasher.update(fingerprint_headers(&stable_authentication_headers));
return hasher.finalize().into();
}
}
// Fail closed when no trusted generation is available. Rebinding after an
// access-token change is preferable to sending a continuation over a
// socket authenticated with a credential that may have been replaced.
fingerprint_headers(authentication_headers)
}
fn is_planner_owned_codex_bearer(
decision: &AiExecutionDecision,
name: &str,
effective_value: &str,
) -> bool {
name.eq_ignore_ascii_case("authorization")
&& decision
.auth_header
.as_deref()
.is_some_and(|header| header.eq_ignore_ascii_case(name))
&& decision.auth_value.as_deref() == Some(effective_value)
&& effective_value
.get(.."bearer ".len())
.is_some_and(|scheme| scheme.eq_ignore_ascii_case("bearer "))
}
/// 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
@@ -313,18 +378,18 @@ mod tests {
assert_ne!(identity, changed_identity);
}
let mut rotated = base.clone();
rotated
let mut static_secret_rotated = base.clone();
static_secret_rotated
.provider_request_headers
.insert("Authorization".to_string(), "Bearer rotated".to_string());
assert_eq!(
assert_ne!(
identity,
UpstreamBindingIdentity::from_decision(adapter, &rotated).unwrap()
UpstreamBindingIdentity::from_decision(adapter, &static_secret_rotated).unwrap()
);
}
#[test]
fn stable_key_identity_ignores_custom_auth_value_rotation() {
fn stable_key_identity_rejects_custom_static_auth_value_rotation() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
let mut base = decision();
base.auth_header = Some("X-Provider-Token".to_string());
@@ -335,43 +400,125 @@ mod tests {
);
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!(
assert_ne!(
identity,
UpstreamBindingIdentity::from_decision(adapter, &rotated).unwrap()
);
}
#[test]
fn missing_key_identity_fingerprints_authentication_values() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
fn codex_access_token_refresh_reuses_the_same_credential_generation() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
let mut first = decision();
first.key_id = None;
first.provider_type = Some("codex".to_string());
first.report_context = Some(json!({
"codex_credential_generation": "credential-generation-1"
}));
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(
let mut access_token_refreshed = first;
access_token_refreshed.auth_value = Some("Bearer refreshed-access-token".to_string());
access_token_refreshed.provider_request_headers.insert(
"Authorization".to_string(),
"Bearer different-account-or-token".to_string(),
"Bearer refreshed-access-token".to_string(),
);
let changed_identity =
UpstreamBindingIdentity::from_decision(adapter, &same_account_rotation).unwrap();
assert_ne!(first_identity, changed_identity);
assert_eq!(
first_identity,
UpstreamBindingIdentity::from_decision(adapter, &access_token_refreshed).unwrap()
);
}
let mut non_auth_change = first;
non_auth_change
.provider_request_headers
.insert("X-Client".to_string(), "other-client".to_string());
#[test]
fn codex_authorization_override_changes_binding_with_the_same_generation() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
let mut first = decision();
first.provider_type = Some("codex".to_string());
first.report_context = Some(json!({
"codex_credential_generation": "credential-generation-1"
}));
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
// The planner-owned auth value remains unchanged while an effective
// auth-config/header override replaces the actual handshake value.
first.provider_request_headers.insert(
"Authorization".to_string(),
"Bearer endpoint-override".to_string(),
);
assert_ne!(
first_identity,
UpstreamBindingIdentity::from_decision(adapter, &non_auth_change).unwrap()
UpstreamBindingIdentity::from_decision(adapter, &first).unwrap()
);
}
#[test]
fn codex_credential_replacement_changes_binding_for_the_same_key_id() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
let mut first = decision();
first.provider_type = Some("codex".to_string());
first.report_context = Some(json!({
"codex_credential_generation": "credential-generation-1"
}));
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
let mut replaced = first;
replaced.provider_request_headers.insert(
"Authorization".to_string(),
"Bearer replacement-access-token".to_string(),
);
replaced.report_context = Some(json!({
"codex_credential_generation": "credential-generation-2"
}));
assert_ne!(
first_identity,
UpstreamBindingIdentity::from_decision(adapter, &replaced).unwrap()
);
}
#[test]
fn codex_custom_auth_rotation_changes_binding_with_the_same_generation() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
let mut first = decision();
first.provider_type = Some("codex".to_string());
first.auth_header = Some("X-Provider-Token".to_string());
first.provider_request_headers.insert(
"X-Provider-Token".to_string(),
"provider-token-1".to_string(),
);
first.report_context = Some(json!({
"codex_credential_generation": "credential-generation-1"
}));
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
first.provider_request_headers.insert(
"X-Provider-Token".to_string(),
"provider-token-2".to_string(),
);
assert_ne!(
first_identity,
UpstreamBindingIdentity::from_decision(adapter, &first).unwrap()
);
}
#[test]
fn missing_codex_credential_generation_fails_closed_on_auth_rotation() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
let mut first = decision();
first.provider_type = Some("codex".to_string());
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
first.provider_request_headers.insert(
"Authorization".to_string(),
"Bearer possibly-replaced-credential".to_string(),
);
assert_ne!(
first_identity,
UpstreamBindingIdentity::from_decision(adapter, &first).unwrap()
);
}
File diff suppressed because it is too large Load Diff
@@ -5,20 +5,18 @@ use std::time::Duration;
use axum::extract::ws::{Message as AxumWsMessage, WebSocket};
use futures_util::{SinkExt, StreamExt};
use serde_json::Value;
use tokio::time::sleep;
use wreq::ws::message::Message as WreqWsMessage;
use super::adapter::ResponsesWebSocketRelayDirective;
use super::client::{adapter_drain_ready, forward_client_message, RelayDisposition};
use super::frame::ParsedResponsesWebSocketFrame;
use super::frame::{encode_opaque_websocket_event, ParsedResponsesWebSocketFrame};
use super::lifecycle::{
await_pending_adapter_observation, finalize_active_turn, queue_turn_finalization,
settle_turn_finalization, ActiveProviderAttempt, PreviousAttemptSettled,
settle_turn_finalization, spawn_bounded_adapter_observation, PreviousAttemptSettled,
};
use super::quota::{
active_continuation_can_retry_from_full_input, detach_exhausted_upstream,
is_usage_limit_error_event, mark_active_response_retry_unsafe,
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, fatal_relay_policy, FatalRelaySignal, QuotaRelayAction, QuotaRelayFacts,
@@ -31,8 +29,7 @@ use super::turn::{
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,
wait_for_optional_deadline, CLOSE_INTERNAL_ERROR, CLOSE_TRY_AGAIN, WEBSOCKET_LOG_TRANSPORT,
};
use crate::handlers::proxy::websocket::transport::{
close_client_socket, send_client_message, send_gateway_error_with_status,
@@ -64,30 +61,10 @@ pub(super) async fn relay_bound_connection(
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.turn_state.attempt().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;
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;
@@ -117,31 +94,6 @@ pub(super) async fn relay_bound_connection(
).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;
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(
@@ -169,7 +121,15 @@ pub(super) async fn relay_bound_connection(
close_bound_upstream(bound).await;
break;
};
match forward_client_message(client_message, bound, client_socket, state, context).await {
match Box::pin(forward_client_message(
client_message,
bound,
client_socket,
state,
context,
))
.await
{
RelayDisposition::Continue => {}
RelayDisposition::Close => {
finalize_active_turn(
@@ -288,7 +248,7 @@ pub(super) async fn relay_bound_connection(
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 {
bound.pending_adapter_observation = Some(spawn_bounded_adapter_observation(async move {
adapter
.persist_upstream_observation(
&state_for_observation,
@@ -380,19 +340,19 @@ pub(super) async fn relay_bound_connection(
drain_ready: drain_for_adapter,
retry_current_turn: bound
.pending_adapter_drain
.is_some_and(|directive| directive.retry_current_turn),
.is_some_and(|directive| directive.retry_current_turn)
&& bound
.turn_state
.logical()
.is_some_and(|turn| turn.quota_retry_block_reason().is_none()),
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) {
// detach_attempt 保留 logical turn:重试是同一轮请求的下一个 attempt。
let retry_turn = bound
.turn_state
.detach_attempt()
.map(ActiveProviderAttempt::disarm);
let retry_turn = bound.turn_state.detach_attempt();
// 先结算旧 attempt 并等它落地,再规划下一个 attempt。两个理由:
//
// 1. 规划要读 health / adaptive / pool 状态,而这些正是旧
@@ -417,7 +377,14 @@ pub(super) async fn relay_bound_connection(
}
None => PreviousAttemptSettled::nothing_to_settle(),
};
if retry_active_turn_after_quota_exhaustion(bound, state, context, settled).await
// Planning and binding a replacement carries the complete
// scheduler/provider state machine. Keep that large future
// off the relay task's stack; the default Tokio/test worker
// stack is otherwise easy to exhaust on this rare branch.
if Box::pin(retry_active_turn_after_quota_exhaustion(
bound, state, context, settled,
))
.await
{
continue;
}
@@ -430,42 +397,9 @@ pub(super) async fn relay_bound_connection(
..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.turn_state.end().map(ActiveProviderAttempt::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;
}
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
continue;
}
if matches!(quota_relay_action, QuotaRelayAction::ForwardQuotaAndDetach) {
let detach_after_forward =
matches!(quota_relay_action, QuotaRelayAction::ForwardQuotaAndDetach);
if detach_after_forward && is_close {
let directive = bound
.pending_adapter_drain
.expect("adapter drain state should be present");
@@ -486,22 +420,127 @@ pub(super) async fn relay_bound_connection(
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
continue;
}
// 响应侧还原:HTTP 在把响应体交给客户端之前会把占位符换回真实值
// (`privacy::restore_sync_response_body` /
// `privacy::StreamingResponseRestorer`),这里是 WS 的同一个位置
// ——最后一跳之前,并且在 `capture_client_frame` 之前,所以审计与
// 终态观测继续消费脱敏态的事件。没有命中还原时保持上游原字节。
let restored_client_frame = parsed_upstream_frame
// Standard Responses frames cross the gateway byte-for-byte unless PII
// restoration has something to replace. Codex may wrap public events with
// provider-private side-channel chunks; only that explicit envelope is
// peeled, and each retained event is serialized as a complete opaque Value.
// Observation and capture continue to consume the redacted event, while the
// final client hop receives restored text.
let relay_directive = parsed_upstream_frame
.as_ref()
.map(ParsedResponsesWebSocketFrame::event)
.and_then(|event| {
bound.redaction_restorer.restore_provider_frame_text(event)
.map(|frame| {
bound
.adapter
.relay_directive_for_upstream_event(frame.event())
});
let client_frame = match restored_client_frame {
Some(restored) => AxumWsMessage::Text(restored.into()),
None => upstream_message_to_client(upstream_message.clone()),
};
if let Err(error) = send_client_message(client_socket, client_frame).await {
let mut relay_send_error = None;
let mut relay_serialization_failed = false;
match relay_directive {
Some(ResponsesWebSocketRelayDirective::ForwardOriginal) => {
let restored = parsed_upstream_frame
.as_ref()
.and_then(|frame| {
bound
.redaction_restorer
.restore_provider_frame_text(frame.event())
});
let client_frame = match restored {
Some(text) => AxumWsMessage::Text(text.into()),
None => upstream_message_to_client(upstream_message.clone()),
};
match send_client_message(client_socket, client_frame).await {
Ok(()) => {
if let (Some(turn), Some(frame)) = (
bound.turn_state.attempt_mut(),
parsed_upstream_frame.as_ref(),
) {
turn.capture_client_frame(frame.event());
}
}
Err(error) => relay_send_error = Some(error),
}
}
Some(ResponsesWebSocketRelayDirective::ForwardEvents(events)) => {
for event in events {
let text = match bound
.redaction_restorer
.restore_provider_frame_text(event)
{
Some(restored) => restored,
None => match encode_opaque_websocket_event(event) {
Ok(encoded) => encoded,
Err(_) => {
relay_serialization_failed = true;
break;
}
},
};
match send_client_message(
client_socket,
AxumWsMessage::Text(text.into()),
)
.await
{
Ok(()) => {
if let Some(turn) = bound.turn_state.attempt_mut() {
turn.capture_client_frame(event);
}
}
Err(error) => {
relay_send_error = Some(error);
break;
}
}
}
}
Some(ResponsesWebSocketRelayDirective::SuppressProviderPrivate) => {}
None => {
if let Err(error) = send_client_message(
client_socket,
upstream_message_to_client(upstream_message.clone()),
)
.await
{
relay_send_error = Some(error);
}
}
}
if relay_serialization_failed {
warn!(
event_name = "responses_websocket_provider_event_serialization_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
provider_terminal_reached = terminal_outcome.is_some(),
"gateway could not serialize an opaque provider event"
);
bound
.turn_state
.record_client_delivery_aborted(CLIENT_DELIVERY_FAILED_REASON);
finalize_active_turn(
bound,
state,
settle_signal_for_client_delivery_failure(terminal_outcome),
)
.await;
send_gateway_error_with_status(
client_socket,
502,
"responses_websocket_event_serialization_failed",
"Gateway could not relay the provider event",
)
.await;
close_bound_upstream(bound).await;
close_client_socket(
client_socket,
CLOSE_INTERNAL_ERROR,
"provider_event_serialization_failed",
)
.await;
break;
}
if let Some(error) = relay_send_error {
warn!(
event_name = "responses_websocket_client_send_failed",
log_type = "ops",
@@ -525,11 +564,6 @@ pub(super) async fn relay_bound_connection(
close_bound_upstream(bound).await;
break;
}
if let (Some(turn), Some(frame)) =
(bound.turn_state.attempt_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;
} else if is_close {
@@ -540,6 +574,21 @@ pub(super) async fn relay_bound_connection(
)
.await;
}
if detach_after_forward {
let directive = bound
.pending_adapter_drain
.expect("adapter drain state should be present");
if bound.turn_state.response_in_flight() {
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::provider_quota_exhausted(),
)
.await;
}
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
continue;
}
if drain_for_adapter {
let directive = bound
.pending_adapter_drain
@@ -556,7 +605,9 @@ pub(super) async fn relay_bound_connection(
}
}
async fn wait_for_connection_permit_loss(permit: Option<&aether_runtime::AdmissionPermit>) {
pub(super) async fn wait_for_connection_permit_loss(
permit: Option<&aether_runtime::AdmissionPermit>,
) {
let Some(permit) = permit else {
std::future::pending::<()>().await;
return;
@@ -0,0 +1,173 @@
//! Per-turn control-plane refresh for long-lived Responses WebSockets.
//!
//! An Upgrade authenticates the connection, but it must not freeze API-key,
//! wallet, IP, model, or RPM policy for up to an hour. This module produces one
//! live decision and its exact strong API-key snapshot for every
//! `response.create`; the caller uses that pair consistently for rate limiting,
//! redaction, model authorization, planning, admission, balance, and retries.
use axum::http::StatusCode;
use serde_json::Value;
use crate::ai_serving::GatewayAuthApiKeySnapshot;
use crate::control::{
refresh_execution_runtime_auth_context_with_snapshot, request_model_local_rejection,
GatewayControlDecision, GatewayLocalAuthRejection,
};
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
use crate::handlers::proxy::websocket::session::WEBSOCKET_LOG_TRANSPORT;
use crate::handlers::shared::ip_rules_allow;
use crate::{AppState, GatewayError};
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
macro_rules! warn {
($($arg:tt)*) => {
tracing::warn!(target: LOG_TARGET, $($arg)*)
};
}
#[derive(Debug, Clone)]
pub(super) struct ResponsesWebSocketTurnControl {
pub(super) decision: GatewayControlDecision,
pub(super) auth_snapshot: Option<GatewayAuthApiKeySnapshot>,
pub(super) rpm_bypassed: bool,
}
pub(super) async fn resolve_responses_websocket_turn_control(
state: &AppState,
context: &WebSocketRequestContext,
parts: &http::request::Parts,
client_event: &Value,
) -> Result<ResponsesWebSocketTurnControl, GatewayError> {
if state
.admin_security_ip_blacklisted(context.client_ip)
.await?
{
return Err(GatewayError::Client {
status: StatusCode::FORBIDDEN,
message: "The current IP is blocked".to_string(),
});
}
let mut decision = context.decision.clone();
let auth_snapshot = if let Some(auth_context) = decision.auth_context.take() {
let (refreshed, snapshot) = refresh_execution_runtime_auth_context_with_snapshot(
state,
auth_context,
decision.auth_endpoint_signature.as_deref(),
)
.await?;
decision.local_auth_rejection = refreshed.local_rejection.clone();
decision.auth_context = Some(refreshed);
snapshot
} else {
None
};
// Model-directive configuration is mutable policy too; do not retain the
// Upgrade-time snapshot for the lifetime of the socket.
decision.model_directive_policy =
crate::system_features::ModelDirectivePolicySnapshot::load(state).await;
if let Some(rejection) = decision.local_auth_rejection.clone() {
return Err(websocket_auth_rejection_error(rejection));
}
let Some(auth_context) = decision.auth_context.as_ref() else {
return Err(websocket_auth_rejection_error(
GatewayLocalAuthRejection::InvalidApiKey,
));
};
if !auth_context.access_allowed
|| auth_context.user_id.trim().is_empty()
|| auth_context.api_key_id.trim().is_empty()
{
return Err(websocket_auth_rejection_error(
GatewayLocalAuthRejection::InvalidApiKey,
));
}
if !ip_rules_allow(auth_context.ip_rules.as_deref(), context.client_ip) {
return Err(websocket_auth_rejection_error(
GatewayLocalAuthRejection::IpNotAllowed {
remote_ip: context.client_ip.to_string(),
},
));
}
let body = serde_json::to_vec(client_event)
.map(axum::body::Bytes::from)
.map_err(|error| GatewayError::Internal(error.to_string()))?;
if let Some(rejection) =
request_model_local_rejection(state, Some(&decision), &parts.uri, &parts.headers, &body)
.await?
{
return Err(websocket_auth_rejection_error(rejection));
}
let rpm_bypassed = match state.admin_security_ip_whitelisted(context.client_ip).await {
Ok(value) => value,
Err(error) => {
warn!(
event_name = "responses_websocket_turn_ip_whitelist_check_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
client_ip = %context.client_ip,
error = ?error,
"gateway applied ordinary WebSocket RPM after the live IP whitelist check failed"
);
false
}
};
Ok(ResponsesWebSocketTurnControl {
decision,
auth_snapshot,
rpm_bypassed,
})
}
fn websocket_auth_rejection_error(rejection: GatewayLocalAuthRejection) -> GatewayError {
let (status, message) = match rejection {
GatewayLocalAuthRejection::InvalidApiKey => {
(StatusCode::UNAUTHORIZED, "The API key is invalid")
}
GatewayLocalAuthRejection::LockedApiKey => (
StatusCode::FORBIDDEN,
"The API key is locked and cannot be used",
),
GatewayLocalAuthRejection::WalletUnavailable => {
(StatusCode::FORBIDDEN, "The account wallet is unavailable")
}
GatewayLocalAuthRejection::BalanceDenied { remaining } => {
let message = match remaining {
Some(remaining) => format!("Insufficient balance (remaining: ${remaining:.2})"),
None => "Insufficient balance".to_string(),
};
return GatewayError::Client {
status: StatusCode::TOO_MANY_REQUESTS,
message,
};
}
GatewayLocalAuthRejection::ProviderNotAllowed { .. } => (
StatusCode::FORBIDDEN,
"The provider is not allowed for this API key",
),
GatewayLocalAuthRejection::ApiFormatNotAllowed { .. } => (
StatusCode::FORBIDDEN,
"The API format is not allowed for this API key",
),
GatewayLocalAuthRejection::ModelNotAllowed { .. } => (
StatusCode::FORBIDDEN,
"The requested model is not allowed for this API key",
),
GatewayLocalAuthRejection::IpNotAllowed { .. } => (
StatusCode::UNAUTHORIZED,
"The current IP is not allowed for this API key",
),
};
GatewayError::Client {
status,
message: message.to_string(),
}
}
@@ -119,6 +119,17 @@ impl<'a> ParsedResponsesWebSocketFrame<'a> {
}
}
/// Encodes one event peeled from a provider-private envelope without applying
/// an event-type or field projection.
///
/// Direct provider events should use [`ParsedResponsesWebSocketFrame::raw_text`]
/// so their bytes remain identical. This helper exists only for batch
/// envelopes that cannot be relayed as a whole: serializing the complete
/// [`Value`] preserves every known and future JSON member.
pub(super) fn encode_opaque_websocket_event(event: &Value) -> serde_json::Result<String> {
serde_json::to_string(event)
}
/// 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> {
@@ -148,21 +159,6 @@ fn event_is_started(event: &Value) -> bool {
)
}
/// `response.incomplete` 的合法终态 reason 白名单。
///
/// 这些 reason 表示上游按规则正常结束了本轮响应(写满 `max_output_tokens`、
/// 命中内容过滤、按工具调用截断),标准流解析里它们会变成 `length` /
/// `content_filter` / `tool_calls` 这类正常 finish,和
/// `openai_responses_incomplete_finish_reason` 的既有映射保持一致,因此不能
/// 当成 provider failure 记账。
const LEGITIMATE_RESPONSES_INCOMPLETE_REASONS: [&str; 5] = [
"max_output_tokens",
"max_tokens",
"content_filter",
"tool_calls",
"function_call",
];
/// 读取 `response.incomplete` 携带的 `incomplete_details.reason`。
///
/// 标准位置是 `response.incomplete_details.reason`;批量封装偶尔把
@@ -179,17 +175,33 @@ fn responses_incomplete_reason(event: &Value) -> Option<&str> {
.find(|reason| !reason.is_empty())
}
/// 判断一个 `response.incomplete` 是否是合法终态。
fn responses_incomplete_has_explicit_error(event: &Value) -> bool {
[event.get("error"), event.pointer("/response/error")]
.into_iter()
.flatten()
.any(|error| !error.is_null())
}
/// Derives only the fallback status for `response.incomplete`.
///
/// reason 缺失或不在白名单内(例如 `error`、`server_error`)时继续按
/// provider failure 处理:这类 incomplete 说明上游确实没能正常收尾,仍应扣
/// 供应商健康分。
fn responses_incomplete_is_legitimate_terminal(event: &Value) -> bool {
responses_incomplete_reason(event).is_some_and(|reason| {
LEGITIMATE_RESPONSES_INCOMPLETE_REASONS
.iter()
.any(|candidate| reason.eq_ignore_ascii_case(candidate))
})
/// A non-empty reason is provider-owned protocol data. Treating it as a fixed
/// allowlist would turn every future legitimate reason into a synthetic 502
/// and incorrectly penalize provider health. Missing/malformed reasons and
/// explicit error markers still fail closed; numeric status and recognized
/// error codes continue to override this fallback in
/// [`websocket_event_status_code`].
fn responses_incomplete_default_status(event: &Value) -> u16 {
match responses_incomplete_reason(event) {
None => 502,
Some(reason)
if reason.eq_ignore_ascii_case("error")
|| reason.eq_ignore_ascii_case("server_error") =>
{
502
}
Some(_) if responses_incomplete_has_explicit_error(event) => 502,
Some(_) => 200,
}
}
fn terminal_for_event(event: &Value) -> Option<ResponsesWebSocketFrameTerminal> {
@@ -198,18 +210,13 @@ fn terminal_for_event(event: &Value) -> Option<ResponsesWebSocketFrameTerminal>
status_code: websocket_event_status_code(event, 200),
cancelled: false,
}),
// 合法 incomplete(例如写满 max_output_tokens)是正常终态,默认按 200
// 记账,不再一律当 502 provider failure;reason 缺失或未知时保留原来的
// 502 默认值。显式 `status_code` 和 error code 映射仍然优先于默认值,
// 所以带 `rate_limit_exceeded` 的 incomplete 依旧是 429。
// A non-empty provider reason is a normal terminal by default, including
// future reasons Aether does not yet know. Explicit status/error data
// still wins, so quota and server failures retain their failure status.
"response.incomplete" => Some(ResponsesWebSocketFrameTerminal {
status_code: websocket_event_status_code(
event,
if responses_incomplete_is_legitimate_terminal(event) {
200
} else {
502
},
responses_incomplete_default_status(event),
),
cancelled: false,
}),
@@ -282,7 +289,9 @@ fn safe_websocket_event_label(value: &str) -> String {
#[cfg(test)]
mod tests {
use super::ParsedResponsesWebSocketFrame;
use serde_json::json;
use super::{encode_opaque_websocket_event, ParsedResponsesWebSocketFrame};
#[test]
fn parses_started_frame_once_with_raw_text_and_event_metadata() {
@@ -298,6 +307,37 @@ mod tests {
assert_eq!(frame.event_type_for_log(), "response.in_progress");
}
#[test]
fn future_response_event_keeps_its_exact_original_text_and_unknown_fields() {
let raw = "{ \n \"future_top_level\": {\"nested\": [1, true, null]}, \n \"type\": \"response.future_capability.delta\", \n \"delta\": {\"new_wire_shape\": \"opaque\"}\n}";
let frame = ParsedResponsesWebSocketFrame::parse(raw).expect("valid future event");
assert_eq!(frame.raw_text(), raw);
assert_eq!(
frame.event()["future_top_level"],
json!({"nested": [1, true, null]})
);
assert_eq!(frame.event()["delta"], json!({"new_wire_shape": "opaque"}));
assert!(!frame.is_terminal());
}
#[test]
fn peeled_batch_event_encoding_preserves_the_complete_opaque_value() {
let frame = ParsedResponsesWebSocketFrame::parse(
r#"{"chunks":[{"type":"response.future.done","future_capability":{"mode":"new"},"response":{"id":"resp_future","future_usage":{"novel_tokens":7}}}]}"#,
)
.expect("valid private envelope");
let events = frame.protocol_events();
let event = events.first().expect("one future response event");
let encoded = encode_opaque_websocket_event(event).expect("Value serialization succeeds");
let round_trip: serde_json::Value =
serde_json::from_str(&encoded).expect("encoded event stays valid JSON");
assert_eq!(round_trip, **event);
assert_eq!(round_trip["future_capability"], json!({"mode": "new"}));
assert_eq!(round_trip["response"]["future_usage"]["novel_tokens"], 7);
}
#[test]
fn classifies_terminal_status_and_cancellation() {
let completed = ParsedResponsesWebSocketFrame::parse(
@@ -373,7 +413,7 @@ mod tests {
}
#[test]
fn an_incomplete_without_a_legitimate_reason_stays_a_provider_failure() {
fn an_incomplete_without_a_reason_or_with_a_failure_reason_stays_a_provider_failure() {
for raw in [
r#"{"type":"response.incomplete"}"#,
r#"{"type":"response.incomplete","response":{"incomplete_details":null}}"#,
@@ -386,11 +426,26 @@ mod tests {
assert_eq!(
frame.status(),
Some(502),
"an incomplete without a known-good reason must stay a provider failure: {raw}"
"an incomplete without a usable reason must stay a provider failure: {raw}"
);
}
}
#[test]
fn a_future_incomplete_reason_is_forward_compatible_without_hiding_explicit_errors() {
let future = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"response.incomplete","response":{"incomplete_details":{"reason":"future_context_boundary"}}}"#,
)
.expect("valid frame");
assert_eq!(future.status(), Some(200));
let future_with_error = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"response.incomplete","response":{"error":{"code":"future_provider_error"},"incomplete_details":{"reason":"future_context_boundary"}}}"#,
)
.expect("valid frame");
assert_eq!(future_with_error.status(), Some(502));
}
#[test]
fn a_legitimate_incomplete_still_respects_an_explicit_provider_status() {
let explicit = ParsedResponsesWebSocketFrame::parse(
@@ -6,13 +6,13 @@
use std::time::Duration;
use axum::extract::ws::WebSocket;
use axum::http::StatusCode;
use tokio::task::JoinHandle;
use tokio::time::timeout;
use super::state::BoundResponsesConnection;
use super::turn::{
spawn_responses_websocket_turn_finalization, ResponsesProviderAttempt,
ResponsesWebSocketTurnOutcome,
begin_unowned_responses_websocket_turn, ResponsesProviderAttempt, ResponsesWebSocketTurnOutcome,
};
use crate::handlers::proxy::websocket::session::{
CLOSE_INTERNAL_ERROR, CLOSE_POLICY_VIOLATION, CLOSE_TRY_AGAIN, WEBSOCKET_LOG_TRANSPORT,
@@ -76,11 +76,90 @@ impl std::ops::DerefMut for ActiveProviderAttempt {
}
}
/// Starts a turn and arms its cancellation fallback before control returns to
/// code that can await an upstream bind or socket write.
pub(super) async fn begin_responses_websocket_turn(
state: &AppState,
trace_id: &str,
parts: http::request::Parts,
control_decision: &crate::control::GatewayControlDecision,
decision: crate::ai_serving::AiExecutionDecision,
client_event: &serde_json::Value,
) -> Result<ActiveProviderAttempt, GatewayError> {
let state = state.clone();
let trace_id = trace_id.to_string();
let owner_timeout = state
.frontdoor_runtime_guards
.local_execution_planning_timeout;
let control_decision = control_decision.clone();
let client_event = client_event.clone();
// Beginning an attempt performs several indispensable async writes before
// an `ActiveProviderAttempt` can exist (balance/admission, Pending usage,
// and candidate state). Run that whole transition in an owned task. If the
// relay/session future is cancelled while awaiting it, Tokio detaches this
// task; it still reaches either an explicitly cleaned-up error or an armed
// guard whose dropped output finalizes the attempt.
await_owned_turn_begin(
async move {
let turn = begin_unowned_responses_websocket_turn(
&state,
&parts,
&control_decision,
decision,
&client_event,
)
.await?;
Ok(ActiveProviderAttempt::new(&state, turn))
},
owner_timeout,
trace_id,
)
.await
}
async fn await_owned_turn_begin<T>(
begin: impl std::future::Future<Output = Result<T, GatewayError>> + Send + 'static,
owner_timeout: Duration,
trace_id: String,
) -> Result<T, GatewayError>
where
T: Send + 'static,
{
await_owned_turn_begin_with_timeout(begin, owner_timeout, trace_id).await
}
async fn await_owned_turn_begin_with_timeout<T>(
begin: impl std::future::Future<Output = Result<T, GatewayError>> + Send + 'static,
owner_timeout: Duration,
trace_id: String,
) -> Result<T, GatewayError>
where
T: Send + 'static,
{
tokio::spawn(async move {
tokio::time::timeout(owner_timeout, begin)
.await
.map_err(|_| GatewayError::LocalExecutionPlanningTimeout {
trace_id,
phase: "responses_websocket_turn_begin_owner",
timeout_ms: owner_timeout.as_millis() as u64,
})?
})
.await
.map_err(|error| {
GatewayError::Internal(format!(
"Responses WebSocket turn begin task failed before ownership transfer: {error}"
))
})?
}
impl Drop for ActiveProviderAttempt {
fn drop(&mut self) {
let Some(turn) = self.turn.take() else {
return;
};
let outcome = turn.abandonment_outcome();
let state = self.state.clone();
// No runtime means the process is going down; the spawn could not
// complete anyway.
@@ -93,11 +172,7 @@ impl Drop for ActiveProviderAttempt {
"gateway finalized a Responses WebSocket turn whose relay task went away"
);
handle.spawn(async move {
turn.finalize_detached(
&state,
ResponsesWebSocketTurnOutcome::relay_task_abandoned(),
)
.await;
turn.finalize_detached(&state, outcome).await;
});
}
}
@@ -113,20 +188,23 @@ pub(super) async fn finalize_active_turn(
outcome: ResponsesWebSocketTurnOutcome,
) {
if let Some(turn) = bound.turn_state.end() {
queue_turn_finalization(bound, state, turn.disarm(), outcome).await;
queue_turn_finalization(bound, state, turn, outcome).await;
}
}
pub(super) async fn queue_turn_finalization(
bound: &mut BoundResponsesConnection,
state: &AppState,
turn: ResponsesProviderAttempt,
turn: ActiveProviderAttempt,
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);
bound.pending_turn_finalization = Some(spawn_guarded_turn_finalization(
state.clone(),
turn,
outcome,
));
}
/// 「上一个 attempt 已经结算完毕」的凭证。
@@ -153,7 +231,7 @@ impl PreviousAttemptSettled {
pub(super) async fn settle_turn_finalization(
bound: &mut BoundResponsesConnection,
state: &AppState,
turn: ResponsesProviderAttempt,
turn: ActiveProviderAttempt,
outcome: ResponsesWebSocketTurnOutcome,
) -> PreviousAttemptSettled {
queue_turn_finalization(bound, state, turn, outcome).await;
@@ -161,42 +239,68 @@ pub(super) async fn settle_turn_finalization(
PreviousAttemptSettled(())
}
pub(super) fn spawn_bounded_adapter_observation(
observation: impl std::future::Future<Output = ()> + Send + 'static,
) -> JoinHandle<()> {
spawn_bounded_adapter_observation_with_timeout(
observation,
RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT,
)
}
fn spawn_bounded_adapter_observation_with_timeout(
observation: impl std::future::Future<Output = ()> + Send + 'static,
owner_timeout: Duration,
) -> JoinHandle<()> {
tokio::spawn(async move {
if timeout(owner_timeout, observation).await.is_err() {
warn!(
event_name = "responses_websocket_adapter_observation_timeout",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
timeout_ms = owner_timeout.as_millis() as u64,
"gateway stopped a timed-out Responses WebSocket adapter observation"
);
}
})
}
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"
);
}
if let Some(handle) = bound.pending_adapter_observation.take() {
if let Err(error) = handle.await {
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"
);
}
}
}
pub(super) async fn finalize_unbound_turn(
pub(super) fn finalize_unbound_turn(
state: AppState,
turn: ResponsesProviderAttempt,
turn: ActiveProviderAttempt,
outcome: ResponsesWebSocketTurnOutcome,
) -> JoinHandle<()> {
spawn_responses_websocket_turn_finalization(state, turn, outcome).await
spawn_guarded_turn_finalization(state, turn, outcome)
}
fn spawn_guarded_turn_finalization(
state: AppState,
turn: ActiveProviderAttempt,
outcome: ResponsesWebSocketTurnOutcome,
) -> JoinHandle<()> {
// Spawn synchronously while the armed guard is still owned here. Caller
// cancellation cannot drop an unguarded attempt between cleanup awaits.
tokio::spawn(async move {
let mut turn = turn;
turn.release_admission().await;
turn.disarm().finalize_detached(&state, outcome).await;
})
}
pub(super) async fn await_turn_finalization_handle(handle: JoinHandle<()>) {
@@ -228,6 +332,7 @@ pub(super) async fn send_responses_websocket_turn_start_error(
client_socket: &mut WebSocket,
error: &GatewayError,
) {
let status_code = responses_websocket_turn_start_http_status(error);
match error {
GatewayError::Client { status, message } => {
let (error_type, code) = if status.as_u16() == 429 {
@@ -235,19 +340,13 @@ pub(super) async fn send_responses_websocket_turn_start_error(
} else {
("invalid_request_error", "gateway_request_not_allowed")
};
send_responses_websocket_error(
client_socket,
status.as_u16(),
error_type,
code,
message,
)
.await;
send_responses_websocket_error(client_socket, status_code, error_type, code, message)
.await;
}
GatewayError::AdmissionTimeout { .. } => {
send_responses_websocket_error(
client_socket,
503,
status_code,
"server_error",
"gateway_admission_timeout",
"Gateway capacity is busy; retry this response",
@@ -257,7 +356,7 @@ pub(super) async fn send_responses_websocket_turn_start_error(
GatewayError::LocalExecutionPlanningTimeout { .. } => {
send_responses_websocket_error(
client_socket,
504,
status_code,
"server_error",
"gateway_planning_timeout",
"Gateway planning timed out; retry this response",
@@ -267,7 +366,7 @@ pub(super) async fn send_responses_websocket_turn_start_error(
_ => {
send_responses_websocket_error(
client_socket,
500,
status_code,
"server_error",
"responses_websocket_turn_start_failed",
"Gateway could not start this response",
@@ -277,6 +376,15 @@ pub(super) async fn send_responses_websocket_turn_start_error(
}
}
fn responses_websocket_turn_start_http_status(error: &GatewayError) -> u16 {
match error {
GatewayError::Client { status, .. } => status.as_u16(),
GatewayError::AdmissionTimeout { .. } => StatusCode::TOO_MANY_REQUESTS.as_u16(),
GatewayError::LocalExecutionPlanningTimeout { .. } => StatusCode::GATEWAY_TIMEOUT.as_u16(),
_ => StatusCode::INTERNAL_SERVER_ERROR.as_u16(),
}
}
pub(super) fn responses_websocket_turn_start_close(error: &GatewayError) -> (u16, &'static str) {
match error {
GatewayError::Client { .. } => (CLOSE_POLICY_VIOLATION, "request_not_allowed"),
@@ -292,7 +400,27 @@ mod tests {
use std::sync::Arc;
use std::time::Duration;
use super::await_turn_finalization_handle;
use super::{
await_owned_turn_begin, await_owned_turn_begin_with_timeout,
await_turn_finalization_handle, responses_websocket_turn_start_close,
responses_websocket_turn_start_http_status, spawn_bounded_adapter_observation_with_timeout,
};
use crate::GatewayError;
#[test]
fn admission_timeout_uses_http_429_and_keeps_the_retry_later_close_code() {
let error = GatewayError::AdmissionTimeout {
trace_id: "turn-admission".to_string(),
gate: "gateway_upstream_execution",
queue_budget_ms: 25,
};
assert_eq!(responses_websocket_turn_start_http_status(&error), 429);
assert_eq!(
responses_websocket_turn_start_close(&error),
(1013, "gateway_busy")
);
}
/// C6 依赖的性质:结算是「等到落地」而不是「排进队列」。
///
@@ -348,4 +476,113 @@ mod tests {
let handle = tokio::spawn(async { panic!("settlement task exploded") });
await_turn_finalization_handle(handle).await;
}
#[tokio::test]
async fn cancelling_the_caller_does_not_cancel_turn_begin_or_drop_an_unowned_result() {
struct DropProbe(Arc<AtomicBool>);
impl Drop for DropProbe {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
let begin_finished = Arc::new(AtomicBool::new(false));
let result_dropped = Arc::new(AtomicBool::new(false));
let finished = Arc::clone(&begin_finished);
let dropped = Arc::clone(&result_dropped);
let caller = tokio::spawn(async move {
await_owned_turn_begin(
async move {
tokio::time::sleep(Duration::from_millis(60)).await;
finished.store(true, Ordering::SeqCst);
Ok(DropProbe(dropped))
},
Duration::from_secs(1),
"turn-begin-cancel".to_string(),
)
.await
});
tokio::time::sleep(Duration::from_millis(10)).await;
caller.abort();
let _ = caller.await;
tokio::time::sleep(Duration::from_millis(120)).await;
assert!(
begin_finished.load(Ordering::SeqCst),
"the owned begin task must outlive its cancelled relay caller"
);
assert!(
result_dropped.load(Ordering::SeqCst),
"an undeliverable armed result must be dropped so its cleanup guard runs"
);
}
#[tokio::test]
async fn turn_begin_owner_deadline_drops_stalled_work_and_its_guards() {
struct DropProbe(Arc<AtomicBool>);
impl Drop for DropProbe {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
let dropped = Arc::new(AtomicBool::new(false));
let task_dropped = Arc::clone(&dropped);
let result: Result<(), GatewayError> = await_owned_turn_begin_with_timeout(
async move {
let _probe = DropProbe(task_dropped);
std::future::pending::<()>().await;
Ok(())
},
Duration::from_millis(20),
"turn-begin-deadline".to_string(),
)
.await;
assert!(matches!(
result,
Err(GatewayError::LocalExecutionPlanningTimeout {
trace_id,
phase: "responses_websocket_turn_begin_owner",
timeout_ms: 20,
}) if trace_id == "turn-begin-deadline"
));
assert!(
dropped.load(Ordering::SeqCst),
"owner timeout must drop the stalled begin future so RAII cleanup runs"
);
}
#[tokio::test]
async fn cancelling_observation_waiter_cannot_bypass_the_owner_timeout() {
struct DropProbe(Arc<AtomicBool>);
impl Drop for DropProbe {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
let dropped = Arc::new(AtomicBool::new(false));
let task_dropped = Arc::clone(&dropped);
let observation = async move {
let _probe = DropProbe(task_dropped);
std::future::pending::<()>().await;
};
let owner =
spawn_bounded_adapter_observation_with_timeout(observation, Duration::from_millis(20));
let waiter = tokio::spawn(async move {
let _ = owner.await;
});
waiter.abort();
let _ = waiter.await;
tokio::time::timeout(Duration::from_secs(1), async {
while !dropped.load(Ordering::SeqCst) {
tokio::task::yield_now().await;
}
})
.await
.expect("detached observation owner must enforce its own timeout");
}
}
@@ -12,9 +12,11 @@ mod admission;
mod binding;
mod client;
mod connection;
mod control;
mod frame;
mod lifecycle;
mod observation;
mod ownership;
mod quota;
mod redaction;
mod relay_policy;
@@ -61,5 +63,4 @@ pub(crate) async fn responses_websocket(
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",
};
@@ -37,7 +37,11 @@ impl ResponsesStructuredTerminalObserver {
/// 第一个被拒绝的事件就停止推进并把摘要标成 parser_error:解析器的状态机是
/// 有顺序的,跳过一个事件继续喂后面的只会得到更没意义的摘要。
pub(super) fn observe_events(&mut self, report_context: &Value, events: &[&Value]) {
for event in events {
for event in events
.iter()
.copied()
.filter(|event| event_is_relevant_to_terminal_observation(event))
{
if let Err(error) = self.inner.push_event(report_context, event) {
self.inner.disable_with_error(error.to_string());
break;
@@ -61,6 +65,28 @@ impl ResponsesStructuredTerminalObserver {
}
}
/// The WebSocket relay is not a Responses schema gateway. It forwards all
/// events opaquely, while this observer consumes only identity/terminal
/// snapshots needed for usage and settlement. In particular, a future
/// `response.*` delta must not become an observation failure merely because
/// Aether's canonical streaming parser does not know it yet.
fn event_is_relevant_to_terminal_observation(event: &Value) -> bool {
matches!(
event.get("type").and_then(Value::as_str),
Some(
"response.created"
| "response.in_progress"
| "response.queued"
| "response.completed"
| "response.done"
| "response.failed"
| "response.incomplete"
| "response.cancelled"
| "error"
)
)
}
#[cfg(test)]
mod tests {
use serde_json::json;
@@ -141,6 +167,45 @@ mod tests {
assert_eq!(usage.dimensions.get("total_tokens"), Some(&json!(4)));
}
#[test]
fn future_and_provider_private_events_are_ignored_only_by_the_side_observer() {
let context = report_context();
let private = json!({
"type": "codex.response.metadata",
"private_future_field": {"shape": "unknown"},
});
let future = json!({
"type": "response.future_capability.delta",
"future_capability": {"nested": [1, 2, 3]},
});
let completed = json!({
"type": "response.completed",
"response": {
"id": "resp_future",
"model": "future-model",
"status": "completed",
"future_response_field": {"also": "unknown"},
"usage": {"input_tokens": 5, "output_tokens": 2, "total_tokens": 7},
},
});
let mut observer = ResponsesStructuredTerminalObserver::default();
observer.observe_events(&context, &[&private, &future, &completed]);
let summary = observer.finish(&context);
assert!(summary.observed_finish);
assert_eq!(summary.response_id.as_deref(), Some("resp_future"));
assert_eq!(summary.unknown_event_count, 0);
assert_eq!(
summary
.standardized_usage
.as_ref()
.map(|usage| (usage.input_tokens, usage.output_tokens)),
Some((5, 2))
);
assert!(summary.parser_error.is_none());
}
#[test]
fn a_disabled_observer_reports_the_parser_error() {
let context = report_context();
@@ -0,0 +1,245 @@
//! Cancellation-safe ownership handoff for WebSocket planning leases.
//!
//! The relay races every turn against connection and response deadlines. A
//! planner future therefore cannot directly own a distributed pool-key lease:
//! losing the race would drop the future between scheduler selection and turn
//! startup, leaving that key unavailable until the lease TTL elapsed.
use std::collections::BTreeSet;
use std::time::Duration;
use serde_json::Value;
use tokio::task::JoinHandle;
use super::lifecycle::{begin_responses_websocket_turn, ActiveProviderAttempt};
use crate::ai_serving::{
maybe_build_responses_websocket_decision, AiExecutionDecision, GatewayAuthApiKeySnapshot,
ResponsesWebSocketDecision, ResponsesWebSocketPinnedCandidate,
};
use crate::control::GatewayControlDecision;
use crate::orchestration::release_pool_key_lease_from_report_context;
use crate::{AppState, GatewayError};
/// Owns a selected pool-key lease until the attempt lifecycle has taken over
/// the decision report context.
pub(super) struct PlannedPoolKeyLeaseGuard {
state: AppState,
report_context: Option<Value>,
}
/// Planner output coupled to both its request parts and lease guard.
pub(super) struct OwnedResponsesWebSocketDecision {
pub(super) planned: ResponsesWebSocketDecision,
pub(super) planning_parts: http::request::Parts,
pub(super) planned_lease: PlannedPoolKeyLeaseGuard,
}
/// Runs planning in an owner task. Dropping the caller's waiter detaches this
/// task; an unobserved successful output drops its guard and releases the
/// selected pool-key lease.
#[allow(clippy::too_many_arguments)]
pub(super) fn spawn_owned_responses_websocket_plan(
state: AppState,
parts: http::request::Parts,
trace_id: String,
control_decision: GatewayControlDecision,
auth_snapshot: Option<GatewayAuthApiKeySnapshot>,
client_event: Value,
excluded_key_ids: Option<BTreeSet<String>>,
excluded_codex_account_ids: Option<BTreeSet<String>>,
pinned_candidate: Option<ResponsesWebSocketPinnedCandidate>,
) -> JoinHandle<Result<Option<OwnedResponsesWebSocketDecision>, GatewayError>> {
let owner_timeout = state
.frontdoor_runtime_guards
.local_execution_planning_timeout;
tokio::spawn(async move {
let planned = await_owned_planning_deadline(
maybe_build_responses_websocket_decision(
&state,
&parts,
&trace_id,
&control_decision,
auth_snapshot.as_ref(),
&client_event,
excluded_key_ids.as_ref(),
excluded_codex_account_ids.as_ref(),
pinned_candidate.as_ref(),
),
owner_timeout,
)
.await
.map_err(|_| GatewayError::LocalExecutionPlanningTimeout {
trace_id: trace_id.clone(),
phase: "responses_websocket_plan_owner",
timeout_ms: owner_timeout.as_millis() as u64,
})??;
Ok(planned.map(|planned| {
let planned_lease =
PlannedPoolKeyLeaseGuard::new(&state, planned.execution.report_context.as_ref());
OwnedResponsesWebSocketDecision {
planned,
planning_parts: parts,
planned_lease,
}
}))
})
}
async fn await_owned_planning_deadline<F, T>(
planning: F,
deadline: Duration,
) -> Result<T, tokio::time::error::Elapsed>
where
F: std::future::Future<Output = T>,
{
tokio::time::timeout(deadline, planning).await
}
pub(super) async fn await_owned_responses_websocket_plan(
handle: JoinHandle<Result<Option<OwnedResponsesWebSocketDecision>, GatewayError>>,
) -> Result<Option<OwnedResponsesWebSocketDecision>, GatewayError> {
handle.await.map_err(|error| {
GatewayError::Internal(format!(
"Responses WebSocket planning task failed before ownership transfer: {error}"
))
})?
}
impl PlannedPoolKeyLeaseGuard {
fn new(state: &AppState, report_context: Option<&Value>) -> Self {
Self {
state: state.clone(),
report_context: report_context.cloned(),
}
}
pub(super) async fn release(mut self) {
release_pool_key_lease_from_report_context(&self.state, self.report_context.as_ref()).await;
self.report_context = None;
}
fn disarm(&mut self) {
self.report_context = None;
}
}
impl Drop for PlannedPoolKeyLeaseGuard {
fn drop(&mut self) {
let Some(report_context) = self.report_context.take() else {
return;
};
let state = self.state.clone();
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn(async move {
release_pool_key_lease_from_report_context(&state, Some(&report_context)).await;
});
}
}
}
/// Keeps the planning guard in the same detached owner task as lifecycle
/// startup. If the relay loses a deadline race while awaiting startup, the
/// task completes the handoff (or releases the lease on failure) without a
/// cancellation gap.
pub(super) async fn begin_responses_websocket_turn_with_planned_lease(
state: &AppState,
trace_id: &str,
parts: http::request::Parts,
control_decision: &GatewayControlDecision,
decision: AiExecutionDecision,
client_event: &Value,
mut planned_lease: PlannedPoolKeyLeaseGuard,
) -> Result<ActiveProviderAttempt, GatewayError> {
let state = state.clone();
let trace_id = trace_id.to_string();
let control_decision = control_decision.clone();
let client_event = client_event.clone();
tokio::spawn(async move {
let turn = begin_responses_websocket_turn(
&state,
&trace_id,
parts,
&control_decision,
decision,
&client_event,
)
.await?;
// ActiveProviderAttempt now owns the report context containing the
// lease. No await occurs between that handoff and disarming the guard.
planned_lease.disarm();
Ok(turn)
})
.await
.map_err(|error| {
GatewayError::Internal(format!(
"Responses WebSocket guarded turn startup task failed: {error}"
))
})?
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
struct DropProbe(Arc<AtomicUsize>);
impl Drop for DropProbe {
fn drop(&mut self) {
self.0.fetch_add(1, Ordering::SeqCst);
}
}
#[tokio::test]
async fn dropping_a_planning_waiter_detaches_the_owner_and_drops_its_output() {
let started = Arc::new(tokio::sync::Notify::new());
let release = Arc::new(tokio::sync::Notify::new());
let dropped = Arc::new(AtomicUsize::new(0));
let task_started = Arc::clone(&started);
let task_release = Arc::clone(&release);
let task_dropped = Arc::clone(&dropped);
let owner = tokio::spawn(async move {
task_started.notify_one();
task_release.notified().await;
DropProbe(task_dropped)
});
started.notified().await;
let waiter = tokio::spawn(async move {
let _ = owner.await;
});
waiter.abort();
let _ = waiter.await;
release.notify_one();
tokio::time::timeout(Duration::from_secs(1), async {
while dropped.load(Ordering::SeqCst) == 0 {
tokio::task::yield_now().await;
}
})
.await
.expect("detached owner output should be dropped after it finishes");
assert_eq!(dropped.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn planning_owner_deadline_drops_stalled_work_and_its_guards() {
let dropped = Arc::new(AtomicUsize::new(0));
let task_dropped = Arc::clone(&dropped);
let planning = async move {
let _probe = DropProbe(task_dropped);
std::future::pending::<()>().await;
};
let result =
super::await_owned_planning_deadline(planning, Duration::from_millis(20)).await;
assert!(
result.is_err(),
"stalled planning must hit its owner deadline"
);
assert_eq!(dropped.load(Ordering::SeqCst), 1);
}
}
@@ -1,7 +1,5 @@
//! 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;
@@ -10,28 +8,21 @@ use super::adapter::{
resolve_responses_websocket_adapter, ResponsesWebSocketDrainDirective,
ResponsesWebSocketRebindSafety,
};
use super::lifecycle::{queue_turn_finalization, ActiveProviderAttempt, PreviousAttemptSettled};
use super::request::{
build_planning_parts, planned_response_create_event, response_create_has_previous_response_id,
use super::lifecycle::{queue_turn_finalization, PreviousAttemptSettled};
use super::ownership::{
await_owned_responses_websocket_plan, begin_responses_websocket_turn_with_planned_lease,
spawn_owned_responses_websocket_plan, OwnedResponsesWebSocketDecision,
};
use super::request::{build_planning_parts, planned_response_create_event};
use super::state::BoundResponsesConnection;
use super::turn::{
begin_responses_websocket_turn, prepare_responses_websocket_turn_decision,
ResponsesWebSocketTurnOutcome,
};
use super::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::handlers::proxy::websocket::transport::close_upstream_socket;
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 {
@@ -132,6 +123,17 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
active.retry_attempted = true;
active.turn_attempt = active.turn_attempt.saturating_add(1);
let client_event = active.client_event.clone();
let Some(turn_control) = active.turn_control.clone() else {
warn!(
event_name = "responses_websocket_quota_retry_control_missing",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
"gateway refused to retry a WebSocket turn without its live authorization snapshot"
);
return false;
};
let turn_index = active.turn_index;
let logical_turn_id = active.logical_turn_id.clone();
let turn_attempt = active.turn_attempt;
@@ -147,18 +149,20 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
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_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_codex_account_ids.is_empty()).then_some(excluded_codex_account_ids);
let planned = match await_owned_responses_websocket_plan(spawn_owned_responses_websocket_plan(
state.clone(),
planning_parts,
turn_request_id.clone(),
turn_control.decision.clone(),
turn_control.auth_snapshot.clone(),
client_event.clone(),
excluded_key_ids,
excluded_codex_account_ids,
)
None,
))
.await
{
Ok(Some(decision)) => decision,
@@ -188,11 +192,16 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
return false;
}
};
let OwnedResponsesWebSocketDecision {
planned,
planning_parts,
planned_lease,
} = planned;
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;
planned_lease.release().await;
warn!(
event_name = "responses_websocket_quota_retry_selected_exhausted_key",
log_type = "ops",
@@ -212,8 +221,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
) {
Ok(event) => event,
Err(code) => {
release_pool_key_lease_from_report_context(state, decision.report_context.as_ref())
.await;
planned_lease.release().await;
warn!(
event_name = "responses_websocket_quota_retry_normalization_failed",
log_type = "ops",
@@ -237,12 +245,14 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
&logical_turn_id,
turn_attempt,
);
let mut turn = match begin_responses_websocket_turn(
let mut turn = match begin_responses_websocket_turn_with_planned_lease(
state,
&planning_parts,
&context.decision,
&context.trace_id,
planning_parts,
&turn_control.decision,
turn_decision,
&client_event,
planned_lease,
)
.await
{
@@ -309,10 +319,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
// 同一个 logical turn 的下一个 attempt 就位。状态不符时把 attempt 交回
// drop guard 结算并让调用方走「透明重试失败」分支,不静默丢弃一条已经写了
// pending usage 行、占着 candidate 和 pool key lease 的 attempt。
if let Err(orphan) = bound
.turn_state
.resume(ActiveProviderAttempt::new(state, turn))
{
if let Err(orphan) = bound.turn_state.resume(turn) {
drop(orphan);
return false;
}
@@ -334,15 +341,6 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
true
}
pub(super) fn active_continuation_can_retry_from_full_input(
bound: &BoundResponsesConnection,
) -> bool {
bound.turn_state.logical().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")
@@ -355,27 +353,6 @@ pub(super) fn is_usage_limit_error_event(event: &Value) -> bool {
.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,
@@ -305,8 +305,8 @@ mod tests {
remote_addr: "127.0.0.1:65000"
.parse::<SocketAddr>()
.expect("remote address should parse"),
client_ip: "127.0.0.1".parse().expect("client IP should parse"),
decision,
rpm_bypassed: false,
websocket_connection_permit: None,
}
}
@@ -81,15 +81,12 @@ pub struct QuotaRelayFacts {
/// 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.
/// and could not bind an alternate upstream. This prevents retry loops and
/// causes the original upstream quota event to be relayed to the client.
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,
}
@@ -97,7 +94,6 @@ pub struct QuotaRelayFacts {
pub enum QuotaRelayAction {
None,
AttemptTransparentRetry,
RequestFullContinuationRetry,
ForwardQuotaAndDetach,
}
@@ -113,13 +109,7 @@ pub const fn classify_quota_relay(facts: QuotaRelayFacts) -> QuotaRelayAction {
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 {
if facts.usage_limit_error || facts.upstream_closed {
return QuotaRelayAction::ForwardQuotaAndDetach;
}
QuotaRelayAction::None
@@ -177,7 +167,6 @@ mod tests {
retry_current_turn: true,
transparent_retry_failed: false,
usage_limit_error: true,
continuation_retry_eligible: false,
upstream_closed: false,
});
assert_eq!(first, QuotaRelayAction::AttemptTransparentRetry);
@@ -190,50 +179,22 @@ mod tests {
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() {
fn quota_error_is_forwarded_after_transparent_retry_is_unavailable() {
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
QuotaRelayAction::ForwardQuotaAndDetach
);
}
@@ -277,7 +238,6 @@ mod tests {
retry_current_turn: true,
transparent_retry_failed: false,
usage_limit_error: false,
continuation_retry_eligible: false,
upstream_closed: false,
}),
QuotaRelayAction::None
@@ -289,7 +249,6 @@ mod tests {
retry_current_turn: false,
transparent_retry_failed: true,
usage_limit_error: false,
continuation_retry_eligible: false,
upstream_closed: true,
}),
QuotaRelayAction::ForwardQuotaAndDetach
@@ -13,6 +13,25 @@ use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
use crate::headers::request_origin_from_headers_and_remote_addr;
use crate::privacy::RedactionSessionSlot;
/// Model identifiers are copied into planner diagnostics. Bound them before
/// any planning/logging so a single 16 MiB WebSocket frame cannot amplify into
/// repeated multi-megabyte log records.
pub(super) const MAX_RESPONSES_WEBSOCKET_MODEL_BYTES: usize = 256;
pub(super) fn validated_response_create_model(value: &Value) -> Result<&str, &'static str> {
let Some(model) = value
.as_str()
.map(str::trim)
.filter(|model| !model.is_empty())
else {
return Err("invalid_response_create_model");
};
if model.len() > MAX_RESPONSES_WEBSOCKET_MODEL_BYTES {
return Err("invalid_response_create_model");
}
Ok(model)
}
/// 把一条 WebSocket turn 还原成 planner 需要的 HTTP 形状请求头部。
///
/// 这里必须和 HTTP 前门(`handlers/proxy/mod.rs`)保持同一份 extension 契约:
@@ -73,11 +92,12 @@ pub(super) fn planned_response_create_event(
/// 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.
/// `previous_response_id` is on the Codex unsupported-field list, Codex HTTP
/// normalization may force `store`, and `generate` is not an HTTP body option
/// at all. Those fields are WebSocket protocol state, so an explicitly supplied
/// value (including `null`) must be re-grafted verbatim from the client event.
/// `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,
@@ -89,13 +109,9 @@ fn finish_response_create_event(
"type".to_string(),
Value::String("response.create".to_string()),
);
for field in ["previous_response_id", "generate"] {
for field in ["store", "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.insert(field.to_string(), value.clone());
}
}
object.remove("stream");
@@ -109,13 +125,6 @@ pub(super) fn response_create_has_previous_response_id(event: &Value) -> bool {
.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,
@@ -126,13 +135,7 @@ pub(super) fn changed_followup_response_create_model(
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");
};
let model = validated_response_create_model(model)?;
if model.eq_ignore_ascii_case(current_client_model) {
Ok(None)
} else {
@@ -154,14 +157,10 @@ pub(super) fn response_create_model_or_current(
);
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())
let model = validated_response_create_model(model)?;
let model = model.to_string();
object.insert("model".to_string(), Value::String(model.clone()));
Ok(model)
}
pub(super) fn provider_model_from_decision(decision: &AiExecutionDecision) -> Option<String> {
@@ -217,7 +216,7 @@ mod tests {
use std::net::SocketAddr;
use axum::http::{HeaderMap, Uri};
use serde_json::json;
use serde_json::{json, Value};
use super::{
build_planning_parts, normalize_followup_response_create,
@@ -236,6 +235,7 @@ mod tests {
remote_addr: "127.0.0.1:65001"
.parse::<SocketAddr>()
.expect("remote address should parse"),
client_ip: "127.0.0.1".parse().expect("client IP should parse"),
decision: GatewayControlDecision::synthetic(
"/v1/responses".to_string(),
Some("ai_public".to_string()),
@@ -243,7 +243,6 @@ mod tests {
Some("responses_websocket".to_string()),
Some("openai:responses".to_string()),
),
rpm_bypassed: false,
websocket_connection_permit: None,
}
}
@@ -311,6 +310,56 @@ mod tests {
assert!(normalized.get("background").is_none());
}
#[test]
fn explicit_store_and_previous_response_id_are_forwarded_opaquely() {
let event = json!({
"type": "response.create",
"model": "public-model",
"store": true,
"previous_response_id": {"future": "opaque"},
"input": [],
});
let normalized = normalized_continuation(
&event,
&ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_provider_type_for_tests("codex"),
);
// Codex HTTP normalization normally forces `store: false` and removes
// `previous_response_id`. WebSocket framing restores exactly what the
// client sent so the upstream owns validation and continuation lookup.
assert_eq!(normalized["store"], true);
assert_eq!(
normalized["previous_response_id"],
json!({"future": "opaque"})
);
}
#[test]
fn explicit_null_websocket_protocol_state_is_not_rewritten() {
let event = json!({
"type": "response.create",
"model": "public-model",
"store": null,
"previous_response_id": null,
"generate": null,
"input": [],
});
let normalized = normalized_continuation(
&event,
&ResponsesWebSocketBodyNormalization::for_tests("provider-model")
.with_provider_type_for_tests("codex"),
);
assert!(normalized.get("store").is_some_and(Value::is_null));
assert!(normalized
.get("previous_response_id")
.is_some_and(Value::is_null));
assert!(normalized.get("generate").is_some_and(Value::is_null));
}
#[test]
fn continuation_strips_fields_the_codex_backend_rejects() {
// The point of the fix: before it, turns 2..N reached Codex with the
@@ -8,7 +8,6 @@
//! request may replace it, but a continuation must stay on the original
//! connection and account.
use axum::body::Bytes;
use axum::extract::ws::{Message as AxumWsMessage, WebSocket};
use futures_util::{SinkExt, StreamExt};
use serde_json::Value;
@@ -16,23 +15,27 @@ use uuid::Uuid;
use super::adapter::resolve_responses_websocket_adapter;
use super::client::consume_response_create_rate_limit;
use super::connection::relay_bound_connection;
use super::connection::{relay_bound_connection, wait_for_connection_permit_loss};
use super::control::resolve_responses_websocket_turn_control;
use super::lifecycle::{
await_pending_adapter_observation, await_pending_turn_finalization,
await_turn_finalization_handle, finalize_unbound_turn, responses_websocket_turn_start_close,
send_responses_websocket_turn_start_error, ActiveProviderAttempt,
await_turn_finalization_handle, finalize_active_turn, finalize_unbound_turn,
responses_websocket_turn_start_close, send_responses_websocket_turn_start_error,
};
use super::ownership::{
await_owned_responses_websocket_plan, begin_responses_websocket_turn_with_planned_lease,
spawn_owned_responses_websocket_plan, OwnedResponsesWebSocketDecision,
};
use super::redaction::redact_responses_websocket_client_event;
use super::request::{build_planning_parts, planned_response_create_event};
use super::turn::{
begin_responses_websocket_turn, prepare_responses_websocket_turn_decision,
ResponsesWebSocketTurnOutcome,
use super::relay_policy::{fatal_relay_policy, FatalRelaySignal};
use super::request::{
build_planning_parts, planned_response_create_event, validated_response_create_model,
};
use super::state::BoundResponsesConnection;
use super::turn::{prepare_responses_websocket_turn_decision, ResponsesWebSocketTurnOutcome};
use super::turn_state::LogicalTurn;
use super::upstream::bind_responses_upstream;
use super::upstream::{bind_responses_upstream, close_bound_upstream};
use crate::ai_serving::maybe_build_responses_websocket_decision;
use crate::control::request_model_local_rejection;
use crate::handlers::proxy::websocket::ingress::{
WebSocketConnectionLog, WebSocketConnectionLogSpec, WebSocketRequestContext,
};
@@ -43,7 +46,6 @@ use crate::handlers::proxy::websocket::session::{
use crate::handlers::proxy::websocket::transport::{
close_client_socket, send_gateway_error, send_gateway_error_with_status,
};
use crate::orchestration::release_pool_key_lease_from_report_context;
use crate::AppState;
const RESPONSES_WEBSOCKET_LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
@@ -71,6 +73,13 @@ enum InitialMessageError {
InvalidJson,
MissingResponseCreate,
MissingModel,
InvalidModel,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ConnectionTermination {
ConnectionLimitReached,
ConnectionAdmissionLost,
}
impl InitialMessageError {
@@ -83,6 +92,7 @@ impl InitialMessageError {
Self::InvalidJson => "invalid_response_create",
Self::MissingResponseCreate => "expected_response_create",
Self::MissingModel => "response_create_model_required",
Self::InvalidModel => "invalid_response_create_model",
}
}
@@ -91,7 +101,9 @@ impl InitialMessageError {
Self::TimedOut => CLOSE_TRY_AGAIN,
Self::ClientClosed => 1000,
Self::ClientRead | Self::UnsupportedFrame | Self::InvalidJson => CLOSE_POLICY_VIOLATION,
Self::MissingResponseCreate | Self::MissingModel => CLOSE_POLICY_VIOLATION,
Self::MissingResponseCreate | Self::MissingModel | Self::InvalidModel => {
CLOSE_POLICY_VIOLATION
}
}
}
}
@@ -101,46 +113,151 @@ pub(super) async fn run_responses_websocket(
state: AppState,
mut context: WebSocketRequestContext,
) {
// The public connection limit starts when the HTTP Upgrade hands us the
// socket, not after provider planning and its upstream handshake finish.
let connection_deadline =
tokio::time::Instant::now() + RESPONSES_WEBSOCKET_SESSION_LIMITS.max_connection_duration;
let connection_permit = context.websocket_connection_permit.take();
let connection_log = WebSocketConnectionLog::new(&context, RESPONSES_CONNECTION_LOG_SPEC);
connection_log.log_opened();
let (first_text, first_event) = match receive_initial_response_create(&mut client_socket).await
{
Ok(value) => value,
Err(error) => {
if !matches!(error, InitialMessageError::ClientClosed) {
send_gateway_error(
&mut client_socket,
error.code(),
"WebSocket must start with a valid response.create event",
)
.await;
close_client_socket(
&mut client_socket,
error.close_code(),
"invalid_initial_event",
)
.await;
}
let bootstrap_result = supervise_responses_websocket_phase(
bootstrap_responses_websocket(&mut client_socket, state.clone(), &context),
connection_deadline,
connection_permit.as_ref(),
)
.await;
let mut bound = match bootstrap_result {
Ok(Some(bound)) => bound,
Ok(None) => return,
Err(termination) => {
// The bootstrap future (and therefore every lease/attempt guard it
// owned) has been dropped before the permit or client socket is
// touched here.
drop(connection_permit);
close_terminated_bootstrap(&mut client_socket, &context, termination).await;
return;
}
};
let planning_parts = build_planning_parts(&context);
match consume_response_create_rate_limit(&state, &context.decision, context.rpm_bypassed).await
let relay_result = supervise_responses_websocket_phase(
relay_bound_connection(&mut client_socket, &mut bound, &state, &context),
connection_deadline,
connection_permit.as_ref(),
)
.await;
if let Err(termination) = relay_result {
// The relay future has been dropped before cleanup starts. This makes
// cancellation safe even when a client-frame branch was awaiting a
// live auth refresh, redaction, planning, admission, or rebind.
let outcome = match termination {
ConnectionTermination::ConnectionLimitReached => {
ResponsesWebSocketTurnOutcome::connection_limit_reached()
}
ConnectionTermination::ConnectionAdmissionLost => {
ResponsesWebSocketTurnOutcome::connection_admission_lost()
}
};
finalize_active_turn(&mut bound, &state, outcome).await;
close_bound_upstream(&mut bound).await;
drop(connection_permit);
close_terminated_relay(&mut client_socket, &context, termination).await;
} else {
drop(connection_permit);
}
await_pending_turn_finalization(&mut bound).await;
await_pending_adapter_observation(&mut bound).await;
}
async fn supervise_responses_websocket_phase<F, T>(
phase: F,
connection_deadline: tokio::time::Instant,
connection_permit: Option<&aether_runtime::AdmissionPermit>,
) -> Result<T, ConnectionTermination>
where
F: std::future::Future<Output = T>,
{
tokio::pin!(phase);
tokio::select! {
biased;
_ = tokio::time::sleep_until(connection_deadline) => {
Err(ConnectionTermination::ConnectionLimitReached)
}
_ = wait_for_connection_permit_loss(connection_permit) => {
Err(ConnectionTermination::ConnectionAdmissionLost)
}
output = &mut phase => Ok(output),
}
}
async fn bootstrap_responses_websocket(
client_socket: &mut WebSocket,
state: AppState,
context: &WebSocketRequestContext,
) -> Option<BoundResponsesConnection> {
let (_first_text, first_event) = match receive_initial_response_create(client_socket).await {
Ok(value) => value,
Err(error) => {
if !matches!(error, InitialMessageError::ClientClosed) {
send_gateway_error(
client_socket,
error.code(),
"WebSocket must start with a valid response.create event",
)
.await;
close_client_socket(client_socket, error.close_code(), "invalid_initial_event")
.await;
}
return None;
}
};
let planning_parts = build_planning_parts(context);
let turn_control = match resolve_responses_websocket_turn_control(
&state,
context,
&planning_parts,
&first_event,
)
.await
{
Ok(control) => control,
Err(error) => {
warn!(
event_name = "responses_websocket_initial_turn_control_rejected",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error = ?error,
"gateway rejected a Responses WebSocket initial turn after live policy refresh"
);
send_responses_websocket_turn_start_error(client_socket, &error).await;
let (close_code, close_reason) = responses_websocket_turn_start_close(&error);
close_client_socket(client_socket, close_code, close_reason).await;
return None;
}
};
match consume_response_create_rate_limit(
&state,
&turn_control.decision,
turn_control.rpm_bypassed,
)
.await
{
Ok(true) => {}
Ok(false) => {
send_gateway_error_with_status(
&mut client_socket,
client_socket,
429,
"rate_limit_exceeded",
"Too many response.create events; retry later",
)
.await;
close_client_socket(&mut client_socket, CLOSE_TRY_AGAIN, "rate_limit_exceeded").await;
return;
close_client_socket(client_socket, CLOSE_TRY_AGAIN, "rate_limit_exceeded").await;
return None;
}
Err(()) => {
warn!(
@@ -152,68 +269,19 @@ pub(super) async fn run_responses_websocket(
"gateway failed to consume WebSocket response rate limit"
);
send_gateway_error_with_status(
&mut client_socket,
client_socket,
503,
"gateway_rate_limit_unavailable",
"Gateway could not evaluate the response rate limit",
)
.await;
close_client_socket(
&mut client_socket,
client_socket,
CLOSE_INTERNAL_ERROR,
"rate_limit_unavailable",
)
.await;
return;
}
}
match request_model_local_rejection(
&state,
Some(&context.decision),
&planning_parts.uri,
&planning_parts.headers,
&Bytes::from(first_text.into_bytes()),
)
.await
{
Ok(Some(_)) => {
send_gateway_error(
&mut client_socket,
"model_not_allowed",
"The requested model is not available to this API key",
)
.await;
close_client_socket(
&mut client_socket,
CLOSE_POLICY_VIOLATION,
"model_not_allowed",
)
.await;
return;
}
Ok(None) => {}
Err(_) => {
warn!(
event_name = "responses_websocket_model_access_check_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
"gateway failed to evaluate WebSocket model access policy"
);
send_gateway_error(
&mut client_socket,
"gateway_auth_unavailable",
"Gateway could not evaluate request access",
)
.await;
close_client_socket(
&mut client_socket,
CLOSE_INTERNAL_ERROR,
"gateway_auth_unavailable",
)
.await;
return;
return None;
}
}
@@ -223,7 +291,7 @@ pub(super) async fn run_responses_websocket(
let redacted_first_event = redact_responses_websocket_client_event(
&state,
&planning_parts,
&context.decision,
&turn_control.decision,
&first_event,
)
.await;
@@ -243,49 +311,51 @@ pub(super) async fn run_responses_websocket(
"gateway could not apply chat PII redaction to the initial Responses WebSocket event"
);
send_gateway_error_with_status(
&mut client_socket,
client_socket,
500,
"responses_websocket_redaction_unavailable",
"Gateway could not apply the configured PII redaction",
)
.await;
close_client_socket(
&mut client_socket,
client_socket,
CLOSE_INTERNAL_ERROR,
"responses_websocket_redaction_unavailable",
)
.await;
return;
return None;
}
};
let planned = match maybe_build_responses_websocket_decision(
&state,
&planning_parts,
&context.trace_id,
&context.decision,
&first_event,
let planned = match await_owned_responses_websocket_plan(spawn_owned_responses_websocket_plan(
state.clone(),
planning_parts,
context.trace_id.clone(),
turn_control.decision.clone(),
turn_control.auth_snapshot.clone(),
first_event.clone(),
None,
None,
)
None,
))
.await
{
Ok(Some(decision)) => decision,
Ok(None) => {
send_gateway_error_with_status(
&mut client_socket,
client_socket,
503,
"responses_provider_unavailable",
"No eligible WebSocket-enabled Responses provider is available",
)
.await;
close_client_socket(
&mut client_socket,
client_socket,
CLOSE_TRY_AGAIN,
"responses_provider_unavailable",
)
.await;
return;
return None;
}
Err(_) => {
warn!(
@@ -297,22 +367,27 @@ pub(super) async fn run_responses_websocket(
"gateway failed to plan Responses WebSocket provider request"
);
send_gateway_error_with_status(
&mut client_socket,
client_socket,
503,
"responses_provider_unavailable",
"Gateway could not prepare a Provider connection",
)
.await;
close_client_socket(
&mut client_socket,
client_socket,
CLOSE_INTERNAL_ERROR,
"responses_planning_failed",
)
.await;
return;
return None;
}
};
let OwnedResponsesWebSocketDecision {
planned,
planning_parts,
planned_lease,
} = planned;
let adapter = resolve_responses_websocket_adapter(planned.adapter);
let normalization = planned.normalization;
let decision = planned.execution;
@@ -322,8 +397,7 @@ pub(super) async fn run_responses_websocket(
}) {
Ok(event) => event,
Err(code) => {
release_pool_key_lease_from_report_context(&state, decision.report_context.as_ref())
.await;
planned_lease.release().await;
warn!(
event_name = "responses_websocket_initial_event_normalization_failed",
log_type = "ops",
@@ -334,13 +408,13 @@ pub(super) async fn run_responses_websocket(
"gateway could not normalize the initial Responses WebSocket event"
);
send_gateway_error(
&mut client_socket,
client_socket,
code,
"Gateway could not prepare the Responses response.create event",
)
.await;
close_client_socket(&mut client_socket, CLOSE_POLICY_VIOLATION, code).await;
return;
close_client_socket(client_socket, CLOSE_POLICY_VIOLATION, code).await;
return None;
}
};
let first_logical_turn_id = Uuid::new_v4().to_string();
@@ -355,12 +429,14 @@ pub(super) async fn run_responses_websocket(
&first_logical_turn_id,
1,
);
let mut first_turn = match begin_responses_websocket_turn(
let mut first_turn = match begin_responses_websocket_turn_with_planned_lease(
&state,
&planning_parts,
&context.decision,
&context.trace_id,
planning_parts,
&turn_control.decision,
first_turn_decision,
&first_event,
planned_lease,
)
.await
{
@@ -375,10 +451,10 @@ pub(super) async fn run_responses_websocket(
error = ?error,
"gateway could not start Responses WebSocket usage/audit lifecycle"
);
send_responses_websocket_turn_start_error(&mut client_socket, &error).await;
send_responses_websocket_turn_start_error(client_socket, &error).await;
let (close_code, close_reason) = responses_websocket_turn_start_close(&error);
close_client_socket(&mut client_socket, close_code, close_reason).await;
return;
close_client_socket(client_socket, close_code, close_reason).await;
return None;
}
};
@@ -390,8 +466,7 @@ pub(super) async fn run_responses_websocket(
state.clone(),
first_turn,
ResponsesWebSocketTurnOutcome::upstream_connect_failed(code),
)
.await;
);
warn!(
event_name = "responses_websocket_upstream_connect_failed",
log_type = "ops",
@@ -402,15 +477,15 @@ pub(super) async fn run_responses_websocket(
"gateway failed to establish Responses WebSocket upstream"
);
send_gateway_error_with_status(
&mut client_socket,
client_socket,
502,
code,
"Gateway could not establish the Provider connection",
)
.await;
close_client_socket(&mut client_socket, CLOSE_TRY_AGAIN, code).await;
close_client_socket(client_socket, CLOSE_TRY_AGAIN, code).await;
await_turn_finalization_handle(finalizer).await;
return;
return None;
}
};
first_turn.mark_upstream_request_sent();
@@ -419,20 +494,103 @@ pub(super) async fn run_responses_websocket(
bound.redaction_restorer.register(session);
}
bound.turn_state.begin(
LogicalTurn::new(first_event, 1, first_logical_turn_id),
ActiveProviderAttempt::new(&state, first_turn),
LogicalTurn::new(first_event, 1, first_logical_turn_id).with_turn_control(turn_control),
first_turn,
);
relay_bound_connection(
&mut client_socket,
&mut bound,
&state,
&context,
connection_permit,
)
.await;
await_pending_turn_finalization(&mut bound).await;
await_pending_adapter_observation(&mut bound).await;
Some(bound)
}
async fn close_terminated_bootstrap(
client_socket: &mut WebSocket,
context: &WebSocketRequestContext,
termination: ConnectionTermination,
) {
match termination {
ConnectionTermination::ConnectionLimitReached => {
warn!(
event_name = "responses_websocket_bootstrap_connection_limit_reached",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
"gateway stopped a Responses WebSocket bootstrap at the absolute connection deadline"
);
send_gateway_error_with_status(
client_socket,
503,
"websocket_connection_limit_reached",
"WebSocket connection duration limit reached; reconnect to continue",
)
.await;
close_client_socket(client_socket, CLOSE_TRY_AGAIN, "connection_limit_reached").await;
}
ConnectionTermination::ConnectionAdmissionLost => {
let policy = fatal_relay_policy(FatalRelaySignal::ConnectionAdmissionLost);
warn!(
event_name = "responses_websocket_bootstrap_connection_admission_lost",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
"gateway stopped a Responses WebSocket bootstrap after its connection admission became unhealthy"
);
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;
}
}
}
async fn close_terminated_relay(
client_socket: &mut WebSocket,
context: &WebSocketRequestContext,
termination: ConnectionTermination,
) {
match termination {
ConnectionTermination::ConnectionLimitReached => {
warn!(
event_name = "responses_websocket_connection_limit_reached",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
"gateway stopped a Responses WebSocket relay at the absolute connection deadline"
);
send_gateway_error_with_status(
client_socket,
503,
"websocket_connection_limit_reached",
"WebSocket connection duration limit reached; reconnect to continue",
)
.await;
close_client_socket(client_socket, CLOSE_TRY_AGAIN, "connection_limit_reached").await;
}
ConnectionTermination::ConnectionAdmissionLost => {
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 stopped a Responses WebSocket relay after its connection admission became unhealthy"
);
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;
}
}
}
/// 等待客户端发送第一条 response.create 事件。
@@ -477,9 +635,9 @@ where
let message = message.map_err(|_| InitialMessageError::ClientRead)?;
match message {
AxumWsMessage::Ping(payload) => {
socket
.send(AxumWsMessage::Pong(payload))
tokio::time::timeout_at(deadline, socket.send(AxumWsMessage::Pong(payload)))
.await
.map_err(|_| InitialMessageError::TimedOut)?
.map_err(|_| InitialMessageError::ClientRead)?;
}
AxumWsMessage::Pong(_) => {}
@@ -501,21 +659,17 @@ fn validate_initial_response_create(event: &Value) -> Result<(), InitialMessageE
if object.get("type").and_then(Value::as_str) != Some("response.create") {
return Err(InitialMessageError::MissingResponseCreate);
}
if object
let model = object
.get("model")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.is_none()
{
return Err(InitialMessageError::MissingModel);
}
.ok_or(InitialMessageError::MissingModel)?;
validated_response_create_model(model).map_err(|_| InitialMessageError::InvalidModel)?;
Ok(())
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
@@ -525,15 +679,13 @@ mod tests {
use super::super::binding::UpstreamBindingIdentity;
use super::super::client::adapter_drain_ready;
use super::super::quota::{
active_continuation_can_retry_from_full_input, is_usage_limit_error_event,
observe_active_response_rebind_safety, record_exhausted_bound_key,
should_request_full_continuation_retry,
is_usage_limit_error_event, observe_active_response_rebind_safety,
record_exhausted_bound_key,
};
use super::super::redaction::ResponsesWebSocketRedactionRestorer;
use super::super::request::{
changed_followup_response_create_model, continuation_requires_same_upstream,
normalize_followup_response_create, planned_response_create_event,
response_create_model_or_current,
changed_followup_response_create_model, normalize_followup_response_create,
planned_response_create_event, response_create_model_or_current,
};
use super::super::state::{BoundResponsesConnection, ExhaustedResponsesWebSocketExclusions};
use super::super::turn::{
@@ -569,6 +721,82 @@ mod tests {
event: serde_json::Value,
}
struct BootstrapDropProbe(Arc<AtomicUsize>);
impl Drop for BootstrapDropProbe {
fn drop(&mut self) {
self.0.fetch_add(1, Ordering::SeqCst);
}
}
struct TestAdmissionHealth(Arc<AtomicBool>);
impl aether_runtime::AdmissionPermitHealth for TestAdmissionHealth {
fn is_healthy(&self) -> bool {
self.0.load(Ordering::Acquire)
}
}
#[tokio::test]
async fn phase_supervisor_drops_work_at_the_absolute_connection_deadline() {
let dropped = Arc::new(AtomicUsize::new(0));
let task_dropped = Arc::clone(&dropped);
let bootstrap = async move {
let _probe = BootstrapDropProbe(task_dropped);
std::future::pending::<()>().await;
};
let result = tokio::time::timeout(
Duration::from_secs(1),
super::supervise_responses_websocket_phase(
bootstrap,
tokio::time::Instant::now() + Duration::from_millis(20),
None,
),
)
.await
.expect("bootstrap supervisor should honor its absolute deadline");
assert_eq!(
result,
Err(super::ConnectionTermination::ConnectionLimitReached)
);
assert_eq!(dropped.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn phase_supervisor_drops_work_when_connection_admission_is_lost() {
let health = Arc::new(AtomicBool::new(false));
let permit = aether_runtime::AdmissionPermit::from_parts(
None,
Some(TestAdmissionHealth(Arc::clone(&health))),
)
.expect("distributed test health should create a permit");
let dropped = Arc::new(AtomicUsize::new(0));
let task_dropped = Arc::clone(&dropped);
let bootstrap = async move {
let _probe = BootstrapDropProbe(task_dropped);
std::future::pending::<()>().await;
};
let result = tokio::time::timeout(
Duration::from_secs(1),
super::supervise_responses_websocket_phase(
bootstrap,
tokio::time::Instant::now() + Duration::from_secs(30),
Some(&permit),
),
)
.await
.expect("unhealthy connection admission should stop bootstrap");
assert_eq!(
result,
Err(super::ConnectionTermination::ConnectionAdmissionLost)
);
assert_eq!(dropped.load(Ordering::SeqCst), 1);
}
#[test]
fn adapter_drain_waits_for_an_active_turn_terminal_event() {
let directive = Some(ResponsesWebSocketDrainDirective {
@@ -719,45 +947,7 @@ mod tests {
}
#[test]
fn continuation_requires_the_existing_upstream_connection_and_account() {
let continuation = json!({
"type": "response.create",
"previous_response_id": "resp-previous",
});
assert!(!continuation_requires_same_upstream(&continuation, true));
assert!(continuation_requires_same_upstream(&continuation, false));
assert!(!continuation_requires_same_upstream(
&json!({"type": "response.create"}),
false,
));
}
#[test]
fn quota_error_can_request_a_full_retry_only_before_public_response_state() {
let mut bound = sample_bound_for_rebind_safety();
bound.turn_state = ResponsesTurnState::Replanning {
logical: LogicalTurn::new(
json!({
"type": "response.create",
"previous_response_id": "resp-previous",
}),
2,
"logical-turn".to_string(),
),
};
assert!(active_continuation_can_retry_from_full_input(&bound));
bound
.turn_state
.logical_mut()
.expect("active request")
.mark_retry_unsafe("standard_response_event");
assert!(!active_continuation_can_retry_from_full_input(&bound));
}
#[test]
fn only_an_actual_usage_limit_error_requests_full_retry() {
fn only_an_actual_usage_limit_error_requests_transparent_retry() {
assert!(is_usage_limit_error_event(&json!({
"type": "error",
"error": {"type": "usage_limit_reached"},
@@ -773,46 +963,6 @@ mod tests {
})));
}
#[test]
fn full_continuation_retry_does_not_consume_a_successful_terminal_event() {
let mut bound = sample_bound_for_rebind_safety();
bound.turn_state = ResponsesTurnState::Replanning {
logical: LogicalTurn::new(
json!({
"type": "response.create",
"previous_response_id": "resp-previous",
}),
2,
"logical-turn".to_string(),
),
};
assert!(should_request_full_continuation_retry(
&bound,
true,
Some(&json!({
"type": "error",
"error": {"type": "usage_limit_reached"},
})),
));
assert!(!should_request_full_continuation_retry(
&bound,
true,
Some(&json!({
"type": "response.completed",
"response": {"id": "resp-completed"},
})),
));
assert!(!should_request_full_continuation_retry(
&bound,
false,
Some(&json!({
"type": "error",
"error": {"type": "usage_limit_reached"},
})),
));
}
#[test]
fn followup_rewrites_the_provider_model_and_removes_http_stream_fields() {
let event = json!({
@@ -856,6 +1006,41 @@ mod tests {
);
}
#[test]
fn oversized_model_is_rejected_before_initial_or_followup_planning() {
let oversized = "m".repeat(257);
let initial = json!({
"type": "response.create",
"model": oversized,
"input": [],
});
assert!(matches!(
super::validate_initial_response_create(&initial),
Err(super::InitialMessageError::InvalidModel)
));
assert_eq!(
changed_followup_response_create_model(&initial, "current-model"),
Err("invalid_response_create_model")
);
}
#[test]
fn model_at_the_identifier_limit_remains_valid() {
let model = "m".repeat(256);
let event = json!({
"type": "response.create",
"model": model,
"input": [],
});
assert!(super::validate_initial_response_create(&event).is_ok());
assert_eq!(
changed_followup_response_create_model(&event, "current-model"),
Ok(Some(model))
);
}
#[test]
fn followup_without_a_model_reuses_the_current_connection_model() {
let event = json!({
@@ -1282,6 +1467,77 @@ mod tests {
}
}
/// Emits one Ping and then permanently backpressures every Pong write.
/// This models a peer that keeps the TCP connection open but never reads.
struct StalledPongSocket {
emitted_ping: bool,
}
impl futures_util::Stream for StalledPongSocket {
type Item = Result<axum::extract::ws::Message, axum::Error>;
fn poll_next(
mut self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
if self.emitted_ping {
std::task::Poll::Pending
} else {
self.emitted_ping = true;
std::task::Poll::Ready(Some(Ok(axum::extract::ws::Message::Ping(vec![7].into()))))
}
}
}
impl futures_util::Sink<axum::extract::ws::Message> for StalledPongSocket {
type Error = axum::Error;
fn poll_ready(
self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
std::task::Poll::Pending
}
fn start_send(
self: std::pin::Pin<&mut Self>,
_item: axum::extract::ws::Message,
) -> Result<(), Self::Error> {
unreachable!("a permanently backpressured sink is never ready")
}
fn poll_flush(
self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
std::task::Poll::Pending
}
fn poll_close(
self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
std::task::Poll::Pending
}
}
#[tokio::test]
async fn initial_message_deadline_also_bounds_a_stalled_pong_write() {
use super::{receive_initial_response_create_with_deadline, InitialMessageError};
let mut socket = StalledPongSocket {
emitted_ping: false,
};
let result = tokio::time::timeout(
Duration::from_millis(300),
receive_initial_response_create_with_deadline(&mut socket, Duration::from_millis(30)),
)
.await
.expect("the initial-message deadline must cancel a stalled Pong write");
assert!(matches!(result, Err(InitialMessageError::TimedOut)));
}
/// 验证 receive_initial_response_create_with_deadline 的绝对 deadline:
/// 客户端周期性发送 Ping 帧不会重置计时器,deadline 到期后返回 TimedOut。
/// 这直接驱动真实的循环逻辑,如果改回每次迭代 timeout(budget, ...) 则会变红。
@@ -38,8 +38,7 @@ use super::settlement::attempt_facts_for_outcome;
use crate::ai_serving::{build_openai_responses_stream_plan_from_decision, AiExecutionDecision};
use crate::clock::current_unix_ms;
use crate::control::{
execution_plan_balance_capacity_rejection, refresh_execution_runtime_auth_context,
request_model_local_rejection, GatewayControlDecision, GatewayLocalAuthRejection,
execution_plan_balance_capacity_rejection, GatewayControlDecision, GatewayLocalAuthRejection,
};
use crate::execution_runtime::attach_provider_response_headers_to_report_context;
use crate::execution_runtime::attempt_lifecycle::{
@@ -208,6 +207,20 @@ impl ResponsesWebSocketTurnOutcome {
}
}
pub(super) const fn relay_task_abandoned_before_upstream_send() -> Self {
Self::Cancelled {
reason: "gateway abandoned the turn before sending response.create upstream",
}
}
pub(super) const fn relay_task_abandonment(upstream_request_sent: bool) -> Self {
if upstream_request_sent {
Self::relay_task_abandoned()
} else {
Self::relay_task_abandoned_before_upstream_send()
}
}
const fn status_code(self) -> u16 {
match self {
Self::ProviderTerminal { status_code, .. } | Self::Failure { status_code, .. } => {
@@ -226,6 +239,8 @@ pub(super) struct ResponsesProviderAttempt {
/// 记账三段(pending / started / terminal)由共享的 transport 中立生命周期负责。
lifecycle: ExecutionAttemptLifecycle,
started_at: Instant,
provider_request_started_at_unix_ms: u64,
provider_request_order_id: String,
provider_headers: BTreeMap<String, String>,
observer: ResponsesStructuredTerminalObserver,
provider_capture: AttemptBodyCapture,
@@ -241,6 +256,10 @@ pub(super) struct ResponsesProviderAttempt {
provider_outcome: Option<AttemptProviderOutcome>,
/// 这一个 attempt 的内容是否完整交付给了客户端。与 provider 终态正交。
client_delivery: AttemptClientDelivery,
/// True only after `response.create` has been accepted by the upstream
/// socket writer. Cancellation before this point must not be projected as
/// provider failure or billed usage.
upstream_request_sent: bool,
}
/// 组装一轮 turn 的 decision。
@@ -282,7 +301,7 @@ pub(super) fn prepare_responses_websocket_turn_decision(
decision
}
pub(super) async fn begin_responses_websocket_turn(
pub(super) async fn begin_unowned_responses_websocket_turn(
state: &AppState,
parts: &http::request::Parts,
control_decision: &GatewayControlDecision,
@@ -290,17 +309,6 @@ pub(super) async fn begin_responses_websocket_turn(
client_event: &Value,
) -> Result<ResponsesProviderAttempt, GatewayError> {
let planned_report_context = decision.report_context.clone();
let effective_control_decision =
match refresh_websocket_turn_auth_context(state, control_decision, parts, client_event)
.await
{
Ok(decision) => decision,
Err(error) => {
release_pool_key_lease_from_report_context(state, planned_report_context.as_ref())
.await;
return Err(error);
}
};
let attempt = match build_openai_responses_stream_plan_from_decision(
parts,
client_event,
@@ -344,7 +352,7 @@ pub(super) async fn begin_responses_websocket_turn(
let balance_rejection = execution_plan_balance_capacity_rejection(
state,
&effective_control_decision,
control_decision,
&plan,
report_context.as_ref(),
)
@@ -375,6 +383,7 @@ pub(super) async fn begin_responses_websocket_turn(
return Err(websocket_auth_rejection_error(rejection));
}
let candidate_started_at_unix_ms = current_unix_ms();
ensure_execution_request_candidate_slot(state, &mut plan, &mut report_context).await;
let admission = match ResponsesWebSocketTurnAdmission::acquire(
state,
@@ -385,12 +394,21 @@ pub(super) async fn begin_responses_websocket_turn(
{
Ok(admission) => admission,
Err(error) => {
release_local_pool_key_lease(
state,
LocalExecutionEffectContext {
plan: &plan,
report_context: report_context.as_ref(),
},
release_then_record_responses_websocket_admission_failure(
release_local_pool_key_lease(
state,
LocalExecutionEffectContext {
plan: &plan,
report_context: report_context.as_ref(),
},
),
record_responses_websocket_admission_failure(
state,
&plan,
report_context.as_ref(),
candidate_started_at_unix_ms,
&error,
),
)
.await;
return Err(error);
@@ -413,6 +431,8 @@ pub(super) async fn begin_responses_websocket_turn(
Ok(ResponsesProviderAttempt {
lifecycle,
started_at: Instant::now(),
provider_request_started_at_unix_ms: current_unix_ms(),
provider_request_order_id: uuid::Uuid::now_v7().to_string(),
provider_headers: BTreeMap::new(),
observer: ResponsesStructuredTerminalObserver::default(),
provider_capture: AttemptBodyCapture::default(),
@@ -425,40 +445,72 @@ pub(super) async fn begin_responses_websocket_turn(
terminal_error_body: None,
provider_outcome: None,
client_delivery: AttemptClientDelivery::Complete,
upstream_request_sent: false,
})
}
async fn refresh_websocket_turn_auth_context(
state: &AppState,
control_decision: &GatewayControlDecision,
parts: &http::request::Parts,
client_event: &Value,
) -> Result<GatewayControlDecision, GatewayError> {
let mut effective = control_decision.clone();
if let Some(auth_context) = effective.auth_context.take() {
let refreshed = refresh_execution_runtime_auth_context(
state,
auth_context,
effective.auth_endpoint_signature.as_deref(),
)
.await?;
effective.local_auth_rejection = refreshed.local_rejection.clone();
effective.auth_context = Some(refreshed);
}
if let Some(rejection) = effective.local_auth_rejection.clone() {
return Err(websocket_auth_rejection_error(rejection));
}
async fn release_then_record_responses_websocket_admission_failure(
release_pool_lease: impl Future<Output = ()>,
record_candidate_failure: impl Future<Output = ()>,
) {
// Lease cleanup protects live routing capacity and must not sit behind a
// slow candidate writer. The candidate write still follows immediately so
// the seeded row reaches a terminal state on the ordinary error path.
release_pool_lease.await;
record_candidate_failure.await;
}
let body = serde_json::to_vec(client_event)
.map(axum::body::Bytes::from)
.map_err(|error| GatewayError::Internal(error.to_string()))?;
if let Some(rejection) =
request_model_local_rejection(state, Some(&effective), &parts.uri, &parts.headers, &body)
.await?
{
return Err(websocket_auth_rejection_error(rejection));
async fn record_responses_websocket_admission_failure(
state: &AppState,
plan: &ExecutionPlan,
report_context: Option<&Value>,
candidate_started_at_unix_ms: u64,
error: &GatewayError,
) {
let terminal_at_unix_ms = current_unix_ms();
record_local_request_candidate_status(
state,
plan,
report_context,
responses_websocket_admission_failure_update(
candidate_started_at_unix_ms,
terminal_at_unix_ms,
error,
),
)
.await;
}
fn responses_websocket_admission_failure_update(
candidate_started_at_unix_ms: u64,
terminal_at_unix_ms: u64,
error: &GatewayError,
) -> SchedulerRequestCandidateStatusUpdate {
let (status_code, error_type, error_message) = match error {
GatewayError::AdmissionTimeout {
gate,
queue_budget_ms,
..
} => (
StatusCode::TOO_MANY_REQUESTS.as_u16(),
"gateway_admission_timeout",
format!("gateway admission gate {gate} timed out after {queue_budget_ms}ms"),
),
other => (
StatusCode::INTERNAL_SERVER_ERROR.as_u16(),
"gateway_admission_failed",
format!("{other:?}"),
),
};
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed,
status_code: Some(status_code),
error_type: Some(error_type.to_string()),
error_message: Some(error_message),
latency_ms: Some(terminal_at_unix_ms.saturating_sub(candidate_started_at_unix_ms)),
started_at_unix_ms: Some(candidate_started_at_unix_ms),
finished_at_unix_ms: Some(terminal_at_unix_ms),
}
Ok(effective)
}
fn websocket_auth_rejection_error(rejection: GatewayLocalAuthRejection) -> GatewayError {
@@ -526,9 +578,13 @@ impl ResponsesProviderAttempt {
}
pub(super) fn set_provider_response_headers(&mut self, headers: BTreeMap<String, String>) {
let observed_at_unix_ms = current_unix_ms();
let report_context = attach_provider_response_headers_to_report_context(
self.lifecycle.take_report_context(),
&headers,
self.provider_request_started_at_unix_ms,
observed_at_unix_ms,
&self.provider_request_order_id,
);
self.lifecycle.set_report_context(report_context);
self.provider_headers = headers;
@@ -537,10 +593,21 @@ impl ResponsesProviderAttempt {
/// Starts the per-turn response deadlines only after the corresponding
/// `response.create` has been accepted by the upstream socket writer.
pub(super) fn mark_upstream_request_sent(&mut self) {
self.upstream_request_sent = true;
self.started_at = Instant::now();
self.provider_request_started_at_unix_ms = current_unix_ms();
self.provider_request_order_id = uuid::Uuid::now_v7().to_string();
self.first_event_elapsed_ms = None;
}
/// Selects a cancellation-safe fallback for an attempt whose owner task
/// disappeared. Before the upstream write this is a void cancellation;
/// after the write it remains a gateway relay failure because provider
/// work may already have started.
pub(super) const fn abandonment_outcome(&self) -> ResponsesWebSocketTurnOutcome {
ResponsesWebSocketTurnOutcome::relay_task_abandonment(self.upstream_request_sent)
}
pub(super) fn deadline(&self) -> ResponsesWebSocketTurnDeadline {
let (phase, timeout) = if self.first_event_elapsed_ms.is_some() {
(
@@ -770,17 +837,6 @@ impl ResponsesProviderAttempt {
}
}
pub(super) async fn spawn_responses_websocket_turn_finalization(
state: AppState,
mut turn: ResponsesProviderAttempt,
outcome: ResponsesWebSocketTurnOutcome,
) -> tokio::task::JoinHandle<()> {
turn.release_admission().await;
tokio::spawn(async move {
turn.settle(&state, outcome).await;
})
}
/// 把一轮 turn 的事实写进审计/用量 report context。
///
/// `effective_client_event` 是脱敏后的客户端事件(未启用脱敏时就是原事件)。
@@ -945,9 +1001,12 @@ fn elapsed_ms(started_at: Instant) -> u64 {
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicU8, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use aether_contracts::ExecutionTimeouts;
use aether_data_contracts::repository::candidates::RequestCandidateStatus;
use serde_json::json;
use super::super::observation::ResponsesStructuredTerminalObserver;
@@ -958,7 +1017,8 @@ mod tests {
};
use super::{
attach_client_delivery_to_report_context, prepare_websocket_report_context,
provider_terminal_outcome, resolve_responses_websocket_turn_timeouts,
provider_terminal_outcome, release_then_record_responses_websocket_admission_failure,
resolve_responses_websocket_turn_timeouts, responses_websocket_admission_failure_update,
websocket_event_as_sse_line, ResponsesWebSocketTurnDeadline, ResponsesWebSocketTurnOutcome,
ResponsesWebSocketTurnTimeoutPhase,
};
@@ -966,6 +1026,72 @@ mod tests {
classify_attempt_settlement, AttemptBilling, AttemptCandidateError, AttemptCandidateStatus,
AttemptClientDelivery, AttemptProviderEffect, AttemptSettlementInputs,
};
use crate::GatewayError;
#[tokio::test]
async fn admission_failure_releases_pool_lease_before_recording_candidate_terminal() {
let phase = Arc::new(AtomicU8::new(0));
let release_phase = Arc::clone(&phase);
let record_phase = Arc::clone(&phase);
release_then_record_responses_websocket_admission_failure(
async move {
assert_eq!(release_phase.swap(1, Ordering::SeqCst), 0);
},
async move {
assert_eq!(record_phase.swap(2, Ordering::SeqCst), 1);
},
)
.await;
assert_eq!(phase.load(Ordering::SeqCst), 2);
}
#[test]
fn admission_timeout_terminalizes_the_seeded_candidate() {
let update = responses_websocket_admission_failure_update(
1_000,
1_025,
&GatewayError::AdmissionTimeout {
trace_id: "turn-1".to_string(),
gate: "gateway_upstream_execution",
queue_budget_ms: 25,
},
);
assert_eq!(update.status, RequestCandidateStatus::Failed);
assert_eq!(update.status_code, Some(429));
assert_eq!(
update.error_type.as_deref(),
Some("gateway_admission_timeout")
);
assert_eq!(
update.error_message.as_deref(),
Some("gateway admission gate gateway_upstream_execution timed out after 25ms")
);
assert_eq!(update.latency_ms, Some(25));
assert_eq!(update.started_at_unix_ms, Some(1_000));
assert_eq!(update.finished_at_unix_ms, Some(1_025));
}
#[test]
fn non_timeout_admission_failure_still_terminalizes_the_seeded_candidate() {
let update = responses_websocket_admission_failure_update(
50,
40,
&GatewayError::Internal("admission gate closed".to_string()),
);
assert_eq!(update.status, RequestCandidateStatus::Failed);
assert_eq!(update.status_code, Some(500));
assert_eq!(
update.error_type.as_deref(),
Some("gateway_admission_failed")
);
assert_eq!(update.latency_ms, Some(0));
assert_eq!(update.started_at_unix_ms, Some(50));
assert_eq!(update.finished_at_unix_ms, Some(40));
}
#[test]
fn followup_context_uses_a_fresh_request_and_candidate() {
@@ -1321,11 +1447,8 @@ mod tests {
}
#[test]
fn an_abandoned_turn_is_recorded_as_a_gateway_failure_not_a_cancellation() {
// A turn reclaimed by the Drop guard must not look like a client
// cancellation: cancelled turns skip the stream report entirely, which
// would defeat the point of reclaiming it.
let outcome = ResponsesWebSocketTurnOutcome::relay_task_abandoned();
fn an_abandoned_turn_after_upstream_send_is_recorded_as_a_gateway_failure() {
let outcome = ResponsesWebSocketTurnOutcome::relay_task_abandonment(true);
let facts = attempt_facts_for_outcome(None, AttemptClientDelivery::Complete, outcome);
let settlement = classify_attempt_settlement(AttemptSettlementInputs {
facts,
@@ -1341,6 +1464,31 @@ mod tests {
assert!(settlement.submit_execution_report);
}
#[test]
fn an_abandoned_turn_before_upstream_send_is_void_and_does_not_penalize_provider() {
let outcome = ResponsesWebSocketTurnOutcome::relay_task_abandonment(false);
let facts = attempt_facts_for_outcome(None, AttemptClientDelivery::Complete, outcome);
let settlement = classify_attempt_settlement(AttemptSettlementInputs {
facts,
report_represents_failure: false,
observed_finish: false,
has_parser_error: false,
});
assert_eq!(settlement.status_code, 499);
assert_eq!(settlement.billing, AttemptBilling::Void);
assert_eq!(
settlement.candidate_status,
AttemptCandidateStatus::Cancelled
);
assert_eq!(settlement.candidate_error, AttemptCandidateError::Cancelled);
assert_eq!(
settlement.provider_effect,
AttemptProviderEffect::ReleasePoolKeyLease
);
assert!(!settlement.submit_execution_report);
}
#[test]
fn turn_timeouts_reuse_provider_first_byte_and_request_deadlines() {
let (first_event, terminal) =
@@ -7,6 +7,7 @@
use serde_json::Value;
use super::control::ResponsesWebSocketTurnControl;
use super::lifecycle::ActiveProviderAttempt;
use super::request::response_create_has_previous_response_id;
@@ -24,6 +25,10 @@ pub(super) struct LogicalTurn {
pub(super) turn_attempt: u32,
pub(super) retry_attempted: bool,
pub(super) retry_unsafe_reason: Option<&'static str>,
/// Exact live control decision and strong auth snapshot used to authorize
/// this logical turn. Quota retries reuse it instead of falling back to the
/// connection's Upgrade-time authorization snapshot.
pub(super) turn_control: Option<ResponsesWebSocketTurnControl>,
}
impl LogicalTurn {
@@ -35,9 +40,15 @@ impl LogicalTurn {
turn_attempt: 1,
retry_attempted: false,
retry_unsafe_reason: None,
turn_control: None,
}
}
pub(super) fn with_turn_control(mut self, turn_control: ResponsesWebSocketTurnControl) -> Self {
self.turn_control = Some(turn_control);
self
}
pub(super) fn quota_retry_block_reason(&self) -> Option<&'static str> {
if self.retry_attempted {
Some("quota_retry_already_attempted")
@@ -10,9 +10,9 @@ 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,
PROXY_AUTHORIZATION, TE, TRAILER, TRANSFER_ENCODING, UPGRADE,
};
use axum::http::HeaderMap;
use axum::http::{HeaderMap, HeaderName};
use futures_util::{SinkExt, TryFutureExt};
use serde_json::json;
use url::Url;
@@ -113,8 +113,20 @@ pub(crate) fn websocket_handshake_headers(
provider_headers: &BTreeMap<String, String>,
invalid_code: &'static str,
) -> Result<HeaderMap, &'static str> {
// `build_request_headers` already strips `Connection` itself. Read the
// dynamic hop-by-hop names from the source map first, otherwise a header
// named by `Connection: keep-alive, x-provider-hop` would survive.
let connection_scoped_names = provider_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case(CONNECTION.as_str()))
.flat_map(|(_, value)| value.split(','))
.filter_map(|name| HeaderName::from_bytes(name.trim().as_bytes()).ok())
.collect::<Vec<_>>();
let mut headers =
build_request_headers(provider_headers, None, false).map_err(|_| invalid_code)?;
for name in connection_scoped_names {
headers.remove(name);
}
for header in [
ACCEPT,
ACCEPT_ENCODING,
@@ -123,11 +135,29 @@ pub(crate) fn websocket_handshake_headers(
CONTENT_LENGTH,
CONTENT_TYPE,
HOST,
PROXY_AUTHORIZATION,
TE,
TRAILER,
TRANSFER_ENCODING,
UPGRADE,
] {
headers.remove(header);
}
for header in ["keep-alive", "proxy-connection"] {
headers.remove(header);
}
// The WebSocket client owns every Sec-WebSocket-* field, including
// extensions introduced after this gateway was built. Passing a
// downstream handshake field through here can corrupt negotiation or
// disclose the client's nonce/subprotocol to a different upstream.
let websocket_managed_names = headers
.keys()
.filter(|name| name.as_str().starts_with("sec-websocket-"))
.cloned()
.collect::<Vec<_>>();
for name in websocket_managed_names {
headers.remove(name);
}
Ok(headers)
}
@@ -261,13 +291,6 @@ pub(crate) fn upstream_message_to_client(message: WreqWsMessage) -> AxumWsMessag
}
}
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.
@@ -332,9 +355,10 @@ pub(crate) async fn close_client_socket(client_socket: &mut WebSocket, code: u16
#[cfg(test)]
mod tests {
use super::{
bounded_send, responses_websocket_error_event, websocket_upstream_url, WebSocketWriteError,
RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT,
bounded_send, responses_websocket_error_event, websocket_handshake_headers,
websocket_upstream_url, WebSocketWriteError, RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT,
};
use std::collections::BTreeMap;
use std::time::Duration;
#[tokio::test]
@@ -403,4 +427,74 @@ mod tests {
fn rejects_upstream_url_with_credentials() {
assert!(websocket_upstream_url("https://[email protected]/responses", "invalid").is_err());
}
#[test]
fn upstream_handshake_keeps_provider_auth_but_drops_transport_managed_headers() {
let provider_headers = BTreeMap::from([
(
"authorization".to_string(),
"Bearer provider-token".to_string(),
),
("x-api-key".to_string(), "provider-api-key".to_string()),
(
"cookie".to_string(),
"provider_session=provider-cookie".to_string(),
),
(
"connection".to_string(),
"keep-alive, x-provider-hop".to_string(),
),
("x-provider-hop".to_string(), "must-not-pass".to_string()),
("upgrade".to_string(), "websocket".to_string()),
(
"sec-websocket-key".to_string(),
"downstream-nonce".to_string(),
),
(
"sec-websocket-future-field".to_string(),
"future-value".to_string(),
),
(
"proxy-authorization".to_string(),
"Basic must-not-pass".to_string(),
),
("x-provider-header".to_string(), "safe".to_string()),
]);
let headers = websocket_handshake_headers(&provider_headers, "invalid")
.expect("provider headers should be valid");
assert_eq!(
headers
.get("authorization")
.and_then(|value| value.to_str().ok()),
Some("Bearer provider-token")
);
assert_eq!(
headers
.get("x-api-key")
.and_then(|value| value.to_str().ok()),
Some("provider-api-key")
);
assert_eq!(
headers.get("cookie").and_then(|value| value.to_str().ok()),
Some("provider_session=provider-cookie")
);
assert_eq!(
headers
.get("x-provider-header")
.and_then(|value| value.to_str().ok()),
Some("safe")
);
for name in [
"connection",
"x-provider-hop",
"upgrade",
"sec-websocket-key",
"sec-websocket-future-field",
"proxy-authorization",
] {
assert!(headers.get(name).is_none(), "{name} must not survive");
}
}
}