use super::{
error::QueryConversionError,
operator::{
Filter, GroupBy, KnnBatch, KnnProjection, Limit, Projection, Rank, Scan, ScanToProtoError,
Select,
},
};
use crate::{
chroma_proto,
operator::{Key, RankExpr},
validators::{validate_group_by, validate_rank, validate_search_payload},
Where,
};
use serde::{Deserialize, Serialize};
use thiserror::Error;
#[cfg(feature = "utoipa")]
use utoipa::{
openapi::{
schema::{Schema, SchemaType},
ArrayBuilder, Object, ObjectBuilder, RefOr, Type,
},
PartialSchema,
};
use validator::Validate;
#[derive(Error, Debug)]
pub enum PlanToProtoError {
#[error("Failed to convert scan to proto: {0}")]
Scan(#[from] ScanToProtoError),
}
#[derive(Clone)]
pub struct Count {
pub scan: Scan,
pub read_level: ReadLevel,
}
impl TryFrom<chroma_proto::CountPlan> for Count {
type Error = QueryConversionError;
fn try_from(value: chroma_proto::CountPlan) -> Result<Self, Self::Error> {
let read_level = value.read_level().into();
Ok(Self {
scan: value
.scan
.ok_or(QueryConversionError::field("scan"))?
.try_into()?,
read_level,
})
}
}
impl TryFrom<Count> for chroma_proto::CountPlan {
type Error = PlanToProtoError;
fn try_from(value: Count) -> Result<Self, Self::Error> {
Ok(Self {
scan: Some(value.scan.try_into()?),
read_level: chroma_proto::ReadLevel::from(value.read_level).into(),
})
}
}
#[derive(Clone, Debug)]
pub struct Get {
pub scan: Scan,
pub filter: Filter,
pub limit: Limit,
pub proj: Projection,
}
impl TryFrom<chroma_proto::GetPlan> for Get {
type Error = QueryConversionError;
fn try_from(value: chroma_proto::GetPlan) -> Result<Self, Self::Error> {
Ok(Self {
scan: value
.scan
.ok_or(QueryConversionError::field("scan"))?
.try_into()?,
filter: value
.filter
.ok_or(QueryConversionError::field("filter"))?
.try_into()?,
limit: value
.limit
.ok_or(QueryConversionError::field("limit"))?
.into(),
proj: value
.projection
.ok_or(QueryConversionError::field("projection"))?
.into(),
})
}
}
impl TryFrom<Get> for chroma_proto::GetPlan {
type Error = QueryConversionError;
fn try_from(value: Get) -> Result<Self, Self::Error> {
Ok(Self {
scan: Some(value.scan.try_into()?),
filter: Some(value.filter.try_into()?),
limit: Some(value.limit.into()),
projection: Some(value.proj.into()),
})
}
}
#[derive(Clone, Debug)]
pub struct Knn {
pub scan: Scan,
pub filter: Filter,
pub knn: KnnBatch,
pub proj: KnnProjection,
}
impl TryFrom<chroma_proto::KnnPlan> for Knn {
type Error = QueryConversionError;
fn try_from(value: chroma_proto::KnnPlan) -> Result<Self, Self::Error> {
Ok(Self {
scan: value
.scan
.ok_or(QueryConversionError::field("scan"))?
.try_into()?,
filter: value
.filter
.ok_or(QueryConversionError::field("filter"))?
.try_into()?,
knn: value
.knn
.ok_or(QueryConversionError::field("knn"))?
.try_into()?,
proj: value
.projection
.ok_or(QueryConversionError::field("projection"))?
.try_into()?,
})
}
}
impl TryFrom<Knn> for chroma_proto::KnnPlan {
type Error = QueryConversionError;
fn try_from(value: Knn) -> Result<Self, Self::Error> {
Ok(Self {
scan: Some(value.scan.try_into()?),
filter: Some(value.filter.try_into()?),
knn: Some(value.knn.try_into()?),
projection: Some(value.proj.into()),
})
}
}
#[derive(Clone, Debug, Default, Deserialize, Serialize, Validate)]
#[validate(schema(function = "validate_search_payload"))]
pub struct SearchPayload {
#[serde(default)]
pub filter: Filter,
#[serde(default)]
#[validate(custom(function = "validate_rank"))]
pub rank: Rank,
#[serde(default)]
#[validate(custom(function = "validate_group_by"))]
pub group_by: GroupBy,
#[serde(default)]
pub limit: Limit,
#[serde(default)]
pub select: Select,
}
impl SearchPayload {
pub fn limit(mut self, limit: Option<u32>, offset: u32) -> Self {
self.limit.limit = limit;
self.limit.offset = offset;
self
}
pub fn rank(mut self, expr: RankExpr) -> Self {
self.rank.expr = Some(expr);
self
}
pub fn select<I, T>(mut self, keys: I) -> Self
where
I: IntoIterator<Item = T>,
T: Into<Key>,
{
self.select.keys = keys.into_iter().map(Into::into).collect();
self
}
pub fn r#where(mut self, r#where: Where) -> Self {
self.filter.where_clause = Some(r#where);
self
}
pub fn group_by(mut self, group_by: GroupBy) -> Self {
self.group_by = group_by;
self
}
}
#[cfg(feature = "utoipa")]
impl PartialSchema for SearchPayload {
fn schema() -> RefOr<Schema> {
RefOr::T(Schema::Object(
ObjectBuilder::new()
.schema_type(SchemaType::Type(Type::Object))
.property(
"filter",
ObjectBuilder::new()
.schema_type(SchemaType::Type(Type::Object))
.property(
"query_ids",
ArrayBuilder::new()
.items(Object::with_type(SchemaType::Type(Type::String))),
)
.property(
"where_clause",
Object::with_type(SchemaType::Type(Type::Object)),
),
)
.property("rank", Object::with_type(SchemaType::Type(Type::Object)))
.property(
"group_by",
ObjectBuilder::new()
.schema_type(SchemaType::Type(Type::Object))
.property(
"keys",
ArrayBuilder::new()
.items(Object::with_type(SchemaType::Type(Type::String))),
)
.property(
"aggregate",
Object::with_type(SchemaType::Type(Type::Object)),
),
)
.property(
"limit",
ObjectBuilder::new()
.schema_type(SchemaType::Type(Type::Object))
.property("offset", Object::with_type(SchemaType::Type(Type::Integer)))
.property("limit", Object::with_type(SchemaType::Type(Type::Integer))),
)
.property(
"select",
ObjectBuilder::new()
.schema_type(SchemaType::Type(Type::Object))
.property(
"keys",
ArrayBuilder::new()
.items(Object::with_type(SchemaType::Type(Type::String))),
),
)
.build(),
))
}
}
#[cfg(feature = "utoipa")]
impl utoipa::ToSchema for SearchPayload {}
impl TryFrom<chroma_proto::SearchPayload> for SearchPayload {
type Error = QueryConversionError;
fn try_from(value: chroma_proto::SearchPayload) -> Result<Self, Self::Error> {
Ok(Self {
filter: value
.filter
.ok_or(QueryConversionError::field("filter"))?
.try_into()?,
rank: value
.rank
.ok_or(QueryConversionError::field("rank"))?
.try_into()?,
group_by: value
.group_by
.map(TryInto::try_into)
.transpose()?
.unwrap_or_default(),
limit: value
.limit
.ok_or(QueryConversionError::field("limit"))?
.into(),
select: value
.select
.ok_or(QueryConversionError::field("select"))?
.try_into()?,
})
}
}
impl TryFrom<SearchPayload> for chroma_proto::SearchPayload {
type Error = QueryConversionError;
fn try_from(value: SearchPayload) -> Result<Self, Self::Error> {
Ok(Self {
filter: Some(value.filter.try_into()?),
rank: Some(value.rank.try_into()?),
group_by: Some(value.group_by.try_into()?),
limit: Some(value.limit.into()),
select: Some(value.select.try_into()?),
})
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[cfg_attr(feature = "utoipa", derive(utoipa::ToSchema))]
pub enum ReadLevel {
#[default]
IndexAndWal,
IndexOnly,
IndexAndBoundedWal,
}
impl From<chroma_proto::ReadLevel> for ReadLevel {
fn from(value: chroma_proto::ReadLevel) -> Self {
match value {
chroma_proto::ReadLevel::IndexAndWal => ReadLevel::IndexAndWal,
chroma_proto::ReadLevel::IndexOnly => ReadLevel::IndexOnly,
chroma_proto::ReadLevel::IndexAndBoundedWal => ReadLevel::IndexAndBoundedWal,
}
}
}
impl From<ReadLevel> for chroma_proto::ReadLevel {
fn from(value: ReadLevel) -> Self {
match value {
ReadLevel::IndexAndWal => chroma_proto::ReadLevel::IndexAndWal,
ReadLevel::IndexOnly => chroma_proto::ReadLevel::IndexOnly,
ReadLevel::IndexAndBoundedWal => chroma_proto::ReadLevel::IndexAndBoundedWal,
}
}
}
#[derive(Clone, Debug)]
pub struct Search {
pub scan: Scan,
pub payloads: Vec<SearchPayload>,
pub read_level: ReadLevel,
}
impl TryFrom<chroma_proto::SearchPlan> for Search {
type Error = QueryConversionError;
fn try_from(value: chroma_proto::SearchPlan) -> Result<Self, Self::Error> {
let read_level = value.read_level().into();
Ok(Self {
scan: value
.scan
.ok_or(QueryConversionError::field("scan"))?
.try_into()?,
payloads: value
.payloads
.into_iter()
.map(TryInto::try_into)
.collect::<Result<Vec<_>, _>>()?,
read_level,
})
}
}
impl TryFrom<Search> for chroma_proto::SearchPlan {
type Error = QueryConversionError;
fn try_from(value: Search) -> Result<Self, Self::Error> {
Ok(Self {
scan: Some(value.scan.try_into()?),
payloads: value
.payloads
.into_iter()
.map(TryInto::try_into)
.collect::<Result<Vec<_>, _>>()?,
read_level: chroma_proto::ReadLevel::from(value.read_level).into(),
})
}
}