From 5e38a087418f5d6a1323be709a58b527b3d01d49 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 00:01:44 +0000 Subject: [PATCH] feat(cache): select Rust caching through explicit cache objects (#43601) * refactor(cache): organize v2 cache as a package * docs: clarify experimental v2 guidance * fix(cache): verify cache-hit accounting and preserve logging metadata * refactor(cache): separate execution facts from host accounting * refactor(rust): build messages routes with named dependencies * wip * fix(cache): preserve facade policy and preflight fallback * refactor(cache): defer shared Python logging changes * test(gateway-inference): allow dead code in shared test helpers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cache): key prepared requests and honor facade controls * feat(cache): use Python caches from Rust Messages inference * refactor(cache): separate native and Python cache adapters * refactor(cache): enforce shared composition and adapter boundaries * fix(cache): let Python key delegated Rust Messages entries --------- Co-authored-by: Yujong Lee Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm-rust/Cargo.lock | 12 + litellm-rust/crates/cache-gcs/tests/cache.rs | 1 + litellm-rust/crates/cache-response/AGENTS.md | 29 + litellm-rust/crates/cache-response/Cargo.toml | 3 + litellm-rust/crates/cache-response/README.md | 51 - litellm-rust/crates/cache-response/src/lib.rs | 6 + .../crates/cache-response/src/response.rs | 29 +- .../crates/cache-response/src/service.rs | 177 +++ .../crates/cache-response/tests/response.rs | 128 +- .../crates/cache-response/tests/service.rs | 176 +++ .../callbacks-legacy-python/src/adapter.rs | 62 +- .../callbacks-legacy-python/src/mapping.rs | 9 +- litellm-rust/crates/core/AGENTS.md | 12 + litellm-rust/crates/core/Cargo.toml | 5 + litellm-rust/crates/core/src/caching.rs | 318 +++++ .../core/src/chat_completions/handler.rs | 127 +- .../crates/core/src/chat_completions/mod.rs | 58 +- .../crates/core/src/chat_completions/route.rs | 56 +- litellm-rust/crates/core/src/lib.rs | 25 + .../crates/core/src/messages/handler.rs | 91 +- litellm-rust/crates/core/src/messages/mod.rs | 126 +- .../crates/core/src/messages/route.rs | 31 +- .../crates/core/src/responses/handler.rs | 137 +- litellm-rust/crates/core/src/responses/mod.rs | 58 +- .../crates/core/src/responses/prepare.rs | 29 +- .../crates/core/src/responses/route.rs | 39 +- litellm-rust/crates/core/tests/caching.rs | 1217 +++++++++++++++++ .../crates/core/tests/chat_completions.rs | 12 +- .../crates/core/tests/messages/response.rs | 108 +- .../crates/core/tests/ocr/lifecycle.rs | 1 + litellm-rust/crates/core/tests/ocr/machine.rs | 6 + litellm-rust/crates/core/tests/responses.rs | 39 +- litellm-rust/crates/core/tests/support/mod.rs | 10 +- .../crates/gateway-inference/Cargo.toml | 4 + .../crates/gateway-inference/src/caching.rs | 109 ++ .../gateway-inference/src/chat_completions.rs | 19 +- .../crates/gateway-inference/src/lib.rs | 16 +- .../crates/gateway-inference/src/messages.rs | 16 +- .../crates/gateway-inference/src/responses.rs | 16 +- .../crates/gateway-inference/tests/caching.rs | 185 +++ .../gateway-inference/tests/support/mod.rs | 102 +- litellm-rust/crates/host-native/src/driver.rs | 5 + litellm-rust/crates/host-python/src/driver.rs | 137 +- .../host-python/src/hooks/chain/adapter.rs | 6 +- .../host-python/src/hooks/chain/dispatch.rs | 13 +- .../crates/host-python/src/services.rs | 16 + .../crates/host-python/tests/hook_chain.rs | 4 +- litellm-rust/crates/host/AGENTS.md | 2 + litellm-rust/crates/host/src/hooks.rs | 6 +- litellm-rust/crates/host/src/interceptors.rs | 26 + litellm-rust/crates/host/src/lifecycle.rs | 12 +- .../crates/host/src/machine/context.rs | 11 + litellm-rust/crates/host/src/protocol.rs | 4 + .../crates/python-bridge/src/cache/AGENTS.md | 4 + .../crates/python-bridge/src/cache/handle.rs | 363 ----- .../crates/python-bridge/src/cache/mod.rs | 21 +- .../python-bridge/src/cache/native/AGENTS.md | 9 + .../src/cache/{ => native}/activation.rs | 6 +- .../cache/{native.rs => native/backend.rs} | 54 +- .../src/cache/{ => native}/config.rs | 14 +- .../src/cache/{ => native}/embedder.rs | 14 +- .../src/cache/{ => native}/facade.rs | 40 +- .../src/cache/{ => native}/identity.rs | 9 +- .../python-bridge/src/cache/native/mod.rs | 11 + .../src/cache/{ => native}/request.rs | 8 +- .../src/cache/{ => native}/semantic.rs | 4 +- .../python-bridge/src/cache/native/v2.rs | 345 +++++ .../python-bridge/src/cache/python/AGENTS.md | 9 + .../src/cache/{ => python}/callback.rs | 29 +- .../python-bridge/src/cache/python/host.rs | 141 ++ .../python-bridge/src/cache/python/mod.rs | 8 + .../python-bridge/src/cache/python/service.rs | 107 ++ .../python-bridge/src/cache/resolver.rs | 25 - .../src/cache/{binding.rs => runtime.rs} | 20 +- .../python-bridge/src/cache/selection.rs | 176 +++ litellm-rust/crates/python-bridge/src/lib.rs | 12 +- .../src/routes/chat_completions.rs | 17 +- .../python-bridge/src/routes/messages/host.rs | 52 +- .../python-bridge/src/routes/messages/mod.rs | 41 +- .../python-bridge/src/routes/responses.rs | 17 +- litellm/_v2/AGENTS.md | 3 + litellm/_v2/__init__.py | 3 + litellm/_v2/cache/AGENTS.md | 11 + litellm/_v2/cache/__init__.py | 72 + litellm/caching/caching.py | 13 +- litellm/rust_bridge/_native.pyi | 101 +- .../rust_bridge/callbacks_legacy_python.py | 29 +- litellm/rust_bridge/catalog.py | 30 +- litellm/rust_bridge/public_call.py | 6 +- litellm/rust_bridge/response_cache.py | 47 +- litellm/rust_bridge/response_metadata.py | 7 +- .../cache/test_azure_blob.py | 64 +- tests/test_litellm_rust/cache/test_disk.py | 56 +- tests/test_litellm_rust/cache/test_facade.py | 147 +- tests/test_litellm_rust/cache/test_gcs.py | 246 +--- .../cache/test_qdrant_semantic.py | 76 +- tests/test_litellm_rust/cache/test_redis.py | 54 +- .../cache/test_redis_semantic.py | 98 +- tests/test_litellm_rust/cache/test_rollout.py | 35 +- tests/test_litellm_rust/cache/test_s3.py | 70 +- tests/test_litellm_rust/cache/test_v2.py | 794 +++++++++++ .../cache/test_valkey_semantic.py | 195 +-- tests/test_litellm_rust/support/cache.py | 38 +- tests/test_litellm_rust/support/fake_gcs.py | 152 -- .../test_callbacks_legacy_python.py | 29 + tests/unit/rust_bridge/test_catalog.py | 26 +- tests/unit/rust_bridge/test_dispatch.py | 7 +- tests/unit/rust_bridge/test_runtime.py | 18 +- 108 files changed, 5844 insertions(+), 2036 deletions(-) create mode 100644 litellm-rust/crates/cache-response/AGENTS.md delete mode 100644 litellm-rust/crates/cache-response/README.md create mode 100644 litellm-rust/crates/cache-response/src/service.rs create mode 100644 litellm-rust/crates/cache-response/tests/service.rs create mode 100644 litellm-rust/crates/core/src/caching.rs create mode 100644 litellm-rust/crates/core/tests/caching.rs create mode 100644 litellm-rust/crates/gateway-inference/src/caching.rs create mode 100644 litellm-rust/crates/gateway-inference/tests/caching.rs delete mode 100644 litellm-rust/crates/python-bridge/src/cache/handle.rs create mode 100644 litellm-rust/crates/python-bridge/src/cache/native/AGENTS.md rename litellm-rust/crates/python-bridge/src/cache/{ => native}/activation.rs (97%) rename litellm-rust/crates/python-bridge/src/cache/{native.rs => native/backend.rs} (95%) rename litellm-rust/crates/python-bridge/src/cache/{ => native}/config.rs (99%) rename litellm-rust/crates/python-bridge/src/cache/{ => native}/embedder.rs (88%) rename litellm-rust/crates/python-bridge/src/cache/{ => native}/facade.rs (94%) rename litellm-rust/crates/python-bridge/src/cache/{ => native}/identity.rs (98%) create mode 100644 litellm-rust/crates/python-bridge/src/cache/native/mod.rs rename litellm-rust/crates/python-bridge/src/cache/{ => native}/request.rs (96%) rename litellm-rust/crates/python-bridge/src/cache/{ => native}/semantic.rs (98%) create mode 100644 litellm-rust/crates/python-bridge/src/cache/native/v2.rs create mode 100644 litellm-rust/crates/python-bridge/src/cache/python/AGENTS.md rename litellm-rust/crates/python-bridge/src/cache/{ => python}/callback.rs (84%) create mode 100644 litellm-rust/crates/python-bridge/src/cache/python/host.rs create mode 100644 litellm-rust/crates/python-bridge/src/cache/python/mod.rs create mode 100644 litellm-rust/crates/python-bridge/src/cache/python/service.rs delete mode 100644 litellm-rust/crates/python-bridge/src/cache/resolver.rs rename litellm-rust/crates/python-bridge/src/cache/{binding.rs => runtime.rs} (95%) create mode 100644 litellm-rust/crates/python-bridge/src/cache/selection.rs create mode 100644 litellm/_v2/AGENTS.md create mode 100644 litellm/_v2/__init__.py create mode 100644 litellm/_v2/cache/AGENTS.md create mode 100644 litellm/_v2/cache/__init__.py create mode 100644 tests/test_litellm_rust/cache/test_v2.py delete mode 100644 tests/test_litellm_rust/support/fake_gcs.py diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 0350aa2f24a..1f3790c7b61 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3552,8 +3552,10 @@ name = "litellm-cache-response" version = "0.1.0" dependencies = [ "litellm-cache", + "litellm-cache-gcs", "litellm-cache-memory", "litellm-cache-redis", + "litellm-http", "py_literal", "redis", "redis-test", @@ -3562,6 +3564,7 @@ dependencies = [ "serde_json", "sha2 0.10.9", "tokio", + "wiremock", ] [[package]] @@ -3646,7 +3649,11 @@ dependencies = [ "litellm-auth", "litellm-auth-aws", "litellm-auth-gcp", + "litellm-cache", + "litellm-cache-memory", + "litellm-cache-response", "litellm-core-utils", + "litellm-framing", "litellm-host", "litellm-host-native", "litellm-http", @@ -3669,6 +3676,7 @@ dependencies = [ "time", "tokio", "tokio-tungstenite", + "tokio-util", "tracing", "url", "veil", @@ -3808,8 +3816,11 @@ dependencies = [ "bytes", "futures-util", "litellm-auth", + "litellm-cache-memory", + "litellm-cache-response", "litellm-core", "litellm-gateway-auth", + "litellm-host", "litellm-host-http", "litellm-http", "litellm-llms", @@ -3817,6 +3828,7 @@ dependencies = [ "litellm-secrets", "litellm-types", "rstest", + "serde", "serde_json", "thiserror 2.0.19", "tokio", diff --git a/litellm-rust/crates/cache-gcs/tests/cache.rs b/litellm-rust/crates/cache-gcs/tests/cache.rs index 12bb5344570..096b691ea16 100644 --- a/litellm-rust/crates/cache-gcs/tests/cache.rs +++ b/litellm-rust/crates/cache-gcs/tests/cache.rs @@ -56,6 +56,7 @@ async fn set_writes_encoded_object_and_headers(#[future(awt)] server: MockServer )] #[case::missing("missing", ResponseTemplate::new(404), Ok(None))] #[case::server_error("server-error", ResponseTemplate::new(500), Err(Error::Unavailable))] +#[case::unauthorized("unauthorized", ResponseTemplate::new(401), Err(Error::Unavailable))] #[case::invalid( "invalid", ResponseTemplate::new(200).set_body_string("not json"), diff --git a/litellm-rust/crates/cache-response/AGENTS.md b/litellm-rust/crates/cache-response/AGENTS.md new file mode 100644 index 00000000000..d86fe6cc588 --- /dev/null +++ b/litellm-rust/crates/cache-response/AGENTS.md @@ -0,0 +1,29 @@ +# Response caching + +Design this crate for shared Rust execution used by the Python SDK and the Rust gateway. The Python SDK will remain, with more core execution moving to Rust and Python callbacks staying in Python. The Rust gateway is still evolving and is intended to replace the Python proxy. Keep response-cache policy independent of Python, HTTP serving, and either proxy's configuration format + +Separate what is cached, how a hit is matched, and where entries are stored. Chat Completions, Messages, Responses, and embeddings are API workloads. Exact and semantic matching are lookup behaviors. Memory, Redis, disk, and object stores are storage choices. Embeddings are inference too, so do not use an inference-cache name to imply a category that excludes embeddings. Consult the existing Python cache and caching handler for behavior and compatibility contracts without copying their class structure + +Storage traits, codecs, and backend capabilities belong in `litellm-cache` and the storage crates. Keep storage reusable for value types beyond LLM responses. This crate owns response entries, matching and freshness semantics, the Python-compatible response codec, and deferred-write policy. Core owns route-specific request identity, response encoding and reconstruction, embedding partial-hit orchestration, and stream capture and replay. Boundaries own configuration translation, resource construction, and caller identity + +Construct and inject the response-cache service at the Python bridge or gateway boundary, as with the HTTP client. Reuse it across calls. Core and provider code must not discover cache configuration through Python globals, process configuration, or backend-specific factories + +Keep `ResponseCache` generic over its storage backend. Preserve typed backend contexts and capability bounds internally. Inject an object-safe service into core for runtime backend selection, so storage types do not spread through route and host types. Keep API request and response types statically typed. Add a generic parameter only where it preserves a useful type relationship or capability + +Keep the core service contract narrow. Lookup and store must not require connection testing, ping, flush, deletion, counters, queues, or scripts. Require batch operations where a consumer needs partial hits, and keep management capabilities on their own interfaces. An exact-only adapter must remain explicit about its matching restriction. Supporting semantic matching requires a defined lookup-context and embedding execution contract, not just a renamed trait + +Separate reusable resources from per-call policy. Backend configuration, namespace, default expiry, and entry limits belong to the configured service or backend. Read/write controls, expiry and freshness overrides, and authenticated caller scope belong to the call. Passing call options must not replace or mutate the route's configured service + +Keep cache misses and storage failures distinguishable in return values. Core owns the decision to continue with provider execution after a cache failure. A read can reject an entry for freshness while the backend still retains it. Preserve the timestamp at which a response was produced when writing it later + +Define lookup placement explicitly relative to authorization, deployment and credential resolution, and request-transforming callbacks. Cache identity must account for every input that affects reuse, including API surface and caller scope, while preserving intentional Python caching groups. Preserve existing keys and response formats unless changing them is an explicit migration decision + +Cache normalized provider results before caller-specific response transformations. Hits must still run the applicable response processing, success callbacks, and cache-hit accounting. Keep callback execution in the host. Python cache implementations and semantic embedders that require the caller's task must use the existing host-operation mechanism rather than Python calls from a Rust worker. Preserve legacy fallback until that contract is supported + +Keep unary caching independent of stream-only methods. Store streams only after successful exhaustion and protocol completion. Errors, incomplete streams, cancellation, and oversized entries must not populate the cache. Embedding batches need ordered partial results and reconstruction around the uncached inputs + +Test each contract in its owner: storage capabilities in backend tests, envelopes and freshness here, reuse and replay in core, Python callback and fallback behavior at the bridge, and HTTP behavior at the gateway. Run backend contract checks and Python response-codec fixtures before exposing a new backend + +`ScopedCache` requires an explicit shared or isolated scope at construction. `CacheOptions` has no default sharing policy. Callers may override policy per invocation without replacing the attached service. Versioned native envelopes reject incompatible API surfaces and versions as misses; this envelope is distinct from the legacy Python response codec + +Response storage is not the source of budget or rate-limit coordination dependencies. Keep counters, reservations, and atomic admission operations out of `ResponseCacheService`, including when both services happen to use Redis diff --git a/litellm-rust/crates/cache-response/Cargo.toml b/litellm-rust/crates/cache-response/Cargo.toml index 42a1afb2ba0..1379573e505 100644 --- a/litellm-rust/crates/cache-response/Cargo.toml +++ b/litellm-rust/crates/cache-response/Cargo.toml @@ -13,9 +13,12 @@ serde_json.workspace = true sha2.workspace = true [dev-dependencies] +litellm-cache-gcs.workspace = true +litellm-http = { workspace = true, features = ["test-support"] } litellm-cache-memory.workspace = true litellm-cache-redis.workspace = true redis = "1.7.0" redis-test = "1.0.4" rstest.workspace = true tokio.workspace = true +wiremock = "0.6.5" diff --git a/litellm-rust/crates/cache-response/README.md b/litellm-rust/crates/cache-response/README.md deleted file mode 100644 index dbad474c9e7..00000000000 --- a/litellm-rust/crates/cache-response/README.md +++ /dev/null @@ -1,51 +0,0 @@ -# Response cache - -`ResponseCache` adds request keys, independent read/write controls, response envelopes, and freshness checks to any `B: BaseCache` - -## Ownership - -`litellm-cache` defines typed storage, codec, and capability traits. `BaseCache` is only get, set, TTL, and pipeline writes. Everything else is an optional capability a backend implements only where its Python class defines the method: `DisconnectCache`, `ConnectionCache` (`test_connection`), `PingCache`, `BatchCache`, `DeleteCache`, `FlushCache`, counters, queues, TTL, scan, and scripts. Memory, Redis, disk, S3, GCS, and Azure Blob implement those traits without depending on response policy, so other consumers can store their own value types in the same backends - -Semantic backends (Redis, Valkey, Qdrant) are generic over their embedder and codec, and share one prompt and embedding contract from `litellm_cache::semantic`. They take a `SemanticCacheContext`, so `ResponseCache` drives them the same way it drives exact backends - -`litellm-cache-response` owns response keys, controls, entries, the Python-compatible response codec, and `WriteBuffer`, the backend-neutral deferred-write policy. It has no runtime dependency on a specific cache backend or Python - -`ExactResponseCache` is the object-safe view of a `ResponseCache` over an exact backend. `ConnectionProbe` is the object-safe `test_connection`, implemented only when the backend implements `ConnectionCache`, so a host holds one next to its `ExactResponseCache` and reports the operation as unsupported otherwise, as Python's `BaseCache` does. Lookup, store, batch, and flush never require it - -## Native Rust use - -```rust -use std::{sync::Arc, time::Duration}; -use litellm_cache_memory::InMemoryCache; -use litellm_cache_response::{CacheKeyInput, ResponseCache, ResponseCacheRequest}; -use serde_json::json; - -let cache = ResponseCache::new(Arc::new(InMemoryCache::default())); -let request = ResponseCacheRequest::new(CacheKeyInput { - preset: Some("example:key".into()), - ..Default::default() -}); -let now = Duration::from_secs(100); -cache.store(&request, json!({"answer": 7}), now)?; -assert_eq!(cache.async_lookup(&request, now).await?, Some(json!({"answer": 7}))); -``` - -For Redis, inject `RedisCache::new(url, ttl, ResponseCacheCodec)` instead. Namespaces are optional and existing namespace prefixes are preserved - -Callers supply Unix time for response freshness. Backend TTL uses its own clock. A read can reject an entry through `max_age` even while the backend still retains it - -## Python integration - -The bridge activates backends through the Rust catalog in `litellm/rust_bridge/catalog.py`. Every cache rule ships as `PYTHON_ONLY`, so SDK, Router, and proxy calls stay on Python and construct no native cache resources until a rule is changed - -When a rule selects a backend, the Python `Cache` facade builds the native runtime from its own configuration and routes its storage calls (sync and async lookup and store, and pipelined batch store) to it. Stream replay, embedding partial-hit merging, response reconstruction, and callbacks stay in Python on top of that native store. The Python backend object remains for its direct API - -Object responses are written as they are, and every other response shape is written as a serialized string, which is the pair of shapes Python reads. A string on the wire is therefore always a serialized response, so string-valued responses round trip. Typed backends such as memory never pass through the codec - -Native cache handles must be recreated after fork. Native errors propagate to the host, which owns the existing fail-open and logging policy - -## Adding another backend - -Implement `BaseCache` for the backend with its associated value type and the capability traits its Python class supports, and accept a `CacheCodec` when wire serialization is needed. `ResponseCache` then works without another response implementation - -Run the `litellm-cache-testing` contract checks the backend's capabilities allow, and run response fixtures with `ResponseCacheCodec`, including both Python envelope encodings, before adding a catalog rule diff --git a/litellm-rust/crates/cache-response/src/lib.rs b/litellm-rust/crates/cache-response/src/lib.rs index a6a4bb3eb64..ebabcf70c9f 100644 --- a/litellm-rust/crates/cache-response/src/lib.rs +++ b/litellm-rust/crates/cache-response/src/lib.rs @@ -4,6 +4,7 @@ mod codec; mod embedding; mod exact; mod response; +mod service; pub use buffer::WriteBuffer; pub use caching::{ @@ -14,3 +15,8 @@ pub use codec::ResponseCacheCodec; pub use embedding::PartialHits; pub use exact::{ConnectionProbe, ExactResponseCache}; pub use response::{ResponseCache, ResponseCacheRequest}; + +pub use service::{ + CacheOptions, CacheScope, ResponseCacheConfig, ResponseCacheService, ResponseEnvelope, + ScopedCache, +}; diff --git a/litellm-rust/crates/cache-response/src/response.rs b/litellm-rust/crates/cache-response/src/response.rs index e761c7157db..a5bbef99a3a 100644 --- a/litellm-rust/crates/cache-response/src/response.rs +++ b/litellm-rust/crates/cache-response/src/response.rs @@ -7,7 +7,9 @@ use litellm_cache::{ }; use serde_json::Value; -use crate::{CacheControls, CacheEntry, CacheKeyInput, PartialHits, cache_key}; +use crate::{ + CacheControls, CacheEntry, CacheKeyInput, PartialHits, ResponseCacheConfig, cache_key, +}; #[derive(Clone)] pub struct ResponseCacheRequest { @@ -50,6 +52,7 @@ where B::Context: Default + PartialEq, { backend: Arc, + config: ResponseCacheConfig, } impl ResponseCache @@ -58,7 +61,18 @@ where B::Context: Default + PartialEq, { pub fn new(backend: Arc) -> Self { - Self { backend } + Self { + backend, + config: ResponseCacheConfig::default(), + } + } + + pub fn with_config(self, config: ResponseCacheConfig) -> Self { + Self { config, ..self } + } + + pub fn config(&self) -> &ResponseCacheConfig { + &self.config } pub fn backend(&self) -> &B { @@ -221,7 +235,7 @@ where response: Value, now: Duration, ) -> Result<(), Error> { - if !request.controls.writes() { + if !request.controls.writes() || !self.fits(&response) { return Ok(()); } self.backend.set_cache( @@ -240,7 +254,7 @@ where response: Value, now: Duration, ) -> Result<(), Error> { - if !request.controls.writes() { + if !request.controls.writes() || !self.fits(&response) { return Ok(()); } self.backend @@ -277,7 +291,7 @@ where ) -> Result<(), Error> { let writable = entries .into_iter() - .filter(|(request, _, _)| request.controls.writes()) + .filter(|(request, response, _)| request.controls.writes() && self.fits(response)) .map(|(request, response, now)| { ( cache_key(&request.key), @@ -312,6 +326,11 @@ where Ok(()) } + fn fits(&self, response: &Value) -> bool { + self.config.max_entry_bytes == usize::MAX + || response.to_string().len() <= self.config.max_entry_bytes + } + fn partial_hits( requests: &[ResponseCacheRequest], readable: Vec<(usize, &ResponseCacheRequest)>, diff --git a/litellm-rust/crates/cache-response/src/service.rs b/litellm-rust/crates/cache-response/src/service.rs new file mode 100644 index 00000000000..51359a9a8d6 --- /dev/null +++ b/litellm-rust/crates/cache-response/src/service.rs @@ -0,0 +1,177 @@ +use std::{future::Future, pin::Pin, time::Duration}; + +use litellm_cache::{BaseCache, Error, ExactCacheContext}; +use serde_json::Value; + +use crate::{ + CacheControls, CacheEntry, CacheKeyField, CacheKeyInput, ResponseCache, ResponseCacheRequest, +}; + +type CacheFuture<'a, T> = Pin> + Send + 'a>>; + +#[derive(Clone)] +pub struct ResponseCacheConfig { + pub namespace: String, + pub max_entry_bytes: usize, +} + +impl Default for ResponseCacheConfig { + fn default() -> Self { + Self { + namespace: String::new(), + max_entry_bytes: usize::MAX, + } + } +} + +pub trait ResponseCacheService: Send + Sync { + fn config(&self) -> &ResponseCacheConfig; + + fn lookup<'a>( + &'a self, + request: &'a ResponseCacheRequest, + now: Duration, + ) -> CacheFuture<'a, Option>; + + fn store<'a>( + &'a self, + request: &'a ResponseCacheRequest, + response: Value, + now: Duration, + ) -> CacheFuture<'a, ()>; +} + +impl ResponseCacheService for ResponseCache +where + B: BaseCache, +{ + fn config(&self) -> &ResponseCacheConfig { + self.config() + } + + fn lookup<'a>( + &'a self, + request: &'a ResponseCacheRequest, + now: Duration, + ) -> CacheFuture<'a, Option> { + Box::pin(self.async_lookup(request, now)) + } + + fn store<'a>( + &'a self, + request: &'a ResponseCacheRequest, + response: Value, + now: Duration, + ) -> CacheFuture<'a, ()> { + Box::pin(self.async_store(request, response, now)) + } +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum CacheScope { + Shared, + Isolated(String), +} + +#[derive(Clone)] +pub struct CacheOptions { + pub caching: Option, + pub no_cache: bool, + pub no_store: bool, + pub ttl: Option, + pub max_age: Option, + pub scope: CacheScope, +} + +impl CacheOptions { + pub fn new(scope: CacheScope) -> Self { + Self { + caching: None, + no_cache: false, + no_store: false, + ttl: None, + max_age: None, + scope, + } + } + + pub fn enabled(&self) -> bool { + self.caching != Some(false) && !(self.no_cache && self.no_store) + } + + pub fn request(self, namespace: &str, surface: &str, mut input: Value) -> ResponseCacheRequest { + input.sort_all_objects(); + let scope = match self.scope { + CacheScope::Shared => String::new(), + CacheScope::Isolated(scope) => serde_json::json!(["isolated", scope]).to_string(), + }; + ResponseCacheRequest { + key: CacheKeyInput { + namespace: Some(format!("{namespace}:inference-v2")), + fields: [ + ("surface", surface.to_owned()), + ("scope", scope), + ("request", input.to_string()), + ] + .into_iter() + .map(|(name, value)| CacheKeyField { + name: name.into(), + value: Some(value), + api_parameter: true, + internal_parameter: false, + }) + .collect(), + ..Default::default() + }, + controls: CacheControls { + configured: true, + supported_call_type: true, + native_backend: true, + default_on: true, + caching: self.caching, + no_cache: self.no_cache, + no_store: self.no_store, + ..Default::default() + }, + context: ExactCacheContext { ttl: self.ttl }, + max_age: self.max_age, + } + } +} + +#[derive(serde::Serialize, serde::Deserialize)] +pub struct ResponseEnvelope { + version: u32, + surface: String, + output: T, +} + +impl ResponseEnvelope { + pub fn new(surface: &str, output: T) -> Self { + Self { + version: 1, + surface: surface.into(), + output, + } + } + + pub fn decode(self, surface: &str) -> Option { + (self.version == 1 && self.surface == surface).then_some(self.output) + } +} + +#[derive(Clone)] +pub struct ScopedCache { + pub service: std::sync::Arc, + pub scope: CacheScope, +} + +impl ScopedCache { + pub fn new(service: std::sync::Arc, scope: CacheScope) -> Self { + Self { service, scope } + } + + pub fn options(&self, overrides: Option) -> CacheOptions { + overrides.unwrap_or_else(|| CacheOptions::new(self.scope.clone())) + } +} diff --git a/litellm-rust/crates/cache-response/tests/response.rs b/litellm-rust/crates/cache-response/tests/response.rs index ec5e16f1367..655fcb8a46a 100644 --- a/litellm-rust/crates/cache-response/tests/response.rs +++ b/litellm-rust/crates/cache-response/tests/response.rs @@ -18,7 +18,7 @@ use litellm_cache_response::{ WriteBuffer, cache_key, }; use redis_test::MockCmd; -use rstest::rstest; +use rstest::{fixture, rstest}; use serde_json::{Value, json}; use support::{keyed, memory, redis, request}; @@ -648,3 +648,129 @@ async fn write_buffer_clear_drops_pending_entries(memory: Memory, request: Respo assert_eq!(memory.lookup(&request, now).unwrap(), None); assert_eq!(memory.lookup(&other, now).unwrap(), None); } + +#[rstest] +#[case::python_sync("{'timestamp': 100.0, 'response': '{\"answer\": 7}'}")] +#[case::python_async(r#"{"timestamp":100.0,"response":{"answer":7}}"#)] +#[case::bare_response(r#"{"answer":7}"#)] +#[tokio::test] +async fn gcs_reads_python_entries_and_writes_python_compatible_envelopes( + #[case] encoded: &str, + #[values(false, true)] asynchronous: bool, + #[future(awt)] gcs: (wiremock::MockServer, Gcs), +) { + use wiremock::{ + Mock, ResponseTemplate, + matchers::{body_json, header, method, path, query_param}, + }; + + let (server, cache) = gcs; + let response = json!({"answer": 7}); + Mock::given(method("GET")) + .and(path("/storage/v1/b/bucket/o/cache%2Fpython")) + .and(query_param("alt", "media")) + .and(header("authorization", "Bearer token")) + .respond_with(ResponseTemplate::new(200).set_body_string(encoded)) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/upload/storage/v1/b/bucket/o")) + .and(query_param("uploadType", "media")) + .and(query_param("name", "cache/native")) + .and(header("authorization", "Bearer token")) + .and(header("content-type", "application/json")) + .and(body_json(json!({"timestamp": 102.0, "response": response}))) + .respond_with(ResponseTemplate::new(200)) + .expect(1) + .mount(&server) + .await; + let lookup = if asynchronous { + cache + .async_lookup(&keyed("python"), Duration::from_secs(102)) + .await + } else { + cache.lookup(&keyed("python"), Duration::from_secs(102)) + }; + assert_eq!(lookup.unwrap(), Some(response.clone())); + let request = ResponseCacheRequest { + context: litellm_cache::ExactCacheContext { + ttl: Some(Duration::from_secs(12)), + }, + ..keyed("native") + }; + let stored = if asynchronous { + cache + .async_store(&request, response, Duration::from_secs(102)) + .await + } else { + cache.store(&request, response, Duration::from_secs(102)) + }; + assert_eq!(stored, Ok(())); + let requests = server.received_requests().await.unwrap(); + let upload = requests + .iter() + .find(|request| request.method.as_str() == "POST") + .unwrap(); + assert_eq!( + upload.url.query(), + Some("uploadType=media&name=cache%2Fnative") + ); +} + +#[rstest] +#[tokio::test] +async fn gcs_batch_reads_preserve_order_and_treat_invalid_entries_as_misses( + #[future(awt)] gcs: (wiremock::MockServer, Gcs), +) { + use wiremock::{ + Mock, ResponseTemplate, + matchers::{method, path}, + }; + + let (server, cache) = gcs; + Mock::given(method("GET")) + .and(path("/storage/v1/b/bucket/o/cache%2Fhit")) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(json!({"timestamp": 100.0, "response": {"answer":7}})), + ) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/storage/v1/b/bucket/o/cache%2Finvalid")) + .respond_with(ResponseTemplate::new(200).set_body_string("not an entry")) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/storage/v1/b/bucket/o/cache%2Fmissing")) + .respond_with(ResponseTemplate::new(404)) + .mount(&server) + .await; + let requests = [keyed("hit"), keyed("missing"), keyed("invalid")]; + let partial = cache + .async_lookup_batch(&requests, Duration::from_secs(102)) + .await + .unwrap(); + assert_eq!(partial.values, vec![Some(json!({"answer":7})), None, None]); + assert_eq!(partial.missing_indices, vec![1, 2]); +} + +type Gcs = ResponseCache>; + +#[fixture] +async fn gcs() -> (wiremock::MockServer, Gcs) { + let server = wiremock::MockServer::start().await; + let cache = ResponseCache::new(Arc::new(litellm_cache_gcs::GcsCache::with_token_source( + litellm_cache_gcs::GcsConfig { + bucket_name: "bucket".into(), + gcs_path: Some("cache".into()), + path_service_account: None, + endpoint: server.uri(), + }, + litellm_http::Client::plain_for_test(), + litellm_cache_response::ResponseCacheCodec, + Arc::new(litellm_cache_gcs::StaticTokenSource("token".into())), + ))); + (server, cache) +} diff --git a/litellm-rust/crates/cache-response/tests/service.rs b/litellm-rust/crates/cache-response/tests/service.rs new file mode 100644 index 00000000000..d4532776719 --- /dev/null +++ b/litellm-rust/crates/cache-response/tests/service.rs @@ -0,0 +1,176 @@ +use std::{ + sync::{ + Arc, + atomic::{AtomicU64, Ordering}, + }, + time::Duration, +}; + +use litellm_cache::ExactCacheContext; +use litellm_cache_memory::InMemoryCache; +use litellm_cache_response::{ + CacheEntry, CacheKeyInput, ResponseCache, ResponseCacheConfig, ResponseCacheRequest, + ResponseCacheService, +}; +use rstest::rstest; +use serde_json::json; + +#[rstest] +#[tokio::test] +async fn service_honors_per_call_expiry_and_freshness() { + let clock = Arc::new(AtomicU64::new(0)); + let cache_clock = clock.clone(); + let cache: Arc = Arc::new(ResponseCache::new(Arc::new( + InMemoryCache::with_clock(Some(100), Some(Duration::from_secs(60)), move || { + Duration::from_secs(cache_clock.load(Ordering::SeqCst)) + }), + ))); + let request = ResponseCacheRequest { + context: ExactCacheContext { + ttl: Some(Duration::from_secs(5)), + }, + ..ResponseCacheRequest::new(CacheKeyInput { + preset: Some("entry".into()), + ..Default::default() + }) + }; + cache + .store(&request, json!({"answer":7}), Duration::ZERO) + .await + .unwrap(); + assert_eq!( + cache.lookup(&request, Duration::ZERO).await.unwrap(), + Some(json!({"answer":7})) + ); + let stale_request = ResponseCacheRequest { + max_age: Some(Duration::from_secs(1)), + ..request.clone() + }; + clock.store(2, Ordering::SeqCst); + assert_eq!( + cache + .lookup(&stale_request, Duration::from_secs(2)) + .await + .unwrap(), + None + ); + assert!( + cache + .lookup(&request, Duration::from_secs(2)) + .await + .unwrap() + .is_some() + ); + clock.store(6, Ordering::SeqCst); + assert_eq!( + cache + .lookup(&request, Duration::from_secs(6)) + .await + .unwrap(), + None + ); +} + +#[rstest] +#[tokio::test] +async fn entry_limit_applies_to_sync_async_and_batch_writes() { + let storage = Arc::new(InMemoryCache::::default()); + let cache = ResponseCache::new(storage.clone()).with_config(ResponseCacheConfig { + namespace: "service-test".into(), + max_entry_bytes: json!({"answer":7}).to_string().len(), + }); + let small = json!({"answer":7}); + let large = json!({"answer":"too large"}); + let request = |key: &str| { + ResponseCacheRequest::new(CacheKeyInput { + preset: Some(key.into()), + ..Default::default() + }) + }; + cache + .store(&request("sync"), large.clone(), Duration::ZERO) + .unwrap(); + cache + .async_store(&request("async"), large.clone(), Duration::ZERO) + .await + .unwrap(); + cache + .async_store_batch( + vec![ + (request("batch-large"), large), + (request("batch-small"), small.clone()), + ], + Duration::ZERO, + ) + .await + .unwrap(); + let service: Arc = Arc::new(cache); + service + .store(&request("service"), small.clone(), Duration::ZERO) + .await + .unwrap(); + for key in ["sync", "async", "batch-large"] { + assert!(storage.get_cache(key).unwrap().is_none()); + } + for key in ["batch-small", "service"] { + assert_eq!( + service.lookup(&request(key), Duration::ZERO).await.unwrap(), + Some(small.clone()) + ); + } +} + +#[rstest] +#[case::same_scope("tenant-a", "tenant-a", true)] +#[case::different_scope("tenant-a", "tenant-b", false)] +#[case::empty_isolated_scope("", "", true)] +#[tokio::test] +async fn isolated_policy_controls_actual_entry_reuse( + #[case] first: &str, + #[case] second: &str, + #[case] hit: bool, +) { + use litellm_cache_response::{CacheOptions, CacheScope}; + let service = ResponseCache::new(Arc::new(InMemoryCache::::default())); + let request = + |scope| CacheOptions::new(scope).request("test", "messages", json!({"prompt":"hello"})); + service + .async_store( + &request(CacheScope::Isolated(first.into())), + json!({"answer":7}), + Duration::ZERO, + ) + .await + .unwrap(); + assert_eq!( + service + .async_lookup( + &request(CacheScope::Isolated(second.into())), + Duration::ZERO + ) + .await + .unwrap(), + hit.then(|| json!({"answer":7})) + ); + assert_eq!( + service + .async_lookup(&request(CacheScope::Shared), Duration::ZERO) + .await + .unwrap(), + None + ); +} + +#[rstest] +#[case::valid(1, "messages", Some(7))] +#[case::unknown_version(2, "messages", None)] +#[case::another_surface(1, "responses", None)] +fn envelopes_require_a_matching_surface_and_version( + #[case] version: u32, + #[case] surface: &str, + #[case] expected: Option, +) { + let envelope: litellm_cache_response::ResponseEnvelope = + serde_json::from_value(json!({"version":version,"surface":surface,"output":7})).unwrap(); + assert_eq!(envelope.decode("messages"), expected); +} diff --git a/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs b/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs index 00d92168285..a2505588761 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs @@ -56,6 +56,7 @@ pub struct LegacyLogging { stream: Option, asynchronous: bool, internal: bool, + cache_key: Option, } fn datetime(py: Python<'_>, epoch_seconds: f64) -> PyResult> { @@ -80,6 +81,7 @@ impl LegacyLogging { stream: None, asynchronous, internal: false, + cache_key: None, } } @@ -207,10 +209,10 @@ impl LegacyLogging { logger.object(py), billing.url_route, billing.endpoint_type, - &self - .request - .as_ref() - .map(|request| request.body.clone_ref(py)), + &self.request.as_ref().map_or_else( + || self.call.kwargs().clone_ref(py), + |request| request.body.clone_ref(py), + ), &stream.chunks, &self.start, &self.end, @@ -246,10 +248,10 @@ impl LegacyLogging { ( logger.object(py), billing.endpoint_type, - &self - .request - .as_ref() - .map(|request| request.body.clone_ref(py)), + &self.request.as_ref().map_or_else( + || self.call.kwargs().clone_ref(py), + |request| request.body.clone_ref(py), + ), &stream.chunks, error, ), @@ -428,6 +430,40 @@ impl LegacyLogging { self.finalize(py) } + pub(crate) fn result_ready( + &mut self, + py: Python<'_>, + facts: &litellm_host::interceptors::ExecutionFacts, + ) -> PyResult> { + use litellm_host::interceptors::ResultSource; + + let logger = self.logger()?.object(py); + let params = logger + .getattr("litellm_params")? + .cast_into::()? + .copy()?; + params.set_item("custom_llm_provider", &facts.provider.provider)?; + crate::python::Logging::Update.call( + py, + ( + &logger, + self.call.kwargs(), + &facts.provider.model, + logger.getattr("optional_params")?, + params, + &facts.provider.provider, + ), + )?; + let details = logger.getattr("model_call_details")?; + self.cache_key = match &facts.source { + ResultSource::Provider => None, + ResultSource::Cache { key } => Some(key.clone()), + }; + details.set_item("cache_hit", self.cache_key.is_some())?; + details.set_item("cache_key", self.cache_key.as_deref())?; + Ok(HookStep::Ready(())) + } + pub(crate) fn post_call( &mut self, py: Python<'_>, @@ -485,10 +521,14 @@ impl LegacyLogging { self.dispatch_failure(py) } - pub(crate) fn stream_opened(&mut self, py: Python<'_>) -> PyResult<()> { + pub(crate) fn stream_opened(&mut self, py: Python<'_>, head: &Py) -> PyResult<()> { if self.stream_billing().is_none() { return Err(missing_state()); } + if let Some(key) = &self.cache_key { + head.bind(py).set_item("cache_key", key)?; + head.bind(py).set_item("cache_hit", true)?; + } Streaming::Opened.call(py, (self.logger()?.object(py),))?; self.stream = Some(DeliveredStream { chunks: PyList::empty(py).unbind(), @@ -1726,7 +1766,9 @@ assert logger.calls[1][1] is response operation: litellm_types::Operation::Messages, ..logged(py, &locals, true) }; - logging.on_stream_open(py).unwrap(); + logging + .on_stream_open(py, &pyo3::types::PyDict::new(py).into_any().unbind()) + .unwrap(); logging .on_stream_chunk(py, &local(&locals, "first").unbind()) .unwrap(); diff --git a/litellm-rust/crates/callbacks-legacy-python/src/mapping.rs b/litellm-rust/crates/callbacks-legacy-python/src/mapping.rs index 8321e63e196..a4aee4eb4da 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/mapping.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/mapping.rs @@ -56,7 +56,7 @@ type After = fn(&mut LegacyLogging, Python<'_>, &RawResponse) -> Step<()>; type Transform = fn(&mut LegacyLogging, Python<'_>, Py, Timing) -> Step>; type Success = fn(&mut LegacyLogging, Python<'_>, Timing, &Py) -> Step<()>; type Failure = fn(&mut LegacyLogging, Python<'_>, Timing, FailureOrigin, &PyErr) -> Step<()>; -type Open = fn(&mut LegacyLogging, Python<'_>) -> PyResult<()>; +type Open = fn(&mut LegacyLogging, Python<'_>, &Py) -> PyResult<()>; type Chunk = fn(&mut LegacyLogging, Python<'_>, &Py) -> PyResult<()>; const PREPARE: Binding = Binding { @@ -173,6 +173,9 @@ impl CallHooks for LegacyLogging { PythonCallEvent::Started { .. } | PythonCallEvent::Cancelled { .. } => { Ok(HookStep::Ready(())) } + PythonCallEvent::Execution(ExecutionEvent::ResultReady { facts }) => { + self.result_ready(py, &facts) + } PythonCallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => { (AFTER.invoke)(self, py, raw) } @@ -187,8 +190,8 @@ impl CallHooks for LegacyLogging { } } - fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> { - (OPEN.invoke)(self, py) + fn on_stream_open(&mut self, py: Python<'_>, head: &Py) -> PyResult<()> { + (OPEN.invoke)(self, py, head) } fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py) -> PyResult<()> { diff --git a/litellm-rust/crates/core/AGENTS.md b/litellm-rust/crates/core/AGENTS.md index f316e6f7799..217fdfc5e11 100644 --- a/litellm-rust/crates/core/AGENTS.md +++ b/litellm-rust/crates/core/AGENTS.md @@ -33,3 +33,15 @@ Scope follows the concept, not the first caller. An error type under `litellm-ll `litellm_llms::Error` (`crates/llms/src/error.rs`) is the one transformation error for every provider and API. `base_llm/ocr/error.rs` is the recorded exception until OCR folds into it Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback execution of any kind. Core runs each route as a machine that yields host operations and call events; which integrations consume those events is the host's business. + +## Response caching and accounting boundary + +Attach a `litellm_cache_response::ScopedCache` with `route.with_cache(cache)`. Cached and uncached routes use the same `execute` and `machine` methods. `CallOptions` carries per-call cache overrides and observation; attaching a service does not change the execution contract + +Core owns request identity, typed response reconstruction and stream capture/replay. `cache-response` owns cache policy, namespacing, scope encoding, versioned envelopes and freshness. The SDK explicitly chooses shared scope. The gateway derives isolated scope from authenticated identity before attaching its service + +Core delivers `ExecutionFacts` through the awaited `ResultReady` host operation for both provider and cached results, before public response processing or stream opening. Facts carry resolved model/provider and result source, including the hit key. Usage remains in the typed response or delivered stream, where completion and cancellation determine what was actually reported. Passive observation is not an accounting delivery mechanism + +Core does not calculate prices, charge budgets, or update rate-limit counters. The legacy Python callback adapter translates execution facts into the existing Python logging contract; Python remains the accounting owner on that path. Native gateway accounting belongs to gateway dependencies, independently of `host-python`. Response-cache services expose no coordination counters or reservation APIs. A shared Redis deployment does not make response storage and accounting coordination the same dependency + +Cache lookup follows provider preparation, credential resolution and the request interceptor. Keys describe the effective provider URL, authenticated headers and rewritten body. Signed requests bypass caching until the signing identity has a stable cache representation diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index 023267d56ef..56f9a0c4163 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -6,6 +6,10 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-cache.workspace = true +litellm-cache-response.workspace = true +litellm-framing.workspace = true +tokio-util = { version = "0.7", features = ["codec"] } litellm-secrets.workspace = true litellm-types.workspace = true litellm-core-utils.workspace = true @@ -36,6 +40,7 @@ url.workspace = true veil.workspace = true [dev-dependencies] +litellm-cache-memory.workspace = true litellm-http = { workspace = true, features = ["test-support"] } litellm-auth-gcp.workspace = true litellm-host-native.workspace = true diff --git a/litellm-rust/crates/core/src/caching.rs b/litellm-rust/crates/core/src/caching.rs new file mode 100644 index 00000000000..d182ba94543 --- /dev/null +++ b/litellm-rust/crates/core/src/caching.rs @@ -0,0 +1,318 @@ +use std::{ + future::Future, + sync::Arc, + time::{Duration, SystemTime, UNIX_EPOCH}, +}; + +use bytes::{Bytes, BytesMut}; +use futures_util::{StreamExt, TryStreamExt, stream}; +use litellm_cache_response::{ + CacheOptions, ResponseCacheRequest, ResponseCacheService, ResponseEnvelope, cache_key, +}; +use litellm_host::{ + call::{CallOutput, OutputOf}, + interceptors::{ExecutionFacts, Interceptors, ProviderIdentity, ResultSource, WireRequest}, + lifecycle::{CallEvent, ExecutionEvent}, + observation::ObservationSender, + protocol::Protocol, +}; +use serde::{Deserialize, Serialize, de::DeserializeOwned}; +use serde_json::Value; +use tokio_util::codec::Decoder; + +use crate::RouteError; + +pub trait Cachable: Protocol { + const SURFACE: &'static str; + + fn reusable(_response: &Self::Response) -> bool { + true + } +} + +pub struct CacheRequest { + pub identity: ProviderIdentity, + pub input: Value, +} + +impl CacheRequest { + pub fn from_wire(identity: ProviderIdentity, wire: Option<&WireRequest>) -> Self { + Self { + input: wire.map_or(Value::Null, |wire| { + serde_json::json!({ + "provider": identity.provider, + "model": identity.model, + "url": wire.url, + "headers": wire.headers, + "body": wire.body, + }) + }), + identity, + } + } +} + +pub trait StreamCachable: Cachable { + const TERMINAL_EVENT: &'static str; + + fn replay(data: Bytes) -> Option>; + fn bytes(chunk: &Self::Chunk) -> &[u8]; +} + +#[derive(Serialize, Deserialize)] +#[serde(tag = "kind", content = "value")] +pub enum CachedOutput { + Response(R), + Stream(String), +} + +struct CacheSession { + service: Arc, + request: ResponseCacheRequest, +} + +impl CacheSession { + fn prepare( + service: Option>, + options: Option, + request: &CacheRequest, + ) -> Option { + let options = options.filter(CacheOptions::enabled)?; + let service = service?; + let input = request.input.clone(); + let request = options.request(&service.config().namespace, P::SURFACE, input); + Some(Self { service, request }) + } + + async fn lookup(&self) -> Option> + where + P::Response: DeserializeOwned, + { + if !self.request.controls.reads() { + return None; + } + match self.service.lookup(&self.request, now()).await { + Ok(Some(value)) => { + serde_json::from_value::>>(value) + .ok() + .and_then(|entry| entry.decode(P::SURFACE)) + } + Ok(None) => None, + Err(_) => { + tracing::warn!("response cache lookup failed"); + None + } + } + } + + async fn store(&self, entry: Value) { + if !self.request.controls.writes() { + return; + } + if self + .service + .store(&self.request, entry, now()) + .await + .is_err() + { + tracing::warn!("response cache write failed"); + } + } + + async fn store_response(&self, response: &P::Response) + where + P::Response: Serialize, + { + if !self.request.controls.writes() || !P::reusable(response) { + return; + } + if let Ok(value) = serde_json::to_value(response) + && let Ok(entry) = serde_json::to_value(ResponseEnvelope::new( + P::SURFACE, + CachedOutput::Response(value), + )) + { + self.store(entry).await; + } + } +} + +pub async fn execute_unary( + request: CacheRequest, + cache: Option>, + options: Option, + interceptors: &impl Interceptors, + observers: Option<&ObservationSender>, + provider: F, +) -> Result +where + P: Cachable, + P::Response: Serialize + DeserializeOwned, + F: FnOnce() -> Fut, + Fut: Future>, +{ + let identity = request.identity.clone(); + crate::diagnostic::provider(&identity.model, &identity.provider); + let session = CacheSession::prepare::

(cache, options, &request); + let hit = match &session { + Some(session) => session.lookup::

().await.and_then(|entry| match entry { + CachedOutput::Response(response) => Some((response, cache_key(&session.request.key))), + CachedOutput::Stream(_) => None, + }), + None => None, + }; + let (response, source) = match hit { + Some((response, key)) => (response, ResultSource::Cache { key }), + None => (provider().await?, ResultSource::Provider), + }; + let from_provider = source == ResultSource::Provider; + publish( + ExecutionFacts { + provider: identity, + source, + }, + interceptors, + observers, + ) + .await?; + if from_provider && let Some(session) = session { + session.store_response::

(&response).await; + } + Ok(response) +} + +pub async fn execute_streaming( + request: CacheRequest, + cache: Option>, + options: Option, + interceptors: &impl Interceptors, + observers: Option<&ObservationSender>, + provider: F, +) -> Result, RouteError> +where + P: StreamCachable, + P::Response: Serialize + DeserializeOwned, + F: FnOnce() -> Fut, + Fut: Future, RouteError>>, +{ + let identity = request.identity.clone(); + crate::diagnostic::provider(&identity.model, &identity.provider); + let session = CacheSession::prepare::

(cache, options, &request); + let hit = match &session { + Some(session) => session.lookup::

().await.and_then(|entry| { + let output = match entry { + CachedOutput::Response(response) => Some(CallOutput::Complete(response)), + CachedOutput::Stream(data) => P::replay(Bytes::from(data)), + }; + output.map(|output| (output, cache_key(&session.request.key))) + }), + None => None, + }; + let (output, source) = match hit { + Some((output, key)) => (output, ResultSource::Cache { key }), + None => (provider().await?, ResultSource::Provider), + }; + let from_provider = source == ResultSource::Provider; + publish( + ExecutionFacts { + provider: identity, + source, + }, + interceptors, + observers, + ) + .await?; + let Some(session) = + session.filter(|session| from_provider && session.request.controls.writes()) + else { + return Ok(output); + }; + match output { + CallOutput::Complete(response) => { + session.store_response::

(&response).await; + Ok(CallOutput::Complete(response)) + } + CallOutput::Stream { head, chunks } => { + let captured = stream::try_unfold( + (chunks, Some(Vec::::new()), session), + |(mut chunks, captured, session)| async move { + match chunks.try_next().await? { + Some(chunk) => { + let captured = captured.and_then(|mut data| { + let bytes = P::bytes(&chunk); + if data.len().saturating_add(bytes.len()) + > session.service.config().max_entry_bytes + { + return None; + } + data.extend_from_slice(bytes); + Some(data) + }); + Ok(Some((chunk, (chunks, captured, session)))) + } + None => { + if let Some(data) = captured + && let Ok(text) = String::from_utf8(data) + && successful_stream(&text, P::TERMINAL_EVENT) + && let Ok(entry) = serde_json::to_value(ResponseEnvelope::new( + P::SURFACE, + CachedOutput::::Stream(text), + )) + { + session.store(entry).await; + } + Ok::<_, RouteError>(None) + } + } + }, + ) + .boxed(); + Ok(CallOutput::Stream { + head, + chunks: captured, + }) + } + } +} + +fn now() -> Duration { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() +} + +fn successful_stream(text: &str, terminal: &str) -> bool { + let mut pending = BytesMut::from(text.as_bytes()); + let mut codec = litellm_framing::sse::SseCodec::default(); + let mut complete = false; + loop { + let event = match codec.decode(&mut pending) { + Ok(Some(event)) => event, + Ok(None) => return complete && pending.is_empty(), + Err(_) => return false, + }; + let Ok(value) = serde_json::from_str::(&event.data) else { + return false; + }; + let Some(kind) = value.get("type").and_then(Value::as_str) else { + return false; + }; + if matches!(kind, "error" | "response.failed" | "response.incomplete") { + return false; + } + complete |= kind == terminal; + } +} + +async fn publish( + facts: ExecutionFacts, + interceptors: &impl Interceptors, + observers: Option<&ObservationSender>, +) -> Result<(), RouteError> { + if let Some(observers) = observers { + observers.emit(CallEvent::Execution(ExecutionEvent::ResultReady { + facts: facts.clone(), + })); + } + interceptors.result_ready(facts).await +} diff --git a/litellm-rust/crates/core/src/chat_completions/handler.rs b/litellm-rust/crates/core/src/chat_completions/handler.rs index dd54c3057f1..8cfee9b59cf 100644 --- a/litellm-rust/crates/core/src/chat_completions/handler.rs +++ b/litellm-rust/crates/core/src/chat_completions/handler.rs @@ -22,6 +22,8 @@ pub(super) async fn execute( http: &Client, auth: &AuthServices, request: ProviderChatCompletionsRequest, + cache: Option, + cache_options: Option, interceptors: &impl Interceptors, observers: Option<&ObservationSender>, ) -> Result { @@ -45,6 +47,10 @@ pub(super) async fn execute( api_key, }; let authenticated = resolve_auth(auth, environment, &|key| secrets.get(key)).await?; + let identity = litellm_host::interceptors::ProviderIdentity { + model: context.model.clone(), + provider: context.custom_llm_provider.clone(), + }; let wire = interceptors .before_provider_request( WireRequest { @@ -55,59 +61,72 @@ pub(super) async fn execute( context, ) .await?; - let outbound = outbound_request( - Authenticated { - headers: wire.headers, - signer: authenticated.signer, + let cache = cache.filter(|_| authenticated.signer.is_none()); + let cache_request = + crate::caching::CacheRequest::from_wire(identity, cache.as_ref().map(|_| &wire)); + crate::caching::execute_unary::( + cache_request, + cache.as_ref().map(|cache| cache.service.clone()), + cache.as_ref().map(|cache| cache.options(cache_options)), + interceptors, + observers, + || async move { + let outbound = outbound_request( + Authenticated { + headers: wire.headers, + signer: authenticated.signer, + }, + wire.url, + &wire.body, + timeout, + )?; + + let response = crate::outbound::send(outbound, http).await.map_err(|err| { + // Failing to establish the connection means the request never went out, + // so the host can still serve it. Everything else here, a timeout + // above all, may have reached the provider and been answered. + if err.is_connect() || err.is_builder() { + Error::Transport(litellm_http::transport::Error::Connect(err.to_string())) + } else { + Error::Transport(litellm_http::transport::Error::Network(err.to_string())) + } + })?; + + let status = response.status(); + let text = response.text().await.map_err(|err| { + Error::Transport(litellm_http::transport::Error::Network(err.to_string())) + })?; + + if !status.is_success() { + return Err(Error::Transport(litellm_http::transport::Error::Http { + status: status.as_u16(), + body: truncate_error_body(&text), + })); + } + let raw = RawResponse { body: text.clone() }; + if let Some(observers) = observers { + observers.emit(litellm_host::lifecycle::CallEvent::Execution( + ExecutionEvent::ProviderResponseReceived { raw: raw.clone() }, + )); + } + interceptors + .after_provider_response(raw) + .await + .map_err(Error::post_call)?; + + let body: Value = serde_json::from_str(&text).map_err(|err| { + Error::InvalidResponse(litellm_llms::ErrorDetail::invalid( + "chat completions response JSON", + err, + )) + })?; + config + .transform_response(&model, ProviderChatResponseData { body }) + .map_err(Error::from) + .map_err(as_response_error) }, - wire.url, - &wire.body, - timeout, - )?; - - let response = crate::outbound::send(outbound, http).await.map_err(|err| { - // Failing to establish the connection means the request never went out, - // so the host can still serve it. Everything else here, a timeout - // above all, may have reached the provider and been answered. - if err.is_connect() || err.is_builder() { - Error::Transport(litellm_http::transport::Error::Connect(err.to_string())) - } else { - Error::Transport(litellm_http::transport::Error::Network(err.to_string())) - } - })?; - - let status = response.status(); - let text = response.text().await.map_err(|err| { - Error::Transport(litellm_http::transport::Error::Network(err.to_string())) - })?; - - if !status.is_success() { - return Err(Error::Transport(litellm_http::transport::Error::Http { - status: status.as_u16(), - body: truncate_error_body(&text), - })); - } - let raw = RawResponse { body: text.clone() }; - if let Some(observers) = observers { - observers.emit(litellm_host::lifecycle::CallEvent::Execution( - ExecutionEvent::ProviderResponseReceived { raw: raw.clone() }, - )); - } - interceptors - .after_provider_response(raw) - .await - .map_err(Error::post_call)?; - - let body: Value = serde_json::from_str(&text).map_err(|err| { - Error::InvalidResponse(litellm_llms::ErrorDetail::invalid( - "chat completions response JSON", - err, - )) - })?; - config - .transform_response(&model, ProviderChatResponseData { body }) - .map_err(Error::from) - .map_err(as_response_error) + ) + .await } /// Re-tag an error raised while normalizing a response the provider already @@ -232,6 +251,8 @@ mod tests { &Client::plain_for_test(), &AuthServices::default(), prepared(&upstream.uri()), + None, + None, &interceptors, None, ) @@ -271,6 +292,8 @@ mod tests { &Client::plain_for_test(), &AuthServices::default(), prepared(&upstream.uri()), + None, + None, &interceptors, None, ) diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index 9e350e757d3..a64249fa185 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -18,6 +18,7 @@ pub struct ChatCompletionsRoute { http: litellm_http::Client, auth: Arc, secrets: Arc, + cache: Option, } impl ChatCompletionsRoute { @@ -30,6 +31,14 @@ impl ChatCompletionsRoute { http, auth, secrets, + cache: None, + } + } + + pub fn with_cache(self, cache: litellm_cache_response::ScopedCache) -> Self { + Self { + cache: Some(cache), + ..self } } @@ -37,49 +46,48 @@ impl ChatCompletionsRoute { &self, request: ChatCompletionsRequest<'_>, interceptors: &impl litellm_host::interceptors::Interceptors, - observers: Option, + options: impl Into, ) -> Result { + let crate::CallOptions { + cache: cache_options, + observers, + } = options.into(); litellm_host::lifecycle::observe_unary( observers.clone(), - self.run(request, interceptors, observers.as_ref()), + self.run_call( + request.into(), + cache_options, + interceptors, + observers.as_ref(), + ), ) .await } - #[tracing::instrument(name = "litellm.route", skip_all, fields( - route = "chat_completions", - model = %request.model, - provider, - resolved_model, - stream = false, - outcome - ))] async fn run( &self, request: ChatCompletionsRequest<'_>, + cache_options: Option, interceptors: &impl litellm_host::interceptors::Interceptors, observers: Option<&ObservationSender>, ) -> Result { - crate::diagnostic::unary(async { - let resolved = resolve_request(request)?; - let snapshot = self - .secrets - .resolve(&resolved.config.secret_names()) - .await?; - let prepared = prepare_provider_request(resolved, snapshot)?; - crate::diagnostic::provider(&prepared.model, &prepared.custom_llm_provider); - let execute: futures_util::future::BoxFuture< - '_, - Result, - > = Box::pin(handler::execute( + let resolved = resolve_request(request)?; + let snapshot = self + .secrets + .resolve(&resolved.config.secret_names()) + .await?; + let prepared = prepare_provider_request(resolved, snapshot)?; + crate::diagnostic::provider(&prepared.model, &prepared.custom_llm_provider); + let execute: futures_util::future::BoxFuture<'_, Result> = + Box::pin(handler::execute( &self.http, &self.auth, prepared, + self.cache.clone(), + cache_options, interceptors, observers, )); - execute.await - }) - .await + execute.await } } diff --git a/litellm-rust/crates/core/src/chat_completions/route.rs b/litellm-rust/crates/core/src/chat_completions/route.rs index 0816850bcc1..d43d2bf9eef 100644 --- a/litellm-rust/crates/core/src/chat_completions/route.rs +++ b/litellm-rust/crates/core/src/chat_completions/route.rs @@ -27,26 +27,56 @@ impl ChatCompletionsRoute { pub fn machine( self, call: ChatCompletionsCall, - observers: Option, + options: impl Into, ) -> HostedMachine { + let crate::CallOptions { + cache: cache_options, + observers, + } = options.into(); hosted_call( call, observers, - move |call: ChatCompletionsCall, _, interceptors, observers| async move { - let request = ChatCompletionsRequest { - model: &call.model, - messages: call.messages, - optional_params: call.optional_params, - api_key: call.api_key.as_deref(), - api_base: call.api_base.as_deref(), - custom_llm_provider: call.custom_llm_provider.as_deref(), - extra_headers: call.extra_headers, - timeout: call.timeout, - }; - self.run(request, &interceptors, observers.as_ref()) + move |call, _, interceptors, observers| async move { + self.run_call(call, cache_options, &interceptors, observers.as_ref()) .await .map(CallOutput::Complete) }, ) } + + #[tracing::instrument(name = "litellm.route", skip_all, fields( + route = "chat_completions", + model = %call.model, + provider, + resolved_model, + stream = false, + outcome + ))] + pub(super) async fn run_call( + &self, + call: ChatCompletionsCall, + cache_options: Option, + interceptors: &impl litellm_host::interceptors::Interceptors, + observers: Option<&ObservationSender>, + ) -> Result { + crate::diagnostic::unary(async { + let request = ChatCompletionsRequest { + model: &call.model, + messages: call.messages, + optional_params: call.optional_params, + api_key: call.api_key.as_deref(), + api_base: call.api_base.as_deref(), + custom_llm_provider: call.custom_llm_provider.as_deref(), + extra_headers: call.extra_headers, + timeout: call.timeout, + }; + self.run(request, cache_options, interceptors, observers) + .await + }) + .await + } +} + +impl crate::caching::Cachable for ChatCompletions { + const SURFACE: &'static str = "chat_completions"; } diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index fe487f41544..dbdfc63e929 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -1,6 +1,7 @@ mod diagnostic; pub mod audio_transcription; +pub mod caching; pub mod chat_completions; pub mod constants; pub mod error; @@ -12,3 +13,27 @@ pub mod resources; pub mod responses; pub use error::RouteError; + +#[derive(Clone, Default)] +pub struct CallOptions { + pub cache: Option, + pub observers: Option, +} + +impl From> for CallOptions { + fn from(observers: Option) -> Self { + Self { + cache: None, + observers, + } + } +} + +impl From for CallOptions { + fn from(cache: litellm_cache_response::CacheOptions) -> Self { + Self { + cache: Some(cache), + observers: None, + } + } +} diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index d061f456b2a..4d379e27ca4 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -27,6 +27,8 @@ pub(super) async fn execute( http: &litellm_http::Client, auth: &AuthServices, request: ProviderMessagesRequest, + cache: Option, + cache_options: Option, interceptors: &impl Interceptors, observers: Option<&ObservationSender>, ) -> Result { @@ -47,6 +49,10 @@ pub(super) async fn execute( api_key, }; let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?; + let identity = litellm_host::interceptors::ProviderIdentity { + model: context.model.clone(), + provider: context.custom_llm_provider.clone(), + }; let wire = interceptors .before_provider_request( WireRequest { @@ -57,44 +63,57 @@ pub(super) async fn execute( context, ) .await?; - let provider_name = provider.as_str(); - log_request_body(provider_name, stream, &wire.body); - let response = send( - http, - Authenticated { - headers: wire.headers, - signer: authenticated.signer, + let cache = cache.filter(|_| authenticated.signer.is_none()); + let cache_request = + crate::caching::CacheRequest::from_wire(identity, cache.as_ref().map(|_| &wire)); + crate::caching::execute_streaming::( + cache_request, + cache.as_ref().map(|cache| cache.service.clone()), + cache.as_ref().map(|cache| cache.options(cache_options)), + interceptors, + observers, + || async move { + let provider_name = provider.as_str(); + log_request_body(provider_name, stream, &wire.body); + let response = send( + http, + Authenticated { + headers: wire.headers, + signer: authenticated.signer, + }, + &wire.url, + &wire.body, + timeout, + ) + .await?; + if !response.status().is_success() { + return Err(provider_error(response).await); + } + let config = provider.config(); + if stream { + return Ok(streaming_response( + response, + config.stream_decoder(), + provider_name, + )); + } + let text = response.text().await.map_err(network)?; + log_response_body(&text); + let raw = RawResponse { body: text.clone() }; + if let Some(observers) = observers { + observers.emit(litellm_host::lifecycle::CallEvent::Execution( + ExecutionEvent::ProviderResponseReceived { raw: raw.clone() }, + )); + } + interceptors + .after_provider_response(raw) + .await + .map_err(Error::post_call)?; + decode_response(config, &body.model, &text) + .map(|message| MessagesResponse::Complete(Box::new(message))) }, - &wire.url, - &wire.body, - timeout, ) - .await?; - if !response.status().is_success() { - return Err(provider_error(response).await); - } - let config = provider.config(); - if stream { - return Ok(streaming_response( - response, - config.stream_decoder(), - provider_name, - )); - } - let text = response.text().await.map_err(network)?; - log_response_body(&text); - let raw = RawResponse { body: text.clone() }; - if let Some(observers) = observers { - observers.emit(litellm_host::lifecycle::CallEvent::Execution( - ExecutionEvent::ProviderResponseReceived { raw: raw.clone() }, - )); - } - interceptors - .after_provider_response(raw) - .await - .map_err(Error::post_call)?; - decode_response(config, &body.model, &text) - .map(|message| MessagesResponse::Complete(Box::new(message))) + .await } fn serialize_failure(err: serde_json::Error) -> Error { diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index 07f1fff8d45..23e0a3fb624 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -17,18 +17,84 @@ pub struct MessagesRoute { http: litellm_http::Client, auth: Arc, secrets: Arc, + cache: Option, +} + +#[must_use] +#[derive(Clone, Default)] +pub struct MessagesRouteBuilder { + http: Http, + auth: Auth, + secrets: Secrets, + cache: Option, +} + +impl MessagesRouteBuilder { + pub fn with_http( + self, + http: litellm_http::Client, + ) -> MessagesRouteBuilder { + MessagesRouteBuilder { + http, + auth: self.auth, + secrets: self.secrets, + cache: self.cache, + } + } + + pub fn with_auth( + self, + auth: Arc, + ) -> MessagesRouteBuilder, Secrets> { + MessagesRouteBuilder { + http: self.http, + auth, + secrets: self.secrets, + cache: self.cache, + } + } + + pub fn with_secrets( + self, + secrets: Arc, + ) -> MessagesRouteBuilder> { + MessagesRouteBuilder { + http: self.http, + auth: self.auth, + secrets, + cache: self.cache, + } + } + + pub fn with_cache(self, cache: litellm_cache_response::ScopedCache) -> Self { + Self { + cache: Some(cache), + ..self + } + } +} + +impl MessagesRouteBuilder, Arc> { + pub fn build(self) -> MessagesRoute { + MessagesRoute { + http: self.http, + auth: self.auth, + secrets: self.secrets, + cache: self.cache, + } + } } impl MessagesRoute { - pub fn new( - http: litellm_http::Client, - auth: Arc, - secrets: Arc, - ) -> Self { + pub fn builder() -> MessagesRouteBuilder { + MessagesRouteBuilder::default() + } + + #[must_use] + pub fn with_cache(self, cache: litellm_cache_response::ScopedCache) -> Self { Self { - http, - auth, - secrets, + cache: Some(cache), + ..self } } @@ -36,11 +102,15 @@ impl MessagesRoute { &self, call: MessagesCall, interceptors: &impl litellm_host::interceptors::Interceptors, - observers: Option, + options: impl Into, ) -> Result { + let crate::CallOptions { + cache: cache_options, + observers, + } = options.into(); litellm_host::lifecycle::observe_call( observers.clone(), - self.run(call, interceptors, observers.as_ref()), + self.run(call, cache_options, interceptors, observers.as_ref()), ) .await } @@ -56,22 +126,36 @@ impl MessagesRoute { async fn run( &self, call: MessagesCall, + cache_options: Option, interceptors: &impl litellm_host::interceptors::Interceptors, observers: Option<&ObservationSender>, ) -> Result { crate::diagnostic::call(async { - let request = prepare::prepare(call, self.secrets.as_ref()).await?; - crate::diagnostic::provider(&request.body.model, request.provider.as_str()); - let execute: futures_util::future::BoxFuture<'_, Result> = - Box::pin(handler::execute( - &self.http, - &self.auth, - request, - interceptors, - observers, - )); - execute.await + self.run_provider(call, cache_options, interceptors, observers) + .await }) .await } + + async fn run_provider( + &self, + call: MessagesCall, + cache_options: Option, + interceptors: &impl litellm_host::interceptors::Interceptors, + observers: Option<&ObservationSender>, + ) -> Result { + let request = prepare::prepare(call, self.secrets.as_ref()).await?; + crate::diagnostic::provider(&request.body.model, request.provider.as_str()); + let execute: futures_util::future::BoxFuture<'_, Result> = + Box::pin(handler::execute( + &self.http, + &self.auth, + request, + self.cache.clone(), + cache_options, + interceptors, + observers, + )); + execute.await + } } diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 954c31be9d8..1d2f95da957 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -1,4 +1,3 @@ -use litellm_host::observation::ObservationSender; use std::convert::Infallible; use bytes::Bytes; @@ -34,14 +33,40 @@ impl super::MessagesRoute { pub fn machine( self, request: super::MessagesCall, - observers: Option, + options: impl Into, ) -> MessagesMachine { + let crate::CallOptions { + cache: cache_options, + observers, + } = options.into(); hosted_call( request, observers, move |call, _, interceptors, observers| async move { - self.run(call, &interceptors, observers.as_ref()).await + self.run(call, cache_options, &interceptors, observers.as_ref()) + .await }, ) } } + +impl crate::caching::Cachable for Messages { + const SURFACE: &'static str = "messages"; +} + +impl crate::caching::StreamCachable for Messages { + const TERMINAL_EVENT: &'static str = "message_stop"; + + fn replay(data: bytes::Bytes) -> Option> { + Some(litellm_host::call::CallOutput::Stream { + head: MessagesStreamHead { + headers: Vec::new(), + }, + chunks: Box::pin(futures_util::stream::iter([Ok(data)])), + }) + } + + fn bytes(chunk: &Self::Chunk) -> &[u8] { + chunk.as_ref() + } +} diff --git a/litellm-rust/crates/core/src/responses/handler.rs b/litellm-rust/crates/core/src/responses/handler.rs index 71b3b268d74..b90e6af594a 100644 --- a/litellm-rust/crates/core/src/responses/handler.rs +++ b/litellm-rust/crates/core/src/responses/handler.rs @@ -15,10 +15,16 @@ pub(super) async fn execute( http: &litellm_http::Client, auth: &litellm_auth::AuthServices, request: ProviderResponsesRequest, + cache: Option, + cache_options: Option, interceptors: &impl Interceptors, observers: Option<&ObservationSender>, ) -> Result { let authenticated = resolve_auth(auth, request.environment, &|_| None).await?; + let identity = litellm_host::interceptors::ProviderIdentity { + model: request.context.model.clone(), + provider: request.context.custom_llm_provider.clone(), + }; let wire = interceptors .before_provider_request( WireRequest { @@ -29,65 +35,80 @@ pub(super) async fn execute( request.context, ) .await?; - let stream = match wire.body.get("stream") { - None => false, - Some(serde_json::Value::Bool(value)) => *value, - Some(_) => return Err(Error::InvalidRequest("stream must be a boolean".into())), - }; - let outbound = crate::outbound::outbound_request( - Authenticated { - headers: wire.headers, - signer: authenticated.signer, + let cache = cache.filter(|_| authenticated.signer.is_none()); + let cache_request = + crate::caching::CacheRequest::from_wire(identity, cache.as_ref().map(|_| &wire)); + crate::caching::execute_streaming::( + cache_request, + cache.as_ref().map(|cache| cache.service.clone()), + cache.as_ref().map(|cache| cache.options(cache_options)), + interceptors, + observers, + || async move { + let stream = match wire.body.get("stream") { + None => false, + Some(serde_json::Value::Bool(value)) => *value, + Some(_) => return Err(Error::InvalidRequest("stream must be a boolean".into())), + }; + let outbound = crate::outbound::outbound_request( + Authenticated { + headers: wire.headers, + signer: authenticated.signer, + }, + wire.url, + &wire.body, + Some(request.timeout.unwrap_or(Duration::from_secs(600))), + )?; + let response = crate::outbound::send(outbound, http) + .await + .map_err(network)?; + let status = response.status().as_u16(); + if !response.status().is_success() { + let body = response.text().await.map_err(network)?; + return Err(litellm_http::transport::Error::Http { + status, + body: litellm_http::request::truncate_error_body(&body), + } + .into()); + } + if stream { + let headers = response + .headers() + .iter() + .filter_map(|(name, value)| { + Some((name.to_string(), value.to_str().ok()?.to_owned())) + }) + .collect(); + let chunks = response + .bytes_stream() + .map(|chunk| chunk.map_err(network)) + .boxed(); + return Ok(ResponsesOutput::Stream { + head: ResponsesStreamHead { headers }, + chunks, + }); + } + let body = response.text().await.map_err(network)?; + let raw = RawResponse { body: body.clone() }; + if let Some(observers) = observers { + observers.emit(litellm_host::lifecycle::CallEvent::Execution( + ExecutionEvent::ProviderResponseReceived { raw: raw.clone() }, + )); + } + interceptors + .after_provider_response(raw) + .await + .map_err(Error::post_call)?; + let value = serde_json::from_str(&body) + .map_err(|error| Error::InvalidResponse(error.to_string().into()))?; + request + .config + .transform_response_api_response(value) + .map(ResponsesOutput::Complete) + .map_err(Error::from) }, - wire.url, - &wire.body, - Some(request.timeout.unwrap_or(Duration::from_secs(600))), - )?; - let response = crate::outbound::send(outbound, http) - .await - .map_err(network)?; - let status = response.status().as_u16(); - if !response.status().is_success() { - let body = response.text().await.map_err(network)?; - return Err(litellm_http::transport::Error::Http { - status, - body: litellm_http::request::truncate_error_body(&body), - } - .into()); - } - if stream { - let headers = response - .headers() - .iter() - .filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_owned()))) - .collect(); - let chunks = response - .bytes_stream() - .map(|chunk| chunk.map_err(network)) - .boxed(); - return Ok(ResponsesOutput::Stream { - head: ResponsesStreamHead { headers }, - chunks, - }); - } - let body = response.text().await.map_err(network)?; - let raw = RawResponse { body: body.clone() }; - if let Some(observers) = observers { - observers.emit(litellm_host::lifecycle::CallEvent::Execution( - ExecutionEvent::ProviderResponseReceived { raw: raw.clone() }, - )); - } - interceptors - .after_provider_response(raw) - .await - .map_err(Error::post_call)?; - let value = serde_json::from_str(&body) - .map_err(|error| Error::InvalidResponse(error.to_string().into()))?; - request - .config - .transform_response_api_response(value) - .map(ResponsesOutput::Complete) - .map_err(Error::from) + ) + .await } fn network(error: reqwest::Error) -> Error { diff --git a/litellm-rust/crates/core/src/responses/mod.rs b/litellm-rust/crates/core/src/responses/mod.rs index 8997aea8269..f388df25c7a 100644 --- a/litellm-rust/crates/core/src/responses/mod.rs +++ b/litellm-rust/crates/core/src/responses/mod.rs @@ -19,6 +19,7 @@ pub struct ResponsesRoute { http: litellm_http::Client, auth: Arc, secrets: Arc, + cache: Option, } impl ResponsesRoute { @@ -31,6 +32,14 @@ impl ResponsesRoute { http, auth, secrets, + cache: None, + } + } + + pub fn with_cache(self, cache: litellm_cache_response::ScopedCache) -> Self { + Self { + cache: Some(cache), + ..self } } @@ -38,11 +47,15 @@ impl ResponsesRoute { &self, call: ResponsesCall, interceptors: &impl Interceptors, - observers: Option, + options: impl Into, ) -> Result { + let crate::CallOptions { + cache: cache_options, + observers, + } = options.into(); litellm_host::lifecycle::observe_call( observers.clone(), - self.run(call, interceptors, observers.as_ref()), + self.run(call, cache_options, interceptors, observers.as_ref()), ) .await } @@ -58,25 +71,36 @@ impl ResponsesRoute { async fn run( &self, call: ResponsesCall, - interceptors: &impl Interceptors, + cache_options: Option, + interceptors: &impl litellm_host::interceptors::Interceptors, observers: Option<&ObservationSender>, ) -> Result { crate::diagnostic::call(async { - let request = prepare::prepare(call, self.secrets.as_ref()).await?; - crate::diagnostic::provider( - &request.context.model, - &request.context.custom_llm_provider, - ); - let execute: futures_util::future::BoxFuture<'_, Result> = - Box::pin(handler::execute( - &self.http, - &self.auth, - request, - interceptors, - observers, - )); - execute.await + self.run_provider(call, cache_options, interceptors, observers) + .await }) .await } + + async fn run_provider( + &self, + call: ResponsesCall, + cache_options: Option, + interceptors: &impl Interceptors, + observers: Option<&ObservationSender>, + ) -> Result { + let request = prepare::prepare(call, self.secrets.as_ref()).await?; + crate::diagnostic::provider(&request.context.model, &request.context.custom_llm_provider); + let execute: futures_util::future::BoxFuture<'_, Result> = + Box::pin(handler::execute( + &self.http, + &self.auth, + request, + self.cache.clone(), + cache_options, + interceptors, + observers, + )); + execute.await + } } diff --git a/litellm-rust/crates/core/src/responses/prepare.rs b/litellm-rust/crates/core/src/responses/prepare.rs index 2d58d9097d8..dcc112b48bb 100644 --- a/litellm-rust/crates/core/src/responses/prepare.rs +++ b/litellm-rust/crates/core/src/responses/prepare.rs @@ -16,14 +16,9 @@ pub(super) async fn prepare( call: ResponsesCall, secrets: &dyn SecretSource, ) -> Result { - let provider = call.custom_llm_provider.as_deref().unwrap_or("openai"); - if provider != "openai" { - return Err(Error::Unsupported("native HTTP responses provider")); - } - let model = call.model.strip_prefix("openai/").unwrap_or(&call.model); - if model.is_empty() || model.contains('/') { - return Err(Error::InvalidProvider(call.model)); - } + let identity = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?; + let provider = identity.provider.as_str(); + let model = identity.model.as_str(); let config: &'static dyn BaseResponsesApiConfig = &OpenAiResponsesApiConfig; let snapshot = secrets .resolve(config.secret_names(call.api_key.as_deref(), call.api_base.as_deref())) @@ -56,3 +51,21 @@ pub(super) async fn prepare( timeout: call.timeout, }) } + +pub(super) fn resolve_provider( + model: &str, + custom_llm_provider: Option<&str>, +) -> Result { + let provider = custom_llm_provider.unwrap_or("openai"); + if provider != "openai" { + return Err(Error::Unsupported("native HTTP responses provider")); + } + let resolved = model.strip_prefix("openai/").unwrap_or(model); + if resolved.is_empty() || resolved.contains('/') { + return Err(Error::InvalidProvider(model.into())); + } + Ok(litellm_host::interceptors::ProviderIdentity { + model: resolved.into(), + provider: provider.into(), + }) +} diff --git a/litellm-rust/crates/core/src/responses/route.rs b/litellm-rust/crates/core/src/responses/route.rs index 3d7f545f443..cc641a22acc 100644 --- a/litellm-rust/crates/core/src/responses/route.rs +++ b/litellm-rust/crates/core/src/responses/route.rs @@ -1,4 +1,3 @@ -use litellm_host::observation::ObservationSender; use std::convert::Infallible; use bytes::Bytes; @@ -28,14 +27,48 @@ impl ResponsesRoute { pub fn machine( self, call: ResponsesCall, - observers: Option, + options: impl Into, ) -> HostedMachine { + let crate::CallOptions { + cache: cache_options, + observers, + } = options.into(); hosted_call( call, observers, move |call, _, interceptors, observers| async move { - self.run(call, &interceptors, observers.as_ref()).await + self.run(call, cache_options, &interceptors, observers.as_ref()) + .await }, ) } } + +impl crate::caching::Cachable for Responses { + const SURFACE: &'static str = "responses"; + + fn reusable(response: &Self::Response) -> bool { + response + .extra + .get("status") + .and_then(serde_json::Value::as_str) + == Some("completed") + } +} + +impl crate::caching::StreamCachable for Responses { + const TERMINAL_EVENT: &'static str = "response.completed"; + + fn replay(data: bytes::Bytes) -> Option> { + Some(litellm_host::call::CallOutput::Stream { + head: ResponsesStreamHead { + headers: Vec::new(), + }, + chunks: Box::pin(futures_util::stream::iter([Ok(data)])), + }) + } + + fn bytes(chunk: &Self::Chunk) -> &[u8] { + chunk.as_ref() + } +} diff --git a/litellm-rust/crates/core/tests/caching.rs b/litellm-rust/crates/core/tests/caching.rs new file mode 100644 index 00000000000..77c6b4cde1d --- /dev/null +++ b/litellm-rust/crates/core/tests/caching.rs @@ -0,0 +1,1217 @@ +use std::{ + convert::Infallible, + num::NonZeroUsize, + sync::{ + Arc, OnceLock, + atomic::{AtomicUsize, Ordering}, + }, + time::Duration, +}; + +use bytes::Bytes; +use futures_util::{StreamExt, TryStreamExt, stream}; +use litellm_cache_memory::InMemoryCache; +use litellm_cache_response::{ + CacheOptions, CacheScope, ResponseCache, ResponseCacheConfig, ResponseCacheService, + ResponseEnvelope, +}; +use litellm_core::{ + RouteError, + caching::{Cachable, CacheRequest, StreamCachable, execute_streaming, execute_unary}, +}; +use litellm_host::{ + call::{CallOutput, OutputOf}, + interceptors::{ + ExecutionFacts, Interceptors, ProviderIdentity, RawResponse, RequestContext, ResultSource, + WireRequest, + }, + lifecycle::{CallEvent, ExecutionEvent}, + observation::observation_channel, + protocol::Protocol, +}; +use rstest::{fixture, rstest}; +use serde_json::{Value, json}; + +struct TestRoute; + +impl Protocol for TestRoute { + type Request = Value; + type Response = Value; + type Error = RouteError; + type HostCall = Infallible; + type Chunk = Bytes; + type StreamHead = (); +} + +impl Cachable for TestRoute { + const SURFACE: &'static str = "test"; +} + +impl StreamCachable for TestRoute { + const TERMINAL_EVENT: &'static str = "message_stop"; + + fn replay(data: Bytes) -> Option> { + Some(CallOutput::Stream { + head: (), + chunks: stream::iter([Ok(data)]).boxed(), + }) + } + fn bytes(chunk: &Bytes) -> &[u8] { + chunk + } +} + +fn cache_request(input: Value) -> CacheRequest { + CacheRequest { + identity: ProviderIdentity { + model: "test-model".into(), + provider: "test-provider".into(), + }, + input, + } +} + +#[fixture] +fn cache() -> Arc { + cache_with_limit(4096) +} + +fn cache_with_limit(max_entry_bytes: usize) -> Arc { + Arc::new( + ResponseCache::new(Arc::new(InMemoryCache::new( + Some(100), + Some(Duration::from_secs(60)), + ))) + .with_config(ResponseCacheConfig { + namespace: "test".into(), + max_entry_bytes, + }), + ) +} + +async fn call( + cache: &Arc, + options: Option, + calls: &AtomicUsize, + request: Value, +) -> Value { + let output = execute_streaming::( + cache_request(request), + Some(cache.clone()), + options, + &(), + None, + || async { + Ok(CallOutput::Complete( + json!({"call": calls.fetch_add(1, Ordering::SeqCst)}), + )) + }, + ) + .await + .unwrap(); + let CallOutput::Complete(response) = output else { + panic!("expected a response"); + }; + response +} + +#[rstest] +#[case::normal(CacheOptions::new(CacheScope::Shared), true, true)] +#[case::no_cache(CacheOptions { no_cache: true, ..CacheOptions::new(CacheScope::Shared) }, false, true)] +#[case::no_store(CacheOptions { no_store: true, ..CacheOptions::new(CacheScope::Shared) }, true, false)] +#[case::disabled(CacheOptions { caching: Some(false), ..CacheOptions::new(CacheScope::Shared) }, false, false)] +#[tokio::test] +async fn cache_controls_apply_to_both_reads_and_writes( + cache: Arc, + #[case] options: CacheOptions, + #[case] reads: bool, + #[case] writes: bool, +) { + let calls = AtomicUsize::new(0); + let options = Some(options); + let first = call(&cache, options.clone(), &calls, json!({"model":"test"})).await; + let second = call( + &cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + json!({"model":"test"}), + ) + .await; + assert_eq!(first == second, writes); + let third = call(&cache, options, &calls, json!({"model":"test"})).await; + assert_eq!(second == third, reads); + assert_eq!( + calls.load(Ordering::SeqCst), + 1 + usize::from(!writes) + usize::from(!reads) + ); +} + +#[rstest] +#[tokio::test] +async fn request_identity_is_canonical_and_scoped(cache: Arc) { + let calls = AtomicUsize::new(0); + let first = call( + &cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + json!({"model":"m", "input":{"a":1,"b":2}}), + ) + .await; + let second = call( + &cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + json!({"input":{"b":2,"a":1}, "model":"m"}), + ) + .await; + assert_eq!(first, second); + let other = call( + &cache, + Some(CacheOptions { + scope: CacheScope::Isolated("other-tenant".into()), + ..CacheOptions::new(CacheScope::Shared) + }), + &calls, + json!({"model":"m", "input":{"a":1,"b":2}}), + ) + .await; + assert_ne!(first, other); + let changed = call( + &cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + json!({"model":"m", "input":{"a":2,"b":2}}), + ) + .await; + assert_ne!(first, changed); +} + +async fn streamed( + cache: &Arc, + calls: &AtomicUsize, + text: &str, + fail: bool, +) -> OutputOf { + execute_streaming::( + cache_request(json!({"stream":true})), + Some(cache.clone()), + Some(CacheOptions::new(CacheScope::Shared)), + &(), + None, + || async { + calls.fetch_add(1, Ordering::SeqCst); + let chunks = text + .as_bytes() + .chunks(3) + .map(|bytes| Ok(Bytes::copy_from_slice(bytes))) + .collect::>(); + let ending = fail.then_some(Err(RouteError::Unsupported("test transport failure"))); + Ok(CallOutput::Stream { + head: (), + chunks: stream::iter(chunks.into_iter().chain(ending)).boxed(), + }) + }, + ) + .await + .unwrap() +} + +async fn consume(output: OutputOf) -> Result, RouteError> { + let CallOutput::Stream { chunks, .. } = output else { + panic!("expected a stream"); + }; + chunks + .try_fold(Vec::new(), |mut bytes, chunk| async move { + bytes.extend_from_slice(&chunk); + Ok(bytes) + }) + .await +} + +#[rstest] +#[case::complete("data: {\"type\":\"message_stop\"}\n\n", false, true)] +#[case::truncated("data: {\"type\":\"content_block_delta\"}\n\n", false, false)] +#[case::error_then_stop( + "data: {\"type\":\"error\"}\n\ndata: {\"type\":\"message_stop\"}\n\n", + false, + false +)] +#[case::trailing_incomplete("data: {\"type\":\"message_stop\"}\n\ndata: {", false, false)] +#[case::transport_failure("data: {\"type\":\"message_stop\"}\n\n", true, false)] +#[tokio::test] +async fn stream_replay_requires_successful_exhaustion( + cache: Arc, + #[case] text: &str, + #[case] fail: bool, + #[case] cached: bool, +) { + let calls = AtomicUsize::new(0); + let first = consume(streamed(&cache, &calls, text, fail).await).await; + assert_eq!(first.is_err(), fail); + let second = consume(streamed(&cache, &calls, text, fail).await).await; + assert_eq!(second.is_err(), fail); + if !fail { + assert_eq!(first.unwrap(), second.unwrap()); + } + assert_eq!(calls.load(Ordering::SeqCst), if cached { 1 } else { 2 }); +} + +#[rstest] +#[tokio::test] +async fn abandoning_a_partially_consumed_stream_does_not_store( + cache: Arc, +) { + let calls = AtomicUsize::new(0); + let text = "data: {\"type\":\"message_stop\"}\n\n"; + let CallOutput::Stream { mut chunks, .. } = streamed(&cache, &calls, text, false).await else { + panic!(); + }; + assert!(chunks.next().await.unwrap().is_ok()); + drop(chunks); + assert_eq!( + consume(streamed(&cache, &calls, text, false).await) + .await + .unwrap(), + text.as_bytes() + ); + assert_eq!( + consume(streamed(&cache, &calls, text, false).await) + .await + .unwrap(), + text.as_bytes() + ); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +#[rstest] +#[tokio::test] +async fn oversized_streams_are_delivered_without_being_stored() { + let cache = cache_with_limit(8); + let calls = AtomicUsize::new(0); + let text = "data: {\"type\":\"message_stop\"}\n\n"; + assert_eq!( + consume(streamed(&cache, &calls, text, false).await) + .await + .unwrap(), + text.as_bytes() + ); + assert_eq!( + consume(streamed(&cache, &calls, text, false).await) + .await + .unwrap(), + text.as_bytes() + ); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +#[rstest] +#[tokio::test] +async fn a_provider_failure_never_populates_the_cache(cache: Arc) { + let calls = AtomicUsize::new(0); + let first = execute_streaming::( + cache_request(json!({})), + Some(cache.clone()), + Some(CacheOptions::new(CacheScope::Shared)), + &(), + None, + || async { + calls.fetch_add(1, Ordering::SeqCst); + Err(RouteError::Unsupported("test provider failure")) + }, + ) + .await; + assert!(first.is_err()); + let successful = call( + &cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + json!({}), + ) + .await; + let replayed = call( + &cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + json!({}), + ) + .await; + assert_eq!(successful, replayed); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +struct InvalidEntryCache( + ResponseCache>, + Value, +); + +impl ResponseCacheService for InvalidEntryCache { + fn config(&self) -> &ResponseCacheConfig { + self.0.config() + } + + fn lookup<'a>( + &'a self, + request: &'a litellm_cache_response::ResponseCacheRequest, + now: Duration, + ) -> futures_util::future::BoxFuture<'a, Result, litellm_cache::Error>> { + Box::pin(async move { + Ok(self + .0 + .async_lookup(request, now) + .await? + .or_else(|| Some(self.1.clone()))) + }) + } + + fn store<'a>( + &'a self, + request: &'a litellm_cache_response::ResponseCacheRequest, + response: Value, + now: Duration, + ) -> futures_util::future::BoxFuture<'a, Result<(), litellm_cache::Error>> { + Box::pin(self.0.async_store(request, response, now)) + } +} + +#[rstest] +#[case::legacy(json!({"unexpected":"old-format"}))] +#[case::wrong_version(json!({"version":2,"surface":"test","output":{"kind":"Response","value":{"call":100}}}))] +#[case::wrong_surface(json!({"version":1,"surface":"other","output":{"kind":"Response","value":{"call":100}}}))] +#[tokio::test] +async fn an_invalid_cached_envelope_is_replaced_by_a_provider_result(#[case] poisoned: Value) { + let cache: Arc = Arc::new(InvalidEntryCache( + ResponseCache::new(Arc::new(InMemoryCache::default())), + poisoned, + )); + let request = json!({"input":"hello"}); + let calls = AtomicUsize::new(0); + let first = call( + &cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + request.clone(), + ) + .await; + let second = call( + &cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + request, + ) + .await; + assert_eq!(first, second); + assert_eq!(calls.load(Ordering::SeqCst), 1); +} + +#[rstest] +#[case::chat_completion(json!({"kind":"Response","value":{"id":"chat-1","model":"test","choices":[],"usage":{"prompt_tokens":3,"completion_tokens":2}}}))] +#[case::wrong_envelope(json!({"kind":"Stream","value":"data: [DONE]\n\n"}))] +#[tokio::test] +async fn responses_refetches_instead_of_deserializing_another_api_response( + #[case] poisoned: Value, +) { + use litellm_core::responses::route::Responses; + use litellm_types::responses::main::ResponsesApiResponse; + + let cache: Arc = Arc::new(InvalidEntryCache( + ResponseCache::new(Arc::new(InMemoryCache::default())), + serde_json::to_value(ResponseEnvelope::new("responses", poisoned)).unwrap(), + )); + let calls = AtomicUsize::new(0); + for _ in 0..2 { + let response = execute_unary::( + cache_request(json!({"input":"hello"})), + Some(cache.clone()), + Some(CacheOptions::new(CacheScope::Shared)), + &(), + None, + || async { + calls.fetch_add(1, Ordering::SeqCst); + Ok(ResponsesApiResponse { + id: "fresh-response".into(), + model: "test".into(), + output: vec![ + json!({"type":"message","content":[{"type":"output_text","text":"fresh"}]}), + ], + extra: [("status".into(), json!("completed"))] + .into_iter() + .collect(), + }) + }, + ) + .await + .unwrap(); + assert_eq!(response.id, "fresh-response"); + assert_eq!(response.output[0]["content"][0]["text"], "fresh"); + } + assert_eq!(calls.load(Ordering::SeqCst), 1); +} + +#[rstest] +#[case::system("system", json!("answer ALPHA"), json!("answer BETA"))] +#[case::stop_sequences("stop_sequences", json!(["STOP"]), json!(["END"]))] +#[case::top_k("top_k", json!(5), json!(10))] +#[case::tools("tools", json!([{"name":"a","input_schema":{"type":"object"}}]), json!([{"name":"b","input_schema":{"type":"object"}}]))] +#[case::tool_choice("tool_choice", json!({"type":"auto"}), json!({"type":"none"}))] +#[tokio::test] +async fn messages_cache_identity_includes_provider_native_parameters( + cache: Arc, + #[case] field: &str, + #[case] original: Value, + #[case] changed: Value, +) { + use litellm_core::messages::route::Messages; + use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; + + let calls = AtomicUsize::new(0); + for (value, expected_call) in [(original.clone(), 0), (changed, 1), (original, 0)] { + let response = + execute_unary::( + CacheRequest::from_wire( + ProviderIdentity { + model: "test".into(), + provider: "anthropic".into(), + }, + Some(&WireRequest { + url: "https://example.test/v1/messages".into(), + headers: vec![], + body: json!({ + "model":"test", "messages":[{"role":"user","content":"hello"}], + "max_tokens":32, (field):value + }), + }), + ), + Some(cache.clone()), + Some(CacheOptions::new(CacheScope::Shared)), + &(), + None, + || async { + let call = calls.fetch_add(1, Ordering::SeqCst); + Ok(Box::new(serde_json::from_value::(json!({ + "id":call.to_string(), "type":"message", "role":"assistant", "model":"test", + "content":[{"type":"text","text":format!("answer {call}")}], + "stop_reason":"end_turn", "stop_sequence":null + })).unwrap())) + }, + ) + .await + .unwrap(); + assert_eq!(response.id, expected_call.to_string()); + assert_eq!( + response.content[0]["text"], + format!("answer {expected_call}") + ); + } + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +struct UnavailableCache; + +impl litellm_cache::BaseCache for UnavailableCache { + type Value = litellm_cache_response::CacheEntry; + type Context = litellm_cache::ExactCacheContext; + + fn get_ttl(&self, _: &Self::Context) -> Option { + Some(Duration::from_secs(60)) + } + + fn get_cache( + &self, + _: &str, + _: &Self::Context, + ) -> Result, litellm_cache::Error> { + Err(litellm_cache::Error::Unavailable) + } + + fn set_cache( + &self, + _: &str, + _: Self::Value, + _: &Self::Context, + ) -> Result<(), litellm_cache::Error> { + Err(litellm_cache::Error::Unavailable) + } +} + +#[rstest] +#[tokio::test] +async fn backend_failures_do_not_fail_inference() { + let cache: Arc = Arc::new( + ResponseCache::new(Arc::new(UnavailableCache)).with_config(ResponseCacheConfig { + namespace: "test".into(), + max_entry_bytes: 4096, + }), + ); + let calls = AtomicUsize::new(0); + let first = call( + &cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + json!({}), + ) + .await; + let second = call( + &cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + json!({}), + ) + .await; + assert_ne!(first, second); + assert_eq!(calls.load(Ordering::SeqCst), 2); +} + +struct UnaryTestRoute; + +#[derive(Default)] +struct CacheHitAccounting { + calls: AtomicUsize, + key: OnceLock, + reject: bool, +} + +impl Interceptors for CacheHitAccounting { + async fn result_ready(&self, facts: ExecutionFacts) -> Result<(), RouteError> { + self.calls.fetch_add(1, Ordering::SeqCst); + let ResultSource::Cache { key } = facts.source else { + panic!("expected cache source") + }; + assert_eq!( + facts.provider, + ProviderIdentity { + model: "test-model".into(), + provider: "test-provider".into() + } + ); + self.key.set(key).unwrap(); + if self.reject { + return Err(RouteError::Unsupported("cache accounting rejected")); + } + Ok(()) + } + + async fn before_provider_request( + &self, + wire: WireRequest, + _: RequestContext, + ) -> Result { + Ok(wire) + } + + async fn after_provider_response(&self, _: RawResponse) -> Result<(), RouteError> { + Ok(()) + } +} + +#[rstest] +#[case::unary(false)] +#[case::stream_replay(true)] +#[tokio::test] +async fn cache_hits_notify_accounting_once_and_propagate_its_failure( + cache: Arc, + #[case] streaming_route: bool, + #[values(false, true)] reject: bool, +) { + let provider_calls = AtomicUsize::new(0); + let accounting = CacheHitAccounting { + reject, + ..Default::default() + }; + let (observer, mut events) = observation_channel(NonZeroUsize::new(4).unwrap()); + let request = if streaming_route { + json!({"stream":true}) + } else { + json!({"input":"hello"}) + }; + let expected = if streaming_route { + json!( + consume( + streamed( + &cache, + &provider_calls, + "data: {\"type\":\"message_stop\"}\n\n", + false, + ) + .await + ) + .await + .unwrap() + ) + } else { + unary_call( + &cache, + Some(CacheOptions::new(CacheScope::Shared)), + &provider_calls, + request.clone(), + ) + .await + }; + let result = if streaming_route { + match execute_streaming::( + cache_request(request), + Some(cache), + Some(CacheOptions::new(CacheScope::Shared)), + &accounting, + Some(&observer), + || async { panic!("a cache hit must not call the provider") }, + ) + .await + { + Ok(output) => consume(output).await.map(|bytes| json!(bytes)), + Err(error) => Err(error), + } + } else { + execute_unary::( + cache_request(request), + Some(cache), + Some(CacheOptions::new(CacheScope::Shared)), + &accounting, + Some(&observer), + || async { panic!("a cache hit must not call the provider") }, + ) + .await + }; + if reject { + assert!(matches!( + result, + Err(RouteError::Unsupported("cache accounting rejected")) + )); + } else { + assert_eq!(result.unwrap(), expected); + } + assert_eq!(provider_calls.load(Ordering::SeqCst), 1); + assert_eq!(accounting.calls.load(Ordering::SeqCst), 1); + let key = accounting.key.get().unwrap(); + assert!(!key.is_empty()); + assert!(matches!( + events.try_recv().unwrap(), + CallEvent::Execution(ExecutionEvent::ResultReady { facts }) if facts.source == ResultSource::Cache { key: key.clone() } + )); + assert!(events.try_recv().is_err()); +} + +impl Protocol for UnaryTestRoute { + type Request = Value; + type Response = Value; + type Error = RouteError; + type HostCall = Infallible; + type Chunk = Infallible; + type StreamHead = Infallible; +} + +impl Cachable for UnaryTestRoute { + const SURFACE: &'static str = "unary-test"; +} + +async fn unary_call( + cache: &Arc, + options: Option, + calls: &AtomicUsize, + request: Value, +) -> Value { + execute_unary::( + cache_request(request), + Some(cache.clone()), + options, + &(), + None, + || async { Ok(json!({"call":calls.fetch_add(1, Ordering::SeqCst)})) }, + ) + .await + .unwrap() +} + +#[rstest] +#[case::normal(CacheOptions::new(CacheScope::Shared), true, true)] +#[case::no_cache(CacheOptions { no_cache: true, ..CacheOptions::new(CacheScope::Shared) }, false, true)] +#[case::no_store(CacheOptions { no_store: true, ..CacheOptions::new(CacheScope::Shared) }, true, false)] +#[case::disabled(CacheOptions { caching: Some(false), ..CacheOptions::new(CacheScope::Shared) }, false, false)] +#[tokio::test] +async fn unary_cache_controls_do_not_change_the_shared_service( + cache: Arc, + #[case] options: CacheOptions, + #[case] reads: bool, + #[case] writes: bool, +) { + let calls = AtomicUsize::new(0); + let options = Some(options); + let first = unary_call(&cache, options.clone(), &calls, json!({"input":"hello"})).await; + let second = unary_call( + &cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + json!({"input":"hello"}), + ) + .await; + assert_eq!(first == second, writes); + let third = unary_call(&cache, options, &calls, json!({"input":"hello"})).await; + assert_eq!(second == third, reads); + let fourth = unary_call( + &cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + json!({"input":"hello"}), + ) + .await; + assert_eq!(fourth, if !reads && writes { third } else { second }); + assert_eq!( + calls.load(Ordering::SeqCst), + 1 + usize::from(!writes) + usize::from(!reads) + ); +} + +#[rstest] +#[tokio::test] +async fn namespaces_and_surfaces_isolate_entries_on_shared_storage() { + let storage = Arc::new(InMemoryCache::default()); + let first_cache: Arc = Arc::new( + ResponseCache::new(storage.clone()).with_config(ResponseCacheConfig { + namespace: "first".into(), + max_entry_bytes: 4096, + }), + ); + let second_cache: Arc = Arc::new( + ResponseCache::new(storage).with_config(ResponseCacheConfig { + namespace: "second".into(), + max_entry_bytes: 4096, + }), + ); + let calls = AtomicUsize::new(0); + let first = call( + &first_cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + json!({"input":"hello"}), + ) + .await; + let different_namespace = call( + &second_cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + json!({"input":"hello"}), + ) + .await; + let different_surface = unary_call( + &first_cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + json!({"input":"hello"}), + ) + .await; + assert_ne!(first, different_namespace); + assert_ne!(first, different_surface); + assert_eq!( + call( + &first_cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + json!({"input":"hello"}) + ) + .await, + first + ); + assert_eq!( + call( + &second_cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + json!({"input":"hello"}) + ) + .await, + different_namespace + ); + assert_eq!( + unary_call( + &first_cache, + Some(CacheOptions::new(CacheScope::Shared)), + &calls, + json!({"input":"hello"}) + ) + .await, + different_surface + ); + assert_eq!(calls.load(Ordering::SeqCst), 3); +} + +#[rstest] +#[case::completed("completed", 1)] +#[case::incomplete("incomplete", 2)] +#[tokio::test] +async fn responses_cache_only_reuses_completed_responses( + cache: Arc, + #[case] status: &str, + #[case] expected_calls: usize, +) { + use litellm_core::responses::route::Responses; + use litellm_types::responses::main::ResponsesApiResponse; + + let calls = AtomicUsize::new(0); + for _ in 0..2 { + let response = execute_unary::( + cache_request(json!({"input":"hello"})), + Some(cache.clone()), + Some(CacheOptions::new(CacheScope::Shared)), + &(), + None, + || async { + let call = calls.fetch_add(1, Ordering::SeqCst); + Ok(ResponsesApiResponse { + id: call.to_string(), + model: "test".into(), + output: Vec::new(), + extra: [("status".into(), json!(status))].into_iter().collect(), + }) + }, + ) + .await + .unwrap(); + assert_eq!(response.extra.get("status"), Some(&json!(status))); + } + assert_eq!(calls.load(Ordering::SeqCst), expected_calls); +} + +mod support; +use support::traces; + +#[rstest] +#[case::without_cache(false)] +#[case::with_cache(true)] +#[tokio::test] +async fn the_same_route_entrypoint_reports_facts_with_or_without_caching( + cache: Arc, + #[case] caching: bool, + traces: support::TraceCapture, +) { + use litellm_cache_response::ScopedCache; + use litellm_core::chat_completions::types::ChatCompletionsRequest; + use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; + + let upstream = MockServer::start().await; + let body = json!({"id":"msg-test","type":"message","role":"assistant","model":"cache-test-model", + "content":[{"type":"text","text":"cached answer"}],"stop_reason":"end_turn", + "stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}); + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(200).set_body_json(body)) + .expect(if caching { 1 } else { 2 }) + .mount(&upstream) + .await; + let route = support::chat_completions_route(); + let route = if caching { + route.with_cache(ScopedCache::new(cache, CacheScope::Shared)) + } else { + route + }; + let (observer, mut events) = observation_channel(NonZeroUsize::new(16).unwrap()); + let base = upstream.uri(); + for _ in 0..2 { + let response = traces + .logger() + .instrument(route.execute( + ChatCompletionsRequest { + model: "anthropic/cache-test-model", + messages: json!([{"role":"user","content":"hello"}]), + optional_params: [("max_tokens".into(), json!(16))].into_iter().collect(), + api_key: Some("test-key"), + api_base: Some(&base), + custom_llm_provider: None, + extra_headers: None, + timeout: None, + }, + &(), + Some(observer.clone()), + )) + .await + .unwrap(); + assert_eq!( + serde_json::to_value(response).unwrap()["usage"]["total_tokens"], + 15 + ); + } + let facts: Vec<_> = std::iter::from_fn(|| events.try_recv().ok()) + .filter_map(|event| match event { + CallEvent::Execution(ExecutionEvent::ResultReady { facts }) => Some(facts), + _ => None, + }) + .collect(); + assert_eq!(facts.len(), 2); + assert_eq!( + facts[0].provider, + ProviderIdentity { + model: "cache-test-model".into(), + provider: "anthropic".into() + } + ); + assert_eq!(facts[1].provider, facts[0].provider); + assert_eq!(facts[0].source, ResultSource::Provider); + match &facts[1].source { + ResultSource::Provider => assert!(!caching), + ResultSource::Cache { key } => { + assert!(caching); + assert!(!key.is_empty()); + } + } + let summaries = traces.summaries("litellm.route"); + assert_eq!(summaries.len(), 2); + for summary in summaries { + assert_eq!(summary["provider"], "anthropic"); + assert_eq!(summary["resolved_model"], "cache-test-model"); + assert_eq!(summary["outcome"], "success"); + } + upstream.verify().await; +} + +struct ChangingSecrets { + revision: AtomicUsize, + endpoints: [String; 2], + change_credentials: bool, +} + +impl litellm_secrets::source::SecretSource for ChangingSecrets { + fn get_secret_str<'a>( + &'a self, + name: &'a str, + ) -> futures_util::future::BoxFuture< + 'a, + Result, litellm_secrets::Error>, + > { + Box::pin(async move { + let revision = self.revision.load(Ordering::SeqCst); + let value = if name.ends_with("_API_KEY") { + Some(format!( + "key-{}", + if self.change_credentials { revision } else { 0 } + )) + } else if name.ends_with("_API_BASE") { + Some(self.endpoints[revision].clone()) + } else { + None + }; + Ok(value.map(litellm_secrets::SecretValue::new)) + }) + } +} + +#[derive(Default)] +struct ChangingHooks { + calls: AtomicUsize, + rewrite: bool, + facts: std::sync::Mutex>, +} + +impl Interceptors for ChangingHooks { + async fn before_provider_request( + &self, + mut wire: WireRequest, + _: RequestContext, + ) -> Result { + let call = self.calls.fetch_add(1, Ordering::SeqCst); + if self.rewrite { + wire.body["temperature"] = json!(if call < 2 { 0.1 } else { 0.8 }); + } + Ok(wire) + } + + async fn after_provider_response(&self, _: RawResponse) -> Result<(), RouteError> { + Ok(()) + } + + async fn result_ready(&self, facts: ExecutionFacts) -> Result<(), RouteError> { + self.facts.lock().unwrap().push(facts); + Ok(()) + } +} + +#[rstest] +#[case::chat_credentials("chat", "credentials")] +#[case::chat_endpoint("chat", "endpoint")] +#[case::chat_callback("chat", "callback")] +#[case::messages_credentials("messages", "credentials")] +#[case::messages_endpoint("messages", "endpoint")] +#[case::messages_callback("messages", "callback")] +#[case::responses_credentials("responses", "credentials")] +#[case::responses_endpoint("responses", "endpoint")] +#[case::responses_callback("responses", "callback")] +#[tokio::test] +async fn cache_identity_follows_resolved_configuration_and_request_callbacks( + cache: Arc, + #[case] surface: &str, + #[case] change: &str, +) { + use litellm_cache_response::ScopedCache; + use litellm_core::{ + chat_completions::{ChatCompletionsRoute, types::ChatCompletionsRequest}, + messages::MessagesCall, + responses::types::ResponsesCall, + }; + use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; + + let first = MockServer::start().await; + let second = MockServer::start().await; + let response = if surface == "responses" { + json!({"id":"response-test", "model":"test", "output":[], "status":"completed"}) + } else { + json!({"id":"message-test", "type":"message", "role":"assistant", "model":"test", + "content":[{"type":"text", "text":"answer"}], "stop_reason":"end_turn", "stop_sequence":null, + "usage":{"input_tokens":3,"output_tokens":2}}) + }; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(200).set_body_json(response.clone())) + .expect(if change == "endpoint" { 1 } else { 2 }) + .mount(&first) + .await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(200).set_body_json(response)) + .expect(if change == "endpoint" { 1 } else { 0 }) + .mount(&second) + .await; + let secrets = Arc::new(ChangingSecrets { + revision: AtomicUsize::new(0), + endpoints: [ + first.uri(), + if change == "endpoint" { + second.uri() + } else { + first.uri() + }, + ], + change_credentials: change == "credentials", + }); + let hooks = ChangingHooks { + rewrite: change == "callback", + ..Default::default() + }; + for call in 0..4 { + secrets + .revision + .store(usize::from(call >= 2), Ordering::SeqCst); + let cache = ScopedCache::new(cache.clone(), CacheScope::Shared); + match surface { + "chat" => { + ChatCompletionsRoute::new( + litellm_http::Client::plain_for_test(), + Arc::new(Default::default()), + secrets.clone(), + ) + .with_cache(cache) + .execute( + ChatCompletionsRequest { + model: "anthropic/cache-test-model", + messages: json!([{"role":"user","content":"hello"}]), + optional_params: [("max_tokens".into(), json!(32))].into_iter().collect(), + api_key: None, + api_base: None, + custom_llm_provider: None, + extra_headers: None, + timeout: None, + }, + &hooks, + None, + ) + .await + .unwrap(); + } + "messages" => { + support::messages_route(secrets.clone()).with_cache(cache).execute(MessagesCall { + body: serde_json::from_value(json!({"model":"anthropic/cache-test-model","messages":[{"role":"user","content":"hello"}],"max_tokens":32})).unwrap(), + api_key:None,api_base:None,custom_llm_provider:None,extra_headers:None,provider_specific_header:None,timeout:None,shaping:Default::default(), + }, &hooks, None).await.unwrap(); + } + "responses" => { + support::responses_route(secrets.clone()) + .with_cache(cache) + .execute( + ResponsesCall { + model: "test".into(), + input: json!("hello"), + optional_params: Default::default(), + api_key: None, + api_base: None, + custom_llm_provider: None, + extra_headers: None, + timeout: None, + }, + &hooks, + None, + ) + .await + .unwrap(); + } + _ => unreachable!(), + } + } + assert_eq!(hooks.calls.load(Ordering::SeqCst), 4); + { + let facts = hooks.facts.lock().unwrap(); + assert_eq!(facts[0].source, ResultSource::Provider); + assert_eq!(facts[2].source, ResultSource::Provider); + let (ResultSource::Cache { key: first_key }, ResultSource::Cache { key: second_key }) = + (&facts[1].source, &facts[3].source) + else { + panic!("unchanged effective requests must hit the cache"); + }; + assert_ne!(first_key, second_key); + } + let requests = first.received_requests().await.unwrap(); + if change == "credentials" { + let header = if surface == "responses" { + "authorization" + } else { + "x-api-key" + }; + assert_ne!(requests[0].headers[header], requests[1].headers[header]); + } + if change == "callback" { + assert_eq!( + serde_json::from_slice::(&requests[0].body).unwrap()["temperature"], + 0.1 + ); + assert_eq!( + serde_json::from_slice::(&requests[1].body).unwrap()["temperature"], + 0.8 + ); + } + first.verify().await; + second.verify().await; +} + +#[rstest] +#[tokio::test] +async fn signed_requests_bypass_response_caching(cache: Arc) { + use litellm_cache_response::ScopedCache; + use litellm_core::chat_completions::types::ChatCompletionsRequest; + use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; + + let upstream = MockServer::start().await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "output":{"message":{"role":"assistant","content":[{"text":"answer"}]}}, + "stopReason":"end_turn", "usage":{"inputTokens":3,"outputTokens":2,"totalTokens":5} + }))) + .expect(2) + .mount(&upstream) + .await; + let route = + support::chat_completions_route().with_cache(ScopedCache::new(cache, CacheScope::Shared)); + let hooks = ChangingHooks::default(); + for _ in 0..2 { + let response = route.execute(ChatCompletionsRequest { + model:"bedrock/anthropic.cache-test-model", + messages:json!([{"role":"user","content":"hello"}]), + optional_params:json!({"aws_access_key_id":"test-access","aws_secret_access_key":"test-secret","aws_region_name":"eu-west-1"}).as_object().unwrap().clone(), + api_key:None,api_base:Some(&upstream.uri()),custom_llm_provider:None,extra_headers:None,timeout:None, + }, &hooks, None).await.unwrap(); + assert_eq!( + serde_json::to_value(response).unwrap()["usage"]["total_tokens"], + 5 + ); + } + assert!( + hooks + .facts + .lock() + .unwrap() + .iter() + .all(|facts| facts.source == ResultSource::Provider) + ); + upstream.verify().await; +} diff --git a/litellm-rust/crates/core/tests/chat_completions.rs b/litellm-rust/crates/core/tests/chat_completions.rs index bb62fd000a8..fa9bd731809 100644 --- a/litellm-rust/crates/core/tests/chat_completions.rs +++ b/litellm-rust/crates/core/tests/chat_completions.rs @@ -1,4 +1,8 @@ use litellm_host::interceptors::RawResponse; +use litellm_host::{ + interceptors::{ExecutionFacts, ResultSource}, + lifecycle::ExecutionEvent, +}; use std::time::Duration; use litellm_core::chat_completions::{Error, types::ChatCompletionsRequest}; @@ -310,7 +314,13 @@ async fn direct_and_hosted_calls_share_hooks_and_lifecycle( &events[..], [ CallEvent::Started { .. }, - CallEvent::Execution(_), + CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { .. }), + CallEvent::Execution(ExecutionEvent::ResultReady { + facts: ExecutionFacts { + source: ResultSource::Provider, + .. + } + }), CallEvent::Succeeded { .. } ] )); diff --git a/litellm-rust/crates/core/tests/messages/response.rs b/litellm-rust/crates/core/tests/messages/response.rs index 7d2fffc5beb..7ef669599fb 100644 --- a/litellm-rust/crates/core/tests/messages/response.rs +++ b/litellm-rust/crates/core/tests/messages/response.rs @@ -1,4 +1,8 @@ use litellm_core::messages::{MessagesResponse, messages_body}; +use litellm_host::{ + interceptors::{ExecutionFacts, ResultSource}, + lifecycle::ExecutionEvent, +}; use litellm_http::transport::Error as TransportError; use rstest::rstest; @@ -60,7 +64,13 @@ async fn calls_defer_execution_until_polled( true, [ CallEvent::Started { .. }, - CallEvent::Execution(_), + CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { .. }), + CallEvent::Execution(ExecutionEvent::ResultReady { + facts: ExecutionFacts { + source: ResultSource::Provider, + .. + } + }), CallEvent::Succeeded { .. } ] ) @@ -69,7 +79,13 @@ async fn calls_defer_execution_until_polled( true, [ CallEvent::Started { .. }, - CallEvent::Execution(_), + CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { .. }), + CallEvent::Execution(ExecutionEvent::ResultReady { + facts: ExecutionFacts { + source: ResultSource::Provider, + .. + } + }), CallEvent::Succeeded { .. } ] ) @@ -253,22 +269,25 @@ async fn the_facade_sends_through_the_injected_http_pool_configuration(call: Mes }; let resources = support::resources(); - let response = litellm_core::messages::MessagesRoute::new( - provider_http(&resources, &Resolution::from(&settings).config), - resources.auth, - no_secrets(), - ) - .execute( - MessagesCall { - api_key: Some("sk-ant".into()), - api_base: Some(base), - ..call - }, - &(), - None, - ) - .await - .expect("messages request succeeds"); + let response = litellm_core::messages::MessagesRoute::builder() + .with_http(provider_http( + &resources, + &Resolution::from(&settings).config, + )) + .with_auth(resources.auth) + .with_secrets(no_secrets()) + .build() + .execute( + MessagesCall { + api_key: Some("sk-ant".into()), + api_base: Some(base), + ..call + }, + &(), + None, + ) + .await + .expect("messages request succeeds"); let MessagesResponse::Complete(message) = response else { panic!("a non-streaming request returns a message"); @@ -321,3 +340,56 @@ async fn message_route_summary_excludes_payload_diagnostics( assert!(summaries[0].get("body").is_none()); assert!(!format!("{:?}", traces.records()).contains("private-key-sentinel")); } + +#[rstest] +#[case::uncached(false, 2)] +#[case::cached(true, 1)] +#[tokio::test] +async fn builder_preserves_dependencies_and_optional_cache( + #[case] caching: bool, + #[case] expected_requests: usize, +) { + use litellm_cache_memory::InMemoryCache; + use litellm_cache_response::{CacheScope, ResponseCache, ScopedCache}; + use litellm_core::messages::MessagesRoute; + + let upstream = upstream([message_response(), message_response()]).await; + let resources = resources(); + let builder = MessagesRoute::builder(); + let builder = if caching { + builder.with_cache(ScopedCache::new( + Arc::new(ResponseCache::new(Arc::new(InMemoryCache::new( + Some(100), + Some(Duration::from_secs(60)), + )))), + CacheScope::Shared, + )) + } else { + builder + }; + let route = builder + .with_secrets(Arc::new(RecordingSecrets::new([( + "ANTHROPIC_API_KEY", + "builder-key", + )]))) + .with_auth(resources.auth.clone()) + .with_http(provider_http(&resources, &http_config())) + .build(); + for _ in 0..2 { + let request = MessagesCall { + api_base: Some(upstream.uri()), + ..super::call() + }; + let MessagesResponse::Complete(response) = route.execute(request, &(), None).await.unwrap() + else { + panic!("expected a completed message"); + }; + assert_eq!( + response.content, + message_body()["content"].as_array().unwrap().as_slice() + ); + } + let requests = received(&upstream).await; + assert_eq!(requests.len(), expected_requests); + assert_eq!(requests[0].header("x-api-key"), Some("builder-key")); +} diff --git a/litellm-rust/crates/core/tests/ocr/lifecycle.rs b/litellm-rust/crates/core/tests/ocr/lifecycle.rs index d9492c4173a..dd94e44df37 100644 --- a/litellm-rust/crates/core/tests/ocr/lifecycle.rs +++ b/litellm-rust/crates/core/tests/ocr/lifecycle.rs @@ -19,6 +19,7 @@ use super::*; pub(crate) fn event_name(event: &CallEvent) -> &'static str { match event { + CallEvent::Execution(ExecutionEvent::ResultReady { .. }) => "result_ready", CallEvent::Started { .. } => "started", CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { .. }) => "response", CallEvent::Succeeded { .. } => "success", diff --git a/litellm-rust/crates/core/tests/ocr/machine.rs b/litellm-rust/crates/core/tests/ocr/machine.rs index b3653b65f2b..3d7a4140635 100644 --- a/litellm-rust/crates/core/tests/ocr/machine.rs +++ b/litellm-rust/crates/core/tests/ocr/machine.rs @@ -41,6 +41,11 @@ async fn drive_until( Err(error) => break Err(error), }; let answer = match op { + HostRequest::Intercept(InterceptRequest::ResultReady { facts, reply }) => host + .result_ready(facts) + .await + .map(|()| reply.send(())) + .map_err(HostFailure::Error), HostRequest::Stream(stream) => match stream { litellm_host::protocol::StreamDelivery::Open(head, _) => match head {}, litellm_host::protocol::StreamDelivery::Chunk(chunk, _) => match chunk {}, @@ -80,6 +85,7 @@ async fn drive_until_notified(machine: &mut OcrMachine, host: &LocalOcrHost, sto _ = stop.notified() => break, step = machine.resume() => { match step.unwrap() { + MachineStep::Suspended(HostRequest::Intercept(InterceptRequest::ResultReady { reply, .. })) => reply.send(()), MachineStep::Suspended(HostRequest::HostCall(op)) => host.handle_host_call(op).await.unwrap(), MachineStep::Suspended(HostRequest::Intercept(InterceptRequest::BeforeProviderRequest { wire, reply, .. })) => reply.send(*wire), MachineStep::Suspended(HostRequest::Intercept(InterceptRequest::AfterProviderResponse { reply, .. })) => reply.send(()), diff --git a/litellm-rust/crates/core/tests/responses.rs b/litellm-rust/crates/core/tests/responses.rs index 73d3208afdf..cd2a0c4734c 100644 --- a/litellm-rust/crates/core/tests/responses.rs +++ b/litellm-rust/crates/core/tests/responses.rs @@ -1,3 +1,7 @@ +use litellm_host::{ + interceptors::{ExecutionFacts, ResultSource}, + lifecycle::ExecutionEvent, +}; use std::sync::Arc; use futures_util::TryStreamExt; @@ -70,7 +74,13 @@ async fn http_responses_share_execution_and_hooks(call: ResponsesCall, #[case] h &host.events.0.lock().unwrap()[..], [ CallEvent::Started { .. }, - CallEvent::Execution(_), + CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { .. }), + CallEvent::Execution(ExecutionEvent::ResultReady { + facts: ExecutionFacts { + source: ResultSource::Provider, + .. + } + }), CallEvent::Succeeded { .. } ] )); @@ -97,7 +107,8 @@ async fn streaming_keeps_headers_and_bytes_and_finishes_after_consumption( let (headers, bytes) = if hosted { assert_eq!( litellm_host_native::in_process::run_hosted( - responses_route(no_secrets()).machine(host.request().unwrap(), None,), + responses_route(no_secrets()) + .machine(host.request().unwrap(), Some(host.events.0.sender.clone())), host.runtime(), ) .await @@ -117,7 +128,18 @@ async fn streaming_keeps_headers_and_bytes_and_finishes_after_consumption( else { panic!() }; - assert_eq!(host.events.0.lock().unwrap().len(), 1); + assert!(matches!( + &host.events.0.lock().unwrap()[..], + [ + CallEvent::Started { .. }, + CallEvent::Execution(ExecutionEvent::ResultReady { + facts: ExecutionFacts { + source: ResultSource::Provider, + .. + } + }), + ] + )); ( head.headers, chunks.try_collect::>().await.unwrap().concat(), @@ -127,7 +149,16 @@ async fn streaming_keeps_headers_and_bytes_and_finishes_after_consumption( assert_eq!(bytes, body.as_bytes()); assert!(matches!( &host.events.0.lock().unwrap()[..], - [CallEvent::Started { .. }, CallEvent::Succeeded { .. }] + [ + CallEvent::Started { .. }, + CallEvent::Execution(ExecutionEvent::ResultReady { + facts: ExecutionFacts { + source: ResultSource::Provider, + .. + } + }), + CallEvent::Succeeded { .. } + ] )); } diff --git a/litellm-rust/crates/core/tests/support/mod.rs b/litellm-rust/crates/core/tests/support/mod.rs index 1dd53114293..5ba1eb3ca46 100644 --- a/litellm-rust/crates/core/tests/support/mod.rs +++ b/litellm-rust/crates/core/tests/support/mod.rs @@ -44,11 +44,11 @@ pub fn provider_http( pub fn messages_route(secrets: Arc) -> litellm_core::messages::MessagesRoute { let resources = resources(); - litellm_core::messages::MessagesRoute::new( - provider_http(&resources, &http_config()), - resources.auth, - secrets, - ) + litellm_core::messages::MessagesRoute::builder() + .with_http(provider_http(&resources, &http_config())) + .with_auth(resources.auth) + .with_secrets(secrets) + .build() } pub fn chat_completions_route() -> litellm_core::chat_completions::ChatCompletionsRoute { diff --git a/litellm-rust/crates/gateway-inference/Cargo.toml b/litellm-rust/crates/gateway-inference/Cargo.toml index fd7c99204f7..e890679ff23 100644 --- a/litellm-rust/crates/gateway-inference/Cargo.toml +++ b/litellm-rust/crates/gateway-inference/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-cache-response.workspace = true axum = { workspace = true, features = ["json", "multipart", "original-uri"] } base64.workspace = true bytes.workspace = true @@ -13,15 +14,18 @@ litellm-auth.workspace = true litellm-gateway-auth.workspace = true litellm-core.workspace = true litellm-host-http.workspace = true +litellm-host.workspace = true litellm-http.workspace = true litellm-llms.workspace = true litellm-router.workspace = true litellm-secrets.workspace = true litellm-types.workspace = true +serde.workspace = true serde_json.workspace = true thiserror.workspace = true [dev-dependencies] +litellm-cache-memory.workspace = true futures-util.workspace = true tokio = { workspace = true, features = ["io-util"] } rstest.workspace = true diff --git a/litellm-rust/crates/gateway-inference/src/caching.rs b/litellm-rust/crates/gateway-inference/src/caching.rs new file mode 100644 index 00000000000..020f942ec19 --- /dev/null +++ b/litellm-rust/crates/gateway-inference/src/caching.rs @@ -0,0 +1,109 @@ +use std::time::Duration; + +use litellm_cache_response::{CacheOptions, CacheScope}; +use litellm_gateway_auth::AuthenticatedRequest; +use serde::Deserialize; +use serde_json::{Map, Value}; + +use crate::Error; + +#[derive(Default, Deserialize)] +#[serde(default, deny_unknown_fields)] +struct Controls { + #[serde(rename = "no-cache")] + no_cache: bool, + #[serde(rename = "no-store")] + no_store: bool, + ttl: Option, + #[serde(rename = "s-maxage", alias = "s-max-age")] + max_age: Option, +} + +type Prepared = (Map, CacheOptions); + +pub(crate) fn prepare( + identity: &AuthenticatedRequest, + body: Map, +) -> Result { + let controls: Controls = match body.get("cache").filter(|value| !value.is_null()) { + Some(value) => serde_json::from_value(value.clone()) + .map_err(|error| Error::InvalidBody(error.to_string()))?, + None => Controls::default(), + }; + let caching: Option = body + .get("caching") + .filter(|value| !value.is_null()) + .map(|value| serde_json::from_value(value.clone())) + .transpose() + .map_err(|error| Error::InvalidBody(error.to_string()))?; + let caller = identity.caller(); + let options = CacheOptions { + caching, + no_cache: controls.no_cache, + no_store: controls.no_store, + ttl: controls.ttl.map(duration).transpose()?, + max_age: controls.max_age.map(duration).transpose()?, + scope: CacheScope::Isolated( + serde_json::json!([ + caller.principal().authority(), + caller.principal().subject(), + caller.authentication().credential_id + ]) + .to_string(), + ), + }; + Ok(( + body.into_iter() + .filter(|(name, _)| !matches!(name.as_str(), "cache" | "caching")) + .collect(), + options, + )) +} + +fn duration(seconds: f64) -> Result { + Duration::try_from_secs_f64(seconds) + .ok() + .filter(|duration| !duration.is_zero()) + .ok_or_else(|| Error::InvalidBody("cache durations must be finite and positive".into())) +} + +#[derive(Clone, Default)] +pub(crate) struct CacheHeaders(std::sync::Arc>); + +impl litellm_host::interceptors::Interceptors for CacheHeaders { + async fn before_provider_request( + &self, + wire: litellm_host::interceptors::WireRequest, + _: litellm_host::interceptors::RequestContext, + ) -> Result { + Ok(wire) + } + + async fn after_provider_response( + &self, + _: litellm_host::interceptors::RawResponse, + ) -> Result<(), litellm_core::RouteError> { + Ok(()) + } + + async fn result_ready( + &self, + facts: litellm_host::interceptors::ExecutionFacts, + ) -> Result<(), litellm_core::RouteError> { + if let litellm_host::interceptors::ResultSource::Cache { key } = facts.source { + let _ = self.0.set(key); + } + Ok(()) + } +} + +impl CacheHeaders { + pub(crate) fn apply(&self, mut response: axum::response::Response) -> axum::response::Response { + if let Some(key) = self.0.get() + && let Ok(value) = axum::http::HeaderValue::from_str(key) + { + response.headers_mut().insert("x-litellm-cache-key", value); + } + response + } +} diff --git a/litellm-rust/crates/gateway-inference/src/chat_completions.rs b/litellm-rust/crates/gateway-inference/src/chat_completions.rs index 85b1f990a40..27b8e856b7e 100644 --- a/litellm-rust/crates/gateway-inference/src/chat_completions.rs +++ b/litellm-rust/crates/gateway-inference/src/chat_completions.rs @@ -42,9 +42,20 @@ async fn handle( ) -> Result { let deployment = request::resolve_deployment(gateway, &body)?; request::authorize_model(identity, deployment, &body).await?; + let (body, cache_options) = crate::caching::prepare(identity, body)?; + let route = gateway.chat_completions.clone(); + let route = match &gateway.cache { + Some(cache) => route.with_cache(litellm_cache_response::ScopedCache::new( + cache.clone(), + cache_options.scope.clone(), + )), + None => route, + }; + let messages = body.get("messages").cloned().unwrap_or_default(); + let headers = crate::caching::CacheHeaders::default(); let response = litellm_host_http::serve_unary( - gateway.chat_completions.clone().machine( + route.machine( ChatCompletionsCall { model: deployment.model.clone(), messages, @@ -58,13 +69,13 @@ async fn handle( extra_headers: None, timeout: deployment.timeout, }, - None, + cache_options, ), (), - (), + headers.clone(), litellm_host_http::Unary::new(Json), None, ) .await?; - Ok(response) + Ok(headers.apply(response)) } diff --git a/litellm-rust/crates/gateway-inference/src/lib.rs b/litellm-rust/crates/gateway-inference/src/lib.rs index e9ffd14c257..b669fb1ced1 100644 --- a/litellm-rust/crates/gateway-inference/src/lib.rs +++ b/litellm-rust/crates/gateway-inference/src/lib.rs @@ -4,6 +4,7 @@ //! maps a public model name to its deployment and runs the core route. mod audio_transcription; +mod caching; mod chat_completions; mod error; pub mod messages; @@ -27,6 +28,7 @@ pub use litellm_router::{Deployment, Router as ModelRouter}; pub use request::{JsonObject, RequestId}; pub struct Gateway { + cache: Option>, pub audio_transcription: AudioTranscriptionRoute, pub chat_completions: ChatCompletionsRoute, pub messages: MessagesRoute, @@ -39,6 +41,13 @@ pub struct Gateway { } impl Gateway { + pub fn with_cache(self, cache: Arc) -> Self { + Self { + cache: Some(cache), + ..self + } + } + pub fn new( resources: CoreResources, http: HttpClientConfig, @@ -48,6 +57,7 @@ impl Gateway { let provider = resources.pool.client(&http, ClientVariant::Provider)?; let auth = resources.auth.clone(); Ok(Self { + cache: None, audio_transcription: AudioTranscriptionRoute::new( provider.clone(), auth.clone(), @@ -58,7 +68,11 @@ impl Gateway { auth.clone(), secrets.clone(), ), - messages: MessagesRoute::new(provider.clone(), auth.clone(), secrets.clone()), + messages: MessagesRoute::builder() + .with_http(provider.clone()) + .with_auth(auth.clone()) + .with_secrets(secrets.clone()) + .build(), responses: ResponsesRoute::new(provider, auth.clone(), secrets.clone()), ocr: OcrRoute::new(OcrClient::new( &resources.pool, diff --git a/litellm-rust/crates/gateway-inference/src/messages.rs b/litellm-rust/crates/gateway-inference/src/messages.rs index 14fa0865b6e..6be5921ac6d 100644 --- a/litellm-rust/crates/gateway-inference/src/messages.rs +++ b/litellm-rust/crates/gateway-inference/src/messages.rs @@ -43,11 +43,23 @@ async fn handle( ) -> Result { let deployment = request::resolve_deployment(gateway, &body)?; request::authorize_model(identity, deployment, &body).await?; + let (body, cache_options) = crate::caching::prepare(identity, body)?; + let route = gateway.messages.clone(); + let route = match &gateway.cache { + Some(cache) => route.with_cache(litellm_cache_response::ScopedCache::new( + cache.clone(), + cache_options.scope.clone(), + )), + None => route, + }; + let call = project(deployment, body, headers)?; - let machine = gateway.messages.clone().machine(call, None); + let machine = route.machine(call, cache_options); let stream = Sse::::new(Json, |error| Bytes::from(Error::from(error).sse_frame())); - Ok(litellm_host_http::serve(machine, (), (), stream, None).await?) + let headers = crate::caching::CacheHeaders::default(); + let response = litellm_host_http::serve(machine, (), headers.clone(), stream, None).await?; + Ok(headers.apply(response)) } fn project( diff --git a/litellm-rust/crates/gateway-inference/src/responses.rs b/litellm-rust/crates/gateway-inference/src/responses.rs index 3aca2a0c7f7..5b324d74172 100644 --- a/litellm-rust/crates/gateway-inference/src/responses.rs +++ b/litellm-rust/crates/gateway-inference/src/responses.rs @@ -15,6 +15,16 @@ pub(crate) async fn create( ) -> Result { let deployment = request::resolve_deployment(&gateway, &body)?; request::authorize_model(&identity, deployment, &body).await?; + let (body, cache_options) = crate::caching::prepare(&identity, body)?; + let route = gateway.responses.clone(); + let route = match &gateway.cache { + Some(cache) => route.with_cache(litellm_cache_response::ScopedCache::new( + cache.clone(), + cache_options.scope.clone(), + )), + None => route, + }; + let call = ResponsesCall { model: deployment.model.clone(), input: body.get("input").cloned().unwrap_or_default(), @@ -28,7 +38,7 @@ pub(crate) async fn create( extra_headers: None, timeout: deployment.timeout, }; - let machine = gateway.responses.clone().machine(call, None); + let machine = route.machine(call, cache_options); let stream = Sse::::new(Json, |error| { let error = Error::from(error); Bytes::from(format!( @@ -36,5 +46,7 @@ pub(crate) async fn create( json!({"type": "error", "code": error.status().as_u16().to_string(), "message": error.to_string(), "param": null}) )) }); - Ok(litellm_host_http::serve(machine, (), (), stream, None).await?) + let headers = crate::caching::CacheHeaders::default(); + let response = litellm_host_http::serve(machine, (), headers.clone(), stream, None).await?; + Ok(headers.apply(response)) } diff --git a/litellm-rust/crates/gateway-inference/tests/caching.rs b/litellm-rust/crates/gateway-inference/tests/caching.rs new file mode 100644 index 00000000000..143bc1ba6db --- /dev/null +++ b/litellm-rust/crates/gateway-inference/tests/caching.rs @@ -0,0 +1,185 @@ +mod support; + +use std::{sync::Arc, time::Duration}; + +use axum::body::to_bytes; +use litellm_cache_memory::InMemoryCache; +use litellm_cache_response::{ + CacheKeyInput, ResponseCache, ResponseCacheConfig, ResponseCacheRequest, ResponseCacheService, +}; +use rstest::rstest; +use serde_json::{Value, json}; +use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; + +#[rstest] +#[case::chat("/v1/chat/completions", "anthropic/test-model", false)] +#[case::messages("/v1/messages", "anthropic/test-model", false)] +#[case::responses("/v1/responses", "openai/test-model", false)] +#[case::messages_stream("/v1/messages", "anthropic/test-model", true)] +#[case::responses_stream("/v1/responses", "openai/test-model", true)] +#[tokio::test] +async fn all_inference_endpoints_share_native_cache( + #[case] path: &str, + #[case] model: &str, + #[case] stream: bool, + #[values("s-maxage", "s-max-age")] max_age: &str, +) { + let upstream = MockServer::start().await; + let is_responses = path.ends_with("responses"); + let provider_body = if is_responses { + json!({"id":"response-1", "model":"test-model", "status":"completed", "output":[]}) + } else { + json!({"id":"message-1", "model":"test-model", "type":"message", "role":"assistant", "content":[{"type":"text","text":"hello"}], "stop_reason":"end_turn", "usage":{"input_tokens":1,"output_tokens":1}}) + }; + let terminal = if is_responses { + "response.completed" + } else { + "message_stop" + }; + let events = format!("event: {terminal}\ndata: {{\"type\":\"{terminal}\"}}\n\n"); + let template = if stream { + ResponseTemplate::new(200).set_body_raw(events.clone(), "text/event-stream") + } else { + ResponseTemplate::new(200).set_body_json(provider_body) + }; + Mock::given(method("POST")) + .respond_with(template) + .expect(2) + .mount(&upstream) + .await; + let cache: Arc = Arc::new( + ResponseCache::new(Arc::new(InMemoryCache::new( + Some(100), + Some(Duration::from_secs(60)), + ))) + .with_config(ResponseCacheConfig { + namespace: "gateway-test".into(), + max_entry_bytes: 4096, + }), + ); + let app = support::app_with_cache(model, &upstream.uri(), cache.clone()); + let request = if is_responses { + json!({"model":"public/model", "input":"hello", "stream":stream, "cache":{(max_age):600}}) + } else { + json!({"model":"public/model", "messages":[{"role":"user","content":"hello"}], "max_tokens":16, "stream":stream, "cache":{(max_age):600}}) + }; + let first = support::post(app.clone(), path, request.clone()).await; + assert_eq!(first.status(), 200); + assert!(!first.headers().contains_key("x-litellm-cache-key")); + let first = to_bytes(first.into_body(), 4096).await.unwrap(); + let second = support::post(app.clone(), path, request.clone()).await; + assert_eq!(second.status(), 200); + let cache_key = second.headers().get("x-litellm-cache-key").unwrap().clone(); + assert!(!cache_key.as_bytes().is_empty()); + let stored = cache + .lookup( + &ResponseCacheRequest::new(CacheKeyInput { + preset: Some(cache_key.to_str().unwrap().into()), + ..Default::default() + }), + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap(), + ) + .await + .unwrap(); + assert!( + stored.is_some(), + "the header must identify the stored entry" + ); + let second = to_bytes(second.into_body(), 4096).await.unwrap(); + if stream { + assert_eq!(first, events); + assert_eq!(second, first); + } else { + assert_eq!( + serde_json::from_slice::(&first).unwrap(), + serde_json::from_slice::(&second).unwrap() + ); + } + let bypass_request = Value::Object( + request + .as_object() + .unwrap() + .iter() + .map(|(name, value)| { + ( + name.clone(), + if name == "cache" { + json!({"no-cache": true, "no-store": true}) + } else { + value.clone() + }, + ) + }) + .collect(), + ); + let bypassed = support::post(app.clone(), path, bypass_request).await; + assert_eq!(bypassed.status(), 200); + assert!(!bypassed.headers().contains_key("x-litellm-cache-key")); + to_bytes(bypassed.into_body(), 4096).await.unwrap(); + let restored = support::post(app, path, request).await; + assert_eq!(restored.status(), 200); + assert_eq!( + restored.headers().get("x-litellm-cache-key"), + Some(&cache_key) + ); + assert_eq!(to_bytes(restored.into_body(), 4096).await.unwrap(), second); +} + +#[rstest] +#[case::different_subject("issuer", "tenant-b")] +#[case::different_authority("other-issuer", "tenant-a")] +#[tokio::test] +async fn authenticated_callers_do_not_share_cached_responses( + #[case] authority: &str, + #[case] subject: &str, +) { + use litellm_gateway_auth::{Principal, PrincipalKind}; + + let upstream = MockServer::start().await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "id":"response-1", "model":"test-model", "status":"completed", "output":[] + }))) + .expect(2) + .mount(&upstream) + .await; + let cache: Arc = Arc::new(ResponseCache::new(Arc::new( + InMemoryCache::new(Some(100), Some(Duration::from_secs(60))), + ))); + let first_caller = support::app_with_cache_for_principal( + "openai/test-model", + &upstream.uri(), + cache.clone(), + Principal::new("issuer".into(), "tenant-a".into(), PrincipalKind::Service), + ); + let second_caller = support::app_with_cache_for_principal( + "openai/test-model", + &upstream.uri(), + cache, + Principal::new(authority.into(), subject.into(), PrincipalKind::Service), + ); + let body = json!({"model":"public/model","input":"same prompt"}); + let first = support::post(first_caller.clone(), "/v1/responses", body.clone()).await; + assert_eq!(first.status(), 200); + assert!(!first.headers().contains_key("x-litellm-cache-key")); + let first_hit = support::post(first_caller.clone(), "/v1/responses", body.clone()).await; + assert_eq!(first_hit.status(), 200); + let first_key = first_hit.headers().get("x-litellm-cache-key").unwrap(); + let second = support::post(second_caller.clone(), "/v1/responses", body.clone()).await; + assert_eq!(second.status(), 200); + assert!(!second.headers().contains_key("x-litellm-cache-key")); + let second_hit = support::post(second_caller, "/v1/responses", body.clone()).await; + assert_eq!(second_hit.status(), 200); + assert_ne!( + second_hit.headers().get("x-litellm-cache-key").unwrap(), + first_key + ); + let first_again = support::post(first_caller, "/v1/responses", body).await; + assert_eq!(first_again.status(), 200); + assert_eq!( + first_again.headers().get("x-litellm-cache-key"), + Some(first_key) + ); +} diff --git a/litellm-rust/crates/gateway-inference/tests/support/mod.rs b/litellm-rust/crates/gateway-inference/tests/support/mod.rs index b59e335334e..30c009e86f8 100644 --- a/litellm-rust/crates/gateway-inference/tests/support/mod.rs +++ b/litellm-rust/crates/gateway-inference/tests/support/mod.rs @@ -1,3 +1,6 @@ +// Shared across integration-test targets; each target uses a different subset. +#![allow(dead_code)] + use std::{sync::Arc, time::Duration}; use axum::{ @@ -33,39 +36,83 @@ pub fn app_with_permissions( model: &str, api_base: &str, permissions: litellm_gateway_auth::Permissions, +) -> Router { + configured_app(model, api_base, permissions, None, None) +} + +pub fn app_with_cache( + model: &str, + api_base: &str, + cache: Arc, +) -> Router { + configured_app( + model, + api_base, + litellm_gateway_auth::Permissions::All, + Some(cache), + None, + ) +} + +pub fn app_with_cache_for_principal( + model: &str, + api_base: &str, + cache: Arc, + principal: litellm_gateway_auth::Principal, +) -> Router { + configured_app( + model, + api_base, + litellm_gateway_auth::Permissions::All, + Some(cache), + Some(principal), + ) +} + +fn configured_app( + model: &str, + api_base: &str, + permissions: litellm_gateway_auth::Permissions, + cache: Option>, + principal: Option, ) -> Router { let pool = Arc::new(HttpClientPool::new(Arc::new(PublicDnsResolver))); let http = Resolution::from(&HttpSettings::default()).config; let secrets = Arc::new(NoSecrets); let resources = CoreResources::new(pool); - router(Arc::new( - Gateway::new( - resources, - http, - secrets, - [( - "public/model".into(), - Deployment { - model: model.into(), - api_base: Some(api_base.into()), - api_key: Some("test-key".into()), - timeout: Some(Duration::from_secs(5)), - ..Default::default() - }, - )] - .into_iter() - .collect(), - ) - .unwrap(), - )) - .layer(axum::middleware::from_fn_with_state( - permissions, + let gateway = Gateway::new( + resources, + http, + secrets, + [( + "public/model".into(), + Deployment { + model: model.into(), + api_base: Some(api_base.into()), + api_key: Some("test-key".into()), + timeout: Some(Duration::from_secs(5)), + ..Default::default() + }, + )] + .into_iter() + .collect(), + ) + .unwrap(); + let gateway = match cache { + Some(cache) => gateway.with_cache(cache), + None => gateway, + }; + router(Arc::new(gateway)).layer(axum::middleware::from_fn_with_state( + (permissions, principal), test_identity, )) } async fn test_identity( - axum::extract::State(permissions): axum::extract::State, + axum::extract::State((permissions, principal)): axum::extract::State<( + litellm_gateway_auth::Permissions, + Option, + )>, mut request: axum::extract::Request, next: axum::middleware::Next, ) -> Response { @@ -74,7 +121,7 @@ async fn test_identity( Some(SecretValue::new("test-inbound-key")), Arc::new(NoSecrets), )), - Arc::new(TestPermissions(permissions)), + Arc::new(TestPermissions(permissions, principal)), Arc::new(litellm_gateway_auth::NoAdditionalPolicy), Arc::new(litellm_gateway_auth::SystemClock), ); @@ -101,7 +148,10 @@ pub async fn json(response: Response) -> Value { serde_json::from_slice(&to_bytes(response.into_body(), 1024 * 1024).await.unwrap()).unwrap() } -struct TestPermissions(litellm_gateway_auth::Permissions); +struct TestPermissions( + litellm_gateway_auth::Permissions, + Option, +); impl litellm_gateway_auth::IdentityResolver for TestPermissions { fn resolve<'a>( @@ -110,7 +160,7 @@ impl litellm_gateway_auth::IdentityResolver for TestPermissions { ) -> litellm_gateway_auth::AuthFuture<'a, litellm_gateway_auth::ResolvedIdentity> { Box::pin(async move { Ok(litellm_gateway_auth::ResolvedIdentity { - principal: identity.principal.clone(), + principal: self.1.clone().unwrap_or_else(|| identity.principal.clone()), permissions: self.0.clone(), }) }) diff --git a/litellm-rust/crates/host-native/src/driver.rs b/litellm-rust/crates/host-native/src/driver.rs index b559de6e29f..7254721da87 100644 --- a/litellm-rust/crates/host-native/src/driver.rs +++ b/litellm-rust/crates/host-native/src/driver.rs @@ -66,6 +66,11 @@ where MachineStep::Suspended(request) => request, }; let answered = match request { + HostRequest::Intercept(InterceptRequest::ResultReady { facts, reply }) => self + .interceptors + .result_ready(facts) + .await + .map(|()| reply.send(())), HostRequest::HostCall(call) => self.services.handle_host_call(call).await, HostRequest::Intercept(InterceptRequest::BeforeProviderRequest { wire, diff --git a/litellm-rust/crates/host-python/src/driver.rs b/litellm-rust/crates/host-python/src/driver.rs index 07586e9db07..c70f2a5f1a0 100644 --- a/litellm-rust/crates/host-python/src/driver.rs +++ b/litellm-rust/crates/host-python/src/driver.rs @@ -64,6 +64,7 @@ enum EventNext { } enum Pending { + Host, Native, Arguments(HookResume>), Wire(HookResume>, Reply), @@ -199,6 +200,19 @@ where Err(error) => self.hook_failed(py, error), } } + (Some(Pending::Host), Some(result)) => { + match self.binding.resume_host_call(py, result) { + Ok(Some(awaitable)) => { + self.pending = Some(Pending::Host); + Ok(ExecutionStep::Await(awaitable)) + } + Ok(None) => self.resume_machine(py, None), + Err(InvokeError::Python(error)) => self.interrupt(py, error), + Err(InvokeError::Native(error)) => { + self.resume_machine(py, Some(HostFailure::Error(error))) + } + } + } (Some(Pending::Native), Some(Ok(_))) => { let result = self.native.take_result()?; self.run_steps(py, NativePoll::Ready(result)) @@ -405,7 +419,13 @@ where Err(error) => return self.machine_failed(py, error).map(Next::Return), }; let answered = match op { - HostRequest::HostCall(op) => answered(self.binding.handle_host_call(py, op)), + HostRequest::HostCall(op) => match self.binding.begin_host_call(py, op) { + Ok(Some(awaitable)) => { + self.pending = Some(Pending::Host); + return Ok(Next::Return(ExecutionStep::Await(awaitable))); + } + result => answered(result.map(|_| ())), + }, HostRequest::Intercept(InterceptRequest::BeforeProviderRequest { wire, context, @@ -427,6 +447,21 @@ where HostRequest::Stream(StreamDelivery::Chunk(chunk, reply)) => { return self.delivered(py, chunk, reply).map(Next::Return); } + HostRequest::Intercept(InterceptRequest::ResultReady { facts, reply }) => { + let event = PythonCallEvent::Execution(ExecutionEvent::ResultReady { facts }); + self.observe(&event); + match self.hooks.on_event(py, event) { + Ok(HookStep::Ready(())) => { + reply.send(()); + Ok(Ok(())) + } + Ok(HookStep::Await(awaitable, resume)) => { + self.pending = Some(Pending::Event(resume, EventNext::Emitted(reply))); + return Ok(Next::Return(ExecutionStep::Await(awaitable))); + } + Err(error) => Err(error), + } + } HostRequest::Intercept(InterceptRequest::AfterProviderResponse { raw, reply }) => { let event = PythonCallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw: &raw, @@ -466,7 +501,7 @@ where Ok(head) => head, Err(error) => return self.interrupt(py, error), }; - match self.hooks.on_stream_open(py) { + match self.hooks.on_stream_open(py, &head) { Ok(()) => { self.pending = Some(Pending::Consumer(reply)); Ok(ExecutionStep::Open(head)) @@ -770,12 +805,15 @@ mod tests { RejectNatively, RejectRequestNatively, RaiseRequestPython, + AwaitAnswer, + AwaitFailure, } struct SyntheticBinding { log: Log, op: OpScript, classifier_fails: bool, + pending_reply: Option>, } /// The fake route's public exception, kept as a value so a test sees what `classify` @@ -792,7 +830,7 @@ mod tests { impl SyntheticBinding { fn answer(&self, value: impl FnOnce() -> String) -> Result> { match self.op { - OpScript::Answer => Ok(value()), + OpScript::Answer | OpScript::AwaitAnswer | OpScript::AwaitFailure => Ok(value()), OpScript::RaisePython | OpScript::RaiseRequestPython => { Err(PyValueError::new_err("op failed").into()) } @@ -867,6 +905,41 @@ mod tests { self.answer(|| op.to_string()) .map(|answer| reply.send(answer)) } + + fn begin_host_call( + &mut self, + py: Python<'_>, + (op, reply): (&'static str, Reply), + ) -> Result>, InvokeError> { + if !matches!(self.op, OpScript::AwaitAnswer | OpScript::AwaitFailure) { + return self.handle_host_call(py, (op, reply)).map(|()| None); + } + self.pending_reply = Some(reply); + let module = PyModule::from_code( + py, + pyo3::ffi::c_str!( + "async def answer(fail):\n if fail:\n raise LookupError('async host failed')\n return 'awaited'\n" + ), + pyo3::ffi::c_str!("host_op.py"), + pyo3::ffi::c_str!("host_op"), + )?; + Ok(Some( + module + .getattr("answer")? + .call1((matches!(self.op, OpScript::AwaitFailure),))? + .unbind(), + )) + } + + fn resume_host_call( + &mut self, + py: Python<'_>, + result: PyResult>, + ) -> Result>, InvokeError> { + let answer = result?.extract::(py)?; + self.pending_reply.take().unwrap().send(answer); + Ok(None) + } } impl PythonOwned for SyntheticBinding { @@ -988,6 +1061,9 @@ mod tests { return Err(PyValueError::new_err("callback failed")); } self.log.push(match event { + PythonCallEvent::Execution(ExecutionEvent::ResultReady { .. }) => { + "cache_hit".into() + } PythonCallEvent::Started { .. } => "started".into(), PythonCallEvent::Cancelled { .. } => "cancelled".into(), PythonCallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => { @@ -1003,7 +1079,7 @@ mod tests { Ok(HookStep::Ready(())) } - fn on_stream_open(&mut self, _: Python<'_>) -> PyResult<()> { + fn on_stream_open(&mut self, _: Python<'_>, _: &Py) -> PyResult<()> { self.log.push("opened"); Ok(()) } @@ -1037,6 +1113,7 @@ mod tests { log: Log::default(), op, classifier_fails: false, + pending_reply: None, }, script, asynchronous, @@ -1060,6 +1137,50 @@ mod tests { ) } + #[rstest::rstest] + #[case::success(OpScript::AwaitAnswer)] + #[case::failure(OpScript::AwaitFailure)] + fn asynchronous_host_operations_resume_the_same_machine(#[case] op: OpScript) { + let _guard = PYTHON_GLOBALS + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + Python::initialize(); + Python::attach(|py| { + install_lifecycle_module(py); + let (result, log) = run_scripted( + py, + |_| { + CallMachine::::new(None, |host| { + Box::pin(async move { host.services.call(|reply| ("read", reply)).await }) + }) + }, + op, + HookScript::Plain, + true, + ); + match op { + OpScript::AwaitAnswer => { + assert_eq!(result.unwrap().extract::(py).unwrap(), "awaited") + } + OpScript::AwaitFailure => { + assert!( + result + .unwrap_err() + .is_instance_of::(py) + ); + assert_eq!( + log.iter() + .filter(|entry| entry.starts_with("failed:")) + .count(), + 1 + ); + assert!(!log.iter().any(|entry| entry.starts_with("succeeded:"))); + } + _ => unreachable!(), + } + }); + } + #[rstest::rstest] #[case::synchronous(false)] #[case::asynchronous(true)] @@ -1086,6 +1207,7 @@ mod tests { log: Log(log.0.clone()), op: OpScript::Answer, classifier_fails: false, + pending_reply: None, }, hooks, PyDict::new(py).unbind(), @@ -1223,6 +1345,7 @@ mod tests { log: Log::default(), op: OpScript::Answer, classifier_fails: false, + pending_reply: None, }, HookScript::ReplaceResponse, std::convert::identity, @@ -1282,6 +1405,7 @@ mod tests { log: Log::default(), op: OpScript::Answer, classifier_fails: false, + pending_reply: None, }, script, std::convert::identity, @@ -1787,6 +1911,7 @@ mod tests { log: Log::default(), op: OpScript::Answer, classifier_fails: true, + pending_reply: None, }, HookScript::Plain, false, @@ -1902,6 +2027,7 @@ mod tests { log: Log(log.0.clone()), op: OpScript::Answer, classifier_fails: false, + pending_reply: None, }, crate::HookChain::new() .with(SyntheticHooks { @@ -1984,6 +2110,7 @@ mod tests { log: Log(log.0.clone()), op: OpScript::Answer, classifier_fails: false, + pending_reply: None, }, SyntheticHooks { log: Log(log.0.clone()), @@ -2020,6 +2147,7 @@ mod tests { log: Log::default(), op: OpScript::Answer, classifier_fails: false, + pending_reply: None, }, HookScript::Plain, |hooks| { @@ -2062,6 +2190,7 @@ mod tests { log: Log::default(), op: OpScript::Answer, classifier_fails: false, + pending_reply: None, }, HookScript::Plain, |hooks| { diff --git a/litellm-rust/crates/host-python/src/hooks/chain/adapter.rs b/litellm-rust/crates/host-python/src/hooks/chain/adapter.rs index 87f77c8428e..b08a0e4ebcf 100644 --- a/litellm-rust/crates/host-python/src/hooks/chain/adapter.rs +++ b/litellm-rust/crates/host-python/src/hooks/chain/adapter.rs @@ -89,7 +89,7 @@ pub(super) trait ChainHooks: PythonOwned { result: PyResult>, ) -> PyResult>; fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py) -> PyResult<()>; - fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()>; + fn on_stream_open(&mut self, py: Python<'_>, head: &Py) -> PyResult<()>; fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py) -> PyResult<()>; } @@ -181,8 +181,8 @@ impl ChainHooks for HookAdapter { self.hooks.arguments_prepared(py, arguments) } - fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> { - self.hooks.on_stream_open(py) + fn on_stream_open(&mut self, py: Python<'_>, head: &Py) -> PyResult<()> { + self.hooks.on_stream_open(py, head) } fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py) -> PyResult<()> { diff --git a/litellm-rust/crates/host-python/src/hooks/chain/dispatch.rs b/litellm-rust/crates/host-python/src/hooks/chain/dispatch.rs index 2f7ed0d778a..968345bde97 100644 --- a/litellm-rust/crates/host-python/src/hooks/chain/dispatch.rs +++ b/litellm-rust/crates/host-python/src/hooks/chain/dispatch.rs @@ -236,6 +236,9 @@ fn notification_result( fn retain_event(py: Python<'_>, event: PythonCallEvent<'_>) -> OwnedEvent { match event { + CallEvent::Execution(ExecutionEvent::ResultReady { facts }) => { + CallEvent::Execution(ExecutionEvent::ResultReady { facts }) + } CallEvent::Started { start_time } => CallEvent::Started { start_time }, CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => { CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw: raw.clone() }) @@ -263,6 +266,12 @@ fn dispatch( event: &OwnedEvent, ) -> PyResult> { match event { + CallEvent::Execution(ExecutionEvent::ResultReady { facts }) => hooks.on_event( + py, + CallEvent::Execution(ExecutionEvent::ResultReady { + facts: facts.clone(), + }), + ), CallEvent::Started { start_time } => hooks.on_event( py, CallEvent::Started { @@ -350,10 +359,10 @@ impl CallHooks for HookChain { Ok(HookStep::Ready(())) } - fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> { + fn on_stream_open(&mut self, py: Python<'_>, head: &Py) -> PyResult<()> { self.hooks .iter_mut() - .try_for_each(|hooks| hooks.on_stream_open(py)) + .try_for_each(|hooks| hooks.on_stream_open(py, head)) } fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py) -> PyResult<()> { diff --git a/litellm-rust/crates/host-python/src/services.rs b/litellm-rust/crates/host-python/src/services.rs index 0a06118926e..3d8faa2a5aa 100644 --- a/litellm-rust/crates/host-python/src/services.rs +++ b/litellm-rust/crates/host-python/src/services.rs @@ -8,4 +8,20 @@ pub trait PythonHostCalls: PythonOwned { py: Python<'_>, call: P::HostCall, ) -> Result<(), InvokeError>; + + fn begin_host_call( + &mut self, + py: Python<'_>, + call: P::HostCall, + ) -> Result>, InvokeError> { + self.handle_host_call(py, call).map(|()| None) + } + + fn resume_host_call( + &mut self, + _: Python<'_>, + result: PyResult>, + ) -> Result>, InvokeError> { + result.map(|_| None).map_err(InvokeError::Python) + } } diff --git a/litellm-rust/crates/host-python/tests/hook_chain.rs b/litellm-rust/crates/host-python/tests/hook_chain.rs index c1efa941227..ca28bce5d17 100644 --- a/litellm-rust/crates/host-python/tests/hook_chain.rs +++ b/litellm-rust/crates/host-python/tests/hook_chain.rs @@ -158,7 +158,7 @@ impl CallHooks for ScriptHooks { } } - fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> { + fn on_stream_open(&mut self, py: Python<'_>, _head: &Py) -> PyResult<()> { self.object.call_method1(py, "stream", (py.None(),))?; Ok(()) } @@ -307,7 +307,7 @@ fn transformations_feed_each_other_and_notifications_share_final_values( ) .unwrap(); finish(py, &mut hooks, step).unwrap(); - hooks.on_stream_open(py).unwrap(); + hooks.on_stream_open(py, &py.None()).unwrap(); hooks.on_stream_chunk(py, &response).unwrap(); let locals = scripts.bind(py); locals.set_item("arguments", arguments).unwrap(); diff --git a/litellm-rust/crates/host/AGENTS.md b/litellm-rust/crates/host/AGENTS.md index e2034e76c6f..317a43516c7 100644 --- a/litellm-rust/crates/host/AGENTS.md +++ b/litellm-rust/crates/host/AGENTS.md @@ -25,3 +25,5 @@ Rust handlers answer suspensions through `litellm-host-native::Driver`, which `l Keep API policy in gateway-inference and python-bridge, and legacy callback policy in callbacks-legacy-python. Python bindings and hooks expose retained references through `PythonOwned`, with idempotent close and GC traversal. Runtime machinery stays in driver, native, handle and runtime modules Interceptors run inline and can rewrite values or fail execution. Observers consume owned `CallEvent` snapshots from `observation_channel`; its bounded `ObservationSender` never waits for delivery and counts events dropped when the queue is full or closed. The host owns receiver processing and draining. Pass the same publisher to machine construction and the driver when one receiver should collect execution and lifecycle events. Legacy Python callbacks retain their existing awaited, fallible behavior through the Python adapter + +`ExecutionFacts` and `ResultSource` describe execution without pricing or budget policy. `Interceptors::result_ready` delivers these facts through an awaited `InterceptRequest::ResultReady`; hosts receive them before response transformation or stream delivery. `ExecutionEvent::ResultReady` is the matching lifecycle event and can also be published as a passive snapshot. Accounting must consume the awaited path rather than a lossy observation queue diff --git a/litellm-rust/crates/host/src/hooks.rs b/litellm-rust/crates/host/src/hooks.rs index 4444aeec551..da16bef258c 100644 --- a/litellm-rust/crates/host/src/hooks.rs +++ b/litellm-rust/crates/host/src/hooks.rs @@ -61,7 +61,11 @@ pub trait CallHooks: Sized { Ok(R::ready(())) } - fn on_stream_open(&mut self, _runtime: R::Context<'_>) -> Result<(), R::Error> { + fn on_stream_open( + &mut self, + _runtime: R::Context<'_>, + _head: &R::Response, + ) -> Result<(), R::Error> { Ok(()) } diff --git a/litellm-rust/crates/host/src/interceptors.rs b/litellm-rust/crates/host/src/interceptors.rs index 0044e842e4a..3633fc7ee39 100644 --- a/litellm-rust/crates/host/src/interceptors.rs +++ b/litellm-rust/crates/host/src/interceptors.rs @@ -30,7 +30,29 @@ pub struct RawResponse { pub body: String, } +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct ProviderIdentity { + pub model: String, + pub provider: String, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum ResultSource { + Provider, + Cache { key: String }, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct ExecutionFacts { + pub provider: ProviderIdentity, + pub source: ResultSource, +} + pub trait Interceptors: Send + Sync { + fn result_ready(&self, _facts: ExecutionFacts) -> impl Future> + Send { + async { Ok(()) } + } + fn before_provider_request( &self, wire: WireRequest, @@ -44,6 +66,10 @@ pub trait Interceptors: Send + Sync { } impl + ?Sized> Interceptors for &T { + fn result_ready(&self, facts: ExecutionFacts) -> impl Future> + Send { + (**self).result_ready(facts) + } + fn before_provider_request( &self, wire: WireRequest, diff --git a/litellm-rust/crates/host/src/lifecycle.rs b/litellm-rust/crates/host/src/lifecycle.rs index f16a99fbcef..e6b31381494 100644 --- a/litellm-rust/crates/host/src/lifecycle.rs +++ b/litellm-rust/crates/host/src/lifecycle.rs @@ -52,7 +52,12 @@ pub enum CallEvent { #[derive(Clone, Debug, PartialEq, Eq)] pub enum ExecutionEvent { - ProviderResponseReceived { raw: Raw }, + ResultReady { + facts: crate::interceptors::ExecutionFacts, + }, + ProviderResponseReceived { + raw: Raw, + }, } impl> CallEvent { @@ -66,6 +71,11 @@ impl> CallEvent { + CallEvent::Execution(ExecutionEvent::ResultReady { + facts: facts.clone(), + }) + } Self::Succeeded { timing, .. } => CallEvent::Succeeded { timing: *timing, response: (), diff --git a/litellm-rust/crates/host/src/machine/context.rs b/litellm-rust/crates/host/src/machine/context.rs index 39b2b82d502..b2c9e95c563 100644 --- a/litellm-rust/crates/host/src/machine/context.rs +++ b/litellm-rust/crates/host/src/machine/context.rs @@ -101,6 +101,17 @@ where .await } + async fn result_ready( + &self, + facts: crate::interceptors::ExecutionFacts, + ) -> Result<(), P::Error> { + self.0 + .request_reply(|reply| { + HostRequest::Intercept(InterceptRequest::ResultReady { facts, reply }) + }) + .await + } + async fn after_provider_response(&self, raw: RawResponse) -> Result<(), P::Error> { self.0 .request_reply(|reply| { diff --git a/litellm-rust/crates/host/src/protocol.rs b/litellm-rust/crates/host/src/protocol.rs index c9db375cde3..9e5aff8aae8 100644 --- a/litellm-rust/crates/host/src/protocol.rs +++ b/litellm-rust/crates/host/src/protocol.rs @@ -20,6 +20,10 @@ pub enum HostRequest { } pub enum InterceptRequest { + ResultReady { + facts: crate::interceptors::ExecutionFacts, + reply: Reply<()>, + }, BeforeProviderRequest { wire: Box, context: Box, diff --git a/litellm-rust/crates/python-bridge/src/cache/AGENTS.md b/litellm-rust/crates/python-bridge/src/cache/AGENTS.md index 8c6f31a7780..63223a6816c 100644 --- a/litellm-rust/crates/python-bridge/src/cache/AGENTS.md +++ b/litellm-rust/crates/python-bridge/src/cache/AGENTS.md @@ -2,6 +2,10 @@ This folder owns Python cache API compatibility: argument projection, facade identity, public result construction, Python embedding calls and per-operation composition of native backends. Cache algorithms, storage protocols and response-cache semantics belong to their cache crates +`mod.rs` exposes the cache boundary to routes and module registration; adapter directories remain private. `selection.rs` owns global cache selection, route admission and inference protocol composition for both adapters. `runtime.rs` exposes the Python-facing runtime that can wrap either native storage or a Python callback. `future.rs` converts cache results into ready Futures + +`python/` delegates operations to the selected Python cache without discovering configuration. `native/` owns native backend construction, configuration projection, facade validation, embedding and storage bindings, including experimental V2 handles. Neither adapter depends on shared selection or the other adapter. Shared composition depends on the adapters, and routes use only the parent module's exports + `SemanticExecution` belongs here because its steps select cache operations and invoke the Python embedder. Use the shared `Execution` handle and inline lifecycle driver; do not duplicate coroutine state validation, runtime waiting or GIL machinery. Python embedding awaits stay in the caller's task, and cancellation must prevent later backend or batch operations from starting Resolved asyncio Future construction is generic host machinery. Use `litellm-host-python::ready_future` with an already constructed Python value. Keep cache-specific conversion and disabled-cache return values here. Preserve the Future-returning API and running-loop requirement diff --git a/litellm-rust/crates/python-bridge/src/cache/handle.rs b/litellm-rust/crates/python-bridge/src/cache/handle.rs deleted file mode 100644 index fcc8aa6218a..00000000000 --- a/litellm-rust/crates/python-bridge/src/cache/handle.rs +++ /dev/null @@ -1,363 +0,0 @@ -use crate::http::host_client; -use crate::logger::run_sync_value; -use litellm_auth_aws::AwsAuthConfig; -use litellm_cache_gcs::{DEFAULT_ENDPOINT, GcsConfig}; -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; -use litellm_http::ClientVariant; -use pyo3::{ - PyTraverseError, PyVisit, - exceptions::{PyRuntimeError, PyTypeError}, - prelude::*, -}; -use url::Url; - -use super::{ - cache_error, - config::{QdrantSemanticCacheConfig, project_redis_semantic}, - embedder::PythonEmbedder, - facade::FacadeGuard, - native::NativeResponseCache, - request::duration, -}; - -#[pyclass(frozen, name = "_CacheTestHandle")] -pub(crate) struct CacheTestHandle { - service: NativeResponseCache, - pub(super) guard: Option, - pid: u32, -} - -impl CacheTestHandle { - pub(super) fn service(&self) -> PyResult { - if self.pid != std::process::id() { - return Err(PyRuntimeError::new_err( - "native cache handles must be recreated after fork", - )); - } - Ok(self.service.clone()) - } -} - -#[pymethods] -impl CacheTestHandle { - #[staticmethod] - #[pyo3(signature = (*, capacity=200, ttl_seconds=600.0, max_entry_bytes=1048576))] - fn memory(capacity: usize, ttl_seconds: f64, max_entry_bytes: usize) -> PyResult { - Ok(Self { - service: NativeResponseCache::memory(capacity, duration(ttl_seconds)?, max_entry_bytes), - guard: None, - pid: std::process::id(), - }) - } - - #[staticmethod] - #[pyo3(signature = (url, *, ttl_seconds=60.0, namespace=None, startup_nodes=None))] - fn redis( - py: Python<'_>, - url: String, - ttl_seconds: f64, - namespace: Option, - startup_nodes: Option>, - ) -> PyResult { - let ttl = Some(duration(ttl_seconds)?); - let topology = match startup_nodes { - None => RedisTopology::Standalone, - Some(nodes) => RedisTopology::Cluster { - startup_nodes: nodes - .into_iter() - .map(|(host, port)| RedisNode { host, port }) - .collect(), - }, - }; - let service = release_gil(py, move || { - NativeResponseCache::redis(&url, &topology, ttl, namespace) - }) - .map_err(cache_error)?; - Ok(Self { - service, - guard: None, - pid: std::process::id(), - }) - } - - #[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 http = host_client(py, ClientVariant::NoRedirect)?; - let service = run_sync_value(py, async move { - Ok(NativeResponseCache::s3(config, http).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( - py: Python<'_>, - bucket_name: String, - gcs_path: Option, - path_service_account: Option, - endpoint: Option, - token: Option, - ) -> PyResult { - let config = GcsConfig { - bucket_name, - gcs_path, - path_service_account, - endpoint: endpoint.unwrap_or_else(|| DEFAULT_ENDPOINT.to_string()), - }; - let client = host_client(py, ClientVariant::NoRedirect)?; - let service = NativeResponseCache::gcs(config, client, token); - Ok(Self { - service, - guard: None, - pid: std::process::id(), - }) - } - - #[staticmethod] - #[pyo3(signature = (directory))] - fn disk(py: Python<'_>, directory: String) -> PyResult { - let service = - release_gil(py, move || NativeResponseCache::disk(&directory)).map_err(cache_error)?; - Ok(Self { - service, - guard: None, - pid: std::process::id(), - }) - } - - #[staticmethod] - #[pyo3(signature = (url, *, collection_name, similarity_threshold, vector_size, embedding_model="text-embedding-3-small", api_key=None, embedding_api_key=None, embedding_api_base=None, embedding_timeout_seconds=None, quantization="binary"))] - #[expect( - clippy::too_many_arguments, - reason = "the test handle exposes the complete Qdrant constructor" - )] - fn qdrant_semantic( - py: Python<'_>, - url: String, - collection_name: String, - similarity_threshold: f64, - vector_size: u64, - embedding_model: &str, - api_key: Option, - embedding_api_key: Option, - embedding_api_base: Option, - embedding_timeout_seconds: Option, - quantization: &str, - ) -> PyResult { - let parsed = Url::parse(&url).map_err(|_| { - pyo3::exceptions::PyValueError::new_err( - "native Qdrant requires the default REST port so the gRPC port can be derived", - ) - })?; - if !matches!(parsed.scheme(), "http" | "https") - || (!parsed.path().is_empty() && parsed.path() != "/") - || parsed.query().is_some() - || parsed.host_str().is_none() - || 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", - )); - } - let mut grpc_url = parsed; - grpc_url.set_port(Some(6334)).map_err(|_| { - pyo3::exceptions::PyValueError::new_err( - "native Qdrant requires the default REST port so the gRPC port can be derived", - ) - })?; - grpc_url.set_path(""); - grpc_url.set_query(None); - let embedding_api_key = embedding_api_key - .or_else(|| { - std::env::var("OPENAI_API_KEY") - .ok() - .filter(|value| !value.is_empty()) - }) - .ok_or_else(|| { - pyo3::exceptions::PyValueError::new_err( - "native semantic embedding requires an OpenAI API key", - ) - })?; - let embedding_api_base = embedding_api_base.unwrap_or_else(|| { - 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 { - "binary" => Quantization::Binary, - "scalar" => Quantization::Scalar, - "product" => Quantization::Product, - _ => { - return Err(pyo3::exceptions::PyValueError::new_err( - "unsupported Qdrant quantization", - )); - } - }; - let config = QdrantSemanticCacheConfig { - grpc_url: grpc_url.to_string().trim_end_matches('/').to_owned(), - api_key, - collection_name, - similarity_threshold, - vector_size, - embedding: OpenAiEmbedderConfig { - api_base: embedding_api_base, - api_key: embedding_api_key, - model: embedding_model.to_owned(), - timeout: embedding_timeout_seconds.map(duration).transpose()?, - }, - quantization, - }; - let client = host_client(py, ClientVariant::Provider)?; - let service = run_sync_value(py, async move { - let handle = tokio::runtime::Handle::current(); - NativeResponseCache::qdrant_semantic(config, client, handle) - .await - .map_err(cache_error) - })?; - Ok(Self { - service, - guard: None, - pid: std::process::id(), - }) - } - - #[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::new(embedder.clone().unbind()); - 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 { - let http = host_client(py, ClientVariant::NoRedirect)?; - let service = run_sync_value(py, async move { - NativeResponseCache::azure_blob(&account_url, &container, http) - .await - .map_err(cache_error) - })?; - Ok(Self { - service, - guard: None, - pid: std::process::id(), - }) - } - - #[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() - } - - 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_scope( - facade - .getattr("semantic_cache_scope")? - .extract::()?, - ) - .with_redis_flush_size( - facade - .getattr("redis_flush_size")? - .extract::>()?, - ); - let handle = Py::new( - py, - Self { - service, - guard: Some(guard), - pid: self.pid, - }, - )?; - facade.setattr("_native_cache_handle", handle) - } - - fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> { - self.service.traverse(&visit)?; - if let Some(guard) = &self.guard { - guard.traverse(visit)?; - } - Ok(()) - } -} diff --git a/litellm-rust/crates/python-bridge/src/cache/mod.rs b/litellm-rust/crates/python-bridge/src/cache/mod.rs index 00b0c71684a..179f16c4a1f 100644 --- a/litellm-rust/crates/python-bridge/src/cache/mod.rs +++ b/litellm-rust/crates/python-bridge/src/cache/mod.rs @@ -1,16 +1,13 @@ -mod activation; -mod binding; -mod callback; -mod config; -mod embedder; -mod facade; mod future; -mod handle; -mod identity; mod native; -mod request; -mod resolver; -mod semantic; +mod python; +mod runtime; +mod selection; + +pub(crate) use native::NativeCacheHandle; +pub(crate) use python::{CacheCall, PythonCache}; +pub(crate) use runtime::ResolvedCache; +pub(crate) use selection::{Cached, Selection, admit_native, configure, configured_native}; use litellm_cache::Error; use pyo3::{ @@ -18,8 +15,6 @@ use pyo3::{ prelude::*, }; -pub(crate) use self::{binding::ResolvedCache, handle::CacheTestHandle, resolver::CacheResolver}; - fn cache_error(error: Error) -> PyErr { match error { Error::InvalidEntry => PyValueError::new_err(error.to_string()), diff --git a/litellm-rust/crates/python-bridge/src/cache/native/AGENTS.md b/litellm-rust/crates/python-bridge/src/cache/native/AGENTS.md new file mode 100644 index 00000000000..65804a18564 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/cache/native/AGENTS.md @@ -0,0 +1,9 @@ +# Native cache bindings + +This directory constructs and exposes Rust cache backends to Python. It owns backend configuration projection, facade validation, native request conversion, semantic embedding integration and experimental V2 handles. Cache algorithms and storage protocols remain in their cache crates + +Accept the cache object or projected configuration selected by the parent module. Do not read global `litellm.cache`, decide route admission, or select the Python cache adapter here + +Keep Python-facing cache classes and method signatures stable when reorganizing modules. Native internals stay private to this directory unless the shared cache boundary or Python module registration needs them. Python embedding awaits use the existing inline lifecycle driver, preserving caller task identity and cancellation + +Verify changes with the existing backend and facade tests using a freshly built extension. Test cache behavior, not module paths or file structure diff --git a/litellm-rust/crates/python-bridge/src/cache/activation.rs b/litellm-rust/crates/python-bridge/src/cache/native/activation.rs similarity index 97% rename from litellm-rust/crates/python-bridge/src/cache/activation.rs rename to litellm-rust/crates/python-bridge/src/cache/native/activation.rs index 58735679554..20f8179ffb0 100644 --- a/litellm-rust/crates/python-bridge/src/cache/activation.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native/activation.rs @@ -1,3 +1,4 @@ +use crate::cache::cache_error; use crate::logger::run_sync_value; use litellm_cache_gcs::{DEFAULT_ENDPOINT, GcsConfig}; use litellm_cache_redis_semantic::RedisSemanticConfig; @@ -6,10 +7,9 @@ use litellm_http::ClientVariant; use pyo3::prelude::*; use super::{ - cache_error, + backend::NativeResponseCache, config::{CacheBackendConfig, NativeCacheConfig, UnsupportedCacheConfig}, embedder::PythonEmbedder, - native::NativeResponseCache, }; use crate::errors::RustBridgeDeclined; use crate::http::host_client; @@ -20,7 +20,7 @@ fn declined(reason: UnsupportedCacheConfig) -> PyErr { /// Builds the native backend a `Cache` facade's projected configuration describes. `backend` is /// the facade's `.cache` object, which owns embedding for the Python-embedded semantic caches. -pub(super) fn activate( +pub(in crate::cache) fn activate( py: Python<'_>, backend: &Bound<'_, PyAny>, config: NativeCacheConfig, diff --git a/litellm-rust/crates/python-bridge/src/cache/native.rs b/litellm-rust/crates/python-bridge/src/cache/native/backend.rs similarity index 95% rename from litellm-rust/crates/python-bridge/src/cache/native.rs rename to litellm-rust/crates/python-bridge/src/cache/native/backend.rs index 460136baa1f..03897d0ddf1 100644 --- a/litellm-rust/crates/python-bridge/src/cache/native.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native/backend.rs @@ -1,3 +1,4 @@ +use crate::cache::cache_error; use std::{sync::Arc, time::Duration}; use litellm_cache::{CacheCodec, CacheConnectionResult, Error, semantic::SemanticLookup}; @@ -14,7 +15,7 @@ use litellm_cache_response::{ }; use litellm_cache_s3::{S3Cache, S3CacheConfig}; use litellm_cache_valkey_semantic::{ValkeySemanticCache, ValkeySemanticConfig}; -use pyo3::{PyTraverseError, PyVisit, prelude::*}; +use pyo3::prelude::*; use serde_json::Value; use super::{ @@ -26,13 +27,13 @@ use super::{ }; /// What the Python embedder receives for one semantic request. -pub(super) struct EmbeddingInput { - pub(super) prompt: String, - pub(super) metadata: Option, +pub(in crate::cache) struct EmbeddingInput { + pub(in crate::cache) prompt: String, + pub(in crate::cache) metadata: Option, } /// An exact-match backend behind one pointer, with the identity its facade must reproduce. -pub(super) struct ExactService { +pub(in crate::cache) struct ExactService { cache: Arc, probe: Option>, buffer: Option, @@ -40,7 +41,7 @@ pub(super) struct ExactService { } #[derive(Clone)] -pub(super) enum NativeResponseCache { +pub(in crate::cache) enum NativeResponseCache { Exact(Arc), ValkeySemantic { cache: Arc>>, @@ -240,7 +241,7 @@ impl NativeResponseCache { }) } - pub async fn qdrant_semantic( + pub(super) async fn qdrant_semantic( config: QdrantSemanticCacheConfig, client: litellm_http::Client, runtime: tokio::runtime::Handle, @@ -285,10 +286,6 @@ impl NativeResponseCache { } } - pub fn kind(&self) -> &'static str { - self.identity().kind() - } - pub fn with_redis_flush_size(self, flush_size: Option) -> Self { match self { Self::Exact(service) if matches!(service.identity, BackendIdentity::Redis { .. }) => { @@ -324,7 +321,10 @@ impl NativeResponseCache { } /// The prompt and metadata this backend would embed for `request`, if it has a prompt. - pub(super) fn embedding_input(&self, request: &NativeRequest) -> Option { + pub(in crate::cache) fn embedding_input( + &self, + request: &NativeRequest, + ) -> Option { let context = match self { Self::ValkeySemantic { scope, .. } => request.scoped_semantic(scope).context, Self::RedisSemantic { .. } => request.semantic().context, @@ -462,7 +462,7 @@ impl NativeResponseCache { } } - pub(super) fn async_lookup_semantic_py<'py>( + pub(in crate::cache) fn async_lookup_semantic_py<'py>( &self, py: Python<'py>, request: NativeRequest, @@ -478,7 +478,7 @@ impl NativeResponseCache { .await .map(SemanticReply::from) }, - super::cache_error, + cache_error, ) } Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } => { @@ -487,7 +487,7 @@ impl NativeResponseCache { } } - pub(super) fn async_lookup_py<'py>( + pub(in crate::cache) fn async_lookup_py<'py>( &self, py: Python<'py>, request: NativeRequest, @@ -498,7 +498,7 @@ impl NativeResponseCache { crate::logger::run_async( py, async move { service.async_lookup(&request, now()).await }, - super::cache_error, + cache_error, ) } Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } => { @@ -541,7 +541,7 @@ impl NativeResponseCache { } } - pub(super) fn async_store_py<'py>( + pub(in crate::cache) fn async_store_py<'py>( &self, py: Python<'py>, request: NativeRequest, @@ -553,7 +553,7 @@ impl NativeResponseCache { crate::logger::run_async( py, async move { service.async_store(&request, response, now()).await }, - super::cache_error, + cache_error, ) } Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } => { @@ -611,7 +611,7 @@ impl NativeResponseCache { } } - pub(super) fn async_store_batch_py<'py>( + pub(in crate::cache) fn async_store_batch_py<'py>( &self, py: Python<'py>, entries: Vec<(NativeRequest, Value)>, @@ -622,7 +622,7 @@ impl NativeResponseCache { crate::logger::run_async( py, async move { service.async_store_batch(entries, now()).await }, - super::cache_error, + cache_error, ) } Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } => { @@ -656,15 +656,6 @@ impl NativeResponseCache { } } } - - pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - match self { - Self::ValkeySemantic { embedder, .. } | Self::RedisSemantic { embedder, .. } => { - embedder.traverse(visit) - } - Self::Exact(_) | Self::QdrantSemantic(_) => Ok(()), - } - } } fn exact_requests(requests: &[NativeRequest]) -> Vec { @@ -673,7 +664,10 @@ fn exact_requests(requests: &[NativeRequest]) -> Vec, pub(super) Option); +pub(in crate::cache) struct SemanticReply( + pub(in crate::cache) Option, + pub(in crate::cache) Option, +); impl From> for SemanticReply { fn from(lookup: SemanticLookup) -> Self { diff --git a/litellm-rust/crates/python-bridge/src/cache/config.rs b/litellm-rust/crates/python-bridge/src/cache/native/config.rs similarity index 99% rename from litellm-rust/crates/python-bridge/src/cache/config.rs rename to litellm-rust/crates/python-bridge/src/cache/native/config.rs index e58902b07ee..aa1cb1c70fc 100644 --- a/litellm-rust/crates/python-bridge/src/cache/config.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native/config.rs @@ -11,7 +11,7 @@ use pyo3::{ types::{PyAny, PyBool, PyDict, PyList, PyString}, }; -use super::{identity::BackendIdentity, native::NativeResponseCache, request::duration}; +use super::{backend::NativeResponseCache, identity::BackendIdentity, request::duration}; pub(super) struct CachePolicy { pub(super) redis_flush_size: Option, @@ -233,12 +233,12 @@ pub(super) enum CacheBackendConfig { QdrantSemantic(Box), } -pub(super) struct NativeCacheConfig { +pub(in crate::cache) struct NativeCacheConfig { pub(super) policy: CachePolicy, pub(super) backend: CacheBackendConfig, } -pub(super) enum UnsupportedCacheConfig { +pub(in crate::cache) enum UnsupportedCacheConfig { Backend, RedisTopology, RedisCredentials, @@ -263,7 +263,7 @@ pub(super) enum UnsupportedCacheConfig { } impl UnsupportedCacheConfig { - pub(super) fn message(&self) -> &'static str { + pub(in crate::cache) fn message(&self) -> &'static str { match self { Self::Backend => "native cache backend is not implemented", Self::RedisTopology => "native Redis topology is not implemented", @@ -305,14 +305,14 @@ impl UnsupportedCacheConfig { } } -pub(super) enum CacheConfigProjection { +pub(in crate::cache) enum CacheConfigProjection { Native(Box), Unsupported(UnsupportedCacheConfig), } impl NativeCacheConfig { #[inline(never)] - pub(super) fn project(facade: &Bound<'_, PyAny>) -> PyResult { + pub(in crate::cache) fn project(facade: &Bound<'_, PyAny>) -> PyResult { let backend_name = facade.getattr("type")?.extract::()?; let policy = CachePolicy { redis_flush_size: facade @@ -1170,7 +1170,7 @@ mod tests { GcsCacheConfig, NativeCacheConfig, REDIS_PY_DEFAULT_MAX_CONNECTIONS, RedisConnectionConfig, RedisProtocol, RedisSemanticCacheConfig, RedisTlsConfig, UnsupportedCacheConfig, }; - use crate::cache::{embedder::PythonEmbedder, native::NativeResponseCache}; + use crate::cache::native::{backend::NativeResponseCache, embedder::PythonEmbedder}; fn cluster_facade<'py>(py: Python<'py>, startup_nodes: &str, hook: &str) -> Bound<'py, PyAny> { facade( diff --git a/litellm-rust/crates/python-bridge/src/cache/embedder.rs b/litellm-rust/crates/python-bridge/src/cache/native/embedder.rs similarity index 88% rename from litellm-rust/crates/python-bridge/src/cache/embedder.rs rename to litellm-rust/crates/python-bridge/src/cache/native/embedder.rs index 7eadd9bc4b0..cce36531efd 100644 --- a/litellm-rust/crates/python-bridge/src/cache/embedder.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native/embedder.rs @@ -11,7 +11,7 @@ tokio::task_local! { /// Runs `future` with the vector the Python embedder already produced, so the backend's /// `async_embed` never has to call back into Python from the runtime. -pub(super) fn with_prepared_embedding( +pub(in crate::cache) fn with_prepared_embedding( vector: Result, Error>, future: F, ) -> impl Future { @@ -19,7 +19,7 @@ pub(super) fn with_prepared_embedding( } /// The Python object that owns embedding for a semantic backend. -pub(super) struct PythonEmbedder(Py); +pub(in crate::cache) struct PythonEmbedder(Py); impl Clone for PythonEmbedder { fn clone(&self) -> Self { @@ -28,15 +28,15 @@ impl Clone for PythonEmbedder { } impl PythonEmbedder { - pub(super) fn new(object: Py) -> Self { + pub(in crate::cache) fn new(object: Py) -> Self { Self(object) } - pub(super) fn object(&self) -> &Py { + pub(in crate::cache) fn object(&self) -> &Py { &self.0 } - pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + pub(in crate::cache) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { visit.call(&self.0) } @@ -50,7 +50,7 @@ impl PythonEmbedder { } /// The awaitable of `_get_async_embedding(prompt, metadata=...)`, to run in the caller's loop. - pub(super) fn async_embedding( + pub(in crate::cache) fn async_embedding( &self, py: Python<'_>, prompt: &str, @@ -63,7 +63,7 @@ impl PythonEmbedder { .map(Bound::unbind) } - pub(super) fn extract(vector: Bound<'_, PyAny>) -> PyResult> { + pub(in crate::cache) fn extract(vector: Bound<'_, PyAny>) -> PyResult> { Ok(vector .extract::>()? .into_iter() diff --git a/litellm-rust/crates/python-bridge/src/cache/facade.rs b/litellm-rust/crates/python-bridge/src/cache/native/facade.rs similarity index 94% rename from litellm-rust/crates/python-bridge/src/cache/facade.rs rename to litellm-rust/crates/python-bridge/src/cache/native/facade.rs index d1bddef67ff..0402381777c 100644 --- a/litellm-rust/crates/python-bridge/src/cache/facade.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native/facade.rs @@ -9,10 +9,9 @@ use pyo3::{ use serde_json::Value; use super::{ + backend::NativeResponseCache, config::{CacheConfigProjection, NativeCacheConfig}, - handle::CacheTestHandle, identity::BackendIdentity, - native::NativeResponseCache, }; struct ClassGuard { @@ -84,7 +83,7 @@ const VALKEY_POOL: RedisPoolAttributes = STANDALONE_POOL; /// `Cache._native_cache` holds the runtime `Cache.__init__` resolved. const INSTANCE_STATE: &[&str] = &["_native_cache"]; -pub(super) struct FacadeGuard { +pub(in crate::cache) struct FacadeGuard { outer: ObjectGuard, backend: ObjectGuard, disk_store: Option, @@ -354,7 +353,7 @@ impl ConnectionGuard { } impl FacadeGuard { - pub(super) fn capture( + pub(in crate::cache) fn capture( py: Python<'_>, facade: &Bound<'_, PyAny>, service: &NativeResponseCache, @@ -472,7 +471,11 @@ impl FacadeGuard { }) } - pub(super) fn matches(&self, py: Python<'_>, facade: &Bound<'_, PyAny>) -> PyResult { + pub(in crate::cache) fn matches( + &self, + py: Python<'_>, + facade: &Bound<'_, PyAny>, + ) -> PyResult { if !self.outer.matches(py, facade)? { return Ok(false); } @@ -488,7 +491,7 @@ impl FacadeGuard { self.connection.matches(py, &backend) } - pub(super) fn traverse(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> { + pub(in crate::cache) fn traverse(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> { self.outer.traverse(&visit)?; self.backend.traverse(&visit)?; if let Some(guard) = &self.disk_store { @@ -497,28 +500,3 @@ impl FacadeGuard { self.connection.traverse(&visit) } } - -pub(super) fn resolve( - py: Python<'_>, - facade: &Bound<'_, PyAny>, -) -> PyResult> { - let Ok(dict) = facade - .getattr("__dict__") - .and_then(|dict| dict.cast_into::().map_err(Into::into)) - else { - return Ok(None); - }; - let Some(handle) = dict.get_item("_native_cache_handle")? else { - return Ok(None); - }; - let Ok(handle) = handle.extract::>() else { - return Ok(None); - }; - let Some(guard) = &handle.guard else { - return Ok(None); - }; - if !guard.matches(py, facade).unwrap_or(false) { - return Ok(None); - } - handle.service().map(Some) -} diff --git a/litellm-rust/crates/python-bridge/src/cache/identity.rs b/litellm-rust/crates/python-bridge/src/cache/native/identity.rs similarity index 98% rename from litellm-rust/crates/python-bridge/src/cache/identity.rs rename to litellm-rust/crates/python-bridge/src/cache/native/identity.rs index 835bafd3ff1..3d4868dbe7b 100644 --- a/litellm-rust/crates/python-bridge/src/cache/identity.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native/identity.rs @@ -6,7 +6,7 @@ use litellm_cache_redis::RedisTopology; /// observe on the Python object, captured once so facade projection and native construction /// compare plain data instead of reaching into each backend type. #[derive(Clone, Debug, PartialEq)] -pub(super) enum BackendIdentity { +pub(in crate::cache) enum BackendIdentity { Memory { capacity: usize, max_entry_bytes: Option, @@ -55,8 +55,7 @@ pub(super) enum BackendIdentity { const TYPES: &str = "facade and native backend types must match"; impl BackendIdentity { - /// The native backend name reported to Python through `_CacheTestHandle.backend`. - pub(super) fn kind(&self) -> &'static str { + pub(in crate::cache) fn kind(&self) -> &'static str { match self { Self::Memory { .. } => "memory", Self::Redis { .. } => "redis", @@ -71,7 +70,7 @@ impl BackendIdentity { } /// The `LiteLLMCacheType` value a facade of this backend carries in `Cache.type`. - pub(super) fn cache_type(&self) -> &'static str { + pub(in crate::cache) fn cache_type(&self) -> &'static str { match self { Self::Memory { .. } => "local", Self::Redis { .. } => "redis", @@ -87,7 +86,7 @@ impl BackendIdentity { /// The first difference between the facade's configuration (`self`) and the native /// backend (`native`), in the order Python users see the attributes. - pub(super) fn mismatch(&self, native: &Self) -> Option<&'static str> { + pub(in crate::cache) fn mismatch(&self, native: &Self) -> Option<&'static str> { let mut differences: Vec<(bool, &'static str)> = Vec::new(); let mut differs = |condition: bool, message: &'static str| { differences.push((condition, message)); diff --git a/litellm-rust/crates/python-bridge/src/cache/native/mod.rs b/litellm-rust/crates/python-bridge/src/cache/native/mod.rs new file mode 100644 index 00000000000..ad117dd9eaf --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/cache/native/mod.rs @@ -0,0 +1,11 @@ +pub(super) mod activation; +pub(super) mod backend; +pub(super) mod config; +mod embedder; +pub(super) mod facade; +mod identity; +pub(super) mod request; +mod semantic; +pub(super) mod v2; + +pub(crate) use v2::NativeCacheHandle; diff --git a/litellm-rust/crates/python-bridge/src/cache/request.rs b/litellm-rust/crates/python-bridge/src/cache/native/request.rs similarity index 96% rename from litellm-rust/crates/python-bridge/src/cache/request.rs rename to litellm-rust/crates/python-bridge/src/cache/native/request.rs index 627bf9f1840..c009882a4ad 100644 --- a/litellm-rust/crates/python-bridge/src/cache/request.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native/request.rs @@ -23,7 +23,7 @@ struct RequestInput { } #[derive(Clone)] -pub(super) struct NativeRequest { +pub(in crate::cache) struct NativeRequest { pub(super) key: CacheKeyInput, pub(super) controls: CacheControls, pub(super) ttl: Option, @@ -129,7 +129,7 @@ fn semantic_key(request: &NativeRequest, scope: &str) -> CacheKeyInput { key } -pub(super) fn request(value: &Bound<'_, PyAny>) -> PyResult { +pub(in crate::cache) fn request(value: &Bound<'_, PyAny>) -> PyResult { let input: RequestInput = from_py(value)?; request_input(input) } @@ -152,7 +152,7 @@ fn request_input(input: RequestInput) -> PyResult { }) } -pub(super) fn requests(value: &Bound<'_, PyAny>) -> PyResult> { +pub(in crate::cache) fn requests(value: &Bound<'_, PyAny>) -> PyResult> { from_py::>(value)? .into_iter() .map(request_input) @@ -164,7 +164,7 @@ pub(super) fn duration(seconds: f64) -> PyResult { .map_err(|_| PyValueError::new_err("cache durations must be finite and nonnegative")) } -pub(super) fn now() -> Duration { +pub(in crate::cache) fn now() -> Duration { SystemTime::now() .duration_since(UNIX_EPOCH) .unwrap_or_default() diff --git a/litellm-rust/crates/python-bridge/src/cache/semantic.rs b/litellm-rust/crates/python-bridge/src/cache/native/semantic.rs similarity index 98% rename from litellm-rust/crates/python-bridge/src/cache/semantic.rs rename to litellm-rust/crates/python-bridge/src/cache/native/semantic.rs index 4a1cc0dfe8e..b1c43e1f602 100644 --- a/litellm-rust/crates/python-bridge/src/cache/semantic.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native/semantic.rs @@ -1,3 +1,4 @@ +use crate::cache::cache_error; use crate::logger::run_async; use std::{collections::VecDeque, time::Duration}; @@ -11,9 +12,8 @@ use pyo3::{ use serde_json::Value; use super::{ - cache_error, + backend::{NativeResponseCache, SemanticReply}, embedder::{PythonEmbedder, with_prepared_embedding}, - native::{NativeResponseCache, SemanticReply}, request::{NativeRequest, now}, }; diff --git a/litellm-rust/crates/python-bridge/src/cache/native/v2.rs b/litellm-rust/crates/python-bridge/src/cache/native/v2.rs new file mode 100644 index 00000000000..e355d0d698a --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/cache/native/v2.rs @@ -0,0 +1,345 @@ +use crate::cache::cache_error; +use std::{sync::Arc, time::Duration}; + +use litellm_cache::{DeleteCache, DisconnectCache, PingCache}; +use litellm_host_python::{from_py, release_gil, to_py}; +use serde_json::Value; + +use litellm_cache_memory::InMemoryCache; +use litellm_cache_redis::{RedisCache, RedisTopology}; +use litellm_cache_response::{ + CacheEntry, CacheKeyInput, ExactResponseCache, ResponseCache, ResponseCacheCodec, + ResponseCacheConfig, ResponseCacheRequest, ResponseCacheService, +}; +use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict}; + +#[pyclass( + frozen, + name = "NativeCacheHandle", + module = "litellm.rust_bridge._native" +)] +pub(crate) struct NativeCacheHandle { + service: Arc, + backend: Arc, + storage: Storage, + pid: u32, +} + +#[derive(Clone)] +enum Storage { + Memory(Arc>), + Redis(Arc>), +} + +impl NativeCacheHandle { + fn check_process(&self) -> PyResult<()> { + if self.pid != std::process::id() { + return Err(pyo3::exceptions::PyRuntimeError::new_err( + "recreate the v2 cache after fork", + )); + } + Ok(()) + } +} + +fn request(key: String, ttl: Option) -> PyResult { + let mut request: ResponseCacheRequest = ResponseCacheRequest::new(CacheKeyInput { + preset: Some(key), + ..Default::default() + }); + request.context.ttl = ttl.map(duration).transpose()?; + Ok(request) +} + +#[pymethods] +impl NativeCacheHandle { + #[staticmethod] + #[pyo3(signature = (*, ttl=600.0, capacity=200, max_entry_bytes=4194304))] + fn memory(ttl: f64, capacity: usize, max_entry_bytes: usize) -> PyResult { + let ttl = duration(ttl)?; + if capacity == 0 || max_entry_bytes == 0 { + return Err(PyValueError::new_err("cache limits must be positive")); + } + let storage = Arc::new(InMemoryCache::new(Some(capacity), Some(ttl))); + let backend = Arc::new(ResponseCache::new(storage.clone()).with_config( + ResponseCacheConfig { + namespace: "sdk".into(), + max_entry_bytes, + }, + )); + Ok(Self { + service: backend.clone(), + backend, + storage: Storage::Memory(storage), + pid: std::process::id(), + }) + } + + #[staticmethod] + #[pyo3(signature = (url, *, namespace, ttl=600.0, max_entry_bytes=4194304))] + fn redis( + py: Python<'_>, + url: &str, + namespace: String, + ttl: f64, + max_entry_bytes: usize, + ) -> PyResult { + let ttl = duration(ttl)?; + if namespace.is_empty() || max_entry_bytes == 0 { + return Err(PyValueError::new_err( + "namespace and a positive cache limit are required", + )); + } + let storage = Arc::new( + release_gil(py, || { + RedisCache::connect( + url, + &RedisTopology::Standalone, + Some(ttl), + ResponseCacheCodec, + ) + }) + .map_err(cache_error)? + .with_namespace(Some(namespace.clone())), + ); + let backend = Arc::new(ResponseCache::new(storage.clone()).with_config( + ResponseCacheConfig { + namespace, + max_entry_bytes, + }, + )); + Ok(Self { + service: backend.clone(), + backend, + storage: Storage::Redis(storage), + pid: std::process::id(), + }) + } + fn get(&self, py: Python<'_>, key: String) -> PyResult> { + self.check_process()?; + let request = request(key, None)?; + let value = release_gil(py, || self.backend.lookup(&request, super::request::now())) + .map_err(cache_error)?; + to_py(py, &value) + } + + #[pyo3(signature = (key, value, *, ttl=None))] + fn set( + &self, + py: Python<'_>, + key: String, + value: &Bound<'_, PyAny>, + ttl: Option, + ) -> PyResult<()> { + self.check_process()?; + let request = request(key, ttl)?; + let value: Value = from_py(value)?; + release_gil(py, || { + self.backend.store(&request, value, super::request::now()) + }) + .map_err(cache_error) + } + + fn async_get<'py>(&self, py: Python<'py>, key: String) -> PyResult> { + self.check_process()?; + let request = request(key, None)?; + let backend = self.backend.clone(); + crate::logger::run_async( + py, + async move { backend.async_lookup(&request, super::request::now()).await }, + cache_error, + ) + } + + #[pyo3(signature = (key, value, *, ttl=None))] + fn async_set<'py>( + &self, + py: Python<'py>, + key: String, + value: &Bound<'_, PyAny>, + ttl: Option, + ) -> PyResult> { + self.check_process()?; + let request = request(key, ttl)?; + let value: Value = from_py(value)?; + let backend = self.backend.clone(); + crate::logger::run_async( + py, + async move { + backend + .async_store(&request, value, super::request::now()) + .await + }, + cache_error, + ) + } + + #[pyo3(signature = (entries, *, ttl=None))] + fn async_set_many<'py>( + &self, + py: Python<'py>, + entries: &Bound<'_, PyAny>, + ttl: Option, + ) -> PyResult> { + self.check_process()?; + let entries: Vec<(String, Value)> = from_py(entries)?; + let entries = entries + .into_iter() + .map(|(key, value)| Ok((request(key, ttl)?, value))) + .collect::>>()?; + let backend = self.backend.clone(); + crate::logger::run_async( + py, + async move { + backend + .async_store_batch(entries, super::request::now()) + .await + }, + cache_error, + ) + } + + fn flush(&self, py: Python<'_>) -> PyResult> { + self.check_process()?; + let backend = self.backend.clone(); + crate::logger::run_sync(py, async move { backend.async_flush().await }, cache_error) + } + + fn async_flush<'py>(&self, py: Python<'py>) -> PyResult> { + self.check_process()?; + let backend = self.backend.clone(); + crate::logger::run_async(py, async move { backend.async_flush().await }, cache_error) + } + + fn ping<'py>(&self, py: Python<'py>) -> PyResult> { + self.check_process()?; + let storage = self.storage.clone(); + crate::logger::run_async( + py, + async move { + match storage { + Storage::Memory(_) => Ok(true), + Storage::Redis(cache) => cache.ping().await, + } + }, + cache_error, + ) + } + + fn disconnect<'py>(&self, py: Python<'py>) -> PyResult> { + self.check_process()?; + let storage = self.storage.clone(); + crate::logger::run_async( + py, + async move { + match storage { + Storage::Memory(cache) => cache.disconnect().await, + Storage::Redis(cache) => cache.disconnect().await, + } + }, + cache_error, + ) + } + + fn delete<'py>(&self, py: Python<'py>, keys: Vec) -> PyResult> { + self.check_process()?; + let storage = self.storage.clone(); + crate::logger::run_async( + py, + async move { + for key in keys { + match &storage { + Storage::Memory(cache) => cache.async_delete_cache(&key).await?, + Storage::Redis(cache) => cache.async_delete_cache(&key).await?, + } + } + Ok(()) + }, + cache_error, + ) + } +} + +fn duration(seconds: f64) -> PyResult { + Duration::try_from_secs_f64(seconds) + .ok() + .filter(|value| !value.is_zero()) + .ok_or_else(|| PyValueError::new_err("cache durations must be finite and positive")) +} + +pub(in crate::cache) fn native_handle<'py>( + configured: &Bound<'py, PyAny>, +) -> PyResult>> { + Ok(configured + .getattr_opt("cache")? + .map(|backend| backend.getattr_opt("native_handle")) + .transpose()? + .flatten() + .filter(|handle| handle.is_instance_of::())) +} + +pub(in crate::cache) fn configured( + configured: &Bound<'_, PyAny>, + kwargs: &Bound<'_, PyDict>, +) -> PyResult<( + Option>, + litellm_cache_response::CacheOptions, +)> { + let handle = native_handle(configured)?.ok_or_else(|| { + pyo3::exceptions::PyRuntimeError::new_err( + "the configured cache changed to a Python cache after native admission", + ) + })?; + let cache = handle.extract::>()?; + cache.check_process()?; + let controls = kwargs.get_item("cache")?.filter(|value| !value.is_none()); + let controls = controls + .as_ref() + .map(|value| value.cast::()) + .transpose()?; + if let Some(controls) = controls { + for name in controls.keys() { + let name = name.extract::()?; + if !matches!( + name.as_str(), + "no-cache" | "no-store" | "ttl" | "s-maxage" | "s-max-age" | "use-cache" + ) { + return Err(PyValueError::new_err(format!( + "unsupported v2 cache control: {name}" + ))); + } + } + } + let boolean = |name: &str| -> PyResult { + controls + .map(|values| values.get_item(name)) + .transpose()? + .flatten() + .map(|value| value.extract()) + .transpose() + .map(|value| value.unwrap_or(false)) + }; + let seconds = |name: &str| -> PyResult> { + controls + .map(|values| values.get_item(name)) + .transpose()? + .flatten() + .map(|value| duration(value.extract()?)) + .transpose() + }; + Ok(( + Some(cache.service.clone()), + litellm_cache_response::CacheOptions { + caching: kwargs + .get_item("caching")? + .filter(|value| !value.is_none()) + .map(|value| value.extract()) + .transpose()?, + no_cache: boolean("no-cache")?, + no_store: boolean("no-store")?, + ttl: seconds("ttl")?, + max_age: seconds("s-max-age")?.or(seconds("s-maxage")?), + scope: litellm_cache_response::CacheScope::Shared, + }, + )) +} diff --git a/litellm-rust/crates/python-bridge/src/cache/python/AGENTS.md b/litellm-rust/crates/python-bridge/src/cache/python/AGENTS.md new file mode 100644 index 00000000000..ca4bd68dd38 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/cache/python/AGENTS.md @@ -0,0 +1,9 @@ +# Python cache delegation + +This directory lets Rust inference use a selected Python cache. `service.rs` implements the injected Rust response-cache service and yields typed cache operations. `host.rs` calls the Python cache's sync or async API and delivers the result back to Rust. `callback.rs` provides Python cache delegation for the Python-facing cache runtime + +Keep the shared inference protocol wrapper and adapter selection in the parent module. Receive the configured cache and prepared arguments from the parent module. Do not discover global configuration, choose native backends, or move inference to Python. Core remains independent of Python objects and cache implementation details + +Await asynchronous cache operations through the existing host driver in the caller's task. Do not create another asyncio task or event loop. Cancellation must prevent subsequent provider requests and cache writes. Preserve ordinary cache failure handling without swallowing cancellation or other Python base exceptions + +Traverse retained Python references for GC and release pending replies when the call closes. Regression tests must require native inference with Python fallback disabled and assert observable hits, provider request counts, cache-key headers, task identity and cancellation diff --git a/litellm-rust/crates/python-bridge/src/cache/callback.rs b/litellm-rust/crates/python-bridge/src/cache/python/callback.rs similarity index 84% rename from litellm-rust/crates/python-bridge/src/cache/callback.rs rename to litellm-rust/crates/python-bridge/src/cache/python/callback.rs index 492e0329672..fd8ebe508bc 100644 --- a/litellm-rust/crates/python-bridge/src/cache/callback.rs +++ b/litellm-rust/crates/python-bridge/src/cache/python/callback.rs @@ -5,16 +5,16 @@ use pyo3::{ types::{PyDict, PyList, PyTuple}, }; -use super::future::ready_none; +use crate::cache::future::ready_none; -pub(super) struct PythonCallback(Py); +pub(in crate::cache) struct PythonCallback(Py); impl PythonCallback { - pub(super) fn new(object: Py) -> Self { + pub(in crate::cache) fn new(object: Py) -> Self { Self(object) } - pub(super) fn lookup<'py>( + pub(in crate::cache) fn lookup<'py>( &self, py: Python<'py>, kwargs: Option<&Bound<'py, PyDict>>, @@ -24,7 +24,7 @@ impl PythonCallback { .call_method("get_cache", (), Some(callback_kwargs(kwargs)?)) } - pub(super) fn async_lookup<'py>( + pub(in crate::cache) fn async_lookup<'py>( &self, py: Python<'py>, kwargs: Option<&Bound<'py, PyDict>>, @@ -34,7 +34,7 @@ impl PythonCallback { .call_method("async_get_cache", (), Some(callback_kwargs(kwargs)?)) } - pub(super) fn store( + pub(in crate::cache) fn store( &self, py: Python<'_>, response: &Bound<'_, PyAny>, @@ -46,7 +46,7 @@ impl PythonCallback { .map(|_| ()) } - pub(super) fn async_store<'py>( + pub(in crate::cache) fn async_store<'py>( &self, py: Python<'py>, response: &Bound<'py, PyAny>, @@ -59,7 +59,7 @@ impl PythonCallback { ) } - pub(super) fn lookup_batch<'py>( + pub(in crate::cache) fn lookup_batch<'py>( &self, py: Python<'py>, requests: &Bound<'py, PyAny>, @@ -76,7 +76,7 @@ impl PythonCallback { Ok(results.into_any()) } - pub(super) fn async_lookup_batch<'py>( + pub(in crate::cache) fn async_lookup_batch<'py>( &self, py: Python<'py>, requests: &Bound<'py, PyAny>, @@ -94,7 +94,7 @@ impl PythonCallback { .call_method1("gather", PyTuple::new(py, awaitables)?) } - pub(super) fn async_store_batch<'py>( + pub(in crate::cache) fn async_store_batch<'py>( &self, py: Python<'py>, result: Option<&Bound<'py, PyAny>>, @@ -110,7 +110,10 @@ impl PythonCallback { ) } - pub(super) fn async_flush<'py>(&self, py: Python<'py>) -> PyResult> { + pub(in crate::cache) fn async_flush<'py>( + &self, + py: Python<'py>, + ) -> PyResult> { let object = self.0.bind(py); let backend = match object.getattr_opt("cache")? { Some(backend) if !backend.is_none() => backend, @@ -123,11 +126,11 @@ impl PythonCallback { ready_none(py) } - pub(super) fn ping<'py>(&self, py: Python<'py>) -> PyResult> { + pub(in crate::cache) fn ping<'py>(&self, py: Python<'py>) -> PyResult> { self.0.bind(py).call_method0("ping") } - pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + pub(in crate::cache) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { visit.call(&self.0) } } diff --git a/litellm-rust/crates/python-bridge/src/cache/python/host.rs b/litellm-rust/crates/python-bridge/src/cache/python/host.rs new file mode 100644 index 00000000000..d8f8ac7c181 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/cache/python/host.rs @@ -0,0 +1,141 @@ +use litellm_cache::Error; +use litellm_host::protocol::Reply; +use litellm_host_python::{from_py, to_py}; +use pyo3::{ + gc::{PyTraverseError, PyVisit}, + prelude::*, + types::PyDict, +}; +use serde_json::Value; + +use super::service::CacheCall; + +enum Pending { + Lookup(Reply, Error>>), + Store(Reply>), +} + +pub(crate) struct PythonCache { + cache: Option>, + arguments: Option>, + pending: Option, + asynchronous: bool, +} + +impl PythonCache { + pub fn new(asynchronous: bool) -> Self { + Self { + cache: None, + arguments: None, + pending: None, + asynchronous, + } + } + + pub(in crate::cache) fn bind( + &mut self, + cache: Bound<'_, PyAny>, + arguments: &Bound<'_, PyDict>, + ) { + self.cache = Some(cache.unbind()); + self.arguments = Some(arguments.clone().unbind()); + } + + pub fn begin(&mut self, py: Python<'_>, call: CacheCall) -> PyResult>> { + let Some(cache) = self.cache.as_ref() else { + return Err(pyo3::exceptions::PyRuntimeError::new_err( + "cache operation without configured cache", + )); + }; + let arguments = self + .arguments + .as_ref() + .ok_or_else(|| { + pyo3::exceptions::PyRuntimeError::new_err("cache arguments unavailable") + })? + .bind(py) + .copy()?; + let (method, result) = match call { + CacheCall::Lookup { reply } => { + self.pending = Some(Pending::Lookup(reply)); + ( + if self.asynchronous { + "async_get_cache" + } else { + "get_cache" + }, + None, + ) + } + CacheCall::Store { value, reply } => { + self.pending = Some(Pending::Store(reply)); + ( + if self.asynchronous { + "async_add_cache" + } else { + "add_cache" + }, + Some(to_py(py, &value)?), + ) + } + }; + let result = match result { + Some(value) => cache + .bind(py) + .call_method(method, (value,), Some(&arguments)), + None => cache.bind(py).call_method(method, (), Some(&arguments)), + } + .map(Bound::unbind); + if self.asynchronous && result.is_ok() { + return result.map(Some); + } + self.resume(py, result) + } + + pub fn resume( + &mut self, + py: Python<'_>, + result: PyResult>, + ) -> PyResult>> { + if let Err(error) = &result + && !error.is_instance_of::(py) + { + self.pending = None; + return Err(result.err().unwrap()); + } + match self.pending.take() { + Some(Pending::Lookup(reply)) => { + let value = result.map_err(|_| Error::Unavailable).and_then(|value| { + if value.bind(py).is_none() { + Ok(None) + } else { + from_py(value.bind(py)) + .map(Some) + .map_err(|_| Error::InvalidEntry) + } + }); + reply.send(value); + } + Some(Pending::Store(reply)) => { + reply.send(result.map(|_| ()).map_err(|_| Error::Unavailable)); + } + None => { + return Err(pyo3::exceptions::PyRuntimeError::new_err( + "cache reply without pending operation", + )); + } + } + Ok(None) + } + + pub fn close(&mut self) { + self.pending = None; + self.cache = None; + self.arguments = None; + } + + pub fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + visit.call(&self.cache)?; + visit.call(&self.arguments) + } +} diff --git a/litellm-rust/crates/python-bridge/src/cache/python/mod.rs b/litellm-rust/crates/python-bridge/src/cache/python/mod.rs new file mode 100644 index 00000000000..25550c2ac73 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/cache/python/mod.rs @@ -0,0 +1,8 @@ +mod callback; +mod host; +mod service; + +pub(super) use callback::PythonCallback; +pub(crate) use host::PythonCache; +pub(crate) use service::CacheCall; +pub(super) use service::service; diff --git a/litellm-rust/crates/python-bridge/src/cache/python/service.rs b/litellm-rust/crates/python-bridge/src/cache/python/service.rs new file mode 100644 index 00000000000..14310d952b2 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/cache/python/service.rs @@ -0,0 +1,107 @@ +use std::{future::Future, pin::Pin, time::Duration}; + +use litellm_cache::Error; +use litellm_cache_response::{ + ResponseCacheConfig, ResponseCacheRequest, ResponseCacheService, ResponseEnvelope, +}; +use litellm_core::caching::CachedOutput; +use litellm_host::{ + machine::{HostServices, MachineFault}, + protocol::{Protocol, Reply}, +}; +use serde_json::Value; + +const STREAM_EVENTS_KEY: &str = "litellm_cached_anthropic_sse_events"; + +fn from_python(value: Value) -> Result { + let output = match value.get(STREAM_EVENTS_KEY) { + Some(events) => { + let events: Vec = + serde_json::from_value(events.clone()).map_err(|_| Error::InvalidEntry)?; + CachedOutput::Stream(events.concat()) + } + None => CachedOutput::Response(value), + }; + serde_json::to_value(ResponseEnvelope::new("messages", output)).map_err(|_| Error::InvalidEntry) +} + +fn to_python(value: Value) -> Result { + let envelope: ResponseEnvelope> = + serde_json::from_value(value).map_err(|_| Error::InvalidEntry)?; + match envelope.decode("messages").ok_or(Error::InvalidEntry)? { + CachedOutput::Response(response) => Ok(response), + CachedOutput::Stream(text) => Ok(serde_json::json!({ + STREAM_EVENTS_KEY: text.split_inclusive("\n\n").collect::>() + })), + } +} + +pub(crate) enum CacheCall { + Lookup { + reply: Reply, Error>>, + }, + Store { + value: Value, + reply: Reply>, + }, +} + +struct PythonCacheService { + services: HostServices

, + config: ResponseCacheConfig, +} + +pub(in crate::cache) fn service>( + services: HostServices

, + namespace: String, +) -> std::sync::Arc +where + P::Error: From, +{ + std::sync::Arc::new(PythonCacheService { + services, + config: ResponseCacheConfig { + namespace, + ..Default::default() + }, + }) +} + +impl> ResponseCacheService for PythonCacheService

+where + P::Error: From, +{ + fn config(&self) -> &ResponseCacheConfig { + &self.config + } + + fn lookup<'a>( + &'a self, + _: &'a ResponseCacheRequest, + _: Duration, + ) -> Pin, Error>> + Send + 'a>> { + Box::pin(async move { + self.services + .call(|reply| CacheCall::Lookup { reply }) + .await + .map_err(|_| Error::Unavailable)?? + .map(from_python) + .transpose() + }) + } + + fn store<'a>( + &'a self, + _: &'a ResponseCacheRequest, + value: Value, + _: Duration, + ) -> Pin> + Send + 'a>> { + Box::pin(async move { + let value = to_python(value)?; + self.services + .call(|reply| CacheCall::Store { value, reply }) + .await + .map_err(|_| Error::Unavailable)? + }) + } +} diff --git a/litellm-rust/crates/python-bridge/src/cache/resolver.rs b/litellm-rust/crates/python-bridge/src/cache/resolver.rs deleted file mode 100644 index 3baaada4b17..00000000000 --- a/litellm-rust/crates/python-bridge/src/cache/resolver.rs +++ /dev/null @@ -1,25 +0,0 @@ -use pyo3::{PyTraverseError, PyVisit, prelude::*}; - -use super::binding::ResolvedCache; - -#[pyclass(frozen, name = "_CacheResolver")] -pub(crate) struct CacheResolver { - namespace: Py, -} - -#[pymethods] -impl CacheResolver { - #[new] - fn new(namespace: Py) -> Self { - Self { namespace } - } - - pub(crate) fn resolve(&self, py: Python<'_>) -> PyResult { - let object = self.namespace.bind(py).getattr("cache")?; - ResolvedCache::from_selected(&object) - } - - fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> { - visit.call(&self.namespace) - } -} diff --git a/litellm-rust/crates/python-bridge/src/cache/binding.rs b/litellm-rust/crates/python-bridge/src/cache/runtime.rs similarity index 95% rename from litellm-rust/crates/python-bridge/src/cache/binding.rs rename to litellm-rust/crates/python-bridge/src/cache/runtime.rs index 6b5de7f029b..3bcefad1f1c 100644 --- a/litellm-rust/crates/python-bridge/src/cache/binding.rs +++ b/litellm-rust/crates/python-bridge/src/cache/runtime.rs @@ -10,13 +10,13 @@ use pyo3::{ use serde_json::Value; use super::{ - activation::activate, cache_error, - callback::PythonCallback, - config::{CacheConfigProjection, NativeCacheConfig}, future::{ready_none, ready_value}, - native::{NativeResponseCache, SemanticReply}, - request::{now, request, requests}, + native::activation::activate, + native::backend::{NativeResponseCache, SemanticReply}, + native::config::{CacheConfigProjection, NativeCacheConfig}, + native::request::{now, request, requests}, + python::PythonCallback, }; use crate::errors::RustBridgeDeclined; @@ -29,7 +29,7 @@ pub(super) enum CacheBinding { #[pyclass(frozen, name = "_ResponseCacheRuntime")] pub(crate) struct ResolvedCache { binding: CacheBinding, - guard: Option, + guard: Option, pid: u32, } @@ -42,7 +42,7 @@ impl ResolvedCache { } } - pub(super) fn with_guard(mut self, guard: super::facade::FacadeGuard) -> Self { + pub(super) fn with_guard(mut self, guard: super::native::facade::FacadeGuard) -> Self { self.guard = Some(guard); self } @@ -90,10 +90,6 @@ impl ResolvedCache { let py = cache.py(); let binding = if cache.is_none() { CacheBinding::Disabled - } else if let Ok(handle) = cache.extract::>() { - CacheBinding::Native(handle.service()?) - } else if let Some(service) = super::facade::resolve(py, cache)? { - CacheBinding::Native(service) } else if let Some(runtime) = cache .getattr_opt("_native_cache")? .filter(|value| !value.is_none()) @@ -134,7 +130,7 @@ impl ResolvedCache { let service = activate(cache.py(), &backend, config)?; let resolved = Self::new(CacheBinding::Native(service.clone())); Ok( - match super::facade::FacadeGuard::capture(cache.py(), cache, &service) { + match super::native::facade::FacadeGuard::capture(cache.py(), cache, &service) { Ok(guard) => resolved.with_guard(guard), Err(_) => resolved, }, diff --git a/litellm-rust/crates/python-bridge/src/cache/selection.rs b/litellm-rust/crates/python-bridge/src/cache/selection.rs new file mode 100644 index 00000000000..723ff703f64 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/cache/selection.rs @@ -0,0 +1,176 @@ +use super::{native, python}; +use litellm_cache_response::{CacheOptions, CacheScope, ResponseCacheService, ScopedCache}; +use litellm_host::{ + machine::{HostServices, MachineFault}, + protocol::Protocol, +}; +use pyo3::{prelude::*, types::PyDict}; +use std::sync::Arc; + +pub(crate) struct Cached

(std::marker::PhantomData

); + +impl Protocol for Cached

{ + type Request = (P::Request, Selection); + type Response = P::Response; + type Error = P::Error; + type HostCall = python::CacheCall; + type Chunk = P::Chunk; + type StreamHead = P::StreamHead; +} + +enum Backend { + Disabled, + Native(Arc), + Python { namespace: String }, +} + +pub(crate) struct Selection { + backend: Backend, + options: CacheOptions, +} + +impl Selection { + pub(crate) fn attach>( + self, + services: HostServices

, + ) -> (Option, CacheOptions) + where + P::Error: From, + { + let service = match self.backend { + Backend::Disabled => None, + Backend::Native(service) => Some(service), + Backend::Python { namespace } => Some(python::service(services, namespace)), + }; + ( + service.map(|service| ScopedCache::new(service, CacheScope::Shared)), + self.options, + ) + } +} + +fn selected_cache<'py>( + py: Python<'py>, + kwargs: &Bound<'py, PyDict>, + call_type: &str, +) -> PyResult>> { + let configured = py.import("litellm")?.getattr("cache")?; + if configured.is_none() + || kwargs + .get_item("caching")? + .is_some_and(|value| value.is(pyo3::types::PyBool::new(py, false))) + { + return Ok(None); + } + let supported = configured.getattr("supported_call_types")?; + if supported.is_none() || !supported.contains(call_type)? { + return Ok(None); + } + Ok(Some(configured)) +} + +pub(crate) fn admit_native( + py: Python<'_>, + kwargs: &Bound<'_, PyDict>, + call_type: &str, +) -> PyResult<()> { + if let Some(configured) = selected_cache(py, kwargs, call_type)? + && native::v2::native_handle(&configured)?.is_none() + { + return Err(crate::errors::RustBridgeDeclined::new_err( + "the configured cache requires Python inference", + )); + } + Ok(()) +} + +pub(crate) fn configured_native( + py: Python<'_>, + kwargs: &Bound<'_, PyDict>, + call_type: &str, +) -> PyResult<( + Option>, + litellm_cache_response::CacheOptions, +)> { + let Some(configured) = selected_cache(py, kwargs, call_type)? else { + return Ok(( + None, + litellm_cache_response::CacheOptions::new(litellm_cache_response::CacheScope::Shared), + )); + }; + native_configuration(&configured, kwargs) +} + +fn native_configuration( + configured: &Bound<'_, PyAny>, + kwargs: &Bound<'_, PyDict>, +) -> PyResult<(Option>, CacheOptions)> { + if !configured + .call_method("should_use_cache", (), Some(kwargs))? + .extract::()? + { + return Ok(( + None, + litellm_cache_response::CacheOptions::new(litellm_cache_response::CacheScope::Shared), + )); + } + native::v2::configured(configured, kwargs) +} + +pub(crate) fn configure( + python: &mut python::PythonCache, + py: Python<'_>, + arguments: &Bound<'_, PyDict>, + call_type: &str, +) -> PyResult { + let selected = selected_cache(py, arguments, call_type)?; + let Some(cache) = selected else { + return Ok(Selection { + backend: Backend::Disabled, + options: CacheOptions::new(CacheScope::Shared), + }); + }; + if native::v2::native_handle(&cache)?.is_some() { + let (native, options) = native_configuration(&cache, arguments)?; + return Ok(Selection { + backend: native.map_or(Backend::Disabled, Backend::Native), + options, + }); + } + let enabled = cache + .call_method("should_use_cache", (), Some(arguments))? + .extract::()?; + let controls = arguments + .get_item("cache")? + .filter(|value| !value.is_none()); + let boolean = |name: &str| -> PyResult { + controls + .as_ref() + .map(|value| value.cast::()?.get_item(name)) + .transpose()? + .flatten() + .map(|value| value.extract()) + .transpose() + .map(|value| value.unwrap_or(false)) + }; + let options = CacheOptions { + no_cache: boolean("no-cache")?, + no_store: boolean("no-store")?, + ..CacheOptions::new(CacheScope::Shared) + }; + let namespace = cache + .getattr_opt("namespace")? + .filter(|value| !value.is_none()) + .map(|value| value.extract()) + .transpose()? + .unwrap_or_default(); + python.bind(cache, arguments); + Ok(Selection { + backend: if enabled { + Backend::Python { namespace } + } else { + Backend::Disabled + }, + options, + }) +} diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index b90b313e799..c8b0fd2f8bc 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -16,7 +16,7 @@ mod tokenizer; #[pymodule(gil_used = true)] mod _native { - use crate::cache::{CacheResolver, CacheTestHandle, ResolvedCache}; + use crate::cache::ResolvedCache; #[cfg(feature = "panic-test")] #[pymodule_export] use crate::diagnostics::_panic_for_test; @@ -55,9 +55,10 @@ mod _native { fn init(module: &Bound<'_, PyModule>) -> PyResult<()> { let py = module.py(); let dict = module.dict(); - dict.set_item("_CacheTestHandle", py.get_type::())?; - dict.set_item("_CacheResolver", py.get_type::())?; - dict.set_item("_CacheTestResolver", py.get_type::())?; + dict.set_item( + "NativeCacheHandle", + py.get_type::(), + )?; dict.set_item("_ResponseCacheRuntime", py.get_type::())?; dict.set_item( "_SecretManagerRuntime", @@ -77,11 +78,12 @@ pub(crate) fn native_module(py: Python<'_>) -> Bound<'_, PyModule> { mod tests { use super::*; - #[test] + #[rstest::rstest] fn module_registration_preserves_the_public_surface() { Python::initialize(); Python::attach(|py| { let mut expected = vec![ + "NativeCacheHandle", "RustBridgeDeclined", "RustUpstreamError", "ForkedAfterNativeRuntimeStarted", diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs index 850d6526d7f..37d3420a285 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -142,6 +142,12 @@ fn run_public( request.clone().unbind(), "litellm.rust_bridge.chat_completions.route_host", ); + let cache_call_type = if asynchronous { + "acompletion" + } else { + "completion" + }; + crate::cache::admit_native(py, &kwargs, cache_call_type)?; let (arguments, hooks) = crate::routes::call_hooks( py, Operation::Completion, @@ -160,7 +166,16 @@ fn run_public( crate::http::resources().auth.clone(), crate::secrets::source(py)?, ); - Ok(route.machine(request, None)) + let (cache, cache_options) = + crate::cache::configured_native(py, arguments, cache_call_type)?; + let route = match cache { + Some(cache) => route.with_cache(litellm_cache_response::ScopedCache::new( + cache, + litellm_cache_response::CacheScope::Shared, + )), + None => route, + }; + Ok(route.machine(request, cache_options)) }, host::ChatCompletionsPythonHost(host), hooks, diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index d05dbb75d73..d03be3ebb49 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -1,5 +1,5 @@ +use crate::cache::{CacheCall, Cached, PythonCache, Selection}; use litellm_host_python::{PythonHostCalls, PythonOwned}; -use std::convert::Infallible; use bytes::Bytes; use litellm_core::messages::{ @@ -89,11 +89,15 @@ fn native_error(py: Python<'_>, error: Error) -> PyResult { /// public response, chunks and exceptions. pub(super) struct MessagesPythonHost { request: Py, + cache: PythonCache, } impl MessagesPythonHost { - pub(super) fn new(request: Py) -> Self { - Self { request } + pub(super) fn new(request: Py, asynchronous: bool) -> Self { + Self { + request, + cache: PythonCache::new(asynchronous), + } } fn projection( @@ -223,17 +227,21 @@ impl MessagesPythonHost { } impl PythonBinding for MessagesPythonHost { - type Protocol = Messages; + type Protocol = Cached; type Failure = PyErr; fn decode_request( &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, - ) -> Result> { + ) -> Result<(MessagesCall, Selection), InvokeError> { + let selection = + crate::cache::configure(&mut self.cache, py, arguments, "anthropic_messages") + .map_err(InvokeError::Python)?; self.projection(py, arguments) .map_err(|error| InvokeError::Python(self.map_failure(py, error)))? .map_err(InvokeError::Native) + .map(|request| (request, selection)) } fn encode_response( @@ -278,20 +286,42 @@ impl PythonBinding for MessagesPythonHost { } } -impl PythonHostCalls for MessagesPythonHost { +impl PythonHostCalls> for MessagesPythonHost { fn handle_host_call( &mut self, - _: Python<'_>, - op: Infallible, + py: Python<'_>, + op: CacheCall, ) -> Result<(), InvokeError> { - match op {} + self.cache + .begin(py, op) + .map(|_| ()) + .map_err(InvokeError::Python) + } + + fn begin_host_call( + &mut self, + py: Python<'_>, + op: CacheCall, + ) -> Result>, InvokeError> { + self.cache.begin(py, op).map_err(InvokeError::Python) + } + + fn resume_host_call( + &mut self, + py: Python<'_>, + result: PyResult>, + ) -> Result>, InvokeError> { + self.cache.resume(py, result).map_err(InvokeError::Python) } } impl PythonOwned for MessagesPythonHost { - fn close(&mut self, _: Python<'_>) {} + fn close(&mut self, _: Python<'_>) { + self.cache.close(); + } fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - visit.call(&self.request) + visit.call(&self.request)?; + self.cache.traverse(visit) } } diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index d13be311b10..1bf2b1e0ba1 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -26,15 +26,40 @@ fn run_messages( py, arguments, move |py, arguments, request| { - let route = litellm_core::messages::MessagesRoute::new( - crate::http::provider_client(py, arguments, asynchronous)? - .map_err(crate::http::client_error)?, - crate::http::resources().auth.clone(), - crate::secrets::source(py)?, - ); - Ok(route.machine(request, None)) + let builder = litellm_core::messages::MessagesRoute::builder() + .with_http( + crate::http::provider_client(py, arguments, asynchronous)? + .map_err(crate::http::client_error)?, + ) + .with_auth(crate::http::resources().auth.clone()) + .with_secrets(crate::secrets::source(py)?); + let route = builder.build(); + Ok(litellm_host::call::hosted_call( + request, + None, + move |(call, selection): (_, crate::cache::Selection), + services, + interceptors, + observers| async move { + let (cache, options) = selection.attach(services); + let route = match cache { + Some(cache) => route.with_cache(cache), + None => route, + }; + route + .execute( + call, + &interceptors, + litellm_core::CallOptions { + cache: Some(options), + observers, + }, + ) + .await + }, + )) }, - MessagesPythonHost::new(request.unbind()), + MessagesPythonHost::new(request.unbind(), asynchronous), hooks, asynchronous, ) diff --git a/litellm-rust/crates/python-bridge/src/routes/responses.rs b/litellm-rust/crates/python-bridge/src/routes/responses.rs index e65f74ec1b4..9b21ef13324 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses.rs @@ -63,6 +63,12 @@ fn run_public( "native Python responses streaming", )); } + let cache_call_type = if asynchronous { + "aresponses" + } else { + "responses" + }; + crate::cache::admit_native(py, &kwargs, cache_call_type)?; let (arguments, hooks) = crate::routes::call_hooks( py, Operation::Responses, @@ -81,7 +87,16 @@ fn run_public( crate::http::resources().auth.clone(), crate::secrets::source(py)?, ); - Ok(route.machine(request, None)) + let (cache, cache_options) = + crate::cache::configured_native(py, arguments, cache_call_type)?; + let route = match cache { + Some(cache) => route.with_cache(litellm_cache_response::ScopedCache::new( + cache, + litellm_cache_response::CacheScope::Shared, + )), + None => route, + }; + Ok(route.machine(request, cache_options)) }, host::ResponsesPythonHost(host), hooks, diff --git a/litellm/_v2/AGENTS.md b/litellm/_v2/AGENTS.md new file mode 100644 index 00000000000..55998814a7f --- /dev/null +++ b/litellm/_v2/AGENTS.md @@ -0,0 +1,3 @@ +Everything here is experimental and should not be documented + +Use this directory to explore alternative APIs where the Rust migration makes backward compatibility difficult. The gateway can use them for performance, but keep them behind the v2 flag to avoid breaking SDK users diff --git a/litellm/_v2/__init__.py b/litellm/_v2/__init__.py new file mode 100644 index 00000000000..79156cf0f80 --- /dev/null +++ b/litellm/_v2/__init__.py @@ -0,0 +1,3 @@ +from litellm._v2.cache import Cache + +__all__ = ("Cache",) diff --git a/litellm/_v2/cache/AGENTS.md b/litellm/_v2/cache/AGENTS.md new file mode 100644 index 00000000000..308419b7253 --- /dev/null +++ b/litellm/_v2/cache/AGENTS.md @@ -0,0 +1,11 @@ +# Python v2 cache + +Keep `litellm._v2.cache.Cache` import-compatible when reorganizing this package. This package owns Python cache factories and the adapter between the existing `BaseCache` interface and `NativeCacheHandle` + +Construct native handles at this boundary and inject the adapter through the existing cache facade's `_backend` parameter. Keep the facade's `type`, namespace, and TTL consistent with the configured native backend + +Keep storage implementation in the Rust storage crates and response-cache policy in `litellm-cache-response` and core. Do not duplicate cache-key generation, freshness rules, response encoding, or inference orchestration here + +Preserve synchronous and asynchronous cache operations, including TTL forwarding and lifecycle methods. Validate Python values before passing them to typed native interfaces. Keep native extension imports lazy so importing the package does not require loading the extension + +Extend the existing v2 cache tests in `tests/test_litellm_rust/test_v2.py` for behavioral changes, following that directory's `AGENTS.md`. Test observable cache behavior rather than package layout or implementation structure diff --git a/litellm/_v2/cache/__init__.py b/litellm/_v2/cache/__init__.py new file mode 100644 index 00000000000..5a3860b7a6e --- /dev/null +++ b/litellm/_v2/cache/__init__.py @@ -0,0 +1,72 @@ +from __future__ import annotations + +from collections.abc import Sequence +from typing import TYPE_CHECKING, Final + +from pydantic import TypeAdapter + +from litellm.caching.base_cache import BaseCache +from litellm.caching.caching import Cache as CacheFacade +from litellm.types.caching import LiteLLMCacheType + +if TYPE_CHECKING: + from litellm.rust_bridge._native import NativeCacheHandle + +_DURATION: Final[TypeAdapter[float | None]] = TypeAdapter(float | None) + + +class NativeBackend(BaseCache): + def __init__(self, handle: NativeCacheHandle) -> None: + self.native_handle = handle + + def get_cache(self, key: str, **kwargs: object) -> object: + return self.native_handle.get(key) + + async def async_get_cache(self, key: str, **kwargs: object) -> object: + return await self.native_handle.async_get(key) + + def set_cache(self, key: str, value: object, **kwargs: object) -> None: + self.native_handle.set(key, value, ttl=_DURATION.validate_python(kwargs.get("ttl"))) + + async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None: + await self.native_handle.async_set(key, value, ttl=_DURATION.validate_python(kwargs.get("ttl"))) + + async def async_set_cache_pipeline(self, cache_list: Sequence[tuple[str, object]], **kwargs: object) -> None: + await self.native_handle.async_set_many(cache_list, ttl=_DURATION.validate_python(kwargs.get("ttl"))) + + async def batch_cache_write(self, key: str, value: object, **kwargs: object) -> None: + await self.async_set_cache(key, value, **kwargs) + + def flush_cache(self) -> None: + self.native_handle.flush() + + async def async_flush_cache(self) -> None: + await self.native_handle.async_flush() + + async def ping(self) -> bool: + return await self.native_handle.ping() + + async def disconnect(self) -> None: + await self.native_handle.disconnect() + + async def delete_cache_keys(self, keys: Sequence[str]) -> None: + await self.native_handle.delete(keys) + + async def test_connection(self) -> dict[str, str]: + return {"status": "success" if await self.ping() else "failed"} + + +class Cache: + @staticmethod + def memory(*, ttl: float = 600, capacity: int = 200, max_entry_bytes: int = 4194304) -> CacheFacade: + from litellm.rust_bridge._native import NativeCacheHandle + + handle: Final = NativeCacheHandle.memory(ttl=ttl, capacity=capacity, max_entry_bytes=max_entry_bytes) + return CacheFacade(type=LiteLLMCacheType.LOCAL, ttl=ttl, _backend=NativeBackend(handle)) + + @staticmethod + def redis(url: str, *, namespace: str, ttl: float = 600, max_entry_bytes: int = 4194304) -> CacheFacade: + from litellm.rust_bridge._native import NativeCacheHandle + + handle: Final = NativeCacheHandle.redis(url, namespace=namespace, ttl=ttl, max_entry_bytes=max_entry_bytes) + return CacheFacade(type=LiteLLMCacheType.REDIS, namespace=namespace, ttl=ttl, _backend=NativeBackend(handle)) diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index d766d1a58bc..e157730779b 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -127,6 +127,7 @@ class Cache: # GCP IAM authentication parameters gcp_service_account: str | None = None, gcp_ssl_ca_certs: str | None = None, + _backend: BaseCache | None = None, **kwargs, ): """ @@ -183,7 +184,9 @@ class Cache: Returns: None. Cache is set as a litellm param """ - if type == LiteLLMCacheType.REDIS: + if _backend is not None: + self.cache: BaseCache = _backend + elif type == LiteLLMCacheType.REDIS: # Check REDIS_CLUSTER_NODES env var if no explicit startup nodes if not redis_startup_nodes: _env_cluster_nodes: Final = litellm.get_secret("REDIS_CLUSTER_NODES") @@ -205,7 +208,7 @@ class Cache: if gcp_ssl_ca_certs is not None: cluster_kwargs["gcp_ssl_ca_certs"] = gcp_ssl_ca_certs - self.cache: BaseCache = RedisClusterCache(**cluster_kwargs) + self.cache = RedisClusterCache(**cluster_kwargs) else: self.cache = RedisCache( host=host, @@ -314,12 +317,6 @@ class Cache: if self.namespace is not None and isinstance(self.cache, RedisCache): self.cache.namespace = self.namespace - from litellm.rust_bridge.response_cache import resolve_response_cache - - # The Rust catalog picks the store per backend. When it selects Rust, the storage calls - # below go to the native runtime and the Python backend stays only for its direct API. - self._native_cache = resolve_response_cache(self) - # Params whose values carry prompt content. Excluded from semantic-cache # scope keys so differently worded prompts share a bucket and match via # vector similarity rather than being split into per-wording buckets. diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index 4fe95d040a5..6a579889869 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -202,85 +202,6 @@ class _ResponseCacheRuntime: 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 _CacheResolver: - def __new__(cls, namespace: object) -> _CacheResolver: ... - def resolve(self) -> _ResponseCacheRuntime: ... - -@final -class _CacheTestResolver: - def __new__(cls, namespace: object) -> _CacheTestResolver: ... - def resolve(self) -> _ResponseCacheRuntime: ... - @final class TokenCounter: @staticmethod @@ -458,3 +379,25 @@ class _SecretManagerRuntime: self, secret_name: str, optional_params: Mapping[str, object] | None = None, timeout: float | httpx.Timeout | None = None, primary_secret_name: str | None = None, ) -> Future[JsonValue]: ... + +@final +class NativeCacheHandle: + def __new__(cls, _uninstantiable: Never, /) -> Never: ... + @staticmethod + def memory( + *, ttl: float = 600.0, capacity: int = 200, max_entry_bytes: int = 4194304, + ) -> NativeCacheHandle: ... + @staticmethod + def redis( + url: str, *, namespace: str, ttl: float = 600.0, max_entry_bytes: int = 4194304, + ) -> NativeCacheHandle: ... + def get(self, key: str) -> object: ... + def set(self, key: str, value: object, *, ttl: float | None = None) -> None: ... + def async_get(self, key: str) -> Future[object]: ... + def async_set(self, key: str, value: object, *, ttl: float | None = None) -> Future[None]: ... + def async_set_many(self, entries: Sequence[tuple[str, object]], *, ttl: float | None = None) -> Future[None]: ... + def flush(self) -> None: ... + def async_flush(self) -> Future[None]: ... + def ping(self) -> Future[bool]: ... + def disconnect(self) -> Future[None]: ... + def delete(self, keys: Sequence[str]) -> Future[None]: ... diff --git a/litellm/rust_bridge/callbacks_legacy_python.py b/litellm/rust_bridge/callbacks_legacy_python.py index 8ce11491277..4f1f2b3c9fd 100644 --- a/litellm/rust_bridge/callbacks_legacy_python.py +++ b/litellm/rust_bridge/callbacks_legacy_python.py @@ -85,9 +85,18 @@ def finalize( MetadataUpdater, response_metadata.update_response_metadata ) update(response, logger, model if isinstance(model, str) else None, kwargs, start_time, end_time) + cache_key: Final = logger.model_call_details.get("cache_key") + if logger.model_call_details.get("cache_hit") is True and isinstance(cache_key, str): + from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict + + hidden: Final = get_hidden_params_dict(response, create=True) + hidden.update({"cache_key": cache_key, "cache_hit": True}) class LoggingSurface(Protocol): + @property + def model_call_details(self) -> Mapping[str, object]: ... + @property def litellm_params(self) -> Mapping[str, object]: ... @@ -225,7 +234,12 @@ def defer_success(logger: LoggingSurface, pending: object) -> None: def sync_success_for_async_call( logger: LoggingSurface, response: object, start: datetime.datetime, end: datetime.datetime ) -> None: - logger.handle_sync_success_callbacks_for_async_calls(result=response, start_time=start, end_time=end) + logger.handle_sync_success_callbacks_for_async_calls( + result=response, + start_time=start, + end_time=end, + cache_hit=True if logger.model_call_details.get("cache_hit") is True else None, + ) def failure_handler( @@ -245,13 +259,22 @@ def failure_handler( def submit_success(logger: LoggingSurface, response: object, start: datetime.datetime, end: datetime.datetime) -> None: from litellm.litellm_core_utils.litellm_logging import executor - executor.submit(contextvars.copy_context().run, logger.success_handler, response, start, end) + executor.submit( + contextvars.copy_context().run, + logger.success_handler, + response, + start, + end, + cache_hit=True if logger.model_call_details.get("cache_hit") is True else None, + ) def async_success_handler( logger: LoggingSurface, response: object, start: datetime.datetime, end: datetime.datetime ) -> Coroutine[object, object, None]: - return logger.async_success_handler(response, start, end) + return logger.async_success_handler( + response, start, end, cache_hit=True if logger.model_call_details.get("cache_hit") is True else None + ) def enqueue_logging(coroutine: Coroutine[object, object, None]) -> None: diff --git a/litellm/rust_bridge/catalog.py b/litellm/rust_bridge/catalog.py index 32926a8464b..40c6456431d 100644 --- a/litellm/rust_bridge/catalog.py +++ b/litellm/rust_bridge/catalog.py @@ -1,4 +1,4 @@ -"""Ordered rollout policy for routes, cache backends, and secret managers. +"""Ordered rollout policy for routes, loggers, and secret managers. The first matching rule wins; unmatched contexts stay on Python. Native admission separately decides whether the selected implementation can execute. @@ -12,7 +12,6 @@ from typing import Final, TypeAlias from litellm.rust_bridge.configuration import Decision, Rollout from litellm.rust_bridge.configuration import decision as _decision -from litellm.types.caching import LiteLLMCacheType from litellm.types.secret_managers.main import KeyManagementSystem @@ -50,20 +49,6 @@ class RouteRule: ) -@dataclass(frozen=True, slots=True) -class CacheContext: - backend: str - - -@dataclass(frozen=True, slots=True) -class CacheRule: - rollout: Rollout - backends: frozenset[str] | None = None - - def matches(self, context: Context) -> bool: - return isinstance(context, CacheContext) and (self.backends is None or context.backend in self.backends) - - @dataclass(frozen=True, slots=True) class SecretManagerContext: system: str @@ -91,8 +76,8 @@ class LoggerRule: return isinstance(context, LoggerContext) -Context: TypeAlias = RouteContext | CacheContext | SecretManagerContext | LoggerContext -Rule: TypeAlias = RouteRule | CacheRule | SecretManagerRule | LoggerRule +Context: TypeAlias = RouteContext | SecretManagerContext | LoggerContext +Rule: TypeAlias = RouteRule | SecretManagerRule | LoggerRule Rules: TypeAlias = tuple[Rule, ...] RULES: Final[Rules] = ( @@ -106,15 +91,6 @@ RULES: Final[Rules] = ( RouteRule(Route.TOKEN_COUNTER, Rollout.PYTHON_ONLY), RouteRule(Route.TOKENIZER, Rollout.PYTHON_ONLY), RouteRule(Route.TRANSCRIPTION, Rollout.RUST_REQUIRED, providers=frozenset({"bedrock"})), - CacheRule(Rollout.PYTHON_ONLY, backends=frozenset({LiteLLMCacheType.LOCAL})), - CacheRule(Rollout.PYTHON_ONLY, backends=frozenset({LiteLLMCacheType.REDIS})), - CacheRule(Rollout.PYTHON_ONLY, backends=frozenset({LiteLLMCacheType.REDIS_SEMANTIC})), - CacheRule(Rollout.PYTHON_ONLY, backends=frozenset({LiteLLMCacheType.VALKEY_SEMANTIC})), - CacheRule(Rollout.PYTHON_ONLY, backends=frozenset({LiteLLMCacheType.S3})), - CacheRule(Rollout.PYTHON_ONLY, backends=frozenset({LiteLLMCacheType.DISK})), - CacheRule(Rollout.PYTHON_ONLY, backends=frozenset({LiteLLMCacheType.QDRANT_SEMANTIC})), - CacheRule(Rollout.PYTHON_ONLY, backends=frozenset({LiteLLMCacheType.AZURE_BLOB})), - CacheRule(Rollout.PYTHON_ONLY, backends=frozenset({LiteLLMCacheType.GCS})), SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({KeyManagementSystem.GOOGLE_KMS.value})), SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({KeyManagementSystem.AZURE_KEY_VAULT.value})), SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({KeyManagementSystem.AWS_SECRET_MANAGER.value})), diff --git a/litellm/rust_bridge/public_call.py b/litellm/rust_bridge/public_call.py index d107483ecc0..3cf19026de1 100644 --- a/litellm/rust_bridge/public_call.py +++ b/litellm/rust_bridge/public_call.py @@ -70,11 +70,13 @@ def optional_sequence(value: object) -> Sequence[object] | None: def inference_decline_reason(parameters: tuple[str, ...], kwargs: Mapping[str, object]) -> str | None: - if litellm.cache is not None or litellm.drop_params or litellm.modify_params: - return "native inference does not implement the configured cache or parameter rewrites" + if litellm.drop_params or litellm.modify_params: + return "native inference does not implement the configured parameter rewrites" for name, value in kwargs.items(): if value is None: continue + if name in {"cache", "caching"}: + continue if name not in parameters and name not in _INFERENCE_CONTEXT: return f"native inference does not implement {name}" return None diff --git a/litellm/rust_bridge/response_cache.py b/litellm/rust_bridge/response_cache.py index a6fc121a3b5..f100d9eff57 100644 --- a/litellm/rust_bridge/response_cache.py +++ b/litellm/rust_bridge/response_cache.py @@ -5,11 +5,7 @@ from collections.abc import Awaitable, Mapping, Sequence from dataclasses import dataclass from typing import Final, Protocol, cast -from typing_extensions import ReadOnly, Required, TypedDict, assert_never - -from litellm.rust_bridge.bindings import NativeBinding, native_exception_types -from litellm.rust_bridge.catalog import CacheContext, Rules, decision -from litellm.rust_bridge.configuration import Decision +from typing_extensions import ReadOnly, Required, TypedDict class CacheFacade(Protocol): @@ -62,18 +58,6 @@ class NativeResponseCacheRuntime(Protocol): def ping(self) -> Awaitable[object]: ... -class NativeResponseCacheRuntimeFactory(Protocol): - @staticmethod - def from_cache(cache: CacheFacade) -> NativeResponseCacheRuntime: ... - - -def _runtime_factory(value: object) -> NativeResponseCacheRuntimeFactory | None: - return cast(NativeResponseCacheRuntimeFactory, value) if callable(getattr(value, "from_cache", None)) else None - - -_RUNTIME: Final = NativeBinding("_ResponseCacheRuntime", validate=_runtime_factory) - - @dataclass(frozen=True, slots=True) class ResponseCacheRuntime: native: NativeResponseCacheRuntime @@ -148,35 +132,6 @@ class ResponseCacheRuntime: await self.native.async_flush() -def resolve_response_cache( - cache: CacheFacade, - rules: Rules | None = None, -) -> ResponseCacheRuntime | None: - backend_value: Final = cache.type - backend: Final = str.__str__(backend_value) if isinstance(backend_value, str) else str(backend_value) - selected: Final = decision(CacheContext(backend=backend), rules) - match selected: - case Decision.PYTHON: - return None - case Decision.RUST_WITH_FALLBACK | Decision.RUST_REQUIRED: - factory: Final = _RUNTIME.load() - if factory is None: - if selected is Decision.RUST_REQUIRED: - raise RuntimeError("Rust response cache runtime is unavailable") - return None - try: - return ResponseCacheRuntime(factory.from_cache(cache)) - except Exception as error: - exceptions: Final = native_exception_types() - if exceptions is None or not isinstance(error, exceptions[0]): - raise - if selected is Decision.RUST_REQUIRED: - raise RuntimeError(f"Rust response cache runtime declined the cache: {error}") from error - return None - case _: - assert_never(selected) - - def _duration(value: object) -> float | None: if isinstance(value, bool) or not isinstance(value, int | float): return None diff --git a/litellm/rust_bridge/response_metadata.py b/litellm/rust_bridge/response_metadata.py index ae459710b34..1ef7b8c6595 100644 --- a/litellm/rust_bridge/response_metadata.py +++ b/litellm/rust_bridge/response_metadata.py @@ -2,11 +2,16 @@ from typing import Final, TypeVar from litellm.router_utils.add_retry_fallback_headers import ( _add_headers_to_response, # pyright: ignore[reportPrivateUsage] # reuse the proxy's identity-preserving response metadata writer + get_hidden_params_dict, ) ResultT: Final = TypeVar("ResultT") def mark_rust_response(response: ResultT) -> ResultT: - _add_headers_to_response(response, {"x-litellm-rust": "true"}) + cache_key: Final = get_hidden_params_dict(response).get("cache_key") + _add_headers_to_response( + response, + {"x-litellm-rust": "true", **({"x-litellm-cache-key": cache_key} if isinstance(cache_key, str) else {})}, + ) return response diff --git a/tests/test_litellm_rust/cache/test_azure_blob.py b/tests/test_litellm_rust/cache/test_azure_blob.py index bbbab22baca..064458ae9b0 100644 --- a/tests/test_litellm_rust/cache/test_azure_blob.py +++ b/tests/test_litellm_rust/cache/test_azure_blob.py @@ -16,12 +16,11 @@ from litellm.rust_bridge import _native from litellm.types.caching import LiteLLMCacheType from tests.test_litellm_rust.support.cache import ( CacheLookup, - CacheTestHandle, CacheTestResolver, + activate_native, assert_native_runtime, completion_kwargs, request, - require_rust, ) from tests.test_litellm_rust.support.isolation import rebound @@ -49,26 +48,11 @@ def azure_blob_facade() -> Generator[Cache]: asyncio.run(backend.disconnect()) -def azure_blob_handle(facade: Cache) -> _native._CacheTestHandle: - backend: Final = facade.cache - assert isinstance(backend, AzureBlobCache) - return CacheTestHandle.azure_blob( - backend.container_client.url.removesuffix(f"/{backend.container_client.container_name}"), - backend.container_client.container_name, - ) - - def test_azure_blob_facade_serves_natively_and_python_reads_the_same_blobs(azure_blob_facade: Cache) -> None: backend: Final = azure_blob_facade.cache assert isinstance(backend, AzureBlobCache) - handle: Final = azure_blob_handle(azure_blob_facade) - assert handle.backend == "azure-blob" + activate_native(azure_blob_facade) account_url: Final = backend.container_client.url.removesuffix(f"/{backend.container_client.container_name}") - with pytest.raises(TypeError, match="containers must match"): - 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 = CacheTestResolver(SimpleNamespace(cache=azure_blob_facade)) native: Final = resolver.resolve() assert native.kind == "native" @@ -86,44 +70,50 @@ def test_azure_blob_facade_serves_natively_and_python_reads_the_same_blobs(azure assert stored["response"] == response assert isinstance(stored["timestamp"], float) assert native.lookup(request("sync")) == response - assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="sync") == response + with rebound(azure_blob_facade, "_native_cache", None): + assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="sync") == response backend.set_cache("python", {"timestamp": time.time(), "response": response}) backend.set_cache("legacy", "bare legacy value") backend.container_client.upload_blob("invalid", b"{not json", overwrite=True) assert native.lookup(request("python")) == response - assert native.lookup(request("legacy")) == cast(CacheLookup, azure_blob_facade).get_cache(cache_key="legacy") + with rebound(azure_blob_facade, "_native_cache", None): + assert native.lookup(request("legacy")) == cast(CacheLookup, azure_blob_facade).get_cache(cache_key="legacy") assert native.lookup_batch([request("python"), request("missing"), request("invalid"), request("sync")]) == { "values": [response, None, None, response], "missing_indices": [1, 2], } with rebound(azure_blob_facade, "ttl", 12): - assert resolver.resolve().kind == "python_callback" + with pytest.raises(_native.RustBridgeDeclined): + resolver.resolve() with rebound(backend, "container_client", ContainerClient.from_container_url(backend.container_client.url)): - assert resolver.resolve().kind == "python_callback" + with pytest.raises(_native.RustBridgeDeclined): + resolver.resolve() def custom_get(*_args: object, **_kwargs: object) -> None: return None with rebound(backend, "get_cache", custom_get): - assert resolver.resolve().kind == "python_callback" - assert resolver.resolve().kind == "python_callback" - assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="sync") == response + with pytest.raises(_native.RustBridgeDeclined): + resolver.resolve() + with pytest.raises(_native.RustBridgeDeclined): + resolver.resolve() + with rebound(azure_blob_facade, "_native_cache", None): + assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="sync") == response class CustomBlobCache(AzureBlobCache): pass with rebound(azure_blob_facade, "cache", CustomBlobCache(account_url, backend.container_client.container_name)): - assert resolver.resolve().kind == "python_callback" - with pytest.raises(TypeError): - azure_blob_handle(azure_blob_facade)._bind_facade(azure_blob_facade) + with pytest.raises(_native.RustBridgeDeclined): + resolver.resolve() async def test_azure_blob_native_async_writes_overwrite_batch_and_flush_like_python(azure_blob_facade: Cache) -> None: backend: Final = azure_blob_facade.cache assert isinstance(backend, AzureBlobCache) - azure_blob_handle(azure_blob_facade)._bind_facade(azure_blob_facade) + activate_native(azure_blob_facade) binding: Final = CacheTestResolver(SimpleNamespace(cache=azure_blob_facade)).resolve() assert binding.kind == "native" ping: Final = cast(dict[str, object], await binding.ping()) @@ -136,7 +126,8 @@ async def test_azure_blob_native_async_writes_overwrite_batch_and_flush_like_pyt 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} + with rebound(azure_blob_facade, "_native_cache", None): + 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}]) assert await binding.async_lookup_batch([request("second"), request("missing"), request("first")]) == { @@ -148,17 +139,18 @@ async def test_azure_blob_native_async_writes_overwrite_batch_and_flush_like_pyt assert await binding.async_lookup(request("async")) is None -async def test_azure_blob_rust_required_rule_activates_natively(monkeypatch: pytest.MonkeyPatch) -> None: +async def test_azure_blob_explicit_selection_activates_natively(monkeypatch: pytest.MonkeyPatch) -> None: account_url: Final = os.environ.get("AZURE_BLOB_CACHE_ACCOUNT_URL") if account_url is None: pytest.skip( "live Azure Blob parity needs AZURE_BLOB_CACHE_ACCOUNT_URL plus DefaultAzureCredential inputs in the environment" ) - require_rust(monkeypatch, LiteLLMCacheType.AZURE_BLOB) - facade: Final = Cache( - type=LiteLLMCacheType.AZURE_BLOB, - azure_account_url=account_url, - azure_blob_container=f"litellm-parity-{uuid.uuid4().hex[:12]}", + facade: Final = activate_native( + Cache( + type=LiteLLMCacheType.AZURE_BLOB, + azure_account_url=account_url, + azure_blob_container=f"litellm-parity-{uuid.uuid4().hex[:12]}", + ) ) backend: Final = facade.cache assert isinstance(backend, AzureBlobCache) diff --git a/tests/test_litellm_rust/cache/test_disk.py b/tests/test_litellm_rust/cache/test_disk.py index 4f2907e6a09..e2faaa50221 100644 --- a/tests/test_litellm_rust/cache/test_disk.py +++ b/tests/test_litellm_rust/cache/test_disk.py @@ -10,8 +10,9 @@ import pytest from litellm.caching.caching import Cache from litellm.caching.disk_cache import DiskCache +from litellm.rust_bridge import _native from litellm.types.caching import LiteLLMCacheType -from tests.test_litellm_rust.support.cache import CacheTestHandle, CacheTestResolver, request +from tests.test_litellm_rust.support.cache import CacheTestResolver, activate_native, native_runtime, request from tests.test_litellm_rust.support.isolation import rebound pytestmark: Final = pytest.mark.requires_rust_extension @@ -31,7 +32,7 @@ async def test_disk_reads_python_entries_and_python_reads_native_entries(tmp_pat "large", {"timestamp": time.time(), "response": {"text": "x" * 70_000}}, ) - binding: Final = CacheTestResolver(SimpleNamespace(cache=CacheTestHandle.disk(str(tmp_path)))).resolve() + binding: Final = native_runtime(Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path))) assert binding.lookup(request("sync")) == response assert await binding.async_lookup(request("async")) == response @@ -52,10 +53,10 @@ async def test_disk_reads_python_entries_and_python_reads_native_entries(tmp_pat async def test_disk_entries_survive_a_fresh_handle_and_expire_on_time(tmp_path: Path) -> None: - first: Final = CacheTestResolver(SimpleNamespace(cache=CacheTestHandle.disk(str(tmp_path)))).resolve() + first: Final = native_runtime(Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path))) await first.async_store(request("persistent"), {"value": "persistent"}) await first.async_store({**request("expiring"), "ttl_seconds": 0.3}, {"value": "expiring"}) - fresh: Final = CacheTestResolver(SimpleNamespace(cache=CacheTestHandle.disk(str(tmp_path)))).resolve() + fresh: Final = native_runtime(Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path))) assert fresh.lookup(request("persistent")) == {"value": "persistent"} assert fresh.lookup(request("expiring")) == {"value": "expiring"} await asyncio.sleep(0.4) @@ -63,39 +64,36 @@ async def test_disk_entries_survive_a_fresh_handle_and_expire_on_time(tmp_path: assert fresh.lookup(request("persistent")) == {"value": "persistent"} -def test_disk_facade_registers_and_store_changes_fall_back(tmp_path: Path) -> None: - facade: Final = Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path)) - with pytest.raises(TypeError, match="directories must match"): - CacheTestHandle.disk(str(tmp_path / "other"))._bind_facade(facade) - handle: Final = CacheTestHandle.disk(str(tmp_path)) - handle._bind_facade(facade) - resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade)) - binding: Final = resolver.resolve() - assert binding.kind == "native" - binding.store(request("native"), {"value": "native"}) +def test_selected_disk_runtime_declines_store_changes(tmp_path: Path) -> None: + facade: Final = activate_native(Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path))) + selected: Final = CacheTestResolver(SimpleNamespace(cache=facade)) + native: Final = selected.resolve() + native.store(request("native"), {"value": "native"}) assert facade.get_cache(cache_key="native") == {"value": "native"} - - with rebound(facade.cache, "disk_cache", diskcache.Cache(str(tmp_path))): - assert resolver.resolve().kind == "python_callback" - assert resolver.resolve().kind == "native" - - class CustomDiskCache(DiskCache): - pass - - with rebound(facade, "cache", CustomDiskCache(disk_cache_dir=str(tmp_path))): - assert resolver.resolve().kind == "python_callback" + replacement: Final = diskcache.Cache(str(tmp_path)) + try: + with rebound(facade.cache, "disk_cache", replacement): + with pytest.raises(_native.RustBridgeDeclined): + selected.resolve() + assert selected.resolve().kind == "native" + finally: + replacement.close() class CustomStore(diskcache.Cache): pass - custom_facade: Final = Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path)) - custom_facade.cache.disk_cache = CustomStore(str(tmp_path)) - with pytest.raises(TypeError, match="built-in diskcache store"): - CacheTestHandle.disk(str(tmp_path))._bind_facade(custom_facade) + unsupported: Final = Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path)) + store: Final = CustomStore(str(tmp_path)) + try: + unsupported.cache.disk_cache = store + with pytest.raises(_native.RustBridgeDeclined, match="built-in diskcache store"): + native_runtime(unsupported) + finally: + store.close() async def test_disk_native_batch_lookup_and_store_report_partial_hits(tmp_path: Path) -> None: - binding: Final = CacheTestResolver(SimpleNamespace(cache=CacheTestHandle.disk(str(tmp_path)))).resolve() + binding: Final = native_runtime(Cache(type=LiteLLMCacheType.DISK, disk_cache_dir=str(tmp_path))) requests: Final = [request("hit"), request("miss"), request("disabled")] requests[2]["controls"] = { "supported_call_type": True, diff --git a/tests/test_litellm_rust/cache/test_facade.py b/tests/test_litellm_rust/cache/test_facade.py index d99ea4e2baa..1b91b7edb0c 100644 --- a/tests/test_litellm_rust/cache/test_facade.py +++ b/tests/test_litellm_rust/cache/test_facade.py @@ -11,11 +11,15 @@ import litellm from litellm.caching.caching import Cache, disable_cache, enable_cache, update_cache from litellm.caching.in_memory_cache import InMemoryCache from litellm.rust_bridge import _native -from litellm.rust_bridge.catalog import CacheRule, Route, RouteRule, SecretManagerRule -from litellm.rust_bridge.configuration import Rollout -from litellm.rust_bridge.response_cache import ResponseCacheRuntime, resolve_response_cache +from litellm.rust_bridge.response_cache import ResponseCacheRuntime from litellm.types.caching import LiteLLMCacheType -from tests.test_litellm_rust.support.cache import CacheLookup, CacheTestHandle, CacheTestResolver, request +from tests.test_litellm_rust.support.cache import ( + CacheLookup, + CacheTestResolver, + activate_native, + native_runtime, + request, +) from tests.test_litellm_rust.support.isolation import rebound pytestmark: Final = pytest.mark.requires_rust_extension @@ -24,8 +28,6 @@ pytestmark: Final = pytest.mark.requires_rust_extension def test_existing_constructor_and_global_are_unchanged() -> None: facade: Final = Cache(type=LiteLLMCacheType.LOCAL) assert type(facade.cache) is InMemoryCache - assert "_native_cache_handle" not in vars(facade) - assert resolve_response_cache(facade) is None with rebound(litellm, "cache", facade): resolver: Final = CacheTestResolver(litellm) assert resolver.resolve().kind == "python_callback" @@ -33,14 +35,9 @@ def test_existing_constructor_and_global_are_unchanged() -> None: assert cast(CacheLookup, facade).get_cache(cache_key="key") == {"answer": 7} -async def test_catalog_constructs_native_runtime_from_public_cache_configuration() -> None: - rules: Final = ( - RouteRule(Route.OCR, Rollout.PYTHON_ONLY), - SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({"local"})), - CacheRule(Rollout.RUST_REQUIRED, backends=frozenset({"local"})), - ) +async def test_explicit_selection_constructs_native_runtime_from_public_cache_configuration() -> None: facade: Final = Cache(type=LiteLLMCacheType.LOCAL) - runtime: Final = resolve_response_cache(facade, rules) + runtime: Final = ResponseCacheRuntime(_native._ResponseCacheRuntime.from_cache(facade)) assert isinstance(runtime, ResponseCacheRuntime) assert runtime.kind == "native" @@ -70,17 +67,12 @@ async def test_catalog_constructs_native_runtime_from_public_cache_configuration async def test_inference_resolver_uses_the_configured_native_cache_directly() -> None: - rules: Final = ( - RouteRule(Route.OCR, Rollout.PYTHON_ONLY), - SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({"local"})), - CacheRule(Rollout.RUST_REQUIRED, backends=frozenset({"local"})), - ) facade: Final = Cache(type=LiteLLMCacheType.LOCAL) - runtime: Final = resolve_response_cache(facade, rules) + runtime: Final = ResponseCacheRuntime(_native._ResponseCacheRuntime.from_cache(facade)) assert isinstance(runtime, ResponseCacheRuntime) facade._native_cache = runtime - selected: Final = _native._CacheResolver(SimpleNamespace(cache=facade)).resolve() + selected: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() assert selected.kind == "native" request: Final = runtime.request(facade, {"cache_key": "inference-native"}) assert request is not None @@ -90,7 +82,7 @@ async def test_inference_resolver_uses_the_configured_native_cache_directly() -> assert facade.cache.get_cache("inference-native") is None facade._native_cache = None - fallback: Final = _native._CacheResolver(SimpleNamespace(cache=facade)).resolve() + fallback: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() assert fallback.kind == "python_callback" await fallback.async_store(None, {"answer": 7}, callback_kwargs={"cache_key": "inference-python"}) assert facade.get_cache(cache_key="inference-python") == {"answer": 7} @@ -98,13 +90,8 @@ async def test_inference_resolver_uses_the_configured_native_cache_directly() -> async def test_inference_resolver_declines_a_native_runtime_whose_facade_changed() -> None: - rules: Final = ( - RouteRule(Route.OCR, Rollout.PYTHON_ONLY), - SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({"local"})), - CacheRule(Rollout.RUST_REQUIRED, backends=frozenset({"local"})), - ) facade: Final = Cache(type=LiteLLMCacheType.LOCAL) - runtime: Final = resolve_response_cache(facade, rules) + runtime: Final = ResponseCacheRuntime(_native._ResponseCacheRuntime.from_cache(facade)) assert isinstance(runtime, ResponseCacheRuntime) facade._native_cache = runtime stale_request: Final = runtime.request(facade, {"cache_key": "stale-only"}) @@ -114,7 +101,7 @@ async def test_inference_resolver_declines_a_native_runtime_whose_facade_changed replacement: Final = InMemoryCache() facade.cache = replacement with pytest.raises(_native.RustBridgeDeclined): - _native._CacheResolver(SimpleNamespace(cache=facade)).resolve() + CacheTestResolver(SimpleNamespace(cache=facade)).resolve() assert await runtime.async_lookup(stale_request) == {"answer": "stale"} assert replacement.get_cache("stale-only") is None assert replacement.get_cache("swapped-backend") is None @@ -144,13 +131,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=CacheTestHandle.memory()) + namespace: Final = SimpleNamespace(cache=activate_native(Cache(type=LiteLLMCacheType.LOCAL))) 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", CacheTestHandle.memory()): + with rebound(namespace, "cache", activate_native(Cache(type=LiteLLMCacheType.LOCAL))): replacement: Final = resolver.resolve() await selected.async_store(request(), {"answer": 2}) assert replacement.lookup(request()) is None @@ -217,62 +204,30 @@ async def test_callback_cancellation_stays_in_the_callers_task() -> None: assert finished.is_set() -def test_registered_facade_uses_native_and_instance_overrides_fall_back() -> None: - facade: Final = Cache(type=LiteLLMCacheType.LOCAL) - handle: Final = CacheTestHandle.memory() - handle._bind_facade(facade) - resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade)) - native: Final = resolver.resolve() - assert native.kind == "native" +@pytest.mark.parametrize("method", ("get_cache", "get_cache_key", "async_get_cache")) +def test_selected_native_runtime_declines_instance_overrides(method: str) -> None: + facade: Final = activate_native(Cache(type=LiteLLMCacheType.LOCAL)) + selected: Final = CacheTestResolver(SimpleNamespace(cache=facade)) + native: Final = selected.resolve() native.store(request(), {"source": "native"}) + + def override(**_kwargs: object) -> None: + return None + + with rebound(facade, method, override): + with pytest.raises(_native.RustBridgeDeclined): + selected.resolve() assert native.lookup(request()) == {"source": "native"} - assert cast(CacheLookup, facade).get_cache(cache_key="key") is None - sentinel: Final = object() - - def outer_override(**_kwargs: object) -> object: - return sentinel - - def backend_override(*_args: object, **_kwargs: object) -> dict[str, str]: - return {"source": "override"} - - with rebound(facade, "get_cache", outer_override): - fallback: Final = resolver.resolve() - assert fallback.kind == "python_callback" - assert fallback.lookup(None, callback_kwargs={"cache_key": "key"}) is sentinel - assert resolver.resolve().kind == "python_callback" - delattr(facade, "get_cache") - assert resolver.resolve().kind == "native" - with rebound(facade.cache, "get_cache", backend_override): - backend_fallback: Final = resolver.resolve() - assert backend_fallback.kind == "python_callback" - assert backend_fallback.lookup(None, callback_kwargs={"cache_key": "key"}) == {"source": "override"} -def test_facade_subclasses_backend_replacement_and_configuration_changes_are_not_bypassed() -> None: - class CustomCache(Cache): - pass - - 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 = CacheTestResolver(SimpleNamespace(cache=facade)) - with rebound(facade, "cache", InMemoryCache()): - assert resolver.resolve().kind == "python_callback" - with rebound(facade, "ttl", 12): - assert resolver.resolve().kind == "python_callback" - with rebound(facade, "semantic_cache_scope", "end_user"): - assert resolver.resolve().kind == "python_callback" - - def custom_key(**_kwargs: object) -> str: - return "custom" - - with rebound(facade, "get_cache_key", custom_key): - assert resolver.resolve().kind == "python_callback" - assert resolver.resolve().kind == "python_callback" - delattr(facade, "get_cache_key") - assert resolver.resolve().kind == "native" +@pytest.mark.parametrize(("attribute", "value"), (("ttl", 12), ("semantic_cache_scope", "end_user"))) +def test_selected_native_runtime_declines_policy_changes(attribute: str, value: object) -> None: + facade: Final = activate_native(Cache(type=LiteLLMCacheType.LOCAL)) + selected: Final = CacheTestResolver(SimpleNamespace(cache=facade)) + with rebound(facade, attribute, value): + with pytest.raises(_native.RustBridgeDeclined): + selected.resolve() + assert selected.resolve().kind == "native" def test_resolver_and_callback_cycles_can_be_collected() -> None: @@ -292,31 +247,41 @@ def test_resolver_and_callback_cycles_can_be_collected() -> None: def test_invalid_duration_and_request_shape_fail_before_storage() -> None: - binding: Final = CacheTestResolver(SimpleNamespace(cache=CacheTestHandle.memory())).resolve() + binding: Final = CacheTestResolver( + SimpleNamespace(cache=activate_native(Cache(type=LiteLLMCacheType.LOCAL))) + ).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 + facade: Final = Cache(type=LiteLLMCacheType.LOCAL) + facade.cache = InMemoryCache(default_ttl=-1) with pytest.raises(ValueError, match="cache durations must be finite and nonnegative"): - CacheTestHandle.memory(ttl_seconds=-1) + native_runtime(facade) async def test_memory_size_policy_is_applied_by_the_native_host() -> None: - handle: Final = CacheTestHandle.memory(capacity=2, max_entry_bytes=128) + facade: Final = Cache(type=LiteLLMCacheType.LOCAL) + facade.cache = InMemoryCache(max_size_in_memory=2, max_size_per_item=1) + handle: Final = activate_native(facade) 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}) + await binding.async_store(request("large"), {"answer": "x" * 2048}) assert binding.lookup(request("large")) is None assert binding.lookup(request("small")) == small - disabled: Final = CacheTestResolver(SimpleNamespace(cache=CacheTestHandle.memory(capacity=0))).resolve() + disabled_facade: Final = Cache(type=LiteLLMCacheType.LOCAL) + disabled_facade.cache = InMemoryCache(max_size_in_memory=0) + disabled: Final = native_runtime(disabled_facade) 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 = CacheTestResolver(SimpleNamespace(cache=CacheTestHandle.memory())).resolve() + binding: Final = CacheTestResolver( + SimpleNamespace(cache=activate_native(Cache(type=LiteLLMCacheType.LOCAL))) + ).resolve() requests: Final = [request("hit"), request("miss"), request("disabled")] requests[2]["controls"] = { "supported_call_type": True, @@ -389,9 +354,3 @@ async def test_unmodified_builtin_cache_callbacks_can_ping_and_flush() -> None: assert await binding.ping() == "pong" await binding.async_flush() assert cache.cache.get_cache("key") is None - - -def test_facade_registration_rejects_mismatched_capacity() -> None: - facade: Final = Cache(type=LiteLLMCacheType.LOCAL) - with pytest.raises(TypeError, match="capacities must match"): - CacheTestHandle.memory(capacity=7)._bind_facade(facade) diff --git a/tests/test_litellm_rust/cache/test_gcs.py b/tests/test_litellm_rust/cache/test_gcs.py index bfc9ebbb4d7..5b81e48bc1d 100644 --- a/tests/test_litellm_rust/cache/test_gcs.py +++ b/tests/test_litellm_rust/cache/test_gcs.py @@ -1,242 +1,46 @@ -import json -import time -from collections.abc import Generator from types import SimpleNamespace -from typing import Final, cast +from typing import Final import pytest from litellm.caching.caching import Cache -from litellm.caching.gcs_cache import GCSCache +from litellm.rust_bridge import _native from litellm.types.caching import LiteLLMCacheType -from tests.test_litellm_rust.support.cache import CacheLookup, CacheTestHandle, CacheTestResolver, request -from tests.test_litellm_rust.support.fake_gcs import FakeGcs +from tests.test_litellm_rust.support.cache import CacheTestResolver, activate_native, native_runtime from tests.test_litellm_rust.support.isolation import rebound pytestmark: Final = pytest.mark.requires_rust_extension -@pytest.fixture -def fake_gcs() -> Generator[FakeGcs]: - server: Final = FakeGcs() - try: - yield server - finally: - server.close() - - -async def test_gcs_reads_python_entries_and_writes_python_compatible_objects( - fake_gcs: FakeGcs, monkeypatch: pytest.MonkeyPatch +@pytest.mark.parametrize( + ("attribute", "replacement"), + (("bucket_name", "other"), ("key_prefix", "other/"), ("path_service_account", "other.json")), +) +def test_selected_gcs_runtime_declines_backend_configuration_changes( + monkeypatch: pytest.MonkeyPatch, attribute: str, replacement: str ) -> None: monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) monkeypatch.delenv("GCS_BUCKET_NAME", raising=False) - response: Final = {"choices": [{"text": "cached"}], "usage": {"total_tokens": 3}, "flag": True, "empty": None} - fake_gcs.put( - "bucket", - "cache/sync", - json.dumps({"timestamp": time.time(), "response": json.dumps(response)}).encode(), - ) - fake_gcs.put("bucket", "cache/async", json.dumps({"timestamp": time.time(), "response": response}).encode()) - fake_gcs.put("bucket", "cache/raw", json.dumps(response).encode()) - fake_gcs.put("bucket", "cache/invalid", b"not a cache entry") - binding: Final = CacheTestResolver( - SimpleNamespace( - cache=CacheTestHandle.gcs( - "bucket", - gcs_path="cache", - endpoint=fake_gcs.url, - token=fake_gcs.token, - ) - ) - ).resolve() - - assert binding.lookup(request("sync")) == response - assert await binding.async_lookup(request("async")) == response - assert binding.lookup(request("raw")) == response - assert await binding.async_lookup(request("invalid")) is None - assert binding.lookup(request("missing")) is None - - await binding.async_store({**request("native"), "ttl_seconds": 12.0}, response) - stored: Final = fake_gcs.objects[("bucket", "cache/native")] - stored_value: Final = cast(dict[str, object], json.loads(stored)) - assert stored_value["response"] == response - assert isinstance(stored_value["timestamp"], float) - upload: Final = next(item for item in fake_gcs.requests if item.method == "POST") - assert upload.path == "/upload/storage/v1/b/bucket/o" - assert upload.query == "uploadType=media&name=cache%2Fnative" - assert upload.headers["Authorization"] == f"Bearer {fake_gcs.token}" - assert upload.headers["Content-Type"] == "application/json" - upload_text: Final = f"{upload.path}?{upload.query}{upload.headers}" - assert "ttl" not in upload_text.lower() - assert "expiry" not in upload_text.lower() - download: Final = next(item for item in fake_gcs.requests if item.path.endswith("/cache%2Fsync")) - assert download.path == "/storage/v1/b/bucket/o/cache%2Fsync" - assert download.query == "alt=media" - - binding.store(request("sync2"), response) - assert binding.lookup(request("sync2")) == response - assert GCSCache(bucket_name="bucket", gcs_path="cache").key_prefix == "cache/" - assert GCSCache(bucket_name="bucket", gcs_path="cache/").key_prefix == "cache/" - assert GCSCache(bucket_name="bucket").key_prefix == "" + facade: Final = activate_native(Cache(type=LiteLLMCacheType.GCS, gcs_bucket_name="bucket", gcs_path="cache/")) + selected: Final = CacheTestResolver(SimpleNamespace(cache=facade)) + assert selected.resolve().kind == "native" + with rebound(facade.cache, attribute, replacement): + with pytest.raises(_native.RustBridgeDeclined): + selected.resolve() + assert selected.resolve().kind == "native" -async def test_gcs_batch_lookup_preserves_order_and_treats_malformed_entries_as_misses(fake_gcs: FakeGcs) -> None: - fake_gcs.put("bucket", "cache/hit", json.dumps({"timestamp": time.time(), "response": {"value": 1}}).encode()) - fake_gcs.put("bucket", "cache/invalid", b"not a cache entry") - binding: Final = CacheTestResolver( - SimpleNamespace( - cache=CacheTestHandle.gcs( - "bucket", - gcs_path="cache", - endpoint=fake_gcs.url, - token=fake_gcs.token, - ) - ) - ).resolve() - requests: Final = [request("hit"), request("missing"), request("invalid")] - expected: Final = {"values": [{"value": 1}, None, None], "missing_indices": [1, 2]} - - assert await binding.async_lookup_batch(requests) == expected - assert binding.lookup_batch(requests) == expected - await binding.async_store_batch([request("first"), request("second")], [{"value": 1}, {"value": 2}]) - assert ("bucket", "cache/first") in fake_gcs.objects - assert ("bucket", "cache/second") in fake_gcs.objects - - -async def test_gcs_facade_binds_only_exact_matching_configuration( - fake_gcs: FakeGcs, monkeypatch: pytest.MonkeyPatch -) -> None: +async def test_gcs_runtime_flush_is_a_no_op_and_ping_is_not_implemented(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) monkeypatch.delenv("GCS_BUCKET_NAME", raising=False) - monkeypatch.setenv("GOOGLE_APPLICATION_CREDENTIALS", "/nonexistent") - facade: Final = Cache(type=LiteLLMCacheType.GCS, gcs_bucket_name="bucket", gcs_path="cache/") - assert type(facade.cache) is GCSCache - - mismatched_bucket: Final = CacheTestHandle.gcs( - "other", - gcs_path="cache", - endpoint=fake_gcs.url, - token=fake_gcs.token, - ) - with pytest.raises(TypeError, match="buckets must match"): - mismatched_bucket._bind_facade(facade) - mismatched_prefix: Final = CacheTestHandle.gcs( - "bucket", - gcs_path="x", - endpoint=fake_gcs.url, - token=fake_gcs.token, - ) - with pytest.raises(TypeError, match="key prefixes must match"): - mismatched_prefix._bind_facade(facade) - mismatched_credentials: Final = CacheTestHandle.gcs( - "bucket", - gcs_path="cache", - path_service_account="sa.json", - endpoint=fake_gcs.url, - token=fake_gcs.token, - ) - with pytest.raises(TypeError, match="credentials must match"): - mismatched_credentials._bind_facade(facade) - with pytest.raises(TypeError, match="types must match"): - CacheTestHandle.memory()._bind_facade(facade) - - matching: Final = CacheTestHandle.gcs( - "bucket", - gcs_path="cache", - endpoint=fake_gcs.url, - token=fake_gcs.token, - ) - matching._bind_facade(facade) - resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade)) - binding: Final = resolver.resolve() - assert binding.kind == "native" - await binding.async_store(request("native"), {"value": "native"}) - assert await binding.async_lookup(request("native")) == {"value": "native"} - assert cast(CacheLookup, facade).get_cache(cache_key="native") is None - - with rebound(facade.cache, "bucket_name", "other"): - assert resolver.resolve().kind == "python_callback" - with rebound(facade.cache, "key_prefix", "x/"): - 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 - - with rebound(facade.cache, "get_cache", no_get_cache): - assert resolver.resolve().kind == "python_callback" - with rebound(facade, "ttl", 12): - assert resolver.resolve().kind == "python_callback" - - class CustomGcs(GCSCache): - pass - - with rebound(facade, "cache", CustomGcs(bucket_name="bucket", gcs_path="cache/")): - assert resolver.resolve().kind == "python_callback" - custom_facade: Final = Cache(type=LiteLLMCacheType.GCS, gcs_bucket_name="bucket", gcs_path="cache/") - with rebound(custom_facade, "cache", CustomGcs(bucket_name="bucket", gcs_path="cache/")): - with pytest.raises(TypeError, match="types must match"): - matching._bind_facade(custom_facade) - - missing_bucket: Final = Cache(type=LiteLLMCacheType.GCS) - with pytest.raises(TypeError, match="requires a configured bucket name"): - matching._bind_facade(missing_bucket) - - -async def test_gcs_flush_is_a_no_op_and_ping_is_not_implemented( - fake_gcs: FakeGcs, monkeypatch: pytest.MonkeyPatch -) -> None: - monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) - monkeypatch.delenv("GCS_BUCKET_NAME", raising=False) - binding: Final = CacheTestResolver( - SimpleNamespace( - cache=CacheTestHandle.gcs( - "bucket", - gcs_path="cache", - endpoint=fake_gcs.url, - token=fake_gcs.token, - ) - ) - ).resolve() - await binding.async_store(request("key"), {"value": "stored"}) - await binding.async_flush() - assert ("bucket", "cache/key") in fake_gcs.objects - assert await binding.async_lookup(request("key")) == {"value": "stored"} + runtime: Final = native_runtime(Cache(type=LiteLLMCacheType.GCS, gcs_bucket_name="bucket")) + await runtime.async_flush() with pytest.raises(NotImplementedError): - await binding.ping() - - facade: Final = Cache(type=LiteLLMCacheType.GCS, gcs_bucket_name="bucket", gcs_path="cache/") - with pytest.raises(AttributeError): - await facade.ping() - assert cast(CacheLookup, facade.cache).flush_cache() is None + await runtime.ping() -async def test_gcs_unauthorized_and_server_errors_surface_as_runtime_errors(fake_gcs: FakeGcs) -> None: - wrong_token: Final = CacheTestResolver( - SimpleNamespace( - cache=CacheTestHandle.gcs( - "bucket", - gcs_path="cache", - endpoint=fake_gcs.url, - token="wrong-token", - ) - ) - ).resolve() - with pytest.raises(RuntimeError): - wrong_token.lookup(request("missing")) - assert not fake_gcs.objects - - binding: Final = CacheTestResolver( - SimpleNamespace( - cache=CacheTestHandle.gcs( - "bucket", - gcs_path="cache", - endpoint=fake_gcs.url, - token=fake_gcs.token, - ) - ) - ).resolve() - with pytest.raises(RuntimeError): - binding.lookup(request("server-error")) - assert binding.lookup(request("missing")) is None +def test_gcs_runtime_declines_missing_bucket_configuration(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) + monkeypatch.delenv("GCS_BUCKET_NAME", raising=False) + with pytest.raises(_native.RustBridgeDeclined, match="requires a configured bucket name"): + native_runtime(Cache(type=LiteLLMCacheType.GCS)) diff --git a/tests/test_litellm_rust/cache/test_qdrant_semantic.py b/tests/test_litellm_rust/cache/test_qdrant_semantic.py index 160089c9002..529ec66bdd8 100644 --- a/tests/test_litellm_rust/cache/test_qdrant_semantic.py +++ b/tests/test_litellm_rust/cache/test_qdrant_semantic.py @@ -13,13 +13,14 @@ from uuid import uuid4 import pytest from litellm.caching.caching import Cache +from litellm.rust_bridge import _native from litellm.types.caching import LiteLLMCacheType from tests.test_litellm_rust.support.cache import ( - CacheTestHandle, CacheTestResolver, + activate_native, assert_native_runtime, + native_runtime, request, - require_rust, ) pytestmark: Final = pytest.mark.requires_rust_extension @@ -111,13 +112,7 @@ def test_qdrant_semantic_facade_binds_native_and_shares_entries(qdrant_url: str, {"timestamp": time.time(), "response": json.dumps({"id": "py"})}, messages=messages, ) - handle: Final = CacheTestHandle.qdrant_semantic( - qdrant_url, - collection_name=collection, - similarity_threshold=0.99, - vector_size=8, - ) - handle._bind_facade(facade) + activate_native(facade) binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() assert binding.kind == "native" assert binding.lookup(qdrant_request("python-key", messages)) == {"id": "py"} @@ -137,13 +132,7 @@ async def test_qdrant_semantic_async_parity(qdrant_url: str, fake_embedding_endp messages: Final = [{"role": "user", "content": "async prompt"}] collection: Final = f"cache_{uuid4().hex}" facade: Final = qdrant_facade(qdrant_url, collection) - handle: Final = CacheTestHandle.qdrant_semantic( - qdrant_url, - collection_name=collection, - similarity_threshold=0.99, - vector_size=8, - ) - handle._bind_facade(facade) + activate_native(facade) binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() await facade.cache.async_set_cache( "python-key", @@ -161,13 +150,7 @@ async def test_qdrant_semantic_async_store_batch_shares_entries(qdrant_url: str, del fake_embedding_endpoint collection: Final = f"cache_{uuid4().hex}" facade: Final = qdrant_facade(qdrant_url, collection) - handle: Final = CacheTestHandle.qdrant_semantic( - qdrant_url, - collection_name=collection, - similarity_threshold=0.99, - vector_size=8, - ) - handle._bind_facade(facade) + activate_native(facade) binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() entries: Final = [ qdrant_request("batch-one", [{"role": "user", "content": "first batch prompt"}]), @@ -192,13 +175,7 @@ async def test_qdrant_semantic_malformed_entries_and_unsupported_operations( messages: Final = [{"role": "user", "content": "malformed prompt"}] collection: Final = f"cache_{uuid4().hex}" facade: Final = qdrant_facade(qdrant_url, collection) - handle: Final = CacheTestHandle.qdrant_semantic( - qdrant_url, - collection_name=collection, - similarity_threshold=0.99, - vector_size=8, - ) - handle._bind_facade(facade) + activate_native(facade) binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() key: Final = "malformed-key" response: Final = { @@ -233,13 +210,7 @@ def test_qdrant_semantic_ignores_request_expiry(qdrant_url: str, fake_embedding_ messages: Final = [{"role": "user", "content": "persistent prompt"}] collection: Final = f"cache_{uuid4().hex}" facade: Final = qdrant_facade(qdrant_url, collection) - handle: Final = CacheTestHandle.qdrant_semantic( - qdrant_url, - collection_name=collection, - similarity_threshold=0.99, - vector_size=8, - ) - handle._bind_facade(facade) + activate_native(facade) binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() binding.store(qdrant_request("persistent-key", messages, ttl_seconds=1.0), {"id": "persistent"}) time.sleep(1.2) @@ -249,37 +220,34 @@ def test_qdrant_semantic_ignores_request_expiry(qdrant_url: str, fake_embedding_ assert python_value["response"] == {"id": "persistent"} -def test_qdrant_semantic_mutation_and_projection_fallback(qdrant_url: str, fake_embedding_endpoint: str) -> None: +def test_qdrant_runtime_declines_mutation_and_unsupported_configuration( + qdrant_url: str, fake_embedding_endpoint: str +) -> None: del fake_embedding_endpoint collection: Final = f"cache_{uuid4().hex}" facade: Final = qdrant_facade(qdrant_url, collection) - handle: Final = CacheTestHandle.qdrant_semantic( - qdrant_url, - collection_name=collection, - similarity_threshold=0.99, - vector_size=8, - ) - handle._bind_facade(facade) + activate_native(facade) facade.cache.qdrant_api_key = "rotated" - assert CacheTestResolver(SimpleNamespace(cache=facade)).resolve().kind == "python_callback" + with pytest.raises(_native.RustBridgeDeclined): + CacheTestResolver(SimpleNamespace(cache=facade)).resolve() facade.cache.similarity_threshold = 0.5 - assert CacheTestResolver(SimpleNamespace(cache=facade)).resolve().kind == "python_callback" + with pytest.raises(_native.RustBridgeDeclined): + CacheTestResolver(SimpleNamespace(cache=facade)).resolve() unsupported: Final = qdrant_facade(qdrant_url, f"cache_{uuid4().hex}") unsupported.cache.embedding_max_input_tokens = 100 - with pytest.raises(TypeError, match="requires Python"): - handle._bind_facade(unsupported) + with pytest.raises(_native.RustBridgeDeclined, match="requires Python"): + native_runtime(unsupported) unsupported.cache.embedding_max_input_tokens = None unsupported.cache.qdrant_api_base = "http://127.0.0.1:7777" - with pytest.raises(TypeError, match="gRPC"): - handle._bind_facade(unsupported) + with pytest.raises(_native.RustBridgeDeclined, match="gRPC"): + native_runtime(unsupported) -def test_qdrant_semantic_rust_required_rule_activates_natively( +def test_qdrant_semantic_explicit_selection_activates_natively( qdrant_url: str, fake_embedding_endpoint: str, monkeypatch: pytest.MonkeyPatch ) -> None: del fake_embedding_endpoint - require_rust(monkeypatch, LiteLLMCacheType.QDRANT_SEMANTIC) - facade: Final = qdrant_facade(qdrant_url, f"cache_{uuid4().hex}") + facade: Final = activate_native(qdrant_facade(qdrant_url, f"cache_{uuid4().hex}")) assert_native_runtime(facade) kwargs: Final = {"model": "gpt-4o", "messages": [{"role": "user", "content": "qdrant activation"}]} facade.add_cache({"answer": "qdrant"}, **kwargs) diff --git a/tests/test_litellm_rust/cache/test_redis.py b/tests/test_litellm_rust/cache/test_redis.py index dd88145ef21..881db9b1f2e 100644 --- a/tests/test_litellm_rust/cache/test_redis.py +++ b/tests/test_litellm_rust/cache/test_redis.py @@ -11,17 +11,15 @@ import redis import litellm from litellm.caching.caching import Cache from litellm.caching.redis_cluster_cache import RedisClusterCache -from litellm.rust_bridge import catalog -from litellm.rust_bridge.catalog import CacheRule -from litellm.rust_bridge.configuration import Rollout +from litellm.rust_bridge import _native from litellm.types.caching import LiteLLMCacheType from tests.test_litellm_rust.support.cache import ( - CacheTestHandle, CacheTestResolver, + activate_native, assert_native_runtime, completion_kwargs, + native_runtime, request, - require_rust, ) from tests.test_litellm_rust.support.isolation import rebound @@ -38,8 +36,7 @@ def cluster_nodes() -> tuple[tuple[str, int], ...]: 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=CacheTestHandle.redis(redis_url, namespace="team")) - binding: Final = CacheTestResolver(namespace).resolve() + binding: Final = native_runtime(redis_facade(redis_url, namespace="team")) 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)) @@ -69,20 +66,18 @@ async def test_redis_facade_buffers_native_async_writes(redis_url: str) -> None: port=str(parsed.port), redis_flush_size=2, ) - with pytest.raises(TypeError, match="default TTLs must match"): - CacheTestHandle.redis(redis_url, ttl_seconds=61)._bind_facade(facade) - with pytest.raises(TypeError, match="namespaces must match"): - CacheTestHandle.redis(redis_url, namespace="other")._bind_facade(facade) - CacheTestHandle.redis(redis_url, ttl_seconds=60)._bind_facade(facade) + activate_native(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 CacheTestResolver(SimpleNamespace(cache=facade)).resolve().kind == "python_callback" + with pytest.raises(_native.RustBridgeDeclined): + CacheTestResolver(SimpleNamespace(cache=facade)).resolve() pool: Final = facade.cache.redis_client.connection_pool with rebound(pool, "connection_kwargs", {**pool.connection_kwargs, "db": 1}): - assert CacheTestResolver(SimpleNamespace(cache=facade)).resolve().kind == "python_callback" + with pytest.raises(_native.RustBridgeDeclined): + CacheTestResolver(SimpleNamespace(cache=facade)).resolve() await binding.async_store(request("first"), {"value": 1}) assert client.get("first") is None @@ -98,21 +93,20 @@ async def test_redis_cluster_facade_serves_multi_slot_batches_and_scoped_flush_n cluster_nodes: tuple[tuple[str, int], ...], ) -> None: startup_nodes: Final = [{"host": host, "port": port} for host, port in cluster_nodes] - url: Final = f"redis://{cluster_nodes[0][0]}:{cluster_nodes[0][1]}" with rebound(litellm, "default_redis_ttl", 60): facade: Final = Cache(type=LiteLLMCacheType.REDIS, redis_startup_nodes=startup_nodes, namespace="parity") assert type(facade.cache) is RedisClusterCache - with pytest.raises(TypeError, match="types must match"): - CacheTestHandle.redis(url, namespace="parity")._bind_facade(facade) - CacheTestHandle.redis(url, namespace="parity", startup_nodes=list(cluster_nodes))._bind_facade(facade) + activate_native(facade) resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade)) assert resolver.resolve().kind == "native" manager: Final = facade.cache.redis_client.nodes_manager with rebound(manager, "connection_kwargs", {**manager.connection_kwargs, "db": 1}): - assert resolver.resolve().kind == "python_callback" + with pytest.raises(_native.RustBridgeDeclined): + resolver.resolve() with rebound(facade.cache, "redis_kwargs", {**facade.cache.redis_kwargs, "startup_nodes": startup_nodes[:1]}): - assert resolver.resolve().kind == "python_callback" + with pytest.raises(_native.RustBridgeDeclined): + resolver.resolve() binding: Final = resolver.resolve() assert binding.kind == "native" @@ -190,19 +184,16 @@ def redis_facade(redis_url: str, **settings: object) -> Cache: def test_redis_settings_the_native_client_cannot_honor_decline( redis_url: str, monkeypatch: pytest.MonkeyPatch, settings: dict[str, object], message: str ) -> None: - require_rust(monkeypatch, LiteLLMCacheType.REDIS) - with pytest.raises(RuntimeError, match=f"declined the cache: native Redis.*{message}"): - redis_facade(redis_url, **settings) + with pytest.raises(_native.RustBridgeDeclined, match=f"native Redis.*{message}"): + activate_native(redis_facade(redis_url, **settings)) def test_redis_verified_tls_activates_natively(redis_url: str, monkeypatch: pytest.MonkeyPatch) -> None: - require_rust(monkeypatch, LiteLLMCacheType.REDIS) - assert_native_runtime(redis_facade(redis_url, ssl=True, ssl_check_hostname=True)) + assert_native_runtime(activate_native(redis_facade(redis_url, ssl=True, ssl_check_hostname=True))) async def test_redis_flush_size_buffers_native_facade_writes(redis_url: str, monkeypatch: pytest.MonkeyPatch) -> None: - require_rust(monkeypatch, LiteLLMCacheType.REDIS) - facade: Final = redis_facade(redis_url, redis_flush_size=2, namespace="team") + facade: Final = activate_native(redis_facade(redis_url, redis_flush_size=2, namespace="team")) assert_native_runtime(facade) client: Final = redis.Redis.from_url(redis_url) first: Final = completion_kwargs("first") @@ -217,12 +208,5 @@ async def test_redis_flush_size_buffers_native_facade_writes(redis_url: str, mon client.close() -def test_rust_with_fallback_keeps_python_when_the_native_client_declines( - redis_url: str, monkeypatch: pytest.MonkeyPatch -) -> None: - monkeypatch.setattr( - catalog, - "RULES", - (CacheRule(Rollout.RUST_OPT_OUT, backends=frozenset({LiteLLMCacheType.REDIS})),), - ) +def test_legacy_constructor_accepts_python_only_settings(redis_url: str, monkeypatch: pytest.MonkeyPatch) -> None: assert redis_facade(redis_url, socket_timeout=1.0)._native_cache is None # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor diff --git a/tests/test_litellm_rust/cache/test_redis_semantic.py b/tests/test_litellm_rust/cache/test_redis_semantic.py index 279330d9060..a8b0c174d7e 100644 --- a/tests/test_litellm_rust/cache/test_redis_semantic.py +++ b/tests/test_litellm_rust/cache/test_redis_semantic.py @@ -16,15 +16,15 @@ import redis import litellm from litellm.caching.caching import Cache from litellm.caching.redis_semantic_cache import RedisSemanticCache +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.cache import ( - CacheTestHandle, CacheTestResolver, + activate_native, assert_native_runtime, request, - require_rust, ) from tests.test_litellm_rust.support.isolation import rebound @@ -198,7 +198,7 @@ def semantic_facade(url: str, index: str, *, similarity_threshold: float = 0.8) redis_semantic_cache_embedding_model=SEMANTIC_EMBEDDING_MODEL, redis_semantic_cache_index_name=index, ) - CacheTestHandle.redis_semantic(facade.cache)._bind_facade(facade) + activate_native(facade) return facade @@ -214,9 +214,7 @@ def test_redis_semantic_constructor_identity_and_provenance( 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" + assert_native_runtime(facade) binding: Final = CacheTestResolver(SimpleNamespace(cache=facade)).resolve() assert binding.kind == "native" @@ -508,7 +506,7 @@ def test_redis_semantic_scope_overrides_the_tag_and_isolates_entries( client.close() -def test_redis_semantic_configuration_drift_falls_back_to_python( +def test_selected_redis_semantic_runtime_declines_configuration_drift( redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding, monkeypatch: pytest.MonkeyPatch, @@ -519,86 +517,42 @@ def test_redis_semantic_configuration_drift_falls_back_to_python( assert resolver.resolve().kind == "native" with rebound(facade.cache, "similarity_threshold", 0.5): - assert resolver.resolve().kind == "python_callback" + with pytest.raises(_native.RustBridgeDeclined): + resolver.resolve() with rebound(facade, "semantic_cache_scope", "end_user"): - assert resolver.resolve().kind == "python_callback" + with pytest.raises(_native.RustBridgeDeclined): + resolver.resolve() with rebound(facade.cache, "embedding_model", "other-model"): - assert resolver.resolve().kind == "python_callback" + with pytest.raises(_native.RustBridgeDeclined): + resolver.resolve() with rebound(facade.cache, "_index_name", "other-index"): - assert resolver.resolve().kind == "python_callback" + with pytest.raises(_native.RustBridgeDeclined): + resolver.resolve() with rebound(facade.cache, "CACHE_KEY_FIELD_NAME", "other-field"): - assert resolver.resolve().kind == "python_callback" + with pytest.raises(_native.RustBridgeDeclined): + resolver.resolve() 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" + with pytest.raises(_native.RustBridgeDeclined): + resolver.resolve() -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) - - -async def test_redis_semantic_rust_required_rule_activates_natively( +async def test_redis_semantic_explicit_selection_activates_natively( redis_stack: tuple[str, str], semantic_embedding: DeterministicEmbedding, monkeypatch: pytest.MonkeyPatch ) -> None: del semantic_embedding url, index = redis_stack - require_rust(monkeypatch, LiteLLMCacheType.REDIS_SEMANTIC) - 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, + facade: Final = activate_native( + 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, + ) ) assert_native_runtime(facade) kwargs: Final = {"model": "gpt-4o", "messages": semantic_messages("name a primary color")} diff --git a/tests/test_litellm_rust/cache/test_rollout.py b/tests/test_litellm_rust/cache/test_rollout.py index 7f33e31599f..a290857c8bb 100644 --- a/tests/test_litellm_rust/cache/test_rollout.py +++ b/tests/test_litellm_rust/cache/test_rollout.py @@ -1,7 +1,6 @@ import asyncio from collections.abc import Callable from pathlib import Path -from types import SimpleNamespace from typing import Final, TypeAlias, cast from urllib.parse import urlparse from uuid import uuid4 @@ -9,10 +8,11 @@ from uuid import uuid4 import pytest from litellm.caching.caching import Cache -from litellm.rust_bridge.response_cache import NativeResponseCacheRuntime, ResponseCacheRuntime, resolve_response_cache +from litellm.rust_bridge import _native +from litellm.rust_bridge.response_cache import NativeResponseCacheRuntime, ResponseCacheRuntime from litellm.types.caching import LiteLLMCacheType from litellm.types.utils import EmbeddingResponse -from tests.test_litellm_rust.support.cache import assert_native_runtime, completion_kwargs, require_rust +from tests.test_litellm_rust.support.cache import activate_native, assert_native_runtime, completion_kwargs from tests.test_litellm_rust.support.s3_stub import S3Stub pytestmark: Final = pytest.mark.requires_rust_extension @@ -69,11 +69,6 @@ ROUND_TRIP_BACKENDS: Final = ( SHARED_STORE_BACKENDS: Final = (LiteLLMCacheType.DISK, LiteLLMCacheType.REDIS, LiteLLMCacheType.S3) -@pytest.mark.parametrize("backend", list(LiteLLMCacheType)) -def test_shipped_rules_keep_every_backend_on_python(backend: LiteLLMCacheType) -> None: - assert resolve_response_cache(cast(Cache, SimpleNamespace(type=backend))) is None - - @pytest.mark.parametrize( "cache_factory", [ @@ -87,7 +82,7 @@ def test_shipped_rules_keep_every_backend_on_python(backend: LiteLLMCacheType) - ], indirect=True, ) -def test_shipped_rules_construct_python_backed_facades(cache_factory: CacheFactory) -> None: +def test_legacy_constructor_keeps_python_backends(cache_factory: CacheFactory) -> None: assert cache_factory()._native_cache is None # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor @@ -104,19 +99,17 @@ def test_shipped_rules_construct_python_backed_facades(cache_factory: CacheFacto ], indirect=True, ) -def test_rust_required_rule_activates_the_native_backend( +def test_explicit_selection_activates_the_native_backend( cache_factory: CacheFactory, monkeypatch: pytest.MonkeyPatch, request: pytest.FixtureRequest ) -> None: - require_rust(monkeypatch, cast(LiteLLMCacheType, request.node.callspec.params["cache_factory"])) - assert_native_runtime(cache_factory()) + assert_native_runtime(activate_native(cache_factory())) @pytest.mark.parametrize("cache_factory", ROUND_TRIP_BACKENDS, indirect=True) async def test_facade_storage_calls_round_trip_through_the_native_backend( cache_factory: CacheFactory, monkeypatch: pytest.MonkeyPatch, request: pytest.FixtureRequest ) -> None: - require_rust(monkeypatch, cast(LiteLLMCacheType, request.node.callspec.params["cache_factory"])) - facade: Final = cache_factory() + facade: Final = activate_native(cache_factory()) assert_native_runtime(facade) sync_kwargs: Final = completion_kwargs("sync") @@ -130,8 +123,7 @@ async def test_facade_storage_calls_round_trip_through_the_native_backend( async def test_memory_facade_writes_bypass_the_python_backend(monkeypatch: pytest.MonkeyPatch) -> None: - require_rust(monkeypatch, LiteLLMCacheType.LOCAL) - facade: Final = Cache(type=LiteLLMCacheType.LOCAL) + facade: Final = activate_native(Cache(type=LiteLLMCacheType.LOCAL)) assert_native_runtime(facade) kwargs: Final = completion_kwargs("memory") facade.add_cache({"answer": 1}, **kwargs) @@ -145,8 +137,7 @@ async def test_native_and_python_facades_share_one_wire_format( ) -> None: python_facade: Final = cache_factory() assert python_facade._native_cache is None # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor - require_rust(monkeypatch, cast(LiteLLMCacheType, request.node.callspec.params["cache_factory"])) - native_facade: Final = cache_factory() + native_facade: Final = activate_native(cache_factory()) assert_native_runtime(native_facade) native_written: Final = completion_kwargs("native") @@ -170,8 +161,7 @@ async def test_native_and_python_facades_share_one_wire_format( async def test_embedding_pipeline_stores_one_native_entry_per_input( cache_factory: CacheFactory, monkeypatch: pytest.MonkeyPatch, request: pytest.FixtureRequest ) -> None: - require_rust(monkeypatch, cast(LiteLLMCacheType, request.node.callspec.params["cache_factory"])) - facade: Final = cache_factory() + facade: Final = activate_native(cache_factory()) assert_native_runtime(facade) inputs: Final = [f"alpha {uuid4().hex}", f"beta {uuid4().hex}"] result: Final = EmbeddingResponse( @@ -224,9 +214,8 @@ async def test_embedding_pipeline_stores_one_native_entry_per_input( def test_semantic_settings_the_native_client_cannot_honor_decline( monkeypatch: pytest.MonkeyPatch, backend: LiteLLMCacheType, settings: dict[str, object], message: str ) -> None: - require_rust(monkeypatch, backend) - with pytest.raises(RuntimeError, match=f"declined the cache: {message}"): - Cache(type=backend, **settings) + with pytest.raises(_native.RustBridgeDeclined, match=message): + activate_native(Cache(type=backend, **settings)) class _SemanticHit: diff --git a/tests/test_litellm_rust/cache/test_s3.py b/tests/test_litellm_rust/cache/test_s3.py index 044bfc39f8d..d7b36ad1b2e 100644 --- a/tests/test_litellm_rust/cache/test_s3.py +++ b/tests/test_litellm_rust/cache/test_s3.py @@ -11,8 +11,9 @@ import pytest from litellm.caching.caching import Cache from litellm.caching.s3_cache import S3Cache +from litellm.rust_bridge import _native from litellm.types.caching import LiteLLMCacheType -from tests.test_litellm_rust.support.cache import CacheTestHandle, CacheTestResolver, request +from tests.test_litellm_rust.support.cache import CacheTestResolver, activate_native, native_runtime, request from tests.test_litellm_rust.support.isolation import rebound from tests.test_litellm_rust.support.s3_stub import S3Stub @@ -30,6 +31,18 @@ def python_s3(url: str) -> S3Cache: ) +def s3_facade(url: str) -> Cache: + return Cache( + type=LiteLLMCacheType.S3, + 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}} @@ -41,18 +54,7 @@ async def test_s3_reads_python_entries_and_writes_with_python_metadata(s3_stub: json.dumps({"timestamp": time.time(), "response": response}).encode(), {"expires": "Thu, 01 Jan 1970 00:00:00 GMT"}, ) - binding: Final = CacheTestResolver( - SimpleNamespace( - cache=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() + binding: Final = native_runtime(s3_facade(s3_stub.url)) assert binding.lookup(request("sync:key")) == response assert await binding.async_lookup(request("plain")) == response @@ -79,7 +81,7 @@ async def test_s3_reads_python_entries_and_writes_with_python_metadata(s3_stub: 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: +def test_selected_s3_runtime_declines_backend_mutation(s3_stub: S3Stub) -> None: facade: Final = Cache( type=LiteLLMCacheType.S3, s3_bucket_name="cache-bucket", @@ -89,21 +91,7 @@ def test_s3_facade_binds_only_exact_configuration_and_falls_back_on_mutation(s3_ s3_aws_secret_access_key="secret", s3_path="team", ) - handle: Final = 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"): - CacheTestHandle.s3("other", region="us-east-1", endpoint_url=s3_stub.url)._bind_facade(facade) - with pytest.raises(TypeError, match="key prefixes must match"): - CacheTestHandle.s3( - "cache-bucket", region="us-east-1", endpoint_url=s3_stub.url, key_prefix="other/" - )._bind_facade(facade) - handle._bind_facade(facade) + activate_native(facade) resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade)) binding: Final = resolver.resolve() assert binding.kind == "native" @@ -116,7 +104,8 @@ def test_s3_facade_binds_only_exact_configuration_and_falls_back_on_mutation(s3_ assert "team/native" in s3_stub.objects with rebound(facade.cache, "bucket_name", "other"): - assert resolver.resolve().kind == "python_callback" + with pytest.raises(_native.RustBridgeDeclined): + resolver.resolve() other_client: Final = boto3.client( "s3", region_name="us-east-1", @@ -125,7 +114,8 @@ def test_s3_facade_binds_only_exact_configuration_and_falls_back_on_mutation(s3_ aws_secret_access_key="secret", ) with rebound(facade.cache, "s3_client", other_client): - assert resolver.resolve().kind == "python_callback" + with pytest.raises(_native.RustBridgeDeclined): + resolver.resolve() class CustomS3Cache(S3Cache): pass @@ -147,20 +137,10 @@ def test_s3_facade_binds_only_exact_configuration_and_falls_back_on_mutation(s3_ s3_aws_secret_access_key="secret", s3_path="team", ) - with pytest.raises(TypeError): - handle._bind_facade(subclassed) assert CacheTestResolver(SimpleNamespace(cache=subclassed)).resolve().kind == "python_callback" def test_s3_facade_rejects_configurations_that_require_python(s3_stub: S3Stub) -> None: - handle: Final = 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", @@ -171,8 +151,8 @@ def test_s3_facade_rejects_configurations_that_require_python(s3_stub: S3Stub) - s3_path="team", s3_verify=False, ) - with pytest.raises(TypeError, match="requires Python"): - handle._bind_facade(unverified) + with pytest.raises(_native.RustBridgeDeclined, match="requires Python"): + native_runtime(unverified) proxied: Final = Cache( type=LiteLLMCacheType.S3, s3_bucket_name="cache-bucket", @@ -183,5 +163,5 @@ def test_s3_facade_rejects_configurations_that_require_python(s3_stub: S3Stub) - s3_path="team", s3_config=botocore.config.Config(proxies={"https": "http://proxy.test"}), ) - with pytest.raises(TypeError, match="requires Python"): - handle._bind_facade(proxied) + with pytest.raises(_native.RustBridgeDeclined, match="requires Python"): + native_runtime(proxied) diff --git a/tests/test_litellm_rust/cache/test_v2.py b/tests/test_litellm_rust/cache/test_v2.py new file mode 100644 index 00000000000..d086a02c600 --- /dev/null +++ b/tests/test_litellm_rust/cache/test_v2.py @@ -0,0 +1,794 @@ +import asyncio +from collections.abc import AsyncIterator, Mapping +from types import MappingProxyType +from typing import Final, Literal + +import pytest +from pydantic import BaseModel, TypeAdapter + +import litellm +from litellm import _v2 +from litellm._v2.cache import NativeBackend +from litellm.caching.caching import Cache, CacheMode +from litellm.caching.caching_handler import ( + _PENDING_CACHE_WRITES, # pyright: ignore[reportPrivateUsage] # await the existing background cache writer before the next request +) +from litellm.proxy._types import Litellm_EntityType, UserAPIKeyAuth +from litellm.proxy.hooks.model_max_budget_limiter import ( + _PROXY_VirtualKeyModelMaxBudgetLimiter, + model_budget_spend_cache_key, +) +from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 +from litellm.proxy.utils import InternalUsageCache +from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict +from litellm.rust_bridge import runtime +from litellm.rust_bridge.catalog import Route, RouteContext, RouteRule +from litellm.rust_bridge.chat_completions.entrypoints import NATIVE_ACOMPLETION, LiteLLMChatCompletionsRequest +from litellm.rust_bridge.configuration import Rollout +from litellm.rust_bridge.dispatch import call_hook +from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES, LiteLLMMessagesRequest +from litellm.rust_bridge.responses.entrypoints import NATIVE_ARESPONSES, LiteLLMResponsesRequest +from litellm.types.caching import CachingSupportedCallTypes +from litellm.types.utils import ModelResponse +from tests.test_litellm_rust.support.callback_recorder import RecordingLogger, drain_logging +from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec +from tests.test_litellm_rust.support.requests import MESSAGES, MESSAGES_EVENTS, MESSAGES_MODEL, MESSAGES_RESPONSE +from tests.test_litellm_rust.test_inference import RESPONSES_MODEL, RESPONSES_RESPONSE + +pytestmark = pytest.mark.requires_rust_extension + + +def payload(value: object) -> object: + if isinstance(value, ModelResponse): + return value.model_dump_json(exclude=MappingProxyType({"id": True, "created": True})) + if isinstance(value, dict): + fields: Final = TypeAdapter(dict[str, object]).validate_python(value) + return {name: field for name, field in fields.items() if name != "_hidden_params"} + return value.model_dump_json() if isinstance(value, BaseModel) else value + + +def cache_key(response: object) -> object: + hidden: Final = get_hidden_params_dict(response) + headers: Final = TypeAdapter(dict[str, object]).validate_python(hidden.get("additional_headers", {})) + return headers.get("x-litellm-cache-key") + + +async def invoke( + route: Literal["chat", "messages", "responses"], + server: RecordingServer, + options: Mapping[str, object], + native: bool = True, +) -> object: + common: Final = {"api_key": "test-key", "api_base": server.base_url, **options} + if route == "responses": + server.default_response = ResponseSpec(body=RESPONSES_RESPONSE) + arguments: Final = {"model": RESPONSES_MODEL, "input": "hello", **common} + if not native: + return await litellm.aresponses(**arguments) + request: Final = LiteLLMResponsesRequest( + RESPONSES_MODEL, "hello", None, "test-key", server.base_url, "openai", None, arguments + ) + return await runtime.arun( + RouteContext(Route.RESPONSES), + binding=NATIVE_ARESPONSES, + native=lambda hook: call_hook(hook, request, (), arguments), + python=runtime.NO_PYTHON, + rules=(RouteRule(Route.RESPONSES, Rollout.RUST_REQUIRED),), + ) + server.default_response = ( + ResponseSpec(body=None, events=MESSAGES_EVENTS) + if options.get("stream") + else ResponseSpec(body=MESSAGES_RESPONSE) + ) + parameters: Final = {"model": MESSAGES_MODEL, "messages": list(MESSAGES), "max_tokens": 32, **common} + if route == "chat": + if not native: + return await litellm.acompletion(**parameters) + chat: Final = LiteLLMChatCompletionsRequest( + MESSAGES_MODEL, list(MESSAGES), None, "test-key", server.base_url, None, None, parameters + ) + return await runtime.arun( + RouteContext(Route.CHAT_COMPLETIONS), + binding=NATIVE_ACOMPLETION, + native=lambda hook: call_hook(hook, chat, (), parameters), + python=runtime.NO_PYTHON, + rules=(RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),), + ) + if not native: + return await litellm.anthropic_messages(**parameters) + messages: Final = LiteLLMMessagesRequest( + MESSAGES_MODEL, list(MESSAGES), 32, None, "test-key", server.base_url, "anthropic", parameters + ) + return await runtime.arun( + RouteContext(Route.MESSAGES), + binding=NATIVE_AMESSAGES, + native=lambda hook: call_hook(hook, messages, (), parameters), + python=runtime.NO_PYTHON, + rules=(RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),), + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("route", ("chat", "messages", "responses")) +@pytest.mark.parametrize("backend", ("memory", "redis")) +async def test_v2_cache_skips_provider_and_reports_one_success_per_call( + recording_server: RecordingServer, + route: Literal["chat", "messages", "responses"], + backend: Literal["memory", "redis"], + redis_url: str, +) -> None: + recording_server.expected_requests = 2 + litellm.cache = _v2.Cache.memory() if backend == "memory" else _v2.Cache.redis(redis_url, namespace="headers") + recorder: Final = RecordingLogger() + first: Final = await invoke(route, recording_server, {"callbacks": [recorder]}) + await recorder.wait_for_async("async_log_success_event") + second: Final = await invoke(route, recording_server, {"callbacks": [recorder]}) + assert payload(first) == payload(second) + assert cache_key(first) is None + key: Final = cache_key(second) + assert isinstance(key, str) + assert key == get_hidden_params_dict(second)["cache_key"] + assert len(recording_server.requests) == 1 + await drain_logging() + successes: Final = await recorder.wait_for_async("async_log_success_event", count=2) + assert len(successes) == 2 + cached_log: Final = TypeAdapter(dict[str, object]).validate_python(successes[-1].kwargs) + assert cached_log["cache_hit"] is True + assert cached_log["response_cost"] == 0 + await litellm.cache.delete_cache_keys([key]) + refreshed: Final = await invoke(route, recording_server, {"callbacks": [recorder]}) + assert cache_key(refreshed) is None + assert len(recording_server.requests) == 2 + assert len(await recorder.wait_for_async("async_log_success_event", count=3)) == 3 + await litellm.cache.disconnect() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("route", "stream", "legacy"), + ( + ("chat", False, False), + ("messages", False, False), + ("responses", False, False), + ("messages", True, False), + ("messages", False, True), + ("messages", True, True), + ), +) +@pytest.mark.parametrize("native", (False, True), ids=("python", "rust")) +async def test_cache_hit_keeps_model_budget_spend_but_accounts_for_usage( + recording_server: RecordingServer, + route: Literal["chat", "messages", "responses"], + stream: bool, + native: bool, + monkeypatch: pytest.MonkeyPatch, + legacy: bool, +) -> None: + monkeypatch.setenv("LITELLM_RUST", "1" if native else "0") + litellm.cache = Cache() if legacy else _v2.Cache.memory() + counters: Final = litellm.DualCache() + budget: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(counters) + limiter: Final = _PROXY_MaxParallelRequestsHandler_v3( + InternalUsageCache(counters), model_group_resolver=lambda model: model + ) + recorder: Final = RecordingLogger() + key_hash: Final = "a" * 64 + metadata: Final = { + "user_api_key": key_hash, + "model_group": "cached-model", + "user_api_key_model_max_budget": {"cached-model": {"max_budget": 1, "budget_duration": "1h"}}, + } + options: Final = { + "callbacks": [budget, limiter, recorder], + "metadata": metadata, + "stream": stream, + } + spend_key: Final = model_budget_spend_cache_key(Litellm_EntityType.KEY, key_hash, "cached-model", "1h") + token_key: Final = limiter.create_rate_limit_keys("api_key", key_hash, "tokens") + first: Final = await invoke(route, recording_server, options, native=native) + if stream: + await collect(first) + await asyncio.gather(*tuple(_PENDING_CACHE_WRITES)) + await drain_logging() + first_events: Final = await recorder.wait_for_async("async_log_success_event") + first_log: Final = TypeAdapter(dict[str, object]).validate_python(first_events[0].kwargs) + first_payload: Final = TypeAdapter(dict[str, object]).validate_python(first_log["standard_logging_object"]) + expected_cost: Final = TypeAdapter(float).validate_python(first_log["response_cost"]) + usage: Final = RESPONSES_RESPONSE["usage"] if route == "responses" else MESSAGES_RESPONSE["usage"] + expected_tokens: Final = usage["input_tokens"] + usage["output_tokens"] + assert expected_cost > 0 + assert counters.get_cache(spend_key) == pytest.approx(expected_cost) + assert counters.get_cache(token_key) == first_payload["total_tokens"] == expected_tokens + + second: Final = await invoke(route, recording_server, options, native=native) + if stream: + await collect(second) + await drain_logging() + successes: Final = await recorder.wait_for_async("async_log_success_event", count=2) + cached_log: Final = TypeAdapter(dict[str, object]).validate_python(successes[-1].kwargs) + cached_payload: Final = TypeAdapter(dict[str, object]).validate_python(cached_log["standard_logging_object"]) + assert len(recording_server.requests) == 1 + assert len(successes) == 2 + assert cached_log["cache_hit"] is True + assert cached_log["response_cost"] == cached_payload["response_cost"] == 0 + assert cached_payload["cache_hit"] is True + assert cached_payload["id"] != first_payload["id"] + assert cached_payload["custom_llm_provider"] == first_payload["custom_llm_provider"] + assert cached_payload["custom_llm_provider"] == ("openai" if route == "responses" else "anthropic"), { + "miss_provider": first_log.get("custom_llm_provider"), + "hit_provider": cached_log.get("custom_llm_provider"), + } + assert cached_payload["total_tokens"] == first_payload["total_tokens"] + assert counters.get_cache(spend_key) == pytest.approx(expected_cost) + assert counters.get_cache(token_key) == 2 * expected_tokens + + +@pytest.mark.asyncio +@pytest.mark.parametrize("native", (False, True), ids=("python", "rust")) +@pytest.mark.parametrize("backend", ("disabled", "memory", "redis")) +async def test_response_cache_backend_does_not_control_coordination( + recording_server: RecordingServer, + native: bool, + monkeypatch: pytest.MonkeyPatch, + backend: Literal["disabled", "memory", "redis"], + redis_url: str, +) -> None: + monkeypatch.setenv("LITELLM_RUST", "1" if native else "0") + litellm.cache = ( + None + if backend == "disabled" + else _v2.Cache.memory() + if backend == "memory" + else _v2.Cache.redis(redis_url, namespace="independent-coordination") + ) + recording_server.expected_requests = 2 if backend == "disabled" else 1 + counters: Final = litellm.DualCache() + budget: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(counters) + limiter: Final = _PROXY_MaxParallelRequestsHandler_v3( + InternalUsageCache(counters), model_group_resolver=lambda model: model + ) + key_hash: Final = "b" * 64 + identity: Final = UserAPIKeyAuth(api_key=key_hash, rpm_limit=2, tpm_limit=1000, max_parallel_requests=1) + spend_key: Final = model_budget_spend_cache_key(Litellm_EntityType.KEY, key_hash, "cached-model", "1h") + request_key: Final = limiter.create_rate_limit_keys("api_key", key_hash, "requests") + token_key: Final = limiter.create_rate_limit_keys("api_key", key_hash, "tokens") + parallel_key: Final = limiter.create_rate_limit_keys("api_key", key_hash, "max_parallel_requests") + expected_tokens: Final = MESSAGES_RESPONSE["usage"]["input_tokens"] + MESSAGES_RESPONSE["usage"]["output_tokens"] + recorder: Final = RecordingLogger() + + async def request(call_id: str, successes: int) -> object: + data: Final = { + "model": MESSAGES_MODEL, + "messages": list(MESSAGES), + "litellm_call_id": call_id, + "max_tokens": 32, + "metadata": { + "user_api_key": key_hash, + "model_group": "cached-model", + "user_api_key_model_max_budget": {"cached-model": {"max_budget": 1, "budget_duration": "1h"}}, + }, + } + await limiter.async_pre_call_hook(identity, counters, data, "acompletion") + assert len(TypeAdapter(dict[str, float]).validate_python(counters.get_cache(parallel_key))) == 1 + response: Final = await invoke( + "chat", recording_server, {**data, "callbacks": [budget, limiter, recorder]}, native=native + ) + await asyncio.gather(*tuple(_PENDING_CACHE_WRITES)) + await recorder.wait_for_async("async_log_success_event", count=successes) + return response + + await asyncio.create_task(request("cache-miss", 1)) + assert counters.get_cache(parallel_key) == {} + assert counters.get_cache(request_key) == 1 + assert counters.get_cache(token_key) == expected_tokens + first_events: Final = await recorder.wait_for_async("async_log_success_event") + first_cost: Final = TypeAdapter(float).validate_python(first_events[0].kwargs["response_cost"]) + assert first_cost > 0 + assert counters.get_cache(spend_key) == pytest.approx(first_cost) + await asyncio.create_task(request("cache-hit", 2)) + expected_spend: Final = first_cost * recording_server.expected_requests + assert counters.get_cache(spend_key) == pytest.approx(expected_spend) + assert len(recording_server.requests) == recording_server.expected_requests + assert counters.get_cache(parallel_key) == {} + assert counters.get_cache(request_key) == 2 + assert counters.get_cache(token_key) == 2 * expected_tokens + with pytest.raises(litellm.RateLimitError): + await asyncio.create_task(request("over-rpm-limit", 3)) + assert len(recording_server.requests) == recording_server.expected_requests + assert counters.get_cache(parallel_key) == {} + assert counters.get_cache(token_key) == 2 * expected_tokens + + assert counters.get_cache(spend_key) == pytest.approx(expected_spend) + if litellm.cache is not None: + await litellm.cache.disconnect() + + +@pytest.mark.asyncio +async def test_v2_global_cache_leaves_legacy_only_calls_usable() -> None: + litellm.cache = _v2.Cache.memory() + response: Final = await litellm.aembedding( + model="openai/cache-test-embedding", + input=["hello"], + api_key="test-key", + mock_response=[0.25, 0.75], + ) + assert response.model_dump(include={"data"}) == { + "data": [{"embedding": [0.25, 0.75], "index": 0, "object": "embedding"}] + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("route", "legacy"), (("chat", False), ("messages", False), ("responses", False), ("messages", True)) +) +async def test_cache_controls_and_backend_credential_key_semantics( + recording_server: RecordingServer, + route: Literal["chat", "messages", "responses"], + legacy: bool, +) -> None: + recording_server.expected_requests = 3 if legacy else 4 + litellm.cache = Cache() if legacy else _v2.Cache.memory() + await invoke(route, recording_server, {"cache": {"no-store": True}}) + await invoke(route, recording_server, {}) + await invoke(route, recording_server, {}) + assert len(recording_server.requests) == 2 + await invoke(route, recording_server, {"cache": {"no-cache": True}}) + await invoke(route, recording_server, {"api_key": "another-key"}) + assert len(recording_server.requests) == recording_server.expected_requests + + +async def collect(stream: object) -> bytes: + assert isinstance(stream, AsyncIterator) + return b"".join([chunk_bytes(chunk) async for chunk in stream]) + + +def chunk_bytes(value: object) -> bytes: + assert isinstance(value, bytes) + return value + + +@pytest.mark.asyncio +@pytest.mark.parametrize("legacy", (False, True)) +async def test_v2_messages_replays_a_completed_stream(recording_server: RecordingServer, legacy: bool) -> None: + recording_server.default_response = ResponseSpec(body=None, events=MESSAGES_EVENTS) + litellm.cache = Cache() if legacy else _v2.Cache.memory() + recorder: Final = RecordingLogger() + parameters: Final = { + "model": MESSAGES_MODEL, + "messages": list(MESSAGES), + "max_tokens": 32, + "api_key": "test-key", + "api_base": recording_server.base_url, + "stream": True, + "callbacks": [recorder], + } + first_stream: Final = await litellm.anthropic_messages(**parameters) + assert cache_key(first_stream) is None + first: Final = await collect(first_stream) + await recorder.wait_for_async("async_log_success_event") + second_stream: Final = await litellm.anthropic_messages(**parameters) + assert isinstance(cache_key(second_stream), str) + assert cache_key(second_stream) == get_hidden_params_dict(second_stream)["cache_key"] + second: Final = await collect(second_stream) + assert payload(first) == payload(second) + assert first == b"".join(recording_server.default_response.payloads()) + assert len(recording_server.requests) == 1 + await drain_logging() + successes: Final = await recorder.wait_for_async("async_log_success_event", count=2) + cached_log: Final = TypeAdapter(dict[str, object]).validate_python(successes[-1].kwargs) + assert cached_log["cache_hit"] is True + assert cached_log["response_cost"] == 0 + + +@pytest.mark.parametrize("route", ("chat", "responses")) +def test_v2_cache_works_through_python_inference( + recording_server: RecordingServer, route: Literal["chat", "messages", "responses"], monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_RUST", "0") + litellm.cache = _v2.Cache.memory() + common: Final = {"api_key": "test-key", "api_base": recording_server.base_url} + if route == "responses": + recording_server.default_response = ResponseSpec(body=RESPONSES_RESPONSE) + parameters: Final = {"model": RESPONSES_MODEL, "input": "hello", **common} + first: Final = litellm.responses(**parameters) + second: Final = litellm.responses(**parameters) + assert payload(first) == payload(second) + else: + recording_server.default_response = ResponseSpec(body=MESSAGES_RESPONSE) + arguments: Final = {"model": MESSAGES_MODEL, "messages": list(MESSAGES), "max_tokens": 32, **common} + initial: Final = litellm.completion(**arguments) + cached: Final = litellm.completion(**arguments) + assert isinstance(initial, ModelResponse) and isinstance(cached, ModelResponse) + assert ( + initial.choices[0].message.content + == cached.choices[0].message.content + == MESSAGES_RESPONSE["content"][0]["text"] + ) + assert len(recording_server.requests) == 1 + + +@pytest.mark.asyncio +async def test_v2_facade_and_backend_share_storage_and_management() -> None: + cache: Final = _v2.Cache.memory() + await cache.async_add_cache({"answer": 7}, cache_key="shared") + assert cache.get_cache(cache_key="shared") == {"answer": 7} + assert await cache.ping() is True + await cache.delete_cache_keys(["shared"]) + assert await cache.async_get_cache(cache_key="shared") is None + cache.add_cache({"answer": 8}, cache_key="flush") + backend: Final = cache.cache + assert isinstance(backend, NativeBackend) + backend.flush_cache() + assert cache.get_cache(cache_key="flush") is None + await cache.disconnect() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("control", ("s-maxage", "s-max-age")) +async def test_v2_native_cache_accepts_existing_freshness_aliases( + recording_server: RecordingServer, control: str +) -> None: + litellm.cache = _v2.Cache.memory() + first: Final = await invoke("responses", recording_server, {}) + second: Final = await invoke("responses", recording_server, {"cache": {control: 600}}) + assert payload(first) == payload(second) + assert len(recording_server.requests) == 1 + + +@pytest.mark.asyncio +async def test_v2_cache_does_not_force_native_responses_streaming( + recording_server: RecordingServer, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_RUST", "0") + litellm.cache = _v2.Cache.memory() + recording_server.default_response = ResponseSpec( + body=None, + events=( + ("response.created", {"type": "response.created", "sequence_number": 0, "response": RESPONSES_RESPONSE}), + ( + "response.completed", + {"type": "response.completed", "sequence_number": 1, "response": RESPONSES_RESPONSE}, + ), + ), + ) + response: Final = await litellm.aresponses( + model=RESPONSES_MODEL, + input="hello", + stream=True, + caching=False, + api_key="test-key", + api_base=recording_server.base_url, + ) + assert isinstance(response, AsyncIterator) + chunks: Final = [chunk async for chunk in response] + assert chunks[-1].type == "response.completed" + assert chunks[-1].response.output[0].content[0].text == "native response" + + +@pytest.mark.asyncio +async def test_v2_cache_works_through_python_messages( + recording_server: RecordingServer, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_RUST", "0") + litellm.cache = _v2.Cache.memory() + recording_server.default_response = ResponseSpec(body=MESSAGES_RESPONSE) + parameters: Final = { + "model": MESSAGES_MODEL, + "messages": list(MESSAGES), + "max_tokens": 32, + "api_key": "test-key", + "api_base": recording_server.base_url, + } + first: Final = await litellm.anthropic_messages(**parameters) + await asyncio.gather(*tuple(_PENDING_CACHE_WRITES)) + second: Final = await litellm.anthropic_messages(**parameters) + assert payload(first) == payload(second) + assert len(recording_server.requests) == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("backend", ("memory", "redis")) +async def test_rust_messages_uses_a_legacy_cache_without_python_inference( + recording_server: RecordingServer, + monkeypatch: pytest.MonkeyPatch, + backend: Literal["memory", "redis"], + redis_url: str, +) -> None: + from litellm.caching.caching import Cache + + monkeypatch.setenv("LITELLM_RUST", "1") + litellm.cache = Cache() if backend == "memory" else Cache(type="redis", url=redis_url, namespace="rust-host") + logger: Final = RecordingLogger() + litellm.callbacks = [logger] + recording_server.default_response = ResponseSpec(body=MESSAGES_RESPONSE) + parameters: Final = { + "model": MESSAGES_MODEL, + "messages": list(MESSAGES), + "max_tokens": 32, + "api_key": "test-key", + "api_base": recording_server.base_url, + } + first: Final = await invoke("messages", recording_server, parameters) + second: Final = await invoke("messages", recording_server, parameters) + assert cache_key(second) + assert cache_key(first) is None + assert payload(first) == payload(second) + assert len(recording_server.requests) == 1 + + await logger.wait_for_async("async_log_success_event", count=2) + assert logger.names.count("async_log_success_event") == 2 + assert "log_failure_event" not in logger.names + assert "async_log_failure_event" not in logger.names + + +@pytest.mark.asyncio +@pytest.mark.parametrize("route", ("chat", "messages", "responses")) +@pytest.mark.parametrize("native", (False, True)) +@pytest.mark.parametrize("excluded", (None, [], ["embedding"])) +async def test_v2_cache_honors_supported_call_types_for_reads_and_writes( + recording_server: RecordingServer, + monkeypatch: pytest.MonkeyPatch, + route: Literal["chat", "messages", "responses"], + native: bool, + excluded: list[CachingSupportedCallTypes] | None, +) -> None: + monkeypatch.setenv("LITELLM_RUST", "1" if native else "0") + litellm.cache = _v2.Cache.memory() + call_type: Final[CachingSupportedCallTypes] = ( + "acompletion" if route == "chat" else "anthropic_messages" if route == "messages" else "aresponses" + ) + recording_server.expected_requests = 4 + litellm.cache.supported_call_types = excluded + await invoke(route, recording_server, {}, native=native) + await invoke(route, recording_server, {}, native=native) + assert len(recording_server.requests) == 2 + litellm.cache.supported_call_types = [call_type] + await invoke(route, recording_server, {}, native=native) + assert len(recording_server.requests) == 3 + await asyncio.gather(*tuple(_PENDING_CACHE_WRITES)) + await invoke(route, recording_server, {}, native=native) + assert len(recording_server.requests) == 3 + litellm.cache.supported_call_types = excluded + await invoke(route, recording_server, {}, native=native) + assert len(recording_server.requests) == 4 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_v2_redis_flush_only_removes_its_namespace( + redis_url: str, recording_server: RecordingServer, asynchronous: bool +) -> None: + own: Final = _v2.Cache.redis(redis_url, namespace="flush-own") + other: Final = _v2.Cache.redis(redis_url, namespace="flush-other") + litellm.cache = own + recording_server.expected_requests = 2 + await own.async_add_cache({"answer": "own"}, cache_key="shared") + await other.async_add_cache({"answer": "other"}, cache_key="shared") + await invoke("responses", recording_server, {}) + hit: Final = await invoke("responses", recording_server, {}) + assert isinstance(cache_key(hit), str) + assert await own.async_get_cache(cache_key="shared") == {"answer": "own"} + backend: Final = own.cache + assert isinstance(backend, NativeBackend) + if asynchronous: + await backend.async_flush_cache() + else: + backend.flush_cache() + assert await own.async_get_cache(cache_key="shared") is None + assert await other.async_get_cache(cache_key="shared") == {"answer": "other"} + refreshed: Final = await invoke("responses", recording_server, {}) + assert cache_key(refreshed) is None + assert len(recording_server.requests) == 2 + await own.disconnect() + await other.disconnect() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("native", (False, True)) +@pytest.mark.parametrize( + ("route", "legacy"), (("chat", False), ("messages", False), ("responses", False), ("messages", True)) +) +async def test_v2_default_off_requires_opt_in_even_for_existing_entries( + recording_server: RecordingServer, + monkeypatch: pytest.MonkeyPatch, + route: Literal["chat", "messages", "responses"], + native: bool, + legacy: bool, +) -> None: + monkeypatch.setenv("LITELLM_RUST", "1" if native else "0") + litellm.cache = Cache() if legacy else _v2.Cache.memory() + litellm.cache.mode = CacheMode.default_off + recording_server.expected_requests = 4 + await invoke(route, recording_server, {}, native=native) + await invoke(route, recording_server, {}, native=native) + assert len(recording_server.requests) == 2 + await invoke(route, recording_server, {"cache": {"use-cache": True}}, native=native) + assert len(recording_server.requests) == 3 + await asyncio.gather(*tuple(_PENDING_CACHE_WRITES)) + await invoke(route, recording_server, {"cache": {"use-cache": True}}, native=native) + assert len(recording_server.requests) == 3 + await invoke(route, recording_server, {}, native=native) + assert len(recording_server.requests) == 4 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("route", "legacy"), (("chat", False), ("messages", False), ("responses", False), ("messages", True)) +) +async def test_cache_lookup_uses_backend_request_callback_semantics( + recording_server: RecordingServer, + route: Literal["chat", "messages", "responses"], + legacy: bool, +) -> None: + from tests.test_litellm_rust.support.requests import request_body + + class Rewrite(RecordingLogger): + temperature = 0.1 + + def log_pre_api_call(self, model: str, messages: object, kwargs: dict[str, object]) -> None: + request_body(kwargs)["temperature"] = self.temperature + super().log_pre_api_call(model, messages, kwargs) + + logger: Final = Rewrite() + litellm.cache = Cache() if legacy else _v2.Cache.memory() + recording_server.expected_requests = 1 if legacy else 2 + await invoke(route, recording_server, {"callbacks": [logger]}) + first_hit: Final = await invoke(route, recording_server, {"callbacks": [logger]}) + logger.temperature = 0.8 + await invoke(route, recording_server, {"callbacks": [logger]}) + second_hit: Final = await invoke(route, recording_server, {"callbacks": [logger]}) + assert logger.names.count("log_pre_api_call") == 4 + assert len(recording_server.requests) == recording_server.expected_requests + assert recording_server.requests[0].body["temperature"] == 0.1 + if not legacy: + assert recording_server.requests[1].body["temperature"] == 0.8 + assert isinstance(cache_key(first_hit), str) + assert isinstance(cache_key(second_hit), str) + if not legacy: + assert cache_key(first_hit) != cache_key(second_hit) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("cancel_lookup", (False, True)) +async def test_python_cache_operations_stay_in_the_rust_callers_task( + recording_server: RecordingServer, + cancel_lookup: bool, +) -> None: + from litellm.caching.base_cache import BaseCache + from litellm.caching.in_memory_cache import InMemoryCache + + caller: Final = asyncio.current_task() + entered: Final = asyncio.Event() + release: Final = asyncio.Event() + storage: Final = InMemoryCache() + + class CallerCache(BaseCache): + async def async_set_cache_pipeline( + self, cache_list: list[tuple[str, object]], ttl: float | None = None + ) -> None: + await storage.async_set_cache_pipeline(cache_list, ttl=ttl) + + async def async_get_cache(self, key: str, **kwargs: object) -> object: + if cancel_lookup: + entered.set() + await release.wait() + else: + assert asyncio.current_task() is caller + return storage.get_cache(key, **kwargs) + + async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None: + assert asyncio.current_task() is caller + await asyncio.sleep(0) + storage.set_cache(key, value, **kwargs) + + litellm.cache = Cache(_backend=CallerCache()) + if cancel_lookup: + recording_server.expected_requests = 0 + task: Final = asyncio.create_task(invoke("messages", recording_server, {})) + await asyncio.wait_for(entered.wait(), timeout=5) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + release.set() + await asyncio.sleep(0) + assert len(recording_server.requests) == 0 + assert storage.cache_dict == {} + return + first: Final = await invoke("messages", recording_server, {}) + second: Final = await invoke("messages", recording_server, {}) + assert payload(first) == payload(second) + assert cache_key(second) + assert len(recording_server.requests) == 1 + + +def test_sync_rust_messages_calls_python_cache(recording_server: RecordingServer) -> None: + from litellm.rust_bridge.messages.entrypoints import NATIVE_MESSAGES + + litellm.cache = Cache() + recording_server.default_response = ResponseSpec(body=MESSAGES_RESPONSE) + arguments: Final = { + "model": MESSAGES_MODEL, + "messages": list(MESSAGES), + "max_tokens": 32, + "api_key": "test-key", + "api_base": recording_server.base_url, + } + request: Final = LiteLLMMessagesRequest( + MESSAGES_MODEL, list(MESSAGES), 32, None, "test-key", recording_server.base_url, "anthropic", arguments + ) + + def call() -> object: + return runtime.run( + RouteContext(Route.MESSAGES), + binding=NATIVE_MESSAGES, + native=lambda hook: call_hook(hook, request, (), arguments), + python=runtime.NO_PYTHON, + rules=(RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),), + ) + + first: Final = call() + second: Final = call() + assert payload(first) == payload(second) + assert cache_key(second) + assert len(recording_server.requests) == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("namespace_source", ("cache", "metadata")) +async def test_rust_messages_legacy_cache_honors_request_namespaces( + recording_server: RecordingServer, namespace_source: str +) -> None: + litellm.cache = Cache() + recording_server.expected_requests = 2 + first_options: Final = ( + {"cache": {"namespace": "first"}} if namespace_source == "cache" else {"metadata": {"redis_namespace": "first"}} + ) + second_options: Final = ( + {"cache": {"namespace": "second"}} + if namespace_source == "cache" + else {"metadata": {"redis_namespace": "second"}} + ) + await invoke("messages", recording_server, first_options) + second: Final = await invoke("messages", recording_server, second_options) + first_hit: Final = await invoke("messages", recording_server, first_options) + second_hit: Final = await invoke("messages", recording_server, second_options) + assert cache_key(second) is None + assert cache_key(first_hit) is not None + assert cache_key(second_hit) is not None + assert len(recording_server.requests) == 2 + + +@pytest.mark.asyncio +async def test_rust_messages_legacy_semantic_cache_preserves_python_scope( + recording_server: RecordingServer, +) -> None: + from litellm.caching.in_memory_cache import InMemoryCache + from litellm.types.caching import LiteLLMCacheType + + litellm.cache = Cache(type=LiteLLMCacheType.REDIS_SEMANTIC, _backend=InMemoryCache()) + first_options: Final = {"messages": [{"role": "user", "content": "hello"}]} + second_options: Final = {"messages": [{"role": "user", "content": "hi"}]} + first: Final = await invoke("messages", recording_server, first_options) + second: Final = await invoke("messages", recording_server, second_options) + assert cache_key(first) is None + assert cache_key(second) is not None + assert payload(first) == payload(second) + assert len(recording_server.requests) == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("rust_first", (False, True), ids=("python_to_rust", "rust_to_python")) +@pytest.mark.parametrize("stream", (False, True), ids=("response", "stream")) +async def test_legacy_cache_keeps_public_messages_responses_compatible( + recording_server: RecordingServer, monkeypatch: pytest.MonkeyPatch, rust_first: bool, stream: bool +) -> None: + litellm.cache = Cache() + monkeypatch.setenv("LITELLM_RUST", "0") + options: Final = {"litellm_params": {"preset_cache_key": "shared-messages"}, "stream": stream} + first: Final = await invoke("messages", recording_server, options, native=rust_first) + first_payload: Final = await collect(first) if stream else payload(first) + await asyncio.gather(*tuple(_PENDING_CACHE_WRITES)) + second: Final = await invoke("messages", recording_server, options, native=not rust_first) + second_payload: Final = await collect(second) if stream else payload(second) + assert second_payload == first_payload + assert len(recording_server.requests) == 1 diff --git a/tests/test_litellm_rust/cache/test_valkey_semantic.py b/tests/test_litellm_rust/cache/test_valkey_semantic.py index 046f3a70ae8..ebd00695d95 100644 --- a/tests/test_litellm_rust/cache/test_valkey_semantic.py +++ b/tests/test_litellm_rust/cache/test_valkey_semantic.py @@ -15,11 +15,10 @@ import redis from litellm.caching.caching import Cache from litellm.caching.valkey_semantic_cache import ValkeySemanticCache -from litellm.rust_bridge import _native, catalog -from litellm.rust_bridge.catalog import CacheRule -from litellm.rust_bridge.configuration import Rollout +from litellm.rust_bridge import _native from litellm.rust_bridge.response_cache import ResponseCacheRuntime from litellm.types.caching import LiteLLMCacheType +from tests.test_litellm_rust.support.cache import CacheTestResolver, activate_native, native_runtime pytestmark: Final = pytest.mark.requires_rust_extension embedding_context: Final = contextvars.ContextVar("embedding_context") @@ -92,7 +91,7 @@ def _field_request( def _facade( url: str, index_name: str, - embeddings: Mapping[str, list[float]], + embeddings: Mapping[str, list[float]] | None = None, *, namespace: str | None = None, ) -> Cache: @@ -103,7 +102,7 @@ def _facade( valkey_semantic_cache_index_name=index_name, namespace=namespace, ) - vectors: Final = embeddings + vectors: Final = embeddings or {"semantic cache prompt": [1.0, 0.0]} def embed(prompt: str, metadata: Mapping[str, object] | None = None) -> list[float]: return vectors[prompt] @@ -116,43 +115,15 @@ def _facade( 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) + facade: Final = _facade(valkey_url, index_name) + backend: Final = cast(ValkeySemanticCache, facade.cache) 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() + binding: Final = native_runtime(facade) assert binding.lookup(_request()) == response @@ -160,14 +131,9 @@ 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() + facade: Final = _facade(valkey_url, index_name) + backend: Final = cast(ValkeySemanticCache, facade.cache) + binding: Final = native_runtime(facade) 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"])) @@ -178,14 +144,8 @@ 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() + facade: Final = _facade(valkey_url, index_name) + binding: Final = native_runtime(facade) request: Final = {**_request(), "ttl_seconds": 2.0} await binding.async_store(request, {"answer": "async"}) assert await binding.async_lookup(request) == {"answer": "async"} @@ -195,7 +155,8 @@ async def test_disabled_cache_controls_skip_async_embedding( valkey_url: str, index_name: str, ) -> None: - backend: Final = _backend(valkey_url, index_name) + facade: Final = _facade(valkey_url, index_name) + backend: Final = cast(ValkeySemanticCache, facade.cache) calls: Final = [] async def fail_embedding(prompt: str, metadata: dict[str, object] | None = None) -> list[float]: @@ -203,13 +164,7 @@ async def test_disabled_cache_controls_skip_async_embedding( 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() + binding: Final = native_runtime(facade) controls: Final = { "supported_call_type": True, "configured": True, @@ -234,7 +189,8 @@ async def test_async_embedding_runs_inline_in_caller_task( valkey_url: str, index_name: str, ) -> None: - backend: Final = _backend(valkey_url, index_name) + facade: Final = _facade(valkey_url, index_name) + backend: Final = cast(ValkeySemanticCache, facade.cache) observed: dict[str, object] = {} async def async_embedding(prompt: str, metadata: dict[str, object] | None = None) -> list[float]: @@ -245,13 +201,7 @@ async def test_async_embedding_runs_inline_in_caller_task( 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() + binding: Final = native_runtime(facade) request: Final = {**_request(), "ttl_seconds": 2.0} caller_task: Final = asyncio.current_task() caller_thread: Final = threading.get_ident() @@ -267,7 +217,7 @@ async def test_async_embedding_runs_inline_in_caller_task( embedding_context.reset(token) -def test_facade_activation_and_mutation_fallback( +def test_selected_valkey_runtime_declines_threshold_mutation( valkey_url: str, index_name: str, ) -> None: @@ -277,31 +227,20 @@ def test_facade_activation_and_mutation_fallback( 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)) + activate_native(facade) + resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade)) assert resolver.resolve().kind == "native" facade.cache.similarity_threshold = 0.7 - assert resolver.resolve().kind == "python_callback" + with pytest.raises(_native.RustBridgeDeclined): + resolver.resolve() 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() + facade: Final = _facade(valkey_url, index_name) + binding: Final = native_runtime(facade) with pytest.raises(NotImplementedError): binding.lookup_batch([_request()]) @@ -310,9 +249,8 @@ 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() + facade: Final = _facade(valkey_url, index_name) + binding: Final = native_runtime(facade) 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}:*")) @@ -326,9 +264,9 @@ 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() + facade: Final = _facade(valkey_url, index_name) + backend: Final = cast(ValkeySemanticCache, facade.cache) + binding: Final = native_runtime(facade) response: Final = {"answer": "persistent"} binding.store(_request(), response) client: Final = redis.Redis.from_url(valkey_url) @@ -343,13 +281,13 @@ def test_below_threshold_misses_on_native_and_python( valkey_url: str, index_name: str, ) -> None: - backend: Final = _backend( + facade: Final = _facade( 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() + backend: Final = cast(ValkeySemanticCache, facade.cache) + binding: Final = native_runtime(facade) 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 @@ -359,7 +297,8 @@ 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) + facade: Final = _facade(valkey_url, index_name) + backend: Final = cast(ValkeySemanticCache, facade.cache) client: Final = redis.Redis.from_url(valkey_url) scope: Final = hashlib.sha256(b"key").hexdigest() document: Final = f"{index_name}:{scope}:{uuid4().hex}" @@ -372,8 +311,7 @@ def test_malformed_entry_is_a_miss_on_native_and_python( "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() + binding: Final = native_runtime(facade) assert binding.lookup(_request()) is None assert backend.get_cache("key", messages=_request()["messages"]) is None @@ -382,13 +320,13 @@ def test_mixed_content_parts_match_python_semantic_behavior( valkey_url: str, index_name: str, ) -> None: - backend: Final = _backend(valkey_url, index_name) + facade: Final = _facade(valkey_url, index_name) + backend: Final = cast(ValkeySemanticCache, facade.cache) 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() + binding: Final = native_runtime(facade) request: Final = {**_request(), "messages": messages} binding.store(request, {"answer": "mixed"}) assert binding.lookup(request) is None @@ -401,11 +339,12 @@ async def test_async_store_batch_and_lookup( valkey_url: str, index_name: str, ) -> None: - backend: Final = _backend( + facade: Final = _facade( valkey_url, index_name, {"prompt A": [1.0, 0.0], "prompt B": [0.0, 1.0]}, ) + backend: Final = cast(ValkeySemanticCache, facade.cache) sync_calls: Final = [] async_tasks: Final = [] @@ -422,8 +361,7 @@ async def test_async_store_batch_and_lookup( 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() + binding: Final = native_runtime(facade) requests: Final = [_request("prompt A"), _request("prompt B")] responses: Final = [{"answer": "A"}, {"answer": "B"}] caller_task: Final = asyncio.current_task() @@ -449,7 +387,7 @@ def test_subclass_backend_falls_back_to_python( 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)) + resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade)) assert resolver.resolve().kind == "python_callback" @@ -464,13 +402,7 @@ def test_field_key_matches_python_semantic_scope( 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: Final = native_runtime(facade) 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}:*")) @@ -492,13 +424,7 @@ def test_field_key_reads_all_python_tenant_metadata_sources( 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: Final = native_runtime(facade) binding.store( _field_request( "semantic cache prompt", @@ -536,13 +462,7 @@ def test_namespace_isolates_semantic_entries( {"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() + binding: Final = native_runtime(facade) 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"}) @@ -563,13 +483,7 @@ def test_field_key_isolates_tenant_scope( 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: Final = native_runtime(facade) binding.store( _field_request("semantic cache prompt", {"user_api_key": "k1"}), {"answer": "tenant one"}, @@ -587,7 +501,7 @@ def test_tls_valkey_facade_falls_back_to_python( similarity_threshold=0.8, valkey_semantic_cache_index_name=index_name, ) - resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade)) + resolver: Final = CacheTestResolver(SimpleNamespace(cache=facade)) assert resolver.resolve().kind == "python_callback" @@ -595,24 +509,19 @@ 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() + facade: Final = _facade(valkey_url, index_name) + binding: Final = native_runtime(facade) with pytest.raises(NotImplementedError): await binding.ping() -async def test_rust_required_rule_activates_the_facade_natively( +async def test_explicit_selection_activates_the_facade_natively( valkey_url: str, index_name: str, monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr( - catalog, - "RULES", - (CacheRule(Rollout.RUST_REQUIRED, backends=frozenset({LiteLLMCacheType.VALKEY_SEMANTIC})),), - ) facade: Final = _facade(valkey_url, index_name, {"semantic cache prompt": [1.0, 0.0]}) + facade._native_cache = ResponseCacheRuntime(_native._ResponseCacheRuntime.from_cache(facade)) # pyright: ignore[reportPrivateUsage] # explicitly select the runtime under test runtime: Final = facade._native_cache # pyright: ignore[reportPrivateUsage] # the activation under test has no public accessor assert isinstance(runtime, ResponseCacheRuntime) assert runtime.kind == "native" diff --git a/tests/test_litellm_rust/support/cache.py b/tests/test_litellm_rust/support/cache.py index 41eb4d25257..c5f2844170a 100644 --- a/tests/test_litellm_rust/support/cache.py +++ b/tests/test_litellm_rust/support/cache.py @@ -1,19 +1,28 @@ -from typing import Final, Protocol +from __future__ import annotations + +from dataclasses import dataclass +from typing import Final, Protocol, TypeAlias from uuid import uuid4 -import pytest - from litellm.caching.caching import Cache -from litellm.rust_bridge import _native, catalog -from litellm.rust_bridge.catalog import CacheRule -from litellm.rust_bridge.configuration import Rollout +from litellm.rust_bridge import _native from litellm.rust_bridge.response_cache import ResponseCacheRuntime -from litellm.types.caching import LiteLLMCacheType - -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 +CacheRuntime: TypeAlias = _native._ResponseCacheRuntime # pyright: ignore[reportPrivateUsage] # private runtime under test + + +class CacheNamespace(Protocol): + @property + def cache(self) -> object: ... + + +@dataclass(frozen=True, slots=True) +class CacheTestResolver: + namespace: CacheNamespace + + def resolve(self) -> CacheRuntime: + return CacheRuntime.from_selected(self.namespace.cache) class CacheLookup(Protocol): @@ -25,8 +34,13 @@ def request(key: str = "key") -> dict[str, object]: return {"key": {"preset": key}} -def require_rust(monkeypatch: pytest.MonkeyPatch, backend: LiteLLMCacheType) -> None: - monkeypatch.setattr(catalog, "RULES", (CacheRule(Rollout.RUST_REQUIRED, backends=frozenset({backend})),)) +def native_runtime(facade: Cache) -> CacheRuntime: + return CacheRuntime.from_cache(facade) + + +def activate_native(facade: Cache) -> Cache: + facade._native_cache = ResponseCacheRuntime(_native._ResponseCacheRuntime.from_cache(facade)) # pyright: ignore[reportPrivateUsage] # explicitly select the runtime under test + return facade def assert_native_runtime(facade: Cache) -> ResponseCacheRuntime: diff --git a/tests/test_litellm_rust/support/fake_gcs.py b/tests/test_litellm_rust/support/fake_gcs.py deleted file mode 100644 index 67eb61798b9..00000000000 --- a/tests/test_litellm_rust/support/fake_gcs.py +++ /dev/null @@ -1,152 +0,0 @@ -from __future__ import annotations - -import json -import threading -from collections.abc import Mapping -from dataclasses import dataclass -from functools import partial -from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer -from socket import socket -from types import MappingProxyType -from typing import Final, cast -from urllib.parse import unquote, urlsplit - - -@dataclass(frozen=True, slots=True) -class RecordedRequest: - method: str - path: str - query: str - headers: Mapping[str, str] - body: bytes - - -class _FakeGcsHandler(BaseHTTPRequestHandler): - def __init__( - self, - request: socket | tuple[bytes, socket], - client_address: tuple[str, int], - server: ThreadingHTTPServer, - *, - fake: FakeGcs, - ) -> None: - self._fake: Final = fake - super().__init__(request, client_address, server) - - def _handle(self) -> None: - parsed: Final = urlsplit(self.path) - content_length: Final = int(self.headers.get("Content-Length", "0")) - body: Final = self.rfile.read(content_length) if content_length else b"" - headers: Final = MappingProxyType( - {name.title(): value for name, value in self.headers.items()} - ) - self._fake.record( - RecordedRequest( - method=self.command, - path=parsed.path, - query=parsed.query, - headers=headers, - body=body, - ) - ) - if self.headers.get("Authorization") != f"Bearer {self._fake.token}": - self._send_json(401, {"error": "unauthorized"}) - return - - upload_prefix: Final = "/upload/storage/v1/b/" - download_prefix: Final = "/storage/v1/b/" - if parsed.path.startswith(upload_prefix) and parsed.path.endswith("/o"): - self._upload(parsed.path[len(upload_prefix) : -2], parsed.query, body) - return - if parsed.path.startswith(download_prefix): - self._download(parsed.path[len(download_prefix) :], parsed.query) - return - self._send_json(404, {"error": "not found"}) - - def _upload(self, path: str, query: str, body: bytes) -> None: - values: Final = { - unquote(pair.partition("=")[0]): unquote(pair.partition("=")[2]) - for pair in query.split("&") - if pair - } - if not path or values.get("uploadType") != "media" or "name" not in values: - self._send_json(404, {"error": "not found"}) - return - self._fake.put_object(path, values["name"], body) - self._send_json(200, {"name": values["name"], "bucket": path}) - - def _download(self, path: str, query: str) -> None: - bucket, separator, encoded_name = path.partition("/o/") - if not separator or query != "alt=media": - self._send_json(404, {"error": "not found"}) - return - name: Final = unquote(encoded_name) - if name.endswith("/server-error") or name == "server-error": - self._send_json(500, {"error": "server error"}) - return - body: Final = self._fake.get_object(bucket, name) - if body is None: - self._send_json(404, {"error": "not found"}) - return - self._send(200, body, "application/octet-stream") - - def _send_json(self, status: int, value: object) -> None: - payload: Final = json.dumps(value).encode() - self._send(status, payload, "application/json") - - def _send(self, status: int, body: bytes, content_type: str) -> None: - self.send_response(status) - self.send_header("Content-Type", content_type) - self.send_header("Content-Length", str(len(body))) - self.end_headers() - self.wfile.write(body) - - def log_message(self, format: str, *args: object) -> None: - pass - - do_GET = _handle - do_POST = _handle - - -class FakeGcs: - def __init__(self) -> None: - self._objects: dict[tuple[str, str], bytes] = {} # mutable-ok: fake object store - self._requests: list[RecordedRequest] = [] # mutable-ok: recorded request history - self._server = ThreadingHTTPServer( - ("127.0.0.1", 0), - partial(_FakeGcsHandler, fake=self), - ) - self._worker = threading.Thread(target=self._server.serve_forever, daemon=True) - self._worker.start() - self.token: Final = "test-token" - - @property - def url(self) -> str: - address: Final = cast(tuple[str, int], self._server.server_address) - host, port = address - return f"http://{host}:{port}" - - @property - def objects(self) -> Mapping[tuple[str, str], bytes]: - return MappingProxyType(self._objects) - - @property - def requests(self) -> tuple[RecordedRequest, ...]: - return tuple(self._requests) - - def put(self, bucket: str, name: str, body: bytes) -> None: - self.put_object(bucket, name, body) - - def close(self) -> None: - self._server.shutdown() - self._server.server_close() - self._worker.join(timeout=5) - - def record(self, request: RecordedRequest) -> None: - self._requests.append(request) - - def put_object(self, bucket: str, name: str, body: bytes) -> None: - self._objects[(bucket, name)] = body - - def get_object(self, bucket: str, name: str) -> bytes | None: - return self._objects.get((bucket, name)) diff --git a/tests/unit/rust_bridge/test_callbacks_legacy_python.py b/tests/unit/rust_bridge/test_callbacks_legacy_python.py index 7365679a28c..045dcb52ce9 100644 --- a/tests/unit/rust_bridge/test_callbacks_legacy_python.py +++ b/tests/unit/rust_bridge/test_callbacks_legacy_python.py @@ -12,6 +12,7 @@ from litellm._internal_context import is_internal_call from litellm.litellm_core_utils.litellm_logging import Logging from litellm.rust_bridge import callbacks_legacy_python as legacy from litellm.rust_bridge.callbacks_legacy_python import failure_handler, setup +from litellm.types.utils import ModelResponse _OCR_KWARGS: Final = MappingProxyType( { @@ -41,6 +42,34 @@ def test_setup_reuses_a_supplied_logger() -> None: assert result.logger is supplied +@pytest.mark.parametrize("explicit_provider", (None, "openai")) +def test_cache_hit_finalization_preserves_execution_provider_attribution(explicit_provider: str | None) -> None: + now: Final = datetime.datetime.now() + kwargs: Final = { + "model": "openai/cache-test-model", + "messages": [{"role": "user", "content": "hello"}], + "custom_llm_provider": explicit_provider, + "metadata": {"user_api_key": "key-hash"}, + } + prepared: Final = setup("acompletion", (), kwargs, now, asynchronous=True) + legacy.update_logging( + prepared.logger, + prepared.kwargs, + "resolved-cache-model", + {}, + {**prepared.logger.litellm_params, "custom_llm_provider": "azure"}, + "azure", + ) + prepared.logger.model_call_details.update({"cache_hit": True, "cache_key": "cached-response"}) + response: Final = ModelResponse(model="cache-test-model") + legacy.finalize(response, prepared.logger, prepared.kwargs, now, now) + assert prepared.logger.model_call_details["custom_llm_provider"] == "azure" + assert prepared.logger.model_call_details["model"] == "resolved-cache-model" + assert prepared.logger.litellm_params["metadata"]["user_api_key"] == "key-hash" + assert response._hidden_params["cache_key"] == "cached-response" + assert response._hidden_params["response_cost"] == 0 + + @pytest.mark.parametrize( "call_type, kwargs", [ diff --git a/tests/unit/rust_bridge/test_catalog.py b/tests/unit/rust_bridge/test_catalog.py index d3cab8679ef..95e98fe98da 100644 --- a/tests/unit/rust_bridge/test_catalog.py +++ b/tests/unit/rust_bridge/test_catalog.py @@ -7,8 +7,6 @@ import pytest from litellm.rust_bridge import catalog, configuration from litellm.rust_bridge.catalog import ( - CacheContext, - CacheRule, Context, LoggerContext, Route, @@ -19,7 +17,6 @@ from litellm.rust_bridge.catalog import ( SecretManagerRule, ) from litellm.rust_bridge.configuration import Decision, Rollout -from litellm.types.caching import LiteLLMCacheType from litellm.types.secret_managers.main import KeyManagementSystem @@ -71,9 +68,7 @@ def test_missing_rule_stays_on_python_even_when_rust_is_enabled(monkeypatch: pyt @pytest.mark.parametrize( "context", ( - *(CacheContext(backend.value) for backend in LiteLLMCacheType), *(SecretManagerContext(system.value) for system in KeyManagementSystem), - CacheContext("custom"), SecretManagerContext("unknown"), ), ) @@ -94,16 +89,6 @@ def test_logger_rollout_obeys_the_global_switch() -> None: assert catalog.decision(LoggerContext()) is Decision.RUST_WITH_FALLBACK -def test_response_cache_rules_select_the_whole_backend_runtime() -> None: - rules: Final = ( - CacheRule(Rollout.RUST_REQUIRED, backends=frozenset({"local"})), - CacheRule(Rollout.PYTHON_ONLY), - ) - - assert catalog.decision(CacheContext(backend="local"), rules) is Decision.RUST_REQUIRED - assert catalog.decision(CacheContext(backend="redis"), rules) is Decision.PYTHON - - @pytest.mark.parametrize( ("context", "expected"), ( @@ -149,16 +134,12 @@ def test_ocr_has_no_python_path_to_opt_out_to( (RouteContext(Route.OCR, provider="local"), Decision.RUST_REQUIRED), (RouteContext(Route.OCR, provider="other"), Decision.PYTHON), (RouteContext(Route.MESSAGES, provider="local"), Decision.PYTHON), - (CacheContext("local"), Decision.RUST_WITH_FALLBACK), - (CacheContext("other"), Decision.PYTHON), (SecretManagerContext("local"), Decision.PYTHON), (SecretManagerContext("other"), Decision.RUST_REQUIRED), ), ) def test_mixed_rules_select_only_the_matching_domain(context: Context, expected: Decision) -> None: rules: Final[Rules] = ( - CacheRule(Rollout.RUST_OPT_OUT, backends=frozenset({"local"})), - CacheRule(Rollout.PYTHON_ONLY), SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({"local"})), SecretManagerRule(Rollout.RUST_REQUIRED), RouteRule(Route.OCR, Rollout.RUST_REQUIRED, providers=frozenset({"local"})), @@ -168,7 +149,7 @@ def test_mixed_rules_select_only_the_matching_domain(context: Context, expected: assert catalog.decision(context, rules) is expected -@pytest.mark.parametrize("context", (RouteContext(Route.OCR), CacheContext("local"), SecretManagerContext("local"))) +@pytest.mark.parametrize("context", (RouteContext(Route.OCR), SecretManagerContext("local"))) @pytest.mark.parametrize( ("rollout", "process", "environment", "expected"), ( @@ -195,10 +176,8 @@ def test_all_domains_share_rollout_switches_and_first_match( monkeypatch.setenv("LITELLM_RUST", environment) rules: Final[Rules] = ( RouteRule(Route.OCR, rollout), - CacheRule(rollout), SecretManagerRule(rollout), RouteRule(Route.OCR, Rollout.RUST_REQUIRED), - CacheRule(Rollout.RUST_REQUIRED), SecretManagerRule(Rollout.RUST_REQUIRED), ) @@ -206,11 +185,10 @@ def test_all_domains_share_rollout_switches_and_first_match( assert catalog.decision(context, ()) is Decision.PYTHON -@pytest.mark.parametrize("context", (RouteContext(Route.OCR), CacheContext("local"), SecretManagerContext("local"))) +@pytest.mark.parametrize("context", (RouteContext(Route.OCR), SecretManagerContext("local"))) def test_empty_constraints_match_nothing(context: Context) -> None: rules: Final[Rules] = ( RouteRule(Route.OCR, Rollout.RUST_REQUIRED, providers=frozenset()), - CacheRule(Rollout.RUST_REQUIRED, backends=frozenset()), SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset()), ) diff --git a/tests/unit/rust_bridge/test_dispatch.py b/tests/unit/rust_bridge/test_dispatch.py index 9c3e35edda4..974e943ea9f 100644 --- a/tests/unit/rust_bridge/test_dispatch.py +++ b/tests/unit/rust_bridge/test_dispatch.py @@ -6,7 +6,7 @@ import pytest from litellm.rust_bridge import configuration from litellm.rust_bridge.bindings import NativeBinding -from litellm.rust_bridge.catalog import CacheRule, Route, RouteContext, RouteRule, Rules, SecretManagerRule +from litellm.rust_bridge.catalog import Route, RouteContext, RouteRule, Rules, SecretManagerRule from litellm.rust_bridge.configuration import Rollout from litellm.rust_bridge.dispatch import PublicDispatch from litellm.rust_bridge.runtime import NO_PYTHON, NoPythonImplementationError @@ -23,7 +23,7 @@ def binding() -> NativeBinding[object]: return bound -@pytest.mark.parametrize("rules", ((), (CacheRule(Rollout.RUST_REQUIRED), SecretManagerRule(Rollout.RUST_REQUIRED)))) +@pytest.mark.parametrize("rules", ((), (SecretManagerRule(Rollout.RUST_REQUIRED),))) def test_route_without_rules_forwards_before_request_projection(rules: Rules) -> None: stream: Final[Iterator[int]] = iter((1, 2)) @@ -97,7 +97,6 @@ def test_native_stream_result_is_not_consumed_or_wrapped() -> None: request: Final = Request(model="streaming-model") stream: Final[Iterator[int]] = iter((1, 2)) rules: Final[Rules] = ( - CacheRule(Rollout.PYTHON_ONLY), SecretManagerRule(Rollout.PYTHON_ONLY), RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED), ) @@ -126,7 +125,7 @@ def test_native_stream_result_is_not_consumed_or_wrapped() -> None: @pytest.mark.asyncio -@pytest.mark.parametrize("rules", ((), (CacheRule(Rollout.RUST_REQUIRED), SecretManagerRule(Rollout.RUST_REQUIRED)))) +@pytest.mark.parametrize("rules", ((), (SecretManagerRule(Rollout.RUST_REQUIRED),))) async def test_async_route_without_rules_preserves_async_iterator_result(rules: Rules) -> None: async def chunks() -> AsyncGenerator[int, None]: yield 1 diff --git a/tests/unit/rust_bridge/test_runtime.py b/tests/unit/rust_bridge/test_runtime.py index fd84d629937..bc3c7b43d75 100644 --- a/tests/unit/rust_bridge/test_runtime.py +++ b/tests/unit/rust_bridge/test_runtime.py @@ -229,8 +229,15 @@ async def test_python_fallback_does_not_claim_rust_execution(missing: bool) -> N @pytest.mark.asyncio @pytest.mark.parametrize("shape", ("model", "dict")) @pytest.mark.parametrize("asynchronous", (False, True)) -async def test_native_response_marker_reaches_caller_with_existing_metadata(shape: str, asynchronous: bool) -> None: - hidden: Final = {"additional_headers": {"x-request-id": "upstream"}, "response_cost": 0.01} +@pytest.mark.parametrize("cache_key", (None, "test-cache-key")) +async def test_native_response_marker_reaches_caller_with_existing_metadata( + shape: str, asynchronous: bool, cache_key: str | None +) -> None: + hidden: Final = { + "additional_headers": {"x-request-id": "upstream"}, + "response_cost": 0.01, + **({"cache_key": cache_key} if cache_key is not None else {}), + } response: Final[OCRResponse | dict[str, object]] = ( OCRResponse(pages=[], model="native") if shape == "model" else {"content": "native", "_hidden_params": hidden} ) @@ -258,7 +265,12 @@ async def test_native_response_marker_reaches_caller_with_existing_metadata(shap assert result is response assert get_hidden_params_dict(result) == { "response_cost": 0.01, - "additional_headers": {"x-request-id": "upstream", "x-litellm-rust": "true"}, + "additional_headers": { + "x-request-id": "upstream", + "x-litellm-rust": "true", + **({"x-litellm-cache-key": cache_key} if cache_key is not None else {}), + }, + **({"cache_key": cache_key} if cache_key is not None else {}), }