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
|
||||
@@ -0,0 +1,132 @@
|
||||
"""SQLite 数据库结构和迁移定义。
|
||||
|
||||
每个版本对应一组按顺序执行的 SQL。新增数据库版本时,只能在
|
||||
``MIGRATIONS`` 末尾增加版本,不能修改已经发布的迁移。
|
||||
"""
|
||||
|
||||
SCHEMA_VERSION = 1
|
||||
|
||||
|
||||
MIGRATION_1 = (
|
||||
"""
|
||||
CREATE TABLE pdd_tasks (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
remote_task_id TEXT NOT NULL UNIQUE,
|
||||
task_type TEXT NOT NULL
|
||||
CHECK (task_type IN ('collect', 'purchase')),
|
||||
goods_id TEXT,
|
||||
goods_url TEXT NOT NULL,
|
||||
title TEXT,
|
||||
target_color TEXT,
|
||||
target_size TEXT,
|
||||
price_cent INTEGER
|
||||
CHECK (price_cent IS NULL OR price_cent >= 0),
|
||||
quantity INTEGER
|
||||
CHECK (quantity IS NULL OR quantity > 0),
|
||||
status TEXT NOT NULL DEFAULT 'claimed'
|
||||
CHECK (status IN (
|
||||
'claimed', 'running',
|
||||
'result_pending', 'retry_wait',
|
||||
'manual_review', 'succeeded',
|
||||
'failed', 'cancelled'
|
||||
)),
|
||||
current_step TEXT,
|
||||
priority INTEGER NOT NULL DEFAULT 0,
|
||||
version INTEGER NOT NULL DEFAULT 1 CHECK (version > 0),
|
||||
admin_payload TEXT NOT NULL DEFAULT '{}',
|
||||
pdd_data TEXT,
|
||||
retry_count INTEGER NOT NULL DEFAULT 0 CHECK (retry_count >= 0),
|
||||
last_error_code TEXT,
|
||||
last_error_message TEXT,
|
||||
received_at TEXT NOT NULL,
|
||||
started_at TEXT,
|
||||
finished_at TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL
|
||||
)
|
||||
""",
|
||||
"""
|
||||
CREATE INDEX idx_pdd_tasks_list
|
||||
ON pdd_tasks(updated_at DESC, id DESC)
|
||||
""",
|
||||
"""
|
||||
CREATE INDEX idx_pdd_tasks_status_type
|
||||
ON pdd_tasks(status, task_type)
|
||||
""",
|
||||
"""
|
||||
CREATE INDEX idx_pdd_tasks_goods_id
|
||||
ON pdd_tasks(goods_id)
|
||||
""",
|
||||
"""
|
||||
CREATE TABLE task_runs (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
task_id INTEGER NOT NULL,
|
||||
attempt_id TEXT NOT NULL UNIQUE,
|
||||
attempt_no INTEGER NOT NULL CHECK (attempt_no > 0),
|
||||
device_address TEXT NOT NULL,
|
||||
run_status TEXT NOT NULL
|
||||
CHECK (run_status IN (
|
||||
'running', 'succeeded', 'failed',
|
||||
'cancelled', 'manual_review'
|
||||
)),
|
||||
current_step TEXT,
|
||||
started_at TEXT NOT NULL,
|
||||
finished_at TEXT,
|
||||
irreversible_action_at TEXT,
|
||||
order_submitted_at TEXT,
|
||||
error_code TEXT,
|
||||
error_message TEXT,
|
||||
diagnostics_json TEXT,
|
||||
artifact_directory TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL,
|
||||
FOREIGN KEY (task_id) REFERENCES pdd_tasks(id) ON DELETE CASCADE,
|
||||
UNIQUE (task_id, attempt_no)
|
||||
)
|
||||
""",
|
||||
"""
|
||||
CREATE INDEX idx_task_runs_task
|
||||
ON task_runs(task_id, attempt_no DESC)
|
||||
""",
|
||||
"""
|
||||
CREATE TABLE outbox_events (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
task_id INTEGER NOT NULL,
|
||||
event_type TEXT NOT NULL
|
||||
CHECK (event_type IN (
|
||||
'collect_result', 'purchase_result',
|
||||
'task_failure'
|
||||
)),
|
||||
idempotency_key TEXT NOT NULL UNIQUE,
|
||||
payload_json TEXT NOT NULL,
|
||||
status TEXT NOT NULL DEFAULT 'pending'
|
||||
CHECK (status IN (
|
||||
'pending', 'sending', 'sent', 'failed'
|
||||
)),
|
||||
attempt_count INTEGER NOT NULL DEFAULT 0
|
||||
CHECK (attempt_count >= 0),
|
||||
next_retry_at TEXT,
|
||||
last_error TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL,
|
||||
sent_at TEXT,
|
||||
FOREIGN KEY (task_id) REFERENCES pdd_tasks(id) ON DELETE CASCADE
|
||||
)
|
||||
""",
|
||||
"""
|
||||
CREATE INDEX idx_outbox_pending
|
||||
ON outbox_events(status, next_retry_at, id)
|
||||
""",
|
||||
"""
|
||||
CREATE TABLE app_settings (
|
||||
setting_key TEXT PRIMARY KEY,
|
||||
value_json TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL
|
||||
)
|
||||
""",
|
||||
)
|
||||
|
||||
|
||||
MIGRATIONS = {
|
||||
1: MIGRATION_1,
|
||||
}
|
||||
@@ -0,0 +1,194 @@
|
||||
"""SQLite 初始化和 v1 数据库结构测试。"""
|
||||
|
||||
import sqlite3
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from src.db import DatabaseVersionError, initialize_database, open_database
|
||||
|
||||
|
||||
EXPECTED_TABLES = {
|
||||
"pdd_tasks",
|
||||
"task_runs",
|
||||
"outbox_events",
|
||||
"app_settings",
|
||||
}
|
||||
|
||||
EXPECTED_INDEXES = {
|
||||
"idx_pdd_tasks_list",
|
||||
"idx_pdd_tasks_status_type",
|
||||
"idx_pdd_tasks_goods_id",
|
||||
"idx_task_runs_task",
|
||||
"idx_outbox_pending",
|
||||
}
|
||||
|
||||
|
||||
class DatabaseInitializationTests(unittest.TestCase):
|
||||
"""每个测试都使用独立临时数据库,不接触真实 data 目录。"""
|
||||
|
||||
def setUp(self) -> None:
|
||||
self._temporary_directory = tempfile.TemporaryDirectory()
|
||||
self.db_path = (
|
||||
Path(self._temporary_directory.name) / "nested" / "client.db"
|
||||
)
|
||||
|
||||
def tearDown(self) -> None:
|
||||
self._temporary_directory.cleanup()
|
||||
|
||||
def test_initialize_creates_v1_tables_and_indexes(self) -> None:
|
||||
result_path = initialize_database(self.db_path)
|
||||
|
||||
self.assertEqual(result_path, self.db_path)
|
||||
self.assertTrue(self.db_path.is_file())
|
||||
|
||||
connection = open_database(self.db_path)
|
||||
try:
|
||||
tables = {
|
||||
row[0]
|
||||
for row in connection.execute(
|
||||
"SELECT name FROM sqlite_master WHERE type = 'table'"
|
||||
)
|
||||
}
|
||||
indexes = {
|
||||
row[0]
|
||||
for row in connection.execute(
|
||||
"SELECT name FROM sqlite_master WHERE type = 'index'"
|
||||
)
|
||||
}
|
||||
version = connection.execute("PRAGMA user_version").fetchone()[0]
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
self.assertTrue(EXPECTED_TABLES.issubset(tables))
|
||||
self.assertTrue(EXPECTED_INDEXES.issubset(indexes))
|
||||
self.assertEqual(version, 1)
|
||||
|
||||
def test_initialize_can_run_twice_without_losing_data(self) -> None:
|
||||
initialize_database(self.db_path)
|
||||
connection = open_database(self.db_path)
|
||||
try:
|
||||
with connection:
|
||||
connection.execute(
|
||||
"INSERT INTO app_settings"
|
||||
" (setting_key, value_json, updated_at) VALUES (?, ?, ?)",
|
||||
("automation.dry_run", "true", "2026-08-06T00:00:00Z"),
|
||||
)
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
initialize_database(self.db_path)
|
||||
connection = open_database(self.db_path)
|
||||
try:
|
||||
value = connection.execute(
|
||||
"SELECT value_json FROM app_settings WHERE setting_key = ?",
|
||||
("automation.dry_run",),
|
||||
).fetchone()[0]
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
self.assertEqual(value, "true")
|
||||
|
||||
def test_new_connection_uses_required_pragmas(self) -> None:
|
||||
initialize_database(self.db_path)
|
||||
connection = open_database(self.db_path)
|
||||
try:
|
||||
foreign_keys = connection.execute("PRAGMA foreign_keys").fetchone()[0]
|
||||
journal_mode = connection.execute("PRAGMA journal_mode").fetchone()[0]
|
||||
busy_timeout = connection.execute("PRAGMA busy_timeout").fetchone()[0]
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
self.assertEqual(foreign_keys, 1)
|
||||
self.assertEqual(journal_mode.lower(), "wal")
|
||||
self.assertEqual(busy_timeout, 5000)
|
||||
|
||||
def test_pdd_task_check_constraints_are_enforced(self) -> None:
|
||||
initialize_database(self.db_path)
|
||||
connection = open_database(self.db_path)
|
||||
try:
|
||||
invalid_values = (
|
||||
("unknown", "claimed", None, None),
|
||||
("collect", "unknown", None, None),
|
||||
("collect", "claimed", -1, None),
|
||||
("purchase", "claimed", None, 0),
|
||||
)
|
||||
for index, values in enumerate(invalid_values):
|
||||
with self.subTest(values=values):
|
||||
with self.assertRaises(sqlite3.IntegrityError):
|
||||
connection.execute(
|
||||
"INSERT INTO pdd_tasks"
|
||||
" (remote_task_id, task_type, goods_url, price_cent,"
|
||||
" quantity, status, received_at, created_at, updated_at)"
|
||||
" VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)",
|
||||
(
|
||||
f"TASK-{index}",
|
||||
values[0],
|
||||
"https://example.test/goods",
|
||||
values[2],
|
||||
values[3],
|
||||
values[1],
|
||||
"2026-08-06T00:00:00Z",
|
||||
"2026-08-06T00:00:00Z",
|
||||
"2026-08-06T00:00:00Z",
|
||||
),
|
||||
)
|
||||
connection.rollback()
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
def test_task_run_and_outbox_foreign_keys_are_enforced(self) -> None:
|
||||
initialize_database(self.db_path)
|
||||
connection = open_database(self.db_path)
|
||||
try:
|
||||
with self.assertRaises(sqlite3.IntegrityError):
|
||||
connection.execute(
|
||||
"INSERT INTO task_runs"
|
||||
" (task_id, attempt_id, attempt_no, device_address, run_status,"
|
||||
" started_at, created_at, updated_at)"
|
||||
" VALUES (?, ?, ?, ?, ?, ?, ?, ?)",
|
||||
(
|
||||
999,
|
||||
"attempt-1",
|
||||
1,
|
||||
"device-1",
|
||||
"running",
|
||||
"2026-08-06T00:00:00Z",
|
||||
"2026-08-06T00:00:00Z",
|
||||
"2026-08-06T00:00:00Z",
|
||||
),
|
||||
)
|
||||
connection.rollback()
|
||||
|
||||
with self.assertRaises(sqlite3.IntegrityError):
|
||||
connection.execute(
|
||||
"INSERT INTO outbox_events"
|
||||
" (task_id, event_type, idempotency_key, payload_json,"
|
||||
" created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?)",
|
||||
(
|
||||
999,
|
||||
"collect_result",
|
||||
"TASK-1:attempt-1:result-v1",
|
||||
"{}",
|
||||
"2026-08-06T00:00:00Z",
|
||||
"2026-08-06T00:00:00Z",
|
||||
),
|
||||
)
|
||||
connection.rollback()
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
def test_newer_database_version_is_rejected(self) -> None:
|
||||
initialize_database(self.db_path)
|
||||
connection = open_database(self.db_path)
|
||||
try:
|
||||
connection.execute("PRAGMA user_version = 99")
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
with self.assertRaisesRegex(DatabaseVersionError, "数据库版本 99"):
|
||||
initialize_database(self.db_path)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user