128 lines
4.1 KiB
Rust
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)
|
|
}
|
|
}
|