use std::collections::HashSet;
use std::fmt;
pub use rusqlite;
use rusqlite::Connection;
use rusqlite::types::{Value as SqlValue, ValueRef};
use crate::ast::{BinaryOp, Consumer, Expr, LogicalOp, Query, UnaryOp};
use crate::plan::{partition_pushable, residual_query};
use crate::planner::{Plan, QueryPlanner};
use crate::value::{Object, Value};
type RowMapper = Box<dyn Fn(Object) -> Value>;
pub struct SqliteTable<'c> {
db: &'c Connection,
table: String,
columns: HashSet<String>,
json_columns: Vec<String>,
map: Option<RowMapper>,
}
#[derive(Clone, Debug, PartialEq)]
pub struct Compiled {
pub sql: String,
pub params: Vec<SqlValue>,
pub residual: Query,
}
impl<'c> SqliteTable<'c> {
pub fn new<S: AsRef<str>>(db: &'c Connection, table: impl Into<String>, columns: &[S]) -> Self {
Self {
db,
table: table.into(),
columns: columns.iter().map(|c| c.as_ref().to_owned()).collect(),
json_columns: Vec::new(),
map: None,
}
}
pub fn json_columns<S: AsRef<str>>(mut self, columns: &[S]) -> Self {
self.json_columns = columns.iter().map(|c| c.as_ref().to_owned()).collect();
self
}
pub fn map(mut self, map: impl Fn(Object) -> Value + 'static) -> Self {
self.map = Some(Box::new(map));
self
}
pub fn table(&self) -> &str {
&self.table
}
pub fn is_column(&self, name: &str) -> bool {
self.columns.contains(name)
}
pub fn compile(&self, query: &Query, params: &[Value]) -> Option<Compiled> {
match &query.source {
Expr::Ident { name } if *name == self.table => {}
_ => return None,
}
if !query.from.is_empty() || query.follow.is_some() {
return None;
}
let (pushed, residual) =
partition_pushable(query.r#where.as_ref(), |e| self.translatable(e));
let mut sql_params = Vec::new();
let where_sql = pushed
.iter()
.map(|e| self.translate(e, params, &mut sql_params))
.collect::<Vec<_>>()
.join(" AND ");
let mut tail = "";
if query.consumer == Consumer::First
&& residual.is_none()
&& query.order_by.is_none()
&& query.limit.is_none()
&& query.offset.is_none()
{
tail = " LIMIT 1";
}
let mut sql = format!("SELECT * FROM {}", quote_ident(&self.table));
if !where_sql.is_empty() {
sql.push_str(" WHERE ");
sql.push_str(&where_sql);
}
sql.push_str(tail);
Some(Compiled {
sql,
params: sql_params,
residual: residual_query(query, residual),
})
}
pub fn try_plan(&self, query: &Query, params: &[Value]) -> rusqlite::Result<Option<Plan>> {
let Some(compiled) = self.compile(query, params) else {
return Ok(None);
};
let mut stmt = self.db.prepare(&compiled.sql)?;
let names: Vec<String> = stmt
.column_names()
.iter()
.map(|s| (*s).to_owned())
.collect();
let mut rows = stmt.query(rusqlite::params_from_iter(compiled.params.iter()))?;
let mut out = Vec::new();
while let Some(row) = rows.next()? {
let mut raw = Object::with_capacity(names.len());
for (i, name) in names.iter().enumerate() {
raw.insert(name.clone(), from_sql(row.get_ref(i)?));
}
out.push(self.map_row(raw));
}
Ok(Some(Plan::new(out, compiled.residual)))
}
fn translatable(&self, e: &Expr) -> bool {
match e {
Expr::Lit(_) | Expr::Binding { .. } => true,
Expr::Ident { name } => self.columns.contains(name),
Expr::Unary { expr, .. } => self.translatable(expr),
Expr::Logical { left, right, .. } => {
self.translatable(left) && self.translatable(right)
}
Expr::Binary { left, right, .. } => self.translatable(left) && self.translatable(right),
_ => false,
}
}
fn translate(&self, e: &Expr, params: &[Value], out: &mut Vec<SqlValue>) -> String {
match e {
Expr::Lit(v) => {
out.push(to_sql_param(v));
"?".to_owned()
}
Expr::Binding { index } => {
out.push(to_sql_param(
params.get(*index).unwrap_or(&Value::Undefined),
));
"?".to_owned()
}
Expr::Ident { name } => format!("{}.{}", quote_ident(&self.table), quote_ident(name)),
Expr::Unary { op, expr } => {
let inner = self.translate(expr, params, out);
match op {
UnaryOp::Not => format!("(NOT {inner})"),
UnaryOp::Neg => format!("(-{inner})"),
}
}
Expr::Logical { op, left, right } => {
let l = self.translate(left, params, out);
let r = self.translate(right, params, out);
let op = match op {
LogicalOp::And => "AND",
LogicalOp::Or => "OR",
};
format!("({l} {op} {r})")
}
Expr::Binary { op, left, right } => {
let l = self.translate(left, params, out);
let r = self.translate(right, params, out);
format!("({l} {} {r})", binary_sql(*op))
}
other => unreachable!("sqlite: not translatable: {other:?}"),
}
}
fn map_row(&self, raw: Object) -> Value {
if let Some(map) = &self.map {
return map(raw);
}
if self.json_columns.is_empty() {
return Value::Object(raw);
}
let mut out = raw;
for c in &self.json_columns {
let parsed = match out.get(c) {
Some(Value::Str(text)) => serde_json::from_str::<serde_json::Value>(text)
.ok()
.map(Value::from_json),
_ => None,
};
if let Some(v) = parsed {
out.insert(c.clone(), v);
}
}
Value::Object(out)
}
}
impl QueryPlanner for SqliteTable<'_> {
fn plan(&self, query: &Query, params: &[Value]) -> Option<Plan> {
match self.try_plan(query, params) {
Ok(plan) => plan,
Err(e) => {
let sql = self
.compile(query, params)
.map_or_else(String::new, |c| c.sql);
panic!("oqx sqlite adapter: {e} (statement: {sql})")
}
}
}
}
impl fmt::Debug for SqliteTable<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let mut columns: Vec<&str> = self.columns.iter().map(String::as_str).collect();
columns.sort_unstable();
f.debug_struct("SqliteTable")
.field("table", &self.table)
.field("columns", &columns)
.field("json_columns", &self.json_columns)
.field("map", &self.map.as_ref().map(|_| "<fn>"))
.finish_non_exhaustive()
}
}
fn binary_sql(op: BinaryOp) -> &'static str {
match op {
BinaryOp::Eq => "=",
BinaryOp::Ne => "<>",
BinaryOp::Lt => "<",
BinaryOp::Le => "<=",
BinaryOp::Gt => ">",
BinaryOp::Ge => ">=",
BinaryOp::Add => "+",
BinaryOp::Sub => "-",
BinaryOp::Mul => "*",
BinaryOp::Div => "/",
BinaryOp::Mod => "%",
}
}
fn quote_ident(name: &str) -> String {
format!("\"{}\"", name.replace('"', "\"\""))
}
fn to_sql_param(v: &Value) -> SqlValue {
match v {
Value::Undefined | Value::Null => SqlValue::Null,
Value::Bool(b) => SqlValue::Integer(i64::from(*b)),
Value::Number(n) => SqlValue::Real(*n),
Value::Str(s) => SqlValue::Text(s.clone()),
Value::Array(_) | Value::Object(_) | Value::Range(_) => SqlValue::Text(v.to_string()),
}
}
fn from_sql(v: ValueRef<'_>) -> Value {
match v {
ValueRef::Null => Value::Null,
ValueRef::Integer(i) => Value::Number(i as f64),
ValueRef::Real(f) => Value::Number(f),
ValueRef::Text(bytes) => Value::Str(String::from_utf8_lossy(bytes).into_owned()),
ValueRef::Blob(bytes) => {
Value::Array(bytes.iter().map(|b| Value::Number(f64::from(*b))).collect())
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::value::Range;
#[test]
fn params_bind_as_node_sqlite_would() {
assert_eq!(to_sql_param(&Value::Undefined), SqlValue::Null);
assert_eq!(to_sql_param(&Value::Null), SqlValue::Null);
assert_eq!(to_sql_param(&Value::Bool(true)), SqlValue::Integer(1));
assert_eq!(to_sql_param(&Value::Bool(false)), SqlValue::Integer(0));
assert_eq!(to_sql_param(&Value::Number(5.0)), SqlValue::Real(5.0));
assert_eq!(
to_sql_param(&Value::from("x")),
SqlValue::Text("x".to_owned())
);
assert_eq!(
to_sql_param(&Value::Array(vec![Value::Number(1.0), Value::from("a")])),
SqlValue::Text("1,a".to_owned())
);
assert_eq!(
to_sql_param(&Value::Object(Object::new())),
SqlValue::Text("[object Object]".to_owned())
);
assert_eq!(
to_sql_param(&Value::from(Range {
lo: Some(Value::Number(1.0)),
hi: None,
exclusive_end: false,
})),
SqlValue::Text("1..".to_owned())
);
}
#[test]
fn columns_map_back_to_values() {
assert_eq!(from_sql(ValueRef::Null), Value::Null);
assert_eq!(from_sql(ValueRef::Integer(7)), Value::Number(7.0));
assert_eq!(from_sql(ValueRef::Real(1.5)), Value::Number(1.5));
assert_eq!(from_sql(ValueRef::Text(b"hi")), Value::from("hi"));
assert_eq!(
from_sql(ValueRef::Text(b"a\xffb")),
Value::from("a\u{fffd}b")
);
assert_eq!(
from_sql(ValueRef::Blob(&[0, 255])),
Value::Array(vec![Value::Number(0.0), Value::Number(255.0)])
);
}
#[test]
fn identifiers_are_quoted_and_escaped() {
assert_eq!(quote_ident("dept"), "\"dept\"");
assert_eq!(quote_ident("we\"ird"), "\"we\"\"ird\"");
}
}