1use std::collections::HashSet;
17
18const ADVISORY_LOCK_KEY: i64 = 0x5947_4752_4153_494C;
22
23const 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 ];
103
104#[derive(Debug)]
106pub enum MigrateError {
107 Pool(deadpool_postgres::PoolError),
109 Query(tokio_postgres::Error),
111 Apply {
113 version: String,
114 source: tokio_postgres::Error,
115 },
116 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 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
184pub async fn run_on_conn(conn: &mut deadpool_postgres::Object) -> Result<(), MigrateError> {
199 conn.execute("SELECT pg_advisory_lock($1)", &[&ADVISORY_LOCK_KEY])
202 .await?;
203
204 let result = run_inner(conn).await;
214
215 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
226async fn run_inner(conn: &mut deadpool_postgres::Object) -> Result<(), MigrateError> {
228 ensure_versions_table(conn).await?;
230
231 let applied = applied_versions(conn).await?;
233
234 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
253async 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
265async 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
279async 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 if let Err(e) = tx.batch_execute(sql).await {
289 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 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
327async 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 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 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
382fn 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
410fn 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 & underline</p><pre><code><span class="base"></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 & overline">overline & 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 & overline">overline & 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 #[test]
499 fn migrations_match_files_on_disk() {
500 use std::collections::HashSet;
501 use std::fs;
502
503 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 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 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 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}