use std::borrow::Cow;
use std::fmt;
use crate::dialect::Dialect;
use crate::error::Result;
use crate::expr::{IntoExpr, RawArg};
use crate::value::{ToValue, Value};
use crate::writer::{Expression, build, build_from};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
#[non_exhaustive]
pub enum QueryType {
#[default]
Unknown,
Select,
Insert,
Update,
Delete,
Merge,
}
impl fmt::Display for QueryType {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(match self {
QueryType::Unknown => "UNKNOWN",
QueryType::Select => "SELECT",
QueryType::Insert => "INSERT",
QueryType::Update => "UPDATE",
QueryType::Delete => "DELETE",
QueryType::Merge => "MERGE",
})
}
}
pub trait Query: Expression {
fn query_type(&self) -> QueryType;
fn dialect(&self) -> &dyn Dialect;
fn build(&self) -> Result<(String, Vec<Value>)> {
build(self.dialect(), self)
}
fn build_from(&self, start: usize) -> Result<(String, Vec<Value>)> {
build_from(self.dialect(), start, self)
}
}
pub trait QueryExtensions<Hook, Loader, MapperMod>: Query {
fn hooks(&self) -> &[Hook] {
&[]
}
fn loaders(&self) -> &[Loader] {
&[]
}
fn mapper_mods(&self) -> &[MapperMod] {
&[]
}
}
#[derive(Debug, Clone)]
pub struct RawQuery<D> {
sql: Cow<'static, str>,
args: Vec<RawArg>,
dialect: D,
query_type: QueryType,
}
impl<D> RawQuery<D> {
pub fn new(dialect: D, sql: impl Into<Cow<'static, str>>) -> Self {
RawQuery {
sql: sql.into(),
args: Vec::new(),
dialect,
query_type: QueryType::Unknown,
}
}
#[must_use]
pub fn bind(mut self, value: impl ToValue) -> Self {
self.args.push(RawArg::value(value));
self
}
#[must_use]
pub fn bind_all<V: ToValue>(mut self, values: impl IntoIterator<Item = V>) -> Self {
self.args.extend(values.into_iter().map(RawArg::value));
self
}
#[must_use]
pub fn bind_expr(mut self, expression: impl IntoExpr) -> Self {
self.args.push(RawArg::expr(expression));
self
}
#[must_use]
pub fn kind(mut self, query_type: QueryType) -> Self {
self.query_type = query_type;
self
}
}
impl<D: fmt::Debug + Send + Sync> Expression for RawQuery<D> {
fn write_sql(&self, w: &mut crate::writer::SqlWriter<'_>) {
crate::expr::template(self.sql.clone(), self.args.iter().cloned()).write_sql(w);
}
}
impl<D: Dialect> Query for RawQuery<D> {
fn query_type(&self) -> QueryType {
self.query_type
}
fn dialect(&self) -> &dyn Dialect {
&self.dialect
}
}
impl<D: Dialect, H, L, M> QueryExtensions<H, L, M> for RawQuery<D> {}
#[cfg(test)]
mod tests {
use keelson_sqlcheck::testing::assert_stmt_sql;
use super::*;
use crate::dialect::testing::Numbered;
use crate::error::Error;
use crate::writer::SqlWriter;
#[test]
fn display_matches_the_sql_keyword() {
assert_eq!(QueryType::Select.to_string(), "SELECT");
assert_eq!(QueryType::Insert.to_string(), "INSERT");
assert_eq!(QueryType::Update.to_string(), "UPDATE");
assert_eq!(QueryType::Delete.to_string(), "DELETE");
assert_eq!(QueryType::Merge.to_string(), "MERGE");
assert_eq!(QueryType::Unknown.to_string(), "UNKNOWN");
assert_eq!(QueryType::default(), QueryType::Unknown);
}
#[derive(Debug)]
struct Select {
table: &'static str,
min_age: i32,
}
impl Expression for Select {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
w.push_str("SELECT * FROM ");
w.push_quoted(&[self.table]);
w.push_str(" WHERE ");
w.push_quoted(&["age"]);
w.push_str(" >= ");
w.push_arg(self.min_age);
}
}
impl Query for Select {
fn query_type(&self) -> QueryType {
QueryType::Select
}
fn dialect(&self) -> &dyn Dialect {
&Numbered
}
}
#[test]
fn a_query_builds_itself_without_being_told_the_dialect() {
let q = Select {
table: "users",
min_age: 21,
};
let (sql, args) = q.build().unwrap();
assert_stmt_sql(&sql, r#"SELECT * FROM "users" WHERE "age" >= $1"#);
assert_eq!(args, vec![Value::I32(21)]);
assert_eq!(q.query_type(), QueryType::Select);
let (sql, _) = q.build_from(4).unwrap();
assert_eq!(sql, r#"SELECT * FROM "users" WHERE "age" >= $4"#);
}
#[test]
fn a_query_is_usable_erased() {
let q: Box<dyn Query> = Box::new(Select {
table: "users",
min_age: 1,
});
assert_eq!(q.query_type(), QueryType::Select);
assert!(q.build().is_ok());
}
#[test]
fn a_query_nested_in_another_shares_the_numbering() {
#[derive(Debug)]
struct Wrapper(Select);
impl Expression for Wrapper {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
w.push_str("SELECT * FROM (");
w.write_expr(&self.0);
w.push_str(") AS \"u\" WHERE \"u\".\"id\" = ");
w.push_arg(9i32);
}
}
let (sql, args) = build(
&Numbered,
&Wrapper(Select {
table: "users",
min_age: 21,
}),
)
.unwrap();
assert_stmt_sql(
&sql,
concat!(
r#"SELECT * FROM (SELECT * FROM "users" WHERE "age" >= $1) AS "u" "#,
r#"WHERE "u"."id" = $2"#
),
);
assert_eq!(args, vec![Value::I32(21), Value::I32(9)]);
}
impl<H, L, M> QueryExtensions<H, L, M> for Select {}
#[test]
fn extension_points_default_to_none_without_any_downcasting() {
let q = Select {
table: "users",
min_age: 1,
};
let q: &dyn QueryExtensions<&'static str, u8, u8> = &q;
assert!(q.hooks().is_empty());
assert!(q.loaders().is_empty());
assert!(q.mapper_mods().is_empty());
}
#[test]
fn a_hand_written_statement_binds_and_rewrites_its_placeholders() {
let q = RawQuery::new(Numbered, "SELECT * FROM \"users\" WHERE \"age\" >= ?")
.bind(21)
.kind(QueryType::Select);
let (sql, args) = q.build().unwrap();
assert_stmt_sql(&sql, r#"SELECT * FROM "users" WHERE "age" >= $1"#);
assert_eq!(args, vec![Value::I32(21)]);
assert_eq!(q.query_type(), QueryType::Select);
}
#[test]
fn the_statement_kind_is_unknown_until_it_is_declared() {
let q = RawQuery::new(Numbered, "SELECT 1");
assert_eq!(q.query_type(), QueryType::Unknown);
}
#[test]
fn binding_more_than_the_placeholders_is_an_error_not_a_misbound_statement() {
let q = RawQuery::new(Numbered, "SELECT * FROM \"users\" WHERE \"age\" >= ?")
.bind(21)
.bind(22);
let err = q.build().unwrap_err();
assert!(matches!(err, Error::RawArgCount { .. }), "{err}");
}
#[test]
fn an_expression_can_be_spliced_where_a_value_would_go() {
let q = RawQuery::new(
Numbered,
"SELECT * FROM \"users\" WHERE \"id\" IN (?) AND \"age\" >= ?",
)
.bind_expr(crate::expr::args([1, 2, 3]))
.bind(21);
let (sql, args) = q.build().unwrap();
assert_stmt_sql(
&sql,
r#"SELECT * FROM "users" WHERE "id" IN ($1, $2, $3) AND "age" >= $4"#,
);
assert_eq!(
args,
vec![Value::I32(1), Value::I32(2), Value::I32(3), Value::I32(21)]
);
}
#[test]
fn a_hand_written_statement_nests_in_a_built_one() {
let (sql, args) = build(
&Numbered,
&Select {
table: "users",
min_age: 21,
}
.wrapped_around(
RawQuery::new(Numbered, "SELECT \"id\" FROM \"posts\" WHERE \"views\" > ?")
.bind(100),
),
)
.unwrap();
assert_stmt_sql(
&sql,
concat!(
r#"SELECT * FROM "users" WHERE "age" >= $1 AND "id" IN "#,
r#"(SELECT "id" FROM "posts" WHERE "views" > $2)"#
),
);
assert_eq!(args, vec![Value::I32(21), Value::I32(100)]);
}
#[derive(Debug)]
struct Wrapped(Select, RawQuery<Numbered>);
impl Select {
fn wrapped_around(self, inner: RawQuery<Numbered>) -> Wrapped {
Wrapped(self, inner)
}
}
impl Expression for Wrapped {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
self.0.write_sql(w);
w.push_str(" AND ");
w.push_quoted(&["id"]);
w.push_str(" IN (");
self.1.write_sql(w);
w.push_str(")");
}
}
}