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" => 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.stats.record_request(original_size, compressed_size);
72 }
73
74 let tokens_saved = original_size.saturating_sub(compressed_size) as u64 / 4;
75 super::metrics::record_request(tokens_saved, compressed_size as u64);
76
77 let model = parsed
78 .as_ref()
79 .and_then(|v| v.get("model"))
80 .and_then(|m| m.as_str());
81 super::cost::record(
82 model,
83 tokens_saved,
84 original_size as u64,
85 compressed_size as u64,
86 );
87
88 let upstream_url = build_upstream_url(&parts, upstream_base, default_path);
89 let response = send_upstream(
90 &state,
91 &parts,
92 &upstream_url,
93 prepared.body,
94 provider_label,
95 preserve_content_encoding,
96 )
97 .await?;
98
99 let usage_provider = super::usage::Provider::from_label(provider_label);
102 let url_model = if usage_provider == super::usage::Provider::Gemini {
103 super::usage::gemini_model_from_path(parts.uri.path())
104 } else {
105 None
106 };
107
108 build_response(
109 response,
110 extra_stream_types,
111 usage_provider,
112 url_model,
113 cohort,
114 )
115 .await
116}
117
118fn cohort_arm(
122 parsed: &serde_json::Value,
123 provider_label: &str,
124 default_path: &str,
125) -> Option<super::holdout::Arm> {
126 let holdout = crate::core::config::Config::load()
127 .proxy
128 .output_holdout_fraction();
129 if holdout <= 0.0 {
130 return None;
131 }
132 let key = match provider_label {
133 "Anthropic" => super::holdout::anthropic_key(parsed),
134 "OpenAI" => {
135 if default_path.contains("responses") {
136 super::holdout::openai_responses_key(parsed)
137 } else {
138 super::holdout::openai_chat_key(parsed)
139 }
140 }
141 _ => super::holdout::google_key(parsed),
142 };
143 Some(super::holdout::assign(&key, holdout))
144}
145
146struct PreparedRequestBody {
147 body: Vec<u8>,
148 parsed: Option<serde_json::Value>,
149 original_size: usize,
150 compressed_size: usize,
151 compression_candidate: bool,
152 preserve_content_encoding: bool,
153}
154
155#[derive(Clone, Copy, Debug, Eq, PartialEq)]
156enum RequestBodyEncoding {
157 Identity,
158 Gzip,
159 Zstd,
160 Passthrough,
161}
162
163fn prepare_request_body(
164 parts: &Parts,
165 body_bytes: &[u8],
166 compress_body: impl FnOnce(serde_json::Value, usize) -> (Vec<u8>, usize, usize),
167) -> Result<PreparedRequestBody, StatusCode> {
168 let encoding = request_body_encoding(parts);
169 let decoded = match encoding {
170 RequestBodyEncoding::Identity => Cow::Borrowed(body_bytes),
171 RequestBodyEncoding::Gzip => Cow::Owned(decode_gzip_bounded(body_bytes, max_body_bytes())?),
172 RequestBodyEncoding::Zstd => Cow::Owned(decode_zstd_bounded(body_bytes, max_body_bytes())?),
173 RequestBodyEncoding::Passthrough => {
174 return Ok(PreparedRequestBody {
175 body: body_bytes.to_vec(),
176 parsed: None,
177 original_size: body_bytes.len(),
178 compressed_size: body_bytes.len(),
179 compression_candidate: false,
180 preserve_content_encoding: true,
181 });
182 }
183 };
184
185 let Some(parsed) = serde_json::from_slice::<serde_json::Value>(&decoded).ok() else {
186 return Ok(PreparedRequestBody {
187 body: body_bytes.to_vec(),
188 parsed: None,
189 original_size: body_bytes.len(),
190 compressed_size: body_bytes.len(),
191 compression_candidate: false,
192 preserve_content_encoding: encoding != RequestBodyEncoding::Identity,
193 });
194 };
195
196 let original_size = decoded.len();
197 let (logical_body, _, compressed_size) = compress_body(parsed.clone(), original_size);
198 let body = match encoding {
199 RequestBodyEncoding::Identity => logical_body,
200 RequestBodyEncoding::Gzip => encode_gzip(&logical_body)?,
201 RequestBodyEncoding::Zstd => encode_zstd(&logical_body)?,
202 RequestBodyEncoding::Passthrough => unreachable!("passthrough returned above"),
203 };
204
205 Ok(PreparedRequestBody {
206 body,
207 parsed: Some(parsed),
208 original_size,
209 compressed_size,
210 compression_candidate: true,
211 preserve_content_encoding: encoding != RequestBodyEncoding::Identity,
212 })
213}
214
215fn build_upstream_url(parts: &Parts, base: &str, default_path: &str) -> String {
216 format!(
217 "{base}{}",
218 parts
219 .uri
220 .path_and_query()
221 .map_or(default_path, axum::http::uri::PathAndQuery::as_str)
222 )
223}
224
225pub(super) const ALLOWED_REQUEST_HEADERS: &[&str] = &[
233 "authorization",
234 "x-api-key",
235 "content-type",
236 "accept",
237 "user-agent",
238 "originator",
239 "anthropic-version",
240 "anthropic-beta",
241 "anthropic-dangerous-direct-browser-access",
242 "openai-organization",
243 "openai-project",
244 "openai-beta",
245 "chatgpt-account-id",
246 "x-openai-fedramp",
247 "x-openai-internal-codex-residency",
248 "x-openai-internal-codex-responses-lite",
249 "x-openai-product-sku",
250 "oai-product-sku",
251 "x-oai-attestation",
252 "x-client-request-id",
253 "x-codex-beta-features",
254 "x-codex-installation-id",
255 "x-codex-parent-thread-id",
256 "x-openai-subagent",
257 "x-codex-turn-state",
258 "x-codex-turn-metadata",
259 "x-codex-window-id",
260 "x-openai-memgen-request",
261 "x-responsesapi-include-timing-metrics",
262 "mcp-session-id",
263 "last-event-id",
264 "cache-control",
265 "x-goog-api-key",
266 "x-goog-api-client",
267];
268
269pub(super) fn is_allowed_request_header(name: &str) -> bool {
270 ALLOWED_REQUEST_HEADERS.contains(&name)
271}
272
273fn should_forward_request_header(name: &str, preserve_content_encoding: bool) -> bool {
274 is_allowed_request_header(name)
275 || (preserve_content_encoding && name.eq_ignore_ascii_case("content-encoding"))
276}
277
278fn request_body_encoding(parts: &Parts) -> RequestBodyEncoding {
279 let Some(value) = parts
280 .headers
281 .get(axum::http::header::CONTENT_ENCODING)
282 .and_then(|value| value.to_str().ok())
283 else {
284 return RequestBodyEncoding::Identity;
285 };
286
287 let encodings = value
288 .split(',')
289 .map(str::trim)
290 .filter(|part| !part.is_empty() && !part.eq_ignore_ascii_case("identity"))
291 .collect::<Vec<_>>();
292 match encodings.as_slice() {
293 [] => RequestBodyEncoding::Identity,
294 [encoding] if encoding.eq_ignore_ascii_case("gzip") => RequestBodyEncoding::Gzip,
295 [encoding] if encoding.eq_ignore_ascii_case("zstd") => RequestBodyEncoding::Zstd,
296 _ => RequestBodyEncoding::Passthrough,
297 }
298}
299
300fn decode_zstd_bounded(data: &[u8], max_bytes: usize) -> Result<Vec<u8>, StatusCode> {
301 let decoder = zstd::Decoder::new(data).map_err(|e| {
302 tracing::warn!("lean-ctx proxy: invalid zstd request body: {e}");
303 StatusCode::BAD_REQUEST
304 })?;
305 read_bounded(decoder, max_bytes).inspect_err(|e| {
306 tracing::warn!("lean-ctx proxy: zstd request decode failed: {e}");
307 })
308}
309
310fn encode_zstd(data: &[u8]) -> Result<Vec<u8>, StatusCode> {
311 zstd::encode_all(data, 3).map_err(|e| {
312 tracing::error!("lean-ctx proxy: zstd request encode failed: {e}");
313 StatusCode::INTERNAL_SERVER_ERROR
314 })
315}
316
317fn decode_gzip_bounded(data: &[u8], max_bytes: usize) -> Result<Vec<u8>, StatusCode> {
318 read_bounded(GzDecoder::new(data), max_bytes).inspect_err(|e| {
319 tracing::warn!("lean-ctx proxy: gzip request decode failed: {e}");
320 })
321}
322
323fn encode_gzip(data: &[u8]) -> Result<Vec<u8>, StatusCode> {
324 let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
325 encoder.write_all(data).map_err(|e| {
326 tracing::error!("lean-ctx proxy: gzip request encode failed: {e}");
327 StatusCode::INTERNAL_SERVER_ERROR
328 })?;
329 encoder.finish().map_err(|e| {
330 tracing::error!("lean-ctx proxy: gzip request encode failed: {e}");
331 StatusCode::INTERNAL_SERVER_ERROR
332 })
333}
334
335fn read_bounded<R: Read>(reader: R, max_bytes: usize) -> Result<Vec<u8>, StatusCode> {
336 let mut limited = reader.take(max_bytes as u64 + 1);
337 let mut out = Vec::new();
338 limited
339 .read_to_end(&mut out)
340 .map_err(|_| StatusCode::BAD_REQUEST)?;
341 if out.len() > max_bytes {
342 return Err(StatusCode::PAYLOAD_TOO_LARGE);
343 }
344 Ok(out)
345}
346
347async fn send_upstream(
348 state: &ProxyState,
349 parts: &Parts,
350 url: &str,
351 body: Vec<u8>,
352 provider_label: &str,
353 preserve_content_encoding: bool,
354) -> Result<reqwest::Response, StatusCode> {
355 let mut req = state.client.request(parts.method.clone(), url);
356
357 for (key, value) in &parts.headers {
358 let k = key.as_str().to_lowercase();
359 if should_forward_request_header(&k, preserve_content_encoding) {
360 req = req.header(key.clone(), value.clone());
361 }
362 }
363
364 req.body(body).send().await.map_err(|e| {
365 tracing::error!("lean-ctx proxy: {provider_label} upstream error: {e}");
366 StatusCode::BAD_GATEWAY
367 })
368}
369
370pub(super) const FORWARDED_HEADERS: &[&str] = &[
371 "content-type",
372 "content-encoding",
373 "mcp-session-id",
374 "x-request-id",
375 "x-oai-request-id",
376 "cf-ray",
377 "x-openai-authorization-error",
378 "x-error-json",
379 "openai-organization",
380 "openai-model",
381 "openai-processing-ms",
382 "openai-version",
383 "x-models-etag",
384 "x-reasoning-included",
385 "anthropic-ratelimit-requests-limit",
386 "anthropic-ratelimit-requests-remaining",
387 "anthropic-ratelimit-tokens-limit",
388 "anthropic-ratelimit-tokens-remaining",
389 "retry-after",
390 "x-ratelimit-limit-requests",
391 "x-ratelimit-remaining-requests",
392 "x-ratelimit-limit-tokens",
393 "x-ratelimit-remaining-tokens",
394 "cache-control",
395];
396
397pub(super) fn is_forwarded_response_header(name: &str) -> bool {
398 FORWARDED_HEADERS.contains(&name)
399 || name.starts_with("x-codex-")
400 || name.starts_with("x-ratelimit-")
401}
402
403async fn build_response(
404 response: reqwest::Response,
405 extra_stream_types: &[&str],
406 usage_provider: super::usage::Provider,
407 url_model: Option<String>,
408 cohort: Option<super::holdout::Arm>,
409) -> Result<Response, StatusCode> {
410 let status = StatusCode::from_u16(response.status().as_u16()).unwrap_or(StatusCode::OK);
411 let resp_headers = response.headers().clone();
412
413 let is_stream = resp_headers
414 .get("content-type")
415 .and_then(|v| v.to_str().ok())
416 .is_some_and(|ct| {
417 ct.contains("text/event-stream") || extra_stream_types.iter().any(|t| ct.contains(t))
418 });
419
420 if is_stream {
421 let scanner = super::usage::Scanner::new(usage_provider, url_model).with_cohort(cohort);
425 let inner = Box::pin(response.bytes_stream());
426 let body = Body::from_stream(super::usage::tee_stream(inner, scanner));
427 let mut resp = Response::builder().status(status);
428 for (k, v) in &resp_headers {
429 let ks = k.as_str().to_lowercase();
430 if is_forwarded_response_header(&ks) {
431 resp = resp.header(k, v);
432 }
433 }
434 return resp
435 .body(body)
436 .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR);
437 }
438
439 let resp_bytes = response
440 .bytes()
441 .await
442 .map_err(|_| StatusCode::BAD_GATEWAY)?;
443
444 let mut scanner = super::usage::Scanner::new(usage_provider, url_model).with_cohort(cohort);
446 scanner.feed_body(&resp_bytes);
447 if let Some(usage) = scanner.finalize() {
448 super::usage_meter::record(&usage);
449 }
450
451 let mut resp = Response::builder().status(status);
452 for (k, v) in &resp_headers {
453 let ks = k.as_str().to_lowercase();
454 if is_forwarded_response_header(&ks) {
455 resp = resp.header(k, v);
456 }
457 }
458 resp.body(Body::from(resp_bytes))
459 .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)
460}
461
462#[cfg(test)]
463mod tests {
464 use super::*;
465
466 fn parts_for(uri: &str) -> Parts {
467 Request::builder().uri(uri).body(()).unwrap().into_parts().0
468 }
469
470 fn add_test_marker(
471 mut value: serde_json::Value,
472 original_size: usize,
473 ) -> (Vec<u8>, usize, usize) {
474 value["lean_ctx_touched"] = serde_json::Value::Bool(true);
475 let out = serde_json::to_vec(&value).unwrap();
476 let compressed_size = out.len();
477 (out, original_size, compressed_size)
478 }
479
480 #[test]
481 fn zstd_request_bodies_are_rewritten_and_reencoded() {
482 let body = serde_json::json!({"model": "gpt-5", "input": []});
483 let json = serde_json::to_vec(&body).unwrap();
484 let encoded = encode_zstd(&json).unwrap();
485 let parts = Request::builder()
486 .uri("/backend-api/codex/responses")
487 .header(axum::http::header::CONTENT_ENCODING, "zstd")
488 .body(())
489 .unwrap()
490 .into_parts()
491 .0;
492
493 let prepared = prepare_request_body(&parts, &encoded, add_test_marker).unwrap();
494 assert_eq!(request_body_encoding(&parts), RequestBodyEncoding::Zstd);
495 assert_eq!(prepared.original_size, json.len());
496 assert!(prepared.compression_candidate);
497 assert!(prepared.preserve_content_encoding);
498 assert!(should_forward_request_header("content-encoding", true));
499 assert!(!should_forward_request_header("content-encoding", false));
500
501 let decoded = zstd::decode_all(prepared.body.as_slice()).unwrap();
502 let parsed: serde_json::Value = serde_json::from_slice(&decoded).unwrap();
503 assert_eq!(parsed["lean_ctx_touched"], true);
504 assert_eq!(parsed["model"], "gpt-5");
505 }
506
507 #[test]
508 fn gzip_request_bodies_are_rewritten_and_reencoded() {
509 let body = serde_json::json!({"model": "gpt-5", "input": []});
510 let json = serde_json::to_vec(&body).unwrap();
511 let encoded = encode_gzip(&json).unwrap();
512 let parts = Request::builder()
513 .uri("/backend-api/codex/responses")
514 .header(axum::http::header::CONTENT_ENCODING, "gzip")
515 .body(())
516 .unwrap()
517 .into_parts()
518 .0;
519
520 let prepared = prepare_request_body(&parts, &encoded, add_test_marker).unwrap();
521 assert_eq!(request_body_encoding(&parts), RequestBodyEncoding::Gzip);
522 assert_eq!(prepared.original_size, json.len());
523 assert!(prepared.compression_candidate);
524 assert!(prepared.preserve_content_encoding);
525
526 let decoded = decode_gzip_bounded(&prepared.body, max_body_bytes()).unwrap();
527 let parsed: serde_json::Value = serde_json::from_slice(&decoded).unwrap();
528 assert_eq!(parsed["lean_ctx_touched"], true);
529 assert_eq!(parsed["model"], "gpt-5");
530 }
531
532 #[test]
533 fn identity_content_encoding_can_be_rewritten_as_json() {
534 let parts = Request::builder()
535 .uri("/v1/responses")
536 .header(axum::http::header::CONTENT_ENCODING, "identity")
537 .body(())
538 .unwrap()
539 .into_parts()
540 .0;
541
542 assert_eq!(request_body_encoding(&parts), RequestBodyEncoding::Identity);
543 }
544
545 #[test]
546 fn unknown_encoded_request_bodies_stay_passthrough() {
547 let parts = Request::builder()
548 .uri("/v1/responses")
549 .header(axum::http::header::CONTENT_ENCODING, "br")
550 .body(())
551 .unwrap()
552 .into_parts()
553 .0;
554 let body = b"not-json";
555
556 let prepared = prepare_request_body(&parts, body, |_, _| {
557 panic!("unknown encodings must not be JSON-rewritten")
558 })
559 .unwrap();
560
561 assert_eq!(
562 request_body_encoding(&parts),
563 RequestBodyEncoding::Passthrough
564 );
565 assert_eq!(prepared.body, body);
566 assert!(prepared.parsed.is_none());
567 assert!(!prepared.compression_candidate);
568 assert!(prepared.preserve_content_encoding);
569 }
570
571 #[test]
572 fn invalid_json_request_bodies_are_not_compression_candidates() {
573 let parts = Request::builder()
574 .uri("/v1/responses")
575 .body(())
576 .unwrap()
577 .into_parts()
578 .0;
579 let body = b"not-json";
580
581 let prepared = prepare_request_body(&parts, body, |_, _| {
582 panic!("invalid JSON must not enter the compression pipeline")
583 })
584 .unwrap();
585
586 assert_eq!(request_body_encoding(&parts), RequestBodyEncoding::Identity);
587 assert_eq!(prepared.body, body);
588 assert!(prepared.parsed.is_none());
589 assert!(!prepared.compression_candidate);
590 assert!(!prepared.preserve_content_encoding);
591 }
592
593 #[test]
594 fn upstream_url_preserves_subpath() {
595 let base = "https://api.anthropic.com";
596 let parts = parts_for("/v1/messages/count_tokens");
597 assert_eq!(
598 build_upstream_url(&parts, base, "/v1/messages"),
599 "https://api.anthropic.com/v1/messages/count_tokens"
600 );
601 }
602
603 #[test]
604 fn upstream_url_preserves_batches_subpath() {
605 let base = "https://api.anthropic.com";
606 let parts = parts_for("/v1/messages/batches/batch_123/results");
607 assert_eq!(
608 build_upstream_url(&parts, base, "/v1/messages"),
609 "https://api.anthropic.com/v1/messages/batches/batch_123/results"
610 );
611 }
612
613 #[test]
614 fn upstream_url_exact_path() {
615 let base = "https://api.anthropic.com";
616 let parts = parts_for("/v1/messages");
617 assert_eq!(
618 build_upstream_url(&parts, base, "/v1/messages"),
619 "https://api.anthropic.com/v1/messages"
620 );
621 }
622
623 #[test]
624 fn upstream_url_preserves_query_params() {
625 let base = "https://api.anthropic.com";
626 let parts = parts_for("/v1/messages/count_tokens?model=claude-4");
627 assert_eq!(
628 build_upstream_url(&parts, base, "/v1/messages"),
629 "https://api.anthropic.com/v1/messages/count_tokens?model=claude-4"
630 );
631 }
632
633 #[test]
634 fn forwards_openai_project_and_auth_headers() {
635 for required in ["authorization", "openai-project", "openai-organization"] {
639 assert!(
640 ALLOWED_REQUEST_HEADERS.contains(&required),
641 "request header `{required}` must be forwarded upstream"
642 );
643 }
644 }
645
646 #[test]
647 fn forwards_chatgpt_codex_oauth_headers() {
648 for required in [
649 "authorization",
650 "chatgpt-account-id",
651 "x-openai-fedramp",
652 "x-openai-internal-codex-residency",
653 "x-openai-product-sku",
654 "oai-product-sku",
655 "x-client-request-id",
656 "x-codex-installation-id",
657 "x-codex-turn-metadata",
658 "x-openai-subagent",
659 "x-codex-turn-state",
660 "originator",
661 ] {
662 assert!(
663 is_allowed_request_header(required),
664 "request header `{required}` must be forwarded upstream"
665 );
666 }
667 }
668
669 #[test]
670 fn forwards_streamable_http_mcp_headers() {
671 for required in ["mcp-session-id", "last-event-id"] {
672 assert!(
673 ALLOWED_REQUEST_HEADERS.contains(&required),
674 "request header `{required}` must be forwarded upstream"
675 );
676 }
677 assert!(
678 is_forwarded_response_header("mcp-session-id"),
679 "MCP session id response header must be forwarded downstream"
680 );
681 }
682
683 #[test]
684 fn forwards_codex_state_response_headers() {
685 for required in [
686 "x-codex-turn-state",
687 "x-codex-primary-used-percent",
688 "openai-model",
689 "x-models-etag",
690 "x-reasoning-included",
691 "x-oai-request-id",
692 "cf-ray",
693 "x-openai-authorization-error",
694 "x-error-json",
695 ] {
696 assert!(
697 is_forwarded_response_header(required),
698 "response header `{required}` must be forwarded downstream"
699 );
700 }
701 }
702}