From d125f1de47f7146913de082274002e341474a0dd Mon Sep 17 00:00:00 2001 From: daniel Date: Wed, 23 Sep 2026 11:43:53 -0700 Subject: [PATCH 1/2] feat(openai): add session-scoped Responses WebSocket transport --- Cargo.lock | 79 +- crates/agentkit-provider-openai/Cargo.toml | 3 +- crates/agentkit-provider-openai/README.md | 80 ++ crates/agentkit-provider-openai/src/lib.rs | 2 +- .../agentkit-provider-openai/src/responses.rs | 272 ++++++- .../src/responses/websocket.rs | 339 ++++++++ .../src/responses/websocket/tests.rs | 763 ++++++++++++++++++ 7 files changed, 1496 insertions(+), 42 deletions(-) create mode 100644 crates/agentkit-provider-openai/src/responses/websocket.rs create mode 100644 crates/agentkit-provider-openai/src/responses/websocket/tests.rs diff --git a/Cargo.lock b/Cargo.lock index 4161a72..852ed00 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -353,7 +353,7 @@ dependencies = [ [[package]] name = "agentkit-provider-openai" -version = "0.10.9" +version = "0.10.10" dependencies = [ "agentkit-adapter-completions", "agentkit-core", @@ -367,6 +367,7 @@ dependencies = [ "serde_json", "thiserror 2.0.18", "tokio", + "tokio-tungstenite", "zeroize", ] @@ -1169,6 +1170,12 @@ dependencies = [ "syn 3.0.3", ] +[[package]] +name = "data-encoding" +version = "2.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4583a4551df46e2792f82ceeac45e850d2e2d5debba0b91f102385cda5b11f06" + [[package]] name = "defmt" version = "1.1.1" @@ -2811,10 +2818,20 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "22f6172bdec972074665ed81ed53b71da00bfc44b65a753cfde883ec4c702a1a" dependencies = [ "libc", - "rand_chacha", + "rand_chacha 0.3.1", "rand_core 0.6.4", ] +[[package]] +name = "rand" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b9ef1d0d795eb7d84685bca4f72f3649f064e6641543d3a8c415898726a57b41" +dependencies = [ + "rand_chacha 0.9.0", + "rand_core 0.9.5", +] + [[package]] name = "rand" version = "0.10.2" @@ -2836,6 +2853,16 @@ dependencies = [ "rand_core 0.6.4", ] +[[package]] +name = "rand_chacha" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3022b5f1df60f26e1ffddd6c66e8aa15de382ae63b3a0c1bfc0e4d3e3f325cb" +dependencies = [ + "ppv-lite86", + "rand_core 0.9.5", +] + [[package]] name = "rand_core" version = "0.6.4" @@ -2845,6 +2872,15 @@ dependencies = [ "getrandom 0.2.17", ] +[[package]] +name = "rand_core" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "76afc826de14238e6e8c374ddcc1fa19e374fd8dd986b0d2af0d02377261d83c" +dependencies = [ + "getrandom 0.3.4", +] + [[package]] name = "rand_core" version = "0.10.1" @@ -3462,6 +3498,17 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "sha1" +version = "0.10.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a978451301f4db1d02937a4ab3ccce137717b81826e79b7d49ffe3244a13c3b8" +dependencies = [ + "cfg-if", + "cpufeatures 0.2.17", + "digest 0.10.7", +] + [[package]] name = "sha2" version = "0.10.9" @@ -3826,6 +3873,18 @@ dependencies = [ "tokio", ] +[[package]] +name = "tokio-tungstenite" +version = "0.29.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f72a05e828585856dacd553fba484c242c46e391fb0e58917c942ee9202915c" +dependencies = [ + "futures-util", + "log", + "tokio", + "tungstenite", +] + [[package]] name = "tokio-util" version = "0.7.19" @@ -3964,6 +4023,22 @@ version = "0.2.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" +[[package]] +name = "tungstenite" +version = "0.29.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6c01152af293afb9c7c2a57e4b559c5620b421f6d133261c60dd2d0cdb38e6b8" +dependencies = [ + "bytes", + "data-encoding", + "http", + "httparse", + "log", + "rand 0.9.5", + "sha1", + "thiserror 2.0.18", +] + [[package]] name = "typeid" version = "1.0.3" diff --git a/crates/agentkit-provider-openai/Cargo.toml b/crates/agentkit-provider-openai/Cargo.toml index 17e1545..5da3482 100644 --- a/crates/agentkit-provider-openai/Cargo.toml +++ b/crates/agentkit-provider-openai/Cargo.toml @@ -7,9 +7,10 @@ edition.workspace = true license.workspace = true repository.workspace = true rust-version.workspace = true -version = "0.10.9" +version = "0.10.10" [dependencies] +tokio-tungstenite = { version = "=0.29.0", default-features = false, features = ["handshake"] } agentkit-adapter-completions = { version = "0.10.8", path = "../agentkit-adapter-completions" } agentkit-core = { version = "0.10.5", path = "../agentkit-core" } agentkit-http = { version = "0.10.6", path = "../agentkit-http" } diff --git a/crates/agentkit-provider-openai/README.md b/crates/agentkit-provider-openai/README.md index cfea30c..d1e5060 100644 --- a/crates/agentkit-provider-openai/README.md +++ b/crates/agentkit-provider-openai/README.md @@ -76,6 +76,86 @@ these with `OpenAIResponsesLimits` and Limits must be non-zero; the per-field bound must fit both request and attempt bounds, and the per-attempt bound must fit the aggregate wire bound. +## Responses WebSocket transport + +HTTP/SSE remains the default. Select a transport explicitly on the Responses +configuration (public API and ChatGPT private profiles both support selection): + +```rust +use agentkit_provider_openai::{OpenAIResponsesConfig, OpenAIResponsesTransport}; + +let config = OpenAIResponsesConfig::new("token", "gpt-5") + .with_transport(OpenAIResponsesTransport::Auto); +``` + +- `Http`: the existing HTTP/SSE path, including custom `Http` clients. +- `WebSocket`: require an HTTP/1 WebSocket upgrade; never silently use SSE. +- `Auto`: try WebSocket. An upgrade response of HTTP **426** switches this + session to HTTP/SSE for all remaining turns. No fallback is attempted after a + `response.create` has been sent, or for arbitrary authentication/protocol errors. + +The configured `https://.../responses` endpoint is upgraded on that same path +(the wire equivalent of `wss://.../responses`); `http://` endpoints support local +mock servers. Authentication comes from the same refresh-capable provider used +by HTTP. The handshake sends `OpenAI-Beta: responses_websockets=2026-02-06` and +requests are JSON text messages with `type: "response.create"` plus the normal +Responses fields, excluding HTTP-only `stream` and `background` controls as +required by the [public WebSocket guide](https://developers.openai.com/api/docs/guides/websocket-mode). JSON events use the same semantic decoder as SSE, preserving +reasoning, tool, usage, image, and credential-bound continuation metadata. +No browser or separate authentication flow is required. + +Each `ModelSession` owns an independent connection slot. A turn checks out that +slot and releases it **when `Finished` is delivered**, not when the owned turn +value is eventually dropped. Only one active turn is allowed in a WebSocket/Auto +session; a concurrent `begin_turn` is rejected. A validated terminal response +makes the socket reusable. Cancellation, dropping an unfinished turn, malformed +messages, missing completion, and terminal errors discard it. Buffered unsolicited +frames force reconnection before another request; known completed response IDs +are rejected if they reappear on a reused socket. Cancellation never +sends `response.cancel`. Changed authentication headers or credential binding +force a new connection; a 401 can refresh once only if the binding is unchanged. + +WebSocket retries are intentionally more conservative than HTTP retries: + +- Handshake status failures use the existing bounded retry policy and observer. +- A wrapped HTTP error before `response.created` (and before visible output) + may reconnect and retry. An accepted response is never automatically replayed. +- An interrupted send, socket EOF, receive failure, or timeout after sending is + **not replayed**: the server may already have accepted the request. +- Visible WebSocket output is never superseded/replayed, even when the consumer + opts into HTTP response-attempt supersession. No WebSocket idempotency guarantee + is assumed. Wrapped error statuses and allowlisted retry headers are retained + in normal retry observations; raw provider error messages are not exposed. +- Dropping a pending `next_event` future during authentication refresh or a + WebSocket opening operation makes a retained turn fail safely on its next + poll. It does not resume with stale credentials or replay an uncertain send. + +Request, per-attempt, aggregate wire, field, and item bounds remain in force. +Frame/message sizes are bounded before JSON decoding; binary messages are +rejected and consecutive control frames are bounded. Upgrade and send operations +have a 30-second ceiling; configured attempt, idle, logical retry-budget, and +cancellation bounds also apply. As with HTTP, configure `with_resilience` to set +stream idle and whole-turn deadlines. + +### Current limitations + +The **full credential-bound transcript is always authoritative and always sent**. +This release deliberately does not send `previous_response_id` or maintain an +incremental response-ID cache: reconstructing an exact prefix from normalized +text/tool/reasoning/image output without losing provider fields requires a +separate lossless compatibility proof. Reuse therefore saves connection setup, +not request transcript bytes. Compaction, changed inputs, and reconnection do not +risk a stale server-side prefix. + +WebSocket upgrades use a dedicated reqwest HTTP/1 client with redirects and +implicit HTTP retries disabled, using the existing reqwest TLS stack. A custom +`Http` passed to `with_client` applies only to HTTP/SSE, **not** to WebSocket +upgrades; applications requiring custom transport middleware should keep `Http`. + +Wire contract reference: OpenAI Codex commit +[`6824dabe0393337a38cb257d5fe75ae5ca168470`](https://github.com/openai/codex/tree/6824dabe0393337a38cb257d5fe75ae5ca168470), +`codex-rs/codex-api/src/endpoint/responses_websocket.rs` and `common.rs`. + ## Retry observations and typed failures Responses emits `agentkit_loop::ProviderRetryEvent` through diff --git a/crates/agentkit-provider-openai/src/lib.rs b/crates/agentkit-provider-openai/src/lib.rs index 2847e83..0916972 100644 --- a/crates/agentkit-provider-openai/src/lib.rs +++ b/crates/agentkit-provider-openai/src/lib.rs @@ -31,7 +31,7 @@ mod responses; pub use responses::{ OpenAIResponsesAdapter, OpenAIResponsesConfig, OpenAIResponsesError, OpenAIResponsesLimits, OpenAIResponsesProfile, OpenAIResponsesRequestPolicy, OpenAIResponsesSession, - OpenAIResponsesTurn, + OpenAIResponsesTransport, OpenAIResponsesTurn, }; use std::fmt; diff --git a/crates/agentkit-provider-openai/src/responses.rs b/crates/agentkit-provider-openai/src/responses.rs index 9c5dc57..16090b3 100644 --- a/crates/agentkit-provider-openai/src/responses.rs +++ b/crates/agentkit-provider-openai/src/responses.rs @@ -32,6 +32,7 @@ use thiserror::Error; use zeroize::{Zeroize, Zeroizing}; mod retry; +mod websocket; use retry::{RetryTracker, local_error, provider_error, stream_classification}; const PUBLIC_ENDPOINT: &str = "https://api.openai.com/v1/responses"; @@ -56,6 +57,17 @@ pub enum OpenAIResponsesProfile { ChatGptPrivate, } +/// Transport selection for a Responses session. HTTP/SSE remains the default. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum OpenAIResponsesTransport { + #[default] + Http, + /// Require a real WebSocket upgrade; never fall back to HTTP. + WebSocket, + /// Try WebSocket, then use HTTP for this session after HTTP 426. + Auto, +} + /// Bounds serialized requests, streamed responses, counts, and individual fields. #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub struct OpenAIResponsesLimits { @@ -132,6 +144,7 @@ pub struct OpenAIResponsesConfig { limits: OpenAIResponsesLimits, user_agent: Option, originator: Option, + transport: OpenAIResponsesTransport, } impl fmt::Debug for OpenAIResponsesConfig { @@ -141,6 +154,7 @@ impl fmt::Debug for OpenAIResponsesConfig { .field("model", &self.model) .field("endpoint", &self.endpoint) .field("profile", &self.profile) + .field("transport", &self.transport) .field("header_names", &self.headers.keys().collect::>()) .field("request_policy", &self.request_policy) .field("reasoning_effort", &self.reasoning_effort) @@ -182,6 +196,7 @@ impl OpenAIResponsesConfig { limits: OpenAIResponsesLimits::default(), user_agent: None, originator: None, + transport: OpenAIResponsesTransport::Http, } } @@ -207,9 +222,16 @@ impl OpenAIResponsesConfig { limits: OpenAIResponsesLimits::default(), user_agent: None, originator: None, + transport: OpenAIResponsesTransport::Http, } } + /// Selects the transport without changing request or authentication semantics. + pub fn with_transport(mut self, transport: OpenAIResponsesTransport) -> Self { + self.transport = transport; + self + } + pub fn with_endpoint(mut self, endpoint: impl Into) -> Self { self.endpoint = endpoint.into(); self @@ -352,6 +374,7 @@ impl ModelAdapter for OpenAIResponsesAdapter { config: self.config.clone(), session: config, retry_observer: None, + websocket: Arc::new(Mutex::new(websocket::Session::default())), }) } @@ -366,6 +389,7 @@ pub struct OpenAIResponsesSession { config: Arc, session: SessionConfig, retry_observer: Option>, + websocket: Arc>, } #[async_trait] @@ -381,6 +405,11 @@ impl ModelSession for OpenAIResponsesSession { request: TurnRequest, cancellation: Option, ) -> Result { + let websocket = if self.config.transport != OpenAIResponsesTransport::Http { + Some(websocket::Lease::checkout(self.websocket.clone())?) + } else { + None + }; let mut tracker = RetryTracker::new(self.config.profile, self.retry_observer.clone()); let prepared = async { if cancelled(cancellation.as_ref()) { @@ -455,6 +484,8 @@ impl ModelSession for OpenAIResponsesSession { refreshed: false, wire_bytes: 0, tracker, + websocket, + websocket_sent: false, }, supersession_enabled, cancellation.as_ref(), @@ -479,6 +510,7 @@ pub struct OpenAIResponsesTurn { attempt_output_emitted: bool, pending_reopen: bool, pending_delay: Duration, + interrupted_operation: Option, finished: bool, } @@ -488,6 +520,7 @@ impl ModelTurn for OpenAIResponsesTurn { if !self.finished { self.finished = true; self.attempt = None; + self.context.websocket = None; let _ = self.context.tracker.finish(LoopError::Cancelled); } } @@ -503,12 +536,19 @@ impl ModelTurn for OpenAIResponsesTurn { Ok(event) => { if matches!(event, Some(ModelTurnEvent::Finished(_))) { self.context.tracker.succeed(); + if let Some(mut lease) = self.context.websocket.take() + && let Some(attempt) = self.attempt.take() + && let LiveBody::WebSocket(connection) = attempt.body + { + lease.complete(connection, attempt.decoder.state.response_id.as_deref()); + } } Ok(event) } Err(error) => { self.finished = true; self.attempt = None; + self.context.websocket = None; Err(self.context.tracker.finish(error)) } } @@ -528,6 +568,7 @@ impl OpenAIResponsesTurn { attempt_output_emitted: false, pending_reopen: true, pending_delay: Duration::ZERO, + interrupted_operation: None, finished: false, }; if let Err(error) = turn.reopen(cancellation).await { @@ -537,9 +578,13 @@ impl OpenAIResponsesTurn { } async fn reopen(&mut self, cancellation: Option<&TurnCancellation>) -> Result<(), LoopError> { + // Dropping a pending next_event future can interrupt a send without + // dropping the owned turn. Do not replay that ambiguous attempt on repoll. + if self.context.websocket_sent { + return Err(local_error(ProviderFailureReason::ReplayUnsafe)); + } if !self.pending_delay.is_zero() { let delay = self.pending_delay; - self.pending_delay = Duration::ZERO; cancellable( run_bounded_http( async { @@ -554,9 +599,17 @@ impl OpenAIResponsesTurn { ) .await? .map_err(http_loop_error)?; + self.pending_delay = Duration::ZERO; self.context.tracker.completed_wait(delay); } + // An owned turn can outlive a dropped next_event future. Conservatively + // terminate an interrupted WS opening operation on repoll: it may have + // suspended in authentication, backoff, or a partially delivered send. + if self.context.websocket.is_some() { + self.interrupted_operation = Some(ProviderFailureReason::ReplayUnsafe); + } self.attempt = Some(open_live_attempt(&mut self.context, cancellation).await?); + self.interrupted_operation = None; self.pending_reopen = false; Ok(()) } @@ -566,6 +619,9 @@ impl OpenAIResponsesTurn { cancellation: Option<&TurnCancellation>, ) -> Result, LoopError> { loop { + if let Some(reason) = self.interrupted_operation { + return Err(local_error(reason)); + } if cancelled(cancellation) { return Err(LoopError::Cancelled); } @@ -644,7 +700,7 @@ impl OpenAIResponsesTurn { .min(); let chunk = { let attempt = self.attempt.as_mut().expect("live attempt is open"); - cancellable(next_body_chunk(&mut attempt.body, timeout), cancellation).await + cancellable(attempt.body.next_chunk(timeout), cancellation).await }; match chunk { Err(error) => Err(nonretryable(error)), @@ -683,7 +739,16 @@ impl OpenAIResponsesTurn { } else { let attempt = self.attempt.as_mut().expect("live attempt is open"); attempt.truncated.observe(&chunk); - attempt.decoder.push(&chunk) + match &attempt.body { + LiveBody::Http(_) => attempt.decoder.push(&chunk), + LiveBody::WebSocket(_) => { + let result = attempt.decoder.push_json(&chunk); + if result.is_ok() && attempt.decoder.state.terminal { + attempt.eof = true; + } + result + } + } } } else { Err(protocol_failure("Responses wire-byte count overflowed")) @@ -698,6 +763,43 @@ impl OpenAIResponsesTurn { }; if let Err(failure) = result { self.context.tracker.note_failure(&failure.error); + let websocket = self + .attempt + .as_ref() + .is_some_and(|attempt| matches!(attempt.body, LiveBody::WebSocket(_))); + // A sent WebSocket request has no idempotency guarantee. Only explicit + // provider rejections before visible output can be retried. + if websocket + && (self.attempt_output_emitted + || !self + .attempt + .as_ref() + .is_some_and(|attempt| attempt.decoder.request_rejected)) + { + return Err(*failure.error); + } + if websocket { + self.context.websocket_sent = false; + } + if websocket + && failure + .error + .provider_failure() + .is_some_and(|f| f.upstream.http_status == Some(401)) + && !self.context.refreshed + { + self.context + .tracker + .scheduled(&failure.error, Duration::ZERO); + // Commit a safe repoll state before discarding the attempt + // and awaiting a refresh that may itself be cancelled by Drop. + self.interrupted_operation = Some(ProviderFailureReason::Authentication); + self.attempt = None; + refresh_authentication(&mut self.context, cancellation).await?; + self.pending_reopen = true; + self.interrupted_operation = None; + continue; + } if !failure.retryable || self.context.retries >= self @@ -1634,10 +1736,31 @@ struct ResponsesRequestContext { refreshed: bool, wire_bytes: usize, tracker: RetryTracker, + websocket: Option, + websocket_sent: bool, +} + +enum LiveBody { + Http(agentkit_http::BodyStream), + WebSocket(Box), +} + +impl LiveBody { + async fn next_chunk( + &mut self, + timeout: Option, + ) -> Result, HttpError> { + match self { + Self::Http(body) => next_body_chunk(body, timeout).await, + Self::WebSocket(connection) => { + run_bounded_http(connection.recv(), timeout, None, "WebSocket receive").await + } + } + } } struct LiveAttempt { - body: agentkit_http::BodyStream, + body: LiveBody, truncated: TruncatedStreamDetector, decoder: ResponsesSseDecoder, deadline: Option, @@ -1670,6 +1793,43 @@ fn stopped_attempt(context: &ResponsesRequestContext, failure: AttemptFailure) - } } +async fn refresh_authentication( + context: &mut ResponsesRequestContext, + cancellation: Option<&TurnCancellation>, +) -> Result<(), LoopError> { + let binding = context.auth.binding().map(str::to_owned); + let refreshed = cancellable( + run_bounded_http( + context + .config + .authentication + .authenticate(Some(&context.auth)), + context + .config + .resilience + .as_ref() + .and_then(|config| config.attempt_timeout), + context.deadline.as_ref(), + "OpenAI reauthentication", + ), + cancellation, + ) + .await? + .map_err(|error| match error { + HttpError::Timeout { + operation: "logical request retry budget", + .. + } => http_loop_error(error), + _ => local_error(ProviderFailureReason::Authentication), + })?; + if refreshed.binding() != binding.as_deref() { + return Err(local_error(ProviderFailureReason::Authentication)); + } + context.auth = refreshed; + context.refreshed = true; + Ok(()) +} + async fn open_live_attempt( context: &mut ResponsesRequestContext, cancellation: Option<&TurnCancellation>, @@ -1682,13 +1842,20 @@ async fn open_live_attempt( .and_then(|config| config.attempt_timeout); let attempt_deadline = attempt_timeout.map(LogicalDeadline::new); let logical_deadline = context.deadline.clone(); - let result = attempt_with_timeout( + context.websocket_sent = false; + let mut result = attempt_with_timeout( send_live_attempt(context, cancellation), attempt_timeout, logical_deadline.as_ref(), cancellation, ) .await; + if context.websocket_sent + && let Err(failure) = &mut result + { + // A timeout may have interrupted the send future after bytes left. + failure.retryable = false; + } if let Err(failure) = &result { context.tracker.note_failure(&failure.error); } @@ -1699,36 +1866,7 @@ async fn open_live_attempt( } Err(failure) if is_unauthorized(&failure.error) && !context.refreshed => { context.tracker.scheduled(&failure.error, Duration::ZERO); - let binding = context.auth.binding().map(str::to_owned); - let refreshed = cancellable( - run_bounded_http( - context - .config - .authentication - .authenticate(Some(&context.auth)), - context - .config - .resilience - .as_ref() - .and_then(|config| config.attempt_timeout), - context.deadline.as_ref(), - "OpenAI reauthentication", - ), - cancellation, - ) - .await? - .map_err(|error| match error { - HttpError::Timeout { - operation: "logical request retry budget", - .. - } => http_loop_error(error), - _ => local_error(ProviderFailureReason::Authentication), - })?; - if refreshed.binding() != binding.as_deref() { - return Err(local_error(ProviderFailureReason::Authentication)); - } - context.auth = refreshed; - context.refreshed = true; + refresh_authentication(context, cancellation).await?; } Err(failure) if failure.retryable @@ -1831,6 +1969,15 @@ async fn send_live_attempt( None }; headers.extend(context.auth.headers().clone()); + if context + .websocket + .as_ref() + .is_some_and(|lease| !lease.http_only()) + && let Some(attempt) = websocket::send(context, headers.clone()).await? + { + return Ok(attempt); + } + // Auto: rejected upgrade (426), never a sent response.create. let request = context .client .post(&context.config.endpoint) @@ -1904,7 +2051,7 @@ async fn send_live_attempt( } let truncated = TruncatedStreamDetector::from_headers(response.headers()); Ok(LiveAttempt { - body: response.bytes_stream(), + body: LiveBody::Http(response.bytes_stream()), truncated, decoder: ResponsesSseDecoder::with_policy( &context.config.model, @@ -1988,6 +2135,7 @@ struct ResponsesSseDecoder { buffer: Zeroizing>, buffer_start: usize, received: usize, + request_rejected: bool, max_attempt_bytes: usize, state: ResponsesState, } @@ -2019,6 +2167,7 @@ impl ResponsesSseDecoder { buffer: Zeroizing::new(Vec::new()), buffer_start: 0, received: 0, + request_rejected: false, max_attempt_bytes: limits.max_attempt_bytes, state: ResponsesState::new( model, @@ -2152,7 +2301,21 @@ impl ResponsesSseDecoder { "Responses SSE used an unsupported terminal marker", )); } - let mut value: Value = serde_json::from_str(data.as_str()) + self.consume_json(data.as_bytes(), event) + } + + fn push_json(&mut self, bytes: &[u8]) -> Result<(), AttemptFailure> { + self.received = self.received.saturating_add(bytes.len()); + if self.received > self.max_attempt_bytes { + return Err(protocol_failure( + "Responses WebSocket attempt exceeds byte limit", + )); + } + self.consume_json(bytes, None) + } + + fn consume_json(&mut self, data: &[u8], event: Option<&str>) -> Result<(), AttemptFailure> { + let mut value: Value = serde_json::from_slice(data) .map_err(|_| protocol_failure("Responses SSE data is malformed JSON"))?; let kind = value .get("type") @@ -2163,7 +2326,40 @@ impl ResponsesSseDecoder { zeroize_encrypted_content(&mut value); return Err(protocol_failure("Responses SSE event name/type mismatch")); } - let result = self.state.consume(&kind, &value); + let mut result = self.state.consume(&kind, &value); + if kind == "error" + && let Err(failure) = &mut result + { + if let Some(status) = value + .get("status") + .or_else(|| value.get("status_code")) + .and_then(Value::as_u64) + .and_then(|s| u16::try_from(s).ok()) + .and_then(|s| StatusCode::from_u16(s).ok()) + { + // A wrapped HTTP error before response.created is an explicit + // request rejection. A failure of an accepted response is not. + self.request_rejected = !self.state.created && !status.is_success(); + let mut classification = stream_classification(&value, &kind); + classification.http_status = Some(status.as_u16()); + *failure.error = + provider_error(ProviderFailureReason::ResponseFailed, classification); + failure.retryable = retryable_response_status(status, self.state.profile); + } + if let Some(headers) = value.get("headers").and_then(Value::as_object) { + let mut parsed = HeaderMap::new(); + for name in ["retry-after", "retry-after-ms"] { + if let Some(value) = headers + .get(name) + .and_then(Value::as_str) + .and_then(|v| HeaderValue::from_str(v).ok()) + { + parsed.insert(name, value); + } + } + failure.headers = retry_headers(&parsed); + } + } zeroize_encrypted_content(&mut value); result } diff --git a/crates/agentkit-provider-openai/src/responses/websocket.rs b/crates/agentkit-provider-openai/src/responses/websocket.rs new file mode 100644 index 0000000..4238d3e --- /dev/null +++ b/crates/agentkit-provider-openai/src/responses/websocket.rs @@ -0,0 +1,339 @@ +//! A session owns at most one leased socket. No lock is held across I/O. +use super::*; +use futures_util::{FutureExt, SinkExt, StreamExt}; +use tokio_tungstenite::{ + WebSocketStream, + tungstenite::{ + Message, + handshake::{client::generate_key, derive_accept_key}, + protocol::{Role, WebSocketConfig}, + }, +}; + +const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(30); + +#[derive(Default)] +pub(super) struct Session { + busy: bool, + http_only: bool, + idle: Option>, +} + +/// Writers: checkout, sticky fallback, completion and Drop. Checkout commits the +/// busy bit and takes the socket under one guard. Only that lease can return it. +/// Failure/cancellation/unwind drops the private socket; Drop clears the claim. +/// Poison is fail-closed, never recovered. No callbacks/I/O/destructors run under +/// the guard. A completed owned turn releases its claim before returning Finished. +pub(super) struct Lease { + session: Arc>, + connection: Option>, + reusable: bool, + http_only: bool, +} + +impl Lease { + pub(super) fn checkout(session: Arc>) -> Result { + let (connection, http_only) = { + let mut state = session + .lock() + .map_err(|_| local_error(ProviderFailureReason::Protocol))?; + if state.busy { + return Err(LoopError::Provider( + "Responses WebSocket session already has an active turn".into(), + )); + } + state.busy = true; + (state.idle.take(), state.http_only) + }; + Ok(Self { + session, + connection, + reusable: false, + http_only, + }) + } + + pub(super) fn http_only(&self) -> bool { + self.http_only + } + + fn fallback(&mut self) -> Result<(), AttemptFailure> { + self.session + .lock() + .map_err(|_| protocol_failure("WebSocket session lock poisoned"))? + .http_only = true; + self.http_only = true; + Ok(()) + } + + pub(super) fn complete(&mut self, mut connection: Box, response_id: Option<&str>) { + // Keep bounded replay correlation, not a server-side continuation cache. + if let Some(id) = response_id + && connection.completed_ids.len() < 1024 + { + connection.completed_ids.insert(id.to_owned()); + self.connection = Some(connection); + self.reusable = true; + } + } +} + +impl Drop for Lease { + fn drop(&mut self) { + let returned = if self.reusable { + self.connection.take() + } else { + None + }; + // No poison recovery: a poisoned session rejects all future checkouts. + let old = if let Ok(mut state) = self.session.lock() { + state.busy = false; + std::mem::replace(&mut state.idle, returned) + } else { + returned + }; + drop(old); + } +} + +pub(super) struct Connection { + stream: WebSocketStream, + auth_headers: HeaderMap, + binding: Option, + completed_ids: BTreeSet, +} + +impl Connection { + pub(super) async fn recv(&mut self) -> Result, HttpError> { + // Bound control traffic as well as data, including peers that only ping. + for _ in 0..128 { + match self.stream.next().await { + Some(Ok(Message::Text(text))) => { + if !self.completed_ids.is_empty() { + let mut value: Value = serde_json::from_str(&text) + .map_err(|_| HttpError::Other("invalid WebSocket JSON".into()))?; + let stale = value + .pointer("/response/id") + .or_else(|| value.get("response_id")) + .and_then(Value::as_str) + .is_some_and(|id| self.completed_ids.contains(id)); + zeroize_encrypted_content(&mut value); + if stale { + return Err(HttpError::Other( + "stale Responses WebSocket response".into(), + )); + } + } + return Ok(Some(agentkit_http::Bytes::copy_from_slice(text.as_bytes()))); + } + Some(Ok(Message::Ping(_))) => self + .stream + .flush() + .await + .map_err(|_| HttpError::Other("WebSocket pong failed".into()))?, + Some(Ok(Message::Pong(_))) => {} + Some(Ok(Message::Close(_))) | None => return Ok(None), + _ => { + return Err(HttpError::Other( + "invalid or failed Responses WebSocket message".into(), + )); + } + } + } + Err(HttpError::Other( + "excessive WebSocket control frames".into(), + )) + } +} + +pub(super) async fn send( + context: &mut ResponsesRequestContext, + mut headers: HeaderMap, +) -> Result, AttemptFailure> { + let previous = context + .websocket + .as_mut() + .expect("WebSocket lease") + .connection + .take(); + let previous = previous.and_then(|mut connection| { + // Polling mutates the framed reader, so Option::filter cannot be used. + // Unsolicited data/close/error at a boundary forces a fresh connection. + let clean = connection.stream.next().now_or_never().is_none(); + clean.then_some(connection) + }); + let mut connection = if let Some(connection) = previous.filter(|connection| { + connection.auth_headers == *context.auth.headers() + && connection.binding.as_deref() == context.auth.binding() + }) { + context.tracker.accounting.attempts = context.tracker.accounting.attempts.saturating_add(1); + connection + } else { + headers.remove("idempotency-key"); // WebSocket response.create has no idempotency contract. + headers.remove("accept"); + headers.remove("content-type"); + headers.insert( + "openai-beta", + HeaderValue::from_static("responses_websockets=2026-02-06"), + ); + let key = generate_key(); + headers.insert("connection", HeaderValue::from_static("Upgrade")); + headers.insert("upgrade", HeaderValue::from_static("websocket")); + headers.insert("sec-websocket-version", HeaderValue::from_static("13")); + headers.insert( + "sec-websocket-key", + HeaderValue::from_str(&key).map_err(|_| protocol_failure("invalid WebSocket nonce"))?, + ); + let client = reqwest::Client::builder() + .http1_only() + .redirect(reqwest::redirect::Policy::none()) + .retry(reqwest::retry::never()) + .connect_timeout(HANDSHAKE_TIMEOUT) + .timeout(HANDSHAKE_TIMEOUT) + .build() + .map_err(|_| protocol_failure("could not build WebSocket upgrade client"))?; + context.tracker.accounting.attempts = context.tracker.accounting.attempts.saturating_add(1); + // HTTP/1 upgrade on https uses reqwest's existing rustls trust/proxy stack; + // it is the same wire operation as connecting to the corresponding wss URL. + let response = client + .get(&context.config.endpoint) + .headers(headers) + .send() + .await + .map_err(|error| transport_failure(HttpError::request(error)))?; + let status = response.status(); + if status == StatusCode::UPGRADE_REQUIRED + && context.config.transport == OpenAIResponsesTransport::Auto + { + context + .websocket + .as_mut() + .expect("WebSocket lease") + .fallback()?; + return Ok(None); + } + if status != StatusCode::SWITCHING_PROTOCOLS { + return Err(AttemptFailure { + error: Box::new(provider_error( + ProviderFailureReason::HttpStatus, + ProviderClassification { + http_status: Some(status.as_u16()), + ..ProviderClassification::default() + }, + )), + retryable: retryable_response_status(status, context.config.profile), + headers: retry_headers(response.headers()), + }); + } + validate_handshake(response.headers(), &key)?; + let turn_state = if context.config.profile == OpenAIResponsesProfile::ChatGptPrivate { + validated_turn_state_header(response.headers())? + } else { + None + }; + if let Some(captured) = turn_state { + let mut state = context + .turn_state + .lock() + .map_err(|_| protocol_failure("turn-state lock poisoned"))?; + if state.as_ref().is_some_and(|expected| expected != captured) { + return Err(protocol_failure("provider changed x-codex-turn-state")); + } + *state = Some(captured); + } + let socket = response + .upgrade() + .await + .map_err(|_| protocol_failure("WebSocket upgrade failed"))?; + let limit = context.config.limits.max_attempt_bytes; + let config = WebSocketConfig::default() + .max_message_size(Some(limit)) + .max_frame_size(Some(limit)) + .write_buffer_size(0); + Box::new(Connection { + stream: WebSocketStream::from_raw_socket(socket, Role::Client, Some(config)).await, + auth_headers: context.auth.headers().clone(), + binding: context.auth.binding().map(str::to_owned), + completed_ids: BTreeSet::new(), + }) + }; + // Always send the authoritative, credential-bound full transcript. Incremental + // previous_response_id is intentionally not used without a lossless prefix proof. + let mut value: Value = serde_json::from_slice(&context.body) + .map_err(|_| protocol_failure("invalid encoded Responses request"))?; + let fields = value + .as_object_mut() + .ok_or_else(|| protocol_failure("invalid Responses request object"))?; + // Public Responses WebSocket mode excludes HTTP transport controls. + // https://developers.openai.com/api/docs/guides/websocket-mode + fields.remove("stream"); + fields.remove("background"); + fields.insert("type".into(), json!("response.create")); + let serialized = serde_json::to_string(&value); + zeroize_encrypted_content(&mut value); + let request = Zeroizing::new( + serialized.map_err(|_| protocol_failure("could not serialize WebSocket request"))?, + ); + if request.len() > context.config.limits.max_request_bytes { + return Err(protocol_failure( + "Responses WebSocket request exceeds byte limit", + )); + } + // Once send is polled, delivery is ambiguous. Never retry send/timeout errors. + context.websocket_sent = true; + run_bounded_http( + async { + connection + .stream + .send(Message::Text(request.as_str().into())) + .await + .map_err(|_| HttpError::Other("WebSocket request send failed".into())) + }, + Some(HANDSHAKE_TIMEOUT), + context.deadline.as_ref(), + "WebSocket send", + ) + .await + .map_err(|e| nonretryable(http_loop_error(e)))?; + Ok(Some(LiveAttempt { + body: LiveBody::WebSocket(connection), + truncated: TruncatedStreamDetector::from_headers(&HeaderMap::new()), + decoder: ResponsesSseDecoder::with_policy( + &context.config.model, + &context.session_id, + context.config.profile, + context.config.request_policy.include_encrypted_reasoning, + context.auth.binding(), + context.turn_state.clone(), + context.config.limits, + ), + deadline: None, + eof: false, + closed: false, + })) +} + +fn validate_handshake(headers: &HeaderMap, key: &str) -> Result<(), AttemptFailure> { + let has_token = |name: &str, token: &str| { + headers + .get_all(name) + .iter() + .filter_map(|v| v.to_str().ok()) + .flat_map(|v| v.split(',')) + .any(|v| v.trim().eq_ignore_ascii_case(token)) + }; + let accepts: Vec<_> = headers.get_all("sec-websocket-accept").iter().collect(); + if !has_token("connection", "upgrade") + || !has_token("upgrade", "websocket") + || accepts.len() != 1 + || accepts[0].as_bytes() != derive_accept_key(key.as_bytes()).as_bytes() + || headers.contains_key("sec-websocket-extensions") + || headers.contains_key("sec-websocket-protocol") + { + return Err(protocol_failure("invalid Responses WebSocket handshake")); + } + Ok(()) +} + +#[cfg(test)] +mod tests; diff --git a/crates/agentkit-provider-openai/src/responses/websocket/tests.rs b/crates/agentkit-provider-openai/src/responses/websocket/tests.rs new file mode 100644 index 0000000..a51e597 --- /dev/null +++ b/crates/agentkit-provider-openai/src/responses/websocket/tests.rs @@ -0,0 +1,763 @@ +//! Loopback-only transport tests; no live inference. +use super::*; +use agentkit_core::{SessionId, TurnId}; +use std::io::{Read, Write}; +use std::net::{TcpListener, TcpStream}; +use std::thread; +use tokio_tungstenite::tungstenite::{self, WebSocket}; + +// Same text/tool/reasoning/usage fixture as the HTTP decoder tests. +const SUCCESS: &str = r#"event: response.created +data: {"type":"response.created","sequence_number":1,"response":{"id":"resp-1","model":"gpt-test"}} + +event: response.output_item.added +data: {"type":"response.output_item.added","sequence_number":2,"output_index":0,"item":{"id":"msg-1","type":"message"}} + +event: response.content_part.added +data: {"type":"response.content_part.added","sequence_number":3,"item_id":"msg-1","output_index":0,"content_index":0,"part":{"type":"output_text"}} + +event: response.output_text.delta +data: {"type":"response.output_text.delta","sequence_number":4,"item_id":"msg-1","output_index":0,"content_index":0,"delta":"hello"} + +event: response.output_text.done +data: {"type":"response.output_text.done","sequence_number":5,"item_id":"msg-1","output_index":0,"content_index":0,"text":"hello"} + +event: response.content_part.done +data: {"type":"response.content_part.done","sequence_number":6,"item_id":"msg-1","output_index":0,"content_index":0,"part":{"type":"output_text","text":"hello"}} + +event: response.output_item.done +data: {"type":"response.output_item.done","sequence_number":7,"output_index":0,"item":{"id":"msg-1","type":"message","role":"assistant","content":[{"type":"output_text","text":"hello"}]}} + +event: response.output_item.added +data: {"type":"response.output_item.added","sequence_number":8,"output_index":1,"item":{"id":"reason-1","type":"reasoning"}} + +event: response.reasoning_summary_part.added +data: {"type":"response.reasoning_summary_part.added","sequence_number":9,"item_id":"reason-1","output_index":1,"summary_index":0,"part":{"type":"summary_text"}} + +event: response.reasoning_summary_text.delta +data: {"type":"response.reasoning_summary_text.delta","sequence_number":10,"item_id":"reason-1","output_index":1,"summary_index":0,"delta":"brief"} + +event: response.reasoning_summary_text.done +data: {"type":"response.reasoning_summary_text.done","sequence_number":11,"item_id":"reason-1","output_index":1,"summary_index":0,"text":"brief"} + +event: response.reasoning_summary_part.done +data: {"type":"response.reasoning_summary_part.done","sequence_number":12,"item_id":"reason-1","output_index":1,"summary_index":0,"part":{"type":"summary_text","text":"brief"}} + +event: response.output_item.done +data: {"type":"response.output_item.done","sequence_number":13,"output_index":1,"item":{"id":"reason-1","type":"reasoning","summary":[{"type":"summary_text","text":"brief"}],"encrypted_content":"opaque"}} + +event: response.output_item.added +data: {"type":"response.output_item.added","sequence_number":14,"output_index":2,"item":{"id":"call-item","type":"function_call"}} + +event: response.function_call_arguments.delta +data: {"type":"response.function_call_arguments.delta","sequence_number":15,"item_id":"call-item","output_index":2,"delta":"{\"q\":1}"} + +event: response.function_call_arguments.done +data: {"type":"response.function_call_arguments.done","sequence_number":16,"item_id":"call-item","output_index":2,"arguments":"{\"q\":1}"} + +event: response.output_item.done +data: {"type":"response.output_item.done","sequence_number":17,"output_index":2,"item":{"id":"call-item","type":"function_call","call_id":"call-1","name":"lookup","arguments":"{\"q\":1}"}} + +event: response.completed +data: {"type":"response.completed","sequence_number":18,"response":{"id":"resp-1","model":"gpt-test","usage":{"input_tokens":3,"output_tokens":5,"output_tokens_details":{"reasoning_tokens":2}}}} + +"#; + +const WAIT: Duration = Duration::from_secs(5); +type Socket = WebSocket; + +fn server(run: F) -> (String, thread::JoinHandle<()>) +where + F: FnOnce(TcpListener) + Send + 'static, +{ + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let endpoint = format!("http://{}/v1/responses", listener.local_addr().unwrap()); + listener.set_nonblocking(true).unwrap(); + (endpoint, thread::spawn(move || run(listener))) +} +fn accept(listener: &TcpListener) -> TcpStream { + let start = Instant::now(); + loop { + match listener.accept() { + Ok((stream, _)) => { + stream.set_nonblocking(false).unwrap(); + stream.set_read_timeout(Some(WAIT)).unwrap(); + stream.set_write_timeout(Some(WAIT)).unwrap(); + return stream; + } + Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => { + assert!(start.elapsed() < WAIT, "missing expected connection"); + thread::sleep(Duration::from_millis(2)); + } + Err(e) => panic!("accept: {e}"), + } + } +} +fn socket(listener: &TcpListener) -> Socket { + tungstenite::accept(accept(listener)).unwrap() +} +fn receive(ws: &mut Socket) -> Value { + let message = ws.read().unwrap(); + let value: Value = serde_json::from_str(message.to_text().unwrap()).unwrap(); + assert_eq!(value["type"], "response.create"); + assert!(value.get("previous_response_id").is_none()); + assert!( + value.get("stream").is_none(), + "HTTP stream field on WS wire" + ); + assert!( + value.get("background").is_none(), + "HTTP background field on WS wire" + ); + assert!(value.get("model").is_some_and(Value::is_string)); + assert!(value.get("input").is_some_and(Value::is_array)); + value +} +fn success(ws: &mut Socket) { + success_id(ws, "resp-1"); +} +fn success_id(ws: &mut Socket, id: &str) { + for line in SUCCESS + .replace("resp-1", id) + .lines() + .filter_map(|line| line.strip_prefix("data: ")) + { + ws.send(Message::Text(line.into())).unwrap(); + } +} +// Read one request without over-reading into the next. Close each HTTP response. +fn http(listener: &TcpListener, status: &str, body: &str) -> (String, Value) { + let mut stream = accept(listener); + let mut raw = Vec::new(); + while !raw.ends_with(b"\r\n\r\n") { + let mut byte = [0]; + stream.read_exact(&mut byte).unwrap(); + raw.push(byte[0]); + assert!(raw.len() < 64 * 1024); + } + let headers = String::from_utf8(raw).unwrap(); + let length = headers + .lines() + .find_map(|line| { + let (key, value) = line.split_once(':')?; + key.eq_ignore_ascii_case("content-length") + .then(|| value.trim().parse::().unwrap()) + }) + .unwrap_or(0); + let mut request = vec![0; length]; + stream.read_exact(&mut request).unwrap(); + write!(stream, "HTTP/1.1 {status}\r\nContent-Type: text/event-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}", body.len()).unwrap(); + ( + headers, + serde_json::from_slice(&request).unwrap_or(Value::Null), + ) +} +fn config(endpoint: &str, transport: OpenAIResponsesTransport) -> OpenAIResponsesConfig { + OpenAIResponsesConfig::new("loopback-test", "gpt-test") + .with_endpoint(endpoint) + .with_transport(transport) + .with_resilience(ResilienceConfig { + max_retries: 0, + retry_budget: WAIT, + attempt_timeout: Some(Duration::from_secs(2)), + stream_idle_timeout: Some(Duration::from_secs(2)), + initial_backoff: Duration::ZERO, + max_backoff: Duration::ZERO, + }) +} +fn request() -> TurnRequest { + TurnRequest { + session_id: SessionId::new("session"), + turn_id: TurnId::new("turn"), + transcript: vec![Item::text(ItemKind::User, "hello")], + available_tools: vec![], + cache: None, + metadata: MetadataMap::new(), + } +} +async fn drain(turn: &mut OpenAIResponsesTurn) -> Result, LoopError> { + tokio::time::timeout(WAIT, async { + let mut events = vec![]; + while let Some(event) = turn.next_event(None).await? { + events.push(event); + } + Ok(events) + }) + .await + .expect("turn hung") +} +fn finished(events: &[ModelTurnEvent]) { + assert_eq!( + events + .iter() + .filter(|event| matches!(event, ModelTurnEvent::Finished(_))) + .count(), + 1 + ); + assert!(events.iter().any(|event| matches!(event, + ModelTurnEvent::Delta(Delta::AppendText {chunk, ..}) if chunk == "hello"))); +} +#[tokio::test] +async fn reuses_completed_socket_releases_claim_and_isolates_sessions() { + let (endpoint, peer) = server(|listener| { + let mut first = socket(&listener); + let one = receive(&mut first); + success(&mut first); + let two = receive(&mut first); + assert_eq!(one["input"], two["input"]); + success_id(&mut first, "resp-2"); + let mut second = socket(&listener); + receive(&mut second); + success(&mut second); + }); + let adapter = + OpenAIResponsesAdapter::new(config(&endpoint, OpenAIResponsesTransport::WebSocket)) + .unwrap(); + let mut session = adapter + .start_session(SessionConfig::new("session")) + .await + .unwrap(); + let mut turn = session.begin_turn(request(), None).await.unwrap(); + assert!(session.begin_turn(request(), None).await.is_err()); + finished(&drain(&mut turn).await.unwrap()); + // Finished, not Drop, releases the claim: keep the completed turn alive. + let mut next = session.begin_turn(request(), None).await.unwrap(); + finished(&drain(&mut next).await.unwrap()); + let mut other = adapter + .start_session(SessionConfig::new("other")) + .await + .unwrap(); + let mut third = other.begin_turn(request(), None).await.unwrap(); + finished(&drain(&mut third).await.unwrap()); + peer.join().unwrap(); +} +#[tokio::test] +async fn http_and_websocket_have_identical_text_tool_reasoning_usage_events() { + let (endpoint, peer) = server(|listener| { + let mut ws = socket(&listener); + let ws_request = receive(&mut ws); + success(&mut ws); + let (headers, http_request) = http(&listener, "200 OK", SUCCESS); + assert!(headers.starts_with("POST /v1/responses")); + assert_eq!(http_request["stream"], true); + assert!(http_request.get("type").is_none()); + // Compare semantic request fields, not transport-specific wire envelopes. + for (key, value) in http_request.as_object().unwrap() { + if key != "stream" && key != "background" { + assert_eq!(ws_request.get(key), Some(value), "field {key}"); + } + } + for key in ws_request.as_object().unwrap().keys() { + assert!(key == "type" || http_request.get(key).is_some()); + } + }); + let mut outputs = vec![]; + let base = config(&endpoint, OpenAIResponsesTransport::WebSocket); + for transport in [ + OpenAIResponsesTransport::WebSocket, + OpenAIResponsesTransport::Http, + ] { + let adapter = OpenAIResponsesAdapter::new(base.clone().with_transport(transport)).unwrap(); + let mut session = adapter + .start_session(SessionConfig::new("session")) + .await + .unwrap(); + let mut turn = session.begin_turn(request(), None).await.unwrap(); + let events = drain(&mut turn).await.unwrap(); + finished(&events); + outputs.push(format!("{events:?}")); + } + assert_eq!(outputs[0], outputs[1]); + peer.join().unwrap(); +} +#[tokio::test] +async fn auto_426_fallback_is_sticky_but_explicit_websocket_fails() { + let (endpoint, peer) = server(|listener| { + let (headers, _) = http(&listener, "426 Upgrade Required", ""); + assert!(headers.starts_with("GET /v1/responses")); + for _ in 0..2 { + let (headers, _) = http(&listener, "200 OK", SUCCESS); + assert!( + headers.starts_with("POST /v1/responses"), + "fallback must remain sticky" + ); + } + let (headers, _) = http(&listener, "426 Upgrade Required", ""); + assert!(headers.starts_with("GET /v1/responses")); + }); + let adapter = + OpenAIResponsesAdapter::new(config(&endpoint, OpenAIResponsesTransport::Auto)).unwrap(); + let mut session = adapter + .start_session(SessionConfig::new("session")) + .await + .unwrap(); + for _ in 0..2 { + let mut turn = session.begin_turn(request(), None).await.unwrap(); + finished(&drain(&mut turn).await.unwrap()); + } + let adapter = + OpenAIResponsesAdapter::new(config(&endpoint, OpenAIResponsesTransport::WebSocket)) + .unwrap(); + let mut session = adapter + .start_session(SessionConfig::new("session")) + .await + .unwrap(); + assert!(session.begin_turn(request(), None).await.is_err()); + peer.join().unwrap(); +} +#[tokio::test] +async fn eof_malformed_binary_and_oversized_frames_never_complete() { + for frame in [ + None, + Some(Message::Text("not json".into())), + Some(Message::Binary(vec![0, 1].into())), + Some(Message::Text("x".repeat(1025).into())), + ] { + let (endpoint, peer) = server(move |listener| { + let mut ws = socket(&listener); + receive(&mut ws); + if let Some(frame) = frame { + ws.send(frame).unwrap(); + } + let _ = ws.close(None); + let mut fresh = socket(&listener); + receive(&mut fresh); + fresh.send(Message::Text(r#"{"type":"response.created","sequence_number":1,"response":{"id":"fresh","model":"gpt-test"}}"#.into())).unwrap(); + fresh.send(Message::Text(r#"{"type":"response.completed","sequence_number":2,"response":{"id":"fresh","model":"gpt-test","status":"completed","output":[]}}"#.into())).unwrap(); + }); + let mut cfg = config(&endpoint, OpenAIResponsesTransport::WebSocket); + cfg.limits.max_attempt_bytes = 1024; + cfg.limits.max_text_bytes = 1024; + let adapter = OpenAIResponsesAdapter::new(cfg).unwrap(); + let mut session = adapter + .start_session(SessionConfig::new("session")) + .await + .unwrap(); + let mut turn = session.begin_turn(request(), None).await.unwrap(); + assert!(drain(&mut turn).await.is_err()); + let mut fresh = session.begin_turn(request(), None).await.unwrap(); + let events = drain(&mut fresh).await.unwrap(); + assert!( + events + .iter() + .any(|e| matches!(e, ModelTurnEvent::Finished(_))) + ); + peer.join().unwrap(); + } +} + +#[tokio::test] +async fn dropped_and_cancelled_turns_discard_late_frames_and_release_claim() { + for cancel in [false, true] { + let (release, released) = std::sync::mpsc::channel(); + let (endpoint, peer) = server(move |listener| { + let mut old = socket(&listener); + receive(&mut old); + released.recv_timeout(WAIT).unwrap(); + // This old response must never satisfy the next turn, even if the + // kernel accepts the write after the client has dropped its socket. + let late = SUCCESS + .lines() + .find_map(|l| l.strip_prefix("data: ")) + .unwrap(); + let _ = old.send(Message::Text(late.replace("resp-1", "late").into())); + let mut fresh = socket(&listener); + receive(&mut fresh); + success_id(&mut fresh, "fresh"); + }); + let adapter = + OpenAIResponsesAdapter::new(config(&endpoint, OpenAIResponsesTransport::WebSocket)) + .unwrap(); + let mut session = adapter + .start_session(SessionConfig::new("session")) + .await + .unwrap(); + let mut turn = session.begin_turn(request(), None).await.unwrap(); + if cancel { + let handle = agentkit_core::CancellationController::new(); + let cancellation = handle.handle().checkpoint(); + handle.interrupt(); + assert!(turn.next_event(Some(cancellation)).await.is_err()); + } else { + drop(turn); + } + release.send(()).unwrap(); + let mut fresh = session.begin_turn(request(), None).await.unwrap(); + let events = drain(&mut fresh).await.unwrap(); + finished(&events); + assert!(events.iter().any(|event| matches!(event, + ModelTurnEvent::Finished(result) if result.response_id.as_deref() == Some("fresh")))); + peer.join().unwrap(); + } +} + +#[tokio::test] +async fn attempt_and_logical_deadlines_bound_silent_socket() { + for logical in [false, true] { + let (release, released) = std::sync::mpsc::channel(); + let (endpoint, peer) = server(move |listener| { + let mut ws = socket(&listener); + receive(&mut ws); + released.recv_timeout(WAIT).unwrap(); + }); + let mut cfg = config(&endpoint, OpenAIResponsesTransport::WebSocket); + cfg.resilience = Some(ResilienceConfig { + max_retries: 0, + retry_budget: if logical { + Duration::from_millis(100) + } else { + WAIT + }, + attempt_timeout: if logical { + None + } else { + Some(Duration::from_millis(100)) + }, + stream_idle_timeout: None, + initial_backoff: Duration::ZERO, + max_backoff: Duration::ZERO, + }); + let adapter = OpenAIResponsesAdapter::new(cfg).unwrap(); + let mut session = adapter + .start_session(SessionConfig::new("session")) + .await + .unwrap(); + let mut turn = session.begin_turn(request(), None).await.unwrap(); + let start = Instant::now(); + let error = drain(&mut turn).await.unwrap_err(); + assert!(start.elapsed() < Duration::from_secs(2)); + assert!(error.provider_failure().is_some()); + release.send(()).unwrap(); + peer.join().unwrap(); + } +} + +#[tokio::test] +async fn wrapped_retryable_error_reconnects_before_acceptance_but_not_after() { + for mode in 0..3 { + let accepted = mode != 0; + let (endpoint, peer) = server(move |listener| { + let mut ws = socket(&listener); + let original = receive(&mut ws); + if accepted { + ws.send(Message::Text(r#"{"type":"response.created","sequence_number":1,"response":{"id":"accepted","model":"gpt-test"}}"#.into())).unwrap(); + } + if mode == 2 { + ws.send(Message::Text(r#"{"type":"response.failed","sequence_number":2,"response":{"id":"accepted","error":{"code":"server_error","type":"server_error"}}}"#.into())).unwrap(); + } else { + ws.send(Message::Text(r#"{"type":"error","status":429,"error":{"type":"rate_limit_error","code":"rate_limit_exceeded","message":"retry later"}}"#.into())).unwrap(); + } + if !accepted { + let mut retry = socket(&listener); + assert_eq!(original, receive(&mut retry)); + success(&mut retry); + } + }); + let mut cfg = config(&endpoint, OpenAIResponsesTransport::WebSocket); + cfg.resilience.as_mut().unwrap().max_retries = 1; + let adapter = OpenAIResponsesAdapter::new(cfg).unwrap(); + let mut session = adapter + .start_session(SessionConfig::new("session")) + .await + .unwrap(); + let mut turn = session.begin_turn(request(), None).await.unwrap(); + let result = drain(&mut turn).await; + if accepted { + let error = result.unwrap_err(); + assert_eq!(error.provider_failure().unwrap().accounting.attempts, 1); + } else { + finished(&result.unwrap()); + } + peer.join().unwrap(); + } +} + +#[test] +fn poisoned_session_is_fail_closed() { + let session = Arc::new(Mutex::new(Session::default())); + let poisoned = session.clone(); + assert!( + std::panic::catch_unwind(move || { + let _guard = poisoned.lock().unwrap(); + panic!("poison the session"); + }) + .is_err() + ); + assert!(Lease::checkout(session.clone()).is_err()); + assert!(Lease::checkout(session).is_err()); +} + +#[tokio::test] +async fn observer_unwind_does_not_poison_session_or_strand_lease() { + use futures_util::FutureExt; + let (endpoint, peer) = server(|listener| { + let mut first = socket(&listener); + receive(&mut first); + first.send(Message::Text(r#"{"type":"error","status":429,"error":{"type":"rate_limit_error","code":"rate_limit_exceeded"}}"#.into())).unwrap(); + let mut fresh = socket(&listener); + receive(&mut fresh); + success(&mut fresh); + }); + let mut cfg = config(&endpoint, OpenAIResponsesTransport::WebSocket); + cfg.resilience.as_mut().unwrap().max_retries = 1; + let adapter = OpenAIResponsesAdapter::new(cfg).unwrap(); + let mut session = adapter + .start_session(SessionConfig::new("session")) + .await + .unwrap(); + let armed = Arc::new(std::sync::atomic::AtomicBool::new(false)); + let callback_armed = armed.clone(); + session.set_retry_observer(Some(Arc::new(move |_: ProviderRetryEvent| { + if callback_armed.load(std::sync::atomic::Ordering::SeqCst) { + panic!("observer panic"); + } + }))); + let mut turn = session.begin_turn(request(), None).await.unwrap(); + armed.store(true, std::sync::atomic::Ordering::SeqCst); + assert!( + std::panic::AssertUnwindSafe(turn.next_event(None)) + .catch_unwind() + .await + .is_err() + ); + drop(turn); + session.set_retry_observer(None); + let mut fresh = session.begin_turn(request(), None).await.unwrap(); + finished(&drain(&mut fresh).await.unwrap()); + peer.join().unwrap(); +} + +#[tokio::test] +async fn stale_completed_response_id_is_rejected_on_reused_socket() { + let (endpoint, peer) = server(|listener| { + let mut ws = socket(&listener); + receive(&mut ws); + success(&mut ws); + receive(&mut ws); + ws.send(Message::Text(r#"{"type":"response.created","sequence_number":1,"response":{"id":"resp-1","model":"gpt-test"}}"#.into())).unwrap(); + }); + let adapter = + OpenAIResponsesAdapter::new(config(&endpoint, OpenAIResponsesTransport::WebSocket)) + .unwrap(); + let mut session = adapter + .start_session(SessionConfig::new("session")) + .await + .unwrap(); + let mut first = session.begin_turn(request(), None).await.unwrap(); + finished(&drain(&mut first).await.unwrap()); + let mut second = session.begin_turn(request(), None).await.unwrap(); + assert!(drain(&mut second).await.is_err()); + peer.join().unwrap(); +} + +struct RefreshAuth; +#[async_trait] +impl AuthenticationProvider for RefreshAuth { + async fn authenticate( + &self, + previous: Option<&AuthenticationAttempt>, + ) -> Result { + let mut headers = HeaderMap::new(); + headers.insert( + "authorization", + HeaderValue::from_static(if previous.is_some() { + "Bearer refreshed" + } else { + "Bearer initial" + }), + ); + Ok(AuthenticationAttempt::stateless(headers).with_binding("loopback-identity")) + } +} + +#[tokio::test] +#[allow( + clippy::result_large_err, + reason = "tungstenite fixes the handshake callback error type" +)] +async fn handshake_and_wrapped_401_refresh_credentials_on_fresh_connection() { + for wrapped in [false, true] { + let (endpoint, peer) = server(move |listener| { + let original = if wrapped { + let mut first = socket(&listener); + let request = receive(&mut first); + first.send(Message::Text(r#"{"type":"error","status":401,"error":{"type":"authentication_error","code":"invalid_api_key"}}"#.into())).unwrap(); + Some(request) + } else { + let (headers, _) = http(&listener, "401 Unauthorized", ""); + assert!( + headers + .to_ascii_lowercase() + .contains("authorization: bearer initial") + ); + None + }; + let mut refreshed = tungstenite::accept_hdr( + accept(&listener), + |request: &tungstenite::handshake::server::Request, + response: tungstenite::handshake::server::Response| { + assert_eq!(request.headers()["authorization"], "Bearer refreshed"); + Ok(response) + }, + ) + .unwrap(); + let replay = receive(&mut refreshed); + if let Some(original) = original { + assert_eq!(original, replay); + } + success(&mut refreshed); + }); + let cfg = config(&endpoint, OpenAIResponsesTransport::WebSocket) + .with_authentication_provider(RefreshAuth); + let adapter = OpenAIResponsesAdapter::new(cfg).unwrap(); + let mut session = adapter + .start_session(SessionConfig::new("session")) + .await + .unwrap(); + let mut turn = session.begin_turn(request(), None).await.unwrap(); + finished(&drain(&mut turn).await.unwrap()); + peer.join().unwrap(); + } +} + +#[tokio::test] +async fn unsolicited_buffered_frames_force_reconnect_before_next_send() { + let (endpoint, peer) = server(|listener| { + let mut ws = socket(&listener); + receive(&mut ws); + for line in SUCCESS + .lines() + .filter_map(|line| line.strip_prefix("data: ")) + { + ws.write(Message::Text(line.into())).unwrap(); + } + ws.write(Message::Text( + r#"{"type":"error","status":429,"error":{"code":"rate_limit_exceeded"}}"#.into(), + )) + .unwrap(); + ws.flush().unwrap(); + let mut fresh = socket(&listener); + receive(&mut fresh); + success_id(&mut fresh, "fresh-response"); + }); + let adapter = + OpenAIResponsesAdapter::new(config(&endpoint, OpenAIResponsesTransport::WebSocket)) + .unwrap(); + let mut session = adapter + .start_session(SessionConfig::new("session")) + .await + .unwrap(); + let mut first = session.begin_turn(request(), None).await.unwrap(); + finished(&drain(&mut first).await.unwrap()); + let mut second = session.begin_turn(request(), None).await.unwrap(); + finished(&drain(&mut second).await.unwrap()); + peer.join().unwrap(); +} + +struct SuspendedRefresh { + started: Arc, + calls: Arc, +} + +#[async_trait] +impl AuthenticationProvider for SuspendedRefresh { + async fn authenticate( + &self, + previous: Option<&AuthenticationAttempt>, + ) -> Result { + if previous.is_some() { + self.calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + self.started + .store(true, std::sync::atomic::Ordering::SeqCst); + futures_util::future::pending().await + } else { + let mut headers = HeaderMap::new(); + headers.insert("authorization", HeaderValue::from_static("Bearer initial")); + Ok(AuthenticationAttempt::stateless(headers).with_binding("loopback-identity")) + } + } +} + +#[tokio::test] +async fn dropped_refresh_future_repolls_safely_without_replaying_stale_credentials() { + use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; + use std::task::Poll; + for handshake_refresh in [false, true] { + let (check_tx, check_rx) = std::sync::mpsc::channel(); + let (checked_tx, checked_rx) = std::sync::mpsc::channel(); + let (endpoint, peer) = server(move |listener| { + let mut first = socket(&listener); + receive(&mut first); + let status = if handshake_refresh { 429 } else { 401 }; + first + .send(Message::Text( + json!({"type":"error", "status": status, + "error":{"code":"rate_limit_exceeded"}}) + .to_string() + .into(), + )) + .unwrap(); + if handshake_refresh { + http(&listener, "401 Unauthorized", ""); + } + // The rejected socket is discarded before entering authentication. + assert!(first.read().is_err()); + check_rx.recv_timeout(WAIT).unwrap(); + assert_eq!( + listener.accept().unwrap_err().kind(), + std::io::ErrorKind::WouldBlock, + "repoll must not reconnect with stale credentials" + ); + checked_tx.send(()).unwrap(); + let mut fresh = socket(&listener); + receive(&mut fresh); + success(&mut fresh); + }); + let started = Arc::new(AtomicBool::new(false)); + let calls = Arc::new(AtomicUsize::new(0)); + let mut cfg = config(&endpoint, OpenAIResponsesTransport::WebSocket) + .with_authentication_provider(SuspendedRefresh { + started: started.clone(), + calls: calls.clone(), + }); + cfg.resilience.as_mut().unwrap().max_retries = 1; + let adapter = OpenAIResponsesAdapter::new(cfg).unwrap(); + let mut session = adapter + .start_session(SessionConfig::new("session")) + .await + .unwrap(); + let mut turn = session.begin_turn(request(), None).await.unwrap(); + let mut pending = Box::pin(turn.next_event(None)); + tokio::time::timeout( + WAIT, + futures_util::future::poll_fn(|cx| { + assert!(pending.as_mut().poll(cx).is_pending()); + if started.load(Ordering::SeqCst) { + Poll::Ready(()) + } else { + Poll::Pending + } + }), + ) + .await + .unwrap(); + drop(pending); + let error = turn.next_event(None).await.unwrap_err(); + assert_eq!( + error.provider_failure().unwrap().reason, + if handshake_refresh { + ProviderFailureReason::ReplayUnsafe + } else { + ProviderFailureReason::Authentication + } + ); + assert_eq!(calls.load(Ordering::SeqCst), 1); + assert!(turn.next_event(None).await.unwrap().is_none()); + check_tx.send(()).unwrap(); + checked_rx.recv_timeout(WAIT).unwrap(); + // Terminal failure released the claim even though the owned turn remains alive. + let mut fresh = session.begin_turn(request(), None).await.unwrap(); + finished(&drain(&mut fresh).await.unwrap()); + peer.join().unwrap(); + } +} From 5d238ca5a35973765ff4ab1d376f74a77f0bb03b Mon Sep 17 00:00:00 2001 From: daniel Date: Wed, 23 Sep 2026 12:07:16 -0700 Subject: [PATCH 2/2] fix(openai): prevent uncorrelated websocket rejection replays --- crates/agentkit-provider-openai/README.md | 5 +- .../agentkit-provider-openai/src/responses.rs | 12 ++-- .../src/responses/websocket.rs | 6 ++ .../src/responses/websocket/tests.rs | 63 +++++++++++++++++++ 4 files changed, 80 insertions(+), 6 deletions(-) diff --git a/crates/agentkit-provider-openai/README.md b/crates/agentkit-provider-openai/README.md index d1e5060..a34bb99 100644 --- a/crates/agentkit-provider-openai/README.md +++ b/crates/agentkit-provider-openai/README.md @@ -119,7 +119,10 @@ WebSocket retries are intentionally more conservative than HTTP retries: - Handshake status failures use the existing bounded retry policy and observer. - A wrapped HTTP error before `response.created` (and before visible output) - may reconnect and retry. An accepted response is never automatically replayed. + may reconnect and retry only on a fresh socket. On reused sockets, errors lack + reliable request correlation and may belong to a previous turn, so they never + trigger automatic replay or authentication refresh. An accepted response is + never automatically replayed. - An interrupted send, socket EOF, receive failure, or timeout after sending is **not replayed**: the server may already have accepted the request. - Visible WebSocket output is never superseded/replayed, even when the consumer diff --git a/crates/agentkit-provider-openai/src/responses.rs b/crates/agentkit-provider-openai/src/responses.rs index 16090b3..3a32f68 100644 --- a/crates/agentkit-provider-openai/src/responses.rs +++ b/crates/agentkit-provider-openai/src/responses.rs @@ -768,13 +768,15 @@ impl OpenAIResponsesTurn { .as_ref() .is_some_and(|attempt| matches!(attempt.body, LiveBody::WebSocket(_))); // A sent WebSocket request has no idempotency guarantee. Only explicit - // provider rejections before visible output can be retried. + // provider rejections before visible output on a fresh socket + // can be retried; reused sockets can deliver uncorrelated stale errors. if websocket && (self.attempt_output_emitted - || !self - .attempt - .as_ref() - .is_some_and(|attempt| attempt.decoder.request_rejected)) + || !self.attempt.as_ref().is_some_and(|attempt| { + attempt.decoder.request_rejected + && matches!(&attempt.body, LiveBody::WebSocket(connection) + if connection.can_retry_rejection()) + })) { return Err(*failure.error); } diff --git a/crates/agentkit-provider-openai/src/responses/websocket.rs b/crates/agentkit-provider-openai/src/responses/websocket.rs index 4238d3e..6dadf12 100644 --- a/crates/agentkit-provider-openai/src/responses/websocket.rs +++ b/crates/agentkit-provider-openai/src/responses/websocket.rs @@ -104,6 +104,12 @@ pub(super) struct Connection { } impl Connection { + pub(super) fn can_retry_rejection(&self) -> bool { + // Wrapped errors have no reliable request correlation. On a reused + // socket they may belong to a previous turn, even after a pending read. + self.completed_ids.is_empty() + } + pub(super) async fn recv(&mut self) -> Result, HttpError> { // Bound control traffic as well as data, including peers that only ping. for _ in 0..128 { diff --git a/crates/agentkit-provider-openai/src/responses/websocket/tests.rs b/crates/agentkit-provider-openai/src/responses/websocket/tests.rs index a51e597..73a9349 100644 --- a/crates/agentkit-provider-openai/src/responses/websocket/tests.rs +++ b/crates/agentkit-provider-openai/src/responses/websocket/tests.rs @@ -654,6 +654,69 @@ async fn unsolicited_buffered_frames_force_reconnect_before_next_send() { peer.join().unwrap(); } +#[tokio::test] +async fn partial_stale_error_on_reused_socket_never_replays_new_request() { + let (partial_tx, partial_rx) = tokio::sync::oneshot::channel(); + let (done_tx, done_rx) = std::sync::mpsc::channel(); + let (endpoint, peer) = server(move |listener| { + let mut ws = socket(&listener); + receive(&mut ws); + success(&mut ws); + let error = br#"{"type":"error","status":429,"error":{"code":"rate_limit_exceeded"}}"#; + assert!(error.len() < 126); + // An unmasked text frame, split before the next request is sent. + ws.get_mut().write_all(&[0x81, error.len() as u8]).unwrap(); + let split = error.len() / 2; + ws.get_mut().write_all(&error[..split]).unwrap(); + partial_tx.send(()).unwrap(); + receive(&mut ws); + ws.get_mut().write_all(&error[split..]).unwrap(); + let start = Instant::now(); + loop { + if done_rx.try_recv().is_ok() { + assert!( + matches!(listener.accept(), Err(e) if e.kind() == std::io::ErrorKind::WouldBlock) + ); + break; + } + match listener.accept() { + Ok((stream, _)) => { + stream.set_nonblocking(false).unwrap(); + stream.set_read_timeout(Some(WAIT)).unwrap(); + stream.set_write_timeout(Some(WAIT)).unwrap(); + let mut replay = tungstenite::accept(stream).unwrap(); + receive(&mut replay); + success_id(&mut replay, "unexpected-replay"); + panic!("replayed a request after an uncorrelated stale error"); + } + Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {} + Err(e) => panic!("accept: {e}"), + } + assert!(start.elapsed() < WAIT); + thread::sleep(Duration::from_millis(2)); + } + }); + let mut cfg = config(&endpoint, OpenAIResponsesTransport::WebSocket); + cfg.resilience.as_mut().unwrap().max_retries = 1; + let adapter = OpenAIResponsesAdapter::new(cfg).unwrap(); + let mut session = adapter + .start_session(SessionConfig::new("session")) + .await + .unwrap(); + let mut first = session.begin_turn(request(), None).await.unwrap(); + finished(&drain(&mut first).await.unwrap()); + partial_rx.await.unwrap(); + let mut second = session.begin_turn(request(), None).await.unwrap(); + let result = drain(&mut second).await; + done_tx.send(()).ok(); + peer.join().unwrap(); + let error = result.unwrap_err(); + assert_eq!( + error.provider_failure().unwrap().upstream.http_status, + Some(429) + ); +} + struct SuspendedRefresh { started: Arc, calls: Arc,