Skip to main content

claude_codex/providers/codex/
native.rs

1use std::io;
2use std::pin::Pin;
3use std::sync::{Arc, Mutex};
4use std::time::Duration;
5
6use axum::body::Body;
7use axum::response::{IntoResponse, Response};
8use bytes::Bytes;
9use futures_util::{Stream, StreamExt};
10use http::{HeaderMap, HeaderName, StatusCode};
11use serde_json::{Map, Value, json};
12
13use crate::anthropic::sse::parse_sse_events;
14use crate::provider::RequestContext;
15use crate::traffic::{
16    MAX_SSE_CAPTURE_BYTES, MAX_STREAM_CAPTURE_EVENT_BYTES, MAX_STREAM_CAPTURE_EVENTS,
17    MAX_STREAM_CAPTURE_FRAME_BYTES,
18};
19
20use super::client::{CodexError, CodexHttpClient};
21use super::translate::model_allowlist::{
22    ALLOWED_MODELS, MODEL_ALIASES, assert_allowed_model, full_lane_web_search_model,
23    uses_responses_lite,
24};
25
26pub struct CodexNativeBackend {
27    client: Arc<CodexHttpClient>,
28}
29
30impl Default for CodexNativeBackend {
31    fn default() -> Self {
32        Self::new()
33    }
34}
35
36impl CodexNativeBackend {
37    pub fn new() -> Self {
38        Self {
39            client: Arc::new(CodexHttpClient::new()),
40        }
41    }
42
43    pub async fn handle(&self, mut body: Value, ctx: RequestContext) -> Response {
44        let resolved = match shape_native_request(&mut body) {
45            Ok(resolved) => resolved,
46            Err(response) => return response,
47        };
48        if let Some(monitor) = ctx.monitor.as_ref() {
49            monitor.model_resolved(&ctx.req_id, &resolved.model);
50            monitor.upstream_started(&ctx.req_id);
51        }
52
53        let upstream = match self
54            .client
55            .post_native_responses(&body, &ctx, resolved.use_responses_lite, resolved.stream)
56            .await
57        {
58            Ok(response) => response,
59            Err(error) => return local_codex_error(error),
60        };
61
62        passthrough_response(upstream, ctx, self.client.body_idle_timeout_ms())
63    }
64}
65
66struct NativeResolved {
67    model: String,
68    use_responses_lite: bool,
69    stream: bool,
70}
71
72#[allow(clippy::result_large_err)]
73pub fn validate_native_request_model(body: &Value) -> Result<String, Response> {
74    let object = body.as_object().ok_or_else(|| {
75        openai_error(
76            StatusCode::BAD_REQUEST,
77            "invalid_request_error",
78            "Request body must be a JSON object",
79            None,
80            None,
81        )
82    })?;
83    let requested = object
84        .get("model")
85        .and_then(Value::as_str)
86        .filter(|model| !model.is_empty())
87        .map(str::to_string)
88        .ok_or_else(|| {
89            openai_error(
90                StatusCode::BAD_REQUEST,
91                "invalid_request_error",
92                "Missing or invalid 'model' in request body",
93                Some("model"),
94                None,
95            )
96        })?;
97    let (resolved, _) = resolve_native_model(&requested);
98    if let Err(error) = assert_allowed_model(&resolved) {
99        return Err(openai_error(
100            StatusCode::BAD_REQUEST,
101            "invalid_request_error",
102            format!(
103                "Model '{requested}' resolves to unsupported model '{}'. Supported: {}",
104                error.model,
105                ALLOWED_MODELS.join(", ")
106            ),
107            Some("model"),
108            Some("model_not_supported"),
109        ));
110    }
111    Ok(requested)
112}
113
114#[allow(clippy::result_large_err)]
115fn shape_native_request(body: &mut Value) -> Result<NativeResolved, Response> {
116    let requested = validate_native_request_model(body)?;
117    let object = body
118        .as_object_mut()
119        .expect("validated native Responses body must be an object");
120
121    let (mut model, priority) = resolve_native_model(&requested);
122
123    let hosted_web_search = has_native_hosted_web_search(object);
124    if hosted_web_search {
125        model = full_lane_web_search_model(&model).to_string();
126    }
127    object.insert("model".to_string(), Value::String(model.clone()));
128    if priority && !object.contains_key("service_tier") {
129        object.insert("service_tier".to_string(), json!("priority"));
130    }
131
132    Ok(NativeResolved {
133        use_responses_lite: uses_responses_lite(&model) && !hosted_web_search,
134        model,
135        stream: object
136            .get("stream")
137            .and_then(Value::as_bool)
138            .unwrap_or(false),
139    })
140}
141
142fn resolve_native_model(requested: &str) -> (String, bool) {
143    let (requested, priority) = match requested.strip_suffix("-fast") {
144        Some(base) if ALLOWED_MODELS.contains(&base) => (base, true),
145        _ => (requested, false),
146    };
147    let model = MODEL_ALIASES
148        .iter()
149        .find(|(alias, _)| *alias == requested)
150        .map(|(_, target)| *target)
151        .unwrap_or(requested);
152    (model.to_string(), priority)
153}
154
155fn has_native_hosted_web_search(object: &Map<String, Value>) -> bool {
156    object
157        .get("tools")
158        .and_then(Value::as_array)
159        .is_some_and(|tools| {
160            tools.iter().any(|tool| {
161                matches!(
162                    tool.get("type").and_then(Value::as_str),
163                    Some("web_search" | "web_search_preview")
164                )
165            })
166        })
167}
168
169fn local_codex_error(error: CodexError) -> Response {
170    let status = match error.status {
171        401 => StatusCode::UNAUTHORIZED,
172        403 => StatusCode::FORBIDDEN,
173        429 => StatusCode::TOO_MANY_REQUESTS,
174        400..=599 => StatusCode::from_u16(error.status).unwrap_or(StatusCode::BAD_GATEWAY),
175        _ => StatusCode::BAD_GATEWAY,
176    };
177    let kind = match status {
178        StatusCode::UNAUTHORIZED => "authentication_error",
179        StatusCode::FORBIDDEN => "permission_error",
180        StatusCode::TOO_MANY_REQUESTS => "rate_limit_error",
181        _ => "api_error",
182    };
183    let message = error.detail.as_deref().unwrap_or(&error.message);
184    let response = openai_error(status, kind, message, None, None);
185    if let Some(retry_after) = error.retry_after
186        && let Ok(value) = http::HeaderValue::from_str(&retry_after)
187    {
188        let (mut parts, body) = response.into_parts();
189        parts.headers.insert(http::header::RETRY_AFTER, value);
190        return Response::from_parts(parts, body);
191    }
192    response
193}
194
195pub fn openai_error(
196    status: StatusCode,
197    kind: &str,
198    message: impl Into<String>,
199    param: Option<&str>,
200    code: Option<&str>,
201) -> Response {
202    (
203        status,
204        axum::Json(json!({
205            "error": {
206                "message": message.into(),
207                "type": kind,
208                "param": param,
209                "code": code,
210            }
211        })),
212    )
213        .into_response()
214}
215
216#[derive(Clone, Default)]
217pub struct NativeResponseOutcome {
218    failure: Arc<Mutex<Option<String>>>,
219}
220
221impl NativeResponseOutcome {
222    pub fn failure(&self) -> Option<String> {
223        self.failure.lock().ok().and_then(|failure| failure.clone())
224    }
225
226    pub(crate) fn fail(&self, message: String) {
227        if let Ok(mut failure) = self.failure.lock()
228            && failure.is_none()
229        {
230            *failure = Some(message);
231        }
232    }
233}
234
235fn passthrough_response(
236    upstream: reqwest::Response,
237    ctx: RequestContext,
238    body_idle_timeout_ms: u64,
239) -> Response {
240    let status = upstream.status();
241    let headers = passthrough_headers(upstream.headers());
242    let is_sse = upstream
243        .headers()
244        .get(http::header::CONTENT_TYPE)
245        .and_then(|value| value.to_str().ok())
246        .is_some_and(|value| value.starts_with("text/event-stream"));
247    let outcome = NativeResponseOutcome::default();
248    let observer = NativeResponseObserver::new(ctx, is_sse, outcome.clone());
249    let state = Some(NativeBodyState {
250        stream: Box::pin(upstream.bytes_stream()),
251        observer,
252        body_idle_timeout_ms,
253    });
254    let stream = futures_util::stream::unfold(state, |state| async move {
255        let mut state = state?;
256        match tokio::time::timeout(
257            Duration::from_millis(state.body_idle_timeout_ms),
258            state.stream.next(),
259        )
260        .await
261        {
262            Ok(Some(Ok(chunk))) => {
263                state.observer.observe(&chunk);
264                Some((Ok::<Bytes, io::Error>(chunk), Some(state)))
265            }
266            Ok(Some(Err(error))) => {
267                let message = format!("Native Responses body read failed: {error}");
268                state.observer.finish("read_error");
269                Some((Err(io::Error::other(message)), None))
270            }
271            Ok(None) => {
272                state.observer.finish("complete");
273                None
274            }
275            Err(_) => {
276                let message = format!(
277                    "Timed out waiting {}ms for the next Codex response body chunk",
278                    state.body_idle_timeout_ms
279                );
280                state.observer.finish("idle_timeout");
281                Some((Err(io::Error::new(io::ErrorKind::TimedOut, message)), None))
282            }
283        }
284    });
285
286    let mut response = Response::new(Body::from_stream(stream));
287    *response.status_mut() = status;
288    *response.headers_mut() = headers;
289    response.extensions_mut().insert(outcome);
290    response
291}
292
293fn passthrough_headers(upstream: &HeaderMap) -> HeaderMap {
294    let mut headers = HeaderMap::new();
295    for (name, value) in upstream {
296        if native_response_header_allowed(name) {
297            headers.append(name.clone(), value.clone());
298        }
299    }
300    headers
301}
302
303fn native_response_header_allowed(name: &HeaderName) -> bool {
304    matches!(
305        name.as_str(),
306        "content-type"
307            | "cache-control"
308            | "retry-after"
309            | "x-request-id"
310            | "openai-processing-ms"
311            | "openai-version"
312    ) || name.as_str().starts_with("x-ratelimit-")
313}
314
315type UpstreamByteStream =
316    Pin<Box<dyn Stream<Item = Result<Bytes, reqwest::Error>> + Send + 'static>>;
317
318struct NativeBodyState {
319    stream: UpstreamByteStream,
320    observer: NativeResponseObserver,
321    body_idle_timeout_ms: u64,
322}
323
324struct NativeResponseObserver {
325    ctx: RequestContext,
326    is_sse: bool,
327    generation_started: bool,
328    raw: Vec<u8>,
329    raw_truncated: u64,
330    pending: Vec<u8>,
331    pending_scan: usize,
332    discarding_oversized_frame: bool,
333    pending_truncated: bool,
334    captured_events: Vec<Value>,
335    captured_event_bytes: usize,
336    captured_events_truncated: u64,
337    input_tokens: Option<u64>,
338    output_tokens: Option<u64>,
339    outcome: NativeResponseOutcome,
340    finished: bool,
341}
342
343impl NativeResponseObserver {
344    fn new(ctx: RequestContext, is_sse: bool, outcome: NativeResponseOutcome) -> Self {
345        Self {
346            ctx,
347            is_sse,
348            generation_started: false,
349            raw: Vec::with_capacity(64 * 1024),
350            raw_truncated: 0,
351            pending: Vec::new(),
352            pending_scan: 0,
353            discarding_oversized_frame: false,
354            pending_truncated: false,
355            captured_events: Vec::new(),
356            captured_event_bytes: 0,
357            captured_events_truncated: 0,
358            input_tokens: None,
359            output_tokens: None,
360            outcome,
361            finished: false,
362        }
363    }
364
365    fn observe(&mut self, chunk: &[u8]) {
366        if !chunk.is_empty() && !self.generation_started {
367            if let Some(monitor) = self.ctx.monitor.as_ref() {
368                monitor.generation_started(&self.ctx.req_id);
369            }
370            self.generation_started = true;
371        }
372        self.capture_raw(chunk);
373
374        let events = if self.is_sse {
375            self.pending.extend_from_slice(chunk);
376            self.drain_sse_events()
377        } else {
378            0
379        };
380        if let Some(monitor) = self.ctx.monitor.as_ref() {
381            monitor.stream_progress(
382                &self.ctx.req_id,
383                chunk.len() as u64,
384                events,
385                self.input_tokens,
386                self.output_tokens,
387            );
388        }
389    }
390
391    fn capture_raw(&mut self, chunk: &[u8]) {
392        let remaining = MAX_SSE_CAPTURE_BYTES.saturating_sub(self.raw.len());
393        let captured = remaining.min(chunk.len());
394        self.raw.extend_from_slice(&chunk[..captured]);
395        if captured < chunk.len() {
396            self.raw_truncated = self
397                .raw_truncated
398                .saturating_add((chunk.len() - captured) as u64);
399        }
400    }
401
402    fn drain_sse_events(&mut self) -> u64 {
403        if self.discarding_oversized_frame {
404            let Some((end, separator_len)) = find_sse_boundary(&self.pending) else {
405                retain_boundary_prefix(&mut self.pending);
406                return 0;
407            };
408            self.pending.drain(..end + separator_len);
409            self.discarding_oversized_frame = false;
410            self.pending_scan = 0;
411        }
412
413        let mut consumed = 0;
414        let mut parsed = Vec::new();
415        while let Some((relative_end, separator_len)) =
416            find_sse_boundary_from(&self.pending, self.pending_scan)
417        {
418            let end = relative_end + separator_len;
419            parsed.extend(parse_sse_events(&self.pending[consumed..end]));
420            consumed = end;
421            self.pending_scan = consumed;
422        }
423        if consumed > 0 {
424            self.pending.drain(..consumed);
425            self.pending_scan = 0;
426        } else {
427            self.pending_scan = self.pending.len().saturating_sub(3);
428        }
429
430        let mut count = 0_u64;
431        for event in parsed {
432            count += 1;
433            if event.data == "[DONE]" {
434                continue;
435            }
436            match serde_json::from_str::<Value>(&event.data) {
437                Ok(value) => self.record_event(event.event.as_deref(), value),
438                Err(_) => self.capture_event(json!({
439                    "event": event.event,
440                    "unparseable": true,
441                    "bytes": event.data.len(),
442                })),
443            }
444        }
445
446        if self.pending.len() > MAX_STREAM_CAPTURE_FRAME_BYTES {
447            self.pending_truncated = true;
448            self.discarding_oversized_frame = true;
449            retain_boundary_prefix(&mut self.pending);
450            self.pending_scan = 0;
451        }
452        count
453    }
454
455    fn record_event(&mut self, event: Option<&str>, value: Value) {
456        self.update_usage(&value);
457        self.update_outcome(&value);
458        let mut captured = value;
459        if let Some(event) = event
460            && let Some(object) = captured.as_object_mut()
461        {
462            object
463                .entry("_sse_event")
464                .or_insert_with(|| Value::String(event.to_string()));
465        }
466        self.capture_event(captured);
467    }
468
469    fn capture_event(&mut self, value: Value) {
470        if self.ctx.traffic.is_none() {
471            return;
472        }
473        let bytes = serde_json::to_vec(&value).map_or(0, |value| value.len());
474        if self.captured_events.len() < MAX_STREAM_CAPTURE_EVENTS
475            && self.captured_event_bytes.saturating_add(bytes) <= MAX_STREAM_CAPTURE_EVENT_BYTES
476        {
477            self.captured_event_bytes += bytes;
478            self.captured_events.push(value);
479        } else {
480            self.captured_events_truncated = self.captured_events_truncated.saturating_add(1);
481        }
482    }
483
484    fn update_outcome(&self, value: &Value) {
485        let event_type = value.get("type").and_then(Value::as_str);
486        let has_error = value.get("error").is_some_and(|error| !error.is_null());
487        let failed_status = value.get("status").and_then(Value::as_str) == Some("failed");
488        if matches!(
489            event_type,
490            Some("response.failed" | "response.error" | "error")
491        ) || has_error
492            || failed_status
493        {
494            let message = value
495                .pointer("/response/error/message")
496                .or_else(|| value.pointer("/error/message"))
497                .or_else(|| value.get("message"))
498                .and_then(Value::as_str)
499                .unwrap_or("Native Responses stream failed");
500            self.outcome.fail(message.to_string());
501        }
502    }
503
504    fn update_usage(&mut self, value: &Value) {
505        let usage = value
506            .pointer("/response/usage")
507            .or_else(|| value.get("usage"));
508        if let Some(usage) = usage {
509            self.input_tokens = usage
510                .get("input_tokens")
511                .and_then(Value::as_u64)
512                .or(self.input_tokens);
513            self.output_tokens = usage
514                .get("output_tokens")
515                .and_then(Value::as_u64)
516                .or(self.output_tokens);
517        }
518    }
519
520    fn finish(&mut self, outcome: &str) {
521        if self.finished {
522            return;
523        }
524        self.finished = true;
525        if !self.is_sse
526            && let Ok(value) = serde_json::from_slice::<Value>(&self.raw)
527        {
528            self.update_usage(&value);
529            self.update_outcome(&value);
530            self.capture_event(value);
531            if let Some(monitor) = self.ctx.monitor.as_ref() {
532                monitor.usage_updated(&self.ctx.req_id, self.input_tokens, self.output_tokens);
533            }
534        }
535        self.write_capture(outcome);
536    }
537
538    fn write_capture(&self, outcome: &str) {
539        let Some(traffic) = self.ctx.traffic.as_deref() else {
540            return;
541        };
542        if !self.raw.is_empty() {
543            traffic.write_bytes(
544                if self.is_sse {
545                    "032-upstream-response-body.sse"
546                } else {
547                    "032-upstream-response-body.json"
548                },
549                &self.raw,
550            );
551        }
552        for event in &self.captured_events {
553            traffic.write_json_event("040-upstream-event", event);
554        }
555        traffic.write_json(
556            "033-native-response-capture",
557            &json!({
558                "outcome": outcome,
559                "capturedBytes": self.raw.len(),
560                "truncatedBytes": self.raw_truncated,
561                "pendingFrameTruncated": self.pending_truncated,
562                "capturedEvents": self.captured_events.len(),
563                "capturedEventBytes": self.captured_event_bytes,
564                "truncatedEvents": self.captured_events_truncated,
565                "inputTokens": self.input_tokens,
566                "outputTokens": self.output_tokens,
567            }),
568        );
569    }
570}
571
572impl Drop for NativeResponseObserver {
573    fn drop(&mut self) {
574        if !self.finished {
575            self.finish("downstream_cancelled");
576        }
577    }
578}
579
580fn find_sse_boundary(bytes: &[u8]) -> Option<(usize, usize)> {
581    find_sse_boundary_from(bytes, 0)
582}
583
584fn find_sse_boundary_from(bytes: &[u8], start: usize) -> Option<(usize, usize)> {
585    for index in start.min(bytes.len())..bytes.len() {
586        if bytes[index..].starts_with(b"\r\n\r\n") {
587            return Some((index, 4));
588        }
589        if bytes[index..].starts_with(b"\n\n") || bytes[index..].starts_with(b"\r\r") {
590            return Some((index, 2));
591        }
592    }
593    None
594}
595
596fn retain_boundary_prefix(bytes: &mut Vec<u8>) {
597    let keep = bytes.len().min(3);
598    if bytes.len() > keep {
599        bytes.drain(..bytes.len() - keep);
600    }
601}
602
603#[cfg(test)]
604mod tests {
605    use super::*;
606
607    fn request(body: Value) -> Value {
608        body
609    }
610
611    fn observer_context() -> RequestContext {
612        RequestContext {
613            req_id: "native-test".into(),
614            session_id: None,
615            session_seq: None,
616            provider: "codex".into(),
617            traffic: None,
618            monitor: None,
619            passthrough: None,
620        }
621    }
622
623    #[test]
624    fn failed_sse_event_records_native_outcome() {
625        let outcome = NativeResponseOutcome::default();
626        let mut observer = NativeResponseObserver::new(observer_context(), true, outcome.clone());
627        observer.observe(
628            b"event: response.failed\ndata: {\"type\":\"response.failed\",\"response\":{\"error\":{\"message\":\"generation failed\"}}}\n\n",
629        );
630
631        assert_eq!(outcome.failure().as_deref(), Some("generation failed"));
632    }
633
634    #[test]
635    fn completed_json_with_null_error_stays_successful() {
636        let outcome = NativeResponseOutcome::default();
637        let mut observer = NativeResponseObserver::new(observer_context(), false, outcome.clone());
638        observer
639            .observe(br#"{"id":"resp_ok","object":"response","status":"completed","error":null}"#);
640        observer.finish("complete");
641
642        assert_eq!(outcome.failure(), None);
643    }
644
645    #[test]
646    fn failed_json_records_error_message() {
647        let outcome = NativeResponseOutcome::default();
648        let mut observer = NativeResponseObserver::new(observer_context(), false, outcome.clone());
649        observer.observe(
650            br#"{"id":"resp_failed","object":"response","status":"failed","error":{"message":"request failed"}}"#,
651        );
652        observer.finish("complete");
653
654        assert_eq!(outcome.failure().as_deref(), Some("request failed"));
655    }
656
657    #[test]
658    fn response_error_event_records_failure() {
659        let outcome = NativeResponseOutcome::default();
660        let mut observer = NativeResponseObserver::new(observer_context(), true, outcome.clone());
661        observer.observe(
662            b"event: response.error\ndata: {\"type\":\"response.error\",\"response\":{\"error\":{\"message\":\"stream error\"}}}\n\n",
663        );
664
665        assert_eq!(outcome.failure().as_deref(), Some("stream error"));
666    }
667
668    #[test]
669    fn event_capture_obeys_count_limit() {
670        let temp = tempfile::TempDir::new().unwrap();
671        let outcome = NativeResponseOutcome::default();
672        let mut context = observer_context();
673        context.traffic = Some(Arc::new(crate::traffic::test_capture(
674            temp.path().to_path_buf(),
675        )));
676        let mut observer = NativeResponseObserver::new(context, true, outcome);
677        for index in 0..MAX_STREAM_CAPTURE_EVENTS + 10 {
678            observer.capture_event(json!({"index": index}));
679        }
680
681        assert_eq!(observer.captured_events.len(), MAX_STREAM_CAPTURE_EVENTS);
682        assert_eq!(observer.captured_events_truncated, 10);
683        observer.finished = true;
684    }
685
686    #[test]
687    fn native_request_requires_object_and_model() {
688        assert!(shape_native_request(&mut json!([])).is_err());
689        assert!(shape_native_request(&mut json!({})).is_err());
690        assert!(shape_native_request(&mut json!({"model": 7})).is_err());
691    }
692
693    #[test]
694    fn native_request_resolves_alias_and_fast_tier() {
695        let mut body = request(json!({"model":"claude-opus-5","input":[]}));
696        let resolved = shape_native_request(&mut body).unwrap();
697        assert_eq!(resolved.model, "gpt-5.6-sol");
698        assert_eq!(body["model"], "gpt-5.6-sol");
699        assert!(resolved.use_responses_lite);
700
701        let mut fast = request(json!({"model":"gpt-5.4-fast","input":[]}));
702        let resolved = shape_native_request(&mut fast).unwrap();
703        assert_eq!(resolved.model, "gpt-5.4");
704        assert_eq!(fast["service_tier"], "priority");
705    }
706
707    #[test]
708    fn native_request_preserves_parallel_tool_calls() {
709        for parallel in [false, true] {
710            let mut body = request(json!({
711                "model":"gpt-5.4",
712                "input":[],
713                "parallel_tool_calls":parallel
714            }));
715            shape_native_request(&mut body).unwrap();
716            assert_eq!(body["parallel_tool_calls"], parallel);
717        }
718    }
719
720    #[test]
721    fn explicit_service_tier_is_preserved() {
722        let mut body = request(json!({
723            "model":"gpt-5.4-fast",
724            "service_tier":"flex",
725            "input":[]
726        }));
727        shape_native_request(&mut body).unwrap();
728        assert_eq!(body["service_tier"], "flex");
729    }
730
731    #[test]
732    fn hosted_search_uses_full_lane_and_upgrades_luna() {
733        for tool_type in ["web_search", "web_search_preview"] {
734            let mut body = request(json!({
735                "model":"gpt-5.6-luna",
736                "tools":[{"type":tool_type}],
737                "input":[]
738            }));
739            let resolved = shape_native_request(&mut body).unwrap();
740            assert_eq!(resolved.model, "gpt-5.6-sol");
741            assert!(!resolved.use_responses_lite);
742        }
743    }
744
745    #[tokio::test]
746    async fn openai_error_has_native_envelope() {
747        let response = openai_error(
748            StatusCode::BAD_REQUEST,
749            "invalid_request_error",
750            "bad model",
751            Some("model"),
752            Some("invalid"),
753        );
754        let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
755            .await
756            .unwrap();
757        let value: Value = serde_json::from_slice(&bytes).unwrap();
758        assert!(value.get("type").is_none());
759        assert_eq!(value["error"]["type"], "invalid_request_error");
760        assert_eq!(value["error"]["param"], "model");
761        assert_eq!(value["error"]["code"], "invalid");
762    }
763
764    #[test]
765    fn response_headers_use_allowlist() {
766        let mut upstream = HeaderMap::new();
767        upstream.insert(
768            http::header::CONTENT_TYPE,
769            "text/event-stream".parse().unwrap(),
770        );
771        upstream.insert(http::header::SET_COOKIE, "secret=1".parse().unwrap());
772        upstream.insert(http::header::CONTENT_LENGTH, "12".parse().unwrap());
773        upstream.insert("x-request-id", "req_1".parse().unwrap());
774        upstream.insert("x-ratelimit-remaining-requests", "2".parse().unwrap());
775
776        let headers = passthrough_headers(&upstream);
777        assert_eq!(
778            headers.get(http::header::CONTENT_TYPE).unwrap(),
779            "text/event-stream"
780        );
781        assert_eq!(headers.get("x-request-id").unwrap(), "req_1");
782        assert_eq!(headers.get("x-ratelimit-remaining-requests").unwrap(), "2");
783        assert!(headers.get(http::header::SET_COOKIE).is_none());
784        assert!(headers.get(http::header::CONTENT_LENGTH).is_none());
785    }
786
787    #[test]
788    fn sse_boundary_handles_lf_and_crlf() {
789        assert_eq!(find_sse_boundary(b"data: {}\n\nnext"), Some((8, 2)));
790        assert_eq!(find_sse_boundary(b"data: {}\r\n\r\nnext"), Some((8, 4)));
791        assert_eq!(find_sse_boundary(b"data: {}"), None);
792    }
793}