Skip to main content

safe_migrate/engine/
engine.rs

1// FILE: src/engine/engine.rs
2use 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        // Phase 10.6: Deterministic violation ordering
83        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            // Capture raw statement text for sql field on violations (strip leading comments)
161            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                            // Downgrade tier at push time if confidence was already tainted
231                            // BEFORE this mutation was applied.
232                            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    /// Pre-process SQL to handle EXECUTE '...' which Squawk's parser does not
248    /// recognize (top-level EXECUTE expects a prepared-statement name, not a
249    /// string literal).  Rewriting to DO '...' is semantically equivalent at
250    /// the top level and lets the parser produce a proper DoBlock node.
251    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}