mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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<T> 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.
This commit is contained in:
parent
04a95f60c1
commit
3d2abf9a44
16 changed files with 364 additions and 135 deletions
13
litellm-rust/Cargo.lock
generated
13
litellm-rust/Cargo.lock
generated
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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" }
|
||||
|
|
|
|||
|
|
@ -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"] }
|
||||
|
|
|
|||
|
|
@ -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<String>,
|
||||
#[serde(default)]
|
||||
#[redact(with = "[REDACTED]")]
|
||||
pub aws_secret_access_key: Option<String>,
|
||||
#[serde(default)]
|
||||
#[redact(with = "[REDACTED]")]
|
||||
pub aws_session_token: Option<String>,
|
||||
#[serde(default)]
|
||||
pub aws_region_name: Option<String>,
|
||||
|
|
@ -36,10 +39,12 @@ pub struct AwsParams {
|
|||
#[serde(default)]
|
||||
pub aws_role_name: Option<String>,
|
||||
#[serde(default)]
|
||||
#[redact(with = "[REDACTED]")]
|
||||
pub aws_web_identity_token: Option<String>,
|
||||
#[serde(default)]
|
||||
pub aws_sts_endpoint: Option<String>,
|
||||
#[serde(default)]
|
||||
#[redact(with = "[REDACTED]")]
|
||||
pub aws_external_id: Option<String>,
|
||||
#[serde(default)]
|
||||
pub aws_bedrock_runtime_endpoint: Option<String>,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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<bool>,
|
||||
|
|
@ -28,68 +28,3 @@ impl fmt::Debug for Model {
|
|||
.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()
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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<SecretValue>,
|
||||
pub database: Option<String>,
|
||||
pub retention_days: Option<NumberOrString>,
|
||||
pub retention_days: Option<Spelled<f64>>,
|
||||
}
|
||||
|
||||
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<Flag>,
|
||||
pub ssl_verify: Option<Spelled<bool>>,
|
||||
pub ssl_certificate: Option<String>,
|
||||
pub ssl_security_level: Option<String>,
|
||||
pub ssl_ecdh_curve: Option<String>,
|
||||
|
|
@ -216,14 +216,17 @@ pub struct LiteLlmSettings {
|
|||
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 drop_params: Option<Spelled<bool>>,
|
||||
pub request_timeout: Option<Spelled<f64>>,
|
||||
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>>,
|
||||
#[serde(deserialize_with = "one_or_many")]
|
||||
pub callbacks: Option<Vec<Value>>,
|
||||
#[serde(deserialize_with = "one_or_many")]
|
||||
pub success_callback: Option<Vec<Value>>,
|
||||
#[serde(deserialize_with = "one_or_many")]
|
||||
pub failure_callback: Option<Vec<Value>>,
|
||||
pub json_logs: Option<bool>,
|
||||
pub set_verbose: Option<bool>,
|
||||
#[serde(flatten)]
|
||||
|
|
|
|||
|
|
@ -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<String, Value>;
|
||||
|
|
@ -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<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,
|
||||
}
|
||||
}
|
||||
/// 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<Option<Vec<Value>>, D::Error> {
|
||||
Ok(
|
||||
Option::<Value>::deserialize(deserializer)?.map(|value| match value {
|
||||
Value::Sequence(values) => values,
|
||||
one => vec![one],
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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}}}"),
|
||||
|
|
|
|||
|
|
@ -40,9 +40,9 @@ pub fn build_inference(config: &Config) -> Result<Arc<Gateway>, 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(),
|
||||
|
|
|
|||
16
litellm-rust/crates/router-types/Cargo.toml
Normal file
16
litellm-rust/crates/router-types/Cargo.toml
Normal file
|
|
@ -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
|
||||
5
litellm-rust/crates/router-types/src/lib.rs
Normal file
5
litellm-rust/crates/router-types/src/lib.rs
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
mod litellm_params;
|
||||
mod value;
|
||||
|
||||
pub use litellm_params::{ExtraParams, LitellmParams};
|
||||
pub use value::Spelled;
|
||||
92
litellm-rust/crates/router-types/src/litellm_params.rs
Normal file
92
litellm-rust/crates/router-types/src/litellm_params.rs
Normal file
|
|
@ -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<SecretValue>,
|
||||
pub api_base: Option<String>,
|
||||
pub api_version: Option<String>,
|
||||
#[serde(flatten)]
|
||||
pub aws: AwsParams,
|
||||
#[serde(flatten)]
|
||||
pub vertex: VertexParams,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub timeout: Option<Spelled<f64>>,
|
||||
pub stream_timeout: Option<Spelled<f64>>,
|
||||
pub max_retries: Option<Spelled<f64>>,
|
||||
pub tpm: Option<Spelled<f64>>,
|
||||
pub rpm: Option<Spelled<f64>>,
|
||||
pub itpm: Option<Spelled<f64>>,
|
||||
pub otpm: Option<Spelled<f64>>,
|
||||
pub max_parallel_requests: Option<u64>,
|
||||
#[serde_as(as = "Option<OneOrMany<_, PreferMany>>")]
|
||||
pub organization: Option<Vec<String>>,
|
||||
pub drop_params: Option<Spelled<bool>>,
|
||||
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 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<Item = &'static ParamSpec> {
|
||||
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<Item = &'static str> {
|
||||
["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<String, Value>);
|
||||
|
||||
impl Deref for ExtraParams {
|
||||
type Target = BTreeMap<String, Value>;
|
||||
|
||||
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()
|
||||
}
|
||||
}
|
||||
10
litellm-rust/crates/router-types/src/value.rs
Normal file
10
litellm-rust/crates/router-types/src/value.rs
Normal file
|
|
@ -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<T> {
|
||||
Value(T),
|
||||
Text(String),
|
||||
}
|
||||
158
litellm-rust/crates/router-types/tests/litellm_params.rs
Normal file
158
litellm-rust/crates/router-types/tests/litellm_params.rs
Normal file
|
|
@ -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::<Vec<_>>(),
|
||||
["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::<LitellmParams>(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::<Map<_, _>>()
|
||||
.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"));
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue