systemprompt_database/lifecycle/installation/
migration_refs.rs1use 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
210fn 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}