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