1use crate::completion::NormalizeCompletionResponse;
8use crate::completion::{self, CompletionError};
9use crate::http_client::HttpClientExt;
10use crate::providers::internal::adapter::{TriagedFrame, triage_frame};
11use crate::providers::openai::responses_api::streaming::{
12 ItemChunk, RawChoiceAccumulator, ResponseChunk, ResponseChunkKind, ResponsesStreamOptions,
13 StreamingCompletionChunk, classify_responses_frame, completion_response_from_raw_choices,
14};
15use crate::wasm_compat::{WasmCompatSend, WasmCompatSync};
16use futures::{SinkExt, StreamExt};
17use serde::{Deserialize, Serialize};
18use serde_json::{Map, Value};
19use std::time::Duration;
20use tokio::net::TcpStream;
21use tokio_tungstenite::{
22 MaybeTlsStream, WebSocketStream, connect_async,
23 tungstenite::{self, Message, client::IntoClientRequest},
24};
25use url::Url;
26
27use super::{CompletionResponse, ResponseStatus, ResponsesCompletionModel, ResponsesUsage};
28
29type OpenAIWebSocket = WebSocketStream<MaybeTlsStream<TcpStream>>;
30type WebSocketRawChoice = crate::streaming::RawStreamingChoice<
31 crate::providers::openai::responses_api::streaming::StreamingCompletionResponse,
32>;
33const DEFAULT_CONNECT_TIMEOUT: Duration = Duration::from_secs(30);
34const REQUEST_ID_HEADER: Option<&'static str> =
38 <crate::providers::openai::OpenAIResponsesExt as super::ResponsesProviderExt>::REQUEST_ID_HEADER;
39
40#[derive(Debug, Clone, Default, Serialize, Deserialize)]
42pub struct ResponsesWebSocketCreateOptions {
43 #[serde(skip_serializing_if = "Option::is_none")]
47 pub generate: Option<bool>,
48}
49
50impl ResponsesWebSocketCreateOptions {
51 #[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: super::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#[derive(Debug, Clone, Serialize, Deserialize)]
78pub struct ResponsesWebSocketErrorEvent {
79 #[serde(rename = "type")]
81 pub kind: ResponsesWebSocketErrorEventKind,
82 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#[derive(Debug, Clone, Serialize, Deserialize)]
94pub enum ResponsesWebSocketErrorEventKind {
95 #[serde(rename = "error")]
96 Error,
97}
98
99#[derive(Debug, Clone, Default, Serialize, Deserialize)]
101pub struct ResponsesWebSocketErrorPayload {
102 #[serde(skip_serializing_if = "Option::is_none")]
104 pub code: Option<String>,
105 #[serde(skip_serializing_if = "Option::is_none")]
107 pub message: Option<String>,
108 #[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#[derive(Debug, Clone, Serialize, Deserialize)]
126pub struct ResponsesWebSocketDoneEvent {
127 #[serde(rename = "type")]
129 pub kind: ResponsesWebSocketDoneEventKind,
130 pub response: Value,
132}
133
134impl ResponsesWebSocketDoneEvent {
135 #[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#[derive(Debug, Clone, Serialize, Deserialize)]
155pub enum ResponsesWebSocketDoneEventKind {
156 #[serde(rename = "response.done")]
157 ResponseDone,
158}
159
160#[derive(Debug, Clone)]
162pub enum ResponsesWebSocketEvent {
163 Response(Box<ResponseChunk>),
165 Item(ItemChunk),
167 Error(ResponsesWebSocketErrorEvent),
169 Done(ResponsesWebSocketDoneEvent),
171 Unknown(crate::streaming::UnknownPayload),
175}
176
177impl ResponsesWebSocketEvent {
178 #[must_use]
180 pub fn response_id(&self) -> Option<&str> {
181 match self {
182 Self::Response(chunk) => Some(&chunk.response.id),
183 Self::Done(done) => done.response_id(),
184 Self::Item(_) | Self::Error(_) | Self::Unknown(_) => None,
185 }
186 }
187
188 #[must_use]
190 pub fn is_terminal(&self) -> bool {
191 match self {
192 Self::Response(chunk) => matches!(
193 chunk.kind,
194 ResponseChunkKind::ResponseCompleted
195 | ResponseChunkKind::ResponseFailed
196 | ResponseChunkKind::ResponseIncomplete
197 ),
198 Self::Error(_) | Self::Done(_) => true,
199 Self::Item(_) | Self::Unknown(_) => false,
200 }
201 }
202}
203
204pub struct ResponsesWebSocketSessionBuilder<H = reqwest::Client> {
209 model: ResponsesCompletionModel<H>,
210 connect_timeout: Option<Duration>,
211 event_timeout: Option<Duration>,
212}
213
214impl<H> ResponsesWebSocketSessionBuilder<H> {
215 pub(crate) fn new(model: ResponsesCompletionModel<H>) -> Self {
216 Self {
217 model,
218 connect_timeout: Some(DEFAULT_CONNECT_TIMEOUT),
219 event_timeout: None,
220 }
221 }
222
223 #[must_use]
225 pub fn connect_timeout(mut self, timeout: Duration) -> Self {
226 self.connect_timeout = Some(timeout);
227 self
228 }
229
230 #[must_use]
232 pub fn without_connect_timeout(mut self) -> Self {
233 self.connect_timeout = None;
234 self
235 }
236
237 #[must_use]
239 pub fn event_timeout(mut self, timeout: Duration) -> Self {
240 self.event_timeout = Some(timeout);
241 self
242 }
243
244 #[must_use]
246 pub fn without_event_timeout(mut self) -> Self {
247 self.event_timeout = None;
248 self
249 }
250}
251
252impl<H> ResponsesWebSocketSessionBuilder<H>
253where
254 H: HttpClientExt
255 + Clone
256 + std::fmt::Debug
257 + Default
258 + WasmCompatSend
259 + WasmCompatSync
260 + 'static,
261{
262 pub async fn connect(self) -> Result<ResponsesWebSocketSession<H>, CompletionError> {
264 ResponsesWebSocketSession::connect_with_timeouts(
265 self.model,
266 self.connect_timeout,
267 self.event_timeout,
268 )
269 .await
270 }
271}
272
273pub struct ResponsesWebSocketSession<H = reqwest::Client> {
282 model: ResponsesCompletionModel<H>,
283 previous_response_id: Option<String>,
284 pending_done_response_id: Option<String>,
285 socket: OpenAIWebSocket,
286 in_flight: bool,
287 event_timeout: Option<Duration>,
288 closed: bool,
289 failed: bool,
290}
291
292impl<H> ResponsesWebSocketSession<H>
293where
294 H: HttpClientExt
295 + Clone
296 + std::fmt::Debug
297 + Default
298 + WasmCompatSend
299 + WasmCompatSync
300 + 'static,
301{
302 async fn connect_with_timeouts(
303 model: ResponsesCompletionModel<H>,
304 connect_timeout: Option<Duration>,
305 event_timeout: Option<Duration>,
306 ) -> Result<Self, CompletionError> {
307 let url = websocket_url(model.client.base_url())?;
308 let request = websocket_request(&url, model.client.headers())?;
309 let socket = connect_websocket(request, connect_timeout).await?;
310
311 Ok(Self {
312 model,
313 previous_response_id: None,
314 pending_done_response_id: None,
315 socket,
316 in_flight: false,
317 event_timeout,
318 closed: false,
319 failed: false,
320 })
321 }
322
323 #[must_use]
325 pub fn previous_response_id(&self) -> Option<&str> {
326 self.previous_response_id.as_deref()
327 }
328
329 pub fn clear_previous_response_id(&mut self) {
331 self.previous_response_id = None;
332 }
333
334 pub async fn send(
336 &mut self,
337 completion_request: crate::completion::CompletionRequest,
338 ) -> Result<(), CompletionError> {
339 self.send_with_options(
340 completion_request,
341 ResponsesWebSocketCreateOptions::default(),
342 )
343 .await
344 }
345
346 pub async fn send_with_options(
348 &mut self,
349 completion_request: crate::completion::CompletionRequest,
350 options: ResponsesWebSocketCreateOptions,
351 ) -> Result<(), CompletionError> {
352 self.ensure_open()?;
353
354 if self.in_flight {
355 return Err(CompletionError::ProviderError(
356 "An OpenAI websocket response is already in flight on this session".to_string(),
357 ));
358 }
359
360 completion_request.validate_message_content()?;
366
367 let payload = ResponsesWebSocketClientEvent {
368 kind: ResponsesWebSocketClientEventKind::ResponseCreate,
369 request: self.prepare_request(completion_request)?,
370 generate: options.generate,
371 };
372
373 crate::providers::internal::trace_json(
374 crate::providers::internal::LogTarget::Completions,
375 "OpenAI websocket request",
376 &payload,
377 );
378
379 let payload = serde_json::to_string(&payload)?;
380
381 if let Err(error) = self.socket.send(Message::text(payload)).await {
382 return Err(self.fail_session(websocket_provider_error(error)));
383 }
384 self.in_flight = true;
385
386 Ok(())
387 }
388
389 pub async fn next_event(&mut self) -> Result<ResponsesWebSocketEvent, CompletionError> {
391 self.ensure_open()?;
392
393 if !self.in_flight {
394 return Err(CompletionError::ProviderError(
395 "No OpenAI websocket response is currently in flight on this session".to_string(),
396 ));
397 }
398
399 loop {
400 let message = match self.read_next_message().await {
401 Ok(message) => message,
402 Err(error) => return Err(error),
403 };
404
405 let Some(message) = message else {
406 self.mark_closed();
407 return Err(CompletionError::ProviderError(
408 "The OpenAI websocket connection closed before the turn finished".to_string(),
409 ));
410 };
411
412 let message = match message {
413 Ok(message) => message,
414 Err(error) => return Err(self.fail_session(websocket_provider_error(error))),
415 };
416 let payload = match websocket_message_to_text(message) {
417 Ok(Some(payload)) => payload,
418 Ok(None) => continue,
419 Err(error) => return Err(self.fail_session(error)),
420 };
421 let event = match parse_server_event(&payload) {
422 Ok(Some(event)) => event,
423 Ok(None) => continue,
424 Err(error) => return Err(self.fail_session(error)),
425 };
426 if let ResponsesWebSocketEvent::Done(done) = &event {
427 if self.pending_done_response_id.as_deref() == done.response_id() {
430 self.pending_done_response_id = None;
431 continue;
432 }
433 }
434 self.update_state_for_event(&event);
435 return Ok(event);
436 }
437 }
438
439 pub async fn warmup(
441 &mut self,
442 completion_request: crate::completion::CompletionRequest,
443 ) -> Result<String, CompletionError> {
444 self.send_with_options(
445 completion_request,
446 ResponsesWebSocketCreateOptions::warmup(),
447 )
448 .await?;
449 let response = self.wait_for_completed_response().await?;
450 Ok(response.id)
451 }
452
453 pub async fn completion(
459 &mut self,
460 completion_request: crate::completion::CompletionRequest,
461 ) -> Result<completion::CompletionResponse, CompletionError> {
462 let provider = self.model.provider_name();
463 self.send(completion_request).await?;
464 let (response, raw_choices) = self.wait_for_terminal_response().await?;
465 match completion_response_from_raw_choices(provider, raw_choices, &response).await? {
471 Some(normalized) => Ok(normalized),
472 None => response.normalize(provider),
473 }
474 }
475
476 pub async fn raw_completion(
482 &mut self,
483 completion_request: crate::completion::CompletionRequest,
484 ) -> Result<CompletionResponse, CompletionError> {
485 self.send(completion_request).await?;
486 self.wait_for_completed_response().await
487 }
488
489 pub async fn close(&mut self) -> Result<(), CompletionError> {
494 if self.closed {
495 return Ok(());
496 }
497
498 let result = self
499 .socket
500 .close(None)
501 .await
502 .map_err(websocket_provider_error);
503 self.mark_closed();
504 result
505 }
506
507 fn prepare_request(
508 &self,
509 completion_request: crate::completion::CompletionRequest,
510 ) -> Result<super::CompletionRequest, CompletionError> {
511 let mut request = self.model.create_completion_request(completion_request)?;
512
513 request.stream = None;
516 request.additional_parameters.background = None;
517
518 if request.additional_parameters.previous_response_id.is_none() {
519 request.additional_parameters.previous_response_id = self.previous_response_id.clone();
520 }
521
522 Ok(request)
523 }
524
525 async fn wait_for_completed_response(&mut self) -> Result<CompletionResponse, CompletionError> {
526 Ok(self.wait_for_terminal_response().await?.0)
527 }
528
529 async fn wait_for_terminal_response(
555 &mut self,
556 ) -> Result<(CompletionResponse, Vec<WebSocketRawChoice>), CompletionError> {
557 let mut accumulator = RawChoiceAccumulator::new(ResponsesUsage::new());
558 let mut raw_choices = Vec::new();
559 loop {
560 match self.next_event().await? {
561 ResponsesWebSocketEvent::Response(chunk) => {
562 if matches!(
563 chunk.kind,
564 ResponseChunkKind::ResponseCompleted
565 | ResponseChunkKind::ResponseFailed
566 | ResponseChunkKind::ResponseIncomplete
567 ) {
568 return finish_terminal_response(accumulator, chunk.response, raw_choices);
569 }
570 }
571 ResponsesWebSocketEvent::Done(done) => {
572 if let Some(response) = done.as_completion_response() {
573 return finish_terminal_response(accumulator, response, raw_choices);
574 }
575
576 let message = if let Some(response_id) = done.response_id() {
577 format!(
578 "OpenAI websocket turn ended with response.done before a terminal response body was available (response_id={response_id})"
579 )
580 } else {
581 "OpenAI websocket turn ended with response.done before a terminal response body was available"
582 .to_string()
583 };
584
585 return Err(CompletionError::ProviderError(message));
586 }
587 ResponsesWebSocketEvent::Error(error) => {
588 return Err(provider_error_from_event(error));
593 }
594 ResponsesWebSocketEvent::Item(chunk) => {
595 raw_choices.extend(
596 accumulator.decode_item_chunk(chunk, ResponsesStreamOptions::strict()),
597 );
598 }
599 ResponsesWebSocketEvent::Unknown(value) => {
600 raw_choices.push(crate::streaming::RawStreamingChoice::Unknown(value));
604 }
605 }
606 }
607 }
608
609 fn update_state_for_event(&mut self, event: &ResponsesWebSocketEvent) {
610 match event {
611 ResponsesWebSocketEvent::Response(chunk) => match chunk.kind {
612 ResponseChunkKind::ResponseCompleted | ResponseChunkKind::ResponseIncomplete => {
616 let response_id = chunk.response.id.clone();
617 self.previous_response_id = Some(response_id.clone());
618 self.pending_done_response_id = Some(response_id);
619 self.in_flight = false;
620 }
621 ResponseChunkKind::ResponseFailed => {
622 self.pending_done_response_id = Some(chunk.response.id.clone());
623 self.previous_response_id = None;
624 self.in_flight = false;
625 }
626 ResponseChunkKind::ResponseCreated | ResponseChunkKind::ResponseInProgress => {}
627 },
628 ResponsesWebSocketEvent::Done(done) => {
629 match done.status() {
630 Some(ResponseStatus::Completed) | Some(ResponseStatus::Incomplete) => {
631 if let Some(response_id) = done.response_id() {
632 self.previous_response_id = Some(response_id.to_string());
633 }
634 }
635 Some(ResponseStatus::Failed)
636 | Some(ResponseStatus::Cancelled)
637 | Some(ResponseStatus::Other(_)) => {
638 self.previous_response_id = None;
639 }
640 Some(ResponseStatus::InProgress | ResponseStatus::Queued) | None => {}
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 ResponsesWebSocketEvent::Item(_) | ResponsesWebSocketEvent::Unknown(_) => {}
652 }
653 }
654
655 fn abort_turn(&mut self) {
656 self.previous_response_id = None;
657 self.pending_done_response_id = None;
658 self.in_flight = false;
659 }
660
661 fn mark_closed(&mut self) {
662 self.abort_turn();
663 self.closed = true;
664 self.failed = false;
665 }
666
667 fn mark_failed(&mut self) {
668 self.abort_turn();
669 self.failed = true;
670 }
671
672 fn ensure_open(&self) -> Result<(), CompletionError> {
673 if self.closed || self.failed {
674 return Err(CompletionError::ProviderError(
675 "The OpenAI websocket session is closed".to_string(),
676 ));
677 }
678
679 Ok(())
680 }
681
682 fn fail_session(&mut self, error: CompletionError) -> CompletionError {
683 self.mark_failed();
684 error
685 }
686
687 async fn read_next_message(
688 &mut self,
689 ) -> Result<Option<Result<Message, tungstenite::Error>>, CompletionError> {
690 if let Some(timeout_duration) = self.event_timeout {
691 match tokio::time::timeout(timeout_duration, self.socket.next()).await {
692 Ok(message) => Ok(message),
693 Err(_) => Err(self.fail_session(event_timeout_error(timeout_duration))),
694 }
695 } else {
696 Ok(self.socket.next().await)
697 }
698 }
699}
700
701impl<H> Drop for ResponsesWebSocketSession<H> {
702 fn drop(&mut self) {
703 if !self.closed {
704 tracing::warn!(
705 target: "rig::completions",
706 in_flight = self.in_flight,
707 "Dropping an OpenAI websocket session without calling close(); the connection will end without a close handshake"
708 );
709 }
710 }
711}
712
713fn finish_terminal_response(
716 mut accumulator: RawChoiceAccumulator,
717 response: CompletionResponse,
718 mut raw_choices: Vec<WebSocketRawChoice>,
719) -> Result<(CompletionResponse, Vec<WebSocketRawChoice>), CompletionError> {
720 let response = terminal_response_result(response)?;
721 let kind = if matches!(response.status, ResponseStatus::Incomplete) {
725 ResponseChunkKind::ResponseIncomplete
726 } else {
727 ResponseChunkKind::ResponseCompleted
728 };
729 accumulator.record_response_chunk(kind, response.clone(), "")?;
730 raw_choices.extend(accumulator.finish());
731 Ok((response, raw_choices))
732}
733
734fn terminal_response_result(
735 response: CompletionResponse,
736) -> Result<CompletionResponse, CompletionError> {
737 match response.status {
738 ResponseStatus::Completed => Ok(response),
739 ResponseStatus::Failed => match response.error.as_ref() {
750 Some(error) => Err(CompletionError::from_provider_body(
751 serde_json::to_string(&response).unwrap_or_else(|_| error.message.clone()),
752 )),
753 None => Err(CompletionError::ProviderError(response_error_message(
754 "failed response",
755 ))),
756 },
757 ResponseStatus::Incomplete => Ok(response),
762 other => Err(CompletionError::ProviderError(format!(
763 "OpenAI websocket response ended in state {other:?}"
764 ))),
765 }
766}
767
768fn response_error_message(fallback: &str) -> String {
769 format!("OpenAI websocket returned a {fallback}")
770}
771
772fn provider_error_from_event(error: ResponsesWebSocketErrorEvent) -> CompletionError {
779 CompletionError::from_provider_body(
780 serde_json::to_string(&error).unwrap_or_else(|_| error.to_string()),
781 )
782}
783
784fn parse_server_event(payload: &str) -> Result<Option<ResponsesWebSocketEvent>, CompletionError> {
792 #[derive(Deserialize)]
793 struct EventType {
794 #[serde(rename = "type")]
795 kind: String,
796 }
797
798 let event_type = serde_json::from_str::<EventType>(payload)?;
799 match event_type.kind.as_str() {
800 "error" => serde_json::from_str(payload)
801 .map(|e| Some(ResponsesWebSocketEvent::Error(e)))
802 .map_err(CompletionError::from),
803 "response.done" => serde_json::from_str(payload)
804 .map(|d| Some(ResponsesWebSocketEvent::Done(d)))
805 .map_err(CompletionError::from),
806 _ => Ok(Some(
810 match triage_frame(classify_responses_frame(payload))? {
811 TriagedFrame::Event(StreamingCompletionChunk::Response(response)) => {
812 ResponsesWebSocketEvent::Response(response)
813 }
814 TriagedFrame::Event(StreamingCompletionChunk::Delta(item)) => {
815 ResponsesWebSocketEvent::Item(item)
816 }
817 TriagedFrame::Unknown(value) => ResponsesWebSocketEvent::Unknown(value),
818 },
819 )),
820 }
821}
822
823fn websocket_message_to_text(message: Message) -> Result<Option<String>, CompletionError> {
824 match message {
825 Message::Text(text) => Ok(Some(text.to_string())),
826 Message::Binary(bytes) => String::from_utf8(bytes.to_vec())
827 .map(Some)
828 .map_err(|error| CompletionError::ResponseError(error.to_string())),
829 Message::Ping(_) | Message::Pong(_) | Message::Frame(_) => Ok(None),
830 Message::Close(frame) => {
831 let reason = frame
832 .map(|frame| frame.reason.to_string())
833 .filter(|reason| !reason.is_empty())
834 .unwrap_or_else(|| "without a close reason".to_string());
835 Err(CompletionError::ProviderError(format!(
836 "The OpenAI websocket connection closed {reason}"
837 )))
838 }
839 }
840}
841
842fn websocket_url(base_url: &str) -> Result<String, CompletionError> {
843 let mut url = Url::parse(base_url)?;
844 match url.scheme() {
845 "https" => {
846 url.set_scheme("wss").map_err(|_| {
847 CompletionError::ProviderError("Failed to convert https URL to wss".to_string())
848 })?;
849 }
850 "http" => {
851 url.set_scheme("ws").map_err(|_| {
852 CompletionError::ProviderError("Failed to convert http URL to ws".to_string())
853 })?;
854 }
855 scheme => {
856 return Err(CompletionError::ProviderError(format!(
857 "Unsupported base URL scheme for OpenAI websocket mode: {scheme}"
858 )));
859 }
860 }
861
862 let path = format!("{}/responses", url.path().trim_end_matches('/'));
863 url.set_path(&path);
864 Ok(url.to_string())
865}
866
867fn websocket_request(
868 url: &str,
869 headers: &http::HeaderMap,
870) -> Result<http::Request<()>, CompletionError> {
871 let mut request = url.into_client_request().map_err(|error| {
872 CompletionError::ProviderError(format!("Failed to build OpenAI websocket request: {error}"))
873 })?;
874
875 for (name, value) in headers {
876 request.headers_mut().insert(name, value.clone());
877 }
878
879 Ok(request)
880}
881
882async fn connect_websocket(
883 request: http::Request<()>,
884 connect_timeout: Option<Duration>,
885) -> Result<OpenAIWebSocket, CompletionError> {
886 if let Some(timeout_duration) = connect_timeout {
887 match tokio::time::timeout(timeout_duration, connect_async(request)).await {
888 Ok(result) => result
889 .map(|(socket, _)| socket)
890 .map_err(websocket_provider_error),
891 Err(_) => Err(connect_timeout_error(timeout_duration)),
892 }
893 } else {
894 connect_async(request)
895 .await
896 .map(|(socket, _)| socket)
897 .map_err(websocket_provider_error)
898 }
899}
900
901fn connect_timeout_error(timeout: Duration) -> CompletionError {
902 CompletionError::ProviderError(format!(
903 "Timed out connecting to the OpenAI websocket after {timeout:?}"
904 ))
905}
906
907fn event_timeout_error(timeout: Duration) -> CompletionError {
908 CompletionError::ProviderError(format!(
909 "Timed out waiting for the next OpenAI websocket event after {timeout:?}"
910 ))
911}
912
913fn websocket_provider_error(error: tungstenite::Error) -> CompletionError {
943 let tungstenite::Error::Http(response) = error else {
944 return CompletionError::ProviderError(error.to_string());
945 };
946
947 let (parts, body) = (*response).into_parts();
948 let provider_request_id = REQUEST_ID_HEADER
949 .and_then(|header| parts.headers.get(header))
950 .and_then(|value| value.to_str().ok())
951 .filter(|value| !value.is_empty())
952 .map(str::to_string);
953 let body = body
957 .map(|body| String::from_utf8_lossy(&body).into_owned())
958 .unwrap_or_default();
959
960 CompletionError::from_http_response_with_request_id(parts.status, body, provider_request_id)
961 .with_response_headers(Some(Box::new(parts.headers)))
963}
964
965#[cfg(test)]
966mod tests {
967 use super::{CompletionError, tungstenite};
968 use super::{
969 ResponsesWebSocketCreateOptions, ResponsesWebSocketDoneEvent, ResponsesWebSocketEvent,
970 parse_server_event, terminal_response_result, websocket_provider_error, websocket_url,
971 };
972 use crate::client::CompletionClient;
973 use crate::completion::CompletionModel;
974 use crate::providers::openai::responses_api::{
975 CompletionResponse, IncompleteDetailsReason, Output, ResponseError, ResponseObject,
976 ResponseStatus, ResponsesUsage,
977 };
978 use futures::{SinkExt, StreamExt};
979 use serde_json::json;
980
981 fn handshake_rejection(
984 status: u16,
985 request_id: Option<&str>,
986 body: Option<&str>,
987 ) -> tungstenite::Error {
988 handshake_rejection_with_headers(status, request_id, body, &[])
989 }
990
991 fn handshake_rejection_with_headers(
994 status: u16,
995 request_id: Option<&str>,
996 body: Option<&str>,
997 headers: &[(&str, &str)],
998 ) -> tungstenite::Error {
999 let mut response = http::Response::builder().status(status);
1000 if let Some(request_id) = request_id {
1001 response = response.header("x-request-id", request_id);
1002 }
1003 for (name, value) in headers {
1004 response = response.header(*name, *value);
1005 }
1006 tungstenite::Error::Http(Box::new(
1007 response
1008 .body(body.map(|body| body.as_bytes().to_vec()))
1009 .expect("response should build"),
1010 ))
1011 }
1012
1013 const REJECTION_BODY: &str = r#"{"error":{"message":"Incorrect API key provided: sk-inval***-key.","type":"invalid_request_error","code":"invalid_api_key","param":null},"status":401}"#;
1016
1017 #[test]
1018 fn websocket_provider_error_preserves_status_body_and_request_id() {
1019 let error = websocket_provider_error(handshake_rejection(
1020 401,
1021 Some("req_websocket_1"),
1022 Some(REJECTION_BODY),
1023 ));
1024
1025 assert!(matches!(error, CompletionError::ProviderResponse(_)));
1026 assert_eq!(
1027 error.provider_response_status(),
1028 Some(http::StatusCode::UNAUTHORIZED)
1029 );
1030 assert_eq!(error.provider_response_body(), Some(REJECTION_BODY));
1031 assert_eq!(error.provider_request_id(), Some("req_websocket_1"));
1032 assert_eq!(
1033 error
1034 .provider_response_json()
1035 .expect("body should be valid JSON")
1036 .expect("parsed JSON should be present")["error"]["code"],
1037 "invalid_api_key"
1038 );
1039 }
1040
1041 #[test]
1044 fn websocket_provider_error_without_a_request_id_keeps_the_rest() {
1045 let error = websocket_provider_error(handshake_rejection(429, None, Some("slow down")));
1046
1047 assert_eq!(
1048 error.provider_response_status(),
1049 Some(http::StatusCode::TOO_MANY_REQUESTS)
1050 );
1051 assert_eq!(error.provider_response_body(), Some("slow down"));
1052 assert_eq!(error.provider_request_id(), None);
1053 }
1054
1055 #[test]
1058 fn websocket_provider_error_treats_an_empty_request_id_as_absent() {
1059 let error = websocket_provider_error(handshake_rejection(401, Some(""), Some("nope")));
1060
1061 assert_eq!(error.provider_request_id(), None);
1062 assert_eq!(
1063 error.provider_response_status(),
1064 Some(http::StatusCode::UNAUTHORIZED)
1065 );
1066 }
1067
1068 #[test]
1071 fn websocket_provider_error_without_a_body_keeps_the_status() {
1072 let error = websocket_provider_error(handshake_rejection(403, Some("req_2"), None));
1073
1074 assert_eq!(
1075 error.provider_response_status(),
1076 Some(http::StatusCode::FORBIDDEN)
1077 );
1078 assert_eq!(error.provider_response_body(), Some(""));
1079 assert_eq!(error.provider_request_id(), Some("req_2"));
1080 }
1081
1082 #[test]
1087 fn websocket_provider_error_preserves_every_rejection_status() {
1088 for status in [200u16, 302, 400, 401, 403, 404, 429, 500, 503] {
1089 let error = websocket_provider_error(handshake_rejection(status, None, Some("body")));
1090 assert_eq!(
1091 error.provider_response_status().map(|s| s.as_u16()),
1092 Some(status),
1093 "status {status} must survive"
1094 );
1095 }
1096 }
1097
1098 #[test]
1110 fn websocket_provider_error_leaves_every_transport_failure_alone() {
1111 let cases: Vec<tungstenite::Error> = vec![
1112 tungstenite::Error::ConnectionClosed,
1113 tungstenite::Error::AlreadyClosed,
1114 tungstenite::Error::Io(std::io::Error::other("connection reset")),
1115 tungstenite::Error::Capacity(tungstenite::error::CapacityError::TooManyHeaders),
1116 tungstenite::Error::Protocol(tungstenite::error::ProtocolError::HandshakeIncomplete),
1117 tungstenite::Error::WriteBufferFull(Box::new(tungstenite::Message::Text(
1118 "queued".into(),
1119 ))),
1120 tungstenite::Error::AttackAttempt,
1121 tungstenite::Error::Url(tungstenite::error::UrlError::NoPathOrQuery),
1122 tungstenite::Error::HttpFormat(
1123 http::header::HeaderName::from_bytes(b"not a header")
1124 .expect_err("an invalid header name should not parse")
1125 .into(),
1126 ),
1127 ];
1128
1129 for error in cases {
1130 let expected = error.to_string();
1131 let mapped = websocket_provider_error(error);
1132
1133 assert!(
1134 matches!(mapped, CompletionError::ProviderError(_)),
1135 "a failure with no provider response must stay a ProviderError: {mapped:?}"
1136 );
1137 assert_eq!(mapped.to_string(), format!("ProviderError: {expected}"));
1138 assert_eq!(mapped.provider_response_status(), None);
1139 assert_eq!(mapped.provider_response_body(), None);
1140 assert_eq!(mapped.provider_request_id(), None);
1141 }
1142 }
1143
1144 #[test]
1149 fn websocket_provider_error_preserves_the_rejections_headers() {
1150 let error = websocket_provider_error(handshake_rejection_with_headers(
1151 429,
1152 Some("req_rate_limited"),
1153 Some(r#"{"error":{"code":"rate_limit_exceeded"}}"#),
1154 &[("retry-after", "20"), ("x-ratelimit-remaining", "0")],
1155 ));
1156
1157 let headers = error
1158 .provider_response_headers()
1159 .expect("a rejection's headers must survive onto the error");
1160 assert_eq!(
1161 headers
1162 .get(http::header::RETRY_AFTER)
1163 .and_then(|value| value.to_str().ok()),
1164 Some("20")
1165 );
1166 assert_eq!(
1167 headers
1168 .get("x-ratelimit-remaining")
1169 .and_then(|value| value.to_str().ok()),
1170 Some("0")
1171 );
1172 assert_eq!(
1174 error.provider_response_status(),
1175 Some(http::StatusCode::TOO_MANY_REQUESTS)
1176 );
1177 assert_eq!(error.provider_request_id(), Some("req_rate_limited"));
1178 }
1179
1180 #[test]
1183 fn websocket_provider_error_reports_no_headers_for_a_transport_failure() {
1184 let error = websocket_provider_error(tungstenite::Error::ConnectionClosed);
1185
1186 assert!(error.provider_response_headers().is_none());
1187 }
1188
1189 #[test]
1193 fn websocket_provider_error_no_longer_flattens_a_rejection_to_a_string() {
1194 let error = websocket_provider_error(handshake_rejection(
1195 401,
1196 Some("req_websocket_1"),
1197 Some(REJECTION_BODY),
1198 ));
1199
1200 assert_ne!(
1201 error.to_string(),
1202 "ProviderError: HTTP error: 401 Unauthorized",
1203 "the pre-fix behavior discarded the status, body and request id"
1204 );
1205 }
1206 use std::time::Duration;
1207 use tokio::net::TcpListener;
1208 use tokio::time::sleep;
1209 use tokio_tungstenite::{accept_async, tungstenite::Message};
1210
1211 #[test]
1212 fn websocket_error_event_preserves_provider_payload_as_json() {
1213 let mut extra = serde_json::Map::new();
1214 extra.insert(
1215 "type".to_string(),
1216 serde_json::Value::String("invalid_request_error".to_string()),
1217 );
1218 let event = super::ResponsesWebSocketErrorEvent {
1219 kind: super::ResponsesWebSocketErrorEventKind::Error,
1220 error: super::ResponsesWebSocketErrorPayload {
1221 code: Some("rate_limit_exceeded".to_string()),
1222 message: Some("slow down".to_string()),
1223 extra,
1224 },
1225 };
1226
1227 let err = super::provider_error_from_event(event);
1228
1229 assert_eq!(err.provider_response_status(), None);
1232 let json = err
1233 .provider_response_json()
1234 .expect("preserved body should be valid JSON")
1235 .expect("provider response body should be present");
1236 assert_eq!(json["error"]["code"], "rate_limit_exceeded");
1237 assert_eq!(json["error"]["message"], "slow down");
1238 assert_eq!(json["error"]["type"], "invalid_request_error");
1239 }
1240
1241 fn sample_response(status: ResponseStatus) -> CompletionResponse {
1242 CompletionResponse {
1243 id: "resp_123".to_string(),
1244 object: ResponseObject::Response,
1245 provider_request_id: None,
1246 created_at: 0,
1247 status,
1248 error: None,
1249 incomplete_details: None,
1250 instructions: None,
1251 max_output_tokens: None,
1252 model: "gpt-5.4".to_string(),
1253 usage: Some(ResponsesUsage {
1254 input_tokens: 1,
1255 input_tokens_details: None,
1256 output_tokens: 2,
1257 output_tokens_details: Some(
1258 crate::providers::openai::responses_api::OutputTokensDetails {
1259 reasoning_tokens: 0,
1260 },
1261 ),
1262 total_tokens: 3,
1263 }),
1264 output: Vec::new(),
1265 tools: Vec::new(),
1266 additional_parameters: Default::default(),
1267 provider_reasoning: None,
1268 reasoning_metadata: None,
1269 reasoning_context: None,
1270 }
1271 }
1272
1273 #[test]
1274 fn warmup_options_serialize_generate_false() {
1275 let options = ResponsesWebSocketCreateOptions::warmup();
1276 let json = serde_json::to_value(options).expect("options should serialize");
1277
1278 assert_eq!(json, json!({ "generate": false }));
1279 }
1280
1281 #[test]
1282 fn websocket_url_converts_https_to_wss() {
1283 let url = websocket_url("https://api.openai.com/v1").expect("url should convert");
1284 assert_eq!(url, "wss://api.openai.com/v1/responses");
1285 }
1286
1287 #[test]
1288 fn parse_done_event_exposes_response_id() {
1289 let payload = json!({
1290 "type": "response.done",
1291 "response": {
1292 "id": "resp_done_1",
1293 "status": "completed"
1294 }
1295 });
1296
1297 let event = parse_server_event(&payload.to_string())
1298 .expect("done event should deserialize")
1299 .expect("done event should not be skipped");
1300
1301 assert!(matches!(
1302 event,
1303 ResponsesWebSocketEvent::Done(ResponsesWebSocketDoneEvent { .. })
1304 ));
1305 assert_eq!(event.response_id(), Some("resp_done_1"));
1306 assert!(event.is_terminal());
1307 }
1308
1309 #[test]
1310 fn parse_response_completed_event_is_terminal() {
1311 let payload = json!({
1312 "type": "response.completed",
1313 "sequence_number": 12,
1314 "response": {
1315 "id": "resp_completed_1",
1316 "object": "response",
1317 "created_at": 0,
1318 "status": "completed",
1319 "error": null,
1320 "incomplete_details": null,
1321 "instructions": null,
1322 "max_output_tokens": null,
1323 "model": "gpt-5.4",
1324 "usage": null,
1325 "output": [],
1326 "tools": []
1327 }
1328 });
1329
1330 let event = parse_server_event(&payload.to_string())
1331 .expect("response event should deserialize")
1332 .expect("response event should not be skipped");
1333
1334 assert!(matches!(event, ResponsesWebSocketEvent::Response(_)));
1335 assert!(event.is_terminal());
1336 assert_eq!(event.response_id(), Some("resp_completed_1"));
1337 }
1338
1339 #[test]
1340 fn parse_live_output_item_added_event() {
1341 let payload = json!({
1342 "type": "response.output_item.added",
1343 "item": {
1344 "id": "msg_036471c3a72c147b0069ae7848d68881959773fd2d99e3d98a",
1345 "type": "message",
1346 "status": "in_progress",
1347 "content": [],
1348 "role": "assistant"
1349 },
1350 "output_index": 0,
1351 "sequence_number": 2
1352 });
1353
1354 let event = parse_server_event(&payload.to_string())
1355 .expect("output item event should parse")
1356 .expect("output item event should not be skipped");
1357
1358 assert!(matches!(event, ResponsesWebSocketEvent::Item(_)));
1359 }
1360
1361 #[test]
1362 fn parse_live_content_part_added_event() {
1363 let payload = json!({
1364 "type": "response.content_part.added",
1365 "content_index": 0,
1366 "item_id": "msg_036471c3a72c147b0069ae7848d68881959773fd2d99e3d98a",
1367 "output_index": 0,
1368 "part": {
1369 "type": "output_text",
1370 "annotations": [],
1371 "logprobs": [],
1372 "text": ""
1373 },
1374 "sequence_number": 3
1375 });
1376
1377 let event = parse_server_event(&payload.to_string())
1378 .expect("content part event should parse")
1379 .expect("content part event should not be skipped");
1380
1381 assert!(matches!(event, ResponsesWebSocketEvent::Item(_)));
1382 }
1383
1384 #[test]
1385 fn parse_live_output_text_delta_event() {
1386 let payload = json!({
1387 "type": "response.output_text.delta",
1388 "content_index": 0,
1389 "delta": "Web",
1390 "item_id": "msg_023af0f0a91bc2a90069ae788612e881958345bb156915ba29",
1391 "logprobs": [],
1392 "obfuscation": "2YYErYq7jkqqM",
1393 "output_index": 0,
1394 "sequence_number": 4
1395 });
1396
1397 let event = parse_server_event(&payload.to_string())
1398 .expect("output text delta event should parse")
1399 .expect("output text delta event should not be skipped");
1400
1401 assert!(matches!(event, ResponsesWebSocketEvent::Item(_)));
1402 }
1403
1404 #[test]
1405 fn terminal_response_requires_completed_status() {
1406 let completed = terminal_response_result(sample_response(ResponseStatus::Completed))
1407 .expect("completed response should succeed");
1408 assert_eq!(completed.id, "resp_123");
1409
1410 let failed = terminal_response_result(sample_response(ResponseStatus::Failed))
1411 .expect_err("failed response should error");
1412 assert!(failed.to_string().contains("failed response"));
1413 }
1414
1415 #[tokio::test]
1416 async fn incomplete_turn_keeps_streamed_partial_output() {
1417 let listener = TcpListener::bind("127.0.0.1:0")
1418 .await
1419 .expect("listener should bind");
1420 let address = listener.local_addr().expect("listener should have address");
1421
1422 let server = tokio::spawn(async move {
1423 let (stream, _) = listener.accept().await.expect("server should accept");
1424 let mut socket = accept_async(stream)
1425 .await
1426 .expect("server should upgrade websocket");
1427
1428 let request = socket
1429 .next()
1430 .await
1431 .expect("request should exist")
1432 .expect("request should be valid");
1433 let payload = request.into_text().expect("request should be text");
1434 assert!(
1435 payload.contains("\"type\":\"response.create\""),
1436 "expected response.create payload, got {payload}"
1437 );
1438
1439 socket
1443 .send(Message::text(
1444 json!({
1445 "type": "response.output_text.delta",
1446 "content_index": 0,
1447 "delta": "partial",
1448 "item_id": "msg_incomplete_1",
1449 "logprobs": [],
1450 "output_index": 0,
1451 "sequence_number": 1
1452 })
1453 .to_string(),
1454 ))
1455 .await
1456 .expect("delta event should send");
1457
1458 let mut response = sample_response(ResponseStatus::Incomplete);
1459 response.incomplete_details = Some(IncompleteDetailsReason {
1460 reason: "max_output_tokens".to_string(),
1461 });
1462 let response = serde_json::to_value(response).expect("response should serialize");
1463
1464 socket
1465 .send(Message::text(
1466 json!({
1467 "type": "response.incomplete",
1468 "sequence_number": 2,
1469 "response": response,
1470 })
1471 .to_string(),
1472 ))
1473 .await
1474 .expect("incomplete event should send");
1475 });
1476
1477 let base_url = format!("http://{address}/v1");
1478 let client = crate::providers::openai::Client::builder()
1479 .api_key("test-key")
1480 .base_url(&base_url)
1481 .build()
1482 .expect("client should build");
1483 let model = client.completion_model("gpt-4o");
1484 let mut session = client
1485 .responses_websocket("gpt-4o")
1486 .await
1487 .expect("session should connect");
1488
1489 let normalized = session
1490 .completion(model.completion_request("hello").build())
1491 .await
1492 .expect("incomplete turn should be a successful terminal");
1493
1494 assert_eq!(
1497 normalized.finish_reason(),
1498 Some(crate::completion::FinishReason::Length)
1499 );
1500 assert_eq!(normalized.usage.input_tokens, 1);
1501 assert_eq!(normalized.usage.output_tokens, 2);
1502 assert_eq!(normalized.usage.total_tokens, 3);
1503 assert!(matches!(
1504 normalized.choice.first(),
1505 Some(crate::completion::AssistantContent::Text(text)) if text.text == "partial"
1506 ));
1507
1508 server.await.expect("server task should finish");
1509 }
1510
1511 #[tokio::test]
1515 async fn same_item_text_resumes_as_one_part_across_interleaved_reasoning() {
1516 let listener = TcpListener::bind("127.0.0.1:0")
1517 .await
1518 .expect("listener should bind");
1519 let address = listener.local_addr().expect("listener should have address");
1520
1521 let server = tokio::spawn(async move {
1522 let (stream, _) = listener.accept().await.expect("server should accept");
1523 let mut socket = accept_async(stream)
1524 .await
1525 .expect("server should upgrade websocket");
1526
1527 socket
1528 .next()
1529 .await
1530 .expect("request should exist")
1531 .expect("request should be valid");
1532
1533 let events = [
1534 json!({
1535 "type": "response.output_text.delta",
1536 "content_index": 0,
1537 "delta": "hello ",
1538 "item_id": "msg_1",
1539 "logprobs": [],
1540 "output_index": 0,
1541 "sequence_number": 1
1542 }),
1543 json!({
1544 "type": "response.reasoning_summary_text.delta",
1545 "delta": "because",
1546 "item_id": "rs_2",
1547 "output_index": 1,
1548 "summary_index": 0,
1549 "sequence_number": 2
1550 }),
1551 json!({
1552 "type": "response.output_text.delta",
1553 "content_index": 0,
1554 "delta": "world",
1555 "item_id": "msg_1",
1556 "logprobs": [],
1557 "output_index": 0,
1558 "sequence_number": 3
1559 }),
1560 json!({
1561 "type": "response.completed",
1562 "sequence_number": 4,
1563 "response": serde_json::to_value(sample_response(ResponseStatus::Completed))
1564 .expect("response should serialize"),
1565 }),
1566 ];
1567 for event in events {
1568 socket
1569 .send(Message::text(event.to_string()))
1570 .await
1571 .expect("event should send");
1572 }
1573 });
1574
1575 let base_url = format!("http://{address}/v1");
1576 let client = crate::providers::openai::Client::builder()
1577 .api_key("test-key")
1578 .base_url(&base_url)
1579 .build()
1580 .expect("client should build");
1581 let model = client.completion_model("gpt-4o");
1582 let mut session = client
1583 .responses_websocket("gpt-4o")
1584 .await
1585 .expect("session should connect");
1586
1587 let normalized = session
1588 .completion(model.completion_request("hello").build())
1589 .await
1590 .expect("interleaved turn should normalize");
1591
1592 let texts: Vec<_> = normalized
1593 .choice
1594 .iter()
1595 .filter_map(|content| match content {
1596 crate::completion::AssistantContent::Text(text) => Some(text.text.clone()),
1597 _ => None,
1598 })
1599 .collect();
1600 assert_eq!(
1601 texts,
1602 ["hello world"],
1603 "same-item text must aggregate as one part around the reasoning"
1604 );
1605 assert!(
1606 normalized.choice.iter().any(|content| matches!(
1607 content,
1608 crate::completion::AssistantContent::Reasoning(_)
1609 )),
1610 "the interleaved reasoning must survive"
1611 );
1612
1613 server.await.expect("server task should finish");
1614 }
1615
1616 #[tokio::test]
1617 async fn completed_turn_without_deltas_falls_back_to_terminal_body() {
1618 let listener = TcpListener::bind("127.0.0.1:0")
1619 .await
1620 .expect("listener should bind");
1621 let address = listener.local_addr().expect("listener should have address");
1622
1623 let server = tokio::spawn(async move {
1624 let (stream, _) = listener.accept().await.expect("server should accept");
1625 let mut socket = accept_async(stream)
1626 .await
1627 .expect("server should upgrade websocket");
1628
1629 let request = socket
1630 .next()
1631 .await
1632 .expect("request should exist")
1633 .expect("request should be valid");
1634 let payload = request.into_text().expect("request should be text");
1635 assert!(
1636 payload.contains("\"type\":\"response.create\""),
1637 "expected response.create payload, got {payload}"
1638 );
1639
1640 let mut response = sample_response(ResponseStatus::Completed);
1643 response.output = vec![
1644 serde_json::from_value::<Output>(json!({
1645 "type": "message",
1646 "id": "msg_terminal_1",
1647 "status": "completed",
1648 "role": "assistant",
1649 "content": [{ "type": "output_text", "annotations": [], "text": "hello there" }]
1650 }))
1651 .expect("output message should deserialize"),
1652 ];
1653 let response = serde_json::to_value(response).expect("response should serialize");
1654
1655 socket
1656 .send(Message::text(
1657 json!({
1658 "type": "response.completed",
1659 "sequence_number": 1,
1660 "response": response,
1661 })
1662 .to_string(),
1663 ))
1664 .await
1665 .expect("completed event should send");
1666 });
1667
1668 let base_url = format!("http://{address}/v1");
1669 let client = crate::providers::openai::Client::builder()
1670 .api_key("test-key")
1671 .base_url(&base_url)
1672 .build()
1673 .expect("client should build");
1674 let model = client.completion_model("gpt-4o");
1675 let mut session = client
1676 .responses_websocket("gpt-4o")
1677 .await
1678 .expect("session should connect");
1679
1680 let normalized = session
1681 .completion(model.completion_request("hello").build())
1682 .await
1683 .expect("completed turn should normalize");
1684
1685 assert!(matches!(
1686 normalized.choice.first(),
1687 Some(crate::completion::AssistantContent::Text(text)) if text.text == "hello there"
1688 ));
1689 assert_eq!(normalized.message_id.as_deref(), Some("msg_terminal_1"));
1690
1691 server.await.expect("server task should finish");
1692 }
1693
1694 #[tokio::test]
1695 async fn incomplete_turn_without_deltas_normalizes_terminal_body_output() {
1696 let listener = TcpListener::bind("127.0.0.1:0")
1697 .await
1698 .expect("listener should bind");
1699 let address = listener.local_addr().expect("listener should have address");
1700
1701 let server = tokio::spawn(async move {
1702 let (stream, _) = listener.accept().await.expect("server should accept");
1703 let mut socket = accept_async(stream)
1704 .await
1705 .expect("server should upgrade websocket");
1706
1707 let request = socket
1708 .next()
1709 .await
1710 .expect("request should exist")
1711 .expect("request should be valid");
1712 let payload = request.into_text().expect("request should be text");
1713 assert!(
1714 payload.contains("\"type\":\"response.create\""),
1715 "expected response.create payload, got {payload}"
1716 );
1717
1718 let mut response = sample_response(ResponseStatus::Incomplete);
1722 response.incomplete_details = Some(IncompleteDetailsReason {
1723 reason: "max_output_tokens".to_string(),
1724 });
1725 response.output = vec![
1726 serde_json::from_value::<Output>(json!({
1727 "type": "message",
1728 "id": "msg_body_only_1",
1729 "status": "incomplete",
1730 "role": "assistant",
1731 "content": [{ "type": "output_text", "annotations": [], "text": "partial from body" }]
1732 }))
1733 .expect("output message should deserialize"),
1734 ];
1735 let response = serde_json::to_value(response).expect("response should serialize");
1736
1737 socket
1738 .send(Message::text(
1739 json!({
1740 "type": "response.incomplete",
1741 "sequence_number": 1,
1742 "response": response,
1743 })
1744 .to_string(),
1745 ))
1746 .await
1747 .expect("incomplete event should send");
1748 });
1749
1750 let base_url = format!("http://{address}/v1");
1751 let client = crate::providers::openai::Client::builder()
1752 .api_key("test-key")
1753 .base_url(&base_url)
1754 .build()
1755 .expect("client should build");
1756 let model = client.completion_model("gpt-4o");
1757 let mut session = client
1758 .responses_websocket("gpt-4o")
1759 .await
1760 .expect("session should connect");
1761
1762 let normalized = session
1763 .completion(model.completion_request("hello").build())
1764 .await
1765 .expect("incomplete turn with body output should normalize");
1766
1767 assert!(matches!(
1768 normalized.choice.first(),
1769 Some(crate::completion::AssistantContent::Text(text)) if text.text == "partial from body"
1770 ));
1771 assert_eq!(
1772 normalized.finish_reason(),
1773 Some(crate::completion::FinishReason::Length)
1774 );
1775 assert_eq!(normalized.message_id.as_deref(), Some("msg_body_only_1"));
1776
1777 server.await.expect("server task should finish");
1778 }
1779
1780 #[test]
1781 fn terminal_failed_response_with_error_preserves_raw_payload() {
1782 let mut response = sample_response(ResponseStatus::Failed);
1783 response.error = Some(ResponseError {
1784 code: "server_error".to_string(),
1785 message: "the model failed to generate a response".to_string(),
1786 });
1787
1788 let err = match terminal_response_result(response) {
1789 Ok(_) => panic!("failed response with an error object should fail"),
1790 Err(e) => e,
1791 };
1792
1793 assert_eq!(err.provider_response_status(), None);
1798
1799 let json = err
1800 .provider_response_json()
1801 .expect("preserved body should parse as JSON")
1802 .expect("preserved body should not be empty");
1803 assert_eq!(
1804 json["error"]["message"],
1805 "the model failed to generate a response"
1806 );
1807 assert_eq!(json["error"]["code"], "server_error");
1808 }
1809
1810 #[test]
1811 fn terminal_failed_response_without_error_is_rig_diagnostic() {
1812 let err = match terminal_response_result(sample_response(ResponseStatus::Failed)) {
1813 Ok(_) => panic!("failed response should fail"),
1814 Err(e) => e,
1815 };
1816
1817 assert_eq!(err.provider_response_body(), None);
1820 assert!(err.to_string().contains("failed response"));
1821 }
1822
1823 #[tokio::test]
1824 async fn malformed_known_event_rejects_reuse_and_allows_close() {
1825 let listener = TcpListener::bind("127.0.0.1:0")
1826 .await
1827 .expect("listener should bind");
1828 let address = listener.local_addr().expect("listener should have address");
1829
1830 let server = tokio::spawn(async move {
1831 let (stream, _) = listener.accept().await.expect("server should accept");
1832 let mut socket = accept_async(stream)
1833 .await
1834 .expect("server should upgrade websocket");
1835
1836 let request = socket
1837 .next()
1838 .await
1839 .expect("request should exist")
1840 .expect("request should be valid");
1841 let payload = request.into_text().expect("request should be text");
1842 assert!(
1843 payload.contains("\"type\":\"response.create\""),
1844 "expected response.create payload, got {payload}"
1845 );
1846
1847 socket
1848 .send(Message::text(
1849 json!({
1850 "type": "response.completed"
1851 })
1852 .to_string(),
1853 ))
1854 .await
1855 .expect("malformed known event should send");
1856
1857 let message = socket
1858 .next()
1859 .await
1860 .expect("close frame should arrive")
1861 .expect("close frame should be valid");
1862 assert!(
1863 matches!(message, Message::Close(_)),
1864 "expected close frame, got {message:?}"
1865 );
1866 });
1867
1868 let base_url = format!("http://{address}/v1");
1869 let client = crate::providers::openai::Client::builder()
1870 .api_key("test-key")
1871 .base_url(&base_url)
1872 .build()
1873 .expect("client should build");
1874 let model = client.completion_model("gpt-4o");
1875 let mut session = client
1876 .responses_websocket("gpt-4o")
1877 .await
1878 .expect("session should connect");
1879
1880 session
1881 .send(model.completion_request("hello").build())
1882 .await
1883 .expect("request should send");
1884
1885 let error = session
1886 .next_event()
1887 .await
1888 .expect_err("malformed known event should fail");
1889 assert!(
1890 error.to_string().contains("StreamingCompletionChunk"),
1891 "expected strict decode failure, got {error}"
1892 );
1893
1894 let closed = session
1895 .send(model.completion_request("retry").build())
1896 .await
1897 .expect_err("session should close after fatal parse error");
1898 assert!(
1899 closed.to_string().contains("session is closed"),
1900 "expected closed-session error, got {closed}"
1901 );
1902
1903 session
1904 .close()
1905 .await
1906 .expect("explicit close after fatal parse error should succeed");
1907
1908 server.await.expect("server task should finish");
1909 }
1910
1911 #[tokio::test]
1912 async fn event_timeout_rejects_reuse_and_allows_close() {
1913 let listener = TcpListener::bind("127.0.0.1:0")
1914 .await
1915 .expect("listener should bind");
1916 let address = listener.local_addr().expect("listener should have address");
1917
1918 let server = tokio::spawn(async move {
1919 let (stream, _) = listener.accept().await.expect("server should accept");
1920 let mut socket = accept_async(stream)
1921 .await
1922 .expect("server should upgrade websocket");
1923
1924 let request = socket
1925 .next()
1926 .await
1927 .expect("request should exist")
1928 .expect("request should be valid");
1929 let payload = request.into_text().expect("request should be text");
1930 assert!(
1931 payload.contains("\"type\":\"response.create\""),
1932 "expected response.create payload, got {payload}"
1933 );
1934
1935 sleep(Duration::from_millis(60)).await;
1936 let message = socket
1937 .next()
1938 .await
1939 .expect("close frame should arrive")
1940 .expect("close frame should be valid");
1941 assert!(
1942 matches!(message, Message::Close(_)),
1943 "expected close frame, got {message:?}"
1944 );
1945 });
1946
1947 let base_url = format!("http://{address}/v1");
1948 let client = crate::providers::openai::Client::builder()
1949 .api_key("test-key")
1950 .base_url(&base_url)
1951 .build()
1952 .expect("client should build");
1953 let model = client.completion_model("gpt-4o");
1954 let mut session = client
1955 .responses_websocket_builder("gpt-4o")
1956 .event_timeout(Duration::from_millis(20))
1957 .connect()
1958 .await
1959 .expect("session should connect");
1960
1961 session
1962 .send(model.completion_request("hello").build())
1963 .await
1964 .expect("request should send");
1965
1966 let error = session
1967 .next_event()
1968 .await
1969 .expect_err("next_event should time out");
1970 assert!(
1971 error
1972 .to_string()
1973 .contains("Timed out waiting for the next OpenAI websocket event"),
1974 "expected timeout error, got {error}"
1975 );
1976
1977 let closed = session
1978 .send(model.completion_request("retry").build())
1979 .await
1980 .expect_err("timed-out session should close");
1981 assert!(
1982 closed.to_string().contains("session is closed"),
1983 "expected closed-session error, got {closed}"
1984 );
1985
1986 session
1987 .close()
1988 .await
1989 .expect("explicit close after timeout should succeed");
1990
1991 server.await.expect("server task should finish");
1992 }
1993
1994 #[tokio::test]
1995 async fn late_response_done_is_ignored_on_next_turn() {
1996 let listener = TcpListener::bind("127.0.0.1:0")
1997 .await
1998 .expect("listener should bind");
1999 let address = listener.local_addr().expect("listener should have address");
2000
2001 let server = tokio::spawn(async move {
2002 let (stream, _) = listener.accept().await.expect("server should accept");
2003 let mut socket = accept_async(stream)
2004 .await
2005 .expect("server should upgrade websocket");
2006
2007 for (index, response_id) in ["resp_1", "resp_2"].iter().enumerate() {
2008 let request = socket
2009 .next()
2010 .await
2011 .expect("request should exist")
2012 .expect("request should be valid");
2013 let payload = request.into_text().expect("request should be text");
2014 assert!(
2015 payload.contains("\"type\":\"response.create\""),
2016 "expected response.create payload, got {payload}"
2017 );
2018
2019 let response = sample_response(ResponseStatus::Completed);
2020 let response = serde_json::to_value(CompletionResponse {
2021 id: (*response_id).to_string(),
2022 ..response
2023 })
2024 .expect("response should serialize");
2025
2026 socket
2027 .send(Message::text(
2028 json!({
2029 "type": "response.completed",
2030 "sequence_number": (index * 2) + 1,
2031 "response": response,
2032 })
2033 .to_string(),
2034 ))
2035 .await
2036 .expect("completed event should send");
2037 socket
2038 .send(Message::text(
2039 json!({
2040 "type": "response.done",
2041 "response": {
2042 "id": response_id,
2043 "status": "completed",
2044 },
2045 })
2046 .to_string(),
2047 ))
2048 .await
2049 .expect("done event should send");
2050 }
2051 });
2052
2053 let base_url = format!("http://{address}/v1");
2054 let client = crate::providers::openai::Client::builder()
2055 .api_key("test-key")
2056 .base_url(&base_url)
2057 .build()
2058 .expect("client should build");
2059 let model = client.completion_model("gpt-4o");
2060 let mut session = client
2061 .responses_websocket("gpt-4o")
2062 .await
2063 .expect("session should connect");
2064
2065 session
2066 .send(model.completion_request("first").build())
2067 .await
2068 .expect("first request should send");
2069 let first = session
2070 .wait_for_completed_response()
2071 .await
2072 .expect("first response should complete");
2073 assert_eq!(first.id, "resp_1");
2074 assert_eq!(session.previous_response_id(), Some("resp_1"));
2075
2076 session
2077 .send(model.completion_request("second").build())
2078 .await
2079 .expect("second request should send");
2080 let second = session
2081 .wait_for_completed_response()
2082 .await
2083 .expect("second response should complete");
2084 assert_eq!(second.id, "resp_2");
2085 assert_eq!(session.previous_response_id(), Some("resp_2"));
2086
2087 server.await.expect("server task should finish");
2088 }
2089
2090 #[tokio::test]
2091 async fn clearing_previous_response_id_does_not_disable_late_done_filter() {
2092 let listener = TcpListener::bind("127.0.0.1:0")
2093 .await
2094 .expect("listener should bind");
2095 let address = listener.local_addr().expect("listener should have address");
2096
2097 let server = tokio::spawn(async move {
2098 let (stream, _) = listener.accept().await.expect("server should accept");
2099 let mut socket = accept_async(stream)
2100 .await
2101 .expect("server should upgrade websocket");
2102
2103 for response_id in ["resp_1", "resp_2"] {
2104 let request = socket
2105 .next()
2106 .await
2107 .expect("request should exist")
2108 .expect("request should be valid");
2109 let payload = request.into_text().expect("request should be text");
2110 assert!(
2111 payload.contains("\"type\":\"response.create\""),
2112 "expected response.create payload, got {payload}"
2113 );
2114
2115 let response = sample_response(ResponseStatus::Completed);
2116 let response = serde_json::to_value(CompletionResponse {
2117 id: response_id.to_string(),
2118 ..response
2119 })
2120 .expect("response should serialize");
2121
2122 socket
2123 .send(Message::text(
2124 json!({
2125 "type": "response.completed",
2126 "sequence_number": 1,
2127 "response": response,
2128 })
2129 .to_string(),
2130 ))
2131 .await
2132 .expect("completed event should send");
2133 socket
2134 .send(Message::text(
2135 json!({
2136 "type": "response.done",
2137 "response": {
2138 "id": response_id,
2139 "status": "completed",
2140 },
2141 })
2142 .to_string(),
2143 ))
2144 .await
2145 .expect("done event should send");
2146 }
2147 });
2148
2149 let base_url = format!("http://{address}/v1");
2150 let client = crate::providers::openai::Client::builder()
2151 .api_key("test-key")
2152 .base_url(&base_url)
2153 .build()
2154 .expect("client should build");
2155 let model = client.completion_model("gpt-4o");
2156 let mut session = client
2157 .responses_websocket("gpt-4o")
2158 .await
2159 .expect("session should connect");
2160
2161 session
2162 .send(model.completion_request("first").build())
2163 .await
2164 .expect("first request should send");
2165 let first = session
2166 .wait_for_completed_response()
2167 .await
2168 .expect("first response should complete");
2169 assert_eq!(first.id, "resp_1");
2170
2171 session.clear_previous_response_id();
2172 assert_eq!(session.previous_response_id(), None);
2173
2174 session
2175 .send(model.completion_request("second").build())
2176 .await
2177 .expect("second request should send");
2178 let second = session
2179 .wait_for_completed_response()
2180 .await
2181 .expect("second response should complete");
2182 assert_eq!(second.id, "resp_2");
2183
2184 server.await.expect("server task should finish");
2185 }
2186
2187 #[tokio::test]
2188 async fn failed_turn_keeps_late_done_out_of_next_request() {
2189 let listener = TcpListener::bind("127.0.0.1:0")
2190 .await
2191 .expect("listener should bind");
2192 let address = listener.local_addr().expect("listener should have address");
2193
2194 let server = tokio::spawn(async move {
2195 let (stream, _) = listener.accept().await.expect("server should accept");
2196 let mut socket = accept_async(stream)
2197 .await
2198 .expect("server should upgrade websocket");
2199
2200 let first_request = socket
2201 .next()
2202 .await
2203 .expect("request should exist")
2204 .expect("request should be valid");
2205 let payload = first_request
2206 .into_text()
2207 .expect("failed request should be text");
2208 assert!(
2209 payload.contains("\"type\":\"response.create\""),
2210 "expected response.create payload, got {payload}"
2211 );
2212
2213 let failed_response = serde_json::to_value(CompletionResponse {
2214 id: "resp_failed".to_string(),
2215 status: ResponseStatus::Failed,
2216 ..sample_response(ResponseStatus::Completed)
2217 })
2218 .expect("failed response should serialize");
2219
2220 socket
2221 .send(Message::text(
2222 json!({
2223 "type": "response.failed",
2224 "sequence_number": 1,
2225 "response": failed_response,
2226 })
2227 .to_string(),
2228 ))
2229 .await
2230 .expect("failed event should send");
2231 socket
2232 .send(Message::text(
2233 json!({
2234 "type": "response.done",
2235 "response": {
2236 "id": "resp_failed",
2237 "status": "failed",
2238 },
2239 })
2240 .to_string(),
2241 ))
2242 .await
2243 .expect("done event should send");
2244
2245 let second_request = socket
2246 .next()
2247 .await
2248 .expect("request should exist")
2249 .expect("request should be valid");
2250 let payload = second_request
2251 .into_text()
2252 .expect("second request should be text");
2253 assert!(
2254 payload.contains("\"type\":\"response.create\""),
2255 "expected response.create payload, got {payload}"
2256 );
2257
2258 let response = sample_response(ResponseStatus::Completed);
2259 let response = serde_json::to_value(CompletionResponse {
2260 id: "resp_2".to_string(),
2261 ..response
2262 })
2263 .expect("response should serialize");
2264
2265 socket
2266 .send(Message::text(
2267 json!({
2268 "type": "response.completed",
2269 "sequence_number": 2,
2270 "response": response,
2271 })
2272 .to_string(),
2273 ))
2274 .await
2275 .expect("completed event should send");
2276 socket
2277 .send(Message::text(
2278 json!({
2279 "type": "response.done",
2280 "response": {
2281 "id": "resp_2",
2282 "status": "completed",
2283 },
2284 })
2285 .to_string(),
2286 ))
2287 .await
2288 .expect("done event should send");
2289 });
2290
2291 let base_url = format!("http://{address}/v1");
2292 let client = crate::providers::openai::Client::builder()
2293 .api_key("test-key")
2294 .base_url(&base_url)
2295 .build()
2296 .expect("client should build");
2297 let model = client.completion_model("gpt-4o");
2298 let mut session = client
2299 .responses_websocket("gpt-4o")
2300 .await
2301 .expect("session should connect");
2302
2303 session
2304 .send(model.completion_request("first").build())
2305 .await
2306 .expect("first request should send");
2307 let error = session
2308 .wait_for_completed_response()
2309 .await
2310 .expect_err("failed response should error");
2311 assert!(error.to_string().contains("failed response"));
2312 assert_eq!(session.previous_response_id(), None);
2313
2314 session
2315 .send(model.completion_request("second").build())
2316 .await
2317 .expect("second request should send");
2318 let second = session
2319 .wait_for_completed_response()
2320 .await
2321 .expect("second response should complete");
2322 assert_eq!(second.id, "resp_2");
2323
2324 server.await.expect("server task should finish");
2325 }
2326
2327 #[tokio::test]
2328 async fn done_first_completed_turn_updates_previous_response_id() {
2329 let listener = TcpListener::bind("127.0.0.1:0")
2330 .await
2331 .expect("listener should bind");
2332 let address = listener.local_addr().expect("listener should have address");
2333
2334 let server = tokio::spawn(async move {
2335 let (stream, _) = listener.accept().await.expect("server should accept");
2336 let mut socket = accept_async(stream)
2337 .await
2338 .expect("server should upgrade websocket");
2339
2340 for response_id in ["resp_1", "resp_2"] {
2341 let request = socket
2342 .next()
2343 .await
2344 .expect("request should exist")
2345 .expect("request should be valid");
2346 let payload = request.into_text().expect("request should be text");
2347 assert!(
2348 payload.contains("\"type\":\"response.create\""),
2349 "expected response.create payload, got {payload}"
2350 );
2351
2352 if response_id == "resp_2" {
2353 assert!(
2354 payload.contains("\"previous_response_id\":\"resp_1\""),
2355 "expected chained previous_response_id in payload, got {payload}"
2356 );
2357 }
2358
2359 let response = serde_json::to_value(CompletionResponse {
2360 id: response_id.to_string(),
2361 ..sample_response(ResponseStatus::Completed)
2362 })
2363 .expect("response should serialize");
2364
2365 socket
2366 .send(Message::text(
2367 json!({
2368 "type": "response.done",
2369 "response": response,
2370 })
2371 .to_string(),
2372 ))
2373 .await
2374 .expect("done event should send");
2375 }
2376 });
2377
2378 let base_url = format!("http://{address}/v1");
2379 let client = crate::providers::openai::Client::builder()
2380 .api_key("test-key")
2381 .base_url(&base_url)
2382 .build()
2383 .expect("client should build");
2384 let model = client.completion_model("gpt-4o");
2385 let mut session = client
2386 .responses_websocket("gpt-4o")
2387 .await
2388 .expect("session should connect");
2389
2390 session
2391 .send(model.completion_request("first").build())
2392 .await
2393 .expect("first request should send");
2394 let first = session
2395 .wait_for_completed_response()
2396 .await
2397 .expect("first response should complete");
2398 assert_eq!(first.id, "resp_1");
2399 assert_eq!(session.previous_response_id(), Some("resp_1"));
2400
2401 session
2402 .send(model.completion_request("second").build())
2403 .await
2404 .expect("second request should send");
2405 let second = session
2406 .wait_for_completed_response()
2407 .await
2408 .expect("second response should complete");
2409 assert_eq!(second.id, "resp_2");
2410 assert_eq!(session.previous_response_id(), Some("resp_2"));
2411
2412 server.await.expect("server task should finish");
2413 }
2414
2415 #[tokio::test]
2416 async fn done_first_failed_turn_does_not_chain_next_request() {
2417 let listener = TcpListener::bind("127.0.0.1:0")
2418 .await
2419 .expect("listener should bind");
2420 let address = listener.local_addr().expect("listener should have address");
2421
2422 let server = tokio::spawn(async move {
2423 let (stream, _) = listener.accept().await.expect("server should accept");
2424 let mut socket = accept_async(stream)
2425 .await
2426 .expect("server should upgrade websocket");
2427
2428 let first_request = socket
2429 .next()
2430 .await
2431 .expect("request should exist")
2432 .expect("request should be valid");
2433 let payload = first_request
2434 .into_text()
2435 .expect("first request should be text");
2436 assert!(
2437 payload.contains("\"type\":\"response.create\""),
2438 "expected response.create payload, got {payload}"
2439 );
2440 assert!(
2441 !payload.contains("\"previous_response_id\""),
2442 "did not expect previous_response_id in first payload, got {payload}"
2443 );
2444
2445 let failed_response = serde_json::to_value(CompletionResponse {
2446 id: "resp_failed".to_string(),
2447 status: ResponseStatus::Failed,
2448 ..sample_response(ResponseStatus::Completed)
2449 })
2450 .expect("failed response should serialize");
2451
2452 socket
2453 .send(Message::text(
2454 json!({
2455 "type": "response.done",
2456 "response": failed_response,
2457 })
2458 .to_string(),
2459 ))
2460 .await
2461 .expect("done event should send");
2462
2463 let second_request = socket
2464 .next()
2465 .await
2466 .expect("request should exist")
2467 .expect("request should be valid");
2468 let payload = second_request
2469 .into_text()
2470 .expect("second request should be text");
2471 assert!(
2472 payload.contains("\"type\":\"response.create\""),
2473 "expected response.create payload, got {payload}"
2474 );
2475 assert!(
2476 !payload.contains("\"previous_response_id\""),
2477 "did not expect chained previous_response_id in payload, got {payload}"
2478 );
2479
2480 let response = serde_json::to_value(CompletionResponse {
2481 id: "resp_2".to_string(),
2482 ..sample_response(ResponseStatus::Completed)
2483 })
2484 .expect("response should serialize");
2485
2486 socket
2487 .send(Message::text(
2488 json!({
2489 "type": "response.done",
2490 "response": response,
2491 })
2492 .to_string(),
2493 ))
2494 .await
2495 .expect("done event should send");
2496 });
2497
2498 let base_url = format!("http://{address}/v1");
2499 let client = crate::providers::openai::Client::builder()
2500 .api_key("test-key")
2501 .base_url(&base_url)
2502 .build()
2503 .expect("client should build");
2504 let model = client.completion_model("gpt-4o");
2505 let mut session = client
2506 .responses_websocket("gpt-4o")
2507 .await
2508 .expect("session should connect");
2509
2510 session
2511 .send(model.completion_request("first").build())
2512 .await
2513 .expect("first request should send");
2514 let error = session
2515 .wait_for_completed_response()
2516 .await
2517 .expect_err("failed response should error");
2518 assert!(error.to_string().contains("failed response"));
2519 assert_eq!(session.previous_response_id(), None);
2520
2521 session
2522 .send(model.completion_request("second").build())
2523 .await
2524 .expect("second request should send");
2525 let second = session
2526 .wait_for_completed_response()
2527 .await
2528 .expect("second response should complete");
2529 assert_eq!(second.id, "resp_2");
2530 assert_eq!(session.previous_response_id(), Some("resp_2"));
2531
2532 server.await.expect("server task should finish");
2533 }
2534
2535 #[test]
2536 fn websocket_url_converts_http_to_ws() {
2537 let url = websocket_url("http://localhost:8080/v1").expect("url should convert");
2538 assert_eq!(url, "ws://localhost:8080/v1/responses");
2539 }
2540
2541 #[test]
2542 fn websocket_url_rejects_unsupported_scheme() {
2543 let result = websocket_url("ftp://example.com/v1");
2544 assert!(result.is_err());
2545 }
2546
2547 #[test]
2548 fn websocket_url_trims_trailing_slash() {
2549 let url = websocket_url("https://api.openai.com/v1/").expect("url should convert");
2550 assert_eq!(url, "wss://api.openai.com/v1/responses");
2551 }
2552
2553 #[test]
2554 fn unknown_event_type_is_forwarded_raw() {
2555 let payload = json!({
2556 "type": "response.some_future_event",
2557 "data": "hello"
2558 });
2559
2560 let result =
2561 parse_server_event(&payload.to_string()).expect("unknown event should not error");
2562 match result {
2565 Some(ResponsesWebSocketEvent::Unknown(value)) => assert_eq!(value, payload.into()),
2566 other => panic!("expected the raw Unknown passthrough event, got {other:?}"),
2567 }
2568 }
2569
2570 #[test]
2571 fn malformed_known_event_returns_error() {
2572 let payload = json!({
2573 "type": "response.completed"
2574 });
2575
2576 let error = parse_server_event(&payload.to_string())
2577 .expect_err("malformed known event should error");
2578 assert!(
2579 error.to_string().contains("StreamingCompletionChunk"),
2580 "expected strict decode failure, got {error}"
2581 );
2582 }
2583
2584 #[tokio::test]
2585 async fn close_is_idempotent() {
2586 let listener = TcpListener::bind("127.0.0.1:0")
2587 .await
2588 .expect("listener should bind");
2589 let address = listener.local_addr().expect("listener should have address");
2590
2591 let server = tokio::spawn(async move {
2592 let (stream, _) = listener.accept().await.expect("server should accept");
2593 let mut socket = accept_async(stream)
2594 .await
2595 .expect("server should upgrade websocket");
2596
2597 let message = socket
2598 .next()
2599 .await
2600 .expect("close frame should arrive")
2601 .expect("close frame should be valid");
2602 assert!(
2603 matches!(message, Message::Close(_)),
2604 "expected close frame, got {message:?}"
2605 );
2606 });
2607
2608 let base_url = format!("http://{address}/v1");
2609 let client = crate::providers::openai::Client::builder()
2610 .api_key("test-key")
2611 .base_url(&base_url)
2612 .build()
2613 .expect("client should build");
2614 let mut session = client
2615 .responses_websocket("gpt-4o")
2616 .await
2617 .expect("session should connect");
2618
2619 session.close().await.expect("first close should succeed");
2620 session.close().await.expect("second close should succeed");
2621
2622 server.await.expect("server task should finish");
2623 }
2624
2625 #[tokio::test]
2626 async fn send_while_in_flight_returns_error() {
2627 let listener = TcpListener::bind("127.0.0.1:0")
2628 .await
2629 .expect("listener should bind");
2630 let address = listener.local_addr().expect("listener should have address");
2631
2632 let server = tokio::spawn(async move {
2633 let (stream, _) = listener.accept().await.expect("server should accept");
2634 let mut socket = accept_async(stream)
2635 .await
2636 .expect("server should upgrade websocket");
2637
2638 let _request = socket
2640 .next()
2641 .await
2642 .expect("request should exist")
2643 .expect("request should be valid");
2644
2645 sleep(Duration::from_millis(100)).await;
2647 let _ = socket.close(None).await;
2648 });
2649
2650 let base_url = format!("http://{address}/v1");
2651 let client = crate::providers::openai::Client::builder()
2652 .api_key("test-key")
2653 .base_url(&base_url)
2654 .build()
2655 .expect("client should build");
2656 let model = client.completion_model("gpt-4o");
2657 let mut session = client
2658 .responses_websocket("gpt-4o")
2659 .await
2660 .expect("session should connect");
2661
2662 session
2663 .send(model.completion_request("first").build())
2664 .await
2665 .expect("first request should send");
2666
2667 let error = session
2668 .send(model.completion_request("second").build())
2669 .await
2670 .expect_err("second send while in-flight should error");
2671 assert!(
2672 error.to_string().contains("already in flight"),
2673 "expected in-flight error, got {error}"
2674 );
2675
2676 server.await.expect("server task should finish");
2677 }
2678
2679 #[tokio::test]
2680 async fn send_after_close_returns_error() {
2681 let listener = TcpListener::bind("127.0.0.1:0")
2682 .await
2683 .expect("listener should bind");
2684 let address = listener.local_addr().expect("listener should have address");
2685
2686 let server = tokio::spawn(async move {
2687 let (stream, _) = listener.accept().await.expect("server should accept");
2688 let _socket = accept_async(stream)
2689 .await
2690 .expect("server should upgrade websocket");
2691 sleep(Duration::from_millis(100)).await;
2692 });
2693
2694 let base_url = format!("http://{address}/v1");
2695 let client = crate::providers::openai::Client::builder()
2696 .api_key("test-key")
2697 .base_url(&base_url)
2698 .build()
2699 .expect("client should build");
2700 let model = client.completion_model("gpt-4o");
2701 let mut session = client
2702 .responses_websocket("gpt-4o")
2703 .await
2704 .expect("session should connect");
2705
2706 session.close().await.expect("close should succeed");
2707
2708 let error = session
2709 .send(model.completion_request("after close").build())
2710 .await
2711 .expect_err("send after close should error");
2712 assert!(
2713 error.to_string().contains("session is closed"),
2714 "expected closed-session error, got {error}"
2715 );
2716
2717 server.await.expect("server task should finish");
2718 }
2719
2720 #[tokio::test]
2721 async fn next_event_without_send_returns_error() {
2722 let listener = TcpListener::bind("127.0.0.1:0")
2723 .await
2724 .expect("listener should bind");
2725 let address = listener.local_addr().expect("listener should have address");
2726
2727 let server = tokio::spawn(async move {
2728 let (stream, _) = listener.accept().await.expect("server should accept");
2729 let _socket = accept_async(stream)
2730 .await
2731 .expect("server should upgrade websocket");
2732 sleep(Duration::from_millis(100)).await;
2733 });
2734
2735 let base_url = format!("http://{address}/v1");
2736 let client = crate::providers::openai::Client::builder()
2737 .api_key("test-key")
2738 .base_url(&base_url)
2739 .build()
2740 .expect("client should build");
2741 let mut session = client
2742 .responses_websocket("gpt-4o")
2743 .await
2744 .expect("session should connect");
2745
2746 let error = session
2747 .next_event()
2748 .await
2749 .expect_err("next_event without send should error");
2750 assert!(
2751 error
2752 .to_string()
2753 .contains("No OpenAI websocket response is currently in flight"),
2754 "expected not-in-flight error, got {error}"
2755 );
2756
2757 server.await.expect("server task should finish");
2758 }
2759
2760 #[tokio::test]
2761 async fn unknown_event_is_skipped_and_reasoning_metadata_is_preserved() {
2762 let listener = TcpListener::bind("127.0.0.1:0")
2763 .await
2764 .expect("listener should bind");
2765 let address = listener.local_addr().expect("listener should have address");
2766
2767 let server = tokio::spawn(async move {
2768 let (stream, _) = listener.accept().await.expect("server should accept");
2769 let mut socket = accept_async(stream)
2770 .await
2771 .expect("server should upgrade websocket");
2772
2773 let _request = socket
2774 .next()
2775 .await
2776 .expect("request should exist")
2777 .expect("request should be valid");
2778
2779 socket
2781 .send(Message::text(
2782 json!({
2783 "type": "response.some_future_event",
2784 "data": "should be skipped"
2785 })
2786 .to_string(),
2787 ))
2788 .await
2789 .expect("unknown event should send");
2790
2791 let mut response = sample_response(ResponseStatus::Completed);
2794 response.id = "resp_after_unknown".to_string();
2795 response.reasoning_metadata = Some(
2796 json!({
2797 "context": "all_turns",
2798 "effort": "ultra",
2799 "summary": null,
2800 "future_control": true
2801 })
2802 .as_object()
2803 .expect("reasoning metadata should be an object")
2804 .clone(),
2805 );
2806 response.reasoning_context = Some("all_turns".to_string());
2807 let response = serde_json::to_value(response).expect("response should serialize");
2808
2809 socket
2810 .send(Message::text(
2811 json!({
2812 "type": "response.completed",
2813 "sequence_number": 1,
2814 "response": response,
2815 })
2816 .to_string(),
2817 ))
2818 .await
2819 .expect("completed event should send");
2820 });
2821
2822 let base_url = format!("http://{address}/v1");
2823 let client = crate::providers::openai::Client::builder()
2824 .api_key("test-key")
2825 .base_url(&base_url)
2826 .build()
2827 .expect("client should build");
2828 let model = client.completion_model("gpt-4o");
2829 let mut session = client
2830 .responses_websocket("gpt-4o")
2831 .await
2832 .expect("session should connect");
2833
2834 session
2835 .send(model.completion_request("hello").build())
2836 .await
2837 .expect("send should succeed");
2838 let response = session
2839 .wait_for_completed_response()
2840 .await
2841 .expect("response should complete despite unknown event");
2842 assert_eq!(response.id, "resp_after_unknown");
2843 assert_eq!(response.reasoning_context.as_deref(), Some("all_turns"));
2844 assert_eq!(
2845 response.reasoning_metadata.as_ref(),
2846 json!({
2847 "context": "all_turns",
2848 "effort": "ultra",
2849 "summary": null,
2850 "future_control": true
2851 })
2852 .as_object()
2853 );
2854
2855 server.await.expect("server task should finish");
2856 }
2857
2858 fn ws_messages_from_sse_frames<'a>(
2862 frames: impl IntoIterator<Item = &'a bytes::Bytes>,
2863 ) -> Vec<String> {
2864 frames
2865 .into_iter()
2866 .flat_map(|frame| {
2867 std::str::from_utf8(frame)
2868 .expect("SSE fixture frames should be UTF-8")
2869 .lines()
2870 .filter_map(|line| line.strip_prefix("data:").map(str::trim))
2871 .filter(|data| !data.is_empty() && *data != "[DONE]")
2872 .map(ToOwned::to_owned)
2873 .collect::<Vec<_>>()
2874 })
2875 .collect()
2876 }
2877
2878 fn spawn_ws_server_with_messages(
2879 listener: TcpListener,
2880 messages: Vec<String>,
2881 ) -> tokio::task::JoinHandle<()> {
2882 tokio::spawn(async move {
2883 let (stream, _) = listener.accept().await.expect("server should accept");
2884 let mut socket = accept_async(stream)
2885 .await
2886 .expect("server should upgrade websocket");
2887
2888 let request = socket
2889 .next()
2890 .await
2891 .expect("request should exist")
2892 .expect("request should be valid");
2893 let payload = request.into_text().expect("request should be text");
2894 assert!(
2895 payload.contains("\"type\":\"response.create\""),
2896 "expected response.create payload, got {payload}"
2897 );
2898
2899 for message in messages {
2900 socket
2901 .send(Message::text(message))
2902 .await
2903 .expect("event should send");
2904 }
2905 })
2906 }
2907
2908 #[tokio::test]
2915 async fn websocket_conformance_replays_sse_fixture_frames() {
2916 let fixture =
2917 crate::test_utils::streaming_conformance::fixtures::openai_responses::fixture();
2918 let byte_frame = |frame: &crate::test_utils::streaming_conformance::WireInput| {
2920 frame
2921 .as_bytes()
2922 .cloned()
2923 .expect("the Responses fixture scripts byte frames")
2924 };
2925 let mut frames: Vec<bytes::Bytes> = Vec::new();
2926 frames.extend(fixture.text_frames.iter().map(byte_frame));
2927 frames.extend(fixture.tool_call_frames.iter().map(byte_frame));
2928 frames.extend(fixture.unknown_event_frame.iter().map(byte_frame));
2929 frames.extend(fixture.terminal_frames.iter().map(byte_frame));
2930 let messages = ws_messages_from_sse_frames(frames.iter());
2931
2932 let listener = TcpListener::bind("127.0.0.1:0")
2933 .await
2934 .expect("listener should bind");
2935 let address = listener.local_addr().expect("listener should have address");
2936 let server = spawn_ws_server_with_messages(listener, messages);
2937
2938 let base_url = format!("http://{address}/v1");
2939 let client = crate::providers::openai::Client::builder()
2940 .api_key("test-key")
2941 .base_url(&base_url)
2942 .build()
2943 .expect("client should build");
2944 let model = client.completion_model("gpt-4o");
2945 let mut session = client
2946 .responses_websocket("gpt-4o")
2947 .await
2948 .expect("session should connect");
2949
2950 let normalized = session
2951 .completion(model.completion_request("hello").build())
2952 .await
2953 .expect("fixture turn should normalize");
2954
2955 let texts: Vec<&str> = normalized
2956 .choice
2957 .iter()
2958 .filter_map(|content| match content {
2959 crate::completion::AssistantContent::Text(text) => Some(text.text.as_str()),
2960 _ => None,
2961 })
2962 .collect();
2963 assert_eq!(texts, fixture.expected_texts);
2964 let tool_names: Vec<&str> = normalized
2965 .choice
2966 .iter()
2967 .filter_map(|content| match content {
2968 crate::completion::AssistantContent::ToolCall(call) => {
2969 Some(call.function.name.as_str())
2970 }
2971 _ => None,
2972 })
2973 .collect();
2974 assert_eq!(tool_names, vec![fixture.expected_tool_name]);
2975 assert_eq!(normalized.usage.total_tokens, fixture.expected_usage_total);
2976 assert_eq!(
2980 normalized.finish_reason(),
2981 Some(crate::completion::FinishReason::ToolCalls)
2982 );
2983
2984 server.await.expect("server task should finish");
2985 }
2986
2987 #[tokio::test]
2992 async fn reasoning_text_delta_arrives_over_websocket() {
2993 let messages = vec![
2994 json!({
2995 "type": "response.reasoning_text.delta",
2996 "item_id": "rs_1",
2997 "output_index": 0,
2998 "content_index": 0,
2999 "sequence_number": 1,
3000 "delta": "thinking hard",
3001 })
3002 .to_string(),
3003 json!({
3004 "type": "response.output_text.delta",
3005 "content_index": 0,
3006 "delta": "answer",
3007 "item_id": "msg_1",
3008 "output_index": 0,
3009 "sequence_number": 2,
3010 })
3011 .to_string(),
3012 json!({
3013 "type": "response.completed",
3014 "sequence_number": 3,
3015 "response": serde_json::to_value(sample_response(ResponseStatus::Completed))
3016 .expect("response should serialize"),
3017 })
3018 .to_string(),
3019 ];
3020
3021 let listener = TcpListener::bind("127.0.0.1:0")
3022 .await
3023 .expect("listener should bind");
3024 let address = listener.local_addr().expect("listener should have address");
3025 let server = spawn_ws_server_with_messages(listener, messages);
3026
3027 let base_url = format!("http://{address}/v1");
3028 let client = crate::providers::openai::Client::builder()
3029 .api_key("test-key")
3030 .base_url(&base_url)
3031 .build()
3032 .expect("client should build");
3033 let model = client.completion_model("gpt-4o");
3034 let mut session = client
3035 .responses_websocket("gpt-4o")
3036 .await
3037 .expect("session should connect");
3038
3039 let normalized = session
3040 .completion(model.completion_request("hello").build())
3041 .await
3042 .expect("turn with reasoning deltas should normalize");
3043
3044 assert!(
3045 normalized.choice.iter().any(|content| matches!(
3046 content,
3047 crate::completion::AssistantContent::Reasoning(reasoning)
3048 if reasoning.content.iter().any(|block| matches!(
3049 block,
3050 crate::message::ReasoningContent::Text { text, .. }
3051 if text.contains("thinking hard")
3052 ))
3053 )),
3054 "reasoning delta should survive over websocket, got {:?}",
3055 normalized.choice
3056 );
3057 assert!(
3058 normalized.choice.iter().any(|content| matches!(
3059 content,
3060 crate::completion::AssistantContent::Text(text) if text.text == "answer"
3061 )),
3062 "text delta should survive alongside reasoning, got {:?}",
3063 normalized.choice
3064 );
3065
3066 server.await.expect("server task should finish");
3067 }
3068
3069 #[test]
3070 fn parse_reasoning_text_delta_event_is_item() {
3071 let payload = json!({
3072 "type": "response.reasoning_text.delta",
3073 "item_id": "rs_1",
3074 "output_index": 0,
3075 "content_index": 0,
3076 "sequence_number": 1,
3077 "delta": "thinking",
3078 });
3079
3080 let event = parse_server_event(&payload.to_string())
3081 .expect("reasoning delta should parse")
3082 .expect("reasoning delta should not be skipped");
3083
3084 assert!(matches!(event, ResponsesWebSocketEvent::Item(_)));
3085 assert!(!event.is_terminal());
3086 }
3087}