Skip to main content

safe_migrate/engine/
engine.rs

1// FILE: src/engine/engine.rs
2use crate::analysis::mutations::{AlterTableActionMutation, Mutation};
3use crate::analysis::resolver::Resolver;
4use crate::analysis::state::AnalysisState;
5use crate::ast::identifiers::ObjectId;
6use crate::ast::visitor::AstVisitor;
7use crate::engine::config::Config;
8use crate::model::relation::{RelationOverlay, RelationState};
9use crate::report::violations::Violation;
10use crate::rules::Rule;
11use crate::rules::constraints::BlockingConstraintRule;
12use crate::rules::destructive::{CascadingDropRule, SizeAwareAddColumnRule, TypeChangeRewriteRule};
13use crate::rules::expressions::VolatileDefaultRule;
14use crate::rules::idempotency::IdempotencyRule;
15use crate::rules::indexes::ConcurrentIndexRule;
16use crate::rules::opaque::OpaqueDynamicSqlRule;
17use crate::rules::partitions::PartitionLockRule;
18use crate::rules::transactions::{ConcurrentInsideTransactionRule, VacuumFullRule};
19use crate::rules::views::MaterializedViewRefreshRule;
20use squawk_syntax::ast::{AstNode, SourceFile};
21use std::collections::{HashMap, HashSet};
22
23pub struct SafeMigrateEngine {
24    config: Config,
25    rules: Vec<Box<dyn Rule>>,
26}
27
28impl SafeMigrateEngine {
29    pub fn new(config: Config) -> Self {
30        Self {
31            config,
32            rules: vec![
33                Box::new(CascadingDropRule),
34                Box::new(SizeAwareAddColumnRule),
35                Box::new(TypeChangeRewriteRule),
36                Box::new(BlockingConstraintRule),
37                Box::new(ConcurrentIndexRule),
38                Box::new(MaterializedViewRefreshRule),
39                Box::new(PartitionLockRule),
40                Box::new(IdempotencyRule),
41                Box::new(ConcurrentInsideTransactionRule),
42                Box::new(VacuumFullRule),
43                Box::new(OpaqueDynamicSqlRule),
44                Box::new(VolatileDefaultRule),
45            ],
46        }
47    }
48
49    fn parse_directives(
50        text: &str,
51        file_ignores: &mut HashSet<String>,
52        stmt_ignores: &mut HashSet<String>,
53    ) {
54        let mut search = text;
55        while let Some(idx) = search.find("safe-migrate: ignore-file(") {
56            let start = idx + "safe-migrate: ignore-file(".len();
57            if let Some(end) = search[start..].find(')') {
58                file_ignores.insert(search[start..start + end].trim().to_string());
59                search = &search[start + end + 1..];
60            } else {
61                break;
62            }
63        }
64
65        let mut search = text;
66        while let Some(idx) = search.find("safe-migrate: ignore(") {
67            let start = idx + "safe-migrate: ignore(".len();
68            if let Some(end) = search[start..].find(')') {
69                stmt_ignores.insert(search[start..start + end].trim().to_string());
70                search = &search[start + end + 1..];
71            } else {
72                break;
73            }
74        }
75    }
76
77    pub fn analyze(
78        &self,
79        sql: &str,
80        state: &mut AnalysisState,
81    ) -> Result<Vec<Violation>, Vec<String>> {
82        let parsed = SourceFile::parse(sql);
83        let errors: Vec<String> = parsed.errors().iter().map(|e| e.to_string()).collect();
84        if !errors.is_empty() {
85            return Err(errors);
86        }
87
88        let mut all_violations = Vec::new();
89        let mut warned_keys = HashSet::new();
90
91        let mut file_ignores = HashSet::new();
92        for token in parsed
93            .tree()
94            .syntax()
95            .descendants_with_tokens()
96            .filter_map(|it| it.into_token())
97        {
98            let mut dummy = HashSet::new();
99            Self::parse_directives(token.text(), &mut file_ignores, &mut dummy);
100        }
101
102        for stmt in parsed.tree().stmts() {
103            let mut stmt_ignores = HashSet::new();
104
105            let mut prev = stmt.syntax().prev_sibling_or_token();
106            while let Some(element) = prev {
107                if element.as_node().is_some() {
108                    break;
109                }
110                if let Some(token) = element.as_token() {
111                    let mut dummy = HashSet::new();
112                    Self::parse_directives(token.text(), &mut dummy, &mut stmt_ignores);
113                }
114                prev = element.prev_sibling_or_token();
115            }
116
117            for token in stmt
118                .syntax()
119                .descendants_with_tokens()
120                .filter_map(|it| it.into_token())
121            {
122                let mut dummy = HashSet::new();
123                Self::parse_directives(token.text(), &mut dummy, &mut stmt_ignores);
124            }
125
126            if let Some(fact) = AstVisitor::extract(&stmt) {
127                let mutations = Resolver::resolve(&fact, state);
128
129                for mutation in mutations {
130                    let pre_relations: HashMap<ObjectId, RelationState> = match &mutation {
131                        Mutation::AlterTable(a) => {
132                            let mut snapshot = HashMap::new();
133                            if let Some(RelationOverlay::Present(r)) =
134                                state.local.relations.get(&a.id)
135                            {
136                                snapshot.insert(a.id.clone(), r.clone());
137                            }
138
139                            if let AlterTableActionMutation::AddForeignKey { to_table, .. } =
140                                &a.action
141                                && let Some(RelationOverlay::Present(parent_rel)) =
142                                    state.local.relations.get(to_table)
143                            {
144                                snapshot.insert(to_table.clone(), parent_rel.clone());
145                            }
146                            snapshot
147                        }
148                        Mutation::CreateIndex(c) => state
149                            .local
150                            .relations
151                            .get(&c.table)
152                            .and_then(|o| {
153                                if let RelationOverlay::Present(r) = o {
154                                    Some((c.table.clone(), r.clone()))
155                                } else {
156                                    None
157                                }
158                            })
159                            .into_iter()
160                            .collect(),
161                        Mutation::RefreshMaterializedView(r) => state
162                            .local
163                            .relations
164                            .get(&r.id)
165                            .and_then(|o| {
166                                if let RelationOverlay::Present(rel) = o {
167                                    Some((r.id.clone(), rel.clone()))
168                                } else {
169                                    None
170                                }
171                            })
172                            .into_iter()
173                            .collect(),
174                        Mutation::DropIndex(d) => state
175                            .local
176                            .graph
177                            .is_referenced_by_index(&d.id)
178                            .into_iter()
179                            .filter_map(|tid| {
180                                state.local.relations.get(tid).and_then(|o| {
181                                    if let RelationOverlay::Present(r) = o {
182                                        Some((tid.clone(), r.clone()))
183                                    } else {
184                                        None
185                                    }
186                                })
187                            })
188                            .collect(),
189                        _ => HashMap::new(),
190                    };
191
192                    let pre_cascade = match &mutation {
193                        Mutation::DropTable(d) if d.cascade => {
194                            Some(state.get_cascade_closure(&d.id))
195                        }
196                        _ => None,
197                    };
198
199                    let result = state.apply(&mutation, pre_cascade.as_ref());
200
201                    for rule in &self.rules {
202                        if file_ignores.contains(rule.id())
203                            || stmt_ignores.contains(rule.id())
204                            || self.config.is_rule_disabled(rule.id())
205                        {
206                            continue;
207                        }
208
209                        let violations = rule.evaluate(
210                            &mutation,
211                            &result,
212                            &pre_relations,
213                            state,
214                            &self.config,
215                            pre_cascade.as_ref(),
216                        );
217
218                        for v in violations {
219                            if let Some(key) = &v.dedup_key
220                                && !warned_keys.insert(key.clone())
221                            {
222                                continue;
223                            }
224                            all_violations.push(v);
225                        }
226                    }
227                }
228            }
229        }
230
231        Ok(all_violations)
232    }
233}