Files
mcp-memory/server/src/vector_db.rs
T

128 lines
4.1 KiB
Rust

use qdrant_client::qdrant::{CreateCollectionBuilder, Distance, PointStruct, VectorParamsBuilder, UpsertPointsBuilder};
use qdrant_client::Qdrant;
use std::sync::Arc;
use std::error::Error;
use uuid::Uuid;
use tracing::info;
use serde::{Deserialize, Serialize};
#[derive(Clone)]
pub struct VectorDB {
client: Arc<Qdrant>,
collection_name: String,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct VectorSearchResult {
pub id: String,
pub doc_type: String,
pub text: String,
pub score: f32,
}
impl VectorDB {
pub async fn new(url: &str, collection_name: &str) -> Result<Self, Box<dyn Error + Send + Sync>> {
let client = Qdrant::from_url(url).build()?;
let db = Self {
client: Arc::new(client),
collection_name: collection_name.to_string(),
};
db.init_collection().await?;
Ok(db)
}
async fn init_collection(&self) -> Result<(), Box<dyn Error + Send + Sync>> {
// Fastembed AllMiniLML6V2 uses 384 dimensions
let vector_params = VectorParamsBuilder::new(384, Distance::Cosine).build();
let collection_exists = self.client.collection_exists(&self.collection_name).await?;
if !collection_exists {
self.client
.create_collection(
CreateCollectionBuilder::new(&self.collection_name)
.vectors_config(vector_params)
)
.await?;
info!("Created Qdrant collection: {}", self.collection_name);
} else {
info!("Qdrant collection {} already exists", self.collection_name);
}
Ok(())
}
pub async fn index_document(
&self,
id: &str,
doc_type: &str,
text: &str,
vector: Vec<f32>,
) -> Result<(), Box<dyn Error + Send + Sync>> {
let point_id = match Uuid::parse_str(id) {
Ok(u) => u.to_string(),
Err(_) => {
// If it's not a valid UUID, let's create a deterministic UUID based on the string
let namespace = Uuid::NAMESPACE_OID;
Uuid::new_v5(&namespace, id.as_bytes()).to_string()
}
};
let mut payload: std::collections::HashMap<String, serde_json::Value> = std::collections::HashMap::new();
payload.insert("doc_type".to_string(), serde_json::Value::String(doc_type.to_string()));
payload.insert("text".to_string(), serde_json::Value::String(text.to_string()));
payload.insert("original_id".to_string(), serde_json::Value::String(id.to_string()));
let point = PointStruct::new(point_id, vector, payload);
self.client
.upsert_points(UpsertPointsBuilder::new(&self.collection_name, vec![point]))
.await?;
Ok(())
}
pub async fn search(
&self,
query_vector: Vec<f32>,
limit: u64,
) -> Result<Vec<VectorSearchResult>, Box<dyn Error + Send + Sync>> {
use qdrant_client::qdrant::SearchPointsBuilder;
let search_result = self.client
.search_points(
SearchPointsBuilder::new(&self.collection_name, query_vector, limit)
.with_payload(true)
)
.await?;
let mut results = Vec::new();
for point in search_result.result {
let id = point.payload.get("original_id")
.and_then(|v| v.as_str())
.map(|s| s.to_string())
.unwrap_or_default();
let doc_type = point.payload.get("doc_type")
.and_then(|v| v.as_str())
.map(|s| s.to_string())
.unwrap_or_default();
let text = point.payload.get("text")
.and_then(|v| v.as_str())
.map(|s| s.to_string())
.unwrap_or_default();
results.push(VectorSearchResult {
id,
doc_type,
text,
score: point.score,
});
}
Ok(results)
}
}