import os import sys import unittest sys.path.insert(0, os.path.dirname(__file__)) from _helpers import TempDirMixin from app import db class DbTests(TempDirMixin, unittest.TestCase): def test_init_db_is_idempotent_and_sets_pragmas(self): with self.make_temp_dir() as temp_dir: db_path = os.path.join(temp_dir, "cmshopee.db") db.init_db(db_path) db.init_db(db_path) conn = db.connect(db_path) try: self.assertEqual(1, conn.execute("PRAGMA foreign_keys").fetchone()[0]) self.assertEqual("wal", conn.execute("PRAGMA journal_mode").fetchone()[0]) self.assertEqual(5000, conn.execute("PRAGMA busy_timeout").fetchone()[0]) tables = { row["name"] for row in conn.execute( "SELECT name FROM sqlite_master WHERE type = 'table'" ).fetchall() } self.assertTrue({"batches", "accounts", "tasks"}.issubset(tables)) finally: conn.close() self.assert_removed(temp_dir) def test_account_batch_task_lifecycle(self): with self.make_temp_dir() as temp_dir: db_path = os.path.join(temp_dir, "cmshopee.db") db.init_db(db_path) batch_id = db.create_batch(["input.xlsx"], note="导入", path=db_path) batch = db.get_batch(batch_id, path=db_path) self.assertEqual("导入", batch.note) self.assertEqual([os.path.abspath("input.xlsx")], batch.source_files) account = db.add_account( "shop", "alias", "shopee.tw", 9222, note="备注", path=db_path, ) self.assertEqual("alias", account.alias) self.assertEqual("alias_cdb6fdbe", account.slug) db.update_account("alias", path=db_path, debug_port=9333) self.assertEqual(9333, db.get_account_by_alias("alias", path=db_path).debug_port) self.assertEqual(1, len(db.list_accounts(path=db_path))) count = db.insert_tasks( batch_id, [ { "source_file": "input.xlsx", "source_file_abs": os.path.abspath("input.xlsx"), "source_sheet": "Sheet1", "source_row": 2, "account_name": "shop", "alias": "alias", "item_id": "51100639510", } ], path=db_path, ) self.assertEqual(1, count) task = db.list_tasks(batch_id=batch_id, alias="alias", path=db_path)[0] self.assertEqual("imported", task.stage) self.assertEqual("pending", task.status) db.mark_running(task.id, "collect", path=db_path) self.assertEqual("running", db.list_tasks(path=db_path)[0].status) db.mark_failed(task.id, "collect", "采集失败", path=db_path) failed = db.list_tasks(path=db_path)[0] self.assertEqual("imported", failed.stage) self.assertEqual("failed", failed.status) self.assertEqual(1, failed.collect_attempts) db.set_collected(task.id, "旧标题", "old.jpg", path=db_path) collected = db.list_tasks(path=db_path)[0] self.assertEqual("collected", collected.stage) self.assertEqual("success", collected.status) self.assertEqual("旧标题", collected.old_title) db.set_generated(task.id, "新标题", "new.jpg", path=db_path) generated = db.list_tasks(path=db_path)[0] self.assertEqual("generated", generated.stage) self.assertEqual("新标题", generated.new_title) db.set_applied(task.id, False, "按钮禁用", path=db_path) apply_failed = db.list_tasks(path=db_path)[0] self.assertEqual("generated", apply_failed.stage) self.assertEqual("failed", apply_failed.status) self.assertEqual(1, apply_failed.apply_attempts) db.set_applied(task.id, True, path=db_path) applied = db.list_tasks(path=db_path)[0] self.assertEqual("applied", applied.stage) self.assertEqual("success", applied.status) self.assertEqual(1, applied.committed) self.assertEqual(2, applied.apply_attempts) self.assert_removed(temp_dir) def test_reset_generated_and_apply_status_keep_local_history(self): with self.make_temp_dir() as temp_dir: db_path = os.path.join(temp_dir, "cmshopee.db") db.init_db(db_path) batch_id = db.create_batch(["input.xlsx"], path=db_path) db.insert_tasks( batch_id, [ { "source_file_abs": os.path.join(temp_dir, "input.xlsx"), "source_sheet": "Sheet1", "source_row": 2, "account_name": "shop", "alias": "alias", "item_id": "51100639510", } ], path=db_path, ) task = db.list_tasks(batch_id=batch_id, path=db_path)[0] new_cover = os.path.join(temp_dir, "new.jpg") with open(new_cover, "wb") as fh: fh.write(b"jpeg") db.set_collected(task.id, "旧标题", "old.jpg", path=db_path) db.set_generated(task.id, "新标题", new_cover, path=db_path) db.set_applied(task.id, True, path=db_path) reset_apply = db.reset_apply_status(task.id, path=db_path) self.assertEqual("applied", reset_apply["before"].stage) after_apply = reset_apply["after"] self.assertEqual("generated", after_apply.stage) self.assertEqual("pending", after_apply.status) self.assertEqual("新标题", after_apply.new_title) self.assertEqual(new_cover, after_apply.new_cover_path) self.assertEqual(1, after_apply.committed) reset_generated = db.reset_generated(task.id, path=db_path) after_generated = reset_generated["after"] self.assertEqual(new_cover, reset_generated["new_cover_path"]) self.assertIsNone(reset_generated["deleted_file"]) self.assertTrue(os.path.exists(new_cover)) self.assertEqual("collected", after_generated.stage) self.assertEqual("success", after_generated.status) self.assertIsNone(after_generated.new_title) self.assertIsNone(after_generated.new_cover_path) self.assertIsNone(after_generated.last_error) self.assertEqual(1, after_generated.committed) self.assert_removed(temp_dir) def test_duplicate_task_and_invalid_update_raise_clear_errors(self): with self.make_temp_dir() as temp_dir: db_path = os.path.join(temp_dir, "cmshopee.db") db.init_db(db_path) batch_id = db.create_batch(["input.xlsx"], path=db_path) row = { "source_file": "input.xlsx", "source_file_abs": os.path.abspath("input.xlsx"), "source_sheet": "Sheet1", "source_row": 2, "alias": "alias", "item_id": "51100639510", } db.insert_tasks(batch_id, [row], path=db_path) with self.assertRaises(db.DbError): db.insert_tasks(batch_id, [row], path=db_path) with self.assertRaises(db.DbError): db.update_account("alias", path=db_path, unknown_field=True) self.assert_removed(temp_dir) def test_run_logs_and_events_are_persisted(self): with self.make_temp_dir() as temp_dir: db_path = os.path.join(temp_dir, "cmshopee.db") db.init_db(db_path) run_id = db.create_run_log( "apply", dry_run=True, total=2, options={"api_key": "secret", "mode": "preview"}, path=db_path, ) db.add_run_log_event( run_id, "dry-run 预览任务", task_id=7, alias="alias", item_id="51100639510", path=db_path, ) db.finish_run_log( run_id, status="done", done=2, success_count=1, skipped_count=1, failed_count=0, summary_json={"password": "secret", "done": 2}, path=db_path, ) run = db.list_run_logs(path=db_path)[0] self.assertEqual(run_id, run.id) self.assertEqual("apply", run.run_type) self.assertEqual(1, run.dry_run) self.assertEqual("done", run.status) self.assertEqual(2, run.done) self.assertEqual("***", run.options["api_key"]) self.assertEqual("***", run.summary["password"]) events = db.list_run_log_events(run_id, path=db_path) self.assertEqual(1, len(events)) self.assertEqual("alias", events[0].alias) self.assertEqual("51100639510", events[0].item_id) self.assertIn("dry-run", events[0].message) self.assert_removed(temp_dir) if __name__ == "__main__": unittest.main()