mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
feat(cache): add a guarded native response-cache resolver foundation (#42769)
* feat(cache): resolve configured backend for native inference * fix(cache): reuse the resolved native runtime only while its facade guard matches Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cache): decline native inference when the resolved runtime no longer matches its facade Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Yujong Lee <yujong@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
ccee9e77ce
commit
02d1e2c579
7 changed files with 130 additions and 26 deletions
|
|
@ -29,6 +29,7 @@ pub(super) enum CacheBinding {
|
|||
#[pyclass(frozen, name = "_ResponseCacheRuntime")]
|
||||
pub(crate) struct ResolvedCache {
|
||||
binding: CacheBinding,
|
||||
guard: Option<super::facade::FacadeGuard>,
|
||||
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<Option<NativeResponseCache>> {
|
||||
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<Self> {
|
||||
let py = cache.py();
|
||||
let binding = if cache.is_none() {
|
||||
CacheBinding::Disabled
|
||||
} else if let Ok(handle) = cache.extract::<PyRef<'_, super::handle::CacheTestHandle>>() {
|
||||
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::<PyRef<'_, ResolvedCache>>()?;
|
||||
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<Self> {
|
||||
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(())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -472,7 +472,7 @@ impl FacadeGuard {
|
|||
})
|
||||
}
|
||||
|
||||
fn matches(&self, py: Python<'_>, facade: &Bound<'_, PyAny>) -> PyResult<bool> {
|
||||
pub(super) fn matches(&self, py: Python<'_>, facade: &Bound<'_, PyAny>) -> PyResult<bool> {
|
||||
if !self.outer.matches(py, facade)? {
|
||||
return Ok(false);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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<PyAny>,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl CacheTestResolver {
|
||||
impl CacheResolver {
|
||||
#[new]
|
||||
fn new(namespace: Py<PyAny>) -> Self {
|
||||
Self { namespace }
|
||||
|
|
@ -21,16 +16,7 @@ impl CacheTestResolver {
|
|||
|
||||
pub(crate) fn resolve(&self, py: Python<'_>) -> PyResult<ResolvedCache> {
|
||||
let object = self.namespace.bind(py).getattr("cache")?;
|
||||
let binding = if object.is_none() {
|
||||
CacheBinding::Disabled
|
||||
} else if let Ok(handle) = object.extract::<PyRef<'_, CacheTestHandle>>() {
|
||||
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> {
|
||||
|
|
|
|||
|
|
@ -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::<CacheTestHandle>())?;
|
||||
dict.set_item("_CacheTestResolver", py.get_type::<CacheTestResolver>())?;
|
||||
dict.set_item("_CacheResolver", py.get_type::<CacheResolver>())?;
|
||||
dict.set_item("_CacheTestResolver", py.get_type::<CacheResolver>())?;
|
||||
dict.set_item("_ResponseCacheRuntime", py.get_type::<ResolvedCache>())?;
|
||||
dict.set_item(
|
||||
"_SecretManagerRuntime",
|
||||
|
|
|
|||
|
|
@ -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: ...
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue