Skip to main content

starweaver_model/providers/client/
adapter_impl.rs

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                    /*fallback*/ 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                    /*fallback*/ 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}