mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-05 00:47:48 +08:00
fix(ws): preserve Codex continuation bindings
This commit is contained in:
@@ -70,6 +70,7 @@ impl UpstreamBindingIdentity {
|
|||||||
.map_err(|_| UpstreamBindingIdentityError::InvalidUpstreamUrl)?
|
.map_err(|_| UpstreamBindingIdentityError::InvalidUpstreamUrl)?
|
||||||
.to_string();
|
.to_string();
|
||||||
|
|
||||||
|
let adapter_kind = adapter.kind();
|
||||||
let headers = websocket_handshake_headers(&decision.provider_request_headers, "invalid")
|
let headers = websocket_handshake_headers(&decision.provider_request_headers, "invalid")
|
||||||
.map_err(|_| UpstreamBindingIdentityError::InvalidHandshakeHeaders)?;
|
.map_err(|_| UpstreamBindingIdentityError::InvalidHandshakeHeaders)?;
|
||||||
let authentication_header_names = authentication_header_names(decision);
|
let authentication_header_names = authentication_header_names(decision);
|
||||||
@@ -82,7 +83,7 @@ impl UpstreamBindingIdentity {
|
|||||||
.map_err(|_| UpstreamBindingIdentityError::InvalidHandshakeHeaders)?;
|
.map_err(|_| UpstreamBindingIdentityError::InvalidHandshakeHeaders)?;
|
||||||
if authentication_header_names.contains(name.as_str()) {
|
if authentication_header_names.contains(name.as_str()) {
|
||||||
authentication_headers.insert(name, value.to_string());
|
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());
|
handshake_headers.insert(name, value.to_string());
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -90,7 +91,7 @@ impl UpstreamBindingIdentity {
|
|||||||
credential_binding_fingerprint(decision, &authentication_headers);
|
credential_binding_fingerprint(decision, &authentication_headers);
|
||||||
|
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
adapter_kind: adapter.kind(),
|
adapter_kind,
|
||||||
provider_id: decision.provider_id.clone(),
|
provider_id: decision.provider_id.clone(),
|
||||||
endpoint_id: decision.endpoint_id.clone(),
|
endpoint_id: decision.endpoint_id.clone(),
|
||||||
key_id: decision.key_id.clone(),
|
key_id: decision.key_id.clone(),
|
||||||
@@ -101,6 +102,58 @@ impl UpstreamBindingIdentity {
|
|||||||
transport_profile: decision.transport_profile.clone(),
|
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<String> {
|
||||||
|
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::<BTreeSet<_>>()
|
||||||
|
{
|
||||||
|
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
|
/// 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]
|
#[test]
|
||||||
fn codex_authorization_override_changes_binding_with_the_same_generation() {
|
fn codex_authorization_override_changes_binding_with_the_same_generation() {
|
||||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
|
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
|
||||||
|
|||||||
@@ -29,7 +29,9 @@ use super::turn::{
|
|||||||
ResponsesWebSocketTurnOutcome,
|
ResponsesWebSocketTurnOutcome,
|
||||||
};
|
};
|
||||||
use super::turn_state::LogicalTurn;
|
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::ai_serving::ResponsesWebSocketPinnedCandidate;
|
||||||
use crate::clock::current_unix_secs;
|
use crate::clock::current_unix_secs;
|
||||||
use crate::control::GatewayControlDecision;
|
use crate::control::GatewayControlDecision;
|
||||||
@@ -496,9 +498,15 @@ async fn forward_pinned_continuation(
|
|||||||
let normalization = planned.normalization;
|
let normalization = planned.normalization;
|
||||||
let decision = planned.execution;
|
let decision = planned.execution;
|
||||||
let planned_provider_model = provider_model_from_decision(&decision);
|
let planned_provider_model = provider_model_from_decision(&decision);
|
||||||
if !decision_reuses_bound_upstream(bound, adapter, &decision)
|
let reuses_bound_upstream = decision_reuses_bound_upstream(bound, adapter, &decision);
|
||||||
|| planned_provider_model.as_deref() != Some(bound.provider_model.as_str())
|
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;
|
planned_lease.release().await;
|
||||||
warn!(
|
warn!(
|
||||||
event_name = "responses_websocket_continuation_binding_changed",
|
event_name = "responses_websocket_continuation_binding_changed",
|
||||||
@@ -507,6 +515,8 @@ async fn forward_pinned_continuation(
|
|||||||
websocket = true,
|
websocket = true,
|
||||||
trace_id = %context.trace_id,
|
trace_id = %context.trace_id,
|
||||||
key_id = ?decision.key_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"
|
"gateway rejected a continuation after the pinned candidate's physical binding changed"
|
||||||
);
|
);
|
||||||
send_gateway_error_with_status(
|
send_gateway_error_with_status(
|
||||||
|
|||||||
@@ -64,7 +64,7 @@ macro_rules! warn {
|
|||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||||
enum InitialMessageError {
|
enum InitialMessageError {
|
||||||
TimedOut,
|
TimedOut,
|
||||||
ClientClosed,
|
ClientClosed,
|
||||||
@@ -76,6 +76,69 @@ enum InitialMessageError {
|
|||||||
InvalidModel,
|
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<InitialMessageFrameMetadata>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl InitialMessageFailure {
|
||||||
|
const fn new(
|
||||||
|
error: InitialMessageError,
|
||||||
|
last_frame: Option<InitialMessageFrameMetadata>,
|
||||||
|
) -> 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<usize>,
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||||
enum ConnectionTermination {
|
enum ConnectionTermination {
|
||||||
ConnectionLimitReached,
|
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 = "<unplanned>",
|
||||||
|
endpoint_id = "<unplanned>",
|
||||||
|
key_id = "<unplanned>",
|
||||||
|
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(
|
pub(super) async fn run_responses_websocket(
|
||||||
@@ -113,6 +270,7 @@ pub(super) async fn run_responses_websocket(
|
|||||||
state: AppState,
|
state: AppState,
|
||||||
mut context: WebSocketRequestContext,
|
mut context: WebSocketRequestContext,
|
||||||
) {
|
) {
|
||||||
|
let upgraded_at = std::time::Instant::now();
|
||||||
// The public connection limit starts when the HTTP Upgrade hands us the
|
// The public connection limit starts when the HTTP Upgrade hands us the
|
||||||
// socket, not after provider planning and its upstream handshake finish.
|
// socket, not after provider planning and its upstream handshake finish.
|
||||||
let connection_deadline =
|
let connection_deadline =
|
||||||
@@ -122,7 +280,7 @@ pub(super) async fn run_responses_websocket(
|
|||||||
connection_log.log_opened();
|
connection_log.log_opened();
|
||||||
|
|
||||||
let bootstrap_result = supervise_responses_websocket_phase(
|
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_deadline,
|
||||||
connection_permit.as_ref(),
|
connection_permit.as_ref(),
|
||||||
)
|
)
|
||||||
@@ -196,19 +354,20 @@ async fn bootstrap_responses_websocket(
|
|||||||
client_socket: &mut WebSocket,
|
client_socket: &mut WebSocket,
|
||||||
state: AppState,
|
state: AppState,
|
||||||
context: &WebSocketRequestContext,
|
context: &WebSocketRequestContext,
|
||||||
|
upgraded_at: std::time::Instant,
|
||||||
) -> Option<BoundResponsesConnection> {
|
) -> Option<BoundResponsesConnection> {
|
||||||
let (_first_text, first_event) = match receive_initial_response_create(client_socket).await {
|
let (_first_text, first_event) = match receive_initial_response_create(client_socket).await {
|
||||||
Ok(value) => value,
|
Ok(value) => value,
|
||||||
Err(error) => {
|
Err(failure) => {
|
||||||
if !matches!(error, InitialMessageError::ClientClosed) {
|
if let Some(client_message) = failure.error.client_message() {
|
||||||
send_gateway_error(
|
log_initial_message_failure(context, failure, upgraded_at);
|
||||||
|
send_gateway_error(client_socket, failure.error.code(), client_message).await;
|
||||||
|
close_client_socket(
|
||||||
client_socket,
|
client_socket,
|
||||||
error.code(),
|
failure.error.close_code(),
|
||||||
"WebSocket must start with a valid response.create event",
|
"invalid_initial_event",
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
close_client_socket(client_socket, error.close_code(), "invalid_initial_event")
|
|
||||||
.await;
|
|
||||||
}
|
}
|
||||||
return None;
|
return None;
|
||||||
}
|
}
|
||||||
@@ -598,7 +757,7 @@ async fn close_terminated_relay(
|
|||||||
/// 但不会重置计时器。防止客户端通过周期性 Ping 无限占用 connection permit。
|
/// 但不会重置计时器。防止客户端通过周期性 Ping 无限占用 connection permit。
|
||||||
async fn receive_initial_response_create(
|
async fn receive_initial_response_create(
|
||||||
client_socket: &mut WebSocket,
|
client_socket: &mut WebSocket,
|
||||||
) -> Result<(String, Value), InitialMessageError> {
|
) -> Result<(String, Value), InitialMessageFailure> {
|
||||||
receive_initial_response_create_with_deadline(
|
receive_initial_response_create_with_deadline(
|
||||||
client_socket,
|
client_socket,
|
||||||
RESPONSES_WEBSOCKET_SESSION_LIMITS.initial_message_timeout,
|
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<S>(
|
async fn receive_initial_response_create_with_deadline<S>(
|
||||||
socket: &mut S,
|
socket: &mut S,
|
||||||
deadline_budget: std::time::Duration,
|
deadline_budget: std::time::Duration,
|
||||||
) -> Result<(String, Value), InitialMessageError>
|
) -> Result<(String, Value), InitialMessageFailure>
|
||||||
where
|
where
|
||||||
S: futures_util::Stream<Item = Result<AxumWsMessage, axum::Error>>
|
S: futures_util::Stream<Item = Result<AxumWsMessage, axum::Error>>
|
||||||
+ futures_util::Sink<AxumWsMessage, Error = axum::Error>
|
+ futures_util::Sink<AxumWsMessage, Error = axum::Error>
|
||||||
@@ -625,29 +784,51 @@ where
|
|||||||
|
|
||||||
// 绝对 deadline:入口计算一次,后续所有迭代共享,Ping/Pong 不会重启
|
// 绝对 deadline:入口计算一次,后续所有迭代共享,Ping/Pong 不会重启
|
||||||
let deadline = tokio::time::Instant::now() + deadline_budget;
|
let deadline = tokio::time::Instant::now() + deadline_budget;
|
||||||
|
let mut last_frame = None;
|
||||||
loop {
|
loop {
|
||||||
let message = tokio::time::timeout_at(deadline, socket.next())
|
let message = tokio::time::timeout_at(deadline, socket.next())
|
||||||
.await
|
.await
|
||||||
.map_err(|_| InitialMessageError::TimedOut)?;
|
.map_err(|_| InitialMessageFailure::new(InitialMessageError::TimedOut, last_frame))?;
|
||||||
let Some(message) = message else {
|
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 {
|
match message {
|
||||||
AxumWsMessage::Ping(payload) => {
|
AxumWsMessage::Ping(payload) => {
|
||||||
tokio::time::timeout_at(deadline, socket.send(AxumWsMessage::Pong(payload)))
|
tokio::time::timeout_at(deadline, socket.send(AxumWsMessage::Pong(payload)))
|
||||||
.await
|
.await
|
||||||
.map_err(|_| InitialMessageError::TimedOut)?
|
.map_err(|_| {
|
||||||
.map_err(|_| InitialMessageError::ClientRead)?;
|
InitialMessageFailure::new(InitialMessageError::TimedOut, last_frame)
|
||||||
|
})?
|
||||||
|
.map_err(|_| {
|
||||||
|
InitialMessageFailure::new(InitialMessageError::ClientRead, last_frame)
|
||||||
|
})?;
|
||||||
}
|
}
|
||||||
AxumWsMessage::Pong(_) => {}
|
AxumWsMessage::Pong(_) => {}
|
||||||
AxumWsMessage::Close(_) => return Err(InitialMessageError::ClientClosed),
|
AxumWsMessage::Close(_) => {
|
||||||
AxumWsMessage::Binary(_) => return Err(InitialMessageError::UnsupportedFrame),
|
return Err(InitialMessageFailure::new(
|
||||||
|
InitialMessageError::ClientClosed,
|
||||||
|
last_frame,
|
||||||
|
));
|
||||||
|
}
|
||||||
|
AxumWsMessage::Binary(_) => {
|
||||||
|
return Err(InitialMessageFailure::new(
|
||||||
|
InitialMessageError::UnsupportedFrame,
|
||||||
|
last_frame,
|
||||||
|
));
|
||||||
|
}
|
||||||
AxumWsMessage::Text(text) => {
|
AxumWsMessage::Text(text) => {
|
||||||
let text = text.to_string();
|
let text = text.to_string();
|
||||||
let event: Value =
|
let event: Value = serde_json::from_str(&text).map_err(|_| {
|
||||||
serde_json::from_str(&text).map_err(|_| InitialMessageError::InvalidJson)?;
|
InitialMessageFailure::new(InitialMessageError::InvalidJson, last_frame)
|
||||||
validate_initial_response_create(&event)?;
|
})?;
|
||||||
|
validate_initial_response_create(&event)
|
||||||
|
.map_err(|error| InitialMessageFailure::new(error, last_frame))?;
|
||||||
return Ok((text, event));
|
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]
|
#[tokio::test]
|
||||||
async fn phase_supervisor_drops_work_at_the_absolute_connection_deadline() {
|
async fn phase_supervisor_drops_work_at_the_absolute_connection_deadline() {
|
||||||
let dropped = Arc::new(AtomicUsize::new(0));
|
let dropped = Arc::new(AtomicUsize::new(0));
|
||||||
@@ -1535,7 +1819,10 @@ mod tests {
|
|||||||
.await
|
.await
|
||||||
.expect("the initial-message deadline must cancel a stalled Pong write");
|
.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:
|
/// 验证 receive_initial_response_create_with_deadline 的绝对 deadline:
|
||||||
@@ -1583,7 +1870,10 @@ mod tests {
|
|||||||
ping_task.abort();
|
ping_task.abort();
|
||||||
|
|
||||||
assert!(
|
assert!(
|
||||||
matches!(result, Err(InitialMessageError::TimedOut)),
|
matches!(
|
||||||
|
&result,
|
||||||
|
Err(failure) if failure.error == InitialMessageError::TimedOut
|
||||||
|
),
|
||||||
"expected TimedOut after absolute deadline, got: {result:?}"
|
"expected TimedOut after absolute deadline, got: {result:?}"
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -156,6 +156,20 @@ pub(super) fn decision_reuses_bound_upstream(
|
|||||||
.unwrap_or(false)
|
.unwrap_or(false)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub(super) fn decision_bound_upstream_change_fields(
|
||||||
|
bound: &BoundResponsesConnection,
|
||||||
|
adapter: &'static dyn ResponsesWebSocketProtocolAdapter,
|
||||||
|
decision: &AiExecutionDecision,
|
||||||
|
) -> Vec<String> {
|
||||||
|
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)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use std::time::Duration;
|
use std::time::Duration;
|
||||||
|
|||||||
Reference in New Issue
Block a user