From 0a98c7be10243d7b267788c49e1a1a02e2c15961 Mon Sep 17 00:00:00 2001 From: yujonglee Date: Thu, 8 Oct 2026 14:24:49 -0700 Subject: [PATCH] refactor(rust): derive strum VariantArray and string conversions (#45434) * refactor(rust): derive strum VariantArray and string conversions Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(rust): use strum conversions directly with explicit spellings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(rust): use rstest values for key management systems Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../.agents/skills/rust-string-enums/SKILL.md | 4 +- litellm-rust/AGENTS.md | 4 ++ litellm-rust/Cargo.lock | 5 ++ litellm-rust/crates/cache/Cargo.toml | 1 + litellm-rust/crates/cache/src/cache_type.rs | 57 +++++++---------- litellm-rust/crates/cache/tests/cache_type.rs | 16 +++-- litellm-rust/crates/config/Cargo.toml | 2 + litellm-rust/crates/config/src/mcp.rs | 48 ++++++--------- litellm-rust/crates/config/tests/mcp.rs | 39 ++++++++++++ .../core-utils/src/get_llm_provider_logic.rs | 11 +++- .../src/prompt_templates/factory.rs | 9 +-- .../crates/gateway-mcp/src/configured.rs | 4 +- litellm-rust/crates/http/Cargo.toml | 1 + litellm-rust/crates/http/src/tls.rs | 61 ++++++++++++++----- .../src/formats/chat_completions.rs | 41 ++++--------- .../src/formats/messages/metadata.rs | 15 ++++- .../src/formats/messages/request.rs | 45 ++++++++------ .../llms-types/src/providers/anthropic.rs | 47 ++++++++------ .../llms/src/anthropic/chat/transformation.rs | 2 +- .../crates/llms/src/anthropic/common_utils.rs | 14 ++--- .../llms/src/anthropic/messages/thinking.rs | 10 ++- .../src/anthropic/messages/transformation.rs | 2 +- .../llms/src/aws_textract/ocr/common_utils.rs | 4 +- .../src/bedrock/audio_transcription/mod.rs | 13 ++-- .../bedrock/chat/converse_transformation.rs | 2 +- .../llms/src/bedrock/chat/invoke_handler.rs | 38 ++++++++---- .../llms/src/cohere/ocr/transformation.rs | 10 ++- .../python-bridge/src/python_settings.rs | 16 ++--- .../crates/python-bridge/src/routes/traces.rs | 2 +- .../crates/secrets-types/src/config.rs | 23 ++++++- .../crates/secrets-types/tests/context.rs | 18 +++--- litellm-rust/crates/secrets/src/oidc.rs | 7 ++- litellm-rust/crates/testkit/Cargo.toml | 1 + .../crates/testkit/src/agent/configure.rs | 2 +- .../crates/testkit/tests/configure.rs | 6 +- .../crates/traces-clickhouse/src/insert.rs | 46 ++++++++------ .../crates/traces-clickhouse/src/query.rs | 8 ++- .../crates/traces-clickhouse/src/table.rs | 4 +- .../crates/traces-clickhouse/tests/queries.rs | 4 +- .../src/normalize/format/claude_code.rs | 38 ++++++++---- .../traces/src/normalize/format/genai.rs | 8 ++- .../crates/traces/src/normalize/metadata.rs | 41 ++++++++++++- .../crates/traces/src/normalize/mod.rs | 14 ++++- litellm-rust/crates/traces/src/query.rs | 17 +++++- 44 files changed, 490 insertions(+), 270 deletions(-) create mode 100644 litellm-rust/crates/config/tests/mcp.rs diff --git a/litellm-rust/.agents/skills/rust-string-enums/SKILL.md b/litellm-rust/.agents/skills/rust-string-enums/SKILL.md index fc6988639d0..5fc067790a0 100644 --- a/litellm-rust/.agents/skills/rust-string-enums/SKILL.md +++ b/litellm-rust/.agents/skills/rust-string-enums/SKILL.md @@ -35,7 +35,7 @@ pub enum EventType { } ``` -Use `#[strum(serialize_all = "snake_case")]` or another supported case style when it exactly matches the contract. Use explicit variant spellings otherwise. With multiple accepted aliases, set `to_string` to the existing canonical output: Strum Display otherwise selects the longest `serialize` spelling +Spell each variant explicitly with `#[strum(serialize = "...")]`; do not use `serialize_all`. With multiple accepted aliases, set `to_string` to the existing canonical output: Strum Display otherwise selects the longest `serialize` spelling Use only the derives the contract needs. A deserialize-only type should remain deserialize-only. Do not add an unknown variant to a closed enum, or derive Serde for a type that currently has no serialization contract @@ -43,7 +43,7 @@ Plain Serde derives with `rename` or `rename_all` remain appropriate for closed ## Preserve behavior during migration -Read the type, its callers, and existing tests before changing it. Preserve canonical output, accepted aliases, case sensitivity, whitespace handling, unknown values, malformed-input rejection, and existing public conversion APIs. Keep conversions needed by callers or compatibility even when Serde no longer uses them +Read the type, its callers, and existing tests before changing it. Preserve canonical output, accepted aliases, case sensitivity, whitespace handling, unknown values, malformed-input rejection, and existing public conversion contracts. Wrappers that only delegate to a Strum-derived trait are removed rather than kept: callers use `FromStr` and `From for &'static str` directly. Keep conversions that add behavior needed by callers or compatibility even when Serde no longer uses them Use the workspace dependencies and enable `serde_with.workspace = true` in a crate only when needed. Check the versions and enabled features in `Cargo.toml` and `Cargo.lock` rather than upgrading dependencies for this refactor diff --git a/litellm-rust/AGENTS.md b/litellm-rust/AGENTS.md index 2c099c6f7a2..c22f9575f74 100644 --- a/litellm-rust/AGENTS.md +++ b/litellm-rust/AGENTS.md @@ -4,6 +4,10 @@ For diagnostic tracing changes, follow [.agents/skills/rust-tracing/SKILL.md](.a For string-valued enums and their Serde conversions, follow [.agents/skills/rust-string-enums/SKILL.md](.agents/skills/rust-string-enums/SKILL.md) +For fieldless enums, derive `strum::VariantArray` and use `VARIANTS` instead of a hand-listed `ALL` array; derive Strum string conversions instead of hand-written variant-to-string matches + +Use the derived conversions directly (`<&'static str>::from(x)` / `.into()`, `str::parse`) with no `as_str`/`parse` wrapper that only delegates, and spell each variant with explicit `#[strum(serialize = "...")]` instead of `serialize_all` + ## Test placement - Never create a `tests.rs` (or `test.rs`) file under `src/`, and never `#[path = "tests.rs"] mod tests;` diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 057faa279e3..9d511774b04 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3485,6 +3485,7 @@ dependencies = [ "rstest", "serde", "serde_json", + "strum", "thiserror 2.0.19", "tokio", ] @@ -3692,7 +3693,9 @@ dependencies = [ "litellm-auth-types", "rstest", "serde", + "serde_json", "serde_yaml_ng", + "strum", "tempfile", "thiserror 2.0.19", ] @@ -3990,6 +3993,7 @@ dependencies = [ "rustls-native-certs", "serde", "serde_json", + "strum", "tempfile", "thiserror 2.0.19", "tokio", @@ -4539,6 +4543,7 @@ dependencies = [ "serde", "serde_json", "sha2 0.10.9", + "strum", "tar", "target-lexicon", "tempfile", diff --git a/litellm-rust/crates/cache/Cargo.toml b/litellm-rust/crates/cache/Cargo.toml index f18dbd9cb26..51bbe12a85e 100644 --- a/litellm-rust/crates/cache/Cargo.toml +++ b/litellm-rust/crates/cache/Cargo.toml @@ -8,6 +8,7 @@ repository.workspace = true [dependencies] serde.workspace = true serde_json = { workspace = true, features = ["preserve_order"] } +strum.workspace = true thiserror.workspace = true [dev-dependencies] diff --git a/litellm-rust/crates/cache/src/cache_type.rs b/litellm-rust/crates/cache/src/cache_type.rs index 22d8c8c7cb5..80d040db487 100644 --- a/litellm-rust/crates/cache/src/cache_type.rs +++ b/litellm-rust/crates/cache/src/cache_type.rs @@ -1,57 +1,44 @@ use serde::{Deserialize, Serialize}; -#[derive(Clone, Copy, Debug, Deserialize, Serialize, PartialEq, Eq, Hash)] +#[derive( + Clone, + Copy, + Debug, + Deserialize, + Serialize, + PartialEq, + Eq, + Hash, + strum::EnumString, + strum::IntoStaticStr, + strum::VariantArray, +)] pub enum CacheType { #[serde(rename = "local")] + #[strum(serialize = "local")] Local, #[serde(rename = "redis")] + #[strum(serialize = "redis")] Redis, #[serde(rename = "redis-semantic")] + #[strum(serialize = "redis-semantic")] RedisSemantic, #[serde(rename = "valkey-semantic")] + #[strum(serialize = "valkey-semantic")] ValkeySemantic, #[serde(rename = "s3")] + #[strum(serialize = "s3")] S3, #[serde(rename = "disk")] + #[strum(serialize = "disk")] Disk, #[serde(rename = "qdrant-semantic")] + #[strum(serialize = "qdrant-semantic")] QdrantSemantic, #[serde(rename = "azure-blob")] + #[strum(serialize = "azure-blob")] AzureBlob, #[serde(rename = "gcs")] + #[strum(serialize = "gcs")] Gcs, } - -impl CacheType { - pub const ALL: [Self; 9] = [ - Self::Local, - Self::Redis, - Self::RedisSemantic, - Self::ValkeySemantic, - Self::S3, - Self::Disk, - Self::QdrantSemantic, - Self::AzureBlob, - Self::Gcs, - ]; - - pub const fn as_python_name(self) -> &'static str { - match self { - Self::Local => "local", - Self::Redis => "redis", - Self::RedisSemantic => "redis-semantic", - Self::ValkeySemantic => "valkey-semantic", - Self::S3 => "s3", - Self::Disk => "disk", - Self::QdrantSemantic => "qdrant-semantic", - Self::AzureBlob => "azure-blob", - Self::Gcs => "gcs", - } - } - - pub fn from_python_name(value: &str) -> Option { - Self::ALL - .into_iter() - .find(|cache_type| cache_type.as_python_name() == value) - } -} diff --git a/litellm-rust/crates/cache/tests/cache_type.rs b/litellm-rust/crates/cache/tests/cache_type.rs index 24aaba8e5fd..49b3e8b27ad 100644 --- a/litellm-rust/crates/cache/tests/cache_type.rs +++ b/litellm-rust/crates/cache/tests/cache_type.rs @@ -1,5 +1,6 @@ use litellm_cache::CacheType; use rstest::rstest; +use strum::VariantArray; #[rstest] #[case(CacheType::Local, "local")] @@ -15,16 +16,16 @@ fn every_python_cache_type_has_one_round_trip_identity( #[case] cache_type: CacheType, #[case] name: &str, ) { - assert_eq!(cache_type.as_python_name(), name); - assert_eq!(CacheType::from_python_name(name), Some(cache_type)); + assert_eq!(<&'static str>::from(cache_type), name); + assert_eq!(name.parse::().ok(), Some(cache_type)); assert_eq!( serde_json::to_value(cache_type).unwrap(), serde_json::Value::from(name) ); assert_eq!( - CacheType::ALL + CacheType::VARIANTS .iter() - .filter(|candidate| candidate.as_python_name() == name) + .filter(|candidate| <&'static str>::from(**candidate) == name) .count(), 1 ); @@ -33,7 +34,10 @@ fn every_python_cache_type_has_one_round_trip_identity( #[rstest] fn python_cache_types_are_listed_in_python_order() { assert_eq!( - CacheType::ALL.map(CacheType::as_python_name), + CacheType::VARIANTS + .iter() + .map(|cache_type| <&'static str>::from(*cache_type)) + .collect::>(), [ "local", "redis", @@ -52,5 +56,5 @@ fn python_cache_types_are_listed_in_python_order() { #[case::unknown("memcached")] #[case::case_sensitive("Redis")] fn unknown_python_names_have_no_cache_type(#[case] name: &str) { - assert_eq!(CacheType::from_python_name(name), None); + assert!(name.parse::().is_err()); } diff --git a/litellm-rust/crates/config/Cargo.toml b/litellm-rust/crates/config/Cargo.toml index 36bd68fe2a0..e23739997e2 100644 --- a/litellm-rust/crates/config/Cargo.toml +++ b/litellm-rust/crates/config/Cargo.toml @@ -9,8 +9,10 @@ repository.workspace = true litellm-auth-types.workspace = true serde.workspace = true serde_yaml_ng = "0.10.0" +strum.workspace = true thiserror.workspace = true [dev-dependencies] rstest.workspace = true +serde_json.workspace = true tempfile.workspace = true diff --git a/litellm-rust/crates/config/src/mcp.rs b/litellm-rust/crates/config/src/mcp.rs index 522744e4fe4..cbe992f9956 100644 --- a/litellm-rust/crates/config/src/mcp.rs +++ b/litellm-rust/crates/config/src/mcp.rs @@ -41,57 +41,43 @@ impl fmt::Debug for McpServer { } } -#[derive(Clone, Copy, Debug, Default, Deserialize, PartialEq, Eq)] +#[derive(Clone, Copy, Debug, Default, Deserialize, PartialEq, Eq, strum::IntoStaticStr)] #[serde(rename_all = "snake_case")] pub enum McpTransport { #[default] + #[strum(serialize = "http")] Http, + #[strum(serialize = "sse")] Sse, + #[strum(serialize = "stdio")] Stdio, } -impl McpTransport { - pub fn as_str(self) -> &'static str { - match self { - Self::Http => "http", - Self::Sse => "sse", - Self::Stdio => "stdio", - } - } -} - -#[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq)] +#[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq, strum::IntoStaticStr)] #[serde(rename_all = "snake_case")] pub enum McpAuth { + #[strum(serialize = "none")] None, + #[strum(serialize = "api_key")] ApiKey, + #[strum(serialize = "bearer_token")] BearerToken, + #[strum(serialize = "basic")] Basic, + #[strum(serialize = "authorization")] Authorization, + #[strum(serialize = "token")] Token, + #[strum(serialize = "oauth2")] Oauth2, + #[strum(serialize = "aws_sigv4")] AwsSigv4, + #[strum(serialize = "oauth2_token_exchange")] Oauth2TokenExchange, + #[strum(serialize = "oauth2_id_jag")] Oauth2IdJag, + #[strum(serialize = "true_passthrough")] TruePassthrough, + #[strum(serialize = "oauth_delegate")] OauthDelegate, } - -impl McpAuth { - pub fn as_str(self) -> &'static str { - match self { - Self::None => "none", - Self::ApiKey => "api_key", - Self::BearerToken => "bearer_token", - Self::Basic => "basic", - Self::Authorization => "authorization", - Self::Token => "token", - Self::Oauth2 => "oauth2", - Self::AwsSigv4 => "aws_sigv4", - Self::Oauth2TokenExchange => "oauth2_token_exchange", - Self::Oauth2IdJag => "oauth2_id_jag", - Self::TruePassthrough => "true_passthrough", - Self::OauthDelegate => "oauth_delegate", - } - } -} diff --git a/litellm-rust/crates/config/tests/mcp.rs b/litellm-rust/crates/config/tests/mcp.rs new file mode 100644 index 00000000000..9ac6ad02bbb --- /dev/null +++ b/litellm-rust/crates/config/tests/mcp.rs @@ -0,0 +1,39 @@ +use litellm_config::{McpAuth, McpTransport}; +use rstest::rstest; +use serde_json::json; + +#[rstest] +#[case::http(McpTransport::Http, "http")] +#[case::sse(McpTransport::Sse, "sse")] +#[case::stdio(McpTransport::Stdio, "stdio")] +fn mcp_transport_as_str_matches_the_serde_spelling( + #[case] transport: McpTransport, + #[case] name: &str, +) { + assert_eq!(<&'static str>::from(transport), name); + assert_eq!( + serde_json::from_value::(json!(name)).unwrap(), + transport + ); +} + +#[rstest] +#[case::none(McpAuth::None, "none")] +#[case::api_key(McpAuth::ApiKey, "api_key")] +#[case::bearer_token(McpAuth::BearerToken, "bearer_token")] +#[case::basic(McpAuth::Basic, "basic")] +#[case::authorization(McpAuth::Authorization, "authorization")] +#[case::token(McpAuth::Token, "token")] +#[case::oauth2(McpAuth::Oauth2, "oauth2")] +#[case::aws_sigv4(McpAuth::AwsSigv4, "aws_sigv4")] +#[case::oauth2_token_exchange(McpAuth::Oauth2TokenExchange, "oauth2_token_exchange")] +#[case::oauth2_id_jag(McpAuth::Oauth2IdJag, "oauth2_id_jag")] +#[case::true_passthrough(McpAuth::TruePassthrough, "true_passthrough")] +#[case::oauth_delegate(McpAuth::OauthDelegate, "oauth_delegate")] +fn mcp_auth_as_str_matches_the_serde_spelling(#[case] auth: McpAuth, #[case] name: &str) { + assert_eq!(<&'static str>::from(auth), name); + assert_eq!( + serde_json::from_value::(json!(name)).unwrap(), + auth + ); +} diff --git a/litellm-rust/crates/core-utils/src/get_llm_provider_logic.rs b/litellm-rust/crates/core-utils/src/get_llm_provider_logic.rs index e71ecdffa44..a2bb462c2a9 100644 --- a/litellm-rust/crates/core-utils/src/get_llm_provider_logic.rs +++ b/litellm-rust/crates/core-utils/src/get_llm_provider_logic.rs @@ -7,17 +7,26 @@ pub struct CustomLlmProvider<'a> { } #[derive(Debug, Clone, Copy, PartialEq, Eq, EnumString, IntoStaticStr)] -#[strum(serialize_all = "snake_case")] pub enum LlmProviders { + #[strum(serialize = "anthropic")] Anthropic, + #[strum(serialize = "aws_textract")] AwsTextract, + #[strum(serialize = "azure_ai")] AzureAi, + #[strum(serialize = "bedrock")] Bedrock, + #[strum(serialize = "cohere")] Cohere, + #[strum(serialize = "mistral")] Mistral, + #[strum(serialize = "openai")] Openai, + #[strum(serialize = "openai_like")] OpenaiLike, + #[strum(serialize = "reducto")] Reducto, + #[strum(serialize = "vertex_ai")] VertexAi, } diff --git a/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs b/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs index 10a0d719e9d..dcb98ed17f9 100644 --- a/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs +++ b/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs @@ -17,18 +17,13 @@ pub const EMPTY_TEXT_PLACEHOLDER: &str = "[System: Empty message content sanitised to satisfy protocol]"; #[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq)] -#[strum(serialize_all = "snake_case")] pub enum TurnRole { + #[strum(serialize = "user")] User, + #[strum(serialize = "assistant")] Assistant, } -impl TurnRole { - pub fn as_str(self) -> &'static str { - self.into() - } -} - #[derive(Clone, Debug, PartialEq, Eq)] pub struct Turn { pub role: TurnRole, diff --git a/litellm-rust/crates/gateway-mcp/src/configured.rs b/litellm-rust/crates/gateway-mcp/src/configured.rs index 06a1779a274..bb4f5594f1f 100644 --- a/litellm-rust/crates/gateway-mcp/src/configured.rs +++ b/litellm-rust/crates/gateway-mcp/src/configured.rs @@ -49,8 +49,8 @@ fn info(name: &str, config: &McpServer) -> ServerInfo { let identity = format!( "{name}|{}|{}|{}|{}", config.url.as_ref().map_or("", SecretValue::expose), - config.transport.as_str(), - config.auth_type.map_or("", McpAuth::as_str), + <&'static str>::from(config.transport), + config.auth_type.map_or("", <&'static str>::from), config.alias.as_deref().unwrap_or("") ); ServerInfo { diff --git a/litellm-rust/crates/http/Cargo.toml b/litellm-rust/crates/http/Cargo.toml index 3d880be49f9..7d64467f2f7 100644 --- a/litellm-rust/crates/http/Cargo.toml +++ b/litellm-rust/crates/http/Cargo.toml @@ -21,6 +21,7 @@ rustls-native-certs.workspace = true tokio-tungstenite.workspace = true serde.workspace = true serde_json.workspace = true +strum.workspace = true thiserror.workspace = true tokio.workspace = true tracing.workspace = true diff --git a/litellm-rust/crates/http/src/tls.rs b/litellm-rust/crates/http/src/tls.rs index c58076607e4..0b1c44bf6fe 100644 --- a/litellm-rust/crates/http/src/tls.rs +++ b/litellm-rust/crates/http/src/tls.rs @@ -42,30 +42,28 @@ impl KeyExchangeGroup { } } -#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)] +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord, strum::EnumString)] +#[strum( + parse_err_ty = Unsupported, + parse_err_fn = unsupported_cipher_token +)] pub enum Tls12CipherSuite { + #[strum(serialize = "ECDHE-ECDSA-AES128-GCM-SHA256")] EcdheEcdsaAes128Gcm, + #[strum(serialize = "ECDHE-ECDSA-AES256-GCM-SHA384")] EcdheEcdsaAes256Gcm, + #[strum(serialize = "ECDHE-ECDSA-CHACHA20-POLY1305")] EcdheEcdsaChacha20, + #[strum(serialize = "ECDHE-RSA-AES128-GCM-SHA256")] EcdheRsaAes128Gcm, + #[strum(serialize = "ECDHE-RSA-AES256-GCM-SHA384")] EcdheRsaAes256Gcm, + #[strum(serialize = "ECDHE-RSA-CHACHA20-POLY1305")] EcdheRsaChacha20, } -impl FromStr for Tls12CipherSuite { - type Err = Unsupported; - - fn from_str(name: &str) -> Result { - match name { - "ECDHE-ECDSA-AES128-GCM-SHA256" => Ok(Self::EcdheEcdsaAes128Gcm), - "ECDHE-ECDSA-AES256-GCM-SHA384" => Ok(Self::EcdheEcdsaAes256Gcm), - "ECDHE-ECDSA-CHACHA20-POLY1305" => Ok(Self::EcdheEcdsaChacha20), - "ECDHE-RSA-AES128-GCM-SHA256" => Ok(Self::EcdheRsaAes128Gcm), - "ECDHE-RSA-AES256-GCM-SHA384" => Ok(Self::EcdheRsaAes256Gcm), - "ECDHE-RSA-CHACHA20-POLY1305" => Ok(Self::EcdheRsaChacha20), - _ => Err(Unsupported::CipherToken(name.to_owned())), - } - } +fn unsupported_cipher_token(name: &str) -> Unsupported { + Unsupported::CipherToken(name.to_owned()) } impl Tls12CipherSuite { @@ -367,6 +365,39 @@ mod tests { ); } + #[rstest] + #[case::ecdhe_ecdsa_aes128( + "ECDHE-ECDSA-AES128-GCM-SHA256", + Tls12CipherSuite::EcdheEcdsaAes128Gcm + )] + #[case::ecdhe_ecdsa_aes256( + "ECDHE-ECDSA-AES256-GCM-SHA384", + Tls12CipherSuite::EcdheEcdsaAes256Gcm + )] + #[case::ecdhe_ecdsa_chacha20( + "ECDHE-ECDSA-CHACHA20-POLY1305", + Tls12CipherSuite::EcdheEcdsaChacha20 + )] + #[case::ecdhe_rsa_aes128("ECDHE-RSA-AES128-GCM-SHA256", Tls12CipherSuite::EcdheRsaAes128Gcm)] + #[case::ecdhe_rsa_aes256("ECDHE-RSA-AES256-GCM-SHA384", Tls12CipherSuite::EcdheRsaAes256Gcm)] + #[case::ecdhe_rsa_chacha20("ECDHE-RSA-CHACHA20-POLY1305", Tls12CipherSuite::EcdheRsaChacha20)] + fn tls12_cipher_suite_parses_each_openssl_name( + #[case] name: &str, + #[case] expected: Tls12CipherSuite, + ) { + assert_eq!(name.parse::().unwrap(), expected); + } + + #[rstest] + #[case::unknown("AES128-SHA")] + #[case::case_sensitive("ecdhe-rsa-aes256-gcm-sha384")] + fn unsupported_cipher_token_keeps_the_verbatim_name(#[case] name: &str) { + assert_eq!( + name.parse::().unwrap_err(), + Unsupported::CipherToken(name.to_owned()) + ); + } + #[test] fn named_suites_are_the_only_tls12_suites_offered_and_tls13_stays() { let tls = ClientConfig::try_from(&config(HttpSettings { diff --git a/litellm-rust/crates/llms-types/src/formats/chat_completions.rs b/litellm-rust/crates/llms-types/src/formats/chat_completions.rs index ca4ffa3c718..8f154dd7fb5 100644 --- a/litellm-rust/crates/llms-types/src/formats/chat_completions.rs +++ b/litellm-rust/crates/llms-types/src/formats/chat_completions.rs @@ -9,41 +9,25 @@ pub use content::{ /// Reasoning effort level accepted or applied by the model. #[macro_rules_attribute::apply(crate::wire_type)] -#[derive(Copy, Eq, IntoStaticStr)] +#[derive(Copy, Eq, IntoStaticStr, strum::EnumString, strum::VariantArray)] #[serde(rename_all = "snake_case")] -#[strum(serialize_all = "snake_case")] pub enum ReasoningEffort { + #[strum(serialize = "none")] None, + #[strum(serialize = "minimal")] Minimal, + #[strum(serialize = "low")] Low, + #[strum(serialize = "medium")] Medium, + #[strum(serialize = "high")] High, + #[strum(serialize = "xhigh")] Xhigh, + #[strum(serialize = "max")] Max, } -impl ReasoningEffort { - pub const ALL: [Self; 7] = [ - Self::None, - Self::Minimal, - Self::Low, - Self::Medium, - Self::High, - Self::Xhigh, - Self::Max, - ]; - - pub fn as_str(self) -> &'static str { - self.into() - } - - pub fn parse(value: &str) -> Option { - Self::ALL - .into_iter() - .find(|effort| effort.as_str() == value) - } -} - #[macro_rules_attribute::apply(crate::wire_type)] #[serde(untagged)] pub enum ChatMessageContent { @@ -195,6 +179,7 @@ pub struct ChatCompletionChunk { #[cfg(test)] mod tests { use rstest::rstest; + use strum::VariantArray; use super::*; @@ -213,10 +198,10 @@ mod tests { ) { assert_eq!( serde_json::to_value(effort).unwrap(), - Value::String(effort.as_str().to_string()) + Value::String(<&'static str>::from(effort).to_string()) ); - assert_eq!(ReasoningEffort::parse(effort.as_str()), Some(effort)); - assert!(ReasoningEffort::ALL.contains(&effort)); + assert_eq!(<&'static str>::from(effort).parse(), Ok(effort)); + assert!(ReasoningEffort::VARIANTS.contains(&effort)); } #[rstest] @@ -224,6 +209,6 @@ mod tests { #[case::uppercase("HIGH")] #[case::empty("")] fn reasoning_effort_parse_rejects(#[case] value: &str) { - assert_eq!(ReasoningEffort::parse(value), None); + assert!(value.parse::().is_err()); } } diff --git a/litellm-rust/crates/llms-types/src/formats/messages/metadata.rs b/litellm-rust/crates/llms-types/src/formats/messages/metadata.rs index 9db13cab498..6e5ebb2b98d 100644 --- a/litellm-rust/crates/llms-types/src/formats/messages/metadata.rs +++ b/litellm-rust/crates/llms-types/src/formats/messages/metadata.rs @@ -165,10 +165,12 @@ pub struct Safeguard { )] #[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] #[cfg_attr(feature = "schema", schemars(from = "String", into = "String"))] -#[strum(serialize_all = "snake_case")] pub enum MessageRole { + #[strum(serialize = "user")] User, + #[strum(serialize = "assistant")] Assistant, + #[strum(serialize = "system")] System, #[strum(default, transparent)] Other(String), @@ -205,8 +207,8 @@ impl From for String { )] #[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] #[cfg_attr(feature = "schema", schemars(from = "String", into = "String"))] -#[strum(serialize_all = "snake_case")] pub enum MessageType { + #[strum(serialize = "message")] Message, #[strum(default, transparent)] Other(String), @@ -243,15 +245,22 @@ impl From for String { )] #[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] #[cfg_attr(feature = "schema", schemars(from = "String", into = "String"))] -#[strum(serialize_all = "snake_case")] pub enum StopReason { + #[strum(serialize = "end_turn")] EndTurn, + #[strum(serialize = "max_tokens")] MaxTokens, + #[strum(serialize = "stop_sequence")] StopSequence, + #[strum(serialize = "tool_use")] ToolUse, + #[strum(serialize = "refusal")] Refusal, + #[strum(serialize = "compaction")] Compaction, + #[strum(serialize = "pause_turn")] PauseTurn, + #[strum(serialize = "model_context_window_exceeded")] ModelContextWindowExceeded, #[strum(default, transparent)] Other(String), diff --git a/litellm-rust/crates/llms-types/src/formats/messages/request.rs b/litellm-rust/crates/llms-types/src/formats/messages/request.rs index c4a31142fb6..d065f45d41c 100644 --- a/litellm-rust/crates/llms-types/src/formats/messages/request.rs +++ b/litellm-rust/crates/llms-types/src/formats/messages/request.rs @@ -30,16 +30,24 @@ pub enum MessageContent { )] #[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] #[cfg_attr(feature = "schema", schemars(from = "String", into = "String"))] -#[strum(serialize_all = "snake_case")] pub enum ContentBlockType { + #[strum(serialize = "text")] Text, + #[strum(serialize = "thinking")] Thinking, + #[strum(serialize = "redacted_thinking")] RedactedThinking, + #[strum(serialize = "tool_use")] ToolUse, + #[strum(serialize = "server_tool_use")] ServerToolUse, + #[strum(serialize = "tool_result")] ToolResult, + #[strum(serialize = "compaction")] Compaction, + #[strum(serialize = "advisor_tool_result")] AdvisorToolResult, + #[strum(serialize = "web_search_tool_result")] WebSearchToolResult, #[strum(default, transparent)] Other(String), @@ -124,23 +132,21 @@ pub struct Message { } #[macro_rules_attribute::apply(crate::wire_type)] -#[derive(Copy, Hash, IntoStaticStr, Eq)] +#[derive(Copy, Hash, IntoStaticStr, Eq, strum::VariantArray)] #[serde(rename_all = "lowercase")] -#[strum(serialize_all = "lowercase")] pub enum EffortLevel { + #[strum(serialize = "low")] Low, + #[strum(serialize = "medium")] Medium, + #[strum(serialize = "high")] High, + #[strum(serialize = "xhigh")] Xhigh, + #[strum(serialize = "max")] Max, } -impl EffortLevel { - pub fn as_str(self) -> &'static str { - self.into() - } -} - impl From for ReasoningEffort { fn from(level: EffortLevel) -> Self { match level { @@ -156,18 +162,13 @@ impl From for ReasoningEffort { #[macro_rules_attribute::apply(crate::wire_type)] #[derive(Copy, IntoStaticStr, Eq)] #[serde(rename_all = "lowercase")] -#[strum(serialize_all = "lowercase")] pub enum Speed { + #[strum(serialize = "fast")] Fast, + #[strum(serialize = "standard")] Standard, } -impl Speed { - pub fn as_str(self) -> &'static str { - self.into() - } -} - /// The tools whose presence changes how the request is sent. Every other tool, custom or /// server, deserializes as `Recognized::Unrecognized` and passes through verbatim. #[macro_rules_attribute::apply(crate::wire_type)] @@ -691,7 +692,10 @@ mod tests { #[rstest] fn speed_names_match_the_wire(#[values(Speed::Fast, Speed::Standard)] speed: Speed) { - assert_eq!(serde_json::to_value(speed).unwrap(), json!(speed.as_str())); + assert_eq!( + serde_json::to_value(speed).unwrap(), + json!(<&'static str>::from(speed)) + ); } #[rstest] @@ -705,10 +709,13 @@ mod tests { )] level: EffortLevel, ) { - assert_eq!(serde_json::to_value(level).unwrap(), json!(level.as_str())); + assert_eq!( + serde_json::to_value(level).unwrap(), + json!(<&'static str>::from(level)) + ); assert_eq!( serde_json::to_value(ReasoningEffort::from(level)).unwrap(), - json!(level.as_str()) + json!(<&'static str>::from(level)) ); } } diff --git a/litellm-rust/crates/llms-types/src/providers/anthropic.rs b/litellm-rust/crates/llms-types/src/providers/anthropic.rs index c2dcefce082..c7e58f9a60a 100644 --- a/litellm-rust/crates/llms-types/src/providers/anthropic.rs +++ b/litellm-rust/crates/llms-types/src/providers/anthropic.rs @@ -7,34 +7,39 @@ use std::{ str::FromStr, }; +use strum::VariantArray; + /// A provider column of `litellm/anthropic_beta_headers_config.json`: which betas a host accepts /// and under which name. #[derive( - Clone, Copy, Debug, PartialEq, Eq, Hash, strum::AsRefStr, strum::Display, strum::EnumString, + Clone, + Copy, + Debug, + PartialEq, + Eq, + Hash, + strum::AsRefStr, + strum::Display, + strum::EnumString, + strum::VariantArray, )] -#[strum(serialize_all = "snake_case")] pub enum BetaProvider { + #[strum(serialize = "anthropic")] Anthropic, + #[strum(serialize = "azure_ai")] AzureAi, + #[strum(serialize = "bedrock_converse")] BedrockConverse, + #[strum(serialize = "bedrock")] Bedrock, + #[strum(serialize = "bedrock_mantle")] BedrockMantle, + #[strum(serialize = "vertex_ai")] VertexAi, + #[strum(serialize = "databricks")] Databricks, } -impl BetaProvider { - pub const ALL: &[Self] = &[ - Self::Anthropic, - Self::AzureAi, - Self::BedrockConverse, - Self::Bedrock, - Self::BedrockMantle, - Self::VertexAi, - Self::Databricks, - ]; -} - /// One value of the `anthropic-beta` header. Equality, ordering and hashing follow the wire /// string, so a value parsed from a caller's header never disagrees with the matching variant. #[derive(Clone, Debug, strum::AsRefStr, strum::Display, strum::EnumString)] @@ -203,7 +208,7 @@ impl AnthropicBeta { | Self::McpServers20251204 | Self::StructuredOutput20240301 | Self::TextEditor20241022 - | Self::TextEditor20250124 => BetaProvider::ALL, + | Self::TextEditor20250124 => BetaProvider::VARIANTS, Self::ClaudeCode20250219 | Self::ToolExamples20251029 => &[ BetaProvider::Anthropic, BetaProvider::AzureAi, @@ -271,7 +276,7 @@ impl AnthropicBeta { BetaProvider::BedrockConverse, BetaProvider::Databricks, ], - Self::Other(_) => BetaProvider::ALL, + Self::Other(_) => BetaProvider::VARIANTS, } } } @@ -462,15 +467,17 @@ mod tests { #[case::bedrock_mantle(BetaProvider::BedrockMantle)] #[case::vertex_ai(BetaProvider::VertexAi)] #[case::databricks(BetaProvider::Databricks)] - fn provider_columns_match_beta_provider_all(#[case] provider: BetaProvider) { + fn provider_columns_match_beta_provider_variants(#[case] provider: BetaProvider) { let config = beta_headers_config(); let config_columns: Vec = config .keys() .filter(|column| column.as_str() != "description") .cloned() .collect(); - let beta_provider_columns: Vec = - BetaProvider::ALL.iter().map(ToString::to_string).collect(); + let beta_provider_columns: Vec = BetaProvider::VARIANTS + .iter() + .map(ToString::to_string) + .collect(); assert_eq!(config_columns, beta_provider_columns); assert_eq!( provider.to_string().parse::().unwrap(), @@ -481,7 +488,7 @@ mod tests { #[rstest] fn on_matches_every_config_cell() { let config = beta_headers_config(); - for provider in BetaProvider::ALL { + for provider in BetaProvider::VARIANTS { for beta in &AnthropicBeta::KNOWN { let expected = config .get(provider.as_ref()) diff --git a/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs b/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs index ecc6cfaf83d..e1fd1956167 100644 --- a/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs @@ -241,7 +241,7 @@ fn anthropic_body( .iter() .map(|turn| { json!({ - "role": turn.role.as_str(), + "role": <&'static str>::from(turn.role), "content": turn.texts.iter().map(|text| text_block(text)).collect::>(), }) }) diff --git a/litellm-rust/crates/llms/src/anthropic/common_utils.rs b/litellm-rust/crates/llms/src/anthropic/common_utils.rs index 6e2fa3b8785..ccf3a6d15f2 100644 --- a/litellm-rust/crates/llms/src/anthropic/common_utils.rs +++ b/litellm-rust/crates/llms/src/anthropic/common_utils.rs @@ -596,17 +596,10 @@ mod tests { use crate::base_llm::messages::context::SupportedEffortTiers; use rstest::{fixture, rstest}; use serde_json::json; + use strum::VariantArray; use super::*; - const ALL_LEVELS: [EffortLevel; 5] = [ - EffortLevel::Low, - EffortLevel::Medium, - EffortLevel::High, - EffortLevel::Xhigh, - EffortLevel::Max, - ]; - fn apply(sanitizer: fn(Vec) -> Vec, messages: Value) -> Value { let parsed: Vec = serde_json::from_value(messages).unwrap(); serde_json::to_value(sanitizer(parsed)).unwrap() @@ -1712,7 +1705,10 @@ mod tests { ..unmapped }; assert_eq!( - ALL_LEVELS.map(|level| supports_effort_tier(&capabilities, level)), + EffortLevel::VARIANTS + .iter() + .map(|level| supports_effort_tier(&capabilities, *level)) + .collect::>(), expected ); } diff --git a/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs b/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs index 2162c39c229..df1ae13fd28 100644 --- a/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs +++ b/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs @@ -10,6 +10,7 @@ use litellm_llms_types::{ }; use litellm_python_compat::{json::from_json, repr::repr, truthy::truthy}; use serde_json::Value; +use strum::VariantArray; use crate::base_llm::messages::context::{ MessagesModelCapabilities, ThinkingBudgets, ThinkingContext, @@ -54,14 +55,17 @@ fn unmapped_effort(effort: &Value) -> Error { Error::InvalidRequest(crate::ErrorDetail::InvalidChoice { field: "reasoning effort", actual: repr(&from_json(effort.clone())), - choices: ReasoningEffort::ALL.map(|effort| effort.as_str()).into(), + choices: ReasoningEffort::VARIANTS + .iter() + .map(|effort| <&'static str>::from(*effort)) + .collect(), }) } fn unsupported_effort(level: EffortLevel, model: &str) -> Error { Error::InvalidRequest(crate::ErrorDetail::UnsupportedValue { field: "effort", - value: level.as_str(), + value: <&'static str>::from(level), model: model.into(), }) } @@ -134,7 +138,7 @@ fn legacy_reasoning_effort( Some(Recognized::Known(level)) => Ok((*level).into()), Some(Recognized::Unrecognized(value)) if truthy(&from_json(value.clone())) => value .as_str() - .and_then(ReasoningEffort::parse) + .and_then(|text| text.parse().ok()) .ok_or_else(|| unmapped_effort(value)), None | Some(Recognized::Unrecognized(_)) => Ok(ReasoningEffort::Medium), } diff --git a/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs b/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs index b9cf6c37272..74ac35752ab 100644 --- a/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs @@ -270,7 +270,7 @@ fn drop_unsupported_params( fn speed_text(speed: &Recognized) -> String { match speed { - Recognized::Known(speed) => speed.as_str().to_string(), + Recognized::Known(speed) => <&'static str>::from(*speed).to_string(), Recognized::Unrecognized(Value::String(text)) => text.clone(), Recognized::Unrecognized(other) => other.to_string(), } diff --git a/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs b/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs index 90f5b97322f..4139ee860cf 100644 --- a/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs +++ b/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs @@ -23,9 +23,11 @@ const HEALTH_CHECK_IMAGE_DATA_URI: &str = "data:image/png;base64,iVBORw0KGgoAAAA /// Textract has operations rather than models; the model slot of /// `aws_textract/` names the one to call. #[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, VariantNames, PartialEq, Eq)] -#[strum(serialize_all = "kebab-case", ascii_case_insensitive)] +#[strum(ascii_case_insensitive)] pub enum TextractOperation { + #[strum(serialize = "detect-document-text")] DetectDocumentText, + #[strum(serialize = "analyze-document")] AnalyzeDocument, } diff --git a/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs b/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs index cee906e77a4..c0b2c7aa3b7 100644 --- a/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs +++ b/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs @@ -30,20 +30,17 @@ pub struct BedrockAudioTranscriptionConfig; #[derive(Clone, Copy, Deserialize, IntoStaticStr)] #[serde(rename_all = "lowercase")] -#[strum(serialize_all = "lowercase")] enum AudioFormat { + #[strum(serialize = "wav")] Wav, + #[strum(serialize = "mp3")] Mp3, + #[strum(serialize = "flac")] Flac, + #[strum(serialize = "ogg")] Ogg, } -impl AudioFormat { - fn as_str(self) -> &'static str { - self.into() - } -} - struct AudioInput { data: String, format: AudioFormat, @@ -118,7 +115,7 @@ impl BaseAudioTranscriptionConfig for BedrockAudioTranscriptionConfig { "messages": [{ "role": "user", "content": [ - {"audio": {"format": audio.format.as_str(), "source": {"bytes": audio.data}}}, + {"audio": {"format": <&'static str>::from(audio.format), "source": {"bytes": audio.data}}}, {"text": instruction} ] }], diff --git a/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs b/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs index 3a13e388a4b..bcd1077b830 100644 --- a/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs +++ b/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs @@ -385,7 +385,7 @@ fn converse_body(conversation: &Conversation, optional_params: &Map::from(turn.role), "content": turn.texts.iter().map(|text| json!({"text": text})).collect::>(), }) }) diff --git a/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs b/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs index ad7025ab6af..373ad4ca4f5 100644 --- a/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs +++ b/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs @@ -64,27 +64,23 @@ pub fn invoke_anthropic_event_stream(bytes: ByteStream) -> EventStream { Box::pin(invoke_chunk_stream(bytes).map(|chunk| decode_invoke_anthropic_chunk(chunk?))) } -#[derive(Clone, Copy)] +#[derive(Clone, Copy, Debug, PartialEq, Eq, strum::EnumString)] enum InvokeProvider { + #[strum(serialize = "anthropic")] Anthropic, + #[strum(serialize = "deepseek_r1")] DeepseekR1, + #[strum(serialize = "moonshot")] Moonshot, + #[strum(disabled)] Unsupported, } -impl From<&str> for InvokeProvider { - fn from(value: &str) -> Self { - match value { - "anthropic" => Self::Anthropic, - "deepseek_r1" => Self::DeepseekR1, - "moonshot" => Self::Moonshot, - _ => Self::Unsupported, - } - } -} - pub fn invoke_chat_stream(invoke_provider: &str, shape: StreamShape) -> Result { - match InvokeProvider::from(invoke_provider) { + match invoke_provider + .parse() + .unwrap_or(InvokeProvider::Unsupported) + { InvokeProvider::Anthropic => Ok(ChatStream::new( invoke_anthropic_event_stream, ModelResponseIterator::new(shape), @@ -116,6 +112,22 @@ mod tests { const TEXT_DELTA: &str = r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hello"}}"#; + #[rstest::rstest] + #[case::anthropic("anthropic", InvokeProvider::Anthropic)] + #[case::deepseek_r1("deepseek_r1", InvokeProvider::DeepseekR1)] + #[case::moonshot("moonshot", InvokeProvider::Moonshot)] + #[case::unknown("qwen", InvokeProvider::Unsupported)] + #[case::case_sensitive("Anthropic", InvokeProvider::Unsupported)] + fn invoke_provider_maps_each_model_family_name( + #[case] name: &str, + #[case] expected: InvokeProvider, + ) { + assert_eq!( + name.parse().unwrap_or(InvokeProvider::Unsupported), + expected + ); + } + fn aws_wire(chunk: &str) -> Vec { let payload = serde_json::json!({"bytes": STANDARD.encode(chunk)}); let message = Message::new(Bytes::from(serde_json::to_vec(&payload).unwrap())).add_header( diff --git a/litellm-rust/crates/llms/src/cohere/ocr/transformation.rs b/litellm-rust/crates/llms/src/cohere/ocr/transformation.rs index 67b81ec6feb..dadab413fe4 100644 --- a/litellm-rust/crates/llms/src/cohere/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/cohere/ocr/transformation.rs @@ -25,11 +25,13 @@ const COHERE_API_KEY_ENV: &str = "COHERE_API_KEY"; const COHERE_PARSE_HEALTH_CHECK_IMAGE_DATA_URI: &str = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGP4//8/AAX+Av4N70a4AAAAAElFTkSuQmCC"; -#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)] +#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize, strum::IntoStaticStr)] #[serde(rename_all = "lowercase")] pub enum OutputFormat { #[default] + #[strum(serialize = "markdown")] Markdown, + #[strum(serialize = "blocks")] Blocks, } @@ -267,11 +269,7 @@ fn build_request(model: &str, image_url: String, params: &CohereOptions) -> Cohe CohereRequest { model: model.into(), document: CohereParseDocument::ImageUrl { image_url }, - output_format: match params.output_format.unwrap_or_default() { - OutputFormat::Markdown => "markdown", - OutputFormat::Blocks => "blocks", - } - .into(), + output_format: <&'static str>::from(params.output_format.unwrap_or_default()).into(), } } diff --git a/litellm-rust/crates/python-bridge/src/python_settings.rs b/litellm-rust/crates/python-bridge/src/python_settings.rs index 8f49444f730..99afc772b99 100644 --- a/litellm-rust/crates/python-bridge/src/python_settings.rs +++ b/litellm-rust/crates/python-bridge/src/python_settings.rs @@ -6,13 +6,16 @@ use crate::coercion::{FieldSpec, ProjectionError}; const MODULE: &str = "litellm.rust_bridge.settings"; #[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq)] -#[strum(serialize_all = "snake_case")] pub(crate) enum PythonSettings { #[strum(serialize = "http_settings")] Http, + #[strum(serialize = "url_policy")] UrlPolicy, + #[strum(serialize = "provider_defaults")] ProviderDefaults, + #[strum(serialize = "secret_manager")] SecretManager, + #[strum(serialize = "secret_manager_binding")] SecretManagerBinding, } @@ -23,17 +26,16 @@ pub(crate) struct Snapshot<'py> { impl Snapshot<'_> { pub(crate) fn read(&self, spec: &FieldSpec) -> Result { - spec.read(&self.value, self.group.name()) + spec.read(&self.value, self.group.into()) } } impl PythonSettings { - pub(crate) fn name(self) -> &'static str { - self.into() - } - pub(crate) fn read(self, py: Python<'_>) -> PyResult> { - let value = py.import(MODULE)?.getattr(self.name())?.call0()?; + let value = py + .import(MODULE)? + .getattr(<&'static str>::from(self))? + .call0()?; Ok(Snapshot { group: self, value }) } diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index 61066e322a7..413245c1ec2 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -172,7 +172,7 @@ impl NativeTraceStorage { BTreeMap, >, ) -> PyResult> { - let table = InsertTable::parse(table).map_err(map_error)?; + let table = table.parse::().map_err(map_error)?; let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; let connection = self.config.storage().writer().clone(); let database = self.config.storage().database().to_owned(); diff --git a/litellm-rust/crates/secrets-types/src/config.rs b/litellm-rust/crates/secrets-types/src/config.rs index 82d48f7b2e1..9b6e82c3a6b 100644 --- a/litellm-rust/crates/secrets-types/src/config.rs +++ b/litellm-rust/crates/secrets-types/src/config.rs @@ -5,18 +5,37 @@ use strum::IntoStaticStr; use crate::SecretValue; -#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, IntoStaticStr, PartialEq, Serialize)] +#[derive( + Clone, + Copy, + Debug, + Deserialize, + Eq, + Hash, + IntoStaticStr, + PartialEq, + Serialize, + strum::VariantArray, +)] #[serde(rename_all = "snake_case")] -#[strum(serialize_all = "snake_case")] pub enum KeyManagementSystem { + #[strum(serialize = "google_kms")] GoogleKms, + #[strum(serialize = "azure_key_vault")] AzureKeyVault, + #[strum(serialize = "aws_secret_manager")] AwsSecretManager, + #[strum(serialize = "google_secret_manager")] GoogleSecretManager, + #[strum(serialize = "hashicorp_vault")] HashicorpVault, + #[strum(serialize = "cyberark")] Cyberark, + #[strum(serialize = "local")] Local, + #[strum(serialize = "aws_kms")] AwsKms, + #[strum(serialize = "custom")] Custom, } diff --git a/litellm-rust/crates/secrets-types/tests/context.rs b/litellm-rust/crates/secrets-types/tests/context.rs index a167e87954f..8f8eeabdb62 100644 --- a/litellm-rust/crates/secrets-types/tests/context.rs +++ b/litellm-rust/crates/secrets-types/tests/context.rs @@ -128,17 +128,21 @@ fn rotation_write_context_preserves_the_operation_context(aws_context: SecretOpe fn provider_context_accepts_only_its_owner( #[case] owner: KeyManagementSystem, #[case] context: SecretOperationContext, -) { - for system in [ - KeyManagementSystem::AwsSecretManager, + #[values( + KeyManagementSystem::GoogleKms, KeyManagementSystem::AzureKeyVault, + KeyManagementSystem::AwsSecretManager, KeyManagementSystem::GoogleSecretManager, KeyManagementSystem::HashicorpVault, KeyManagementSystem::Cyberark, - ] { - assert_eq!(context.validate_for(system).is_ok(), system == owner); - assert!(SecretOperationContext::Default.validate_for(system).is_ok()); - } + KeyManagementSystem::Local, + KeyManagementSystem::AwsKms, + KeyManagementSystem::Custom + )] + system: KeyManagementSystem, +) { + assert_eq!(context.validate_for(system).is_ok(), system == owner); + assert!(SecretOperationContext::Default.validate_for(system).is_ok()); if matches!( owner, KeyManagementSystem::AzureKeyVault | KeyManagementSystem::GoogleSecretManager diff --git a/litellm-rust/crates/secrets/src/oidc.rs b/litellm-rust/crates/secrets/src/oidc.rs index b6e8dbc123b..adeb6c3db00 100644 --- a/litellm-rust/crates/secrets/src/oidc.rs +++ b/litellm-rust/crates/secrets/src/oidc.rs @@ -23,17 +23,22 @@ const OIDC_ALLOWED_CREDENTIAL_DIRS: &str = "LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS const DEFAULT_CREDENTIAL_DIRS: &str = "/var/run/secrets,/run/secrets"; #[derive(Clone, Copy, Debug, Eq, Hash, PartialEq, strum::EnumString, strum::AsRefStr)] -#[strum(serialize_all = "snake_case")] pub enum OidcProvider { + #[strum(serialize = "google")] Google, #[strum(serialize = "circleci")] CircleCi, #[strum(serialize = "circleci_v2")] CircleCiV2, + #[strum(serialize = "github")] Github, + #[strum(serialize = "azure")] Azure, + #[strum(serialize = "file")] File, + #[strum(serialize = "env")] Env, + #[strum(serialize = "env_path")] EnvPath, } diff --git a/litellm-rust/crates/testkit/Cargo.toml b/litellm-rust/crates/testkit/Cargo.toml index 98a36a1e87f..97efb753384 100644 --- a/litellm-rust/crates/testkit/Cargo.toml +++ b/litellm-rust/crates/testkit/Cargo.toml @@ -13,6 +13,7 @@ serde.workspace = true semver.workspace = true serde_json.workspace = true sha2.workspace = true +strum.workspace = true tar.workspace = true target-lexicon.workspace = true thiserror.workspace = true diff --git a/litellm-rust/crates/testkit/src/agent/configure.rs b/litellm-rust/crates/testkit/src/agent/configure.rs index a095aacc92f..a41bec9379a 100644 --- a/litellm-rust/crates/testkit/src/agent/configure.rs +++ b/litellm-rust/crates/testkit/src/agent/configure.rs @@ -5,7 +5,7 @@ use semver::Version; use crate::Error; -#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[derive(Clone, Copy, Debug, PartialEq, Eq, strum::VariantArray)] pub enum Wire { ChatCompletions, Messages, diff --git a/litellm-rust/crates/testkit/tests/configure.rs b/litellm-rust/crates/testkit/tests/configure.rs index ca3587c3474..87f53d2f9c0 100644 --- a/litellm-rust/crates/testkit/tests/configure.rs +++ b/litellm-rust/crates/testkit/tests/configure.rs @@ -2,6 +2,7 @@ use std::path::Path; use litellm_testkit::{ClaudeCode, Codex, Configure, Error, Opencode, Settings, Version, Wire}; use rstest::rstest; +use strum::VariantArray; fn settings(wire: Wire) -> Settings { Settings { @@ -121,7 +122,10 @@ fn opencode_uses_a_different_provider_package_for_every_wire() { .unwrap() .to_owned() }; - let packages = [Wire::ChatCompletions, Wire::Responses, Wire::Messages].map(package); + let packages = Wire::VARIANTS + .iter() + .map(|wire| package(*wire)) + .collect::>(); assert_eq!( packages diff --git a/litellm-rust/crates/traces-clickhouse/src/insert.rs b/litellm-rust/crates/traces-clickhouse/src/insert.rs index ed8db188c14..18c077977cc 100644 --- a/litellm-rust/crates/traces-clickhouse/src/insert.rs +++ b/litellm-rust/crates/traces-clickhouse/src/insert.rs @@ -30,29 +30,19 @@ fn max_insert_bytes() -> Result { pub type InsertRow = BTreeMap>; +#[derive(Clone, Copy, Debug, PartialEq, Eq, strum::EnumString, strum::IntoStaticStr)] +#[strum(parse_err_ty = Error, parse_err_fn = invalid_table)] pub enum InsertTable { + #[strum(serialize = "otel_traces")] OtelTraces, + #[strum(serialize = "spend_logs")] SpendLogs, + #[strum(serialize = "lens_feedback")] LensFeedback, } -impl InsertTable { - pub fn parse(value: &str) -> Result { - match value { - "otel_traces" => Ok(Self::OtelTraces), - "spend_logs" => Ok(Self::SpendLogs), - "lens_feedback" => Ok(Self::LensFeedback), - _ => Err(Error::InvalidTable), - } - } - - fn name(&self) -> &'static str { - match self { - Self::OtelTraces => "otel_traces", - Self::SpendLogs => "spend_logs", - Self::LensFeedback => "lens_feedback", - } - } +fn invalid_table(_name: &str) -> Error { + Error::InvalidTable } pub async fn insert_rows( @@ -81,7 +71,7 @@ pub async fn insert_shared_rows( client, connection, database, - table.name(), + <&'static str>::from(table), &token, body, ) @@ -242,7 +232,25 @@ mod tests { use rstest::rstest; use serde_json::json; - use super::{Error, shared_rows, write_rows}; + use super::{Error, InsertTable, shared_rows, write_rows}; + + #[rstest] + #[case::otel_traces("otel_traces", InsertTable::OtelTraces)] + #[case::spend_logs("spend_logs", InsertTable::SpendLogs)] + #[case::lens_feedback("lens_feedback", InsertTable::LensFeedback)] + fn insert_table_parses_each_table_name(#[case] name: &str, #[case] expected: InsertTable) { + assert_eq!(name.parse::().unwrap(), expected); + } + + #[rstest] + #[case::unknown("events")] + #[case::case_sensitive("OTEL_TRACES")] + fn insert_table_rejects_unknown_names(#[case] name: &str) { + assert!(matches!( + name.parse::(), + Err(Error::InvalidTable) + )); + } #[rstest] fn encoded_limit_counts_utf8_bytes_across_rows() { diff --git a/litellm-rust/crates/traces-clickhouse/src/query.rs b/litellm-rust/crates/traces-clickhouse/src/query.rs index a08affb0cad..cd1a96eef31 100644 --- a/litellm-rust/crates/traces-clickhouse/src/query.rs +++ b/litellm-rust/crates/traces-clickhouse/src/query.rs @@ -56,15 +56,21 @@ enum PathPart { #[macro_rules_attribute::apply(crate::response_type)] #[derive(Clone, Copy, Debug, Eq, Ord, PartialEq, PartialOrd, strum::Display)] #[serde(rename_all = "lowercase")] -#[strum(serialize_all = "lowercase")] #[cfg_attr(feature = "schema", schemars(rename = "MetadataValueType"))] enum JsonKind { + #[strum(serialize = "array")] Array, + #[strum(serialize = "boolean")] Boolean, + #[strum(serialize = "integer")] Integer, + #[strum(serialize = "null")] Null, + #[strum(serialize = "number")] Number, + #[strum(serialize = "object")] Object, + #[strum(serialize = "string")] String, } diff --git a/litellm-rust/crates/traces-clickhouse/src/table.rs b/litellm-rust/crates/traces-clickhouse/src/table.rs index 60ab4943f5e..c9039ad8bed 100644 --- a/litellm-rust/crates/traces-clickhouse/src/table.rs +++ b/litellm-rust/crates/traces-clickhouse/src/table.rs @@ -4,9 +4,11 @@ Clone, Copy, Debug, strum::Display, strum::AsRefStr, strum::EnumIter, strum::IntoStaticStr, )] #[serde(rename_all = "snake_case")] -#[strum(serialize_all = "snake_case")] pub enum TraceTable { + #[strum(serialize = "otel_traces")] OtelTraces, + #[strum(serialize = "agent_traces_by_key")] AgentTracesByKey, + #[strum(serialize = "spend_logs")] SpendLogs, } diff --git a/litellm-rust/crates/traces-clickhouse/tests/queries.rs b/litellm-rust/crates/traces-clickhouse/tests/queries.rs index 6e63adfc347..5cb24dc5548 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/queries.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/queries.rs @@ -126,10 +126,12 @@ async fn lens_sample_keeps_spans_before_window_start_and_excludes_old_only_trace } #[derive(Clone, Copy, strum::AsRefStr)] -#[strum(serialize_all = "snake_case")] enum ScopeCase { + #[strum(serialize = "admin")] Admin, + #[strum(serialize = "team")] Team, + #[strum(serialize = "other_team")] OtherTeam, } diff --git a/litellm-rust/crates/traces/src/normalize/format/claude_code.rs b/litellm-rust/crates/traces/src/normalize/format/claude_code.rs index 37fb2d01330..d849d0fea87 100644 --- a/litellm-rust/crates/traces/src/normalize/format/claude_code.rs +++ b/litellm-rust/crates/traces/src/normalize/format/claude_code.rs @@ -15,14 +15,23 @@ use crate::{ /// Claude Code's built-in tracing, identified by its instrumentation scope. pub(crate) struct ClaudeCode; +#[derive(Debug, PartialEq, Eq, strum::EnumString)] enum SpanType { + #[strum(serialize = "assistant_response")] AssistantResponse, + #[strum(serialize = "tool_result")] ToolResult, + #[strum(serialize = "api_request_body")] ApiRequestBody, + #[strum(serialize = "compaction")] Compaction, + #[strum(serialize = "interaction")] Interaction, + #[strum(serialize = "llm_request")] LlmRequest, + #[strum(serialize = "tool")] Tool, + #[strum(disabled)] Other, } @@ -33,16 +42,7 @@ fn span_type(name: &str, attributes: &BTreeMap) -> SpanType { } else { kind }; - match kind { - "assistant_response" => SpanType::AssistantResponse, - "tool_result" => SpanType::ToolResult, - "api_request_body" => SpanType::ApiRequestBody, - "compaction" => SpanType::Compaction, - "interaction" => SpanType::Interaction, - "llm_request" => SpanType::LlmRequest, - "tool" => SpanType::Tool, - _ => SpanType::Other, - } + kind.parse().unwrap_or(SpanType::Other) } /// `agent:custom:search_agent` -> `search_agent`: the subagent a request ran for. @@ -318,7 +318,7 @@ mod tests { use rstest::rstest; use serde_json::Value; - use super::CLAUDE_CODE_SCOPE; + use super::{CLAUDE_CODE_SCOPE, SpanType, span_type}; use crate::{ Error, normalize::{Normalization, NormalizedSpan, ObservationType}, @@ -355,6 +355,22 @@ mod tests { .collect() } + #[rstest] + #[case::assistant_response("assistant_response", SpanType::AssistantResponse)] + #[case::tool_result("tool_result", SpanType::ToolResult)] + #[case::api_request_body("api_request_body", SpanType::ApiRequestBody)] + #[case::compaction("compaction", SpanType::Compaction)] + #[case::interaction("interaction", SpanType::Interaction)] + #[case::llm_request("llm_request", SpanType::LlmRequest)] + #[case::tool("tool", SpanType::Tool)] + #[case::unknown("surprise", SpanType::Other)] + fn span_type_maps_each_recorded_kind(#[case] kind: &str, #[case] expected: SpanType) { + assert_eq!( + span_type("anything", &attributes(&[("span.type", kind)])), + expected + ); + } + #[rstest] fn notification_prompts_keep_user_provenance_and_compaction_is_system() { let prompt_text = diff --git a/litellm-rust/crates/traces/src/normalize/format/genai.rs b/litellm-rust/crates/traces/src/normalize/format/genai.rs index 5ba9b449737..2bfd6ce5370 100644 --- a/litellm-rust/crates/traces/src/normalize/format/genai.rs +++ b/litellm-rust/crates/traces/src/normalize/format/genai.rs @@ -11,18 +11,24 @@ use crate::{ pub(crate) struct GenAi; #[derive(strum::EnumString)] -#[strum(serialize_all = "snake_case")] pub(crate) enum Operation { + #[strum(serialize = "create_agent")] CreateAgent, + #[strum(serialize = "invoke_agent")] InvokeAgent, + #[strum(serialize = "invoke_workflow")] InvokeWorkflow, + #[strum(serialize = "chat")] Chat, #[strum(serialize = "text_completion", serialize = "completion")] TextCompletion, + #[strum(serialize = "generate_content")] GenerateContent, + #[strum(serialize = "execute_tool")] ExecuteTool, #[strum(serialize = "embeddings", serialize = "embedding")] Embeddings, + #[strum(serialize = "retrieval")] Retrieval, } diff --git a/litellm-rust/crates/traces/src/normalize/metadata.rs b/litellm-rust/crates/traces/src/normalize/metadata.rs index 100008714c5..600bfce4287 100644 --- a/litellm-rust/crates/traces/src/normalize/metadata.rs +++ b/litellm-rust/crates/traces/src/normalize/metadata.rs @@ -24,32 +24,56 @@ pub enum AgentType { serde_with::DeserializeFromStr, serde_with::SerializeDisplay, )] -#[strum(serialize_all = "kebab-case")] pub enum Integration { + #[strum(serialize = "claude-code")] ClaudeCode, + #[strum(serialize = "claude-agent-sdk")] ClaudeAgentSdk, + #[strum(serialize = "openai-codex")] OpenaiCodex, + #[strum(serialize = "deepagents-code")] DeepagentsCode, + #[strum(serialize = "cursor")] Cursor, + #[strum(serialize = "pi")] Pi, + #[strum(serialize = "opencode")] Opencode, + #[strum(serialize = "copilot")] Copilot, + #[strum(serialize = "langchain")] Langchain, + #[strum(serialize = "langgraph")] Langgraph, + #[strum(serialize = "deepagents")] Deepagents, + #[strum(serialize = "autogen")] Autogen, + #[strum(serialize = "crewai")] Crewai, + #[strum(serialize = "google-adk")] GoogleAdk, + #[strum(serialize = "llama-index")] LlamaIndex, + #[strum(serialize = "mastra")] Mastra, + #[strum(serialize = "microsoft-agent-framework")] MicrosoftAgentFramework, + #[strum(serialize = "openai-agents")] OpenaiAgents, + #[strum(serialize = "pydantic-ai")] PydanticAi, + #[strum(serialize = "semantic-kernel")] SemanticKernel, + #[strum(serialize = "strands")] Strands, + #[strum(serialize = "vercel-ai-sdk")] VercelAiSdk, + #[strum(serialize = "instructor")] Instructor, + #[strum(serialize = "n8n")] N8n, + #[strum(serialize = "temporal")] Temporal, #[strum(default)] Other(String), @@ -138,23 +162,36 @@ impl AgentMetadata { } #[derive(strum::EnumString, strum::IntoStaticStr)] -#[strum(serialize_all = "snake_case")] enum MetadataField { + #[strum(serialize = "lc_agent_name")] LcAgentName, + #[strum(serialize = "ls_integration")] LsIntegration, + #[strum(serialize = "ls_agent_type")] LsAgentType, + #[strum(serialize = "ls_agent_purpose")] LsAgentPurpose, + #[strum(serialize = "ls_agent_runtime")] LsAgentRuntime, #[strum(serialize = "ls_agent_runtime_version", to_string = "ls_agent_version")] LsAgentVersion, + #[strum(serialize = "ls_trace_schema_version")] LsTraceSchemaVersion, + #[strum(serialize = "thread_id")] ThreadId, + #[strum(serialize = "ls_subagent_id")] LsSubagentId, + #[strum(serialize = "ls_subagent_type")] LsSubagentType, + #[strum(serialize = "ls_tool_name")] LsToolName, + #[strum(serialize = "ls_model_name")] LsModelName, + #[strum(serialize = "ls_provider")] LsProvider, + #[strum(serialize = "git_branch")] GitBranch, + #[strum(serialize = "git_commit_sha")] GitCommitSha, #[strum(serialize = "repository_url", to_string = "git_repo_url")] GitRepoUrl, diff --git a/litellm-rust/crates/traces/src/normalize/mod.rs b/litellm-rust/crates/traces/src/normalize/mod.rs index 4d11e76e19b..894fcd1bb03 100644 --- a/litellm-rust/crates/traces/src/normalize/mod.rs +++ b/litellm-rust/crates/traces/src/normalize/mod.rs @@ -32,20 +32,32 @@ pub use metadata::{AgentMetadata, AgentType, Integration}; #[macro_rules_attribute::apply(crate::wire_type)] #[derive(Clone, Copy, Debug, Eq, PartialEq, strum::EnumString)] #[serde(rename_all = "lowercase")] -#[strum(serialize_all = "lowercase", ascii_case_insensitive)] +#[strum(ascii_case_insensitive)] #[cfg_attr(feature = "schema", schemars(rename = "SpanType"))] pub enum ObservationType { + #[strum(serialize = "agent")] Agent, + #[strum(serialize = "llm")] Llm, + #[strum(serialize = "tool")] Tool, + #[strum(serialize = "chain")] Chain, + #[strum(serialize = "framework")] Framework, + #[strum(serialize = "retriever")] Retriever, + #[strum(serialize = "embedding")] Embedding, + #[strum(serialize = "reranker")] Reranker, + #[strum(serialize = "guardrail")] Guardrail, + #[strum(serialize = "evaluator")] Evaluator, + #[strum(serialize = "prompt")] Prompt, + #[strum(serialize = "decision")] Decision, } diff --git a/litellm-rust/crates/traces/src/query.rs b/litellm-rust/crates/traces/src/query.rs index 26eaa33d5ac..33dcbb31abf 100644 --- a/litellm-rust/crates/traces/src/query.rs +++ b/litellm-rust/crates/traces/src/query.rs @@ -2,23 +2,38 @@ pub mod guide; pub mod named; #[derive(Clone, Copy, Debug, Eq, PartialEq, strum::EnumString, strum::Display, strum::AsRefStr)] -#[strum(serialize_all = "snake_case")] pub enum ReadQuery { + #[strum(serialize = "list_traces")] ListTraces, + #[strum(serialize = "trace_agents")] TraceAgents, + #[strum(serialize = "trace_spans")] TraceSpans, + #[strum(serialize = "trace_page_spans")] TracePageSpans, + #[strum(serialize = "trace_identity")] TraceIdentity, + #[strum(serialize = "span_detail")] SpanDetail, + #[strum(serialize = "span_error")] SpanError, + #[strum(serialize = "spend_by_response_ids")] SpendByResponseIds, + #[strum(serialize = "availability")] Availability, + #[strum(serialize = "agents")] Agents, + #[strum(serialize = "sample")] Sample, + #[strum(serialize = "content")] Content, + #[strum(serialize = "evidence")] Evidence, + #[strum(serialize = "feedback_target")] FeedbackTarget, + #[strum(serialize = "feedback")] Feedback, + #[strum(serialize = "feedback_summary")] FeedbackSummary, }