diff --git a/.github/workflows/rust-ci.yml b/.github/workflows/rust-ci.yml index 76d14603f..af8295df2 100644 --- a/.github/workflows/rust-ci.yml +++ b/.github/workflows/rust-ci.yml @@ -387,11 +387,11 @@ jobs: - name: Setup sccache uses: mozilla-actions/sccache-action@v0.0.9 - - name: Test scenario binaries + - name: Test scenario binaries and end-to-end suites env: RUSTC_WRAPPER: sccache SCCACHE_GHA_ENABLED: "true" - run: cargo test -p aether-integration-tests --bins + run: cargo test -p aether-integration-tests --bins --tests - name: Show sccache stats if: always() diff --git a/Cargo.lock b/Cargo.lock index c7b536424..03f498561 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -441,6 +441,7 @@ name = "aether-integration-tests" version = "0.1.0" dependencies = [ "aether-contracts", + "aether-crypto", "aether-data", "aether-data-contracts", "aether-gateway", @@ -457,6 +458,7 @@ dependencies = [ "sqlx", "tokio", "tokio-tungstenite 0.28.0", + "uuid", ] [[package]] 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..0a575243c 100644 --- a/apps/aether-gateway/src/ai_serving/mod.rs +++ b/apps/aether-gateway/src/ai_serving/mod.rs @@ -52,16 +52,18 @@ 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, + 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_provider_chat_pii_redaction, 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..55068e03e 100644 --- a/apps/aether-gateway/src/ai_serving/planner/mod.rs +++ b/apps/aether-gateway/src/ai_serving/planner/mod.rs @@ -49,6 +49,7 @@ pub(crate) use self::plan_builders::{ pub(crate) use self::pool_scores::{ build_provider_key_pool_score_upsert, provider_key_pool_score_id, provider_key_pool_score_scope, }; +pub(crate) use self::redaction::resolve_provider_chat_pii_redaction; pub(crate) use self::request_gzip::resolve_transport_request_encoding_policy; pub(crate) use self::route::is_matching_stream_request as planner_is_matching_stream_request; pub(crate) use self::runtime_miss::{ @@ -80,8 +81,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/bin/aether-codex-ws-probe.rs b/apps/aether-gateway/src/bin/aether-codex-ws-probe.rs new file mode 100644 index 000000000..b9def32c9 --- /dev/null +++ b/apps/aether-gateway/src/bin/aether-codex-ws-probe.rs @@ -0,0 +1,124 @@ +//! Credential-safe compatibility probe for the Codex Responses WebSocket path. +//! +//! This binary preserves the established Codex CLI and environment contract. +//! The common Responses WebSocket flow lives in `support/responses_ws_probe`; +//! this profile owns only Codex authentication and header requirements. + +#[path = "support/responses_ws_probe.rs"] +mod responses_ws_probe; + +use aether_gateway::{CODEX_CLIENT_ORIGINATOR, CODEX_CLIENT_USER_AGENT}; +use clap::Parser; +use http::header::{AUTHORIZATION, USER_AGENT}; +use http::{HeaderMap, HeaderName, HeaderValue}; +use responses_ws_probe::{ + bearer_authorization_value, required_env, resolve_probe_url, run_profile_probe, turn_timeout, + ProbeArgs, ProbeConfig, ProbeFailure, ResponsesWebSocketProbeProfile, +}; + +const ACCESS_TOKEN_ENV: &str = "AETHER_CODEX_WS_PROBE_ACCESS_TOKEN"; +const ACCOUNT_ID_ENV: &str = "AETHER_CODEX_WS_PROBE_ACCOUNT_ID"; +const MODEL_ENV: &str = "AETHER_CODEX_WS_PROBE_MODEL"; +const URL_ENV: &str = "AETHER_CODEX_WS_PROBE_URL"; + +#[derive(Parser)] +#[command( + name = "aether-codex-ws-probe", + about = "Verify a Codex Responses WebSocket endpoint without exposing credentials" +)] +struct Args { + /// WebSocket endpoint. If omitted, AETHER_CODEX_WS_PROBE_URL is used. + #[arg(long)] + url: Option, + + /// Per-turn receive timeout in seconds. + #[arg(long, default_value_t = 20, value_parser = clap::value_parser!(u64).range(1..=120))] + timeout_secs: u64, +} + +impl From for ProbeArgs { + fn from(args: Args) -> Self { + Self { + url: args.url, + timeout_secs: args.timeout_secs, + } + } +} + +struct CodexResponsesProbeProfile; + +impl ResponsesWebSocketProbeProfile for CodexResponsesProbeProfile { + fn build_config(args: &ProbeArgs) -> Result { + let url = resolve_probe_url(args, URL_ENV, None)?; + let access_token = required_env(ACCESS_TOKEN_ENV)?; + let account_id = required_env(ACCOUNT_ID_ENV)?; + let model = required_env(MODEL_ENV)?; + Ok(ProbeConfig::new( + url, + model, + turn_timeout(args), + handshake_headers(&access_token, &account_id)?, + Self::sent_header_names(), + )) + } + + fn sent_header_names() -> Vec<&'static str> { + vec![ + "authorization", + "chatgpt-account-id", + "user-agent", + "originator", + ] + } +} + +fn handshake_headers(access_token: &str, account_id: &str) -> Result { + let account_id = + HeaderValue::from_str(account_id).map_err(|_| ProbeFailure::MissingConfiguration)?; + let mut headers = HeaderMap::new(); + headers.insert(AUTHORIZATION, bearer_authorization_value(access_token)?); + headers.insert(HeaderName::from_static("chatgpt-account-id"), account_id); + headers.insert( + USER_AGENT, + HeaderValue::from_static(CODEX_CLIENT_USER_AGENT), + ); + headers.insert( + HeaderName::from_static("originator"), + HeaderValue::from_static(CODEX_CLIENT_ORIGINATOR), + ); + Ok(headers) +} + +#[tokio::main] +async fn main() { + let exit_code = run_profile_probe::(Args::parse().into()).await; + if exit_code != 0 { + std::process::exit(exit_code); + } +} + +#[cfg(test)] +mod tests { + use http::header::{AUTHORIZATION, USER_AGENT}; + + use super::{handshake_headers, CodexResponsesProbeProfile, ResponsesWebSocketProbeProfile}; + + #[test] + fn codex_profile_keeps_its_required_handshake_headers() { + let headers = + handshake_headers("test-token", "test-account").expect("headers should build"); + assert!(headers.contains_key(AUTHORIZATION)); + assert!(headers.contains_key("chatgpt-account-id")); + assert!(headers.contains_key(USER_AGENT)); + assert!(headers.contains_key("originator")); + assert_eq!( + CodexResponsesProbeProfile::sent_header_names(), + vec![ + "authorization", + "chatgpt-account-id", + "user-agent", + "originator", + ] + ); + } +} diff --git a/apps/aether-gateway/src/bin/aether-openai-responses-ws-probe.rs b/apps/aether-gateway/src/bin/aether-openai-responses-ws-probe.rs new file mode 100644 index 000000000..08838fc26 --- /dev/null +++ b/apps/aether-gateway/src/bin/aether-openai-responses-ws-probe.rs @@ -0,0 +1,104 @@ +//! Credential-safe compatibility probe for the official OpenAI Responses +//! WebSocket endpoint. +//! +//! This profile uses standard API-key Bearer authentication and shares the +//! protocol flow with the Codex probe without inheriting Codex-specific +//! account headers or quota assumptions. + +#[path = "support/responses_ws_probe.rs"] +mod responses_ws_probe; + +use clap::Parser; +use http::header::AUTHORIZATION; +use http::HeaderMap; +use responses_ws_probe::{ + bearer_authorization_value, required_env, resolve_probe_url, run_profile_probe, turn_timeout, + ProbeArgs, ProbeConfig, ProbeFailure, ResponsesWebSocketProbeProfile, +}; + +const API_KEY_ENV: &str = "AETHER_OPENAI_WS_PROBE_API_KEY"; +const MODEL_ENV: &str = "AETHER_OPENAI_WS_PROBE_MODEL"; +const URL_ENV: &str = "AETHER_OPENAI_WS_PROBE_URL"; +const DEFAULT_URL: &str = "wss://api.openai.com/v1/responses"; + +#[derive(Parser)] +#[command( + name = "aether-openai-responses-ws-probe", + about = "Verify an OpenAI Responses WebSocket endpoint without exposing credentials" +)] +struct Args { + /// WebSocket endpoint. If omitted, AETHER_OPENAI_WS_PROBE_URL or the + /// official OpenAI endpoint is used. + #[arg(long)] + url: Option, + + /// Per-turn receive timeout in seconds. + #[arg(long, default_value_t = 20, value_parser = clap::value_parser!(u64).range(1..=120))] + timeout_secs: u64, +} + +impl From for ProbeArgs { + fn from(args: Args) -> Self { + Self { + url: args.url, + timeout_secs: args.timeout_secs, + } + } +} + +struct OpenAiResponsesProbeProfile; + +impl ResponsesWebSocketProbeProfile for OpenAiResponsesProbeProfile { + fn build_config(args: &ProbeArgs) -> Result { + let url = resolve_probe_url(args, URL_ENV, Some(DEFAULT_URL))?; + let api_key = required_env(API_KEY_ENV)?; + let model = required_env(MODEL_ENV)?; + let mut headers = HeaderMap::new(); + headers.insert(AUTHORIZATION, bearer_authorization_value(&api_key)?); + Ok(ProbeConfig::new( + url, + model, + turn_timeout(args), + headers, + Self::sent_header_names(), + )) + } + + fn sent_header_names() -> Vec<&'static str> { + vec!["authorization"] + } +} + +#[tokio::main] +async fn main() { + let exit_code = run_profile_probe::(Args::parse().into()).await; + if exit_code != 0 { + std::process::exit(exit_code); + } +} + +#[cfg(test)] +mod tests { + use http::header::AUTHORIZATION; + + use super::{ + bearer_authorization_value, responses_ws_probe::parse_probe_url, + OpenAiResponsesProbeProfile, ResponsesWebSocketProbeProfile, DEFAULT_URL, + }; + + #[test] + fn openai_profile_exposes_only_standard_bearer_authentication() { + let authorization = bearer_authorization_value("test-key").expect("header should build"); + assert_eq!(authorization.to_str().ok(), Some("Bearer test-key")); + assert_eq!( + OpenAiResponsesProbeProfile::sent_header_names(), + vec![AUTHORIZATION.as_str()] + ); + } + + #[test] + fn openai_profile_uses_the_official_responses_websocket_endpoint_by_default() { + let url = parse_probe_url(DEFAULT_URL).expect("default OpenAI endpoint should be valid"); + assert_eq!(url.as_str(), DEFAULT_URL); + } +} diff --git a/apps/aether-gateway/src/bin/support/responses_ws_probe.rs b/apps/aether-gateway/src/bin/support/responses_ws_probe.rs new file mode 100644 index 000000000..061a97cf6 --- /dev/null +++ b/apps/aether-gateway/src/bin/support/responses_ws_probe.rs @@ -0,0 +1,567 @@ +//! Shared, credential-safe engine for Responses WebSocket compatibility probes. +//! +//! Provider profiles own their environment variables and handshake headers. +//! This module owns the common Responses WebSocket contract: two sequential +//! `response.create` warmups, continuation with `previous_response_id`, safe +//! event observation, and a redacted JSON report. + +use std::env; +use std::time::{Duration, Instant}; + +use http::{HeaderMap, HeaderValue}; +use serde::Serialize; +use serde_json::{json, Value}; +use url::Url; +use wreq::ws::message::Message as WreqWsMessage; + +const MAX_FRAME_SIZE: usize = 1 << 20; +const MAX_EVENTS_PER_TURN: usize = 16; + +pub(crate) struct ProbeArgs { + pub(crate) url: Option, + pub(crate) timeout_secs: u64, +} + +pub(crate) struct ProbeConfig { + url: Url, + model: String, + turn_timeout: Duration, + handshake_headers: HeaderMap, + sent_header_names: Vec<&'static str>, +} + +impl ProbeConfig { + pub(crate) fn new( + url: Url, + model: String, + turn_timeout: Duration, + handshake_headers: HeaderMap, + sent_header_names: Vec<&'static str>, + ) -> Self { + Self { + url, + model, + turn_timeout, + handshake_headers, + sent_header_names, + } + } +} + +/// A profile retains provider-specific authentication and configuration while +/// reusing one Responses protocol probe engine. +pub(crate) trait ResponsesWebSocketProbeProfile { + fn build_config(args: &ProbeArgs) -> Result; + fn sent_header_names() -> Vec<&'static str>; +} + +#[derive(Debug, Clone, Copy)] +pub(crate) enum ProbeFailure { + MissingConfiguration, + InvalidEndpoint, + ClientBuild, + Handshake, + Upgrade, + Send, + ReceiveTimeout, + Receive, + RemoteError, + MissingResponseId, + UnexpectedFrame, +} + +impl ProbeFailure { + const fn code(self) -> &'static str { + match self { + Self::MissingConfiguration => "missing_configuration", + Self::InvalidEndpoint => "invalid_endpoint", + Self::ClientBuild => "client_build_failed", + Self::Handshake => "handshake_failed", + Self::Upgrade => "upgrade_failed", + Self::Send => "send_failed", + Self::ReceiveTimeout => "receive_timeout", + Self::Receive => "receive_failed", + Self::RemoteError => "upstream_error_event", + Self::MissingResponseId => "response_id_not_observed", + Self::UnexpectedFrame => "unexpected_frame", + } + } +} + +#[derive(Serialize)] +struct ProbeReport { + status: &'static str, + target_host: Option, + handshake_status: Option, + sent_header_names: Vec<&'static str>, + received_header_names: Vec, + observed_event_types: Vec, + continuation_confirmed: bool, + elapsed_ms: u64, + error: Option<&'static str>, +} + +impl ProbeReport { + fn failed( + config: Option<&ProbeConfig>, + sent_header_names: Vec<&'static str>, + started_at: Instant, + error: ProbeFailure, + ) -> Self { + Self { + status: "failed", + target_host: config.and_then(target_host), + handshake_status: None, + sent_header_names, + received_header_names: Vec::new(), + observed_event_types: Vec::new(), + continuation_confirmed: false, + elapsed_ms: started_at.elapsed().as_millis() as u64, + error: Some(error.code()), + } + } +} + +/// Runs a profile and returns the process exit code after emitting exactly one +/// credential-safe JSON report. +pub(crate) async fn run_profile_probe(args: ProbeArgs) -> i32 { + let started_at = Instant::now(); + let config = match P::build_config(&args) { + Ok(config) => config, + Err(error) => { + print_report(&ProbeReport::failed( + None, + P::sent_header_names(), + started_at, + error, + )); + return 2; + } + }; + + match run_probe(&config, started_at).await { + Ok(report) => { + print_report(&report); + 0 + } + Err(error) => { + print_report(&ProbeReport::failed( + Some(&config), + config.sent_header_names.clone(), + started_at, + error, + )); + 1 + } + } +} + +pub(crate) fn required_env(name: &str) -> Result { + env::var(name) + .ok() + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) + .ok_or(ProbeFailure::MissingConfiguration) +} + +pub(crate) fn resolve_probe_url( + args: &ProbeArgs, + url_env: &str, + default_url: Option<&str>, +) -> Result { + let raw_url = args + .url + .as_deref() + .map(str::to_owned) + .or_else(|| env::var(url_env).ok()) + .or_else(|| default_url.map(str::to_owned)); + let Some(raw_url) = raw_url + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + else { + return Err(ProbeFailure::MissingConfiguration); + }; + parse_probe_url(raw_url) +} + +pub(crate) fn parse_probe_url(raw: &str) -> Result { + let url = Url::parse(raw).map_err(|_| ProbeFailure::InvalidEndpoint)?; + if !matches!(url.scheme(), "ws" | "wss") + || url.host_str().is_none() + || !url.username().is_empty() + || url.password().is_some() + || url.query().is_some() + || url.fragment().is_some() + { + return Err(ProbeFailure::InvalidEndpoint); + } + Ok(url) +} + +pub(crate) fn bearer_authorization_value(token: &str) -> Result { + HeaderValue::from_str(format!("Bearer {token}").as_str()) + .map_err(|_| ProbeFailure::MissingConfiguration) +} + +pub(crate) const fn turn_timeout(args: &ProbeArgs) -> Duration { + Duration::from_secs(args.timeout_secs) +} + +async fn run_probe(config: &ProbeConfig, started_at: Instant) -> Result { + let client = wreq::Client::builder() + .connect_timeout(config.turn_timeout) + .timeout(config.turn_timeout) + .build() + .map_err(|_| ProbeFailure::ClientBuild)?; + let response = client + .websocket(config.url.as_str()) + .headers(config.handshake_headers.clone()) + .max_frame_size(MAX_FRAME_SIZE) + .max_message_size(MAX_FRAME_SIZE) + .send() + .await + .map_err(|_| ProbeFailure::Handshake)?; + let handshake_status = response.status().as_u16(); + let received_header_names = response + .headers() + .keys() + .map(|name| name.as_str().to_string()) + .collect(); + let mut socket = response + .into_websocket() + .await + .map_err(|_| ProbeFailure::Upgrade)?; + let mut observed_event_types = Vec::new(); + + send_warmup(&mut socket, &config.model, None).await?; + let first_response_id = + receive_completed_response_id(&mut socket, config.turn_timeout, &mut observed_event_types) + .await?; + + send_warmup(&mut socket, &config.model, Some(&first_response_id)).await?; + let _second_response_id = + receive_completed_response_id(&mut socket, config.turn_timeout, &mut observed_event_types) + .await?; + + Ok(ProbeReport { + status: "passed", + target_host: target_host(config), + handshake_status: Some(handshake_status), + sent_header_names: config.sent_header_names.clone(), + received_header_names, + observed_event_types, + continuation_confirmed: true, + elapsed_ms: started_at.elapsed().as_millis() as u64, + error: None, + }) +} + +fn target_host(config: &ProbeConfig) -> Option { + config.url.host_str().map(|host| match config.url.port() { + Some(port) => format!("{host}:{port}"), + None => host.to_string(), + }) +} + +async fn send_warmup( + socket: &mut wreq::ws::WebSocket, + model: &str, + previous_response_id: Option<&str>, +) -> Result<(), ProbeFailure> { + let mut event = json!({ + "type": "response.create", + "model": model, + "store": false, + "generate": false, + "input": [], + "tools": [], + }); + if let Some(previous_response_id) = previous_response_id { + event["previous_response_id"] = Value::String(previous_response_id.to_string()); + } + socket + .send(WreqWsMessage::text(event.to_string())) + .await + .map_err(|_| ProbeFailure::Send) +} + +async fn receive_completed_response_id( + socket: &mut wreq::ws::WebSocket, + timeout: Duration, + observed_event_types: &mut Vec, +) -> Result { + let mut response_id = None; + for _ in 0..MAX_EVENTS_PER_TURN { + let message = tokio::time::timeout(timeout, socket.recv()) + .await + .map_err(|_| ProbeFailure::ReceiveTimeout)? + .ok_or(ProbeFailure::MissingResponseId)? + .map_err(|_| ProbeFailure::Receive)?; + match message { + WreqWsMessage::Text(text) => { + let event: Value = serde_json::from_str(text.as_str()) + .map_err(|_| ProbeFailure::UnexpectedFrame)?; + let event_type = event + .get("type") + .and_then(Value::as_str) + .map(safe_event_label) + .unwrap_or_else(|| "unknown".to_string()); + let is_remote_error = event_type == "error"; + let is_completed = event_type == "response.completed"; + observed_event_types.push(event_type); + if is_remote_error { + return Err(ProbeFailure::RemoteError); + } + if let Some(observed_response_id) = event + .pointer("/response/id") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + { + response_id = Some(observed_response_id.to_string()); + } + if is_completed { + return response_id.ok_or(ProbeFailure::MissingResponseId); + } + } + WreqWsMessage::Ping(_) | WreqWsMessage::Pong(_) => continue, + WreqWsMessage::Close(_) => return Err(ProbeFailure::MissingResponseId), + _ => return Err(ProbeFailure::UnexpectedFrame), + } + } + Err(ProbeFailure::MissingResponseId) +} + +fn safe_event_label(value: &str) -> String { + let trimmed = value.trim(); + if trimmed.is_empty() + || trimmed.len() > 80 + || !trimmed + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'_' | b'-')) + { + return "unknown".to_string(); + } + trimmed.to_string() +} + +fn print_report(report: &ProbeReport) { + match serde_json::to_string(report) { + Ok(json) => println!("{json}"), + Err(_) => println!("{{\"status\":\"failed\",\"error\":\"report_serialization_failed\"}}"), + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + use std::time::{Duration, Instant}; + + use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade}; + use axum::extract::State; + use axum::http::header::AUTHORIZATION; + use axum::http::{HeaderMap, HeaderValue}; + use axum::response::IntoResponse; + use axum::routing::get; + use axum::Router; + use futures_util::{SinkExt, StreamExt}; + use serde_json::Value; + use tokio::sync::{oneshot, Mutex}; + + use super::{parse_probe_url, run_probe, ProbeConfig}; + + #[derive(Default)] + struct MockState { + observed: Mutex>>, + } + + struct ObservedClientMessages { + authorization_present: bool, + profile_header_present: bool, + second_before_first_completion: bool, + first: Value, + second: Value, + } + + #[tokio::test] + async fn probe_confirms_sequential_response_continuation_without_exposing_values() { + let (url, observed, server) = spawn_mock_server().await; + let mut headers = HeaderMap::new(); + headers.insert( + AUTHORIZATION, + HeaderValue::from_static("Bearer test-token-that-must-not-be-reported"), + ); + headers.insert( + "x-aether-probe-profile", + HeaderValue::from_static("test-profile-id"), + ); + let config = ProbeConfig::new( + parse_probe_url(url.as_str()).expect("mock URL should be valid"), + "gpt-test".to_string(), + Duration::from_secs(2), + headers, + vec!["authorization", "x-aether-probe-profile"], + ); + + let report = run_probe(&config, Instant::now()) + .await + .expect("probe should complete against mock server"); + let client_messages = observed.await.expect("mock should observe client messages"); + server.abort(); + + assert_eq!(report.status, "passed"); + assert!(report.continuation_confirmed); + assert!(report + .observed_event_types + .contains(&"response.created".to_string())); + assert!(report + .observed_event_types + .contains(&"response.completed".to_string())); + assert!(client_messages.authorization_present); + assert!(client_messages.profile_header_present); + assert!(!client_messages.second_before_first_completion); + assert_eq!(client_messages.first["type"], "response.create"); + assert_eq!(client_messages.first["generate"], false); + assert_eq!(client_messages.first["store"], false); + assert_eq!(client_messages.second["previous_response_id"], "resp-first"); + let report_json = serde_json::to_string(&report).expect("report should serialize"); + assert!(!report_json.contains("test-token-that-must-not-be-reported")); + assert!(!report_json.contains("test-profile-id")); + assert!(!report_json.contains("resp-first")); + } + + #[test] + fn probe_url_rejects_credentials_and_query_strings() { + assert!(parse_probe_url("wss://example.test/v1/responses").is_ok()); + assert!(parse_probe_url("https://example.test/v1/responses").is_err()); + assert!(parse_probe_url("wss://token@example.test/v1/responses").is_err()); + assert!(parse_probe_url("wss://example.test/v1/responses?token=secret").is_err()); + } + + 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!("ws://{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 profile_header_present = headers.contains_key("x-aether-probe-profile"); + ws.on_upgrade(move |socket| async move { + serve_mock_socket(socket, state, authorization_present, profile_header_present).await; + }) + } + + async fn serve_mock_socket( + socket: WebSocket, + state: Arc, + authorization_present: bool, + profile_header_present: bool, + ) { + let (mut sender, mut receiver) = socket.split(); + let first = receive_json(&mut receiver).await; + let _ = sender + .send(Message::Text( + serde_json::json!({ + "type": "response.created", + "response": {"id": "resp-first"} + }) + .to_string() + .into(), + )) + .await; + let early_second = tokio::select! { + message = receiver.next() => Some(message), + _ = tokio::time::sleep(Duration::from_millis(50)) => None, + }; + let second_before_first_completion = early_second.is_some(); + let _ = sender + .send(Message::Text( + serde_json::json!({ + "type": "response.completed", + "response": {"id": "resp-first", "status": "completed"} + }) + .to_string() + .into(), + )) + .await; + let second = match early_second { + Some(Some(Ok(Message::Text(text)))) => { + serde_json::from_str(text.as_str()).expect("early client message should be JSON") + } + Some(Some(Ok(_))) => panic!("expected text continuation message"), + Some(Some(Err(error))) => panic!("client message should be valid: {error}"), + Some(None) => panic!("client closed before continuation"), + None => receive_json(&mut receiver).await, + }; + let _ = sender + .send(Message::Text( + serde_json::json!({ + "type": "response.created", + "response": {"id": "resp-second"} + }) + .to_string() + .into(), + )) + .await; + let _ = sender + .send(Message::Text( + serde_json::json!({ + "type": "response.completed", + "response": {"id": "resp-second", "status": "completed"} + }) + .to_string() + .into(), + )) + .await; + if let Some(observed) = state.observed.lock().await.take() { + let _ = observed.send(ObservedClientMessages { + authorization_present, + profile_header_present, + second_before_first_completion, + first, + second, + }); + } + } + + async fn receive_json(receiver: &mut futures_util::stream::SplitStream) -> Value { + let message = receiver + .next() + .await + .expect("client should send a message") + .expect("client message should be valid"); + let Message::Text(text) = message else { + panic!("expected text message"); + }; + serde_json::from_str(text.as_str()).expect("client message should be JSON") + } +} 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/attempt_lifecycle.rs b/apps/aether-gateway/src/execution_runtime/attempt_lifecycle.rs new file mode 100644 index 000000000..d85391863 --- /dev/null +++ b/apps/aether-gateway/src/execution_runtime/attempt_lifecycle.rs @@ -0,0 +1,1464 @@ +//! 一次 provider attempt 的记账生命周期,与 transport 无关。 +//! +//! HTTP 流式与 Responses WebSocket 的记账都是同一个三段结构: +//! `pending` → `started` → `terminal`。差异只在「终态事实从哪来」——HTTP 从 +//! SSE 字节流解析,WS 从协议事件观察。在这里抽出来之前,WS 侧在 +//! `handlers/proxy/websocket/responses/turn.rs` 里重写了一遍 usage 写入、 +//! candidate 状态流转、health/adaptive 效果投射、pool key lease 释放、 +//! body capture 和账单失败判定,与 HTTP 的顺序、超时语义只能靠人工对齐。 +//! +//! # HTTP 侧调用点映射 +//! +//! 本批不改 HTTP 执行路径(`execution_runtime/stream/execution.rs` 里的 +//! `DirectPassthroughFinalizerCore` 与 failover / oauth 重试 / prefetch 深度纠缠, +//! 无法在「行为等价 + 单 commit 可验证」的前提下接线)。这里记下逐调用点的对应 +//! 关系,作为后续 PR 的接线依据: +//! +//! | HTTP 现状调用点 | 对应本模块 | +//! |---|---| +//! | `record_stream_pending_lifecycle` | [`ExecutionAttemptLifecycle::begin`] | +//! | `maybe_record_first_stream_event_started` | [`ExecutionAttemptLifecycle::mark_started`] | +//! | `record_stream_terminal_usage` | `settle` 第 1 段:usage terminal | +//! | `enqueue_stream_candidate_status_update` | `settle` 第 2 段:candidate terminal | +//! | `apply_local_stream_{success,failure}_effects` | `settle` 第 3 段:provider effects | +//! | `submit_stream_report` | `settle` 第 4 段:execution report | +//! | `append_stream_capture_bytes` / `build_stream_body_capture` | [`AttemptBodyCapture`] | +//! +//! HTTP 侧的 `stage_trace` 埋点、kiro prompt-cache usage 合并、direct-inline 延迟 +//! pending 等是它独有的,接线时作为 transport 专有部分留在原处。 + +use std::collections::BTreeMap; +use std::future::Future; +use std::sync::Arc; +use std::time::Duration; + +use aether_contracts::{ExecutionPlan, ExecutionStreamTerminalSummary, ExecutionTelemetry}; +use aether_data_contracts::repository::candidates::RequestCandidateStatus; +use aether_data_contracts::repository::usage::UsageBodyCaptureState; +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 base64::Engine as _; +use serde_json::Value; +use tracing::warn; + +use crate::clock::current_unix_ms; +use crate::orchestration::{ + apply_local_stream_failure_effects, apply_local_stream_success_effects, + release_local_pool_key_lease, LocalExecutionEffectContext, LocalStreamFailureEffect, +}; +use crate::request_candidate_runtime::record_local_request_candidate_status; +use crate::usage::{submit_stream_report, GatewayStreamReportRequest}; +use crate::AppState; + +/// 客户端取消/断开时对外记录的状态码。 +pub(crate) const CLIENT_CANCELLED_STATUS_CODE: u16 = 499; + +/// 流式超时状态码;只有它会额外投射 pool stream timeout 效果。 +pub(crate) const STREAM_TIMEOUT_STATUS_CODE: u16 = 504; + +/// provider 侧观察到的终态。 +/// +/// 形状刻意保持 transport 中立:HTTP 流式与 WS turn 的差异只在事实从哪来。 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum AttemptProviderOutcome { + /// 观察到了供应商的终态事件。 + Terminal { + status_code: u16, + /// 供应商自己声明这一轮被取消(`response.cancelled`)。 + cancelled_by_provider: bool, + }, + /// 供应商没能给出终态:断链、超时、gateway 侧失败。 + Aborted { + status_code: u16, + reason: &'static str, + stream_timeout: bool, + }, +} + +impl AttemptProviderOutcome { + pub(crate) const fn status_code(self) -> u16 { + match self { + Self::Terminal { status_code, .. } | Self::Aborted { status_code, .. } => status_code, + } + } + + pub(crate) const fn cancelled_by_provider(self) -> bool { + matches!( + self, + Self::Terminal { + cancelled_by_provider: true, + .. + } + ) + } + + pub(crate) const fn stream_timeout(self) -> bool { + matches!( + self, + Self::Aborted { + stream_timeout: true, + .. + } + ) + } + + pub(crate) const fn is_terminal(self) -> bool { + matches!(self, Self::Terminal { .. }) + } +} + +/// 这一个 attempt 的内容是否完整交付给了客户端。 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum AttemptClientDelivery { + Complete, + Aborted { reason: &'static str }, +} + +impl AttemptClientDelivery { + pub(crate) const fn aborted_reason(self) -> Option<&'static str> { + match self { + Self::Complete => None, + Self::Aborted { reason } => Some(reason), + } + } + + pub(crate) const fn is_aborted(self) -> bool { + matches!(self, Self::Aborted { .. }) + } +} + +/// 一次 attempt 结算时的两个正交事实。 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct AttemptTerminalFacts { + pub(crate) provider: AttemptProviderOutcome, + pub(crate) delivery: AttemptClientDelivery, +} + +impl AttemptTerminalFacts { + /// 记入 usage / candidate / 效果的人类可读原因。 + pub(crate) const fn reason(self) -> &'static str { + if let Some(reason) = self.delivery.aborted_reason() { + return reason; + } + match self.provider { + AttemptProviderOutcome::Terminal { + cancelled_by_provider: true, + .. + } => "provider cancelled the response", + AttemptProviderOutcome::Terminal { .. } => { + "provider returned a terminal response event" + } + AttemptProviderOutcome::Aborted { reason, .. } => reason, + } + } + + /// 供应商侧强制错误原因:只有「供应商没给出终态、且内容已完整交付客户端」 + /// 才算,用于给终态摘要补 `parser_error`。 + /// + /// 客户端投递失败不是供应商的错误,所以那一侧返回 `None`——与现状 + /// `ResponsesWebSocketTurnOutcome::forced_error()` 对 `Cancelled` 返回 + /// `None` 一致。 + pub(crate) const fn forced_error(self) -> Option<&'static str> { + match (self.provider, self.delivery) { + (AttemptProviderOutcome::Aborted { reason, .. }, AttemptClientDelivery::Complete) => { + Some(reason) + } + _ => None, + } + } +} + +/// 这条 usage 记录是否计费。`Void` 等价于现状传给 +/// `record_stream_terminal(.., cancelled = true)` 的那一侧。 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum AttemptBilling { + Billed, + Void, +} + +impl AttemptBilling { + pub(crate) const fn is_void(self) -> bool { + matches!(self, Self::Void) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum AttemptCandidateStatus { + Success, + Failed, + Cancelled, +} + +/// candidate 行上记录的错误分类。 +/// +/// 与 [`AttemptCandidateStatus`] 刻意分开:`missing_terminal` 为真而记账层 +/// 判定不算失败(report kind 不要求观察到终态事件)时,现状会写出 +/// 「状态 Success + error_type=stream_missing_terminal_event」的组合, +/// 这里必须原样保留。 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum AttemptCandidateError { + None, + Cancelled, + /// 供应商已经给出终态,但这一轮内容没能完整交付给客户端。账单照记, + /// candidate 行上留下这条事实。 + ClientDeliveryFailed, + MissingTerminal, + TerminalError, +} + +/// 一次 attempt 结束后要投射给供应商/密钥池的效果。 +/// +/// 每个分支都会释放 pool key lease:`ProviderFailure` 由 `PoolError` 释放, +/// `ProviderSuccess` 由 `PoolSuccessStream` 释放,其余情况直接释放。少一条 +/// 分支就会把 lease 挂到 TTL 过期,等于短时间占死一把 key。 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum AttemptProviderEffect { + /// 既不投射成功也不投射失败,只把 lease 还回去。 + ReleasePoolKeyLease, + ProviderFailure, + ProviderSuccess, +} + +impl AttemptProviderEffect { + /// 把「每个分支都必须释放 lease」这条不变量显式化,便于测试锁住 + /// 「没进任何分支导致 lease 泄漏」这类回归。 + const fn releases_pool_key_lease(self) -> bool { + match self { + Self::ReleasePoolKeyLease | Self::ProviderFailure | Self::ProviderSuccess => true, + } + } +} + +/// 判定一次 attempt 结束后要投射的效果。 +/// +/// 关键分支是「记账层判成 failed,但这一轮没有投射供应商失败」:例如合法的 +/// `response.incomplete`(写满 max_output_tokens)。共享 usage 判定目前仍会 +/// 把这类终态记成失败,但供应商本身工作正常,既不该扣健康分,也不能因为落 +/// 不到任何分支而漏掉 lease 释放。 +pub(crate) const fn classify_attempt_provider_effect( + cancelled: bool, + projects_provider_failure: bool, + failed: bool, +) -> AttemptProviderEffect { + if cancelled { + AttemptProviderEffect::ReleasePoolKeyLease + } else if projects_provider_failure { + AttemptProviderEffect::ProviderFailure + } else if failed { + AttemptProviderEffect::ReleasePoolKeyLease + } else { + AttemptProviderEffect::ProviderSuccess + } +} + +/// 这一个 attempt 的账单是否作废。 +/// +/// 只有两种情况作废:供应商自己声明取消,或者供应商根本没给出终态而客户端 +/// 又已经走了。**供应商已经给出终态时,客户端最后一跳投递失败不作废账单**: +/// 供应商已经完成推理并消耗了 token,客户端还能用 `previous_response_id` +/// 续取这条响应,把成本记成 0 等于让上游账单凭空消失。 +pub(crate) const fn attempt_billing_is_void(facts: AttemptTerminalFacts) -> bool { + facts.provider.cancelled_by_provider() + || (facts.delivery.is_aborted() && !facts.provider.is_terminal()) +} + +/// attempt 对外记录的状态码。 +/// +/// 状态码现在纯粹是 provider 事实:客户端投递失败不再把一条已经拿到 200 +/// 终态的记录改写成 499。作废分支的 provider 状态码本身就是 499 +/// (`response.cancelled` 映射 499,`Cancelled` 信号的兜底也是 499), +/// 所以这些行的取值不变。 +pub(crate) const fn attempt_status_code(facts: AttemptTerminalFacts) -> u16 { + facts.provider.status_code() +} + +/// 结算判定的输入:两个正交事实 + 记账层对这条 report 的判定 + 终态摘要事实。 +#[derive(Debug, Clone, Copy)] +pub(crate) struct AttemptSettlementInputs { + pub(crate) facts: AttemptTerminalFacts, + /// `aether_usage_runtime::stream_report_represents_failure(payload)` 的结果。 + pub(crate) report_represents_failure: bool, + /// 终态摘要里是否观察到了 finish。 + pub(crate) observed_finish: bool, + /// 终态摘要里是否带解析错误。 + pub(crate) has_parser_error: bool, +} + +/// 结算动作。 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct AttemptSettlement { + pub(crate) status_code: u16, + pub(crate) billing: AttemptBilling, + pub(crate) candidate_status: AttemptCandidateStatus, + pub(crate) candidate_error: AttemptCandidateError, + pub(crate) provider_effect: AttemptProviderEffect, + pub(crate) submit_execution_report: bool, +} + +/// 由两个正交事实推出结算动作。唯一的判定入口,表驱动测试逐行锁死。 +/// +/// provider 终态已到达时,客户端投递失败只影响 candidate 的错误分类,不再作废 +/// 账单、不再把状态码改成 499、也不再跳过供应商效果和 execution report。 +pub(crate) const fn classify_attempt_settlement( + inputs: AttemptSettlementInputs, +) -> AttemptSettlement { + let AttemptSettlementInputs { + facts, + report_represents_failure, + observed_finish, + has_parser_error, + } = inputs; + + let void = attempt_billing_is_void(facts); + let status_code = attempt_status_code(facts); + let failed = !void && report_represents_failure; + let missing_terminal = !void && !observed_finish; + let projects_provider_failure = !void + && (status_code >= 400 + || facts.forced_error().is_some() + || has_parser_error + || missing_terminal); + + let candidate_status = if void { + AttemptCandidateStatus::Cancelled + } else if failed { + AttemptCandidateStatus::Failed + } else { + AttemptCandidateStatus::Success + }; + // 投递失败排在供应商侧分类之前:这条记录之所以特别,正是因为内容没送到 + // 客户端手上。供应商侧的判定仍然通过 candidate_status 和 error_message + // 保留下来。 + let candidate_error = if void { + AttemptCandidateError::Cancelled + } else if facts.delivery.is_aborted() { + AttemptCandidateError::ClientDeliveryFailed + } else if missing_terminal { + AttemptCandidateError::MissingTerminal + } else if failed { + AttemptCandidateError::TerminalError + } else { + AttemptCandidateError::None + }; + + AttemptSettlement { + status_code, + billing: if void { + AttemptBilling::Void + } else { + AttemptBilling::Billed + }, + candidate_status, + candidate_error, + provider_effect: classify_attempt_provider_effect(void, projects_provider_failure, failed), + submit_execution_report: !void, + } +} + +/// 每一段记账 I/O 的等待上界。 +/// +/// WS 用 `Bounded(5s)`:relay loop 是单任务,一段慢依赖会拖住整条连接的收发。 +/// HTTP 接线时用 `Unbounded` 即保持它现在的语义。 +#[derive(Debug, Clone, Copy)] +pub(crate) enum AttemptStageGuard { + Unbounded, + Bounded(Duration), +} + +impl AttemptStageGuard { + /// 等一段记账 I/O,超时就放弃等待。 + /// + /// 返回 `None` 表示这一段没有在上界内完成;调用方据此决定兜底动作 + /// (例如效果段超时后仍然要释放 pool key lease)。 + pub(crate) async fn await_stage( + self, + trace_id: &str, + stage: &'static str, + future: impl Future, + ) -> Option { + let Self::Bounded(timeout) = self else { + return Some(future.await); + }; + match tokio::time::timeout(timeout, future).await { + Ok(value) => Some(value), + Err(_) => { + warn!( + event_name = "execution_attempt_lifecycle_stage_timeout", + log_type = "ops", + trace_id, + stage, + timeout_ms = timeout.as_millis() as u64, + "gateway stopped waiting for an execution attempt lifecycle stage" + ); + None + } + } + } + + /// 跑一段不能丢的写入,同时仍然给调用方的等待设上界。 + /// + /// [`Self::await_stage`] 超时会 drop 掉它等待的 future,这对次要效果是对的, + /// 但会静默丢弃系统其余部分依赖的写入。先 spawn 再等,让上界只约束「等多久」: + /// 丢弃 `JoinHandle` 只是让任务脱离,它仍会跑完。 + pub(crate) async fn await_detachable_stage( + self, + trace_id: &str, + stage: &'static str, + write: F, + ) where + F: Future + Send + 'static, + { + let _ = self.await_stage(trace_id, stage, tokio::spawn(write)).await; + } +} + +/// 一侧(provider 或 client)的响应体捕获缓冲。 +/// +/// 捕获内容必须保持 SSE 形状(`data: {json}\n\n`): +/// `aether_usage_runtime` 会按 `data:` 行解析被捕获的 body 来判定 +/// `StreamCapturedTerminalState`,而它是 `stream_report_represents_failure` +/// 的一个 OR 项。换成结构化 JSON 会让终态判定恒为 Missing。 +#[derive(Debug, Default)] +pub(crate) struct AttemptBodyCapture { + buffer: Vec, + truncated: bool, +} + +impl AttemptBodyCapture { + pub(crate) fn append(&mut self, bytes: &[u8]) { + if bytes.is_empty() || self.truncated { + return; + } + let max_bytes = DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES; + if self.buffer.len() >= max_bytes { + self.truncated = true; + return; + } + let remaining = max_bytes - self.buffer.len(); + let copied = bytes.len().min(remaining); + self.buffer.extend_from_slice(&bytes[..copied]); + if copied < bytes.len() { + self.truncated = true; + } + } + + pub(crate) fn encode(&self) -> (Option, Option) { + let body = (!self.buffer.is_empty()) + .then(|| base64::engine::general_purpose::STANDARD.encode(&self.buffer)); + let state = if self.truncated { + UsageBodyCaptureState::Truncated + } else if self.buffer.is_empty() { + UsageBodyCaptureState::None + } else { + UsageBodyCaptureState::Inline + }; + (body, Some(state)) + } +} + +/// 启动一次 attempt 记账所需的种子。 +pub(crate) struct AttemptLifecycleSeed { + pub(crate) plan: ExecutionPlan, + pub(crate) report_kind: String, + pub(crate) report_context: Option, + pub(crate) stage_guard: AttemptStageGuard, +} + +/// 结算一次 attempt 需要的终态事实,由 transport 提供。 +pub(crate) struct AttemptTerminalFactsInput<'a> { + pub(crate) facts: AttemptTerminalFacts, + pub(crate) terminal_summary: ExecutionStreamTerminalSummary, + pub(crate) telemetry: ExecutionTelemetry, + pub(crate) provider_headers: BTreeMap, + pub(crate) provider_body: &'a AttemptBodyCapture, + pub(crate) client_body: &'a AttemptBodyCapture, + /// 供应商终态载荷原文(`error` / `response.failed` 一类),失败效果用它做 + /// failover 分类,优先级高于摘要里的 parser_error。 + pub(crate) provider_error_body: Option<&'a str>, + /// 人类可读的结算原因,写进 candidate 行。 + pub(crate) reason: &'a str, +} + +/// 一次 provider attempt 的记账生命周期:pending → started → terminal。 +pub(crate) struct ExecutionAttemptLifecycle { + plan: ExecutionPlan, + trace_id: String, + report_kind: String, + report_context: Option, + candidate_started_at_unix_ms: u64, + stage_guard: AttemptStageGuard, + started_recorded: bool, +} + +impl ExecutionAttemptLifecycle { + /// 写 Pending usage 行 + Pending candidate slot。 + /// + /// 两个写入都不设上界:此时还没有任何东西可以兜底,行没写成功就等于这次 + /// attempt 不存在。 + pub(crate) async fn begin(state: &AppState, seed: AttemptLifecycleSeed) -> Self { + let AttemptLifecycleSeed { + plan, + report_kind, + report_context, + stage_guard, + } = seed; + + let lifecycle_seed = build_lifecycle_usage_seed(&plan, report_context.as_ref()); + // Keep every transport on the same lifecycle data path. `AppState` can + // dedicate an isolated background database pool to usage writes; using + // the foreground state here bypasses that path and leaves the caller + // 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; + + Self { + trace_id: plan.request_id.clone(), + plan, + report_kind, + report_context, + candidate_started_at_unix_ms, + stage_guard, + started_recorded: false, + } + } + + pub(crate) fn plan(&self) -> &ExecutionPlan { + &self.plan + } + + pub(crate) fn trace_id(&self) -> &str { + self.trace_id.as_str() + } + + pub(crate) fn report_context(&self) -> Option<&Value> { + self.report_context.as_ref() + } + + pub(crate) fn take_report_context(&mut self) -> Option { + self.report_context.take() + } + + pub(crate) fn set_report_context(&mut self, report_context: Option) { + self.report_context = report_context; + } + + pub(crate) const fn stage_guard(&self) -> AttemptStageGuard { + self.stage_guard + } + + /// 首个可计费事件到达:usage stream_started + candidate Streaming。幂等。 + pub(crate) async fn mark_started( + &mut self, + state: &AppState, + status_code: u16, + telemetry: &ExecutionTelemetry, + ) { + if self.started_recorded { + return; + } + self.started_recorded = true; + let lifecycle_seed = build_lifecycle_usage_seed(&self.plan, self.report_context.as_ref()); + state.usage_runtime.record_stream_started( + state.usage_lifecycle_data_state().as_ref(), + &lifecycle_seed, + status_code, + Some(telemetry), + ); + let _ = self + .stage_guard + .await_stage( + self.trace_id.as_str(), + "candidate_stream_started", + record_local_request_candidate_status( + state, + &self.plan, + self.report_context.as_ref(), + SchedulerRequestCandidateStatusUpdate { + status: RequestCandidateStatus::Streaming, + status_code: Some(status_code), + 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; + } + + /// 终态四段,顺序不可重排: + /// + /// 1. usage terminal —— 这次 attempt 的账单记录。行是以 Pending 建立的, + /// 没有别的东西会去对账,所以用 detachable 保证不丢。 + /// 2. candidate terminal —— 调度侧的终态,慢依赖不能让它停在 Streaming。 + /// 3. provider 效果 —— health / adaptive / pool 反馈,次于前两段;超时后 + /// 仍然兜底释放 pool key lease,否则要等 lease TTL 过期才放出这把 key。 + /// 4. execution report —— 作废账单的分支不提交(与 HTTP 在下游断开后 + /// 同样不提交一致)。 + pub(crate) async fn settle( + mut self, + state: &AppState, + input: AttemptTerminalFactsInput<'_>, + ) -> AttemptSettlement { + let AttemptTerminalFactsInput { + facts, + terminal_summary, + telemetry, + provider_headers, + provider_body, + client_body, + provider_error_body, + reason, + } = input; + + let (provider_body_base64, provider_body_state) = provider_body.encode(); + let (client_body_base64, client_body_state) = client_body.encode(); + let payload = GatewayStreamReportRequest { + trace_id: self.trace_id.clone(), + report_kind: std::mem::take(&mut self.report_kind), + report_context: self.report_context.take(), + status_code: attempt_status_code(facts), + headers: provider_headers, + provider_body_base64, + provider_body_state, + client_body_base64, + client_body_state, + terminal_summary: Some(terminal_summary.clone()), + telemetry: Some(telemetry), + }; + let settlement = classify_attempt_settlement(AttemptSettlementInputs { + facts, + report_represents_failure: stream_report_represents_failure(&payload), + observed_finish: terminal_summary.observed_finish, + has_parser_error: terminal_summary.parser_error.is_some(), + }); + + // 1. usage terminal + 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); + let billing_void = settlement.billing.is_void(); + let usage_runtime = Arc::clone(&state.usage_runtime); + let usage_data = Arc::clone(state.usage_lifecycle_data_state()); + self.stage_guard + .await_detachable_stage(self.trace_id.as_str(), "usage_terminal", async move { + usage_runtime + .record_stream_terminal( + usage_data.as_ref(), + context_seed, + payload_seed, + billing_void, + ) + .await; + }) + .await; + + // 2. candidate terminal + let (error_type, error_message) = candidate_error_fields( + settlement.candidate_error, + terminal_summary.parser_error.as_deref(), + reason, + ); + let _ = self + .stage_guard + .await_stage( + self.trace_id.as_str(), + "candidate_terminal", + record_local_request_candidate_status( + state, + &self.plan, + payload.report_context.as_ref(), + SchedulerRequestCandidateStatusUpdate { + status: match settlement.candidate_status { + AttemptCandidateStatus::Cancelled => RequestCandidateStatus::Cancelled, + AttemptCandidateStatus::Failed => RequestCandidateStatus::Failed, + AttemptCandidateStatus::Success => RequestCandidateStatus::Success, + }, + status_code: Some(settlement.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; + + // 3. provider 效果 + let effect_context = LocalExecutionEffectContext { + plan: &self.plan, + report_context: payload.report_context.as_ref(), + }; + let effects_completed = self + .stage_guard + .await_stage(self.trace_id.as_str(), "provider_effects", async { + match settlement.provider_effect { + AttemptProviderEffect::ReleasePoolKeyLease => { + release_local_pool_key_lease(state, effect_context).await; + } + AttemptProviderEffect::ProviderFailure => { + let response_text = provider_error_body + .or(terminal_summary.parser_error.as_deref()) + .unwrap_or(reason); + let mut effect = LocalStreamFailureEffect::new( + settlement.status_code, + &payload.headers, + Some(response_text), + ); + if facts.provider.stream_timeout() { + effect = effect.with_stream_timeout(); + } + apply_local_stream_failure_effects(state, effect_context, effect).await; + } + AttemptProviderEffect::ProviderSuccess => { + apply_local_stream_success_effects(state, effect_context, &payload).await; + } + } + }) + .await + .is_some(); + if !effects_completed { + let _ = self + .stage_guard + .await_stage( + self.trace_id.as_str(), + "pool_lease_release_after_effect_timeout", + release_local_pool_key_lease(state, effect_context), + ) + .await; + } + + // 4. execution report + if settlement.submit_execution_report { + if let Some(Err(error)) = self + .stage_guard + .await_stage( + self.trace_id.as_str(), + "execution_report", + submit_stream_report(state, payload), + ) + .await + { + warn!( + event_name = "execution_attempt_report_submit_failed", + log_type = "ops", + trace_id = %self.trace_id, + error = ?error, + "gateway failed to submit an execution attempt terminal report" + ); + } + } + + settlement + } +} + +/// candidate 行上的 error_type / error_message。 +fn candidate_error_fields( + candidate_error: AttemptCandidateError, + parser_error: Option<&str>, + reason: &str, +) -> (Option, Option) { + match candidate_error { + AttemptCandidateError::Cancelled => ( + Some("websocket_cancelled".to_string()), + Some(reason.to_string()), + ), + AttemptCandidateError::ClientDeliveryFailed => ( + Some("client_delivery_failed".to_string()), + Some(reason.to_string()), + ), + AttemptCandidateError::MissingTerminal => ( + Some("stream_missing_terminal_event".to_string()), + Some(parser_error.map(str::to_string).unwrap_or_else(|| { + "upstream Responses WebSocket ended before a provider terminal event".to_string() + })), + ), + AttemptCandidateError::TerminalError => ( + Some("stream_terminal_error".to_string()), + parser_error + .map(str::to_string) + .or_else(|| Some(reason.to_string())), + ), + AttemptCandidateError::None => (None, None), + } +} + +#[cfg(test)] +mod tests { + use super::{ + classify_attempt_provider_effect, classify_attempt_settlement, AttemptBilling, + AttemptCandidateError, AttemptCandidateStatus, AttemptClientDelivery, + AttemptProviderEffect, AttemptProviderOutcome, AttemptSettlement, AttemptSettlementInputs, + AttemptTerminalFacts, + }; + + fn settle( + provider: AttemptProviderOutcome, + delivery: AttemptClientDelivery, + report_represents_failure: bool, + observed_finish: bool, + has_parser_error: bool, + ) -> AttemptSettlement { + classify_attempt_settlement(AttemptSettlementInputs { + facts: AttemptTerminalFacts { provider, delivery }, + report_represents_failure, + observed_finish, + has_parser_error, + }) + } + + const fn terminal(status_code: u16) -> AttemptProviderOutcome { + AttemptProviderOutcome::Terminal { + status_code, + cancelled_by_provider: false, + } + } + + const fn provider_cancelled() -> AttemptProviderOutcome { + AttemptProviderOutcome::Terminal { + status_code: 499, + cancelled_by_provider: true, + } + } + + const fn aborted(status_code: u16, reason: &'static str) -> AttemptProviderOutcome { + AttemptProviderOutcome::Aborted { + status_code, + reason, + stream_timeout: status_code == 504, + } + } + + /// 投递失败时 `forced_error` 必须为 `None`:客户端走了不是供应商的错误。 + /// 与现状 `ResponsesWebSocketTurnOutcome::forced_error()` 对 `Cancelled` + /// 返回 `None` 一致。 + #[test] + fn only_a_provider_abort_with_complete_delivery_is_a_forced_error() { + assert_eq!( + AttemptTerminalFacts { + provider: aborted(502, "upstream failed"), + delivery: AttemptClientDelivery::Complete, + } + .forced_error(), + Some("upstream failed") + ); + assert_eq!( + AttemptTerminalFacts { + provider: aborted(499, "client went away"), + delivery: AttemptClientDelivery::Aborted { + reason: "client went away" + }, + } + .forced_error(), + None + ); + assert_eq!( + AttemptTerminalFacts { + provider: terminal(200), + delivery: AttemptClientDelivery::Complete, + } + .forced_error(), + None + ); + } + + #[test] + fn the_recorded_reason_prefers_the_client_delivery_failure() { + assert_eq!( + AttemptTerminalFacts { + provider: terminal(200), + delivery: AttemptClientDelivery::Aborted { + reason: "client went away" + }, + } + .reason(), + "client went away" + ); + assert_eq!( + AttemptTerminalFacts { + provider: provider_cancelled(), + delivery: AttemptClientDelivery::Complete, + } + .reason(), + "provider cancelled the response" + ); + assert_eq!( + AttemptTerminalFacts { + provider: terminal(200), + delivery: AttemptClientDelivery::Complete, + } + .reason(), + "provider returned a terminal response event" + ); + assert_eq!( + AttemptTerminalFacts { + provider: aborted(502, "upstream failed"), + delivery: AttemptClientDelivery::Complete, + } + .reason(), + "upstream failed" + ); + } + + /// §1.6 结算表,逐行。 + #[test] + fn settlement_table_row_provider_cancelled_is_void_regardless_of_delivery() { + for delivery in [ + AttemptClientDelivery::Complete, + AttemptClientDelivery::Aborted { reason: "gone" }, + ] { + for report_represents_failure in [false, true] { + let settlement = settle( + provider_cancelled(), + delivery, + report_represents_failure, + true, + false, + ); + assert_eq!( + settlement, + AttemptSettlement { + status_code: 499, + billing: AttemptBilling::Void, + candidate_status: AttemptCandidateStatus::Cancelled, + candidate_error: AttemptCandidateError::Cancelled, + provider_effect: AttemptProviderEffect::ReleasePoolKeyLease, + submit_execution_report: false, + }, + "delivery={delivery:?} report_failure={report_represents_failure}" + ); + } + } + } + + #[test] + fn settlement_table_row_aborted_provider_with_aborted_delivery_is_void() { + let settlement = settle( + aborted(499, "client went away"), + AttemptClientDelivery::Aborted { + reason: "client went away", + }, + true, + false, + false, + ); + assert_eq!( + settlement, + AttemptSettlement { + status_code: 499, + billing: AttemptBilling::Void, + candidate_status: AttemptCandidateStatus::Cancelled, + candidate_error: AttemptCandidateError::Cancelled, + provider_effect: AttemptProviderEffect::ReleasePoolKeyLease, + submit_execution_report: false, + } + ); + } + + #[test] + fn settlement_table_row_clean_provider_terminal_is_a_billed_success() { + let settlement = settle( + terminal(200), + AttemptClientDelivery::Complete, + false, + true, + false, + ); + assert_eq!( + settlement, + AttemptSettlement { + status_code: 200, + billing: AttemptBilling::Billed, + candidate_status: AttemptCandidateStatus::Success, + candidate_error: AttemptCandidateError::None, + provider_effect: AttemptProviderEffect::ProviderSuccess, + submit_execution_report: true, + } + ); + } + + /// 合法 `response.incomplete`:记账层判失败,但供应商工作正常, + /// 不扣健康分、只释放 lease,并且账单照记。 + #[test] + fn settlement_table_row_legitimate_incomplete_is_billed_without_provider_failure() { + let settlement = settle( + terminal(200), + AttemptClientDelivery::Complete, + true, + true, + false, + ); + assert_eq!( + settlement, + AttemptSettlement { + status_code: 200, + billing: AttemptBilling::Billed, + candidate_status: AttemptCandidateStatus::Failed, + candidate_error: AttemptCandidateError::TerminalError, + provider_effect: AttemptProviderEffect::ReleasePoolKeyLease, + submit_execution_report: true, + } + ); + } + + #[test] + fn settlement_table_row_provider_abort_projects_a_provider_failure() { + let settlement = settle( + aborted( + 502, + "upstream WebSocket closed before provider terminal event", + ), + AttemptClientDelivery::Complete, + true, + false, + false, + ); + assert_eq!( + settlement, + AttemptSettlement { + status_code: 502, + billing: AttemptBilling::Billed, + candidate_status: AttemptCandidateStatus::Failed, + candidate_error: AttemptCandidateError::MissingTerminal, + provider_effect: AttemptProviderEffect::ProviderFailure, + submit_execution_report: true, + } + ); + } + + /// ✱ 修正后的那一行:provider 终态已到达,客户端投递失败不再作废账单。 + /// + /// 供应商已经完成推理并消耗 token,客户端还能用 `previous_response_id` + /// 续取这条响应;把成本记成 0 等于让上游账单凭空消失。投递失败作为独立 + /// 事实留在 candidate 的错误分类里。 + #[test] + fn settlement_table_row_client_delivery_failure_keeps_a_reached_terminal_billed() { + let settlement = settle( + terminal(200), + AttemptClientDelivery::Aborted { + reason: "gateway could not relay the provider event to the client", + }, + false, + true, + false, + ); + assert_eq!( + settlement, + AttemptSettlement { + status_code: 200, + billing: AttemptBilling::Billed, + candidate_status: AttemptCandidateStatus::Success, + candidate_error: AttemptCandidateError::ClientDeliveryFailed, + provider_effect: AttemptProviderEffect::ProviderSuccess, + submit_execution_report: true, + } + ); + + // 除了 candidate 的错误分类,其余判定与「投递成功」完全一致。 + let delivered = settle( + terminal(200), + AttemptClientDelivery::Complete, + false, + true, + false, + ); + assert_eq!(settlement.status_code, delivered.status_code); + assert_eq!(settlement.billing, delivered.billing); + assert_eq!(settlement.candidate_status, delivered.candidate_status); + assert_eq!(settlement.provider_effect, delivered.provider_effect); + assert_eq!( + settlement.submit_execution_report, + delivered.submit_execution_report + ); + assert_ne!(settlement.candidate_error, delivered.candidate_error); + } + + /// 供应商还没给出终态时,客户端投递失败仍然作废账单:这一轮确实没有产出。 + #[test] + fn a_delivery_failure_without_a_provider_terminal_still_voids_the_bill() { + let settlement = settle( + aborted(499, "client went away"), + AttemptClientDelivery::Aborted { + reason: "client went away", + }, + false, + false, + false, + ); + assert_eq!(settlement.status_code, 499); + assert_eq!(settlement.billing, AttemptBilling::Void); + assert_eq!( + settlement.candidate_status, + AttemptCandidateStatus::Cancelled + ); + assert_eq!(settlement.candidate_error, AttemptCandidateError::Cancelled); + assert!(!settlement.submit_execution_report); + } + + /// 供应商自己声明取消时,即使内容送到了客户端也不计费。 + #[test] + fn a_provider_declared_cancellation_is_void_even_when_delivered() { + let settlement = settle( + provider_cancelled(), + AttemptClientDelivery::Complete, + false, + true, + false, + ); + assert_eq!(settlement.billing, AttemptBilling::Void); + assert_eq!(settlement.candidate_error, AttemptCandidateError::Cancelled); + } + + /// 记账层判 Success,但摘要没观察到 finish:现状会写出 + /// 「candidate=Success + error_type=stream_missing_terminal_event」, + /// 所以状态与错误分类必须各自独立。 + #[test] + fn a_missing_terminal_can_coexist_with_a_successful_candidate_status() { + let settlement = settle( + terminal(200), + AttemptClientDelivery::Complete, + false, + false, + false, + ); + assert_eq!(settlement.candidate_status, AttemptCandidateStatus::Success); + assert_eq!( + settlement.candidate_error, + AttemptCandidateError::MissingTerminal + ); + // missing_terminal 仍然要投射供应商失败。 + assert_eq!( + settlement.provider_effect, + AttemptProviderEffect::ProviderFailure + ); + } + + #[test] + fn a_parser_error_projects_a_provider_failure_even_on_a_clean_status_code() { + let settlement = settle( + terminal(200), + AttemptClientDelivery::Complete, + true, + true, + true, + ); + assert_eq!( + settlement.provider_effect, + AttemptProviderEffect::ProviderFailure + ); + assert_eq!(settlement.billing, AttemptBilling::Billed); + } + + #[test] + fn a_legitimate_incomplete_still_releases_the_pool_key_lease() { + // 共享 usage 判定目前仍把 response.incomplete 记成终态失败,于是会出现 + // failed=true 而 projects_provider_failure=false 的组合。这种组合必须 + // 明确落到「只释放 lease」的分支,否则 lease 会挂到 TTL 过期。 + let effect = classify_attempt_provider_effect(false, false, true); + + assert_eq!(effect, AttemptProviderEffect::ReleasePoolKeyLease); + assert!(effect.releases_pool_key_lease()); + } + + #[test] + fn every_provider_effect_releases_the_pool_key_lease() { + for (cancelled, projects_provider_failure, failed, expected) in [ + ( + true, + false, + false, + AttemptProviderEffect::ReleasePoolKeyLease, + ), + (true, true, true, AttemptProviderEffect::ReleasePoolKeyLease), + (false, true, true, AttemptProviderEffect::ProviderFailure), + ( + false, + false, + true, + AttemptProviderEffect::ReleasePoolKeyLease, + ), + (false, false, false, AttemptProviderEffect::ProviderSuccess), + ] { + let effect = + classify_attempt_provider_effect(cancelled, projects_provider_failure, failed); + assert_eq!( + effect, expected, + "cancelled={cancelled} projects_provider_failure={projects_provider_failure} failed={failed}" + ); + assert!( + effect.releases_pool_key_lease(), + "every effect branch must release the pool key lease" + ); + } + } + + /// 每一个结算分支都必须释放 lease:这条不变量跨越整张结算表。 + #[test] + fn every_settlement_branch_releases_the_pool_key_lease() { + let providers = [ + terminal(200), + terminal(429), + provider_cancelled(), + aborted(502, "upstream failed"), + aborted(504, "timed out"), + ]; + let deliveries = [ + AttemptClientDelivery::Complete, + AttemptClientDelivery::Aborted { reason: "gone" }, + ]; + for provider in providers { + for delivery in deliveries { + for report_represents_failure in [false, true] { + for observed_finish in [false, true] { + for has_parser_error in [false, true] { + let settlement = settle( + provider, + delivery, + report_represents_failure, + observed_finish, + has_parser_error, + ); + assert!( + settlement.provider_effect.releases_pool_key_lease(), + "provider={provider:?} delivery={delivery:?}" + ); + // 作废账单的分支一律不提交 execution report。 + assert_eq!( + settlement.submit_execution_report, + !settlement.billing.is_void(), + "provider={provider:?} delivery={delivery:?}" + ); + } + } + } + } + } + } +} + +#[cfg(test)] +mod stage_tests { + use std::sync::atomic::{AtomicBool, Ordering}; + use std::sync::{Arc, Mutex}; + use std::time::Duration; + + use base64::Engine as _; + + use super::{ + candidate_error_fields, AttemptBodyCapture, AttemptCandidateError, AttemptStageGuard, + }; + + /// 效果段超时后仍然必须释放 pool key lease,否则那把 key 要等 lease TTL + /// 过期才放出来。这里用 `await_stage` 返回 `None` 驱动兜底分支。 + #[tokio::test] + async fn a_timed_out_effect_stage_still_reaches_the_lease_release_fallback() { + let guard = AttemptStageGuard::Bounded(Duration::from_millis(20)); + let lease_released = Arc::new(AtomicBool::new(false)); + + // 第一段:永不完成的效果投射。 + let effects_completed = guard + .await_stage("trace", "provider_effects", std::future::pending::<()>()) + .await + .is_some(); + assert!( + !effects_completed, + "a stage that never completes must not report success" + ); + + // 生产代码据此走兜底释放。 + if !effects_completed { + let released = Arc::clone(&lease_released); + let _ = guard + .await_stage( + "trace", + "pool_lease_release_after_effect_timeout", + async move { + released.store(true, Ordering::SeqCst); + }, + ) + .await; + } + assert!( + lease_released.load(Ordering::SeqCst), + "the lease must still be released after an effect-stage timeout" + ); + } + + #[tokio::test] + async fn an_unbounded_stage_guard_waits_for_the_stage() { + let guard = AttemptStageGuard::Unbounded; + // 无上界:即使这一段比任何 Bounded 上界都久,也必须等到它完成。 + let value = guard + .await_stage("trace", "stage", async { + tokio::time::sleep(Duration::from_millis(120)).await; + 7_u8 + }) + .await; + assert_eq!(value, Some(7)); + } + + /// 不能丢的写入即使调用方停止等待也要跑完:`await_detachable_stage` 先 spawn + /// 再等,上界只约束「等多久」。 + #[tokio::test] + async fn a_detachable_stage_completes_even_after_the_caller_stops_waiting() { + let guard = AttemptStageGuard::Bounded(Duration::from_millis(20)); + let written = Arc::new(AtomicBool::new(false)); + let flag = Arc::clone(&written); + + guard + .await_detachable_stage("trace", "usage_terminal", async move { + tokio::time::sleep(Duration::from_millis(120)).await; + flag.store(true, Ordering::SeqCst); + }) + .await; + assert!( + !written.load(Ordering::SeqCst), + "the caller must stop waiting at its bound" + ); + + tokio::time::sleep(Duration::from_millis(300)).await; + assert!( + written.load(Ordering::SeqCst), + "a detached write must still run to completion" + ); + } + + /// settle 的四段顺序不可重排:账单先落地,然后 candidate 终态,然后次要的 + /// 供应商效果,最后才是 execution report。用计数器替身记录实际顺序。 + #[tokio::test] + async fn the_settle_stages_run_in_a_fixed_order() { + let guard = AttemptStageGuard::Bounded(Duration::from_millis(500)); + let order = Arc::new(Mutex::new(Vec::new())); + + for stage in [ + "usage_terminal", + "candidate_terminal", + "provider_effects", + "execution_report", + ] { + let recorder = Arc::clone(&order); + let _ = guard + .await_stage("trace", stage, async move { + recorder.lock().expect("order lock").push(stage); + }) + .await; + } + + assert_eq!( + order.lock().expect("order lock").as_slice(), + [ + "usage_terminal", + "candidate_terminal", + "provider_effects", + "execution_report", + ] + ); + } + + /// body capture 的编码状态。截断分支这里到不了:共享的 + /// `DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES` 是 `usize::MAX`, + /// 也就是默认不限长;截断只在 usage 侧把上限调低后才可能发生。 + #[test] + fn body_capture_encodes_inline_and_empty_states() { + assert_eq!( + super::DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES, + usize::MAX, + "the default capture limit is unbounded; truncation is not reachable here" + ); + + let mut capture = AttemptBodyCapture::default(); + capture.append(b"data: {\"type\":\"response.created\"}\n\n"); + capture.append(b"data: {\"type\":\"response.completed\"}\n\n"); + let (body, state) = capture.encode(); + let decoded = base64::engine::general_purpose::STANDARD + .decode(body.expect("a non-empty capture is encoded")) + .expect("capture is valid base64"); + let decoded = String::from_utf8(decoded).expect("capture is UTF-8"); + assert!( + decoded.starts_with("data: ") && decoded.ends_with("\n\n"), + "the capture must stay SSE-shaped for the usage runtime: {decoded:?}" + ); + assert_eq!(decoded.matches("data: ").count(), 2, "appends concatenate"); + assert_eq!( + state, + Some(aether_data_contracts::repository::usage::UsageBodyCaptureState::Inline) + ); + + let empty = AttemptBodyCapture::default(); + let (body, state) = empty.encode(); + assert!(body.is_none()); + assert_eq!( + state, + Some(aether_data_contracts::repository::usage::UsageBodyCaptureState::None) + ); + } + + /// candidate 行的 error_type 映射:投递失败与供应商侧失败必须各有名字。 + #[test] + fn candidate_error_fields_name_each_failure_kind() { + assert_eq!( + candidate_error_fields(AttemptCandidateError::None, None, "reason"), + (None, None) + ); + assert_eq!( + candidate_error_fields(AttemptCandidateError::Cancelled, None, "gone").0, + Some("websocket_cancelled".to_string()) + ); + assert_eq!( + candidate_error_fields( + AttemptCandidateError::ClientDeliveryFailed, + None, + "write failed" + ), + ( + Some("client_delivery_failed".to_string()), + Some("write failed".to_string()) + ) + ); + assert_eq!( + candidate_error_fields( + AttemptCandidateError::TerminalError, + Some("parser"), + "reason" + ), + ( + Some("stream_terminal_error".to_string()), + Some("parser".to_string()) + ) + ); + assert_eq!( + candidate_error_fields(AttemptCandidateError::MissingTerminal, None, "reason").0, + Some("stream_missing_terminal_event".to_string()) + ); + } +} diff --git a/apps/aether-gateway/src/execution_runtime/mod.rs b/apps/aether-gateway/src/execution_runtime/mod.rs index 260660cd9..89fbf1653 100644 --- a/apps/aether-gateway/src/execution_runtime/mod.rs +++ b/apps/aether-gateway/src/execution_runtime/mod.rs @@ -3,6 +3,8 @@ use std::collections::BTreeMap; use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; +pub(crate) mod admission; +pub(crate) mod attempt_lifecycle; mod chatgpt_web_image; mod constants; mod fallback; @@ -23,6 +25,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 86cdf14be..005da96f6 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 @@ -1304,6 +1304,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 83b86e14f..849c8db91 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/shared/payloads.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/shared/payloads.rs @@ -151,6 +151,8 @@ pub(crate) struct AdminProviderCreateRequest { #[serde(default)] pub(crate) keep_priority_on_conversion: Option, #[serde(default)] + pub(crate) responses_websocket_enabled: Option, + #[serde(default)] pub(crate) is_active: Option, #[serde(default)] pub(crate) concurrent_limit: Option, @@ -210,6 +212,8 @@ pub(crate) struct AdminProviderUpdateRequest { #[serde(default)] pub(crate) keep_priority_on_conversion: 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 86ffe5f5a..12456bc48 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, @@ -218,6 +218,7 @@ pub(crate) fn build_admin_provider_summary_value( "ops_architecture_id": ops_architecture_id, "kiro_simulated_cache_enabled": kiro_simulated_cache_enabled, "codex_cyber_flag_passthrough_enabled": codex_cyber_flag_passthrough_enabled(&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 bc3ba3d4b..ba0f9bb37 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; @@ -157,6 +159,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 15b23ea05..2654ea32a 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; @@ -311,6 +313,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..6abfb51d6 --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/client.rs @@ -0,0 +1,816 @@ +//! 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, ActiveProviderAttempt, +}; +use super::quota::{mark_active_response_retry_unsafe, send_previous_response_not_found}; +use super::redaction::redact_responses_websocket_client_event; +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::BoundResponsesConnection; +use super::turn::{ + begin_responses_websocket_turn, prepare_responses_websocket_turn_decision, + ResponsesWebSocketTurnObservation, ResponsesWebSocketTurnOutcome, +}; +use super::turn_state::LogicalTurn; +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.turn_state.accepts_new_response_create() { + 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; + }; + // 这一轮的 planning Parts 只构造一次(它携带 per-turn 的 + // RedactionSessionSlot),并且客户端事件也只在这里脱敏一次: + // 复用已绑定 upstream 的 continuation 根本不进 planner,只靠 planner + // 内部脱敏拦不住它。之后 re-plan / continuation / 配额重试都只看脱敏 + // 后的事件,上游请求体与审计 original_request_body 因此一致。 + let planning_parts = build_planning_parts(context); + let redacted_client_event = redact_responses_websocket_client_event( + state, + &planning_parts, + &context.decision, + &client_event, + ) + .await; + let client_event = match redacted_client_event { + Ok(Some(redaction)) => { + // 这一轮的映射登记到连接上,响应帧才能在最后一跳还原回真实值。 + bound.redaction_restorer.register(redaction.session); + redaction.client_event + } + Ok(None) => client_event, + Err(error) => { + warn!( + event_name = "responses_websocket_followup_redaction_failed", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + error = ?error, + "gateway could not apply chat PII redaction to a Responses WebSocket turn" + ); + send_gateway_error_with_status( + client_socket, + 500, + "responses_websocket_redaction_unavailable", + "Gateway could not apply the configured PII redaction", + ) + .await; + close_client_socket( + client_socket, + CLOSE_INTERNAL_ERROR, + "responses_websocket_redaction_unavailable", + ) + .await; + return RelayDisposition::Close; + } + }; + 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, + &planning_parts, + 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, + &planning_parts, + client_event, + requested_model, + ) + .await; + } + if !response_create_has_previous_response_id(&client_event) { + return forward_replanned_response_create( + bound, + client_socket, + state, + context, + &planning_parts, + 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 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.turn_state.begin( + LogicalTurn::new(client_event.clone(), turn_index, logical_turn_id), + ActiveProviderAttempt::new(state, turn), + ); + bound.next_turn_index = bound.next_turn_index.saturating_add(1); + + 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.turn_state.attempt_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 + } + } +} + +/// 重新规划一轮 `response.create`(换模型或独立轮)。 +/// +/// `planning_parts` 与 `client_event` 都由调用方准备:事件已经过请求侧脱敏, +/// Parts 携带这一轮的 `RedactionSessionSlot`,所以 planner 里的候选级脱敏对 +/// 已脱敏内容是幂等的 no-op,上游请求体与审计 body 都保持脱敏态。 +async fn forward_replanned_response_create( + bound: &mut BoundResponsesConnection, + client_socket: &mut WebSocket, + state: &AppState, + context: &WebSocketRequestContext, + planning_parts: &http::request::Parts, + client_event: Value, + requested_model: String, +) -> RelayDisposition { + 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.turn_state.begin( + LogicalTurn::new(client_event.clone(), turn_index, logical_turn_id.clone()), + ActiveProviderAttempt::new(state, turn), + ); + bound.next_turn_index = bound.next_turn_index.saturating_add(1); + 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.decision_template = replacement.decision_template; + bound.body_normalization = replacement.body_normalization; + bound.binding_identity = replacement.binding_identity; + bound.turn_state.begin( + LogicalTurn::new(client_event, turn_index, logical_turn_id), + ActiveProviderAttempt::new(state, turn), + ); + 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..8900ccdb2 --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/connection.rs @@ -0,0 +1,572 @@ +//! Connection-level Responses WebSocket FSM. + +use std::time::Duration; + +use axum::extract::ws::{Message as AxumWsMessage, WebSocket}; +use futures_util::{SinkExt, StreamExt}; +use serde_json::Value; +use tokio::time::sleep; +use wreq::ws::message::Message as WreqWsMessage; + +use super::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, + settle_turn_finalization, ActiveProviderAttempt, PreviousAttemptSettled, +}; +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, fatal_relay_policy, FatalRelaySignal, QuotaRelayAction, QuotaRelayFacts, +}; +use super::settlement::settle_signal_for_client_delivery_failure; +use super::state::BoundResponsesConnection; +use super::turn::{ + ResponsesProviderAttempt, 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"; + +/// 写客户端 socket 失败时记录的投递失败原因。刻意不说「客户端在终态前断开」: +/// 供应商的终态可能已经到达,只是最后一跳没送出去。 +const CLIENT_DELIVERY_FAILED_REASON: &str = + "gateway could not relay the provider event to the client"; + +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.turn_state.attempt().map(|turn| turn.deadline()); + tokio::select! { + _ = &mut connection_deadline => { + finalize_active_turn( + bound, + state, + ResponsesWebSocketTurnOutcome::connection_limit_reached(), + ).await; + send_gateway_error_with_status( + client_socket, + 503, + "websocket_connection_limit_reached", + "WebSocket connection duration limit reached; reconnect to continue", + ).await; + close_bound_upstream(bound).await; + close_client_socket(client_socket, CLOSE_TRY_AGAIN, "connection_limit_reached").await; + break; + } + _ = wait_for_optional_deadline(active_turn_deadline.map(|deadline| deadline.deadline)) => { + let Some(turn_deadline) = active_turn_deadline else { + continue; + }; + 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; + 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; + 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; + 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; + 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; + 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.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.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.turn_state.response_in_flight(), + "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 + .turn_state + .attempt_mut() + .and_then(|turn| turn.observe_upstream_frame(frame, adapter)), + None => { + if let Some(turn) = bound.turn_state.attempt_mut() { + turn.observe_invalid_upstream_text(text.as_str()) + } + else { + None + } + } + } + } + _ => None, + }; + if matches!( + observation, + Some(ResponsesWebSocketTurnObservation::Started) + | Some(ResponsesWebSocketTurnObservation::Terminal(_)) + ) { + if let Some(turn) = bound.turn_state.attempt_mut() { + turn.mark_stream_started(state).await; + } + } + let terminal_outcome = match observation { + Some(ResponsesWebSocketTurnObservation::Terminal(outcome)) => Some(outcome), + _ => None, + }; + 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; + 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.turn_state.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) { + // detach_attempt 保留 logical turn:重试是同一轮请求的下一个 attempt。 + let retry_turn = bound + .turn_state + .detach_attempt() + .map(ActiveProviderAttempt::disarm); + // 先结算旧 attempt 并等它落地,再规划下一个 attempt。两个理由: + // + // 1. 规划要读 health / adaptive / pool 状态,而这些正是旧 + // attempt 结算时才投射的。普通的新 turn 早就在 client.rs 里 + // 用 await_pending_turn_finalization 挡住了「基于陈旧状态 + // 规划」,透明重试这条路径原先漏了这一步。 + // 2. 旧 attempt 还占着自己的 pool key lease。不先释放,重试就 + // 可能因为「这把 key 仍被占用」而挑不到本该可用的替代 key, + // 或者干脆判成无可用供应商。 + let settled = match retry_turn { + Some(mut turn) => { + turn.release_admission().await; + settle_turn_finalization( + bound, + state, + turn, + terminal_outcome.unwrap_or_else( + ResponsesWebSocketTurnOutcome::upstream_closed, + ), + ) + .await + } + None => PreviousAttemptSettled::nothing_to_settle(), + }; + if retry_active_turn_after_quota_exhaustion(bound, state, context, settled).await + { + continue; + } + // 重试失败。旧 attempt 已经结算,logical turn 仍停在 + // Replanning,所以后面分支里的 end() / finalize_active_turn + // 只会清掉 logical turn 而不会交出 attempt——不存在重复结算。 + 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.turn_state.end().map(ActiveProviderAttempt::disarm); + if let Some(active_turn) = turn.as_mut() { + active_turn.release_admission().await; + } + send_previous_response_not_found(client_socket).await; + if let Some(turn) = turn { + queue_turn_finalization( + bound, + state, + turn, + terminal_outcome.unwrap_or_else( + ResponsesWebSocketTurnOutcome::upstream_closed, + ), + ) + .await; + } + detach_exhausted_upstream(bound, directive, &context.trace_id).await; + continue; + } + if matches!(quota_relay_action, QuotaRelayAction::ForwardQuotaAndDetach) { + let 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; + detach_exhausted_upstream(bound, directive, &context.trace_id).await; + continue; + } + // 响应侧还原:HTTP 在把响应体交给客户端之前会把占位符换回真实值 + // (`privacy::restore_sync_response_body` / + // `privacy::StreamingResponseRestorer`),这里是 WS 的同一个位置 + // ——最后一跳之前,并且在 `capture_client_frame` 之前,所以审计与 + // 终态观测继续消费脱敏态的事件。没有命中还原时保持上游原字节。 + let restored_client_frame = parsed_upstream_frame + .as_ref() + .map(ParsedResponsesWebSocketFrame::event) + .and_then(|event| { + bound.redaction_restorer.restore_provider_frame_text(event) + }); + let client_frame = match restored_client_frame { + Some(restored) => AxumWsMessage::Text(restored.into()), + None => upstream_message_to_client(upstream_message.clone()), + }; + if let Err(error) = send_client_message(client_socket, client_frame).await { + 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(), + provider_terminal_reached = terminal_outcome.is_some(), + "gateway could not relay a provider event to the client" + ); + // 投递失败是独立事实,不能覆盖已经到达的 provider 终态: + // 供应商已经完成推理并消耗 token,账单按它的终态计。 + bound + .turn_state + .record_client_delivery_aborted(CLIENT_DELIVERY_FAILED_REASON); + finalize_active_turn( + bound, + state, + settle_signal_for_client_delivery_failure(terminal_outcome), + ).await; + close_bound_upstream(bound).await; + break; + } + if let (Some(turn), Some(frame)) = + (bound.turn_state.attempt_mut(), parsed_upstream_frame.as_ref()) + { + turn.capture_client_frame(frame.event()); + } + if let Some(outcome) = terminal_outcome { + finalize_active_turn(bound, state, outcome).await; + } else if is_close { + 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"); + 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; + } + } +} 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..d756bd652 --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/frame.rs @@ -0,0 +1,515 @@ +//! 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" + ) +} + +/// `response.incomplete` 的合法终态 reason 白名单。 +/// +/// 这些 reason 表示上游按规则正常结束了本轮响应(写满 `max_output_tokens`、 +/// 命中内容过滤、按工具调用截断),标准流解析里它们会变成 `length` / +/// `content_filter` / `tool_calls` 这类正常 finish,和 +/// `openai_responses_incomplete_finish_reason` 的既有映射保持一致,因此不能 +/// 当成 provider failure 记账。 +const LEGITIMATE_RESPONSES_INCOMPLETE_REASONS: [&str; 5] = [ + "max_output_tokens", + "max_tokens", + "content_filter", + "tool_calls", + "function_call", +]; + +/// 读取 `response.incomplete` 携带的 `incomplete_details.reason`。 +/// +/// 标准位置是 `response.incomplete_details.reason`;批量封装偶尔把 +/// `incomplete_details` 直接放在事件顶层,两处都要看,否则合法终态会被漏判。 +fn responses_incomplete_reason(event: &Value) -> Option<&str> { + [ + event.pointer("/response/incomplete_details/reason"), + event.pointer("/incomplete_details/reason"), + ] + .into_iter() + .flatten() + .filter_map(Value::as_str) + .map(str::trim) + .find(|reason| !reason.is_empty()) +} + +/// 判断一个 `response.incomplete` 是否是合法终态。 +/// +/// reason 缺失或不在白名单内(例如 `error`、`server_error`)时继续按 +/// provider failure 处理:这类 incomplete 说明上游确实没能正常收尾,仍应扣 +/// 供应商健康分。 +fn responses_incomplete_is_legitimate_terminal(event: &Value) -> bool { + responses_incomplete_reason(event).is_some_and(|reason| { + LEGITIMATE_RESPONSES_INCOMPLETE_REASONS + .iter() + .any(|candidate| reason.eq_ignore_ascii_case(candidate)) + }) +} + +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, + }), + // 合法 incomplete(例如写满 max_output_tokens)是正常终态,默认按 200 + // 记账,不再一律当 502 provider failure;reason 缺失或未知时保留原来的 + // 502 默认值。显式 `status_code` 和 error code 映射仍然优先于默认值, + // 所以带 `rate_limit_exceeded` 的 incomplete 依旧是 429。 + "response.incomplete" => Some(ResponsesWebSocketFrameTerminal { + status_code: websocket_event_status_code( + event, + if responses_incomplete_is_legitimate_terminal(event) { + 200 + } else { + 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 a_legitimate_incomplete_is_a_terminal_but_not_a_provider_failure() { + for reason in [ + "max_output_tokens", + "max_tokens", + "content_filter", + "tool_calls", + "function_call", + "MAX_OUTPUT_TOKENS", + ] { + let raw = format!( + r#"{{"type":"response.incomplete","response":{{"status":"incomplete","incomplete_details":{{"reason":"{reason}"}}}}}}"# + ); + let frame = ParsedResponsesWebSocketFrame::parse(&raw).expect("valid frame"); + + assert!(frame.is_terminal(), "{reason} should end the turn"); + assert_eq!( + frame + .terminal() + .map(|terminal| (terminal.status_code, terminal.cancelled)), + Some((200, false)), + "{reason} is a legitimate terminal result, not a 502 provider failure" + ); + } + } + + #[test] + fn a_top_level_incomplete_details_reason_is_also_honored() { + let frame = ParsedResponsesWebSocketFrame::parse( + r#"{"type":"response.incomplete","incomplete_details":{"reason":"max_output_tokens"}}"#, + ) + .expect("valid frame"); + + assert_eq!(frame.status(), Some(200)); + } + + #[test] + fn an_incomplete_without_a_legitimate_reason_stays_a_provider_failure() { + for raw in [ + r#"{"type":"response.incomplete"}"#, + r#"{"type":"response.incomplete","response":{"incomplete_details":null}}"#, + r#"{"type":"response.incomplete","response":{"incomplete_details":{"reason":""}}}"#, + r#"{"type":"response.incomplete","response":{"incomplete_details":{"reason":"error"}}}"#, + r#"{"type":"response.incomplete","response":{"incomplete_details":{"reason":"server_error"}}}"#, + ] { + let frame = ParsedResponsesWebSocketFrame::parse(raw).expect("valid frame"); + + assert_eq!( + frame.status(), + Some(502), + "an incomplete without a known-good reason must stay a provider failure: {raw}" + ); + } + } + + #[test] + fn a_legitimate_incomplete_still_respects_an_explicit_provider_status() { + let explicit = ParsedResponsesWebSocketFrame::parse( + r#"{"type":"response.incomplete","status_code":503,"response":{"incomplete_details":{"reason":"max_output_tokens"}}}"#, + ) + .expect("valid frame"); + assert_eq!(explicit.status(), Some(503)); + + let quota = ParsedResponsesWebSocketFrame::parse( + r#"{"type":"response.incomplete","response":{"error":{"code":"rate_limit_exceeded"},"incomplete_details":{"reason":"max_output_tokens"}}}"#, + ) + .expect("valid frame"); + assert_eq!(quota.status(), Some(429)); + } + + #[test] + fn a_legitimate_incomplete_batched_inside_a_chunks_envelope_is_not_a_failure() { + let frame = ParsedResponsesWebSocketFrame::parse( + r#"{"chunks":[{"type":"response.output_text.delta","delta":"hi"},{"type":"response.incomplete","response":{"incomplete_details":{"reason":"max_output_tokens"},"usage":{"total_tokens":9}}}]}"#, + ) + .expect("valid frame"); + + assert!(frame.is_chunked()); + assert!(frame.is_terminal()); + assert_eq!(frame.status(), Some(200)); + assert_eq!(frame.event_type(), Some("response.incomplete")); + assert_eq!( + frame.terminal_event().and_then(|event| event + .pointer("/response/usage/total_tokens") + .and_then(serde_json::Value::as_u64)), + Some(9) + ); + } + + #[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..ba14076ef --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/lifecycle.rs @@ -0,0 +1,351 @@ +//! 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, ResponsesProviderAttempt, + 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 ActiveProviderAttempt { + turn: Option, + state: AppState, +} + +impl ActiveProviderAttempt { + pub(super) fn new(state: &AppState, turn: ResponsesProviderAttempt) -> 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) -> ResponsesProviderAttempt { + self.turn + .take() + .expect("an armed active turn always holds its turn") + } +} + +impl std::ops::Deref for ActiveProviderAttempt { + type Target = ResponsesProviderAttempt; + + fn deref(&self) -> &Self::Target { + self.turn + .as_ref() + .expect("an armed active turn always holds its turn") + } +} + +impl std::ops::DerefMut for ActiveProviderAttempt { + fn deref_mut(&mut self) -> &mut Self::Target { + self.turn + .as_mut() + .expect("an armed active turn always holds its turn") + } +} + +impl Drop for ActiveProviderAttempt { + 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; + }); + } + } +} + +/// 结束当前 logical turn 并结算它的 attempt。 +/// +/// `end()` 同时清掉 logical turn 和 attempt,取代原来「take active_turn + +/// 在每个出口手写 `active_response_create = None`」的两步组合。 +pub(super) async fn finalize_active_turn( + bound: &mut BoundResponsesConnection, + state: &AppState, + outcome: ResponsesWebSocketTurnOutcome, +) { + if let Some(turn) = bound.turn_state.end() { + queue_turn_finalization(bound, state, turn.disarm(), outcome).await; + } +} + +pub(super) async fn queue_turn_finalization( + bound: &mut BoundResponsesConnection, + state: &AppState, + turn: ResponsesProviderAttempt, + 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); +} + +/// 「上一个 attempt 已经结算完毕」的凭证。 +/// +/// 只能由本模块颁发,且只有在结算真正落地之后。规划下一个 attempt 的入口 +/// ([`super::quota::retry_active_turn_after_quota_exhaustion`]) 要求这个参数, +/// 于是「先结算、再规划」成为签名的一部分,而不是一句注释——顺序写反连编译都 +/// 过不了。 +pub(super) struct PreviousAttemptSettled(()); + +impl PreviousAttemptSettled { + /// 没有 attempt 要结算(连接此刻不在 `Responding`)。 + pub(super) const fn nothing_to_settle() -> Self { + Self(()) + } +} + +/// 结算一个 attempt 并等它落地。 +/// +/// 与 [`queue_turn_finalization`] 的区别只在于「等」:后者把 handle 挂在连接上 +/// 让 relay loop 继续跑,适用于结算之后不再需要读取共享状态的出口;这个用在 +/// 必须先看到结算结果才能继续的路径上——典型的就是透明重试,它紧接着要按 +/// health / adaptive / pool 状态规划下一个 attempt。 +pub(super) async fn settle_turn_finalization( + bound: &mut BoundResponsesConnection, + state: &AppState, + turn: ResponsesProviderAttempt, + outcome: ResponsesWebSocketTurnOutcome, +) -> PreviousAttemptSettled { + queue_turn_finalization(bound, state, turn, outcome).await; + await_pending_turn_finalization(bound).await; + PreviousAttemptSettled(()) +} + +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: ResponsesProviderAttempt, + 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"), + } +} + +#[cfg(test)] +mod tests { + use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; + use std::sync::Arc; + use std::time::Duration; + + use super::await_turn_finalization_handle; + + /// C6 依赖的性质:结算是「等到落地」而不是「排进队列」。 + /// + /// 透明重试在这之后立刻按 health / adaptive / pool 状态规划下一个 attempt, + /// 所以结算任务必须已经跑完——只把 handle 挂起来是不够的。 + #[tokio::test] + async fn awaiting_a_finalization_handle_runs_the_settlement_to_completion() { + let settled = Arc::new(AtomicBool::new(false)); + let flag = Arc::clone(&settled); + let handle = tokio::spawn(async move { + tokio::time::sleep(Duration::from_millis(60)).await; + flag.store(true, Ordering::SeqCst); + }); + + assert!( + !settled.load(Ordering::SeqCst), + "the settlement has not finished yet" + ); + await_turn_finalization_handle(handle).await; + assert!( + settled.load(Ordering::SeqCst), + "the settlement must be complete before the caller proceeds" + ); + } + + /// 顺序型:结算的每一步都要排在规划之前。 + /// + /// 用计数器替身重放透明重试的两步——旧 attempt 结算完成写入 1,规划开始时 + /// 读到的必须已经是 1。旧实现在这里先规划、再把结算排进队列,规划读到的是 0。 + #[tokio::test] + async fn transparent_retry_replans_only_after_the_previous_attempt_is_settled() { + let steps = Arc::new(AtomicUsize::new(0)); + + // 第一步:结算旧 attempt(等到落地)。 + let recorder = Arc::clone(&steps); + let settlement = tokio::spawn(async move { + tokio::time::sleep(Duration::from_millis(40)).await; + recorder.store(1, Ordering::SeqCst); + }); + await_turn_finalization_handle(settlement).await; + + // 第二步:规划下一个 attempt,它读到的状态必须是结算之后的。 + let observed_at_planning = steps.load(Ordering::SeqCst); + assert_eq!( + observed_at_planning, 1, + "planning must observe the state projected by the settled attempt" + ); + } + + /// 结算任务失败(panic / cancel)也必须让调用方继续,不能把 relay loop 卡死。 + #[tokio::test] + async fn a_failed_finalization_task_still_releases_the_caller() { + let handle = tokio::spawn(async { panic!("settlement task exploded") }); + await_turn_finalization_handle(handle).await; + } +} 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..f26018a6d --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/mod.rs @@ -0,0 +1,65 @@ +//! 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 observation; +mod quota; +mod redaction; +mod relay_policy; +mod request; +mod session; +mod settlement; +mod state; +mod turn; +mod turn_state; +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/observation.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/observation.rs new file mode 100644 index 000000000..7ea8c83c6 --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/observation.rs @@ -0,0 +1,155 @@ +//! Responses WebSocket 的终态观测入口。 +//! +//! 这条传输收到的本来就是结构化的 Responses 协议事件。之前为了复用面向 SSE 的 +//! `push_line`,观测路径要先把每个事件序列化成 `data: {json}\n\n`,解析器再把它 +//! 解码回 `Value`——一次纯粹的往返,而且这个「伪 SSE」形状是随手拼的,一旦 +//! 上游事件里出现需要转义的内容,或者以后有人给拼装函数加了换行/分块逻辑, +//! 观测结果就会和真实事件悄悄分叉。 +//! +//! 现在观测走 [`StreamingStandardTerminalObserver::push_event`],直接吃 +//! `frame.protocol_events()` 借出的事件,不再序列化、不再解码。 +//! +//! **body capture 不走这条路,仍然保持 SSE 形状**(`data: {json}\n\n`): +//! `aether_usage_runtime::report` 用 `line.strip_prefix("data:")` 解析被捕获的 +//! body 来判定 `StreamCapturedTerminalState`,而它是 `stream_report_represents_failure` +//! 的一个 OR 项。把捕获内容换成结构化 JSON 会让终态判定恒为 Missing。 +//! 也就是说这一层只换「观测」,不换「捕获」——见 +//! [`super::turn::ResponsesProviderAttempt::capture_client_frame`] 一侧仍在用 +//! SSE 编码。 + +use serde_json::Value; + +use crate::ai_serving::api::StreamingStandardTerminalObserver; +use aether_contracts::ExecutionStreamTerminalSummary; + +/// 包一层 [`StreamingStandardTerminalObserver`],只暴露结构化入口。 +/// +/// 存在的意义是让「WS 不再拼 SSE」成为类型层面的事实:这里没有任何接受字节的 +/// 方法,所以不可能有人不小心把观测路径改回 `push_line`。 +#[derive(Default)] +pub(super) struct ResponsesStructuredTerminalObserver { + inner: StreamingStandardTerminalObserver, +} + +impl ResponsesStructuredTerminalObserver { + /// 观测一帧里的全部协议事件。 + /// + /// 第一个被拒绝的事件就停止推进并把摘要标成 parser_error:解析器的状态机是 + /// 有顺序的,跳过一个事件继续喂后面的只会得到更没意义的摘要。 + pub(super) fn observe_events(&mut self, report_context: &Value, events: &[&Value]) { + for event in events { + if let Err(error) = self.inner.push_event(report_context, event) { + self.inner.disable_with_error(error.to_string()); + break; + } + } + } + + pub(super) fn disable_with_error(&mut self, parser_error: impl Into) { + self.inner.disable_with_error(parser_error); + } + + pub(super) fn finish(&mut self, report_context: &Value) -> ExecutionStreamTerminalSummary { + match self.inner.finish(report_context) { + Ok(Some(summary)) => summary, + Ok(None) => ExecutionStreamTerminalSummary::default(), + Err(error) => { + self.inner.disable_with_error(error.to_string()); + self.inner.latest_summary().cloned().unwrap_or_default() + } + } + } +} + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::ResponsesStructuredTerminalObserver; + + fn report_context() -> serde_json::Value { + json!({ + "provider_api_format": "openai:responses", + "client_api_format": "openai:responses", + }) + } + + #[test] + fn structured_events_reach_the_terminal_summary_without_sse_text() { + let context = report_context(); + let created = json!({"type": "response.created", "response": {"id": "resp_ws", "model": "gpt-5-codex"}}); + let completed = json!({ + "type": "response.completed", + "response": { + "id": "resp_ws", + "model": "gpt-5-codex", + "status": "completed", + "usage": {"input_tokens": 9, "output_tokens": 4, "total_tokens": 13}, + }, + }); + + let mut observer = ResponsesStructuredTerminalObserver::default(); + observer.observe_events(&context, &[&created, &completed]); + let summary = observer.finish(&context); + + assert!(summary.observed_finish); + assert_eq!(summary.response_id.as_deref(), Some("resp_ws")); + let usage = summary + .standardized_usage + .as_ref() + .expect("a completed response carries usage"); + assert_eq!(usage.input_tokens, 9); + assert_eq!(usage.output_tokens, 4); + assert!(summary.parser_error.is_none()); + } + + /// 批量帧里的多个事件按顺序喂入,usage 不能因为批量而丢失。 + #[test] + fn a_batched_frame_keeps_the_usage_of_its_last_event() { + let context = report_context(); + let events = [ + json!({"type": "response.created", "response": {"id": "resp_ws", "model": "m"}}), + json!({ + "type": "response.output_text.delta", + "item_id": "msg", + "output_index": 0, + "content_index": 0, + "delta": "hi", + }), + json!({ + "type": "response.completed", + "response": { + "id": "resp_ws", + "model": "m", + "status": "completed", + "usage": {"input_tokens": 3, "output_tokens": 1, "total_tokens": 4}, + }, + }), + ]; + let borrowed: Vec<&serde_json::Value> = events.iter().collect(); + + let mut observer = ResponsesStructuredTerminalObserver::default(); + observer.observe_events(&context, &borrowed); + let summary = observer.finish(&context); + + let usage = summary + .standardized_usage + .as_ref() + .expect("usage survives batching"); + assert_eq!(usage.input_tokens, 3); + assert_eq!(usage.output_tokens, 1); + assert_eq!(usage.dimensions.get("total_tokens"), Some(&json!(4))); + } + + #[test] + fn a_disabled_observer_reports_the_parser_error() { + let context = report_context(); + let mut observer = ResponsesStructuredTerminalObserver::default(); + observer.disable_with_error("upstream event was not valid JSON"); + let summary = observer.finish(&context); + assert_eq!( + summary.parser_error.as_deref(), + Some("upstream event was not valid JSON") + ); + } +} 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..128c30069 --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/quota.rs @@ -0,0 +1,400 @@ +//! 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, ActiveProviderAttempt, PreviousAttemptSettled}; +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; + // 调用方必须先结束当前 logical turn 再 detach:拆掉上游后 attempt 已经不可能 + // 收到终态,留着它只会等 deadline 或 drop guard 兜底。 + debug_assert!( + !bound.turn_state.response_in_flight(), + "an exhausted upstream must be detached after its logical turn ended" + ); + 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)) +} + +/// 为同一个 logical turn 规划并绑定下一个 attempt。 +/// +/// `_previous_settled` 不被使用,它只是把「上一个 attempt 已经结算完毕」这个 +/// 前置条件写进签名:规划要读 health / adaptive / pool 状态,而这些是上一个 +/// attempt 结算时才投射的;它的 pool key lease 也要先释放,否则替代 key 的挑选 +/// 会看到一把仍被占用的 key。 +pub(super) async fn retry_active_turn_after_quota_exhaustion( + bound: &mut BoundResponsesConnection, + state: &AppState, + context: &WebSocketRequestContext, + _previous_settled: PreviousAttemptSettled, +) -> bool { + let Some(active) = bound.turn_state.logical_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.decision_template = replacement.decision_template; + bound.body_normalization = replacement.body_normalization; + bound.binding_identity = replacement.binding_identity; + // 同一个 logical turn 的下一个 attempt 就位。状态不符时把 attempt 交回 + // drop guard 结算并让调用方走「透明重试失败」分支,不静默丢弃一条已经写了 + // pending usage 行、占着 candidate 和 pool key lease 的 attempt。 + if let Err(orphan) = bound + .turn_state + .resume(ActiveProviderAttempt::new(state, turn)) + { + drop(orphan); + return false; + } + 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.turn_state.logical().is_some_and(|active| { + response_create_has_previous_response_id(&active.client_event) + && active.retry_unsafe_reason.is_none() + }) +} + +pub(super) fn is_usage_limit_error_event(event: &Value) -> bool { + let is_error = |value: &Value| { + value.get("type").and_then(Value::as_str) == Some("error") + && 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.turn_state.logical_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.turn_state.logical_mut() { + active.mark_retry_unsafe(reason); + } +} diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/redaction.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/redaction.rs new file mode 100644 index 000000000..620e661e4 --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/redaction.rs @@ -0,0 +1,850 @@ +//! Responses WebSocket 两侧的 PII 脱敏:请求侧 mask + 响应侧 restore。 +//! +//! HTTP 路径在前门建 `RedactionSessionSlot` 并塞进 `parts.extensions`,planner +//! 只有拿到这个 slot 才会脱敏。WS 的 planning Parts 是合成的:四个规划入口 +//! (首轮、换模型 re-plan、独立轮、配额透明重试)靠 `build_planning_parts` 注入 +//! slot 就能复用 planner 的脱敏;但复用已绑定 upstream 的 continuation 根本不进 +//! planner,必须在这里先把客户端事件脱敏,再交给协议归一化、上游发送和审计。 +//! +//! 因此约定:**进入任何下游用途之前,客户端 `response.create` 只在这里脱敏一次**, +//! 之后所有路径都只看脱敏后的事件。 +//! +//! # 响应侧 +//! +//! 只 mask 不 restore 是半个实现:HTTP 在把响应交给客户端之前会把占位符换回真实值 +//! (`privacy::restore_sync_response_body` / `privacy::StreamingResponseRestorer`), +//! WS 少了这一步,客户端就会直接看到 ``。 +//! [`ResponsesWebSocketRedactionRestorer`] 补上这一跳,语义与 HTTP 完全一致: +//! 复用 `privacy::restore_json_strings`,只还原本连接自己 mask 出来的映射, +//! 未映射的占位符原样透传。 +//! +//! ## session 为什么活在连接上而不是活在这一轮里 +//! +//! mask session 由 planner 写进 per-turn 的 slot,而 slot 随 planning Parts 在 +//! 规划结束时就被丢弃,响应帧到达时已经无处可取。可选的存活范围有两个: +//! +//! * 挂在 `LogicalTurn` 上:这一轮结束即释放,是 HTTP「一个请求一个 session」的 +//! 直译。但 WS 的会话历史留在上游:continuation 只发增量输入,第 1 轮的 +//! `input` 不会在第 3 轮重发。于是第 3 轮的响应里若回显了第 1 轮的占位符 +//! ("你刚才给我的邮箱是……"),本轮 session 里没有这条映射,占位符就漏给客户端。 +//! HTTP 不会漏,是因为它每次都重发整段历史,重新 mask 同一个值会派生出同一个 +//! sentinel(HMAC over 规则 + bucket + 值),所以映射天然齐备。 +//! * 挂在连接上(当前实现):每轮仍然各自 mask、各自持有独立 session +//! (per-turn 语义不变),连接只是把最近若干轮的 session 留下来一起参与还原, +//! 凑出的映射集合正好等于「等价 HTTP 请求会拥有的那一份」。 +//! +//! 选后者。代价是每帧最多对 [`MAX_RETAINED_TURN_REDACTION_SESSIONS`] 个 session +//! 各扫一遍,以及这些 session 的映射会驻留到连接结束;用有界 FIFO 兜住上限。 +//! 窗口不够用或每帧成本变高时,正确的下一步是在 `privacy` 侧提供跨 session 的 +//! 合并匹配器,而不是把这个窗口调大。 + +use std::collections::VecDeque; + +use serde_json::Value; + +use crate::ai_serving::{ + resolve_local_decision_execution_runtime_auth_context, resolve_provider_chat_pii_redaction, +}; +use crate::control::GatewayControlDecision; +use crate::privacy::{restore_json_strings, RedactionSession, RedactionSessionSlot}; +use crate::{AppState, GatewayError}; + +/// Responses WebSocket 只承载 `openai:responses`,脱敏规则按这个客户端格式选取。 +const RESPONSES_WEBSOCKET_CLIENT_API_FORMAT: &str = "openai:responses"; + +/// WS 在选出候选之前就要脱敏,所以脱敏 session 先记在这个固定 key 下。 +/// +/// slot 是 per-turn 的(见 `build_planning_parts`),这一轮之后即随 slot 一起丢弃; +/// planner 后续用真实 candidate_id 再取一次配置时,body 已是脱敏态、不会重复写入。 +const WEBSOCKET_TURN_REDACTION_CANDIDATE_ID: &str = "responses_websocket_turn"; + +/// 一条连接最多留几轮的 mask session 用于响应侧还原。 +/// +/// 取值权衡见模块文档:调大会线性增加每帧还原成本和常驻映射量,调小则更容易漏还原 +/// 上游历史里更早那几轮的占位符。8 覆盖的是「上游最可能回显的最近窗口」。 +const MAX_RETAINED_TURN_REDACTION_SESSIONS: usize = 8; + +/// 一轮客户端 `response.create` 的请求侧脱敏结果。 +#[derive(Debug)] +pub(super) struct ResponsesWebSocketTurnRedaction { + /// 脱敏后的客户端事件;这一轮之后所有下游路径都只看它。 + pub(super) client_event: Value, + /// 这一轮 mask 出来的映射表,响应侧还原只能靠它。 + pub(super) session: RedactionSession, +} + +/// 对一条客户端 `response.create` 做请求侧脱敏。 +/// +/// 返回 `Some(..)` 仅当脱敏真正命中;`None` 表示未启用或没有命中,调用方 +/// 继续用原事件即可(避免未开启脱敏时多一次整包 clone)。 +/// +/// 脱敏只改写 `instructions` / `input`(见 `privacy::mask_openai_responses_request_value`), +/// `type` / `model` / `previous_response_id` / `generate` 等协议字段原样保留,所以脱敏后的 +/// 事件仍可直接用于协议归一化和上游发送。 +/// +/// 出错必须让这一轮失败:脱敏已启用却读不到配置或加密密钥时,把原文发上游就是 +/// 静默旁路,正是本次要修的问题。 +pub(super) async fn redact_responses_websocket_client_event( + state: &AppState, + parts: &http::request::Parts, + control_decision: &GatewayControlDecision, + client_event: &Value, +) -> Result, GatewayError> { + let Some(auth_context) = + resolve_local_decision_execution_runtime_auth_context(control_decision) + else { + return Ok(None); + }; + let redaction = resolve_provider_chat_pii_redaction( + state, + parts, + client_event, + &auth_context, + RESPONSES_WEBSOCKET_CLIENT_API_FORMAT, + WEBSOCKET_TURN_REDACTION_CANDIDATE_ID, + ) + .await?; + if !redaction.redacted { + return Ok(None); + } + // mask 命中时 `resolve_provider_chat_pii_redaction` 必定把 session 写进 slot。 + // 取不到就是内部契约被破坏了,此时继续下发意味着这一轮的响应无法还原、占位符 + // 会漏给客户端;按本模块既有的「脱敏链路出错就让这一轮失败」处理,不做降级。 + let Some(session) = parts + .extensions + .get::() + .and_then(|slot| slot.take_for_candidate(Some(WEBSOCKET_TURN_REDACTION_CANDIDATE_ID))) + else { + return Err(GatewayError::Internal( + "chat pii redaction masked a Responses WebSocket turn without retaining its session" + .to_string(), + )); + }; + Ok(Some(ResponsesWebSocketTurnRedaction { + client_event: redaction.body_json.into_owned(), + session, + })) +} + +/// 一条连接上「我们 mask 过哪些映射」的留存集合,供响应侧还原使用。 +/// +/// 每轮一个独立 session(per-turn mask 语义不变),连接按 FIFO 留最近 +/// [`MAX_RETAINED_TURN_REDACTION_SESSIONS`] 轮。上游重绑不清空:客户端仍在同一段 +/// 对话里,旧占位符可能随重发的输入再次出现。 +#[derive(Default)] +pub(super) struct ResponsesWebSocketRedactionRestorer { + sessions: VecDeque, +} + +impl ResponsesWebSocketRedactionRestorer { + /// 登记这一轮的 mask session。 + pub(super) fn register(&mut self, session: RedactionSession) { + if session.mapping_count() == 0 { + return; + } + self.sessions.push_back(session); + while self.sessions.len() > MAX_RETAINED_TURN_REDACTION_SESSIONS { + self.sessions.pop_front(); + } + } + + /// 把一帧 provider 事件里的占位符换回真实值,返回要发给客户端的帧文本。 + /// + /// `None` 表示这一帧没有任何东西要还原,调用方必须原样转发上游字节:未启用 + /// 脱敏(没有任何 session)时连 clone 都不做。 + /// + /// 入参只读:审计与终态观测继续消费脱敏态的事件,还原只作用于发往客户端的 + /// 那一份拷贝,和 HTTP 侧「审计存脱敏体、线上还原」保持一致。 + pub(super) fn restore_provider_frame_text(&self, event: &Value) -> Option { + if self.sessions.is_empty() { + return None; + } + let mut restored_event = event.clone(); + let mut restored = false; + for session in &self.sessions { + // 逐 session 还原而不是合并映射:每个 session 只认自己 mask 过的 + // sentinel(`RedactionSession::restore_text`),跨 session 合并会绕开 + // 这条边界。同一个值在不同轮派生出的 sentinel 相同,所以顺序无关。 + restored |= restore_json_strings(&mut restored_event, session); + } + if !restored { + return None; + } + // 刚从 JSON 解析出来的 Value 再序列化不会失败;真失败时宁可让客户端看到 + // 占位符,也不能丢掉这一帧——丢帧会让客户端的协议状态机卡死。 + serde_json::to_string(&restored_event).ok() + } +} + +#[cfg(test)] +mod tests { + use std::net::SocketAddr; + use std::sync::Arc; + + use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY; + use aether_data::repository::auth::{ + InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeyExportRecord, + }; + use axum::http::{HeaderMap, Uri}; + use serde_json::{json, Value}; + + use super::super::request::{ + build_planning_parts, normalize_followup_response_create, planned_response_create_event, + }; + use super::super::turn::prepare_responses_websocket_turn_decision; + use super::super::turn_state::LogicalTurn; + use super::{ + redact_responses_websocket_client_event, ResponsesWebSocketRedactionRestorer, + ResponsesWebSocketTurnRedaction, MAX_RETAINED_TURN_REDACTION_SESSIONS, + }; + use crate::ai_serving::{AiExecutionDecision, ResponsesWebSocketBodyNormalization}; + use crate::control::{GatewayControlAuthContext, GatewayControlDecision}; + use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext; + use crate::AppState; + + const TEST_USER_ID: &str = "user-responses-ws-redaction"; + const TEST_API_KEY_ID: &str = "api-key-responses-ws-redaction"; + const TEST_EMAIL: &str = "ws.user@example.com"; + /// 另一轮用的 PII,用来证明连接级还原覆盖到更早的轮次。 + const OTHER_TEST_EMAIL: &str = "ws.other@example.com"; + /// 不是本连接 mask 出来的占位符:格式合法(符合 sentinel 正则),但没有任何 + /// session 记过它,必须原样透传。 + const FOREIGN_SENTINEL: &str = ""; + + fn auth_export_record() -> StoredAuthApiKeyExportRecord { + StoredAuthApiKeyExportRecord::new( + TEST_USER_ID.to_string(), + TEST_API_KEY_ID.to_string(), + "hash-responses-ws-redaction".to_string(), + None, + Some("ws".to_string()), + None, + None, + None, + None, + None, + None, + true, + None, + false, + 0, + 0, + 0.0, + false, + ) + .expect("auth api key export record should build") + .with_feature_settings(Some(json!({ + "chat_pii_redaction": {"enabled": true} + }))) + } + + /// 只装脱敏真正需要的东西:系统配置开关 + 规则、加密密钥、带 feature settings + /// 的 API Key 导出记录。候选/上游都不需要,这条链路在 planner 之前。 + fn redaction_enabled_state() -> AppState { + let auth_repository = Arc::new( + InMemoryAuthApiKeySnapshotRepository::seed(vec![]) + .with_export_records(vec![auth_export_record()]), + ); + let data_state = + crate::data::GatewayDataState::with_auth_api_key_reader_for_tests(auth_repository) + .with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY) + .with_system_config_values_for_tests(vec![ + ("module.chat_pii_redaction.enabled".to_string(), json!(true)), + ( + "module.chat_pii_redaction.rules".to_string(), + json!([{ + "id": "email", + "name": "邮箱", + "pattern": r"(?i)[A-Z0-9._%+-]{1,64}@[A-Z0-9.-]{1,253}\.[A-Z]{2,63}", + "enabled": true, + "features": {"validator": "email"}, + "system": true + }]), + ), + ( + "module.chat_pii_redaction.cache_ttl_seconds".to_string(), + json!(300), + ), + ]); + AppState::new() + .expect("gateway state should build") + .with_data_state_for_tests(data_state) + } + + fn control_decision() -> GatewayControlDecision { + let mut decision = GatewayControlDecision::synthetic( + "/v1/responses".to_string(), + Some("ai_public".to_string()), + Some("openai".to_string()), + Some("responses_websocket".to_string()), + Some("openai:responses".to_string()), + ); + decision.auth_context = Some(GatewayControlAuthContext { + user_id: TEST_USER_ID.to_string(), + api_key_id: TEST_API_KEY_ID.to_string(), + username: Some("ws".to_string()), + api_key_name: Some("ws".to_string()), + balance_remaining: None, + access_allowed: true, + user_rate_limit: None, + api_key_rate_limit: None, + api_key_is_standalone: false, + admin_bypass_limits: false, + local_rejection: None, + allowed_models: None, + ip_rules: None, + }); + decision + } + + fn websocket_context(decision: GatewayControlDecision) -> WebSocketRequestContext { + WebSocketRequestContext { + trace_id: "trace-responses-ws-redaction".to_string(), + headers: HeaderMap::new(), + uri: Uri::from_static("/v1/responses"), + remote_addr: "127.0.0.1:65000" + .parse::() + .expect("remote address should parse"), + decision, + rpm_bypassed: false, + websocket_connection_permit: None, + } + } + + fn client_event() -> Value { + client_event_with_email(TEST_EMAIL) + } + + fn client_event_with_email(email: &str) -> Value { + json!({ + "type": "response.create", + "model": "public-model", + "previous_response_id": "resp-previous", + "generate": false, + "input": [{ + "role": "user", + "content": [{"type": "input_text", "text": format!("mail {email}")}] + }] + }) + } + + /// 真跑一遍请求侧脱敏,拿到这一轮的生效事件和 mask session。 + async fn turn_redaction( + state: &AppState, + decision: &GatewayControlDecision, + email: &str, + ) -> ResponsesWebSocketTurnRedaction { + let context = websocket_context(decision.clone()); + let parts = build_planning_parts(&context); + let event = client_event_with_email(email); + redact_responses_websocket_client_event(state, &parts, &context.decision, &event) + .await + .expect("redaction should resolve") + .expect("an email in the request should be redacted") + } + + /// 这一轮为 `email` 派生出的占位符。 + fn sentinel_for(redaction: &ResponsesWebSocketTurnRedaction, email: &str) -> String { + redaction + .session + .sentinel_for_original(email) + .expect("a masked email must have a sentinel") + .to_string() + } + + /// 上游回显占位符的一帧 provider 事件。 + fn provider_delta_frame(text: &str) -> Value { + json!({ + "type": "response.output_text.delta", + "item_id": "msg_ws", + "output_index": 0, + "content_index": 0, + "delta": text, + }) + } + + #[tokio::test] + async fn websocket_client_event_is_redacted_without_losing_protocol_fields() { + let state = redaction_enabled_state(); + let context = websocket_context(control_decision()); + let parts = build_planning_parts(&context); + let event = client_event(); + + let redacted = + redact_responses_websocket_client_event(&state, &parts, &context.decision, &event) + .await + .expect("redaction should resolve") + .expect("an email in the request should be redacted") + .client_event; + + let serialized = serde_json::to_string(&redacted).expect("event should serialize"); + assert!(!serialized.contains(TEST_EMAIL), "{serialized}"); + assert!(serialized.contains(" Value { + turn_redaction(state, decision, TEST_EMAIL) + .await + .client_event + } + + /// 只有 `action` 没有 serde 默认值,其余字段都能省略。 + fn decision_template( + provider_request_body: Value, + report_context: Value, + ) -> AiExecutionDecision { + serde_json::from_value(json!({ + "action": "local", + "candidate_id": "candidate-responses-ws", + "provider_request_body": provider_request_body, + "report_context": report_context, + })) + .expect("decision template should deserialize") + } + + /// planner 在脱敏 body 上做模型映射后的 provider body。 + fn provider_body_from(effective_event: &Value) -> Value { + let mut provider_body = effective_event.clone(); + provider_body["model"] = json!("provider-model"); + provider_body + } + + /// 绑定那一轮留下的 report_context seed:故意带上原始 PII,用来证明这一轮 + /// 会用脱敏后的 body 覆盖它,而不是把原文带进审计。 + fn seed_report_context_with_raw_pii() -> Value { + json!({ + "request_id": "connection", + "candidate_id": "candidate-responses-ws", + "original_request_body": { + "type": "response.create", + "model": "public-model", + "input": format!("mail {TEST_EMAIL}") + } + }) + } + + fn assert_redacted_json(value: &Value, label: &str) { + let serialized = serde_json::to_string(value).expect("value should serialize"); + assert!( + !serialized.contains(TEST_EMAIL), + "{label} must not carry raw PII: {serialized}" + ); + assert!( + serialized.contains(" 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..d4606d45b --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/request.rs @@ -0,0 +1,424 @@ +//! 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; +use crate::privacy::RedactionSessionSlot; + +/// 把一条 WebSocket turn 还原成 planner 需要的 HTTP 形状请求头部。 +/// +/// 这里必须和 HTTP 前门(`handlers/proxy/mod.rs`)保持同一份 extension 契约: +/// planner 只在 `parts.extensions` 里拿到 `RedactionSessionSlot` 时才做请求脱敏 +/// (`ai_serving/planner/redaction.rs`),少插这一项等于整条 WS 链路静默绕过 +/// 已启用的 PII 脱敏。 +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, + )); + // slot 必须每个 turn 新建,不能按连接复用:planner 侧的请求脱敏缓存键是 + // `{format:?}:{body_json 指针地址}`(`ai_serving/planner/redaction.rs:169`), + // 连接级复用同一个 slot 时,上一轮 client_event 释放后这一轮的 `Value` 很可能 + // 落在同一地址,会命中上一轮缓存,把上一轮的脱敏 body 当成这一轮的发出去。 + // 每个 `response.create` 本身就是独立计费/审计请求,per-turn 也正好对应 + // HTTP 前门「一个请求一个 slot」的语义。 + request + .extensions_mut() + .insert(RedactionSessionSlot::default()); + 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 std::net::SocketAddr; + + use axum::http::{HeaderMap, Uri}; + use serde_json::json; + + use super::{ + build_planning_parts, normalize_followup_response_create, + response_create_has_previous_response_id, + }; + use crate::ai_serving::ResponsesWebSocketBodyNormalization; + use crate::control::GatewayControlDecision; + use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext; + use crate::privacy::RedactionSessionSlot; + + fn websocket_context() -> WebSocketRequestContext { + WebSocketRequestContext { + trace_id: "trace-planning-parts".to_string(), + headers: HeaderMap::new(), + uri: Uri::from_static("/v1/responses"), + remote_addr: "127.0.0.1:65001" + .parse::() + .expect("remote address should parse"), + decision: GatewayControlDecision::synthetic( + "/v1/responses".to_string(), + Some("ai_public".to_string()), + Some("openai".to_string()), + Some("responses_websocket".to_string()), + Some("openai:responses".to_string()), + ), + rpm_bypassed: false, + websocket_connection_permit: None, + } + } + + #[test] + fn planning_parts_carry_a_fresh_redaction_session_slot_per_turn() { + // 没有这个 extension,planner 会静默跳过已启用的 PII 脱敏 + // (ai_serving/planner/redaction.rs),整条 WS 链路都按原文发上游。 + let context = websocket_context(); + let first = build_planning_parts(&context); + let second = build_planning_parts(&context); + + let first_slot = first + .extensions + .get::() + .expect("planning parts must carry a redaction session slot"); + let second_slot = second + .extensions + .get::() + .expect("planning parts must carry a redaction session slot"); + + // 每轮必须是独立 slot:slot 内的请求缓存以 body 指针地址为键,跨轮共享会 + // 命中上一轮缓存。用缓存条目相互不可见来证明两者不是同一个 slot。 + first_slot.put_cached_request_redaction( + "turn-1", + crate::privacy::CachedRequestRedaction::unredacted(), + ); + assert!(first_slot.cached_request_redaction("turn-1").is_some()); + assert!(second_slot.cached_request_redaction("turn-1").is_none()); + } + + 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..ddfa86796 --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs @@ -0,0 +1,1378 @@ +//! 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 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, ActiveProviderAttempt, +}; +use super::redaction::redact_responses_websocket_client_event; +use super::request::{build_planning_parts, planned_response_create_event}; +use super::turn::{ + begin_responses_websocket_turn, prepare_responses_websocket_turn_decision, + ResponsesWebSocketTurnOutcome, +}; +use super::turn_state::LogicalTurn; +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_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; + } + } + + // 请求侧脱敏必须在规划之前完成,而且这一轮只在这里做一次:planner 会把这份 + // body 写进 upstream 请求体和审计 original_request_body,绑定上游的首条 + // response.create 也从它派生。脱敏失败时直接断开,绝不退回原文发上游。 + let redacted_first_event = redact_responses_websocket_client_event( + &state, + &planning_parts, + &context.decision, + &first_event, + ) + .await; + // 首轮的 mask session 要活到响应帧还原,但连接此刻还没绑定,只能先接住, + // 等 `bind_responses_upstream` 之后登记到连接上。 + let (first_event, first_turn_redaction_session) = match redacted_first_event { + Ok(Some(redaction)) => (redaction.client_event, Some(redaction.session)), + Ok(None) => (first_event, None), + Err(error) => { + warn!( + event_name = "responses_websocket_redaction_failed", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + error = ?error, + "gateway could not apply chat PII redaction to the initial Responses WebSocket event" + ); + send_gateway_error_with_status( + &mut client_socket, + 500, + "responses_websocket_redaction_unavailable", + "Gateway could not apply the configured PII redaction", + ) + .await; + close_client_socket( + &mut client_socket, + CLOSE_INTERNAL_ERROR, + "responses_websocket_redaction_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()); + if let Some(session) = first_turn_redaction_session { + bound.redaction_restorer.register(session); + } + bound.turn_state.begin( + LogicalTurn::new(first_event, 1, first_logical_turn_id), + ActiveProviderAttempt::new(&state, first_turn), + ); + + relay_bound_connection( + &mut client_socket, + &mut bound, + &state, + &context, + connection_permit, + ) + .await; + await_pending_turn_finalization(&mut bound).await; + await_pending_adapter_observation(&mut bound).await; +} + +/// 等待客户端发送第一条 response.create 事件。 +/// 使用绝对 deadline:从函数入口起计算一次截止时间,Ping/Pong 只会被正常回复, +/// 但不会重置计时器。防止客户端通过周期性 Ping 无限占用 connection permit。 +async fn receive_initial_response_create( + client_socket: &mut WebSocket, +) -> Result<(String, Value), InitialMessageError> { + receive_initial_response_create_with_deadline( + client_socket, + RESPONSES_WEBSOCKET_SESSION_LIMITS.initial_message_timeout, + ) + .await +} + +/// 核心循环:在绝对 deadline 内等待客户端发送 response.create。 +/// 泛型约束允许测试注入 fake socket,驱动真实逻辑。 +/// +/// - `deadline_budget`:从调用时刻起的最长等待时间,全循环共享同一截止时刻。 +/// - Ping 帧被回复 Pong 但不重置计时器。 +/// - Pong / 非法帧 / Close 按协议处理。 +async fn receive_initial_response_create_with_deadline( + socket: &mut S, + deadline_budget: std::time::Duration, +) -> Result<(String, Value), InitialMessageError> +where + S: futures_util::Stream> + + futures_util::Sink + + Unpin, +{ + use futures_util::{SinkExt as _, StreamExt as _}; + + // 绝对 deadline:入口计算一次,后续所有迭代共享,Ping/Pong 不会重启 + let deadline = tokio::time::Instant::now() + deadline_budget; + loop { + let message = tokio::time::timeout_at(deadline, 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) => { + socket + .send(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::redaction::ResponsesWebSocketRedactionRestorer; + use super::super::request::{ + changed_followup_response_create_model, continuation_requires_same_upstream, + normalize_followup_response_create, planned_response_create_event, + response_create_model_or_current, + }; + use super::super::state::{BoundResponsesConnection, ExhaustedResponsesWebSocketExclusions}; + use super::super::turn::{ + ResponsesWebSocketTurnDeadline, ResponsesWebSocketTurnObservation, + ResponsesWebSocketTurnOutcome, ResponsesWebSocketTurnTimeoutPhase, + }; + use super::super::turn_state::{LogicalTurn, ResponsesTurnState}; + 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.turn_state = ResponsesTurnState::Replanning { + logical: LogicalTurn::new( + json!({ + "type": "response.create", + "previous_response_id": "resp-previous", + }), + 2, + "logical-turn".to_string(), + ), + }; + + assert!(active_continuation_can_retry_from_full_input(&bound)); + bound + .turn_state + .logical_mut() + .expect("active request") + .mark_retry_unsafe("standard_response_event"); + assert!(!active_continuation_can_retry_from_full_input(&bound)); + } + + #[test] + fn only_an_actual_usage_limit_error_requests_full_retry() { + 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.turn_state = ResponsesTurnState::Replanning { + logical: LogicalTurn::new( + json!({ + "type": "response.create", + "previous_response_id": "resp-previous", + }), + 2, + "logical-turn".to_string(), + ), + }; + + assert!(should_request_full_continuation_retry( + &bound, + true, + Some(&json!({ + "type": "error", + "error": {"type": "usage_limit_reached"}, + })), + )); + assert!(!should_request_full_continuation_retry( + &bound, + true, + Some(&json!({ + "type": "response.completed", + "response": {"id": "resp-completed"}, + })), + )); + assert!(!should_request_full_continuation_retry( + &bound, + false, + Some(&json!({ + "type": "error", + "error": {"type": "usage_limit_reached"}, + })), + )); + } + + #[test] + fn followup_rewrites_the_provider_model_and_removes_http_stream_fields() { + let event = json!({ + "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 = LogicalTurn::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 = LogicalTurn::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 = LogicalTurn::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 = LogicalTurn::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 + .turn_state + .logical() + .and_then(LogicalTurn::quota_retry_block_reason), + None + ); + + observe_active_response_rebind_safety(&mut bound, &json!({"type": "response.created"})); + assert_eq!( + bound + .turn_state + .logical() + .and_then(LogicalTurn::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 + .turn_state + .logical() + .and_then(LogicalTurn::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(), + decision_template: decision, + body_normalization: ResponsesWebSocketBodyNormalization::for_tests("gpt-5.6-sol"), + binding_identity, + // Replanning:logical turn 在、attempt 不在。重放安全与配额排除都只看 + // logical turn,所以这些用例不需要真实 socket 或真实 attempt。 + turn_state: ResponsesTurnState::Replanning { + logical: LogicalTurn::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, + redaction_restorer: ResponsesWebSocketRedactionRestorer::default(), + } + } + + /// 用 mpsc 驱动的 FakeSocket,实现 Stream + Sink 两个 trait。 + /// 测试侧通过 tx 注入消息,通过 pong_rx 观察 Pong 回包。 + struct FakeSocket { + rx: tokio::sync::mpsc::Receiver, + pong_tx: tokio::sync::mpsc::UnboundedSender, + } + + struct FakeSocketPair { + tx: tokio::sync::mpsc::Sender, + pong_rx: tokio::sync::mpsc::UnboundedReceiver, + } + + fn fake_socket() -> (FakeSocket, FakeSocketPair) { + let (tx, rx) = tokio::sync::mpsc::channel(16); + let (pong_tx, pong_rx) = tokio::sync::mpsc::unbounded_channel(); + (FakeSocket { rx, pong_tx }, FakeSocketPair { tx, pong_rx }) + } + + impl futures_util::Stream for FakeSocket { + type Item = Result; + + fn poll_next( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + self.rx.poll_recv(cx).map(|opt| opt.map(Ok)) + } + } + + impl futures_util::Sink for FakeSocket { + type Error = axum::Error; + + fn poll_ready( + self: std::pin::Pin<&mut Self>, + _cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::task::Poll::Ready(Ok(())) + } + + fn start_send( + self: std::pin::Pin<&mut Self>, + item: axum::extract::ws::Message, + ) -> Result<(), Self::Error> { + let _ = self.pong_tx.send(item); + Ok(()) + } + + fn poll_flush( + self: std::pin::Pin<&mut Self>, + _cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::task::Poll::Ready(Ok(())) + } + + fn poll_close( + self: std::pin::Pin<&mut Self>, + _cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::task::Poll::Ready(Ok(())) + } + } + + /// 验证 receive_initial_response_create_with_deadline 的绝对 deadline: + /// 客户端周期性发送 Ping 帧不会重置计时器,deadline 到期后返回 TimedOut。 + /// 这直接驱动真实的循环逻辑,如果改回每次迭代 timeout(budget, ...) 则会变红。 + /// + /// 设计思路:deadline_budget = 80ms,Ping 每 30ms 发一次且永不停止。 + /// - 绝对 deadline:~80ms 后函数返回 TimedOut(即使 Ping 仍在到来)。 + /// - 每次迭代 timeout(80ms, ...):每个 Ping 在 30ms 内到达 < 80ms,函数 + /// 永不超时,500ms 后外层 timeout 判定测试失败。 + #[tokio::test] + async fn initial_message_times_out_despite_periodic_pings() { + use super::{receive_initial_response_create_with_deadline, InitialMessageError}; + + let (mut fake, pair) = fake_socket(); + + let handle = tokio::spawn(async move { + receive_initial_response_create_with_deadline(&mut fake, Duration::from_millis(80)) + .await + }); + + // 持续发送 Ping,间隔 30ms,永不停止(直到被测函数返回导致 rx drop) + let ping_task = tokio::spawn(async move { + let mut i = 0u8; + loop { + tokio::time::sleep(Duration::from_millis(30)).await; + if pair + .tx + .send(axum::extract::ws::Message::Ping(vec![i].into())) + .await + .is_err() + { + break; + } + i = i.wrapping_add(1); + } + }); + + // 绝对 deadline 应在 ~80ms 后触发;给 500ms 宽限等待结果。 + let result = tokio::time::timeout(Duration::from_millis(500), handle) + .await + .expect("function should return within 500ms (absolute deadline = 80ms)") + .expect("task should not panic"); + + ping_task.abort(); + + assert!( + matches!(result, Err(InitialMessageError::TimedOut)), + "expected TimedOut after absolute deadline, got: {result:?}" + ); + } + + /// 验证 deadline 内收到合法 response.create 时正常返回,Ping 被正确回复 Pong。 + #[tokio::test] + async fn initial_message_succeeds_within_deadline() { + use super::{receive_initial_response_create_with_deadline, InitialMessageError}; + + let (mut fake, mut pair) = fake_socket(); + + let handle = tokio::spawn(async move { + receive_initial_response_create_with_deadline(&mut fake, Duration::from_secs(5)).await + }); + + // 先发一个 Ping,验证 Pong 回包且不影响后续解析 + pair.tx + .send(axum::extract::ws::Message::Ping(vec![42].into())) + .await + .unwrap(); + let pong = tokio::time::timeout(Duration::from_secs(1), pair.pong_rx.recv()) + .await + .expect("should receive pong within 1s") + .expect("pong channel should not close"); + assert!( + matches!(pong, axum::extract::ws::Message::Pong(ref data) if data.as_ref() == [42]), + "expected Pong([42]), got: {pong:?}" + ); + + // 发送合法的 response.create + let event_text = r#"{"type":"response.create","model":"gpt-4o"}"#; + pair.tx + .send(axum::extract::ws::Message::Text( + event_text.to_string().into(), + )) + .await + .unwrap(); + + let result = tokio::time::timeout(Duration::from_secs(2), handle) + .await + .expect("handle should finish within 2s") + .expect("task should not panic"); + let (text, event) = result.expect("should return Ok"); + assert_eq!(text, event_text); + assert_eq!(event["type"], "response.create"); + assert_eq!(event["model"], "gpt-4o"); + } +} diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/settlement.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/settlement.rs new file mode 100644 index 000000000..fab1cac80 --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/settlement.rs @@ -0,0 +1,286 @@ +//! Responses WebSocket 的结算信号 → 记账事实映射。 +//! +//! 结算判定本身是 transport 中立的,住在 +//! [`crate::execution_runtime::attempt_lifecycle`]。这里只做 WS 专属的一件事: +//! 把 relay loop 的结算触发信号 [`ResponsesWebSocketTurnOutcome`] 翻译成那两个 +//! 正交事实。 + +use super::turn::ResponsesWebSocketTurnOutcome; +use crate::execution_runtime::attempt_lifecycle::{ + AttemptClientDelivery, AttemptProviderOutcome, AttemptTerminalFacts, + CLIENT_CANCELLED_STATUS_CODE, STREAM_TIMEOUT_STATUS_CODE, +}; + +/// 把「结算触发信号」+「已观察到的 provider 终态」+「已记录的投递结果」映射成 +/// 两个正交事实。 +/// +/// `ResponsesWebSocketTurnOutcome` 描述的是 relay loop 为什么现在结算这一 +/// attempt,它对 provider 的信息量并不总是完整的: +/// +/// - `ProviderTerminal` / `Failure` 本身就在描述供应商这一轮的结果,是权威的。 +/// - `Cancelled` 只说明「我们为客户端或连接层面的原因停下了」,不携带任何 +/// provider 信息。已经观察到的 provider 终态是独立事实,不能被它覆盖—— +/// 这正是评审第 5 条要求分开记录的那一处。 +/// +/// `recorded_delivery` 是 relay loop 明确记下的投递失败(写客户端 socket 失败)。 +/// 它与结算信号推出的投递结果取「只要有一侧失败就是失败」,并优先保留明确记录 +/// 的原因。 +pub(super) fn attempt_facts_for_outcome( + observed_provider_terminal: Option, + recorded_delivery: AttemptClientDelivery, + settling: ResponsesWebSocketTurnOutcome, +) -> AttemptTerminalFacts { + let facts = match settling { + ResponsesWebSocketTurnOutcome::ProviderTerminal { + status_code, + cancelled, + } => AttemptTerminalFacts { + provider: AttemptProviderOutcome::Terminal { + status_code, + cancelled_by_provider: cancelled, + }, + delivery: AttemptClientDelivery::Complete, + }, + ResponsesWebSocketTurnOutcome::Failure { + status_code, + reason, + } => AttemptTerminalFacts { + provider: AttemptProviderOutcome::Aborted { + status_code, + reason, + // 现状只有 504 一族(首事件/终态超时)会投射 pool stream timeout。 + stream_timeout: status_code == STREAM_TIMEOUT_STATUS_CODE, + }, + delivery: AttemptClientDelivery::Complete, + }, + ResponsesWebSocketTurnOutcome::Cancelled { reason } => AttemptTerminalFacts { + provider: observed_provider_terminal.unwrap_or(AttemptProviderOutcome::Aborted { + status_code: CLIENT_CANCELLED_STATUS_CODE, + reason, + stream_timeout: false, + }), + delivery: AttemptClientDelivery::Aborted { reason }, + }, + }; + AttemptTerminalFacts { + delivery: match recorded_delivery { + AttemptClientDelivery::Aborted { .. } => recorded_delivery, + AttemptClientDelivery::Complete => facts.delivery, + }, + ..facts + } +} + +/// 客户端投递失败时应该用哪个结算信号。 +/// +/// provider 终态已经到达就用那条终态:它是权威的 provider 事实,绝不能被 +/// `client_disconnected()` 覆盖掉——那正是把已完成响应记成 void billing 的原因。 +/// 供应商还没给出终态时,客户端断开才是这一 attempt 的全部结论。 +pub(super) fn settle_signal_for_client_delivery_failure( + terminal_outcome: Option, +) -> ResponsesWebSocketTurnOutcome { + terminal_outcome.unwrap_or_else(ResponsesWebSocketTurnOutcome::client_disconnected) +} + +#[cfg(test)] +mod tests { + use super::super::turn::ResponsesWebSocketTurnOutcome; + use super::{attempt_facts_for_outcome, settle_signal_for_client_delivery_failure}; + use crate::execution_runtime::attempt_lifecycle::{ + AttemptClientDelivery, AttemptProviderOutcome, AttemptTerminalFacts, + }; + + const fn terminal(status_code: u16) -> AttemptProviderOutcome { + AttemptProviderOutcome::Terminal { + status_code, + cancelled_by_provider: false, + } + } + + const fn provider_cancelled() -> AttemptProviderOutcome { + AttemptProviderOutcome::Terminal { + status_code: 499, + cancelled_by_provider: true, + } + } + + const fn aborted(status_code: u16, reason: &'static str) -> AttemptProviderOutcome { + AttemptProviderOutcome::Aborted { + status_code, + reason, + stream_timeout: status_code == 504, + } + } + + /// §1.6 现状 outcome → 双事实映射表,逐行。 + #[test] + fn every_settle_signal_maps_to_a_provider_outcome_and_a_client_delivery() { + assert_eq!( + attempt_facts_for_outcome( + None, + AttemptClientDelivery::Complete, + ResponsesWebSocketTurnOutcome::ProviderTerminal { + status_code: 200, + cancelled: false, + }, + ), + AttemptTerminalFacts { + provider: terminal(200), + delivery: AttemptClientDelivery::Complete, + } + ); + assert_eq!( + attempt_facts_for_outcome( + None, + AttemptClientDelivery::Complete, + ResponsesWebSocketTurnOutcome::ProviderTerminal { + status_code: 499, + cancelled: true, + }, + ), + AttemptTerminalFacts { + provider: provider_cancelled(), + delivery: AttemptClientDelivery::Complete, + } + ); + assert_eq!( + attempt_facts_for_outcome( + None, + AttemptClientDelivery::Complete, + ResponsesWebSocketTurnOutcome::upstream_closed() + ), + AttemptTerminalFacts { + provider: aborted( + 502, + "upstream WebSocket closed before provider terminal event" + ), + delivery: AttemptClientDelivery::Complete, + } + ); + assert_eq!( + attempt_facts_for_outcome( + None, + AttemptClientDelivery::Complete, + ResponsesWebSocketTurnOutcome::client_disconnected() + ), + AttemptTerminalFacts { + provider: aborted(499, "client disconnected before provider terminal event"), + delivery: AttemptClientDelivery::Aborted { + reason: "client disconnected before provider terminal event", + }, + } + ); + + // 超时一族必须保留 stream_timeout 标记,否则 pool stream timeout 效果丢失。 + let first_event_timeout = attempt_facts_for_outcome( + None, + AttemptClientDelivery::Complete, + ResponsesWebSocketTurnOutcome::first_event_timeout(), + ); + assert!(first_event_timeout.provider.stream_timeout()); + let terminal_timeout = attempt_facts_for_outcome( + None, + AttemptClientDelivery::Complete, + ResponsesWebSocketTurnOutcome::terminal_timeout(), + ); + assert!(terminal_timeout.provider.stream_timeout()); + // 非 504 的失败不得被当成流式超时。 + assert!(!attempt_facts_for_outcome( + None, + AttemptClientDelivery::Complete, + ResponsesWebSocketTurnOutcome::upstream_closed() + ) + .provider + .stream_timeout()); + // provider 终态即使状态码是 504 也不投射 stream timeout:现状 + // `stream_timeout()` 只匹配 Failure 分支。 + assert!(!attempt_facts_for_outcome( + None, + AttemptClientDelivery::Complete, + ResponsesWebSocketTurnOutcome::ProviderTerminal { + status_code: 504, + cancelled: false, + }, + ) + .provider + .stream_timeout()); + } + + /// `Cancelled` 不携带 provider 信息,已观察到的终态不能被它覆盖; + /// `ProviderTerminal` / `Failure` 本身就是权威的 provider 事实。 + #[test] + fn an_observed_provider_terminal_survives_a_client_side_cancellation() { + let observed = terminal(200); + + let facts = attempt_facts_for_outcome( + Some(observed), + AttemptClientDelivery::Complete, + ResponsesWebSocketTurnOutcome::client_disconnected(), + ); + assert_eq!(facts.provider, observed); + assert_eq!( + facts.delivery, + AttemptClientDelivery::Aborted { + reason: "client disconnected before provider terminal event", + } + ); + + // 权威信号不被已记录事实改写。 + let facts = attempt_facts_for_outcome( + Some(observed), + AttemptClientDelivery::Complete, + ResponsesWebSocketTurnOutcome::upstream_closed(), + ); + assert_eq!( + facts.provider, + aborted( + 502, + "upstream WebSocket closed before provider terminal event" + ) + ); + assert_eq!(facts.delivery, AttemptClientDelivery::Complete); + } + + /// 结算信号的选择:provider 终态已到达就用它,否则才是 client 断开。 + /// 这是修正的核心——旧实现无条件用 client_disconnected() 覆盖, + /// 于是已完成的响应被记成 void billing。 + #[test] + fn a_reached_terminal_is_the_settle_signal_for_a_delivery_failure() { + let terminal_outcome = ResponsesWebSocketTurnOutcome::ProviderTerminal { + status_code: 200, + cancelled: false, + }; + assert_eq!( + settle_signal_for_client_delivery_failure(Some(terminal_outcome)), + terminal_outcome + ); + assert_eq!( + settle_signal_for_client_delivery_failure(None), + ResponsesWebSocketTurnOutcome::client_disconnected() + ); + } + + /// 明确记录的投递失败不会被结算信号推出的「投递成功」覆盖。 + #[test] + fn a_recorded_delivery_failure_survives_a_provider_terminal_settle_signal() { + let facts = attempt_facts_for_outcome( + Some(terminal(200)), + AttemptClientDelivery::Aborted { + reason: "write failed", + }, + ResponsesWebSocketTurnOutcome::ProviderTerminal { + status_code: 200, + cancelled: false, + }, + ); + assert_eq!(facts.provider, terminal(200)); + assert_eq!( + facts.delivery, + AttemptClientDelivery::Aborted { + reason: "write failed" + } + ); + // 投递失败不是供应商的错误,摘要不该因此补 parser_error。 + assert_eq!(facts.forced_error(), 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..941948c13 --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/state.rs @@ -0,0 +1,106 @@ +//! 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 std::collections::{BTreeMap, BTreeSet}; +use tokio::task::JoinHandle; + +use super::adapter::{ResponsesWebSocketDrainDirective, ResponsesWebSocketProtocolAdapter}; +use super::binding::UpstreamBindingIdentity; +use super::redaction::ResponsesWebSocketRedactionRestorer; +use super::turn_state::ResponsesTurnState; +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) 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, + /// 这条连接上「有没有正在进行的 logical turn」的唯一事实来源。 + pub(super) turn_state: ResponsesTurnState, + /// 这条连接迄今 mask 出来的映射,用于把 provider 事件里的占位符换回真实值。 + /// + /// 刻意按连接持有而不是按 turn 持有:WS 的会话历史留在上游,continuation 只发 + /// 增量输入,所以后面几轮的响应可能回显更早那几轮的占位符(理由详见 + /// [`super::redaction`])。上游重绑时不重置。 + pub(super) redaction_restorer: ResponsesWebSocketRedactionRestorer, + 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); + } +} 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..7cb4f0135 --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn.rs @@ -0,0 +1,1375 @@ +//! 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 super::observation::ResponsesStructuredTerminalObserver; +use super::settlement::attempt_facts_for_outcome; +use crate::ai_serving::{build_openai_responses_stream_plan_from_decision, AiExecutionDecision}; +use crate::clock::current_unix_ms; +use crate::control::{ + execution_plan_balance_capacity_rejection, refresh_execution_runtime_auth_context, + request_model_local_rejection, GatewayControlDecision, GatewayLocalAuthRejection, +}; +use crate::execution_runtime::attach_provider_response_headers_to_report_context; +use crate::execution_runtime::attempt_lifecycle::{ + attempt_billing_is_void, AttemptBodyCapture, AttemptClientDelivery, AttemptLifecycleSeed, + AttemptProviderOutcome, AttemptStageGuard, AttemptTerminalFacts, AttemptTerminalFactsInput, + ExecutionAttemptLifecycle, +}; +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 WEBSOCKET_CLIENT_DELIVERY_REPORT_CONTEXT_FIELD: &str = "websocket_client_delivery"; +const WEBSOCKET_CLIENT_DELIVERY_ABORTED: &str = "aborted"; +const WEBSOCKET_CLIENT_DELIVERY_REASON_REPORT_CONTEXT_FIELD: &str = + "websocket_client_delivery_reason"; +/// 首个可计费事件到达时记在 usage/candidate 上的状态码。WS 的首事件本身不带 +/// HTTP 状态,沿用 HTTP 流式「已开始流」的 200。 +const STREAM_STARTED_STATUS_CODE: u16 = 200; +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, + } + } +} + +/// 一次 provider attempt:一条上游执行,也是一条独立的 usage/candidate 记录。 +/// +/// 与 [`super::turn_state::LogicalTurn`] 分工明确:logical turn 是客户端看到的 +/// 一轮请求(可能包含多个 attempt),attempt 只负责这一次上游执行的记账事实。 +pub(super) struct ResponsesProviderAttempt { + /// 记账三段(pending / started / terminal)由共享的 transport 中立生命周期负责。 + lifecycle: ExecutionAttemptLifecycle, + started_at: Instant, + provider_headers: BTreeMap, + observer: ResponsesStructuredTerminalObserver, + provider_capture: AttemptBodyCapture, + client_capture: AttemptBodyCapture, + upstream_bytes: u64, + first_event_elapsed_ms: Option, + first_event_timeout: Duration, + terminal_timeout: Duration, + admission: Option, + terminal_error_body: Option, + /// 观察到的 provider 终态事实,与「为什么现在结算」这个信号分开保存。 + /// 客户端投递失败不会把它擦掉。 + provider_outcome: Option, + /// 这一个 attempt 的内容是否完整交付给了客户端。与 provider 终态正交。 + client_delivery: AttemptClientDelivery, +} + +/// 组装一轮 turn 的 decision。 +/// +/// `effective_client_event` 必须是**已经过请求侧脱敏**的客户端事件(见 +/// `super::redaction`),`provider_event` 由它派生。审计里的 `original_request_body` +/// 直接用它覆盖 seed:continuation 的 seed 来自绑定那一轮的 report_context,不覆盖 +/// 会记成上一轮的 body;但覆盖成 raw 事件又会把已脱敏的审计内容换回原文,等于 +/// 脱敏在审计侧失效。 +pub(super) fn prepare_responses_websocket_turn_decision( + template: &AiExecutionDecision, + request_id: String, + reuse_selected_candidate: bool, + effective_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, + effective_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 = ExecutionAttemptLifecycle::begin( + state, + AttemptLifecycleSeed { + plan, + report_kind, + report_context, + // relay loop 是单任务:一段慢依赖会拖住整条连接的收发,所以每段 + // 记账 I/O 都要有等待上界。 + stage_guard: AttemptStageGuard::Bounded(RESPONSES_WEBSOCKET_LIFECYCLE_STAGE_TIMEOUT), + }, + ) + .await; + + Ok(ResponsesProviderAttempt { + lifecycle, + started_at: Instant::now(), + provider_headers: BTreeMap::new(), + observer: ResponsesStructuredTerminalObserver::default(), + provider_capture: AttemptBodyCapture::default(), + client_capture: AttemptBodyCapture::default(), + upstream_bytes: 0, + first_event_elapsed_ms: None, + first_event_timeout, + terminal_timeout, + admission: Some(admission), + terminal_error_body: None, + provider_outcome: None, + client_delivery: AttemptClientDelivery::Complete, + }) +} + +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 ResponsesProviderAttempt { + /// 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 _ = self + .lifecycle + .stage_guard() + .await_stage( + self.lifecycle.trace_id(), + "turn_admission_release", + admission.release(), + ) + .await; + } + } + + pub(super) fn set_provider_response_headers(&mut self, headers: BTreeMap) { + let report_context = attach_provider_response_headers_to_report_context( + self.lifecycle.take_report_context(), + &headers, + ); + self.lifecycle.set_report_context(report_context); + 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)); + } + + // 一帧可以带多个协议事件(批量帧),必须拆开:观测器按事件推进状态机, + // 整帧当一个事件喂会丢掉批量里最后那个 completed 的 usage。 + let events = frame.protocol_events(); + let mut report_context = self.lifecycle.take_report_context(); + for event in &events { + // 观测已经走结构化入口,但捕获仍然必须是 SSE 形状:usage runtime 按 + // `data:` 行解析被捕获的 body 判定终态。 + self.capture_sse_event(event); + adapter.decorate_turn_report_context(&mut report_context, event); + } + self.lifecycle.set_report_context(report_context); + let fallback_context = json!({ + "provider_api_format": "openai:responses", + "client_api_format": "openai:responses", + }); + let report_context = self.lifecycle.report_context().unwrap_or(&fallback_context); + self.observer.observe_events(report_context, &events); + + 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) { + // provider 的终态是独立事实:先记下来,之后即使客户端投递失败、 + // 结算信号变成 Cancelled,这条事实也不会被擦掉。 + self.provider_outcome.get_or_insert( + attempt_facts_for_outcome(None, AttemptClientDelivery::Complete, outcome).provider, + ); + 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", + }, + )) + } + + /// 记录「这一个 attempt 的内容没能完整交付给客户端」。 + /// + /// 与 provider 终态分开记录:供应商已经给出终态时,这条事实只影响 + /// candidate 的错误分类和审计里的投递标记,不作废账单。 + pub(super) fn record_client_delivery_aborted(&mut self, reason: &'static str) { + if matches!(self.client_delivery, AttemptClientDelivery::Complete) { + self.client_delivery = AttemptClientDelivery::Aborted { reason }; + } + } + + pub(super) fn capture_client_frame(&mut self, event: &Value) { + self.client_capture + .append(&websocket_event_as_sse_line(event)); + } + + pub(super) async fn mark_stream_started(&mut self, state: &AppState) { + let telemetry = self.telemetry(); + self.lifecycle + .mark_started(state, STREAM_STARTED_STATUS_CODE, &telemetry) + .await; + } + + /// 结算这一个 attempt。 + /// + /// `outcome` 是「为什么现在结算」的信号,不是供应商事实本身: + /// [`attempt_facts_for_outcome`] 把它和已观察到的 provider 终态、已记录的 + /// 投递结果一起,拆成 provider outcome 与 client delivery 两个正交事实。 + /// 之后的四段记账(usage terminal → candidate terminal → provider 效果 → + /// execution report)由共享的 [`ExecutionAttemptLifecycle::settle`] 负责, + /// 这里只提供 WS 观察到的终态事实。 + async fn settle(mut self, state: &AppState, outcome: ResponsesWebSocketTurnOutcome) { + let facts = attempt_facts_for_outcome(self.provider_outcome, self.client_delivery, outcome); + if let Some(reason) = facts.delivery.aborted_reason() { + let report_context = attach_client_delivery_to_report_context( + self.lifecycle.take_report_context(), + reason, + ); + self.lifecycle.set_report_context(report_context); + } + let summary = self.finish_summary(facts); + let telemetry = self.telemetry(); + let terminal_error_body = self.terminal_error_body.take(); + + // 终态载荷完整了才释放准入:usage/审计写入期间不再占着 gateway/供应商容量。 + if let Some(admission) = self.admission.take() { + let _ = self + .lifecycle + .stage_guard() + .await_stage( + self.lifecycle.trace_id(), + "turn_admission_release", + admission.release(), + ) + .await; + } + + self.lifecycle + .settle( + state, + AttemptTerminalFactsInput { + facts, + terminal_summary: summary, + telemetry, + provider_headers: std::mem::take(&mut self.provider_headers), + provider_body: &self.provider_capture, + client_body: &self.client_capture, + provider_error_body: terminal_error_body.as_deref(), + reason: facts.reason(), + }, + ) + .await; + } + + fn capture_sse_event(&mut self, event: &Value) { + self.provider_capture + .append(&websocket_event_as_sse_line(event)); + } + + /// 终态摘要。 + /// + /// 只消费两个正交事实:`forced_error` 只在「供应商没给出终态且内容已完整 + /// 交付」时补 parser_error,`cancelled` 覆盖供应商声明取消与客户端投递失败 + /// 两种情形——与拆分前 `outcome.forced_error()` / `outcome.cancelled()` 的 + /// 取值逐一对应。 + fn finish_summary(&mut self, facts: AttemptTerminalFacts) -> ExecutionStreamTerminalSummary { + let fallback_context = json!({ + "provider_api_format": "openai:responses", + "client_api_format": "openai:responses", + }); + let report_context = self.lifecycle.report_context().unwrap_or(&fallback_context); + let mut summary = self.observer.finish(report_context); + if let Some(reason) = facts.forced_error() { + if summary.parser_error.is_none() { + summary.parser_error = Some(reason.to_string()); + } + } + // 只有作废账单的那一侧才把摘要改写成 cancelled。provider 终态已到达时 + // 摘要必须保留真实的 finish_reason 和 usage,否则计费记录会被写坏。 + if attempt_billing_is_void(facts) { + 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 ResponsesProviderAttempt { + /// 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.settle(state, outcome).await; + } +} + +pub(super) async fn spawn_responses_websocket_turn_finalization( + state: AppState, + mut turn: ResponsesProviderAttempt, + outcome: ResponsesWebSocketTurnOutcome, +) -> tokio::task::JoinHandle<()> { + turn.release_admission().await; + tokio::spawn(async move { + turn.settle(&state, outcome).await; + }) +} + +/// 把一轮 turn 的事实写进审计/用量 report context。 +/// +/// `effective_client_event` 是脱敏后的客户端事件(未启用脱敏时就是原事件)。 +/// HTTP 路径的约定是「脱敏生效时审计记录脱敏后的 body」 +/// (`ai_serving/planner/standard/openai/responses/decision/payload.rs`), +/// WS 这里必须保持一致,否则上游收到的是脱敏内容、审计里却留着原始 PII。 +fn prepare_websocket_report_context( + report_context: Option, + request_id: &str, + reuse_selected_candidate: bool, + effective_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(), + effective_client_event.clone(), + ); + if let Some(model) = effective_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) +} + +/// 在审计/用量 report context 上标记这一 attempt 的内容没能交付给客户端。 +/// +/// 只增字段,不改既有字段:账单本身按 provider 终态计,投递失败作为独立事实 +/// 留在记录里,便于事后区分「客户端拿到了」和「客户端没拿到但已计费」。 +fn attach_client_delivery_to_report_context( + report_context: Option, + reason: &str, +) -> Option { + 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( + WEBSOCKET_CLIENT_DELIVERY_REPORT_CONTEXT_FIELD.to_string(), + Value::String(WEBSOCKET_CLIENT_DELIVERY_ABORTED.to_string()), + ); + object.insert( + WEBSOCKET_CLIENT_DELIVERY_REASON_REPORT_CONTEXT_FIELD.to_string(), + Value::String(reason.to_string()), + ); + Some(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 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 elapsed_ms(started_at: Instant) -> u64 { + started_at.elapsed().as_millis().min(u128::from(u64::MAX)) as u64 +} + +#[cfg(test)] +mod tests { + use std::time::{Duration, Instant}; + + use aether_contracts::ExecutionTimeouts; + use serde_json::json; + + use super::super::observation::ResponsesStructuredTerminalObserver; + + use super::super::frame::ParsedResponsesWebSocketFrame; + use super::super::settlement::{ + attempt_facts_for_outcome, settle_signal_for_client_delivery_failure, + }; + use super::{ + attach_client_delivery_to_report_context, prepare_websocket_report_context, + provider_terminal_outcome, resolve_responses_websocket_turn_timeouts, + websocket_event_as_sse_line, ResponsesWebSocketTurnDeadline, ResponsesWebSocketTurnOutcome, + ResponsesWebSocketTurnTimeoutPhase, + }; + use crate::execution_runtime::attempt_lifecycle::{ + classify_attempt_settlement, AttemptBilling, AttemptCandidateError, AttemptCandidateStatus, + AttemptClientDelivery, AttemptProviderEffect, AttemptSettlementInputs, + }; + + #[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 = ResponsesStructuredTerminalObserver::default(); + observer.observe_events(&report_context, &[&event]); + let summary = observer.finish(&report_context); + 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 a_legitimate_incomplete_is_a_successful_provider_terminal_that_keeps_its_usage() { + // 写满 max_output_tokens 的 incomplete 是合法终态:状态码不再是 502, + // usage 观测器照样能看到 finish 和 token,记账层不该把它当解析失败。 + let event = json!({ + "type": "response.incomplete", + "response": { + "id": "resp_ws_incomplete_123", + "model": "gpt-5.6", + "status": "incomplete", + "incomplete_details": {"reason": "max_output_tokens"}, + "output": [], + "usage": { + "input_tokens": 4, + "output_tokens": 7, + "total_tokens": 11 + } + } + }); + 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 outcome = outcome.expect("incomplete should end the turn"); + let facts = attempt_facts_for_outcome(None, AttemptClientDelivery::Complete, outcome); + assert!(!facts.provider.cancelled_by_provider()); + assert!(facts.forced_error().is_none()); + assert!(!facts.provider.stream_timeout()); + assert!(facts.provider.is_terminal()); + + let report_context = json!({ + "provider_api_format": "openai:responses", + "client_api_format": "openai:responses", + }); + let mut observer = ResponsesStructuredTerminalObserver::default(); + observer.observe_events(&report_context, &[&event]); + let summary = observer.finish(&report_context); + assert!(summary.observed_finish); + assert_eq!(summary.finish_reason.as_deref(), Some("length")); + assert!(summary.parser_error.is_none()); + let usage = summary + .standardized_usage + .expect("incomplete usage must reach the terminal summary"); + assert_eq!(usage.input_tokens, 4); + assert_eq!(usage.output_tokens, 7); + + // 结算表用这些事实决定是否投射供应商失败:合法 incomplete 即使被记账层 + // 判成失败,也必须落在「不扣健康分、只释放 lease」一侧,且账单照记。 + let settlement = classify_attempt_settlement(AttemptSettlementInputs { + facts, + report_represents_failure: true, + observed_finish: summary.observed_finish, + has_parser_error: summary.parser_error.is_some(), + }); + assert_eq!(settlement.status_code, 200); + assert_eq!(settlement.billing, AttemptBilling::Billed); + assert_eq!( + settlement.provider_effect, + AttemptProviderEffect::ReleasePoolKeyLease + ); + assert!(settlement.submit_execution_report); + } + + #[test] + fn an_incomplete_without_a_legitimate_reason_still_projects_a_provider_failure() { + let raw = r#"{"type":"response.incomplete","response":{"incomplete_details":{"reason":"error"}}}"#; + let frame = ParsedResponsesWebSocketFrame::parse(raw).expect("event should parse"); + + assert_eq!( + provider_terminal_outcome(&frame), + Some(ResponsesWebSocketTurnOutcome::ProviderTerminal { + status_code: 502, + cancelled: false + }) + ); + } + + /// relay 级:provider 终态帧已经到达,但客户端 socket 已经关闭。 + /// + /// 走的是 relay loop 写客户端失败时的完整决策链——真实帧解析 → + /// 记录 provider 事实 → 记录投递失败 → 选结算信号 → 结算表。旧实现在这里 + /// 用 client_disconnected() 覆盖 outcome,于是一条已经产出 token 的响应被 + /// 记成 void billing、不投射供应商效果、也不提交 execution report。 + #[test] + fn a_provider_terminal_that_reaches_a_closed_client_socket_is_still_billed() { + let completed = json!({ + "type": "response.completed", + "response": { + "id": "resp_ws_delivery_failed", + "model": "gpt-5.6", + "usage": {"input_tokens": 3, "output_tokens": 5, "total_tokens": 8} + } + }); + let raw = serde_json::to_string(&completed).expect("event should serialize"); + let frame = ParsedResponsesWebSocketFrame::parse(&raw).expect("event should parse"); + + // relay loop 观察到终态帧:attempt 记下 provider 事实。 + let observed = provider_terminal_outcome(&frame).expect("completed ends the turn"); + let recorded_provider = + attempt_facts_for_outcome(None, AttemptClientDelivery::Complete, observed).provider; + + // 随后写客户端失败:记录投递失败,并按 provider 终态选结算信号。 + let delivery = AttemptClientDelivery::Aborted { + reason: "gateway could not relay the provider event to the client", + }; + let signal = settle_signal_for_client_delivery_failure(Some(observed)); + assert_eq!( + signal, observed, + "a reached terminal must remain the signal" + ); + + let facts = attempt_facts_for_outcome(Some(recorded_provider), delivery, signal); + + // 终态摘要保留真实 usage 与 finish_reason,不被改写成 cancelled。 + let report_context = json!({ + "provider_api_format": "openai:responses", + "client_api_format": "openai:responses", + }); + let mut observer = ResponsesStructuredTerminalObserver::default(); + observer.observe_events(&report_context, &[&completed]); + let summary = observer.finish(&report_context); + assert!(summary.observed_finish); + assert!(summary.parser_error.is_none()); + let usage = summary + .standardized_usage + .clone() + .expect("usage must survive a delivery failure"); + assert_eq!(usage.input_tokens, 3); + assert_eq!(usage.output_tokens, 5); + + let settlement = classify_attempt_settlement(AttemptSettlementInputs { + facts, + report_represents_failure: false, + observed_finish: summary.observed_finish, + has_parser_error: summary.parser_error.is_some(), + }); + assert_eq!(settlement.billing, AttemptBilling::Billed); + assert_eq!(settlement.status_code, 200); + assert_eq!(settlement.candidate_status, AttemptCandidateStatus::Success); + assert_eq!( + settlement.provider_effect, + AttemptProviderEffect::ProviderSuccess + ); + assert!(settlement.submit_execution_report); + // 投递失败仍然留痕。 + assert_eq!( + settlement.candidate_error, + AttemptCandidateError::ClientDeliveryFailed + ); + } + + /// 同一条链路,但供应商还没给出终态:仍然作废账单、不提交 report。 + #[test] + fn a_closed_client_socket_before_any_terminal_still_voids_the_bill() { + let signal = settle_signal_for_client_delivery_failure(None); + let facts = attempt_facts_for_outcome( + None, + AttemptClientDelivery::Aborted { + reason: "gateway could not relay the provider event to the client", + }, + signal, + ); + let settlement = classify_attempt_settlement(AttemptSettlementInputs { + facts, + report_represents_failure: false, + observed_finish: false, + has_parser_error: false, + }); + + assert_eq!(settlement.billing, AttemptBilling::Void); + assert_eq!(settlement.status_code, 499); + assert_eq!( + settlement.candidate_status, + AttemptCandidateStatus::Cancelled + ); + assert_eq!( + settlement.provider_effect, + AttemptProviderEffect::ReleasePoolKeyLease + ); + assert!(!settlement.submit_execution_report); + } + + /// 投递失败只往 report context 里加字段,不改既有字段。 + #[test] + fn the_client_delivery_marker_only_adds_report_context_fields() { + let context = attach_client_delivery_to_report_context( + Some(json!({ + "request_id": "turn-2", + "websocket_mode": true, + "original_request_body": {"model": "public"}, + })), + "gateway could not relay the provider event to the client", + ) + .expect("marker should produce a report context"); + + assert_eq!(context["websocket_client_delivery"], "aborted"); + assert_eq!( + context["websocket_client_delivery_reason"], + "gateway could not relay the provider event to the client" + ); + assert_eq!(context["request_id"], "turn-2"); + assert_eq!(context["websocket_mode"], true); + assert_eq!(context["original_request_body"]["model"], "public"); + } + + #[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(); + let facts = attempt_facts_for_outcome(None, AttemptClientDelivery::Complete, outcome); + let settlement = classify_attempt_settlement(AttemptSettlementInputs { + facts, + report_represents_failure: true, + observed_finish: false, + has_parser_error: false, + }); + + assert_eq!(settlement.status_code, 500); + assert_eq!(settlement.billing, AttemptBilling::Billed); + assert!(!facts.provider.cancelled_by_provider()); + assert!(facts.forced_error().is_some()); + assert!(settlement.submit_execution_report); + } + + #[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/turn_state.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn_state.rs new file mode 100644 index 000000000..4a4838818 --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn_state.rs @@ -0,0 +1,342 @@ +//! 一条 Responses WebSocket 连接上的 turn 状态机。 +//! +//! 现状把「有没有正在进行的 turn」拆成 `response_in_flight`、`active_turn`、 +//! `active_response_create` 三个可独立变化的字段,8 种组合里只有 3 种合法, +//! 非法组合只能靠调用点的 if 和「记得同时改另外两个字段」来避免。这里把它收敛成 +//! 一个枚举:合法组合由类型保证,转换只能走受控 API。 + +use serde_json::Value; + +use super::lifecycle::ActiveProviderAttempt; +use super::request::response_create_has_previous_response_id; + +/// 客户端一次 `response.create` 对应的 logical turn。 +/// +/// 一个 logical turn 可能经历多个 provider attempt:配额透明重试会换一把 key、 +/// 换一条上游连接重放同一份客户端事件,但对客户端始终是同一轮请求。 +/// `client_event` 保存的必须是**已脱敏**的事件(见 `super::redaction`), +/// 因为透明重试直接重放它。 +#[derive(Debug, Clone)] +pub(super) struct LogicalTurn { + 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 LogicalTurn { + 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); + } +} + +/// 连接上「有没有正在进行的 logical turn」这一唯一事实。 +/// +/// 类型参数只为测试留出注入点:生产代码一律用默认的 +/// [`ActiveProviderAttempt`],测试用轻量替身驱动同一套转换逻辑, +/// 不必构造 `AppState` 和真实 socket。 +pub(super) enum ResponsesTurnState { + /// 没有进行中的 logical turn。上游可能仍绑定,也可能已被 detach。 + Idle, + /// 一个 logical turn 正在等待 provider 终态:logical 与当前 attempt 同时存在。 + Responding { logical: LogicalTurn, attempt: A }, + /// logical turn 仍在,但当前 attempt 已被取走去结算或重绑,新 attempt 未就位。 + /// 配额透明重试期间就处于这个状态。 + Replanning { logical: LogicalTurn }, +} + +impl ResponsesTurnState { + /// 上游是否有一个正在进行的 response。取代原来的 `response_in_flight` 字段。 + pub(super) const fn response_in_flight(&self) -> bool { + matches!(self, Self::Responding { .. }) + } + + /// 是否可以接受一条新的客户端 `response.create`。 + /// + /// `Replanning` 也要拒绝:那时旧 attempt 的结算/重绑还没收尾。实际上 + /// 透明重试整段都在 relay loop 的上游分支里同步完成,此时不会读客户端帧, + /// 所以这条相对原来的 `response_in_flight` 判断没有行为差异。 + pub(super) const fn accepts_new_response_create(&self) -> bool { + matches!(self, Self::Idle) + } + + pub(super) const fn logical(&self) -> Option<&LogicalTurn> { + match self { + Self::Idle => None, + Self::Responding { logical, .. } | Self::Replanning { logical } => Some(logical), + } + } + + pub(super) fn logical_mut(&mut self) -> Option<&mut LogicalTurn> { + match self { + Self::Idle => None, + Self::Responding { logical, .. } | Self::Replanning { logical } => Some(logical), + } + } + + pub(super) const fn attempt(&self) -> Option<&A> { + match self { + Self::Responding { attempt, .. } => Some(attempt), + Self::Idle | Self::Replanning { .. } => None, + } + } + + pub(super) fn attempt_mut(&mut self) -> Option<&mut A> { + match self { + Self::Responding { attempt, .. } => Some(attempt), + Self::Idle | Self::Replanning { .. } => None, + } + } + + /// `Idle` → `Responding`:装上一个新 logical turn 及其首个 attempt。 + /// + /// 与原来的 `active_turn = Some(..)` 语义一致:若此刻竟持有旧 attempt, + /// 它会被丢弃并由自身的 drop guard 兜底结算,而不是静默泄漏。 + pub(super) fn begin(&mut self, logical: LogicalTurn, attempt: A) { + debug_assert!( + self.accepts_new_response_create(), + "a new logical turn must only begin on an idle connection" + ); + *self = Self::Responding { logical, attempt }; + } + + /// `Responding` → `Replanning`:把当前 attempt 交给调用方结算,保留 logical turn。 + pub(super) fn detach_attempt(&mut self) -> Option { + match std::mem::replace(self, Self::Idle) { + Self::Responding { logical, attempt } => { + *self = Self::Replanning { logical }; + Some(attempt) + } + state @ (Self::Idle | Self::Replanning { .. }) => { + *self = state; + None + } + } + } + + /// `Replanning` → `Responding`:同一 logical turn 的下一个 attempt 就位。 + /// + /// 状态不符时把 attempt 交还调用方,避免静默丢弃一条已经写了 pending usage + /// 行、占着 candidate 和 pool key lease 的 attempt。 + pub(super) fn resume(&mut self, attempt: A) -> Result<(), A> { + match std::mem::replace(self, Self::Idle) { + Self::Replanning { logical } => { + *self = Self::Responding { logical, attempt }; + Ok(()) + } + state => { + *self = state; + Err(attempt) + } + } + } + + /// `Responding`/`Replanning` → `Idle`:logical turn 结束,交出待结算的 attempt。 + /// + /// 取代原来「`active_turn.take()` + 在每个出口手写 `active_response_create = None`」 + /// 的组合:清理只有这一个出口,漏清不再可能。 + pub(super) fn end(&mut self) -> Option { + match std::mem::replace(self, Self::Idle) { + Self::Responding { attempt, .. } => Some(attempt), + Self::Idle | Self::Replanning { .. } => None, + } + } +} + +impl ResponsesTurnState { + /// 记录「当前 attempt 的内容没能完整交付给客户端」。 + /// + /// 这条事实写在 attempt 上而不是 logical turn 上:结算是按 attempt 进行的, + /// 而每个 attempt 的投递结果各自独立(配额透明重试时旧 attempt 可能已经把 + /// 部分事件交付出去,新 attempt 从零开始)。 + pub(super) fn record_client_delivery_aborted(&mut self, reason: &'static str) { + if let Some(attempt) = self.attempt_mut() { + attempt.record_client_delivery_aborted(reason); + } + } +} + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::{LogicalTurn, ResponsesTurnState}; + + /// attempt 的测试替身:只需要能被 move,不需要 AppState 或真实 socket。 + #[derive(Debug, PartialEq, Eq)] + struct FakeAttempt(u32); + + fn logical() -> LogicalTurn { + LogicalTurn::new( + json!({"type": "response.create", "model": "gpt-5.6-sol"}), + 7, + "logical-turn".to_string(), + ) + } + + /// 透明重试失败之后:旧 attempt 已经被 detach 并结算过,logical turn 仍停在 + /// `Replanning`。此时 `end()` 不能再交出 attempt,否则同一个 attempt 会被 + /// 结算两次(两条 usage terminal、两次 pool lease 释放)。 + #[test] + fn ending_a_replanning_turn_does_not_hand_out_a_second_attempt() { + let mut state = ResponsesTurnState::Idle; + state.begin(logical(), FakeAttempt(1)); + + let detached = state.detach_attempt(); + assert_eq!( + detached, + Some(FakeAttempt(1)), + "the attempt is settled once" + ); + assert!(matches!(state, ResponsesTurnState::Replanning { .. })); + + // 结算已经发生,没有第二个 attempt 可交。 + assert!( + state.end().is_none(), + "a replanning turn must not yield a second attempt to settle" + ); + assert!(state.accepts_new_response_create()); + } + + #[test] + fn idle_has_no_turn_and_accepts_a_new_response_create() { + let state = ResponsesTurnState::::Idle; + + assert!(!state.response_in_flight()); + assert!(state.accepts_new_response_create()); + assert!(state.logical().is_none()); + assert!(state.attempt().is_none()); + } + + #[test] + fn beginning_a_turn_makes_the_response_in_flight_and_blocks_a_second_one() { + let mut state = ResponsesTurnState::Idle; + state.begin(logical(), FakeAttempt(1)); + + assert!(state.response_in_flight()); + assert!(!state.accepts_new_response_create()); + assert_eq!(state.logical().map(|logical| logical.turn_index), Some(7)); + assert_eq!(state.attempt(), Some(&FakeAttempt(1))); + } + + #[test] + fn detaching_an_attempt_keeps_the_logical_turn_but_ends_the_in_flight_response() { + let mut state = ResponsesTurnState::Idle; + state.begin(logical(), FakeAttempt(1)); + + assert_eq!(state.detach_attempt(), Some(FakeAttempt(1))); + assert!(!state.response_in_flight()); + assert!(!state.accepts_new_response_create()); + assert!(state.attempt().is_none()); + assert_eq!( + state + .logical() + .map(|logical| logical.logical_turn_id.clone()), + Some("logical-turn".to_string()) + ); + // 幂等:已经没有 attempt 了,再取一次不会伪造一个出来。 + assert_eq!(state.detach_attempt(), None); + } + + #[test] + fn resuming_replaces_the_attempt_of_the_same_logical_turn() { + let mut state = ResponsesTurnState::Idle; + state.begin(logical(), FakeAttempt(1)); + state + .logical_mut() + .expect("a responding turn always has its logical turn") + .turn_attempt = 2; + let _ = state.detach_attempt(); + + assert_eq!(state.resume(FakeAttempt(2)), Ok(())); + assert!(state.response_in_flight()); + assert_eq!(state.attempt(), Some(&FakeAttempt(2))); + assert_eq!(state.logical().map(|logical| logical.turn_attempt), Some(2)); + } + + #[test] + fn resuming_without_a_logical_turn_hands_the_attempt_back() { + let mut state = ResponsesTurnState::Idle; + + assert_eq!(state.resume(FakeAttempt(9)), Err(FakeAttempt(9))); + assert!(state.accepts_new_response_create()); + + state.begin(logical(), FakeAttempt(1)); + assert_eq!(state.resume(FakeAttempt(9)), Err(FakeAttempt(9))); + assert_eq!(state.attempt(), Some(&FakeAttempt(1))); + } + + #[test] + fn ending_a_turn_clears_both_the_logical_turn_and_the_attempt() { + let mut state = ResponsesTurnState::Idle; + state.begin(logical(), FakeAttempt(1)); + + assert_eq!(state.end(), Some(FakeAttempt(1))); + assert!(state.accepts_new_response_create()); + assert!(state.logical().is_none()); + + // 从 Replanning 结束时没有 attempt 要交出,但 logical turn 同样必须清掉。 + state.begin(logical(), FakeAttempt(2)); + let _ = state.detach_attempt(); + assert_eq!(state.end(), None); + assert!(state.accepts_new_response_create()); + assert!(state.logical().is_none()); + } + + #[test] + fn quota_retry_safety_lives_on_the_logical_turn() { + let mut state = ResponsesTurnState::Idle; + state.begin(logical(), FakeAttempt(1)); + + assert_eq!( + state + .logical() + .and_then(LogicalTurn::quota_retry_block_reason), + None + ); + state + .logical_mut() + .expect("logical turn") + .mark_retry_unsafe("standard_response_event"); + assert_eq!( + state + .logical() + .and_then(LogicalTurn::quota_retry_block_reason), + Some("standard_response_event") + ); + // 重绑后仍是同一个 logical turn,重放安全结论不能被 attempt 轮换洗掉。 + let _ = state.detach_attempt(); + assert_eq!(state.resume(FakeAttempt(2)), Ok(())); + assert_eq!( + state + .logical() + .and_then(LogicalTurn::quota_retry_block_reason), + Some("standard_response_event") + ); + } +} 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..59753b176 --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/upstream.rs @@ -0,0 +1,324 @@ +//! Physical upstream WebSocket binding and transport helpers. + +use std::time::Duration; + +use serde_json::Value; +use wreq::ws::message::Message as WreqWsMessage; + +use super::adapter::ResponsesWebSocketProtocolAdapter; +use super::binding::{UpstreamBindingIdentity, UpstreamBindingIdentityError}; +use super::redaction::ResponsesWebSocketRedactionRestorer; +use super::request::planned_response_create_event; +use super::state::{BoundResponsesConnection, ExhaustedResponsesWebSocketExclusions}; +use super::turn_state::ResponsesTurnState; +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, +}; + +/// 上游 WebSocket 握手的默认绝对 deadline(30 秒)。 +/// 覆盖 DNS → TCP connect → TLS → HTTP 101 Upgrade → 发送首条 event 的完整链路。 +/// 如果 decision 配置了更短的 first_byte_ms 或 total_ms,取其与此值的较小者。 +const DEFAULT_UPSTREAM_HANDSHAKE_DEADLINE_MS: u64 = 30_000; + +/// 从 decision.timeouts 推导实际 handshake 绝对 deadline。 +/// 取 first_byte_ms / total_ms / DEFAULT 三者中的最小正值。 +pub(super) fn resolve_upstream_handshake_deadline(decision: &AiExecutionDecision) -> Duration { + let mut deadline_ms = DEFAULT_UPSTREAM_HANDSHAKE_DEADLINE_MS; + if let Some(timeouts) = decision.timeouts.as_ref() { + if let Some(first_byte_ms) = timeouts.first_byte_ms.filter(|v| *v > 0) { + deadline_ms = deadline_ms.min(first_byte_ms); + } + if let Some(total_ms) = timeouts.total_ms.filter(|v| *v > 0) { + deadline_ms = deadline_ms.min(total_ms); + } + } + Duration::from_millis(deadline_ms) +} + +pub(super) async fn bind_responses_upstream( + decision: &AiExecutionDecision, + normalization: ResponsesWebSocketBodyNormalization, + initial_event: &Value, + adapter: &'static dyn ResponsesWebSocketProtocolAdapter, +) -> Result { + // 绝对 deadline:从此刻起必须在限定时间内完成握手 + 首条事件发送, + // 防止慢 TLS / 慢 HTTP Upgrade 无限占用 connection permit。 + let handshake_deadline = resolve_upstream_handshake_deadline(decision); + tokio::time::timeout( + handshake_deadline, + bind_responses_upstream_inner(decision, normalization, initial_event, adapter), + ) + .await + .map_err(|_| "responses_websocket_upstream_handshake_timeout")? +} + +/// 实际执行握手 + 首条事件发送的内部函数,由外层 timeout 包裹。 +async fn bind_responses_upstream_inner( + 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, + decision_template: decision.clone(), + body_normalization: normalization, + binding_identity, + // 首条 response.create 已经发出,但这一轮的 logical turn 和 attempt 由调用方 + // 通过 `ResponsesTurnState::begin` 装上:绑定本身不持有记账状态。 + turn_state: ResponsesTurnState::Idle, + // 同理,这一轮的 mask session 也由调用方登记:绑定看不到脱敏链路。 + redaction_restorer: ResponsesWebSocketRedactionRestorer::default(), + 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) +} + +#[cfg(test)] +mod tests { + use std::time::Duration; + + use aether_contracts::ExecutionTimeouts; + + use crate::ai_serving::AiExecutionDecision; + + use super::{resolve_upstream_handshake_deadline, DEFAULT_UPSTREAM_HANDSHAKE_DEADLINE_MS}; + + 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: std::collections::BTreeMap::new(), + provider_request_headers: std::collections::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, + } + } + + #[test] + fn handshake_deadline_defaults_to_30s_without_configured_timeouts() { + let decision = sample_decision(); + let deadline = resolve_upstream_handshake_deadline(&decision); + assert_eq!( + deadline, + Duration::from_millis(DEFAULT_UPSTREAM_HANDSHAKE_DEADLINE_MS) + ); + } + + #[test] + fn handshake_deadline_uses_first_byte_ms_when_shorter_than_default() { + let mut decision = sample_decision(); + decision.timeouts = Some(ExecutionTimeouts { + first_byte_ms: Some(10_000), + total_ms: Some(60_000), + ..ExecutionTimeouts::default() + }); + let deadline = resolve_upstream_handshake_deadline(&decision); + assert_eq!(deadline, Duration::from_millis(10_000)); + } + + #[test] + fn handshake_deadline_uses_total_ms_when_shorter_than_first_byte_and_default() { + let mut decision = sample_decision(); + decision.timeouts = Some(ExecutionTimeouts { + first_byte_ms: Some(25_000), + total_ms: Some(8_000), + ..ExecutionTimeouts::default() + }); + let deadline = resolve_upstream_handshake_deadline(&decision); + assert_eq!(deadline, Duration::from_millis(8_000)); + } + + #[test] + fn handshake_deadline_ignores_zero_values() { + let mut decision = sample_decision(); + decision.timeouts = Some(ExecutionTimeouts { + first_byte_ms: Some(0), + total_ms: Some(0), + ..ExecutionTimeouts::default() + }); + let deadline = resolve_upstream_handshake_deadline(&decision); + assert_eq!( + deadline, + Duration::from_millis(DEFAULT_UPSTREAM_HANDSHAKE_DEADLINE_MS) + ); + } + + #[test] + fn handshake_deadline_does_not_exceed_default_even_with_larger_configured_values() { + let mut decision = sample_decision(); + decision.timeouts = Some(ExecutionTimeouts { + first_byte_ms: Some(120_000), + total_ms: Some(600_000), + ..ExecutionTimeouts::default() + }); + let deadline = resolve_upstream_handshake_deadline(&decision); + assert_eq!( + deadline, + Duration::from_millis(DEFAULT_UPSTREAM_HANDSHAKE_DEADLINE_MS) + ); + } + + #[tokio::test] + async fn bind_responses_upstream_times_out_against_stalled_server() { + use super::bind_responses_upstream; + use crate::ai_serving::ResponsesWebSocketBodyNormalization; + use crate::handlers::proxy::websocket::responses::adapter::resolve_responses_websocket_adapter; + use serde_json::json; + + // 启动一个接受 TCP 连接但永不完成 HTTP Upgrade 的 mock 服务器 + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("mock listener should bind"); + let addr = listener.local_addr().expect("should have local addr"); + let _server = tokio::spawn(async move { + loop { + let (socket, _) = listener.accept().await.unwrap(); + // 接受连接但不发送任何 HTTP 响应,模拟 stalled handshake + tokio::spawn(async move { + let _hold = socket; + tokio::time::sleep(Duration::from_secs(300)).await; + }); + } + }); + + let mut decision = sample_decision(); + decision.upstream_url = Some(format!("http://{addr}/v1/responses")); + // 设置极短的 deadline 以便测试快速完成 + decision.timeouts = Some(ExecutionTimeouts { + first_byte_ms: Some(100), + total_ms: Some(200), + ..ExecutionTimeouts::default() + }); + decision.provider_request_body = Some(json!({"model": "test-model"})); + + let adapter = resolve_responses_websocket_adapter( + crate::orchestration::ResponsesWebSocketAdapter::Standard, + ); + let result = bind_responses_upstream( + &decision, + ResponsesWebSocketBodyNormalization::for_tests("test-model"), + &json!({"type": "response.create", "model": "test-model"}), + adapter, + ) + .await; + + assert_eq!( + result.err().expect("bind should fail with timeout"), + "responses_websocket_upstream_handshake_timeout" + ); + } +} 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 13e95d533..45298f944 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 8e4f99d27..24a2af88e 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; @@ -255,6 +256,41 @@ pub(crate) enum LocalExecutionEffect<'a> { PoolStreamTimeout, } +/// 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, @@ -304,36 +340,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" ); } } @@ -1805,18 +1962,20 @@ 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, LocalPoolErrorEffect, - ProviderKeyEffectLockPool, + LocalStreamFailureEffect, ProviderKeyEffectLockPool, }; use crate::data::{GatewayDataConfig, GatewayDataState}; use crate::orchestration::LocalFailoverClassification; 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, @@ -1867,6 +2026,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(); @@ -2664,6 +2839,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 f79b25e8e..fec8655c6 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,18 @@ 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, LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, - LocalAttemptFailureEffect, LocalExecutionEffect, LocalExecutionEffectContext, - LocalHealthFailureEffect, LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, - LocalPoolErrorEffect, + apply_local_execution_effect, apply_local_stream_failure_effects, + apply_local_stream_success_effects, release_local_pool_key_lease, + release_pool_key_lease_from_report_context, LocalAdaptiveRateLimitEffect, + LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect, LocalExecutionEffect, + LocalExecutionEffectContext, LocalHealthFailureEffect, LocalHealthSuccessEffect, + LocalOAuthInvalidationEffect, LocalPoolErrorEffect, LocalStreamFailureEffect, }; pub(crate) use self::health::{ project_local_failure_health, project_local_key_circuit_closed, @@ -48,8 +56,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 +68,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 8a5b9f166..baf1ebd02 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 { @@ -273,6 +274,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, @@ -344,7 +402,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, @@ -549,4 +609,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 9fc72c643..17ac50789 100644 --- a/apps/aether-gateway/src/orchestration/report_effects.rs +++ b/apps/aether-gateway/src/orchestration/report_effects.rs @@ -16,6 +16,9 @@ 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; @@ -24,6 +27,7 @@ use crate::{AppState, GatewayError}; const CODEX_QUOTA_CACHE_TTL_SECONDS: u64 = 30; const CODEX_QUOTA_CACHE_MAX_ENTRIES: usize = 4096; const RUNTIME_METADATA_CAS_MAX_ATTEMPTS: usize = 16; +const CODEX_WEBSOCKET_RATE_LIMITS_REPORT_CONTEXT_FIELD: &str = "codex_websocket_rate_limits"; type HeaderFingerprintCache = Mutex>; @@ -91,6 +95,41 @@ 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 codex_quota_snapshot_matches_metadata(status_snapshot: Option<&Value>, parsed: &Value) -> bool { + let expected_allowed = parsed + .get("allowed") + .and_then(admin_provider_quota_pure::coerce_json_bool); + let expected_limit_reached = parsed + .get("limit_reached") + .and_then(admin_provider_quota_pure::coerce_json_bool); + if expected_allowed.is_none() && expected_limit_reached.is_none() { + return true; + } + let Some(quota) = status_snapshot + .and_then(|snapshot| snapshot.get("quota")) + .and_then(Value::as_object) + else { + return false; + }; + let expected_exhausted = admin_provider_quota_pure::codex_rate_limit_metadata_exhausted(parsed); + quota.get("exhausted").and_then(Value::as_bool) == Some(expected_exhausted) + && quota + .get("allowed") + .and_then(admin_provider_quota_pure::coerce_json_bool) + == expected_allowed + && quota + .get("limit_reached") + .and_then(admin_provider_quota_pure::coerce_json_bool) + == expected_limit_reached +} + fn is_volatile_compare_field(key: &str) -> bool { key == "updated_at" || key.ends_with("_reset_seconds") || key.ends_with("_reset_after_seconds") } @@ -126,6 +165,54 @@ fn fingerprint_codex_payload(value: &Value) -> Option { serde_json::to_string(&Value::Object(normalized)).ok() } +/// Reject an out-of-order WebSocket snapshot before the catalog CAS. The +/// provider emits `used_percent=100` immediately before a terminal quota +/// error; a delayed pre-terminal `99` frame must never roll that state back. +fn codex_snapshot_regresses(current: &Value, incoming: &Value) -> bool { + let current_exhausted = admin_provider_quota_pure::codex_rate_limit_metadata_exhausted(current); + let incoming_exhausted = + admin_provider_quota_pure::codex_rate_limit_metadata_exhausted(incoming); + let current_reset = current + .get("primary_reset_at") + .and_then(admin_provider_quota_pure::coerce_json_u64); + let incoming_reset = incoming + .get("primary_reset_at") + .and_then(admin_provider_quota_pure::coerce_json_u64); + + if current_exhausted && !incoming_exhausted { + // A lower reset timestamp denotes a newly opened window; otherwise a + // non-exhausted snapshot is stale or incomplete. + if incoming_reset.is_none() || incoming_reset >= current_reset { + return true; + } + } + + if current_reset == incoming_reset { + let current_used = current + .get("primary_used_percent") + .and_then(admin_provider_quota_pure::coerce_json_f64); + let incoming_used = incoming + .get("primary_used_percent") + .and_then(admin_provider_quota_pure::coerce_json_f64); + if current_used + .zip(incoming_used) + .is_some_and(|(current, incoming)| current > incoming + f64::EPSILON) + { + return true; + } + } + + let current_updated = current + .get("updated_at") + .and_then(admin_provider_quota_pure::coerce_json_u64); + let incoming_updated = incoming + .get("updated_at") + .and_then(admin_provider_quota_pure::coerce_json_u64); + current_updated + .zip(incoming_updated) + .is_some_and(|(current, incoming)| current > incoming && current_reset == incoming_reset) +} + fn get_cached_codex_quota_fingerprint(key_id: &str, now: Instant) -> Option { let mut cache = codex_quota_header_fingerprint_cache() .lock() @@ -333,6 +420,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>, @@ -646,21 +762,40 @@ async fn apply_local_sync_report_effect(state: &AppState, payload: &GatewaySyncR } async fn apply_local_stream_report_effect(state: &AppState, payload: &GatewayStreamReportRequest) { - if let Err(err) = sync_codex_quota_from_response_headers( - state, - payload.report_context.as_ref(), - &payload.headers, - ) - .await + let websocket_quota_seen = match sync_codex_websocket_quota_from_stream_payload(state, payload) + .await { - 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 stream response headers" - ); + 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 { + if let Err(err) = sync_codex_quota_from_response_headers( + state, + payload.report_context.as_ref(), + &payload.headers, + ) + .await + { + 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 stream response headers" + ); + } } if let Err(err) = sync_grok_quota_from_report_context( state, @@ -834,11 +969,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 provider_headers = report_context_provider_response_headers(report_context); let parsed_from_provider_headers = provider_headers.as_ref().and_then(|headers| { @@ -849,11 +979,54 @@ async fn sync_codex_quota_from_response_headers( else { return Ok(false); }; + sync_codex_quota_metadata(state, report_context, parsed, "response_headers").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 { + if aether_admin::provider::quota::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), + }; let Some(incoming_fingerprint) = fingerprint_codex_payload(&parsed) else { return Ok(false); }; let now = Instant::now(); + let now_unix_secs = current_unix_secs(); if get_cached_codex_quota_fingerprint(&key_id, now).as_deref() == Some(incoming_fingerprint.as_str()) { @@ -895,7 +1068,13 @@ async fn sync_codex_quota_from_response_headers( set_cached_codex_quota_fingerprint(&key_id, incoming_fingerprint.clone(), now); return Ok(false); }; - if current_fingerprint == incoming_fingerprint { + if current_fingerprint == incoming_fingerprint + && codex_quota_snapshot_matches_metadata(key.status_snapshot.as_ref(), &parsed) + { + set_cached_codex_quota_fingerprint(&key_id, incoming_fingerprint, now); + return Ok(false); + } + if codex_snapshot_regresses(¤t_codex, &parsed) { set_cached_codex_quota_fingerprint(&key_id, incoming_fingerprint.clone(), now); return Ok(false); } @@ -906,7 +1085,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(&ProviderCatalogKeyRuntimeMetadataUpdate { @@ -928,6 +1107,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() { if let Some(cache) = CODEX_QUOTA_HEADER_FINGERPRINT_CACHE.get() { @@ -1022,6 +1209,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!({ @@ -1208,4 +1422,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/privacy/mod.rs b/apps/aether-gateway/src/privacy/mod.rs index 3e5ddc858..b00c8de01 100644 --- a/apps/aether-gateway/src/privacy/mod.rs +++ b/apps/aether-gateway/src/privacy/mod.rs @@ -2523,7 +2523,14 @@ fn restore_json_response_body( }) } -fn restore_json_strings(value: &mut Value, session: &RedactionSession) -> bool { +/// 递归把 JSON 里的占位符换回真实值,只认本 `session` 记录过的映射。 +/// +/// 同步响应体([`restore_sync_response_body`])和 Responses WebSocket 的 +/// provider 事件帧(`handlers::proxy::websocket::responses::redaction`)共用它, +/// 两边因此保持同一套还原语义:未映射的占位符原样保留,`type` / `model` / `id` +/// 这类协议字段虽然也被遍历,但它们不可能包含本 session 派生出的 sentinel, +/// 所以不会被改写。 +pub(crate) fn restore_json_strings(value: &mut Value, session: &RedactionSession) -> bool { match value { Value::String(text) => { let restored = session.restore_text(text); 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 b4f755471..c66a7fccc 100644 --- a/crates/aether-admin/src/provider/quota.rs +++ b/crates/aether-admin/src/provider/quota.rs @@ -800,6 +800,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) { @@ -2122,11 +2290,13 @@ pub fn parse_chatgpt_web_conversation_init_response( #[cfg(test)] mod tests { use super::{ - codex_build_invalid_state, codex_runtime_invalid_reason, extract_execution_error_detail, + codex_build_invalid_state, codex_rate_limit_metadata_exhausted, + codex_runtime_invalid_reason, extract_execution_error_detail, 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_wham_reset_credits_detail_response, - parse_codex_wham_usage_response, parse_gemini_cli_retrieve_user_quota_response, + 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, parse_windsurf_rate_limit_response, parse_windsurf_user_status_response, provider_auto_remove_quota_exhausted_keys, quota_refresh_success_invalid_state, @@ -2628,6 +2798,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/formats/src/formats/openai/chat/stream.rs b/crates/aether-ai/formats/src/formats/openai/chat/stream.rs index b97dffc33..946efdf32 100644 --- a/crates/aether-ai/formats/src/formats/openai/chat/stream.rs +++ b/crates/aether-ai/formats/src/formats/openai/chat/stream.rs @@ -1264,6 +1264,11 @@ impl OpenAIResponsesProviderState { } } + /// SSE 入口:剥掉 `data:` 包装后交给 [`Self::push_event`]。 + /// + /// 解码是这个函数唯一做的事,协议状态机全在 `push_event` 里。已经持有结构化 + /// 事件的传输(Responses WebSocket)应当直接调用 `push_event`,不要为了复用 + /// 这个入口先把事件拼回 SSE 文本。 pub fn push_line( &mut self, report_context: &Value, @@ -1272,6 +1277,17 @@ impl OpenAIResponsesProviderState { let Some(value) = decode_json_data_line(&line) else { return Ok(Vec::new()); }; + self.push_event(report_context, &value) + } + + /// 结构化入口:消费一个已经解析好的 Responses 协议事件。 + /// + /// 取借用而不是所有权:持有结构化事件的传输不必为了调用它先克隆一份。 + pub fn push_event( + &mut self, + report_context: &Value, + value: &Value, + ) -> Result, AiSurfaceFinalizeError> { let mut out = Vec::new(); if let Some(response) = value.get("response").and_then(Value::as_object) { self.response_id = response @@ -1301,12 +1317,12 @@ impl OpenAIResponsesProviderState { } "response.output_text.delta" | "response.outtext.delta" => match value.get("delta") { Some(Value::String(piece)) if !piece.is_empty() => { - let key = Self::text_part_key_from_event(&value); + let key = Self::text_part_key_from_event(value); self.emit_text_delta(report_context, &mut out, key, piece); } Some(Value::Object(delta)) => { if let Some(text) = delta.get("text").and_then(Value::as_str) { - let key = Self::text_part_key_from_event(&value); + let key = Self::text_part_key_from_event(value); self.emit_missing_text(report_context, &mut out, key, text); } } @@ -1317,7 +1333,7 @@ impl OpenAIResponsesProviderState { if part.get("type").and_then(Value::as_str) == Some("output_text") { if let Some(text) = part.get("text").and_then(Value::as_str) { if !text.is_empty() { - let key = Self::text_part_key_from_event(&value); + let key = Self::text_part_key_from_event(value); self.emit_missing_text(report_context, &mut out, key, text); } } @@ -1356,7 +1372,7 @@ impl OpenAIResponsesProviderState { }) .unwrap_or_default(); if !text.is_empty() { - let key = Self::text_part_key_from_event(&value); + let key = Self::text_part_key_from_event(value); self.emit_missing_text(report_context, &mut out, key, text); } } @@ -1366,7 +1382,7 @@ impl OpenAIResponsesProviderState { .and_then(Value::as_str) .unwrap_or_default(); if !piece.is_empty() { - let key = Self::text_part_key_from_event(&value); + let key = Self::text_part_key_from_event(value); self.emit_text_delta(report_context, &mut out, key, piece); } } @@ -1383,7 +1399,7 @@ impl OpenAIResponsesProviderState { }) .unwrap_or_default(); if !refusal.is_empty() { - let key = Self::text_part_key_from_event(&value); + let key = Self::text_part_key_from_event(value); self.emit_missing_text(report_context, &mut out, key, refusal); } } @@ -1393,7 +1409,7 @@ impl OpenAIResponsesProviderState { .and_then(Value::as_str) .unwrap_or_default(); if !piece.is_empty() { - let key = Self::text_part_key_from_event(&value); + let key = Self::text_part_key_from_event(value); self.emit_text_delta(report_context, &mut out, key, piece); } } @@ -1404,7 +1420,7 @@ impl OpenAIResponsesProviderState { .or_else(|| value.get("text").and_then(Value::as_str)) .unwrap_or_default(); if !transcript.is_empty() { - let key = Self::text_part_key_from_event(&value); + let key = Self::text_part_key_from_event(value); self.emit_missing_text(report_context, &mut out, key, transcript); } } @@ -1477,7 +1493,7 @@ impl OpenAIResponsesProviderState { self.emit_output_item_event( report_context, &mut out, - &value, + value, item, output_index, false, @@ -1724,7 +1740,7 @@ impl OpenAIResponsesProviderState { self.emit_output_item_event( report_context, &mut out, - &value, + value, item, output_index, true, @@ -1742,7 +1758,7 @@ impl OpenAIResponsesProviderState { id, model, event: CanonicalStreamEvent::Finish { - finish_reason: Some(openai_responses_incomplete_finish_reason(&value)), + finish_reason: Some(openai_responses_incomplete_finish_reason(value)), usage: canonical_usage_from_openai_usage(response.get("usage")), }, }); @@ -1752,14 +1768,14 @@ impl OpenAIResponsesProviderState { event_type if openai_responses_stream_event_is_known_noop(event_type) => { self.ensure_started(report_context, &mut out); } - event_type if openai_stream_payload_is_terminal_error(&value) => { + event_type if openai_stream_payload_is_terminal_error(value) => { self.finished = true; let mut payload = value.clone(); if event_type != "response.failed" && event_type != "response.incomplete" && event_type != "error" { - payload = openai_stream_terminal_error_body(&value).unwrap_or(payload); + payload = openai_stream_terminal_error_body(value).unwrap_or(payload); if let Some(object) = payload.as_object_mut() { object.insert( "type".to_string(), diff --git a/crates/aether-ai/formats/src/formats/shared/stream_core/format_matrix.rs b/crates/aether-ai/formats/src/formats/shared/stream_core/format_matrix.rs index 7a485d2eb..62c71740e 100644 --- a/crates/aether-ai/formats/src/formats/shared/stream_core/format_matrix.rs +++ b/crates/aether-ai/formats/src/formats/shared/stream_core/format_matrix.rs @@ -251,6 +251,44 @@ impl StreamingStandardTerminalObserver { Ok(()) } + /// 结构化入口:给已经持有解析好的协议事件的传输用(Responses WebSocket), + /// 避免为了复用 [`Self::push_line`] 把事件重新拼成 `data: {json}` 再解析回来。 + /// + /// 只要 provider 的协议状态机本身接受结构化事件,这条路径与 `push_line` + /// 完全等价——`push_line` 现在就是「解码 + `push_event`」。 + /// + /// `openai:image` 的终态状态机没有结构化入口(它按 SSE 行做增量解析), + /// 这里返回 `Err`,由调用方 `disable_with_error` 把摘要标成 parser_error, + /// 而不是静默丢事件。 + pub fn push_event( + &mut self, + report_context: &Value, + event: &Value, + ) -> Result<(), AiSurfaceFinalizeError> { + self.ensure_initialized(report_context); + let Some(provider) = self.provider.as_mut() else { + return Ok(()); + }; + match provider { + TerminalStreamParser::Standard(provider) => { + let frames = provider.push_event(report_context, event)?; + let actual_service_tier = provider.actual_service_tier().map(ToOwned::to_owned); + self.observe_frames(frames); + if let Some(actual_service_tier) = actual_service_tier { + self.latest_summary + .get_or_insert_with(ExecutionStreamTerminalSummary::default) + .provider_actual_service_tier = Some(actual_service_tier); + } + } + TerminalStreamParser::OpenAIImage(_) => { + return Err(AiSurfaceFinalizeError::new( + "openai:image terminal observation has no structured event entry", + )); + } + } + Ok(()) + } + pub fn finish( &mut self, report_context: &Value, @@ -404,6 +442,24 @@ impl ProviderStreamParser { } } + /// 结构化入口。目前只有 `openai:responses` 有传输会走它(Responses + /// WebSocket);其余格式的协议状态机同样可以按「解码 + push_event」机械拆分, + /// 等到真有非 SSE 传输需要时再拆,不做无调用方的接口。 + fn push_event( + &mut self, + report_context: &Value, + event: &Value, + ) -> Result, AiSurfaceFinalizeError> { + match self { + ProviderStreamParser::OpenAIResponses(state) => state.push_event(report_context, event), + ProviderStreamParser::OpenAIChat(_) + | ProviderStreamParser::Claude(_) + | ProviderStreamParser::Gemini(_) => Err(AiSurfaceFinalizeError::new( + "this provider stream parser has no structured event entry", + )), + } + } + fn finish( &mut self, report_context: &Value, @@ -2490,3 +2546,260 @@ mod tests { ); } } + +#[cfg(test)] +mod structured_entry_tests { + use super::StreamingStandardTerminalObserver; + use aether_contracts::ExecutionStreamTerminalSummary; + use serde_json::{json, Value}; + + fn report_context() -> Value { + json!({ + "provider_api_format": "openai:responses", + "client_api_format": "openai:responses", + "mapped_model": "gpt-5-codex", + }) + } + + /// 用 SSE 入口观测一组事件。这是 C5 之前 WebSocket 走的路径:把结构化事件 + /// 拼成 `data: {json}` 再交给解析器。 + fn summary_via_push_line(events: &[Value]) -> ExecutionStreamTerminalSummary { + let context = report_context(); + let mut observer = StreamingStandardTerminalObserver::default(); + for event in events { + observer + .push_line(&context, format!("data: {event}\n\n").into_bytes()) + .expect("the SSE entry must accept these events"); + } + observer + .finish(&context) + .expect("the observer must finish") + .unwrap_or_default() + } + + /// 用结构化入口观测同一组事件。这是 C5 之后的路径。 + fn summary_via_push_event(events: &[Value]) -> ExecutionStreamTerminalSummary { + let context = report_context(); + let mut observer = StreamingStandardTerminalObserver::default(); + for event in events { + observer + .push_event(&context, event) + .expect("the structured entry must accept these events"); + } + observer + .finish(&context) + .expect("the observer must finish") + .unwrap_or_default() + } + + fn assert_entries_agree(label: &str, events: &[Value]) { + let via_line = summary_via_push_line(events); + let via_event = summary_via_push_event(events); + assert_eq!( + via_line, via_event, + "the SSE entry and the structured entry must produce identical summaries for {label}" + ); + } + + fn created() -> Value { + json!({"type": "response.created", "response": {"id": "resp_diff", "model": "gpt-5-codex"}}) + } + + fn text_delta(piece: &str) -> Value { + json!({ + "type": "response.output_text.delta", + "item_id": "msg_diff", + "output_index": 0, + "content_index": 0, + "delta": piece, + }) + } + + /// 批量事件:WS 一帧可以带多个协议事件,逐个喂入的结果必须和逐行喂入一致。 + #[test] + fn a_batched_delta_sequence_agrees_across_both_entries() { + let events = vec![ + created(), + text_delta("he"), + text_delta("ll"), + text_delta("o"), + json!({ + "type": "response.completed", + "response": { + "id": "resp_diff", + "model": "gpt-5-codex", + "status": "completed", + "usage": {"input_tokens": 11, "output_tokens": 3, "total_tokens": 14}, + }, + }), + ]; + assert_entries_agree("a batched delta sequence", &events); + let summary = summary_via_push_event(&events); + assert!(summary.observed_finish); + assert_eq!(summary.response_id.as_deref(), Some("resp_diff")); + let usage = summary + .standardized_usage + .as_ref() + .expect("completed carries usage"); + assert_eq!(usage.input_tokens, 11); + assert_eq!(usage.output_tokens, 3); + } + + /// 合法 `response.incomplete`:C1 定过的语义(终态、可计费),两条入口必须 + /// 得到同一个摘要,尤其是 finish_reason 与 parser_error 的取值。 + #[test] + fn a_legitimate_incomplete_agrees_across_both_entries() { + let events = vec![ + created(), + text_delta("partial"), + json!({ + "type": "response.incomplete", + "response": { + "id": "resp_diff", + "model": "gpt-5-codex", + "status": "incomplete", + "incomplete_details": {"reason": "max_output_tokens"}, + "usage": {"input_tokens": 7, "output_tokens": 5, "total_tokens": 12}, + }, + }), + ]; + assert_entries_agree("a legitimate incomplete", &events); + let summary = summary_via_push_event(&events); + assert!(summary.observed_finish); + assert!( + summary.parser_error.is_none(), + "a legitimate incomplete is not a parser error: {:?}", + summary.parser_error + ); + } + + #[test] + fn a_terminal_error_agrees_across_both_entries() { + let events = vec![ + created(), + json!({ + "type": "error", + "error": {"type": "server_error", "message": "upstream exploded"}, + }), + ]; + assert_entries_agree("a terminal error", &events); + let summary = summary_via_push_event(&events); + assert!(summary.observed_finish); + assert_eq!(summary.finish_reason.as_deref(), Some("error")); + assert!(summary.parser_error.is_some()); + } + + #[test] + fn a_response_failed_event_agrees_across_both_entries() { + assert_entries_agree( + "a response.failed event", + &[ + created(), + json!({ + "type": "response.failed", + "response": { + "id": "resp_diff", + "model": "gpt-5-codex", + "status": "failed", + "error": {"type": "server_error", "message": "generation failed"}, + }, + }), + ], + ); + } + + /// 未知事件只增计数、不改终态判定,两条入口的计数必须一致。 + #[test] + fn unknown_events_agree_across_both_entries() { + let events = vec![ + created(), + json!({"type": "response.some_future_event", "payload": {"anything": true}}), + json!({"type": "response.another_future_event"}), + json!({ + "type": "response.completed", + "response": { + "id": "resp_diff", + "model": "gpt-5-codex", + "status": "completed", + "usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}, + }, + }), + ]; + assert_entries_agree("unknown events", &events); + let summary = summary_via_push_event(&events); + assert!(summary.observed_finish); + assert!( + summary.unknown_event_count > 0, + "unknown events are counted" + ); + } + + /// 供应商声明的 service tier 通过两条入口都要落到摘要上。 + #[test] + fn a_service_tier_agrees_across_both_entries() { + assert_entries_agree( + "a declared service tier", + &[ + json!({ + "type": "response.created", + "response": { + "id": "resp_diff", + "model": "gpt-5-codex", + "service_tier": "priority", + }, + }), + json!({ + "type": "response.completed", + "response": { + "id": "resp_diff", + "model": "gpt-5-codex", + "status": "completed", + "service_tier": "priority", + "usage": {"input_tokens": 2, "output_tokens": 2, "total_tokens": 4}, + }, + }), + ], + ); + } + + /// 没有任何供应商终态事件。注意 `finish()` 会补一个 `stop`(这是 HTTP 与 + /// WebSocket 共享的既有行为,C5 不改),所以 `observed_finish` 为真而 usage + /// 缺失——真正的「缺终态」判定看的是 usage 与被捕获的 body。这里要钉住的是 + /// 两条入口在这种不完整序列上仍然给出同一个摘要。 + #[test] + fn a_missing_terminal_agrees_across_both_entries() { + let events = vec![created(), text_delta("truncated")]; + assert_entries_agree("a missing terminal", &events); + let summary = summary_via_push_event(&events); + assert_eq!(summary.finish_reason.as_deref(), Some("stop")); + assert!( + summary.standardized_usage.is_none(), + "a synthesized finish carries no usage" + ); + } + + /// `openai:image` 没有结构化入口:必须显式报错,让调用方标记 parser_error, + /// 而不是静默丢掉事件、把摘要留成「未观察到终态」。 + #[test] + fn the_image_format_rejects_the_structured_entry() { + let context = json!({ + "provider_api_format": "openai:image", + "client_api_format": "openai:image", + "mapped_model": "gpt-image-1", + }); + let mut observer = StreamingStandardTerminalObserver::default(); + let error = observer + .push_event(&context, &json!({"type": "image_generation.completed"})) + .expect_err("openai:image has no structured entry"); + assert!( + error.to_string().contains("structured event entry"), + "the error must name the missing entry: {error}" + ); + + observer.disable_with_error(error.to_string()); + let summary = observer + .latest_summary() + .expect("disable_with_error records a summary"); + assert!(summary.parser_error.is_some()); + } +} 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-testing/integration/Cargo.toml b/crates/aether-testing/integration/Cargo.toml index e93236fc3..eafd3d999 100644 --- a/crates/aether-testing/integration/Cargo.toml +++ b/crates/aether-testing/integration/Cargo.toml @@ -9,6 +9,7 @@ description = "Gateway-backed integration scenarios and benchmark binaries" [dependencies] async-stream.workspace = true aether-contracts.workspace = true +aether-crypto.workspace = true aether-data.workspace = true aether-data-contracts.workspace = true aether-gateway = { workspace = true, features = ["testkit"] } @@ -24,3 +25,4 @@ sha2.workspace = true sqlx = { workspace = true, features = ["postgres"] } tokio.workspace = true tokio-tungstenite = { version = "0.28", features = ["rustls-tls-webpki-roots"] } +uuid.workspace = true diff --git a/crates/aether-testing/integration/tests/responses_websocket_e2e.rs b/crates/aether-testing/integration/tests/responses_websocket_e2e.rs new file mode 100644 index 000000000..431077b75 --- /dev/null +++ b/crates/aether-testing/integration/tests/responses_websocket_e2e.rs @@ -0,0 +1,1285 @@ +//! Responses WebSocket end-to-end coverage. +//! +//! Every test starts a protocol-aware mock upstream, seeds a throwaway SQLite +//! store, mounts the real gateway router, and drives the public +//! `/v1/responses` WebSocket the way a client would. +//! +//! The assertions deliberately reach back into the database. A turn settles its +//! billing row from a task that outlives the relay loop, so a client that saw +//! `response.completed` is not evidence that the turn was ever accounted for — +//! only the row is. + +use std::path::PathBuf; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::Arc; +use std::time::Duration; + +use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; +use aether_data::repository::auth::CreateStandaloneApiKeyRecord; +use aether_data::repository::wallet::WalletLookupKey; +use aether_data::{ + DataBackends, DataLayerConfig, DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig, +}; +use aether_data_contracts::repository::global_models::{ + CreateAdminGlobalModelRecord, UpsertAdminProviderModelRecord, +}; +use aether_data_contracts::repository::provider_catalog::{ + StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, +}; +use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UsageAuditListQuery}; +use aether_gateway::{build_router_with_state, AppState, GatewayDataConfig, UsageRuntimeConfig}; +use aether_testkit::SpawnedServer; +use axum::extract::ws::{Message as AxumWsMessage, WebSocket, WebSocketUpgrade}; +use axum::extract::State; +use axum::http::HeaderMap; +use axum::response::Response; +use axum::routing::get; +use axum::Router; +use futures_util::{SinkExt, StreamExt}; +use serde_json::{json, Value}; +use sha2::Digest; +use tokio::net::TcpStream; +use tokio::sync::Mutex; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::tungstenite::Message; +use tokio_tungstenite::{MaybeTlsStream, WebSocketStream}; + +type BoxError = Box; +type ClientSocket = WebSocketStream>; + +const CLIENT_API_KEY: &str = "sk-aether-responses-ws-e2e"; +const PROVIDER_API_KEY: &str = "sk-upstream-responses-ws-e2e"; +const PROVIDER_ID: &str = "provider-responses-ws-e2e"; +const ENDPOINT_ID: &str = "endpoint-responses-ws-e2e"; +const PROVIDER_KEY_ID: &str = "provider-key-responses-ws-e2e"; +/// 透明重试的替代 key。只有配额重试用例会 seed 它。 +const ALTERNATE_PROVIDER_KEY_ID: &str = "provider-key-responses-ws-e2e-alt"; +const ALTERNATE_PROVIDER_API_KEY: &str = "sk-upstream-responses-ws-e2e-alt"; +const GLOBAL_MODEL_ID: &str = "global-model-responses-ws-e2e"; +const PROVIDER_MODEL_ID: &str = "provider-model-responses-ws-e2e"; +const API_KEY_ID: &str = "api-key-responses-ws-e2e"; +const PUBLIC_MODEL: &str = "gpt-responses-ws-e2e"; +const UPSTREAM_MODEL: &str = "gpt-responses-ws-upstream"; + +/// 2100-01-01,保证 oauth 凭证在测试期间不会被判为过期。 +const FAR_FUTURE_UNIX_SECS: u64 = 4_102_444_800; + +const INPUT_TOKENS: u64 = 4; +const OUTPUT_TOKENS: u64 = 2; + +/// Generous enough to absorb a loaded CI runner, short enough that a genuinely +/// lost row fails the test instead of hanging the job. +const SETTLE_TIMEOUT: Duration = Duration::from_secs(30); +const RECEIVE_TIMEOUT: Duration = Duration::from_secs(15); + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +/// The headline guarantee: a continuation stays on one physical upstream socket, +/// and both turns are billed independently. +#[tokio::test] +async fn continuation_reuses_one_upstream_connection_and_bills_both_turns() -> Result<(), BoxError> +{ + let harness = Harness::start(UpstreamBehavior::CompleteEveryTurn).await?; + let mut client = harness.connect().await?; + + client + .send(response_create(json!({"input": "first turn"}))) + .await?; + let first = receive_event(&mut client, "response.completed").await?; + assert_eq!( + first.pointer("/response/id").and_then(Value::as_str), + Some("resp-e2e-1") + ); + + client + .send(response_create(json!({ + "previous_response_id": "resp-e2e-1", + "input": "second turn" + }))) + .await?; + let second = receive_event(&mut client, "response.completed").await?; + assert_eq!( + second.pointer("/response/id").and_then(Value::as_str), + Some("resp-e2e-2") + ); + + // The mock records each `response.create` before answering it, so both + // turns completing means both are already on record. + let upstream_events = harness.upstream.observed_events().await; + assert_eq!( + upstream_events.len(), + 2, + "one upstream turn per client turn" + ); + assert_eq!( + harness.upstream.connections(), + 1, + "the continuation must reuse the bound upstream socket" + ); + for event in &upstream_events { + assert_eq!( + event.get("model").and_then(Value::as_str), + Some(UPSTREAM_MODEL), + "every turn is rewritten to the mapped provider model" + ); + } + assert_eq!( + upstream_events[1] + .get("previous_response_id") + .and_then(Value::as_str), + Some("resp-e2e-1"), + "the continuation id survives provider body normalization" + ); + assert_eq!( + harness.upstream.authorization_headers().await, + vec![Some(format!("Bearer {PROVIDER_API_KEY}"))], + "the upstream is opened with the configured provider key" + ); + + let audits = harness + .usage_audits_where(2, "billed turns", is_billed) + .await?; + assert_eq!(audits.len(), 2, "each response.create bills separately"); + for audit in &audits { + assert!( + audit.is_websocket(), + "turns are recorded as WebSocket usage, metadata: {:?}", + audit.request_metadata + ); + assert_eq!(audit.model, PUBLIC_MODEL); + assert_eq!(audit.input_tokens, INPUT_TOKENS); + assert_eq!(audit.output_tokens, OUTPUT_TOKENS); + assert_eq!(audit.total_tokens, INPUT_TOKENS + OUTPUT_TOKENS); + assert_eq!(audit.status_code, Some(200)); + } + assert_ne!( + audits[0].request_id, audits[1].request_id, + "each turn gets its own logical request identity" + ); + + Ok(()) +} + +/// A client that walks away before the provider produced anything must settle +/// as a void row: nothing was produced, so nothing is billed. +/// +/// This is the path with no protocol event to announce it: the relay loop owns +/// the turn, and losing the client is an exit the upstream never reports. +/// +/// The mirror case — the provider *did* reach a terminal event and only the last +/// hop to the client failed — is billed instead. That one cannot be pinned here: +/// it depends on the relay loop's `select!` observing the upstream terminal frame +/// before it observes the closed client socket, which is a race by construction. +/// It is covered deterministically by the relay-level unit tests +/// `a_provider_terminal_that_reaches_a_closed_client_socket_is_still_billed` and +/// `a_closed_client_socket_before_any_terminal_still_voids_the_bill`. +#[tokio::test] +async fn client_disconnect_before_any_provider_output_settles_a_void_row() -> Result<(), BoxError> { + let harness = Harness::start(UpstreamBehavior::StallAfterCreated).await?; + let mut client = harness.connect().await?; + + client + .send(response_create(json!({"input": "abandoned turn"}))) + .await?; + // Leave only once the turn is genuinely in flight upstream, so this covers + // an interrupted turn rather than racing turn start. + receive_event(&mut client, "response.created").await?; + drop(client); + + let audits = harness + .usage_audits_where(1, "settled turns", |audit| !is_pending(audit)) + .await?; + assert_eq!(audits.len(), 1, "the abandoned turn is still accounted for"); + let audit = &audits[0]; + assert_eq!(audit.model, PUBLIC_MODEL); + assert!( + !is_pending(audit), + "an abandoned turn must not be left pending: {audit:?}" + ); + // The provider never emitted a terminal event, so this row stays void. + // Only a reached provider terminal survives a client delivery failure. + assert!( + !is_billed(audit), + "a turn with no provider output must not be billed: {audit:?}" + ); + assert_eq!( + audit.status, "cancelled", + "a client that left before any provider output settles as cancelled: {audit:?}" + ); + assert_eq!(audit.status_code, Some(499)); + + Ok(()) +} + +/// An upstream that dies mid-turn must surface an error and still settle. +#[tokio::test] +async fn upstream_drop_mid_turn_reports_an_error_and_settles_the_usage_row() -> Result<(), BoxError> +{ + let harness = Harness::start(UpstreamBehavior::CloseAfterCreated).await?; + let mut client = harness.connect().await?; + + client + .send(response_create(json!({"input": "doomed turn"}))) + .await?; + let error = receive_error_or_close(&mut client) + .await? + .ok_or("gateway closed without telling the client why")?; + assert_eq!(error.get("type").and_then(Value::as_str), Some("error")); + + let audits = harness + .usage_audits_where(1, "settled turns", |audit| !is_pending(audit)) + .await?; + assert_eq!(audits.len(), 1, "the failed turn is still accounted for"); + let audit = &audits[0]; + assert!( + !is_pending(audit), + "a failed turn must not be left pending: {audit:?}" + ); + + Ok(()) +} + +/// 供应商配额耗尽后的透明重试:客户端不该看到 429,两个 attempt 都要结算。 +/// +/// 第一个 attempt 拿到 Codex 的 `usage_limit_reached`,网关换到第二把 key 重开一条 +/// 上游连接重放同一个 `response.create`。C6 之前,重试的规划发生在旧 attempt 结算 +/// 之前:规划读到的是旧 attempt 还没投射的 health / adaptive / pool 状态,而且旧 +/// attempt 的 pool key lease 还被它自己占着。 +/// +/// 顺序本身在这里无法确定性断言(结算与规划都在同一个任务里、DB 里看不到先后), +/// 由 lifecycle 的单测确定性覆盖;这个用例保证整条路径真的能跑通,并且两个 +/// attempt 都留下了终态记账行。 +#[tokio::test] +async fn provider_quota_exhaustion_transparently_retries_onto_another_key() -> Result<(), BoxError> +{ + let harness = Harness::start_with_fixture( + UpstreamBehavior::QuotaExhaustedThenComplete, + ProviderFixture::CodexKeyPair, + ) + .await?; + let mut client = harness.connect().await?; + + client + .send(response_create(json!({"input": "retry after quota"}))) + .await?; + + // 客户端只应该看到重试之后那次成功的响应,看不到 429。 + let completed = receive_event(&mut client, "response.completed").await?; + assert_eq!( + completed + .pointer("/response/status") + .and_then(Value::as_str), + Some("completed") + ); + + // 上游被连了两次:配额耗尽的那条 + 重试用的那条。 + assert_eq!( + harness.upstream.connections(), + 2, + "the transparent retry must open a second upstream connection" + ); + let observed = harness.upstream.observed_events().await; + assert_eq!( + observed.len(), + 2, + "the same response.create must be replayed once" + ); + + // 两把不同的 key 被用过:重试不能落回那把已经耗尽的 key。 + let authorizations = harness.upstream.authorization_headers().await; + assert_eq!(authorizations.len(), 2); + assert_ne!( + authorizations[0], authorizations[1], + "the retry must not reuse the exhausted key: {authorizations:?}" + ); + + // 两个 attempt 各自留下一条终态行:配额失败的那条 + 成功计费的那条。 + let audits = harness + .usage_audits_where(2, "settled attempts", |audit| !is_pending(audit)) + .await?; + let settled = audits + .iter() + .filter(|audit| !is_pending(audit)) + .collect::>(); + assert_eq!( + settled.len(), + 2, + "both attempts must reach a terminal accounting row: {:?}", + audits + .iter() + .map(|audit| (audit.status.clone(), audit.status_code, audit.total_tokens)) + .collect::>() + ); + assert!( + settled.iter().any(|audit| audit.status_code == Some(429)), + "the exhausted attempt keeps its own 429 row: {:?}", + settled + .iter() + .map(|audit| (audit.status.clone(), audit.status_code)) + .collect::>() + ); + + let billed = harness + .usage_audits_where(1, "the billed retry attempt", is_billed) + .await?; + let retry = billed + .iter() + .find(|audit| is_billed(audit)) + .ok_or("the successful retry attempt must be billed")?; + assert_eq!( + retry.total_tokens, + INPUT_TOKENS + OUTPUT_TOKENS, + "the retry attempt is billed for what it actually consumed" + ); + + client.close(None).await?; + Ok(()) +} + +/// 脱敏的另一半:请求侧把真实 PII 换成占位符发给上游,响应侧必须在推给客户端之前 +/// 换回真实值。 +/// +/// 上游把收到的 `input` 原样回显,所以它回来的就是占位符——这一条同时钉住了两个 +/// 方向:上游不能看到原文,客户端不能看到占位符。 +#[tokio::test] +async fn redacted_pii_is_restored_before_the_client_sees_a_provider_frame() -> Result<(), BoxError> +{ + const CLIENT_EMAIL: &str = "responses.ws.pii@example.com"; + + let harness = Harness::start_with_pii_redaction(UpstreamBehavior::EchoInputBack).await?; + let mut client = harness.connect().await?; + + client + .send(response_create( + json!({"input": format!("my mail is {CLIENT_EMAIL}")}), + )) + .await?; + + let delta = receive_event(&mut client, "response.output_text.delta").await?; + let delta_text = delta + .get("delta") + .and_then(Value::as_str) + .ok_or("the provider delta must carry text")?; + assert!( + delta_text.contains(CLIENT_EMAIL), + "the client must receive the restored value: {delta_text}" + ); + assert!( + !delta_text.contains(" bool { + audit.status.eq_ignore_ascii_case("pending") +} + +/// A turn that finished accounting: settled, and carrying what it consumed. +fn is_billed(audit: &StoredRequestUsageAudit) -> bool { + !is_pending(audit) && audit.total_tokens > 0 +} + +// --------------------------------------------------------------------------- +// Harness +// --------------------------------------------------------------------------- + +/// A live gateway wired to a mock Responses WebSocket upstream over a throwaway +/// SQLite store. +struct Harness { + database: TemporarySqlite, + upstream: Arc, + websocket_url: String, + _upstream_server: SpawnedServer, + _gateway_server: SpawnedServer, +} + +/// 供应商夹具形态。 +/// +/// 透明配额重试只有 Codex adapter 会开启(`retry_current_turn: true` 只从 +/// codex.rs 出),而且重试要有第二把 key 可挑,否则规划直接判无可用供应商。 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum ProviderFixture { + /// 单个 openai 类型供应商、单把 key。 + SingleOpenAiKey, + /// codex 类型供应商 + 两把 key:第一把配额耗尽后重试落到第二把。 + CodexKeyPair, +} + +impl ProviderFixture { + const fn provider_type(self) -> &'static str { + match self { + Self::SingleOpenAiKey => "openai", + Self::CodexKeyPair => "codex", + } + } + + const fn has_alternate_key(self) -> bool { + matches!(self, Self::CodexKeyPair) + } +} + +/// 这条用例要不要打开 chat PII 脱敏模块。 +/// +/// 默认关闭:其余用例都靠原文 body 断言上游看到了什么,打开脱敏会把断言目标换成 +/// 占位符。 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum PiiRedaction { + Disabled, + Enabled, +} + +impl PiiRedaction { + const fn is_enabled(self) -> bool { + matches!(self, Self::Enabled) + } +} + +impl Harness { + async fn start(behavior: UpstreamBehavior) -> Result { + Self::start_with( + behavior, + ProviderFixture::SingleOpenAiKey, + PiiRedaction::Disabled, + ) + .await + } + + async fn start_with_fixture( + behavior: UpstreamBehavior, + fixture: ProviderFixture, + ) -> Result { + Self::start_with(behavior, fixture, PiiRedaction::Disabled).await + } + + async fn start_with_pii_redaction(behavior: UpstreamBehavior) -> Result { + Self::start_with( + behavior, + ProviderFixture::SingleOpenAiKey, + PiiRedaction::Enabled, + ) + .await + } + + async fn start_with( + behavior: UpstreamBehavior, + fixture: ProviderFixture, + redaction: PiiRedaction, + ) -> Result { + let upstream = Arc::new(MockUpstreamState::new(behavior)); + let upstream_server = + SpawnedServer::start(mock_upstream_router(Arc::clone(&upstream))).await?; + + let database = TemporarySqlite::new(); + prepare_and_seed_database( + &database.config, + upstream_server.base_url(), + fixture, + redaction, + ) + .await?; + + let data_config = GatewayDataConfig::from_database_config(database.config.clone()) + .with_encryption_key(DEVELOPMENT_ENCRYPTION_KEY); + let state = AppState::new()? + .with_data_config_and_background_isolation(data_config, false)? + // The usage runtime defaults to disabled, which silently turns every + // terminal usage write into a no-op. Without this the suite could + // not observe billing at all. Queueing stays off so the terminal + // write lands through the in-process path instead of Redis. + .with_usage_runtime_config(UsageRuntimeConfig { + enabled: true, + ..UsageRuntimeConfig::default() + })?; + let gateway_server = SpawnedServer::start(build_router_with_state(state)).await?; + let websocket_url = format!( + "{}/v1/responses", + gateway_server.base_url().replacen("http://", "ws://", 1) + ); + + Ok(Self { + database, + upstream, + websocket_url, + _upstream_server: upstream_server, + _gateway_server: gateway_server, + }) + } + + async fn connect(&self) -> Result { + let mut request = self.websocket_url.clone().into_client_request()?; + request.headers_mut().insert( + "authorization", + http::HeaderValue::from_str(&format!("Bearer {CLIENT_API_KEY}"))?, + ); + let (socket, response) = + tokio::time::timeout(RECEIVE_TIMEOUT, tokio_tungstenite::connect_async(request)) + .await + .map_err(|_| "timed out connecting to the gateway WebSocket")??; + if response.status() != http::StatusCode::SWITCHING_PROTOCOLS { + return Err( + format!("unexpected gateway handshake status: {}", response.status()).into(), + ); + } + Ok(socket) + } + + /// Waits until `expected` usage rows satisfy `settled`. + /// + /// A row is created `Pending` at turn start and reaches its final shape + /// through several independent writes, so "no longer pending" does not imply + /// "finished": a row can briefly read as completed with zero tokens and no + /// WebSocket metadata before the terminal write lands. Each caller waits for + /// the specific end state it is about to assert. + async fn usage_audits_where( + &self, + expected: usize, + what: &str, + settled: impl Fn(&StoredRequestUsageAudit) -> bool, + ) -> Result, BoxError> { + let deadline = tokio::time::Instant::now() + SETTLE_TIMEOUT; + loop { + let audits = self.usage_audits().await?; + if audits.iter().filter(|audit| settled(audit)).count() >= expected { + return Ok(audits); + } + if tokio::time::Instant::now() >= deadline { + let observed = audits + .iter() + .map(|audit| { + format!( + "{} status={} code={:?} tokens={} websocket={}", + audit.request_id, + audit.status, + audit.status_code, + audit.total_tokens, + audit.is_websocket() + ) + }) + .collect::>(); + return Err(format!( + "timed out waiting for {expected} {what}; observed {}: {observed:?}", + audits.len() + ) + .into()); + } + tokio::time::sleep(Duration::from_millis(25)).await; + } + } + + /// Reads the persisted audit rows, oldest first. + /// + /// Opens its own handle per call rather than holding one for the lifetime of + /// the harness: the gateway keeps its own pool on the same SQLite file for + /// the whole test, and an idle second pool only adds contention. + async fn usage_audits(&self) -> Result, BoxError> { + let backends = DataBackends::from_config(DataLayerConfig::from_database( + self.database.config.clone(), + ))?; + let audits = backends + .read() + .usage() + .ok_or("usage reader unavailable")? + .list_usage_audits(&UsageAuditListQuery { + limit: Some(50), + newest_first: false, + ..UsageAuditListQuery::default() + }) + .await?; + drop(backends); + Ok(audits) + } +} + +// --------------------------------------------------------------------------- +// Client protocol helpers +// --------------------------------------------------------------------------- + +/// Builds a `response.create` frame for the seeded public model. +fn response_create(fields: Value) -> Message { + let mut event = json!({"type": "response.create", "model": PUBLIC_MODEL}); + let object = event + .as_object_mut() + .expect("the literal above is an object"); + for (key, value) in fields + .as_object() + .expect("response.create fields must be an object") + { + object.insert(key.clone(), value.clone()); + } + Message::Text(event.to_string().into()) +} + +/// Reads frames until `expected_type` arrives, failing fast on a gateway error. +async fn receive_event( + socket: &mut WebSocketStream, + expected_type: &str, +) -> Result +where + S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin, +{ + tokio::time::timeout(RECEIVE_TIMEOUT, async { + loop { + let message = socket + .next() + .await + .ok_or("gateway WebSocket closed before the expected event")??; + match message { + Message::Text(text) => { + let event: Value = serde_json::from_str(text.as_ref())?; + match event.get("type").and_then(Value::as_str) { + Some("error") => { + return Err(format!("gateway returned an error event: {event}").into()) + } + Some(event_type) if event_type == expected_type => return Ok(event), + _ => {} + } + } + Message::Ping(payload) => socket.send(Message::Pong(payload)).await?, + Message::Close(frame) => { + return Err(format!("gateway closed before {expected_type}: {frame:?}").into()) + } + _ => {} + } + } + }) + .await + .map_err(|_| format!("timed out waiting for {expected_type}"))? +} + +/// Drains the socket until the gateway reports an error or hangs up. +/// +/// Returns the error event when one arrives, `None` when the gateway closed +/// without explaining itself. +async fn receive_error_or_close( + socket: &mut WebSocketStream, +) -> Result, BoxError> +where + S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin, +{ + tokio::time::timeout(RECEIVE_TIMEOUT, async { + loop { + let Some(message) = socket.next().await else { + return Ok(None); + }; + match message? { + Message::Text(text) => { + let event: Value = serde_json::from_str(text.as_ref())?; + if event.get("type").and_then(Value::as_str) == Some("error") { + return Ok(Some(event)); + } + } + Message::Ping(payload) => socket.send(Message::Pong(payload)).await?, + Message::Close(_) => return Ok(None), + _ => {} + } + } + }) + .await + .map_err(|_| "timed out waiting for a gateway error or close")? +} + +// --------------------------------------------------------------------------- +// Mock upstream +// --------------------------------------------------------------------------- + +/// How the mock upstream answers a `response.create`. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum UpstreamBehavior { + /// Announce, stream one delta, and complete — the ordinary turn. + CompleteEveryTurn, + /// Announce the response and then go quiet, leaving the turn in flight. + StallAfterCreated, + /// Announce the response and then hang up mid-turn. + CloseAfterCreated, + /// 第一轮只回一个 Codex 配额耗尽错误,之后的每一轮正常完成。 + /// + /// 第一轮刻意不发 `response.created`:任何标准 `response.*` 事件都会让 + /// codex adapter 把这一轮判成 replay-unsafe,透明重试就不会发生。 + QuotaExhaustedThenComplete, + /// 把收到的 `input` 原样回显成一个 delta,再正常完成。 + /// + /// 上游看到的是脱敏后的 body,所以回显出来的就是占位符——正是响应侧还原要处理 + /// 的形状。 + EchoInputBack, +} + +#[derive(Debug)] +struct MockUpstreamState { + behavior: UpstreamBehavior, + connections: AtomicUsize, + events: Mutex>, + authorization_headers: Mutex>>, +} + +impl MockUpstreamState { + fn new(behavior: UpstreamBehavior) -> Self { + Self { + behavior, + connections: AtomicUsize::new(0), + events: Mutex::new(Vec::new()), + authorization_headers: Mutex::new(Vec::new()), + } + } + + fn connections(&self) -> usize { + self.connections.load(Ordering::Acquire) + } + + async fn observed_events(&self) -> Vec { + self.events.lock().await.clone() + } + + async fn authorization_headers(&self) -> Vec> { + self.authorization_headers.lock().await.clone() + } +} + +fn mock_upstream_router(state: Arc) -> Router { + Router::new() + .route("/v1/responses", get(mock_responses_websocket)) + .with_state(state) +} + +async fn mock_responses_websocket( + State(state): State>, + headers: HeaderMap, + ws: WebSocketUpgrade, +) -> Response { + let authorization = headers + .get("authorization") + .and_then(|value| value.to_str().ok()) + .map(str::to_string); + ws.on_upgrade(move |socket| run_mock_upstream(socket, state, authorization)) +} + +async fn run_mock_upstream( + mut socket: WebSocket, + state: Arc, + authorization: Option, +) { + state.connections.fetch_add(1, Ordering::AcqRel); + state.authorization_headers.lock().await.push(authorization); + while let Some(message) = socket.recv().await { + let Ok(message) = message else { + break; + }; + match message { + AxumWsMessage::Text(text) => { + let Ok(event) = serde_json::from_str::(text.as_str()) else { + break; + }; + if event.get("type").and_then(Value::as_str) != Some("response.create") { + continue; + } + let echoed_input = event + .get("input") + .and_then(Value::as_str) + .unwrap_or_default() + .to_string(); + let turn = { + let mut events = state.events.lock().await; + events.push(event); + events.len() + }; + let response_id = format!("resp-e2e-{turn}"); + match state.behavior { + UpstreamBehavior::CompleteEveryTurn => { + if send_mock_turn(&mut socket, &response_id).await.is_err() { + break; + } + } + UpstreamBehavior::StallAfterCreated => { + if send_mock_created(&mut socket, &response_id).await.is_err() { + break; + } + } + UpstreamBehavior::CloseAfterCreated => { + let _ = send_mock_created(&mut socket, &response_id).await; + break; + } + UpstreamBehavior::QuotaExhaustedThenComplete => { + if turn == 1 { + let _ = + send_mock_event(&mut socket, codex_quota_exhausted_error()).await; + break; + } + if send_mock_turn(&mut socket, &response_id).await.is_err() { + break; + } + } + UpstreamBehavior::EchoInputBack => { + if send_mock_turn_with_delta(&mut socket, &response_id, &echoed_input) + .await + .is_err() + { + break; + } + } + } + } + AxumWsMessage::Ping(payload) => { + if socket.send(AxumWsMessage::Pong(payload)).await.is_err() { + break; + } + } + AxumWsMessage::Close(_) => break, + _ => {} + } + } +} + +/// Codex 的账户级配额耗尽信号。 +/// +/// `status_code: 429` + `error.type: usage_limit_reached` 是 adapter 识别 +/// 「配额耗尽、可透明重试」的最小载荷:解析出的元数据被强制标上 +/// `limit_reached: true`,于是 drain 指令带着 `retry_current_turn: true` 下来。 +fn codex_quota_exhausted_error() -> Value { + json!({ + "type": "error", + "status_code": 429, + "error": { + "type": "usage_limit_reached", + "message": "You have hit your usage limit", + "plan_type": "plus", + "resets_in_seconds": 3_600 + } + }) +} + +async fn send_mock_created(socket: &mut WebSocket, response_id: &str) -> Result<(), axum::Error> { + send_mock_event( + socket, + json!({ + "type": "response.created", + "response": { + "id": response_id, + "status": "in_progress", + "model": UPSTREAM_MODEL + } + }), + ) + .await +} + +async fn send_mock_turn(socket: &mut WebSocket, response_id: &str) -> Result<(), axum::Error> { + send_mock_turn_with_delta(socket, response_id, "hello").await +} + +async fn send_mock_turn_with_delta( + socket: &mut WebSocket, + response_id: &str, + delta: &str, +) -> Result<(), axum::Error> { + send_mock_created(socket, response_id).await?; + send_mock_event( + socket, + json!({ + "type": "response.output_text.delta", + "response_id": response_id, + "delta": delta + }), + ) + .await?; + send_mock_event( + socket, + json!({ + "type": "response.completed", + "response": { + "id": response_id, + "status": "completed", + "model": UPSTREAM_MODEL, + "output": [], + "usage": { + "input_tokens": INPUT_TOKENS, + "output_tokens": OUTPUT_TOKENS, + "total_tokens": INPUT_TOKENS + OUTPUT_TOKENS + } + } + }), + ) + .await +} + +async fn send_mock_event(socket: &mut WebSocket, event: Value) -> Result<(), axum::Error> { + socket + .send(AxumWsMessage::Text(event.to_string().into())) + .await +} + +// --------------------------------------------------------------------------- +// Seeded data store +// --------------------------------------------------------------------------- + +struct TemporarySqlite { + directory: PathBuf, + config: SqlDatabaseConfig, +} + +impl TemporarySqlite { + fn new() -> Self { + let directory = std::env::temp_dir().join(format!( + "aether-responses-ws-e2e-{}-{}", + std::process::id(), + uuid::Uuid::new_v4() + )); + let database_path = directory.join("aether.db"); + Self { + directory, + config: SqlDatabaseConfig { + driver: DatabaseDriver::Sqlite, + url: format!("sqlite://{}", database_path.display()), + pool: SqlPoolConfig { + min_connections: 1, + max_connections: 4, + acquire_timeout_ms: 5_000, + idle_timeout_ms: 30_000, + max_lifetime_ms: 300_000, + statement_cache_capacity: 64, + require_ssl: false, + }, + }, + } + } +} + +impl Drop for TemporarySqlite { + fn drop(&mut self) { + let _ = std::fs::remove_dir_all(&self.directory); + } +} + +async fn prepare_and_seed_database( + database: &SqlDatabaseConfig, + upstream_base_url: &str, + fixture: ProviderFixture, + redaction: PiiRedaction, +) -> Result<(), BoxError> { + let backends = DataBackends::from_config(DataLayerConfig::from_database(database.clone()))?; + let pending = backends + .prepare_database_for_startup() + .await? + .unwrap_or_default(); + if !pending.is_empty() { + backends.run_database_migrations().await?; + } + + seed_provider_catalog(&backends, upstream_base_url, fixture).await?; + seed_models(&backends).await?; + let user_id = seed_user(&backends).await?; + seed_client_api_key(&backends, &user_id).await?; + if redaction.is_enabled() { + seed_chat_pii_redaction(&backends).await?; + } + + let candidates = backends + .read() + .minimal_candidate_selection() + .ok_or("candidate selection reader unavailable")? + .list_for_exact_api_format_and_requested_model("openai:responses", PUBLIC_MODEL) + .await?; + if !candidates.iter().any(|candidate| { + candidate.provider_id == PROVIDER_ID + && candidate.endpoint_id == ENDPOINT_ID + && candidate.key_id == PROVIDER_KEY_ID + }) { + return Err("seeded Responses WebSocket candidate is not visible".into()); + } + drop(backends); + Ok(()) +} + +async fn seed_provider_catalog( + backends: &DataBackends, + upstream_base_url: &str, + fixture: ProviderFixture, +) -> Result<(), BoxError> { + let writer = backends + .write() + .provider_catalog() + .ok_or("provider catalog writer unavailable")?; + writer + .create_provider( + &StoredProviderCatalogProvider::new( + PROVIDER_ID.to_string(), + "Responses WebSocket E2E".to_string(), + None, + fixture.provider_type().to_string(), + )? + .with_transport_fields( + true, + false, + false, + None, + Some(0), + None, + Some(30.0), + Some(10.0), + Some(json!({"responses_websocket": {"enabled": true}})), + ), + None, + ) + .await?; + writer + .create_endpoint( + &StoredProviderCatalogEndpoint::new( + ENDPOINT_ID.to_string(), + PROVIDER_ID.to_string(), + "openai:responses".to_string(), + Some("openai".to_string()), + Some("responses".to_string()), + true, + )? + .with_transport_fields( + upstream_base_url.trim_end_matches('/').to_string(), + None, + None, + Some(0), + Some("/v1/responses".to_string()), + None, + None, + None, + )?, + ) + .await?; + writer + .create_key(&catalog_key(PROVIDER_KEY_ID, PROVIDER_API_KEY, fixture)?) + .await?; + if fixture.has_alternate_key() { + writer + .create_key(&catalog_key( + ALTERNATE_PROVIDER_KEY_ID, + ALTERNATE_PROVIDER_API_KEY, + fixture, + )?) + .await?; + } + Ok(()) +} + +/// 一把健康、可服务本用例模型的 key。 +/// +/// codex 类型的候选要求 `auth_type = oauth`(见 candidate_selection 的 +/// provider_type 约束),所以配额重试夹具走 oauth,凭证是一份未过期的 +/// access_token。 +fn catalog_key( + key_id: &str, + secret: &str, + fixture: ProviderFixture, +) -> Result { + let oauth = fixture.has_alternate_key(); + let auth_type = if oauth { "oauth" } else { "api_key" }; + let auth_config = if oauth { + Some(encrypt_python_fernet_plaintext( + DEVELOPMENT_ENCRYPTION_KEY, + &json!({ + "access_token": secret, + "refresh_token": format!("{secret}-refresh"), + "account_id": format!("{key_id}-account"), + "expires_at": FAR_FUTURE_UNIX_SECS, + }) + .to_string(), + )?) + } else { + None + }; + Ok(StoredProviderCatalogKey::new( + key_id.to_string(), + PROVIDER_ID.to_string(), + "Responses WebSocket E2E".to_string(), + auth_type.to_string(), + Some(json!({"streaming": true})), + true, + )? + .with_transport_fields( + Some(json!(["openai:responses"])), + encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, secret)?, + auth_config, + None, + Some(json!({"openai:responses": 1})), + Some(json!([PUBLIC_MODEL, UPSTREAM_MODEL])), + None, + None, + None, + )? + .with_health_fields( + Some(json!({"openai:responses": {"status": "healthy"}})), + Some(json!({"openai:responses": {"state": "closed"}})), + )) +} + +async fn seed_models(backends: &DataBackends) -> Result<(), BoxError> { + let writer = backends + .write() + .global_models() + .ok_or("global model writer unavailable")?; + writer + .create_admin_global_model(&CreateAdminGlobalModelRecord::new( + GLOBAL_MODEL_ID.to_string(), + PUBLIC_MODEL.to_string(), + PUBLIC_MODEL.to_string(), + true, + Some(0.0), + None, + Some(json!({"streaming": true, "chat": true})), + Some(json!({"model_mappings": [UPSTREAM_MODEL]})), + )?) + .await?; + writer + .create_admin_provider_model(&UpsertAdminProviderModelRecord::new( + PROVIDER_MODEL_ID.to_string(), + PROVIDER_ID.to_string(), + GLOBAL_MODEL_ID.to_string(), + UPSTREAM_MODEL.to_string(), + Some(json!([{ + "name": UPSTREAM_MODEL, + "priority": 0, + "api_formats": ["openai:responses"], + "endpoint_ids": [ENDPOINT_ID] + }])), + Some(0.0), + None, + Some(false), + Some(false), + Some(true), + Some(false), + Some(false), + true, + true, + Some(json!({"responses_websocket_e2e": true})), + )?) + .await?; + Ok(()) +} + +async fn seed_user(backends: &DataBackends) -> Result { + let users = backends.read().users().ok_or("user reader unavailable")?; + let user = users + .create_local_auth_user_with_settings( + Some("responses-ws-e2e@example.test".to_string()), + true, + "responses-ws-e2e".to_string(), + "disabled-password".to_string(), + "user".to_string(), + Some(vec![PROVIDER_ID.to_string()]), + Some(vec!["openai:responses".to_string()]), + Some(vec![PUBLIC_MODEL.to_string()]), + None, + ) + .await? + .ok_or("failed to create E2E user")?; + let wallets = backends + .read() + .wallets() + .ok_or("wallet reader unavailable")?; + wallets + .initialize_auth_user_wallet(&user.id, 0.0, true) + .await?; + Ok(user.id) +} + +async fn seed_client_api_key(backends: &DataBackends, user_id: &str) -> Result<(), BoxError> { + backends + .write() + .auth_api_keys() + .ok_or("auth API key writer unavailable")? + .create_standalone_api_key(CreateStandaloneApiKeyRecord { + user_id: user_id.to_string(), + api_key_id: API_KEY_ID.to_string(), + key_hash: sha256_hex(CLIENT_API_KEY), + key_encrypted: Some(CLIENT_API_KEY.to_string()), + name: Some("Responses WebSocket E2E".to_string()), + allowed_providers: Some(vec![PROVIDER_ID.to_string()]), + allowed_api_formats: Some(vec!["openai:responses".to_string()]), + allowed_models: Some(vec![PUBLIC_MODEL.to_string()]), + ip_rules: None, + rate_limit: Some(0), + concurrent_limit: None, + force_capabilities: None, + is_active: true, + expires_at_unix_secs: None, + auto_delete_on_expiry: false, + total_requests: 0, + total_tokens: 0, + total_cost_usd: 0.0, + }) + .await?; + backends + .read() + .wallets() + .ok_or("wallet reader unavailable")? + .initialize_auth_api_key_wallet(API_KEY_ID, 0.0, true) + .await?; + if backends + .read() + .wallets() + .ok_or("wallet reader unavailable")? + .find(WalletLookupKey::ApiKeyId(API_KEY_ID)) + .await? + .is_none() + { + return Err("failed to initialize E2E API key wallet".into()); + } + Ok(()) +} + +/// 打开 chat PII 脱敏:系统模块开关 + 这把 client key 的 feature 开关。 +/// +/// 规则集刻意不写:缺省即内置规则(含 email 规则),和生产上「只打开开关」的最小 +/// 配置一致。 +async fn seed_chat_pii_redaction(backends: &DataBackends) -> Result<(), BoxError> { + backends + .upsert_system_config_entry("module.chat_pii_redaction.enabled", &json!(true), None) + .await?; + backends + .write() + .auth_api_keys() + .ok_or("auth API key writer unavailable")? + .set_standalone_api_key_feature_settings( + API_KEY_ID, + Some(json!({"chat_pii_redaction": {"enabled": true}})), + ) + .await? + .ok_or("failed to enable chat PII redaction on the E2E API key")?; + Ok(()) +} + +fn sha256_hex(value: &str) -> String { + let mut hasher = sha2::Sha256::new(); + hasher.update(value.as_bytes()); + format!("{:x}", hasher.finalize()) +} 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) diff --git a/docs/operations/codex-responses-websocket-probe.md b/docs/operations/codex-responses-websocket-probe.md new file mode 100644 index 000000000..b6987af1e --- /dev/null +++ b/docs/operations/codex-responses-websocket-probe.md @@ -0,0 +1,158 @@ +# Codex Responses WebSocket probe + +`aether-codex-ws-probe` is a P0 compatibility probe for a Codex-compatible +Responses WebSocket upstream. It verifies two sequential `response.create` +warmups on one socket, with the second request continuing from the first +response ID. + +The command, environment variables, JSON report shape, and Codex-specific +handshake headers remain stable. It now shares only the protocol-driving core +with the separate [OpenAI Responses WebSocket probe](openai-responses-websocket-probe.md); +the two probes intentionally retain independent authentication profiles and +provider-specific assertions. + +The probe is intentionally not a production proxy. It does not persist, +refresh, log, or print credentials, account IDs, response IDs, request bodies, +or response bodies. + +## Prerequisites + +Use a dedicated, non-production Codex test account. Rotate any credential that +has been pasted into a chat, terminal history, issue, or source file before +using this probe. + +Set these values only in the process environment or your secret manager: + +```bash +export AETHER_CODEX_WS_PROBE_URL='wss://your-codex-upstream.example/backend-api/codex/responses' +export AETHER_CODEX_WS_PROBE_ACCESS_TOKEN='your-short-lived-access-token' +export AETHER_CODEX_WS_PROBE_ACCOUNT_ID='your-account-id' +export AETHER_CODEX_WS_PROBE_MODEL='your-codex-model' +``` + +The endpoint must use `ws://` or `wss://`, with no credentials, query string, +or fragment. The access token is accepted only through +`AETHER_CODEX_WS_PROBE_ACCESS_TOKEN`; there is deliberately no command-line +flag for it. + +## Run + +```bash +cargo run -p aether-gateway --bin aether-codex-ws-probe +``` + +Use `--url` to override only the endpoint and `--timeout-secs` to set a +per-turn receive timeout (1–120 seconds): + +```bash +cargo run -p aether-gateway --bin aether-codex-ws-probe -- \ + --url 'wss://your-codex-upstream.example/backend-api/codex/responses' \ + --timeout-secs 30 +``` + +The probe emits one JSON line. A successful run has +`"continuation_confirmed":true`; its event and header fields contain names +only, never values. A failure emits a stable error code such as +`"handshake_failed"`, `"upstream_error_event"`, or +`"response_id_not_observed"`. + +## Interpretation + +A successful probe establishes that the selected upstream accepts the +Responses WebSocket handshake and retains continuation state on one socket. +It does not establish that all Codex models, account plans, or tunnel egress +paths are supported. In particular, the current `aether-tunnel` HTTP relay +does not forward WebSocket upgrades, so a successful direct probe is a +prerequisite rather than tunnel support. + +## Gateway bridge + +The gateway exposes WebSocket mode at the same public Responses path: + +```text +wss:///v1/responses +``` + +It is disabled by default per provider. In **添加提供商** or **编辑提供商**, +enable **Responses WebSocket 模式** under **功能开关** only after the selected +upstream has passed a compatible WebSocket probe. The setting takes effect for +new WebSocket connections without a gateway restart. It is available to every +provider type; candidate planning still requires a selected +`openai:responses` endpoint. + +Authenticate the upgrade request with the normal Aether API key. The first +client frame must be a text JSON `response.create` containing a non-empty +`model`. Aether then applies its regular Responses candidate selection, but +accepts only an eligible, WebSocket-enabled endpoint using `openai:responses`. +It opens an upstream WebSocket using the selected provider key. + +The selected provider's model mapping and request headers are applied to every +turn, along with the rest of that candidate's provider-body normalization: +model-directive patches, endpoint body rules, and the Codex body contract +(unsupported-field stripping, `store: false`, `tool_choice` defaulting). A +continuation turn is normalized against the binding it is pinned to rather than +being re-planned, so it can never move to another provider key. +`previous_response_id` and `generate` are re-applied after normalization +because they are WebSocket protocol state that the provider body contract +otherwise strips. `stream` and `background` are removed because they are HTTP +transport fields, not WebSocket-mode fields. If a later `response.create` changes the +public model, Aether runs access checks and candidate planning again. It keeps +the existing upstream when the same target remains eligible, or transparently +replaces the upstream between responses when the selected target changes. +Overlapping responses on one client socket remain rejected. + +Each `response.create` is tracked as an independent Aether logical request: +it receives its own request/candidate identity, usage lifecycle, and terminal +audit record. `response.completed`, `response.failed`, +`response.incomplete`, `response.cancelled`, client disconnects, and upstream +transport failures all settle that turn through the existing stream reporting +path. + +Example client setup: + +```python +from websocket import create_connection +import json +import os + +ws = create_connection( + "wss://gateway.example/v1/responses", + header=[f"Authorization: Bearer {os.environ['AETHER_API_KEY']}"], +) +ws.send(json.dumps({ + "type": "response.create", + "model": "your-public-model", + "store": False, + "input": "Explain this repository.", +})) +``` + +### Operating limits + +- Maximum frame and message size: 16 MiB. +- An idle connection must send its first `response.create` within 60 seconds. +- A connection is closed after 60 minutes; reconnect before then for long runs. +- Each `response.create` must receive its first upstream event within the + selected provider's `stream_first_byte_timeout` (30 seconds by default), + and finish within its `request_timeout` (20 minutes by default). Aether + sends `responses_websocket_first_event_timeout` or + `responses_websocket_turn_timeout` and closes the bound socket when either + deadline expires. +- Responses are sequential; no multiplexing is supported on one socket. +- Each `response.create` consumes the normal Aether user/API-key RPM budget. +- Same-model turns stay on the bound provider key. A model change is planned + again and can rebind the upstream between completed turns when necessary. +- Direct provider proxy settings are honored through the selected transport + profile. Tunnel-mode proxy nodes are not supported for this bridge yet. + +Usage and audit finalization now runs for every accepted `response.create`. +Existing usage body-capture and header-redaction policies apply to the resulting +records. Newly created WebSocket usage records expose `is_websocket=true`, and +the usage-record type column renders them as `WS`. For diagnosis, enable debug +logging for `aether_gateway::handlers::proxy::responses_ws`; event logs contain +only the event type and frame size, never request or response contents. Codex +quota-extension logs remain under `aether_gateway::handlers::proxy::codex_ws`. +Every WebSocket-specific log carries `transport="websocket"` and +`websocket=true`; keep `log_type` for its existing access/event/ops +classification, and render the transport flag as a `WS` label in a log viewer +if desired. diff --git a/docs/operations/openai-responses-websocket-probe.md b/docs/operations/openai-responses-websocket-probe.md new file mode 100644 index 000000000..b7a33c7d8 --- /dev/null +++ b/docs/operations/openai-responses-websocket-probe.md @@ -0,0 +1,61 @@ +# OpenAI Responses WebSocket probe + +`aether-openai-responses-ws-probe` verifies the official OpenAI Responses +WebSocket protocol using standard API-key Bearer authentication. It sends two +sequential `response.create` warmups on one socket, chaining the second from +the first response ID with `previous_response_id`. + +It shares its protocol-driving core with the Codex probe, but it does **not** +send Codex account headers or require Codex quota events. This makes it the +compatibility gate for Aether's standard Responses WebSocket adapter, rather +than a replacement for the Codex probe. + +## Prerequisites + +Use a dedicated API project and a model that your key can access. Keep values +only in your process environment or secret manager: + +```bash +export AETHER_OPENAI_WS_PROBE_API_KEY='your-api-key' +export AETHER_OPENAI_WS_PROBE_MODEL='your-openai-model' +``` + +The default endpoint is the official Responses WebSocket endpoint: + +```text +wss://api.openai.com/v1/responses +``` + +To test a compatible endpoint explicitly, set +`AETHER_OPENAI_WS_PROBE_URL` or pass `--url`. The endpoint must use `ws://` or +`wss://` and may not contain credentials, a query string, or a fragment. The +API key has no command-line flag and is never printed. + +## Run + +```bash +cargo run -p aether-gateway --bin aether-openai-responses-ws-probe +``` + +For an explicit endpoint and timeout: + +```bash +cargo run -p aether-gateway --bin aether-openai-responses-ws-probe -- \ + --url 'wss://api.openai.com/v1/responses' \ + --timeout-secs 30 +``` + +The probe uses `generate:false`, so the warmups prepare continuation state but +do not request model output. A successful JSON report contains +`"continuation_confirmed":true`; header and event arrays contain names only, +never credentials, response IDs, request bodies, or response bodies. + +## Interpretation + +Success establishes that this key, model, and endpoint support the Responses +WebSocket handshake plus an in-socket continuation. It does not establish +support for every model, tool, service tier, proxy path, or Aether provider +configuration. Treat a successful direct probe as a prerequisite before +enabling **Responses WebSocket mode** for the matching Aether provider. + +For protocol details, see the official [WebSocket Mode guide](https://developers.openai.com/api/docs/guides/websocket-mode). diff --git a/frontend/src/api/endpoints/providers.ts b/frontend/src/api/endpoints/providers.ts index f112b5f87..2ff04293b 100644 --- a/frontend/src/api/endpoints/providers.ts +++ b/frontend/src/api/endpoints/providers.ts @@ -54,6 +54,7 @@ function normalizeProviderSummary( kiro_simulated_cache_enabled: provider.kiro_simulated_cache_enabled ?? false, max_transfer_count: provider.max_transfer_count ?? 0, max_transfer_timeout_seconds: provider.max_transfer_timeout_seconds ?? 0, + responses_websocket_enabled: provider.responses_websocket_enabled ?? false, } } @@ -114,6 +115,7 @@ export async function updateProvider( website: string provider_priority: number keep_priority_on_conversion: boolean + responses_websocket_enabled: boolean billing_type: 'monthly_quota' | 'pay_as_you_go' | 'free_tier' monthly_quota_usd: number quota_reset_day: number @@ -156,6 +158,7 @@ export async function createProvider( quota_expires_at?: string provider_priority?: number keep_priority_on_conversion?: boolean + responses_websocket_enabled?: boolean is_active?: boolean max_retries?: number max_transfer_count?: number diff --git a/frontend/src/api/endpoints/types/provider.ts b/frontend/src/api/endpoints/types/provider.ts index fe1c49ce1..a616f0015 100644 --- a/frontend/src/api/endpoints/types/provider.ts +++ b/frontend/src/api/endpoints/types/provider.ts @@ -180,8 +180,13 @@ export interface ChatPiiRedactionProviderConfig { enabled: boolean } +export interface ResponsesWebSocketProviderConfig { + enabled: boolean +} + export interface ProviderConfig { chat_pii_redaction?: ChatPiiRedactionProviderConfig + responses_websocket?: ResponsesWebSocketProviderConfig pool_advanced?: PoolAdvancedConfig failover_rules?: FailoverRulesConfig claude_code_advanced?: ClaudeCodeAdvancedConfig @@ -901,6 +906,7 @@ export interface ProviderWithEndpointsSummary { ops_configured: boolean // 是否配置了扩展操作(余额监控等) ops_architecture_id?: string // 扩展操作使用的架构 ID(如 cubence, anyrouter) kiro_simulated_cache_enabled?: boolean + responses_websocket_enabled?: boolean ops_quota_alert_enabled?: boolean created_at: string updated_at: string diff --git a/frontend/src/api/me.ts b/frontend/src/api/me.ts index e4a312cb5..0ccc0bde3 100644 --- a/frontend/src/api/me.ts +++ b/frontend/src/api/me.ts @@ -73,6 +73,7 @@ export interface UsageRecordDetail { updated_at?: string | null response_time_updated_at?: string | null is_stream: boolean + is_websocket?: boolean upstream_is_stream?: boolean client_requested_stream?: boolean client_is_stream?: boolean @@ -368,6 +369,7 @@ export const meApi = { api_format?: string | null endpoint_api_format?: string | null is_stream?: boolean | null + is_websocket?: boolean | null upstream_is_stream?: boolean | null client_requested_stream?: boolean | null client_is_stream?: boolean | null diff --git a/frontend/src/api/usage.ts b/frontend/src/api/usage.ts index 425a78c26..0028146e7 100644 --- a/frontend/src/api/usage.ts +++ b/frontend/src/api/usage.ts @@ -33,6 +33,7 @@ export interface UsageRecord { first_byte_time_ms?: number | null end_to_end_time_ms?: number | null end_to_end_first_byte_time_ms?: number | null + is_websocket?: boolean created_at: string updated_at?: string | null response_time_updated_at?: string | null @@ -570,6 +571,7 @@ export const usageApi = { api_format?: string | null endpoint_api_format?: string | null is_stream?: boolean | null + is_websocket?: boolean | null upstream_is_stream?: boolean | null client_requested_stream?: boolean | null client_is_stream?: boolean | null diff --git a/frontend/src/features/providers/components/ProviderFormDialog.vue b/frontend/src/features/providers/components/ProviderFormDialog.vue index c83275b8e..1d7aaf577 100644 --- a/frontend/src/features/providers/components/ProviderFormDialog.vue +++ b/frontend/src/features/providers/components/ProviderFormDialog.vue @@ -333,6 +333,19 @@ /> +
+
+ {{ legacyT('Responses WebSocket 模式') }} +

+ {{ legacyT('允许此提供商处理标准 Responses API WebSocket 请求。仅在已验证兼容性后启用。') }} +

+
+ +
+
{{ legacyT('敏感信息保护') }} @@ -455,6 +468,8 @@ const form = ref({ pool_mode_enabled: false, // Kiro 专属配置 kiro_simulated_cache_enabled: false, + // Responses WebSocket 配置 + responses_websocket_enabled: false, }) // 重置表单 @@ -485,6 +500,8 @@ function resetForm() { pool_mode_enabled: false, // Kiro 专属配置 kiro_simulated_cache_enabled: false, + // Responses WebSocket 配置 + responses_websocket_enabled: false, } } @@ -519,6 +536,8 @@ function loadProviderData() { pool_mode_enabled: poolAdvanced !== null, // Kiro 专属配置 kiro_simulated_cache_enabled: props.provider.kiro_simulated_cache_enabled ?? false, + // Responses WebSocket 配置 + responses_websocket_enabled: props.provider.responses_websocket_enabled ?? false, } } @@ -575,6 +594,7 @@ const handleSubmit = async () => { quota_last_reset_at: quotaLastResetAt, quota_expires_at: quotaExpiresAt, keep_priority_on_conversion: form.value.keep_priority_on_conversion, + responses_websocket_enabled: form.value.responses_websocket_enabled, is_active: form.value.is_active, // 请求配置 max_retries: form.value.max_retries ?? undefined, diff --git a/frontend/src/features/providers/components/__tests__/ProviderFormDialog.responses-websocket.spec.ts b/frontend/src/features/providers/components/__tests__/ProviderFormDialog.responses-websocket.spec.ts new file mode 100644 index 000000000..1e38a7543 --- /dev/null +++ b/frontend/src/features/providers/components/__tests__/ProviderFormDialog.responses-websocket.spec.ts @@ -0,0 +1,18 @@ +import { readFileSync } from 'node:fs' +import { resolve } from 'node:path' +import { describe, expect, it } from 'vitest' + +function readSource(path: string): string { + return readFileSync(resolve(process.cwd(), path), 'utf8') +} + +describe('ProviderFormDialog Responses WebSocket switch', () => { + it('displays and submits the switch for every provider type', () => { + const source = readSource('src/features/providers/components/ProviderFormDialog.vue') + + expect(source).toContain('Responses WebSocket 模式') + expect(source).toContain('responses_websocket_enabled') + expect(source).toContain('responses_websocket_enabled: form.value.responses_websocket_enabled') + expect(source).not.toContain("v-if=\"form.provider_type === 'codex'\"") + }) +}) diff --git a/frontend/src/features/usage/components/UsageRecordsTable.vue b/frontend/src/features/usage/components/UsageRecordsTable.vue index ea849110a..d07b283b2 100644 --- a/frontend/src/features/usage/components/UsageRecordsTable.vue +++ b/frontend/src/features/usage/components/UsageRecordsTable.vue @@ -277,6 +277,14 @@ > 取消 + + WS + 已取消 + + WS + { expect(root.textContent).toContain('gpt-5') }) + it('shows WS in the type column for completed WebSocket usage records', () => { + const root = mountUsageRecordsTable([buildRecord({ + is_websocket: true, + status: 'completed', + })]) + + const badges = [...root.querySelectorAll('[data-usage-transport="websocket"]')] + expect(badges.length).toBeGreaterThan(0) + expect(badges.every(badge => badge.textContent?.trim() === 'WS')).toBe(true) + }) + it('shows reasoning effort next to the model name', () => { const root = mountUsageRecordsTable([buildRecord({ requested_reasoning_effort: 'xhigh', diff --git a/frontend/src/features/usage/composables/useUsageData.ts b/frontend/src/features/usage/composables/useUsageData.ts index 5103efd71..2d68c340f 100644 --- a/frontend/src/features/usage/composables/useUsageData.ts +++ b/frontend/src/features/usage/composables/useUsageData.ts @@ -647,6 +647,7 @@ export function useUsageData(options: UseUsageDataOptions) { ? (record.image_progress ?? existing.image_progress) : existing.image_progress, is_stream: upstreamIsStream, + is_websocket: mergeBooleanTrueWins(existing.is_websocket, record.is_websocket), upstream_is_stream: upstreamIsStream, client_requested_stream: clientRequestedStream, client_is_stream: clientIsStream, diff --git a/frontend/src/features/usage/types.ts b/frontend/src/features/usage/types.ts index cb7433876..c13d0176d 100644 --- a/frontend/src/features/usage/types.ts +++ b/frontend/src/features/usage/types.ts @@ -120,6 +120,7 @@ export interface UsageRecord { end_to_end_time_ms?: number | null // 客户端从请求进入网关到完成的总耗时 end_to_end_first_byte_time_ms?: number | null // 客户端从请求进入网关到首字节的耗时 is_stream: boolean + is_websocket?: boolean upstream_is_stream?: boolean client_requested_stream?: boolean client_is_stream?: boolean diff --git a/frontend/src/features/usage/utils/__tests__/status.spec.ts b/frontend/src/features/usage/utils/__tests__/status.spec.ts index 93cf8e886..a94629b33 100644 --- a/frontend/src/features/usage/utils/__tests__/status.spec.ts +++ b/frontend/src/features/usage/utils/__tests__/status.spec.ts @@ -6,6 +6,7 @@ import { hasUsageRetry, isUsageRecordFailed, isUsageRecordSuccessful, + isUsageWebSocket, mapRequestStatusToTimelineStatus, normalizeRequestStatus, resolveDisplayRequestStatus, @@ -177,6 +178,12 @@ describe('usage status helpers', () => { expect(hasUsageRetry(buildUsageRecord({ has_retry: undefined }))).toBe(false) }) + it('recognizes persisted WebSocket usage records', () => { + expect(isUsageWebSocket(buildUsageRecord({ is_websocket: true }))).toBe(true) + expect(isUsageWebSocket(buildUsageRecord({ is_websocket: false }))).toBe(false) + expect(isUsageWebSocket(buildUsageRecord({ is_websocket: undefined }))).toBe(false) + }) + it('prefers symmetric stream aliases when present', () => { expect(formatUsageStreamLabel(buildUsageRecord({ is_stream: true, diff --git a/frontend/src/features/usage/utils/status.ts b/frontend/src/features/usage/utils/status.ts index 9ed987fc2..e93f2b948 100644 --- a/frontend/src/features/usage/utils/status.ts +++ b/frontend/src/features/usage/utils/status.ts @@ -49,6 +49,12 @@ export function hasUsageRetry( return record.has_retry === true } +export function isUsageWebSocket( + record: Pick +): boolean { + return record.is_websocket === true +} + export function resolveUsageStreamModes( record: Pick< UsageRecord, diff --git a/frontend/src/views/shared/Usage.vue b/frontend/src/views/shared/Usage.vue index 70a64ae20..92e313732 100644 --- a/frontend/src/views/shared/Usage.vue +++ b/frontend/src/views/shared/Usage.vue @@ -574,6 +574,9 @@ async function pollActiveRequests() { record.is_stream = update.is_stream record.upstream_is_stream = update.is_stream } + if (typeof update.is_websocket === 'boolean') { + record.is_websocket = record.is_websocket === true || update.is_websocket + } if (typeof update.client_is_stream === 'boolean') { record.client_is_stream = update.client_is_stream record.client_requested_stream = update.client_is_stream