Skip to main content

lean_ctx/proxy/forward/
mod.rs

1//! Shared upstream forward path for OpenAI-compatible providers.
2
3use 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)] // re-exported for proxy::* and tests
26pub(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// Unit tests import these via `use super::*`.
33#[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
53/// Check whether an incoming request was already compressed by Headroom.
54pub(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
62/// Default request-body ceiling (MiB). A large-codebase refactor with several
63/// big files in context easily exceeds the old 10 MiB cap, which surfaced to the
64/// agent as a hard `400` mid-task. Raised and made configurable via
65/// `LEAN_CTX_PROXY_MAX_BODY_MB`.
66const 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    // Org-policy gate (enterprise#25): under a signed + trusted + enforced org
117    // policy, refuse models outside the ceiling and requests over a hard
118    // budget — before any routing/compression work. No policy → no-op.
119    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    // Active router (enterprise#13): may rewrite `model` in the parsed body
141    // (before compression, so exactly one serialization) and re-target the
142    // upstream within the same wire shape. Fail-open: any miss routes nothing.
143    // An org policy may exempt specific projects from downgrades (#25).
144    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    // Cross-shape translation (enterprise#16) only exists for the exact
155    // messages-create call — count_tokens/batches subpaths have no OpenAI
156    // equivalent and must stay within-shape.
157    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    // Apply the routing decision to the wire: re-target the upstream and — for
191    // registry providers holding their own key — swap the credential headers.
192    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    // #895 Track B: assign output-savings holdout from the same pristine parsed
210    // body that each provider's compressor receives. Only when active.
211    let cohort = parsed
212        .as_ref()
213        .and_then(|p| prepare::cohort_arm(p, provider_label, default_path));
214    if compression_candidate {
215        // Shape label drives compression/routing; stats identity may differ —
216        // Grok registry routes speak OpenAI shape but meter under "Grok".
217        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    // Context Kernel: record identity, coverage, ETPAO for this request.
231    {
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        // Evidence pipeline: proxy data → envelope → normalizer → receipt chain.
257        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    // Cross-shape route (enterprise#16): the body now speaks OpenAI Chat
296    // Completions — address the matching endpoint instead of the caller's
297    // `/v1/messages` path, and scan the response with the OpenAI parser.
298    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    // Measured usage: read the real model + billed tokens from the response.
366    // Gemini puts the model in the URL path, not the request/response body.
367    // Translated requests get OpenAI-shape responses regardless of the label.
368    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    // Gateway context (enterprise#11/#17/#18): identity tags from the auth
380    // guard + wire savings + baseline inputs, stamped onto the usage record.
381    // A routed request is attributed to the provider actually serving it, and
382    // carries the originally requested model as routed_from (enterprise#13).
383    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        // Registry route targets carry their own local-inference flag
397        // (shadow-rate billing); built-in targets keep the URL heuristic.
398        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}