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
17pub(crate) mod enterprise_headers;
18mod headers;
19mod prepare;
20pub mod trace_id;
21mod transport;
22
23#[cfg(test)]
24mod tests;
25
26#[allow(unused_imports)] // re-exported for proxy::* and tests
27pub(super) use headers::{
28    ALLOWED_REQUEST_HEADERS, FORWARDED_HEADERS, is_allowed_request_header,
29    is_forwarded_response_header,
30};
31pub(super) use transport::xlat_stream_body;
32
33// Unit tests import these via `use super::*`.
34#[cfg(test)]
35#[allow(unused_imports)]
36use super::codec::{
37    RequestBodyEncoding, decode_gzip_bounded, encode_gzip, encode_zstd, is_retryable_status,
38    request_body_encoding,
39};
40#[cfg(test)]
41#[allow(unused_imports)]
42use axum::http::request::Parts;
43#[cfg(test)]
44#[allow(unused_imports)]
45use headers::should_forward_request_header;
46#[cfg(test)]
47#[allow(unused_imports)]
48pub(super) use prepare::{cohort_arm, prepare_request_body, wire_context};
49
50const HEADROOM_COMPRESSED_HEADER: &str = "x-headroom-compressed";
51const OCLA_BUDGET_SCOPE_HEADER: &str = "x-ocla-budget-scope";
52const ESTIMATED_CHARS_PER_TOKEN: u64 = 4;
53
54/// Check whether an incoming request was already compressed by Headroom.
55pub(super) fn is_headroom_compressed(parts: &axum::http::request::Parts) -> bool {
56    parts
57        .headers
58        .get(HEADROOM_COMPRESSED_HEADER)
59        .and_then(|v| v.to_str().ok())
60        .is_some_and(|v| !v.is_empty() && v != "0" && !v.eq_ignore_ascii_case("false"))
61}
62
63/// Default request-body ceiling (MiB). A large-codebase refactor with several
64/// big files in context easily exceeds the old 10 MiB cap, which surfaced to the
65/// agent as a hard `400` mid-task. Raised and made configurable via
66/// `LEAN_CTX_PROXY_MAX_BODY_MB`.
67const DEFAULT_MAX_BODY_MB: usize = 64;
68
69pub(super) fn max_body_bytes() -> usize {
70    std::env::var("LEAN_CTX_PROXY_MAX_BODY_MB")
71        .ok()
72        .and_then(|v| v.trim().parse::<usize>().ok())
73        .filter(|mb| *mb > 0)
74        .unwrap_or(DEFAULT_MAX_BODY_MB)
75        .saturating_mul(1024 * 1024)
76}
77
78fn apply_ocla_budget_admission(
79    parts: &axum::http::request::Parts,
80    estimated_bytes: usize,
81) -> Result<(), StatusCode> {
82    let Some(scope) = parts
83        .headers
84        .get(OCLA_BUDGET_SCOPE_HEADER)
85        .and_then(|value| value.to_str().ok())
86        .map(str::trim)
87        .filter(|value| !value.is_empty())
88    else {
89        return Ok(());
90    };
91    let estimated_tokens = (estimated_bytes as u64).saturating_add(ESTIMATED_CHARS_PER_TOKEN - 1)
92        / ESTIMATED_CHARS_PER_TOKEN;
93    crate::core::ocla::wire_api::admit_budgeted_request(scope, estimated_tokens, 0.0)
94        .map_err(|_| StatusCode::PAYMENT_REQUIRED)
95}
96
97pub async fn forward_request(
98    State(state): State<ProxyState>,
99    req: Request<Body>,
100    upstream_base: &str,
101    default_path: &str,
102    compress_body: impl FnOnce(serde_json::Value, usize) -> (Vec<u8>, usize, usize),
103    provider_label: &str,
104    extra_stream_types: &[&str],
105) -> Result<Response, StatusCode> {
106    let (mut parts, body) = req.into_parts();
107    let trace_id = trace_id::extract_or_generate_trace_id(&parts.headers);
108    let body_limit = super::bedrock::request_body_limit(&parts).unwrap_or_else(max_body_bytes);
109    let body_bytes = axum::body::to_bytes(body, body_limit)
110        .await
111        .map_err(|_| StatusCode::PAYLOAD_TOO_LARGE)?;
112    let mut lineage = super::lineage::from_trusted_request(&parts, &body_bytes);
113    if let Some(context) = lineage.as_mut() {
114        context.trace_id.clone_from(&trace_id);
115    }
116
117    // Org-policy gate (enterprise#25): under a signed + trusted + enforced org
118    // policy, refuse models outside the ceiling and requests over a hard
119    // budget — before any routing/compression work. No policy → no-op.
120    #[cfg(feature = "enterprise")]
121    let gate_rules = super::policy_gate::active_rules();
122    #[cfg(feature = "enterprise")]
123    if let Some(rules) = &gate_rules {
124        let tags = parts
125            .extensions
126            .get::<super::gateway_identity::GatewayTags>()
127            .cloned()
128            .unwrap_or_default();
129        let requested_model = prepare::requested_model_of(&parts, &body_bytes);
130        if let Err(refusal) = super::policy_gate::enforce(rules, requested_model.as_deref(), &tags)
131        {
132            tracing::warn!(
133                "lean-ctx gateway: org policy refused request ({refusal:?}) \
134                 person={:?} project={:?}",
135                tags.person,
136                tags.project
137            );
138            let mut response = super::policy_gate::refusal_response(&refusal, provider_label);
139            trace_id::inject_trace_id(&mut response, &trace_id);
140            return Ok(response);
141        }
142    }
143    // Active router (enterprise#13): may rewrite `model` in the parsed body
144    // (before compression, so exactly one serialization) and re-target the
145    // upstream within the same wire shape. Fail-open: any miss routes nothing.
146    // An org policy may exempt specific projects from downgrades (#25).
147    let routing_rules = crate::core::config::Config::load().proxy.routing.clone();
148    #[cfg(feature = "enterprise")]
149    let downgrade_forbidden = gate_rules.as_ref().is_some_and(|rules| {
150        let project = parts
151            .extensions
152            .get::<super::gateway_identity::GatewayTags>()
153            .and_then(|t| t.project.clone());
154        super::policy_gate::downgrade_forbidden(rules, project.as_deref())
155    });
156    #[cfg(not(feature = "enterprise"))]
157    let downgrade_forbidden = false;
158    let route_upstreams =
159        (routing_rules.is_active() && !downgrade_forbidden).then(|| state.upstream_snapshot());
160    // Cross-shape translation (enterprise#16) only exists for the exact
161    // messages-create call — count_tokens/batches subpaths have no OpenAI
162    // equivalent and must stay within-shape.
163    let xlat_ok = cfg!(feature = "shape-xlat")
164        && provider_label == "Anthropic"
165        && parts
166            .uri
167            .path()
168            .trim_end_matches('/')
169            .ends_with("/v1/messages");
170    let route_hook = |parsed: &mut serde_json::Value| {
171        route_upstreams.as_ref().and_then(|up| {
172            super::routing::route_request(parsed, provider_label, up, &routing_rules, xlat_ok)
173        })
174    };
175    if is_headroom_compressed(&parts) {
176        super::anthropic::set_headroom_request(true);
177        super::prefix_cache_stats::record_headroom_compat();
178    }
179    let prepared = prepare::prepare_request_body(
180        &parts,
181        &body_bytes,
182        compress_body,
183        route_hook,
184        upstream_base,
185        provider_label == "OpenAI",
186    )?;
187    apply_ocla_budget_admission(&parts, prepared.body.len())?;
188    let original_size = prepared.original_size;
189    let compressed_size = prepared.compressed_size;
190    let compression_candidate = prepared.compression_candidate;
191    let preserve_content_encoding = prepared.preserve_content_encoding;
192    let route = prepared.route;
193    let parsed = prepared.parsed;
194    let intent_classification =
195        classify_and_store_proxy_intent(&mut parts, parsed.as_ref(), lineage.as_ref(), &body_bytes);
196    // Apply the routing decision to the wire: re-target the upstream and — for
197    // registry providers holding their own key — swap the credential headers.
198    let upstream_base = route
199        .as_ref()
200        .and_then(|r| r.upstream_base.as_deref())
201        .unwrap_or(upstream_base);
202    if let Some(provider) = route.as_ref().and_then(|r| r.credential.as_ref()) {
203        super::providers::inject_gateway_credential(provider, &mut parts.headers)?;
204    }
205    schedule_provider_connector(&parts, lineage.as_ref(), route.as_ref(), provider_label);
206    if let Some(ref parsed) = parsed {
207        let provider = match provider_label {
208            "Anthropic" | "Bedrock" => super::introspect::Provider::Anthropic,
209            "OpenAI" | "ChatGPT" => super::introspect::Provider::OpenAi,
210            _ => super::introspect::Provider::Gemini,
211        };
212        let breakdown = super::introspect::analyze_request(parsed, provider);
213        state.introspect.record(breakdown);
214    }
215    // #895 Track B: assign output-savings holdout from the same pristine parsed
216    // body that each provider's compressor receives. Only when active.
217    let cohort = parsed
218        .as_ref()
219        .and_then(|p| prepare::cohort_arm(p, provider_label, default_path));
220    if compression_candidate {
221        // Shape label drives compression/routing; stats identity may differ —
222        // Grok registry routes speak OpenAI shape but meter under "Grok".
223        let registry_id = parts
224            .extensions
225            .get::<super::providers::RegistryProviderId>()
226            .map(|r| r.id.as_str());
227        let stats_label = super::providers::stats_label(registry_id, provider_label);
228        state
229            .stats
230            .record_provider_request(stats_label, original_size, compressed_size);
231    }
232
233    let tokens_saved = original_size.saturating_sub(compressed_size) as u64 / 4;
234    super::metrics::record_request(tokens_saved, compressed_size as u64);
235
236    // Context Kernel: record identity, coverage, ETPAO for this request.
237    {
238        let proxy_headers: Vec<(String, String)> = parts
239            .headers
240            .iter()
241            .filter_map(|(k, v)| {
242                v.to_str()
243                    .ok()
244                    .map(|v| (k.as_str().to_owned(), v.to_owned()))
245            })
246            .collect();
247        let kernel_data = crate::core::context_kernel::proxy_bridge::ProxyRequestData {
248            headers: proxy_headers,
249            input_tokens: original_size / 4,
250            output_tokens: 0,
251            tokens_saved: tokens_saved as usize,
252            model: parsed
253                .as_ref()
254                .and_then(|v| v.get("model"))
255                .and_then(|m| m.as_str())
256                .map(String::from),
257            provider: Some(provider_label.to_owned()),
258            request_count: 1,
259            ..Default::default()
260        };
261
262        // Evidence pipeline: proxy data → envelope → normalizer → receipt chain.
263        let kernel_result =
264            crate::core::context_kernel::proxy_bridge::process_proxy_request(&kernel_data);
265        crate::core::context_kernel::envelope_wiring::process_proxy_evidence(
266            &kernel_data,
267            &kernel_result,
268        );
269    }
270
271    let model = parsed
272        .as_ref()
273        .and_then(|v| v.get("model"))
274        .and_then(|m| m.as_str());
275    let cache_prompt_hash = super::ocla_cache_bridge::prompt_hash(&body_bytes);
276    if let (Some(cache), Some(model)) = (&state.ocla_cache, model)
277        && let Some(cached) = cache.try_cache_hit(model, &cache_prompt_hash, 0.0, 0)
278    {
279        if let Some(route_decision) = &route {
280            crate::proxy::routing_feedback::global_feedback().record_outcome_for_decision(
281                &route_decision.decision_id,
282                None,
283                tokens_saved,
284                0,
285            );
286        }
287        let mut response = Response::builder()
288            .status(cached.status)
289            .body(Body::from(cached.body))
290            .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
291        trace_id::inject_trace_id(&mut response, &trace_id);
292        return Ok(response);
293    }
294    super::cost::record(
295        model,
296        tokens_saved,
297        original_size as u64,
298        compressed_size as u64,
299    );
300
301    // Cross-shape route (enterprise#16): the body now speaks OpenAI Chat
302    // Completions — address the matching endpoint instead of the caller's
303    // `/v1/messages` path, and scan the response with the OpenAI parser.
304    let xlat = route.as_ref().is_some_and(|r| r.xlat);
305    let upstream_url = if xlat {
306        format!("{upstream_base}/v1/chat/completions")
307    } else {
308        crate::proxy::codec::build_upstream_url(&parts, upstream_base, default_path)
309    };
310
311    let counterfactual = if provider_label == "Anthropic" && !xlat {
312        super::counterfactual::maybe_spawn_probe(
313            &state.client,
314            &parts,
315            upstream_base,
316            parsed.as_ref(),
317            route.as_ref().map(|r| r.routed_from.as_str()),
318            compressed_size < original_size,
319        )
320    } else {
321        None
322    };
323
324    let forwarded_body = super::bedrock::finalize_request(
325        provider_label,
326        &mut parts,
327        &body_bytes,
328        prepared.body,
329        body_limit,
330        &upstream_url,
331    )?;
332
333    if let Some(ref pre) = parsed {
334        let cfg_replay = crate::core::config::Config::load();
335        if matches!(
336            cfg_replay.proxy.resolved_proxy_mode(),
337            crate::core::config::ProxyMode::Cache
338        ) {
339            let system_val = pre.get("system");
340            if let Some(msgs) = pre.get("messages").and_then(|m| m.as_array()) {
341                let conv_id = super::prefix_replay::conversation_id(system_val, msgs);
342                super::prefix_replay::record_forwarded(
343                    conv_id,
344                    forwarded_body.clone(),
345                    msgs,
346                    msgs.len(),
347                );
348            }
349        }
350    }
351
352    // Enterprise Suite: inject x-leanctx-* metadata headers before dispatch.
353    {
354        let enterprise_cfg = crate::core::config::Config::load().enterprise.clone();
355        if enterprise_cfg.should_inject_headers() {
356            let agent_id = lineage.as_ref().map(|l| l.agent_id.clone());
357            let session_id = lineage.as_ref().map(|l| l.session_id.clone());
358            let task_class = intent_classification
359                .as_ref()
360                .map(|ic| ic._decision.intent.clone());
361            let meta = enterprise_headers::RuntimeMetadata {
362                original_tokens: original_size / 4,
363                compressed_tokens: compressed_size / 4,
364                agent_id,
365                session_id,
366                task_class,
367            };
368            enterprise_headers::inject(&mut parts, &enterprise_cfg, &meta);
369        }
370    }
371    let response = transport::send_upstream(
372        &state,
373        &parts,
374        &upstream_url,
375        forwarded_body,
376        provider_label,
377        preserve_content_encoding,
378    )
379    .await?;
380
381    if let Some(route_decision) = &route {
382        crate::proxy::routing_feedback::global_feedback().record_outcome_for_decision(
383            &route_decision.decision_id,
384            None,
385            tokens_saved,
386            0,
387        );
388    }
389
390    // Measured usage: read the real model + billed tokens from the response.
391    // Gemini puts the model in the URL path, not the request/response body.
392    // Translated requests get OpenAI-shape responses regardless of the label.
393    let usage_provider = if xlat {
394        super::usage::Provider::OpenAi
395    } else {
396        super::usage::Provider::from_label(provider_label)
397    };
398    let url_model = if usage_provider == super::usage::Provider::Gemini {
399        super::usage::gemini_model_from_path(parts.uri.path())
400    } else {
401        None
402    };
403
404    // Gateway context (enterprise#11/#17/#18): identity tags from the auth
405    // guard + wire savings + baseline inputs, stamped onto the usage record.
406    // A routed request is attributed to the provider actually serving it, and
407    // carries the originally requested model as routed_from (enterprise#13).
408    let mut wire = prepare::wire_context(
409        &parts,
410        provider_label,
411        upstream_base,
412        tokens_saved,
413        original_size,
414        lineage,
415    );
416    if let Some(route) = &route {
417        wire.routed_from = Some(route.routed_from.clone());
418        if let Some(id) = &route.provider_id {
419            wire.provider = id.clone();
420        }
421        // Registry route targets carry their own local-inference flag
422        // (shadow-rate billing); built-in targets keep the URL heuristic.
423        if let Some(local) = route.local {
424            wire.is_local = local;
425        }
426    }
427    wire.counterfactual = counterfactual;
428    let wire = Some(wire);
429    let mut response = transport::build_response(
430        response,
431        extra_stream_types,
432        usage_provider,
433        url_model,
434        cohort,
435        wire,
436        xlat,
437        state.ocla_cache.as_deref(),
438        model,
439        &cache_prompt_hash,
440    )
441    .await?;
442    trace_id::inject_trace_id(&mut response, &trace_id);
443    Ok(response)
444}