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;
19mod 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        let mut response = Response::builder()
274            .status(cached.status)
275            .body(Body::from(cached.body))
276            .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
277        trace_id::inject_trace_id(&mut response, &trace_id);
278        return Ok(response);
279    }
280    super::cost::record(
281        model,
282        tokens_saved,
283        original_size as u64,
284        compressed_size as u64,
285    );
286
287    // Cross-shape route (enterprise#16): the body now speaks OpenAI Chat
288    // Completions — address the matching endpoint instead of the caller's
289    // `/v1/messages` path, and scan the response with the OpenAI parser.
290    let xlat = route.as_ref().is_some_and(|r| r.xlat);
291    let upstream_url = if xlat {
292        format!("{upstream_base}/v1/chat/completions")
293    } else {
294        crate::proxy::codec::build_upstream_url(&parts, upstream_base, default_path)
295    };
296
297    let counterfactual = if provider_label == "Anthropic" && !xlat {
298        super::counterfactual::maybe_spawn_probe(
299            &state.client,
300            &parts,
301            upstream_base,
302            parsed.as_ref(),
303            route.as_ref().map(|r| r.routed_from.as_str()),
304            compressed_size < original_size,
305        )
306    } else {
307        None
308    };
309
310    let forwarded_body = super::bedrock::finalize_request(
311        provider_label,
312        &mut parts,
313        &body_bytes,
314        prepared.body,
315        body_limit,
316        &upstream_url,
317    )?;
318
319    if let Some(ref pre) = parsed {
320        let cfg_replay = crate::core::config::Config::load();
321        if matches!(
322            cfg_replay.proxy.resolved_proxy_mode(),
323            crate::core::config::ProxyMode::Cache
324        ) {
325            let system_val = pre.get("system");
326            if let Some(msgs) = pre.get("messages").and_then(|m| m.as_array()) {
327                let conv_id = super::prefix_replay::conversation_id(system_val, msgs);
328                super::prefix_replay::record_forwarded(
329                    conv_id,
330                    forwarded_body.clone(),
331                    msgs,
332                    msgs.len(),
333                );
334            }
335        }
336    }
337
338    let response = transport::send_upstream(
339        &state,
340        &parts,
341        &upstream_url,
342        forwarded_body,
343        provider_label,
344        preserve_content_encoding,
345    )
346    .await?;
347
348    // Measured usage: read the real model + billed tokens from the response.
349    // Gemini puts the model in the URL path, not the request/response body.
350    // Translated requests get OpenAI-shape responses regardless of the label.
351    let usage_provider = if xlat {
352        super::usage::Provider::OpenAi
353    } else {
354        super::usage::Provider::from_label(provider_label)
355    };
356    let url_model = if usage_provider == super::usage::Provider::Gemini {
357        super::usage::gemini_model_from_path(parts.uri.path())
358    } else {
359        None
360    };
361
362    // Gateway context (enterprise#11/#17/#18): identity tags from the auth
363    // guard + wire savings + baseline inputs, stamped onto the usage record.
364    // A routed request is attributed to the provider actually serving it, and
365    // carries the originally requested model as routed_from (enterprise#13).
366    let mut wire = prepare::wire_context(
367        &parts,
368        provider_label,
369        upstream_base,
370        tokens_saved,
371        original_size,
372        lineage,
373    );
374    if let Some(route) = &route {
375        wire.routed_from = Some(route.routed_from.clone());
376        if let Some(id) = &route.provider_id {
377            wire.provider = id.clone();
378        }
379        // Registry route targets carry their own local-inference flag
380        // (shadow-rate billing); built-in targets keep the URL heuristic.
381        if let Some(local) = route.local {
382            wire.is_local = local;
383        }
384    }
385    wire.counterfactual = counterfactual;
386    let wire = Some(wire);
387    let mut response = transport::build_response(
388        response,
389        extra_stream_types,
390        usage_provider,
391        url_model,
392        cohort,
393        wire,
394        xlat,
395        state.ocla_cache.as_deref(),
396        model,
397        &cache_prompt_hash,
398    )
399    .await?;
400    trace_id::inject_trace_id(&mut response, &trace_id);
401    Ok(response)
402}