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)?
|
||||
.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<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
|
||||
@@ -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);
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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<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)]
|
||||
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 = "<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(
|
||||
@@ -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<BoundResponsesConnection> {
|
||||
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<S>(
|
||||
socket: &mut S,
|
||||
deadline_budget: std::time::Duration,
|
||||
) -> Result<(String, Value), InitialMessageError>
|
||||
) -> Result<(String, Value), InitialMessageFailure>
|
||||
where
|
||||
S: futures_util::Stream<Item = Result<AxumWsMessage, axum::Error>>
|
||||
+ futures_util::Sink<AxumWsMessage, Error = axum::Error>
|
||||
@@ -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:?}"
|
||||
);
|
||||
}
|
||||
|
||||
@@ -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<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)]
|
||||
mod tests {
|
||||
use std::time::Duration;
|
||||
|
||||
Reference in New Issue
Block a user