use anyhow::Context;
use anyhow::Result;
use qdrant_client::qdrant::{SearchPoints, WithPayloadSelector};
use super::{VectorDb, utils};
#[expect(
async_fn_in_trait,
reason = "async trait method is required for the db interfaces"
)]
pub trait VectorSearchExt {
async fn search(
&self,
vector: &[f32],
limit: usize,
repo_name: Option<&str>,
) -> Result<Vec<serde_json::Value>>;
}
impl VectorSearchExt for VectorDb {
async fn search(
&self,
vector: &[f32],
limit: usize,
repo_name: Option<&str>,
) -> Result<Vec<serde_json::Value>> {
let filter = repo_name.map(|repo| qdrant_client::qdrant::Filter {
must: vec![qdrant_client::qdrant::Condition {
condition_one_of: Some(qdrant_client::qdrant::condition::ConditionOneOf::Field(
qdrant_client::qdrant::FieldCondition {
key: "repo_name".to_string(),
r#match: Some(qdrant_client::qdrant::Match {
match_value: Some(qdrant_client::qdrant::r#match::MatchValue::Keyword(
repo.to_string(),
)),
}),
..Default::default()
},
)),
}],
..Default::default()
});
let search_request = SearchPoints {
collection_name: self.collection.clone(),
vector: vector.to_vec(),
limit: limit as u64,
with_payload: Some(WithPayloadSelector {
selector_options: Some(
qdrant_client::qdrant::with_payload_selector::SelectorOptions::Enable(true),
),
}),
filter,
..Default::default()
};
let search_result = self
.client
.search_points(search_request)
.await
.context("Failed to search Qdrant")?;
let results = search_result
.result
.into_iter()
.filter_map(|scored_point| {
if !scored_point.payload.is_empty() {
let mut json_obj = serde_json::json!({});
for (key, value) in scored_point.payload {
json_obj[&key] = utils::qdrant_value_to_json(&value);
}
Some(json_obj)
} else {
None
}
})
.collect();
Ok(results)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::db::vector::connection::VectorConnectExt;
#[ignore = "requires local Qdrant instance running on http://localhost:6334"]
#[tokio::test]
async fn test_search_vector() {
let vector_db = VectorDb::connect("http://localhost:6334", "test_collection_search", 384)
.await
.expect("Failed to connect to Qdrant");
let query_vector = vec![0.5; 384];
let result = vector_db.search(&query_vector, 10, None).await;
assert!(result.is_ok());
let results = result.unwrap();
assert!(results.is_empty() || !results.is_empty()); }
#[ignore = "requires local Qdrant instance running on http://localhost:6334"]
#[tokio::test]
async fn test_search_vector_with_repo_filter() {
let vector_db = VectorDb::connect(
"http://localhost:6334",
"test_collection_search_filter",
384,
)
.await
.expect("Failed to connect to Qdrant");
let query_vector = vec![0.5; 384];
let result = vector_db.search(&query_vector, 10, Some("test-repo")).await;
assert!(result.is_ok());
let results = result.unwrap();
assert!(results.is_empty() || !results.is_empty()); }
#[ignore = "requires local Qdrant instance running on http://localhost:6334"]
#[tokio::test]
async fn test_search_zero_limit() {
let vector_db =
VectorDb::connect("http://localhost:6334", "test_collection_search_zero", 384)
.await
.expect("Failed to connect to Qdrant");
let query_vector = vec![0.5; 384];
let result = vector_db.search(&query_vector, 0, None).await;
assert!(result.is_ok());
}
#[ignore = "requires local Qdrant instance running on http://localhost:6334"]
#[tokio::test]
async fn test_search_large_limit() {
let vector_db =
VectorDb::connect("http://localhost:6334", "test_collection_search_large", 384)
.await
.expect("Failed to connect to Qdrant");
let query_vector = vec![0.5; 384];
let result = vector_db.search(&query_vector, 1000, None).await;
assert!(result.is_ok());
}
}