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