diff --git a/client/src/settings_repository.py b/client/src/settings_repository.py new file mode 100644 index 0000000..25ba768 --- /dev/null +++ b/client/src/settings_repository.py @@ -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 diff --git a/client/src/task_models.py b/client/src/task_models.py new file mode 100644 index 0000000..f7ed055 --- /dev/null +++ b/client/src/task_models.py @@ -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 diff --git a/client/src/task_repository.py b/client/src/task_repository.py new file mode 100644 index 0000000..6c469db --- /dev/null +++ b/client/src/task_repository.py @@ -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 diff --git a/client/test/test_settings_repository.py b/client/test/test_settings_repository.py new file mode 100644 index 0000000..cc0295f --- /dev/null +++ b/client/test/test_settings_repository.py @@ -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() diff --git a/client/test/test_task_repository.py b/client/test/test_task_repository.py new file mode 100644 index 0000000..b7bc263 --- /dev/null +++ b/client/test/test_task_repository.py @@ -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()