use crate::chroma::create_client;
use crate::error::RuChatError;
use anyhow::Result;
use chromadb::collection::{ChromaCollection, GetOptions, GetResult, QueryOptions, QueryResult};
use clap::Parser;
use serde_json::json;
#[derive(Parser, Debug, Clone, PartialEq)]
pub struct SimilarityArgs {
#[clap(short, long)]
pub(crate) query: String,
#[clap(short, long, default_value = "1")]
pub(crate) count: usize,
#[clap(short, long, default_value = "5")]
pub(crate) similarity_count: usize,
#[clap(short, long, default_value = "default")]
pub(crate) collection: String,
#[clap(short, long)]
pub(crate) metadata: Option<String>,
#[clap(short = 'C', long, default_value = "http://localhost:8000")]
pub(crate) chroma_server: String,
#[clap(short = 'd', long, default_value = "default")]
pub(crate) chroma_database: String,
#[clap(short = 't', long)]
pub(crate) chroma_token: Option<String>,
}
pub(crate) async fn similarity_search(args: &SimilarityArgs) -> Result<(), RuChatError> {
let client = create_client(
args.chroma_token.as_deref(),
&args.chroma_server,
&args.chroma_database,
)
.await?;
let collection: ChromaCollection = client
.get_or_create_collection(&args.collection, None)
.await?;
let metadata = args.metadata.as_deref().map(|md| md.into());
let where_document = json!({
"$contains": args.query.as_str()
});
let get_query = GetOptions {
ids: vec![],
where_metadata: metadata,
limit: Some(args.count),
offset: None,
where_document: Some(where_document),
include: Some(vec!["documents".into(), "embeddings".into()]),
};
let get_result: GetResult = collection.get(get_query).await?;
let query = QueryOptions {
query_texts: None,
query_embeddings: get_result
.embeddings
.map(|embeddings| embeddings.into_iter().flatten().collect()),
where_metadata: None,
where_document: None,
n_results: Some(args.similarity_count),
include: None,
};
let query_result: QueryResult = collection.query(query, None).await?;
println!("Query result: {:?}", query_result);
Ok(())
}