Skip to main content

rig_core/providers/openai/responses_api/
websocket.rs

1//! WebSocket session support for the OpenAI Responses API.
2//!
3//! This module implements OpenAI's `/v1/responses` WebSocket mode as a stateful,
4//! sequential session. Each connection supports a single in-flight response at a
5//! time, which matches OpenAI's current protocol constraints.
6
7use crate::completion::{self, CompletionError};
8use crate::http_client::HttpClientExt;
9use crate::providers::openai::responses_api::streaming::{
10    ItemChunk, ResponseChunk, ResponseChunkKind, StreamingCompletionChunk,
11};
12use crate::wasm_compat::{WasmCompatSend, WasmCompatSync};
13use futures::{SinkExt, StreamExt};
14use serde::{Deserialize, Serialize};
15use serde_json::{Map, Value};
16use std::time::Duration;
17use tokio::net::TcpStream;
18use tokio_tungstenite::{
19    MaybeTlsStream, WebSocketStream, connect_async,
20    tungstenite::{self, Message, client::IntoClientRequest},
21};
22use tracing::Level;
23use url::Url;
24
25use super::{CompletionResponse, ResponseStatus, ResponsesCompletionModel};
26
27type OpenAIWebSocket = WebSocketStream<MaybeTlsStream<TcpStream>>;
28const DEFAULT_CONNECT_TIMEOUT: Duration = Duration::from_secs(30);
29
30/// Options for a `response.create` message sent over OpenAI WebSocket mode.
31#[derive(Debug, Clone, Default, Serialize, Deserialize)]
32pub struct ResponsesWebSocketCreateOptions {
33    /// When set to `false`, OpenAI prepares request state without generating a model output.
34    ///
35    /// This is the "warmup" mode described in the OpenAI WebSocket mode guide.
36    #[serde(skip_serializing_if = "Option::is_none")]
37    pub generate: Option<bool>,
38}
39
40impl ResponsesWebSocketCreateOptions {
41    /// Creates warmup options equivalent to `generate: false`.
42    #[must_use]
43    pub fn warmup() -> Self {
44        Self {
45            generate: Some(false),
46        }
47    }
48}
49
50#[derive(Debug, Clone, Serialize)]
51struct ResponsesWebSocketClientEvent {
52    #[serde(rename = "type")]
53    kind: ResponsesWebSocketClientEventKind,
54    #[serde(flatten)]
55    request: super::CompletionRequest,
56    #[serde(skip_serializing_if = "Option::is_none")]
57    generate: Option<bool>,
58}
59
60#[derive(Debug, Clone, Serialize)]
61enum ResponsesWebSocketClientEventKind {
62    #[serde(rename = "response.create")]
63    ResponseCreate,
64}
65
66/// A protocol error event emitted by OpenAI WebSocket mode.
67#[derive(Debug, Clone, Serialize, Deserialize)]
68pub struct ResponsesWebSocketErrorEvent {
69    /// The event type.
70    #[serde(rename = "type")]
71    pub kind: ResponsesWebSocketErrorEventKind,
72    /// The provider error payload.
73    pub error: ResponsesWebSocketErrorPayload,
74}
75
76impl std::fmt::Display for ResponsesWebSocketErrorEvent {
77    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
78        self.error.fmt(f)
79    }
80}
81
82/// The event kind for an OpenAI WebSocket protocol error.
83#[derive(Debug, Clone, Serialize, Deserialize)]
84pub enum ResponsesWebSocketErrorEventKind {
85    #[serde(rename = "error")]
86    Error,
87}
88
89/// The payload carried by an OpenAI WebSocket protocol error event.
90#[derive(Debug, Clone, Default, Serialize, Deserialize)]
91pub struct ResponsesWebSocketErrorPayload {
92    /// Provider-specific error code when supplied.
93    #[serde(skip_serializing_if = "Option::is_none")]
94    pub code: Option<String>,
95    /// Human-readable error message.
96    #[serde(skip_serializing_if = "Option::is_none")]
97    pub message: Option<String>,
98    /// Any extra fields supplied by the provider.
99    #[serde(flatten, default)]
100    pub extra: Map<String, Value>,
101}
102
103impl std::fmt::Display for ResponsesWebSocketErrorPayload {
104    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
105        match (&self.code, &self.message) {
106            (Some(code), Some(message)) => write!(f, "{code}: {message}"),
107            (None, Some(message)) => f.write_str(message),
108            (Some(code), None) => f.write_str(code),
109            (None, None) => f.write_str("OpenAI websocket error"),
110        }
111    }
112}
113
114/// The optional `response.done` event emitted by OpenAI WebSocket mode.
115#[derive(Debug, Clone, Serialize, Deserialize)]
116pub struct ResponsesWebSocketDoneEvent {
117    /// The event type.
118    #[serde(rename = "type")]
119    pub kind: ResponsesWebSocketDoneEventKind,
120    /// The provider payload for the finished response.
121    pub response: Value,
122}
123
124impl ResponsesWebSocketDoneEvent {
125    /// Returns the response ID if the payload includes one.
126    #[must_use]
127    pub fn response_id(&self) -> Option<&str> {
128        self.response.get("id").and_then(Value::as_str)
129    }
130
131    fn status(&self) -> Option<ResponseStatus> {
132        self.response
133            .get("status")
134            .cloned()
135            .and_then(|status| serde_json::from_value(status).ok())
136    }
137
138    fn as_completion_response(&self) -> Option<CompletionResponse> {
139        serde_json::from_value(self.response.clone()).ok()
140    }
141}
142
143/// The event kind for the terminal websocket event.
144#[derive(Debug, Clone, Serialize, Deserialize)]
145pub enum ResponsesWebSocketDoneEventKind {
146    #[serde(rename = "response.done")]
147    ResponseDone,
148}
149
150/// A server event emitted by OpenAI WebSocket mode.
151#[derive(Debug, Clone)]
152pub enum ResponsesWebSocketEvent {
153    /// A response lifecycle event such as `response.created` or `response.completed`.
154    Response(Box<ResponseChunk>),
155    /// A streaming item/delta event such as `response.output_text.delta`.
156    Item(ItemChunk),
157    /// A protocol-level websocket error event.
158    Error(ResponsesWebSocketErrorEvent),
159    /// An optional `response.done` event emitted by OpenAI over WebSockets.
160    Done(ResponsesWebSocketDoneEvent),
161}
162
163impl ResponsesWebSocketEvent {
164    /// Returns the response ID when the event includes one.
165    #[must_use]
166    pub fn response_id(&self) -> Option<&str> {
167        match self {
168            Self::Response(chunk) => Some(&chunk.response.id),
169            Self::Done(done) => done.response_id(),
170            Self::Item(_) | Self::Error(_) => None,
171        }
172    }
173
174    /// Returns `true` when this event ends the current in-flight websocket turn.
175    #[must_use]
176    pub fn is_terminal(&self) -> bool {
177        match self {
178            Self::Response(chunk) => matches!(
179                chunk.kind,
180                ResponseChunkKind::ResponseCompleted
181                    | ResponseChunkKind::ResponseFailed
182                    | ResponseChunkKind::ResponseIncomplete
183            ),
184            Self::Error(_) | Self::Done(_) => true,
185            Self::Item(_) => false,
186        }
187    }
188}
189
190/// A builder for an OpenAI Responses WebSocket session.
191///
192/// The default builder applies a 30 second connection timeout and leaves the
193/// per-event timeout disabled.
194pub struct ResponsesWebSocketSessionBuilder<H = reqwest::Client> {
195    model: ResponsesCompletionModel<H>,
196    connect_timeout: Option<Duration>,
197    event_timeout: Option<Duration>,
198}
199
200impl<H> ResponsesWebSocketSessionBuilder<H> {
201    pub(crate) fn new(model: ResponsesCompletionModel<H>) -> Self {
202        Self {
203            model,
204            connect_timeout: Some(DEFAULT_CONNECT_TIMEOUT),
205            event_timeout: None,
206        }
207    }
208
209    /// Sets the timeout for establishing the websocket connection.
210    #[must_use]
211    pub fn connect_timeout(mut self, timeout: Duration) -> Self {
212        self.connect_timeout = Some(timeout);
213        self
214    }
215
216    /// Disables the websocket connection timeout.
217    #[must_use]
218    pub fn without_connect_timeout(mut self) -> Self {
219        self.connect_timeout = None;
220        self
221    }
222
223    /// Sets the timeout for waiting on the next websocket event.
224    #[must_use]
225    pub fn event_timeout(mut self, timeout: Duration) -> Self {
226        self.event_timeout = Some(timeout);
227        self
228    }
229
230    /// Disables the websocket event timeout.
231    #[must_use]
232    pub fn without_event_timeout(mut self) -> Self {
233        self.event_timeout = None;
234        self
235    }
236}
237
238impl<H> ResponsesWebSocketSessionBuilder<H>
239where
240    H: HttpClientExt
241        + Clone
242        + std::fmt::Debug
243        + Default
244        + WasmCompatSend
245        + WasmCompatSync
246        + 'static,
247{
248    /// Opens the websocket session using the configured builder options.
249    pub async fn connect(self) -> Result<ResponsesWebSocketSession<H>, CompletionError> {
250        ResponsesWebSocketSession::connect_with_timeouts(
251            self.model,
252            self.connect_timeout,
253            self.event_timeout,
254        )
255        .await
256    }
257}
258
259/// A stateful OpenAI Responses WebSocket session.
260///
261/// This session keeps track of the most recent successful `response.id` so later
262/// turns can automatically chain via `previous_response_id` unless the request
263/// explicitly sets a different one.
264///
265/// Call [`ResponsesWebSocketSession::close`] when you are finished with the
266/// session so the websocket can complete a close handshake cleanly.
267pub struct ResponsesWebSocketSession<H = reqwest::Client> {
268    model: ResponsesCompletionModel<H>,
269    previous_response_id: Option<String>,
270    pending_done_response_id: Option<String>,
271    socket: OpenAIWebSocket,
272    in_flight: bool,
273    event_timeout: Option<Duration>,
274    closed: bool,
275    failed: bool,
276}
277
278impl<H> ResponsesWebSocketSession<H>
279where
280    H: HttpClientExt
281        + Clone
282        + std::fmt::Debug
283        + Default
284        + WasmCompatSend
285        + WasmCompatSync
286        + 'static,
287{
288    async fn connect_with_timeouts(
289        model: ResponsesCompletionModel<H>,
290        connect_timeout: Option<Duration>,
291        event_timeout: Option<Duration>,
292    ) -> Result<Self, CompletionError> {
293        let url = websocket_url(model.client.base_url())?;
294        let request = websocket_request(&url, model.client.headers())?;
295        let socket = connect_websocket(request, connect_timeout).await?;
296
297        Ok(Self {
298            model,
299            previous_response_id: None,
300            pending_done_response_id: None,
301            socket,
302            in_flight: false,
303            event_timeout,
304            closed: false,
305            failed: false,
306        })
307    }
308
309    /// Returns the most recent successful `response.id` tracked by this session.
310    #[must_use]
311    pub fn previous_response_id(&self) -> Option<&str> {
312        self.previous_response_id.as_deref()
313    }
314
315    /// Clears the cached `previous_response_id` so the next turn starts a fresh chain.
316    pub fn clear_previous_response_id(&mut self) {
317        self.previous_response_id = None;
318    }
319
320    /// Sends a `response.create` event for a Rig completion request.
321    pub async fn send(
322        &mut self,
323        completion_request: crate::completion::CompletionRequest,
324    ) -> Result<(), CompletionError> {
325        self.send_with_options(
326            completion_request,
327            ResponsesWebSocketCreateOptions::default(),
328        )
329        .await
330    }
331
332    /// Sends a `response.create` event with explicit websocket-mode options.
333    pub async fn send_with_options(
334        &mut self,
335        completion_request: crate::completion::CompletionRequest,
336        options: ResponsesWebSocketCreateOptions,
337    ) -> Result<(), CompletionError> {
338        self.ensure_open()?;
339
340        if self.in_flight {
341            return Err(CompletionError::ProviderError(
342                "An OpenAI websocket response is already in flight on this session".to_string(),
343            ));
344        }
345
346        let payload = ResponsesWebSocketClientEvent {
347            kind: ResponsesWebSocketClientEventKind::ResponseCreate,
348            request: self.prepare_request(completion_request)?,
349            generate: options.generate,
350        };
351
352        if tracing::enabled!(Level::TRACE) {
353            tracing::trace!(
354                target: "rig::completions",
355                "OpenAI websocket request: {}",
356                serde_json::to_string_pretty(&payload)?
357            );
358        }
359
360        let payload = serde_json::to_string(&payload)?;
361
362        if let Err(error) = self.socket.send(Message::text(payload)).await {
363            return Err(self.fail_session(websocket_provider_error(error)));
364        }
365        self.in_flight = true;
366
367        Ok(())
368    }
369
370    /// Reads the next server event for the current in-flight turn.
371    pub async fn next_event(&mut self) -> Result<ResponsesWebSocketEvent, CompletionError> {
372        self.ensure_open()?;
373
374        if !self.in_flight {
375            return Err(CompletionError::ProviderError(
376                "No OpenAI websocket response is currently in flight on this session".to_string(),
377            ));
378        }
379
380        loop {
381            let message = match self.read_next_message().await {
382                Ok(message) => message,
383                Err(error) => return Err(error),
384            };
385
386            let Some(message) = message else {
387                self.mark_closed();
388                return Err(CompletionError::ProviderError(
389                    "The OpenAI websocket connection closed before the turn finished".to_string(),
390                ));
391            };
392
393            let message = match message {
394                Ok(message) => message,
395                Err(error) => return Err(self.fail_session(websocket_provider_error(error))),
396            };
397            let payload = match websocket_message_to_text(message) {
398                Ok(Some(payload)) => payload,
399                Ok(None) => continue,
400                Err(error) => return Err(self.fail_session(error)),
401            };
402            let event = match parse_server_event(&payload) {
403                Ok(Some(event)) => event,
404                Ok(None) => continue,
405                Err(error) => return Err(self.fail_session(error)),
406            };
407            if let ResponsesWebSocketEvent::Done(done) = &event {
408                // OpenAI may emit `response.done` after the turn has already ended at
409                // `response.completed`. Ignore that trailing event on the next turn.
410                if self.pending_done_response_id.as_deref() == done.response_id() {
411                    self.pending_done_response_id = None;
412                    continue;
413                }
414            }
415            self.update_state_for_event(&event);
416            return Ok(event);
417        }
418    }
419
420    /// Sends a warmup turn (`generate: false`) and returns the resulting response ID.
421    pub async fn warmup(
422        &mut self,
423        completion_request: crate::completion::CompletionRequest,
424    ) -> Result<String, CompletionError> {
425        self.send_with_options(
426            completion_request,
427            ResponsesWebSocketCreateOptions::warmup(),
428        )
429        .await?;
430        let response = self.wait_for_completed_response().await?;
431        Ok(response.id)
432    }
433
434    /// Sends a completion turn and collects the final OpenAI response.
435    pub async fn completion(
436        &mut self,
437        completion_request: crate::completion::CompletionRequest,
438    ) -> Result<completion::CompletionResponse<CompletionResponse>, CompletionError> {
439        self.send(completion_request).await?;
440        let response = self.wait_for_completed_response().await?;
441        response.try_into()
442    }
443
444    /// Closes the websocket connection.
445    ///
446    /// Call this when you are finished with the session so the websocket can
447    /// terminate with a clean close handshake.
448    pub async fn close(&mut self) -> Result<(), CompletionError> {
449        if self.closed {
450            return Ok(());
451        }
452
453        let result = self
454            .socket
455            .close(None)
456            .await
457            .map_err(websocket_provider_error);
458        self.mark_closed();
459        result
460    }
461
462    fn prepare_request(
463        &self,
464        completion_request: crate::completion::CompletionRequest,
465    ) -> Result<super::CompletionRequest, CompletionError> {
466        let mut request = self.model.create_completion_request(completion_request)?;
467
468        // WebSocket mode is always event-driven, so these HTTP/SSE-specific flags
469        // are ignored by the provider and only add noise to the payload.
470        request.stream = None;
471        request.additional_parameters.background = None;
472
473        if request.additional_parameters.previous_response_id.is_none() {
474            request.additional_parameters.previous_response_id = self.previous_response_id.clone();
475        }
476
477        Ok(request)
478    }
479
480    async fn wait_for_completed_response(&mut self) -> Result<CompletionResponse, CompletionError> {
481        loop {
482            match self.next_event().await? {
483                ResponsesWebSocketEvent::Response(chunk) => {
484                    if matches!(
485                        chunk.kind,
486                        ResponseChunkKind::ResponseCompleted
487                            | ResponseChunkKind::ResponseFailed
488                            | ResponseChunkKind::ResponseIncomplete
489                    ) {
490                        return terminal_response_result(chunk.response);
491                    }
492                }
493                ResponsesWebSocketEvent::Done(done) => {
494                    if let Some(response) = done.as_completion_response() {
495                        return terminal_response_result(response);
496                    }
497
498                    let message = if let Some(response_id) = done.response_id() {
499                        format!(
500                            "OpenAI websocket turn ended with response.done before a terminal response body was available (response_id={response_id})"
501                        )
502                    } else {
503                        "OpenAI websocket turn ended with response.done before a terminal response body was available"
504                            .to_string()
505                    };
506
507                    return Err(CompletionError::ProviderError(message));
508                }
509                ResponsesWebSocketEvent::Error(error) => {
510                    // Genuine provider error event: preserve the serialized payload
511                    // (code + message + any extra fields) so provider_response_json()
512                    // parses it, matching the response.failed path. No HTTP status on
513                    // the websocket stream, so status: None.
514                    return Err(provider_error_from_event(error));
515                }
516                ResponsesWebSocketEvent::Item(_) => {}
517            }
518        }
519    }
520
521    fn update_state_for_event(&mut self, event: &ResponsesWebSocketEvent) {
522        match event {
523            ResponsesWebSocketEvent::Response(chunk) => match chunk.kind {
524                ResponseChunkKind::ResponseCompleted => {
525                    let response_id = chunk.response.id.clone();
526                    self.previous_response_id = Some(response_id.clone());
527                    self.pending_done_response_id = Some(response_id);
528                    self.in_flight = false;
529                }
530                ResponseChunkKind::ResponseFailed | ResponseChunkKind::ResponseIncomplete => {
531                    self.pending_done_response_id = Some(chunk.response.id.clone());
532                    self.previous_response_id = None;
533                    self.in_flight = false;
534                }
535                ResponseChunkKind::ResponseCreated | ResponseChunkKind::ResponseInProgress => {}
536            },
537            ResponsesWebSocketEvent::Done(done) => {
538                match done.status() {
539                    Some(ResponseStatus::Completed) => {
540                        if let Some(response_id) = done.response_id() {
541                            self.previous_response_id = Some(response_id.to_string());
542                        }
543                    }
544                    Some(ResponseStatus::Failed)
545                    | Some(ResponseStatus::Incomplete)
546                    | Some(ResponseStatus::Cancelled) => {
547                        self.previous_response_id = None;
548                    }
549                    Some(ResponseStatus::InProgress | ResponseStatus::Queued) | None => {}
550                }
551                self.pending_done_response_id = None;
552                self.in_flight = false;
553            }
554            ResponsesWebSocketEvent::Error(_) => {
555                self.previous_response_id = None;
556                self.pending_done_response_id = None;
557                self.in_flight = false;
558            }
559            ResponsesWebSocketEvent::Item(_) => {}
560        }
561    }
562
563    fn abort_turn(&mut self) {
564        self.previous_response_id = None;
565        self.pending_done_response_id = None;
566        self.in_flight = false;
567    }
568
569    fn mark_closed(&mut self) {
570        self.abort_turn();
571        self.closed = true;
572        self.failed = false;
573    }
574
575    fn mark_failed(&mut self) {
576        self.abort_turn();
577        self.failed = true;
578    }
579
580    fn ensure_open(&self) -> Result<(), CompletionError> {
581        if self.closed || self.failed {
582            return Err(CompletionError::ProviderError(
583                "The OpenAI websocket session is closed".to_string(),
584            ));
585        }
586
587        Ok(())
588    }
589
590    fn fail_session(&mut self, error: CompletionError) -> CompletionError {
591        self.mark_failed();
592        error
593    }
594
595    async fn read_next_message(
596        &mut self,
597    ) -> Result<Option<Result<Message, tungstenite::Error>>, CompletionError> {
598        if let Some(timeout_duration) = self.event_timeout {
599            match tokio::time::timeout(timeout_duration, self.socket.next()).await {
600                Ok(message) => Ok(message),
601                Err(_) => Err(self.fail_session(event_timeout_error(timeout_duration))),
602            }
603        } else {
604            Ok(self.socket.next().await)
605        }
606    }
607}
608
609impl<H> Drop for ResponsesWebSocketSession<H> {
610    fn drop(&mut self) {
611        if !self.closed {
612            tracing::warn!(
613                target: "rig::completions",
614                in_flight = self.in_flight,
615                "Dropping an OpenAI websocket session without calling close(); the connection will end without a close handshake"
616            );
617        }
618    }
619}
620
621fn terminal_response_result(
622    response: CompletionResponse,
623) -> Result<CompletionResponse, CompletionError> {
624    match response.status {
625        ResponseStatus::Completed => Ok(response),
626        // Deliberate two-tier behaviour: when the provider supplies its own error
627        // object we preserve the full failed-response envelope through
628        // `from_provider_body` (status: None, no HTTP status on the websocket
629        // stream) so `provider_response_json()` parses it — consistent with the
630        // `error` event and the streaming paths. The body is re-serialized from
631        // the parsed response (not byte-identical to the wire bytes, which aren't
632        // retained past parsing) — semantically the provider's payload. When the
633        // object is absent we have nothing provider-authored to surface, so we
634        // emit a Rig-authored `ProviderError` diagnostic (provider_response_body()
635        // is None).
636        ResponseStatus::Failed => match response.error.as_ref() {
637            Some(error) => Err(CompletionError::from_provider_body(
638                serde_json::to_string(&response).unwrap_or_else(|_| error.message.clone()),
639            )),
640            None => Err(CompletionError::ProviderError(response_error_message(
641                "failed response",
642            ))),
643        },
644        ResponseStatus::Incomplete => {
645            let reason = response
646                .incomplete_details
647                .as_ref()
648                .map(|details| details.reason.as_str())
649                .unwrap_or("unknown reason");
650            Err(CompletionError::ProviderError(format!(
651                "OpenAI websocket response was incomplete: {reason}"
652            )))
653        }
654        other => Err(CompletionError::ProviderError(format!(
655            "OpenAI websocket response ended in state {other:?}"
656        ))),
657    }
658}
659
660fn response_error_message(fallback: &str) -> String {
661    format!("OpenAI websocket returned a {fallback}")
662}
663
664/// Maps a provider `error` event into a [`CompletionError`] that preserves the
665/// raw error payload as JSON (code + message + any extra provider fields) so the
666/// `provider_response_*` helpers can inspect it. The websocket stream carries no
667/// HTTP status, so `status` is `None`. The body is the event re-serialized from
668/// the parsed representation (not byte-identical to the original wire bytes,
669/// which are not retained past parsing) — semantically the provider's payload.
670fn provider_error_from_event(error: ResponsesWebSocketErrorEvent) -> CompletionError {
671    CompletionError::from_provider_body(
672        serde_json::to_string(&error).unwrap_or_else(|_| error.to_string()),
673    )
674}
675
676fn is_known_streaming_event(kind: &str) -> bool {
677    matches!(
678        kind,
679        "response.created"
680            | "response.in_progress"
681            | "response.completed"
682            | "response.failed"
683            | "response.incomplete"
684            | "response.output_item.added"
685            | "response.output_item.done"
686            | "response.content_part.added"
687            | "response.content_part.done"
688            | "response.output_text.delta"
689            | "response.output_text.done"
690            | "response.refusal.delta"
691            | "response.refusal.done"
692            | "response.function_call_arguments.delta"
693            | "response.function_call_arguments.done"
694            | "response.reasoning_summary_part.added"
695            | "response.reasoning_summary_part.done"
696            | "response.reasoning_summary_text.delta"
697            | "response.reasoning_summary_text.done"
698    )
699}
700
701fn parse_server_event(payload: &str) -> Result<Option<ResponsesWebSocketEvent>, CompletionError> {
702    #[derive(Deserialize)]
703    struct EventType {
704        #[serde(rename = "type")]
705        kind: String,
706    }
707
708    let event_type = serde_json::from_str::<EventType>(payload)?;
709    match event_type.kind.as_str() {
710        "error" => serde_json::from_str(payload)
711            .map(|e| Some(ResponsesWebSocketEvent::Error(e)))
712            .map_err(CompletionError::from),
713        "response.done" => serde_json::from_str(payload)
714            .map(|d| Some(ResponsesWebSocketEvent::Done(d)))
715            .map_err(CompletionError::from),
716        kind if is_known_streaming_event(kind) => match serde_json::from_str(payload)? {
717            StreamingCompletionChunk::Response(response) => {
718                Ok(Some(ResponsesWebSocketEvent::Response(response)))
719            }
720            StreamingCompletionChunk::Delta(item) => Ok(Some(ResponsesWebSocketEvent::Item(item))),
721        },
722        _ => {
723            tracing::debug!(
724                target: "rig::completions",
725                event_type = event_type.kind.as_str(),
726                "Skipping unrecognised OpenAI websocket event"
727            );
728            Ok(None)
729        }
730    }
731}
732
733fn websocket_message_to_text(message: Message) -> Result<Option<String>, CompletionError> {
734    match message {
735        Message::Text(text) => Ok(Some(text.to_string())),
736        Message::Binary(bytes) => String::from_utf8(bytes.to_vec())
737            .map(Some)
738            .map_err(|error| CompletionError::ResponseError(error.to_string())),
739        Message::Ping(_) | Message::Pong(_) | Message::Frame(_) => Ok(None),
740        Message::Close(frame) => {
741            let reason = frame
742                .map(|frame| frame.reason.to_string())
743                .filter(|reason| !reason.is_empty())
744                .unwrap_or_else(|| "without a close reason".to_string());
745            Err(CompletionError::ProviderError(format!(
746                "The OpenAI websocket connection closed {reason}"
747            )))
748        }
749    }
750}
751
752fn websocket_url(base_url: &str) -> Result<String, CompletionError> {
753    let mut url = Url::parse(base_url)?;
754    match url.scheme() {
755        "https" => {
756            url.set_scheme("wss").map_err(|_| {
757                CompletionError::ProviderError("Failed to convert https URL to wss".to_string())
758            })?;
759        }
760        "http" => {
761            url.set_scheme("ws").map_err(|_| {
762                CompletionError::ProviderError("Failed to convert http URL to ws".to_string())
763            })?;
764        }
765        scheme => {
766            return Err(CompletionError::ProviderError(format!(
767                "Unsupported base URL scheme for OpenAI websocket mode: {scheme}"
768            )));
769        }
770    }
771
772    let path = format!("{}/responses", url.path().trim_end_matches('/'));
773    url.set_path(&path);
774    Ok(url.to_string())
775}
776
777fn websocket_request(
778    url: &str,
779    headers: &http::HeaderMap,
780) -> Result<http::Request<()>, CompletionError> {
781    let mut request = url.into_client_request().map_err(|error| {
782        CompletionError::ProviderError(format!("Failed to build OpenAI websocket request: {error}"))
783    })?;
784
785    for (name, value) in headers {
786        request.headers_mut().insert(name, value.clone());
787    }
788
789    Ok(request)
790}
791
792async fn connect_websocket(
793    request: http::Request<()>,
794    connect_timeout: Option<Duration>,
795) -> Result<OpenAIWebSocket, CompletionError> {
796    if let Some(timeout_duration) = connect_timeout {
797        match tokio::time::timeout(timeout_duration, connect_async(request)).await {
798            Ok(result) => result
799                .map(|(socket, _)| socket)
800                .map_err(websocket_provider_error),
801            Err(_) => Err(connect_timeout_error(timeout_duration)),
802        }
803    } else {
804        connect_async(request)
805            .await
806            .map(|(socket, _)| socket)
807            .map_err(websocket_provider_error)
808    }
809}
810
811fn connect_timeout_error(timeout: Duration) -> CompletionError {
812    CompletionError::ProviderError(format!(
813        "Timed out connecting to the OpenAI websocket after {timeout:?}"
814    ))
815}
816
817fn event_timeout_error(timeout: Duration) -> CompletionError {
818    CompletionError::ProviderError(format!(
819        "Timed out waiting for the next OpenAI websocket event after {timeout:?}"
820    ))
821}
822
823fn websocket_provider_error(error: tungstenite::Error) -> CompletionError {
824    CompletionError::ProviderError(error.to_string())
825}
826
827#[cfg(test)]
828mod tests {
829    use super::{
830        ResponsesWebSocketCreateOptions, ResponsesWebSocketDoneEvent, ResponsesWebSocketEvent,
831        parse_server_event, terminal_response_result, websocket_url,
832    };
833    use crate::client::CompletionClient;
834    use crate::completion::CompletionModel;
835    use crate::providers::openai::responses_api::{
836        CompletionResponse, ResponseError, ResponseObject, ResponseStatus, ResponsesUsage,
837    };
838    use futures::{SinkExt, StreamExt};
839    use serde_json::json;
840    use std::time::Duration;
841    use tokio::net::TcpListener;
842    use tokio::time::sleep;
843    use tokio_tungstenite::{accept_async, tungstenite::Message};
844
845    #[test]
846    fn websocket_error_event_preserves_provider_payload_as_json() {
847        let mut extra = serde_json::Map::new();
848        extra.insert(
849            "type".to_string(),
850            serde_json::Value::String("invalid_request_error".to_string()),
851        );
852        let event = super::ResponsesWebSocketErrorEvent {
853            kind: super::ResponsesWebSocketErrorEventKind::Error,
854            error: super::ResponsesWebSocketErrorPayload {
855                code: Some("rate_limit_exceeded".to_string()),
856                message: Some("slow down".to_string()),
857                extra,
858            },
859        };
860
861        let err = super::provider_error_from_event(event);
862
863        // No HTTP status on the websocket stream, and the raw payload round-trips
864        // through provider_response_json() (code + message + extra all preserved).
865        assert_eq!(err.provider_response_status(), None);
866        let json = err
867            .provider_response_json()
868            .expect("preserved body should be valid JSON")
869            .expect("provider response body should be present");
870        assert_eq!(json["error"]["code"], "rate_limit_exceeded");
871        assert_eq!(json["error"]["message"], "slow down");
872        assert_eq!(json["error"]["type"], "invalid_request_error");
873    }
874
875    fn sample_response(status: ResponseStatus) -> CompletionResponse {
876        CompletionResponse {
877            id: "resp_123".to_string(),
878            object: ResponseObject::Response,
879            created_at: 0,
880            status,
881            error: None,
882            incomplete_details: None,
883            instructions: None,
884            max_output_tokens: None,
885            model: "gpt-5.4".to_string(),
886            usage: Some(ResponsesUsage {
887                input_tokens: 1,
888                input_tokens_details: None,
889                output_tokens: 2,
890                output_tokens_details: Some(
891                    crate::providers::openai::responses_api::OutputTokensDetails {
892                        reasoning_tokens: 0,
893                    },
894                ),
895                total_tokens: 3,
896            }),
897            output: Vec::new(),
898            tools: Vec::new(),
899            additional_parameters: Default::default(),
900            provider_reasoning: None,
901            reasoning_metadata: None,
902            reasoning_context: None,
903        }
904    }
905
906    #[test]
907    fn warmup_options_serialize_generate_false() {
908        let options = ResponsesWebSocketCreateOptions::warmup();
909        let json = serde_json::to_value(options).expect("options should serialize");
910
911        assert_eq!(json, json!({ "generate": false }));
912    }
913
914    #[test]
915    fn websocket_url_converts_https_to_wss() {
916        let url = websocket_url("https://api.openai.com/v1").expect("url should convert");
917        assert_eq!(url, "wss://api.openai.com/v1/responses");
918    }
919
920    #[test]
921    fn parse_done_event_exposes_response_id() {
922        let payload = json!({
923            "type": "response.done",
924            "response": {
925                "id": "resp_done_1",
926                "status": "completed"
927            }
928        });
929
930        let event = parse_server_event(&payload.to_string())
931            .expect("done event should deserialize")
932            .expect("done event should not be skipped");
933
934        assert!(matches!(
935            event,
936            ResponsesWebSocketEvent::Done(ResponsesWebSocketDoneEvent { .. })
937        ));
938        assert_eq!(event.response_id(), Some("resp_done_1"));
939        assert!(event.is_terminal());
940    }
941
942    #[test]
943    fn parse_response_completed_event_is_terminal() {
944        let payload = json!({
945            "type": "response.completed",
946            "sequence_number": 12,
947            "response": {
948                "id": "resp_completed_1",
949                "object": "response",
950                "created_at": 0,
951                "status": "completed",
952                "error": null,
953                "incomplete_details": null,
954                "instructions": null,
955                "max_output_tokens": null,
956                "model": "gpt-5.4",
957                "usage": null,
958                "output": [],
959                "tools": []
960            }
961        });
962
963        let event = parse_server_event(&payload.to_string())
964            .expect("response event should deserialize")
965            .expect("response event should not be skipped");
966
967        assert!(matches!(event, ResponsesWebSocketEvent::Response(_)));
968        assert!(event.is_terminal());
969        assert_eq!(event.response_id(), Some("resp_completed_1"));
970    }
971
972    #[test]
973    fn parse_live_output_item_added_event() {
974        let payload = json!({
975            "type": "response.output_item.added",
976            "item": {
977                "id": "msg_036471c3a72c147b0069ae7848d68881959773fd2d99e3d98a",
978                "type": "message",
979                "status": "in_progress",
980                "content": [],
981                "role": "assistant"
982            },
983            "output_index": 0,
984            "sequence_number": 2
985        });
986
987        let event = parse_server_event(&payload.to_string())
988            .expect("output item event should parse")
989            .expect("output item event should not be skipped");
990
991        assert!(matches!(event, ResponsesWebSocketEvent::Item(_)));
992    }
993
994    #[test]
995    fn parse_live_content_part_added_event() {
996        let payload = json!({
997            "type": "response.content_part.added",
998            "content_index": 0,
999            "item_id": "msg_036471c3a72c147b0069ae7848d68881959773fd2d99e3d98a",
1000            "output_index": 0,
1001            "part": {
1002                "type": "output_text",
1003                "annotations": [],
1004                "logprobs": [],
1005                "text": ""
1006            },
1007            "sequence_number": 3
1008        });
1009
1010        let event = parse_server_event(&payload.to_string())
1011            .expect("content part event should parse")
1012            .expect("content part event should not be skipped");
1013
1014        assert!(matches!(event, ResponsesWebSocketEvent::Item(_)));
1015    }
1016
1017    #[test]
1018    fn parse_live_output_text_delta_event() {
1019        let payload = json!({
1020            "type": "response.output_text.delta",
1021            "content_index": 0,
1022            "delta": "Web",
1023            "item_id": "msg_023af0f0a91bc2a90069ae788612e881958345bb156915ba29",
1024            "logprobs": [],
1025            "obfuscation": "2YYErYq7jkqqM",
1026            "output_index": 0,
1027            "sequence_number": 4
1028        });
1029
1030        let event = parse_server_event(&payload.to_string())
1031            .expect("output text delta event should parse")
1032            .expect("output text delta event should not be skipped");
1033
1034        assert!(matches!(event, ResponsesWebSocketEvent::Item(_)));
1035    }
1036
1037    #[test]
1038    fn terminal_response_requires_completed_status() {
1039        let completed = terminal_response_result(sample_response(ResponseStatus::Completed))
1040            .expect("completed response should succeed");
1041        assert_eq!(completed.id, "resp_123");
1042
1043        let failed = terminal_response_result(sample_response(ResponseStatus::Failed))
1044            .expect_err("failed response should error");
1045        assert!(failed.to_string().contains("failed response"));
1046    }
1047
1048    #[test]
1049    fn terminal_failed_response_with_error_preserves_raw_payload() {
1050        let mut response = sample_response(ResponseStatus::Failed);
1051        response.error = Some(ResponseError {
1052            code: "server_error".to_string(),
1053            message: "the model failed to generate a response".to_string(),
1054        });
1055
1056        let err = match terminal_response_result(response) {
1057            Ok(_) => panic!("failed response with an error object should fail"),
1058            Err(e) => e,
1059        };
1060
1061        // The full failed-response envelope is preserved as a ProviderResponse with
1062        // no HTTP status (the websocket stream carries none), so the raw JSON parses
1063        // back with the provider error nested under `error` — proving the whole
1064        // envelope is kept, not just the error object.
1065        assert_eq!(err.provider_response_status(), None);
1066
1067        let json = err
1068            .provider_response_json()
1069            .expect("preserved body should parse as JSON")
1070            .expect("preserved body should not be empty");
1071        assert_eq!(
1072            json["error"]["message"],
1073            "the model failed to generate a response"
1074        );
1075        assert_eq!(json["error"]["code"], "server_error");
1076    }
1077
1078    #[test]
1079    fn terminal_failed_response_without_error_is_rig_diagnostic() {
1080        let err = match terminal_response_result(sample_response(ResponseStatus::Failed)) {
1081            Ok(_) => panic!("failed response should fail"),
1082            Err(e) => e,
1083        };
1084
1085        // No provider error object, so this is a Rig-authored diagnostic and exposes
1086        // no preserved provider response body.
1087        assert_eq!(err.provider_response_body(), None);
1088        assert!(err.to_string().contains("failed response"));
1089    }
1090
1091    #[tokio::test]
1092    async fn malformed_known_event_rejects_reuse_and_allows_close() {
1093        let listener = TcpListener::bind("127.0.0.1:0")
1094            .await
1095            .expect("listener should bind");
1096        let address = listener.local_addr().expect("listener should have address");
1097
1098        let server = tokio::spawn(async move {
1099            let (stream, _) = listener.accept().await.expect("server should accept");
1100            let mut socket = accept_async(stream)
1101                .await
1102                .expect("server should upgrade websocket");
1103
1104            let request = socket
1105                .next()
1106                .await
1107                .expect("request should exist")
1108                .expect("request should be valid");
1109            let payload = request.into_text().expect("request should be text");
1110            assert!(
1111                payload.contains("\"type\":\"response.create\""),
1112                "expected response.create payload, got {payload}"
1113            );
1114
1115            socket
1116                .send(Message::text(
1117                    json!({
1118                        "type": "response.completed"
1119                    })
1120                    .to_string(),
1121                ))
1122                .await
1123                .expect("malformed known event should send");
1124
1125            let message = socket
1126                .next()
1127                .await
1128                .expect("close frame should arrive")
1129                .expect("close frame should be valid");
1130            assert!(
1131                matches!(message, Message::Close(_)),
1132                "expected close frame, got {message:?}"
1133            );
1134        });
1135
1136        let base_url = format!("http://{address}/v1");
1137        let client = crate::providers::openai::Client::builder()
1138            .api_key("test-key")
1139            .base_url(&base_url)
1140            .build()
1141            .expect("client should build");
1142        let model = client.completion_model("gpt-4o");
1143        let mut session = client
1144            .responses_websocket("gpt-4o")
1145            .await
1146            .expect("session should connect");
1147
1148        session
1149            .send(model.completion_request("hello").build())
1150            .await
1151            .expect("request should send");
1152
1153        let error = session
1154            .next_event()
1155            .await
1156            .expect_err("malformed known event should fail");
1157        assert!(
1158            error.to_string().contains("StreamingCompletionChunk"),
1159            "expected strict decode failure, got {error}"
1160        );
1161
1162        let closed = session
1163            .send(model.completion_request("retry").build())
1164            .await
1165            .expect_err("session should close after fatal parse error");
1166        assert!(
1167            closed.to_string().contains("session is closed"),
1168            "expected closed-session error, got {closed}"
1169        );
1170
1171        session
1172            .close()
1173            .await
1174            .expect("explicit close after fatal parse error should succeed");
1175
1176        server.await.expect("server task should finish");
1177    }
1178
1179    #[tokio::test]
1180    async fn event_timeout_rejects_reuse_and_allows_close() {
1181        let listener = TcpListener::bind("127.0.0.1:0")
1182            .await
1183            .expect("listener should bind");
1184        let address = listener.local_addr().expect("listener should have address");
1185
1186        let server = tokio::spawn(async move {
1187            let (stream, _) = listener.accept().await.expect("server should accept");
1188            let mut socket = accept_async(stream)
1189                .await
1190                .expect("server should upgrade websocket");
1191
1192            let request = socket
1193                .next()
1194                .await
1195                .expect("request should exist")
1196                .expect("request should be valid");
1197            let payload = request.into_text().expect("request should be text");
1198            assert!(
1199                payload.contains("\"type\":\"response.create\""),
1200                "expected response.create payload, got {payload}"
1201            );
1202
1203            sleep(Duration::from_millis(60)).await;
1204            let message = socket
1205                .next()
1206                .await
1207                .expect("close frame should arrive")
1208                .expect("close frame should be valid");
1209            assert!(
1210                matches!(message, Message::Close(_)),
1211                "expected close frame, got {message:?}"
1212            );
1213        });
1214
1215        let base_url = format!("http://{address}/v1");
1216        let client = crate::providers::openai::Client::builder()
1217            .api_key("test-key")
1218            .base_url(&base_url)
1219            .build()
1220            .expect("client should build");
1221        let model = client.completion_model("gpt-4o");
1222        let mut session = client
1223            .responses_websocket_builder("gpt-4o")
1224            .event_timeout(Duration::from_millis(20))
1225            .connect()
1226            .await
1227            .expect("session should connect");
1228
1229        session
1230            .send(model.completion_request("hello").build())
1231            .await
1232            .expect("request should send");
1233
1234        let error = session
1235            .next_event()
1236            .await
1237            .expect_err("next_event should time out");
1238        assert!(
1239            error
1240                .to_string()
1241                .contains("Timed out waiting for the next OpenAI websocket event"),
1242            "expected timeout error, got {error}"
1243        );
1244
1245        let closed = session
1246            .send(model.completion_request("retry").build())
1247            .await
1248            .expect_err("timed-out session should close");
1249        assert!(
1250            closed.to_string().contains("session is closed"),
1251            "expected closed-session error, got {closed}"
1252        );
1253
1254        session
1255            .close()
1256            .await
1257            .expect("explicit close after timeout should succeed");
1258
1259        server.await.expect("server task should finish");
1260    }
1261
1262    #[tokio::test]
1263    async fn late_response_done_is_ignored_on_next_turn() {
1264        let listener = TcpListener::bind("127.0.0.1:0")
1265            .await
1266            .expect("listener should bind");
1267        let address = listener.local_addr().expect("listener should have address");
1268
1269        let server = tokio::spawn(async move {
1270            let (stream, _) = listener.accept().await.expect("server should accept");
1271            let mut socket = accept_async(stream)
1272                .await
1273                .expect("server should upgrade websocket");
1274
1275            for (index, response_id) in ["resp_1", "resp_2"].iter().enumerate() {
1276                let request = socket
1277                    .next()
1278                    .await
1279                    .expect("request should exist")
1280                    .expect("request should be valid");
1281                let payload = request.into_text().expect("request should be text");
1282                assert!(
1283                    payload.contains("\"type\":\"response.create\""),
1284                    "expected response.create payload, got {payload}"
1285                );
1286
1287                let response = sample_response(ResponseStatus::Completed);
1288                let response = serde_json::to_value(CompletionResponse {
1289                    id: (*response_id).to_string(),
1290                    ..response
1291                })
1292                .expect("response should serialize");
1293
1294                socket
1295                    .send(Message::text(
1296                        json!({
1297                            "type": "response.completed",
1298                            "sequence_number": (index * 2) + 1,
1299                            "response": response,
1300                        })
1301                        .to_string(),
1302                    ))
1303                    .await
1304                    .expect("completed event should send");
1305                socket
1306                    .send(Message::text(
1307                        json!({
1308                            "type": "response.done",
1309                            "response": {
1310                                "id": response_id,
1311                                "status": "completed",
1312                            },
1313                        })
1314                        .to_string(),
1315                    ))
1316                    .await
1317                    .expect("done event should send");
1318            }
1319        });
1320
1321        let base_url = format!("http://{address}/v1");
1322        let client = crate::providers::openai::Client::builder()
1323            .api_key("test-key")
1324            .base_url(&base_url)
1325            .build()
1326            .expect("client should build");
1327        let model = client.completion_model("gpt-4o");
1328        let mut session = client
1329            .responses_websocket("gpt-4o")
1330            .await
1331            .expect("session should connect");
1332
1333        session
1334            .send(model.completion_request("first").build())
1335            .await
1336            .expect("first request should send");
1337        let first = session
1338            .wait_for_completed_response()
1339            .await
1340            .expect("first response should complete");
1341        assert_eq!(first.id, "resp_1");
1342        assert_eq!(session.previous_response_id(), Some("resp_1"));
1343
1344        session
1345            .send(model.completion_request("second").build())
1346            .await
1347            .expect("second request should send");
1348        let second = session
1349            .wait_for_completed_response()
1350            .await
1351            .expect("second response should complete");
1352        assert_eq!(second.id, "resp_2");
1353        assert_eq!(session.previous_response_id(), Some("resp_2"));
1354
1355        server.await.expect("server task should finish");
1356    }
1357
1358    #[tokio::test]
1359    async fn clearing_previous_response_id_does_not_disable_late_done_filter() {
1360        let listener = TcpListener::bind("127.0.0.1:0")
1361            .await
1362            .expect("listener should bind");
1363        let address = listener.local_addr().expect("listener should have address");
1364
1365        let server = tokio::spawn(async move {
1366            let (stream, _) = listener.accept().await.expect("server should accept");
1367            let mut socket = accept_async(stream)
1368                .await
1369                .expect("server should upgrade websocket");
1370
1371            for response_id in ["resp_1", "resp_2"] {
1372                let request = socket
1373                    .next()
1374                    .await
1375                    .expect("request should exist")
1376                    .expect("request should be valid");
1377                let payload = request.into_text().expect("request should be text");
1378                assert!(
1379                    payload.contains("\"type\":\"response.create\""),
1380                    "expected response.create payload, got {payload}"
1381                );
1382
1383                let response = sample_response(ResponseStatus::Completed);
1384                let response = serde_json::to_value(CompletionResponse {
1385                    id: response_id.to_string(),
1386                    ..response
1387                })
1388                .expect("response should serialize");
1389
1390                socket
1391                    .send(Message::text(
1392                        json!({
1393                            "type": "response.completed",
1394                            "sequence_number": 1,
1395                            "response": response,
1396                        })
1397                        .to_string(),
1398                    ))
1399                    .await
1400                    .expect("completed event should send");
1401                socket
1402                    .send(Message::text(
1403                        json!({
1404                            "type": "response.done",
1405                            "response": {
1406                                "id": response_id,
1407                                "status": "completed",
1408                            },
1409                        })
1410                        .to_string(),
1411                    ))
1412                    .await
1413                    .expect("done event should send");
1414            }
1415        });
1416
1417        let base_url = format!("http://{address}/v1");
1418        let client = crate::providers::openai::Client::builder()
1419            .api_key("test-key")
1420            .base_url(&base_url)
1421            .build()
1422            .expect("client should build");
1423        let model = client.completion_model("gpt-4o");
1424        let mut session = client
1425            .responses_websocket("gpt-4o")
1426            .await
1427            .expect("session should connect");
1428
1429        session
1430            .send(model.completion_request("first").build())
1431            .await
1432            .expect("first request should send");
1433        let first = session
1434            .wait_for_completed_response()
1435            .await
1436            .expect("first response should complete");
1437        assert_eq!(first.id, "resp_1");
1438
1439        session.clear_previous_response_id();
1440        assert_eq!(session.previous_response_id(), None);
1441
1442        session
1443            .send(model.completion_request("second").build())
1444            .await
1445            .expect("second request should send");
1446        let second = session
1447            .wait_for_completed_response()
1448            .await
1449            .expect("second response should complete");
1450        assert_eq!(second.id, "resp_2");
1451
1452        server.await.expect("server task should finish");
1453    }
1454
1455    #[tokio::test]
1456    async fn failed_turn_keeps_late_done_out_of_next_request() {
1457        let listener = TcpListener::bind("127.0.0.1:0")
1458            .await
1459            .expect("listener should bind");
1460        let address = listener.local_addr().expect("listener should have address");
1461
1462        let server = tokio::spawn(async move {
1463            let (stream, _) = listener.accept().await.expect("server should accept");
1464            let mut socket = accept_async(stream)
1465                .await
1466                .expect("server should upgrade websocket");
1467
1468            let first_request = socket
1469                .next()
1470                .await
1471                .expect("request should exist")
1472                .expect("request should be valid");
1473            let payload = first_request
1474                .into_text()
1475                .expect("failed request should be text");
1476            assert!(
1477                payload.contains("\"type\":\"response.create\""),
1478                "expected response.create payload, got {payload}"
1479            );
1480
1481            let failed_response = serde_json::to_value(CompletionResponse {
1482                id: "resp_failed".to_string(),
1483                status: ResponseStatus::Failed,
1484                ..sample_response(ResponseStatus::Completed)
1485            })
1486            .expect("failed response should serialize");
1487
1488            socket
1489                .send(Message::text(
1490                    json!({
1491                        "type": "response.failed",
1492                        "sequence_number": 1,
1493                        "response": failed_response,
1494                    })
1495                    .to_string(),
1496                ))
1497                .await
1498                .expect("failed event should send");
1499            socket
1500                .send(Message::text(
1501                    json!({
1502                        "type": "response.done",
1503                        "response": {
1504                            "id": "resp_failed",
1505                            "status": "failed",
1506                        },
1507                    })
1508                    .to_string(),
1509                ))
1510                .await
1511                .expect("done event should send");
1512
1513            let second_request = socket
1514                .next()
1515                .await
1516                .expect("request should exist")
1517                .expect("request should be valid");
1518            let payload = second_request
1519                .into_text()
1520                .expect("second request should be text");
1521            assert!(
1522                payload.contains("\"type\":\"response.create\""),
1523                "expected response.create payload, got {payload}"
1524            );
1525
1526            let response = sample_response(ResponseStatus::Completed);
1527            let response = serde_json::to_value(CompletionResponse {
1528                id: "resp_2".to_string(),
1529                ..response
1530            })
1531            .expect("response should serialize");
1532
1533            socket
1534                .send(Message::text(
1535                    json!({
1536                        "type": "response.completed",
1537                        "sequence_number": 2,
1538                        "response": response,
1539                    })
1540                    .to_string(),
1541                ))
1542                .await
1543                .expect("completed event should send");
1544            socket
1545                .send(Message::text(
1546                    json!({
1547                        "type": "response.done",
1548                        "response": {
1549                            "id": "resp_2",
1550                            "status": "completed",
1551                        },
1552                    })
1553                    .to_string(),
1554                ))
1555                .await
1556                .expect("done event should send");
1557        });
1558
1559        let base_url = format!("http://{address}/v1");
1560        let client = crate::providers::openai::Client::builder()
1561            .api_key("test-key")
1562            .base_url(&base_url)
1563            .build()
1564            .expect("client should build");
1565        let model = client.completion_model("gpt-4o");
1566        let mut session = client
1567            .responses_websocket("gpt-4o")
1568            .await
1569            .expect("session should connect");
1570
1571        session
1572            .send(model.completion_request("first").build())
1573            .await
1574            .expect("first request should send");
1575        let error = session
1576            .wait_for_completed_response()
1577            .await
1578            .expect_err("failed response should error");
1579        assert!(error.to_string().contains("failed response"));
1580        assert_eq!(session.previous_response_id(), None);
1581
1582        session
1583            .send(model.completion_request("second").build())
1584            .await
1585            .expect("second request should send");
1586        let second = session
1587            .wait_for_completed_response()
1588            .await
1589            .expect("second response should complete");
1590        assert_eq!(second.id, "resp_2");
1591
1592        server.await.expect("server task should finish");
1593    }
1594
1595    #[tokio::test]
1596    async fn done_first_completed_turn_updates_previous_response_id() {
1597        let listener = TcpListener::bind("127.0.0.1:0")
1598            .await
1599            .expect("listener should bind");
1600        let address = listener.local_addr().expect("listener should have address");
1601
1602        let server = tokio::spawn(async move {
1603            let (stream, _) = listener.accept().await.expect("server should accept");
1604            let mut socket = accept_async(stream)
1605                .await
1606                .expect("server should upgrade websocket");
1607
1608            for response_id in ["resp_1", "resp_2"] {
1609                let request = socket
1610                    .next()
1611                    .await
1612                    .expect("request should exist")
1613                    .expect("request should be valid");
1614                let payload = request.into_text().expect("request should be text");
1615                assert!(
1616                    payload.contains("\"type\":\"response.create\""),
1617                    "expected response.create payload, got {payload}"
1618                );
1619
1620                if response_id == "resp_2" {
1621                    assert!(
1622                        payload.contains("\"previous_response_id\":\"resp_1\""),
1623                        "expected chained previous_response_id in payload, got {payload}"
1624                    );
1625                }
1626
1627                let response = serde_json::to_value(CompletionResponse {
1628                    id: response_id.to_string(),
1629                    ..sample_response(ResponseStatus::Completed)
1630                })
1631                .expect("response should serialize");
1632
1633                socket
1634                    .send(Message::text(
1635                        json!({
1636                            "type": "response.done",
1637                            "response": response,
1638                        })
1639                        .to_string(),
1640                    ))
1641                    .await
1642                    .expect("done event should send");
1643            }
1644        });
1645
1646        let base_url = format!("http://{address}/v1");
1647        let client = crate::providers::openai::Client::builder()
1648            .api_key("test-key")
1649            .base_url(&base_url)
1650            .build()
1651            .expect("client should build");
1652        let model = client.completion_model("gpt-4o");
1653        let mut session = client
1654            .responses_websocket("gpt-4o")
1655            .await
1656            .expect("session should connect");
1657
1658        session
1659            .send(model.completion_request("first").build())
1660            .await
1661            .expect("first request should send");
1662        let first = session
1663            .wait_for_completed_response()
1664            .await
1665            .expect("first response should complete");
1666        assert_eq!(first.id, "resp_1");
1667        assert_eq!(session.previous_response_id(), Some("resp_1"));
1668
1669        session
1670            .send(model.completion_request("second").build())
1671            .await
1672            .expect("second request should send");
1673        let second = session
1674            .wait_for_completed_response()
1675            .await
1676            .expect("second response should complete");
1677        assert_eq!(second.id, "resp_2");
1678        assert_eq!(session.previous_response_id(), Some("resp_2"));
1679
1680        server.await.expect("server task should finish");
1681    }
1682
1683    #[tokio::test]
1684    async fn done_first_failed_turn_does_not_chain_next_request() {
1685        let listener = TcpListener::bind("127.0.0.1:0")
1686            .await
1687            .expect("listener should bind");
1688        let address = listener.local_addr().expect("listener should have address");
1689
1690        let server = tokio::spawn(async move {
1691            let (stream, _) = listener.accept().await.expect("server should accept");
1692            let mut socket = accept_async(stream)
1693                .await
1694                .expect("server should upgrade websocket");
1695
1696            let first_request = socket
1697                .next()
1698                .await
1699                .expect("request should exist")
1700                .expect("request should be valid");
1701            let payload = first_request
1702                .into_text()
1703                .expect("first request should be text");
1704            assert!(
1705                payload.contains("\"type\":\"response.create\""),
1706                "expected response.create payload, got {payload}"
1707            );
1708            assert!(
1709                !payload.contains("\"previous_response_id\""),
1710                "did not expect previous_response_id in first payload, got {payload}"
1711            );
1712
1713            let failed_response = serde_json::to_value(CompletionResponse {
1714                id: "resp_failed".to_string(),
1715                status: ResponseStatus::Failed,
1716                ..sample_response(ResponseStatus::Completed)
1717            })
1718            .expect("failed response should serialize");
1719
1720            socket
1721                .send(Message::text(
1722                    json!({
1723                        "type": "response.done",
1724                        "response": failed_response,
1725                    })
1726                    .to_string(),
1727                ))
1728                .await
1729                .expect("done event should send");
1730
1731            let second_request = socket
1732                .next()
1733                .await
1734                .expect("request should exist")
1735                .expect("request should be valid");
1736            let payload = second_request
1737                .into_text()
1738                .expect("second request should be text");
1739            assert!(
1740                payload.contains("\"type\":\"response.create\""),
1741                "expected response.create payload, got {payload}"
1742            );
1743            assert!(
1744                !payload.contains("\"previous_response_id\""),
1745                "did not expect chained previous_response_id in payload, got {payload}"
1746            );
1747
1748            let response = serde_json::to_value(CompletionResponse {
1749                id: "resp_2".to_string(),
1750                ..sample_response(ResponseStatus::Completed)
1751            })
1752            .expect("response should serialize");
1753
1754            socket
1755                .send(Message::text(
1756                    json!({
1757                        "type": "response.done",
1758                        "response": response,
1759                    })
1760                    .to_string(),
1761                ))
1762                .await
1763                .expect("done event should send");
1764        });
1765
1766        let base_url = format!("http://{address}/v1");
1767        let client = crate::providers::openai::Client::builder()
1768            .api_key("test-key")
1769            .base_url(&base_url)
1770            .build()
1771            .expect("client should build");
1772        let model = client.completion_model("gpt-4o");
1773        let mut session = client
1774            .responses_websocket("gpt-4o")
1775            .await
1776            .expect("session should connect");
1777
1778        session
1779            .send(model.completion_request("first").build())
1780            .await
1781            .expect("first request should send");
1782        let error = session
1783            .wait_for_completed_response()
1784            .await
1785            .expect_err("failed response should error");
1786        assert!(error.to_string().contains("failed response"));
1787        assert_eq!(session.previous_response_id(), None);
1788
1789        session
1790            .send(model.completion_request("second").build())
1791            .await
1792            .expect("second request should send");
1793        let second = session
1794            .wait_for_completed_response()
1795            .await
1796            .expect("second response should complete");
1797        assert_eq!(second.id, "resp_2");
1798        assert_eq!(session.previous_response_id(), Some("resp_2"));
1799
1800        server.await.expect("server task should finish");
1801    }
1802
1803    #[test]
1804    fn websocket_url_converts_http_to_ws() {
1805        let url = websocket_url("http://localhost:8080/v1").expect("url should convert");
1806        assert_eq!(url, "ws://localhost:8080/v1/responses");
1807    }
1808
1809    #[test]
1810    fn websocket_url_rejects_unsupported_scheme() {
1811        let result = websocket_url("ftp://example.com/v1");
1812        assert!(result.is_err());
1813    }
1814
1815    #[test]
1816    fn websocket_url_trims_trailing_slash() {
1817        let url = websocket_url("https://api.openai.com/v1/").expect("url should convert");
1818        assert_eq!(url, "wss://api.openai.com/v1/responses");
1819    }
1820
1821    #[test]
1822    fn unknown_event_type_is_skipped() {
1823        let payload = json!({
1824            "type": "response.some_future_event",
1825            "data": "hello"
1826        });
1827
1828        let result =
1829            parse_server_event(&payload.to_string()).expect("unknown event should not error");
1830        assert!(result.is_none(), "unknown event should be skipped");
1831    }
1832
1833    #[test]
1834    fn malformed_known_event_returns_error() {
1835        let payload = json!({
1836            "type": "response.completed"
1837        });
1838
1839        let error = parse_server_event(&payload.to_string())
1840            .expect_err("malformed known event should error");
1841        assert!(
1842            error.to_string().contains("StreamingCompletionChunk"),
1843            "expected strict decode failure, got {error}"
1844        );
1845    }
1846
1847    #[tokio::test]
1848    async fn close_is_idempotent() {
1849        let listener = TcpListener::bind("127.0.0.1:0")
1850            .await
1851            .expect("listener should bind");
1852        let address = listener.local_addr().expect("listener should have address");
1853
1854        let server = tokio::spawn(async move {
1855            let (stream, _) = listener.accept().await.expect("server should accept");
1856            let mut socket = accept_async(stream)
1857                .await
1858                .expect("server should upgrade websocket");
1859
1860            let message = socket
1861                .next()
1862                .await
1863                .expect("close frame should arrive")
1864                .expect("close frame should be valid");
1865            assert!(
1866                matches!(message, Message::Close(_)),
1867                "expected close frame, got {message:?}"
1868            );
1869        });
1870
1871        let base_url = format!("http://{address}/v1");
1872        let client = crate::providers::openai::Client::builder()
1873            .api_key("test-key")
1874            .base_url(&base_url)
1875            .build()
1876            .expect("client should build");
1877        let mut session = client
1878            .responses_websocket("gpt-4o")
1879            .await
1880            .expect("session should connect");
1881
1882        session.close().await.expect("first close should succeed");
1883        session.close().await.expect("second close should succeed");
1884
1885        server.await.expect("server task should finish");
1886    }
1887
1888    #[tokio::test]
1889    async fn send_while_in_flight_returns_error() {
1890        let listener = TcpListener::bind("127.0.0.1:0")
1891            .await
1892            .expect("listener should bind");
1893        let address = listener.local_addr().expect("listener should have address");
1894
1895        let server = tokio::spawn(async move {
1896            let (stream, _) = listener.accept().await.expect("server should accept");
1897            let mut socket = accept_async(stream)
1898                .await
1899                .expect("server should upgrade websocket");
1900
1901            // Read the first request but don't respond — keep it in-flight
1902            let _request = socket
1903                .next()
1904                .await
1905                .expect("request should exist")
1906                .expect("request should be valid");
1907
1908            // Wait for client to finish its test
1909            sleep(Duration::from_millis(100)).await;
1910            let _ = socket.close(None).await;
1911        });
1912
1913        let base_url = format!("http://{address}/v1");
1914        let client = crate::providers::openai::Client::builder()
1915            .api_key("test-key")
1916            .base_url(&base_url)
1917            .build()
1918            .expect("client should build");
1919        let model = client.completion_model("gpt-4o");
1920        let mut session = client
1921            .responses_websocket("gpt-4o")
1922            .await
1923            .expect("session should connect");
1924
1925        session
1926            .send(model.completion_request("first").build())
1927            .await
1928            .expect("first request should send");
1929
1930        let error = session
1931            .send(model.completion_request("second").build())
1932            .await
1933            .expect_err("second send while in-flight should error");
1934        assert!(
1935            error.to_string().contains("already in flight"),
1936            "expected in-flight error, got {error}"
1937        );
1938
1939        server.await.expect("server task should finish");
1940    }
1941
1942    #[tokio::test]
1943    async fn send_after_close_returns_error() {
1944        let listener = TcpListener::bind("127.0.0.1:0")
1945            .await
1946            .expect("listener should bind");
1947        let address = listener.local_addr().expect("listener should have address");
1948
1949        let server = tokio::spawn(async move {
1950            let (stream, _) = listener.accept().await.expect("server should accept");
1951            let _socket = accept_async(stream)
1952                .await
1953                .expect("server should upgrade websocket");
1954            sleep(Duration::from_millis(100)).await;
1955        });
1956
1957        let base_url = format!("http://{address}/v1");
1958        let client = crate::providers::openai::Client::builder()
1959            .api_key("test-key")
1960            .base_url(&base_url)
1961            .build()
1962            .expect("client should build");
1963        let model = client.completion_model("gpt-4o");
1964        let mut session = client
1965            .responses_websocket("gpt-4o")
1966            .await
1967            .expect("session should connect");
1968
1969        session.close().await.expect("close should succeed");
1970
1971        let error = session
1972            .send(model.completion_request("after close").build())
1973            .await
1974            .expect_err("send after close should error");
1975        assert!(
1976            error.to_string().contains("session is closed"),
1977            "expected closed-session error, got {error}"
1978        );
1979
1980        server.await.expect("server task should finish");
1981    }
1982
1983    #[tokio::test]
1984    async fn next_event_without_send_returns_error() {
1985        let listener = TcpListener::bind("127.0.0.1:0")
1986            .await
1987            .expect("listener should bind");
1988        let address = listener.local_addr().expect("listener should have address");
1989
1990        let server = tokio::spawn(async move {
1991            let (stream, _) = listener.accept().await.expect("server should accept");
1992            let _socket = accept_async(stream)
1993                .await
1994                .expect("server should upgrade websocket");
1995            sleep(Duration::from_millis(100)).await;
1996        });
1997
1998        let base_url = format!("http://{address}/v1");
1999        let client = crate::providers::openai::Client::builder()
2000            .api_key("test-key")
2001            .base_url(&base_url)
2002            .build()
2003            .expect("client should build");
2004        let mut session = client
2005            .responses_websocket("gpt-4o")
2006            .await
2007            .expect("session should connect");
2008
2009        let error = session
2010            .next_event()
2011            .await
2012            .expect_err("next_event without send should error");
2013        assert!(
2014            error
2015                .to_string()
2016                .contains("No OpenAI websocket response is currently in flight"),
2017            "expected not-in-flight error, got {error}"
2018        );
2019
2020        server.await.expect("server task should finish");
2021    }
2022
2023    #[tokio::test]
2024    async fn unknown_event_is_skipped_and_reasoning_metadata_is_preserved() {
2025        let listener = TcpListener::bind("127.0.0.1:0")
2026            .await
2027            .expect("listener should bind");
2028        let address = listener.local_addr().expect("listener should have address");
2029
2030        let server = tokio::spawn(async move {
2031            let (stream, _) = listener.accept().await.expect("server should accept");
2032            let mut socket = accept_async(stream)
2033                .await
2034                .expect("server should upgrade websocket");
2035
2036            let _request = socket
2037                .next()
2038                .await
2039                .expect("request should exist")
2040                .expect("request should be valid");
2041
2042            // Send an unknown event type first
2043            socket
2044                .send(Message::text(
2045                    json!({
2046                        "type": "response.some_future_event",
2047                        "data": "should be skipped"
2048                    })
2049                    .to_string(),
2050                ))
2051                .await
2052                .expect("unknown event should send");
2053
2054            // Then send the real completed response, including reasoning
2055            // metadata to verify that the WebSocket path preserves it.
2056            let mut response = sample_response(ResponseStatus::Completed);
2057            response.id = "resp_after_unknown".to_string();
2058            response.reasoning_metadata = Some(
2059                json!({
2060                    "context": "all_turns",
2061                    "effort": "ultra",
2062                    "summary": null,
2063                    "future_control": true
2064                })
2065                .as_object()
2066                .expect("reasoning metadata should be an object")
2067                .clone(),
2068            );
2069            response.reasoning_context = Some("all_turns".to_string());
2070            let response = serde_json::to_value(response).expect("response should serialize");
2071
2072            socket
2073                .send(Message::text(
2074                    json!({
2075                        "type": "response.completed",
2076                        "sequence_number": 1,
2077                        "response": response,
2078                    })
2079                    .to_string(),
2080                ))
2081                .await
2082                .expect("completed event should send");
2083        });
2084
2085        let base_url = format!("http://{address}/v1");
2086        let client = crate::providers::openai::Client::builder()
2087            .api_key("test-key")
2088            .base_url(&base_url)
2089            .build()
2090            .expect("client should build");
2091        let model = client.completion_model("gpt-4o");
2092        let mut session = client
2093            .responses_websocket("gpt-4o")
2094            .await
2095            .expect("session should connect");
2096
2097        session
2098            .send(model.completion_request("hello").build())
2099            .await
2100            .expect("send should succeed");
2101        let response = session
2102            .wait_for_completed_response()
2103            .await
2104            .expect("response should complete despite unknown event");
2105        assert_eq!(response.id, "resp_after_unknown");
2106        assert_eq!(response.reasoning_context.as_deref(), Some("all_turns"));
2107        assert_eq!(
2108            response.reasoning_metadata.as_ref(),
2109            json!({
2110                "context": "all_turns",
2111                "effort": "ultra",
2112                "summary": null,
2113                "future_control": true
2114            })
2115            .as_object()
2116        );
2117
2118        server.await.expect("server task should finish");
2119    }
2120}