use axum::{
body::Body,
extract::{Path, State},
http::{HeaderMap, HeaderValue, Request, StatusCode},
response::Response,
};
use super::{ProxyState, forward};
use crate::core::config::{ResolvedProvider, WireShape};
use crate::core::ocla::types::{ConnectorJob, OclaResult, ScheduledJob};
pub fn schedule_connector(job: &ConnectorJob, sequence: u64) -> OclaResult<ScheduledJob> {
job.context.validate()?;
if job.connector_id.trim().is_empty() {
return Err(crate::core::ocla::types::OclaError::InvalidRequest(
"connector_id is required".into(),
));
}
let project_root = std::env::current_dir().unwrap_or_else(|_| std::path::PathBuf::from("."));
crate::core::providers::init::init_with_project_root(Some(&project_root));
let available = crate::core::providers::global_registry().available_provider_ids();
let task = format!("{} {}", job.connector_id, job.payload_ref);
let mut bandit =
crate::core::provider_bandit::ProviderBandit::load(project_root.to_str().unwrap_or("."));
let (provider_id, action) = select_provider(&job.connector_id, &task, &available, &mut bandit);
Ok(ScheduledJob {
job_ref: format!("job:{}:{sequence}", job.connector_id),
queue_ref: format!("provider:{provider_id}:{action}"),
})
}
fn select_provider(
requested: &str,
task: &str,
available: &[String],
bandit: &mut crate::core::provider_bandit::ProviderBandit,
) -> (String, String) {
if available.iter().any(|id| id == requested) {
return (requested.to_string(), "dispatch".into());
}
if let Some(prediction) =
crate::core::active_inference::predict_preloads(task, available, bandit, 1)
.into_iter()
.next()
{
return (prediction.provider_id, prediction.action);
}
(requested.to_string(), "dispatch".into())
}
#[derive(Debug, Clone)]
pub(super) struct RegistryProviderId {
pub id: String,
pub local: bool,
}
pub(super) fn is_grok_provider_id(id: &str) -> bool {
matches!(
id.trim().to_ascii_lowercase().as_str(),
"grok-chat" | "xai" | "grok"
)
}
pub(super) fn is_commandcode_provider_id(id: &str) -> bool {
matches!(
id.trim().to_ascii_lowercase().as_str(),
"commandcode" | "command-code"
)
}
pub(super) fn stats_label<'a>(registry_id: Option<&str>, shape_label: &'a str) -> &'a str {
match registry_id {
Some(id) if is_grok_provider_id(id) => "Grok",
Some(id) if is_commandcode_provider_id(id) => "CommandCode",
_ => shape_label,
}
}
pub(super) fn is_openai_responses_path(path: &str) -> bool {
let path = path.trim_end_matches('/');
path == "/responses"
|| path == "/v1/responses"
|| path.starts_with("/responses/")
|| path.starts_with("/v1/responses/")
}
pub async fn handler(
State(state): State<ProxyState>,
Path((id, rest)): Path<(String, String)>,
mut req: Request<Body>,
) -> Result<Response, StatusCode> {
let Some(provider) = state.upstream_snapshot().provider_by_id(&id).cloned() else {
tracing::warn!("lean-ctx proxy: unknown registry provider '{id}' (404)");
return Err(StatusCode::NOT_FOUND);
};
req.extensions_mut().insert(RegistryProviderId {
id: provider.id.clone(),
local: provider.local,
});
let path = format!("/{rest}");
let uri = match req.uri().query() {
Some(q) => format!("{path}?{q}").parse::<axum::http::Uri>(),
None => path.parse::<axum::http::Uri>(),
}
.map_err(|_| StatusCode::BAD_REQUEST)?;
*req.uri_mut() = uri;
if provider.shape == WireShape::Bedrock {
super::bedrock::validate_invoke_request(&req)?;
super::bedrock::attach_signing_context(&provider, &mut req)?;
} else if provider.api_key_env.is_some() {
inject_gateway_credential(&provider, req.headers_mut())?;
}
match provider.shape {
WireShape::Anthropic => {
forward::forward_request(
State(state),
req,
&provider.base_url,
"/v1/messages",
super::bedrock::passthrough_request_body,
"Anthropic",
&[],
)
.await
}
WireShape::OpenAi => {
if is_openai_responses_path(req.uri().path()) {
forward::forward_request(
State(state),
req,
&provider.base_url,
"/v1/responses",
super::openai_responses::compress_request_body,
"OpenAI",
&[],
)
.await
} else {
forward::forward_request(
State(state),
req,
&provider.base_url,
"/v1/chat/completions",
super::openai::compress_request_body,
"OpenAI",
&[],
)
.await
}
}
WireShape::Gemini => {
let model = super::usage::gemini_model_from_path(req.uri().path());
forward::forward_request(
State(state),
req,
&provider.base_url,
"/",
move |body, size| {
super::google::compress_request_body(body, size, model.as_deref())
},
"Gemini",
&["application/x-ndjson"],
)
.await
}
WireShape::Bedrock => {
forward::forward_request(
State(state),
req,
&provider.base_url,
"/",
super::bedrock::passthrough_request_body,
"Bedrock",
&["application/vnd.amazon.eventstream"],
)
.await
}
}
}
pub(super) fn inject_gateway_credential(
provider: &ResolvedProvider,
headers: &mut HeaderMap,
) -> Result<(), StatusCode> {
let env_name = provider
.api_key_env
.as_deref()
.expect("caller checked api_key_env");
let key = std::env::var(env_name)
.ok()
.filter(|k| !k.trim().is_empty());
let Some(key) = key else {
tracing::error!(
"lean-ctx proxy: provider '{}' configures api_key_env='{env_name}' but the \
variable is unset/empty — cannot authenticate upstream (502)",
provider.id
);
return Err(StatusCode::BAD_GATEWAY);
};
for h in ["authorization", "x-api-key", "api-key", "x-goog-api-key"] {
headers.remove(h);
}
let value = |v: String| {
HeaderValue::from_str(&v).map_err(|_| {
tracing::error!(
"lean-ctx proxy: provider '{}' key from {env_name} contains invalid header bytes",
provider.id
);
StatusCode::BAD_GATEWAY
})
};
match provider.shape {
WireShape::Anthropic => {
headers.insert("x-api-key", value(key)?);
}
WireShape::OpenAi => {
headers.insert("api-key", value(key.clone())?);
headers.insert("authorization", value(format!("Bearer {key}"))?);
}
WireShape::Gemini => {
headers.insert("x-goog-api-key", value(key)?);
}
WireShape::Bedrock => unreachable!("Bedrock uses SigV4, not api_key_env"),
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn select_provider_honors_available_explicit_connector() {
let available = vec!["github".to_string(), "jira".to_string()];
let mut bandit = crate::core::provider_bandit::ProviderBandit::new();
assert_eq!(
select_provider("jira", "bug investigation", &available, &mut bandit),
("jira".to_string(), "dispatch".to_string())
);
}
#[test]
fn select_provider_uses_active_inference_priority() {
let available = vec!["github".to_string(), "jira".to_string()];
let mut bandit = crate::core::provider_bandit::ProviderBandit::new();
assert_eq!(
select_provider("automatic", "investigate a bug", &available, &mut bandit),
("github".to_string(), "issues".to_string())
);
}
#[test]
fn stats_label_maps_grok_rails_to_grok_bucket() {
assert_eq!(stats_label(Some("grok-chat"), "OpenAI"), "Grok");
assert_eq!(stats_label(Some("xai"), "OpenAI"), "Grok");
assert_eq!(stats_label(Some("GROK"), "OpenAI"), "Grok");
assert_eq!(stats_label(Some("foundry"), "OpenAI"), "OpenAI");
assert_eq!(stats_label(None, "OpenAI"), "OpenAI");
assert_eq!(stats_label(Some("local"), "OpenAI"), "OpenAI");
}
#[test]
fn is_grok_provider_id_accepts_dual_rail_ids() {
assert!(is_grok_provider_id("grok-chat"));
assert!(is_grok_provider_id("xai"));
assert!(is_grok_provider_id(" Grok "));
assert!(!is_grok_provider_id("openai"));
assert!(!is_grok_provider_id("foundry"));
}
#[test]
fn stats_label_maps_commandcode_rail_to_commandcode_bucket() {
assert_eq!(stats_label(Some("commandcode"), "OpenAI"), "CommandCode");
assert_eq!(stats_label(Some("command-code"), "OpenAI"), "CommandCode");
assert_eq!(stats_label(Some(" CommandCode "), "OpenAI"), "CommandCode");
}
#[test]
fn is_commandcode_provider_id_accepts_rail_ids() {
assert!(is_commandcode_provider_id("commandcode"));
assert!(is_commandcode_provider_id("command-code"));
assert!(is_commandcode_provider_id(" COMMANDCODE "));
assert!(!is_commandcode_provider_id("openai"));
assert!(!is_commandcode_provider_id("grok"));
}
#[test]
fn is_openai_responses_path_detects_responses_api() {
assert!(is_openai_responses_path("/v1/responses"));
assert!(is_openai_responses_path("/v1/responses/"));
assert!(is_openai_responses_path("/responses"));
assert!(is_openai_responses_path(
"/v1/responses/resp_123/input_items"
));
assert!(!is_openai_responses_path("/v1/chat/completions"));
assert!(!is_openai_responses_path("/v1/models"));
assert!(!is_openai_responses_path("/v1/responsesx"));
}
fn provider(shape: WireShape, api_key_env: Option<&str>) -> ResolvedProvider {
ResolvedProvider {
id: "test".into(),
shape,
base_url: "https://example.invalid".into(),
api_key_env: api_key_env.map(str::to_string),
aws_region: None,
local: false,
}
}
#[test]
fn injection_replaces_caller_credentials_per_shape() {
let _lock = crate::core::data_dir::test_env_lock();
crate::test_env::set_var("LC_TEST_PROVIDER_KEY", "sk-upstream");
for (shape, expect_header, expect_value) in [
(WireShape::Anthropic, "x-api-key", "sk-upstream"),
(WireShape::OpenAi, "authorization", "Bearer sk-upstream"),
(WireShape::Gemini, "x-goog-api-key", "sk-upstream"),
] {
let mut headers = HeaderMap::new();
headers.insert("authorization", "Bearer lean-ctx-token".parse().unwrap());
headers.insert("x-api-key", "caller-key".parse().unwrap());
inject_gateway_credential(&provider(shape, Some("LC_TEST_PROVIDER_KEY")), &mut headers)
.expect("key present");
assert_eq!(
headers.get(expect_header).unwrap().to_str().unwrap(),
expect_value,
"{shape:?} must carry the gateway key in its native header"
);
let leaked = headers
.iter()
.any(|(_, v)| v.to_str().is_ok_and(|v| v.contains("lean-ctx-token")));
assert!(!leaked, "caller bearer token leaked upstream for {shape:?}");
if shape != WireShape::Anthropic {
assert!(
headers.get("x-api-key").is_none(),
"stale caller x-api-key must be stripped for {shape:?}"
);
}
}
crate::test_env::remove_var("LC_TEST_PROVIDER_KEY");
}
#[test]
fn openai_shape_also_sets_azure_api_key_header() {
let _lock = crate::core::data_dir::test_env_lock();
crate::test_env::set_var("LC_TEST_PROVIDER_KEY2", "fk-123");
let mut headers = HeaderMap::new();
inject_gateway_credential(
&provider(WireShape::OpenAi, Some("LC_TEST_PROVIDER_KEY2")),
&mut headers,
)
.unwrap();
assert_eq!(headers.get("api-key").unwrap(), "fk-123");
crate::test_env::remove_var("LC_TEST_PROVIDER_KEY2");
}
#[test]
fn missing_key_env_is_a_loud_bad_gateway() {
let _lock = crate::core::data_dir::test_env_lock();
crate::test_env::remove_var("LC_TEST_PROVIDER_KEY_MISSING");
let mut headers = HeaderMap::new();
headers.insert("authorization", "Bearer lean-ctx-token".parse().unwrap());
let err = inject_gateway_credential(
&provider(WireShape::OpenAi, Some("LC_TEST_PROVIDER_KEY_MISSING")),
&mut headers,
)
.unwrap_err();
assert_eq!(err, StatusCode::BAD_GATEWAY);
}
}