1use async_trait::async_trait;
2use serde_json::{Value, json};
3use std::sync::{
4 Arc, Mutex,
5 atomic::{AtomicBool, Ordering},
6};
7
8use crate::{
9 ModelAdapter, ModelError, ModelResponseEventStream, ModelResponseStreamEvent, ModelRunSession,
10 StreamDiagnostic,
11 adapter::{ModelRequestContext, ModelRequestParameters, allow_real_model_requests},
12 message::{ModelMessage, ModelResponse},
13 profile::{ModelProfile, ProtocolFamily},
14 request::prepare_model_request,
15 settings::{ModelSettings, ResponseStreamTransport},
16 transport::{
17 HttpRequest, ModelEventStream, ModelWebSocketEventSession, build_http_request,
18 send_event_stream_with_retries, send_websocket_event_stream_with_retries,
19 send_websocket_session_event_stream_with_retries, send_with_retries,
20 should_fallback_websocket_to_http,
21 },
22};
23
24use super::ProtocolModelClient;
25
26#[derive(Clone, Copy, Debug, Eq, PartialEq)]
27enum ResolvedResponseStreamTransport {
28 HttpOnly,
29 WebSocketOnly,
30 WebSocketThenHttpOnRetryable,
31}
32
33#[async_trait]
34impl ModelAdapter for ProtocolModelClient {
35 fn model_name(&self) -> &str {
36 &self.model_name
37 }
38
39 fn provider_name(&self) -> Option<&str> {
40 Some(&self.provider_name)
41 }
42
43 fn profile(&self) -> &ModelProfile {
44 &self.profile
45 }
46
47 fn default_settings(&self) -> Option<&ModelSettings> {
48 self.default_settings.as_ref()
49 }
50
51 fn start_run_session(&self) -> Box<dyn ModelRunSession + '_> {
52 Box::new(ProtocolModelClientRunSession::new(self))
53 }
54
55 async fn request(
56 &self,
57 messages: Vec<ModelMessage>,
58 settings: Option<ModelSettings>,
59 params: ModelRequestParameters,
60 context: ModelRequestContext,
61 ) -> Result<ModelResponse, ModelError> {
62 let prepared = prepare_model_request(
63 messages,
64 self.default_settings.as_ref(),
65 settings,
66 params,
67 &self.profile,
68 );
69 let wire_body = self.build_wire_body(
70 &prepared.normalized_messages,
71 prepared.settings.as_ref(),
72 &prepared.params,
73 )?;
74 let options = self.request_options(&context, prepared.settings.as_ref(), &prepared.params);
75 let mut request = build_http_request(&self.http_config, &options, wire_body);
76 request.cancellation_token = context.cancellation_token();
77 self.finalize_http_request(&mut request)?;
78 if let Some(audit) = self.request_audit.as_ref() {
79 audit.record(&self.provider_name, &self.model_name, false, &request);
80 }
81 if !allow_real_model_requests() {
82 return Err(ModelError::RealModelRequestBlocked { url: request.url });
83 }
84 let response = send_with_retries(
85 self.http_client.as_ref(),
86 self.sleeper.as_ref(),
87 request,
88 &self.http_config.retry_policy,
89 )
90 .await?;
91 self.parse_wire_response(&response.body)
92 }
93
94 async fn request_stream(
95 &self,
96 messages: Vec<ModelMessage>,
97 settings: Option<ModelSettings>,
98 params: ModelRequestParameters,
99 context: ModelRequestContext,
100 ) -> Result<Vec<ModelResponseStreamEvent>, ModelError> {
101 let mut stream = self
102 .request_stream_incremental(messages, settings, params, context)
103 .await?;
104 let mut events = Vec::new();
105 while let Some(event) = stream.recv().await {
106 events.push(event?);
107 }
108 Ok(events)
109 }
110
111 async fn request_stream_incremental(
112 &self,
113 messages: Vec<ModelMessage>,
114 settings: Option<ModelSettings>,
115 params: ModelRequestParameters,
116 context: ModelRequestContext,
117 ) -> Result<ModelResponseEventStream, ModelError> {
118 let mut session = ProtocolModelClientRunSession::new(self);
119 session
120 .request_stream_incremental(messages, settings, params, context)
121 .await
122 }
123}
124
125#[derive(Clone, Copy, Debug)]
126enum ResponseStreamRequestKind {
127 HttpSse,
128 WebSocket,
129}
130
131struct ProtocolModelClientRunSession<'a> {
132 client: &'a ProtocolModelClient,
133 websocket_session: Box<dyn ModelWebSocketEventSession + 'a>,
134 last_websocket_request: Option<Value>,
135 last_response: Arc<Mutex<LastResponseSlot>>,
136 fallback_to_http: Arc<AtomicBool>,
137}
138
139#[derive(Clone, Debug)]
140struct LastResponse {
141 response_id: String,
142 replay_items: Vec<Value>,
143}
144
145#[derive(Debug, Default)]
146struct LastResponseSlot {
147 response: Option<LastResponse>,
148 failed: bool,
149}
150
151struct SessionFallbackPlan {
152 http_client: crate::transport::DynHttpClient,
153 request_audit: Option<crate::transport::ProviderRequestAuditCapture>,
154 provider_name: String,
155 model_name: String,
156 http_request: HttpRequest,
157 sleeper: crate::transport::DynSleeper,
158 retry_policy: crate::transport::RetryPolicy,
159 fallback_to_http: Arc<AtomicBool>,
160 last_response: Arc<Mutex<LastResponseSlot>>,
161}
162
163impl<'a> ProtocolModelClientRunSession<'a> {
164 fn new(client: &'a ProtocolModelClient) -> Self {
165 Self {
166 client,
167 websocket_session: client.http_client.websocket_event_session(),
168 last_websocket_request: None,
169 last_response: Arc::new(Mutex::new(LastResponseSlot::default())),
170 fallback_to_http: Arc::new(AtomicBool::new(false)),
171 }
172 }
173
174 fn fallback_plan(&self, http_request: HttpRequest) -> SessionFallbackPlan {
175 SessionFallbackPlan {
176 http_client: self.client.http_client.clone(),
177 request_audit: self.client.request_audit.clone(),
178 provider_name: self.client.provider_name.clone(),
179 model_name: self.client.model_name.clone(),
180 http_request,
181 sleeper: self.client.sleeper.clone(),
182 retry_policy: self.client.http_config.retry_policy.clone(),
183 fallback_to_http: Arc::clone(&self.fallback_to_http),
184 last_response: Arc::clone(&self.last_response),
185 }
186 }
187
188 fn required_fallback_plan(
189 &self,
190 http_request: Option<HttpRequest>,
191 ) -> Result<SessionFallbackPlan, ModelError> {
192 http_request
193 .map(|request| self.fallback_plan(request))
194 .ok_or_else(|| {
195 ModelError::Transport("HTTP request is required for websocket fallback".to_string())
196 })
197 }
198
199 fn current_last_response(&mut self) -> Option<LastResponse> {
200 let mut slot = self
201 .last_response
202 .lock()
203 .unwrap_or_else(std::sync::PoisonError::into_inner);
204 if slot.failed {
205 slot.failed = false;
206 slot.response = None;
207 self.last_websocket_request = None;
208 return None;
209 }
210 slot.response.clone()
211 }
212
213 fn prepare_websocket_request_body(&mut self, logical_body: &Value) -> Value {
214 let Some(last_response) = self.current_last_response() else {
215 return logical_body.clone();
216 };
217 let Some(incremental_items) =
218 self.websocket_incremental_items(logical_body, &last_response)
219 else {
220 return logical_body.clone();
221 };
222 if last_response.response_id.is_empty() {
223 return logical_body.clone();
224 }
225 let mut body = logical_body.clone();
226 if let Some(object) = body.as_object_mut() {
227 object.insert(
228 "previous_response_id".to_string(),
229 json!(last_response.response_id),
230 );
231 object.insert("input".to_string(), Value::Array(incremental_items));
232 }
233 body
234 }
235
236 fn websocket_incremental_items(
237 &self,
238 current: &Value,
239 last_response: &LastResponse,
240 ) -> Option<Vec<Value>> {
241 if current.get("previous_response_id").is_some() || current.get("conversation").is_some() {
242 return None;
243 }
244 let previous = self.last_websocket_request.as_ref()?;
245 if !responses_websocket_request_properties_match(previous, current) {
246 return None;
247 }
248 let previous_input = response_input_items(previous)?;
249 let current_input = response_input_items(current)?;
250 let after_previous = strip_value_prefix(current_input, previous_input)?;
251 let incremental = strip_value_prefix(after_previous, &last_response.replay_items)?;
252 Some(incremental.to_vec())
253 }
254
255 async fn request_websocket_stream(
256 &mut self,
257 mut websocket_request: HttpRequest,
258 http_request: Option<HttpRequest>,
259 request_settings: Option<ModelSettings>,
260 fallback: bool,
261 ) -> Result<ModelResponseEventStream, ModelError> {
262 let cancellation_token = websocket_request.cancellation_token.clone();
263 let logical_body = websocket_request.body.clone();
264 websocket_request.body = self.prepare_websocket_request_body(&logical_body);
265 ProtocolModelClient::ensure_real_model_request_allowed(&websocket_request)?;
266 let diagnostic = transport_selected_diagnostic(
267 &self.client.provider_name,
268 &self.client.model_name,
269 "websocket",
270 );
271 if let Some(audit) = self.client.request_audit.as_ref() {
272 audit.record(
273 &self.client.provider_name,
274 &self.client.model_name,
275 true,
276 &websocket_request,
277 );
278 }
279 match send_websocket_session_event_stream_with_retries(
280 self.websocket_session.as_mut(),
281 self.client.sleeper.as_ref(),
282 websocket_request,
283 &self.client.http_config.retry_policy,
284 )
285 .await
286 {
287 Ok(events) => {
288 let fallback_plan = if fallback {
289 Some(self.required_fallback_plan(http_request)?)
290 } else {
291 None
292 };
293 self.last_websocket_request = Some(logical_body);
294 Ok(canonical_openai_response_stream_for_session(
295 events,
296 cancellation_token,
297 Some(diagnostic),
298 Arc::clone(&self.last_response),
299 request_settings,
300 fallback_plan,
301 ))
302 }
303 Err(error) if fallback => {
304 mark_last_response_failed(&self.last_response);
305 self.last_websocket_request = None;
306 let plan = self.required_fallback_plan(http_request)?;
307 Ok(openai_response_stream_from_websocket_setup_error(
308 error,
309 cancellation_token,
310 diagnostic,
311 plan,
312 ))
313 }
314 Err(error) => {
315 mark_last_response_failed(&self.last_response);
316 self.last_websocket_request = None;
317 Err(error)
318 }
319 }
320 }
321}
322
323#[async_trait]
324impl ModelRunSession for ProtocolModelClientRunSession<'_> {
325 async fn request_stream_incremental(
326 &mut self,
327 messages: Vec<ModelMessage>,
328 settings: Option<ModelSettings>,
329 params: ModelRequestParameters,
330 context: ModelRequestContext,
331 ) -> Result<ModelResponseEventStream, ModelError> {
332 let cancellation_token = context.cancellation_token();
333 if self.client.profile.protocol != ProtocolFamily::OpenAiResponses {
334 let response = self
335 .client
336 .request(messages, settings, params, context)
337 .await?;
338 let (sender, receiver) = tokio::sync::mpsc::channel(1);
339 let _ = sender
340 .send(Ok(ModelResponseStreamEvent::FinalResult(Box::new(
341 response,
342 ))))
343 .await;
344 return Ok(ModelResponseEventStream::new_with_cancellation(
345 receiver,
346 cancellation_token,
347 ));
348 }
349
350 let prepared = prepare_model_request(
351 messages,
352 self.client.default_settings.as_ref(),
353 settings,
354 params,
355 &self.client.profile,
356 );
357 let request_settings = prepared.settings.clone();
358 let wire_body = self.client.build_wire_body(
359 &prepared.normalized_messages,
360 prepared.settings.as_ref(),
361 &prepared.params,
362 )?;
363 let options =
364 self.client
365 .request_options(&context, prepared.settings.as_ref(), &prepared.params);
366 let mut transport = self
367 .client
368 .resolve_response_stream_transport(prepared.settings.as_ref());
369 if self.fallback_to_http.load(Ordering::Relaxed)
370 && matches!(
371 transport,
372 ResolvedResponseStreamTransport::WebSocketThenHttpOnRetryable
373 )
374 {
375 transport = ResolvedResponseStreamTransport::HttpOnly;
376 }
377 match transport {
378 ResolvedResponseStreamTransport::HttpOnly => {
379 let http_request = self.client.build_response_stream_request(
380 wire_body,
381 &options,
382 cancellation_token,
383 ResponseStreamRequestKind::HttpSse,
384 )?;
385 ProtocolModelClient::ensure_real_model_request_allowed(&http_request)?;
386 self.client
387 .openai_response_stream_from_request(
388 http_request,
389 ResponseStreamRequestKind::HttpSse,
390 )
391 .await
392 }
393 ResolvedResponseStreamTransport::WebSocketOnly => {
394 let websocket_request = self.client.build_response_stream_request(
395 wire_body,
396 &options,
397 cancellation_token,
398 ResponseStreamRequestKind::WebSocket,
399 )?;
400 self.request_websocket_stream(
401 websocket_request,
402 None,
403 request_settings,
404 false,
405 )
406 .await
407 }
408 ResolvedResponseStreamTransport::WebSocketThenHttpOnRetryable => {
409 let websocket_request = self.client.build_response_stream_request(
410 wire_body.clone(),
411 &options,
412 cancellation_token.clone(),
413 ResponseStreamRequestKind::WebSocket,
414 )?;
415 let http_request = self.client.build_response_stream_request(
416 wire_body,
417 &options,
418 cancellation_token,
419 ResponseStreamRequestKind::HttpSse,
420 )?;
421 self.request_websocket_stream(
422 websocket_request,
423 Some(http_request),
424 request_settings,
425 true,
426 )
427 .await
428 }
429 }
430 }
431
432 async fn close(&mut self) {
433 self.websocket_session.reset().await;
434 self.last_websocket_request = None;
435 mark_last_response_failed(&self.last_response);
436 }
437}
438
439impl ProtocolModelClient {
440 fn resolve_response_stream_transport(
441 &self,
442 settings: Option<&ModelSettings>,
443 ) -> ResolvedResponseStreamTransport {
444 let configured_transport = settings
445 .and_then(|settings| settings.provider_settings.openai_responses.as_ref())
446 .and_then(|settings| settings.stream_transport);
447 match configured_transport {
448 Some(ResponseStreamTransport::WebSocket) => {
449 ResolvedResponseStreamTransport::WebSocketOnly
450 }
451 Some(ResponseStreamTransport::Auto) => {
452 ResolvedResponseStreamTransport::WebSocketThenHttpOnRetryable
453 }
454 None if self.is_codex_oauth_provider() => {
455 ResolvedResponseStreamTransport::WebSocketThenHttpOnRetryable
456 }
457 Some(ResponseStreamTransport::Http) | None => ResolvedResponseStreamTransport::HttpOnly,
458 }
459 }
460
461 fn is_codex_oauth_provider(&self) -> bool {
462 self.provider_name == "codex"
463 || self
464 .http_config
465 .metadata
466 .get("oauth_provider")
467 .and_then(Value::as_str)
468 .is_some_and(|provider| provider == "codex")
469 }
470
471 fn build_response_stream_request(
472 &self,
473 wire_body: Value,
474 options: &crate::transport::HttpRequestOptions,
475 cancellation_token: starweaver_core::CancellationToken,
476 kind: ResponseStreamRequestKind,
477 ) -> Result<HttpRequest, ModelError> {
478 let body = match kind {
479 ResponseStreamRequestKind::HttpSse => response_http_sse_body(wire_body),
480 ResponseStreamRequestKind::WebSocket => response_websocket_body(wire_body),
481 };
482 let mut request = build_http_request(&self.http_config, options, body);
483 request.cancellation_token = cancellation_token;
484 match kind {
485 ResponseStreamRequestKind::HttpSse => {
486 request.metadata.insert(
487 "starweaver.response_stream_transport".to_string(),
488 json!("http"),
489 );
490 }
491 ResponseStreamRequestKind::WebSocket => {
492 request.metadata.insert(
493 "starweaver.response_stream_transport".to_string(),
494 json!("websocket"),
495 );
496 }
497 }
498 self.finalize_http_request(&mut request)?;
499 Ok(request)
500 }
501
502 fn ensure_real_model_request_allowed(request: &HttpRequest) -> Result<(), ModelError> {
503 if allow_real_model_requests() {
504 Ok(())
505 } else {
506 Err(ModelError::RealModelRequestBlocked {
507 url: request.url.clone(),
508 })
509 }
510 }
511
512 fn record_stream_request_audit(&self, request: &HttpRequest) {
513 if let Some(audit) = self.request_audit.as_ref() {
514 audit.record(&self.provider_name, &self.model_name, true, request);
515 }
516 }
517
518 async fn openai_response_stream_from_request(
519 &self,
520 request: HttpRequest,
521 kind: ResponseStreamRequestKind,
522 ) -> Result<ModelResponseEventStream, ModelError> {
523 let cancellation_token = request.cancellation_token.clone();
524 let diagnostic = matches!(kind, ResponseStreamRequestKind::WebSocket).then(|| {
525 transport_selected_diagnostic(
526 &self.provider_name,
527 &self.model_name,
528 transport_name(kind),
529 )
530 });
531 self.record_stream_request_audit(&request);
532 let events = match kind {
533 ResponseStreamRequestKind::HttpSse => {
534 send_event_stream_with_retries(
535 self.http_client.as_ref(),
536 self.sleeper.as_ref(),
537 request,
538 &self.http_config.retry_policy,
539 )
540 .await?
541 }
542 ResponseStreamRequestKind::WebSocket => {
543 send_websocket_event_stream_with_retries(
544 self.http_client.as_ref(),
545 self.sleeper.as_ref(),
546 request,
547 &self.http_config.retry_policy,
548 )
549 .await?
550 }
551 };
552 Ok(canonical_openai_response_stream(
553 events,
554 cancellation_token,
555 diagnostic,
556 ))
557 }
558}
559
560fn response_http_sse_body(mut body: Value) -> Value {
561 if let Some(object) = body.as_object_mut() {
562 object.insert("stream".to_string(), Value::Bool(true));
563 }
564 body
565}
566
567fn response_websocket_body(body: Value) -> Value {
568 let mut envelope = serde_json::Map::new();
569 if let Value::Object(mut object) = body {
570 object.remove("background");
571 envelope.extend(object);
572 }
573 envelope.insert(
574 "type".to_string(),
575 Value::String("response.create".to_string()),
576 );
577 envelope.insert("stream".to_string(), Value::Bool(true));
578 Value::Object(envelope)
579}
580
581const fn transport_name(kind: ResponseStreamRequestKind) -> &'static str {
582 match kind {
583 ResponseStreamRequestKind::HttpSse => "http",
584 ResponseStreamRequestKind::WebSocket => "websocket",
585 }
586}
587
588fn transport_selected_diagnostic(
589 provider_name: &str,
590 model_name: &str,
591 transport: &str,
592) -> StreamDiagnostic {
593 StreamDiagnostic::new(
594 "model_transport_selected",
595 json!({
596 "provider": provider_name,
597 "model": model_name,
598 "transport": transport,
599 "message": format!("model transport: {transport}"),
600 }),
601 )
602}
603
604fn transport_fallback_diagnostic(
605 provider_name: &str,
606 model_name: &str,
607 error: &ModelError,
608) -> StreamDiagnostic {
609 StreamDiagnostic::new(
610 "model_transport_fallback",
611 json!({
612 "provider": provider_name,
613 "model": model_name,
614 "from": "websocket",
615 "to": "http",
616 "reason": transport_fallback_reason(error),
617 "detail": error.to_string(),
618 "message": format!(
619 "model transport: websocket -> http fallback ({})",
620 transport_fallback_reason(error)
621 ),
622 }),
623 )
624}
625
626fn transport_fallback_reason(error: &ModelError) -> &'static str {
627 match error {
628 ModelError::ProviderStatus { body, .. } if websocket_connection_limit_reached(body) => {
629 "websocket_connection_limit_reached"
630 }
631 ModelError::ProviderStatus { .. } => "provider_status",
632 ModelError::RetryExhausted { source, .. } => transport_fallback_reason(source),
633 ModelError::Transport(_) => "websocket_transport_error",
634 ModelError::Cancelled { .. } => "cancelled",
635 ModelError::MessageMapping(_)
636 | ModelError::ResponseParsing(_)
637 | ModelError::RealModelRequestBlocked { .. }
638 | ModelError::UnsupportedResponse(_) => "model_error",
639 }
640}
641
642fn websocket_connection_limit_reached(body: &Value) -> bool {
643 body.get("error")
644 .and_then(|error| error.get("code"))
645 .and_then(Value::as_str)
646 .is_some_and(|code| code == "websocket_connection_limit_reached")
647 || body
648 .get("code")
649 .and_then(Value::as_str)
650 .is_some_and(|code| code == "websocket_connection_limit_reached")
651}
652
653fn canonical_openai_response_stream(
654 events: ModelEventStream,
655 cancellation_token: starweaver_core::CancellationToken,
656 diagnostic: Option<StreamDiagnostic>,
657) -> ModelResponseEventStream {
658 let drop_abort_token = events.drop_abort_token();
659 let (sender, receiver) = tokio::sync::mpsc::channel(32);
660 tokio::spawn(async move {
661 if let Some(diagnostic) = diagnostic
662 && sender
663 .send(Ok(ModelResponseStreamEvent::Diagnostic(diagnostic)))
664 .await
665 .is_err()
666 {
667 return;
668 }
669 let mut emitted_any_event = false;
670 if let Err(error) =
671 forward_openai_response_events(events, &sender, &mut emitted_any_event).await
672 {
673 let _ = sender.send(Err(error)).await;
674 }
675 });
676 ModelResponseEventStream::new_with_cancellation_and_drop_abort(
677 receiver,
678 cancellation_token,
679 drop_abort_token,
680 )
681}
682
683fn canonical_openai_response_stream_for_session(
684 events: ModelEventStream,
685 cancellation_token: starweaver_core::CancellationToken,
686 diagnostic: Option<StreamDiagnostic>,
687 last_response: Arc<Mutex<LastResponseSlot>>,
688 request_settings: Option<ModelSettings>,
689 fallback_plan: Option<SessionFallbackPlan>,
690) -> ModelResponseEventStream {
691 let drop_abort_token = events.drop_abort_token();
692 let (sender, receiver) = tokio::sync::mpsc::channel(32);
693 tokio::spawn(async move {
694 if let Some(diagnostic) = diagnostic
695 && sender
696 .send(Ok(ModelResponseStreamEvent::Diagnostic(diagnostic)))
697 .await
698 .is_err()
699 {
700 return;
701 }
702 let mut emitted_any_event = false;
703 let result = forward_openai_response_events_tracked(
704 events,
705 &sender,
706 &mut emitted_any_event,
707 &last_response,
708 request_settings.as_ref(),
709 )
710 .await;
711 match result {
712 Ok(()) => {}
713 Err(error) if !emitted_any_event => {
714 if let Some(plan) = fallback_plan {
715 if should_fallback_websocket_to_http(&error) {
716 forward_http_fallback(plan, &sender, &error, &mut emitted_any_event).await;
717 } else {
718 mark_last_response_failed(&last_response);
719 let _ = sender.send(Err(error)).await;
720 }
721 } else {
722 mark_last_response_failed(&last_response);
723 let _ = sender.send(Err(error)).await;
724 }
725 }
726 Err(error) => {
727 mark_last_response_failed(&last_response);
728 let _ = sender.send(Err(error)).await;
729 }
730 }
731 });
732 ModelResponseEventStream::new_with_cancellation_and_drop_abort(
733 receiver,
734 cancellation_token,
735 drop_abort_token,
736 )
737}
738
739fn openai_response_stream_from_websocket_setup_error(
740 error: ModelError,
741 cancellation_token: starweaver_core::CancellationToken,
742 diagnostic: StreamDiagnostic,
743 plan: SessionFallbackPlan,
744) -> ModelResponseEventStream {
745 let (sender, receiver) = tokio::sync::mpsc::channel(32);
746 tokio::spawn(async move {
747 if sender
748 .send(Ok(ModelResponseStreamEvent::Diagnostic(diagnostic)))
749 .await
750 .is_err()
751 {
752 return;
753 }
754 if should_fallback_websocket_to_http(&error) {
755 let mut emitted_any_event = false;
756 forward_http_fallback(plan, &sender, &error, &mut emitted_any_event).await;
757 } else {
758 mark_last_response_failed(&plan.last_response);
759 let _ = sender.send(Err(error)).await;
760 }
761 });
762 ModelResponseEventStream::new_with_cancellation(receiver, cancellation_token)
763}
764
765async fn forward_http_fallback(
766 plan: SessionFallbackPlan,
767 sender: &tokio::sync::mpsc::Sender<Result<ModelResponseStreamEvent, ModelError>>,
768 websocket_error: &ModelError,
769 emitted_any_event: &mut bool,
770) {
771 plan.fallback_to_http.store(true, Ordering::Relaxed);
772 mark_last_response_failed(&plan.last_response);
773 let _ = sender
774 .send(Ok(ModelResponseStreamEvent::Diagnostic(
775 transport_fallback_diagnostic(&plan.provider_name, &plan.model_name, websocket_error),
776 )))
777 .await;
778 if let Some(audit) = plan.request_audit.as_ref() {
779 audit.record(
780 &plan.provider_name,
781 &plan.model_name,
782 true,
783 &plan.http_request,
784 );
785 }
786 match send_event_stream_with_retries(
787 plan.http_client.as_ref(),
788 plan.sleeper.as_ref(),
789 plan.http_request,
790 &plan.retry_policy,
791 )
792 .await
793 {
794 Ok(events) => {
795 if let Err(error) =
796 forward_openai_response_events(events, sender, emitted_any_event).await
797 {
798 let _ = sender.send(Err(error)).await;
799 }
800 }
801 Err(error) => {
802 let _ = sender.send(Err(error)).await;
803 }
804 }
805}
806
807async fn forward_openai_response_events(
808 mut events: ModelEventStream,
809 sender: &tokio::sync::mpsc::Sender<Result<ModelResponseStreamEvent, ModelError>>,
810 emitted_any_event: &mut bool,
811) -> Result<(), ModelError> {
812 let mut parser = crate::providers::openai_responses::OpenAiResponsesStreamParser::default();
813 while let Some(event) = events.recv().await {
814 let event = event?;
815 let stream_events = parser.push_event(&event)?;
816 for stream_event in stream_events {
817 if sender.send(Ok(stream_event)).await.is_err() {
818 return Ok(());
819 }
820 *emitted_any_event = true;
821 }
822 }
823 for stream_event in parser.finish()? {
824 if sender.send(Ok(stream_event)).await.is_err() {
825 return Ok(());
826 }
827 *emitted_any_event = true;
828 }
829 Ok(())
830}
831
832async fn forward_openai_response_events_tracked(
833 mut events: ModelEventStream,
834 sender: &tokio::sync::mpsc::Sender<Result<ModelResponseStreamEvent, ModelError>>,
835 emitted_any_event: &mut bool,
836 last_response: &Arc<Mutex<LastResponseSlot>>,
837 request_settings: Option<&ModelSettings>,
838) -> Result<(), ModelError> {
839 let mut parser = crate::providers::openai_responses::OpenAiResponsesStreamParser::default();
840 while let Some(event) = events.recv().await {
841 let event = event?;
842 let stream_events = parser.push_event(&event)?;
843 for stream_event in stream_events {
844 if let ModelResponseStreamEvent::FinalResult(response) = &stream_event {
845 record_last_response(last_response, response, request_settings);
846 }
847 if sender.send(Ok(stream_event)).await.is_err() {
848 return Ok(());
849 }
850 *emitted_any_event = true;
851 }
852 }
853 for stream_event in parser.finish()? {
854 if let ModelResponseStreamEvent::FinalResult(response) = &stream_event {
855 record_last_response(last_response, response, request_settings);
856 }
857 if sender.send(Ok(stream_event)).await.is_err() {
858 return Ok(());
859 }
860 *emitted_any_event = true;
861 }
862 Ok(())
863}
864
865fn record_last_response(
866 slot: &Arc<Mutex<LastResponseSlot>>,
867 response: &crate::message::ModelResponse,
868 request_settings: Option<&ModelSettings>,
869) {
870 let response_id = response
871 .provider
872 .as_ref()
873 .and_then(|provider| provider.response_id.clone())
874 .unwrap_or_default();
875 let replay_items =
876 crate::providers::openai_responses::OpenAiResponsesAdapter::response_replay_items(
877 response,
878 request_settings,
879 );
880 let mut slot = slot
881 .lock()
882 .unwrap_or_else(std::sync::PoisonError::into_inner);
883 slot.response = Some(LastResponse {
884 response_id,
885 replay_items,
886 });
887 slot.failed = false;
888}
889
890fn mark_last_response_failed(slot: &Arc<Mutex<LastResponseSlot>>) {
891 let mut slot = slot
892 .lock()
893 .unwrap_or_else(std::sync::PoisonError::into_inner);
894 slot.response = None;
895 slot.failed = true;
896}
897
898fn response_input_items(value: &Value) -> Option<&[Value]> {
899 value
900 .get("input")
901 .and_then(Value::as_array)
902 .map(Vec::as_slice)
903}
904
905fn strip_value_prefix<'a>(items: &'a [Value], prefix: &[Value]) -> Option<&'a [Value]> {
906 if items.len() < prefix.len() || &items[..prefix.len()] != prefix {
907 return None;
908 }
909 Some(&items[prefix.len()..])
910}
911
912fn responses_websocket_request_properties_match(previous: &Value, current: &Value) -> bool {
913 let Some(previous) = previous.as_object() else {
914 return false;
915 };
916 let Some(current) = current.as_object() else {
917 return false;
918 };
919 let mut previous = previous.clone();
920 let mut current = current.clone();
921 previous.remove("input");
922 current.remove("input");
923 previous == current
924}
925
926#[cfg(test)]
927mod tests {
928 use std::{
929 collections::VecDeque,
930 sync::{Arc, Mutex},
931 };
932
933 use async_trait::async_trait;
934 use serde_json::{Value, json};
935 use starweaver_core::{ConversationId, RunId};
936
937 use super::*;
938 use crate::{
939 adapter::ModelAdapter,
940 message::{ModelRequest, ModelResponse},
941 profile::ProtocolFamily,
942 settings::{OpenAiResponsesSettings, ProviderReplaySettings, ProviderSettings},
943 transport::{
944 HttpModelConfig, HttpResponse, ModelHttpClient, ModelWebSocketEventSession, NoopSleeper,
945 },
946 };
947
948 #[test]
949 fn websocket_body_wraps_response_create_and_forces_stream_true() {
950 let body = response_websocket_body(json!({
951 "model": "gpt-5-codex",
952 "input": [{"role": "user", "content": "hello"}],
953 "stream": false,
954 "background": false,
955 "store": false
956 }));
957
958 assert_eq!(
959 body,
960 json!({
961 "type": "response.create",
962 "model": "gpt-5-codex",
963 "input": [{"role": "user", "content": "hello"}],
964 "stream": true,
965 "store": false
966 })
967 );
968 }
969
970 #[test]
971 fn resolves_transport_defaults_and_explicit_settings() {
972 let native = test_client("openai", None, Arc::new(FakeStreamClient::default()));
973 assert_eq!(
974 native.resolve_response_stream_transport(None),
975 ResolvedResponseStreamTransport::HttpOnly
976 );
977
978 let codex_by_name = test_client("codex", None, Arc::new(FakeStreamClient::default()));
979 assert_eq!(
980 codex_by_name.resolve_response_stream_transport(None),
981 ResolvedResponseStreamTransport::WebSocketThenHttpOnRetryable
982 );
983
984 let codex_by_metadata = test_client(
985 "openai",
986 Some(json!({"oauth_provider": "codex"})),
987 Arc::new(FakeStreamClient::default()),
988 );
989 assert_eq!(
990 codex_by_metadata.resolve_response_stream_transport(None),
991 ResolvedResponseStreamTransport::WebSocketThenHttpOnRetryable
992 );
993
994 let explicit_http = settings_with_transport(ResponseStreamTransport::Http);
995 assert_eq!(
996 codex_by_name.resolve_response_stream_transport(Some(&explicit_http)),
997 ResolvedResponseStreamTransport::HttpOnly
998 );
999
1000 let explicit_websocket = settings_with_transport(ResponseStreamTransport::WebSocket);
1001 assert_eq!(
1002 native.resolve_response_stream_transport(Some(&explicit_websocket)),
1003 ResolvedResponseStreamTransport::WebSocketOnly
1004 );
1005
1006 let explicit_auto = settings_with_transport(ResponseStreamTransport::Auto);
1007 assert_eq!(
1008 native.resolve_response_stream_transport(Some(&explicit_auto)),
1009 ResolvedResponseStreamTransport::WebSocketThenHttpOnRetryable
1010 );
1011 }
1012
1013 #[tokio::test]
1014 async fn auto_transport_falls_back_to_http_for_pre_event_retryable_websocket_error() {
1015 let fake = Arc::new(FakeStreamClient::new(
1016 WebSocketBehavior::ImmediateConnectionLimit,
1017 vec![completed_text_event("from http")],
1018 ));
1019 let client = test_client("codex", None, fake.clone());
1020
1021 let stream = client
1022 .request_stream_incremental(
1023 vec![ModelMessage::Request(ModelRequest::user_text("hello"))],
1024 None,
1025 ModelRequestParameters::default(),
1026 test_context(),
1027 )
1028 .await;
1029 let mut stream = result_or_panic(stream, "stream should be created");
1030
1031 let mut final_text = None;
1032 while let Some(event) = stream.recv().await {
1033 if let ModelResponseStreamEvent::FinalResult(response) =
1034 result_or_panic(event, "event should parse")
1035 {
1036 final_text = Some(response.text_output());
1037 }
1038 }
1039
1040 assert_eq!(final_text.as_deref(), Some("from http"));
1041 assert_eq!(
1042 fake.calls(),
1043 vec![
1044 FakeCallKind::WebSocket,
1045 FakeCallKind::WebSocket,
1046 FakeCallKind::WebSocket,
1047 FakeCallKind::WebSocket,
1048 FakeCallKind::WebSocket,
1049 FakeCallKind::Http
1050 ]
1051 );
1052 let bodies = fake.bodies();
1053 assert_eq!(
1054 bodies[0].get("type").and_then(Value::as_str),
1055 Some("response.create")
1056 );
1057 assert_eq!(bodies[1].get("stream"), Some(&Value::Bool(true)));
1058 }
1059
1060 #[tokio::test]
1061 async fn explicit_websocket_transport_does_not_fallback() {
1062 let fake = Arc::new(FakeStreamClient::new(
1063 WebSocketBehavior::ImmediateConnectionLimit,
1064 vec![completed_text_event("from http")],
1065 ));
1066 let client = test_client("codex", None, fake.clone());
1067 let settings = settings_with_transport(ResponseStreamTransport::WebSocket);
1068
1069 let result = client
1070 .request_stream_incremental(
1071 vec![ModelMessage::Request(ModelRequest::user_text("hello"))],
1072 Some(settings),
1073 ModelRequestParameters::default(),
1074 test_context(),
1075 )
1076 .await;
1077
1078 assert!(matches!(
1079 result,
1080 Err(ModelError::RetryExhausted {
1081 attempts: 5,
1082 source
1083 }) if matches!(
1084 source.as_ref(),
1085 ModelError::ProviderStatus {
1086 status: 400,
1087 retryable: true,
1088 ..
1089 }
1090 )
1091 ));
1092 assert_eq!(
1093 fake.calls(),
1094 vec![
1095 FakeCallKind::WebSocket,
1096 FakeCallKind::WebSocket,
1097 FakeCallKind::WebSocket,
1098 FakeCallKind::WebSocket,
1099 FakeCallKind::WebSocket
1100 ]
1101 );
1102 }
1103
1104 #[tokio::test]
1105 async fn websocket_error_after_canonical_event_is_not_fallback_safe() {
1106 let fake = Arc::new(FakeStreamClient::new(
1107 WebSocketBehavior::TextThenConnectionLimit,
1108 vec![completed_text_event("from http")],
1109 ));
1110 let client = test_client("codex", None, fake.clone());
1111
1112 let stream = client
1113 .request_stream_incremental(
1114 vec![ModelMessage::Request(ModelRequest::user_text("hello"))],
1115 None,
1116 ModelRequestParameters::default(),
1117 test_context(),
1118 )
1119 .await;
1120 let mut stream = result_or_panic(stream, "stream should be created");
1121
1122 let first = next_non_diagnostic(&mut stream).await;
1123 assert!(matches!(first, ModelResponseStreamEvent::PartStart(_)));
1124 let second = next_non_diagnostic(&mut stream).await;
1125 assert!(matches!(
1126 second,
1127 ModelResponseStreamEvent::PartDelta(crate::PartDelta {
1128 delta: crate::StreamDelta::Text { .. },
1129 ..
1130 })
1131 ));
1132 let error = option_or_panic(stream.recv().await, "expected websocket error");
1133 assert!(matches!(
1134 error,
1135 Err(ModelError::ProviderStatus {
1136 status: 400,
1137 retryable: true,
1138 ..
1139 })
1140 ));
1141 assert_eq!(fake.calls(), vec![FakeCallKind::WebSocket]);
1142 }
1143
1144 #[tokio::test]
1145 async fn run_session_reuses_websocket_session_and_sends_incremental_create() {
1146 let fake = Arc::new(SessionFakeClient::new(
1147 vec![
1148 Ok(vec![Ok(completed_text_event_with_ids(
1149 "resp_1",
1150 "msg_1",
1151 "assistant output",
1152 ))]),
1153 Ok(vec![Ok(completed_text_event_with_ids(
1154 "resp_2", "msg_2", "done",
1155 ))]),
1156 ],
1157 Vec::new(),
1158 ));
1159 let client = test_client_with_session_fake("codex", fake.clone());
1160 let mut session = client.start_run_session();
1161 let settings = settings_with_transport(ResponseStreamTransport::WebSocket);
1162
1163 let first_response = final_response_from_stream(
1164 session
1165 .request_stream_incremental(
1166 vec![ModelMessage::Request(ModelRequest::user_text("hello"))],
1167 Some(settings.clone()),
1168 ModelRequestParameters::default(),
1169 test_context(),
1170 )
1171 .await,
1172 )
1173 .await;
1174 let second_response = final_response_from_stream(
1175 session
1176 .request_stream_incremental(
1177 vec![
1178 ModelMessage::Request(ModelRequest::user_text("hello")),
1179 ModelMessage::Response(first_response),
1180 ModelMessage::Request(ModelRequest::user_text("second")),
1181 ],
1182 Some(settings),
1183 ModelRequestParameters::default(),
1184 test_context(),
1185 )
1186 .await,
1187 )
1188 .await;
1189
1190 assert_eq!(second_response.text_output(), "done");
1191 assert_eq!(fake.websocket_sessions(), 1);
1192 assert_eq!(
1193 fake.calls(),
1194 vec![FakeCallKind::WebSocket, FakeCallKind::WebSocket]
1195 );
1196 let bodies = fake.bodies();
1197 assert_eq!(bodies.len(), 2);
1198 assert_eq!(
1199 bodies[1]
1200 .get("previous_response_id")
1201 .and_then(Value::as_str),
1202 Some("resp_1")
1203 );
1204 let second_input = option_or_panic(
1205 bodies[1].get("input").and_then(Value::as_array),
1206 "second request input should be an array",
1207 );
1208 assert_eq!(second_input.len(), 1);
1209 assert_eq!(
1210 second_input[0].get("role").and_then(Value::as_str),
1211 Some("user")
1212 );
1213 }
1214
1215 #[tokio::test]
1216 async fn run_session_reuses_websocket_but_full_creates_on_non_prefix_input() {
1217 let fake = Arc::new(SessionFakeClient::new(
1218 vec![
1219 Ok(vec![Ok(completed_text_event_with_ids(
1220 "resp_1",
1221 "msg_1",
1222 "assistant output",
1223 ))]),
1224 Ok(vec![Ok(completed_text_event_with_ids(
1225 "resp_2", "msg_2", "done",
1226 ))]),
1227 ],
1228 Vec::new(),
1229 ));
1230 let client = test_client_with_session_fake("codex", fake.clone());
1231 let mut session = client.start_run_session();
1232 let settings = settings_with_transport(ResponseStreamTransport::WebSocket);
1233
1234 let _ = final_response_from_stream(
1235 session
1236 .request_stream_incremental(
1237 vec![ModelMessage::Request(ModelRequest::user_text("hello"))],
1238 Some(settings.clone()),
1239 ModelRequestParameters::default(),
1240 test_context(),
1241 )
1242 .await,
1243 )
1244 .await;
1245 let _ = final_response_from_stream(
1246 session
1247 .request_stream_incremental(
1248 vec![ModelMessage::Request(ModelRequest::user_text("different"))],
1249 Some(settings),
1250 ModelRequestParameters::default(),
1251 test_context(),
1252 )
1253 .await,
1254 )
1255 .await;
1256
1257 assert_eq!(fake.websocket_sessions(), 1);
1258 let bodies = fake.bodies();
1259 assert_eq!(bodies.len(), 2);
1260 assert!(bodies[1].get("previous_response_id").is_none());
1261 assert_eq!(
1262 bodies[1]
1263 .get("input")
1264 .and_then(Value::as_array)
1265 .map(Vec::len),
1266 Some(1)
1267 );
1268 assert_eq!(
1269 bodies[1]["input"][0]
1270 .get("content")
1271 .and_then(Value::as_array)
1272 .and_then(|content| content.first())
1273 .and_then(|content| content.get("text"))
1274 .and_then(Value::as_str),
1275 Some("different")
1276 );
1277 }
1278
1279 #[tokio::test]
1280 async fn run_session_auto_fallback_uses_http_for_remaining_requests() {
1281 let fake = Arc::new(SessionFakeClient::new(
1282 vec![
1283 Err(connection_limit_error()),
1284 Err(connection_limit_error()),
1285 Err(connection_limit_error()),
1286 Err(connection_limit_error()),
1287 Err(connection_limit_error()),
1288 ],
1289 vec![
1290 vec![Ok(completed_text_event_with_ids(
1291 "resp_http_1",
1292 "msg_http_1",
1293 "from http one",
1294 ))],
1295 vec![Ok(completed_text_event_with_ids(
1296 "resp_http_2",
1297 "msg_http_2",
1298 "from http two",
1299 ))],
1300 ],
1301 ));
1302 let client = test_client_with_session_fake("codex", fake.clone());
1303 let mut session = client.start_run_session();
1304
1305 let first_response = final_response_from_stream(
1306 session
1307 .request_stream_incremental(
1308 vec![ModelMessage::Request(ModelRequest::user_text("hello"))],
1309 None,
1310 ModelRequestParameters::default(),
1311 test_context(),
1312 )
1313 .await,
1314 )
1315 .await;
1316 let second_response = final_response_from_stream(
1317 session
1318 .request_stream_incremental(
1319 vec![
1320 ModelMessage::Request(ModelRequest::user_text("hello")),
1321 ModelMessage::Response(first_response),
1322 ModelMessage::Request(ModelRequest::user_text("second")),
1323 ],
1324 None,
1325 ModelRequestParameters::default(),
1326 test_context(),
1327 )
1328 .await,
1329 )
1330 .await;
1331
1332 assert_eq!(second_response.text_output(), "from http two");
1333 assert_eq!(
1334 fake.calls(),
1335 vec![
1336 FakeCallKind::WebSocket,
1337 FakeCallKind::WebSocket,
1338 FakeCallKind::WebSocket,
1339 FakeCallKind::WebSocket,
1340 FakeCallKind::WebSocket,
1341 FakeCallKind::Http,
1342 FakeCallKind::Http
1343 ]
1344 );
1345 }
1346
1347 #[tokio::test]
1348 async fn run_session_does_not_delta_when_request_uses_conversation_state() {
1349 let fake = Arc::new(SessionFakeClient::new(
1350 vec![
1351 Ok(vec![Ok(completed_text_event_with_ids(
1352 "resp_1",
1353 "msg_1",
1354 "assistant output",
1355 ))]),
1356 Ok(vec![Ok(completed_text_event_with_ids(
1357 "resp_2", "msg_2", "done",
1358 ))]),
1359 ],
1360 Vec::new(),
1361 ));
1362 let client = test_client_with_session_fake("codex", fake.clone());
1363 let mut session = client.start_run_session();
1364 let websocket_settings = settings_with_transport(ResponseStreamTransport::WebSocket);
1365
1366 let first_response = final_response_from_stream(
1367 session
1368 .request_stream_incremental(
1369 vec![ModelMessage::Request(ModelRequest::user_text("hello"))],
1370 Some(websocket_settings),
1371 ModelRequestParameters::default(),
1372 test_context(),
1373 )
1374 .await,
1375 )
1376 .await;
1377 let mut conversation_settings = settings_with_transport(ResponseStreamTransport::WebSocket);
1378 conversation_settings.provider_replay = Some(ProviderReplaySettings {
1379 conversation_id: Some("conv_manual".to_string()),
1380 ..ProviderReplaySettings::default()
1381 });
1382 let _ = final_response_from_stream(
1383 session
1384 .request_stream_incremental(
1385 vec![
1386 ModelMessage::Request(ModelRequest::user_text("hello")),
1387 ModelMessage::Response(first_response),
1388 ModelMessage::Request(ModelRequest::user_text("second")),
1389 ],
1390 Some(conversation_settings),
1391 ModelRequestParameters::default(),
1392 test_context(),
1393 )
1394 .await,
1395 )
1396 .await;
1397
1398 let bodies = fake.bodies();
1399 assert_eq!(bodies.len(), 2);
1400 assert_eq!(bodies[1].get("conversation"), Some(&json!("conv_manual")));
1401 assert!(bodies[1].get("previous_response_id").is_none());
1402 assert_eq!(
1403 bodies[1]
1404 .get("input")
1405 .and_then(Value::as_array)
1406 .map(Vec::len),
1407 Some(3)
1408 );
1409 }
1410
1411 #[tokio::test]
1412 async fn run_session_drops_incremental_state_after_stream_error() {
1413 let fake = Arc::new(SessionFakeClient::new(
1414 vec![
1415 Ok(vec![
1416 Ok(text_delta_event("partial")),
1417 Err(connection_limit_error()),
1418 ]),
1419 Ok(vec![Ok(completed_text_event_with_ids(
1420 "resp_2", "msg_2", "done",
1421 ))]),
1422 ],
1423 Vec::new(),
1424 ));
1425 let client = test_client_with_session_fake("codex", fake.clone());
1426 let mut session = client.start_run_session();
1427 let settings = settings_with_transport(ResponseStreamTransport::WebSocket);
1428
1429 let mut first_stream = result_or_panic(
1430 session
1431 .request_stream_incremental(
1432 vec![ModelMessage::Request(ModelRequest::user_text("hello"))],
1433 Some(settings.clone()),
1434 ModelRequestParameters::default(),
1435 test_context(),
1436 )
1437 .await,
1438 "first stream should be created",
1439 );
1440 let mut saw_error = false;
1441 while let Some(event) = first_stream.recv().await {
1442 if matches!(
1443 event,
1444 Err(ModelError::ProviderStatus {
1445 status: 400,
1446 retryable: true,
1447 ..
1448 })
1449 ) {
1450 saw_error = true;
1451 break;
1452 }
1453 }
1454 assert!(saw_error);
1455
1456 let second_response = final_response_from_stream(
1457 session
1458 .request_stream_incremental(
1459 vec![ModelMessage::Request(ModelRequest::user_text("next"))],
1460 Some(settings),
1461 ModelRequestParameters::default(),
1462 test_context(),
1463 )
1464 .await,
1465 )
1466 .await;
1467
1468 assert_eq!(second_response.text_output(), "done");
1469 let bodies = fake.bodies();
1470 assert_eq!(bodies.len(), 2);
1471 assert!(bodies[1].get("previous_response_id").is_none());
1472 assert_eq!(
1473 bodies[1]["input"][0]
1474 .get("content")
1475 .and_then(Value::as_array)
1476 .and_then(|content| content.first())
1477 .and_then(|content| content.get("text"))
1478 .and_then(Value::as_str),
1479 Some("next")
1480 );
1481 }
1482
1483 fn test_client(
1484 provider_name: &str,
1485 metadata: Option<Value>,
1486 http_client: Arc<FakeStreamClient>,
1487 ) -> ProtocolModelClient {
1488 let mut config = HttpModelConfig::new("https://api.openai.com/v1", "responses");
1489 if let Some(Value::Object(metadata)) = metadata {
1490 config.metadata = metadata;
1491 }
1492 ProtocolModelClient::new(
1493 provider_name,
1494 "gpt-5-codex",
1495 ModelProfile::for_protocol(ProtocolFamily::OpenAiResponses),
1496 config,
1497 http_client,
1498 )
1499 .with_sleeper(Arc::new(NoopSleeper))
1500 }
1501
1502 fn test_client_with_session_fake(
1503 provider_name: &str,
1504 http_client: Arc<SessionFakeClient>,
1505 ) -> ProtocolModelClient {
1506 ProtocolModelClient::new(
1507 provider_name,
1508 "gpt-5-codex",
1509 ModelProfile::for_protocol(ProtocolFamily::OpenAiResponses),
1510 HttpModelConfig::new("https://api.openai.com/v1", "responses"),
1511 http_client,
1512 )
1513 .with_sleeper(Arc::new(NoopSleeper))
1514 }
1515
1516 fn settings_with_transport(transport: ResponseStreamTransport) -> ModelSettings {
1517 ModelSettings {
1518 provider_settings: ProviderSettings {
1519 openai_responses: Some(OpenAiResponsesSettings {
1520 stream_transport: Some(transport),
1521 ..OpenAiResponsesSettings::default()
1522 }),
1523 ..ProviderSettings::default()
1524 },
1525 ..ModelSettings::default()
1526 }
1527 }
1528
1529 fn test_context() -> ModelRequestContext {
1530 ModelRequestContext::new(RunId::new(), ConversationId::new())
1531 }
1532
1533 fn completed_text_event(text: &str) -> Value {
1534 completed_text_event_with_ids("resp_test", "msg_test", text)
1535 }
1536
1537 fn completed_text_event_with_ids(response_id: &str, message_id: &str, text: &str) -> Value {
1538 json!({
1539 "type": "response.completed",
1540 "response": {
1541 "id": response_id,
1542 "model": "gpt-5-codex",
1543 "status": "completed",
1544 "output": [{
1545 "id": message_id,
1546 "type": "message",
1547 "role": "assistant",
1548 "status": "completed",
1549 "content": [{"type": "output_text", "text": text}]
1550 }],
1551 "usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}
1552 }
1553 })
1554 }
1555
1556 async fn final_response_from_stream(
1557 stream: Result<ModelResponseEventStream, ModelError>,
1558 ) -> ModelResponse {
1559 let mut stream = result_or_panic(stream, "stream should be created");
1560 while let Some(event) = stream.recv().await {
1561 if let ModelResponseStreamEvent::FinalResult(response) =
1562 result_or_panic(event, "event should parse")
1563 {
1564 return *response;
1565 }
1566 }
1567 panic!("stream ended without final response")
1568 }
1569
1570 fn text_delta_event(text: &str) -> Value {
1571 json!({
1572 "type": "response.output_text.delta",
1573 "delta": text
1574 })
1575 }
1576
1577 fn connection_limit_error() -> ModelError {
1578 ModelError::ProviderStatus {
1579 status: 400,
1580 body: json!({"error": {"code": "websocket_connection_limit_reached"}}),
1581 retryable: true,
1582 }
1583 }
1584
1585 #[derive(Clone, Copy, Debug, Eq, PartialEq)]
1586 enum FakeCallKind {
1587 Http,
1588 WebSocket,
1589 }
1590
1591 #[derive(Clone, Copy, Debug, Eq, PartialEq)]
1592 enum WebSocketBehavior {
1593 ImmediateConnectionLimit,
1594 TextThenConnectionLimit,
1595 }
1596
1597 #[derive(Debug)]
1598 struct FakeStreamClient {
1599 websocket_behavior: WebSocketBehavior,
1600 http_events: Vec<Value>,
1601 calls: Mutex<Vec<FakeCallKind>>,
1602 bodies: Mutex<Vec<Value>>,
1603 }
1604
1605 impl Default for FakeStreamClient {
1606 fn default() -> Self {
1607 Self::new(WebSocketBehavior::ImmediateConnectionLimit, Vec::new())
1608 }
1609 }
1610
1611 impl FakeStreamClient {
1612 fn new(websocket_behavior: WebSocketBehavior, http_events: Vec<Value>) -> Self {
1613 Self {
1614 websocket_behavior,
1615 http_events,
1616 calls: Mutex::new(Vec::new()),
1617 bodies: Mutex::new(Vec::new()),
1618 }
1619 }
1620
1621 fn calls(&self) -> Vec<FakeCallKind> {
1622 lock_or_panic(self.calls.lock(), "calls lock should not be poisoned").clone()
1623 }
1624
1625 fn bodies(&self) -> Vec<Value> {
1626 lock_or_panic(self.bodies.lock(), "bodies lock should not be poisoned").clone()
1627 }
1628
1629 fn record(&self, kind: FakeCallKind, request: &HttpRequest) {
1630 lock_or_panic(self.calls.lock(), "calls lock should not be poisoned").push(kind);
1631 lock_or_panic(self.bodies.lock(), "bodies lock should not be poisoned")
1632 .push(request.body.clone());
1633 }
1634 }
1635
1636 type SessionFakeWebSocketEvents = Vec<Result<Value, ModelError>>;
1637 type SessionFakeWebSocketResult = Result<SessionFakeWebSocketEvents, ModelError>;
1638 type SessionFakeHttpEvents = Vec<Result<Value, ModelError>>;
1639
1640 #[derive(Debug)]
1641 struct SessionFakeClient {
1642 websocket_results: Mutex<VecDeque<SessionFakeWebSocketResult>>,
1643 http_results: Mutex<VecDeque<SessionFakeHttpEvents>>,
1644 calls: Mutex<Vec<FakeCallKind>>,
1645 bodies: Mutex<Vec<Value>>,
1646 websocket_sessions: Mutex<usize>,
1647 }
1648
1649 impl SessionFakeClient {
1650 fn new(
1651 websocket_results: Vec<SessionFakeWebSocketResult>,
1652 http_results: Vec<SessionFakeHttpEvents>,
1653 ) -> Self {
1654 Self {
1655 websocket_results: Mutex::new(VecDeque::from(websocket_results)),
1656 http_results: Mutex::new(VecDeque::from(http_results)),
1657 calls: Mutex::new(Vec::new()),
1658 bodies: Mutex::new(Vec::new()),
1659 websocket_sessions: Mutex::new(0),
1660 }
1661 }
1662
1663 fn calls(&self) -> Vec<FakeCallKind> {
1664 lock_or_panic(self.calls.lock(), "calls lock should not be poisoned").clone()
1665 }
1666
1667 fn bodies(&self) -> Vec<Value> {
1668 lock_or_panic(self.bodies.lock(), "bodies lock should not be poisoned").clone()
1669 }
1670
1671 fn websocket_sessions(&self) -> usize {
1672 *lock_or_panic(
1673 self.websocket_sessions.lock(),
1674 "websocket sessions lock should not be poisoned",
1675 )
1676 }
1677
1678 fn record(&self, kind: FakeCallKind, request: &HttpRequest) {
1679 lock_or_panic(self.calls.lock(), "calls lock should not be poisoned").push(kind);
1680 lock_or_panic(self.bodies.lock(), "bodies lock should not be poisoned")
1681 .push(request.body.clone());
1682 }
1683 }
1684
1685 #[async_trait]
1686 impl ModelHttpClient for SessionFakeClient {
1687 async fn send(&self, _request: HttpRequest) -> Result<HttpResponse, ModelError> {
1688 Err(ModelError::Transport(
1689 "send is not used by session transport tests".to_string(),
1690 ))
1691 }
1692
1693 async fn send_event_stream_incremental(
1694 &self,
1695 request: HttpRequest,
1696 ) -> Result<ModelEventStream, ModelError> {
1697 self.record(FakeCallKind::Http, &request);
1698 let events = lock_or_panic(
1699 self.http_results.lock(),
1700 "http results lock should not be poisoned",
1701 )
1702 .pop_front()
1703 .unwrap_or_else(|| {
1704 vec![Err(ModelError::Transport(
1705 "missing fake HTTP response".to_string(),
1706 ))]
1707 });
1708 Ok(stream_from_results(events))
1709 }
1710
1711 async fn send_websocket_event_stream_incremental(
1712 &self,
1713 request: HttpRequest,
1714 ) -> Result<ModelEventStream, ModelError> {
1715 self.record(FakeCallKind::WebSocket, &request);
1716 let result = lock_or_panic(
1717 self.websocket_results.lock(),
1718 "websocket results lock should not be poisoned",
1719 )
1720 .pop_front()
1721 .unwrap_or_else(|| {
1722 Err(ModelError::Transport(
1723 "missing fake websocket response".to_string(),
1724 ))
1725 });
1726 result.map(stream_from_results)
1727 }
1728
1729 fn websocket_event_session(&self) -> Box<dyn ModelWebSocketEventSession + '_> {
1730 *lock_or_panic(
1731 self.websocket_sessions.lock(),
1732 "websocket sessions lock should not be poisoned",
1733 ) += 1;
1734 Box::new(SessionFakeWebSocketSession { client: self })
1735 }
1736 }
1737
1738 struct SessionFakeWebSocketSession<'a> {
1739 client: &'a SessionFakeClient,
1740 }
1741
1742 #[async_trait]
1743 impl ModelWebSocketEventSession for SessionFakeWebSocketSession<'_> {
1744 async fn send_websocket_event_stream_incremental(
1745 &mut self,
1746 request: HttpRequest,
1747 ) -> Result<ModelEventStream, ModelError> {
1748 self.client
1749 .send_websocket_event_stream_incremental(request)
1750 .await
1751 }
1752 }
1753
1754 #[async_trait]
1755 impl ModelHttpClient for FakeStreamClient {
1756 async fn send(&self, _request: HttpRequest) -> Result<HttpResponse, ModelError> {
1757 Err(ModelError::Transport(
1758 "send is not used by stream transport tests".to_string(),
1759 ))
1760 }
1761
1762 async fn send_event_stream_incremental(
1763 &self,
1764 request: HttpRequest,
1765 ) -> Result<ModelEventStream, ModelError> {
1766 self.record(FakeCallKind::Http, &request);
1767 Ok(stream_from_results(
1768 self.http_events.iter().cloned().map(Ok).collect(),
1769 ))
1770 }
1771
1772 async fn send_websocket_event_stream_incremental(
1773 &self,
1774 request: HttpRequest,
1775 ) -> Result<ModelEventStream, ModelError> {
1776 self.record(FakeCallKind::WebSocket, &request);
1777 match self.websocket_behavior {
1778 WebSocketBehavior::ImmediateConnectionLimit => Err(connection_limit_error()),
1779 WebSocketBehavior::TextThenConnectionLimit => Ok(stream_from_results(vec![
1780 Ok(text_delta_event("partial")),
1781 Err(connection_limit_error()),
1782 ])),
1783 }
1784 }
1785 }
1786
1787 fn result_or_panic<T, E: std::fmt::Debug>(result: Result<T, E>, message: &str) -> T {
1788 match result {
1789 Ok(value) => value,
1790 Err(error) => panic!("{message}: {error:?}"),
1791 }
1792 }
1793
1794 async fn next_non_diagnostic(
1795 stream: &mut ModelResponseEventStream,
1796 ) -> ModelResponseStreamEvent {
1797 while let Some(event) = stream.recv().await {
1798 let event = result_or_panic(event, "event should parse");
1799 if matches!(event, ModelResponseStreamEvent::Diagnostic(_)) {
1800 continue;
1801 }
1802 return event;
1803 }
1804 panic!("expected non-diagnostic stream event")
1805 }
1806
1807 fn option_or_panic<T>(value: Option<T>, message: &str) -> T {
1808 value.unwrap_or_else(|| panic!("{message}"))
1809 }
1810
1811 fn lock_or_panic<'a, T>(
1812 lock: std::sync::LockResult<std::sync::MutexGuard<'a, T>>,
1813 message: &str,
1814 ) -> std::sync::MutexGuard<'a, T> {
1815 lock.unwrap_or_else(|error| panic!("{message}: {error}"))
1816 }
1817
1818 fn stream_from_results(events: Vec<Result<Value, ModelError>>) -> ModelEventStream {
1819 let (sender, receiver) = tokio::sync::mpsc::channel(32);
1820 tokio::spawn(async move {
1821 for event in events {
1822 if sender.send(event).await.is_err() {
1823 break;
1824 }
1825 }
1826 });
1827 ModelEventStream::new(receiver)
1828 }
1829}