mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
fix(cache): align native semantic cache scope keys
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
0d09e9d892
commit
4509eb9914
5 changed files with 286 additions and 3 deletions
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -2700,6 +2700,7 @@ dependencies = [
|
|||
"rstest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sha2 0.10.9",
|
||||
"tokio",
|
||||
"tokio-tungstenite",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -46,6 +46,7 @@ tokio = { workspace = true, features = ["sync"] }
|
|||
criterion.workspace = true
|
||||
futures-util.workspace = true
|
||||
rstest.workspace = true
|
||||
sha2.workspace = true
|
||||
tokio-tungstenite.workspace = true
|
||||
|
||||
[[bench]]
|
||||
|
|
|
|||
|
|
@ -322,9 +322,17 @@ fn project_valkey_semantic(
|
|||
) -> PyResult<Result<ValkeySemanticCacheConfig, UnsupportedCacheConfig>> {
|
||||
let client = backend.getattr("sync_client")?;
|
||||
let pool = client.getattr("connection_pool")?;
|
||||
let Ok((resolved, _is_tls)) = project_connection_pool(&pool)? else {
|
||||
let Ok((resolved, is_tls)) = project_connection_pool(&pool)? else {
|
||||
return Ok(Err(UnsupportedCacheConfig::RedisConnection));
|
||||
};
|
||||
for key in ["credential_provider", "redis_connect_func"] {
|
||||
if has_value(&resolved, key)? {
|
||||
return Ok(Err(UnsupportedCacheConfig::RedisCredentials));
|
||||
}
|
||||
}
|
||||
if is_tls {
|
||||
return Ok(Err(UnsupportedCacheConfig::RedisConnection));
|
||||
}
|
||||
let connection = RedisConnectionConfig {
|
||||
host: required_string(&resolved, "host")?,
|
||||
port: u16::try_from(required_i64(&resolved, "port")?)
|
||||
|
|
@ -682,6 +690,53 @@ mod tests {
|
|||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn valkey_semantic_tls_stays_on_python() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let facade = facade(
|
||||
py,
|
||||
"pool = ConnectionPool()\n\
|
||||
pool.connection_class = SSLConnection\n\
|
||||
pool.connection_kwargs = {'host': 'cache.internal', 'port': 6390}\n\
|
||||
client = SimpleNamespace(connection_pool=pool)\n\
|
||||
backend = SimpleNamespace(similarity_threshold=0.85, index_name='semantic_idx', embedding_model='text-embedding-3-small', sync_client=client)\n\
|
||||
facade = SimpleNamespace(type='valkey-semantic', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)",
|
||||
);
|
||||
let CacheConfigProjection::Unsupported(reason) =
|
||||
NativeCacheConfig::project(&facade).unwrap()
|
||||
else {
|
||||
panic!("TLS Valkey semantic cache should stay on Python");
|
||||
};
|
||||
assert_eq!(
|
||||
reason.message(),
|
||||
"native Redis connection type is not implemented"
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn valkey_semantic_dynamic_auth_stays_on_python() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let facade = facade(
|
||||
py,
|
||||
"pool = ConnectionPool()\n\
|
||||
pool.connection_class = Connection\n\
|
||||
pool.connection_kwargs = {'host': 'cache.internal', 'port': 6390, 'credential_provider': object()}\n\
|
||||
client = SimpleNamespace(connection_pool=pool)\n\
|
||||
backend = SimpleNamespace(similarity_threshold=0.85, index_name='semantic_idx', embedding_model='text-embedding-3-small', sync_client=client)\n\
|
||||
facade = SimpleNamespace(type='valkey-semantic', mode='default-on', ttl=None, namespace=None, supported_call_types=None, redis_flush_size=None, semantic_cache_scope='key', cache=backend)",
|
||||
);
|
||||
let CacheConfigProjection::Unsupported(reason) =
|
||||
NativeCacheConfig::project(&facade).unwrap()
|
||||
else {
|
||||
panic!("dynamic Valkey authentication must stay on Python");
|
||||
};
|
||||
assert_eq!(reason.message(), "native Redis credentials require Python");
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dynamic_redis_auth_stays_on_python() {
|
||||
Python::initialize();
|
||||
|
|
|
|||
|
|
@ -6,13 +6,50 @@ use litellm_cache::{
|
|||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_redis::RedisCache;
|
||||
use litellm_cache_response::{
|
||||
CacheEntry, PartialHits, ResponseCache, ResponseCacheCodec, ResponseCacheRequest, WriteBuffer,
|
||||
CacheEntry, CacheKeyField, PartialHits, ResponseCache, ResponseCacheCodec,
|
||||
ResponseCacheRequest, WriteBuffer,
|
||||
};
|
||||
use litellm_cache_valkey_semantic::{ValkeySemanticCache, ValkeySemanticConfig};
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{embedder::PythonEmbedder, request::NativeRequest};
|
||||
|
||||
fn semantic_key(request: &NativeRequest, scope: &str) -> litellm_cache_response::CacheKeyInput {
|
||||
let mut key = request.key.clone();
|
||||
if key.preset.is_some() {
|
||||
return key;
|
||||
}
|
||||
key.fields
|
||||
.retain(|field| !matches!(field.name.as_str(), "messages" | "prompt" | "input"));
|
||||
const TENANT: [&str; 3] = [
|
||||
"user_api_key",
|
||||
"user_api_key_team_id",
|
||||
"user_api_key_org_id",
|
||||
];
|
||||
let end_user = (scope == "end_user").then_some("user_api_key_end_user_id");
|
||||
for name in TENANT.into_iter().chain(end_user) {
|
||||
let Some(value) = request
|
||||
.metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.get(name))
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
let value = match value {
|
||||
Value::Null => continue,
|
||||
Value::String(text) => text.clone(),
|
||||
other => other.to_string(),
|
||||
};
|
||||
key.fields.push(CacheKeyField {
|
||||
name: name.to_owned(),
|
||||
value: Some(value),
|
||||
api_parameter: true,
|
||||
internal_parameter: false,
|
||||
});
|
||||
}
|
||||
key
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(super) enum NativeResponseCache {
|
||||
Memory(Arc<ResponseCache<InMemoryCache<CacheEntry>>>),
|
||||
|
|
@ -88,7 +125,7 @@ impl NativeResponseCache {
|
|||
scope: &str,
|
||||
) -> ResponseCacheRequest<SemanticCacheContext> {
|
||||
ResponseCacheRequest {
|
||||
key: request.key.clone(),
|
||||
key: semantic_key(request, scope),
|
||||
controls: request.controls,
|
||||
context: SemanticCacheContext {
|
||||
input: request.input.clone(),
|
||||
|
|
@ -329,3 +366,77 @@ impl NativeResponseCache {
|
|||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_cache_response::{CacheControls, CacheKeyInput, cache_key};
|
||||
use serde_json::json;
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
use super::*;
|
||||
|
||||
fn native_request(key: CacheKeyInput, metadata: Value) -> NativeRequest {
|
||||
NativeRequest {
|
||||
key,
|
||||
controls: CacheControls::default(),
|
||||
ttl: None,
|
||||
max_age: None,
|
||||
messages: Some(json!([{"role": "user", "content": "prompt"}])),
|
||||
input: None,
|
||||
metadata: Some(metadata),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn semantic_key_matches_python_scope_material() {
|
||||
let key = CacheKeyInput {
|
||||
fields: vec![
|
||||
CacheKeyField {
|
||||
name: "model".to_owned(),
|
||||
value: Some("gpt-4.1".to_owned()),
|
||||
api_parameter: true,
|
||||
internal_parameter: false,
|
||||
},
|
||||
CacheKeyField {
|
||||
name: "messages".to_owned(),
|
||||
value: Some("prompt".to_owned()),
|
||||
api_parameter: true,
|
||||
internal_parameter: false,
|
||||
},
|
||||
],
|
||||
..Default::default()
|
||||
};
|
||||
let request = native_request(
|
||||
key,
|
||||
json!({"user_api_key": "k1", "user_api_key_team_id": null}),
|
||||
);
|
||||
let expected = format!("{:x}", Sha256::digest(b"model: gpt-4.1user_api_key: k1"));
|
||||
assert_eq!(cache_key(&semantic_key(&request, "key")), expected);
|
||||
|
||||
let end_user_request = native_request(
|
||||
request.key.clone(),
|
||||
json!({"user_api_key": "k1", "user_api_key_end_user_id": "u1"}),
|
||||
);
|
||||
let expected = format!(
|
||||
"{:x}",
|
||||
Sha256::digest(b"model: gpt-4.1user_api_key: k1user_api_key_end_user_id: u1")
|
||||
);
|
||||
assert_eq!(
|
||||
cache_key(&semantic_key(&end_user_request, "end_user")),
|
||||
expected
|
||||
);
|
||||
|
||||
let preset_request = native_request(
|
||||
CacheKeyInput {
|
||||
preset: Some("preset-key".to_owned()),
|
||||
..Default::default()
|
||||
},
|
||||
json!({"user_api_key": "k1"}),
|
||||
);
|
||||
assert_eq!(
|
||||
semantic_key(&preset_request, "end_user").preset.as_deref(),
|
||||
Some("preset-key")
|
||||
);
|
||||
assert!(semantic_key(&preset_request, "end_user").fields.is_empty());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -46,6 +46,56 @@ def _request(prompt: str = "semantic cache prompt") -> dict[str, object]:
|
|||
}
|
||||
|
||||
|
||||
def _field_request(
|
||||
prompt: str,
|
||||
metadata: Mapping[str, object],
|
||||
) -> dict[str, object]:
|
||||
return {
|
||||
"key": {
|
||||
"fields": [
|
||||
{
|
||||
"name": "model",
|
||||
"value": "gpt-4.1",
|
||||
"api_parameter": True,
|
||||
"internal_parameter": False,
|
||||
},
|
||||
{
|
||||
"name": "messages",
|
||||
"value": prompt,
|
||||
"api_parameter": True,
|
||||
"internal_parameter": False,
|
||||
},
|
||||
]
|
||||
},
|
||||
"messages": [{"role": "user", "content": prompt}],
|
||||
"metadata": dict(metadata),
|
||||
}
|
||||
|
||||
|
||||
def _facade(
|
||||
url: str,
|
||||
index_name: str,
|
||||
embeddings: Mapping[str, list[float]],
|
||||
) -> Cache:
|
||||
facade: Final = Cache(
|
||||
type=LiteLLMCacheType.VALKEY_SEMANTIC,
|
||||
redis_url=url,
|
||||
similarity_threshold=0.8,
|
||||
valkey_semantic_cache_index_name=index_name,
|
||||
)
|
||||
vectors: Final = embeddings
|
||||
|
||||
def embed(prompt: str, metadata: Mapping[str, object] | None = None) -> list[float]:
|
||||
return vectors[prompt]
|
||||
|
||||
async def async_embedding(prompt: str, metadata: dict[str, object] | None = None) -> list[float]:
|
||||
return vectors[prompt]
|
||||
|
||||
facade.cache._get_embedding = embed
|
||||
facade.cache._get_async_embedding = async_embedding
|
||||
return facade
|
||||
|
||||
|
||||
def _backend(
|
||||
url: str,
|
||||
index_name: str,
|
||||
|
|
@ -268,6 +318,71 @@ def test_subclass_backend_falls_back_to_python(
|
|||
assert resolver.resolve().kind == "python_callback"
|
||||
|
||||
|
||||
def test_field_key_matches_python_semantic_scope(
|
||||
valkey_url: str,
|
||||
index_name: str,
|
||||
) -> None:
|
||||
facade: Final = _facade(valkey_url, index_name, {"semantic cache prompt": [1.0, 0.0]})
|
||||
metadata: Final = {"user_api_key": "k1"}
|
||||
expected: Final = facade.get_cache_key(
|
||||
model="gpt-4.1",
|
||||
messages=[{"role": "user", "content": "semantic cache prompt"}],
|
||||
metadata=metadata,
|
||||
)
|
||||
handle: Final = _native._CacheTestHandle.valkey_semantic(
|
||||
valkey_url,
|
||||
0.8,
|
||||
index_name,
|
||||
facade.cache,
|
||||
)
|
||||
binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve()
|
||||
binding.store(_field_request("semantic cache prompt", metadata), {"answer": "scoped"})
|
||||
client: Final = redis.Redis.from_url(valkey_url)
|
||||
documents: Final = list(client.scan_iter(f"{index_name}:*"))
|
||||
assert len(documents) == 1
|
||||
document_parts: Final = documents[0].decode().split(":")
|
||||
assert document_parts[1] == hashlib.sha256(expected.encode()).hexdigest()
|
||||
client.close()
|
||||
|
||||
|
||||
def test_field_key_isolates_tenant_scope(
|
||||
valkey_url: str,
|
||||
index_name: str,
|
||||
) -> None:
|
||||
facade: Final = _facade(valkey_url, index_name, {"semantic cache prompt": [1.0, 0.0]})
|
||||
handle: Final = _native._CacheTestHandle.valkey_semantic(
|
||||
valkey_url,
|
||||
0.8,
|
||||
index_name,
|
||||
facade.cache,
|
||||
)
|
||||
binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=handle)).resolve()
|
||||
binding.store(
|
||||
_field_request("semantic cache prompt", {"user_api_key": "k1"}),
|
||||
{"answer": "tenant one"},
|
||||
)
|
||||
assert (
|
||||
binding.lookup(_field_request("semantic cache prompt", {"user_api_key": "k2"}))
|
||||
is None
|
||||
)
|
||||
assert binding.lookup(_field_request("semantic cache prompt", {"user_api_key": "k1"})) == {
|
||||
"answer": "tenant one"
|
||||
}
|
||||
|
||||
|
||||
def test_tls_valkey_facade_falls_back_to_python(
|
||||
index_name: str,
|
||||
) -> None:
|
||||
facade: Final = Cache(
|
||||
type=LiteLLMCacheType.VALKEY_SEMANTIC,
|
||||
redis_url="rediss://127.0.0.1:6390/0",
|
||||
similarity_threshold=0.8,
|
||||
valkey_semantic_cache_index_name=index_name,
|
||||
)
|
||||
resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=facade))
|
||||
assert resolver.resolve().kind == "python_callback"
|
||||
|
||||
|
||||
async def test_ping_maps_unsupported_native_operation_to_not_implemented(
|
||||
valkey_url: str,
|
||||
index_name: str,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue