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:
Yujong Lee 2026-09-21 21:22:01 +00:00
parent 0d09e9d892
commit 4509eb9914
5 changed files with 286 additions and 3 deletions

View file

@ -2700,6 +2700,7 @@ dependencies = [
"rstest",
"serde",
"serde_json",
"sha2 0.10.9",
"tokio",
"tokio-tungstenite",
]

View file

@ -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]]

View file

@ -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();

View file

@ -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());
}
}

View file

@ -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,