Skip to main content

lean_ctx/proxy/
forward.rs

1use axum::{
2    body::Body,
3    extract::State,
4    http::{Request, StatusCode, request::Parts},
5    response::Response,
6};
7
8use flate2::{Compression, read::GzDecoder, write::GzEncoder};
9use std::borrow::Cow;
10use std::io::{Read, Write};
11
12use super::ProxyState;
13
14/// Default request-body ceiling (MiB). A large-codebase refactor with several
15/// big files in context easily exceeds the old 10 MiB cap, which surfaced to the
16/// agent as a hard `400` mid-task. Raised and made configurable via
17/// `LEAN_CTX_PROXY_MAX_BODY_MB`.
18const DEFAULT_MAX_BODY_MB: usize = 64;
19
20pub(super) fn max_body_bytes() -> usize {
21    std::env::var("LEAN_CTX_PROXY_MAX_BODY_MB")
22        .ok()
23        .and_then(|v| v.trim().parse::<usize>().ok())
24        .filter(|mb| *mb > 0)
25        .unwrap_or(DEFAULT_MAX_BODY_MB)
26        .saturating_mul(1024 * 1024)
27}
28
29/// Transforms the already-parsed JSON request body (parsed once upstream, so the
30/// compressor never re-parses) into the serialized — possibly compressed — body,
31/// its original size, and its compressed size. A plain `fn` from the static
32/// providers or a closure that captures request-derived context (e.g. Gemini's
33/// path-encoded model) both satisfy this bound.
34pub async fn forward_request(
35    State(state): State<ProxyState>,
36    req: Request<Body>,
37    upstream_base: &str,
38    default_path: &str,
39    compress_body: impl FnOnce(serde_json::Value, usize) -> (Vec<u8>, usize, usize),
40    provider_label: &str,
41    extra_stream_types: &[&str],
42) -> Result<Response, StatusCode> {
43    let (parts, body) = req.into_parts();
44    let body_bytes = axum::body::to_bytes(body, max_body_bytes())
45        .await
46        .map_err(|_| StatusCode::PAYLOAD_TOO_LARGE)?;
47
48    let prepared = prepare_request_body(&parts, &body_bytes, compress_body)?;
49    let original_size = prepared.original_size;
50    let compressed_size = prepared.compressed_size;
51    let compression_candidate = prepared.compression_candidate;
52    let preserve_content_encoding = prepared.preserve_content_encoding;
53    let parsed = prepared.parsed;
54    if let Some(ref parsed) = parsed {
55        let provider = match provider_label {
56            "Anthropic" => super::introspect::Provider::Anthropic,
57            "OpenAI" | "ChatGPT" => super::introspect::Provider::OpenAi,
58            _ => super::introspect::Provider::Gemini,
59        };
60        let breakdown = super::introspect::analyze_request(parsed, provider);
61        state.introspect.record(breakdown);
62    }
63
64    // #895 Track B: assign output-savings holdout from the same pristine parsed
65    // body that each provider's compressor receives. Only when active.
66    let cohort = parsed
67        .as_ref()
68        .and_then(|p| cohort_arm(p, provider_label, default_path));
69
70    if compression_candidate {
71        state
72            .stats
73            .record_provider_request(provider_label, original_size, compressed_size);
74    }
75
76    let tokens_saved = original_size.saturating_sub(compressed_size) as u64 / 4;
77    super::metrics::record_request(tokens_saved, compressed_size as u64);
78
79    let model = parsed
80        .as_ref()
81        .and_then(|v| v.get("model"))
82        .and_then(|m| m.as_str());
83    super::cost::record(
84        model,
85        tokens_saved,
86        original_size as u64,
87        compressed_size as u64,
88    );
89
90    let upstream_url = build_upstream_url(&parts, upstream_base, default_path);
91    let response = send_upstream(
92        &state,
93        &parts,
94        &upstream_url,
95        prepared.body,
96        provider_label,
97        preserve_content_encoding,
98    )
99    .await?;
100
101    // Measured usage: read the real model + billed tokens from the response.
102    // Gemini puts the model in the URL path, not the request/response body.
103    let usage_provider = super::usage::Provider::from_label(provider_label);
104    let url_model = if usage_provider == super::usage::Provider::Gemini {
105        super::usage::gemini_model_from_path(parts.uri.path())
106    } else {
107        None
108    };
109
110    build_response(
111        response,
112        extra_stream_types,
113        usage_provider,
114        url_model,
115        cohort,
116    )
117    .await
118}
119
120/// Output-savings arm (#895) for a request body, or `None` when no holdout is
121/// active. Keyed per provider; OpenAI's Chat vs Responses bodies are
122/// distinguished by the request path so each uses the matching cohort key.
123fn cohort_arm(
124    parsed: &serde_json::Value,
125    provider_label: &str,
126    default_path: &str,
127) -> Option<super::holdout::Arm> {
128    let holdout = crate::core::config::Config::load()
129        .proxy
130        .output_holdout_fraction();
131    if holdout <= 0.0 {
132        return None;
133    }
134    let key = match provider_label {
135        "Anthropic" => super::holdout::anthropic_key(parsed),
136        "OpenAI" | "ChatGPT" => {
137            if default_path.contains("responses") {
138                super::holdout::openai_responses_key(parsed)
139            } else {
140                super::holdout::openai_chat_key(parsed)
141            }
142        }
143        _ => super::holdout::google_key(parsed),
144    };
145    Some(super::holdout::assign(&key, holdout))
146}
147
148struct PreparedRequestBody {
149    body: Vec<u8>,
150    parsed: Option<serde_json::Value>,
151    original_size: usize,
152    compressed_size: usize,
153    compression_candidate: bool,
154    preserve_content_encoding: bool,
155}
156
157#[derive(Clone, Copy, Debug, Eq, PartialEq)]
158enum RequestBodyEncoding {
159    Identity,
160    Gzip,
161    Zstd,
162    Passthrough,
163}
164
165fn prepare_request_body(
166    parts: &Parts,
167    body_bytes: &[u8],
168    compress_body: impl FnOnce(serde_json::Value, usize) -> (Vec<u8>, usize, usize),
169) -> Result<PreparedRequestBody, StatusCode> {
170    let encoding = request_body_encoding(parts);
171    let decoded = match encoding {
172        RequestBodyEncoding::Identity => Cow::Borrowed(body_bytes),
173        RequestBodyEncoding::Gzip => Cow::Owned(decode_gzip_bounded(body_bytes, max_body_bytes())?),
174        RequestBodyEncoding::Zstd => Cow::Owned(decode_zstd_bounded(body_bytes, max_body_bytes())?),
175        RequestBodyEncoding::Passthrough => {
176            return Ok(PreparedRequestBody {
177                body: body_bytes.to_vec(),
178                parsed: None,
179                original_size: body_bytes.len(),
180                compressed_size: body_bytes.len(),
181                compression_candidate: false,
182                preserve_content_encoding: true,
183            });
184        }
185    };
186
187    let Some(parsed) = serde_json::from_slice::<serde_json::Value>(&decoded).ok() else {
188        return Ok(PreparedRequestBody {
189            body: body_bytes.to_vec(),
190            parsed: None,
191            original_size: body_bytes.len(),
192            compressed_size: body_bytes.len(),
193            compression_candidate: false,
194            preserve_content_encoding: encoding != RequestBodyEncoding::Identity,
195        });
196    };
197
198    let original_size = decoded.len();
199    let (logical_body, _, compressed_size) = compress_body(parsed.clone(), original_size);
200    let body = match encoding {
201        RequestBodyEncoding::Identity => logical_body,
202        RequestBodyEncoding::Gzip => encode_gzip(&logical_body)?,
203        RequestBodyEncoding::Zstd => encode_zstd(&logical_body)?,
204        RequestBodyEncoding::Passthrough => unreachable!("passthrough returned above"),
205    };
206
207    Ok(PreparedRequestBody {
208        body,
209        parsed: Some(parsed),
210        original_size,
211        compressed_size,
212        compression_candidate: true,
213        preserve_content_encoding: encoding != RequestBodyEncoding::Identity,
214    })
215}
216
217fn build_upstream_url(parts: &Parts, base: &str, default_path: &str) -> String {
218    format!(
219        "{base}{}",
220        parts
221            .uri
222            .path_and_query()
223            .map_or(default_path, axum::http::uri::PathAndQuery::as_str)
224    )
225}
226
227/// Request headers forwarded verbatim to the upstream provider. Anything not
228/// listed here is stripped before the request leaves the loopback proxy.
229///
230/// `openai-project` (and `openai-organization`) must be forwarded: OpenCode and
231/// the OpenAI SDK send the project scope via this header for project-scoped API
232/// keys when calling the Responses API (`/responses`). Dropping it makes OpenAI
233/// reject the request with `Missing scopes: api.responses.write` (#366).
234pub(super) const ALLOWED_REQUEST_HEADERS: &[&str] = &[
235    "authorization",
236    "x-api-key",
237    "content-type",
238    "accept",
239    "user-agent",
240    "originator",
241    "anthropic-version",
242    "anthropic-beta",
243    "anthropic-dangerous-direct-browser-access",
244    "openai-organization",
245    "openai-project",
246    "openai-beta",
247    "chatgpt-account-id",
248    "x-openai-fedramp",
249    "x-openai-internal-codex-residency",
250    "x-openai-internal-codex-responses-lite",
251    "x-openai-product-sku",
252    "oai-product-sku",
253    "x-oai-attestation",
254    "x-client-request-id",
255    "x-codex-beta-features",
256    "x-codex-installation-id",
257    "x-codex-parent-thread-id",
258    "x-openai-subagent",
259    "x-codex-turn-state",
260    "x-codex-turn-metadata",
261    "x-codex-window-id",
262    "x-openai-memgen-request",
263    "x-responsesapi-include-timing-metrics",
264    "mcp-session-id",
265    "last-event-id",
266    "cache-control",
267    "x-goog-api-key",
268    "x-goog-api-client",
269];
270
271pub(super) fn is_allowed_request_header(name: &str) -> bool {
272    ALLOWED_REQUEST_HEADERS.contains(&name)
273}
274
275fn should_forward_request_header(name: &str, preserve_content_encoding: bool) -> bool {
276    is_allowed_request_header(name)
277        || (preserve_content_encoding && name.eq_ignore_ascii_case("content-encoding"))
278}
279
280fn request_body_encoding(parts: &Parts) -> RequestBodyEncoding {
281    let Some(value) = parts
282        .headers
283        .get(axum::http::header::CONTENT_ENCODING)
284        .and_then(|value| value.to_str().ok())
285    else {
286        return RequestBodyEncoding::Identity;
287    };
288
289    let encodings = value
290        .split(',')
291        .map(str::trim)
292        .filter(|part| !part.is_empty() && !part.eq_ignore_ascii_case("identity"))
293        .collect::<Vec<_>>();
294    match encodings.as_slice() {
295        [] => RequestBodyEncoding::Identity,
296        [encoding] if encoding.eq_ignore_ascii_case("gzip") => RequestBodyEncoding::Gzip,
297        [encoding] if encoding.eq_ignore_ascii_case("zstd") => RequestBodyEncoding::Zstd,
298        _ => RequestBodyEncoding::Passthrough,
299    }
300}
301
302fn decode_zstd_bounded(data: &[u8], max_bytes: usize) -> Result<Vec<u8>, StatusCode> {
303    let decoder = zstd::Decoder::new(data).map_err(|e| {
304        tracing::warn!("lean-ctx proxy: invalid zstd request body: {e}");
305        StatusCode::BAD_REQUEST
306    })?;
307    read_bounded(decoder, max_bytes).inspect_err(|e| {
308        tracing::warn!("lean-ctx proxy: zstd request decode failed: {e}");
309    })
310}
311
312fn encode_zstd(data: &[u8]) -> Result<Vec<u8>, StatusCode> {
313    zstd::encode_all(data, 3).map_err(|e| {
314        tracing::error!("lean-ctx proxy: zstd request encode failed: {e}");
315        StatusCode::INTERNAL_SERVER_ERROR
316    })
317}
318
319fn decode_gzip_bounded(data: &[u8], max_bytes: usize) -> Result<Vec<u8>, StatusCode> {
320    read_bounded(GzDecoder::new(data), max_bytes).inspect_err(|e| {
321        tracing::warn!("lean-ctx proxy: gzip request decode failed: {e}");
322    })
323}
324
325fn encode_gzip(data: &[u8]) -> Result<Vec<u8>, StatusCode> {
326    let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
327    encoder.write_all(data).map_err(|e| {
328        tracing::error!("lean-ctx proxy: gzip request encode failed: {e}");
329        StatusCode::INTERNAL_SERVER_ERROR
330    })?;
331    encoder.finish().map_err(|e| {
332        tracing::error!("lean-ctx proxy: gzip request encode failed: {e}");
333        StatusCode::INTERNAL_SERVER_ERROR
334    })
335}
336
337fn read_bounded<R: Read>(reader: R, max_bytes: usize) -> Result<Vec<u8>, StatusCode> {
338    let mut limited = reader.take(max_bytes as u64 + 1);
339    let mut out = Vec::new();
340    limited
341        .read_to_end(&mut out)
342        .map_err(|_| StatusCode::BAD_REQUEST)?;
343    if out.len() > max_bytes {
344        return Err(StatusCode::PAYLOAD_TOO_LARGE);
345    }
346    Ok(out)
347}
348
349async fn send_upstream(
350    state: &ProxyState,
351    parts: &Parts,
352    url: &str,
353    body: Vec<u8>,
354    provider_label: &str,
355    preserve_content_encoding: bool,
356) -> Result<reqwest::Response, StatusCode> {
357    let mut req = state.client.request(parts.method.clone(), url);
358
359    for (key, value) in &parts.headers {
360        let k = key.as_str().to_lowercase();
361        if should_forward_request_header(&k, preserve_content_encoding) {
362            req = req.header(key.clone(), value.clone());
363        }
364    }
365
366    req.body(body).send().await.map_err(|e| {
367        tracing::error!("lean-ctx proxy: {provider_label} upstream error: {e}");
368        StatusCode::BAD_GATEWAY
369    })
370}
371
372pub(super) const FORWARDED_HEADERS: &[&str] = &[
373    "content-type",
374    "content-encoding",
375    "mcp-session-id",
376    "x-request-id",
377    "x-oai-request-id",
378    "cf-ray",
379    "x-openai-authorization-error",
380    "x-error-json",
381    "openai-organization",
382    "openai-model",
383    "openai-processing-ms",
384    "openai-version",
385    "x-models-etag",
386    "x-reasoning-included",
387    "anthropic-ratelimit-requests-limit",
388    "anthropic-ratelimit-requests-remaining",
389    "anthropic-ratelimit-tokens-limit",
390    "anthropic-ratelimit-tokens-remaining",
391    "retry-after",
392    "x-ratelimit-limit-requests",
393    "x-ratelimit-remaining-requests",
394    "x-ratelimit-limit-tokens",
395    "x-ratelimit-remaining-tokens",
396    "cache-control",
397];
398
399pub(super) fn is_forwarded_response_header(name: &str) -> bool {
400    FORWARDED_HEADERS.contains(&name)
401        || name.starts_with("x-codex-")
402        || name.starts_with("x-ratelimit-")
403}
404
405async fn build_response(
406    response: reqwest::Response,
407    extra_stream_types: &[&str],
408    usage_provider: super::usage::Provider,
409    url_model: Option<String>,
410    cohort: Option<super::holdout::Arm>,
411) -> Result<Response, StatusCode> {
412    let status = StatusCode::from_u16(response.status().as_u16()).unwrap_or(StatusCode::OK);
413    let resp_headers = response.headers().clone();
414
415    let is_stream = resp_headers
416        .get("content-type")
417        .and_then(|v| v.to_str().ok())
418        .is_some_and(|ct| {
419            ct.contains("text/event-stream") || extra_stream_types.iter().any(|t| ct.contains(t))
420        });
421
422    if is_stream {
423        // Tee the stream through a usage Scanner: each chunk is forwarded
424        // byte-for-byte while the real model + billed tokens are extracted from
425        // the final event and recorded when the stream ends.
426        let scanner = super::usage::Scanner::new(usage_provider, url_model).with_cohort(cohort);
427        let inner = Box::pin(response.bytes_stream());
428        let body = Body::from_stream(super::usage::tee_stream(inner, scanner));
429        let mut resp = Response::builder().status(status);
430        for (k, v) in &resp_headers {
431            let ks = k.as_str().to_lowercase();
432            if is_forwarded_response_header(&ks) {
433                resp = resp.header(k, v);
434            }
435        }
436        return resp
437            .body(body)
438            .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR);
439    }
440
441    let resp_bytes = response
442        .bytes()
443        .await
444        .map_err(|_| StatusCode::BAD_GATEWAY)?;
445
446    // Non-streaming: the whole body is one JSON object carrying `usage`.
447    let mut scanner = super::usage::Scanner::new(usage_provider, url_model).with_cohort(cohort);
448    scanner.feed_body(&resp_bytes);
449    if let Some(usage) = scanner.finalize() {
450        super::usage_meter::record(&usage);
451    }
452
453    let mut resp = Response::builder().status(status);
454    for (k, v) in &resp_headers {
455        let ks = k.as_str().to_lowercase();
456        if is_forwarded_response_header(&ks) {
457            resp = resp.header(k, v);
458        }
459    }
460    resp.body(Body::from(resp_bytes))
461        .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)
462}
463
464#[cfg(test)]
465mod tests {
466    use super::*;
467
468    fn parts_for(uri: &str) -> Parts {
469        Request::builder().uri(uri).body(()).unwrap().into_parts().0
470    }
471
472    fn add_test_marker(
473        mut value: serde_json::Value,
474        original_size: usize,
475    ) -> (Vec<u8>, usize, usize) {
476        value["lean_ctx_touched"] = serde_json::Value::Bool(true);
477        let out = serde_json::to_vec(&value).unwrap();
478        let compressed_size = out.len();
479        (out, original_size, compressed_size)
480    }
481
482    #[test]
483    fn zstd_request_bodies_are_rewritten_and_reencoded() {
484        let body = serde_json::json!({"model": "gpt-5", "input": []});
485        let json = serde_json::to_vec(&body).unwrap();
486        let encoded = encode_zstd(&json).unwrap();
487        let parts = Request::builder()
488            .uri("/backend-api/codex/responses")
489            .header(axum::http::header::CONTENT_ENCODING, "zstd")
490            .body(())
491            .unwrap()
492            .into_parts()
493            .0;
494
495        let prepared = prepare_request_body(&parts, &encoded, add_test_marker).unwrap();
496        assert_eq!(request_body_encoding(&parts), RequestBodyEncoding::Zstd);
497        assert_eq!(prepared.original_size, json.len());
498        assert!(prepared.compression_candidate);
499        assert!(prepared.preserve_content_encoding);
500        assert!(should_forward_request_header("content-encoding", true));
501        assert!(!should_forward_request_header("content-encoding", false));
502
503        let decoded = zstd::decode_all(prepared.body.as_slice()).unwrap();
504        let parsed: serde_json::Value = serde_json::from_slice(&decoded).unwrap();
505        assert_eq!(parsed["lean_ctx_touched"], true);
506        assert_eq!(parsed["model"], "gpt-5");
507    }
508
509    #[test]
510    fn gzip_request_bodies_are_rewritten_and_reencoded() {
511        let body = serde_json::json!({"model": "gpt-5", "input": []});
512        let json = serde_json::to_vec(&body).unwrap();
513        let encoded = encode_gzip(&json).unwrap();
514        let parts = Request::builder()
515            .uri("/backend-api/codex/responses")
516            .header(axum::http::header::CONTENT_ENCODING, "gzip")
517            .body(())
518            .unwrap()
519            .into_parts()
520            .0;
521
522        let prepared = prepare_request_body(&parts, &encoded, add_test_marker).unwrap();
523        assert_eq!(request_body_encoding(&parts), RequestBodyEncoding::Gzip);
524        assert_eq!(prepared.original_size, json.len());
525        assert!(prepared.compression_candidate);
526        assert!(prepared.preserve_content_encoding);
527
528        let decoded = decode_gzip_bounded(&prepared.body, max_body_bytes()).unwrap();
529        let parsed: serde_json::Value = serde_json::from_slice(&decoded).unwrap();
530        assert_eq!(parsed["lean_ctx_touched"], true);
531        assert_eq!(parsed["model"], "gpt-5");
532    }
533
534    #[test]
535    fn identity_content_encoding_can_be_rewritten_as_json() {
536        let parts = Request::builder()
537            .uri("/v1/responses")
538            .header(axum::http::header::CONTENT_ENCODING, "identity")
539            .body(())
540            .unwrap()
541            .into_parts()
542            .0;
543
544        assert_eq!(request_body_encoding(&parts), RequestBodyEncoding::Identity);
545    }
546
547    #[test]
548    fn unknown_encoded_request_bodies_stay_passthrough() {
549        let parts = Request::builder()
550            .uri("/v1/responses")
551            .header(axum::http::header::CONTENT_ENCODING, "br")
552            .body(())
553            .unwrap()
554            .into_parts()
555            .0;
556        let body = b"not-json";
557
558        let prepared = prepare_request_body(&parts, body, |_, _| {
559            panic!("unknown encodings must not be JSON-rewritten")
560        })
561        .unwrap();
562
563        assert_eq!(
564            request_body_encoding(&parts),
565            RequestBodyEncoding::Passthrough
566        );
567        assert_eq!(prepared.body, body);
568        assert!(prepared.parsed.is_none());
569        assert!(!prepared.compression_candidate);
570        assert!(prepared.preserve_content_encoding);
571    }
572
573    #[test]
574    fn invalid_json_request_bodies_are_not_compression_candidates() {
575        let parts = Request::builder()
576            .uri("/v1/responses")
577            .body(())
578            .unwrap()
579            .into_parts()
580            .0;
581        let body = b"not-json";
582
583        let prepared = prepare_request_body(&parts, body, |_, _| {
584            panic!("invalid JSON must not enter the compression pipeline")
585        })
586        .unwrap();
587
588        assert_eq!(request_body_encoding(&parts), RequestBodyEncoding::Identity);
589        assert_eq!(prepared.body, body);
590        assert!(prepared.parsed.is_none());
591        assert!(!prepared.compression_candidate);
592        assert!(!prepared.preserve_content_encoding);
593    }
594
595    #[test]
596    fn upstream_url_preserves_subpath() {
597        let base = "https://api.anthropic.com";
598        let parts = parts_for("/v1/messages/count_tokens");
599        assert_eq!(
600            build_upstream_url(&parts, base, "/v1/messages"),
601            "https://api.anthropic.com/v1/messages/count_tokens"
602        );
603    }
604
605    #[test]
606    fn upstream_url_preserves_batches_subpath() {
607        let base = "https://api.anthropic.com";
608        let parts = parts_for("/v1/messages/batches/batch_123/results");
609        assert_eq!(
610            build_upstream_url(&parts, base, "/v1/messages"),
611            "https://api.anthropic.com/v1/messages/batches/batch_123/results"
612        );
613    }
614
615    #[test]
616    fn upstream_url_exact_path() {
617        let base = "https://api.anthropic.com";
618        let parts = parts_for("/v1/messages");
619        assert_eq!(
620            build_upstream_url(&parts, base, "/v1/messages"),
621            "https://api.anthropic.com/v1/messages"
622        );
623    }
624
625    #[test]
626    fn upstream_url_preserves_query_params() {
627        let base = "https://api.anthropic.com";
628        let parts = parts_for("/v1/messages/count_tokens?model=claude-4");
629        assert_eq!(
630            build_upstream_url(&parts, base, "/v1/messages"),
631            "https://api.anthropic.com/v1/messages/count_tokens?model=claude-4"
632        );
633    }
634
635    #[test]
636    fn forwards_openai_project_and_auth_headers() {
637        // #366: project-scoped OpenAI keys carry the scope via `OpenAI-Project`.
638        // It must be forwarded upstream, otherwise the Responses API rejects the
639        // call with `Missing scopes: api.responses.write`.
640        for required in ["authorization", "openai-project", "openai-organization"] {
641            assert!(
642                ALLOWED_REQUEST_HEADERS.contains(&required),
643                "request header `{required}` must be forwarded upstream"
644            );
645        }
646    }
647
648    #[test]
649    fn forwards_chatgpt_codex_oauth_headers() {
650        for required in [
651            "authorization",
652            "chatgpt-account-id",
653            "x-openai-fedramp",
654            "x-openai-internal-codex-residency",
655            "x-openai-product-sku",
656            "oai-product-sku",
657            "x-client-request-id",
658            "x-codex-installation-id",
659            "x-codex-turn-metadata",
660            "x-openai-subagent",
661            "x-codex-turn-state",
662            "originator",
663        ] {
664            assert!(
665                is_allowed_request_header(required),
666                "request header `{required}` must be forwarded upstream"
667            );
668        }
669    }
670
671    #[test]
672    fn forwards_streamable_http_mcp_headers() {
673        for required in ["mcp-session-id", "last-event-id"] {
674            assert!(
675                ALLOWED_REQUEST_HEADERS.contains(&required),
676                "request header `{required}` must be forwarded upstream"
677            );
678        }
679        assert!(
680            is_forwarded_response_header("mcp-session-id"),
681            "MCP session id response header must be forwarded downstream"
682        );
683    }
684
685    #[test]
686    fn forwards_codex_state_response_headers() {
687        for required in [
688            "x-codex-turn-state",
689            "x-codex-primary-used-percent",
690            "openai-model",
691            "x-models-etag",
692            "x-reasoning-included",
693            "x-oai-request-id",
694            "cf-ray",
695            "x-openai-authorization-error",
696            "x-error-json",
697        ] {
698            assert!(
699                is_forwarded_response_header(required),
700                "response header `{required}` must be forwarded downstream"
701            );
702        }
703    }
704
705    #[test]
706    fn chatgpt_responses_use_openai_responses_holdout_key() {
707        let _iso = crate::core::data_dir::isolated_data_dir();
708        crate::core::config::Config::update_global(|c| {
709            c.proxy.output_holdout = Some(1.0);
710        })
711        .unwrap();
712
713        let body = serde_json::json!({
714            "model": "gpt-5",
715            "input": "same conversation",
716        });
717
718        assert_eq!(
719            cohort_arm(&body, "ChatGPT", "/backend-api/codex/responses"),
720            cohort_arm(&body, "OpenAI", "/v1/responses")
721        );
722    }
723}