feat: stabilize multi-region runtime sync and worker orchestration

This commit is contained in:
root
2026-04-27 15:48:12 +08:00
parent 7cbde2aa78
commit 215a364891
137 changed files with 31931 additions and 1943 deletions

View File

@@ -0,0 +1,94 @@
from __future__ import annotations
import tempfile
import unittest
from pathlib import Path
from unittest.mock import patch
from app.services.import_worker_service import import_domains_from_path
class _FakeCursor:
def __init__(self) -> None:
self._fetchall_result = []
self._fetchone_result = None
self.inserted_domains: list[str] = []
self.inserted_detect_tasks: list[int] = []
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
return False
def execute(self, sql: str, params=None) -> None:
normalized_sql = " ".join(str(sql or "").split()).lower()
params = params or ()
if normalized_sql.startswith("select domain from domains where domain = any"):
self._fetchall_result = []
return
if normalized_sql.startswith("insert into domains"):
domain = params[0]
self.inserted_domains.append(domain)
self._fetchone_result = (len(self.inserted_domains),)
return
if normalized_sql.startswith("insert into detect_tasks"):
self.inserted_detect_tasks.append(int(params[0]))
return
raise AssertionError(f"unexpected sql: {sql}")
def fetchall(self):
return list(self._fetchall_result)
def fetchone(self):
return self._fetchone_result
class _FakeConnection:
def __init__(self) -> None:
self.cursor_instance = _FakeCursor()
self.committed = False
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
return False
def cursor(self):
return self.cursor_instance
def commit(self) -> None:
self.committed = True
class ImportWorkerServiceTests(unittest.TestCase):
@patch("app.services.import_worker_service.get_db")
def test_import_domains_from_path_skips_duplicate_domains_in_same_batch(self, mock_get_db) -> None:
fake_conn = _FakeConnection()
mock_get_db.return_value = fake_conn
with tempfile.TemporaryDirectory() as tmpdir:
path = Path(tmpdir) / "domains.txt"
path.write_text("a.com\na.com\nb.net\ninvalid-domain\n", encoding="utf-8")
result = import_domains_from_path(path, source_type=7)
self.assertEqual(["a.com", "b.net"], fake_conn.cursor_instance.inserted_domains)
self.assertEqual([1, 2], fake_conn.cursor_instance.inserted_detect_tasks)
self.assertTrue(fake_conn.committed)
self.assertEqual(
{
"total": 4,
"valid": 3,
"added": 2,
"exists": 1,
"invalid": 1,
"failed": 0,
},
result["stats"],
)
if __name__ == "__main__":
unittest.main()