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