Skip to main content

systemprompt_database/lifecycle/installation/
migration_refs.rs

1//! Refuses, before any database write, a migration that names a trigger or
2//! view only a declarative schema file creates.
3//!
4//! Migrations run before the dependent phase that applies declarative
5//! triggers and views, so such a reference works on every database that has
6//! already booted on the schema shipping the object and fails the first
7//! upgrade from one that has not — which is never the author's workstation.
8//! Functions need no check: the routine pre-pass applies every declarative
9//! function before any migration runs. `DO $$ … $$` bodies are opaque to the
10//! parser, which makes a catalog-guarded reference inside one the sanctioned
11//! way for a migration to touch a declarative object.
12//!
13//! Copyright (c) systemprompt.io — Business Source License 1.1.
14//! See <https://systemprompt.io> for licensing details.
15
16use std::collections::HashSet;
17use std::sync::Arc;
18
19use pg_query::protobuf::{AlterTableType, ObjectType};
20use pg_query::{Context, NodeEnum};
21use systemprompt_extension::{Extension, LoaderError};
22use tracing::warn;
23
24#[derive(Debug, Clone, Copy, PartialEq, Eq)]
25enum ObjectKind {
26    Trigger,
27    View,
28}
29
30impl ObjectKind {
31    const fn label(self) -> &'static str {
32        match self {
33            Self::Trigger => "trigger",
34            Self::View => "view",
35        }
36    }
37}
38
39#[derive(Default)]
40struct Objects {
41    triggers: HashSet<String>,
42    views: HashSet<String>,
43}
44
45impl Objects {
46    fn contains(&self, kind: ObjectKind, name: &str) -> bool {
47        match kind {
48            ObjectKind::Trigger => self.triggers.contains(name),
49            ObjectKind::View => self.views.contains(name),
50        }
51    }
52
53    fn absorb(&mut self, other: Self) {
54        self.triggers.extend(other.triggers);
55        self.views.extend(other.views);
56    }
57}
58
59struct Reference {
60    kind: ObjectKind,
61    name: String,
62    how: &'static str,
63}
64
65struct ParsedMigration {
66    extension: String,
67    migration: String,
68    creates: Objects,
69    references: Vec<Reference>,
70}
71
72pub fn check_migration_references(extensions: &[Arc<dyn Extension>]) -> Result<(), LoaderError> {
73    let mut declared = Objects::default();
74    let mut migrated = Objects::default();
75    let mut migrations = Vec::new();
76
77    for ext in extensions {
78        let extension = ext.id().to_owned();
79        for schema in ext.schemas() {
80            declared.absorb(created_objects(&extension, &schema.sql)?);
81        }
82        for migration in ext.migrations().into_iter().filter(|m| !m.tombstone) {
83            let label = format!("{:03}_{}", migration.version, migration.name);
84            let parsed = pg_query::parse(migration.sql).map_err(|e| {
85                LoaderError::SchemaInstallationFailed {
86                    extension: extension.clone(),
87                    message: format!("migration {label}: SQL parse failed: {e}"),
88                }
89            })?;
90            let creates = created_objects_of(&parsed);
91            migrated.triggers.extend(creates.triggers.iter().cloned());
92            migrated.views.extend(creates.views.iter().cloned());
93            migrations.push(ParsedMigration {
94                extension: extension.clone(),
95                migration: label,
96                creates,
97                references: references_of(&parsed),
98            });
99        }
100    }
101
102    let mut first: Option<LoaderError> = None;
103    for m in &migrations {
104        for r in &m.references {
105            if !declared.contains(r.kind, &r.name)
106                || migrated.contains(r.kind, &r.name)
107                || m.creates.contains(r.kind, &r.name)
108            {
109                continue;
110            }
111            let err = LoaderError::MigrationReferencesDeclarativeObject {
112                extension: m.extension.clone(),
113                migration: m.migration.clone(),
114                kind: r.kind.label().to_owned(),
115                object: r.name.clone(),
116                how: r.how.to_owned(),
117            };
118            if first.is_some() {
119                warn!(error = %err, "Further declarative-object reference in a migration");
120            } else {
121                first = Some(err);
122            }
123        }
124    }
125    first.map_or(Ok(()), Err)
126}
127
128fn created_objects(extension: &str, sql: &str) -> Result<Objects, LoaderError> {
129    let parsed = pg_query::parse(sql).map_err(|e| LoaderError::SchemaInstallationFailed {
130        extension: extension.to_owned(),
131        message: format!("SQL parse failed: {e}"),
132    })?;
133    Ok(created_objects_of(&parsed))
134}
135
136fn created_objects_of(parsed: &pg_query::ParseResult) -> Objects {
137    let mut objects = Objects::default();
138    for node in top_level(parsed) {
139        match node {
140            NodeEnum::CreateTrigStmt(t) => {
141                objects.triggers.insert(t.trigname.to_lowercase());
142            },
143            NodeEnum::ViewStmt(v) => {
144                if let Some(view) = &v.view {
145                    objects.views.insert(view.relname.to_lowercase());
146                }
147            },
148            _ => {},
149        }
150    }
151    objects
152}
153
154fn references_of(parsed: &pg_query::ParseResult) -> Vec<Reference> {
155    let mut refs = Vec::new();
156    for node in top_level(parsed) {
157        match node {
158            NodeEnum::AlterTableStmt(alter) => {
159                for cmd in &alter.cmds {
160                    let Some(NodeEnum::AlterTableCmd(cmd)) = &cmd.node else {
161                        continue;
162                    };
163                    if matches!(
164                        AlterTableType::try_from(cmd.subtype),
165                        Ok(AlterTableType::AtEnableTrig
166                            | AlterTableType::AtEnableAlwaysTrig
167                            | AlterTableType::AtEnableReplicaTrig
168                            | AlterTableType::AtDisableTrig)
169                    ) {
170                        refs.push(Reference {
171                            kind: ObjectKind::Trigger,
172                            name: cmd.name.to_lowercase(),
173                            how: "ALTER TABLE … ENABLE/DISABLE TRIGGER",
174                        });
175                    }
176                }
177            },
178            NodeEnum::DropStmt(drop) if !drop.missing_ok => {
179                let kind = match ObjectType::try_from(drop.remove_type) {
180                    Ok(ObjectType::ObjectTrigger) => ObjectKind::Trigger,
181                    Ok(ObjectType::ObjectView) => ObjectKind::View,
182                    _ => continue,
183                };
184                for object in &drop.objects {
185                    if let Some(name) = dropped_name(object) {
186                        refs.push(Reference {
187                            kind,
188                            name,
189                            how: "DROP without IF EXISTS",
190                        });
191                    }
192                }
193            },
194            _ => {},
195        }
196    }
197    for (table, context) in &parsed.tables {
198        if matches!(context, Context::Select | Context::DML) {
199            let name = table.rsplit('.').next().unwrap_or(table).to_lowercase();
200            refs.push(Reference {
201                kind: ObjectKind::View,
202                name,
203                how: "a query over the view",
204            });
205        }
206    }
207    refs
208}
209
210// Why: a dropped object is a List of name parts (schema, table, trigger) or
211// a bare RangeVar for a view; the last part is the object's own name.
212fn dropped_name(object: &pg_query::protobuf::Node) -> Option<String> {
213    match object.node.as_ref()? {
214        NodeEnum::List(list) => list.items.iter().rev().find_map(|item| match &item.node {
215            Some(NodeEnum::String(s)) => Some(s.sval.to_lowercase()),
216            _ => None,
217        }),
218        NodeEnum::RangeVar(range) => Some(range.relname.to_lowercase()),
219        NodeEnum::String(s) => Some(s.sval.to_lowercase()),
220        _ => None,
221    }
222}
223
224fn top_level(parsed: &pg_query::ParseResult) -> impl Iterator<Item = &NodeEnum> {
225    parsed
226        .protobuf
227        .stmts
228        .iter()
229        .filter_map(|raw| raw.stmt.as_ref().and_then(|s| s.node.as_ref()))
230}