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>
This commit is contained in:
yujonglee 2026-10-08 14:24:49 -07:00 • committed by GitHub
parent 4ed0267b5a
commit 0a98c7be10
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
44 changed files with 490 additions and 270 deletions

View file

@ -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<Enum> 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

View file

@ -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;`

View file

@ -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",

View file

@ -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]

View file

@ -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> {
Self::ALL
.into_iter()
.find(|cache_type| cache_type.as_python_name() == value)
}
}

View file

@ -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::<CacheType>().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::<Vec<_>>(),
[
"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::<CacheType>().is_err());
}

View file

@ -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

View file

@ -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",
}
}
}

View file

@ -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::<McpTransport>(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::<McpAuth>(json!(name)).unwrap(),
auth
);
}

View file

@ -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,
}

View file

@ -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,

View file

@ -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 {

View file

@ -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

View file

@ -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<Self, Self::Err> {
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::<Tls12CipherSuite>().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::<Tls12CipherSuite>().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 {

View file

@ -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> {
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::<ReasoningEffort>().is_err());
}
}

View file

@ -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<MessageRole> 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<MessageType> 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),

View file

@ -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<EffortLevel> for ReasoningEffort {
fn from(level: EffortLevel) -> Self {
match level {
@ -156,18 +162,13 @@ impl From<EffortLevel> 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))
);
}
}

View file

@ -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<String> = config
.keys()
.filter(|column| column.as_str() != "description")
.cloned()
.collect();
let beta_provider_columns: Vec<String> =
BetaProvider::ALL.iter().map(ToString::to_string).collect();
let beta_provider_columns: Vec<String> = BetaProvider::VARIANTS
.iter()
.map(ToString::to_string)
.collect();
assert_eq!(config_columns, beta_provider_columns);
assert_eq!(
provider.to_string().parse::<BetaProvider>().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())

View file

@ -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::<Vec<_>>(),
})
})

View file

@ -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<Message>) -> Vec<Message>, messages: Value) -> Value {
let parsed: Vec<Message> = 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::<Vec<_>>(),
expected
);
}

View file

@ -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),
}

View file

@ -270,7 +270,7 @@ fn drop_unsupported_params(
fn speed_text(speed: &Recognized<Speed>) -> 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(),
}

View file

@ -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/<model>` 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,
}

View file

@ -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}
]
}],

View file

@ -385,7 +385,7 @@ fn converse_body(conversation: &Conversation, optional_params: &Map<String, Valu
.iter()
.map(|turn| {
json!({
"role": turn.role.as_str(),
"role": <&'static str>::from(turn.role),
"content": turn.texts.iter().map(|text| json!({"text": text})).collect::<Vec<_>>(),
})
})

View file

@ -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<ChatStream, Error> {
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<u8> {
let payload = serde_json::json!({"bytes": STANDARD.encode(chunk)});
let message = Message::new(Bytes::from(serde_json::to_vec(&payload).unwrap())).add_header(

View file

@ -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(),
}
}

View file

@ -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<T>(&self, spec: &FieldSpec<T>) -> Result<T, ProjectionError> {
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<Snapshot<'_>> {
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 })
}

View file

@ -172,7 +172,7 @@ impl NativeTraceStorage {
BTreeMap<String, serde_json::Value>,
>,
) -> PyResult<Bound<'py, PyAny>> {
let table = InsertTable::parse(table).map_err(map_error)?;
let table = table.parse::<InsertTable>().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();

View file

@ -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,
}

View file

@ -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

View file

@ -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,
}

View file

@ -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

View file

@ -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,

View file

@ -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::<Vec<_>>();
assert_eq!(
packages

View file

@ -30,29 +30,19 @@ fn max_insert_bytes() -> Result<usize, Error> {
pub type InsertRow = BTreeMap<String, Shared<Value>>;
#[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<Self, Error> {
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::<InsertTable>().unwrap(), expected);
}
#[rstest]
#[case::unknown("events")]
#[case::case_sensitive("OTEL_TRACES")]
fn insert_table_rejects_unknown_names(#[case] name: &str) {
assert!(matches!(
name.parse::<InsertTable>(),
Err(Error::InvalidTable)
));
}
#[rstest]
fn encoded_limit_counts_utf8_bytes_across_rows() {

View file

@ -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,
}

View file

@ -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,
}

View file

@ -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,
}

View file

@ -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<String, String>) -> 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 =

View file

@ -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,
}

View file

@ -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,

View file

@ -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,
}

View file

@ -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,
}