pub mod cols {
pub const DELETED: &str = "_deleted";
pub const VERSION: &str = "_version";
}
#[derive(Debug, Clone)]
pub enum Scalar {
Null,
Bool(bool),
Int(i64),
UInt(u64),
Float(f64),
Str(String),
}
impl PartialEq for Scalar {
fn eq(&self, other: &Self) -> bool {
match (self, other) {
(Scalar::Null, Scalar::Null) => true,
(Scalar::Bool(a), Scalar::Bool(b)) => a == b,
(Scalar::Int(a), Scalar::Int(b)) => a == b,
(Scalar::UInt(a), Scalar::UInt(b)) => a == b,
(Scalar::Float(a), Scalar::Float(b)) => a.to_bits() == b.to_bits(),
(Scalar::Str(a), Scalar::Str(b)) => a == b,
_ => false,
}
}
}
impl Eq for Scalar {}
impl Scalar {
pub fn from_value(value: &uni_common::Value) -> Option<Scalar> {
match value {
uni_common::Value::String(s) => Some(Scalar::Str(s.clone())),
uni_common::Value::Int(n) => Some(Scalar::Int(*n)),
uni_common::Value::Float(f) => Some(Scalar::Float(*f)),
uni_common::Value::Bool(b) => Some(Scalar::Bool(*b)),
_ => None,
}
}
}
impl std::hash::Hash for Scalar {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
std::mem::discriminant(self).hash(state);
match self {
Scalar::Null => {}
Scalar::Bool(b) => b.hash(state),
Scalar::Int(i) => i.hash(state),
Scalar::UInt(u) => u.hash(state),
Scalar::Float(f) => f.to_bits().hash(state),
Scalar::Str(s) => s.hash(state),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum CmpOp {
Eq,
NotEq,
Lt,
LtEq,
Gt,
GtEq,
}
impl CmpOp {
fn as_sql(self) -> &'static str {
match self {
CmpOp::Eq => "=",
CmpOp::NotEq => "!=",
CmpOp::Lt => "<",
CmpOp::LtEq => "<=",
CmpOp::Gt => ">",
CmpOp::GtEq => ">=",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum StringMatchKind {
Contains,
StartsWith,
EndsWith,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ToSqlError {
Unsupported(String),
}
impl std::fmt::Display for ToSqlError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ToSqlError::Unsupported(why) => write!(f, "cannot render filter to SQL: {why}"),
}
}
}
impl std::error::Error for ToSqlError {}
#[derive(Debug, Clone, PartialEq)]
pub enum FilterExpr {
Literal(bool),
And(Vec<FilterExpr>),
Or(Vec<FilterExpr>),
Not(Box<FilterExpr>),
Compare {
column: String,
op: CmpOp,
value: Scalar,
},
In { column: String, values: Vec<Scalar> },
ArrayContains { column: String, value: Scalar },
StringMatch {
column: String,
kind: StringMatchKind,
pattern: String,
},
IsNull(String),
IsNotNull(String),
Raw(String),
}
impl Default for FilterExpr {
fn default() -> Self {
FilterExpr::Literal(true)
}
}
impl FilterExpr {
pub fn not_deleted() -> Self {
FilterExpr::Compare {
column: cols::DELETED.to_string(),
op: CmpOp::Eq,
value: Scalar::Bool(false),
}
}
pub fn version_at_most(v: u64) -> Self {
FilterExpr::Compare {
column: cols::VERSION.to_string(),
op: CmpOp::LtEq,
value: Scalar::UInt(v),
}
}
pub fn compare(column: impl Into<String>, op: CmpOp, value: Scalar) -> Self {
FilterExpr::Compare {
column: column.into(),
op,
value,
}
}
pub fn equals(column: impl Into<String>, value: Scalar) -> Self {
Self::compare(column, CmpOp::Eq, value)
}
pub fn one_of(column: impl Into<String>, values: impl IntoIterator<Item = Scalar>) -> Self {
FilterExpr::In {
column: column.into(),
values: values.into_iter().collect(),
}
}
pub fn array_contains(column: impl Into<String>, value: Scalar) -> Self {
FilterExpr::ArrayContains {
column: column.into(),
value,
}
}
pub fn all(parts: impl IntoIterator<Item = FilterExpr>) -> Self {
let kept: Vec<_> = parts
.into_iter()
.filter(|p| !matches!(p, FilterExpr::Literal(true)))
.collect();
match kept.len() {
0 => FilterExpr::Literal(true),
1 => kept.into_iter().next().expect("len checked"),
_ => FilterExpr::And(kept),
}
}
pub fn any_of(parts: impl IntoIterator<Item = FilterExpr>) -> Self {
let kept: Vec<_> = parts
.into_iter()
.filter(|p| !matches!(p, FilterExpr::Literal(false)))
.collect();
match kept.len() {
0 => FilterExpr::Literal(false),
1 => kept.into_iter().next().expect("len checked"),
_ => FilterExpr::Or(kept),
}
}
pub fn negate(inner: FilterExpr) -> Self {
FilterExpr::Not(Box::new(inner))
}
pub fn sql_pushable(&self) -> FilterExpr {
if self.to_sql().is_ok() {
return self.clone();
}
match self {
FilterExpr::And(parts) => FilterExpr::all(parts.iter().map(Self::sql_pushable)),
_ => FilterExpr::Literal(true),
}
}
pub fn is_trivially_true(&self) -> bool {
match self {
FilterExpr::Literal(b) => *b,
FilterExpr::And(parts) => parts.iter().all(Self::is_trivially_true),
FilterExpr::Or(parts) => parts.iter().any(Self::is_trivially_true),
_ => false,
}
}
pub fn to_sql(&self) -> Result<String, ToSqlError> {
match self {
FilterExpr::Literal(true) => Ok("true".to_string()),
FilterExpr::Literal(false) => Ok("false".to_string()),
FilterExpr::And(parts) => {
if parts.is_empty() {
return Ok("true".to_string());
}
let rendered: Result<Vec<_>, _> =
parts.iter().map(|p| p.to_sql().map(paren)).collect();
Ok(rendered?.join(" AND "))
}
FilterExpr::Or(parts) => {
if parts.is_empty() {
return Ok("false".to_string());
}
let rendered: Result<Vec<_>, _> =
parts.iter().map(|p| p.to_sql().map(paren)).collect();
Ok(rendered?.join(" OR "))
}
FilterExpr::Not(inner) => Ok(format!("NOT {}", paren(inner.to_sql()?))),
FilterExpr::Compare { column, op, value } => Ok(format!(
"{} {} {}",
column,
op.as_sql(),
scalar_to_sql(value)
)),
FilterExpr::In { column, values } => {
if values.is_empty() {
return Ok("false".to_string());
}
let (nulls, non_nulls): (Vec<_>, Vec<_>) =
values.iter().partition(|v| matches!(v, Scalar::Null));
if nulls.is_empty() {
let items: Vec<_> = non_nulls.iter().map(|v| scalar_to_sql(v)).collect();
return Ok(format!("{} IN ({})", column, items.join(", ")));
}
if non_nulls.is_empty() {
return Ok(SQL_UNKNOWN.to_string());
}
let items: Vec<_> = non_nulls.iter().map(|v| scalar_to_sql(v)).collect();
Ok(format!(
"({} IN ({}) OR {})",
column,
items.join(", "),
SQL_UNKNOWN
))
}
FilterExpr::ArrayContains { column, value } => Ok(format!(
"array_contains({}, {})",
column,
scalar_to_sql(value)
)),
FilterExpr::StringMatch {
column,
kind,
pattern,
} => {
if pattern.contains('%') || pattern.contains('_') {
return Err(ToSqlError::Unsupported(format!(
"LIKE pattern {pattern:?} contains a SQL wildcard and this \
dialect has no ESCAPE clause"
)));
}
let escaped = pattern.replace('\'', "''");
Ok(match kind {
StringMatchKind::Contains => format!("{column} LIKE '%{escaped}%'"),
StringMatchKind::StartsWith => format!("{column} LIKE '{escaped}%'"),
StringMatchKind::EndsWith => format!("{column} LIKE '%{escaped}'"),
})
}
FilterExpr::IsNull(column) => Ok(format!("{column} IS NULL")),
FilterExpr::IsNotNull(column) => Ok(format!("{column} IS NOT NULL")),
FilterExpr::Raw(s) => Ok(s.clone()),
}
}
}
fn paren(s: String) -> String {
if s == "true" || s == "false" {
s
} else {
format!("({s})")
}
}
const SQL_UNKNOWN: &str = "CAST(NULL AS BOOLEAN)";
fn scalar_to_sql(v: &Scalar) -> String {
match v {
Scalar::Null => "NULL".to_string(),
Scalar::Bool(b) => b.to_string(),
Scalar::Int(i) => i.to_string(),
Scalar::UInt(u) => u.to_string(),
Scalar::Float(f) => f.to_string(),
Scalar::Str(s) => format!("'{}'", s.replace('\'', "''")),
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct VectorQueryOpts {
pub nprobes: Option<usize>,
pub refine_factor: Option<u32>,
pub ef: Option<usize>,
}
#[derive(Debug, Clone)]
pub enum ColumnProjection {
Columns(Vec<String>),
All,
}
#[derive(Debug, Clone)]
pub struct ScanRequest {
pub table_name: String,
pub columns: ColumnProjection,
pub filter: FilterExpr,
pub limit: Option<usize>,
pub branch: Option<String>,
}
impl ScanRequest {
pub fn all(table_name: impl Into<String>) -> Self {
Self {
table_name: table_name.into(),
columns: ColumnProjection::All,
filter: FilterExpr::Literal(true),
limit: None,
branch: None,
}
}
pub fn with_columns(mut self, columns: Vec<String>) -> Self {
self.columns = ColumnProjection::Columns(columns);
self
}
pub fn with_filter(mut self, filter: FilterExpr) -> Self {
self.filter = filter;
self
}
pub fn with_limit(mut self, limit: usize) -> Self {
self.limit = Some(limit);
self
}
pub fn with_branch(mut self, branch: impl Into<String>) -> Self {
self.branch = Some(branch.into());
self
}
pub fn with_optional_branch(mut self, branch: Option<String>) -> Self {
self.branch = branch;
self
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum WriteMode {
Append,
Overwrite,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DistanceMetric {
L2,
Cosine,
Dot,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum VectorIndexKind {
Flat,
IvfFlat { num_partitions: u32 },
IvfPq {
num_partitions: u32,
num_sub_vectors: u32,
num_bits: u8,
},
IvfSq { num_partitions: u32 },
IvfRq {
num_partitions: u32,
num_bits: Option<u8>,
},
HnswFlat {
m: u32,
ef_construction: u32,
num_partitions: u32,
},
HnswSq {
m: u32,
ef_construction: u32,
num_partitions: u32,
},
HnswPq {
m: u32,
ef_construction: u32,
num_sub_vectors: u32,
num_partitions: u32,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct VectorIndexParams {
pub metric: DistanceMetric,
pub kind: VectorIndexKind,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ScalarIndexType {
BTree,
Bitmap,
LabelList,
}
#[derive(Debug, Clone)]
pub struct IndexInfo {
pub name: String,
pub columns: Vec<String>,
pub index_type: String,
}