1use crate::SQLConstraintKind;
2use crate::prelude::*;
3use crate::{Dialect, Param, Placeholder, SQLParam, sql::tokens::Token};
4
5#[derive(Clone, Copy, Debug, PartialEq, Eq)]
12pub struct EnumVariantRef {
13 pub label: &'static str,
15 pub discriminant: i64,
17}
18
19pub trait SQLEnumVariants {
26 const VARIANTS: &'static [EnumVariantRef];
28}
29
30#[derive(Clone, Copy, Debug, PartialEq, Eq)]
35pub enum ColumnDialect {
36 SQLite {
37 autoincrement: bool,
38 default: Option<&'static str>,
39 generated_expression: Option<&'static str>,
40 generated_stored: bool,
41 collate: Option<&'static str>,
42 enum_variants: Option<&'static [EnumVariantRef]>,
44 },
45 PostgreSQL {
46 postgres_type: &'static str,
47 dimensions: Option<i32>,
48 is_serial: bool,
49 is_bigserial: bool,
50 is_generated_identity: bool,
51 is_identity_always: bool,
52 default: Option<&'static str>,
53 generated_expression: Option<&'static str>,
54 generated_stored: bool,
55 collate: Option<&'static str>,
56 comment: Option<&'static str>,
57 enum_variants: Option<&'static [EnumVariantRef]>,
59 },
60 MySQL {
61 auto_increment: bool,
62 default: Option<&'static str>,
63 generated_expression: Option<&'static str>,
64 generated_stored: bool,
65 charset: Option<&'static str>,
66 collate: Option<&'static str>,
67 on_update: Option<&'static str>,
68 },
69}
70
71#[derive(Clone, Copy, Debug, PartialEq, Eq)]
75pub enum TableDialect {
76 PostgreSQL {
77 is_unlogged: bool,
78 is_temporary: bool,
79 inherits: Option<&'static str>,
80 tablespace: Option<&'static str>,
81 is_rls_enabled: bool,
82 comment: Option<&'static str>,
83 },
84 SQLite {
85 without_rowid: bool,
86 strict: bool,
87 },
88 MySQL {
89 is_temporary: bool,
90 engine: Option<&'static str>,
91 charset: Option<&'static str>,
92 collate: Option<&'static str>,
93 comment: Option<&'static str>,
94 },
95}
96
97impl Default for TableDialect {
98 fn default() -> Self {
99 Self::PostgreSQL {
100 is_unlogged: false,
101 is_temporary: false,
102 inherits: None,
103 tablespace: None,
104 is_rls_enabled: false,
105 comment: None,
106 }
107 }
108}
109
110#[derive(Clone, Copy, Debug, PartialEq, Eq)]
115pub struct ForeignKeyRef {
116 pub name: &'static str,
117 pub name_explicit: bool,
118 pub target_table: &'static str,
119 pub target_schema: &'static str,
120 pub source_columns: &'static [&'static str],
121 pub target_columns: &'static [&'static str],
122 pub on_delete: Option<&'static str>,
123 pub on_update: Option<&'static str>,
124 pub deferrable: bool,
125 pub initially_deferred: bool,
126}
127
128#[derive(Clone, Copy, Debug, PartialEq, Eq)]
130pub struct PrimaryKeyRef {
131 pub columns: &'static [&'static str],
132}
133
134#[derive(Clone, Copy, Debug, PartialEq, Eq)]
136pub struct ConstraintRef {
137 pub name: Option<&'static str>,
138 pub name_explicit: bool,
139 pub kind: SQLConstraintKind,
140 pub columns: &'static [&'static str],
141 pub check_expression: Option<&'static str>,
142 pub deferrable: bool,
143 pub initially_deferred: bool,
144}
145
146#[derive(Clone, Copy, Debug, PartialEq, Eq)]
154pub struct TableRef {
155 pub name: &'static str,
157 pub column_names: &'static [&'static str],
158
159 pub schema: Option<&'static str>,
161 pub qualified_name: &'static str,
162 pub columns: &'static [ColumnRef],
163 pub primary_key: Option<PrimaryKeyRef>,
164 pub foreign_keys: &'static [ForeignKeyRef],
165 pub constraints: &'static [ConstraintRef],
166 pub dependency_names: &'static [&'static str],
167
168 pub dialect: TableDialect,
170}
171
172impl TableRef {
173 #[must_use]
179 pub const fn sql(name: &'static str, column_names: &'static [&'static str]) -> Self {
180 Self {
181 name,
182 column_names,
183 schema: None,
184 qualified_name: "",
185 columns: &[],
186 primary_key: None,
187 foreign_keys: &[],
188 constraints: &[],
189 dependency_names: &[],
190 dialect: TableDialect::PostgreSQL {
191 is_unlogged: false,
192 is_temporary: false,
193 inherits: None,
194 tablespace: None,
195 is_rls_enabled: false,
196 comment: None,
197 },
198 }
199 }
200}
201
202#[derive(Clone, Copy, Debug, PartialEq, Eq)]
204pub struct TableSqlRef {
205 pub schema: Option<&'static str>,
206 pub name: &'static str,
207 pub column_names: &'static [&'static str],
208}
209
210impl TableSqlRef {
211 #[inline]
213 #[must_use]
214 pub const fn from_table_ref(table: TableRef) -> Self {
215 Self {
216 schema: table.schema,
217 name: table.name,
218 column_names: table.column_names,
219 }
220 }
221
222 #[inline]
224 #[must_use]
225 pub const fn from_table_ref_ref(table: &TableRef) -> Self {
226 Self {
227 schema: table.schema,
228 name: table.name,
229 column_names: table.column_names,
230 }
231 }
232}
233
234impl From<&TableRef> for TableSqlRef {
235 #[inline]
236 fn from(value: &TableRef) -> Self {
237 Self::from_table_ref_ref(value)
238 }
239}
240
241impl From<TableRef> for TableSqlRef {
242 #[inline]
243 fn from(value: TableRef) -> Self {
244 Self::from_table_ref(value)
245 }
246}
247
248#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
262pub struct ColumnFlags(u8);
263
264impl ColumnFlags {
265 pub const NOT_NULL: Self = Self(1 << 0);
267 pub const PRIMARY_KEY: Self = Self(1 << 1);
269 pub const UNIQUE: Self = Self(1 << 2);
271 pub const HAS_DEFAULT: Self = Self(1 << 3);
273
274 #[must_use]
276 pub const fn empty() -> Self {
277 Self(0)
278 }
279
280 #[must_use]
282 pub const fn from_bits(bits: u8) -> Self {
283 Self(bits)
284 }
285
286 #[must_use]
288 pub const fn bits(self) -> u8 {
289 self.0
290 }
291
292 #[must_use]
294 pub const fn contains(self, other: Self) -> bool {
295 (self.0 & other.0) == other.0
296 }
297
298 #[must_use]
300 pub const fn union(self, other: Self) -> Self {
301 Self(self.0 | other.0)
302 }
303}
304
305impl core::ops::BitOr for ColumnFlags {
306 type Output = Self;
307 fn bitor(self, rhs: Self) -> Self {
308 self.union(rhs)
309 }
310}
311
312impl core::ops::BitOrAssign for ColumnFlags {
313 fn bitor_assign(&mut self, rhs: Self) {
314 *self = self.union(rhs);
315 }
316}
317
318#[derive(Clone, Copy, Debug, PartialEq, Eq)]
324pub struct ColumnRef {
325 pub table: &'static str,
327 pub name: &'static str,
328
329 pub sql_type: &'static str,
331 pub flags: ColumnFlags,
332
333 pub dialect: ColumnDialect,
335}
336
337impl ColumnRef {
338 #[must_use]
344 pub const fn sql(table: &'static str, name: &'static str) -> Self {
345 Self {
346 table,
347 name,
348 sql_type: "",
349 flags: ColumnFlags::empty(),
350 dialect: ColumnDialect::SQLite {
351 autoincrement: false,
352 default: None,
353 generated_expression: None,
354 generated_stored: false,
355 collate: None,
356 enum_variants: None,
357 },
358 }
359 }
360
361 #[must_use]
363 pub const fn not_null(&self) -> bool {
364 self.flags.contains(ColumnFlags::NOT_NULL)
365 }
366
367 #[must_use]
369 pub const fn primary_key(&self) -> bool {
370 self.flags.contains(ColumnFlags::PRIMARY_KEY)
371 }
372
373 #[must_use]
375 pub const fn unique(&self) -> bool {
376 self.flags.contains(ColumnFlags::UNIQUE)
377 }
378
379 #[must_use]
381 pub const fn has_default(&self) -> bool {
382 self.flags.contains(ColumnFlags::HAS_DEFAULT)
383 }
384}
385
386#[derive(Clone, Copy, Debug, PartialEq, Eq)]
388pub struct ColumnSqlRef {
389 pub table: &'static str,
390 pub name: &'static str,
391}
392
393impl ColumnSqlRef {
394 #[inline]
396 #[must_use]
397 pub const fn from_column_ref(column: ColumnRef) -> Self {
398 Self {
399 table: column.table,
400 name: column.name,
401 }
402 }
403
404 #[inline]
406 #[must_use]
407 pub const fn from_column_ref_ref(column: &ColumnRef) -> Self {
408 Self {
409 table: column.table,
410 name: column.name,
411 }
412 }
413}
414
415impl From<&ColumnRef> for ColumnSqlRef {
416 #[inline]
417 fn from(value: &ColumnRef) -> Self {
418 Self::from_column_ref_ref(value)
419 }
420}
421
422impl From<ColumnRef> for ColumnSqlRef {
423 #[inline]
424 fn from(value: ColumnRef) -> Self {
425 Self::from_column_ref(value)
426 }
427}
428
429#[inline]
445pub fn write_quoted_ident(buf: &mut impl core::fmt::Write, name: &str) {
446 write_dialect_quoted_ident(Dialect::SQLite, buf, name);
447}
448
449#[inline]
455pub(crate) fn write_dialect_quoted_ident(
456 dialect: Dialect,
457 buf: &mut impl core::fmt::Write,
458 name: &str,
459) {
460 let delimiter = match dialect {
461 Dialect::MySQL => '`',
462 Dialect::SQLite | Dialect::PostgreSQL => '"',
463 };
464
465 let _ = buf.write_char(delimiter);
466 if name.contains(delimiter) {
467 for ch in name.chars() {
468 if ch == delimiter {
469 let _ = buf.write_char(delimiter);
470 let _ = buf.write_char(delimiter);
471 } else {
472 let _ = buf.write_char(ch);
473 }
474 }
475 } else {
476 let _ = buf.write_str(name);
477 }
478 let _ = buf.write_char(delimiter);
479}
480
481#[derive(Clone)]
488pub enum SQLChunk<'a, V: SQLParam> {
489 Token(Token),
491
492 Ident(Cow<'a, str>),
495
496 Raw(Cow<'a, str>),
499
500 Number(usize),
503
504 Param(Param<'a, V>),
507
508 Table(TableSqlRef),
511
512 Column(ColumnSqlRef),
514}
515
516impl<'a, V: SQLParam> SQLChunk<'a, V> {
517 #[inline]
521 #[must_use]
522 pub const fn token(t: Token) -> Self {
523 Self::Token(t)
524 }
525
526 #[inline]
528 #[must_use]
529 pub const fn ident_static(name: &'static str) -> Self {
530 Self::Ident(Cow::Borrowed(name))
531 }
532
533 #[inline]
535 #[must_use]
536 pub const fn raw_static(text: &'static str) -> Self {
537 Self::Raw(Cow::Borrowed(text))
538 }
539
540 #[inline]
542 #[must_use]
543 pub const fn table(table: TableRef) -> Self {
544 Self::Table(TableSqlRef::from_table_ref(table))
545 }
546
547 #[inline]
549 #[must_use]
550 pub const fn column(column: ColumnRef) -> Self {
551 Self::Column(ColumnSqlRef::from_column_ref(column))
552 }
553
554 #[inline]
556 pub const fn param_borrowed(value: &'a V, placeholder: Placeholder) -> Self {
557 Self::Param(Param {
558 value: Some(Cow::Borrowed(value)),
559 placeholder,
560 })
561 }
562
563 #[inline]
567 pub fn ident(name: impl Into<Cow<'a, str>>) -> Self {
568 Self::Ident(name.into())
569 }
570
571 #[inline]
573 pub fn raw(text: impl Into<Cow<'a, str>>) -> Self {
574 Self::Raw(text.into())
575 }
576
577 #[inline]
579 #[must_use]
580 pub const fn number(value: usize) -> Self {
581 Self::Number(value)
582 }
583
584 #[inline]
586 pub fn param(value: impl Into<Cow<'a, V>>, placeholder: Placeholder) -> Self {
587 Self::Param(Param {
588 value: Some(value.into()),
589 placeholder,
590 })
591 }
592
593 #[inline]
597 pub(crate) fn write(&self, buf: &mut impl core::fmt::Write) {
598 match self {
599 SQLChunk::Token(token) => {
600 let _ = buf.write_str(token.as_str());
601 }
602 SQLChunk::Ident(name) => {
603 write_dialect_quoted_ident(V::DIALECT, buf, name);
604 }
605 SQLChunk::Raw(text) => {
606 let _ = buf.write_str(text);
607 }
608 SQLChunk::Number(value) => {
609 let _ = write!(buf, "{value}");
610 }
611 SQLChunk::Param(Param { placeholder, .. }) => {
612 let _ = write!(buf, "{placeholder}");
613 }
614 SQLChunk::Table(t) => {
615 if let Some(schema) = t.schema {
616 write_dialect_quoted_ident(V::DIALECT, buf, schema);
617 let _ = buf.write_char('.');
618 }
619 write_dialect_quoted_ident(V::DIALECT, buf, t.name);
620 }
621 SQLChunk::Column(c) => {
622 write_dialect_quoted_ident(V::DIALECT, buf, c.table);
623 let _ = buf.write_char('.');
624 write_dialect_quoted_ident(V::DIALECT, buf, c.name);
625 }
626 }
627 }
628
629 #[inline]
632 pub(crate) const fn is_word_like(&self) -> bool {
633 match self {
634 SQLChunk::Token(t) => !t.is_punctuation() && !t.is_operator(),
635 SQLChunk::Ident(_)
636 | SQLChunk::Raw(_)
637 | SQLChunk::Number(_)
638 | SQLChunk::Param(_)
639 | SQLChunk::Table(_)
640 | SQLChunk::Column(_) => true,
641 }
642 }
643}
644
645impl<V: SQLParam + core::fmt::Debug> core::fmt::Debug for SQLChunk<'_, V> {
646 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
647 match self {
648 SQLChunk::Token(token) => f.debug_tuple("Token").field(token).finish(),
649 SQLChunk::Ident(name) => f.debug_tuple("Ident").field(name).finish(),
650 SQLChunk::Raw(text) => f.debug_tuple("Raw").field(text).finish(),
651 SQLChunk::Number(value) => f.debug_tuple("Number").field(value).finish(),
652 SQLChunk::Param(param) => f.debug_tuple("Param").field(param).finish(),
653 SQLChunk::Table(t) => f
654 .debug_tuple("Table")
655 .field(&t.schema)
656 .field(&t.name)
657 .finish(),
658 SQLChunk::Column(c) => f
659 .debug_tuple("Column")
660 .field(&format!("{}.{}", c.table, c.name))
661 .finish(),
662 }
663 }
664}
665
666impl<V: SQLParam> From<Token> for SQLChunk<'_, V> {
669 #[inline]
670 fn from(value: Token) -> Self {
671 Self::Token(value)
672 }
673}
674
675impl<V: SQLParam> From<TableRef> for SQLChunk<'_, V> {
676 #[inline]
677 fn from(value: TableRef) -> Self {
678 Self::Table(value.into())
679 }
680}
681
682impl<V: SQLParam> From<ColumnRef> for SQLChunk<'_, V> {
683 #[inline]
684 fn from(value: ColumnRef) -> Self {
685 Self::Column(value.into())
686 }
687}
688
689impl<'a, V: SQLParam> From<Param<'a, V>> for SQLChunk<'a, V> {
690 #[inline]
691 fn from(value: Param<'a, V>) -> Self {
692 Self::Param(value)
693 }
694}
695
696#[cfg(test)]
697mod tests {
698 use super::*;
699 use crate::dialect::{Dialect, MySQLDialect, SQLiteDialect};
700 use core::mem::size_of;
701
702 #[allow(dead_code)]
703 #[derive(Clone, Debug)]
704 struct TestParam([usize; 4]);
705
706 impl SQLParam for TestParam {
707 const DIALECT: Dialect = Dialect::SQLite;
708 type DialectMarker = SQLiteDialect;
709 }
710
711 #[derive(Clone, Debug)]
712 struct MySQLTestParam;
713
714 impl SQLParam for MySQLTestParam {
715 const DIALECT: Dialect = Dialect::MySQL;
716 type DialectMarker = MySQLDialect;
717 }
718
719 #[test]
720 fn sql_chunk_stays_slim() {
721 assert!(size_of::<SQLChunk<'static, TestParam>>() <= 64);
723 }
724
725 #[test]
726 fn quoted_ident_uses_the_dialect_delimiter() {
727 let mut sqlite = String::new();
728 write_dialect_quoted_ident(Dialect::SQLite, &mut sqlite, "account\"owner");
729 assert_eq!(sqlite, "\"account\"\"owner\"");
730
731 let mut postgres = String::new();
732 write_dialect_quoted_ident(Dialect::PostgreSQL, &mut postgres, "account\"owner");
733 assert_eq!(postgres, "\"account\"\"owner\"");
734
735 let mut mysql = String::new();
736 write_dialect_quoted_ident(Dialect::MySQL, &mut mysql, "account`owner");
737 assert_eq!(mysql, "`account``owner`");
738 }
739
740 #[test]
741 fn quoted_ident_keeps_injection_text_inside_the_identifier() {
742 let mut mysql = String::new();
743 write_dialect_quoted_ident(Dialect::MySQL, &mut mysql, "users`; DROP TABLE audit; --");
744 assert_eq!(mysql, "`users``; DROP TABLE audit; --`");
745 }
746
747 #[test]
748 fn table_chunk_preserves_structured_mysql_database_qualification() {
749 let table = TableRef {
750 schema: Some("tenant`db"),
751 ..TableRef::sql("user`accounts", &["id"])
752 };
753 let chunk = SQLChunk::<MySQLTestParam>::table(table);
754 let mut sql = String::new();
755
756 chunk.write(&mut sql);
757
758 assert_eq!(sql, "`tenant``db`.`user``accounts`");
759 }
760}