systemprompt_database/lifecycle/installation/
migration_cost.rs1use std::sync::Arc;
44
45use pg_query::NodeEnum;
46use pg_query::protobuf::AlterTableType;
47use systemprompt_extension::Extension;
48use systemprompt_extension::cost::{self, CostDirective};
49use systemprompt_identifiers::ExtensionId;
50
51pub const HOT_TABLES: &[&str] = &[
52 "ai_requests",
53 "ai_request_messages",
54 "ai_request_payloads",
55 "ai_request_client_evidence",
56 "ai_request_tool_calls",
57 "analytics_events",
58 "event_outbox",
59 "logs",
60 "user_sessions",
61];
62
63#[derive(Debug, Clone, PartialEq, Eq)]
66pub struct ExpensiveStatement {
67 pub position: usize,
68 pub table: String,
69 pub form: &'static str,
70}
71
72#[derive(Debug, Clone)]
75pub struct MigrationCost {
76 pub extension: ExtensionId,
77 pub migration: String,
78 pub statements: Vec<ExpensiveStatement>,
79 pub declared: Option<CostDirective>,
80 pub malformed: Option<String>,
81}
82
83impl MigrationCost {
84 #[must_use]
85 pub const fn is_undeclared(&self) -> bool {
86 !self.statements.is_empty() && self.declared.is_none()
87 }
88
89 #[must_use]
90 pub fn label(&self) -> String {
91 format!("{}/{}", self.extension, self.migration)
92 }
93
94 #[must_use]
95 pub fn statement_summary(&self) -> String {
96 self.statements
97 .iter()
98 .map(|s| format!("statement {} {} {}", s.position, s.form, s.table))
99 .collect::<Vec<_>>()
100 .join("; ")
101 }
102}
103
104#[must_use]
105pub fn audit_migration_cost(extensions: &[Arc<dyn Extension>], hot: &[&str]) -> Vec<MigrationCost> {
106 let mut out = Vec::new();
107 for ext in extensions {
108 let extension = ExtensionId::new(ext.id());
109 for migration in ext.migrations().into_iter().filter(|m| !m.tombstone) {
110 let label = format!("{:03}_{}", migration.version, migration.name);
111 if let Some(cost) = audit_one(&extension, &label, migration.sql, hot) {
112 out.push(cost);
113 }
114 }
115 }
116 out
117}
118
119#[must_use]
120pub fn audit_one(
121 extension: &ExtensionId,
122 migration: &str,
123 sql: &str,
124 hot: &[&str],
125) -> Option<MigrationCost> {
126 let (declared, malformed) = match cost::parse(sql) {
127 Ok(found) => (found, None),
128 Err(e) => (None, Some(e.to_string())),
129 };
130 let statements = expensive_statements(sql, hot);
131 if statements.is_empty() && declared.is_none() && malformed.is_none() {
132 return None;
133 }
134 Some(MigrationCost {
135 extension: extension.clone(),
136 migration: migration.to_owned(),
137 statements,
138 declared,
139 malformed,
140 })
141}
142
143fn expensive_statements(sql: &str, hot: &[&str]) -> Vec<ExpensiveStatement> {
147 let Ok(parsed) = pg_query::parse(sql) else {
148 return Vec::new();
149 };
150 let mut out = Vec::new();
151 for (index, node) in parsed
152 .protobuf
153 .stmts
154 .iter()
155 .filter_map(|raw| raw.stmt.as_ref().and_then(|s| s.node.as_ref()))
156 .enumerate()
157 {
158 let position = index + 1;
159 if let Some((table, form)) = classify(node)
160 && hot.contains(&table.as_str())
161 {
162 out.push(ExpensiveStatement {
163 position,
164 table,
165 form,
166 });
167 }
168 }
169 out
170}
171
172fn is_select_driven(select: Option<&pg_query::protobuf::Node>) -> bool {
173 let Some(NodeEnum::SelectStmt(select)) = select.and_then(|n| n.node.as_ref()) else {
174 return false;
175 };
176 select.values_lists.is_empty()
177}
178
179fn classify(node: &NodeEnum) -> Option<(String, &'static str)> {
180 match node {
181 NodeEnum::UpdateStmt(stmt) => Some((stmt.relation.as_ref()?.relname.clone(), "UPDATE on")),
182 NodeEnum::DeleteStmt(stmt) => {
183 Some((stmt.relation.as_ref()?.relname.clone(), "DELETE from"))
184 },
185 NodeEnum::InsertStmt(stmt) if is_select_driven(stmt.select_stmt.as_deref()) => Some((
191 stmt.relation.as_ref()?.relname.clone(),
192 "INSERT … SELECT into",
193 )),
194 NodeEnum::IndexStmt(stmt) if !stmt.concurrent => Some((
197 stmt.relation.as_ref()?.relname.clone(),
198 "CREATE INDEX (not CONCURRENTLY) on",
199 )),
200 NodeEnum::AlterTableStmt(stmt) => {
201 let table = stmt.relation.as_ref()?.relname.clone();
202 let form = stmt.cmds.iter().find_map(|cmd| match cmd.node.as_ref() {
203 Some(NodeEnum::AlterTableCmd(c)) => scanning_alter(c.subtype),
204 _ => None,
205 })?;
206 Some((table, form))
207 },
208 _ => None,
209 }
210}
211
212const fn scanning_alter(subtype: i32) -> Option<&'static str> {
215 if subtype == AlterTableType::AtValidateConstraint as i32 {
216 return Some("ALTER TABLE … VALIDATE CONSTRAINT on");
217 }
218 if subtype == AlterTableType::AtSetNotNull as i32 {
219 return Some("ALTER TABLE … SET NOT NULL on");
220 }
221 None
222}