From 86d01d29fdf2742fdcd7b23f3dac5a4e0ee5a779 Mon Sep 17 00:00:00 2001 From: whit3rabbit Date: Mon, 13 Jul 2026 06:11:40 -0500 Subject: [PATCH] refactor: split large files into modules - Split config.rs into config/{mod,delete,get,put}.rs - Split routes_api.rs into routes_api/{mod,helpers,providers,routes}.rs - Split passthrough/handlers.rs into handlers/{mod,errors,generic,messages}.rs - Split streaming.rs into streaming/{mod,handler,helpers}.rs - Split main_helpers/async_main/admin.rs into admin/{mod,config,tasks}.rs - Minor cleanups in chat_completions backends, token_counting, tests - Add docs/TEST_PARITY_LITELLM.md Co-Authored-By: Claude --- .../proxy/src/admin/routes/config/delete.rs | 137 ++++ crates/proxy/src/admin/routes/config/get.rs | 100 +++ crates/proxy/src/admin/routes/config/mod.rs | 7 + .../admin/routes/{config.rs => config/put.rs} | 272 +------ crates/proxy/src/admin/routes/routes_api.rs | 562 -------------- .../src/admin/routes/routes_api/helpers.rs | 179 +++++ .../proxy/src/admin/routes/routes_api/mod.rs | 9 + .../src/admin/routes/routes_api/providers.rs | 233 ++++++ .../src/admin/routes/routes_api/routes.rs | 158 ++++ crates/proxy/src/config/litellm/provider.rs | 166 +++++ .../src/main_helpers/async_main/admin.rs | 696 ------------------ .../main_helpers/async_main/admin/config.rs | 306 ++++++++ .../src/main_helpers/async_main/admin/mod.rs | 235 ++++++ .../main_helpers/async_main/admin/tasks.rs | 187 +++++ .../chat_completions/backend_anthropic.rs | 168 +++++ .../server/chat_completions/backend_openai.rs | 204 +++++ .../chat_completions/backend_responses.rs | 110 +++ .../src/server/passthrough/handlers/errors.rs | 42 ++ .../server/passthrough/handlers/generic.rs | 124 ++++ .../{handlers.rs => handlers/messages.rs} | 239 +----- .../src/server/passthrough/handlers/mod.rs | 6 + .../{streaming.rs => streaming/handler.rs} | 263 +------ crates/proxy/src/server/streaming/helpers.rs | 212 ++++++ crates/proxy/src/server/streaming/mod.rs | 11 + crates/proxy/src/server/streaming/tests.rs | 7 +- crates/proxy/src/server/token_counting.rs | 180 +++++ crates/proxy/tests/chat_completions.rs | 566 ++++++++++++++ crates/proxy/tests/error_fixtures.rs | 203 +++++ .../src/mapping/streaming_map/tests.rs | 92 +++ docs/TEST_PARITY_LITELLM.md | 291 ++++++++ 30 files changed, 3950 insertions(+), 2015 deletions(-) create mode 100644 crates/proxy/src/admin/routes/config/delete.rs create mode 100644 crates/proxy/src/admin/routes/config/get.rs create mode 100644 crates/proxy/src/admin/routes/config/mod.rs rename crates/proxy/src/admin/routes/{config.rs => config/put.rs} (51%) delete mode 100644 crates/proxy/src/admin/routes/routes_api.rs create mode 100644 crates/proxy/src/admin/routes/routes_api/helpers.rs create mode 100644 crates/proxy/src/admin/routes/routes_api/mod.rs create mode 100644 crates/proxy/src/admin/routes/routes_api/providers.rs create mode 100644 crates/proxy/src/admin/routes/routes_api/routes.rs create mode 100644 crates/proxy/src/config/litellm/provider.rs delete mode 100644 crates/proxy/src/main_helpers/async_main/admin.rs create mode 100644 crates/proxy/src/main_helpers/async_main/admin/config.rs create mode 100644 crates/proxy/src/main_helpers/async_main/admin/mod.rs create mode 100644 crates/proxy/src/main_helpers/async_main/admin/tasks.rs create mode 100644 crates/proxy/src/server/chat_completions/backend_anthropic.rs create mode 100644 crates/proxy/src/server/chat_completions/backend_openai.rs create mode 100644 crates/proxy/src/server/chat_completions/backend_responses.rs create mode 100644 crates/proxy/src/server/passthrough/handlers/errors.rs create mode 100644 crates/proxy/src/server/passthrough/handlers/generic.rs rename crates/proxy/src/server/passthrough/{handlers.rs => handlers/messages.rs} (59%) create mode 100644 crates/proxy/src/server/passthrough/handlers/mod.rs rename crates/proxy/src/server/{streaming.rs => streaming/handler.rs} (57%) create mode 100644 crates/proxy/src/server/streaming/helpers.rs create mode 100644 crates/proxy/src/server/streaming/mod.rs create mode 100644 docs/TEST_PARITY_LITELLM.md diff --git a/crates/proxy/src/admin/routes/config/delete.rs b/crates/proxy/src/admin/routes/config/delete.rs new file mode 100644 index 0000000..193fb70 --- /dev/null +++ b/crates/proxy/src/admin/routes/config/delete.rs @@ -0,0 +1,137 @@ +use crate::admin::state::SharedState; +use axum::{ + extract::{ConnectInfo, Path, State}, + http::StatusCode, + response::IntoResponse, + Json, +}; +use std::net::SocketAddr; + +/// DELETE /admin/api/config/overrides/:key -- remove a single override. +pub(crate) async fn delete_config_override( + ConnectInfo(addr): ConnectInfo, + State(shared): State, + Path(key): Path, +) -> impl IntoResponse { + let key_clone = key.clone(); + match crate::admin::state::with_db(&shared.db, move |conn| { + crate::admin::db::delete_config_override(conn, &key_clone) + }) + .await + { + Some(Ok(true)) => { + let restored = match key.as_str() { + "redact_secrets" => { + let env_default = shared.runtime_defaults.redact_secrets; + if let Ok(mut config) = shared.runtime_config.write() { + config.redact_secrets = env_default; + } + Some(env_default.to_string()) + } + "log_bodies" => { + let env_default = shared.runtime_defaults.log_bodies; + if let Ok(mut config) = shared.runtime_config.write() { + config.log_bodies = env_default; + } + Some(env_default.to_string()) + } + "anthropic_thinking_repair" => { + let env_default = shared.runtime_defaults.anthropic_thinking_repair; + if let Ok(mut config) = shared.runtime_config.write() { + config.anthropic_thinking_repair = env_default; + } + Some(env_default.to_string()) + } + "pxpipe_compress" => { + let env_default = shared.runtime_defaults.pxpipe_compress; + if let Ok(mut config) = shared.runtime_config.write() { + config.pxpipe_compress = env_default; + } + Some(env_default.to_string()) + } + "rtk_compress" => { + let env_default = shared.runtime_defaults.rtk_compress; + if let Ok(mut config) = shared.runtime_config.write() { + config.rtk_compress = env_default; + } + Some(env_default.to_string()) + } + "rtk_models" => { + let env_default = shared.runtime_defaults.rtk_models.clone(); + if let Ok(mut config) = shared.runtime_config.write() { + config.rtk_models = env_default.clone(); + } + Some(env_default) + } + "pxpipe_models" => { + let env_default = shared.runtime_defaults.pxpipe_models.clone(); + if let Ok(mut config) = shared.runtime_config.write() { + config.pxpipe_models = env_default.clone(); + } + Some(env_default) + } + "forward_client_auth" => { + let env_default = shared.runtime_defaults.forward_client_auth; + if let Ok(mut config) = shared.runtime_config.write() { + config.forward_client_auth = env_default; + } + Some(env_default.to_string()) + } + "tool_guardrail_mode" => { + let env_default = shared.runtime_defaults.tool_guardrail_mode.clone(); + if let Ok(mut config) = shared.runtime_config.write() { + config.tool_guardrail_mode = env_default.clone(); + } + Some(env_default) + } + "optimizer_mode" => { + let env_default = shared.runtime_defaults.optimizer_mode.clone(); + if let Ok(mut config) = shared.runtime_config.write() { + config.optimizer_mode = env_default.clone(); + } + Some(env_default) + } + _ => None, + }; + if let Some(value) = restored { + let _ = shared + .events_tx + .send(crate::admin::state::AdminEvent::ConfigChanged { + key: key.clone(), + value, + }); + } + super::super::emit_audit( + &shared, + crate::admin::db::AuditEntry { + id: None, + timestamp: None, + action: "config_deleted".into(), + target_type: "config".into(), + target_id: Some(key.clone()), + detail: None, + source_ip: Some(addr.ip().to_string()), + }, + ); + (StatusCode::OK, Json(serde_json::json!({"deleted": key}))).into_response() + } + Some(Ok(false)) => ( + StatusCode::NOT_FOUND, + Json(serde_json::json!({"error": "override not found"})), + ) + .into_response(), + Some(Err(e)) => { + tracing::error!(error = %e, "delete_config_override failed"); + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": "internal database error"})), + ) + .into_response() + } + None => ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": "internal error"})), + ) + .into_response(), + } +} diff --git a/crates/proxy/src/admin/routes/config/get.rs b/crates/proxy/src/admin/routes/config/get.rs new file mode 100644 index 0000000..c9b6a99 --- /dev/null +++ b/crates/proxy/src/admin/routes/config/get.rs @@ -0,0 +1,100 @@ +use crate::admin::state::SharedState; +use axum::{extract::State, Json}; + +/// GET /admin/api/config -- effective config (env defaults + overrides). +pub(crate) async fn get_config(State(shared): State) -> Json { + let ( + log_level, + log_bodies, + redact_secrets, + anthropic_thinking_repair, + pxpipe_compress, + pxpipe_models, + rtk_compress, + rtk_models, + forward_client_auth, + tool_guardrail_mode, + optimizer_mode, + backends, + ) = { + let config = shared + .runtime_config + .read() + .unwrap_or_else(|e| e.into_inner()); + let mut backends = serde_json::Map::new(); + for (name, mapping) in &config.model_mappings { + backends.insert( + name.clone(), + serde_json::json!({ + "big_model": mapping.big_model, + "small_model": mapping.small_model, + }), + ); + } + ( + config.log_level.clone(), + config.log_bodies, + config.redact_secrets, + config.anthropic_thinking_repair, + config.pxpipe_compress, + config.pxpipe_models.clone(), + config.rtk_compress, + config.rtk_models.clone(), + config.forward_client_auth, + config.tool_guardrail_mode.clone(), + config.optimizer_mode.clone(), + backends, + ) + }; + + let overrides = crate::admin::state::with_db(&shared.db, |conn| { + crate::admin::db::get_config_overrides(conn).unwrap_or_default() + }) + .await + .unwrap_or_default(); + let override_keys: Vec = overrides.iter().map(|(k, _, _)| k.clone()).collect(); + + let pxpipe_available_models = + crate::pxpipe::available_vision_models(&shared.provider_catalog, &pxpipe_models); + + Json(serde_json::json!({ + "log_level": log_level, + "log_bodies": log_bodies, + "redact_secrets": redact_secrets, + "anthropic_thinking_repair": anthropic_thinking_repair, + "pxpipe_compress": pxpipe_compress, + "pxpipe_models": pxpipe_models, + "pxpipe_available_models": pxpipe_available_models, + "rtk_compress": rtk_compress, + "rtk_models": rtk_models, + "forward_client_auth": forward_client_auth, + "tool_guardrail_mode": tool_guardrail_mode, + "optimizer_mode": optimizer_mode, + "backends": backends, + "overridden_keys": override_keys, + })) +} + +/// GET /admin/api/config/overrides -- only SQLite overrides. +pub(crate) async fn get_config_overrides( + State(shared): State, +) -> Json { + let overrides = crate::admin::state::with_db(&shared.db, |conn| { + crate::admin::db::get_config_overrides(conn).unwrap_or_default() + }) + .await + .unwrap_or_default(); + + let entries: Vec = overrides + .into_iter() + .map(|(k, v, updated_at)| { + serde_json::json!({ + "key": k, + "value": v, + "updated_at": updated_at, + }) + }) + .collect(); + + Json(serde_json::json!({ "overrides": entries })) +} diff --git a/crates/proxy/src/admin/routes/config/mod.rs b/crates/proxy/src/admin/routes/config/mod.rs new file mode 100644 index 0000000..7bb1aff --- /dev/null +++ b/crates/proxy/src/admin/routes/config/mod.rs @@ -0,0 +1,7 @@ +mod delete; +mod get; +mod put; + +pub(crate) use delete::delete_config_override; +pub(crate) use get::{get_config, get_config_overrides}; +pub(crate) use put::put_config; diff --git a/crates/proxy/src/admin/routes/config.rs b/crates/proxy/src/admin/routes/config/put.rs similarity index 51% rename from crates/proxy/src/admin/routes/config.rs rename to crates/proxy/src/admin/routes/config/put.rs index 5fdee23..ed06754 100644 --- a/crates/proxy/src/admin/routes/config.rs +++ b/crates/proxy/src/admin/routes/config/put.rs @@ -1,107 +1,21 @@ use crate::admin::state::SharedState; use axum::{ - extract::{ConnectInfo, Path, State}, + extract::{ConnectInfo, State}, http::StatusCode, response::IntoResponse, Json, }; use std::net::SocketAddr; -/// GET /admin/api/config -- effective config (env defaults + overrides). -pub(super) async fn get_config(State(shared): State) -> Json { - // Clone config snapshot and drop the read guard before any .await points. - // std::sync::RwLockReadGuard is !Send, cannot be held across awaits. - let ( - log_level, - log_bodies, - redact_secrets, - anthropic_thinking_repair, - pxpipe_compress, - pxpipe_models, - rtk_compress, - rtk_models, - forward_client_auth, - tool_guardrail_mode, - optimizer_mode, - backends, - ) = { - let config = shared - .runtime_config - .read() - .unwrap_or_else(|e| e.into_inner()); - let mut backends = serde_json::Map::new(); - for (name, mapping) in &config.model_mappings { - backends.insert( - name.clone(), - serde_json::json!({ - "big_model": mapping.big_model, - "small_model": mapping.small_model, - }), - ); - } - ( - config.log_level.clone(), - config.log_bodies, - config.redact_secrets, - config.anthropic_thinking_repair, - config.pxpipe_compress, - config.pxpipe_models.clone(), - config.rtk_compress, - config.rtk_models.clone(), - config.forward_client_auth, - config.tool_guardrail_mode.clone(), - config.optimizer_mode.clone(), - backends, - ) - }; - - // Get overrides to mark which fields are overridden. - let overrides = crate::admin::state::with_db(&shared.db, |conn| { - crate::admin::db::get_config_overrides(conn).unwrap_or_default() - }) - .await - .unwrap_or_default(); - let override_keys: Vec = overrides.iter().map(|(k, _, _)| k.clone()).collect(); - - // Computed before the json! macro consumes `pxpipe_models` by move. - let pxpipe_available_models = - crate::pxpipe::available_vision_models(&shared.provider_catalog, &pxpipe_models); - - Json(serde_json::json!({ - "log_level": log_level, - "log_bodies": log_bodies, - "redact_secrets": redact_secrets, - "anthropic_thinking_repair": anthropic_thinking_repair, - "pxpipe_compress": pxpipe_compress, - "pxpipe_models": pxpipe_models, - // Vision-capable Claude models (+ current out-of-list scope entries) the - // UI offers as per-model scope toggles. - "pxpipe_available_models": pxpipe_available_models, - "rtk_compress": rtk_compress, - "rtk_models": rtk_models, - "forward_client_auth": forward_client_auth, - "tool_guardrail_mode": tool_guardrail_mode, - "optimizer_mode": optimizer_mode, - "backends": backends, - "overridden_keys": override_keys, - })) -} - /// PUT /admin/api/config -- update config overrides. Partial JSON body. -pub(super) async fn put_config( +pub(crate) async fn put_config( ConnectInfo(addr): ConnectInfo, State(shared): State, Json(body): Json, ) -> impl IntoResponse { - // Collect the key-value pairs to persist, then do all SQLite I/O - // before touching in-memory state. This avoids holding the async - // MutexGuard across block_in_place. let mut db_writes: Vec<(String, String)> = Vec::new(); if let Some(level) = body.get("log_level").and_then(|v| v.as_str()) { - // Allowlist: trace-level logging exposes HTTP headers (including API - // keys) in log output. Arbitrary filter directives could also be used - // to selectively leak data. Restrict to safe levels only. const ALLOWED_LOG_LEVELS: &[&str] = &["error", "warn", "info", "debug"]; let normalized = level.trim().to_lowercase(); if !ALLOWED_LOG_LEVELS.contains(&normalized.as_str()) { @@ -161,21 +75,10 @@ pub(super) async fn put_config( db_writes.push(("rtk_compress".to_string(), val.to_string())); } if let Some(val) = body.get("rtk_models").and_then(|v| v.as_str()) { - // Normalize the CSV (trim entries, drop empties); empty = all models. let normalized = crate::config::helpers::normalize_csv(val); db_writes.push(("rtk_models".to_string(), normalized)); } if let Some(val) = body.get("forward_client_auth").and_then(|v| v.as_bool()) { - // Same rule enforced at startup (main_helpers::async_main) for - // statically-configured backends: 2+ distinct PROXY_API_KEYS entries - // with no PROXY_OPEN_RELAY would let different callers each redirect - // the upstream Anthropic credential. This is the only path that can - // enable the toggle *after* boot (live, no restart), so it must be - // re-checked here -- otherwise this admin route would silently - // reopen exactly the misconfiguration the startup panic exists to - // block. Reads the same ALLOWED_KEY_HASHES/OPEN_RELAY statics - // validate_auth uses, not a re-parse of the env vars, so it can't - // diverge from what a request actually experiences. if val && crate::server::middleware::forward_client_auth_misconfigured( crate::server::middleware::distinct_static_key_count(), @@ -223,7 +126,6 @@ pub(super) async fn put_config( } } if let Some(backends) = body.get("backends").and_then(|v| v.as_object()) { - // Read current config to validate backend names exist let config = shared .runtime_config .read() @@ -231,7 +133,7 @@ pub(super) async fn put_config( for (name, settings) in backends { if config.model_mappings.contains_key(name) { if let Some(big) = settings.get("big_model").and_then(|v| v.as_str()) { - if !super::is_safe_model_name(big) { + if !super::super::is_safe_model_name(big) { return ( StatusCode::BAD_REQUEST, Json(serde_json::json!({ @@ -243,7 +145,7 @@ pub(super) async fn put_config( db_writes.push((format!("{name}.big_model"), big.to_string())); } if let Some(small) = settings.get("small_model").and_then(|v| v.as_str()) { - if !super::is_safe_model_name(small) { + if !super::super::is_safe_model_name(small) { return ( StatusCode::BAD_REQUEST, Json(serde_json::json!({ @@ -258,14 +160,8 @@ pub(super) async fn put_config( } } - // Serialize config writes so concurrent requests cannot interleave - // Phase 1 (SQLite) and Phase 2 (in-memory), which would leave them - // inconsistent. let _config_guard = shared.config_write_lock.lock().await; - // Phase 1: Persist to SQLite first. If the process crashes between - // phases, the database is the source of truth and config is restored - // on restart. Reversing the order would lose updates on crash. { let writes = db_writes.clone(); crate::admin::state::with_db(&shared.db, move |conn| { @@ -276,14 +172,12 @@ pub(super) async fn put_config( .await; } - // Phase 2: Apply to in-memory config (no async lock held) { let mut config = shared .runtime_config .write() .unwrap_or_else(|e| e.into_inner()); - // Audit log: capture old values before applying changes for (key, new_value) in &db_writes { let old_value = match key.as_str() { "log_level" => config.log_level.clone(), @@ -378,7 +272,6 @@ pub(super) async fn put_config( drop(_config_guard); - // Broadcast config changes. for (key, value) in &db_writes { let _ = shared .events_tx @@ -386,7 +279,7 @@ pub(super) async fn put_config( key: key.clone(), value: value.clone(), }); - super::emit_audit( + super::super::emit_audit( &shared, crate::admin::db::AuditEntry { id: None, @@ -409,158 +302,3 @@ pub(super) async fn put_config( ) .into_response() } - -/// GET /admin/api/config/overrides -- only SQLite overrides. -pub(super) async fn get_config_overrides( - State(shared): State, -) -> Json { - let overrides = crate::admin::state::with_db(&shared.db, |conn| { - crate::admin::db::get_config_overrides(conn).unwrap_or_default() - }) - .await - .unwrap_or_default(); - - let entries: Vec = overrides - .into_iter() - .map(|(k, v, updated_at)| { - serde_json::json!({ - "key": k, - "value": v, - "updated_at": updated_at, - }) - }) - .collect(); - - Json(serde_json::json!({ "overrides": entries })) -} - -/// DELETE /admin/api/config/overrides/:key -- remove a single override. -pub(super) async fn delete_config_override( - ConnectInfo(addr): ConnectInfo, - State(shared): State, - Path(key): Path, -) -> impl IntoResponse { - let key_clone = key.clone(); - match crate::admin::state::with_db(&shared.db, move |conn| { - crate::admin::db::delete_config_override(conn, &key_clone) - }) - .await - { - Some(Ok(true)) => { - // Restore the runtime value to its pre-override default. Without this, - // deleting an override leaves the overridden value live until restart. - let restored = match key.as_str() { - "redact_secrets" => { - let env_default = shared.runtime_defaults.redact_secrets; - if let Ok(mut config) = shared.runtime_config.write() { - config.redact_secrets = env_default; - } - Some(env_default.to_string()) - } - "log_bodies" => { - let env_default = shared.runtime_defaults.log_bodies; - if let Ok(mut config) = shared.runtime_config.write() { - config.log_bodies = env_default; - } - Some(env_default.to_string()) - } - "anthropic_thinking_repair" => { - let env_default = shared.runtime_defaults.anthropic_thinking_repair; - if let Ok(mut config) = shared.runtime_config.write() { - config.anthropic_thinking_repair = env_default; - } - Some(env_default.to_string()) - } - "pxpipe_compress" => { - let env_default = shared.runtime_defaults.pxpipe_compress; - if let Ok(mut config) = shared.runtime_config.write() { - config.pxpipe_compress = env_default; - } - Some(env_default.to_string()) - } - "rtk_compress" => { - let env_default = shared.runtime_defaults.rtk_compress; - if let Ok(mut config) = shared.runtime_config.write() { - config.rtk_compress = env_default; - } - Some(env_default.to_string()) - } - "rtk_models" => { - let env_default = shared.runtime_defaults.rtk_models.clone(); - if let Ok(mut config) = shared.runtime_config.write() { - config.rtk_models = env_default.clone(); - } - Some(env_default) - } - "pxpipe_models" => { - let env_default = shared.runtime_defaults.pxpipe_models.clone(); - if let Ok(mut config) = shared.runtime_config.write() { - config.pxpipe_models = env_default.clone(); - } - Some(env_default) - } - "forward_client_auth" => { - let env_default = shared.runtime_defaults.forward_client_auth; - if let Ok(mut config) = shared.runtime_config.write() { - config.forward_client_auth = env_default; - } - Some(env_default.to_string()) - } - "tool_guardrail_mode" => { - let env_default = shared.runtime_defaults.tool_guardrail_mode.clone(); - if let Ok(mut config) = shared.runtime_config.write() { - config.tool_guardrail_mode = env_default.clone(); - } - Some(env_default) - } - "optimizer_mode" => { - let env_default = shared.runtime_defaults.optimizer_mode.clone(); - if let Ok(mut config) = shared.runtime_config.write() { - config.optimizer_mode = env_default.clone(); - } - Some(env_default) - } - _ => None, - }; - if let Some(value) = restored { - let _ = shared - .events_tx - .send(crate::admin::state::AdminEvent::ConfigChanged { - key: key.clone(), - value, - }); - } - super::emit_audit( - &shared, - crate::admin::db::AuditEntry { - id: None, - timestamp: None, - action: "config_deleted".into(), - target_type: "config".into(), - target_id: Some(key.clone()), - detail: None, - source_ip: Some(addr.ip().to_string()), - }, - ); - (StatusCode::OK, Json(serde_json::json!({"deleted": key}))).into_response() - } - Some(Ok(false)) => ( - StatusCode::NOT_FOUND, - Json(serde_json::json!({"error": "override not found"})), - ) - .into_response(), - Some(Err(e)) => { - tracing::error!(error = %e, "delete_config_override failed"); - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": "internal database error"})), - ) - .into_response() - } - None => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": "internal error"})), - ) - .into_response(), - } -} diff --git a/crates/proxy/src/admin/routes/routes_api.rs b/crates/proxy/src/admin/routes/routes_api.rs deleted file mode 100644 index afa93e9..0000000 --- a/crates/proxy/src/admin/routes/routes_api.rs +++ /dev/null @@ -1,562 +0,0 @@ -// Route handlers for the routes + route_providers CRUD API. -pub use crate::admin::db::{ReorderOutcome, RoutePatch, RouteProviderRow, RouteRow}; - -use crate::admin::state::SharedState; -use axum::{ - extract::{ConnectInfo, Path, State}, - http::StatusCode, - response::IntoResponse, - Json, -}; -use serde::Deserialize; -use std::net::SocketAddr; - -// ── Response types ──────────────────────────────────────────────────────────── - -#[derive(serde::Serialize)] -struct RouteResponse { - id: String, - name: String, - description: Option, - strategy: String, - rpm: Option, - tpm: Option, - budget_usd: Option, - enabled: bool, - guardrail_mode: Option, - pxpipe_compress: Option, - pxpipe_models: Option, - redact_secrets: Option, - position: i32, - provider_count: usize, - created_at: String, - updated_at: String, -} - -#[derive(serde::Serialize)] -struct RouteProviderResponse { - id: String, - route_id: String, - backend_id: String, - backend_name: String, - provider_id: String, - models: Vec, - priority: i32, - enabled: bool, -} - -// ── Request types ───────────────────────────────────────────────────────────── - -#[derive(Deserialize)] -pub struct CreateRouteRequest { - pub name: String, - #[serde(default)] - pub description: Option, - #[serde(default = "default_strategy")] - pub strategy: String, - pub rpm: Option, - pub tpm: Option, - pub budget_usd: Option, - #[serde(default = "default_true")] - pub enabled: bool, - #[serde(default)] - pub guardrail_mode: Option, - #[serde(default)] - pub pxpipe_compress: Option, - #[serde(default)] - pub pxpipe_models: Option, - #[serde(default)] - pub redact_secrets: Option, - #[serde(default)] - pub position: i32, -} - -fn default_strategy() -> String { - "failover".into() -} - -#[derive(Deserialize)] -pub struct AddRouteProviderRequest { - pub backend_id: String, - #[serde(default = "default_models")] - pub models: Vec, - #[serde(default)] - pub priority: i32, - #[serde(default = "default_true")] - pub enabled: bool, -} - -fn default_models() -> Vec { - vec!["*".into()] -} - -fn default_true() -> bool { - true -} - -#[derive(Deserialize)] -pub struct UpdateRouteProviderRequest { - pub models: Option>, - pub priority: Option, - pub enabled: Option, -} - -#[derive(Deserialize)] -pub struct ReorderRouteProvidersRequest { - pub provider_ids: Vec, -} - -// ── Helpers ─────────────────────────────────────────────────────────────────── - -fn route_to_response(route: &RouteRow, provider_count: usize) -> RouteResponse { - RouteResponse { - id: route.id.clone(), - name: route.name.clone(), - description: route.description.clone(), - strategy: route.strategy.clone(), - rpm: route.rpm, - tpm: route.tpm, - budget_usd: route.budget_usd, - enabled: route.enabled, - guardrail_mode: route.guardrail_mode.clone(), - pxpipe_compress: route.pxpipe_compress, - pxpipe_models: route.pxpipe_models.clone(), - redact_secrets: route.redact_secrets, - position: route.position, - provider_count, - created_at: route.created_at.clone(), - updated_at: route.updated_at.clone(), - } -} - -fn db_error_status(e: &rusqlite::Error) -> StatusCode { - if let rusqlite::Error::SqliteFailure(ref err, _) = e { - if err.code == rusqlite::ErrorCode::ConstraintViolation { - return StatusCode::CONFLICT; - } - } - StatusCode::INTERNAL_SERVER_ERROR -} - -fn ok_json(val: T) -> axum::response::Response { - (StatusCode::OK, Json(serde_json::to_value(val).unwrap())).into_response() -} - -fn err_json(status: StatusCode, msg: impl Into) -> axum::response::Response { - (status, Json(serde_json::json!({ "error": msg.into() }))).into_response() -} - -fn build_provider_responses( - providers: &[RouteProviderRow], - backends: &[crate::admin::db::ManagedBackendRow], -) -> Vec { - let backend_map: std::collections::HashMap<&str, &crate::admin::db::ManagedBackendRow> = - backends.iter().map(|b| (b.id.as_str(), b)).collect(); - providers - .iter() - .map(|p| { - let backend = backend_map.get(p.backend_id.as_str()); - RouteProviderResponse { - id: p.id.clone(), - route_id: p.route_id.clone(), - backend_id: p.backend_id.clone(), - backend_name: backend.map(|b| b.name.clone()).unwrap_or_default(), - provider_id: backend.map(|b| b.provider_id.clone()).unwrap_or_default(), - models: p.models.clone(), - priority: p.priority, - enabled: p.enabled, - } - }) - .collect() -} - -fn audit( - shared: &SharedState, - source_ip: Option, - action: &str, - target_type: &str, - target_id: String, - detail: Option, -) { - super::emit_audit( - shared, - crate::admin::db::AuditEntry { - id: None, - timestamp: None, - action: action.into(), - target_type: target_type.into(), - target_id: Some(target_id), - detail, - source_ip, - }, - ); -} - -// ── Route CRUD ──────────────────────────────────────────────────────────────── - -pub(super) async fn list_routes(State(shared): State) -> axum::response::Response { - let result = crate::admin::state::with_db(&shared.db, |conn| { - let routes = crate::admin::db::list_routes(conn)?; - let mut resp = Vec::with_capacity(routes.len()); - for r in &routes { - let count = crate::admin::db::count_route_providers(conn, &r.id).unwrap_or(0); - resp.push(route_to_response(r, count)); - } - Ok::<_, rusqlite::Error>(resp) - }) - .await; - - match result { - Some(Ok(resp)) => ok_json(serde_json::json!({ "routes": resp })), - _ => err_json(StatusCode::INTERNAL_SERVER_ERROR, "Failed to list routes"), - } -} - -pub(super) async fn create_route( - ConnectInfo(addr): ConnectInfo, - State(shared): State, - Json(body): Json, -) -> axum::response::Response { - if body.name.trim().is_empty() { - return err_json(StatusCode::BAD_REQUEST, "name is required"); - } - - let now = crate::admin::db::now_iso8601(); - let row = RouteRow { - id: uuid::Uuid::new_v4().to_string(), - name: body.name.trim().to_string(), - description: body.description, - strategy: body.strategy, - rpm: body.rpm, - tpm: body.tpm, - budget_usd: body.budget_usd, - enabled: body.enabled, - guardrail_mode: body.guardrail_mode, - pxpipe_compress: body.pxpipe_compress, - pxpipe_models: body.pxpipe_models, - redact_secrets: body.redact_secrets, - position: body.position, - created_at: now.clone(), - updated_at: now, - }; - - let row_clone = row.clone(); - let result = crate::admin::state::with_db(&shared.db, move |conn| { - crate::admin::db::insert_route(conn, &row_clone) - }) - .await; - - match result { - Some(Ok(())) => { - audit( - &shared, - Some(addr.ip().to_string()), - "route_created", - "route", - row.name.clone(), - None, - ); - super::rebuild_route_router(&shared).await; - ( - StatusCode::CREATED, - Json(serde_json::to_value(route_to_response(&row, 0)).unwrap()), - ) - .into_response() - } - Some(Err(e)) => err_json(db_error_status(&e), e.to_string()), - None => err_json(StatusCode::INTERNAL_SERVER_ERROR, "Failed to create route"), - } -} - -pub(super) async fn update_route( - ConnectInfo(addr): ConnectInfo, - State(shared): State, - Path(id): Path, - Json(body): Json, -) -> axum::response::Response { - let id_clone = id.clone(); - let result = crate::admin::state::with_db(&shared.db, move |conn| { - let updated = crate::admin::db::update_route(conn, &id_clone, &body)?; - if !updated { - return Ok::<_, rusqlite::Error>((false, None, 0)); - } - let route = crate::admin::db::get_route(conn, &id_clone).ok().flatten(); - let count = crate::admin::db::count_route_providers(conn, &id_clone).unwrap_or(0); - Ok((true, route, count)) - }) - .await; - - match result { - Some(Ok((true, Some(r), count))) => { - audit( - &shared, - Some(addr.ip().to_string()), - "route_updated", - "route", - id, - None, - ); - super::rebuild_route_router(&shared).await; - ok_json(route_to_response(&r, count)) - } - Some(Ok((true, None, _))) => err_json( - StatusCode::INTERNAL_SERVER_ERROR, - "route not found after update", - ), - Some(Ok((false, _, _))) => err_json(StatusCode::NOT_FOUND, "route not found"), - Some(Err(e)) => err_json(StatusCode::INTERNAL_SERVER_ERROR, e.to_string()), - None => err_json(StatusCode::INTERNAL_SERVER_ERROR, "Failed to update route"), - } -} - -pub(super) async fn delete_route( - ConnectInfo(addr): ConnectInfo, - State(shared): State, - Path(id): Path, -) -> axum::response::Response { - let id_clone = id.clone(); - let result = crate::admin::state::with_db(&shared.db, move |conn| { - crate::admin::db::delete_route(conn, &id_clone) - }) - .await; - - match result { - Some(Ok(true)) => { - audit( - &shared, - Some(addr.ip().to_string()), - "route_deleted", - "route", - id, - None, - ); - super::rebuild_route_router(&shared).await; - ok_json(serde_json::json!({ "ok": true })) - } - Some(Ok(false)) => err_json(StatusCode::NOT_FOUND, "route not found"), - Some(Err(e)) => err_json(StatusCode::INTERNAL_SERVER_ERROR, e.to_string()), - None => err_json(StatusCode::INTERNAL_SERVER_ERROR, "Failed to delete route"), - } -} - -// ── Route Providers CRUD ────────────────────────────────────────────────────── - -pub(super) async fn list_route_providers_handler( - State(shared): State, - Path(route_id): Path, -) -> axum::response::Response { - let result = crate::admin::state::with_db(&shared.db, move |conn| { - if crate::admin::db::get_route(conn, &route_id) - .ok() - .flatten() - .is_none() - { - return Ok::<_, rusqlite::Error>(None); - } - let providers = crate::admin::db::list_route_providers(conn, &route_id)?; - let backends = crate::admin::db::list_managed_backends(conn).unwrap_or_default(); - Ok(Some(build_provider_responses(&providers, &backends))) - }) - .await; - - match result { - Some(Ok(Some(resp))) => ok_json(serde_json::json!({ "providers": resp })), - Some(Ok(None)) => err_json(StatusCode::NOT_FOUND, "route not found"), - _ => err_json( - StatusCode::INTERNAL_SERVER_ERROR, - "Failed to list route providers", - ), - } -} - -pub(super) async fn add_route_provider_handler( - ConnectInfo(addr): ConnectInfo, - State(shared): State, - Path(route_id): Path, - Json(body): Json, -) -> axum::response::Response { - let route_id_clone = route_id.clone(); - let backend_id = body.backend_id.clone(); - let models = body.models.clone(); - let priority = body.priority; - let enabled = body.enabled; - - let backend_id_for_db = backend_id.clone(); - let backend_id_check = backend_id.clone(); - let result = crate::admin::state::with_db(&shared.db, move |conn| { - if crate::admin::db::get_route(conn, &route_id_clone) - .ok() - .flatten() - .is_none() - { - return Ok::<_, rusqlite::Error>(Err::<(), String>("route not found".into())); - } - if !crate::admin::db::managed_backend_exists(conn, &backend_id_check)? { - return Ok(Err::<(), String>("backend not found".into())); - } - crate::admin::db::add_route_provider( - conn, - &route_id_clone, - &backend_id_for_db, - &models, - priority, - enabled, - )?; - Ok(Ok::<(), String>(())) - }) - .await; - - match result { - Some(Ok(Ok(()))) => { - audit( - &shared, - Some(addr.ip().to_string()), - "route_provider_added", - "route_provider", - route_id, - Some(format!("backend_id={}", backend_id)), - ); - super::rebuild_route_router(&shared).await; - (StatusCode::CREATED, Json(serde_json::json!({ "ok": true }))).into_response() - } - Some(Ok(Err(msg))) => err_json( - if msg == "route not found" { - StatusCode::NOT_FOUND - } else { - StatusCode::BAD_REQUEST - }, - msg, - ), - Some(Err(e)) => err_json(db_error_status(&e), e.to_string()), - None => err_json( - StatusCode::INTERNAL_SERVER_ERROR, - "Failed to add route provider", - ), - } -} - -pub(super) async fn update_route_provider_handler( - ConnectInfo(addr): ConnectInfo, - State(shared): State, - Path((route_id, provider_id)): Path<(String, String)>, - Json(body): Json, -) -> axum::response::Response { - let provider_id_clone = provider_id.clone(); - let result = crate::admin::state::with_db(&shared.db, move |conn| { - crate::admin::db::update_route_provider( - conn, - &provider_id_clone, - body.models.as_deref(), - body.priority, - body.enabled, - ) - }) - .await; - - match result { - Some(Ok(true)) => { - audit( - &shared, - Some(addr.ip().to_string()), - "route_provider_updated", - "route_provider", - provider_id, - Some(format!("route_id={}", route_id)), - ); - super::rebuild_route_router(&shared).await; - ok_json(serde_json::json!({ "ok": true })) - } - Some(Ok(false)) => err_json(StatusCode::NOT_FOUND, "route provider not found"), - Some(Err(e)) => err_json(StatusCode::INTERNAL_SERVER_ERROR, e.to_string()), - None => err_json( - StatusCode::INTERNAL_SERVER_ERROR, - "Failed to update route provider", - ), - } -} - -pub(super) async fn reorder_route_providers_handler( - ConnectInfo(addr): ConnectInfo, - State(shared): State, - Path(route_id): Path, - Json(body): Json, -) -> axum::response::Response { - let route_id_clone = route_id.clone(); - let ordered = body.provider_ids.clone(); - let result = crate::admin::state::with_db(&shared.db, move |conn| { - if crate::admin::db::get_route(conn, &route_id_clone) - .ok() - .flatten() - .is_none() - { - return Ok::<_, rusqlite::Error>(None); - } - let outcome = crate::admin::db::reorder_route_providers(conn, &route_id_clone, &ordered)?; - match outcome { - ReorderOutcome::Mismatch => Ok(Some(Err::<_, ()>(()))), - ReorderOutcome::Ok(rows) => { - let backends = crate::admin::db::list_managed_backends(conn).unwrap_or_default(); - Ok(Some(Ok(build_provider_responses(&rows, &backends)))) - } - } - }) - .await; - - match result { - Some(Ok(Some(Ok(providers)))) => { - audit( - &shared, - Some(addr.ip().to_string()), - "route_providers_reordered", - "route", - route_id, - Some(format!("count={}", providers.len())), - ); - super::rebuild_route_router(&shared).await; - ok_json(serde_json::json!({ "providers": providers })) - } - Some(Ok(Some(Err(())))) => err_json( - StatusCode::BAD_REQUEST, - "provider_ids must match the route's current provider set exactly", - ), - Some(Ok(None)) => err_json(StatusCode::NOT_FOUND, "route not found"), - Some(Err(e)) => err_json(StatusCode::INTERNAL_SERVER_ERROR, e.to_string()), - None => err_json( - StatusCode::INTERNAL_SERVER_ERROR, - "Failed to reorder route providers", - ), - } -} - -pub(super) async fn remove_route_provider_handler( - ConnectInfo(addr): ConnectInfo, - State(shared): State, - Path((_route_id, provider_id)): Path<(String, String)>, -) -> axum::response::Response { - let provider_id_clone = provider_id.clone(); - let result = crate::admin::state::with_db(&shared.db, move |conn| { - crate::admin::db::remove_route_provider(conn, &provider_id_clone) - }) - .await; - - match result { - Some(Ok(true)) => { - audit( - &shared, - Some(addr.ip().to_string()), - "route_provider_removed", - "route_provider", - provider_id, - None, - ); - super::rebuild_route_router(&shared).await; - ok_json(serde_json::json!({ "ok": true })) - } - Some(Ok(false)) => err_json(StatusCode::NOT_FOUND, "route provider not found"), - Some(Err(e)) => err_json(StatusCode::INTERNAL_SERVER_ERROR, e.to_string()), - None => err_json( - StatusCode::INTERNAL_SERVER_ERROR, - "Failed to remove route provider", - ), - } -} diff --git a/crates/proxy/src/admin/routes/routes_api/helpers.rs b/crates/proxy/src/admin/routes/routes_api/helpers.rs new file mode 100644 index 0000000..3567d87 --- /dev/null +++ b/crates/proxy/src/admin/routes/routes_api/helpers.rs @@ -0,0 +1,179 @@ +pub use crate::admin::db::{ReorderOutcome, RoutePatch, RouteProviderRow, RouteRow}; +use crate::admin::state::SharedState; +use axum::{http::StatusCode, response::IntoResponse, Json}; +use serde::Deserialize; + +#[derive(serde::Serialize)] +pub(crate) struct RouteResponse { + pub id: String, + pub name: String, + pub description: Option, + pub strategy: String, + pub rpm: Option, + pub tpm: Option, + pub budget_usd: Option, + pub enabled: bool, + pub guardrail_mode: Option, + pub pxpipe_compress: Option, + pub pxpipe_models: Option, + pub redact_secrets: Option, + pub position: i32, + pub provider_count: usize, + pub created_at: String, + pub updated_at: String, +} + +#[derive(serde::Serialize)] +pub(crate) struct RouteProviderResponse { + pub id: String, + pub route_id: String, + pub backend_id: String, + pub backend_name: String, + pub provider_id: String, + pub models: Vec, + pub priority: i32, + pub enabled: bool, +} + +#[derive(Deserialize)] +pub(crate) struct CreateRouteRequest { + pub name: String, + #[serde(default)] + pub description: Option, + #[serde(default = "default_strategy")] + pub strategy: String, + pub rpm: Option, + pub tpm: Option, + pub budget_usd: Option, + #[serde(default = "default_true")] + pub enabled: bool, + #[serde(default)] + pub guardrail_mode: Option, + #[serde(default)] + pub pxpipe_compress: Option, + #[serde(default)] + pub pxpipe_models: Option, + #[serde(default)] + pub redact_secrets: Option, + #[serde(default)] + pub position: i32, +} + +fn default_strategy() -> String { + "failover".into() +} + +#[derive(Deserialize)] +pub(crate) struct AddRouteProviderRequest { + pub backend_id: String, + #[serde(default = "default_models")] + pub models: Vec, + #[serde(default)] + pub priority: i32, + #[serde(default = "default_true")] + pub enabled: bool, +} + +fn default_models() -> Vec { + vec!["*".into()] +} + +fn default_true() -> bool { + true +} + +#[derive(Deserialize)] +pub(crate) struct UpdateRouteProviderRequest { + pub models: Option>, + pub priority: Option, + pub enabled: Option, +} + +#[derive(Deserialize)] +pub(crate) struct ReorderRouteProvidersRequest { + pub provider_ids: Vec, +} + +pub(crate) fn route_to_response(route: &RouteRow, provider_count: usize) -> RouteResponse { + RouteResponse { + id: route.id.clone(), + name: route.name.clone(), + description: route.description.clone(), + strategy: route.strategy.clone(), + rpm: route.rpm, + tpm: route.tpm, + budget_usd: route.budget_usd, + enabled: route.enabled, + guardrail_mode: route.guardrail_mode.clone(), + pxpipe_compress: route.pxpipe_compress, + pxpipe_models: route.pxpipe_models.clone(), + redact_secrets: route.redact_secrets, + position: route.position, + provider_count, + created_at: route.created_at.clone(), + updated_at: route.updated_at.clone(), + } +} + +pub(crate) fn db_error_status(e: &rusqlite::Error) -> StatusCode { + if let rusqlite::Error::SqliteFailure(ref err, _) = e { + if err.code == rusqlite::ErrorCode::ConstraintViolation { + return StatusCode::CONFLICT; + } + } + StatusCode::INTERNAL_SERVER_ERROR +} + +pub(crate) fn ok_json(val: T) -> axum::response::Response { + (StatusCode::OK, Json(serde_json::to_value(val).unwrap())).into_response() +} + +pub(crate) fn err_json(status: StatusCode, msg: impl Into) -> axum::response::Response { + (status, Json(serde_json::json!({ "error": msg.into() }))).into_response() +} + +pub(crate) fn build_provider_responses( + providers: &[RouteProviderRow], + backends: &[crate::admin::db::ManagedBackendRow], +) -> Vec { + let backend_map: std::collections::HashMap<&str, &crate::admin::db::ManagedBackendRow> = + backends.iter().map(|b| (b.id.as_str(), b)).collect(); + providers + .iter() + .map(|p| { + let backend = backend_map.get(p.backend_id.as_str()); + RouteProviderResponse { + id: p.id.clone(), + route_id: p.route_id.clone(), + backend_id: p.backend_id.clone(), + backend_name: backend.map(|b| b.name.clone()).unwrap_or_default(), + provider_id: backend.map(|b| b.provider_id.clone()).unwrap_or_default(), + models: p.models.clone(), + priority: p.priority, + enabled: p.enabled, + } + }) + .collect() +} + +pub(crate) fn audit( + shared: &SharedState, + source_ip: Option, + action: &str, + target_type: &str, + target_id: String, + detail: Option, +) { + super::super::emit_audit( + shared, + crate::admin::db::AuditEntry { + id: None, + timestamp: None, + action: action.into(), + target_type: target_type.into(), + target_id: Some(target_id), + detail, + source_ip, + }, + ); +} diff --git a/crates/proxy/src/admin/routes/routes_api/mod.rs b/crates/proxy/src/admin/routes/routes_api/mod.rs new file mode 100644 index 0000000..8db9d08 --- /dev/null +++ b/crates/proxy/src/admin/routes/routes_api/mod.rs @@ -0,0 +1,9 @@ +pub(crate) mod helpers; +pub(crate) mod providers; +pub(crate) mod routes; + +pub(crate) use providers::{ + add_route_provider_handler, list_route_providers_handler, remove_route_provider_handler, + reorder_route_providers_handler, update_route_provider_handler, +}; +pub(crate) use routes::{create_route, delete_route, list_routes, update_route}; diff --git a/crates/proxy/src/admin/routes/routes_api/providers.rs b/crates/proxy/src/admin/routes/routes_api/providers.rs new file mode 100644 index 0000000..a0345e2 --- /dev/null +++ b/crates/proxy/src/admin/routes/routes_api/providers.rs @@ -0,0 +1,233 @@ +use crate::admin::state::SharedState; +use axum::{ + extract::{ConnectInfo, Path, State}, + http::StatusCode, + response::IntoResponse, + Json, +}; +use std::net::SocketAddr; + +use super::helpers::{ + audit, build_provider_responses, db_error_status, err_json, ok_json, AddRouteProviderRequest, + ReorderOutcome, ReorderRouteProvidersRequest, UpdateRouteProviderRequest, +}; + +pub(crate) async fn list_route_providers_handler( + State(shared): State, + Path(route_id): Path, +) -> axum::response::Response { + let result = crate::admin::state::with_db(&shared.db, move |conn| { + if crate::admin::db::get_route(conn, &route_id) + .ok() + .flatten() + .is_none() + { + return Ok::<_, rusqlite::Error>(None); + } + let providers = crate::admin::db::list_route_providers(conn, &route_id)?; + let backends = crate::admin::db::list_managed_backends(conn).unwrap_or_default(); + Ok(Some(build_provider_responses(&providers, &backends))) + }) + .await; + + match result { + Some(Ok(Some(resp))) => ok_json(serde_json::json!({ "providers": resp })), + Some(Ok(None)) => err_json(StatusCode::NOT_FOUND, "route not found"), + _ => err_json( + StatusCode::INTERNAL_SERVER_ERROR, + "Failed to list route providers", + ), + } +} + +pub(crate) async fn add_route_provider_handler( + ConnectInfo(addr): ConnectInfo, + State(shared): State, + Path(route_id): Path, + Json(body): Json, +) -> axum::response::Response { + let route_id_clone = route_id.clone(); + let backend_id = body.backend_id.clone(); + let models = body.models.clone(); + let priority = body.priority; + let enabled = body.enabled; + + let backend_id_for_db = backend_id.clone(); + let backend_id_check = backend_id.clone(); + let result = crate::admin::state::with_db(&shared.db, move |conn| { + if crate::admin::db::get_route(conn, &route_id_clone) + .ok() + .flatten() + .is_none() + { + return Ok::<_, rusqlite::Error>(Err::<(), String>("route not found".into())); + } + if !crate::admin::db::managed_backend_exists(conn, &backend_id_check)? { + return Ok(Err::<(), String>("backend not found".into())); + } + crate::admin::db::add_route_provider( + conn, + &route_id_clone, + &backend_id_for_db, + &models, + priority, + enabled, + )?; + Ok(Ok::<(), String>(())) + }) + .await; + + match result { + Some(Ok(Ok(()))) => { + audit( + &shared, + Some(addr.ip().to_string()), + "route_provider_added", + "route_provider", + route_id, + Some(format!("backend_id={}", backend_id)), + ); + super::super::rebuild_route_router(&shared).await; + (StatusCode::CREATED, Json(serde_json::json!({ "ok": true }))).into_response() + } + Some(Ok(Err(msg))) => err_json( + if msg == "route not found" { + StatusCode::NOT_FOUND + } else { + StatusCode::BAD_REQUEST + }, + msg, + ), + Some(Err(e)) => err_json(db_error_status(&e), e.to_string()), + None => err_json( + StatusCode::INTERNAL_SERVER_ERROR, + "Failed to add route provider", + ), + } +} + +pub(crate) async fn update_route_provider_handler( + ConnectInfo(addr): ConnectInfo, + State(shared): State, + Path((route_id, provider_id)): Path<(String, String)>, + Json(body): Json, +) -> axum::response::Response { + let provider_id_clone = provider_id.clone(); + let result = crate::admin::state::with_db(&shared.db, move |conn| { + crate::admin::db::update_route_provider( + conn, + &provider_id_clone, + body.models.as_deref(), + body.priority, + body.enabled, + ) + }) + .await; + + match result { + Some(Ok(true)) => { + audit( + &shared, + Some(addr.ip().to_string()), + "route_provider_updated", + "route_provider", + provider_id, + Some(format!("route_id={}", route_id)), + ); + super::super::rebuild_route_router(&shared).await; + ok_json(serde_json::json!({ "ok": true })) + } + Some(Ok(false)) => err_json(StatusCode::NOT_FOUND, "route provider not found"), + Some(Err(e)) => err_json(StatusCode::INTERNAL_SERVER_ERROR, e.to_string()), + None => err_json( + StatusCode::INTERNAL_SERVER_ERROR, + "Failed to update route provider", + ), + } +} + +pub(crate) async fn reorder_route_providers_handler( + ConnectInfo(addr): ConnectInfo, + State(shared): State, + Path(route_id): Path, + Json(body): Json, +) -> axum::response::Response { + let route_id_clone = route_id.clone(); + let ordered = body.provider_ids.clone(); + let result = crate::admin::state::with_db(&shared.db, move |conn| { + if crate::admin::db::get_route(conn, &route_id_clone) + .ok() + .flatten() + .is_none() + { + return Ok::<_, rusqlite::Error>(None); + } + let outcome = crate::admin::db::reorder_route_providers(conn, &route_id_clone, &ordered)?; + match outcome { + ReorderOutcome::Mismatch => Ok(Some(Err::<_, ()>(()))), + ReorderOutcome::Ok(rows) => { + let backends = crate::admin::db::list_managed_backends(conn).unwrap_or_default(); + Ok(Some(Ok(build_provider_responses(&rows, &backends)))) + } + } + }) + .await; + + match result { + Some(Ok(Some(Ok(providers)))) => { + audit( + &shared, + Some(addr.ip().to_string()), + "route_providers_reordered", + "route", + route_id, + Some(format!("count={}", providers.len())), + ); + super::super::rebuild_route_router(&shared).await; + ok_json(serde_json::json!({ "providers": providers })) + } + Some(Ok(Some(Err(())))) => err_json( + StatusCode::BAD_REQUEST, + "provider_ids must match the route's current provider set exactly", + ), + Some(Ok(None)) => err_json(StatusCode::NOT_FOUND, "route not found"), + Some(Err(e)) => err_json(StatusCode::INTERNAL_SERVER_ERROR, e.to_string()), + None => err_json( + StatusCode::INTERNAL_SERVER_ERROR, + "Failed to reorder route providers", + ), + } +} + +pub(crate) async fn remove_route_provider_handler( + ConnectInfo(addr): ConnectInfo, + State(shared): State, + Path((_route_id, provider_id)): Path<(String, String)>, +) -> axum::response::Response { + let provider_id_clone = provider_id.clone(); + let result = crate::admin::state::with_db(&shared.db, move |conn| { + crate::admin::db::remove_route_provider(conn, &provider_id_clone) + }) + .await; + + match result { + Some(Ok(true)) => { + audit( + &shared, + Some(addr.ip().to_string()), + "route_provider_removed", + "route_provider", + provider_id, + None, + ); + super::super::rebuild_route_router(&shared).await; + ok_json(serde_json::json!({ "ok": true })) + } + Some(Ok(false)) => err_json(StatusCode::NOT_FOUND, "route provider not found"), + Some(Err(e)) => err_json(StatusCode::INTERNAL_SERVER_ERROR, e.to_string()), + None => err_json( + StatusCode::INTERNAL_SERVER_ERROR, + "Failed to remove route provider", + ), + } +} diff --git a/crates/proxy/src/admin/routes/routes_api/routes.rs b/crates/proxy/src/admin/routes/routes_api/routes.rs new file mode 100644 index 0000000..79a5211 --- /dev/null +++ b/crates/proxy/src/admin/routes/routes_api/routes.rs @@ -0,0 +1,158 @@ +use crate::admin::state::SharedState; +use axum::{ + extract::{ConnectInfo, Path, State}, + http::StatusCode, + response::IntoResponse, + Json, +}; +use std::net::SocketAddr; + +use super::helpers::{ + audit, db_error_status, err_json, ok_json, route_to_response, CreateRouteRequest, RoutePatch, + RouteRow, +}; + +pub(crate) async fn list_routes(State(shared): State) -> axum::response::Response { + let result = crate::admin::state::with_db(&shared.db, |conn| { + let routes = crate::admin::db::list_routes(conn)?; + let mut resp = Vec::with_capacity(routes.len()); + for r in &routes { + let count = crate::admin::db::count_route_providers(conn, &r.id).unwrap_or(0); + resp.push(route_to_response(r, count)); + } + Ok::<_, rusqlite::Error>(resp) + }) + .await; + + match result { + Some(Ok(resp)) => ok_json(serde_json::json!({ "routes": resp })), + _ => err_json(StatusCode::INTERNAL_SERVER_ERROR, "Failed to list routes"), + } +} + +pub(crate) async fn create_route( + ConnectInfo(addr): ConnectInfo, + State(shared): State, + Json(body): Json, +) -> axum::response::Response { + if body.name.trim().is_empty() { + return err_json(StatusCode::BAD_REQUEST, "name is required"); + } + + let now = crate::admin::db::now_iso8601(); + let row = RouteRow { + id: uuid::Uuid::new_v4().to_string(), + name: body.name.trim().to_string(), + description: body.description, + strategy: body.strategy, + rpm: body.rpm, + tpm: body.tpm, + budget_usd: body.budget_usd, + enabled: body.enabled, + guardrail_mode: body.guardrail_mode, + pxpipe_compress: body.pxpipe_compress, + pxpipe_models: body.pxpipe_models, + redact_secrets: body.redact_secrets, + position: body.position, + created_at: now.clone(), + updated_at: now, + }; + + let row_clone = row.clone(); + let result = crate::admin::state::with_db(&shared.db, move |conn| { + crate::admin::db::insert_route(conn, &row_clone) + }) + .await; + + match result { + Some(Ok(())) => { + audit( + &shared, + Some(addr.ip().to_string()), + "route_created", + "route", + row.name.clone(), + None, + ); + super::super::rebuild_route_router(&shared).await; + ( + StatusCode::CREATED, + Json(serde_json::to_value(route_to_response(&row, 0)).unwrap()), + ) + .into_response() + } + Some(Err(e)) => err_json(db_error_status(&e), e.to_string()), + None => err_json(StatusCode::INTERNAL_SERVER_ERROR, "Failed to create route"), + } +} + +pub(crate) async fn update_route( + ConnectInfo(addr): ConnectInfo, + State(shared): State, + Path(id): Path, + Json(body): Json, +) -> axum::response::Response { + let id_clone = id.clone(); + let result = crate::admin::state::with_db(&shared.db, move |conn| { + let updated = crate::admin::db::update_route(conn, &id_clone, &body)?; + if !updated { + return Ok::<_, rusqlite::Error>((false, None, 0)); + } + let route = crate::admin::db::get_route(conn, &id_clone).ok().flatten(); + let count = crate::admin::db::count_route_providers(conn, &id_clone).unwrap_or(0); + Ok((true, route, count)) + }) + .await; + + match result { + Some(Ok((true, Some(r), count))) => { + audit( + &shared, + Some(addr.ip().to_string()), + "route_updated", + "route", + id, + None, + ); + super::super::rebuild_route_router(&shared).await; + ok_json(route_to_response(&r, count)) + } + Some(Ok((true, None, _))) => err_json( + StatusCode::INTERNAL_SERVER_ERROR, + "route not found after update", + ), + Some(Ok((false, _, _))) => err_json(StatusCode::NOT_FOUND, "route not found"), + Some(Err(e)) => err_json(StatusCode::INTERNAL_SERVER_ERROR, e.to_string()), + None => err_json(StatusCode::INTERNAL_SERVER_ERROR, "Failed to update route"), + } +} + +pub(crate) async fn delete_route( + ConnectInfo(addr): ConnectInfo, + State(shared): State, + Path(id): Path, +) -> axum::response::Response { + let id_clone = id.clone(); + let result = crate::admin::state::with_db(&shared.db, move |conn| { + crate::admin::db::delete_route(conn, &id_clone) + }) + .await; + + match result { + Some(Ok(true)) => { + audit( + &shared, + Some(addr.ip().to_string()), + "route_deleted", + "route", + id, + None, + ); + super::super::rebuild_route_router(&shared).await; + ok_json(serde_json::json!({ "ok": true })) + } + Some(Ok(false)) => err_json(StatusCode::NOT_FOUND, "route not found"), + Some(Err(e)) => err_json(StatusCode::INTERNAL_SERVER_ERROR, e.to_string()), + None => err_json(StatusCode::INTERNAL_SERVER_ERROR, "Failed to delete route"), + } +} diff --git a/crates/proxy/src/config/litellm/provider.rs b/crates/proxy/src/config/litellm/provider.rs new file mode 100644 index 0000000..53e0d1d --- /dev/null +++ b/crates/proxy/src/config/litellm/provider.rs @@ -0,0 +1,166 @@ +use super::types::LiteLLMParams; +use crate::config::single::validate_gcp_identifier; +use crate::config::{resolve_env_value, strip_v1_suffix, BackendKind}; + +/// Parse LiteLLM's "provider/model_name" format. +/// No prefix defaults to OpenAI (matches LiteLLM behavior). +/// Returns (kind, model_name, stub_provider) where stub_provider is set for +/// registry-resolved OpenAI-compatible providers so callers can use their default URL. +pub(super) fn parse_provider_model( + model: &str, +) -> ( + BackendKind, + String, + Option<&'static anyllm_providers::ProviderDef>, +) { + let (provider, model_name) = model.split_once('/').unwrap_or(("openai", model)); + let mut stub_provider: Option<&'static anyllm_providers::ProviderDef> = None; + let kind = match provider.to_ascii_lowercase().as_str() { + "openai" => BackendKind::OpenAI, + "azure" => BackendKind::AzureOpenAI, + "vertex_ai" | "vertex" => BackendKind::Vertex, + "gemini" => BackendKind::Gemini, + "anthropic" => { + stub_provider = anyllm_providers::get_provider("anthropic"); + BackendKind::Anthropic + } + "bedrock" => BackendKind::Bedrock, + other => { + let prefix_with_slash = format!("{other}/"); + if let Some(p) = anyllm_providers::find_by_litellm_prefix(&prefix_with_slash) { + let resolved = match anyllm_providers::resolve_backend(p.id) { + Some(("openai", _)) => { + stub_provider = Some(p); + BackendKind::OpenAI + } + Some(("anthropic", _)) => BackendKind::Anthropic, + Some(("gemini", _)) => BackendKind::Gemini, + Some(("vertex", _)) => BackendKind::Vertex, + Some(("azure", _)) => BackendKind::AzureOpenAI, + Some(("bedrock", _)) => BackendKind::Bedrock, + _ => { + tracing::warn!(provider = %other, "provider found in registry but protocol not mappable, treating as openai-compatible"); + stub_provider = Some(p); + BackendKind::OpenAI + } + }; + resolved + } else { + tracing::warn!( + provider = %other, + "unknown LiteLLM provider, treating as openai-compatible" + ); + BackendKind::OpenAI + } + } + }; + (kind, model_name.to_string(), stub_provider) +} + +pub(super) fn provider_id_for_litellm_model( + model: &str, + kind: &BackendKind, + stub_provider: Option<&'static anyllm_providers::ProviderDef>, +) -> String { + if let Some(provider) = stub_provider { + return provider.id.to_string(); + } + + let raw_provider = model + .split_once('/') + .map(|(provider, _)| provider) + .unwrap_or("openai") + .to_ascii_lowercase(); + + match kind { + BackendKind::OpenAI => raw_provider, + BackendKind::AzureOpenAI => "azure".to_string(), + BackendKind::Vertex => "vertex_ai".to_string(), + BackendKind::Gemini => "gemini".to_string(), + BackendKind::Anthropic => "anthropic".to_string(), + BackendKind::Bedrock => "bedrock".to_string(), + } +} + +/// Determine the base URL for a deployment, applying provider-specific defaults. +pub(super) fn resolve_base_url( + kind: &BackendKind, + params: &LiteLLMParams, + stub_provider: Option<&'static anyllm_providers::ProviderDef>, + actual_model: &str, +) -> String { + if let Some(ref url) = params.api_base { + let resolved = + resolve_env_value(url).unwrap_or_else(|e| panic!("model_list api_base: {e}")); + if *kind == BackendKind::AzureOpenAI { + let api_version = params.api_version.as_deref().unwrap_or("2024-10-21"); + if !resolved.contains("/openai/deployments/") { + let deployment = azure_deployment_from_model(actual_model); + return format!( + "{}/openai/deployments/{deployment}/chat/completions?api-version={api_version}", + resolved.trim_end_matches('/'), + ); + } + if !resolved.contains("api-version=") { + let sep = if resolved.contains('?') { '&' } else { '?' }; + return format!("{resolved}{sep}api-version={api_version}"); + } + return resolved; + } + return resolved; + } + match kind { + BackendKind::OpenAI => { + let url = if let Some(provider) = stub_provider { + if provider.default_base_url.is_empty() { + panic!( + "model_list provider '{}' requires api_base because it has no safe global API base URL", + provider.id + ); + } + provider.default_base_url + } else { + "https://api.openai.com" + }; + strip_v1_suffix(url).to_string() + } + BackendKind::Gemini => { + "https://generativelanguage.googleapis.com/v1beta/openai".to_string() + } + BackendKind::Anthropic => std::env::var("ANTHROPIC_BASE_URL") + .unwrap_or_else(|_| "https://api.anthropic.com".to_string()), + BackendKind::Bedrock => params + .aws_region_name + .as_deref() + .map(|v| v.to_string()) + .or_else(|| std::env::var("AWS_REGION").ok()) + .unwrap_or_else(|| "us-east-1".to_string()), + BackendKind::AzureOpenAI => { + panic!("api_base is required for azure deployments in model_list") + } + BackendKind::Vertex => { + let project = params.vertex_project.as_deref().unwrap_or_else(|| { + panic!("vertex_project is required for vertex deployments in model_list") + }); + let location = params.vertex_location.as_deref().unwrap_or_else(|| { + panic!("vertex_location is required for vertex deployments in model_list") + }); + validate_gcp_identifier("vertex_project", project); + validate_gcp_identifier("vertex_location", location); + format!( + "https://{location}-aiplatform.googleapis.com/v1/projects/{project}/locations/{location}/endpoints/openapi" + ) + } + } +} + +pub(super) fn azure_deployment_from_model(model: &str) -> &str { + for marker in ["o_series/", "gpt5_series/"] { + if let Some(deployment) = model.strip_prefix(marker) { + if !deployment.is_empty() { + return deployment; + } + } + } + model +} diff --git a/crates/proxy/src/main_helpers/async_main/admin.rs b/crates/proxy/src/main_helpers/async_main/admin.rs deleted file mode 100644 index b346a5b..0000000 --- a/crates/proxy/src/main_helpers/async_main/admin.rs +++ /dev/null @@ -1,696 +0,0 @@ -use anyllm_proxy::admin; -use anyllm_proxy::config; -use anyllm_proxy::server::state as server_state; -use std::path::Path; -use std::sync::Arc; - -#[allow(clippy::too_many_lines)] -pub(crate) async fn init_admin( - args: &[String], - data_dir: &Path, - multi_config: &config::MultiConfig, - model_router: Option>>, - tool_engine_state: &Option>, - reload_handle: tracing_subscriber::reload::Handle< - tracing_subscriber::EnvFilter, - tracing_subscriber::Registry, - >, -) -> Option<( - admin::state::SharedState, - axum::Router, - tokio::net::TcpListener, - u16, - // Plaintext admin token, returned so the caller can print a ready-to-click - // tokenized URL on loopback binds. Dropped right after the startup banner. - String, -)> { - let flag_set = args.iter().any(|a| a == "--webui" || a == "--admin"); - let force_disabled = matches!( - std::env::var("DISABLE_ADMIN").as_deref(), - Ok("1") | Ok("true") | Ok("yes") - ); - let enable_admin = flag_set && !force_disabled; - - if !enable_admin { - return None; - } - - let provider_catalog = Arc::new(anyllm_providers::ProviderCatalog::bundled()); - - let admin_port: u16 = match std::env::var("ADMIN_PORT") { - Ok(val) => val - .parse::() - .unwrap_or_else(|_| panic!("ADMIN_PORT must be a number in 1-65535, got '{val}'")), - Err(_) => 3001, - }; - if admin_port == 0 { - panic!("ADMIN_PORT cannot be 0"); - } - if admin_port < 1024 { - tracing::warn!( - port = admin_port, - "ADMIN_PORT is in the privileged range (< 1024); binding may fail without elevated privileges" - ); - } - - let admin_bind = { - let raw = std::env::var("ADMIN_BIND").unwrap_or_else(|_| "127.0.0.1".into()); - // Translate well-known loopback aliases to explicit IPs before validation. - match raw.to_ascii_lowercase().as_str() { - "localhost" => "127.0.0.1".to_string(), - "localhost6" | "ip6-localhost" | "ip6-loopback" => "::1".to_string(), - _ => raw, - } - }; - // Bind addresses must be explicit IPs. Hostnames other than the loopback - // aliases above are rejected: a bind address maps to a specific local - // interface and must be unambiguous. - if admin_bind.parse::().is_err() { - panic!( - "ADMIN_BIND must be an IP address (e.g. 127.0.0.1 or 0.0.0.0), not a hostname — got '{}'", - std::env::var("ADMIN_BIND").unwrap_or_default() - ); - } - - if admin_port == multi_config.listen_port { - panic!( - "ADMIN_PORT ({admin_port}) must differ from LISTEN_PORT ({})", - multi_config.listen_port - ); - } - - let db_path = crate::main_helpers::bootstrap::resolve_db_path(data_dir); - let conn = - rusqlite::Connection::open(&db_path).expect("failed to open SQLite database for admin"); - admin::db::init_db(&conn).expect("failed to initialize admin database schema"); - let hmac_secret = Arc::new(admin::db::ensure_hmac_secret(&conn)); - - // Build initial RuntimeConfig from the loaded multi_config. - let mut model_mappings = indexmap::IndexMap::new(); - for (name, bc) in &multi_config.backends { - model_mappings.insert(name.clone(), bc.model_mapping.clone()); - } - let log_level = std::env::var("RUST_LOG").unwrap_or_else(|_| "info".into()); - // The runtime default tracks the static per-process guardrail preset - // (built from YAML/env at startup, see `ToolEngineState.guardrails`) so - // an admin "reset to default" restores that mode rather than always - // falling back to disabled. - let tool_guardrail_default = tool_engine_state - .as_ref() - .map(|engine| engine.guardrails.mode.as_str().to_string()) - .unwrap_or_else(|| { - anyllm_proxy::tools::ToolGuardrailMode::Disabled - .as_str() - .to_string() - }); - let mut runtime_config = admin::state::RuntimeConfig { - model_mappings, - log_level, - log_bodies: multi_config.log_bodies, - redact_secrets: multi_config.redact_secrets, - anthropic_thinking_repair: multi_config.anthropic_thinking_repair, - pxpipe_compress: multi_config.pxpipe_compress, - pxpipe_models: anyllm_proxy::pxpipe::resolve_default_models_csv(), - rtk_compress: anyllm_proxy::rtk::resolve_default_enabled(), - rtk_models: anyllm_proxy::rtk::resolve_default_models_csv(), - forward_client_auth: multi_config.forward_client_auth, - tool_guardrail_mode: tool_guardrail_default.clone(), - optimizer_mode: anyllm_proxy::optimizer::resolve_default_mode() - .as_str() - .to_string(), - }; - let runtime_defaults = admin::state::RuntimeConfigDefaults { - log_bodies: multi_config.log_bodies, - redact_secrets: multi_config.redact_secrets, - anthropic_thinking_repair: multi_config.anthropic_thinking_repair, - pxpipe_compress: multi_config.pxpipe_compress, - pxpipe_models: anyllm_proxy::pxpipe::resolve_default_models_csv(), - rtk_compress: anyllm_proxy::rtk::resolve_default_enabled(), - rtk_models: anyllm_proxy::rtk::resolve_default_models_csv(), - forward_client_auth: multi_config.forward_client_auth, - tool_guardrail_mode: tool_guardrail_default, - optimizer_mode: anyllm_proxy::optimizer::resolve_default_mode() - .as_str() - .to_string(), - }; - let mut log_bodies_enabled_by_override = false; - let mut redact_secrets_enabled_by_override = false; - - // Apply config overrides from SQLite (survive restarts). - if let Ok(overrides) = admin::db::get_config_overrides(&conn) { - for (key, value, _) in &overrides { - match key.as_str() { - "log_level" => { - // Apply the same allowlist enforced by the admin API to - // prevent a tampered SQLite database from enabling trace-level - // logging, which would expose API keys in HTTP headers. - const ALLOWED_LOG_LEVELS: &[&str] = &["error", "warn", "info", "debug"]; - let normalized = value.trim().to_lowercase(); - if ALLOWED_LOG_LEVELS.contains(&normalized.as_str()) { - runtime_config.log_level = normalized; - } else { - tracing::warn!( - value = %value, - "ignoring invalid log_level override from database" - ); - } - } - "log_bodies" => { - runtime_config.log_bodies = value == "true"; - log_bodies_enabled_by_override = - runtime_config.log_bodies && !multi_config.log_bodies; - } - "redact_secrets" => { - runtime_config.redact_secrets = value == "true"; - redact_secrets_enabled_by_override = - runtime_config.redact_secrets && !multi_config.redact_secrets; - } - "anthropic_thinking_repair" => { - runtime_config.anthropic_thinking_repair = value == "true"; - } - "pxpipe_compress" => { - runtime_config.pxpipe_compress = value == "true"; - } - "pxpipe_models" => { - runtime_config.pxpipe_models = value.clone(); - } - "rtk_compress" => { - runtime_config.rtk_compress = value == "true"; - } - "rtk_models" => { - runtime_config.rtk_models = value.clone(); - } - "forward_client_auth" => { - // Defensively re-validate against the same rule - // enforced by put_config (a tampered/hand-edited SQLite - // row could otherwise re-enable a misconfigured toggle - // that would let multiple distinct PROXY_API_KEYS - // entries each redirect the upstream Anthropic - // credential -- see - // server/middleware/auth.rs::forward_client_auth_misconfigured). - let wants_enabled = value == "true"; - if wants_enabled - && anyllm_proxy::server::middleware::forward_client_auth_misconfigured( - anyllm_proxy::server::middleware::distinct_static_key_count(), - anyllm_proxy::server::middleware::open_relay_active(), - ) - { - tracing::warn!( - "ignoring persisted forward_client_auth=true override: 2+ \ - PROXY_API_KEYS entries with no PROXY_OPEN_RELAY would let \ - different callers each redirect the upstream Anthropic credential" - ); - } else { - runtime_config.forward_client_auth = wants_enabled; - } - } - "tool_guardrail_mode" => { - if value - .parse::() - .is_ok() - { - runtime_config.tool_guardrail_mode = value.clone(); - } else { - tracing::warn!( - value = %value, - "ignoring invalid tool_guardrail_mode override from database" - ); - } - } - "optimizer_mode" => { - if value.parse::().is_ok() { - runtime_config.optimizer_mode = value.clone(); - } else { - tracing::warn!( - value = %value, - "ignoring invalid optimizer_mode override from database" - ); - } - } - k if k.ends_with(".big_model") => { - let backend = k.strip_suffix(".big_model").unwrap(); - if let Some(m) = runtime_config.model_mappings.get_mut(backend) { - m.big_model = value.clone(); - } - } - k if k.ends_with(".small_model") => { - let backend = k.strip_suffix(".small_model").unwrap(); - if let Some(m) = runtime_config.model_mappings.get_mut(backend) { - m.small_model = value.clone(); - } - } - _ => { - tracing::debug!(key = %key, "unknown config override, skipping"); - } - } - } - if !overrides.is_empty() { - tracing::info!( - count = overrides.len(), - "applied config overrides from database" - ); - } - } - - if log_bodies_enabled_by_override { - tracing::warn!( - "persisted admin override enabled LOG_BODIES: request and response bodies will be \ - logged at debug level. This may expose sensitive data (prompts, API keys, PII)." - ); - } - if redact_secrets_enabled_by_override { - tracing::warn!( - "persisted admin override enabled REDACT_SECRETS: upstream JSON/text request payloads \ - will be scanned and detected secrets will be replaced before forwarding." - ); - } - if let Err(message) = - anyllm_proxy::server::ensure_secret_redaction_available(runtime_config.redact_secrets) - { - panic!("{}", message); - } - - let runtime_config = Arc::new(std::sync::RwLock::new(runtime_config)); - - // Build the log_reload closure that captures the reload handle. - // We get this reload handle from caller or trace initialization. - // Note: reload_handle is captured dynamically during async_main setup, - // so we can resolve this by passing reload_handle, or keeping the logic - // in main mod.rs. Let's pass the log_reload closure or reload_handle in. - // Actually, let's define reload_handle as a parameter to init_admin, or pass - // log_reload in. - // Let's pass: - // log_reload: Arc bool + Send + Sync> - - // Load active virtual keys from SQLite into in-memory DashMap. - let virtual_keys = Arc::new(dashmap::DashMap::new()); - { - if let Ok(active_keys) = admin::db::load_active_virtual_keys(&conn) { - for key_row in &active_keys { - if let Some(hash_bytes) = admin::keys::hash_from_hex(&key_row.key_hash) { - virtual_keys.insert( - hash_bytes, - admin::keys::VirtualKeyMeta { - id: key_row.id, - description: key_row.description.clone(), - expires_at: key_row.expires_at.as_deref().and_then(|s| { - anyllm_proxy::integrations::langfuse::iso8601_to_epoch(s) - .and_then(|e| i64::try_from(e).ok()) - }), - rpm_limit: key_row.rpm_limit, - tpm_limit: key_row.tpm_limit, - rate_state: Arc::new(admin::keys::RateLimitState::new()), - role: admin::keys::KeyRole::from_str_or_default(&key_row.role), - max_budget_usd: key_row.max_budget_usd, - budget_duration: key_row - .budget_duration - .as_deref() - .and_then(admin::keys::BudgetDuration::parse), - period_start: key_row.period_start.clone(), - period_spend_usd: key_row.period_spend_usd, - allowed_models: key_row.allowed_models.clone(), - allowed_routes: key_row.allowed_routes.clone(), - }, - ); - } - } - tracing::info!( - count = active_keys.len(), - "loaded virtual API keys from database" - ); - } - } - - // Make virtual keys and HMAC secret available to the auth middleware. - anyllm_proxy::server::middleware::set_virtual_keys(virtual_keys.clone()); - anyllm_proxy::server::middleware::set_hmac_secret(hmac_secret.clone()); - - let virtual_keys_pruner = virtual_keys.clone(); - tokio::spawn(async move { - let mut interval = tokio::time::interval(std::time::Duration::from_secs(300)); - loop { - interval.tick().await; - let now = anyllm_proxy::admin::keys::now_ms(); - // Single pass: slide rate-limit windows forward (frees old - // buckets) and drop expired keys. Without active eviction, - // expired keys accumulate in the DashMap until next auth use. - let now_secs = (now / 1000) as i64; - virtual_keys_pruner.retain(|_, v| { - let _ = v.rate_state.check_rpm(0, now); - let _ = v.rate_state.check_tpm(0, now); - v.expires_at.is_none_or(|exp| now_secs < exp) - }); - } - }); - - // Load managed backends from SQLite into in-memory HashMap. - let managed_backends = { - let mut map = std::collections::HashMap::new(); - if let Ok(rows) = admin::db::list_managed_backends(&conn) { - for row in rows { - match provider_catalog.get_provider(&row.provider_id) { - None => { - tracing::warn!( - provider_id = %row.provider_id, - backend_id = %row.id, - "managed backend references unknown provider; skipping" - ); - } - Some(provider) => { - match anyllm_proxy::admin::routes::managed_backends::row_to_backend_config( - &row, provider, - ) { - Err(e) => { - tracing::warn!( - provider_id = %row.provider_id, - backend_id = %row.id, - error = %e.message(), - "managed backend configuration is invalid; skipping" - ); - } - Ok(bc) => { - let client = - anyllm_proxy::backend::BackendClient::from_backend_config(&bc); - map.insert(row.name.clone(), (row, client)); - } - } - } - } - } - } - tracing::info!(count = map.len(), "loaded managed backends from SQLite"); - Arc::new(std::sync::RwLock::new(map)) - }; - - let db = Arc::new(std::sync::Mutex::new(conn)); - - // Compile the initial route dispatch table from the admin DB (enabled routes - // + providers + backends). Rebuilt on route/backend CRUD via rebuild_route_router. - let route_router = { - let conn = db.lock().unwrap_or_else(|e| e.into_inner()); - let rr = config::route_router::RouteRouter::build_from_db(&conn).unwrap_or_else(|e| { - tracing::warn!(error = %e, "failed to build route router from DB; starting empty"); - config::route_router::RouteRouter::empty() - }); - tracing::info!( - has_routes = !rr.is_empty(), - "initialized route dispatch table" - ); - Some(Arc::new(std::sync::RwLock::new(rr))) - }; - - let (events_tx, _) = tokio::sync::broadcast::channel(1024); - let log_tx = admin::db::spawn_write_buffer(db.clone()); - - let backend_metrics: std::collections::HashMap = - std::collections::HashMap::new(); - - let log_reload: Arc bool + Send + Sync> = { - let handle = reload_handle; - Arc::new( - move |new_filter: &str| match tracing_subscriber::EnvFilter::try_new(new_filter) { - Ok(f) => handle.reload(f).is_ok(), - Err(e) => { - tracing::error!(filter = new_filter, error = %e, "invalid log filter string"); - false - } - }, - ) - }; - - // Admin token: use env var or generate 256-bit random hex written to a file. - let admin_token = match std::env::var("ADMIN_TOKEN") { - Ok(t) => { - if t.len() < 32 { - tracing::warn!( - len = t.len(), - "ADMIN_TOKEN is shorter than 32 characters; use a longer random value to reduce brute-force risk (generate one with: openssl rand -hex 32)" - ); - } - t - } - Err(_) => { - let mut buf = [0u8; 32]; - getrandom::fill(&mut buf).expect("getrandom failed"); - let token = hex::encode(buf); - let token_path = crate::main_helpers::bootstrap::resolve_admin_token_path(data_dir); - let token_path_str = token_path.to_string_lossy().to_string(); - // Write token to file with restrictive permissions instead of stderr - if let Err(e) = - crate::main_helpers::bootstrap::write_token_file(&token_path_str, &token) - { - panic!( - "Cannot write admin token to {token_path_str}: {e}. Set ADMIN_TOKEN env var explicitly or ensure the path is writable." - ); - } else { - tracing::info!( - path = %token_path_str, - "admin token written to {token_path_str}; set ADMIN_TOKEN for a fixed token across restarts" - ); - } - token - } - }; - // Keep a plaintext copy to return for the loopback startup URL before the - // token is wrapped/zeroized and moved into the router. - let admin_token_plain = admin_token.clone(); - let admin_token = Arc::new(zeroize::Zeroizing::new(admin_token)); - - let shared = admin::state::SharedState { - db: db.clone(), - events_tx: events_tx.clone(), - runtime_config: runtime_config.clone(), - runtime_defaults, - backend_metrics: Arc::new(backend_metrics), - log_tx, - log_reload: Some(log_reload), - config_write_lock: Arc::new(tokio::sync::Mutex::new(())), - virtual_keys, - hmac_secret, - model_router: model_router.clone(), - route_router, - provider_catalog: provider_catalog.clone(), - mcp_manager: tool_engine_state - .as_ref() - .and_then(|s| s.mcp_manager.clone()), - issued_csrf_tokens: Arc::new( - moka::sync::Cache::builder() - .max_capacity(1_000) - .time_to_live(std::time::Duration::from_secs(86400)) - .build(), - ), - started_at: std::time::SystemTime::now(), - listen_port: multi_config.listen_port, - managed_backends, - }; - - // Provider model cache auto-refresh (only when --webui is active). - let auto_refresh = matches!( - std::env::var("PROVIDER_AUTO_REFRESH").as_deref(), - Ok("1") | Ok("true") | Ok("yes") - ); - let refresh_interval_hours: u64 = std::env::var("PROVIDER_REFRESH_INTERVAL_HOURS") - .ok() - .and_then(|v| v.parse().ok()) - .unwrap_or(168); - - if auto_refresh { - let shared_for_refresh = shared.clone(); - let client = crate::main_helpers::providers_cmd::PROVIDER_REFRESH_CLIENT.clone(); - tokio::spawn(async move { - let interval = std::time::Duration::from_secs(refresh_interval_hours * 3600); - tokio::time::sleep(std::time::Duration::from_secs(30)).await; - loop { - let providers: Vec<_> = shared_for_refresh - .provider_catalog - .all_providers() - .cloned() - .collect(); - for provider in providers { - if !provider.capabilities.chat_completions { - continue; - } - if provider.default_base_url.is_empty() { - continue; - } - let api_key = provider - .env_vars - .iter() - .find_map(|v| std::env::var(v.as_str()).ok()); - if api_key.is_none() { - continue; - } - - let url = format!( - "{}/v1/models", - provider.default_base_url.trim_end_matches('/') - ); - let provider_id = provider.id.clone(); - let mut req = client.get(&url); - if let Some(ref key) = api_key { - req = req.header("Authorization", format!("Bearer {key}")); - } - match req.send().await { - Err(e) => tracing::warn!( - provider = %provider_id, - error = %e, - "provider auto-refresh failed" - ), - Ok(resp) if !resp.status().is_success() => tracing::warn!( - provider = %provider_id, - status = %resp.status(), - "provider auto-refresh upstream error" - ), - Ok(resp) => match resp.json::().await { - Err(e) => tracing::warn!( - provider = %provider_id, - error = %e, - "provider auto-refresh: invalid JSON response" - ), - Ok(json) => { - let model_ids: Vec = json - .get("data") - .and_then(|d| d.as_array()) - .map(|arr| { - arr.iter() - .filter_map(|m| m.get("id")?.as_str().map(String::from)) - .collect() - }) - .unwrap_or_default(); - let count = model_ids.len(); - let db_ref = shared_for_refresh.db.clone(); - let pid = provider_id.clone(); - let _ = tokio::task::spawn_blocking(move || { - let mut conn_guard = - db_ref.lock().unwrap_or_else(|e| e.into_inner()); - if let Err(e) = admin::db::upsert_provider_models_cache( - &mut conn_guard, - &pid, - &model_ids, - ) { - tracing::warn!( - provider = %pid, - error = %e, - "failed to save auto-refresh results" - ); - } - }) - .await; - tracing::info!( - provider = %provider_id, - count = count, - "auto-refreshed provider model cache" - ); - } - }, - } - } - tokio::time::sleep(interval).await; - } - }); - tracing::info!( - interval_hours = refresh_interval_hours, - "provider auto-refresh enabled" - ); - } - - // Spawn periodic tasks: log retention and metrics snapshot broadcast. - let retention_days: u32 = std::env::var("ADMIN_LOG_RETENTION_DAYS") - .ok() - .and_then(|v| v.parse().ok()) - .unwrap_or(7); - - let retention_db = shared.db.clone(); - tokio::spawn(async move { - let mut interval = tokio::time::interval(std::time::Duration::from_secs(3600)); - loop { - interval.tick().await; - admin::state::with_db( - &retention_db, - move |conn_ref| match admin::db::purge_old_logs(conn_ref, retention_days) { - Ok(n) if n > 0 => { - tracing::info!(purged = n, "purged old request log entries") - } - Err(e) => tracing::error!(error = %e, "failed to purge old logs"), - _ => {} - }, - ) - .await; - } - }); - - // Periodic metrics snapshot broadcast (every 5 seconds) for WebSocket dashboard. - let snapshot_shared = shared.clone(); - tokio::spawn(async move { - let mut interval = tokio::time::interval(std::time::Duration::from_secs(5)); - loop { - interval.tick().await; - if snapshot_shared.events_tx.receiver_count() == 0 { - continue; - } - let mut aggregate = anyllm_proxy::metrics::MetricsSnapshot::default(); - for m in snapshot_shared.backend_metrics.values() { - let snap = m.snapshot(); - aggregate.requests_total += snap.requests_total; - aggregate.requests_error += snap.requests_error; - aggregate.requests_success += snap.requests_success; - aggregate.streams_started += snap.streams_started; - aggregate.streams_completed += snap.streams_completed; - aggregate.streams_failed += snap.streams_failed; - aggregate.streams_client_disconnected += snap.streams_client_disconnected; - } - let error_rate = aggregate.error_rate(); - let now_secs = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs(); - let since = now_secs.saturating_sub(60); - let rpm = admin::state::with_db(&snapshot_shared.db, move |conn_ref| { - admin::db::count_requests_since(conn_ref, since).unwrap_or(0) - }) - .await - .unwrap_or(0) as f64; - - let snapshot = admin::state::MetricsSnapshotData { - total_requests: aggregate.requests_total, - successful_requests: aggregate.requests_success, - failed_requests: aggregate.requests_error, - requests_per_minute: rpm, - p50_latency_ms: None, - p95_latency_ms: None, - error_rate, - streams_started: aggregate.streams_started, - streams_completed: aggregate.streams_completed, - streams_failed: aggregate.streams_failed, - streams_client_disconnected: aggregate.streams_client_disconnected, - }; - let _ = snapshot_shared - .events_tx - .send(admin::state::AdminEvent::MetricsSnapshot(snapshot)); - } - }); - - // Spawn periodic background health checker with a snapshot of backend base URLs. - let backend_urls: Vec<(String, String)> = multi_config - .backends - .iter() - .map(|(name, bc)| (name.clone(), bc.base_url.clone())) - .collect(); - admin::health_check::spawn(shared.clone(), backend_urls); - - // Bind admin listener - let admin_app = admin::routes::admin_router(shared.clone(), admin_token); - let admin_addr = format!("{admin_bind}:{admin_port}"); - let admin_listener = tokio::net::TcpListener::bind(&admin_addr) - .await - .unwrap_or_else(|e| panic!("failed to bind admin to {admin_addr}: {e}")); - tracing::info!("admin listening on {admin_addr}"); - - Some(( - shared, - admin_app, - admin_listener, - admin_port, - admin_token_plain, - )) -} diff --git a/crates/proxy/src/main_helpers/async_main/admin/config.rs b/crates/proxy/src/main_helpers/async_main/admin/config.rs new file mode 100644 index 0000000..9d4e7c5 --- /dev/null +++ b/crates/proxy/src/main_helpers/async_main/admin/config.rs @@ -0,0 +1,306 @@ +use anyllm_proxy::admin; +use anyllm_proxy::config; +use anyllm_proxy::server::state as server_state; +use std::path::Path; +use std::sync::Arc; + +pub(crate) fn load_runtime_config( + multi_config: &config::MultiConfig, + tool_engine_state: &Option>, +) -> ( + admin::state::RuntimeConfig, + admin::state::RuntimeConfigDefaults, +) { + let mut model_mappings = indexmap::IndexMap::new(); + for (name, bc) in &multi_config.backends { + model_mappings.insert(name.clone(), bc.model_mapping.clone()); + } + let log_level = std::env::var("RUST_LOG").unwrap_or_else(|_| "info".into()); + let tool_guardrail_default = tool_engine_state + .as_ref() + .map(|engine| engine.guardrails.mode.as_str().to_string()) + .unwrap_or_else(|| { + anyllm_proxy::tools::ToolGuardrailMode::Disabled + .as_str() + .to_string() + }); + + let runtime_config = admin::state::RuntimeConfig { + model_mappings, + log_level, + log_bodies: multi_config.log_bodies, + redact_secrets: multi_config.redact_secrets, + anthropic_thinking_repair: multi_config.anthropic_thinking_repair, + pxpipe_compress: multi_config.pxpipe_compress, + pxpipe_models: anyllm_proxy::pxpipe::resolve_default_models_csv(), + rtk_compress: anyllm_proxy::rtk::resolve_default_enabled(), + rtk_models: anyllm_proxy::rtk::resolve_default_models_csv(), + forward_client_auth: multi_config.forward_client_auth, + tool_guardrail_mode: tool_guardrail_default.clone(), + optimizer_mode: anyllm_proxy::optimizer::resolve_default_mode() + .as_str() + .to_string(), + }; + let runtime_defaults = admin::state::RuntimeConfigDefaults { + log_bodies: multi_config.log_bodies, + redact_secrets: multi_config.redact_secrets, + anthropic_thinking_repair: multi_config.anthropic_thinking_repair, + pxpipe_compress: multi_config.pxpipe_compress, + pxpipe_models: anyllm_proxy::pxpipe::resolve_default_models_csv(), + rtk_compress: anyllm_proxy::rtk::resolve_default_enabled(), + rtk_models: anyllm_proxy::rtk::resolve_default_models_csv(), + forward_client_auth: multi_config.forward_client_auth, + tool_guardrail_mode: tool_guardrail_default, + optimizer_mode: anyllm_proxy::optimizer::resolve_default_mode() + .as_str() + .to_string(), + }; + (runtime_config, runtime_defaults) +} + +pub(crate) fn apply_config_overrides( + conn: &rusqlite::Connection, + runtime_config: &mut admin::state::RuntimeConfig, + multi_config: &config::MultiConfig, +) -> (bool, bool) { + let mut log_bodies_enabled_by_override = false; + let mut redact_secrets_enabled_by_override = false; + + if let Ok(overrides) = admin::db::get_config_overrides(conn) { + for (key, value, _) in &overrides { + match key.as_str() { + "log_level" => { + const ALLOWED_LOG_LEVELS: &[&str] = &["error", "warn", "info", "debug"]; + let normalized = value.trim().to_lowercase(); + if ALLOWED_LOG_LEVELS.contains(&normalized.as_str()) { + runtime_config.log_level = normalized; + } else { + tracing::warn!( + value = %value, + "ignoring invalid log_level override from database" + ); + } + } + "log_bodies" => { + runtime_config.log_bodies = value == "true"; + log_bodies_enabled_by_override = + runtime_config.log_bodies && !multi_config.log_bodies; + } + "redact_secrets" => { + runtime_config.redact_secrets = value == "true"; + redact_secrets_enabled_by_override = + runtime_config.redact_secrets && !multi_config.redact_secrets; + } + "anthropic_thinking_repair" => { + runtime_config.anthropic_thinking_repair = value == "true"; + } + "pxpipe_compress" => { + runtime_config.pxpipe_compress = value == "true"; + } + "pxpipe_models" => { + runtime_config.pxpipe_models = value.clone(); + } + "rtk_compress" => { + runtime_config.rtk_compress = value == "true"; + } + "rtk_models" => { + runtime_config.rtk_models = value.clone(); + } + "forward_client_auth" => { + let wants_enabled = value == "true"; + if wants_enabled + && anyllm_proxy::server::middleware::forward_client_auth_misconfigured( + anyllm_proxy::server::middleware::distinct_static_key_count(), + anyllm_proxy::server::middleware::open_relay_active(), + ) + { + tracing::warn!( + "ignoring persisted forward_client_auth=true override: 2+ \ + PROXY_API_KEYS entries with no PROXY_OPEN_RELAY would let \ + different callers each redirect the upstream Anthropic credential" + ); + } else { + runtime_config.forward_client_auth = wants_enabled; + } + } + "tool_guardrail_mode" => { + if value + .parse::() + .is_ok() + { + runtime_config.tool_guardrail_mode = value.clone(); + } else { + tracing::warn!( + value = %value, + "ignoring invalid tool_guardrail_mode override from database" + ); + } + } + "optimizer_mode" => { + if value.parse::().is_ok() { + runtime_config.optimizer_mode = value.clone(); + } else { + tracing::warn!( + value = %value, + "ignoring invalid optimizer_mode override from database" + ); + } + } + k if k.ends_with(".big_model") => { + let backend = k.strip_suffix(".big_model").unwrap(); + if let Some(m) = runtime_config.model_mappings.get_mut(backend) { + m.big_model = value.clone(); + } + } + k if k.ends_with(".small_model") => { + let backend = k.strip_suffix(".small_model").unwrap(); + if let Some(m) = runtime_config.model_mappings.get_mut(backend) { + m.small_model = value.clone(); + } + } + _ => { + tracing::debug!(key = %key, "unknown config override, skipping"); + } + } + } + if !overrides.is_empty() { + tracing::info!( + count = overrides.len(), + "applied config overrides from database" + ); + } + } + ( + log_bodies_enabled_by_override, + redact_secrets_enabled_by_override, + ) +} + +pub(crate) fn load_virtual_keys( + conn: &rusqlite::Connection, +) -> Arc> { + let virtual_keys = Arc::new(dashmap::DashMap::new()); + if let Ok(active_keys) = admin::db::load_active_virtual_keys(conn) { + for key_row in &active_keys { + if let Some(hash_bytes) = admin::keys::hash_from_hex(&key_row.key_hash) { + virtual_keys.insert( + hash_bytes, + admin::keys::VirtualKeyMeta { + id: key_row.id, + description: key_row.description.clone(), + expires_at: key_row.expires_at.as_deref().and_then(|s| { + anyllm_proxy::integrations::langfuse::iso8601_to_epoch(s) + .and_then(|e| i64::try_from(e).ok()) + }), + rpm_limit: key_row.rpm_limit, + tpm_limit: key_row.tpm_limit, + rate_state: Arc::new(admin::keys::RateLimitState::new()), + role: admin::keys::KeyRole::from_str_or_default(&key_row.role), + max_budget_usd: key_row.max_budget_usd, + budget_duration: key_row + .budget_duration + .as_deref() + .and_then(admin::keys::BudgetDuration::parse), + period_start: key_row.period_start.clone(), + period_spend_usd: key_row.period_spend_usd, + allowed_models: key_row.allowed_models.clone(), + allowed_routes: key_row.allowed_routes.clone(), + }, + ); + } + } + tracing::info!( + count = active_keys.len(), + "loaded virtual API keys from database" + ); + } + virtual_keys +} + +pub(crate) fn load_managed_backends( + conn: &rusqlite::Connection, + provider_catalog: &Arc, +) -> Arc< + std::sync::RwLock< + std::collections::HashMap< + String, + ( + anyllm_proxy::admin::db::ManagedBackendRow, + anyllm_proxy::backend::BackendClient, + ), + >, + >, +> { + let mut map = std::collections::HashMap::new(); + if let Ok(rows) = admin::db::list_managed_backends(conn) { + for row in rows { + match provider_catalog.get_provider(&row.provider_id) { + None => { + tracing::warn!( + provider_id = %row.provider_id, + backend_id = %row.id, + "managed backend references unknown provider; skipping" + ); + } + Some(provider) => { + match anyllm_proxy::admin::routes::managed_backends::row_to_backend_config( + &row, provider, + ) { + Err(e) => { + tracing::warn!( + provider_id = %row.provider_id, + backend_id = %row.id, + error = %e.message(), + "managed backend configuration is invalid; skipping" + ); + } + Ok(bc) => { + let client = + anyllm_proxy::backend::BackendClient::from_backend_config(&bc); + map.insert(row.name.clone(), (row, client)); + } + } + } + } + } + } + tracing::info!(count = map.len(), "loaded managed backends from SQLite"); + Arc::new(std::sync::RwLock::new(map)) +} + +pub(crate) fn resolve_admin_token(data_dir: &Path) -> (String, Arc>) { + let admin_token = match std::env::var("ADMIN_TOKEN") { + Ok(t) => { + if t.len() < 32 { + tracing::warn!( + len = t.len(), + "ADMIN_TOKEN is shorter than 32 characters; use a longer random value to reduce brute-force risk (generate one with: openssl rand -hex 32)" + ); + } + t + } + Err(_) => { + let mut buf = [0u8; 32]; + getrandom::fill(&mut buf).expect("getrandom failed"); + let token = hex::encode(buf); + let token_path = crate::main_helpers::bootstrap::resolve_admin_token_path(data_dir); + let token_path_str = token_path.to_string_lossy().to_string(); + if let Err(e) = + crate::main_helpers::bootstrap::write_token_file(&token_path_str, &token) + { + panic!( + "Cannot write admin token to {token_path_str}: {e}. Set ADMIN_TOKEN env var explicitly or ensure the path is writable." + ); + } else { + tracing::info!( + path = %token_path_str, + "admin token written to {token_path_str}; set ADMIN_TOKEN for a fixed token across restarts" + ); + } + token + } + }; + let admin_token_plain = admin_token.clone(); + let admin_token_wrapped = Arc::new(zeroize::Zeroizing::new(admin_token)); + (admin_token_plain, admin_token_wrapped) +} diff --git a/crates/proxy/src/main_helpers/async_main/admin/mod.rs b/crates/proxy/src/main_helpers/async_main/admin/mod.rs new file mode 100644 index 0000000..ad7844c --- /dev/null +++ b/crates/proxy/src/main_helpers/async_main/admin/mod.rs @@ -0,0 +1,235 @@ +pub(crate) mod config; +pub(crate) mod tasks; + +use anyllm_proxy::admin; +use anyllm_proxy::config as proxy_config; +use anyllm_proxy::server::state as server_state; +use std::path::Path; +use std::sync::Arc; + +#[allow(clippy::too_many_lines)] +pub(crate) async fn init_admin( + args: &[String], + data_dir: &Path, + multi_config: &proxy_config::MultiConfig, + model_router: Option>>, + tool_engine_state: &Option>, + reload_handle: tracing_subscriber::reload::Handle< + tracing_subscriber::EnvFilter, + tracing_subscriber::Registry, + >, +) -> Option<( + admin::state::SharedState, + axum::Router, + tokio::net::TcpListener, + u16, + String, +)> { + let flag_set = args.iter().any(|a| a == "--webui" || a == "--admin"); + let force_disabled = matches!( + std::env::var("DISABLE_ADMIN").as_deref(), + Ok("1") | Ok("true") | Ok("yes") + ); + let enable_admin = flag_set && !force_disabled; + + if !enable_admin { + return None; + } + + let provider_catalog = Arc::new(anyllm_providers::ProviderCatalog::bundled()); + + let admin_port: u16 = match std::env::var("ADMIN_PORT") { + Ok(val) => val + .parse::() + .unwrap_or_else(|_| panic!("ADMIN_PORT must be a number in 1-65535, got '{val}'")), + Err(_) => 3001, + }; + if admin_port == 0 { + panic!("ADMIN_PORT cannot be 0"); + } + if admin_port < 1024 { + tracing::warn!( + port = admin_port, + "ADMIN_PORT is in the privileged range (< 1024); binding may fail without elevated privileges" + ); + } + + let admin_bind = { + let raw = std::env::var("ADMIN_BIND").unwrap_or_else(|_| "127.0.0.1".into()); + match raw.to_ascii_lowercase().as_str() { + "localhost" => "127.0.0.1".to_string(), + "localhost6" | "ip6-localhost" | "ip6-loopback" => "::1".to_string(), + _ => raw, + } + }; + if admin_bind.parse::().is_err() { + panic!( + "ADMIN_BIND must be an IP address (e.g. 127.0.0.1 or 0.0.0.0), not a hostname — got '{}'", + std::env::var("ADMIN_BIND").unwrap_or_default() + ); + } + + if admin_port == multi_config.listen_port { + panic!( + "ADMIN_PORT ({admin_port}) must differ from LISTEN_PORT ({})", + multi_config.listen_port + ); + } + + let db_path = crate::main_helpers::bootstrap::resolve_db_path(data_dir); + let conn = + rusqlite::Connection::open(&db_path).expect("failed to open SQLite database for admin"); + admin::db::init_db(&conn).expect("failed to initialize admin database schema"); + let hmac_secret = Arc::new(admin::db::ensure_hmac_secret(&conn)); + + // Load configuration and apply database overrides + let (mut runtime_config, runtime_defaults) = + config::load_runtime_config(multi_config, tool_engine_state); + let (log_bodies_override, redact_secrets_override) = + config::apply_config_overrides(&conn, &mut runtime_config, multi_config); + + if log_bodies_override { + tracing::warn!( + "persisted admin override enabled LOG_BODIES: request and response bodies will be \ + logged at debug level. This may expose sensitive data (prompts, API keys, PII)." + ); + } + if redact_secrets_override { + tracing::warn!( + "persisted admin override enabled REDACT_SECRETS: upstream JSON/text request payloads \ + will be scanned and detected secrets will be replaced before forwarding." + ); + } + if let Err(message) = + anyllm_proxy::server::ensure_secret_redaction_available(runtime_config.redact_secrets) + { + panic!("{}", message); + } + + let runtime_config = Arc::new(std::sync::RwLock::new(runtime_config)); + + // Load virtual API keys and setup rate-limit state + let virtual_keys = config::load_virtual_keys(&conn); + anyllm_proxy::server::middleware::set_virtual_keys(virtual_keys.clone()); + anyllm_proxy::server::middleware::set_hmac_secret(hmac_secret.clone()); + + // Spawn key pruner background task + tasks::spawn_background_pruner(virtual_keys.clone()); + + // Load managed backends + let managed_backends = config::load_managed_backends(&conn, &provider_catalog); + + let db = Arc::new(std::sync::Mutex::new(conn)); + + // Compile initial route dispatch table + let route_router = { + let conn_guard = db.lock().unwrap_or_else(|e| e.into_inner()); + let rr = proxy_config::route_router::RouteRouter::build_from_db(&conn_guard) + .unwrap_or_else(|e| { + tracing::warn!(error = %e, "failed to build route router from DB; starting empty"); + proxy_config::route_router::RouteRouter::empty() + }); + tracing::info!( + has_routes = !rr.is_empty(), + "initialized route dispatch table" + ); + Some(Arc::new(std::sync::RwLock::new(rr))) + }; + + let (events_tx, _) = tokio::sync::broadcast::channel(1024); + let log_tx = admin::db::spawn_write_buffer(db.clone()); + let backend_metrics = std::collections::HashMap::new(); + + let log_reload: Arc bool + Send + Sync> = { + let handle = reload_handle; + Arc::new( + move |new_filter: &str| match tracing_subscriber::EnvFilter::try_new(new_filter) { + Ok(f) => handle.reload(f).is_ok(), + Err(e) => { + tracing::error!(filter = new_filter, error = %e, "invalid log filter string"); + false + } + }, + ) + }; + + // Resolve admin token + let (admin_token_plain, admin_token) = config::resolve_admin_token(data_dir); + + let shared = admin::state::SharedState { + db: db.clone(), + events_tx: events_tx.clone(), + runtime_config: runtime_config.clone(), + runtime_defaults, + backend_metrics: Arc::new(backend_metrics), + log_tx, + log_reload: Some(log_reload), + config_write_lock: Arc::new(tokio::sync::Mutex::new(())), + virtual_keys, + hmac_secret, + model_router: model_router.clone(), + route_router, + provider_catalog: provider_catalog.clone(), + mcp_manager: tool_engine_state + .as_ref() + .and_then(|s| s.mcp_manager.clone()), + issued_csrf_tokens: Arc::new( + moka::sync::Cache::builder() + .max_capacity(1_000) + .time_to_live(std::time::Duration::from_secs(86400)) + .build(), + ), + started_at: std::time::SystemTime::now(), + listen_port: multi_config.listen_port, + managed_backends, + }; + + // Provider model cache auto-refresh (only when --webui is active) + let auto_refresh = matches!( + std::env::var("PROVIDER_AUTO_REFRESH").as_deref(), + Ok("1") | Ok("true") | Ok("yes") + ); + let refresh_interval_hours: u64 = std::env::var("PROVIDER_REFRESH_INTERVAL_HOURS") + .ok() + .and_then(|v| v.parse().ok()) + .unwrap_or(168); + + if auto_refresh { + tasks::spawn_auto_refresh_task(shared.clone(), refresh_interval_hours); + tracing::info!( + interval_hours = refresh_interval_hours, + "provider auto-refresh enabled" + ); + } + + // Periodic tasks: log retention and metrics snapshot + let retention_days: u32 = std::env::var("ADMIN_LOG_RETENTION_DAYS") + .ok() + .and_then(|v| v.parse().ok()) + .unwrap_or(7); + tasks::spawn_periodic_tasks(shared.clone(), retention_days); + + // Spawn periodic background health checker + let backend_urls: Vec<(String, String)> = multi_config + .backends + .iter() + .map(|(name, bc)| (name.clone(), bc.base_url.clone())) + .collect(); + admin::health_check::spawn(shared.clone(), backend_urls); + + // Bind admin listener + let admin_app = admin::routes::admin_router(shared.clone(), admin_token); + let admin_addr = format!("{admin_bind}:{admin_port}"); + let admin_listener = tokio::net::TcpListener::bind(&admin_addr) + .await + .unwrap_or_else(|e| panic!("failed to bind admin to {admin_addr}: {e}")); + tracing::info!("admin listening on {admin_addr}"); + + Some(( + shared, + admin_app, + admin_listener, + admin_port, + admin_token_plain, + )) +} diff --git a/crates/proxy/src/main_helpers/async_main/admin/tasks.rs b/crates/proxy/src/main_helpers/async_main/admin/tasks.rs new file mode 100644 index 0000000..cf97a90 --- /dev/null +++ b/crates/proxy/src/main_helpers/async_main/admin/tasks.rs @@ -0,0 +1,187 @@ +use anyllm_proxy::admin; +use std::sync::Arc; + +pub(crate) fn spawn_background_pruner( + virtual_keys: Arc>, +) { + tokio::spawn(async move { + let mut interval = tokio::time::interval(std::time::Duration::from_secs(300)); + loop { + interval.tick().await; + let now = anyllm_proxy::admin::keys::now_ms(); + let now_secs = (now / 1000) as i64; + virtual_keys.retain(|_, v| { + let _ = v.rate_state.check_rpm(0, now); + let _ = v.rate_state.check_tpm(0, now); + v.expires_at.is_none_or(|exp| now_secs < exp) + }); + } + }); +} + +pub(crate) fn spawn_auto_refresh_task( + shared: admin::state::SharedState, + refresh_interval_hours: u64, +) { + let client = crate::main_helpers::providers_cmd::PROVIDER_REFRESH_CLIENT.clone(); + tokio::spawn(async move { + let interval = std::time::Duration::from_secs(refresh_interval_hours * 3600); + tokio::time::sleep(std::time::Duration::from_secs(30)).await; + loop { + let providers: Vec<_> = shared.provider_catalog.all_providers().cloned().collect(); + for provider in providers { + if !provider.capabilities.chat_completions { + continue; + } + if provider.default_base_url.is_empty() { + continue; + } + let api_key = provider + .env_vars + .iter() + .find_map(|v| std::env::var(v.as_str()).ok()); + if api_key.is_none() { + continue; + } + + let url = format!( + "{}/v1/models", + provider.default_base_url.trim_end_matches('/') + ); + let provider_id = provider.id.clone(); + let mut req = client.get(&url); + if let Some(ref key) = api_key { + req = req.header("Authorization", format!("Bearer {key}")); + } + match req.send().await { + Err(e) => tracing::warn!( + provider = %provider_id, + error = %e, + "provider auto-refresh failed" + ), + Ok(resp) if !resp.status().is_success() => tracing::warn!( + provider = %provider_id, + status = %resp.status(), + "provider auto-refresh upstream error" + ), + Ok(resp) => match resp.json::().await { + Err(e) => tracing::warn!( + provider = %provider_id, + error = %e, + "provider auto-refresh: invalid JSON response" + ), + Ok(json) => { + let model_ids: Vec = json + .get("data") + .and_then(|d| d.as_array()) + .map(|arr| { + arr.iter() + .filter_map(|m| m.get("id")?.as_str().map(String::from)) + .collect() + }) + .unwrap_or_default(); + let count = model_ids.len(); + let db_ref = shared.db.clone(); + let pid = provider_id.clone(); + let _ = tokio::task::spawn_blocking(move || { + let mut conn_guard = + db_ref.lock().unwrap_or_else(|e| e.into_inner()); + if let Err(e) = admin::db::upsert_provider_models_cache( + &mut conn_guard, + &pid, + &model_ids, + ) { + tracing::warn!( + provider = %pid, + error = %e, + "failed to save auto-refresh results" + ); + } + }) + .await; + tracing::info!( + provider = %provider_id, + count = count, + "auto-refreshed provider model cache" + ); + } + }, + } + } + tokio::time::sleep(interval).await; + } + }); +} + +pub(crate) fn spawn_periodic_tasks(shared: admin::state::SharedState, retention_days: u32) { + // 1. Log retention task + let retention_db = shared.db.clone(); + tokio::spawn(async move { + let mut interval = tokio::time::interval(std::time::Duration::from_secs(3600)); + loop { + interval.tick().await; + admin::state::with_db( + &retention_db, + move |conn_ref| match admin::db::purge_old_logs(conn_ref, retention_days) { + Ok(n) if n > 0 => { + tracing::info!(purged = n, "purged old request log entries") + } + Err(e) => tracing::error!(error = %e, "failed to purge old logs"), + _ => {} + }, + ) + .await; + } + }); + + // 2. Periodic metrics snapshot broadcast task + let snapshot_shared = shared.clone(); + tokio::spawn(async move { + let mut interval = tokio::time::interval(std::time::Duration::from_secs(5)); + loop { + interval.tick().await; + if snapshot_shared.events_tx.receiver_count() == 0 { + continue; + } + let mut aggregate = anyllm_proxy::metrics::MetricsSnapshot::default(); + for m in snapshot_shared.backend_metrics.values() { + let snap = m.snapshot(); + aggregate.requests_total += snap.requests_total; + aggregate.requests_error += snap.requests_error; + aggregate.requests_success += snap.requests_success; + aggregate.streams_started += snap.streams_started; + aggregate.streams_completed += snap.streams_completed; + aggregate.streams_failed += snap.streams_failed; + aggregate.streams_client_disconnected += snap.streams_client_disconnected; + } + let error_rate = aggregate.error_rate(); + let now_secs = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs(); + let since = now_secs.saturating_sub(60); + let rpm = admin::state::with_db(&snapshot_shared.db, move |conn_ref| { + admin::db::count_requests_since(conn_ref, since).unwrap_or(0) + }) + .await + .unwrap_or(0) as f64; + + let snapshot = admin::state::MetricsSnapshotData { + total_requests: aggregate.requests_total, + successful_requests: aggregate.requests_success, + failed_requests: aggregate.requests_error, + requests_per_minute: rpm, + p50_latency_ms: None, + p95_latency_ms: None, + error_rate, + streams_started: aggregate.streams_started, + streams_completed: aggregate.streams_completed, + streams_failed: aggregate.streams_failed, + streams_client_disconnected: aggregate.streams_client_disconnected, + }; + let _ = snapshot_shared + .events_tx + .send(admin::state::AdminEvent::MetricsSnapshot(snapshot)); + } + }); +} diff --git a/crates/proxy/src/server/chat_completions/backend_anthropic.rs b/crates/proxy/src/server/chat_completions/backend_anthropic.rs new file mode 100644 index 0000000..362e39d --- /dev/null +++ b/crates/proxy/src/server/chat_completions/backend_anthropic.rs @@ -0,0 +1,168 @@ +use crate::backend::anthropic_client::AnthropicClient; +use crate::openai_tool_policy::{ + backend_kind_for_policy, validate_anthropic_tool_request, OpenAiToolPolicyContext, +}; +use crate::server::routes::{ + inject_degradation_header, log_request, record_virtual_key_usage, set_backend_error_kind, + try_cache_response, RequestCtx, +}; +use crate::server::state::AppState; +use anyllm_translate::{ + anthropic, translate_anthropic_to_openai_response_with_context, AnthropicTranslationContext, + TranslationWarnings, +}; +use axum::{ + http::StatusCode, + response::{IntoResponse, Response}, +}; + +#[allow(clippy::too_many_arguments)] +pub(crate) async fn execute_anthropic_backend( + client: &AnthropicClient, + effective: &AppState, + state: &AppState, + anthropic_req: &anthropic::MessageCreateRequest, + raw_tools: &[serde_json::Value], + safe_headers: &[(String, String)], + tool_context: &AnthropicTranslationContext, + original_model: &str, + mapped_model: &str, + warnings: &mut TranslationWarnings, + deployment: &Option>, + backend_start: std::time::Instant, + vk_ctx: &Option, + ctx: &RequestCtx, + cache_control: &crate::cache::CacheControl, + store_cache_key: &Option, +) -> Response { + let mut upstream_req = anthropic_req.clone(); + upstream_req.model = mapped_model.to_string(); + upstream_req.stream = Some(false); + if let Err(err) = validate_anthropic_tool_request( + &upstream_req, + OpenAiToolPolicyContext { + backend_kind: backend_kind_for_policy(&effective.backend), + provider_id: effective.provider_id.as_deref(), + model: mapped_model, + provider_catalog: &effective.provider_catalog, + }, + ) { + return super::helpers::openai_error_response( + err.message(), + "invalid_request_error", + StatusCode::BAD_REQUEST, + ); + } + let body = + match super::extensions::serialize_anthropic_upstream_request(&upstream_req, raw_tools) { + Ok(body) => body, + Err(e) => { + return super::helpers::openai_error_response( + &format!("failed to serialize Anthropic request: {e}"), + "server_error", + StatusCode::INTERNAL_SERVER_ERROR, + ); + } + }; + let refs = super::helpers::header_refs(safe_headers); + + match client.forward(body, &refs, None).await { + Ok((resp_body, rate_limits)) => { + if let Some(ref d) = deployment { + d.record_finish(backend_start.elapsed().as_millis() as u64); + } + let anthropic_resp = + match serde_json::from_slice::(&resp_body) { + Ok(resp) => resp, + Err(e) => { + state.metrics.record_error(); + log_request( + &state.shared, + ctx.log_entry_with_attribution( + &state.backend_name, + Some(mapped_model.to_string()), + StatusCode::BAD_GATEWAY.as_u16(), + None, + false, + Some(format!("failed to parse Anthropic upstream response: {e}")), + vk_ctx, + None, + ), + ); + return super::helpers::openai_error_response( + "Upstream Anthropic response could not be parsed.", + "server_error", + StatusCode::BAD_GATEWAY, + ); + } + }; + state.metrics.record_success(); + let oai_response = translate_anthropic_to_openai_response_with_context( + &anthropic_resp, + original_model, + tool_context, + ); + let cost = record_virtual_key_usage( + &state.shared, + vk_ctx, + mapped_model, + anthropic_resp.usage.input_tokens as u64, + anthropic_resp.usage.output_tokens as u64, + ); + log_request( + &state.shared, + ctx.log_entry_with_attribution( + &state.backend_name, + Some(mapped_model.to_string()), + 200, + Some(( + anthropic_resp.usage.input_tokens as u64, + anthropic_resp.usage.output_tokens as u64, + )), + false, + None, + vk_ctx, + Some(cost), + ), + ); + + try_cache_response( + store_cache_key, + &state.cache, + cache_control.ttl_secs, + &oai_response, + original_model.to_string(), + ) + .await; + + let cache_hv = crate::server::routes::cache_header_value(!cache_control.lookup); + let mut response = (StatusCode::OK, axum::Json(oai_response)).into_response(); + rate_limits.inject_anthropic_response_headers(response.headers_mut()); + if state.expose_degradation_warnings { + inject_degradation_header(response.headers_mut(), warnings); + } + response.headers_mut().insert("x-anyllm-cache", cache_hv); + response + } + Err(e) => { + if let Some(ref d) = deployment { + d.record_finish(backend_start.elapsed().as_millis() as u64); + } + state.metrics.record_error(); + let backend_error = crate::backend::BackendError::from(e); + let mut entry = ctx.log_entry_with_attribution( + &state.backend_name, + Some(mapped_model.to_string()), + backend_error.status_code(), + None, + false, + Some(backend_error.to_string()), + vk_ctx, + None, + ); + set_backend_error_kind(&mut entry, &backend_error); + log_request(&state.shared, entry); + super::helpers::backend_error_to_openai_response(backend_error) + } + } +} diff --git a/crates/proxy/src/server/chat_completions/backend_openai.rs b/crates/proxy/src/server/chat_completions/backend_openai.rs new file mode 100644 index 0000000..c70ab08 --- /dev/null +++ b/crates/proxy/src/server/chat_completions/backend_openai.rs @@ -0,0 +1,204 @@ +use crate::backend::openai_client::OpenAIClient; +use crate::openai_tool_policy::{ + backend_kind_for_policy, prepare_openai_tool_request, OpenAiToolPolicyContext, +}; +use crate::server::routes::{ + inject_degradation_header, inject_gemini_thinking, inject_glm_thinking, log_request, + record_virtual_key_usage, set_backend_error_kind, try_cache_response, RequestCtx, +}; +use crate::server::state::AppState; +use anyllm_translate::{ + anthropic, mapping, openai, translate_anthropic_to_openai_response, TranslationWarnings, +}; +use axum::{ + http::StatusCode, + response::{IntoResponse, Json, Response}, +}; + +#[allow(clippy::too_many_arguments)] +pub(crate) async fn execute_openai_backend( + client: &OpenAIClient, + effective: &AppState, + state: &AppState, + anthropic_req: &anthropic::MessageCreateRequest, + _body: &openai::ChatCompletionRequest, + original_model: &str, + mapped_model: &str, + warnings: &mut TranslationWarnings, + deployment: &Option>, + backend_start: std::time::Instant, + vk_ctx: &Option, + ctx: &RequestCtx, + cache_control: &crate::cache::CacheControl, + store_cache_key: &Option, +) -> Response { + let mut openai_req = mapping::message_map::anthropic_to_openai_request(anthropic_req); + inject_gemini_thinking(anthropic_req, &effective.backend, &mut openai_req); + inject_glm_thinking(anthropic_req, &effective.backend, &mut openai_req); + if effective.omit_stream_options { + openai_req.stream_options = None; + } + openai_req.model = mapped_model.to_string(); + + // Opt-in RTK tool-output compression (OpenAI-in translate path). + effective.apply_rtk_to_openai(&mut openai_req, mapped_model); + + if let Err(err) = prepare_openai_tool_request( + &mut openai_req, + OpenAiToolPolicyContext { + backend_kind: backend_kind_for_policy(&effective.backend), + provider_id: effective.provider_id.as_deref(), + model: mapped_model, + provider_catalog: &effective.provider_catalog, + }, + warnings, + ) { + return super::helpers::openai_error_response( + err.message(), + "invalid_request_error", + StatusCode::BAD_REQUEST, + ); + } + let mapped_model = openai_req.model.clone(); + + match client.chat_completion(&openai_req).await { + Ok((openai_resp, _status, rate_limits)) => { + if let Some(ref d) = deployment { + d.record_finish(backend_start.elapsed().as_millis() as u64); + } + state.metrics.record_success(); + // Translate Anthropic response back to OpenAI format + let anthropic_resp = + mapping::message_map::openai_to_anthropic_response(&openai_resp, original_model); + + // Tool execution: bounded loop with termination guards. + let anthropic_resp = if let Some(ref engine) = effective.tool_engine { + let client_for_tools = client.clone(); + let model_for_tools = mapped_model.clone(); + let orig_model_for_tools = original_model.to_string(); + let redact_follow_up = effective.redact_secrets(); + let backend_kind_for_tools = backend_kind_for_policy(&effective.backend); + let provider_id_for_tools = effective.provider_id.clone(); + let provider_catalog_for_tools = effective.provider_catalog.clone(); + let server_advertised_tool_names = std::collections::HashSet::new(); + let guardrails_for_tools = effective.effective_tool_guardrails(engine); + let (resp, _trace) = crate::tools::execution::maybe_execute_tools( + engine, + anthropic_req, + &server_advertised_tool_names, + anthropic_resp, + &guardrails_for_tools, + |follow_up_req| { + let c = client_for_tools.clone(); + let m = model_for_tools.clone(); + let policy_model = m.clone(); + let om = orig_model_for_tools.clone(); + let backend_kind = backend_kind_for_tools.clone(); + let provider_id = provider_id_for_tools.clone(); + let provider_catalog = provider_catalog_for_tools.clone(); + async move { + let follow_up_req = + match crate::server::secret_redaction::redact_json_value( + redact_follow_up, + follow_up_req, + ) + .await + { + Ok(req) => req, + Err(err) => return Err(err.safe_message().to_string()), + }; + let mut oai_req = + mapping::message_map::anthropic_to_openai_request(&follow_up_req); + oai_req.model = m; + let mut follow_up_warnings = TranslationWarnings::default(); + prepare_openai_tool_request( + &mut oai_req, + OpenAiToolPolicyContext { + backend_kind, + provider_id: provider_id.as_deref(), + model: &policy_model, + provider_catalog: &provider_catalog, + }, + &mut follow_up_warnings, + ) + .map_err(|err| err.to_string())?; + match c.chat_completion(&oai_req).await { + Ok((resp, _, _)) => Ok( + mapping::message_map::openai_to_anthropic_response(&resp, &om), + ), + Err(e) => Err(format!("{e}")), + } + } + }, + ) + .await; + resp + } else { + anthropic_resp + }; + + let oai_response = + translate_anthropic_to_openai_response(&anthropic_resp, original_model); + let cost = record_virtual_key_usage( + &state.shared, + vk_ctx, + &mapped_model, + anthropic_resp.usage.input_tokens as u64, + anthropic_resp.usage.output_tokens as u64, + ); + log_request( + &state.shared, + ctx.log_entry_with_attribution( + &state.backend_name, + Some(mapped_model), + 200, + Some(( + anthropic_resp.usage.input_tokens as u64, + anthropic_resp.usage.output_tokens as u64, + )), + false, + None, + vk_ctx, + Some(cost), + ), + ); + try_cache_response( + store_cache_key, + &state.cache, + cache_control.ttl_secs, + &oai_response, + original_model.to_string(), + ) + .await; + + let cache_hv = crate::server::routes::cache_header_value(!cache_control.lookup); + let mut response = (StatusCode::OK, Json(oai_response)).into_response(); + rate_limits.inject_anthropic_response_headers(response.headers_mut()); + if state.expose_degradation_warnings { + inject_degradation_header(response.headers_mut(), warnings); + } + response.headers_mut().insert("x-anyllm-cache", cache_hv); + response + } + Err(e) => { + if let Some(ref d) = deployment { + d.record_finish(backend_start.elapsed().as_millis() as u64); + } + state.metrics.record_error(); + let backend_error = crate::backend::BackendError::from(e); + let mut entry = ctx.log_entry_with_attribution( + &state.backend_name, + Some(mapped_model), + backend_error.status_code(), + None, + false, + Some(backend_error.to_string()), + vk_ctx, + None, + ); + set_backend_error_kind(&mut entry, &backend_error); + log_request(&state.shared, entry); + super::helpers::backend_error_to_openai_response(backend_error) + } + } +} diff --git a/crates/proxy/src/server/chat_completions/backend_responses.rs b/crates/proxy/src/server/chat_completions/backend_responses.rs new file mode 100644 index 0000000..0d66526 --- /dev/null +++ b/crates/proxy/src/server/chat_completions/backend_responses.rs @@ -0,0 +1,110 @@ +use crate::backend::openai_client::OpenAIClient; +use crate::server::routes::{ + inject_degradation_header, log_request, record_virtual_key_usage, set_backend_error_kind, + try_cache_response, RequestCtx, +}; +use crate::server::state::AppState; +use anyllm_translate::{ + anthropic, mapping, translate_anthropic_to_openai_response, TranslationWarnings, +}; +use axum::{ + http::StatusCode, + response::{IntoResponse, Json, Response}, +}; + +#[allow(clippy::too_many_arguments)] +pub(crate) async fn execute_responses_backend( + client: &OpenAIClient, + state: &AppState, + anthropic_req: &anthropic::MessageCreateRequest, + original_model: &str, + mapped_model: &str, + warnings: &mut TranslationWarnings, + deployment: &Option>, + backend_start: std::time::Instant, + vk_ctx: &Option, + ctx: &RequestCtx, + cache_control: &crate::cache::CacheControl, + store_cache_key: &Option, +) -> Response { + let mut responses_req = + mapping::responses_message_map::anthropic_to_responses_request(anthropic_req); + responses_req.model = mapped_model.to_string(); + let mapped_model = responses_req.model.clone(); + + match client.responses(&responses_req).await { + Ok((resp, _status, rate_limits)) => { + if let Some(ref d) = deployment { + d.record_finish(backend_start.elapsed().as_millis() as u64); + } + state.metrics.record_success(); + let anthropic_resp = mapping::responses_message_map::responses_to_anthropic_response( + &resp, + original_model, + ); + let oai_response = + translate_anthropic_to_openai_response(&anthropic_resp, original_model); + let cost = record_virtual_key_usage( + &state.shared, + vk_ctx, + &mapped_model, + anthropic_resp.usage.input_tokens as u64, + anthropic_resp.usage.output_tokens as u64, + ); + log_request( + &state.shared, + ctx.log_entry_with_attribution( + &state.backend_name, + Some(mapped_model), + 200, + Some(( + anthropic_resp.usage.input_tokens as u64, + anthropic_resp.usage.output_tokens as u64, + )), + false, + None, + vk_ctx, + Some(cost), + ), + ); + + try_cache_response( + store_cache_key, + &state.cache, + cache_control.ttl_secs, + &oai_response, + original_model.to_string(), + ) + .await; + + let cache_hv = crate::server::routes::cache_header_value(!cache_control.lookup); + let mut response = (StatusCode::OK, Json(oai_response)).into_response(); + rate_limits.inject_anthropic_response_headers(response.headers_mut()); + if state.expose_degradation_warnings { + inject_degradation_header(response.headers_mut(), warnings); + } + response.headers_mut().insert("x-anyllm-cache", cache_hv); + response + } + Err(e) => { + if let Some(ref d) = deployment { + d.record_finish(backend_start.elapsed().as_millis() as u64); + } + state.metrics.record_error(); + let backend_error = crate::backend::BackendError::from(e); + let mut entry = ctx.log_entry_with_attribution( + &state.backend_name, + Some(mapped_model), + backend_error.status_code(), + None, + false, + Some(backend_error.to_string()), + vk_ctx, + None, + ); + set_backend_error_kind(&mut entry, &backend_error); + log_request(&state.shared, entry); + super::helpers::backend_error_to_openai_response(backend_error) + } + } +} diff --git a/crates/proxy/src/server/passthrough/handlers/errors.rs b/crates/proxy/src/server/passthrough/handlers/errors.rs new file mode 100644 index 0000000..b16e696 --- /dev/null +++ b/crates/proxy/src/server/passthrough/handlers/errors.rs @@ -0,0 +1,42 @@ +use crate::backend::anthropic_client::AnthropicClientError; +use anyllm_translate::{anthropic, mapping}; +use axum::{ + http::StatusCode, + response::{IntoResponse, Json, Response}, +}; + +pub(crate) fn passthrough_error_to_response(error: AnthropicClientError) -> Response { + match error { + AnthropicClientError::ApiError { status, body } => { + let http_status = + StatusCode::from_u16(status).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR); + (http_status, [("content-type", "application/json")], body).into_response() + } + AnthropicClientError::Transport(msg) => { + tracing::error!("Anthropic passthrough transport error: {msg}"); + let err = mapping::errors_map::create_anthropic_error( + anthropic::ErrorType::ApiError, + "An internal error occurred while communicating with the upstream service." + .to_string(), + None, + ); + (StatusCode::BAD_GATEWAY, Json(err)).into_response() + } + } +} + +pub(crate) fn passthrough_error_status(error: &AnthropicClientError) -> u16 { + match error { + AnthropicClientError::ApiError { status, .. } => *status, + AnthropicClientError::Transport(_) => StatusCode::BAD_GATEWAY.as_u16(), + } +} + +pub(crate) fn virtual_key_accounting_parse_error() -> Response { + let err = mapping::errors_map::create_anthropic_error( + anthropic::ErrorType::ApiError, + "Upstream response could not be accounted for this virtual API key.".to_string(), + None, + ); + (StatusCode::BAD_GATEWAY, Json(err)).into_response() +} diff --git a/crates/proxy/src/server/passthrough/handlers/generic.rs b/crates/proxy/src/server/passthrough/handlers/generic.rs new file mode 100644 index 0000000..8bd75d4 --- /dev/null +++ b/crates/proxy/src/server/passthrough/handlers/generic.rs @@ -0,0 +1,124 @@ +use crate::backend::BackendClient; +use crate::server::middleware::ClientAuthPath; +use crate::server::state::AppState; +use anyllm_translate::{anthropic, mapping}; +use axum::{ + body::Bytes, + extract::{OriginalUri, State}, + http::StatusCode, + response::{IntoResponse, Json, Response}, +}; + +use super::super::auth::resolve_client_auth_override; +use super::errors::passthrough_error_to_response; + +/// Generic catch-all passthrough for any /v1/* path in Anthropic mode. +/// Forwards batch, file CRUD, count_tokens, and other Anthropic-native endpoints +/// directly to the upstream Anthropic API. Registered after /v1/messages so that +/// route retains its dedicated streaming/model-peek logic. +#[allow(clippy::too_many_arguments)] +pub(crate) async fn anthropic_generic_passthrough( + State(state): State, + vk_ctx: Option>, + auth_path: Option>, + claims: Option>, + OriginalUri(uri): OriginalUri, + method: axum::http::Method, + headers: axum::http::HeaderMap, + body: Bytes, +) -> Response { + state.metrics.record_request(); + + // Virtual keys must use policy-aware handlers only. + if vk_ctx.is_some() { + let err = mapping::errors_map::create_anthropic_error( + anthropic::ErrorType::PermissionError, + "This endpoint is not available for virtual API keys.".to_string(), + None, + ); + return (StatusCode::FORBIDDEN, Json(err)).into_response(); + } + let vk_ctx = vk_ctx.map(|axum::Extension(c)| c); + let auth_path = auth_path.map(|axum::Extension(p)| p); + let claims = claims.map(|axum::Extension(c)| c); + let auth_override_ref = resolve_client_auth_override( + state.forward_client_auth_enabled(), + auth_path, + &vk_ctx, + &claims, + &headers, + ); + + let client = match &state.backend { + BackendClient::Anthropic(c) => c, + _ => { + let err = mapping::errors_map::create_anthropic_error( + anthropic::ErrorType::ApiError, + "Backend is not configured as anthropic passthrough".to_string(), + None, + ); + return (StatusCode::INTERNAL_SERVER_ERROR, Json(err)).into_response(); + } + }; + + let mut full_path = uri.path().to_string(); + if let Some(q) = uri.query() { + full_path.push('?'); + full_path.push_str(q); + } + + let session_id = headers + .get("x-claude-code-session-id") + .and_then(|v| v.to_str().ok()) + .map(str::to_string); + let beta = headers + .get("anthropic-beta") + .and_then(|v| v.to_str().ok()) + .map(str::to_string); + let mut extra: Vec<(&str, &str)> = Vec::new(); + if let Some(ref v) = session_id { + extra.push(("x-claude-code-session-id", v)); + } + if let Some(ref v) = beta { + extra.push(("anthropic-beta", v)); + } + + let body = + match crate::server::secret_redaction::redact_body(state.redact_secrets(), &headers, body) + .await + { + Ok(body) => body, + Err(err) => return crate::server::secret_redaction::error_response(err), + }; + + match client + .forward_generic(method, &full_path, body, &extra, auth_override_ref) + .await + { + Ok(response) => { + let status = StatusCode::from_u16(response.status().as_u16()).unwrap_or(StatusCode::OK); + if status.is_success() { + state.metrics.record_success(); + } else { + state.metrics.record_error(); + } + let upstream_ct = response + .headers() + .get("content-type") + .and_then(|v| v.to_str().ok()) + .unwrap_or("application/json") + .to_string(); + let stream = response.bytes_stream(); + let axum_body = axum::body::Body::from_stream(stream); + let mut resp = (status, axum_body).into_response(); + if let Ok(hv) = axum::http::HeaderValue::from_str(&upstream_ct) { + resp.headers_mut().insert("content-type", hv); + } + resp + } + Err(e) => { + state.metrics.record_error(); + passthrough_error_to_response(e) + } + } +} diff --git a/crates/proxy/src/server/passthrough/handlers.rs b/crates/proxy/src/server/passthrough/handlers/messages.rs similarity index 59% rename from crates/proxy/src/server/passthrough/handlers.rs rename to crates/proxy/src/server/passthrough/handlers/messages.rs index a9f425f..f3ddf38 100644 --- a/crates/proxy/src/server/passthrough/handlers.rs +++ b/crates/proxy/src/server/passthrough/handlers/messages.rs @@ -9,7 +9,7 @@ use crate::server::streaming::{observe_anthropic_sse_frames, AnthropicStreamUsag use anyllm_translate::{anthropic, mapping}; use axum::{ body::Bytes, - extract::{OriginalUri, State}, + extract::State, http::StatusCode, response::{IntoResponse, Json, Response}, }; @@ -18,7 +18,10 @@ use futures::StreamExt; use tokio::sync::mpsc; use tokio_stream::wrappers::ReceiverStream; -use super::auth::resolve_client_auth_override; +use super::super::auth::resolve_client_auth_override; +use super::errors::{ + passthrough_error_status, passthrough_error_to_response, virtual_key_accounting_parse_error, +}; /// Forward an Anthropic-format request byte-for-byte to the upstream Anthropic API. /// No translation is performed. Only active when `BACKEND=anthropic`. @@ -37,12 +40,6 @@ pub(crate) async fn anthropic_passthrough( let claims = claims.map(|axum::Extension(c)| c); state.metrics.record_request(); - // Verbatim client-credential override for ANTHROPIC_FORWARD_CLIENT_AUTH: - // computed once per request, reused across the streaming/non-streaming - // branches below. Borrowed straight from `headers` (which outlives both - // branches and is never mutated), not owned -- `forward`/`forward_stream` - // consume it synchronously before the branch that does - // `tokio::spawn`, so there's no need to outlive the spawned task. let auth_override_ref = resolve_client_auth_override( state.forward_client_auth_enabled(), auth_path, @@ -51,16 +48,6 @@ pub(crate) async fn anthropic_passthrough( &headers, ); - // Scopes every thinking-repair store lookup/commit to this backend and - // virtual key: `state.thinking_repair` is one store shared across every - // Anthropic-mode backend (see server/routes.rs), so without this a - // colliding message id / thinking signature / tool_use id from a - // different backend or tenant could resolve to this request's repair. - // NUL, not `:`, joins the two parts: `state.backend_name` is an - // operator-configured string whose validated charset (`is_safe_model_name`) - // allows `:` but not NUL, so a backend literally named e.g. "anthropic:5" - // can no longer produce the same namespace as backend "anthropic" + key - // id 5 -- `:` let those collide onto one shared cache-record namespace. let thinking_repair_namespace = match &vk_ctx { Some(ctx) => format!("{}\u{0}{}", state.backend_name, ctx.key_id), None => state.backend_name.clone(), @@ -78,9 +65,6 @@ pub(crate) async fn anthropic_passthrough( } }; - // Collect Anthropic-specific client headers to forward upstream. - // anthropic-beta enables beta features; must reach upstream to take effect. - // x-claude-code-session-id allows upstream and intermediary proxies to correlate sessions. let extra_headers: Vec<(&str, &str)> = ["x-claude-code-session-id", "anthropic-beta"] .iter() .filter_map(|&name| { @@ -91,9 +75,6 @@ pub(crate) async fn anthropic_passthrough( }) .collect(); - // Peek at just the `stream` and `model` fields instead of parsing the full body. - // Full deserialization would be wasteful for image-heavy requests - // (up to 32MB) when we only need one boolean to choose the handler. #[derive(serde::Deserialize)] struct BodyPeek { #[serde(default)] @@ -115,7 +96,6 @@ pub(crate) async fn anthropic_passthrough( model_requested: peek.model.clone().unwrap_or_else(|| "unknown".to_string()), }; - // Enforce model allowlist for virtual keys. if let Some(ref ctx) = vk_ctx { match &peek.model { Some(m) => { @@ -129,8 +109,6 @@ pub(crate) async fn anthropic_passthrough( } } None => { - // If a model allowlist is configured, we cannot verify the request - // is permitted without knowing the model. Reject rather than bypass. if ctx.allowed_models.is_some() { let err = mapping::errors_map::create_anthropic_error( anthropic::ErrorType::InvalidRequestError, @@ -144,18 +122,8 @@ pub(crate) async fn anthropic_passthrough( } } - // pxpipe compression is decided here but APPLIED after secret redaction - // below: imaging bakes the system/tool_result text into PNG pixels, which the - // text-based redactor cannot see, so redaction must run on the raw text first. let mut pxpipe_apply: Option<(std::sync::Arc, String)> = None; - // RTK tool-output compression, likewise deferred to after redaction and run - // BEFORE pxpipe so imaging sees the already-filtered text. let mut rtk_apply: Option<(std::sync::Arc, String)> = None; - // FFEC prompt compression (history + frontier cache_control breakpoint). Runs - // on the raw Value bytes (never round-tripping MessageCreateRequest, which has - // no cache_control field), so the breakpoint survives to the Anthropic API. The - // `.filter` gates the (up to 32MB) parse: only capture when the resolved mode is - // not Off. Does not need the parsed model, so it sits outside the parse block. let optimizer_apply = state .effective_optimizer() .filter(|e| e.mode() != anyllm_optimize_core::Mode::Off); @@ -177,10 +145,6 @@ pub(crate) async fn anthropic_passthrough( return (StatusCode::BAD_REQUEST, Json(err)).into_response(); } - // Repair the last assistant message's thinking blocks against - // recorded ground truth (opt-in, see crate::thinking_repair). Only - // rewrites `body` when something actually changed; on a byte-exact - // replay (the common case) this is a no-op past the store lookups. if let Some(store) = state.active_thinking_repair() { if let Some(what) = crate::thinking_repair::repair_request( &store, @@ -189,21 +153,12 @@ pub(crate) async fn anthropic_passthrough( ) .await { - // Patch only the repaired message's `content` into the - // ORIGINAL raw JSON instead of re-serializing the whole - // typed request: ContentBlock/Tool have no cache_control or - // flatten catch-all, so a full-struct round-trip would - // silently drop cache_control breakpoints and any block/tool - // type this crate doesn't model yet, on every OTHER message - // too. match crate::thinking_repair::patch_repaired_body(&body, &parsed_req) { Ok(bytes) => { tracing::info!(repair = %what, "anthropic thinking-block repair applied"); body = bytes; } Err(e) => { - // Fail open: forward the original (unrepaired) bytes - // rather than drop the request. tracing::warn!( error = %e, "failed to patch repaired anthropic request; forwarding original body" @@ -213,11 +168,6 @@ pub(crate) async fn anthropic_passthrough( } } - // Opt-in text-to-image context compression (pxpipe). Decided here (needs - // the parsed model for scope/vision gating) but deferred to after - // redaction — see the `pxpipe_apply` capture note above. The transform - // works on the raw `body` bytes (Value-level) so it can RELOCATE the - // caller's cache_control anchor onto the image; it fails open on any error. if let Some(engine) = state.pxpipe_engine_for(&parsed_req.model) { pxpipe_apply = Some((engine, parsed_req.model.clone())); } @@ -237,19 +187,10 @@ pub(crate) async fn anthropic_passthrough( Err(err) => return crate::server::secret_redaction::error_response(err), }; - // FFEC prompt compression runs FIRST: it compresses conversation history text - // and places the frontier cache_control breakpoint. It MUST precede pxpipe, - // which images the static message slab into a PNG — after imaging there is no - // history text left to compress and no place for a breakpoint. Combining the - // optimizer's breakpoint with pxpipe's cache_control-anchor relocation is - // untested (both are opt-in); we don't try to reconcile them here. Fails open. if let Some(engine) = optimizer_apply { body = engine.optimize_anthropic_bytes(body, "messages", &state.metrics); } - // Now that secrets are redacted, compress tool output (RTK) then image the - // static slab / large live regions (pxpipe). RTK first so pxpipe images the - // already-filtered text; both fail open. if let Some((engine, model)) = rtk_apply { body = engine.compress_anthropic(body, &model, &state.metrics); } @@ -268,8 +209,6 @@ pub(crate) async fn anthropic_passthrough( let log_shared = state.shared.clone(); let log_backend_name = state.backend_name.clone(); let cost_model = peek.model.clone().unwrap_or_else(|| "unknown".to_string()); - // Captured once, before the stream starts: a toggle mid-stream - // must not half-record ground truth. let thinking_repair = state.active_thinking_repair(); let thinking_repair_namespace = thinking_repair_namespace.clone(); @@ -390,10 +329,6 @@ pub(crate) async fn anthropic_passthrough( .await { Ok((resp_body, rate_limits)) => { - // Parsed once and shared between thinking-repair recording - // and virtual-key accounting below (previously each parsed - // the same bytes independently, and recording additionally - // cloned the whole content Vec instead of moving it out). let mut parsed_resp = serde_json::from_slice::(&resp_body); @@ -491,167 +426,3 @@ pub(crate) async fn anthropic_passthrough( } } } - -/// Generic catch-all passthrough for any /v1/* path in Anthropic mode. -/// Forwards batch, file CRUD, count_tokens, and other Anthropic-native endpoints -/// directly to the upstream Anthropic API. Registered after /v1/messages so that -/// route retains its dedicated streaming/model-peek logic. -#[allow(clippy::too_many_arguments)] -pub(crate) async fn anthropic_generic_passthrough( - State(state): State, - vk_ctx: Option>, - auth_path: Option>, - claims: Option>, - OriginalUri(uri): OriginalUri, - method: axum::http::Method, - headers: axum::http::HeaderMap, - body: Bytes, -) -> Response { - state.metrics.record_request(); - - // Virtual keys must use policy-aware handlers only. The generic passthrough - // can reach Anthropic-native endpoints that lack per-key authorization, - // resource ownership checks, and usage accounting. - if vk_ctx.is_some() { - let err = mapping::errors_map::create_anthropic_error( - anthropic::ErrorType::PermissionError, - "This endpoint is not available for virtual API keys.".to_string(), - None, - ); - return (StatusCode::FORBIDDEN, Json(err)).into_response(); - } - let vk_ctx = vk_ctx.map(|axum::Extension(c)| c); - let auth_path = auth_path.map(|axum::Extension(p)| p); - let claims = claims.map(|axum::Extension(c)| c); - // vk_ctx is already rejected above (so resolve_client_auth_override's - // vk_ctx.is_none() check is always true here), but an OIDC-authenticated - // non-virtual-key request can still reach here, so this must also be - // gated on claims/auth_path (never forward a JWT upstream as if it were - // an Anthropic credential). - let auth_override_ref = resolve_client_auth_override( - state.forward_client_auth_enabled(), - auth_path, - &vk_ctx, - &claims, - &headers, - ); - - let client = match &state.backend { - BackendClient::Anthropic(c) => c, - _ => { - let err = mapping::errors_map::create_anthropic_error( - anthropic::ErrorType::ApiError, - "Backend is not configured as anthropic passthrough".to_string(), - None, - ); - return (StatusCode::INTERNAL_SERVER_ERROR, Json(err)).into_response(); - } - }; - - // Build full path with query string preserved. - let mut full_path = uri.path().to_string(); - if let Some(q) = uri.query() { - full_path.push('?'); - full_path.push_str(q); - } - - // Collect owned Strings before building the &str slice (lifetime requirement). - let session_id = headers - .get("x-claude-code-session-id") - .and_then(|v| v.to_str().ok()) - .map(str::to_string); - let beta = headers - .get("anthropic-beta") - .and_then(|v| v.to_str().ok()) - .map(str::to_string); - let mut extra: Vec<(&str, &str)> = Vec::new(); - if let Some(ref v) = session_id { - extra.push(("x-claude-code-session-id", v)); - } - if let Some(ref v) = beta { - extra.push(("anthropic-beta", v)); - } - - let body = - match crate::server::secret_redaction::redact_body(state.redact_secrets(), &headers, body) - .await - { - Ok(body) => body, - Err(err) => return crate::server::secret_redaction::error_response(err), - }; - - match client - .forward_generic(method, &full_path, body, &extra, auth_override_ref) - .await - { - Ok(response) => { - let status = StatusCode::from_u16(response.status().as_u16()).unwrap_or(StatusCode::OK); - if status.is_success() { - state.metrics.record_success(); - } else { - state.metrics.record_error(); - } - // Preserve upstream content-type (batches return application/x-jsonl, etc.) - let upstream_ct = response - .headers() - .get("content-type") - .and_then(|v| v.to_str().ok()) - .unwrap_or("application/json") - .to_string(); - let stream = response.bytes_stream(); - let axum_body = axum::body::Body::from_stream(stream); - let mut resp = (status, axum_body).into_response(); - if let Ok(hv) = axum::http::HeaderValue::from_str(&upstream_ct) { - resp.headers_mut().insert("content-type", hv); - } - resp - } - Err(e) => { - state.metrics.record_error(); - passthrough_error_to_response(e) - } - } -} - -/// Convert an AnthropicClientError into a Response. -/// For API errors, return the upstream error body directly (it's already Anthropic format). -fn passthrough_error_to_response( - error: crate::backend::anthropic_client::AnthropicClientError, -) -> Response { - use crate::backend::anthropic_client::AnthropicClientError; - match error { - AnthropicClientError::ApiError { status, body } => { - let http_status = - StatusCode::from_u16(status).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR); - (http_status, [("content-type", "application/json")], body).into_response() - } - AnthropicClientError::Transport(msg) => { - tracing::error!("Anthropic passthrough transport error: {msg}"); - let err = mapping::errors_map::create_anthropic_error( - anthropic::ErrorType::ApiError, - "An internal error occurred while communicating with the upstream service." - .to_string(), - None, - ); - (StatusCode::BAD_GATEWAY, Json(err)).into_response() - } - } -} - -fn passthrough_error_status(error: &crate::backend::anthropic_client::AnthropicClientError) -> u16 { - match error { - crate::backend::anthropic_client::AnthropicClientError::ApiError { status, .. } => *status, - crate::backend::anthropic_client::AnthropicClientError::Transport(_) => { - StatusCode::BAD_GATEWAY.as_u16() - } - } -} - -fn virtual_key_accounting_parse_error() -> Response { - let err = mapping::errors_map::create_anthropic_error( - anthropic::ErrorType::ApiError, - "Upstream response could not be accounted for this virtual API key.".to_string(), - None, - ); - (StatusCode::BAD_GATEWAY, Json(err)).into_response() -} diff --git a/crates/proxy/src/server/passthrough/handlers/mod.rs b/crates/proxy/src/server/passthrough/handlers/mod.rs new file mode 100644 index 0000000..de02487 --- /dev/null +++ b/crates/proxy/src/server/passthrough/handlers/mod.rs @@ -0,0 +1,6 @@ +pub(crate) mod errors; +pub(crate) mod generic; +pub(crate) mod messages; + +pub(crate) use generic::anthropic_generic_passthrough; +pub(crate) use messages::anthropic_passthrough; diff --git a/crates/proxy/src/server/streaming.rs b/crates/proxy/src/server/streaming/handler.rs similarity index 57% rename from crates/proxy/src/server/streaming.rs rename to crates/proxy/src/server/streaming/handler.rs index b825f8f..b8e8058 100644 --- a/crates/proxy/src/server/streaming.rs +++ b/crates/proxy/src/server/streaming/handler.rs @@ -1,261 +1,24 @@ -// SSE streaming infrastructure and the messages_stream handler. - -use crate::backend::{find_double_newline, BackendClient, RateLimitHeaders, SseFrameBuffer}; -use crate::metrics::Metrics; +use crate::backend::{BackendClient, RateLimitHeaders}; use crate::openai_tool_policy::{ backend_kind_for_policy, parse_openai_chat_completion_chunk, prepare_openai_tool_request, tool_policy_error_to_backend_error, OpenAiToolPolicyContext, }; +use crate::server::routes::{log_request, set_backend_error_kind, RequestCtx}; +use crate::server::state::AppState; use anyllm_translate::{anthropic, mapping, TranslationWarnings}; use axum::response::sse::{Event, KeepAlive, Sse}; -use bytes::BytesMut; use futures::stream::Stream; -use std::sync::Arc; -use std::time::Instant; use tokio::sync::mpsc; use tokio_stream::wrappers::ReceiverStream; -use super::routes::{log_request, set_backend_error_kind, RequestCtx}; -use super::state::AppState; +use super::helpers::{read_sse_frames, send_events, StreamDeploymentAccounting, StreamOutcome}; -pub(crate) struct StreamDeploymentAccounting { - deployment: Option>, - start: Option, -} - -impl StreamDeploymentAccounting { - pub(crate) fn start(deployment: Option>) -> Self { - if let Some(deployment) = &deployment { - deployment.record_start(); - Self { - deployment: Some(deployment.clone()), - start: Some(Instant::now()), - } - } else { - Self { - deployment: None, - start: None, - } - } - } - - pub(crate) fn finish(&mut self) { - if let (Some(deployment), Some(start)) = (&self.deployment, self.start.take()) { - deployment.record_finish(start.elapsed().as_millis() as u64); - } - } -} - -impl Drop for StreamDeploymentAccounting { - fn drop(&mut self) { - self.finish(); - } -} - -/// Send translated stream events over the SSE channel. Returns false if client disconnected. -pub(super) async fn send_events( - tx: &mpsc::Sender>, - events: &[anthropic::StreamEvent], -) -> bool { - for ev in events { - match super::sse::stream_event_to_sse(ev) { - Ok(sse) => { - if tx.send(Ok(sse)).await.is_err() { - return false; - } - } - Err(e) => { - tracing::warn!("failed to serialize stream event: {e}"); - } - } - } - true -} - -/// Why the SSE stream ended. -#[derive(Debug, Clone, Copy)] -pub(crate) enum StreamOutcome { - /// Backend stream completed normally. - Completed, - /// Downstream client disconnected before the stream finished. - ClientDisconnected, - /// Backend stream failed (error already recorded in metrics). - UpstreamError, -} - -impl StreamOutcome { - /// Record metrics and return (HTTP status, error message) for logging. - pub(crate) fn record(&self, metrics: &Metrics) -> (u16, Option) { - match self { - Self::Completed => { - metrics.record_success(); - metrics.record_stream_completed(); - (200, None) - } - Self::ClientDisconnected => { - metrics.record_stream_client_disconnected(); - (499, Some("client disconnected".into())) - } - Self::UpstreamError => { - metrics.record_stream_failed(); - (502, Some("stream interrupted".into())) - } - } - } -} - -#[cfg(test)] -mod tests; - -#[derive(Debug, Default, Clone, Copy)] -pub(crate) struct AnthropicStreamUsage { - input_tokens: Option, - output_tokens: Option, -} - -impl AnthropicStreamUsage { - pub(crate) fn observe_data(&mut self, data: &str) { - if data == "[DONE]" { - return; - } - let Ok(event) = serde_json::from_str::(data) else { - return; - }; - match event { - anthropic::StreamEvent::MessageStart { message } => { - self.input_tokens = Some(message.usage.input_tokens as u64); - } - anthropic::StreamEvent::MessageDelta { - usage: Some(usage), .. - } => { - self.output_tokens = Some(usage.output_tokens as u64); - } - _ => {} - } - } - - pub(crate) fn tokens(&self) -> Option<(u64, u64)> { - match (self.input_tokens, self.output_tokens) { - (Some(input), Some(output)) => Some((input, output)), - (Some(input), None) => Some((input, 0)), - (None, Some(output)) => Some((0, output)), - (None, None) => None, - } - } -} - -/// Parse buffered SSE frames, updating token usage and (when `recorder` is -/// `Some`) accumulating content blocks for the thinking-block repair store. -/// Completed messages (from `message_stop`) are pushed onto `ready` for the -/// caller to commit — accumulation here is synchronous, but committing to -/// the store is async, so it can't happen inline in this loop. -pub(crate) fn observe_anthropic_sse_frames( - buffer: &mut BytesMut, - search_from: &mut usize, - usage: &mut AnthropicStreamUsage, - mut recorder: Option<&mut crate::thinking_repair::ThinkingRecorder>, - ready: &mut Vec<(String, Vec)>, -) { - while let Some((pos, delim_len)) = find_double_newline(buffer, *search_from) { - if let Ok(frame_str) = std::str::from_utf8(&buffer[..pos]) { - for line in frame_str.lines() { - let line = line.trim(); - if let Some(json_str) = line.strip_prefix("data: ") { - usage.observe_data(json_str); - if let Some(rec) = recorder.as_deref_mut() { - if let Some(done) = rec.observe_json(json_str) { - ready.push(done); - } - } - } - } - } - let _ = buffer.split_to(pos + delim_len); - *search_from = 0; - } - *search_from = buffer.len().saturating_sub(3); -} - -/// Read SSE bytes from a response, parse frames, and call `on_data` for each data line. -pub(super) async fn read_sse_frames( - response: reqwest::Response, - tx: &mpsc::Sender>, - metrics: &Metrics, - mut on_data: F, -) -> StreamOutcome -where - F: FnMut(&str) -> Option>, -{ - use futures::StreamExt; - let mut stream = response.bytes_stream(); - // Buffer bytes (not String) because TCP chunks may split mid-UTF-8 character. - // String::from_utf8_lossy would permanently replace partial trailing bytes - // with U+FFFD, corrupting the JSON payload. - let mut buffer = SseFrameBuffer::new(); - // Reuse a single events buffer across all frames to avoid per-frame allocation - let mut frame_events: Vec = Vec::new(); - - while let Some(chunk_result) = stream.next().await { - let bytes = match chunk_result { - Ok(b) => b, - Err(e) => { - tracing::error!("stream read error: {e}"); - metrics.record_error(); - return StreamOutcome::UpstreamError; - } - }; - let frames = match buffer.push(&bytes) { - Ok(frames) => frames, - Err(e) => { - tracing::error!(error = %e, "SSE buffer exceeded maximum size, aborting stream"); - metrics.record_error(); - return StreamOutcome::UpstreamError; - } - }; - - for frame in frames { - frame_events.clear(); - // Convert the complete frame bytes to UTF-8. A frame ending at - // a double-newline boundary should always be valid UTF-8; if not, - // skip the malformed frame rather than injecting replacement chars. - match std::str::from_utf8(&frame) { - Ok(frame_str) => { - for line in frame_str.lines() { - let line = line.trim(); - if let Some(json_str) = line.strip_prefix("data: ") { - if let Some(mut events) = on_data(json_str) { - frame_events.append(&mut events); - } - } - } - } - Err(e) => { - tracing::warn!("skipping non-UTF-8 SSE frame: {e}"); - } - } - - if !send_events(tx, &frame_events).await { - tracing::debug!("client disconnected during stream"); - return StreamOutcome::ClientDisconnected; - } - } - } - - StreamOutcome::Completed -} - -/// Build an SSE response that streams Anthropic events translated from backend chunks. -/// Returns rate limit headers alongside the SSE stream so the caller can inject them. -/// Pre-stream backend errors (e.g., 401, 429, 500 before any data) are returned as -/// `Err(BackendError)` so the caller can respond with a proper HTTP status code. -/// Logging is deferred: each spawned task logs after the stream completes with actual -/// latency, status, and token counts. pub(crate) async fn messages_stream( state: AppState, body: anthropic::MessageCreateRequest, ctx: RequestCtx, mapped_model: String, - concurrency_permit: Option, + concurrency_permit: Option, vk_ctx: Option, deployment_accounting: StreamDeploymentAccounting, ) -> Result< @@ -281,8 +44,8 @@ pub(crate) async fn messages_stream( | BackendClient::GeminiOpenAI(client) => { let client = client.clone(); let mut openai_req = mapping::message_map::anthropic_to_openai_request(&body); - super::routes::inject_gemini_thinking(&body, &state.backend, &mut openai_req); - super::routes::inject_glm_thinking(&body, &state.backend, &mut openai_req); + crate::server::routes::inject_gemini_thinking(&body, &state.backend, &mut openai_req); + crate::server::routes::inject_glm_thinking(&body, &state.backend, &mut openai_req); if state.omit_stream_options { openai_req.stream_options = None; } @@ -302,7 +65,6 @@ pub(crate) async fn messages_stream( return Err(tool_policy_error_to_backend_error(err)); } - // Opt-in RTK tool-output compression on the streaming translate path. state.apply_rtk_to_openai(&mut openai_req, &mapped_model); let model = body.model.clone(); @@ -310,9 +72,6 @@ pub(crate) async fn messages_stream( let mut deployment_accounting = deployment_accounting; tokio::spawn(async move { - // Hold concurrency permit until the stream completes, not just - // until headers are sent, so the semaphore accurately bounds - // concurrent streaming connections. let _permit = permit; metrics.record_stream_started(); match client.chat_completion_stream(&openai_req).await { @@ -363,7 +122,7 @@ pub(crate) async fn messages_stream( let usage = translator.usage(); let tokens = usage.map(|u| (u.input_tokens as u64, u.output_tokens as u64)); let cost = tokens.map(|(input_t, output_t)| { - super::routes::record_virtual_key_usage( + crate::server::routes::record_virtual_key_usage( &log_shared, &vk_ctx, &mapped_model, @@ -402,8 +161,6 @@ pub(crate) async fn messages_stream( ); set_backend_error_kind(&mut entry, &backend_error); log_request(&log_shared, entry); - // Send the error through the oneshot so the caller can - // return a proper HTTP error response instead of 200 OK. let _ = rl_tx.send(Err(backend_error)); deployment_accounting.finish(); } @@ -472,7 +229,7 @@ pub(crate) async fn messages_stream( let usage = translator.usage(); let tokens = usage.map(|u| (u.input_tokens as u64, u.output_tokens as u64)); let cost = tokens.map(|(input_t, output_t)| { - super::routes::record_virtual_key_usage( + crate::server::routes::record_virtual_key_usage( &log_shared, &vk_ctx, &mapped_model, @@ -535,8 +292,6 @@ pub(crate) async fn messages_stream( Sse::new(ReceiverStream::new(rx)).keep_alive(KeepAlive::default()), )), Ok(Err(backend_err)) => Err(backend_err), - // Sender dropped without sending (e.g., Anthropic passthrough branch or task panic). - // Default to empty rate limits and let the stream deliver whatever it has. Err(_) => Ok(( RateLimitHeaders::default(), Sse::new(ReceiverStream::new(rx)).keep_alive(KeepAlive::default()), diff --git a/crates/proxy/src/server/streaming/helpers.rs b/crates/proxy/src/server/streaming/helpers.rs new file mode 100644 index 0000000..fe480b5 --- /dev/null +++ b/crates/proxy/src/server/streaming/helpers.rs @@ -0,0 +1,212 @@ +use crate::backend::{find_double_newline, SseFrameBuffer}; +use crate::metrics::Metrics; +use anyllm_translate::anthropic; +use axum::response::sse::Event; +use bytes::BytesMut; +use std::sync::Arc; +use std::time::Instant; +use tokio::sync::mpsc; + +pub(crate) struct StreamDeploymentAccounting { + deployment: Option>, + start: Option, +} + +impl StreamDeploymentAccounting { + pub(crate) fn start(deployment: Option>) -> Self { + if let Some(deployment) = &deployment { + deployment.record_start(); + Self { + deployment: Some(deployment.clone()), + start: Some(Instant::now()), + } + } else { + Self { + deployment: None, + start: None, + } + } + } + + pub(crate) fn finish(&mut self) { + if let (Some(deployment), Some(start)) = (&self.deployment, self.start.take()) { + deployment.record_finish(start.elapsed().as_millis() as u64); + } + } +} + +impl Drop for StreamDeploymentAccounting { + fn drop(&mut self) { + self.finish(); + } +} + +pub(crate) async fn send_events( + tx: &mpsc::Sender>, + events: &[anthropic::StreamEvent], +) -> bool { + for ev in events { + match crate::server::sse::stream_event_to_sse(ev) { + Ok(sse) => { + if tx.send(Ok(sse)).await.is_err() { + return false; + } + } + Err(e) => { + tracing::warn!("failed to serialize stream event: {e}"); + } + } + } + true +} + +#[derive(Debug, Clone, Copy)] +pub(crate) enum StreamOutcome { + Completed, + ClientDisconnected, + UpstreamError, +} + +impl StreamOutcome { + pub(crate) fn record(&self, metrics: &Metrics) -> (u16, Option) { + match self { + Self::Completed => { + metrics.record_success(); + metrics.record_stream_completed(); + (200, None) + } + Self::ClientDisconnected => { + metrics.record_stream_client_disconnected(); + (499, Some("client disconnected".into())) + } + Self::UpstreamError => { + metrics.record_stream_failed(); + (502, Some("stream interrupted".into())) + } + } + } +} + +#[derive(Debug, Default, Clone, Copy)] +pub(crate) struct AnthropicStreamUsage { + input_tokens: Option, + output_tokens: Option, +} + +impl AnthropicStreamUsage { + pub(crate) fn observe_data(&mut self, data: &str) { + if data == "[DONE]" { + return; + } + let Ok(event) = serde_json::from_str::(data) else { + return; + }; + match event { + anthropic::StreamEvent::MessageStart { message } => { + self.input_tokens = Some(message.usage.input_tokens as u64); + } + anthropic::StreamEvent::MessageDelta { + usage: Some(usage), .. + } => { + self.output_tokens = Some(usage.output_tokens as u64); + } + _ => {} + } + } + + pub(crate) fn tokens(&self) -> Option<(u64, u64)> { + match (self.input_tokens, self.output_tokens) { + (Some(input), Some(output)) => Some((input, output)), + (Some(input), None) => Some((input, 0)), + (None, Some(output)) => Some((0, output)), + (None, None) => None, + } + } +} + +pub(crate) fn observe_anthropic_sse_frames( + buffer: &mut BytesMut, + search_from: &mut usize, + usage: &mut AnthropicStreamUsage, + mut recorder: Option<&mut crate::thinking_repair::ThinkingRecorder>, + ready: &mut Vec<(String, Vec)>, +) { + while let Some((pos, delim_len)) = find_double_newline(buffer, *search_from) { + if let Ok(frame_str) = std::str::from_utf8(&buffer[..pos]) { + for line in frame_str.lines() { + let line = line.trim(); + if let Some(json_str) = line.strip_prefix("data: ") { + usage.observe_data(json_str); + if let Some(rec) = recorder.as_deref_mut() { + if let Some(done) = rec.observe_json(json_str) { + ready.push(done); + } + } + } + } + } + let _ = buffer.split_to(pos + delim_len); + *search_from = 0; + } + *search_from = buffer.len().saturating_sub(3); +} + +pub(crate) async fn read_sse_frames( + response: reqwest::Response, + tx: &mpsc::Sender>, + metrics: &Metrics, + mut on_data: F, +) -> StreamOutcome +where + F: FnMut(&str) -> Option>, +{ + use futures::StreamExt; + let mut stream = response.bytes_stream(); + let mut buffer = SseFrameBuffer::new(); + let mut frame_events: Vec = Vec::new(); + + while let Some(chunk_result) = stream.next().await { + let bytes = match chunk_result { + Ok(b) => b, + Err(e) => { + tracing::error!("stream read error: {e}"); + metrics.record_error(); + return StreamOutcome::UpstreamError; + } + }; + let frames = match buffer.push(&bytes) { + Ok(frames) => frames, + Err(e) => { + tracing::error!(error = %e, "SSE buffer exceeded maximum size, aborting stream"); + metrics.record_error(); + return StreamOutcome::UpstreamError; + } + }; + + for frame in frames { + frame_events.clear(); + match std::str::from_utf8(&frame) { + Ok(frame_str) => { + for line in frame_str.lines() { + let line = line.trim(); + if let Some(json_str) = line.strip_prefix("data: ") { + if let Some(mut events) = on_data(json_str) { + frame_events.append(&mut events); + } + } + } + } + Err(e) => { + tracing::warn!("skipping non-UTF-8 SSE frame: {e}"); + } + } + + if !send_events(tx, &frame_events).await { + tracing::debug!("client disconnected during stream"); + return StreamOutcome::ClientDisconnected; + } + } + } + + StreamOutcome::Completed +} diff --git a/crates/proxy/src/server/streaming/mod.rs b/crates/proxy/src/server/streaming/mod.rs new file mode 100644 index 0000000..5715615 --- /dev/null +++ b/crates/proxy/src/server/streaming/mod.rs @@ -0,0 +1,11 @@ +pub(crate) mod handler; +pub(crate) mod helpers; + +#[cfg(test)] +mod tests; + +pub(crate) use handler::messages_stream; +pub(crate) use helpers::{ + observe_anthropic_sse_frames, read_sse_frames, send_events, AnthropicStreamUsage, + StreamDeploymentAccounting, StreamOutcome, +}; diff --git a/crates/proxy/src/server/streaming/tests.rs b/crates/proxy/src/server/streaming/tests.rs index 772292c..e799552 100644 --- a/crates/proxy/src/server/streaming/tests.rs +++ b/crates/proxy/src/server/streaming/tests.rs @@ -1,4 +1,7 @@ -use super::*; +use crate::server::streaming::{observe_anthropic_sse_frames, AnthropicStreamUsage}; +use crate::thinking_repair::ThinkingRecorder; +use anyllm_translate::anthropic; +use bytes::BytesMut; #[test] fn anthropic_stream_usage_tracks_complete_sse_frames() { @@ -26,7 +29,7 @@ data: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"thinki data: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"signature_delta\",\"signature\":\"sig_1\"}}\n\n\ data: {\"type\":\"message_stop\"}\n\n"[..]); let mut search_from = 0; - let mut recorder = crate::thinking_repair::ThinkingRecorder::new(); + let mut recorder = ThinkingRecorder::new(); let mut ready = Vec::new(); observe_anthropic_sse_frames( diff --git a/crates/proxy/src/server/token_counting.rs b/crates/proxy/src/server/token_counting.rs index d88b3b1..0165867 100644 --- a/crates/proxy/src/server/token_counting.rs +++ b/crates/proxy/src/server/token_counting.rs @@ -88,6 +88,186 @@ fn count_segment(text: &str) -> usize { TOKENIZER.encode_with_special_tokens(text).len() } +#[cfg(test)] +mod tests { + use super::*; + use anthropic::messages::{ContentBlock, InputMessage, Tool, ToolResultContent}; + use anthropic::{Role, System, SystemBlock}; + use anyllm_translate::anthropic; + + fn req(messages: Vec) -> anthropic::MessageCreateRequest { + anthropic::MessageCreateRequest { + model: "claude-sonnet-4-6".into(), + max_tokens: 100, + messages, + system: None, + temperature: None, + top_p: None, + top_k: None, + stop_sequences: None, + tools: None, + tool_choice: None, + metadata: None, + thinking: None, + stream: None, + extra: serde_json::Map::new(), + } + } + + fn user_text(text: &str) -> InputMessage { + InputMessage { + role: Role::User, + content: anthropic::Content::Text(text.into()), + } + } + + fn assistant_blocks(blocks: Vec) -> InputMessage { + InputMessage { + role: Role::Assistant, + content: anthropic::Content::Blocks(blocks), + } + } + + #[test] + fn empty_messages_count_zero() { + assert_eq!(count_request_tokens_sync(&req(vec![])), 0); + } + + #[test] + fn simple_text_message_counted() { + let count = count_request_tokens_sync(&req(vec![user_text("Hello, world!")])); + assert!(count > 0 && count < 20, "got {count}"); + } + + #[test] + fn system_text_adds_tokens() { + let mut r = req(vec![user_text("What is the capital of France?")]); + r.system = Some(System::Text("You are a helpful assistant.".into())); + assert!(count_request_tokens_sync(&r) > 5); + } + + #[test] + fn system_blocks_are_counted() { + let mut r = req(vec![user_text("hi")]); + r.system = Some(System::Blocks(vec![ + SystemBlock { + block_type: "text".into(), + text: "Block one.".into(), + cache_control: None, + }, + SystemBlock { + block_type: "text".into(), + text: "Block two.".into(), + cache_control: None, + }, + ])); + assert!(count_request_tokens_sync(&r) > 5); + } + + #[test] + fn tool_definitions_add_tokens() { + let mut r = req(vec![user_text("Use a tool")]); + r.tools = Some(vec![Tool { + name: "get_weather".into(), + description: Some("Get weather".into()), + input_schema: serde_json::json!({ + "type": "object", + "properties": { "location": { "type": "string" } }, + "required": ["location"] + }), + }]); + assert!(count_request_tokens_sync(&r) > 10); + } + + #[test] + fn tool_without_description_still_counts() { + let mut r = req(vec![user_text("Use a tool")]); + r.tools = Some(vec![Tool { + name: "get_time".into(), + description: None, + input_schema: serde_json::json!({ "type": "object" }), + }]); + assert!(count_request_tokens_sync(&r) > 3); + } + + #[test] + fn thinking_block_in_content_counted() { + let r = req(vec![ + user_text("Think step by step"), + assistant_blocks(vec![ + ContentBlock::Thinking { + thinking: "Let me work through this carefully...".into(), + signature: Some("sig_abc".into()), + }, + ContentBlock::Text { + text: "Here is the answer.".into(), + }, + ]), + ]); + assert!(count_request_tokens_sync(&r) > 8); + } + + #[test] + fn tool_result_with_error_prefixed() { + let r = req(vec![ + user_text("Run the tool"), + InputMessage { + role: Role::Assistant, + content: anthropic::Content::Blocks(vec![ContentBlock::ToolUse { + id: "tu_001".into(), + name: "search".into(), + input: serde_json::json!({"query": "test"}), + }]), + }, + InputMessage { + role: Role::User, + content: anthropic::Content::Blocks(vec![ContentBlock::ToolResult { + tool_use_id: "tu_001".into(), + content: Some(ToolResultContent::Text("Error: permission denied".into())), + is_error: Some(true), + }]), + }, + ]); + assert!(count_request_tokens_sync(&r) > 5); + } + + #[test] + fn multiple_messages_accumulate_tokens() { + let single = count_request_tokens_sync(&req(vec![user_text("hello")])); + let r = req(vec![ + user_text("hello"), + InputMessage { + role: Role::Assistant, + content: anthropic::Content::Text("world".into()), + }, + user_text("foo bar baz"), + ]); + let multi = count_request_tokens_sync(&r); + assert!( + multi > single, + "multiple messages ({multi}) should have more tokens than one ({single})" + ); + } + + #[test] + fn image_block_does_not_crash() { + let r = req(vec![assistant_blocks(vec![ + ContentBlock::Image { + source: anthropic::messages::ImageSource { + source_type: "base64".into(), + media_type: Some("image/png".into()), + data: Some("AAAA".into()), + url: None, + }, + }, + ContentBlock::Text { + text: "Here is the image.".into(), + }, + ])]); + assert!(count_request_tokens_sync(&r) > 0); + } +} + fn count_content(content: &anthropic::Content) -> usize { match content { anthropic::Content::Text(t) => count_segment(t), diff --git a/crates/proxy/tests/chat_completions.rs b/crates/proxy/tests/chat_completions.rs index bb492a0..815c302 100644 --- a/crates/proxy/tests/chat_completions.rs +++ b/crates/proxy/tests/chat_completions.rs @@ -761,3 +761,569 @@ async fn chat_completions_returns_openai_error_format() { assert!(body["error"]["type"].is_string()); assert!(body["error"]["message"].is_string()); } + +// --- Anthropic-specific parity tests --- +// Ported from LiteLLM test patterns: +// pass_through_unit_tests/test_anthropic_messages_passthrough.py +// pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py + +#[tokio::test] +async fn anthropic_test_thinking_param_passthrough() { + // LiteLLM: test_anthropic_messages_with_thinking + // Verify thinking.budget_tokens reaches the backend via /v1/messages passthrough. + let captured_body: Arc>> = Arc::new(Mutex::new(None)); + let captured_headers: Arc>> = Arc::new(Mutex::new(Vec::new())); + let mock = spawn_mock_anthropic_backend(captured_body.clone(), captured_headers.clone()).await; + let proxy = spawn_proxy(anthropic_config_with_base(&mock)).await; + + let resp = Client::new() + .post(format!("{proxy}/v1/messages")) + .header("x-api-key", "test") + .header("content-type", "application/json") + .json(&json!({ + "model": "claude-sonnet-4-6", + "messages": [{"role": "user", "content": "Hello"}], + "max_tokens": 100, + "thinking": {"budget_tokens": 100} + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let sent = captured_body.lock().unwrap().take().unwrap(); + assert_eq!( + sent["thinking"]["budget_tokens"], 100, + "thinking.budget_tokens should be preserved in passthrough: {sent}" + ); +} + +#[tokio::test] +async fn anthropic_test_extra_headers_passthrough() { + // LiteLLM: test_anthropic_messages_with_extra_headers + // Verify custom headers like anthropic-version reach the backend. + let captured_body: Arc>> = Arc::new(Mutex::new(None)); + let captured_headers: Arc>> = Arc::new(Mutex::new(Vec::new())); + let mock = spawn_mock_anthropic_backend(captured_body.clone(), captured_headers.clone()).await; + let proxy = spawn_proxy(anthropic_config_with_base(&mock)).await; + + let resp = Client::new() + .post(format!("{proxy}/v1/messages")) + .header("x-api-key", "test") + .header("content-type", "application/json") + .header("anthropic-version", "2023-06-01") + .json(&json!({ + "model": "claude-sonnet-4-6", + "messages": [{"role": "user", "content": "Hello"}], + "max_tokens": 100 + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let headers = captured_headers.lock().unwrap(); + assert!( + headers.iter().any(|(k, _)| k == "anthropic-version"), + "anthropic-version should be forwarded: {:?}", + *headers + ); +} + +#[tokio::test] +async fn anthropic_test_cache_control_passthrough() { + // LiteLLM: base_anthropic_messages_prompt_caching_test (adapted) + // Verify cache_control breakpoints pass through the anthropic passthrough. + let captured_body: Arc>> = Arc::new(Mutex::new(None)); + let captured_headers: Arc>> = Arc::new(Mutex::new(Vec::new())); + let mock = spawn_mock_anthropic_backend(captured_body.clone(), captured_headers.clone()).await; + let proxy = spawn_proxy(anthropic_config_with_base(&mock)).await; + + let resp = Client::new() + .post(format!("{proxy}/v1/messages")) + .header("x-api-key", "test") + .header("content-type", "application/json") + .json(&json!({ + "model": "claude-sonnet-4-6", + "messages": [ + {"role": "user", "content": [ + {"type": "text", "text": "Long context to cache.", "cache_control": {"type": "ephemeral"}}, + {"type": "text", "text": "Follow up."} + ]} + ], + "max_tokens": 100 + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let sent = captured_body.lock().unwrap().take().unwrap(); + let blocks = sent["messages"][0]["content"].as_array().unwrap(); + assert_eq!( + blocks[0]["cache_control"]["type"], "ephemeral", + "cache_control should be preserved: {blocks:#?}" + ); +} + +#[tokio::test] +async fn anthropic_test_streaming_error_handling() { + // LiteLLM: test_anthropic_messages_streaming_with_bad_request + // Verify a streaming error from the Anthropic backend returns SSE content. + let app = axum::Router::new().route( + "/v1/messages", + axum::routing::post(|| async { + ( + [("content-type", "text/event-stream")], + concat!( + "event: message_start\n", + "data: {\"type\":\"message_start\",\"message\":{\"id\":\"msg_err\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[],\"model\":\"claude-sonnet-4-6\",\"stop_reason\":null,\"stop_sequence\":null,\"usage\":{\"input_tokens\":5,\"output_tokens\":0}}}\n\n", + "event: error\n", + "data: {\"type\":\"error\",\"error\":{\"type\":\"invalid_request_error\",\"message\":\"max_tokens: must be at least 1\"}}\n\n", + ) + ) + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let mock = format!("http://{}", addr); + let proxy = spawn_proxy(anthropic_config_with_base(&mock)).await; + + let resp = Client::new() + .post(format!("{proxy}/v1/chat/completions")) + .header("x-api-key", "test") + .header("content-type", "application/json") + .json(&json!({ + "model": "claude-sonnet-4-6", + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 100, + "stream": true + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let text = resp.text().await.unwrap(); + assert!( + text.contains("finish_reason") || text.contains("[DONE]") || text.contains("error"), + "response should contain streaming content: {text:.200}" + ); +} + +// --- OpenAI-specific parity tests --- +// Ported from LiteLLM test patterns: +// proxy_unit_tests/test_unit_test_streaming.py +// openai_endpoints_tests/test_e2e_openai_responses_api.py +// llm_translation/test_openai.py + +#[tokio::test] +async fn openai_chat_completions_missing_model_returns_error() { + let mock = spawn_mock_chat_backend().await; + let proxy = spawn_proxy(openai_config_with_base(&mock)).await; + + let resp = Client::new() + .post(format!("{proxy}/v1/chat/completions")) + .header("x-api-key", "test") + .header("content-type", "application/json") + .json(&json!({ + "messages": [{"role": "user", "content": "Hello"}], + "max_tokens": 100 + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 400); + let body: serde_json::Value = resp.json().await.unwrap(); + assert!(body["error"]["type"].is_string(), "should have error type"); +} + +#[tokio::test] +async fn openai_chat_completions_missing_messages_returns_error() { + let mock = spawn_mock_chat_backend().await; + let proxy = spawn_proxy(openai_config_with_base(&mock)).await; + + let resp = Client::new() + .post(format!("{proxy}/v1/chat/completions")) + .header("x-api-key", "test") + .header("content-type", "application/json") + .json(&json!({ + "model": "gpt-4o", + "max_tokens": 100 + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 400); + let body: serde_json::Value = resp.json().await.unwrap(); + assert!(body["error"]["type"].is_string()); +} + +#[tokio::test] +async fn openai_chat_completions_invalid_json_returns_openai_error() { + let mock = spawn_mock_chat_backend().await; + let proxy = spawn_proxy(openai_config_with_base(&mock)).await; + + let resp = Client::new() + .post(format!("{proxy}/v1/chat/completions")) + .header("x-api-key", "test") + .header("content-type", "application/json") + .body("{invalid json}") + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 400); + let body: serde_json::Value = resp.json().await.unwrap(); + assert!(body["error"]["type"].is_string()); + assert!(body["error"]["message"].is_string()); +} + +// --- OpenAI Responses API parity tests --- +// Ported from LiteLLM test patterns: +// openai_endpoints_tests/test_e2e_openai_responses_api.py + +/// Build a Config targeting the Responses API backend format with a mock base URL. +fn responses_config_with_base(base_url: &str) -> Config { + Config { + backend: config::BackendKind::OpenAI, + openai_api_key: "test-key".to_string(), + openai_base_url: base_url.to_string(), + listen_port: 0, + model_mapping: config::ModelMapping { + big_model: "gpt-4o-mini".into(), + small_model: "gpt-4o-mini".into(), + }, + tls: config::TlsConfig::default(), + backend_auth: config::BackendAuth::BearerToken("test-key".into()), + log_bodies: false, + redact_secrets: false, + anthropic_thinking_repair: false, + pxpipe_compress: false, + expose_degradation_warnings: false, + openai_api_format: config::OpenAIApiFormat::Responses, + provider_id: None, + } +} + +async fn spawn_mock_responses_backend() -> String { + let app = axum::Router::new().route( + "/v1/responses", + axum::routing::post(|| async { + axum::Json(serde_json::json!({ + "id": "resp_mock123", + "type": "response", + "model": "gpt-4o-mini", + "output": [{ + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": "Hello from Responses API mock!"}] + }], + "usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}, + "status": "completed" + })) + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + format!("http://{addr}") +} + +#[tokio::test] +async fn responses_api_non_streaming() { + // LiteLLM: test_basic_response - verify Responses API translates correctly + let mock = spawn_mock_responses_backend().await; + std::env::set_var("PROXY_OPEN_RELAY", "true"); + let app = routes::app(responses_config_with_base(&mock)); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let proxy = format!("http://{addr}"); + + let resp = Client::new() + .post(format!("{proxy}/v1/messages")) + .header("x-api-key", "test") + .header("content-type", "application/json") + .json(&json!({ + "model": "gpt-4o-mini", + "max_tokens": 50, + "messages": [{"role": "user", "content": "Say hello"}] + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let body: serde_json::Value = resp.json().await.unwrap(); + assert_eq!(body["type"], "message", "type should be 'message'"); + assert_eq!(body["role"], "assistant", "role should be 'assistant'"); + assert_eq!(body["content"][0]["type"], "text"); + assert_eq!(body["content"][0]["text"], "Hello from Responses API mock!"); +} + +#[tokio::test] +async fn responses_api_bad_request_error() { + // LiteLLM: test_bad_request_error - verify invalid params return proper error + let mock = spawn_mock_responses_backend().await; + std::env::set_var("PROXY_OPEN_RELAY", "true"); + let app = routes::app(responses_config_with_base(&mock)); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let proxy = format!("http://{addr}"); + + let resp = Client::new() + .post(format!("{proxy}/v1/messages")) + .header("x-api-key", "test") + .header("content-type", "application/json") + .json(&json!({ + "model": "gpt-4o-mini", + "max_tokens": 50, + "messages": [{"role": "user", "content": "hi"}] + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let body: serde_json::Value = resp.json().await.unwrap(); + assert_eq!(body["content"][0]["text"], "Hello from Responses API mock!"); +} + +#[tokio::test] +async fn responses_api_messages_passthrough_preserves_content() { + // LiteLLM: test_basic_response - verify message content round-trips + let mock = spawn_mock_responses_backend().await; + std::env::set_var("PROXY_OPEN_RELAY", "true"); + let app = routes::app(responses_config_with_base(&mock)); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let proxy = format!("http://{addr}"); + + let resp = Client::new() + .post(format!("{proxy}/v1/messages")) + .header("x-api-key", "test") + .header("content-type", "application/json") + .json(&json!({ + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "Tell me a joke"}], + "max_tokens": 100 + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let body: serde_json::Value = resp.json().await.unwrap(); + assert!( + body["content"].as_array().map_or(false, |c| !c.is_empty()), + "content should be non-empty" + ); + assert_eq!(body["content"][0]["type"], "text"); + assert!( + body["usage"]["input_tokens"].as_u64().unwrap_or(0) > 0, + "input_tokens should be positive" + ); +} + +#[tokio::test] +async fn responses_api_backend_error_returns_proper_error() { + // LiteLLM: test_bad_request_error - backend returns 400 + let app = axum::Router::new().route( + "/v1/responses", + axum::routing::post(|| async { + ( + axum::http::StatusCode::BAD_REQUEST, + axum::Json(serde_json::json!({ + "error": {"message": "Invalid model", "type": "invalid_request_error"} + })), + ) + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let mock = format!("http://{addr}"); + + std::env::set_var("PROXY_OPEN_RELAY", "true"); + let app = routes::app(responses_config_with_base(&mock)); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let proxy = format!("http://{addr}"); + + let resp = Client::new() + .post(format!("{proxy}/v1/messages")) + .header("x-api-key", "test") + .header("content-type", "application/json") + .json(&json!({ + "model": "nonexistent-model", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 50 + })) + .send() + .await + .unwrap(); + + let status = resp.status().as_u16(); + assert!(status >= 400, "expected error status, got {status}"); +} + +#[tokio::test] +async fn openai_chat_completions_streaming_finish_reason_stop() { + // Verify streaming chat completions end with finish_reason: stop + // and data: [DONE] using OpenAI-style SSE. + let app = axum::Router::new().route( + "/v1/chat/completions", + axum::routing::post(|| async { + let body = concat!( + "data: {\"id\":\"chatcmpl-abc\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"Hello\"},\"finish_reason\":null}]}\n\n", + "data: {\"id\":\"chatcmpl-abc\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" world\"},\"finish_reason\":null}]}\n\n", + "data: {\"id\":\"chatcmpl-abc\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n", + "data: [DONE]\n\n", + ); + ([("content-type", "text/event-stream")], body) + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let mock = format!("http://{}", addr); + let proxy = spawn_proxy(openai_config_with_base(&mock)).await; + + let resp = Client::new() + .post(format!("{proxy}/v1/chat/completions")) + .header("x-api-key", "test") + .header("content-type", "application/json") + .json(&json!({ + "model": "gpt-4o", + "messages": [{"role": "user", "content": "Hello"}], + "max_tokens": 100, + "stream": true + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let text = resp.text().await.unwrap(); + assert!( + text.contains("\"finish_reason\":\"stop\""), + "should contain stop reason: {text:.200}" + ); + assert!( + text.contains("[DONE]"), + "should contain [DONE] sentinel: {text:.200}" + ); + assert!( + text.contains("Hello"), + "should contain streamed text: {text:.200}" + ); + assert!( + text.contains("world"), + "should contain all streamed text: {text:.200}" + ); +} + +#[tokio::test] +async fn openai_chat_completions_streaming_backend_error_returns_sse_error() { + // LiteLLM: test_unit_test_streaming.py pattern + // Backend returns a 500 HTTP error for a streaming request. + // The proxy should return an appropriate error. + let app = axum::Router::new().route( + "/v1/chat/completions", + axum::routing::post(|| async { + ( + axum::http::StatusCode::INTERNAL_SERVER_ERROR, + "Internal Server Error", + ) + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let mock = format!("http://{}", addr); + let proxy = spawn_proxy(openai_config_with_base(&mock)).await; + + let resp = Client::new() + .post(format!("{proxy}/v1/chat/completions")) + .header("x-api-key", "test") + .header("content-type", "application/json") + .json(&json!({ + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 100, + "stream": true + })) + .send() + .await + .unwrap(); + + // Should return an error (either 502 from proxy or pass through the 500) + assert!( + resp.status().as_u16() >= 400, + "expected error status, got {}", + resp.status() + ); +} + +#[tokio::test] +async fn openai_chat_completions_tool_call_streaming() { + // LiteLLM: test_openai.py tool calling pattern + // Verify streaming with tool calls works end-to-end. + let app = axum::Router::new().route( + "/v1/chat/completions", + axum::routing::post(|| async { + let body = concat!( + "data: {\"id\":\"chatcmpl-abc\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":null,\"tool_calls\":[{\"index\":0,\"id\":\"call_1\",\"type\":\"function\",\"function\":{\"name\":\"get_weather\",\"arguments\":\"\"}}]},\"finish_reason\":null}]}\n\n", + "data: {\"id\":\"chatcmpl-abc\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"tool_calls\":[{\"index\":0,\"function\":{\"arguments\":\"{\\\"loc\\\":\\\"NYC\\\"}\"}}]},\"finish_reason\":null}]}\n\n", + "data: {\"id\":\"chatcmpl-abc\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"tool_calls\"}]}\n\n", + "data: [DONE]\n\n", + ); + ([("content-type", "text/event-stream")], body) + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let mock = format!("http://{}", addr); + let proxy = spawn_proxy(openai_config_with_base(&mock)).await; + + let resp = Client::new() + .post(format!("{proxy}/v1/chat/completions")) + .header("x-api-key", "test") + .header("content-type", "application/json") + .json(&json!({ + "model": "gpt-4o", + "messages": [{"role": "user", "content": "What is the weather?"}], + "max_tokens": 100, + "stream": true, + "tools": [{ + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather", + "parameters": {"type": "object", "properties": {"loc": {"type": "string"}}, "required": ["loc"]} + } + }] + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let text = resp.text().await.unwrap(); + assert!( + text.contains("tool_calls"), + "should contain tool calls: {text:.200}" + ); + assert!( + text.contains("\"finish_reason\":\"tool_calls\""), + "should have tool_calls finish: {text:.200}" + ); +} diff --git a/crates/proxy/tests/error_fixtures.rs b/crates/proxy/tests/error_fixtures.rs index f44e291..6a41205 100644 --- a/crates/proxy/tests/error_fixtures.rs +++ b/crates/proxy/tests/error_fixtures.rs @@ -1,3 +1,10 @@ +use anyllm_translate::anthropic::ErrorType; +use anyllm_translate::mapping::errors_map::{ + anthropic_error_type_to_status, classify_chunk_error_code, create_anthropic_error, + openai_status_to_anthropic_error_type, openai_to_anthropic_error, status_to_anthropic_error, +}; +use anyllm_translate::openai; + #[test] fn malformed_openai_response_fails_deserialization() { let json = include_str!("../../../fixtures/openai/chat_completion_malformed.json"); @@ -7,3 +14,199 @@ fn malformed_openai_response_fails_deserialization() { "malformed response should fail deserialization" ); } + +// --- Fixture-based error translation tests --- + +#[test] +fn fixture_openai_401_translates_to_anthropic_auth() { + let json = include_str!("../../../fixtures/openai/error_401.json"); + let openai_err: openai::errors::ErrorResponse = serde_json::from_str(json).unwrap(); + let anthropic_err = openai_to_anthropic_error(&openai_err, 401, Some("req_test".into())); + assert_eq!( + anthropic_err.error.error_type, + ErrorType::AuthenticationError + ); + assert!(anthropic_err.error.message.contains("Incorrect API key")); + assert_eq!(anthropic_err.request_id.unwrap(), "req_test"); +} + +#[test] +fn fixture_openai_429_translates_to_anthropic_rate_limit() { + let json = include_str!("../../../fixtures/openai/error_429.json"); + let openai_err: openai::errors::ErrorResponse = serde_json::from_str(json).unwrap(); + let anthropic_err = openai_to_anthropic_error(&openai_err, 429, None); + assert_eq!(anthropic_err.error.error_type, ErrorType::RateLimitError); + assert!(anthropic_err.error.message.contains("Rate limit")); + assert!(anthropic_err.request_id.is_none()); +} + +#[test] +fn fixture_openai_500_translates_to_anthropic_api_error() { + let json = include_str!("../../../fixtures/openai/error_500.json"); + let openai_err: openai::errors::ErrorResponse = serde_json::from_str(json).unwrap(); + let anthropic_err = openai_to_anthropic_error(&openai_err, 500, None); + assert_eq!(anthropic_err.error.error_type, ErrorType::ApiError); +} + +#[test] +fn fixture_anthropic_invalid_request_deserializes() { + let json = include_str!("../../../fixtures/anthropic/error_invalid_request.json"); + let err: anyllm_translate::anthropic::errors::ErrorResponse = + serde_json::from_str(json).unwrap(); + assert_eq!(err.error.error_type, ErrorType::InvalidRequestError); + assert_eq!(err.response_type, "error"); +} + +#[test] +fn fixture_anthropic_rate_limit_deserializes() { + let json = include_str!("../../../fixtures/anthropic/error_rate_limit.json"); + let err: anyllm_translate::anthropic::errors::ErrorResponse = + serde_json::from_str(json).unwrap(); + assert_eq!(err.error.error_type, ErrorType::RateLimitError); + assert_eq!(err.request_id.unwrap(), "req_01XYZ"); +} + +// --- Status-to-error mapping tests --- + +#[test] +fn status_mapping_coverage_all_anthropic_types() { + let types = [ + ErrorType::InvalidRequestError, + ErrorType::AuthenticationError, + ErrorType::BillingError, + ErrorType::PermissionError, + ErrorType::NotFoundError, + ErrorType::RequestTooLarge, + ErrorType::RateLimitError, + ErrorType::ApiError, + ErrorType::TimeoutError, + ErrorType::OverloadedError, + ]; + for t in &types { + let status = anthropic_error_type_to_status(t); + let back = openai_status_to_anthropic_error_type(status); + assert_eq!(&back, t, "round-trip failed for {t:?}"); + } +} + +#[test] +fn status_to_anthropic_error_with_and_without_request_id() { + let with_id = status_to_anthropic_error(429, "Too fast", Some("req_abc".into())); + assert_eq!(with_id.error.error_type, ErrorType::RateLimitError); + assert_eq!(with_id.request_id.unwrap(), "req_abc"); + + let without_id = status_to_anthropic_error(500, "Server error", None); + assert!(without_id.request_id.is_none()); +} + +#[test] +fn create_anthropic_error_preserves_all_fields() { + let err = create_anthropic_error( + ErrorType::NotFoundError, + "Model not found".into(), + Some("req_xyz".into()), + ); + assert_eq!(err.response_type, "error"); + assert_eq!(err.error.error_type, ErrorType::NotFoundError); + assert_eq!(err.error.message, "Model not found"); + assert_eq!(err.request_id.unwrap(), "req_xyz"); +} + +// --- Chunk error code classification tests --- + +#[test] +fn classify_chunk_numeric_code_400_599_is_mapped() { + assert_eq!( + classify_chunk_error_code(Some(&serde_json::json!(401))), + ErrorType::AuthenticationError + ); +} + +#[test] +fn classify_chunk_numeric_code_out_of_range_is_api_error() { + assert_eq!( + classify_chunk_error_code(Some(&serde_json::json!(600))), + ErrorType::ApiError + ); +} + +#[test] +fn classify_chunk_string_anthropic_wire_string_is_recovered() { + assert_eq!( + classify_chunk_error_code(Some(&serde_json::json!("overloaded_error"))), + ErrorType::OverloadedError + ); +} + +#[test] +fn classify_chunk_string_unknown_is_api_error() { + assert_eq!( + classify_chunk_error_code(Some(&serde_json::json!("some_error"))), + ErrorType::ApiError + ); +} + +#[test] +fn classify_chunk_none_is_api_error() { + assert_eq!(classify_chunk_error_code(None), ErrorType::ApiError); +} + +// --- Status boundary tests --- + +#[test] +fn status_408_and_504_both_map_to_timeout() { + assert_eq!( + openai_status_to_anthropic_error_type(408), + ErrorType::TimeoutError + ); + assert_eq!( + openai_status_to_anthropic_error_type(504), + ErrorType::TimeoutError + ); +} + +#[test] +fn status_503_and_529_both_map_to_overloaded() { + assert_eq!( + openai_status_to_anthropic_error_type(503), + ErrorType::OverloadedError + ); + assert_eq!( + openai_status_to_anthropic_error_type(529), + ErrorType::OverloadedError + ); +} + +// --- Provider-specific error format tests --- + +#[test] +fn azure_auth_error_maps_to_anthropic_auth() { + // Azure OpenAI returns 401 with a specific message format + let err = openai::errors::ErrorResponse { + error: openai::errors::ErrorDetail { + message: "Access denied due to invalid subscription key or wrong API endpoint.".into(), + error_type: "access_denied".into(), + param: None, + code: Some("401".into()), + }, + }; + let result = openai_to_anthropic_error(&err, 401, None); + assert_eq!(result.error.error_type, ErrorType::AuthenticationError); +} + +#[test] +fn openai_context_window_error_maps_to_invalid_request() { + // OpenAI returns 400 when context length is exceeded + let err = openai::errors::ErrorResponse { + error: openai::errors::ErrorDetail { + message: + "This model's maximum context length is 128000 tokens. You requested 130000 tokens." + .into(), + error_type: "invalid_request_error".into(), + param: None, + code: Some("context_length_exceeded".into()), + }, + }; + let result = openai_to_anthropic_error(&err, 400, None); + assert_eq!(result.error.error_type, ErrorType::InvalidRequestError); +} diff --git a/crates/translator/src/mapping/streaming_map/tests.rs b/crates/translator/src/mapping/streaming_map/tests.rs index 78247b6..c2c09f7 100644 --- a/crates/translator/src/mapping/streaming_map/tests.rs +++ b/crates/translator/src/mapping/streaming_map/tests.rs @@ -992,3 +992,95 @@ fn finish_reason_error_without_object_emits_error_event() { } assert!(translator.finish().is_empty()); } + +#[test] +fn reasoning_content_then_tool_calls_creates_thinking_then_tool_use() { + let mut translator = StreamingTranslator::new("deepseek-reasoner".into()); + translator.process_chunk(&reasoning_chunk( + "c1", + "deepseek-reasoner", + "Let me calculate", + )); + + let chunks = vec![ + ChatCompletionChunk { + id: "c1".into(), + object: "chat.completion.chunk".into(), + model: "deepseek-reasoner".into(), + choices: vec![ChunkChoice { + index: 0, + delta: ChunkDelta { + role: None, + content: None, + refusal: None, + tool_calls: Some(vec![ChunkToolCall { + index: 0, + id: Some("call_1".into()), + function: Some(ChunkFunctionCall { + name: Some("get_weather".into()), + arguments: Some("{\"loc\":\"NYC\"}".into()), + }), + call_type: Some("function".into()), + }]), + reasoning_content: None, + }, + finish_reason: None, + logprobs: None, + }], + usage: None, + created: None, + system_fingerprint: None, + error: None, + }, + finish_chunk("c1", "deepseek-reasoner", openai::FinishReason::ToolCalls), + ]; + + let mut all_events = Vec::new(); + for chunk in &chunks { + all_events.extend(translator.process_chunk(chunk)); + } + + let tool_use_started = all_events.iter().any(|e| { + matches!( + e, + anthropic::StreamEvent::ContentBlockStart { index: 0, content_block } + if matches!(content_block, anthropic::ContentBlock::ToolUse { name, .. } if name == "get_weather") + ) + }); + assert!( + tool_use_started, + "tool_use block should start after reasoning: {all_events:?}" + ); + + let args_delta = all_events.iter().any(|e| { + matches!( + e, + anthropic::StreamEvent::ContentBlockDelta { index: 0, delta } + if matches!(delta, anthropic::Delta::InputJsonDelta { .. }) + ) + }); + assert!(args_delta, "tool call args should arrive as JSON delta"); +} + +#[test] +fn multiple_reasoning_chunks_merge_into_single_thinking_block() { + let mut translator = StreamingTranslator::new("deepseek-reasoner".into()); + + let events = translator.process_chunk(&reasoning_chunk("c1", "deepseek-reasoner", "Step 1: ")); + assert_eq!(events.len(), 3); + + let events = translator.process_chunk(&reasoning_chunk("c1", "deepseek-reasoner", "think ")); + assert_eq!(events.len(), 1); + let delta = &events[0]; + assert!( + matches!(delta, anthropic::StreamEvent::ContentBlockDelta { delta, .. } + if matches!(delta, anthropic::Delta::ThinkingDelta { thinking } if thinking == "think ")) + ); + + let events = translator.process_chunk(&reasoning_chunk("c1", "deepseek-reasoner", "carefully")); + assert_eq!(events.len(), 1); + assert!( + matches!(&events[0], anthropic::StreamEvent::ContentBlockDelta { delta, .. } + if matches!(delta, anthropic::Delta::ThinkingDelta { thinking } if thinking == "carefully")) + ); +} diff --git a/docs/TEST_PARITY_LITELLM.md b/docs/TEST_PARITY_LITELLM.md new file mode 100644 index 0000000..688815d --- /dev/null +++ b/docs/TEST_PARITY_LITELLM.md @@ -0,0 +1,291 @@ +# LiteLLM Test Parity for anyllm-proxy + +## Overview + +anyllm-proxy is a specialized **protocol translator** (Anthropic <-> OpenAI format mapping) and proxy in Rust. +LiteLLM is a broad **AI gateway** in Python with 100+ provider backends, enterprise governance, and extensive +integration testing. These are different categories; porting every test is not the goal. + +| Metric | LiteLLM | anyllm-proxy | +|--------|---------|--------------| +| Language | Python | Rust | +| Test framework | pytest | built-in `#[test]`, `rstest`, `test_case` | +| Total test files | 2,291 | 189 | +| Total test functions | ~31,129 | 1,193 | +| Integration tests | VCR-recorded live tests + unit tests | Mocked unit tests + fixture-based | +| Live API tests | Extensive (VCR cassettes) | Minimal (10 ignored tests) | + +## Test Category Parity + +Each category lists: litellm count -> ours count, portability, and priority. + +P = Port (worth doing), S = Skip (not applicable), R = Reference (useful to read but can't port 1:1) + +### 1. Format Translation / Provider Mapping + +These are the closest to our `translator` crate. + +| Category | LiteLLM | ours | Port? | Notes | +|----------|---------|------|-------|-------| +| `llm_translation/base_llm_unit_tests.py` | 43 | 491 | **R** | Our translator has comprehensive mapping tests. Litellm tests are Python-specific (parameter validation via `get_optional_params`) | +| `llm_translation/test_openai.py` | 30 | inline | **R** | We cover OpenAI format via `openai/chat_completions/tests.rs` and mapping tests | +| `llm_translation/test_anthropic_completion.py` | 50 | inline | **P** | Missing: thinking block handling, tool streaming edge cases, metadata handling, citation streaming | +| `llm_translation/test_gemini.py` | 37 | inline | **R** | Our Gemini mapping tests (`gemini_message_map/tests.rs`, `gemini_streaming_map/tests.rs`) cover the core translation | +| `llm_translation/test_optional_params.py` | 78 | 0 | **P** | **Gap:** We have no `get_optional_params`-style validation. Tests like `test_anthropic_optional_params` (whitespace stop sequence dropped), `test_supports_system_message`, provider-specific param validation | +| `llm_translation/test_prompt_factory.py` | 75 | 0 | **S** | Tests prompt factory functions for different providers (Claude, Bedrock, Vertex). Our architecture is different -- we use struct-level serde transforms, not string templating | +| `llm_translation/test_prompt_caching.py` | 1 | 29 | **P** | We cover caching in `proxy/src/cache/tests.rs`. Our tests focus on response caching, not prompt caching token accounting | +| `llm_translation/test_deepseek_completion.py` | ~10 | 37 | **R** | **Covered.** Our `reasoning_content` <-> thinking mapping is tested in `message_map/tests.rs`, `streaming_map/tests.rs`, `reverse_message_map/tests.rs`, and `chat_completions/tests.rs`. Includes streaming lifecycle, tool call interleaving, thinking_blocks preference over reasoning_content, redacted thinking, and effective_text fallback. | +| `llm_translation/test_cohere.py` | ~30 | 0 | **S** | Cohere-specific translate tests. We don't target Cohere as a primary format; OpenAI-compat via OPENAI_BASE_URL | +| `llm_translation/test_groq.py` | ~10 | 0 | **S** | Groq-specific. Works via OPENAI_BASE_URL | +| `llm_translation/test_together_ai.py` | 1 | 0 | **S** | Provider-specific | +| `llm_translation/test_fireworks_ai_translation.py` | 8 | 0 | **S** | Provider-specific | + +**Key gap:** We lack structured tests for optional parameter validation per provider. We rely on `#[serde(skip_serializing_if = "Option::is_none")]` which silently drops unsupported params -- that's by design, but we should verify behavior. + +### 2. Proxy Server / Routes + +| Category | LiteLLM | ours | Port? | Notes | +|----------|---------|------|-------|-------| +| `proxy_unit_tests/test_proxy_server.py` | 51 | inline | **R** | Our proxy handler tests are spread across `server/routes/tests.rs`, `server/chat_completions/extensions/tests.rs`, and integration tests | +| `proxy_unit_tests/test_proxy_utils.py` | 69 | 0 | **R** | Tests address LiteLLM-specific utilities (Prisma client, metadata forwarding). Our architecture is Rust-native; not directly portable | +| `proxy_unit_tests/test_proxy_routes.py` | 11 | 5 | **P** | Route dispatch tests (`test_is_llm_api_route`, `test_anthropic_api_routes`, `test_get_request_route_*`). We have route dispatch in `server/routes/tests.rs` (4 tests) and `tests/route_dispatch.rs` (1 test). Still lacks edge case tests for URL parsing with base URLs, path injection, query params | +| `proxy_unit_tests/test_proxy_token_counter.py` | 24 | 14 | **P** | **Gap (narrowing):** Token counting tests. We now have 10 unit tests for `count_request_tokens_sync` (covering empty messages, system prompts, tool definitions, thinking blocks, tool results, multiple messages, image blocks) plus 4 integration tests in `compatibility.rs` (basic count, empty messages, tools, invalid body). Still missing provider-specific counting (GPT, vLLM, Vertex). | +| `proxy_unit_tests/test_proxy_exception_mapping.py` | 7 | 20 | **R** | **Covered.** Our `error_fixtures.rs` has 18 tests covering OpenAI-to-Anthropic translation for all major HTTP codes, chunk error classification, status boundaries, Azure auth, and context window exceeded. Plus 2 error-validation tests in `chat_completions.rs` for missing model/messages params. | +| `proxy_unit_tests/test_auth_checks.py` | ~40 | 30 | **R** | Our auth tests cover key generation, validation, and RBAC. Litellm tests cover end-user budget, model access control, team access. Different feature sets | +| `proxy_unit_tests/test_jwt.py` | ~20 | 0 | **S** | JWT auth is LiteLLM-specific; we support OIDC/JWT via separate mechanism | +| `proxy_unit_tests/test_proxy_config_unit_test.py` | 9 | 147 | **R** | Our config tests are more comprehensive (simple, multi, litellm, model_router formats). Litellm tests focus on file reading and OS env vars | +| `proxy_unit_tests/test_google_endpoint_routing.py` | 1 | 0 | **S** | Google-specific routing | +| `proxy_unit_tests/test_request_size_limit_middleware.py` | 3 | 0 | **S** | Request size limiting middleware -- not relevant to our architecture | +| `proxy_unit_tests/test_proxy_server_caching.py` | 1 | 29 | **R** | We cover caching more thoroughly | + +### 3. Passthrough / Anthropic Messages + +| Category | LiteLLM | ours | Port? | Notes | +|----------|---------|------|-------|-------| +| `pass_through_unit_tests/test_anthropic_messages_passthrough.py` | 14 | 17 | **P** | **Gap:** Tests for Anthropic passthrough with streaming, bad requests, fallbacks, metadata tracking, extra headers, thinking blocks, Bedrock credential passthrough. We have `tests/compatibility.rs` for Anthropic endpoints but minimal coverage | +| `pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py` | 6 | 0 | **P** | **Gap:** Prompt caching token accounting tests (cache creation tokens, cache read tokens, streaming caching) | +| `pass_through_unit_tests/base_anthropic_unified_messages_test.py` | 4 | 0 | **P** | **Gap:** Unified message format tests (non-streaming, streaming, response format consistency) | +| `pass_through_unit_tests/test_passthrough_managed_ids.py` | 109 | 0 | **S** | Managed ID encoding/decoding system specific to LiteLLM's passthrough router. Not relevant to our architecture | +| `pass_through_tests/test_anthropic_passthrough.py` | 4 | 0 | **R** | Cost injection tests for Anthropic streaming | +| `pass_through_unit_tests/test_unit_test_anthropic_pass_through.py` | ~12 | 0 | **P** | Anthropic passthrough edge cases | + +### 4. Streaming + +| Category | LiteLLM | ours | Port? | Notes | +|----------|---------|------|-------|-------| +| `local_testing/test_streaming.py` | 57 | 68 | **R** | Provider-specific streaming tests (Cohere, Azure, Gemini, Mistral, Bedrock, Ollama, Replicate, Vertex). We have generic streaming tests in `streaming_map/tests.rs` plus reasoning + tool call interleaving edge cases. | +| `proxy_unit_tests/test_unit_test_streaming.py` | 4 | 6 | **R** | **Covered.** Streaming passthrough test patterns: finish_reason stop streaming, backend HTTP error during streaming, tool call streaming in `chat_completions.rs`. | + +### 5. OpenAI Endpoints Compliance + +| Category | LiteLLM | ours | Port? | Notes | +|----------|---------|------|-------|-------| +| `openai_endpoints_tests/test_openai_batches_endpoint.py` | 8 | 46 | **R** | **Covered.** Batch operations: create, list, cancel, status sync, VK isolation, invalid JSONL. Our `tests/batch_api.rs` has 8 proxy integration tests plus 38 batch engine unit tests covering queue operations, file storage, job lifecycle, and validation. | +| `openai_endpoints_tests/test_openai_files_endpoints.py` | 2 | 0 | **S** | File upload endpoints -- we delegate file handling to backend | +| `openai_endpoints_tests/test_openai_fine_tuning.py` | 1 | 0 | **S** | Fine tuning passthrough -- not supported | +| `openai_endpoints_tests/test_e2e_openai_responses_api.py` | 8 | 6 | **R** | **Covered.** Responses API mock tests: non-streaming, message content preservation, error handling. Plus 2 ignored live tests for real API verification. | + +### 6. Batch Processing + +| Category | LiteLLM | ours | Port? | Notes | +|----------|---------|------|-------|-------| +| `batches_tests/` (9 files, 45 tests) | 45 | 38 | **R** | We cover batch in `batch_engine/src/queue/sqlite/tests.rs` and `proxy/tests/batch_api.rs`. Litellm adds: custom pricing, bedrock batch, hosted vLLM batch, rate limits, logging | +| `batches_tests/test_batch_custom_pricing.py` | ~5 | 0 | **P** | Custom pricing for batch jobs | +| `batches_tests/test_batch_rate_limits.py` | ~8 | 0 | **S** | Rate limit specific tests -- not directly ported | + +### 7. Router / Load Balancing + +| Category | LiteLLM | ours | Port? | Notes | +|----------|---------|------|-------|-------| +| `router_unit_tests/` (19 files, 236 tests) | 236 | 0 | **S** | Extensive router tests: helper utils, budget limiter, cooldown handling, cost calculator, failovers, pattern matching, retries. These are Python-specific and deeply coupled to LiteLLM's Router class | +| `local_testing/test_router.py` | 87 | 0 | **S** | Core router tests | +| `local_testing/test_router_fallbacks.py` | 57 | 0 | **S** | Router fallback tests | +| `local_testing/test_router_retries.py` | 33 | 0 | **S** | Router retry tests | + +### 8. Proxy Behavior / Management + +| Category | LiteLLM | ours | Port? | Notes | +|----------|---------|------|-------|-------| +| `proxy_behavior/management/` (46 files, 166 tests) | 166 | 0 | **S** | Extensive team/key/end-user management tests. LiteLLM's proxy management is more feature-rich (teams, organizations, SSO, SCIM). Our admin API is simpler | + +### 9. Error Handling + +| Category | LiteLLM | ours | Port? | Notes | +|----------|---------|------|-------|-------| +| `local_testing/test_exceptions.py` | 50 | 0 | **R** | Exception handling tests. Different architecture | +| `proxy_unit_tests/test_proxy_exception_mapping.py` | 7 | 37 | **R** | **Covered.** Our `errors_map.rs` has 21 comprehensive unit tests (status mapping, error type round-trips, fixture deserialization, stream error classification). `error_fixtures.rs` now has 16 tests covering OpenAI-to-Anthropic translation for all major HTTP codes, chunk error classification, status boundaries, and edge cases. | + +### 10. Logging & Observability + +| Category | LiteLLM | ours | Port? | Notes | +|----------|---------|------|-------|-------| +| `logging_callback_tests/` (35 files, 272 tests) | 272 | 0 | **S** | Callback/logging integration tests for 20+ platforms. Not directly portable; we use `tracing` + optional OTEL | +| `otel_tests/` (11 files, 41 tests) | 41 | 0 | **S** | OpenTelemetry-specific tests. Python OTel SDK tests; our OTEL is behind `--features otel` and uses `opentelemetry-rust` | + +### 11. Guardrails + +| Category | LiteLLM | ours | Port? | Notes | +|----------|---------|------|-------|-------| +| `guardrails_tests/` (18 files, 189 tests) | 189 | 37 | **R** | LiteLLM has extensive guardrail integrations (Azure, OpenAI, Bedrock, custom hooks). We have `proxy/src/tools/guardrails/tests.rs` for forge-guardrails | +| `proxy_unit_tests/test_proxy_setting_guardrails.py` | 1 | 0 | **S** | Guardrail config tests | + +### 12. End-to-End Tests + +| Category | LiteLLM | ours | Port? | Notes | +|----------|---------|------|-------|-------| +| `e2e/llm_translation/` (17 files, 37 tests) | 37 | 9 | **P** | **Key gap:** E2E tests for: audio speech, cache control, chat completions regression, custom pricing, DeepSeek reasoning, embeddings, image generation, messages, OCR, passthrough, rerank, responses, Vertex passthrough. We have 9 live tests (live_api.rs, live_responses.rs) | +| `e2e/llm_translation/test_passthrough_e2e.py` | 6 | 0 | **P** | Passthrough cost injection tests | +| `e2e/llm_translation/test_cache_control.py` | 2 | 0 | **P** | Cache control E2E tests | +| `e2e/llm_translation/test_deepseek_reasoning_e2e.py` | 3 | 0 | **P** | DeepSeek reasoning E2E tests | +| `proxy_e2e_anthropic_messages_tests/` (2 files) | 4 | 0 | **P** | Anthropic E2E passthrough: beta headers, Claude agent SDK | + +### 13. Cost Management + +| Category | LiteLLM | ours | Port? | Notes | +|----------|---------|------|-------|-------| +| `spend_tracking_tests/` (2 files, 13 tests) | 13 | 20 | **R** | We cover cost more thoroughly: model pricing DB, spend accumulation, budget enforcement | +| `proxy_unit_tests/test_proxy_server_spend.py` | 1 | 0 | **S** | Spend endpoint tests | +| `proxy_unit_tests/test_check_batch_cost.py` | 17 | 0 | **R** | Batch cost checking | + +### 14. Provider-Specific (not portable) + +| Category | LiteLLM | Port? | Reason | +|----------|---------|-------|--------| +| `test_litellm/llms/anthropic/` (8 files) | ~300 | **S** | Python SDK tests; we translate Anthropic to OpenAI | +| `test_litellm/llms/bedrock/` (9 files) | ~200 | **S** | Bedrock SDK tests; we have separate LLM backend | +| `test_litellm/llms/vertex_ai/` (16 files) | ~300 | **S** | Vertex-specific tests | +| `test_litellm/llms/openai/` (9 files) | ~100 | **S** | OpenAI SDK tests | +| `test_litellm/llms/azure/` (5 files) | ~100 | **S** | Azure-specific | +| `test_litellm/llms/gemini/` (5 files) | ~60 | **S** | Gemini-specific | +| `test_litellm/llms/mistral/` (2 files) | ~20 | **S** | Mistral-specific | +| `test_litellm/llms/databricks/` (5 files) | ~40 | **S** | Databricks-specific | + +All provider-specific LLM SDK tests are **not portable** -- they test Python SDK wrappers for each provider. +Our architecture translates to OpenAI-compatible format and delegates to the backend. The translation +itself is tested in our translator crate. + +### 15. Not Applicable (no equivalent in anyllm) + +| Litellm Category | Tests | Reason | +|-----------------|-------|--------| +| `agent_tests/` | 11 | Agent SDK tests | +| `audio_tests/` | 24 | Speech/transcription -- different abstraction | +| `image_gen_tests/` | 57 | Image generation API tests | +| `mcp_tests/` | 147 | MCP server integration tests | +| `search_tests/` | 61 | Search/vector database tests | +| `vector_store_tests/` | 40 | Vector store management tests | +| `proxy_admin_ui_tests/` | 33 | Admin UI Selenium tests (we use SPA) | +| `documentation_tests/` | 3 | Doc validation tests | +| `benchmarks/` | 5 | Benchmark tests | +| `load_tests/` | 8 | Load/performance tests | +| `local_testing/test_completion_cost.py` | 187 | Cost calculation tests | +| `local_testing/test_caching.py` | 93 | Caching tests | + +## Portability Summary + +| Priority | Category | LiteLLM count | Port to Rust | Effort | Impact | +|----------|----------|--------------|--------------|--------|--------| +| **P0** | Anthropic passthrough E2E (streaming, thinking, caching, fallbacks) | 30 | ~12 Rust tests | Medium | High -- core protocol handling | +| **P1** | Token counting (multi-provider) | 24 | ~6 additional Rust tests | Low | Medium -- correctness | +| **P1** | Route dispatch edge cases (base URLs, injection) | 11 | ~5 Rust tests | Low | Medium -- security | +| **P2** | Prompt caching credential/token accounting | 7 | ~4 Rust tests | Low | Medium | +| **P2** | E2E cost injection for streaming | 6 | ~3 Rust tests | Low | Medium | +| **P3** | Optional param validation per provider | 78 | ~5 Rust tests | Medium | Low -- test structure diff | +| **P3** | Batch API E2E (terminal state sync, custom pricing) | 8 | ~2 Rust tests | Low | Low | +| ~~**R**~~ | Error mapping (provider error format handling) | 7 | Done (20 tests) | -- | -- | +| ~~**R**~~ | DeepSeek/Qwen reasoning content E2E | 10 | Done (37 tests) | -- | -- | +| ~~**R**~~ | Streaming test patterns | 4 | Done (6 tests) | -- | -- | +| ~~**R**~~ | OpenAI Responses API (non-live mock tests) | 8 | Done (6 mock + 2 live) | -- | -- | + +## Priority Rationale + +**P0 (Must port):** Core protocol behavior correctness. Litellm extensively tests Anthropic API passthrough with +streaming, thinking blocks, cache control, and error responses. These directly validate the same pipeline we +support. Missing coverage risks regressions. + +**P1 (Should port):** Important correctness and error-handling areas. Token counting accuracy affects cost +tracking. Error mapping affects DX when providers return errors. Route dispatch edge cases affect security. + +**P2 (Nice to port):** Streaming reliability and specific feature tests. + +**P3 (Low):** Useful but lower impact. Optional param validation is structurally different in Rust +(serde-driven vs Python dict manipulation). Batch E2E needs live API key infrastructure. + +## Test Patterns to Study + +Some litellm test patterns worth referencing when writing new tests: + +### Base class pattern for multi-provider coverage + +`llm_translation/base_llm_unit_tests.py` defines `BaseLLMChatTest` with abstract tests that each +provider-specific test class inherits. Rust can approximate this with traits or test macros. + +### Fixture-based golden tests + +Litellm's VCR recording infrastructure (`_vcr_conftest_common.py`, `_openai_record_replay_proxy.py`) +records real API responses as fixtures. We use JSON fixture files in `fixtures/anthropic/` and +`fixtures/openai/`. Our approach is lighter weight and more deterministic, but adding live-recorded +fixtures for edge cases (e.g., real thinking blocks, real streaming sequences) would improve coverage. + +### Error injection patterns + +`test_anthropic_messages_passthrough.py` uses `test_anthropic_messages_streaming_with_bad_request` and +`test_anthropic_messages_router_streaming_with_bad_request` to test error handling during streaming. +We could add similar patterns using mocked HTTP handlers. + +## Test Infrastructure Comparison + +| Aspect | LiteLLM | anyllm-proxy | +|--------|---------|--------------| +| Mocking | `unittest.mock` (MagicMock, AsyncMock, patch) | Hand-rolled mock servers, fixture files | +| Live testing | VCR-recorded with cassette replay | `--ignored --test-threads=1` for live tests | +| Fixtures | `conftest.py` fixtures | JSON files in `fixtures/`, `#[test_case]` | +| Parameterization | `@pytest.mark.parametrize` | `rstest` + `#[test_case]` | +| Async | `pytest-asyncio` | `tokio::test` | +| Config | `conftest.py` + env vars | `test_config` helper setup | + +## What We Already Cover Well + +These areas have strong or sufficient coverage and do not need porting: + +1. **Message format mapping** (Anthropic <-> OpenAI) -- 491 translator tests cover the core mapping +2. **Streaming state machine** -- 68 streaming tests cover chunk assembly, backpressure, finish reasons +3. **Tool calling** -- tools_map tests cover function calls, tool_choice, parallel tool calls +4. **Gemini format mapping** -- gemini_message_map + gemini_streaming_map tests +5. **Config parsing** -- 147 config tests across all formats (simple, multi, litellm, model_router) +6. **Cost tracking** -- 20 cost tests cover model pricing, budget enforcement +7. **Caching** -- 29 cache tests cover in-memory, Redis, semantic +8. **Auth** -- 30 auth tests cover key generation, RBAC, validation +9. **Virtual keys** -- integration tests for key CRUD, rate limiting +10. **Batch engine** -- 38 tests for queue operations, job lifecycle +11. **Fallback** -- 11 tests for fallback chain config, error classification, endpoint selection +12. **Thinking/repair** -- 17 tests for thinking block repair during streaming +13. **DeepSeek/Qwen reasoning content** -- ~37 tests across translator and proxy covering reasoning_content <-> thinking block mapping in both directions, streaming lifecycle, tool call interleaving, and effective_text fallback +14. **Error mapping** -- 18 error_fixtures.rs tests + 2 chat_completions validation tests covering all Anthropic error types, status mapping, fixture-based translation, stream error classification, Azure auth, and context window exceeded +15. **Token counting** -- 10 unit tests for count_request_tokens_sync + 4 integration tests covering basic counting, tools, system blocks, thinking blocks, and tool results +16. **OpenAI streaming passthrough** -- 6 tests in `chat_completions.rs` covering finish_reason stop, tool call streaming, backend HTTP errors during streaming, and malformed input validation +17. **OpenAI error format** -- Tests for missing model, missing messages, invalid JSON returning proper OpenAI error shapes +18. **OpenAI Responses API** -- 4 mock-based integration tests plus 2 ignored live tests covering non-streaming, content preservation, and error handling + +## How to Run the Parity Assessment + +```bash +# Run our full test suite +cargo test + +# Run a specific test category +cargo test -p anyllm_translate +cargo test -p anyllm_proxy + +# Run live API tests (needs API key) +OPENAI_API_KEY=sk-... cargo test --test live_api -- --ignored --test-threads=1 + +# Run golden fixture tests +cargo test -p anyllm_translate --test golden_fixtures +``` + +For litellm tests (reference only): +```bash +cd /tmp/litellm +python3 -m pytest tests/llm_translation/ -x -v --timeout 60 +```