From 3d2abf9a441babc9d6646fe2f824aa269b01cb98 Mon Sep 17 00:00:00 2001 From: yujonglee Date: Fri, 9 Oct 2026 16:29:45 -0700 Subject: [PATCH] feat(rust): one LitellmParams type for a config and a caller (#45642) litellm-router-types mirrors litellm/types/router.py: LitellmParams holds the model, the credentials, one flattened group per credential family from auth-types (aws, vertex), the deployment settings a config spells, and every other key in extra, as Python's extra="allow" keeps it. Config parses it directly, so config::LiteLlmParams and its untyped additional_fields are gone and a mistyped provider param is a parse error. fields() and specs() derive from the families' param specs for hosts that project kwargs and fold module globals. Spelled is the one untagged shape for a typed value or the text a config spells, organization takes the list Python's router expands, and the callback shorthand keeps a value-level one-or-many. No credentials wrapper group: serde's flatten only consumes keys for struct-shaped children. --- litellm-rust/Cargo.lock | 13 ++ litellm-rust/Cargo.toml | 1 + litellm-rust/crates/auth-aws/Cargo.toml | 1 + .../crates/auth-types/src/params/aws.rs | 7 +- litellm-rust/crates/config/Cargo.toml | 1 + litellm-rust/crates/config/src/lib.rs | 5 +- litellm-rust/crates/config/src/model.rs | 71 +------- litellm-rust/crates/config/src/settings.rs | 19 ++- litellm-rust/crates/config/src/value.rs | 47 ++---- litellm-rust/crates/config/tests/config.rs | 47 ++++-- litellm-rust/crates/gateway/src/lib.rs | 6 +- litellm-rust/crates/router-types/Cargo.toml | 16 ++ litellm-rust/crates/router-types/src/lib.rs | 5 + .../crates/router-types/src/litellm_params.rs | 92 ++++++++++ litellm-rust/crates/router-types/src/value.rs | 10 ++ .../router-types/tests/litellm_params.rs | 158 ++++++++++++++++++ 16 files changed, 364 insertions(+), 135 deletions(-) create mode 100644 litellm-rust/crates/router-types/Cargo.toml create mode 100644 litellm-rust/crates/router-types/src/lib.rs create mode 100644 litellm-rust/crates/router-types/src/litellm_params.rs create mode 100644 litellm-rust/crates/router-types/src/value.rs create mode 100644 litellm-rust/crates/router-types/tests/litellm_params.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 4cbde047a71..ec8d10b58de 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3383,6 +3383,7 @@ dependencies = [ "sha2 0.10.9", "thiserror 2.0.19", "tokio", + "veil", ] [[package]] @@ -3645,6 +3646,7 @@ name = "litellm-config" version = "0.1.0" dependencies = [ "litellm-auth-types", + "litellm-router-types", "rstest", "serde", "serde_json", @@ -4271,6 +4273,17 @@ dependencies = [ "rstest", ] +[[package]] +name = "litellm-router-types" +version = "0.1.0" +dependencies = [ + "litellm-auth-types", + "rstest", + "serde", + "serde_json", + "serde_with", +] + [[package]] name = "litellm-secrets" version = "0.1.0" diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 794070366af..b94988eb8e1 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -11,6 +11,7 @@ repository = "https://github.com/BerriAI/litellm" [workspace.dependencies] litellm-config = { path = "crates/config" } litellm-router = { path = "crates/router" } +litellm-router-types = { path = "crates/router-types" } litellm-tracing = { path = "crates/tracing" } litellm-spend-clickhouse = { path = "crates/spend-clickhouse" } litellm-storage-clickhouse = { path = "crates/storage-clickhouse" } diff --git a/litellm-rust/crates/auth-aws/Cargo.toml b/litellm-rust/crates/auth-aws/Cargo.toml index bf757c6744e..65b411ffaef 100644 --- a/litellm-rust/crates/auth-aws/Cargo.toml +++ b/litellm-rust/crates/auth-aws/Cargo.toml @@ -14,6 +14,7 @@ serde.workspace = true serde_json = { workspace = true, features = ["preserve_order"] } sha2.workspace = true thiserror.workspace = true +veil.workspace = true aws-config = { version = "1.9.0", default-features = false, features = ["rustls", "rt-tokio"] } aws-credential-types = { version = "1.3.0", features = ["hardcoded-credentials"] } diff --git a/litellm-rust/crates/auth-types/src/params/aws.rs b/litellm-rust/crates/auth-types/src/params/aws.rs index f3653e969d6..25fe52b2ce1 100644 --- a/litellm-rust/crates/auth-types/src/params/aws.rs +++ b/litellm-rust/crates/auth-types/src/params/aws.rs @@ -2,6 +2,7 @@ use std::sync::LazyLock; use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; +use veil::Redact; use super::ParamSpec; @@ -19,13 +20,15 @@ const fn spec( } /// The `aws_*` fields of Python's `GenericLiteLLMParams`. -#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] +#[derive(Redact, Clone, Default, PartialEq, Eq, Serialize, Deserialize)] pub struct AwsParams { #[serde(default)] pub aws_access_key_id: Option, #[serde(default)] + #[redact(with = "[REDACTED]")] pub aws_secret_access_key: Option, #[serde(default)] + #[redact(with = "[REDACTED]")] pub aws_session_token: Option, #[serde(default)] pub aws_region_name: Option, @@ -36,10 +39,12 @@ pub struct AwsParams { #[serde(default)] pub aws_role_name: Option, #[serde(default)] + #[redact(with = "[REDACTED]")] pub aws_web_identity_token: Option, #[serde(default)] pub aws_sts_endpoint: Option, #[serde(default)] + #[redact(with = "[REDACTED]")] pub aws_external_id: Option, #[serde(default)] pub aws_bedrock_runtime_endpoint: Option, diff --git a/litellm-rust/crates/config/Cargo.toml b/litellm-rust/crates/config/Cargo.toml index e23739997e2..23d53923362 100644 --- a/litellm-rust/crates/config/Cargo.toml +++ b/litellm-rust/crates/config/Cargo.toml @@ -7,6 +7,7 @@ repository.workspace = true [dependencies] litellm-auth-types.workspace = true +litellm-router-types.workspace = true serde.workspace = true serde_yaml_ng = "0.10.0" strum.workspace = true diff --git a/litellm-rust/crates/config/src/lib.rs b/litellm-rust/crates/config/src/lib.rs index fed60ab1a4f..8603ef3279c 100644 --- a/litellm-rust/crates/config/src/lib.rs +++ b/litellm-rust/crates/config/src/lib.rs @@ -10,13 +10,14 @@ use std::{fmt, path::Path}; use serde::Deserialize; pub use error::Error; +pub use litellm_router_types::{LitellmParams, Spelled}; pub use mcp::{McpAuth, McpServer, McpTransport}; -pub use model::{LiteLlmParams, Model}; +pub use model::Model; pub use settings::{ ClickHouseStoreSettings, GeneralSettings, LiteLlmSettings, RouterSettings, TracingSettings, TracingStoreSettings, }; -pub use value::{AdditionalFields, Flag, NumberOrString, Object, OneOrMany, Value}; +pub use value::{AdditionalFields, Object, Value}; #[derive(Clone, Default, Deserialize)] #[serde(default)] diff --git a/litellm-rust/crates/config/src/model.rs b/litellm-rust/crates/config/src/model.rs index cd05bb4b7c8..c0f8516103d 100644 --- a/litellm-rust/crates/config/src/model.rs +++ b/litellm-rust/crates/config/src/model.rs @@ -1,14 +1,14 @@ use std::fmt; -use litellm_auth_types::SecretValue; +use litellm_router_types::LitellmParams; use serde::Deserialize; -use crate::{AdditionalFields, Flag, NumberOrString, Object}; +use crate::{AdditionalFields, Object}; #[derive(Clone, Deserialize)] pub struct Model { pub model_name: String, - pub litellm_params: LiteLlmParams, + pub litellm_params: LitellmParams, #[serde(default)] pub model_info: Object, pub blocked: Option, @@ -28,68 +28,3 @@ impl fmt::Debug for Model { .finish() } } - -#[derive(Clone, Deserialize)] -pub struct LiteLlmParams { - pub model: String, - pub api_key: Option, - pub api_base: Option, - pub api_version: Option, - pub custom_llm_provider: Option, - pub timeout: Option, - pub stream_timeout: Option, - pub max_retries: Option, - pub tpm: Option, - pub rpm: Option, - pub itpm: Option, - pub otpm: Option, - pub max_parallel_requests: Option, - pub organization: Option, - pub drop_params: Option, - pub tags: Option>, - pub tag_regex: Option>, - pub max_budget: Option, - pub budget_duration: Option, - pub default_api_key_tpm_limit: Option, - pub default_api_key_rpm_limit: Option, - pub use_in_pass_through: Option, - pub use_chat_completions_api: Option, - pub litellm_credential_name: Option, - pub provider_affinity_header: Option, - #[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() - } -} diff --git a/litellm-rust/crates/config/src/settings.rs b/litellm-rust/crates/config/src/settings.rs index b1ead35e55b..1094de1638e 100644 --- a/litellm-rust/crates/config/src/settings.rs +++ b/litellm-rust/crates/config/src/settings.rs @@ -3,7 +3,7 @@ use std::fmt; use litellm_auth_types::SecretValue; use serde::Deserialize; -use crate::{AdditionalFields, Flag, NumberOrString, Object, OneOrMany, Value}; +use crate::{AdditionalFields, Object, Spelled, Value, value::one_or_many}; #[derive(Clone, Debug, Deserialize)] #[serde(rename_all = "lowercase")] @@ -18,7 +18,7 @@ pub struct ClickHouseStoreSettings { pub kind: TracingStoreKind, pub url: Option, pub database: Option, - pub retention_days: Option, + pub retention_days: Option>, } impl fmt::Debug for ClickHouseStoreSettings { @@ -207,7 +207,7 @@ impl fmt::Debug for RouterSettings { #[derive(Clone, Default, Deserialize)] #[serde(default)] pub struct LiteLlmSettings { - pub ssl_verify: Option, + pub ssl_verify: Option>, pub ssl_certificate: Option, pub ssl_security_level: Option, pub ssl_ecdh_curve: Option, @@ -216,14 +216,17 @@ pub struct LiteLlmSettings { pub aiohttp_trust_env: Option, pub disable_aiohttp_trust_env: Option, pub disable_aiohttp_transport: Option, - pub drop_params: Option, - pub request_timeout: Option, + pub drop_params: Option>, + pub request_timeout: Option>, pub num_retries: Option, pub cache: Option, pub cache_params: Option, - pub callbacks: Option>, - pub success_callback: Option>, - pub failure_callback: Option>, + #[serde(deserialize_with = "one_or_many")] + pub callbacks: Option>, + #[serde(deserialize_with = "one_or_many")] + pub success_callback: Option>, + #[serde(deserialize_with = "one_or_many")] + pub failure_callback: Option>, pub json_logs: Option, pub set_verbose: Option, #[serde(flatten)] diff --git a/litellm-rust/crates/config/src/value.rs b/litellm-rust/crates/config/src/value.rs index 6b68b5bc2d0..5becdc87b70 100644 --- a/litellm-rust/crates/config/src/value.rs +++ b/litellm-rust/crates/config/src/value.rs @@ -1,6 +1,6 @@ use std::{collections::BTreeMap, fmt, ops::Deref}; -use serde::Deserialize; +use serde::{Deserialize, Deserializer}; pub type Value = serde_yaml_ng::Value; pub type AdditionalFields = BTreeMap; @@ -40,39 +40,14 @@ impl fmt::Debug for Object { } } -#[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 { - Many(Box<[T]>), - One(T), -} - -impl OneOrMany { - 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, - } - } +/// Python's callback shorthand: one entry or a list of them, kept as the values written. +pub(crate) fn one_or_many<'de, D: Deserializer<'de>>( + deserializer: D, +) -> Result>, D::Error> { + Ok( + Option::::deserialize(deserializer)?.map(|value| match value { + Value::Sequence(values) => values, + one => vec![one], + }), + ) } diff --git a/litellm-rust/crates/config/tests/config.rs b/litellm-rust/crates/config/tests/config.rs index ab9f3403a01..d425b62efc2 100644 --- a/litellm-rust/crates/config/tests/config.rs +++ b/litellm-rust/crates/config/tests/config.rs @@ -1,4 +1,4 @@ -use litellm_config::{Config, Error, Flag, NumberOrString, TracingStoreSettings}; +use litellm_config::{Config, Error, Spelled, TracingStoreSettings}; use rstest::{fixture, rstest}; use tempfile::TempDir; @@ -75,6 +75,9 @@ fn config_debug_redacts_api_keys() { #[case::malformed_yaml("model_list: [")] #[case::missing_params("model_list: [{model_name: assistant}]")] #[case::missing_model("model_list: [{model_name: assistant, litellm_params: {api_key: key}}]")] +#[case::provider_param_of_the_wrong_type( + "model_list: [{model_name: assistant, litellm_params: {model: test, aws_region_name: [eu-central-1]}}]" +)] fn rejects_malformed_and_incomplete_config(#[case] yaml: &str) { assert!(matches!(Config::from_yaml(yaml), Err(Error::Parse(_)))); } @@ -128,7 +131,7 @@ fn tracing_settings_are_typed_and_redact_the_url() { "https://writer:password@example.com" ); assert_eq!(store.database.as_deref(), Some("analytics")); - assert_eq!(store.retention_days, Some(NumberOrString::Number(7.0))); + assert_eq!(store.retention_days, Some(Spelled::Value(7.0))); assert!(!format!("{config:?}").contains("password")); } @@ -145,9 +148,7 @@ fn tracing_settings_accept_environment_references() { }; assert_eq!( store.retention_days, - Some(NumberOrString::String( - "os.environ/RETENTION_DAYS".to_owned() - )) + Some(Spelled::Text("os.environ/RETENTION_DAYS".to_owned())) ); } @@ -190,6 +191,9 @@ model_list: rpm: 5 drop_params: "true" vertex_project: test-project + vertex_credentials: + type: service_account + future_param: 1 model_info: mode: chat access_groups: [internal] @@ -214,25 +218,26 @@ future_section: let model = &config.model_list[0]; assert_eq!( model.litellm_params.timeout, - Some(NumberOrString::String( - "os.environ/REQUEST_TIMEOUT".to_string() - )) + Some(Spelled::Text("os.environ/REQUEST_TIMEOUT".to_string())) ); assert_eq!( model.litellm_params.tpm, - Some(NumberOrString::String("os.environ/TPM_LIMIT".to_string())) + Some(Spelled::Text("os.environ/TPM_LIMIT".to_string())) ); - assert_eq!(model.litellm_params.rpm, Some(NumberOrString::Number(5.0))); + assert_eq!(model.litellm_params.rpm, Some(Spelled::Value(5.0))); assert_eq!( model.litellm_params.drop_params, - Some(Flag::String("true".to_string())) + Some(Spelled::Text("true".to_string())) ); - assert!( - model - .litellm_params - .additional_fields - .contains_key("vertex_project") + assert_eq!( + model.litellm_params.vertex.vertex_project.as_deref(), + Some("test-project") ); + assert_eq!( + model.litellm_params.vertex.vertex_credentials.as_deref(), + Some(r#"{"type":"service_account"}"#) + ); + assert!(model.litellm_params.extra.contains_key("future_param")); assert!(model.additional_fields.contains_key("access_groups")); assert_eq!(config.router_settings.allowed_fails, Some(2)); assert!( @@ -345,7 +350,9 @@ fn accepts_python_callback_shorthand(#[case] setting: &str, #[case] expected_len #[rstest] #[case::api_key("api_key", "provider-secret")] -#[case::provider_extension("aws_secret_access_key", "aws-secret")] +#[case::aws_secret("aws_secret_access_key", "aws-secret")] +#[case::vertex_credentials("vertex_credentials", "vertex-secret")] +#[case::provider_extension("azure_ad_token", "azure-secret")] #[case::general_extension("custom_auth_secret", "auth-secret")] #[case::router_extension("redis_password", "redis-secret")] #[case::litellm_extension("callback_token", "callback-secret")] @@ -358,6 +365,12 @@ fn debug_output_does_not_expose_config_values(#[case] field: &str, #[case] secre "aws_secret_access_key" => format!( "model_list: [{{model_name: assistant, litellm_params: {{model: test, aws_secret_access_key: {secret}}}}}]" ), + "vertex_credentials" => format!( + "model_list: [{{model_name: assistant, litellm_params: {{model: test, vertex_credentials: {secret}}}}}]" + ), + "azure_ad_token" => format!( + "model_list: [{{model_name: assistant, litellm_params: {{model: test, azure_ad_token: {secret}}}}}]" + ), "custom_auth_secret" => format!("general_settings: {{{field}: {secret}}}"), "redis_password" => format!("router_settings: {{{field}: {secret}}}"), "callback_token" => format!("litellm_settings: {{{field}: {secret}}}"), diff --git a/litellm-rust/crates/gateway/src/lib.rs b/litellm-rust/crates/gateway/src/lib.rs index d711eb78caa..0b3aba81181 100644 --- a/litellm-rust/crates/gateway/src/lib.rs +++ b/litellm-rust/crates/gateway/src/lib.rs @@ -40,9 +40,9 @@ pub fn build_inference(config: &Config) -> Result, Error> { HttpSettingsLayer::from_environment(&lookup), HttpSettingsLayer { ssl_verify: settings.ssl_verify.as_ref().map(|value| match value { - litellm_config::Flag::Boolean(true) => SslVerify::Enabled, - litellm_config::Flag::Boolean(false) => SslVerify::Disabled, - litellm_config::Flag::String(value) => SslVerify::parse(value), + litellm_config::Spelled::Value(true) => SslVerify::Enabled, + litellm_config::Spelled::Value(false) => SslVerify::Disabled, + litellm_config::Spelled::Text(value) => SslVerify::parse(value), }), ssl_certificate: settings.ssl_certificate.as_ref().map(Into::into), ssl_security_level: settings.ssl_security_level.clone(), diff --git a/litellm-rust/crates/router-types/Cargo.toml b/litellm-rust/crates/router-types/Cargo.toml new file mode 100644 index 00000000000..2e9fec28685 --- /dev/null +++ b/litellm-rust/crates/router-types/Cargo.toml @@ -0,0 +1,16 @@ +[package] +name = "litellm-router-types" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true +description = "The types of litellm/types/router.py: a deployment's litellm_params as a config or a caller spells them" + +[dependencies] +litellm-auth-types.workspace = true +serde.workspace = true +serde_json.workspace = true +serde_with.workspace = true + +[dev-dependencies] +rstest.workspace = true diff --git a/litellm-rust/crates/router-types/src/lib.rs b/litellm-rust/crates/router-types/src/lib.rs new file mode 100644 index 00000000000..2235e40a891 --- /dev/null +++ b/litellm-rust/crates/router-types/src/lib.rs @@ -0,0 +1,5 @@ +mod litellm_params; +mod value; + +pub use litellm_params::{ExtraParams, LitellmParams}; +pub use value::Spelled; diff --git a/litellm-rust/crates/router-types/src/litellm_params.rs b/litellm-rust/crates/router-types/src/litellm_params.rs new file mode 100644 index 00000000000..186a1d98036 --- /dev/null +++ b/litellm-rust/crates/router-types/src/litellm_params.rs @@ -0,0 +1,92 @@ +use std::{collections::BTreeMap, fmt, ops::Deref}; + +use litellm_auth_types::{AwsParams, ParamSpec, SecretValue, VertexParams}; +use serde::Deserialize; +use serde_json::Value; +use serde_with::{OneOrMany, formats::PreferMany, serde_as}; + +use crate::Spelled; + +/// Python's `LiteLLM_Params`: a deployment's `litellm_params`, as a config spells them and as +/// a caller passes them. Each credential family is one flattened group owned by its auth +/// crate, so a config reaches for `litellm_params.aws` or `litellm_params.vertex`, never a +/// map. Every key no field names lands in [`Self::extra`], as Python's `extra="allow"` keeps it. +/// +/// A key a field names but cannot hold is an error, as it is for Python's pydantic model; +/// the coercions Python's validators apply (`max_retries` from text, `drop_params` from a flag +/// word) are left to the readers, so the shapes here are the ones the YAML carries. An +/// `organization` list is what Python's router expands into one deployment per entry. +#[serde_as] +#[derive(Clone, Debug, Default, PartialEq, Deserialize)] +pub struct LitellmParams { + pub model: String, + pub api_key: Option, + pub api_base: Option, + pub api_version: Option, + #[serde(flatten)] + pub aws: AwsParams, + #[serde(flatten)] + pub vertex: VertexParams, + pub custom_llm_provider: Option, + pub timeout: Option>, + pub stream_timeout: Option>, + pub max_retries: Option>, + pub tpm: Option>, + pub rpm: Option>, + pub itpm: Option>, + pub otpm: Option>, + pub max_parallel_requests: Option, + #[serde_as(as = "Option>")] + pub organization: Option>, + pub drop_params: Option>, + pub tags: Option>, + pub tag_regex: Option>, + pub max_budget: Option, + pub budget_duration: Option, + pub default_api_key_tpm_limit: Option, + pub default_api_key_rpm_limit: Option, + pub use_in_pass_through: Option, + pub use_chat_completions_api: Option, + pub litellm_credential_name: Option, + pub provider_affinity_header: Option, + #[serde(flatten)] + pub extra: ExtraParams, +} + +impl LitellmParams { + /// Every connection param spec the credential families declare, for hosts that fold a + /// spec's module global in when a call names none of its spellings. + pub fn specs() -> impl Iterator { + AwsParams::SPECS.iter().chain(VertexParams::SPECS.iter()) + } + + /// The wire names a host projects out of a caller's kwargs: the model, the credentials + /// and each credential family. The routing and limit fields are a deployment's + /// settings, which only a config spells. + pub fn fields() -> impl Iterator { + ["model", "api_key", "api_base", "api_version"] + .into_iter() + .chain(AwsParams::fields()) + .chain(VertexParams::fields()) + } +} + +/// The `litellm_params` keys no field names, kept as written. Its `Debug` lists the keys +/// only, since a provider's secret may sit among the values. +#[derive(Clone, Default, PartialEq, Deserialize)] +#[serde(transparent)] +pub struct ExtraParams(BTreeMap); + +impl Deref for ExtraParams { + type Target = BTreeMap; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl fmt::Debug for ExtraParams { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.debug_list().entries(self.0.keys()).finish() + } +} diff --git a/litellm-rust/crates/router-types/src/value.rs b/litellm-rust/crates/router-types/src/value.rs new file mode 100644 index 00000000000..e52cf60c0ef --- /dev/null +++ b/litellm-rust/crates/router-types/src/value.rs @@ -0,0 +1,10 @@ +use serde::Deserialize; + +/// A typed value, or the text a config spells in its place: an `os.environ/NAME` reference +/// the secrets layer resolves later, or a word Python's coercion reads the value from. +#[derive(Clone, Debug, Deserialize, PartialEq)] +#[serde(untagged)] +pub enum Spelled { + Value(T), + Text(String), +} diff --git a/litellm-rust/crates/router-types/tests/litellm_params.rs b/litellm-rust/crates/router-types/tests/litellm_params.rs new file mode 100644 index 00000000000..7d021fd857a --- /dev/null +++ b/litellm-rust/crates/router-types/tests/litellm_params.rs @@ -0,0 +1,158 @@ +use litellm_auth_types::{AwsParams, SecretValue, VertexParams}; +use litellm_router_types::{LitellmParams, Spelled}; +use rstest::rstest; +use serde_json::{Map, Value, json}; + +fn params(value: Value) -> LitellmParams { + serde_json::from_value(value).unwrap() +} + +#[rstest] +#[case::aws_group( + json!({"model": "m", "aws_region_name": "eu-west-1", "aws_access_key_id": "AKIA"}), + LitellmParams { + model: "m".into(), + aws: AwsParams { + aws_region_name: Some("eu-west-1".into()), + aws_access_key_id: Some("AKIA".into()), + ..AwsParams::default() + }, + ..LitellmParams::default() + }, +)] +#[case::vertex_group( + json!({"model": "m", "vertex_project": "p", "vertex_ai_location": "eu", "vertex_credentials": {"type": "service_account"}}), + LitellmParams { + model: "m".into(), + vertex: VertexParams { + vertex_project: Some("p".into()), + vertex_ai_location: Some("eu".into()), + vertex_credentials: Some(r#"{"type":"service_account"}"#.into()), + ..VertexParams::default() + }, + ..LitellmParams::default() + }, +)] +#[case::both_groups_and_the_credentials( + json!({"model": "m", "api_key": "k", "api_base": "b", "aws_region_name": "eu-west-1", "vertex_location": "us-east5"}), + LitellmParams { + model: "m".into(), + api_key: Some(SecretValue::new("k")), + api_base: Some("b".into()), + aws: AwsParams { + aws_region_name: Some("eu-west-1".into()), + ..AwsParams::default() + }, + vertex: VertexParams { + vertex_location: Some("us-east5".into()), + ..VertexParams::default() + }, + ..LitellmParams::default() + }, +)] +#[case::deployment_settings( + json!({"model": "m", "timeout": "os.environ/T", "rpm": 5, "drop_params": "true", "tags": ["a"], "max_budget": 1.5}), + LitellmParams { + model: "m".into(), + timeout: Some(Spelled::Text("os.environ/T".into())), + rpm: Some(Spelled::Value(5.0)), + drop_params: Some(Spelled::Text("true".into())), + tags: Some(Box::from(["a".to_string()])), + max_budget: Some(1.5), + ..LitellmParams::default() + }, +)] +#[case::explicit_null_is_absent( + json!({"model": "m", "aws_region_name": null, "vertex_project": null, "api_key": null, "tpm": null}), + LitellmParams { model: "m".into(), ..LitellmParams::default() }, +)] +#[case::one_organization( + json!({"model": "m", "organization": "org-a"}), + LitellmParams { + model: "m".into(), + organization: Some(vec!["org-a".into()]), + ..LitellmParams::default() + }, +)] +#[case::organizations_the_router_expands( + json!({"model": "m", "organization": ["org-a", "org-b"]}), + LitellmParams { + model: "m".into(), + organization: Some(vec!["org-a".into(), "org-b".into()]), + ..LitellmParams::default() + }, +)] +fn deserializes_each_group_from_a_config_or_a_call( + #[case] value: Value, + #[case] expected: LitellmParams, +) { + assert_eq!(params(value), expected); +} + +#[rstest] +fn keys_no_field_names_land_in_extra_and_nothing_else_does() { + let typed = params(json!({ + "model": "m", + "api_key": "k", + "aws_region_name": "eu-west-1", + "vertex_project": "p", + "rpm": 5, + "azure_ad_token": "t", + "messages": [], + })); + + assert_eq!( + typed.extra.keys().collect::>(), + ["azure_ad_token", "messages"] + ); + assert_eq!(typed.extra["azure_ad_token"], "t"); + assert_eq!(typed.aws.aws_region_name.as_deref(), Some("eu-west-1")); + assert_eq!(typed.vertex.vertex_project.as_deref(), Some("p")); +} + +#[rstest] +#[case::model_missing(json!({"api_key": "k"}))] +#[case::aws(json!({"model": "m", "aws_region_name": 7}))] +#[case::vertex(json!({"model": "m", "vertex_project": ["p"]}))] +#[case::vertex_credentials(json!({"model": "m", "vertex_credentials": 7}))] +#[case::rate_limit(json!({"model": "m", "max_parallel_requests": "many"}))] +#[case::tags(json!({"model": "m", "tags": "not-a-list"}))] +fn a_param_of_the_wrong_type_is_rejected(#[case] value: Value) { + assert!(serde_json::from_value::(value).is_err()); +} + +#[rstest] +fn fields_fill_the_model_the_credentials_and_every_group() { + let filled: Value = LitellmParams::fields() + .map(|name| (name.to_string(), Value::from(name))) + .collect::>() + .into(); + let typed = params(filled); + let nulls = |group: Value| { + group + .as_object() + .unwrap() + .values() + .filter(|value| value.is_null()) + .count() + }; + + assert_eq!(typed.model, "model"); + assert_eq!(typed.api_key, Some(SecretValue::new("api_key"))); + assert_eq!(typed.api_base.as_deref(), Some("api_base")); + assert_eq!(typed.api_version.as_deref(), Some("api_version")); + assert_eq!(nulls(serde_json::to_value(&typed.aws).unwrap()), 0); + assert_eq!(nulls(serde_json::to_value(&typed.vertex).unwrap()), 0); + assert!(typed.extra.is_empty()); +} + +#[rstest] +fn debug_shows_neither_the_api_key_nor_the_extra_values() { + let typed = + params(json!({"model": "m", "api_key": "sk-secret", "azure_ad_token": "token-secret"})); + let debug = format!("{typed:?}"); + + assert!(!debug.contains("sk-secret")); + assert!(!debug.contains("token-secret")); + assert!(debug.contains("azure_ad_token")); +}