feat(cache-qdrant-semantic): add native Qdrant semantic cache backend

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Yujong Lee 2026-09-21 20:27:06 +00:00
parent fc3844e991
commit 80c0ceb5e6
8 changed files with 553 additions and 0 deletions

View file

@ -547,6 +547,49 @@ dependencies = [
"tracing",
]
[[package]]
name = "axum"
version = "0.8.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "31b698c5f9a010f6573133b09e0de5408834d0c82f8d7475a89fc1867a71cd90"
dependencies = [
"axum-core",
"bytes",
"futures-util",
"http 1.4.2",
"http-body 1.1.0",
"http-body-util",
"itoa",
"matchit",
"memchr",
"mime",
"percent-encoding",
"pin-project-lite",
"serde_core",
"sync_wrapper",
"tower",
"tower-layer",
"tower-service",
]
[[package]]
name = "axum-core"
version = "0.5.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "08c78f31d7b1291f7ee735c1c6780ccde7785daae9a9206026862dab7d8792d1"
dependencies = [
"bytes",
"futures-core",
"http 1.4.2",
"http-body 1.1.0",
"http-body-util",
"mime",
"pin-project-lite",
"sync_wrapper",
"tower-layer",
"tower-service",
]
[[package]]
name = "azure_core"
version = "1.1.0"
@ -2480,6 +2523,25 @@ dependencies = [
"tokio",
]
[[package]]
name = "litellm-cache-qdrant-semantic"
version = "0.1.0"
dependencies = [
"futures-util",
"litellm-cache",
"litellm-cache-response",
"qdrant-client",
"reqwest 0.12.28",
"rstest",
"serde",
"serde_json",
"thiserror 2.0.19",
"tokio",
"tonic",
"tonic-prost",
"uuid",
]
[[package]]
name = "litellm-cache-redis"
version = "0.1.0"
@ -2896,6 +2958,12 @@ version = "0.2.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "fc04a4c58212d57930a24bf47d3fa87485264a3a054e9c10e042eb373573ad3c"
[[package]]
name = "matchit"
version = "0.8.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "47e1ffaa40ddd1f3ed91f717a33c8c0ee23fff369e3aa8772b9605cc1d22f4c3"
[[package]]
name = "memchr"
version = "2.8.3"
@ -3509,6 +3577,27 @@ dependencies = [
"serde",
]
[[package]]
name = "qdrant-client"
version = "1.19.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "dddc19df129bad7346ebd027288621ab1ac7e52678371f906b9a8622d7aaf87e"
dependencies = [
"anyhow",
"derive_builder",
"futures",
"parking_lot",
"prost",
"prost-types",
"semver",
"serde",
"serde_json",
"thiserror 2.0.19",
"tokio",
"tonic",
"tonic-prost",
]
[[package]]
name = "quick-error"
version = "1.2.3"
@ -4866,8 +4955,12 @@ version = "0.14.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ac2a5518c70fa84342385732db33fb3f44bc4cc748936eb5833d2df34d6445ef"
dependencies = [
"async-trait",
"axum",
"base64 0.22.1",
"bytes",
"flate2",
"h2 0.4.15",
"http 1.4.2",
"http-body 1.1.0",
"http-body-util",
@ -4877,6 +4970,7 @@ dependencies = [
"percent-encoding",
"pin-project",
"rustls-native-certs",
"socket2 0.6.5",
"sync_wrapper",
"tokio",
"tokio-rustls 0.26.4",

View file

@ -31,6 +31,7 @@ litellm-cache = { path = "crates/cache" }
litellm-cache-memory = { path = "crates/cache-memory" }
litellm-cache-redis = { path = "crates/cache-redis" }
litellm-cache-response = { path = "crates/cache-response" }
litellm-cache-qdrant-semantic = { path = "crates/cache-qdrant-semantic" }
litellm-token-counter = { path = "crates/token-counter" }
litellm-token-counter-fast = { path = "crates/token-counter-fast" }
litellm-token-counter-huggingface = { path = "crates/token-counter-huggingface" }
@ -48,6 +49,8 @@ pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] }
pythonize = "0.29.0"
rand = "0.8"
reqwest = { version = "0.12", default-features = false, features = ["json", "multipart", "rustls-tls", "http2", "stream"] }
qdrant-client = { version = "1.19.0", default-features = false }
uuid = { version = "1", features = ["v4"] }
rstest = "0.26.1"
rstest_reuse = "0.7.0"
rustls = { version = "0.23", default-features = false, features = ["ring", "std", "tls12"] }

View file

@ -0,0 +1,23 @@
[package]
name = "litellm-cache-qdrant-semantic"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
futures-util.workspace = true
litellm-cache.workspace = true
qdrant-client = { workspace = true, features = ["serde"] }
reqwest.workspace = true
serde.workspace = true
serde_json.workspace = true
thiserror.workspace = true
tokio.workspace = true
uuid.workspace = true
[dev-dependencies]
litellm-cache-response.workspace = true
rstest.workspace = true
tonic = "0.14"
tonic-prost = "0.14"

View file

@ -0,0 +1,73 @@
use std::time::Duration;
use litellm_cache::Error;
use reqwest::Client;
use serde_json::Value;
use crate::Embedder;
pub struct OpenAiEmbedder {
client: Client,
api_base: String,
api_key: String,
model: String,
}
pub struct OpenAiEmbedderConfig {
pub api_base: String,
pub api_key: String,
pub model: String,
pub timeout: Option<Duration>,
}
impl OpenAiEmbedder {
pub fn new(config: OpenAiEmbedderConfig) -> Result<Self, Error> {
let mut builder = Client::builder();
if let Some(timeout) = config.timeout {
builder = builder.timeout(timeout);
}
let client = builder.build().map_err(|_| Error::Unavailable)?;
Ok(Self {
client,
api_base: config.api_base.trim_end_matches('/').to_owned(),
api_key: config.api_key,
model: config.model,
})
}
}
impl Embedder for OpenAiEmbedder {
fn model(&self) -> &str {
&self.model
}
async fn embed(&self, input: &str) -> Result<Vec<f32>, Error> {
let response = self
.client
.post(format!("{}/embeddings", self.api_base))
.bearer_auth(&self.api_key)
.json(&serde_json::json!({
"model": self.model,
"input": input,
"encoding_format": "float",
}))
.send()
.await
.map_err(|_| Error::Unavailable)?
.error_for_status()
.map_err(|_| Error::Unavailable)?;
let body: Value = response.json().await.map_err(|_| Error::Unavailable)?;
body.get("data")
.and_then(Value::as_array)
.and_then(|data| data.first())
.and_then(|item| item.get("embedding"))
.and_then(Value::as_array)
.and_then(|embedding| {
embedding
.iter()
.map(|value| value.as_f64().map(|value| value as f32))
.collect::<Option<Vec<_>>>()
})
.ok_or(Error::Unavailable)
}
}

View file

@ -0,0 +1,7 @@
mod embedder;
mod prompt;
mod semantic;
pub use embedder::{OpenAiEmbedder, OpenAiEmbedderConfig};
pub use prompt::prompt_from_messages;
pub use semantic::{Embedder, QdrantSemanticCache, QdrantSemanticConfig, Quantization};

View file

@ -0,0 +1,59 @@
use serde_json::Value;
fn search_results_text(search_results: Option<&Value>) -> String {
let Some(Value::Array(results)) = search_results else {
return String::new();
};
results
.iter()
.filter_map(Value::as_object)
.flat_map(|result| {
let source = result
.get("source")
.and_then(Value::as_str)
.map(str::to_owned);
let title = result
.get("title")
.and_then(Value::as_str)
.map(str::to_owned);
let content = result
.get("content")
.and_then(Value::as_array)
.into_iter()
.flatten()
.filter_map(Value::as_object)
.filter_map(|block| block.get("text").and_then(Value::as_str).map(str::to_owned));
let citations = result
.get("citations")
.filter(|value| !value.is_null())
.map(|value| serde_json::to_string(value).unwrap_or_default());
source
.into_iter()
.chain(title)
.chain(content)
.chain(citations)
})
.collect()
}
pub fn prompt_from_messages(messages: &[Value]) -> String {
messages
.iter()
.filter_map(Value::as_object)
.map(|message| {
let content = match message.get("content") {
Some(Value::String(content)) => content.clone(),
Some(Value::Array(parts)) => parts
.iter()
.filter_map(Value::as_object)
.filter_map(|part| part.get("text").and_then(Value::as_str))
.collect(),
_ => String::new(),
};
format!(
"{content}{}",
search_results_text(message.get("search_results"))
)
})
.collect()
}

View file

@ -0,0 +1,256 @@
use std::future::Future;
use futures_util::future::try_join_all;
use litellm_cache::{BaseCache, CacheCodec, CacheConnectionResult, Error, SemanticCacheContext};
use qdrant_client::{
Payload, Qdrant,
qdrant::{
BinaryQuantizationBuilder, CompressionRatio, Condition, CreateCollectionBuilder,
CreateFieldIndexCollectionBuilder, Distance, FieldType, Filter, PointStruct,
ProductQuantizationBuilder, QuantizationSearchParamsBuilder, ScalarQuantizationBuilder,
SearchParamsBuilder, SearchPointsBuilder, UpsertPointsBuilder, VectorParamsBuilder,
},
};
use serde_json::{Map, Value, json};
use uuid::Uuid;
use crate::prompt_from_messages;
pub trait Embedder: Send + Sync + 'static {
fn model(&self) -> &str;
fn embed(&self, input: &str) -> impl Future<Output = Result<Vec<f32>, Error>> + Send;
}
#[derive(Clone, Debug, PartialEq)]
pub enum Quantization {
Binary,
Scalar,
Product,
}
pub struct QdrantSemanticConfig {
pub collection_name: String,
pub similarity_threshold: f64,
pub vector_size: u64,
pub quantization: Quantization,
}
pub struct QdrantSemanticCache<E: Embedder, C: CacheCodec> {
client: Qdrant,
embedder: E,
codec: C,
config: QdrantSemanticConfig,
runtime: tokio::runtime::Handle,
}
impl<E: Embedder, C: CacheCodec> QdrantSemanticCache<E, C> {
pub async fn connect(
client: Qdrant,
embedder: E,
codec: C,
config: QdrantSemanticConfig,
runtime: tokio::runtime::Handle,
) -> Result<Self, Error> {
let exists = client
.collection_exists(config.collection_name.clone())
.await
.map_err(|_| Error::Unavailable)?;
if !exists {
client
.create_collection(
CreateCollectionBuilder::new(config.collection_name.clone())
.vectors_config(VectorParamsBuilder::new(
config.vector_size,
Distance::Cosine,
))
.quantization_config(quantization(&config.quantization)),
)
.await
.map_err(|_| Error::Unavailable)?;
}
let _ = client
.create_field_index(CreateFieldIndexCollectionBuilder::new(
config.collection_name.clone(),
"litellm_cache_key".to_owned(),
FieldType::Keyword,
))
.await;
Ok(Self {
client,
embedder,
codec,
config,
runtime,
})
}
pub fn collection_name(&self) -> &str {
&self.config.collection_name
}
pub fn similarity_threshold(&self) -> f64 {
self.config.similarity_threshold
}
pub fn vector_size(&self) -> u64 {
self.config.vector_size
}
pub fn embedder(&self) -> &E {
&self.embedder
}
fn prompt(context: &SemanticCacheContext) -> Result<String, Error> {
if context.messages.is_empty() {
return Err(Error::MissingPrompt);
}
Ok(prompt_from_messages(&context.messages))
}
async fn set(
&self,
key: &str,
value: C::Value,
context: &SemanticCacheContext,
) -> Result<(), Error> {
let prompt = Self::prompt(context)?;
let vector = self.embedder.embed(&prompt).await?;
let response =
String::from_utf8(self.codec.encode(&value)?).map_err(|_| Error::InvalidEntry)?;
let payload = Payload::try_from(json!({
"litellm_cache_key": key,
"text": prompt,
"response": response,
}))
.map_err(|_| Error::InvalidEntry)?;
self.client
.upsert_points(UpsertPointsBuilder::new(
self.collection_name(),
vec![PointStruct::new(
Uuid::new_v4().to_string(),
vector,
payload,
)],
))
.await
.map_err(|_| Error::Unavailable)?;
Ok(())
}
async fn get(
&self,
key: &str,
context: &SemanticCacheContext,
) -> Result<Option<C::Value>, Error> {
let prompt = Self::prompt(context)?;
let vector = self.embedder.embed(&prompt).await?;
let result = self
.client
.search_points(
SearchPointsBuilder::new(self.collection_name(), vector, 1)
.with_payload(true)
.filter(Filter::must([Condition::matches(
"litellm_cache_key",
key.to_owned(),
)]))
.params(
SearchParamsBuilder::default().quantization(
QuantizationSearchParamsBuilder::default()
.ignore(false)
.rescore(true)
.oversampling(3.0),
),
),
)
.await
.map_err(|_| Error::Unavailable)?;
let Some(point) = result.result.into_iter().next() else {
return Ok(None);
};
if f64::from(point.score) < self.config.similarity_threshold {
return Ok(None);
}
let payload: Map<String, Value> = Payload::from(point.payload).into();
if payload.get("litellm_cache_key").and_then(Value::as_str) != Some(key) {
return Ok(None);
}
let response = payload
.get("response")
.and_then(Value::as_str)
.ok_or(Error::InvalidEntry)?;
self.codec.decode(response.as_bytes()).map(Some)
}
}
fn quantization(value: &Quantization) -> qdrant_client::qdrant::quantization_config::Quantization {
match value {
Quantization::Binary => BinaryQuantizationBuilder::new(false).into(),
Quantization::Scalar => ScalarQuantizationBuilder::default()
.quantile(0.99)
.always_ram(false)
.into(),
Quantization::Product => ProductQuantizationBuilder::new(CompressionRatio::X16.into())
.always_ram(false)
.into(),
}
}
impl<E: Embedder, C: CacheCodec> BaseCache for QdrantSemanticCache<E, C> {
type Value = C::Value;
type Context = SemanticCacheContext;
fn get_ttl(&self, _: &Self::Context) -> Option<std::time::Duration> {
None
}
fn set_cache(
&self,
key: &str,
value: Self::Value,
context: &Self::Context,
) -> Result<(), Error> {
self.runtime.block_on(self.set(key, value, context))
}
fn get_cache(&self, key: &str, context: &Self::Context) -> Result<Option<Self::Value>, Error> {
self.runtime.block_on(self.get(key, context))
}
async fn async_set_cache(
&self,
key: &str,
value: Self::Value,
context: Self::Context,
) -> Result<(), Error> {
self.set(key, value, &context).await
}
async fn async_get_cache(
&self,
key: &str,
context: &Self::Context,
) -> Result<Option<Self::Value>, Error> {
self.get(key, context).await
}
async fn async_set_cache_pipeline(
&self,
entries: Vec<(String, Self::Value)>,
context: Self::Context,
) -> Result<(), Error> {
try_join_all(entries.into_iter().map(|(key, value)| {
let context = context.clone();
async move { self.async_set_cache(&key, value, context).await }
}))
.await
.map(|_| ())
}
async fn disconnect(&self) -> Result<(), Error> {
Ok(())
}
async fn test_connection(&self) -> Result<CacheConnectionResult, Error> {
Err(Error::UnsupportedOperation)
}
}

View file

@ -0,0 +1,38 @@
use litellm_cache_qdrant_semantic::prompt_from_messages;
use serde_json::json;
#[test]
fn prompt_matches_python_message_content_rules() {
let messages = vec![
json!({"role": "user", "content": "hello"}),
json!({
"role": "user",
"content": [
{"type": "text", "text": "world"},
{"type": "image_url", "image_url": {"url": "ignored"}},
{"type": "text", "text": "!"},
],
}),
];
assert_eq!(prompt_from_messages(&messages), "helloworld!");
}
#[test]
fn prompt_includes_search_result_text_and_compact_citations() {
let messages = vec![json!({
"role": "tool",
"content": null,
"search_results": [{
"source": "source",
"title": "title",
"content": [{"text": "body"}],
"citations": {"page": 1, "section": "intro"},
}],
})];
assert_eq!(
prompt_from_messages(&messages),
r#"sourcetitlebody{"page":1,"section":"intro"}"#
);
}