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}