refactor: establish database migration foundation

This commit is contained in:
leefer
2026-07-29 17:47:09 +08:00
parent 831291c818
commit 6ac9571ca0
9 changed files with 302 additions and 17 deletions
+11
View File
@@ -0,0 +1,11 @@
from .connection import ManagedConnection, SQLiteConnectionFactory
from .migrations import MIGRATIONS, Migration, MigrationError, MigrationRunner
__all__ = [
"MIGRATIONS",
"ManagedConnection",
"Migration",
"MigrationError",
"MigrationRunner",
"SQLiteConnectionFactory",
]
+33
View File
@@ -0,0 +1,33 @@
from __future__ import annotations
import sqlite3
from dataclasses import dataclass
from pathlib import Path
class ManagedConnection(sqlite3.Connection):
"""Commit or roll back, then release the SQLite handle on context exit."""
def __exit__(self, exc_type, exc_value, traceback):
try:
return super().__exit__(exc_type, exc_value, traceback)
finally:
self.close()
@dataclass(frozen=True)
class SQLiteConnectionFactory:
path: Path
timeout_seconds: float = 20
def connect(self) -> sqlite3.Connection:
connection = sqlite3.connect(
self.path,
timeout=self.timeout_seconds,
factory=ManagedConnection,
)
connection.row_factory = sqlite3.Row
connection.execute("PRAGMA journal_mode=WAL")
connection.execute("PRAGMA foreign_keys=ON")
connection.execute("PRAGMA busy_timeout=20000")
return connection
+6
View File
@@ -0,0 +1,6 @@
from .m0001_adopt_legacy import MIGRATION as M0001_ADOPT_LEGACY
from .runner import Migration, MigrationError, MigrationRunner
MIGRATIONS = (M0001_ADOPT_LEGACY,)
__all__ = ["MIGRATIONS", "Migration", "MigrationError", "MigrationRunner"]
@@ -0,0 +1,42 @@
from __future__ import annotations
import sqlite3
from backend.database.migrations.runner import Migration, MigrationError
REQUIRED_TABLES = frozenset(
{
"users",
"user_sessions",
"dashboard_snapshots",
"watchlist",
"review_notes",
"stock_master",
"daily_bars",
"screener_runs",
"mentor_messages",
"trade_entries",
"heaven_readings",
}
)
def adopt_legacy_schema(connection: sqlite3.Connection) -> None:
tables = {
str(row["name"])
for row in connection.execute(
"SELECT name FROM sqlite_master WHERE type = 'table'"
)
}
missing = sorted(REQUIRED_TABLES - tables)
if missing:
raise MigrationError(f"Legacy schema is incomplete: {', '.join(missing)}")
MIGRATION = Migration(
version="0001",
name="adopt_legacy_schema",
action=adopt_legacy_schema,
signature="required-tables:v1:" + ",".join(sorted(REQUIRED_TABLES)),
)
+98
View File
@@ -0,0 +1,98 @@
from __future__ import annotations
import hashlib
import sqlite3
from collections.abc import Callable, Iterable
from dataclasses import dataclass
from datetime import datetime, timezone
MigrationAction = Callable[[sqlite3.Connection], None]
class MigrationError(RuntimeError):
pass
@dataclass(frozen=True)
class Migration:
version: str
name: str
action: MigrationAction
signature: str
@property
def checksum(self) -> str:
return hashlib.sha256(self.signature.encode("utf-8")).hexdigest()
class MigrationRunner:
def apply(
self,
connection: sqlite3.Connection,
migrations: Iterable[Migration],
) -> tuple[str, ...]:
ordered = sorted(migrations, key=lambda item: item.version)
versions = [item.version for item in ordered]
if versions != sorted(set(versions)):
raise MigrationError("Migration versions must be unique and ordered")
self._ensure_ledger(connection)
applied = {
str(row["version"]): str(row["checksum"])
for row in connection.execute(
"SELECT version, checksum FROM schema_migrations ORDER BY version"
)
}
known = set(versions)
unknown = sorted(set(applied) - known)
if unknown:
raise MigrationError(f"Database contains unknown migrations: {', '.join(unknown)}")
completed: list[str] = []
for migration in ordered:
existing = applied.get(migration.version)
if existing:
if existing != migration.checksum:
raise MigrationError(
f"Migration checksum changed: {migration.version} {migration.name}"
)
continue
savepoint = f"migration_{migration.version.replace('-', '_')}"
connection.execute(f"SAVEPOINT {savepoint}")
try:
migration.action(connection)
connection.execute(
"""
INSERT INTO schema_migrations
(version, name, checksum, applied_at)
VALUES (?, ?, ?, ?)
""",
(
migration.version,
migration.name,
migration.checksum,
datetime.now(timezone.utc).isoformat(),
),
)
connection.execute(f"RELEASE SAVEPOINT {savepoint}")
except Exception as exc:
connection.execute(f"ROLLBACK TO SAVEPOINT {savepoint}")
connection.execute(f"RELEASE SAVEPOINT {savepoint}")
raise MigrationError(
f"Migration failed: {migration.version} {migration.name}"
) from exc
completed.append(migration.version)
return tuple(completed)
@staticmethod
def _ensure_ledger(connection: sqlite3.Connection) -> None:
connection.execute(
"""
CREATE TABLE IF NOT EXISTS schema_migrations (
version TEXT PRIMARY KEY,
name TEXT NOT NULL,
checksum TEXT NOT NULL,
applied_at TEXT NOT NULL
)
"""
)
+5 -15
View File
@@ -6,6 +6,8 @@ from datetime import datetime, timezone
from pathlib import Path
from typing import Any
from backend.database import MIGRATIONS, MigrationRunner, SQLiteConnectionFactory
def _optional_float(value: Any) -> float | None:
if value in (None, ""):
@@ -16,28 +18,15 @@ def _optional_float(value: Any) -> float | None:
return None
class ManagedConnection(sqlite3.Connection):
"""Commit or roll back, then release the SQLite file handle on context exit."""
def __exit__(self, exc_type, exc_value, traceback):
try:
return super().__exit__(exc_type, exc_value, traceback)
finally:
self.close()
class ReviewDatabase:
def __init__(self, path: Path) -> None:
self.path = path
self.path.parent.mkdir(parents=True, exist_ok=True)
self.connection_factory = SQLiteConnectionFactory(self.path)
self._initialize()
def connect(self) -> sqlite3.Connection:
connection = sqlite3.connect(self.path, timeout=20, factory=ManagedConnection)
connection.row_factory = sqlite3.Row
connection.execute("PRAGMA journal_mode=WAL")
connection.execute("PRAGMA foreign_keys=ON")
return connection
return self.connection_factory.connect()
def _initialize(self) -> None:
with self.connect() as connection:
@@ -656,6 +645,7 @@ class ReviewDatabase:
ON screener_runs(user_id, mode, trade_date DESC, id DESC)
"""
)
MigrationRunner().apply(connection, MIGRATIONS)
def count_users(self) -> int:
with self.connect() as connection:
+2 -2
View File
@@ -280,8 +280,8 @@
},
{
"path": "database.py",
"bytes": 121468,
"lines": 2839
"bytes": 121153,
"lines": 2829
},
{
"path": "screener.py",
@@ -0,0 +1,28 @@
# Stage 08: Database Connection and Migration Foundation
Date: 2026-07-29
## Result
- Centralized SQLite connection policy in `SQLiteConnectionFactory`.
- Preserved WAL, foreign-key enforcement, row mapping, handle cleanup, and the existing
20-second contention tolerance.
- Added an ordered migration runner with immutable checksums and an applied-migration ledger.
- Added savepoint rollback so a failed migration cannot be recorded or leave partial schema.
- Adopted existing databases as version `0001` only after verifying the required legacy
tables.
- Kept the legacy idempotent bootstrap in place for compatibility with databases created by
every previous application version.
## Forward Rule
All schema changes after this stage must be a new immutable module under
`backend/database/migrations`. Editing an applied migration is rejected by checksum. A
database containing a migration unknown to the running code is rejected rather than opened
with an older schema interpretation.
## Residual Risk
The historical inline bootstrap remains a compatibility facade during repository migration.
It may be removed only after legacy upgrade fixtures cover every supported historical shape.
No user table or row is rewritten in this stage.
+77
View File
@@ -0,0 +1,77 @@
from __future__ import annotations
import sqlite3
import tempfile
import unittest
from pathlib import Path
from backend.database import Migration, MigrationError, MigrationRunner
from database import ReviewDatabase
class DatabaseMigrationTests(unittest.TestCase):
def test_fresh_database_records_the_adopted_schema_once(self) -> None:
with tempfile.TemporaryDirectory() as root:
path = Path(root) / "review.db"
database = ReviewDatabase(path)
with database.connect() as connection:
rows = connection.execute(
"SELECT version, name FROM schema_migrations"
).fetchall()
self.assertEqual(
[(row["version"], row["name"]) for row in rows],
[("0001", "adopt_legacy_schema")],
)
ReviewDatabase(path)
with database.connect() as connection:
count = connection.execute(
"SELECT COUNT(*) AS count FROM schema_migrations"
).fetchone()["count"]
self.assertEqual(count, 1)
def test_connection_factory_enables_required_pragmas(self) -> None:
with tempfile.TemporaryDirectory() as root:
database = ReviewDatabase(Path(root) / "review.db")
with database.connect() as connection:
self.assertEqual(connection.execute("PRAGMA foreign_keys").fetchone()[0], 1)
self.assertEqual(connection.execute("PRAGMA journal_mode").fetchone()[0], "wal")
self.assertEqual(connection.execute("PRAGMA busy_timeout").fetchone()[0], 20000)
def test_failed_migration_rolls_back_and_is_not_recorded(self) -> None:
connection = sqlite3.connect(":memory:")
self.addCleanup(connection.close)
connection.row_factory = sqlite3.Row
def fail(conn: sqlite3.Connection) -> None:
conn.execute("CREATE TABLE should_rollback (id INTEGER)")
raise RuntimeError("stop")
migration = Migration("9000", "failure", fail, "failure:v1")
with self.assertRaises(MigrationError):
MigrationRunner().apply(connection, (migration,))
tables = {
row["name"]
for row in connection.execute(
"SELECT name FROM sqlite_master WHERE type = 'table'"
)
}
self.assertNotIn("should_rollback", tables)
self.assertEqual(
connection.execute("SELECT COUNT(*) FROM schema_migrations").fetchone()[0],
0,
)
def test_applied_migration_checksum_is_immutable(self) -> None:
connection = sqlite3.connect(":memory:")
self.addCleanup(connection.close)
connection.row_factory = sqlite3.Row
first = Migration("9001", "example", lambda conn: None, "example:v1")
changed = Migration("9001", "example", lambda conn: None, "example:v2")
runner = MigrationRunner()
runner.apply(connection, (first,))
with self.assertRaises(MigrationError):
runner.apply(connection, (changed,))
if __name__ == "__main__":
unittest.main()