mirror of
https://github.com/whit3rabbit/anyllm-proxy.git
synced 2026-09-22 08:00:51 +00:00
Merge branch 'feat/model-routing-ux'
This commit is contained in:
@@ -101,7 +101,10 @@ OPENAI_API_KEY=sk-... cargo run -p anyllm_proxy
|
||||
- `OTEL_EXPORTER_OTLP_ENDPOINT`: OTLP collector endpoint (default: `http://localhost:4318`). Only effective when built with `--features otel`.
|
||||
- `OTEL_SERVICE_NAME`: Service name for exported traces. Only effective when built with `--features otel`.
|
||||
- `OTEL_TRACES_SAMPLER`: Sampling strategy (default: `parentbased_always_on`). Only effective when built with `--features otel`.
|
||||
- `PROXY_CONFIG`: Path to config file. TOML for multi-backend config, or `.yaml`/`.yml` for LiteLLM-compatible config with model_list routing.
|
||||
- `PROXY_CONFIG`: Path to config file. Three formats accepted:
|
||||
- **Simple YAML** (`.yaml`/`.yml` with top-level `models:` key): ergonomic native format; provider API keys from env vars; supports routing strategies and string shorthand (e.g. `- openai/gpt-4o`).
|
||||
- **LiteLLM YAML** (`.yaml`/`.yml` with top-level `model_list:` key): LiteLLM-compatible format with `litellm_params:` nesting.
|
||||
- **TOML** (any other extension): multi-backend TOML config.
|
||||
- `IP_ALLOWLIST`: Comma-separated CIDR ranges for IP allowlisting (e.g., `192.168.1.0/24,10.0.0.0/8`). Bare IPs also accepted. When set, only matching IPs can access the proxy.
|
||||
- `TRUST_PROXY_HEADERS`: Set to `true` or `1` to use `X-Forwarded-For` header for client IP when behind a reverse proxy. Only effective when `IP_ALLOWLIST` is set.
|
||||
- `WEBHOOK_URLS`: Comma-separated webhook URLs for request completion notifications. Fire-and-forget HTTP POST with `RequestLogEntry` JSON payload.
|
||||
@@ -205,6 +208,46 @@ Client (Anthropic format) -> proxy (axum)
|
||||
- Virtual key CRUD integration tests are in `crates/proxy/tests/virtual_keys.rs`. They use a shared `OnceLock<DashMap>` to avoid fighting over the global `set_virtual_keys` OnceLock.
|
||||
- The `PROXY_OPEN_RELAY=true` env var enables dev mode (any non-empty key accepted). Without it and without `PROXY_API_KEYS`, the proxy rejects all requests.
|
||||
|
||||
## Simple Config Format
|
||||
|
||||
Ergonomic native alternative to the LiteLLM format. Activated when the config file has a top-level `models:` key. Set via `PROXY_CONFIG=/path/to/anyllm.yaml`.
|
||||
|
||||
```yaml
|
||||
# anyllm.yaml
|
||||
routing_strategy: latency-based # round-robin (default) | least-busy | latency-based | weighted | cost-based
|
||||
listen_port: 3000 # optional
|
||||
log_bodies: false # optional
|
||||
|
||||
models:
|
||||
# String shorthand: bare model name defaults to openai
|
||||
- gpt-4o
|
||||
# String shorthand with provider prefix
|
||||
- openai/gpt-4o-mini
|
||||
- anthropic/claude-3-5-sonnet-20241022
|
||||
# Full form: virtual name, actual model, weight, rate limits
|
||||
- name: smart # virtual name clients send in requests
|
||||
model: gpt-4o
|
||||
provider: openai
|
||||
weight: 3
|
||||
rpm: 1000
|
||||
tpm: 500000
|
||||
- name: smart # second deployment for "smart" (round-robin / failover)
|
||||
model: claude-3-5-sonnet-20241022
|
||||
provider: anthropic
|
||||
weight: 1
|
||||
```
|
||||
|
||||
Provider API key defaults (used when `api_key` is not specified in the entry):
|
||||
|
||||
| provider | env var |
|
||||
|------------|-----------------------------------------------|
|
||||
| openai | `OPENAI_API_KEY` |
|
||||
| anthropic | `ANTHROPIC_API_KEY` |
|
||||
| gemini | `GEMINI_API_KEY` |
|
||||
| vertex | `VERTEX_API_KEY` or `GOOGLE_ACCESS_TOKEN` |
|
||||
| azure | `AZURE_OPENAI_API_KEY` |
|
||||
| bedrock | `AWS_ACCESS_KEY_ID` + `AWS_SECRET_ACCESS_KEY` |
|
||||
|
||||
## References
|
||||
|
||||
- OpenAI API spec: https://github.com/openai/openai-openapi/blob/manual_spec/openapi.yaml (very large, ~70k+ lines). See https://simonwillison.net/2024/Dec/22/openai-openapi/ for context on the spec's size and structure. Do not attempt to load the full spec into context; reference specific sections as needed.
|
||||
|
||||
@@ -84,3 +84,4 @@ optional = true
|
||||
pretty_assertions = "1"
|
||||
reqwest = { version = "0.12", default-features = false, features = ["json", "native-tls", "multipart"] }
|
||||
tokio = { version = "1", features = ["full"] }
|
||||
tempfile = "3"
|
||||
|
||||
@@ -75,7 +75,7 @@ struct RouterSettings {
|
||||
}
|
||||
|
||||
/// Map LiteLLM routing_strategy string to our enum.
|
||||
fn parse_routing_strategy(s: &str) -> RoutingStrategy {
|
||||
pub(crate) fn parse_routing_strategy_str(s: &str) -> RoutingStrategy {
|
||||
match s.to_ascii_lowercase().replace('_', "-").as_str() {
|
||||
"simple-shuffle" | "round-robin" => RoutingStrategy::RoundRobin,
|
||||
"least-busy" => RoutingStrategy::LeastBusy,
|
||||
@@ -286,7 +286,7 @@ pub fn parse_litellm_yaml(yaml: &str) -> LiteLLMParsed {
|
||||
.router_settings
|
||||
.as_ref()
|
||||
.and_then(|rs| rs.routing_strategy.as_deref())
|
||||
.map(parse_routing_strategy)
|
||||
.map(parse_routing_strategy_str)
|
||||
.unwrap_or_default();
|
||||
|
||||
if strategy != RoutingStrategy::RoundRobin {
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
pub mod env_aliases;
|
||||
pub mod litellm;
|
||||
pub mod model_router;
|
||||
pub mod simple;
|
||||
mod tls;
|
||||
mod url_validation;
|
||||
|
||||
@@ -506,18 +507,44 @@ impl MultiConfig {
|
||||
/// Load configuration.
|
||||
///
|
||||
/// Detection order:
|
||||
/// 1. `PROXY_CONFIG` with `.yaml`/`.yml` extension: parse as LiteLLM config
|
||||
/// 1. `PROXY_CONFIG` with `.yaml`/`.yml` extension:
|
||||
/// - If root `models:` key is present: simple native format (`simple::parse_simple_yaml`)
|
||||
/// - Otherwise (`model_list:` key): LiteLLM-compatible format (`litellm::parse_litellm_yaml`)
|
||||
/// 2. `PROXY_CONFIG` with any other extension: parse as TOML
|
||||
/// 3. No `PROXY_CONFIG`: env-var-based single-backend config
|
||||
///
|
||||
/// The model router is only set for LiteLLM configs (model_list routing).
|
||||
/// The model router is set for both YAML config formats (simple and LiteLLM).
|
||||
/// `litellm_master_key` is returned (not applied) so the caller can
|
||||
/// consolidate all `set_var` calls into a single pre-runtime block.
|
||||
pub fn load() -> LoadResult {
|
||||
if let Ok(path) = std::env::var("PROXY_CONFIG") {
|
||||
if path.ends_with(".yaml") || path.ends_with(".yml") {
|
||||
let yaml = std::fs::read_to_string(&path)
|
||||
.unwrap_or_else(|e| panic!("failed to read LiteLLM config '{path}': {e}"));
|
||||
.unwrap_or_else(|e| panic!("failed to read config '{path}': {e}"));
|
||||
|
||||
// Detect format: "models:" key = simple native format, "model_list:" = LiteLLM.
|
||||
let probe: serde_yaml::Value = serde_yaml::from_str(&yaml)
|
||||
.unwrap_or_else(|e| panic!("invalid YAML in '{path}': {e}"));
|
||||
|
||||
if probe.get("models").is_some() {
|
||||
// Simple native format.
|
||||
let parsed = simple::parse_simple_yaml(&yaml);
|
||||
return LoadResult {
|
||||
multi_config: parsed.multi_config,
|
||||
model_router: Some(Arc::new(std::sync::RwLock::new(parsed.router))),
|
||||
litellm_master_key: None,
|
||||
};
|
||||
}
|
||||
|
||||
// LiteLLM format requires model_list: key.
|
||||
if probe.get("model_list").is_none() {
|
||||
panic!(
|
||||
"config file '{path}' must contain either a top-level 'models:' key \
|
||||
(simple format) or 'model_list:' key (LiteLLM format)"
|
||||
);
|
||||
}
|
||||
|
||||
// LiteLLM format (model_list: + litellm_params:).
|
||||
let parsed = litellm::parse_litellm_yaml(&yaml);
|
||||
|
||||
// Wire up webhook callbacks and named integrations from litellm_settings.callbacks.
|
||||
|
||||
@@ -0,0 +1,660 @@
|
||||
//! Simple native YAML config format for anyllm-proxy.
|
||||
//!
|
||||
//! Activated when the config file contains a top-level `models:` key
|
||||
//! (as opposed to LiteLLM's `model_list:`).
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use indexmap::IndexMap;
|
||||
use serde::Deserialize;
|
||||
|
||||
use super::litellm::parse_routing_strategy_str;
|
||||
use super::model_router::{Deployment, ModelRouter, RoutingStrategy};
|
||||
use super::{
|
||||
validate_base_url, BackendAuth, BackendConfig, BackendKind, ModelMapping, MultiConfig,
|
||||
OpenAIApiFormat, TlsConfig,
|
||||
};
|
||||
|
||||
/// Top-level simple config document.
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct SimpleConfig {
|
||||
/// Routing strategy for all models. Case-insensitive.
|
||||
/// Accepted values: round-robin, least-busy, latency-based, weighted, cost-based.
|
||||
/// Default: round-robin.
|
||||
#[serde(default)]
|
||||
pub routing_strategy: Option<String>,
|
||||
/// Proxy listen port. Default: 3000.
|
||||
#[serde(default)]
|
||||
pub listen_port: Option<u16>,
|
||||
/// Log request/response bodies at debug level. Default: false.
|
||||
#[serde(default)]
|
||||
pub log_bodies: Option<bool>,
|
||||
/// List of model deployments.
|
||||
#[serde(default)]
|
||||
pub models: Vec<SimpleModelEntry>,
|
||||
}
|
||||
|
||||
/// A model entry: either a string shorthand or a full struct.
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum SimpleModelEntry {
|
||||
/// String shorthand: "model-name" or "provider/model-name".
|
||||
Shorthand(String),
|
||||
/// Full form with all fields. Boxed to reduce enum size.
|
||||
Full(Box<SimpleModelFull>),
|
||||
}
|
||||
|
||||
/// Full model entry with all optional fields.
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct SimpleModelFull {
|
||||
/// Virtual model name clients send in requests. Defaults to `model` if omitted.
|
||||
#[serde(default)]
|
||||
pub name: Option<String>,
|
||||
/// Actual model name forwarded to the backend.
|
||||
pub model: String,
|
||||
/// Backend provider. Default: "openai".
|
||||
#[serde(default)]
|
||||
pub provider: Option<String>,
|
||||
/// Static weight for weighted routing. Default: 1.
|
||||
#[serde(default)]
|
||||
pub weight: Option<u32>,
|
||||
/// Per-deployment requests-per-minute limit.
|
||||
#[serde(default)]
|
||||
pub rpm: Option<u32>,
|
||||
/// Per-deployment tokens-per-minute limit.
|
||||
#[serde(default)]
|
||||
pub tpm: Option<u64>,
|
||||
/// API key override. When absent, falls back to the standard env var for the provider.
|
||||
#[serde(default)]
|
||||
pub api_key: Option<String>,
|
||||
/// Base URL override. When absent, uses the provider default.
|
||||
#[serde(default)]
|
||||
pub api_base: Option<String>,
|
||||
// Azure-specific
|
||||
#[serde(default)]
|
||||
pub deployment: Option<String>,
|
||||
#[serde(default)]
|
||||
pub api_version: Option<String>,
|
||||
// Vertex-specific
|
||||
#[serde(default)]
|
||||
pub project: Option<String>,
|
||||
#[serde(default)]
|
||||
pub region: Option<String>,
|
||||
// Bedrock-specific
|
||||
#[serde(default)]
|
||||
pub aws_region: Option<String>,
|
||||
#[serde(default)]
|
||||
pub aws_access_key_id: Option<String>,
|
||||
#[serde(default)]
|
||||
pub aws_secret_access_key: Option<String>,
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Parsed result + public parser
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Result from parsing a simple YAML config file.
|
||||
pub struct SimpleParsed {
|
||||
pub multi_config: MultiConfig,
|
||||
pub router: ModelRouter,
|
||||
}
|
||||
|
||||
/// Parse a simple YAML config string and produce a `MultiConfig + ModelRouter`.
|
||||
///
|
||||
/// # Panics
|
||||
/// On invalid YAML, empty model list, or unresolvable required values.
|
||||
pub fn parse_simple_yaml(yaml: &str) -> SimpleParsed {
|
||||
let config: SimpleConfig =
|
||||
serde_yaml::from_str(yaml).unwrap_or_else(|e| panic!("invalid simple config YAML: {e}"));
|
||||
|
||||
if config.models.is_empty() {
|
||||
panic!("simple config must define at least one model");
|
||||
}
|
||||
|
||||
let listen_port = config
|
||||
.listen_port
|
||||
.or_else(|| {
|
||||
std::env::var("LISTEN_PORT")
|
||||
.ok()
|
||||
.and_then(|v| v.parse().ok())
|
||||
})
|
||||
.unwrap_or(3000);
|
||||
|
||||
let log_bodies = config.log_bodies.unwrap_or_else(|| {
|
||||
std::env::var("LOG_BODIES")
|
||||
.map(|v| v == "true" || v == "1")
|
||||
.unwrap_or(false)
|
||||
});
|
||||
|
||||
let tls = TlsConfig::from_env();
|
||||
|
||||
#[derive(Hash, PartialEq, Eq)]
|
||||
struct BackendKey {
|
||||
kind: String,
|
||||
base_url: String,
|
||||
api_key_hash: u64,
|
||||
}
|
||||
|
||||
fn hash_str(s: &str) -> u64 {
|
||||
use std::hash::{Hash, Hasher};
|
||||
let mut h = std::collections::hash_map::DefaultHasher::new();
|
||||
s.hash(&mut h);
|
||||
h.finish()
|
||||
}
|
||||
|
||||
struct DepSpec {
|
||||
backend_name: String,
|
||||
actual_model: String,
|
||||
rpm: Option<u32>,
|
||||
tpm: Option<u64>,
|
||||
weight: u32,
|
||||
}
|
||||
|
||||
let mut backend_map: HashMap<BackendKey, (String, BackendConfig)> = HashMap::new();
|
||||
let mut backend_counter = 0u32;
|
||||
let mut model_deployments: HashMap<String, Vec<DepSpec>> = HashMap::new();
|
||||
|
||||
for entry in &config.models {
|
||||
let norm = normalize_entry(entry);
|
||||
let kind = parse_kind(&norm.provider);
|
||||
let api_key = norm
|
||||
.api_key
|
||||
.clone()
|
||||
.unwrap_or_else(|| default_api_key_for_provider(&norm.provider, &kind));
|
||||
let base_url = if norm.api_base.is_some() && kind == BackendKind::AzureOpenAI {
|
||||
// Azure: build the full deployment URL from api_base + deployment + api_version
|
||||
default_base_url(&kind, &norm)
|
||||
} else {
|
||||
norm.api_base
|
||||
.clone()
|
||||
.unwrap_or_else(|| default_base_url(&kind, &norm))
|
||||
};
|
||||
|
||||
if kind != BackendKind::Bedrock {
|
||||
if let Err(e) = validate_base_url(&base_url) {
|
||||
panic!("model '{}' base_url rejected: {e}", norm.virtual_name);
|
||||
}
|
||||
}
|
||||
|
||||
let bk = BackendKey {
|
||||
kind: format!("{kind:?}"),
|
||||
base_url: base_url.clone(),
|
||||
api_key_hash: hash_str(&api_key),
|
||||
};
|
||||
|
||||
let backend_name = if let Some((name, _)) = backend_map.get(&bk) {
|
||||
name.clone()
|
||||
} else {
|
||||
let name = format!("simple_{backend_counter}");
|
||||
backend_counter += 1;
|
||||
let bc =
|
||||
build_backend_config(&name, &kind, &api_key, &base_url, &norm, &tls, log_bodies);
|
||||
backend_map.insert(bk, (name.clone(), bc));
|
||||
name
|
||||
};
|
||||
|
||||
model_deployments
|
||||
.entry(norm.virtual_name.clone())
|
||||
.or_default()
|
||||
.push(DepSpec {
|
||||
backend_name,
|
||||
actual_model: norm.actual_model.clone(),
|
||||
rpm: norm.rpm,
|
||||
tpm: norm.tpm,
|
||||
weight: norm.weight.unwrap_or(1),
|
||||
});
|
||||
}
|
||||
|
||||
let mut backends = IndexMap::new();
|
||||
for (name, bc) in backend_map.values() {
|
||||
backends.insert(name.clone(), bc.clone());
|
||||
}
|
||||
let default_backend = backends.keys().next().cloned().expect("at least one backend");
|
||||
let multi = MultiConfig {
|
||||
listen_port,
|
||||
log_bodies,
|
||||
default_backend,
|
||||
backends,
|
||||
};
|
||||
|
||||
let strategy: RoutingStrategy = config
|
||||
.routing_strategy
|
||||
.as_deref()
|
||||
.map(parse_routing_strategy_str)
|
||||
.unwrap_or_default();
|
||||
|
||||
let mut routes: HashMap<String, Vec<Arc<Deployment>>> = HashMap::new();
|
||||
for (virtual_name, specs) in model_deployments {
|
||||
let deployments = specs
|
||||
.into_iter()
|
||||
.map(|s| {
|
||||
Arc::new(Deployment::with_weight(
|
||||
s.backend_name,
|
||||
s.actual_model,
|
||||
s.rpm,
|
||||
s.tpm,
|
||||
s.weight,
|
||||
))
|
||||
})
|
||||
.collect();
|
||||
routes.insert(virtual_name, deployments);
|
||||
}
|
||||
let router = ModelRouter::with_strategy(routes, strategy);
|
||||
|
||||
SimpleParsed {
|
||||
multi_config: multi,
|
||||
router,
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Internal helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
struct NormalizedEntry {
|
||||
virtual_name: String,
|
||||
provider: String,
|
||||
actual_model: String,
|
||||
api_key: Option<String>,
|
||||
api_base: Option<String>,
|
||||
weight: Option<u32>,
|
||||
rpm: Option<u32>,
|
||||
tpm: Option<u64>,
|
||||
deployment: Option<String>,
|
||||
api_version: Option<String>,
|
||||
project: Option<String>,
|
||||
region: Option<String>,
|
||||
aws_region: Option<String>,
|
||||
aws_access_key_id: Option<String>,
|
||||
aws_secret_access_key: Option<String>,
|
||||
}
|
||||
|
||||
fn normalize_entry(entry: &SimpleModelEntry) -> NormalizedEntry {
|
||||
match entry {
|
||||
SimpleModelEntry::Shorthand(s) => {
|
||||
let (provider, model) = s
|
||||
.split_once('/')
|
||||
.map(|(p, m)| (p.to_string(), m.to_string()))
|
||||
.unwrap_or_else(|| ("openai".to_string(), s.clone()));
|
||||
NormalizedEntry {
|
||||
virtual_name: model.clone(),
|
||||
provider,
|
||||
actual_model: model,
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
weight: None,
|
||||
rpm: None,
|
||||
tpm: None,
|
||||
deployment: None,
|
||||
api_version: None,
|
||||
project: None,
|
||||
region: None,
|
||||
aws_region: None,
|
||||
aws_access_key_id: None,
|
||||
aws_secret_access_key: None,
|
||||
}
|
||||
}
|
||||
SimpleModelEntry::Full(f) => {
|
||||
let provider = f.provider.clone().unwrap_or_else(|| "openai".to_string());
|
||||
let virtual_name = f.name.clone().unwrap_or_else(|| f.model.clone());
|
||||
NormalizedEntry {
|
||||
virtual_name,
|
||||
provider,
|
||||
actual_model: f.model.clone(),
|
||||
api_key: f.api_key.clone(),
|
||||
api_base: f.api_base.clone(),
|
||||
weight: f.weight,
|
||||
rpm: f.rpm,
|
||||
tpm: f.tpm,
|
||||
deployment: f.deployment.clone(),
|
||||
api_version: f.api_version.clone(),
|
||||
project: f.project.clone(),
|
||||
region: f.region.clone(),
|
||||
aws_region: f.aws_region.clone(),
|
||||
aws_access_key_id: f.aws_access_key_id.clone(),
|
||||
aws_secret_access_key: f.aws_secret_access_key.clone(),
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_kind(provider: &str) -> BackendKind {
|
||||
match provider.to_ascii_lowercase().as_str() {
|
||||
"openai" => BackendKind::OpenAI,
|
||||
"azure" => BackendKind::AzureOpenAI,
|
||||
"vertex_ai" | "vertex" => BackendKind::Vertex,
|
||||
"gemini" => BackendKind::Gemini,
|
||||
"anthropic" => BackendKind::Anthropic,
|
||||
"bedrock" => BackendKind::Bedrock,
|
||||
other => {
|
||||
tracing::warn!(provider = %other, "unknown provider, treating as openai-compatible");
|
||||
BackendKind::OpenAI
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn default_api_key_for_provider(provider: &str, kind: &BackendKind) -> String {
|
||||
let var = match kind {
|
||||
BackendKind::OpenAI => "OPENAI_API_KEY",
|
||||
BackendKind::Anthropic => "ANTHROPIC_API_KEY",
|
||||
BackendKind::Gemini => "GEMINI_API_KEY",
|
||||
BackendKind::Vertex => {
|
||||
return std::env::var("VERTEX_API_KEY")
|
||||
.or_else(|_| std::env::var("GOOGLE_ACCESS_TOKEN"))
|
||||
.unwrap_or_default();
|
||||
}
|
||||
BackendKind::AzureOpenAI => "AZURE_OPENAI_API_KEY",
|
||||
BackendKind::Bedrock => return String::new(),
|
||||
};
|
||||
std::env::var(var).unwrap_or_else(|_| {
|
||||
tracing::warn!(
|
||||
provider = %provider,
|
||||
env_var = %var,
|
||||
"provider API key env var not set; backend calls will likely fail"
|
||||
);
|
||||
String::new()
|
||||
})
|
||||
}
|
||||
|
||||
fn default_base_url(kind: &BackendKind, entry: &NormalizedEntry) -> String {
|
||||
match kind {
|
||||
BackendKind::OpenAI => std::env::var("OPENAI_BASE_URL")
|
||||
.unwrap_or_else(|_| "https://api.openai.com".to_string()),
|
||||
BackendKind::Gemini => {
|
||||
let base = std::env::var("GEMINI_BASE_URL")
|
||||
.unwrap_or_else(|_| "https://generativelanguage.googleapis.com/v1beta".to_string());
|
||||
format!("{base}/openai")
|
||||
}
|
||||
BackendKind::Anthropic => "https://api.anthropic.com".to_string(),
|
||||
BackendKind::Vertex => {
|
||||
let project = entry
|
||||
.project
|
||||
.as_deref()
|
||||
.map(|s| s.to_string())
|
||||
.unwrap_or_else(|| {
|
||||
std::env::var("VERTEX_PROJECT").expect(
|
||||
"project field (or VERTEX_PROJECT env var) required for vertex provider",
|
||||
)
|
||||
});
|
||||
let region = entry
|
||||
.region
|
||||
.as_deref()
|
||||
.map(|s| s.to_string())
|
||||
.unwrap_or_else(|| {
|
||||
std::env::var("VERTEX_REGION").expect(
|
||||
"region field (or VERTEX_REGION env var) required for vertex provider",
|
||||
)
|
||||
});
|
||||
format!(
|
||||
"https://{region}-aiplatform.googleapis.com/v1/projects/{project}/locations/{region}/endpoints/openapi"
|
||||
)
|
||||
}
|
||||
BackendKind::AzureOpenAI => {
|
||||
let endpoint = entry
|
||||
.api_base
|
||||
.as_deref()
|
||||
.map(|s| s.to_string())
|
||||
.unwrap_or_else(|| {
|
||||
std::env::var("AZURE_OPENAI_ENDPOINT").expect(
|
||||
"api_base field (or AZURE_OPENAI_ENDPOINT env var) required for azure provider",
|
||||
)
|
||||
});
|
||||
let dep = entry.deployment.as_deref().unwrap_or("chat");
|
||||
let version = entry.api_version.as_deref().unwrap_or("2024-10-21");
|
||||
format!(
|
||||
"{}/openai/deployments/{dep}/chat/completions?api-version={version}",
|
||||
endpoint.trim_end_matches('/')
|
||||
)
|
||||
}
|
||||
BackendKind::Bedrock => {
|
||||
// Bedrock doesn't use a URL — the region string is stored in base_url
|
||||
// and used by the Bedrock client directly for SigV4 endpoint construction.
|
||||
entry
|
||||
.aws_region
|
||||
.as_deref()
|
||||
.map(|s| s.to_string())
|
||||
.unwrap_or_else(|| {
|
||||
std::env::var("AWS_REGION").unwrap_or_else(|_| "us-east-1".to_string())
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn build_backend_config(
|
||||
name: &str,
|
||||
kind: &BackendKind,
|
||||
api_key: &str,
|
||||
base_url: &str,
|
||||
entry: &NormalizedEntry,
|
||||
tls: &TlsConfig,
|
||||
log_bodies: bool,
|
||||
) -> BackendConfig {
|
||||
let backend_auth = match kind {
|
||||
BackendKind::AzureOpenAI => BackendAuth::AzureApiKey(api_key.to_string()),
|
||||
BackendKind::Gemini | BackendKind::Vertex => BackendAuth::GoogleApiKey(api_key.to_string()),
|
||||
_ => BackendAuth::BearerToken(api_key.to_string()),
|
||||
};
|
||||
|
||||
let bedrock_credentials = if *kind == BackendKind::Bedrock {
|
||||
let access_key = entry
|
||||
.aws_access_key_id
|
||||
.as_deref()
|
||||
.map(|s| s.to_string())
|
||||
.or_else(|| std::env::var("AWS_ACCESS_KEY_ID").ok())
|
||||
.unwrap_or_else(|| {
|
||||
panic!("backend '{name}': aws_access_key_id required for bedrock")
|
||||
});
|
||||
|
||||
let secret_key = entry
|
||||
.aws_secret_access_key
|
||||
.as_deref()
|
||||
.map(|s| s.to_string())
|
||||
.or_else(|| std::env::var("AWS_SECRET_ACCESS_KEY").ok())
|
||||
.unwrap_or_else(|| {
|
||||
panic!("backend '{name}': aws_secret_access_key required for bedrock")
|
||||
});
|
||||
|
||||
Some(aws_credential_types::Credentials::new(
|
||||
access_key,
|
||||
secret_key,
|
||||
None,
|
||||
None,
|
||||
"simple-config",
|
||||
))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
BackendConfig {
|
||||
kind: kind.clone(),
|
||||
api_key: api_key.to_string(),
|
||||
base_url: base_url.to_string(),
|
||||
api_format: OpenAIApiFormat::Chat,
|
||||
model_mapping: ModelMapping {
|
||||
big_model: String::new(),
|
||||
small_model: String::new(),
|
||||
},
|
||||
tls: tls.clone(),
|
||||
backend_auth,
|
||||
log_bodies,
|
||||
omit_stream_options: false,
|
||||
bedrock_credentials,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn simple_config_roundtrip_string_shorthand() {
|
||||
let yaml = r#"
|
||||
models:
|
||||
- gpt-4o
|
||||
- openai/gpt-4o-mini
|
||||
- anthropic/claude-3-5-sonnet-20241022
|
||||
"#;
|
||||
let cfg: SimpleConfig = serde_yaml::from_str(yaml).unwrap();
|
||||
assert_eq!(cfg.models.len(), 3);
|
||||
match &cfg.models[0] {
|
||||
SimpleModelEntry::Shorthand(s) => assert_eq!(s, "gpt-4o"),
|
||||
SimpleModelEntry::Full(_) => panic!("expected shorthand"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn simple_config_roundtrip_full_entry() {
|
||||
let yaml = r#"
|
||||
routing_strategy: weighted
|
||||
models:
|
||||
- name: smart
|
||||
model: gpt-4o
|
||||
provider: openai
|
||||
weight: 3
|
||||
rpm: 1000
|
||||
"#;
|
||||
let cfg: SimpleConfig = serde_yaml::from_str(yaml).unwrap();
|
||||
assert_eq!(cfg.routing_strategy.as_deref(), Some("weighted"));
|
||||
assert_eq!(cfg.models.len(), 1);
|
||||
match &cfg.models[0] {
|
||||
SimpleModelEntry::Full(f) => {
|
||||
assert_eq!(f.name.as_deref(), Some("smart"));
|
||||
assert_eq!(f.model, "gpt-4o");
|
||||
assert_eq!(f.provider.as_deref(), Some("openai"));
|
||||
assert_eq!(f.weight, Some(3));
|
||||
assert_eq!(f.rpm, Some(1000));
|
||||
}
|
||||
SimpleModelEntry::Shorthand(_) => panic!("expected full entry"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn simple_config_mixed_entries() {
|
||||
let yaml = r#"
|
||||
models:
|
||||
- gpt-4o
|
||||
- name: my-model
|
||||
model: claude-3-5-sonnet-20241022
|
||||
provider: anthropic
|
||||
"#;
|
||||
let cfg: SimpleConfig = serde_yaml::from_str(yaml).unwrap();
|
||||
assert_eq!(cfg.models.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_single_openai_model() {
|
||||
unsafe { std::env::set_var("OPENAI_API_KEY", "sk-test") };
|
||||
let yaml = r#"
|
||||
models:
|
||||
- gpt-4o
|
||||
"#;
|
||||
let parsed = parse_simple_yaml(yaml);
|
||||
assert_eq!(parsed.multi_config.backends.len(), 1);
|
||||
assert!(parsed.router.has_model("gpt-4o"));
|
||||
let routed = parsed.router.route("gpt-4o").unwrap();
|
||||
assert_eq!(routed.actual_model, "gpt-4o");
|
||||
unsafe { std::env::remove_var("OPENAI_API_KEY") };
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_provider_slash_model_shorthand() {
|
||||
unsafe {
|
||||
std::env::set_var("OPENAI_API_KEY", "sk-openai");
|
||||
std::env::set_var("ANTHROPIC_API_KEY", "sk-anthropic");
|
||||
};
|
||||
let yaml = r#"
|
||||
models:
|
||||
- openai/gpt-4o
|
||||
- anthropic/claude-3-5-sonnet-20241022
|
||||
"#;
|
||||
let parsed = parse_simple_yaml(yaml);
|
||||
assert_eq!(parsed.multi_config.backends.len(), 2);
|
||||
assert!(parsed.router.has_model("gpt-4o"));
|
||||
assert!(parsed.router.has_model("claude-3-5-sonnet-20241022"));
|
||||
unsafe {
|
||||
std::env::remove_var("OPENAI_API_KEY");
|
||||
std::env::remove_var("ANTHROPIC_API_KEY");
|
||||
};
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_full_entry_with_virtual_name() {
|
||||
unsafe { std::env::set_var("OPENAI_API_KEY", "sk-test") };
|
||||
let yaml = r#"
|
||||
models:
|
||||
- name: smart
|
||||
model: gpt-4o
|
||||
provider: openai
|
||||
weight: 3
|
||||
"#;
|
||||
let parsed = parse_simple_yaml(yaml);
|
||||
assert!(parsed.router.has_model("smart"));
|
||||
assert!(!parsed.router.has_model("gpt-4o"));
|
||||
let routed = parsed.router.route("smart").unwrap();
|
||||
assert_eq!(routed.actual_model, "gpt-4o");
|
||||
unsafe { std::env::remove_var("OPENAI_API_KEY") };
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_routing_strategy_latency() {
|
||||
unsafe { std::env::set_var("OPENAI_API_KEY", "sk-test") };
|
||||
let yaml = r#"
|
||||
routing_strategy: latency-based
|
||||
models:
|
||||
- gpt-4o
|
||||
"#;
|
||||
let parsed = parse_simple_yaml(yaml);
|
||||
assert_eq!(
|
||||
parsed.router.strategy(),
|
||||
crate::config::model_router::RoutingStrategy::LatencyBased
|
||||
);
|
||||
unsafe { std::env::remove_var("OPENAI_API_KEY") };
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_weighted_two_deployments_same_virtual_name() {
|
||||
unsafe { std::env::set_var("OPENAI_API_KEY", "sk-test") };
|
||||
let yaml = r#"
|
||||
routing_strategy: weighted
|
||||
models:
|
||||
- name: smart
|
||||
model: gpt-4o
|
||||
provider: openai
|
||||
weight: 3
|
||||
- name: smart
|
||||
model: gpt-4o-mini
|
||||
provider: openai
|
||||
weight: 1
|
||||
"#;
|
||||
let parsed = parse_simple_yaml(yaml);
|
||||
assert!(parsed.router.has_model("smart"));
|
||||
let list = parsed.router.list_models();
|
||||
let (_, count) = list.iter().find(|(n, _)| *n == "smart").unwrap();
|
||||
assert_eq!(*count, 2);
|
||||
unsafe { std::env::remove_var("OPENAI_API_KEY") };
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_api_key_inline_overrides_env() {
|
||||
unsafe { std::env::set_var("OPENAI_API_KEY", "sk-from-env") };
|
||||
let yaml = r#"
|
||||
models:
|
||||
- name: my-model
|
||||
model: gpt-4o
|
||||
provider: openai
|
||||
api_key: sk-inline-key
|
||||
"#;
|
||||
let parsed = parse_simple_yaml(yaml);
|
||||
let bc = parsed.multi_config.backends.values().next().unwrap();
|
||||
assert_eq!(bc.api_key, "sk-inline-key");
|
||||
unsafe { std::env::remove_var("OPENAI_API_KEY") };
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "must define at least one model")]
|
||||
fn parse_empty_models_panics() {
|
||||
let yaml = "models: []\n";
|
||||
parse_simple_yaml(yaml);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
//! Integration tests: MultiConfig::load() dispatches simple YAML format.
|
||||
|
||||
use anyllm_proxy::config::MultiConfig;
|
||||
|
||||
#[test]
|
||||
fn load_dispatches_simple_format_by_models_key() {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let path = dir.path().join("routing.yaml");
|
||||
std::fs::write(
|
||||
&path,
|
||||
r#"
|
||||
models:
|
||||
- openai/gpt-4o
|
||||
- openai/gpt-4o-mini
|
||||
routing_strategy: least-busy
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
// SAFETY: test-only env var mutation
|
||||
unsafe { std::env::set_var("PROXY_CONFIG", path.to_str().unwrap()) };
|
||||
unsafe { std::env::set_var("OPENAI_API_KEY", "sk-test") };
|
||||
|
||||
let result = MultiConfig::load();
|
||||
assert!(result.model_router.is_some(), "simple format must produce a model router");
|
||||
|
||||
let router_arc = result.model_router.unwrap();
|
||||
let router = router_arc.read().unwrap();
|
||||
assert!(router.has_model("gpt-4o"));
|
||||
assert!(router.has_model("gpt-4o-mini"));
|
||||
assert_eq!(
|
||||
router.strategy(),
|
||||
anyllm_proxy::config::model_router::RoutingStrategy::LeastBusy
|
||||
);
|
||||
|
||||
unsafe {
|
||||
std::env::remove_var("PROXY_CONFIG");
|
||||
std::env::remove_var("OPENAI_API_KEY");
|
||||
};
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn load_litellm_format_unchanged() {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let path = dir.path().join("litellm.yaml");
|
||||
std::fs::write(
|
||||
&path,
|
||||
r#"
|
||||
model_list:
|
||||
- model_name: gpt-4o
|
||||
litellm_params:
|
||||
model: openai/gpt-4o
|
||||
api_key: sk-test
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
unsafe { std::env::set_var("PROXY_CONFIG", path.to_str().unwrap()) };
|
||||
|
||||
let result = MultiConfig::load();
|
||||
assert!(result.model_router.is_some());
|
||||
let router = result.model_router.unwrap();
|
||||
assert!(router.read().unwrap().has_model("gpt-4o"));
|
||||
|
||||
unsafe { std::env::remove_var("PROXY_CONFIG") };
|
||||
}
|
||||
Reference in New Issue
Block a user