feat: 定义任务模型与 SQLite Repository (#8)

This commit is contained in:
chengma
2026-08-06 16:21:39 +08:00
parent d64ebef922
commit e60b0c91ac
5 changed files with 770 additions and 0 deletions
+96
View File
@@ -0,0 +1,96 @@
"""非敏感应用设置的 SQLite Repository。"""
import json
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Optional, Union
from .db import initialize_database, open_database
from .task_models import AppSettingRecord
PathValue = Union[str, Path]
def _utc_now_iso() -> str:
return datetime.now(timezone.utc).isoformat(timespec="seconds").replace(
"+00:00", "Z"
)
class SettingsRepository:
"""用 JSON 保存非敏感设置;凭据不得传入本类。"""
def __init__(self, db_path: Optional[PathValue] = None):
self._db_path = initialize_database(db_path)
def get(self, setting_key: str, default: Any = None) -> Any:
"""读取设置值;键不存在时返回 default。"""
record = self.get_record(setting_key)
return default if record is None else record.value
def get_record(self, setting_key: str) -> Optional[AppSettingRecord]:
"""读取包含更新时间的完整设置记录。"""
key = self._validate_key(setting_key)
connection = open_database(self._db_path)
try:
row = connection.execute(
"SELECT setting_key, value_json, updated_at"
" FROM app_settings WHERE setting_key = ?",
(key,),
).fetchone()
finally:
connection.close()
if row is None:
return None
return AppSettingRecord(
setting_key=row["setting_key"],
value=json.loads(row["value_json"]),
updated_at=row["updated_at"],
)
def set(
self, setting_key: str, value: Any, updated_at: Optional[str] = None
) -> AppSettingRecord:
"""新增或更新一个非敏感设置并返回保存后的记录。"""
key = self._validate_key(setting_key)
now = updated_at or _utc_now_iso()
value_json = json.dumps(value, ensure_ascii=False)
connection = open_database(self._db_path)
try:
with connection:
connection.execute(
"INSERT INTO app_settings (setting_key, value_json, updated_at)"
" VALUES (?, ?, ?)"
" ON CONFLICT(setting_key) DO UPDATE SET"
" value_json = excluded.value_json,"
" updated_at = excluded.updated_at",
(key, value_json, now),
)
finally:
connection.close()
return AppSettingRecord(key, value, now)
def delete(self, setting_key: str) -> bool:
"""删除设置;确实删除了一条记录时返回 True。"""
key = self._validate_key(setting_key)
connection = open_database(self._db_path)
try:
with connection:
cursor = connection.execute(
"DELETE FROM app_settings WHERE setting_key = ?", (key,)
)
return cursor.rowcount > 0
finally:
connection.close()
@staticmethod
def _validate_key(setting_key: str) -> str:
key = setting_key.strip()
if not key:
raise ValueError("setting_key 不能为空")
return key
+193
View File
@@ -0,0 +1,193 @@
"""Client 任务数据库使用的枚举和简单数据对象。
本文件不依赖 Qt、Admin 客户端或 uiautomator2,只表达数据含义。
"""
from dataclasses import dataclass, field
from enum import Enum
from typing import Any, Dict, Mapping, Optional
class TaskType(str, Enum):
"""PDD 任务类型。"""
COLLECT = "collect"
PURCHASE = "purchase"
class TaskStatus(str, Enum):
"""PDD 任务本地状态。"""
CLAIMED = "claimed"
RUNNING = "running"
RESULT_PENDING = "result_pending"
RETRY_WAIT = "retry_wait"
MANUAL_REVIEW = "manual_review"
SUCCEEDED = "succeeded"
FAILED = "failed"
CANCELLED = "cancelled"
class RunStatus(str, Enum):
"""单次任务执行状态。"""
RUNNING = "running"
SUCCEEDED = "succeeded"
FAILED = "failed"
CANCELLED = "cancelled"
MANUAL_REVIEW = "manual_review"
class OutboxEventType(str, Enum):
"""等待提交 Admin 的事件类型。"""
COLLECT_RESULT = "collect_result"
PURCHASE_RESULT = "purchase_result"
TASK_FAILURE = "task_failure"
class OutboxStatus(str, Enum):
"""Outbox 事件发送状态。"""
PENDING = "pending"
SENDING = "sending"
SENT = "sent"
FAILED = "failed"
@dataclass(frozen=True)
class NewClaimedTask:
"""刚从 Admin 领取、准备写入本地数据库的任务。"""
remote_task_id: str
task_type: TaskType
goods_url: str
goods_id: Optional[str] = None
title: Optional[str] = None
target_color: Optional[str] = None
target_size: Optional[str] = None
price_cent: Optional[int] = None
quantity: Optional[int] = None
priority: int = 0
version: int = 1
admin_payload: Mapping[str, Any] = field(default_factory=dict)
def __post_init__(self) -> None:
if not isinstance(self.task_type, TaskType):
raise ValueError("task_type 必须是 TaskType")
if not self.remote_task_id.strip():
raise ValueError("remote_task_id 不能为空")
if not self.goods_url.strip():
raise ValueError("goods_url 不能为空")
if self.price_cent is not None and self.price_cent < 0:
raise ValueError("price_cent 不能小于 0")
if self.quantity is not None and self.quantity <= 0:
raise ValueError("quantity 必须大于 0")
if self.version <= 0:
raise ValueError("version 必须大于 0")
@dataclass(frozen=True)
class TaskFilters:
"""本地任务列表的可选筛选条件。"""
task_type: Optional[TaskType] = None
status: Optional[TaskStatus] = None
keyword: str = ""
@dataclass(frozen=True)
class TaskSummary:
"""任务表格使用的轻量数据,不包含完整 JSON。"""
id: int
remote_task_id: str
task_type: TaskType
goods_id: Optional[str]
title: Optional[str]
target_color: Optional[str]
target_size: Optional[str]
price_cent: Optional[int]
quantity: Optional[int]
status: TaskStatus
updated_at: str
@dataclass(frozen=True)
class TaskDetail:
"""一条 PDD 任务的完整本地记录。"""
id: int
remote_task_id: str
task_type: TaskType
goods_id: Optional[str]
goods_url: str
title: Optional[str]
target_color: Optional[str]
target_size: Optional[str]
price_cent: Optional[int]
quantity: Optional[int]
status: TaskStatus
current_step: Optional[str]
priority: int
version: int
admin_payload: Dict[str, Any]
pdd_data: Optional[Dict[str, Any]]
retry_count: int
last_error_code: Optional[str]
last_error_message: Optional[str]
received_at: str
started_at: Optional[str]
finished_at: Optional[str]
created_at: str
updated_at: str
@dataclass(frozen=True)
class TaskRunRecord:
"""一次任务执行的数据库记录。"""
id: int
task_id: int
attempt_id: str
attempt_no: int
device_address: str
run_status: RunStatus
current_step: Optional[str]
started_at: str
finished_at: Optional[str]
irreversible_action_at: Optional[str]
order_submitted_at: Optional[str]
error_code: Optional[str]
error_message: Optional[str]
diagnostics_json: Optional[Dict[str, Any]]
artifact_directory: Optional[str]
created_at: str
updated_at: str
@dataclass(frozen=True)
class OutboxEventRecord:
"""一条等待提交 Admin 的可靠事件记录。"""
id: int
task_id: int
event_type: OutboxEventType
idempotency_key: str
payload_json: Dict[str, Any]
status: OutboxStatus
attempt_count: int
next_retry_at: Optional[str]
last_error: Optional[str]
created_at: str
updated_at: str
sent_at: Optional[str]
@dataclass(frozen=True)
class AppSettingRecord:
"""一条非敏感应用设置。"""
setting_key: str
value: Any
updated_at: str
+236
View File
@@ -0,0 +1,236 @@
"""PDD 任务的 SQLite Repository。
Repository 是数据库访问入口。界面和自动化代码不应自行拼接任务 SQL。
"""
import json
import sqlite3
from datetime import datetime, timezone
from pathlib import Path
from typing import Dict, List, Optional, Tuple, Union
from .db import initialize_database, open_database
from .task_models import (
NewClaimedTask,
TaskDetail,
TaskFilters,
TaskStatus,
TaskSummary,
TaskType,
)
PathValue = Union[str, Path]
MAX_PAGE_SIZE = 500
class DuplicateTaskError(ValueError):
"""相同远程任务编号已经存在,不能覆盖。"""
def utc_now_iso() -> str:
"""返回精确到秒的 UTC ISO 8601 时间。"""
return datetime.now(timezone.utc).isoformat(timespec="seconds").replace(
"+00:00", "Z"
)
class TaskRepository:
"""保存和查询本机已经领取的 PDD 任务。"""
def __init__(self, db_path: Optional[PathValue] = None):
self._db_path = initialize_database(db_path)
def add_claimed_task(
self, task: NewClaimedTask, received_at: Optional[str] = None
) -> int:
"""写入一条新领取任务,返回本地自增编号。
相同 ``remote_task_id`` 已存在时抛出 ``DuplicateTaskError``,
不覆盖已经保存的本地状态。
"""
now = received_at or utc_now_iso()
payload = json.dumps(task.admin_payload, ensure_ascii=False)
connection = open_database(self._db_path)
try:
try:
with connection:
cursor = connection.execute(
"INSERT INTO pdd_tasks ("
" remote_task_id, task_type, goods_id, goods_url, title,"
" target_color, target_size, price_cent, quantity, status,"
" priority, version, admin_payload, received_at, created_at,"
" updated_at"
") VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
(
task.remote_task_id.strip(),
task.task_type.value,
task.goods_id,
task.goods_url.strip(),
task.title,
task.target_color,
task.target_size,
task.price_cent,
task.quantity,
TaskStatus.CLAIMED.value,
task.priority,
task.version,
payload,
now,
now,
now,
),
)
return int(cursor.lastrowid)
except sqlite3.IntegrityError as exc:
if "pdd_tasks.remote_task_id" in str(exc):
raise DuplicateTaskError(
f"任务 {task.remote_task_id} 已经存在"
) from exc
raise
finally:
connection.close()
def list_tasks(
self,
filters: Optional[TaskFilters] = None,
limit: int = 50,
offset: int = 0,
) -> List[TaskSummary]:
"""分页查询任务摘要,默认按更新时间和本地编号倒序。"""
self._validate_page(limit, offset)
where_sql, parameters = self._build_where(filters or TaskFilters())
parameters.extend((limit, offset))
connection = open_database(self._db_path)
try:
rows = connection.execute(
"SELECT id, remote_task_id, task_type, goods_id, title,"
" target_color, target_size, price_cent, quantity, status, updated_at"
" FROM pdd_tasks"
f"{where_sql}"
" ORDER BY updated_at DESC, id DESC"
" LIMIT ? OFFSET ?",
parameters,
).fetchall()
finally:
connection.close()
return [self._to_summary(row) for row in rows]
def count_tasks(self, filters: Optional[TaskFilters] = None) -> int:
"""返回符合筛选条件的任务总数。"""
where_sql, parameters = self._build_where(filters or TaskFilters())
connection = open_database(self._db_path)
try:
row = connection.execute(
f"SELECT COUNT(*) FROM pdd_tasks{where_sql}", parameters
).fetchone()
finally:
connection.close()
return int(row[0])
def get_task(self, remote_task_id: str) -> Optional[TaskDetail]:
"""按稳定远程编号读取完整任务;不存在时返回 None。"""
connection = open_database(self._db_path)
try:
row = connection.execute(
"SELECT * FROM pdd_tasks WHERE remote_task_id = ?",
(remote_task_id,),
).fetchone()
finally:
connection.close()
return self._to_detail(row) if row is not None else None
@staticmethod
def _validate_page(limit: int, offset: int) -> None:
if not 1 <= limit <= MAX_PAGE_SIZE:
raise ValueError(f"limit 必须在 1 到 {MAX_PAGE_SIZE} 之间")
if offset < 0:
raise ValueError("offset 不能小于 0")
@staticmethod
def _build_where(filters: TaskFilters) -> Tuple[str, List[object]]:
clauses = []
parameters: List[object] = []
if filters.task_type is not None:
clauses.append("task_type = ?")
parameters.append(filters.task_type.value)
if filters.status is not None:
clauses.append("status = ?")
parameters.append(filters.status.value)
if filters.keyword.strip():
keyword = TaskRepository._escape_like(filters.keyword.strip())
pattern = f"%{keyword}%"
clauses.append(
"(remote_task_id LIKE ? ESCAPE '\\'"
" OR COALESCE(goods_id, '') LIKE ? ESCAPE '\\'"
" OR COALESCE(title, '') LIKE ? ESCAPE '\\')"
)
parameters.extend((pattern, pattern, pattern))
return (" WHERE " + " AND ".join(clauses) if clauses else "", parameters)
@staticmethod
def _escape_like(value: str) -> str:
return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
@staticmethod
def _to_summary(row: sqlite3.Row) -> TaskSummary:
return TaskSummary(
id=row["id"],
remote_task_id=row["remote_task_id"],
task_type=TaskType(row["task_type"]),
goods_id=row["goods_id"],
title=row["title"],
target_color=row["target_color"],
target_size=row["target_size"],
price_cent=row["price_cent"],
quantity=row["quantity"],
status=TaskStatus(row["status"]),
updated_at=row["updated_at"],
)
@staticmethod
def _to_detail(row: sqlite3.Row) -> TaskDetail:
admin_payload = TaskRepository._load_json_object(row["admin_payload"])
pdd_data = (
TaskRepository._load_json_object(row["pdd_data"])
if row["pdd_data"] is not None
else None
)
return TaskDetail(
id=row["id"],
remote_task_id=row["remote_task_id"],
task_type=TaskType(row["task_type"]),
goods_id=row["goods_id"],
goods_url=row["goods_url"],
title=row["title"],
target_color=row["target_color"],
target_size=row["target_size"],
price_cent=row["price_cent"],
quantity=row["quantity"],
status=TaskStatus(row["status"]),
current_step=row["current_step"],
priority=row["priority"],
version=row["version"],
admin_payload=admin_payload,
pdd_data=pdd_data,
retry_count=row["retry_count"],
last_error_code=row["last_error_code"],
last_error_message=row["last_error_message"],
received_at=row["received_at"],
started_at=row["started_at"],
finished_at=row["finished_at"],
created_at=row["created_at"],
updated_at=row["updated_at"],
)
@staticmethod
def _load_json_object(text: str) -> Dict[str, object]:
value = json.loads(text)
if not isinstance(value, dict):
raise ValueError("数据库 JSON 字段必须是对象")
return value
+57
View File
@@ -0,0 +1,57 @@
"""SettingsRepository 测试。"""
import tempfile
import unittest
from pathlib import Path
from src.settings_repository import SettingsRepository
class SettingsRepositoryTests(unittest.TestCase):
def setUp(self) -> None:
self._temporary_directory = tempfile.TemporaryDirectory()
self.db_path = Path(self._temporary_directory.name) / "client.db"
self.repository = SettingsRepository(self.db_path)
def tearDown(self) -> None:
self._temporary_directory.cleanup()
def test_get_returns_default_for_missing_key(self) -> None:
self.assertEqual(
self.repository.get("automation.max_retries", default=3), 3
)
def test_set_get_and_update_json_value(self) -> None:
first = self.repository.set(
"client.info",
{"device_id": "CLIENT-001", "device_name": "办公室电脑"},
"2026-08-06T08:00:00Z",
)
self.assertEqual(first.updated_at, "2026-08-06T08:00:00Z")
self.assertEqual(
self.repository.get("client.info")["device_name"], "办公室电脑"
)
self.repository.set(
"client.info",
{"device_id": "CLIENT-001", "device_name": "仓库电脑"},
"2026-08-06T09:00:00Z",
)
record = self.repository.get_record("client.info")
self.assertEqual(record.value["device_name"], "仓库电脑")
self.assertEqual(record.updated_at, "2026-08-06T09:00:00Z")
def test_delete_reports_whether_record_existed(self) -> None:
self.repository.set("automation.dry_run", True)
self.assertTrue(self.repository.delete("automation.dry_run"))
self.assertFalse(self.repository.delete("automation.dry_run"))
self.assertIsNone(self.repository.get("automation.dry_run"))
def test_empty_setting_key_is_rejected(self) -> None:
with self.assertRaisesRegex(ValueError, "setting_key"):
self.repository.set(" ", True)
if __name__ == "__main__":
unittest.main()
+188
View File
@@ -0,0 +1,188 @@
"""任务数据模型和 TaskRepository 测试。"""
import tempfile
import unittest
from pathlib import Path
from src.db import open_database
from src.task_models import (
NewClaimedTask,
OutboxEventType,
OutboxStatus,
RunStatus,
TaskFilters,
TaskStatus,
TaskType,
)
from src.task_repository import DuplicateTaskError, TaskRepository
class TaskRepositoryTests(unittest.TestCase):
def setUp(self) -> None:
self._temporary_directory = tempfile.TemporaryDirectory()
self.db_path = Path(self._temporary_directory.name) / "client.db"
self.repository = TaskRepository(self.db_path)
def tearDown(self) -> None:
self._temporary_directory.cleanup()
@staticmethod
def _task(
remote_task_id: str,
task_type: TaskType = TaskType.COLLECT,
title: str = "测试商品",
goods_id: str = "10001",
) -> NewClaimedTask:
return NewClaimedTask(
remote_task_id=remote_task_id,
task_type=task_type,
goods_id=goods_id,
goods_url=f"https://example.test/goods/{goods_id}",
title=title,
target_color="黑色" if task_type == TaskType.PURCHASE else None,
target_size="L" if task_type == TaskType.PURCHASE else None,
price_cent=3990,
quantity=2 if task_type == TaskType.PURCHASE else None,
admin_payload={"schema_version": 1, "task_id": remote_task_id},
)
def test_enum_values_match_database_constraints(self) -> None:
self.assertEqual({item.value for item in TaskType}, {"collect", "purchase"})
self.assertEqual(
{item.value for item in TaskStatus},
{
"claimed",
"running",
"result_pending",
"retry_wait",
"manual_review",
"succeeded",
"failed",
"cancelled",
},
)
self.assertEqual(
{item.value for item in RunStatus},
{"running", "succeeded", "failed", "cancelled", "manual_review"},
)
self.assertEqual(
{item.value for item in OutboxEventType},
{"collect_result", "purchase_result", "task_failure"},
)
self.assertEqual(
{item.value for item in OutboxStatus},
{"pending", "sending", "sent", "failed"},
)
def test_add_claimed_task_and_read_detail(self) -> None:
task_id = self.repository.add_claimed_task(
self._task("TASK-001"), "2026-08-06T08:00:00Z"
)
connection = open_database(self.db_path)
try:
with connection:
connection.execute(
"UPDATE pdd_tasks SET pdd_data = ? WHERE id = ?",
('{"schema_version": 1, "goods": {"title": "测试商品"}}', task_id),
)
finally:
connection.close()
detail = self.repository.get_task("TASK-001")
self.assertIsNotNone(detail)
self.assertEqual(detail.id, task_id)
self.assertEqual(detail.status, TaskStatus.CLAIMED)
self.assertEqual(detail.task_type, TaskType.COLLECT)
self.assertEqual(detail.admin_payload["task_id"], "TASK-001")
self.assertEqual(detail.pdd_data["schema_version"], 1)
def test_duplicate_remote_task_id_does_not_overwrite(self) -> None:
self.repository.add_claimed_task(self._task("TASK-001"))
with self.assertRaises(DuplicateTaskError):
self.repository.add_claimed_task(
self._task("TASK-001", title="不应覆盖的新标题")
)
detail = self.repository.get_task("TASK-001")
self.assertEqual(detail.title, "测试商品")
def test_list_tasks_uses_stable_paging_order(self) -> None:
same_time = "2026-08-06T08:00:00Z"
self.repository.add_claimed_task(self._task("TASK-001"), same_time)
self.repository.add_claimed_task(self._task("TASK-002"), same_time)
self.repository.add_claimed_task(
self._task("TASK-003"), "2026-08-06T09:00:00Z"
)
first_page = self.repository.list_tasks(limit=2)
second_page = self.repository.list_tasks(limit=2, offset=2)
self.assertEqual(
[task.remote_task_id for task in first_page], ["TASK-003", "TASK-002"]
)
self.assertEqual(
[task.remote_task_id for task in second_page], ["TASK-001"]
)
self.assertFalse(hasattr(first_page[0], "pdd_data"))
self.assertFalse(hasattr(first_page[0], "admin_payload"))
def test_filters_use_and_relationship_and_count_matches(self) -> None:
self.repository.add_claimed_task(
self._task("COLLECT-BLACK", title="黑色短袖", goods_id="20001")
)
self.repository.add_claimed_task(
self._task(
"PURCHASE-BLACK",
task_type=TaskType.PURCHASE,
title="黑色长裙",
goods_id="20002",
)
)
self.repository.add_claimed_task(
self._task("COLLECT-WHITE", title="白色短袖", goods_id="20003")
)
connection = open_database(self.db_path)
try:
with connection:
connection.execute(
"UPDATE pdd_tasks SET status = 'running'"
" WHERE remote_task_id = 'COLLECT-BLACK'"
)
finally:
connection.close()
filters = TaskFilters(
task_type=TaskType.COLLECT,
status=TaskStatus.RUNNING,
keyword="黑色",
)
tasks = self.repository.list_tasks(filters)
self.assertEqual([task.remote_task_id for task in tasks], ["COLLECT-BLACK"])
self.assertEqual(self.repository.count_tasks(filters), 1)
def test_keyword_treats_percent_as_normal_text(self) -> None:
self.repository.add_claimed_task(self._task("TASK-100%", title="百分号"))
self.repository.add_claimed_task(self._task("TASK-OTHER", title="普通商品"))
tasks = self.repository.list_tasks(TaskFilters(keyword="100%"))
self.assertEqual([task.remote_task_id for task in tasks], ["TASK-100%"])
def test_invalid_model_and_page_parameters_are_rejected(self) -> None:
with self.assertRaisesRegex(ValueError, "price_cent"):
NewClaimedTask(
remote_task_id="TASK-001",
task_type=TaskType.COLLECT,
goods_url="https://example.test/goods",
price_cent=-1,
)
with self.assertRaisesRegex(ValueError, "limit"):
self.repository.list_tasks(limit=0)
with self.assertRaisesRegex(ValueError, "offset"):
self.repository.list_tasks(offset=-1)
if __name__ == "__main__":
unittest.main()