Skip to main content

yggdrasil/db/
migrate.rs

1//! 数据库迁移运行器。
2//!
3//! 在服务器启动时(`dioxus::server::serve()` 之前)自动执行迁移。
4//! 设计要点:
5//! - 迁移 SQL 通过 `include_str!` 内联进二进制,部署只需单个二进制。
6//! - `schema_migrations` 表记录已应用版本,避免重复执行。
7//! - 每个迁移在独立事务里执行,失败自动回滚,版本行不会写入。
8//! - 咨询锁(`pg_advisory_lock`)保证多实例启动时只有一个进程执行迁移。
9//!
10//! 前置条件:目标数据库的存在性由 [`crate::db::pool::ensure_database`] 在本模块
11//! 被调用前保证(它先连 `postgres` 维护库做 `CREATE DATABASE IF NOT EXISTS` 等价逻辑)。
12//! 本模块负责 schema 与一次性存储格式迁移,不再关心"库不存在"。
13//!
14//! 仅在 `feature = "server"` 时编译。
15
16use std::collections::HashSet;
17
18/// 咨询锁的固定 key。Postgres 咨询锁是数据库级唯一的;
19/// 这里用一个项目专属的大整数,避免与同库其它应用冲突。
20/// 该值无语义,仅要求唯一性。
21const ADVISORY_LOCK_KEY: i64 = 0x5947_4752_4153_494C;
22
23/// 所有迁移的 (version, sql) 列表,按 version 升序排列。
24///
25/// 新增迁移时:
26/// 1. 在 `migrations/` 下创建 `NNN_描述.sql`
27/// 2. 在本数组末尾追加一行 `("NNN", include_str!("../../migrations/NNN_描述.sql"))`
28///
29/// version 用字符串而非整数,与文件名前缀直接对应,且未来可支持日期戳/语义版本。
30const MIGRATIONS: &[(&str, &str)] = &[
31    ("001", include_str!("../../migrations/001_init.sql")),
32    ("002", include_str!("../../migrations/002_posts.sql")),
33    ("003", include_str!("../../migrations/003_indexes.sql")),
34    ("004", include_str!("../../migrations/004_search_trgm.sql")),
35    ("005", include_str!("../../migrations/005_comments.sql")),
36    ("006", include_str!("../../migrations/006_add_toc_html.sql")),
37    ("007", include_str!("../../migrations/007_settings.sql")),
38    (
39        "008",
40        include_str!("../../migrations/008_comments_cascade.sql"),
41    ),
42    (
43        "009",
44        include_str!("../../migrations/009_cleanup_duplicate_indexes.sql"),
45    ),
46    (
47        "010",
48        include_str!("../../migrations/010_post_word_counts.sql"),
49    ),
50    ("011", include_str!("../../migrations/011_perf_indexes.sql")),
51    (
52        "012",
53        include_str!("../../migrations/012_session_generation.sql"),
54    ),
55    (
56        "013",
57        include_str!("../../migrations/013_comment_content_hash_index.sql"),
58    ),
59    (
60        "014",
61        include_str!("../../migrations/014_drop_ineffective_trgm_index.sql"),
62    ),
63    ("015", include_str!("../../migrations/015_assets.sql")),
64    (
65        "016",
66        include_str!("../../migrations/016_assets_content_hash.sql"),
67    ),
68    ("017", include_str!("../../migrations/017_mcp_tokens.sql")),
69    (
70        "018",
71        include_str!("../../migrations/018_session_generation_trigger.sql"),
72    ),
73    (
74        "019",
75        include_str!("../../migrations/019_restore_search_trgm_index.sql"),
76    ),
77    ("020", include_str!("../../migrations/020_friend_links.sql")),
78    (
79        "021",
80        include_str!("../../migrations/021_site_settings.sql"),
81    ),
82    (
83        "022",
84        include_str!("../../migrations/022_backup_settings.sql"),
85    ),
86    ("023", include_str!("../../migrations/023_user_profile.sql")),
87    (
88        "024",
89        include_str!("../../migrations/024_comments_user_id.sql"),
90    ),
91    ("025", include_str!("../../migrations/025_logs.sql")),
92    (
93        "026",
94        include_str!("../../migrations/026_katex_html_classes.sql"),
95    ),
96    ("027", include_str!("../../migrations/027_notes.sql")),
97    (
98        "028",
99        include_str!("../../migrations/028_notebook_archive.sql"),
100    ),
101    // 新增迁移在此追加,同时在 migrations/ 下创建对应 .sql 文件。
102];
103
104/// 迁移执行错误。
105#[derive(Debug)]
106pub enum MigrateError {
107    /// 无法从连接池获取连接。
108    Pool(deadpool_postgres::PoolError),
109    /// 执行查询(建表、查版本、咨询锁等)失败。
110    Query(tokio_postgres::Error),
111    /// 某个具体迁移执行失败(包含版本号便于定位)。
112    Apply {
113        version: String,
114        source: tokio_postgres::Error,
115    },
116    /// 迁移 026 的 HTML 解析/改写失败;整笔迁移回滚,不丢弃正文。
117    KatexHtml {
118        table: &'static str,
119        id: i64,
120        source: lol_html::errors::RewritingError,
121    },
122}
123
124impl From<deadpool_postgres::PoolError> for MigrateError {
125    fn from(e: deadpool_postgres::PoolError) -> Self {
126        MigrateError::Pool(e)
127    }
128}
129
130impl From<tokio_postgres::Error> for MigrateError {
131    fn from(e: tokio_postgres::Error) -> Self {
132        MigrateError::Query(e)
133    }
134}
135
136impl std::fmt::Display for MigrateError {
137    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
138        // 主动展开 source 链:tokio_postgres::Error 的 Display 对 DB 侧错误只打印
139        // 无信息量的 `db error`,真正的 message/SQLSTATE/约束名藏在 source() 里。
140        match self {
141            MigrateError::Pool(e) => {
142                write!(
143                    f,
144                    "database pool error: {}",
145                    crate::db::format_with_sources(e)
146                )
147            }
148            MigrateError::Query(e) => {
149                write!(
150                    f,
151                    "database query error: {}",
152                    crate::db::format_with_sources(e)
153                )
154            }
155            MigrateError::Apply { version, source } => {
156                write!(
157                    f,
158                    "migration {} failed: {}",
159                    version,
160                    crate::db::format_with_sources(source)
161                )
162            }
163            MigrateError::KatexHtml { table, id, source } => {
164                write!(
165                    f,
166                    "migration 026 failed rewriting {table} row {id}: {source}"
167                )
168            }
169        }
170    }
171}
172
173impl std::error::Error for MigrateError {
174    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
175        match self {
176            MigrateError::Pool(e) => Some(e),
177            MigrateError::Query(e) => Some(e),
178            MigrateError::Apply { source, .. } => Some(source),
179            MigrateError::KatexHtml { source, .. } => Some(source),
180        }
181    }
182}
183
184/// 在**已获取**的连接上执行迁移主体逻辑(咨询锁 + 建表 + 应用迁移 + 解锁)。
185///
186/// 调用方负责自行控制连接获取策略——例如 `main.rs` 启动时用
187/// [`get_conn_for_startup`](crate::db::pool::get_conn_for_startup)(长重试窗口)
188/// 拿到连接后再调用本函数,以应对 DB 尚未就绪的场景。
189///
190/// 流程:
191/// 1. 抢咨询锁(多实例启动时串行化)
192/// 2. 确保 `schema_migrations` 表存在
193/// 3. 查询已应用版本集合
194/// 4. 按序应用未应用的迁移(每个一个事务)
195/// 5. 释放咨询锁
196///
197/// 失败时返回错误;调用方(`main.rs`)应让进程退出,避免启动半残服务。
198pub async fn run_on_conn(conn: &mut deadpool_postgres::Object) -> Result<(), MigrateError> {
199    // 抢咨询锁:多实例滚动发布时只有一个进程能进入迁移循环,
200    // 其余实例在此等待;锁释放后它们查版本表发现已全部应用,直接返回。
201    conn.execute("SELECT pg_advisory_lock($1)", &[&ADVISORY_LOCK_KEY])
202        .await?;
203
204    // 在已持锁连接上执行迁移主体逻辑。
205    // 锁释放策略:
206    // - 正常返回 / 返回 Err:下面的显式 pg_advisory_unlock 释放锁。
207    // - 进程被强杀(SIGKILL 等):连接断开,Postgres 在检测到会话断开后释放
208    //   session 级咨询锁。
209    // - `main.rs` 在迁移失败时用 `std::process::exit(1)` 终止进程(不再 panic):
210    //   `exit(1)` 不会 unwind,但同样会关闭进程持有的所有 socket / 池连接,
211    //   效果等价于会话断开——Postgres 会释放 session 级咨询锁。
212    //   因此把 `.expect()` 改成 `exit(1)` 不破坏原有的锁安全保证。
213    let result = run_inner(conn).await;
214
215    // 无论成功失败都尝试显式释放锁;释放失败不应掩盖原始错误,仅记录告警。
216    if let Err(unlock_err) = conn
217        .execute("SELECT pg_advisory_unlock($1)", &[&ADVISORY_LOCK_KEY])
218        .await
219    {
220        tracing::warn!("failed to release migration advisory lock: {}", unlock_err);
221    }
222
223    result
224}
225
226/// 在已持有咨询锁的连接上执行迁移主体逻辑。
227async fn run_inner(conn: &mut deadpool_postgres::Object) -> Result<(), MigrateError> {
228    // 确保版本表存在(独立语句,不在事务里,否则建表失败无法记录)。
229    ensure_versions_table(conn).await?;
230
231    // 查询已应用的版本集合。
232    let applied = applied_versions(conn).await?;
233
234    // 按序应用未应用的迁移。
235    let mut applied_count = 0usize;
236    for (version, sql) in MIGRATIONS {
237        if applied.contains(*version) {
238            continue;
239        }
240        tracing::info!("applying migration {}", version);
241        apply_one(conn, version, sql).await?;
242        applied_count += 1;
243    }
244
245    if applied_count == 0 {
246        tracing::info!("database is up to date, 0 migrations applied");
247    } else {
248        tracing::info!("successfully applied {} migration(s)", applied_count);
249    }
250    Ok(())
251}
252
253/// 创建 `schema_migrations` 表(若不存在)。
254async fn ensure_versions_table(conn: &deadpool_postgres::Object) -> Result<(), MigrateError> {
255    conn.batch_execute(
256        "CREATE TABLE IF NOT EXISTS schema_migrations (
257            version    TEXT PRIMARY KEY,
258            applied_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
259        )",
260    )
261    .await?;
262    Ok(())
263}
264
265/// 查询已应用的版本集合。
266async fn applied_versions(
267    conn: &deadpool_postgres::Object,
268) -> Result<HashSet<String>, MigrateError> {
269    let rows = conn
270        .query("SELECT version FROM schema_migrations", &[])
271        .await?;
272    let mut set = HashSet::with_capacity(rows.len());
273    for row in rows {
274        set.insert(row.get::<_, String>(0));
275    }
276    Ok(set)
277}
278
279/// 在一个事务内应用单个迁移:执行 SQL + 写入版本行,失败则回滚。
280async fn apply_one(
281    conn: &mut deadpool_postgres::Object,
282    version: &str,
283    sql: &str,
284) -> Result<(), MigrateError> {
285    let tx = conn.transaction().await.map_err(MigrateError::Query)?;
286
287    // batch_execute 执行整段 SQL(可含多条语句)。
288    if let Err(e) = tx.batch_execute(sql).await {
289        // 显式回滚以尽早释放事务(而非等 Transaction drop 的隐式回滚);
290        // 回滚本身的错误丢弃,因为已有更具信息量的 apply 错误要上报。
291        let _ = tx.rollback().await;
292        return Err(MigrateError::Apply {
293            version: version.to_string(),
294            source: e,
295        });
296    }
297
298    if version == "026" {
299        if let Err(error) = migrate_katex_html(&tx).await {
300            let _ = tx.rollback().await;
301            return Err(error);
302        }
303    }
304
305    // 记录版本行。显式构造 Apply 错误(不能用 ?,否则会被 blanket From 映射成 Query)。
306    if let Err(e) = tx
307        .execute(
308            "INSERT INTO schema_migrations (version) VALUES ($1)",
309            &[&version],
310        )
311        .await
312    {
313        let _ = tx.rollback().await;
314        return Err(MigrateError::Apply {
315            version: version.to_string(),
316            source: e,
317        });
318    }
319
320    tx.commit().await.map_err(|e| MigrateError::Apply {
321        version: version.to_string(),
322        source: e,
323    })?;
324    Ok(())
325}
326
327/// 一次性格式迁移:不重渲染 Markdown、不触碰原文/目录/时间戳/审核状态。
328/// SQL 已锁住两张表,批次仅限制内存;所有批次与版本行仍属于同一事务。
329async fn migrate_katex_html(tx: &deadpool_postgres::Transaction<'_>) -> Result<(), MigrateError> {
330    let db_error = |source| MigrateError::Apply {
331        version: "026".to_string(),
332        source,
333    };
334    // 只暂停时间戳触发器;外键等完整性约束照常执行。失败时 DDL 一同回滚。
335    tx.batch_execute("ALTER TABLE comments DISABLE TRIGGER trg_comments_updated_at")
336        .await
337        .map_err(db_error)?;
338
339    for table in ["posts", "comments"] {
340        // table 仅来自上面的常量列表,绝不接受外部输入。
341        let select = tx
342            .prepare(&format!(
343                "SELECT id::bigint, content_html FROM {table} \
344                 WHERE ($1::bigint IS NULL OR id > $1) AND content_html LIKE '%katex%' \
345                 ORDER BY id LIMIT 100"
346            ))
347            .await
348            .map_err(db_error)?;
349        let update = tx
350            .prepare(&format!(
351                "UPDATE {table} SET content_html = $1 WHERE id = $2::bigint"
352            ))
353            .await
354            .map_err(db_error)?;
355        let mut last_id: Option<i64> = None;
356        loop {
357            let rows = tx.query(&select, &[&last_id]).await.map_err(db_error)?;
358            if rows.is_empty() {
359                break;
360            }
361            for row in rows {
362                let id: i64 = row.get(0);
363                let html: &str = row.get(1);
364                if let Some(rewritten) = migrate_katex_classes(html)
365                    .map_err(|source| MigrateError::KatexHtml { table, id, source })?
366                {
367                    tx.execute(&update, &[&rewritten, &id])
368                        .await
369                        .map_err(db_error)?;
370                }
371                last_id = Some(id);
372            }
373        }
374    }
375
376    tx.batch_execute("ALTER TABLE comments ENABLE TRIGGER trg_comments_updated_at")
377        .await
378        .map_err(db_error)?;
379    Ok(())
380}
381
382/// KaTeX 0.18 的精确 class 改名表(上游 commit 6f5c44f,katex-rs 0.3)。
383/// `hbox` 的未使用 CSS 被上游删除,没有对应的旧渲染节点需要迁移。
384fn is_old_katex_class(class: &str) -> bool {
385    matches!(
386        class,
387        "accent"
388            | "base"
389            | "fix"
390            | "hdashline"
391            | "hline"
392            | "inner"
393            | "newline"
394            | "overlay"
395            | "overline"
396            | "root"
397            | "rule"
398            | "sizing"
399            | "smash"
400            | "sout"
401            | "stretchy"
402            | "strut"
403            | "tag"
404            | "thinbox"
405            | "underline"
406            | "vbox"
407    )
408}
409
410/// 仅修改公式视觉层 span 的完整 class token,保留文字、其它属性和注释。
411/// lol_html 不重新序列化未改动的节点;无旧 class 返回 None,避免无效 UPDATE。
412fn migrate_katex_classes(html: &str) -> Result<Option<String>, lol_html::errors::RewritingError> {
413    let mut changed = false;
414    let settings = lol_html::RewriteStrSettings::new().append_element_content_handler(
415        lol_html::element!("span.katex > span.katex-html span[class]", |el| {
416            let Some(classes) = el.get_attribute("class") else {
417                return Ok(());
418            };
419            if !classes.split_ascii_whitespace().any(is_old_katex_class) {
420                return Ok(());
421            }
422            let mut rewritten = String::with_capacity(classes.len() + 6);
423            for class in classes.split_ascii_whitespace() {
424                if !rewritten.is_empty() {
425                    rewritten.push(' ');
426                }
427                if is_old_katex_class(class) {
428                    rewritten.push_str("katex-");
429                }
430                rewritten.push_str(class);
431            }
432            el.set_attribute("class", &rewritten)?;
433            changed = true;
434            Ok(())
435        }),
436    );
437    let rewritten = lol_html::rewrite_str(html, settings)?;
438    Ok(changed.then_some(rewritten))
439}
440
441#[cfg(all(test, feature = "server"))]
442mod tests {
443    use super::*;
444
445    #[test]
446    fn katex_migration_preserves_non_formula_html() {
447        let before = r#"<!-- base class="overline" --><p class="base underline" title="overline">base &amp; underline</p><pre><code>&lt;span class="base"&gt;</code></pre>"#;
448        let formula = r#"<span class="katex"><span class="katex-html" aria-hidden="true"><span class="base"><span class="mord overline" title="base &amp; overline">overline &amp; base</span></span></span></span>"#;
449        let after = r#"<span class="overline">unchanged</span><span class="katex"><span class="base">not a rendered visual layer</span></span>"#;
450        let input = format!("{before}{formula}{after}");
451        let output = migrate_katex_classes(&input).unwrap().unwrap();
452        assert!(output.starts_with(before));
453        assert!(output.ends_with(after));
454        assert!(output.contains(r#"class="katex-base""#));
455        assert!(output.contains(
456            r#"class="mord katex-overline" title="base &amp; overline">overline &amp; base</span>"#
457        ));
458        assert_eq!(migrate_katex_classes(&output).unwrap(), None);
459    }
460
461    #[test]
462    fn katex_migration_renames_only_complete_layout_tokens() {
463        let input = r#"<span class="katex"><span class="katex-html"><span class="accent base fix hdashline hline inner newline overlay overline root rule sizing smash sout stretchy strut tag thinbox underline vbox mord overline-line underline-line accent-body reset-size6 size3 mtight my-base katex-base">x</span></span></span>"#;
464        let expected = r#"<span class="katex"><span class="katex-html"><span class="katex-accent katex-base katex-fix katex-hdashline katex-hline katex-inner katex-newline katex-overlay katex-overline katex-root katex-rule katex-sizing katex-smash katex-sout katex-stretchy katex-strut katex-tag katex-thinbox katex-underline katex-vbox mord overline-line underline-line accent-body reset-size6 size3 mtight my-base katex-base">x</span></span></span>"#;
465        assert_eq!(
466            migrate_katex_classes(input).unwrap().as_deref(),
467            Some(expected)
468        );
469    }
470
471    #[test]
472    fn migrations_are_sorted_ascending() {
473        let mut sorted = MIGRATIONS.iter().map(|(v, _)| *v).collect::<Vec<_>>();
474        sorted.sort_unstable();
475        let original: Vec<&str> = MIGRATIONS.iter().map(|(v, _)| *v).collect();
476        assert_eq!(
477            original, sorted,
478            "MIGRATIONS must be in ascending version order"
479        );
480    }
481
482    #[test]
483    fn migrations_have_unique_versions() {
484        let mut versions: Vec<&str> = MIGRATIONS.iter().map(|(v, _)| *v).collect();
485        let total = versions.len();
486        versions.sort_unstable();
487        versions.dedup();
488        assert_eq!(
489            versions.len(),
490            total,
491            "MIGRATIONS has duplicate version strings"
492        );
493    }
494
495    /// 防止"新建了 .sql 但忘记在 MIGRATIONS 加行"的脚枪。
496    /// 扫描 migrations/ 目录,断言每个 .sql 文件都在 MIGRATIONS 里有对应版本行。
497    /// 仅在 server feature + test 下运行(WASM 无文件系统)。
498    #[test]
499    fn migrations_match_files_on_disk() {
500        use std::collections::HashSet;
501        use std::fs;
502
503        // CARGO_MANIFEST_DIR 指向 crate 根目录(yggdrasil/)。
504        let manifest_dir = env!("CARGO_MANIFEST_DIR");
505        let migrations_dir = std::path::Path::new(manifest_dir).join("migrations");
506
507        let mut files_on_disk: HashSet<String> = HashSet::new();
508        for entry in fs::read_dir(&migrations_dir)
509            .unwrap_or_else(|e| panic!("failed to read {}: {}", migrations_dir.display(), e))
510        {
511            let entry = entry.unwrap();
512            let path = entry.path();
513            if path.extension().and_then(|e| e.to_str()) == Some("sql") {
514                let filename = path
515                    .file_name()
516                    .and_then(|n| n.to_str())
517                    .unwrap_or_else(|| panic!("non-utf8 filename: {}", path.display()));
518                // 文件名形如 "001_init.sql",取前 3 位数字作为 version。
519                let version = filename
520                    .split('_')
521                    .next()
522                    .unwrap_or_else(|| panic!("filename has no '_' separator: {}", filename));
523                files_on_disk.insert(version.to_string());
524            }
525        }
526
527        let versions_in_array: HashSet<String> =
528            MIGRATIONS.iter().map(|(v, _)| v.to_string()).collect();
529
530        // 磁盘上有但数组里没有 → 忘记加行(会静默不执行该迁移)。
531        let missing_in_array: Vec<&String> = files_on_disk.difference(&versions_in_array).collect();
532        assert!(
533            missing_in_array.is_empty(),
534            "migrations/*.sql files not registered in MIGRATIONS: {:?}. \
535             Add a row for each in src/db/migrate.rs.",
536            missing_in_array
537        );
538
539        // 数组里有但磁盘上没有 → include_str! 本就会编译失败,这里只是双保险。
540        let missing_on_disk: Vec<&String> = versions_in_array.difference(&files_on_disk).collect();
541        assert!(
542            missing_on_disk.is_empty(),
543            "MIGRATIONS rows without a corresponding .sql file: {:?}",
544            missing_on_disk
545        );
546    }
547}