Skip to main content

pylon_core/migration/
mod.rs

1//
2// This source file is part of the Pylon open source project.
3//
4// Copyright (c) 2026 Jaldis B.V.
5//
6// Licensed under the MIT OR Apache-2.0 license (the "License");
7// you may not use this file except in compliance with the License.
8// You may obtain a copy of the License at
9//
10//     https://opensource.org/licenses/MIT
11//     https://www.apache.org/licenses/LICENSE-2.0
12//
13// Unless required by applicable law or agreed to in writing, software
14// distributed under the License is distributed on an "AS IS" BASIS,
15// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
16// See the License for the specific language governing permissions and
17// limitations under the License.
18//
19
20//! Migration file parsing, ID computation, and chain validation.
21//!
22//! Every migration file has a two-line header:
23//!   -- migration: m2<38-hex>
24//!   -- onto: m2<38-hex> | initial
25//!
26//! The migration ID is `m2` + 19 bytes (38 hex chars) of a SHA-256 over the
27//! body *and* the migration's position in the chain (`onto`, plus any
28//! squashed IDs). The short ID used in filenames is the same digest cut to 6
29//! bytes (12 hex chars). `m1` — a body-only hash with a 3-byte short form —
30//! is still accepted for files already on disk; see `verify_integrity`.
31
32use sha2::{Digest, Sha256};
33
34// ── Error ─────────────────────────────────────────────────────────────────────
35
36#[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// ── Data types ────────────────────────────────────────────────────────────────
59
60/// A parsed migration file.
61#[derive(Debug, Clone)]
62pub struct MigrationFile {
63    /// Full migration ID: a version prefix (`m2`, or `m1` for a file
64    /// written before the format change) + 38 hex chars.
65    pub id: String,
66    /// Parent's full ID, or the literal string `"initial"`.
67    pub onto: String,
68    /// IDs this migration squashes (§12); empty for normal migrations.
69    pub squashed: Vec<String>,
70    /// Original filename (stem only, e.g. `00001_m2a3f9bc12de4`).
71    pub filename: String,
72    /// Everything after the two header lines.
73    pub body: String,
74}
75
76impl MigrationFile {
77    /// Short ID used in filenames — the leading prefix of `id`, whose length
78    /// depends on which format wrote the file: `m1` carried 6 hex chars,
79    /// `m2` carries 12 (see `compute_short_id`).
80    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    /// Whether this is the "first" migration (onto == "initial").
89    pub fn is_first(&self) -> bool {
90        self.onto == "initial"
91    }
92}
93
94// ── ID computation ────────────────────────────────────────────────────────────
95
96/// Current ID format. `m1` hashed the body alone, which meant two migrations
97/// with identical bodies at different points in the chain collided on one ID
98/// — and `_pylon."Migrations"` keys on that ID, so the second apply
99/// overwrote the first's tracking row. `m2` folds `onto` and the squash list
100/// in, making the ID identify a migration's position in the chain and not
101/// just its text.
102const ID_VERSION: &str = "m2";
103
104/// The digest every ID is derived from: the body, plus the chain position
105/// that body sits at.
106///
107/// Field-separated with a byte that can't occur in any of the inputs, so
108/// `(onto = "a", body = "bc")` and `(onto = "ab", body = "c")` can't hash
109/// alike.
110fn 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
122/// Compute the full migration ID (§5): `m2` + 38 hex chars of the digest.
123pub 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]); // 19 bytes = 38 hex chars
126    format!("{ID_VERSION}{hex}")
127}
128
129/// Compute the short migration ID used as the filename component.
130///
131/// 12 hex chars (48 bits), not the 6 (24 bits) the `m1` format used: at 24
132/// bits two migrations share a filename stem with ~50% probability by the
133/// ~4,800th migration, and a few percent within the first few hundred.
134pub 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]); // 6 bytes = 12 hex chars
137    format!("{ID_VERSION}{hex}")
138}
139
140/// The `m1` full ID for `body` — body-only hash, kept solely so migration
141/// files written before the `m2` format still verify.
142fn legacy_compute_id(body: &str) -> String {
143    let digest = Sha256::digest(body.as_bytes());
144    format!("m1{}", hex::encode(&digest[..19]))
145}
146
147// ── Parsing ───────────────────────────────────────────────────────────────────
148
149/// Parse a migration file's content. `filename` is used for error messages.
150pub 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(""); // body (may be empty)
160
161    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    // The body starts after the second newline.
167    // Strip exactly one leading newline that separates headers from body.
168    let body = rest.to_string();
169
170    // Parse optional squashed: lines (§12)
171    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
182/// Parse `-- key: value` and return `value`.
183fn 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
189/// Extract `-- squashed: id1, id2, ...` from the header block.
190fn parse_squashed_lines(content: &str) -> Vec<String> {
191    for line in content.lines() {
192        if !line.starts_with("-- ") {
193            break; // end of header block
194        }
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
206// ── Integrity verification ─────────────────────────────────────────────────────
207
208/// Re-hash a migration and verify it matches the ID in the header (§9.2).
209///
210/// Which hash to check against comes from the ID's own version prefix, so a
211/// project with `m1` files on disk keeps verifying against the body-only
212/// hash those files were written with. Only newly created migrations get
213/// `m2` — there is no rewrite step and no flag day.
214pub 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
230// ── Chain validation ──────────────────────────────────────────────────────────
231
232/// Validate that `migrations` form a single unbroken linear chain (§6).
233///
234/// Returns the migrations in chain order, from oldest (onto=initial) to newest (tip).
235/// Errors on forks, cycles, broken references, or an empty input.
236pub 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    // Build a map from ID → migration.
244    let by_id: HashMap<&str, &MigrationFile> = migrations.iter().map(|m| (m.id.as_str(), m)).collect();
245
246    // Build a map from onto → child. Detect forks.
247    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    // Find the root (onto == "initial").
259    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    // Walk the chain from root to tip.
272    let mut chain: Vec<&MigrationFile> = vec![];
273    let mut current = roots[0];
274    loop {
275        // Verify onto reference exists (unless initial).
276        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, // reached the tip
285            Some(next) => current = next,
286        }
287    }
288
289    // Sanity: chain length should match input length (no orphans/cycles).
290    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
301// ── Tip resolution ────────────────────────────────────────────────────────────
302
303/// Return the tip migration ID from an ordered chain, or `"initial"` if empty.
304pub fn chain_tip<'a>(chain: &[&'a MigrationFile]) -> &'a str {
305    chain.last().map(|m| m.id.as_str()).unwrap_or("initial")
306}
307
308// ── Blank migration body ──────────────────────────────────────────────────────
309
310/// Generate the stub body for a blank migration (§8.5).
311/// Includes the leading blank line that separates the header from SQL content.
312pub fn blank_body() -> &'static str {
313    "\n-- TODO: write this migration's SQL by hand\n"
314}
315
316/// Build the full file content for a migration given `onto` and `body`.
317///
318/// `body` must include the leading blank line (the separator between header and SQL),
319/// so that `parse` recovers the same bytes that `compute_id` hashed.
320pub 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
330// ── Step splitting ────────────────────────────────────────────────────────────
331
332/// Split a migration body on `-- pylon:step` markers into `(transactional,
333/// sql)` pairs, in order. A step is transactional unless its *preceding*
334/// marker was `-- pylon:step non-transactional` (that marker applies to the
335/// step it introduces, not the one it ends — matching the Python
336/// implementation this replaces exactly, including the leading segment
337/// before any marker always being transactional).
338pub 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
360// ── Statement splitting ───────────────────────────────────────────────────────
361
362/// Split one step's SQL into individual statements on top-level `;`.
363///
364/// Needed because `apply --dev-mode` has to be able to skip a *single*
365/// already-applied statement rather than discarding the whole step (see
366/// `migrate::apply_one`). A naive `split(';')` would corrupt every step that
367/// contains a dollar-quoted function body — which is most of them, since
368/// trigger and constraint DDL is emitted as `... AS $$ ... ; ... $$`.
369///
370/// Recognises the four places a `;` can appear without ending a statement:
371/// single-quoted strings (`''` escapes), quoted identifiers (`""` escapes),
372/// dollar-quoted bodies (`$$` or `$tag$`), and comments (`--` to end of line,
373/// `/* */` which nest in PostgreSQL). Empty statements are dropped, so a
374/// trailing `;` or a stray blank line never produces one.
375pub 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                        // A doubled quote is an escaped quote, not the end.
389                        if bytes.get(i + 1) == Some(&quote) {
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                    // Scan to the matching closing tag; an unterminated body
423                    // runs to end of input rather than looping forever.
424                    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
450/// If a dollar-quote tag opens at `at`, return it (including both `$`s).
451/// `$$` and `$tag$` open one; `$1` (a parameter) and a bare `$` do not.
452fn 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        // A tag can't start with a digit — that's a `$1` placeholder.
456        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// ── Tests ─────────────────────────────────────────────────────────────────────
476
477#[cfg(test)]
478mod tests {
479    use super::*;
480
481    fn make_body(tag: &str) -> String {
482        // Body includes the leading blank line (header-to-SQL separator).
483        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); // "m2" + 38 hex
496    }
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); // "m2" + 12 hex
503    }
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        // The `m1` collision: same DDL added, reverted, then re-added later
515        // hashed to one ID, and `_pylon."Migrations"` keys on that ID.
516        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        // Without a separator these two would hash the same bytes.
538        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        // Pass in reverse order — validate_chain should still sort correctly.
588        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        // Blank body must be stable so its hash is consistent.
621        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); // leading empty segment, transactional
664        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    // ── split_statements ──────────────────────────────────────────────────
674
675    #[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        // This is the case a naive split(';') corrupts: the function body has
746        // two internal semicolons that must not end the CREATE FUNCTION.
747        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        // Exactly what `export::emit_enum` produces.
768        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}