mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 16:37:46 +08:00
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:
@@ -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");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user