feat: stabilize multi-region runtime sync and worker orchestration
This commit is contained in:
94
domain-api/tests/test_import_worker_service.py
Normal file
94
domain-api/tests/test_import_worker_service.py
Normal 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()
|
||||
Reference in New Issue
Block a user