feat: 实现 SQLite 初始化与 v1 迁移 (#7)
This commit is contained in:
@@ -0,0 +1,86 @@
|
||||
"""Client SQLite 数据库的路径、连接和初始化函数。"""
|
||||
|
||||
import sqlite3
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Optional, Union
|
||||
|
||||
from .db_schema import MIGRATIONS, SCHEMA_VERSION
|
||||
|
||||
|
||||
PathValue = Union[str, Path]
|
||||
|
||||
|
||||
class DatabaseVersionError(RuntimeError):
|
||||
"""数据库版本比当前程序支持的版本新。"""
|
||||
|
||||
|
||||
def data_dir() -> Path:
|
||||
"""返回可写数据目录;目录不存在时自动创建。"""
|
||||
|
||||
if getattr(sys, "frozen", False):
|
||||
directory = Path(sys.executable).resolve().parent / "data"
|
||||
else:
|
||||
directory = Path(__file__).resolve().parents[1] / "data"
|
||||
directory.mkdir(parents=True, exist_ok=True)
|
||||
return directory
|
||||
|
||||
|
||||
def default_database_path() -> Path:
|
||||
"""返回默认数据库文件路径。"""
|
||||
|
||||
return data_dir() / "client.db"
|
||||
|
||||
|
||||
def open_database(db_path: Optional[PathValue] = None) -> sqlite3.Connection:
|
||||
"""打开一个独立连接,并启用项目要求的 SQLite 参数。
|
||||
|
||||
调用者用完后必须关闭连接。不同线程不能共享同一个连接。
|
||||
"""
|
||||
|
||||
path = Path(db_path) if db_path is not None else default_database_path()
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
connection = sqlite3.connect(str(path), timeout=5.0)
|
||||
connection.row_factory = sqlite3.Row
|
||||
connection.execute("PRAGMA foreign_keys = ON")
|
||||
connection.execute("PRAGMA journal_mode = WAL")
|
||||
connection.execute("PRAGMA busy_timeout = 5000")
|
||||
return connection
|
||||
|
||||
|
||||
def migrate_database(connection: sqlite3.Connection) -> None:
|
||||
"""把数据库从当前版本依次迁移到最新版本。"""
|
||||
|
||||
current_version = int(connection.execute("PRAGMA user_version").fetchone()[0])
|
||||
if current_version > SCHEMA_VERSION:
|
||||
raise DatabaseVersionError(
|
||||
f"数据库版本 {current_version} 高于程序支持版本 {SCHEMA_VERSION}"
|
||||
)
|
||||
|
||||
for version in range(current_version + 1, SCHEMA_VERSION + 1):
|
||||
statements = MIGRATIONS.get(version)
|
||||
if statements is None:
|
||||
raise RuntimeError(f"缺少数据库迁移版本 {version}")
|
||||
|
||||
try:
|
||||
connection.execute("BEGIN IMMEDIATE")
|
||||
for statement in statements:
|
||||
connection.execute(statement)
|
||||
connection.execute(f"PRAGMA user_version = {version}")
|
||||
connection.commit()
|
||||
except Exception:
|
||||
connection.rollback()
|
||||
raise
|
||||
|
||||
|
||||
def initialize_database(db_path: Optional[PathValue] = None) -> Path:
|
||||
"""创建或升级数据库,完成后返回数据库文件路径。"""
|
||||
|
||||
path = Path(db_path) if db_path is not None else default_database_path()
|
||||
connection = open_database(path)
|
||||
try:
|
||||
migrate_database(connection)
|
||||
finally:
|
||||
connection.close()
|
||||
return path
|
||||
Reference in New Issue
Block a user