feat(rust): expand gateway configuration parsing (#43460)

* feat(rust): expand gateway configuration parsing

* fix(config): accept environment references for model rate limits

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Yujong Lee <yujong@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-27 14:54:04 -07:00 • committed by GitHub
parent 268e8bb735
commit b94f5bdbed
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 746 additions and 33 deletions

View file

@ -0,0 +1,61 @@
use std::{
collections::{BTreeSet, VecDeque},
path::{Path, PathBuf},
};
use serde::Deserialize;
use serde_yaml_ng::{Mapping, Value};
use crate::Error;
#[derive(Deserialize)]
struct Includes {
#[serde(default)]
include: Vec<String>,
}
fn read(path: &Path) -> Result<Mapping, Error> {
Ok(serde_yaml_ng::from_str(&std::fs::read_to_string(path)?)?)
}
fn entries(config: &Mapping, path: &Path) -> Result<Vec<(String, PathBuf)>, Error> {
let includes: Includes = serde_yaml_ng::from_value(Value::Mapping(config.clone()))?;
Ok(includes
.include
.into_iter()
.map(|entry| (entry, path.to_owned()))
.collect())
}
pub(super) fn load(path: &Path) -> Result<Value, Error> {
let root = path.canonicalize()?;
let mut merged = read(&root)?;
let mut pending: VecDeque<_> = entries(&merged, &root)?.into();
let mut loaded = BTreeSet::from([root.clone()]);
merged.remove(Value::String("include".into()));
while let Some((entry, declaring)) = pending.pop_front() {
let declared = declaring.parent().unwrap_or(Path::new(".")).join(&entry);
let fallback = root.parent().unwrap_or(Path::new(".")).join(&entry);
let location = if declared.exists() {
declared
} else {
fallback
}
.canonicalize()?;
if !loaded.insert(location.clone()) {
continue;
}
let mut included = read(&location)?;
pending.extend(entries(&included, &location)?);
included.remove(Value::String("include".into()));
for (key, value) in included {
match (merged.get_mut(&key), value) {
(Some(Value::Sequence(base)), Value::Sequence(extra)) => base.extend(extra),
(_, value) => {
merged.insert(key, value);
}
}
}
}
Ok(Value::Mapping(merged))
}

View file

@ -1,40 +1,78 @@
mod error;
mod includes;
mod model;
mod settings;
mod value;
use std::path::Path;
use std::{fmt, path::Path};
use litellm_auth_types::SecretValue;
use serde::Deserialize;
pub use error::Error;
pub use model::{LiteLlmParams, Model};
pub use settings::{GeneralSettings, LiteLlmSettings, RouterSettings};
pub use value::{AdditionalFields, Flag, NumberOrString, Object, OneOrMany, Value};
#[derive(Clone, Debug, Deserialize)]
#[serde(deny_unknown_fields)]
#[derive(Clone, Default, Deserialize)]
#[serde(default)]
pub struct Config {
pub model_list: Box<[Model]>,
#[serde(default)]
pub general_settings: GeneralSettings,
pub router_settings: RouterSettings,
pub litellm_settings: LiteLlmSettings,
pub environment_variables: Object,
pub callback_settings: Object,
pub assistant_settings: Object,
pub default_vertex_config: Object,
pub credential_list: Box<[Object]>,
pub guardrails: Box<[Object]>,
pub prompts: Box<[Object]>,
pub sandbox_tools: Box<[Object]>,
pub search_tools: Box<[Object]>,
pub files_settings: Box<[Object]>,
pub finetune_settings: Box<[Object]>,
pub mcp_tools: Box<[Object]>,
pub vector_store_registry: Box<[Object]>,
pub worker_registry: Box<[Object]>,
pub agents: Box<[Object]>,
pub agent_list: Box<[Object]>,
pub policies: Object,
pub policy_attachments: Box<[Object]>,
pub include: Box<[String]>,
#[serde(flatten)]
pub additional_fields: AdditionalFields,
}
#[derive(Clone, Debug, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct GeneralSettings {
pub master_key: Option<SecretValue>,
}
#[derive(Clone, Debug, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct Model {
pub model_name: String,
pub litellm_params: LiteLlmParams,
}
#[derive(Clone, Debug, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct LiteLlmParams {
pub model: String,
pub api_key: Option<SecretValue>,
pub api_base: Option<String>,
pub custom_llm_provider: Option<String>,
impl fmt::Debug for Config {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("Config")
.field("model_list", &self.model_list)
.field("general_settings", &self.general_settings)
.field("router_settings", &self.router_settings)
.field("litellm_settings", &self.litellm_settings)
.field("environment_variables", &self.environment_variables)
.field("callback_settings", &self.callback_settings)
.field("assistant_settings", &self.assistant_settings)
.field("default_vertex_config", &self.default_vertex_config)
.field("credential_list", &self.credential_list)
.field("guardrails", &self.guardrails)
.field("prompts", &self.prompts)
.field("sandbox_tools", &self.sandbox_tools)
.field("search_tools", &self.search_tools)
.field("files_settings", &self.files_settings)
.field("finetune_settings", &self.finetune_settings)
.field("mcp_tools", &self.mcp_tools)
.field("vector_store_registry", &self.vector_store_registry)
.field("worker_registry", &self.worker_registry)
.field("agents", &self.agents)
.field("agent_list", &self.agent_list)
.field("policies", &self.policies)
.field("policy_attachments", &self.policy_attachments)
.field("include", &self.include)
.field("additional_fields", &self.additional_fields.keys())
.finish()
}
}
impl Config {
@ -43,6 +81,6 @@ impl Config {
}
pub fn load(path: impl AsRef<Path>) -> Result<Self, Error> {
Self::from_yaml(&std::fs::read_to_string(path)?)
Ok(serde_yaml_ng::from_value(includes::load(path.as_ref())?)?)
}
}

View file

@ -0,0 +1,95 @@
use std::fmt;
use litellm_auth_types::SecretValue;
use serde::Deserialize;
use crate::{AdditionalFields, Flag, NumberOrString, Object};
#[derive(Clone, Deserialize)]
pub struct Model {
pub model_name: String,
pub litellm_params: LiteLlmParams,
#[serde(default)]
pub model_info: Object,
pub blocked: Option<bool>,
#[serde(flatten)]
pub additional_fields: AdditionalFields,
}
impl fmt::Debug for Model {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("Model")
.field("model_name", &self.model_name)
.field("litellm_params", &self.litellm_params)
.field("model_info", &self.model_info)
.field("blocked", &self.blocked)
.field("additional_fields", &self.additional_fields.keys())
.finish()
}
}
#[derive(Clone, Deserialize)]
pub struct LiteLlmParams {
pub model: String,
pub api_key: Option<SecretValue>,
pub api_base: Option<String>,
pub api_version: Option<String>,
pub custom_llm_provider: Option<String>,
pub timeout: Option<NumberOrString>,
pub stream_timeout: Option<NumberOrString>,
pub max_retries: Option<NumberOrString>,
pub tpm: Option<NumberOrString>,
pub rpm: Option<NumberOrString>,
pub itpm: Option<NumberOrString>,
pub otpm: Option<NumberOrString>,
pub max_parallel_requests: Option<u64>,
pub organization: Option<serde_yaml_ng::Value>,
pub drop_params: Option<Flag>,
pub tags: Option<Box<[String]>>,
pub tag_regex: Option<Box<[String]>>,
pub max_budget: Option<f64>,
pub budget_duration: Option<String>,
pub default_api_key_tpm_limit: Option<u64>,
pub default_api_key_rpm_limit: Option<u64>,
pub use_in_pass_through: Option<bool>,
pub use_chat_completions_api: Option<bool>,
pub litellm_credential_name: Option<String>,
pub provider_affinity_header: Option<String>,
#[serde(flatten)]
pub additional_fields: AdditionalFields,
}
impl fmt::Debug for LiteLlmParams {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("LiteLlmParams")
.field("model", &self.model)
.field("api_key", &self.api_key)
.field("api_base", &self.api_base)
.field("api_version", &self.api_version)
.field("custom_llm_provider", &self.custom_llm_provider)
.field("timeout", &self.timeout)
.field("stream_timeout", &self.stream_timeout)
.field("max_retries", &self.max_retries)
.field("tpm", &self.tpm)
.field("rpm", &self.rpm)
.field("itpm", &self.itpm)
.field("otpm", &self.otpm)
.field("max_parallel_requests", &self.max_parallel_requests)
.field("organization", &self.organization)
.field("drop_params", &self.drop_params)
.field("tags", &self.tags)
.field("tag_regex", &self.tag_regex)
.field("max_budget", &self.max_budget)
.field("budget_duration", &self.budget_duration)
.field("default_api_key_tpm_limit", &self.default_api_key_tpm_limit)
.field("default_api_key_rpm_limit", &self.default_api_key_rpm_limit)
.field("use_in_pass_through", &self.use_in_pass_through)
.field("use_chat_completions_api", &self.use_chat_completions_api)
.field("litellm_credential_name", &self.litellm_credential_name)
.field("provider_affinity_header", &self.provider_affinity_header)
.field("additional_fields", &self.additional_fields.keys())
.finish()
}
}

View file

@ -0,0 +1,217 @@
use std::fmt;
use litellm_auth_types::SecretValue;
use serde::Deserialize;
use crate::{AdditionalFields, Flag, NumberOrString, Object, OneOrMany, Value};
#[derive(Clone, Deserialize)]
#[serde(default)]
pub struct GeneralSettings {
pub completion_model: Option<String>,
pub max_in_flight_requests_per_worker: Option<u64>,
pub max_queued_requests_per_worker: Option<u64>,
pub admission_queue_timeout_seconds: f64,
pub master_key: Option<SecretValue>,
pub database_url: Option<SecretValue>,
pub database_connection_pool_limit: Option<u64>,
pub database_connection_timeout: Option<f64>,
pub database_connect_timeout: Option<f64>,
pub database_socket_timeout: Option<f64>,
pub database_max_idle_connection_lifetime: Option<f64>,
pub max_parallel_requests: Option<u64>,
pub global_max_parallel_requests: Option<u64>,
pub max_request_size_mb: Option<u64>,
pub max_response_size_mb: Option<u64>,
pub proxy_config_reload_interval_seconds: u64,
pub background_health_checks: Option<bool>,
pub health_check_interval: u64,
pub health_check_concurrency: Option<u64>,
pub store_model_in_db: Option<bool>,
pub forward_client_headers_to_llm_api: Option<bool>,
pub cancel_on_disconnect: Option<bool>,
pub infer_model_from_keys: Option<bool>,
pub enable_public_model_hub: bool,
pub dangerously_permit_weak_or_unset_master_key: Option<bool>,
pub plugins: Option<Box<[Object]>>,
pub coordination_redis: Option<Object>,
#[serde(flatten)]
pub additional_fields: AdditionalFields,
}
impl Default for GeneralSettings {
fn default() -> Self {
Self {
completion_model: None,
max_in_flight_requests_per_worker: None,
max_queued_requests_per_worker: None,
admission_queue_timeout_seconds: 1.0,
master_key: None,
database_url: None,
database_connection_pool_limit: Some(10),
database_connection_timeout: Some(60.0),
database_connect_timeout: None,
database_socket_timeout: None,
database_max_idle_connection_lifetime: Some(60.0),
max_parallel_requests: None,
global_max_parallel_requests: None,
max_request_size_mb: None,
max_response_size_mb: None,
proxy_config_reload_interval_seconds: 30,
background_health_checks: None,
health_check_interval: 300,
health_check_concurrency: None,
store_model_in_db: None,
forward_client_headers_to_llm_api: None,
cancel_on_disconnect: None,
infer_model_from_keys: None,
enable_public_model_hub: false,
dangerously_permit_weak_or_unset_master_key: None,
plugins: None,
coordination_redis: None,
additional_fields: AdditionalFields::new(),
}
}
}
impl fmt::Debug for GeneralSettings {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("GeneralSettings")
.field("completion_model", &self.completion_model)
.field(
"max_in_flight_requests_per_worker",
&self.max_in_flight_requests_per_worker,
)
.field(
"max_queued_requests_per_worker",
&self.max_queued_requests_per_worker,
)
.field(
"admission_queue_timeout_seconds",
&self.admission_queue_timeout_seconds,
)
.field("master_key", &self.master_key)
.field("database_url", &self.database_url)
.field("store_model_in_db", &self.store_model_in_db)
.field("additional_fields", &self.additional_fields.keys())
.finish_non_exhaustive()
}
}
#[derive(Clone, Default, Deserialize)]
#[serde(default)]
pub struct RouterSettings {
pub routing_strategy: Option<String>,
pub routing_strategy_args: Option<Object>,
pub routing_groups: Option<Box<[Object]>>,
pub retry_policy: Option<Object>,
pub model_group_retry_policy: Option<Object>,
pub model_group_affinity_config: Option<Object>,
pub allowed_fails: Option<u64>,
pub cooldown_time: Option<f64>,
pub num_retries: Option<u64>,
pub timeout: Option<f64>,
pub max_retries: Option<u64>,
pub retry_after: Option<f64>,
pub fallbacks: Option<Box<[Object]>>,
pub context_window_fallbacks: Option<Box<[Object]>>,
pub model_group_alias: Option<Object>,
pub enable_tag_filtering: Option<bool>,
pub weights: Option<Object>,
pub tag_routing_prefix: Option<String>,
pub optional_pre_call_checks: Option<Box<[String]>>,
#[serde(flatten)]
pub additional_fields: AdditionalFields,
}
impl fmt::Debug for RouterSettings {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("RouterSettings")
.field("routing_strategy", &self.routing_strategy)
.field("routing_strategy_args", &self.routing_strategy_args)
.field("routing_groups", &self.routing_groups)
.field("retry_policy", &self.retry_policy)
.field("model_group_retry_policy", &self.model_group_retry_policy)
.field(
"model_group_affinity_config",
&self.model_group_affinity_config,
)
.field("allowed_fails", &self.allowed_fails)
.field("cooldown_time", &self.cooldown_time)
.field("num_retries", &self.num_retries)
.field("timeout", &self.timeout)
.field("max_retries", &self.max_retries)
.field("retry_after", &self.retry_after)
.field("fallbacks", &self.fallbacks)
.field("context_window_fallbacks", &self.context_window_fallbacks)
.field("model_group_alias", &self.model_group_alias)
.field("enable_tag_filtering", &self.enable_tag_filtering)
.field("weights", &self.weights)
.field("tag_routing_prefix", &self.tag_routing_prefix)
.field("optional_pre_call_checks", &self.optional_pre_call_checks)
.field("additional_fields", &self.additional_fields.keys())
.finish()
}
}
#[derive(Clone, Default, Deserialize)]
#[serde(default)]
pub struct LiteLlmSettings {
pub ssl_verify: Option<Flag>,
pub ssl_certificate: Option<String>,
pub ssl_security_level: Option<String>,
pub ssl_ecdh_curve: Option<String>,
pub force_ipv4: Option<bool>,
pub http2: Option<bool>,
pub aiohttp_trust_env: Option<bool>,
pub disable_aiohttp_trust_env: Option<bool>,
pub disable_aiohttp_transport: Option<bool>,
pub drop_params: Option<Flag>,
pub request_timeout: Option<NumberOrString>,
pub num_retries: Option<u64>,
pub cache: Option<bool>,
pub cache_params: Option<Object>,
pub callbacks: Option<OneOrMany<Value>>,
pub success_callback: Option<OneOrMany<Value>>,
pub failure_callback: Option<OneOrMany<Value>>,
pub json_logs: Option<bool>,
pub set_verbose: Option<bool>,
#[serde(flatten)]
pub additional_fields: AdditionalFields,
}
impl fmt::Debug for LiteLlmSettings {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("LiteLlmSettings")
.field("drop_params", &self.drop_params)
.field("request_timeout", &self.request_timeout)
.field("num_retries", &self.num_retries)
.field("cache", &self.cache)
.field("cache_params", &self.cache_params)
.field(
"callbacks",
&self.callbacks.as_ref().map(|callbacks| callbacks.len()),
)
.field(
"success_callback",
&self
.success_callback
.as_ref()
.map(|callbacks| callbacks.len()),
)
.field(
"failure_callback",
&self
.failure_callback
.as_ref()
.map(|callbacks| callbacks.len()),
)
.field("json_logs", &self.json_logs)
.field("set_verbose", &self.set_verbose)
.field("additional_fields", &self.additional_fields.keys())
.finish()
}
}

View file

@ -0,0 +1,78 @@
use std::{collections::BTreeMap, fmt, ops::Deref};
use serde::Deserialize;
pub type Value = serde_yaml_ng::Value;
pub type AdditionalFields = BTreeMap<String, Value>;
#[derive(Clone, Default, Deserialize)]
#[serde(transparent)]
pub struct Object(BTreeMap<String, Value>);
impl Object {
pub fn get(&self, key: &str) -> Option<&Value> {
self.0.get(key)
}
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
pub fn len(&self) -> usize {
self.0.len()
}
}
impl Deref for Object {
type Target = BTreeMap<String, Value>;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl fmt::Debug for Object {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("Object")
.field("keys", &self.0.keys())
.finish()
}
}
#[derive(Clone, Debug, Deserialize, PartialEq)]
#[serde(untagged)]
pub enum NumberOrString {
Number(f64),
String(String),
}
#[derive(Clone, Debug, Deserialize, PartialEq, Eq)]
#[serde(untagged)]
pub enum Flag {
Boolean(bool),
String(String),
}
#[derive(Clone, Debug, Deserialize)]
#[serde(untagged)]
pub enum OneOrMany<T> {
Many(Box<[T]>),
One(T),
}
impl<T> OneOrMany<T> {
pub fn len(&self) -> usize {
match self {
Self::Many(values) => values.len(),
Self::One(_) => 1,
}
}
pub fn is_empty(&self) -> bool {
match self {
Self::Many(values) => values.is_empty(),
Self::One(_) => false,
}
}
}

View file

@ -1,4 +1,4 @@
use litellm_config::{Config, Error};
use litellm_config::{Config, Error, Flag, NumberOrString};
use rstest::{fixture, rstest};
use tempfile::TempDir;
@ -73,14 +73,9 @@ fn config_debug_redacts_api_keys() {
#[rstest]
#[case::malformed_yaml("model_list: [")]
#[case::missing_model_list("{}")]
#[case::missing_params("model_list: [{model_name: assistant}]")]
#[case::missing_model("model_list: [{model_name: assistant, litellm_params: {api_key: key}}]")]
#[case::unsupported_settings("model_list: []\ngeneral_settings: {unknown: true}")]
#[case::misspelled_param(
"model_list: [{model_name: assistant, litellm_params: {model: test, api_bsae: url}}]"
)]
fn rejects_malformed_incomplete_and_unsupported_config(#[case] yaml: &str) {
fn rejects_malformed_and_incomplete_config(#[case] yaml: &str) {
assert!(matches!(Config::from_yaml(yaml), Err(Error::Parse(_))));
}
@ -117,3 +112,232 @@ fn missing_general_settings_has_no_master_key() {
let config = Config::from_yaml("model_list: []").unwrap();
assert!(config.general_settings.master_key.is_none());
}
#[rstest]
fn empty_config_matches_python_defaults() {
let config = Config::from_yaml("{}").unwrap();
assert!(config.model_list.is_empty());
assert_eq!(config.general_settings.admission_queue_timeout_seconds, 1.0);
assert_eq!(
config.general_settings.database_connection_pool_limit,
Some(10)
);
assert_eq!(
config.general_settings.proxy_config_reload_interval_seconds,
30
);
assert_eq!(config.general_settings.health_check_interval, 300);
}
#[rstest]
fn parses_typed_settings_and_preserves_extension_fields() {
let config = Config::from_yaml(
r#"
model_list:
- model_name: assistant
litellm_params:
model: vertex_ai/test-model
timeout: os.environ/REQUEST_TIMEOUT
tpm: os.environ/TPM_LIMIT
rpm: 5
drop_params: "true"
vertex_project: test-project
model_info:
mode: chat
access_groups: [internal]
general_settings:
master_key: secret-master-key
store_model_in_db: true
custom_auth: auth.py
router_settings:
routing_strategy: simple-shuffle
allowed_fails: 2
redis_host: cache.internal
litellm_settings:
drop_params: true
cache: true
custom_callback_name: audit
future_section:
enabled: true
"#,
)
.unwrap();
let model = &config.model_list[0];
assert_eq!(
model.litellm_params.timeout,
Some(NumberOrString::String(
"os.environ/REQUEST_TIMEOUT".to_string()
))
);
assert_eq!(
model.litellm_params.tpm,
Some(NumberOrString::String("os.environ/TPM_LIMIT".to_string()))
);
assert_eq!(model.litellm_params.rpm, Some(NumberOrString::Number(5.0)));
assert_eq!(
model.litellm_params.drop_params,
Some(Flag::String("true".to_string()))
);
assert!(
model
.litellm_params
.additional_fields
.contains_key("vertex_project")
);
assert!(model.additional_fields.contains_key("access_groups"));
assert_eq!(config.router_settings.allowed_fails, Some(2));
assert!(
config
.router_settings
.additional_fields
.contains_key("redis_host")
);
assert!(
config
.litellm_settings
.additional_fields
.contains_key("custom_callback_name")
);
assert!(config.additional_fields.contains_key("future_section"));
}
#[rstest]
fn parses_python_config_sections() {
let config = Config::from_yaml(
r#"
environment_variables:
REDIS_PORT: 6379
callback_settings:
otel:
message_logging: false
assistant_settings:
custom_llm_provider: openai
credential_list:
- credential_name: bedrock
credential_values:
aws_region_name: us-east-1
guardrails:
- guardrail_name: pii
litellm_params:
guardrail: presidio
prompts:
- prompt_id: support
sandbox_tools:
- sandbox_tool_name: e2b
search_tools:
- search_tool_name: web
files_settings:
- custom_llm_provider: openai
finetune_settings:
- custom_llm_provider: openai
mcp_tools:
- name: lookup
vector_store_registry:
- vector_store_name: docs
worker_registry:
- worker_id: regional
agents:
- agent_name: reviewer
agent_list:
- agent_name: legacy-reviewer
policies:
safe:
guardrails:
add: [pii]
policy_attachments:
- policy_id: safe
include:
- models.yaml
"#,
)
.unwrap();
assert_eq!(config.environment_variables.len(), 1);
assert_eq!(config.credential_list.len(), 1);
assert_eq!(config.guardrails.len(), 1);
assert_eq!(config.prompts.len(), 1);
assert_eq!(config.sandbox_tools.len(), 1);
assert_eq!(config.search_tools.len(), 1);
assert_eq!(config.files_settings.len(), 1);
assert_eq!(config.finetune_settings.len(), 1);
assert_eq!(config.mcp_tools.len(), 1);
assert_eq!(config.vector_store_registry.len(), 1);
assert_eq!(config.worker_registry.len(), 1);
assert_eq!(config.agents.len(), 1);
assert_eq!(config.agent_list.len(), 1);
assert!(config.policies.contains_key("safe"));
assert_eq!(config.policy_attachments.len(), 1);
assert_eq!(&*config.include, &["models.yaml"]);
}
#[rstest]
#[case::policy_pipeline("../../../litellm/proxy/example_config_yaml/test_pipeline_config.yaml")]
#[case::gateway("../../../tests/e2e/gateway/stage_mirror_ci_config.yml")]
fn parses_representative_python_configs(#[case] relative_path: &str) {
let path = std::path::Path::new(env!("CARGO_MANIFEST_DIR")).join(relative_path);
let config = Config::load(path).unwrap();
assert!(!config.model_list.is_empty());
}
#[rstest]
#[case::one("callbacks: custom_callbacks.logger", 1)]
#[case::many("callbacks: [prometheus, otel]", 2)]
fn accepts_python_callback_shorthand(#[case] setting: &str, #[case] expected_len: usize) {
let config = Config::from_yaml(&format!("litellm_settings:\n {setting}")).unwrap();
assert_eq!(
config.litellm_settings.callbacks.as_ref().unwrap().len(),
expected_len
);
}
#[rstest]
#[case::api_key("api_key", "provider-secret")]
#[case::provider_extension("aws_secret_access_key", "aws-secret")]
#[case::general_extension("custom_auth_secret", "auth-secret")]
#[case::router_extension("redis_password", "redis-secret")]
#[case::litellm_extension("callback_token", "callback-secret")]
#[case::root_extension("private_token", "root-secret")]
fn debug_output_does_not_expose_config_values(#[case] field: &str, #[case] secret: &str) {
let yaml = match field {
"api_key" => format!(
"model_list: [{{model_name: assistant, litellm_params: {{model: test, api_key: {secret}}}}}]"
),
"aws_secret_access_key" => format!(
"model_list: [{{model_name: assistant, litellm_params: {{model: test, aws_secret_access_key: {secret}}}}}]"
),
"custom_auth_secret" => format!("general_settings: {{{field}: {secret}}}"),
"redis_password" => format!("router_settings: {{{field}: {secret}}}"),
"callback_token" => format!("litellm_settings: {{{field}: {secret}}}"),
"private_token" => format!("{field}: {secret}"),
_ => unreachable!(),
};
let config = Config::from_yaml(&yaml).unwrap();
assert!(!format!("{config:?}").contains(secret));
}
#[rstest]
fn resolves_nested_includes_once_in_breadth_first_order() {
let directory = TempDir::new().unwrap();
let root = directory.path();
std::fs::create_dir(root.join("nested")).unwrap();
std::fs::write(root.join("config.yaml"), "include: [nested/first.yaml, second.yaml]\nmodel_list: [{model_name: root, litellm_params: {model: root}}]\n").unwrap();
std::fs::write(root.join("nested/first.yaml"), "include: [third.yaml]\nmodel_list: [{model_name: first, litellm_params: {model: first}}]\n").unwrap();
std::fs::write(root.join("second.yaml"), "general_settings: {master_key: second}\nmodel_list: [{model_name: second, litellm_params: {model: second}}]\n").unwrap();
std::fs::write(root.join("nested/third.yaml"), "include: [../config.yaml]\ngeneral_settings: {master_key: third}\nmodel_list: [{model_name: third, litellm_params: {model: third}}]\n").unwrap();
let config = Config::load(root.join("config.yaml")).unwrap();
assert_eq!(
config
.model_list
.iter()
.map(|model| model.model_name.as_str())
.collect::<Vec<_>>(),
["root", "first", "second", "third"]
);
assert_eq!(
config.general_settings.master_key.unwrap().expose(),
"third"
);
assert!(config.include.is_empty());
}