1#![expect(
3 clippy::pub_with_shorthand,
4 reason = "Rustfmt normalizes crate-private helper visibility to pub(crate)."
5)]
6
7use alloc::collections::BTreeMap;
8
9use serde::{Deserialize, Serialize};
10
11pub const SNAPSHOT_VERSION: u32 = 3;
13
14#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
16#[serde(deny_unknown_fields)]
17pub struct Snapshot {
18 pub version: u32,
19 pub schemas: Vec<Schema>,
20}
21
22impl Snapshot {
23 #[expect(
28 clippy::impl_trait_in_params,
29 reason = "Path arguments accept standard owned and borrowed path types."
30 )]
31 pub fn write_to(&self, path: impl AsRef<std::path::Path>) -> Result<(), crate::Error> {
32 let mut bytes = serde_json::to_vec_pretty(self)?;
33 bytes.push(b'\n');
34 crate::write_if_changed(path.as_ref(), &bytes)
35 }
36}
37
38#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
39#[serde(deny_unknown_fields)]
40pub struct Schema {
41 pub name: String,
42 pub enums: Vec<Enum>,
43 pub composites: Vec<Composite>,
44 pub tables: Vec<Table>,
45 pub functions: Vec<Function>,
46}
47
48#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
49#[serde(deny_unknown_fields)]
50pub struct Enum {
51 pub name: String,
52 pub variants: Vec<String>,
53}
54
55#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
56#[serde(deny_unknown_fields)]
57pub struct Composite {
58 pub name: String,
59 pub fields: Vec<Column>,
60}
61
62#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
63#[serde(deny_unknown_fields)]
64pub struct Table {
65 pub name: String,
66 pub kind: TableKind,
67 pub columns: Vec<Column>,
68 #[serde(deserialize_with = "Deserialize::deserialize")]
69 pub primary_key: Option<Vec<String>>,
70 pub unique_keys: Vec<Vec<String>>,
71 pub foreign_keys: Vec<ForeignKey>,
72 pub is_partition: bool,
73}
74
75#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
77#[serde(deny_unknown_fields)]
78pub struct RelationRef {
79 pub schema: String,
80 pub name: String,
81}
82
83#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
85#[serde(deny_unknown_fields)]
86pub struct ForeignKey {
87 pub name: String,
88 pub columns: Vec<String>,
89 pub referenced_relation: RelationRef,
90 pub referenced_columns: Vec<String>,
91}
92
93#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
94#[serde(rename_all = "snake_case")]
95pub enum TableKind {
96 Table,
97 View,
98 MaterializedView,
99}
100
101#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
102#[serde(deny_unknown_fields)]
103pub struct Column {
104 pub name: String,
105 pub ty: PgType,
106 pub nullable: bool,
107 pub has_default: bool,
108 pub generated: bool,
109 pub identity: Identity,
110}
111
112#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
113#[serde(rename_all = "snake_case")]
114pub enum Identity {
115 None,
116 Always,
117 ByDefault,
118}
119
120#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
121#[serde(
122 tag = "kind",
123 content = "value",
124 rename_all = "snake_case",
125 deny_unknown_fields
126)]
127pub enum PgType {
128 Builtin(String),
129 Named {
130 schema: String,
131 name: String,
132 },
133 Array(Box<Self>),
134 Domain {
135 schema: String,
136 name: String,
137 base: Box<Self>,
138 },
139}
140
141#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
142#[serde(deny_unknown_fields)]
143pub struct Function {
144 pub name: String,
145 pub arguments: Vec<Argument>,
146 pub returns: ReturnType,
147 pub returns_set: bool,
148}
149
150#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
151#[serde(deny_unknown_fields)]
152pub struct Argument {
153 pub name: String,
154 pub ty: PgType,
155 pub has_default: bool,
156 pub nullable: bool,
158}
159
160#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
161#[serde(
162 tag = "kind",
163 content = "value",
164 rename_all = "snake_case",
165 deny_unknown_fields
166)]
167pub enum ReturnType {
168 Type(PgType),
169 Record(Vec<Column>),
170}
171
172pub(crate) fn set_not_null(
173 columns: &mut [Column],
174 fields: &[String],
175 target: &str,
176) -> Result<(), crate::Error> {
177 for field in fields {
178 let column = columns
179 .iter_mut()
180 .find(|column| &column.name == field)
181 .ok_or_else(|| {
182 crate::Error::Invalid(format!("unknown @not_null field {target}.{field}"))
183 })?;
184 column.nullable = false;
185 }
186 Ok(())
187}
188
189#[cfg(any(feature = "database", test))]
190pub(crate) fn annotation(comment: Option<&str>, tag: &str) -> Result<Vec<String>, crate::Error> {
191 let mut fields = Vec::new();
192 for line in comment.unwrap_or_default().lines() {
193 let line = line.trim();
194 if let Some(rest) = line.strip_prefix(tag) {
195 if !rest.starts_with(char::is_whitespace) {
196 return Err(crate::Error::Invalid(format!("invalid {tag} annotation")));
197 }
198 let names: Vec<_> = rest
199 .split([',', ' ', '\t'])
200 .filter(|name| !name.is_empty())
201 .map(str::to_owned)
202 .collect();
203 if names.is_empty() {
204 return Err(crate::Error::Invalid(format!("empty {tag} annotation")));
205 }
206 fields.extend(names);
207 }
208 }
209 Ok(fields)
210}
211
212pub(crate) fn apply_not_null(
213 snapshot: &mut Snapshot,
214 targets: &BTreeMap<String, Vec<String>>,
215) -> Result<(), crate::Error> {
216 for (target, fields) in targets {
217 let mut found = false;
218 for schema in &mut snapshot.schemas {
219 let schema_name = crate::emitter::ident(&schema.name, false)?.to_string();
220 for table in &mut schema.tables {
221 if target
222 == &format!(
223 "{}.tables.{}.Row",
224 schema_name,
225 crate::emitter::ident(&table.name, false)?
226 )
227 && table.kind != TableKind::Table
228 {
229 set_not_null(&mut table.columns, fields, target)?;
230 found = true;
231 }
232 }
233 for composite in &mut schema.composites {
234 if target
235 == &format!(
236 "{}.composites.{}",
237 schema_name,
238 crate::emitter::ident(&composite.name, true)?
239 )
240 {
241 set_not_null(&mut composite.fields, fields, target)?;
242 found = true;
243 }
244 }
245 let mut groups: BTreeMap<String, Vec<&mut Function>> = BTreeMap::new();
246 for function in &mut schema.functions {
247 groups
248 .entry(function.name.clone())
249 .or_default()
250 .push(function);
251 }
252 for (name, mut functions) in groups {
253 functions.sort_by_key(|function| {
254 function
255 .arguments
256 .iter()
257 .map(|argument| format!("{}:{:?}", argument.name, argument.ty))
258 .collect::<Vec<_>>()
259 });
260 let count = functions.len();
261 let base = crate::emitter::ident(&name, false)?.to_string();
262 for (index, function) in functions.into_iter().enumerate() {
263 let name = if count == 1 {
264 base.clone()
265 } else {
266 format!("{}_{index}", base.strip_prefix("r#").unwrap_or(&base))
267 };
268 if target == &format!("{schema_name}.functions.{name}.Record") {
269 if let ReturnType::Record(columns) = &mut function.returns {
270 set_not_null(columns, fields, target)?;
271 found = true;
272 }
273 }
274 }
275 }
276 }
277 if !found {
278 return Err(crate::Error::Invalid(format!(
279 "unknown not_null target {target:?}"
280 )));
281 }
282 }
283 Ok(())
284}
285
286#[cfg(any(feature = "database", test))]
287#[expect(
290 clippy::wildcard_enum_match_arm,
291 reason = "SQL expressions are a conservative whitelist; every other third-party AST form provides no inference."
292)]
293pub(crate) fn view_columns(sql: &str) -> Option<(Vec<String>, Vec<Option<String>>)> {
294 use sqlparser::{
295 ast::{Expr, GroupByExpr, SelectItem, SetExpr, Statement, TableFactor},
296 dialect::PostgreSqlDialect,
297 parser::Parser,
298 };
299 let statements = Parser::parse_sql(&PostgreSqlDialect {}, sql).ok()?;
300 let [Statement::Query(query)] = statements.as_slice() else {
301 return None;
302 };
303 if query.with.is_some() {
304 return None;
305 }
306 let SetExpr::Select(select) = query.body.as_ref() else {
307 return None;
308 };
309 if !matches!(&select.group_by, GroupByExpr::Expressions(expressions, modifiers) if expressions.is_empty() && modifiers.is_empty())
310 {
311 return None;
312 }
313 let [from] = select.from.as_slice() else {
314 return None;
315 };
316 if !from.joins.is_empty() {
317 return None;
318 }
319 let TableFactor::Table {
320 name,
321 alias,
322 args: None,
323 ..
324 } = &from.relation
325 else {
326 return None;
327 };
328 if alias
329 .as_ref()
330 .is_some_and(|alias| !alias.columns.is_empty())
331 {
332 return None;
333 }
334 let relation: Vec<String> = name
335 .0
336 .iter()
337 .map(|part| part.as_ident().map(|name| name.value.clone()))
338 .collect::<Option<_>>()?;
339 let qualifier = alias
340 .as_ref()
341 .map(|alias| alias.name.value.as_str())
342 .or_else(|| relation.last().map(String::as_str))?;
343 let columns = select
344 .projection
345 .iter()
346 .map(|item| {
347 let (SelectItem::UnnamedExpr(expression)
348 | SelectItem::ExprWithAlias {
349 expr: expression, ..
350 }) = item
351 else {
352 return None;
353 };
354 match expression {
355 Expr::Identifier(name) => Some(name.value.clone()),
356 Expr::CompoundIdentifier(names)
357 if names.len() == 2
358 && names.first().is_some_and(|name| name.value == qualifier) =>
359 {
360 names.last().map(|name| name.value.clone())
361 }
362 _ => None,
363 }
364 })
365 .collect();
366 Some((relation, columns))
367}
368
369#[cfg(any(feature = "database", test))]
370#[expect(
373 clippy::wildcard_enum_match_arm,
374 reason = "Only exact literal membership forms qualify; other third-party AST expressions and values are rejected."
375)]
376pub(crate) fn check_values(sql: &str, column: &str) -> Option<Vec<String>> {
377 use sqlparser::{
378 ast::{BinaryOperator, Expr, SelectItem, SetExpr, Statement, Value},
379 dialect::PostgreSqlDialect,
380 parser::Parser,
381 };
382 fn unnest(expression: &Expr) -> Option<&Expr> {
383 match expression {
384 Expr::Nested(inner) => unnest(inner),
385 Expr::Cast {
386 expr, data_type, ..
387 } if matches!(
388 data_type.to_string().as_str(),
389 "TEXT"
390 | "VARCHAR"
391 | "TEXT[]"
392 | "VARCHAR[]"
393 | "pg_catalog.text"
394 | "pg_catalog.text[]"
395 ) =>
396 {
397 unnest(expr)
398 }
399 Expr::Cast { .. } => None,
400 other => Some(other),
401 }
402 }
403 let statements = Parser::parse_sql(&PostgreSqlDialect {}, &format!("SELECT {sql}")).ok()?;
404 let [Statement::Query(query)] = statements.as_slice() else {
405 return None;
406 };
407 let SetExpr::Select(select) = query.body.as_ref() else {
408 return None;
409 };
410 let [SelectItem::UnnamedExpr(expression)] = select.projection.as_slice() else {
411 return None;
412 };
413 let (left, values) = match unnest(expression)? {
414 Expr::InList {
415 expr,
416 list,
417 negated: false,
418 } => (expr.as_ref(), list.as_slice()),
419 Expr::AnyOp {
420 left,
421 compare_op: BinaryOperator::Eq,
422 right,
423 ..
424 } => {
425 let Expr::Array(array) = unnest(right)? else {
426 return None;
427 };
428 (left.as_ref(), array.elem.as_slice())
429 }
430 _ => return None,
431 };
432 if !matches!(unnest(left)?, Expr::Identifier(name) if name.value == column) || values.is_empty()
433 {
434 return None;
435 }
436 let mut labels = values
437 .iter()
438 .map(|value| match unnest(value)? {
439 Expr::Value(value) => match &value.value {
440 Value::SingleQuotedString(label) | Value::EscapedStringLiteral(label) => {
443 Some(label.clone())
444 }
445 _ => None,
446 },
447 _ => None,
448 })
449 .collect::<Option<Vec<_>>>()?;
450 labels.sort();
451 labels.dedup();
452 Some(labels)
453}
454#[cfg(test)]
455#[expect(
456 clippy::unwrap_used,
457 reason = "Contract tests retain failure diagnostics."
458)]
459mod contract_tests {
460 use super::*;
461
462 #[test]
463 fn view_inference_rejects_nullable_sql_shapes() {
464 assert_eq!(
465 view_columns("SELECT t.id, t.id + 1 AS computed FROM public.t t"),
466 Some((
467 vec!["public".into(), "t".into()],
468 vec![Some("id".into()), None]
469 ))
470 );
471 for sql in [
472 "SELECT a.id FROM a LEFT JOIN b ON a.id = b.id",
473 "SELECT id FROM a UNION SELECT id FROM b",
474 "SELECT id FROM a GROUP BY ROLLUP(id)",
475 "WITH a AS (SELECT NULL AS id) SELECT id FROM a",
476 ] {
477 assert!(view_columns(sql).is_none());
478 }
479 }
480
481 #[test]
482 fn annotations_reject_empty_contracts() {
483 assert_eq!(
484 annotation(Some("@not_null id, name"), "@not_null").unwrap(),
485 ["id", "name"]
486 );
487 annotation(Some("@not_null"), "@not_null").unwrap_err();
488 }
489
490 #[test]
491 fn check_membership_accepts_only_literal_exact_forms() {
492 let labels = Some(vec!["organization".into(), "user".into()]);
493 assert_eq!(
494 check_values("owner_type IN ('user', 'organization')", "owner_type"),
495 labels
496 );
497 assert_eq!(
498 check_values(
499 "((owner_type = ANY (ARRAY['user'::text, 'organization'::text])))",
500 "owner_type"
501 ),
502 labels
503 );
504 for expression in [
505 "owner_type NOT IN ('user')",
506 "owner_type IN ('user', other)",
507 "owner_type IN ('user') OR owner_type IS NULL",
508 "lower(owner_type) IN ('user')",
509 "owner_type::uuid IN ('user')",
510 ] {
511 assert!(check_values(expression, "owner_type").is_none());
512 }
513 }
514
515 #[test]
516 fn check_membership_decodes_postgres_escape_strings_once() {
517 assert_eq!(
518 check_values(
519 r"kind IN (E'a\\b', E'line\nbreak', E'\u0061', 'a\b')",
520 "kind"
521 ),
522 Some(vec!["a".into(), "a\\b".into(), "line\nbreak".into()])
523 );
524 assert!(check_values(r"kind IN (E'\xc3\xa9')", "kind").is_none());
527 assert!(check_values("kind::char IN ('a')", "kind").is_none());
528 }
529}