fix(ws): harden Responses continuation state

This commit is contained in:
ZheFox
2026-08-20 08:51:16 +08:00
parent bef282cfee
commit 654f798d25
45 changed files with 6403 additions and 493 deletions
@@ -3162,7 +3162,7 @@ async fn provider_query_execute_standard_test_candidate(
| "openai:rerank"
| "jina:rerank" => {
let Some(mut provider_request_body) =
crate::ai_serving::build_standard_request_body_with_model_directives_and_request_headers(
crate::ai_serving::build_standard_request_body_with_model_directives_and_request_headers_and_reasoning_replay_policy(
&request_body,
client_api_format,
request_model,
@@ -143,17 +143,162 @@ impl UpstreamBindingIdentity {
}
changed
}
/// Returns a versioned, one-way identity suitable for a continuation
/// registry. The digest covers every field used by `PartialEq`, including
/// effective credential and transport state, without persisting header or
/// credential values themselves.
pub(super) fn continuation_fingerprint(&self) -> [u8; 32] {
let mut digest = Sha256::new();
digest.update(b"aether-responses-websocket-binding-v1");
digest.update([match self.adapter_kind {
ResponsesWebSocketAdapter::Standard => 0,
ResponsesWebSocketAdapter::Codex => 1,
}]);
update_optional_string_digest(&mut digest, self.provider_id.as_deref());
update_optional_string_digest(&mut digest, self.endpoint_id.as_deref());
update_optional_string_digest(&mut digest, self.key_id.as_deref());
update_string_digest(&mut digest, self.upstream_url.as_str());
update_string_map_digest(&mut digest, &self.handshake_headers);
digest.update(self.credential_fingerprint);
update_proxy_digest(&mut digest, self.proxy.as_ref());
update_transport_profile_digest(&mut digest, self.transport_profile.as_ref());
digest.finalize().into()
}
pub(super) fn adapter_kind(&self) -> ResponsesWebSocketAdapter {
self.adapter_kind
}
}
/// `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 update_bytes_digest(digest: &mut Sha256, value: &[u8]) {
digest.update((value.len() as u64).to_be_bytes());
digest.update(value);
}
fn update_string_digest(digest: &mut Sha256, value: &str) {
update_bytes_digest(digest, value.as_bytes());
}
fn update_optional_string_digest(digest: &mut Sha256, value: Option<&str>) {
match value {
Some(value) => {
digest.update([1]);
update_string_digest(digest, value);
}
None => digest.update([0]),
}
}
fn update_optional_bool_digest(digest: &mut Sha256, value: Option<bool>) {
digest.update([match value {
None => 0,
Some(false) => 1,
Some(true) => 2,
}]);
}
fn update_string_map_digest(digest: &mut Sha256, values: &BTreeMap<String, String>) {
digest.update((values.len() as u64).to_be_bytes());
for (name, value) in values {
update_string_digest(digest, name);
update_string_digest(digest, value);
}
}
fn update_optional_json_digest(digest: &mut Sha256, value: Option<&serde_json::Value>) {
match value {
Some(value) => {
digest.update([1]);
update_json_digest(digest, value);
}
None => digest.update([0]),
}
}
fn update_json_digest(digest: &mut Sha256, value: &serde_json::Value) {
use serde_json::Value;
match value {
Value::Null => digest.update(b"n"),
Value::Bool(value) => digest.update(if *value { b"t" } else { b"f" }),
Value::Number(value) => {
digest.update(b"d");
update_string_digest(digest, value.to_string().as_str());
}
Value::String(value) => {
digest.update(b"s");
update_string_digest(digest, value);
}
Value::Array(values) => {
digest.update(b"[");
digest.update((values.len() as u64).to_be_bytes());
for value in values {
update_json_digest(digest, value);
}
digest.update(b"]");
}
Value::Object(values) => {
digest.update(b"{");
digest.update((values.len() as u64).to_be_bytes());
let mut keys = values.keys().collect::<Vec<_>>();
keys.sort_unstable();
for key in keys {
update_string_digest(digest, key);
update_json_digest(digest, &values[key]);
}
digest.update(b"}");
}
}
}
fn update_proxy_digest(digest: &mut Sha256, proxy: Option<&ProxySnapshot>) {
let Some(proxy) = proxy else {
digest.update([0]);
return;
};
digest.update([1]);
update_optional_bool_digest(digest, proxy.enabled);
update_optional_string_digest(digest, proxy.mode.as_deref());
update_optional_string_digest(digest, proxy.node_id.as_deref());
update_optional_string_digest(digest, proxy.label.as_deref());
update_optional_string_digest(digest, proxy.url.as_deref());
update_optional_json_digest(digest, proxy.extra.as_ref());
}
fn update_transport_profile_digest(
digest: &mut Sha256,
profile: Option<&ResolvedTransportProfile>,
) {
let Some(profile) = profile else {
digest.update([0]);
return;
};
digest.update([1]);
update_string_digest(digest, profile.profile_id.as_str());
update_string_digest(digest, profile.backend.as_str());
update_string_digest(digest, profile.http_mode.as_str());
update_string_digest(digest, profile.pool_scope.as_str());
update_optional_json_digest(digest, profile.header_fingerprint.as_ref());
update_optional_json_digest(digest, profile.extra.as_ref());
}
/// Excludes headers whose values identify one downstream request/turn rather
/// than the provider connection contract.
///
/// `x-trace-id` is generated by Aether's access middleware for every
/// downstream request. A cross-socket continuation therefore receives a new
/// value even when it revalidates to the exact same provider binding.
///
/// `x-codex-turn-metadata` likewise 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. 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")
name.eq_ignore_ascii_case(crate::constants::TRACE_ID_HEADER)
|| (adapter_kind == ResponsesWebSocketAdapter::Codex
&& name.eq_ignore_ascii_case("x-codex-turn-metadata"))
}
/// Header names that carry credentials in the provider handshake. The
@@ -386,6 +531,51 @@ mod tests {
);
}
#[test]
fn request_trace_id_does_not_change_standard_or_codex_binding_identity() {
for adapter_kind in [
ResponsesWebSocketAdapter::Standard,
ResponsesWebSocketAdapter::Codex,
] {
let adapter = resolve_responses_websocket_adapter(adapter_kind);
let mut first = decision();
if adapter_kind == ResponsesWebSocketAdapter::Codex {
first.provider_type = Some("codex".to_string());
first.report_context = Some(json!({
"codex_credential_generation": "credential-generation-1"
}));
}
first
.provider_request_headers
.insert("x-trace-id".to_string(), "request-1".to_string());
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
assert!(!first_identity.handshake_headers.contains_key("x-trace-id"));
let mut continuation = first;
continuation
.provider_request_headers
.insert("x-trace-id".to_string(), "request-2".to_string());
let continuation_identity =
UpstreamBindingIdentity::from_decision(adapter, &continuation).unwrap();
assert_eq!(first_identity, continuation_identity);
assert_eq!(
first_identity.continuation_fingerprint(),
continuation_identity.continuation_fingerprint()
);
continuation
.provider_request_headers
.insert("x-correlation-id".to_string(), "changed".to_string());
let changed_identity =
UpstreamBindingIdentity::from_decision(adapter, &continuation).unwrap();
assert_ne!(first_identity, changed_identity);
assert_eq!(
first_identity.changed_field_names(&changed_identity),
vec!["handshake_header:x-correlation-id".to_string()]
);
}
}
#[test]
fn identity_changes_when_physical_binding_changes() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
@@ -441,6 +631,25 @@ mod tests {
);
}
#[test]
fn continuation_fingerprint_changes_with_binding_and_is_not_secret_text() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
let base = decision();
let identity = UpstreamBindingIdentity::from_decision(adapter, &base).unwrap();
let mut changed = base;
changed
.provider_request_headers
.insert("X-Client".to_string(), "different-client".to_string());
let changed_identity = UpstreamBindingIdentity::from_decision(adapter, &changed).unwrap();
assert_ne!(
identity.continuation_fingerprint(),
changed_identity.continuation_fingerprint()
);
let digest = format!("{:?}", identity.continuation_fingerprint());
assert!(!digest.contains("secret"));
assert!(!digest.contains("api.example"));
}
#[test]
fn stable_key_identity_rejects_custom_static_auth_value_rotation() {
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
@@ -528,7 +737,7 @@ mod tests {
}
#[test]
fn only_codex_turn_metadata_is_excluded_from_binding_headers() {
fn unknown_codex_headers_remain_part_of_the_binding_identity() {
let codex_adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
let mut first = decision();
first.provider_type = Some("codex".to_string());
@@ -17,11 +17,14 @@ use super::ownership::{
spawn_owned_responses_websocket_plan, OwnedResponsesWebSocketDecision,
};
use super::quota::mark_active_response_retry_unsafe;
use super::redaction::redact_responses_websocket_client_event;
use super::redaction::redact_responses_websocket_client_event_with_reasoning_replay_policy;
use super::request::{
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,
build_planning_parts, changed_followup_response_create_model,
planned_request_uses_codex_responses_lite, planned_response_create_event,
prepare_responses_lite_continuation, provider_model_from_decision,
response_create_has_previous_response_id, response_create_model_or_current,
validate_response_create_previous_response_id, validate_response_create_stream_id_support,
validated_named_stream_id, ResponsesLiteStaticConfig,
};
use super::state::BoundResponsesConnection;
use super::turn::{
@@ -39,8 +42,10 @@ 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::{
close_client_socket, close_upstream_socket, send_client_message, send_gateway_error,
send_gateway_error_with_status, send_upstream_message,
send_gateway_error_with_status, send_gateway_error_with_stream_id,
send_responses_websocket_error_with_param, send_upstream_message,
};
use crate::privacy::RedactionSession;
use crate::rate_limit::FrontdoorUserRpmOutcome;
use crate::AppState;
@@ -80,20 +85,51 @@ pub(super) fn adapter_drain_ready(
))
}
fn parse_response_create_event(text: &str) -> Result<Value, &'static str> {
let event = serde_json::from_str::<Value>(text).map_err(|_| "invalid_response_create")?;
#[derive(Debug, Clone, PartialEq, Eq)]
struct ResponseCreateParseError {
code: &'static str,
stream_id: Option<String>,
}
impl ResponseCreateParseError {
fn new(code: &'static str, event: Option<&Value>) -> Self {
// A syntactically valid named lane is safe to reflect on every
// request-scoped error, even when another response.create field fails
// validation first. Invalid lane values never pass this helper and
// therefore cannot be reflected.
let stream_id = event
.and_then(validated_named_stream_id)
.map(str::to_string);
Self { code, stream_id }
}
}
fn parse_response_create_event(text: &str) -> Result<Value, ResponseCreateParseError> {
let event = serde_json::from_str::<Value>(text)
.map_err(|_| ResponseCreateParseError::new("invalid_response_create", None))?;
if event.as_object().is_none() {
return Err("invalid_response_create");
return Err(ResponseCreateParseError::new(
"invalid_response_create",
Some(&event),
));
}
if event.get("type").and_then(Value::as_str) != Some("response.create") {
return Err("expected_response_create");
return Err(ResponseCreateParseError::new(
"expected_response_create",
Some(&event),
));
}
validate_response_create_previous_response_id(&event)
.map_err(|code| ResponseCreateParseError::new(code, Some(&event)))?;
validate_response_create_stream_id_support(&event)
.map_err(|code| ResponseCreateParseError::new(code, Some(&event)))?;
Ok(event)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ContinuationConstraint {
Pinned,
UnknownResponseId,
UpstreamUnavailable,
ModelChangeUnsupported,
}
@@ -106,10 +142,14 @@ fn continuation_constraint(
event: &Value,
current_client_model: &str,
upstream_available: bool,
response_id_owned_by_connection: bool,
) -> Result<Option<ContinuationConstraint>, &'static str> {
if !response_create_has_previous_response_id(event) {
return Ok(None);
}
if !response_id_owned_by_connection {
return Ok(Some(ContinuationConstraint::UnknownResponseId));
}
if changed_followup_response_create_model(event, current_client_model)?.is_some() {
return Ok(Some(ContinuationConstraint::ModelChangeUnsupported));
}
@@ -151,11 +191,24 @@ pub(super) async fn forward_client_message(
let text = text.to_string();
let mut client_event = match parse_response_create_event(&text) {
Ok(event) => event,
Err(code) => {
send_gateway_error(
Err(error) => {
let message = match error.code {
"invalid_response_create_previous_response_id" => {
"response.create.previous_response_id must be null or a non-empty string"
}
"invalid_response_create_stream_id" => {
"response.create.stream_id must be 1-256 ASCII letters, numbers, underscores, hyphens, or periods"
}
"responses_websocket_named_stream_unsupported" => {
"Aether currently supports only the implicit default WebSocket lane; omit response.create.stream_id"
}
_ => "WebSocket client text events must be response.create JSON objects",
};
send_gateway_error_with_stream_id(
client_socket,
code,
"WebSocket client text events must be response.create JSON objects",
error.code,
message,
error.stream_id.as_deref(),
)
.await;
return RelayDisposition::Continue;
@@ -256,25 +309,116 @@ pub(super) async fn forward_client_message(
}
}
// 这一轮的 planning Parts 只构造一次(它携带 per-turn 的
// RedactionSessionSlot),并且客户端事件也只在这里脱敏一次:
// 复用已绑定 upstream 的 continuation 根本不进 planner,只靠 planner
// 内部脱敏拦不住它。之后 re-plan / continuation / 配额重试都只看脱敏
// 后的事件,上游请求体与审计 original_request_body 因此一致。
let redacted_client_event = redact_responses_websocket_client_event(
state,
&planning_parts,
&turn_control.decision,
let response_id_owned_by_connection = client_event
.get("previous_response_id")
.and_then(Value::as_str)
.is_none_or(|response_id| bound.continuation_response_ids.contains(response_id));
let is_pinned_continuation = match continuation_constraint(
&client_event,
)
.await;
let client_event = match redacted_client_event {
Ok(Some(redaction)) => {
// 这一轮的映射登记到连接上,响应帧才能在最后一跳还原回真实值。
bound.redaction_restorer.register(redaction.session);
redaction.client_event
&bound.client_model,
bound.upstream.is_some(),
response_id_owned_by_connection,
) {
Ok(Some(ContinuationConstraint::Pinned)) => true,
Ok(Some(ContinuationConstraint::UnknownResponseId)) => {
send_responses_websocket_error_with_param(
client_socket,
400,
"invalid_request_error",
"previous_response_not_found",
"The previous response is unavailable on this authenticated WebSocket connection",
"previous_response_id",
)
.await;
return RelayDisposition::Continue;
}
Ok(None) => client_event,
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;
}
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) => false,
Err(code) => {
send_gateway_error(
client_socket,
code,
"response.create.model must be a non-empty string",
)
.await;
return RelayDisposition::Continue;
}
};
// Static Responses Lite configuration belongs to the response
// chain, not to a redaction TTL bucket. Compare and strip it from
// continuation input while it is still the raw client value; only
// the incremental event is redacted below. Independent turns keep
// a raw hash so a later sentinel rotation cannot look like a tools
// or instructions change.
let raw_responses_lite_static_config = (!is_pinned_continuation)
.then(|| ResponsesLiteStaticConfig::from_response_create(&client_event));
let client_event = if is_pinned_continuation {
if let Some(static_config) = bound.responses_lite_static_config.as_ref() {
match prepare_responses_lite_continuation(&client_event, static_config) {
Ok(event) => event,
Err("responses_lite_continuation_static_config_changed") => {
send_gateway_error_with_status(
client_socket,
409,
"responses_lite_continuation_static_config_changed",
"Responses Lite tools or instructions changed; start a new response without previous_response_id",
)
.await;
return RelayDisposition::Continue;
}
Err(code) => {
send_gateway_error(
client_socket,
code,
"Gateway could not validate the Responses Lite continuation",
)
.await;
return RelayDisposition::Continue;
}
}
} else {
client_event
}
} else {
client_event
};
// 这一轮的 planning Parts 只构造一次(它携带 per-turn 的
// RedactionSessionSlot),并且客户端事件也只在这里脱敏一次。
// continuation 已经去掉继承的 static prefix,所以只 mask 新增 input;
// 后续 re-plan / continuation / 配额重试都只看脱敏后的增量事件。
let redacted_client_event =
redact_responses_websocket_client_event_with_reasoning_replay_policy(
state,
&planning_parts,
&turn_control.decision,
&client_event,
bound.body_normalization.reasoning_replay_policy(),
)
.await;
let (client_event, turn_redaction_session) = match redacted_client_event {
Ok(Some(redaction)) => (redaction.client_event, Some(redaction.session)),
Ok(None) => (client_event, None),
Err(error) => {
warn!(
event_name = "responses_websocket_followup_redaction_failed",
@@ -301,53 +445,18 @@ pub(super) async fn forward_client_message(
return RelayDisposition::Close;
}
};
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;
}
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,
"response.create.model must be a non-empty string",
)
.await;
return RelayDisposition::Continue;
}
if is_pinned_continuation {
return forward_pinned_continuation(
bound,
client_socket,
state,
context,
planning_parts,
client_event,
turn_control,
turn_redaction_session,
)
.await;
}
forward_replanned_response_create(
bound,
@@ -358,6 +467,9 @@ pub(super) async fn forward_client_message(
client_event,
requested_model,
turn_control,
raw_responses_lite_static_config
.expect("independent turns always retain their raw static config"),
turn_redaction_session,
)
.await
}
@@ -402,6 +514,7 @@ async fn forward_pinned_continuation(
planning_parts: http::request::Parts,
client_event: Value,
turn_control: ResponsesWebSocketTurnControl,
turn_redaction_session: Option<RedactionSession>,
) -> RelayDisposition {
let Some(pinned_candidate) =
ResponsesWebSocketPinnedCandidate::from_decision(&bound.decision_template)
@@ -426,8 +539,8 @@ async fn forward_pinned_continuation(
// 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.
// injected copy never replaces the already de-duplicated, redacted 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,
@@ -497,6 +610,25 @@ async fn forward_pinned_continuation(
let adapter = resolve_responses_websocket_adapter(planned.adapter);
let normalization = planned.normalization;
let decision = planned.execution;
let bound_uses_responses_lite = bound.responses_lite_static_config.is_some();
let planned_uses_responses_lite =
planned_request_uses_codex_responses_lite(&decision, &normalization);
if bound_uses_responses_lite != planned_uses_responses_lite
|| (bound_uses_responses_lite
&& !bound
.body_normalization
.has_same_responses_lite_static_contract(&normalization))
{
planned_lease.release().await;
send_gateway_error_with_status(
client_socket,
409,
"responses_lite_continuation_contract_changed",
"The Responses Lite provider contract changed; start a new response without previous_response_id",
)
.await;
return RelayDisposition::Continue;
}
let planned_provider_model = provider_model_from_decision(&decision);
let reuses_bound_upstream = decision_reuses_bound_upstream(bound, adapter, &decision);
let provider_model_changed =
@@ -530,10 +662,12 @@ async fn forward_pinned_continuation(
}
let provider_event =
match planned_response_create_event(&decision, &client_event).and_then(|event| {
serde_json::from_str::<Value>(&event)
.map_err(|_| "response_create_serialization_failed")
}) {
match planned_response_create_event(&decision, &normalization, &client_event).and_then(
|event| {
serde_json::from_str::<Value>(&event)
.map_err(|_| "response_create_serialization_failed")
},
) {
Ok(event) => event,
Err(code) => {
planned_lease.release().await;
@@ -624,11 +758,16 @@ async fn forward_pinned_continuation(
turn.mark_upstream_request_sent();
turn.set_provider_response_headers(bound.upstream_response_headers.clone());
if let Some(session) = turn_redaction_session {
bound.redaction_restorer.register(session);
}
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),
LogicalTurn::new(client_event, turn_index, logical_turn_id)
.with_provider_store(provider_event.get("store") == Some(&Value::Bool(true)))
.with_turn_control(turn_control),
turn,
);
bound.next_turn_index = bound.next_turn_index.saturating_add(1);
@@ -662,6 +801,8 @@ async fn forward_replanned_response_create(
client_event: Value,
requested_model: String,
turn_control: ResponsesWebSocketTurnControl,
raw_responses_lite_static_config: ResponsesLiteStaticConfig,
turn_redaction_session: Option<RedactionSession>,
) -> RelayDisposition {
let turn_request_id = Uuid::new_v4().to_string();
let logical_turn_id = Uuid::new_v4().to_string();
@@ -726,10 +867,12 @@ async fn forward_replanned_response_create(
let decision = planned.execution;
let reuses_bound_upstream = decision_reuses_bound_upstream(bound, adapter, &decision);
let provider_event =
match planned_response_create_event(&decision, &client_event).and_then(|event| {
serde_json::from_str::<Value>(&event)
.map_err(|_| "response_create_serialization_failed")
}) {
match planned_response_create_event(&decision, &normalization, &client_event).and_then(
|event| {
serde_json::from_str::<Value>(&event)
.map_err(|_| "response_create_serialization_failed")
},
) {
Ok(event) => event,
Err(code) => {
planned_lease.release().await;
@@ -826,18 +969,30 @@ async fn forward_replanned_response_create(
return RelayDisposition::UpstreamError("responses_websocket_send_failed");
}
// A response.create without previous_response_id starts a new chain.
// IDs from the preceding chain must not be accepted merely because
// the planner can reuse the same physical provider socket.
bound.continuation_response_ids.clear();
bound
.redaction_restorer
.start_new_chain(turn_redaction_session);
turn.mark_upstream_request_sent();
turn.set_provider_response_headers(bound.upstream_response_headers.clone());
let provider_model =
provider_model_from_decision(&decision).unwrap_or_else(|| bound.provider_model.clone());
let previous_client_model = std::mem::replace(&mut bound.client_model, requested_model);
let previous_provider_model = std::mem::replace(&mut bound.provider_model, provider_model);
let uses_responses_lite =
planned_request_uses_codex_responses_lite(&decision, &normalization);
bound.decision_template = decision;
// The re-plan keeps this upstream but resolved a new model, so later
// continuations must normalize against the new plan, not the old one.
bound.responses_lite_static_config =
uses_responses_lite.then_some(raw_responses_lite_static_config);
bound.body_normalization = normalization;
bound.turn_state.begin(
LogicalTurn::new(client_event.clone(), turn_index, logical_turn_id.clone())
.with_provider_store(provider_event.get("store") == Some(&Value::Bool(true)))
.with_turn_control(turn_control),
turn,
);
@@ -891,6 +1046,9 @@ async fn forward_replanned_response_create(
return RelayDisposition::Continue;
}
};
if replacement.responses_lite_static_config.is_some() {
replacement.responses_lite_static_config = Some(raw_responses_lite_static_config);
}
turn.mark_upstream_request_sent();
turn.set_provider_response_headers(replacement.upstream_response_headers.clone());
@@ -903,14 +1061,23 @@ async fn forward_replanned_response_create(
if let Some(mut previous_upstream) = bound.upstream.replace(replacement_upstream) {
close_upstream_socket(&mut previous_upstream, None).await;
}
// Provider connection-local response state cannot survive a physical
// rebind, and this request starts an independent chain in any case.
bound.continuation_response_ids.clear();
bound
.redaction_restorer
.start_new_chain(turn_redaction_session);
bound.adapter = replacement.adapter;
bound.client_model = replacement.client_model;
bound.provider_model = replacement.provider_model;
bound.decision_template = replacement.decision_template;
bound.body_normalization = replacement.body_normalization;
bound.responses_lite_static_config = replacement.responses_lite_static_config;
bound.binding_identity = replacement.binding_identity;
bound.turn_state.begin(
LogicalTurn::new(client_event, turn_index, logical_turn_id).with_turn_control(turn_control),
LogicalTurn::new(client_event, turn_index, logical_turn_id)
.with_provider_store(provider_event.get("store") == Some(&Value::Bool(true)))
.with_turn_control(turn_control),
turn,
);
bound.next_turn_index = bound.next_turn_index.saturating_add(1);
@@ -965,17 +1132,66 @@ mod tests {
#[test]
fn invalid_client_text_does_not_poison_the_next_response_create() {
assert_eq!(
parse_response_create_event("not-json"),
Err("invalid_response_create")
parse_response_create_event("not-json")
.expect_err("invalid JSON")
.code,
"invalid_response_create"
);
assert_eq!(
parse_response_create_event("[]"),
Err("invalid_response_create")
parse_response_create_event("[]")
.expect_err("non-object JSON")
.code,
"invalid_response_create"
);
assert_eq!(
parse_response_create_event(r#"{"type":"response.cancel"}"#),
Err("expected_response_create")
parse_response_create_event(r#"{"type":"response.cancel"}"#)
.expect_err("wrong event type")
.code,
"expected_response_create"
);
for invalid_previous_response_id in [
r#"{"type":"response.create","previous_response_id":""}"#,
r#"{"type":"response.create","previous_response_id":42}"#,
r#"{"type":"response.create","previous_response_id":{"id":"resp_1"}}"#,
] {
assert_eq!(
parse_response_create_event(invalid_previous_response_id)
.expect_err("invalid previous_response_id")
.code,
"invalid_response_create_previous_response_id"
);
}
let invalid_previous_on_named_lane = parse_response_create_event(
r#"{"type":"response.create","stream_id":"main","previous_response_id":""}"#,
)
.expect_err("invalid previous_response_id must retain a valid lane identity");
assert_eq!(
invalid_previous_on_named_lane.code,
"invalid_response_create_previous_response_id"
);
assert_eq!(
invalid_previous_on_named_lane.stream_id.as_deref(),
Some("main")
);
let named_stream = parse_response_create_event(
r#"{"type":"response.create","stream_id":"main-lane_1.test"}"#,
)
.expect_err("named streams are not implemented");
assert_eq!(
named_stream.code,
"responses_websocket_named_stream_unsupported"
);
assert_eq!(named_stream.stream_id.as_deref(), Some("main-lane_1.test"));
for invalid_stream_id in [
r#"{"type":"response.create","stream_id":null}"#,
r#"{"type":"response.create","stream_id":""}"#,
r#"{"type":"response.create","stream_id":"not/a/lane"}"#,
] {
let error = parse_response_create_event(invalid_stream_id)
.expect_err("invalid stream_id must be rejected");
assert_eq!(error.code, "invalid_response_create_stream_id");
assert_eq!(error.stream_id, None);
}
let valid = parse_response_create_event(
r#"{"type":"response.create","model":"gpt-test","store":true}"#,
@@ -993,11 +1209,11 @@ mod tests {
});
assert_eq!(
continuation_constraint(&continuation, "gpt-current", false),
continuation_constraint(&continuation, "gpt-current", false, true),
Ok(Some(ContinuationConstraint::UpstreamUnavailable))
);
assert_eq!(
continuation_constraint(&continuation, "gpt-current", true),
continuation_constraint(&continuation, "gpt-current", true, true),
Ok(Some(ContinuationConstraint::Pinned))
);
assert_eq!(
@@ -1005,6 +1221,7 @@ mod tests {
&json!({"type": "response.create", "model": "gpt-current"}),
"gpt-current",
false,
false,
),
Ok(None)
);
@@ -1019,11 +1236,25 @@ mod tests {
});
assert_eq!(
continuation_constraint(&continuation, "gpt-current", true),
continuation_constraint(&continuation, "gpt-current", true, true),
Ok(Some(ContinuationConstraint::ModelChangeUnsupported))
);
}
#[test]
fn unknown_same_socket_response_id_never_reaches_the_pinned_provider() {
let continuation = json!({
"type": "response.create",
"model": "gpt-current",
"previous_response_id": "resp_from_another_principal",
});
assert_eq!(
continuation_constraint(&continuation, "gpt-current", true, false),
Ok(Some(ContinuationConstraint::UnknownResponseId))
);
}
#[test]
fn pinned_planning_uses_the_canonical_bound_model() {
let client_event = json!({
@@ -9,6 +9,9 @@ use wreq::ws::message::Message as WreqWsMessage;
use super::adapter::ResponsesWebSocketRelayDirective;
use super::client::{adapter_drain_ready, forward_client_message, RelayDisposition};
use super::continuation::{
ResponsesWebSocketContinuationRecord, ResponsesWebSocketContinuationRegistry,
};
use super::frame::{encode_opaque_websocket_event, ParsedResponsesWebSocketFrame};
use super::lifecycle::{
await_pending_adapter_observation, finalize_active_turn, queue_turn_finalization,
@@ -26,6 +29,7 @@ use super::state::BoundResponsesConnection;
use super::turn::{
ResponsesProviderAttempt, ResponsesWebSocketTurnObservation, ResponsesWebSocketTurnOutcome,
};
use super::turn_state::LogicalTurn;
use super::upstream::{close_bound_upstream, receive_optional_upstream};
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
use crate::handlers::proxy::websocket::session::{
@@ -38,6 +42,7 @@ use crate::handlers::proxy::websocket::transport::{
use crate::AppState;
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
const CONTINUATION_REGISTRATION_TIMEOUT: Duration = Duration::from_millis(500);
/// 写客户端 socket 失败时记录的投递失败原因。刻意不说「客户端在终态前断开」:
/// 供应商的终态可能已经到达,只是最后一跳没送出去。
@@ -207,6 +212,55 @@ pub(super) async fn relay_bound_connection(
}
_ => None,
};
if parsed_upstream_frame
.as_ref()
.is_some_and(ParsedResponsesWebSocketFrame::carries_stream_id)
{
// The session currently owns only the implicit default
// lane. Reject a provider-side named identity before the
// adapter, usage observer, continuation cache, PII
// restorer, or logical-turn state can attribute an
// interleaved event to the sole default-lane attempt.
let policy = fatal_relay_policy(
FatalRelaySignal::UnexpectedUpstreamStreamId,
);
let frame = parsed_upstream_frame
.as_ref()
.expect("the stream-id guard requires a parsed frame");
warn!(
event_name = "responses_websocket_unexpected_upstream_stream_id",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
event_type = %frame.event_type_for_log(),
frame_bytes = frame.raw_text().len(),
chunked = frame.is_chunked(),
"gateway rejected a named-lane provider event on a default-lane Responses WebSocket"
);
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::upstream_receive_failed(),
)
.await;
send_responses_websocket_error(
client_socket,
policy.status_code,
"server_error",
policy.error_code,
policy.client_message,
)
.await;
close_bound_upstream(bound).await;
close_client_socket(
client_socket,
policy.close_code,
policy.close_reason,
)
.await;
break;
}
let parsed_upstream_event = parsed_upstream_frame
.as_ref()
.map(ParsedResponsesWebSocketFrame::event);
@@ -300,6 +354,46 @@ pub(super) async fn relay_bound_connection(
Some(ResponsesWebSocketTurnObservation::Terminal(outcome)) => Some(outcome),
_ => None,
};
if let Some(frame) = parsed_upstream_frame.as_ref() {
if let Some(response_id) = evicted_default_lane_continuation_response_id(
bound.turn_state.logical(),
frame,
) {
// A 4xx/5xx continuation terminal evicts the referenced
// ID from the provider's implicit default-lane cache.
// Do not keep claiming local ownership and replay it.
bound
.continuation_response_ids
.forget_connection_local(response_id);
}
if let Some(response_id) = connection_local_terminal_response_id(
bound.turn_state.logical(),
frame,
) {
// Remember every successful response on this physical
// socket, including store=false responses that exist
// only in the provider's connection-local cache.
bound
.continuation_response_ids
.remember_connection_local(response_id);
}
if let Some(registration) =
prepare_persisted_continuation_registration(bound, context, frame)
{
if let Some(response_id) =
register_persisted_continuation_before_terminal_delivery(
state,
context,
registration,
)
.await
{
bound
.continuation_response_ids
.remember_persisted(response_id.as_str());
}
}
}
if matches!(&upstream_message, WreqWsMessage::Text(_))
&& parsed_upstream_frame.is_none()
{
@@ -605,6 +699,203 @@ pub(super) async fn relay_bound_connection(
}
}
struct PendingContinuationRegistration {
user_id: String,
api_key_id: String,
response_id: String,
record: ResponsesWebSocketContinuationRecord,
}
fn prepare_persisted_continuation_registration(
bound: &BoundResponsesConnection,
context: &WebSocketRequestContext,
frame: &ParsedResponsesWebSocketFrame<'_>,
) -> Option<PendingContinuationRegistration> {
let logical = bound.turn_state.logical();
let Some(response_id) = persistable_terminal_response_id(logical, frame).map(str::to_string)
else {
return None;
};
let logical = logical.expect("a persistable terminal requires an active logical turn");
let Some(auth_context) = logical
.turn_control
.as_ref()
.and_then(|control| control.decision.auth_context.as_ref())
else {
warn!(
event_name = "responses_websocket_continuation_registration_skipped",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
reason = "missing_live_auth_context",
"gateway did not register a persisted Responses continuation"
);
return None;
};
let Some(pinned_candidate) =
crate::ai_serving::ResponsesWebSocketPinnedCandidate::from_decision(
&bound.decision_template,
)
else {
warn!(
event_name = "responses_websocket_continuation_registration_skipped",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
reason = "missing_binding_identity",
"gateway did not register a persisted Responses continuation"
);
return None;
};
let record = match ResponsesWebSocketContinuationRecord::from_binding(
pinned_candidate,
bound.client_model.as_str(),
bound.provider_model.as_str(),
&bound.binding_identity,
&bound.body_normalization,
bound.redaction_restorer.has_sessions(),
bound.responses_lite_static_config.clone(),
) {
Ok(record) => record,
Err(error) => {
warn!(
event_name = "responses_websocket_continuation_registration_skipped",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
provider_id = %bound.decision_template.provider_id.as_deref().unwrap_or("-"),
endpoint_id = %bound.decision_template.endpoint_id.as_deref().unwrap_or("-"),
key_id = %bound.decision_template.key_id.as_deref().unwrap_or("-"),
reason = error.kind(),
"gateway could not build a persisted Responses continuation record"
);
return None;
}
};
Some(PendingContinuationRegistration {
user_id: auth_context.user_id.clone(),
api_key_id: auth_context.api_key_id.clone(),
response_id,
record,
})
}
fn persistable_terminal_response_id<'a>(
logical: Option<&LogicalTurn>,
frame: &'a ParsedResponsesWebSocketFrame<'_>,
) -> Option<&'a str> {
let logical = logical?;
// `provider_store` is derived from the final framed provider event. False
// or absent is ZDR/connection-local and must never create a 24-hour KV
// record, even if the provider happens to return an ID.
if !logical.provider_store {
return None;
}
connection_local_terminal_response_id(Some(logical), frame)
}
fn connection_local_terminal_response_id<'a>(
logical: Option<&LogicalTurn>,
frame: &'a ParsedResponsesWebSocketFrame<'_>,
) -> Option<&'a str> {
logical?;
frame.continuation_response_id()
}
fn evicted_default_lane_continuation_response_id<'a>(
logical: Option<&'a LogicalTurn>,
frame: &ParsedResponsesWebSocketFrame<'_>,
) -> Option<&'a str> {
let logical = logical?;
let terminal = frame.terminal()?;
if terminal.status_code < 400 || terminal.cancelled {
return None;
}
logical
.client_event
.get("previous_response_id")
.and_then(Value::as_str)
.filter(|response_id| !response_id.trim().is_empty())
}
async fn register_persisted_continuation_before_terminal_delivery(
state: &AppState,
context: &WebSocketRequestContext,
registration: PendingContinuationRegistration,
) -> Option<String> {
let PendingContinuationRegistration {
user_id,
api_key_id,
response_id,
record,
} = registration;
let registry = ResponsesWebSocketContinuationRegistry::new(state.runtime_state.as_ref());
match tokio::time::timeout(
CONTINUATION_REGISTRATION_TIMEOUT,
registry.register(
user_id.as_str(),
api_key_id.as_str(),
response_id.as_str(),
&record,
),
)
.await
{
Ok(Ok(())) => {
debug!(
event_name = "responses_websocket_continuation_registered",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
user_id = %user_id,
api_key_id = %api_key_id,
provider_id = %record.pinned_candidate().provider_id(),
endpoint_id = %record.pinned_candidate().endpoint_id(),
key_id = %record.pinned_candidate().key_id(),
client_model = %record.client_model(),
provider_model = %record.provider_model(),
"gateway registered a persisted Responses continuation before terminal delivery"
);
Some(response_id)
}
Ok(Err(error)) => {
warn!(
event_name = "responses_websocket_continuation_registration_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
provider_id = %record.pinned_candidate().provider_id(),
endpoint_id = %record.pinned_candidate().endpoint_id(),
key_id = %record.pinned_candidate().key_id(),
reason = error.kind(),
"gateway failed to register a persisted Responses continuation"
);
None
}
Err(_) => {
warn!(
event_name = "responses_websocket_continuation_registration_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
provider_id = %record.pinned_candidate().provider_id(),
endpoint_id = %record.pinned_candidate().endpoint_id(),
key_id = %record.pinned_candidate().key_id(),
reason = "timeout",
timeout_ms = CONTINUATION_REGISTRATION_TIMEOUT.as_millis() as u64,
"gateway timed out registering a persisted Responses continuation"
);
None
}
}
}
pub(super) async fn wait_for_connection_permit_loss(
permit: Option<&aether_runtime::AdmissionPermit>,
) {
@@ -621,3 +912,79 @@ pub(super) async fn wait_for_connection_permit_loss(
}
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::{
connection_local_terminal_response_id, evicted_default_lane_continuation_response_id,
persistable_terminal_response_id, LogicalTurn, ParsedResponsesWebSocketFrame,
};
#[test]
fn cross_connection_registration_requires_explicit_store_true_and_a_success_terminal() {
let completed = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"response.completed","response":{"id":"resp_persisted"}}"#,
)
.expect("valid completed event");
let failed = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"response.failed","response":{"id":"resp_failed"}}"#,
)
.expect("valid failed event");
let omitted_or_false = LogicalTurn::new(
json!({"type": "response.create", "store": false}),
1,
"logical-local".to_string(),
);
assert_eq!(
persistable_terminal_response_id(Some(&omitted_or_false), &completed),
None,
"store=false or an omitted provider-side store must remain connection-local"
);
assert_eq!(persistable_terminal_response_id(None, &completed), None);
assert_eq!(
connection_local_terminal_response_id(Some(&omitted_or_false), &completed),
Some("resp_persisted"),
"store=false continuations remain valid on the same physical socket"
);
assert_eq!(
connection_local_terminal_response_id(None, &completed),
None
);
let persisted = LogicalTurn::new(
json!({"type": "response.create", "store": true}),
1,
"logical-persisted".to_string(),
)
.with_provider_store(true);
assert_eq!(
persistable_terminal_response_id(Some(&persisted), &completed),
Some("resp_persisted")
);
assert_eq!(
persistable_terminal_response_id(Some(&persisted), &failed),
None,
"a failed terminal must never establish cross-connection ownership"
);
let continuation = LogicalTurn::new(
json!({
"type": "response.create",
"previous_response_id": "resp_parent"
}),
2,
"logical-continuation".to_string(),
);
assert_eq!(
evicted_default_lane_continuation_response_id(Some(&continuation), &failed),
Some("resp_parent")
);
assert_eq!(
evicted_default_lane_continuation_response_id(Some(&continuation), &completed),
None
);
}
}
@@ -0,0 +1,781 @@
//! Short-lived ownership registry for Responses WebSocket continuations.
//!
//! OpenAI response IDs are opaque bearer-like references to provider state. A
//! response created on one physical provider binding must never be resumed on
//! a scheduler-selected replacement. This registry stores only non-secret
//! routing metadata and one-way contract fingerprints. Raw response IDs,
//! downstream credentials and upstream credentials are never persisted.
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use aether_runtime_state::{RuntimeLockLease, RuntimeState};
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use super::binding::UpstreamBindingIdentity;
use super::request::{ResponsesLiteStaticConfig, MAX_RESPONSES_WEBSOCKET_RESPONSE_ID_BYTES};
use crate::ai_serving::{ResponsesWebSocketBodyNormalization, ResponsesWebSocketPinnedCandidate};
use crate::orchestration::ResponsesWebSocketAdapter;
const CONTINUATION_RECORD_SCHEMA_VERSION: u16 = 1;
const CONTINUATION_KEY_PREFIX: &str = "responses_ws:continuation:v1:";
const CONTINUATION_KEY_DOMAIN: &[u8] = b"aether-responses-websocket-continuation-key-v1";
const CONTINUATION_INDEX_PREFIX: &str = "responses_ws:continuation_index:v1:";
const CONTINUATION_INDEX_DOMAIN: &[u8] = b"aether-responses-websocket-continuation-index-v1";
const CONTINUATION_LOCK_PREFIX: &str = "responses_ws:continuation_lock:v1:";
const CONTINUATION_RECORD_TTL: Duration = Duration::from_secs(24 * 60 * 60);
const CONTINUATION_INDEX_LOCK_TTL: Duration = Duration::from_secs(2);
const CONTINUATION_INDEX_LOCK_ACQUIRE_TIMEOUT: Duration = Duration::from_millis(250);
const CONTINUATION_INDEX_LOCK_INITIAL_RETRY_DELAY: Duration = Duration::from_millis(5);
const CONTINUATION_INDEX_LOCK_MAX_RETRY_DELAY: Duration = Duration::from_millis(50);
const CONTINUATION_INDEX_LOCK_OWNER: &str = "responses_ws_continuation_registry";
const MAX_CONTINUATION_RECORDS_PER_PRINCIPAL: usize = 1_024;
const MAX_SERIALIZED_CONTINUATION_RECORD_BYTES: usize = 16 * 1024;
const MAX_CONTINUATION_PRINCIPAL_BYTES: usize = 256;
const MAX_CONTINUATION_RECORD_ID_BYTES: usize = 256;
#[derive(Debug, thiserror::Error)]
pub(super) enum ResponsesWebSocketContinuationRegistryError {
#[error("invalid Responses WebSocket continuation identity: {0}")]
InvalidIdentity(&'static str),
#[error("invalid Responses WebSocket continuation record: {0}")]
InvalidRecord(&'static str),
#[error("Responses WebSocket continuation registry serialization failed")]
Serialization(#[source] serde_json::Error),
#[error("Responses WebSocket continuation registry contains a corrupt record")]
CorruptRecord(#[source] serde_json::Error),
#[error("Responses WebSocket continuation registry record is too large")]
RecordTooLarge,
#[error("Responses WebSocket continuation registry ownership conflict")]
OwnershipConflict,
#[error("Responses WebSocket continuation registry capacity lock is busy")]
CapacityLockBusy,
#[error("Responses WebSocket continuation registry storage is unavailable")]
Storage(#[source] aether_runtime_state::DataLayerError),
}
impl ResponsesWebSocketContinuationRegistryError {
pub(super) const fn kind(&self) -> &'static str {
match self {
Self::InvalidIdentity(_) => "invalid_identity",
Self::InvalidRecord(_) => "invalid_record",
Self::Serialization(_) => "serialization_failed",
Self::CorruptRecord(_) => "corrupt_record",
Self::RecordTooLarge => "record_too_large",
Self::OwnershipConflict => "ownership_conflict",
Self::CapacityLockBusy => "capacity_lock_busy",
Self::Storage(_) => "storage_unavailable",
}
}
}
/// Non-secret metadata required to prove ownership of a persisted response.
///
/// The record deliberately contains no raw response ID. Its RuntimeState key
/// is derived from the live authenticated principal plus a SHA-256 digest of
/// the opaque response ID.
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub(super) struct ResponsesWebSocketContinuationRecord {
schema_version: u16,
pinned_candidate: ResponsesWebSocketPinnedCandidate,
client_model: String,
provider_model: String,
adapter: ResponsesWebSocketAdapter,
binding_fingerprint: [u8; 32],
normalization_fingerprint: [u8; 32],
/// Server-derived replay contract for a provider whose id-less reasoning
/// state must remain byte-identical. This is trusted only because the
/// record is created after binding an authenticated provider candidate;
/// request JSON can never set it.
#[serde(default)]
deepseek_opaque_reasoning_replay: bool,
/// A prior turn stored PII sentinels whose restore mapping exists only on
/// the original downstream socket. Such a chain cannot safely resume on a
/// new socket without leaking sentinels, so lookup succeeds but bootstrap
/// rejects it before contacting the provider.
#[serde(default)]
has_connection_local_redaction: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
responses_lite_static_config: Option<ResponsesLiteStaticConfig>,
}
impl ResponsesWebSocketContinuationRecord {
pub(super) fn from_binding(
pinned_candidate: ResponsesWebSocketPinnedCandidate,
client_model: &str,
provider_model: &str,
binding: &UpstreamBindingIdentity,
normalization: &ResponsesWebSocketBodyNormalization,
has_connection_local_redaction: bool,
responses_lite_static_config: Option<ResponsesLiteStaticConfig>,
) -> Result<Self, ResponsesWebSocketContinuationRegistryError> {
let record = Self {
schema_version: CONTINUATION_RECORD_SCHEMA_VERSION,
pinned_candidate,
client_model: client_model.to_string(),
provider_model: provider_model.to_string(),
adapter: binding.adapter_kind(),
binding_fingerprint: binding.continuation_fingerprint(),
normalization_fingerprint: normalization.continuation_fingerprint(),
deepseek_opaque_reasoning_replay: matches!(
normalization.reasoning_replay_policy(),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
),
has_connection_local_redaction,
responses_lite_static_config,
};
record.validate()?;
Ok(record)
}
pub(super) fn pinned_candidate(&self) -> &ResponsesWebSocketPinnedCandidate {
&self.pinned_candidate
}
pub(super) fn client_model(&self) -> &str {
self.client_model.as_str()
}
pub(super) fn provider_model(&self) -> &str {
self.provider_model.as_str()
}
pub(super) fn adapter(&self) -> ResponsesWebSocketAdapter {
self.adapter
}
pub(super) fn responses_lite_static_config(&self) -> Option<&ResponsesLiteStaticConfig> {
self.responses_lite_static_config.as_ref()
}
pub(super) fn has_connection_local_redaction(&self) -> bool {
self.has_connection_local_redaction
}
pub(super) fn reasoning_replay_policy(
&self,
) -> crate::ai_serving::OpenAiResponsesReasoningReplayPolicy {
if self.deepseek_opaque_reasoning_replay {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
} else {
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds
}
}
pub(super) fn matches_contract(
&self,
binding: &UpstreamBindingIdentity,
normalization: &ResponsesWebSocketBodyNormalization,
) -> bool {
self.adapter == binding.adapter_kind()
&& self.binding_fingerprint == binding.continuation_fingerprint()
&& self.normalization_fingerprint == normalization.continuation_fingerprint()
}
fn validate(&self) -> Result<(), ResponsesWebSocketContinuationRegistryError> {
if self.schema_version != CONTINUATION_RECORD_SCHEMA_VERSION {
return Err(ResponsesWebSocketContinuationRegistryError::InvalidRecord(
"unsupported_schema_version",
));
}
validate_record_identifier(self.pinned_candidate.provider_id(), "invalid_provider_id")?;
validate_record_identifier(self.pinned_candidate.endpoint_id(), "invalid_endpoint_id")?;
validate_record_identifier(self.pinned_candidate.key_id(), "invalid_key_id")?;
validate_record_identifier(self.client_model.as_str(), "invalid_client_model")?;
validate_record_identifier(self.provider_model.as_str(), "invalid_provider_model")?;
Ok(())
}
}
pub(super) struct ResponsesWebSocketContinuationRegistry<'a> {
runtime_state: &'a RuntimeState,
ttl: Duration,
max_records_per_principal: usize,
}
impl<'a> ResponsesWebSocketContinuationRegistry<'a> {
pub(super) fn new(runtime_state: &'a RuntimeState) -> Self {
Self {
runtime_state,
ttl: CONTINUATION_RECORD_TTL,
max_records_per_principal: MAX_CONTINUATION_RECORDS_PER_PRINCIPAL,
}
}
#[cfg(test)]
fn with_limits(runtime_state: &'a RuntimeState, ttl: Duration, max_records: usize) -> Self {
Self {
runtime_state,
ttl,
max_records_per_principal: max_records,
}
}
pub(super) async fn register(
&self,
user_id: &str,
api_key_id: &str,
response_id: &str,
record: &ResponsesWebSocketContinuationRecord,
) -> Result<(), ResponsesWebSocketContinuationRegistryError> {
let key = continuation_registry_key(user_id, api_key_id, response_id)?;
let index_key = continuation_registry_index_key(user_id, api_key_id)?;
let lock_key = continuation_registry_lock_key(user_id, api_key_id)?;
record.validate()?;
let serialized = serde_json::to_string(record)
.map_err(ResponsesWebSocketContinuationRegistryError::Serialization)?;
if serialized.len() > MAX_SERIALIZED_CONTINUATION_RECORD_BYTES {
return Err(ResponsesWebSocketContinuationRegistryError::RecordTooLarge);
}
let lease = self.acquire_capacity_lock(&lock_key).await?;
let result = self
.register_under_capacity_lock(&key, &index_key, serialized, record)
.await;
let release_result = self.runtime_state.lock_release(&lease).await;
match (result, release_result) {
(Err(error), _) => Err(error),
(Ok(()), Ok(_)) => Ok(()),
(Ok(()), Err(error)) => {
Err(ResponsesWebSocketContinuationRegistryError::Storage(error))
}
}
}
async fn acquire_capacity_lock(
&self,
lock_key: &str,
) -> Result<RuntimeLockLease, ResponsesWebSocketContinuationRegistryError> {
let deadline = tokio::time::Instant::now() + CONTINUATION_INDEX_LOCK_ACQUIRE_TIMEOUT;
let mut retry_delay = CONTINUATION_INDEX_LOCK_INITIAL_RETRY_DELAY;
loop {
if let Some(lease) = self
.runtime_state
.lock_try_acquire(
lock_key,
CONTINUATION_INDEX_LOCK_OWNER,
CONTINUATION_INDEX_LOCK_TTL,
)
.await
.map_err(ResponsesWebSocketContinuationRegistryError::Storage)?
{
return Ok(lease);
}
let now = tokio::time::Instant::now();
if now >= deadline {
return Err(ResponsesWebSocketContinuationRegistryError::CapacityLockBusy);
}
tokio::time::sleep(retry_delay.min(deadline.saturating_duration_since(now))).await;
retry_delay = retry_delay
.saturating_mul(2)
.min(CONTINUATION_INDEX_LOCK_MAX_RETRY_DELAY);
}
}
async fn register_under_capacity_lock(
&self,
key: &str,
index_key: &str,
serialized: String,
record: &ResponsesWebSocketContinuationRecord,
) -> Result<(), ResponsesWebSocketContinuationRegistryError> {
if let Some(existing) = self
.runtime_state
.kv_get(key)
.await
.map_err(ResponsesWebSocketContinuationRegistryError::Storage)?
{
let existing = serde_json::from_str::<ResponsesWebSocketContinuationRecord>(&existing)
.map_err(ResponsesWebSocketContinuationRegistryError::CorruptRecord)?;
if existing != *record {
return Err(ResponsesWebSocketContinuationRegistryError::OwnershipConflict);
}
}
self.runtime_state
.kv_set(key, serialized, Some(self.ttl))
.await
.map_err(ResponsesWebSocketContinuationRegistryError::Storage)?;
let score = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_millis() as f64;
if let Err(error) = self.runtime_state.score_set(index_key, key, score).await {
let _ = self.runtime_state.kv_delete(key).await;
return Err(ResponsesWebSocketContinuationRegistryError::Storage(error));
}
if let Err(error) = self.runtime_state.key_expire(index_key, self.ttl).await {
let _ = self.runtime_state.score_remove(index_key, key).await;
let _ = self.runtime_state.kv_delete(key).await;
return Err(ResponsesWebSocketContinuationRegistryError::Storage(error));
}
let members = self
.runtime_state
.score_range_by_min(index_key, f64::NEG_INFINITY)
.await
.map_err(ResponsesWebSocketContinuationRegistryError::Storage)?;
let overflow = members.len().saturating_sub(self.max_records_per_principal);
for oldest_key in members.into_iter().take(overflow) {
self.runtime_state
.kv_delete(&oldest_key)
.await
.map_err(ResponsesWebSocketContinuationRegistryError::Storage)?;
self.runtime_state
.score_remove(index_key, &oldest_key)
.await
.map_err(ResponsesWebSocketContinuationRegistryError::Storage)?;
}
Ok(())
}
pub(super) async fn lookup(
&self,
user_id: &str,
api_key_id: &str,
response_id: &str,
) -> Result<
Option<ResponsesWebSocketContinuationRecord>,
ResponsesWebSocketContinuationRegistryError,
> {
let key = continuation_registry_key(user_id, api_key_id, response_id)?;
let Some(serialized) = self
.runtime_state
.kv_get(&key)
.await
.map_err(ResponsesWebSocketContinuationRegistryError::Storage)?
else {
return Ok(None);
};
let record = serde_json::from_str::<ResponsesWebSocketContinuationRecord>(&serialized)
.map_err(ResponsesWebSocketContinuationRegistryError::CorruptRecord)?;
record.validate()?;
Ok(Some(record))
}
}
fn validate_record_identifier(
value: &str,
error: &'static str,
) -> Result<(), ResponsesWebSocketContinuationRegistryError> {
if value.trim().is_empty() || value.len() > MAX_CONTINUATION_RECORD_ID_BYTES {
return Err(ResponsesWebSocketContinuationRegistryError::InvalidRecord(
error,
));
}
Ok(())
}
fn validate_key_component(
value: &str,
max_bytes: usize,
error: &'static str,
) -> Result<(), ResponsesWebSocketContinuationRegistryError> {
if value.is_empty() || value.len() > max_bytes {
return Err(ResponsesWebSocketContinuationRegistryError::InvalidIdentity(error));
}
Ok(())
}
fn continuation_registry_key(
user_id: &str,
api_key_id: &str,
response_id: &str,
) -> Result<String, ResponsesWebSocketContinuationRegistryError> {
validate_key_component(user_id, MAX_CONTINUATION_PRINCIPAL_BYTES, "invalid_user_id")?;
validate_key_component(
api_key_id,
MAX_CONTINUATION_PRINCIPAL_BYTES,
"invalid_api_key_id",
)?;
validate_key_component(
response_id,
MAX_RESPONSES_WEBSOCKET_RESPONSE_ID_BYTES,
"invalid_response_id",
)?;
let digest = digest_key_components(
CONTINUATION_KEY_DOMAIN,
[
user_id.as_bytes(),
api_key_id.as_bytes(),
response_id.as_bytes(),
],
);
Ok(format!("{CONTINUATION_KEY_PREFIX}{digest}"))
}
fn continuation_registry_index_key(
user_id: &str,
api_key_id: &str,
) -> Result<String, ResponsesWebSocketContinuationRegistryError> {
validate_key_component(user_id, MAX_CONTINUATION_PRINCIPAL_BYTES, "invalid_user_id")?;
validate_key_component(
api_key_id,
MAX_CONTINUATION_PRINCIPAL_BYTES,
"invalid_api_key_id",
)?;
let digest = digest_key_components(
CONTINUATION_INDEX_DOMAIN,
[user_id.as_bytes(), api_key_id.as_bytes()],
);
Ok(format!("{CONTINUATION_INDEX_PREFIX}{digest}"))
}
fn continuation_registry_lock_key(
user_id: &str,
api_key_id: &str,
) -> Result<String, ResponsesWebSocketContinuationRegistryError> {
let index = continuation_registry_index_key(user_id, api_key_id)?;
Ok(format!(
"{CONTINUATION_LOCK_PREFIX}{}",
index
.strip_prefix(CONTINUATION_INDEX_PREFIX)
.unwrap_or(index.as_str())
))
}
fn digest_key_components<const N: usize>(domain: &[u8], components: [&[u8]; N]) -> String {
let mut digest = Sha256::new();
digest.update(domain);
for component in components {
digest.update((component.len() as u64).to_be_bytes());
digest.update(component);
}
format!("{:x}", digest.finalize())
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use aether_runtime_state::{MemoryRuntimeStateConfig, RuntimeState};
use serde_json::json;
use super::*;
const USER_ID: &str = "user-live-0198";
const API_KEY_ID: &str = "api-key-live-0198";
const RESPONSE_ID: &str = "resp_opaque_super_secret_reference";
fn runtime_state() -> RuntimeState {
RuntimeState::memory(MemoryRuntimeStateConfig::default())
}
fn record() -> ResponsesWebSocketContinuationRecord {
ResponsesWebSocketContinuationRecord {
schema_version: CONTINUATION_RECORD_SCHEMA_VERSION,
pinned_candidate: ResponsesWebSocketPinnedCandidate::new(
"provider-1",
"endpoint-1",
"key-1",
)
.expect("candidate"),
client_model: "public-model".to_string(),
provider_model: "provider-model".to_string(),
adapter: ResponsesWebSocketAdapter::Codex,
binding_fingerprint: [7; 32],
normalization_fingerprint: [9; 32],
deepseek_opaque_reasoning_replay: false,
has_connection_local_redaction: false,
responses_lite_static_config: Some(ResponsesLiteStaticConfig::from_response_create(
&json!({
"tools": [{"type": "function", "name": "lookup"}],
"instructions": "Do not persist this plaintext"
}),
)),
}
}
#[tokio::test]
async fn same_principal_gets_the_exact_pinned_candidate_and_other_principals_miss() {
let runtime = runtime_state();
let registry = ResponsesWebSocketContinuationRegistry::new(&runtime);
let expected = record();
registry
.register(USER_ID, API_KEY_ID, RESPONSE_ID, &expected)
.await
.expect("register");
let found = registry
.lookup(USER_ID, API_KEY_ID, RESPONSE_ID)
.await
.expect("lookup")
.expect("same authenticated principal must find its record");
assert_eq!(found, expected);
assert_eq!(found.pinned_candidate(), expected.pinned_candidate());
for (user_id, api_key_id, response_id) in [
("other-user", API_KEY_ID, RESPONSE_ID),
(USER_ID, "other-api-key", RESPONSE_ID),
(USER_ID, API_KEY_ID, "resp_other"),
] {
assert_eq!(
registry
.lookup(user_id, api_key_id, response_id)
.await
.expect("isolated lookup"),
None
);
}
}
#[tokio::test]
async fn corrupt_or_unsupported_records_fail_closed() {
let runtime = runtime_state();
let registry = ResponsesWebSocketContinuationRegistry::new(&runtime);
let key = continuation_registry_key(USER_ID, API_KEY_ID, RESPONSE_ID).expect("key");
runtime
.kv_set(&key, "not-json", Some(Duration::from_secs(60)))
.await
.expect("seed corrupt record");
assert!(matches!(
registry.lookup(USER_ID, API_KEY_ID, RESPONSE_ID).await,
Err(ResponsesWebSocketContinuationRegistryError::CorruptRecord(
_
))
));
let mut unsupported = serde_json::to_value(record()).expect("serialize record");
unsupported["schema_version"] = json!(CONTINUATION_RECORD_SCHEMA_VERSION + 1);
runtime
.kv_set(&key, unsupported.to_string(), Some(Duration::from_secs(60)))
.await
.expect("seed unsupported record");
assert!(matches!(
registry.lookup(USER_ID, API_KEY_ID, RESPONSE_ID).await,
Err(ResponsesWebSocketContinuationRegistryError::InvalidRecord(
"unsupported_schema_version"
))
));
}
#[tokio::test]
async fn expired_records_are_not_returned() {
let runtime = runtime_state();
let registry = ResponsesWebSocketContinuationRegistry::with_limits(
&runtime,
Duration::from_millis(5),
MAX_CONTINUATION_RECORDS_PER_PRINCIPAL,
);
registry
.register(USER_ID, API_KEY_ID, RESPONSE_ID, &record())
.await
.expect("register");
tokio::time::sleep(Duration::from_millis(20)).await;
assert_eq!(
registry
.lookup(USER_ID, API_KEY_ID, RESPONSE_ID)
.await
.expect("lookup"),
None
);
}
#[tokio::test]
async fn registry_evicts_the_oldest_record_per_authenticated_principal() {
let runtime = runtime_state();
let registry = ResponsesWebSocketContinuationRegistry::with_limits(
&runtime,
Duration::from_secs(60),
2,
);
for response_id in ["resp_oldest", "resp_middle", "resp_newest"] {
registry
.register(USER_ID, API_KEY_ID, response_id, &record())
.await
.expect("register");
tokio::time::sleep(Duration::from_millis(2)).await;
}
assert_eq!(
registry
.lookup(USER_ID, API_KEY_ID, "resp_oldest")
.await
.expect("lookup"),
None
);
for response_id in ["resp_middle", "resp_newest"] {
assert_eq!(
registry
.lookup(USER_ID, API_KEY_ID, response_id)
.await
.expect("lookup"),
Some(record())
);
}
let index_key = continuation_registry_index_key(USER_ID, API_KEY_ID).expect("index key");
assert_eq!(runtime.score_len(&index_key).await.expect("index len"), 2);
}
#[tokio::test]
async fn registry_retries_a_briefly_contended_principal_capacity_lock() {
let runtime = runtime_state();
let registry = ResponsesWebSocketContinuationRegistry::new(&runtime);
let lock_key = continuation_registry_lock_key(USER_ID, API_KEY_ID).expect("lock key");
let held = runtime
.lock_try_acquire(
&lock_key,
"continuation-registry-contention-test",
CONTINUATION_INDEX_LOCK_TTL,
)
.await
.expect("acquire test lock")
.expect("test lock should be uncontended");
let release_held_lock = async {
tokio::time::sleep(Duration::from_millis(20)).await;
assert!(runtime
.lock_release(&held)
.await
.expect("release test lock"));
};
let expected = record();
let (registration, ()) = tokio::join!(
registry.register(USER_ID, API_KEY_ID, RESPONSE_ID, &expected),
release_held_lock,
);
registration.expect("registration should retry until the short contention clears");
assert!(registry
.lookup(USER_ID, API_KEY_ID, RESPONSE_ID)
.await
.expect("lookup")
.is_some());
}
#[tokio::test]
async fn same_response_id_cannot_be_rebound_to_a_different_owner_record() {
let runtime = runtime_state();
let registry = ResponsesWebSocketContinuationRegistry::new(&runtime);
let original = record();
registry
.register(USER_ID, API_KEY_ID, RESPONSE_ID, &original)
.await
.expect("register");
let mut replacement = original.clone();
replacement.provider_model = "different-provider-model".to_string();
assert!(matches!(
registry
.register(USER_ID, API_KEY_ID, RESPONSE_ID, &replacement)
.await,
Err(ResponsesWebSocketContinuationRegistryError::OwnershipConflict)
));
assert_eq!(
registry
.lookup(USER_ID, API_KEY_ID, RESPONSE_ID)
.await
.expect("lookup"),
Some(original)
);
}
#[test]
fn registry_key_is_length_delimited_hashed_and_contains_no_plaintext() {
let key = continuation_registry_key(USER_ID, API_KEY_ID, RESPONSE_ID).expect("key");
assert!(key.starts_with(CONTINUATION_KEY_PREFIX));
assert_eq!(key.len(), CONTINUATION_KEY_PREFIX.len() + 64);
for secret in [USER_ID, API_KEY_ID, RESPONSE_ID, "super_secret_reference"] {
assert!(!key.contains(secret));
}
let index_key = continuation_registry_index_key(USER_ID, API_KEY_ID).expect("index key");
for secret in [USER_ID, API_KEY_ID] {
assert!(!index_key.contains(secret));
}
assert_ne!(
continuation_registry_key("ab", "c", "d").expect("key"),
continuation_registry_key("a", "bc", "d").expect("key")
);
}
#[test]
fn invalid_or_oversized_identity_never_produces_a_cache_key() {
for (user_id, api_key_id, response_id) in [
("", API_KEY_ID, RESPONSE_ID),
(USER_ID, "", RESPONSE_ID),
(USER_ID, API_KEY_ID, ""),
] {
assert!(continuation_registry_key(user_id, api_key_id, response_id).is_err());
}
let oversized = "x".repeat(MAX_RESPONSES_WEBSOCKET_RESPONSE_ID_BYTES + 1);
assert!(continuation_registry_key(USER_ID, API_KEY_ID, &oversized).is_err());
}
#[test]
fn serialized_record_contains_only_digests_for_static_and_binding_state() {
let serialized = serde_json::to_string(&record()).expect("serialize");
for plaintext in [
RESPONSE_ID,
"Do not persist this plaintext",
"lookup",
"upstream-oauth-token",
] {
assert!(!serialized.contains(plaintext));
}
let decoded: ResponsesWebSocketContinuationRecord =
serde_json::from_str(&serialized).expect("deserialize");
assert_eq!(decoded, record());
}
#[test]
fn serialized_record_preserves_only_the_server_derived_reasoning_replay_policy_bit() {
let mut expected = record();
expected.deepseek_opaque_reasoning_replay = true;
let serialized = serde_json::to_string(&expected).expect("serialize");
let decoded: ResponsesWebSocketContinuationRecord =
serde_json::from_str(&serialized).expect("deserialize");
assert_eq!(
decoded.reasoning_replay_policy(),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
);
let mut legacy = serde_json::to_value(&expected).expect("serialize legacy fixture");
legacy
.as_object_mut()
.expect("record is an object")
.remove("deepseek_opaque_reasoning_replay");
let legacy: ResponsesWebSocketContinuationRecord =
serde_json::from_value(legacy).expect("legacy record should remain readable");
assert_eq!(
legacy.reasoning_replay_policy(),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds,
"old records must fail closed instead of inferring trust from request shape"
);
}
#[test]
fn record_derives_deepseek_replay_policy_from_the_bound_normalization() {
let decision: crate::ai_serving::AiExecutionDecision = serde_json::from_value(json!({
"action": "local",
"provider_id": "provider-1",
"endpoint_id": "endpoint-1",
"key_id": "key-1",
"upstream_url": "https://api.deepseek.com/v1/responses",
"provider_request_headers": {}
}))
.expect("minimal decision");
let adapter = super::super::adapter::resolve_responses_websocket_adapter(
ResponsesWebSocketAdapter::Standard,
);
let binding =
UpstreamBindingIdentity::from_decision(adapter, &decision).expect("binding identity");
let normalization = ResponsesWebSocketBodyNormalization::for_tests("deepseek-reasoner")
.with_reasoning_replay_policy_for_tests(
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque,
);
let record = ResponsesWebSocketContinuationRecord::from_binding(
ResponsesWebSocketPinnedCandidate::new("provider-1", "endpoint-1", "key-1")
.expect("pinned candidate"),
"public-model",
"deepseek-reasoner",
&binding,
&normalization,
false,
None,
)
.expect("continuation record");
assert_eq!(
record.reasoning_replay_policy(),
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque
);
}
}
@@ -7,6 +7,8 @@
use serde_json::Value;
use super::request::MAX_RESPONSES_WEBSOCKET_RESPONSE_ID_BYTES;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) struct ResponsesWebSocketFrameTerminal {
pub(super) status_code: u16,
@@ -89,6 +91,23 @@ impl<'a> ParsedResponsesWebSocketFrame<'a> {
&self.event
}
/// Whether the provider attached a named-lane identity to this frame.
///
/// Aether's current relay owns only the implicit default lane. OpenAI's
/// named-lane contract places `stream_id` directly on each public event;
/// Codex may additionally batch those events under `chunks`. Inspect both
/// wire positions without looking through arbitrary response payloads so
/// a model-produced field named `stream_id` is not mistaken for transport
/// framing.
pub(super) fn carries_stream_id(&self) -> bool {
self.event.get("stream_id").is_some()
|| self
.event
.get("chunks")
.and_then(Value::as_array)
.is_some_and(|chunks| chunks.iter().any(|event| event.get("stream_id").is_some()))
}
pub(super) fn event_type(&self) -> Option<&str> {
self.event_type.as_deref()
}
@@ -109,6 +128,30 @@ impl<'a> ParsedResponsesWebSocketFrame<'a> {
self.terminal
}
/// Returns the opaque response ID only for a successfully observed,
/// provider-persistable terminal. Failed, cancelled, malformed and
/// oversized terminal IDs must never establish continuation ownership.
pub(super) fn continuation_response_id(&self) -> Option<&str> {
let terminal = self.terminal?;
if terminal.status_code >= 400 || terminal.cancelled {
return None;
}
let event = self.terminal_event.as_ref()?;
if !matches!(
event_type_of(event),
Some("response.completed" | "response.done" | "response.incomplete")
) {
return None;
}
let response_id = event
.pointer("/response/id")
.or_else(|| event.get("response_id"))
.and_then(Value::as_str)?;
(!response_id.trim().is_empty()
&& response_id.len() <= MAX_RESPONSES_WEBSOCKET_RESPONSE_ID_BYTES)
.then_some(response_id)
}
/// Return a bounded label suitable for structured logs. Event payloads
/// are never inserted directly into a log field.
pub(super) fn event_type_for_log(&self) -> String {
@@ -206,7 +249,7 @@ fn responses_incomplete_default_status(event: &Value) -> u16 {
fn terminal_for_event(event: &Value) -> Option<ResponsesWebSocketFrameTerminal> {
match event_type_of(event).unwrap_or_default() {
"response.completed" => Some(ResponsesWebSocketFrameTerminal {
"response.completed" | "response.done" => Some(ResponsesWebSocketFrameTerminal {
status_code: websocket_event_status_code(event, 200),
cancelled: false,
}),
@@ -307,6 +350,29 @@ mod tests {
assert_eq!(frame.event_type_for_log(), "response.in_progress");
}
#[test]
fn continuation_id_requires_a_successful_persistable_terminal() {
for event_type in ["response.completed", "response.done"] {
let raw = format!(r#"{{"type":"{event_type}","response":{{"id":"resp_persisted"}}}}"#);
let frame = ParsedResponsesWebSocketFrame::parse(&raw).expect("valid terminal");
assert_eq!(frame.continuation_response_id(), Some("resp_persisted"));
}
let incomplete = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"response.incomplete","response":{"id":"resp_partial","incomplete_details":{"reason":"max_output_tokens"}}}"#,
)
.expect("valid incomplete terminal");
assert_eq!(incomplete.continuation_response_id(), Some("resp_partial"));
for raw in [
r#"{"type":"response.failed","response":{"id":"resp_failed"}}"#,
r#"{"type":"response.cancelled","response":{"id":"resp_cancelled"}}"#,
r#"{"type":"response.completed","status_code":500,"response":{"id":"resp_error"}}"#,
] {
let frame = ParsedResponsesWebSocketFrame::parse(raw).expect("valid terminal");
assert_eq!(frame.continuation_response_id(), None);
}
}
#[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}";
@@ -338,6 +404,36 @@ mod tests {
assert_eq!(round_trip["response"]["future_usage"]["novel_tokens"], 7);
}
#[test]
fn detects_a_top_level_named_lane_identity() {
let frame = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"response.in_progress","stream_id":"planner","response":{"id":"resp_1"}}"#,
)
.expect("valid named-lane event");
assert!(frame.carries_stream_id());
}
#[test]
fn detects_a_named_lane_identity_inside_a_batch() {
let frame = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"codex.response.metadata","chunks":[{"type":"codex.rate_limits"},{"type":"response.completed","stream_id":"planner","response":{"id":"resp_1"}}]}"#,
)
.expect("valid named-lane batch");
assert!(frame.carries_stream_id());
}
#[test]
fn response_payload_fields_are_not_mistaken_for_lane_framing() {
let frame = ParsedResponsesWebSocketFrame::parse(
r#"{"type":"response.completed","response":{"id":"resp_1","metadata":{"stream_id":"model-owned"}}}"#,
)
.expect("valid default-lane event");
assert!(!frame.carries_stream_id());
}
#[test]
fn classifies_terminal_status_and_cancellation() {
let completed = ParsedResponsesWebSocketFrame::parse(
@@ -12,6 +12,7 @@ mod admission;
mod binding;
mod client;
mod connection;
mod continuation;
mod control;
mod frame;
mod lifecycle;
@@ -13,7 +13,9 @@ 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::request::{
build_planning_parts, planned_response_create_event, ResponsesLiteStaticConfig,
};
use super::state::BoundResponsesConnection;
use super::turn::{prepare_responses_websocket_turn_decision, ResponsesWebSocketTurnOutcome};
use super::upstream::{bind_responses_upstream, close_bound_upstream};
@@ -102,6 +104,14 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
context: &WebSocketRequestContext,
_previous_settled: PreviousAttemptSettled,
) -> bool {
// `LogicalTurn::client_event` is intentionally redacted before it is
// retained for replay. The binding, however, keeps the hash of the raw
// client-side Responses Lite tools/instructions so a later continuation
// can compare the client's plaintext configuration before redaction.
// Preserve that chain identity across this transparent rebind instead of
// replacing it with the hash `bind_responses_upstream` derives from the
// redacted replay event.
let responses_lite_static_config = bound.responses_lite_static_config.clone();
let Some(active) = bound.turn_state.logical_mut() else {
return false;
};
@@ -213,12 +223,14 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
);
return false;
}
let provider_event = match planned_response_create_event(&decision, &client_event).and_then(
|event| {
serde_json::from_str::<Value>(&event)
.map_err(|_| "response_create_serialization_failed")
},
) {
let provider_event = match planned_response_create_event(
&decision,
&normalization,
&client_event,
)
.and_then(|event| {
serde_json::from_str::<Value>(&event).map_err(|_| "response_create_serialization_failed")
}) {
Ok(event) => event,
Err(code) => {
planned_lease.release().await;
@@ -234,6 +246,7 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
return false;
}
};
let replacement_provider_store = provider_event.get("store") == Some(&Value::Bool(true));
let turn_decision = prepare_responses_websocket_turn_decision(
&decision,
turn_request_id,
@@ -309,12 +322,20 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
if let Some(mut previous_upstream) = bound.upstream.replace(replacement_upstream) {
close_upstream_socket(&mut previous_upstream, None).await;
}
// The replacement socket has no access to response IDs cached only on
// the exhausted physical connection. Continuation turns are never quota
// replayed, so clearing here cannot discard the active turn's parent.
bound.continuation_response_ids.clear();
let previous_key_id = bound.decision_template.key_id.clone();
bound.adapter = replacement.adapter;
bound.client_model = replacement.client_model;
bound.provider_model = replacement.provider_model;
bound.decision_template = replacement.decision_template;
bound.body_normalization = replacement.body_normalization;
bound.responses_lite_static_config = responses_lite_static_config_after_rebind(
responses_lite_static_config,
replacement.responses_lite_static_config,
);
bound.binding_identity = replacement.binding_identity;
// 同一个 logical turn 的下一个 attempt 就位。状态不符时把 attempt 交回
// drop guard 结算并让调用方走「透明重试失败」分支,不静默丢弃一条已经写了
@@ -323,6 +344,9 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
drop(orphan);
return false;
}
if let Some(logical) = bound.turn_state.logical_mut() {
logical.provider_store = replacement_provider_store;
}
bound.upstream_response_headers = replacement.upstream_response_headers;
bound.pending_adapter_drain = None;
debug!(
@@ -341,6 +365,13 @@ pub(super) async fn retry_active_turn_after_quota_exhaustion(
true
}
fn responses_lite_static_config_after_rebind(
previous: Option<ResponsesLiteStaticConfig>,
replacement: Option<ResponsesLiteStaticConfig>,
) -> Option<ResponsesLiteStaticConfig> {
replacement.map(|replacement| previous.unwrap_or(replacement))
}
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")
@@ -375,3 +406,97 @@ pub(super) fn mark_active_response_retry_unsafe(
active.mark_retry_unsafe(reason);
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::responses_lite_static_config_after_rebind;
use crate::handlers::proxy::websocket::responses::request::{
prepare_responses_lite_continuation, ResponsesLiteStaticConfig,
};
#[test]
fn quota_rebind_preserves_the_raw_responses_lite_static_hash() {
let raw = json!({
"type": "response.create",
"model": "gpt-5.6-sol",
"instructions": "Contact [email protected] before using the tool",
"tools": [{
"type": "function",
"name": "lookup",
"description": "Look up [email protected]",
"parameters": {"type": "object", "properties": {}}
}],
"input": [{"role": "user", "content": "hello"}]
});
let redacted_replay = json!({
"type": "response.create",
"model": "gpt-5.6-sol",
"instructions": "Contact <EMAIL_2> before using the tool",
"tools": [{
"type": "function",
"name": "lookup",
"description": "Look up <EMAIL_1>",
"parameters": {"type": "object", "properties": {}}
}],
"input": [{"role": "user", "content": "hello"}]
});
let raw_static_config = ResponsesLiteStaticConfig::from_response_create(&raw);
let redacted_static_config =
ResponsesLiteStaticConfig::from_response_create(&redacted_replay);
assert_ne!(raw_static_config, redacted_static_config);
let rebound_static_config = responses_lite_static_config_after_rebind(
Some(raw_static_config.clone()),
Some(redacted_static_config),
)
.expect("the replacement still uses Responses Lite");
assert_eq!(rebound_static_config, raw_static_config);
let continuation = json!({
"type": "response.create",
"model": "gpt-5.6-sol",
"previous_response_id": "resp_after_quota_retry",
"instructions": "Contact [email protected] before using the tool",
"tools": [{
"type": "function",
"name": "lookup",
"description": "Look up [email protected]",
"parameters": {"type": "object", "properties": {}}
}],
"input": [{
"type": "function_call_output",
"call_id": "call_1",
"output": "ok"
}]
});
let prepared = prepare_responses_lite_continuation(&continuation, &rebound_static_config)
.expect("an unchanged plaintext continuation must survive a quota retry");
assert!(prepared.get("tools").is_none());
assert!(prepared.get("instructions").is_none());
}
#[test]
fn quota_rebind_still_tracks_the_replacement_contract() {
let raw = ResponsesLiteStaticConfig::from_response_create(&json!({
"type": "response.create",
"tools": [{"type": "function", "name": "lookup", "parameters": {}}]
}));
let replacement = ResponsesLiteStaticConfig::from_response_create(&json!({
"type": "response.create",
"instructions": "replacement"
}));
assert_eq!(
responses_lite_static_config_after_rebind(Some(raw), None),
None,
"a non-Lite replacement must clear the Lite chain marker"
);
assert_eq!(
responses_lite_static_config_after_rebind(None, Some(replacement.clone())),
Some(replacement),
"a newly selected Lite contract has no earlier raw hash to preserve"
);
}
}
@@ -29,12 +29,14 @@
//! ("你刚才给我的邮箱是……"),本轮 session 里没有这条映射,占位符就漏给客户端。
//! HTTP 不会漏,是因为它每次都重发整段历史,重新 mask 同一个值会派生出同一个
//! sentinel(HMAC over 规则 + bucket + 值),所以映射天然齐备。
//! * 挂在连接上(当前实现):每轮仍然各自 mask、各自持有独立 session
//! * 挂在当前 response chain 上(当前实现):每轮仍然各自 mask、各自持有独立 session
//! (per-turn 语义不变),连接只是把最近若干轮的 session 留下来一起参与还原,
//! 凑出的映射集合正好等于「等价 HTTP 请求会拥有的那一份」。
//!
//! 选后者。代价是每帧最多对 [`MAX_RETAINED_TURN_REDACTION_SESSIONS`] 个 session
//! 各扫一遍,以及这些 session 的映射会驻留到连接结束;用有界 FIFO 兜住上限。
//! 选后者。省略(或置空)`previous_response_id` 会开始一条独立 response chain,
//! 此时必须丢弃旧链的映射,否则新链里偶然出现的旧 sentinel 会被还原成旧链 PII。
//! 代价是每帧最多对 [`MAX_RETAINED_TURN_REDACTION_SESSIONS`] 个 session 各扫一遍,
//! 以及这些 session 的映射会驻留到当前链结束;用有界 FIFO 兜住上限。
//! 窗口不够用或每帧成本变高时,正确的下一步是在 `privacy` 侧提供跨 session 的
//! 合并匹配器,而不是把这个窗口调大。
@@ -89,6 +91,27 @@ pub(super) async fn redact_responses_websocket_client_event(
parts: &http::request::Parts,
control_decision: &GatewayControlDecision,
client_event: &Value,
) -> Result<Option<ResponsesWebSocketTurnRedaction>, GatewayError> {
redact_responses_websocket_client_event_with_reasoning_replay_policy(
state,
parts,
control_decision,
client_event,
crate::ai_serving::OpenAiResponsesReasoningReplayPolicy::OpenAiItemIds,
)
.await
}
/// Variant used only after the gateway has selected and authenticated the
/// provider binding. The replay policy comes from that trusted binding, never
/// from client JSON, so a forged reasoning-item shape cannot opt itself into
/// byte-opaque PII handling.
pub(super) async fn redact_responses_websocket_client_event_with_reasoning_replay_policy(
state: &AppState,
parts: &http::request::Parts,
control_decision: &GatewayControlDecision,
client_event: &Value,
reasoning_replay_policy: crate::ai_serving::OpenAiResponsesReasoningReplayPolicy,
) -> Result<Option<ResponsesWebSocketTurnRedaction>, GatewayError> {
let Some(auth_context) =
resolve_local_decision_execution_runtime_auth_context(control_decision)
@@ -101,6 +124,7 @@ pub(super) async fn redact_responses_websocket_client_event(
client_event,
&auth_context,
RESPONSES_WEBSOCKET_CLIENT_API_FORMAT,
reasoning_replay_policy,
WEBSOCKET_TURN_REDACTION_CANDIDATE_ID,
)
.await?;
@@ -126,17 +150,22 @@ pub(super) async fn redact_responses_websocket_client_event(
}))
}
/// 一条连接上「我们 mask 过哪些映射」的留存集合,供响应侧还原使用。
/// 当前 response chain 上「我们 mask 过哪些映射」的留存集合,供响应侧还原使用。
///
/// 每轮一个独立 session(per-turn mask 语义不变),连接按 FIFO 留最近
/// [`MAX_RETAINED_TURN_REDACTION_SESSIONS`] 轮。上游重绑不清空:客户端仍在同一段
/// 对话里,旧占位符可能随重发的输入再次出现。
/// 每轮一个独立 session(per-turn mask 语义不变),当前链按 FIFO 留最近
/// [`MAX_RETAINED_TURN_REDACTION_SESSIONS`] 轮。物理上游重绑本身不决定生命周期;
/// `previous_response_id` 决定是否延续旧链。独立请求成功发出时由调用方通过
/// [`Self::start_new_chain`] 原子替换为新链的首轮 session。
#[derive(Default)]
pub(super) struct ResponsesWebSocketRedactionRestorer {
sessions: VecDeque<RedactionSession>,
}
impl ResponsesWebSocketRedactionRestorer {
pub(super) fn has_sessions(&self) -> bool {
!self.sessions.is_empty()
}
/// 登记这一轮的 mask session。
pub(super) fn register(&mut self, session: RedactionSession) {
if session.mapping_count() == 0 {
@@ -148,6 +177,18 @@ impl ResponsesWebSocketRedactionRestorer {
}
}
/// Commits a successfully started independent response chain.
///
/// Keep this transition next to the successful upstream send/bind. A
/// rejected independent request has not replaced the active chain and
/// therefore must not discard the old chain's restore mappings.
pub(super) fn start_new_chain(&mut self, session: Option<RedactionSession>) {
self.sessions.clear();
if let Some(session) = session {
self.register(session);
}
}
/// 把一帧 provider 事件里的占位符换回真实值,返回要发给客户端的帧文本。
///
/// `None` 表示这一帧没有任何东西要还原,调用方必须原样转发上游字节:未启用
@@ -497,8 +538,9 @@ mod tests {
seed_report_context_with_raw_pii(),
);
// 首轮实际发上游的事件由 decision.provider_request_body 派生。
let normalization = ResponsesWebSocketBodyNormalization::for_tests("provider-model");
let provider_event: Value = serde_json::from_str(
&planned_response_create_event(&template, &effective_event)
&planned_response_create_event(&template, &normalization, &effective_event)
.expect("first provider event should serialize"),
)
.expect("first provider event should parse");
@@ -602,8 +644,9 @@ mod tests {
provider_body_from(&active.client_event),
seed_report_context_with_raw_pii(),
);
let normalization = ResponsesWebSocketBodyNormalization::for_tests("provider-model");
let provider_event: Value = serde_json::from_str(
&planned_response_create_event(&template, &active.client_event)
&planned_response_create_event(&template, &normalization, &active.client_event)
.expect("retry provider event should serialize"),
)
.expect("retry provider event should parse");
@@ -814,6 +857,53 @@ mod tests {
assert!(!restored.contains(&second_sentinel), "{restored}");
}
/// Omitting `previous_response_id` starts a new response chain. Restore
/// mappings from the prior chain must not leak into that independent
/// response, while the new chain's first-turn mapping remains available.
#[tokio::test]
async fn an_independent_chain_replaces_prior_restore_mappings() {
let state = redaction_enabled_state();
let decision = control_decision();
let prior = turn_redaction(&state, &decision, TEST_EMAIL).await;
let current = turn_redaction(&state, &decision, OTHER_TEST_EMAIL).await;
let prior_sentinel = sentinel_for(&prior, TEST_EMAIL);
let current_sentinel = sentinel_for(&current, OTHER_TEST_EMAIL);
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
restorer.register(prior.session);
restorer.start_new_chain(Some(current.session));
assert!(
restorer
.restore_provider_frame_text(&provider_delta_frame(&prior_sentinel))
.is_none(),
"an independent chain must not restore PII from its predecessor"
);
let restored = restorer
.restore_provider_frame_text(&provider_delta_frame(&current_sentinel))
.expect("the new chain's first-turn mapping must remain available");
assert!(restored.contains(OTHER_TEST_EMAIL), "{restored}");
assert!(!restored.contains(&current_sentinel), "{restored}");
}
#[tokio::test]
async fn an_unredacted_independent_chain_clears_prior_restore_mappings() {
let state = redaction_enabled_state();
let prior = turn_redaction(&state, &control_decision(), TEST_EMAIL).await;
let prior_sentinel = sentinel_for(&prior, TEST_EMAIL);
let mut restorer = ResponsesWebSocketRedactionRestorer::default();
restorer.register(prior.session);
assert!(restorer.has_sessions());
restorer.start_new_chain(None);
assert!(!restorer.has_sessions());
assert!(restorer
.restore_provider_frame_text(&provider_delta_frame(&prior_sentinel))
.is_none());
}
/// 留存窗口是有界的:长连接不能无限累积映射,代价是更早的轮次会退回
/// 「占位符原样透传」而不是被错误还原成别的值。
#[tokio::test]
@@ -10,6 +10,7 @@
pub enum FatalRelaySignal {
ConnectionAdmissionLost,
InvalidUpstreamText,
UnexpectedUpstreamStreamId,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -40,6 +41,13 @@ pub const fn fatal_relay_policy(signal: FatalRelaySignal) -> FatalRelayPolicy {
client_message: "Provider returned an invalid WebSocket event",
close_reason: "invalid_upstream_event",
},
FatalRelaySignal::UnexpectedUpstreamStreamId => FatalRelayPolicy {
status_code: 502,
close_code: 1011,
error_code: "responses_websocket_unexpected_upstream_stream_id",
client_message: "Provider returned a named-lane event on a default-lane connection",
close_reason: "unexpected_upstream_stream_id",
},
}
}
@@ -230,6 +238,20 @@ mod tests {
assert_eq!(upstream.next(), None);
}
#[test]
fn unexpected_named_lane_is_a_terminal_provider_protocol_error() {
assert_eq!(
fatal_relay_policy(FatalRelaySignal::UnexpectedUpstreamStreamId),
FatalRelayPolicy {
status_code: 502,
close_code: 1011,
error_code: "responses_websocket_unexpected_upstream_stream_id",
client_message: "Provider returned a named-lane event on a default-lane connection",
close_reason: "unexpected_upstream_stream_id",
}
);
}
#[test]
fn quota_snapshot_without_definitive_error_does_not_trigger_retry() {
assert_eq!(
File diff suppressed because it is too large Load Diff
@@ -14,8 +14,12 @@ use serde_json::Value;
use uuid::Uuid;
use super::adapter::resolve_responses_websocket_adapter;
use super::binding::UpstreamBindingIdentity;
use super::client::consume_response_create_rate_limit;
use super::connection::{relay_bound_connection, wait_for_connection_permit_loss};
use super::continuation::{
ResponsesWebSocketContinuationRecord, ResponsesWebSocketContinuationRegistry,
};
use super::control::resolve_responses_websocket_turn_control;
use super::lifecycle::{
await_pending_adapter_observation, await_pending_turn_finalization,
@@ -26,16 +30,20 @@ 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::redaction::redact_responses_websocket_client_event_with_reasoning_replay_policy;
use super::relay_policy::{fatal_relay_policy, FatalRelaySignal};
use super::request::{
build_planning_parts, planned_response_create_event, validated_response_create_model,
build_planning_parts, planned_request_uses_codex_responses_lite, planned_response_create_event,
prepare_responses_lite_continuation, validate_response_create_previous_response_id,
validate_response_create_stream_id_support, validated_named_stream_id,
validated_response_create_model, ResponsesLiteStaticConfig,
};
use super::state::BoundResponsesConnection;
use super::turn::{prepare_responses_websocket_turn_decision, ResponsesWebSocketTurnOutcome};
use super::turn_state::LogicalTurn;
use super::upstream::{bind_responses_upstream, close_bound_upstream};
use crate::ai_serving::ResponsesWebSocketPinnedCandidate;
use crate::handlers::proxy::websocket::ingress::{
WebSocketConnectionLog, WebSocketConnectionLogSpec, WebSocketRequestContext,
};
@@ -45,10 +53,13 @@ use crate::handlers::proxy::websocket::session::{
};
use crate::handlers::proxy::websocket::transport::{
close_client_socket, send_gateway_error, send_gateway_error_with_status,
send_gateway_error_with_stream_id, send_responses_websocket_error_with_param,
};
use crate::privacy::RedactionSession;
use crate::AppState;
const RESPONSES_WEBSOCKET_LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
const CONTINUATION_LOOKUP_TIMEOUT: std::time::Duration = std::time::Duration::from_millis(500);
const RESPONSES_CONNECTION_LOG_SPEC: WebSocketConnectionLogSpec = WebSocketConnectionLogSpec {
opened_event_name: "responses_websocket_connection_opened",
closed_event_name: "responses_websocket_connection_closed",
@@ -74,6 +85,9 @@ enum InitialMessageError {
MissingResponseCreate,
MissingModel,
InvalidModel,
InvalidPreviousResponseId,
InvalidStreamId,
UnsupportedStreamId,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -113,10 +127,11 @@ impl InitialMessageFrameMetadata {
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[derive(Debug, Clone, PartialEq, Eq)]
struct InitialMessageFailure {
error: InitialMessageError,
last_frame: Option<InitialMessageFrameMetadata>,
stream_id: Option<String>,
}
impl InitialMessageFailure {
@@ -124,7 +139,27 @@ impl InitialMessageFailure {
error: InitialMessageError,
last_frame: Option<InitialMessageFrameMetadata>,
) -> Self {
Self { error, last_frame }
Self {
error,
last_frame,
stream_id: None,
}
}
fn for_event(
error: InitialMessageError,
last_frame: Option<InitialMessageFrameMetadata>,
event: &Value,
) -> Self {
// Preserve a valid lane identity for any request-scoped validation
// error. Invalid `stream_id` values are rejected by the grammar helper
// and are never reflected back to the client.
let stream_id = validated_named_stream_id(event).map(str::to_string);
Self {
error,
last_frame,
stream_id,
}
}
}
@@ -156,6 +191,9 @@ impl InitialMessageError {
Self::MissingResponseCreate => "expected_response_create",
Self::MissingModel => "response_create_model_required",
Self::InvalidModel => "invalid_response_create_model",
Self::InvalidPreviousResponseId => "invalid_response_create_previous_response_id",
Self::InvalidStreamId => "invalid_response_create_stream_id",
Self::UnsupportedStreamId => "responses_websocket_named_stream_unsupported",
}
}
@@ -167,6 +205,9 @@ impl InitialMessageError {
Self::MissingResponseCreate | Self::MissingModel | Self::InvalidModel => {
CLOSE_POLICY_VIOLATION
}
Self::InvalidPreviousResponseId => CLOSE_POLICY_VIOLATION,
Self::InvalidStreamId => CLOSE_POLICY_VIOLATION,
Self::UnsupportedStreamId => CLOSE_POLICY_VIOLATION,
}
}
@@ -190,6 +231,15 @@ impl InitialMessageError {
Self::InvalidModel => {
Some("response.create.model must be a non-empty string no longer than 256 bytes")
}
Self::InvalidPreviousResponseId => {
Some("response.create.previous_response_id must be null or a non-empty string")
}
Self::InvalidStreamId => Some(
"response.create.stream_id must be 1-256 ASCII letters, numbers, underscores, hyphens, or periods",
),
Self::UnsupportedStreamId => Some(
"Aether currently supports only the implicit default WebSocket lane; omit response.create.stream_id",
),
}
}
@@ -203,6 +253,9 @@ impl InitialMessageError {
Self::MissingResponseCreate => "unexpected_event_type",
Self::MissingModel => "missing_model",
Self::InvalidModel => "invalid_model",
Self::InvalidPreviousResponseId => "invalid_previous_response_id",
Self::InvalidStreamId => "invalid_stream_id",
Self::UnsupportedStreamId => "unsupported_stream_id",
}
}
@@ -211,7 +264,7 @@ impl InitialMessageError {
}
}
fn initial_message_diagnostic(failure: InitialMessageFailure) -> InitialMessageDiagnostic {
fn initial_message_diagnostic(failure: &InitialMessageFailure) -> InitialMessageDiagnostic {
InitialMessageDiagnostic {
error_code: failure.error.code(),
error_kind: failure.error.kind(),
@@ -225,7 +278,7 @@ fn initial_message_diagnostic(failure: InitialMessageFailure) -> InitialMessageD
fn log_initial_message_failure(
context: &WebSocketRequestContext,
failure: InitialMessageFailure,
failure: &InitialMessageFailure,
upgraded_at: std::time::Instant,
) {
let diagnostic = initial_message_diagnostic(failure);
@@ -360,8 +413,14 @@ async fn bootstrap_responses_websocket(
Ok(value) => value,
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;
log_initial_message_failure(context, &failure, upgraded_at);
send_gateway_error_with_stream_id(
client_socket,
failure.error.code(),
client_message,
failure.stream_id.as_deref(),
)
.await;
close_client_socket(
client_socket,
failure.error.close_code(),
@@ -373,6 +432,18 @@ async fn bootstrap_responses_websocket(
}
};
let initial_previous_response_id = first_event
.get("previous_response_id")
.and_then(Value::as_str)
.filter(|value| !value.trim().is_empty())
.map(str::to_string);
// Keep the response-chain identity on the client's raw configuration.
// Redaction sentinels may rotate between turns, but that must not turn
// identical plaintext tools/instructions into a synthetic config change.
let raw_responses_lite_static_config =
ResponsesLiteStaticConfig::from_response_create(&first_event);
let planning_parts = build_planning_parts(context);
let turn_control = match resolve_responses_websocket_turn_control(
&state,
@@ -418,7 +489,7 @@ async fn bootstrap_responses_websocket(
close_client_socket(client_socket, CLOSE_TRY_AGAIN, "rate_limit_exceeded").await;
return None;
}
Err(()) => {
Err(_) => {
warn!(
event_name = "responses_websocket_rate_limit_check_failed",
log_type = "ops",
@@ -444,16 +515,152 @@ async fn bootstrap_responses_websocket(
}
}
// A cross-socket response chain must be owned by this exact live
// authenticated principal. Missing, expired, corrupt or unavailable state
// fails closed; allowing the normal scheduler to choose a provider/key
// would disclose an opaque response ID to an unrelated account.
let continuation_record = if let Some(previous_response_id) =
initial_previous_response_id.as_deref()
{
let Some(auth_context) = turn_control.decision.auth_context.as_ref() else {
reject_initial_previous_response(client_socket).await;
return None;
};
let registry = ResponsesWebSocketContinuationRegistry::new(state.runtime_state.as_ref());
match tokio::time::timeout(
CONTINUATION_LOOKUP_TIMEOUT,
registry.lookup(
auth_context.user_id.as_str(),
auth_context.api_key_id.as_str(),
previous_response_id,
),
)
.await
{
Ok(Ok(Some(record)))
if record.client_model()
== first_event
.get("model")
.and_then(Value::as_str)
.map(str::trim)
.unwrap_or_default()
&& !record.has_connection_local_redaction() =>
{
Some(record)
}
Ok(Ok(Some(record))) => {
warn!(
event_name = "responses_websocket_continuation_registry_rejected",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
user_id = %auth_context.user_id,
api_key_id = %auth_context.api_key_id,
provider_id = %record.pinned_candidate().provider_id(),
endpoint_id = %record.pinned_candidate().endpoint_id(),
key_id = %record.pinned_candidate().key_id(),
reason = if record.has_connection_local_redaction() {
"connection_local_redaction_state_unavailable"
} else {
"client_model_mismatch"
},
"gateway rejected a cross-socket Responses continuation"
);
reject_initial_previous_response(client_socket).await;
return None;
}
Ok(Ok(None)) => {
warn!(
event_name = "responses_websocket_continuation_registry_miss",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
user_id = %auth_context.user_id,
api_key_id = %auth_context.api_key_id,
reason = "not_found_or_expired",
"gateway could not prove ownership of a cross-socket Responses continuation"
);
reject_initial_previous_response(client_socket).await;
return None;
}
Ok(Err(error)) => {
warn!(
event_name = "responses_websocket_continuation_registry_lookup_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
user_id = %auth_context.user_id,
api_key_id = %auth_context.api_key_id,
reason = error.kind(),
"gateway failed closed while looking up a cross-socket Responses continuation"
);
reject_initial_previous_response(client_socket).await;
return None;
}
Err(_) => {
warn!(
event_name = "responses_websocket_continuation_registry_lookup_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
user_id = %auth_context.user_id,
api_key_id = %auth_context.api_key_id,
reason = "timeout",
timeout_ms = CONTINUATION_LOOKUP_TIMEOUT.as_millis() as u64,
"gateway timed out looking up a cross-socket Responses continuation"
);
reject_initial_previous_response(client_socket).await;
return None;
}
}
} else {
None
};
// Validate and remove any repeated Lite static prefix against the stored
// chain identity before PII redaction rotates per-turn sentinels.
let first_event = match continuation_record
.as_ref()
.and_then(ResponsesWebSocketContinuationRecord::responses_lite_static_config)
{
Some(stored) => match prepare_responses_lite_continuation(&first_event, stored) {
Ok(prepared) => prepared,
Err(_) => {
warn!(
event_name = "responses_websocket_continuation_static_contract_rejected",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
"gateway rejected changed Responses Lite static configuration on a continuation"
);
reject_initial_previous_response(client_socket).await;
return None;
}
},
None => first_event,
};
// 请求侧脱敏必须在规划之前完成,而且这一轮只在这里做一次:planner 会把这份
// body 写进 upstream 请求体和审计 original_request_body,绑定上游的首条
// response.create 也从它派生。脱敏失败时直接断开,绝不退回原文发上游。
let redacted_first_event = redact_responses_websocket_client_event(
&state,
&planning_parts,
&turn_control.decision,
&first_event,
)
.await;
let reasoning_replay_policy = continuation_record
.as_ref()
.map(ResponsesWebSocketContinuationRecord::reasoning_replay_policy)
.unwrap_or_default();
let redacted_first_event =
redact_responses_websocket_client_event_with_reasoning_replay_policy(
&state,
&planning_parts,
&turn_control.decision,
&first_event,
reasoning_replay_policy,
)
.await;
// 首轮的 mask session 要活到响应帧还原,但连接此刻还没绑定,只能先接住,
// 等 `bind_responses_upstream` 之后登记到连接上。
let (first_event, first_turn_redaction_session) = match redacted_first_event {
@@ -486,6 +693,21 @@ async fn bootstrap_responses_websocket(
}
};
let pinned_candidate = match initial_continuation_planning_candidate(
initial_previous_response_id.is_some(),
continuation_record
.as_ref()
.map(|record| record.pinned_candidate().clone()),
) {
Ok(candidate) => candidate,
Err(_) => {
// Keep this invariant next to the planner boundary as a second
// fail-closed guard: an unproved response ID must never enter the
// ordinary scheduler and land on an unrelated provider/key.
reject_initial_previous_response(client_socket).await;
return None;
}
};
let planned = match await_owned_responses_websocket_plan(spawn_owned_responses_websocket_plan(
state.clone(),
planning_parts,
@@ -495,12 +717,24 @@ async fn bootstrap_responses_websocket(
first_event.clone(),
None,
None,
None,
pinned_candidate,
))
.await
{
Ok(Some(decision)) => decision,
Ok(None) => {
if continuation_record.is_some() {
warn!(
event_name = "responses_websocket_continuation_pinned_candidate_unavailable",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
"gateway could not revalidate the registered Responses continuation binding"
);
reject_initial_previous_response(client_socket).await;
return None;
}
send_gateway_error_with_status(
client_socket,
503,
@@ -547,35 +781,85 @@ async fn bootstrap_responses_websocket(
planning_parts,
planned_lease,
} = planned;
let adapter = resolve_responses_websocket_adapter(planned.adapter);
let adapter_kind = planned.adapter;
let adapter = resolve_responses_websocket_adapter(adapter_kind);
let normalization = planned.normalization;
let decision = planned.execution;
let first_provider_event = match planned_response_create_event(&decision, &first_event)
.and_then(|event| {
serde_json::from_str::<Value>(&event).map_err(|_| "responses_websocket_request_invalid")
}) {
Ok(event) => event,
Err(code) => {
if let Some(record) = continuation_record.as_ref() {
let planned_candidate = ResponsesWebSocketPinnedCandidate::from_decision(&decision);
let planned_provider_model = decision
.provider_request_body
.as_ref()
.and_then(|body| body.get("model"))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.or_else(|| {
decision
.mapped_model
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
});
let planned_binding = UpstreamBindingIdentity::from_decision(adapter, &decision).ok();
let planned_uses_responses_lite =
planned_request_uses_codex_responses_lite(&decision, &normalization);
let matches_record = record.adapter() == adapter_kind
&& planned_candidate.as_ref() == Some(record.pinned_candidate())
&& planned_provider_model == Some(record.provider_model())
&& responses_lite_contract_modes_match(
record.responses_lite_static_config().is_some(),
planned_uses_responses_lite,
)
&& planned_binding
.as_ref()
.is_some_and(|binding| record.matches_contract(binding, &normalization));
if !matches_record {
planned_lease.release().await;
warn!(
event_name = "responses_websocket_initial_event_normalization_failed",
event_name = "responses_websocket_continuation_binding_rejected",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error_code = code,
"gateway could not normalize the initial Responses WebSocket event"
provider_id = %record.pinned_candidate().provider_id(),
endpoint_id = %record.pinned_candidate().endpoint_id(),
key_id = %record.pinned_candidate().key_id(),
"gateway rejected a cross-socket continuation after pinned planning changed its contract"
);
send_gateway_error(
client_socket,
code,
"Gateway could not prepare the Responses response.create event",
)
.await;
close_client_socket(client_socket, CLOSE_POLICY_VIOLATION, code).await;
reject_initial_previous_response(client_socket).await;
return None;
}
};
}
let first_provider_event =
match planned_response_create_event(&decision, &normalization, &first_event).and_then(
|event| {
serde_json::from_str::<Value>(&event)
.map_err(|_| "responses_websocket_request_invalid")
},
) {
Ok(event) => event,
Err(code) => {
planned_lease.release().await;
warn!(
event_name = "responses_websocket_initial_event_normalization_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error_code = code,
"gateway could not normalize the initial Responses WebSocket event"
);
send_gateway_error(
client_socket,
code,
"Gateway could not prepare the Responses response.create event",
)
.await;
close_client_socket(client_socket, CLOSE_POLICY_VIOLATION, code).await;
return None;
}
};
let first_logical_turn_id = Uuid::new_v4().to_string();
let first_turn_decision = prepare_responses_websocket_turn_decision(
&decision,
@@ -647,19 +931,83 @@ async fn bootstrap_responses_websocket(
return None;
}
};
if bound.responses_lite_static_config.is_some() {
bound.responses_lite_static_config = continuation_record
.as_ref()
.and_then(ResponsesWebSocketContinuationRecord::responses_lite_static_config)
.cloned()
.or(Some(raw_responses_lite_static_config));
}
if let Some(previous_response_id) = initial_previous_response_id.as_deref() {
// Reaching this point means the principal-scoped registry record and
// the newly planned physical binding were both proved above. Keep the
// parent as persisted ownership: the new physical socket may hydrate
// it, but a failed attempt can evict only its connection-local copy.
bound
.continuation_response_ids
.remember_persisted(previous_response_id);
}
first_turn.mark_upstream_request_sent();
first_turn.set_provider_response_headers(bound.upstream_response_headers.clone());
if let Some(session) = first_turn_redaction_session {
bound.redaction_restorer.register(session);
register_initial_redaction_session(&mut bound, session);
}
bound.turn_state.begin(
LogicalTurn::new(first_event, 1, first_logical_turn_id).with_turn_control(turn_control),
LogicalTurn::new(first_event, 1, first_logical_turn_id)
.with_provider_store(first_provider_event.get("store") == Some(&Value::Bool(true)))
.with_turn_control(turn_control),
first_turn,
);
Some(bound)
}
fn register_initial_redaction_session(
bound: &mut BoundResponsesConnection,
mut session: RedactionSession,
) {
session.set_reasoning_replay_policy(bound.body_normalization.reasoning_replay_policy());
bound.redaction_restorer.register(session);
}
fn responses_lite_contract_modes_match(
stored_chain_uses_responses_lite: bool,
planned_request_uses_responses_lite: bool,
) -> bool {
stored_chain_uses_responses_lite == planned_request_uses_responses_lite
}
fn initial_continuation_planning_candidate(
has_previous_response_id: bool,
registered_candidate: Option<ResponsesWebSocketPinnedCandidate>,
) -> Result<Option<ResponsesWebSocketPinnedCandidate>, &'static str> {
match (has_previous_response_id, registered_candidate) {
(false, None) => Ok(None),
(true, Some(candidate)) => Ok(Some(candidate)),
// A registry miss/corruption and an impossible stray record both fail
// closed. Neither state is allowed to turn into an unpinned plan.
(true, None) | (false, Some(_)) => Err("previous_response_not_found"),
}
}
async fn reject_initial_previous_response(client_socket: &mut WebSocket) {
send_responses_websocket_error_with_param(
client_socket,
400,
"invalid_request_error",
"previous_response_not_found",
"The previous response is unavailable for this authenticated WebSocket connection",
"previous_response_id",
)
.await;
close_client_socket(
client_socket,
CLOSE_POLICY_VIOLATION,
"previous_response_not_found",
)
.await;
}
async fn close_terminated_bootstrap(
client_socket: &mut WebSocket,
context: &WebSocketRequestContext,
@@ -828,7 +1176,7 @@ where
InitialMessageFailure::new(InitialMessageError::InvalidJson, last_frame)
})?;
validate_initial_response_create(&event)
.map_err(|error| InitialMessageFailure::new(error, last_frame))?;
.map_err(|error| InitialMessageFailure::for_event(error, last_frame, &event))?;
return Ok((text, event));
}
}
@@ -844,6 +1192,15 @@ fn validate_initial_response_create(event: &Value) -> Result<(), InitialMessageE
.get("model")
.ok_or(InitialMessageError::MissingModel)?;
validated_response_create_model(model).map_err(|_| InitialMessageError::InvalidModel)?;
validate_response_create_previous_response_id(event)
.map_err(|_| InitialMessageError::InvalidPreviousResponseId)?;
match validate_response_create_stream_id_support(event) {
Ok(()) => {}
Err("invalid_response_create_stream_id") => {
return Err(InitialMessageError::InvalidStreamId);
}
Err(_) => return Err(InitialMessageError::UnsupportedStreamId),
}
Ok(())
}
@@ -875,11 +1232,15 @@ mod tests {
};
use super::super::turn_state::{LogicalTurn, ResponsesTurnState};
use super::super::upstream::bind_responses_upstream;
use crate::ai_serving::{AiExecutionDecision, ResponsesWebSocketBodyNormalization};
use crate::ai_serving::{
AiExecutionDecision, OpenAiResponsesReasoningReplayPolicy,
ResponsesWebSocketBodyNormalization,
};
use crate::handlers::proxy::websocket::session::wait_for_optional_deadline;
use crate::handlers::proxy::websocket::transport::{
websocket_handshake_headers, websocket_timeouts, websocket_upstream_url,
};
use crate::privacy::{RedactionSession, RedactionSessionConfig};
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
use axum::extract::State;
use axum::http::header::{AUTHORIZATION, CONTENT_TYPE};
@@ -987,10 +1348,38 @@ mod tests {
1008,
false,
),
(
InitialMessageError::InvalidPreviousResponseId,
"invalid_response_create_previous_response_id",
"invalid_previous_response_id",
Some("response.create.previous_response_id must be null or a non-empty string"),
1008,
false,
),
(
InitialMessageError::InvalidStreamId,
"invalid_response_create_stream_id",
"invalid_stream_id",
Some(
"response.create.stream_id must be 1-256 ASCII letters, numbers, underscores, hyphens, or periods",
),
1008,
false,
),
(
InitialMessageError::UnsupportedStreamId,
"responses_websocket_named_stream_unsupported",
"unsupported_stream_id",
Some(
"Aether currently supports only the implicit default WebSocket lane; omit response.create.stream_id",
),
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));
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);
@@ -1011,7 +1400,7 @@ mod tests {
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(
let diagnostic = initial_message_diagnostic(&InitialMessageFailure::new(
InitialMessageError::MissingResponseCreate,
Some(metadata),
));
@@ -1211,8 +1600,10 @@ mod tests {
"stream": true,
"background": true,
}));
let normalization = ResponsesWebSocketBodyNormalization::for_tests("provider-model");
let event = planned_response_create_event(
&decision,
&normalization,
&json!({
"type": "response.create",
"model": "public-model",
@@ -1309,6 +1700,168 @@ mod tests {
);
}
#[test]
fn malformed_initial_previous_response_id_is_rejected_before_planning() {
for previous_response_id in [json!(""), json!(42), json!({"id": "resp_1"})] {
let initial = json!({
"type": "response.create",
"model": "gpt-5.6-sol",
"previous_response_id": previous_response_id,
"input": [],
});
assert!(matches!(
super::validate_initial_response_create(&initial),
Err(super::InitialMessageError::InvalidPreviousResponseId)
));
}
let named = json!({
"type": "response.create",
"model": "gpt-5.6-sol",
"stream_id": "main",
"previous_response_id": "",
"input": [],
});
let failure = super::InitialMessageFailure::for_event(
super::InitialMessageError::InvalidPreviousResponseId,
None,
&named,
);
assert_eq!(failure.stream_id.as_deref(), Some("main"));
}
#[test]
fn valid_previous_response_can_start_a_new_socket() {
let initial = json!({
"type": "response.create",
"model": "gpt-5.6-sol",
"previous_response_id": "resp_existing",
"input": [{"type": "function_call_output", "call_id": "call_1", "output": "ok"}],
});
assert!(super::validate_initial_response_create(&initial).is_ok());
let null_previous = json!({
"type": "response.create",
"model": "gpt-5.6-sol",
"previous_response_id": null,
"input": [{"role": "user", "content": "new chain"}],
});
assert!(super::validate_initial_response_create(&null_previous).is_ok());
}
#[test]
fn cross_socket_registry_miss_never_falls_through_to_random_provider_planning() {
let pinned = crate::ai_serving::ResponsesWebSocketPinnedCandidate::new(
"provider-original",
"endpoint-original",
"key-original",
)
.expect("valid pinned candidate");
assert_eq!(
super::initial_continuation_planning_candidate(false, None),
Ok(None),
"a genuinely new response may use the ordinary planner"
);
assert_eq!(
super::initial_continuation_planning_candidate(true, Some(pinned.clone())),
Ok(Some(pinned)),
"a proved continuation must retain its exact provider/endpoint/key"
);
assert_eq!(
super::initial_continuation_planning_candidate(true, None),
Err("previous_response_not_found"),
"a registry miss must be rejected before an unpinned planner can choose another key"
);
}
#[test]
fn cross_socket_continuation_rejects_an_effective_lite_mode_change() {
let mut decision: AiExecutionDecision = serde_json::from_value(json!({
"action": "local",
"provider_type": "codex",
"provider_api_format": "openai:responses",
"provider_request_headers": {}
}))
.expect("minimal Codex decision");
decision.provider_request_headers.insert(
crate::ai_serving::CODEX_RESPONSES_LITE_HEADER.to_string(),
"true".to_string(),
);
let normalization = ResponsesWebSocketBodyNormalization::for_tests("gpt-5.6-sol")
.with_provider_type_for_tests("codex");
let effective_lite =
super::planned_request_uses_codex_responses_lite(&decision, &normalization);
assert!(effective_lite);
assert!(super::responses_lite_contract_modes_match(
true,
effective_lite
));
// A non-null context_management object suppresses the converged Lite
// contract/header even though the model capability remains enabled.
// A chain whose stored prefix used Lite must not cross that boundary.
decision.provider_request_body = Some(json!({
"model": "gpt-5.6-sol",
"context_management": {"compact_threshold": 1_000}
}));
let effective_lite =
super::planned_request_uses_codex_responses_lite(&decision, &normalization);
assert!(!effective_lite);
assert!(!super::responses_lite_contract_modes_match(
true,
effective_lite
));
assert!(super::responses_lite_contract_modes_match(
false,
effective_lite
));
}
#[test]
fn initial_named_stream_is_rejected_before_planning() {
let initial = json!({
"type": "response.create",
"model": "gpt-5.6-sol",
"stream_id": "main",
"input": [],
});
assert!(matches!(
super::validate_initial_response_create(&initial),
Err(super::InitialMessageError::UnsupportedStreamId)
));
let failure = super::InitialMessageFailure::for_event(
super::InitialMessageError::UnsupportedStreamId,
None,
&initial,
);
assert_eq!(failure.stream_id.as_deref(), Some("main"));
}
#[test]
fn malformed_initial_stream_id_is_rejected_before_planning() {
for stream_id in [json!(null), json!(""), json!("not/a/lane"), json!(42)] {
let initial = json!({
"type": "response.create",
"model": "gpt-5.6-sol",
"stream_id": stream_id,
"input": [],
});
assert!(matches!(
super::validate_initial_response_create(&initial),
Err(super::InitialMessageError::InvalidStreamId)
));
let failure = super::InitialMessageFailure::for_event(
super::InitialMessageError::InvalidStreamId,
None,
&initial,
);
assert_eq!(failure.stream_id, None);
}
}
#[test]
fn model_at_the_identifier_limit_remains_valid() {
let model = "m".repeat(256);
@@ -1669,7 +2222,9 @@ mod tests {
provider_model: "gpt-5.6-sol".to_string(),
decision_template: decision,
body_normalization: ResponsesWebSocketBodyNormalization::for_tests("gpt-5.6-sol"),
responses_lite_static_config: None,
binding_identity,
continuation_response_ids: Default::default(),
// Replanning:logical turn 在、attempt 不在。重放安全与配额排除都只看
// logical turn,所以这些用例不需要真实 socket 或真实 attempt。
turn_state: ResponsesTurnState::Replanning {
@@ -1689,6 +2244,57 @@ mod tests {
}
}
#[test]
fn initial_deepseek_binding_upgrades_redaction_restore_policy() {
let mut session = RedactionSession::new(RedactionSessionConfig::new(
b"initial-deepseek-redaction-test".to_vec(),
300,
600,
));
let sentinel = session.redact_text("[email protected]").text;
let provider_event = json!({
"type": "response.completed",
"response": {
"output": [{
"type": "reasoning",
"encrypted_content": "provider-owned-state",
"content": [{
"type": "reasoning_text",
"text": format!("opaque replay {sentinel}")
}]
}]
}
});
let mut ordinary_bound = sample_bound_for_rebind_safety();
super::register_initial_redaction_session(&mut ordinary_bound, session.clone());
let restored = ordinary_bound
.redaction_restorer
.restore_provider_frame_text(&provider_event)
.expect("ordinary OpenAI replay policy should restore response text");
let restored: serde_json::Value =
serde_json::from_str(&restored).expect("restored provider event should remain JSON");
assert_eq!(
restored["response"]["output"][0]["content"][0]["text"],
"opaque replay [email protected]"
);
let mut deepseek_bound = sample_bound_for_rebind_safety();
deepseek_bound.body_normalization =
ResponsesWebSocketBodyNormalization::for_tests("deepseek-reasoner")
.with_reasoning_replay_policy_for_tests(
OpenAiResponsesReasoningReplayPolicy::DeepSeekOpaque,
);
super::register_initial_redaction_session(&mut deepseek_bound, session);
assert!(
deepseek_bound
.redaction_restorer
.restore_provider_frame_text(&provider_event)
.is_none(),
"the authenticated DeepSeek binding must keep opaque reasoning state byte-identical"
);
}
/// 用 mpsc 驱动的 FakeSocket,实现 Stream + Sink 两个 trait。
/// 测试侧通过 tx 注入消息,通过 pong_rx 观察 Pong 回包。
struct FakeSocket {
@@ -4,16 +4,93 @@
//! connection may survive many `response.create` turns, while the turn
//! lifecycle and upstream binding are replaced independently.
use std::collections::{BTreeMap, BTreeSet};
use std::collections::{BTreeMap, BTreeSet, VecDeque};
use tokio::task::JoinHandle;
use super::adapter::{ResponsesWebSocketDrainDirective, ResponsesWebSocketProtocolAdapter};
use super::binding::UpstreamBindingIdentity;
use super::redaction::ResponsesWebSocketRedactionRestorer;
use super::request::ResponsesLiteStaticConfig;
use super::turn_state::ResponsesTurnState;
use crate::ai_serving::{AiExecutionDecision, ResponsesWebSocketBodyNormalization};
const EXHAUSTED_KEY_EXCLUSION_FALLBACK_SECONDS: u64 = 300;
const MAX_CONNECTION_LOCAL_CONTINUATION_RESPONSE_IDS: usize = 1_024;
#[derive(Debug, Default)]
struct BoundedContinuationResponseIds {
ids: BTreeSet<String>,
insertion_order: VecDeque<String>,
}
impl BoundedContinuationResponseIds {
fn contains(&self, response_id: &str) -> bool {
self.ids.contains(response_id)
}
fn remember(&mut self, response_id: &str) {
if !self.ids.insert(response_id.to_string()) {
return;
}
self.insertion_order.push_back(response_id.to_string());
while self.ids.len() > MAX_CONNECTION_LOCAL_CONTINUATION_RESPONSE_IDS {
let Some(oldest) = self.insertion_order.pop_front() else {
break;
};
self.ids.remove(oldest.as_str());
}
}
fn forget(&mut self, response_id: &str) {
if self.ids.remove(response_id) {
self.insertion_order
.retain(|remembered| remembered != response_id);
}
}
fn clear(&mut self) {
self.ids.clear();
self.insertion_order.clear();
}
}
/// Principal-proved response IDs for the currently bound response chain.
///
/// Connection-local IDs prove that the provider's current physical socket can
/// resolve a parent, including `store=false` responses. Persisted IDs are kept
/// separately after a successful principal-scoped registry write (or a proved
/// cross-socket bootstrap), so a 4xx/5xx eviction of the provider's local cache
/// does not suppress the provider's documented `store=true` hydration fallback.
/// Starting an independent chain or replacing the physical upstream clears
/// both bounded sets.
#[derive(Debug, Default)]
pub(super) struct ContinuationResponseIds {
connection_local: BoundedContinuationResponseIds,
persisted: BoundedContinuationResponseIds,
}
impl ContinuationResponseIds {
pub(super) fn contains(&self, response_id: &str) -> bool {
self.connection_local.contains(response_id) || self.persisted.contains(response_id)
}
pub(super) fn remember_connection_local(&mut self, response_id: &str) {
self.connection_local.remember(response_id);
}
pub(super) fn remember_persisted(&mut self, response_id: &str) {
self.persisted.remember(response_id);
}
pub(super) fn forget_connection_local(&mut self, response_id: &str) {
self.connection_local.forget(response_id);
}
pub(super) fn clear(&mut self) {
self.connection_local.clear();
self.persisted.clear();
}
}
/// All mutable state associated with the physical upstream connection.
pub(super) struct BoundResponsesConnection {
@@ -26,7 +103,15 @@ pub(super) struct BoundResponsesConnection {
/// turns, which must not re-enter the planner. Replaced whenever the
/// binding or its decision is replaced.
pub(super) body_normalization: ResponsesWebSocketBodyNormalization,
/// Static Responses Lite configuration already represented in the current
/// response chain. A continuation may repeat it, but may not append a
/// changed synthetic prefix to the inherited history.
pub(super) responses_lite_static_config: Option<ResponsesLiteStaticConfig>,
pub(super) binding_identity: UpstreamBindingIdentity,
/// IDs observed on the current physical upstream and current independent
/// response chain. A continuation must reference one of these IDs; the
/// cross-socket bootstrap path seeds the already registry-proved parent.
pub(super) continuation_response_ids: ContinuationResponseIds,
/// 这条连接上「有没有正在进行的 logical turn」的唯一事实来源。
pub(super) turn_state: ResponsesTurnState,
/// 这条连接迄今 mask 出来的映射,用于把 provider 事件里的占位符换回真实值。
@@ -104,3 +189,63 @@ impl ExhaustedResponsesWebSocketExclusions {
.retain(|_, expires_at| *expires_at > now_unix_secs);
}
}
#[cfg(test)]
mod tests {
use super::{ContinuationResponseIds, MAX_CONNECTION_LOCAL_CONTINUATION_RESPONSE_IDS};
#[test]
fn continuation_ids_are_bounded_and_clear_with_the_chain() {
let mut ids = ContinuationResponseIds::default();
for index in 0..=MAX_CONNECTION_LOCAL_CONTINUATION_RESPONSE_IDS {
ids.remember_connection_local(format!("resp_{index}").as_str());
}
assert!(!ids.contains("resp_0"));
assert!(ids.contains("resp_1"));
assert!(
ids.contains(format!("resp_{MAX_CONNECTION_LOCAL_CONTINUATION_RESPONSE_IDS}").as_str())
);
assert_eq!(
ids.connection_local.ids.len(),
MAX_CONNECTION_LOCAL_CONTINUATION_RESPONSE_IDS
);
assert_eq!(
ids.connection_local.insertion_order.len(),
MAX_CONNECTION_LOCAL_CONTINUATION_RESPONSE_IDS
);
ids.forget_connection_local("resp_1");
assert!(!ids.contains("resp_1"));
assert_eq!(
ids.connection_local.insertion_order.len(),
MAX_CONNECTION_LOCAL_CONTINUATION_RESPONSE_IDS - 1
);
ids.clear();
assert!(ids.connection_local.ids.is_empty());
assert!(ids.connection_local.insertion_order.is_empty());
assert!(ids.persisted.ids.is_empty());
assert!(ids.persisted.insertion_order.is_empty());
}
#[test]
fn local_eviction_keeps_only_registered_persisted_ownership() {
let mut ids = ContinuationResponseIds::default();
ids.remember_connection_local("resp_store_false");
ids.remember_connection_local("resp_store_true");
ids.remember_persisted("resp_store_true");
ids.forget_connection_local("resp_store_false");
ids.forget_connection_local("resp_store_true");
assert!(
!ids.contains("resp_store_false"),
"store=false has no persisted hydration fallback after local eviction"
);
assert!(
ids.contains("resp_store_true"),
"a registered store=true parent remains eligible for persisted hydration"
);
}
}
@@ -20,6 +20,10 @@ use super::request::response_create_has_previous_response_id;
#[derive(Debug, Clone)]
pub(super) struct LogicalTurn {
pub(super) client_event: Value,
/// Effective `store` after provider body rules and WebSocket framing. Only
/// an explicit provider-side `true` permits cross-connection registry
/// state; false or absent remains ZDR/connection-local.
pub(super) provider_store: bool,
pub(super) turn_index: u64,
pub(super) logical_turn_id: String,
pub(super) turn_attempt: u32,
@@ -35,6 +39,7 @@ impl LogicalTurn {
pub(super) fn new(client_event: Value, turn_index: u64, logical_turn_id: String) -> Self {
Self {
client_event,
provider_store: false,
turn_index,
logical_turn_id,
turn_attempt: 1,
@@ -49,6 +54,11 @@ impl LogicalTurn {
self
}
pub(super) fn with_provider_store(mut self, provider_store: bool) -> Self {
self.provider_store = provider_store;
self
}
pub(super) fn quota_retry_block_reason(&self) -> Option<&'static str> {
if self.retry_attempted {
Some("quota_retry_already_attempted")
@@ -8,8 +8,10 @@ use wreq::ws::message::Message as WreqWsMessage;
use super::adapter::ResponsesWebSocketProtocolAdapter;
use super::binding::{UpstreamBindingIdentity, UpstreamBindingIdentityError};
use super::redaction::ResponsesWebSocketRedactionRestorer;
use super::request::planned_response_create_event;
use super::state::{BoundResponsesConnection, ExhaustedResponsesWebSocketExclusions};
use super::request::{planned_request_uses_codex_responses_lite, planned_response_create_event};
use super::state::{
BoundResponsesConnection, ContinuationResponseIds, ExhaustedResponsesWebSocketExclusions,
};
use super::turn_state::ResponsesTurnState;
use crate::ai_serving::{AiExecutionDecision, ResponsesWebSocketBodyNormalization};
use crate::handlers::proxy::websocket::session::RESPONSES_WEBSOCKET_SESSION_LIMITS;
@@ -79,7 +81,7 @@ async fn bind_responses_upstream_inner(
adapter.upstream_errors(),
)
.await?;
let first_event = planned_response_create_event(decision, initial_event)?;
let first_event = planned_response_create_event(decision, &normalization, initial_event)?;
send_upstream_message(&mut upstream.socket, WreqWsMessage::text(first_event))
.await
.map_err(|_| "responses_websocket_initial_send_failed")?;
@@ -108,6 +110,11 @@ async fn bind_responses_upstream_inner(
.ok_or("responses_websocket_mapped_model_missing")?
.to_string();
let responses_lite_static_config =
planned_request_uses_codex_responses_lite(decision, &normalization).then(|| {
super::request::ResponsesLiteStaticConfig::from_response_create(initial_event)
});
Ok(BoundResponsesConnection {
upstream: Some(upstream.socket),
adapter,
@@ -115,7 +122,9 @@ async fn bind_responses_upstream_inner(
provider_model,
decision_template: decision.clone(),
body_normalization: normalization,
responses_lite_static_config,
binding_identity,
continuation_response_ids: ContinuationResponseIds::default(),
// 首条 response.create 已经发出,但这一轮的 logical turn 和 attempt 由调用方
// 通过 `ResponsesTurnState::begin` 装上:绑定本身不持有记账状态。
turn_state: ResponsesTurnState::Idle,
@@ -82,6 +82,7 @@ pub(crate) async fn connect_upstream_websocket(
fn websocket_response_headers(headers: &HeaderMap) -> BTreeMap<String, String> {
headers
.iter()
.filter(|(name, _)| websocket_response_header_is_safe_to_retain(name))
.filter_map(|(name, value)| {
value
.to_str()
@@ -91,6 +92,24 @@ fn websocket_response_headers(headers: &HeaderMap) -> BTreeMap<String, String> {
.collect()
}
fn websocket_response_header_is_safe_to_retain(name: &HeaderName) -> bool {
!matches!(
name.as_str(),
"authorization"
| "proxy-authorization"
| "www-authenticate"
| "proxy-authenticate"
| "authentication-info"
| "proxy-authentication-info"
| "cookie"
| "set-cookie"
| "set-cookie2"
| "x-api-key"
| "api-key"
| "x-goog-api-key"
)
}
pub(crate) fn websocket_upstream_url(
raw: &str,
invalid_code: &'static str,
@@ -300,7 +319,20 @@ pub(crate) fn responses_websocket_error_event(
code: &str,
message: &str,
) -> serde_json::Value {
json!({
responses_websocket_error_event_with_stream_id(status, error_type, code, message, None)
}
/// Builds a request-scoped Responses error. Callers must supply `stream_id`
/// only after validating the protocol's named-lane grammar; untrusted or
/// malformed identifiers must never be reflected into a provider event.
pub(crate) fn responses_websocket_error_event_with_stream_id(
status: u16,
error_type: &str,
code: &str,
message: &str,
stream_id: Option<&str>,
) -> serde_json::Value {
let mut event = json!({
"type": "error",
"status": status,
"error": {
@@ -308,7 +340,17 @@ pub(crate) fn responses_websocket_error_event(
"code": code,
"message": message,
},
})
});
if let Some(stream_id) = stream_id {
event
.as_object_mut()
.expect("Responses error events are JSON objects")
.insert(
"stream_id".to_string(),
serde_json::Value::String(stream_id.to_string()),
);
}
event
}
pub(crate) async fn send_responses_websocket_error(
@@ -318,7 +360,49 @@ pub(crate) async fn send_responses_websocket_error(
code: &str,
message: &str,
) {
let event = responses_websocket_error_event(status, error_type, code, message);
send_responses_websocket_error_with_stream_id(
client_socket,
status,
error_type,
code,
message,
None,
)
.await;
}
/// Sends a standard invalid-request error with a bounded, server-owned
/// parameter name. This is used for protocol fields such as
/// `previous_response_id`; no untrusted value is reflected.
pub(crate) async fn send_responses_websocket_error_with_param(
client_socket: &mut WebSocket,
status: u16,
error_type: &str,
code: &str,
message: &str,
param: &'static str,
) {
let mut event = responses_websocket_error_event(status, error_type, code, message);
event["error"]["param"] = serde_json::Value::String(param.to_string());
send_teardown_message(
client_socket
.send(AxumWsMessage::Text(event.to_string().into()))
.map_err(|_| ()),
)
.await;
}
pub(crate) async fn send_responses_websocket_error_with_stream_id(
client_socket: &mut WebSocket,
status: u16,
error_type: &str,
code: &str,
message: &str,
stream_id: Option<&str>,
) {
let event = responses_websocket_error_event_with_stream_id(
status, error_type, code, message, stream_id,
);
send_teardown_message(
client_socket
.send(AxumWsMessage::Text(event.to_string().into()))
@@ -331,13 +415,41 @@ pub(crate) async fn send_gateway_error(client_socket: &mut WebSocket, code: &str
send_gateway_error_with_status(client_socket, 400, code, message).await;
}
pub(crate) async fn send_gateway_error_with_stream_id(
client_socket: &mut WebSocket,
code: &str,
message: &str,
stream_id: Option<&str>,
) {
send_gateway_error_with_status_and_stream_id(client_socket, 400, code, message, stream_id)
.await;
}
pub(crate) async fn send_gateway_error_with_status(
client_socket: &mut WebSocket,
status: u16,
code: &str,
message: &str,
) {
send_responses_websocket_error(client_socket, status, "gateway_error", code, message).await;
send_gateway_error_with_status_and_stream_id(client_socket, status, code, message, None).await;
}
pub(crate) async fn send_gateway_error_with_status_and_stream_id(
client_socket: &mut WebSocket,
status: u16,
code: &str,
message: &str,
stream_id: Option<&str>,
) {
send_responses_websocket_error_with_stream_id(
client_socket,
status,
"gateway_error",
code,
message,
stream_id,
)
.await;
}
pub(crate) async fn close_client_socket(client_socket: &mut WebSocket, code: u16, reason: &str) {
@@ -355,9 +467,12 @@ 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_handshake_headers,
websocket_upstream_url, WebSocketWriteError, RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT,
bounded_send, responses_websocket_error_event,
responses_websocket_error_event_with_stream_id, websocket_handshake_headers,
websocket_response_headers, websocket_upstream_url, WebSocketWriteError,
RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT,
};
use axum::http::HeaderMap;
use std::collections::BTreeMap;
use std::time::Duration;
@@ -391,6 +506,32 @@ mod tests {
assert!(TEARDOWN_WRITE_TIMEOUT < RELAY_WRITE_TIMEOUT);
}
#[test]
fn upstream_handshake_observability_drops_credential_bearing_headers() {
let mut headers = HeaderMap::new();
headers.insert("x-codex-primary-used-percent", "10".parse().unwrap());
headers.insert("x-request-id", "request-123".parse().unwrap());
headers.insert("set-cookie", "session=secret".parse().unwrap());
headers.insert("www-authenticate", "Bearer secret".parse().unwrap());
headers.insert("authentication-info", "nextnonce=secret".parse().unwrap());
let retained = websocket_response_headers(&headers);
assert_eq!(
retained
.get("x-codex-primary-used-percent")
.map(String::as_str),
Some("10")
);
assert_eq!(
retained.get("x-request-id").map(String::as_str),
Some("request-123")
);
assert!(!retained.contains_key("set-cookie"));
assert!(!retained.contains_key("www-authenticate"));
assert!(!retained.contains_key("authentication-info"));
}
#[test]
fn builds_a_client_compatible_responses_error_event() {
let event = responses_websocket_error_event(
@@ -408,6 +549,24 @@ mod tests {
event["error"]["message"],
"Previous response was not found."
);
assert!(event.get("stream_id").is_none());
}
#[test]
fn request_scoped_responses_errors_include_the_validated_named_stream() {
let event = responses_websocket_error_event_with_stream_id(
400,
"gateway_error",
"responses_websocket_named_stream_unsupported",
"Named streams are not supported.",
Some("main-lane_1.test"),
);
assert_eq!(event["stream_id"], "main-lane_1.test");
assert_eq!(
event["error"]["code"],
"responses_websocket_named_stream_unsupported"
);
}
#[test]
@@ -146,9 +146,8 @@ async fn load_codex_model_cards(
event_name = "codex_catalog_aggregate_incomplete",
client_version = %client_version.as_str(),
target_count = targets.len(),
"Codex catalog aggregation was incomplete; returning an empty remote catalog so the client can use its bundled fallback"
"Codex catalog aggregation was incomplete; serving cards from available last-known-good snapshots"
);
return (Vec::new(), None);
}
let mut seen_global_models = BTreeSet::new();
let possible_inference_catalogs = rows
@@ -210,9 +209,8 @@ async fn load_codex_model_cards(
expected_model_count = expected_global_models.len(),
projected_model_count = cards.len(),
missing_model_count,
"Codex upstream catalogs omitted authorized mappings; returning an empty remote catalog so the client can use its bundled fallback"
"Codex upstream catalogs omitted authorized mappings; serving the available cards without fabricating missing model metadata"
);
return (Vec::new(), None);
}
if !codex_projected_catalog_fits_response_limits(&cards) {
warn!(