Skip to main content

rig_core/providers/openai/responses_api/
websocket.rs

1//! Stateful Responses WebSocket sessions over caller-supplied connections.
2//! Sessions permit one in-flight turn and chain completed or incomplete response IDs.
3//!
4//! ```
5//! use rig_core::providers::openai::responses_api::websocket::ResponsesWebSocketCreateOptions;
6//! let options = ResponsesWebSocketCreateOptions::warmup();
7//! assert_eq!(options.generate, Some(false));
8//! ```
9
10use crate::completion;
11use crate::driver::Model;
12use crate::error::{EncodeError, ProviderError};
13use crate::http_client::{self, NoBody};
14use crate::operation::{Completion, Turn};
15use crate::providers::openai::responses_api::streaming::{
16    ItemChunk, ResponseChunk, ResponseChunkKind, ResponsesDecoder, StreamingCompletionChunk,
17    classify_responses_frame,
18};
19use crate::providers::openai::responses_api::wire::Responses;
20use crate::streaming::Item;
21use crate::wire::{Flow, Reply, Shared, Wire, WireFrame};
22use crate::ws_client::{
23    BoxedWebSocketConnection, ConnectOptions, Frame, WebSocketClientExt, WebSocketConnection,
24};
25use serde::{Deserialize, Serialize};
26use serde_json::{Map, Value};
27use std::time::Duration;
28
29use crate::providers::openai::responses_api::{CompletionResponse, ResponseStatus};
30
31/// The websocket endpoint's path, appended to the client's configured base URL.
32const WEBSOCKET_PATH: &str = "responses";
33
34const DEFAULT_CONNECT_TIMEOUT: Duration = Duration::from_secs(30);
35
36/// Request-ID header read from rejected WebSocket upgrades.
37const REQUEST_ID_HEADER: Option<&'static str> =
38    crate::providers::openai::wire::OPENAI.request_id_header;
39
40/// Options for a `response.create` message sent over OpenAI WebSocket mode.
41#[derive(Debug, Clone, Default, Serialize, Deserialize)]
42pub struct ResponsesWebSocketCreateOptions {
43    /// When set to `false`, OpenAI prepares request state without generating a model output.
44    ///
45    /// This is the "warmup" mode described in the OpenAI WebSocket mode guide.
46    #[serde(skip_serializing_if = "Option::is_none")]
47    pub generate: Option<bool>,
48}
49
50impl ResponsesWebSocketCreateOptions {
51    /// Creates warmup options equivalent to `generate: false`.
52    #[must_use]
53    pub fn warmup() -> Self {
54        Self {
55            generate: Some(false),
56        }
57    }
58}
59
60#[derive(Debug, Clone, Serialize)]
61struct ResponsesWebSocketClientEvent {
62    #[serde(rename = "type")]
63    kind: ResponsesWebSocketClientEventKind,
64    #[serde(flatten)]
65    request: crate::providers::openai::responses_api::CompletionRequest,
66    #[serde(skip_serializing_if = "Option::is_none")]
67    generate: Option<bool>,
68}
69
70#[derive(Debug, Clone, Serialize)]
71enum ResponsesWebSocketClientEventKind {
72    #[serde(rename = "response.create")]
73    ResponseCreate,
74}
75
76/// A protocol error event emitted by OpenAI WebSocket mode.
77#[derive(Debug, Clone, Serialize, Deserialize)]
78pub struct ResponsesWebSocketErrorEvent {
79    /// The event type.
80    #[serde(rename = "type")]
81    pub kind: ResponsesWebSocketErrorEventKind,
82    /// The provider error payload.
83    pub error: ResponsesWebSocketErrorPayload,
84}
85
86impl std::fmt::Display for ResponsesWebSocketErrorEvent {
87    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
88        self.error.fmt(f)
89    }
90}
91
92/// The event kind for an OpenAI WebSocket protocol error.
93#[derive(Debug, Clone, Serialize, Deserialize)]
94pub enum ResponsesWebSocketErrorEventKind {
95    #[serde(rename = "error")]
96    Error,
97}
98
99/// The payload carried by an OpenAI WebSocket protocol error event.
100#[derive(Debug, Clone, Default, Serialize, Deserialize)]
101pub struct ResponsesWebSocketErrorPayload {
102    /// Provider-specific error code when supplied.
103    #[serde(skip_serializing_if = "Option::is_none")]
104    pub code: Option<String>,
105    /// Human-readable error message.
106    #[serde(skip_serializing_if = "Option::is_none")]
107    pub message: Option<String>,
108    /// Any extra fields supplied by the provider.
109    #[serde(flatten, default)]
110    pub extra: Map<String, Value>,
111}
112
113impl std::fmt::Display for ResponsesWebSocketErrorPayload {
114    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
115        match (&self.code, &self.message) {
116            (Some(code), Some(message)) => write!(f, "{code}: {message}"),
117            (None, Some(message)) => f.write_str(message),
118            (Some(code), None) => f.write_str(code),
119            (None, None) => f.write_str("OpenAI websocket error"),
120        }
121    }
122}
123
124/// The optional `response.done` event emitted by OpenAI WebSocket mode.
125#[derive(Debug, Clone, Serialize, Deserialize)]
126pub struct ResponsesWebSocketDoneEvent {
127    /// The event type.
128    #[serde(rename = "type")]
129    pub kind: ResponsesWebSocketDoneEventKind,
130    /// The provider payload for the finished response.
131    pub response: Value,
132}
133
134impl ResponsesWebSocketDoneEvent {
135    /// Returns the response ID if the payload includes one.
136    #[must_use]
137    pub fn response_id(&self) -> Option<&str> {
138        self.response.get("id").and_then(Value::as_str)
139    }
140
141    fn status(&self) -> Option<ResponseStatus> {
142        self.response
143            .get("status")
144            .cloned()
145            .and_then(|status| serde_json::from_value(status).ok())
146    }
147
148    fn as_completion_response(&self) -> Option<CompletionResponse> {
149        serde_json::from_value(self.response.clone()).ok()
150    }
151}
152
153/// The event kind for the terminal websocket event.
154#[derive(Debug, Clone, Serialize, Deserialize)]
155pub enum ResponsesWebSocketDoneEventKind {
156    #[serde(rename = "response.done")]
157    ResponseDone,
158}
159
160/// A server event emitted by OpenAI WebSocket mode.
161#[derive(Debug, Clone)]
162pub enum ResponsesWebSocketEvent {
163    /// A response lifecycle event such as `response.created` or `response.completed`.
164    Response(ResponseChunk),
165    /// A streaming item/delta event such as `response.output_text.delta`.
166    Item(ItemChunk),
167    /// A protocol-level websocket error event.
168    Error(ResponsesWebSocketErrorEvent),
169    /// An optional `response.done` event emitted by OpenAI over WebSockets.
170    Done(ResponsesWebSocketDoneEvent),
171    /// Unrecognized event retained for [`Item::Unknown`] passthrough.
172    Unknown(crate::streaming::UnknownPayload),
173}
174
175impl ResponsesWebSocketEvent {
176    /// Returns the response ID when the event includes one.
177    #[must_use]
178    pub fn response_id(&self) -> Option<&str> {
179        match self {
180            Self::Response(chunk) => Some(&chunk.response.id),
181            Self::Done(done) => done.response_id(),
182            Self::Item(_) | Self::Error(_) | Self::Unknown(_) => None,
183        }
184    }
185
186    /// Returns `true` when this event ends the current in-flight websocket turn.
187    #[must_use]
188    pub fn is_terminal(&self) -> bool {
189        match self {
190            Self::Response(chunk) => matches!(
191                chunk.kind,
192                ResponseChunkKind::ResponseCompleted
193                    | ResponseChunkKind::ResponseFailed
194                    | ResponseChunkKind::ResponseIncomplete
195            ),
196            Self::Error(_) | Self::Done(_) => true,
197            Self::Item(_) | Self::Unknown(_) => false,
198        }
199    }
200}
201
202/// A builder for an OpenAI Responses WebSocket session.
203///
204/// The default builder applies a 30 second connection timeout and leaves the
205/// per-event timeout disabled.
206pub struct ResponsesWebSocketSessionBuilder {
207    wire: Responses,
208    connect_timeout: Option<Duration>,
209    event_timeout: Option<Duration>,
210}
211
212impl ResponsesWebSocketSessionBuilder {
213    pub(crate) fn new(wire: Responses) -> Self {
214        Self {
215            wire,
216            connect_timeout: Some(DEFAULT_CONNECT_TIMEOUT),
217            event_timeout: None,
218        }
219    }
220
221    /// Sets the timeout for establishing the websocket connection.
222    #[must_use]
223    pub fn connect_timeout(mut self, timeout: Duration) -> Self {
224        self.connect_timeout = Some(timeout);
225        self
226    }
227
228    /// Disables the websocket connection timeout.
229    #[must_use]
230    pub fn without_connect_timeout(mut self) -> Self {
231        self.connect_timeout = None;
232        self
233    }
234
235    /// Sets the timeout for waiting on the next websocket event.
236    #[must_use]
237    pub fn event_timeout(mut self, timeout: Duration) -> Self {
238        self.event_timeout = Some(timeout);
239        self
240    }
241
242    /// Disables the websocket event timeout.
243    #[must_use]
244    pub fn without_event_timeout(mut self) -> Self {
245        self.event_timeout = None;
246        self
247    }
248}
249
250impl ResponsesWebSocketSessionBuilder {
251    /// Open a session over the bundled tungstenite backend with the configured
252    /// timeouts. Return handshake construction, transport, or provider errors.
253    #[cfg(all(feature = "tungstenite", not(target_family = "wasm")))]
254    #[cfg_attr(docsrs, doc(cfg(feature = "tungstenite")))]
255    pub async fn connect(self) -> Result<ResponsesWebSocketSession, ProviderError> {
256        self.connect_with(&rig_tungstenite::TungsteniteClient::new())
257            .await
258    }
259
260    /// Open a session over `backend` with the configured timeouts.
261    /// Return handshake construction, transport, or provider errors.
262    pub async fn connect_with<W>(
263        self,
264        backend: &W,
265    ) -> Result<ResponsesWebSocketSession, ProviderError>
266    where
267        W: WebSocketClientExt,
268    {
269        ResponsesWebSocketSession::connect_with_timeouts(
270            backend,
271            self.wire,
272            self.connect_timeout,
273            self.event_timeout,
274        )
275        .await
276    }
277}
278
279/// Sequential Responses session with automatic response-ID chaining.
280/// Completed and incomplete responses update the chain unless a request supplies
281/// its own `previous_response_id`. Call [`Self::close`] to perform a close handshake.
282pub struct ResponsesWebSocketSession {
283    wire: Responses,
284    previous_response_id: Option<String>,
285    pending_done_response_id: Option<String>,
286    socket: BoxedWebSocketConnection,
287    in_flight: bool,
288    event_timeout: Option<Duration>,
289    closed: bool,
290    failed: bool,
291}
292
293impl ResponsesWebSocketSession {
294    async fn connect_with_timeouts<W>(
295        backend: &W,
296        wire: Responses,
297        connect_timeout: Option<Duration>,
298        event_timeout: Option<Duration>,
299    ) -> Result<Self, ProviderError>
300    where
301        W: WebSocketClientExt,
302    {
303        let request = websocket_request(&wire)?;
304        let socket = backend
305            .connect(request, ConnectOptions::new().with_timeout(connect_timeout))
306            .await
307            .map_err(websocket_provider_error)?;
308
309        Ok(Self::from_connection(wire, socket, event_timeout))
310    }
311
312    /// Build a session over an already-open, authenticated connection.
313    /// `event_timeout: None` waits indefinitely for each event.
314    pub fn from_connection(
315        wire: Responses,
316        connection: BoxedWebSocketConnection,
317        event_timeout: Option<Duration>,
318    ) -> Self {
319        Self {
320            wire,
321            previous_response_id: None,
322            pending_done_response_id: None,
323            socket: connection,
324            in_flight: false,
325            event_timeout,
326            closed: false,
327            failed: false,
328        }
329    }
330
331    /// Return the response ID retained for automatic chaining, if any.
332    #[must_use]
333    pub fn previous_response_id(&self) -> Option<&str> {
334        self.previous_response_id.as_deref()
335    }
336
337    /// Clears the cached `previous_response_id` so the next turn starts a fresh chain.
338    pub fn clear_previous_response_id(&mut self) {
339        self.previous_response_id = None;
340    }
341
342    /// Sends a `response.create` event for a Rig completion request.
343    pub async fn send(
344        &mut self,
345        completion_request: crate::completion::CompletionRequest,
346    ) -> Result<(), ProviderError> {
347        self.send_with_options(
348            completion_request,
349            ResponsesWebSocketCreateOptions::default(),
350        )
351        .await
352    }
353
354    /// Sends a `response.create` event with explicit websocket-mode options.
355    pub async fn send_with_options(
356        &mut self,
357        completion_request: crate::completion::CompletionRequest,
358        options: ResponsesWebSocketCreateOptions,
359    ) -> Result<(), ProviderError> {
360        self.ensure_open()?;
361
362        if self.in_flight {
363            return Err(ProviderError::Provider(
364                "An OpenAI websocket response is already in flight on this session".to_string(),
365            ));
366        }
367
368        let payload = ResponsesWebSocketClientEvent {
369            kind: ResponsesWebSocketClientEventKind::ResponseCreate,
370            request: self.prepare_request(completion_request)?,
371            generate: options.generate,
372        };
373
374        crate::providers::internal::trace_json(
375            crate::providers::internal::LogTarget::Completions,
376            "OpenAI websocket request",
377            &payload,
378        );
379
380        let payload = serde_json::to_string(&payload)?;
381
382        if let Err(error) = self.socket.send(Frame::Text(payload)).await {
383            return Err(self.fail_session(websocket_provider_error(error)));
384        }
385        self.in_flight = true;
386
387        Ok(())
388    }
389
390    /// Reads the next server event for the current in-flight turn.
391    pub async fn next_event(&mut self) -> Result<ResponsesWebSocketEvent, ProviderError> {
392        self.next_event_with_payload().await.map(|(event, _)| event)
393    }
394
395    /// Reads the next lifecycle event and retains its payload for content decoding
396    /// by [`ResponsesDecoder`]. Returns the same session errors as [`Self::next_event`].
397    async fn next_event_with_payload(
398        &mut self,
399    ) -> Result<(ResponsesWebSocketEvent, String), ProviderError> {
400        self.ensure_open()?;
401
402        if !self.in_flight {
403            return Err(ProviderError::Provider(
404                "No OpenAI websocket response is currently in flight on this session".to_string(),
405            ));
406        }
407
408        loop {
409            let message = match self.read_next_frame().await? {
410                Ok(message) => message,
411                Err(error) => return Err(self.fail_session(websocket_provider_error(error))),
412            };
413
414            let Some(message) = message else {
415                self.mark_closed();
416                return Err(ProviderError::Provider(
417                    "The OpenAI websocket connection closed before the turn finished".to_string(),
418                ));
419            };
420
421            let payload = match websocket_frame_to_text(message) {
422                Ok(Some(payload)) => payload,
423                Ok(None) => continue,
424                Err(error) => return Err(self.fail_session(error)),
425            };
426            let event = match parse_server_event(&payload) {
427                Ok(Some(event)) => event,
428                Ok(None) => continue,
429                Err(error) => return Err(self.fail_session(error)),
430            };
431            if let ResponsesWebSocketEvent::Done(done) = &event {
432                // OpenAI may emit `response.done` after the turn has already ended at
433                // `response.completed`. Ignore that trailing event on the next turn.
434                if self.pending_done_response_id.as_deref() == done.response_id() {
435                    self.pending_done_response_id = None;
436                    continue;
437                }
438            }
439            self.update_state_for_event(&event);
440            return Ok((event, payload));
441        }
442    }
443
444    /// Sends a warmup turn (`generate: false`) and returns the resulting response ID.
445    pub async fn warmup(
446        &mut self,
447        completion_request: crate::completion::CompletionRequest,
448    ) -> Result<String, ProviderError> {
449        self.send_with_options(
450            completion_request,
451            ResponsesWebSocketCreateOptions::warmup(),
452        )
453        .await?;
454        let response = self.wait_for_completed_response().await?;
455        Ok(response.id)
456    }
457
458    /// Sends a completion turn and collects the final OpenAI response,
459    /// normalized; its `raw` is the provider's own terminal response object.
460    pub async fn completion(
461        &mut self,
462        completion_request: crate::completion::CompletionRequest,
463    ) -> Result<completion::CompletionResponse, ProviderError> {
464        let provider = self.wire.describe().name.to_owned();
465        self.send(completion_request).await?;
466        let (response, folded) = self.wait_for_terminal_response().await?;
467        if folded.choice.is_empty() {
468            // The turn carried no content events but its terminal body
469            // restates `output[]` (the shape a warmed-up or replayed session
470            // answers with): fold that body, through the same decoder's
471            // unary variant.
472            return super::wire::fold_body(&provider, response);
473        }
474        Ok(folded)
475    }
476
477    /// Closes the websocket connection.
478    ///
479    /// Call this when you are finished with the session so the websocket can
480    /// terminate with a clean close handshake.
481    pub async fn close(&mut self) -> Result<(), ProviderError> {
482        if self.closed {
483            return Ok(());
484        }
485
486        let result = self
487            .socket
488            .close(None)
489            .await
490            .map_err(websocket_provider_error);
491        self.mark_closed();
492        result
493    }
494
495    fn prepare_request(
496        &self,
497        completion_request: crate::completion::CompletionRequest,
498    ) -> Result<crate::providers::openai::responses_api::CompletionRequest, ProviderError> {
499        // The session sends without the driver, so it runs the driver's check.
500        completion_request.validate_message_content()?;
501        let (completion_request, issuers) = crate::providers::openai::wire::scope_reasoning(
502            &self.wire.provider.dialect,
503            &self.wire.model,
504            completion_request,
505        )?;
506        let mut request = self
507            .wire
508            .responses_request(completion_request, issuers, false)?;
509
510        // WebSocket mode is always event-driven, so these HTTP/SSE-specific flags
511        // are ignored by the provider and only add noise to the payload.
512        request.stream = None;
513        request.additional_parameters.background = None;
514
515        if request.additional_parameters.previous_response_id.is_none() {
516            request
517                .additional_parameters
518                .previous_response_id
519                .clone_from(&self.previous_response_id);
520        }
521
522        Ok(request)
523    }
524
525    async fn wait_for_completed_response(&mut self) -> Result<CompletionResponse, ProviderError> {
526        Ok(self.wait_for_terminal_response().await?.0)
527    }
528
529    /// Decode the turn's events as they arrive, and return the provider's
530    /// completed or incomplete response with the reply they fold into.
531    /// Transport, protocol, and decoder failures return an error. A terminal
532    /// event without a response body is an error.
533    async fn wait_for_terminal_response(
534        &mut self,
535    ) -> Result<(CompletionResponse, completion::CompletionResponse), ProviderError> {
536        let provider = self.wire.describe().name.to_owned();
537        let wire = self.wire.clone();
538        // The reply's state and its decoder live for this turn only; the
539        // decoder's handles are branded with the borrow of that state.
540        let reply = std::sync::Mutex::new(Shared::new(Turn::new(provider.clone())));
541        let mut decoder = wire.decoder();
542        loop {
543            let (event, payload) = self.next_event_with_payload().await?;
544            match event {
545                ResponsesWebSocketEvent::Response(chunk) => {
546                    let terminal = matches!(
547                        chunk.kind,
548                        ResponseChunkKind::ResponseCompleted
549                            | ResponseChunkKind::ResponseFailed
550                            | ResponseChunkKind::ResponseIncomplete
551                    );
552                    if !terminal {
553                        feed(&mut decoder, &reply, payload)?;
554                        continue;
555                    }
556                    // A failed turn is reported from its own envelope; only a
557                    // completed or incomplete one reaches the decoder, whose
558                    // end closes the turn.
559                    let response = terminal_response_result(chunk.response)?;
560                    let ended = feed(&mut decoder, &reply, payload)?;
561                    let folded = fold_reply(reply, ended, &provider, &response)?;
562                    return Ok((response, folded));
563                }
564                ResponsesWebSocketEvent::Done(done) => {
565                    if let Some(response) = done.as_completion_response() {
566                        // A failed turn is reported from its own envelope, as
567                        // on the `response.failed` path.
568                        let response = terminal_response_result(response)?;
569                        // `response.done` carries the response object itself,
570                        // which is the decoder's whole-body shape: hand it
571                        // over as the frame it is.
572                        let body = serde_json::to_string(&done.response)?;
573                        let ended = feed(&mut decoder, &reply, body)?;
574                        let folded = fold_reply(reply, ended, &provider, &response)?;
575                        return Ok((response, folded));
576                    }
577
578                    let message = if let Some(response_id) = done.response_id() {
579                        format!(
580                            "OpenAI websocket turn ended with response.done before a terminal response body was available (response_id={response_id})"
581                        )
582                    } else {
583                        "OpenAI websocket turn ended with response.done before a terminal response body was available"
584                            .to_string()
585                    };
586
587                    return Err(ProviderError::Provider(message));
588                }
589                ResponsesWebSocketEvent::Error(error) => {
590                    // Genuine provider error event: preserve the serialized payload
591                    // (code + message + any extra fields) so provider_response_json()
592                    // parses it, matching the response.failed path. No HTTP status on
593                    // the websocket stream, so status: None.
594                    return Err(provider_error_from_event(&error));
595                }
596                // Unknown frames keep their raw payload.
597                ResponsesWebSocketEvent::Item(_) | ResponsesWebSocketEvent::Unknown(_) => {
598                    feed(&mut decoder, &reply, payload)?;
599                }
600            }
601        }
602    }
603
604    fn update_state_for_event(&mut self, event: &ResponsesWebSocketEvent) {
605        match event {
606            ResponsesWebSocketEvent::Response(chunk) => match chunk.kind {
607                // An incomplete turn still produced a response the next turn
608                // can chain from, so it keeps `previous_response_id` like a
609                // completed one.
610                ResponseChunkKind::ResponseCompleted | ResponseChunkKind::ResponseIncomplete => {
611                    let response_id = chunk.response.id.clone();
612                    self.previous_response_id = Some(response_id.clone());
613                    self.pending_done_response_id = Some(response_id);
614                    self.in_flight = false;
615                }
616                ResponseChunkKind::ResponseFailed => {
617                    self.pending_done_response_id = Some(chunk.response.id.clone());
618                    self.previous_response_id = None;
619                    self.in_flight = false;
620                }
621                ResponseChunkKind::ResponseCreated | ResponseChunkKind::ResponseInProgress => {}
622            },
623            ResponsesWebSocketEvent::Done(done) => {
624                match done.status() {
625                    Some(ResponseStatus::Completed) | Some(ResponseStatus::Incomplete) => {
626                        if let Some(response_id) = done.response_id() {
627                            self.previous_response_id = Some(response_id.to_string());
628                        }
629                    }
630                    Some(ResponseStatus::Failed)
631                    | Some(ResponseStatus::Cancelled)
632                    | Some(ResponseStatus::Other(_)) => {
633                        self.previous_response_id = None;
634                    }
635                    Some(ResponseStatus::InProgress | ResponseStatus::Queued) | None => {}
636                }
637                self.pending_done_response_id = None;
638                self.in_flight = false;
639            }
640            ResponsesWebSocketEvent::Error(_) => {
641                self.previous_response_id = None;
642                self.pending_done_response_id = None;
643                self.in_flight = false;
644            }
645            // An unknown frame carries no turn-lifecycle signal.
646            ResponsesWebSocketEvent::Item(_) | ResponsesWebSocketEvent::Unknown(_) => {}
647        }
648    }
649
650    fn abort_turn(&mut self) {
651        self.previous_response_id = None;
652        self.pending_done_response_id = None;
653        self.in_flight = false;
654    }
655
656    fn mark_closed(&mut self) {
657        self.abort_turn();
658        self.closed = true;
659        self.failed = false;
660    }
661
662    fn mark_failed(&mut self) {
663        self.abort_turn();
664        self.failed = true;
665    }
666
667    fn ensure_open(&self) -> Result<(), ProviderError> {
668        if self.closed || self.failed {
669            return Err(ProviderError::Provider(
670                "The OpenAI websocket session is closed".to_string(),
671            ));
672        }
673
674        Ok(())
675    }
676
677    fn fail_session(&mut self, error: ProviderError) -> ProviderError {
678        self.mark_failed();
679        error
680    }
681
682    /// Read a frame with a WASM-compatible timeout.
683    /// Timeout failure marks the session failed; transport results remain nested.
684    async fn read_next_frame(
685        &mut self,
686    ) -> Result<http_client::Result<Option<Frame>>, ProviderError> {
687        let Some(timeout_duration) = self.event_timeout else {
688            return Ok(self.socket.recv().await);
689        };
690
691        match crate::wasm_compat::timeout(timeout_duration, self.socket.recv()).await {
692            Ok(message) => Ok(message),
693            Err(_) => Err(self.fail_session(event_timeout_error(timeout_duration))),
694        }
695    }
696}
697
698impl Drop for ResponsesWebSocketSession {
699    fn drop(&mut self) {
700        if !self.closed {
701            tracing::warn!(
702                target: "rig::completions",
703                in_flight = self.in_flight,
704                "Dropping an OpenAI websocket session without calling close(); the connection will end without a close handshake"
705            );
706        }
707    }
708}
709
710/// Feed one message to the turn's decoder. Returns whether it ended the
711/// turn.
712fn feed<'id>(
713    decoder: &mut ResponsesDecoder<'id>,
714    reply: &'id std::sync::Mutex<Shared<Completion>>,
715    payload: String,
716) -> Result<bool, ProviderError> {
717    crate::driver::step(decoder, reply, WireFrame::Text(payload))
718        .map(|step| matches!(step, Flow::Ended(_)))
719}
720
721/// Fold the turn's events and end into the normalized response, retaining
722/// the terminal body as raw JSON.
723fn fold_reply(
724    reply: std::sync::Mutex<Shared<Completion>>,
725    ended: bool,
726    provider: &str,
727    response: &CompletionResponse,
728) -> Result<completion::CompletionResponse, ProviderError> {
729    let fed = if ended {
730        Ok(())
731    } else {
732        Err(ProviderError::Truncated)
733    };
734    let reply_of = Reply {
735        provider: provider.to_owned(),
736        raw: serde_json::to_value(response)?,
737        // The websocket carries no reply headers past the handshake.
738        provider_request_id: None,
739    };
740    crate::driver::settle(reply, fed, reply_of).outcome
741}
742
743fn terminal_response_result(
744    response: CompletionResponse,
745) -> Result<CompletionResponse, ProviderError> {
746    match response.status {
747        ResponseStatus::Completed => Ok(response),
748        // Preserve provider error envelopes as reserialized JSON without an HTTP status.
749        // Without an error object, return a local diagnostic instead.
750        ResponseStatus::Failed => match response.error.as_ref() {
751            Some(error) => Err(ProviderError::from_provider_body(
752                serde_json::to_string(&response).unwrap_or_else(|_| error.message.clone()),
753            )),
754            None => Err(ProviderError::Provider(response_error_message(
755                "failed response",
756            ))),
757        },
758        // An incomplete response (e.g. hitting `max_output_tokens`) is a
759        // genuine terminal: the partial output and usage are kept, and the
760        // normalization path maps the status/incomplete_details to a finish
761        // reason via `map_finish_reason`, matching the unary and SSE paths.
762        ResponseStatus::Incomplete => Ok(response),
763        other => Err(ProviderError::Provider(format!(
764            "OpenAI websocket response ended in state {other:?}"
765        ))),
766    }
767}
768
769fn response_error_message(fallback: &str) -> String {
770    format!("OpenAI websocket returned a {fallback}")
771}
772
773/// Preserve an error event as reserialized provider JSON without an HTTP status.
774/// Fall back to its display text if serialization fails.
775fn provider_error_from_event(error: &ResponsesWebSocketErrorEvent) -> ProviderError {
776    ProviderError::from_provider_body(
777        serde_json::to_string(&error).unwrap_or_else(|_| error.to_string()),
778    )
779}
780
781/// Decode WebSocket error and done events or delegate to Responses classification.
782/// Return parsing and triage errors; preserve unknown payloads.
783fn parse_server_event(payload: &str) -> Result<Option<ResponsesWebSocketEvent>, ProviderError> {
784    #[derive(Deserialize)]
785    struct EventType {
786        #[serde(rename = "type")]
787        kind: String,
788    }
789
790    let event_type = serde_json::from_str::<EventType>(payload)?;
791    match event_type.kind.as_str() {
792        "error" => serde_json::from_str(payload)
793            .map(|e| Some(ResponsesWebSocketEvent::Error(e)))
794            .map_err(ProviderError::from),
795        "response.done" => serde_json::from_str(payload)
796            .map(|d| Some(ResponsesWebSocketEvent::Done(d)))
797            .map_err(ProviderError::from),
798        _ => Ok(Some(
799            match crate::driver::triage(classify_responses_frame(payload))? {
800                Item::Event(StreamingCompletionChunk::Response(response)) => {
801                    ResponsesWebSocketEvent::Response(response)
802                }
803                Item::Event(StreamingCompletionChunk::Delta(item)) => {
804                    ResponsesWebSocketEvent::Item(item)
805                }
806                Item::Unknown(value) => ResponsesWebSocketEvent::Unknown(value),
807            },
808        )),
809    }
810}
811
812/// Lower one websocket frame onto the JSON payload the protocol carries.
813///
814/// `Ok(None)` is a frame with no protocol payload (a keepalive), which the
815/// session skips; a close frame mid-turn is an error naming the peer's reason.
816fn websocket_frame_to_text(frame: Frame) -> Result<Option<String>, ProviderError> {
817    match frame {
818        Frame::Text(text) => Ok(Some(text)),
819        Frame::Binary(bytes) => String::from_utf8(bytes.to_vec())
820            .map(Some)
821            .map_err(|error| ProviderError::Response(error.to_string())),
822        Frame::Ping(_) | Frame::Pong(_) => Ok(None),
823        Frame::Close(frame) => {
824            let reason = frame
825                .map(|frame| frame.reason)
826                .filter(|reason| !reason.is_empty())
827                .unwrap_or_else(|| "without a close reason".to_string());
828            Err(ProviderError::Provider(format!(
829                "The OpenAI websocket connection closed {reason}"
830            )))
831        }
832    }
833}
834
835/// Build the handshake request: the websocket URL derived from the client's
836/// base URL, carrying the client's own auth headers.
837///
838/// The backend supplies the websocket-specific handshake headers; this only
839/// states where to connect and who is connecting.
840fn websocket_request(wire: &Responses) -> Result<http_client::Request<NoBody>, EncodeError> {
841    let url = crate::ws_client::websocket_url(&wire.provider.base_url, WEBSOCKET_PATH)
842        .map_err(EncodeError::request)?;
843
844    let request = wire.provider.headers(
845        http_client::Request::builder()
846            .method(http::Method::GET)
847            .uri(url),
848    );
849
850    request.body(NoBody).map_err(|error| {
851        EncodeError::request(format!("Failed to build OpenAI websocket request: {error}"))
852    })
853}
854
855fn event_timeout_error(timeout: Duration) -> ProviderError {
856    ProviderError::Provider(format!(
857        "Timed out waiting for the next OpenAI websocket event after {timeout:?}"
858    ))
859}
860
861/// Convert transport errors, retaining rejected-upgrade status, body, and request ID.
862/// Failures without a provider response retain transport error classification.
863fn websocket_provider_error(error: http_client::Error) -> ProviderError {
864    let provider_request_id = error.non_success_headers().and_then(|headers| {
865        crate::providers::internal::request_id_from_headers(headers, REQUEST_ID_HEADER)
866    });
867    ProviderError::from_transport_error(error).with_provider_request_id(provider_request_id)
868}
869
870impl<T> Model<Responses, T> {
871    /// Start configuring a Responses WebSocket session for this model. Open
872    /// it with [`connect_with`](ResponsesWebSocketSessionBuilder::connect_with)
873    /// and a backend, or, under the `tungstenite` feature, with
874    /// [`connect`](ResponsesWebSocketSessionBuilder::connect).
875    pub fn responses_websocket(&self) -> ResponsesWebSocketSessionBuilder {
876        ResponsesWebSocketSessionBuilder::new(self.wire.clone())
877    }
878}
879
880/// Native sessions and builders satisfy Send and Sync.
881#[cfg(not(target_family = "wasm"))]
882const _: fn() = || {
883    fn assert_send_sync<T: Send + Sync>() {}
884    assert_send_sync::<ResponsesWebSocketSession>();
885    assert_send_sync::<ResponsesWebSocketSessionBuilder>();
886};
887
888#[cfg(test)]
889#[allow(clippy::expect_used, clippy::panic)]
890mod tests;