use sqlparser::ast::{BinaryOperator, Expr, SetExpr, Statement, Value as SqlValue};
use sqlparser::dialect::GenericDialect;
use sqlparser::parser::Parser;
use quatzal_schema::{Row, UaceError, UaceResult, Value};
use quatzal_storage::Engine;
#[derive(Debug, PartialEq)]
pub enum QueryResult {
TableCreated,
Inserted(usize),
Rows(Vec<Row>),
}
pub fn execute(engine: &Engine, sql: &str) -> UaceResult<QueryResult> {
let dialect = GenericDialect {};
let statements =
Parser::parse_sql(&dialect, sql).map_err(|e| UaceError::Parse(e.to_string()))?;
let statement = statements
.into_iter()
.next()
.ok_or_else(|| UaceError::Parse("empty statement".into()))?;
match statement {
Statement::CreateTable(_) => Ok(QueryResult::TableCreated),
Statement::Insert(insert) => execute_insert(engine, insert),
Statement::Query(query) => execute_select(engine, *query),
other => Err(UaceError::Parse(format!("unsupported statement: {other}"))),
}
}
fn execute_insert(engine: &Engine, insert: sqlparser::ast::Insert) -> UaceResult<QueryResult> {
let columns: Vec<String> = insert.columns.iter().map(|c| c.value.clone()).collect();
let key_pos = columns
.iter()
.position(|c| c.eq_ignore_ascii_case("key"))
.ok_or_else(|| UaceError::Parse("INSERT must include a `key` column".into()))?;
let source = insert
.source
.ok_or_else(|| UaceError::Parse("INSERT requires VALUES".into()))?;
let SetExpr::Values(values) = *source.body else {
return Err(UaceError::Parse("INSERT requires a VALUES clause".into()));
};
let mut inserted = 0usize;
for value_row in values.rows {
if value_row.len() != columns.len() {
return Err(UaceError::Parse(
"column/value count mismatch in INSERT".into(),
));
}
let mut row = Row::new(Vec::new());
for (col, expr) in columns.iter().zip(value_row.iter()) {
let value = expr_to_value(expr)?;
if col.eq_ignore_ascii_case("key") {
row.key = match &value {
Value::Text(s) => s.clone().into_bytes(),
Value::I64(n) => n.to_string().into_bytes(),
other => return Err(UaceError::Parse(format!("invalid key value: {other:?}"))),
};
} else {
row.scalars.insert(col.clone(), value);
}
}
if key_pos >= columns.len() {
return Err(UaceError::Parse("key column not found".into()));
}
engine.put(row)?;
inserted += 1;
}
Ok(QueryResult::Inserted(inserted))
}
fn execute_select(engine: &Engine, query: sqlparser::ast::Query) -> UaceResult<QueryResult> {
let SetExpr::Select(select) = *query.body else {
return Err(UaceError::Parse("only simple SELECT is supported".into()));
};
let selection = select
.selection
.ok_or_else(|| UaceError::Parse("SELECT requires WHERE key = <literal>".into()))?;
let Expr::BinaryOp { left, op, right } = &selection else {
return Err(UaceError::Parse(
"only `WHERE key = <literal>` is supported".into(),
));
};
if !matches!(op, BinaryOperator::Eq) {
return Err(UaceError::Parse("only `=` predicates are supported".into()));
}
let Expr::Identifier(ident) = left.as_ref() else {
return Err(UaceError::Parse(
"left side of WHERE must be a column name".into(),
));
};
if !ident.value.eq_ignore_ascii_case("key") {
return Err(UaceError::Parse(
"Phase 1 only supports point lookups on `key`".into(),
));
}
let key_value = expr_to_value(right)?;
let key_bytes = match key_value {
Value::Text(s) => s.into_bytes(),
Value::I64(n) => n.to_string().into_bytes(),
other => return Err(UaceError::Parse(format!("invalid key literal: {other:?}"))),
};
let row = engine.get(&key_bytes)?;
Ok(QueryResult::Rows(row.into_iter().collect()))
}
fn expr_to_value(expr: &Expr) -> UaceResult<Value> {
match expr {
Expr::Value(v) => sql_value_to_value(v),
Expr::UnaryOp {
op: sqlparser::ast::UnaryOperator::Minus,
expr,
} => match expr.as_ref() {
Expr::Value(v) => match v {
SqlValue::Number(s, _) => parse_number(&format!("-{s}")),
other => Err(UaceError::Parse(format!("unsupported literal: {other:?}"))),
},
other => Err(UaceError::Parse(format!(
"unsupported expression: {other:?}"
))),
},
other => Err(UaceError::Parse(format!(
"unsupported expression: {other:?}"
))),
}
}
fn sql_value_to_value(value: &SqlValue) -> UaceResult<Value> {
match value {
SqlValue::Number(s, _) => parse_number(s),
SqlValue::SingleQuotedString(s) | SqlValue::DoubleQuotedString(s) => {
Ok(Value::Text(s.clone()))
}
SqlValue::Boolean(b) => Ok(Value::Bool(*b)),
SqlValue::Null => Ok(Value::Null),
other => Err(UaceError::Parse(format!("unsupported literal: {other:?}"))),
}
}
fn parse_number(s: &str) -> UaceResult<Value> {
if let Ok(i) = s.parse::<i64>() {
Ok(Value::I64(i))
} else if let Ok(f) = s.parse::<f64>() {
Ok(Value::F64(f))
} else {
Err(UaceError::Parse(format!("invalid numeric literal: {s}")))
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::TempDir;
fn open_engine() -> (TempDir, Engine) {
let tmp = TempDir::new().unwrap();
let engine = Engine::open(tmp.path(), Some(1)).unwrap();
(tmp, engine)
}
#[test]
fn create_insert_select_round_trip() {
let (_tmp, engine) = open_engine();
let created = execute(
&engine,
"CREATE TABLE docs (\"key\" TEXT, title TEXT, views BIGINT)",
)
.unwrap();
assert_eq!(created, QueryResult::TableCreated);
let inserted = execute(
&engine,
"INSERT INTO docs (\"key\", title, views) VALUES ('doc-1', 'Hello', 42)",
)
.unwrap();
assert_eq!(inserted, QueryResult::Inserted(1));
let rows = execute(&engine, "SELECT * FROM docs WHERE \"key\" = 'doc-1'").unwrap();
match rows {
QueryResult::Rows(rows) => {
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].key, b"doc-1".to_vec());
assert_eq!(
rows[0].scalars.get("title"),
Some(&Value::Text("Hello".into()))
);
assert_eq!(rows[0].scalars.get("views"), Some(&Value::I64(42)));
}
other => panic!("expected Rows, got {other:?}"),
}
}
#[test]
fn select_missing_key_returns_empty() {
let (_tmp, engine) = open_engine();
let rows = execute(&engine, "SELECT * FROM docs WHERE \"key\" = 'nope'").unwrap();
assert_eq!(rows, QueryResult::Rows(vec![]));
}
}