diff --git a/.github/workflows/test-rust.yml b/.github/workflows/test-rust.yml index 6f8599daf75..4bf41dc249c 100644 --- a/.github/workflows/test-rust.yml +++ b/.github/workflows/test-rust.yml @@ -130,7 +130,7 @@ jobs: - name: Test secret manager feature combinations run: | cargo test -p litellm-auth-gcp --locked --no-default-features - for features in '' aws google azure cyberark aws,google aws,azure google,azure aws,google,azure aws,google,cyberark aws,google,azure,cyberark; do + for features in '' aws google hashicorp azure cyberark aws,google aws,azure google,azure aws,google,azure aws,google,cyberark aws,google,azure,cyberark aws,google,hashicorp,azure,cyberark; do cargo test -p litellm-secrets --locked --no-default-features --features "$features" done diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260919000000_add_autorouter_user_session_rollup/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260919000000_add_autorouter_user_session_rollup/migration.sql new file mode 100644 index 00000000000..2b864131ab2 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260919000000_add_autorouter_user_session_rollup/migration.sql @@ -0,0 +1,42 @@ +CREATE TABLE IF NOT EXISTS "LiteLLM_AutoRouterUserSession" ( + "user_id" TEXT NOT NULL, + "api_key" TEXT NOT NULL, + "session_id" TEXT NOT NULL, + "router_name" TEXT NOT NULL, + "router_type" TEXT NOT NULL, + "first_turn_at" TIMESTAMP(3) NOT NULL, + "last_turn_at" TIMESTAMP(3) NOT NULL, + "last_model" TEXT NOT NULL, + "models" JSONB NOT NULL DEFAULT '{}', + "turns" INTEGER NOT NULL DEFAULT 0, + "unordered_turns" INTEGER NOT NULL DEFAULT 0, + "covered_turns" INTEGER NOT NULL DEFAULT 0, + "cache_hits" INTEGER NOT NULL DEFAULT 0, + "same_model_turns" INTEGER NOT NULL DEFAULT 0, + "same_model_hits" INTEGER NOT NULL DEFAULT 0, + "first_visit_turns" INTEGER NOT NULL DEFAULT 0, + "first_visit_hits" INTEGER NOT NULL DEFAULT 0, + "return_turns" INTEGER NOT NULL DEFAULT 0, + "return_hits" INTEGER NOT NULL DEFAULT 0, + "return_expired_misses" INTEGER NOT NULL DEFAULT 0, + "return_within_ttl_misses" INTEGER NOT NULL DEFAULT 0, + "ttl_5m_turns" INTEGER NOT NULL DEFAULT 0, + "ttl_1h_turns" INTEGER NOT NULL DEFAULT 0, + "total_tokens" BIGINT NOT NULL DEFAULT 0, + "spend" DOUBLE PRECISION NOT NULL DEFAULT 0, + "saved_spend" DOUBLE PRECISION NOT NULL DEFAULT 0, + "savings_estimated_turns" INTEGER NOT NULL DEFAULT 0, + "savings_estimated_actual_spend" DOUBLE PRECISION NOT NULL DEFAULT 0, + "savings_estimated_saved_spend" DOUBLE PRECISION NOT NULL DEFAULT 0, + "savings_estimated_baseline_models" JSONB NOT NULL DEFAULT '{}', + "classifier_cost" DOUBLE PRECISION NOT NULL DEFAULT 0, + "classifier_cost_recorded_turns" INTEGER NOT NULL DEFAULT 0, + "tier_turns" JSONB NOT NULL DEFAULT '{}', + "baseline_models" JSONB NOT NULL DEFAULT '{}', + + CONSTRAINT "LiteLLM_AutoRouterUserSession_pkey" PRIMARY KEY ("user_id", "api_key", "session_id", "router_name") +); + +CREATE INDEX IF NOT EXISTS "idx_autorouter_user_session_last_turn" ON "LiteLLM_AutoRouterUserSession"("last_turn_at"); + +CREATE INDEX IF NOT EXISTS "idx_autorouter_user_session_user_last_turn" ON "LiteLLM_AutoRouterUserSession"("user_id", "last_turn_at"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 368864f9cd1..31fc76f7a89 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1623,6 +1623,47 @@ model LiteLLM_AutoRouterSession { @@index([last_turn_at], map: "idx_autorouter_session_last_turn") } +model LiteLLM_AutoRouterUserSession { + user_id String + api_key String + session_id String + router_name String + router_type String + first_turn_at DateTime + last_turn_at DateTime + last_model String + models Json @default("{}") + turns Int @default(0) + unordered_turns Int @default(0) + covered_turns Int @default(0) + cache_hits Int @default(0) + same_model_turns Int @default(0) + same_model_hits Int @default(0) + first_visit_turns Int @default(0) + first_visit_hits Int @default(0) + return_turns Int @default(0) + return_hits Int @default(0) + return_expired_misses Int @default(0) + return_within_ttl_misses Int @default(0) + ttl_5m_turns Int @default(0) + ttl_1h_turns Int @default(0) + total_tokens BigInt @default(0) + spend Float @default(0) + saved_spend Float @default(0) + savings_estimated_turns Int @default(0) + savings_estimated_actual_spend Float @default(0) + savings_estimated_saved_spend Float @default(0) + savings_estimated_baseline_models Json @default("{}") + classifier_cost Float @default(0) + classifier_cost_recorded_turns Int @default(0) + tier_turns Json @default("{}") + baseline_models Json @default("{}") + + @@id([user_id, api_key, session_id, router_name]) + @@index([last_turn_at], map: "idx_autorouter_user_session_last_turn") + @@index([user_id, last_turn_at], map: "idx_autorouter_user_session_user_last_turn") +} + // Shadow eval: evaluation of an auto-router against one or more keys' live traffic, in // either direction. forward duplicates the requests the keys did not route through the // router through it, answering whether they should adopt it; reverse duplicates the diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index dd2566a6f84..464edb3b104 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -40,6 +40,12 @@ dependencies = [ "cc", ] +[[package]] +name = "allocator-api2" +version = "0.2.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "683d7910e743518b0e34f1186f92494becacb047c7b6bf616c96772180fef923" + [[package]] name = "android_system_properties" version = "0.1.6" @@ -230,6 +236,7 @@ dependencies = [ "aws-credential-types", "aws-sigv4", "aws-smithy-async", + "aws-smithy-eventstream", "aws-smithy-http", "aws-smithy-runtime", "aws-smithy-runtime-api", @@ -238,7 +245,9 @@ dependencies = [ "bytes", "bytes-utils", "fastrand", + "http 0.2.12", "http 1.4.2", + "http-body 0.4.6", "http-body 1.1.0", "percent-encoding", "pin-project-lite", @@ -272,6 +281,43 @@ dependencies = [ "tracing", ] +[[package]] +name = "aws-sdk-s3" +version = "1.146.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2cd651b4400d4011b8927b83a9552bf90ff11e6e5da0b9f0a7583247aceec971" +dependencies = [ + "arc-swap", + "aws-credential-types", + "aws-runtime", + "aws-sigv4", + "aws-smithy-async", + "aws-smithy-checksums", + "aws-smithy-eventstream", + "aws-smithy-http", + "aws-smithy-json", + "aws-smithy-observability", + "aws-smithy-runtime", + "aws-smithy-runtime-api", + "aws-smithy-schema", + "aws-smithy-types", + "aws-smithy-xml 0.62.1", + "aws-types", + "bytes", + "fastrand", + "hex", + "hmac", + "http 0.2.12", + "http 1.4.2", + "http-body 1.1.0", + "lru", + "percent-encoding", + "regex-lite", + "sha2 0.11.0", + "tracing", + "url", +] + [[package]] name = "aws-sdk-secretsmanager" version = "1.117.0" @@ -316,7 +362,7 @@ dependencies = [ "aws-smithy-runtime-api", "aws-smithy-schema", "aws-smithy-types", - "aws-smithy-xml", + "aws-smithy-xml 0.61.1", "aws-types", "fastrand", "http 0.2.12", @@ -332,6 +378,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "31d955e76ff96acd555bf06fa0fa6d5bf9335fa84ae7c64481b20ae61d231f70" dependencies = [ "aws-credential-types", + "aws-smithy-eventstream", "aws-smithy-http", "aws-smithy-runtime-api", "aws-smithy-types", @@ -359,10 +406,31 @@ dependencies = [ ] [[package]] -name = "aws-smithy-eventstream" -version = "0.61.1" +name = "aws-smithy-checksums" +version = "0.65.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5a9381123ab62d20c13082b151f30f962a3b112b727345394536dfa39a482944" +checksum = "b67ecd999972b58e67cab052f5129906c08c25883bd0788ceefc55ef97d61307" +dependencies = [ + "aws-smithy-http", + "aws-smithy-types", + "bytes", + "crc-fast", + "hex", + "http 1.4.2", + "http-body 1.1.0", + "http-body-util", + "md-5", + "pin-project-lite", + "sha1 0.11.0", + "sha2 0.11.0", + "tracing", +] + +[[package]] +name = "aws-smithy-eventstream" +version = "0.61.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "80c2051c2f1016fb8e6548dd07b8bc2ac9c3fe583721444b92f515e856d31609" dependencies = [ "aws-smithy-types", "bytes", @@ -375,6 +443,7 @@ version = "0.64.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "37843d9add67c3aff5856f409c6dc315d3cdff60f9c0cb5b670dab1e9920306d" dependencies = [ + "aws-smithy-eventstream", "aws-smithy-runtime-api", "aws-smithy-types", "bytes", @@ -554,6 +623,18 @@ dependencies = [ "xmlparser", ] +[[package]] +name = "aws-smithy-xml" +version = "0.62.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b932c8d6dc127fc980eecd78f8694ae9b9551b69a93a7def2a199c1c0033daf" +dependencies = [ + "aws-smithy-runtime-api", + "aws-smithy-schema", + "aws-smithy-types", + "xmlparser", +] + [[package]] name = "aws-types" version = "1.6.0" @@ -1023,6 +1104,16 @@ dependencies = [ "libc", ] +[[package]] +name = "crc-fast" +version = "1.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e75b2483e97a5a7da73ac68a05b629f9c53cff58d8ed1c77866079e18b00dba5" +dependencies = [ + "digest 0.10.7", + "spin", +] + [[package]] name = "crc16" version = "0.4.0" @@ -1968,6 +2059,8 @@ version = "0.17.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ed5909b6e89a2db4456e54cd5f673791d7eca6732202bbf2a9cc504fe2f9b84a" dependencies = [ + "allocator-api2", + "equivalent", "foldhash", ] @@ -2711,6 +2804,21 @@ dependencies = [ "tokio", ] +[[package]] +name = "litellm-cache-redis-semantic" +version = "0.1.0" +dependencies = [ + "litellm-cache", + "litellm-cache-redis", + "litellm-cache-response", + "r2d2", + "redis", + "redis-test", + "serde_json", + "sha2 0.10.9", + "tokio", +] + [[package]] name = "litellm-cache-response" version = "0.1.0" @@ -2727,6 +2835,37 @@ dependencies = [ "tokio", ] +[[package]] +name = "litellm-cache-s3" +version = "0.1.0" +dependencies = [ + "aws-credential-types", + "aws-sdk-s3", + "aws-smithy-types", + "aws-types", + "litellm-auth-aws", + "litellm-cache", + "serde_json", + "tokio", + "wiremock", +] + +[[package]] +name = "litellm-cache-valkey-semantic" +version = "0.1.0" +dependencies = [ + "litellm-cache", + "litellm-cache-redis", + "litellm-cache-response", + "redis", + "redis-test", + "rstest", + "serde_json", + "sha2 0.10.9", + "tokio", + "uuid", +] + [[package]] name = "litellm-callbacks-legacy-python" version = "0.1.0" @@ -2889,6 +3028,7 @@ dependencies = [ "criterion", "futures-util", "litellm-auth", + "litellm-auth-aws", "litellm-auth-gcp", "litellm-cache", "litellm-cache-azure-blob", @@ -2897,7 +3037,10 @@ dependencies = [ "litellm-cache-memory", "litellm-cache-qdrant-semantic", "litellm-cache-redis", + "litellm-cache-redis-semantic", "litellm-cache-response", + "litellm-cache-s3", + "litellm-cache-valkey-semantic", "litellm-callbacks-legacy-python", "litellm-core", "litellm-core-utils", @@ -2909,11 +3052,13 @@ dependencies = [ "pyo3", "pyo3-async-runtimes", "qdrant-client", + "redis", "reqwest 0.12.28", "rstest", "serde", "serde_json", "serde_with", + "sha2 0.10.9", "tokio", "tokio-tungstenite", "url", @@ -2933,6 +3078,7 @@ dependencies = [ "litellm-secrets-azure", "litellm-secrets-cyberark", "litellm-secrets-google", + "litellm-secrets-hashicorp", "litellm-secrets-types", "moka", "reqwest 0.12.28", @@ -3030,6 +3176,26 @@ dependencies = [ "wiremock", ] +[[package]] +name = "litellm-secrets-hashicorp" +version = "0.1.0" +dependencies = [ + "litellm-core-utils", + "litellm-secrets-types", + "moka", + "rstest", + "rustify", + "rustify_derive", + "serde", + "serde_json", + "tempfile", + "thiserror 2.0.19", + "tokio", + "vaultrs", + "veil", + "wiremock", +] + [[package]] name = "litellm-secrets-types" version = "0.1.0" @@ -3121,6 +3287,15 @@ version = "0.4.33" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0ceec5bc11778974d1bcb055b18002eba7f4b3518b6a0081b3af5f21666da9ad" +[[package]] +name = "lru" +version = "0.18.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff9840bcc50b71349309900da0ce7279aa336ae71d73250b07998932c7d97c25" +dependencies = [ + "hashbrown 0.17.1", +] + [[package]] name = "lru-slab" version = "0.1.2" @@ -3149,6 +3324,16 @@ version = "0.8.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "47e1ffaa40ddd1f3ed91f717a33c8c0ee23fff369e3aa8772b9605cc1d22f4c3" +[[package]] +name = "md-5" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "69b6441f590336821bb897fb28fc622898ccceb1d6cea3fde5ea86b090c4de98" +dependencies = [ + "cfg-if", + "digest 0.11.3", +] + [[package]] name = "memchr" version = "2.8.3" @@ -4295,6 +4480,40 @@ dependencies = [ "semver", ] +[[package]] +name = "rustify" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4800ce4c1cc2fec12c559dae2ddbf0e17fcee7569b796e6d75898efef443368b" +dependencies = [ + "anyhow", + "async-trait", + "bytes", + "http 1.4.2", + "reqwest 0.13.5", + "rustify_derive", + "serde", + "serde_json", + "serde_urlencoded", + "thiserror 1.0.69", + "tracing", + "url", +] + +[[package]] +name = "rustify_derive" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "78ea7fda74240f7410d0198b603a8a2f662acc7d76b6667a49f9b162cd8d9b4f" +dependencies = [ + "proc-macro2", + "quote", + "regex", + "serde_urlencoded", + "syn 1.0.109", + "synstructure 0.12.6", +] + [[package]] name = "rustix" version = "1.1.5" @@ -4649,6 +4868,17 @@ dependencies = [ "digest 0.10.7", ] +[[package]] +name = "sha1" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "aacc4cc499359472b4abe1bf11d0b12e688af9a805fa5e3016f9a386dc2d0214" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", + "digest 0.11.3", +] + [[package]] name = "sha1_smol" version = "1.0.1" @@ -4765,6 +4995,12 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "spin" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "023a211cb3138dbc438680b32560ad89f699977624c9f8dbb95a47d5b4c07dd3" + [[package]] name = "spm_precompiled" version = "0.1.4" @@ -4847,6 +5083,17 @@ version = "2.6.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" +[[package]] +name = "syn" +version = "1.0.109" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72b64191b275b66ffe2469e8af2c1cfe3bafa67b529ead792a6d0160888b4237" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + [[package]] name = "syn" version = "2.0.119" @@ -4878,6 +5125,18 @@ dependencies = [ "futures-core", ] +[[package]] +name = "synstructure" +version = "0.12.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f36bdaa60a83aca3921b5259d5400cbf5e90fc51931376a9bd4a0eb79aa7210f" +dependencies = [ + "proc-macro2", + "quote", + "syn 1.0.109", + "unicode-xid", +] + [[package]] name = "synstructure" version = "0.13.2" @@ -5297,6 +5556,7 @@ version = "0.1.44" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100" dependencies = [ + "log", "pin-project-lite", "tracing-attributes", "tracing-core", @@ -5379,7 +5639,7 @@ dependencies = [ "rand 0.8.7", "rustls 0.23.42", "rustls-pki-types", - "sha1", + "sha1 0.10.7", "thiserror 1.0.69", "utf-8", ] @@ -5481,6 +5741,12 @@ version = "1.13.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c6f5d3c3b1bf09027a88a6bc961fc00497d651009560b5463668dc81b0fa87a8" +[[package]] +name = "unicode-xid" +version = "0.2.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebc1c04c71510c7f702b52b7c350734c9ff1295c464a03335b00bb84fc54f853" + [[package]] name = "unicode_categories" version = "0.1.1" @@ -5540,6 +5806,25 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ba73ea9cf16a25df0c8caa16c51acb937d5712a8429db78a3ee29d5dcacd3a65" +[[package]] +name = "vaultrs" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "30ffcc0e81025065dda612ec1e26a3d81bb16ef3062354873d17a35965d68522" +dependencies = [ + "async-trait", + "derive_builder", + "http 1.4.2", + "reqwest 0.13.5", + "rustify", + "rustify_derive", + "serde", + "serde_json", + "thiserror 2.0.19", + "tracing", + "url", +] + [[package]] name = "vcpkg" version = "0.2.15" @@ -5989,7 +6274,7 @@ dependencies = [ "proc-macro2", "quote", "syn 2.0.119", - "synstructure", + "synstructure 0.13.2", ] [[package]] @@ -6030,7 +6315,7 @@ dependencies = [ "proc-macro2", "quote", "syn 2.0.119", - "synstructure", + "synstructure 0.13.2", ] [[package]] diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 0d8e77b4f96..0e8941cb8e6 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -22,6 +22,7 @@ litellm-secrets = { path = "crates/secrets" } litellm-secrets-types = { path = "crates/secrets-types" } litellm-secrets-aws = { path = "crates/secrets-aws" } litellm-secrets-google = { path = "crates/secrets-google" } +litellm-secrets-hashicorp = { path = "crates/secrets-hashicorp" } litellm-secrets-azure = { path = "crates/secrets-azure" } litellm-secrets-cyberark = { path = "crates/secrets-cyberark" } litellm-http = { path = "crates/http" } @@ -32,8 +33,10 @@ litellm-cache = { path = "crates/cache" } litellm-cache-azure-blob = { path = "crates/cache-azure-blob" } litellm-cache-memory = { path = "crates/cache-memory" } litellm-cache-redis = { path = "crates/cache-redis" } +litellm-cache-s3 = { path = "crates/cache-s3" } litellm-cache-gcs = { path = "crates/cache-gcs" } litellm-cache-disk = { path = "crates/cache-disk" } +litellm-cache-redis-semantic = { path = "crates/cache-redis-semantic" } litellm-cache-response = { path = "crates/cache-response" } litellm-cache-qdrant-semantic = { path = "crates/cache-qdrant-semantic" } litellm-token-counter = { path = "crates/token-counter" } @@ -58,6 +61,9 @@ uuid = { version = "1", features = ["v4"] } rstest = "0.26.1" rstest_reuse = "0.7.0" rustls = { version = "0.23", default-features = false, features = ["ring", "std", "tls12"] } +rustify = "=0.7.0" +rustify_derive = "=0.5.5" +vaultrs = { version = "=0.8.0", default-features = false, features = ["rustls"] } rustls-native-certs = "0.8" serde = { version = "1.0", features = ["derive"] } serde_json = { version = "1.0", features = ["float_roundtrip"] } diff --git a/litellm-rust/crates/cache-qdrant-semantic/src/semantic.rs b/litellm-rust/crates/cache-qdrant-semantic/src/semantic.rs index cc82118d280..fb165ed5a8e 100644 --- a/litellm-rust/crates/cache-qdrant-semantic/src/semantic.rs +++ b/litellm-rust/crates/cache-qdrant-semantic/src/semantic.rs @@ -101,10 +101,13 @@ impl QdrantSemanticCache { } fn prompt(context: &SemanticCacheContext) -> Result { - if context.messages.is_empty() { + let Some(messages) = context.messages.as_ref().and_then(Value::as_array) else { + return Err(Error::MissingPrompt); + }; + if messages.is_empty() { return Err(Error::MissingPrompt); } - Ok(prompt_from_messages(&context.messages)) + Ok(prompt_from_messages(messages)) } async fn set( diff --git a/litellm-rust/crates/cache-qdrant-semantic/tests/qdrant.rs b/litellm-rust/crates/cache-qdrant-semantic/tests/qdrant.rs index d806a9cb455..c7522c0b313 100644 --- a/litellm-rust/crates/cache-qdrant-semantic/tests/qdrant.rs +++ b/litellm-rust/crates/cache-qdrant-semantic/tests/qdrant.rs @@ -3,9 +3,7 @@ mod support; use std::{collections::HashMap, sync::Arc, time::Duration}; -use litellm_cache::{ - BaseCache, CacheCodec, CacheContext, Error, SemanticCacheContext, SemanticCacheScope, -}; +use litellm_cache::{BaseCache, CacheCodec, CacheContext, Error, SemanticCacheContext}; use litellm_cache_qdrant_semantic::{ Embedder, QdrantSemanticCache, QdrantSemanticConfig, Quantization, }; @@ -60,8 +58,7 @@ fn config(quantization: Quantization) -> QdrantSemanticConfig { fn context(prompt: &str) -> SemanticCacheContext { SemanticCacheContext { - messages: vec![json!({"role": "user", "content": prompt})], - scope: SemanticCacheScope::default(), + messages: Some(json!([{"role": "user", "content": prompt}])), ..Default::default() } } diff --git a/litellm-rust/crates/cache-redis-semantic/Cargo.toml b/litellm-rust/crates/cache-redis-semantic/Cargo.toml new file mode 100644 index 00000000000..9a8755a189e --- /dev/null +++ b/litellm-rust/crates/cache-redis-semantic/Cargo.toml @@ -0,0 +1,21 @@ +[package] +name = "litellm-cache-redis-semantic" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +litellm-cache.workspace = true +litellm-cache-redis.workspace = true +litellm-cache-response.workspace = true +redis = { version = "1.7.0", features = ["tls-rustls"] } +r2d2 = "0.8.10" +serde_json.workspace = true +sha2.workspace = true +tokio.workspace = true + +[dev-dependencies] +redis-test = "1.0.4" +serde_json.workspace = true +tokio.workspace = true diff --git a/litellm-rust/crates/cache-redis-semantic/src/cache.rs b/litellm-rust/crates/cache-redis-semantic/src/cache.rs new file mode 100644 index 00000000000..e0ac31f3630 --- /dev/null +++ b/litellm-rust/crates/cache-redis-semantic/src/cache.rs @@ -0,0 +1,618 @@ +use std::{ + future::Future, + sync::{Arc, OnceLock}, + time::{Duration, SystemTime, UNIX_EPOCH}, +}; + +use litellm_cache::{ + BaseCache, CacheCodec, CacheConnectionResult, CacheConnectionStatus, Error, + SemanticCacheContext, +}; +use litellm_cache_redis::{ + RedisTopology, + connection::{ConnectionRef, Connections}, +}; +use litellm_cache_response::{CacheEntry, ResponseCacheCodec}; +use serde_json::Value; +use sha2::{Digest, Sha256}; + +use crate::prompt::prompt_from_context; + +const CACHE_KEY_FIELD: &str = "litellm_cache_key"; +const VECTOR_FIELD: &str = "prompt_vector"; + +pub trait Embedder: Send + Sync + 'static { + fn embed(&self, prompt: &str, metadata: Option<&Value>) -> Result, Error>; + + fn async_embed( + &self, + prompt: &str, + metadata: Option<&Value>, + ) -> impl Future, Error>> + Send; +} + +#[derive(Clone, Debug)] +pub struct RedisSemanticConfig { + pub index_name: String, + pub similarity_threshold: f32, +} + +struct Inner { + index_name: String, + distance_threshold: f64, + resolved_index: OnceLock, + codec: ResponseCacheCodec, + clock: fn() -> f64, +} + +impl Inner { + fn new(config: RedisSemanticConfig) -> Self { + Self { + index_name: config.index_name, + distance_threshold: 1.0 - f64::from(config.similarity_threshold), + resolved_index: OnceLock::new(), + codec: ResponseCacheCodec, + clock: timestamp, + } + } + + fn ensure_index( + &self, + connection: &mut ConnectionRef<'_>, + dims: usize, + ) -> Result { + if let Some(name) = self.resolved_index.get() { + return Ok(name.clone()); + } + let name = match index_compatible(connection, &self.index_name, dims)? { + Some(true) => self.index_name.clone(), + Some(false) => self.isolated_index(connection, dims)?, + None => match create_index(connection, &self.index_name, dims) { + Ok(()) => self.index_name.clone(), + Err(_) => match index_compatible(connection, &self.index_name, dims)? { + Some(true) => self.index_name.clone(), + Some(false) => self.isolated_index(connection, dims)?, + None => return Err(Error::Unavailable), + }, + }, + }; + let _ = self.resolved_index.set(name.clone()); + Ok(name) + } + + fn isolated_index( + &self, + connection: &mut ConnectionRef<'_>, + dims: usize, + ) -> Result { + let name = format!("{}_isolated", self.index_name); + match index_compatible(connection, &name, dims)? { + Some(true) => Ok(name), + Some(false) => { + redis::cmd("FT.DROPINDEX") + .arg(&name) + .query::<()>(connection) + .map_err(|_| Error::Unavailable)?; + create_index(connection, &name, dims)?; + Ok(name) + } + None => { + create_index(connection, &name, dims)?; + Ok(name) + } + } + } + + fn store( + &self, + connection: &mut ConnectionRef<'_>, + tag: &str, + value: &CacheEntry, + prompt: &str, + vector: &[f32], + ttl: Option, + ) -> Result<(), Error> { + let index = self.ensure_index(connection, vector.len())?; + let entry_id = entry_id(prompt, tag); + let hash_key = format!("{index}:{entry_id}"); + let response = self.codec.encode(value)?; + redis::cmd("HSET") + .arg(&hash_key) + .arg("entry_id") + .arg(&entry_id) + .arg("prompt") + .arg(prompt) + .arg("response") + .arg(response) + .arg(VECTOR_FIELD) + .arg(vector_buffer(vector)) + .arg("inserted_at") + .arg(format!("{}", (self.clock)())) + .arg("updated_at") + .arg(format!("{}", (self.clock)())) + .arg(CACHE_KEY_FIELD) + .arg(tag) + .query::<()>(connection) + .map_err(|_| Error::Unavailable)?; + if let Some(ttl) = ttl { + redis::cmd("EXPIRE") + .arg(&hash_key) + .arg(ttl_seconds(ttl)) + .query::<()>(connection) + .map_err(|_| Error::Unavailable)?; + } + Ok(()) + } + + fn lookup( + &self, + connection: &mut ConnectionRef<'_>, + tag: &str, + vector: &[f32], + ) -> Result, Error> { + let index = self.ensure_index(connection, vector.len())?; + let query = format!( + "(@{CACHE_KEY_FIELD}:{{{}}})=>[KNN 1 @{VECTOR_FIELD} $vector AS vector_distance]", + escape_tag(tag) + ); + let result = redis::cmd("FT.SEARCH") + .arg(&index) + .arg(query) + .arg("RETURN") + .arg(8) + .arg("entry_id") + .arg("prompt") + .arg("response") + .arg("inserted_at") + .arg("updated_at") + .arg("metadata") + .arg(CACHE_KEY_FIELD) + .arg("vector_distance") + .arg("SORTBY") + .arg("vector_distance") + .arg("ASC") + .arg("DIALECT") + .arg(2) + .arg("LIMIT") + .arg(0) + .arg(1) + .arg("PARAMS") + .arg(2) + .arg("vector") + .arg(vector_buffer(vector)) + .query::(connection) + .map_err(|_| Error::Unavailable)?; + let Some(fields) = first_document(&result) else { + return Ok(None); + }; + if string_field(fields, CACHE_KEY_FIELD).as_deref() != Some(tag) { + return Ok(None); + } + if number_field(fields, "vector_distance") + .is_none_or(|distance| distance > self.distance_threshold) + { + return Ok(None); + } + let Some(response) = bytes_field(fields, "response") else { + return Ok(None); + }; + self.codec.decode(&response).map(Some) + } +} + +pub struct RedisSemanticCache { + connections: Arc>, + embedder: E, + inner: Arc, +} + +impl RedisSemanticCache { + pub fn new(url: &str, embedder: E, config: RedisSemanticConfig) -> Result { + Ok(Self { + connections: Arc::new(Connections::open(url, &RedisTopology::Standalone)?), + embedder, + inner: Arc::new(Inner::new(config)), + }) + } +} + +impl RedisSemanticCache { + pub fn with_connection(connection: C, embedder: E, config: RedisSemanticConfig) -> Self { + Self { + connections: Arc::new(Connections::fixed(connection)), + embedder, + inner: Arc::new(Inner::new(config)), + } + } + + pub fn with_clock(self, clock: fn() -> f64) -> Self { + Self { + inner: Arc::new(Inner { + index_name: self.inner.index_name.clone(), + distance_threshold: self.inner.distance_threshold, + resolved_index: OnceLock::new(), + codec: self.inner.codec, + clock, + }), + ..self + } + } + + pub fn embedder(&self) -> &E { + &self.embedder + } + + pub fn index_name(&self) -> &str { + &self.inner.index_name + } + + pub fn similarity_threshold(&self) -> f32 { + (1.0 - self.inner.distance_threshold) as f32 + } + + fn tag<'a>(key: &'a str, context: &'a SemanticCacheContext) -> &'a str { + context.scope.as_deref().unwrap_or(key) + } +} + +impl BaseCache + for RedisSemanticCache +{ + type Value = CacheEntry; + type Context = SemanticCacheContext; + + fn get_ttl(&self, context: &Self::Context) -> Option { + context.ttl + } + + fn set_cache( + &self, + key: &str, + value: Self::Value, + context: &Self::Context, + ) -> Result<(), Error> { + let Some(prompt) = prompt_from_context(context) else { + return Ok(()); + }; + let vector = self.embedder.embed(&prompt, context.metadata.as_ref())?; + let tag = Self::tag(key, context).to_string(); + self.connections.execute(|connection| { + self.inner + .store(connection, &tag, &value, &prompt, &vector, context.ttl) + }) + } + + fn get_cache(&self, key: &str, context: &Self::Context) -> Result, Error> { + let Some(prompt) = prompt_from_context(context) else { + return Ok(None); + }; + let vector = self.embedder.embed(&prompt, context.metadata.as_ref())?; + let tag = Self::tag(key, context).to_string(); + self.connections + .execute(|connection| self.inner.lookup(connection, &tag, &vector)) + } + + async fn async_set_cache( + &self, + key: &str, + value: Self::Value, + context: Self::Context, + ) -> Result<(), Error> { + let Some(prompt) = prompt_from_context(&context) else { + return Ok(()); + }; + let vector = self + .embedder + .async_embed(&prompt, context.metadata.as_ref()) + .await?; + let tag = Self::tag(key, &context).to_string(); + let inner = Arc::clone(&self.inner); + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { + inner.store(connection, &tag, &value, &prompt, &vector, context.ttl) + }) + .await + } + + async fn async_get_cache( + &self, + key: &str, + context: &Self::Context, + ) -> Result, Error> { + let Some(prompt) = prompt_from_context(context) else { + return Ok(None); + }; + let vector = self + .embedder + .async_embed(&prompt, context.metadata.as_ref()) + .await?; + let tag = Self::tag(key, context).to_string(); + let inner = Arc::clone(&self.inner); + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { + inner.lookup(connection, &tag, &vector) + }) + .await + } + + async fn disconnect(&self) -> Result<(), Error> { + Ok(()) + } + + async fn test_connection(&self) -> Result { + match Connections::run_blocking(Arc::clone(&self.connections), |connection| { + Ok(match redis::cmd("PING").query::(connection) { + Ok(_) => CacheConnectionResult { + status: CacheConnectionStatus::Success, + message: "Redis cache connection test successful".into(), + error: None, + }, + Err(error) => CacheConnectionResult { + status: CacheConnectionStatus::Failed, + message: format!("Redis connection failed: {error}"), + error: Some(error.to_string()), + }, + }) + }) + .await + { + Ok(result) => Ok(result), + Err(error) => Ok(CacheConnectionResult { + status: CacheConnectionStatus::Failed, + message: format!("Redis connection failed: {error}"), + error: Some(error.to_string()), + }), + } + } +} + +fn timestamp() -> f64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|duration| duration.as_secs_f64()) + .unwrap_or_default() +} + +fn entry_id(prompt: &str, tag: &str) -> String { + let mut digest = Sha256::new(); + digest.update(prompt.as_bytes()); + digest.update(CACHE_KEY_FIELD.as_bytes()); + digest.update(tag.as_bytes()); + format!("{:x}", digest.finalize()) +} + +fn vector_buffer(vector: &[f32]) -> Vec { + vector + .iter() + .flat_map(|component| component.to_le_bytes()) + .collect() +} + +fn escape_tag(value: &str) -> String { + value + .chars() + .flat_map(|ch| { + if matches!( + ch, + ',' | '.' + | '<' + | '>' + | '{' + | '}' + | '[' + | ']' + | '\\' + | '"' + | '\'' + | ':' + | ';' + | '!' + | '@' + | '#' + | '$' + | '%' + | '^' + | '&' + | '*' + | '(' + | ')' + | '-' + | '+' + | '=' + | '~' + | '|' + | '/' + | ' ' + | '?' + ) { + vec!['\\', ch] + } else { + vec![ch] + } + }) + .collect() +} + +fn create_index(connection: &mut ConnectionRef<'_>, name: &str, dims: usize) -> Result<(), Error> { + redis::cmd("FT.CREATE") + .arg(name) + .arg("ON") + .arg("HASH") + .arg("PREFIX") + .arg(1) + .arg(name) + .arg("SCORE") + .arg(1.0) + .arg("SCHEMA") + .arg("prompt") + .arg("TEXT") + .arg("WEIGHT") + .arg(1) + .arg("response") + .arg("TEXT") + .arg("WEIGHT") + .arg(1) + .arg("inserted_at") + .arg("NUMERIC") + .arg("updated_at") + .arg("NUMERIC") + .arg(VECTOR_FIELD) + .arg("VECTOR") + .arg("FLAT") + .arg(6) + .arg("TYPE") + .arg("FLOAT32") + .arg("DIM") + .arg(dims) + .arg("DISTANCE_METRIC") + .arg("COSINE") + .arg(CACHE_KEY_FIELD) + .arg("TAG") + .arg("SEPARATOR") + .arg(",") + .query::<()>(connection) + .map_err(|_| Error::Unavailable) +} + +fn index_compatible( + connection: &mut ConnectionRef<'_>, + name: &str, + dims: usize, +) -> Result, Error> { + let info = match redis::cmd("FT.INFO") + .arg(name) + .query::(connection) + { + Ok(info) => info, + Err(error) if unknown_index(&error) => return Ok(None), + Err(_) => return Err(Error::Unavailable), + }; + Ok(Some(schema_compatible(&info, dims))) +} + +fn unknown_index(error: &redis::RedisError) -> bool { + let message = error.to_string().to_lowercase(); + message.contains("unknown") && message.contains("index") +} + +fn schema_compatible(info: &redis::Value, dims: usize) -> bool { + let redis::Value::Array(entries) = info else { + return false; + }; + let attributes = entries + .as_chunks::<2>() + .0 + .iter() + .find(|pair| string_value(&pair[0]).as_deref() == Some("attributes")) + .map(|pair| &pair[1]); + let Some(redis::Value::Array(attributes)) = attributes else { + return false; + }; + let fields = attributes + .iter() + .map(|attribute| { + let redis::Value::Array(attribute) = attribute else { + return (None, None, None, None, None); + }; + let mut name = None; + let mut field_type = None; + let mut dim = None; + let mut data_type = None; + let mut distance_metric = None; + for pair in attribute.as_chunks::<2>().0 { + match string_value(&pair[0]).as_deref() { + Some("identifier") => name = string_value(&pair[1]), + Some("type") => field_type = string_value(&pair[1]), + Some("dim") => dim = number_value(&pair[1]), + Some("data_type") => data_type = string_value(&pair[1]), + Some("distance_metric") => distance_metric = string_value(&pair[1]), + _ => {} + } + } + (name, field_type, dim, data_type, distance_metric) + }) + .collect::>(); + let has_field = |name: &str, field_type: &str| { + fields + .iter() + .any(|(n, t, ..)| n.as_deref() == Some(name) && t.as_deref() == Some(field_type)) + }; + has_field("prompt", "TEXT") + && has_field("response", "TEXT") + && has_field("inserted_at", "NUMERIC") + && has_field("updated_at", "NUMERIC") + && has_field(CACHE_KEY_FIELD, "TAG") + && fields.iter().any(|(n, t, d, data, metric)| { + n.as_deref() == Some(VECTOR_FIELD) + && t.as_deref() == Some("VECTOR") + && *d == Some(dims as f64) + && data + .as_deref() + .is_some_and(|data| data.eq_ignore_ascii_case("float32")) + && metric + .as_deref() + .is_some_and(|metric| metric.eq_ignore_ascii_case("cosine")) + }) +} + +fn string_value(value: &redis::Value) -> Option { + match value { + redis::Value::BulkString(bytes) => String::from_utf8(bytes.clone()).ok(), + redis::Value::SimpleString(text) => Some(text.clone()), + redis::Value::VerbatimString { text, .. } => Some(text.clone()), + _ => None, + } +} + +fn number_value(value: &redis::Value) -> Option { + match value { + redis::Value::Int(number) => Some(*number as f64), + redis::Value::Double(number) => Some(*number), + _ => string_value(value).and_then(|text| text.parse().ok()), + } +} + +fn first_document(result: &redis::Value) -> Option<&[redis::Value]> { + let redis::Value::Array(items) = result else { + return None; + }; + let [count, _document_id, fields, ..] = items.as_slice() else { + return None; + }; + if !matches!(count, redis::Value::Int(count) if *count > 0) { + return None; + } + match fields { + redis::Value::Array(fields) => Some(fields.as_slice()), + _ => None, + } +} + +fn field_value<'a>(fields: &'a [redis::Value], name: &str) -> Option<&'a redis::Value> { + fields + .as_chunks::<2>() + .0 + .iter() + .find(|pair| string_value(&pair[0]).as_deref() == Some(name)) + .map(|pair| &pair[1]) +} + +fn string_field(fields: &[redis::Value], name: &str) -> Option { + field_value(fields, name).and_then(string_value) +} + +fn number_field(fields: &[redis::Value], name: &str) -> Option { + field_value(fields, name).and_then(number_value) +} + +fn bytes_field(fields: &[redis::Value], name: &str) -> Option> { + match field_value(fields, name)? { + redis::Value::BulkString(bytes) => Some(bytes.clone()), + redis::Value::SimpleString(text) => Some(text.clone().into_bytes()), + _ => None, + } +} + +fn ttl_seconds(ttl: Duration) -> u64 { + ttl.as_secs() + .saturating_add(u64::from(ttl.subsec_nanos() > 0)) + .max(1) +} diff --git a/litellm-rust/crates/cache-redis-semantic/src/lib.rs b/litellm-rust/crates/cache-redis-semantic/src/lib.rs new file mode 100644 index 00000000000..51d0b4ba5f3 --- /dev/null +++ b/litellm-rust/crates/cache-redis-semantic/src/lib.rs @@ -0,0 +1,5 @@ +mod cache; +mod prompt; + +pub use cache::{Embedder, RedisSemanticCache, RedisSemanticConfig}; +pub use prompt::prompt_from_context; diff --git a/litellm-rust/crates/cache-redis-semantic/src/prompt.rs b/litellm-rust/crates/cache-redis-semantic/src/prompt.rs new file mode 100644 index 00000000000..b9c38e98d77 --- /dev/null +++ b/litellm-rust/crates/cache-redis-semantic/src/prompt.rs @@ -0,0 +1,97 @@ +use litellm_cache::SemanticCacheContext; +use serde_json::Value; + +pub fn prompt_from_context(context: &SemanticCacheContext) -> Option { + if let Some(messages) = context.messages.as_ref().and_then(Value::as_array) + && !messages.is_empty() + { + return Some(messages_text(messages)); + } + let input = context.input.as_ref()?; + let mut parts = Vec::new(); + collect_input_text(input, &mut parts); + let prompt = parts.join("\n").trim().to_string(); + (!prompt.is_empty()).then_some(prompt) +} + +fn messages_text(messages: &[Value]) -> String { + let mut text = String::new(); + for message in messages { + let Some(message) = message.as_object() else { + continue; + }; + match message.get("content") { + Some(Value::String(content)) => text.push_str(content), + Some(Value::Array(parts)) => { + for part in parts { + if let Some(text_content) = part.get("text").and_then(Value::as_str) { + text.push_str(text_content); + } + } + } + _ => {} + } + text.push_str(&search_results_text(message.get("search_results"))); + } + text +} + +fn search_results_text(search_results: Option<&Value>) -> String { + let Some(Value::Array(results)) = search_results else { + return String::new(); + }; + let mut text = String::new(); + for result in results { + let Some(result) = result.as_object() else { + continue; + }; + for key in ["source", "title"] { + if let Some(value) = result.get(key).and_then(Value::as_str) { + text.push_str(value); + } + } + if let Some(Value::Array(content)) = result.get("content") { + for block in content { + if let Some(value) = block.get("text").and_then(Value::as_str) { + text.push_str(value); + } + } + } + if let Some(citations) = result.get("citations") { + text.push_str(&citations.to_string()); + } + } + text +} + +fn collect_input_text(value: &Value, parts: &mut Vec) { + match value { + Value::String(text) => { + let trimmed = text.trim(); + if !trimmed.is_empty() { + parts.push(trimmed.to_string()); + } + } + Value::Array(items) => { + for item in items { + collect_input_text(item, parts); + } + } + Value::Object(map) => { + if let Some(content) = map.get("content").filter(|content| !content.is_null()) { + collect_input_text(content, parts); + return; + } + for key in ["text", "output", "input_text", "output_text"] { + if let Some(Value::String(text)) = map.get(key) { + let trimmed = text.trim(); + if !trimmed.is_empty() { + parts.push(trimmed.to_string()); + return; + } + } + } + } + _ => {} + } +} diff --git a/litellm-rust/crates/cache-redis-semantic/tests/cache.rs b/litellm-rust/crates/cache-redis-semantic/tests/cache.rs new file mode 100644 index 00000000000..233b87ec52f --- /dev/null +++ b/litellm-rust/crates/cache-redis-semantic/tests/cache.rs @@ -0,0 +1,1003 @@ +use std::{ + collections::HashMap, + sync::{Arc, Mutex}, + time::Duration, +}; + +use litellm_cache::{BaseCache, CacheCodec, Error, SemanticCacheContext}; +use litellm_cache_redis_semantic::{Embedder, RedisSemanticCache, RedisSemanticConfig}; +use litellm_cache_response::{CacheEntry, ResponseCacheCodec}; +use redis_test::{MockCmd, MockRedisConnection}; +use serde_json::{Value, json}; +use sha2::{Digest, Sha256}; + +const INDEX: &str = "litellm_semantic_cache_index"; + +struct FakeEmbedder { + vectors: HashMap>, + calls: Arc>>, +} + +impl FakeEmbedder { + fn new(vectors: &[(&str, &[f32])]) -> (Self, Arc>>) { + let calls = Arc::new(Mutex::new(Vec::new())); + ( + Self { + vectors: vectors + .iter() + .map(|(prompt, vector)| (prompt.to_string(), vector.to_vec())) + .collect(), + calls: Arc::clone(&calls), + }, + calls, + ) + } +} + +impl Embedder for FakeEmbedder { + fn embed(&self, prompt: &str, _: Option<&Value>) -> Result, Error> { + self.calls.lock().unwrap().push(prompt.to_string()); + + Ok(self + .vectors + .get(prompt) + .cloned() + .unwrap_or_else(|| vec![0.1, 0.2, 0.3])) + } + + async fn async_embed(&self, prompt: &str, metadata: Option<&Value>) -> Result, Error> { + self.embed(prompt, metadata) + } +} + +fn config() -> RedisSemanticConfig { + RedisSemanticConfig { + index_name: INDEX.into(), + similarity_threshold: 0.9, + } +} + +fn messages_context(messages: Vec) -> SemanticCacheContext { + SemanticCacheContext { + messages: Some(Value::Array(messages)), + ..Default::default() + } +} + +fn entry() -> CacheEntry { + CacheEntry { + timestamp: Some(1.0), + response: json!({"answer": "yes"}), + } +} + +fn encoded(entry: &CacheEntry) -> Vec { + ResponseCacheCodec.encode(entry).unwrap() +} + +fn vector_bytes(vector: &[f32]) -> Vec { + vector + .iter() + .flat_map(|component| component.to_le_bytes()) + .collect() +} + +fn entry_id(prompt: &str, tag: &str) -> String { + let mut digest = Sha256::new(); + digest.update(prompt.as_bytes()); + digest.update(b"litellm_cache_key"); + digest.update(tag.as_bytes()); + format!("{:x}", digest.finalize()) +} + +fn s(value: &str) -> redis::Value { + redis::Value::BulkString(value.as_bytes().to_vec()) +} + +fn unknown_index_error() -> redis::RedisError { + redis::RedisError::from((redis::ErrorKind::Extension, "Unknown index name")) +} + +fn attribute(name: &str, field_type: &str, extra: Vec) -> redis::Value { + let mut parts = vec![ + s("identifier"), + s(name), + s("attribute"), + s(name), + s("type"), + s(field_type), + ]; + parts.extend(extra); + redis::Value::Array(parts) +} + +fn index_info(attributes: Vec) -> redis::Value { + redis::Value::Array(vec![ + s("index_name"), + s(INDEX), + s("attributes"), + redis::Value::Array(attributes), + ]) +} + +fn vector_attribute_with(dims: i64, data_type: &str, distance_metric: &str) -> redis::Value { + attribute( + "prompt_vector", + "VECTOR", + vec![ + s("algorithm"), + s("FLAT"), + s("data_type"), + s(data_type), + s("dim"), + redis::Value::Int(dims), + s("distance_metric"), + s(distance_metric), + ], + ) +} + +fn vector_attribute(dims: i64) -> redis::Value { + vector_attribute_with(dims, "FLOAT32", "COSINE") +} + +fn info_with_vector(vector: redis::Value) -> redis::Value { + index_info(vec![ + attribute("prompt", "TEXT", vec![]), + attribute("response", "TEXT", vec![]), + attribute("inserted_at", "NUMERIC", vec![]), + attribute("updated_at", "NUMERIC", vec![]), + vector, + attribute("litellm_cache_key", "TAG", vec![]), + ]) +} + +fn compatible_info(dims: i64) -> redis::Value { + info_with_vector(vector_attribute(dims)) +} + +fn unscoped_info(dims: i64) -> redis::Value { + index_info(vec![ + attribute("prompt", "TEXT", vec![]), + attribute("response", "TEXT", vec![]), + attribute("inserted_at", "NUMERIC", vec![]), + attribute("updated_at", "NUMERIC", vec![]), + vector_attribute(dims), + ]) +} + +fn create_index_command(name: &str, dims: usize) -> redis::Cmd { + let mut command = redis::cmd("FT.CREATE"); + command + .arg(name) + .arg("ON") + .arg("HASH") + .arg("PREFIX") + .arg(1) + .arg(name) + .arg("SCORE") + .arg(1.0) + .arg("SCHEMA") + .arg("prompt") + .arg("TEXT") + .arg("WEIGHT") + .arg(1) + .arg("response") + .arg("TEXT") + .arg("WEIGHT") + .arg(1) + .arg("inserted_at") + .arg("NUMERIC") + .arg("updated_at") + .arg("NUMERIC") + .arg("prompt_vector") + .arg("VECTOR") + .arg("FLAT") + .arg(6) + .arg("TYPE") + .arg("FLOAT32") + .arg("DIM") + .arg(dims) + .arg("DISTANCE_METRIC") + .arg("COSINE") + .arg("litellm_cache_key") + .arg("TAG") + .arg("SEPARATOR") + .arg(","); + command +} + +fn search_command(index: &str, tag: &str, vector: &[f32]) -> redis::Cmd { + let mut command = redis::cmd("FT.SEARCH"); + command + .arg(index) + .arg(format!( + "(@litellm_cache_key:{{{tag}}})=>[KNN 1 @prompt_vector $vector AS vector_distance]" + )) + .arg("RETURN") + .arg(8) + .arg("entry_id") + .arg("prompt") + .arg("response") + .arg("inserted_at") + .arg("updated_at") + .arg("metadata") + .arg("litellm_cache_key") + .arg("vector_distance") + .arg("SORTBY") + .arg("vector_distance") + .arg("ASC") + .arg("DIALECT") + .arg(2) + .arg("LIMIT") + .arg(0) + .arg(1) + .arg("PARAMS") + .arg(2) + .arg("vector") + .arg(vector_bytes(vector)); + command +} + +fn hit_fields(tag: &str, distance: &str, response: Vec) -> redis::Value { + redis::Value::Array(vec![ + s("entry_id"), + s("stored-id"), + s("prompt"), + s("hello prompt"), + s("response"), + redis::Value::BulkString(response), + s("inserted_at"), + s("1700000000.5"), + s("updated_at"), + s("1700000000.5"), + s("litellm_cache_key"), + s(tag), + s("vector_distance"), + s(distance), + ]) +} + +fn search_result(fields: redis::Value) -> redis::Value { + redis::Value::Array(vec![ + redis::Value::Int(1), + s("litellm_semantic_cache_index:stored-id"), + fields, + ]) +} + +fn empty_result() -> redis::Value { + redis::Value::Array(vec![redis::Value::Int(0)]) +} + +#[test] +fn store_creates_index_and_writes_hash_with_expire() { + let vector = vec![0.1f32, 0.2, 0.3]; + let prompt = "hello prompt"; + let tag = "key1"; + let hash_key = format!("{INDEX}:{}", entry_id(prompt, tag)); + let value = entry(); + let connection = MockRedisConnection::new([ + MockCmd::new( + redis::cmd("FT.INFO").arg(INDEX), + Err::(unknown_index_error()), + ), + MockCmd::new(create_index_command(INDEX, 3), Ok("OK")), + MockCmd::new( + redis::cmd("HSET") + .arg(&hash_key) + .arg("entry_id") + .arg(entry_id(prompt, tag)) + .arg("prompt") + .arg(prompt) + .arg("response") + .arg(encoded(&value)) + .arg("prompt_vector") + .arg(vector_bytes(&vector)) + .arg("inserted_at") + .arg("1700000000.5") + .arg("updated_at") + .arg("1700000000.5") + .arg("litellm_cache_key") + .arg(tag), + Ok(7), + ), + MockCmd::new(redis::cmd("EXPIRE").arg(&hash_key).arg(5), Ok(1)), + ]) + .assert_all_commands_consumed(); + let (embedder, _) = FakeEmbedder::new(&[(prompt, &vector)]); + let cache = RedisSemanticCache::with_connection(connection, embedder, config()) + .with_clock(|| 1700000000.5); + + let context = SemanticCacheContext { + ttl: Some(Duration::from_secs(5)), + ..messages_context(vec![json!({"role": "user", "content": prompt})]) + }; + cache.set_cache(tag, value, &context).unwrap(); +} + +#[test] +fn store_without_ttl_skips_expire() { + let prompt = "hello prompt"; + let connection = MockRedisConnection::new([ + MockCmd::new(redis::cmd("FT.INFO").arg(INDEX), Ok(compatible_info(3))), + MockCmd::new( + redis::cmd("HSET") + .arg(format!("{INDEX}:{}", entry_id(prompt, "key1"))) + .arg("entry_id") + .arg(entry_id(prompt, "key1")) + .arg("prompt") + .arg(prompt) + .arg("response") + .arg(encoded(&entry())) + .arg("prompt_vector") + .arg(vector_bytes(&[0.1f32, 0.2, 0.3])) + .arg("inserted_at") + .arg("1700000000.5") + .arg("updated_at") + .arg("1700000000.5") + .arg("litellm_cache_key") + .arg("key1"), + Ok(7), + ), + ]) + .assert_all_commands_consumed(); + let (embedder, _) = FakeEmbedder::new(&[]); + let cache = RedisSemanticCache::with_connection(connection, embedder, config()) + .with_clock(|| 1700000000.5); + + cache + .set_cache( + "key1", + entry(), + &messages_context(vec![json!({"role": "user", "content": prompt})]), + ) + .unwrap(); +} + +#[test] +fn lookup_returns_hit_below_distance_threshold() { + let vector = vec![0.1f32, 0.2, 0.3]; + let value = entry(); + let connection = MockRedisConnection::new([ + MockCmd::new(redis::cmd("FT.INFO").arg(INDEX), Ok(compatible_info(3))), + MockCmd::new( + search_command(INDEX, "key1", &vector), + Ok(search_result(hit_fields("key1", "0.05", encoded(&value)))), + ), + ]) + .assert_all_commands_consumed(); + let (embedder, _) = FakeEmbedder::new(&[]); + let cache = RedisSemanticCache::with_connection(connection, embedder, config()); + + let hit = cache + .get_cache( + "key1", + &messages_context(vec![json!({"role": "user", "content": "hello prompt"})]), + ) + .unwrap(); + assert_eq!(hit, Some(value)); +} + +#[test] +fn lookup_misses_above_distance_threshold_and_on_tag_mismatch() { + let vector = vec![0.1f32, 0.2, 0.3]; + let connection = MockRedisConnection::new([ + MockCmd::new(redis::cmd("FT.INFO").arg(INDEX), Ok(compatible_info(3))), + MockCmd::new( + search_command(INDEX, "key1", &vector), + Ok(search_result(hit_fields("key1", "0.5", encoded(&entry())))), + ), + MockCmd::new( + search_command(INDEX, "key1", &vector), + Ok(search_result(hit_fields( + "other", + "0.05", + encoded(&entry()), + ))), + ), + ]) + .assert_all_commands_consumed(); + let (embedder, _) = FakeEmbedder::new(&[]); + let cache = RedisSemanticCache::with_connection(connection, embedder, config()); + let context = messages_context(vec![json!({"role": "user", "content": "hello prompt"})]); + + assert_eq!(cache.get_cache("key1", &context).unwrap(), None); + assert_eq!(cache.get_cache("key1", &context).unwrap(), None); +} + +#[test] +fn lookup_returns_invalid_entry_on_malformed_response() { + let vector = vec![0.1f32, 0.2, 0.3]; + let connection = MockRedisConnection::new([ + MockCmd::new(redis::cmd("FT.INFO").arg(INDEX), Ok(compatible_info(3))), + MockCmd::new( + search_command(INDEX, "key1", &vector), + Ok(search_result(hit_fields( + "key1", + "0.05", + b"not json!".to_vec(), + ))), + ), + ]) + .assert_all_commands_consumed(); + let (embedder, _) = FakeEmbedder::new(&[]); + let cache = RedisSemanticCache::with_connection(connection, embedder, config()); + + assert_eq!( + cache + .get_cache( + "key1", + &messages_context(vec![json!({"role": "user", "content": "hello prompt"})]) + ) + .unwrap_err(), + Error::InvalidEntry + ); +} + +#[test] +fn missing_prompt_is_noop_and_never_embeds() { + let connection = MockRedisConnection::new(Vec::::new()).assert_all_commands_consumed(); + let (embedder, calls) = FakeEmbedder::new(&[]); + let cache = RedisSemanticCache::with_connection(connection, embedder, config()); + + let context = SemanticCacheContext::default(); + cache.set_cache("key1", entry(), &context).unwrap(); + assert_eq!(cache.get_cache("key1", &context).unwrap(), None); + assert!(calls.lock().unwrap().is_empty()); +} + +#[test] +fn scope_overrides_key_as_filter_tag() { + let vector = vec![0.1f32, 0.2, 0.3]; + let prompt = "hello prompt"; + let value = entry(); + let connection = MockRedisConnection::new([ + MockCmd::new(redis::cmd("FT.INFO").arg(INDEX), Ok(compatible_info(3))), + MockCmd::new( + redis::cmd("HSET") + .arg(format!("{INDEX}:{}", entry_id(prompt, "scope-a"))) + .arg("entry_id") + .arg(entry_id(prompt, "scope-a")) + .arg("prompt") + .arg(prompt) + .arg("response") + .arg(encoded(&value)) + .arg("prompt_vector") + .arg(vector_bytes(&vector)) + .arg("inserted_at") + .arg("1700000000.5") + .arg("updated_at") + .arg("1700000000.5") + .arg("litellm_cache_key") + .arg("scope-a"), + Ok(7), + ), + MockCmd::new( + search_command(INDEX, "scope\\-a", &vector), + Ok(search_result(hit_fields( + "scope-a", + "0.05", + encoded(&value), + ))), + ), + ]) + .assert_all_commands_consumed(); + let (embedder, _) = FakeEmbedder::new(&[]); + let cache = RedisSemanticCache::with_connection(connection, embedder, config()) + .with_clock(|| 1700000000.5); + let context = SemanticCacheContext { + scope: Some("scope-a".into()), + ..messages_context(vec![json!({"role": "user", "content": prompt})]) + }; + + cache.set_cache("key1", value.clone(), &context).unwrap(); + assert_eq!(cache.get_cache("key1", &context).unwrap(), Some(value)); +} + +#[test] +fn incompatible_schema_falls_back_to_isolated_index() { + let prompt = "hello prompt"; + let tag = "key1"; + let isolated = format!("{INDEX}_isolated"); + let connection = MockRedisConnection::new([ + MockCmd::new(redis::cmd("FT.INFO").arg(INDEX), Ok(unscoped_info(3))), + MockCmd::new( + redis::cmd("FT.INFO").arg(&isolated), + Err::(unknown_index_error()), + ), + MockCmd::new(create_index_command(&isolated, 3), Ok("OK")), + MockCmd::new( + redis::cmd("HSET") + .arg(format!("{isolated}:{}", entry_id(prompt, tag))) + .arg("entry_id") + .arg(entry_id(prompt, tag)) + .arg("prompt") + .arg(prompt) + .arg("response") + .arg(encoded(&entry())) + .arg("prompt_vector") + .arg(vector_bytes(&[0.1f32, 0.2, 0.3])) + .arg("inserted_at") + .arg("1700000000.5") + .arg("updated_at") + .arg("1700000000.5") + .arg("litellm_cache_key") + .arg(tag), + Ok(7), + ), + ]) + .assert_all_commands_consumed(); + let (embedder, _) = FakeEmbedder::new(&[]); + let cache = RedisSemanticCache::with_connection(connection, embedder, config()) + .with_clock(|| 1700000000.5); + + cache + .set_cache( + tag, + entry(), + &messages_context(vec![json!({"role": "user", "content": prompt})]), + ) + .unwrap(); +} + +#[test] +fn create_index_race_rechecks_schema_and_stores() { + let prompt = "hello prompt"; + let tag = "key1"; + let connection = MockRedisConnection::new([ + MockCmd::new( + redis::cmd("FT.INFO").arg(INDEX), + Err::(unknown_index_error()), + ), + MockCmd::new( + create_index_command(INDEX, 3), + Err::<&str, _>(redis::RedisError::from(( + redis::ErrorKind::Extension, + "Index already exists", + ))), + ), + MockCmd::new(redis::cmd("FT.INFO").arg(INDEX), Ok(compatible_info(3))), + MockCmd::new( + redis::cmd("HSET") + .arg(format!("{INDEX}:{}", entry_id(prompt, tag))) + .arg("entry_id") + .arg(entry_id(prompt, tag)) + .arg("prompt") + .arg(prompt) + .arg("response") + .arg(encoded(&entry())) + .arg("prompt_vector") + .arg(vector_bytes(&[0.1f32, 0.2, 0.3])) + .arg("inserted_at") + .arg("1700000000.5") + .arg("updated_at") + .arg("1700000000.5") + .arg("litellm_cache_key") + .arg(tag), + Ok(7), + ), + ]) + .assert_all_commands_consumed(); + let (embedder, _) = FakeEmbedder::new(&[]); + let cache = RedisSemanticCache::with_connection(connection, embedder, config()) + .with_clock(|| 1700000000.5); + + cache + .set_cache( + tag, + entry(), + &messages_context(vec![json!({"role": "user", "content": prompt})]), + ) + .unwrap(); +} + +#[test] +fn wrong_distance_metric_falls_back_to_isolated_index() { + let prompt = "hello prompt"; + let tag = "key1"; + let isolated = format!("{INDEX}_isolated"); + let connection = MockRedisConnection::new([ + MockCmd::new( + redis::cmd("FT.INFO").arg(INDEX), + Ok(info_with_vector(vector_attribute_with(3, "FLOAT32", "L2"))), + ), + MockCmd::new( + redis::cmd("FT.INFO").arg(&isolated), + Err::(unknown_index_error()), + ), + MockCmd::new(create_index_command(&isolated, 3), Ok("OK")), + MockCmd::new( + redis::cmd("HSET") + .arg(format!("{isolated}:{}", entry_id(prompt, tag))) + .arg("entry_id") + .arg(entry_id(prompt, tag)) + .arg("prompt") + .arg(prompt) + .arg("response") + .arg(encoded(&entry())) + .arg("prompt_vector") + .arg(vector_bytes(&[0.1f32, 0.2, 0.3])) + .arg("inserted_at") + .arg("1700000000.5") + .arg("updated_at") + .arg("1700000000.5") + .arg("litellm_cache_key") + .arg(tag), + Ok(7), + ), + ]) + .assert_all_commands_consumed(); + let (embedder, _) = FakeEmbedder::new(&[]); + let cache = RedisSemanticCache::with_connection(connection, embedder, config()) + .with_clock(|| 1700000000.5); + + cache + .set_cache( + tag, + entry(), + &messages_context(vec![json!({"role": "user", "content": prompt})]), + ) + .unwrap(); +} + +#[test] +fn tag_special_characters_are_escaped_in_search_filter() { + let vector = vec![0.1f32, 0.2, 0.3]; + let tag = "a:b, c|d"; + let connection = MockRedisConnection::new([ + MockCmd::new(redis::cmd("FT.INFO").arg(INDEX), Ok(compatible_info(3))), + MockCmd::new( + search_command(INDEX, "a\\:b\\,\\ c\\|d", &vector), + Ok(empty_result()), + ), + ]) + .assert_all_commands_consumed(); + let (embedder, _) = FakeEmbedder::new(&[]); + let cache = RedisSemanticCache::with_connection(connection, embedder, config()); + + assert_eq!( + cache + .get_cache( + tag, + &messages_context(vec![json!({"role": "user", "content": "hello prompt"})]) + ) + .unwrap(), + None + ); +} + +#[test] +fn prompt_extraction_matches_python_message_and_input_shapes() { + let vector = vec![0.1f32, 0.2, 0.3]; + let lookups = 5; + let mut commands = vec![MockCmd::new( + redis::cmd("FT.INFO").arg(INDEX), + Ok(compatible_info(3)), + )]; + for _ in 0..lookups { + commands.push(MockCmd::new( + search_command(INDEX, "key1", &vector), + Ok(empty_result()), + )); + } + let connection = MockRedisConnection::new(commands).assert_all_commands_consumed(); + let (embedder, calls) = FakeEmbedder::new(&[]); + let cache = RedisSemanticCache::with_connection(connection, embedder, config()); + + cache + .get_cache( + "key1", + &messages_context(vec![ + json!({"role": "user", "content": [{"type": "text", "text": "first"}, {"type": "text", "text": "second"}]}), + json!({"role": "assistant", "content": "reply"}), + ]), + ) + .unwrap(); + cache + .get_cache( + "key1", + &SemanticCacheContext { + input: Some(json!(" plain input ")), + ..Default::default() + }, + ) + .unwrap(); + cache + .get_cache( + "key1", + &SemanticCacheContext { + input: Some( + json!([{"content": [{"type": "input_text", "text": "nested"}]}, "tail"]), + ), + ..Default::default() + }, + ) + .unwrap(); + cache + .get_cache( + "key1", + &SemanticCacheContext { + input: Some(json!({"output_text": " result text "})), + ..Default::default() + }, + ) + .unwrap(); + cache + .get_cache( + "key1", + &messages_context(vec![json!({ + "role": "user", + "content": "question", + "search_results": [{"source": "src", "title": "t", "content": [{"text": "found"}], "citations": {"a": 1}}], + })]), + ) + .unwrap(); + + assert_eq!( + *calls.lock().unwrap(), + vec![ + "firstsecondreply", + "plain input", + "nested\ntail", + "result text", + "questionsrctfound{\"a\":1}", + ] + ); +} + +#[test] +fn ttl_passes_through_context_only() { + let (embedder, _) = FakeEmbedder::new(&[]); + let cache = RedisSemanticCache::with_connection( + MockRedisConnection::new(Vec::::new()), + embedder, + config(), + ); + assert_eq!(cache.get_ttl(&SemanticCacheContext::default()), None); + assert_eq!( + cache.get_ttl(&SemanticCacheContext { + ttl: Some(Duration::from_secs(9)), + ..Default::default() + }), + Some(Duration::from_secs(9)) + ); +} + +#[tokio::test] +async fn async_paths_embed_then_run_blocking_redis_work() { + let vector = vec![0.1f32, 0.2, 0.3]; + let prompt = "hello prompt"; + let tag = "key1"; + let hash_key = format!("{INDEX}:{}", entry_id(prompt, tag)); + let value = entry(); + let connection = MockRedisConnection::new([ + MockCmd::new(redis::cmd("FT.INFO").arg(INDEX), Ok(compatible_info(3))), + MockCmd::new( + redis::cmd("HSET") + .arg(&hash_key) + .arg("entry_id") + .arg(entry_id(prompt, tag)) + .arg("prompt") + .arg(prompt) + .arg("response") + .arg(encoded(&value)) + .arg("prompt_vector") + .arg(vector_bytes(&vector)) + .arg("inserted_at") + .arg("1700000000.5") + .arg("updated_at") + .arg("1700000000.5") + .arg("litellm_cache_key") + .arg(tag), + Ok(7), + ), + MockCmd::new( + search_command(INDEX, tag, &vector), + Ok(search_result(hit_fields(tag, "0.05", encoded(&value)))), + ), + ]) + .assert_all_commands_consumed(); + let (embedder, _) = FakeEmbedder::new(&[]); + let cache = RedisSemanticCache::with_connection(connection, embedder, config()) + .with_clock(|| 1700000000.5); + let context = messages_context(vec![json!({"role": "user", "content": prompt})]); + + cache + .async_set_cache(tag, value.clone(), context.clone()) + .await + .unwrap(); + assert_eq!( + cache.async_get_cache(tag, &context).await.unwrap(), + Some(value) + ); +} + +#[test] +fn shared_base_index_across_dimensions_replaces_the_isolated_index() { + // Pins parity with Python's `_isolated` + overwrite=True flow. + let prompt = "shared prompt"; + let tag = "key1"; + let isolated = format!("{INDEX}_isolated"); + let value = entry(); + let context = || messages_context(vec![json!({"role": "user", "content": prompt})]); + let store_hash = |index: &str, vector: &[f32]| { + MockCmd::new( + redis::cmd("HSET") + .arg(format!("{index}:{}", entry_id(prompt, tag))) + .arg("entry_id") + .arg(entry_id(prompt, tag)) + .arg("prompt") + .arg(prompt) + .arg("response") + .arg(encoded(&value)) + .arg("prompt_vector") + .arg(vector_bytes(vector)) + .arg("inserted_at") + .arg("1700000000.5") + .arg("updated_at") + .arg("1700000000.5") + .arg("litellm_cache_key") + .arg(tag), + Ok(7), + ) + }; + + let vector_a = vec![0.1f32; 8]; + let connection_a = MockRedisConnection::new([ + MockCmd::new( + redis::cmd("FT.INFO").arg(INDEX), + Err::(unknown_index_error()), + ), + MockCmd::new(create_index_command(INDEX, 8), Ok("OK")), + store_hash(INDEX, &vector_a), + ]) + .assert_all_commands_consumed(); + let (embedder_a, _) = FakeEmbedder::new(&[(prompt, &vector_a)]); + let worker_a = RedisSemanticCache::with_connection(connection_a, embedder_a, config()) + .with_clock(|| 1700000000.5); + worker_a.set_cache(tag, value.clone(), &context()).unwrap(); + + let vector_b = vec![0.2f32; 4]; + let connection_b = MockRedisConnection::new([ + MockCmd::new(redis::cmd("FT.INFO").arg(INDEX), Ok(compatible_info(8))), + MockCmd::new( + redis::cmd("FT.INFO").arg(&isolated), + Err::(unknown_index_error()), + ), + MockCmd::new(create_index_command(&isolated, 4), Ok("OK")), + store_hash(&isolated, &vector_b), + MockCmd::new( + search_command(&isolated, tag, &vector_b), + Ok(search_result(hit_fields(tag, "0.0", encoded(&value)))), + ), + MockCmd::new( + search_command(&isolated, tag, &vector_b), + Err::(redis::RedisError::from(( + redis::ErrorKind::Extension, + "Vector dimension mismatch", + ))), + ), + ]) + .assert_all_commands_consumed(); + let (embedder_b, _) = FakeEmbedder::new(&[(prompt, &vector_b)]); + let worker_b = RedisSemanticCache::with_connection(connection_b, embedder_b, config()) + .with_clock(|| 1700000000.5); + worker_b.set_cache(tag, value.clone(), &context()).unwrap(); + assert_eq!( + worker_b.get_cache(tag, &context()).unwrap(), + Some(value.clone()) + ); + + let vector_c = vec![0.3f32; 16]; + let connection_c = MockRedisConnection::new([ + MockCmd::new(redis::cmd("FT.INFO").arg(INDEX), Ok(compatible_info(8))), + MockCmd::new(redis::cmd("FT.INFO").arg(&isolated), Ok(compatible_info(4))), + MockCmd::new(redis::cmd("FT.DROPINDEX").arg(&isolated), Ok("OK")), + MockCmd::new(create_index_command(&isolated, 16), Ok("OK")), + store_hash(&isolated, &vector_c), + ]) + .assert_all_commands_consumed(); + let (embedder_c, _) = FakeEmbedder::new(&[(prompt, &vector_c)]); + let worker_c = RedisSemanticCache::with_connection(connection_c, embedder_c, config()) + .with_clock(|| 1700000000.5); + worker_c.set_cache(tag, value.clone(), &context()).unwrap(); + + assert_eq!( + worker_b.get_cache(tag, &context()).unwrap_err(), + Error::Unavailable + ); +} + +#[test] +fn live_shared_index_is_replaced_across_dimensions() { + let Ok(url) = std::env::var("LITELLM_REDIS_STACK_URL") else { + return; + }; + // Pins parity with Python's `_isolated` + overwrite=True flow. + let base = format!("rust_semantic_shared_{}", std::process::id()); + let isolated = format!("{base}_isolated"); + let prompt = "shared live prompt"; + let tag = "key1"; + let context = || messages_context(vec![json!({"role": "user", "content": prompt})]); + let value = entry(); + let worker = |vector: Vec| { + let (embedder, _) = FakeEmbedder::new(&[(prompt, vector.as_slice())]); + RedisSemanticCache::new( + &url, + embedder, + RedisSemanticConfig { + index_name: base.clone(), + similarity_threshold: 0.9, + }, + ) + .unwrap() + }; + + let worker_a = worker(vec![0.1f32; 8]); + worker_a.set_cache(tag, value.clone(), &context()).unwrap(); + + let worker_b = worker(vec![0.2f32; 4]); + worker_b.set_cache(tag, value.clone(), &context()).unwrap(); + assert_eq!( + worker_b.get_cache(tag, &context()).unwrap(), + Some(value.clone()) + ); + + let worker_c = worker(vec![0.3f32; 16]); + worker_c.set_cache(tag, value.clone(), &context()).unwrap(); + + assert_eq!( + worker_b.get_cache(tag, &context()).unwrap_err(), + Error::Unavailable + ); + + let mut connection = redis::Client::open(url).unwrap().get_connection().unwrap(); + for index in [&base, &isolated] { + let _: Result<(), _> = redis::cmd("FT.DROPINDEX") + .arg(index) + .arg("DD") + .query(&mut connection); + } +} + +#[test] +fn live_store_lookup_and_ttl_against_redis_stack() { + let Ok(url) = std::env::var("LITELLM_REDIS_STACK_URL") else { + return; + }; + let vector = vec![0.1f32, 0.2, 0.3, 0.4]; + let prompt = "rust semantic cache live prompt"; + let tag = "live-key"; + let index_name = format!("rust_semantic_test_{}", std::process::id()); + let (embedder, _) = FakeEmbedder::new(&[(prompt, &vector)]); + let cache = RedisSemanticCache::new( + &url, + embedder, + RedisSemanticConfig { + index_name: index_name.clone(), + similarity_threshold: 0.9, + }, + ) + .unwrap(); + let context = SemanticCacheContext { + ttl: Some(Duration::from_secs(120)), + ..messages_context(vec![json!({"role": "user", "content": prompt})]) + }; + let value = entry(); + + cache.set_cache(tag, value.clone(), &context).unwrap(); + assert_eq!(cache.get_cache(tag, &context).unwrap(), Some(value)); + assert_eq!(cache.get_cache("other-key", &context).unwrap(), None); + + let mut connection = redis::Client::open(url).unwrap().get_connection().unwrap(); + let ttl: i64 = redis::Commands::ttl( + &mut connection, + format!("{index_name}:{}", entry_id(prompt, tag)), + ) + .unwrap(); + assert!( + ttl > 0, + "expected stored hash to carry an expiry, got {ttl}" + ); +} diff --git a/litellm-rust/crates/cache-redis/src/cache.rs b/litellm-rust/crates/cache-redis/src/cache.rs index e2e2656fcbb..24399c9b2f9 100644 --- a/litellm-rust/crates/cache-redis/src/cache.rs +++ b/litellm-rust/crates/cache-redis/src/cache.rs @@ -14,7 +14,7 @@ use crate::topology::RedisTopology; mod connection; mod operations; -pub(crate) use connection::ConnectionRef; +pub use connection::ConnectionRef; use connection::{ClusterConnectionManager, ConnectionManager}; pub use operations::{ @@ -40,7 +40,8 @@ const CLAIM_SCRIPT: &str = concat!( ); const CLAIM_ATTEMPTS: usize = 8; -enum Connections { +#[allow(private_interfaces)] +pub enum Connections { Pool(r2d2::Pool), Cluster(r2d2::Pool), Fixed(Mutex), @@ -50,7 +51,7 @@ impl Connections where C: redis::ConnectionLike + Send + 'static, { - fn execute( + pub fn execute( &self, operation: impl FnOnce(&mut ConnectionRef<'_>) -> Result, ) -> Result { @@ -73,6 +74,29 @@ where } } } + + pub async fn run_blocking(connections: Arc, operation: F) -> Result + where + T: Send + 'static, + F: FnOnce(&mut ConnectionRef<'_>) -> Result + Send + 'static, + { + tokio::task::spawn_blocking(move || connections.execute(operation)) + .await + .map_err(|_| Error::Unavailable)? + } + + pub fn fixed(connection: C) -> Self { + Self::Fixed(Mutex::new(connection)) + } + + pub fn open(url: &str, topology: &RedisTopology) -> Result { + match topology { + RedisTopology::Standalone => Ok(Self::Pool(pool(ConnectionManager::open(url)?)?)), + RedisTopology::Cluster { startup_nodes } => Ok(Self::Cluster(pool( + ClusterConnectionManager::open(url, startup_nodes)?, + )?)), + } + } } pub struct RedisCache { @@ -94,12 +118,7 @@ impl RedisCache { default_ttl: Option, codec: S, ) -> Result { - let connections = match topology { - RedisTopology::Standalone => Connections::Pool(pool(ConnectionManager::open(url)?)?), - RedisTopology::Cluster { startup_nodes } => { - Connections::Cluster(pool(ClusterConnectionManager::open(url, startup_nodes)?)?) - } - }; + let connections = Connections::open(url, topology)?; Ok(Self { connections: Arc::new(connections), default_ttl: default_ttl.unwrap_or(DEFAULT_TTL), @@ -127,7 +146,7 @@ where { pub fn with_connection(connection: C, default_ttl: Option, codec: S) -> Self { Self { - connections: Arc::new(Connections::Fixed(Mutex::new(connection))), + connections: Arc::new(Connections::fixed(connection)), default_ttl: default_ttl.unwrap_or(DEFAULT_TTL), codec, namespace: None, @@ -203,16 +222,6 @@ where .saturating_add(u64::from(ttl.subsec_nanos() > 0)) .max(1) } - - async fn run_blocking(connections: Arc>, operation: F) -> Result - where - T: Send + 'static, - F: FnOnce(&mut ConnectionRef<'_>) -> Result + Send + 'static, - { - tokio::task::spawn_blocking(move || connections.execute(operation)) - .await - .map_err(|_| Error::Unavailable)? - } } fn namespaced_key(namespace: Option<&str>, key: &str) -> String { @@ -271,7 +280,7 @@ where let payload = self.codec.encode(&value)?; let key = self.namespaced_key(key); let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl)); - Self::run_blocking(Arc::clone(&self.connections), move |connection| { + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { connection .set_ex::<_, _, ()>(key, payload, ttl) .map_err(|_| Error::Unavailable) @@ -285,7 +294,7 @@ where _: &ExactCacheContext, ) -> Result, Error> { let key = self.namespaced_key(key); - let value = Self::run_blocking(Arc::clone(&self.connections), move |connection| { + let value = Connections::run_blocking(Arc::clone(&self.connections), move |connection| { connection .get::<_, redis::Value>(key) .map_err(|_| Error::Unavailable) @@ -311,7 +320,7 @@ where if entries.is_empty() { return Ok(()); } - Self::run_blocking(Arc::clone(&self.connections), move |connection| { + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { let commands = entries .into_iter() .map(|(key, payload)| { @@ -330,7 +339,7 @@ where } async fn test_connection(&self) -> Result { - match Self::run_blocking(Arc::clone(&self.connections), |connection| { + match Connections::run_blocking(Arc::clone(&self.connections), |connection| { Ok(match connection.ping() { Ok(_) => CacheConnectionResult { status: CacheConnectionStatus::Success, @@ -391,7 +400,7 @@ where .iter() .map(|key| self.namespaced_key(key)) .collect::>(); - let values = Self::run_blocking(Arc::clone(&self.connections), move |connection| { + let values = Connections::run_blocking(Arc::clone(&self.connections), move |connection| { redis::cmd("MGET") .arg(keys) .query::>(connection) @@ -418,7 +427,7 @@ where async fn async_delete_cache(&self, key: &str) -> Result<(), Error> { let key = self.namespaced_key(key); - Self::run_blocking(Arc::clone(&self.connections), move |connection| { + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { connection.del::<_, ()>(key).map_err(|_| Error::Unavailable) }) .await @@ -438,7 +447,7 @@ where async fn async_flush_cache(&self) -> Result<(), Error> { let pattern = self.namespaced_pattern()?; - Self::run_blocking(Arc::clone(&self.connections), move |connection| { + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { Self::flush_matching(connection, &pattern) }) .await @@ -470,7 +479,7 @@ where ) -> Result { let key = self.namespaced_key(key); let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl)); - Self::run_blocking(Arc::clone(&self.connections), move |connection| { + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { increment(connection, key, amount, ttl) }) .await @@ -581,7 +590,7 @@ where let key = self.namespaced_key(key); let ttl = Self::ttl_seconds(self.get_ttl(&context).unwrap_or(self.default_ttl)); let codec = self.codec.clone(); - Self::run_blocking(Arc::clone(&self.connections), move |connection| { + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { claim(connection, &codec, &key, candidate, &eligible, ttl) }) .await diff --git a/litellm-rust/crates/cache-redis/src/cache/connection.rs b/litellm-rust/crates/cache-redis/src/cache/connection.rs index 1834f1d94e5..013bf055f89 100644 --- a/litellm-rust/crates/cache-redis/src/cache/connection.rs +++ b/litellm-rust/crates/cache-redis/src/cache/connection.rs @@ -12,7 +12,7 @@ use redis::{ use super::REDIS_TIMEOUT; use crate::topology::RedisNode; -pub(super) struct PooledConnection { +pub struct PooledConnection { pub(super) connection: C, pub(super) failed: bool, } @@ -20,7 +20,7 @@ pub(super) struct PooledConnection { /// Pools connections without a checkout PING, which would double every operation's round trips. /// A timed-out command leaves its reply on the socket while redis still reports the connection /// open, so any connection whose operation failed is discarded instead of being reused. -pub(super) struct ConnectionManager(redis::Client); +pub struct ConnectionManager(redis::Client); impl ConnectionManager { pub(super) fn open(url: &str) -> Result { @@ -54,7 +54,7 @@ impl r2d2::ManageConnection for ConnectionManager { } } -pub(super) struct ClusterConnectionManager(ClusterClient); +pub struct ClusterConnectionManager(ClusterClient); impl ClusterConnectionManager { pub(super) fn open(url: &str, startup_nodes: &[RedisNode]) -> Result { @@ -117,7 +117,7 @@ impl r2d2::ManageConnection for ClusterConnectionManager { } } -pub(crate) enum ConnectionRef<'a> { +pub enum ConnectionRef<'a> { Node(&'a mut dyn redis::ConnectionLike), Cluster(&'a mut ClusterConnection), } diff --git a/litellm-rust/crates/cache-redis/src/cache/operations.rs b/litellm-rust/crates/cache-redis/src/cache/operations.rs index 9a7023338bf..4345ee879b3 100644 --- a/litellm-rust/crates/cache-redis/src/cache/operations.rs +++ b/litellm-rust/crates/cache-redis/src/cache/operations.rs @@ -144,7 +144,7 @@ where .into_iter() .map(|key| self.namespaced_key(&key)) .collect::>(); - Self::run_blocking(Arc::clone(&self.connections), move |connection| { + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { connection.del(keys).map_err(|_| Error::Unavailable) }) .await @@ -172,7 +172,7 @@ where .iter() .map(|key| self.namespaced_key(key)) .collect::>(); - let values = Self::run_blocking(Arc::clone(&self.connections), move |connection| { + let values = Connections::run_blocking(Arc::clone(&self.connections), move |connection| { redis::cmd("MGET") .arg(keys) .query::>(connection) @@ -188,7 +188,7 @@ where } pub async fn ping(&self) -> Result { - Self::run_blocking(Arc::clone(&self.connections), |connection| { + Connections::run_blocking(Arc::clone(&self.connections), |connection| { connection.ping().map_err(|_| Error::Unavailable) }) .await @@ -196,7 +196,7 @@ where pub async fn async_get_ttl(&self, key: &str) -> Result, Error> { let key = self.namespaced_key(key); - let ttl = Self::run_blocking(Arc::clone(&self.connections), move |connection| { + let ttl = Connections::run_blocking(Arc::clone(&self.connections), move |connection| { redis::cmd("TTL") .arg(key) .query::(connection) @@ -208,7 +208,7 @@ where pub async fn async_scan_iter(&self, pattern: &str, count: usize) -> Result, Error> { let pattern = format!("{}*", self.namespaced_key(pattern)); - Self::run_blocking(Arc::clone(&self.connections), move |connection| { + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { let mut matches = Vec::new(); connection.scan(&pattern, count, |_, keys| { matches.extend(keys); @@ -231,7 +231,7 @@ where } let key = self.namespaced_key(key); let ttl = Self::ttl_seconds(ttl.unwrap_or(self.default_ttl)); - Self::run_blocking(Arc::clone(&self.connections), move |connection| { + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { let mut sadd = redis::cmd("SADD"); sadd.arg(&key).arg(values); let mut expire = redis::cmd("EXPIRE"); @@ -253,7 +253,7 @@ where return Err(Error::InvalidEntry); } let key = self.namespaced_key(key); - Self::run_blocking(Arc::clone(&self.connections), move |connection| { + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { redis::cmd("RPUSH") .arg(key) .arg(values) @@ -279,7 +279,7 @@ where if operations.is_empty() { return Ok(Vec::new()); } - Self::run_blocking(Arc::clone(&self.connections), move |connection| { + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { let commands = operations .into_iter() .map(|(key, values)| { @@ -304,7 +304,7 @@ where ) -> Result { let key = self.namespaced_key(key); let multiple = count.is_some(); - let value = Self::run_blocking(Arc::clone(&self.connections), move |connection| { + let value = Connections::run_blocking(Arc::clone(&self.connections), move |connection| { let mut command = redis::cmd("LPOP"); command.arg(key); if let Some(count) = count { @@ -333,7 +333,7 @@ where .iter() .map(|(_, count)| count.is_some()) .collect::>(); - let values = Self::run_blocking(Arc::clone(&self.connections), move |connection| { + let values = Connections::run_blocking(Arc::clone(&self.connections), move |connection| { let commands = operations .into_iter() .map(|(key, count)| { @@ -365,7 +365,7 @@ where .into_iter() .map(|key| self.namespaced_key(&key)) .collect::>(); - Self::run_blocking(Arc::clone(&self.connections), move |connection| { + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { redis::cmd("EVAL") .arg(script) .arg(keys.len()) @@ -426,7 +426,7 @@ where if operations.is_empty() { return Ok(Vec::new()); } - Self::run_blocking(Arc::clone(&self.connections), move |connection| { + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { let mut commands = Vec::with_capacity(operations.len() * 2); let mut increments = Vec::with_capacity(operations.len()); for (key, amount, ttl) in operations { @@ -460,7 +460,7 @@ where ) -> Result { let key = self.namespaced_key(key); let ttl = Self::ttl_seconds(ttl); - Self::run_blocking(Arc::clone(&self.connections), move |connection| { + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { increment_with_floor(connection, key, amount, ttl) }) .await @@ -474,7 +474,7 @@ where ) -> Result { let key = self.namespaced_key(key); let ttl = Self::ttl_seconds(ttl.unwrap_or(self.default_ttl)); - Self::run_blocking(Arc::clone(&self.connections), move |connection| { + Connections::run_blocking(Arc::clone(&self.connections), move |connection| { redis::cmd("EVAL") .arg(SET_MAX_SCRIPT) .arg(1) diff --git a/litellm-rust/crates/cache-redis/src/lib.rs b/litellm-rust/crates/cache-redis/src/lib.rs index 98f6bfd8ce5..efb0db931ac 100644 --- a/litellm-rust/crates/cache-redis/src/lib.rs +++ b/litellm-rust/crates/cache-redis/src/lib.rs @@ -1,6 +1,10 @@ mod cache; mod topology; +pub mod connection { + pub use crate::cache::{ConnectionRef, Connections}; +} + pub use cache::{ RedisArg, RedisCache, RedisLpopOperation, RedisLpopResult, RedisRpushOperation, RedisScript, }; diff --git a/litellm-rust/crates/cache-response/README.md b/litellm-rust/crates/cache-response/README.md index 56c1646d343..46e561ddad1 100644 --- a/litellm-rust/crates/cache-response/README.md +++ b/litellm-rust/crates/cache-response/README.md @@ -58,4 +58,4 @@ Verify typed values, TTL precedence, missing entries, serialization failures, na Public SDK, Router, and proxy activation still need constructor parity, stream replay, embedding partial-batch integration, response reconstruction, callback scheduling, and failure-policy integration. This foundation does not switch those request paths -Redis cluster, disk, cloud stores, and semantic caching remain follow-ups. The generic dual cache takes read, write, and remote-failure policies, runs its async operations through the async L2 methods, and provides L2-first counters and atomic affinity claims. Errors propagate by default, and `RemoteFailurePolicy::UseLocal` opts key-value operations and claims into the local tier when L2 is unavailable. Claims compare decoded values, so a pin written by Python still matches. Public Router integration remains follow-up work. Reservations and pubsub still need explicit capabilities owned by their consuming features. Adding a cache backend does not establish those guarantees +Redis cluster, disk, and cloud stores remain follow-ups. Semantic backends plug in through `SemanticCacheContext`, which carries the prompt inputs and metadata alongside the cache TTL. The generic dual cache takes read, write, and remote-failure policies, runs its async operations through the async L2 methods, and provides L2-first counters and atomic affinity claims. Errors propagate by default, and `RemoteFailurePolicy::UseLocal` opts key-value operations and claims into the local tier when L2 is unavailable. Claims compare decoded values, so a pin written by Python still matches. Public Router integration remains follow-up work. Reservations and pubsub still need explicit capabilities owned by their consuming features. Adding a cache backend does not establish those guarantees diff --git a/litellm-rust/crates/cache-response/src/response.rs b/litellm-rust/crates/cache-response/src/response.rs index eedbf2caf1a..5088402f125 100644 --- a/litellm-rust/crates/cache-response/src/response.rs +++ b/litellm-rust/crates/cache-response/src/response.rs @@ -1,15 +1,14 @@ use std::{sync::Arc, time::Duration}; use litellm_cache::{ - BaseCache, BatchCache, BatchEntry, CacheConnectionResult, CacheContext, Error, - ExactCacheContext, FlushCache, + BaseCache, BatchCache, BatchEntry, CacheConnectionResult, CacheContext, Error, FlushCache, }; use serde_json::Value; use crate::{CacheControls, CacheEntry, CacheKeyInput, PartialHits, cache_key}; #[derive(Clone)] -pub struct ResponseCacheRequest { +pub struct ResponseCacheRequest { pub key: CacheKeyInput, pub controls: CacheControls, pub context: C, @@ -64,6 +63,10 @@ where &self.backend } + pub fn backend_arc(&self) -> &Arc { + &self.backend + } + pub fn default_ttl(&self) -> Option { self.backend.get_ttl(&B::Context::default()) } @@ -267,9 +270,9 @@ where Ok(()) } - fn partial_hits( - requests: &[ResponseCacheRequest], - readable: Vec<(usize, &ResponseCacheRequest)>, + fn partial_hits( + requests: &[ResponseCacheRequest], + readable: Vec<(usize, &ResponseCacheRequest)>, entries: Vec>, now: Duration, ) -> Result { diff --git a/litellm-rust/crates/cache-response/tests/response.rs b/litellm-rust/crates/cache-response/tests/response.rs index 56589063291..dcfc0301148 100644 --- a/litellm-rust/crates/cache-response/tests/response.rs +++ b/litellm-rust/crates/cache-response/tests/response.rs @@ -89,7 +89,7 @@ fn semantic_context_reaches_backend_for_store_and_lookup() { }); let cache = ResponseCache::new(backend.clone()); let context = SemanticCacheContext { - messages: vec![json!({"role": "user", "content": "hello"})], + messages: Some(json!([{"role": "user", "content": "hello"}])), ..Default::default() }; let request = request().with_context(context.clone()); diff --git a/litellm-rust/crates/cache-s3/Cargo.toml b/litellm-rust/crates/cache-s3/Cargo.toml new file mode 100644 index 00000000000..cdc17e732cb --- /dev/null +++ b/litellm-rust/crates/cache-s3/Cargo.toml @@ -0,0 +1,20 @@ +[package] +name = "litellm-cache-s3" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +litellm-cache.workspace = true +litellm-auth-aws.workspace = true +aws-sdk-s3 = { version = "1.146.1", default-features = false, features = ["rustls", "rt-tokio"] } +aws-credential-types = "1.3.0" +aws-smithy-types = "1.6.0" +aws-types = "1.6.0" +tokio.workspace = true + +[dev-dependencies] +wiremock = "0.6.5" +serde_json.workspace = true +tokio = { workspace = true, features = ["macros", "rt-multi-thread"] } diff --git a/litellm-rust/crates/cache-s3/src/auth.rs b/litellm-rust/crates/cache-s3/src/auth.rs new file mode 100644 index 00000000000..b7ca722cea3 --- /dev/null +++ b/litellm-rust/crates/cache-s3/src/auth.rs @@ -0,0 +1,101 @@ +use aws_credential_types::{ + Credentials as AwsCredentials, + provider::{ProvideCredentials, error::CredentialsError, future}, +}; +use litellm_auth_aws::{AwsAuthConfig, resolve_credentials}; + +#[derive(Clone)] +pub(crate) struct Credentials { + config: AwsAuthConfig, + env: fn(&str) -> Option, +} + +impl Credentials { + pub(crate) fn new(config: AwsAuthConfig) -> Self { + Self::with_env(config, |name| std::env::var(name).ok()) + } + + pub(crate) fn with_env(config: AwsAuthConfig, env: fn(&str) -> Option) -> Self { + Self { config, env } + } +} + +impl ProvideCredentials for Credentials { + fn provide_credentials<'a>(&'a self) -> future::ProvideCredentials<'a> + where + Self: 'a, + { + future::ProvideCredentials::new(async { + if let (Some(access_key_id), Some(secret_access_key)) = ( + self.config.access_key_id.clone(), + self.config.secret_access_key.clone(), + ) { + return Ok(AwsCredentials::new( + access_key_id, + secret_access_key, + self.config.session_token.clone(), + None, + "litellm-s3-cache", + )); + } + resolve_credentials(self.config.clone(), &self.env) + .await + .map_err(|_| CredentialsError::provider_error("S3 cache authentication failed")) + }) + } +} + +impl std::fmt::Debug for Credentials { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("Credentials").finish_non_exhaustive() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn explicit_keys_ignore_an_ambient_session_token() { + let provider = Credentials::with_env( + AwsAuthConfig { + access_key_id: Some("key".to_string()), + secret_access_key: Some("secret".to_string()), + region_name: Some("us-east-1".to_string()), + ..Default::default() + }, + |name| (name == "AWS_SESSION_TOKEN").then(|| "ambient".to_string()), + ); + let credentials = provider.provide_credentials().await.unwrap(); + assert_eq!(credentials.access_key_id(), "key"); + assert_eq!(credentials.secret_access_key(), "secret"); + assert_eq!(credentials.session_token(), None); + } + + #[tokio::test] + async fn explicit_keys_keep_their_session_token() { + let provider = Credentials::new(AwsAuthConfig { + access_key_id: Some("key".to_string()), + secret_access_key: Some("secret".to_string()), + session_token: Some("t".to_string()), + region_name: Some("us-east-1".to_string()), + ..Default::default() + }); + let credentials = provider.provide_credentials().await.unwrap(); + assert_eq!(credentials.session_token(), Some("t")); + } + + #[tokio::test] + async fn environment_keys_resolve_with_their_session_token() { + let provider = Credentials::with_env(AwsAuthConfig::default(), |name| match name { + "AWS_ACCESS_KEY_ID" => Some("env-key".to_string()), + "AWS_SECRET_ACCESS_KEY" => Some("env-secret".to_string()), + "AWS_SESSION_TOKEN" => Some("env-token".to_string()), + _ => None, + }); + let credentials = provider.provide_credentials().await.unwrap(); + assert_eq!(credentials.access_key_id(), "env-key"); + assert_eq!(credentials.secret_access_key(), "env-secret"); + assert_eq!(credentials.session_token(), Some("env-token")); + } +} diff --git a/litellm-rust/crates/cache-s3/src/cache.rs b/litellm-rust/crates/cache-s3/src/cache.rs new file mode 100644 index 00000000000..9c791f42c3b --- /dev/null +++ b/litellm-rust/crates/cache-s3/src/cache.rs @@ -0,0 +1,220 @@ +use std::{ + future::Future, + sync::Arc, + time::{Duration, SystemTime}, +}; + +use aws_sdk_s3::{ + config::{BehaviorVersion, Region, RequestChecksumCalculation, ResponseChecksumValidation}, + error::SdkError, + primitives::ByteStream, +}; +use aws_smithy_types::{DateTime, date_time::Format}; +use litellm_auth_aws::AwsAuthConfig; +use litellm_cache::{ + BaseCache, BatchCache, CacheCodec, CacheConnectionResult, Error, ExactCacheContext, FlushCache, +}; +use tokio::runtime::Handle; + +use crate::auth::Credentials; + +pub struct S3Endpoint { + pub url: String, +} + +pub struct S3CacheConfig { + pub bucket: String, + pub key_prefix: String, + pub region: String, + pub endpoint: Option, + pub auth: AwsAuthConfig, +} + +pub struct S3Cache { + client: aws_sdk_s3::Client, + codec: C, + runtime: Handle, + bucket: Arc, + key_prefix: Arc, + region: Arc, + endpoint: Option>, +} + +impl S3Cache { + pub fn new(config: S3CacheConfig, codec: C, runtime: Handle) -> Self { + let endpoint_url: Option = config.endpoint.map(|endpoint| endpoint.url); + let base = aws_sdk_s3::Config::builder() + .behavior_version(BehaviorVersion::latest()) + .region(Region::new(config.region.clone())) + .credentials_provider(Credentials::new(config.auth)) + .request_checksum_calculation(RequestChecksumCalculation::WhenRequired) + .response_checksum_validation(ResponseChecksumValidation::WhenRequired); + let builder = match &endpoint_url { + Some(url) => base.endpoint_url(url).force_path_style(true), + None => base, + }; + Self { + client: aws_sdk_s3::Client::from_conf(builder.build()), + codec, + runtime, + bucket: config.bucket.into(), + key_prefix: config.key_prefix.into(), + region: config.region.into(), + endpoint: endpoint_url.map(Into::into), + } + } + + pub fn bucket(&self) -> &str { + &self.bucket + } + + pub fn key_prefix(&self) -> &str { + &self.key_prefix + } + + pub fn region(&self) -> &str { + &self.region + } + + pub fn endpoint(&self) -> Option<&str> { + self.endpoint.as_deref() + } + + pub fn to_s3_key(&self, key: &str) -> String { + format!("{}{}", self.key_prefix, key.replace(':', "/")) + } + + fn block_on(&self, future: F) -> F::Output { + if Handle::try_current().is_ok() { + tokio::task::block_in_place(|| self.runtime.block_on(future)) + } else { + self.runtime.block_on(future) + } + } + + async fn put( + &self, + key: &str, + value: C::Value, + context: &ExactCacheContext, + ) -> Result<(), Error> { + let s3_key = self.to_s3_key(key); + let body = self.codec.encode(&value)?; + let request = self + .client + .put_object() + .bucket(self.bucket.as_ref()) + .key(&s3_key) + .body(ByteStream::from(body)) + .content_type("application/json") + .content_language("en") + .content_disposition(format!("inline; filename=\"{s3_key}.json\"")); + let request = match context.ttl { + Some(ttl) => { + let seconds = ttl.as_secs_f64(); + request + .cache_control(format!("immutable, max-age={seconds}, s-maxage={seconds}")) + .expires(DateTime::from(SystemTime::now() + ttl)) + } + None => request.cache_control("immutable, max-age=31536000, s-maxage=31536000"), + }; + request.send().await.map_err(|_| Error::Unavailable)?; + Ok(()) + } + + async fn get(&self, key: &str) -> Result, Error> { + let output = match self + .client + .get_object() + .bucket(self.bucket.as_ref()) + .key(self.to_s3_key(key)) + .send() + .await + { + Ok(output) => output, + Err(error) => { + if let SdkError::ServiceError(service) = &error { + let status = error + .raw_response() + .map(|response| response.status().as_u16()); + let not_found = service.err().is_no_such_key() + || service.err().meta().code() == Some("AccessDenied") + || status == Some(404) + || status == Some(403); + if not_found { + return Ok(None); + } + } + return Err(Error::Unavailable); + } + }; + if let Some(expires) = output.expires_string() + && let Ok(expires) = DateTime::from_str(expires, Format::HttpDate) + && expires < DateTime::from(SystemTime::now()) + { + return Ok(None); + } + let bytes = output + .body + .collect() + .await + .map_err(|_| Error::Unavailable)? + .into_bytes(); + self.codec.decode(&bytes).map(Some) + } +} + +impl BaseCache for S3Cache { + type Value = C::Value; + type Context = ExactCacheContext; + + fn get_ttl(&self, context: &Self::Context) -> Option { + context.ttl + } + + fn set_cache( + &self, + key: &str, + value: Self::Value, + context: &Self::Context, + ) -> Result<(), Error> { + self.block_on(self.put(key, value, context)) + } + + fn get_cache(&self, key: &str, _context: &Self::Context) -> Result, Error> { + self.block_on(self.get(key)) + } + + async fn async_set_cache( + &self, + key: &str, + value: Self::Value, + context: Self::Context, + ) -> Result<(), Error> { + self.put(key, value, &context).await + } + + async fn async_get_cache( + &self, + key: &str, + _context: &Self::Context, + ) -> Result, Error> { + self.get(key).await + } + + async fn disconnect(&self) -> Result<(), Error> { + Ok(()) + } + + async fn test_connection(&self) -> Result { + Err(Error::UnsupportedOperation) + } +} + +impl BatchCache for S3Cache {} + +impl FlushCache for S3Cache { + fn flush_cache(&self) -> Result<(), Error> { + Ok(()) + } +} diff --git a/litellm-rust/crates/cache-s3/src/lib.rs b/litellm-rust/crates/cache-s3/src/lib.rs new file mode 100644 index 00000000000..f6126dfa908 --- /dev/null +++ b/litellm-rust/crates/cache-s3/src/lib.rs @@ -0,0 +1,4 @@ +mod auth; +mod cache; + +pub use cache::{S3Cache, S3CacheConfig, S3Endpoint}; diff --git a/litellm-rust/crates/cache-s3/tests/cache.rs b/litellm-rust/crates/cache-s3/tests/cache.rs new file mode 100644 index 00000000000..9a71656286b --- /dev/null +++ b/litellm-rust/crates/cache-s3/tests/cache.rs @@ -0,0 +1,278 @@ +use std::time::{Duration, SystemTime, UNIX_EPOCH}; + +use litellm_auth_aws::AwsAuthConfig; +use litellm_cache::{ + BaseCache, BatchCache, BatchEntry, Error, ExactCacheContext, FlushCache, JsonCodec, +}; +use litellm_cache_s3::{S3Cache, S3CacheConfig, S3Endpoint}; +use serde_json::{Value, json}; +use tokio::runtime::Handle; +use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{method, path}, +}; + +fn config(endpoint: String) -> S3CacheConfig { + S3CacheConfig { + bucket: "cache-bucket".to_string(), + key_prefix: "team/".to_string(), + region: "us-east-1".to_string(), + endpoint: Some(S3Endpoint { url: endpoint }), + auth: AwsAuthConfig { + access_key_id: Some("key".to_string()), + secret_access_key: Some("secret".to_string()), + region_name: Some("us-east-1".to_string()), + ..Default::default() + }, + } +} + +fn cache(endpoint: &str) -> S3Cache> { + S3Cache::new( + config(endpoint.to_string()), + JsonCodec::::new(), + Handle::current(), + ) +} + +async fn mock_server() -> MockServer { + let server = MockServer::start().await; + Mock::given(method("PUT")) + .respond_with(ResponseTemplate::new(200).insert_header("etag", "\"etag\"")) + .mount(&server) + .await; + server +} + +fn http_date_from(headers: &wiremock::http::HeaderMap, name: &str) -> Option { + use aws_smithy_types::{DateTime, date_time::Format}; + headers + .get(name) + .and_then(|value| DateTime::from_str(value.to_str().ok()?, Format::HttpDate).ok()) + .map(|date| UNIX_EPOCH + Duration::new(date.secs() as u64, date.subsec_nanos())) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn set_writes_python_metadata_with_and_without_ttl() { + let server = mock_server().await; + let cache = cache(&server.uri()); + let context = ExactCacheContext { + ttl: Some(Duration::from_secs(90)), + }; + cache + .set_cache("alpha:beta", json!({"answer": 1}), &context) + .unwrap(); + cache + .set_cache("plain", json!({"answer": 2}), &ExactCacheContext::default()) + .unwrap(); + + let requests = server.received_requests().await.unwrap(); + let ttl_request = requests + .iter() + .find(|request| request.url.path() == "/cache-bucket/team/alpha/beta") + .expect("ttl write should hit the converted S3 key"); + assert_eq!( + ttl_request.headers["cache-control"].to_str().unwrap(), + "immutable, max-age=90, s-maxage=90" + ); + assert_eq!( + ttl_request.headers["content-type"].to_str().unwrap(), + "application/json" + ); + assert_eq!( + ttl_request.headers["content-language"].to_str().unwrap(), + "en" + ); + assert_eq!( + ttl_request.headers["content-disposition"].to_str().unwrap(), + "inline; filename=\"team/alpha/beta.json\"" + ); + let expires = http_date_from(&ttl_request.headers, "expires").expect("ttl write sets Expires"); + let remaining = expires.duration_since(SystemTime::now()).unwrap(); + assert!(remaining > Duration::from_secs(60) && remaining <= Duration::from_secs(91)); + assert_eq!( + serde_json::from_slice::(&ttl_request.body).unwrap(), + json!({"answer": 1}) + ); + + let plain = requests + .iter() + .find(|request| request.url.path() == "/cache-bucket/team/plain") + .expect("no-ttl write should hit the converted S3 key"); + assert_eq!( + plain.headers["cache-control"].to_str().unwrap(), + "immutable, max-age=31536000, s-maxage=31536000" + ); + assert!(plain.headers.get("expires").is_none()); + assert_eq!( + plain.headers["content-disposition"].to_str().unwrap(), + "inline; filename=\"team/plain.json\"" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn get_hit_miss_expired_and_invalid_entries() { + let server = mock_server().await; + Mock::given(method("GET")) + .and(path("/cache-bucket/team/hit")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"answer": 3}))) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/cache-bucket/team/missing")) + .respond_with( + ResponseTemplate::new(404).set_body_string("NoSuchKey"), + ) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/cache-bucket/team/denied")) + .respond_with( + ResponseTemplate::new(403).set_body_string("AccessDenied"), + ) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/cache-bucket/team/expired")) + .respond_with( + ResponseTemplate::new(200) + .insert_header("expires", "Thu, 01 Jan 1970 00:00:00 GMT") + .set_body_json(json!({"answer": 4})), + ) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/cache-bucket/team/malformed")) + .respond_with(ResponseTemplate::new(200).set_body_string("not a cache entry")) + .mount(&server) + .await; + let cache = cache(&server.uri()); + let context = ExactCacheContext::default(); + + assert_eq!( + cache.get_cache("hit", &context).unwrap(), + Some(json!({"answer": 3})) + ); + assert_eq!(cache.get_cache("missing", &context).unwrap(), None); + assert_eq!(cache.get_cache("denied", &context).unwrap(), None); + assert_eq!(cache.get_cache("expired", &context).unwrap(), None); + assert_eq!( + cache.get_cache("malformed", &context), + Err(Error::InvalidEntry) + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn batch_get_preserves_order_with_hits_misses_and_invalid() { + let server = mock_server().await; + for (key, status, body) in [ + ("first", 200, "{\"answer\": 1}"), + ("invalid", 200, "garbage"), + ] { + Mock::given(method("GET")) + .and(path(format!("/cache-bucket/team/{key}"))) + .respond_with(ResponseTemplate::new(status).set_body_string(body)) + .mount(&server) + .await; + } + Mock::given(method("GET")) + .and(path("/cache-bucket/team/miss")) + .respond_with(ResponseTemplate::new(404)) + .mount(&server) + .await; + let cache = cache(&server.uri()); + let context = ExactCacheContext::default(); + let keys = vec![ + "first".to_string(), + "miss".to_string(), + "invalid".to_string(), + ]; + + let entries = cache.batch_get_cache(&keys, &context).unwrap(); + + assert_eq!( + entries, + vec![ + BatchEntry::Hit(json!({"answer": 1})), + BatchEntry::Miss, + BatchEntry::Invalid, + ] + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn unsupported_and_noop_capabilities_match_python() { + let server = mock_server().await; + let cache = cache(&server.uri()); + + assert_eq!( + cache.test_connection().await, + Err(Error::UnsupportedOperation) + ); + cache.flush_cache().unwrap(); + cache.disconnect().await.unwrap(); + assert_eq!(cache.get_ttl(&ExactCacheContext::default()), None); + assert_eq!( + cache.get_ttl(&ExactCacheContext { + ttl: Some(Duration::from_secs(45)), + }), + Some(Duration::from_secs(45)) + ); + assert!(server.received_requests().await.unwrap().is_empty()); +} + +#[test] +fn key_conversion_prefixes_and_splits_colons() { + let runtime = tokio::runtime::Builder::new_multi_thread() + .worker_threads(1) + .enable_all() + .build() + .unwrap(); + let _guard = runtime.enter(); + let cache = S3Cache::new( + S3CacheConfig { + key_prefix: "team/".to_string(), + ..config("http://localhost".to_string()) + }, + JsonCodec::::new(), + runtime.handle().clone(), + ); + + assert_eq!(cache.bucket(), "cache-bucket"); + assert_eq!(cache.key_prefix(), "team/"); + assert_eq!(cache.to_s3_key("a:b:c"), "team/a/b/c"); + assert_eq!(cache.to_s3_key("plain"), "team/plain"); + + let unprefixed = S3Cache::new( + S3CacheConfig { + key_prefix: String::new(), + ..config("http://localhost".to_string()) + }, + JsonCodec::::new(), + runtime.handle().clone(), + ); + assert_eq!(unprefixed.to_s3_key("a:b"), "a/b"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn sync_methods_block_inside_and_outside_the_runtime() { + let server = mock_server().await; + Mock::given(method("GET")) + .and(path("/cache-bucket/team/key")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"answer": 9}))) + .mount(&server) + .await; + let uri = server.uri(); + let cache = tokio::task::spawn_blocking(move || { + let cache = cache(&uri); + let context = ExactCacheContext::default(); + cache + .set_cache("key", json!({"answer": 9}), &context) + .unwrap(); + cache.get_cache("key", &context).unwrap() + }) + .await + .unwrap(); + + assert_eq!(cache, Some(json!({"answer": 9}))); +} diff --git a/litellm-rust/crates/cache-valkey-semantic/Cargo.toml b/litellm-rust/crates/cache-valkey-semantic/Cargo.toml new file mode 100644 index 00000000000..f98bb5a5fa8 --- /dev/null +++ b/litellm-rust/crates/cache-valkey-semantic/Cargo.toml @@ -0,0 +1,20 @@ +[package] +name = "litellm-cache-valkey-semantic" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +litellm-cache.workspace = true +litellm-cache-redis.workspace = true +litellm-cache-response.workspace = true +redis = { version = "1.7.0", features = ["tls-rustls"] } +serde_json.workspace = true +sha2.workspace = true +tokio.workspace = true +uuid = { version = "1", features = ["v4"] } + +[dev-dependencies] +redis-test = "1.0.4" +rstest.workspace = true diff --git a/litellm-rust/crates/cache-valkey-semantic/src/lib.rs b/litellm-rust/crates/cache-valkey-semantic/src/lib.rs new file mode 100644 index 00000000000..6062ccc842c --- /dev/null +++ b/litellm-rust/crates/cache-valkey-semantic/src/lib.rs @@ -0,0 +1,1153 @@ +use std::{ + future::Future, + sync::{Arc, Mutex}, + time::Duration, +}; + +use litellm_cache::{BaseCache, CacheCodec, CacheConnectionResult, Error, SemanticCacheContext}; +use litellm_cache_redis::{ + RedisTopology, + connection::{ConnectionRef, Connections}, +}; +use litellm_cache_response::CacheEntry; +use serde_json::Value; +use sha2::{Digest, Sha256}; +use uuid::Uuid; + +pub trait Embedder: Send + Sync + 'static { + fn embed(&self, prompt: &str, metadata: Option<&Value>) -> Result, Error>; + + fn async_embed( + &self, + prompt: &str, + metadata: Option<&Value>, + ) -> impl Future, Error>> + Send; +} + +pub struct PreparedEmbedding(pub Vec); + +impl Embedder for PreparedEmbedding { + fn embed(&self, _prompt: &str, _metadata: Option<&Value>) -> Result, Error> { + Ok(self.0.clone()) + } + + async fn async_embed( + &self, + _prompt: &str, + _metadata: Option<&Value>, + ) -> Result, Error> { + Ok(self.0.clone()) + } +} + +#[derive(Clone, Debug, PartialEq)] +pub struct ValkeySemanticConfig { + pub similarity_threshold: f64, + pub index_name: String, +} + +pub const DEFAULT_INDEX_NAME: &str = "litellm_semantic_cache_index"; + +#[derive(Clone)] +struct IndexState { + name: String, + prefix: String, + dimension: Arc>>, + similarity_threshold: f64, +} + +pub struct ValkeySemanticCache< + E: Embedder, + S: CacheCodec, + C = redis::Connection, +> { + connections: Arc>, + embedder: E, + codec: S, + config: ValkeySemanticConfig, + index_dimension: Arc>>, +} + +impl ValkeySemanticCache +where + E: Embedder, + S: CacheCodec, +{ + pub fn new( + url: &str, + embedder: E, + codec: S, + config: ValkeySemanticConfig, + ) -> Result { + Ok(Self { + connections: Arc::new(Connections::open(url, &RedisTopology::Standalone)?), + embedder, + codec, + config, + index_dimension: Arc::new(Mutex::new(None)), + }) + } +} + +impl ValkeySemanticCache +where + E: Embedder, + S: CacheCodec, + C: redis::ConnectionLike + Send + 'static, +{ + pub fn with_connection( + connection: C, + embedder: E, + codec: S, + config: ValkeySemanticConfig, + ) -> Self { + Self { + connections: Arc::new(Connections::fixed(connection)), + embedder, + codec, + config, + index_dimension: Arc::new(Mutex::new(None)), + } + } + + pub fn similarity_threshold(&self) -> f64 { + self.config.similarity_threshold + } + + pub fn index_name(&self) -> &str { + &self.config.index_name + } + + fn index_state(&self) -> IndexState { + IndexState { + name: self.config.index_name.clone(), + prefix: format!("{}:", self.config.index_name), + dimension: Arc::clone(&self.index_dimension), + similarity_threshold: self.config.similarity_threshold, + } + } +} + +impl ValkeySemanticCache +where + E: Embedder, + S: CacheCodec + Clone, + C: redis::ConnectionLike + Send + 'static, +{ + pub fn with_embedder(&self, embedder: E2) -> ValkeySemanticCache { + ValkeySemanticCache { + connections: Arc::clone(&self.connections), + embedder, + codec: self.codec.clone(), + config: self.config.clone(), + index_dimension: Arc::clone(&self.index_dimension), + } + } +} + +impl BaseCache for ValkeySemanticCache +where + E: Embedder, + S: CacheCodec, + C: redis::ConnectionLike + Send + 'static, +{ + type Value = CacheEntry; + type Context = SemanticCacheContext; + + fn get_ttl(&self, context: &Self::Context) -> Option { + context.ttl + } + + fn set_cache( + &self, + key: &str, + value: Self::Value, + context: &Self::Context, + ) -> Result<(), Error> { + let Some(prompt) = prompt_from_context(context) else { + return Ok(()); + }; + let embedding = self.embedder.embed(&prompt, context.metadata.as_ref())?; + let scope = scope_tag(key); + let response = self.codec.encode(&value)?; + let vector = embedding_bytes(&embedding); + let index = self.index_state(); + self.connections.execute(|connection| { + write_document( + connection, + &index, + &scope, + &prompt, + response, + vector, + self.get_ttl(context), + ) + }) + } + + fn get_cache(&self, key: &str, context: &Self::Context) -> Result, Error> { + let Some(prompt) = prompt_from_context(context) else { + return Ok(None); + }; + let embedding = self.embedder.embed(&prompt, context.metadata.as_ref())?; + let scope = scope_tag(key); + let vector = embedding_bytes(&embedding); + let index = self.index_state(); + let response = self.connections.execute(|connection| { + search_document(connection, &index, &scope, vector, embedding.len()) + })?; + let Some(response) = response else { + return Ok(None); + }; + self.codec.decode(&response).map(Some) + } + + fn async_set_cache( + &self, + key: &str, + value: Self::Value, + context: Self::Context, + ) -> impl Future> + Send { + let key = key.to_owned(); + let prompt = prompt_from_context(&context); + let metadata = context.metadata.clone(); + async move { + let Some(prompt) = prompt else { + return Ok(()); + }; + let embedding = self + .embedder + .async_embed(&prompt, metadata.as_ref()) + .await?; + let connections = Arc::clone(&self.connections); + let index = self.index_state(); + let response = self.codec.encode(&value)?; + let vector = embedding_bytes(&embedding); + let scope = scope_tag(&key); + let ttl = context.ttl; + Connections::run_blocking(connections, move |connection| { + write_document(connection, &index, &scope, &prompt, response, vector, ttl) + }) + .await + } + } + + fn async_get_cache( + &self, + key: &str, + context: &Self::Context, + ) -> impl Future, Error>> + Send { + let key = key.to_owned(); + let prompt = prompt_from_context(context); + let metadata = context.metadata.clone(); + async move { + let Some(prompt) = prompt else { + return Ok(None); + }; + let embedding = self + .embedder + .async_embed(&prompt, metadata.as_ref()) + .await?; + let connections = Arc::clone(&self.connections); + let index = self.index_state(); + Connections::run_blocking(connections, move |connection| { + let scope = scope_tag(&key); + let vector = embedding_bytes(&embedding); + search_document(connection, &index, &scope, vector, embedding.len()) + }) + .await + .and_then(|response| response.map(|bytes| self.codec.decode(&bytes)).transpose()) + } + } + + async fn disconnect(&self) -> Result<(), Error> { + Ok(()) + } + + async fn test_connection(&self) -> Result { + Err(Error::UnsupportedOperation) + } +} + +pub fn prompt_from_context(context: &SemanticCacheContext) -> Option { + if let Some(Value::Array(messages)) = context.messages.as_ref() + && !messages.is_empty() + { + return messages + .iter() + .filter_map(Value::as_object) + .map(message_text) + .collect(); + } + let input = context.input.as_ref()?; + let mut parts = Vec::new(); + collect_input_text(input, &mut parts); + let prompt = parts.join("\n").trim().to_owned(); + (!prompt.is_empty()).then_some(prompt) +} + +fn message_text(message: &serde_json::Map) -> Option { + let content = match message.get("content") { + Some(Value::String(value)) => value.clone(), + Some(Value::Array(parts)) => { + let mut content = String::new(); + for part in parts { + let part = part.as_object()?; + if let Some(text) = part.get("text").and_then(Value::as_str) { + content.push_str(text); + } + } + content + } + _ => String::new(), + }; + Some(format!( + "{content}{}", + search_results_text(message.get("search_results")) + )) +} + +fn search_results_text(value: Option<&Value>) -> String { + let Some(Value::Array(results)) = value else { + return String::new(); + }; + results + .iter() + .filter_map(Value::as_object) + .map(|result| { + let source = result.get("source").and_then(Value::as_str).unwrap_or(""); + let title = result.get("title").and_then(Value::as_str).unwrap_or(""); + let content = result + .get("content") + .and_then(Value::as_array) + .map(|blocks| { + blocks + .iter() + .filter_map(Value::as_object) + .filter_map(|block| block.get("text").and_then(Value::as_str)) + .collect::() + }) + .unwrap_or_default(); + let citations = result + .get("citations") + .filter(|value| !value.is_null()) + .and_then(|value| serde_json::to_string(value).ok()) + .unwrap_or_default(); + format!("{source}{title}{content}{citations}") + }) + .collect() +} + +fn collect_input_text(value: &Value, parts: &mut Vec) { + match value { + Value::String(value) => { + let value = value.trim(); + if !value.is_empty() { + parts.push(value.to_owned()); + } + } + Value::Array(values) => values + .iter() + .for_each(|value| collect_input_text(value, parts)), + Value::Object(object) => { + if let Some(content) = object.get("content").filter(|value| !value.is_null()) { + collect_input_text(content, parts); + return; + } + for key in ["text", "output", "input_text", "output_text"] { + if let Some(Value::String(value)) = object.get(key) { + let value = value.trim(); + if !value.is_empty() { + parts.push(value.to_owned()); + return; + } + } + } + } + _ => {} + } +} + +fn scope_tag(key: &str) -> String { + let digest = Sha256::digest(key.as_bytes()); + digest.iter().map(|byte| format!("{byte:02x}")).collect() +} + +fn embedding_bytes(embedding: &[f32]) -> Vec { + embedding + .iter() + .flat_map(|value| value.to_le_bytes()) + .collect() +} + +fn write_document( + connection: &mut ConnectionRef<'_>, + index: &IndexState, + scope: &str, + prompt: &str, + response: Vec, + vector: Vec, + ttl: Option, +) -> Result<(), Error> { + let dimension = vector.len() / std::mem::size_of::(); + ensure_index( + connection, + &index.name, + &index.prefix, + &index.dimension, + dimension, + )?; + let document = format!("{}{scope}:{}", index.prefix, Uuid::new_v4()); + let mut pipeline = redis::pipe(); + pipeline + .cmd("HSET") + .arg(&document) + .arg("litellm_cache_key") + .arg(scope) + .arg("prompt") + .arg(prompt) + .arg("response") + .arg(response) + .arg("embedding") + .arg(vector) + .ignore(); + if let Some(ttl) = ttl { + pipeline + .cmd("EXPIRE") + .arg(&document) + .arg(ttl.as_secs()) + .ignore(); + } + pipeline + .query::<()>(connection) + .map_err(|_| Error::Unavailable) +} + +fn search_document( + connection: &mut ConnectionRef<'_>, + index: &IndexState, + scope: &str, + vector: Vec, + dimension: usize, +) -> Result>, Error> { + ensure_index( + connection, + &index.name, + &index.prefix, + &index.dimension, + dimension, + )?; + let query = + format!("(@litellm_cache_key:{{{scope}}})=>[KNN 1 @embedding $vec AS vector_distance]"); + let response = redis::cmd("FT.SEARCH") + .arg(&index.name) + .arg(query) + .arg("PARAMS") + .arg(2) + .arg("vec") + .arg(vector) + .arg("RETURN") + .arg(2) + .arg("response") + .arg("vector_distance") + .arg("DIALECT") + .arg(2) + .query::(connection) + .map_err(|_| Error::Unavailable)?; + let Some(fields) = search_fields(response)? else { + return Ok(None); + }; + let response = fields + .iter() + .find_map(|(name, value)| (name == "response").then(|| value.clone())) + .ok_or(Error::InvalidEntry)?; + let distance = fields + .iter() + .find_map(|(name, value)| (name == "vector_distance").then(|| value.clone())) + .ok_or(Error::InvalidEntry)?; + let distance = parse_f64(&distance)?; + if 1.0 - distance < index.similarity_threshold { + return Ok(None); + } + Ok(Some(response)) +} + +fn ensure_index( + connection: &mut ConnectionRef<'_>, + index_name: &str, + prefix: &str, + index_dimension: &Mutex>, + dimension: usize, +) -> Result<(), Error> { + if index_dimension + .lock() + .map_err(|_| Error::Unavailable)? + .is_some_and(|existing| existing == dimension) + { + return Ok(()); + } + let create = redis::cmd("FT.CREATE") + .arg(index_name) + .arg("ON") + .arg("HASH") + .arg("PREFIX") + .arg(1) + .arg(prefix) + .arg("SCHEMA") + .arg("litellm_cache_key") + .arg("TAG") + .arg("embedding") + .arg("VECTOR") + .arg("HNSW") + .arg(6) + .arg("TYPE") + .arg("FLOAT32") + .arg("DIM") + .arg(dimension) + .arg("DISTANCE_METRIC") + .arg("COSINE") + .query::(connection) + .map(|_| ()) + .map_err(|error| error.to_string()); + if let Err(message) = create { + if !message.to_ascii_lowercase().contains("already exists") { + return Err(Error::Unavailable); + } + let info = redis::cmd("FT.INFO") + .arg(index_name) + .query::(connection) + .map_err(|_| Error::Unavailable)?; + let existing = index_dimension_from_info(&info).ok_or(Error::Unavailable)?; + if existing != dimension { + return Err(Error::Unavailable); + } + } + *index_dimension.lock().map_err(|_| Error::Unavailable)? = Some(dimension); + Ok(()) +} + +fn index_dimension_from_info(value: &redis::Value) -> Option { + let redis::Value::Array(values) = value else { + return None; + }; + let attributes = values.windows(2).find_map(|pair| { + (value_text(&pair[0]).as_deref() == Some("attributes")).then_some(&pair[1]) + })?; + let redis::Value::Array(fields) = attributes else { + return None; + }; + fields.iter().find_map(|field| { + let redis::Value::Array(values) = field else { + return None; + }; + let flattened = values.iter().flat_map(|value| match value { + redis::Value::Array(values) => values.as_slice(), + _ => std::slice::from_ref(value), + }); + let values = flattened.collect::>(); + values.windows(2).find_map(|pair| { + if value_text(pair[0]).as_deref() == Some("dimensions") { + return value_text(pair[1]).and_then(|value| value.parse().ok()); + } + None + }) + }) +} + +type SearchFields = Vec<(String, Vec)>; + +fn search_fields(value: redis::Value) -> Result, Error> { + let redis::Value::Array(values) = value else { + return Err(Error::InvalidEntry); + }; + let total = parse_i64(values.first().ok_or(Error::InvalidEntry)?)?; + if total <= 0 || values.len() < 3 { + return Ok(None); + } + let redis::Value::Array(fields) = &values[2] else { + return Err(Error::InvalidEntry); + }; + let (pairs, remainder) = fields.as_chunks::<2>(); + if !remainder.is_empty() { + return Err(Error::InvalidEntry); + } + let pairs = pairs + .iter() + .map(|pair| { + Ok(( + value_text(&pair[0]).ok_or(Error::InvalidEntry)?, + value_bytes(&pair[1])?, + )) + }) + .collect::, Error>>()?; + Ok(Some(pairs)) +} + +fn parse_i64(value: &redis::Value) -> Result { + value_text(value) + .ok_or(Error::InvalidEntry)? + .parse() + .map_err(|_| Error::InvalidEntry) +} + +fn parse_f64(value: &[u8]) -> Result { + std::str::from_utf8(value) + .map_err(|_| Error::InvalidEntry)? + .parse() + .map_err(|_| Error::InvalidEntry) +} + +fn value_text(value: &redis::Value) -> Option { + match value { + redis::Value::BulkString(bytes) => String::from_utf8(bytes.clone()).ok(), + redis::Value::SimpleString(value) => Some(value.clone()), + redis::Value::Int(value) => Some(value.to_string()), + _ => None, + } +} + +fn value_bytes(value: &redis::Value) -> Result, Error> { + match value { + redis::Value::BulkString(bytes) => Ok(bytes.clone()), + redis::Value::SimpleString(value) => Ok(value.as_bytes().to_vec()), + redis::Value::Int(value) => Ok(value.to_string().into_bytes()), + _ => Err(Error::InvalidEntry), + } +} + +#[cfg(test)] +mod tests { + use std::{ + collections::VecDeque, + sync::{Arc, Mutex}, + time::Duration, + }; + + use litellm_cache::{BaseCache, CacheCodec}; + use litellm_cache_response::{ + CacheEntry, CacheKeyInput, ResponseCache, ResponseCacheCodec, ResponseCacheRequest, + }; + use redis_test::MockRedisConnection; + use rstest::rstest; + use serde_json::{Value, json}; + + use super::{ + Embedder, PreparedEmbedding, ValkeySemanticCache, ValkeySemanticConfig, + index_dimension_from_info, prompt_from_context, scope_tag, + }; + + #[derive(Clone)] + struct FixedEmbedder { + vector: Vec, + calls: EmbedderCalls, + } + + type EmbedderCalls = Arc)>>>; + type RecordingCache = + ValkeySemanticCache; + type RecordingSetup = (RecordingCache, Arc>>>, EmbedderCalls); + + impl Embedder for FixedEmbedder { + fn embed(&self, prompt: &str, metadata: Option<&Value>) -> Result, super::Error> { + self.calls + .lock() + .unwrap() + .push((prompt.into(), metadata.cloned())); + Ok(self.vector.clone()) + } + + async fn async_embed( + &self, + prompt: &str, + metadata: Option<&Value>, + ) -> Result, super::Error> { + self.embed(prompt, metadata) + } + } + + struct RecordingConnection { + requests: Arc>>>, + replies: Mutex>>, + } + + impl RecordingConnection { + fn new(replies: impl IntoIterator>) -> Self { + Self { + requests: Arc::default(), + replies: Mutex::new(replies.into_iter().collect()), + } + } + + fn requests(&self) -> Arc>>> { + Arc::clone(&self.requests) + } + + fn reply(&self) -> redis::RedisResult { + self.replies + .lock() + .unwrap() + .pop_front() + .unwrap_or_else(|| Ok(redis::Value::SimpleString("OK".into()))) + } + } + + impl redis::ConnectionLike for RecordingConnection { + fn req_packed_command(&mut self, command: &[u8]) -> redis::RedisResult { + self.requests.lock().unwrap().push(command.to_vec()); + self.reply() + } + + fn req_packed_commands( + &mut self, + command: &[u8], + _offset: usize, + count: usize, + ) -> redis::RedisResult> { + self.requests.lock().unwrap().push(command.to_vec()); + (0..count).map(|_| self.reply()).collect() + } + + fn get_db(&self) -> i64 { + 0 + } + + fn check_connection(&mut self) -> bool { + true + } + + fn is_open(&self) -> bool { + true + } + } + + fn context( + messages: Option, + input: Option, + ) -> litellm_cache::SemanticCacheContext { + litellm_cache::SemanticCacheContext { + messages, + input, + ..Default::default() + } + } + + #[rstest] + #[case(json!([{"content": "hello"}]), None, Some("hello"))] + #[case(json!([{"content": [{"text": "hello"}, {"text": " world"}]}]), None, Some("hello world"))] + #[case(json!([{"content": ["raw", {"text": "hello"}]}]), None, None)] + #[case(json!([{"search_results": [{"source": "s", "title": "t", "content": [{"text": "c"}], "citations": ["x"]}]}]), None, Some(r#"stc["x"]"#))] + #[case(Value::Array(vec![]), Some(json!(" hello ")), Some("hello"))] + #[case(Value::Array(vec![]), Some(json!([{"content": "first"}, {"text": "second"}])), Some("first\nsecond"))] + #[case(Value::Array(vec![]), Some(json!(" ")), None)] + fn prompt_shapes( + #[case] messages: Value, + #[case] input: Option, + #[case] expected: Option<&str>, + ) { + assert_eq!( + prompt_from_context(&context(Some(messages), input)), + expected.map(str::to_owned) + ); + } + + #[test] + fn scope_tags_are_lowercase_sha256() { + assert_eq!( + scope_tag("key"), + "2c70e12b7a0646f92279f427c7b38e7334d8e5389cff167a1dc30e73f826b683" + ); + } + + #[test] + fn existing_index_dimension_is_read_from_attributes() { + let info = redis::Value::Array(vec![ + redis::Value::SimpleString("attributes".into()), + redis::Value::Array(vec![redis::Value::Array(vec![ + redis::Value::SimpleString("identifier".into()), + redis::Value::SimpleString("embedding".into()), + redis::Value::Array(vec![ + redis::Value::SimpleString("dimensions".into()), + redis::Value::SimpleString("2".into()), + ]), + ])]), + ]); + assert_eq!(index_dimension_from_info(&info), Some(2)); + } + + #[tokio::test] + async fn unsupported_connection_test_is_reported() { + let cache = ValkeySemanticCache::with_connection( + MockRedisConnection::new([]).assert_all_commands_consumed(), + FixedEmbedder { + vector: vec![1.0, 0.0], + calls: Arc::default(), + }, + ResponseCacheCodec, + ValkeySemanticConfig { + similarity_threshold: 0.8, + index_name: "test".into(), + }, + ); + assert_eq!( + cache.test_connection().await, + Err(super::Error::UnsupportedOperation) + ); + } + + #[tokio::test] + async fn prepared_embedding_returns_its_vector_for_any_prompt() { + let embedding = PreparedEmbedding(vec![1.0, 2.0]); + assert_eq!( + embedding + .async_embed("different prompt", None) + .await + .unwrap(), + vec![1.0, 2.0] + ); + } + + #[test] + fn with_embedder_shares_index_state_and_connections() { + let entry = CacheEntry { + timestamp: Some(1.0), + response: json!({"answer": "ok"}), + }; + let encoded = ResponseCacheCodec.encode(&entry).unwrap(); + let cache = ValkeySemanticCache::with_connection( + RecordingConnection::new([ok(), ok(), Ok(search_hit(encoded, "0.1"))]), + FixedEmbedder { + vector: vec![1.0, 0.0], + calls: Arc::default(), + }, + ResponseCacheCodec, + ValkeySemanticConfig { + similarity_threshold: 0.8, + index_name: "test".into(), + }, + ); + cache + .set_cache("key", entry.clone(), &semantic_context(None)) + .unwrap(); + let prepared = cache.with_embedder(PreparedEmbedding(vec![1.0, 0.0])); + assert_eq!( + prepared.get_cache("key", &semantic_context(None)).unwrap(), + Some(entry) + ); + } + + #[test] + fn missing_prompt_does_not_touch_redis() { + let cache = ValkeySemanticCache::with_connection( + MockRedisConnection::new([]).assert_all_commands_consumed(), + FixedEmbedder { + vector: vec![1.0, 0.0], + calls: Arc::default(), + }, + ResponseCacheCodec, + ValkeySemanticConfig { + similarity_threshold: 0.8, + index_name: "test".into(), + }, + ); + assert_eq!(cache.get_cache("key", &context(None, None)).unwrap(), None); + assert_eq!(cache.get_ttl(&context(None, None)), None); + } + + fn semantic_context(ttl: Option) -> litellm_cache::SemanticCacheContext { + litellm_cache::SemanticCacheContext { + messages: Some(json!([{"role": "user", "content": "hello"}])), + metadata: Some(json!({"source": "test"})), + ttl, + ..Default::default() + } + } + + fn cache_with_recording( + replies: impl IntoIterator>, + vector: Vec, + threshold: f64, + ) -> RecordingSetup { + let connection = RecordingConnection::new(replies); + let requests = connection.requests(); + let calls: EmbedderCalls = Arc::default(); + let cache = ValkeySemanticCache::with_connection( + connection, + FixedEmbedder { + vector, + calls: Arc::clone(&calls), + }, + ResponseCacheCodec, + ValkeySemanticConfig { + similarity_threshold: threshold, + index_name: "test".into(), + }, + ); + (cache, requests, calls) + } + + fn ok() -> redis::RedisResult { + Ok(redis::Value::SimpleString("OK".into())) + } + + fn already_exists() -> redis::RedisResult { + Err(redis::RedisError::from(( + redis::ErrorKind::Io, + "already exists", + ))) + } + + fn info_dimension(dimension: usize) -> redis::Value { + redis::Value::Array(vec![ + redis::Value::SimpleString("attributes".into()), + redis::Value::Array(vec![redis::Value::Array(vec![ + redis::Value::SimpleString("embedding".into()), + redis::Value::Array(vec![ + redis::Value::SimpleString("dimensions".into()), + redis::Value::Int(dimension as i64), + ]), + ])]), + ]) + } + + fn search_hit(response: Vec, distance: &str) -> redis::Value { + redis::Value::Array(vec![ + redis::Value::Int(1), + redis::Value::BulkString(b"test:document".to_vec()), + redis::Value::Array(vec![ + redis::Value::BulkString(b"response".to_vec()), + redis::Value::BulkString(response), + redis::Value::BulkString(b"vector_distance".to_vec()), + redis::Value::BulkString(distance.as_bytes().to_vec()), + ]), + ]) + } + + fn requests_text(requests: &Arc>>>) -> String { + requests + .lock() + .unwrap() + .iter() + .map(|request| String::from_utf8_lossy(request)) + .collect::>() + .join("\n") + } + + #[test] + fn set_without_ttl_writes_hset_without_expire() { + let (cache, requests, calls) = cache_with_recording([ok()], vec![1.0, 0.0], 0.8); + cache + .set_cache( + "key", + CacheEntry { + timestamp: None, + response: json!({"answer": "ok"}), + }, + &semantic_context(None), + ) + .unwrap(); + let text = requests_text(&requests); + assert!(text.contains("FT.CREATE")); + assert!(text.contains("HSET")); + assert!( + text.contains("test:2c70e12b7a0646f92279f427c7b38e7334d8e5389cff167a1dc30e73f826b683:") + ); + assert!(!text.contains("EXPIRE")); + assert_eq!( + *calls.lock().unwrap(), + vec![("hello".into(), Some(json!({"source": "test"})))] + ); + } + + #[test] + fn set_with_ttl_truncates_expire_seconds() { + let (cache, requests, _) = cache_with_recording([ok()], vec![1.0, 0.0], 0.8); + cache + .set_cache( + "key", + CacheEntry { + timestamp: None, + response: json!({"answer": "ok"}), + }, + &semantic_context(Some(Duration::from_millis(1900))), + ) + .unwrap(); + let text = requests_text(&requests); + assert!(text.contains("EXPIRE")); + assert!(text.contains("\r\n$1\r\n1\r\n")); + } + + #[test] + fn second_set_skips_create_after_dimension_is_cached() { + let (cache, requests, _) = cache_with_recording([ok()], vec![1.0, 0.0], 0.8); + let context = semantic_context(None); + let entry = CacheEntry { + timestamp: None, + response: json!({"answer": "ok"}), + }; + cache.set_cache("key", entry.clone(), &context).unwrap(); + cache.set_cache("key", entry, &context).unwrap(); + let text = requests_text(&requests); + assert_eq!(text.matches("FT.CREATE").count(), 1); + assert_eq!(text.matches("HSET").count(), 2); + } + + #[test] + fn existing_index_dimension_must_match_embedding() { + let (cache, _, _) = cache_with_recording( + [already_exists(), Ok(info_dimension(2))], + vec![1.0, 0.0], + 0.8, + ); + cache + .set_cache( + "key", + CacheEntry { + timestamp: None, + response: json!({"answer": "ok"}), + }, + &semantic_context(None), + ) + .unwrap(); + + let (cache, _, _) = cache_with_recording( + [already_exists(), Ok(info_dimension(3))], + vec![1.0, 0.0], + 0.8, + ); + assert_eq!( + cache.set_cache( + "key", + CacheEntry { + timestamp: None, + response: json!({"answer": "ok"}), + }, + &semantic_context(None), + ), + Err(super::Error::Unavailable) + ); + } + + #[test] + fn get_applies_threshold_and_decodes_entry() { + let entry = CacheEntry { + timestamp: Some(1.0), + response: json!({"answer": "ok"}), + }; + let encoded = ResponseCacheCodec.encode(&entry).unwrap(); + let (cache, _, _) = cache_with_recording( + [ok(), Ok(search_hit(encoded.clone(), "0.1"))], + vec![1.0, 0.0], + 0.8, + ); + assert_eq!( + cache.get_cache("key", &semantic_context(None)).unwrap(), + Some(entry) + ); + + let (cache, _, _) = + cache_with_recording([ok(), Ok(search_hit(encoded, "0.5"))], vec![1.0, 0.0], 0.8); + assert_eq!( + cache.get_cache("key", &semantic_context(None)).unwrap(), + None + ); + } + + #[test] + fn get_zero_docs_is_a_miss() { + let (cache, _, _) = cache_with_recording( + [ok(), Ok(redis::Value::Array(vec![redis::Value::Int(0)]))], + vec![1.0, 0.0], + 0.8, + ); + assert_eq!( + cache.get_cache("key", &semantic_context(None)).unwrap(), + None + ); + } + + #[rstest] + #[case(redis::Value::Array(vec![ + redis::Value::Int(1), + redis::Value::BulkString(b"document".to_vec()), + redis::Value::Array(vec![ + redis::Value::BulkString(b"vector_distance".to_vec()), + redis::Value::BulkString(b"0.1".to_vec()), + ]), + ]))] + #[case(redis::Value::Array(vec![ + redis::Value::Int(1), + redis::Value::BulkString(b"document".to_vec()), + redis::Value::Array(vec![ + redis::Value::BulkString(b"response".to_vec()), + redis::Value::BulkString(b"not-json".to_vec()), + redis::Value::BulkString(b"vector_distance".to_vec()), + redis::Value::BulkString(b"abc".to_vec()), + ]), + ]))] + fn malformed_entries_are_invalid(#[case] search: redis::Value) { + let (cache, _, _) = cache_with_recording([ok(), Ok(search)], vec![1.0, 0.0], 0.8); + assert_eq!( + cache.get_cache("key", &semantic_context(None)), + Err(super::Error::InvalidEntry) + ); + } + + #[test] + fn response_cache_turns_invalid_entries_into_misses() { + let (cache, _, _) = cache_with_recording( + [ + ok(), + Ok(redis::Value::Array(vec![ + redis::Value::Int(1), + redis::Value::BulkString(b"document".to_vec()), + redis::Value::Array(vec![ + redis::Value::BulkString(b"response".to_vec()), + redis::Value::BulkString(b"not-json".to_vec()), + redis::Value::BulkString(b"vector_distance".to_vec()), + redis::Value::BulkString(b"0.1".to_vec()), + ]), + ])), + ], + vec![1.0, 0.0], + 0.8, + ); + let service = ResponseCache::new(Arc::new(cache)); + let request = ResponseCacheRequest { + key: CacheKeyInput { + preset: Some("key".into()), + ..Default::default() + }, + context: semantic_context(None), + ..ResponseCacheRequest::new(CacheKeyInput::default()) + }; + assert_eq!(service.lookup(&request, Duration::ZERO).unwrap(), None); + } + + #[tokio::test] + async fn async_set_and_get_use_shared_document_helpers() { + let entry = CacheEntry { + timestamp: Some(1.0), + response: json!({"answer": "ok"}), + }; + let encoded = ResponseCacheCodec.encode(&entry).unwrap(); + let (cache, requests, calls) = cache_with_recording( + [ok(), ok(), ok(), Ok(search_hit(encoded, "0.1"))], + vec![1.0, 0.0], + 0.8, + ); + let context = semantic_context(Some(Duration::from_millis(1900))); + cache + .async_set_cache("key", entry.clone(), context.clone()) + .await + .unwrap(); + assert_eq!( + cache.async_get_cache("key", &context).await.unwrap(), + Some(entry) + ); + let text = requests_text(&requests); + assert!(text.contains("FT.CREATE")); + assert!(text.contains("HSET")); + assert!(text.contains("EXPIRE")); + assert_eq!(calls.lock().unwrap().len(), 2); + } +} diff --git a/litellm-rust/crates/cache/src/base_cache.rs b/litellm-rust/crates/cache/src/base_cache.rs index 8bd69ba5ad6..5c10e7fd5c3 100644 --- a/litellm-rust/crates/cache/src/base_cache.rs +++ b/litellm-rust/crates/cache/src/base_cache.rs @@ -32,6 +32,28 @@ impl CacheContext for ExactCacheContext { } } +#[derive(Clone, Debug, Default, PartialEq)] +pub struct SemanticCacheContext { + pub input: Option, + pub messages: Option, + pub metadata: Option, + pub scope: Option, + pub ttl: Option, +} + +impl CacheContext for SemanticCacheContext { + fn ttl(&self) -> Option { + self.ttl + } + + fn with_ttl(&self, ttl: Option) -> Self { + Self { + ttl, + ..self.clone() + } + } +} + #[derive(Clone, Copy, Debug, Deserialize, Serialize, PartialEq, Eq)] #[serde(rename_all = "lowercase")] pub enum CacheConnectionStatus { @@ -105,3 +127,31 @@ pub trait BaseCache: Send + Sync { fn test_connection(&self) -> impl Future> + Send; } + +#[cfg(test)] +mod tests { + use std::time::Duration; + + use serde_json::json; + + use super::{CacheContext, SemanticCacheContext}; + + #[test] + fn semantic_context_with_ttl_only_replaces_ttl() { + let context = SemanticCacheContext { + input: Some(json!({"input": "hello"})), + messages: Some(json!([{"role": "user", "content": "hello"}])), + metadata: Some(json!({"tenant": "team"})), + scope: Some("scope".into()), + ttl: Some(Duration::from_secs(10)), + }; + + let updated = context.with_ttl(Some(Duration::from_secs(20))); + + assert_eq!(updated.ttl, Some(Duration::from_secs(20))); + assert_eq!(updated.input, context.input); + assert_eq!(updated.messages, context.messages); + assert_eq!(updated.metadata, context.metadata); + assert_eq!(updated.scope, context.scope); + } +} diff --git a/litellm-rust/crates/cache/src/lib.rs b/litellm-rust/crates/cache/src/lib.rs index b9c720fa3ee..8364c635e3a 100644 --- a/litellm-rust/crates/cache/src/lib.rs +++ b/litellm-rust/crates/cache/src/lib.rs @@ -5,11 +5,10 @@ mod capabilities; mod codec; mod dual; mod error; -mod semantic; pub use base_cache::{ BaseCache, BatchEntry, CacheConnectionResult, CacheConnectionStatus, CacheContext, - ExactCacheContext, + ExactCacheContext, SemanticCacheContext, }; pub use cache_type::CacheType; pub use caching::{Cache, CacheBackend, get_cache, set_cache}; @@ -20,4 +19,3 @@ pub use capabilities::{ pub use codec::{CacheCodec, JsonCodec}; pub use dual::{DualCache, ReadPolicy, RemoteFailurePolicy, WritePolicy}; pub use error::Error; -pub use semantic::{SemanticCacheContext, SemanticCacheScope}; diff --git a/litellm-rust/crates/cache/src/semantic.rs b/litellm-rust/crates/cache/src/semantic.rs deleted file mode 100644 index 61f9023fa4c..00000000000 --- a/litellm-rust/crates/cache/src/semantic.rs +++ /dev/null @@ -1,72 +0,0 @@ -use std::time::Duration; - -use serde::{Deserialize, Serialize}; -use serde_json::{Map, Value}; - -use crate::CacheContext; - -#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize, PartialEq, Eq)] -#[serde(rename_all = "snake_case")] -pub enum SemanticCacheScope { - #[default] - Key, - EndUser, -} - -#[derive(Clone, Debug, Default, PartialEq)] -pub struct SemanticCacheContext { - pub input: Option, - pub messages: Vec, - pub metadata: Map, - pub scope: SemanticCacheScope, - pub ttl: Option, -} - -impl CacheContext for SemanticCacheContext { - fn ttl(&self) -> Option { - self.ttl - } - - fn with_ttl(&self, ttl: Option) -> Self { - Self { - ttl, - ..self.clone() - } - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn with_ttl_keeps_request_fields() { - let context = SemanticCacheContext { - input: Some("query".to_owned()), - messages: vec![serde_json::json!({"role": "user", "content": "hi"})], - metadata: Map::from_iter([("user".to_owned(), Value::from("u1"))]), - scope: SemanticCacheScope::EndUser, - ttl: None, - }; - - let updated = context.with_ttl(Some(Duration::from_secs(5))); - - assert_eq!(updated.ttl(), Some(Duration::from_secs(5))); - assert_eq!(updated.input, context.input); - assert_eq!(updated.messages, context.messages); - assert_eq!(updated.metadata, context.metadata); - assert_eq!(updated.scope, SemanticCacheScope::EndUser); - } - - #[test] - fn scope_serializes_like_python_cache_scope() { - assert_eq!( - serde_json::to_value(SemanticCacheScope::EndUser).unwrap(), - Value::from("end_user") - ); - assert_eq!( - serde_json::from_value::(Value::from("key")).unwrap(), - SemanticCacheScope::Key - ); - } -} diff --git a/litellm-rust/crates/cache/tests/caching.rs b/litellm-rust/crates/cache/tests/caching.rs index 9180ee9d0dc..36307ac9b33 100644 --- a/litellm-rust/crates/cache/tests/caching.rs +++ b/litellm-rust/crates/cache/tests/caching.rs @@ -1,7 +1,8 @@ use std::{sync::Mutex, time::Duration}; use litellm_cache::{ - BaseCache, CacheConnectionResult, CacheContext, Error, ExactCacheContext, get_cache, + BaseCache, CacheConnectionResult, CacheContext, Error, ExactCacheContext, SemanticCacheContext, + get_cache, }; struct TestCache { @@ -126,6 +127,24 @@ fn associated_context_preserves_backend_specific_lookup_inputs() { ); } +#[test] +fn semantic_context_with_ttl_preserves_lookup_inputs() { + let context = SemanticCacheContext { + input: Some(serde_json::json!("text")), + messages: Some(serde_json::json!([{"role": "user", "content": "hi"}])), + metadata: Some(serde_json::json!({"key": "value"})), + scope: Some("scope".into()), + ttl: None, + }; + let updated = context.with_ttl(Some(Duration::from_secs(30))); + assert_eq!(updated.ttl(), Some(Duration::from_secs(30))); + assert_eq!(updated.input, context.input); + assert_eq!(updated.messages, context.messages); + assert_eq!(updated.metadata, context.metadata); + assert_eq!(updated.scope, context.scope); + assert_eq!(context.with_ttl(None).ttl(), None); +} + #[tokio::test] async fn default_batch_operations_use_async_writes_and_stop_on_failure() { let cache = TestCache { diff --git a/litellm-rust/crates/framer/Cargo.toml b/litellm-rust/crates/framer/Cargo.toml index d22502f871a..62bfcc7da3d 100644 --- a/litellm-rust/crates/framer/Cargo.toml +++ b/litellm-rust/crates/framer/Cargo.toml @@ -11,7 +11,7 @@ aws = ["dep:aws-smithy-eventstream", "dep:aws-smithy-types"] sse = ["dep:sse-stream"] [dependencies] -aws-smithy-eventstream = { version = "=0.61.1", optional = true } +aws-smithy-eventstream = { version = "=0.61.4", optional = true } aws-smithy-types = { version = "1.6.1", optional = true } bytes = "1" futures-util.workspace = true diff --git a/litellm-rust/crates/llms/Cargo.toml b/litellm-rust/crates/llms/Cargo.toml index 0cc7af1836f..f04b78feee1 100644 --- a/litellm-rust/crates/llms/Cargo.toml +++ b/litellm-rust/crates/llms/Cargo.toml @@ -34,7 +34,7 @@ tokio = { workspace = true, features = ["sync"] } url.workspace = true [dev-dependencies] -aws-smithy-eventstream = "=0.61.1" +aws-smithy-eventstream = "=0.61.4" aws-smithy-types = "1.6.1" rstest.workspace = true tokio.workspace = true diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index bace8d61f58..6b4476d897c 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -24,13 +24,17 @@ litellm-cache.workspace = true litellm-cache-azure-blob.workspace = true litellm-cache-memory.workspace = true litellm-cache-redis.workspace = true +litellm-cache-s3.workspace = true litellm-cache-gcs.workspace = true litellm-cache-disk.workspace = true +litellm-cache-redis-semantic.workspace = true litellm-cache-response.workspace = true litellm-cache-qdrant-semantic.workspace = true qdrant-client.workspace = true +litellm-cache-valkey-semantic = { path = "../cache-valkey-semantic" } serde.workspace = true litellm-auth.workspace = true +litellm-auth-aws.workspace = true litellm-callbacks-legacy-python.workspace = true litellm-core.workspace = true litellm-core-utils.workspace = true @@ -43,9 +47,10 @@ litellm-token-counter = { path = "../token-counter", default-features = false } pyo3.workspace = true pyo3-async-runtimes.workspace = true reqwest.workspace = true +redis = { version = "1.7.0", features = ["tls-rustls"] } serde_json.workspace = true url.workspace = true -tokio = { workspace = true, features = ["sync"] } +tokio = { workspace = true, features = ["rt", "sync"] } [dev-dependencies] serde.workspace = true @@ -53,6 +58,7 @@ serde_with.workspace = true criterion.workspace = true futures-util.workspace = true rstest.workspace = true +sha2.workspace = true tokio-tungstenite.workspace = true [[bench]] diff --git a/litellm-rust/crates/python-bridge/src/cache/binding.rs b/litellm-rust/crates/python-bridge/src/cache/binding.rs index ad64b24d3c1..2ff73238202 100644 --- a/litellm-rust/crates/python-bridge/src/cache/binding.rs +++ b/litellm-rust/crates/python-bridge/src/cache/binding.rs @@ -56,12 +56,7 @@ impl ResolvedCache { CacheBinding::Disabled => ready_none(py)?, CacheBinding::Native(service) => { let request = request(input)?; - let service = service.clone(); - run_async( - py, - async move { service.async_lookup(&request, now()).await }, - cache_error, - )? + service.async_lookup_py(py, request)? } CacheBinding::PythonCallback(callback) => callback.async_lookup(py, kwargs)?, }; @@ -179,12 +174,7 @@ impl ResolvedCache { CacheBinding::Native(service) => { let request = self::request(request)?; let response: Value = from_py(response)?; - let service = service.clone(); - run_async( - py, - async move { service.async_store(&request, response, now()).await }, - cache_error, - ) + service.async_store_py(py, request, response) } CacheBinding::PythonCallback(callback) => { callback.async_store(py, response, callback_kwargs) @@ -241,12 +231,7 @@ impl ResolvedCache { )); } let entries = requests.into_iter().zip(responses).collect(); - let service = service.clone(); - run_async( - py, - async move { service.async_store_batch(entries, now()).await }, - cache_error, - ) + service.async_store_batch_py(py, entries) } CacheBinding::PythonCallback(callback) => { callback.async_store_batch(py, callback_result, callback_kwargs) diff --git a/litellm-rust/crates/python-bridge/src/cache/config.rs b/litellm-rust/crates/python-bridge/src/cache/config.rs index 8c26a6033c0..d161b401928 100644 --- a/litellm-rust/crates/python-bridge/src/cache/config.rs +++ b/litellm-rust/crates/python-bridge/src/cache/config.rs @@ -1,12 +1,14 @@ -use std::{env, path::PathBuf, time::Duration}; +use std::{path::PathBuf, time::Duration}; +use litellm_auth_aws::AwsAuthConfig; use litellm_cache::CacheType; use litellm_cache_qdrant_semantic::{OpenAiEmbedderConfig, QdrantSemanticConfig, Quantization}; use litellm_cache_redis::{RedisNode, RedisTopology}; +use litellm_cache_s3::{S3CacheConfig, S3Endpoint}; use pyo3::{ - exceptions::{PyTypeError, PyValueError}, + exceptions::{PyAttributeError, PyTypeError, PyValueError}, prelude::*, - types::{PyAny, PyDict, PyList, PyString}, + types::{PyAny, PyBool, PyDict, PyList, PyString}, }; use super::{native::NativeResponseCache, request::duration}; @@ -92,6 +94,19 @@ pub(super) struct AzureBlobCacheConfig { pub(super) container: String, } +#[allow( + dead_code, + reason = "embedding settings are projected so drift falls back to Python" +)] +pub(super) struct RedisSemanticCacheConfig { + pub(super) redis_url: String, + pub(super) index_name: String, + pub(super) similarity_threshold: f64, + pub(super) embedding_model: String, + pub(super) embedding_max_input_tokens: Option, + pub(super) embedding_timeout: Option, +} + struct RedisClientProjection<'py> { topology: RedisTopology, host: String, @@ -103,6 +118,14 @@ struct RedisClientProjection<'py> { const REDIS_PY_DEFAULT_MAX_CONNECTIONS: usize = 1 << 31; +#[allow(dead_code, reason = "consumed by the cache activation follow-up")] +pub(super) struct ValkeySemanticCacheConfig { + pub(super) similarity_threshold: f64, + pub(super) index_name: String, + pub(super) embedding_model: String, + pub(super) connection: RedisConnectionConfig, +} + pub(super) struct QdrantSemanticCacheConfig { pub(super) grpc_url: String, pub(super) api_key: Option, @@ -123,13 +146,17 @@ impl QdrantSemanticCacheConfig { } } } + pub(super) enum CacheBackendConfig { Memory(MemoryCacheConfig), Redis(Box), + S3(Box), Gcs(GcsCacheConfig), + ValkeySemantic(Box), Disk(DiskCacheConfig), - QdrantSemantic(Box), AzureBlob(AzureBlobCacheConfig), + RedisSemantic(Box), + QdrantSemantic(Box), } #[allow(dead_code, reason = "consumed by the cache activation follow-up")] @@ -144,6 +171,9 @@ pub(super) enum UnsupportedCacheConfig { RedisCredentials, RedisConnection, RedisOption, + S3Client, + S3Credentials, + S3Option, GcsBucket, DiskStore, QdrantEndpoint, @@ -158,6 +188,9 @@ impl UnsupportedCacheConfig { Self::RedisCredentials => "native Redis credentials require Python", Self::RedisConnection => "native Redis connection type is not implemented", Self::RedisOption => "native Redis configuration requires Python", + Self::S3Client => "native S3 client type is not implemented", + Self::S3Credentials => "native S3 credentials require Python", + Self::S3Option => "native S3 configuration requires Python", Self::GcsBucket => "native GCS cache requires a configured bucket name", Self::DiskStore => "native disk cache requires the built-in diskcache store", Self::QdrantEndpoint => { @@ -206,6 +239,13 @@ impl NativeCacheConfig { }))), Err(reason) => Ok(CacheConfigProjection::Unsupported(reason)), }, + Some(CacheType::S3) => match project_s3(&backend)? { + Ok(backend) => Ok(CacheConfigProjection::Native(Box::new(Self { + policy, + backend: CacheBackendConfig::S3(Box::new(backend)), + }))), + Err(reason) => Ok(CacheConfigProjection::Unsupported(reason)), + }, Some(CacheType::Gcs) => match project_gcs(&backend)? { Ok(backend) => Ok(CacheConfigProjection::Native(Box::new(Self { policy, @@ -213,6 +253,13 @@ impl NativeCacheConfig { }))), Err(reason) => Ok(CacheConfigProjection::Unsupported(reason)), }, + Some(CacheType::ValkeySemantic) => match project_valkey_semantic(&backend)? { + Ok(backend) => Ok(CacheConfigProjection::Native(Box::new(Self { + policy, + backend: CacheBackendConfig::ValkeySemantic(Box::new(backend)), + }))), + Err(reason) => Ok(CacheConfigProjection::Unsupported(reason)), + }, Some(CacheType::Disk) => match project_disk(&backend)? { Ok(backend) => Ok(CacheConfigProjection::Native(Box::new(Self { policy, @@ -233,16 +280,81 @@ impl NativeCacheConfig { backend: CacheBackendConfig::AzureBlob(backend), })) }), - Some(CacheType::RedisSemantic | CacheType::ValkeySemantic | CacheType::S3) | None => { - Ok(CacheConfigProjection::Unsupported( - UnsupportedCacheConfig::Backend, - )) - } + Some(CacheType::RedisSemantic) => project_redis_semantic(&backend).map(|backend| { + CacheConfigProjection::Native(Box::new(Self { + policy, + backend: CacheBackendConfig::RedisSemantic(Box::new(backend)), + })) + }), + None => Ok(CacheConfigProjection::Unsupported( + UnsupportedCacheConfig::Backend, + )), } } pub(super) fn service_mismatch(&self, service: &NativeResponseCache) -> Option<&'static str> { + let default_ttl = match &self.backend { + CacheBackendConfig::Memory(config) => Some(config.default_ttl), + CacheBackendConfig::Redis(config) => Some(config.default_ttl), + CacheBackendConfig::S3(_) => None, + CacheBackendConfig::ValkeySemantic(_) => Some(Duration::ZERO), + CacheBackendConfig::Disk(_) + | CacheBackendConfig::AzureBlob(_) + | CacheBackendConfig::Gcs(_) + | CacheBackendConfig::RedisSemantic(_) + | CacheBackendConfig::QdrantSemantic(_) => None, + }; + if !matches!(self.backend, CacheBackendConfig::ValkeySemantic(_)) + && service.default_ttl() != default_ttl + { + return Some("facade and native backend default TTLs must match"); + } match &self.backend { + CacheBackendConfig::Memory(config) if service.kind() != "memory" => { + Some("facade and native backend types must match") + } + CacheBackendConfig::Memory(config) if service.capacity() != Some(config.capacity) => { + Some("facade and native backend capacities must match") + } + CacheBackendConfig::Memory(config) + if service.max_entry_bytes() != Some(config.max_entry_bytes) => + { + Some("facade and native backend item limits must match") + } + CacheBackendConfig::Memory(_) => None, + CacheBackendConfig::Redis(_) if service.kind() != "redis" => { + Some("facade and native backend types must match") + } + CacheBackendConfig::Redis(config) if service.topology() != Some(&config.topology) => { + Some("facade and native backend topologies must match") + } + CacheBackendConfig::Redis(config) => (service.namespace() + != config.namespace.as_deref()) + .then_some("facade and native backend namespaces must match"), + CacheBackendConfig::S3(_) if service.kind() != "s3" => { + Some("facade and native backend types must match") + } + CacheBackendConfig::S3(config) if service.bucket() != Some(config.bucket.as_str()) => { + Some("facade and native backend buckets must match") + } + CacheBackendConfig::S3(config) + if service.key_prefix() != Some(config.key_prefix.as_str()) => + { + Some("facade and native backend key prefixes must match") + } + CacheBackendConfig::S3(config) if service.region() != Some(config.region.as_str()) => { + Some("facade and native backend regions must match") + } + CacheBackendConfig::S3(config) + if service.endpoint() + != config + .endpoint + .as_ref() + .map(|endpoint| endpoint.url.as_str()) => + { + Some("facade and native backend endpoints must match") + } + CacheBackendConfig::S3(_) => None, CacheBackendConfig::Gcs(_) if service.kind() != "gcs" => { Some("facade and native backend types must match") } @@ -268,36 +380,16 @@ impl NativeCacheConfig { Some("facade and native backend credentials must match") } CacheBackendConfig::Gcs(_) => None, - CacheBackendConfig::Memory(_) if service.kind() != "memory" => { - Some("facade and native backend types must match") + CacheBackendConfig::ValkeySemantic(config) => { + if service.kind() != "valkey-semantic" { + return Some("facade and native backend types must match"); + } + let Some((threshold, index_name)) = service.semantic_config() else { + return Some("facade and native backend types must match"); + }; + (threshold != config.similarity_threshold || index_name != config.index_name) + .then_some("facade and native semantic settings must match") } - CacheBackendConfig::Memory(config) - if service.default_ttl() != Some(config.default_ttl) => - { - Some("facade and native backend default TTLs must match") - } - CacheBackendConfig::Memory(config) if service.capacity() != Some(config.capacity) => { - Some("facade and native backend capacities must match") - } - CacheBackendConfig::Memory(config) - if service.max_entry_bytes() != Some(config.max_entry_bytes) => - { - Some("facade and native backend item limits must match") - } - CacheBackendConfig::Memory(_) => None, - CacheBackendConfig::Redis(config) if service.kind() != "redis" => { - Some("facade and native backend types must match") - } - CacheBackendConfig::Redis(config) if service.topology() != Some(&config.topology) => { - Some("facade and native backend topologies must match") - } - CacheBackendConfig::Redis(config) => (service.namespace() - != config.namespace.as_deref()) - .then_some("facade and native backend namespaces must match") - .or_else(|| { - (service.default_ttl() != Some(config.default_ttl)) - .then_some("facade and native backend default TTLs must match") - }), CacheBackendConfig::Disk(_) if service.kind() != "disk" => { Some("facade and native backend types must match") } @@ -309,6 +401,20 @@ impl NativeCacheConfig { let facade = std::fs::canonicalize(&config.directory).ok(); (native != facade).then_some("facade and native backend directories must match") } + CacheBackendConfig::RedisSemantic(_) if service.kind() != "redis_semantic" => { + Some("facade and native backend types must match") + } + CacheBackendConfig::RedisSemantic(config) + if service.index_name() != Some(config.index_name.as_str()) => + { + Some("facade and native backend index names must match") + } + CacheBackendConfig::RedisSemantic(config) + if service.similarity_threshold() != Some(config.similarity_threshold) => + { + Some("facade and native backend similarity thresholds must match") + } + CacheBackendConfig::RedisSemantic(_) => None, CacheBackendConfig::QdrantSemantic(config) if service.kind() != "qdrant_semantic" => { Some("facade and native backend types must match") } @@ -340,9 +446,6 @@ impl NativeCacheConfig { { Some("facade and native backend containers must match") } - Some(_) if service.default_ttl().is_some() => { - Some("facade and native backend default TTLs must match") - } Some(_) => None, }, } @@ -359,7 +462,7 @@ fn project_qdrant_semantic( Err(_) => return Ok(Err(UnsupportedCacheConfig::QdrantEndpoint)), }; if !matches!(parsed.scheme(), "http" | "https") - || !parsed.path().is_empty() && parsed.path() != "/" + || (!parsed.path().is_empty() && parsed.path() != "/") || parsed.query().is_some() || parsed.host_str().is_none() || parsed.port() != Some(6333) @@ -373,8 +476,9 @@ fn project_qdrant_semantic( grpc_url.set_path(""); grpc_url.set_query(None); - let embedding_max_input_tokens = optional_attribute_i64(backend, "embedding_max_input_tokens")?; - if embedding_max_input_tokens.is_some() { + if optional_attribute(backend, "embedding_max_input_tokens")? + .is_some_and(|value| !value.is_none()) + { return Ok(Err(UnsupportedCacheConfig::SemanticEmbedding)); } let configured_model = backend.getattr("embedding_model")?.extract::()?; @@ -404,16 +508,19 @@ fn project_qdrant_semantic( return Ok(Err(UnsupportedCacheConfig::SemanticEmbedding)); } } - let Ok(embedding_api_key) = env::var("OPENAI_API_KEY") else { + let Ok(embedding_api_key) = std::env::var("OPENAI_API_KEY") else { return Ok(Err(UnsupportedCacheConfig::SemanticEmbedding)); }; if embedding_api_key.is_empty() { return Ok(Err(UnsupportedCacheConfig::SemanticEmbedding)); } - let embedding_api_base = env::var("OPENAI_BASE_URL") - .or_else(|_| env::var("OPENAI_API_BASE")) + let embedding_api_base = std::env::var("OPENAI_BASE_URL") + .or_else(|_| std::env::var("OPENAI_API_BASE")) .unwrap_or_else(|_| "https://api.openai.com/v1".to_owned()); - let timeout = optional_attribute_f64(backend, "embedding_timeout")? + let timeout = optional_attribute(backend, "embedding_timeout")? + .map(|value| value.extract::>()) + .transpose()? + .flatten() .map(duration) .transpose()?; Ok(Ok(QdrantSemanticCacheConfig { @@ -459,6 +566,39 @@ fn project_azure_blob(backend: &Bound<'_, PyAny>) -> PyResult, +) -> PyResult { + Ok(RedisSemanticCacheConfig { + redis_url: backend.getattr("_redis_url")?.extract::()?, + index_name: backend + .getattr("_index_name")? + .extract::>()? + .unwrap_or_else(|| "litellm_semantic_cache_index".into()), + similarity_threshold: backend.getattr("similarity_threshold")?.extract::()?, + embedding_model: backend.getattr("embedding_model")?.extract::()?, + embedding_max_input_tokens: backend + .getattr("embedding_max_input_tokens")? + .extract::>()?, + embedding_timeout: backend + .getattr("embedding_timeout")? + .extract::>()?, + }) +} + +#[inline(never)] +fn project_memory(backend: &Bound<'_, PyAny>) -> PyResult { + let max_size_kib = backend.getattr("max_size_per_item")?.extract::()?; + Ok(MemoryCacheConfig { + default_ttl: duration(backend.getattr("default_ttl")?.extract::()?)?, + capacity: backend.getattr("max_size_in_memory")?.extract::()?, + max_entry_bytes: max_size_kib + .checked_mul(1024) + .ok_or_else(|| PyValueError::new_err("memory cache item limit is too large"))?, + }) +} + #[inline(never)] fn project_gcs( backend: &Bound<'_, PyAny>, @@ -476,18 +616,6 @@ fn project_gcs( })) } -#[inline(never)] -fn project_memory(backend: &Bound<'_, PyAny>) -> PyResult { - let max_size_kib = backend.getattr("max_size_per_item")?.extract::()?; - Ok(MemoryCacheConfig { - default_ttl: duration(backend.getattr("default_ttl")?.extract::()?)?, - capacity: backend.getattr("max_size_in_memory")?.extract::()?, - max_entry_bytes: max_size_kib - .checked_mul(1024) - .ok_or_else(|| PyValueError::new_err("memory cache item limit is too large"))?, - }) -} - #[inline(never)] fn project_disk( backend: &Bound<'_, PyAny>, @@ -597,6 +725,77 @@ fn project_redis( })) } +#[inline(never)] +fn project_s3( + backend: &Bound<'_, PyAny>, +) -> PyResult> { + let client = backend.getattr("s3_client")?; + if !instance_class_is(&client, "botocore.client", "S3")? { + return Ok(Err(UnsupportedCacheConfig::S3Client)); + } + let meta = client.getattr("meta")?; + let Some(region) = optional_string(meta.getattr("region_name")?)? else { + return Ok(Err(UnsupportedCacheConfig::S3Option)); + }; + let Some(endpoint_url) = optional_string(meta.getattr("endpoint_url")?)? else { + return Ok(Err(UnsupportedCacheConfig::S3Option)); + }; + let client_config = meta.getattr("config")?; + for name in ["s3", "proxies", "client_cert"] { + if optional_attribute(&client_config, name)?.is_some_and(|value| !value.is_none()) { + return Ok(Err(UnsupportedCacheConfig::S3Option)); + } + } + let signature = match optional_attribute(&client_config, "signature_version")? { + Some(value) => value.extract::>()?, + None => None, + }; + if signature.as_deref() != Some("s3v4") { + return Ok(Err(UnsupportedCacheConfig::S3Option)); + } + let insecure = endpoint_url.starts_with("http://"); + let verify = optional_attribute_chain(&client, &["_endpoint", "http_session", "_verify"])?; + let verified = verify + .and_then(|value| value.cast::().ok().map(|value| value.is_true())) + .unwrap_or(false); + if !verified && !insecure { + return Ok(Err(UnsupportedCacheConfig::S3Option)); + } + let credentials = optional_attribute_chain(&client, &["_request_signer", "_credentials"])? + .ok_or(UnsupportedCacheConfig::S3Credentials); + let credentials = match credentials { + Ok(credentials) if !credentials.is_none() => credentials, + _ => return Ok(Err(UnsupportedCacheConfig::S3Credentials)), + }; + let auth = if credentials.getattr("method")?.extract::()?.as_str() == "explicit" { + AwsAuthConfig { + access_key_id: credentials + .getattr("access_key")? + .extract::>()?, + secret_access_key: credentials + .getattr("secret_key")? + .extract::>()?, + session_token: credentials.getattr("token")?.extract::>()?, + region_name: Some(region.clone()), + ..Default::default() + } + } else { + AwsAuthConfig { + region_name: Some(region.clone()), + ..Default::default() + } + }; + let default_endpoint = endpoint_url == format!("https://s3.{region}.amazonaws.com") + || (region == "us-east-1" && endpoint_url == "https://s3.amazonaws.com"); + Ok(Ok(S3CacheConfig { + bucket: backend.getattr("bucket_name")?.extract::()?, + key_prefix: backend.getattr("key_prefix")?.extract::()?, + region, + endpoint: (!default_endpoint).then_some(S3Endpoint { url: endpoint_url }), + auth, + })) +} + #[inline(never)] fn project_standalone_client<'py>( client: &Bound<'py, PyAny>, @@ -631,7 +830,7 @@ fn project_standalone_client<'py>( #[inline(never)] fn project_cluster_client<'py>( - source: &Bound<'py, PyDict>, + source: &Bound<'_, PyDict>, client: &Bound<'py, PyAny>, ) -> PyResult, UnsupportedCacheConfig>> { let Some(startup_nodes) = startup_nodes(source)? else { @@ -719,6 +918,71 @@ fn port(value: i64) -> PyResult { u16::try_from(value).map_err(|_| PyValueError::new_err("invalid Redis port")) } +#[inline(never)] +fn project_valkey_semantic( + backend: &Bound<'_, PyAny>, +) -> PyResult> { + let client = backend.getattr("sync_client")?; + let pool = client.getattr("connection_pool")?; + let Ok((resolved, is_tls)) = project_connection_pool(&pool)? else { + return Ok(Err(UnsupportedCacheConfig::RedisConnection)); + }; + for key in ["credential_provider", "redis_connect_func"] { + if has_value(&resolved, key)? { + return Ok(Err(UnsupportedCacheConfig::RedisCredentials)); + } + } + if is_tls { + return Ok(Err(UnsupportedCacheConfig::RedisConnection)); + } + let connection = RedisConnectionConfig { + host: required_string(&resolved, "host")?, + port: u16::try_from(required_i64(&resolved, "port")?) + .map_err(|_| PyValueError::new_err("invalid Redis port"))?, + database: optional_i64(&resolved, "db")?.unwrap_or(0), + username: optional_dict_string(&resolved, "username")?, + password: optional_dict_string(&resolved, "password")?, + protocol: RedisProtocol::Resp2, + pool_size: pool.getattr("max_connections")?.extract::()?, + read_timeout: None, + connect_timeout: None, + socket_keepalive: None, + health_check_interval: Duration::ZERO, + client_name: None, + tls: None, + }; + if connection.host.is_empty() { + return Ok(Err(UnsupportedCacheConfig::RedisConnection)); + } + Ok(Ok(ValkeySemanticCacheConfig { + similarity_threshold: backend.getattr("similarity_threshold")?.extract()?, + index_name: backend.getattr("index_name")?.extract()?, + embedding_model: backend.getattr("embedding_model")?.extract()?, + connection, + })) +} + +#[inline(never)] +fn project_connection_pool<'py>( + pool: &Bound<'py, PyAny>, +) -> PyResult, bool), UnsupportedCacheConfig>> { + if !instance_class_is(pool, "redis.connection", "ConnectionPool")? { + return Ok(Err(UnsupportedCacheConfig::RedisConnection)); + } + let resolved = pool.getattr("connection_kwargs")?.cast_into::()?; + let connection_class = resolved + .get_item("connection_class")? + .unwrap_or(pool.getattr("connection_class")?); + let is_tls = if class_is(&connection_class, "redis.connection", "Connection")? { + false + } else if class_is(&connection_class, "redis.connection", "SSLConnection")? { + true + } else { + return Ok(Err(UnsupportedCacheConfig::RedisConnection)); + }; + Ok(Ok((resolved, is_tls))) +} + #[inline(never)] fn project_tls(values: &Bound<'_, PyDict>) -> PyResult { Ok(RedisTlsConfig { @@ -801,25 +1065,28 @@ fn optional_attribute_string(value: &Bound<'_, PyAny>, name: &str) -> PyResult, name: &str) -> PyResult> { +fn optional_attribute<'py>( + value: &Bound<'py, PyAny>, + name: &str, +) -> PyResult>> { match value.getattr(name) { - Ok(attribute) => attribute.extract::>(), - Err(error) if error.is_instance_of::(value.py()) => { - Ok(None) - } + Ok(value) => Ok(Some(value)), + Err(error) if error.is_instance_of::(value.py()) => Ok(None), Err(error) => Err(error), } } #[inline(never)] -fn optional_attribute_f64(value: &Bound<'_, PyAny>, name: &str) -> PyResult> { - match value.getattr(name) { - Ok(attribute) => attribute.extract::>(), - Err(error) if error.is_instance_of::(value.py()) => { - Ok(None) - } - Err(error) => Err(error), - } +fn optional_attribute_chain<'py>( + value: &Bound<'py, PyAny>, + names: &[&str], +) -> PyResult>> { + names + .iter() + .try_fold(Some(value.clone()), |current, name| match current { + Some(current) => optional_attribute(¤t, name), + None => Ok(None), + }) } #[inline(never)] @@ -905,11 +1172,6 @@ fn optional_dict_duration(values: &Bound<'_, PyDict>, key: &str) -> PyResult &'static Mutex<()> { - static LOCK: OnceLock> = OnceLock::new(); - LOCK.get_or_init(|| Mutex::new(())) - } - - fn qdrant_facade<'py>(py: Python<'py>, extra: &str) -> Bound<'py, PyAny> { - install_fake_litellm(py); - facade( - py, - &format!( - "backend = SimpleNamespace(qdrant_api_base='https://qdrant.example:6333', qdrant_api_key='qdrant-key', collection_name='cache', similarity_threshold=0.99, embedding_model='openai/text-embedding-3-small', vector_size=8, embedding_max_input_tokens=None, embedding_timeout=None)\n\ - facade = SimpleNamespace(type='qdrant-semantic', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)\n\ - {extra}" - ), - ) - } - - fn install_fake_litellm(py: Python<'_>) { - py.run( - c" -import sys -import types -litellm = types.ModuleType('litellm') -litellm.api_key = None -litellm.openai_key = None -litellm.api_base = None -litellm.__path__ = [] -caching = types.ModuleType('litellm.caching') -caching.__path__ = [] -embedding_router = types.ModuleType('litellm.caching._embedding_router') -embedding_router.resolve_embedding_router = lambda *_args: None -caching._embedding_router = embedding_router -litellm.caching = caching -sys.modules['litellm'] = litellm -sys.modules['litellm.caching'] = caching -sys.modules['litellm.caching._embedding_router'] = embedding_router -", - None, - None, - ) - .unwrap(); - } - - fn configure_embedding_environment<'py>( - py: Python<'py>, - key: Option<&str>, - ) -> PyResult> { - let environ = py.import("os")?.getattr("environ")?; - let prior = environ.call_method1("get", ("OPENAI_API_KEY",))?; - match key { - Some(key) => environ.set_item("OPENAI_API_KEY", key)?, - None => { - environ.call_method1("pop", ("OPENAI_API_KEY", py.None()))?; - } - } - Ok(prior) - } - - fn restore_embedding_environment(py: Python<'_>, prior: Bound<'_, PyAny>) -> PyResult<()> { - let environ = py.import("os")?.getattr("environ")?; - if prior.is_none() { - environ.call_method1("pop", ("OPENAI_API_KEY", py.None()))?; - } else { - environ.set_item("OPENAI_API_KEY", prior)?; - } - Ok(()) - } - fn facade<'py>(py: Python<'py>, body: &str) -> Bound<'py, PyAny> { let locals = PyDict::new(py); py.run( @@ -1104,6 +1298,87 @@ sys.modules['litellm.caching._embedding_router'] = embedding_router }); } + #[test] + fn projects_valkey_semantic_configuration() { + Python::initialize(); + Python::attach(|py| { + let facade = facade( + py, + "pool = ConnectionPool()\n\ + pool.connection_class = Connection\n\ + pool.max_connections = 12\n\ + pool.connection_kwargs = {'host': 'cache.internal', 'port': 6390, 'db': 2}\n\ + client = SimpleNamespace(connection_pool=pool)\n\ + backend = SimpleNamespace(similarity_threshold=0.85, index_name='semantic_idx', embedding_model='text-embedding-3-small', sync_client=client)\n\ + facade = SimpleNamespace(type='valkey-semantic', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)", + ); + let CacheConfigProjection::Native(config) = + NativeCacheConfig::project(&facade).unwrap() + else { + panic!("Valkey semantic cache should be supported"); + }; + let CacheBackendConfig::ValkeySemantic(valkey) = config.backend else { + panic!("expected Valkey semantic configuration"); + }; + assert_eq!(valkey.similarity_threshold, 0.85); + assert_eq!(valkey.index_name, "semantic_idx"); + assert_eq!(valkey.embedding_model, "text-embedding-3-small"); + assert_eq!(valkey.connection.host, "cache.internal"); + assert_eq!(valkey.connection.port, 6390); + assert_eq!(valkey.connection.database, 2); + assert_eq!(valkey.connection.pool_size, 12); + assert_eq!(valkey.connection.protocol, RedisProtocol::Resp2); + assert!(valkey.connection.tls.is_none()); + }); + } + + #[test] + fn valkey_semantic_tls_stays_on_python() { + Python::initialize(); + Python::attach(|py| { + let facade = facade( + py, + "pool = ConnectionPool()\n\ + pool.connection_class = SSLConnection\n\ + pool.connection_kwargs = {'host': 'cache.internal', 'port': 6390}\n\ + client = SimpleNamespace(connection_pool=pool)\n\ + backend = SimpleNamespace(similarity_threshold=0.85, index_name='semantic_idx', embedding_model='text-embedding-3-small', sync_client=client)\n\ + facade = SimpleNamespace(type='valkey-semantic', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)", + ); + let CacheConfigProjection::Unsupported(reason) = + NativeCacheConfig::project(&facade).unwrap() + else { + panic!("TLS Valkey semantic cache should stay on Python"); + }; + assert_eq!( + reason.message(), + "native Redis connection type is not implemented" + ); + }); + } + + #[test] + fn valkey_semantic_dynamic_auth_stays_on_python() { + Python::initialize(); + Python::attach(|py| { + let facade = facade( + py, + "pool = ConnectionPool()\n\ + pool.connection_class = Connection\n\ + pool.connection_kwargs = {'host': 'cache.internal', 'port': 6390, 'credential_provider': object()}\n\ + client = SimpleNamespace(connection_pool=pool)\n\ + backend = SimpleNamespace(similarity_threshold=0.85, index_name='semantic_idx', embedding_model='text-embedding-3-small', sync_client=client)\n\ + facade = SimpleNamespace(type='valkey-semantic', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)", + ); + let CacheConfigProjection::Unsupported(reason) = + NativeCacheConfig::project(&facade).unwrap() + else { + panic!("dynamic Valkey authentication must stay on Python"); + }; + assert_eq!(reason.message(), "native Redis credentials require Python"); + }); + } + #[test] fn dynamic_redis_auth_stays_on_python() { Python::initialize(); @@ -1122,136 +1397,6 @@ sys.modules['litellm.caching._embedding_router'] = embedding_router }); } - #[test] - fn projects_gcs_configuration() { - Python::initialize(); - Python::attach(|py| { - let facade = facade( - py, - "backend = SimpleNamespace(bucket_name='bucket', key_prefix='cache/', path_service_account='credentials.json')\n\ - facade = SimpleNamespace(type='gcs', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)", - ); - let CacheConfigProjection::Native(config) = - NativeCacheConfig::project(&facade).unwrap() - else { - panic!("GCS cache should be supported"); - }; - let CacheBackendConfig::Gcs(gcs) = config.backend else { - panic!("expected GCS configuration"); - }; - assert_eq!( - gcs, - GcsCacheConfig { - bucket_name: "bucket".into(), - key_prefix: "cache/".into(), - path_service_account: Some("credentials.json".into()), - } - ); - }); - } - - #[test] - fn rejects_gcs_without_a_bucket_name() { - Python::initialize(); - Python::attach(|py| { - let facade = facade( - py, - "backend = SimpleNamespace(bucket_name=None, key_prefix='', path_service_account=None)\n\ - facade = SimpleNamespace(type='gcs', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)", - ); - let CacheConfigProjection::Unsupported(reason) = - NativeCacheConfig::project(&facade).unwrap() - else { - panic!("GCS cache without a bucket should be unsupported"); - }; - assert!(matches!(&reason, UnsupportedCacheConfig::GcsBucket)); - assert_eq!( - reason.message(), - "native GCS cache requires a configured bucket name" - ); - }); - } - - #[test] - fn projects_builtin_disk_configuration_and_rejects_custom_stores() { - Python::initialize(); - Python::attach(|py| { - let root = - std::env::temp_dir().join(format!("litellm-disk-config-{}", std::process::id())); - let directory = root.to_string_lossy(); - let disk_facade = facade( - py, - &format!( - "Cache = type('Cache', (), {{'__module__': 'diskcache.core'}})\n\ - Disk = type('Disk', (), {{'__module__': 'diskcache.core'}})\n\ - store = Cache()\n\ - store._disk = Disk()\n\ - store.directory = {directory:?}\n\ - backend = SimpleNamespace(disk_cache=store)\n\ - facade = SimpleNamespace(type='disk', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)" - ), - ); - let CacheConfigProjection::Native(config) = - NativeCacheConfig::project(&disk_facade).unwrap() - else { - panic!("disk cache should be supported"); - }; - let CacheBackendConfig::Disk(disk) = config.backend else { - panic!("expected disk configuration"); - }; - assert_eq!(disk.directory, root); - let matching = NativeResponseCache::disk(&directory).unwrap(); - assert_eq!( - (NativeCacheConfig { - policy: config.policy, - backend: CacheBackendConfig::Disk(disk), - }) - .service_mismatch(&matching), - None - ); - let other = NativeResponseCache::disk(&root.join("other").to_string_lossy()).unwrap(); - let mismatch = NativeCacheConfig { - policy: CachePolicy { - mode: "default-on".into(), - ttl: None, - namespace: None, - supported_call_types: None, - redis_flush_size: None, - semantic_cache_scope: "key".into(), - }, - backend: CacheBackendConfig::Disk(DiskCacheConfig { - directory: root.clone(), - }), - }; - assert_eq!( - mismatch.service_mismatch(&other), - Some("facade and native backend directories must match") - ); - - let custom = facade( - py, - &format!( - "CustomCache = type('CustomCache', (), {{'__module__': 'mypkg'}})\n\ - CustomDisk = type('CustomDisk', (), {{'__module__': 'mypkg'}})\n\ - store = CustomCache()\n\ - store._disk = CustomDisk()\n\ - store.directory = {directory:?}\n\ - backend = SimpleNamespace(disk_cache=store)\n\ - facade = SimpleNamespace(type='disk', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)" - ), - ); - let CacheConfigProjection::Unsupported(reason) = - NativeCacheConfig::project(&custom).unwrap() - else { - panic!("custom disk store must stay on Python"); - }; - assert_eq!( - reason.message(), - "native disk cache requires the built-in diskcache store" - ); - }); - } - #[test] fn projects_cluster_startup_nodes_as_redis_topology() { Python::initialize(); @@ -1298,6 +1443,71 @@ sys.modules['litellm.caching._embedding_router'] = embedding_router }); } + #[test] + fn projects_gcs_configuration() { + Python::initialize(); + Python::attach(|py| { + let facade = facade( + py, + "backend = SimpleNamespace(bucket_name='bucket', key_prefix='cache/', path_service_account='credentials.json')\n\ + facade = SimpleNamespace(type='gcs', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)", + ); + let CacheConfigProjection::Native(config) = + NativeCacheConfig::project(&facade).unwrap() + else { + panic!("GCS cache should be supported"); + }; + let CacheBackendConfig::Gcs(gcs) = config.backend else { + panic!("expected GCS configuration"); + }; + assert_eq!( + gcs, + GcsCacheConfig { + bucket_name: "bucket".into(), + key_prefix: "cache/".into(), + path_service_account: Some("credentials.json".into()), + } + ); + let matching = NativeResponseCache::gcs( + litellm_cache_gcs::GcsConfig { + bucket_name: "bucket".into(), + gcs_path: Some("cache/".into()), + path_service_account: Some("credentials.json".into()), + endpoint: litellm_cache_gcs::DEFAULT_ENDPOINT.into(), + }, + Some("token".into()), + ) + .unwrap(); + let matching_config = NativeCacheConfig { + policy: config.policy, + backend: CacheBackendConfig::Gcs(gcs), + }; + assert_eq!(matching_config.service_mismatch(&matching), None); + }); + } + + #[test] + fn rejects_gcs_without_a_bucket_name() { + Python::initialize(); + Python::attach(|py| { + let facade = facade( + py, + "backend = SimpleNamespace(bucket_name=None, key_prefix='', path_service_account=None)\n\ + facade = SimpleNamespace(type='gcs', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)", + ); + let CacheConfigProjection::Unsupported(reason) = + NativeCacheConfig::project(&facade).unwrap() + else { + panic!("GCS cache without a bucket should be unsupported"); + }; + assert!(matches!(&reason, UnsupportedCacheConfig::GcsBucket)); + assert_eq!( + reason.message(), + "native GCS cache requires a configured bucket name" + ); + }); + } + #[test] fn malformed_startup_nodes_and_foreign_connect_hooks_stay_on_python() { Python::initialize(); @@ -1334,175 +1544,4 @@ sys.modules['litellm.caching._embedding_router'] = embedding_router } }); } - #[test] - fn projects_qdrant_configuration_from_python() { - let _guard = env_lock().lock().unwrap_or_else(|error| error.into_inner()); - Python::initialize(); - Python::attach(|py| { - let prior = configure_embedding_environment(py, Some("embedding-key")).unwrap(); - let facade = qdrant_facade(py, ""); - let CacheConfigProjection::Native(config) = - NativeCacheConfig::project(&facade).unwrap() - else { - panic!("Qdrant cache should be supported"); - }; - let CacheBackendConfig::QdrantSemantic(config) = config.backend else { - panic!("expected Qdrant configuration"); - }; - assert_eq!(config.grpc_url, "https://qdrant.example:6334"); - assert_eq!(config.api_key.as_deref(), Some("qdrant-key")); - assert_eq!(config.collection_name, "cache"); - assert_eq!(config.vector_size, 8); - assert_eq!(config.embedding.api_key, "embedding-key"); - assert_eq!(config.embedding.model, "text-embedding-3-small"); - restore_embedding_environment(py, prior).unwrap(); - }); - } - - #[test] - fn qdrant_projection_rejects_non_default_port() { - let _guard = env_lock().lock().unwrap_or_else(|error| error.into_inner()); - Python::initialize(); - Python::attach(|py| { - let prior = configure_embedding_environment(py, Some("embedding-key")).unwrap(); - for endpoint in [ - "https://qdrant.example:6332", - "https://qdrant.example", - "http://qdrant.example", - ] { - let facade = qdrant_facade(py, &format!("backend.qdrant_api_base = '{endpoint}'")); - let CacheConfigProjection::Unsupported(reason) = - NativeCacheConfig::project(&facade).unwrap() - else { - panic!("unsupported Qdrant endpoint should stay on Python"); - }; - assert!(matches!(reason, UnsupportedCacheConfig::QdrantEndpoint)); - } - let facade = qdrant_facade(py, ""); - let CacheConfigProjection::Native(config) = - NativeCacheConfig::project(&facade).unwrap() - else { - panic!("default Qdrant endpoint should use native"); - }; - let CacheBackendConfig::QdrantSemantic(config) = config.backend else { - panic!("expected Qdrant configuration"); - }; - assert!(config.grpc_url.ends_with(":6334")); - restore_embedding_environment(py, prior).unwrap(); - }); - } - - #[test] - fn qdrant_projection_passes_configured_embedding_model_to_router() { - let _guard = env_lock().lock().unwrap_or_else(|error| error.into_inner()); - Python::initialize(); - Python::attach(|py| { - let prior = configure_embedding_environment(py, Some("embedding-key")).unwrap(); - let facade = qdrant_facade(py, ""); - py.run( - c" -import sys -import types -proxy_server = types.ModuleType('litellm.proxy.proxy_server') -proxy_server.llm_router = None -proxy_server.llm_model_list = None -sys.modules['litellm.proxy.proxy_server'] = proxy_server -embedding_router = sys.modules['litellm.caching._embedding_router'] -embedding_router.resolve_embedding_router = lambda model, *_args: object() if model == 'openai/text-embedding-3-small' else None -", - None, - None, - ) - .unwrap(); - let CacheConfigProjection::Unsupported(reason) = - NativeCacheConfig::project(&facade).unwrap() - else { - panic!("router-backed embedding should stay on Python"); - }; - assert!(matches!(reason, UnsupportedCacheConfig::SemanticEmbedding)); - py.run( - c" -import sys -sys.modules.pop('litellm.proxy.proxy_server', None) -", - None, - None, - ) - .unwrap(); - restore_embedding_environment(py, prior).unwrap(); - }); - } - - #[test] - fn qdrant_projection_rejects_python_embedding_features() { - let _guard = env_lock().lock().unwrap_or_else(|error| error.into_inner()); - Python::initialize(); - Python::attach(|py| { - let prior = configure_embedding_environment(py, Some("embedding-key")).unwrap(); - let cases = [ - ("backend.embedding_max_input_tokens = 100", "semantic"), - ("backend.embedding_model = 'cohere/embed'", "semantic"), - ]; - for (extra, _) in cases { - let facade = qdrant_facade(py, extra); - let CacheConfigProjection::Unsupported(reason) = - NativeCacheConfig::project(&facade).unwrap() - else { - panic!("unsupported embedding should stay on Python"); - }; - assert!(matches!(reason, UnsupportedCacheConfig::SemanticEmbedding)); - } - restore_embedding_environment(py, prior).unwrap(); - }); - } - - #[test] - fn qdrant_projection_rejects_configured_litellm_base_or_missing_key() { - let _guard = env_lock().lock().unwrap_or_else(|error| error.into_inner()); - Python::initialize(); - Python::attach(|py| { - let prior = configure_embedding_environment(py, Some("embedding-key")).unwrap(); - let facade = qdrant_facade(py, ""); - let litellm = py.import("litellm").unwrap(); - litellm - .setattr("api_base", "https://proxy.example") - .unwrap(); - let CacheConfigProjection::Unsupported(reason) = - NativeCacheConfig::project(&facade).unwrap() - else { - panic!("configured LiteLLM base should stay on Python"); - }; - assert!(matches!(reason, UnsupportedCacheConfig::SemanticEmbedding)); - litellm.setattr("api_base", py.None()).unwrap(); - configure_embedding_environment(py, None).unwrap(); - let CacheConfigProjection::Unsupported(reason) = - NativeCacheConfig::project(&facade).unwrap() - else { - panic!("missing embedding key should stay on Python"); - }; - assert!(matches!(reason, UnsupportedCacheConfig::SemanticEmbedding)); - restore_embedding_environment(py, prior).unwrap(); - }); - } - - #[test] - fn qdrant_service_mismatch_reports_type_before_starting_qdrant() { - let _guard = env_lock().lock().unwrap_or_else(|error| error.into_inner()); - Python::initialize(); - Python::attach(|py| { - let prior = configure_embedding_environment(py, Some("embedding-key")).unwrap(); - let facade = qdrant_facade(py, ""); - let CacheConfigProjection::Native(config) = - NativeCacheConfig::project(&facade).unwrap() - else { - panic!("Qdrant cache should be supported"); - }; - let service = NativeResponseCache::memory(1, Duration::from_secs(1), 1024); - assert_eq!( - config.service_mismatch(&service), - Some("facade and native backend types must match") - ); - restore_embedding_environment(py, prior).unwrap(); - }); - } } diff --git a/litellm-rust/crates/python-bridge/src/cache/embedder.rs b/litellm-rust/crates/python-bridge/src/cache/embedder.rs new file mode 100644 index 00000000000..9398e5a862b --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/cache/embedder.rs @@ -0,0 +1,156 @@ +use std::future::Future; + +use litellm_cache::Error; +use litellm_host_python::to_py; +use pyo3::{PyTraverseError, PyVisit, prelude::*, types::PyDict}; +use serde_json::Value; + +tokio::task_local! { + static PREPARED_EMBEDDING: Result, Error>; +} + +pub(super) fn with_prepared_embedding( + vector: Result, Error>, + future: F, +) -> impl Future { + PREPARED_EMBEDDING.scope(vector, future) +} + +pub(super) struct PythonEmbedder(Py); + +impl Clone for PythonEmbedder { + fn clone(&self) -> Self { + Python::attach(|py| Self(self.0.clone_ref(py))) + } +} + +impl PythonEmbedder { + pub(super) fn new(object: Py) -> Self { + Self(object) + } + + pub(super) fn from_backend(backend: &Bound<'_, PyAny>) -> PyResult { + Ok(Self(backend.clone().unbind())) + } + + pub(super) fn object(&self) -> &Py { + &self.0 + } + + pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + visit.call(&self.0) + } + + pub(super) fn async_embed_awaitable<'py>( + &self, + py: Python<'py>, + prompt: &str, + metadata: &Option, + ) -> PyResult> { + let metadata = to_py(py, metadata)?; + self.0 + .bind(py) + .call_method1("_get_async_embedding", (prompt, metadata)) + } + + fn metadata_kwargs<'py>( + py: Python<'py>, + metadata: Option<&Value>, + ) -> PyResult> { + let kwargs = PyDict::new(py); + kwargs.set_item("metadata", to_py(py, &metadata)?)?; + Ok(kwargs) + } + + pub(super) fn async_embedding_coroutine( + &self, + py: Python<'_>, + prompt: &str, + metadata: Option<&Value>, + ) -> PyResult> { + let kwargs = Self::metadata_kwargs(py, metadata)?; + self.0 + .bind(py) + .call_method("_get_async_embedding", (prompt,), Some(&kwargs)) + .map(Bound::unbind) + } + + pub(super) fn extract(vector: Bound<'_, PyAny>) -> PyResult> { + Ok(vector + .extract::>()? + .into_iter() + .map(|value| value as f32) + .collect()) + } +} + +impl litellm_cache_valkey_semantic::Embedder for PythonEmbedder { + fn embed(&self, prompt: &str, metadata: Option<&Value>) -> Result, Error> { + let result = Python::attach(|py| -> PyResult> { + let metadata = to_py(py, &metadata)?; + self.0 + .bind(py) + .call_method1("_get_embedding", (prompt, metadata))? + .extract() + }) + .map_err(|_| Error::Unavailable)?; + Ok(result.into_iter().map(|value| value as f32).collect()) + } + + fn async_embed( + &self, + _prompt: &str, + _metadata: Option<&Value>, + ) -> impl Future, Error>> + Send { + let seeded = PREPARED_EMBEDDING + .try_with(Clone::clone) + .unwrap_or(Err(Error::Unavailable)); + std::future::ready(seeded) + } +} + +impl litellm_cache_redis_semantic::Embedder for PythonEmbedder { + fn embed(&self, prompt: &str, metadata: Option<&Value>) -> Result, Error> { + Python::attach(|py| { + let kwargs = Self::metadata_kwargs(py, metadata)?; + Self::extract(self.0.bind(py).call_method( + "_get_embedding", + (prompt,), + Some(&kwargs), + )?) + }) + .map_err(|_| Error::Unavailable) + } + + fn async_embed( + &self, + _prompt: &str, + _metadata: Option<&Value>, + ) -> impl Future, Error>> + Send { + let seeded = PREPARED_EMBEDDING + .try_with(Clone::clone) + .unwrap_or(Err(Error::Unavailable)); + std::future::ready(seeded) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn async_embed_returns_the_seeded_vector_or_unavailable() { + Python::initialize(); + let object = Python::attach(|py| py.None()); + let embedder = PythonEmbedder::new(object); + let scoped_embedder = embedder.clone(); + let scoped = with_prepared_embedding(Ok(vec![0.25]), async move { + litellm_cache_redis_semantic::Embedder::async_embed(&scoped_embedder, "prompt", None) + .await + }); + assert_eq!(scoped.await, Ok(vec![0.25])); + let unscoped = + litellm_cache_redis_semantic::Embedder::async_embed(&embedder, "prompt", None).await; + assert_eq!(unscoped, Err(Error::Unavailable)); + } +} diff --git a/litellm-rust/crates/python-bridge/src/cache/facade.rs b/litellm-rust/crates/python-bridge/src/cache/facade.rs index 3a998d511fd..43b53584943 100644 --- a/litellm-rust/crates/python-bridge/src/cache/facade.rs +++ b/litellm-rust/crates/python-bridge/src/cache/facade.rs @@ -31,6 +31,7 @@ struct RedisPoolGuard { connection_class: Py, connection_kwargs: Py, max_connections: Option, + client_name: &'static str, attributes: RedisPoolAttributes, } @@ -46,12 +47,18 @@ struct AzureBlobClientGuard { container_name: String, } +struct S3ClientGuard { + reference: Py, +} + enum ConnectionGuard { None, RedisPool(RedisPoolGuard), AzureBlob(AzureBlobClientGuard), + S3(S3ClientGuard), } +#[derive(Clone, Copy)] struct RedisPoolAttributes { pool: &'static str, connection_class: &'static str, @@ -70,6 +77,8 @@ const CLUSTER_POOL: RedisPoolAttributes = RedisPoolAttributes { max_connections: None, }; +const VALKEY_POOL: RedisPoolAttributes = STANDALONE_POOL; + pub(super) struct FacadeGuard { outer: ObjectGuard, backend: ObjectGuard, @@ -78,34 +87,6 @@ pub(super) struct FacadeGuard { } impl ObjectGuard { - fn class_behaviors(class: &Bound<'_, PyType>) -> PyResult)>> { - let py = class.py(); - let builtins = py.import("builtins")?; - let property_type = builtins.getattr("property")?; - let staticmethod_type = builtins.getattr("staticmethod")?; - let classmethod_type = builtins.getattr("classmethod")?; - class - .getattr("__dict__")? - .call_method0("items")? - .try_iter()? - .map(|item| { - let item = item?; - let (name, value): (String, Py) = item.extract()?; - let value_bound = value.bind(py); - let is_behavior = value_bound.is_callable() - || value_bound.is_instance(&property_type)? - || value_bound.is_instance(&staticmethod_type)? - || value_bound.is_instance(&classmethod_type)?; - Ok(is_behavior.then_some((name, value))) - }) - .filter_map(|result| match result { - Ok(Some(attribute)) => Some(Ok(attribute)), - Ok(None) => None, - Err(error) => Some(Err(error)), - }) - .collect() - } - fn capture( py: Python<'_>, object: &Bound<'_, PyAny>, @@ -118,7 +99,12 @@ impl ObjectGuard { .iter() .map(|class| { let class = class.cast_into::()?; - let attributes = Self::class_behaviors(&class)?; + let attributes = class + .getattr("__dict__")? + .call_method0("items")? + .try_iter()? + .map(|item| item?.extract::<(String, Py)>()) + .collect::>>()?; Ok(ClassGuard { class: class.unbind(), attributes, @@ -171,20 +157,16 @@ impl ObjectGuard { } let instance = object.getattr("__dict__")?.cast_into::()?; for (class, expected) in mro.iter().zip(&self.classes) { - let class = class.cast_into::()?; if !class.is(expected.class.bind(py)) { return Ok(false); } - let attributes = Self::class_behaviors(&class)?; - if attributes.len() != expected.attributes.len() { + let attributes = class.getattr("__dict__")?; + if attributes.len()? != expected.attributes.len() { return Ok(false); } - for ((name, value), (expected_name, expected_value)) in - attributes.iter().zip(&expected.attributes) - { - if name != expected_name - || instance.contains(name)? - || !value.bind(py).is(expected_value.bind(py)) + for (name, value) in &expected.attributes { + if (instance.contains(name)? && !self.config_names.contains(&name.as_str())) + || !attributes.get_item(name)?.is(value.bind(py)) { return Ok(false); } @@ -206,8 +188,12 @@ impl ObjectGuard { } impl RedisPoolGuard { - fn capture(backend: &Bound<'_, PyAny>, attributes: RedisPoolAttributes) -> PyResult { - let pool = backend.getattr("redis_client")?.getattr(attributes.pool)?; + fn capture( + backend: &Bound<'_, PyAny>, + client_name: &'static str, + attributes: RedisPoolAttributes, + ) -> PyResult { + let pool = backend.getattr(client_name)?.getattr(attributes.pool)?; Ok(Self { reference: pool.clone().unbind(), connection_class: pool.getattr(attributes.connection_class)?.unbind(), @@ -215,31 +201,30 @@ impl RedisPoolGuard { .getattr("connection_kwargs")? .call_method0("copy")? .unbind(), - max_connections: Self::max_connections(&pool, &attributes)?, + max_connections: attributes + .max_connections + .map(|name| pool.getattr(name)?.extract::()) + .transpose()?, + client_name, attributes, }) } - fn max_connections( - pool: &Bound<'_, PyAny>, - attributes: &RedisPoolAttributes, - ) -> PyResult> { - attributes - .max_connections - .map(|name| pool.getattr(name)?.extract::()) - .transpose() - } - fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult { let pool = backend - .getattr("redis_client")? + .getattr(self.client_name)? .getattr(self.attributes.pool)?; Ok(self.reference.bind(py).is(&pool) && self .connection_class .bind(py) .is(&pool.getattr(self.attributes.connection_class)?) - && self.max_connections == Self::max_connections(&pool, &self.attributes)? + && self.max_connections + == self + .attributes + .max_connections + .map(|name| pool.getattr(name)?.extract::()) + .transpose()? && self .connection_kwargs .bind(py) @@ -301,12 +286,43 @@ impl AzureBlobClientGuard { } } +impl S3ClientGuard { + fn capture(backend: &Bound<'_, PyAny>) -> PyResult { + Ok(Self { + reference: backend.getattr("s3_client")?.unbind(), + }) + } + + fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult { + Ok(self.reference.bind(py).is(&backend.getattr("s3_client")?)) + } + + fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + visit.call(&self.reference) + } +} + impl ConnectionGuard { fn capture(kind: &str, cluster: bool, backend: &Bound<'_, PyAny>) -> PyResult { Ok(match (kind, cluster) { - ("redis", false) => Self::RedisPool(RedisPoolGuard::capture(backend, STANDALONE_POOL)?), - ("redis", true) => Self::RedisPool(RedisPoolGuard::capture(backend, CLUSTER_POOL)?), + ("redis", false) => Self::RedisPool(RedisPoolGuard::capture( + backend, + "redis_client", + STANDALONE_POOL, + )?), + ("redis", true) => Self::RedisPool(RedisPoolGuard::capture( + backend, + "redis_client", + CLUSTER_POOL, + )?), + ("valkey-semantic", _) => Self::RedisPool(RedisPoolGuard::capture( + backend, + "sync_client", + VALKEY_POOL, + )?), + ("disk", _) => Self::None, ("azure-blob", _) => Self::AzureBlob(AzureBlobClientGuard::capture(backend)?), + ("s3", _) => Self::S3(S3ClientGuard::capture(backend)?), _ => Self::None, }) } @@ -316,6 +332,7 @@ impl ConnectionGuard { Self::None => Ok(true), Self::RedisPool(guard) => guard.matches(py, backend), Self::AzureBlob(guard) => guard.matches(py, backend), + Self::S3(guard) => guard.matches(py, backend), } } @@ -324,6 +341,7 @@ impl ConnectionGuard { Self::None => Ok(()), Self::RedisPool(guard) => guard.traverse(visit), Self::AzureBlob(guard) => guard.traverse(visit), + Self::S3(guard) => guard.traverse(visit), } } } @@ -345,23 +363,34 @@ impl FacadeGuard { let (module, name, cache_kind) = match (kind, cluster) { ("memory", _) => ("litellm.caching.in_memory_cache", "InMemoryCache", "local"), ("redis", false) => ("litellm.caching.redis_cache", "RedisCache", "redis"), - ("redis", true) => ( - "litellm.caching.redis_cluster_cache", - "RedisClusterCache", - "redis", + ("redis_semantic", _) => ( + "litellm.caching.redis_semantic_cache", + "RedisSemanticCache", + "redis-semantic", ), ("qdrant_semantic", _) => ( "litellm.caching.qdrant_semantic_cache", "QdrantSemanticCache", "qdrant-semantic", ), + ("redis", true) => ( + "litellm.caching.redis_cluster_cache", + "RedisClusterCache", + "redis", + ), ("gcs", _) => ("litellm.caching.gcs_cache", "GCSCache", "gcs"), + ("valkey-semantic", false) => ( + "litellm.caching.valkey_semantic_cache", + "ValkeySemanticCache", + "valkey-semantic", + ), ("disk", _) => ("litellm.caching.disk_cache", "DiskCache", "disk"), ("azure-blob", _) => ( "litellm.caching.azure_blob_cache", "AzureBlobCache", "azure-blob", ), + ("s3", _) => ("litellm.caching.s3_cache", "S3Cache", "s3"), _ => unreachable!(), }; let backend = facade.getattr("cache")?; @@ -381,30 +410,15 @@ impl FacadeGuard { if let Some(message) = config.service_mismatch(service) { return Err(PyTypeError::new_err(message)); } - let backend_config_names = match kind { - "memory" | "redis" | "azure-blob" | "disk" | "gcs" => &[ - "namespace", - "default_ttl", - "max_size_in_memory", - "max_size_per_item", - "redis_kwargs", - "redis_flush_size", - "bucket_name", - "key_prefix", - "path_service_account", - ][..], - "qdrant_semantic" => &[ - "qdrant_api_base", - "qdrant_api_key", - "collection_name", - "similarity_threshold", - "embedding_model", - "vector_size", - "embedding_max_input_tokens", - "embedding_timeout", - ][..], - _ => unreachable!(), - }; + if kind == "redis_semantic" + && service + .embedder_object() + .is_none_or(|embedder| !backend.is(embedder.bind(py))) + { + return Err(PyTypeError::new_err( + "facade backend must be the native embedder", + )); + } Ok(Self { outer: ObjectGuard::capture( py, @@ -419,7 +433,37 @@ impl FacadeGuard { "semantic_cache_scope", ], )?, - backend: ObjectGuard::capture(py, &backend, backend_config_names)?, + backend: ObjectGuard::capture( + py, + &backend, + &[ + "namespace", + "default_ttl", + "max_size_in_memory", + "max_size_per_item", + "redis_kwargs", + "redis_flush_size", + "similarity_threshold", + "distance_threshold", + "embedding_model", + "embedding_max_input_tokens", + "embedding_timeout", + "qdrant_api_base", + "qdrant_api_key", + "collection_name", + "vector_size", + "_index_name", + "_redis_url", + "similarity_threshold", + "embedding_model", + "index_name", + "embedding_max_input_tokens", + "embedding_timeout", + "bucket_name", + "key_prefix", + "path_service_account", + ], + )?, disk_store: (kind == "disk") .then(|| DiskStoreGuard::capture(&backend)) .transpose()?, @@ -477,51 +521,3 @@ pub(super) fn resolve( } handle.service().map(Some) } - -#[cfg(test)] -mod tests { - use super::ObjectGuard; - use pyo3::{prelude::*, types::PyDict}; - - #[test] - fn class_data_shadowing_is_ignored_but_method_mutations_are_rejected() { - Python::initialize(); - Python::attach(|py| { - let namespace = PyDict::new(py); - py.run( - c"class Example:\n data = 1\n def method(self):\n return 1\nobject = Example()\nobject.data = 2", - None, - Some(&namespace), - ) - .unwrap(); - let object = namespace.get_item("object").unwrap().unwrap(); - let guard = ObjectGuard::capture(py, &object, &[]).unwrap(); - - assert!(guard.matches(py, &object).unwrap()); - - py.run(c"object.method = lambda: 2", None, Some(&namespace)) - .unwrap(); - assert!(!guard.matches(py, &object).unwrap()); - }); - } - - #[test] - fn class_method_replacement_is_rejected() { - Python::initialize(); - Python::attach(|py| { - let namespace = PyDict::new(py); - py.run( - c"class Example:\n def method(self):\n return 1\nobject = Example()", - None, - Some(&namespace), - ) - .unwrap(); - let object = namespace.get_item("object").unwrap().unwrap(); - let guard = ObjectGuard::capture(py, &object, &[]).unwrap(); - - py.run(c"Example.method = lambda self: 2", None, Some(&namespace)) - .unwrap(); - assert!(!guard.matches(py, &object).unwrap()); - }); - } -} diff --git a/litellm-rust/crates/python-bridge/src/cache/handle.rs b/litellm-rust/crates/python-bridge/src/cache/handle.rs index 87f0ff98f67..51e1b02c405 100644 --- a/litellm-rust/crates/python-bridge/src/cache/handle.rs +++ b/litellm-rust/crates/python-bridge/src/cache/handle.rs @@ -1,19 +1,27 @@ -use std::env; - +use litellm_auth_aws::AwsAuthConfig; use litellm_cache_gcs::{DEFAULT_ENDPOINT, GcsConfig}; -use litellm_cache_redis::{RedisNode, RedisTopology}; -use litellm_http::ClientVariant; - use litellm_cache_qdrant_semantic::{OpenAiEmbedderConfig, Quantization}; +use litellm_cache_redis::{RedisNode, RedisTopology}; +use litellm_cache_redis_semantic::RedisSemanticConfig; +use litellm_cache_s3::{S3CacheConfig, S3Endpoint}; use litellm_host_python::{release_gil, run_sync_value}; -use pyo3::{PyTraverseError, PyVisit, exceptions::PyRuntimeError, prelude::*, types::PyDict}; +use litellm_http::ClientVariant; +use pyo3::{ + PyTraverseError, PyVisit, + exceptions::{PyRuntimeError, PyTypeError}, + prelude::*, + types::PyDict, +}; use url::Url; use super::{ - cache_error, config::QdrantSemanticCacheConfig, facade::FacadeGuard, - native::NativeResponseCache, request::duration, + cache_error, + config::{QdrantSemanticCacheConfig, project_redis_semantic}, + embedder::PythonEmbedder, + facade::FacadeGuard, + native::NativeResponseCache, + request::duration, }; -use crate::http; #[pyclass(frozen, name = "_CacheTestHandle")] pub(crate) struct CacheTestHandle { @@ -75,6 +83,40 @@ impl CacheTestHandle { }) } + #[staticmethod] + #[allow(clippy::too_many_arguments)] + #[pyo3(signature = (bucket, *, region, endpoint_url=None, key_prefix="", access_key_id=None, secret_access_key=None, session_token=None))] + fn s3( + py: Python<'_>, + bucket: String, + region: String, + endpoint_url: Option, + key_prefix: &str, + access_key_id: Option, + secret_access_key: Option, + session_token: Option, + ) -> PyResult { + let config = S3CacheConfig { + bucket, + key_prefix: key_prefix.to_string(), + region: region.clone(), + endpoint: endpoint_url.map(|url| S3Endpoint { url }), + auth: AwsAuthConfig { + access_key_id, + secret_access_key, + session_token, + region_name: Some(region), + ..Default::default() + }, + }; + let service = run_sync_value(py, async move { Ok(NativeResponseCache::s3(config).await) })?; + Ok(Self { + service, + guard: None, + pid: std::process::id(), + }) + } + #[staticmethod] #[pyo3(signature = (bucket_name, *, gcs_path=None, path_service_account=None, endpoint=None, token=None))] fn gcs( @@ -140,7 +182,7 @@ impl CacheTestHandle { || (!parsed.path().is_empty() && parsed.path() != "/") || parsed.query().is_some() || parsed.host_str().is_none() - || parsed.port().is_some_and(|port| port != 6333) + || parsed.port() != Some(6333) { return Err(pyo3::exceptions::PyValueError::new_err( "native Qdrant requires the default REST port so the gRPC port can be derived", @@ -156,7 +198,7 @@ impl CacheTestHandle { grpc_url.set_query(None); let embedding_api_key = embedding_api_key .or_else(|| { - env::var("OPENAI_API_KEY") + std::env::var("OPENAI_API_KEY") .ok() .filter(|value| !value.is_empty()) }) @@ -166,8 +208,8 @@ impl CacheTestHandle { ) })?; let embedding_api_base = embedding_api_base.unwrap_or_else(|| { - env::var("OPENAI_BASE_URL") - .or_else(|_| env::var("OPENAI_API_BASE")) + std::env::var("OPENAI_BASE_URL") + .or_else(|_| std::env::var("OPENAI_API_BASE")) .unwrap_or_else(|_| "https://api.openai.com/v1".to_owned()) }); let quantization = match quantization { @@ -194,10 +236,10 @@ impl CacheTestHandle { }, quantization, }; - let http_config = http::call_config(py, &PyDict::new(py), true)?; - let client = http::pool() + let http_config = crate::http::call_config(py, &PyDict::new(py), true)?; + let client = crate::http::pool() .client(&http_config, ClientVariant::Provider) - .map_err(http::client_error)?; + .map_err(crate::http::client_error)?; let service = run_sync_value(py, async move { let handle = tokio::runtime::Handle::current(); NativeResponseCache::qdrant_semantic(config, client, handle) @@ -211,6 +253,29 @@ impl CacheTestHandle { }) } + #[staticmethod] + #[pyo3(signature = (url, similarity_threshold, index_name, embedder))] + fn valkey_semantic( + url: String, + similarity_threshold: f64, + index_name: String, + embedder: &Bound<'_, PyAny>, + ) -> PyResult { + let python_embedder = PythonEmbedder::from_backend(embedder)?; + let service = NativeResponseCache::valkey_semantic( + &url, + similarity_threshold, + index_name, + python_embedder, + ) + .map_err(cache_error)?; + Ok(Self { + service, + guard: None, + pid: std::process::id(), + }) + } + #[staticmethod] #[pyo3(signature = (account_url, container))] fn azure_blob(py: Python<'_>, account_url: String, container: String) -> PyResult { @@ -226,6 +291,36 @@ impl CacheTestHandle { }) } + #[staticmethod] + fn redis_semantic(py: Python<'_>, backend: Bound<'_, PyAny>) -> PyResult { + let class = py + .import("litellm.caching.redis_semantic_cache")? + .getattr("RedisSemanticCache")?; + if !backend.get_type().is(&class) { + return Err(PyTypeError::new_err( + "native redis-semantic handles require the built-in RedisSemanticCache", + )); + } + let config = project_redis_semantic(&backend)?; + let embedder = PythonEmbedder::new(backend.unbind()); + let service = release_gil(py, move || { + NativeResponseCache::redis_semantic( + &config.redis_url, + embedder, + RedisSemanticConfig { + index_name: config.index_name, + similarity_threshold: config.similarity_threshold as f32, + }, + ) + }) + .map_err(cache_error)?; + Ok(Self { + service, + guard: None, + pid: std::process::id(), + }) + } + #[getter] fn backend(&self) -> &'static str { self.service.kind() @@ -234,11 +329,17 @@ impl CacheTestHandle { fn _bind_facade(&self, py: Python<'_>, facade: &Bound<'_, PyAny>) -> PyResult<()> { let service = self.service()?; let guard = FacadeGuard::capture(py, facade, &service)?; - let service = service.with_redis_flush_size( - facade - .getattr("redis_flush_size")? - .extract::>()?, - ); + let service = service + .with_scope( + facade + .getattr("semantic_cache_scope")? + .extract::()?, + ) + .with_redis_flush_size( + facade + .getattr("redis_flush_size")? + .extract::>()?, + ); let handle = Py::new( py, Self { @@ -251,6 +352,7 @@ impl CacheTestHandle { } fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> { + self.service.traverse(&visit)?; if let Some(guard) = &self.guard { guard.traverse(visit)?; } diff --git a/litellm-rust/crates/python-bridge/src/cache/mod.rs b/litellm-rust/crates/python-bridge/src/cache/mod.rs index ac4494150d9..fa028518559 100644 --- a/litellm-rust/crates/python-bridge/src/cache/mod.rs +++ b/litellm-rust/crates/python-bridge/src/cache/mod.rs @@ -1,12 +1,15 @@ mod binding; mod callback; mod config; +mod embedder; mod facade; mod future; mod handle; mod native; mod request; mod resolver; +mod semantic; +mod semantic_step; use litellm_cache::Error; use pyo3::{ diff --git a/litellm-rust/crates/python-bridge/src/cache/native.rs b/litellm-rust/crates/python-bridge/src/cache/native.rs index 7e1a4c27bc6..7ca62f10ec8 100644 --- a/litellm-rust/crates/python-bridge/src/cache/native.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native.rs @@ -1,18 +1,80 @@ use std::{path::Path, sync::Arc, time::Duration}; -use litellm_cache::{CacheCodec, CacheConnectionResult, Error, SemanticCacheContext}; +use litellm_cache::{ + CacheCodec, CacheConnectionResult, Error, ExactCacheContext, SemanticCacheContext, +}; use litellm_cache_azure_blob::AzureBlobCache; use litellm_cache_disk::DiskCache; use litellm_cache_gcs::{GcsCache, GcsConfig, StaticTokenSource}; use litellm_cache_memory::InMemoryCache; use litellm_cache_qdrant_semantic::{Embedder, OpenAiEmbedder, QdrantSemanticCache}; use litellm_cache_redis::{RedisCache, RedisTopology}; +use litellm_cache_redis_semantic::{RedisSemanticCache, RedisSemanticConfig}; use litellm_cache_response::{ - CacheEntry, PartialHits, ResponseCache, ResponseCacheCodec, ResponseCacheRequest, WriteBuffer, + CacheEntry, CacheKeyField, PartialHits, ResponseCache, ResponseCacheCodec, + ResponseCacheRequest, WriteBuffer, }; +use litellm_cache_s3::{S3Cache, S3CacheConfig}; +use litellm_cache_valkey_semantic::{ValkeySemanticCache, ValkeySemanticConfig}; +use pyo3::{PyTraverseError, PyVisit, prelude::*}; use serde_json::Value; -use super::{config::QdrantSemanticCacheConfig, request::exact}; +use super::{ + config::QdrantSemanticCacheConfig, + embedder::PythonEmbedder, + request::NativeRequest, + semantic::{SemanticBody, SemanticOperation, drive}, + semantic_step::{SemanticEmbedExecution, drive_semantic}, +}; + +fn semantic_key(request: &NativeRequest, scope: &str) -> litellm_cache_response::CacheKeyInput { + let mut key = request.key.clone(); + if key.preset.is_some() { + return key; + } + key.fields + .retain(|field| !matches!(field.name.as_str(), "messages" | "prompt" | "input")); + const TENANT: [&str; 3] = [ + "user_api_key", + "user_api_key_team_id", + "user_api_key_org_id", + ]; + let end_user = (scope == "end_user").then_some("user_api_key_end_user_id"); + for name in TENANT.into_iter().chain(end_user) { + let sources = [ + request.metadata.as_ref(), + request.litellm_metadata.as_ref(), + request + .litellm_params + .as_ref() + .and_then(|params| params.get("metadata")), + request + .litellm_params + .as_ref() + .and_then(|params| params.get("litellm_metadata")), + ]; + let Some(value) = sources.into_iter().flatten().find_map(|source| { + source + .as_object() + .and_then(|values| values.get(name)) + .filter(|value| !value.is_null()) + }) else { + continue; + }; + let value = match value { + Value::Null => continue, + Value::String(text) => text.clone(), + other => other.to_string(), + }; + key.fields.push(CacheKeyField { + name: name.to_owned(), + value: Some(value), + api_parameter: true, + internal_parameter: false, + }); + } + key +} #[derive(Clone)] pub(super) enum NativeResponseCache { @@ -21,9 +83,19 @@ pub(super) enum NativeResponseCache { cache: Arc>>, buffer: Option>, }, + S3(Arc>>), Gcs(Arc>>), - Disk(Arc>>), + ValkeySemantic { + cache: Arc>>, + embedder: PythonEmbedder, + scope: String, + }, + RedisSemantic { + cache: Arc>>, + embedder: PythonEmbedder, + }, QdrantSemantic(Arc>>), + Disk(Arc>>), AzureBlob(Arc>>), } @@ -56,21 +128,47 @@ impl NativeResponseCache { }) } - pub fn disk(directory: &str) -> Result { - let cache = DiskCache::open(directory, ResponseCacheCodec)?; - Ok(Self::Disk(Arc::new(ResponseCache::new(Arc::new(cache))))) + pub async fn s3(config: S3CacheConfig) -> Self { + let runtime = tokio::runtime::Handle::current(); + Self::S3(Arc::new(ResponseCache::new(Arc::new(S3Cache::new( + config, + ResponseCacheCodec, + runtime, + ))))) } - pub fn gcs(config: GcsConfig, token: Option) -> Result { - let backend = match token { - Some(token) => GcsCache::with_token_source( - config, - ResponseCacheCodec, - Arc::new(StaticTokenSource(token)), - )?, - None => GcsCache::new(config, ResponseCacheCodec)?, - }; - Ok(Self::Gcs(Arc::new(ResponseCache::new(Arc::new(backend))))) + pub fn valkey_semantic( + url: &str, + similarity_threshold: f64, + index_name: String, + embedder: PythonEmbedder, + ) -> Result { + let backend = ValkeySemanticCache::new( + url, + embedder.clone(), + ResponseCacheCodec, + ValkeySemanticConfig { + similarity_threshold, + index_name, + }, + )?; + Ok(Self::ValkeySemantic { + cache: Arc::new(ResponseCache::new(Arc::new(backend))), + embedder, + scope: String::from("key"), + }) + } + + pub fn redis_semantic( + url: &str, + embedder: PythonEmbedder, + config: RedisSemanticConfig, + ) -> Result { + let backend = RedisSemanticCache::new(url, embedder.clone(), config)?; + Ok(Self::RedisSemantic { + cache: Arc::new(ResponseCache::new(Arc::new(backend))), + embedder, + }) } pub async fn qdrant_semantic( @@ -98,6 +196,23 @@ impl NativeResponseCache { )))) } + pub fn disk(directory: &str) -> Result { + let cache = DiskCache::open(directory, ResponseCacheCodec)?; + Ok(Self::Disk(Arc::new(ResponseCache::new(Arc::new(cache))))) + } + + pub fn gcs(config: GcsConfig, token: Option) -> Result { + let backend = match token { + Some(token) => GcsCache::with_token_source( + config, + ResponseCacheCodec, + Arc::new(StaticTokenSource(token)), + )?, + None => GcsCache::new(config, ResponseCacheCodec)?, + }; + Ok(Self::Gcs(Arc::new(ResponseCache::new(Arc::new(backend))))) + } + pub async fn azure_blob(account_url: &str, container: &str) -> Result { let backend = AzureBlobCache::connect( account_url, @@ -119,21 +234,92 @@ impl NativeResponseCache { )), Self::Memory(_) | Self::Redis { .. } - | Self::Gcs(_) + | Self::S3(_) + | Self::ValkeySemantic { .. } + | Self::RedisSemantic { .. } + | Self::QdrantSemantic(_) | Self::Disk(_) - | Self::QdrantSemantic(_) => None, + | Self::Gcs(_) => None, + } + } + + fn exact(request: &NativeRequest) -> ResponseCacheRequest { + ResponseCacheRequest { + key: request.key.clone(), + controls: request.controls, + context: ExactCacheContext { ttl: request.ttl }, + max_age: request.max_age, + } + } + + pub(super) fn semantic_request( + request: &NativeRequest, + ) -> ResponseCacheRequest { + ResponseCacheRequest { + key: request.key.clone(), + controls: request.controls, + context: SemanticCacheContext { + input: request.input.clone(), + messages: request.messages.clone(), + metadata: request.metadata.clone(), + scope: request.scope.clone(), + ttl: request.ttl, + }, + max_age: request.max_age, + } + } + + fn semantic( + request: &NativeRequest, + scope: &str, + ) -> ResponseCacheRequest { + ResponseCacheRequest { + key: semantic_key(request, scope), + controls: request.controls, + context: SemanticCacheContext { + input: request.input.clone(), + messages: request.messages.clone(), + metadata: request.metadata.clone(), + scope: Some(scope.to_owned()), + ttl: request.ttl, + }, + max_age: request.max_age, + } + } + + pub fn with_redis_flush_size(self, flush_size: Option) -> Self { + match self { + Self::Redis { cache, .. } => Self::Redis { + cache, + buffer: flush_size.map(|size| Arc::new(WriteBuffer::new(size))), + }, + value => value, + } + } + + pub fn with_scope(self, scope: String) -> Self { + match self { + Self::ValkeySemantic { + cache, embedder, .. + } => Self::ValkeySemantic { + cache, + embedder, + scope, + }, + value => value, } } -} -impl NativeResponseCache { pub fn kind(&self) -> &'static str { match self { Self::Memory(_) => "memory", Self::Redis { .. } => "redis", + Self::S3(_) => "s3", Self::Gcs(_) => "gcs", - Self::Disk(_) => "disk", + Self::ValkeySemantic { .. } => "valkey-semantic", + Self::RedisSemantic { .. } => "redis_semantic", Self::QdrantSemantic(_) => "qdrant_semantic", + Self::Disk(_) => "disk", Self::AzureBlob(_) => "azure-blob", } } @@ -142,18 +328,52 @@ impl NativeResponseCache { match self { Self::Memory(cache) => cache.default_ttl(), Self::Redis { cache, .. } => cache.default_ttl(), + Self::S3(cache) => cache.default_ttl(), Self::Gcs(cache) => cache.default_ttl(), - Self::Disk(cache) => cache.default_ttl(), + Self::ValkeySemantic { cache, .. } => cache.default_ttl(), + Self::RedisSemantic { cache, .. } => cache.default_ttl(), Self::QdrantSemantic(_) => None, + Self::Disk(cache) => cache.default_ttl(), Self::AzureBlob(cache) => cache.default_ttl(), } } + pub fn bucket(&self) -> Option<&str> { + match self { + Self::S3(cache) => Some(cache.backend().bucket()), + _ => None, + } + } + + pub fn key_prefix(&self) -> Option<&str> { + match self { + Self::S3(cache) => Some(cache.backend().key_prefix()), + _ => None, + } + } + + pub fn region(&self) -> Option<&str> { + match self { + Self::S3(cache) => Some(cache.backend().region()), + _ => None, + } + } + + pub fn endpoint(&self) -> Option<&str> { + match self { + Self::S3(cache) => cache.backend().endpoint(), + _ => None, + } + } + pub fn namespace(&self) -> Option<&str> { match self { Self::Memory(_) - | Self::Disk(_) + | Self::S3(_) + | Self::ValkeySemantic { .. } + | Self::RedisSemantic { .. } | Self::QdrantSemantic(_) + | Self::Disk(_) | Self::AzureBlob(_) | Self::Gcs(_) => None, Self::Redis { cache, .. } => cache.backend().namespace(), @@ -163,8 +383,11 @@ impl NativeResponseCache { pub fn topology(&self) -> Option<&RedisTopology> { match self { Self::Memory(_) - | Self::Disk(_) + | Self::S3(_) + | Self::ValkeySemantic { .. } + | Self::RedisSemantic { .. } | Self::QdrantSemantic(_) + | Self::Disk(_) | Self::AzureBlob(_) | Self::Gcs(_) => None, Self::Redis { cache, .. } => Some(cache.backend().topology()), @@ -174,28 +397,28 @@ impl NativeResponseCache { pub fn capacity(&self) -> Option { match self { Self::Memory(cache) => Some(cache.backend().max_size_in_memory()), - Self::Redis { .. } | Self::Gcs(_) | Self::Disk(_) => None, - Self::QdrantSemantic(_) => None, - Self::AzureBlob(_) => None, + Self::Redis { .. } + | Self::S3(_) + | Self::ValkeySemantic { .. } + | Self::RedisSemantic { .. } + | Self::QdrantSemantic(_) + | Self::Disk(_) + | Self::AzureBlob(_) + | Self::Gcs(_) => None, } } pub fn max_entry_bytes(&self) -> Option { match self { Self::Memory(cache) => cache.backend().max_entry_bytes(), - Self::Redis { .. } | Self::Gcs(_) | Self::Disk(_) => None, - Self::QdrantSemantic(_) => None, - Self::AzureBlob(_) => None, - } - } - - pub fn with_redis_flush_size(self, flush_size: Option) -> Self { - match self { - Self::Redis { cache, .. } => Self::Redis { - cache, - buffer: flush_size.map(|flush_size| Arc::new(WriteBuffer::new(flush_size))), - }, - other => other, + Self::Redis { .. } + | Self::S3(_) + | Self::ValkeySemantic { .. } + | Self::RedisSemantic { .. } + | Self::QdrantSemantic(_) + | Self::Disk(_) + | Self::AzureBlob(_) + | Self::Gcs(_) => None, } } @@ -204,26 +427,53 @@ impl NativeResponseCache { Self::Disk(cache) => Some(cache.backend().directory()), Self::Memory(_) | Self::Redis { .. } - | Self::Gcs(_) + | Self::S3(_) + | Self::ValkeySemantic { .. } + | Self::RedisSemantic { .. } | Self::QdrantSemantic(_) - | Self::AzureBlob(_) => None, + | Self::AzureBlob(_) + | Self::Gcs(_) => None, } } - pub fn collection_name(&self) -> Option<&str> { + pub fn semantic_config(&self) -> Option<(f64, &str)> { match self { - Self::QdrantSemantic(cache) => Some(cache.backend().collection_name()), + Self::ValkeySemantic { cache, .. } => Some(( + cache.backend().similarity_threshold(), + cache.backend().index_name(), + )), + Self::RedisSemantic { cache, .. } => Some(( + f64::from(cache.backend().similarity_threshold()), + cache.backend().index_name(), + )), + _ => None, + } + } + + pub fn index_name(&self) -> Option<&str> { + match self { + Self::RedisSemantic { cache, .. } => Some(cache.backend().index_name()), _ => None, } } pub fn similarity_threshold(&self) -> Option { match self { + Self::RedisSemantic { cache, .. } => { + Some(f64::from(cache.backend().similarity_threshold())) + } Self::QdrantSemantic(cache) => Some(cache.backend().similarity_threshold()), _ => None, } } + pub fn collection_name(&self) -> Option<&str> { + match self { + Self::QdrantSemantic(cache) => Some(cache.backend().collection_name()), + _ => None, + } + } + pub fn vector_size(&self) -> Option { match self { Self::QdrantSemantic(cache) => Some(cache.backend().vector_size()), @@ -238,198 +488,444 @@ impl NativeResponseCache { } } - pub fn lookup( - &self, - request: &ResponseCacheRequest, - now: Duration, - ) -> Result, Error> { + pub fn semantic_embedder(&self) -> Option<&PythonEmbedder> { match self { - Self::Memory(cache) => cache.lookup(&exact(request), now), - Self::Redis { cache, .. } => cache.lookup(&exact(request), now), - Self::Gcs(cache) => cache.lookup(&exact(request), now), - Self::Disk(cache) => cache.lookup(&exact(request), now), - Self::QdrantSemantic(cache) => cache.lookup(request, now), - Self::AzureBlob(cache) => cache.lookup(&exact(request), now), + Self::RedisSemantic { embedder, .. } => Some(embedder), + _ => None, + } + } + + pub fn embedder_object(&self) -> Option<&Py> { + match self { + Self::RedisSemantic { embedder, .. } => Some(embedder.object()), + _ => None, + } + } + + pub fn lookup(&self, request: &NativeRequest, now: Duration) -> Result, Error> { + match self { + Self::Memory(cache) => cache.lookup(&Self::exact(request), now), + Self::Redis { cache, .. } => cache.lookup(&Self::exact(request), now), + Self::S3(cache) => cache.lookup(&Self::exact(request), now), + Self::ValkeySemantic { cache, scope, .. } => { + cache.lookup(&Self::semantic(request, scope), now) + } + Self::RedisSemantic { cache, .. } => { + cache.lookup(&Self::semantic_request(request), now) + } + Self::Gcs(cache) => cache.lookup(&Self::exact(request), now), + Self::QdrantSemantic(cache) => cache.lookup(&Self::semantic_request(request), now), + Self::Disk(cache) => cache.lookup(&Self::exact(request), now), + Self::AzureBlob(cache) => cache.lookup(&Self::exact(request), now), } } pub fn store( &self, - request: &ResponseCacheRequest, + request: &NativeRequest, response: Value, now: Duration, ) -> Result<(), Error> { match self { - Self::Memory(cache) => cache.store(&exact(request), response, now), - Self::Redis { cache, .. } => cache.store(&exact(request), response, now), - Self::Gcs(cache) => cache.store(&exact(request), response, now), - Self::Disk(cache) => cache.store(&exact(request), response, now), - Self::QdrantSemantic(cache) => cache.store(request, response, now), - Self::AzureBlob(cache) => cache.store(&exact(request), response, now), + Self::Memory(cache) => cache.store(&Self::exact(request), response, now), + Self::Redis { cache, .. } => cache.store(&Self::exact(request), response, now), + Self::S3(cache) => cache.store(&Self::exact(request), response, now), + Self::ValkeySemantic { cache, scope, .. } => { + cache.store(&Self::semantic(request, scope), response, now) + } + Self::RedisSemantic { cache, .. } => { + cache.store(&Self::semantic_request(request), response, now) + } + Self::Gcs(cache) => cache.store(&Self::exact(request), response, now), + Self::QdrantSemantic(cache) => { + cache.store(&Self::semantic_request(request), response, now) + } + Self::Disk(cache) => cache.store(&Self::exact(request), response, now), + Self::AzureBlob(cache) => cache.store(&Self::exact(request), response, now), } } pub fn lookup_batch( &self, - requests: &[ResponseCacheRequest], + requests: &[NativeRequest], now: Duration, ) -> Result { match self { Self::Memory(cache) => { - cache.lookup_batch(&requests.iter().map(exact).collect::>(), now) + let requests = requests.iter().map(Self::exact).collect::>(); + cache.lookup_batch(&requests, now) } Self::Redis { cache, .. } => { - cache.lookup_batch(&requests.iter().map(exact).collect::>(), now) + let requests = requests.iter().map(Self::exact).collect::>(); + cache.lookup_batch(&requests, now) + } + Self::S3(cache) => { + cache.lookup_batch(&requests.iter().map(Self::exact).collect::>(), now) + } + Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } | Self::QdrantSemantic(_) => { + Err(Error::UnsupportedOperation) } Self::Gcs(cache) => { - cache.lookup_batch(&requests.iter().map(exact).collect::>(), now) + cache.lookup_batch(&requests.iter().map(Self::exact).collect::>(), now) } Self::Disk(cache) => { - cache.lookup_batch(&requests.iter().map(exact).collect::>(), now) + cache.lookup_batch(&requests.iter().map(Self::exact).collect::>(), now) } - Self::QdrantSemantic(_) => Err(Error::UnsupportedOperation), Self::AzureBlob(cache) => { - cache.lookup_batch(&requests.iter().map(exact).collect::>(), now) + cache.lookup_batch(&requests.iter().map(Self::exact).collect::>(), now) } } } pub async fn async_lookup( &self, - request: &ResponseCacheRequest, + request: &NativeRequest, now: Duration, ) -> Result, Error> { match self { - Self::Memory(cache) => cache.async_lookup(&exact(request), now).await, - Self::Redis { cache, .. } => cache.async_lookup(&exact(request), now).await, - Self::Gcs(cache) => cache.async_lookup(&exact(request), now).await, - Self::Disk(cache) => cache.async_lookup(&exact(request), now).await, - Self::QdrantSemantic(cache) => cache.async_lookup(request, now).await, - Self::AzureBlob(cache) => cache.async_lookup(&exact(request), now).await, + Self::Memory(cache) => cache.async_lookup(&Self::exact(request), now).await, + Self::Redis { cache, .. } => cache.async_lookup(&Self::exact(request), now).await, + Self::S3(cache) => cache.async_lookup(&Self::exact(request), now).await, + Self::ValkeySemantic { cache, scope, .. } => { + cache + .async_lookup(&Self::semantic(request, scope), now) + .await + } + Self::RedisSemantic { cache, .. } => { + cache + .async_lookup(&Self::semantic_request(request), now) + .await + } + Self::QdrantSemantic(cache) => { + cache + .async_lookup(&Self::semantic_request(request), now) + .await + } + Self::Gcs(cache) => cache.async_lookup(&Self::exact(request), now).await, + Self::Disk(cache) => cache.async_lookup(&Self::exact(request), now).await, + Self::AzureBlob(cache) => cache.async_lookup(&Self::exact(request), now).await, + } + } + + pub(super) fn async_lookup_py<'py>( + &self, + py: Python<'py>, + request: NativeRequest, + ) -> PyResult> { + match self { + Self::Memory(_) + | Self::Redis { .. } + | Self::S3(_) + | Self::Disk(_) + | Self::AzureBlob(_) + | Self::Gcs(_) => { + let service = self.clone(); + litellm_host_python::run_async( + py, + async move { service.async_lookup(&request, super::request::now()).await }, + super::cache_error, + ) + } + Self::ValkeySemantic { + cache, + embedder, + scope, + } => drive_semantic( + py, + SemanticEmbedExecution::lookup( + Arc::clone(cache.backend_arc()), + embedder.clone(), + Self::semantic(&request, scope), + ), + ), + Self::RedisSemantic { .. } => drive( + py, + SemanticBody::new(self.clone(), SemanticOperation::Lookup(request)), + ), + Self::QdrantSemantic(_) => { + let service = self.clone(); + litellm_host_python::run_async( + py, + async move { service.async_lookup(&request, super::request::now()).await }, + super::cache_error, + ) + } } } pub async fn async_store( &self, - request: &ResponseCacheRequest, + request: &NativeRequest, response: Value, now: Duration, ) -> Result<(), Error> { match self { - Self::Memory(cache) => cache.async_store(&exact(request), response, now).await, + Self::Memory(cache) => { + cache + .async_store(&Self::exact(request), response, now) + .await + } Self::Redis { cache, buffer: None, - } => cache.async_store(&exact(request), response, now).await, + } => { + cache + .async_store(&Self::exact(request), response, now) + .await + } Self::Redis { cache, buffer: Some(buffer), } => { - let request = exact(request); - buffer.async_store(cache, &request, response, now).await + buffer + .async_store(cache, &Self::exact(request), response, now) + .await + } + Self::S3(cache) => { + cache + .async_store(&Self::exact(request), response, now) + .await + } + Self::ValkeySemantic { cache, scope, .. } => { + cache + .async_store(&Self::semantic(request, scope), response, now) + .await + } + Self::RedisSemantic { cache, .. } => { + cache + .async_store(&Self::semantic_request(request), response, now) + .await + } + Self::QdrantSemantic(cache) => { + cache + .async_store(&Self::semantic_request(request), response, now) + .await + } + Self::Gcs(cache) => { + cache + .async_store(&Self::exact(request), response, now) + .await + } + Self::Disk(cache) => { + cache + .async_store(&Self::exact(request), response, now) + .await + } + Self::AzureBlob(cache) => { + cache + .async_store(&Self::exact(request), response, now) + .await + } + } + } + + pub(super) fn async_store_py<'py>( + &self, + py: Python<'py>, + request: NativeRequest, + response: Value, + ) -> PyResult> { + match self { + Self::Memory(_) + | Self::Redis { .. } + | Self::S3(_) + | Self::Disk(_) + | Self::AzureBlob(_) + | Self::Gcs(_) => { + let service = self.clone(); + litellm_host_python::run_async( + py, + async move { + service + .async_store(&request, response, super::request::now()) + .await + }, + super::cache_error, + ) + } + Self::ValkeySemantic { + cache, + embedder, + scope, + } => drive_semantic( + py, + SemanticEmbedExecution::store( + Arc::clone(cache.backend_arc()), + embedder.clone(), + Self::semantic(&request, scope), + response, + ), + ), + Self::RedisSemantic { .. } => drive( + py, + SemanticBody::new(self.clone(), SemanticOperation::Store(request, response)), + ), + Self::QdrantSemantic(_) => { + let service = self.clone(); + litellm_host_python::run_async( + py, + async move { + service + .async_store(&request, response, super::request::now()) + .await + }, + super::cache_error, + ) } - Self::Gcs(cache) => cache.async_store(&exact(request), response, now).await, - Self::Disk(cache) => cache.async_store(&exact(request), response, now).await, - Self::QdrantSemantic(cache) => cache.async_store(request, response, now).await, - Self::AzureBlob(cache) => cache.async_store(&exact(request), response, now).await, } } pub async fn async_lookup_batch( &self, - requests: &[ResponseCacheRequest], + requests: &[NativeRequest], now: Duration, ) -> Result { match self { Self::Memory(cache) => { - let requests = requests.iter().map(exact).collect::>(); + let requests = requests.iter().map(Self::exact).collect::>(); cache.async_lookup_batch(&requests, now).await } Self::Redis { cache, .. } => { - let requests = requests.iter().map(exact).collect::>(); + let requests = requests.iter().map(Self::exact).collect::>(); cache.async_lookup_batch(&requests, now).await } + Self::S3(cache) => { + cache + .async_lookup_batch(&requests.iter().map(Self::exact).collect::>(), now) + .await + } + Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } | Self::QdrantSemantic(_) => { + Err(Error::UnsupportedOperation) + } Self::Gcs(cache) => { - let requests = requests.iter().map(exact).collect::>(); - cache.async_lookup_batch(&requests, now).await + cache + .async_lookup_batch(&requests.iter().map(Self::exact).collect::>(), now) + .await } Self::Disk(cache) => { - let requests = requests.iter().map(exact).collect::>(); - cache.async_lookup_batch(&requests, now).await + cache + .async_lookup_batch(&requests.iter().map(Self::exact).collect::>(), now) + .await } - Self::QdrantSemantic(_) => Err(Error::UnsupportedOperation), Self::AzureBlob(cache) => { - let requests = requests.iter().map(exact).collect::>(); - cache.async_lookup_batch(&requests, now).await + cache + .async_lookup_batch(&requests.iter().map(Self::exact).collect::>(), now) + .await } } } pub async fn async_store_batch( &self, - entries: Vec<(ResponseCacheRequest, Value)>, + entries: Vec<(NativeRequest, Value)>, now: Duration, ) -> Result<(), Error> { match self { Self::Memory(cache) => { - cache - .async_store_batch( - entries - .into_iter() - .map(|(request, value)| (exact(&request), value)) - .collect(), - now, - ) - .await + let entries = entries + .into_iter() + .map(|(request, value)| (Self::exact(&request), value)) + .collect(); + cache.async_store_batch(entries, now).await } Self::Redis { cache, .. } => { - cache - .async_store_batch( - entries - .into_iter() - .map(|(request, value)| (exact(&request), value)) - .collect(), - now, - ) - .await + let entries = entries + .into_iter() + .map(|(request, value)| (Self::exact(&request), value)) + .collect(); + cache.async_store_batch(entries, now).await + } + Self::S3(cache) => { + let entries = entries + .into_iter() + .map(|(request, value)| (Self::exact(&request), value)) + .collect(); + cache.async_store_batch(entries, now).await + } + Self::ValkeySemantic { cache, scope, .. } => { + let entries = entries + .into_iter() + .map(|(request, value)| (Self::semantic(&request, scope), value)) + .collect(); + cache.async_store_batch(entries, now).await + } + Self::RedisSemantic { .. } | Self::QdrantSemantic(_) => { + Err(Error::UnsupportedOperation) } Self::Gcs(cache) => { - cache - .async_store_batch( - entries - .into_iter() - .map(|(request, value)| (exact(&request), value)) - .collect(), - now, - ) - .await + let entries = entries + .into_iter() + .map(|(request, value)| (Self::exact(&request), value)) + .collect(); + cache.async_store_batch(entries, now).await } Self::Disk(cache) => { - cache - .async_store_batch( - entries - .into_iter() - .map(|(request, value)| (exact(&request), value)) - .collect(), - now, - ) - .await + let entries = entries + .into_iter() + .map(|(request, value)| (Self::exact(&request), value)) + .collect(); + cache.async_store_batch(entries, now).await } - Self::QdrantSemantic(cache) => cache.async_store_batch(entries, now).await, Self::AzureBlob(cache) => { - cache - .async_store_batch( - entries - .into_iter() - .map(|(request, value)| (exact(&request), value)) - .collect(), - now, - ) - .await + let entries = entries + .into_iter() + .map(|(request, value)| (Self::exact(&request), value)) + .collect(); + cache.async_store_batch(entries, now).await } } } + pub(super) fn async_store_batch_py<'py>( + &self, + py: Python<'py>, + entries: Vec<(NativeRequest, Value)>, + ) -> PyResult> { + match self { + Self::Memory(_) + | Self::Redis { .. } + | Self::S3(_) + | Self::Disk(_) + | Self::AzureBlob(_) + | Self::Gcs(_) => { + let service = self.clone(); + litellm_host_python::run_async( + py, + async move { + service + .async_store_batch(entries, super::request::now()) + .await + }, + super::cache_error, + ) + } + Self::ValkeySemantic { + cache, + embedder, + scope, + } => { + let (requests, responses): (Vec<_>, Vec<_>) = entries + .into_iter() + .map(|(request, response)| (Self::semantic(&request, scope), response)) + .unzip(); + drive_semantic( + py, + SemanticEmbedExecution::store_batch( + Arc::clone(cache.backend_arc()), + embedder.clone(), + requests, + responses, + ), + ) + } + Self::RedisSemantic { .. } => drive( + py, + SemanticBody::new(self.clone(), SemanticOperation::StoreBatch(entries.into())), + ), + Self::QdrantSemantic(_) => Err(super::cache_error(Error::UnsupportedOperation)), + } + } + pub async fn async_flush(&self) -> Result<(), Error> { match self { Self::Memory(cache) => cache.async_flush().await, @@ -439,9 +935,12 @@ impl NativeResponseCache { } cache.async_flush().await } + Self::S3(cache) => cache.async_flush().await, + Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } | Self::QdrantSemantic(_) => { + Err(Error::UnsupportedOperation) + } Self::Gcs(cache) => cache.async_flush().await, Self::Disk(cache) => cache.async_flush().await, - Self::QdrantSemantic(_) => Err(Error::UnsupportedOperation), Self::AzureBlob(cache) => cache.async_flush().await, } } @@ -450,13 +949,26 @@ impl NativeResponseCache { match self { Self::Memory(cache) => cache.test_connection().await, Self::Redis { cache, .. } => cache.test_connection().await, + Self::S3(cache) => cache.test_connection().await, + Self::ValkeySemantic { cache, .. } => cache.test_connection().await, + Self::RedisSemantic { .. } | Self::QdrantSemantic(_) => { + Err(Error::UnsupportedOperation) + } Self::Gcs(cache) => cache.test_connection().await, Self::Disk(cache) => cache.test_connection().await, - Self::QdrantSemantic(_) => Err(Error::UnsupportedOperation), Self::AzureBlob(cache) => cache.test_connection().await, } } + pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + match self { + Self::ValkeySemantic { embedder, .. } => embedder.traverse(visit)?, + Self::RedisSemantic { embedder, .. } => embedder.traverse(visit)?, + _ => {} + } + Ok(()) + } + pub fn gcs_backend(&self) -> Option<&GcsCache> { match self { Self::Gcs(cache) => Some(cache.backend()), @@ -464,3 +976,80 @@ impl NativeResponseCache { } } } + +#[cfg(test)] +mod tests { + use litellm_cache_response::{CacheControls, CacheKeyInput, cache_key}; + use serde_json::json; + use sha2::{Digest, Sha256}; + + use super::*; + + fn native_request(key: CacheKeyInput, metadata: Value) -> NativeRequest { + NativeRequest { + key, + controls: CacheControls::default(), + ttl: None, + max_age: None, + messages: Some(json!([{"role": "user", "content": "prompt"}])), + input: None, + metadata: Some(metadata), + litellm_metadata: None, + litellm_params: None, + scope: None, + } + } + + #[test] + fn semantic_key_matches_python_scope_material() { + let key = CacheKeyInput { + fields: vec![ + CacheKeyField { + name: "model".to_owned(), + value: Some("gpt-4.1".to_owned()), + api_parameter: true, + internal_parameter: false, + }, + CacheKeyField { + name: "messages".to_owned(), + value: Some("prompt".to_owned()), + api_parameter: true, + internal_parameter: false, + }, + ], + ..Default::default() + }; + let request = native_request( + key, + json!({"user_api_key": "k1", "user_api_key_team_id": null}), + ); + let expected = format!("{:x}", Sha256::digest(b"model: gpt-4.1user_api_key: k1")); + assert_eq!(cache_key(&semantic_key(&request, "key")), expected); + + let end_user_request = native_request( + request.key.clone(), + json!({"user_api_key": "k1", "user_api_key_end_user_id": "u1"}), + ); + let expected = format!( + "{:x}", + Sha256::digest(b"model: gpt-4.1user_api_key: k1user_api_key_end_user_id: u1") + ); + assert_eq!( + cache_key(&semantic_key(&end_user_request, "end_user")), + expected + ); + + let preset_request = native_request( + CacheKeyInput { + preset: Some("preset-key".to_owned()), + ..Default::default() + }, + json!({"user_api_key": "k1"}), + ); + assert_eq!( + semantic_key(&preset_request, "end_user").preset.as_deref(), + Some("preset-key") + ); + assert!(semantic_key(&preset_request, "end_user").fields.is_empty()); + } +} diff --git a/litellm-rust/crates/python-bridge/src/cache/request.rs b/litellm-rust/crates/python-bridge/src/cache/request.rs index 0f50f33e6b9..3b4b910c1f0 100644 --- a/litellm-rust/crates/python-bridge/src/cache/request.rs +++ b/litellm-rust/crates/python-bridge/src/cache/request.rs @@ -1,11 +1,11 @@ use std::time::{Duration, SystemTime, UNIX_EPOCH}; -use litellm_cache::{ExactCacheContext, SemanticCacheContext, SemanticCacheScope}; +use litellm_cache::ExactCacheContext; use litellm_cache_response::{CacheControls, CacheKeyInput, ResponseCacheRequest}; use litellm_host_python::from_py; use pyo3::{exceptions::PyValueError, prelude::*}; use serde::Deserialize; -use serde_json::{Map, Value}; +use serde_json::Value; #[derive(Deserialize)] #[serde(deny_unknown_fields)] @@ -14,51 +14,58 @@ struct RequestInput { controls: Option, ttl_seconds: Option, max_age_seconds: Option, - messages: Option>, - input: Option, - metadata: Option>, - scope: Option, + messages: Option, + input: Option, + metadata: Option, + litellm_metadata: Option, + litellm_params: Option, + scope: Option, } -pub(super) fn request( - value: &Bound<'_, PyAny>, -) -> PyResult> { +#[derive(Clone)] +pub(super) struct NativeRequest { + pub(super) key: CacheKeyInput, + pub(super) controls: CacheControls, + pub(super) ttl: Option, + pub(super) max_age: Option, + pub(super) messages: Option, + pub(super) input: Option, + pub(super) metadata: Option, + pub(super) litellm_metadata: Option, + pub(super) litellm_params: Option, + pub(super) scope: Option, +} + +pub(super) fn request(value: &Bound<'_, PyAny>) -> PyResult { let input: RequestInput = from_py(value)?; request_input(input) } -fn request_input(input: RequestInput) -> PyResult> { - let mut request: ResponseCacheRequest = - ResponseCacheRequest::new(input.key); - if let Some(controls) = input.controls { - request.controls = controls; - } - request.context.ttl = input.ttl_seconds.map(duration).transpose()?; - request.context.messages = input.messages.unwrap_or_default(); - request.context.input = input.input; - request.context.metadata = input.metadata.unwrap_or_default(); - request.context.scope = input.scope.unwrap_or_default(); - request.max_age = input.max_age_seconds.map(duration).transpose()?; - Ok(request) +fn request_input(input: RequestInput) -> PyResult { + let controls = input.controls.unwrap_or_else(|| { + ResponseCacheRequest::::new(input.key.clone()).controls + }); + Ok(NativeRequest { + key: input.key, + controls, + ttl: input.ttl_seconds.map(duration).transpose()?, + max_age: input.max_age_seconds.map(duration).transpose()?, + messages: input.messages, + input: input.input, + metadata: input.metadata, + litellm_metadata: input.litellm_metadata, + litellm_params: input.litellm_params, + scope: input.scope, + }) } -pub(super) fn requests( - value: &Bound<'_, PyAny>, -) -> PyResult>> { +pub(super) fn requests(value: &Bound<'_, PyAny>) -> PyResult> { from_py::>(value)? .into_iter() .map(request_input) .collect() } -pub(super) fn exact( - request: &ResponseCacheRequest, -) -> ResponseCacheRequest { - request.clone().with_context(ExactCacheContext { - ttl: request.context.ttl, - }) -} - pub(super) fn duration(seconds: f64) -> PyResult { Duration::try_from_secs_f64(seconds) .map_err(|_| PyValueError::new_err("cache durations must be finite and nonnegative")) diff --git a/litellm-rust/crates/python-bridge/src/cache/semantic.rs b/litellm-rust/crates/python-bridge/src/cache/semantic.rs new file mode 100644 index 00000000000..934de01e721 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/cache/semantic.rs @@ -0,0 +1,175 @@ +use std::collections::VecDeque; + +use litellm_cache::Error; +use litellm_cache_redis_semantic::prompt_from_context; +use litellm_host_python::{Execution, ExecutionBody, ExecutionStep, run_async}; +use pyo3::{ + PyTraverseError, PyVisit, + exceptions::{PyException, PyRuntimeError}, + prelude::*, +}; +use serde_json::Value; + +use super::{ + cache_error, + embedder::{PythonEmbedder, with_prepared_embedding}, + native::NativeResponseCache, + request::{NativeRequest, now}, +}; + +pub(super) enum SemanticOperation { + Lookup(NativeRequest), + Store(NativeRequest, Value), + StoreBatch(VecDeque<(NativeRequest, Value)>), +} + +enum Phase { + Start, + AwaitingEmbedding, + AwaitingBackend, +} + +pub(super) struct SemanticBody { + service: NativeResponseCache, + operation: SemanticOperation, + pending: Option<(NativeRequest, Option)>, + phase: Phase, +} + +impl SemanticBody { + pub(super) fn new(service: NativeResponseCache, operation: SemanticOperation) -> Self { + Self { + service, + operation, + pending: None, + phase: Phase::Start, + } + } + + fn backend_step( + &mut self, + py: Python<'_>, + seed: Result, Error>, + ) -> PyResult { + self.phase = Phase::AwaitingBackend; + let (request, response) = self.pending.take().ok_or_else(|| { + PyRuntimeError::new_err("semantic execution resumed without a pending operation") + })?; + let service = self.service.clone(); + let future = async move { + match response { + None => service.async_lookup(&request, now()).await, + Some(response) => service + .async_store(&request, response, now()) + .await + .map(|_| None), + } + }; + let awaitable = run_async(py, with_prepared_embedding(seed, future), cache_error)?; + Ok(ExecutionStep::Await(awaitable.unbind())) + } +} + +impl ExecutionBody for SemanticBody { + fn resume(&mut self, mut result: Option>>) -> PyResult { + Python::attach(|py| { + loop { + match self.phase { + Phase::Start => { + if result.is_some() { + return Err(PyRuntimeError::new_err( + "semantic execution received a result before starting", + )); + } + if self.pending.is_none() { + match &mut self.operation { + SemanticOperation::Lookup(request) => { + self.pending = Some((request.clone(), None)); + } + SemanticOperation::Store(request, response) => { + let response = std::mem::replace(response, Value::Null); + self.pending = Some((request.clone(), Some(response))); + } + SemanticOperation::StoreBatch(queue) => { + let Some((request, response)) = queue.pop_front() else { + return Ok(ExecutionStep::Return(py.None())); + }; + self.pending = Some((request, Some(response))); + } + } + } + let (request, _) = self.pending.as_ref().ok_or_else(|| { + PyRuntimeError::new_err("semantic execution has no pending operation") + })?; + let semantic = NativeResponseCache::semantic_request(request); + let Some(prompt) = prompt_from_context(&semantic.context) else { + return self.backend_step(py, Err(Error::Unavailable)); + }; + let embedder = self.service.semantic_embedder().ok_or_else(|| { + PyRuntimeError::new_err( + "semantic execution requires a redis-semantic backend", + ) + })?; + let coroutine = embedder.async_embedding_coroutine( + py, + &prompt, + semantic.context.metadata.as_ref(), + )?; + self.phase = Phase::AwaitingEmbedding; + return Ok(ExecutionStep::Await(coroutine)); + } + Phase::AwaitingEmbedding => { + let result = result.take().ok_or_else(|| { + PyRuntimeError::new_err( + "semantic execution expected an embedding result", + ) + })?; + let seed = match result { + Ok(value) => PythonEmbedder::extract(value.into_bound(py)) + .map_err(|_| Error::Unavailable), + Err(error) => { + if !error.is_instance_of::(py) { + return Err(error); + } + Err(Error::Unavailable) + } + }; + return self.backend_step(py, seed); + } + Phase::AwaitingBackend => { + let result = result.take().ok_or_else(|| { + PyRuntimeError::new_err("semantic execution expected a backend result") + })?; + let value = match result { + Ok(value) => value, + Err(error) => return Err(error), + }; + let more = matches!( + &self.operation, + SemanticOperation::StoreBatch(queue) if !queue.is_empty() + ); + if more { + self.phase = Phase::Start; + continue; + } + return Ok(ExecutionStep::Return(value)); + } + } + } + }) + } + + fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + if let Some(embedder) = self.service.semantic_embedder() { + embedder.traverse(visit)?; + } + Ok(()) + } +} + +pub(super) fn drive(py: Python<'_>, body: SemanticBody) -> PyResult> { + let execution = Py::new(py, Execution::new(body))?; + py.import("litellm.rust_bridge.lifecycle")? + .getattr("drive")? + .call1((execution,)) +} diff --git a/litellm-rust/crates/python-bridge/src/cache/semantic_step.rs b/litellm-rust/crates/python-bridge/src/cache/semantic_step.rs new file mode 100644 index 00000000000..24caf3374d6 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/cache/semantic_step.rs @@ -0,0 +1,249 @@ +use std::{sync::Arc, time::Duration}; + +use litellm_cache::SemanticCacheContext; +use litellm_cache_response::{ResponseCache, ResponseCacheCodec, ResponseCacheRequest}; +use litellm_cache_valkey_semantic::{PreparedEmbedding, ValkeySemanticCache, prompt_from_context}; +use litellm_host_python::{Execution, ExecutionBody, ExecutionStep, run_async}; +use pyo3::{PyTraverseError, PyVisit, exceptions::PyRuntimeError, prelude::*}; +use serde_json::Value; + +use super::{cache_error, embedder::PythonEmbedder}; + +pub(super) enum Op { + Lookup, + Store(Value), + StoreBatch(Vec), +} + +#[derive(Clone, Copy)] +enum State { + Start, + AwaitingEmbedding, + AwaitingStorage, + Done, +} + +pub(super) struct SemanticEmbedExecution { + backend: Arc>, + embedder: PythonEmbedder, + requests: Vec>, + op: Op, + now: Option, + prepared: Vec>>, + index: usize, + state: State, +} + +impl SemanticEmbedExecution { + pub(super) fn lookup( + backend: Arc>, + embedder: PythonEmbedder, + request: ResponseCacheRequest, + ) -> Self { + Self { + backend, + embedder, + requests: vec![request], + op: Op::Lookup, + now: None, + prepared: vec![None], + index: 0, + state: State::Start, + } + } + + pub(super) fn store( + backend: Arc>, + embedder: PythonEmbedder, + request: ResponseCacheRequest, + response: Value, + ) -> Self { + Self { + backend, + embedder, + requests: vec![request], + op: Op::Store(response), + now: None, + prepared: vec![None], + index: 0, + state: State::Start, + } + } + + pub(super) fn store_batch( + backend: Arc>, + embedder: PythonEmbedder, + requests: Vec>, + responses: Vec, + ) -> Self { + Self { + backend, + embedder, + prepared: vec![None; requests.len()], + requests, + op: Op::StoreBatch(responses), + now: None, + index: 0, + state: State::Start, + } + } + + fn start(&mut self, py: Python<'_>) -> PyResult { + if self.now.is_none() { + self.now = Some(super::request::now()); + } + while self.index < self.requests.len() { + let request = &self.requests[self.index]; + let enabled = match &self.op { + Op::Lookup => request.controls.reads(), + Op::Store(_) | Op::StoreBatch(_) => request.controls.writes(), + }; + if !enabled { + self.index += 1; + continue; + } + let Some(prompt) = prompt_from_context(&request.context) else { + self.index += 1; + continue; + }; + let metadata = request.context.metadata.clone(); + let awaitable = self + .embedder + .async_embed_awaitable(py, &prompt, &metadata)?; + self.state = State::AwaitingEmbedding; + return Ok(ExecutionStep::Await(awaitable.unbind())); + } + self.state = State::AwaitingStorage; + self.storage_step(py) + } + + fn storage_step(&self, py: Python<'_>) -> PyResult { + let requests = self.requests.clone(); + let prepared = self.prepared.clone(); + let backend = Arc::clone(&self.backend); + let now = self + .now + .ok_or_else(|| PyRuntimeError::new_err("semantic cache timestamp is unavailable"))?; + let awaitable = match &self.op { + Op::Lookup => { + let Some(request) = requests.into_iter().next() else { + return Err(PyRuntimeError::new_err( + "semantic lookup requires one request", + )); + }; + match prepared.into_iter().next().flatten() { + Some(values) => { + let backend = backend.with_embedder(PreparedEmbedding(values)); + let cache = Arc::new(ResponseCache::new(Arc::new(backend))); + run_async( + py, + async move { cache.async_lookup(&request, now).await }, + cache_error, + )? + } + None => { + let cache = Arc::new(ResponseCache::new(backend)); + run_async( + py, + async move { cache.async_lookup(&request, now).await }, + cache_error, + )? + } + } + } + Op::Store(response) => { + let Some(request) = requests.into_iter().next() else { + return Err(PyRuntimeError::new_err( + "semantic store requires one request", + )); + }; + let response = response.clone(); + match prepared.into_iter().next().flatten() { + Some(values) => { + let backend = backend.with_embedder(PreparedEmbedding(values)); + let cache = Arc::new(ResponseCache::new(Arc::new(backend))); + run_async( + py, + async move { cache.async_store(&request, response, now).await }, + cache_error, + )? + } + None => { + let cache = Arc::new(ResponseCache::new(backend)); + run_async( + py, + async move { cache.async_store(&request, response, now).await }, + cache_error, + )? + } + } + } + Op::StoreBatch(responses) => { + let responses = responses.clone(); + run_async( + py, + async move { + for ((request, response), prepared) in + requests.into_iter().zip(responses).zip(prepared) + { + let Some(values) = prepared else { + continue; + }; + let backend = backend.with_embedder(PreparedEmbedding(values)); + let cache = ResponseCache::new(Arc::new(backend)); + cache.async_store(&request, response, now).await?; + } + Ok(()) + }, + cache_error, + )? + } + }; + Ok(ExecutionStep::Await(awaitable.unbind())) + } + + fn resume_py( + &mut self, + py: Python<'_>, + result: Option>>, + ) -> PyResult { + match (self.state, result) { + (State::Start, None) => self.start(py), + (State::AwaitingEmbedding, Some(Ok(value))) => { + let values = value.bind(py).extract::>()?; + self.prepared[self.index] = + Some(values.into_iter().map(|value| value as f32).collect()); + self.index += 1; + self.start(py) + } + (State::AwaitingStorage, Some(Ok(value))) => { + self.state = State::Done; + Ok(ExecutionStep::Return(value)) + } + (_, Some(Err(error))) => Err(error), + _ => Err(PyRuntimeError::new_err( + "invalid semantic cache execution state", + )), + } + } +} + +impl ExecutionBody for SemanticEmbedExecution { + fn resume(&mut self, result: Option>>) -> PyResult { + Python::attach(|py| self.resume_py(py, result)) + } + + fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + self.embedder.traverse(visit) + } +} + +pub(super) fn drive_semantic<'py>( + py: Python<'py>, + body: SemanticEmbedExecution, +) -> PyResult> { + let execution = Py::new(py, Execution::new(body))?; + py.import("litellm.rust_bridge.lifecycle")? + .getattr("drive")? + .call1((execution,)) +} diff --git a/litellm-rust/crates/secrets-hashicorp/Cargo.toml b/litellm-rust/crates/secrets-hashicorp/Cargo.toml new file mode 100644 index 00000000000..c646e02ef09 --- /dev/null +++ b/litellm-rust/crates/secrets-hashicorp/Cargo.toml @@ -0,0 +1,25 @@ +[package] +name = "litellm-secrets-hashicorp" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +litellm-core-utils.workspace = true +litellm-secrets-types.workspace = true +moka.workspace = true +rustify.workspace = true +rustify_derive.workspace = true +serde.workspace = true +serde_json.workspace = true +thiserror.workspace = true +tokio.workspace = true +vaultrs.workspace = true +veil.workspace = true + +[dev-dependencies] +rstest.workspace = true +tempfile = "3" +tokio.workspace = true +wiremock = "0.6.5" diff --git a/litellm-rust/crates/secrets-hashicorp/src/cert_login.rs b/litellm-rust/crates/secrets-hashicorp/src/cert_login.rs new file mode 100644 index 00000000000..f99df62db12 --- /dev/null +++ b/litellm-rust/crates/secrets-hashicorp/src/cert_login.rs @@ -0,0 +1,22 @@ +#[derive(Debug, rustify_derive::Endpoint)] +#[endpoint(path = "/auth/{self.mount}/login", method = "POST")] +pub struct CertLoginRequest { + #[endpoint(skip)] + pub mount: String, + #[endpoint(raw)] + body: Vec, +} + +impl CertLoginRequest { + pub fn new(name: Option<&str>) -> Self { + let body: Vec = match name { + Some(name) => serde_json::to_vec(&serde_json::json!({ "name": name })) + .expect("json object serialization is infallible"), + None => b"{}".to_vec(), + }; + Self { + mount: "cert".to_owned(), + body, + } + } +} diff --git a/litellm-rust/crates/secrets-hashicorp/src/config.rs b/litellm-rust/crates/secrets-hashicorp/src/config.rs new file mode 100644 index 00000000000..d32491199c4 --- /dev/null +++ b/litellm-rust/crates/secrets-hashicorp/src/config.rs @@ -0,0 +1,161 @@ +use std::{path::PathBuf, time::Duration}; + +use litellm_core_utils::settings::Lookup; +use litellm_secrets_types::SecretValue; + +use crate::Error; + +const DEFAULT_ADDRESS: &str = "http://127.0.0.1:8200"; +const DEFAULT_MOUNT: &str = "secret"; +const DEFAULT_APPROLE_MOUNT_PATH: &str = "approle"; +const DEFAULT_REFRESH_INTERVAL: Duration = Duration::from_secs(86400); +const HCP_VAULT_ADDR: &str = "HCP_VAULT_ADDR"; +const HCP_VAULT_TOKEN: &str = "HCP_VAULT_TOKEN"; +const HCP_VAULT_NAMESPACE: &str = "HCP_VAULT_NAMESPACE"; +const HCP_VAULT_LOGIN_NAMESPACE: &str = "HCP_VAULT_LOGIN_NAMESPACE"; +const HCP_VAULT_SECRET_NAMESPACE: &str = "HCP_VAULT_SECRET_NAMESPACE"; +const HCP_VAULT_MOUNT_NAME: &str = "HCP_VAULT_MOUNT_NAME"; +const HCP_VAULT_PATH_PREFIX: &str = "HCP_VAULT_PATH_PREFIX"; +const HCP_VAULT_APPROLE_ROLE_ID: &str = "HCP_VAULT_APPROLE_ROLE_ID"; +const HCP_VAULT_APPROLE_SECRET_ID: &str = "HCP_VAULT_APPROLE_SECRET_ID"; +const HCP_VAULT_APPROLE_MOUNT_PATH: &str = "HCP_VAULT_APPROLE_MOUNT_PATH"; +const HCP_VAULT_CLIENT_CERT: &str = "HCP_VAULT_CLIENT_CERT"; +const HCP_VAULT_CLIENT_KEY: &str = "HCP_VAULT_CLIENT_KEY"; +const HCP_VAULT_CERT_ROLE: &str = "HCP_VAULT_CERT_ROLE"; +const HCP_VAULT_REFRESH_INTERVAL: &str = "HCP_VAULT_REFRESH_INTERVAL"; +const SECRET_MANAGER_REFRESH_INTERVAL: &str = "SECRET_MANAGER_REFRESH_INTERVAL"; + +#[derive(Clone, Debug)] +pub struct AppRoleAuth { + pub role_id: String, + pub secret_id: SecretValue, + pub mount_path: String, +} + +#[derive(Clone, Debug)] +pub struct TlsCertAuth { + pub cert_path: PathBuf, + pub key_path: PathBuf, + pub role: Option, +} + +#[derive(Clone, Debug)] +pub struct HashicorpVaultConfig { + pub address: String, + pub token: Option, + pub namespace: Option, + pub login_namespace: Option, + pub secret_namespace: Option, + pub mount: String, + pub path_prefix: Option, + pub approle: Option, + pub tls_cert: Option, + pub refresh_interval: Duration, +} + +impl HashicorpVaultConfig { + pub fn from_environment(environment: &dyn Lookup) -> Result { + let address: String = environment + .get(HCP_VAULT_ADDR) + .and_then(|value| nonempty(value.trim())) + .map(|value| value.trim_end_matches('/').to_owned()) + .filter(|value| !value.is_empty()) + .unwrap_or_else(|| DEFAULT_ADDRESS.to_owned()); + let token: Option = environment + .get(HCP_VAULT_TOKEN) + .and_then(nonempty) + .map(SecretValue::new); + let namespace: Option = path_component(environment.get(HCP_VAULT_NAMESPACE)); + let login_namespace: Option = + path_component(environment.get(HCP_VAULT_LOGIN_NAMESPACE)); + let secret_namespace: Option = + path_component(environment.get(HCP_VAULT_SECRET_NAMESPACE)); + let mount: String = path_component(environment.get(HCP_VAULT_MOUNT_NAME)) + .unwrap_or_else(|| DEFAULT_MOUNT.to_owned()); + let path_prefix: Option = path_component(environment.get(HCP_VAULT_PATH_PREFIX)); + let approle: Option = match ( + environment + .get(HCP_VAULT_APPROLE_ROLE_ID) + .and_then(nonempty), + environment + .get(HCP_VAULT_APPROLE_SECRET_ID) + .and_then(nonempty) + .map(SecretValue::new), + ) { + (Some(role_id), Some(secret_id)) => Some(AppRoleAuth { + role_id, + secret_id, + mount_path: path_component(environment.get(HCP_VAULT_APPROLE_MOUNT_PATH)) + .unwrap_or_else(|| DEFAULT_APPROLE_MOUNT_PATH.to_owned()), + }), + _ => None, + }; + let tls_cert: Option = match ( + environment.get(HCP_VAULT_CLIENT_CERT).and_then(nonempty), + environment.get(HCP_VAULT_CLIENT_KEY).and_then(nonempty), + ) { + (Some(cert_path), Some(key_path)) => Some(TlsCertAuth { + cert_path: PathBuf::from(cert_path), + key_path: PathBuf::from(key_path), + role: environment.get(HCP_VAULT_CERT_ROLE).and_then(nonempty), + }), + _ => None, + }; + let refresh_interval: Duration = refresh_interval(environment)?; + Ok(Self { + address, + token, + namespace, + login_namespace, + secret_namespace, + mount, + path_prefix, + approle, + tls_cert, + refresh_interval, + }) + } + + pub fn login_namespace(&self) -> Option<&str> { + self.login_namespace + .as_deref() + .or(self.namespace.as_deref()) + } + + pub fn secret_namespace(&self) -> Option<&str> { + self.secret_namespace + .as_deref() + .or(self.namespace.as_deref()) + } +} + +fn nonempty(value: impl AsRef) -> Option { + let value: &str = value.as_ref(); + (!value.is_empty()).then(|| value.to_owned()) +} + +fn path_component(value: Option) -> Option { + value + .and_then(|value| nonempty(value.trim())) + .map(|value| value.trim_matches('/').to_owned()) + .filter(|value| !value.is_empty()) +} + +fn refresh_interval(environment: &dyn Lookup) -> Result { + let value: Option = environment + .get(HCP_VAULT_REFRESH_INTERVAL) + .and_then(nonempty) + .or_else(|| { + environment + .get(SECRET_MANAGER_REFRESH_INTERVAL) + .and_then(nonempty) + }); + let Some(value) = value else { + return Ok(DEFAULT_REFRESH_INTERVAL); + }; + let seconds: i64 = value.parse().map_err(|_| Error::RefreshInterval)?; + if seconds < 0 { + return Err(Error::RefreshInterval); + } + Ok(Duration::from_secs(seconds as u64)) +} diff --git a/litellm-rust/crates/secrets-hashicorp/src/error.rs b/litellm-rust/crates/secrets-hashicorp/src/error.rs new file mode 100644 index 00000000000..e26033af085 --- /dev/null +++ b/litellm-rust/crates/secrets-hashicorp/src/error.rs @@ -0,0 +1,34 @@ +#[derive(thiserror::Error, veil::Redact)] +pub enum Error { + #[error("HashiCorp Vault requires an enterprise license")] + EnterpriseRequired, + #[error("invalid secret name")] + InvalidSecretName(#[from] litellm_secrets_types::Error), + #[error("HashiCorp Vault client failed")] + Client( + #[from] + #[redact] + vaultrs::error::ClientError, + ), + #[error("HashiCorp Vault client settings are invalid: {message}")] + ClientSettings { message: String }, + #[error("HashiCorp Vault TLS identity could not be configured for {path}: {message}")] + TlsIdentity { + path: std::path::PathBuf, + message: String, + }, + #[error("HashiCorp Vault login returned HTTP {status}")] + LoginStatus { status: u16 }, + #[error("HashiCorp Vault login response is malformed")] + MalformedLogin, + #[error("HashiCorp Vault authentication is not configured")] + NoAuthConfigured, + #[error("HashiCorp Vault returned HTTP {status}")] + Status { status: u16 }, + #[error("HashiCorp Vault response payload is malformed")] + MalformedPayload, + #[error("HashiCorp Vault secret value is not a string")] + NonStringValue, + #[error("invalid HashiCorp Vault refresh interval")] + RefreshInterval, +} diff --git a/litellm-rust/crates/secrets-hashicorp/src/lib.rs b/litellm-rust/crates/secrets-hashicorp/src/lib.rs new file mode 100644 index 00000000000..0c2b05647f8 --- /dev/null +++ b/litellm-rust/crates/secrets-hashicorp/src/lib.rs @@ -0,0 +1,10 @@ +#![forbid(unsafe_code)] + +mod cert_login; +mod config; +mod error; +pub mod secret_manager; + +pub use config::{AppRoleAuth, HashicorpVaultConfig, TlsCertAuth}; +pub use error::Error; +pub use secret_manager::{HashicorpVault, SecretLocation}; diff --git a/litellm-rust/crates/secrets-hashicorp/src/secret_manager.rs b/litellm-rust/crates/secrets-hashicorp/src/secret_manager.rs new file mode 100644 index 00000000000..3ad9a549438 --- /dev/null +++ b/litellm-rust/crates/secrets-hashicorp/src/secret_manager.rs @@ -0,0 +1,359 @@ +use std::{ + collections::HashMap, + fmt, + sync::Arc, + time::{Duration, Instant}, +}; + +use litellm_core_utils::settings::Lookup; +use litellm_secrets_types::{ + BaseSecretManager, SecretValue, async_rotate_secret, validate_secret_name, +}; +use moka::future::Cache; +use rustify::errors::ClientError as RustifyClientError; +use serde_json::Value; +use tokio::sync::Mutex; +use vaultrs::{ + api, + auth::approle, + client::{Identity, VaultClient, VaultClientSettingsBuilder}, + error::ClientError, + kv2, +}; + +use crate::{Error, HashicorpVaultConfig, TlsCertAuth, cert_login::CertLoginRequest}; + +const CACHE_CAPACITY: u64 = 200; + +#[derive(Clone)] +struct CachedClient { + client: Arc, + expires_at: Option, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct SecretLocation { + pub namespace: Option, + pub mount: String, + pub path: String, +} + +#[derive(Clone)] +pub struct HashicorpVault { + config: HashicorpVaultConfig, + cache: Cache, + auth_client: Arc>>, +} + +impl fmt::Debug for HashicorpVault { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("HashicorpVault") + .field("config", &self.config) + .finish_non_exhaustive() + } +} + +impl HashicorpVault { + pub fn new( + environment: Arc, + enterprise_enabled: bool, + ) -> Result { + let config: HashicorpVaultConfig = + HashicorpVaultConfig::from_environment(environment.as_ref())?; + Self::from_config(config, enterprise_enabled) + } + + pub fn from_config( + config: HashicorpVaultConfig, + enterprise_enabled: bool, + ) -> Result { + if !enterprise_enabled { + return Err(Error::EnterpriseRequired); + } + let cache: Cache = Cache::builder() + .max_capacity(CACHE_CAPACITY) + .time_to_live(config.refresh_interval) + .build(); + Ok(Self { + config, + cache, + auth_client: Arc::new(Mutex::new(None)), + }) + } + + pub fn secret_location(&self, secret_name: &str) -> Result { + validate_secret_name(secret_name).map_err(Error::InvalidSecretName)?; + let path: String = [ + self.config.path_prefix.clone(), + Some(secret_name.to_owned()), + ] + .into_iter() + .flatten() + .collect::>() + .join("/"); + Ok(SecretLocation { + namespace: self.config.secret_namespace().map(str::to_owned), + mount: self.config.mount.clone(), + path, + }) + } + + pub fn config(&self) -> &HashicorpVaultConfig { + &self.config + } + + pub async fn async_read_secret(&self, secret_name: &str) -> Result, Error> { + let location: SecretLocation = self.secret_location(secret_name)?; + let cache_key: String = cache_key(&location); + if let Some(value) = self.cache.get(&cache_key).await { + return Ok(Some(value)); + } + let client: Arc = self.vault_client().await?; + let data: HashMap = + match kv2::read(client.as_ref(), &location.mount, &location.path).await { + Ok(data) => data, + Err(error) if api_status(&error) == Some(404) => return Ok(None), + Err(error) => return Err(map_api_error(error, ErrorContext::Read)), + }; + let Some(value) = data.get("key") else { + return Ok(None); + }; + let value: &str = value.as_str().ok_or(Error::NonStringValue)?; + let value: SecretValue = SecretValue::new(value); + self.cache.insert(cache_key, value.clone()).await; + Ok(Some(value)) + } + + pub async fn async_write_secret( + &self, + secret_name: &str, + value: SecretValue, + description: Option<&str>, + ) -> Result { + let location: SecretLocation = self.secret_location(secret_name)?; + let cache_key: String = cache_key(&location); + let data: HashMap = match description { + Some(description) => [ + ("key".to_owned(), Value::String(value.expose().to_owned())), + ( + "description".to_owned(), + Value::String(description.to_owned()), + ), + ] + .into_iter() + .collect(), + None => [("key".to_owned(), Value::String(value.expose().to_owned()))] + .into_iter() + .collect(), + }; + let client: Arc = self.vault_client().await?; + let metadata = kv2::set(client.as_ref(), &location.mount, &location.path, &data) + .await + .map_err(|error| map_api_error(error, ErrorContext::Secret))?; + self.cache.invalidate(&cache_key).await; + serde_json::to_value(metadata) + .map_err(|source| Error::Client(ClientError::JsonParseError { source })) + } + + pub async fn async_delete_secret(&self, secret_name: &str) -> Result<(), Error> { + let location: SecretLocation = self.secret_location(secret_name)?; + let cache_key: String = cache_key(&location); + let client: Arc = self.vault_client().await?; + kv2::delete_latest(client.as_ref(), &location.mount, &location.path) + .await + .map_err(|error| map_api_error(error, ErrorContext::Secret))?; + self.cache.invalidate(&cache_key).await; + Ok(()) + } + + pub async fn async_rotate_secret( + &self, + current_name: &str, + new_name: &str, + value: &SecretValue, + ) -> Result { + async_rotate_secret(self, current_name, new_name, value).await + } + + async fn vault_client(&self) -> Result, Error> { + let mut cached = self.auth_client.lock().await; + if let Some(entry) = cached.as_ref() + && entry + .expires_at + .is_none_or(|expires_at| expires_at > Instant::now()) + { + return Ok(entry.client.clone()); + } + + let (client, expires_at): (VaultClient, Option) = + match (self.config.approle.as_ref(), self.config.tls_cert.as_ref()) { + (Some(approle), _) => { + let login_client: VaultClient = + self.build_client(self.config.login_namespace(), "")?; + let auth = approle::login( + &login_client, + &approle.mount_path, + &approle.role_id, + approle.secret_id.expose(), + ) + .await + .map_err(|error| map_api_error(error, ErrorContext::Login))?; + ( + self.build_client(self.config.secret_namespace(), &auth.client_token)?, + token_expiry(auth.lease_duration), + ) + } + (None, Some(tls)) => { + let login_client: VaultClient = + self.build_client(self.config.login_namespace(), "")?; + let endpoint: CertLoginRequest = CertLoginRequest::new(tls.role.as_deref()); + let auth = api::auth(&login_client, endpoint) + .await + .map_err(|error| map_api_error(error, ErrorContext::Login))?; + ( + self.build_client(self.config.secret_namespace(), &auth.client_token)?, + token_expiry(auth.lease_duration), + ) + } + (None, None) => { + let token: SecretValue = + self.config.token.clone().ok_or(Error::NoAuthConfigured)?; + ( + self.build_client(self.config.secret_namespace(), token.expose())?, + None, + ) + } + }; + let client: Arc = Arc::new(client); + *cached = Some(CachedClient { + client: client.clone(), + expires_at, + }); + Ok(client) + } + + fn build_client(&self, namespace: Option<&str>, token: &str) -> Result { + let settings = VaultClientSettingsBuilder::default() + .address(&self.config.address) + .token(token.to_owned()) + .namespace(namespace.map(str::to_owned)) + .identity(identity_for(self.config.tls_cert.as_ref())?) + .ca_certs(Vec::new()) + .verify(true) + .build() + .map_err(|message| Error::ClientSettings { + message: message.to_string(), + })?; + VaultClient::new(settings).map_err(Error::Client) + } +} + +impl BaseSecretManager for HashicorpVault { + type Error = Error; + type WriteResponse = Value; + type DeleteResponse = (); + + async fn async_read_secret(&self, name: &str) -> Result, Error> { + HashicorpVault::async_read_secret(self, name).await + } + + async fn async_write_secret( + &self, + name: &str, + value: &SecretValue, + description: Option<&str>, + ) -> Result { + HashicorpVault::async_write_secret(self, name, value.clone(), description).await + } + + async fn async_delete_secret( + &self, + name: &str, + _recovery_window_in_days: i64, + ) -> Result<(), Error> { + HashicorpVault::async_delete_secret(self, name).await + } +} + +#[derive(Clone, Copy)] +enum ErrorContext { + Login, + Read, + Secret, +} + +fn cache_key(location: &SecretLocation) -> String { + format!( + "{:?}/{}/{}", + location.namespace, location.mount, location.path + ) +} + +fn identity_for(tls: Option<&TlsCertAuth>) -> Result, Error> { + tls.map(|tls| { + let cert: Vec = std::fs::read(&tls.cert_path).map_err(|source| Error::TlsIdentity { + path: tls.cert_path.clone(), + message: source.to_string(), + })?; + let key: Vec = std::fs::read(&tls.key_path).map_err(|source| Error::TlsIdentity { + path: tls.key_path.clone(), + message: source.to_string(), + })?; + Identity::from_pem(&[cert.as_slice(), key.as_slice()].concat()).map_err(|source| { + Error::TlsIdentity { + path: tls.cert_path.clone(), + message: source.to_string(), + } + }) + }) + .transpose() +} + +fn map_api_error(error: ClientError, context: ErrorContext) -> Error { + match error { + ClientError::APIError { code, .. } => match context { + ErrorContext::Login => Error::LoginStatus { status: code }, + ErrorContext::Read | ErrorContext::Secret => Error::Status { status: code }, + }, + ClientError::JsonParseError { source } => match context { + ErrorContext::Login => Error::MalformedLogin, + ErrorContext::Read => Error::MalformedPayload, + ErrorContext::Secret => Error::Client(ClientError::JsonParseError { source }), + }, + ClientError::ResponseEmptyError | ClientError::ResponseDataEmptyError => { + malformed_response(context) + } + ClientError::RestClientError { source } => match source { + RustifyClientError::ServerResponseError { code, .. } => match context { + ErrorContext::Login => Error::LoginStatus { status: code }, + ErrorContext::Read | ErrorContext::Secret => Error::Status { status: code }, + }, + RustifyClientError::ResponseParseError { .. } => malformed_response(context), + source => Error::Client(ClientError::RestClientError { source }), + }, + error => Error::Client(error), + } +} + +fn api_status(error: &ClientError) -> Option { + match error { + ClientError::APIError { code, .. } => Some(*code), + ClientError::RestClientError { + source: RustifyClientError::ServerResponseError { code, .. }, + } => Some(*code), + _ => None, + } +} + +fn malformed_response(context: ErrorContext) -> Error { + match context { + ErrorContext::Login => Error::MalformedLogin, + ErrorContext::Read => Error::MalformedPayload, + ErrorContext::Secret => Error::Client(ClientError::ResponseDataEmptyError), + } +} + +fn token_expiry(lease_duration: u64) -> Option { + (lease_duration > 0).then(|| Instant::now() + Duration::from_secs(lease_duration)) +} diff --git a/litellm-rust/crates/secrets-hashicorp/tests/secret_manager.rs b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager.rs new file mode 100644 index 00000000000..c52db46e41e --- /dev/null +++ b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager.rs @@ -0,0 +1,602 @@ +use std::{collections::HashMap, sync::Arc, time::Duration}; + +use litellm_core_utils::settings::Lookup; +use litellm_secrets_hashicorp::{Error, HashicorpVault, HashicorpVaultConfig}; +use litellm_secrets_types::SecretValue; +use serde::Deserialize; +use serde_json::json; +use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{body_json, header, method, path}, +}; + +fn config(server: &MockServer, values: &[(&str, &str)]) -> HashicorpVaultConfig { + let mut environment_values: HashMap = values + .iter() + .map(|(name, value)| ((*name).to_owned(), (*value).to_owned())) + .collect(); + environment_values.insert("HCP_VAULT_ADDR".to_owned(), server.uri()); + let environment: Arc = + Arc::new(move |name: &str| environment_values.get(name).cloned()); + HashicorpVaultConfig::from_environment(environment.as_ref()).unwrap() +} + +fn manager(server: &MockServer, values: &[(&str, &str)]) -> HashicorpVault { + HashicorpVault::from_config(config(server, values), true).unwrap() +} + +fn auth_response(token: &str, lease_duration: u64) -> serde_json::Value { + json!({ + "auth": { + "client_token": token, + "accessor": "", + "policies": [], + "token_policies": [], + "metadata": null, + "lease_duration": lease_duration, + "renewable": false, + "entity_id": "", + "token_type": "service", + "orphan": false + }, + "lease_id": "", + "lease_duration": lease_duration, + "renewable": false, + "request_id": "", + "warnings": null, + "wrap_info": null + }) +} + +fn read_response(data: serde_json::Value) -> serde_json::Value { + json!({ + "data": { + "data": data, + "metadata": { + "created_time": "", + "deletion_time": "", + "custom_metadata": null, + "destroyed": false, + "version": 1 + } + }, + "lease_id": "", + "lease_duration": 0, + "renewable": false, + "request_id": "", + "warnings": null, + "wrap_info": null + }) +} + +#[tokio::test] +async fn token_reads_use_vault_headers_and_cache_values() { + let server: MockServer = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/name")) + .and(header("X-Vault-Token", "token")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), + ) + .expect(1) + .mount(&server) + .await; + let manager: HashicorpVault = manager(&server, &[("HCP_VAULT_TOKEN", "token")]); + + assert_eq!( + manager + .async_read_secret("name") + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); + let requests = server.received_requests().await.unwrap(); + assert!( + requests + .iter() + .all(|request| !request.headers.contains_key("X-Vault-Namespace")) + ); + assert_eq!( + manager + .async_read_secret("name") + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); +} + +#[tokio::test] +async fn namespace_mount_and_prefix_are_sanitized_in_the_url() { + let server: MockServer = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/kv-prod/data/virtual-keys/name")) + .and(header("X-Vault-Namespace", "team-a")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), + ) + .expect(1) + .mount(&server) + .await; + let manager: HashicorpVault = manager( + &server, + &[ + ("HCP_VAULT_TOKEN", "token"), + ("HCP_VAULT_SECRET_NAMESPACE", " /team-a/ "), + ("HCP_VAULT_MOUNT_NAME", " /kv-prod/ "), + ("HCP_VAULT_PATH_PREFIX", " /virtual-keys/ "), + ], + ); + + let location = manager.secret_location("name").unwrap(); + assert_eq!(location.namespace.as_deref(), Some("team-a")); + assert_eq!(location.mount, "kv-prod"); + assert_eq!(location.path, "virtual-keys/name"); + assert!(manager.async_read_secret("name").await.unwrap().is_some()); +} + +#[test] +fn trailing_address_slashes_are_removed() { + let environment: Arc = Arc::new(|name: &str| match name { + "HCP_VAULT_ADDR" => Some("http://vault.test:8200///".to_owned()), + "HCP_VAULT_TOKEN" => Some("token".to_owned()), + _ => None, + }); + let config: HashicorpVaultConfig = + HashicorpVaultConfig::from_environment(environment.as_ref()).unwrap(); + let manager: HashicorpVault = HashicorpVault::from_config(config, true).unwrap(); + + assert_eq!( + manager.secret_location("name").unwrap(), + litellm_secrets_hashicorp::SecretLocation { + namespace: None, + mount: "secret".to_owned(), + path: "name".to_owned(), + } + ); +} + +#[rstest::rstest] +#[case("-1")] +#[case("not-a-number")] +fn invalid_refresh_intervals_are_rejected(#[case] value: &str) { + let environment: Arc = Arc::new(move |name: &str| match name { + "HCP_VAULT_REFRESH_INTERVAL" => Some(value.to_owned()), + _ => None, + }); + + assert!(matches!( + HashicorpVaultConfig::from_environment(environment.as_ref()), + Err(Error::RefreshInterval) + )); +} + +#[tokio::test] +async fn approle_login_uses_namespace_and_reuses_the_token() { + let server: MockServer = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/auth/custom-approle/login")) + .and(header("X-Vault-Namespace", "login-root")) + .and(body_json(json!({"role_id": "role", "secret_id": "secret"}))) + .respond_with(ResponseTemplate::new(200).set_body_json(auth_response("login-token", 3600))) + .expect(1) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/name")) + .and(header("X-Vault-Token", "login-token")) + .and(header("X-Vault-Namespace", "secret-root")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), + ) + .expect(1) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/name-2")) + .respond_with(ResponseTemplate::new(404).set_body_json(json!({"errors": ["missing"]}))) + .expect(1) + .mount(&server) + .await; + let manager: HashicorpVault = manager( + &server, + &[ + ("HCP_VAULT_APPROLE_ROLE_ID", "role"), + ("HCP_VAULT_APPROLE_SECRET_ID", "secret"), + ("HCP_VAULT_APPROLE_MOUNT_PATH", "custom-approle"), + ("HCP_VAULT_NAMESPACE", "secret-root"), + ("HCP_VAULT_LOGIN_NAMESPACE", "login-root"), + ], + ); + + assert!(manager.async_read_secret("name").await.unwrap().is_some()); + assert!(manager.async_read_secret("name-2").await.unwrap().is_none()); +} + +#[tokio::test] +async fn approle_tokens_expire_after_the_vault_lease() { + let server: MockServer = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/auth/approle/login")) + .respond_with(ResponseTemplate::new(200).set_body_json(auth_response("login-token", 1))) + .expect(2) + .mount(&server) + .await; + Mock::given(method("GET")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), + ) + .expect(2) + .mount(&server) + .await; + let manager: HashicorpVault = manager( + &server, + &[ + ("HCP_VAULT_APPROLE_ROLE_ID", "role"), + ("HCP_VAULT_APPROLE_SECRET_ID", "secret"), + ("HCP_VAULT_REFRESH_INTERVAL", "0"), + ], + ); + + assert!(manager.async_read_secret("first").await.unwrap().is_some()); + tokio::time::sleep(Duration::from_secs(1) + Duration::from_millis(50)).await; + assert!(manager.async_read_secret("second").await.unwrap().is_some()); +} + +#[tokio::test] +async fn tls_login_posts_the_role_and_uses_the_client_identity() { + let server: MockServer = MockServer::start().await; + let directory: tempfile::TempDir = tempfile::tempdir().unwrap(); + let cert_path = directory.path().join("client.crt"); + let key_path = directory.path().join("client.key"); + std::fs::write(&cert_path, TEST_CERTIFICATE).unwrap(); + std::fs::write(&key_path, TEST_PRIVATE_KEY).unwrap(); + Mock::given(method("POST")) + .and(path("/v1/auth/cert/login")) + .and(header("X-Vault-Namespace", "login-ns")) + .respond_with(ResponseTemplate::new(200).set_body_json(auth_response("cert-token", 0))) + .expect(2) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/name")) + .and(header("X-Vault-Token", "cert-token")) + .and(header("X-Vault-Namespace", "secret-ns")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), + ) + .expect(2) + .mount(&server) + .await; + let role_values: HashMap = HashMap::from([ + ("HCP_VAULT_ADDR".to_owned(), server.uri()), + ( + "HCP_VAULT_CLIENT_CERT".to_owned(), + cert_path.to_str().unwrap().to_owned(), + ), + ( + "HCP_VAULT_CLIENT_KEY".to_owned(), + key_path.to_str().unwrap().to_owned(), + ), + ("HCP_VAULT_CERT_ROLE".to_owned(), "vault-role".to_owned()), + ( + "HCP_VAULT_LOGIN_NAMESPACE".to_owned(), + "login-ns".to_owned(), + ), + ( + "HCP_VAULT_SECRET_NAMESPACE".to_owned(), + "secret-ns".to_owned(), + ), + ]); + let role_environment: Arc = + Arc::new(move |name: &str| role_values.get(name).cloned()); + let role_manager: HashicorpVault = HashicorpVault::new(role_environment, true).unwrap(); + assert!( + role_manager + .async_read_secret("name") + .await + .unwrap() + .is_some() + ); + + let no_role_values: HashMap = HashMap::from([ + ("HCP_VAULT_ADDR".to_owned(), server.uri()), + ( + "HCP_VAULT_CLIENT_CERT".to_owned(), + cert_path.to_str().unwrap().to_owned(), + ), + ( + "HCP_VAULT_CLIENT_KEY".to_owned(), + key_path.to_str().unwrap().to_owned(), + ), + ( + "HCP_VAULT_LOGIN_NAMESPACE".to_owned(), + "login-ns".to_owned(), + ), + ( + "HCP_VAULT_SECRET_NAMESPACE".to_owned(), + "secret-ns".to_owned(), + ), + ]); + let no_role_environment: Arc = + Arc::new(move |name: &str| no_role_values.get(name).cloned()); + let no_role_manager: HashicorpVault = HashicorpVault::new(no_role_environment, true).unwrap(); + assert!( + no_role_manager + .async_read_secret("name") + .await + .unwrap() + .is_some() + ); + let login_bodies: Vec = server + .received_requests() + .await + .unwrap() + .iter() + .filter(|request| request.method.as_str() == "POST") + .map(|request| serde_json::from_slice(&request.body).unwrap()) + .collect(); + assert!(login_bodies.contains(&json!({"name": "vault-role"}))); + assert!(login_bodies.contains(&json!({}))); +} + +#[rstest::rstest] +#[case::missing(404, json!({"errors": ["missing"]}), 0)] +#[case::malformed(200, json!({"data": "invalid"}), 1)] +#[case::missing_key(200, json!({}), 0)] +#[case::non_string(200, json!({"key": 1}), 2)] +#[tokio::test] +async fn read_responses_distinguish_absence_and_malformed_payloads( + #[case] status: u16, + #[case] body: serde_json::Value, + #[case] expected: u8, +) { + let server: MockServer = MockServer::start().await; + Mock::given(method("GET")) + .respond_with(ResponseTemplate::new(status).set_body_json( + if status == 200 && expected != 1 { + read_response(body) + } else { + body + }, + )) + .expect(1) + .mount(&server) + .await; + let result: Result, Error> = + manager(&server, &[("HCP_VAULT_TOKEN", "token")]) + .async_read_secret("name") + .await; + match expected { + 0 => assert!(result.unwrap().is_none()), + 1 => assert!(matches!(result, Err(Error::MalformedPayload))), + 2 => assert!(matches!(result, Err(Error::NonStringValue))), + _ => unreachable!(), + } +} + +#[tokio::test] +async fn write_and_delete_invalidate_the_read_cache() { + let server: MockServer = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/name")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), + ) + .expect(2) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/v1/secret/data/name")) + .and(body_json( + json!({"data": {"key": "updated", "description": "description"}}), + )) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "data": { + "created_time": "", + "deletion_time": "", + "custom_metadata": null, + "destroyed": false, + "version": 2 + }, + "lease_id": "", + "lease_duration": 0, + "renewable": false, + "request_id": "", + "warnings": null, + "wrap_info": null + }))) + .expect(1) + .mount(&server) + .await; + Mock::given(method("DELETE")) + .and(path("/v1/secret/data/name")) + .respond_with(ResponseTemplate::new(204)) + .expect(1) + .mount(&server) + .await; + let manager: HashicorpVault = manager(&server, &[("HCP_VAULT_TOKEN", "token")]); + + assert!(manager.async_read_secret("name").await.unwrap().is_some()); + assert!( + manager + .async_write_secret("name", SecretValue::new("updated"), Some("description")) + .await + .is_ok() + ); + assert!(manager.async_read_secret("name").await.unwrap().is_some()); + manager.async_delete_secret("name").await.unwrap(); +} + +#[tokio::test] +async fn no_auth_and_invalid_names_fail_without_requests() { + let server: MockServer = MockServer::start().await; + let manager: HashicorpVault = manager(&server, &[]); + + assert!(matches!( + manager.async_read_secret("name").await, + Err(Error::NoAuthConfigured) + )); + assert!(matches!( + manager.async_read_secret("../name").await, + Err(Error::InvalidSecretName(_)) + )); + assert!(server.received_requests().await.unwrap().is_empty()); +} + +#[tokio::test] +async fn debug_output_redacts_authentication_values() { + let server: MockServer = MockServer::start().await; + let manager: HashicorpVault = + HashicorpVault::from_config(config(&server, &[("HCP_VAULT_TOKEN", "token-value")]), true) + .unwrap(); + let debug: String = format!("{manager:?}"); + assert!(!debug.contains("token-value")); + assert!(!debug.contains("secret-id")); +} + +#[derive(Deserialize)] +struct ParityCase { + env: HashMap, + expected_secret_url: String, + expected_login_url: Option, + expected_login_namespace: Option, + expected_secret_namespace: Option, + secret_name: String, +} + +#[test] +fn configuration_matches_python_parity_fixture() { + let cases: Vec = serde_json::from_str(include_str!(concat!( + env!("CARGO_MANIFEST_DIR"), + "/../../../tests/test_litellm/secret_managers/hashicorp_vault_parity.json" + ))) + .unwrap(); + for case in cases { + let values: HashMap = case.env.clone(); + let environment: Arc = + Arc::new(move |name: &str| values.get(name).cloned()); + let config: HashicorpVaultConfig = + HashicorpVaultConfig::from_environment(environment.as_ref()).unwrap(); + let manager: HashicorpVault = HashicorpVault::from_config(config.clone(), true).unwrap(); + let location = manager.secret_location(&case.secret_name).unwrap(); + let namespace = location + .namespace + .as_deref() + .map(|namespace| format!("{namespace}/")) + .unwrap_or_default(); + assert_eq!( + format!( + "{}/v1/{}{}/data/{}", + config.address, namespace, location.mount, location.path + ), + case.expected_secret_url + ); + let login_url = config.approle.as_ref().map_or_else( + || { + config + .tls_cert + .as_ref() + .map(|_| format!("{}/v1/auth/cert/login", config.address)) + }, + |approle| { + Some(format!( + "{}/v1/auth/{}/login", + config.address, approle.mount_path + )) + }, + ); + assert_eq!(login_url, case.expected_login_url); + assert_eq!( + manager.config().login_namespace(), + case.expected_login_namespace.as_deref() + ); + assert_eq!( + manager.config().secret_namespace(), + case.expected_secret_namespace.as_deref() + ); + } +} + +#[tokio::test] +#[ignore] +async fn live_vault_round_trip() { + let environment: Arc = + Arc::new(litellm_core_utils::settings::ProcessEnvironment); + let manager: HashicorpVault = HashicorpVault::new(environment, true).unwrap(); + let name: String = std::env::var("LITELLM_VAULT_LIVE_SECRET_NAME").unwrap(); + let value: SecretValue = SecretValue::new("native-live-value"); + let location = manager.secret_location(&name).unwrap(); + println!( + "native provenance: {} vaultrs {} {:?} {} {}", + module_path!(), + manager.config().address, + location.namespace, + location.mount, + location.path + ); + manager + .async_write_secret(&name, value.clone(), None) + .await + .unwrap(); + assert_eq!( + manager.async_read_secret(&name).await.unwrap().unwrap(), + value + ); + manager.async_delete_secret(&name).await.unwrap(); + assert!(manager.async_read_secret(&name).await.unwrap().is_none()); +} + +const TEST_CERTIFICATE: &str = "-----BEGIN CERTIFICATE----- +MIIDDzCCAfegAwIBAgIUeMzLFLM/mRbPGbNAew5N2UTscocwDQYJKoZIhvcNAQEL +BQAwFzEVMBMGA1UEAwwMbGl0ZWxsbS10ZXN0MB4XDTI2MDkyMTIwMjA1OVoXDTI2 +MDkyMjIwMjA1OVowFzEVMBMGA1UEAwwMbGl0ZWxsbS10ZXN0MIIBIjANBgkqhkiG +9w0BAQEFAAOCAQ8AMIIBCgKCAQEAveYoSUJXybmkHmQsBfhBcv2Ob5Oy8ejZu+B3 +vTnrPumW4ANi1XXKBSazRGB3fEtAgr+3KhKeHaSKEQeBwJkAEBfdmQv0tpXICwHs +1kFNtU0owy54HVW5/ia+LMszsFcPzVIoMnbUOuiKr9RaV7P+IEFzILPBVuV4DoYH +yocjD3+9QNqokWgNL8LK37JijmNEFVaKFz0X6SyL2VRDlfPWTEBK52Gp/pvDgA6G +eTSfyI+kCm9h5ECTYUAtmatk9WPVS8sWOqV1EXVanFyYBU+mDxoywAS1/6CHeIPh +bNmCOZjPoO9qWBJ7ZyGhOconBigXY8qnlXymev+44IPHrx4urwIDAQABo1MwUTAd +BgNVHQ4EFgQUvaZrZ6HKtbr3ekeZmgy4b5Pq95QwHwYDVR0jBBgwFoAUvaZrZ6HK +tbr3ekeZmgy4b5Pq95QwDwYDVR0TAQH/BAUwAwEB/zANBgkqhkiG9w0BAQsFAAOC +AQEAEejrD8d1qDxW55XxQ4IC31rufoEvDV955jyvh2kALPaN/i5oWsBGI+UAQZna +aaoQXwzlmHrtDUBWl0LztVTUamIleUep2+PLLauqqt43vxppxMX8Jn2mnPO20YE/ +hIzGx0jN/LBG8PDyLSvHdlgjP9ofA4Vg4rTQugdXRgOvlCE/epnH/MADcg9KYJtJ +C1RObCIkL3LcdUbjStJRCY/U/FeWcgyncEPz95OFDkbrlNDajb6o6CkYfouqvhTc +8XlgjjAVKIbAbRgbVu3elsquuFM97x2DzWDjkrMNmDt1FJ9ubK36gL6B3o0UMaoQ +00R7x/eqvH+EkWa/2ekW9lpleQ== +-----END CERTIFICATE----- +"; + +const TEST_PRIVATE_KEY: &str = "-----BEGIN PRIVATE KEY----- +MIIEvQIBADANBgkqhkiG9w0BAQEFAASCBKcwggSjAgEAAoIBAQC95ihJQlfJuaQe +ZCwF+EFy/Y5vk7Lx6Nm74He9Oes+6ZbgA2LVdcoFJrNEYHd8S0CCv7cqEp4dpIoR +B4HAmQAQF92ZC/S2lcgLAezWQU21TSjDLngdVbn+Jr4syzOwVw/NUigydtQ66Iqv +1FpXs/4gQXMgs8FW5XgOhgfKhyMPf71A2qiRaA0vwsrfsmKOY0QVVooXPRfpLIvZ +VEOV89ZMQErnYan+m8OADoZ5NJ/Ij6QKb2HkQJNhQC2Zq2T1Y9VLyxY6pXURdVqc +XJgFT6YPGjLABLX/oId4g+Fs2YI5mM+g72pYEntnIaE5yicGKBdjyqeVfKZ6/7jg +g8evHi6vAgMBAAECggEAGdJjlP6b8Fa5bdaCM/ebcrbuuNZVJVbb0JPHxGfNSLs7 +pE9hj5QaOdQW2Uviw3h6F61ZCzQH4xD+Iy2po5ZKb2XHYKnDB1bboj+LRGER337T +9aJqe9at2VTMVEv3Rdm40NsEk0QcPLxlK16NQFK90gYEUSSQPDAswJDSG2R/zHn+ +vADI907mW/goEJHeLn8PWGlNlSiR6x+5JJtq+GXCzUzVvJYQSCLGxCSl2x2H+0g7 +NhFI0zPpdzNmO/h+yhzaFb6Rp5U8+ZsnZ3qYjQ/03gw1myTDKJt1YaO9JvArnNYX +hcJQQ8Rt0bHhcrZA16bBOpqZlo5pKCicwI/netgN8QKBgQDcFz7AzdJ26sMSV32V +rwrMgIoggt8qDjO1ARwqW35A1TIge0FoW4M4KpsXQGGfT341uU1esXEcyZ/1L/5X +3ql2gX4DbOYLZLWYzZGR2hq33oi8HkhN98QrEwL9emSH8NqYX3Xxja3PrmCrSYJe +Zbnd9TIm2XkxyMoyXJu6M/QvnwKBgQDc4dzqTbxoGEGa5MuJoGmMwPnqgdG9UM5J +eExVnh7osxc2sOdsiPeRjjQTxs9v2kJwctC359OJoo9yGaaJeSghU4LEWJo1sqnA +fzSCLammYvtVAtniyNv5Mxk/6Uimi4NNDKaAKB+m4K2uSn3U9AmY7KPYMGaSbS9W +XSnobjxm8QKBgC8bPpAvvWs8ZhIn7bY659nLbUT2HeO3dHO6UBf0yzn/J6JyHxbB +93zvCZDZc8uQTRgcmCW7XtVlhjoJUqvl+Wlm39zF0xr/LCsPXKfWAb/2/lcdOCaP +8Emz4QD10EyUTYUtcWYJB/mafhBLRH8F0Nlj4J8WDu2L51MOJTqeYhZLAoGAWffN +icocAbJPlo22sdoa4+/+W5yBF8GAJMDRJtZ+9H1t6SLpQHYRkMIBSETkXUTjZvX9 +Ocs9iIQkNW9pO/mTdO+VBfCo71JUfknR02xR+6m5gYjlws/ZeYlssXGN2/hbhNiw +QOcW7Vv6olFJK6Iy/oz0t6wPO3kpnN3Zogi0paECgYEAwo44M1DdYCtV0snhmYM9 +5u0mPfYt5P2SVLXyUbr+vFTfrTL/WKnXIJgbsnj3Gvf+GIZv9tKcXhSNmEHQCYX4 +X3w9iTPddCHuvZ1fpufi2TyArJh0OkoNtLXJHTKrHjf2N+61AQzFiv5WieJrdE+H +qr32PTUuVGPyO9LyTY4/RL0= +-----END PRIVATE KEY----- +"; diff --git a/litellm-rust/crates/secrets/Cargo.toml b/litellm-rust/crates/secrets/Cargo.toml index 962f66c92d1..e5f30025976 100644 --- a/litellm-rust/crates/secrets/Cargo.toml +++ b/litellm-rust/crates/secrets/Cargo.toml @@ -9,6 +9,7 @@ repository.workspace = true default = [] aws = ["dep:litellm-secrets-aws"] google = ["dep:litellm-secrets-google"] +hashicorp = ["dep:litellm-secrets-hashicorp"] azure = ["dep:litellm-secrets-azure"] cyberark = ["dep:litellm-secrets-cyberark"] @@ -16,6 +17,7 @@ cyberark = ["dep:litellm-secrets-cyberark"] litellm-secrets-types.workspace = true litellm-secrets-aws = { workspace = true, optional = true } litellm-secrets-google = { workspace = true, optional = true } +litellm-secrets-hashicorp = { workspace = true, optional = true } litellm-secrets-azure = { workspace = true, optional = true } litellm-secrets-cyberark = { workspace = true, optional = true } litellm-core-utils.workspace = true diff --git a/litellm-rust/crates/secrets/README.md b/litellm-rust/crates/secrets/README.md index 183a39e15bb..0d333c7d116 100644 --- a/litellm-rust/crates/secrets/README.md +++ b/litellm-rust/crates/secrets/README.md @@ -9,3 +9,5 @@ Backend failures propagate by default. To allow fallback during a backend failur `get_secret` preserves value types. `get_secret_str` accepts a string default and rejects boolean or JSON values with `Error::TypeMismatch`. `get_secret_bool` accepts a boolean default and converts strings containing `true` or `false`, ignoring surrounding whitespace and ASCII case. Other strings and JSON values produce `Error::TypeMismatch`. Conversion failures never activate fallback or replace a found value with the default Provider payloads remain strings unless explicitly selecting a field from an AWS primary JSON secret. Google caches only successfully decoded string payloads, so reads have identical values and types before and after caching. Confirmed absence and failed reads are not cached. AWS resource-not-found responses and Google HTTP 404 responses indicate absence. Other provider errors remain errors, and successful responses without the required payload are malformed responses rather than missing secrets + +The HashiCorp Vault backend is enabled with the `hashicorp` feature and reads KV v2 values from `HCP_VAULT_*` environment variables. It supports static tokens, AppRole authentication, and TLS certificate authentication diff --git a/litellm-rust/crates/secrets/src/error.rs b/litellm-rust/crates/secrets/src/error.rs index 7e03f1f8cbf..1be0adc2bf5 100644 --- a/litellm-rust/crates/secrets/src/error.rs +++ b/litellm-rust/crates/secrets/src/error.rs @@ -30,6 +30,9 @@ pub enum Error { #[cfg(feature = "google")] #[error(transparent)] Google(#[from] litellm_secrets_google::Error), + #[cfg(feature = "hashicorp")] + #[error(transparent)] + Hashicorp(#[from] litellm_secrets_hashicorp::Error), #[cfg(feature = "azure")] #[error(transparent)] Azure(#[from] litellm_secrets_azure::Error), diff --git a/litellm-rust/crates/secrets/src/handler.rs b/litellm-rust/crates/secrets/src/handler.rs index 71214ffc9ce..5ab2caa75d4 100644 --- a/litellm-rust/crates/secrets/src/handler.rs +++ b/litellm-rust/crates/secrets/src/handler.rs @@ -13,6 +13,8 @@ pub enum SecretManager { GoogleKms(crate::google::GoogleKms), #[cfg(feature = "google")] GoogleSecretManager(crate::google::GoogleSecretManager), + #[cfg(feature = "hashicorp")] + HashicorpVault(crate::hashicorp::HashicorpVault), #[cfg(feature = "azure")] AzureKeyVault(crate::azure::AzureKeyVault), #[cfg(feature = "cyberark")] @@ -31,6 +33,8 @@ impl SecretManager { Self::GoogleKms(_) => KeyManagementSystem::GoogleKms, #[cfg(feature = "google")] Self::GoogleSecretManager(_) => KeyManagementSystem::GoogleSecretManager, + #[cfg(feature = "hashicorp")] + Self::HashicorpVault(_) => KeyManagementSystem::HashicorpVault, #[cfg(feature = "azure")] Self::AzureKeyVault(_) => KeyManagementSystem::AzureKeyVault, #[cfg(feature = "cyberark")] @@ -86,6 +90,12 @@ pub async fn get_secret_from_manager( .get_secret_from_google_secret_manager(secret_name) .await .map_err(Error::from), + #[cfg(feature = "hashicorp")] + SecretManager::HashicorpVault(client) => client + .async_read_secret(secret_name) + .await + .map(|value| value.map(Secret::String)) + .map_err(Error::from), #[cfg(feature = "azure")] SecretManager::AzureKeyVault(client) => client .get_secret_from_azure_key_vault(secret_name) diff --git a/litellm-rust/crates/secrets/src/lib.rs b/litellm-rust/crates/secrets/src/lib.rs index dec924abdd0..1acb5269e66 100644 --- a/litellm-rust/crates/secrets/src/lib.rs +++ b/litellm-rust/crates/secrets/src/lib.rs @@ -23,3 +23,5 @@ pub use litellm_secrets_azure as azure; pub use litellm_secrets_cyberark as cyberark; #[cfg(feature = "google")] pub use litellm_secrets_google as google; +#[cfg(feature = "hashicorp")] +pub use litellm_secrets_hashicorp as hashicorp; diff --git a/litellm-rust/crates/secrets/tests/handler.rs b/litellm-rust/crates/secrets/tests/handler.rs index fbe13721491..2a8b7070522 100644 --- a/litellm-rust/crates/secrets/tests/handler.rs +++ b/litellm-rust/crates/secrets/tests/handler.rs @@ -105,6 +105,148 @@ async fn google_handler_requires_canonical_base64_and_preserves_plaintext_whites Err(Error::MissingCiphertext) )); } +#[cfg(feature = "hashicorp")] +#[tokio::test] +async fn hashicorp_handler_resolves_found_missing_and_failed_values() { + use std::sync::Arc; + + use litellm_core_utils::settings::Lookup; + use litellm_secrets::{ + Error, FailurePolicy, KeyManagementSettings, SecretManager, SecretManagerState, + SecretResolver, hashicorp::HashicorpVault, hashicorp::HashicorpVaultConfig, + }; + use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{method, path}, + }; + + let found_server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/KEY")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "data": { + "data": {"key": "remote"}, + "metadata": { + "created_time": "", + "deletion_time": "", + "custom_metadata": null, + "destroyed": false, + "version": 1 + } + }, + "lease_id": "", + "lease_duration": 0, + "renewable": false, + "request_id": "", + "warnings": null, + "wrap_info": null + }))) + .mount(&found_server) + .await; + let found_environment: Arc = Arc::new({ + let address = found_server.uri(); + move |name: &str| match name { + "HCP_VAULT_ADDR" => Some(address.clone()), + "HCP_VAULT_TOKEN" => Some("token".into()), + _ => None, + } + }); + let found_config = HashicorpVaultConfig::from_environment(found_environment.as_ref()).unwrap(); + let found_manager = HashicorpVault::from_config(found_config, true).unwrap(); + let found_resolver = SecretResolver::new( + Arc::new(SecretManagerState::new( + SecretManager::HashicorpVault(found_manager), + KeyManagementSettings { + hosted_keys: Some(vec!["KEY".into()]), + ..Default::default() + }, + )), + Arc::new(|_: &str| None), + litellm_secrets::OidcResolver::default(), + ); + assert_eq!( + found_resolver + .get_secret_str("KEY", None) + .await + .unwrap() + .unwrap() + .expose(), + "remote" + ); + + let missing_server = MockServer::start().await; + Mock::given(method("GET")) + .respond_with( + ResponseTemplate::new(404).set_body_json(serde_json::json!({"errors": ["missing"]})), + ) + .mount(&missing_server) + .await; + let missing_environment: Arc = Arc::new({ + let address = missing_server.uri(); + move |name: &str| match name { + "HCP_VAULT_ADDR" => Some(address.clone()), + "HCP_VAULT_TOKEN" => Some("token".into()), + _ => None, + } + }); + let missing_config = + HashicorpVaultConfig::from_environment(missing_environment.as_ref()).unwrap(); + let missing_manager = HashicorpVault::from_config(missing_config, true).unwrap(); + let missing_state = SecretManagerState::new( + SecretManager::HashicorpVault(missing_manager), + KeyManagementSettings { + hosted_keys: Some(vec!["KEY".into()]), + ..Default::default() + }, + ); + let missing = litellm_secrets::get_secret_from_manager( + missing_state.backend().unwrap(), + "KEY", + missing_state.settings().unwrap(), + &|_: &str| None, + ) + .await + .unwrap(); + assert!(missing.is_none()); + + let failed_server = MockServer::start().await; + Mock::given(method("GET")) + .respond_with( + ResponseTemplate::new(500).set_body_json(serde_json::json!({"errors": ["failed"]})), + ) + .mount(&failed_server) + .await; + let failed_environment: Arc = Arc::new({ + let address = failed_server.uri(); + move |name: &str| match name { + "HCP_VAULT_ADDR" => Some(address.clone()), + "HCP_VAULT_TOKEN" => Some("token".into()), + _ => None, + } + }); + let failed_config = + HashicorpVaultConfig::from_environment(failed_environment.as_ref()).unwrap(); + let failed_manager = HashicorpVault::from_config(failed_config, true).unwrap(); + let failed_state = SecretManagerState::new( + SecretManager::HashicorpVault(failed_manager), + KeyManagementSettings { + hosted_keys: Some(vec!["KEY".into()]), + ..Default::default() + }, + ); + let failed_resolver = SecretResolver::new( + Arc::new(failed_state), + Arc::new(|_: &str| None), + litellm_secrets::OidcResolver::default(), + ) + .with_failure_policy(FailurePolicy::Propagate); + assert!(matches!( + failed_resolver.get_secret_str("KEY", None).await, + Err(Error::Hashicorp( + litellm_secrets::hashicorp::Error::Status { status: 500 } + )) + )); +} #[cfg(feature = "azure")] #[tokio::test] diff --git a/litellm/constants.py b/litellm/constants.py index 895abf75047..1012ca831f2 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1749,6 +1749,12 @@ SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS: Final = float( SPEND_LOG_CLEANUP_RUN_BUDGET_SECONDS: Final = float(os.getenv("SPEND_LOG_CLEANUP_RUN_BUDGET_SECONDS", "300")) SPEND_LOG_CLEANUP_BATCH_TIMEOUT_SECONDS: Final = float(os.getenv("SPEND_LOG_CLEANUP_BATCH_TIMEOUT_SECONDS", "30")) SPEND_LOG_CLEANUP_REMAINING_COUNT_CAP: Final = int(os.getenv("SPEND_LOG_CLEANUP_REMAINING_COUNT_CAP", "100000")) +SCHEDULED_JOB_SHUTDOWN_FINISH_TIMEOUT_SECONDS: Final = float( + os.getenv("SCHEDULED_JOB_SHUTDOWN_FINISH_TIMEOUT_SECONDS", "5") +) +SCHEDULED_JOB_SHUTDOWN_CANCEL_TIMEOUT_SECONDS: Final = float( + os.getenv("SCHEDULED_JOB_SHUTDOWN_CANCEL_TIMEOUT_SECONDS", "5") +) TOOL_SPEND_TOP_TOOLS: Final = 100 SPEND_LOG_PARTITION_INTERVAL: Final = os.getenv("SPEND_LOG_PARTITION_INTERVAL", "day") SPEND_LOG_PARTITION_PRECREATE_AHEAD: Final = int(os.getenv("SPEND_LOG_PARTITION_PRECREATE_AHEAD", 7)) @@ -1860,6 +1866,12 @@ LITELLM_SETTINGS_SAFE_DB_OVERRIDES: Final = [ "max_ui_session_budget", "budget_rollover", "mcp_tool_search", + "turn_off_message_logging", + "datadog_params", + "datadog_llm_observability_params", + "newrelic_params", + "pointfive_params", + "aws_sqs_callback_params", ] SPECIAL_LITELLM_AUTH_TOKEN: Final = ["ui-token"] DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60)) diff --git a/litellm/litellm_core_utils/logging_worker.py b/litellm/litellm_core_utils/logging_worker.py index 5ccc5632646..2f8e7bdccea 100644 --- a/litellm/litellm_core_utils/logging_worker.py +++ b/litellm/litellm_core_utils/logging_worker.py @@ -484,9 +484,13 @@ class LoggingWorker: so it correctly handles items that have been dequeued but whose callback hasn't finished yet — ``queue.empty()`` would return True in that window and cause us to skip the wait. + + ``start()`` runs first so a queue left behind by a previous event loop + is carried onto this one and drained here instead of joined forever. """ if self._queue is None: return + self.start() await self._queue.join() async def clear_queue(self): diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index beb90c2fb8c..3ca9958ad0d 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -43037,21 +43037,21 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro": { - "input_cost_per_token": 8.92272e-07, + "input_cost_per_token": 8.83746e-07, "input_cost_per_token_cache_hit": 4.4e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 1.784544e-06, + "output_cost_per_token": 1.767492e-06, "source": "https://openrouter.ai/api/v1/models", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "cache_read_input_token_cost": 7.4356e-08, + "cache_read_input_token_cost": 7.36455e-08, "supports_audio_input": false, "supports_pdf_input": false, "supports_vision": false, @@ -70299,8 +70299,8 @@ "supports_web_search": true }, "openrouter/meta-llama/llama-4-maverick": { - "input_cost_per_token": 2e-07, - "output_cost_per_token": 8e-07, + "input_cost_per_token": 1.875e-07, + "output_cost_per_token": 6.525e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 16384, @@ -73794,6 +73794,7 @@ "supports_web_search": false }, "openrouter/bytedance-seed/seed-1.6": { + "deprecation_date": "2026-11-11", "input_cost_per_token": 2.5e-07, "input_cost_per_token_above_128k_tokens": 5e-07, "litellm_provider": "openrouter", @@ -73815,6 +73816,7 @@ "supports_web_search": false }, "openrouter/bytedance-seed/seed-1.6-flash": { + "deprecation_date": "2026-11-11", "input_cost_per_token": 7.5e-08, "input_cost_per_token_above_128k_tokens": 1e-07, "litellm_provider": "openrouter", @@ -73855,6 +73857,7 @@ "supports_web_search": false }, "openrouter/bytedance-seed/seed-2.0-code": { + "deprecation_date": "2026-11-11", "input_cost_per_token": 5e-07, "input_cost_per_token_above_128k_tokens": 1e-06, "litellm_provider": "openrouter", @@ -76876,6 +76879,26 @@ "supports_vision": true, "supports_web_search": true }, + "moonshotai.kimi-k3": { + "cache_creation_input_token_cost": 3.75e-06, + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1000000, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true + }, "global.moonshotai.kimi-k3": { "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, @@ -76975,5 +76998,51 @@ "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false + }, + "xiaomi_mimo/mimo-v2.6-pro": { + "cache_read_input_token_cost": 3.6e-09, + "input_cost_per_token": 4.35e-07, + "litellm_provider": "xiaomi_mimo", + "max_input_tokens": 1048576, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 8.7e-07, + "source": "https://platform.xiaomimimo.com/static/docs/price/pay-as-you-go.md", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_audio_input": true, + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true + }, + "xiaomi_mimo/mimo-v2.6-flash": { + "cache_read_input_token_cost": 2.8e-09, + "input_cost_per_token": 1.4e-07, + "litellm_provider": "xiaomi_mimo", + "max_input_tokens": 1048576, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 2.8e-07, + "source": "https://platform.xiaomimimo.com/static/docs/price/pay-as-you-go.md", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_audio_input": true, + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true } } diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 996911cdfaa..010a1b4536e 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -14,7 +14,7 @@ import hashlib import os import re import time -from collections.abc import Awaitable, Callable, Mapping, Sequence +from collections.abc import Awaitable, Callable, Collection, Mapping, Sequence from dataclasses import dataclass from typing import Any, Final, Literal, NoReturn, Protocol, TypeVar, cast @@ -1602,6 +1602,7 @@ class JWTAuthManager: team_object: LiteLLM_TeamTable | None, route: str, request_method: str | None = None, + team_allowed_routes: Collection[str] = (), ) -> bool: normalized_request_method: Final = request_method.upper() if isinstance(request_method, str) else None if not RouteChecks.is_auth_enforced_pass_through_route( @@ -1610,8 +1611,11 @@ class JWTAuthManager: ): return True + if RouteChecks.jwt_team_routes_grant_pass_through(route=route, team_allowed_routes=team_allowed_routes): + return True + # JWT team selection is team-scoped; key metadata is not available here, - # so passthrough access is granted only by the selected team's metadata. + # so beyond the JWT config grant above, only the selected team's metadata grants access. return RouteChecks.check_passthrough_route_access( route=route, user_api_key_dict=UserAPIKeyAuth(team_metadata=(team_object.metadata or {}) if team_object else {}), @@ -1689,6 +1693,7 @@ class JWTAuthManager: team_object=team_object, route=route, request_method=request_method, + team_allowed_routes=jwt_handler.litellm_jwtauth.team_allowed_routes, ): is_allowed = False denied_auth_enforced_pass_through_route = True @@ -2584,6 +2589,7 @@ class JWTAuthManager: team_object=team_object, route=route, request_method=request_method, + team_allowed_routes=handler.litellm_jwtauth.team_allowed_routes, ): JWTAuthManager._raise_team_passthrough_route_denial(route=route) @@ -2653,6 +2659,7 @@ class JWTAuthManager: team_object=team_object, route=route, request_method=request_method, + team_allowed_routes=handler.litellm_jwtauth.team_allowed_routes, ): JWTAuthManager._raise_team_passthrough_route_denial(route=route) elif team_id is None: diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index afddb4866a9..38189a2d07b 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -278,7 +278,11 @@ class RouteChecks: route=route, method=RouteChecks._get_request_method(request=request), ): - RouteChecks._require_auth_pass_through_access(route=route, valid_token=valid_token) + RouteChecks._require_auth_pass_through_access( + route=route, + valid_token=valid_token, + jwt_team_allowed_routes=RouteChecks._jwt_team_allowed_routes(valid_token=valid_token), + ) elif RouteChecks.is_llm_api_route(route=route): pass elif RouteChecks.is_info_route(route=route): @@ -689,16 +693,43 @@ class RouteChecks: ), ) + @staticmethod + def jwt_team_routes_grant_pass_through(route: str, team_allowed_routes: Collection[str]) -> bool: + """ + Explicit paths and trailing-wildcard prefixes grant auth=true pass-through. Blanket grants never do: + a named route group like ``openai_routes`` is only ever compared as a path, and an entry that names + no path segment (``*``, ``/*``) is skipped. + """ + return any( + RouteChecks.route_matches_wildcard_pattern(route=route, pattern=allowed_route) + for allowed_route in team_allowed_routes + if allowed_route.rstrip("*").strip("/") + ) + + @staticmethod + def _jwt_team_allowed_routes(valid_token: UserAPIKeyAuth) -> Collection[str]: + """``team_allowed_routes`` for team tokens built by JWT auth; JWT-mapped virtual keys stay key-scoped.""" + if valid_token.jwt_claims is None or valid_token.token is not None or valid_token.team_id is None: + return () + + from litellm.proxy.proxy_server import jwt_handler + + return jwt_handler.litellm_jwtauth.team_allowed_routes + @staticmethod def _require_auth_pass_through_access( route: str, valid_token: UserAPIKeyAuth, + jwt_team_allowed_routes: Collection[str] = (), ) -> None: """ - Require an explicit ``allowed_passthrough_routes`` match for auth=true pass-through. + Require an explicit grant for auth=true pass-through: ``allowed_passthrough_routes`` on the + key or team, or an explicit JWT ``team_allowed_routes`` entry. """ if RouteChecks.check_passthrough_route_access(route=route, user_api_key_dict=valid_token): return + if RouteChecks.jwt_team_routes_grant_pass_through(route=route, team_allowed_routes=jwt_team_allowed_routes): + return raise RouteChecks._auth_pass_through_denied_exception(route=route) @staticmethod diff --git a/litellm/proxy/config_resolvers/__init__.py b/litellm/proxy/config_resolvers/__init__.py index eee760df458..f031a8d74b9 100644 --- a/litellm/proxy/config_resolvers/__init__.py +++ b/litellm/proxy/config_resolvers/__init__.py @@ -5,6 +5,17 @@ from litellm.proxy.config_resolvers._descriptors import ( FieldSource, resolve_fields, ) -from litellm.proxy.config_resolvers.settings_store import SettingsStore, config_ownership_message +from litellm.proxy.config_resolvers.settings_store import ( + SettingsStore, + config_ownership_message, + source_for, +) -__all__ = ("FieldDescriptor", "FieldSource", "SettingsStore", "config_ownership_message", "resolve_fields") +__all__ = ( + "FieldDescriptor", + "FieldSource", + "SettingsStore", + "config_ownership_message", + "resolve_fields", + "source_for", +) diff --git a/litellm/proxy/config_resolvers/settings_store.py b/litellm/proxy/config_resolvers/settings_store.py index 291000b3b6a..f05af3de03a 100644 --- a/litellm/proxy/config_resolvers/settings_store.py +++ b/litellm/proxy/config_resolvers/settings_store.py @@ -176,3 +176,10 @@ class SettingsStore(MutableMapping[str, JsonValue]): def _resolution_for(self, key: str) -> Resolved: yaml_value: Final[SettingValue] = self._yaml_values.get(key, ABSENT) return resolve(yaml_value, self._db_value(key)) + + +def source_for(settings: SettingsStore, key: str, default: object = None) -> FieldSource: + source: Final = settings.source(key) + if source == "unset": + return "default" if default is not None else "unset" + return source diff --git a/litellm/proxy/db/autorouter_session_rollup.py b/litellm/proxy/db/autorouter_session_rollup.py index 0d812ee812a..dd08cfd1bef 100644 --- a/litellm/proxy/db/autorouter_session_rollup.py +++ b/litellm/proxy/db/autorouter_session_rollup.py @@ -4,7 +4,7 @@ Per-session auto-router benchmarks rollup. At request time the spend writer builds one AutoRouterTurnTransaction per successful auto-routed request (a request whose metadata carries a routing_decision) and queues it on the prisma client. The spend-log flush job drains the queue into -LiteLLM_AutoRouterSession with one conditional upsert per turn: the statement classifies +key and user session rollups with one atomic statement per turn: each upsert classifies the turn (same model, first visit, return to a model the session already used, out of order) against the row's own columns, so nothing is read before the write and concurrent pods compose. The benchmarks endpoint aggregates these rows and never touches @@ -35,10 +35,27 @@ if TYPE_CHECKING: CACHE_TTL_5M_SECONDS: Final = 300 CACHE_TTL_1H_SECONDS: Final = 3600 -AUTOROUTER_BENCHMARKS_SQL: Final = """ +_SESSION_COLUMNS: Final = """ + api_key, session_id, router_name, router_type, first_turn_at, last_turn_at, + last_model, models, turns, unordered_turns, covered_turns, cache_hits, + same_model_turns, same_model_hits, first_visit_turns, first_visit_hits, + return_turns, return_hits, return_expired_misses, return_within_ttl_misses, + ttl_5m_turns, ttl_1h_turns, total_tokens, spend, saved_spend, classifier_cost, classifier_cost_recorded_turns, tier_turns, + baseline_models, savings_estimated_turns, savings_estimated_actual_spend, savings_estimated_saved_spend, + savings_estimated_baseline_models +""" + +AUTOROUTER_BENCHMARKS_SQL: Final = f""" WITH windowed AS ( - SELECT * FROM "LiteLLM_AutoRouterSession" - WHERE last_turn_at >= $1::timestamp + SELECT {_SESSION_COLUMNS} FROM "LiteLLM_AutoRouterSession" + WHERE $4::text IS NULL + AND last_turn_at >= $1::timestamp + AND first_turn_at < $2::timestamp + AND ($3::text IS NULL OR api_key = $3::text) + UNION ALL + SELECT {_SESSION_COLUMNS} FROM "LiteLLM_AutoRouterUserSession" + WHERE (($4::text IS NOT NULL AND user_id = $4::text) OR ($4::text IS NULL AND api_key = '')) + AND last_turn_at >= $1::timestamp AND first_turn_at < $2::timestamp AND ($3::text IS NULL OR api_key = $3::text) ), @@ -53,7 +70,7 @@ tier_maps AS ( ) SELECT agg.*, - COALESCE(tier_maps.tier_turns, '{}'::jsonb) AS tier_turns + COALESCE(tier_maps.tier_turns, '{{}}'::jsonb) AS tier_turns FROM ( SELECT router_name, @@ -111,6 +128,7 @@ class AutoRouterTurnTransaction: savings_estimated_turns: int = 0 savings_estimated_actual_spend: float = 0.0 savings_estimated_saved_spend: float = 0.0 + user_id: str = "" class TurnCacheFacts(NamedTuple): @@ -214,10 +232,11 @@ def build_autorouter_turn_transaction( if not isinstance(routing_decision, Mapping) or not routing_decision: return None router_name: Final = routing_decision.get("router_model_name") or payload.get("model_group") - api_key: Final = payload.get("api_key") + api_key: Final = payload.get("api_key") or "" + user_id: Final = payload.get("user") or "" session_id: Final = payload.get("session_id") model: Final = payload.get("model") - if not (isinstance(router_name, str) and router_name and api_key and session_id and model): + if not (isinstance(router_name, str) and router_name and (api_key or user_id) and session_id and model): return None turn_at: Final = _turn_time_utc(str(payload.get("startTime") or "")) if turn_at is None: @@ -236,6 +255,7 @@ def build_autorouter_turn_transaction( estimated_savings: Final = recorded_estimated_autorouter_savings(metadata) return AutoRouterTurnTransaction( api_key=api_key, + user_id=user_id, session_id=bounded_session_id(session_id), router_name=router_name, router_type=str(routing_decision.get("router_type") or "unknown"), @@ -293,18 +313,18 @@ _RETURN_MISS: Final = ( _IDLE_SECONDS: Final = f"EXTRACT(EPOCH FROM {_TURN_AT}::timestamp) - (t.models -> {_MODEL} ->> 'at')::float8" _CACHE_TOUCHED: Final = f"{_TOUCHED}::int = 1" -UPSERT_AUTOROUTER_SESSION_SQL: Final = f""" -INSERT INTO "LiteLLM_AutoRouterSession" AS t ( - api_key, session_id, router_name, router_type, first_turn_at, last_turn_at, - last_model, models, turns, unordered_turns, covered_turns, cache_hits, - same_model_turns, same_model_hits, first_visit_turns, first_visit_hits, - return_turns, return_hits, return_expired_misses, return_within_ttl_misses, - ttl_5m_turns, ttl_1h_turns, total_tokens, spend, saved_spend, classifier_cost, classifier_cost_recorded_turns, tier_turns, - baseline_models, savings_estimated_turns, savings_estimated_actual_spend, savings_estimated_saved_spend, - savings_estimated_baseline_models + +def _session_upsert_sql(*, user_scoped: bool) -> str: + table_name: Final = "LiteLLM_AutoRouterUserSession" if user_scoped else "LiteLLM_AutoRouterSession" + user_column: Final = "user_id, " if user_scoped else "" + user_value: Final = f"{_p('user_id')}::text, " if user_scoped else "" + required_identity: Final = _p("user_id" if user_scoped else "api_key") + return f""" +INSERT INTO "{table_name}" AS t ( + {user_column}{_SESSION_COLUMNS} ) -VALUES ( - {_p("api_key")}, {_p("session_id")}, {_p("router_name")}, {_p("router_type")}, {_TURN_AT}::timestamp, {_TURN_AT}::timestamp, +SELECT + {user_value}{_p("api_key")}, {_p("session_id")}, {_p("router_name")}, {_p("router_type")}, {_TURN_AT}::timestamp, {_TURN_AT}::timestamp, {_MODEL}, jsonb_build_object({_MODEL}, jsonb_build_object('at', EXTRACT(EPOCH FROM {_TURN_AT}::timestamp), 'ttl', {_CACHE_TTL}::int)), 1, 0, {_COVERED}::int, {_CACHE_HIT}::int, 0, 0, 1, {_CACHE_HIT}::int, @@ -315,8 +335,8 @@ VALUES ( {_p("classifier_cost")}::float8, 1, {_TIER_DELTA}, {_BASELINE_DELTA}, {_p("savings_estimated_turns")}::int, {_p("savings_estimated_actual_spend")}::float8, {_p("savings_estimated_saved_spend")}::float8, {_ESTIMATED_BASELINE_DELTA} -) -ON CONFLICT (api_key, session_id, router_name) DO UPDATE SET +WHERE {required_identity}::text <> '' +ON CONFLICT ({user_column}api_key, session_id, router_name) DO UPDATE SET turns = t.turns + 1, total_tokens = t.total_tokens + EXCLUDED.total_tokens, spend = t.spend + EXCLUDED.spend, @@ -365,6 +385,17 @@ ON CONFLICT (api_key, session_id, router_name) DO UPDATE SET """ +UPSERT_AUTOROUTER_SESSION_SQL: Final = f""" +WITH key_rollup AS ( + {_session_upsert_sql(user_scoped=False)} + RETURNING 1 +) +{_session_upsert_sql(user_scoped=True)} +""" + +UPSERT_AUTOROUTER_USER_SESSION_SQL: Final = _session_upsert_sql(user_scoped=True) + + def _as_sql_param(value: str | float | bool | datetime | None) -> str | float | None: if isinstance(value, bool): return int(value) @@ -377,18 +408,23 @@ def _upsert_params(transaction: AutoRouterTurnTransaction) -> tuple[str | float return tuple(_as_sql_param(getattr(transaction, name)) for name in _UPSERT_PARAM_FIELDS) -async def write_autorouter_turn(db: SupportsExecuteRaw, transaction: AutoRouterTurnTransaction) -> None: - await db.execute_raw(UPSERT_AUTOROUTER_SESSION_SQL, *_upsert_params(transaction)) +async def write_autorouter_turn( + db: SupportsExecuteRaw, + transaction: AutoRouterTurnTransaction, + statement: str = UPSERT_AUTOROUTER_SESSION_SQL, +) -> None: + await db.execute_raw(statement, *_upsert_params(transaction)) async def _upsert_turn_with_retry( prisma_client: PrismaClient, transaction: AutoRouterTurnTransaction, n_retry_times: int, + statement: str, ) -> None: for attempt in range(n_retry_times + 1): try: - await write_autorouter_turn(prisma_client.db, transaction) + await write_autorouter_turn(prisma_client.db, transaction, statement) except DB_RETRY_SAFE_ERROR_TYPES: if attempt >= n_retry_times: raise @@ -397,6 +433,58 @@ async def _upsert_turn_with_retry( return +def _session_partition(transaction: AutoRouterTurnTransaction) -> tuple[str, str, str, str]: + identity: Final = ("key", transaction.api_key) if transaction.api_key else ("user", transaction.user_id) + return (*identity, transaction.session_id, transaction.router_name) + + +async def _drain_session_partition( + prisma_client: PrismaClient, + transactions: tuple[AutoRouterTurnTransaction, ...], + n_retry_times: int, + statement: str, +) -> tuple[AutoRouterTurnTransaction, ...]: + for position, transaction in enumerate(transactions): + try: + await _upsert_turn_with_retry(prisma_client, transaction, n_retry_times, statement) + except Exception as flush_err: # noqa: BLE001 # stop dependent turns without retrying an ambiguous write + verbose_proxy_logger.error( + "Spend tracking - auto-router session rollup flush failed for router %s; " + "%s of %s turn writes stopped in this partition: %s", + transaction.router_name, + len(transactions) - position, + len(transactions), + flush_err, + ) + return transactions[position:] + return () + + +async def _flush_session_partition( + prisma_client: PrismaClient, + transactions: tuple[AutoRouterTurnTransaction, ...], + n_retry_times: int, +) -> None: + failed_suffix: Final = await _drain_session_partition( + prisma_client, transactions, n_retry_times, UPSERT_AUTOROUTER_SESSION_SQL + ) + if not failed_suffix or not failed_suffix[0].api_key: + return + failed_user: Final = failed_suffix[0].user_id + other_users: Final = sorted( + ( + transaction + for transaction in failed_suffix[1:] + if transaction.user_id and transaction.user_id != failed_user + ), + key=lambda transaction: transaction.user_id, + ) + for _, user_turns in groupby(other_users, key=lambda transaction: transaction.user_id): + await _drain_session_partition( + prisma_client, tuple(user_turns), n_retry_times, UPSERT_AUTOROUTER_USER_SESSION_SQL + ) + + async def flush_autorouter_turn_transactions( prisma_client: PrismaClient, transactions: Sequence[AutoRouterTurnTransaction], @@ -407,38 +495,20 @@ async def flush_autorouter_turn_transactions( Statements run sequentially in per-session event order: a turn's classification depends on the turns before it, and Postgres rejects one multi-row INSERT touching the same key twice. Only ConnectError is retried, per statement, because it proves - that statement never reached the database. Any other failure drops the remaining - turns of THAT session only, with an error log, and the flush continues with the - next session: sessions are independent state machines, so one poisoned statement - must not discard unrelated sessions, and a repeated increment is worse than an - undercount. Callers must not add their own retry around this function. + that statement never reached the database. A failed write stops its key and user + histories for this batch. Other users sharing that key can still advance their + independent user histories, with the key projection disabled and the real key + identity preserved. The failed turn is never replayed. Callers must not add their + own retry around this function. """ if not transactions: return ordered: Final = sorted( transactions, - key=lambda transaction: ( - transaction.api_key, - transaction.session_id, - transaction.router_name, - transaction.turn_at, - ), + key=lambda transaction: (*_session_partition(transaction), transaction.turn_at), ) - for session_key, session_group in groupby( + for _, session_group in groupby( ordered, - key=lambda transaction: (transaction.api_key, transaction.session_id, transaction.router_name), + key=_session_partition, ): - session_turns = tuple(session_group) - for position, transaction in enumerate(session_turns): - try: - await _upsert_turn_with_retry(prisma_client, transaction, n_retry_times) - except Exception as flush_err: # noqa: BLE001 # a statement failure drops only its session's remainder by design - verbose_proxy_logger.error( - "Spend tracking - auto-router session rollup flush failed for router %s; " - "%s of %s turn transactions dropped for one session: %s", - session_key[2], - len(session_turns) - position, - len(session_turns), - flush_err, - ) - break + await _flush_session_partition(prisma_client, tuple(session_group), n_retry_times) diff --git a/litellm/proxy/db/baseline_accounting.py b/litellm/proxy/db/baseline_accounting.py index 8622cb9e481..4219102d9aa 100644 --- a/litellm/proxy/db/baseline_accounting.py +++ b/litellm/proxy/db/baseline_accounting.py @@ -171,6 +171,7 @@ class _Change(BaseModel): request_id: str publication: BaselinePublication api_key: str + user_id: str = "" session_id: str router_name: str baseline_model: str @@ -256,42 +257,54 @@ SET publication = x.publication::text FROM jsonb_to_recordset($1::jsonb) AS x(request_id text, publication jsonb) WHERE observations.request_id = x.request_id """ -_UPDATE_SESSIONS: Final = """ + + +def _session_correction_sql(*, user_scoped: bool) -> str: + table_name: Final = "LiteLLM_AutoRouterUserSession" if user_scoped else "LiteLLM_AutoRouterSession" + identity_columns: Final = ("user_id, " if user_scoped else "") + "api_key, session_id, router_name" + user_filter: Final = "WHERE user_id <> ''" if user_scoped else "" + user_match: Final = "session.user_id = totals.user_id AND " if user_scoped else "" + return f""" WITH changes AS ( SELECT * FROM jsonb_to_recordset($1::jsonb) AS x( - api_key text, session_id text, router_name text, baseline_model text, + user_id text, api_key text, session_id text, router_name text, baseline_model text, covered_delta int, actual_delta float8, savings_delta float8 ) + {user_filter} ), totals AS ( - SELECT api_key, session_id, router_name, SUM(covered_delta)::int AS covered_delta, + SELECT {identity_columns}, SUM(covered_delta)::int AS covered_delta, SUM(actual_delta) AS actual_delta, SUM(savings_delta) AS savings_delta - FROM changes GROUP BY api_key, session_id, router_name + FROM changes GROUP BY {identity_columns} ), models AS ( - SELECT api_key, session_id, router_name, jsonb_object_agg(baseline_model, delta) AS deltas + SELECT {identity_columns}, jsonb_object_agg(baseline_model, delta) AS deltas FROM ( - SELECT api_key, session_id, router_name, baseline_model, SUM(covered_delta)::int AS delta - FROM changes GROUP BY api_key, session_id, router_name, baseline_model - ) grouped GROUP BY api_key, session_id, router_name + SELECT {identity_columns}, baseline_model, SUM(covered_delta)::int AS delta + FROM changes GROUP BY {identity_columns}, baseline_model + ) grouped GROUP BY {identity_columns} ) -UPDATE "LiteLLM_AutoRouterSession" AS session +UPDATE "{table_name}" AS session SET saved_spend = session.saved_spend + totals.savings_delta, savings_estimated_turns = session.savings_estimated_turns + totals.covered_delta, savings_estimated_actual_spend = session.savings_estimated_actual_spend + totals.actual_delta, savings_estimated_saved_spend = session.savings_estimated_saved_spend + totals.savings_delta, savings_estimated_baseline_models = ( - SELECT COALESCE(jsonb_object_agg(key, value), '{}'::jsonb) FROM ( + SELECT COALESCE(jsonb_object_agg(key, value), '{{}}'::jsonb) FROM ( SELECT key, SUM(value::int)::int AS value FROM ( SELECT * FROM jsonb_each_text(session.savings_estimated_baseline_models) UNION ALL SELECT * FROM jsonb_each_text(models.deltas) ) combined GROUP BY key HAVING SUM(value::int) > 0 ) counts ) -FROM totals JOIN models USING (api_key, session_id, router_name) -WHERE session.api_key = totals.api_key AND session.session_id = totals.session_id +FROM totals JOIN models USING ({identity_columns}) +WHERE {user_match}session.api_key = totals.api_key AND session.session_id = totals.session_id AND session.router_name = totals.router_name """ +_UPDATE_SESSIONS: Final = _session_correction_sql(user_scoped=False) +_UPDATE_USER_SESSIONS: Final = _session_correction_sql(user_scoped=True) + + def _primary_transaction(client: PrismaClient) -> _TransactionManager: primary: Final = cast(_TransactionalDatabase, writer_wrapper(client.db)) return primary.tx(timeout=_TRANSACTION_TIMEOUT) @@ -308,6 +321,7 @@ def _change(record: BaselineAccountingRecord, old: BaselinePublication | None, n request_id=record.observation.request_id, publication=new, api_key=record.api_key, + user_id=record.turn.user_id if record.turn is not None else "", session_id=record.session_id, router_name=record.router_name, baseline_model=record.baseline_model, @@ -357,6 +371,8 @@ async def _publish(db: SupportsRawQueries, changes: Sequence[_Change]) -> None: serialized: Final = json.dumps(tuple(change.model_dump(mode="json") for change in changes), separators=(",", ":")) await db.execute_raw(_UPDATE_LOGS, serialized) await db.execute_raw(_UPDATE_SESSIONS, serialized) + if any(change.user_id for change in changes): + await db.execute_raw(_UPDATE_USER_SESSIONS, serialized) for entity, table in DAILY_SPEND_TABLES.items(): if adjustments := tuple( change.daily.adjustment(target, change.savings_delta, change.request_id) diff --git a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py index b28a653c9aa..85e19fa8a32 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py +++ b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py @@ -1,5 +1,6 @@ import asyncio import time +from contextvars import ContextVar from dataclasses import dataclass from datetime import datetime, timedelta, timezone from typing import Final, Literal, TypeAlias @@ -40,6 +41,28 @@ class TableCleanupResult: stop_reason: StopReason +class _RunProgress: + """How far one cleanup run has got, reported if that run is cancelled""" + + def __init__(self) -> None: + self.rows_deleted: int = 0 + self.batches: int = 0 + + def record_batch(self, rows_deleted: int) -> None: + self.rows_deleted += rows_deleted + self.batches += 1 + + +_run_progress: ContextVar[_RunProgress] = ContextVar("spend_log_cleanup_run_progress") + + +def _record_run_batch(rows_deleted: int) -> None: + """Count a batch towards the run in progress, if a run is what issued it""" + progress: Final = _run_progress.get(None) + if progress is not None: + progress.record_batch(rows_deleted) + + class _RemainingRow(BaseModel): """One row of the capped outstanding-rows probe, validated out of prisma's untyped result.""" @@ -422,6 +445,7 @@ class SpendLogCleanup: total_deleted += deleted_count run_count += 1 + _record_run_batch(deleted_count) # Add a small sleep to prevent overwhelming the database await asyncio.sleep(0.1) @@ -492,6 +516,18 @@ class SpendLogCleanup: deadline=deadline, ) + async def _delete_old_autorouter_user_session_rows( + self, prisma_client: PrismaClient, cutoff_date: datetime, deadline: float + ) -> TableCleanupResult: + return await self._delete_old_rows_batched( + prisma_client, + cutoff_date, + table_name="LiteLLM_AutoRouterUserSession", + key_columns=("user_id", "api_key", "session_id", "router_name"), + time_column="last_turn_at", + deadline=deadline, + ) + async def _delete_old_health_check_rows( self, prisma_client: PrismaClient, cutoff_date: datetime, deadline: float ) -> TableCleanupResult: @@ -560,9 +596,17 @@ class SpendLogCleanup: ) except Exception: # noqa: BLE001 # retained observations are retried by the next cleanup job verbose_proxy_logger.warning("Auto-router baseline retention remains pending") - sessions_result: Final = await self._delete_old_autorouter_session_rows(prisma_client, session_cutoff, deadline) + sessions_result: Final = await self._delete_old_autorouter_session_rows( + prisma_client, session_cutoff, self._group_deadline(deadline, 2) + ) verbose_proxy_logger.info("Deleted %s expired auto-router session rollup rows", sessions_result.rows_deleted) - return (sessions_result,) + user_sessions_result: Final = await self._delete_old_autorouter_user_session_rows( + prisma_client, session_cutoff, deadline + ) + verbose_proxy_logger.info( + "Deleted %s expired auto-router user session rollup rows", user_sessions_result.rows_deleted + ) + return (sessions_result, user_sessions_result) async def _clean_health_checks( self, prisma_client: PrismaClient, retention_seconds: int, deadline: float @@ -601,6 +645,9 @@ class SpendLogCleanup: If no pod_lock_manager, runs cleanup without distributed locking. """ lock_acquired = False + run_started_at: Final = time.monotonic() + progress: Final = _RunProgress() + progress_token: Final = _run_progress.set(progress) try: verbose_proxy_logger.info("Cleanup job triggered at %s", datetime.now()) self._refresh_bounds() @@ -681,6 +728,15 @@ class SpendLogCleanup: self._run_outcome(spend_log_results + session_results + health_check_results) ) + except asyncio.CancelledError: + verbose_proxy_logger.error( + "Spend log cleanup cancelled after %.2fs (rows_deleted=%d, batches=%d); the next run resumes from here", + time.monotonic() - run_started_at, + progress.rows_deleted, + progress.batches, + ) + SpendLogCleanupMetrics.record_run("aborted") + raise except Exception as e: # .exception() captures the traceback; str(e) alone on a Prisma/DB # timeout is often empty and gives operators no signal to diagnose. @@ -692,6 +748,7 @@ class SpendLogCleanup: SpendLogCleanupMetrics.record_run("aborted") return # Return after error handling finally: + _run_progress.reset(progress_token) # Only release the lock if it was actually acquired if lock_acquired and self.pod_lock_manager and self.pod_lock_manager.redis_cache: await self.pod_lock_manager.release_lock(cronjob_id=SPEND_LOG_CLEANUP_JOB_NAME) diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index d8054a8dc4a..03fc58622ce 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -789,14 +789,18 @@ async def get_auto_router_benchmarks( ] = None, end_date: Annotated[str | None, Query(description="YYYY-MM-DD UTC, inclusive (defaults to today)")] = None, api_key: Annotated[str | None, Query(description="Filter to one virtual key token hash")] = None, + user_id: Annotated[ + str | None, Query(min_length=1, description="Filter to one canonical internal user recorded on each turn") + ] = None, ) -> AutoRouterBenchmarksResponse: """ Benchmarks for the auto-router dashboard: session shape, savings against the configured baseline, and prompt-caching behaviour bucketed by what the router did. - Reads the LiteLLM_AutoRouterSession rollup, folded once per request at spend-write time, - so this endpoint never scans LiteLLM_SpendLogs. A session is in the window when it - overlaps it: its last turn is on or after start_date and its first turn is on or before + Reads session rollups folded once per request at spend-write time, so this endpoint + never scans LiteLLM_SpendLogs. A user filter selects only turns attributed to that + internal user when written; older key-only history remains outside user views. A session + is in the window when it overlaps it: its last turn is on or after start_date and its first turn is on or before end_date. Overall hit rate is over telemetry-bearing turns; each bucket's hit rate is over that bucket's turns. @@ -826,6 +830,7 @@ async def get_auto_router_benchmarks( start_day.isoformat(), (end_day + timedelta(days=1)).isoformat(), api_key, + user_id, ) rows: Final = _SESSION_AGG_ROWS.validate_python(raw_rows or ()) groups: Final = ( diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index c1388e8bb81..ee0e4db291b 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -1078,6 +1078,9 @@ if MCP_AVAILABLE: return {"servers": registry_servers} ## FastAPI Routes + def _mcp_server_display_order(server: LiteLLM_MCPServerTable) -> tuple[str, str]: + return ((server.server_name or server.alias or server.server_id).lower(), server.server_id) + def _get_user_mcp_management_mode() -> UserMCPManagementMode: from litellm.proxy.proxy_server import ( general_settings as proxy_general_settings, @@ -1228,10 +1231,12 @@ if MCP_AVAILABLE: detail="You do not have permission to view MCP servers for this team.", ) - redacted_mcp_servers = await _get_team_scoped_mcp_server_list(sanitized_team_id) + redacted_mcp_servers = sorted( + await _get_team_scoped_mcp_server_list(sanitized_team_id), key=_mcp_server_display_order + ) else: servers: Final = await _resolve_accessible_mcp_servers(user_api_key_dict) - redacted_mcp_servers = _redact_mcp_credentials_list(servers) + redacted_mcp_servers = sorted(_redact_mcp_credentials_list(servers), key=_mcp_server_display_order) if connected_app_view is True and is_ui_session_credential(user_api_key_dict): reachable_ids: Final = await _connected_app_reachable_server_ids(user_api_key_dict) diff --git a/litellm/proxy/management_endpoints/router_settings_endpoints.py b/litellm/proxy/management_endpoints/router_settings_endpoints.py index fc000b1638b..d6d74ada35a 100644 --- a/litellm/proxy/management_endpoints/router_settings_endpoints.py +++ b/litellm/proxy/management_endpoints/router_settings_endpoints.py @@ -8,6 +8,8 @@ GET /router/fields - Get router settings field definitions without values (for U """ import inspect +from collections.abc import Mapping +from types import MappingProxyType from typing import Any, Final, get_args from fastapi import APIRouter, Depends @@ -16,6 +18,7 @@ from pydantic import BaseModel, Field from litellm._logging import verbose_proxy_logger from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.config_resolvers import FieldSource, SettingsStore, source_for from litellm.router import Router from litellm.types.management_endpoints import ( ROUTER_SETTINGS_FIELDS, @@ -30,6 +33,7 @@ class RouterSettingsResponse(BaseModel): fields: list[RouterSettingsField] = Field(description="List of all configurable router settings with metadata") current_values: dict[str, Any] = Field(description="Current values of router settings") routing_strategy_descriptions: dict[str, str] = Field(description="Descriptions for each routing strategy option") + source: dict[str, FieldSource] = Field(description="Source of each current router setting") class RouterFieldsResponse(BaseModel): @@ -39,6 +43,18 @@ class RouterFieldsResponse(BaseModel): routing_strategy_descriptions: dict[str, str] = Field(description="Descriptions for each routing strategy option") +def _router_setting_source( + settings: SettingsStore, + key: str, + current_value: object, + field_default: object, +) -> FieldSource: + source: Final = source_for(settings, key, field_default) + if source != "unset": + return source + return "default" if current_value is not None else "unset" + + def _get_routing_strategies_from_router_class() -> list[str]: """ Dynamically extract routing strategies from the Router class __init__ method. @@ -109,15 +125,29 @@ async def get_router_settings( # Merge with config values (config takes precedence) current_values.update(router_settings_from_config) - # Update field values with current values for field in router_fields: if field.field_name in current_values: field.field_value = current_values[field.field_name] + field_defaults: Final[Mapping[str, object]] = MappingProxyType( + {field.field_name: field.field_default for field in router_fields} + ) + source: Final[Mapping[str, FieldSource]] = MappingProxyType( + { + key: _router_setting_source( + proxy_config.router_settings, + key, + current_values[key], + field_defaults.get(key), + ) + for key in current_values + } + ) return RouterSettingsResponse( fields=router_fields, current_values=current_values, routing_strategy_descriptions=ROUTING_STRATEGY_DESCRIPTIONS, + source=source, ) except Exception as e: verbose_proxy_logger.error("Error fetching router settings: %s", e) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 8111e21e114..534dcf418cc 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -50,6 +50,7 @@ import anyio import websockets import websockets.exceptions from pydantic import BaseModel, Json, JsonValue, TypeAdapter, ValidationError +from pydantic.fields import FieldInfo, PydanticUndefined from typing_extensions import NotRequired, ReadOnly, assert_never from litellm._uuid import uuid @@ -456,7 +457,13 @@ from litellm.proxy.common_utils.user_api_key_cache import ( project_spend_counter_key, tag_cache_key, ) -from litellm.proxy.config_resolvers import SettingsStore, config_ownership_message, resolve_fields +from litellm.proxy.config_resolvers import ( + FieldSource, + SettingsStore, + config_ownership_message, + resolve_fields, + source_for, +) from litellm.proxy.config_resolvers.alerting import ( EMAIL_DESCRIPTORS, MS_TEAMS_DESCRIPTORS, @@ -715,6 +722,11 @@ from litellm.proxy.route_llm_request import route_request from litellm.proxy.route_priority import hot_routes_first from litellm.proxy.search_endpoints.endpoints import router as search_router from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager +from litellm.proxy.shutdown.scheduled_jobs import ( + AwaitableAsyncIOExecutor, + pause_scheduled_jobs, + stop_in_flight_scheduler_jobs, +) from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start from litellm.proxy.spend_tracking.daily_global_spend_rollup import ( run_scheduled_daily_global_spend_reconcile, @@ -1488,6 +1500,10 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: if model_info_scheduler is not scheduler: model_info_scheduler.shutdown(wait=False) + # Shutdown event - stop starting scheduled jobs; the ones already running keep the drain window + if scheduler is not None: + pause_scheduled_jobs(scheduler) + # Shutdown event - drain in-flight requests before tearing down dependencies # so SIGTERM (rolling update, scale-down, liveness kill) doesn't drop them. GracefulShutdownManager.start_shutdown() @@ -1527,6 +1543,13 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: await _drain_spend_event_producer_on_shutdown() + # Shutdown event - finish or cancel in-flight scheduled jobs before the shutdown flushes and the DB disconnect + if scheduler is not None and scheduler_executor is not None: + try: + await stop_in_flight_scheduler_jobs(scheduler, scheduler_executor) + except Exception as e: + verbose_proxy_logger.error("Error stopping in-flight scheduled jobs: %s", e) + await flush_spend_counters_on_shutdown() await _flush_spend_logs_queue_on_shutdown() @@ -2529,6 +2552,7 @@ celery_app_conn: Final = None celery_fn: Final = None # Redis Queue for handling requests scheduler = None +scheduler_executor: AwaitableAsyncIOExecutor | None = None # rebind-ok: bound once the scheduler is built at startup # Global variable for anthropic beta headers reload scheduling last_anthropic_beta_headers_reload = None @@ -4951,6 +4975,12 @@ def _as_settings_mapping(value: object) -> Mapping[str, SettingsJsonValue]: return _SETTINGS_MAPPING.validate_python(value) +def _get_field_default(field_info: FieldInfo) -> JsonValue: + if field_info.default is PydanticUndefined: + return None + return cast(JsonValue, field_info.default) # cast-ok: Pydantic field defaults are JSON values at runtime + + def _bind_general_settings_store(settings: SettingsStore) -> None: global general_settings general_settings = settings # pyright: ignore[reportAssignmentType] # legacy global accepts mappings @@ -10081,7 +10111,7 @@ class ProxyStartupEvent: proxy_logging_obj: ProxyLogging, ) -> ProxyWorkerHeartbeat: """Initializes scheduled background jobs""" - global heuristic_v1_tuning_baselines, store_model_in_db, scheduler # rebind-ok: startup publishes the one read-only baseline snapshot + global heuristic_v1_tuning_baselines, store_model_in_db, scheduler, scheduler_executor # rebind-ok: startup publishes the one read-only baseline snapshot # MEMORY LEAK FIX: Configure scheduler with optimized settings # Memray analysis showed APScheduler's normalize() and _apply_jitter() causing @@ -10090,9 +10120,9 @@ class ProxyStartupEvent: # 1. Remove/minimize jitter to avoid normalize() memory explosion # 2. Use larger misfire_grace_time to prevent backlog calculations # 3. Set replace_existing=True to avoid duplicate jobs - from apscheduler.executors.asyncio import AsyncIOExecutor from apscheduler.jobstores.memory import MemoryJobStore + scheduler_executor = AwaitableAsyncIOExecutor() # rebind-ok: shutdown awaits the jobs this executor runs scheduler = AsyncIOScheduler( job_defaults={ "coalesce": APSCHEDULER_COALESCE, @@ -10105,7 +10135,7 @@ class ProxyStartupEvent: jobstores={"default": MemoryJobStore()}, # explicitly use memory job store # Use simple executor to minimize overhead executors={ - "default": AsyncIOExecutor(), + "default": scheduler_executor, }, # Disable timezone awareness to reduce computation timezone=None, @@ -16006,6 +16036,22 @@ async def model_settings(): #### ALERTING MANAGEMENT ENDPOINTS #### +def _nested_setting_source( + settings: SettingsStore, + db_values: Mapping[str, JsonValue], + parent_key: str, + field_name: str, + field_default: JsonValue, +) -> FieldSource: + unset_source: Final[FieldSource] = "default" if field_default is not None else "unset" + parent_value: Final = settings.config_value(parent_key) + if isinstance(parent_value, Mapping) and field_name in parent_value: + return "config" + if settings.owned_by_config(parent_key): + return unset_source + return "db" if field_name in db_values else unset_source + + @router.get( "/alerting/settings", description="Return the configurable alerting param, description, and current value", @@ -16043,17 +16089,20 @@ async def alerting_settings( where={"param_name": "general_settings"} ) - if db_general_settings is not None and db_general_settings.param_value is not None: - db_general_settings_dict: Final = dict(db_general_settings.param_value) - alerting_args_dict: dict = cast( # cast-ok: ConfigGeneralSettings validates alerting_args as a dict on write - dict[str, JsonValue], db_general_settings_dict.get("alerting_args", {}) - ) - alerting_values: list | None = cast( # cast-ok: ConfigGeneralSettings validates alerting as a list on write - list[JsonValue] | None, db_general_settings_dict.get("alerting") - ) - else: - alerting_args_dict = {} - alerting_values = None + db_general_settings_dict: Final[Mapping[str, JsonValue]] = MappingProxyType( + dict(db_general_settings.param_value) # mutable-ok: Prisma returns the JSON column as a plain dict + if db_general_settings is not None and db_general_settings.param_value is not None + else {} + ) + alerting_args_value: Final = db_general_settings_dict.get("alerting_args") + alerting_args_dict: Final[Mapping[str, JsonValue]] = MappingProxyType( + alerting_args_value if isinstance(alerting_args_value, dict) else {} + ) + alerting_values: Final = cast( # cast-ok: alerting is stored as a JSON list when present + list[JsonValue] | None, db_general_settings_dict.get("alerting") + ) + + settings: Final = proxy_config.settings allowed_args: Final = MappingProxyType( { @@ -16082,9 +16131,9 @@ async def alerting_settings( is_slack_enabled = False - if general_settings.get("alerting") and isinstance(general_settings["alerting"], list): - if "slack" in general_settings["alerting"]: - is_slack_enabled = True + alerting: Final = settings.get("alerting") + if isinstance(alerting, list) and "slack" in alerting: + is_slack_enabled = True _response_obj = ConfigList( field_name="slack_alerting", @@ -16092,6 +16141,7 @@ async def alerting_settings( field_description="Enable slack alerting for monitoring proxy in production: llm outages, budgets, spend tracking failures.", field_value=is_slack_enabled, stored_in_db=True if alerting_values is not None else False, + source=source_for(settings, "alerting"), field_default_value=None, premium_field=False, ) @@ -16099,6 +16149,7 @@ async def alerting_settings( for field_name, field_info in SlackAlertingArgs.model_fields.items(): if field_name in allowed_args: + field_default: JsonValue = _get_field_default(field_info) _stored_in_db: bool | None = None if field_name in alerting_args_dict: _stored_in_db = True @@ -16109,9 +16160,16 @@ async def alerting_settings( field_name=field_name, field_type=allowed_args[field_name], field_description=field_info.description or "", - field_value=_slack_alerting_args_dict.get(field_name, None), + field_value=_slack_alerting_args_dict.get(field_name, field_default), stored_in_db=_stored_in_db, - field_default_value=field_info.default, + source=_nested_setting_source( + settings, + alerting_args_dict, + "alerting_args", + field_name, + field_default, + ), + field_default_value=field_default, premium_field=(True if field_name == "region_outage_alert_ttl" else False), ) return_val.append(_response_obj) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 368864f9cd1..31fc76f7a89 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1623,6 +1623,47 @@ model LiteLLM_AutoRouterSession { @@index([last_turn_at], map: "idx_autorouter_session_last_turn") } +model LiteLLM_AutoRouterUserSession { + user_id String + api_key String + session_id String + router_name String + router_type String + first_turn_at DateTime + last_turn_at DateTime + last_model String + models Json @default("{}") + turns Int @default(0) + unordered_turns Int @default(0) + covered_turns Int @default(0) + cache_hits Int @default(0) + same_model_turns Int @default(0) + same_model_hits Int @default(0) + first_visit_turns Int @default(0) + first_visit_hits Int @default(0) + return_turns Int @default(0) + return_hits Int @default(0) + return_expired_misses Int @default(0) + return_within_ttl_misses Int @default(0) + ttl_5m_turns Int @default(0) + ttl_1h_turns Int @default(0) + total_tokens BigInt @default(0) + spend Float @default(0) + saved_spend Float @default(0) + savings_estimated_turns Int @default(0) + savings_estimated_actual_spend Float @default(0) + savings_estimated_saved_spend Float @default(0) + savings_estimated_baseline_models Json @default("{}") + classifier_cost Float @default(0) + classifier_cost_recorded_turns Int @default(0) + tier_turns Json @default("{}") + baseline_models Json @default("{}") + + @@id([user_id, api_key, session_id, router_name]) + @@index([last_turn_at], map: "idx_autorouter_user_session_last_turn") + @@index([user_id, last_turn_at], map: "idx_autorouter_user_session_user_last_turn") +} + // Shadow eval: evaluation of an auto-router against one or more keys' live traffic, in // either direction. forward duplicates the requests the keys did not route through the // router through it, answering whether they should adopt it; reverse duplicates the diff --git a/litellm/proxy/shutdown/scheduled_jobs.py b/litellm/proxy/shutdown/scheduled_jobs.py new file mode 100644 index 00000000000..7889c35cf4e --- /dev/null +++ b/litellm/proxy/shutdown/scheduled_jobs.py @@ -0,0 +1,79 @@ +# pyright: reportMissingTypeStubs=false # apscheduler ships no type information + +import asyncio +from collections.abc import Collection +from typing import Final, Protocol + +from apscheduler.executors.asyncio import AsyncIOExecutor + +from litellm._logging import verbose_proxy_logger +from litellm.constants import ( + SCHEDULED_JOB_SHUTDOWN_CANCEL_TIMEOUT_SECONDS, + SCHEDULED_JOB_SHUTDOWN_FINISH_TIMEOUT_SECONDS, +) + + +class StoppableScheduler(Protocol): + """The slice of ``AsyncIOScheduler`` shutdown uses, which ships no type information""" + + @property + def running(self) -> bool: ... + + def pause(self) -> None: ... + + def shutdown(self, wait: bool = ...) -> None: ... + + +class AwaitableAsyncIOExecutor(AsyncIOExecutor): # pyright: ignore[reportUntypedBaseClass] # apscheduler ships no type information and is absent from the type-check env + """``AsyncIOExecutor`` whose in-flight job tasks can be awaited after ``shutdown`` cancels them""" + + _pending_futures: Collection["asyncio.Future[object]"] + + def in_flight_jobs(self) -> tuple["asyncio.Future[object]", ...]: + """The job tasks that are running right now, as a snapshot""" + return tuple(future for future in self._pending_futures if not future.done()) + + +def pause_scheduled_jobs(scheduler: StoppableScheduler) -> None: + """Stop the scheduler from starting jobs that shutdown would only cancel; running jobs continue""" + if scheduler.running: + scheduler.pause() + + +async def stop_in_flight_scheduler_jobs( + scheduler: StoppableScheduler, + executor: AwaitableAsyncIOExecutor, + *, + finish_timeout_seconds: float = SCHEDULED_JOB_SHUTDOWN_FINISH_TIMEOUT_SECONDS, + cancel_timeout_seconds: float = SCHEDULED_JOB_SHUTDOWN_CANCEL_TIMEOUT_SECONDS, +) -> None: + """ + Let in-flight jobs finish for up to finish_timeout_seconds, then stop the scheduler and wait, bounded by + cancel_timeout_seconds, for the jobs it cancels. + + Must run before the database is disconnected: a write job that finishes needs its connection, + and a job's cancellation handler is what records the run's outcome. + """ + if not scheduler.running: + return + in_flight: Final = executor.in_flight_jobs() + if in_flight: + verbose_proxy_logger.info( + "Waiting up to %ss for %d in-flight scheduled job(s) to finish", + finish_timeout_seconds, + len(in_flight), + ) + still_running: Final = ( + (await asyncio.wait(in_flight, timeout=finish_timeout_seconds))[1] if in_flight else frozenset() + ) + scheduler.shutdown(wait=False) + if not still_running: + return + verbose_proxy_logger.info("Cancelling %d in-flight scheduled job(s) for shutdown", len(still_running)) + _done, pending = await asyncio.wait(still_running, timeout=cancel_timeout_seconds) + if pending: + verbose_proxy_logger.warning( + "%d scheduled job(s) did not finish within %ss of cancellation; giving up on them", + len(pending), + cancel_timeout_seconds, + ) diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index 7bdadeadf86..f4d4ccf5851 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -15,8 +15,8 @@ from typing import ( from urllib.parse import urlparse from fastapi import APIRouter, Body, Depends, File, HTTPException, UploadFile -from pydantic import ConfigDict, JsonValue, TypeAdapter, ValidationError, create_model -from pydantic.fields import FieldInfo +from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError, create_model +from pydantic.fields import FieldInfo, PydanticUndefined from typing_extensions import NotRequired, ReadOnly, TypedDict import litellm @@ -25,6 +25,7 @@ from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_keys from litellm.proxy._experimental.mcp_server.tool_search import MCP_TOOL_SEARCH_SETTINGS_KEY from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.config_resolvers import FieldSource, SettingsStore, source_for from litellm.proxy.config_resolvers.settings_store import ConfigOwnedKeyError from litellm.proxy.config_resolvers.sso import ( SSO_FIELD_ENV_VARS, @@ -35,7 +36,10 @@ from litellm.proxy.management_endpoints.team_admin_field_permissions import ( SUPPORTED_TEAM_ADMIN_PERMISSIONS, TEAM_ADMIN_EDITABLE_TEAM_FIELDS_SETTING, ) -from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled +from litellm.proxy.spend_tracking.ptu_feature_flag import ( + PTU_COST_ATTRIBUTION_ENV_VAR, + is_ptu_cost_attribution_enabled, +) from litellm.proxy.utils import invalidate_config_param from litellm.repositories.config_repository import ConfigRepository from litellm.repositories.organization_repository import OrganizationRepository @@ -45,6 +49,7 @@ from litellm.repositories.table_repositories import ( UISettingsRepository, ) from litellm.repositories.team_repository import TeamRepository +from litellm.secret_managers.main import get_secret from litellm.types.mcp import MCPToolSearchSettings from litellm.types.proxy.management_endpoints.ui_sso import ( DefaultTeamSSOParams, @@ -199,6 +204,11 @@ class SettingsResponse(BaseModel): """Schema information including descriptions and property types for UI display""" +class _SettingsWithSchema(BaseModel): + values: dict[str, object] + field_schema: dict[str, object] + + class SSOSettingsResponse(SettingsResponse): """Response model for SSO settings""" @@ -330,6 +340,8 @@ class UISettings(BaseModel): class UISettingsResponse(SettingsResponse): """Response model for UI settings""" + source: dict[str, FieldSource] + # Allowlist of UI settings that can be stored ALLOWED_UI_SETTINGS_FIELDS: Final = { @@ -748,6 +760,25 @@ def _root_schema(settings_class: type[BaseModel]) -> _RootSchema: ) +def _model_field_default(settings_class: type[BaseModel], field_name: str) -> object: + field_info: Final = settings_class.model_fields.get(field_name) + if field_info is None or field_info.default is PydanticUndefined: + return None + return cast(object, field_info.default) # cast-ok: Pydantic field defaults are untyped + + +def _ui_setting_source( + key: str, + value: object, + settings: SettingsStore, + settings_class: type[BaseModel], +) -> FieldSource: + if key == ENABLE_PTU_COST_ATTRIBUTION_UI_SETTING: + configured_value: Final = get_secret(PTU_COST_ATTRIBUTION_ENV_VAR, None) + return "config" if configured_value is not None or value is True else "default" + return source_for(settings, key, _model_field_default(settings_class, key)) + + async def _get_settings_with_schema( settings_key: str, settings_class: type[BaseModel], @@ -1705,7 +1736,7 @@ async def get_ui_settings(): Get UI-specific configuration flags. All authenticated users can fetch these settings for client-side behavior. """ - from litellm.proxy.proxy_server import prisma_client + from litellm.proxy.proxy_server import prisma_client, proxy_config if prisma_client is None: raise HTTPException( @@ -1730,20 +1761,43 @@ async def get_ui_settings(): await user_api_key_cache.async_set_cache(key=UI_SETTINGS_CACHE_KEY, value=ui_settings, ttl=UI_SETTINGS_CACHE_TTL) - # Build config-like object for schema helper - config: Final[dict[str, object]] = {"litellm_settings": {"ui_settings": ui_settings}} - - settings: Final = await _get_settings_with_schema( - settings_key="ui_settings", - settings_class=_get_effective_ui_settings_class(), - config=config, + effective_ui_settings: Final[Mapping[str, object]] = MappingProxyType( + { + **ui_settings, + **{key: proxy_config.settings[key] for key in ALLOWED_UI_SETTINGS_FIELDS if key in proxy_config.settings}, + } + ) + config: Final[Mapping[str, object]] = MappingProxyType( + {"litellm_settings": MappingProxyType({"ui_settings": effective_ui_settings})} + ) + settings_class: Final = _get_effective_ui_settings_class() + resolved_settings: Final = _SettingsWithSchema.model_validate( + await _get_settings_with_schema( + settings_key="ui_settings", + settings_class=settings_class, + config=config, + ) + ) + values: Final[Mapping[str, object]] = MappingProxyType( + { + **resolved_settings.values, + ENABLE_PTU_COST_ATTRIBUTION_UI_SETTING: is_ptu_cost_attribution_enabled(), + } + ) + source: Final[Mapping[str, FieldSource]] = MappingProxyType( + { + key: ( + _ui_setting_source(key, values[key], proxy_config.settings, settings_class) + if key in proxy_config.settings or key not in ui_settings + else "db" + ) + for key in values + } ) return UISettingsResponse( - values={ - **settings["values"], - ENABLE_PTU_COST_ATTRIBUTION_UI_SETTING: is_ptu_cost_attribution_enabled(), - }, - field_schema=settings["field_schema"], + values=values, + field_schema=resolved_settings.field_schema, + source=source, ) diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 1537e3a540c..c4dbe68ebe2 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -1321,7 +1321,7 @@ class ComplexityRouterConfig(BaseModel): ) enable_context_window_escalation: bool = Field( - default=True, + default=False, description=( "Escalate a request off a tier whose models provably cannot hold its prompt, before " "dispatch. The classifier scores complexity and never prompt size, so a long agentic " @@ -1331,7 +1331,8 @@ class ComplexityRouterConfig(BaseModel): "moves to the lowest configured tier with a model whose declared window fits; when " "only some of the tier's models fit, the pick is restricted to those and the tier " "keeps the request. Models with no resolvable window are never escalated away from " - "and never escalated onto. Set false to dispatch on complexity alone, as before." + "and never escalated onto. Disabled by default: omit or set false to dispatch on " + "complexity alone; set true to enable context-window escalation." ), ) context_window_escalation_buffer: float = Field( diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index 05a6df6d5af..31623a53dd0 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -93,6 +93,133 @@ class ResponsesWebSocketConnection: def recv_text(self) -> Future[str | None]: ... def close(self) -> Future[None]: ... +@final +class _CacheTestBinding: + @property + def kind(self) -> str: ... + def lookup( + self, + request: object, + *, + callback_kwargs: Mapping[str, object] | Sequence[object] | None = None, + ) -> object: ... + def store( + self, + request: object, + response: object, + *, + callback_kwargs: Mapping[str, object] | None = None, + ) -> None: ... + def lookup_batch( + self, + requests: Sequence[object], + *, + callback_kwargs: Sequence[object] | None = None, + ) -> object: ... + def async_lookup( + self, + request: object, + *, + callback_kwargs: Mapping[str, object] | None = None, + ) -> Future[object]: ... + def async_store( + self, + request: object, + response: object, + *, + callback_kwargs: Mapping[str, object] | None = None, + ) -> Future[None]: ... + def async_lookup_batch( + self, + requests: Sequence[object], + *, + callback_kwargs: Sequence[object] | None = None, + ) -> Future[object]: ... + def async_store_batch( + self, + requests: Sequence[object], + responses: Sequence[object], + *, + callback_result: object = None, + callback_kwargs: Mapping[str, object] | None = None, + ) -> Future[object]: ... + def async_flush(self) -> Future[None]: ... + def ping(self) -> Future[object]: ... + +@final +class _CacheTestHandle: + def __new__(cls, _uninstantiable: Never, /) -> Never: ... + @staticmethod + def memory( + *, + capacity: int = 200, + ttl_seconds: float = 600.0, + max_entry_bytes: int = 1048576, + ) -> _CacheTestHandle: ... + @staticmethod + def redis( + url: str, + *, + ttl_seconds: float = 60.0, + namespace: str | None = None, + startup_nodes: Sequence[tuple[str, int]] | None = None, + ) -> _CacheTestHandle: ... + @staticmethod + def disk(directory: str) -> _CacheTestHandle: ... + @staticmethod + def qdrant_semantic( + url: str, + *, + collection_name: str, + similarity_threshold: float, + vector_size: int, + embedding_model: str = "text-embedding-3-small", + api_key: str | None = None, + embedding_api_key: str | None = None, + embedding_api_base: str | None = None, + embedding_timeout_seconds: float | None = None, + quantization: str = "binary", + ) -> _CacheTestHandle: ... + @staticmethod + def azure_blob(account_url: str, container: str) -> _CacheTestHandle: ... + @staticmethod + def redis_semantic(backend: object) -> _CacheTestHandle: ... + @staticmethod + def valkey_semantic( + url: str, + similarity_threshold: float, + index_name: str, + embedder: object, + ) -> _CacheTestHandle: ... + @staticmethod + def gcs( + bucket_name: str, + *, + gcs_path: str | None = None, + path_service_account: str | None = None, + endpoint: str | None = None, + token: str | None = None, + ) -> _CacheTestHandle: ... + @staticmethod + def s3( + bucket: str, + *, + region: str, + endpoint_url: str | None = None, + key_prefix: str = "", + access_key_id: str | None = None, + secret_access_key: str | None = None, + session_token: str | None = None, + ) -> _CacheTestHandle: ... + @property + def backend(self) -> str: ... + def _bind_facade(self, facade: object) -> None: ... + +@final +class _CacheTestResolver: + def __new__(cls, namespace: object) -> _CacheTestResolver: ... + def resolve(self) -> _CacheTestBinding: ... + @final class TokenCounter: def __new__(cls, tokenizer_json: str) -> TokenCounter: ... diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index beb90c2fb8c..3ca9958ad0d 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -43037,21 +43037,21 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro": { - "input_cost_per_token": 8.92272e-07, + "input_cost_per_token": 8.83746e-07, "input_cost_per_token_cache_hit": 4.4e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 1.784544e-06, + "output_cost_per_token": 1.767492e-06, "source": "https://openrouter.ai/api/v1/models", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "cache_read_input_token_cost": 7.4356e-08, + "cache_read_input_token_cost": 7.36455e-08, "supports_audio_input": false, "supports_pdf_input": false, "supports_vision": false, @@ -70299,8 +70299,8 @@ "supports_web_search": true }, "openrouter/meta-llama/llama-4-maverick": { - "input_cost_per_token": 2e-07, - "output_cost_per_token": 8e-07, + "input_cost_per_token": 1.875e-07, + "output_cost_per_token": 6.525e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 16384, @@ -73794,6 +73794,7 @@ "supports_web_search": false }, "openrouter/bytedance-seed/seed-1.6": { + "deprecation_date": "2026-11-11", "input_cost_per_token": 2.5e-07, "input_cost_per_token_above_128k_tokens": 5e-07, "litellm_provider": "openrouter", @@ -73815,6 +73816,7 @@ "supports_web_search": false }, "openrouter/bytedance-seed/seed-1.6-flash": { + "deprecation_date": "2026-11-11", "input_cost_per_token": 7.5e-08, "input_cost_per_token_above_128k_tokens": 1e-07, "litellm_provider": "openrouter", @@ -73855,6 +73857,7 @@ "supports_web_search": false }, "openrouter/bytedance-seed/seed-2.0-code": { + "deprecation_date": "2026-11-11", "input_cost_per_token": 5e-07, "input_cost_per_token_above_128k_tokens": 1e-06, "litellm_provider": "openrouter", @@ -76876,6 +76879,26 @@ "supports_vision": true, "supports_web_search": true }, + "moonshotai.kimi-k3": { + "cache_creation_input_token_cost": 3.75e-06, + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1000000, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true + }, "global.moonshotai.kimi-k3": { "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, @@ -76975,5 +76998,51 @@ "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false + }, + "xiaomi_mimo/mimo-v2.6-pro": { + "cache_read_input_token_cost": 3.6e-09, + "input_cost_per_token": 4.35e-07, + "litellm_provider": "xiaomi_mimo", + "max_input_tokens": 1048576, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 8.7e-07, + "source": "https://platform.xiaomimimo.com/static/docs/price/pay-as-you-go.md", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_audio_input": true, + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true + }, + "xiaomi_mimo/mimo-v2.6-flash": { + "cache_read_input_token_cost": 2.8e-09, + "input_cost_per_token": 1.4e-07, + "litellm_provider": "xiaomi_mimo", + "max_input_tokens": 1048576, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 2.8e-07, + "source": "https://platform.xiaomimimo.com/static/docs/price/pay-as-you-go.md", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_audio_input": true, + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true } } diff --git a/schema.prisma b/schema.prisma index 368864f9cd1..31fc76f7a89 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1623,6 +1623,47 @@ model LiteLLM_AutoRouterSession { @@index([last_turn_at], map: "idx_autorouter_session_last_turn") } +model LiteLLM_AutoRouterUserSession { + user_id String + api_key String + session_id String + router_name String + router_type String + first_turn_at DateTime + last_turn_at DateTime + last_model String + models Json @default("{}") + turns Int @default(0) + unordered_turns Int @default(0) + covered_turns Int @default(0) + cache_hits Int @default(0) + same_model_turns Int @default(0) + same_model_hits Int @default(0) + first_visit_turns Int @default(0) + first_visit_hits Int @default(0) + return_turns Int @default(0) + return_hits Int @default(0) + return_expired_misses Int @default(0) + return_within_ttl_misses Int @default(0) + ttl_5m_turns Int @default(0) + ttl_1h_turns Int @default(0) + total_tokens BigInt @default(0) + spend Float @default(0) + saved_spend Float @default(0) + savings_estimated_turns Int @default(0) + savings_estimated_actual_spend Float @default(0) + savings_estimated_saved_spend Float @default(0) + savings_estimated_baseline_models Json @default("{}") + classifier_cost Float @default(0) + classifier_cost_recorded_turns Int @default(0) + tier_turns Json @default("{}") + baseline_models Json @default("{}") + + @@id([user_id, api_key, session_id, router_name]) + @@index([last_turn_at], map: "idx_autorouter_user_session_last_turn") + @@index([user_id, last_turn_at], map: "idx_autorouter_user_session_user_last_turn") +} + // Shadow eval: evaluation of an auto-router against one or more keys' live traffic, in // either direction. forward duplicates the requests the keys did not route through the // router through it, answering whether they should adopt it; reverse duplicates the diff --git a/tests/e2e/coverage_registry/llm_conversational.yaml b/tests/e2e/coverage_registry/llm_conversational.yaml index 49d4d92ff0b..64335aa560c 100644 --- a/tests/e2e/coverage_registry/llm_conversational.yaml +++ b/tests/e2e/coverage_registry/llm_conversational.yaml @@ -89,6 +89,9 @@ - {id: llm.chat_completions.together_ai.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together tool result round trip"} - {id: llm.chat_completions.together_ai.basic.nonstream.cost_logged, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: basic, streaming: nonstream, assertions: [cost_logged], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together cost header and spend row match the registry price"} - {id: llm.chat_completions.together_ai.thinking.nonstream.effort_none_disables, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: nonstream, assertions: [effort_none_disables], source: "llm_translation/test_together_ai_e2e.py", rationale: "reasoning_effort=none maps to Together's reasoning disable toggle on hybrid models"} +- {id: llm.chat_completions.xiaomi_mimo.basic.nonstream.cost_logged, module: llm, tier: P1, subject_endpoint: chat_completions, route: xiaomi_mimo, capability: basic, streaming: nonstream, assertions: [cost_logged], source: "llm_translation/test_xiaomi_mimo_e2e.py", rationale: "Native MiMo v2.6 rows price the cost header and spend row from the cost map"} +- {id: llm.chat_completions.xiaomi_mimo.thinking.stream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: xiaomi_mimo, capability: thinking, streaming: stream, assertions: [works], source: "llm_translation/test_xiaomi_mimo_e2e.py", rationale: "MiMo reasoning deltas stream as reasoning_content"} +- {id: llm.chat_completions.xiaomi_mimo.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: xiaomi_mimo, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_xiaomi_mimo_e2e.py", rationale: "MiMo tool calls are not dropped"} - {id: llm.chat_completions.together_ai.structured_output.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: structured_output, streaming: nonstream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "response_format json_schema reaches Together and constrains the reply"} - {id: llm.chat_completions.together_ai.prompt_cache_5m.nonstream.cost_logged, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: prompt_cache_5m, streaming: nonstream, assertions: [cache_hit, cost_logged], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together prefix-cache reads bill at cache_read_input_token_cost, not full input price"} - {id: llm.messages.together_ai.basic.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: together_ai, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together over /v1/messages streaming"} diff --git a/tests/e2e/coverage_registry/schema.py b/tests/e2e/coverage_registry/schema.py index 8b0d38a083c..d9c20d5c588 100644 --- a/tests/e2e/coverage_registry/schema.py +++ b/tests/e2e/coverage_registry/schema.py @@ -58,6 +58,7 @@ LlmRoute = Literal[ "openai", "together_ai", "vertex", + "xiaomi_mimo", ] LlmCapability = Literal[ diff --git a/tests/e2e/llm_translation/test_xiaomi_mimo_e2e.py b/tests/e2e/llm_translation/test_xiaomi_mimo_e2e.py new file mode 100644 index 00000000000..efca216634b --- /dev/null +++ b/tests/e2e/llm_translation/test_xiaomi_mimo_e2e.py @@ -0,0 +1,214 @@ +"""Live e2e: Xiaomi MiMo v2.6 through the gateway on /chat/completions. + +Both native ``xiaomi_mimo/`` v2.6 rows (pro and flash) are registered via +``/model/new`` and driven against Xiaomi's own endpoint. What the gateway owes +us is that the reasoning chain surfaces as ``reasoning_content``, tool calls +survive translation, and the cost header plus spend row follow the proxy's own +cost-map price for the row (read back from ``/model/info``, never pinned here). +Requires XIAOMI_MIMO_API_KEY on the proxy; no skip gate. +""" + +from __future__ import annotations + +from typing import Final + +import pytest +from e2e_config import unique_marker +from e2e_http import StreamingResponse, require_successful_call, unwrap +from lifecycle import ResourceManager +from models import ( + ChatBody, + ChatMessage, + ChatResponse, + ChatTool, + ChatToolFunction, + CostMapEntry, + LiteLLMParamsBody, + OutMessage, + SpendLogRow, +) +from passthrough_client import PassthroughClient +from pydantic import BaseModel + +pytestmark = pytest.mark.e2e + +BACKENDS: Final = ("xiaomi_mimo/mimo-v2.6-pro", "xiaomi_mimo/mimo-v2.6-flash") +ARITHMETIC_PROMPT = "What is 17 + 26? Answer with just the number." +WEATHER_PROMPT = "What is the weather in Paris? Use the tool." +COUNTING_PROMPT = "Count from 1 to 50, one number per line." + +WEATHER_TOOL = ChatTool( + function=ChatToolFunction( + name="get_weather", + description="Get the current weather for a location.", + parameters={ + "type": "object", + "properties": {"location": {"type": "string"}}, + "required": ["location"], + }, + ) +) + + +class _WeatherArgs(BaseModel): + location: str + + +class _StreamDelta(BaseModel): + content: str | None = None + reasoning_content: str | None = None + + +class _StreamChoice(BaseModel): + delta: _StreamDelta | None = None + + +class _StreamChunk(BaseModel): + choices: list[_StreamChoice] = [] + + +def _approx_equal(actual: float, expected: float) -> bool: + return abs(actual - expected) <= max(1e-9, abs(expected) * 1e-2) + + +@pytest.fixture(scope="module") +def registry(client: PassthroughClient) -> dict[str, CostMapEntry]: + return client.proxy.model_cost_map() + + +def _register(client: PassthroughClient, resources: ResourceManager, backend: str) -> tuple[str, str]: + model = f"e2e-xiaomi-{unique_marker()}" + model_id = client.proxy.create_model( + model, LiteLLMParamsBody(model=backend, api_key="os.environ/XIAOMI_MIMO_API_KEY") + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + return model, resources.key() + + +def _message(response: ChatResponse) -> OutMessage: + assert response.choices, f"Xiaomi returned no choices: {response}" + message = response.choices[0].message + assert message is not None, f"Xiaomi choice has no message: {response}" + return message + + +def _deltas(result: StreamingResponse) -> list[_StreamDelta]: + require_successful_call(result) + assert result.is_streaming, f"response was not streamed: {result.headers}" + assert not result.stream_error, f"stream errored: {result.stream_error}" + assert result.stream_done, f"stream never reached [DONE]: {result.stream_events[-3:]}" + return [ + choice.delta + for event in result.stream_events + for choice in _StreamChunk.model_validate_json(event).choices + if choice.delta is not None + ] + + +@pytest.mark.parametrize("backend", BACKENDS) +class TestXiaomiMimoChatCompletions: + @pytest.mark.covers("llm.chat_completions.xiaomi_mimo.basic.nonstream.cost_logged") + def test_cost_header_and_spend_row_match_the_registry_price( + self, + client: PassthroughClient, + resources: ResourceManager, + registry: dict[str, CostMapEntry], + backend: str, + ) -> None: + price = registry.get(backend) + assert price is not None, f"{backend} has no row in the proxy's cost map, so native calls would bill $0" + assert price.litellm_provider == "xiaomi_mimo", f"{backend} is filed under the wrong provider: {price}" + assert price.input_cost_per_token and price.output_cost_per_token, f"{backend} carries no price: {price}" + model, key = _register(client, resources, backend) + + result = client.proxy.transport.send( + "/chat/completions", + headers=client.proxy.transport.bearer(key), + json=ChatBody( + model=model, + messages=[ChatMessage(role="user", content=f"{ARITHMETIC_PROMPT} {unique_marker()}")], + max_tokens=1024, + ), + ) + require_successful_call(result) + response = ChatResponse.model_validate_json(result.body) + message = _message(response) + assert message.content and "43" in message.content, f"answer lost: {message}" + assert message.reasoning_content, f"{backend} reasons, but no reasoning_content came back: {message}" + + usage = response.usage + assert usage is not None and usage.prompt_tokens and usage.completion_tokens, ( + f"response carries no usage, so the cost cannot be real: {result.body[:300]}" + ) + header_cost = result.response_cost + assert header_cost is not None and header_cost > 0, ( + f"x-litellm-response-cost header missing or non-positive: {result.headers}" + ) + cached = (usage.prompt_tokens_details.cached_tokens or 0) if usage.prompt_tokens_details else 0 + expected = ( + (usage.prompt_tokens - cached) * price.input_cost_per_token + + cached * (price.cache_read_input_token_cost or 0.0) + + usage.completion_tokens * price.output_cost_per_token + ) + assert _approx_equal(header_cost, expected), ( + f"header cost {header_cost} disagrees with the registry price for {backend} at {usage}: expected {expected}" + ) + + def _priced(rows: list[SpendLogRow]) -> bool: + return any(row.spend is not None and row.spend > 0 for row in rows) + + rows = client.proxy.poll_logs_for_key(key, predicate=_priced) + priced = [row for row in rows if row.spend is not None and row.spend > 0] + assert priced, f"no priced spend row landed for key {key}; got {rows}" + row = priced[0] + assert row.custom_llm_provider == "xiaomi_mimo", f"spend row misattributed: {row}" + assert row.spend is not None and _approx_equal(row.spend, header_cost), ( + f"logged spend {row.spend} disagrees with the x-litellm-response-cost header {header_cost}" + ) + + @pytest.mark.covers("llm.chat_completions.xiaomi_mimo.thinking.stream.works") + def test_reasoning_and_answer_stream_as_deltas( + self, client: PassthroughClient, resources: ResourceManager, backend: str + ) -> None: + model, key = _register(client, resources, backend) + + deltas = _deltas( + client.proxy.chat_stream( + key, + ChatBody( + model=model, + messages=[ChatMessage(role="user", content=COUNTING_PROMPT)], + max_tokens=2048, + stream=True, + ), + ) + ) + reasoning = "".join(delta.reasoning_content or "" for delta in deltas) + content = "".join(delta.content or "" for delta in deltas) + assert reasoning, f"stream carried no reasoning_content deltas: {deltas[:5]}" + assert "50" in content, f"streamed answer lost: {content[:300]!r}" + + @pytest.mark.covers("llm.chat_completions.xiaomi_mimo.tool_use.nonstream.works") + def test_tool_call_is_returned(self, client: PassthroughClient, resources: ResourceManager, backend: str) -> None: + model, key = _register(client, resources, backend) + + message = _message( + unwrap( + client.proxy.chat( + key, + ChatBody( + model=model, + messages=[ChatMessage(role="user", content=WEATHER_PROMPT)], + tools=[WEATHER_TOOL], + max_tokens=1024, + ), + ) + ) + ) + assert message.tool_calls, f"{backend} dropped the tool call: {message}" + call = message.tool_calls[0] + assert call.id, f"tool call carries no id, so a tool result cannot answer it: {call}" + assert call.function.name == "get_weather", f"wrong tool called: {call}" + assert call.function.arguments, f"tool call carries no arguments: {call}" + args = _WeatherArgs.model_validate_json(call.function.arguments) + assert "paris" in args.location.lower(), f"tool arguments lost the location: {args}" diff --git a/tests/e2e/ui/helpers/roundTrip.ts b/tests/e2e/ui/helpers/roundTrip.ts index 8d6e264e622..91e48b7087c 100644 --- a/tests/e2e/ui/helpers/roundTrip.ts +++ b/tests/e2e/ui/helpers/roundTrip.ts @@ -12,17 +12,111 @@ export async function captureRequestBody( match: { method: string; urlIncludes: string }, action: () => Promise, ): Promise> { - const pending = page.waitForRequest((req) => req.method() === match.method && req.url().includes(match.urlIncludes)); + const pending = page.waitForRequest( + (req) => + req.method() === match.method && req.url().includes(match.urlIncludes), + ); await action(); const request = await pending; return JSON.parse(request.postData() ?? "{}") as Record; } /** Reads an endpoint as the master key, so a failure is bad data and not an expired UI token. */ -export async function readBack(page: Page, endpoint: string): Promise { +export async function readBack( + page: Page, + endpoint: string, +): Promise { const res = await page.request.get(endpoint, { headers: { Authorization: `Bearer ${masterKey()}` }, }); expect(res.ok(), `GET ${endpoint}`).toBe(true); return (await res.json()) as T; } + +type OperationOutcome = + | { readonly status: "success" } + | { readonly status: "failure"; readonly error: unknown }; + +type RunFailure = + | { readonly status: "action_failure"; readonly error: unknown } + | { readonly status: "cleanup_failure"; readonly error: unknown } + | { + readonly status: "action_and_cleanup_failure"; + readonly actionError: unknown; + readonly cleanupError: unknown; + }; + +function toRunFailure( + actionOutcome: OperationOutcome, + cleanupOutcome: OperationOutcome, +): RunFailure | null { + if ( + actionOutcome.status === "failure" && + cleanupOutcome.status === "failure" + ) { + return { + status: "action_and_cleanup_failure", + actionError: actionOutcome.error, + cleanupError: cleanupOutcome.error, + }; + } + if (actionOutcome.status === "failure") { + return { status: "action_failure", error: actionOutcome.error }; + } + if (cleanupOutcome.status === "failure") { + return { status: "cleanup_failure", error: cleanupOutcome.error }; + } + return null; +} + +function raiseRunFailure(failure: RunFailure): never { + switch (failure.status) { + case "action_failure": + throw failure.error; + case "cleanup_failure": + throw failure.error; + case "action_and_cleanup_failure": + throw new AggregateError( + [failure.actionError, failure.cleanupError], + "Action and cleanup failed", + ); + } +} + +async function runAction( + action: () => void | Promise, +): Promise { + return Promise.resolve() + .then(action) + .then( + () => ({ status: "success" as const }), + (error: unknown) => ({ status: "failure" as const, error }), + ); +} + +async function runCleanup( + cleanup: () => boolean | Promise, +): Promise { + return Promise.resolve() + .then(cleanup) + .then( + (succeeded) => + succeeded + ? { status: "success" as const } + : { + status: "failure" as const, + error: new Error("Failed to clean up UI E2E resource"), + }, + (error: unknown) => ({ status: "failure" as const, error }), + ); +} + +export async function runWithCleanup( + action: () => void | Promise, + cleanup: () => boolean | Promise, +): Promise { + const actionOutcome = await runAction(action); + const cleanupOutcome = await runCleanup(cleanup); + const failure = toRunFailure(actionOutcome, cleanupOutcome); + if (failure !== null) raiseRunFailure(failure); +} diff --git a/tests/e2e/ui/tests/prompts/addPrompt.spec.ts b/tests/e2e/ui/tests/prompts/addPrompt.spec.ts new file mode 100644 index 00000000000..9d85236c4a6 --- /dev/null +++ b/tests/e2e/ui/tests/prompts/addPrompt.spec.ts @@ -0,0 +1,75 @@ +import { test, expect } from "@playwright/test"; + +import { ADMIN_STORAGE_PATH } from "../../constants"; +import { Page as DashboardPage } from "../../fixtures/pages"; +import { navigateToPage } from "../../helpers/navigation"; +import { runWithCleanup } from "../../helpers/roundTrip"; +import { masterKey, uniqueSuffix } from "../../helpers/traffic"; + +test.use({ storageState: ADMIN_STORAGE_PATH }); + +test.describe("Prompt upload form", () => { + test("uploads a prompt file and reads the created prompt back", async ({ + page, + }) => { + const promptId = `e2e-prompt-${uniqueSuffix()}`; + const promptContent = "Hello {{name}}"; + + await runWithCleanup( + async () => { + await navigateToPage(page, DashboardPage.Prompts); + await page.getByRole("button", { name: "Upload .prompt File" }).click(); + await expect( + page.getByRole("dialog", { name: "Add New Prompt" }), + ).toBeVisible(); + await page.getByLabel("Prompt ID").fill(promptId); + await page.locator('input[type="file"]').setInputFiles({ + name: "e2e.prompt", + mimeType: "text/plain", + buffer: Buffer.from( + `---\nmodel: fake-openai-gpt-4\n---\n${promptContent}\n`, + ), + }); + await expect(page.getByText("Selected: e2e.prompt")).toBeVisible(); + await page.getByRole("button", { name: "Create Prompt" }).click(); + + await expect + .poll(async () => { + const response = await page.request.get( + `/prompts/${encodeURIComponent(promptId)}/info?environment=development`, + { + headers: { Authorization: `Bearer ${masterKey()}` }, + }, + ); + return response.ok(); + }) + .toBe(true); + await expect + .poll(async () => { + const response = await page.request.get( + `/prompts/${encodeURIComponent(promptId)}/info?environment=development`, + { + headers: { Authorization: `Bearer ${masterKey()}` }, + }, + ); + if (!response.ok()) return undefined; + const promptInfo = (await response.json()) as { + raw_prompt_template?: { content?: string }; + }; + return promptInfo.raw_prompt_template?.content; + }) + .toBe(promptContent); + await expect(page.getByText(promptId, { exact: true })).toBeVisible(); + }, + async () => { + const response = await page.request.delete( + `/prompts/${encodeURIComponent(promptId)}?environment=development`, + { + headers: { Authorization: `Bearer ${masterKey()}` }, + }, + ); + return response.ok(); + }, + ); + }); +}); diff --git a/tests/e2e/ui/tests/tagManagement/tagManagement.spec.ts b/tests/e2e/ui/tests/tagManagement/tagManagement.spec.ts new file mode 100644 index 00000000000..fe659080eab --- /dev/null +++ b/tests/e2e/ui/tests/tagManagement/tagManagement.spec.ts @@ -0,0 +1,84 @@ +import { test, expect } from "@playwright/test"; + +import { ADMIN_STORAGE_PATH } from "../../constants"; +import { Page as DashboardPage } from "../../fixtures/pages"; +import { navigateToPage } from "../../helpers/navigation"; +import { + captureRequestBody, + readBack, + runWithCleanup, +} from "../../helpers/roundTrip"; +import { masterKey, uniqueSuffix } from "../../helpers/traffic"; + +test.use({ storageState: ADMIN_STORAGE_PATH }); + +test.describe("Tag management", () => { + test("creates, edits, reopens, and reads back a tag", async ({ page }) => { + const tagName = `e2e-tag-${uniqueSuffix()}`; + const description = "synthetic tag description"; + const updatedDescription = `${description} updated`; + + await runWithCleanup( + async () => { + await navigateToPage(page, DashboardPage.TagManagement); + await page.getByRole("button", { name: "+ Create New Tag" }).click(); + await expect( + page.getByRole("dialog", { name: "Create New Tag" }), + ).toBeVisible(); + await page.getByLabel("Tag Name").fill(tagName); + await page.getByLabel("Description").fill(description); + await page.getByRole("button", { name: "Create Tag" }).click(); + + await expect + .poll(async () => { + const response = await readBack>( + page, + "/tag/list", + ); + return response.some((tag) => tag.name === tagName); + }) + .toBe(true); + await expect(page.getByText(tagName, { exact: true })).toBeVisible(); + + await page.getByText(tagName, { exact: true }).click(); + await expect(page.getByText("Tag Name:")).toBeVisible(); + await page.getByRole("button", { name: "Edit Tag" }).click(); + await page.getByLabel("Description").fill(updatedDescription); + const updateBody = await captureRequestBody( + page, + { method: "POST", urlIncludes: "/tag/update" }, + () => page.getByRole("button", { name: "Save Changes" }).click(), + ); + expect(updateBody).toMatchObject({ + name: tagName, + description: updatedDescription, + }); + + await expect + .poll(async () => { + const infoResponse = await page.request.post("/tag/info", { + headers: { Authorization: `Bearer ${masterKey()}` }, + data: { names: [tagName] }, + }); + expect(infoResponse.ok()).toBe(true); + const info = (await infoResponse.json()) as Record< + string, + { description?: string } + >; + return info[tagName]?.description; + }) + .toBe(updatedDescription); + }, + async () => { + const response = await page.request.post("/tag/delete", { + headers: { + Authorization: `Bearer ${masterKey()}`, + "Content-Type": "application/json", + }, + data: { name: tagName }, + }); + return response.ok(); + }, + ); + }); +}); diff --git a/tests/proxy_behavior/spend/test_autorouter_session_rollup.py b/tests/proxy_behavior/spend/test_autorouter_session_rollup.py index f3c68b489a5..77549b527d8 100644 --- a/tests/proxy_behavior/spend/test_autorouter_session_rollup.py +++ b/tests/proxy_behavior/spend/test_autorouter_session_rollup.py @@ -6,17 +6,24 @@ tests/test_litellm/proxy/db/test_autorouter_session_rollup.py. """ import asyncio +import time import uuid from datetime import datetime, timedelta, timezone -from typing import Final +from types import SimpleNamespace +from typing import Final, TypedDict, cast import pytest from prisma import Prisma +from prisma.errors import RawQueryError +from typing_extensions import ReadOnly from litellm.proxy.db.autorouter_session_rollup import ( AUTOROUTER_BENCHMARKS_SQL, UPSERT_AUTOROUTER_SESSION_SQL, + AutoRouterTurnTransaction, + flush_autorouter_turn_transactions, ) +from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import SpendLogCleanup pytestmark = pytest.mark.asyncio(loop_scope="session") @@ -45,6 +52,7 @@ async def _turn( tier: "str | None" = None, baseline: "str | None" = None, estimated: bool = True, + user_id: str = "", ) -> None: touched: Final = 1 if (hit or ttl is not None or not covered) else 0 await db.execute_raw( @@ -68,6 +76,7 @@ async def _turn( int(estimated), spend if estimated else 0.0, saved if estimated else 0.0, + user_id, ) @@ -217,7 +226,7 @@ async def test_subtotal_coverage_survives_legacy_and_rolling_writers(db, writers assert row["savings_estimated_actual_spend"] == pytest.approx(0.01 * sum(writers)) assert row["savings_estimated_saved_spend"] == pytest.approx(0.02 * sum(writers)) groups: Final = await db.query_raw( - AUTOROUTER_BENCHMARKS_SQL, T0.isoformat(), (T0 + timedelta(days=1)).isoformat(), key + AUTOROUTER_BENCHMARKS_SQL, T0.isoformat(), (T0 + timedelta(days=1)).isoformat(), key, None ) assert len(groups) == 1 assert groups[0]["classifier_cost"] == row["classifier_cost"] @@ -242,7 +251,7 @@ async def test_unknown_and_legacy_turns_preserve_actual_spend_without_entering_t assert row["saved_spend"] == pytest.approx(-0.03) assert row["savings_estimated_baseline_models"] == {"opus": 1} groups: Final = await db.query_raw( - AUTOROUTER_BENCHMARKS_SQL, T0.isoformat(), (T0 + timedelta(days=1)).isoformat(), key + AUTOROUTER_BENCHMARKS_SQL, T0.isoformat(), (T0 + timedelta(days=1)).isoformat(), key, None ) assert len(groups) == 1 for actual in (row, groups[0]): @@ -277,6 +286,7 @@ async def test_the_benchmarks_aggregate_reads_only_overlapping_sessions(db): (T0 - timedelta(days=1)).isoformat(), (T0 + timedelta(days=1)).isoformat(), None, + None, ) matching = [row for row in rows if row["router_name"] == router] assert len(matching) == 1 @@ -304,6 +314,7 @@ async def test_the_benchmarks_aggregate_can_filter_to_one_key(db): (T0 - timedelta(days=1)).isoformat(), (T0 + timedelta(days=1)).isoformat(), first_key, + None, ) matching = [row for row in rows if row["router_name"] == router] assert len(matching) == 1 @@ -317,10 +328,160 @@ async def test_the_benchmarks_aggregate_can_filter_to_one_key(db): (T0 - timedelta(days=1)).isoformat(), (T0 + timedelta(days=1)).isoformat(), f"k-{uuid.uuid4()}", + None, ) assert [row for row in unknown_key_rows if row["router_name"] == router] == [] +class _BenchmarkRow(TypedDict): + sessions: ReadOnly[int] + turns: ReadOnly[int] + same_model_turns: ReadOnly[int] + first_visit_turns: ReadOnly[int] + spend: ReadOnly[float] + saved_spend: ReadOnly[float] + tier_turns: ReadOnly[dict[str, int]] + cache_hits: ReadOnly[int] + savings_estimated_turns: ReadOnly[int] + savings_estimated_actual_spend: ReadOnly[float] + savings_estimated_saved_spend: ReadOnly[float] + + +async def _scoped_benchmarks( + db: Prisma, router: str, user_id: str | None = None, key: str | None = None +) -> tuple[_BenchmarkRow, ...]: + rows: Final = await db.query_raw( + AUTOROUTER_BENCHMARKS_SQL, + (T0 - timedelta(days=1)).isoformat(), + (T0 + timedelta(days=1)).isoformat(), + key, + user_id, + ) + return tuple(cast(_BenchmarkRow, row) for row in rows if row["router_name"] == router) + + +async def test_users_keep_written_identity_across_shared_keys_and_keyless_sessions(db: Prisma) -> None: + router: Final = f"r-{uuid.uuid4()}" + alice: Final = f"u-{uuid.uuid4()}" + bob: Final = f"u-{uuid.uuid4()}" + first_key: Final = f"k-{uuid.uuid4()}" + second_key: Final = f"k-{uuid.uuid4()}" + await _legacy_turn(db, first_key, T0, router=router) + await _turn(db, first_key, "A", T0 + timedelta(seconds=10), router=router, user_id=alice, tier="simple") + await _turn( + db, first_key, "B", T0 + timedelta(seconds=20), router=router, user_id=bob, spend=0.03, saved=0.06, tier="complex" + ) + await _turn(db, second_key, "C", T0, router=router, user_id=alice, spend=0.02, saved=0.04) + await _turn(db, "", "A", T0, router=router, user_id=alice, ttl=300) + await _turn(db, "", "A", T0 + timedelta(seconds=1), router=router, user_id=alice, hit=1) + await _turn(db, "", "B", T0, router=router, user_id=bob, spend=0.04, saved=0.08) + await _turn(db, second_key, "C", T0 - timedelta(days=40), router=router, user_id=alice, session_id="expired") + + alice_rows: Final = await _scoped_benchmarks(db, router, user_id=alice) + bob_rows: Final = await _scoped_benchmarks(db, router, user_id=bob) + global_rows: Final = await _scoped_benchmarks(db, router) + key_rows: Final = await _scoped_benchmarks(db, router, key=first_key) + intersection: Final = await _scoped_benchmarks(db, router, user_id=alice, key=first_key) + assert len(alice_rows) == len(bob_rows) == len(global_rows) == len(key_rows) == len(intersection) == 1 + assert (alice_rows[0]["sessions"], alice_rows[0]["turns"], alice_rows[0]["same_model_turns"]) == (3, 4, 1) + assert (bob_rows[0]["sessions"], bob_rows[0]["turns"], bob_rows[0]["first_visit_turns"]) == (2, 2, 2) + assert alice_rows[0]["spend"] == pytest.approx(0.05) + assert bob_rows[0]["spend"] == pytest.approx(0.07) + assert alice_rows[0]["tier_turns"] == {"simple": 1} + assert bob_rows[0]["tier_turns"] == {"complex": 1} + assert (alice_rows[0]["cache_hits"], bob_rows[0]["cache_hits"]) == (1, 0) + assert (global_rows[0]["sessions"], global_rows[0]["turns"]) == (4, 7) + assert (alice_rows[0]["savings_estimated_turns"], bob_rows[0]["savings_estimated_turns"]) == (4, 2) + assert global_rows[0]["savings_estimated_turns"] == 6 + for scoped in (alice_rows[0], bob_rows[0]): + assert scoped["savings_estimated_actual_spend"] == pytest.approx(scoped["spend"]) + assert scoped["savings_estimated_saved_spend"] == pytest.approx(scoped["saved_spend"]) + assert global_rows[0]["spend"] == pytest.approx(alice_rows[0]["spend"] + bob_rows[0]["spend"] + 0.01) + assert global_rows[0]["saved_spend"] == pytest.approx(alice_rows[0]["saved_spend"] + bob_rows[0]["saved_spend"] + 0.02) + assert global_rows[0]["tier_turns"] == {"simple": 1, "complex": 1} + assert (key_rows[0]["sessions"], key_rows[0]["turns"]) == (1, 3) + assert key_rows[0]["spend"] == pytest.approx(0.05) + assert (intersection[0]["sessions"], intersection[0]["turns"]) == (1, 1) + assert intersection[0]["spend"] == pytest.approx(0.01) + assert await _scoped_benchmarks(db, router, user_id=bob, key=second_key) == () + assert await _scoped_benchmarks(db, router, user_id=f"u-{uuid.uuid4()}") == () + assert await _scoped_benchmarks(db, router, user_id="") == () + + +async def test_a_failed_user_projection_rolls_back_the_keys_increment(db: Prisma) -> None: + key: Final = f"k-{uuid.uuid4()}" + user_id: Final = "".join(str(uuid.uuid4()) for _ in range(200)) + await _turn(db, key, "A", T0) + before: Final = await _row(db, key) + + with pytest.raises(RawQueryError, match=r"index row (requires|size)"): + await _turn(db, key, "B", T0 + timedelta(seconds=1), user_id=user_id) + + assert await _row(db, key) == before + assert await db.query_raw('SELECT user_id FROM "LiteLLM_AutoRouterUserSession" WHERE user_id = $1', user_id) == [] + + first_user: Final = f"u-{uuid.uuid4()}" + second_user: Final = f"u-{uuid.uuid4()}" + turns: Final = tuple( + AutoRouterTurnTransaction( + api_key=key, + user_id=user, + session_id="s1", + router_name="auto-1", + router_type="complexity", + model=model, + turn_at=T0 + timedelta(seconds=second), + total_tokens=100, + spend=0.01, + saved_spend=0.02, + classifier_cost=0.0, + covered=True, + cache_hit=False, + cache_ttl_seconds=None, + cache_touched=False, + ) + for user, model, second in ( + (first_user, "A", 1), + (user_id, "B", 2), + (first_user, "B", 3), + (second_user, "C", 4), + (first_user, "B", 5), + (second_user, "C", 6), + (user_id, "A", 7), + ) + ) + await flush_autorouter_turn_transactions(SimpleNamespace(db=db), tuple(reversed(turns)), n_retry_times=0) + + key_row: Final = await _row(db, key) + assert (key_row["turns"], key_row["last_model"], key_row["unordered_turns"]) == (2, "A", 0) + assert key_row["spend"] == pytest.approx(0.02) + user_rows: Final = await db.query_raw('SELECT * FROM "LiteLLM_AutoRouterUserSession" WHERE api_key = $1', key) + by_user: Final = {row["user_id"]: row for row in user_rows} + assert set(by_user) == {first_user, second_user} + for user, count, model in ((first_user, 3, "B"), (second_user, 2, "C")): + row: Final = by_user[user] + assert (row["turns"], row["same_model_turns"], row["unordered_turns"], row["last_model"]) == (count, 1, 0, model) + assert row["spend"] == pytest.approx(count * 0.01) + assert row["saved_spend"] == pytest.approx(count * 0.02) + + +async def test_user_session_cleanup_keeps_another_users_recent_keyless_session(db: Prisma) -> None: + router: Final = f"r-{uuid.uuid4()}" + expired_user: Final = f"u-{uuid.uuid4()}" + recent_user: Final = f"u-{uuid.uuid4()}" + await _turn(db, "", "A", T0 - timedelta(days=1), router=router, user_id=expired_user) + await _turn(db, "", "A", T0 + timedelta(days=1), router=router, user_id=recent_user) + cleaner: Final = SpendLogCleanup(general_settings={}) + + await cleaner._delete_old_autorouter_user_session_rows( + SimpleNamespace(db=db), T0.replace(tzinfo=timezone.utc), time.monotonic() + 60 + ) + + assert await db.query_raw( + 'SELECT user_id, turns FROM "LiteLLM_AutoRouterUserSession" WHERE router_name = $1', router + ) == [{"user_id": recent_user, "turns": 1}] + + async def test_a_reconfigured_alias_reports_each_router_type_as_its_own_group(db): key = f"k-{uuid.uuid4()}" router = f"r-{uuid.uuid4()}" @@ -334,6 +495,7 @@ async def test_a_reconfigured_alias_reports_each_router_type_as_its_own_group(db (T0 - timedelta(days=1)).isoformat(), (T0 + timedelta(days=1)).isoformat(), None, + None, ) matching = sorted( (row for row in rows if row["router_name"] == router), @@ -418,6 +580,7 @@ async def test_the_benchmarks_aggregate_sums_tier_turns_across_sessions(db): (T0 - timedelta(days=1)).isoformat(), (T0 + timedelta(days=1)).isoformat(), None, + None, ) grouped = next(row for row in rows if row["router_name"] == router) assert grouped["tier_turns"] == {"simple": 2, "complex": 1} @@ -446,6 +609,7 @@ async def test_tier_maps_stay_separate_per_router_type_on_a_reconfigured_alias(d (T0 - timedelta(days=1)).isoformat(), (T0 + timedelta(days=1)).isoformat(), None, + None, ) by_type = {row["router_type"]: row["tier_turns"] for row in rows if row["router_name"] == router} assert by_type == {"complexity": {"medium": 1}, "quality": {"2": 1}} @@ -461,6 +625,7 @@ async def test_a_window_with_no_tiered_turns_aggregates_to_an_empty_map(db): (T0 - timedelta(days=1)).isoformat(), (T0 + timedelta(days=1)).isoformat(), None, + None, ) grouped = next(row for row in rows if row["router_name"] == router) assert grouped["tier_turns"] == {} diff --git a/tests/proxy_behavior/spend/test_baseline_accounting.py b/tests/proxy_behavior/spend/test_baseline_accounting.py index e187a44c29d..3504751d132 100644 --- a/tests/proxy_behavior/spend/test_baseline_accounting.py +++ b/tests/proxy_behavior/spend/test_baseline_accounting.py @@ -56,7 +56,9 @@ def record() -> Callable[..., BaselineAccountingRecord]: }, ) - def create(label: str = "first", started: float = 10000.0, identical: bool = True) -> BaselineAccountingRecord: + def create( + label: str = "first", started: float = 10000.0, identical: bool = True, user_id: str = "" + ) -> BaselineAccountingRecord: return BaselineAccountingRecord( scope="autorouter-baseline:v3:" + run * 2, api_key=run, session_id=run, router_name="test-router", baseline_model="anthropic/claude-opus-5", @@ -76,6 +78,7 @@ def record() -> Callable[..., BaselineAccountingRecord]: total_tokens=6230, spend=0.17, saved_spend=0.0, classifier_cost=0.0, covered=True, cache_hit=False, cache_ttl_seconds=3600, cache_touched=True, baseline_model="anthropic/claude-opus-5", + user_id=user_id, ), daily=DailyBaselineAttribution( date="2026-09-15", api_key=run, model="claude-opus-5", custom_llm_provider="anthropic", @@ -99,21 +102,36 @@ async def _session(db: Prisma, record: BaselineAccountingRecord): return rows[0] +async def _user_sessions(db: Prisma, record: BaselineAccountingRecord) -> dict[str, dict[str, object]]: + rows: Final = await db.query_raw('SELECT * FROM "LiteLLM_AutoRouterUserSession" WHERE api_key=$1', record.api_key) + return {str(row["user_id"]): row for row in rows} + + async def test_late_replay_updates_all_projections_without_rebilling(db: Prisma, record: Callable[..., BaselineAccountingRecord]) -> None: store: Final = _store(db) - late: Final = record("late", 10001.0) - early: Final = record("early", identical=False) + late: Final = record("late", 10001.0, user_id="late-user") + early: Final = record("early", identical=False, user_id="early-user") await _log(db, late) assert await store.append(late) == "recorded" assert await store.project(late.scope) == "published" before: Final = await _session(db, late) assert before["savings_estimated_actual_spend"] == before["spend"] == 0.17 assert before["saved_spend"] == 0.0 + before_users: Final = await _user_sessions(db, late) + assert set(before_users) == {"late-user"} + assert before_users["late-user"]["savings_estimated_turns"] == 1 + assert before_users["late-user"]["savings_estimated_baseline_models"] == {late.baseline_model: 1} await _log(db, early) assert await store.append(early) == "recorded" pending: Final = await _session(db, late) assert pending["spend"] == 0.34 and pending["savings_estimated_turns"] == 0 assert pending["saved_spend"] == pending["savings_estimated_actual_spend"] == 0.0 + pending_users: Final = await _user_sessions(db, late) + assert set(pending_users) == {"late-user", "early-user"} + for user in pending_users.values(): + assert user["turns"] == 1 and user["spend"] == 0.17 + assert user["savings_estimated_turns"] == user["savings_estimated_actual_spend"] == user["saved_spend"] == 0 + assert user["savings_estimated_baseline_models"] == {} waiting: Final = await db.query_raw('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE request_id=$1', late.observation.request_id) assert waiting[0]["metadata"]["autorouter_savings"] is None assert waiting[0]["metadata"]["autorouter_savings_estimate"]["reason"] == "pending_projection" @@ -125,37 +143,69 @@ async def test_late_replay_updates_all_projections_without_rebilling(db: Prisma, assert logs[0]["spend"] == 0.17 assert logs[0]["metadata"]["autorouter_savings_estimate"]["provenance"] == "modeled" assert after["saved_spend"] == pytest.approx(logs[0]["metadata"]["autorouter_savings"]) + after_users: Final = await _user_sessions(db, late) + assert after_users["early-user"] == pending_users["early-user"] + for field in ( + "saved_spend", "savings_estimated_turns", "savings_estimated_actual_spend", + "savings_estimated_saved_spend", "savings_estimated_baseline_models", + ): + assert after_users["late-user"][field] == after[field] + assert after_users["late-user"]["turns"] == 1 and after_users["late-user"]["spend"] == 0.17 for table in ("DailyUserSpend", "DailyTeamSpend", "DailyOrganizationSpend", "DailyEndUserSpend", "DailyAgentSpend", "DailyTagSpend"): rows: Final = await db.query_raw(f'SELECT spend,api_requests,autorouter_savings_spend FROM "LiteLLM_{table}" WHERE api_key=$1', late.api_key) assert rows[0]["spend"] == rows[0]["api_requests"] == 0 assert rows[0]["autorouter_savings_spend"] == pytest.approx(after["saved_spend"]) -async def test_commit_ack_loss_and_concurrent_duplicate_delivery_are_idempotent(db: Prisma, record: Callable[..., BaselineAccountingRecord]) -> None: - event: Final = record() +@pytest.mark.parametrize("attributed", [True, False]) +async def test_commit_ack_loss_and_concurrent_duplicate_delivery_are_idempotent( + db: Prisma, record: Callable[..., BaselineAccountingRecord], attributed: bool +) -> None: + event: Final = record(user_id="first-user" if attributed else "") + other: Final = record("other", 10001.0, user_id="second-user" if attributed else "") await _log(db, event) assert await _store(db, after_commit=True).append(event) == "unavailable" store: Final = _store(db) assert set(await asyncio.gather(*(store.append(event) for _ in range(4)))) == {"recorded"} + await _log(db, other) + assert await store.append(other) == "recorded" + if not attributed: + await db.execute_raw( + 'UPDATE "LiteLLM_AutoRouterBaselineObservation" SET data=(data::jsonb #- \'{turn,user_id}\')::text WHERE scope=$1', + event.scope, + ) assert await store.project(event.scope) == "published" assert await store.project(event.scope) == "unchanged" session: Final = await _session(db, event) - assert session["turns"] == session["savings_estimated_turns"] == 1 - assert session["spend"] == session["savings_estimated_actual_spend"] == 0.17 + assert session["turns"] == session["savings_estimated_turns"] == 2 + assert session["spend"] == session["savings_estimated_actual_spend"] == 0.34 + users: Final = await _user_sessions(db, event) + assert set(users) == ({"first-user", "second-user"} if attributed else set()) + for user in users.values(): + assert user["turns"] == user["savings_estimated_turns"] == 1 + assert user["spend"] == user["savings_estimated_actual_spend"] == 0.17 + assert user["savings_estimated_baseline_models"] == {event.baseline_model: 1} async def test_publication_rollback_keeps_dirty_revision_for_retry(db: Prisma, record: Callable[..., BaselineAccountingRecord]) -> None: - event: Final = record() + event: Final = record(user_id="rollback-user") await _log(db, event) store: Final = _store(db) assert await store.append(event) == "recorded" assert await _store(db, before_commit=True).project(event.scope) == "unavailable" session: Final = await _session(db, event) assert session["spend"] == 0.17 and session["savings_estimated_turns"] == 0 + before_users: Final = await _user_sessions(db, event) + assert before_users["rollback-user"]["spend"] == 0.17 + assert before_users["rollback-user"]["savings_estimated_turns"] == 0 + assert before_users["rollback-user"]["savings_estimated_baseline_models"] == {} revisions: Final = await db.query_raw('SELECT revision,published_revision FROM "LiteLLM_AutoRouterBaselineComparison" WHERE scope=$1', event.scope) assert revisions[0]["revision"] > revisions[0]["published_revision"] assert await store.project(event.scope) == "published" assert (await _session(db, event))["savings_estimated_turns"] == 1 + after_users: Final = await _user_sessions(db, event) + assert after_users["rollback-user"]["turns"] == after_users["rollback-user"]["savings_estimated_turns"] == 1 + assert after_users["rollback-user"]["spend"] == after_users["rollback-user"]["savings_estimated_actual_spend"] == 0.17 async def test_conflicting_duplicate_cannot_restore_an_observed_estimate(db: Prisma, record: Callable[..., BaselineAccountingRecord]) -> None: @@ -196,6 +246,7 @@ async def test_native_observation_enters_spend_pipeline_once_with_shared_daily_a db: Prisma, record: Callable[..., BaselineAccountingRecord], monkeypatch: pytest.MonkeyPatch, ) -> None: import os + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter from litellm.proxy.hooks.autorouter_baseline_cache import CapturedBaselineObservation diff --git a/tests/test_litellm/litellm_core_utils/test_logging_worker.py b/tests/test_litellm/litellm_core_utils/test_logging_worker.py index 1553e788472..2bb93a58531 100644 --- a/tests/test_litellm/litellm_core_utils/test_logging_worker.py +++ b/tests/test_litellm/litellm_core_utils/test_logging_worker.py @@ -180,6 +180,53 @@ class TestLoggingWorker: assert sorted(fired) == ["first", "second"] + @pytest.mark.parametrize("stranded", ["still_queued", "dequeued_never_started"]) + def test_flush_on_new_loop_drains_tasks_stranded_on_previous_loop(self, stranded): + """ + Regression: ``flush()`` from a new event loop used to ``join()`` the queue bound to the + previous loop, whose unfinished counter nothing on the new loop ever decrements. The first + such flush hung until pytest-timeout killed it and every later one raised + ``RuntimeError: ... is bound to a different event loop`` from the queue's Event. + """ + worker = LoggingWorker(timeout=1.0, max_queue_size=10) + callback = AsyncMock() + + async def enqueue_on_first_loop(): + if stranded == "still_queued": + worker._ensure_queue() + worker.enqueue(callback()) + return + worker.ensure_initialized_and_enqueue(callback()) + + asyncio.run(enqueue_on_first_loop()) + assert worker._queue is not None + expected_shape = (1, 0) if stranded == "still_queued" else (0, 1) + assert (worker._queue.qsize(), len(worker._unstarted_dequeued_tasks())) == expected_shape + assert callback.await_count == 0, "precondition: the callback never ran before the first loop closed" + + async def flush_twice_on_second_loop(): + await asyncio.wait_for(worker.flush(), timeout=5) + await asyncio.wait_for(worker.flush(), timeout=5) + + asyncio.run(flush_twice_on_second_loop()) + + assert callback.await_count == 1 + + def test_flush_starts_a_worker_when_the_queue_has_none(self): + """``flush()`` must drain a queue that exists on the current loop without a running worker.""" + worker = LoggingWorker(timeout=1.0, max_queue_size=10) + callback = AsyncMock() + + async def enqueue_then_flush(): + worker._ensure_queue() + worker.enqueue(callback()) + assert worker._worker_task is None, "precondition: nothing is draining the queue yet" + await asyncio.wait_for(worker.flush(), timeout=3) + + asyncio.run(enqueue_then_flush()) + + assert callback.await_count == 1 + def test_flush_on_exit_swallows_cancellation_and_drains_remaining(self): """A callback raising CancelledError must not abort the atexit flush of later events.""" worker = LoggingWorker(timeout=1.0, max_queue_size=10) diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index e768139f04a..0969a913605 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -291,6 +291,106 @@ async def test_find_team_with_model_access_uses_request_method_for_passthrough_a assert "allowed_passthrough_routes" in exc_info.value.detail +_AUTH_ENFORCED_MODEL_HOST_ROUTES: Final = { + "test-uuid-1:subpath:/model-host/v1/extractor:GET,POST": { + "endpoint_id": "test-uuid-1", + "path": "/model-host/v1/extractor", + "type": "subpath", + "auth": True, + }, +} + + +@pytest.mark.asyncio +async def test_find_team_with_model_access_team_allowed_routes_wildcard_grants_auth_passthrough(): + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_allowed_routes=["openai_routes", "/model-host/*"]) + team_without_passthrough_allowlist = LiteLLM_TeamTable(team_id="team-a", models=["all-proxy-models"], metadata={}) + + with ( + patch( + "litellm.proxy.auth.handle_jwt.get_team_object", + new_callable=AsyncMock, + return_value=team_without_passthrough_allowlist, + ), + patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints._registered_pass_through_routes", + _AUTH_ENFORCED_MODEL_HOST_ROUTES, + ), + patch("litellm.proxy.utils.get_server_root_path", return_value="/"), + ): + team_id, team_obj = await JWTAuthManager.find_team_with_model_access( + team_ids={"team-a"}, + requested_model=None, + route="/model-host/v1/extractor/predict", + jwt_handler=jwt_handler, + prisma_client=None, + user_api_key_cache=MagicMock(), + parent_otel_span=None, + proxy_logging_obj=MagicMock(), + request_method="POST", + ) + + assert team_id == "team-a" + assert team_obj == team_without_passthrough_allowlist + + +@pytest.mark.asyncio +async def test_auth_builder_header_team_allows_auth_passthrough_for_team_allowed_routes_wildcard(): + from litellm.proxy.utils import ProxyLogging + + jwt_handler = JWTHandler() + user_api_key_cache = DualCache() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=user_api_key_cache, + litellm_jwtauth=LiteLLM_JWTAuth( + team_ids_jwt_field="groups", + user_id_jwt_field="sub", + team_allowed_routes=["openai_routes", "/model-host/*"], + ), + ) + + with ( + patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, + patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), + patch.object(JWTAuthManager, "check_admin_access", new_callable=AsyncMock, return_value=None), + patch( + "litellm.proxy.auth.handle_jwt.get_team_object", + new_callable=AsyncMock, + return_value=LiteLLM_TeamTable(team_id="team-2", metadata={}), + ), + patch.object( + JWTAuthManager, + "get_objects", + new_callable=AsyncMock, + return_value=(None, None, None, None, "user-1"), + ), + patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints._registered_pass_through_routes", + _AUTH_ENFORCED_MODEL_HOST_ROUTES, + ), + patch("litellm.proxy.utils.get_server_root_path", return_value="/"), + ): + mock_auth_jwt.return_value = {"sub": "user-1", "scope": "", "groups": ["team-1", "team-2"]} + + result = await JWTAuthManager.auth_builder( + api_key="jwt-token", + jwt_handler=jwt_handler, + request_data={}, + general_settings={}, + route="/model-host/v1/extractor/predict", + prisma_client=None, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=ProxyLogging(user_api_key_cache=user_api_key_cache), + request_headers={"x-litellm-team-id": "team-2"}, + request_method="POST", + ) + + assert result["team_id"] == "team-2" + + @pytest.mark.asyncio async def test_auth_builder_proxy_admin_user_role(): """Test that is_proxy_admin is True when user_object.user_role is PROXY_ADMIN""" @@ -6463,6 +6563,90 @@ async def test_auth_builder_db_fallback_enforces_passthrough_route_access(): assert "passthrough route" in exc_info.value.detail +async def _auth_builder_via_db_team_fallback(team_allowed_routes: list[str]): + user_id = "u_passthrough" + user_object = LiteLLM_UserTable( + user_id=user_id, + user_role=LitellmUserRoles.INTERNAL_USER, + teams=["team_no_passthrough"], + ) + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(fallback_to_db_teams=True, team_allowed_routes=team_allowed_routes) + + async def fake_get_team(team_id, **kwargs): + return LiteLLM_TeamTable(team_id=team_id, metadata={}) + + with ( + patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock, return_value={"sub": user_id, "scope": ""}), + patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), + patch.object(jwt_handler, "get_rbac_role", return_value=None), + patch.object(jwt_handler, "get_scopes", return_value=[]), + patch.object(jwt_handler, "get_object_id", return_value=None), + patch.object( + JWTAuthManager, + "get_user_info", + new_callable=AsyncMock, + return_value=(user_id, "u@example.com", True), + ), + patch.object(jwt_handler, "get_org_id", return_value=None), + patch.object(jwt_handler, "get_end_user_id", return_value=None), + patch.object(JWTAuthManager, "check_admin_access", new_callable=AsyncMock, return_value=None), + patch.object( + JWTAuthManager, + "get_objects", + new_callable=AsyncMock, + return_value=(user_object, None, None, None, user_id), + ), + patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), + patch.object(JWTAuthManager, "validate_object_id", return_value=True), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), + patch( + "litellm.proxy.auth.handle_jwt.get_team_object", + new_callable=AsyncMock, + side_effect=fake_get_team, + ), + patch( + "litellm.proxy.auth.handle_jwt.get_team_membership", + new_callable=AsyncMock, + return_value=None, + ), + patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints._registered_pass_through_routes", + _AUTH_ENFORCED_MODEL_HOST_ROUTES, + ), + patch("litellm.proxy.utils.get_server_root_path", return_value="/"), + ): + return await JWTAuthManager.auth_builder( + api_key="test_jwt_token", + jwt_handler=jwt_handler, + request_data={}, + general_settings={"enforce_rbac": False}, + route="/model-host/v1/extractor/predict", + prisma_client=None, + user_api_key_cache=None, + parent_otel_span=None, + proxy_logging_obj=None, + request_headers=None, + request_method="POST", + ) + + +@pytest.mark.asyncio +async def test_auth_builder_db_fallback_team_allowed_routes_wildcard_grants_auth_passthrough(): + result = await _auth_builder_via_db_team_fallback(team_allowed_routes=["openai_routes", "/model-host/*"]) + + assert result["team_id"] == "team_no_passthrough" + + +@pytest.mark.asyncio +async def test_auth_builder_db_fallback_route_groups_alone_do_not_grant_auth_passthrough(): + with pytest.raises(HTTPException) as exc_info: + await _auth_builder_via_db_team_fallback(team_allowed_routes=["openai_routes", "mapped_pass_through_routes"]) + + assert exc_info.value.status_code == 403, exc_info.value.detail + assert "allowed_passthrough_routes" in exc_info.value.detail + + @pytest.mark.asyncio async def test_sync_user_role_and_teams_singular_claim_reconciles_memberships(): """When fallback_to_db_teams is on but the JWT carries a singular team claim diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index c9b8d6764c9..8da93ba341b 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -1249,6 +1249,122 @@ def test_non_proxy_admin_allows_auth_pass_through_with_team_allowlist(): ) +@pytest.mark.parametrize( + "route, team_allowed_routes, expected", + [ + ("/model-host/v1/extractor/predict", ["/model-host/*"], True), + ("/model-host", ["/model-host/*"], False), + ("/model-host/v1/extractor", ["/model-host/v1/extractor"], True), + ("/model-host/v1/extractor/predict", ["/model-host/v1/extractor"], False), + ("/other/v1/extractor", ["/model-host/*"], False), + ("/model-host/v1/extractor", ["openai_routes", "llm_api_routes", "mapped_pass_through_routes"], False), + ("/model-host/v1/extractor", ["*"], False), + ("/model-host/v1/extractor", ["/*"], False), + ("/model-host/v1/extractor", [], False), + ], +) +def test_jwt_team_routes_grant_pass_through_only_for_explicit_paths(route, team_allowed_routes, expected): + assert ( + RouteChecks.jwt_team_routes_grant_pass_through(route=route, team_allowed_routes=team_allowed_routes) + is expected + ) + + +_AUTH_ENFORCED_MODEL_HOST_ROUTES: Final = { + "test-uuid-1:subpath:/model-host/v1/extractor:GET,POST": { + "endpoint_id": "test-uuid-1", + "path": "/model-host/v1/extractor", + "type": "subpath", + "auth": True, + }, +} + + +def _jwt_handler_with_team_allowed_routes(team_allowed_routes: list[str]): + from litellm.proxy._types import LiteLLM_JWTAuth + from litellm.proxy.auth.handle_jwt import JWTHandler + + jwt_handler: Final = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_allowed_routes=team_allowed_routes) + return jwt_handler + + +def _check_model_host_route_as(valid_token: UserAPIKeyAuth, team_allowed_routes: list[str]) -> None: + with ( + patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints._registered_pass_through_routes", + _AUTH_ENFORCED_MODEL_HOST_ROUTES, + ), + patch("litellm.proxy.utils.get_server_root_path", return_value="/"), + patch( + "litellm.proxy.proxy_server.jwt_handler", + _jwt_handler_with_team_allowed_routes(team_allowed_routes), + ), + ): + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=None, + _user_role=LitellmUserRoles.INTERNAL_USER.value, + route="/model-host/v1/extractor/predict", + request=MagicMock(spec=Request), + valid_token=valid_token, + request_data={}, + ) + + +def test_non_proxy_admin_allows_auth_pass_through_for_jwt_team_allowed_routes_wildcard(): + jwt_token: Final = UserAPIKeyAuth( + user_id="test_user", + user_role=LitellmUserRoles.INTERNAL_USER.value, + team_id="team-a", + jwt_claims={"sub": "test_user"}, + ) + + _check_model_host_route_as(jwt_token, team_allowed_routes=["openai_routes", "/model-host/*"]) + + +def test_non_proxy_admin_denies_auth_pass_through_for_jwt_when_only_route_groups_configured(): + jwt_token: Final = UserAPIKeyAuth( + user_id="test_user", + user_role=LitellmUserRoles.INTERNAL_USER.value, + team_id="team-a", + jwt_claims={"sub": "test_user"}, + ) + + with pytest.raises(HTTPException) as exc_info: + _check_model_host_route_as(jwt_token, team_allowed_routes=["openai_routes", "mapped_pass_through_routes"]) + + assert exc_info.value.status_code == 403, exc_info.value.detail + assert "allowed_passthrough_routes" in exc_info.value.detail + + +@pytest.mark.parametrize( + "api_key, team_id, jwt_claims", + [ + ("sk-test-key", "team-a", None), + ("sk-test-key", "team-a", {"sub": "test_user"}), + (None, "team-a", None), + (None, None, {"sub": "test_user"}), + ], + ids=["plain_virtual_key", "jwt_mapped_virtual_key", "keyless_non_jwt_caller", "jwt_without_team"], +) +def test_non_proxy_admin_jwt_team_allowed_routes_grant_pass_through_only_to_jwt_team_callers( + api_key, team_id, jwt_claims +): + caller: Final = UserAPIKeyAuth( + api_key=api_key, + user_id="test_user", + user_role=LitellmUserRoles.INTERNAL_USER.value, + team_id=team_id, + jwt_claims=jwt_claims, + ) + + with pytest.raises(HTTPException) as exc_info: + _check_model_host_route_as(caller, team_allowed_routes=["openai_routes", "/model-host/*"]) + + assert exc_info.value.status_code == 403, exc_info.value.detail + assert "allowed_passthrough_routes" in exc_info.value.detail + + def test_virtual_key_without_llm_api_routes_cannot_access_pass_through(): """ Test that virtual keys without llm_api_routes permission cannot access registered pass-through endpoints. diff --git a/tests/test_litellm/proxy/db/test_autorouter_session_rollup.py b/tests/test_litellm/proxy/db/test_autorouter_session_rollup.py index acd3dc18b54..c61a489f894 100644 --- a/tests/test_litellm/proxy/db/test_autorouter_session_rollup.py +++ b/tests/test_litellm/proxy/db/test_autorouter_session_rollup.py @@ -56,6 +56,31 @@ def _build(payload: dict | None = None, metadata: dict | None = None): class TestBuildTransaction: + @pytest.mark.parametrize( + "api_key, user_id, included", + [ + ("hashed-key", "canonical-user", True), + ("hashed-key", None, True), + ("hashed-key", "", True), + ("", "canonical-user", True), + ("", None, False), + ("", "", False), + ], + ) + def test_attribution_uses_the_canonical_user_even_without_a_key( + self, api_key: str, user_id: str | None, included: bool + ) -> None: + transaction: Final = _build( + payload=_payload(api_key=api_key, user=user_id), + metadata=_metadata(user="client-user", user_api_key_user_id="metadata-user"), + ) + if not included: + assert transaction is None + return + assert transaction is not None + assert transaction.api_key == api_key + assert transaction.user_id == (user_id or "") + def test_successful_auto_routed_turn_builds_every_field(self): transaction = _build( metadata=_metadata( @@ -205,23 +230,43 @@ class TestBuildTransaction: class _FakeDB: - def __init__(self, failures: "list[Exception] | None" = None, poison_session: str | None = None): + def __init__( + self, + failures: "list[Exception] | None" = None, + poison_session: str | None = None, + poison_user: str | None = None, + commit_then_error_users: frozenset[str] = frozenset(), + ): self.calls: list[tuple] = [] + self.attempts: list[tuple[str, tuple[object, ...]]] = [] self._failures = list(failures or []) self._poison_session = poison_session + self._poison_user = poison_user + self._commit_then_error_users = commit_then_error_users async def execute_raw(self, sql: str, *params: object) -> int: + self.attempts.append((sql, params)) if self._poison_session is not None and params[1] == self._poison_session: raise RuntimeError("index row size exceeds btree maximum") + if self._poison_user is not None and params[19] == self._poison_user: + raise RuntimeError("index row size exceeds btree maximum") if self._failures: raise self._failures.pop(0) self.calls.append((sql, params)) + if params[19] in self._commit_then_error_users: + raise RuntimeError("commit succeeded but acknowledgement was lost") return 1 class _FakeClient: - def __init__(self, failures: "list[Exception] | None" = None, poison_session: str | None = None): - self.db = _FakeDB(failures, poison_session) + def __init__( + self, + failures: "list[Exception] | None" = None, + poison_session: str | None = None, + poison_user: str | None = None, + commit_then_error_users: frozenset[str] = frozenset(), + ): + self.db = _FakeDB(failures, poison_session, poison_user, commit_then_error_users) def _transaction( @@ -229,9 +274,11 @@ def _transaction( at: datetime = datetime(2026, 8, 1, 12, 0, 0), tier: str | None = "medium", baseline_model: str | None = "anthropic/claude-opus-5", + api_key: str = "k1", + user_id: str = "", ) -> AutoRouterTurnTransaction: return AutoRouterTurnTransaction( - api_key="k1", + api_key=api_key, session_id=session_id, router_name="live-auto", router_type="complexity", @@ -247,6 +294,7 @@ def _transaction( cache_touched=False, tier=tier, baseline_model=baseline_model, + user_id=user_id, ) @@ -261,7 +309,7 @@ class TestFlush: def test_params_marshal_in_statement_order(self): client = _FakeClient() - asyncio.run(flush_autorouter_turn_transactions(client, [_transaction()])) + asyncio.run(flush_autorouter_turn_transactions(client, [_transaction(user_id="canonical-user")])) sql, params = client.db.calls[0] assert sql == UPSERT_AUTOROUTER_SESSION_SQL assert params == ( @@ -284,8 +332,65 @@ class TestFlush: 0, 0.0, 0.0, + "canonical-user", ) + def test_a_keys_turns_stay_chronological_when_its_canonical_user_changes(self) -> None: + client: Final = _FakeClient() + earlier: Final = _transaction(user_id="z-user", at=datetime(2026, 8, 1, 12, 0, 0)) + later: Final = _transaction(user_id="a-user", at=datetime(2026, 8, 1, 12, 0, 10)) + asyncio.run(flush_autorouter_turn_transactions(client, [later, earlier])) + assert [(params[5], params[19]) for _, params in client.db.calls] == [ + ("2026-08-01T12:00:00", "z-user"), + ("2026-08-01T12:00:10", "a-user"), + ] + + def test_one_keyless_users_failed_session_does_not_drop_another_users_turn(self) -> None: + client: Final = _FakeClient(poison_user="a-user") + failed: Final = _transaction(api_key="", user_id="a-user") + other: Final = _transaction(api_key="", user_id="b-user", at=datetime(2026, 8, 1, 12, 0, 10)) + asyncio.run(flush_autorouter_turn_transactions(client, [other, failed])) + assert [(params[0], params[1], params[19]) for _, params in client.db.calls] == [("", "s1", "b-user")] + + def test_uncertain_commits_quarantine_only_the_key_and_each_failed_user(self) -> None: + client: Final = _FakeClient(commit_then_error_users=frozenset({"a-failed", "c-failed"})) + turns: Final = tuple( + _transaction(user_id=user, at=datetime(2026, 8, 1, 12, 0, second), api_key=key) + for user, second, key in ( + ("b-healthy", 0, "k1"), + ("a-failed", 1, "k1"), + ("b-healthy", 2, "k1"), + ("c-failed", 3, "k1"), + ("b-healthy", 4, "k1"), + ("d-healthy", 5, "k1"), + ("c-failed", 6, "k1"), + ("d-healthy", 7, "k1"), + ("a-failed", 8, "k1"), + ("", 9, "k1"), + ("z-other", 10, "k2"), + ) + ) + asyncio.run(flush_autorouter_turn_transactions(client, tuple(reversed(turns)))) + + assert client.db.attempts == client.db.calls + assert [ + (params[0], params[19], params[5]) + for sql, params in client.db.calls + if sql == UPSERT_AUTOROUTER_SESSION_SQL + ] == [ + ("k1", "b-healthy", "2026-08-01T12:00:00"), + ("k1", "a-failed", "2026-08-01T12:00:01"), + ("k2", "z-other", "2026-08-01T12:00:10"), + ] + assert [params[19] for _, params in client.db.attempts].count("a-failed") == 1 + assert [params[19] for _, params in client.db.attempts].count("c-failed") == 1 + for user, seconds in (("b-healthy", (2, 4)), ("c-failed", (3,)), ("d-healthy", (5, 7))): + assert [ + (params[0], params[5]) + for sql, params in client.db.calls + if sql != UPSERT_AUTOROUTER_SESSION_SQL and params[19] == user + ] == [("k1", f"2026-08-01T12:00:{second:02d}") for second in seconds] + def test_a_connect_error_retries_the_same_statement(self): client = _FakeClient(failures=[httpx.ConnectError("boom")]) asyncio.run(flush_autorouter_turn_transactions(client, [_transaction()])) diff --git a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py index 931531441d3..a282ae731dd 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py @@ -5,6 +5,7 @@ Unit tests for auto router management endpoints from collections.abc import Mapping, Sequence from functools import partial from pathlib import Path +from types import SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock @@ -781,17 +782,43 @@ class TestAutoRouterBenchmarks: assert _summed_agg_row([complexity, quality]).tier_turns == {} @pytest.mark.asyncio - async def test_non_admin_roles_cannot_read_benchmarks(self): + @pytest.mark.parametrize("user_id", [None, "own-user", "other-user"]) + async def test_non_admin_roles_cannot_read_benchmarks(self, user_id: str | None): from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks with pytest.raises(HTTPException) as err: await get_auto_router_benchmarks( - user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, api_key="sk-x"), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, api_key="sk-x", user_id="own-user" + ), start_date="2026-08-01", end_date="2026-08-02", + user_id=user_id, ) assert err.value.status_code == 403 + @pytest.mark.asyncio + async def test_an_empty_user_filter_is_rejected_before_querying_deployment_data( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + import httpx + from fastapi import FastAPI + + from litellm.proxy import proxy_server + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks + + query: Final = AsyncMock(return_value=[]) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=SimpleNamespace(query_raw=query))) + app: Final = FastAPI() + app.get("/auto_router/benchmarks")(get_auto_router_benchmarks) + app.dependency_overrides[user_api_key_auth] = lambda: ADMIN + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client: + response: Final = await client.get("/auto_router/benchmarks", params={"user_id": ""}) + + assert response.status_code == 422 + query.assert_not_awaited() + @pytest.mark.asyncio async def test_a_reversed_window_is_rejected(self, monkeypatch: pytest.MonkeyPatch): from litellm.proxy import proxy_server @@ -807,7 +834,11 @@ class TestAutoRouterBenchmarks: assert err.value.status_code == 400 @pytest.mark.asyncio - async def test_endpoint_returns_groups_and_totals_from_the_rollup(self, monkeypatch: pytest.MonkeyPatch): + @pytest.mark.parametrize("role", [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY]) + @pytest.mark.parametrize("user_id", [None, "selected-user"]) + async def test_endpoint_returns_groups_and_totals_from_the_rollup( + self, monkeypatch: pytest.MonkeyPatch, role: LitellmUserRoles, user_id: str | None + ): from litellm.proxy import proxy_server from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks @@ -822,12 +853,13 @@ class TestAutoRouterBenchmarks: monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": _DB()})()) response = await get_auto_router_benchmarks( - user_api_key_dict=ADMIN, + user_api_key_dict=UserAPIKeyAuth(user_role=role, api_key="sk-admin", user_id="viewer"), start_date="2026-07-01", end_date="2026-08-01", api_key="key-hash", + user_id=user_id, ) - assert captured["params"] == ("2026-07-01T00:00:00", "2026-08-02T00:00:00", "key-hash") + assert captured["params"] == ("2026-07-01T00:00:00", "2026-08-02T00:00:00", "key-hash", user_id) assert response.routers_in_scope == 1 assert response.groups[0].router_name == "live-auto" assert response.groups[0].saved_pct == response.totals.saved_pct == 75.0 diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 80773f314d8..c66d490f122 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -6,7 +6,7 @@ import logging from contextlib import ExitStack from datetime import datetime, timedelta from types import SimpleNamespace -from typing import List, Optional, cast +from typing import Final, List, Optional, cast from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -1614,6 +1614,60 @@ class TestTeamScopedMCPServerAccess: result = await fetch_all_mcp_servers(user_api_key_dict=mock_user_auth, team_id="any-team-id") assert len(result) == 1 + +class TestFetchAllMCPServersOrdering: + def test_display_order_is_case_insensitive_name_then_id(self) -> None: + servers: Final = ( + LiteLLM_MCPServerTable(server_id="s-2", server_name="GitHub", alias="aaa", transport=MCPTransport.http), + LiteLLM_MCPServerTable(server_id="s-1", alias="github", transport=MCPTransport.http), + LiteLLM_MCPServerTable(server_id="s-0", server_name="Slack", alias="zzz", transport=MCPTransport.http), + LiteLLM_MCPServerTable(server_id="confluence", server_name="", alias="", transport=MCPTransport.http), + ) + + ordered: Final = sorted(servers, key=mgmt_endpoints._mcp_server_display_order) + assert [s.server_id for s in ordered] == ["confluence", "s-1", "s-2", "s-0"] + + @pytest.mark.parametrize("team_id", [None, "team-1"]) + @pytest.mark.parametrize("reverse", [False, True]) + @pytest.mark.asyncio + async def test_list_is_sorted_by_display_name_regardless_of_resolution_order( + self, team_id: str | None, reverse: bool + ) -> None: + mock_user_auth: Final = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="admin_user", + ) + servers: Final = ( + generate_mock_mcp_server_db_record(server_id="s-zeta", alias="zeta"), + generate_mock_mcp_server_db_record(server_id="s-alpha", alias="Alpha"), + generate_mock_mcp_server_db_record(server_id="s-mid", alias="mid"), + ) + resolved: Final = list(reversed(servers) if reverse else servers) + mock_manager: Final = MagicMock() + mock_manager.get_all_allowed_mcp_servers = AsyncMock(return_value=resolved) + with ( + patch( # test-quality-ok: the route reads a module-global manager with no injection seam + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + mock_manager, + ), + patch( # test-quality-ok: admin view is derived from module-global proxy settings + "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + return_value=True, + ), + patch( # test-quality-ok: auth contexts need a live prisma client + "litellm.proxy.management_endpoints.mcp_management_endpoints.build_effective_auth_contexts", + AsyncMock(return_value=[mock_user_auth]), + ), + patch( # test-quality-ok: isolate the route's ordering from team database resolution + "litellm.proxy.management_endpoints.mcp_management_endpoints._get_team_scoped_mcp_server_list", + AsyncMock(return_value=resolved), + ), + ): + result: Final = await mgmt_endpoints.fetch_all_mcp_servers( + user_api_key_dict=mock_user_auth, team_id=team_id + ) + assert [s.server_id for s in result] == ["s-alpha", "s-mid", "s-zeta"] + @pytest.mark.asyncio async def test_restricted_virtual_key_cannot_use_team_id_filter(self): """Restricted virtual keys must not bypass access limits via team_id.""" diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index e6b5fb25c3e..5e6c37c41dd 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -1284,11 +1284,11 @@ class TestUpdateModel: "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", new=AsyncMock(return_value=None), ), - patch( + patch( # test-quality-ok: [TQ008] isolate persistence from encryption implementation "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper", side_effect=lambda value: value, ), - patch( + patch( # test-quality-ok: [TQ008] isolate persistence from router reload implementation "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", new=AsyncMock( return_value=ReconcileOutcome(still_desired=None, live_after=None) @@ -4615,7 +4615,7 @@ class TestPatchModelBlockedAuthGate: "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", new=AsyncMock(return_value=None), ), - patch( + patch( # test-quality-ok: [TQ008] isolate persistence from router reload implementation "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", new=AsyncMock( return_value=ReconcileOutcome(still_desired=None, live_after=None) @@ -7284,3 +7284,337 @@ class TestTeamMemberAutoRouterWrites: assert json.loads(written["model_info"])["member_auto_router"] is True assert appended.await_args.kwargs["data"].models == ["new-personal-router"] assert appended.await_args.kwargs["data"].team_id == "member-team" + + +class TestModelManagementActorEdges: + @pytest.mark.asyncio + async def test_add_model_rejects_non_team_internal_user(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import add_new_model + + actor: Final = UserAPIKeyAuth(user_id="internal-user", user_role=LitellmUserRoles.INTERNAL_USER) + prisma: Final = MagicMock() + deployment: Final = Deployment( + model_name="internal-model", + litellm_params=LiteLLM_Params(model="openai/test-model"), + model_info=ModelInfo(id="internal-model-id"), + ) + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: [TQ008] endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: [TQ008] endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: [TQ008] endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.general_settings", {}), # test-quality-ok: [TQ008] endpoint reads proxy-server state through its only test seam + ): + with pytest.raises(ProxyException) as exc_info: + await add_new_model(model_params=deployment, user_api_key_dict=actor) + + assert str(exc_info.value.code) == "403" + assert "permission" in str(exc_info.value).lower() + prisma.db.litellm_proxymodeltable.create.assert_not_called() + + @pytest.mark.asyncio + async def test_add_model_rejects_proxy_admin_viewer(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import add_new_model + + actor: Final = UserAPIKeyAuth( + user_id="view-only-user", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY + ) + prisma: Final = MagicMock() + deployment: Final = Deployment( + model_name="view-only-model", + litellm_params=LiteLLM_Params(model="openai/test-model"), + model_info=ModelInfo(id="view-only-model-id"), + ) + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: [TQ008] endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: [TQ008] endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: [TQ008] endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.general_settings", {}), # test-quality-ok: [TQ008] endpoint reads proxy-server state through its only test seam + ): + with pytest.raises(ProxyException) as exc_info: + await add_new_model(model_params=deployment, user_api_key_dict=actor) + + assert str(exc_info.value.code) == "403" + assert "view-only" in str(exc_info.value).lower() + prisma.db.litellm_proxymodeltable.create.assert_not_called() + + @pytest.mark.asyncio + async def test_add_model_requires_database_storage(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import add_new_model + + actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + prisma: Final = MagicMock() + deployment: Final = Deployment( + model_name="database-disabled-model", + litellm_params=LiteLLM_Params(model="openai/test-model"), + model_info=ModelInfo(id="database-disabled-model-id"), + ) + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: [TQ008] endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.store_model_in_db", False), # test-quality-ok: [TQ008] endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: [TQ008] endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.general_settings", {}), # test-quality-ok: [TQ008] endpoint reads proxy-server state through its only test seam + ): + with pytest.raises(ProxyException) as exc_info: + await add_new_model(model_params=deployment, user_api_key_dict=actor) + + assert str(exc_info.value.code) == "500" + assert "STORE_MODEL_IN_DB" in str(exc_info.value) + prisma.db.litellm_proxymodeltable.create.assert_not_called() + + @pytest.mark.asyncio + async def test_legacy_model_update_persists_changed_field(self): + from litellm.proxy.management_endpoints.model_management_endpoints import update_model + + model_id: Final = "legacy-update-model-id" + existing_row: Final = MagicMock() + existing_row.litellm_params = {"model": "openai/test-model", "timeout": 30} + existing_row.model_dump.return_value = { + "model_name": "legacy-update-model", + "litellm_params": existing_row.litellm_params, + "model_info": {"id": model_id}, + } + existing_row.model_dump_json.return_value = "{}" + updated_row: Final = MagicMock() + updated_row.model_dump_json.return_value = "{}" + prisma: Final = MagicMock() + prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) + prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row) + router: Final = MagicMock() + router.get_model_ids.return_value = [model_id] + actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: [TQ008] update endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.llm_router", router), # test-quality-ok: [TQ008] update endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: [TQ008] update endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: [TQ008] update endpoint reads proxy-server state through its only test seam + patch( # test-quality-ok: [TQ008] isolate persistence from encryption implementation + "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper", + side_effect=lambda value: value, + ), + patch( # test-quality-ok: [TQ008] isolate persistence from router reload implementation + "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", + new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)), + ), + ): + await update_model( + model_params=updateDeployment( + litellm_params=updateLiteLLMParams(timeout=42), + model_info=ModelInfo(id=model_id), + ), + user_api_key_dict=actor, + ) + + written: Final = json.loads( + prisma.db.litellm_proxymodeltable.update.await_args.kwargs["data"]["litellm_params"] + ) + assert written["timeout"] == 42 + assert written["model"] == "openai/test-model" + + @pytest.mark.asyncio + async def test_legacy_model_update_explicit_null_preserves_existing_field(self): + from litellm.proxy.management_endpoints.model_management_endpoints import update_model + + model_id: Final = "legacy-null-model-id" + existing_row: Final = MagicMock() + existing_row.litellm_params = {"model": "openai/test-model", "timeout": 30} + existing_row.model_dump.return_value = { + "model_name": "legacy-null-model", + "litellm_params": existing_row.litellm_params, + "model_info": {"id": model_id}, + } + existing_row.model_dump_json.return_value = "{}" + updated_row: Final = MagicMock() + updated_row.model_dump_json.return_value = "{}" + prisma: Final = MagicMock() + prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) + prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row) + router: Final = MagicMock() + router.get_model_ids.return_value = [model_id] + actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: [TQ008] update endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.llm_router", router), # test-quality-ok: [TQ008] update endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: [TQ008] update endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: [TQ008] update endpoint reads proxy-server state through its only test seam + patch( # test-quality-ok: [TQ008] isolate persistence from encryption implementation + "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper", + side_effect=lambda value: value, + ), + patch( # test-quality-ok: [TQ008] isolate persistence from router reload implementation + "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", + new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)), + ), + ): + await update_model( + model_params=updateDeployment( + litellm_params=updateLiteLLMParams(timeout=None), + model_info=ModelInfo(id=model_id), + ), + user_api_key_dict=actor, + ) + + written: Final = json.loads( + prisma.db.litellm_proxymodeltable.update.await_args.kwargs["data"]["litellm_params"] + ) + assert written["timeout"] == 30 + + @pytest.mark.asyncio + async def test_patch_model_rejects_config_file_model(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import patch_model + + model_id: Final = "config-model-id" + prisma: Final = MagicMock() + prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=None) + prisma.db.litellm_proxymodeltable.update = AsyncMock() + router: Final = MagicMock() + router.get_deployment.return_value = Deployment( + model_name="config-model", + litellm_params=LiteLLM_Params(model="openai/test-model"), + model_info=ModelInfo(id=model_id), + ) + actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: [TQ008] patch endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.llm_router", router), # test-quality-ok: [TQ008] patch endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: [TQ008] patch endpoint reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: [TQ008] patch endpoint reads proxy-server state through its only test seam + ): + with pytest.raises(ProxyException) as exc_info: + await patch_model( + model_id=model_id, + patch_data=updateDeployment( + litellm_params=updateLiteLLMParams(timeout=42), + model_info=ModelInfo(id=model_id), + ), + user_api_key_dict=actor, + ) + + assert str(exc_info.value.code) == "400" + assert "Cannot edit config-based model" in str(exc_info.value) + prisma.db.litellm_proxymodeltable.update.assert_not_awaited() + + @contextlib.contextmanager + def _client_for(self, actor: UserAPIKeyAuth) -> Iterator[TestClient]: + import litellm.proxy.proxy_server as proxy_server + from litellm.proxy.proxy_server import app + + app.dependency_overrides[proxy_server.user_api_key_auth] = lambda: actor + try: + yield TestClient(app) + finally: + app.dependency_overrides.pop(proxy_server.user_api_key_auth, None) + + def test_post_model_new_binds_to_actor_guard(self): + actor: Final = UserAPIKeyAuth(user_id="internal-user", user_role=LitellmUserRoles.INTERNAL_USER) + prisma: Final = MagicMock() + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: [TQ008] route reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: [TQ008] route reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: [TQ008] route reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.general_settings", {}), # test-quality-ok: [TQ008] route reads proxy-server state through its only test seam + self._client_for(actor) as client, + ): + response: Final = client.post( + "/model/new", + json={ + "model_name": "internal-model", + "litellm_params": {"model": "openai/test-model"}, + "model_info": {"id": "internal-model-id"}, + }, + ) + + assert response.status_code == 403 + assert "permission" in response.text.lower() + prisma.db.litellm_proxymodeltable.create.assert_not_called() + + def test_post_legacy_model_update_binds_to_persistence(self): + model_id: Final = "legacy-route-model-id" + existing_row: Final = LiteLLM_ProxyModelTable( + model_id=model_id, + model_name="legacy-route-model", + litellm_params={"model": "openai/test-model", "timeout": 30}, + model_info={"id": model_id}, + created_by="admin", + updated_by="admin", + ) + updated_row: Final = LiteLLM_ProxyModelTable( + model_id=model_id, + model_name="legacy-route-model", + litellm_params={"model": "openai/test-model", "timeout": 42}, + model_info={"id": model_id}, + created_by="admin", + updated_by="admin", + ) + prisma: Final = MagicMock() + prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) + prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row) + router: Final = MagicMock() + router.get_model_ids.return_value = [model_id] + actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: [TQ008] route reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.llm_router", router), # test-quality-ok: [TQ008] route reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: [TQ008] route reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: [TQ008] route reads proxy-server state through its only test seam + patch( # test-quality-ok: [TQ008] isolate persistence from encryption implementation + "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper", + side_effect=lambda value: value, + ), + patch( # test-quality-ok: [TQ008] isolate persistence from router reload implementation + "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", + new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)), + ), + patch( # test-quality-ok: [TQ008] audit logging is outside the persistence contract + "litellm.proxy.management_endpoints.model_management_endpoints.create_object_audit_log", + new=AsyncMock(return_value=None), + ), + self._client_for(actor) as client, + ): + response: Final = client.post( + "/model/update", + json={ + "litellm_params": {"timeout": 42}, + "model_info": {"id": model_id}, + }, + ) + + assert response.status_code == 200, response.text + written: Final = json.loads( + prisma.db.litellm_proxymodeltable.update.await_args.kwargs["data"]["litellm_params"] + ) + assert written["timeout"] == 42 + + def test_patch_config_model_binds_to_patch_route(self): + model_id: Final = "config-route-model-id" + prisma: Final = MagicMock() + prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=None) + prisma.db.litellm_proxymodeltable.update = AsyncMock() + router: Final = MagicMock() + router.get_deployment.return_value = Deployment( + model_name="config-route-model", + litellm_params=LiteLLM_Params(model="openai/test-model"), + model_info=ModelInfo(id=model_id), + ) + actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: [TQ008] route reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.llm_router", router), # test-quality-ok: [TQ008] route reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: [TQ008] route reads proxy-server state through its only test seam + patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: [TQ008] route reads proxy-server state through its only test seam + self._client_for(actor) as client, + ): + response: Final = client.patch( + f"/model/{model_id}/update", + json={ + "litellm_params": {"timeout": 42}, + "model_info": {"id": model_id}, + }, + ) + + assert response.status_code == 400 + assert "Cannot edit config-based model" in response.text + prisma.db.litellm_proxymodeltable.update.assert_not_awaited() diff --git a/tests/test_litellm/proxy/management_endpoints/test_router_settings_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_router_settings_endpoints.py index 308f4d88f02..3fcda310435 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_router_settings_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_router_settings_endpoints.py @@ -4,6 +4,8 @@ Tests for router settings management endpoints. Tests the GET endpoints for router settings and router fields. """ +from collections.abc import Mapping +from typing import Any, Final from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -15,12 +17,23 @@ from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.management_endpoints.router_settings_endpoints import ( get_router_settings, ) +from litellm.proxy.config_resolvers import SettingsStore from litellm.proxy.proxy_server import app from litellm.router import Router client = TestClient(app) +class _StubProxyConfig: + def __init__(self, router_settings: SettingsStore, config_router_settings: Mapping[str, Any]) -> None: + self.router_settings: Final = router_settings + self._config_router_settings: Final = dict(config_router_settings) + + async def get_config(self, config_file_path: str | None = None) -> dict[str, Any]: + del config_file_path + return {"router_settings": dict(self._config_router_settings)} + + class TestRouterSettingsEndpoints: """Test suite for router settings endpoints""" @@ -75,6 +88,31 @@ class TestRouterSettingsEndpoints: assert isinstance(routing_strategy_field["options"], list) assert len(routing_strategy_field["options"]) > 0 + @pytest.mark.asyncio + async def test_get_router_settings_reports_sources( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + store = SettingsStore("router_settings") + store.load_yaml({"routing_strategy": "simple-shuffle"}) + store.apply_db_row("router_settings", {"num_retries": 3}) + monkeypatch.setattr( + proxy_server, + "proxy_config", + _StubProxyConfig( + store, + {"routing_strategy": "simple-shuffle", "num_retries": 3}, + ), + ) + monkeypatch.setattr(proxy_server, "llm_router", None) + + admin_user = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-x" + ) + response = await get_router_settings(user_api_key_dict=admin_user) + + assert response.source["routing_strategy"] == "config" + assert response.source["num_retries"] == "db" + @pytest.mark.asyncio async def test_get_router_settings_includes_routing_groups_from_live_router( self, monkeypatch @@ -102,12 +140,10 @@ class TestRouterSettingsEndpoints: ) monkeypatch.setattr(proxy_server, "llm_router", llm_router) - - async def fake_get_config(self, config_file_path=None): - return {} - monkeypatch.setattr( - proxy_server.ProxyConfig, "get_config", fake_get_config, raising=True + proxy_server, + "proxy_config", + _StubProxyConfig(SettingsStore("router_settings"), {}), ) admin_user = UserAPIKeyAuth( @@ -116,6 +152,8 @@ class TestRouterSettingsEndpoints: response = await get_router_settings(user_api_key_dict=admin_user) assert response.current_values.get("routing_groups") == groups + assert response.current_values["timeout"] is not None + assert response.source["timeout"] == "default" rg_field = next(f for f in response.fields if f.field_name == "routing_groups") assert rg_field.field_value == groups diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_model_metrics.py b/tests/test_litellm/proxy/proxy_server/test_routes_model_metrics.py index 246e2cbba54..b5536b7618c 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_model_metrics.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_model_metrics.py @@ -11,13 +11,18 @@ Pins (PR2): from __future__ import annotations +from collections.abc import Callable, Mapping +from contextlib import AbstractContextManager from unittest.mock import AsyncMock, MagicMock import pytest +from fastapi.testclient import TestClient import litellm from litellm.proxy import proxy_server from litellm.proxy._types import LitellmUserRoles +from litellm.proxy.config_resolvers.settings_rules import JsonValue +from litellm.proxy.config_resolvers.settings_store import SettingsStore from .conftest import normalize # type: ignore[import-not-found] @@ -179,6 +184,177 @@ def test_model_settings_method_not_allowed(client, auth_as): # --------------------------------------------------------------------------- +def _alerting_client( + monkeypatch: pytest.MonkeyPatch, + *, + yaml_values: Mapping[str, JsonValue], + db_row: Mapping[str, JsonValue], + live_args: Mapping[str, JsonValue], +) -> "SettingsStore": + pc = MagicMock() + row = MagicMock() + row.param_value = db_row + pc.db.litellm_config.find_first = AsyncMock(return_value=row) + monkeypatch.setattr(proxy_server, "prisma_client", pc) + + logging_obj = MagicMock() + args_model = MagicMock() + args_model.model_dump = MagicMock(return_value=live_args) + logging_obj.slack_alerting_instance.alerting_args = args_model + monkeypatch.setattr(proxy_server, "proxy_logging_obj", logging_obj) + + store = SettingsStore("general_settings") + store.load_yaml(yaml_values) + store.apply_db_row("general_settings", db_row) + monkeypatch.setattr(proxy_server.proxy_config, "settings", store) + monkeypatch.setattr(proxy_server, "general_settings", store) + return store + + +def test_alerting_settings_reports_sources( + client: TestClient, + auth_as: Callable[..., AbstractContextManager[None]], + monkeypatch: pytest.MonkeyPatch, +) -> None: + _alerting_client( + monkeypatch, + yaml_values={ + "alerting": ["slack"], + "alerting_args": {"daily_report_frequency": 3, "report_check_interval": 300}, + }, + db_row={"alerting_args": {"daily_report_frequency": 7, "outage_alert_ttl": 4242}}, + live_args={"daily_report_frequency": 3}, + ) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/alerting/settings") + + assert response.status_code == 200 + by_name = {entry["field_name"]: entry for entry in response.json()} + + assert by_name["slack_alerting"]["source"] == "config" + assert by_name["daily_report_frequency"]["source"] == "config" + assert by_name["report_check_interval"]["source"] == "config" + assert by_name["outage_alert_ttl"]["source"] == "default" + assert by_name["budget_alert_ttl"]["source"] == "default" + + +def test_alerting_settings_reports_db_source_when_the_file_omits_alerting_args( + client: TestClient, + auth_as: Callable[..., AbstractContextManager[None]], + monkeypatch: pytest.MonkeyPatch, +) -> None: + store = _alerting_client( + monkeypatch, + yaml_values={"alerting": ["slack"]}, + db_row={ + "alerting_args": { + "outage_alert_ttl": 4242, + "region_outage_alert_ttl": [], + "report_check_interval": None, + } + }, + live_args={"outage_alert_ttl": 4242}, + ) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/alerting/settings") + + assert response.status_code == 200 + by_name = {entry["field_name"]: entry for entry in response.json()} + + assert store.source("alerting_args") == "db" + assert by_name["outage_alert_ttl"]["source"] == "db" + assert by_name["region_outage_alert_ttl"]["source"] == "db" + assert by_name["report_check_interval"]["source"] == "db" + assert by_name["budget_alert_ttl"]["source"] == "default" + + +def test_alerting_settings_reports_config_source_when_db_disagrees( + client: TestClient, + auth_as: Callable[..., AbstractContextManager[None]], + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy.config_resolvers import SettingsStore + + db_alerting_args = {"daily_report_frequency": 7} + + pc = MagicMock() + row = MagicMock() + row.param_value = {"alerting_args": db_alerting_args} + pc.db.litellm_config.find_first = AsyncMock(return_value=row) + monkeypatch.setattr(proxy_server, "prisma_client", pc) + + logging_obj = MagicMock() + args_model = MagicMock() + args_model.model_dump = MagicMock(return_value={"daily_report_frequency": 3}) + logging_obj.slack_alerting_instance.alerting_args = args_model + monkeypatch.setattr(proxy_server, "proxy_logging_obj", logging_obj) + + store = SettingsStore("general_settings") + store.load_yaml({"alerting_args": {"daily_report_frequency": 3}}) + store.apply_db_row("general_settings", {"alerting_args": db_alerting_args}) + monkeypatch.setattr(proxy_server.proxy_config, "settings", store) + monkeypatch.setattr(proxy_server, "general_settings", store) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/alerting/settings") + + assert response.status_code == 200 + by_name = {entry["field_name"]: entry for entry in response.json()} + assert store.source("alerting_args") == "config" + assert by_name["daily_report_frequency"]["field_value"] == 3 + assert by_name["daily_report_frequency"]["source"] == "config" + + +@pytest.mark.parametrize("db_alerting_args", [None, []]) +def test_alerting_settings_handles_empty_db_args( + client: TestClient, + auth_as: Callable[..., AbstractContextManager[None]], + monkeypatch: pytest.MonkeyPatch, + db_alerting_args: JsonValue, +) -> None: + from litellm.proxy.config_resolvers import SettingsStore + + pc = MagicMock() + row = MagicMock() + row.param_value = {"alerting_args": db_alerting_args} + pc.db.litellm_config.find_first = AsyncMock(return_value=row) + monkeypatch.setattr(proxy_server, "prisma_client", pc) + + logging_obj = MagicMock() + args_model = MagicMock() + args_model.model_dump = MagicMock(return_value={}) + logging_obj.slack_alerting_instance.alerting_args = args_model + monkeypatch.setattr(proxy_server, "proxy_logging_obj", logging_obj) + + store = SettingsStore("general_settings") + store.load_yaml({"alerting_args": {"report_check_interval": 300}}) + monkeypatch.setattr(proxy_server.proxy_config, "settings", store) + monkeypatch.setattr(proxy_server, "general_settings", store) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/alerting/settings") + + assert response.status_code == 200 + by_name = {entry["field_name"]: entry for entry in response.json()} + assert by_name["report_check_interval"]["source"] == "config" + assert by_name["budget_alert_ttl"]["source"] == "default" + + +@pytest.mark.parametrize( + ("field_default", "expected"), + [(43200, "default"), (None, "unset")], +) +def test_nested_setting_source_without_a_config_or_db_value(field_default: JsonValue, expected: str) -> None: + store = SettingsStore("general_settings") + store.load_yaml({}) + + assert ( + proxy_server._nested_setting_source(store, {}, "alerting_args", "budget_alert_ttl", field_default) == expected + ) + + def test_alerting_settings_no_db_error(client, auth_as, no_prisma): """Pins ``GET /alerting/settings`` (error: db not connected).""" with auth_as(LitellmUserRoles.PROXY_ADMIN): diff --git a/tests/test_litellm/proxy/shutdown/test_scheduled_jobs.py b/tests/test_litellm/proxy/shutdown/test_scheduled_jobs.py new file mode 100644 index 00000000000..fbce38db39f --- /dev/null +++ b/tests/test_litellm/proxy/shutdown/test_scheduled_jobs.py @@ -0,0 +1,158 @@ +import asyncio +import logging +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager +from datetime import datetime, timedelta + +import pytest +from apscheduler.schedulers.asyncio import AsyncIOScheduler + +from litellm.proxy.shutdown.scheduled_jobs import ( + AwaitableAsyncIOExecutor, + pause_scheduled_jobs, + stop_in_flight_scheduler_jobs, +) + + +class _Job: + """A scheduled job that blocks until cancelled, or for ``work_seconds``, and records what it observed""" + + def __init__(self, swallow_cancellation: bool = False, work_seconds: float | None = None) -> None: + self.started = asyncio.Event() + self.events: list[str] = [] + self.swallow_cancellation = swallow_cancellation + self.work_seconds = work_seconds + + async def run(self) -> None: + self.started.set() + try: + if self.work_seconds is None: + await asyncio.Event().wait() + else: + await asyncio.sleep(self.work_seconds) + self.events.append("committed") + except asyncio.CancelledError: + self.events.append("cancelled") + if self.swallow_cancellation: + await asyncio.Event().wait() + raise + finally: + self.events.append("finished") + + +@asynccontextmanager +async def _running_scheduler(*jobs: _Job) -> AsyncIterator[tuple[AsyncIOScheduler, AwaitableAsyncIOExecutor]]: + """A started scheduler with every job in flight, stopped on the way out whatever the test did""" + executor = AwaitableAsyncIOExecutor() + scheduler = AsyncIOScheduler(executors={"default": executor}) + for index, job in enumerate(jobs): + scheduler.add_job(job.run, id=f"job-{index}", next_run_time=datetime.now()) + scheduler.start() + try: + for job in jobs: + await asyncio.wait_for(job.started.wait(), timeout=5) + yield scheduler, executor + finally: + if scheduler.running: + scheduler.shutdown(wait=False) + stragglers = executor.in_flight_jobs() + for straggler in stragglers: + straggler.cancel() + await asyncio.gather(*stragglers, return_exceptions=True) + + +@pytest.mark.asyncio +async def test_in_flight_jobs_observe_cancellation_before_shutdown_returns(): + """The job's own CancelledError handler records how a run ended, so shutdown must wait for it""" + job = _Job() + async with _running_scheduler(job) as (scheduler, executor): + await stop_in_flight_scheduler_jobs(scheduler, executor) + + assert job.events == ["cancelled", "finished"] + assert scheduler.running is False + assert executor.in_flight_jobs() == () + + +@pytest.mark.asyncio +async def test_a_job_that_is_finishing_is_allowed_to_finish_rather_than_cancelled(): + """A spend write cancelled mid-commit drops the rows it popped, so short jobs get to finish first""" + write = _Job(work_seconds=0.2) + stuck = _Job() + async with _running_scheduler(write, stuck) as (scheduler, executor): + await stop_in_flight_scheduler_jobs(scheduler, executor, finish_timeout_seconds=2.0) + + assert write.events == ["committed", "finished"] + assert stuck.events == ["cancelled", "finished"] + assert scheduler.running is False + + +@pytest.mark.asyncio +async def test_every_in_flight_job_is_cancelled_not_only_the_first(): + first, second = _Job(), _Job() + async with _running_scheduler(first, second) as (scheduler, executor): + await stop_in_flight_scheduler_jobs(scheduler, executor) + + assert first.events == ["cancelled", "finished"] + assert second.events == ["cancelled", "finished"] + + +@pytest.mark.asyncio +async def test_a_job_that_ignores_cancellation_is_abandoned_after_the_timeout(caplog): + """A job that swallows CancelledError must not hold the pod past its termination grace period""" + job = _Job(swallow_cancellation=True) + async with _running_scheduler(job) as (scheduler, executor): + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + await stop_in_flight_scheduler_jobs(scheduler, executor, cancel_timeout_seconds=0.05) + + assert job.events == ["cancelled"] + assert "1 scheduled job(s) did not finish within 0.05s of cancellation" in caplog.text + + +@pytest.mark.asyncio +async def test_shutdown_with_nothing_in_flight_still_stops_the_scheduler(): + async with _running_scheduler() as (scheduler, executor): + await stop_in_flight_scheduler_jobs(scheduler, executor) + await asyncio.sleep(0) + + assert scheduler.running is False + + +@pytest.mark.asyncio +async def test_a_scheduler_that_never_started_is_left_alone(): + """The proxy runs without a scheduler when it has no database""" + executor = AwaitableAsyncIOExecutor() + scheduler = AsyncIOScheduler(executors={"default": executor}) + + await stop_in_flight_scheduler_jobs(scheduler, executor) + + assert scheduler.running is False + + +@pytest.mark.asyncio +async def test_pausing_stops_new_jobs_from_starting_but_leaves_running_ones_alone(): + """A job due during the shutdown drain would only be cancelled, so it must not start at all""" + running = _Job() + async with _running_scheduler(running) as (scheduler, executor): + late = _Job() + scheduler.add_job(late.run, id="late", next_run_time=datetime.now() + timedelta(seconds=0.1)) + + pause_scheduled_jobs(scheduler) + await asyncio.sleep(0.3) + + assert late.started.is_set() is False + assert running.events == [] + assert scheduler.running is True + + await stop_in_flight_scheduler_jobs(scheduler, executor) + + assert running.events == ["cancelled", "finished"] + assert late.started.is_set() is False + + +@pytest.mark.asyncio +async def test_pausing_a_scheduler_that_never_started_is_a_no_op(): + scheduler = AsyncIOScheduler(executors={"default": AwaitableAsyncIOExecutor()}) + + pause_scheduled_jobs(scheduler) + + assert scheduler.running is False diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 950a6cc3c40..2792e176e0e 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -11766,6 +11766,66 @@ def test_prompt_caching_settings_propagate_on_config_reload(monkeypatch, field_n assert getattr(litellm, field_name) == db_value +@pytest.mark.asyncio +async def test_db_stored_datadog_redaction_settings_apply_before_logger_init(monkeypatch: pytest.MonkeyPatch): + """A DB-only litellm_settings row that pairs success_callback: ["datadog"] with + datadog_params.turn_off_message_logging: true must build the DataDogLogger redacted, the + same as the identical block in YAML. Regression for the redaction keys being absent from + the safe-override allowlist while the callback half of the row was honoured.""" + import litellm.proxy.proxy_server as ps + from litellm.integrations.datadog.datadog import DataDogLogger + from litellm.litellm_core_utils import litellm_logging + + monkeypatch.setenv("DD_API_KEY", "test-key") + monkeypatch.setenv("DD_SITE", "us5.datadoghq.com") + monkeypatch.setattr(litellm, "datadog_params", None) + monkeypatch.setattr(litellm, "turn_off_message_logging", False) + monkeypatch.setattr(litellm, "success_callback", []) + monkeypatch.setattr(litellm, "_async_success_callback", []) + monkeypatch.setattr(litellm, "failure_callback", []) + monkeypatch.setattr(litellm, "_async_failure_callback", []) + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm_logging, "_in_memory_loggers", []) + + db_row = { + "success_callback": ["datadog"], + "datadog_params": {"turn_off_message_logging": True}, + "turn_off_message_logging": True, + } + pc = ps.ProxyConfig() + pc._apply_litellm_settings_db_values(pc._prepared_db_settings_values("litellm_settings", db_row)) + pc._add_callbacks_from_db_config({"litellm_settings": db_row}) + + datadog_loggers = [cb for cb in litellm.success_callback if isinstance(cb, DataDogLogger)] + assert len(datadog_loggers) == 1 + assert datadog_loggers[0].turn_off_message_logging is True + assert litellm.turn_off_message_logging is True + + +@pytest.mark.parametrize( + "field_name", + [ + "datadog_params", + "datadog_llm_observability_params", + "newrelic_params", + "pointfive_params", + "aws_sqs_callback_params", + ], +) +def test_db_stored_callback_params_propagate_to_litellm_module(monkeypatch: pytest.MonkeyPatch, field_name: str): + """Every callback init params block stored in the DB litellm_settings row must land on the + litellm module before the matching logger is built, so the DB row behaves like YAML.""" + import litellm.proxy.proxy_server as ps + + monkeypatch.setattr(litellm, field_name, None) + db_value = {"turn_off_message_logging": True} + + pc = ps.ProxyConfig() + pc._apply_litellm_settings_db_values(pc._prepared_db_settings_values("litellm_settings", {field_name: db_value})) + + assert getattr(litellm, field_name) == db_value + + def test_get_config_list_marks_untouched_prompt_caching_flag_as_not_set(monkeypatch): """The flag defaults to False rather than None, so a plain 'is not None' check would report the default as 'In Config' and imply an admin had set it.""" diff --git a/tests/test_litellm/proxy/test_spend_log_cleanup.py b/tests/test_litellm/proxy/test_spend_log_cleanup.py index bf1538183ab..e333da03950 100644 --- a/tests/test_litellm/proxy/test_spend_log_cleanup.py +++ b/tests/test_litellm/proxy/test_spend_log_cleanup.py @@ -7,6 +7,7 @@ import math import time from contextlib import asynccontextmanager from datetime import datetime, timedelta, timezone +from typing import Final from unittest.mock import AsyncMock, MagicMock import pytest @@ -793,18 +794,20 @@ async def test_spend_logs_retention_alone_does_not_touch_the_session_rollup(): tables = [call[0][0] for call in client.db.execute_raw.call_args_list] assert any('"LiteLLM_SpendLogs"' in sql for sql in tables) assert not any('"LiteLLM_AutoRouterSession"' in sql for sql in tables) + assert not any('"LiteLLM_AutoRouterUserSession"' in sql for sql in tables) assert not any('"LiteLLM_HealthCheckTable"' in sql for sql in tables) @pytest.mark.asyncio -async def test_session_retention_alone_cleans_only_the_session_rollup(): - client = _mock_prisma_for_retention([0]) +async def test_session_retention_alone_cleans_both_session_rollups(): + client = _mock_prisma_for_retention([0, 0]) cleaner = SpendLogCleanup(general_settings={"maximum_autorouter_session_retention_period": "365d"}) cleaner.pod_lock_manager = None await cleaner.cleanup_old_spend_logs(client) tables = [call[0][0] for call in client.db.execute_raw.call_args_list] - assert len(tables) == 1 + assert len(tables) == 2 assert '"LiteLLM_AutoRouterSession"' in tables[0] + assert '"LiteLLM_AutoRouterUserSession"' in tables[1] @pytest.mark.asyncio @@ -825,7 +828,7 @@ async def test_health_check_retention_alone_cleans_only_the_health_check_table() @pytest.mark.asyncio async def test_each_retention_key_cuts_off_at_its_own_horizon(): - client = _mock_prisma_for_retention([0, 0, 0, 0]) + client = _mock_prisma_for_retention([0, 0, 0, 0, 0]) cleaner = SpendLogCleanup( general_settings={ "maximum_spend_logs_retention_period": "7d", @@ -839,6 +842,8 @@ async def test_each_retention_key_cuts_off_at_its_own_horizon(): ( "LiteLLM_AutoRouterSession" if '"LiteLLM_AutoRouterSession"' in call[0][0] + else "LiteLLM_AutoRouterUserSession" + if '"LiteLLM_AutoRouterUserSession"' in call[0][0] else "LiteLLM_HealthCheckTable" if '"LiteLLM_HealthCheckTable"' in call[0][0] else "logs" @@ -848,6 +853,7 @@ async def test_each_retention_key_cuts_off_at_its_own_horizon(): now = datetime.now(timezone.utc) assert (now - cutoffs["logs"]).days == 7 assert (now - cutoffs["LiteLLM_AutoRouterSession"]).days == 365 + assert cutoffs["LiteLLM_AutoRouterUserSession"] == cutoffs["LiteLLM_AutoRouterSession"] assert (now - cutoffs["LiteLLM_HealthCheckTable"]).days == 30 @@ -1417,3 +1423,125 @@ def test_the_reported_run_outcome_is_the_most_significant_reason_in_any_order(st """ results = tuple(TableCleanupResult(rows_deleted=0, stop_reason=reason) for reason in stop_reasons) assert SpendLogCleanup._run_outcome(results) == expected + + +_OTHER_OUTCOMES: Final = ("completed", "budget_exhausted", "batch_cap_reached", "skipped_locked", "skipped_disabled") + + +def _runs_recorded(outcome: str) -> float: + """The real ``litellm_spend_log_cleanup_runs_total`` sample for one outcome, 0 when unset""" + from prometheus_client import REGISTRY + + return REGISTRY.get_sample_value("litellm_spend_log_cleanup_runs_total", {"outcome": outcome}) or 0.0 + + +@pytest.mark.asyncio +async def test_a_cancelled_run_records_aborted_and_logs_its_progress_before_re_raising(monkeypatch): + """A run cut short by shutdown must leave its outcome and how far it got behind""" + import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module + + mock_logger = MagicMock() + monkeypatch.setattr(cleanup_module, "verbose_proxy_logger", mock_logger) + aborted_runs_before = _runs_recorded("aborted") + other_runs_before = {outcome: _runs_recorded(outcome) for outcome in _OTHER_OUTCOMES} + + third_batch_reached = asyncio.Event() + + async def _execute_raw(sql, *args): + if third_batch_reached.is_set(): + raise AssertionError("no batch may be issued after the cancelled one") + if _execute_raw.calls < 2: + _execute_raw.calls += 1 + return 150 + third_batch_reached.set() + await asyncio.Event().wait() + + _execute_raw.calls = 0 + mock_prisma_client = MagicMock() + _wire_tx(mock_prisma_client.db) + mock_prisma_client.db.execute_raw = _execute_raw + + cleaner = SpendLogCleanup(general_settings={"maximum_spend_logs_retention_period": "7d"}) + cleaner.pod_lock_manager = MagicMock() + cleaner.pod_lock_manager.redis_cache = MagicMock() + cleaner.pod_lock_manager.acquire_lock = AsyncMock(return_value=True) + cleaner.pod_lock_manager.release_lock = AsyncMock() + + run = asyncio.ensure_future(cleaner.cleanup_old_spend_logs(mock_prisma_client)) + await asyncio.wait_for(third_batch_reached.wait(), timeout=5) + run.cancel() + with pytest.raises(asyncio.CancelledError): + await run + + assert _runs_recorded("aborted") == aborted_runs_before + 1 + assert {outcome: _runs_recorded(outcome) for outcome in _OTHER_OUTCOMES} == other_runs_before + cleaner.pod_lock_manager.release_lock.assert_awaited_once() + mock_logger.exception.assert_not_called() + (error_call,) = mock_logger.error.call_args_list + rendered = error_call[0][0] % error_call[0][1:] + assert rendered.startswith("Spend log cleanup cancelled after ") + assert "s (rows_deleted=300, batches=2)" in rendered + + +@pytest.mark.asyncio +async def test_progress_reported_for_a_cancelled_run_is_that_run_only(monkeypatch): + """The scheduler holds one cleaner for the life of the process, so progress must not carry over""" + import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module + + mock_logger = MagicMock() + monkeypatch.setattr(cleanup_module, "verbose_proxy_logger", mock_logger) + + mock_prisma_client = MagicMock() + _wire_tx(mock_prisma_client.db) + mock_prisma_client.db.execute_raw = AsyncMock(side_effect=[150, 0, 0]) + cleaner = SpendLogCleanup(general_settings={"maximum_spend_logs_retention_period": "7d"}) + cleaner.pod_lock_manager = None + await cleaner.cleanup_old_spend_logs(mock_prisma_client) + + mock_prisma_client.db.execute_raw = AsyncMock(side_effect=[150, asyncio.CancelledError()]) + with pytest.raises(asyncio.CancelledError): + await cleaner.cleanup_old_spend_logs(mock_prisma_client) + + (error_call,) = mock_logger.error.call_args_list + rendered = error_call[0][0] % error_call[0][1:] + assert "(rows_deleted=150, batches=1)" in rendered + + +@pytest.mark.asyncio +async def test_progress_reported_by_an_overlapping_run_is_its_own(monkeypatch): + """With APSCHEDULER_MAX_INSTANCES above one, two runs share the cleaner but not their progress""" + import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module + + mock_logger = MagicMock() + monkeypatch.setattr(cleanup_module, "verbose_proxy_logger", mock_logger) + + first_batch_done = asyncio.Event() + second_run_done = asyncio.Event() + + async def _slow_execute_raw(sql, *args): + first_batch_done.set() + await second_run_done.wait() + return 100 + + slow_client = MagicMock() + _wire_tx(slow_client.db) + slow_client.db.execute_raw = _slow_execute_raw + fast_client = MagicMock() + _wire_tx(fast_client.db) + fast_client.db.execute_raw = AsyncMock(side_effect=[150, 150, 0, 0]) + + cleaner = SpendLogCleanup(general_settings={"maximum_spend_logs_retention_period": "7d"}) + cleaner.pod_lock_manager = None + + slow_run = asyncio.ensure_future(cleaner.cleanup_old_spend_logs(slow_client)) + await asyncio.wait_for(first_batch_done.wait(), timeout=5) + await cleaner.cleanup_old_spend_logs(fast_client) + second_run_done.set() + await asyncio.sleep(0) + slow_run.cancel() + with pytest.raises(asyncio.CancelledError): + await slow_run + + (error_call,) = mock_logger.error.call_args_list + rendered = error_call[0][0] % error_call[0][1:] + assert "(rows_deleted=100, batches=1)" in rendered diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index 0f57af7f82c..eecee2fd0f1 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -1342,6 +1342,48 @@ class TestProxySettingEndpoints: where={"id": "ui_settings"} ) + def test_get_ui_settings_reports_sources(self, monkeypatch: pytest.MonkeyPatch) -> None: + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy import proxy_server + from litellm.proxy.config_resolvers import SettingsStore + + mock_prisma = MagicMock() + mock_db_record = MagicMock() + mock_db_record.ui_settings = { + "disable_model_add_for_internal_users": True, + "require_auth_for_public_ai_hub": True, + } + mock_prisma.db.litellm_uisettings.find_unique = AsyncMock( + return_value=mock_db_record + ) + monkeypatch.setattr(proxy_server, "prisma_client", mock_prisma) + + store = SettingsStore("general_settings") + store.load_yaml( + { + "disable_model_add_for_internal_users": False, + "forward_client_headers_to_llm_api": True, + } + ) + store.apply_db_row( + "ui_settings", + {"disable_model_add_for_internal_users": True}, + ) + monkeypatch.setattr(proxy_server.proxy_config, "settings", store) + monkeypatch.setattr(proxy_server, "general_settings", store) + + response = client.get("/get/ui_settings") + + assert response.status_code == 200 + data = response.json() + assert data["values"]["disable_model_add_for_internal_users"] is False + assert data["values"]["forward_client_headers_to_llm_api"] is True + assert data["values"]["require_auth_for_public_ai_hub"] is True + assert data["source"]["disable_model_add_for_internal_users"] == "config" + assert data["source"]["forward_client_headers_to_llm_api"] == "config" + assert data["source"]["require_auth_for_public_ai_hub"] == "db" + def test_get_ui_settings_schema_description_preserved_with_extensions( self, mock_auth, monkeypatch ): @@ -3477,6 +3519,7 @@ class TestPtuCostAttributionUISetting: assert response.status_code == 200 assert response.json()["values"]["enable_ptu_cost_attribution"] is False + assert response.json()["source"]["enable_ptu_cost_attribution"] == "default" def test_reported_true_once_the_env_var_is_set(self, mock_auth, monkeypatch): from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR @@ -3488,6 +3531,47 @@ class TestPtuCostAttributionUISetting: assert response.status_code == 200 assert response.json()["values"]["enable_ptu_cost_attribution"] is True + assert response.json()["source"]["enable_ptu_cost_attribution"] == "config" + + def test_reported_config_when_secret_manager_enables_the_flag( + self, mock_auth: None, monkeypatch: pytest.MonkeyPatch + ) -> None: + from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR + + monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False) + monkeypatch.setattr( + "litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints.is_ptu_cost_attribution_enabled", + lambda: True, + ) + self._mock_prisma(monkeypatch) + + response = client.get("/get/ui_settings") + + assert response.status_code == 200 + assert response.json()["values"]["enable_ptu_cost_attribution"] is True + assert response.json()["source"]["enable_ptu_cost_attribution"] == "config" + + def test_reported_config_when_secret_manager_disables_the_flag( + self, mock_auth: None, monkeypatch: pytest.MonkeyPatch + ) -> None: + from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR + + monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False) + monkeypatch.setattr( + "litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints.is_ptu_cost_attribution_enabled", + lambda: False, + ) + monkeypatch.setattr( + "litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints.get_secret", + lambda *_args: False, + ) + self._mock_prisma(monkeypatch) + + response = client.get("/get/ui_settings") + + assert response.status_code == 200 + assert response.json()["values"]["enable_ptu_cost_attribution"] is False + assert response.json()["source"]["enable_ptu_cost_attribution"] == "config" def test_a_persisted_true_cannot_forge_the_derived_value(self, mock_auth, monkeypatch): """A row written before the allowlist existed must not be able to turn the feature on.""" diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 83f30dc52a4..97e7a799053 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -13,6 +13,7 @@ import time from collections.abc import AsyncIterator, Mapping, Sequence from copy import deepcopy from functools import partial +from types import MappingProxyType from typing import Dict, Final, List, Literal from unittest.mock import AsyncMock, MagicMock, patch @@ -90,6 +91,7 @@ from litellm.types.router import ( TaggedPreRoutingStrategy, ) from litellm.types.llms.openai import ResponsesAPIResponse +from litellm.types.management_endpoints.auto_router_endpoints import RequestComplexityRouterConfig from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler @@ -6681,6 +6683,7 @@ class TestTierModelAffinity: litellm_router_instance=_windowed_router(_SMALL, _BIG), complexity_router_config={ "tiers": {"SIMPLE": ["small-model", "big-model"]}, + "enable_context_window_escalation": True, "adaptive": adaptive, "deployment_affinity": True, "session_affinity": False, @@ -13878,8 +13881,12 @@ _CJK_TURNS = [ ] -def _tier_config(**overrides) -> Dict: - return {"tiers": {"SIMPLE": "small-model", "COMPLEX": "big-model"}, **overrides} +def _tier_config(**overrides: object) -> dict[str, object]: + return { + "tiers": {"SIMPLE": "small-model", "COMPLEX": "big-model"}, + "enable_context_window_escalation": True, + **overrides, + } class TestContextWindowEscalation: @@ -13938,7 +13945,7 @@ class TestContextWindowEscalation: router = ComplexityRouter( model_name="test-router", litellm_router_instance=_windowed_router(_SMALL, ("mid-model", "openai/gpt-4o-mini", 200000), _BIG), - complexity_router_config={"tiers": {"SIMPLE": ["small-model", "mid-model"], "COMPLEX": "big-model"}}, + complexity_router_config=_tier_config(tiers={"SIMPLE": ["small-model", "mid-model"], "COMPLEX": "big-model"}), ) result = await router.async_pre_routing_hook(model="test-router", request_kwargs={}, messages=_OVERSIZED_TURNS) @@ -13975,7 +13982,7 @@ class TestContextWindowEscalation: }, ] ), - complexity_router_config={"tiers": {"SIMPLE": "mixed-pool", "COMPLEX": "big-model"}}, + complexity_router_config=_tier_config(tiers={"SIMPLE": "mixed-pool", "COMPLEX": "big-model"}), ) result = await router.async_pre_routing_hook(model="test-router", request_kwargs={}, messages=_OVERSIZED_TURNS) @@ -14026,7 +14033,7 @@ class TestContextWindowEscalation: router = ComplexityRouter( model_name="test-router", litellm_router_instance=_windowed_router(*deployments), - complexity_router_config={"tiers": tiers}, + complexity_router_config=_tier_config(tiers=tiers), ) result = await router.async_pre_routing_hook(model="test-router", request_kwargs={}, messages=_OVERSIZED_TURNS) @@ -14035,19 +14042,37 @@ class TestContextWindowEscalation: assert result.model == expected_model @pytest.mark.asyncio - async def test_the_disabled_gate_dispatches_on_complexity_alone(self): - """The escape hatch: enable_context_window_escalation false restores today's behavior.""" - router = ComplexityRouter( + @pytest.mark.parametrize("enabled", (None, False, True), ids=("omitted", "disabled", "enabled")) + @pytest.mark.parametrize("serialized", (False, True), ids=("config", "http-json")) + async def test_context_window_escalation_requires_opt_in(self, enabled: bool | None, serialized: bool) -> None: + setting: Final = ( + MappingProxyType({"enable_context_window_escalation": enabled}) + if enabled is not None + else MappingProxyType({}) + ) + raw_config: Final = RequestComplexityRouterConfig.model_validate( + MappingProxyType( + {"tiers": MappingProxyType({"SIMPLE": "small-model", "COMPLEX": "big-model"}), **setting} + ) + ) + config: Final = ( + RequestComplexityRouterConfig.model_validate_json(raw_config.model_dump_json()) + if serialized + else raw_config + ) + router: Final = ComplexityRouter( model_name="test-router", litellm_router_instance=_windowed_router(_SMALL, _BIG), - complexity_router_config=_tier_config(enable_context_window_escalation=False), + complexity_router_config=config.model_dump(exclude_unset=not serialized, exclude_none=True), ) - result = await router.async_pre_routing_hook(model="test-router", request_kwargs={}, messages=_OVERSIZED_TURNS) + result: Final = await router.async_pre_routing_hook( + model="test-router", request_kwargs={}, messages=_OVERSIZED_TURNS + ) assert result is not None - assert result.model == "small-model" - assert "context_escalated" not in result.routing_decision + assert result.model == ("big-model" if enabled else "small-model") + assert result.routing_decision.get("context_escalated", False) is (enabled is True) @pytest.mark.asyncio async def test_out_of_band_system_and_tools_count_against_the_window(self): @@ -14156,7 +14181,7 @@ class TestContextWindowEscalation: }, ] ), - complexity_router_config={"adaptive": True, "tiers": {"SIMPLE": ["small-model", "mid-model"]}}, + complexity_router_config=_tier_config(adaptive=True, tiers={"SIMPLE": ["small-model", "mid-model"]}), ) result = await router.async_pre_routing_hook(model="test-router", request_kwargs={}, messages=_OVERSIZED_TURNS) @@ -14186,7 +14211,7 @@ class TestContextWindowEscalation: }, ] ), - complexity_router_config={"tiers": {"SIMPLE": "cop-pool", "COMPLEX": "big-model"}}, + complexity_router_config=_tier_config(tiers={"SIMPLE": "cop-pool", "COMPLEX": "big-model"}), ) real_get_llm_provider = litellm.get_llm_provider copilot_resolutions: List = [] @@ -14219,7 +14244,7 @@ class TestContextWindowEscalation: "model_name": "smart-router", "litellm_params": { "model": "auto_router/complexity_router", - "complexity_router_config": {"tiers": {"SIMPLE": "small-model", "COMPLEX": "big-model"}}, + "complexity_router_config": _tier_config(), }, }, { @@ -15139,7 +15164,12 @@ class TestHealthFallbackDispatch: ) -> None: from litellm.types.router import RouterRateLimitError - router: Final = self._router(config={"tiers": {"SIMPLE": "primary", "MEDIUM": "peer", "COMPLEX": "large"}}) + router: Final = self._router( + config={ + "tiers": {"SIMPLE": "primary", "MEDIUM": "peer", "COMPLEX": "large"}, + "enable_context_window_escalation": True, + } + ) router.add_deployment( Deployment( model_name="large", @@ -15216,7 +15246,13 @@ class TestHealthFallbackDispatch: @pytest.mark.asyncio @pytest.mark.parametrize("default_fits", [True, False]) async def test_modality_default_must_also_fit_context(self, default_fits: bool) -> None: - router: Final = self._router(config={"modality_routing": True, "tiers": {"SIMPLE": "primary"}}) + router: Final = self._router( + config={ + "modality_routing": True, + "tiers": {"SIMPLE": "primary"}, + "enable_context_window_escalation": True, + } + ) for deployment in router.model_list: deployment["model_info"]["supports_vision"] = deployment["model_name"] == "fallback" deployment["model_info"]["max_input_tokens"] = 10000 if default_fits else 10 diff --git a/tests/test_litellm/router_strategy/test_lowest_tpm_rpm.py b/tests/test_litellm/router_strategy/test_lowest_tpm_rpm.py new file mode 100644 index 00000000000..7b13b196d5b --- /dev/null +++ b/tests/test_litellm/router_strategy/test_lowest_tpm_rpm.py @@ -0,0 +1,54 @@ +from datetime import datetime, timedelta +from typing import Final + +from litellm import Router +from litellm.types.router import DeploymentTypedDict, LiteLLMParamsTypedDict + +MODEL_GROUP: Final = "lowest-tpm-router" +HIGH_USAGE_DEPLOYMENT_ID: Final = "highest-usage" +LOW_USAGE_DEPLOYMENT_ID: Final = "lowest-usage" + + +def _deployment(deployment_id: str) -> DeploymentTypedDict: + params: LiteLLMParamsTypedDict = { + "model": "gpt-4o", + "api_key": "key", + "mock_response": f"from {deployment_id}", + } + return { + "model_name": MODEL_GROUP, + "litellm_params": params, + "model_info": {"id": deployment_id}, + } + + +def test_usage_based_routing_v1_selects_the_lowest_recorded_tpm() -> None: + router: Final = Router( + model_list=[ + _deployment(HIGH_USAGE_DEPLOYMENT_ID), + _deployment(LOW_USAGE_DEPLOYMENT_ID), + ], + routing_strategy="usage-based-routing", + num_retries=0, + ) + usage_by_deployment: Final = { + HIGH_USAGE_DEPLOYMENT_ID: 100, + LOW_USAGE_DEPLOYMENT_ID: 1, + } + now: Final = datetime.now() + cache_keys: Final = tuple( + f"{MODEL_GROUP}:tpm:{(now + timedelta(minutes=offset)).strftime('%H-%M')}" + for offset in range(60) + ) + + for cache_key in cache_keys: + router.cache.set_cache( + key=cache_key, value=usage_by_deployment, ttl=float("inf") + ) + + deployment: Final = router.get_available_deployment( + model=MODEL_GROUP, + messages=[{"role": "user", "content": "test"}], + ) + + assert deployment["model_info"]["id"] == LOW_USAGE_DEPLOYMENT_ID diff --git a/tests/test_litellm/secret_managers/hashicorp_vault_parity.json b/tests/test_litellm/secret_managers/hashicorp_vault_parity.json new file mode 100644 index 00000000000..f1faefd48e8 --- /dev/null +++ b/tests/test_litellm/secret_managers/hashicorp_vault_parity.json @@ -0,0 +1,85 @@ +[ + { + "name": "defaults", + "env": { + "HCP_VAULT_TOKEN": "token" + }, + "secret_name": "OPENAI_API_KEY", + "expected_secret_url": "http://127.0.0.1:8200/v1/secret/data/OPENAI_API_KEY", + "expected_login_url": null, + "expected_login_namespace": null, + "expected_secret_namespace": null + }, + { + "name": "global_namespace", + "env": { + "HCP_VAULT_ADDR": "http://vault.test:8200", + "HCP_VAULT_TOKEN": "token", + "HCP_VAULT_NAMESPACE": "admin" + }, + "secret_name": "OPENAI_API_KEY", + "expected_secret_url": "http://vault.test:8200/v1/admin/secret/data/OPENAI_API_KEY", + "expected_login_url": null, + "expected_login_namespace": "admin", + "expected_secret_namespace": "admin" + }, + { + "name": "namespace_overrides", + "env": { + "HCP_VAULT_ADDR": "http://vault.test:8200", + "HCP_VAULT_TOKEN": "token", + "HCP_VAULT_NAMESPACE": "admin", + "HCP_VAULT_LOGIN_NAMESPACE": "root", + "HCP_VAULT_SECRET_NAMESPACE": "teams/team-a" + }, + "secret_name": "OPENAI_API_KEY", + "expected_secret_url": "http://vault.test:8200/v1/teams/team-a/secret/data/OPENAI_API_KEY", + "expected_login_url": null, + "expected_login_namespace": "root", + "expected_secret_namespace": "teams/team-a" + }, + { + "name": "custom_mount_and_prefix", + "env": { + "HCP_VAULT_ADDR": "http://vault.test:8200", + "HCP_VAULT_TOKEN": "token", + "HCP_VAULT_MOUNT_NAME": " /kv-prod/ ", + "HCP_VAULT_PATH_PREFIX": " /virtual-keys/ " + }, + "secret_name": "DB_PASSWORD", + "expected_secret_url": "http://vault.test:8200/v1/kv-prod/data/virtual-keys/DB_PASSWORD", + "expected_login_url": null, + "expected_login_namespace": null, + "expected_secret_namespace": null + }, + { + "name": "approle_custom_mount", + "env": { + "HCP_VAULT_ADDR": "http://vault.test:8200", + "HCP_VAULT_APPROLE_ROLE_ID": "role-id", + "HCP_VAULT_APPROLE_SECRET_ID": "secret-id", + "HCP_VAULT_APPROLE_MOUNT_PATH": "custom-approle", + "HCP_VAULT_NAMESPACE": "admin" + }, + "secret_name": "OPENAI_API_KEY", + "expected_secret_url": "http://vault.test:8200/v1/admin/secret/data/OPENAI_API_KEY", + "expected_login_url": "http://vault.test:8200/v1/auth/custom-approle/login", + "expected_login_namespace": "admin", + "expected_secret_namespace": "admin" + }, + { + "name": "tls_cert", + "env": { + "HCP_VAULT_ADDR": "http://vault.test:8200", + "HCP_VAULT_CLIENT_CERT": "/tmp/client.crt", + "HCP_VAULT_CLIENT_KEY": "/tmp/client.key", + "HCP_VAULT_CERT_ROLE": "vault-role", + "HCP_VAULT_NAMESPACE": "admin" + }, + "secret_name": "OPENAI_API_KEY", + "expected_secret_url": "http://vault.test:8200/v1/admin/secret/data/OPENAI_API_KEY", + "expected_login_url": "http://vault.test:8200/v1/auth/cert/login", + "expected_login_namespace": "admin", + "expected_secret_namespace": "admin" + } +] diff --git a/tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py b/tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py index 1676540e4ec..fc18cb8b8f7 100644 --- a/tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py +++ b/tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py @@ -1,4 +1,5 @@ import datetime +import json from collections.abc import Mapping from pathlib import Path from typing import Final @@ -18,6 +19,23 @@ LOGIN_RESPONSE: Final = {"auth": {"client_token": "hvs.login-token", "lease_dura SECRET_RESPONSE: Final = {"data": {"data": {"key": "sk-from-vault", "password": "pw-from-vault"}}} NAMESPACE_ENV_VARS: Final = ("HCP_VAULT_NAMESPACE", "HCP_VAULT_LOGIN_NAMESPACE", "HCP_VAULT_SECRET_NAMESPACE") +PARITY_ENV_VARS: Final = ( + "HCP_VAULT_ADDR", + "HCP_VAULT_TOKEN", + "HCP_VAULT_NAMESPACE", + "HCP_VAULT_LOGIN_NAMESPACE", + "HCP_VAULT_SECRET_NAMESPACE", + "HCP_VAULT_MOUNT_NAME", + "HCP_VAULT_PATH_PREFIX", + "HCP_VAULT_APPROLE_ROLE_ID", + "HCP_VAULT_APPROLE_SECRET_ID", + "HCP_VAULT_APPROLE_MOUNT_PATH", + "HCP_VAULT_CLIENT_CERT", + "HCP_VAULT_CLIENT_KEY", + "HCP_VAULT_CERT_ROLE", + "HCP_VAULT_REFRESH_INTERVAL", + "SECRET_MANAGER_REFRESH_INTERVAL", +) def _build_manager(monkeypatch: pytest.MonkeyPatch, env: Mapping[str, str]) -> HashicorpSecretManager: @@ -236,3 +254,35 @@ def test_tls_login_uses_login_namespace(monkeypatch: pytest.MonkeyPatch, tmp_pat assert manager._auth_via_tls_cert() == "hvs.login-token" assert login_route.calls.last.request.headers["X-Vault-Namespace"] == "root" + + +with Path(__file__).with_name("hashicorp_vault_parity.json").open() as parity_file: + PARITY_CASES: Final = json.load(parity_file) + + +@pytest.mark.parametrize("case", PARITY_CASES, ids=lambda case: case["name"]) +def test_configuration_matches_native_parity_fixture( + monkeypatch: pytest.MonkeyPatch, case: Mapping[str, object] +) -> None: + monkeypatch.setattr(litellm.proxy.proxy_server, "premium_user", True) + for name in PARITY_ENV_VARS: + monkeypatch.delenv(name, raising=False) + for name, value in case["env"].items(): + monkeypatch.setenv(name, value) + + manager: Final = HashicorpSecretManager() + env: Final = case["env"] + expected_login_url: Final = case["expected_login_url"] + if env.get("HCP_VAULT_APPROLE_ROLE_ID") and env.get("HCP_VAULT_APPROLE_SECRET_ID"): + login_url: str | None = ( + f"{manager.vault_addr}/v1/auth/{manager.approle_mount_path}/login" + ) + elif env.get("HCP_VAULT_CLIENT_CERT") and env.get("HCP_VAULT_CLIENT_KEY"): + login_url = f"{manager.vault_addr}/v1/auth/cert/login" + else: + login_url = None + + assert manager.get_url(case["secret_name"]) == case["expected_secret_url"] + assert manager.vault_login_namespace == case["expected_login_namespace"] + assert manager.vault_secret_namespace == case["expected_secret_namespace"] + assert login_url == expected_login_url diff --git a/tests/test_litellm_rust/support/s3_stub.py b/tests/test_litellm_rust/support/s3_stub.py new file mode 100644 index 00000000000..5a683fb78f3 --- /dev/null +++ b/tests/test_litellm_rust/support/s3_stub.py @@ -0,0 +1,112 @@ +"""In-process path-style S3 stub for native cache parity tests.""" + +import threading +from dataclasses import dataclass, field +from email.utils import parsedate_to_datetime +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import Final +from urllib.parse import unquote, urlsplit + +_STORED_HEADERS: Final = ( + "cache-control", + "content-type", + "content-language", + "content-disposition", + "expires", +) + + +@dataclass +class S3Object: + body: bytes + headers: dict[str, str] = field(default_factory=dict) + + +class S3Stub: + """Minimal path-style S3 endpoint serving PUT and GET object operations.""" + + def __init__(self) -> None: + self._objects: dict[str, S3Object] = {} + stub: Final = self + + class Handler(BaseHTTPRequestHandler): + def _key(self) -> str: + parts: Final = urlsplit(self.path).path.lstrip("/").split("/", 1) + return unquote(parts[1]) if len(parts) == 2 else "" + + def _read_body(self) -> bytes: + transfer: Final = self.headers.get("transfer-encoding", "") + if "chunked" not in transfer: + return self.rfile.read(int(self.headers.get("content-length", 0))) + chunks: Final = bytearray() + while True: + size = int(self.rfile.readline().split(b";")[0].strip(), 16) + if size == 0: + while self.rfile.readline().strip(): + pass + return bytes(chunks) + chunks.extend(self.rfile.read(size)) + self.rfile.readline() + + def do_PUT(self) -> None: + body: Final = self._read_body() + headers: Final = {name: self.headers[name] for name in _STORED_HEADERS if name in self.headers} + stub._objects = {**stub._objects, self._key(): S3Object(body=body, headers=headers)} + self.send_response(200) + self.send_header("ETag", '"stub"') + self.send_header("Content-Length", "0") + self.end_headers() + + def do_HEAD(self) -> None: + self._object(send_body=False) + + def do_GET(self) -> None: + self._object(send_body=True) + + def _object(self, send_body: bool) -> None: + entry: Final = stub._objects.get(self._key()) + if entry is None: + self.send_response(404) + self.send_header("Content-Type", "application/xml") + body: Final = b'NoSuchKey' + self.send_header("Content-Length", str(len(body))) + self.end_headers() + if send_body: + self.wfile.write(body) + return + self.send_response(200) + for name, value in entry.headers.items(): + self.send_header(name, value) + self.send_header("ETag", '"stub"') + self.send_header("Content-Length", str(len(entry.body))) + self.end_headers() + if send_body: + self.wfile.write(entry.body) + + def log_message(self, format: str, *args: object) -> None: + pass + + self._server: Final = ThreadingHTTPServer(("127.0.0.1", 0), Handler) + self._worker: Final = threading.Thread(target=self._server.serve_forever, daemon=True) + self._worker.start() + + @property + def url(self) -> str: + host, port = self._server.server_address[:2] + return f"http://{host}:{port}" + + @property + def objects(self) -> dict[str, S3Object]: + return self._objects + + def put_object(self, key: str, body: bytes, headers: dict[str, str] | None = None) -> None: + self._objects = {**self._objects, key: S3Object(body=body, headers=headers or {})} + + def expires(self, key: str) -> object: + header: Final = self._objects[key].headers.get("expires") + return parsedate_to_datetime(header) if header else None + + def close(self) -> None: + self._server.shutdown() + self._server.server_close() + self._worker.join(timeout=5) diff --git a/tests/test_litellm_rust/test_cache.py b/tests/test_litellm_rust/test_cache.py index 1bb2849f335..c4962b533c4 100644 --- a/tests/test_litellm_rust/test_cache.py +++ b/tests/test_litellm_rust/test_cache.py @@ -10,13 +10,18 @@ import threading import time import uuid import weakref -from collections.abc import Generator +from collections.abc import Callable, Generator +from contextlib import ExitStack +from datetime import datetime from pathlib import Path from types import SimpleNamespace from typing import Final, Protocol, cast +from unittest.mock import Mock from urllib.parse import urlparse from uuid import uuid4 +import boto3 +import botocore.config import diskcache import fakeredis import pytest @@ -26,14 +31,22 @@ from azure.storage.blob import ContainerClient import litellm from litellm.caching.azure_blob_cache import AzureBlobCache from litellm.caching.caching import Cache, disable_cache, enable_cache, update_cache -from litellm.caching.gcs_cache import GCSCache from litellm.caching.disk_cache import DiskCache +from litellm.caching.gcs_cache import GCSCache from litellm.caching.in_memory_cache import InMemoryCache from litellm.caching.redis_cluster_cache import RedisClusterCache +from litellm.caching.redis_semantic_cache import RedisSemanticCache +from litellm.caching.s3_cache import S3Cache from litellm.rust_bridge import _native from litellm.types.caching import LiteLLMCacheType +from litellm.types.llms.custom_llm import CustomLLMItem +from litellm.types.utils import EmbeddingResponse from tests.test_litellm_rust.support.fake_gcs import FakeGcs from tests.test_litellm_rust.support.isolation import rebound +from tests.test_litellm_rust.support.s3_stub import S3Stub + +_CacheTestHandle: Final = _native._CacheTestHandle # pyright: ignore[reportPrivateUsage] # test-only handle has no public module name +_CacheTestResolver: Final = _native._CacheTestResolver # pyright: ignore[reportPrivateUsage] # test-only resolver has no public module name pytestmark: Final = pytest.mark.requires_rust_extension @@ -47,7 +60,7 @@ def request(key: str = "key") -> dict[str, object]: return {"key": {"preset": key}} -def semantic_request( +def qdrant_request( key: str, messages: list[dict[str, object]], **kwargs: object, @@ -56,7 +69,7 @@ def semantic_request( def embedding_vector(text: str) -> list[float]: - raw = hashlib.sha256(text.encode()).digest()[:8] + raw: Final = hashlib.sha256(text.encode()).digest()[:8] values: Final = [byte / 127.5 - 1 for byte in raw] norm: Final = math.sqrt(sum(value * value for value in values)) return [value / norm for value in values] @@ -177,14 +190,14 @@ def test_existing_constructor_and_global_are_unchanged() -> None: assert type(facade.cache) is InMemoryCache assert "_native_cache_handle" not in vars(facade) with rebound(litellm, "cache", facade): - resolver: Final = _native._CacheTestResolver(litellm) + resolver: Final = _CacheTestResolver(litellm) assert resolver.resolve().kind == "python_callback" resolver.resolve().store(None, {"answer": 7}, callback_kwargs={"cache_key": "key"}) assert cast(CacheLookup, facade).get_cache(cache_key="key") == {"answer": 7} def test_existing_global_lifecycle_remains_the_resolver_source_of_truth() -> None: - resolver: Final = _native._CacheTestResolver(litellm) + resolver: Final = _CacheTestResolver(litellm) enable_cache(type=LiteLLMCacheType.LOCAL, ttl=30) enabled: Final = litellm.cache @@ -207,13 +220,13 @@ def test_existing_global_lifecycle_remains_the_resolver_source_of_truth() -> Non async def test_native_bindings_survive_replacement_and_capture_writes_before_dispatch() -> None: - namespace: Final = SimpleNamespace(cache=_native._CacheTestHandle.memory()) - resolver: Final = _native._CacheTestResolver(namespace) + namespace: Final = SimpleNamespace(cache=_CacheTestHandle.memory()) + resolver: Final = _CacheTestResolver(namespace) selected: Final = resolver.resolve() assert selected.kind == "native" selected.store(request(), {"answer": 1}) assert await selected.async_lookup(request()) == {"answer": 1} - with rebound(namespace, "cache", _native._CacheTestHandle.memory()): + with rebound(namespace, "cache", _CacheTestHandle.memory()): replacement: Final = resolver.resolve() await selected.async_store(request(), {"answer": 2}) assert replacement.lookup(request()) is None @@ -246,7 +259,7 @@ async def test_python_callback_preserves_identity_caller_task_context_and_errors raise failure namespace: Final = SimpleNamespace(cache=CustomCache()) - binding: Final = _native._CacheTestResolver(namespace).resolve() + binding: Final = _CacheTestResolver(namespace).resolve() assert binding.kind == "python_callback" assert await binding.async_lookup(None, callback_kwargs={"marker": sentinel}) is sentinel assert context.get() == "callback" @@ -267,7 +280,7 @@ async def test_callback_cancellation_stays_in_the_callers_task() -> None: finally: finished.set() - binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=CustomCache())).resolve() + binding: Final = _CacheTestResolver(SimpleNamespace(cache=CustomCache())).resolve() async def lookup() -> object: return await binding.async_lookup(None, callback_kwargs={}) @@ -282,9 +295,9 @@ async def test_callback_cancellation_stays_in_the_callers_task() -> None: def test_registered_facade_uses_native_and_instance_overrides_fall_back() -> None: facade: Final = Cache(type=LiteLLMCacheType.LOCAL) - handle: Final = _native._CacheTestHandle.memory() + handle: Final = _CacheTestHandle.memory() handle._bind_facade(facade) - resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)) + resolver: Final = _CacheTestResolver(SimpleNamespace(cache=facade)) native: Final = resolver.resolve() assert native.kind == "native" native.store(request(), {"source": "native"}) @@ -315,12 +328,12 @@ def test_facade_subclasses_backend_replacement_and_configuration_changes_are_not class CustomCache(Cache): pass - handle: Final = _native._CacheTestHandle.memory() + handle: Final = _CacheTestHandle.memory() with pytest.raises(TypeError): handle._bind_facade(CustomCache(type=LiteLLMCacheType.LOCAL)) facade: Final = Cache(type=LiteLLMCacheType.LOCAL) handle._bind_facade(facade) - resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)) + resolver: Final = _CacheTestResolver(SimpleNamespace(cache=facade)) with rebound(facade, "cache", InMemoryCache()): assert resolver.resolve().kind == "python_callback" with rebound(facade, "ttl", 12): @@ -345,7 +358,7 @@ def test_resolver_and_callback_cycles_can_be_collected() -> None: def cyclic_reference() -> weakref.ReferenceType[CustomCache]: callback: Final = CustomCache() namespace: Final = SimpleNamespace(cache=callback) - binding: Final = _native._CacheTestResolver(namespace).resolve() + binding: Final = _CacheTestResolver(namespace).resolve() setattr(callback, "binding", binding) return weakref.ref(callback) @@ -356,8 +369,8 @@ def test_resolver_and_callback_cycles_can_be_collected() -> None: async def test_redis_reads_python_sync_and_async_entries_and_writes_without_hidden_prefix(redis_url: str) -> None: client: Final = redis.Redis.from_url(redis_url) - namespace: Final = SimpleNamespace(cache=_native._CacheTestHandle.redis(redis_url, namespace="team")) - binding: Final = _native._CacheTestResolver(namespace).resolve() + namespace: Final = SimpleNamespace(cache=_CacheTestHandle.redis(redis_url, namespace="team")) + binding: Final = _CacheTestResolver(namespace).resolve() response: Final = {"choices": [{"text": "cached"}], "usage": {"total_tokens": 3}, "flag": True, "empty": None} envelope: Final = {"timestamp": time.time(), "response": json.dumps(response)} client.set("team:sync", str(envelope)) @@ -379,33 +392,33 @@ async def test_redis_reads_python_sync_and_async_entries_and_writes_without_hidd def test_invalid_duration_and_request_shape_fail_before_storage() -> None: - binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=_native._CacheTestHandle.memory())).resolve() + binding: Final = _CacheTestResolver(SimpleNamespace(cache=_CacheTestHandle.memory())).resolve() for seconds in (-1.0, float("nan"), float("inf")): with pytest.raises(ValueError, match="cache durations must be finite and nonnegative"): binding.store({**request(), "ttl_seconds": seconds}, {"answer": 1}) assert binding.lookup(request()) is None with pytest.raises(ValueError, match="cache durations must be finite and nonnegative"): - _native._CacheTestHandle.memory(ttl_seconds=-1) + _CacheTestHandle.memory(ttl_seconds=-1) async def test_memory_size_policy_is_applied_by_the_native_host() -> None: - handle: Final = _native._CacheTestHandle.memory(capacity=2, max_entry_bytes=128) - binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + handle: Final = _CacheTestHandle.memory(capacity=2, max_entry_bytes=128) + binding: Final = _CacheTestResolver(SimpleNamespace(cache=handle)).resolve() small: Final = {"answer": "ok"} binding.store(request("small"), small) assert await binding.async_lookup(request("small")) == small await binding.async_store(request("large"), {"answer": "x" * 256}) assert binding.lookup(request("large")) is None assert binding.lookup(request("small")) == small - disabled: Final = _native._CacheTestResolver( - SimpleNamespace(cache=_native._CacheTestHandle.memory(capacity=0)) + disabled: Final = _CacheTestResolver( + SimpleNamespace(cache=_CacheTestHandle.memory(capacity=0)) ).resolve() await disabled.async_store(request(), small) assert await disabled.async_lookup(request()) is None async def test_native_batch_lookup_and_store_report_partial_hits() -> None: - binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=_native._CacheTestHandle.memory())).resolve() + binding: Final = _CacheTestResolver(SimpleNamespace(cache=_CacheTestHandle.memory())).resolve() requests: Final = [request("hit"), request("miss"), request("disabled")] requests[2]["controls"] = { "supported_call_type": True, @@ -443,7 +456,7 @@ async def test_python_batch_callbacks_use_the_builtin_cache_api() -> None: ) -> object: return result, kwargs - binding: Final = _native._CacheTestResolver( + binding: Final = _CacheTestResolver( SimpleNamespace(cache=CustomCache(type=LiteLLMCacheType.LOCAL)) ).resolve() assert binding.kind == "python_callback" @@ -473,7 +486,7 @@ async def test_unmodified_builtin_cache_callbacks_can_ping_and_flush() -> None: cache: Final = Cache(type=LiteLLMCacheType.LOCAL) cache.cache.set_cache("key", "value") - binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=cache)).resolve() + binding: Final = _CacheTestResolver(SimpleNamespace(cache=cache)).resolve() assert binding.kind == "python_callback" setattr(cache.cache, "ping", ping) @@ -485,7 +498,7 @@ async def test_unmodified_builtin_cache_callbacks_can_ping_and_flush() -> None: def test_facade_registration_rejects_mismatched_capacity() -> None: facade: Final = Cache(type=LiteLLMCacheType.LOCAL) with pytest.raises(TypeError, match="capacities must match"): - _native._CacheTestHandle.memory(capacity=7)._bind_facade(facade) + _CacheTestHandle.memory(capacity=7)._bind_facade(facade) def test_azure_blob_facade_serves_natively_and_python_reads_the_same_blobs(azure_blob_facade: Cache) -> None: @@ -495,15 +508,20 @@ def test_azure_blob_facade_serves_natively_and_python_reads_the_same_blobs(azure assert handle.backend == "azure-blob" account_url: Final = backend.container_client.url.removesuffix(f"/{backend.container_client.container_name}") with pytest.raises(TypeError, match="containers must match"): - _native._CacheTestHandle.azure_blob(account_url, f"{backend.container_client.container_name}-other")._bind_facade( - azure_blob_facade - ) + _native._CacheTestHandle.azure_blob( + account_url, f"{backend.container_client.container_name}-other" + )._bind_facade(azure_blob_facade) handle._bind_facade(azure_blob_facade) resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=azure_blob_facade)) native: Final = resolver.resolve() assert native.kind == "native" - response: Final = {"choices": [{"text": "caf\u00e9 \u2603"}], "usage": {"total_tokens": 3}, "flag": True, "empty": None} + response: Final = { + "choices": [{"text": "caf\u00e9 \u2603"}], + "usage": {"total_tokens": 3}, + "flag": True, + "empty": None, + } native.store({**request("sync"), "ttl_seconds": 0.001}, response) native.store(request("sync"), {"choices": [{"text": "second"}]}) time.sleep(0.01) @@ -558,7 +576,9 @@ async def test_azure_blob_native_async_writes_overwrite_batch_and_flush_like_pyt await binding.async_store({**request("async"), "ttl_seconds": 0.001}, {"value": 2}) time.sleep(0.01) assert await binding.async_lookup(request("async")) == {"value": 2} - assert await backend.async_get_cache("async") == json.loads(backend.container_client.download_blob("async").readall()) + assert await backend.async_get_cache("async") == json.loads( + backend.container_client.download_blob("async").readall() + ) assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="async") == {"value": 2} await binding.async_store_batch([request("first"), request("second")], [{"value": 3}, {"value": 4}]) @@ -581,19 +601,19 @@ async def test_redis_facade_buffers_native_async_writes(redis_url: str) -> None: redis_flush_size=2, ) with pytest.raises(TypeError, match="default TTLs must match"): - _native._CacheTestHandle.redis(redis_url, ttl_seconds=61)._bind_facade(facade) + _CacheTestHandle.redis(redis_url, ttl_seconds=61)._bind_facade(facade) with pytest.raises(TypeError, match="namespaces must match"): - _native._CacheTestHandle.redis(redis_url, namespace="other")._bind_facade(facade) - _native._CacheTestHandle.redis(redis_url, ttl_seconds=60)._bind_facade(facade) - binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + _CacheTestHandle.redis(redis_url, namespace="other")._bind_facade(facade) + _CacheTestHandle.redis(redis_url, ttl_seconds=60)._bind_facade(facade) + binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() client: Final = redis.Redis.from_url(redis_url) with rebound(facade.cache, "redis_kwargs", {**facade.cache.redis_kwargs, "ssl": True}): - assert _native._CacheTestResolver(SimpleNamespace(cache=facade)).resolve().kind == "python_callback" + assert _CacheTestResolver(SimpleNamespace(cache=facade)).resolve().kind == "python_callback" pool: Final = facade.cache.redis_client.connection_pool with rebound(pool, "connection_kwargs", {**pool.connection_kwargs, "db": 1}): - assert _native._CacheTestResolver(SimpleNamespace(cache=facade)).resolve().kind == "python_callback" + assert _CacheTestResolver(SimpleNamespace(cache=facade)).resolve().kind == "python_callback" await binding.async_store(request("first"), {"value": 1}) assert client.get("first") is None @@ -603,6 +623,8 @@ async def test_redis_facade_buffers_native_async_writes(redis_url: str) -> None: assert client.get("second") is not None await facade.cache.disconnect() client.close() + + async def test_disk_reads_python_entries_and_python_reads_native_entries(tmp_path: Path) -> None: disk_cache: Final = DiskCache(disk_cache_dir=str(tmp_path)) response: Final = {"choices": [{"text": "cached"}], "usage": {"total_tokens": 3}} @@ -711,6 +733,183 @@ async def test_disk_native_batch_lookup_and_store_report_partial_hits(tmp_path: } +@pytest.fixture +def s3_stub() -> Generator[S3Stub]: + stub: Final = S3Stub() + try: + yield stub + finally: + stub.close() + + +def python_s3(url: str) -> S3Cache: + return S3Cache( + s3_bucket_name="cache-bucket", + s3_region_name="us-east-1", + s3_endpoint_url=url, + s3_aws_access_key_id="key", + s3_aws_secret_access_key="secret", + s3_path="team", + ) + + +async def test_s3_reads_python_entries_and_writes_with_python_metadata(s3_stub: S3Stub) -> None: + python_cache: Final = python_s3(s3_stub.url) + response: Final = {"choices": [{"text": "cached"}], "usage": {"total_tokens": 3}} + python_cache.set_cache("sync:key", {"timestamp": time.time(), "response": response}, ttl=90) + python_cache.set_cache("plain", {"timestamp": time.time(), "response": response}) + s3_stub.put_object("team/malformed", b"not a cache entry") + s3_stub.put_object( + "team/expired", + json.dumps({"timestamp": time.time(), "response": response}).encode(), + {"expires": "Thu, 01 Jan 1970 00:00:00 GMT"}, + ) + binding: Final = _native._CacheTestResolver( + SimpleNamespace( + cache=_native._CacheTestHandle.s3( + "cache-bucket", + region="us-east-1", + endpoint_url=s3_stub.url, + key_prefix="team/", + access_key_id="key", + secret_access_key="secret", + ) + ) + ).resolve() + + assert binding.lookup(request("sync:key")) == response + assert await binding.async_lookup(request("plain")) == response + assert binding.lookup(request("malformed")) is None + assert binding.lookup(request("expired")) is None + assert binding.lookup(request("absent")) is None + + binding.store({**request("native:key"), "ttl_seconds": 90.0}, response) + await binding.async_store(request("no_ttl"), response) + stored: Final = s3_stub.objects["team/native/key"] + assert stored.headers["content-type"] == "application/json" + assert stored.headers["content-language"] == "en" + assert stored.headers["content-disposition"] == 'inline; filename="team/native/key.json"' + assert stored.headers["cache-control"] == "immutable, max-age=90, s-maxage=90" + expires: Final = cast(datetime, s3_stub.expires("team/native/key")) + remaining: Final = (expires - datetime.now(expires.tzinfo)).total_seconds() + assert 60 < remaining <= 91 + no_ttl: Final = s3_stub.objects["team/no_ttl"] + assert no_ttl.headers["cache-control"] == "immutable, max-age=31536000, s-maxage=31536000" + assert "expires" not in no_ttl.headers + assert python_cache.get_cache("native:key")["response"] == response + + partial: Final = await binding.async_lookup_batch([request("native:key"), request("absent"), request("malformed")]) + assert partial == {"values": [response, None, None], "missing_indices": [1, 2]} + + +def test_s3_facade_binds_only_exact_configuration_and_falls_back_on_mutation(s3_stub: S3Stub) -> None: + facade: Final = Cache( + type=LiteLLMCacheType.S3, + s3_bucket_name="cache-bucket", + s3_region_name="us-east-1", + s3_endpoint_url=s3_stub.url, + s3_aws_access_key_id="key", + s3_aws_secret_access_key="secret", + s3_path="team", + ) + handle: Final = _native._CacheTestHandle.s3( + "cache-bucket", + region="us-east-1", + endpoint_url=s3_stub.url, + key_prefix="team/", + access_key_id="key", + secret_access_key="secret", + ) + with pytest.raises(TypeError, match="buckets must match"): + _native._CacheTestHandle.s3("other", region="us-east-1", endpoint_url=s3_stub.url)._bind_facade(facade) + with pytest.raises(TypeError, match="key prefixes must match"): + _native._CacheTestHandle.s3( + "cache-bucket", region="us-east-1", endpoint_url=s3_stub.url, key_prefix="other/" + )._bind_facade(facade) + handle._bind_facade(facade) + resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)) + binding: Final = resolver.resolve() + assert binding.kind == "native" + + handler: Final = Mock() + facade.cache.s3_client.meta.events.register("before-call.s3.*", handler) + binding.store(request("native"), {"answer": 1}) + assert binding.lookup(request("native")) == {"answer": 1} + assert handler.call_count == 0 + assert "team/native" in s3_stub.objects + + with rebound(facade.cache, "bucket_name", "other"): + assert resolver.resolve().kind == "python_callback" + other_client: Final = boto3.client( + "s3", + region_name="us-east-1", + endpoint_url=s3_stub.url, + aws_access_key_id="key", + aws_secret_access_key="secret", + ) + with rebound(facade.cache, "s3_client", other_client): + assert resolver.resolve().kind == "python_callback" + + class CustomS3Cache(S3Cache): + pass + + subclassed: Final = Cache( + type=LiteLLMCacheType.S3, + s3_bucket_name="cache-bucket", + s3_region_name="us-east-1", + s3_endpoint_url=s3_stub.url, + s3_aws_access_key_id="key", + s3_aws_secret_access_key="secret", + s3_path="team", + ) + subclassed.cache = CustomS3Cache( + s3_bucket_name="cache-bucket", + s3_region_name="us-east-1", + s3_endpoint_url=s3_stub.url, + s3_aws_access_key_id="key", + s3_aws_secret_access_key="secret", + s3_path="team", + ) + with pytest.raises(TypeError): + handle._bind_facade(subclassed) + assert _native._CacheTestResolver(SimpleNamespace(cache=subclassed)).resolve().kind == "python_callback" + + +def test_s3_facade_rejects_configurations_that_require_python(s3_stub: S3Stub) -> None: + handle: Final = _native._CacheTestHandle.s3( + "cache-bucket", + region="us-east-1", + endpoint_url=s3_stub.url, + key_prefix="team/", + access_key_id="key", + secret_access_key="secret", + ) + unverified: Final = Cache( + type=LiteLLMCacheType.S3, + s3_bucket_name="cache-bucket", + s3_region_name="us-east-1", + s3_endpoint_url="https://s3.example.test", + s3_aws_access_key_id="key", + s3_aws_secret_access_key="secret", + s3_path="team", + s3_verify=False, + ) + with pytest.raises(TypeError, match="requires Python"): + handle._bind_facade(unverified) + proxied: Final = Cache( + type=LiteLLMCacheType.S3, + s3_bucket_name="cache-bucket", + s3_region_name="us-east-1", + s3_endpoint_url=s3_stub.url, + s3_aws_access_key_id="key", + s3_aws_secret_access_key="secret", + s3_path="team", + s3_config=botocore.config.Config(proxies={"https": "http://proxy.test"}), + ) + with pytest.raises(TypeError, match="requires Python"): + handle._bind_facade(proxied) + + async def test_gcs_reads_python_entries_and_writes_python_compatible_objects( fake_gcs: FakeGcs, monkeypatch: pytest.MonkeyPatch ) -> None: @@ -846,6 +1045,7 @@ async def test_gcs_facade_binds_only_exact_matching_configuration( assert resolver.resolve().kind == "python_callback" with rebound(facade.cache, "path_service_account", "sa.json"): assert resolver.resolve().kind == "python_callback" + def no_get_cache(*args: object, **kwargs: object) -> None: return None @@ -980,7 +1180,9 @@ async def test_redis_cluster_facade_serves_multi_slot_batches_and_scoped_flush_n await binding.async_flush() - remaining: Final = tuple(sorted(key for node in client.get_primaries() for key in client.keys("parity:*", target_nodes=node))) + remaining: Final = tuple( + sorted(key for node in client.get_primaries() for key in client.keys("parity:*", target_nodes=node)) + ) assert remaining == (), remaining assert client.get("unscoped") == b"stays" client.delete("unscoped") @@ -988,10 +1190,576 @@ async def test_redis_cluster_facade_serves_multi_slot_batches_and_scoped_flush_n facade.cache.redis_client.close() -def qdrant_facade( - qdrant_url: str, - collection_name: str, -) -> Cache: +PARAPHRASE_MARKER: Final = " (paraphrase)" +SEMANTIC_EMBEDDING_MODEL: Final = "semantic-test/deterministic" +SEMANTIC_INDEX_PREFIX: Final = "litellm_test_semantic_" +SEMANTIC_CONTEXT: Final = contextvars.ContextVar("semantic_test_context", default="unset") + + +def _normalized(vector: list[float]) -> list[float]: + norm: Final = math.sqrt(sum(component * component for component in vector)) + return [component / norm for component in vector] + + +def _base_embedding(prompt: str) -> list[float]: + digest: Final = hashlib.sha256(prompt.encode("utf-8")).digest() + return _normalized([float(digest[index] + 1) for index in range(8)]) + + +def _semantic_embedding(prompt: str) -> list[float]: + if PARAPHRASE_MARKER not in prompt: + return _base_embedding(prompt) + base: Final = _base_embedding(prompt.replace(PARAPHRASE_MARKER, "").strip()) + pivot: Final = min(range(8), key=lambda index: abs(base[index])) + direction: Final = _normalized( + [ + (1.0 - base[pivot] * base[pivot]) if index == pivot else -base[index] * base[pivot] + for index in range(8) + ] + ) + # Rotating an orthogonal unit direction by 0.329 produces ~0.05 cosine distance + return _normalized([base[index] + 0.329 * direction[index] for index in range(8)]) + + +class DeterministicEmbedding(litellm.CustomLLM): + def __init__(self) -> None: + self.calls: list[dict[str, object]] = [] + self.async_calls: list[dict[str, object]] = [] + self.entered = asyncio.Event() + self.gate: asyncio.Event | None = None + + def _respond( + self, + model: str, + input: object, + model_response: EmbeddingResponse, + ) -> EmbeddingResponse: + texts: Final = cast(list[object], input if isinstance(input, list) else [input]) + self.calls.append({"model": model, "input": texts}) + model_response.model = model + model_response.data = [ + {"object": "embedding", "index": index, "embedding": _semantic_embedding(str(text))} + for index, text in enumerate(texts) + ] + return model_response + + def embedding( + self, + model: str, + input: list[object], + model_response: EmbeddingResponse, + print_verbose: Callable[..., object], + logging_obj: object, + optional_params: dict[str, object], + api_key: object = None, + api_base: object = None, + timeout: object = None, + litellm_params: object = None, + ) -> EmbeddingResponse: + return self._respond(model, input, model_response) + + async def aembedding( + self, + model: str, + input: list[object], + model_response: EmbeddingResponse, + print_verbose: Callable[..., object], + logging_obj: object, + optional_params: dict[str, object], + api_key: object = None, + api_base: object = None, + timeout: object = None, + litellm_params: object = None, + ) -> EmbeddingResponse: + texts: Final = cast(list[object], input if isinstance(input, list) else [input]) + self.async_calls.append( + { + "model": model, + "input": texts, + "task": asyncio.current_task(), + "context": SEMANTIC_CONTEXT.get(), + } + ) + SEMANTIC_CONTEXT.set("written-in-aembedding") + self.entered.set() + if self.gate is not None: + await self.gate.wait() + return self._respond(model, input, model_response) + + +@pytest.fixture +def semantic_embedding() -> Generator[DeterministicEmbedding]: + handler: Final = DeterministicEmbedding() + with ExitStack() as stack: + stack.enter_context( + rebound( + litellm, + "custom_provider_map", + [ + *litellm.custom_provider_map, + cast( + CustomLLMItem, + {"provider": "semantic-test", "custom_handler": handler}, + ), + ], + ) + ) + stack.enter_context( + rebound( + litellm, + "_custom_providers", # pyright: ignore[reportPrivateUsage] # no public provider-registration hook + [*litellm._custom_providers, "semantic-test"], # pyright: ignore[reportPrivateUsage] # no public provider-registration hook + ) + ) + stack.enter_context( + rebound(litellm, "provider_list", [*litellm.provider_list, "semantic-test"]) + ) + yield handler + + +@pytest.fixture +def redis_stack() -> Generator[tuple[str, str]]: + url: Final = os.environ.get("LITELLM_REDIS_STACK_URL") + if url is None: + pytest.skip("LITELLM_REDIS_STACK_URL is not set") + index: Final = f"{SEMANTIC_INDEX_PREFIX}{uuid4().hex}" + yield url, index + client: Final = redis.Redis.from_url(url) + try: + client.execute_command("FT.DROPINDEX", index, "DD") # pyright: ignore[reportUnknownMemberType] # redis-py leaves execute_command partially unknown + except redis.RedisError: + pass + client.close() + + +def semantic_request(key: str, prompt: str, **extra: object) -> dict[str, object]: + return { + "key": {"preset": key}, + "messages": [{"role": "user", "content": prompt}], + **extra, + } + + +def semantic_messages(prompt: str) -> list[dict[str, object]]: + return [{"role": "user", "content": prompt}] + + +def semantic_entry_id(prompt: str, tag: str) -> str: + return hashlib.sha256(f"{prompt}litellm_cache_key{tag}".encode()).hexdigest() + + +def semantic_facade(url: str, index: str, *, similarity_threshold: float = 0.8) -> Cache: + facade: Final = Cache( + type=LiteLLMCacheType.REDIS_SEMANTIC, + redis_url=url, + similarity_threshold=similarity_threshold, + redis_semantic_cache_embedding_model=SEMANTIC_EMBEDDING_MODEL, + redis_semantic_cache_index_name=index, + ) + _CacheTestHandle.redis_semantic(facade.cache)._bind_facade(facade) + return facade + + +def test_redis_semantic_constructor_identity_and_provenance( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + backend: Final = cast(RedisSemanticCache, facade.cache) + assert backend.__class__.__module__ == "litellm.caching.redis_semantic_cache" + assert type(backend) is RedisSemanticCache + assert backend._redis_url == url # pyright: ignore[reportPrivateUsage] # provenance check needs the projected config + assert backend._index_name == index # pyright: ignore[reportPrivateUsage] # provenance check needs the projected config + assert backend.similarity_threshold == 0.8 + assert backend.embedding_model == SEMANTIC_EMBEDDING_MODEL + handle: Final = cast(object, getattr(facade, "_native_cache_handle")) + assert isinstance(handle, _CacheTestHandle) + assert handle.backend == "redis_semantic" + binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + assert binding.kind == "native" + + +def test_redis_semantic_native_and_python_sync_entries_share_one_layout( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + client: Final = redis.Redis.from_url(url) + response: Final = {"choices": [{"text": "paris"}], "usage": {"total_tokens": 2}} + + binding.store(semantic_request("geo", "what is the capital of france"), response) + + native_hash_key: Final = f"{index}:{semantic_entry_id('what is the capital of france', 'geo')}" + stored: Final = client.hgetall(native_hash_key) + assert set(stored) == { + b"entry_id", + b"prompt", + b"response", + b"prompt_vector", + b"inserted_at", + b"updated_at", + b"litellm_cache_key", + }, stored + assert stored[b"entry_id"].decode() == native_hash_key.split(":", 1)[1] + assert stored[b"prompt"] == b"what is the capital of france" + assert stored[b"litellm_cache_key"] == b"geo" + assert len(stored[b"prompt_vector"]) == 32 + decoded: Final = cast(dict[str, object], json.loads(stored[b"response"])) + assert decoded["response"] == response + assert ( + cast(RedisSemanticCache, facade.cache).get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class + "geo", messages=semantic_messages("what is the capital of france") + ) + == decoded + ) + assert semantic_embedding.calls == [ + {"model": "deterministic", "input": ["what is the capital of france"]}, + {"model": "deterministic", "input": ["what is the capital of france"]}, + {"model": "deterministic", "input": ["dimension test"]}, + ] + + cast(RedisSemanticCache, facade.cache).set_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class + "math", + json.dumps({"timestamp": 1700000000.0, "response": {"answer": 42}}), + messages=semantic_messages("what is 6 times 7"), + ) + python_hash_key: Final = f"{index}:{semantic_entry_id('what is 6 times 7', 'math')}" + assert json.loads(cast(bytes, client.hget(python_hash_key, "response"))) == { + "timestamp": 1700000000.0, + "response": {"answer": 42}, + } + assert binding.lookup(semantic_request("math", "what is 6 times 7")) == {"answer": 42} + client.close() + + +async def test_redis_semantic_async_paths_and_store_batch_share_one_layout( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + client: Final = redis.Redis.from_url(url) + + await binding.async_store( + semantic_request("async", "name a primary color"), {"answer": "blue"} + ) + hash_key: Final = f"{index}:{semantic_entry_id('name a primary color', 'async')}" + decoded: Final = cast(dict[str, object], json.loads(cast(bytes, client.hget(hash_key, "response")))) + python_read: Final = await cast(RedisSemanticCache, facade.cache).async_get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class + "async", messages=semantic_messages("name a primary color") + ) + assert python_read == decoded + + await binding.async_store_batch( + [ + semantic_request("batch-one", "first batch prompt"), + semantic_request("batch-two", "second batch prompt"), + ], + [{"answer": 1}, {"answer": 2}], + ) + expected: Final = { + key: json.loads( + cast(bytes, client.hget(f"{index}:{semantic_entry_id(prompt, key)}", "response")) + ) + for key, prompt in ( + ("batch-one", "first batch prompt"), + ("batch-two", "second batch prompt"), + ) + } + for key, prompt in ( + ("batch-one", "first batch prompt"), + ("batch-two", "second batch prompt"), + ): + assert cast(RedisSemanticCache, facade.cache).get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class + key, messages=semantic_messages(prompt) + ) == expected[key], key + + cast(RedisSemanticCache, facade.cache).set_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class + "async-python", + json.dumps({"timestamp": 1700000000.0, "response": {"answer": "python"}}), + messages=semantic_messages("python written prompt"), + ) + assert await binding.async_lookup( + semantic_request("async-python", "python written prompt") + ) == {"answer": "python"} + client.close() + + +async def test_native_semantic_async_embedding_runs_inline_in_the_callers_task( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + assert binding.kind == "native" + caller: Final = asyncio.current_task() + SEMANTIC_CONTEXT.set("caller-sentinel") + response: Final = {"choices": [{"text": "paris"}]} + + await binding.async_store( + semantic_request("inline", "what is the capital of france"), response + ) + assert ( + await binding.async_lookup( + semantic_request("inline", f"what is the capital of france{PARAPHRASE_MARKER}") + ) + == response + ) + assert await binding.async_lookup(semantic_request("inline", "python written prompt")) is None + assert SEMANTIC_CONTEXT.get() == "written-in-aembedding" + assert semantic_embedding.async_calls == [ + { + "model": "deterministic", + "input": ["what is the capital of france"], + "task": caller, + "context": "caller-sentinel", + }, + { + "model": "deterministic", + "input": [f"what is the capital of france{PARAPHRASE_MARKER}"], + "task": caller, + "context": "written-in-aembedding", + }, + { + "model": "deterministic", + "input": ["python written prompt"], + "task": caller, + "context": "written-in-aembedding", + }, + ], semantic_embedding.async_calls + + +async def test_native_semantic_cancellation_during_embedding_skips_the_backend( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + assert binding.kind == "native" + semantic_embedding.gate = asyncio.Event() + + async def lookup() -> object: + return await binding.async_lookup( + semantic_request("cancel", "cancelled prompt") + ) + + task: Final = asyncio.create_task(lookup()) + await semantic_embedding.entered.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + semantic_embedding.gate.set() + + assert len(semantic_embedding.async_calls) == 1 + assert ( + await cast(RedisSemanticCache, facade.cache).async_get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class + "cancel", messages=semantic_messages("cancelled prompt") + ) + is None + ) + + +def test_redis_semantic_similarity_tag_and_threshold_boundaries( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + + binding.store(semantic_request("sim", "tell me a joke"), {"answer": "haha"}) + paraphrase: Final = f"tell me a joke{PARAPHRASE_MARKER}" + assert binding.lookup(semantic_request("sim", paraphrase)) == {"answer": "haha"} + assert binding.lookup(semantic_request("sim", "an unrelated question about spreadsheets")) is None + assert binding.lookup(semantic_request("other-key", "tell me a joke")) is None + + strict: Final = semantic_facade(url, index, similarity_threshold=0.99) + strict_binding: Final = _CacheTestResolver(SimpleNamespace(cache=strict)).resolve() + assert strict_binding.lookup(semantic_request("sim", paraphrase)) is None + assert strict_binding.lookup(semantic_request("sim", "tell me a joke")) == {"answer": "haha"} + + +def test_redis_semantic_ttl_is_written_only_when_requested( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + client: Final = redis.Redis.from_url(url) + + binding.store( + {**semantic_request("ttl", "ttl prompt"), "ttl_seconds": 12.0}, {"answer": 1} + ) + expiring: Final = f"{index}:{semantic_entry_id('ttl prompt', 'ttl')}" + assert 0 < client.ttl(expiring) <= 12 + + binding.store(semantic_request("ttl-none", "untimed prompt"), {"answer": 2}) + persistent: Final = f"{index}:{semantic_entry_id('untimed prompt', 'ttl-none')}" + assert client.ttl(persistent) == -1 + + binding.store( + {**semantic_request("ttl-fraction", "fractional prompt"), "ttl_seconds": 1.5}, + {"answer": 3}, + ) + fractional: Final = f"{index}:{semantic_entry_id('fractional prompt', 'ttl-fraction')}" + assert client.ttl(fractional) == 2 + client.close() + + +def test_redis_semantic_malformed_response_is_a_miss_for_both_readers( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + client: Final = redis.Redis.from_url(url) + + binding.store(semantic_request("bad", "corrupt me"), {"answer": 1}) + hash_key: Final = f"{index}:{semantic_entry_id('corrupt me', 'bad')}" + client.hset(hash_key, "response", b"{not json") + assert binding.lookup(semantic_request("bad", "corrupt me")) is None + assert ( + cast(RedisSemanticCache, facade.cache).get_cache( # pyright: ignore[reportUnknownMemberType] # **kwargs stays unknown on the backend class + "bad", messages=semantic_messages("corrupt me") + ) + is None + ) + client.close() + + +async def test_redis_semantic_unsupported_operations_raise_not_implemented( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + + with pytest.raises(NotImplementedError): + binding.lookup_batch([semantic_request("batch", "prompt one")]) + with pytest.raises(NotImplementedError): + await binding.async_lookup_batch([semantic_request("batch", "prompt one")]) + with pytest.raises(NotImplementedError): + await binding.async_flush() + with pytest.raises(NotImplementedError): + await binding.ping() + + +def test_redis_semantic_requests_without_prompt_are_noops( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + client: Final = redis.Redis.from_url(url) + + binding.store(request("plain"), {"answer": 1}) + assert binding.lookup(request("plain")) is None + assert semantic_embedding.calls == [] + assert client.keys(f"{index}:*") == [] + client.close() + + +def test_redis_semantic_scope_overrides_the_tag_and_isolates_entries( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + binding: Final = _CacheTestResolver(SimpleNamespace(cache=facade)).resolve() + client: Final = redis.Redis.from_url(url) + + scoped: Final = {**semantic_request("scoped", "scoped prompt"), "scope": "team-a"} + binding.store(scoped, {"answer": "kept"}) + hash_key: Final = f"{index}:{semantic_entry_id('scoped prompt', 'team-a')}" + assert client.hget(hash_key, "litellm_cache_key") == b"team-a" + assert binding.lookup(scoped) == {"answer": "kept"} + assert binding.lookup(semantic_request("scoped", "scoped prompt")) is None + assert binding.lookup({**scoped, "scope": "team-b"}) is None + client.close() + + +def test_redis_semantic_configuration_drift_falls_back_to_python( + redis_stack: tuple[str, str], + semantic_embedding: DeterministicEmbedding, + monkeypatch: pytest.MonkeyPatch, +) -> None: + url, index = redis_stack + facade: Final = semantic_facade(url, index) + resolver: Final = _CacheTestResolver(SimpleNamespace(cache=facade)) + assert resolver.resolve().kind == "native" + + with rebound(facade.cache, "similarity_threshold", 0.5): + assert resolver.resolve().kind == "python_callback" + with rebound(facade, "semantic_cache_scope", "end_user"): + assert resolver.resolve().kind == "python_callback" + with rebound(facade.cache, "embedding_model", "other-model"): + assert resolver.resolve().kind == "python_callback" + with rebound(facade.cache, "_index_name", "other-index"): + assert resolver.resolve().kind == "python_callback" + with rebound(facade.cache, "CACHE_KEY_FIELD_NAME", "other-field"): + assert resolver.resolve().kind == "python_callback" + + def patched_embedding(self: object, prompt: str, metadata: object = None) -> list[float]: + return _semantic_embedding(prompt) + + monkeypatch.setattr(RedisSemanticCache, "_get_embedding", patched_embedding) + assert resolver.resolve().kind == "python_callback" + + +def test_redis_semantic_handle_rejects_wrong_backends( + redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding +) -> None: + url, index = redis_stack + + class CustomSemanticCache(RedisSemanticCache): + pass + + with pytest.raises(TypeError, match="built-in RedisSemanticCache"): + _CacheTestHandle.redis_semantic(object()) + with pytest.raises(TypeError, match="built-in RedisSemanticCache"): + _CacheTestHandle.redis_semantic( + CustomSemanticCache( + redis_url=url, + similarity_threshold=0.8, + embedding_model=SEMANTIC_EMBEDDING_MODEL, + index_name=f"{index}_subclass", + ) + ) + + facade: Final = semantic_facade(url, index) + with pytest.raises(TypeError, match="backend types must match"): + _CacheTestHandle.redis(url)._bind_facade(facade) + + subclassed_facade: Final = Cache( + type=LiteLLMCacheType.REDIS_SEMANTIC, + redis_url=url, + similarity_threshold=0.8, + redis_semantic_cache_embedding_model=SEMANTIC_EMBEDDING_MODEL, + redis_semantic_cache_index_name=index, + ) + subclassed_facade.cache = CustomSemanticCache( # pyright: ignore[reportAttributeAccessIssue] # facade backend slot is not declared + + redis_url=url, + similarity_threshold=0.8, + embedding_model=SEMANTIC_EMBEDDING_MODEL, + index_name=index, + ) + with pytest.raises(TypeError): + _CacheTestHandle.redis_semantic( + subclassed_facade.cache + )._bind_facade(subclassed_facade) + + replacement_facade: Final = Cache( + type=LiteLLMCacheType.REDIS_SEMANTIC, + redis_url=url, + similarity_threshold=0.8, + redis_semantic_cache_embedding_model=SEMANTIC_EMBEDDING_MODEL, + redis_semantic_cache_index_name=index, + ) + with pytest.raises(TypeError, match="must be the native embedder"): + _CacheTestHandle.redis_semantic(facade.cache)._bind_facade(replacement_facade) + + +def qdrant_facade(qdrant_url: str, collection_name: str) -> Cache: return Cache( type=LiteLLMCacheType.QDRANT_SEMANTIC, qdrant_api_base=qdrant_url, @@ -1003,8 +1771,7 @@ def qdrant_facade( def test_qdrant_semantic_facade_binds_native_and_shares_entries( - qdrant_url: str, - fake_embedding_endpoint: str, + qdrant_url: str, fake_embedding_endpoint: str ) -> None: del fake_embedding_endpoint messages: Final = [{"role": "user", "content": "shared prompt"}] @@ -1024,21 +1791,20 @@ def test_qdrant_semantic_facade_binds_native_and_shares_entries( handle._bind_facade(facade) binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)).resolve() assert binding.kind == "native" - assert binding.lookup(semantic_request("python-key", messages)) == {"id": "py"} - binding.store(semantic_request("native-key", messages), {"id": "native"}) + assert binding.lookup(qdrant_request("python-key", messages)) == {"id": "py"} + binding.store(qdrant_request("native-key", messages), {"id": "native"}) python_value: Final = facade.cache.get_cache("native-key", messages=messages) assert isinstance(python_value, dict) assert python_value["response"] == {"id": "native"} unrelated: Final = [{"role": "user", "content": "unrelated prompt"}] - assert binding.lookup(semantic_request("native-key", unrelated)) is None + assert binding.lookup(qdrant_request("native-key", unrelated)) is None assert facade.cache.get_cache("native-key", messages=unrelated) is None - assert binding.lookup(semantic_request("different-key", messages)) is None + assert binding.lookup(qdrant_request("different-key", messages)) is None assert facade.cache.get_cache("different-key", messages=messages) is None async def test_qdrant_semantic_async_parity( - qdrant_url: str, - fake_embedding_endpoint: str, + qdrant_url: str, fake_embedding_endpoint: str ) -> None: del fake_embedding_endpoint messages: Final = [{"role": "user", "content": "async prompt"}] @@ -1057,17 +1823,15 @@ async def test_qdrant_semantic_async_parity( {"timestamp": time.time(), "response": json.dumps({"id": "py"})}, messages=messages, ) - - assert await binding.async_lookup(semantic_request("python-key", messages)) == {"id": "py"} - await binding.async_store(semantic_request("native-key", messages), {"id": "native"}) + assert await binding.async_lookup(qdrant_request("python-key", messages)) == {"id": "py"} + await binding.async_store(qdrant_request("native-key", messages), {"id": "native"}) python_value: Final = await facade.cache.async_get_cache("native-key", messages=messages) assert isinstance(python_value, dict) assert python_value["response"] == {"id": "native"} async def test_qdrant_semantic_malformed_entries_and_unsupported_operations( - qdrant_url: str, - fake_embedding_endpoint: str, + qdrant_url: str, fake_embedding_endpoint: str ) -> None: del fake_embedding_endpoint messages: Final = [{"role": "user", "content": "malformed prompt"}] @@ -1100,18 +1864,17 @@ async def test_qdrant_semantic_malformed_entries_and_unsupported_operations( headers=facade.cache.headers, json=response, ) - assert binding.lookup(semantic_request(key, messages)) is None - with pytest.raises(RuntimeError, match="does not support"): - binding.lookup_batch([semantic_request(key, messages)]) - with pytest.raises(RuntimeError, match="does not support"): + assert binding.lookup(qdrant_request(key, messages)) is None + with pytest.raises(RuntimeError, match="operation is not supported"): + binding.lookup_batch([qdrant_request(key, messages)]) + with pytest.raises(RuntimeError, match="operation is not supported"): await binding.async_flush() - with pytest.raises(RuntimeError, match="does not support"): + with pytest.raises(RuntimeError, match="operation is not supported"): await binding.ping() def test_qdrant_semantic_ignores_request_expiry( - qdrant_url: str, - fake_embedding_endpoint: str, + qdrant_url: str, fake_embedding_endpoint: str ) -> None: del fake_embedding_endpoint messages: Final = [{"role": "user", "content": "persistent prompt"}] @@ -1125,20 +1888,16 @@ def test_qdrant_semantic_ignores_request_expiry( ) handle._bind_facade(facade) binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)).resolve() - binding.store( - semantic_request("persistent-key", messages, ttl_seconds=1.0), - {"id": "persistent"}, - ) + binding.store(qdrant_request("persistent-key", messages, ttl_seconds=1.0), {"id": "persistent"}) time.sleep(1.2) - assert binding.lookup(semantic_request("persistent-key", messages)) == {"id": "persistent"} + assert binding.lookup(qdrant_request("persistent-key", messages)) == {"id": "persistent"} python_value: Final = facade.cache.get_cache("persistent-key", messages=messages) assert isinstance(python_value, dict) assert python_value["response"] == {"id": "persistent"} def test_qdrant_semantic_mutation_and_projection_fallback( - qdrant_url: str, - fake_embedding_endpoint: str, + qdrant_url: str, fake_embedding_endpoint: str ) -> None: del fake_embedding_endpoint collection: Final = f"cache_{uuid4().hex}" @@ -1162,18 +1921,3 @@ def test_qdrant_semantic_mutation_and_projection_fallback( unsupported.cache.qdrant_api_base = "http://127.0.0.1:7777" with pytest.raises(TypeError, match="gRPC"): handle._bind_facade(unsupported) - from litellm.caching.qdrant_semantic_cache import QdrantSemanticCache - - class CustomQdrantSemanticCache(QdrantSemanticCache): - pass - - subclass_facade: Final = qdrant_facade(qdrant_url, f"cache_{uuid4().hex}") - subclass_facade.cache = CustomQdrantSemanticCache( - qdrant_api_base=qdrant_url, - collection_name=subclass_facade.cache.collection_name, - similarity_threshold=0.99, - embedding_model="text-embedding-3-small", - vector_size=8, - ) - with pytest.raises(TypeError): - handle._bind_facade(subclass_facade) diff --git a/tests/test_litellm_rust/test_valkey_semantic_cache_native.py b/tests/test_litellm_rust/test_valkey_semantic_cache_native.py new file mode 100644 index 00000000000..c87a9f86a80 --- /dev/null +++ b/tests/test_litellm_rust/test_valkey_semantic_cache_native.py @@ -0,0 +1,599 @@ +import asyncio +import contextvars +import hashlib +import os +import struct +import threading +import time +from collections.abc import Generator, Mapping +from types import SimpleNamespace +from typing import Final, cast +from uuid import uuid4 + +import pytest +import redis + +from litellm.caching.caching import Cache +from litellm.caching.valkey_semantic_cache import ValkeySemanticCache +from litellm.rust_bridge import _native +from litellm.types.caching import LiteLLMCacheType + +pytestmark: Final = pytest.mark.requires_rust_extension +embedding_context: Final = contextvars.ContextVar("embedding_context") + + +@pytest.fixture +def valkey_url() -> str: + url: Final = os.environ.get("LITELLM_TEST_VALKEY_URL") + if url is None: + pytest.skip("LITELLM_TEST_VALKEY_URL is not set") + return url + + +@pytest.fixture +def index_name(valkey_url: str) -> Generator[str]: + index: Final = f"litellm_test_{uuid4().hex}" + yield index + client: Final = redis.Redis.from_url(valkey_url) + try: + client.ft(index).dropindex(delete_documents=True) + except redis.ResponseError: + pass + finally: + client.close() + + +def _request(prompt: str = "semantic cache prompt") -> dict[str, object]: + return { + "key": {"preset": "key"}, + "messages": [{"role": "user", "content": prompt}], + } + + +def _field_request( + prompt: str, + metadata: Mapping[str, object], + *, + namespace: str | None = None, + litellm_metadata: Mapping[str, object] | None = None, + litellm_params: Mapping[str, object] | None = None, +) -> dict[str, object]: + request: Final = { + "key": { + "fields": [ + { + "name": "model", + "value": "gpt-4.1", + "api_parameter": True, + "internal_parameter": False, + }, + { + "name": "messages", + "value": prompt, + "api_parameter": True, + "internal_parameter": False, + }, + ], + "namespace": namespace, + }, + "messages": [{"role": "user", "content": prompt}], + "metadata": dict(metadata), + } + if litellm_metadata is not None: + request["litellm_metadata"] = dict(litellm_metadata) + if litellm_params is not None: + request["litellm_params"] = dict(litellm_params) + return request + + +def _facade( + url: str, + index_name: str, + embeddings: Mapping[str, list[float]], + *, + namespace: str | None = None, +) -> Cache: + facade: Final = Cache( + type=LiteLLMCacheType.VALKEY_SEMANTIC, + redis_url=url, + similarity_threshold=0.8, + valkey_semantic_cache_index_name=index_name, + namespace=namespace, + ) + vectors: Final = embeddings + + def embed(prompt: str, metadata: Mapping[str, object] | None = None) -> list[float]: + return vectors[prompt] + + async def async_embedding(prompt: str, metadata: dict[str, object] | None = None) -> list[float]: + return vectors[prompt] + + facade.cache._get_embedding = embed + facade.cache._get_async_embedding = async_embedding + return facade + + +def _backend( + url: str, + index_name: str, + embeddings: Mapping[str, list[float]] | None = None, +) -> ValkeySemanticCache: + vectors: Final = embeddings or {"semantic cache prompt": [1.0, 0.0]} + backend: Final = ValkeySemanticCache( + redis_url=url, + similarity_threshold=0.8, + index_name=index_name, + ) + + def embed(prompt: str, metadata: Mapping[str, object] | None = None) -> list[float]: + return vectors[prompt] + + async def async_embedding(prompt: str, metadata: dict[str, object] | None = None) -> list[float]: + return vectors[prompt] + + backend._get_embedding = embed + backend._get_async_embedding = async_embedding + return backend + + +def test_python_write_native_read( + valkey_url: str, + index_name: str, +) -> None: + backend: Final = _backend(valkey_url, index_name) + response: Final = {"answer": "python"} + backend.set_cache("key", response, messages=_request()["messages"]) + handle: Final = _native._CacheTestHandle.valkey_semantic( + valkey_url, + 0.8, + index_name, + backend, + ) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + assert binding.lookup(_request()) == response + + +def test_native_write_python_read( + valkey_url: str, + index_name: str, +) -> None: + backend: Final = _backend(valkey_url, index_name) + handle: Final = _native._CacheTestHandle.valkey_semantic( + valkey_url, + 0.8, + index_name, + backend, + ) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + response: Final = {"answer": "native"} + binding.store({**_request(), "ttl_seconds": 2.0}, response) + cached: Final = cast(Mapping[str, object], backend.get_cache("key", messages=_request()["messages"])) + assert cached["response"] == response + + +async def test_async_lookup_and_store( + valkey_url: str, + index_name: str, +) -> None: + backend: Final = _backend(valkey_url, index_name) + handle: Final = _native._CacheTestHandle.valkey_semantic( + valkey_url, + 0.8, + index_name, + backend, + ) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + request: Final = {**_request(), "ttl_seconds": 2.0} + await binding.async_store(request, {"answer": "async"}) + assert await binding.async_lookup(request) == {"answer": "async"} + + +async def test_disabled_cache_controls_skip_async_embedding( + valkey_url: str, + index_name: str, +) -> None: + backend: Final = _backend(valkey_url, index_name) + calls: Final = [] + + async def fail_embedding(prompt: str, metadata: dict[str, object] | None = None) -> list[float]: + calls.append(prompt) + raise AssertionError("embedding must not run") + + backend._get_async_embedding = fail_embedding + handle: Final = _native._CacheTestHandle.valkey_semantic( + valkey_url, + 0.8, + index_name, + backend, + ) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + controls: Final = { + "supported_call_type": True, + "configured": True, + "native_backend": True, + "default_on": True, + "caching": True, + "no_cache": False, + "no_store": False, + "use_cache": True, + } + no_read_request: Final = {**_request(), "controls": {**controls, "no_cache": True}} + assert await binding.async_lookup(no_read_request) is None + no_write_request: Final = {**_request(), "controls": {**controls, "no_store": True}} + await binding.async_store(no_write_request, {"answer": "blocked"}) + assert calls == [] + client: Final = redis.Redis.from_url(valkey_url) + assert list(client.scan_iter(f"{index_name}:*")) == [] + client.close() + + +async def test_async_embedding_runs_inline_in_caller_task( + valkey_url: str, + index_name: str, +) -> None: + backend: Final = _backend(valkey_url, index_name) + observed: dict[str, object] = {} + + async def async_embedding(prompt: str, metadata: dict[str, object] | None = None) -> list[float]: + observed["context"] = embedding_context.get("missing") + observed["task"] = asyncio.current_task() + observed["thread"] = threading.get_ident() + embedding_context.set("embedder") + return [1.0, 0.0] + + backend._get_async_embedding = async_embedding + handle: Final = _native._CacheTestHandle.valkey_semantic( + valkey_url, + 0.8, + index_name, + backend, + ) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + request: Final = {**_request(), "ttl_seconds": 2.0} + caller_task: Final = asyncio.current_task() + caller_thread: Final = threading.get_ident() + token: Final = embedding_context.set("caller") + try: + await binding.async_store(request, {"answer": "inline"}) + assert observed["context"] == "caller" + assert observed["task"] is caller_task + assert observed["thread"] == caller_thread + assert embedding_context.get() == "embedder" + assert await binding.async_lookup(request) == {"answer": "inline"} + finally: + embedding_context.reset(token) + + +def test_facade_activation_and_mutation_fallback( + valkey_url: str, + index_name: str, +) -> None: + facade: Final = Cache( + type=LiteLLMCacheType.VALKEY_SEMANTIC, + redis_url=valkey_url, + similarity_threshold=0.8, + valkey_semantic_cache_index_name=index_name, + ) + handle: Final = _native._CacheTestHandle.valkey_semantic( + valkey_url, + 0.8, + index_name, + facade.cache, + ) + handle._bind_facade(facade) + resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)) + assert resolver.resolve().kind == "native" + facade.cache.similarity_threshold = 0.7 + assert resolver.resolve().kind == "python_callback" + + +def test_batch_lookup_is_unsupported( + valkey_url: str, + index_name: str, +) -> None: + backend: Final = _backend(valkey_url, index_name) + handle: Final = _native._CacheTestHandle.valkey_semantic( + valkey_url, + 0.8, + index_name, + backend, + ) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + with pytest.raises(NotImplementedError): + binding.lookup_batch([_request()]) + + +def test_ttl_expiry( + valkey_url: str, + index_name: str, +) -> None: + backend: Final = _backend(valkey_url, index_name) + handle: Final = _native._CacheTestHandle.valkey_semantic(valkey_url, 0.8, index_name, backend) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + binding.store({**_request(), "ttl_seconds": 1.0}, {"answer": "expires"}) + client: Final = redis.Redis.from_url(valkey_url) + documents: Final = list(client.scan_iter(f"{index_name}:*")) + assert len(documents) == 1 + assert client.ttl(documents[0]) > 0 + time.sleep(1.5) + assert binding.lookup(_request()) is None + + +def test_no_ttl_is_persistent_and_python_reads_native_value( + valkey_url: str, + index_name: str, +) -> None: + backend: Final = _backend(valkey_url, index_name) + handle: Final = _native._CacheTestHandle.valkey_semantic(valkey_url, 0.8, index_name, backend) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + response: Final = {"answer": "persistent"} + binding.store(_request(), response) + client: Final = redis.Redis.from_url(valkey_url) + documents: Final = list(client.scan_iter(f"{index_name}:*")) + assert len(documents) == 1 + assert client.ttl(documents[0]) == -1 + cached: Final = cast(Mapping[str, object], backend.get_cache("key", messages=_request()["messages"])) + assert cached["response"] == response + + +def test_below_threshold_misses_on_native_and_python( + valkey_url: str, + index_name: str, +) -> None: + backend: Final = _backend( + valkey_url, + index_name, + {"prompt A": [1.0, 0.0], "prompt B": [0.0, 1.0]}, + ) + handle: Final = _native._CacheTestHandle.valkey_semantic(valkey_url, 0.8, index_name, backend) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + binding.store(_request("prompt A"), {"answer": "A"}) + assert binding.lookup(_request("prompt B")) is None + assert backend.get_cache("key", messages=_request("prompt B")["messages"]) is None + + +def test_malformed_entry_is_a_miss_on_native_and_python( + valkey_url: str, + index_name: str, +) -> None: + backend: Final = _backend(valkey_url, index_name) + client: Final = redis.Redis.from_url(valkey_url) + scope: Final = hashlib.sha256(b"key").hexdigest() + document: Final = f"{index_name}:{scope}:{uuid4().hex}" + client.hset( + document, + mapping={ + "litellm_cache_key": scope, + "prompt": "semantic cache prompt", + "response": "not json", + "embedding": struct.pack("<2f", 1.0, 0.0), + }, + ) + handle: Final = _native._CacheTestHandle.valkey_semantic(valkey_url, 0.8, index_name, backend) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + assert binding.lookup(_request()) is None + assert backend.get_cache("key", messages=_request()["messages"]) is None + + +def test_mixed_content_parts_match_python_semantic_behavior( + valkey_url: str, + index_name: str, +) -> None: + backend: Final = _backend(valkey_url, index_name) + messages: Final = [{"role": "user", "content": ["raw", {"text": "hello"}]}] + backend.set_cache("key", {"answer": "mixed"}, messages=messages) + assert backend.get_cache("key", messages=messages) is None + + handle: Final = _native._CacheTestHandle.valkey_semantic(valkey_url, 0.8, index_name, backend) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + request: Final = {**_request(), "messages": messages} + binding.store(request, {"answer": "mixed"}) + assert binding.lookup(request) is None + client: Final = redis.Redis.from_url(valkey_url) + assert list(client.scan_iter(f"{index_name}:*")) == [] + client.close() + + +async def test_async_store_batch_and_lookup( + valkey_url: str, + index_name: str, +) -> None: + backend: Final = _backend( + valkey_url, + index_name, + {"prompt A": [1.0, 0.0], "prompt B": [0.0, 1.0]}, + ) + sync_calls: Final = [] + async_tasks: Final = [] + + def sync_embedding(prompt: str, metadata: Mapping[str, object] | None = None) -> list[float]: + sync_calls.append(prompt) + return {"prompt A": [1.0, 0.0], "prompt B": [0.0, 1.0]}[prompt] + + async def async_embedding( + prompt: str, + metadata: dict[str, object] | None = None, + ) -> list[float]: + async_tasks.append(asyncio.current_task()) + return {"prompt A": [1.0, 0.0], "prompt B": [0.0, 1.0]}[prompt] + + backend._get_embedding = sync_embedding + backend._get_async_embedding = async_embedding + handle: Final = _native._CacheTestHandle.valkey_semantic(valkey_url, 0.8, index_name, backend) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + requests: Final = [_request("prompt A"), _request("prompt B")] + responses: Final = [{"answer": "A"}, {"answer": "B"}] + caller_task: Final = asyncio.current_task() + await binding.async_store_batch(requests, responses) + assert sync_calls == [] + assert async_tasks + assert all(task is caller_task for task in async_tasks) + assert await binding.async_lookup(requests[0]) == responses[0] + assert await binding.async_lookup(requests[1]) == responses[1] + + +def test_subclass_backend_falls_back_to_python( + valkey_url: str, + index_name: str, +) -> None: + class Custom(ValkeySemanticCache): + pass + + facade: Final = Cache( + type=LiteLLMCacheType.VALKEY_SEMANTIC, + redis_url=valkey_url, + similarity_threshold=0.8, + valkey_semantic_cache_index_name=index_name, + ) + facade.cache = Custom(redis_url=valkey_url, similarity_threshold=0.8, index_name=index_name) + resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)) + assert resolver.resolve().kind == "python_callback" + + +def test_field_key_matches_python_semantic_scope( + valkey_url: str, + index_name: str, +) -> None: + facade: Final = _facade(valkey_url, index_name, {"semantic cache prompt": [1.0, 0.0]}) + metadata: Final = {"user_api_key": "k1"} + expected: Final = facade.get_cache_key( + model="gpt-4.1", + messages=[{"role": "user", "content": "semantic cache prompt"}], + metadata=metadata, + ) + handle: Final = _native._CacheTestHandle.valkey_semantic( + valkey_url, + 0.8, + index_name, + facade.cache, + ) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + binding.store(_field_request("semantic cache prompt", metadata), {"answer": "scoped"}) + client: Final = redis.Redis.from_url(valkey_url) + documents: Final = list(client.scan_iter(f"{index_name}:*")) + assert len(documents) == 1 + document_parts: Final = documents[0].decode().split(":") + assert document_parts[1] == hashlib.sha256(expected.encode()).hexdigest() + client.close() + + +def test_field_key_reads_all_python_tenant_metadata_sources( + valkey_url: str, + index_name: str, +) -> None: + facade: Final = _facade(valkey_url, index_name, {"semantic cache prompt": [1.0, 0.0]}) + params_metadata: Final = {"user_api_key_team_id": "team-from-params"} + expected: Final = facade.get_cache_key( + model="gpt-4.1", + messages=[{"role": "user", "content": "semantic cache prompt"}], + metadata={}, + litellm_params={"metadata": params_metadata}, + ) + handle: Final = _native._CacheTestHandle.valkey_semantic( + valkey_url, + 0.8, + index_name, + facade.cache, + ) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + binding.store( + _field_request( + "semantic cache prompt", + {}, + litellm_params={"metadata": params_metadata}, + ), + {"answer": "params"}, + ) + client: Final = redis.Redis.from_url(valkey_url) + documents: Final = list(client.scan_iter(f"{index_name}:*")) + assert len(documents) == 1 + document_parts: Final = documents[0].decode().split(":") + assert document_parts[1] == hashlib.sha256(expected.encode()).hexdigest() + client.close() + + assert ( + binding.lookup( + _field_request( + "semantic cache prompt", + {}, + litellm_metadata={"user_api_key_team_id": "team-from-litellm"}, + ) + ) + is None + ) + + +def test_namespace_isolates_semantic_entries( + valkey_url: str, + index_name: str, +) -> None: + facade: Final = _facade( + valkey_url, + index_name, + {"semantic cache prompt": [1.0, 0.0]}, + namespace="team-a", + ) + handle: Final = _native._CacheTestHandle.valkey_semantic( + valkey_url, + 0.8, + index_name, + facade.cache, + ) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + team_a: Final = _field_request("semantic cache prompt", {}, namespace="team-a") + team_b: Final = _field_request("semantic cache prompt", {}, namespace="team-b") + binding.store(team_a, {"answer": "team-a"}) + assert binding.lookup(team_b) is None + assert binding.lookup(team_a) == {"answer": "team-a"} + cached: Final = cast( + Mapping[str, object], + facade.get_cache( + model="gpt-4.1", + messages=[{"role": "user", "content": "semantic cache prompt"}], + ), + ) + assert cached == {"answer": "team-a"} + + +def test_field_key_isolates_tenant_scope( + valkey_url: str, + index_name: str, +) -> None: + facade: Final = _facade(valkey_url, index_name, {"semantic cache prompt": [1.0, 0.0]}) + handle: Final = _native._CacheTestHandle.valkey_semantic( + valkey_url, + 0.8, + index_name, + facade.cache, + ) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + binding.store( + _field_request("semantic cache prompt", {"user_api_key": "k1"}), + {"answer": "tenant one"}, + ) + assert binding.lookup(_field_request("semantic cache prompt", {"user_api_key": "k2"})) is None + assert binding.lookup(_field_request("semantic cache prompt", {"user_api_key": "k1"})) == {"answer": "tenant one"} + + +def test_tls_valkey_facade_falls_back_to_python( + index_name: str, +) -> None: + facade: Final = Cache( + type=LiteLLMCacheType.VALKEY_SEMANTIC, + redis_url="rediss://127.0.0.1:6390/0", + similarity_threshold=0.8, + valkey_semantic_cache_index_name=index_name, + ) + resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)) + assert resolver.resolve().kind == "python_callback" + + +async def test_ping_maps_unsupported_native_operation_to_not_implemented( + valkey_url: str, + index_name: str, +) -> None: + backend: Final = _backend(valkey_url, index_name) + handle: Final = _native._CacheTestHandle.valkey_semantic(valkey_url, 0.8, index_name, backend) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve() + with pytest.raises(NotImplementedError): + await binding.ping() diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.test.tsx index 5c7453c1394..a144630cdd0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.test.tsx @@ -434,7 +434,7 @@ describe("AutoRouterBenchmarksTab", () => { mockHook({ data: response([group()]) }); const { dateValue, onDateChange } = renderTab(); - expect(vi.mocked(useAutoRouterBenchmarks)).toHaveBeenCalledWith("sk-test", dateValue, undefined); + expect(vi.mocked(useAutoRouterBenchmarks)).toHaveBeenCalledWith("sk-test", dateValue, undefined, undefined); expect(screen.getByText("Jul 6 – Aug 5 (UTC)")).toBeInTheDocument(); fireEvent.click(screen.getByTestId("date-picker")); @@ -460,7 +460,7 @@ describe("AutoRouterBenchmarksTab", () => { , ); - expect(vi.mocked(useAutoRouterBenchmarks)).toHaveBeenCalledWith("sk-test", dateValue, "key-hash-1"); + expect(vi.mocked(useAutoRouterBenchmarks)).toHaveBeenCalledWith("sk-test", dateValue, "key-hash-1", undefined); expect(screen.getByText("Total estimated savings")).toBeInTheDocument(); expect(screen.queryByRole("tab", { name: "Shadow Evals" })).not.toBeInTheDocument(); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.tsx index ce55b633b60..063598bd46e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.tsx @@ -312,8 +312,7 @@ const BenchmarksBody: React.FC = ({ isPending, error, data, length. Total actual spend includes every turn; savings and baseline spend include only turns with a current estimate, including turns with zero savings. Savings are net of recorded LLM classification cost. Classification cost per 1K turns is averaged over all auto-router turns, including those that skip classification. The range - counts whole sessions that overlap it, so totals can differ slightly from the Overall tab, which buckets savings - by UTC day. + counts whole sessions that overlap it, so totals can differ from savings views that group usage by UTC day.

@@ -333,11 +332,17 @@ interface AutoRouterBenchmarksTabProps { accessToken: string | null; activity: Pick; apiKey?: string; + userId?: string; } -export const AutoRouterUsageView: React.FC = ({ accessToken, activity, apiKey }) => { +export const AutoRouterUsageView: React.FC = ({ + accessToken, + activity, + apiKey, + userId, +}) => { const { dateValue, onDateChange } = activity; - const { data, isPending, error } = useAutoRouterBenchmarks(accessToken, dateValue, apiKey); + const { data, isPending, error } = useAutoRouterBenchmarks(accessToken, dateValue, apiKey, userId); const [selectedKey, setSelectedKey] = useState(ALL_ROUTERS); const { data: autoRouters } = useAutoRouters(); @@ -372,6 +377,12 @@ export const AutoRouterUsageView: React.FC = ({ ac
+ {userId && ( +

+ Usage for this user across API keys and JWT-authenticated requests. Older sessions recorded without a user ID + are not included. +

+ )} +export const useAutoRouterBenchmarks = ( + accessToken: string | null, + range: DateRange, + apiKey?: string, + userId?: string, +) => $api.useQuery( "get", "/auto_router/benchmarks", - { params: { query: { ...benchmarksWindow(range, new Date()), api_key: apiKey } } }, + { params: { query: { ...benchmarksWindow(range, new Date()), api_key: apiKey, user_id: userId } } }, { enabled: Boolean(accessToken && range.from && range.to), retry: false }, ); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.test.tsx index 4059303d5a5..e501cf00b90 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.test.tsx @@ -15,6 +15,8 @@ vi.mock("@/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity", ( isFetchingMore: false, progress: { currentPage: 4, totalPages: 9 }, cancelled: false, + failed: false, + coversRange: true, cancel: mockCancel, }; }, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.ts index 92dd24b8d6d..4eb9f257d30 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.ts @@ -67,14 +67,16 @@ export const useScopedDailyActivityRange = ( args: [accessToken, startTime, endTime, userId, true, apiKey], enabled: !!accessToken && !!startTime && !!endTime, }; - const { data, loading, isFetchingMore, progress, cancelled, failed, cancel } = + const { data, loading, isFetchingMore, progress, cancelled, failed, coversRange, cancel } = usePaginatedDailyActivity(activityQueryOptions); + const readUnavailable = failed || cancelled; + const waitingForRange = activityQueryOptions.enabled && !coversRange && !readUnavailable; return { dateValue, onDateChange, results: data.results as DailyData[], - loading, + loading: loading || waitingForRange, isFetchingMore, progress, cancelled, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx index 091b2f1403f..1217d878489 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx @@ -3,7 +3,8 @@ import { render, waitFor, screen, act, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { describe, it, expect, vi, beforeEach } from "vitest"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; -import MCPServers from "./mcp_servers"; +import MCPServers, { compareServers, type SortKey } from "./mcp_servers"; +import type { MCPServer } from "@/components/mcp_tools/types"; import * as networking from "@/components/networking"; // Mock the networking module @@ -31,6 +32,98 @@ const createQueryClient = () => }, }); +describe("compareServers", () => { + const server = (server_id: string, name: string, created_at = ""): MCPServer => ({ + server_id, + server_name: name, + created_at, + updated_at: created_at, + created_by: "user", + updated_by: "user", + }); + + const shuffled = [server("c", "github"), server("a", "slack"), server("b", "Jira")]; + + it("orders servers without timestamps by name so config.yaml servers render in a stable order", () => { + const byCreated = [...shuffled].sort((a, b) => compareServers(a, b, "created_desc")).map((s) => s.server_id); + const byUpdated = [...shuffled].sort((a, b) => compareServers(a, b, "updated_desc")).map((s) => s.server_id); + const byHealth = [...shuffled].sort((a, b) => compareServers(a, b, "health")).map((s) => s.server_id); + + expect(byCreated).toEqual(["c", "b", "a"]); + expect(byUpdated).toEqual(["c", "b", "a"]); + expect(byHealth).toEqual(["c", "b", "a"]); + }); + + it("keeps newest-first when timestamps differ", () => { + const newest = server("new", "zzz", "2026-02-01T00:00:00Z"); + const oldest = server("old", "aaa", "2026-01-01T00:00:00Z"); + expect([oldest, newest].sort((a, b) => compareServers(a, b, "created_desc")).map((s) => s.server_id)).toEqual([ + "new", + "old", + ]); + }); + + it.each(["created_desc", "updated_desc", "name_asc", "health"])( + "breaks equal timestamps and names by ID for %s regardless of input order", + (sort) => { + const servers = [ + server("b", "GitHub", "2026-01-01T00:00:00Z"), + server("c", "Slack", "2026-01-01T00:00:00Z"), + server("a", "github", "2026-01-01T00:00:00Z"), + ]; + for (const input of [servers, [...servers].reverse()]) { + expect([...input].sort((a, b) => compareServers(a, b, sort)).map((s) => s.server_id)).toEqual(["a", "b", "c"]); + } + }, + ); + + it("uses the display name before alias, then falls back to alias and ID", () => { + const servers: MCPServer[] = [ + { ...server("s-slack", "Slack"), alias: "aaa" }, + { ...server("s-github", ""), server_name: null, alias: "GitHub" }, + { ...server("confluence", ""), alias: "" }, + ]; + for (const input of [servers, [...servers].reverse()]) { + expect([...input].sort((a, b) => compareServers(a, b, "name_asc")).map((s) => s.server_id)).toEqual([ + "confluence", + "s-github", + "s-slack", + ]); + } + }); + + it.each(["created_desc", "updated_desc", "health"])( + "keeps timestamped servers before missing timestamps for %s", + (sort) => { + const servers = [ + server("config", "aaa"), + server("older", "bbb", "2026-01-01T00:00:00Z"), + server("newer", "zzz", "2026-02-01T00:00:00Z"), + ]; + for (const input of [servers, [...servers].reverse()]) { + expect([...input].sort((a, b) => compareServers(a, b, sort)).map((s) => s.server_id)).toEqual([ + "newer", + "older", + "config", + ]); + } + }, + ); + + it("sorts health before recency and display name", () => { + const servers: MCPServer[] = [ + { ...server("healthy", "aaa", "2026-03-01T00:00:00Z"), status: "healthy" }, + { ...server("unknown", "bbb", "2026-02-01T00:00:00Z"), status: "unknown" }, + { ...server("unhealthy", "zzz", "2026-01-01T00:00:00Z"), status: "unhealthy" }, + ]; + expect(servers.sort((a, b) => compareServers(a, b, "health")).map((s) => s.server_id)).toEqual([ + "unhealthy", + "unknown", + "healthy", + ]); + }); +}); + describe("MCPServers", () => { const defaultProps = { accessToken: "123", @@ -74,6 +167,134 @@ describe("MCPServers", () => { const myConnections = await screen.findByRole("link", { name: "My Connections" }); expect(myConnections).toBeVisible(); expect(myConnections).toHaveAttribute("href", "/ui/connect"); + for (const name of ["Semantic Filter", "Tool Search", "Network Settings", "Submitted MCPs"]) { + const tab = screen.queryByRole("tab", { name }); + if (userRole === "Admin") { + expect(tab).toBeVisible(); + } else { + expect(tab).not.toBeInTheDocument(); + } + } + expect( + screen.getByRole("button", { + name: userRole === "Admin" ? "+ Add New MCP Server" : "+ Submit MCP Server", + }), + ).toBeVisible(); + }); + + it.each(["cancel", "success", "failure", "unnamed"])("preserves delete confirmation on %s", async (outcome) => { + const server: MCPServer = { + created_at: "", + updated_at: "", + server_id: "delete-server", + server_name: outcome === "unnamed" ? null : "Delete fixture", + alias: "delete-alias", + url: outcome === "unnamed" ? null : "https://example.com/mcp", + created_by: "user", + updated_by: "user", + }; + vi.mocked(networking.fetchMCPServers).mockResolvedValue([server]); + vi.mocked(networking.fetchMCPServerHealth).mockResolvedValue([]); + let finishDelete: () => void = () => {}; + vi.mocked(networking.deleteMCPServer).mockImplementation( + () => + new Promise((resolve, reject) => { + finishDelete = () => (outcome === "failure" ? reject(new Error("Delete failed")) : resolve(undefined)); + }), + ); + render( + + + , + ); + await userEvent.click(await screen.findByRole("button", { name: "Server actions" })); + await userEvent.click(await screen.findByRole("menuitem", { name: "Delete" })); + const dialog = await screen.findByRole("alertdialog", { name: "Delete MCP Server?" }); + expect(within(dialog).getByText("delete-server")).toBeVisible(); + if (outcome === "unnamed") { + expect(within(dialog).queryByText("Name")).not.toBeInTheDocument(); + expect(within(dialog).queryByText("URL")).not.toBeInTheDocument(); + } else { + expect(within(dialog).getByText("Delete fixture")).toBeVisible(); + expect(within(dialog).getByText("https://example.com/mcp")).toBeVisible(); + } + if (outcome === "cancel") { + await userEvent.click(within(dialog).getByRole("button", { name: "Cancel" })); + expect(networking.deleteMCPServer).not.toHaveBeenCalled(); + } else { + await userEvent.click(within(dialog).getByRole("button", { name: "Delete" })); + expect(within(dialog).getByRole("button", { name: "Deleting..." })).toBeDisabled(); + expect(within(dialog).getByRole("button", { name: "Cancel" })).toBeDisabled(); + expect(networking.deleteMCPServer).toHaveBeenCalledWith("123", "delete-server"); + await act(async () => finishDelete()); + } + await waitFor(() => expect(screen.queryByRole("alertdialog")).not.toBeInTheDocument()); + }); + + it("filters servers by access group", async () => { + const server = { created_by: "user", updated_by: "user" }; + vi.mocked(networking.fetchMCPServers).mockResolvedValue([ + { + ...server, + server_id: "string-group", + server_name: "String group", + alias: "string-alias", + mcp_access_groups: ["shared"], + }, + { + ...server, + server_id: "legacy-group", + server_name: "Legacy group", + alias: "legacy-alias", + mcp_access_groups: ["shared"], + }, + { + ...server, + server_id: "other-group", + server_name: "Other group", + alias: "other-alias", + mcp_access_groups: ["different"], + }, + ]); + vi.mocked(networking.fetchMCPServerHealth).mockResolvedValue([]); + render( + + + , + ); + await screen.findByText("String group"); + await userEvent.click(screen.getByRole("combobox", { name: "Access Group" })); + await userEvent.click(await screen.findByRole("option", { name: "shared" })); + expect(screen.getByText("String group")).toBeVisible(); + expect(screen.getByText("Legacy group")).toBeVisible(); + expect(screen.queryByText("Other group")).not.toBeInTheDocument(); + }); + + it.each(["server_name", "alias", "url", "server_id"] as const)("searches by %s case-insensitively", async (field) => { + const server: MCPServer = { + created_at: "", + updated_at: "", + server_id: "search-server", + server_name: "Search fixture", + created_by: "user", + updated_by: "user", + [field]: "Needle", + }; + vi.mocked(networking.fetchMCPServers).mockResolvedValue([server]); + vi.mocked(networking.fetchMCPServerHealth).mockResolvedValue([]); + render( + + + , + ); + await screen.findByTestId("mcp-servers-grid"); + const search = screen.getByPlaceholderText("Search by name, alias, URL, or ID"); + await userEvent.type(search, " NEEDLE "); + expect(screen.getByTestId("mcp-servers-grid")).toBeVisible(); + await userEvent.clear(search); + await userEvent.type(search, "no-match"); + expect(screen.queryByTestId("mcp-servers-grid")).not.toBeInTheDocument(); + expect(screen.getByText("No servers match the current filters or search.")).toBeVisible(); }); it("should render mocked MCP servers data in the table", async () => { @@ -316,9 +537,7 @@ describe("MCPServers", () => { expect(screen.getByText("Team B Server")).toBeInTheDocument(); expect(screen.getByText("Team A Server 2")).toBeInTheDocument(); - // Find the team select by its "Team" label, then the combobox it labels - const teamLabel = screen.getByText("Team"); - const teamSelect = within(teamLabel.parentElement!).getByRole("combobox"); + const teamSelect = screen.getByRole("combobox", { name: "Team" }); await userEvent.click(teamSelect); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx index 2df304eff96..818b5150650 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx @@ -49,7 +49,7 @@ import { cn } from "@/lib/cva.config"; import UserEnvVarsModal from "./UserEnvVarsModal"; import { listMCPUserEnvVarStatus } from "@/components/networking"; -type SortKey = "created_desc" | "updated_desc" | "name_asc" | "health"; +export type SortKey = "created_desc" | "updated_desc" | "name_asc" | "health"; const SORT_OPTIONS: { value: SortKey; label: string }[] = [ { value: "created_desc", label: "Recently created" }, @@ -64,32 +64,33 @@ const HEALTH_RANK: Record = { healthy: 2, }; -const compareServers = (a: MCPServer, b: MCPServer, sort: SortKey): number => { +const compareByName = (a: MCPServer, b: MCPServer): number => { + const nameA = (a.server_name || a.alias || a.server_id).toLowerCase(); + const nameB = (b.server_name || b.alias || b.server_id).toLowerCase(); + return nameA.localeCompare(nameB) || a.server_id.localeCompare(b.server_id); +}; + +const compareByTimestampDesc = (a: string | null | undefined, b: string | null | undefined): number => { + const ta = a ? new Date(a).getTime() : 0; + const tb = b ? new Date(b).getTime() : 0; + return tb - ta; +}; + +export const compareServers = (a: MCPServer, b: MCPServer, sort: SortKey): number => { switch (sort) { - case "name_asc": { - const nameA = (a.server_name || a.alias || a.server_id).toLowerCase(); - const nameB = (b.server_name || b.alias || b.server_id).toLowerCase(); - return nameA.localeCompare(nameB); - } - case "updated_desc": { - const ta = a.updated_at ? new Date(a.updated_at).getTime() : 0; - const tb = b.updated_at ? new Date(b.updated_at).getTime() : 0; - return tb - ta; - } + case "name_asc": + return compareByName(a, b); + case "updated_desc": + return compareByTimestampDesc(a.updated_at, b.updated_at) || compareByName(a, b); case "health": { const ra = HEALTH_RANK[a.status ?? "unknown"] ?? 1; const rb = HEALTH_RANK[b.status ?? "unknown"] ?? 1; if (ra !== rb) return ra - rb; - const ta = a.created_at ? new Date(a.created_at).getTime() : 0; - const tb = b.created_at ? new Date(b.created_at).getTime() : 0; - return tb - ta; + return compareByTimestampDesc(a.created_at, b.created_at) || compareByName(a, b); } case "created_desc": - default: { - const ta = a.created_at ? new Date(a.created_at).getTime() : 0; - const tb = b.created_at ? new Date(b.created_at).getTime() : 0; - return tb - ta; - } + default: + return compareByTimestampDesc(a.created_at, b.created_at) || compareByName(a, b); } }; @@ -112,6 +113,62 @@ const readToolsOAuthServerId = (): string | null => { } }; +function DeleteServerDialog({ + open, + onOpenChange, + server, + isDeleting, + onConfirm, +}: { + open: boolean; + onOpenChange: (open: boolean) => void; + server: MCPServer | undefined; + isDeleting: boolean; + onConfirm: () => Promise; +}) { + return ( + + + + Delete MCP Server? + +
+

+ This action is permanent and cannot be undone. All associated configurations will be removed. +

+ + {server && ( +
+ {server.server_name && ( +
+
Name
+
{server.server_name}
+
+ )} +
+
ID
+
{server.server_id}
+
+ {server.url && ( +
+
URL
+
{server.url}
+
+ )} +
+ )} +
+ + Cancel + + +
+
+ ); +} + const MCPServers: React.FC = ({ accessToken, userRole, userID, isViewOnly = false }) => { const { data: mcpServers, isLoading: isLoadingServers, refetch } = useMCPServers(); @@ -298,16 +355,12 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID, i } if (group !== "all") { filtered = filtered.filter((server) => - server.mcp_access_groups?.some((g: any) => (typeof g === "string" ? g === group : g && g.name === group)), + server.mcp_access_groups?.some((g: string | { name?: string } | null) => + typeof g === "string" ? g === group : g?.name === group, + ), ); } - const sorted = [...filtered].sort((a, b) => { - if (!a.created_at && !b.created_at) return 0; - if (!a.created_at) return 1; - if (!b.created_at) return -1; - return new Date(b.created_at).getTime() - new Date(a.created_at).getTime(); - }); - setFilteredServers(sorted); + setFilteredServers(filtered); }, [serversWithHealth], ); @@ -338,7 +391,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID, i const alias = (s.alias || "").toLowerCase(); const url = (s.url || "").toLowerCase(); const id = s.server_id.toLowerCase(); - return name.includes(q) || alias.includes(q) || url.includes(q) || id.includes(q); + return [name, alias, url, id].some((value) => value.includes(q)); }) : filteredServers; return [...matches].sort((a, b) => compareServers(a, b, sortKey)); @@ -381,9 +434,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID, i }; // Find the server to delete from the servers list - const serverToDelete = serverIdToDelete - ? (mcpServers || []).find((server) => server.server_id === serverIdToDelete) - : null; + const serverToDelete = mcpServers?.find((server) => server.server_id === serverIdToDelete); const handleCreateSuccess = (newMcpServer: MCPServer) => { setFilteredServers((prev) => [...prev, newMcpServer]); @@ -425,45 +476,13 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID, i return (
- !open && cancelDelete()}> - - - Delete MCP Server? - -
-

- This action is permanent and cannot be undone. All associated configurations will be removed. -

- - {serverToDelete && ( -
- {serverToDelete.server_name && ( -
-
Name
-
{serverToDelete.server_name}
-
- )} -
-
ID
-
{serverToDelete.server_id}
-
- {serverToDelete.url && ( -
-
URL
-
{serverToDelete.url}
-
- )} -
- )} -
- - Cancel - - -
-
+ !open && cancelDelete()} + server={serverToDelete} + isDeleting={isDeletingServer} + onConfirm={confirmDelete} + /> = ({ accessToken, userRole, userID, i My Connections - {isAdminRole(userRole) && ( + {isAdminRole(userRole) ? ( <> - )} - {!isAdminRole(userRole) && ( + ) : (
When a prompt provably cannot fit the decided tier's context windows, route it to the lowest tier whose - window holds it instead of letting the provider reject it. Off means requests dispatch on complexity alone. + window holds it instead of letting the provider reject it. Disabled by default. Off means requests dispatch on + complexity alone. {enabled && (
diff --git a/ui/litellm-dashboard/src/components/add_model/HeuristicKeywordOverrides.tsx b/ui/litellm-dashboard/src/components/add_model/HeuristicKeywordOverrides.tsx new file mode 100644 index 00000000000..185f187bbfc --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/HeuristicKeywordOverrides.tsx @@ -0,0 +1,42 @@ +import React from "react"; +import { MultiSelect } from "@/components/shared/MultiSelect"; +import type { ComplexityRouterConfigValue } from "./ComplexityRouterConfig"; + +const fields = [ + ["code_keywords", "Code keywords"], + ["reasoning_keywords", "Reasoning keywords"], + ["technical_keywords", "Technical keywords"], + ["simple_keywords", "Simple keywords"], +] as const; + +const HeuristicKeywordOverrides: React.FC<{ + value: ComplexityRouterConfigValue; + onChange: (value: ComplexityRouterConfigValue) => void; +}> = ({ value, onChange }) => ( +
+

+ Each list replaces the built-in keyword list of the same name for the heuristic scorer. Leave a list empty to keep + the built-in one. To add technical terms without replacing the list, use custom technical keywords under + Classification Method. +

+ {fields.map(([key, label]) => { + const keywords = value[key] ?? []; + return ( +
+ {label} + ({ label: keyword, value: keyword }))} + value={keywords} + onValueChange={(next) => onChange({ ...value, [key]: next.length > 0 ? next : undefined })} + placeholder={`Add ${label.toLowerCase()}`} + emptyText="Type to add a keyword" + allowCustomValues + className="w-full" + /> +
+ ); + })} +
+); + +export default HeuristicKeywordOverrides; diff --git a/ui/litellm-dashboard/src/components/add_model/HousekeepingRoutingControls.tsx b/ui/litellm-dashboard/src/components/add_model/HousekeepingRoutingControls.tsx new file mode 100644 index 00000000000..2b6caba6291 --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/HousekeepingRoutingControls.tsx @@ -0,0 +1,44 @@ +import React from "react"; +import { MultiSelect } from "@/components/shared/MultiSelect"; +import { Switch } from "@/components/ui/switch"; +import type { ComplexityRouterConfigValue } from "./ComplexityRouterConfig"; + +const HousekeepingRoutingControls: React.FC<{ + value: ComplexityRouterConfigValue; + onChange: (value: ComplexityRouterConfigValue) => void; +}> = ({ value, onChange }) => { + const enabled = value.route_housekeeping_to_cheapest_tier ?? true; + const patterns = value.housekeeping_patterns ?? []; + return ( + <> +
+ onChange({ ...value, route_housekeeping_to_cheapest_tier: next })} + aria-label="Route housekeeping calls to the cheapest tier" + /> + Route housekeeping calls to the cheapest tier +
+ + Conversation-title style calls skip the classifier and go to the cheapest tier. + + Additional housekeeping sentinels + ({ label: pattern, value: pattern }))} + value={patterns} + onValueChange={(next) => onChange({ ...value, housekeeping_patterns: next.length > 0 ? next : undefined })} + placeholder="e.g., conversation title" + emptyText="Type to add a sentinel" + allowCustomValues + disabled={!enabled} + className="w-full" + /> + + Case-sensitive literal strings added to the built-in conversation-title sentinels. + {!enabled && " Turn housekeeping routing on for these to take effect."} + + + ); +}; + +export default HousekeepingRoutingControls; diff --git a/ui/litellm-dashboard/src/components/add_model/PlanModeOverrideControls.tsx b/ui/litellm-dashboard/src/components/add_model/PlanModeOverrideControls.tsx index 83fd0c42d16..e195fb62222 100644 --- a/ui/litellm-dashboard/src/components/add_model/PlanModeOverrideControls.tsx +++ b/ui/litellm-dashboard/src/components/add_model/PlanModeOverrideControls.tsx @@ -1,5 +1,6 @@ import React from "react"; import { Switch } from "@/components/ui/switch"; +import { MultiSelect } from "@/components/shared/MultiSelect"; import TierRowSelect from "./TierRowSelect"; import type { ComplexityRouterConfigValue } from "./ComplexityRouterConfig"; @@ -38,6 +39,23 @@ const PlanModeOverrideControls: React.FC<{ />
)} +
+ Additional plan-mode sentinels + ({ label: pattern, value: pattern }))} + value={value.plan_mode_patterns ?? []} + onValueChange={(patterns) => + onChange({ ...value, plan_mode_patterns: patterns.length > 0 ? patterns : undefined }) + } + placeholder="e.g., enter plan mode" + emptyText="Type to add a sentinel" + allowCustomValues + className="w-full" + /> + + Case-sensitive literal strings added to the built-in Claude Code and Copilot plan-mode markers. + +
); diff --git a/ui/litellm-dashboard/src/components/add_model/ReminderMarkers.tsx b/ui/litellm-dashboard/src/components/add_model/ReminderMarkers.tsx new file mode 100644 index 00000000000..c7f9ee48e0b --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/ReminderMarkers.tsx @@ -0,0 +1,82 @@ +import React from "react"; +import { Plus, Trash2 } from "lucide-react"; +import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import type { ComplexityRouterConfigValue } from "./ComplexityRouterConfig"; +import { getReminderMarkersError, type ReminderMarkerPair } from "./build_complexity_router_config"; + +const ReminderMarkers: React.FC<{ + value: ComplexityRouterConfigValue; + onChange: (value: ComplexityRouterConfigValue) => void; + showValidationErrors?: boolean; +}> = ({ value, onChange, showValidationErrors = false }) => { + const markers = value.reminder_markers ?? []; + const update = (index: number, patch: Partial) => + onChange({ + ...value, + reminder_markers: markers.map((marker, markerIndex) => + markerIndex === index ? { ...marker, ...patch } : marker, + ), + }); + const remove = (index: number) => { + const next = markers.filter((_, markerIndex) => markerIndex !== index); + onChange({ ...value, reminder_markers: next.length > 0 ? next : undefined }); + }; + const error = getReminderMarkersError(value.reminder_markers); + return ( +
+

+ Delimiter pairs that wrap harness-injected reminder blocks, which are stripped before classification. Setting + any pair replaces the built-in pairs, so list every pair your harness emits. Matching is case-insensitive and + values are saved lowercased. +

+
+ {markers.map((marker, index) => ( +
+
+ + update(index, { open: event.target.value })} + /> +
+
+ + update(index, { close: event.target.value })} + /> +
+ +
+ ))} +
+ + {showValidationErrors && error &&

{error}

} +
+ ); +}; + +export default ReminderMarkers; diff --git a/ui/litellm-dashboard/src/components/add_model/ResponseFormatControls.tsx b/ui/litellm-dashboard/src/components/add_model/ResponseFormatControls.tsx index 68dd880a684..6e7b1489a4c 100644 --- a/ui/litellm-dashboard/src/components/add_model/ResponseFormatControls.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ResponseFormatControls.tsx @@ -18,6 +18,18 @@ const ResponseFormatControls: React.FC<{ Return the resolved underlying model name in responses instead of the autorouter alias. +
+ onChange({ ...value, max_tokens_from_tier_model: enabled })} + aria-label="Cap max_tokens at the tier model's output ceiling" + /> + Cap max_tokens at the tier model's output ceiling +
+ + Replace the caller's max_tokens with the routed tier model's output ceiling so one client value fits + every tier. Off forwards the caller's value unchanged. + ); diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx index fdff7ad001c..242893e2ae8 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx @@ -744,7 +744,7 @@ describe("AddAutoRouterTab", () => { ); }); - it("carries a context-window escalation opt-out through to the create payload", async () => { + it("starts context-window escalation disabled and carries an explicit opt-in to the create payload", async () => { const user = userEvent.setup(); vi.mocked(getMissingTiersError).mockReturnValue(null); @@ -754,14 +754,15 @@ describe("AddAutoRouterTab", () => { expandDetailedConfiguration(); await user.click(screen.getByText("Advanced: Context Window Escalation")); const toggle = await screen.findByRole("switch", { name: "Escalate oversized prompts to a tier that fits" }); - expect(toggle).toBeChecked(); + expect(toggle).not.toBeChecked(); + expect(screen.queryByLabelText("Window fit buffer")).not.toBeInTheDocument(); await user.click(toggle); await user.click(screen.getByRole("button", { name: /add auto router/i })); await waitFor(() => expect(handleAddAutoRouterSubmit).toHaveBeenCalled()); expect(vi.mocked(handleAddAutoRouterSubmit).mock.calls.at(-1)?.[0].complexity_router_config).toMatchObject({ - enable_context_window_escalation: false, + enable_context_window_escalation: true, }); }); @@ -774,6 +775,7 @@ describe("AddAutoRouterTab", () => { await user.type(screen.getByPlaceholderText(/smart_router/i), "ctx-buffer-router"); expandDetailedConfiguration(); await user.click(screen.getByText("Advanced: Context Window Escalation")); + await user.click(screen.getByRole("switch", { name: "Escalate oversized prompts to a tier that fits" })); const buffer = await screen.findByLabelText("Window fit buffer"); fireEvent.change(buffer, { target: { value: "1.5" } }); fireEvent.blur(buffer, { target: { value: "1.5" } }); @@ -783,7 +785,7 @@ describe("AddAutoRouterTab", () => { await waitFor(() => expect(handleAddAutoRouterSubmit).toHaveBeenCalled()); const config = vi.mocked(handleAddAutoRouterSubmit).mock.calls.at(-1)?.[0].complexity_router_config; expect(config).toMatchObject({ context_window_escalation_buffer: 1 }); - expect(config).not.toHaveProperty("enable_context_window_escalation"); + expect(config).toHaveProperty("enable_context_window_escalation", true); }); it("clearing the buffer removes it from the payload so the router tracks the backend default", async () => { @@ -795,6 +797,7 @@ describe("AddAutoRouterTab", () => { await user.type(screen.getByPlaceholderText(/smart_router/i), "ctx-clear-router"); expandDetailedConfiguration(); await user.click(screen.getByText("Advanced: Context Window Escalation")); + await user.click(screen.getByRole("switch", { name: "Escalate oversized prompts to a tier that fits" })); const buffer = await screen.findByLabelText("Window fit buffer"); fireEvent.change(buffer, { target: { value: "0.8" } }); fireEvent.blur(buffer, { target: { value: "0.8" } }); diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx index 84e44fee9c3..9be49edf08d 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx @@ -28,10 +28,6 @@ import ComplexityRouterConfig, { effectiveClassifierType, usesLlmClassifier, heuristicScoringRole, - DEFAULT_ADAPTIVE_WEIGHTS, - DEFAULT_SESSION_AFFINITY, - DEFAULT_DEPLOYMENT_AFFINITY, - DEFAULT_TIER_DISTANCE_PENALTY, } from "./ComplexityRouterConfig"; import { KeywordTierRule } from "./KeywordTierRules"; import { customDimensionsError } from "./custom_dimensions"; @@ -48,6 +44,8 @@ import { getKeywordTierRulesError, getClassifierModelError, getHeuristicV2SuccessThresholdError, + getReminderMarkersError, + getClassifierPluginTimeoutError, getClassifierReasoningEffortError, getMissingTiersError, getPlanModeTierError, @@ -55,6 +53,7 @@ import { getTierLabelsError, dryRunRejection, } from "./build_complexity_router_config"; +import { builderParamsFromValue } from "./complexity_router_builder_params"; import { activeTierName, activeTierRows, getCustomTierRowsError, resolveComplexityDefaultModel } from "./tier_rows"; import { tierRowLabel } from "./complexity_router_tiers"; import { buildAutoRouterTestTargets, AutoRouterTestTarget } from "./build_auto_router_test_targets"; @@ -152,6 +151,8 @@ export const getSubmitBlockedReason = ( getKeywordTierRulesError(keywordTierRules, activeTierRows(config)) ?? getClassifierModelError(config) ?? getHeuristicV2SuccessThresholdError(config.heuristic_v2_success_threshold) ?? + getReminderMarkersError(config.reminder_markers) ?? + getClassifierPluginTimeoutError(config.classifier_type, config.classifier_plugin_timeout_ms) ?? (heuristicScoringRole(config) === "decides" ? customDimensionsError(config.custom_dimensions) : null) ?? getClassifierReasoningEffortError(config, modelInfo) ?? getReferencedModelsError(referencedModelsParams, availability) @@ -399,56 +400,25 @@ const AddAutoRouterTab: React.FC = ({ ); const complexityRouterConfigParams: BuildComplexityRouterConfigParams = { - tiers: complexityRouterConfig.tiers, - enableNonReasoningTier: complexityRouterConfig.enable_non_reasoning_tier, - customTierSet: complexityRouterConfig.custom_tier_set, - defaultModel: complexityRouterConfig.default_model, - planModeMinTier: complexityRouterConfig.plan_mode_min_tier, - classificationPrompt: complexityRouterConfig.classification_prompt, - classificationExamples: complexityRouterConfig.classification_examples, - heuristicFirstMaxTier: complexityRouterConfig.heuristic_first_max_tier, - hybridBoundaryMargin: complexityRouterConfig.hybrid_boundary_margin, - classificationMode: complexityRouterConfig.classification_mode, - tierLabels: complexityRouterConfig.tier_labels, - classifierType: complexityRouterConfig.classifier_type, - jevClassifierConfig: complexityRouterConfig.jev_classifier_config, - heuristicV2SuccessThreshold: complexityRouterConfig.heuristic_v2_success_threshold, - capabilityClassifierConfig: complexityRouterConfig.capability_classifier_config, - llmV2Config: complexityRouterConfig.llm_v2_config, - classifierLlmConfig: complexityRouterConfig.classifier_llm_config, - classifierContextWindowSize: complexityRouterConfig.classifier_context_window_size, - classifierContextBudgetChars: complexityRouterConfig.classifier_context_budget_chars, - classifierContextPerTurnChars: complexityRouterConfig.classifier_context_per_turn_chars, - classifierContextIncludeAssistantTurns: complexityRouterConfig.classifier_context_include_assistant_turns, - classifierFallback: complexityRouterConfig.classifier_fallback, - sessionAffinity: complexityRouterConfig.session_affinity ?? DEFAULT_SESSION_AFFINITY, - modalityRouting: complexityRouterConfig.modality_routing ?? false, - modalityPinOverride: complexityRouterConfig.modality_pin_override ?? false, - deploymentAffinity: complexityRouterConfig.deployment_affinity ?? DEFAULT_DEPLOYMENT_AFFINITY, + ...builderParamsFromValue(complexityRouterConfig), customTechnicalKeywords, keywordTierRules, semanticMatchingEnabled, embeddingModel, matchThreshold, escalationKeywords, - stallEscalationEnabled: complexityRouterConfig.stall_escalation_enabled, - stallEscalationWindow: complexityRouterConfig.stall_escalation_window, - stallEscalationRepeatThreshold: complexityRouterConfig.stall_escalation_repeat_threshold, - adaptive: complexityRouterConfig.adaptive ?? false, - adaptiveWeights: complexityRouterConfig.adaptive_weights ?? DEFAULT_ADAPTIVE_WEIGHTS, - tierDistancePenalty: complexityRouterConfig.tier_distance_penalty ?? DEFAULT_TIER_DISTANCE_PENALTY, - adaptiveEligible: complexityRouterConfig.adaptive_eligible ?? "all", - returnRawModelName: complexityRouterConfig.return_raw_model_name ?? false, - tierModelParams: complexityRouterConfig.tier_model_params, - tierBoundaries: complexityRouterConfig.tier_boundaries, - tokenThresholds: complexityRouterConfig.token_thresholds, - dimensionWeights: complexityRouterConfig.dimension_weights, - customDimensions: complexityRouterConfig.custom_dimensions, - reasoningOverrideMinScore: complexityRouterConfig.reasoning_override_min_score, - enableContextWindowEscalation: complexityRouterConfig.enable_context_window_escalation, - contextWindowEscalationBuffer: complexityRouterConfig.context_window_escalation_buffer, - sessionAffinityTtlSeconds: complexityRouterConfig.session_affinity_ttl_seconds, }; + const jevRequestParams = + effectiveClassifierType(complexityRouterConfig) === "jev" + ? { + prompt: JEV_CONNECTION_TEST_PROMPT, + config: buildComplexityRouterConfig(complexityRouterConfigParams), + defaultModel: resolveComplexityDefaultModel(complexityRouterConfig, complexityRouterConfig.default_model), + routerName: watchedName, + teamId: requiresTeamScope ? watchedTeamId ?? undefined : undefined, + } + : undefined; + const jevRequest = jevRequestParams ? buildAutoRouterRoutingTestRequest(jevRequestParams) : undefined; const submitRecommendedRouter = async (name: string) => { // The one answer the submit button reads, so a disabled button and a refused submit cannot @@ -857,20 +827,7 @@ const AddAutoRouterTab: React.FC = ({ testId={connectionTestId} accessToken={accessToken} targets={testTargets} - jevRequest={ - effectiveClassifierType(complexityRouterConfig) === "jev" - ? buildAutoRouterRoutingTestRequest({ - prompt: JEV_CONNECTION_TEST_PROMPT, - config: buildComplexityRouterConfig(complexityRouterConfigParams), - defaultModel: resolveComplexityDefaultModel( - complexityRouterConfig, - complexityRouterConfig.default_model, - ), - routerName: watchedName, - teamId: requiresTeamScope ? watchedTeamId ?? undefined : undefined, - }) - : undefined - } + jevRequest={jevRequest} onTestComplete={() => setIsTestingConnection(false)} />
diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts index 054aba7a6aa..6d458d61797 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts @@ -6,6 +6,8 @@ import { getKeywordTierRulesError, getClassifierModelError, getHeuristicV2SuccessThresholdError, + getReminderMarkersError, + getClassifierPluginTimeoutError, getClassifierReasoningEffortError, getMissingTiersError, hydrateCustomTierSet, @@ -166,7 +168,7 @@ describe("buildComplexityRouterConfig", () => { }); it.each(["capability", "llm_v2", "heuristic"] as const)( - "disables the removed overrides only for forecast creates: %s", + "preserves explicit context-window opt-in beside forecast restrictions: %s", (classifierType) => { const forecast = classifierType !== "heuristic"; const params = { @@ -178,14 +180,10 @@ describe("buildComplexityRouterConfig", () => { }; const config = buildComplexityRouterConfig(params); expect(config.adaptive).toBe(!forecast); - expect(config.enable_context_window_escalation).toBe(!forecast); + expect(config.enable_context_window_escalation).toBe(true); + expect(config.context_window_escalation_buffer).toBe(0.9); expect(config.escalation_keywords).toEqual(forecast ? [] : ["LITELLM ESCALATE"]); - for (const key of [ - "adaptive_weights", - "adaptive_eligible", - "tier_distance_penalty", - "context_window_escalation_buffer", - ]) { + for (const key of ["adaptive_weights", "adaptive_eligible", "tier_distance_penalty"]) { expect(Object.hasOwn(config, key)).toBe(!forecast); } if (forecast) { @@ -224,13 +222,14 @@ describe("buildComplexityRouterConfig", () => { expect(config).toEqual(expected); }); - it("carries an explicit context-window escalation opt-out and buffer, false included", () => { + it.each([undefined, false, true])("preserves the context-window escalation setting: %s", (enabled) => { const config = buildComplexityRouterConfig({ ...baseParams, - enableContextWindowEscalation: false, + enableContextWindowEscalation: enabled, contextWindowEscalationBuffer: 0.9, }); - expect(config.enable_context_window_escalation).toBe(false); + expect(config.enable_context_window_escalation).toBe(enabled); + expect(Object.hasOwn(config, "enable_context_window_escalation")).toBe(enabled !== undefined); expect(config.context_window_escalation_buffer).toBe(0.9); }); @@ -1482,3 +1481,75 @@ describe("classifier vision wire payload", () => { expect(payload.classifier_llm_config).not.toHaveProperty("vision"); }); }); + +describe("advanced complexity router fields", () => { + it("normalizes lists, reminder markers, and explicit false values", () => { + const payload = buildComplexityRouterConfig({ + ...baseParams, + codeKeywords: [" async ", " "], + reasoningKeywords: ["prove"], + technicalKeywords: ["api"], + simpleKeywords: ["hello"], + planModePatterns: [" plan "], + routeHousekeepingToCheapestTier: false, + housekeepingPatterns: [" title "], + reminderMarkers: [{ open: " ", close: " " }], + maxTokensFromTierModel: false, + classifierType: "custom", + classifierPluginTimeoutMs: 3000, + }); + expect(payload).toMatchObject({ + code_keywords: ["async"], + reasoning_keywords: ["prove"], + technical_keywords: ["api"], + simple_keywords: ["hello"], + plan_mode_patterns: ["plan"], + route_housekeeping_to_cheapest_tier: false, + housekeeping_patterns: ["title"], + reminder_markers: [{ open: "", close: "" }], + max_tokens_from_tier_model: false, + classifier_plugin_timeout_ms: 3000, + }); + }); + + it("omits defaults, empty lists, and timeout values for non-custom classifiers", () => { + const payload = buildComplexityRouterConfig({ + ...baseParams, + codeKeywords: [" ", ""], + reminderMarkers: [], + routeHousekeepingToCheapestTier: true, + maxTokensFromTierModel: true, + classifierPluginTimeoutMs: 3000, + }); + expect(payload).not.toHaveProperty("code_keywords"); + expect(payload).not.toHaveProperty("reminder_markers"); + expect(payload).not.toHaveProperty("route_housekeeping_to_cheapest_tier"); + expect(payload).not.toHaveProperty("max_tokens_from_tier_model"); + expect(payload).not.toHaveProperty("classifier_plugin_timeout_ms"); + }); + + it.each([ + "code_keywords", + "reasoning_keywords", + "technical_keywords", + "simple_keywords", + "plan_mode_patterns", + "route_housekeeping_to_cheapest_tier", + "housekeeping_patterns", + "reminder_markers", + "max_tokens_from_tier_model", + "classifier_plugin_timeout_ms", + ])("omits unset advanced field %s", (key) => { + const payload = buildComplexityRouterConfig(baseParams); + expect(payload).not.toHaveProperty(key); + }); + + it("validates marker pairs and custom classifier timeout", () => { + expect(getReminderMarkersError([{ open: " ", close: " " }])).toContain("different"); + expect(getReminderMarkersError([{ open: "", close: "" }])).toContain("needs both"); + expect(getReminderMarkersError([{ open: "", close: "" }])).toBeNull(); + expect(getClassifierPluginTimeoutError("custom", 0)).toContain("whole number"); + expect(getClassifierPluginTimeoutError("custom", 3000)).toBeNull(); + expect(getClassifierPluginTimeoutError("heuristic", 0)).toBeNull(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts index 15ae2b4c0b0..9b4f1980f62 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts @@ -54,6 +54,10 @@ import { export type ClassifierVisionConfig = { enabled?: boolean; max_images?: number }; export type ClassifierLLMConfigWire = ClassifierLLMConfig & { vision?: ClassifierVisionConfig }; +export interface ReminderMarkerPair { + open: string; + close: string; +} /** * Drop an empty system_prompt so the payload carries an override only when there is one. The @@ -181,6 +185,16 @@ export interface StoredComplexityRouterConfig { stall_escalation_enabled?: unknown; stall_escalation_window?: unknown; stall_escalation_repeat_threshold?: unknown; + code_keywords?: unknown; + reasoning_keywords?: unknown; + technical_keywords?: unknown; + simple_keywords?: unknown; + plan_mode_patterns?: unknown; + route_housekeeping_to_cheapest_tier?: unknown; + housekeeping_patterns?: unknown; + reminder_markers?: unknown; + max_tokens_from_tier_model?: unknown; + classifier_plugin_timeout_ms?: unknown; } export interface BuildComplexityRouterConfigParams { @@ -233,6 +247,16 @@ export interface BuildComplexityRouterConfigParams { enableContextWindowEscalation?: boolean; contextWindowEscalationBuffer?: number; sessionAffinityTtlSeconds?: number; + codeKeywords?: string[]; + reasoningKeywords?: string[]; + technicalKeywords?: string[]; + simpleKeywords?: string[]; + planModePatterns?: string[]; + routeHousekeepingToCheapestTier?: boolean; + housekeepingPatterns?: string[]; + reminderMarkers?: ReminderMarkerPair[]; + maxTokensFromTierModel?: boolean; + classifierPluginTimeoutMs?: number; } /** @@ -302,6 +326,16 @@ export interface ComplexityRouterConfigPayload { enable_context_window_escalation?: boolean; context_window_escalation_buffer?: number; tier_model_configs?: Record; + code_keywords?: string[]; + reasoning_keywords?: string[]; + technical_keywords?: string[]; + simple_keywords?: string[]; + plan_mode_patterns?: string[]; + route_housekeeping_to_cheapest_tier?: boolean; + housekeeping_patterns?: string[]; + reminder_markers?: ReminderMarkerPair[]; + max_tokens_from_tier_model?: boolean; + classifier_plugin_timeout_ms?: number; } export const serializeTierLabels = (tierLabels: ComplexityTierLabels | undefined): ComplexityTierLabels | undefined => { @@ -376,6 +410,26 @@ export const getHeuristicV2SuccessThresholdError = (threshold: number | undefine return validProbability ? null : "Success threshold must be a number between 0 and 1"; }; +export const getReminderMarkersError = (pairs: ReminderMarkerPair[] | undefined): string | null => { + for (const [index, pair] of (pairs ?? []).entries()) { + const open = pair.open.trim().toLowerCase(); + const close = pair.close.trim().toLowerCase(); + if (!open || !close) return `Reminder marker pair ${index + 1} needs both an opening and a closing delimiter`; + if (open === close) return `Reminder marker pair ${index + 1} must use different opening and closing delimiters`; + } + return null; +}; + +export const getClassifierPluginTimeoutError = ( + classifierType: ClassifierType, + timeoutMs: number | undefined, +): string | null => { + if (classifierType !== "custom" || timeoutMs === undefined) return null; + return Number.isInteger(timeoutMs) && timeoutMs > 0 + ? null + : "Classifier plugin timeout must be a whole number of milliseconds greater than 0"; +}; + export const getClassifierModelError = ( config: Pick< ComplexityRouterConfigValue, @@ -640,6 +694,16 @@ export const buildComplexityRouterConfig = ({ enableContextWindowEscalation, contextWindowEscalationBuffer, sessionAffinityTtlSeconds, + codeKeywords, + reasoningKeywords, + technicalKeywords, + simpleKeywords, + planModePatterns, + routeHousekeepingToCheapestTier, + housekeepingPatterns, + reminderMarkers, + maxTokensFromTierModel, + classifierPluginTimeoutMs, }: BuildComplexityRouterConfigParams): ComplexityRouterConfigPayload => { const serializedTierModelConfigs = customTierSet ? serializeTierModelConfigs( @@ -672,6 +736,29 @@ export const buildComplexityRouterConfig = ({ }; const effectiveType = effectiveClassifierType({ custom_tier_set: customTierSet, classifier_type: classifierType }); const forecast = isForecastClassifier(effectiveType); + const preserveContextWindowBuffer = !forecast || enableContextWindowEscalation === true; + const cleanList = (items: string[] | undefined): string[] | undefined => { + const cleaned = (items ?? []).map((item) => item.trim()).filter(Boolean); + return cleaned.length > 0 ? cleaned : undefined; + }; + const cleanedReminderMarkers = reminderMarkers?.map(({ open, close }) => ({ + open: open.trim().toLowerCase(), + close: close.trim().toLowerCase(), + })); + const cleanedListValues = { + code_keywords: cleanList(codeKeywords), + reasoning_keywords: cleanList(reasoningKeywords), + technical_keywords: cleanList(technicalKeywords), + simple_keywords: cleanList(simpleKeywords), + plan_mode_patterns: cleanList(planModePatterns), + housekeeping_patterns: cleanList(housekeepingPatterns), + }; + const cleanedLists = Object.fromEntries(Object.entries(cleanedListValues).filter(([, list]) => list !== undefined)); + const hasValidCustomClassifierTimeout = + classifierType === "custom" && + classifierPluginTimeoutMs !== undefined && + Number.isInteger(classifierPluginTimeoutMs) && + classifierPluginTimeoutMs > 0; const supportsOpeningPrompt = !customTierSet && !forecast && usesLlmClassifier(effectiveType); const payload: ComplexityRouterConfigPayload = { @@ -729,17 +816,21 @@ export const buildComplexityRouterConfig = ({ adaptive_eligible: adaptiveEligible, }), ...(returnRawModelName && { return_raw_model_name: true }), - // Omission enables the backend default, so hidden forecast controls need an explicit opt-out. ...((forecast || enableContextWindowEscalation !== undefined) && { - enable_context_window_escalation: forecast ? false : enableContextWindowEscalation, + enable_context_window_escalation: enableContextWindowEscalation ?? false, }), - ...(!forecast && + ...(preserveContextWindowBuffer && contextWindowEscalationBuffer !== undefined && { context_window_escalation_buffer: contextWindowEscalationBuffer, }), ...(sessionAffinityTtlSeconds !== undefined && { session_affinity_ttl_seconds: sessionAffinityTtlSeconds, }), + ...cleanedLists, + ...(routeHousekeepingToCheapestTier === false && { route_housekeeping_to_cheapest_tier: false }), + ...(cleanedReminderMarkers && cleanedReminderMarkers.length > 0 && { reminder_markers: cleanedReminderMarkers }), + ...(maxTokensFromTierModel === false && { max_tokens_from_tier_model: false }), + ...(hasValidCustomClassifierTimeout && { classifier_plugin_timeout_ms: classifierPluginTimeoutMs }), ...scorerKnobs, }; if (!customTierSet) return payload; diff --git a/ui/litellm-dashboard/src/components/add_model/classifier_types.ts b/ui/litellm-dashboard/src/components/add_model/classifier_types.ts index ec88166ed2e..aa9d5619052 100644 --- a/ui/litellm-dashboard/src/components/add_model/classifier_types.ts +++ b/ui/litellm-dashboard/src/components/add_model/classifier_types.ts @@ -6,7 +6,8 @@ export type ClassifierType = | "heuristic_first" | "hybrid" | "capability" - | "llm_v2"; + | "llm_v2" + | "custom"; export const usesLlmClassifier = (classifierType: ClassifierType): boolean => (["llm", "heuristic_first", "hybrid", "capability", "llm_v2"] as const).some((type) => type === classifierType); diff --git a/ui/litellm-dashboard/src/components/add_model/complexity_router_builder_params.ts b/ui/litellm-dashboard/src/components/add_model/complexity_router_builder_params.ts new file mode 100644 index 00000000000..124a85ce9a3 --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/complexity_router_builder_params.ts @@ -0,0 +1,74 @@ +import type { BuildComplexityRouterConfigParams } from "./build_complexity_router_config"; +import type { ComplexityRouterConfigValue } from "./ComplexityRouterConfig"; +import { + DEFAULT_ADAPTIVE_WEIGHTS, + DEFAULT_DEPLOYMENT_AFFINITY, + DEFAULT_SESSION_AFFINITY, + DEFAULT_TIER_DISTANCE_PENALTY, +} from "./ComplexityRouterConfig"; + +export const builderParamsFromValue = ( + value: ComplexityRouterConfigValue, +): Omit< + BuildComplexityRouterConfigParams, + | "customTechnicalKeywords" + | "keywordTierRules" + | "semanticMatchingEnabled" + | "embeddingModel" + | "matchThreshold" + | "escalationKeywords" +> => ({ + tiers: value.tiers, + enableNonReasoningTier: value.enable_non_reasoning_tier, + customTierSet: value.custom_tier_set, + defaultModel: value.default_model, + planModeMinTier: value.plan_mode_min_tier, + classificationPrompt: value.classification_prompt, + classificationExamples: value.classification_examples, + heuristicFirstMaxTier: value.heuristic_first_max_tier, + hybridBoundaryMargin: value.hybrid_boundary_margin, + classificationMode: value.classification_mode, + tierLabels: value.tier_labels, + classifierType: value.classifier_type, + jevClassifierConfig: value.jev_classifier_config, + heuristicV2SuccessThreshold: value.heuristic_v2_success_threshold, + capabilityClassifierConfig: value.capability_classifier_config, + llmV2Config: value.llm_v2_config, + classifierLlmConfig: value.classifier_llm_config, + classifierContextWindowSize: value.classifier_context_window_size, + classifierContextBudgetChars: value.classifier_context_budget_chars, + classifierContextPerTurnChars: value.classifier_context_per_turn_chars, + classifierContextIncludeAssistantTurns: value.classifier_context_include_assistant_turns, + classifierFallback: value.classifier_fallback, + sessionAffinity: value.session_affinity ?? DEFAULT_SESSION_AFFINITY, + sessionAffinityTtlSeconds: value.session_affinity_ttl_seconds, + modalityRouting: value.modality_routing ?? false, + modalityPinOverride: value.modality_pin_override ?? false, + deploymentAffinity: value.deployment_affinity ?? DEFAULT_DEPLOYMENT_AFFINITY, + adaptive: value.adaptive ?? false, + adaptiveWeights: value.adaptive_weights ?? DEFAULT_ADAPTIVE_WEIGHTS, + tierDistancePenalty: value.tier_distance_penalty ?? DEFAULT_TIER_DISTANCE_PENALTY, + adaptiveEligible: value.adaptive_eligible ?? "all", + returnRawModelName: value.return_raw_model_name ?? false, + tierBoundaries: value.tier_boundaries, + tokenThresholds: value.token_thresholds, + dimensionWeights: value.dimension_weights, + customDimensions: value.custom_dimensions, + reasoningOverrideMinScore: value.reasoning_override_min_score, + tierModelParams: value.tier_model_params, + enableContextWindowEscalation: value.enable_context_window_escalation, + contextWindowEscalationBuffer: value.context_window_escalation_buffer, + stallEscalationEnabled: value.stall_escalation_enabled, + stallEscalationWindow: value.stall_escalation_window, + stallEscalationRepeatThreshold: value.stall_escalation_repeat_threshold, + codeKeywords: value.code_keywords, + reasoningKeywords: value.reasoning_keywords, + technicalKeywords: value.technical_keywords, + simpleKeywords: value.simple_keywords, + planModePatterns: value.plan_mode_patterns, + routeHousekeepingToCheapestTier: value.route_housekeeping_to_cheapest_tier, + housekeepingPatterns: value.housekeeping_patterns, + reminderMarkers: value.reminder_markers, + maxTokensFromTierModel: value.max_tokens_from_tier_model, + classifierPluginTimeoutMs: value.classifier_plugin_timeout_ms, +}); diff --git a/ui/litellm-dashboard/src/components/add_model/handle_add_model_submit.test.tsx b/ui/litellm-dashboard/src/components/add_model/handle_add_model_submit.test.tsx index 9d792480c9f..923cf2aa0e3 100644 --- a/ui/litellm-dashboard/src/components/add_model/handle_add_model_submit.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/handle_add_model_submit.test.tsx @@ -101,4 +101,23 @@ describe("prepareModelAddRequest", () => { expect(deployment.litellmParamsObj.litellm_credential_name).toBe("from-json"); expect(deployment.litellmParamsObj.timeout).toBe(5); }); + + it.each([ + ["OpenAI", "openai/*"], + ["Azure_AI_Studio", "azure_ai/*"], + ["Petals", "petals/*"], + ])("composes wildcard names for the all-model selection", async (custom_llm_provider, wildcardModel) => { + const formValues = { + model_mappings: [], + model: "all-wildcard", + custom_llm_provider, + }; + + const deployments = await prepareModelAddRequest({ ...formValues }, "token", null); + + expect(deployments).toHaveLength(1); + const [deployment] = deployments!; + expect(deployment.modelName).toBe(wildcardModel); + expect(deployment.litellmParamsObj.model).toBe(wildcardModel); + }); }); diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts b/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts index 2450f7bce27..162414b7814 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts +++ b/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts @@ -189,14 +189,10 @@ describe("buildUpdatedComplexityRouterConfig keyword matching", () => { const saved = buildUpdatedComplexityRouterConfig(stored, value, undefined, keywordState); const forecast = classifier_type !== "heuristic"; expect(saved.adaptive).toBe(!forecast); - expect(saved.enable_context_window_escalation).toBe(!forecast); + expect(saved.enable_context_window_escalation).toBe(true); + expect(saved.context_window_escalation_buffer).toBe(0.9); expect(saved.escalation_keywords).toEqual(forecast ? [] : stored.escalation_keywords); - for (const key of [ - "adaptive_weights", - "adaptive_eligible", - "tier_distance_penalty", - "context_window_escalation_buffer", - ]) { + for (const key of ["adaptive_weights", "adaptive_eligible", "tier_distance_penalty"]) { expect(Object.hasOwn(saved, key)).toBe(!forecast); } expect(saved.keyword_tier_rules).toEqual(STORED.keyword_tier_rules); @@ -208,6 +204,19 @@ describe("buildUpdatedComplexityRouterConfig keyword matching", () => { }, ); + it.each([undefined, false, true])("preserves stored context-window escalation on save: %s", (enabled) => { + const stored = { + ...STORED, + ...(enabled !== undefined && { enable_context_window_escalation: enabled }), + }; + const value = hydrateComplexityRouterConfig(stored, undefined); + const saved = buildUpdatedComplexityRouterConfig(stored, value, undefined, hydratedState); + const serialized: typeof saved = JSON.parse(JSON.stringify(saved)); + expect(value.enable_context_window_escalation).toBe(enabled); + expect(serialized.enable_context_window_escalation).toBe(enabled); + expect(Object.hasOwn(serialized, "enable_context_window_escalation")).toBe(enabled !== undefined); + }); + it("round-trips an untouched edit without changing any keyword-matching value", () => { // Opening the modal hydrates state from STORED; saving with nothing changed must be a // no-op. These keys are now MANAGED, so a hydration bug silently wipes them. @@ -854,6 +863,16 @@ describe("managed keys survive an untouched open-and-save", () => { reasoning_override_min_score: 0.3, enable_context_window_escalation: false, context_window_escalation_buffer: 0.9, + code_keywords: ["async", "await"], + reasoning_keywords: ["prove"], + technical_keywords: ["api"], + simple_keywords: ["hello"], + plan_mode_patterns: ["plan now"], + route_housekeeping_to_cheapest_tier: false, + housekeeping_patterns: ["conversation title"], + reminder_markers: [{ open: "", close: "" }], + max_tokens_from_tier_model: false, + classifier_plugin_timeout_ms: 3000, }; // tier_definitions and fallback_tier cannot sit beside heuristic_first, which this fixture uses, @@ -864,6 +883,7 @@ describe("managed keys survive an untouched open-and-save", () => { "fallback_tier", "hybrid_boundary_margin", "jev_classifier_config", + "classifier_plugin_timeout_ms", ]); // The stall keys are rejected beside the session pinning and user-turn classification this @@ -894,6 +914,12 @@ describe("managed keys survive an untouched open-and-save", () => { expect(dropped).toEqual([]); }); + it("keeps the custom classifier plugin timeout through an untouched save", () => { + const stored = { ...STORED_ALL_MANAGED, classifier_type: "custom", classifier_plugin_timeout_ms: 3000 }; + const hydrated = hydrateComplexityRouterConfig(stored, undefined); + expect(buildUpdatedComplexityRouterConfig(stored, hydrated).classifier_plugin_timeout_ms).toBe(3000); + }); + it("carries an enabled non-reasoning tier and its models through their own round trip", () => { // `tiers` is rewritten wholesale on save, so this is the regression that matters: opening an // enabled router and saving an unrelated edit must not delete the tier or its pool. diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.integration.test.tsx b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.integration.test.tsx index 34db61483cf..aa6cf92ceb9 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.integration.test.tsx @@ -347,6 +347,77 @@ describe("EditAutoRouterModal keyword matching", () => { }); }); +describe("EditAutoRouterModal advanced field round trips", () => { + const storedAdvancedConfig = { + ...STORED_CONFIG, + route_housekeeping_to_cheapest_tier: false, + housekeeping_patterns: ["conversation title"], + reminder_markers: [{ open: "", close: "" }], + max_tokens_from_tier_model: false, + }; + + const renderAdvancedModal = (props: Partial> = {}) => + renderModal({ + modelData: { + ...MODEL_DATA, + litellm_params: { ...MODEL_DATA.litellm_params, complexity_router_config: storedAdvancedConfig }, + }, + ...props, + }); + + beforeEach(() => { + modelPatchUpdateCall.mockClear(); + }); + + it("hydrates housekeeping and reminder fields, then omits the default max-token value after editing", async () => { + const user = userEvent.setup(); + renderAdvancedModal(); + + await user.click(await screen.findByText("Advanced: Housekeeping Routing")); + expect(screen.getByRole("switch", { name: "Route housekeeping calls to the cheapest tier" })).not.toBeChecked(); + expect(screen.getByRole("combobox", { name: "e.g., conversation title" })).toHaveValue(""); + + await user.click(screen.getByText("Advanced: Reminder Markers")); + expect(screen.getByLabelText("Opening delimiter")).toHaveValue(""); + expect(screen.getByLabelText("Closing delimiter")).toHaveValue(""); + + await user.click(screen.getByText("Advanced: Response Format")); + const maxTokensSwitch = screen.getByRole("switch", { name: "Cap max_tokens at the tier model's output ceiling" }); + await user.click(maxTokensSwitch); + await user.click(screen.getByRole("button", { name: /save changes/i })); + await waitFor(() => expect(modelPatchUpdateCall).toHaveBeenCalledOnce()); + + expect(savedConfig()).not.toHaveProperty("max_tokens_from_tier_model"); + expect(savedConfig()).toMatchObject({ + route_housekeeping_to_cheapest_tier: false, + housekeeping_patterns: ["conversation title"], + reminder_markers: [{ open: "", close: "" }], + }); + }); + + it("does not PATCH when the edit is cancelled", async () => { + const user = userEvent.setup(); + const onCancel = vi.fn(); + renderAdvancedModal({ onCancel }); + await user.click(screen.getByRole("button", { name: /cancel/i })); + expect(onCancel).toHaveBeenCalledOnce(); + expect(modelPatchUpdateCall).not.toHaveBeenCalled(); + }); + + it("preserves all stored advanced fields through an untouched save", async () => { + const user = userEvent.setup(); + renderAdvancedModal(); + await user.click(screen.getByRole("button", { name: /save changes/i })); + await waitFor(() => expect(modelPatchUpdateCall).toHaveBeenCalledOnce()); + expect(savedConfig()).toMatchObject({ + route_housekeeping_to_cheapest_tier: false, + housekeeping_patterns: ["conversation title"], + reminder_markers: [{ open: "", close: "" }], + max_tokens_from_tier_model: false, + }); + }); +}); + describe("EditAutoRouterModal classifier context window", () => { beforeEach(() => { modelPatchUpdateCall.mockClear(); diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx index c88bbb101f7..0dafa9b330a 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx +++ b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx @@ -1,14 +1,7 @@ import AutoRouterClassifierTabs from "../add_model/AutoRouterClassifierTabs"; import { usesClassifierContext } from "../add_model/classifier_types"; -import { defaultJevClassifierConfig, jevClassifierConfigSchema } from "../add_model/jev_classifier_config"; -import type { StoredComplexityRouterConfig } from "../add_model/build_complexity_router_config"; export type { StoredComplexityRouterConfig } from "../add_model/build_complexity_router_config"; -import { - getForecastConfigError, - isForecastClassifier, - capabilitySettingsSchema, - fuseSettingsSchema, -} from "../add_model/forecast_classifier_config"; +import { getForecastConfigError, isForecastClassifier } from "../add_model/forecast_classifier_config"; import React, { useEffect, useMemo, useState } from "react"; import { complexityRouterSchema, @@ -30,13 +23,10 @@ import ModelChoiceCombobox, { type ModelChoice } from "../add_model/ModelChoiceC import { modelAvailableCall, modelPatchUpdateCall, validateAutoRouterConfig } from "../networking"; import { fetchAutoRouterModels, fetchAvailableModels, ModelGroup } from "@/components/llm_calls/fetch_models"; import RouterConfigBuilder, { type RouterConfig, serializeRouterConfig } from "../add_model/RouterConfigBuilder"; -import { hydrateTierModelParams } from "../add_model/complexity_router_tiers"; import { - type ActiveTierSet, CUSTOM_TIER_OMITTED_KEYS, activeTierRows, getCustomTierRowsError, - tierParamsByRowId, resolveComplexityDefaultModel, } from "../add_model/tier_rows"; import { isComplexityRouter } from "../add_model/auto_router_strategies"; @@ -45,16 +35,14 @@ import { buildComplexityRouterConfig, getClassifierModelError, getHeuristicV2SuccessThresholdError, + getReminderMarkersError, + getClassifierPluginTimeoutError, getClassifierReasoningEffortError, getKeywordTierRulesError, getMissingTiersError, getSemanticConfigError, getPlanModeTierError, getTierLabelsError, - hydrateBuiltInTiers, - hydrateCustomTierSet, - hydratePlanModeMinTier, - hydrateTierLabels, dryRunRejection, } from "../add_model/build_complexity_router_config"; import { KeywordTierRule } from "../add_model/KeywordTierRules"; @@ -66,22 +54,14 @@ import { hydrateAutoRouterCompression, } from "../add_model/buildAutoRouterCompression"; import { hydrateKeywordTierRules } from "../add_model/complexity_router_keywords"; -import { customDimensionsError, hydrateCustomDimensions } from "../add_model/custom_dimensions"; -import { - hydrateDimensionWeights, - hydrateReasoningOverrideMinScore, - hydrateTierBoundaries, - hydrateTokenThresholds, -} from "../add_model/heuristic_scoring_knobs"; +import { customDimensionsError } from "../add_model/custom_dimensions"; import ComplexityRouterConfig, { ComplexityRouterConfigValue, effectiveClassifierType, heuristicScoringRole, - DEFAULT_ADAPTIVE_WEIGHTS, - DEFAULT_SESSION_AFFINITY, - DEFAULT_DEPLOYMENT_AFFINITY, - DEFAULT_TIER_DISTANCE_PENALTY, } from "../add_model/ComplexityRouterConfig"; +import { builderParamsFromValue } from "../add_model/complexity_router_builder_params"; +import { hydrateComplexityRouterConfig, hydratePinnedDefaultModel } from "./hydrate_complexity_router_config"; import { Dialog, DialogContent, @@ -104,124 +84,7 @@ interface EditAutoRouterModalProps { // Keys this modal rewrites from its own form state on save. Anything absent from this set is // carried through untouched from the stored config, so a key only belongs here once the modal // actually renders a control that can set it. - -/** - * The stored complexity_router_config as form state. Every key in MANAGED_COMPLEXITY_ROUTER_KEYS is - * rewritten from this state on save, so a key missing here is silently dropped from the saved config. - */ -export const hydrateComplexityRouterConfig = ( - parsedConfig: StoredComplexityRouterConfig, - complexityRouterDefaultModel: string | null | undefined, -): ComplexityRouterConfigValue => { - const builtIn = hydrateBuiltInTiers(parsedConfig.tiers, parsedConfig.enable_non_reasoning_tier); - const { tiers: hydratedTiers, enable_non_reasoning_tier } = builtIn; - const custom_tier_set = hydrateCustomTierSet(parsedConfig); - const activeTiers = { ...builtIn, custom_tier_set }; - - return { - tiers: hydratedTiers, - enable_non_reasoning_tier, - custom_tier_set, - tier_model_params: tierParamsByRowId( - hydrateTierModelParams(parsedConfig.tiers, parsedConfig.tier_model_configs), - activeTierRows(activeTiers), - ), - default_model: hydratePinnedDefaultModel(parsedConfig.default_model, complexityRouterDefaultModel, activeTiers), - plan_mode_min_tier: hydratePlanModeMinTier(parsedConfig.plan_mode_min_tier, custom_tier_set), - tier_labels: hydrateTierLabels(parsedConfig.tier_labels), - classifier_type: parsedConfig.classifier_type || "heuristic", - heuristic_v2_success_threshold: - typeof parsedConfig.heuristic_v2_success_threshold === "number" - ? parsedConfig.heuristic_v2_success_threshold - : undefined, - capability_classifier_config: capabilitySettingsSchema.safeParse(parsedConfig.capability_classifier_config).data, - llm_v2_config: fuseSettingsSchema.safeParse(parsedConfig.llm_v2_config).data, - classifier_llm_config: parsedConfig.classifier_type === "jev" ? undefined : parsedConfig.classifier_llm_config, - jev_classifier_config: - parsedConfig.classifier_type === "jev" - ? jevClassifierConfigSchema.safeParse(parsedConfig.jev_classifier_config ?? {}).data ?? - defaultJevClassifierConfig() - : undefined, - classifier_context_window_size: - typeof parsedConfig.classifier_context_window_size === "number" - ? parsedConfig.classifier_context_window_size - : undefined, - classifier_context_budget_chars: - typeof parsedConfig.classifier_context_budget_chars === "number" - ? parsedConfig.classifier_context_budget_chars - : undefined, - classifier_context_per_turn_chars: - typeof parsedConfig.classifier_context_per_turn_chars === "number" - ? parsedConfig.classifier_context_per_turn_chars - : undefined, - classifier_context_include_assistant_turns: - typeof parsedConfig.classifier_context_include_assistant_turns === "boolean" - ? parsedConfig.classifier_context_include_assistant_turns - : undefined, - classifier_fallback: - parsedConfig.classifier_fallback === "default_model" || parsedConfig.classifier_fallback === "heuristic" - ? parsedConfig.classifier_fallback - : undefined, - classification_prompt: - typeof parsedConfig.classification_prompt === "string" && parsedConfig.classification_prompt.trim() !== "" - ? parsedConfig.classification_prompt - : undefined, - classification_examples: - typeof parsedConfig.classification_examples === "string" && parsedConfig.classification_examples.trim() !== "" - ? parsedConfig.classification_examples - : undefined, - heuristic_first_max_tier: - typeof parsedConfig.heuristic_first_max_tier === "string" && parsedConfig.heuristic_first_max_tier.trim() !== "" - ? parsedConfig.heuristic_first_max_tier - : undefined, - hybrid_boundary_margin: - typeof parsedConfig.hybrid_boundary_margin === "number" ? parsedConfig.hybrid_boundary_margin : undefined, - classification_mode: - parsedConfig.classification_mode === "user_turn" || parsedConfig.classification_mode === "every_request" - ? parsedConfig.classification_mode - : undefined, - tier_boundaries: hydrateTierBoundaries(parsedConfig.tier_boundaries), - token_thresholds: hydrateTokenThresholds(parsedConfig.token_thresholds), - dimension_weights: hydrateDimensionWeights(parsedConfig.dimension_weights), - custom_dimensions: hydrateCustomDimensions(parsedConfig.custom_dimensions), - reasoning_override_min_score: hydrateReasoningOverrideMinScore(parsedConfig.reasoning_override_min_score), - session_affinity: - typeof parsedConfig.session_affinity === "boolean" ? parsedConfig.session_affinity : DEFAULT_SESSION_AFFINITY, - session_affinity_ttl_seconds: - typeof parsedConfig.session_affinity_ttl_seconds === "number" && - Number.isFinite(parsedConfig.session_affinity_ttl_seconds) - ? parsedConfig.session_affinity_ttl_seconds - : undefined, - modality_routing: typeof parsedConfig.modality_routing === "boolean" ? parsedConfig.modality_routing : false, - modality_pin_override: - typeof parsedConfig.modality_pin_override === "boolean" ? parsedConfig.modality_pin_override : false, - deployment_affinity: - typeof parsedConfig.deployment_affinity === "boolean" - ? parsedConfig.deployment_affinity - : DEFAULT_DEPLOYMENT_AFFINITY, - adaptive: parsedConfig.adaptive || false, - adaptive_weights: parsedConfig.adaptive_weights, - tier_distance_penalty: parsedConfig.tier_distance_penalty, - adaptive_eligible: parsedConfig.adaptive_eligible || "all", - return_raw_model_name: parsedConfig.return_raw_model_name || false, - enable_context_window_escalation: - typeof parsedConfig.enable_context_window_escalation === "boolean" - ? parsedConfig.enable_context_window_escalation - : undefined, - context_window_escalation_buffer: - typeof parsedConfig.context_window_escalation_buffer === "number" - ? parsedConfig.context_window_escalation_buffer - : undefined, - stall_escalation_enabled: parsedConfig.stall_escalation_enabled === true || undefined, - stall_escalation_window: - typeof parsedConfig.stall_escalation_window === "number" ? parsedConfig.stall_escalation_window : undefined, - stall_escalation_repeat_threshold: - typeof parsedConfig.stall_escalation_repeat_threshold === "number" - ? parsedConfig.stall_escalation_repeat_threshold - : undefined, - }; -}; - +export { hydrateComplexityRouterConfig, hydratePinnedDefaultModel }; export const MANAGED_COMPLEXITY_ROUTER_KEYS = new Set([ "tiers", "enable_non_reasoning_tier", @@ -266,6 +129,16 @@ export const MANAGED_COMPLEXITY_ROUTER_KEYS = new Set([ "stall_escalation_enabled", "stall_escalation_window", "stall_escalation_repeat_threshold", + "code_keywords", + "reasoning_keywords", + "technical_keywords", + "simple_keywords", + "plan_mode_patterns", + "route_housekeeping_to_cheapest_tier", + "housekeeping_patterns", + "reminder_markers", + "max_tokens_from_tier_model", + "classifier_plugin_timeout_ms", ]); // Managed only when the caller passes the corresponding state. A caller that does not render @@ -285,24 +158,6 @@ const toRecord = (value: unknown): Record => { : {}; }; -// A pin lives in two places: complexity_router_config.default_model (this UI's own marker, added -// by PR #36615) and litellm_params.complexity_router_default_model (what the backend reads). Only -// the marker proves an operator picked it, because before #36615 every save wrote a tier-derived -// value into litellm_params. So with no marker, a litellm_params value counts as a pin only when -// it diverges from what the tiers alone derive; a match stays unpinned and keeps tracking tiers. -export const hydratePinnedDefaultModel = ( - storedConfigDefaultModel: unknown, - litellmParamsDefaultModel: string | null | undefined, - activeTiers: ActiveTierSet, -): string | undefined => { - if (typeof storedConfigDefaultModel === "string" && storedConfigDefaultModel.trim()) { - return storedConfigDefaultModel; - } - const tierDerived = resolveComplexityDefaultModel(activeTiers); - const externalOverride = litellmParamsDefaultModel?.trim(); - return externalOverride && externalOverride !== tierDerived ? externalOverride : undefined; -}; - export interface KeywordMatchingState { keywordTierRules: KeywordTierRule[]; escalationKeywords: string[]; @@ -338,55 +193,13 @@ export const buildUpdatedComplexityRouterConfig = ( ); const builderParams: BuildComplexityRouterConfigParams = { - tiers: value.tiers, - enableNonReasoningTier: value.enable_non_reasoning_tier, - customTierSet: value.custom_tier_set, - defaultModel: value.default_model, - planModeMinTier: value.plan_mode_min_tier, - classificationPrompt: value.classification_prompt, - classificationExamples: value.classification_examples, - heuristicFirstMaxTier: value.heuristic_first_max_tier, - hybridBoundaryMargin: value.hybrid_boundary_margin, - classificationMode: value.classification_mode, - tierLabels: value.tier_labels, - classifierType: value.classifier_type, - jevClassifierConfig: value.jev_classifier_config, - heuristicV2SuccessThreshold: value.heuristic_v2_success_threshold, - capabilityClassifierConfig: value.capability_classifier_config, - llmV2Config: value.llm_v2_config, - classifierLlmConfig: value.classifier_llm_config, - classifierContextWindowSize: value.classifier_context_window_size, - classifierContextBudgetChars: value.classifier_context_budget_chars, - classifierContextPerTurnChars: value.classifier_context_per_turn_chars, - classifierContextIncludeAssistantTurns: value.classifier_context_include_assistant_turns, - classifierFallback: value.classifier_fallback, - sessionAffinity: value.session_affinity ?? DEFAULT_SESSION_AFFINITY, - sessionAffinityTtlSeconds: value.session_affinity_ttl_seconds, - modalityRouting: value.modality_routing ?? false, - modalityPinOverride: value.modality_pin_override ?? false, - deploymentAffinity: value.deployment_affinity ?? DEFAULT_DEPLOYMENT_AFFINITY, + ...builderParamsFromValue(value), customTechnicalKeywords: customTechnicalKeywords ?? [], keywordTierRules: keywordMatching?.keywordTierRules ?? [], semanticMatchingEnabled: keywordMatching?.semanticMatchingEnabled ?? false, embeddingModel: keywordMatching?.embeddingModel, matchThreshold: keywordMatching?.matchThreshold ?? DEFAULT_MATCH_THRESHOLD, escalationKeywords: keywordMatching?.escalationKeywords ?? [], - adaptive: value.adaptive ?? false, - adaptiveWeights: value.adaptive_weights ?? DEFAULT_ADAPTIVE_WEIGHTS, - tierDistancePenalty: value.tier_distance_penalty ?? DEFAULT_TIER_DISTANCE_PENALTY, - adaptiveEligible: value.adaptive_eligible ?? "all", - returnRawModelName: value.return_raw_model_name ?? false, - tierBoundaries: value.tier_boundaries, - tokenThresholds: value.token_thresholds, - dimensionWeights: value.dimension_weights, - customDimensions: value.custom_dimensions, - reasoningOverrideMinScore: value.reasoning_override_min_score, - tierModelParams: value.tier_model_params, - enableContextWindowEscalation: value.enable_context_window_escalation, - contextWindowEscalationBuffer: value.context_window_escalation_buffer, - stallEscalationEnabled: value.stall_escalation_enabled, - stallEscalationWindow: value.stall_escalation_window, - stallEscalationRepeatThreshold: value.stall_escalation_repeat_threshold, }; const built = buildComplexityRouterConfig(builderParams); @@ -557,12 +370,13 @@ const EditAutoRouterModal: React.FC = ({ setRouterConfig(parsedConfig); // Set form values - form.reset({ + const routerFormValues = { auto_router_name: modelData.model_name, auto_router_default_model: modelData.litellm_params?.auto_router_default_model || null, auto_router_embedding_model: modelData.litellm_params?.auto_router_embedding_model || null, model_access_group: modelData.model_info?.access_groups || [], - }); + }; + form.reset(routerFormValues); } catch (error) { console.error("Error parsing auto router config:", error); toast.fromError("Error loading auto router configuration"); @@ -585,6 +399,11 @@ const EditAutoRouterModal: React.FC = ({ const classifierError = getClassifierModelError(complexityRouterConfig) ?? getHeuristicV2SuccessThresholdError(complexityRouterConfig.heuristic_v2_success_threshold) ?? + getReminderMarkersError(complexityRouterConfig.reminder_markers) ?? + getClassifierPluginTimeoutError( + complexityRouterConfig.classifier_type, + complexityRouterConfig.classifier_plugin_timeout_ms, + ) ?? getForecastConfigError(complexityRouterConfig) ?? (heuristicScoringRole(complexityRouterConfig) === "decides" ? customDimensionsError(complexityRouterConfig.custom_dimensions) @@ -635,11 +454,18 @@ const EditAutoRouterModal: React.FC = ({ // Dual write: complexity_router_config.default_model (the pin marker hydratePinnedDefaultModel // reads back) and complexity_router_default_model (what the backend routes on) must always be // written together from the same value. Same pairing in add_auto_router_tab.tsx. + const keywordMatching = { + keywordTierRules, + escalationKeywords, + semanticMatchingEnabled, + embeddingModel, + matchThreshold, + }; const updatedConfig = buildUpdatedComplexityRouterConfig( modelData.litellm_params?.complexity_router_config, complexityRouterConfig, customTechnicalKeywords, - { keywordTierRules, escalationKeywords, semanticMatchingEnabled, embeddingModel, matchThreshold }, + keywordMatching, ); const serverVerdict = await validateAutoRouterConfig(accessToken, updatedConfig, modelData?.model_info?.team_id); const dryRunError = dryRunRejection(serverVerdict); @@ -676,12 +502,13 @@ const EditAutoRouterModal: React.FC = ({ ); toast.success("Auto router configuration updated successfully"); - onSuccess({ + const updatedModelData = { ...modelData, model_name: values.auto_router_name, litellm_params: updatedLitellmParams, model_info: updatedModelInfo, - }); + }; + onSuccess(updatedModelData); onCancel(); return; } diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/hydrate_complexity_router_config.ts b/ui/litellm-dashboard/src/components/edit_auto_router/hydrate_complexity_router_config.ts new file mode 100644 index 00000000000..6dbd2b19b52 --- /dev/null +++ b/ui/litellm-dashboard/src/components/edit_auto_router/hydrate_complexity_router_config.ts @@ -0,0 +1,186 @@ +import { defaultJevClassifierConfig, jevClassifierConfigSchema } from "../add_model/jev_classifier_config"; +import { capabilitySettingsSchema, fuseSettingsSchema } from "../add_model/forecast_classifier_config"; +import type { StoredComplexityRouterConfig } from "../add_model/build_complexity_router_config"; +import { + hydrateBuiltInTiers, + hydrateCustomTierSet, + hydratePlanModeMinTier, + hydrateTierLabels, +} from "../add_model/build_complexity_router_config"; +import { hydrateTierModelParams } from "../add_model/complexity_router_tiers"; +import { hydrateCustomDimensions } from "../add_model/custom_dimensions"; +import { + hydrateDimensionWeights, + hydrateReasoningOverrideMinScore, + hydrateTierBoundaries, + hydrateTokenThresholds, +} from "../add_model/heuristic_scoring_knobs"; +import type { ComplexityRouterConfigValue } from "../add_model/ComplexityRouterConfig"; +import { DEFAULT_DEPLOYMENT_AFFINITY, DEFAULT_SESSION_AFFINITY } from "../add_model/ComplexityRouterConfig"; +import { + type ActiveTierSet, + activeTierRows, + tierParamsByRowId, + resolveComplexityDefaultModel, +} from "../add_model/tier_rows"; + +const isReminderMarkerPair = (input: unknown): input is { open: string; close: string } => { + if (typeof input !== "object" || input === null) { + return false; + } + if (!("open" in input) || !("close" in input)) { + return false; + } + return typeof input.open === "string" && typeof input.close === "string"; +}; + +const stringList = (input: unknown): string[] | undefined => + Array.isArray(input) ? input.filter((item): item is string => typeof item === "string") : undefined; + +export const hydratePinnedDefaultModel = ( + storedConfigDefaultModel: unknown, + litellmParamsDefaultModel: string | null | undefined, + activeTiers: ActiveTierSet, +): string | undefined => { + if (typeof storedConfigDefaultModel === "string" && storedConfigDefaultModel.trim()) { + return storedConfigDefaultModel; + } + const tierDerived = resolveComplexityDefaultModel(activeTiers); + const externalOverride = litellmParamsDefaultModel?.trim(); + return externalOverride && externalOverride !== tierDerived ? externalOverride : undefined; +}; + +export const hydrateComplexityRouterConfig = ( + parsedConfig: StoredComplexityRouterConfig, + complexityRouterDefaultModel: string | null | undefined, +): ComplexityRouterConfigValue => { + const builtIn = hydrateBuiltInTiers(parsedConfig.tiers, parsedConfig.enable_non_reasoning_tier); + const { tiers: hydratedTiers, enable_non_reasoning_tier } = builtIn; + const custom_tier_set = hydrateCustomTierSet(parsedConfig); + const activeTiers = { ...builtIn, custom_tier_set }; + + return { + tiers: hydratedTiers, + enable_non_reasoning_tier, + custom_tier_set, + tier_model_params: tierParamsByRowId( + hydrateTierModelParams(parsedConfig.tiers, parsedConfig.tier_model_configs), + activeTierRows(activeTiers), + ), + default_model: hydratePinnedDefaultModel(parsedConfig.default_model, complexityRouterDefaultModel, activeTiers), + plan_mode_min_tier: hydratePlanModeMinTier(parsedConfig.plan_mode_min_tier, custom_tier_set), + tier_labels: hydrateTierLabels(parsedConfig.tier_labels), + classifier_type: parsedConfig.classifier_type || "heuristic", + heuristic_v2_success_threshold: + typeof parsedConfig.heuristic_v2_success_threshold === "number" + ? parsedConfig.heuristic_v2_success_threshold + : undefined, + capability_classifier_config: capabilitySettingsSchema.safeParse(parsedConfig.capability_classifier_config).data, + llm_v2_config: fuseSettingsSchema.safeParse(parsedConfig.llm_v2_config).data, + classifier_llm_config: parsedConfig.classifier_type === "jev" ? undefined : parsedConfig.classifier_llm_config, + jev_classifier_config: + parsedConfig.classifier_type === "jev" + ? jevClassifierConfigSchema.safeParse(parsedConfig.jev_classifier_config ?? {}).data ?? + defaultJevClassifierConfig() + : undefined, + classifier_context_window_size: + typeof parsedConfig.classifier_context_window_size === "number" + ? parsedConfig.classifier_context_window_size + : undefined, + classifier_context_budget_chars: + typeof parsedConfig.classifier_context_budget_chars === "number" + ? parsedConfig.classifier_context_budget_chars + : undefined, + classifier_context_per_turn_chars: + typeof parsedConfig.classifier_context_per_turn_chars === "number" + ? parsedConfig.classifier_context_per_turn_chars + : undefined, + classifier_context_include_assistant_turns: + typeof parsedConfig.classifier_context_include_assistant_turns === "boolean" + ? parsedConfig.classifier_context_include_assistant_turns + : undefined, + classifier_fallback: + parsedConfig.classifier_fallback === "default_model" || parsedConfig.classifier_fallback === "heuristic" + ? parsedConfig.classifier_fallback + : undefined, + classification_prompt: + typeof parsedConfig.classification_prompt === "string" && parsedConfig.classification_prompt.trim() !== "" + ? parsedConfig.classification_prompt + : undefined, + classification_examples: + typeof parsedConfig.classification_examples === "string" && parsedConfig.classification_examples.trim() !== "" + ? parsedConfig.classification_examples + : undefined, + heuristic_first_max_tier: + typeof parsedConfig.heuristic_first_max_tier === "string" && parsedConfig.heuristic_first_max_tier.trim() !== "" + ? parsedConfig.heuristic_first_max_tier + : undefined, + hybrid_boundary_margin: + typeof parsedConfig.hybrid_boundary_margin === "number" ? parsedConfig.hybrid_boundary_margin : undefined, + classification_mode: + parsedConfig.classification_mode === "user_turn" || parsedConfig.classification_mode === "every_request" + ? parsedConfig.classification_mode + : undefined, + tier_boundaries: hydrateTierBoundaries(parsedConfig.tier_boundaries), + token_thresholds: hydrateTokenThresholds(parsedConfig.token_thresholds), + dimension_weights: hydrateDimensionWeights(parsedConfig.dimension_weights), + custom_dimensions: hydrateCustomDimensions(parsedConfig.custom_dimensions), + reasoning_override_min_score: hydrateReasoningOverrideMinScore(parsedConfig.reasoning_override_min_score), + session_affinity: + typeof parsedConfig.session_affinity === "boolean" ? parsedConfig.session_affinity : DEFAULT_SESSION_AFFINITY, + session_affinity_ttl_seconds: + typeof parsedConfig.session_affinity_ttl_seconds === "number" && + Number.isFinite(parsedConfig.session_affinity_ttl_seconds) + ? parsedConfig.session_affinity_ttl_seconds + : undefined, + modality_routing: typeof parsedConfig.modality_routing === "boolean" ? parsedConfig.modality_routing : false, + modality_pin_override: + typeof parsedConfig.modality_pin_override === "boolean" ? parsedConfig.modality_pin_override : false, + deployment_affinity: + typeof parsedConfig.deployment_affinity === "boolean" + ? parsedConfig.deployment_affinity + : DEFAULT_DEPLOYMENT_AFFINITY, + adaptive: parsedConfig.adaptive || false, + adaptive_weights: parsedConfig.adaptive_weights, + tier_distance_penalty: parsedConfig.tier_distance_penalty, + adaptive_eligible: parsedConfig.adaptive_eligible || "all", + return_raw_model_name: parsedConfig.return_raw_model_name || false, + enable_context_window_escalation: + typeof parsedConfig.enable_context_window_escalation === "boolean" + ? parsedConfig.enable_context_window_escalation + : undefined, + context_window_escalation_buffer: + typeof parsedConfig.context_window_escalation_buffer === "number" + ? parsedConfig.context_window_escalation_buffer + : undefined, + stall_escalation_enabled: parsedConfig.stall_escalation_enabled === true || undefined, + stall_escalation_window: + typeof parsedConfig.stall_escalation_window === "number" ? parsedConfig.stall_escalation_window : undefined, + stall_escalation_repeat_threshold: + typeof parsedConfig.stall_escalation_repeat_threshold === "number" + ? parsedConfig.stall_escalation_repeat_threshold + : undefined, + code_keywords: stringList(parsedConfig.code_keywords), + reasoning_keywords: stringList(parsedConfig.reasoning_keywords), + technical_keywords: stringList(parsedConfig.technical_keywords), + simple_keywords: stringList(parsedConfig.simple_keywords), + plan_mode_patterns: stringList(parsedConfig.plan_mode_patterns), + route_housekeeping_to_cheapest_tier: + typeof parsedConfig.route_housekeeping_to_cheapest_tier === "boolean" + ? parsedConfig.route_housekeeping_to_cheapest_tier + : undefined, + housekeeping_patterns: stringList(parsedConfig.housekeeping_patterns), + reminder_markers: Array.isArray(parsedConfig.reminder_markers) + ? parsedConfig.reminder_markers.filter(isReminderMarkerPair) + : undefined, + max_tokens_from_tier_model: + typeof parsedConfig.max_tokens_from_tier_model === "boolean" + ? parsedConfig.max_tokens_from_tier_model + : undefined, + classifier_plugin_timeout_ms: + typeof parsedConfig.classifier_plugin_timeout_ms === "number" && + Number.isFinite(parsedConfig.classifier_plugin_timeout_ms) + ? parsedConfig.classifier_plugin_timeout_ms + : undefined, + }; +}; diff --git a/ui/litellm-dashboard/src/components/shared/ScopedSavingsTab.tsx b/ui/litellm-dashboard/src/components/shared/ScopedSavingsTab.tsx new file mode 100644 index 00000000000..7a0455f9d7e --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/ScopedSavingsTab.tsx @@ -0,0 +1,133 @@ +"use client"; + +import { useMemo, useState } from "react"; + +import { AreaChart, BarChart, CustomLegend } from "@/components/shared/charts"; +import AdvancedDatePicker from "@/components/shared/advanced_date_picker"; +import SavingsTiles from "@/components/shared/SavingsTiles"; +import { Card, CardAction, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card"; +import { Tabs, TabsList, TabsTrigger } from "@/components/ui/tabs"; +import { + formatRangeLabel, + localIsoDay, + MAX_POINTS_WITH_DOTS, + SAVINGS_COLORS, + SAVINGS_SERIES, + SavingsAccumulation, + SavingsPoint, + savingsSeriesOf, + shortDate, + toCumulative, + usd, + withStartAnchor, +} from "@/app/(dashboard)/cost-optimization/_components/costOptimizationUtils"; +import { + useScopedDailyActivityRange, + type ActivityDateRange, + type DailyActivityScope, +} from "@/app/(dashboard)/cost-optimization/_components/useDailyActivityRange"; + +interface ScopedSavingsTabProps { + accessToken: string | null; + scope: DailyActivityScope; + activity: ActivityDateRange; + entityType: "key" | "user"; + scopeNote?: string; +} + +const ScopedSavingsTab = ({ accessToken, scope, activity, entityType, scopeNote }: ScopedSavingsTabProps) => { + const { dateValue, onDateChange, results, loading, isFetchingMore, failed, cancelled } = useScopedDailyActivityRange( + accessToken, + scope, + activity, + ); + const startTime = dateValue.from; + const endTime = dateValue.to; + + const [accumulation, setAccumulation] = useState("cumulative"); + + const perInterval = useMemo(() => savingsSeriesOf(results), [results]); + + const overTime = useMemo(() => { + if (accumulation !== "cumulative") return perInterval; + const startLabel = startTime ? shortDate(localIsoDay(startTime)) : ""; + return withStartAnchor(toCumulative(perInterval), startLabel); + }, [accumulation, perInterval, startTime]); + + const intervalLabel = "Per day"; + const rangeLabel = formatRangeLabel(startTime, endTime); + const savingsSubtitle = [ + accumulation === "cumulative" ? "Running total saved" : `Saved ${intervalLabel.toLowerCase()}`, + rangeLabel && `${rangeLabel} (UTC)`, + ] + .filter(Boolean) + .join(" · "); + + const isLoading = loading || isFetchingMore; + const unavailable = failed || cancelled; + const showResults = !isLoading && !unavailable; + const hasRows = results.length > 0; + const showEmpty = !unavailable && (isLoading || !hasRows); + const showChart = showResults && hasRows; + const chartProps = { + data: overTime, + index: "date", + categories: SAVINGS_SERIES, + colors: SAVINGS_COLORS, + valueFormatter: usd, + showLegend: false, + }; + + return ( +
+
+ Spend is bucketed by UTC day + +
+ + {scopeNote && ( +

+ {scopeNote} +

+ )} + + {unavailable && ( +

+ Savings are unavailable for this range. Try another date range or reopen this tab. +

+ )} + {showResults && } + + + + Savings + {savingsSubtitle} + + + setAccumulation(value as SavingsAccumulation)}> + + Cumulative + {intervalLabel} + + + + + + {showEmpty && ( +

+ {isLoading ? "Loading savings..." : `No usage recorded for this ${entityType} in this range.`} +

+ )} + {showChart && + (accumulation === "cumulative" ? ( + + ) : ( + + ))} +
+
+
+ ); +}; + +export default ScopedSavingsTab; diff --git a/ui/litellm-dashboard/src/components/shared/SearchSelect.integration.test.tsx b/ui/litellm-dashboard/src/components/shared/SearchSelect.integration.test.tsx index c981010dff9..9f320049b09 100644 --- a/ui/litellm-dashboard/src/components/shared/SearchSelect.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/shared/SearchSelect.integration.test.tsx @@ -111,4 +111,36 @@ describe("SearchSelect", () => { expect(screen.queryByText("Growth")).not.toBeInTheDocument(); expect(onValueChange).not.toHaveBeenCalled(); }); + + it("supports keyboard select, clear, and reselect", async () => { + const onValueChange = vi.fn(); + const user = userEvent.setup(); + function Controlled() { + const [value, setValue] = useState(null); + return ( + { + setValue(next); + onValueChange(next); + }} + /> + ); + } + + render(); + const input = screen.getByRole("combobox"); + await user.tab(); + await user.keyboard("{Enter}"); + await user.keyboard("{ArrowDown}{Enter}"); + expect(onValueChange).toHaveBeenLastCalledWith("team-1"); + const clear = screen.getByRole("button", { name: "Clear" }); + clear.focus(); + await user.keyboard("{Enter}"); + expect(onValueChange).toHaveBeenLastCalledWith(null); + input.focus(); + await user.keyboard("{Enter}{ArrowDown}{Enter}"); + expect(onValueChange).toHaveBeenLastCalledWith("team-1"); + }); }); diff --git a/ui/litellm-dashboard/src/components/templates/KeySavingsTab.tsx b/ui/litellm-dashboard/src/components/templates/KeySavingsTab.tsx index c33529eb042..d1395e0be51 100644 --- a/ui/litellm-dashboard/src/components/templates/KeySavingsTab.tsx +++ b/ui/litellm-dashboard/src/components/templates/KeySavingsTab.tsx @@ -1,132 +1,29 @@ "use client"; -import React, { useMemo, useState } from "react"; - -import { AreaChart, BarChart, CustomLegend } from "@/components/shared/charts"; -import AdvancedDatePicker from "@/components/shared/advanced_date_picker"; -import SavingsTiles from "@/components/shared/SavingsTiles"; -import { Card, CardAction, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card"; -import { Tabs, TabsList, TabsTrigger } from "@/components/ui/tabs"; +import ScopedSavingsTab from "@/components/shared/ScopedSavingsTab"; import { hasProxyWideSpendView, spendScopeUserId } from "@/utils/roles"; -import { - formatRangeLabel, - localIsoDay, - MAX_POINTS_WITH_DOTS, - SAVINGS_COLORS, - SAVINGS_SERIES, - SavingsAccumulation, - SavingsPoint, - savingsSeriesOf, - shortDate, - toCumulative, - usd, - withStartAnchor, -} from "@/app/(dashboard)/cost-optimization/_components/costOptimizationUtils"; -import { - useScopedDailyActivityRange, - type ActivityDateRange, -} from "@/app/(dashboard)/cost-optimization/_components/useDailyActivityRange"; +import type { ActivityDateRange } from "@/app/(dashboard)/cost-optimization/_components/useDailyActivityRange"; interface KeySavingsTabProps { accessToken: string | null; - /** The key's token hash — what spend rows are keyed by, not the one-time plaintext secret. */ keyToken: string; userId: string | null; userRole: string; activity: ActivityDateRange; } -const KeySavingsTab: React.FC = ({ accessToken, keyToken, userId, userRole, activity }) => { - // Proxy admins read the whole key. For anyone else the endpoint applies the caller's own user_id - // alongside the key filter, so the figures cover only that viewer's requests on this key -- said - // plainly in the scope note below rather than left to be misread as the key's total. - const readsWholeKey = hasProxyWideSpendView(userRole); - const { dateValue, onDateChange, results, loading, isFetchingMore } = useScopedDailyActivityRange( - accessToken, - { userId: spendScopeUserId(userRole, userId), apiKey: keyToken }, - activity, - ); - const startTime = dateValue.from ?? null; - const endTime = dateValue.to ?? null; - - const [accumulation, setAccumulation] = useState("cumulative"); - - const perInterval = useMemo(() => savingsSeriesOf(results), [results]); - - const overTime = useMemo(() => { - if (accumulation !== "cumulative") return perInterval; - const startLabel = startTime ? shortDate(localIsoDay(startTime)) : ""; - return withStartAnchor(toCumulative(perInterval), startLabel); - }, [accumulation, perInterval, startTime]); - - const intervalLabel = "Per day"; - const rangeLabel = formatRangeLabel(startTime ?? undefined, endTime ?? undefined); - const savingsSubtitle = [ - accumulation === "cumulative" ? "Running total saved" : `Saved ${intervalLabel.toLowerCase()}`, - rangeLabel && `${rangeLabel} (UTC)`, - ] - .filter(Boolean) - .join(" · "); - - const isLoading = loading || isFetchingMore; - const hasRows = results.length > 0; - const chartProps = { - data: overTime, - index: "date", - categories: SAVINGS_SERIES, - colors: SAVINGS_COLORS, - valueFormatter: usd, - showLegend: false, - }; - - return ( -
-
- Spend is bucketed by UTC day - -
- - {!readsWholeKey && ( -

- Showing your own requests on this key. A key shared across a team will have spend from other members that is - not counted here. -

- )} - - - - - - Savings - {savingsSubtitle} - - - setAccumulation(value as SavingsAccumulation)}> - - Cumulative - {intervalLabel} - - - - - - {/* Distinguishes "still fetching" from "this key genuinely had no traffic": an empty - chart alone reads as a broken panel, and a $0.00 tile reads as a real zero. */} - {!hasRows && ( -

- {isLoading ? "Loading savings..." : "No usage recorded for this key in this range."} -

- )} - {hasRows && accumulation === "cumulative" && ( - - )} - {/* Not stacked: auto-router can go negative on a cold-cache write, and stacking would - draw that below the axis while the rest of the bar still read as the total. */} - {hasRows && accumulation !== "cumulative" && } -
-
-
- ); -}; +const KeySavingsTab = ({ accessToken, keyToken, userId, userRole, activity }: KeySavingsTabProps) => ( + +); export default KeySavingsTab; diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.test.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.test.tsx index 5a35c7ae16b..8d1847e0121 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsFilters.test.tsx @@ -344,4 +344,15 @@ describe("RequestLogsFilters", () => { expect(set).toHaveBeenCalledWith(LOG_FILTER_IDS.CACHE_STATUS, undefined); }); + + it("clears the raw Error Code combobox through the undefined filter contract", async () => { + const user = userEvent.setup(); + const { set } = renderFilters({ [LOG_FILTER_IDS.ERROR_CODE]: "429" }); + const input = await screen.findByPlaceholderText("Select or type an error code"); + + await user.click(input); + await user.click(screen.getByRole("button", { name: "Clear", hidden: true })); + + expect(set).toHaveBeenCalledWith(LOG_FILTER_IDS.ERROR_CODE, undefined); + }); }); diff --git a/ui/litellm-dashboard/src/lib/autorouter_presets.test.ts b/ui/litellm-dashboard/src/lib/autorouter_presets.test.ts index e1d9e6ab5ac..87efb2b1d5e 100644 --- a/ui/litellm-dashboard/src/lib/autorouter_presets.test.ts +++ b/ui/litellm-dashboard/src/lib/autorouter_presets.test.ts @@ -746,7 +746,7 @@ describe("autorouter_presets", () => { expect(prefill.escalationKeywords).toEqual([]); }); - it("carries a preset's context-window escalation opt-out and buffer through the prefill", () => { + it.each([undefined, false, true])("preserves a preset's context-window escalation setting: %s", (enabled) => { const prefill = buildPresetPrefill( { tiers: { SIMPLE: ["gpt-5-nano"], MEDIUM: [], COMPLEX: [], REASONING: [] }, @@ -754,12 +754,12 @@ describe("autorouter_presets", () => { classification_mode: "every_request", session_affinity: false, deployment_affinity: true, - enable_context_window_escalation: false, + enable_context_window_escalation: enabled, context_window_escalation_buffer: 0.9, }, groupsOnly(["gpt-5-nano"]), ); - expect(prefill.complexityRouterConfig.enable_context_window_escalation).toBe(false); + expect(prefill.complexityRouterConfig.enable_context_window_escalation).toBe(enabled); expect(prefill.complexityRouterConfig.context_window_escalation_buffer).toBe(0.9); }); diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 84e3270c48b..58b5cedc2cf 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -1246,9 +1246,10 @@ export interface paths { * @description Benchmarks for the auto-router dashboard: session shape, savings against the configured * baseline, and prompt-caching behaviour bucketed by what the router did. * - * Reads the LiteLLM_AutoRouterSession rollup, folded once per request at spend-write time, - * so this endpoint never scans LiteLLM_SpendLogs. A session is in the window when it - * overlaps it: its last turn is on or after start_date and its first turn is on or before + * Reads session rollups folded once per request at spend-write time, so this endpoint + * never scans LiteLLM_SpendLogs. A user filter selects only turns attributed to that + * internal user when written; older key-only history remains outside user views. A session + * is in the window when it overlaps it: its last turn is on or after start_date and its first turn is on or before * end_date. Overall hit rate is over telemetry-bearing turns; each bucket's hit rate is * over that bucket's turns. * @@ -36884,8 +36885,8 @@ export interface components { embedding_model?: string | null; /** * Enable Context Window Escalation - * @description Escalate a request off a tier whose models provably cannot hold its prompt, before dispatch. The classifier scores complexity and never prompt size, so a long agentic session whose newest ask is trivial lands on a small-window tier and the provider rejects it with a context-window 400 that nothing retries. When every model of the decided tier has a declared window smaller than the estimated prompt, the request moves to the lowest configured tier with a model whose declared window fits; when only some of the tier's models fit, the pick is restricted to those and the tier keeps the request. Models with no resolvable window are never escalated away from and never escalated onto. Set false to dispatch on complexity alone, as before. - * @default true + * @description Escalate a request off a tier whose models provably cannot hold its prompt, before dispatch. The classifier scores complexity and never prompt size, so a long agentic session whose newest ask is trivial lands on a small-window tier and the provider rejects it with a context-window 400 that nothing retries. When every model of the decided tier has a declared window smaller than the estimated prompt, the request moves to the lowest configured tier with a model whose declared window fits; when only some of the tier's models fit, the pick is restricted to those and the tier keeps the request. Models with no resolvable window are never escalated away from and never escalated onto. Disabled by default: omit or set false to dispatch on complexity alone; set true to enable context-window escalation. + * @default false */ enable_context_window_escalation: boolean; /** @@ -37244,6 +37245,13 @@ export interface components { routing_strategy_descriptions: { [key: string]: string; }; + /** + * Source + * @description Source of each current router setting + */ + source: { + [key: string]: "config" | "db" | "env" | "default" | "unset"; + }; }; /** * RoutingGroup @@ -39707,6 +39715,10 @@ export interface components { field_schema: { [key: string]: unknown; }; + /** Source */ + source: { + [key: string]: "config" | "db" | "env" | "default" | "unset"; + }; /** Values */ values: { [key: string]: unknown; @@ -43881,6 +43893,8 @@ export interface operations { end_date?: string | null; /** @description Filter to one virtual key token hash */ api_key?: string | null; + /** @description Filter to one canonical internal user recorded on each turn */ + user_id?: string | null; }; header?: never; path?: never;