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()