diff --git a/litellm-rust/crates/python-bridge/src/cache/binding.rs b/litellm-rust/crates/python-bridge/src/cache/binding.rs index fcf71e10427..6b5de7f029b 100644 --- a/litellm-rust/crates/python-bridge/src/cache/binding.rs +++ b/litellm-rust/crates/python-bridge/src/cache/binding.rs @@ -29,6 +29,7 @@ pub(super) enum CacheBinding { #[pyclass(frozen, name = "_ResponseCacheRuntime")] pub(crate) struct ResolvedCache { binding: CacheBinding, + guard: Option, pid: u32, } @@ -36,10 +37,24 @@ impl ResolvedCache { pub(super) fn new(binding: CacheBinding) -> Self { Self { binding, + guard: None, pid: std::process::id(), } } + pub(super) fn with_guard(mut self, guard: super::facade::FacadeGuard) -> Self { + self.guard = Some(guard); + self + } + + pub(super) fn native_service(&self) -> PyResult> { + self.check_process()?; + Ok(match &self.binding { + CacheBinding::Native(service) => Some(service.clone()), + _ => None, + }) + } + fn check_process(&self) -> PyResult<()> { if matches!(self.binding, CacheBinding::Native(_)) && self.pid != std::process::id() { return Err(PyRuntimeError::new_err( @@ -70,6 +85,43 @@ impl ResolvedCache { #[pymethods] impl ResolvedCache { + #[staticmethod] + pub(crate) fn from_selected(cache: &Bound<'_, PyAny>) -> PyResult { + 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()) + { + let resolved = runtime + .getattr("native")? + .extract::>()?; + match resolved.native_service()? { + Some(service) => { + if !resolved + .guard + .as_ref() + .is_some_and(|guard| guard.matches(py, cache).unwrap_or(false)) + { + return Err(RustBridgeDeclined::new_err( + "native cache runtime no longer matches its facade", + )); + } + CacheBinding::Native(service) + } + None => CacheBinding::PythonCallback(PythonCallback::new(cache.clone().unbind())), + } + } else { + CacheBinding::PythonCallback(PythonCallback::new(cache.clone().unbind())) + }; + Ok(Self::new(binding)) + } + #[staticmethod] fn from_cache(cache: &Bound<'_, PyAny>) -> PyResult { let config = match NativeCacheConfig::project(cache)? { @@ -80,7 +132,13 @@ impl ResolvedCache { }; let backend = cache.getattr("cache")?; let service = activate(cache.py(), &backend, config)?; - Ok(Self::new(CacheBinding::Native(service))) + let resolved = Self::new(CacheBinding::Native(service.clone())); + Ok( + match super::facade::FacadeGuard::capture(cache.py(), cache, &service) { + Ok(guard) => resolved.with_guard(guard), + Err(_) => resolved, + }, + ) } #[getter] @@ -323,6 +381,9 @@ impl ResolvedCache { if let CacheBinding::PythonCallback(callback) = &self.binding { callback.traverse(&visit)?; } + if let Some(guard) = &self.guard { + guard.traverse(visit)?; + } Ok(()) } } diff --git a/litellm-rust/crates/python-bridge/src/cache/facade.rs b/litellm-rust/crates/python-bridge/src/cache/facade.rs index 88fde2f6de0..d1bddef67ff 100644 --- a/litellm-rust/crates/python-bridge/src/cache/facade.rs +++ b/litellm-rust/crates/python-bridge/src/cache/facade.rs @@ -472,7 +472,7 @@ impl FacadeGuard { }) } - fn matches(&self, py: Python<'_>, facade: &Bound<'_, PyAny>) -> PyResult { + pub(super) fn matches(&self, py: Python<'_>, facade: &Bound<'_, PyAny>) -> PyResult { if !self.outer.matches(py, facade)? { return Ok(false); } diff --git a/litellm-rust/crates/python-bridge/src/cache/mod.rs b/litellm-rust/crates/python-bridge/src/cache/mod.rs index 0cfd4ac8138..ac1e00d5273 100644 --- a/litellm-rust/crates/python-bridge/src/cache/mod.rs +++ b/litellm-rust/crates/python-bridge/src/cache/mod.rs @@ -20,9 +20,7 @@ use pyo3::{ types::PyDict, }; -pub(crate) use self::{ - binding::ResolvedCache, handle::CacheTestHandle, resolver::CacheTestResolver, -}; +pub(crate) use self::{binding::ResolvedCache, handle::CacheTestHandle, resolver::CacheResolver}; fn cache_error(error: Error) -> PyErr { match error { diff --git a/litellm-rust/crates/python-bridge/src/cache/resolver.rs b/litellm-rust/crates/python-bridge/src/cache/resolver.rs index ef6f142e0a1..3baaada4b17 100644 --- a/litellm-rust/crates/python-bridge/src/cache/resolver.rs +++ b/litellm-rust/crates/python-bridge/src/cache/resolver.rs @@ -1,19 +1,14 @@ use pyo3::{PyTraverseError, PyVisit, prelude::*}; -use super::{ - binding::{CacheBinding, ResolvedCache}, - callback::PythonCallback, - facade, - handle::CacheTestHandle, -}; +use super::binding::ResolvedCache; -#[pyclass(frozen, name = "_CacheTestResolver")] -pub(crate) struct CacheTestResolver { +#[pyclass(frozen, name = "_CacheResolver")] +pub(crate) struct CacheResolver { namespace: Py, } #[pymethods] -impl CacheTestResolver { +impl CacheResolver { #[new] fn new(namespace: Py) -> Self { Self { namespace } @@ -21,16 +16,7 @@ impl CacheTestResolver { pub(crate) fn resolve(&self, py: Python<'_>) -> PyResult { let object = self.namespace.bind(py).getattr("cache")?; - let binding = if object.is_none() { - CacheBinding::Disabled - } else if let Ok(handle) = object.extract::>() { - CacheBinding::Native(handle.service()?) - } else if let Some(service) = facade::resolve(py, &object)? { - CacheBinding::Native(service) - } else { - CacheBinding::PythonCallback(PythonCallback::new(object.unbind())) - }; - Ok(ResolvedCache::new(binding)) + ResolvedCache::from_selected(&object) } fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> { diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 54b13ba01bb..874e522af15 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -13,7 +13,7 @@ mod tokenizer; #[pymodule(gil_used = true)] mod _native { - use crate::cache::{CacheTestHandle, CacheTestResolver, ResolvedCache}; + use crate::cache::{CacheResolver, CacheTestHandle, ResolvedCache}; #[cfg(feature = "panic-test")] #[pymodule_export] use crate::diagnostics::_panic_for_test; @@ -51,7 +51,8 @@ mod _native { let py = module.py(); let dict = module.dict(); dict.set_item("_CacheTestHandle", py.get_type::())?; - dict.set_item("_CacheTestResolver", py.get_type::())?; + dict.set_item("_CacheResolver", py.get_type::())?; + dict.set_item("_CacheTestResolver", py.get_type::())?; dict.set_item("_ResponseCacheRuntime", py.get_type::())?; dict.set_item( "_SecretManagerRuntime", diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index 60d2e6224c0..0e684d1f10c 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -116,6 +116,8 @@ class ResponsesWebSocketConnection: class _ResponseCacheRuntime: @staticmethod def from_cache(cache: object) -> _ResponseCacheRuntime: ... + @staticmethod + def from_selected(cache: object) -> _ResponseCacheRuntime: ... @property def kind(self) -> str: ... def lookup( @@ -238,6 +240,11 @@ class _CacheTestHandle: 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: ... diff --git a/tests/test_litellm_rust/test_cache.py b/tests/test_litellm_rust/test_cache.py index e2e2f9f1819..96b3674fde3 100644 --- a/tests/test_litellm_rust/test_cache.py +++ b/tests/test_litellm_rust/test_cache.py @@ -236,6 +236,57 @@ async def test_catalog_constructs_native_runtime_from_public_cache_configuration assert await runtime.async_lookup(async_request) is None +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) + assert isinstance(runtime, ResponseCacheRuntime) + facade._native_cache = runtime + + selected: Final = _native._CacheResolver(SimpleNamespace(cache=facade)).resolve() + assert selected.kind == "native" + request: Final = runtime.request(facade, {"cache_key": "inference-native"}) + assert request is not None + await selected.async_store(request, {"answer": 42}) + assert await selected.async_lookup(request) == {"answer": 42} + assert await runtime.async_lookup(request) == {"answer": 42} + assert facade.cache.get_cache("inference-native") is None + + facade._native_cache = None + fallback: Final = _native._CacheResolver(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} + assert facade.cache.get_cache("inference-python") is not None + + +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) + assert isinstance(runtime, ResponseCacheRuntime) + facade._native_cache = runtime + stale_request: Final = runtime.request(facade, {"cache_key": "stale-only"}) + assert stale_request is not None + await runtime.async_store(stale_request, {"answer": "stale"}) + + replacement: Final = InMemoryCache() + facade.cache = replacement + with pytest.raises(_native.RustBridgeDeclined): + _native._CacheResolver(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 + + def test_existing_global_lifecycle_remains_the_resolver_source_of_truth() -> None: resolver: Final = _CacheTestResolver(litellm)