use std::borrow::Cow;
use std::fmt;
use std::sync::Arc;
use crate::dialect::Dialect;
use crate::error::{Error, Result};
use crate::value::{ToValue, Value};
pub trait Expression: fmt::Debug + Send + Sync {
fn write_sql(&self, w: &mut SqlWriter<'_>);
}
pub type DynExpr = Arc<dyn Expression>;
pub fn dyn_expr(e: impl Expression + 'static) -> DynExpr {
Arc::new(e)
}
impl Expression for str {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
w.push_str(self);
}
}
impl Expression for String {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
w.push_str(self);
}
}
impl Expression for Cow<'_, str> {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
w.push_str(self);
}
}
macro_rules! impl_expression_for_number {
($($t:ty),+) => {
$(
impl Expression for $t {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
w.push_str(&self.to_string());
}
}
)+
};
}
impl_expression_for_number!(i8, i16, i32, i64, isize, u8, u16, u32, u64, usize, f32, f64);
impl<T: Expression + ?Sized> Expression for &T {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
(**self).write_sql(w);
}
}
impl<T: Expression + ?Sized> Expression for Box<T> {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
(**self).write_sql(w);
}
}
impl<T: Expression + ?Sized> Expression for Arc<T> {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
(**self).write_sql(w);
}
}
pub struct ExprFn<F>(F);
pub fn expr_fn<F>(f: F) -> ExprFn<F>
where
F: Fn(&mut SqlWriter<'_>) + Send + Sync,
{
ExprFn(f)
}
impl<F> fmt::Debug for ExprFn<F> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("ExprFn")
}
}
impl<F> Expression for ExprFn<F>
where
F: Fn(&mut SqlWriter<'_>) + Send + Sync,
{
fn write_sql(&self, w: &mut SqlWriter<'_>) {
(self.0)(w);
}
}
#[derive(Debug)]
pub struct SqlWriter<'d> {
sql: String,
args: Vec<Value>,
dialect: &'d dyn Dialect,
next_arg: usize,
error: Option<Error>,
}
impl<'d> SqlWriter<'d> {
pub fn new(dialect: &'d dyn Dialect) -> Self {
Self::with_start(dialect, 1)
}
pub fn with_start(dialect: &'d dyn Dialect, start: usize) -> Self {
assert!(start > 0, "placeholder positions are 1-based, got {start}");
SqlWriter {
sql: String::new(),
args: Vec::new(),
dialect,
next_arg: start,
error: None,
}
}
pub fn dialect(&self) -> &'d dyn Dialect {
self.dialect
}
pub fn sql(&self) -> &str {
&self.sql
}
pub fn args(&self) -> &[Value] {
&self.args
}
pub fn arg_position(&self) -> usize {
self.next_arg
}
pub fn error(&self) -> Option<&Error> {
self.error.as_ref()
}
pub fn record_error(&mut self, e: Error) {
if self.error.is_none() {
self.error = Some(e);
}
}
pub fn push_str(&mut self, s: &str) {
self.sql.push_str(s);
}
pub fn push_arg(&mut self, v: impl ToValue) {
let (d, pos) = (self.dialect, self.next_arg);
d.write_arg(self, pos);
self.args.push(v.to_value());
self.next_arg += 1;
}
pub fn push_named_arg(&mut self, name: &str) {
let d = self.dialect;
d.write_named_arg(self, name);
}
pub fn push_quoted<S: AsRef<str>>(&mut self, parts: &[S]) {
let d = self.dialect;
let mut written = 0;
for part in parts {
let part = part.as_ref();
if part.is_empty() {
continue;
}
if written > 0 {
self.sql.push('.');
}
d.write_quoted(self, part);
written += 1;
}
}
pub fn write_expr<E: Expression + ?Sized>(&mut self, e: &E) {
e.write_sql(self);
}
pub fn write_if<E: Expression + ?Sized>(
&mut self,
cond: bool,
prefix: &str,
e: &E,
suffix: &str,
) {
if !cond {
return;
}
self.push_str(prefix);
self.write_expr(e);
self.push_str(suffix);
}
pub fn write_if_some<E: Expression + ?Sized>(
&mut self,
e: Option<&E>,
prefix: &str,
suffix: &str,
) {
if let Some(e) = e {
self.push_str(prefix);
self.write_expr(e);
self.push_str(suffix);
}
}
pub fn write_slice<E: Expression>(
&mut self,
items: &[E],
prefix: &str,
sep: &str,
suffix: &str,
) {
if items.is_empty() {
return;
}
self.push_str(prefix);
for (i, item) in items.iter().enumerate() {
if i > 0 {
self.push_str(sep);
}
self.write_expr(item);
}
self.push_str(suffix);
}
pub fn write_iter<E, I>(&mut self, items: I, prefix: &str, sep: &str, suffix: &str)
where
E: Expression,
I: IntoIterator<Item = E>,
{
let mut it = items.into_iter().peekable();
if it.peek().is_none() {
return;
}
self.push_str(prefix);
for (i, item) in it.enumerate() {
if i > 0 {
self.push_str(sep);
}
self.write_expr(&item);
}
self.push_str(suffix);
}
pub fn write_with_dialect<E: Expression + ?Sized>(&mut self, dialect: &dyn Dialect, e: &E) {
let mut nested = SqlWriter {
sql: std::mem::take(&mut self.sql),
args: std::mem::take(&mut self.args),
dialect,
next_arg: self.next_arg,
error: self.error.take(),
};
e.write_sql(&mut nested);
self.sql = nested.sql;
self.args = nested.args;
self.next_arg = nested.next_arg;
self.error = nested.error;
}
pub fn finish(self) -> Result<(String, Vec<Value>)> {
match self.error {
Some(e) => Err(e),
None => Ok((self.sql, self.args)),
}
}
pub fn into_parts(self) -> (String, Vec<Value>, Option<Error>) {
(self.sql, self.args, self.error)
}
}
impl fmt::Write for SqlWriter<'_> {
fn write_str(&mut self, s: &str) -> fmt::Result {
self.sql.push_str(s);
Ok(())
}
}
pub fn build<E: Expression + ?Sized>(dialect: &dyn Dialect, e: &E) -> Result<(String, Vec<Value>)> {
build_from(dialect, 1, e)
}
pub fn build_from<E: Expression + ?Sized>(
dialect: &dyn Dialect,
start: usize,
e: &E,
) -> Result<(String, Vec<Value>)> {
let mut w = SqlWriter::with_start(dialect, start);
e.write_sql(&mut w);
w.finish()
}
#[cfg(test)]
mod tests {
use std::fmt::Write as _;
use keelson_sqlcheck::testing::assert_frag_sql;
use super::*;
use crate::dialect::testing::{Numbered, Positional, TestDialect};
const COND: &str = r#"SELECT "id" FROM users WHERE {}"#;
const VALUE: &str = r#"SELECT {} FROM users"#;
#[derive(Debug)]
struct Eq(&'static str, i32);
impl Expression for Eq {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
w.push_quoted(&[self.0]);
w.push_str(" = ");
w.push_arg(self.1);
}
}
#[derive(Debug)]
struct Sub(Vec<Eq>);
impl Expression for Sub {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
w.push_str("(SELECT 1 FROM users WHERE ");
w.write_slice(&self.0, "", " AND ", "");
w.push_str(")");
}
}
#[test]
fn placeholders_are_numbered_in_write_order() {
let (sql, args) = build(
&Numbered,
&Sub(vec![Eq("age", 10), Eq("id", 20), Eq("name", 30)]),
)
.unwrap();
assert_frag_sql(
r#"SELECT "id" FROM users WHERE "id" IN {}"#,
&sql,
r#"(SELECT 1 FROM users WHERE "age" = $1 AND "id" = $2 AND "name" = $3)"#,
);
assert_eq!(args, vec![Value::I32(10), Value::I32(20), Value::I32(30)]);
}
#[test]
fn nesting_continues_the_outer_numbering() {
#[derive(Debug)]
struct Outer;
impl Expression for Outer {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
w.write_expr(&Eq("age", 1));
w.push_str(" AND EXISTS ");
w.write_expr(&Sub(vec![Eq("id", 2), Eq("name", 3)]));
w.push_str(" AND ");
w.write_expr(&Eq("email", 4));
}
}
let (sql, args) = build(&Numbered, &Outer).unwrap();
assert_frag_sql(
COND,
&sql,
concat!(
r#""age" = $1 AND EXISTS (SELECT 1 FROM users WHERE "id" = $2 AND "name" = $3)"#,
r#" AND "email" = $4"#
),
);
assert_eq!(args.len(), 4);
assert_eq!(args[3], Value::I32(4));
}
#[test]
fn a_subquery_three_levels_deep_never_restarts_numbering() {
#[derive(Debug)]
struct Nest(usize);
impl Expression for Nest {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
w.push_str("(");
w.push_arg(self.0 as i32);
if self.0 > 1 {
w.push_str(" IN ");
w.write_expr(&Nest(self.0 - 1));
}
w.push_str(")");
}
}
let (sql, args) = build(&Numbered, &Nest(3)).unwrap();
assert_eq!(sql, "($1 IN ($2 IN ($3)))");
assert_eq!(args, vec![Value::I32(3), Value::I32(2), Value::I32(1)]);
}
#[test]
fn interleaved_siblings_and_children_stay_in_write_order() {
#[derive(Debug)]
struct Pair(Box<dyn Expression>, Box<dyn Expression>);
impl Expression for Pair {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
w.push_str("[");
w.write_expr(&self.0);
w.push_str(" ");
w.write_expr(&self.1);
w.push_str("]");
}
}
let tree = Pair(
Box::new(Pair(Box::new(Eq("a", 1)), Box::new(Eq("b", 2)))),
Box::new(Pair(Box::new(Eq("c", 3)), Box::new(Eq("d", 4)))),
);
let (sql, args) = build(&Numbered, &tree).unwrap();
assert_eq!(sql, r#"[["a" = $1 "b" = $2] ["c" = $3 "d" = $4]]"#);
assert_eq!(
args,
vec![Value::I32(1), Value::I32(2), Value::I32(3), Value::I32(4)]
);
}
#[test]
fn build_from_offsets_the_first_placeholder() {
let (sql, args) = build_from(&Numbered, 3, &Sub(vec![Eq("age", 1), Eq("id", 2)])).unwrap();
assert_eq!(
sql,
r#"(SELECT 1 FROM users WHERE "age" = $3 AND "id" = $4)"#
);
assert_eq!(args.len(), 2, "args are still returned from the start");
}
#[test]
#[should_panic(expected = "1-based")]
fn start_zero_is_rejected() {
let _ = build_from(&Numbered, 0, &Eq("a", 1));
}
#[test]
fn positional_dialects_ignore_the_index_but_still_order_args() {
let (sql, args) = build(&Positional, &Sub(vec![Eq("age", 7), Eq("id", 8)])).unwrap();
assert_eq!(sql, "(SELECT 1 FROM users WHERE `age` = ? AND `id` = ?)");
assert_eq!(args, vec![Value::I32(7), Value::I32(8)]);
}
#[test]
fn arg_position_tracks_the_next_placeholder() {
let mut w = SqlWriter::new(&Numbered);
assert_eq!(w.arg_position(), 1);
w.push_arg(1i32);
assert_eq!(w.arg_position(), 2);
w.push_str(" -- not an arg");
assert_eq!(w.arg_position(), 2);
w.push_arg("two");
assert_eq!(w.arg_position(), 3);
}
#[test]
fn raw_strings_of_every_stored_form_are_expressions() {
let (sql, args) = build(&Numbered, "id = 1").unwrap();
assert_frag_sql(COND, &sql, "id = 1");
assert!(args.is_empty());
let (sql, _) = build(&Numbered, &String::from("id = 2")).unwrap();
assert_frag_sql(COND, &sql, "id = 2");
let borrowed: Cow<'static, str> = Cow::Borrowed("id = 3");
let (sql, _) = build(&Numbered, &borrowed).unwrap();
assert_frag_sql(COND, &sql, "id = 3");
let owned: Cow<'static, str> = Cow::Owned(String::from("id = 4"));
let (sql, _) = build(&Numbered, &owned).unwrap();
assert_frag_sql(COND, &sql, "id = 4");
let boxed: Box<dyn Expression> = Box::new(Eq("age", 1));
let (sql, _) = build(&Numbered, &boxed).unwrap();
assert_frag_sql(COND, &sql, r#""age" = $1"#);
let shared: DynExpr = dyn_expr(Eq("id", 2));
let (sql, _) = build(&Numbered, &shared).unwrap();
assert_frag_sql(COND, &sql, r#""id" = $1"#);
}
#[test]
fn numbers_render_as_literals_not_placeholders() {
let (sql, args) = build(&Numbered, &20i64).unwrap();
assert_frag_sql(VALUE, &sql, "20");
assert!(args.is_empty(), "a literal binds nothing");
}
#[test]
fn expr_fn_wraps_a_closure() {
let e = expr_fn(|w: &mut SqlWriter<'_>| {
w.push_str("LIMIT ");
w.push_arg(5i64);
});
let (sql, args) = build(&Numbered, &e).unwrap();
assert_frag_sql(r#"SELECT "id" FROM users {}"#, &sql, "LIMIT $1");
assert_eq!(args, vec![Value::I64(5)]);
}
#[test]
fn write_if_skips_everything_including_the_affixes() {
let mut w = SqlWriter::new(&Numbered);
w.write_if(false, " WHERE ", &Eq("a", 1), "!");
assert_eq!(w.sql(), "");
assert_eq!(w.arg_position(), 1, "a skipped arg must not advance");
w.write_if(true, " WHERE ", &Eq("a", 1), "!");
let (sql, args) = w.finish().unwrap();
assert_eq!(sql, r#" WHERE "a" = $1!"#);
assert_eq!(args.len(), 1);
}
#[test]
fn write_if_some_follows_the_option() {
let mut w = SqlWriter::new(&Numbered);
w.write_if_some(None::<&Eq>, " LIMIT ", "");
assert_eq!(w.sql(), "");
w.write_if_some(Some(&Eq("a", 1)), " WHERE ", ";");
assert_eq!(w.sql(), r#" WHERE "a" = $1;"#);
}
#[test]
fn write_slice_is_a_no_op_when_empty() {
let mut w = SqlWriter::new(&Numbered);
w.write_slice::<Eq>(&[], " WHERE ", " AND ", ";");
assert_eq!(w.sql(), "");
w.write_iter(Vec::<String>::new(), "(", ", ", ")");
assert_eq!(w.sql(), "");
w.write_iter(vec!["a", "b"], "(", ", ", ")");
assert_eq!(w.sql(), "(a, b)");
}
#[test]
fn push_quoted_joins_with_dots_and_drops_empty_parts() {
let mut w = SqlWriter::new(&Numbered);
w.push_quoted(&["users", "id"]);
w.push_str(" ");
w.push_quoted(&["", "id"]);
w.push_str(" ");
w.push_quoted::<&str>(&[]);
w.push_str(" ");
w.push_quoted(&[Cow::Borrowed("a"), Cow::Owned("b".to_owned())]);
assert_eq!(w.sql(), r#""users"."id" "id" "a"."b""#);
}
#[test]
fn named_args_do_not_consume_an_arg_slot() {
let mut w = SqlWriter::new(&TestDialect);
w.push_arg(1i32);
w.push_str(", ");
w.push_named_arg("name");
w.push_str(", ");
w.push_arg(2i32);
let (sql, args) = w.finish().unwrap();
assert_eq!(sql, "?1, :name, ?2");
assert_eq!(args, vec![Value::I32(1), Value::I32(2)]);
}
#[test]
fn a_nested_dialect_shares_the_arg_list_and_counter() {
#[derive(Debug)]
struct Mixed;
impl Expression for Mixed {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
w.write_expr(&Eq("a", 1));
w.push_str(" AND ");
w.write_with_dialect(&Positional, &Eq("b", 2));
w.push_str(" AND ");
w.write_expr(&Eq("c", 3));
}
}
let (sql, args) = build(&Numbered, &Mixed).unwrap();
assert_eq!(sql, r#""a" = $1 AND `b` = ? AND "c" = $3"#);
assert_eq!(
args.len(),
3,
"the counter advanced through the nested part"
);
}
#[test]
fn a_recorded_error_is_surfaced_by_build_and_only_by_build() {
#[derive(Debug)]
struct Bad;
impl Expression for Bad {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
w.push_str("x = ");
w.push_named_arg("nope");
}
}
let mut w = SqlWriter::new(&Numbered);
w.write_expr(&Bad);
assert_eq!(w.sql(), "x = ");
assert!(matches!(w.error(), Some(Error::NoNamedArgs)));
assert!(matches!(build(&Numbered, &Bad), Err(Error::NoNamedArgs)));
let mut w = SqlWriter::new(&Numbered);
w.write_slice(&[Bad], "(", ", ", ")");
assert!(w.finish().is_err());
}
#[test]
fn the_first_recorded_error_wins() {
let mut w = SqlWriter::new(&Numbered);
w.record_error(Error::Incomplete("a table"));
w.record_error(Error::NoNamedArgs);
let (_, _, err) = w.into_parts();
assert!(matches!(err, Some(Error::Incomplete("a table"))));
}
#[test]
fn fmt_write_appends_to_the_same_buffer() {
let mut w = SqlWriter::new(&Numbered);
write!(w, "OFFSET {}", 4).unwrap();
assert_eq!(w.sql(), "OFFSET 4");
}
#[test]
fn the_writer_exposes_its_dialect_to_nested_expressions() {
#[derive(Debug)]
struct UsesDialect;
impl Expression for UsesDialect {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
let d = w.dialect();
d.write_quoted(w, "col");
}
}
let (sql, _) = build(&Positional, &UsesDialect).unwrap();
assert_eq!(sql, "`col`");
}
}