diff --git a/.github/workflows/rust-ci.yml b/.github/workflows/rust-ci.yml index af8295df2..c8943c245 100644 --- a/.github/workflows/rust-ci.yml +++ b/.github/workflows/rust-ci.yml @@ -200,13 +200,13 @@ jobs: RUSTFLAGS: "-C link-arg=-fuse-ld=mold" run: cargo nextest run -p aether-gateway --lib - - name: Test bin + - name: Test bins env: RUSTC_WRAPPER: sccache SCCACHE_GHA_ENABLED: "true" RUST_MIN_STACK: "16777216" RUSTFLAGS: "-C link-arg=-fuse-ld=mold" - run: cargo nextest run -p aether-gateway --bin aether-gateway + run: cargo nextest run -p aether-gateway --bins - name: Show sccache stats if: always() diff --git a/README.md b/README.md index de2e3f10d..207819a70 100644 --- a/README.md +++ b/README.md @@ -137,6 +137,8 @@ Aether Tunnel 是配套的正向代理节点,部署在海外 VPS 上,为墙 - Embeddings: [OpenAI compatible `POST /v1/embeddings`](docs/api/embeddings.md) - Rerank: [OpenAI/Jina compatible `POST /v1/rerank`](docs/api/rerank.md) +- Responses WebSocket mode: [protocol and Aether behavior](docs/WebSocket-Mode.md) +- WebSocket probes: [Codex](docs/operations/codex-responses-websocket-probe.md) · [OpenAI Responses](docs/operations/openai-responses-websocket-probe.md) ## 环境变量 diff --git a/apps/aether-gateway/src/ai_serving/mod.rs b/apps/aether-gateway/src/ai_serving/mod.rs index 0a575243c..1e8dc7e50 100644 --- a/apps/aether-gateway/src/ai_serving/mod.rs +++ b/apps/aether-gateway/src/ai_serving/mod.rs @@ -64,7 +64,7 @@ pub(crate) use self::planner::{ GatewayAuthApiKeySnapshot, GatewayProviderTransportSnapshot, LocalExecutionAttemptSource, LocalExecutionCandidateKind, LocalResolvedOAuthRequestAuth, PlannerAppState, ResponsesWebSocketBodyNormalization, ResponsesWebSocketDecision, - SkippedLocalExecutionCandidate, + ResponsesWebSocketPinnedCandidate, SkippedLocalExecutionCandidate, }; pub(crate) use self::pure::*; pub(crate) use self::response_history::{ diff --git a/apps/aether-gateway/src/ai_serving/planner/decision_input.rs b/apps/aether-gateway/src/ai_serving/planner/decision_input.rs index 007eeeb06..0080375f1 100644 --- a/apps/aether-gateway/src/ai_serving/planner/decision_input.rs +++ b/apps/aether-gateway/src/ai_serving/planner/decision_input.rs @@ -351,6 +351,7 @@ fn apply_codex_oauth_fingerprint_convergence_to_decision( struct GatewayAuthenticatedDecisionInputPort<'a> { state: PlannerAppState<'a>, now_unix_secs: u64, + auth_snapshot_override: Option, model_directive_policy: &'a crate::system_features::ModelDirectivePolicySnapshot, model_directive_base_model: Option, } @@ -367,6 +368,17 @@ impl AiAuthenticatedDecisionInputPort for GatewayAuthenticatedDecisionInputPort< &self, auth_context: &Self::AuthContext, ) -> Result, Self::Error> { + if let Some(snapshot) = self.auth_snapshot_override.as_ref() { + if snapshot.user_id != auth_context.user_id + || snapshot.api_key_id != auth_context.api_key_id + { + return Err(GatewayError::Internal( + "WebSocket auth snapshot identity does not match its control decision" + .to_string(), + )); + } + return Ok(Some(snapshot.clone())); + } self.state .read_auth_api_key_snapshot( &auth_context.user_id, @@ -750,6 +762,27 @@ pub(crate) async fn resolve_local_authenticated_decision_input( requested_model_api_format: Option<&str>, explicit_required_capabilities: Option<&serde_json::Value>, model_directive_policy: &crate::system_features::ModelDirectivePolicySnapshot, +) -> Result, GatewayError> { + resolve_local_authenticated_decision_input_with_snapshot( + state, + auth_context, + None, + requested_model, + requested_model_api_format, + explicit_required_capabilities, + model_directive_policy, + ) + .await +} + +pub(crate) async fn resolve_local_authenticated_decision_input_with_snapshot( + state: &AppState, + auth_context: ExecutionRuntimeAuthContext, + auth_snapshot_override: Option, + requested_model: Option<&str>, + requested_model_api_format: Option<&str>, + explicit_required_capabilities: Option<&serde_json::Value>, + model_directive_policy: &crate::system_features::ModelDirectivePolicySnapshot, ) -> Result, GatewayError> { let model_directive_base_model = match (requested_model, requested_model_api_format) { (Some(model), Some(api_format)) => model_directive_policy @@ -761,6 +794,7 @@ pub(crate) async fn resolve_local_authenticated_decision_input( let port = GatewayAuthenticatedDecisionInputPort { state: PlannerAppState::new(state), now_unix_secs: current_unix_secs(), + auth_snapshot_override, model_directive_policy, model_directive_base_model, }; @@ -1066,6 +1100,52 @@ mod tests { } } + #[tokio::test] + async fn explicit_auth_snapshot_override_does_not_fall_back_to_the_planner_cache() { + // AppState::new has no auth snapshot repository. Without the explicit + // override this resolver returns None; a WebSocket strong snapshot must + // therefore be the exact value used to build the planner input. + let state = AppState::new().expect("test state should build"); + let mut strong_snapshot = sample_auth_snapshot(); + strong_snapshot.api_key_allowed_models = Some(vec!["gpt-live-only".to_string()]); + + let resolved = resolve_local_authenticated_decision_input_with_snapshot( + &state, + sample_auth_context(), + Some(strong_snapshot.clone()), + Some("gpt-live-only"), + Some("openai:responses"), + None, + &Default::default(), + ) + .await + .expect("snapshot override should resolve") + .expect("the explicit snapshot should replace the missing cached value"); + + assert_eq!(resolved.auth_snapshot, strong_snapshot); + } + + #[tokio::test] + async fn explicit_auth_snapshot_override_rejects_an_identity_mismatch() { + let state = AppState::new().expect("test state should build"); + let mut wrong_snapshot = sample_auth_snapshot(); + wrong_snapshot.api_key_id = "another-key".to_string(); + + let error = resolve_local_authenticated_decision_input_with_snapshot( + &state, + sample_auth_context(), + Some(wrong_snapshot), + Some("gpt-live-only"), + Some("openai:responses"), + None, + &Default::default(), + ) + .await + .expect_err("a snapshot for another API key must never be injected"); + + assert!(matches!(error, GatewayError::Internal(message) if message.contains("identity"))); + } + #[tokio::test] async fn explicit_routing_attachment_authorizes_and_caches_per_principal() { let repository = Arc::new(InMemoryRoutingGroupRepository::default()); diff --git a/apps/aether-gateway/src/ai_serving/planner/mod.rs b/apps/aether-gateway/src/ai_serving/planner/mod.rs index 55068e03e..ddb8e5952 100644 --- a/apps/aether-gateway/src/ai_serving/planner/mod.rs +++ b/apps/aether-gateway/src/ai_serving/planner/mod.rs @@ -84,6 +84,7 @@ pub(crate) use self::standard::{ codex_model_capabilities_for_transport, maybe_build_responses_websocket_decision, set_local_openai_chat_execution_exhausted_diagnostic, validate_final_openai_provider_request, ResponsesWebSocketBodyNormalization, ResponsesWebSocketDecision, + ResponsesWebSocketPinnedCandidate, }; 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 9ff88b42a..426eb612f 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/mod.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/mod.rs @@ -50,6 +50,7 @@ pub(crate) use self::openai::{ 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, ResponsesWebSocketBodyNormalization, ResponsesWebSocketDecision, + ResponsesWebSocketPinnedCandidate, }; 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 1af6090ed..0e6111187 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 @@ -26,5 +26,5 @@ pub(crate) use responses::{ maybe_build_responses_websocket_decision, maybe_build_stream_local_openai_responses_decision_payload, maybe_build_sync_local_openai_responses_decision_payload, ResponsesWebSocketBodyNormalization, - ResponsesWebSocketDecision, + ResponsesWebSocketDecision, ResponsesWebSocketPinnedCandidate, }; diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision.rs index 4a72c8849..e3c4a1524 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision.rs @@ -9,7 +9,9 @@ pub(super) use self::payload::maybe_build_local_openai_responses_decision_payloa pub(super) use self::support::{ build_local_openai_responses_candidate_attempt_source, materialize_local_openai_responses_candidate_attempts, - resolve_local_openai_responses_decision_input, LocalOpenAiResponsesCandidateAttempt, - LocalOpenAiResponsesCandidateAttemptSource, LocalOpenAiResponsesDecisionInput, + resolve_local_openai_responses_decision_input, + resolve_local_openai_responses_decision_input_with_snapshot, + LocalOpenAiResponsesCandidateAttempt, LocalOpenAiResponsesCandidateAttemptSource, + LocalOpenAiResponsesDecisionInput, }; pub(super) use crate::ai_serving::LocalOpenAiResponsesSpec; diff --git a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/support.rs b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/support.rs index 3ad3674c0..104c83aa7 100644 --- a/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/support.rs +++ b/apps/aether-gateway/src/ai_serving/planner/standard/openai/responses/decision/support.rs @@ -22,6 +22,7 @@ use crate::ai_serving::planner::common::extract_standard_requested_model; use crate::ai_serving::planner::decision_input::{ attach_routing_policy_to_local_requested_model_input, build_local_requested_model_decision_input, resolve_local_authenticated_decision_input, + resolve_local_authenticated_decision_input_with_snapshot, }; use crate::ai_serving::planner::materialization_policy::{ build_local_candidate_persistence_policy, LocalCandidatePersistencePolicyKind, @@ -32,7 +33,8 @@ use crate::ai_serving::planner::CandidateFailureDiagnostic; use crate::ai_serving::{ ai_local_execution_contract_for_formats, extract_pool_sticky_session_token, openai_responses_request_operation, resolve_local_decision_execution_runtime_auth_context, - ExecutionRuntimeAuthContext, GatewayControlDecision, PlannerAppState, + ExecutionRuntimeAuthContext, GatewayAuthApiKeySnapshot, GatewayControlDecision, + PlannerAppState, }; use crate::client_session_affinity::client_session_affinity_from_parts; use crate::{AppState, GatewayError}; @@ -51,6 +53,21 @@ pub(crate) async fn resolve_local_openai_responses_decision_input( decision: &GatewayControlDecision, body_json: &serde_json::Value, plan_kind: &str, +) -> Result, GatewayError> { + resolve_local_openai_responses_decision_input_with_snapshot( + state, parts, trace_id, decision, body_json, plan_kind, None, + ) + .await +} + +pub(crate) async fn resolve_local_openai_responses_decision_input_with_snapshot( + state: &AppState, + parts: &http::request::Parts, + trace_id: &str, + decision: &GatewayControlDecision, + body_json: &serde_json::Value, + plan_kind: &str, + auth_snapshot_override: Option<&GatewayAuthApiKeySnapshot>, ) -> Result, GatewayError> { let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else { warn!( @@ -87,16 +104,28 @@ pub(crate) async fn resolve_local_openai_responses_decision_input( return Ok(None); }; - let resolved_input = match resolve_local_authenticated_decision_input( - state, - auth_context.clone(), - Some(requested_model.as_str()), - decision.auth_endpoint_signature.as_deref(), - None, - &decision.model_directive_policy, - ) - .await - { + let resolved_input = match if let Some(auth_snapshot) = auth_snapshot_override { + resolve_local_authenticated_decision_input_with_snapshot( + state, + auth_context.clone(), + Some(auth_snapshot.clone()), + Some(requested_model.as_str()), + decision.auth_endpoint_signature.as_deref(), + None, + &decision.model_directive_policy, + ) + .await + } else { + resolve_local_authenticated_decision_input( + state, + auth_context.clone(), + Some(requested_model.as_str()), + decision.auth_endpoint_signature.as_deref(), + None, + &decision.model_directive_policy, + ) + .await + } { Ok(Some(resolved_input)) => resolved_input, Ok(None) => { warn!( 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 c3b33c220..ead7fbd38 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 @@ -12,6 +12,49 @@ use crate::{AiExecutionDecision, AppState, GatewayError}; use aether_runtime_state::RuntimeLockLease; use std::collections::BTreeSet; +/// Releases a scheduler pool-key lease if WebSocket planning is cancelled +/// after candidate selection but before ownership reaches the turn lifecycle. +struct ResponsesWebSocketPlanningLeaseGuard { + state: AppState, + lease: Option, +} + +impl ResponsesWebSocketPlanningLeaseGuard { + fn new(state: &AppState, lease: Option<&RuntimeLockLease>) -> Self { + Self { + state: state.clone(), + lease: lease.cloned(), + } + } + + async fn release(mut self) { + // Keep the lease armed across the await. If the owner task is aborted + // or reaches its hard deadline while the runtime backend is stalled, + // Drop can still hand cleanup to a detached owner. + if release_responses_websocket_planning_lease(&self.state, self.lease.as_ref()).await { + self.lease = None; + } + } + + fn disarm(&mut self) { + self.lease = None; + } +} + +impl Drop for ResponsesWebSocketPlanningLeaseGuard { + fn drop(&mut self) { + let Some(lease) = self.lease.take() else { + return; + }; + let state = self.state.clone(); + if let Ok(handle) = tokio::runtime::Handle::try_current() { + handle.spawn(async move { + let _ = release_responses_websocket_planning_lease(&state, Some(&lease)).await; + }); + } + } +} + mod decision; mod plans; @@ -19,6 +62,7 @@ use self::decision::{ build_local_openai_responses_candidate_attempt_source, maybe_build_local_openai_responses_decision_payload_for_candidate, resolve_local_openai_responses_decision_input, + resolve_local_openai_responses_decision_input_with_snapshot, }; use self::plans::{ build_local_stream_attempt_source, build_local_stream_plan_and_reports, @@ -187,6 +231,45 @@ pub(crate) struct ResponsesWebSocketDecision { pub(crate) normalization: ResponsesWebSocketBodyNormalization, } +/// The scheduler identity a continuation is allowed to reuse. +/// +/// A `previous_response_id` chain cannot move to another provider connection, +/// but it still has to pass the current scheduler runtime checks on every +/// turn. The planner uses this identity as a filter rather than selecting an +/// arbitrary eligible replacement. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ResponsesWebSocketPinnedCandidate { + provider_id: String, + endpoint_id: String, + key_id: String, +} + +impl ResponsesWebSocketPinnedCandidate { + pub(crate) fn from_decision(decision: &AiExecutionDecision) -> Option { + Some(Self { + provider_id: non_empty_decision_identity(decision.provider_id.as_deref())?, + endpoint_id: non_empty_decision_identity(decision.endpoint_id.as_deref())?, + key_id: non_empty_decision_identity(decision.key_id.as_deref())?, + }) + } + + fn matches( + &self, + candidate: &aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate, + ) -> bool { + candidate.provider_id == self.provider_id + && candidate.endpoint_id == self.endpoint_id + && candidate.key_id == self.key_id + } +} + +fn non_empty_decision_identity(value: Option<&str>) -> Option { + value + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) +} + /// Everything needed to re-run provider-body normalization for the candidate a /// socket is already bound to. /// @@ -321,21 +404,24 @@ pub(crate) async fn maybe_build_responses_websocket_decision( parts: &http::request::Parts, trace_id: &str, decision: &GatewayControlDecision, + auth_snapshot: Option<&crate::ai_serving::GatewayAuthApiKeySnapshot>, body_json: &serde_json::Value, excluded_key_ids: Option<&BTreeSet>, excluded_codex_account_ids: Option<&BTreeSet>, + pinned_candidate: Option<&ResponsesWebSocketPinnedCandidate>, ) -> 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( + let Some(input) = resolve_local_openai_responses_decision_input_with_snapshot( state, parts, trace_id, decision, body_json, spec.decision_kind, + auth_snapshot, ) .await? else { @@ -348,18 +434,28 @@ pub(crate) async fn maybe_build_responses_websocket_decision( .await?; while let Some(attempt) = source.next_attempt().await? { - let pool_key_lease = attempt.eligible.orchestration.pool_key_lease.clone(); + // `next_attempt` may return with a distributed pool-key lease. Arm a + // guard before the first await so owner-task timeout/cancellation + // cannot strand that lease until its server-side TTL expires. + let mut planning_lease = ResponsesWebSocketPlanningLeaseGuard::new( + state, + attempt.eligible.orchestration.pool_key_lease.as_ref(), + ); + if pinned_candidate.is_some_and(|pinned| !pinned.matches(&attempt.eligible.candidate)) { + planning_lease.release().await; + continue; + } 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; + planning_lease.release().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; + planning_lease.release().await; continue; }; // Captured before `attempt` is consumed so a later continuation turn can @@ -373,11 +469,11 @@ pub(crate) async fn maybe_build_responses_websocket_decision( { Ok(Some(payload)) => payload, Ok(None) => { - release_responses_websocket_planning_lease(state, pool_key_lease.as_ref()).await; + planning_lease.release().await; continue; } Err(error) => { - release_responses_websocket_planning_lease(state, pool_key_lease.as_ref()).await; + planning_lease.release().await; return Err(error); } }; @@ -393,7 +489,7 @@ pub(crate) async fn maybe_build_responses_websocket_decision( .is_some_and(|account_ids| account_ids.contains(account_id)) }) { - release_responses_websocket_planning_lease(state, pool_key_lease.as_ref()).await; + planning_lease.release().await; continue; } match codex_quota_breaker_blocks_candidate( @@ -405,7 +501,7 @@ pub(crate) async fn maybe_build_responses_websocket_decision( .await { Ok(true) => { - release_responses_websocket_planning_lease(state, pool_key_lease.as_ref()).await; + planning_lease.release().await; continue; } Ok(false) => {} @@ -454,13 +550,17 @@ pub(crate) async fn maybe_build_responses_websocket_decision( .flatten(), mapped_model, }; - return Ok(Some(ResponsesWebSocketDecision { + let decision = ResponsesWebSocketDecision { execution: payload, adapter, normalization, - })); + }; + // The decision report context now carries the lease identity. The + // WebSocket ownership layer takes over before any further await. + planning_lease.disarm(); + return Ok(Some(decision)); } - release_responses_websocket_planning_lease(state, pool_key_lease.as_ref()).await; + planning_lease.release().await; } Ok(None) @@ -469,20 +569,23 @@ pub(crate) async fn maybe_build_responses_websocket_decision( async fn release_responses_websocket_planning_lease( state: &AppState, lease: Option<&RuntimeLockLease>, -) { +) -> bool { let Some(lease) = lease else { - return; + return true; }; - if let Err(error) = - crate::handlers::shared::provider_pool::release_admin_provider_pool_key_lease( - state.runtime_state.as_ref(), - lease, - ) - .await + match 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" - ); + Ok(_) => true, + Err(error) => { + tracing::warn!( + error = ?error, + "gateway Responses WebSocket planner failed to release an unused pool key lease" + ); + false + } } } diff --git a/apps/aether-gateway/src/control/auth/mod.rs b/apps/aether-gateway/src/control/auth/mod.rs index 1aa699566..591e96ce0 100644 --- a/apps/aether-gateway/src/control/auth/mod.rs +++ b/apps/aether-gateway/src/control/auth/mod.rs @@ -11,8 +11,9 @@ pub(crate) use gate::{ should_buffer_request_for_local_auth, trusted_auth_local_rejection, GatewayLocalAuthRejection, }; pub(crate) use resolution::{ - refresh_execution_runtime_auth_context, resolve_execution_runtime_auth_context, - GatewayAdminPrincipalContext, GatewayControlAuthContext, + refresh_execution_runtime_auth_context, refresh_execution_runtime_auth_context_with_snapshot, + resolve_execution_runtime_auth_context, GatewayAdminPrincipalContext, + GatewayControlAuthContext, }; pub(super) use resolution::{resolve_control_decision_auth, ControlDecisionAuthResolution}; pub(crate) use types::GatewayCredentialCarrier; diff --git a/apps/aether-gateway/src/control/auth/resolution.rs b/apps/aether-gateway/src/control/auth/resolution.rs index 6ebe21d03..bd572db6e 100644 --- a/apps/aether-gateway/src/control/auth/resolution.rs +++ b/apps/aether-gateway/src/control/auth/resolution.rs @@ -725,20 +725,47 @@ pub(crate) async fn refresh_execution_runtime_auth_context( auth_context: GatewayControlAuthContext, auth_endpoint_signature: Option<&str>, ) -> Result { + refresh_execution_runtime_auth_context_with_snapshot( + state, + auth_context, + auth_endpoint_signature, + ) + .await + .map(|(auth_context, _)| auth_context) +} + +/// Strongly refreshes the long-lived execution authorization context and +/// returns the exact API-key snapshot that produced it. +/// +/// WebSocket turns need both values: using the refreshed context for RPM and +/// balance checks while letting the planner independently read its normal +/// cache can authorize a different provider/model snapshot for up to the cache +/// TTL. Ordinary HTTP callers keep using [`refresh_execution_runtime_auth_context`]. +pub(crate) async fn refresh_execution_runtime_auth_context_with_snapshot( + state: &AppState, + auth_context: GatewayControlAuthContext, + auth_endpoint_signature: Option<&str>, +) -> Result< + ( + GatewayControlAuthContext, + Option, + ), + GatewayError, +> { if auth_context.local_rejection.is_some() || !auth_context.access_allowed { - return Ok(auth_context); + return Ok((auth_context, None)); } let Some(auth_endpoint_signature) = auth_endpoint_signature .map(str::trim) .filter(|value| !value.is_empty()) else { - return Ok(auth_context); + return Ok((auth_context, None)); }; if !state.has_auth_api_key_reader() || auth_context.user_id.trim().is_empty() || auth_context.api_key_id.trim().is_empty() { - return Ok(auth_context); + return Ok((auth_context, None)); } let snapshot = { @@ -758,19 +785,20 @@ pub(crate) async fn refresh_execution_runtime_auth_context( denied.access_allowed = false; denied.local_rejection = Some(GatewayLocalAuthRejection::InvalidApiKey); denied.balance_remaining = None; - return Ok(denied); + return Ok((denied, None)); }; let wallet_access = resolve_wallet_auth_gate_uncached(state, &snapshot).await?; - Ok(build_data_backed_auth_context( + let refreshed = build_data_backed_auth_context( state, - snapshot, + snapshot.clone(), auth_endpoint_signature, Some(true), auth_context.balance_remaining, wallet_access, ) - .await) + .await; + Ok((refreshed, Some(snapshot))) } fn put_cached_auth_context( diff --git a/apps/aether-gateway/src/control/mod.rs b/apps/aether-gateway/src/control/mod.rs index f62651d4a..ca9531692 100644 --- a/apps/aether-gateway/src/control/mod.rs +++ b/apps/aether-gateway/src/control/mod.rs @@ -9,10 +9,11 @@ mod route; pub(crate) use auth::{ execution_plan_balance_capacity_rejection, extract_requested_model, - refresh_execution_runtime_auth_context, request_model_local_rejection, - resolve_execution_runtime_auth_context, should_buffer_request_for_local_auth, - trusted_auth_local_rejection, GatewayAdminPrincipalContext, GatewayControlAuthContext, - GatewayCredentialCarrier, GatewayLocalAuthRejection, + refresh_execution_runtime_auth_context, refresh_execution_runtime_auth_context_with_snapshot, + request_model_local_rejection, resolve_execution_runtime_auth_context, + should_buffer_request_for_local_auth, trusted_auth_local_rejection, + GatewayAdminPrincipalContext, GatewayControlAuthContext, GatewayCredentialCarrier, + GatewayLocalAuthRejection, }; pub(crate) use execute::{allows_control_execute_emergency, maybe_execute_via_control}; pub(crate) use management_token_permissions::{ diff --git a/apps/aether-gateway/src/execution_runtime/attempt_lifecycle.rs b/apps/aether-gateway/src/execution_runtime/attempt_lifecycle.rs index d85391863..12a7485b0 100644 --- a/apps/aether-gateway/src/execution_runtime/attempt_lifecycle.rs +++ b/apps/aether-gateway/src/execution_runtime/attempt_lifecycle.rs @@ -684,33 +684,35 @@ impl ExecutionAttemptLifecycle { terminal_summary.parser_error.as_deref(), reason, ); - let _ = self - .stage_guard - .await_stage( - self.trace_id.as_str(), - "candidate_terminal", + let candidate_state = state.clone(); + let candidate_plan = self.plan.clone(); + let candidate_report_context = payload.report_context.clone(); + let candidate_update = 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()), + }; + self.stage_guard + .await_detachable_stage(self.trace_id.as_str(), "candidate_terminal", async move { 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()), - }, - ), - ) + &candidate_state, + &candidate_plan, + candidate_report_context.as_ref(), + candidate_update, + ) + .await; + }) .await; // 3. provider 效果 @@ -747,12 +749,28 @@ impl ExecutionAttemptLifecycle { .await .is_some(); if !effects_completed { - let _ = self - .stage_guard - .await_stage( + // The provider-effects future was dropped at the caller's stage bound. Lease + // cleanup must not be dropped by that same bound as well: retain an owned copy of + // the exact report context (including its lease token/fencing token) and let the + // cleanup task finish after the caller stops waiting. The underlying conditional + // release is idempotent, and a context without a lease is a no-op. + let release_state = state.clone(); + let release_plan = self.plan.clone(); + let release_report_context = payload.report_context.clone(); + self.stage_guard + .await_detachable_stage( self.trace_id.as_str(), "pool_lease_release_after_effect_timeout", - release_local_pool_key_lease(state, effect_context), + async move { + release_local_pool_key_lease( + &release_state, + LocalExecutionEffectContext { + plan: &release_plan, + report_context: release_report_context.as_ref(), + }, + ) + .await; + }, ) .await; } @@ -1273,17 +1291,21 @@ mod stage_tests { use std::time::Duration; use base64::Engine as _; + use tokio::sync::Notify; use super::{ candidate_error_fields, AttemptBodyCapture, AttemptCandidateError, AttemptStageGuard, }; /// 效果段超时后仍然必须释放 pool key lease,否则那把 key 要等 lease TTL - /// 过期才放出来。这里用 `await_stage` 返回 `None` 驱动兜底分支。 + /// 过期才放出来。调用方对兜底清理的等待也必须有界,但不能把 owned cleanup + /// 一并取消。 #[tokio::test] - async fn a_timed_out_effect_stage_still_reaches_the_lease_release_fallback() { + async fn a_timed_out_effect_stage_detaches_lease_cleanup_until_it_completes() { let guard = AttemptStageGuard::Bounded(Duration::from_millis(20)); let lease_released = Arc::new(AtomicBool::new(false)); + let allow_release = Arc::new(Notify::new()); + let release_completed = Arc::new(Notify::new()); // 第一段:永不完成的效果投射。 let effects_completed = guard @@ -1295,23 +1317,34 @@ mod stage_tests { "a stage that never completes must not report success" ); - // 生产代码据此走兜底释放。 + // 生产代码据此走 owned/detached 兜底释放。让清理刻意慢于 caller bound, + // 证明调用方先返回之后,清理任务仍然存活并最终完成。 if !effects_completed { let released = Arc::clone(&lease_released); - let _ = guard - .await_stage( + let allow_release = Arc::clone(&allow_release); + let release_completed_task = Arc::clone(&release_completed); + guard + .await_detachable_stage( "trace", "pool_lease_release_after_effect_timeout", async move { + allow_release.notified().await; released.store(true, Ordering::SeqCst); + release_completed_task.notify_one(); }, ) .await; } assert!( - lease_released.load(Ordering::SeqCst), - "the lease must still be released after an effect-stage timeout" + !lease_released.load(Ordering::SeqCst), + "the caller must stop waiting at its bound even while cleanup is pending" ); + + allow_release.notify_one(); + tokio::time::timeout(Duration::from_secs(1), release_completed.notified()) + .await + .expect("detached lease cleanup must eventually complete"); + assert!(lease_released.load(Ordering::SeqCst)); } #[tokio::test] @@ -1327,16 +1360,16 @@ mod stage_tests { assert_eq!(value, Some(7)); } - /// 不能丢的写入即使调用方停止等待也要跑完:`await_detachable_stage` 先 spawn - /// 再等,上界只约束「等多久」。 + /// candidate 终态不能因为调用方的等待上界而丢失:先 spawn 再等,超时只 + /// 停止 relay 对它的等待,后台写入仍然必须完成。 #[tokio::test] - async fn a_detachable_stage_completes_even_after_the_caller_stops_waiting() { + async fn a_detachable_candidate_terminal_completes_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 { + .await_detachable_stage("trace", "candidate_terminal", async move { tokio::time::sleep(Duration::from_millis(120)).await; flag.store(true, Ordering::SeqCst); }) @@ -1349,7 +1382,7 @@ mod stage_tests { tokio::time::sleep(Duration::from_millis(300)).await; assert!( written.load(Ordering::SeqCst), - "a detached write must still run to completion" + "a detached candidate terminal write must still run to completion" ); } diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/ingress.rs b/apps/aether-gateway/src/handlers/proxy/websocket/ingress.rs index 20465da35..bd3626fee 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/ingress.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/ingress.rs @@ -1,11 +1,16 @@ //! Authenticated public WebSocket upgrade admission shared by AI adapters. use std::future::Future; -use std::net::SocketAddr; +use std::net::{IpAddr, SocketAddr}; use axum::body::Body; use axum::extract::ws::{WebSocket, WebSocketUpgrade}; -use axum::http::{HeaderMap, Method, Response, StatusCode, Uri}; +use axum::http::header::{ + AUTHORIZATION, CONNECTION, COOKIE, HOST, PROXY_AUTHORIZATION, TE, TRAILER, TRANSFER_ENCODING, + UPGRADE, +}; +use axum::http::uri::PathAndQuery; +use axum::http::{HeaderMap, HeaderName, Method, Response, StatusCode, Uri}; use tracing::{info, warn}; use crate::api::response::{ @@ -13,7 +18,8 @@ use crate::api::response::{ build_local_overloaded_response, }; use crate::control::{ - trusted_auth_local_rejection, GatewayControlDecision, GatewayLocalAuthRejection, + trusted_auth_local_rejection, GatewayControlDecision, GatewayCredentialCarrier, + GatewayLocalAuthRejection, }; use crate::handlers::proxy::websocket::session::{WebSocketSessionLimits, WEBSOCKET_LOG_TRANSPORT}; use crate::handlers::shared::ip_rules_allow; @@ -28,8 +34,10 @@ pub(crate) struct WebSocketRequestContext { pub(crate) headers: HeaderMap, pub(crate) uri: Uri, pub(crate) remote_addr: SocketAddr, + /// Effective client IP resolved once from the authenticated Upgrade. Every + /// turn re-checks live API-key/admin IP policy against this immutable fact. + pub(crate) client_ip: IpAddr, pub(crate) decision: GatewayControlDecision, - pub(crate) rpm_bypassed: bool, /// Held for the lifetime of the upgraded socket. The Responses session /// polls its health and closes the client when a distributed lease is /// revoked or expires. @@ -40,7 +48,6 @@ pub(crate) struct WebSocketRequestContext { #[derive(Clone, Copy)] pub(crate) struct WebSocketIngressSpec { pub(crate) route_unavailable_message: &'static str, - pub(crate) ip_whitelist_failure_event_name: &'static str, } /// Performs the HTTP-only part of an AI WebSocket request. @@ -81,7 +88,7 @@ where &trace_id, ) .await?; - let Some(decision) = request_context.control_decision else { + let Some(mut decision) = request_context.control_decision else { return build_local_http_error_response( &trace_id, None, @@ -92,6 +99,28 @@ where if let Some(rejection) = trusted_auth_local_rejection(Some(&decision), &headers) { return build_local_auth_rejection_response(&trace_id, Some(&decision), &rejection); } + // Browsers attach cookies to WebSocket handshakes automatically and the + // WebSocket API does not let callers add an Authorization header. A + // cookie-only public upgrade would therefore be vulnerable to cross-site + // WebSocket hijacking unless every deployment maintained an Origin + // allowlist. Explicit API-key/bearer credentials (or trusted internal + // auth resolved by the control plane) remain supported. + if !websocket_credential_carrier_is_allowed(decision.gateway_credential_carrier) { + warn!( + event_name = "ai_websocket_cookie_only_auth_rejected", + log_type = "security", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %trace_id, + client_ip = %client_ip, + "gateway rejected cookie-only public WebSocket authentication" + ); + return build_local_auth_rejection_response( + &trace_id, + Some(&decision), + &GatewayLocalAuthRejection::InvalidApiKey, + ); + } let Some(auth_context) = decision.auth_context.as_ref() else { return build_local_auth_rejection_response( &trace_id, @@ -99,7 +128,10 @@ where &GatewayLocalAuthRejection::InvalidApiKey, ); }; - if !auth_context.access_allowed { + if !auth_context.access_allowed + || auth_context.user_id.trim().is_empty() + || auth_context.api_key_id.trim().is_empty() + { return build_local_auth_rejection_response( &trace_id, Some(&decision), @@ -116,22 +148,6 @@ where ); } - let ip_whitelisted = match state.admin_security_ip_whitelisted(client_ip).await { - Ok(value) => value, - Err(error) => { - warn!( - event_name = spec.ip_whitelist_failure_event_name, - log_type = "ops", - transport = WEBSOCKET_LOG_TRANSPORT, - websocket = true, - trace_id = %trace_id, - client_ip = %client_ip, - error = ?error, - "gateway continued with WebSocket rate limiting after IP whitelist check error" - ); - false - } - }; let request_permit = match state.try_acquire_request_permit().await { Ok(permit) => permit, Err(error) => { @@ -155,13 +171,21 @@ where } }; + // Authentication has consumed the downstream credentials. From this + // point on the URI and headers become planner input, so retain neither an + // API key from the query string nor client authentication/handshake + // headers. Provider authentication is added independently by the + // planner and is therefore unaffected by this boundary. + let uri = websocket_planning_uri(&uri); + decision.public_query_string = uri.query().map(ToOwned::to_owned); + let headers = websocket_planning_headers(headers); let context = WebSocketRequestContext { trace_id, headers, uri, remote_addr, + client_ip, decision, - rpm_bypassed: ip_whitelisted, websocket_connection_permit, }; Ok(ws @@ -173,6 +197,106 @@ where })) } +fn websocket_credential_carrier_is_allowed(carrier: Option) -> bool { + carrier != Some(GatewayCredentialCarrier::CookieHeader) +} + +fn websocket_planning_uri(uri: &Uri) -> Uri { + let Some(query) = uri.query() else { + return uri.clone(); + }; + let mut retained = Vec::new(); + let mut removed_sensitive_value = false; + for (name, value) in url::form_urlencoded::parse(query.as_bytes()) { + if websocket_query_parameter_is_sensitive(name.as_ref()) { + removed_sensitive_value = true; + } else { + retained.push((name.into_owned(), value.into_owned())); + } + } + if !removed_sensitive_value { + return uri.clone(); + } + + let mut serializer = url::form_urlencoded::Serializer::new(String::new()); + serializer.extend_pairs(retained.iter().map(|(name, value)| (name, value))); + let retained_query = serializer.finish(); + let path_and_query = if retained_query.is_empty() { + uri.path().to_string() + } else { + format!("{}?{retained_query}", uri.path()) + }; + let path_and_query = path_and_query + .parse::() + .expect("a valid URI path plus form-encoded query must remain valid"); + let mut parts = uri.clone().into_parts(); + parts.path_and_query = Some(path_and_query); + Uri::from_parts(parts).expect("replacing only path-and-query must preserve a valid URI") +} + +fn websocket_query_parameter_is_sensitive(name: &str) -> bool { + matches!( + name.to_ascii_lowercase().as_str(), + "key" | "api_key" | "api-key" | "access_token" | "authorization" | "token" + ) +} + +fn websocket_planning_headers(mut headers: HeaderMap) -> HeaderMap { + // RFC 9110 permits Connection to name additional hop-by-hop fields. Read + // those names before removing Connection itself. + let connection_scoped_names = headers + .get_all(CONNECTION) + .iter() + .filter_map(|value| value.to_str().ok()) + .flat_map(|value| value.split(',')) + .filter_map(|name| HeaderName::from_bytes(name.trim().as_bytes()).ok()) + .collect::>(); + for name in connection_scoped_names { + headers.remove(name); + } + + for name in [ + AUTHORIZATION, + CONNECTION, + COOKIE, + HOST, + PROXY_AUTHORIZATION, + TE, + TRAILER, + TRANSFER_ENCODING, + UPGRADE, + ] { + headers.remove(name); + } + for name in [ + "api-key", + "keep-alive", + "proxy-connection", + "x-api-key", + "x-goog-api-key", + crate::constants::GATEWAY_HEADER, + crate::constants::TRUSTED_AUTH_USER_ID_HEADER, + crate::constants::TRUSTED_AUTH_API_KEY_ID_HEADER, + crate::constants::TRUSTED_AUTH_BALANCE_HEADER, + crate::constants::TRUSTED_AUTH_ACCESS_ALLOWED_HEADER, + crate::constants::TRUSTED_ADMIN_USER_ID_HEADER, + crate::constants::TRUSTED_ADMIN_USER_ROLE_HEADER, + crate::constants::TRUSTED_ADMIN_SESSION_ID_HEADER, + crate::constants::TRUSTED_ADMIN_MANAGEMENT_TOKEN_ID_HEADER, + ] { + headers.remove(name); + } + let websocket_managed_names = headers + .keys() + .filter(|name| name.as_str().starts_with("sec-websocket-")) + .cloned() + .collect::>(); + for name in websocket_managed_names { + headers.remove(name); + } + headers +} + fn websocket_admission_error_response( trace_id: &str, decision: &GatewayControlDecision, @@ -291,3 +415,113 @@ impl Drop for WebSocketConnectionLog { ); } } + +#[cfg(test)] +mod tests { + use axum::http::header::{ + AUTHORIZATION, CONNECTION, COOKIE, HOST, ORIGIN, SEC_WEBSOCKET_KEY, UPGRADE, USER_AGENT, + }; + use axum::http::{HeaderMap, HeaderValue, Uri}; + + use super::{ + websocket_credential_carrier_is_allowed, websocket_planning_headers, websocket_planning_uri, + }; + use crate::control::GatewayCredentialCarrier; + + #[test] + fn planning_uri_removes_query_credentials_without_losing_safe_parameters() { + let uri: Uri = "/v1/responses?key=downstream-secret&client_hint=a%20b&token=also-secret" + .parse() + .expect("request URI should parse"); + + let sanitized = websocket_planning_uri(&uri); + + assert_eq!(sanitized.path(), "/v1/responses"); + assert_eq!(sanitized.query(), Some("client_hint=a+b")); + assert!(!sanitized.to_string().contains("downstream-secret")); + assert!(!sanitized.to_string().contains("also-secret")); + } + + #[test] + fn planning_uri_leaves_an_uncredentialed_query_byte_for_byte_unchanged() { + let uri: Uri = "/v1/responses?client_hint=a%20b&empty=" + .parse() + .expect("request URI should parse"); + + assert_eq!(websocket_planning_uri(&uri), uri); + } + + #[test] + fn planning_headers_drop_client_auth_cookie_and_websocket_transport_state() { + let mut headers = HeaderMap::new(); + headers.insert( + AUTHORIZATION, + HeaderValue::from_static("Bearer client-secret"), + ); + headers.insert(COOKIE, HeaderValue::from_static("session=client-secret")); + headers.insert("x-api-key", HeaderValue::from_static("client-secret")); + headers.insert(HOST, HeaderValue::from_static("gateway.example")); + headers.insert( + CONNECTION, + HeaderValue::from_static("keep-alive, Upgrade, x-connection-secret"), + ); + headers.insert(UPGRADE, HeaderValue::from_static("websocket")); + headers.insert(SEC_WEBSOCKET_KEY, HeaderValue::from_static("handshake-key")); + headers.insert( + "sec-websocket-future-field", + HeaderValue::from_static("future-handshake-value"), + ); + headers.insert( + "x-connection-secret", + HeaderValue::from_static("connection-secret"), + ); + headers.insert(ORIGIN, HeaderValue::from_static("https://client.example")); + headers.insert(USER_AGENT, HeaderValue::from_static("codex-cli/test")); + headers.insert("x-client-hint", HeaderValue::from_static("safe")); + + let sanitized = websocket_planning_headers(headers); + + for name in [ + AUTHORIZATION.as_str(), + COOKIE.as_str(), + "x-api-key", + HOST.as_str(), + CONNECTION.as_str(), + UPGRADE.as_str(), + SEC_WEBSOCKET_KEY.as_str(), + "sec-websocket-future-field", + "x-connection-secret", + ] { + assert!(sanitized.get(name).is_none(), "{name} must not survive"); + } + assert_eq!( + sanitized.get(ORIGIN), + Some(&HeaderValue::from_static("https://client.example")) + ); + assert_eq!( + sanitized.get(USER_AGENT), + Some(&HeaderValue::from_static("codex-cli/test")) + ); + assert_eq!( + sanitized.get("x-client-hint"), + Some(&HeaderValue::from_static("safe")) + ); + } + + #[test] + fn websocket_auth_requires_an_explicit_credential_instead_of_cookie_only() { + assert!(!websocket_credential_carrier_is_allowed(Some( + GatewayCredentialCarrier::CookieHeader + ))); + for carrier in [ + None, + Some(GatewayCredentialCarrier::AuthorizationBearer), + Some(GatewayCredentialCarrier::XApiKey), + Some(GatewayCredentialCarrier::ApiKey), + Some(GatewayCredentialCarrier::XGoogApiKey), + Some(GatewayCredentialCarrier::QueryKey), + ] { + assert!(websocket_credential_carrier_is_allowed(carrier)); + } + } +} diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/adapter.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/adapter.rs index 7b3f99bf7..3ab4c74b9 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/adapter.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/adapter.rs @@ -46,6 +46,25 @@ pub(super) enum ResponsesWebSocketRebindSafety { Unsafe { reason: &'static str }, } +/// How an upstream text frame crosses the public Responses WebSocket boundary. +/// +/// The normal path is deliberately byte-opaque: callers forward the parsed +/// frame's original text without rebuilding it from a gateway-owned schema. +/// Codex is the only adapter that may peel its documented private batch +/// envelope. Even then, the retained events are borrowed whole so unknown +/// `response.*` event types and unknown fields survive unchanged. +#[derive(Debug, Clone, PartialEq)] +pub(super) enum ResponsesWebSocketRelayDirective<'a> { + /// Forward the provider frame's original text exactly as received. + ForwardOriginal, + /// The provider frame was a private batch envelope. Forward each retained + /// event in document order by serializing the complete borrowed value. + ForwardEvents(Vec<&'a Value>), + /// The entire frame was an explicitly recognized provider-private + /// envelope and therefore has no public event to relay. + SuppressProviderPrivate, +} + /// Boundary between the standard Responses protocol engine and provider /// behavior. Adapters receive already-planned provider requests; they never /// own public WebSocket parsing, turn accounting, or model scheduling. @@ -69,6 +88,15 @@ pub(super) trait ResponsesWebSocketProtocolAdapter: Send + Sync { /// observably ambiguous to the client. fn rebind_safety_for_upstream_event(&self, event: &Value) -> ResponsesWebSocketRebindSafety; + /// Selects the public relay shape without projecting a provider event + /// through an Aether-owned field or event-type allowlist. + fn relay_directive_for_upstream_event<'a>( + &self, + _event: &'a Value, + ) -> ResponsesWebSocketRelayDirective<'a> { + ResponsesWebSocketRelayDirective::ForwardOriginal + } + /// Lets an adapter classify provider-only events. Returning a directive /// asks the shared session to drain after the active standard response. fn observe_upstream_event(&self, event: &Value) @@ -169,7 +197,12 @@ pub(super) fn is_standard_responses_event(event: &Value) -> bool { #[cfg(test)] mod tests { - use super::{resolve_responses_websocket_adapter, ResponsesWebSocketProtocolAdapter}; + use serde_json::json; + + use super::{ + resolve_responses_websocket_adapter, ResponsesWebSocketProtocolAdapter, + ResponsesWebSocketRelayDirective, + }; use crate::orchestration::ResponsesWebSocketAdapter; #[test] @@ -183,4 +216,19 @@ mod tests { "responses_websocket_handshake_failed" ); } + + #[test] + fn standard_adapter_always_forwards_future_events_opaquely() { + let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard); + let event = json!({ + "type": "response.future_capability.delta", + "delta": {"future_shape": [1, {"nested": true}]}, + "unknown_top_level": {"must": "survive"}, + }); + + assert_eq!( + adapter.relay_directive_for_upstream_event(&event), + ResponsesWebSocketRelayDirective::ForwardOriginal + ); + } } 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 index 1db6c4353..059f00559 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/adapters/codex.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/adapters/codex.rs @@ -7,6 +7,7 @@ use super::super::adapter::{ is_standard_responses_event, ResponsesWebSocketAdapterObservation, ResponsesWebSocketDrainDirective, ResponsesWebSocketExclusionIdentity, ResponsesWebSocketProtocolAdapter, ResponsesWebSocketRebindSafety, + ResponsesWebSocketRelayDirective, }; use crate::ai_serving::AiExecutionDecision; use crate::clock::current_unix_secs; @@ -66,19 +67,45 @@ impl ResponsesWebSocketProtocolAdapter for CodexResponsesWebSocketAdapter { } fn rebind_safety_for_upstream_event(&self, event: &Value) -> ResponsesWebSocketRebindSafety { - if let Some(chunks) = event.get("chunks").and_then(Value::as_array) { - if chunks.is_empty() { + let mut saw_event = false; + if event.get("type").and_then(Value::as_str).is_some() { + saw_event = true; + let safety = codex_direct_rebind_safety(event); + if matches!(safety, ResponsesWebSocketRebindSafety::Unsafe { .. }) { + return safety; + } + } + match event.get("chunks") { + Some(Value::Array(chunks)) => { + for chunk in chunks { + saw_event = true; + let safety = codex_direct_rebind_safety(chunk); + if matches!(safety, ResponsesWebSocketRebindSafety::Unsafe { .. }) { + return safety; + } + } + } + Some(_) => { return ResponsesWebSocketRebindSafety::Unsafe { reason: "unrecognized_upstream_event", }; } - return chunks - .iter() - .map(codex_direct_rebind_safety) - .find(|safety| matches!(safety, ResponsesWebSocketRebindSafety::Unsafe { .. })) - .unwrap_or(ResponsesWebSocketRebindSafety::Safe); + None => {} } - codex_direct_rebind_safety(event) + if saw_event { + ResponsesWebSocketRebindSafety::Safe + } else { + ResponsesWebSocketRebindSafety::Unsafe { + reason: "unrecognized_upstream_event", + } + } + } + + fn relay_directive_for_upstream_event<'a>( + &self, + event: &'a Value, + ) -> ResponsesWebSocketRelayDirective<'a> { + codex_relay_directive(event) } fn observe_upstream_event( @@ -148,9 +175,15 @@ fn codex_direct_rebind_safety(event: &Value) -> ResponsesWebSocketRebindSafety { // upstream can safely emit its own current snapshot. return ResponsesWebSocketRebindSafety::Safe; } - if event_type == "error" && parse_codex_rate_limits(event).is_some() { - // The quota error is withheld from the client when the shared - // session successfully rebinds, therefore it remains replay-safe. + if event_type == "error" + && event.pointer("/error/type").and_then(Value::as_str) == Some("usage_limit_reached") + && parse_codex_rate_limits(event).is_some() + { + // This terminal quota event has not been relayed yet. It can trigger + // one transparent attempt on another key as long as no earlier public + // response event made the logical turn unsafe. If replanning fails, + // the connection layer forwards this exact upstream error instead of + // manufacturing a gateway continuation error. return ResponsesWebSocketRebindSafety::Safe; } let reason = if is_standard_responses_event(event) { @@ -161,6 +194,55 @@ fn codex_direct_rebind_safety(event: &Value) -> ResponsesWebSocketRebindSafety { ResponsesWebSocketRebindSafety::Unsafe { reason } } +fn codex_relay_directive(event: &Value) -> ResponsesWebSocketRelayDirective<'_> { + match event.get("chunks") { + Some(Value::Array(chunks)) if is_explicit_codex_batch_envelope(event) => { + let public_events = chunks + .iter() + .filter(|chunk| !is_codex_private_leaf_event(chunk)) + .collect::>(); + if public_events.is_empty() { + ResponsesWebSocketRelayDirective::SuppressProviderPrivate + } else { + ResponsesWebSocketRelayDirective::ForwardEvents(public_events) + } + } + // A malformed or future shape is not proven private. Preserve it + // opaquely rather than guessing at a provider schema. + Some(_) => ResponsesWebSocketRelayDirective::ForwardOriginal, + None if is_codex_private_leaf_event(event) => { + ResponsesWebSocketRelayDirective::SuppressProviderPrivate + } + None => ResponsesWebSocketRelayDirective::ForwardOriginal, + } +} + +/// Recognizes only Codex's private batch container. A type-less object must +/// contain exactly `chunks`; unknown siblings could be future public protocol +/// data and therefore force opaque forwarding. A named Codex private root may +/// carry provider metadata alongside its chunks and is safe to peel. +fn is_explicit_codex_batch_envelope(event: &Value) -> bool { + if is_codex_private_event_type(event) { + return true; + } + event.as_object().is_some_and(|object| { + object.len() == 1 + && object.contains_key("chunks") + && event.get("type").and_then(Value::as_str).is_none() + }) +} + +fn is_codex_private_leaf_event(event: &Value) -> bool { + is_codex_private_event_type(event) && event.get("chunks").is_none() +} + +fn is_codex_private_event_type(event: &Value) -> bool { + matches!( + event.get("type").and_then(Value::as_str), + Some("codex.rate_limits" | "codex.response.metadata") + ) +} + fn parse_codex_rate_limits(event: &Value) -> Option { aether_admin::provider::quota::parse_codex_websocket_rate_limits_response( event, @@ -174,7 +256,7 @@ mod tests { use super::{ CodexResponsesWebSocketAdapter, ResponsesWebSocketProtocolAdapter, - ResponsesWebSocketRebindSafety, + ResponsesWebSocketRebindSafety, ResponsesWebSocketRelayDirective, }; #[test] @@ -245,7 +327,7 @@ mod tests { } #[test] - fn only_known_codex_pre_response_metadata_is_safe_to_rebind() { + fn only_known_codex_pre_response_signals_are_safe_to_rebind() { let adapter = CodexResponsesWebSocketAdapter; assert_eq!( @@ -286,5 +368,105 @@ mod tests { reason: "unrecognized_upstream_event" } ); + assert_eq!( + adapter.rebind_safety_for_upstream_event(&json!({ + "type": "error", + "error": { + "type": "usage_limit_reached", + "plan_type": "plus", + "resets_in_seconds": 3_600 + }, + "status_code": 429 + })), + ResponsesWebSocketRebindSafety::Safe + ); + assert_eq!( + adapter.rebind_safety_for_upstream_event(&json!({ + "type": "error", + "error": {"type": "usage_limit_reached"} + })), + ResponsesWebSocketRebindSafety::Unsafe { + reason: "unrecognized_upstream_event" + } + ); + assert_eq!( + adapter.rebind_safety_for_upstream_event(&json!({ + "type": "response.future_capability.delta", + "chunks": [{"type": "codex.rate_limits"}] + })), + ResponsesWebSocketRebindSafety::Unsafe { + reason: "standard_response_event" + } + ); + } + + #[test] + fn codex_suppresses_only_explicit_private_events_and_envelopes() { + let adapter = CodexResponsesWebSocketAdapter; + + for event in [ + json!({"type": "codex.rate_limits", "rate_limits": {"allowed": true}}), + json!({"type": "codex.response.metadata", "account_hint": "private"}), + json!({"chunks": [ + {"type": "codex.rate_limits"}, + {"type": "codex.response.metadata"} + ]}), + ] { + assert_eq!( + adapter.relay_directive_for_upstream_event(&event), + ResponsesWebSocketRelayDirective::SuppressProviderPrivate + ); + } + + for event in [ + json!({"type": "error", "error": {"type": "usage_limit_reached"}}), + json!({"type": "codex.future_private_maybe", "future": true}), + json!({"chunks": [], "future_envelope_field": {"must": "survive"}}), + json!({"type": "response.future.done", "future_capability": true}), + ] { + assert_eq!( + adapter.relay_directive_for_upstream_event(&event), + ResponsesWebSocketRelayDirective::ForwardOriginal + ); + } + } + + #[test] + fn mixed_codex_batch_forwards_whole_non_private_events_in_order() { + let adapter = CodexResponsesWebSocketAdapter; + let event = json!({ + "chunks": [ + { + "type": "response.created", + "response": {"id": "resp_future"}, + "future_created_field": {"opaque": true} + }, + {"type": "codex.rate_limits", "account_hint": "private"}, + { + "type": "response.future_capability.delta", + "future_capability": {"nested": [1, 2, 3]}, + "sequence_number": 2 + }, + {"provider_future_event": {"unknown": "must be forwarded"}}, + { + "type": "error", + "error": {"type": "future_error", "future_detail": 7} + } + ] + }); + + let ResponsesWebSocketRelayDirective::ForwardEvents(events) = + adapter.relay_directive_for_upstream_event(&event) + else { + panic!("a mixed private envelope must retain all non-private events"); + }; + assert_eq!(events.len(), 4); + assert_eq!(events[0]["future_created_field"], json!({"opaque": true})); + assert_eq!(events[1]["future_capability"], json!({"nested": [1, 2, 3]})); + assert_eq!( + events[2]["provider_future_event"], + json!({"unknown": "must be forwarded"}) + ); + assert_eq!(events[3]["error"]["future_detail"], json!(7)); } } diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/binding.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/binding.rs index 7160eddf6..eed1116ff 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/binding.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/binding.rs @@ -2,10 +2,10 @@ //! //! A Responses continuation carries state that lives on one provider socket. //! Comparing only the selected key is therefore not sufficient: transport -//! settings, stable account headers, and the protocol adapter can all change -//! the connection that would receive the next event. Rotating bearer values -//! are intentionally excluded because they do not change an already-upgraded -//! socket's physical binding. +//! settings, stable account headers, credentials, and the protocol adapter can +//! all change the connection that would receive the next event. Ordinary Codex +//! OAuth access-token refreshes retain the credential generation and therefore +//! do not unnecessarily replace an already-upgraded socket. use std::collections::{BTreeMap, BTreeSet}; use std::fmt; @@ -34,11 +34,15 @@ pub(super) struct UpstreamBindingIdentity { key_id: Option, 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]>, + /// One-way identity for the credential generation used by this socket. + /// + /// A provider key id identifies a catalog row, not the secret currently + /// stored in that row. Codex decisions carry a server-owned credential + /// generation which is stable across access-token refreshes but rotates + /// when the account/static/refresh credential is replaced. Other + /// decisions conservatively fingerprint the effective authentication + /// handshake values. + credential_fingerprint: [u8; 32], proxy: Option, transport_profile: Option, } @@ -82,11 +86,8 @@ impl UpstreamBindingIdentity { handshake_headers.insert(name, value.to_string()); } } - let auth_fingerprint = decision - .key_id - .is_none() - .then(|| fingerprint_headers(&authentication_headers)) - .filter(|_| !authentication_headers.is_empty()); + let credential_fingerprint = + credential_binding_fingerprint(decision, &authentication_headers); Ok(Self { adapter_kind: adapter.kind(), @@ -95,7 +96,7 @@ impl UpstreamBindingIdentity { key_id: decision.key_id.clone(), upstream_url, handshake_headers, - auth_fingerprint, + credential_fingerprint, proxy: effective_proxy_snapshot(decision.proxy.as_ref()), transport_profile: decision.transport_profile.clone(), }) @@ -127,6 +128,7 @@ fn authentication_header_names(decision: &AiExecutionDecision) -> BTreeSet) -> [u8; 32] { let mut hasher = Sha256::new(); + hasher.update(b"aether-responses-websocket-auth-headers-v1"); for (name, value) in headers { hasher.update((name.len() as u64).to_be_bytes()); hasher.update(name.as_bytes()); @@ -136,6 +138,69 @@ fn fingerprint_headers(headers: &BTreeMap) -> [u8; 32] { hasher.finalize().into() } +/// Returns the non-secret credential identity represented by a planner +/// decision. The generation is emitted by Aether's trusted Codex planner from +/// provider-key metadata; it is not sourced from the downstream request. +fn credential_binding_fingerprint( + decision: &AiExecutionDecision, + authentication_headers: &BTreeMap, +) -> [u8; 32] { + if decision + .provider_type + .as_deref() + .is_some_and(|provider_type| provider_type.trim().eq_ignore_ascii_case("codex")) + { + if let Some(generation) = decision + .report_context + .as_ref() + .and_then(|context| context.get("codex_credential_generation")) + .and_then(serde_json::Value::as_str) + .map(str::trim) + .filter(|generation| !generation.is_empty()) + { + let mut hasher = Sha256::new(); + hasher.update(b"aether-responses-websocket-codex-credential-generation-v1"); + hasher.update((generation.len() as u64).to_be_bytes()); + hasher.update(generation.as_bytes()); + // Only a planner-owned Codex bearer access token is expected to + // rotate without changing credential generation. Compare the + // effective handshake value with the decision's original auth + // value: auth-config/routing/header overrides change only the + // former and therefore must force a rebind. + let stable_authentication_headers = authentication_headers + .iter() + .filter(|(name, value)| { + !is_planner_owned_codex_bearer(decision, name.as_str(), value.as_str()) + }) + .map(|(name, value)| (name.clone(), value.clone())) + .collect::>(); + hasher.update(fingerprint_headers(&stable_authentication_headers)); + return hasher.finalize().into(); + } + } + + // Fail closed when no trusted generation is available. Rebinding after an + // access-token change is preferable to sending a continuation over a + // socket authenticated with a credential that may have been replaced. + fingerprint_headers(authentication_headers) +} + +fn is_planner_owned_codex_bearer( + decision: &AiExecutionDecision, + name: &str, + effective_value: &str, +) -> bool { + name.eq_ignore_ascii_case("authorization") + && decision + .auth_header + .as_deref() + .is_some_and(|header| header.eq_ignore_ascii_case(name)) + && decision.auth_value.as_deref() == Some(effective_value) + && effective_value + .get(.."bearer ".len()) + .is_some_and(|scheme| scheme.eq_ignore_ascii_case("bearer ")) +} + /// Normalize only values that are provably direct transport. Keep node/tunnel /// fields even though the current WebSocket builder rejects those proxies: a /// re-plan must not accidentally reuse an already-bound direct socket for a @@ -313,18 +378,18 @@ mod tests { assert_ne!(identity, changed_identity); } - let mut rotated = base.clone(); - rotated + let mut static_secret_rotated = base.clone(); + static_secret_rotated .provider_request_headers .insert("Authorization".to_string(), "Bearer rotated".to_string()); - assert_eq!( + assert_ne!( identity, - UpstreamBindingIdentity::from_decision(adapter, &rotated).unwrap() + UpstreamBindingIdentity::from_decision(adapter, &static_secret_rotated).unwrap() ); } #[test] - fn stable_key_identity_ignores_custom_auth_value_rotation() { + fn stable_key_identity_rejects_custom_static_auth_value_rotation() { let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard); let mut base = decision(); base.auth_header = Some("X-Provider-Token".to_string()); @@ -335,43 +400,125 @@ mod tests { ); let identity = UpstreamBindingIdentity::from_decision(adapter, &base).unwrap(); assert!(!identity.handshake_headers.contains_key("x-provider-token")); - assert!(identity.auth_fingerprint.is_none()); let mut rotated = base; rotated.provider_request_headers.insert( "X-Provider-Token".to_string(), "provider-token-2".to_string(), ); - assert_eq!( + assert_ne!( identity, UpstreamBindingIdentity::from_decision(adapter, &rotated).unwrap() ); } #[test] - fn missing_key_identity_fingerprints_authentication_values() { - let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard); + fn codex_access_token_refresh_reuses_the_same_credential_generation() { + let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex); let mut first = decision(); - first.key_id = None; + first.provider_type = Some("codex".to_string()); + first.report_context = Some(json!({ + "codex_credential_generation": "credential-generation-1" + })); let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap(); - assert!(first_identity.auth_fingerprint.is_some()); - let mut same_account_rotation = first.clone(); - same_account_rotation.provider_request_headers.insert( + let mut access_token_refreshed = first; + access_token_refreshed.auth_value = Some("Bearer refreshed-access-token".to_string()); + access_token_refreshed.provider_request_headers.insert( "Authorization".to_string(), - "Bearer different-account-or-token".to_string(), + "Bearer refreshed-access-token".to_string(), ); - let changed_identity = - UpstreamBindingIdentity::from_decision(adapter, &same_account_rotation).unwrap(); - assert_ne!(first_identity, changed_identity); + assert_eq!( + first_identity, + UpstreamBindingIdentity::from_decision(adapter, &access_token_refreshed).unwrap() + ); + } - let mut non_auth_change = first; - non_auth_change - .provider_request_headers - .insert("X-Client".to_string(), "other-client".to_string()); + #[test] + fn codex_authorization_override_changes_binding_with_the_same_generation() { + let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex); + let mut first = decision(); + first.provider_type = Some("codex".to_string()); + first.report_context = Some(json!({ + "codex_credential_generation": "credential-generation-1" + })); + let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap(); + + // The planner-owned auth value remains unchanged while an effective + // auth-config/header override replaces the actual handshake value. + first.provider_request_headers.insert( + "Authorization".to_string(), + "Bearer endpoint-override".to_string(), + ); assert_ne!( first_identity, - UpstreamBindingIdentity::from_decision(adapter, &non_auth_change).unwrap() + UpstreamBindingIdentity::from_decision(adapter, &first).unwrap() + ); + } + + #[test] + fn codex_credential_replacement_changes_binding_for_the_same_key_id() { + let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex); + let mut first = decision(); + first.provider_type = Some("codex".to_string()); + first.report_context = Some(json!({ + "codex_credential_generation": "credential-generation-1" + })); + let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap(); + + let mut replaced = first; + replaced.provider_request_headers.insert( + "Authorization".to_string(), + "Bearer replacement-access-token".to_string(), + ); + replaced.report_context = Some(json!({ + "codex_credential_generation": "credential-generation-2" + })); + assert_ne!( + first_identity, + UpstreamBindingIdentity::from_decision(adapter, &replaced).unwrap() + ); + } + + #[test] + fn codex_custom_auth_rotation_changes_binding_with_the_same_generation() { + let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex); + let mut first = decision(); + first.provider_type = Some("codex".to_string()); + first.auth_header = Some("X-Provider-Token".to_string()); + first.provider_request_headers.insert( + "X-Provider-Token".to_string(), + "provider-token-1".to_string(), + ); + first.report_context = Some(json!({ + "codex_credential_generation": "credential-generation-1" + })); + let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap(); + + first.provider_request_headers.insert( + "X-Provider-Token".to_string(), + "provider-token-2".to_string(), + ); + assert_ne!( + first_identity, + UpstreamBindingIdentity::from_decision(adapter, &first).unwrap() + ); + } + + #[test] + fn missing_codex_credential_generation_fails_closed_on_auth_rotation() { + let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex); + let mut first = decision(); + first.provider_type = Some("codex".to_string()); + let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap(); + + first.provider_request_headers.insert( + "Authorization".to_string(), + "Bearer possibly-replaced-credential".to_string(), + ); + assert_ne!( + first_identity, + UpstreamBindingIdentity::from_decision(adapter, &first).unwrap() ); } diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/client.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/client.rs index 6abfb51d6..89012b322 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/client.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/client.rs @@ -1,6 +1,5 @@ //! 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; @@ -8,39 +7,43 @@ use uuid::Uuid; use wreq::ws::message::Message as WreqWsMessage; use super::adapter::{resolve_responses_websocket_adapter, ResponsesWebSocketDrainDirective}; +use super::control::{resolve_responses_websocket_turn_control, ResponsesWebSocketTurnControl}; use super::lifecycle::{ await_pending_turn_finalization, queue_turn_finalization, - send_responses_websocket_turn_start_error, ActiveProviderAttempt, + send_responses_websocket_turn_start_error, }; -use super::quota::{mark_active_response_retry_unsafe, send_previous_response_not_found}; +use super::ownership::{ + await_owned_responses_websocket_plan, begin_responses_websocket_turn_with_planned_lease, + spawn_owned_responses_websocket_plan, OwnedResponsesWebSocketDecision, +}; +use super::quota::mark_active_response_retry_unsafe; 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, + build_planning_parts, changed_followup_response_create_model, 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, + 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::ai_serving::ResponsesWebSocketPinnedCandidate; use crate::clock::current_unix_secs; -use crate::control::{request_model_local_rejection, GatewayControlDecision}; +use crate::control::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, + 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"; +const CLOSE_UNSUPPORTED_DATA: u16 = 1003; macro_rules! debug { ($($arg:tt)*) => { @@ -75,6 +78,65 @@ pub(super) fn adapter_drain_ready( )) } +fn parse_response_create_event(text: &str) -> Result { + let event = serde_json::from_str::(text).map_err(|_| "invalid_response_create")?; + 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("expected_response_create"); + } + Ok(event) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum ContinuationConstraint { + Pinned, + UpstreamUnavailable, + ModelChangeUnsupported, +} + +/// Classifies only response-chain turns. Independent turns return `None` and +/// remain free to use the normal planner. A non-null `previous_response_id` +/// must never fall through to that path: its state belongs to the connection +/// and provider key that produced it. +fn continuation_constraint( + event: &Value, + current_client_model: &str, + upstream_available: bool, +) -> Result, &'static str> { + if !response_create_has_previous_response_id(event) { + return Ok(None); + } + if changed_followup_response_create_model(event, current_client_model)?.is_some() { + return Ok(Some(ContinuationConstraint::ModelChangeUnsupported)); + } + if !upstream_available { + return Ok(Some(ContinuationConstraint::UpstreamUnavailable)); + } + Ok(Some(ContinuationConstraint::Pinned)) +} + +fn pinned_continuation_planning_event( + client_event: &Value, + current_client_model: &str, +) -> Result { + let mut planning_event = client_event.clone(); + // Validate an explicitly supplied value before replacing it with the + // canonical bound model. The routing decision has already established + // that this is not a model-changing continuation, but the planner must not + // receive a whitespace or case variant that fails exact mapping lookup. + response_create_model_or_current(&mut planning_event, current_client_model)?; + planning_event + .as_object_mut() + .ok_or("invalid_response_create")? + .insert( + "model".to_string(), + Value::String(current_client_model.to_string()), + ); + Ok(planning_event) +} + pub(super) async fn forward_client_message( client_message: AxumWsMessage, bound: &mut BoundResponsesConnection, @@ -85,39 +147,18 @@ pub(super) async fn forward_client_message( 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() { + let mut client_event = match parse_response_create_event(&text) { + Ok(event) => event, + Err(code) => { send_gateway_error( client_socket, - "responses_websocket_upstream_rebind_required", - "Send a new response.create to select another Provider connection", + code, + "WebSocket client text events must be response.create JSON objects", ) .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( @@ -134,8 +175,55 @@ pub(super) async fn forward_client_message( // 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 + 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; + } + }; + + // Build the per-turn planning shape before any policy check, then + // derive one strong live control snapshot that every stage below + // shares. The connection's Upgrade-time decision is only the + // immutable identity seed. + let planning_parts = build_planning_parts(context); + let turn_control = match resolve_responses_websocket_turn_control( + state, + context, + &planning_parts, + &client_event, + ) + .await + { + Ok(control) => control, + Err(error) => { + warn!( + event_name = "responses_websocket_followup_turn_control_rejected", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + error = ?error, + "gateway rejected a Responses WebSocket follow-up after live policy refresh" + ); + send_responses_websocket_turn_start_error(client_socket, &error).await; + return RelayDisposition::Continue; + } + }; + + match consume_response_create_rate_limit( + state, + &turn_control.decision, + turn_control.rpm_bypassed, + ) + .await { Ok(true) => {} Ok(false) => { @@ -166,25 +254,15 @@ pub(super) async fn forward_client_message( } } - 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, + &turn_control.decision, &client_event, ) .await; @@ -221,230 +299,345 @@ pub(super) async fn forward_client_message( 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; + match continuation_constraint( + &client_event, + &bound.client_model, + bound.upstream.is_some(), + ) { + Ok(Some(ContinuationConstraint::Pinned)) => { + return forward_pinned_continuation( + bound, + client_socket, + state, + context, + planning_parts, + client_event, + turn_control, + ) + .await; + } + Ok(Some(ContinuationConstraint::UpstreamUnavailable)) => { + send_gateway_error_with_status( + client_socket, + 503, + "responses_continuation_provider_unavailable", + "The bound provider connection is unavailable for this continuation", + ) + .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, + Ok(Some(ContinuationConstraint::ModelChangeUnsupported)) => { + send_gateway_error_with_status( + client_socket, + 409, + "responses_continuation_model_change_unsupported", + "A continuation cannot change models on the bound provider connection", + ) + .await; + return RelayDisposition::Continue; + } + Ok(None) => {} Err(code) => { send_gateway_error( client_socket, code, - "Gateway could not prepare the response.create event", + "response.create.model must be a non-empty string", ) .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( + } + forward_replanned_response_create( + bound, + client_socket, state, - &planning_parts, - &context.decision, - turn_decision, - &client_event, + context, + planning_parts, + client_event, + requested_model, + turn_control, ) .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::Binary(_) => { + mark_active_response_retry_unsafe(bound, "client_binary_frame"); + send_gateway_error( + client_socket, + "responses_websocket_binary_frame_unsupported", + "Responses WebSocket mode accepts text events only", + ) + .await; + close_client_socket(client_socket, CLOSE_UNSUPPORTED_DATA, "unsupported_data").await; + RelayDisposition::Close } - 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) => { + AxumWsMessage::Ping(data) => send_client_message(client_socket, AxumWsMessage::Pong(data)) + .await + .map(|()| RelayDisposition::Continue) + .unwrap_or(RelayDisposition::Close), + AxumWsMessage::Pong(_) => RelayDisposition::Continue, + AxumWsMessage::Close(_) => { if let Some(upstream) = bound.upstream.as_mut() { - close_upstream_socket(upstream, client_close_to_upstream(frame)).await; + // Do not forward an untrusted client close reason across the + // provider trust boundary. A neutral close is sufficient to + // tear down the bound upstream transport. + close_upstream_socket(upstream, None).await; } RelayDisposition::Close } } } +/// Revalidates a continuation against the live scheduler while keeping it on +/// the physical provider connection that owns its `previous_response_id` +/// state. This deliberately plans only the pinned provider/endpoint/key: an +/// eligible alternate key is valid for an independent turn, but not for an +/// in-flight response chain. +async fn forward_pinned_continuation( + bound: &mut BoundResponsesConnection, + client_socket: &mut WebSocket, + state: &AppState, + context: &WebSocketRequestContext, + planning_parts: http::request::Parts, + client_event: Value, + turn_control: ResponsesWebSocketTurnControl, +) -> RelayDisposition { + let Some(pinned_candidate) = + ResponsesWebSocketPinnedCandidate::from_decision(&bound.decision_template) + else { + warn!( + event_name = "responses_websocket_continuation_binding_identity_missing", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + "gateway cannot revalidate a continuation whose bound candidate identity is incomplete" + ); + send_gateway_error_with_status( + client_socket, + 503, + "responses_continuation_provider_unavailable", + "The bound provider connection cannot accept this continuation", + ) + .await; + return RelayDisposition::Continue; + }; + + // The public protocol allows follow-ups to omit `model`. The planner still + // needs the effective public model to enumerate the pinned mapping; this + // injected copy never replaces the opaque client event kept for audit and + // protocol-field restoration. + let planning_event = + match pinned_continuation_planning_event(&client_event, bound.client_model.as_str()) { + Ok(event) => event, + Err(code) => { + send_gateway_error( + client_socket, + code, + "response.create.model must be a non-empty string", + ) + .await; + return RelayDisposition::Continue; + } + }; + + let turn_request_id = Uuid::new_v4().to_string(); + let logical_turn_id = Uuid::new_v4().to_string(); + let planned = match await_owned_responses_websocket_plan(spawn_owned_responses_websocket_plan( + state.clone(), + planning_parts, + turn_request_id.clone(), + turn_control.decision.clone(), + turn_control.auth_snapshot.clone(), + planning_event, + None, + None, + Some(pinned_candidate), + )) + .await + { + Ok(Some(decision)) => decision, + Ok(None) => { + send_gateway_error_with_status( + client_socket, + 503, + "responses_continuation_provider_unavailable", + "The bound provider is not currently eligible for this continuation", + ) + .await; + return RelayDisposition::Continue; + } + Err(error) => { + warn!( + event_name = "responses_websocket_continuation_revalidation_failed", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + key_id = ?bound.decision_template.key_id, + error = ?error, + "gateway failed to revalidate the bound Responses WebSocket candidate" + ); + send_gateway_error_with_status( + client_socket, + 503, + "responses_continuation_provider_unavailable", + "Gateway could not revalidate the bound provider", + ) + .await; + return RelayDisposition::Continue; + } + }; + let OwnedResponsesWebSocketDecision { + planned, + planning_parts, + planned_lease, + } = planned; + let adapter = resolve_responses_websocket_adapter(planned.adapter); + let normalization = planned.normalization; + let decision = planned.execution; + let planned_provider_model = provider_model_from_decision(&decision); + if !decision_reuses_bound_upstream(bound, adapter, &decision) + || planned_provider_model.as_deref() != Some(bound.provider_model.as_str()) + { + planned_lease.release().await; + warn!( + event_name = "responses_websocket_continuation_binding_changed", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + key_id = ?decision.key_id, + "gateway rejected a continuation after the pinned candidate's physical binding changed" + ); + send_gateway_error_with_status( + client_socket, + 409, + "responses_continuation_binding_changed", + "The provider binding changed; reconnect or start an independent response", + ) + .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) => { + planned_lease.release().await; + send_gateway_error( + client_socket, + code, + "Gateway could not prepare the response.create event", + ) + .await; + return RelayDisposition::Continue; + } + }; + let outbound = match serde_json::to_string(&provider_event) { + Ok(outbound) => outbound, + Err(_) => { + planned_lease.release().await; + 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_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_with_planned_lease( + state, + &context.trace_id, + planning_parts, + &turn_control.decision, + turn_decision, + &client_event, + planned_lease, + ) + .await + { + Ok(turn) => turn, + Err(error) => { + warn!( + event_name = "responses_websocket_continuation_turn_start_failed", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + error = ?error, + "gateway could not admit a revalidated Responses WebSocket continuation" + ); + send_responses_websocket_turn_start_error(client_socket, &error).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()); + bound.adapter = adapter; + bound.decision_template = decision; + bound.body_normalization = normalization; + bound.turn_state.begin( + LogicalTurn::new(client_event, turn_index, logical_turn_id).with_turn_control(turn_control), + turn, + ); + bound.next_turn_index = bound.next_turn_index.saturating_add(1); + debug!( + event_name = "responses_websocket_continuation_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, + key_id = ?bound.decision_template.key_id, + candidate_revalidated = true, + "gateway forwarded a pinned Responses WebSocket continuation after runtime revalidation" + ); + RelayDisposition::Continue +} + /// 重新规划一轮 `response.create`(换模型或独立轮)。 /// /// `planning_parts` 与 `client_event` 都由调用方准备:事件已经过请求侧脱敏, @@ -455,85 +648,30 @@ async fn forward_replanned_response_create( client_socket: &mut WebSocket, state: &AppState, context: &WebSocketRequestContext, - planning_parts: &http::request::Parts, + planning_parts: http::request::Parts, client_event: Value, requested_model: String, + turn_control: ResponsesWebSocketTurnControl, ) -> 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_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, + (!excluded_codex_account_ids.is_empty()).then_some(excluded_codex_account_ids); + let planned = match await_owned_responses_websocket_plan(spawn_owned_responses_websocket_plan( + state.clone(), planning_parts, - &turn_request_id, - &context.decision, - &client_event, + turn_request_id.clone(), + turn_control.decision.clone(), + turn_control.auth_snapshot.clone(), + client_event.clone(), excluded_key_ids, excluded_codex_account_ids, - ) + None, + )) .await { Ok(Some(decision)) => decision, @@ -568,27 +706,15 @@ async fn forward_replanned_response_create( return RelayDisposition::Continue; } }; + let OwnedResponsesWebSocketDecision { + planned, + planning_parts, + planned_lease, + } = planned; 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) @@ -596,8 +722,7 @@ async fn forward_replanned_response_create( }) { Ok(event) => event, Err(code) => { - release_pool_key_lease_from_report_context(state, decision.report_context.as_ref()) - .await; + planned_lease.release().await; send_gateway_error( client_socket, code, @@ -619,12 +744,14 @@ async fn forward_replanned_response_create( &logical_turn_id, 1, ); - let mut turn = match begin_responses_websocket_turn( + let mut turn = match begin_responses_websocket_turn_with_planned_lease( state, + &context.trace_id, planning_parts, - &context.decision, + &turn_control.decision, turn_decision, &client_event, + planned_lease, ) .await { @@ -700,8 +827,9 @@ async fn forward_replanned_response_create( // 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), + LogicalTurn::new(client_event.clone(), turn_index, logical_turn_id.clone()) + .with_turn_control(turn_control), + turn, ); bound.next_turn_index = bound.next_turn_index.saturating_add(1); debug!( @@ -772,8 +900,8 @@ async fn forward_replanned_response_create( 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), + LogicalTurn::new(client_event, turn_index, logical_turn_id).with_turn_control(turn_control), + turn, ); bound.next_turn_index = bound.next_turn_index.saturating_add(1); bound.upstream_response_headers = replacement.upstream_response_headers; @@ -814,3 +942,90 @@ pub(super) async fn consume_response_create_rate_limit( FrontdoorUserRpmOutcome::Allowed | FrontdoorUserRpmOutcome::NotApplicable => Ok(true), } } + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::{ + continuation_constraint, parse_response_create_event, pinned_continuation_planning_event, + ContinuationConstraint, + }; + + #[test] + fn invalid_client_text_does_not_poison_the_next_response_create() { + assert_eq!( + parse_response_create_event("not-json"), + Err("invalid_response_create") + ); + assert_eq!( + parse_response_create_event("[]"), + Err("invalid_response_create") + ); + assert_eq!( + parse_response_create_event(r#"{"type":"response.cancel"}"#), + Err("expected_response_create") + ); + + let valid = parse_response_create_event( + r#"{"type":"response.create","model":"gpt-test","store":true}"#, + ) + .expect("a valid event after rejected text should still parse"); + assert_eq!(valid["type"], "response.create"); + assert_eq!(valid["store"], true); + } + + #[test] + fn continuation_without_an_upstream_never_falls_through_to_replanning() { + let continuation = json!({ + "type": "response.create", + "previous_response_id": "resp-previous", + }); + + assert_eq!( + continuation_constraint(&continuation, "gpt-current", false), + Ok(Some(ContinuationConstraint::UpstreamUnavailable)) + ); + assert_eq!( + continuation_constraint(&continuation, "gpt-current", true), + Ok(Some(ContinuationConstraint::Pinned)) + ); + assert_eq!( + continuation_constraint( + &json!({"type": "response.create", "model": "gpt-current"}), + "gpt-current", + false, + ), + Ok(None) + ); + } + + #[test] + fn model_changing_continuation_never_falls_through_to_replanning() { + let continuation = json!({ + "type": "response.create", + "model": "gpt-other", + "previous_response_id": "resp-previous", + }); + + assert_eq!( + continuation_constraint(&continuation, "gpt-current", true), + Ok(Some(ContinuationConstraint::ModelChangeUnsupported)) + ); + } + + #[test] + fn pinned_planning_uses_the_canonical_bound_model() { + let client_event = json!({ + "type": "response.create", + "model": " GPT-CURRENT ", + "previous_response_id": "resp-previous", + }); + + let planning = pinned_continuation_planning_event(&client_event, "gpt-current") + .expect("a same-model continuation should normalize"); + assert_eq!(planning["model"], "gpt-current"); + assert_eq!(client_event["model"], " GPT-CURRENT "); + assert_eq!(planning["previous_response_id"], "resp-previous"); + } +} diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/connection.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/connection.rs index 8900ccdb2..dd380a731 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/connection.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/connection.rs @@ -5,20 +5,18 @@ use std::time::Duration; use axum::extract::ws::{Message as AxumWsMessage, WebSocket}; use futures_util::{SinkExt, StreamExt}; use serde_json::Value; -use tokio::time::sleep; use wreq::ws::message::Message as WreqWsMessage; +use super::adapter::ResponsesWebSocketRelayDirective; use super::client::{adapter_drain_ready, forward_client_message, RelayDisposition}; -use super::frame::ParsedResponsesWebSocketFrame; +use super::frame::{encode_opaque_websocket_event, ParsedResponsesWebSocketFrame}; use super::lifecycle::{ await_pending_adapter_observation, finalize_active_turn, queue_turn_finalization, - settle_turn_finalization, ActiveProviderAttempt, PreviousAttemptSettled, + settle_turn_finalization, spawn_bounded_adapter_observation, PreviousAttemptSettled, }; use super::quota::{ - active_continuation_can_retry_from_full_input, detach_exhausted_upstream, - is_usage_limit_error_event, mark_active_response_retry_unsafe, + detach_exhausted_upstream, is_usage_limit_error_event, mark_active_response_retry_unsafe, observe_active_response_rebind_safety, retry_active_turn_after_quota_exhaustion, - send_previous_response_not_found, should_request_full_continuation_retry, }; use super::relay_policy::{ classify_quota_relay, fatal_relay_policy, FatalRelaySignal, QuotaRelayAction, QuotaRelayFacts, @@ -31,8 +29,7 @@ use super::turn::{ use super::upstream::{close_bound_upstream, receive_optional_upstream}; use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext; use crate::handlers::proxy::websocket::session::{ - wait_for_optional_deadline, CLOSE_INTERNAL_ERROR, CLOSE_TRY_AGAIN, - RESPONSES_WEBSOCKET_SESSION_LIMITS, WEBSOCKET_LOG_TRANSPORT, + wait_for_optional_deadline, CLOSE_INTERNAL_ERROR, CLOSE_TRY_AGAIN, WEBSOCKET_LOG_TRANSPORT, }; use crate::handlers::proxy::websocket::transport::{ close_client_socket, send_client_message, send_gateway_error_with_status, @@ -64,30 +61,10 @@ pub(super) async fn relay_bound_connection( bound: &mut BoundResponsesConnection, state: &AppState, context: &WebSocketRequestContext, - connection_permit: Option, ) { - let connection_deadline = sleep(RESPONSES_WEBSOCKET_SESSION_LIMITS.max_connection_duration); - tokio::pin!(connection_deadline); - loop { let active_turn_deadline = bound.turn_state.attempt().map(|turn| turn.deadline()); tokio::select! { - _ = &mut connection_deadline => { - finalize_active_turn( - bound, - state, - ResponsesWebSocketTurnOutcome::connection_limit_reached(), - ).await; - send_gateway_error_with_status( - client_socket, - 503, - "websocket_connection_limit_reached", - "WebSocket connection duration limit reached; reconnect to continue", - ).await; - close_bound_upstream(bound).await; - close_client_socket(client_socket, CLOSE_TRY_AGAIN, "connection_limit_reached").await; - break; - } _ = wait_for_optional_deadline(active_turn_deadline.map(|deadline| deadline.deadline)) => { let Some(turn_deadline) = active_turn_deadline else { continue; @@ -117,31 +94,6 @@ pub(super) async fn relay_bound_connection( ).await; break; } - _ = wait_for_connection_permit_loss(connection_permit.as_ref()) => { - let policy = fatal_relay_policy(FatalRelaySignal::ConnectionAdmissionLost); - warn!( - event_name = "responses_websocket_connection_admission_lost", - log_type = "ops", - transport = WEBSOCKET_LOG_TRANSPORT, - websocket = true, - trace_id = %context.trace_id, - "gateway closed Responses WebSocket after its connection admission became unhealthy" - ); - finalize_active_turn( - bound, - state, - ResponsesWebSocketTurnOutcome::connection_admission_lost(), - ).await; - close_bound_upstream(bound).await; - send_gateway_error_with_status( - client_socket, - policy.status_code, - policy.error_code, - policy.client_message, - ).await; - close_client_socket(client_socket, policy.close_code, policy.close_reason).await; - break; - } client_message = client_socket.next() => { let Some(client_message) = client_message else { finalize_active_turn( @@ -169,7 +121,15 @@ pub(super) async fn relay_bound_connection( close_bound_upstream(bound).await; break; }; - match forward_client_message(client_message, bound, client_socket, state, context).await { + match Box::pin(forward_client_message( + client_message, + bound, + client_socket, + state, + context, + )) + .await + { RelayDisposition::Continue => {} RelayDisposition::Close => { finalize_active_turn( @@ -288,7 +248,7 @@ pub(super) async fn relay_bound_connection( let state_for_observation = state.clone(); let trace_id = context.trace_id.clone(); let report_context = bound.decision_template.report_context.clone(); - bound.pending_adapter_observation = Some(tokio::spawn(async move { + bound.pending_adapter_observation = Some(spawn_bounded_adapter_observation(async move { adapter .persist_upstream_observation( &state_for_observation, @@ -380,19 +340,19 @@ pub(super) async fn relay_bound_connection( drain_ready: drain_for_adapter, retry_current_turn: bound .pending_adapter_drain - .is_some_and(|directive| directive.retry_current_turn), + .is_some_and(|directive| directive.retry_current_turn) + && bound + .turn_state + .logical() + .is_some_and(|turn| turn.quota_retry_block_reason().is_none()), transparent_retry_failed: false, usage_limit_error: parsed_upstream_event.is_some_and(is_usage_limit_error_event), - continuation_retry_eligible: active_continuation_can_retry_from_full_input(bound), upstream_closed: is_close, }; let mut quota_relay_action = classify_quota_relay(quota_facts); if matches!(quota_relay_action, QuotaRelayAction::AttemptTransparentRetry) { // detach_attempt 保留 logical turn:重试是同一轮请求的下一个 attempt。 - let retry_turn = bound - .turn_state - .detach_attempt() - .map(ActiveProviderAttempt::disarm); + let retry_turn = bound.turn_state.detach_attempt(); // 先结算旧 attempt 并等它落地,再规划下一个 attempt。两个理由: // // 1. 规划要读 health / adaptive / pool 状态,而这些正是旧 @@ -417,7 +377,14 @@ pub(super) async fn relay_bound_connection( } None => PreviousAttemptSettled::nothing_to_settle(), }; - if retry_active_turn_after_quota_exhaustion(bound, state, context, settled).await + // Planning and binding a replacement carries the complete + // scheduler/provider state machine. Keep that large future + // off the relay task's stack; the default Tokio/test worker + // stack is otherwise easy to exhaust on this rare branch. + if Box::pin(retry_active_turn_after_quota_exhaustion( + bound, state, context, settled, + )) + .await { continue; } @@ -430,42 +397,9 @@ pub(super) async fn relay_bound_connection( ..quota_facts }); } - if matches!( - quota_relay_action, - QuotaRelayAction::RequestFullContinuationRetry - ) { - let directive = bound - .pending_adapter_drain - .expect("adapter drain state should be present"); - debug!( - event_name = "responses_websocket_continuation_retry_required", - log_type = "event", - transport = WEBSOCKET_LOG_TRANSPORT, - websocket = true, - trace_id = %context.trace_id, - error_code = "previous_response_not_found", - "gateway will ask the client to retry the continuation with complete input" - ); - let mut turn = bound.turn_state.end().map(ActiveProviderAttempt::disarm); - if let Some(active_turn) = turn.as_mut() { - active_turn.release_admission().await; - } - send_previous_response_not_found(client_socket).await; - if let Some(turn) = turn { - queue_turn_finalization( - bound, - state, - turn, - terminal_outcome.unwrap_or_else( - ResponsesWebSocketTurnOutcome::upstream_closed, - ), - ) - .await; - } - detach_exhausted_upstream(bound, directive, &context.trace_id).await; - continue; - } - if matches!(quota_relay_action, QuotaRelayAction::ForwardQuotaAndDetach) { + let detach_after_forward = + matches!(quota_relay_action, QuotaRelayAction::ForwardQuotaAndDetach); + if detach_after_forward && is_close { let directive = bound .pending_adapter_drain .expect("adapter drain state should be present"); @@ -486,22 +420,127 @@ pub(super) async fn relay_bound_connection( detach_exhausted_upstream(bound, directive, &context.trace_id).await; continue; } - // 响应侧还原:HTTP 在把响应体交给客户端之前会把占位符换回真实值 - // (`privacy::restore_sync_response_body` / - // `privacy::StreamingResponseRestorer`),这里是 WS 的同一个位置 - // ——最后一跳之前,并且在 `capture_client_frame` 之前,所以审计与 - // 终态观测继续消费脱敏态的事件。没有命中还原时保持上游原字节。 - let restored_client_frame = parsed_upstream_frame + // Standard Responses frames cross the gateway byte-for-byte unless PII + // restoration has something to replace. Codex may wrap public events with + // provider-private side-channel chunks; only that explicit envelope is + // peeled, and each retained event is serialized as a complete opaque Value. + // Observation and capture continue to consume the redacted event, while the + // final client hop receives restored text. + let relay_directive = parsed_upstream_frame .as_ref() - .map(ParsedResponsesWebSocketFrame::event) - .and_then(|event| { - bound.redaction_restorer.restore_provider_frame_text(event) + .map(|frame| { + bound + .adapter + .relay_directive_for_upstream_event(frame.event()) }); - let client_frame = match restored_client_frame { - Some(restored) => AxumWsMessage::Text(restored.into()), - None => upstream_message_to_client(upstream_message.clone()), - }; - if let Err(error) = send_client_message(client_socket, client_frame).await { + let mut relay_send_error = None; + let mut relay_serialization_failed = false; + match relay_directive { + Some(ResponsesWebSocketRelayDirective::ForwardOriginal) => { + let restored = parsed_upstream_frame + .as_ref() + .and_then(|frame| { + bound + .redaction_restorer + .restore_provider_frame_text(frame.event()) + }); + let client_frame = match restored { + Some(text) => AxumWsMessage::Text(text.into()), + None => upstream_message_to_client(upstream_message.clone()), + }; + match send_client_message(client_socket, client_frame).await { + Ok(()) => { + if let (Some(turn), Some(frame)) = ( + bound.turn_state.attempt_mut(), + parsed_upstream_frame.as_ref(), + ) { + turn.capture_client_frame(frame.event()); + } + } + Err(error) => relay_send_error = Some(error), + } + } + Some(ResponsesWebSocketRelayDirective::ForwardEvents(events)) => { + for event in events { + let text = match bound + .redaction_restorer + .restore_provider_frame_text(event) + { + Some(restored) => restored, + None => match encode_opaque_websocket_event(event) { + Ok(encoded) => encoded, + Err(_) => { + relay_serialization_failed = true; + break; + } + }, + }; + match send_client_message( + client_socket, + AxumWsMessage::Text(text.into()), + ) + .await + { + Ok(()) => { + if let Some(turn) = bound.turn_state.attempt_mut() { + turn.capture_client_frame(event); + } + } + Err(error) => { + relay_send_error = Some(error); + break; + } + } + } + } + Some(ResponsesWebSocketRelayDirective::SuppressProviderPrivate) => {} + None => { + if let Err(error) = send_client_message( + client_socket, + upstream_message_to_client(upstream_message.clone()), + ) + .await + { + relay_send_error = Some(error); + } + } + } + if relay_serialization_failed { + warn!( + event_name = "responses_websocket_provider_event_serialization_failed", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + provider_terminal_reached = terminal_outcome.is_some(), + "gateway could not serialize an opaque provider event" + ); + bound + .turn_state + .record_client_delivery_aborted(CLIENT_DELIVERY_FAILED_REASON); + finalize_active_turn( + bound, + state, + settle_signal_for_client_delivery_failure(terminal_outcome), + ) + .await; + send_gateway_error_with_status( + client_socket, + 502, + "responses_websocket_event_serialization_failed", + "Gateway could not relay the provider event", + ) + .await; + close_bound_upstream(bound).await; + close_client_socket( + client_socket, + CLOSE_INTERNAL_ERROR, + "provider_event_serialization_failed", + ) + .await; + break; + } + if let Some(error) = relay_send_error { warn!( event_name = "responses_websocket_client_send_failed", log_type = "ops", @@ -525,11 +564,6 @@ pub(super) async fn relay_bound_connection( close_bound_upstream(bound).await; break; } - if let (Some(turn), Some(frame)) = - (bound.turn_state.attempt_mut(), parsed_upstream_frame.as_ref()) - { - turn.capture_client_frame(frame.event()); - } if let Some(outcome) = terminal_outcome { finalize_active_turn(bound, state, outcome).await; } else if is_close { @@ -540,6 +574,21 @@ pub(super) async fn relay_bound_connection( ) .await; } + if detach_after_forward { + let directive = bound + .pending_adapter_drain + .expect("adapter drain state should be present"); + if bound.turn_state.response_in_flight() { + finalize_active_turn( + bound, + state, + ResponsesWebSocketTurnOutcome::provider_quota_exhausted(), + ) + .await; + } + detach_exhausted_upstream(bound, directive, &context.trace_id).await; + continue; + } if drain_for_adapter { let directive = bound .pending_adapter_drain @@ -556,7 +605,9 @@ pub(super) async fn relay_bound_connection( } } -async fn wait_for_connection_permit_loss(permit: Option<&aether_runtime::AdmissionPermit>) { +pub(super) async fn wait_for_connection_permit_loss( + permit: Option<&aether_runtime::AdmissionPermit>, +) { let Some(permit) = permit else { std::future::pending::<()>().await; return; diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/control.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/control.rs new file mode 100644 index 000000000..ff9a994a3 --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/control.rs @@ -0,0 +1,173 @@ +//! Per-turn control-plane refresh for long-lived Responses WebSockets. +//! +//! An Upgrade authenticates the connection, but it must not freeze API-key, +//! wallet, IP, model, or RPM policy for up to an hour. This module produces one +//! live decision and its exact strong API-key snapshot for every +//! `response.create`; the caller uses that pair consistently for rate limiting, +//! redaction, model authorization, planning, admission, balance, and retries. + +use axum::http::StatusCode; +use serde_json::Value; + +use crate::ai_serving::GatewayAuthApiKeySnapshot; +use crate::control::{ + refresh_execution_runtime_auth_context_with_snapshot, request_model_local_rejection, + GatewayControlDecision, GatewayLocalAuthRejection, +}; +use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext; +use crate::handlers::proxy::websocket::session::WEBSOCKET_LOG_TRANSPORT; +use crate::handlers::shared::ip_rules_allow; +use crate::{AppState, GatewayError}; + +const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws"; + +macro_rules! warn { + ($($arg:tt)*) => { + tracing::warn!(target: LOG_TARGET, $($arg)*) + }; +} + +#[derive(Debug, Clone)] +pub(super) struct ResponsesWebSocketTurnControl { + pub(super) decision: GatewayControlDecision, + pub(super) auth_snapshot: Option, + pub(super) rpm_bypassed: bool, +} + +pub(super) async fn resolve_responses_websocket_turn_control( + state: &AppState, + context: &WebSocketRequestContext, + parts: &http::request::Parts, + client_event: &Value, +) -> Result { + if state + .admin_security_ip_blacklisted(context.client_ip) + .await? + { + return Err(GatewayError::Client { + status: StatusCode::FORBIDDEN, + message: "The current IP is blocked".to_string(), + }); + } + + let mut decision = context.decision.clone(); + let auth_snapshot = if let Some(auth_context) = decision.auth_context.take() { + let (refreshed, snapshot) = refresh_execution_runtime_auth_context_with_snapshot( + state, + auth_context, + decision.auth_endpoint_signature.as_deref(), + ) + .await?; + decision.local_auth_rejection = refreshed.local_rejection.clone(); + decision.auth_context = Some(refreshed); + snapshot + } else { + None + }; + // Model-directive configuration is mutable policy too; do not retain the + // Upgrade-time snapshot for the lifetime of the socket. + decision.model_directive_policy = + crate::system_features::ModelDirectivePolicySnapshot::load(state).await; + + if let Some(rejection) = decision.local_auth_rejection.clone() { + return Err(websocket_auth_rejection_error(rejection)); + } + let Some(auth_context) = decision.auth_context.as_ref() else { + return Err(websocket_auth_rejection_error( + GatewayLocalAuthRejection::InvalidApiKey, + )); + }; + if !auth_context.access_allowed + || auth_context.user_id.trim().is_empty() + || auth_context.api_key_id.trim().is_empty() + { + return Err(websocket_auth_rejection_error( + GatewayLocalAuthRejection::InvalidApiKey, + )); + } + if !ip_rules_allow(auth_context.ip_rules.as_deref(), context.client_ip) { + return Err(websocket_auth_rejection_error( + GatewayLocalAuthRejection::IpNotAllowed { + remote_ip: context.client_ip.to_string(), + }, + )); + } + + let body = serde_json::to_vec(client_event) + .map(axum::body::Bytes::from) + .map_err(|error| GatewayError::Internal(error.to_string()))?; + if let Some(rejection) = + request_model_local_rejection(state, Some(&decision), &parts.uri, &parts.headers, &body) + .await? + { + return Err(websocket_auth_rejection_error(rejection)); + } + + let rpm_bypassed = match state.admin_security_ip_whitelisted(context.client_ip).await { + Ok(value) => value, + Err(error) => { + warn!( + event_name = "responses_websocket_turn_ip_whitelist_check_failed", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + client_ip = %context.client_ip, + error = ?error, + "gateway applied ordinary WebSocket RPM after the live IP whitelist check failed" + ); + false + } + }; + + Ok(ResponsesWebSocketTurnControl { + decision, + auth_snapshot, + rpm_bypassed, + }) +} + +fn websocket_auth_rejection_error(rejection: GatewayLocalAuthRejection) -> GatewayError { + let (status, message) = match rejection { + GatewayLocalAuthRejection::InvalidApiKey => { + (StatusCode::UNAUTHORIZED, "The API key is invalid") + } + GatewayLocalAuthRejection::LockedApiKey => ( + StatusCode::FORBIDDEN, + "The API key is locked and cannot be used", + ), + GatewayLocalAuthRejection::WalletUnavailable => { + (StatusCode::FORBIDDEN, "The account wallet is unavailable") + } + GatewayLocalAuthRejection::BalanceDenied { remaining } => { + let message = match remaining { + Some(remaining) => format!("Insufficient balance (remaining: ${remaining:.2})"), + None => "Insufficient balance".to_string(), + }; + return GatewayError::Client { + status: StatusCode::TOO_MANY_REQUESTS, + message, + }; + } + GatewayLocalAuthRejection::ProviderNotAllowed { .. } => ( + StatusCode::FORBIDDEN, + "The provider is not allowed for this API key", + ), + GatewayLocalAuthRejection::ApiFormatNotAllowed { .. } => ( + StatusCode::FORBIDDEN, + "The API format is not allowed for this API key", + ), + GatewayLocalAuthRejection::ModelNotAllowed { .. } => ( + StatusCode::FORBIDDEN, + "The requested model is not allowed for this API key", + ), + GatewayLocalAuthRejection::IpNotAllowed { .. } => ( + StatusCode::UNAUTHORIZED, + "The current IP is not allowed for this API key", + ), + }; + GatewayError::Client { + status, + message: message.to_string(), + } +} diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/frame.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/frame.rs index d756bd652..16e0b1d2f 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/frame.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/frame.rs @@ -119,6 +119,17 @@ impl<'a> ParsedResponsesWebSocketFrame<'a> { } } +/// Encodes one event peeled from a provider-private envelope without applying +/// an event-type or field projection. +/// +/// Direct provider events should use [`ParsedResponsesWebSocketFrame::raw_text`] +/// so their bytes remain identical. This helper exists only for batch +/// envelopes that cannot be relayed as a whole: serializing the complete +/// [`Value`] preserves every known and future JSON member. +pub(super) fn encode_opaque_websocket_event(event: &Value) -> serde_json::Result { + serde_json::to_string(event) +} + /// Flattens a frame into the events it carries. An envelope may name its own /// `type` *and* batch further events under `chunks`; both are protocol events. fn protocol_events_of(event: &Value) -> Vec<&Value> { @@ -148,21 +159,6 @@ fn event_is_started(event: &Value) -> bool { ) } -/// `response.incomplete` 的合法终态 reason 白名单。 -/// -/// 这些 reason 表示上游按规则正常结束了本轮响应(写满 `max_output_tokens`、 -/// 命中内容过滤、按工具调用截断),标准流解析里它们会变成 `length` / -/// `content_filter` / `tool_calls` 这类正常 finish,和 -/// `openai_responses_incomplete_finish_reason` 的既有映射保持一致,因此不能 -/// 当成 provider failure 记账。 -const LEGITIMATE_RESPONSES_INCOMPLETE_REASONS: [&str; 5] = [ - "max_output_tokens", - "max_tokens", - "content_filter", - "tool_calls", - "function_call", -]; - /// 读取 `response.incomplete` 携带的 `incomplete_details.reason`。 /// /// 标准位置是 `response.incomplete_details.reason`;批量封装偶尔把 @@ -179,17 +175,33 @@ fn responses_incomplete_reason(event: &Value) -> Option<&str> { .find(|reason| !reason.is_empty()) } -/// 判断一个 `response.incomplete` 是否是合法终态。 +fn responses_incomplete_has_explicit_error(event: &Value) -> bool { + [event.get("error"), event.pointer("/response/error")] + .into_iter() + .flatten() + .any(|error| !error.is_null()) +} + +/// Derives only the fallback status for `response.incomplete`. /// -/// reason 缺失或不在白名单内(例如 `error`、`server_error`)时继续按 -/// provider failure 处理:这类 incomplete 说明上游确实没能正常收尾,仍应扣 -/// 供应商健康分。 -fn responses_incomplete_is_legitimate_terminal(event: &Value) -> bool { - responses_incomplete_reason(event).is_some_and(|reason| { - LEGITIMATE_RESPONSES_INCOMPLETE_REASONS - .iter() - .any(|candidate| reason.eq_ignore_ascii_case(candidate)) - }) +/// A non-empty reason is provider-owned protocol data. Treating it as a fixed +/// allowlist would turn every future legitimate reason into a synthetic 502 +/// and incorrectly penalize provider health. Missing/malformed reasons and +/// explicit error markers still fail closed; numeric status and recognized +/// error codes continue to override this fallback in +/// [`websocket_event_status_code`]. +fn responses_incomplete_default_status(event: &Value) -> u16 { + match responses_incomplete_reason(event) { + None => 502, + Some(reason) + if reason.eq_ignore_ascii_case("error") + || reason.eq_ignore_ascii_case("server_error") => + { + 502 + } + Some(_) if responses_incomplete_has_explicit_error(event) => 502, + Some(_) => 200, + } } fn terminal_for_event(event: &Value) -> Option { @@ -198,18 +210,13 @@ fn terminal_for_event(event: &Value) -> Option status_code: websocket_event_status_code(event, 200), cancelled: false, }), - // 合法 incomplete(例如写满 max_output_tokens)是正常终态,默认按 200 - // 记账,不再一律当 502 provider failure;reason 缺失或未知时保留原来的 - // 502 默认值。显式 `status_code` 和 error code 映射仍然优先于默认值, - // 所以带 `rate_limit_exceeded` 的 incomplete 依旧是 429。 + // A non-empty provider reason is a normal terminal by default, including + // future reasons Aether does not yet know. Explicit status/error data + // still wins, so quota and server failures retain their failure status. "response.incomplete" => Some(ResponsesWebSocketFrameTerminal { status_code: websocket_event_status_code( event, - if responses_incomplete_is_legitimate_terminal(event) { - 200 - } else { - 502 - }, + responses_incomplete_default_status(event), ), cancelled: false, }), @@ -282,7 +289,9 @@ fn safe_websocket_event_label(value: &str) -> String { #[cfg(test)] mod tests { - use super::ParsedResponsesWebSocketFrame; + use serde_json::json; + + use super::{encode_opaque_websocket_event, ParsedResponsesWebSocketFrame}; #[test] fn parses_started_frame_once_with_raw_text_and_event_metadata() { @@ -298,6 +307,37 @@ mod tests { assert_eq!(frame.event_type_for_log(), "response.in_progress"); } + #[test] + fn future_response_event_keeps_its_exact_original_text_and_unknown_fields() { + let raw = "{ \n \"future_top_level\": {\"nested\": [1, true, null]}, \n \"type\": \"response.future_capability.delta\", \n \"delta\": {\"new_wire_shape\": \"opaque\"}\n}"; + let frame = ParsedResponsesWebSocketFrame::parse(raw).expect("valid future event"); + + assert_eq!(frame.raw_text(), raw); + assert_eq!( + frame.event()["future_top_level"], + json!({"nested": [1, true, null]}) + ); + assert_eq!(frame.event()["delta"], json!({"new_wire_shape": "opaque"})); + assert!(!frame.is_terminal()); + } + + #[test] + fn peeled_batch_event_encoding_preserves_the_complete_opaque_value() { + let frame = ParsedResponsesWebSocketFrame::parse( + r#"{"chunks":[{"type":"response.future.done","future_capability":{"mode":"new"},"response":{"id":"resp_future","future_usage":{"novel_tokens":7}}}]}"#, + ) + .expect("valid private envelope"); + let events = frame.protocol_events(); + let event = events.first().expect("one future response event"); + let encoded = encode_opaque_websocket_event(event).expect("Value serialization succeeds"); + let round_trip: serde_json::Value = + serde_json::from_str(&encoded).expect("encoded event stays valid JSON"); + + assert_eq!(round_trip, **event); + assert_eq!(round_trip["future_capability"], json!({"mode": "new"})); + assert_eq!(round_trip["response"]["future_usage"]["novel_tokens"], 7); + } + #[test] fn classifies_terminal_status_and_cancellation() { let completed = ParsedResponsesWebSocketFrame::parse( @@ -373,7 +413,7 @@ mod tests { } #[test] - fn an_incomplete_without_a_legitimate_reason_stays_a_provider_failure() { + fn an_incomplete_without_a_reason_or_with_a_failure_reason_stays_a_provider_failure() { for raw in [ r#"{"type":"response.incomplete"}"#, r#"{"type":"response.incomplete","response":{"incomplete_details":null}}"#, @@ -386,11 +426,26 @@ mod tests { assert_eq!( frame.status(), Some(502), - "an incomplete without a known-good reason must stay a provider failure: {raw}" + "an incomplete without a usable reason must stay a provider failure: {raw}" ); } } + #[test] + fn a_future_incomplete_reason_is_forward_compatible_without_hiding_explicit_errors() { + let future = ParsedResponsesWebSocketFrame::parse( + r#"{"type":"response.incomplete","response":{"incomplete_details":{"reason":"future_context_boundary"}}}"#, + ) + .expect("valid frame"); + assert_eq!(future.status(), Some(200)); + + let future_with_error = ParsedResponsesWebSocketFrame::parse( + r#"{"type":"response.incomplete","response":{"error":{"code":"future_provider_error"},"incomplete_details":{"reason":"future_context_boundary"}}}"#, + ) + .expect("valid frame"); + assert_eq!(future_with_error.status(), Some(502)); + } + #[test] fn a_legitimate_incomplete_still_respects_an_explicit_provider_status() { let explicit = ParsedResponsesWebSocketFrame::parse( diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/lifecycle.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/lifecycle.rs index ba14076ef..708517b66 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/lifecycle.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/lifecycle.rs @@ -6,13 +6,13 @@ use std::time::Duration; use axum::extract::ws::WebSocket; +use axum::http::StatusCode; use tokio::task::JoinHandle; use tokio::time::timeout; use super::state::BoundResponsesConnection; use super::turn::{ - spawn_responses_websocket_turn_finalization, ResponsesProviderAttempt, - ResponsesWebSocketTurnOutcome, + begin_unowned_responses_websocket_turn, ResponsesProviderAttempt, ResponsesWebSocketTurnOutcome, }; use crate::handlers::proxy::websocket::session::{ CLOSE_INTERNAL_ERROR, CLOSE_POLICY_VIOLATION, CLOSE_TRY_AGAIN, WEBSOCKET_LOG_TRANSPORT, @@ -76,11 +76,90 @@ impl std::ops::DerefMut for ActiveProviderAttempt { } } +/// Starts a turn and arms its cancellation fallback before control returns to +/// code that can await an upstream bind or socket write. +pub(super) async fn begin_responses_websocket_turn( + state: &AppState, + trace_id: &str, + parts: http::request::Parts, + control_decision: &crate::control::GatewayControlDecision, + decision: crate::ai_serving::AiExecutionDecision, + client_event: &serde_json::Value, +) -> Result { + let state = state.clone(); + let trace_id = trace_id.to_string(); + let owner_timeout = state + .frontdoor_runtime_guards + .local_execution_planning_timeout; + let control_decision = control_decision.clone(); + let client_event = client_event.clone(); + + // Beginning an attempt performs several indispensable async writes before + // an `ActiveProviderAttempt` can exist (balance/admission, Pending usage, + // and candidate state). Run that whole transition in an owned task. If the + // relay/session future is cancelled while awaiting it, Tokio detaches this + // task; it still reaches either an explicitly cleaned-up error or an armed + // guard whose dropped output finalizes the attempt. + await_owned_turn_begin( + async move { + let turn = begin_unowned_responses_websocket_turn( + &state, + &parts, + &control_decision, + decision, + &client_event, + ) + .await?; + Ok(ActiveProviderAttempt::new(&state, turn)) + }, + owner_timeout, + trace_id, + ) + .await +} + +async fn await_owned_turn_begin( + begin: impl std::future::Future> + Send + 'static, + owner_timeout: Duration, + trace_id: String, +) -> Result +where + T: Send + 'static, +{ + await_owned_turn_begin_with_timeout(begin, owner_timeout, trace_id).await +} + +async fn await_owned_turn_begin_with_timeout( + begin: impl std::future::Future> + Send + 'static, + owner_timeout: Duration, + trace_id: String, +) -> Result +where + T: Send + 'static, +{ + tokio::spawn(async move { + tokio::time::timeout(owner_timeout, begin) + .await + .map_err(|_| GatewayError::LocalExecutionPlanningTimeout { + trace_id, + phase: "responses_websocket_turn_begin_owner", + timeout_ms: owner_timeout.as_millis() as u64, + })? + }) + .await + .map_err(|error| { + GatewayError::Internal(format!( + "Responses WebSocket turn begin task failed before ownership transfer: {error}" + )) + })? +} + impl Drop for ActiveProviderAttempt { fn drop(&mut self) { let Some(turn) = self.turn.take() else { return; }; + let outcome = turn.abandonment_outcome(); let state = self.state.clone(); // No runtime means the process is going down; the spawn could not // complete anyway. @@ -93,11 +172,7 @@ impl Drop for ActiveProviderAttempt { "gateway finalized a Responses WebSocket turn whose relay task went away" ); handle.spawn(async move { - turn.finalize_detached( - &state, - ResponsesWebSocketTurnOutcome::relay_task_abandoned(), - ) - .await; + turn.finalize_detached(&state, outcome).await; }); } } @@ -113,20 +188,23 @@ pub(super) async fn finalize_active_turn( outcome: ResponsesWebSocketTurnOutcome, ) { if let Some(turn) = bound.turn_state.end() { - queue_turn_finalization(bound, state, turn.disarm(), outcome).await; + queue_turn_finalization(bound, state, turn, outcome).await; } } pub(super) async fn queue_turn_finalization( bound: &mut BoundResponsesConnection, state: &AppState, - turn: ResponsesProviderAttempt, + turn: ActiveProviderAttempt, outcome: ResponsesWebSocketTurnOutcome, ) { await_pending_adapter_observation(bound).await; await_pending_turn_finalization(bound).await; - bound.pending_turn_finalization = - Some(spawn_responses_websocket_turn_finalization(state.clone(), turn, outcome).await); + bound.pending_turn_finalization = Some(spawn_guarded_turn_finalization( + state.clone(), + turn, + outcome, + )); } /// 「上一个 attempt 已经结算完毕」的凭证。 @@ -153,7 +231,7 @@ impl PreviousAttemptSettled { pub(super) async fn settle_turn_finalization( bound: &mut BoundResponsesConnection, state: &AppState, - turn: ResponsesProviderAttempt, + turn: ActiveProviderAttempt, outcome: ResponsesWebSocketTurnOutcome, ) -> PreviousAttemptSettled { queue_turn_finalization(bound, state, turn, outcome).await; @@ -161,42 +239,68 @@ pub(super) async fn settle_turn_finalization( PreviousAttemptSettled(()) } +pub(super) fn spawn_bounded_adapter_observation( + observation: impl std::future::Future + Send + 'static, +) -> JoinHandle<()> { + spawn_bounded_adapter_observation_with_timeout( + observation, + RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT, + ) +} + +fn spawn_bounded_adapter_observation_with_timeout( + observation: impl std::future::Future + Send + 'static, + owner_timeout: Duration, +) -> JoinHandle<()> { + tokio::spawn(async move { + if timeout(owner_timeout, observation).await.is_err() { + warn!( + event_name = "responses_websocket_adapter_observation_timeout", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + timeout_ms = owner_timeout.as_millis() as u64, + "gateway stopped a timed-out Responses WebSocket adapter observation" + ); + } + }) +} + pub(super) async fn await_pending_adapter_observation(bound: &mut BoundResponsesConnection) { - if let Some(mut handle) = bound.pending_adapter_observation.take() { - match timeout(RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT, &mut handle).await { - Ok(Err(error)) => { - warn!( - event_name = "responses_websocket_adapter_observation_join_failed", - log_type = "ops", - transport = WEBSOCKET_LOG_TRANSPORT, - websocket = true, - error = ?error, - "gateway Responses WebSocket adapter observation task failed" - ); - } - Ok(Ok(())) => {} - Err(_) => { - handle.abort(); - let _ = handle.await; - warn!( - event_name = "responses_websocket_adapter_observation_timeout", - log_type = "ops", - transport = WEBSOCKET_LOG_TRANSPORT, - websocket = true, - timeout_ms = RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT.as_millis() as u64, - "gateway stopped waiting for a Responses WebSocket adapter observation" - ); - } + if let Some(handle) = bound.pending_adapter_observation.take() { + if let Err(error) = handle.await { + warn!( + event_name = "responses_websocket_adapter_observation_join_failed", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + error = ?error, + "gateway Responses WebSocket adapter observation task failed" + ); } } } -pub(super) async fn finalize_unbound_turn( +pub(super) fn finalize_unbound_turn( state: AppState, - turn: ResponsesProviderAttempt, + turn: ActiveProviderAttempt, outcome: ResponsesWebSocketTurnOutcome, ) -> JoinHandle<()> { - spawn_responses_websocket_turn_finalization(state, turn, outcome).await + spawn_guarded_turn_finalization(state, turn, outcome) +} + +fn spawn_guarded_turn_finalization( + state: AppState, + turn: ActiveProviderAttempt, + outcome: ResponsesWebSocketTurnOutcome, +) -> JoinHandle<()> { + // Spawn synchronously while the armed guard is still owned here. Caller + // cancellation cannot drop an unguarded attempt between cleanup awaits. + tokio::spawn(async move { + let mut turn = turn; + turn.release_admission().await; + turn.disarm().finalize_detached(&state, outcome).await; + }) } pub(super) async fn await_turn_finalization_handle(handle: JoinHandle<()>) { @@ -228,6 +332,7 @@ pub(super) async fn send_responses_websocket_turn_start_error( client_socket: &mut WebSocket, error: &GatewayError, ) { + let status_code = responses_websocket_turn_start_http_status(error); match error { GatewayError::Client { status, message } => { let (error_type, code) = if status.as_u16() == 429 { @@ -235,19 +340,13 @@ pub(super) async fn send_responses_websocket_turn_start_error( } else { ("invalid_request_error", "gateway_request_not_allowed") }; - send_responses_websocket_error( - client_socket, - status.as_u16(), - error_type, - code, - message, - ) - .await; + send_responses_websocket_error(client_socket, status_code, error_type, code, message) + .await; } GatewayError::AdmissionTimeout { .. } => { send_responses_websocket_error( client_socket, - 503, + status_code, "server_error", "gateway_admission_timeout", "Gateway capacity is busy; retry this response", @@ -257,7 +356,7 @@ pub(super) async fn send_responses_websocket_turn_start_error( GatewayError::LocalExecutionPlanningTimeout { .. } => { send_responses_websocket_error( client_socket, - 504, + status_code, "server_error", "gateway_planning_timeout", "Gateway planning timed out; retry this response", @@ -267,7 +366,7 @@ pub(super) async fn send_responses_websocket_turn_start_error( _ => { send_responses_websocket_error( client_socket, - 500, + status_code, "server_error", "responses_websocket_turn_start_failed", "Gateway could not start this response", @@ -277,6 +376,15 @@ pub(super) async fn send_responses_websocket_turn_start_error( } } +fn responses_websocket_turn_start_http_status(error: &GatewayError) -> u16 { + match error { + GatewayError::Client { status, .. } => status.as_u16(), + GatewayError::AdmissionTimeout { .. } => StatusCode::TOO_MANY_REQUESTS.as_u16(), + GatewayError::LocalExecutionPlanningTimeout { .. } => StatusCode::GATEWAY_TIMEOUT.as_u16(), + _ => StatusCode::INTERNAL_SERVER_ERROR.as_u16(), + } +} + pub(super) fn responses_websocket_turn_start_close(error: &GatewayError) -> (u16, &'static str) { match error { GatewayError::Client { .. } => (CLOSE_POLICY_VIOLATION, "request_not_allowed"), @@ -292,7 +400,27 @@ mod tests { use std::sync::Arc; use std::time::Duration; - use super::await_turn_finalization_handle; + use super::{ + await_owned_turn_begin, await_owned_turn_begin_with_timeout, + await_turn_finalization_handle, responses_websocket_turn_start_close, + responses_websocket_turn_start_http_status, spawn_bounded_adapter_observation_with_timeout, + }; + use crate::GatewayError; + + #[test] + fn admission_timeout_uses_http_429_and_keeps_the_retry_later_close_code() { + let error = GatewayError::AdmissionTimeout { + trace_id: "turn-admission".to_string(), + gate: "gateway_upstream_execution", + queue_budget_ms: 25, + }; + + assert_eq!(responses_websocket_turn_start_http_status(&error), 429); + assert_eq!( + responses_websocket_turn_start_close(&error), + (1013, "gateway_busy") + ); + } /// C6 依赖的性质:结算是「等到落地」而不是「排进队列」。 /// @@ -348,4 +476,113 @@ mod tests { let handle = tokio::spawn(async { panic!("settlement task exploded") }); await_turn_finalization_handle(handle).await; } + + #[tokio::test] + async fn cancelling_the_caller_does_not_cancel_turn_begin_or_drop_an_unowned_result() { + struct DropProbe(Arc); + impl Drop for DropProbe { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } + } + + let begin_finished = Arc::new(AtomicBool::new(false)); + let result_dropped = Arc::new(AtomicBool::new(false)); + let finished = Arc::clone(&begin_finished); + let dropped = Arc::clone(&result_dropped); + let caller = tokio::spawn(async move { + await_owned_turn_begin( + async move { + tokio::time::sleep(Duration::from_millis(60)).await; + finished.store(true, Ordering::SeqCst); + Ok(DropProbe(dropped)) + }, + Duration::from_secs(1), + "turn-begin-cancel".to_string(), + ) + .await + }); + + tokio::time::sleep(Duration::from_millis(10)).await; + caller.abort(); + let _ = caller.await; + tokio::time::sleep(Duration::from_millis(120)).await; + + assert!( + begin_finished.load(Ordering::SeqCst), + "the owned begin task must outlive its cancelled relay caller" + ); + assert!( + result_dropped.load(Ordering::SeqCst), + "an undeliverable armed result must be dropped so its cleanup guard runs" + ); + } + + #[tokio::test] + async fn turn_begin_owner_deadline_drops_stalled_work_and_its_guards() { + struct DropProbe(Arc); + impl Drop for DropProbe { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } + } + + let dropped = Arc::new(AtomicBool::new(false)); + let task_dropped = Arc::clone(&dropped); + let result: Result<(), GatewayError> = await_owned_turn_begin_with_timeout( + async move { + let _probe = DropProbe(task_dropped); + std::future::pending::<()>().await; + Ok(()) + }, + Duration::from_millis(20), + "turn-begin-deadline".to_string(), + ) + .await; + + assert!(matches!( + result, + Err(GatewayError::LocalExecutionPlanningTimeout { + trace_id, + phase: "responses_websocket_turn_begin_owner", + timeout_ms: 20, + }) if trace_id == "turn-begin-deadline" + )); + assert!( + dropped.load(Ordering::SeqCst), + "owner timeout must drop the stalled begin future so RAII cleanup runs" + ); + } + + #[tokio::test] + async fn cancelling_observation_waiter_cannot_bypass_the_owner_timeout() { + struct DropProbe(Arc); + impl Drop for DropProbe { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } + } + + let dropped = Arc::new(AtomicBool::new(false)); + let task_dropped = Arc::clone(&dropped); + let observation = async move { + let _probe = DropProbe(task_dropped); + std::future::pending::<()>().await; + }; + let owner = + spawn_bounded_adapter_observation_with_timeout(observation, Duration::from_millis(20)); + let waiter = tokio::spawn(async move { + let _ = owner.await; + }); + waiter.abort(); + let _ = waiter.await; + + tokio::time::timeout(Duration::from_secs(1), async { + while !dropped.load(Ordering::SeqCst) { + tokio::task::yield_now().await; + } + }) + .await + .expect("detached observation owner must enforce its own timeout"); + } } diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/mod.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/mod.rs index f26018a6d..2aa903e78 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/mod.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/mod.rs @@ -12,9 +12,11 @@ mod admission; mod binding; mod client; mod connection; +mod control; mod frame; mod lifecycle; mod observation; +mod ownership; mod quota; mod redaction; mod relay_policy; @@ -61,5 +63,4 @@ pub(crate) async fn responses_websocket( const RESPONSES_WEBSOCKET_INGRESS_SPEC: WebSocketIngressSpec = WebSocketIngressSpec { route_unavailable_message: "WebSocket route is unavailable", - ip_whitelist_failure_event_name: "responses_websocket_ip_whitelist_check_failed", }; diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/observation.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/observation.rs index 7ea8c83c6..035d029d4 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/observation.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/observation.rs @@ -37,7 +37,11 @@ impl ResponsesStructuredTerminalObserver { /// 第一个被拒绝的事件就停止推进并把摘要标成 parser_error:解析器的状态机是 /// 有顺序的,跳过一个事件继续喂后面的只会得到更没意义的摘要。 pub(super) fn observe_events(&mut self, report_context: &Value, events: &[&Value]) { - for event in events { + for event in events + .iter() + .copied() + .filter(|event| event_is_relevant_to_terminal_observation(event)) + { if let Err(error) = self.inner.push_event(report_context, event) { self.inner.disable_with_error(error.to_string()); break; @@ -61,6 +65,28 @@ impl ResponsesStructuredTerminalObserver { } } +/// The WebSocket relay is not a Responses schema gateway. It forwards all +/// events opaquely, while this observer consumes only identity/terminal +/// snapshots needed for usage and settlement. In particular, a future +/// `response.*` delta must not become an observation failure merely because +/// Aether's canonical streaming parser does not know it yet. +fn event_is_relevant_to_terminal_observation(event: &Value) -> bool { + matches!( + event.get("type").and_then(Value::as_str), + Some( + "response.created" + | "response.in_progress" + | "response.queued" + | "response.completed" + | "response.done" + | "response.failed" + | "response.incomplete" + | "response.cancelled" + | "error" + ) + ) +} + #[cfg(test)] mod tests { use serde_json::json; @@ -141,6 +167,45 @@ mod tests { assert_eq!(usage.dimensions.get("total_tokens"), Some(&json!(4))); } + #[test] + fn future_and_provider_private_events_are_ignored_only_by_the_side_observer() { + let context = report_context(); + let private = json!({ + "type": "codex.response.metadata", + "private_future_field": {"shape": "unknown"}, + }); + let future = json!({ + "type": "response.future_capability.delta", + "future_capability": {"nested": [1, 2, 3]}, + }); + let completed = json!({ + "type": "response.completed", + "response": { + "id": "resp_future", + "model": "future-model", + "status": "completed", + "future_response_field": {"also": "unknown"}, + "usage": {"input_tokens": 5, "output_tokens": 2, "total_tokens": 7}, + }, + }); + + let mut observer = ResponsesStructuredTerminalObserver::default(); + observer.observe_events(&context, &[&private, &future, &completed]); + let summary = observer.finish(&context); + + assert!(summary.observed_finish); + assert_eq!(summary.response_id.as_deref(), Some("resp_future")); + assert_eq!(summary.unknown_event_count, 0); + assert_eq!( + summary + .standardized_usage + .as_ref() + .map(|usage| (usage.input_tokens, usage.output_tokens)), + Some((5, 2)) + ); + assert!(summary.parser_error.is_none()); + } + #[test] fn a_disabled_observer_reports_the_parser_error() { let context = report_context(); diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/ownership.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/ownership.rs new file mode 100644 index 000000000..361c264ee --- /dev/null +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/ownership.rs @@ -0,0 +1,245 @@ +//! Cancellation-safe ownership handoff for WebSocket planning leases. +//! +//! The relay races every turn against connection and response deadlines. A +//! planner future therefore cannot directly own a distributed pool-key lease: +//! losing the race would drop the future between scheduler selection and turn +//! startup, leaving that key unavailable until the lease TTL elapsed. + +use std::collections::BTreeSet; +use std::time::Duration; + +use serde_json::Value; +use tokio::task::JoinHandle; + +use super::lifecycle::{begin_responses_websocket_turn, ActiveProviderAttempt}; +use crate::ai_serving::{ + maybe_build_responses_websocket_decision, AiExecutionDecision, GatewayAuthApiKeySnapshot, + ResponsesWebSocketDecision, ResponsesWebSocketPinnedCandidate, +}; +use crate::control::GatewayControlDecision; +use crate::orchestration::release_pool_key_lease_from_report_context; +use crate::{AppState, GatewayError}; + +/// Owns a selected pool-key lease until the attempt lifecycle has taken over +/// the decision report context. +pub(super) struct PlannedPoolKeyLeaseGuard { + state: AppState, + report_context: Option, +} + +/// Planner output coupled to both its request parts and lease guard. +pub(super) struct OwnedResponsesWebSocketDecision { + pub(super) planned: ResponsesWebSocketDecision, + pub(super) planning_parts: http::request::Parts, + pub(super) planned_lease: PlannedPoolKeyLeaseGuard, +} + +/// Runs planning in an owner task. Dropping the caller's waiter detaches this +/// task; an unobserved successful output drops its guard and releases the +/// selected pool-key lease. +#[allow(clippy::too_many_arguments)] +pub(super) fn spawn_owned_responses_websocket_plan( + state: AppState, + parts: http::request::Parts, + trace_id: String, + control_decision: GatewayControlDecision, + auth_snapshot: Option, + client_event: Value, + excluded_key_ids: Option>, + excluded_codex_account_ids: Option>, + pinned_candidate: Option, +) -> JoinHandle, GatewayError>> { + let owner_timeout = state + .frontdoor_runtime_guards + .local_execution_planning_timeout; + tokio::spawn(async move { + let planned = await_owned_planning_deadline( + maybe_build_responses_websocket_decision( + &state, + &parts, + &trace_id, + &control_decision, + auth_snapshot.as_ref(), + &client_event, + excluded_key_ids.as_ref(), + excluded_codex_account_ids.as_ref(), + pinned_candidate.as_ref(), + ), + owner_timeout, + ) + .await + .map_err(|_| GatewayError::LocalExecutionPlanningTimeout { + trace_id: trace_id.clone(), + phase: "responses_websocket_plan_owner", + timeout_ms: owner_timeout.as_millis() as u64, + })??; + + Ok(planned.map(|planned| { + let planned_lease = + PlannedPoolKeyLeaseGuard::new(&state, planned.execution.report_context.as_ref()); + OwnedResponsesWebSocketDecision { + planned, + planning_parts: parts, + planned_lease, + } + })) + }) +} + +async fn await_owned_planning_deadline( + planning: F, + deadline: Duration, +) -> Result +where + F: std::future::Future, +{ + tokio::time::timeout(deadline, planning).await +} + +pub(super) async fn await_owned_responses_websocket_plan( + handle: JoinHandle, GatewayError>>, +) -> Result, GatewayError> { + handle.await.map_err(|error| { + GatewayError::Internal(format!( + "Responses WebSocket planning task failed before ownership transfer: {error}" + )) + })? +} + +impl PlannedPoolKeyLeaseGuard { + fn new(state: &AppState, report_context: Option<&Value>) -> Self { + Self { + state: state.clone(), + report_context: report_context.cloned(), + } + } + + pub(super) async fn release(mut self) { + release_pool_key_lease_from_report_context(&self.state, self.report_context.as_ref()).await; + self.report_context = None; + } + + fn disarm(&mut self) { + self.report_context = None; + } +} + +impl Drop for PlannedPoolKeyLeaseGuard { + fn drop(&mut self) { + let Some(report_context) = self.report_context.take() else { + return; + }; + let state = self.state.clone(); + if let Ok(handle) = tokio::runtime::Handle::try_current() { + handle.spawn(async move { + release_pool_key_lease_from_report_context(&state, Some(&report_context)).await; + }); + } + } +} + +/// Keeps the planning guard in the same detached owner task as lifecycle +/// startup. If the relay loses a deadline race while awaiting startup, the +/// task completes the handoff (or releases the lease on failure) without a +/// cancellation gap. +pub(super) async fn begin_responses_websocket_turn_with_planned_lease( + state: &AppState, + trace_id: &str, + parts: http::request::Parts, + control_decision: &GatewayControlDecision, + decision: AiExecutionDecision, + client_event: &Value, + mut planned_lease: PlannedPoolKeyLeaseGuard, +) -> Result { + let state = state.clone(); + let trace_id = trace_id.to_string(); + let control_decision = control_decision.clone(); + let client_event = client_event.clone(); + tokio::spawn(async move { + let turn = begin_responses_websocket_turn( + &state, + &trace_id, + parts, + &control_decision, + decision, + &client_event, + ) + .await?; + // ActiveProviderAttempt now owns the report context containing the + // lease. No await occurs between that handoff and disarming the guard. + planned_lease.disarm(); + Ok(turn) + }) + .await + .map_err(|error| { + GatewayError::Internal(format!( + "Responses WebSocket guarded turn startup task failed: {error}" + )) + })? +} + +#[cfg(test)] +mod tests { + use std::sync::atomic::{AtomicUsize, Ordering}; + use std::sync::Arc; + use std::time::Duration; + + struct DropProbe(Arc); + + impl Drop for DropProbe { + fn drop(&mut self) { + self.0.fetch_add(1, Ordering::SeqCst); + } + } + + #[tokio::test] + async fn dropping_a_planning_waiter_detaches_the_owner_and_drops_its_output() { + let started = Arc::new(tokio::sync::Notify::new()); + let release = Arc::new(tokio::sync::Notify::new()); + let dropped = Arc::new(AtomicUsize::new(0)); + let task_started = Arc::clone(&started); + let task_release = Arc::clone(&release); + let task_dropped = Arc::clone(&dropped); + let owner = tokio::spawn(async move { + task_started.notify_one(); + task_release.notified().await; + DropProbe(task_dropped) + }); + started.notified().await; + + let waiter = tokio::spawn(async move { + let _ = owner.await; + }); + waiter.abort(); + let _ = waiter.await; + release.notify_one(); + + tokio::time::timeout(Duration::from_secs(1), async { + while dropped.load(Ordering::SeqCst) == 0 { + tokio::task::yield_now().await; + } + }) + .await + .expect("detached owner output should be dropped after it finishes"); + assert_eq!(dropped.load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn planning_owner_deadline_drops_stalled_work_and_its_guards() { + let dropped = Arc::new(AtomicUsize::new(0)); + let task_dropped = Arc::clone(&dropped); + let planning = async move { + let _probe = DropProbe(task_dropped); + std::future::pending::<()>().await; + }; + + let result = + super::await_owned_planning_deadline(planning, Duration::from_millis(20)).await; + + assert!( + result.is_err(), + "stalled planning must hit its owner deadline" + ); + assert_eq!(dropped.load(Ordering::SeqCst), 1); + } +} diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/quota.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/quota.rs index 128c30069..b43fec3ba 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/quota.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/quota.rs @@ -1,7 +1,5 @@ //! Quota exhaustion, replay safety, and upstream replacement policy. -use axum::extract::ws::WebSocket; -use futures_util::SinkExt; use serde_json::Value; use uuid::Uuid; use wreq::ws::message::Message as WreqWsMessage; @@ -10,28 +8,21 @@ use super::adapter::{ resolve_responses_websocket_adapter, ResponsesWebSocketDrainDirective, ResponsesWebSocketRebindSafety, }; -use super::lifecycle::{queue_turn_finalization, ActiveProviderAttempt, PreviousAttemptSettled}; -use super::request::{ - build_planning_parts, planned_response_create_event, response_create_has_previous_response_id, +use super::lifecycle::{queue_turn_finalization, PreviousAttemptSettled}; +use super::ownership::{ + await_owned_responses_websocket_plan, begin_responses_websocket_turn_with_planned_lease, + spawn_owned_responses_websocket_plan, OwnedResponsesWebSocketDecision, }; +use super::request::{build_planning_parts, planned_response_create_event}; use super::state::BoundResponsesConnection; -use super::turn::{ - begin_responses_websocket_turn, prepare_responses_websocket_turn_decision, - ResponsesWebSocketTurnOutcome, -}; +use super::turn::{prepare_responses_websocket_turn_decision, ResponsesWebSocketTurnOutcome}; use super::upstream::{bind_responses_upstream, close_bound_upstream}; -use crate::ai_serving::maybe_build_responses_websocket_decision; use crate::clock::current_unix_secs; use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext; use crate::handlers::proxy::websocket::session::WEBSOCKET_LOG_TRANSPORT; -use crate::handlers::proxy::websocket::transport::{ - close_upstream_socket, send_responses_websocket_error, -}; -use crate::orchestration::release_pool_key_lease_from_report_context; +use crate::handlers::proxy::websocket::transport::close_upstream_socket; use crate::AppState; -const PREVIOUS_RESPONSE_NOT_FOUND_MESSAGE: &str = - "Previous response was not found. Retrying the full request."; const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws"; macro_rules! debug { @@ -132,6 +123,17 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( active.retry_attempted = true; active.turn_attempt = active.turn_attempt.saturating_add(1); let client_event = active.client_event.clone(); + let Some(turn_control) = active.turn_control.clone() else { + warn!( + event_name = "responses_websocket_quota_retry_control_missing", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + "gateway refused to retry a WebSocket turn without its live authorization snapshot" + ); + return false; + }; let turn_index = active.turn_index; let logical_turn_id = active.logical_turn_id.clone(); let turn_attempt = active.turn_attempt; @@ -147,18 +149,20 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( let now_unix_secs = current_unix_secs(); let excluded_key_ids = bound.exhausted_exclusions.key_ids(now_unix_secs); let excluded_codex_account_ids = bound.exhausted_exclusions.codex_account_ids(now_unix_secs); - let excluded_key_ids = (!excluded_key_ids.is_empty()).then_some(&excluded_key_ids); + let excluded_key_ids = (!excluded_key_ids.is_empty()).then_some(excluded_key_ids); let excluded_codex_account_ids = - (!excluded_codex_account_ids.is_empty()).then_some(&excluded_codex_account_ids); - let planned = match maybe_build_responses_websocket_decision( - state, - &planning_parts, - &turn_request_id, - &context.decision, - &client_event, + (!excluded_codex_account_ids.is_empty()).then_some(excluded_codex_account_ids); + let planned = match await_owned_responses_websocket_plan(spawn_owned_responses_websocket_plan( + state.clone(), + planning_parts, + turn_request_id.clone(), + turn_control.decision.clone(), + turn_control.auth_snapshot.clone(), + client_event.clone(), excluded_key_ids, excluded_codex_account_ids, - ) + None, + )) .await { Ok(Some(decision)) => decision, @@ -188,11 +192,16 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( return false; } }; + let OwnedResponsesWebSocketDecision { + planned, + planning_parts, + planned_lease, + } = planned; let adapter = resolve_responses_websocket_adapter(planned.adapter); let normalization = planned.normalization; let decision = planned.execution; if exhausted_key_id.as_deref() == decision.key_id.as_deref() { - release_pool_key_lease_from_report_context(state, decision.report_context.as_ref()).await; + planned_lease.release().await; warn!( event_name = "responses_websocket_quota_retry_selected_exhausted_key", log_type = "ops", @@ -212,8 +221,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( ) { Ok(event) => event, Err(code) => { - release_pool_key_lease_from_report_context(state, decision.report_context.as_ref()) - .await; + planned_lease.release().await; warn!( event_name = "responses_websocket_quota_retry_normalization_failed", log_type = "ops", @@ -237,12 +245,14 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( &logical_turn_id, turn_attempt, ); - let mut turn = match begin_responses_websocket_turn( + let mut turn = match begin_responses_websocket_turn_with_planned_lease( state, - &planning_parts, - &context.decision, + &context.trace_id, + planning_parts, + &turn_control.decision, turn_decision, &client_event, + planned_lease, ) .await { @@ -309,10 +319,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( // 同一个 logical turn 的下一个 attempt 就位。状态不符时把 attempt 交回 // drop guard 结算并让调用方走「透明重试失败」分支,不静默丢弃一条已经写了 // pending usage 行、占着 candidate 和 pool key lease 的 attempt。 - if let Err(orphan) = bound - .turn_state - .resume(ActiveProviderAttempt::new(state, turn)) - { + if let Err(orphan) = bound.turn_state.resume(turn) { drop(orphan); return false; } @@ -334,15 +341,6 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion( true } -pub(super) fn active_continuation_can_retry_from_full_input( - bound: &BoundResponsesConnection, -) -> bool { - bound.turn_state.logical().is_some_and(|active| { - response_create_has_previous_response_id(&active.client_event) - && active.retry_unsafe_reason.is_none() - }) -} - pub(super) fn is_usage_limit_error_event(event: &Value) -> bool { let is_error = |value: &Value| { value.get("type").and_then(Value::as_str) == Some("error") @@ -355,27 +353,6 @@ pub(super) fn is_usage_limit_error_event(event: &Value) -> bool { .is_some_and(|chunks| chunks.iter().any(is_error)) } -pub(super) fn should_request_full_continuation_retry( - bound: &BoundResponsesConnection, - retry_current_turn: bool, - upstream_event: Option<&Value>, -) -> bool { - retry_current_turn - && active_continuation_can_retry_from_full_input(bound) - && upstream_event.is_some_and(is_usage_limit_error_event) -} - -pub(super) async fn send_previous_response_not_found(client_socket: &mut WebSocket) { - send_responses_websocket_error( - client_socket, - 400, - "invalid_request_error", - "previous_response_not_found", - PREVIOUS_RESPONSE_NOT_FOUND_MESSAGE, - ) - .await; -} - pub(super) fn observe_active_response_rebind_safety( bound: &mut BoundResponsesConnection, event: &Value, diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/redaction.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/redaction.rs index 620e661e4..7a698dfad 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/redaction.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/redaction.rs @@ -305,8 +305,8 @@ mod tests { remote_addr: "127.0.0.1:65000" .parse::() .expect("remote address should parse"), + client_ip: "127.0.0.1".parse().expect("client IP should parse"), decision, - rpm_bypassed: false, websocket_connection_permit: None, } } diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/relay_policy.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/relay_policy.rs index 63c3534eb..99b974b4a 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/relay_policy.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/relay_policy.rs @@ -81,15 +81,12 @@ pub struct QuotaRelayFacts { /// The adapter allows a transparent replay of this turn. pub retry_current_turn: bool, /// The session already attempted the adapter-approved transparent replay - /// and could not bind an alternate upstream. A continuation may request - /// complete input only after that first recovery path was exhausted. + /// and could not bind an alternate upstream. This prevents retry loops and + /// causes the original upstream quota event to be relayed to the client. pub transparent_retry_failed: bool, /// The event contains the definitive `usage_limit_reached` error. A /// merely exhausted-looking rate-limit snapshot must not trigger retry. pub usage_limit_error: bool, - /// The active request is a continuation that can be retried from complete - /// input after the old account is detached. - pub continuation_retry_eligible: bool, pub upstream_closed: bool, } @@ -97,7 +94,6 @@ pub struct QuotaRelayFacts { pub enum QuotaRelayAction { None, AttemptTransparentRetry, - RequestFullContinuationRetry, ForwardQuotaAndDetach, } @@ -113,13 +109,7 @@ pub const fn classify_quota_relay(facts: QuotaRelayFacts) -> QuotaRelayAction { if facts.usage_limit_error && facts.retry_current_turn && !facts.transparent_retry_failed { return QuotaRelayAction::AttemptTransparentRetry; } - if facts.usage_limit_error - && facts.continuation_retry_eligible - && (facts.transparent_retry_failed || !facts.retry_current_turn) - { - return QuotaRelayAction::RequestFullContinuationRetry; - } - if facts.upstream_closed { + if facts.usage_limit_error || facts.upstream_closed { return QuotaRelayAction::ForwardQuotaAndDetach; } QuotaRelayAction::None @@ -177,7 +167,6 @@ mod tests { retry_current_turn: true, transparent_retry_failed: false, usage_limit_error: true, - continuation_retry_eligible: false, upstream_closed: false, }); assert_eq!(first, QuotaRelayAction::AttemptTransparentRetry); @@ -190,50 +179,22 @@ mod tests { retry_current_turn: false, transparent_retry_failed: true, usage_limit_error: true, - continuation_retry_eligible: false, upstream_closed: true, }); assert_eq!(after_retry_failure, QuotaRelayAction::ForwardQuotaAndDetach); } #[test] - fn continuation_quota_can_request_full_input_retry_without_replaying_partial_state() { + fn quota_error_is_forwarded_after_transparent_retry_is_unavailable() { assert_eq!( classify_quota_relay(QuotaRelayFacts { drain_ready: true, retry_current_turn: false, transparent_retry_failed: false, usage_limit_error: true, - continuation_retry_eligible: true, - upstream_closed: true, - }), - QuotaRelayAction::RequestFullContinuationRetry - ); - assert_eq!( - classify_quota_relay(QuotaRelayFacts { - drain_ready: true, - retry_current_turn: false, - transparent_retry_failed: true, - usage_limit_error: true, - continuation_retry_eligible: true, - upstream_closed: true, - }), - QuotaRelayAction::RequestFullContinuationRetry - ); - } - - #[test] - fn continuation_quota_without_transparent_retry_support_uses_full_input_retry() { - assert_eq!( - classify_quota_relay(QuotaRelayFacts { - drain_ready: true, - retry_current_turn: false, - transparent_retry_failed: false, - usage_limit_error: true, - continuation_retry_eligible: true, upstream_closed: false, }), - QuotaRelayAction::RequestFullContinuationRetry + QuotaRelayAction::ForwardQuotaAndDetach ); } @@ -277,7 +238,6 @@ mod tests { retry_current_turn: true, transparent_retry_failed: false, usage_limit_error: false, - continuation_retry_eligible: false, upstream_closed: false, }), QuotaRelayAction::None @@ -289,7 +249,6 @@ mod tests { retry_current_turn: false, transparent_retry_failed: true, usage_limit_error: false, - continuation_retry_eligible: false, upstream_closed: true, }), QuotaRelayAction::ForwardQuotaAndDetach diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/request.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/request.rs index d4606d45b..33d187221 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/request.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/request.rs @@ -13,6 +13,25 @@ use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext; use crate::headers::request_origin_from_headers_and_remote_addr; use crate::privacy::RedactionSessionSlot; +/// Model identifiers are copied into planner diagnostics. Bound them before +/// any planning/logging so a single 16 MiB WebSocket frame cannot amplify into +/// repeated multi-megabyte log records. +pub(super) const MAX_RESPONSES_WEBSOCKET_MODEL_BYTES: usize = 256; + +pub(super) fn validated_response_create_model(value: &Value) -> Result<&str, &'static str> { + let Some(model) = value + .as_str() + .map(str::trim) + .filter(|model| !model.is_empty()) + else { + return Err("invalid_response_create_model"); + }; + if model.len() > MAX_RESPONSES_WEBSOCKET_MODEL_BYTES { + return Err("invalid_response_create_model"); + } + Ok(model) +} + /// 把一条 WebSocket turn 还原成 planner 需要的 HTTP 形状请求头部。 /// /// 这里必须和 HTTP 前门(`handlers/proxy/mod.rs`)保持同一份 extension 契约: @@ -73,11 +92,12 @@ pub(super) fn planned_response_create_event( /// Restores the WebSocket protocol framing that provider-body normalization is /// not aware of. /// -/// `previous_response_id` is on the Codex unsupported-field list and `generate` -/// is not an HTTP body option at all, so normalization strips both — yet they -/// are the entire point of WebSocket mode. They must be re-grafted from the -/// client event afterwards. `stream`/`background` go the other way: the -/// normalizer inserts `stream`, and the WebSocket protocol has no use for it. +/// `previous_response_id` is on the Codex unsupported-field list, Codex HTTP +/// normalization may force `store`, and `generate` is not an HTTP body option +/// at all. Those fields are WebSocket protocol state, so an explicitly supplied +/// value (including `null`) must be re-grafted verbatim from the client event. +/// `stream`/`background` go the other way: the normalizer inserts `stream`, and +/// the WebSocket protocol has no use for it. fn finish_response_create_event( mut event: Value, client_event: &Value, @@ -89,13 +109,9 @@ fn finish_response_create_event( "type".to_string(), Value::String("response.create".to_string()), ); - for field in ["previous_response_id", "generate"] { + for field in ["store", "previous_response_id", "generate"] { if let Some(value) = client_event.get(field) { - if value.is_null() { - object.remove(field); - } else { - object.insert(field.to_string(), value.clone()); - } + object.insert(field.to_string(), value.clone()); } } object.remove("stream"); @@ -109,13 +125,6 @@ pub(super) fn response_create_has_previous_response_id(event: &Value) -> bool { .is_some_and(|value| !value.is_null()) } -pub(super) fn continuation_requires_same_upstream( - event: &Value, - reuses_bound_upstream: bool, -) -> bool { - response_create_has_previous_response_id(event) && !reuses_bound_upstream -} - pub(super) fn changed_followup_response_create_model( event: &Value, current_client_model: &str, @@ -126,13 +135,7 @@ pub(super) fn changed_followup_response_create_model( let Some(model) = object.get("model") else { return Ok(None); }; - let Some(model) = model - .as_str() - .map(str::trim) - .filter(|model| !model.is_empty()) - else { - return Err("invalid_response_create_model"); - }; + let model = validated_response_create_model(model)?; if model.eq_ignore_ascii_case(current_client_model) { Ok(None) } else { @@ -154,14 +157,10 @@ pub(super) fn response_create_model_or_current( ); return Ok(current_client_model.to_string()); }; - let Some(model) = model - .as_str() - .map(str::trim) - .filter(|model| !model.is_empty()) - else { - return Err("invalid_response_create_model"); - }; - Ok(model.to_string()) + let model = validated_response_create_model(model)?; + let model = model.to_string(); + object.insert("model".to_string(), Value::String(model.clone())); + Ok(model) } pub(super) fn provider_model_from_decision(decision: &AiExecutionDecision) -> Option { @@ -217,7 +216,7 @@ mod tests { use std::net::SocketAddr; use axum::http::{HeaderMap, Uri}; - use serde_json::json; + use serde_json::{json, Value}; use super::{ build_planning_parts, normalize_followup_response_create, @@ -236,6 +235,7 @@ mod tests { remote_addr: "127.0.0.1:65001" .parse::() .expect("remote address should parse"), + client_ip: "127.0.0.1".parse().expect("client IP should parse"), decision: GatewayControlDecision::synthetic( "/v1/responses".to_string(), Some("ai_public".to_string()), @@ -243,7 +243,6 @@ mod tests { Some("responses_websocket".to_string()), Some("openai:responses".to_string()), ), - rpm_bypassed: false, websocket_connection_permit: None, } } @@ -311,6 +310,56 @@ mod tests { assert!(normalized.get("background").is_none()); } + #[test] + fn explicit_store_and_previous_response_id_are_forwarded_opaquely() { + let event = json!({ + "type": "response.create", + "model": "public-model", + "store": true, + "previous_response_id": {"future": "opaque"}, + "input": [], + }); + + let normalized = normalized_continuation( + &event, + &ResponsesWebSocketBodyNormalization::for_tests("provider-model") + .with_provider_type_for_tests("codex"), + ); + + // Codex HTTP normalization normally forces `store: false` and removes + // `previous_response_id`. WebSocket framing restores exactly what the + // client sent so the upstream owns validation and continuation lookup. + assert_eq!(normalized["store"], true); + assert_eq!( + normalized["previous_response_id"], + json!({"future": "opaque"}) + ); + } + + #[test] + fn explicit_null_websocket_protocol_state_is_not_rewritten() { + let event = json!({ + "type": "response.create", + "model": "public-model", + "store": null, + "previous_response_id": null, + "generate": null, + "input": [], + }); + + let normalized = normalized_continuation( + &event, + &ResponsesWebSocketBodyNormalization::for_tests("provider-model") + .with_provider_type_for_tests("codex"), + ); + + assert!(normalized.get("store").is_some_and(Value::is_null)); + assert!(normalized + .get("previous_response_id") + .is_some_and(Value::is_null)); + assert!(normalized.get("generate").is_some_and(Value::is_null)); + } + #[test] fn continuation_strips_fields_the_codex_backend_rejects() { // The point of the fix: before it, turns 2..N reached Codex with the diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs index ddfa86796..bc6f416f5 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs @@ -8,7 +8,6 @@ //! request may replace it, but a continuation must stay on the original //! connection and account. -use axum::body::Bytes; use axum::extract::ws::{Message as AxumWsMessage, WebSocket}; use futures_util::{SinkExt, StreamExt}; use serde_json::Value; @@ -16,23 +15,27 @@ use uuid::Uuid; use super::adapter::resolve_responses_websocket_adapter; use super::client::consume_response_create_rate_limit; -use super::connection::relay_bound_connection; +use super::connection::{relay_bound_connection, wait_for_connection_permit_loss}; +use super::control::resolve_responses_websocket_turn_control; use super::lifecycle::{ await_pending_adapter_observation, await_pending_turn_finalization, - await_turn_finalization_handle, finalize_unbound_turn, responses_websocket_turn_start_close, - send_responses_websocket_turn_start_error, ActiveProviderAttempt, + await_turn_finalization_handle, finalize_active_turn, finalize_unbound_turn, + responses_websocket_turn_start_close, send_responses_websocket_turn_start_error, +}; +use super::ownership::{ + await_owned_responses_websocket_plan, begin_responses_websocket_turn_with_planned_lease, + spawn_owned_responses_websocket_plan, OwnedResponsesWebSocketDecision, }; use super::redaction::redact_responses_websocket_client_event; -use super::request::{build_planning_parts, planned_response_create_event}; -use super::turn::{ - begin_responses_websocket_turn, prepare_responses_websocket_turn_decision, - ResponsesWebSocketTurnOutcome, +use super::relay_policy::{fatal_relay_policy, FatalRelaySignal}; +use super::request::{ + build_planning_parts, planned_response_create_event, validated_response_create_model, }; +use super::state::BoundResponsesConnection; +use super::turn::{prepare_responses_websocket_turn_decision, ResponsesWebSocketTurnOutcome}; use super::turn_state::LogicalTurn; -use super::upstream::bind_responses_upstream; +use super::upstream::{bind_responses_upstream, close_bound_upstream}; -use crate::ai_serving::maybe_build_responses_websocket_decision; -use crate::control::request_model_local_rejection; use crate::handlers::proxy::websocket::ingress::{ WebSocketConnectionLog, WebSocketConnectionLogSpec, WebSocketRequestContext, }; @@ -43,7 +46,6 @@ use crate::handlers::proxy::websocket::session::{ use crate::handlers::proxy::websocket::transport::{ close_client_socket, send_gateway_error, send_gateway_error_with_status, }; -use crate::orchestration::release_pool_key_lease_from_report_context; use crate::AppState; const RESPONSES_WEBSOCKET_LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws"; @@ -71,6 +73,13 @@ enum InitialMessageError { InvalidJson, MissingResponseCreate, MissingModel, + InvalidModel, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum ConnectionTermination { + ConnectionLimitReached, + ConnectionAdmissionLost, } impl InitialMessageError { @@ -83,6 +92,7 @@ impl InitialMessageError { Self::InvalidJson => "invalid_response_create", Self::MissingResponseCreate => "expected_response_create", Self::MissingModel => "response_create_model_required", + Self::InvalidModel => "invalid_response_create_model", } } @@ -91,7 +101,9 @@ impl InitialMessageError { Self::TimedOut => CLOSE_TRY_AGAIN, Self::ClientClosed => 1000, Self::ClientRead | Self::UnsupportedFrame | Self::InvalidJson => CLOSE_POLICY_VIOLATION, - Self::MissingResponseCreate | Self::MissingModel => CLOSE_POLICY_VIOLATION, + Self::MissingResponseCreate | Self::MissingModel | Self::InvalidModel => { + CLOSE_POLICY_VIOLATION + } } } } @@ -101,46 +113,151 @@ pub(super) async fn run_responses_websocket( state: AppState, mut context: WebSocketRequestContext, ) { + // The public connection limit starts when the HTTP Upgrade hands us the + // socket, not after provider planning and its upstream handshake finish. + let connection_deadline = + tokio::time::Instant::now() + RESPONSES_WEBSOCKET_SESSION_LIMITS.max_connection_duration; let connection_permit = context.websocket_connection_permit.take(); let connection_log = WebSocketConnectionLog::new(&context, RESPONSES_CONNECTION_LOG_SPEC); connection_log.log_opened(); - let (first_text, first_event) = match receive_initial_response_create(&mut client_socket).await - { - Ok(value) => value, - Err(error) => { - if !matches!(error, InitialMessageError::ClientClosed) { - send_gateway_error( - &mut client_socket, - error.code(), - "WebSocket must start with a valid response.create event", - ) - .await; - close_client_socket( - &mut client_socket, - error.close_code(), - "invalid_initial_event", - ) - .await; - } + let bootstrap_result = supervise_responses_websocket_phase( + bootstrap_responses_websocket(&mut client_socket, state.clone(), &context), + connection_deadline, + connection_permit.as_ref(), + ) + .await; + + let mut bound = match bootstrap_result { + Ok(Some(bound)) => bound, + Ok(None) => return, + Err(termination) => { + // The bootstrap future (and therefore every lease/attempt guard it + // owned) has been dropped before the permit or client socket is + // touched here. + drop(connection_permit); + close_terminated_bootstrap(&mut client_socket, &context, termination).await; return; } }; - let planning_parts = build_planning_parts(&context); - match consume_response_create_rate_limit(&state, &context.decision, context.rpm_bypassed).await + let relay_result = supervise_responses_websocket_phase( + relay_bound_connection(&mut client_socket, &mut bound, &state, &context), + connection_deadline, + connection_permit.as_ref(), + ) + .await; + + if let Err(termination) = relay_result { + // The relay future has been dropped before cleanup starts. This makes + // cancellation safe even when a client-frame branch was awaiting a + // live auth refresh, redaction, planning, admission, or rebind. + let outcome = match termination { + ConnectionTermination::ConnectionLimitReached => { + ResponsesWebSocketTurnOutcome::connection_limit_reached() + } + ConnectionTermination::ConnectionAdmissionLost => { + ResponsesWebSocketTurnOutcome::connection_admission_lost() + } + }; + finalize_active_turn(&mut bound, &state, outcome).await; + close_bound_upstream(&mut bound).await; + drop(connection_permit); + close_terminated_relay(&mut client_socket, &context, termination).await; + } else { + drop(connection_permit); + } + await_pending_turn_finalization(&mut bound).await; + await_pending_adapter_observation(&mut bound).await; +} + +async fn supervise_responses_websocket_phase( + phase: F, + connection_deadline: tokio::time::Instant, + connection_permit: Option<&aether_runtime::AdmissionPermit>, +) -> Result +where + F: std::future::Future, +{ + tokio::pin!(phase); + tokio::select! { + biased; + _ = tokio::time::sleep_until(connection_deadline) => { + Err(ConnectionTermination::ConnectionLimitReached) + } + _ = wait_for_connection_permit_loss(connection_permit) => { + Err(ConnectionTermination::ConnectionAdmissionLost) + } + output = &mut phase => Ok(output), + } +} + +async fn bootstrap_responses_websocket( + client_socket: &mut WebSocket, + state: AppState, + context: &WebSocketRequestContext, +) -> Option { + let (_first_text, first_event) = match receive_initial_response_create(client_socket).await { + Ok(value) => value, + Err(error) => { + if !matches!(error, InitialMessageError::ClientClosed) { + send_gateway_error( + client_socket, + error.code(), + "WebSocket must start with a valid response.create event", + ) + .await; + close_client_socket(client_socket, error.close_code(), "invalid_initial_event") + .await; + } + return None; + } + }; + + let planning_parts = build_planning_parts(context); + let turn_control = match resolve_responses_websocket_turn_control( + &state, + context, + &planning_parts, + &first_event, + ) + .await + { + Ok(control) => control, + Err(error) => { + warn!( + event_name = "responses_websocket_initial_turn_control_rejected", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + error = ?error, + "gateway rejected a Responses WebSocket initial turn after live policy refresh" + ); + send_responses_websocket_turn_start_error(client_socket, &error).await; + let (close_code, close_reason) = responses_websocket_turn_start_close(&error); + close_client_socket(client_socket, close_code, close_reason).await; + return None; + } + }; + match consume_response_create_rate_limit( + &state, + &turn_control.decision, + turn_control.rpm_bypassed, + ) + .await { Ok(true) => {} Ok(false) => { send_gateway_error_with_status( - &mut client_socket, + client_socket, 429, "rate_limit_exceeded", "Too many response.create events; retry later", ) .await; - close_client_socket(&mut client_socket, CLOSE_TRY_AGAIN, "rate_limit_exceeded").await; - return; + close_client_socket(client_socket, CLOSE_TRY_AGAIN, "rate_limit_exceeded").await; + return None; } Err(()) => { warn!( @@ -152,68 +269,19 @@ pub(super) async fn run_responses_websocket( "gateway failed to consume WebSocket response rate limit" ); send_gateway_error_with_status( - &mut client_socket, + client_socket, 503, "gateway_rate_limit_unavailable", "Gateway could not evaluate the response rate limit", ) .await; close_client_socket( - &mut client_socket, + client_socket, CLOSE_INTERNAL_ERROR, "rate_limit_unavailable", ) .await; - return; - } - } - match request_model_local_rejection( - &state, - Some(&context.decision), - &planning_parts.uri, - &planning_parts.headers, - &Bytes::from(first_text.into_bytes()), - ) - .await - { - Ok(Some(_)) => { - send_gateway_error( - &mut client_socket, - "model_not_allowed", - "The requested model is not available to this API key", - ) - .await; - close_client_socket( - &mut client_socket, - CLOSE_POLICY_VIOLATION, - "model_not_allowed", - ) - .await; - return; - } - Ok(None) => {} - Err(_) => { - warn!( - event_name = "responses_websocket_model_access_check_failed", - log_type = "ops", - transport = WEBSOCKET_LOG_TRANSPORT, - websocket = true, - trace_id = %context.trace_id, - "gateway failed to evaluate WebSocket model access policy" - ); - send_gateway_error( - &mut client_socket, - "gateway_auth_unavailable", - "Gateway could not evaluate request access", - ) - .await; - close_client_socket( - &mut client_socket, - CLOSE_INTERNAL_ERROR, - "gateway_auth_unavailable", - ) - .await; - return; + return None; } } @@ -223,7 +291,7 @@ pub(super) async fn run_responses_websocket( let redacted_first_event = redact_responses_websocket_client_event( &state, &planning_parts, - &context.decision, + &turn_control.decision, &first_event, ) .await; @@ -243,49 +311,51 @@ pub(super) async fn run_responses_websocket( "gateway could not apply chat PII redaction to the initial Responses WebSocket event" ); send_gateway_error_with_status( - &mut client_socket, + client_socket, 500, "responses_websocket_redaction_unavailable", "Gateway could not apply the configured PII redaction", ) .await; close_client_socket( - &mut client_socket, + client_socket, CLOSE_INTERNAL_ERROR, "responses_websocket_redaction_unavailable", ) .await; - return; + return None; } }; - let planned = match maybe_build_responses_websocket_decision( - &state, - &planning_parts, - &context.trace_id, - &context.decision, - &first_event, + let planned = match await_owned_responses_websocket_plan(spawn_owned_responses_websocket_plan( + state.clone(), + planning_parts, + context.trace_id.clone(), + turn_control.decision.clone(), + turn_control.auth_snapshot.clone(), + first_event.clone(), None, None, - ) + None, + )) .await { Ok(Some(decision)) => decision, Ok(None) => { send_gateway_error_with_status( - &mut client_socket, + client_socket, 503, "responses_provider_unavailable", "No eligible WebSocket-enabled Responses provider is available", ) .await; close_client_socket( - &mut client_socket, + client_socket, CLOSE_TRY_AGAIN, "responses_provider_unavailable", ) .await; - return; + return None; } Err(_) => { warn!( @@ -297,22 +367,27 @@ pub(super) async fn run_responses_websocket( "gateway failed to plan Responses WebSocket provider request" ); send_gateway_error_with_status( - &mut client_socket, + client_socket, 503, "responses_provider_unavailable", "Gateway could not prepare a Provider connection", ) .await; close_client_socket( - &mut client_socket, + client_socket, CLOSE_INTERNAL_ERROR, "responses_planning_failed", ) .await; - return; + return None; } }; + let OwnedResponsesWebSocketDecision { + planned, + planning_parts, + planned_lease, + } = planned; let adapter = resolve_responses_websocket_adapter(planned.adapter); let normalization = planned.normalization; let decision = planned.execution; @@ -322,8 +397,7 @@ pub(super) async fn run_responses_websocket( }) { Ok(event) => event, Err(code) => { - release_pool_key_lease_from_report_context(&state, decision.report_context.as_ref()) - .await; + planned_lease.release().await; warn!( event_name = "responses_websocket_initial_event_normalization_failed", log_type = "ops", @@ -334,13 +408,13 @@ pub(super) async fn run_responses_websocket( "gateway could not normalize the initial Responses WebSocket event" ); send_gateway_error( - &mut client_socket, + client_socket, code, "Gateway could not prepare the Responses response.create event", ) .await; - close_client_socket(&mut client_socket, CLOSE_POLICY_VIOLATION, code).await; - return; + close_client_socket(client_socket, CLOSE_POLICY_VIOLATION, code).await; + return None; } }; let first_logical_turn_id = Uuid::new_v4().to_string(); @@ -355,12 +429,14 @@ pub(super) async fn run_responses_websocket( &first_logical_turn_id, 1, ); - let mut first_turn = match begin_responses_websocket_turn( + let mut first_turn = match begin_responses_websocket_turn_with_planned_lease( &state, - &planning_parts, - &context.decision, + &context.trace_id, + planning_parts, + &turn_control.decision, first_turn_decision, &first_event, + planned_lease, ) .await { @@ -375,10 +451,10 @@ pub(super) async fn run_responses_websocket( error = ?error, "gateway could not start Responses WebSocket usage/audit lifecycle" ); - send_responses_websocket_turn_start_error(&mut client_socket, &error).await; + send_responses_websocket_turn_start_error(client_socket, &error).await; let (close_code, close_reason) = responses_websocket_turn_start_close(&error); - close_client_socket(&mut client_socket, close_code, close_reason).await; - return; + close_client_socket(client_socket, close_code, close_reason).await; + return None; } }; @@ -390,8 +466,7 @@ pub(super) async fn run_responses_websocket( state.clone(), first_turn, ResponsesWebSocketTurnOutcome::upstream_connect_failed(code), - ) - .await; + ); warn!( event_name = "responses_websocket_upstream_connect_failed", log_type = "ops", @@ -402,15 +477,15 @@ pub(super) async fn run_responses_websocket( "gateway failed to establish Responses WebSocket upstream" ); send_gateway_error_with_status( - &mut client_socket, + client_socket, 502, code, "Gateway could not establish the Provider connection", ) .await; - close_client_socket(&mut client_socket, CLOSE_TRY_AGAIN, code).await; + close_client_socket(client_socket, CLOSE_TRY_AGAIN, code).await; await_turn_finalization_handle(finalizer).await; - return; + return None; } }; first_turn.mark_upstream_request_sent(); @@ -419,20 +494,103 @@ pub(super) async fn run_responses_websocket( bound.redaction_restorer.register(session); } bound.turn_state.begin( - LogicalTurn::new(first_event, 1, first_logical_turn_id), - ActiveProviderAttempt::new(&state, first_turn), + LogicalTurn::new(first_event, 1, first_logical_turn_id).with_turn_control(turn_control), + first_turn, ); - relay_bound_connection( - &mut client_socket, - &mut bound, - &state, - &context, - connection_permit, - ) - .await; - await_pending_turn_finalization(&mut bound).await; - await_pending_adapter_observation(&mut bound).await; + Some(bound) +} + +async fn close_terminated_bootstrap( + client_socket: &mut WebSocket, + context: &WebSocketRequestContext, + termination: ConnectionTermination, +) { + match termination { + ConnectionTermination::ConnectionLimitReached => { + warn!( + event_name = "responses_websocket_bootstrap_connection_limit_reached", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + "gateway stopped a Responses WebSocket bootstrap at the absolute connection deadline" + ); + send_gateway_error_with_status( + client_socket, + 503, + "websocket_connection_limit_reached", + "WebSocket connection duration limit reached; reconnect to continue", + ) + .await; + close_client_socket(client_socket, CLOSE_TRY_AGAIN, "connection_limit_reached").await; + } + ConnectionTermination::ConnectionAdmissionLost => { + let policy = fatal_relay_policy(FatalRelaySignal::ConnectionAdmissionLost); + warn!( + event_name = "responses_websocket_bootstrap_connection_admission_lost", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + "gateway stopped a Responses WebSocket bootstrap after its connection admission became unhealthy" + ); + send_gateway_error_with_status( + client_socket, + policy.status_code, + policy.error_code, + policy.client_message, + ) + .await; + close_client_socket(client_socket, policy.close_code, policy.close_reason).await; + } + } +} + +async fn close_terminated_relay( + client_socket: &mut WebSocket, + context: &WebSocketRequestContext, + termination: ConnectionTermination, +) { + match termination { + ConnectionTermination::ConnectionLimitReached => { + warn!( + event_name = "responses_websocket_connection_limit_reached", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + "gateway stopped a Responses WebSocket relay at the absolute connection deadline" + ); + send_gateway_error_with_status( + client_socket, + 503, + "websocket_connection_limit_reached", + "WebSocket connection duration limit reached; reconnect to continue", + ) + .await; + close_client_socket(client_socket, CLOSE_TRY_AGAIN, "connection_limit_reached").await; + } + ConnectionTermination::ConnectionAdmissionLost => { + let policy = fatal_relay_policy(FatalRelaySignal::ConnectionAdmissionLost); + warn!( + event_name = "responses_websocket_connection_admission_lost", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + "gateway stopped a Responses WebSocket relay after its connection admission became unhealthy" + ); + send_gateway_error_with_status( + client_socket, + policy.status_code, + policy.error_code, + policy.client_message, + ) + .await; + close_client_socket(client_socket, policy.close_code, policy.close_reason).await; + } + } } /// 等待客户端发送第一条 response.create 事件。 @@ -477,9 +635,9 @@ where let message = message.map_err(|_| InitialMessageError::ClientRead)?; match message { AxumWsMessage::Ping(payload) => { - socket - .send(AxumWsMessage::Pong(payload)) + tokio::time::timeout_at(deadline, socket.send(AxumWsMessage::Pong(payload))) .await + .map_err(|_| InitialMessageError::TimedOut)? .map_err(|_| InitialMessageError::ClientRead)?; } AxumWsMessage::Pong(_) => {} @@ -501,21 +659,17 @@ fn validate_initial_response_create(event: &Value) -> Result<(), InitialMessageE if object.get("type").and_then(Value::as_str) != Some("response.create") { return Err(InitialMessageError::MissingResponseCreate); } - if object + let model = object .get("model") - .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .is_none() - { - return Err(InitialMessageError::MissingModel); - } + .ok_or(InitialMessageError::MissingModel)?; + validated_response_create_model(model).map_err(|_| InitialMessageError::InvalidModel)?; Ok(()) } #[cfg(test)] mod tests { use std::collections::BTreeMap; + use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use std::sync::Arc; use std::time::{Duration, Instant}; @@ -525,15 +679,13 @@ mod tests { use super::super::binding::UpstreamBindingIdentity; use super::super::client::adapter_drain_ready; use super::super::quota::{ - active_continuation_can_retry_from_full_input, is_usage_limit_error_event, - observe_active_response_rebind_safety, record_exhausted_bound_key, - should_request_full_continuation_retry, + is_usage_limit_error_event, observe_active_response_rebind_safety, + record_exhausted_bound_key, }; use super::super::redaction::ResponsesWebSocketRedactionRestorer; use super::super::request::{ - changed_followup_response_create_model, continuation_requires_same_upstream, - normalize_followup_response_create, planned_response_create_event, - response_create_model_or_current, + changed_followup_response_create_model, normalize_followup_response_create, + planned_response_create_event, response_create_model_or_current, }; use super::super::state::{BoundResponsesConnection, ExhaustedResponsesWebSocketExclusions}; use super::super::turn::{ @@ -569,6 +721,82 @@ mod tests { event: serde_json::Value, } + struct BootstrapDropProbe(Arc); + + impl Drop for BootstrapDropProbe { + fn drop(&mut self) { + self.0.fetch_add(1, Ordering::SeqCst); + } + } + + struct TestAdmissionHealth(Arc); + + impl aether_runtime::AdmissionPermitHealth for TestAdmissionHealth { + fn is_healthy(&self) -> bool { + self.0.load(Ordering::Acquire) + } + } + + #[tokio::test] + async fn phase_supervisor_drops_work_at_the_absolute_connection_deadline() { + let dropped = Arc::new(AtomicUsize::new(0)); + let task_dropped = Arc::clone(&dropped); + let bootstrap = async move { + let _probe = BootstrapDropProbe(task_dropped); + std::future::pending::<()>().await; + }; + + let result = tokio::time::timeout( + Duration::from_secs(1), + super::supervise_responses_websocket_phase( + bootstrap, + tokio::time::Instant::now() + Duration::from_millis(20), + None, + ), + ) + .await + .expect("bootstrap supervisor should honor its absolute deadline"); + + assert_eq!( + result, + Err(super::ConnectionTermination::ConnectionLimitReached) + ); + assert_eq!(dropped.load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn phase_supervisor_drops_work_when_connection_admission_is_lost() { + let health = Arc::new(AtomicBool::new(false)); + let permit = aether_runtime::AdmissionPermit::from_parts( + None, + Some(TestAdmissionHealth(Arc::clone(&health))), + ) + .expect("distributed test health should create a permit"); + let dropped = Arc::new(AtomicUsize::new(0)); + let task_dropped = Arc::clone(&dropped); + let bootstrap = async move { + let _probe = BootstrapDropProbe(task_dropped); + std::future::pending::<()>().await; + }; + + let result = tokio::time::timeout( + Duration::from_secs(1), + super::supervise_responses_websocket_phase( + bootstrap, + tokio::time::Instant::now() + Duration::from_secs(30), + Some(&permit), + ), + ) + .await + .expect("unhealthy connection admission should stop bootstrap"); + + assert_eq!( + result, + Err(super::ConnectionTermination::ConnectionAdmissionLost) + ); + assert_eq!(dropped.load(Ordering::SeqCst), 1); + } + #[test] fn adapter_drain_waits_for_an_active_turn_terminal_event() { let directive = Some(ResponsesWebSocketDrainDirective { @@ -719,45 +947,7 @@ mod tests { } #[test] - fn continuation_requires_the_existing_upstream_connection_and_account() { - let continuation = json!({ - "type": "response.create", - "previous_response_id": "resp-previous", - }); - - assert!(!continuation_requires_same_upstream(&continuation, true)); - assert!(continuation_requires_same_upstream(&continuation, false)); - assert!(!continuation_requires_same_upstream( - &json!({"type": "response.create"}), - false, - )); - } - - #[test] - fn quota_error_can_request_a_full_retry_only_before_public_response_state() { - let mut bound = sample_bound_for_rebind_safety(); - bound.turn_state = ResponsesTurnState::Replanning { - logical: LogicalTurn::new( - json!({ - "type": "response.create", - "previous_response_id": "resp-previous", - }), - 2, - "logical-turn".to_string(), - ), - }; - - assert!(active_continuation_can_retry_from_full_input(&bound)); - bound - .turn_state - .logical_mut() - .expect("active request") - .mark_retry_unsafe("standard_response_event"); - assert!(!active_continuation_can_retry_from_full_input(&bound)); - } - - #[test] - fn only_an_actual_usage_limit_error_requests_full_retry() { + fn only_an_actual_usage_limit_error_requests_transparent_retry() { assert!(is_usage_limit_error_event(&json!({ "type": "error", "error": {"type": "usage_limit_reached"}, @@ -773,46 +963,6 @@ mod tests { }))); } - #[test] - fn full_continuation_retry_does_not_consume_a_successful_terminal_event() { - let mut bound = sample_bound_for_rebind_safety(); - bound.turn_state = ResponsesTurnState::Replanning { - logical: LogicalTurn::new( - json!({ - "type": "response.create", - "previous_response_id": "resp-previous", - }), - 2, - "logical-turn".to_string(), - ), - }; - - assert!(should_request_full_continuation_retry( - &bound, - true, - Some(&json!({ - "type": "error", - "error": {"type": "usage_limit_reached"}, - })), - )); - assert!(!should_request_full_continuation_retry( - &bound, - true, - Some(&json!({ - "type": "response.completed", - "response": {"id": "resp-completed"}, - })), - )); - assert!(!should_request_full_continuation_retry( - &bound, - false, - Some(&json!({ - "type": "error", - "error": {"type": "usage_limit_reached"}, - })), - )); - } - #[test] fn followup_rewrites_the_provider_model_and_removes_http_stream_fields() { let event = json!({ @@ -856,6 +1006,41 @@ mod tests { ); } + #[test] + fn oversized_model_is_rejected_before_initial_or_followup_planning() { + let oversized = "m".repeat(257); + let initial = json!({ + "type": "response.create", + "model": oversized, + "input": [], + }); + + assert!(matches!( + super::validate_initial_response_create(&initial), + Err(super::InitialMessageError::InvalidModel) + )); + assert_eq!( + changed_followup_response_create_model(&initial, "current-model"), + Err("invalid_response_create_model") + ); + } + + #[test] + fn model_at_the_identifier_limit_remains_valid() { + let model = "m".repeat(256); + let event = json!({ + "type": "response.create", + "model": model, + "input": [], + }); + + assert!(super::validate_initial_response_create(&event).is_ok()); + assert_eq!( + changed_followup_response_create_model(&event, "current-model"), + Ok(Some(model)) + ); + } + #[test] fn followup_without_a_model_reuses_the_current_connection_model() { let event = json!({ @@ -1282,6 +1467,77 @@ mod tests { } } + /// Emits one Ping and then permanently backpressures every Pong write. + /// This models a peer that keeps the TCP connection open but never reads. + struct StalledPongSocket { + emitted_ping: bool, + } + + impl futures_util::Stream for StalledPongSocket { + type Item = Result; + + fn poll_next( + mut self: std::pin::Pin<&mut Self>, + _cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + if self.emitted_ping { + std::task::Poll::Pending + } else { + self.emitted_ping = true; + std::task::Poll::Ready(Some(Ok(axum::extract::ws::Message::Ping(vec![7].into())))) + } + } + } + + impl futures_util::Sink for StalledPongSocket { + type Error = axum::Error; + + fn poll_ready( + self: std::pin::Pin<&mut Self>, + _cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::task::Poll::Pending + } + + fn start_send( + self: std::pin::Pin<&mut Self>, + _item: axum::extract::ws::Message, + ) -> Result<(), Self::Error> { + unreachable!("a permanently backpressured sink is never ready") + } + + fn poll_flush( + self: std::pin::Pin<&mut Self>, + _cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::task::Poll::Pending + } + + fn poll_close( + self: std::pin::Pin<&mut Self>, + _cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::task::Poll::Pending + } + } + + #[tokio::test] + async fn initial_message_deadline_also_bounds_a_stalled_pong_write() { + use super::{receive_initial_response_create_with_deadline, InitialMessageError}; + + let mut socket = StalledPongSocket { + emitted_ping: false, + }; + let result = tokio::time::timeout( + Duration::from_millis(300), + receive_initial_response_create_with_deadline(&mut socket, Duration::from_millis(30)), + ) + .await + .expect("the initial-message deadline must cancel a stalled Pong write"); + + assert!(matches!(result, Err(InitialMessageError::TimedOut))); + } + /// 验证 receive_initial_response_create_with_deadline 的绝对 deadline: /// 客户端周期性发送 Ping 帧不会重置计时器,deadline 到期后返回 TimedOut。 /// 这直接驱动真实的循环逻辑,如果改回每次迭代 timeout(budget, ...) 则会变红。 diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn.rs index 7cb4f0135..79f9bc2f3 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn.rs @@ -38,8 +38,7 @@ use super::settlement::attempt_facts_for_outcome; use crate::ai_serving::{build_openai_responses_stream_plan_from_decision, AiExecutionDecision}; use crate::clock::current_unix_ms; use crate::control::{ - execution_plan_balance_capacity_rejection, refresh_execution_runtime_auth_context, - request_model_local_rejection, GatewayControlDecision, GatewayLocalAuthRejection, + execution_plan_balance_capacity_rejection, GatewayControlDecision, GatewayLocalAuthRejection, }; use crate::execution_runtime::attach_provider_response_headers_to_report_context; use crate::execution_runtime::attempt_lifecycle::{ @@ -208,6 +207,20 @@ impl ResponsesWebSocketTurnOutcome { } } + pub(super) const fn relay_task_abandoned_before_upstream_send() -> Self { + Self::Cancelled { + reason: "gateway abandoned the turn before sending response.create upstream", + } + } + + pub(super) const fn relay_task_abandonment(upstream_request_sent: bool) -> Self { + if upstream_request_sent { + Self::relay_task_abandoned() + } else { + Self::relay_task_abandoned_before_upstream_send() + } + } + const fn status_code(self) -> u16 { match self { Self::ProviderTerminal { status_code, .. } | Self::Failure { status_code, .. } => { @@ -226,6 +239,8 @@ pub(super) struct ResponsesProviderAttempt { /// 记账三段(pending / started / terminal)由共享的 transport 中立生命周期负责。 lifecycle: ExecutionAttemptLifecycle, started_at: Instant, + provider_request_started_at_unix_ms: u64, + provider_request_order_id: String, provider_headers: BTreeMap, observer: ResponsesStructuredTerminalObserver, provider_capture: AttemptBodyCapture, @@ -241,6 +256,10 @@ pub(super) struct ResponsesProviderAttempt { provider_outcome: Option, /// 这一个 attempt 的内容是否完整交付给了客户端。与 provider 终态正交。 client_delivery: AttemptClientDelivery, + /// True only after `response.create` has been accepted by the upstream + /// socket writer. Cancellation before this point must not be projected as + /// provider failure or billed usage. + upstream_request_sent: bool, } /// 组装一轮 turn 的 decision。 @@ -282,7 +301,7 @@ pub(super) fn prepare_responses_websocket_turn_decision( decision } -pub(super) async fn begin_responses_websocket_turn( +pub(super) async fn begin_unowned_responses_websocket_turn( state: &AppState, parts: &http::request::Parts, control_decision: &GatewayControlDecision, @@ -290,17 +309,6 @@ pub(super) async fn begin_responses_websocket_turn( client_event: &Value, ) -> Result { let planned_report_context = decision.report_context.clone(); - let effective_control_decision = - match refresh_websocket_turn_auth_context(state, control_decision, parts, client_event) - .await - { - Ok(decision) => decision, - Err(error) => { - release_pool_key_lease_from_report_context(state, planned_report_context.as_ref()) - .await; - return Err(error); - } - }; let attempt = match build_openai_responses_stream_plan_from_decision( parts, client_event, @@ -344,7 +352,7 @@ pub(super) async fn begin_responses_websocket_turn( let balance_rejection = execution_plan_balance_capacity_rejection( state, - &effective_control_decision, + control_decision, &plan, report_context.as_ref(), ) @@ -375,6 +383,7 @@ pub(super) async fn begin_responses_websocket_turn( return Err(websocket_auth_rejection_error(rejection)); } + let candidate_started_at_unix_ms = current_unix_ms(); ensure_execution_request_candidate_slot(state, &mut plan, &mut report_context).await; let admission = match ResponsesWebSocketTurnAdmission::acquire( state, @@ -385,12 +394,21 @@ pub(super) async fn begin_responses_websocket_turn( { Ok(admission) => admission, Err(error) => { - release_local_pool_key_lease( - state, - LocalExecutionEffectContext { - plan: &plan, - report_context: report_context.as_ref(), - }, + release_then_record_responses_websocket_admission_failure( + release_local_pool_key_lease( + state, + LocalExecutionEffectContext { + plan: &plan, + report_context: report_context.as_ref(), + }, + ), + record_responses_websocket_admission_failure( + state, + &plan, + report_context.as_ref(), + candidate_started_at_unix_ms, + &error, + ), ) .await; return Err(error); @@ -413,6 +431,8 @@ pub(super) async fn begin_responses_websocket_turn( Ok(ResponsesProviderAttempt { lifecycle, started_at: Instant::now(), + provider_request_started_at_unix_ms: current_unix_ms(), + provider_request_order_id: uuid::Uuid::now_v7().to_string(), provider_headers: BTreeMap::new(), observer: ResponsesStructuredTerminalObserver::default(), provider_capture: AttemptBodyCapture::default(), @@ -425,40 +445,72 @@ pub(super) async fn begin_responses_websocket_turn( terminal_error_body: None, provider_outcome: None, client_delivery: AttemptClientDelivery::Complete, + upstream_request_sent: false, }) } -async fn refresh_websocket_turn_auth_context( - state: &AppState, - control_decision: &GatewayControlDecision, - parts: &http::request::Parts, - client_event: &Value, -) -> Result { - let mut effective = control_decision.clone(); - if let Some(auth_context) = effective.auth_context.take() { - let refreshed = refresh_execution_runtime_auth_context( - state, - auth_context, - effective.auth_endpoint_signature.as_deref(), - ) - .await?; - effective.local_auth_rejection = refreshed.local_rejection.clone(); - effective.auth_context = Some(refreshed); - } - if let Some(rejection) = effective.local_auth_rejection.clone() { - return Err(websocket_auth_rejection_error(rejection)); - } +async fn release_then_record_responses_websocket_admission_failure( + release_pool_lease: impl Future, + record_candidate_failure: impl Future, +) { + // Lease cleanup protects live routing capacity and must not sit behind a + // slow candidate writer. The candidate write still follows immediately so + // the seeded row reaches a terminal state on the ordinary error path. + release_pool_lease.await; + record_candidate_failure.await; +} - let body = serde_json::to_vec(client_event) - .map(axum::body::Bytes::from) - .map_err(|error| GatewayError::Internal(error.to_string()))?; - if let Some(rejection) = - request_model_local_rejection(state, Some(&effective), &parts.uri, &parts.headers, &body) - .await? - { - return Err(websocket_auth_rejection_error(rejection)); +async fn record_responses_websocket_admission_failure( + state: &AppState, + plan: &ExecutionPlan, + report_context: Option<&Value>, + candidate_started_at_unix_ms: u64, + error: &GatewayError, +) { + let terminal_at_unix_ms = current_unix_ms(); + record_local_request_candidate_status( + state, + plan, + report_context, + responses_websocket_admission_failure_update( + candidate_started_at_unix_ms, + terminal_at_unix_ms, + error, + ), + ) + .await; +} + +fn responses_websocket_admission_failure_update( + candidate_started_at_unix_ms: u64, + terminal_at_unix_ms: u64, + error: &GatewayError, +) -> SchedulerRequestCandidateStatusUpdate { + let (status_code, error_type, error_message) = match error { + GatewayError::AdmissionTimeout { + gate, + queue_budget_ms, + .. + } => ( + StatusCode::TOO_MANY_REQUESTS.as_u16(), + "gateway_admission_timeout", + format!("gateway admission gate {gate} timed out after {queue_budget_ms}ms"), + ), + other => ( + StatusCode::INTERNAL_SERVER_ERROR.as_u16(), + "gateway_admission_failed", + format!("{other:?}"), + ), + }; + SchedulerRequestCandidateStatusUpdate { + status: RequestCandidateStatus::Failed, + status_code: Some(status_code), + error_type: Some(error_type.to_string()), + error_message: Some(error_message), + latency_ms: Some(terminal_at_unix_ms.saturating_sub(candidate_started_at_unix_ms)), + started_at_unix_ms: Some(candidate_started_at_unix_ms), + finished_at_unix_ms: Some(terminal_at_unix_ms), } - Ok(effective) } fn websocket_auth_rejection_error(rejection: GatewayLocalAuthRejection) -> GatewayError { @@ -526,9 +578,13 @@ impl ResponsesProviderAttempt { } pub(super) fn set_provider_response_headers(&mut self, headers: BTreeMap) { + let observed_at_unix_ms = current_unix_ms(); let report_context = attach_provider_response_headers_to_report_context( self.lifecycle.take_report_context(), &headers, + self.provider_request_started_at_unix_ms, + observed_at_unix_ms, + &self.provider_request_order_id, ); self.lifecycle.set_report_context(report_context); self.provider_headers = headers; @@ -537,10 +593,21 @@ impl ResponsesProviderAttempt { /// Starts the per-turn response deadlines only after the corresponding /// `response.create` has been accepted by the upstream socket writer. pub(super) fn mark_upstream_request_sent(&mut self) { + self.upstream_request_sent = true; self.started_at = Instant::now(); + self.provider_request_started_at_unix_ms = current_unix_ms(); + self.provider_request_order_id = uuid::Uuid::now_v7().to_string(); self.first_event_elapsed_ms = None; } + /// Selects a cancellation-safe fallback for an attempt whose owner task + /// disappeared. Before the upstream write this is a void cancellation; + /// after the write it remains a gateway relay failure because provider + /// work may already have started. + pub(super) const fn abandonment_outcome(&self) -> ResponsesWebSocketTurnOutcome { + ResponsesWebSocketTurnOutcome::relay_task_abandonment(self.upstream_request_sent) + } + pub(super) fn deadline(&self) -> ResponsesWebSocketTurnDeadline { let (phase, timeout) = if self.first_event_elapsed_ms.is_some() { ( @@ -770,17 +837,6 @@ impl ResponsesProviderAttempt { } } -pub(super) async fn spawn_responses_websocket_turn_finalization( - state: AppState, - mut turn: ResponsesProviderAttempt, - outcome: ResponsesWebSocketTurnOutcome, -) -> tokio::task::JoinHandle<()> { - turn.release_admission().await; - tokio::spawn(async move { - turn.settle(&state, outcome).await; - }) -} - /// 把一轮 turn 的事实写进审计/用量 report context。 /// /// `effective_client_event` 是脱敏后的客户端事件(未启用脱敏时就是原事件)。 @@ -945,9 +1001,12 @@ fn elapsed_ms(started_at: Instant) -> u64 { #[cfg(test)] mod tests { + use std::sync::atomic::{AtomicU8, Ordering}; + use std::sync::Arc; use std::time::{Duration, Instant}; use aether_contracts::ExecutionTimeouts; + use aether_data_contracts::repository::candidates::RequestCandidateStatus; use serde_json::json; use super::super::observation::ResponsesStructuredTerminalObserver; @@ -958,7 +1017,8 @@ mod tests { }; use super::{ attach_client_delivery_to_report_context, prepare_websocket_report_context, - provider_terminal_outcome, resolve_responses_websocket_turn_timeouts, + provider_terminal_outcome, release_then_record_responses_websocket_admission_failure, + resolve_responses_websocket_turn_timeouts, responses_websocket_admission_failure_update, websocket_event_as_sse_line, ResponsesWebSocketTurnDeadline, ResponsesWebSocketTurnOutcome, ResponsesWebSocketTurnTimeoutPhase, }; @@ -966,6 +1026,72 @@ mod tests { classify_attempt_settlement, AttemptBilling, AttemptCandidateError, AttemptCandidateStatus, AttemptClientDelivery, AttemptProviderEffect, AttemptSettlementInputs, }; + use crate::GatewayError; + + #[tokio::test] + async fn admission_failure_releases_pool_lease_before_recording_candidate_terminal() { + let phase = Arc::new(AtomicU8::new(0)); + let release_phase = Arc::clone(&phase); + let record_phase = Arc::clone(&phase); + + release_then_record_responses_websocket_admission_failure( + async move { + assert_eq!(release_phase.swap(1, Ordering::SeqCst), 0); + }, + async move { + assert_eq!(record_phase.swap(2, Ordering::SeqCst), 1); + }, + ) + .await; + + assert_eq!(phase.load(Ordering::SeqCst), 2); + } + + #[test] + fn admission_timeout_terminalizes_the_seeded_candidate() { + let update = responses_websocket_admission_failure_update( + 1_000, + 1_025, + &GatewayError::AdmissionTimeout { + trace_id: "turn-1".to_string(), + gate: "gateway_upstream_execution", + queue_budget_ms: 25, + }, + ); + + assert_eq!(update.status, RequestCandidateStatus::Failed); + assert_eq!(update.status_code, Some(429)); + assert_eq!( + update.error_type.as_deref(), + Some("gateway_admission_timeout") + ); + assert_eq!( + update.error_message.as_deref(), + Some("gateway admission gate gateway_upstream_execution timed out after 25ms") + ); + assert_eq!(update.latency_ms, Some(25)); + assert_eq!(update.started_at_unix_ms, Some(1_000)); + assert_eq!(update.finished_at_unix_ms, Some(1_025)); + } + + #[test] + fn non_timeout_admission_failure_still_terminalizes_the_seeded_candidate() { + let update = responses_websocket_admission_failure_update( + 50, + 40, + &GatewayError::Internal("admission gate closed".to_string()), + ); + + assert_eq!(update.status, RequestCandidateStatus::Failed); + assert_eq!(update.status_code, Some(500)); + assert_eq!( + update.error_type.as_deref(), + Some("gateway_admission_failed") + ); + assert_eq!(update.latency_ms, Some(0)); + assert_eq!(update.started_at_unix_ms, Some(50)); + assert_eq!(update.finished_at_unix_ms, Some(40)); + } #[test] fn followup_context_uses_a_fresh_request_and_candidate() { @@ -1321,11 +1447,8 @@ mod tests { } #[test] - fn an_abandoned_turn_is_recorded_as_a_gateway_failure_not_a_cancellation() { - // A turn reclaimed by the Drop guard must not look like a client - // cancellation: cancelled turns skip the stream report entirely, which - // would defeat the point of reclaiming it. - let outcome = ResponsesWebSocketTurnOutcome::relay_task_abandoned(); + fn an_abandoned_turn_after_upstream_send_is_recorded_as_a_gateway_failure() { + let outcome = ResponsesWebSocketTurnOutcome::relay_task_abandonment(true); let facts = attempt_facts_for_outcome(None, AttemptClientDelivery::Complete, outcome); let settlement = classify_attempt_settlement(AttemptSettlementInputs { facts, @@ -1341,6 +1464,31 @@ mod tests { assert!(settlement.submit_execution_report); } + #[test] + fn an_abandoned_turn_before_upstream_send_is_void_and_does_not_penalize_provider() { + let outcome = ResponsesWebSocketTurnOutcome::relay_task_abandonment(false); + let facts = attempt_facts_for_outcome(None, AttemptClientDelivery::Complete, outcome); + let settlement = classify_attempt_settlement(AttemptSettlementInputs { + facts, + report_represents_failure: false, + observed_finish: false, + has_parser_error: false, + }); + + assert_eq!(settlement.status_code, 499); + assert_eq!(settlement.billing, AttemptBilling::Void); + assert_eq!( + settlement.candidate_status, + AttemptCandidateStatus::Cancelled + ); + assert_eq!(settlement.candidate_error, AttemptCandidateError::Cancelled); + assert_eq!( + settlement.provider_effect, + AttemptProviderEffect::ReleasePoolKeyLease + ); + assert!(!settlement.submit_execution_report); + } + #[test] fn turn_timeouts_reuse_provider_first_byte_and_request_deadlines() { let (first_event, terminal) = 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 index 4a4838818..859512cec 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn_state.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/turn_state.rs @@ -7,6 +7,7 @@ use serde_json::Value; +use super::control::ResponsesWebSocketTurnControl; use super::lifecycle::ActiveProviderAttempt; use super::request::response_create_has_previous_response_id; @@ -24,6 +25,10 @@ pub(super) struct LogicalTurn { pub(super) turn_attempt: u32, pub(super) retry_attempted: bool, pub(super) retry_unsafe_reason: Option<&'static str>, + /// Exact live control decision and strong auth snapshot used to authorize + /// this logical turn. Quota retries reuse it instead of falling back to the + /// connection's Upgrade-time authorization snapshot. + pub(super) turn_control: Option, } impl LogicalTurn { @@ -35,9 +40,15 @@ impl LogicalTurn { turn_attempt: 1, retry_attempted: false, retry_unsafe_reason: None, + turn_control: None, } } + pub(super) fn with_turn_control(mut self, turn_control: ResponsesWebSocketTurnControl) -> Self { + self.turn_control = Some(turn_control); + self + } + pub(super) fn quota_retry_block_reason(&self) -> Option<&'static str> { if self.retry_attempted { Some("quota_retry_already_attempted") diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/transport.rs b/apps/aether-gateway/src/handlers/proxy/websocket/transport.rs index c0bb22de5..19ebd25f1 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/transport.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/transport.rs @@ -10,9 +10,9 @@ use std::time::Duration; use axum::extract::ws::{CloseFrame as AxumCloseFrame, Message as AxumWsMessage, WebSocket}; use axum::http::header::{ ACCEPT, ACCEPT_ENCODING, CONNECTION, CONTENT_ENCODING, CONTENT_LENGTH, CONTENT_TYPE, HOST, - TRANSFER_ENCODING, UPGRADE, + PROXY_AUTHORIZATION, TE, TRAILER, TRANSFER_ENCODING, UPGRADE, }; -use axum::http::HeaderMap; +use axum::http::{HeaderMap, HeaderName}; use futures_util::{SinkExt, TryFutureExt}; use serde_json::json; use url::Url; @@ -113,8 +113,20 @@ pub(crate) fn websocket_handshake_headers( provider_headers: &BTreeMap, invalid_code: &'static str, ) -> Result { + // `build_request_headers` already strips `Connection` itself. Read the + // dynamic hop-by-hop names from the source map first, otherwise a header + // named by `Connection: keep-alive, x-provider-hop` would survive. + let connection_scoped_names = provider_headers + .iter() + .filter(|(name, _)| name.eq_ignore_ascii_case(CONNECTION.as_str())) + .flat_map(|(_, value)| value.split(',')) + .filter_map(|name| HeaderName::from_bytes(name.trim().as_bytes()).ok()) + .collect::>(); let mut headers = build_request_headers(provider_headers, None, false).map_err(|_| invalid_code)?; + for name in connection_scoped_names { + headers.remove(name); + } for header in [ ACCEPT, ACCEPT_ENCODING, @@ -123,11 +135,29 @@ pub(crate) fn websocket_handshake_headers( CONTENT_LENGTH, CONTENT_TYPE, HOST, + PROXY_AUTHORIZATION, + TE, + TRAILER, TRANSFER_ENCODING, UPGRADE, ] { headers.remove(header); } + for header in ["keep-alive", "proxy-connection"] { + headers.remove(header); + } + // The WebSocket client owns every Sec-WebSocket-* field, including + // extensions introduced after this gateway was built. Passing a + // downstream handshake field through here can corrupt negotiation or + // disclose the client's nonce/subprotocol to a different upstream. + let websocket_managed_names = headers + .keys() + .filter(|name| name.as_str().starts_with("sec-websocket-")) + .cloned() + .collect::>(); + for name in websocket_managed_names { + headers.remove(name); + } Ok(headers) } @@ -261,13 +291,6 @@ pub(crate) fn upstream_message_to_client(message: WreqWsMessage) -> AxumWsMessag } } -pub(crate) fn client_close_to_upstream(frame: Option) -> 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. @@ -332,9 +355,10 @@ pub(crate) async fn close_client_socket(client_socket: &mut WebSocket, code: u16 #[cfg(test)] mod tests { use super::{ - bounded_send, responses_websocket_error_event, websocket_upstream_url, WebSocketWriteError, - RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT, + bounded_send, responses_websocket_error_event, websocket_handshake_headers, + websocket_upstream_url, WebSocketWriteError, RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT, }; + use std::collections::BTreeMap; use std::time::Duration; #[tokio::test] @@ -403,4 +427,74 @@ mod tests { fn rejects_upstream_url_with_credentials() { assert!(websocket_upstream_url("https://token@example.test/responses", "invalid").is_err()); } + + #[test] + fn upstream_handshake_keeps_provider_auth_but_drops_transport_managed_headers() { + let provider_headers = BTreeMap::from([ + ( + "authorization".to_string(), + "Bearer provider-token".to_string(), + ), + ("x-api-key".to_string(), "provider-api-key".to_string()), + ( + "cookie".to_string(), + "provider_session=provider-cookie".to_string(), + ), + ( + "connection".to_string(), + "keep-alive, x-provider-hop".to_string(), + ), + ("x-provider-hop".to_string(), "must-not-pass".to_string()), + ("upgrade".to_string(), "websocket".to_string()), + ( + "sec-websocket-key".to_string(), + "downstream-nonce".to_string(), + ), + ( + "sec-websocket-future-field".to_string(), + "future-value".to_string(), + ), + ( + "proxy-authorization".to_string(), + "Basic must-not-pass".to_string(), + ), + ("x-provider-header".to_string(), "safe".to_string()), + ]); + + let headers = websocket_handshake_headers(&provider_headers, "invalid") + .expect("provider headers should be valid"); + + assert_eq!( + headers + .get("authorization") + .and_then(|value| value.to_str().ok()), + Some("Bearer provider-token") + ); + assert_eq!( + headers + .get("x-api-key") + .and_then(|value| value.to_str().ok()), + Some("provider-api-key") + ); + assert_eq!( + headers.get("cookie").and_then(|value| value.to_str().ok()), + Some("provider_session=provider-cookie") + ); + assert_eq!( + headers + .get("x-provider-header") + .and_then(|value| value.to_str().ok()), + Some("safe") + ); + for name in [ + "connection", + "x-provider-hop", + "upgrade", + "sec-websocket-key", + "sec-websocket-future-field", + "proxy-authorization", + ] { + assert!(headers.get(name).is_none(), "{name} must not survive"); + } + } } diff --git a/apps/aether-gateway/src/orchestration/mod.rs b/apps/aether-gateway/src/orchestration/mod.rs index 27a8bb2db..1818122c5 100644 --- a/apps/aether-gateway/src/orchestration/mod.rs +++ b/apps/aether-gateway/src/orchestration/mod.rs @@ -39,14 +39,13 @@ pub(crate) use self::codex_quota_breaker::{ log_codex_quota_breaker_check_failure, log_codex_quota_breaker_install_failure, }; pub(crate) use self::effects::{ - apply_local_execution_effect, spawn_local_oauth_success_effect, LocalAdaptiveRateLimitEffect, - LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect, LocalExecutionEffect, - LocalExecutionEffectContext, LocalHealthFailureEffect, LocalHealthSuccessEffect, - LocalOAuthInvalidationEffect, LocalOAuthSuccessEffect, LocalPoolErrorEffect, - apply_local_stream_failure_effects, apply_local_stream_success_effects, - release_local_pool_key_lease, - release_pool_key_lease_from_report_context, LocalAdaptiveRateLimitEffect, - LocalStreamFailureEffect, + 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, spawn_local_oauth_success_effect, + LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect, + LocalExecutionEffect, LocalExecutionEffectContext, LocalHealthFailureEffect, + LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, LocalOAuthSuccessEffect, + LocalPoolErrorEffect, LocalStreamFailureEffect, }; pub(crate) use self::health::{ project_local_failure_health, project_local_key_circuit_closed, diff --git a/apps/aether-gateway/src/orchestration/report_effects.rs b/apps/aether-gateway/src/orchestration/report_effects.rs index f68150e7f..12cef61c7 100644 --- a/apps/aether-gateway/src/orchestration/report_effects.rs +++ b/apps/aether-gateway/src/orchestration/report_effects.rs @@ -1254,33 +1254,6 @@ 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!({ @@ -1467,42 +1440,4 @@ 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/crates/aether-admin/src/provider/quota.rs b/crates/aether-admin/src/provider/quota.rs index 5125f3594..00de34d9d 100644 --- a/crates/aether-admin/src/provider/quota.rs +++ b/crates/aether-admin/src/provider/quota.rs @@ -3542,10 +3542,10 @@ pub fn parse_chatgpt_web_conversation_init_response( mod tests { use super::{ codex_build_invalid_state, codex_oauth_success_request_order_is_stale, - codex_rate_limit_metadata_exhausted, - codex_runtime_invalid_reason, extract_execution_error_detail, - merge_codex_quota_metadata_snapshot, normalize_codex_reset_credit_consume_outcome, - parse_antigravity_usage_response, parse_chatgpt_web_conversation_init_response, + codex_rate_limit_metadata_exhausted, codex_runtime_invalid_reason, + extract_execution_error_detail, merge_codex_quota_metadata_snapshot, + normalize_codex_reset_credit_consume_outcome, parse_antigravity_usage_response, + parse_chatgpt_web_conversation_init_response, parse_codex_backend_me_response, parse_codex_usage_headers, parse_codex_websocket_rate_limits_response, parse_codex_wham_reset_credits_detail_response, parse_codex_wham_usage_response, parse_gemini_cli_retrieve_user_quota_response, diff --git a/crates/aether-testing/integration/tests/responses_websocket_e2e.rs b/crates/aether-testing/integration/tests/responses_websocket_e2e.rs index 431077b75..9828d3436 100644 --- a/crates/aether-testing/integration/tests/responses_websocket_e2e.rs +++ b/crates/aether-testing/integration/tests/responses_websocket_e2e.rs @@ -31,7 +31,7 @@ use aether_gateway::{build_router_with_state, AppState, GatewayDataConfig, Usage use aether_testkit::SpawnedServer; use axum::extract::ws::{Message as AxumWsMessage, WebSocket, WebSocketUpgrade}; use axum::extract::State; -use axum::http::HeaderMap; +use axum::http::{HeaderMap, Uri}; use axum::response::Response; use axum::routing::get; use axum::Router; @@ -162,6 +162,304 @@ async fn continuation_reuses_one_upstream_connection_and_bills_both_turns() -> R Ok(()) } +#[tokio::test] +async fn persisted_previous_response_can_continue_on_a_new_client_connection( +) -> Result<(), BoxError> { + let harness = Harness::start(UpstreamBehavior::CompleteEveryTurn).await?; + + let mut first_client = harness.connect().await?; + first_client + .send(response_create(json!({ + "store": true, + "input": "first connection" + }))) + .await?; + let first = receive_event(&mut first_client, "response.completed").await?; + assert_eq!( + first.pointer("/response/id").and_then(Value::as_str), + Some("resp-e2e-1") + ); + first_client.close(None).await?; + + let mut second_client = harness.connect().await?; + second_client + .send(response_create(json!({ + "store": true, + "previous_response_id": "resp-e2e-1", + "input": "second connection" + }))) + .await?; + let second = receive_event(&mut second_client, "response.completed").await?; + assert_eq!( + second.pointer("/response/id").and_then(Value::as_str), + Some("resp-e2e-2") + ); + + let upstream_events = harness.upstream.observed_events().await; + assert_eq!(upstream_events.len(), 2); + assert_eq!(harness.upstream.connections(), 2); + assert_eq!(upstream_events[0]["store"], json!(true)); + assert_eq!(upstream_events[1]["store"], json!(true)); + assert_eq!( + upstream_events[1]["previous_response_id"], + json!("resp-e2e-1") + ); + + let audits = harness + .usage_audits_where(2, "billed cross-connection turns", is_billed) + .await?; + assert_eq!(audits.iter().filter(|audit| is_billed(audit)).count(), 2); + + second_client.close(None).await?; + Ok(()) +} + +#[tokio::test] +async fn unknown_future_request_and_response_fields_round_trip_opaquely() -> Result<(), BoxError> { + let harness = Harness::start(UpstreamBehavior::FutureEventThenComplete).await?; + let mut client = harness.connect().await?; + + client + .send(response_create(json!({ + "input": "future-compatible", + "future_request_capability": { + "mode": "opaque", + "revision": 7 + } + }))) + .await?; + + let future = receive_event(&mut client, "response.future.capability").await?; + assert_eq!(future["future_capability"]["enabled"], json!(true)); + assert_eq!(future["future_capability"]["revision"], json!(7)); + receive_event(&mut client, "response.completed").await?; + + let upstream_events = harness.upstream.observed_events().await; + assert_eq!(upstream_events.len(), 1); + assert_eq!( + upstream_events[0]["future_request_capability"], + json!({"mode": "opaque", "revision": 7}) + ); + + client.close(None).await?; + Ok(()) +} + +#[tokio::test] +async fn invalid_or_control_text_on_a_bound_socket_does_not_poison_the_next_valid_turn( +) -> Result<(), BoxError> { + let harness = Harness::start(UpstreamBehavior::CompleteEveryTurn).await?; + let mut client = harness.connect().await?; + + // Establish the provider connection first. Initial-frame validation has a + // separate handshake policy; this regression targets client events read by + // the live relay loop after the socket is fully bound. + client + .send(response_create(json!({"input": "bind the socket"}))) + .await?; + receive_event(&mut client, "response.completed").await?; + + client + .send(Message::Text("this is not json".into())) + .await?; + let invalid_json = receive_error_or_close(&mut client) + .await? + .ok_or("gateway closed after invalid JSON instead of returning an error")?; + assert_eq!(invalid_json["status"], json!(400)); + assert_eq!( + invalid_json.pointer("/error/code").and_then(Value::as_str), + Some("invalid_response_create") + ); + + client + .send(Message::Text( + json!({"type": "response.cancel"}).to_string().into(), + )) + .await?; + let unsupported_control = receive_error_or_close(&mut client) + .await? + .ok_or("gateway closed after a control event instead of returning an error")?; + assert_eq!(unsupported_control["status"], json!(400)); + assert_eq!( + unsupported_control + .pointer("/error/code") + .and_then(Value::as_str), + Some("expected_response_create") + ); + client + .send(response_create(json!({"input": "valid after errors"}))) + .await?; + let completed = receive_event(&mut client, "response.completed").await?; + assert_eq!( + completed + .pointer("/response/status") + .and_then(Value::as_str), + Some("completed") + ); + + let upstream_events = harness.upstream.observed_events().await; + assert_eq!(upstream_events.len(), 2); + assert_eq!(upstream_events[0]["input"], json!("bind the socket")); + assert_eq!(upstream_events[1]["input"], json!("valid after errors")); + assert_eq!( + harness.upstream.connections(), + 1, + "client protocol errors must not tear down the bound provider socket" + ); + let audits = harness + .usage_audits_where(2, "valid turns surrounding client errors", is_billed) + .await?; + assert_eq!(audits.iter().filter(|audit| is_billed(audit)).count(), 2); + + client.close(None).await?; + Ok(()) +} + +#[tokio::test] +async fn downstream_credentials_and_handshake_headers_do_not_reach_upstream() -> Result<(), BoxError> +{ + const QUERY_CREDENTIAL: &str = "downstream-query-secret"; + const COOKIE_CREDENTIAL: &str = "downstream-cookie-secret"; + const PROXY_CREDENTIAL: &str = "downstream-proxy-secret"; + const WEBSOCKET_FUTURE_VALUE: &str = "downstream-future-websocket-secret"; + const CONNECTION_SECRET: &str = "downstream-connection-secret"; + + let harness = Harness::start(UpstreamBehavior::CompleteEveryTurn).await?; + let mut request = format!( + "{}?key={QUERY_CREDENTIAL}&client_version=0.145.2", + harness.websocket_url + ) + .into_client_request()?; + request.headers_mut().insert( + "authorization", + http::HeaderValue::from_str(&format!("Bearer {CLIENT_API_KEY}"))?, + ); + request.headers_mut().insert( + "cookie", + http::HeaderValue::from_str(&format!("session={COOKIE_CREDENTIAL}"))?, + ); + request.headers_mut().insert( + "proxy-authorization", + http::HeaderValue::from_str(&format!("Bearer {PROXY_CREDENTIAL}"))?, + ); + request.headers_mut().insert( + "sec-websocket-future-capability", + http::HeaderValue::from_static(WEBSOCKET_FUTURE_VALUE), + ); + request.headers_mut().insert( + "connection", + http::HeaderValue::from_static("Upgrade, x-downstream-connection-secret"), + ); + request.headers_mut().insert( + "x-downstream-connection-secret", + http::HeaderValue::from_static(CONNECTION_SECRET), + ); + + let mut client = harness.connect_request(request).await?; + client + .send(response_create(json!({"input": "sanitize handshake"}))) + .await?; + receive_event(&mut client, "response.completed").await?; + + let handshakes = harness.upstream.observed_handshakes().await; + assert_eq!(handshakes.len(), 1); + let handshake = &handshakes[0]; + assert_eq!( + handshake.request_target, "/v1/responses?client_version=0.145.2", + "benign query state must survive while the downstream credential is removed" + ); + + let expected_provider_authorization = format!("Bearer {PROVIDER_API_KEY}"); + assert_eq!( + handshake + .headers + .get("authorization") + .and_then(|value| value.to_str().ok()), + Some(expected_provider_authorization.as_str()), + "provider authentication must be generated independently" + ); + for removed in [ + "cookie", + "proxy-authorization", + "sec-websocket-future-capability", + "x-downstream-connection-secret", + ] { + assert!( + handshake.headers.get(removed).is_none(), + "downstream header {removed} reached the upstream handshake" + ); + } + for (name, value) in &handshake.headers { + let value = value.to_str().unwrap_or_default(); + for secret in [ + CLIENT_API_KEY, + QUERY_CREDENTIAL, + COOKIE_CREDENTIAL, + PROXY_CREDENTIAL, + WEBSOCKET_FUTURE_VALUE, + CONNECTION_SECRET, + ] { + assert!( + !value.contains(secret), + "downstream secret leaked through upstream header {name}: {value}" + ); + } + } + + client.close(None).await?; + Ok(()) +} + +#[tokio::test] +async fn disabling_the_downstream_key_is_enforced_on_the_next_turn_of_the_same_socket( +) -> Result<(), BoxError> { + let harness = Harness::start(UpstreamBehavior::CompleteEveryTurn).await?; + let mut client = harness.connect().await?; + + client + .send(response_create(json!({"input": "before key disable"}))) + .await?; + receive_event(&mut client, "response.completed").await?; + + // Mutate through a repository handle that is independent of the gateway + // state. The next turn must perform a strong refresh instead of trusting + // the Upgrade-time or ordinary cached auth snapshot. + harness.set_client_api_key_active(false).await?; + + client + .send(response_create(json!({"input": "after key disable"}))) + .await?; + let rejection = receive_error_or_close(&mut client) + .await? + .ok_or("gateway closed without reporting the live API-key rejection")?; + assert_eq!(rejection["status"], json!(401)); + assert_eq!( + rejection.pointer("/error/code").and_then(Value::as_str), + Some("gateway_request_not_allowed") + ); + + let upstream_events = harness.upstream.observed_events().await; + assert_eq!( + upstream_events.len(), + 1, + "the disabled key's second turn must not reach the provider" + ); + assert_eq!(upstream_events[0]["input"], json!("before key disable")); + assert_eq!(harness.upstream.connections(), 1); + + let audits = harness + .usage_audits_where(1, "the turn completed before key disable", is_billed) + .await?; + assert_eq!( + audits.len(), + 1, + "a control-plane rejection must not create a provider attempt row" + ); + + client.close(None).await?; + Ok(()) +} + /// A client that walks away before the provider produced anything must settle /// as a void row: nothing was produced, so nothing is billed. /// @@ -546,6 +844,10 @@ impl Harness { "authorization", http::HeaderValue::from_str(&format!("Bearer {CLIENT_API_KEY}"))?, ); + self.connect_request(request).await + } + + async fn connect_request(&self, request: http::Request<()>) -> Result { let (socket, response) = tokio::time::timeout(RECEIVE_TIMEOUT, tokio_tungstenite::connect_async(request)) .await @@ -623,6 +925,21 @@ impl Harness { drop(backends); Ok(audits) } + + async fn set_client_api_key_active(&self, is_active: bool) -> Result<(), BoxError> { + let backends = DataBackends::from_config(DataLayerConfig::from_database( + self.database.config.clone(), + ))?; + backends + .write() + .auth_api_keys() + .ok_or("auth API key writer unavailable")? + .set_standalone_api_key_active(API_KEY_ID, is_active) + .await? + .ok_or("the E2E client API key disappeared before its live-control update")?; + drop(backends); + Ok(()) + } } // --------------------------------------------------------------------------- @@ -736,6 +1053,8 @@ enum UpstreamBehavior { /// 上游看到的是脱敏后的 body,所以回显出来的就是占位符——正是响应侧还原要处理 /// 的形状。 EchoInputBack, + /// Emit an event type and fields the gateway does not know, then complete. + FutureEventThenComplete, } #[derive(Debug)] @@ -744,6 +1063,13 @@ struct MockUpstreamState { connections: AtomicUsize, events: Mutex>, authorization_headers: Mutex>>, + handshakes: Mutex>, +} + +#[derive(Debug, Clone)] +struct ObservedUpstreamHandshake { + request_target: String, + headers: HeaderMap, } impl MockUpstreamState { @@ -753,6 +1079,7 @@ impl MockUpstreamState { connections: AtomicUsize::new(0), events: Mutex::new(Vec::new()), authorization_headers: Mutex::new(Vec::new()), + handshakes: Mutex::new(Vec::new()), } } @@ -767,6 +1094,10 @@ impl MockUpstreamState { async fn authorization_headers(&self) -> Vec> { self.authorization_headers.lock().await.clone() } + + async fn observed_handshakes(&self) -> Vec { + self.handshakes.lock().await.clone() + } } fn mock_upstream_router(state: Arc) -> Router { @@ -777,6 +1108,7 @@ fn mock_upstream_router(state: Arc) -> Router { async fn mock_responses_websocket( State(state): State>, + uri: Uri, headers: HeaderMap, ws: WebSocketUpgrade, ) -> Response { @@ -784,16 +1116,22 @@ async fn mock_responses_websocket( .get("authorization") .and_then(|value| value.to_str().ok()) .map(str::to_string); - ws.on_upgrade(move |socket| run_mock_upstream(socket, state, authorization)) + let handshake = ObservedUpstreamHandshake { + request_target: uri.to_string(), + headers, + }; + ws.on_upgrade(move |socket| run_mock_upstream(socket, state, authorization, handshake)) } async fn run_mock_upstream( mut socket: WebSocket, state: Arc, authorization: Option, + handshake: ObservedUpstreamHandshake, ) { state.connections.fetch_add(1, Ordering::AcqRel); state.authorization_headers.lock().await.push(authorization); + state.handshakes.lock().await.push(handshake); while let Some(message) = socket.recv().await { let Ok(message) = message else { break; @@ -850,6 +1188,27 @@ async fn run_mock_upstream( break; } } + UpstreamBehavior::FutureEventThenComplete => { + if send_mock_event( + &mut socket, + json!({ + "type": "response.future.capability", + "response_id": response_id, + "future_capability": { + "enabled": true, + "revision": 7 + } + }), + ) + .await + .is_err() + { + break; + } + if send_mock_turn(&mut socket, &response_id).await.is_err() { + break; + } + } } } AxumWsMessage::Ping(payload) => { diff --git a/docs/WebSocket-Mode.md b/docs/WebSocket-Mode.md index faa9aee98..d605a871d 100644 --- a/docs/WebSocket-Mode.md +++ b/docs/WebSocket-Mode.md @@ -8,7 +8,7 @@ WebSocket mode is compatible with both Zero Data Retention (ZDR) and `store=fals 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. +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. The [OpenAI WebSocket-mode guide](https://developers.openai.com/api/docs/guides/websocket-mode) reports up to roughly 40% faster end-to-end execution for workloads with 20 or more tool calls; this is an upstream product claim, not an Aether benchmark. ## Connect and create responses diff --git a/docs/operations/codex-responses-websocket-probe.md b/docs/operations/codex-responses-websocket-probe.md index b6987af1e..11065a79b 100644 --- a/docs/operations/codex-responses-websocket-probe.md +++ b/docs/operations/codex-responses-websocket-probe.md @@ -89,17 +89,22 @@ 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. +(unsupported-field stripping, its HTTP `store: false` default, and +`tool_choice` defaulting). An explicitly supplied WebSocket `store` value is +restored unchanged after that HTTP-oriented normalization. A +continuation turn with a non-null `previous_response_id` is revalidated through +the current scheduler and normalized against its pinned binding. It can never +move to another provider key, and it is rejected if that exact candidate is no +longer eligible or its physical binding changed. +`store`, `previous_response_id`, and `generate` are re-applied after +normalization because they are WebSocket protocol state that the provider body +contract may otherwise rewrite or strip. `stream` and `background` are removed +because they are HTTP transport fields, not WebSocket-mode fields. Every +independent `response.create` (one without a non-null `previous_response_id`) +runs access checks and candidate planning again, even when the public model is unchanged. +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 @@ -140,8 +145,9 @@ ws.send(json.dumps({ 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. +- Continuations with a non-null `previous_response_id` stay on the bound + provider key. Independent turns are re-authorized and re-planned each time; + they reuse the socket only when planning selects the same physical target. - Direct provider proxy settings are honored through the selected transport profile. Tunnel-mode proxy nodes are not supported for this bridge yet. diff --git a/frontend/src/features/providers/components/ProviderFormDialog.vue b/frontend/src/features/providers/components/ProviderFormDialog.vue index 387823d55..b075aa1ce 100644 --- a/frontend/src/features/providers/components/ProviderFormDialog.vue +++ b/frontend/src/features/providers/components/ProviderFormDialog.vue @@ -357,15 +357,25 @@ /> -
+
- {{ legacyT('Responses WebSocket 模式') }} +

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

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 index 1e38a7543..61596c0d4 100644 --- a/frontend/src/features/providers/components/__tests__/ProviderFormDialog.responses-websocket.spec.ts +++ b/frontend/src/features/providers/components/__tests__/ProviderFormDialog.responses-websocket.spec.ts @@ -13,6 +13,10 @@ describe('ProviderFormDialog Responses WebSocket switch', () => { 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'\"") + expect(source).toMatch( + /]*data-testid="responses-websocket-setting")(?![^>]*\bv-if=)[^>]*>[\s\S]{0,500}Responses WebSocket 模式/, + ) + expect(source).toContain('id="responses-websocket-enabled"') + expect(source).toContain(':aria-label="legacyT(\'Responses WebSocket 模式\')"') }) })