use crate::prelude::{Cow, Vec};
use crate::{
ColumnRef, PaginationArg, SQL, SQLChunk, SQLSchemaType, SQLTable, ToSQL, Token, expr::Expr,
traits::SQLParam, types::BooleanLike,
};
#[doc(hidden)]
pub const MYSQL_UNBOUNDED_LIMIT: &str = "9223372036854775807";
pub fn select<'a, Value, T>(columns: T) -> SQL<'a, Value>
where
Value: SQLParam,
T: ToSQL<'a, Value>,
{
SQL::from(Token::SELECT).append(columns.into_sql())
}
pub fn select_distinct<'a, Value, T>(columns: T) -> SQL<'a, Value>
where
Value: SQLParam,
T: ToSQL<'a, Value>,
{
SQL::from_iter([Token::SELECT, Token::DISTINCT]).append(columns.into_sql())
}
#[derive(Debug, Default, Clone, Copy)]
struct OperandShape {
has_tail: bool,
has_union_or_except: bool,
has_intersect: bool,
starts_with_cte: bool,
}
impl OperandShape {
fn of<V: SQLParam>(sql: &SQL<'_, V>) -> Self {
let mut shape = Self::default();
let mut depth = 0usize;
let mut leading = true;
for chunk in &sql.chunks {
match chunk {
SQLChunk::Raw(text) if leading && text.trim_start().starts_with("/*") => {
continue;
}
SQLChunk::Token(Token::LPAREN) => depth += 1,
SQLChunk::Token(Token::RPAREN) => depth = depth.saturating_sub(1),
SQLChunk::Token(Token::WITH) if leading => shape.starts_with_cte = true,
SQLChunk::Token(Token::ORDER | Token::LIMIT | Token::OFFSET | Token::FOR)
if depth == 0 =>
{
shape.has_tail = true;
}
SQLChunk::Token(Token::UNION | Token::EXCEPT) if depth == 0 => {
shape.has_union_or_except = true;
}
SQLChunk::Token(Token::INTERSECT) if depth == 0 => shape.has_intersect = true,
_ => {}
}
leading = false;
}
shape
}
const fn is_compound(self) -> bool {
self.has_union_or_except || self.has_intersect
}
}
fn group_set_operand<'a, V: SQLParam>(operand: SQL<'a, V>) -> SQL<'a, V> {
match V::DIALECT {
crate::Dialect::SQLite => {
SQL::from_iter([Token::SELECT, Token::STAR, Token::FROM]).append(operand.parens())
}
crate::Dialect::PostgreSQL | crate::Dialect::MySQL => operand.parens(),
}
}
fn set_op<'a, Value, L, R>(left: L, op: Token, all: bool, right: R) -> SQL<'a, Value>
where
Value: SQLParam,
L: ToSQL<'a, Value>,
R: ToSQL<'a, Value>,
{
let left = left.into_sql();
let right = right.into_sql();
let left_shape = OperandShape::of(&left);
let intersect_binds_tighter = !matches!(Value::DIALECT, crate::Dialect::SQLite);
let left = if left_shape.has_tail
|| (intersect_binds_tighter
&& matches!(op, Token::INTERSECT)
&& left_shape.has_union_or_except)
{
group_set_operand(left)
} else {
left
};
let right_shape = OperandShape::of(&right);
let right = if right_shape.has_tail || right_shape.is_compound() || right_shape.starts_with_cte
{
group_set_operand(right)
} else {
right
};
let op_sql = if all {
SQL::from(op).push(Token::ALL)
} else {
SQL::from(op)
};
left.append(op_sql).append(right)
}
pub fn union<'a, Value, L, R>(left: L, right: R) -> SQL<'a, Value>
where
Value: SQLParam,
L: ToSQL<'a, Value>,
R: ToSQL<'a, Value>,
{
set_op(left, Token::UNION, false, right)
}
pub fn union_all<'a, Value, L, R>(left: L, right: R) -> SQL<'a, Value>
where
Value: SQLParam,
L: ToSQL<'a, Value>,
R: ToSQL<'a, Value>,
{
set_op(left, Token::UNION, true, right)
}
pub fn intersect<'a, Value, L, R>(left: L, right: R) -> SQL<'a, Value>
where
Value: SQLParam,
L: ToSQL<'a, Value>,
R: ToSQL<'a, Value>,
{
set_op(left, Token::INTERSECT, false, right)
}
pub fn intersect_all<'a, Value, L, R>(left: L, right: R) -> SQL<'a, Value>
where
Value: SQLParam,
L: ToSQL<'a, Value>,
R: ToSQL<'a, Value>,
{
set_op(left, Token::INTERSECT, true, right)
}
pub fn except<'a, Value, L, R>(left: L, right: R) -> SQL<'a, Value>
where
Value: SQLParam,
L: ToSQL<'a, Value>,
R: ToSQL<'a, Value>,
{
set_op(left, Token::EXCEPT, false, right)
}
pub fn except_all<'a, Value, L, R>(left: L, right: R) -> SQL<'a, Value>
where
Value: SQLParam,
L: ToSQL<'a, Value>,
R: ToSQL<'a, Value>,
{
set_op(left, Token::EXCEPT, true, right)
}
pub fn insert<'a, Table, Type, Value>(table: &Table) -> SQL<'a, Value>
where
Type: SQLSchemaType,
Value: SQLParam,
Table: SQLTable<'a, Type, Value>,
{
SQL::from_iter([Token::INSERT, Token::INTO]).append(table)
}
#[doc(hidden)]
pub fn insert_values_with_defaults<'a, V: SQLParam>(
rows: Vec<(Cow<'static, [ColumnRef]>, SQL<'a, V>)>,
) -> Option<SQL<'a, V>> {
let mut columns: Vec<ColumnRef> = Vec::new();
for (row_columns, _) in &rows {
for column in row_columns.iter() {
if !columns.contains(column) {
columns.push(*column);
}
}
}
let mut values = SQL::with_capacity_chunks(rows.len().saturating_mul(columns.len() * 2 + 2));
for (index, (row_columns, row_values)) in rows.into_iter().enumerate() {
let mut cells = split_top_level_commas(row_values);
if cells.len() != row_columns.len() {
return None;
}
if index > 0 {
values.push_mut(Token::COMMA);
}
values.push_mut(Token::LPAREN);
for (position, column) in columns.iter().enumerate() {
if position > 0 {
values.push_mut(Token::COMMA);
}
match row_columns.iter().position(|set| set == column) {
Some(cell) => values.append_mut(core::mem::take(&mut cells[cell])),
None => values.push_mut(Token::DEFAULT),
}
}
values.push_mut(Token::RPAREN);
}
Some(
SQL::columns(&columns)
.parens()
.push(Token::VALUES)
.append(values),
)
}
fn split_top_level_commas<'a, V: SQLParam>(sql: SQL<'a, V>) -> Vec<SQL<'a, V>> {
let mut parts = Vec::new();
let mut current = SQL::empty();
let mut depth = 0usize;
for chunk in sql.chunks {
match chunk {
SQLChunk::Token(Token::LPAREN) => depth += 1,
SQLChunk::Token(Token::RPAREN) => depth = depth.saturating_sub(1),
SQLChunk::Token(Token::COMMA) if depth == 0 => {
parts.push(core::mem::take(&mut current));
continue;
}
_ => {}
}
current.chunks.push(chunk);
}
if !current.chunks.is_empty() || !parts.is_empty() {
parts.push(current);
}
parts
}
pub fn from<'a, T, Value>(query: T) -> SQL<'a, Value>
where
T: ToSQL<'a, Value>,
Value: SQLParam,
{
SQL::from(Token::FROM).append(query.into_sql())
}
pub fn r#where<'a, V, E>(condition: E) -> SQL<'a, V>
where
V: SQLParam + 'a,
E: Expr<'a, V>,
E::SQLType: BooleanLike,
{
SQL::from(Token::WHERE).append(condition.into_expr_sql())
}
pub fn group_by<'a, V, I, T>(expressions: I) -> SQL<'a, V>
where
V: SQLParam + 'a,
I: IntoIterator<Item = T>,
T: ToSQL<'a, V>,
{
SQL::from_iter([Token::GROUP, Token::BY]).append(SQL::join(
expressions.into_iter().map(ToSQL::into_sql),
Token::COMMA,
))
}
pub fn group_by_expr<'a, V, T>(expr: T) -> SQL<'a, V>
where
V: SQLParam + 'a,
T: ToSQL<'a, V>,
{
SQL::from_iter([Token::GROUP, Token::BY]).append(expr.into_sql())
}
pub fn having<'a, V, E>(condition: E) -> SQL<'a, V>
where
V: SQLParam + 'a,
E: Expr<'a, V>,
E::SQLType: BooleanLike,
{
SQL::from(Token::HAVING).append(condition.into_expr_sql())
}
pub fn order_by<'a, T, V>(expressions: T) -> SQL<'a, V>
where
T: ToSQL<'a, V>,
V: SQLParam + 'a,
{
SQL::from_iter([Token::ORDER, Token::BY]).append(expressions.into_sql())
}
pub fn set_order_by<'a, T, V>(expressions: T) -> SQL<'a, V>
where
T: ToSQL<'a, V>,
V: SQLParam + 'a,
{
SQL::from_iter([Token::ORDER, Token::BY]).append(unqualified_columns(expressions.into_sql()))
}
pub fn unqualified_columns<'a, V>(mut sql: SQL<'a, V>) -> SQL<'a, V>
where
V: SQLParam + 'a,
{
for chunk in &mut sql.chunks {
if let SQLChunk::Column(column) = chunk {
*chunk = SQLChunk::ident_static(column.name);
}
}
sql
}
#[must_use]
#[track_caller]
pub fn limit<'a, V, P>(value: P) -> SQL<'a, V>
where
V: SQLParam + 'a,
P: PaginationArg<'a, V>,
{
SQL::from(Token::LIMIT).append(value.into_pagination_sql())
}
#[must_use]
#[track_caller]
pub fn offset<'a, V, P>(value: P) -> SQL<'a, V>
where
V: SQLParam + 'a,
P: PaginationArg<'a, V>,
{
SQL::from(Token::OFFSET).append(value.into_pagination_sql())
}
pub fn update<'a, Table, Type, Value>(table: &Table) -> SQL<'a, Value>
where
Table: SQLTable<'a, Type, Value>,
Type: SQLSchemaType,
Value: SQLParam + 'a,
{
SQL::from(Token::UPDATE).append(table)
}
pub fn set<'a, Table, Type, Value>(assignments: &Table::Update) -> SQL<'a, Value>
where
Value: SQLParam + 'a,
Table: SQLTable<'a, Type, Value>,
Type: SQLSchemaType,
{
SQL::from(Token::SET).append(assignments.to_sql())
}
pub fn delete<'a, Table, Type, Value>(table: &Table) -> SQL<'a, Value>
where
Table: SQLTable<'a, Type, Value>,
Type: SQLSchemaType,
Value: SQLParam + 'a,
{
SQL::from_iter([Token::DELETE, Token::FROM]).append(table)
}