95 lines
2.8 KiB
Python
95 lines
2.8 KiB
Python
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()
|