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
14const 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
29pub 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 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 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
120fn 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
227pub(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 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 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 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}