mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 08:27:46 +08:00
fix(ws): harden Responses continuation state
This commit is contained in:
@@ -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(¤t, 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(¤t_sentinel))
|
||||
.expect("the new chain's first-turn mapping must remain available");
|
||||
assert!(restored.contains(OTHER_TEST_EMAIL), "{restored}");
|
||||
assert!(!restored.contains(¤t_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!(
|
||||
|
||||
Reference in New Issue
Block a user