Skip to main content

safe_migrate/rules/
functions.rs

1use crate::analysis::mutations::Mutation;
2use crate::analysis::state::{AnalysisState, CascadeResult, MutationResult, PreState};
3use crate::engine::config::Config;
4use crate::report::violations::{ObjectKind, OperationKind, Violation, ViolationTier};
5use crate::rules::Rule;
6
7pub struct FunctionVolatilityRule;
8
9impl Rule for FunctionVolatilityRule {
10    fn id(&self) -> &'static str {
11        "function-volatility-change"
12    }
13    fn default_tier(&self) -> ViolationTier {
14        ViolationTier::Tier2
15    }
16    fn recipe(&self) -> &'static str {
17        "Changing a function's volatility (e.g., IMMUTABLE -> VOLATILE) can invalidate existing indexes or change query plan stability."
18    }
19
20    fn evaluate(
21        &self,
22        mutation: &Mutation,
23        _result: &MutationResult,
24        pre_state: &PreState,
25        _state: &AnalysisState,
26        _config: &Config,
27        _cascade: Option<&CascadeResult>,
28    ) -> Vec<Violation> {
29        let mut violations = Vec::new();
30
31        if let Mutation::AlterFunction(alter) = mutation
32            && let Some(old_func) = pre_state.functions.get(&alter.id)
33            && let crate::analysis::facts::AlterFunctionAction::OptionsChange(new_opts) =
34                &alter.action
35        {
36            let ov = old_func.volatility.clone();
37            let new_vol = new_opts.iter().find_map(|opt| {
38                if let crate::analysis::facts::FuncOptionFact::Volatility(v) = opt {
39                    match v {
40                        crate::analysis::facts::VolatilityKind::Volatile => {
41                            Some(crate::model::function::Volatility::Volatile)
42                        }
43                        crate::analysis::facts::VolatilityKind::Stable => {
44                            Some(crate::model::function::Volatility::Stable)
45                        }
46                        crate::analysis::facts::VolatilityKind::Immutable => {
47                            Some(crate::model::function::Volatility::Immutable)
48                        }
49                    }
50                } else {
51                    None
52                }
53            });
54
55            if let Some(nv) = new_vol
56                && ov != nv
57            {
58                violations.push(Violation {
59                    source_range: None,
60                    rule_id: self.id(),
61                    operation_kind: OperationKind::AlterFunction,
62                    object_kind: ObjectKind::Function,
63                    object_name: alter.id.to_string(),
64                    tier: self.default_tier(),
65                    reason: format!(
66                        "Function {} volatility changed from {:?} to {:?}",
67                        alter.id, ov, nv
68                    ),
69                    recipe: self.recipe(),
70                    dedup_key: None,
71                    sql: None,
72                    fk_dependency_related: false,
73                });
74            }
75        }
76
77        violations
78    }
79}
80
81pub struct BrokenComputeRule;
82
83impl Rule for BrokenComputeRule {
84    fn id(&self) -> &'static str {
85        "broken-compute"
86    }
87    fn default_tier(&self) -> ViolationTier {
88        ViolationTier::Tier1
89    }
90    fn recipe(&self) -> &'static str {
91        "Dropping a function used by a trigger will cause the trigger to fail at runtime."
92    }
93
94    fn evaluate(
95        &self,
96        mutation: &Mutation,
97        result: &MutationResult,
98        _pre_state: &PreState,
99        state: &AnalysisState,
100        _config: &Config,
101        _cascade_closure: Option<&CascadeResult>,
102    ) -> Vec<Violation> {
103        if *result == MutationResult::Skipped {
104            return vec![];
105        }
106        if let Mutation::DropFunction(drop) = mutation {
107            for sig in &drop.signatures {
108                // Construct ID in same way as during creation
109                let sig_str = format!("{}({})", sig.name.name.resolve(), sig.params.join(","));
110                let schema = state.resolve_function_schema(&sig.name, &sig_str);
111                let function_id = crate::ast::identifiers::ObjectId::new(schema, sig_str);
112
113                println!("DEBUG BrokenComputeRule function_id: {:?}", function_id);
114                let affected = state.local.graph.triggers_for_function(&function_id);
115
116                if !affected.is_empty() {
117                    let triggers_info: Vec<String> = affected
118                        .iter()
119                        .map(|t| format!("trigger {} on table {}", t.trigger_id, t.table_id))
120                        .collect();
121
122                    return vec![Violation {
123                        source_range: None,
124                        rule_id: self.id(),
125                        operation_kind: OperationKind::DropFunction,
126                        object_kind: ObjectKind::Function,
127                        object_name: function_id.to_string(),
128                        tier: self.default_tier(),
129                        reason: format!(
130                            "Broken Compute: Dropping Function Used by Trigger: {}",
131                            triggers_info.join(", ")
132                        ),
133                        recipe: self.recipe(),
134                        dedup_key: None,
135                        sql: None,
136                        fk_dependency_related: false,
137                    }];
138                }
139            }
140        }
141        vec![]
142    }
143}