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 eed1116ff..62ab8af7f 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/binding.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/binding.rs @@ -70,6 +70,7 @@ impl UpstreamBindingIdentity { .map_err(|_| UpstreamBindingIdentityError::InvalidUpstreamUrl)? .to_string(); + let adapter_kind = adapter.kind(); let headers = websocket_handshake_headers(&decision.provider_request_headers, "invalid") .map_err(|_| UpstreamBindingIdentityError::InvalidHandshakeHeaders)?; let authentication_header_names = authentication_header_names(decision); @@ -82,7 +83,7 @@ impl UpstreamBindingIdentity { .map_err(|_| UpstreamBindingIdentityError::InvalidHandshakeHeaders)?; if authentication_header_names.contains(name.as_str()) { authentication_headers.insert(name, value.to_string()); - } else { + } else if !is_turn_scoped_handshake_header(adapter_kind, name.as_str()) { handshake_headers.insert(name, value.to_string()); } } @@ -90,7 +91,7 @@ impl UpstreamBindingIdentity { credential_binding_fingerprint(decision, &authentication_headers); Ok(Self { - adapter_kind: adapter.kind(), + adapter_kind, provider_id: decision.provider_id.clone(), endpoint_id: decision.endpoint_id.clone(), key_id: decision.key_id.clone(), @@ -101,6 +102,58 @@ impl UpstreamBindingIdentity { transport_profile: decision.transport_profile.clone(), }) } + + /// Returns only safe field/header names for operator diagnostics. Header + /// values and credential fingerprints must never reach logs. + pub(super) fn changed_field_names(&self, other: &Self) -> Vec { + let mut changed = Vec::new(); + if self.adapter_kind != other.adapter_kind { + changed.push("adapter_kind".to_string()); + } + if self.provider_id != other.provider_id { + changed.push("provider_id".to_string()); + } + if self.endpoint_id != other.endpoint_id { + changed.push("endpoint_id".to_string()); + } + if self.key_id != other.key_id { + changed.push("key_id".to_string()); + } + if self.upstream_url != other.upstream_url { + changed.push("upstream_url".to_string()); + } + for name in self + .handshake_headers + .keys() + .chain(other.handshake_headers.keys()) + .collect::>() + { + if self.handshake_headers.get(name) != other.handshake_headers.get(name) { + changed.push(format!("handshake_header:{name}")); + } + } + if self.credential_fingerprint != other.credential_fingerprint { + changed.push("credential_fingerprint".to_string()); + } + if self.proxy != other.proxy { + changed.push("proxy".to_string()); + } + if self.transport_profile != other.transport_profile { + changed.push("transport_profile".to_string()); + } + changed + } +} + +/// `x-codex-turn-metadata` describes one logical response turn. Fingerprint +/// convergence intentionally rewrites its `turn_id` and timestamp for every +/// `response.create`, including a continuation on an already-upgraded socket. +/// It therefore cannot identify the physical handshake that owns a +/// `previous_response_id`; the per-turn copy in `client_metadata` still travels +/// in the response.create body. +fn is_turn_scoped_handshake_header(adapter_kind: ResponsesWebSocketAdapter, name: &str) -> bool { + adapter_kind == ResponsesWebSocketAdapter::Codex + && name.eq_ignore_ascii_case("x-codex-turn-metadata") } /// Header names that carry credentials in the provider handshake. The @@ -434,6 +487,81 @@ mod tests { ); } + #[test] + fn codex_turn_metadata_changes_do_not_change_the_physical_binding() { + 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" + })); + first.provider_request_headers.insert( + "x-codex-turn-metadata".to_string(), + json!({ + "session_id": "session-1", + "thread_id": "thread-1", + "turn_id": "turn-1", + "turn_started_at_unix_ms": 1 + }) + .to_string(), + ); + let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap(); + assert!(!first_identity + .handshake_headers + .contains_key("x-codex-turn-metadata")); + + let mut continuation = first; + continuation.provider_request_headers.insert( + "x-codex-turn-metadata".to_string(), + json!({ + "session_id": "session-1", + "thread_id": "thread-1", + "turn_id": "turn-2", + "turn_started_at_unix_ms": 2 + }) + .to_string(), + ); + assert_eq!( + first_identity, + UpstreamBindingIdentity::from_decision(adapter, &continuation).unwrap() + ); + } + + #[test] + fn only_codex_turn_metadata_is_excluded_from_binding_headers() { + let codex_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" + })); + first.provider_request_headers.insert( + "x-codex-turn-metadata".to_string(), + r#"{"turn_id":"turn-1"}"#.to_string(), + ); + let first_identity = UpstreamBindingIdentity::from_decision(codex_adapter, &first).unwrap(); + + let mut changed = first; + changed + .provider_request_headers + .insert("x-codex-window-id".to_string(), "window-2".to_string()); + let changed_identity = + UpstreamBindingIdentity::from_decision(codex_adapter, &changed).unwrap(); + assert_ne!(first_identity, changed_identity); + assert_eq!( + first_identity.changed_field_names(&changed_identity), + vec!["handshake_header:x-codex-window-id".to_string()] + ); + + let standard_adapter = + resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard); + let standard_identity = + UpstreamBindingIdentity::from_decision(standard_adapter, &changed).unwrap(); + assert!(standard_identity + .handshake_headers + .contains_key("x-codex-turn-metadata")); + } + #[test] fn codex_authorization_override_changes_binding_with_the_same_generation() { let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex); 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 89012b322..f7140f84c 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/client.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/client.rs @@ -29,7 +29,9 @@ use super::turn::{ ResponsesWebSocketTurnOutcome, }; use super::turn_state::LogicalTurn; -use super::upstream::{bind_responses_upstream, decision_reuses_bound_upstream}; +use super::upstream::{ + bind_responses_upstream, decision_bound_upstream_change_fields, decision_reuses_bound_upstream, +}; use crate::ai_serving::ResponsesWebSocketPinnedCandidate; use crate::clock::current_unix_secs; use crate::control::GatewayControlDecision; @@ -496,9 +498,15 @@ async fn forward_pinned_continuation( 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()) - { + let reuses_bound_upstream = decision_reuses_bound_upstream(bound, adapter, &decision); + let provider_model_changed = + planned_provider_model.as_deref() != Some(bound.provider_model.as_str()); + if !reuses_bound_upstream || provider_model_changed { + let binding_change_fields = if reuses_bound_upstream { + Vec::new() + } else { + decision_bound_upstream_change_fields(bound, adapter, &decision) + }; planned_lease.release().await; warn!( event_name = "responses_websocket_continuation_binding_changed", @@ -507,6 +515,8 @@ async fn forward_pinned_continuation( websocket = true, trace_id = %context.trace_id, key_id = ?decision.key_id, + binding_change_fields = ?binding_change_fields, + provider_model_changed, "gateway rejected a continuation after the pinned candidate's physical binding changed" ); send_gateway_error_with_status( 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 bc6f416f5..7d7f75c0d 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/session.rs @@ -64,7 +64,7 @@ macro_rules! warn { }; } -#[derive(Debug, Clone, Copy)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] enum InitialMessageError { TimedOut, ClientClosed, @@ -76,6 +76,69 @@ enum InitialMessageError { InvalidModel, } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +struct InitialMessageFrameMetadata { + opcode: &'static str, + bytes: usize, +} + +impl InitialMessageFrameMetadata { + fn from_message(message: &AxumWsMessage) -> Self { + match message { + AxumWsMessage::Text(text) => Self { + opcode: "text", + bytes: text.len(), + }, + AxumWsMessage::Binary(payload) => Self { + opcode: "binary", + bytes: payload.len(), + }, + AxumWsMessage::Ping(payload) => Self { + opcode: "ping", + bytes: payload.len(), + }, + AxumWsMessage::Pong(payload) => Self { + opcode: "pong", + bytes: payload.len(), + }, + AxumWsMessage::Close(frame) => Self { + opcode: "close", + // Only the payload length is retained. The untrusted close + // reason itself must never cross the logging boundary. + bytes: frame + .as_ref() + .map_or(0, |frame| 2usize.saturating_add(frame.reason.len())), + }, + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +struct InitialMessageFailure { + error: InitialMessageError, + last_frame: Option, +} + +impl InitialMessageFailure { + const fn new( + error: InitialMessageError, + last_frame: Option, + ) -> Self { + Self { error, last_frame } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +struct InitialMessageDiagnostic { + error_code: &'static str, + error_kind: &'static str, + client_message: Option<&'static str>, + close_code: u16, + timed_out: bool, + last_frame_opcode: Option<&'static str>, + last_frame_bytes: Option, +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum ConnectionTermination { ConnectionLimitReached, @@ -106,6 +169,100 @@ impl InitialMessageError { } } } + + const fn client_message(self) -> Option<&'static str> { + match self { + Self::TimedOut => Some("Timed out waiting for the initial response.create event"), + Self::ClientClosed => None, + Self::ClientRead => { + Some("Failed to read the initial WebSocket event before response.create") + } + Self::UnsupportedFrame => { + Some("The initial response.create event must be sent as a text WebSocket message") + } + Self::InvalidJson => { + Some("The initial WebSocket text message must be a JSON response.create object") + } + Self::MissingResponseCreate => { + Some("The initial WebSocket JSON object must have type response.create") + } + Self::MissingModel => Some("The initial response.create event must include a model"), + Self::InvalidModel => { + Some("response.create.model must be a non-empty string no longer than 256 bytes") + } + } + } + + const fn kind(self) -> &'static str { + match self { + Self::TimedOut => "timeout", + Self::ClientClosed => "client_closed", + Self::ClientRead => "client_read_failed", + Self::UnsupportedFrame => "unsupported_frame", + Self::InvalidJson => "invalid_json", + Self::MissingResponseCreate => "unexpected_event_type", + Self::MissingModel => "missing_model", + Self::InvalidModel => "invalid_model", + } + } + + const fn timed_out(self) -> bool { + matches!(self, Self::TimedOut) + } +} + +fn initial_message_diagnostic(failure: InitialMessageFailure) -> InitialMessageDiagnostic { + InitialMessageDiagnostic { + error_code: failure.error.code(), + error_kind: failure.error.kind(), + client_message: failure.error.client_message(), + close_code: failure.error.close_code(), + timed_out: failure.error.timed_out(), + last_frame_opcode: failure.last_frame.map(|frame| frame.opcode), + last_frame_bytes: failure.last_frame.map(|frame| frame.bytes), + } +} + +fn log_initial_message_failure( + context: &WebSocketRequestContext, + failure: InitialMessageFailure, + upgraded_at: std::time::Instant, +) { + let diagnostic = initial_message_diagnostic(failure); + let auth_context = context.decision.auth_context.as_ref(); + + // Initial-frame validation intentionally runs before provider planning, so + // no provider identity exists yet. Keep the usual identity fields in the + // event schema and say so explicitly instead of guessing from request data. + warn!( + event_name = "responses_websocket_initial_event_rejected", + log_type = "ops", + transport = WEBSOCKET_LOG_TRANSPORT, + websocket = true, + trace_id = %context.trace_id, + user_id = auth_context.map(|auth| auth.user_id.as_str()).unwrap_or("-"), + api_key_id = auth_context + .map(|auth| auth.api_key_id.as_str()) + .unwrap_or("-"), + provider_selected = false, + provider_id = "", + endpoint_id = "", + key_id = "", + path = %context.uri.path(), + route_class = context.decision.route_class.as_deref().unwrap_or("-"), + route_kind = context.decision.route_kind.as_deref().unwrap_or("-"), + error_code = diagnostic.error_code, + error_kind = diagnostic.error_kind, + close_code = diagnostic.close_code, + timed_out = diagnostic.timed_out, + upgrade_to_initial_outcome_ms = upgraded_at.elapsed().as_millis() as u64, + initial_message_timeout_ms = RESPONSES_WEBSOCKET_SESSION_LIMITS + .initial_message_timeout + .as_millis() as u64, + last_frame_opcode = diagnostic.last_frame_opcode.unwrap_or("none"), + last_frame_bytes = ?diagnostic.last_frame_bytes, + "gateway rejected the initial Responses WebSocket event" + ); } pub(super) async fn run_responses_websocket( @@ -113,6 +270,7 @@ pub(super) async fn run_responses_websocket( state: AppState, mut context: WebSocketRequestContext, ) { + let upgraded_at = std::time::Instant::now(); // 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 = @@ -122,7 +280,7 @@ pub(super) async fn run_responses_websocket( connection_log.log_opened(); let bootstrap_result = supervise_responses_websocket_phase( - bootstrap_responses_websocket(&mut client_socket, state.clone(), &context), + bootstrap_responses_websocket(&mut client_socket, state.clone(), &context, upgraded_at), connection_deadline, connection_permit.as_ref(), ) @@ -196,19 +354,20 @@ async fn bootstrap_responses_websocket( client_socket: &mut WebSocket, state: AppState, context: &WebSocketRequestContext, + upgraded_at: std::time::Instant, ) -> 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( + Err(failure) => { + if let Some(client_message) = failure.error.client_message() { + log_initial_message_failure(context, failure, upgraded_at); + send_gateway_error(client_socket, failure.error.code(), client_message).await; + close_client_socket( client_socket, - error.code(), - "WebSocket must start with a valid response.create event", + failure.error.close_code(), + "invalid_initial_event", ) .await; - close_client_socket(client_socket, error.close_code(), "invalid_initial_event") - .await; } return None; } @@ -598,7 +757,7 @@ async fn close_terminated_relay( /// 但不会重置计时器。防止客户端通过周期性 Ping 无限占用 connection permit。 async fn receive_initial_response_create( client_socket: &mut WebSocket, -) -> Result<(String, Value), InitialMessageError> { +) -> Result<(String, Value), InitialMessageFailure> { receive_initial_response_create_with_deadline( client_socket, RESPONSES_WEBSOCKET_SESSION_LIMITS.initial_message_timeout, @@ -615,7 +774,7 @@ async fn receive_initial_response_create( async fn receive_initial_response_create_with_deadline( socket: &mut S, deadline_budget: std::time::Duration, -) -> Result<(String, Value), InitialMessageError> +) -> Result<(String, Value), InitialMessageFailure> where S: futures_util::Stream> + futures_util::Sink @@ -625,29 +784,51 @@ where // 绝对 deadline:入口计算一次,后续所有迭代共享,Ping/Pong 不会重启 let deadline = tokio::time::Instant::now() + deadline_budget; + let mut last_frame = None; loop { let message = tokio::time::timeout_at(deadline, socket.next()) .await - .map_err(|_| InitialMessageError::TimedOut)?; + .map_err(|_| InitialMessageFailure::new(InitialMessageError::TimedOut, last_frame))?; let Some(message) = message else { - return Err(InitialMessageError::ClientClosed); + return Err(InitialMessageFailure::new( + InitialMessageError::ClientClosed, + last_frame, + )); }; - let message = message.map_err(|_| InitialMessageError::ClientRead)?; + let message = message + .map_err(|_| InitialMessageFailure::new(InitialMessageError::ClientRead, last_frame))?; + last_frame = Some(InitialMessageFrameMetadata::from_message(&message)); match message { AxumWsMessage::Ping(payload) => { tokio::time::timeout_at(deadline, socket.send(AxumWsMessage::Pong(payload))) .await - .map_err(|_| InitialMessageError::TimedOut)? - .map_err(|_| InitialMessageError::ClientRead)?; + .map_err(|_| { + InitialMessageFailure::new(InitialMessageError::TimedOut, last_frame) + })? + .map_err(|_| { + InitialMessageFailure::new(InitialMessageError::ClientRead, last_frame) + })?; } AxumWsMessage::Pong(_) => {} - AxumWsMessage::Close(_) => return Err(InitialMessageError::ClientClosed), - AxumWsMessage::Binary(_) => return Err(InitialMessageError::UnsupportedFrame), + AxumWsMessage::Close(_) => { + return Err(InitialMessageFailure::new( + InitialMessageError::ClientClosed, + last_frame, + )); + } + AxumWsMessage::Binary(_) => { + return Err(InitialMessageFailure::new( + InitialMessageError::UnsupportedFrame, + last_frame, + )); + } AxumWsMessage::Text(text) => { let text = text.to_string(); - let event: Value = - serde_json::from_str(&text).map_err(|_| InitialMessageError::InvalidJson)?; - validate_initial_response_create(&event)?; + let event: Value = serde_json::from_str(&text).map_err(|_| { + InitialMessageFailure::new(InitialMessageError::InvalidJson, last_frame) + })?; + validate_initial_response_create(&event) + .map_err(|error| InitialMessageFailure::new(error, last_frame))?; return Ok((text, event)); } } @@ -737,6 +918,109 @@ mod tests { } } + #[test] + fn initial_message_error_diagnostics_are_stable() { + use super::{initial_message_diagnostic, InitialMessageError, InitialMessageFailure}; + + let cases = [ + ( + InitialMessageError::TimedOut, + "initial_response_create_timeout", + "timeout", + Some("Timed out waiting for the initial response.create event"), + 1013, + true, + ), + ( + InitialMessageError::ClientClosed, + "client_closed", + "client_closed", + None, + 1000, + false, + ), + ( + InitialMessageError::ClientRead, + "client_read_failed", + "client_read_failed", + Some("Failed to read the initial WebSocket event before response.create"), + 1008, + false, + ), + ( + InitialMessageError::UnsupportedFrame, + "initial_response_create_must_be_text", + "unsupported_frame", + Some("The initial response.create event must be sent as a text WebSocket message"), + 1008, + false, + ), + ( + InitialMessageError::InvalidJson, + "invalid_response_create", + "invalid_json", + Some("The initial WebSocket text message must be a JSON response.create object"), + 1008, + false, + ), + ( + InitialMessageError::MissingResponseCreate, + "expected_response_create", + "unexpected_event_type", + Some("The initial WebSocket JSON object must have type response.create"), + 1008, + false, + ), + ( + InitialMessageError::MissingModel, + "response_create_model_required", + "missing_model", + Some("The initial response.create event must include a model"), + 1008, + false, + ), + ( + InitialMessageError::InvalidModel, + "invalid_response_create_model", + "invalid_model", + Some("response.create.model must be a non-empty string no longer than 256 bytes"), + 1008, + false, + ), + ]; + + for (error, error_code, error_kind, client_message, close_code, timed_out) in cases { + let diagnostic = initial_message_diagnostic(InitialMessageFailure::new(error, None)); + assert_eq!(diagnostic.error_code, error_code); + assert_eq!(diagnostic.error_kind, error_kind); + assert_eq!(diagnostic.client_message, client_message); + assert_eq!(diagnostic.close_code, close_code); + assert_eq!(diagnostic.timed_out, timed_out); + assert_eq!(diagnostic.last_frame_opcode, None); + assert_eq!(diagnostic.last_frame_bytes, None); + } + } + + #[test] + fn initial_message_diagnostic_retains_only_safe_frame_shape() { + use super::{ + initial_message_diagnostic, InitialMessageError, InitialMessageFailure, + InitialMessageFrameMetadata, + }; + + let secret_body = r#"{"type":"not-response.create","token":"must-not-log"}"#; + let frame = Message::Text(secret_body.to_string().into()); + let metadata = InitialMessageFrameMetadata::from_message(&frame); + let diagnostic = initial_message_diagnostic(InitialMessageFailure::new( + InitialMessageError::MissingResponseCreate, + Some(metadata), + )); + + assert_eq!(diagnostic.last_frame_opcode, Some("text")); + assert_eq!(diagnostic.last_frame_bytes, Some(secret_body.len())); + assert!(!format!("{diagnostic:?}").contains("must-not-log")); + } + #[tokio::test] async fn phase_supervisor_drops_work_at_the_absolute_connection_deadline() { let dropped = Arc::new(AtomicUsize::new(0)); @@ -1535,7 +1819,10 @@ mod tests { .await .expect("the initial-message deadline must cancel a stalled Pong write"); - assert!(matches!(result, Err(InitialMessageError::TimedOut))); + assert!(matches!( + &result, + Err(failure) if failure.error == InitialMessageError::TimedOut + )); } /// 验证 receive_initial_response_create_with_deadline 的绝对 deadline: @@ -1583,7 +1870,10 @@ mod tests { ping_task.abort(); assert!( - matches!(result, Err(InitialMessageError::TimedOut)), + matches!( + &result, + Err(failure) if failure.error == InitialMessageError::TimedOut + ), "expected TimedOut after absolute deadline, got: {result:?}" ); } diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/responses/upstream.rs b/apps/aether-gateway/src/handlers/proxy/websocket/responses/upstream.rs index 59753b176..31f7384df 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/responses/upstream.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/responses/upstream.rs @@ -156,6 +156,20 @@ pub(super) fn decision_reuses_bound_upstream( .unwrap_or(false) } +pub(super) fn decision_bound_upstream_change_fields( + bound: &BoundResponsesConnection, + adapter: &'static dyn ResponsesWebSocketProtocolAdapter, + decision: &AiExecutionDecision, +) -> Vec { + if bound.upstream.is_none() { + return vec!["upstream_socket".to_string()]; + } + match UpstreamBindingIdentity::from_decision(adapter, decision) { + Ok(identity) => bound.binding_identity.changed_field_names(&identity), + Err(_) => vec!["binding_identity_invalid".to_string()], + } +} + #[cfg(test)] mod tests { use std::time::Duration;