1use crate::analysis::mutations::Mutation;
3use crate::analysis::resolver::Resolver;
4use crate::analysis::state::AnalysisState;
5use crate::ast::visitor::AstVisitor;
6use crate::engine::config::Config;
7use crate::report::violations::Violation;
8use crate::rules::Rule;
9use crate::rules::conflict::ConflictRule;
10use crate::rules::constraints::BlockingConstraintRule;
11use crate::rules::destructive::{
12 CascadingDropRule, CreateTableAsSelectRule, DropDatabaseRule, DropSchemaCascadeRule,
13 GeneralCascadeRule, ReversibilityRule, SizeAwareAddColumnRule, TypeChangeRewriteRule,
14};
15use crate::rules::drift::DriftDetectionRule;
16use crate::rules::expressions::VolatileDefaultRule;
17use crate::rules::functions::{BrokenComputeRule, FunctionVolatilityRule};
18use crate::rules::idempotency::IdempotencyRule;
19use crate::rules::indexes::ConcurrentIndexRule;
20use crate::rules::opaque::OpaqueDynamicSqlRule;
21use crate::rules::partitions::{PartitionLockRule, PartitionStrategyMismatchRule};
22use crate::rules::policies::RestrictivePolicyRule;
23use crate::rules::security::OverbroadGrantRule;
24use crate::rules::transactions::{
25 AlterTypeAddValueRule, ConcurrentInsideTransactionRule, VacuumFullRule,
26};
27use crate::rules::triggers::DisableTriggerRule;
28use crate::rules::views::MaterializedViewRefreshRule;
29use squawk_syntax::ast::{AstNode, SourceFile};
30use std::collections::HashSet;
31
32pub struct SafeMigrateEngine {
33 config: Config,
34 rules: Vec<Box<dyn Rule>>,
35}
36
37impl SafeMigrateEngine {
38 pub fn new(config: Config) -> Self {
39 Self {
40 config,
41 rules: vec![
42 Box::new(ReversibilityRule),
43 Box::new(DropDatabaseRule),
44 Box::new(DropSchemaCascadeRule),
45 Box::new(GeneralCascadeRule),
46 Box::new(CascadingDropRule),
47 Box::new(CreateTableAsSelectRule),
48 Box::new(SizeAwareAddColumnRule),
49 Box::new(TypeChangeRewriteRule),
50 Box::new(BlockingConstraintRule),
51 Box::new(ConcurrentIndexRule),
52 Box::new(MaterializedViewRefreshRule),
53 Box::new(PartitionLockRule),
54 Box::new(PartitionStrategyMismatchRule),
55 Box::new(RestrictivePolicyRule),
56 Box::new(DisableTriggerRule),
57 Box::new(BrokenComputeRule),
58 Box::new(FunctionVolatilityRule),
59 Box::new(IdempotencyRule),
60 Box::new(ConcurrentInsideTransactionRule),
61 Box::new(AlterTypeAddValueRule),
62 Box::new(VacuumFullRule),
63 Box::new(OpaqueDynamicSqlRule),
64 Box::new(VolatileDefaultRule),
65 Box::new(OverbroadGrantRule),
66 Box::new(DriftDetectionRule),
67 Box::new(ConflictRule),
68 ],
69 }
70 }
71
72 pub fn analyze_chain(
73 &self,
74 files: &[(String, String)],
75 state: &mut AnalysisState,
76 ) -> Result<Vec<Violation>, Vec<String>> {
77 let mut all_violations = Vec::new();
78 for (filename, sql) in files {
79 let violations = self.analyze_single_file(filename, sql, state)?;
80 all_violations.extend(violations);
81 }
82 all_violations.sort_by(|a, b| {
84 a.tier
85 .cmp(&b.tier)
86 .then_with(|| match (&a.source_range, &b.source_range) {
87 (Some(ar), Some(br)) => ar
88 .start()
89 .cmp(&br.start())
90 .then_with(|| ar.end().cmp(&br.end())),
91 (Some(_), None) => std::cmp::Ordering::Less,
92 (None, Some(_)) => std::cmp::Ordering::Greater,
93 (None, None) => std::cmp::Ordering::Equal,
94 })
95 .then_with(|| a.object_name.cmp(&b.object_name))
96 .then_with(|| a.rule_id.cmp(b.rule_id))
97 });
98 Ok(all_violations)
99 }
100
101 pub fn analyze(
102 &self,
103 sql: &str,
104 state: &mut AnalysisState,
105 ) -> Result<Vec<Violation>, Vec<String>> {
106 self.analyze_chain(&[("<inline>".to_string(), sql.to_string())], state)
107 }
108
109 fn analyze_single_file(
110 &self,
111 _filename: &str,
112 sql: &str,
113 state: &mut AnalysisState,
114 ) -> Result<Vec<Violation>, Vec<String>> {
115 let sql = Self::normalize_execute(sql);
116 let parsed = SourceFile::parse(&sql);
117 let errors: Vec<String> = parsed.errors().iter().map(|e| e.to_string()).collect();
118 if !errors.is_empty() {
119 return Err(errors);
120 }
121
122 let mut all_violations = Vec::new();
123 let mut warned_keys = HashSet::new();
124
125 let mut file_ignores = HashSet::new();
126 for token in parsed
127 .tree()
128 .syntax()
129 .descendants_with_tokens()
130 .filter_map(|it| it.into_token())
131 {
132 let mut dummy = HashSet::new();
133 Self::parse_directives(token.text(), &mut file_ignores, &mut dummy);
134 }
135
136 for stmt in parsed.tree().stmts() {
137 let mut stmt_ignores = HashSet::new();
138
139 let mut prev = stmt.syntax().prev_sibling_or_token();
140 while let Some(element) = prev {
141 if element.as_node().is_some() {
142 break;
143 }
144 if let Some(token) = element.as_token() {
145 let mut dummy = HashSet::new();
146 Self::parse_directives(token.text(), &mut dummy, &mut stmt_ignores);
147 }
148 prev = element.prev_sibling_or_token();
149 }
150
151 for token in stmt
152 .syntax()
153 .descendants_with_tokens()
154 .filter_map(|it| it.into_token())
155 {
156 let mut dummy = HashSet::new();
157 Self::parse_directives(token.text(), &mut dummy, &mut stmt_ignores);
158 }
159
160 let stmt_text = Self::strip_sql_leading_comments(&stmt.syntax().text().to_string());
162
163 if let Some(fact) = AstVisitor::extract(&stmt) {
164 let mutations = Resolver::resolve(&fact, state);
165
166 for mutation in mutations {
167 let pre_cascade = match &mutation {
168 Mutation::DropTable(d) if d.cascade => {
169 Some(state.get_cascade_closure(&d.id))
170 }
171 _ => None,
172 };
173
174 let pre_state = state.capture_pre_state();
175 let result = state.apply(&mutation, pre_cascade.as_ref());
176
177 for rule in &self.rules {
178 if file_ignores.contains(rule.id())
179 || stmt_ignores.contains(rule.id())
180 || self.config.is_rule_disabled(rule.id())
181 {
182 continue;
183 }
184
185 let violations = rule.evaluate(
186 &mutation,
187 &result,
188 &pre_state,
189 state,
190 &self.config,
191 pre_cascade.as_ref(),
192 );
193
194 for v in violations {
195 if let Some(key) = &v.dedup_key
196 && !warned_keys.insert(key.clone())
197 {
198 continue;
199 }
200 let mut v = v;
201 if v.source_range.is_none() {
202 let start = stmt
203 .syntax()
204 .descendants_with_tokens()
205 .filter_map(|element| element.into_token())
206 .find(|token| {
207 let text = token.text().trim();
208 !text.is_empty()
209 && !text.starts_with("--")
210 && !text.starts_with("/*")
211 })
212 .map(|token| token.text_range().start())
213 .unwrap_or_else(|| stmt.syntax().text_range().start());
214 let end = stmt.syntax().text_range().end();
215 v.source_range = Some(rowan::TextRange::new(start, end));
216 }
217 if v.sql.is_none() {
218 if let Some(range) = v.source_range {
219 let start = usize::from(range.start());
220 let end = usize::from(range.end());
221 if start < sql.len() && end <= sql.len() {
222 v.sql = Some(sql[start..end].trim().to_string());
223 } else {
224 v.sql = Some(stmt_text.trim().to_string());
225 }
226 } else {
227 v.sql = Some(stmt_text.trim().to_string());
228 }
229 }
230 if state.local.confidence == crate::analysis::state::Confidence::Tainted
233 && v.tier == crate::report::violations::ViolationTier::Tier1
234 {
235 v.tier = crate::report::violations::ViolationTier::Tier2;
236 }
237 all_violations.push(v);
238 }
239 }
240 }
241 }
242 }
243
244 Ok(all_violations)
245 }
246
247 fn normalize_execute(sql: &str) -> String {
252 let mut out = String::with_capacity(sql.len());
253 for line in sql.split_inclusive('\n') {
254 let trimmed = line.trim_start();
255 if trimmed.len() > 9 && trimmed[..9].eq_ignore_ascii_case("EXECUTE '") {
256 let indent = &line[..line.len() - trimmed.len()];
257 out.push_str(indent);
258 out.push_str("DO '");
259 out.push_str(&trimmed[9..]);
260 } else if trimmed.len() > 10 && trimmed[..10].eq_ignore_ascii_case("EXECUTE $$") {
261 let indent = &line[..line.len() - trimmed.len()];
262 out.push_str(indent);
263 out.push_str("DO $$");
264 out.push_str(&trimmed[10..]);
265 } else {
266 out.push_str(line);
267 }
268 }
269 out
270 }
271
272 fn parse_directives(
273 text: &str,
274 file_ignores: &mut HashSet<String>,
275 stmt_ignores: &mut HashSet<String>,
276 ) {
277 let marker = "safe-migrate:";
278 let mut pos = 0;
279
280 while let Some(start) = text[pos..].find(marker) {
281 let after = text[pos + start + marker.len()..].trim_start();
282
283 if let Some(rest) = after.strip_prefix("ignore-file") {
284 let rest = rest.trim_start();
285 if let Some(inner) = rest
286 .strip_prefix('(')
287 .and_then(|s| s.find(')').map(|e| &s[..e]))
288 {
289 file_ignores.insert(inner.trim().to_string());
290 }
291 } else if let Some(rest) = after.strip_prefix("ignore") {
292 let rest = rest.trim_start();
293 if let Some(inner) = rest
294 .strip_prefix('(')
295 .and_then(|s| s.find(')').map(|e| &s[..e]))
296 {
297 stmt_ignores.insert(inner.trim().to_string());
298 }
299 }
300
301 pos = pos + start + marker.len();
302 }
303 }
304
305 fn strip_sql_leading_comments(s: &str) -> String {
306 let mut pos = 0;
307 let bytes = s.as_bytes();
308 while pos < bytes.len() {
309 while pos < bytes.len() && bytes[pos].is_ascii_whitespace() {
310 pos += 1;
311 }
312 if pos + 1 < bytes.len() && bytes[pos] == b'-' && bytes[pos + 1] == b'-' {
313 while pos < bytes.len() && bytes[pos] != b'\n' {
314 pos += 1;
315 }
316 continue;
317 }
318 if pos + 1 < bytes.len() && bytes[pos] == b'/' && bytes[pos + 1] == b'*' {
319 pos += 2;
320 while pos + 1 < bytes.len() && !(bytes[pos] == b'*' && bytes[pos + 1] == b'/') {
321 pos += 1;
322 }
323 if pos + 1 < bytes.len() {
324 pos += 2;
325 }
326 continue;
327 }
328 break;
329 }
330 s[pos..].to_string()
331 }
332}