Skip to main content

sql_schema/
lib.rs

1use std::fmt;
2
3use self::ast::Statement;
4
5pub use self::{
6    diff::TreeDiffer,
7    migration::TreeMigrator,
8    parser::{Parse, ParseError},
9};
10
11mod ast;
12pub mod dialect;
13mod diff;
14mod migration;
15pub mod name_gen;
16mod parser;
17pub mod path_template;
18mod sealed;
19
20#[derive(Debug, Clone)]
21pub struct SyntaxTree<Dialect> {
22    dialect: Dialect,
23    pub(crate) tree: Vec<Statement>,
24}
25
26impl<Dialect: Default> SyntaxTree<Dialect> {
27    pub fn empty() -> Self {
28        Self {
29            dialect: Default::default(),
30            tree: Vec::with_capacity(0),
31        }
32    }
33}
34
35impl<Dialect> SyntaxTree<Dialect>
36where
37    Dialect: Parse,
38{
39    pub fn parse<'a>(dialect: Dialect, sql: impl Into<&'a str>) -> Result<Self, ParseError> {
40        let tree = dialect.parse_sql::<Dialect>(sql)?;
41        Ok(Self { dialect, tree })
42    }
43}
44
45pub use diff::DiffError;
46pub use migration::MigrateError;
47
48impl<Dialect> SyntaxTree<Dialect>
49where
50    Dialect: TreeDiffer,
51{
52    pub fn diff(&self, other: &SyntaxTree<Dialect>) -> Result<Option<Self>, DiffError> {
53        Ok(
54            TreeDiffer::diff_tree(&self.dialect, &self.tree, &other.tree)?.map(|tree| Self {
55                dialect: self.dialect.clone(),
56                tree,
57            }),
58        )
59    }
60}
61
62impl<Dialect> SyntaxTree<Dialect>
63where
64    Dialect: TreeMigrator,
65{
66    pub fn migrate(self, other: &SyntaxTree<Dialect>) -> Result<Self, MigrateError> {
67        let tree = TreeMigrator::migrate_tree(&self.dialect, self.tree, &other.tree)?;
68        Ok(Self {
69            dialect: self.dialect.clone(),
70            tree,
71        })
72    }
73}
74
75impl<Dialect> fmt::Display for SyntaxTree<Dialect> {
76    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
77        let mut iter = self.tree.iter().peekable();
78        while let Some(s) = iter.next() {
79            let formatted = sqlformat::format(
80                format!("{s};").as_str(),
81                &sqlformat::QueryParams::None,
82                &sqlformat::FormatOptions::default(),
83            );
84            write!(f, "{formatted}")?;
85            if iter.peek().is_some() {
86                write!(f, "\n\n")?;
87            }
88        }
89        Ok(())
90    }
91}
92
93#[cfg(test)]
94mod tests {
95    use super::dialect::Generic;
96    use super::*;
97
98    macro_rules! test_case {
99        (
100            @dialect($dialect:ty) $(,)?
101
102            $(
103                $test_name:ident { $( $field:ident : $value:literal ),+ $(,)? }
104            ),* $(,)?
105
106            => $test_fn:expr $(,)?
107        ) => {
108            $(
109                #[test]
110                fn $test_name() {
111                    let dialect = <$dialect>::default();
112
113                    let test_case: TestCase<$dialect> = TestCase {
114                        dialect: dialect.clone(),
115                        $( $field : $value ),+
116                    };
117
118                    run_test_case(&test_case, $test_fn);
119                }
120            )*
121        };
122    }
123
124    #[derive(Debug)]
125    struct TestCase<Dialect = Generic> {
126        dialect: Dialect,
127        sql_a: &'static str,
128        sql_b: &'static str,
129        expect: &'static str,
130    }
131
132    fn run_test_case<F, E, Dialect>(tc: &TestCase<Dialect>, testfn: F)
133    where
134        Dialect: Parse + TreeDiffer,
135        E: std::error::Error,
136        F: Fn(SyntaxTree<Dialect>, SyntaxTree<Dialect>) -> Result<Option<SyntaxTree<Dialect>>, E>,
137    {
138        let dialect = tc.dialect.clone();
139        let ast_a = SyntaxTree::parse(dialect.clone(), tc.sql_a).unwrap();
140        let ast_b = SyntaxTree::parse(dialect.clone(), tc.sql_b).unwrap();
141        SyntaxTree::parse(dialect, tc.expect)
142            .unwrap_or_else(|_| panic!("invalid SQL: {:?}", tc.expect));
143        let actual = testfn(ast_a, ast_b)
144            .inspect_err(|err| eprintln!("Error: {err:?}"))
145            .unwrap()
146            .unwrap();
147        assert_eq!(actual.to_string(), tc.expect, "{tc:?}");
148    }
149
150    mod test_diff {
151        use super::*;
152
153        test_case!(
154            @dialect(Generic)
155
156            create_table_a {
157                sql_a: "CREATE TABLE foo(\
158                    id int PRIMARY KEY
159                )",
160                sql_b: "CREATE TABLE foo(\
161                    id int PRIMARY KEY
162                );\
163                    CREATE TABLE bar (id INT PRIMARY KEY);",
164                expect: "CREATE TABLE bar (id INT PRIMARY KEY);",
165            },
166
167            create_table_b {
168                sql_a: "CREATE TABLE foo(\
169                    id int PRIMARY KEY
170                )",
171                sql_b: "CREATE TABLE foo(\
172                    \"id\" int PRIMARY KEY
173                );\
174                    CREATE TABLE bar (id INT PRIMARY KEY);",
175                expect: "CREATE TABLE bar (id INT PRIMARY KEY);",
176            },
177
178            create_table_c {
179                sql_a: "CREATE TABLE foo(\
180                    \"id\" int PRIMARY KEY
181                )",
182                sql_b: "CREATE TABLE foo(\
183                    id int PRIMARY KEY
184                );\
185                    CREATE TABLE bar (id INT PRIMARY KEY);",
186                expect: "CREATE TABLE bar (id INT PRIMARY KEY);",
187            },
188
189            drop_table_a {
190                sql_a: "CREATE TABLE foo(\
191                    id int PRIMARY KEY
192                );\
193                    CREATE TABLE bar (id INT PRIMARY KEY);",
194                sql_b: "CREATE TABLE foo(\
195                    id int PRIMARY KEY
196                )",
197                expect: "DROP TABLE bar;",
198            },
199
200            add_column_a {
201                sql_a: "CREATE TABLE foo(\
202                    id int PRIMARY KEY
203                )",
204                sql_b: "CREATE TABLE foo(\
205                    id int PRIMARY KEY,
206                    bar text
207                )",
208                expect: "ALTER TABLE\n  foo\nADD\n  COLUMN bar TEXT;",
209            },
210
211            drop_column_a {
212                sql_a: "CREATE TABLE foo(\
213                    id int PRIMARY KEY,
214                    bar text
215                )",
216                sql_b: "CREATE TABLE foo(\
217                    id int PRIMARY KEY
218                )",
219                expect: "ALTER TABLE\n  foo DROP COLUMN bar;",
220            },
221
222            create_index_a {
223                sql_a: "CREATE UNIQUE INDEX title_idx ON films (title);",
224                sql_b: "CREATE UNIQUE INDEX title_idx ON films ((lower(title)));",
225                expect: "DROP INDEX title_idx;\n\nCREATE UNIQUE INDEX title_idx ON films((lower(title)));",
226            },
227
228            create_index_b {
229                sql_a: "CREATE UNIQUE INDEX IF NOT EXISTS title_idx ON films (title);",
230                sql_b: "CREATE UNIQUE INDEX IF NOT EXISTS title_idx ON films ((lower(title)));",
231                expect: "DROP INDEX IF EXISTS title_idx;\n\nCREATE UNIQUE INDEX IF NOT EXISTS title_idx ON films((lower(title)));",
232            },
233
234            create_type_a {
235                sql_a: "CREATE TYPE bug_status AS ENUM ('new', 'open');",
236                sql_b: "CREATE TYPE foo AS ENUM ('bar');",
237                expect: "DROP TYPE bug_status;\n\nCREATE TYPE foo AS ENUM ('bar');",
238            },
239
240            create_type_b {
241                sql_a: "CREATE TYPE bug_status AS ENUM ('new', 'open', 'closed');",
242                sql_b: "CREATE TYPE bug_status AS ENUM ('new', 'open', 'assigned', 'closed');",
243                expect: "ALTER TYPE bug_status\nADD\n  VALUE 'assigned'\nAFTER\n  'open';",
244            },
245
246            create_type_c {
247                sql_a: "CREATE TYPE bug_status AS ENUM ('open', 'closed');",
248                sql_b: "CREATE TYPE bug_status AS ENUM ('new', 'open', 'closed');",
249                expect: "ALTER TYPE bug_status\nADD\n  VALUE 'new' BEFORE 'open';",
250            },
251
252            create_type_d {
253                sql_a: "CREATE TYPE bug_status AS ENUM ('new', 'open');",
254                sql_b: "CREATE TYPE bug_status AS ENUM ('new', 'open', 'closed');",
255                expect: "ALTER TYPE bug_status\nADD\n  VALUE 'closed';",
256            },
257
258            create_type_e {
259                sql_a: "CREATE TYPE bug_status AS ENUM ('new', 'open');",
260                sql_b: "CREATE TYPE bug_status AS ENUM ('new', 'open', 'assigned', 'closed');",
261                expect: "ALTER TYPE bug_status\nADD\n  VALUE 'assigned';\n\nALTER TYPE bug_status\nADD\n  VALUE 'closed';",
262            },
263
264            create_type_f {
265                sql_a: "CREATE TYPE bug_status AS ENUM ('open', 'critical');",
266                sql_b: "CREATE TYPE bug_status AS ENUM ('new', 'open', 'assigned', 'closed', 'critical');",
267                expect: "ALTER TYPE bug_status\nADD\n  VALUE 'new' BEFORE 'open';\n\nALTER TYPE bug_status\nADD\n  VALUE 'assigned'\nAFTER\n  'open';\n\nALTER TYPE bug_status\nADD\n  VALUE 'closed'\nAFTER\n  'assigned';",
268            },
269
270            create_type_g {
271                sql_a: "CREATE TYPE bug_status AS ENUM ('open');",
272                sql_b: "CREATE TYPE bug_status AS ENUM ('new', 'open', 'closed');",
273                expect: "ALTER TYPE bug_status\nADD\n  VALUE 'new' BEFORE 'open';\n\nALTER TYPE bug_status\nADD\n  VALUE 'closed';",
274            },
275
276            create_extension_a {
277                sql_a: "CREATE EXTENSION hstore;",
278                sql_b: "CREATE EXTENSION IF NOT EXISTS \"uuid-ossp\";",
279                expect: "DROP EXTENSION hstore;\n\nCREATE EXTENSION IF NOT EXISTS \"uuid-ossp\";",
280            },
281
282            => |ast_a, ast_b| {
283                ast_a.diff(&ast_b)
284            }
285        );
286
287        test_case!(
288            @dialect(Generic)
289
290            create_domain_a {
291                sql_a: "",
292                sql_b: "CREATE DOMAIN email AS VARCHAR(255) CHECK (VALUE ~ '^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\\.[a-zA-Z]{2,}$');",
293                expect: "CREATE DOMAIN email AS VARCHAR(255) CHECK (\n  VALUE ~ '^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\\.[a-zA-Z]{2,}$'\n);",
294            },
295
296            edit_domain_a {
297                sql_a: "CREATE DOMAIN positive_int AS INTEGER CHECK (VALUE > 0);",
298                sql_b: "CREATE DOMAIN positive_int AS BIGINT CHECK (VALUE > 0 AND VALUE < 1000000);",
299                expect: "DROP DOMAIN IF EXISTS positive_int;\n\nCREATE DOMAIN positive_int AS BIGINT CHECK (\n  VALUE > 0\n  AND VALUE < 1000000\n);",
300            },
301
302            => |ast_a, ast_b| {
303                ast_a.diff(&ast_b)
304            }
305        );
306    }
307
308    mod migrate {
309        use crate::dialect::PostgreSQL;
310
311        use super::*;
312
313        test_case!(
314            @dialect(Generic)
315
316            create_table_a {
317                sql_a: "CREATE TABLE bar (id INT PRIMARY KEY);",
318                sql_b: "CREATE TABLE foo (id INT PRIMARY KEY);",
319                expect: "CREATE TABLE bar (id INT PRIMARY KEY);\n\nCREATE TABLE foo (id INT PRIMARY KEY);",
320            },
321
322            drop_table_a {
323                sql_a: "CREATE TABLE bar (id INT PRIMARY KEY)",
324                sql_b: "DROP TABLE bar; CREATE TABLE foo (id INT PRIMARY KEY)",
325                expect: "CREATE TABLE foo (id INT PRIMARY KEY);",
326            },
327
328            alter_table_add_column_a {
329                sql_a: "CREATE TABLE bar (id INT PRIMARY KEY)",
330                sql_b: "ALTER TABLE bar ADD COLUMN bar TEXT",
331                expect: "CREATE TABLE bar (id INT PRIMARY KEY, bar TEXT);",
332            },
333
334            alter_table_drop_column_a {
335                sql_a: "CREATE TABLE bar (bar TEXT, id INT PRIMARY KEY)",
336                sql_b: "ALTER TABLE bar DROP COLUMN bar",
337                expect: "CREATE TABLE bar (id INT PRIMARY KEY);",
338            },
339
340            alter_table_alter_column_a {
341                sql_a: "CREATE TABLE bar (bar TEXT, id INT PRIMARY KEY)",
342                sql_b: "ALTER TABLE bar ALTER COLUMN bar SET NOT NULL",
343                expect: "CREATE TABLE bar (bar TEXT NOT NULL, id INT PRIMARY KEY);",
344            },
345
346            alter_table_alter_column_b {
347                sql_a: "CREATE TABLE bar (bar TEXT NOT NULL, id INT PRIMARY KEY)",
348                sql_b: "ALTER TABLE bar ALTER COLUMN bar DROP NOT NULL",
349                expect: "CREATE TABLE bar (bar TEXT, id INT PRIMARY KEY);",
350            },
351
352            alter_table_alter_column_c {
353                sql_a: "CREATE TABLE bar (bar TEXT NOT NULL DEFAULT 'foo', id INT PRIMARY KEY)",
354                sql_b: "ALTER TABLE bar ALTER COLUMN bar DROP DEFAULT",
355                expect: "CREATE TABLE bar (bar TEXT NOT NULL, id INT PRIMARY KEY);",
356            },
357
358            alter_table_alter_column_d {
359                sql_a: "CREATE TABLE bar (bar TEXT, id INT PRIMARY KEY)",
360                sql_b: "ALTER TABLE bar ALTER COLUMN bar SET DATA TYPE INTEGER",
361                expect: "CREATE TABLE bar (bar INTEGER, id INT PRIMARY KEY);",
362            },
363
364            alter_table_alter_column_f {
365                sql_a: "CREATE TABLE bar (bar INTEGER, id INT PRIMARY KEY)",
366                sql_b: "ALTER TABLE bar ALTER COLUMN bar ADD GENERATED BY DEFAULT AS IDENTITY",
367                expect: "CREATE TABLE bar (\n  bar INTEGER GENERATED BY DEFAULT AS IDENTITY,\n  id INT PRIMARY KEY\n);",
368            },
369
370            alter_table_alter_column_g {
371                sql_a: "CREATE TABLE bar (bar INTEGER, id INT PRIMARY KEY)",
372                sql_b: "ALTER TABLE bar ALTER COLUMN bar ADD GENERATED ALWAYS AS IDENTITY (START WITH 10)",
373                expect: "CREATE TABLE bar (\n  bar INTEGER GENERATED ALWAYS AS IDENTITY (START WITH 10),\n  id INT PRIMARY KEY\n);",
374            },
375
376            create_index_a {
377                sql_a: "CREATE UNIQUE INDEX title_idx ON films (title);",
378                sql_b: "CREATE INDEX code_idx ON films (code);",
379                expect: "CREATE UNIQUE INDEX title_idx ON films(title);\n\nCREATE INDEX code_idx ON films(code);",
380            },
381
382            drop_index_a {
383                sql_a: "CREATE UNIQUE INDEX title_idx ON films (title);",
384                sql_b: "DROP INDEX title_idx;",
385                expect: "",
386            },
387
388            drop_index_b {
389                sql_a: "CREATE UNIQUE INDEX title_idx ON films (title);",
390                sql_b: "DROP INDEX title_idx;CREATE INDEX code_idx ON films (code);",
391                expect: "CREATE INDEX code_idx ON films(code);",
392            },
393
394            create_type_a {
395                sql_a: "CREATE TYPE bug_status AS ENUM ('open', 'closed');",
396                sql_b: "CREATE TYPE compfoo AS (f1 int, f2 text);",
397                expect: "CREATE TYPE bug_status AS ENUM ('open', 'closed');\n\nCREATE TYPE compfoo AS (f1 INT, f2 TEXT);",
398            },
399
400            drop_type_a {
401                sql_a: "CREATE TYPE bug_status AS ENUM ('open', 'closed'); CREATE TYPE compfoo AS (f1 int, f2 text);",
402                sql_b: "DROP TYPE bug_status;",
403                expect: "CREATE TYPE compfoo AS (f1 INT, f2 TEXT);",
404            },
405
406            alter_type_rename_a {
407                sql_a: "CREATE TYPE bug_status AS ENUM ('open', 'closed');",
408                sql_b: "ALTER TYPE bug_status RENAME TO issue_status",
409                expect: "CREATE TYPE issue_status AS ENUM ('open', 'closed');",
410            },
411
412            alter_type_add_value_a {
413                sql_a: "CREATE TYPE bug_status AS ENUM ('open');",
414                sql_b: "ALTER TYPE bug_status ADD VALUE 'new' BEFORE 'open';",
415                expect: "CREATE TYPE bug_status AS ENUM ('new', 'open');",
416            },
417
418            alter_type_add_value_b {
419                sql_a: "CREATE TYPE bug_status AS ENUM ('open');",
420                sql_b: "ALTER TYPE bug_status ADD VALUE 'closed' AFTER 'open';",
421                expect: "CREATE TYPE bug_status AS ENUM ('open', 'closed');",
422            },
423
424            alter_type_add_value_c {
425                sql_a: "CREATE TYPE bug_status AS ENUM ('open');",
426                sql_b: "ALTER TYPE bug_status ADD VALUE 'closed';",
427                expect: "CREATE TYPE bug_status AS ENUM ('open', 'closed');",
428            },
429
430            alter_type_rename_value_a {
431                sql_a: "CREATE TYPE bug_status AS ENUM ('new', 'closed');",
432                sql_b: "ALTER TYPE bug_status RENAME VALUE 'new' TO 'open';",
433                expect: "CREATE TYPE bug_status AS ENUM ('open', 'closed');",
434            },
435
436            create_extension_a {
437                sql_a: "CREATE EXTENSION hstore;",
438                sql_b: "CREATE EXTENSION IF NOT EXISTS \"uuid-ossp\";",
439                expect: "CREATE EXTENSION hstore;\n\nCREATE EXTENSION IF NOT EXISTS \"uuid-ossp\";",
440            },
441
442            drop_extension_a {
443                sql_a: "CREATE EXTENSION hstore; CREATE EXTENSION IF NOT EXISTS \"uuid-ossp\";",
444                sql_b: "DROP EXTENSION hstore;",
445                expect: "CREATE EXTENSION IF NOT EXISTS \"uuid-ossp\";",
446            },
447
448            => |ast_a, ast_b| {
449                Some(ast_a.migrate(&ast_b)).transpose()
450            }
451        );
452
453        test_case!(
454            @dialect(PostgreSQL)
455
456            alter_table_alter_column_e {
457                sql_a: "CREATE TABLE bar (bar TEXT, id INT PRIMARY KEY)",
458                sql_b: "ALTER TABLE bar ALTER COLUMN bar SET DATA TYPE timestamp with time zone\n USING timestamp with time zone 'epoch' + foo_timestamp * interval '1 second'",
459                expect: "CREATE TABLE bar (bar TIMESTAMP WITH TIME ZONE, id INT PRIMARY KEY);",
460            },
461
462            create_domain_a {
463                sql_a: "CREATE DOMAIN positive_int AS INTEGER CHECK (VALUE > 0);",
464                sql_b: "CREATE DOMAIN email AS VARCHAR(255) CHECK (VALUE ~ '^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\\.[a-zA-Z]{2,}$');",
465                expect: "CREATE DOMAIN positive_int AS INTEGER CHECK (VALUE > 0);\n\nCREATE DOMAIN email AS VARCHAR(255) CHECK (\n  VALUE ~ '^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\\.[a-zA-Z]{2,}$'\n);",
466            },
467
468            drop_domain_a {
469                sql_a: "CREATE DOMAIN positive_int AS INTEGER CHECK (VALUE > 0); CREATE DOMAIN above_ten AS INTEGER CHECK (VALUE > 10);",
470                sql_b: "DROP DOMAIN above_ten;",
471                expect: "CREATE DOMAIN positive_int AS INTEGER CHECK (VALUE > 0);",
472            },
473
474            => |ast_a, ast_b| {
475                Some(ast_a.migrate(&ast_b)).transpose()
476            }
477        );
478    }
479}