1use axum::{
4 body::Body,
5 extract::State,
6 http::{Request, StatusCode},
7 response::Response,
8};
9
10use super::ProxyState;
11use super::connector::schedule_provider_connector;
12use super::intent::classify_and_store_proxy_intent;
13
14#[cfg(feature = "shape-xlat")]
15mod xlat;
16
17mod headers;
18mod prepare;
19pub mod trace_id;
20mod transport;
21
22#[cfg(test)]
23mod tests;
24
25#[allow(unused_imports)] pub(super) use headers::{
27 ALLOWED_REQUEST_HEADERS, FORWARDED_HEADERS, is_allowed_request_header,
28 is_forwarded_response_header,
29};
30pub(super) use transport::xlat_stream_body;
31
32#[cfg(test)]
34#[allow(unused_imports)]
35use super::codec::{
36 RequestBodyEncoding, decode_gzip_bounded, encode_gzip, encode_zstd, is_retryable_status,
37 request_body_encoding,
38};
39#[cfg(test)]
40#[allow(unused_imports)]
41use axum::http::request::Parts;
42#[cfg(test)]
43#[allow(unused_imports)]
44use headers::should_forward_request_header;
45#[cfg(test)]
46#[allow(unused_imports)]
47pub(super) use prepare::{cohort_arm, prepare_request_body, wire_context};
48
49const HEADROOM_COMPRESSED_HEADER: &str = "x-headroom-compressed";
50const OCLA_BUDGET_SCOPE_HEADER: &str = "x-ocla-budget-scope";
51const ESTIMATED_CHARS_PER_TOKEN: u64 = 4;
52
53pub(super) fn is_headroom_compressed(parts: &axum::http::request::Parts) -> bool {
55 parts
56 .headers
57 .get(HEADROOM_COMPRESSED_HEADER)
58 .and_then(|v| v.to_str().ok())
59 .is_some_and(|v| !v.is_empty() && v != "0" && !v.eq_ignore_ascii_case("false"))
60}
61
62const DEFAULT_MAX_BODY_MB: usize = 64;
67
68pub(super) fn max_body_bytes() -> usize {
69 std::env::var("LEAN_CTX_PROXY_MAX_BODY_MB")
70 .ok()
71 .and_then(|v| v.trim().parse::<usize>().ok())
72 .filter(|mb| *mb > 0)
73 .unwrap_or(DEFAULT_MAX_BODY_MB)
74 .saturating_mul(1024 * 1024)
75}
76
77fn apply_ocla_budget_admission(
78 parts: &axum::http::request::Parts,
79 estimated_bytes: usize,
80) -> Result<(), StatusCode> {
81 let Some(scope) = parts
82 .headers
83 .get(OCLA_BUDGET_SCOPE_HEADER)
84 .and_then(|value| value.to_str().ok())
85 .map(str::trim)
86 .filter(|value| !value.is_empty())
87 else {
88 return Ok(());
89 };
90 let estimated_tokens = (estimated_bytes as u64).saturating_add(ESTIMATED_CHARS_PER_TOKEN - 1)
91 / ESTIMATED_CHARS_PER_TOKEN;
92 crate::core::ocla::wire_api::admit_budgeted_request(scope, estimated_tokens, 0.0)
93 .map_err(|_| StatusCode::PAYMENT_REQUIRED)
94}
95
96pub async fn forward_request(
97 State(state): State<ProxyState>,
98 req: Request<Body>,
99 upstream_base: &str,
100 default_path: &str,
101 compress_body: impl FnOnce(serde_json::Value, usize) -> (Vec<u8>, usize, usize),
102 provider_label: &str,
103 extra_stream_types: &[&str],
104) -> Result<Response, StatusCode> {
105 let (mut parts, body) = req.into_parts();
106 let trace_id = trace_id::extract_or_generate_trace_id(&parts.headers);
107 let body_limit = super::bedrock::request_body_limit(&parts).unwrap_or_else(max_body_bytes);
108 let body_bytes = axum::body::to_bytes(body, body_limit)
109 .await
110 .map_err(|_| StatusCode::PAYLOAD_TOO_LARGE)?;
111 let mut lineage = super::lineage::from_trusted_request(&parts, &body_bytes);
112 if let Some(context) = lineage.as_mut() {
113 context.trace_id.clone_from(&trace_id);
114 }
115
116 let gate_rules = super::policy_gate::active_rules();
120 if let Some(rules) = &gate_rules {
121 let tags = parts
122 .extensions
123 .get::<super::gateway_identity::GatewayTags>()
124 .cloned()
125 .unwrap_or_default();
126 let requested_model = prepare::requested_model_of(&parts, &body_bytes);
127 if let Err(refusal) = super::policy_gate::enforce(rules, requested_model.as_deref(), &tags)
128 {
129 tracing::warn!(
130 "lean-ctx gateway: org policy refused request ({refusal:?}) \
131 person={:?} project={:?}",
132 tags.person,
133 tags.project
134 );
135 let mut response = super::policy_gate::refusal_response(&refusal, provider_label);
136 trace_id::inject_trace_id(&mut response, &trace_id);
137 return Ok(response);
138 }
139 }
140 let routing_rules = crate::core::config::Config::load().proxy.routing.clone();
145 let downgrade_forbidden = gate_rules.as_ref().is_some_and(|rules| {
146 let project = parts
147 .extensions
148 .get::<super::gateway_identity::GatewayTags>()
149 .and_then(|t| t.project.clone());
150 super::policy_gate::downgrade_forbidden(rules, project.as_deref())
151 });
152 let route_upstreams =
153 (routing_rules.is_active() && !downgrade_forbidden).then(|| state.upstream_snapshot());
154 let xlat_ok = cfg!(feature = "shape-xlat")
158 && provider_label == "Anthropic"
159 && parts
160 .uri
161 .path()
162 .trim_end_matches('/')
163 .ends_with("/v1/messages");
164 let route_hook = |parsed: &mut serde_json::Value| {
165 route_upstreams.as_ref().and_then(|up| {
166 super::routing::route_request(parsed, provider_label, up, &routing_rules, xlat_ok)
167 })
168 };
169 if is_headroom_compressed(&parts) {
170 super::anthropic::set_headroom_request(true);
171 super::prefix_cache_stats::record_headroom_compat();
172 }
173 let prepared = prepare::prepare_request_body(
174 &parts,
175 &body_bytes,
176 compress_body,
177 route_hook,
178 upstream_base,
179 provider_label == "OpenAI",
180 )?;
181 apply_ocla_budget_admission(&parts, prepared.body.len())?;
182 let original_size = prepared.original_size;
183 let compressed_size = prepared.compressed_size;
184 let compression_candidate = prepared.compression_candidate;
185 let preserve_content_encoding = prepared.preserve_content_encoding;
186 let route = prepared.route;
187 let parsed = prepared.parsed;
188 let _intent_classification =
189 classify_and_store_proxy_intent(&mut parts, parsed.as_ref(), lineage.as_ref(), &body_bytes);
190 let upstream_base = route
193 .as_ref()
194 .and_then(|r| r.upstream_base.as_deref())
195 .unwrap_or(upstream_base);
196 if let Some(provider) = route.as_ref().and_then(|r| r.credential.as_ref()) {
197 super::providers::inject_gateway_credential(provider, &mut parts.headers)?;
198 }
199 schedule_provider_connector(&parts, lineage.as_ref(), route.as_ref(), provider_label);
200 if let Some(ref parsed) = parsed {
201 let provider = match provider_label {
202 "Anthropic" | "Bedrock" => super::introspect::Provider::Anthropic,
203 "OpenAI" | "ChatGPT" => super::introspect::Provider::OpenAi,
204 _ => super::introspect::Provider::Gemini,
205 };
206 let breakdown = super::introspect::analyze_request(parsed, provider);
207 state.introspect.record(breakdown);
208 }
209 let cohort = parsed
212 .as_ref()
213 .and_then(|p| prepare::cohort_arm(p, provider_label, default_path));
214 if compression_candidate {
215 let registry_id = parts
218 .extensions
219 .get::<super::providers::RegistryProviderId>()
220 .map(|r| r.id.as_str());
221 let stats_label = super::providers::stats_label(registry_id, provider_label);
222 state
223 .stats
224 .record_provider_request(stats_label, original_size, compressed_size);
225 }
226
227 let tokens_saved = original_size.saturating_sub(compressed_size) as u64 / 4;
228 super::metrics::record_request(tokens_saved, compressed_size as u64);
229
230 {
232 let proxy_headers: Vec<(String, String)> = parts
233 .headers
234 .iter()
235 .filter_map(|(k, v)| {
236 v.to_str()
237 .ok()
238 .map(|v| (k.as_str().to_owned(), v.to_owned()))
239 })
240 .collect();
241 let kernel_data = crate::core::context_kernel::proxy_bridge::ProxyRequestData {
242 headers: proxy_headers,
243 input_tokens: original_size / 4,
244 output_tokens: 0,
245 tokens_saved: tokens_saved as usize,
246 model: parsed
247 .as_ref()
248 .and_then(|v| v.get("model"))
249 .and_then(|m| m.as_str())
250 .map(String::from),
251 provider: Some(provider_label.to_owned()),
252 request_count: 1,
253 ..Default::default()
254 };
255
256 let kernel_result =
258 crate::core::context_kernel::proxy_bridge::process_proxy_request(&kernel_data);
259 crate::core::context_kernel::envelope_wiring::process_proxy_evidence(
260 &kernel_data,
261 &kernel_result,
262 );
263 }
264
265 let model = parsed
266 .as_ref()
267 .and_then(|v| v.get("model"))
268 .and_then(|m| m.as_str());
269 let cache_prompt_hash = super::ocla_cache_bridge::prompt_hash(&body_bytes);
270 if let (Some(cache), Some(model)) = (&state.ocla_cache, model)
271 && let Some(cached) = cache.try_cache_hit(model, &cache_prompt_hash, 0.0, 0)
272 {
273 if let Some(route_decision) = &route {
274 crate::proxy::routing_feedback::global_feedback().record_outcome_for_decision(
275 &route_decision.decision_id,
276 None,
277 tokens_saved,
278 0,
279 );
280 }
281 let mut response = Response::builder()
282 .status(cached.status)
283 .body(Body::from(cached.body))
284 .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
285 trace_id::inject_trace_id(&mut response, &trace_id);
286 return Ok(response);
287 }
288 super::cost::record(
289 model,
290 tokens_saved,
291 original_size as u64,
292 compressed_size as u64,
293 );
294
295 let xlat = route.as_ref().is_some_and(|r| r.xlat);
299 let upstream_url = if xlat {
300 format!("{upstream_base}/v1/chat/completions")
301 } else {
302 crate::proxy::codec::build_upstream_url(&parts, upstream_base, default_path)
303 };
304
305 let counterfactual = if provider_label == "Anthropic" && !xlat {
306 super::counterfactual::maybe_spawn_probe(
307 &state.client,
308 &parts,
309 upstream_base,
310 parsed.as_ref(),
311 route.as_ref().map(|r| r.routed_from.as_str()),
312 compressed_size < original_size,
313 )
314 } else {
315 None
316 };
317
318 let forwarded_body = super::bedrock::finalize_request(
319 provider_label,
320 &mut parts,
321 &body_bytes,
322 prepared.body,
323 body_limit,
324 &upstream_url,
325 )?;
326
327 if let Some(ref pre) = parsed {
328 let cfg_replay = crate::core::config::Config::load();
329 if matches!(
330 cfg_replay.proxy.resolved_proxy_mode(),
331 crate::core::config::ProxyMode::Cache
332 ) {
333 let system_val = pre.get("system");
334 if let Some(msgs) = pre.get("messages").and_then(|m| m.as_array()) {
335 let conv_id = super::prefix_replay::conversation_id(system_val, msgs);
336 super::prefix_replay::record_forwarded(
337 conv_id,
338 forwarded_body.clone(),
339 msgs,
340 msgs.len(),
341 );
342 }
343 }
344 }
345
346 let response = transport::send_upstream(
347 &state,
348 &parts,
349 &upstream_url,
350 forwarded_body,
351 provider_label,
352 preserve_content_encoding,
353 )
354 .await?;
355
356 if let Some(route_decision) = &route {
357 crate::proxy::routing_feedback::global_feedback().record_outcome_for_decision(
358 &route_decision.decision_id,
359 None,
360 tokens_saved,
361 0,
362 );
363 }
364
365 let usage_provider = if xlat {
369 super::usage::Provider::OpenAi
370 } else {
371 super::usage::Provider::from_label(provider_label)
372 };
373 let url_model = if usage_provider == super::usage::Provider::Gemini {
374 super::usage::gemini_model_from_path(parts.uri.path())
375 } else {
376 None
377 };
378
379 let mut wire = prepare::wire_context(
384 &parts,
385 provider_label,
386 upstream_base,
387 tokens_saved,
388 original_size,
389 lineage,
390 );
391 if let Some(route) = &route {
392 wire.routed_from = Some(route.routed_from.clone());
393 if let Some(id) = &route.provider_id {
394 wire.provider = id.clone();
395 }
396 if let Some(local) = route.local {
399 wire.is_local = local;
400 }
401 }
402 wire.counterfactual = counterfactual;
403 let wire = Some(wire);
404 let mut response = transport::build_response(
405 response,
406 extra_stream_types,
407 usage_provider,
408 url_model,
409 cohort,
410 wire,
411 xlat,
412 state.ocla_cache.as_deref(),
413 model,
414 &cache_prompt_hash,
415 )
416 .await?;
417 trace_id::inject_trace_id(&mut response, &trace_id);
418 Ok(response)
419}