1use sha2::{Digest, Sha256};
33
34#[derive(Debug, thiserror::Error)]
37pub enum MigrationError {
38 #[error("migration file has no header: {0}")]
39 MissingHeader(String),
40 #[error("migration file has malformed header line: {0:?}")]
41 MalformedHeader(String),
42 #[error("migration ID mismatch in {filename}: header says {header_id}, body hashes to {computed_id}")]
43 IdMismatch {
44 filename: String,
45 header_id: String,
46 computed_id: String,
47 },
48 #[error("migration chain conflict: {0}")]
49 ChainConflict(String),
50 #[error("migration chain has a fork: both {a} and {b} claim onto {parent}")]
51 Fork { parent: String, a: String, b: String },
52 #[error("no migrations found")]
53 Empty,
54 #[error("unknown onto reference {onto} in migration {id}")]
55 UnknownOnto { id: String, onto: String },
56}
57
58#[derive(Debug, Clone)]
62pub struct MigrationFile {
63 pub id: String,
66 pub onto: String,
68 pub squashed: Vec<String>,
70 pub filename: String,
72 pub body: String,
74}
75
76impl MigrationFile {
77 pub fn short_id(&self) -> &str {
81 if self.id.starts_with("m1") {
82 &self.id[..8]
83 } else {
84 &self.id[..14]
85 }
86 }
87
88 pub fn is_first(&self) -> bool {
90 self.onto == "initial"
91 }
92}
93
94const ID_VERSION: &str = "m2";
103
104fn id_digest(body: &str, onto: &str, squashed: &[String]) -> [u8; 32] {
111 let mut hasher = Sha256::new();
112 hasher.update(body.as_bytes());
113 hasher.update([0x1f]);
114 hasher.update(onto.as_bytes());
115 for id in squashed {
116 hasher.update([0x1f]);
117 hasher.update(id.as_bytes());
118 }
119 hasher.finalize().into()
120}
121
122pub fn compute_id(body: &str, onto: &str, squashed: &[String]) -> String {
124 let digest = id_digest(body, onto, squashed);
125 let hex = hex::encode(&digest[..19]); format!("{ID_VERSION}{hex}")
127}
128
129pub fn compute_short_id(body: &str, onto: &str, squashed: &[String]) -> String {
135 let digest = id_digest(body, onto, squashed);
136 let hex = hex::encode(&digest[..6]); format!("{ID_VERSION}{hex}")
138}
139
140fn legacy_compute_id(body: &str) -> String {
143 let digest = Sha256::digest(body.as_bytes());
144 format!("m1{}", hex::encode(&digest[..19]))
145}
146
147pub fn parse(content: &str, filename: &str) -> Result<MigrationFile, MigrationError> {
151 let mut lines = content.splitn(3, '\n');
152
153 let migration_line = lines
154 .next()
155 .ok_or_else(|| MigrationError::MissingHeader(filename.to_string()))?;
156 let onto_line = lines
157 .next()
158 .ok_or_else(|| MigrationError::MissingHeader(filename.to_string()))?;
159 let rest = lines.next().unwrap_or(""); let id = parse_header_line(migration_line, "migration")
162 .ok_or_else(|| MigrationError::MalformedHeader(migration_line.to_string()))?;
163 let onto =
164 parse_header_line(onto_line, "onto").ok_or_else(|| MigrationError::MalformedHeader(onto_line.to_string()))?;
165
166 let body = rest.to_string();
169
170 let squashed = parse_squashed_lines(content);
172
173 Ok(MigrationFile {
174 id: id.to_string(),
175 onto: onto.to_string(),
176 squashed,
177 filename: filename.to_string(),
178 body,
179 })
180}
181
182fn parse_header_line<'a>(line: &'a str, key: &str) -> Option<&'a str> {
184 let prefix = format!("-- {}:", key);
185 let stripped = line.trim().strip_prefix(prefix.as_str())?;
186 Some(stripped.trim())
187}
188
189fn parse_squashed_lines(content: &str) -> Vec<String> {
191 for line in content.lines() {
192 if !line.starts_with("-- ") {
193 break; }
195 if let Some(rest) = line.trim().strip_prefix("-- squashed:") {
196 return rest
197 .split(',')
198 .map(|s| s.trim().to_string())
199 .filter(|s| !s.is_empty())
200 .collect();
201 }
202 }
203 vec![]
204}
205
206pub fn verify_integrity(m: &MigrationFile) -> Result<(), MigrationError> {
215 let computed = if m.id.starts_with("m1") {
216 legacy_compute_id(&m.body)
217 } else {
218 compute_id(&m.body, &m.onto, &m.squashed)
219 };
220 if computed != m.id {
221 return Err(MigrationError::IdMismatch {
222 filename: m.filename.clone(),
223 header_id: m.id.clone(),
224 computed_id: computed,
225 });
226 }
227 Ok(())
228}
229
230pub fn validate_chain(migrations: &[MigrationFile]) -> Result<Vec<&MigrationFile>, MigrationError> {
237 if migrations.is_empty() {
238 return Ok(vec![]);
239 }
240
241 use std::collections::HashMap;
242
243 let by_id: HashMap<&str, &MigrationFile> = migrations.iter().map(|m| (m.id.as_str(), m)).collect();
245
246 let mut children: HashMap<&str, &MigrationFile> = HashMap::new();
248 for m in migrations {
249 if let Some(existing) = children.insert(m.onto.as_str(), m) {
250 return Err(MigrationError::Fork {
251 parent: m.onto.clone(),
252 a: existing.id.clone(),
253 b: m.id.clone(),
254 });
255 }
256 }
257
258 let roots: Vec<_> = migrations.iter().filter(|m| m.onto == "initial").collect();
260 match roots.len() {
261 0 => return Err(MigrationError::ChainConflict("no migration with onto=initial".into())),
262 2.. => {
263 return Err(MigrationError::ChainConflict(format!(
264 "multiple migrations claim onto=initial: {}",
265 roots.iter().map(|m| m.id.as_str()).collect::<Vec<_>>().join(", ")
266 )));
267 }
268 _ => {}
269 }
270
271 let mut chain: Vec<&MigrationFile> = vec![];
273 let mut current = roots[0];
274 loop {
275 if current.onto != "initial" && !by_id.contains_key(current.onto.as_str()) {
277 return Err(MigrationError::UnknownOnto {
278 id: current.id.clone(),
279 onto: current.onto.clone(),
280 });
281 }
282 chain.push(current);
283 match children.get(current.id.as_str()) {
284 None => break, Some(next) => current = next,
286 }
287 }
288
289 if chain.len() != migrations.len() {
291 return Err(MigrationError::ChainConflict(format!(
292 "chain has {} entries but {} files exist — possible cycle or orphan",
293 chain.len(),
294 migrations.len()
295 )));
296 }
297
298 Ok(chain)
299}
300
301pub fn chain_tip<'a>(chain: &[&'a MigrationFile]) -> &'a str {
305 chain.last().map(|m| m.id.as_str()).unwrap_or("initial")
306}
307
308pub fn blank_body() -> &'static str {
313 "\n-- TODO: write this migration's SQL by hand\n"
314}
315
316pub fn render_file(onto: &str, body: &str, squashed: &[String]) -> String {
321 let id = compute_id(body, onto, squashed);
322 let mut out = format!("-- migration: {}\n-- onto: {}\n", id, onto);
323 if !squashed.is_empty() {
324 out.push_str(&format!("-- squashed: {}\n", squashed.join(", ")));
325 }
326 out.push_str(body);
327 out
328}
329
330pub fn parse_steps(body: &str) -> Vec<(bool, String)> {
339 let mut steps = Vec::new();
340 let mut current_transactional = true;
341 let mut current = String::new();
342
343 for line in body.split_inclusive('\n') {
344 match line.trim() {
345 "-- pylon:step" => {
346 steps.push((current_transactional, std::mem::take(&mut current)));
347 current_transactional = true;
348 }
349 "-- pylon:step non-transactional" => {
350 steps.push((current_transactional, std::mem::take(&mut current)));
351 current_transactional = false;
352 }
353 _ => current.push_str(line),
354 }
355 }
356 steps.push((current_transactional, current));
357 steps
358}
359
360pub fn split_statements(sql: &str) -> Vec<String> {
376 let bytes = sql.as_bytes();
377 let mut out = Vec::new();
378 let mut start = 0usize;
379 let mut i = 0usize;
380
381 while i < bytes.len() {
382 match bytes[i] {
383 b'\'' | b'"' => {
384 let quote = bytes[i];
385 i += 1;
386 while i < bytes.len() {
387 if bytes[i] == quote {
388 if bytes.get(i + 1) == Some("e) {
390 i += 2;
391 continue;
392 }
393 break;
394 }
395 i += 1;
396 }
397 i += 1;
398 }
399 b'-' if bytes.get(i + 1) == Some(&b'-') => {
400 while i < bytes.len() && bytes[i] != b'\n' {
401 i += 1;
402 }
403 }
404 b'/' if bytes.get(i + 1) == Some(&b'*') => {
405 let mut depth = 1usize;
406 i += 2;
407 while i < bytes.len() && depth > 0 {
408 if bytes[i] == b'/' && bytes.get(i + 1) == Some(&b'*') {
409 depth += 1;
410 i += 2;
411 } else if bytes[i] == b'*' && bytes.get(i + 1) == Some(&b'/') {
412 depth -= 1;
413 i += 2;
414 } else {
415 i += 1;
416 }
417 }
418 }
419 b'$' => match dollar_tag(bytes, i) {
420 Some(tag) => {
421 i += tag.len();
422 match find_subslice(&bytes[i..], tag) {
425 Some(offset) => i += offset + tag.len(),
426 None => i = bytes.len(),
427 }
428 }
429 None => i += 1,
430 },
431 b';' => {
432 let stmt = sql[start..i].trim();
433 if !stmt.is_empty() {
434 out.push(stmt.to_string());
435 }
436 i += 1;
437 start = i;
438 }
439 _ => i += 1,
440 }
441 }
442
443 let tail = sql[start..].trim();
444 if !tail.is_empty() {
445 out.push(tail.to_string());
446 }
447 out
448}
449
450fn dollar_tag(bytes: &[u8], at: usize) -> Option<&[u8]> {
453 let mut j = at + 1;
454 while j < bytes.len() && (bytes[j].is_ascii_alphanumeric() || bytes[j] == b'_') {
455 if j == at + 1 && bytes[j].is_ascii_digit() {
457 return None;
458 }
459 j += 1;
460 }
461 if bytes.get(j) == Some(&b'$') {
462 Some(&bytes[at..=j])
463 } else {
464 None
465 }
466}
467
468fn find_subslice(haystack: &[u8], needle: &[u8]) -> Option<usize> {
469 if needle.is_empty() || haystack.len() < needle.len() {
470 return None;
471 }
472 (0..=haystack.len() - needle.len()).find(|&i| &haystack[i..i + needle.len()] == needle)
473}
474
475#[cfg(test)]
478mod tests {
479 use super::*;
480
481 fn make_body(tag: &str) -> String {
482 format!("\nCREATE TABLE \"{}\" ();\n", tag)
484 }
485
486 fn make_file(onto: &str, body: &str) -> MigrationFile {
487 let content = render_file(onto, body, &[]);
488 parse(&content, "test").unwrap()
489 }
490
491 #[test]
492 fn test_compute_id_prefix() {
493 let id = compute_id("hello", "initial", &[]);
494 assert!(id.starts_with("m2"));
495 assert_eq!(id.len(), 40); }
497
498 #[test]
499 fn test_compute_short_id() {
500 let sid = compute_short_id("hello", "initial", &[]);
501 assert!(sid.starts_with("m2"));
502 assert_eq!(sid.len(), 14); }
504
505 #[test]
506 fn test_id_is_prefix_of_short_id() {
507 let full = compute_id("SELECT 1;", "initial", &[]);
508 let short = compute_short_id("SELECT 1;", "initial", &[]);
509 assert!(full.starts_with(&short));
510 }
511
512 #[test]
513 fn identical_bodies_at_different_chain_positions_get_different_ids() {
514 let body = "ALTER TABLE t ADD COLUMN c int8;";
517 let first = compute_id(body, "initial", &[]);
518 let second = compute_id(body, "m2aaaaaaaaaaaa", &[]);
519 assert_ne!(first, second);
520 assert_ne!(
521 compute_short_id(body, "initial", &[]),
522 compute_short_id(body, "m2aaaaaaaaaaaa", &[])
523 );
524 }
525
526 #[test]
527 fn the_squash_list_is_part_of_the_id() {
528 let body = "SELECT 1;";
529 assert_ne!(
530 compute_id(body, "initial", &[]),
531 compute_id(body, "initial", &["m1abc".to_string()])
532 );
533 }
534
535 #[test]
536 fn the_field_separator_stops_boundary_ambiguity() {
537 assert_ne!(compute_id("bc", "a", &[]), compute_id("c", "ab", &[]));
539 }
540
541 #[test]
542 fn a_legacy_m1_file_still_verifies_against_the_body_only_hash() {
543 let body = "\nSELECT 1;\n";
544 let legacy_id = legacy_compute_id(body);
545 let content = format!("-- migration: {legacy_id}\n-- onto: initial\n{body}");
546 let m = parse(&content, "00001_legacy").unwrap();
547 assert!(
548 verify_integrity(&m).is_ok(),
549 "an m1 file written before the format change must keep verifying"
550 );
551 assert_eq!(m.short_id().len(), 8, "m1 short IDs stay 6 hex chars wide");
552 }
553
554 #[test]
555 fn test_parse_and_verify() {
556 let body = make_body("Person");
557 let content = render_file("initial", &body, &[]);
558 let m = parse(&content, "00001_m1abc123").unwrap();
559 assert_eq!(m.onto, "initial");
560 assert!(m.id.starts_with("m2"));
561 verify_integrity(&m).unwrap();
562 }
563
564 #[test]
565 fn test_verify_detects_tampering() {
566 let body = make_body("Person");
567 let content = render_file("initial", &body, &[]);
568 let mut m = parse(&content, "test").unwrap();
569 m.body = "TAMPERED;\n".to_string();
570 assert!(verify_integrity(&m).is_err());
571 }
572
573 #[test]
574 fn test_chain_single() {
575 let m1 = make_file("initial", &make_body("A"));
576 let files = [m1];
577 let chain = validate_chain(&files).unwrap();
578 assert_eq!(chain.len(), 1);
579 assert!(chain[0].is_first());
580 }
581
582 #[test]
583 fn test_chain_ordered() {
584 let m1 = make_file("initial", &make_body("A"));
585 let m2 = make_file(&m1.id, &make_body("B"));
586 let m3 = make_file(&m2.id, &make_body("C"));
587 let files = [m3.clone(), m1.clone(), m2.clone()];
589 let chain = validate_chain(&files).unwrap();
590 assert_eq!(chain[0].id, m1.id);
591 assert_eq!(chain[1].id, m2.id);
592 assert_eq!(chain[2].id, m3.id);
593 }
594
595 #[test]
596 fn test_chain_fork_detected() {
597 let m1 = make_file("initial", &make_body("A"));
598 let m2a = make_file(&m1.id, &make_body("B"));
599 let m2b = make_file(&m1.id, &make_body("C"));
600 let files = [m1, m2a, m2b];
601 assert!(validate_chain(&files).is_err());
602 }
603
604 #[test]
605 fn test_empty_chain() {
606 assert!(validate_chain(&[]).unwrap().is_empty());
607 }
608
609 #[test]
610 fn test_squashed_header() {
611 let ids = vec!["m1aaa".to_string(), "m1bbb".to_string()];
612 let body = make_body("Squashed");
613 let content = render_file("initial", &body, &ids);
614 let m = parse(&content, "test").unwrap();
615 assert_eq!(m.squashed, ids);
616 }
617
618 #[test]
619 fn test_blank_body_stable() {
620 assert_eq!(blank_body(), "\n-- TODO: write this migration's SQL by hand\n");
622 }
623
624 #[test]
625 fn test_parse_steps_no_markers_is_one_transactional_step() {
626 let steps = parse_steps("CREATE TABLE foo ();\n");
627 assert_eq!(steps, vec![(true, "CREATE TABLE foo ();\n".to_string())]);
628 }
629
630 #[test]
631 fn test_parse_steps_splits_on_transactional_marker() {
632 let steps = parse_steps("CREATE TABLE a ();\n-- pylon:step\nCREATE TABLE b ();\n");
633 assert_eq!(
634 steps,
635 vec![
636 (true, "CREATE TABLE a ();\n".to_string()),
637 (true, "CREATE TABLE b ();\n".to_string()),
638 ]
639 );
640 }
641
642 #[test]
643 fn test_parse_steps_non_transactional_marker_applies_to_the_next_step() {
644 let steps = parse_steps(
645 "CREATE TABLE a ();\n-- pylon:step non-transactional\nCREATE INDEX CONCURRENTLY idx ON a (x);\n",
646 );
647 assert_eq!(
648 steps,
649 vec![
650 (true, "CREATE TABLE a ();\n".to_string()),
651 (false, "CREATE INDEX CONCURRENTLY idx ON a (x);\n".to_string()),
652 ]
653 );
654 }
655
656 #[test]
657 fn test_parse_steps_reverts_to_transactional_after_a_plain_marker() {
658 let steps = parse_steps(
659 "-- pylon:step non-transactional\nCREATE INDEX CONCURRENTLY idx ON a (x);\n\
660 -- pylon:step\nCREATE TABLE b ();\n",
661 );
662 assert_eq!(steps.len(), 3);
663 assert!(steps[0].0); assert!(!steps[1].0);
665 assert!(steps[2].0);
666 }
667
668 #[test]
669 fn test_parse_steps_empty_body() {
670 assert_eq!(parse_steps(""), vec![(true, String::new())]);
671 }
672
673 #[test]
676 fn split_statements_splits_on_plain_semicolons() {
677 assert_eq!(
678 split_statements("CREATE TABLE a (id int8); CREATE TABLE b (id int8);"),
679 vec!["CREATE TABLE a (id int8)", "CREATE TABLE b (id int8)"]
680 );
681 }
682
683 #[test]
684 fn split_statements_drops_empty_and_trailing_statements() {
685 assert_eq!(split_statements("SELECT 1;;\n\n;"), vec!["SELECT 1"]);
686 assert_eq!(split_statements(""), Vec::<String>::new());
687 assert_eq!(split_statements(" \n "), Vec::<String>::new());
688 }
689
690 #[test]
691 fn split_statements_keeps_a_statement_without_a_trailing_semicolon() {
692 assert_eq!(split_statements("SELECT 1"), vec!["SELECT 1"]);
693 }
694
695 #[test]
696 fn split_statements_ignores_semicolons_in_string_literals() {
697 assert_eq!(
698 split_statements("INSERT INTO t VALUES ('a;b'); SELECT 1;"),
699 vec!["INSERT INTO t VALUES ('a;b')", "SELECT 1"]
700 );
701 }
702
703 #[test]
704 fn split_statements_handles_doubled_quote_escapes() {
705 assert_eq!(
706 split_statements("SELECT 'it''s; fine'; SELECT 2;"),
707 vec!["SELECT 'it''s; fine'", "SELECT 2"]
708 );
709 }
710
711 #[test]
712 fn split_statements_ignores_semicolons_in_quoted_identifiers() {
713 assert_eq!(
714 split_statements("CREATE TABLE \"weird;name\" (id int8); SELECT 1;"),
715 vec!["CREATE TABLE \"weird;name\" (id int8)", "SELECT 1"]
716 );
717 }
718
719 #[test]
720 fn split_statements_ignores_semicolons_in_line_comments() {
721 assert_eq!(
722 split_statements("SELECT 1; -- trailing; comment\nSELECT 2;"),
723 vec!["SELECT 1", "-- trailing; comment\nSELECT 2"]
724 );
725 }
726
727 #[test]
728 fn split_statements_ignores_semicolons_in_block_comments() {
729 assert_eq!(
730 split_statements("SELECT 1 /* a; b */; SELECT 2;"),
731 vec!["SELECT 1 /* a; b */", "SELECT 2"]
732 );
733 }
734
735 #[test]
736 fn split_statements_handles_nested_block_comments() {
737 assert_eq!(
738 split_statements("SELECT 1 /* a /* b; */ c; */; SELECT 2;"),
739 vec!["SELECT 1 /* a /* b; */ c; */", "SELECT 2"]
740 );
741 }
742
743 #[test]
744 fn split_statements_keeps_a_dollar_quoted_body_intact() {
745 let sql = "CREATE FUNCTION f() RETURNS trigger LANGUAGE plpgsql AS $$\n\
748 BEGIN\n PERFORM 1;\n RETURN NEW;\nEND;\n$$;\nSELECT 1;";
749 let stmts = split_statements(sql);
750 assert_eq!(stmts.len(), 2, "got {stmts:#?}");
751 assert!(stmts[0].starts_with("CREATE FUNCTION f()"));
752 assert!(stmts[0].ends_with("$$"));
753 assert_eq!(stmts[1], "SELECT 1");
754 }
755
756 #[test]
757 fn split_statements_keeps_a_tagged_dollar_quoted_body_intact() {
758 let sql = "CREATE FUNCTION f() RETURNS int8 AS $body$ SELECT 1; $body$ LANGUAGE sql; SELECT 2;";
759 let stmts = split_statements(sql);
760 assert_eq!(stmts.len(), 2, "got {stmts:#?}");
761 assert!(stmts[0].contains("$body$ SELECT 1; $body$"));
762 assert_eq!(stmts[1], "SELECT 2");
763 }
764
765 #[test]
766 fn split_statements_handles_the_do_block_enum_form() {
767 let sql = "DO $$ BEGIN CREATE TYPE s.t AS ENUM ('a'); EXCEPTION WHEN duplicate_object THEN NULL; END $$;";
769 assert_eq!(split_statements(sql).len(), 1);
770 }
771
772 #[test]
773 fn split_statements_treats_dollar_digit_as_a_placeholder_not_a_tag() {
774 assert_eq!(
775 split_statements("SELECT $1; SELECT $2;"),
776 vec!["SELECT $1", "SELECT $2"]
777 );
778 }
779
780 #[test]
781 fn split_statements_does_not_hang_on_an_unterminated_dollar_body() {
782 let stmts = split_statements("CREATE FUNCTION f() AS $$ SELECT 1;");
783 assert_eq!(stmts.len(), 1);
784 }
785
786 #[test]
787 fn split_statements_does_not_hang_on_an_unterminated_string() {
788 let stmts = split_statements("SELECT 'oops;");
789 assert_eq!(stmts.len(), 1);
790 }
791}