safe_migrate/engine/
engine.rs1use 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}