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    // 新增迁移在此追加,同时在 migrations/ 下创建对应 .sql 文件。
79];
80
81/// 迁移执行错误。
82#[derive(Debug)]
83pub enum MigrateError {
84    /// 无法从连接池获取连接。
85    Pool(deadpool_postgres::PoolError),
86    /// 执行查询(建表、查版本、咨询锁等)失败。
87    Query(tokio_postgres::Error),
88    /// 某个具体迁移执行失败(包含版本号便于定位)。
89    Apply {
90        version: String,
91        source: tokio_postgres::Error,
92    },
93}
94
95impl From<deadpool_postgres::PoolError> for MigrateError {
96    fn from(e: deadpool_postgres::PoolError) -> Self {
97        MigrateError::Pool(e)
98    }
99}
100
101impl From<tokio_postgres::Error> for MigrateError {
102    fn from(e: tokio_postgres::Error) -> Self {
103        MigrateError::Query(e)
104    }
105}
106
107impl std::fmt::Display for MigrateError {
108    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
109        // 主动展开 source 链:tokio_postgres::Error 的 Display 对 DB 侧错误只打印
110        // 无信息量的 `db error`,真正的 message/SQLSTATE/约束名藏在 source() 里。
111        match self {
112            MigrateError::Pool(e) => {
113                write!(
114                    f,
115                    "database pool error: {}",
116                    crate::db::format_with_sources(e)
117                )
118            }
119            MigrateError::Query(e) => {
120                write!(
121                    f,
122                    "database query error: {}",
123                    crate::db::format_with_sources(e)
124                )
125            }
126            MigrateError::Apply { version, source } => {
127                write!(
128                    f,
129                    "migration {} failed: {}",
130                    version,
131                    crate::db::format_with_sources(source)
132                )
133            }
134        }
135    }
136}
137
138impl std::error::Error for MigrateError {
139    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
140        match self {
141            MigrateError::Pool(e) => Some(e),
142            MigrateError::Query(e) => Some(e),
143            MigrateError::Apply { source, .. } => Some(source),
144        }
145    }
146}
147
148/// 在**已获取**的连接上执行迁移主体逻辑(咨询锁 + 建表 + 应用迁移 + 解锁)。
149///
150/// 调用方负责自行控制连接获取策略——例如 `main.rs` 启动时用
151/// [`get_conn_for_startup`](crate::db::pool::get_conn_for_startup)(长重试窗口)
152/// 拿到连接后再调用本函数,以应对 DB 尚未就绪的场景。
153///
154/// 流程:
155/// 1. 抢咨询锁(多实例启动时串行化)
156/// 2. 确保 `schema_migrations` 表存在
157/// 3. 查询已应用版本集合
158/// 4. 按序应用未应用的迁移(每个一个事务)
159/// 5. 释放咨询锁
160///
161/// 失败时返回错误;调用方(`main.rs`)应让进程退出,避免启动半残服务。
162pub async fn run_on_conn(conn: &mut deadpool_postgres::Object) -> Result<(), MigrateError> {
163    // 抢咨询锁:多实例滚动发布时只有一个进程能进入迁移循环,
164    // 其余实例在此等待;锁释放后它们查版本表发现已全部应用,直接返回。
165    conn.execute("SELECT pg_advisory_lock($1)", &[&ADVISORY_LOCK_KEY])
166        .await?;
167
168    // 在已持锁连接上执行迁移主体逻辑。
169    // 锁释放策略:
170    // - 正常返回 / 返回 Err:下面的显式 pg_advisory_unlock 释放锁。
171    // - 进程被强杀(SIGKILL 等):连接断开,Postgres 在检测到会话断开后释放
172    //   session 级咨询锁。
173    // - `main.rs` 在迁移失败时用 `std::process::exit(1)` 终止进程(不再 panic):
174    //   `exit(1)` 不会 unwind,但同样会关闭进程持有的所有 socket / 池连接,
175    //   效果等价于会话断开——Postgres 会释放 session 级咨询锁。
176    //   因此把 `.expect()` 改成 `exit(1)` 不破坏原有的锁安全保证。
177    let result = run_inner(conn).await;
178
179    // 无论成功失败都尝试显式释放锁;释放失败不应掩盖原始错误,仅记录告警。
180    if let Err(unlock_err) = conn
181        .execute("SELECT pg_advisory_unlock($1)", &[&ADVISORY_LOCK_KEY])
182        .await
183    {
184        tracing::warn!("failed to release migration advisory lock: {}", unlock_err);
185    }
186
187    result
188}
189
190/// 在已持有咨询锁的连接上执行迁移主体逻辑。
191async fn run_inner(conn: &mut deadpool_postgres::Object) -> Result<(), MigrateError> {
192    // 确保版本表存在(独立语句,不在事务里,否则建表失败无法记录)。
193    ensure_versions_table(conn).await?;
194
195    // 查询已应用的版本集合。
196    let applied = applied_versions(conn).await?;
197
198    // 按序应用未应用的迁移。
199    let mut applied_count = 0usize;
200    for (version, sql) in MIGRATIONS {
201        if applied.contains(*version) {
202            continue;
203        }
204        tracing::info!("applying migration {}", version);
205        apply_one(conn, version, sql).await?;
206        applied_count += 1;
207    }
208
209    if applied_count == 0 {
210        tracing::info!("database is up to date, 0 migrations applied");
211    } else {
212        tracing::info!("successfully applied {} migration(s)", applied_count);
213    }
214    Ok(())
215}
216
217/// 创建 `schema_migrations` 表(若不存在)。
218async fn ensure_versions_table(conn: &deadpool_postgres::Object) -> Result<(), MigrateError> {
219    conn.batch_execute(
220        "CREATE TABLE IF NOT EXISTS schema_migrations (
221            version    TEXT PRIMARY KEY,
222            applied_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
223        )",
224    )
225    .await?;
226    Ok(())
227}
228
229/// 查询已应用的版本集合。
230async fn applied_versions(
231    conn: &deadpool_postgres::Object,
232) -> Result<HashSet<String>, MigrateError> {
233    let rows = conn
234        .query("SELECT version FROM schema_migrations", &[])
235        .await?;
236    let mut set = HashSet::with_capacity(rows.len());
237    for row in rows {
238        set.insert(row.get::<_, String>(0));
239    }
240    Ok(set)
241}
242
243/// 在一个事务内应用单个迁移:执行 SQL + 写入版本行,失败则回滚。
244async fn apply_one(
245    conn: &mut deadpool_postgres::Object,
246    version: &str,
247    sql: &str,
248) -> Result<(), MigrateError> {
249    let tx = conn.transaction().await.map_err(MigrateError::Query)?;
250
251    // batch_execute 执行整段 SQL(可含多条语句)。
252    if let Err(e) = tx.batch_execute(sql).await {
253        // 显式回滚以尽早释放事务(而非等 Transaction drop 的隐式回滚);
254        // 回滚本身的错误丢弃,因为已有更具信息量的 apply 错误要上报。
255        let _ = tx.rollback().await;
256        return Err(MigrateError::Apply {
257            version: version.to_string(),
258            source: e,
259        });
260    }
261
262    // 记录版本行。显式构造 Apply 错误(不能用 ?,否则会被 blanket From 映射成 Query)。
263    if let Err(e) = tx
264        .execute(
265            "INSERT INTO schema_migrations (version) VALUES ($1)",
266            &[&version],
267        )
268        .await
269    {
270        let _ = tx.rollback().await;
271        return Err(MigrateError::Apply {
272            version: version.to_string(),
273            source: e,
274        });
275    }
276
277    tx.commit().await.map_err(|e| MigrateError::Apply {
278        version: version.to_string(),
279        source: e,
280    })?;
281    Ok(())
282}
283
284#[cfg(all(test, feature = "server"))]
285mod tests {
286    use super::*;
287
288    #[test]
289    fn migrations_are_sorted_ascending() {
290        let mut sorted = MIGRATIONS.iter().map(|(v, _)| *v).collect::<Vec<_>>();
291        sorted.sort_unstable();
292        let original: Vec<&str> = MIGRATIONS.iter().map(|(v, _)| *v).collect();
293        assert_eq!(
294            original, sorted,
295            "MIGRATIONS must be in ascending version order"
296        );
297    }
298
299    #[test]
300    fn migrations_have_unique_versions() {
301        let mut versions: Vec<&str> = MIGRATIONS.iter().map(|(v, _)| *v).collect();
302        let total = versions.len();
303        versions.sort_unstable();
304        versions.dedup();
305        assert_eq!(
306            versions.len(),
307            total,
308            "MIGRATIONS has duplicate version strings"
309        );
310    }
311
312    #[test]
313    fn migrations_non_empty() {
314        assert!(!MIGRATIONS.is_empty(), "MIGRATIONS must not be empty");
315    }
316
317    /// 防止"新建了 .sql 但忘记在 MIGRATIONS 加行"的脚枪。
318    /// 扫描 migrations/ 目录,断言每个 .sql 文件都在 MIGRATIONS 里有对应版本行。
319    /// 仅在 server feature + test 下运行(WASM 无文件系统)。
320    #[test]
321    fn migrations_match_files_on_disk() {
322        use std::collections::HashSet;
323        use std::fs;
324
325        // CARGO_MANIFEST_DIR 指向 crate 根目录(yggdrasil/)。
326        let manifest_dir = env!("CARGO_MANIFEST_DIR");
327        let migrations_dir = std::path::Path::new(manifest_dir).join("migrations");
328
329        let mut files_on_disk: HashSet<String> = HashSet::new();
330        for entry in fs::read_dir(&migrations_dir)
331            .unwrap_or_else(|e| panic!("failed to read {}: {}", migrations_dir.display(), e))
332        {
333            let entry = entry.unwrap();
334            let path = entry.path();
335            if path.extension().and_then(|e| e.to_str()) == Some("sql") {
336                let filename = path
337                    .file_name()
338                    .and_then(|n| n.to_str())
339                    .unwrap_or_else(|| panic!("non-utf8 filename: {}", path.display()));
340                // 文件名形如 "001_init.sql",取前 3 位数字作为 version。
341                let version = filename
342                    .split('_')
343                    .next()
344                    .unwrap_or_else(|| panic!("filename has no '_' separator: {}", filename));
345                files_on_disk.insert(version.to_string());
346            }
347        }
348
349        let versions_in_array: HashSet<String> =
350            MIGRATIONS.iter().map(|(v, _)| v.to_string()).collect();
351
352        // 磁盘上有但数组里没有 → 忘记加行(会静默不执行该迁移)。
353        let missing_in_array: Vec<&String> = files_on_disk.difference(&versions_in_array).collect();
354        assert!(
355            missing_in_array.is_empty(),
356            "migrations/*.sql files not registered in MIGRATIONS: {:?}. \
357             Add a row for each in src/db/migrate.rs.",
358            missing_in_array
359        );
360
361        // 数组里有但磁盘上没有 → include_str! 本就会编译失败,这里只是双保险。
362        let missing_on_disk: Vec<&String> = versions_in_array.difference(&files_on_disk).collect();
363        assert!(
364            missing_on_disk.is_empty(),
365            "MIGRATIONS rows without a corresponding .sql file: {:?}",
366            missing_on_disk
367        );
368    }
369}