from __future__ import annotations import unittest from unittest.mock import patch from app.services import domains_service class _FakeDomainsCursor: def __init__(self) -> None: self._fetchone_result = None self._fetchall_result = [] self.executed: list[tuple[str, tuple]] = [] self.updated_detection_params = None 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() tuple_params = tuple(params or ()) self.executed.append((normalized_sql, tuple_params)) if normalized_sql.startswith("select count(*)"): self._fetchone_result = (1,) return if normalized_sql.startswith("select d.id, d.domain"): self._fetchall_result = [ ( 1, "a.com", 0, 0, 0, 0, 1, "", None, "", 0, None, 7, False, None, "", {"status": True, "state": "passed"}, {"status": True, "state": "passed"}, True, {"status": True, "state": "passed"}, {"status": True, "state": "passed"}, {"status": True, "state": "passed"}, {"status": True, "state": "passed"}, {"status": True, "state": "passed"}, {"status": False, "state": "failed", "message": "juziseo failed"}, {"status": False, "state": "blacklisted", "message": "jucha blacklisted"}, ) ] return if normalized_sql.startswith("update domains set"): return if normalized_sql.startswith("select id, baidu_history"): self._fetchone_result = ( 7, {"status": False, "state": "failed", "message": "timeout", "checked_at": "2026-04-20 12:00:00", "step": "baidu_site"}, {"status": True, "state": "passed", "message": "ok", "checked_at": "2026-04-20 12:00:00", "step": "baidu_site"}, False, {"status": False, "state": "failed", "message": "old", "checked_at": "2026-04-20 12:00:00", "step": "qihu360_site"}, {"status": False, "state": "failed", "message": "old", "checked_at": "2026-04-20 12:00:00", "step": "google_site"}, False, ) return if normalized_sql.startswith("update domain_detections set"): self.updated_detection_params = tuple_params return raise AssertionError(f"unexpected sql: {sql}") def fetchone(self): return self._fetchone_result def fetchall(self): return list(self._fetchall_result) class _FakeDomainsConnection: def __init__(self) -> None: self.cursor_instance = _FakeDomainsCursor() 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 DomainsServiceTests(unittest.TestCase): def test_build_domain_query_parts_supports_false_backlink_filter(self) -> None: _from_clause, where_clause, params = domains_service._build_domain_query_parts({"backlink_gt_10": False}) self.assertIn("coalesce(dd.backlink_count_gt_10, false) = %s", where_clause) self.assertEqual([False], params) @patch("app.services.domains_service.get_db") def test_fetch_domains_step_summary_counts_juziseo_and_jucha_results(self, mock_get_db) -> None: fake_conn = _FakeDomainsConnection() mock_get_db.return_value = fake_conn result = domains_service.fetch_domains(page=1, page_size=20) summary = result["list"][0]["step_summary"] self.assertEqual(1, summary["failed_count"]) self.assertEqual(1, summary["blacklisted_count"]) self.assertTrue(summary["has_failed"]) self.assertTrue(summary["has_blacklisted_step"]) @patch("app.services.domains_service.get_db") def test_batch_update_domains_preserves_detection_metadata_shape(self, mock_get_db) -> None: fake_conn = _FakeDomainsConnection() mock_get_db.return_value = fake_conn with patch("app.services.domains_service._now_text", return_value="2026-04-23 15:30:00"): result = domains_service.batch_update_domains([42], {"baidu_site": "否"}) self.assertEqual(1, result["updated_count"]) self.assertIsNotNone(fake_conn.cursor_instance.updated_detection_params) updated_payload = fake_conn.cursor_instance.updated_detection_params[0] self.assertEqual(False, updated_payload["status"]) self.assertEqual("failed", updated_payload["state"]) self.assertEqual("人工批量更新", updated_payload["message"]) self.assertEqual("2026-04-23 15:30:00", updated_payload["checked_at"]) self.assertTrue(updated_payload["manual_override"]) if __name__ == "__main__": unittest.main()