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

@@ -1,5 +1,6 @@
import unittest
import sys
import unittest
from datetime import datetime, timedelta
sys.path.insert(0, "/www/wwwroot/getDomain/domainCheck")
@@ -7,6 +8,7 @@ sys.path.insert(0, "/www/wwwroot/getDomain/domainCheck")
from app.utils.database import _detect_job_item_step_priority
from app.utils.database import _ordered_step_claim_codes
from app.utils.database import _resolve_step_claim_quota
from app.utils.database import _select_preferred_claim_job_ids
class DetectJobItemClaimPriorityTestCase(unittest.TestCase):
@@ -46,9 +48,42 @@ class DetectJobItemClaimPriorityTestCase(unittest.TestCase):
_detect_job_item_step_priority("detect_register"),
)
def test_step_claim_quota_defaults_to_quarter_window_with_floor(self):
def test_step_claim_quota_defaults_to_full_window_for_large_pools(self):
self.assertEqual(64, _resolve_step_claim_quota(120))
self.assertEqual(300, _resolve_step_claim_quota(1200))
self.assertEqual(1200, _resolve_step_claim_quota(1200))
def test_select_preferred_claim_job_ids_prioritizes_running_jobs_even_if_older(self):
now = datetime(2026, 4, 24, 18, 30, 0)
selected = _select_preferred_claim_job_ids(
[
(19, "pending", now - timedelta(hours=60)),
(57, "running", now - timedelta(hours=72)),
(58, "pending", now - timedelta(hours=2)),
],
limit=3,
recent_hours=24,
now=now,
)
self.assertEqual([57, 58], selected)
def test_select_preferred_claim_job_ids_filters_out_stale_pending_jobs(self):
now = datetime(2026, 4, 24, 18, 30, 0)
selected = _select_preferred_claim_job_ids(
[
(19, "pending", now - timedelta(hours=60)),
(20, "pending", now - timedelta(hours=40)),
(58, "pending", now - timedelta(hours=2)),
(59, "pending", now - timedelta(minutes=30)),
],
limit=4,
recent_hours=24,
now=now,
)
self.assertEqual([59, 58], selected)
if __name__ == "__main__":

View File

@@ -0,0 +1,440 @@
import threading
import time
import sys
import unittest
from unittest.mock import MagicMock, patch
sys.path.insert(0, "/www/wwwroot/getDomain/domainCheck")
from psycopg2 import extensions # noqa: E402
from app.utils.database import Database, _build_claim_token # noqa: E402
class _FakeCursor:
def __init__(self):
self.closed = False
self.executed = []
self.fetchone_queue = [(1,)]
self.fetchall_queue = []
def execute(self, sql, params=None):
self.executed.append((sql, params))
def fetchone(self):
if self.fetchone_queue:
return self.fetchone_queue.pop(0)
return (1,)
def fetchall(self):
if self.fetchall_queue:
return self.fetchall_queue.pop(0)
return []
def close(self):
self.closed = True
class _HealthyConn:
def __init__(self):
self.closed = False
self.cursor_calls = 0
def get_transaction_status(self):
return extensions.TRANSACTION_STATUS_IDLE
def rollback(self):
return None
def cursor(self):
self.cursor_calls += 1
return _FakeCursor()
def close(self):
self.closed = True
class _BrokenRollbackConn:
def __init__(self):
self.closed = False
def get_transaction_status(self):
return extensions.TRANSACTION_STATUS_INTRANS
def rollback(self):
raise RuntimeError("rollback failed")
def close(self):
self.closed = True
class DatabaseConnectionPoolTests(unittest.TestCase):
def _build_db(self):
db = Database.__new__(Database)
db.connection_pool = []
db.pool_size = 2
db.pool_idle_keep_max = 1
db.pool_healthcheck_interval = 30.0
db.pool_acquire_timeout = 1.0
db.pool_lock = threading.Lock()
db.pool_condition = threading.Condition(db.pool_lock)
db._pool_initialized = True
db.total_connections = 0
db._connection_last_healthcheck = {}
db.close = Database.close.__get__(db, Database)
db.connect = Database.connect.__get__(db, Database)
db._prepare_pooled_connection = Database._prepare_pooled_connection.__get__(db, Database)
db._discard_connection = Database._discard_connection.__get__(db, Database)
return db
def test_claim_token_is_unique_even_for_same_thread_and_instant(self):
token_a = _build_claim_token("mainland-controller-01", thread_id=123)
token_b = _build_claim_token("mainland-controller-01", thread_id=123)
self.assertNotEqual(token_a, token_b)
self.assertLessEqual(len(token_a), 64)
self.assertLessEqual(len(token_b), 64)
self.assertTrue(token_a.startswith("mainland-controller-01-7b-"))
self.assertTrue(token_b.startswith("mainland-controller-01-7b-"))
def test_connect_discards_bad_pooled_connection_and_creates_new_one(self):
db = self._build_db()
bad_conn = _BrokenRollbackConn()
good_conn = _HealthyConn()
db.connection_pool = [bad_conn]
db.total_connections = 1
db._connection_last_healthcheck[id(bad_conn)] = time.monotonic()
db._create_connection = MagicMock(return_value=good_conn)
conn, cur = db.connect(thread_id=123)
self.assertIs(conn, good_conn)
self.assertIsInstance(cur, _FakeCursor)
self.assertTrue(bad_conn.closed)
self.assertEqual(1, db.total_connections)
db._create_connection.assert_called_once()
def test_close_discards_connection_when_rollback_fails(self):
db = self._build_db()
conn = _BrokenRollbackConn()
db.total_connections = 1
db._connection_last_healthcheck[id(conn)] = time.monotonic()
db.close(conn, None)
self.assertTrue(conn.closed)
self.assertEqual(0, db.total_connections)
self.assertEqual([], db.connection_pool)
def test_ensure_cluster_runtime_tables_skips_ddl_when_schema_is_ready(self):
db = Database.__new__(Database)
db._cluster_runtime_schema_ready = MagicMock(return_value=True)
db.execute = MagicMock()
db.ensure_cluster_runtime_tables = Database.ensure_cluster_runtime_tables.__get__(db, Database)
self.assertTrue(db.ensure_cluster_runtime_tables())
db.execute.assert_not_called()
def test_ensure_cluster_runtime_tables_executes_ddl_when_schema_is_missing(self):
db = Database.__new__(Database)
db._cluster_runtime_schema_ready = MagicMock(return_value=False)
db._cluster_runtime_schema_basics_present = MagicMock(return_value=False)
db.execute = MagicMock(return_value=True)
db.ensure_cluster_runtime_tables = Database.ensure_cluster_runtime_tables.__get__(db, Database)
self.assertTrue(db.ensure_cluster_runtime_tables())
db.execute.assert_called_once()
def test_ensure_cluster_runtime_tables_skips_index_repair_without_explicit_enable(self):
db = Database.__new__(Database)
db._cluster_runtime_schema_ready = MagicMock(return_value=False)
db._cluster_runtime_schema_basics_present = MagicMock(return_value=True)
db._ensure_cluster_runtime_indexes = MagicMock(return_value=True)
db.execute = MagicMock()
db.ensure_cluster_runtime_tables = Database.ensure_cluster_runtime_tables.__get__(db, Database)
db._runtime_index_repair_enabled = Database._runtime_index_repair_enabled.__get__(db, Database)
with patch.dict("os.environ", {}, clear=False):
self.assertFalse(db.ensure_cluster_runtime_tables())
db._ensure_cluster_runtime_indexes.assert_not_called()
db.execute.assert_not_called()
def test_cluster_runtime_schema_ready_requires_claim_step_indexes(self):
db = Database.__new__(Database)
db._cluster_runtime_schema_basics_present = MagicMock(return_value=True)
db._cluster_runtime_missing_indexes = MagicMock(return_value=["idx_detect_job_items_claim_step_ready"])
db._cluster_runtime_schema_ready = Database._cluster_runtime_schema_ready.__get__(db, Database)
self.assertFalse(db._cluster_runtime_schema_ready())
def test_cluster_runtime_missing_indexes_treats_invalid_indexes_as_missing(self):
db = Database.__new__(Database)
db.fetch_all = MagicMock(
return_value=[
{"index_name": "idx_detect_job_items_job_domain_step", "is_valid": True},
{"index_name": "idx_detect_job_items_claim_step_ready", "is_valid": False},
]
)
db._cluster_runtime_missing_indexes = Database._cluster_runtime_missing_indexes.__get__(db, Database)
missing = db._cluster_runtime_missing_indexes()
self.assertIn("idx_detect_job_items_claim_step_ready", missing)
self.assertIn("idx_detect_job_items_claim_job_step_ready", missing)
self.assertNotIn("idx_detect_job_items_job_domain_step", missing)
def test_ensure_cluster_runtime_indexes_skips_invalid_rebuild_by_default(self):
db = Database.__new__(Database)
conn = MagicMock()
cur = _FakeCursor()
cur.fetchone_queue = [(True,)]
db._cluster_runtime_missing_indexes = MagicMock(return_value=["idx_detect_job_items_stalled_job_activity"])
db._cluster_runtime_invalid_indexes = MagicMock(return_value=["idx_detect_job_items_stalled_job_activity"])
db._create_connection = MagicMock(return_value=conn)
conn.cursor.return_value = cur
db._ensure_cluster_runtime_indexes = Database._ensure_cluster_runtime_indexes.__get__(db, Database)
db._runtime_index_repair_enabled = Database._runtime_index_repair_enabled.__get__(db, Database)
with patch.dict("os.environ", {}, clear=False):
self.assertTrue(db._ensure_cluster_runtime_indexes())
executed_sql = [sql for sql, _ in cur.executed]
self.assertFalse(any("DROP INDEX CONCURRENTLY IF EXISTS idx_detect_job_items_stalled_job_activity" in str(sql) for sql in executed_sql))
self.assertFalse(any("CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_detect_job_items_stalled_job_activity" in str(sql) for sql in executed_sql))
def test_ensure_cluster_runtime_indexes_rebuilds_invalid_index_when_enabled(self):
db = Database.__new__(Database)
conn = MagicMock()
cur = _FakeCursor()
cur.fetchone_queue = [(True,)]
db._cluster_runtime_missing_indexes = MagicMock(return_value=["idx_detect_job_items_stalled_job_activity"])
db._cluster_runtime_invalid_indexes = MagicMock(return_value=["idx_detect_job_items_stalled_job_activity"])
db._create_connection = MagicMock(return_value=conn)
conn.cursor.return_value = cur
db._ensure_cluster_runtime_indexes = Database._ensure_cluster_runtime_indexes.__get__(db, Database)
db._runtime_index_repair_enabled = Database._runtime_index_repair_enabled.__get__(db, Database)
with patch.dict("os.environ", {"DOMAINCHECK_RUNTIME_INDEX_REPAIR_ENABLED": "1"}, clear=False):
self.assertTrue(db._ensure_cluster_runtime_indexes())
executed_sql = [str(sql) for sql, _ in cur.executed]
self.assertTrue(any("DROP INDEX CONCURRENTLY IF EXISTS idx_detect_job_items_stalled_job_activity" in sql for sql in executed_sql))
self.assertTrue(any("CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_detect_job_items_stalled_job_activity" in sql for sql in executed_sql))
def test_release_detect_job_items_is_enabled_by_default(self):
db = Database.__new__(Database)
conn = MagicMock()
cur = _FakeCursor()
cur.fetchone_queue = [(0, [])]
db.connect = MagicMock(return_value=(conn, cur))
db.close = MagicMock()
db.release_detect_job_items_for_node = Database.release_detect_job_items_for_node.__get__(db, Database)
self.assertEqual(0, db.release_detect_job_items_for_node("mainland-controller-01-a"))
db.connect.assert_called_once()
executed_sql = [sql for sql, _ in cur.executed]
self.assertTrue(any("FOR UPDATE SKIP LOCKED" in sql for sql in executed_sql))
def test_claim_restart_released_detect_job_items_targets_restart_release_rows(self):
db = Database.__new__(Database)
conn = MagicMock()
cur = _FakeCursor()
cur.fetchall_queue = [[
(
101,
880,
501,
"token-a",
"detect_register",
"domain_pipeline",
"sync-overseas-20249",
"detect_register",
None,
"example.com",
0,
0,
0,
0,
None,
0,
0,
)
]]
db.connect = MagicMock(return_value=(conn, cur))
db.close = MagicMock()
db.claim_restart_released_detect_job_items = Database.claim_restart_released_detect_job_items.__get__(db, Database)
rows = db.claim_restart_released_detect_job_items("mainland-controller-01-a", 880, limit=16, lease_seconds=900)
self.assertEqual(1, len(rows))
executed_sql = [sql for sql, _ in cur.executed]
self.assertTrue(any("released after worker restart" in sql for sql in executed_sql))
self.assertTrue(any("released before execution after worker restart" in sql for sql in executed_sql))
self.assertTrue(any("FOR UPDATE SKIP LOCKED" in sql for sql in executed_sql))
def test_release_detect_job_items_can_be_disabled_explicitly(self):
db = Database.__new__(Database)
db.connect = MagicMock()
db.release_detect_job_items_for_node = Database.release_detect_job_items_for_node.__get__(db, Database)
with patch.dict("os.environ", {"DOMAINCHECK_ENABLE_NODE_ITEM_RELEASE": "0"}, clear=False):
self.assertEqual(0, db.release_detect_job_items_for_node("mainland-controller-01-a"))
db.connect.assert_not_called()
def test_release_detect_job_items_for_node_job_targets_single_job(self):
db = Database.__new__(Database)
conn = MagicMock()
cur = _FakeCursor()
cur.fetchone_queue = [(4, [876])]
db.connect = MagicMock(return_value=(conn, cur))
db.close = MagicMock()
db._refresh_detect_job_status_with_cursor = MagicMock()
db.release_detect_job_items_for_node_job = (
Database.release_detect_job_items_for_node_job.__get__(db, Database)
)
self.assertEqual(4, db.release_detect_job_items_for_node_job("mainland-controller-01-ba", 876))
executed_sql = [sql for sql, _ in cur.executed]
self.assertTrue(any("WHERE claimed_by = %s" in sql and "AND job_id = %s" in sql for sql in executed_sql))
self.assertTrue(any("FOR UPDATE SKIP LOCKED" in sql for sql in executed_sql))
db._refresh_detect_job_status_with_cursor.assert_called_once_with(cur, 876)
def test_release_single_detect_job_item_skips_job_refresh(self):
db = Database.__new__(Database)
conn = MagicMock()
cur = _FakeCursor()
cur.fetchone_queue = [(11,)]
db.connect = MagicMock(return_value=(conn, cur))
db.close = MagicMock()
db._refresh_detect_job_status_with_cursor = MagicMock()
db.release_detect_job_item = Database.release_detect_job_item.__get__(db, Database)
self.assertTrue(db.release_detect_job_item(101, "token-a", reason="session_replaced:7"))
executed_sql = [sql for sql, _ in cur.executed]
self.assertTrue(any("UPDATE detect_job_items" in sql for sql in executed_sql))
db._refresh_detect_job_status_with_cursor.assert_not_called()
def test_release_detect_job_items_batch_skips_job_refresh(self):
db = Database.__new__(Database)
conn = MagicMock()
cur = _FakeCursor()
cur.fetchall_queue = [[(101,), (102,)]]
cur.mogrify = MagicMock(side_effect=lambda sql, params: str(tuple(params)).encode("utf-8"))
db.connect = MagicMock(return_value=(conn, cur))
db.close = MagicMock()
db._refresh_detect_job_status_with_cursor = MagicMock()
db.release_detect_job_items_batch = Database.release_detect_job_items_batch.__get__(db, Database)
released = db.release_detect_job_items_batch(
[
(101, "token-a", "session_replaced:7"),
(102, "token-b", "queued_before_start"),
]
)
self.assertEqual(2, released)
executed_sql = [sql for sql, _ in cur.executed]
self.assertTrue(any("UPDATE detect_job_items" in sql for sql in executed_sql))
db._refresh_detect_job_status_with_cursor.assert_not_called()
def test_recycle_expired_detect_job_items_skips_when_advisory_lock_is_busy(self):
db = Database.__new__(Database)
conn = MagicMock()
cur = _FakeCursor()
cur.fetchone_queue = [(False,)]
db.connect = MagicMock(return_value=(conn, cur))
db.close = MagicMock()
db.recycle_expired_detect_job_items = Database.recycle_expired_detect_job_items.__get__(db, Database)
self.assertEqual(0, db.recycle_expired_detect_job_items())
executed_sql = [sql for sql, _ in cur.executed]
self.assertIn("SELECT pg_try_advisory_lock(%s)", executed_sql)
self.assertNotIn(
"WITH recycled AS (",
" ".join(executed_sql),
)
def test_recycle_expired_detect_job_items_runs_under_advisory_lock(self):
db = Database.__new__(Database)
conn = MagicMock()
cur = _FakeCursor()
cur.fetchone_queue = [(True,), (3, [11, 12])]
db.connect = MagicMock(return_value=(conn, cur))
db.close = MagicMock()
db._refresh_detect_job_status_with_cursor = MagicMock()
db.recycle_expired_detect_job_items = Database.recycle_expired_detect_job_items.__get__(db, Database)
self.assertEqual(3, db.recycle_expired_detect_job_items())
executed_sql = [sql for sql, _ in cur.executed]
self.assertIn("SELECT pg_try_advisory_lock(%s)", executed_sql)
self.assertTrue(any("expired_candidates" in sql and "WITH" in sql for sql in executed_sql))
self.assertTrue(any("FOR UPDATE SKIP LOCKED" in sql for sql in executed_sql))
self.assertIn("SELECT pg_advisory_unlock(%s)", executed_sql)
db._refresh_detect_job_status_with_cursor.assert_any_call(cur, 11)
db._refresh_detect_job_status_with_cursor.assert_any_call(cur, 12)
def test_recycle_stalled_detect_job_items_skips_when_advisory_lock_is_busy(self):
db = Database.__new__(Database)
conn = MagicMock()
cur = _FakeCursor()
cur.fetchone_queue = [(False,)]
db.connect = MagicMock(return_value=(conn, cur))
db.close = MagicMock()
db.recycle_stalled_detect_job_items = Database.recycle_stalled_detect_job_items.__get__(db, Database)
self.assertEqual(0, db.recycle_stalled_detect_job_items(876, stall_seconds=1800))
executed_sql = [sql for sql, _ in cur.executed]
self.assertIn("SELECT pg_try_advisory_lock(%s)", executed_sql)
self.assertNotIn(
"WITH stalled_candidates AS (",
" ".join(executed_sql),
)
def test_recycle_stalled_detect_job_items_runs_under_advisory_lock(self):
db = Database.__new__(Database)
conn = MagicMock()
cur = _FakeCursor()
cur.fetchone_queue = [(True,), (5, [876])]
db.connect = MagicMock(return_value=(conn, cur))
db.close = MagicMock()
db._refresh_detect_job_status_with_cursor = MagicMock()
db.recycle_stalled_detect_job_items = Database.recycle_stalled_detect_job_items.__get__(db, Database)
self.assertEqual(5, db.recycle_stalled_detect_job_items(876, stall_seconds=1800, batch_size=64))
executed_sql = [sql for sql, _ in cur.executed]
self.assertIn("SELECT pg_try_advisory_lock(%s)", executed_sql)
self.assertTrue(any("stalled_candidates" in sql and "WITH" in sql for sql in executed_sql))
self.assertTrue(any("FOR UPDATE SKIP LOCKED" in sql for sql in executed_sql))
self.assertIn("SELECT pg_advisory_unlock(%s)", executed_sql)
db._refresh_detect_job_status_with_cursor.assert_called_once_with(cur, 876)
def test_refresh_detect_job_status_uses_exists_queries_for_final_status(self):
db = Database.__new__(Database)
cur = _FakeCursor()
cur.fetchone_queue = [
("domain_pipeline",),
(False,), # dispatch_active_exists
(False,), # unprocessed_terminal_exists
(False,), # pending_exists
(True,), # failed_exists
(True,), # done_exists
]
Database._refresh_detect_job_status_with_cursor(db, cur, 876)
executed_sql = [sql for sql, _ in cur.executed]
self.assertTrue(any("SELECT COALESCE(task_mode, '')" in sql for sql in executed_sql))
self.assertTrue(any("SELECT EXISTS" in sql and "item.status IN ('claimed', 'running')" in sql for sql in executed_sql))
self.assertTrue(any("SELECT EXISTS" in sql and "controller_processed" in sql for sql in executed_sql))
self.assertTrue(any("SELECT EXISTS" in sql and "item.status = 'pending'" in sql for sql in executed_sql))
self.assertTrue(any("SELECT EXISTS" in sql and "item.status = 'failed'" in sql for sql in executed_sql))
self.assertTrue(any("SELECT EXISTS" in sql and "item.status IN ('completed', 'blacklisted')" in sql for sql in executed_sql))
self.assertTrue(any("UPDATE detect_jobs" in sql for sql in executed_sql))
self.assertTrue(any(params == ('partial_failed', 876) for _, params in cur.executed))
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,200 @@
import sys
import unittest
from unittest.mock import patch
from unittest.mock import MagicMock
import json
import time
sys.path.insert(0, "/www/wwwroot/getDomain/domainCheck")
from app.utils.database import Database # noqa: E402
class _FakeConn:
def __init__(self):
self.committed = False
self.rolled_back = False
def commit(self):
self.committed = True
def rollback(self):
self.rolled_back = True
class _FakeCursor:
def __init__(self, fetchall_sequence):
self.fetchall_sequence = list(fetchall_sequence or [])
self.executed_sql = []
def mogrify(self, template, params):
rendered = []
for value in params:
if isinstance(value, str):
rendered.append(f"'{value}'")
else:
rendered.append(str(value))
return f"({', '.join(rendered)})".encode("utf-8")
def execute(self, sql, params=None):
self.executed_sql.append((sql, params))
def fetchall(self):
if self.fetchall_sequence:
return self.fetchall_sequence.pop(0)
return []
class DatabaseJobStatusRefreshTests(unittest.TestCase):
def test_mark_running_batch_skips_job_status_refresh(self):
db = Database.__new__(Database)
conn = _FakeConn()
cur = _FakeCursor(
fetchall_sequence=[
[(101,), (101,), (102,)],
]
)
db.connect = MagicMock(return_value=(conn, cur))
db.close = MagicMock()
db._refresh_detect_job_status_with_cursor = MagicMock()
updated = db.mark_detect_job_items_running_batch(
[(1, "token-a"), (2, "token-b"), (3, "token-c")]
)
self.assertEqual(3, updated)
self.assertTrue(conn.committed)
db._refresh_detect_job_status_with_cursor.assert_not_called()
def test_release_single_item_skips_job_status_refresh(self):
db = Database.__new__(Database)
conn = _FakeConn()
cur = _FakeCursor(fetchall_sequence=[[(101,)]] )
cur.fetchone = lambda: cur.fetchall_sequence.pop(0)[0] if cur.fetchall_sequence else None
db.connect = MagicMock(return_value=(conn, cur))
db.close = MagicMock()
db._refresh_detect_job_status_with_cursor = MagicMock()
released = db.release_detect_job_item(1, "token-a", reason="session replaced")
self.assertTrue(released)
self.assertTrue(conn.committed)
db._refresh_detect_job_status_with_cursor.assert_not_called()
def test_get_active_detect_job_uses_lightweight_candidate_query(self):
db = Database.__new__(Database)
db.fetch_one = MagicMock(return_value={"id": 9, "job_code": "sync-overseas-9"})
db.redis_client = None
with patch.dict(
"os.environ",
{
"DOMAINCHECK_TAIL_HANDOFF_ENABLED": "1",
"DOMAINCHECK_TAIL_HANDOFF_MAX_ACTIVE_ITEMS": "64",
"DOMAINCHECK_TAIL_HANDOFF_MAX_PENDING_ITEMS": "128",
"DOMAINCHECK_TAIL_HANDOFF_MIN_PENDING_ITEMS": "1",
"DOMAINCHECK_RUNNING_JOB_STALL_SECONDS": "1800",
},
clear=False,
):
result = db.get_active_detect_job()
self.assertEqual({"id": 9, "job_code": "sync-overseas-9"}, result)
sql, params = db.fetch_one.call_args[0]
self.assertIn("WITH tail_config AS", sql)
self.assertIn("candidate_jobs AS", sql)
self.assertIn("tail_handoff_candidate", sql)
self.assertIn("selection_reason", sql)
self.assertIn("LEFT JOIN LATERAL", sql)
self.assertIn("item.status IN ('pending', 'claimed', 'running')", sql)
self.assertIn("running_job_stalled", sql)
self.assertIn("latest_unfinished_activity_at", sql)
self.assertEqual((True, 64, 128, 1, 1800), params)
def test_get_active_detect_job_uses_cache_when_available(self):
db = Database.__new__(Database)
db.redis_client = MagicMock()
db.redis_client.get.return_value = json.dumps({"id": 12, "job_code": "sync-overseas-12"})
db.fetch_one = MagicMock()
result = db.get_active_detect_job()
self.assertEqual({"id": 12, "job_code": "sync-overseas-12"}, result)
db.fetch_one.assert_not_called()
def test_get_active_detect_job_writes_cache_after_query(self):
db = Database.__new__(Database)
db.redis_client = MagicMock()
db.redis_client.get.return_value = None
db.redis_client.set.return_value = True
db.fetch_one = MagicMock(return_value={"id": 13, "job_code": "sync-overseas-13", "items_pending": 5})
result = db.get_active_detect_job()
self.assertEqual({"id": 13, "job_code": "sync-overseas-13", "items_pending": 5}, result)
db.redis_client.setex.assert_called_once()
cache_key, ttl, payload = db.redis_client.setex.call_args[0]
self.assertEqual("domaincheck:active_detect_job_summary:v1", cache_key)
self.assertEqual(3, ttl)
self.assertEqual({"id": 13, "job_code": "sync-overseas-13", "items_pending": 5}, json.loads(payload))
def test_get_active_detect_job_uses_local_cache_after_first_query(self):
db = Database.__new__(Database)
db.redis_client = None
db.fetch_one = MagicMock(return_value={"id": 14, "job_code": "sync-overseas-14"})
first = db.get_active_detect_job()
second = db.get_active_detect_job()
self.assertEqual({"id": 14, "job_code": "sync-overseas-14"}, first)
self.assertEqual({"id": 14, "job_code": "sync-overseas-14"}, second)
db.fetch_one.assert_called_once()
def test_get_active_detect_job_uses_stale_local_cache_when_refresh_lock_is_busy(self):
db = Database.__new__(Database)
db.redis_client = MagicMock()
db.redis_client.get.return_value = None
db.redis_client.set.return_value = False
db.fetch_one = MagicMock()
now_ts = time.time()
db._set_local_active_detect_job_cache(
{"id": 15, "job_code": "sync-overseas-15"},
now_ts=now_ts - 4,
fresh_ttl_seconds=3,
stale_ttl_seconds=10,
)
with patch("time.sleep", return_value=None):
result = db.get_active_detect_job()
self.assertEqual({"id": 15, "job_code": "sync-overseas-15"}, result)
db.fetch_one.assert_not_called()
def test_finalize_batch_skips_job_status_refresh(self):
db = Database.__new__(Database)
conn = _FakeConn()
cur = _FakeCursor(
fetchall_sequence=[
[(1, 101), (2, 101), (3, 102)],
]
)
db.connect = MagicMock(return_value=(conn, cur))
db.close = MagicMock()
db._refresh_detect_job_status_with_cursor = MagicMock()
updated = db.finalize_detect_job_items_batch(
[
{"job_item_id": 1, "claim_token": "token-a", "final_status": "completed"},
{"job_item_id": 2, "claim_token": "token-b", "final_status": "failed"},
{"job_item_id": 3, "claim_token": "token-c", "final_status": "completed"},
]
)
self.assertEqual(3, updated)
self.assertTrue(conn.committed)
db._refresh_detect_job_status_with_cursor.assert_not_called()
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,171 @@
import json
import os
import sys
import tempfile
import unittest
from pathlib import Path
from unittest.mock import patch
sys.path.insert(0, "/www/wwwroot/getDomain/domainCheck")
from app.config import config # noqa: E402
from app.utils.database import _local_config_search_roots, _resolve_db_pool_limits # noqa: E402
class DatabasePoolLimitTests(unittest.TestCase):
def _write_json(self, root: str, name: str, payload: dict) -> None:
Path(root, name).write_text(json.dumps(payload, ensure_ascii=False), encoding="utf-8")
def test_explicit_total_budget_override_keeps_previous_pool_shape(self):
with tempfile.TemporaryDirectory() as temp_dir, patch.object(
config, "NODE_CODE", "mainland-controller-01-a"
), patch.dict(
os.environ,
{
"WORKER_PARENT_NODE_CODE": "mainland-controller-01",
"DOMAINCHECK_DB_POOL_TOTAL_BUDGET": "480",
"DB_POOL_SIZE": "",
"DB_POOL_WARM_SIZE": "",
"DB_POOL_IDLE_KEEP_MAX": "",
},
clear=False,
):
self._write_json(temp_dir, "thread_count.json", {"thread_count": "1000"})
self._write_json(temp_dir, "process_count.json", {"process_count": "80"})
self._write_json(temp_dir, "node_thread_counts.json", {"mainland-controller-01": 1200})
self._write_json(temp_dir, "node_process_counts.json", {"mainland-controller-01": 80})
previous_cwd = os.getcwd()
try:
os.chdir(temp_dir)
limits = _resolve_db_pool_limits()
finally:
os.chdir(previous_cwd)
self.assertEqual(6, limits["pool_size"])
self.assertEqual(1, limits["pool_warm_size"])
self.assertEqual(3, limits["pool_idle_keep_max"])
self.assertEqual(1200, limits["scaling_hints"]["thread_count"])
self.assertEqual(80, limits["scaling_hints"]["process_count"])
def test_pool_budget_scales_down_for_mid_sized_multi_process_workers(self):
with tempfile.TemporaryDirectory() as temp_dir, patch.object(
config, "NODE_CODE", "mainland-worker-01-a"
), patch.dict(
os.environ,
{
"WORKER_PARENT_NODE_CODE": "mainland-worker-01",
"DB_POOL_SIZE": "",
"DB_POOL_WARM_SIZE": "",
"DB_POOL_IDLE_KEEP_MAX": "",
"DOMAINCHECK_DB_POOL_TOTAL_BUDGET": "",
},
clear=False,
):
self._write_json(temp_dir, "thread_count.json", {"thread_count": "1000"})
self._write_json(temp_dir, "process_count.json", {"process_count": "60"})
self._write_json(temp_dir, "node_thread_counts.json", {"mainland-worker-01": 1200})
self._write_json(temp_dir, "node_process_counts.json", {"mainland-worker-01": 60})
previous_cwd = os.getcwd()
try:
os.chdir(temp_dir)
limits = _resolve_db_pool_limits()
finally:
os.chdir(previous_cwd)
self.assertEqual(4, limits["pool_size"])
self.assertEqual(1, limits["pool_warm_size"])
self.assertEqual(2, limits["pool_idle_keep_max"])
self.assertEqual(60, limits["scaling_hints"]["process_count"])
def test_explicit_pool_env_overrides_win_over_auto_scaling(self):
with tempfile.TemporaryDirectory() as temp_dir, patch.object(
config, "NODE_CODE", "mainland-controller-01"
), patch.dict(
os.environ,
{
"DB_POOL_SIZE": "16",
"DB_POOL_WARM_SIZE": "4",
"DB_POOL_IDLE_KEEP_MAX": "8",
},
clear=False,
):
self._write_json(temp_dir, "thread_count.json", {"thread_count": "1000"})
self._write_json(temp_dir, "process_count.json", {"process_count": "80"})
previous_cwd = os.getcwd()
try:
os.chdir(temp_dir)
limits = _resolve_db_pool_limits()
finally:
os.chdir(previous_cwd)
self.assertEqual(16, limits["pool_size"])
self.assertEqual(4, limits["pool_warm_size"])
self.assertEqual(8, limits["pool_idle_keep_max"])
def test_pool_limits_can_resolve_parent_process_override_from_absolute_config_root(self):
with tempfile.TemporaryDirectory() as temp_dir, patch.object(
config, "NODE_CODE", "mainland-controller-01-u"
), patch.dict(
os.environ,
{
"WORKER_PARENT_NODE_CODE": "mainland-controller-01",
"DB_POOL_SIZE": "",
"DB_POOL_WARM_SIZE": "",
"DB_POOL_IDLE_KEEP_MAX": "",
"DOMAINCHECK_DB_POOL_TOTAL_BUDGET": "",
},
clear=False,
):
self._write_json(temp_dir, "thread_count.json", {"thread_count": "1000"})
self._write_json(temp_dir, "process_count.json", {"process_count": "80"})
self._write_json(temp_dir, "node_thread_counts.json", {"mainland-controller-01": 1000})
self._write_json(temp_dir, "node_process_counts.json", {"mainland-controller-01": 80})
with patch(
"app.utils.database._local_config_search_roots",
return_value=[Path("/nonexistent/domainCheck"), Path(temp_dir)],
):
limits = _resolve_db_pool_limits()
self.assertEqual(80, limits["scaling_hints"]["process_count"])
self.assertEqual("mainland-controller-01", limits["scaling_hints"]["parent_node_code"])
self.assertEqual(9, limits["pool_size"])
def test_controller_high_process_pool_keeps_ten_connections_at_hundred_processes(self):
with tempfile.TemporaryDirectory() as temp_dir, patch.object(
config, "NODE_CODE", "mainland-controller-01-aa"
), patch.dict(
os.environ,
{
"WORKER_PARENT_NODE_CODE": "mainland-controller-01",
"DB_POOL_SIZE": "",
"DB_POOL_WARM_SIZE": "",
"DB_POOL_IDLE_KEEP_MAX": "",
"DOMAINCHECK_DB_POOL_TOTAL_BUDGET": "",
},
clear=False,
):
self._write_json(temp_dir, "thread_count.json", {"thread_count": "1000"})
self._write_json(temp_dir, "process_count.json", {"process_count": "100"})
self._write_json(temp_dir, "node_thread_counts.json", {"mainland-controller-01": 1000})
self._write_json(temp_dir, "node_process_counts.json", {"mainland-controller-01": 100})
previous_cwd = os.getcwd()
try:
os.chdir(temp_dir)
limits = _resolve_db_pool_limits()
finally:
os.chdir(previous_cwd)
self.assertEqual(100, limits["scaling_hints"]["process_count"])
self.assertEqual(10, limits["pool_size"])
self.assertEqual(2, limits["pool_warm_size"])
self.assertEqual(5, limits["pool_idle_keep_max"])
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,170 @@
import sys
import unittest
from pathlib import Path
from unittest.mock import MagicMock, patch
PROJECT_ROOT = Path(__file__).resolve().parents[2]
DOMAINCHECK_ROOT = PROJECT_ROOT / "domainCheck"
if str(DOMAINCHECK_ROOT) not in sys.path:
sys.path.insert(0, str(DOMAINCHECK_ROOT))
from app.core.detect_engine import DetectEngine
from app.detectors.rdap_detector import RDAPDetector
class DetectEngineOutcomeTests(unittest.TestCase):
def _build_engine(self):
patchers = [
patch("app.core.detect_engine.Database"),
patch("app.core.detect_engine.RDAPDetector"),
patch("app.core.detect_engine.WaybackDetector"),
patch("app.core.detect_engine.BaiduDetector"),
patch("app.core.detect_engine.Qihu360Detector"),
patch("app.core.detect_engine.GoogleDetector"),
patch("app.core.detect_engine.ChinazDetector"),
patch("app.core.detect_engine.AizhanDetector"),
patch("app.core.detect_engine.JuziseoDetector"),
patch("app.core.detect_engine.JuchaDetector"),
]
started = [patcher.start() for patcher in patchers]
self.addCleanup(lambda: [patcher.stop() for patcher in reversed(patchers)])
engine = DetectEngine()
for started_mock in started:
started_mock.return_value = MagicMock()
return engine
def test_process_task_blacklisted_is_completed_without_retry(self):
engine = self._build_engine()
engine.db.get_task_by_id.return_value = {"id": 7, "domain_id": 42, "retry_count": 1}
engine._detect_domain_with_outcome = MagicMock(return_value=engine.OUTCOME_BLACKLISTED)
success = engine.process_task(7)
self.assertTrue(success)
self.assertEqual(
[(7, 1), (7, 2)],
[call.args for call in engine.db.update_task_status.call_args_list],
)
engine.db.update_task_retry_count.assert_not_called()
def test_process_task_failed_requeues_for_retry(self):
engine = self._build_engine()
engine.db.get_task_by_id.return_value = {"id": 9, "domain_id": 99, "retry_count": 0}
engine._detect_domain_with_outcome = MagicMock(return_value=engine.OUTCOME_FAILED)
success = engine.process_task(9)
self.assertFalse(success)
self.assertEqual(
[(9, 1), (9, 0)],
[call.args for call in engine.db.update_task_status.call_args_list],
)
engine.db.update_task_retry_count.assert_called_once_with(9, 1)
def test_deep_detect_fails_when_any_detector_returns_error(self):
engine = self._build_engine()
engine.baidu_detector.check_history.return_value = {"error": "timeout"}
engine.baidu_detector.check_site.return_value = {"has_收录": False, "subdomains": []}
engine.qihu360_detector.check_site.return_value = {"has_收录": False, "subdomains": []}
engine.google_detector.check_site.return_value = {"has_收录": False}
engine.chinaz_detector.check_domain.return_value = {"title": "", "category": "", "has_sensitive": False}
engine.aizhan_detector.check_domain.return_value = {"title": "", "risk": "", "has_sensitive": False}
engine.juziseo_detector.check_domain.return_value = {"history": {}, "backlink": {}}
engine.jucha_detector.check_domain.return_value = {"whois": {}, "beian": {}, "intercept": {"normal": True}}
engine.db.add_detection_result.return_value = True
outcome = engine._deep_detect(1, "example.com")
self.assertEqual(engine.OUTCOME_FAILED, outcome)
engine.db.add_detection_result.assert_not_called()
def test_deep_detect_fails_when_result_persistence_fails(self):
engine = self._build_engine()
engine.baidu_detector.check_history.return_value = {"has_history": False, "has_gray": False}
engine.baidu_detector.check_site.return_value = {"has_收录": False, "subdomains": []}
engine.qihu360_detector.check_site.return_value = {"has_收录": False, "subdomains": []}
engine.google_detector.check_site.return_value = {"has_收录": False}
engine.chinaz_detector.check_domain.return_value = {"title": "", "category": "", "has_sensitive": False}
engine.aizhan_detector.check_domain.return_value = {"title": "", "risk": "", "has_sensitive": False}
engine.juziseo_detector.check_domain.return_value = {"history": {}, "backlink": {}}
engine.jucha_detector.check_domain.return_value = {"whois": {}, "beian": {}, "intercept": {"normal": True}}
engine.db.add_detection_result.return_value = False
outcome = engine._deep_detect(1, "example.com")
self.assertEqual(engine.OUTCOME_FAILED, outcome)
engine.db.add_detection_result.assert_called_once()
def test_deep_detect_normalizes_results_before_persisting(self):
engine = self._build_engine()
engine.baidu_detector.check_history.return_value = {"has_history": True, "has_gray": False}
engine.baidu_detector.check_site.return_value = {"has_收录": True, "subdomains": ["www"]}
engine.qihu360_detector.check_site.return_value = {"has_收录": False, "subdomains": []}
engine.google_detector.check_site.return_value = {"has_收录": False}
engine.chinaz_detector.check_domain.return_value = {"title": "Example", "category": "", "has_sensitive": False}
engine.aizhan_detector.check_domain.return_value = {"title": "Example", "risk": "", "has_sensitive": False}
engine.juziseo_detector.check_domain.return_value = {
"history": {"has_sensitive": False, "has_baidu_history": True, "has_subdomains": False, "is_simplified": True},
"backlink": {"has_sensitive": False, "has_subdomains": False},
}
engine.jucha_detector.check_domain.return_value = {
"whois": {"status": ""},
"beian": {"has_beian": True, "beian_year": "2024", "is_enterprise": True, "beian_match": True},
"intercept": {"normal": True},
}
engine.db.add_detection_result.return_value = True
outcome = engine._deep_detect(1, "example.com")
self.assertEqual(engine.OUTCOME_SUCCESS, outcome)
persisted_args = engine.db.add_detection_result.call_args.args
self.assertTrue(persisted_args[1]["status"])
self.assertTrue(persisted_args[1]["has_history"])
self.assertTrue(persisted_args[2]["status"])
self.assertTrue(persisted_args[2]["has_收录"])
self.assertIn("state", persisted_args[7]["history"])
self.assertTrue(persisted_args[8]["beian"]["status"])
def test_deep_detect_stops_early_after_blacklist_hit(self):
engine = self._build_engine()
engine.baidu_detector.check_history.return_value = {"has_history": True, "has_gray": True}
engine.baidu_detector.check_site.return_value = {"has_收录": False, "subdomains": []}
engine.db.add_detection_result.return_value = True
outcome = engine._deep_detect(1, "example.com")
self.assertEqual(engine.OUTCOME_BLACKLISTED, outcome)
engine.qihu360_detector.check_site.assert_not_called()
engine.google_detector.check_site.assert_not_called()
engine.db.add_detection_result.assert_called_once()
def test_deep_detect_fails_when_nested_detector_returns_error(self):
engine = self._build_engine()
engine.baidu_detector.check_history.return_value = {"has_history": False, "has_gray": False}
engine.baidu_detector.check_site.return_value = {"has_收录": False, "subdomains": []}
engine.qihu360_detector.check_site.return_value = {"has_收录": False, "subdomains": []}
engine.google_detector.check_site.return_value = {"has_收录": False}
engine.chinaz_detector.check_domain.return_value = {"title": "", "category": "", "has_sensitive": False}
engine.aizhan_detector.check_domain.return_value = {"title": "", "risk": "", "has_sensitive": False}
engine.juziseo_detector.check_domain.return_value = {
"history": {"error": "HTTP 429"},
"backlink": {"has_sensitive": False, "has_subdomains": False},
}
engine.jucha_detector.check_domain.return_value = {"whois": {}, "beian": {}, "intercept": {"normal": True}}
outcome = engine._deep_detect(1, "example.com")
self.assertEqual(engine.OUTCOME_FAILED, outcome)
engine.db.add_detection_result.assert_not_called()
class RDAPDetectorStatusMappingTests(unittest.TestCase):
def test_check_register_status_uses_statuses_field(self):
detector = RDAPDetector()
with patch.object(detector, "check_domain", return_value={"statuses": ["clientHold"]}):
status = detector.check_register_status("example.com")
self.assertEqual(7, status)
if __name__ == "__main__":
unittest.main()

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,44 @@
import sys
import unittest
from pathlib import Path
PROJECT_ROOT = Path(__file__).resolve().parents[2]
DOMAINCHECK_ROOT = PROJECT_ROOT / "domainCheck"
if str(DOMAINCHECK_ROOT) not in sys.path:
sys.path.insert(0, str(DOMAINCHECK_ROOT))
from app.utils.detection_results import (
build_manual_detection_result,
normalize_detector_result,
resolve_detection_status,
)
class DetectionResultSchemaTests(unittest.TestCase):
def test_resolve_detection_status_supports_legacy_key(self):
self.assertTrue(resolve_detection_status({"has_收录": True}, "has_收录"))
self.assertFalse(resolve_detection_status({"has_history": False}, "has_history"))
def test_manual_detection_result_preserves_legacy_key(self):
result = build_manual_detection_result(True, legacy_key="has_收录")
self.assertTrue(result["status"])
self.assertTrue(result["has_收录"])
self.assertEqual("manual", result["state"])
def test_normalize_jucha_info_bubbles_nested_error(self):
normalized = normalize_detector_result(
"jucha_info",
{
"whois": {"error": "HTTP 403"},
"beian": {"has_beian": False, "beian_year": "", "is_enterprise": False, "beian_match": False},
"intercept": {"normal": True},
},
)
self.assertEqual("HTTP 403", normalized["error"])
self.assertEqual("error", normalized["state"])
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,82 @@
import sys
import tempfile
import unittest
from pathlib import Path
from unittest.mock import MagicMock, patch
PROJECT_ROOT = Path(__file__).resolve().parents[2]
DOMAINCHECK_ROOT = PROJECT_ROOT / "domainCheck"
if str(DOMAINCHECK_ROOT) not in sys.path:
sys.path.insert(0, str(DOMAINCHECK_ROOT))
import requests
from domainCheck.detect import geetest2, jucha, juming, juziseo
from domainCheck.detect.locked_pickle import load_pickle_locked, save_pickle_atomic
class GeetestCookieSafetyTests(unittest.TestCase):
def test_geetest_slide_asset_fetch_applies_timeout(self):
response = MagicMock()
response.content = b"binary"
with patch("domainCheck.detect.geetest2.requests.get", return_value=response) as mock_get:
slider = geetest2.slide()
with patch.object(slider, "tp_huanyuan", return_value=b"bg-bytes"):
with patch("domainCheck.detect.geetest2.quekou") as mock_quekou:
mock_quekou.return_value.get_distance.return_value = 12
slider.huak({"bg": "bg.png", "slice": "slice.png"})
self.assertGreaterEqual(mock_get.call_count, 2)
for call in mock_get.call_args_list:
self.assertEqual(slider.asset_timeout, call.kwargs["timeout"])
def test_locked_pickle_roundtrip_is_atomic(self):
with tempfile.TemporaryDirectory() as temp_dir:
target = Path(temp_dir) / "cookies.pkl"
save_pickle_atomic(str(target), {"sid": "abc"})
loaded = load_pickle_locked(str(target), default_factory=dict)
self.assertEqual({"sid": "abc"}, loaded)
def test_juziseo_cookie_roundtrip_uses_locked_pickle(self):
with tempfile.TemporaryDirectory() as temp_dir:
target = Path(temp_dir) / "juziseo.pkl"
detector = juziseo.Juziseo()
jar = requests.cookies.RequestsCookieJar()
jar.set("sid", "value")
detector.cookie = jar
detector.save_cookies(str(target))
loaded = juziseo.Juziseo()
loaded.load_cookies(str(target))
self.assertEqual("value", loaded.cookie.get("sid"))
def test_jucha_and_juming_cookie_roundtrip_use_locked_pickle(self):
with tempfile.TemporaryDirectory() as temp_dir:
jucha_path = Path(temp_dir) / "jucha.pkl"
juming_path = Path(temp_dir) / "juming.pkl"
jc = jucha.JC()
jc_jar = requests.cookies.RequestsCookieJar()
jc_jar.set("jc", "cookie")
jc.cookie = jc_jar
jc.save_cookies(str(jucha_path))
jm = juming.JM()
jm_jar = requests.cookies.RequestsCookieJar()
jm_jar.set("jm", "cookie")
jm.cookie = jm_jar
jm.save_cookies(str(juming_path))
loaded_jc = jucha.JC()
loaded_jc.load_cookies(str(jucha_path))
loaded_jm = juming.JM()
loaded_jm.load_cookies(str(juming_path))
self.assertEqual("cookie", loaded_jc.cookie.get("jc"))
self.assertEqual("cookie", loaded_jm.cookie.get("jm"))
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,60 @@
import unittest
import sys
from pathlib import Path
from unittest.mock import MagicMock
PROJECT_ROOT = Path(__file__).resolve().parents[2]
DOMAINCHECK_ROOT = PROJECT_ROOT / "domainCheck"
if str(DOMAINCHECK_ROOT) not in sys.path:
sys.path.insert(0, str(DOMAINCHECK_ROOT))
from domainCheck.detect import jucha, juming, juziseo
class LegacyDetectorTimeoutTests(unittest.TestCase):
def test_juziseo_request_helper_applies_default_timeout(self):
client = juziseo.Juziseo()
client.session = MagicMock()
client._request("get", "https://example.com/api", headers={"x": "1"})
client.session.get.assert_called_once_with(
"https://example.com/api",
headers={"x": "1"},
timeout=client.request_timeout,
)
def test_jucha_request_helper_applies_default_timeout(self):
client = jucha.JC()
client.session = MagicMock()
client._request("post", "https://example.com/api", data={"a": 1})
client.session.post.assert_called_once_with(
"https://example.com/api",
data={"a": 1},
timeout=client.request_timeout,
)
def test_juming_request_helper_applies_default_timeout(self):
client = juming.JM()
client.session = MagicMock()
client._request("post", "https://example.com/api", json={"a": 1})
client.session.post.assert_called_once_with(
"https://example.com/api",
json={"a": 1},
timeout=client.request_timeout,
)
def test_request_helper_preserves_explicit_timeout_override(self):
client = juming.JM()
client.session = MagicMock()
client._request("get", "https://example.com/download", timeout=60)
client.session.get.assert_called_once_with(
"https://example.com/download",
timeout=60,
)
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,39 @@
import unittest
from unittest.mock import patch
from app.utils.redis_client import get_redis_client, reset_redis_clients_for_tests
class DomainCheckRedisClientTests(unittest.TestCase):
def tearDown(self) -> None:
reset_redis_clients_for_tests()
@patch("app.utils.redis_client.redis.Redis")
@patch("app.utils.redis_client.redis.BlockingConnectionPool")
def test_standard_client_is_cached_per_process(self, mock_pool, mock_redis) -> None:
client = object()
mock_redis.return_value = client
first = get_redis_client()
second = get_redis_client()
self.assertIs(first, client)
self.assertIs(second, client)
mock_pool.assert_called_once()
mock_redis.assert_called_once()
@patch("app.utils.redis_client.redis.Redis")
@patch("app.utils.redis_client.redis.BlockingConnectionPool")
def test_pubsub_role_uses_separate_cached_client(self, mock_pool, mock_redis) -> None:
mock_redis.side_effect = [object(), object()]
standard_client = get_redis_client()
pubsub_client = get_redis_client(role="pubsub")
self.assertIsNot(standard_client, pubsub_client)
self.assertEqual(2, mock_pool.call_count)
self.assertEqual(2, mock_redis.call_count)
if __name__ == "__main__":
unittest.main()

View File

@@ -7,6 +7,9 @@ from domainCheck.detect import register
class RegisterTimeoutConfigTests(unittest.TestCase):
def tearDown(self):
register._PROXY_MANAGERS.clear()
def test_proxy_timeout_uses_bounded_total(self):
with patch.dict("os.environ", {}, clear=False):
timeout = register._resolve_register_timeout({"http": "http://127.0.0.1:8080"})
@@ -59,6 +62,19 @@ class RegisterTimeoutConfigTests(unittest.TestCase):
allow_redirects=True,
)
def test_proxy_manager_cache_evicts_oldest_entry(self):
with patch.dict("os.environ", {"DOMAINCHECK_REGISTER_PROXY_SESSION_CACHE_SIZE": "2"}, clear=False):
first = register._get_http_manager({"http": "http://127.0.0.1:8080"})
second = register._get_http_manager({"http": "http://127.0.0.1:8081"})
third = register._get_http_manager({"http": "http://127.0.0.1:8082"})
self.assertEqual(2, len(register._PROXY_MANAGERS))
self.assertIsNotNone(second)
self.assertIsNotNone(third)
recreated_first = register._get_http_manager({"http": "http://127.0.0.1:8080"})
self.assertIsNot(first, recreated_first)
if __name__ == "__main__":
unittest.main()

View File

@@ -5,6 +5,9 @@ from domainCheck.detect import aizhan, baidu, c360, chinaz
class StepTimeoutBudgetTests(unittest.TestCase):
def tearDown(self):
c360._PROXY_SESSIONS.clear()
def test_baidu_timeout_respects_remaining_budget(self):
with patch.dict("os.environ", {}, clear=False):
timeout = baidu._resolve_baidu_timeout({"http": "http://127.0.0.1:8080"}, budget_seconds=1.1)
@@ -17,6 +20,33 @@ class StepTimeoutBudgetTests(unittest.TestCase):
self.assertLessEqual(timeout, 0.9)
self.assertGreaterEqual(timeout, 0.6)
def test_360_direct_session_disables_env_proxy_and_is_reused(self):
first = c360._get_session()
second = c360._get_session()
self.assertIs(first, second)
self.assertFalse(first.trust_env)
def test_360_proxy_session_is_reused_per_proxy_url(self):
first = c360._get_session({"http": "http://127.0.0.1:8080"})
second = c360._get_session({"https": "http://127.0.0.1:8080"})
third = c360._get_session({"http": "http://127.0.0.1:8081"})
self.assertIs(first, second)
self.assertIsNot(first, third)
self.assertFalse(first.trust_env)
def test_360_proxy_session_cache_evicts_oldest_entry(self):
with patch.dict("os.environ", {"DOMAINCHECK_360_PROXY_SESSION_CACHE_SIZE": "2"}, clear=False):
first = c360._get_session({"http": "http://127.0.0.1:8080"})
second = c360._get_session({"http": "http://127.0.0.1:8081"})
third = c360._get_session({"http": "http://127.0.0.1:8082"})
self.assertEqual(2, len(c360._PROXY_SESSIONS))
self.assertIsNotNone(second)
self.assertIsNotNone(third)
recreated_first = c360._get_session({"http": "http://127.0.0.1:8080"})
self.assertIsNot(first, recreated_first)
def test_chinaz_timeout_respects_remaining_budget(self):
timeout = chinaz._resolve_chinaz_timeout({"http": "http://127.0.0.1:8080"}, budget_seconds=1.3)
self.assertLessEqual(timeout, 1.3)

View File

@@ -92,7 +92,47 @@ class WaybackDetectorRecentYearsTest(unittest.TestCase):
# 最新快照会先独立尝试一次,再进入裁剪后的扫描窗口。
self.assertEqual(4, result["fetched_snapshot_count"])
def test_scan_snapshots_fast_degrades_when_latest_cdx_is_transient_failure(self):
def test_scan_snapshots_fetches_recent_cdx_window_instead_of_full_history(self):
detector = WaybackDetector.__new__(WaybackDetector)
fetch_limits = []
def fake_fetch(domain, limit=None, fast_latest=False):
fetch_limits.append((limit, fast_latest))
if fast_latest:
return {
"records": [{"timestamp": "20260101000000", "digest": "latest"}],
"error": None,
}
return {
"records": [
{"timestamp": "20260101000000", "digest": "latest"},
{"timestamp": "20251201000000", "digest": "older-1"},
{"timestamp": "20251101000000", "digest": "older-2"},
{"timestamp": "20251001000000", "digest": "older-3"},
],
"error": None,
}
detector._fetch_cdx_records_with_meta = fake_fetch
detector._load_cached_records = lambda domain: None
detector._save_cached_records = lambda domain, records: None
detector._save_cached_timestamps = lambda domain, timestamps: None
detector._fetch_snapshot_title = lambda domain, timestamp: {
"timestamp": timestamp,
"title": f"title-{timestamp}",
"ok": True,
}
detector._normalize_title = lambda title: title
detector._find_sensitive_word = lambda title, words: None
detector._log_info = lambda message: None
detector._handle_exception = lambda exc, domain: None
with patch("app.detectors.wayback_detector.config.WAYBACK_MAX_RECORDS", 3):
detector.scan_snapshots("example.com", sensitive_words=[], recent_years=5)
self.assertEqual((-18, False), fetch_limits[1])
def test_scan_snapshots_continues_with_cached_records_when_latest_cdx_is_transient_failure(self):
detector = WaybackDetector.__new__(WaybackDetector)
detector._fetch_cdx_records_with_meta = lambda domain, limit=None, fast_latest=False: {
"records": [],
@@ -104,19 +144,55 @@ class WaybackDetectorRecentYearsTest(unittest.TestCase):
]
detector._save_cached_records = lambda domain, records: None
detector._save_cached_timestamps = lambda domain, timestamps: None
detector._fetch_snapshot_title = lambda domain, timestamp: {
"timestamp": timestamp,
"title": f"title-{timestamp}",
"ok": True,
}
detector._normalize_title = lambda title: title
detector._find_sensitive_word = lambda title, words: None
detector._log_info = lambda message: None
detector._handle_exception = lambda exc, domain: None
detector._trip_transient_backoff = lambda: self.fail("single latest_cdx timeout should not trigger global backoff")
result = detector.scan_snapshots("example.com", sensitive_words=[], recent_years=5)
self.assertEqual(2, result["checked_snapshot_count"])
self.assertEqual(2, result["fetched_snapshot_count"])
self.assertGreaterEqual(result["request_error_count"], 1)
def test_scan_snapshots_skips_records_cdx_when_latest_cdx_transient_fails_without_cache(self):
detector = WaybackDetector.__new__(WaybackDetector)
fetch_calls = []
def fake_fetch(domain, limit=None, fast_latest=False):
fetch_calls.append((limit, fast_latest))
return {
"records": [],
"error": "HTTPSConnectionPool(host='web.archive.org', port=443): Max retries exceeded",
}
detector._fetch_cdx_records_with_meta = fake_fetch
detector._load_cached_records = lambda domain: None
detector._save_cached_records = lambda domain, records: None
detector._save_cached_timestamps = lambda domain, timestamps: None
detector._fetch_snapshot_title = lambda domain, timestamp: self.fail("should not fetch snapshot titles")
detector._normalize_title = lambda title: title
detector._find_sensitive_word = lambda title, words: None
detector._log_info = lambda message: None
detector._handle_exception = lambda exc, domain: None
trip_calls = []
detector._trip_transient_backoff = lambda: trip_calls.append("trip")
result = detector.scan_snapshots("example.com", sensitive_words=[], recent_years=5)
self.assertEqual([(-1, True)], fetch_calls)
self.assertEqual(1, len(trip_calls))
self.assertEqual(0, result["checked_snapshot_count"])
self.assertGreaterEqual(result["failed_snapshot_count"], 1)
self.assertEqual(0, result["fetched_snapshot_count"])
self.assertGreaterEqual(result["request_error_count"], 1)
def test_scan_snapshots_fast_degrades_when_latest_snapshot_is_transient_failure(self):
def test_scan_snapshots_continues_when_latest_snapshot_is_transient_failure(self):
detector = WaybackDetector.__new__(WaybackDetector)
detector._fetch_cdx_records_with_meta = lambda domain, limit=None, fast_latest=False: {
"records": [{"timestamp": "20260101000000", "digest": "latest"}] if fast_latest else [
@@ -128,21 +204,31 @@ class WaybackDetectorRecentYearsTest(unittest.TestCase):
detector._load_cached_records = lambda domain: None
detector._save_cached_records = lambda domain, records: None
detector._save_cached_timestamps = lambda domain, timestamps: None
detector._fetch_snapshot_title = lambda domain, timestamp: {
"timestamp": timestamp,
"title": "",
"ok": False,
"error": "HTTPSConnectionPool(host='web.archive.org', port=443): Max retries exceeded",
}
def fake_fetch(domain, timestamp):
if timestamp == "20260101000000":
return {
"timestamp": timestamp,
"title": "",
"ok": False,
"error": "HTTPSConnectionPool(host='web.archive.org', port=443): Max retries exceeded",
}
return {
"timestamp": timestamp,
"title": "older-title",
"ok": True,
}
detector._fetch_snapshot_title = fake_fetch
detector._normalize_title = lambda title: title
detector._find_sensitive_word = lambda title, words: None
detector._log_info = lambda message: None
detector._handle_exception = lambda exc, domain: None
detector._trip_transient_backoff = lambda: self.fail("single latest_snapshot timeout should not trigger global backoff")
result = detector.scan_snapshots("example.com", sensitive_words=[], recent_years=5)
self.assertEqual(1, result["checked_snapshot_count"])
self.assertEqual(0, result["fetched_snapshot_count"])
self.assertEqual(2, result["checked_snapshot_count"])
self.assertEqual(1, result["fetched_snapshot_count"])
self.assertGreaterEqual(result["failed_snapshot_count"], 1)
self.assertTrue(
any("latest_snapshot:" in item for item in result["request_errors"])
@@ -168,6 +254,87 @@ class WaybackDetectorRecentYearsTest(unittest.TestCase):
self.assertEqual(1, result["request_error_count"])
self.assertTrue(any("wayback_backoff_active:" in item for item in result["request_errors"]))
def test_scan_snapshots_does_not_trip_global_backoff_on_single_snapshot_timeout(self):
detector = WaybackDetector.__new__(WaybackDetector)
detector._fetch_cdx_records_with_meta = lambda domain, limit=None, fast_latest=False: {
"records": [{"timestamp": "20260101000000", "digest": "latest"}] if fast_latest else [
{"timestamp": "20260101000000", "digest": "latest"},
{"timestamp": "20250101000000", "digest": "older"},
],
"error": None,
}
detector._load_cached_records = lambda domain: None
detector._save_cached_records = lambda domain, records: None
detector._save_cached_timestamps = lambda domain, timestamps: None
def fake_fetch(domain, timestamp):
if timestamp == "20260101000000":
return {"timestamp": timestamp, "title": "latest", "ok": True}
return {
"timestamp": timestamp,
"title": "",
"ok": False,
"error": "HTTPSConnectionPool(host='web.archive.org', port=443): Max retries exceeded",
}
detector._fetch_snapshot_title = fake_fetch
detector._normalize_title = lambda title: title
detector._find_sensitive_word = lambda title, words: None
detector._log_info = lambda message: None
detector._handle_exception = lambda exc, domain: None
detector._trip_transient_backoff = lambda: self.fail("single snapshot timeout should not trigger global backoff")
result = detector.scan_snapshots("example.com", sensitive_words=[], recent_years=5)
self.assertEqual(2, result["checked_snapshot_count"])
self.assertEqual(1, result["fetched_snapshot_count"])
self.assertEqual(1, result["failed_snapshot_count"])
def test_scan_snapshots_trips_global_backoff_after_threshold_transient_failures(self):
detector = WaybackDetector.__new__(WaybackDetector)
detector._fetch_cdx_records_with_meta = lambda domain, limit=None, fast_latest=False: {
"records": [{"timestamp": "20260101000000", "digest": "latest"}] if fast_latest else [
{"timestamp": "20260101000000", "digest": "latest"},
{"timestamp": "20250101000000", "digest": "older-1"},
{"timestamp": "20240101000000", "digest": "older-2"},
],
"error": None,
}
detector._load_cached_records = lambda domain: None
detector._save_cached_records = lambda domain, records: None
detector._save_cached_timestamps = lambda domain, timestamps: None
def fake_fetch(domain, timestamp):
if timestamp == "20260101000000":
return {"timestamp": timestamp, "title": "latest", "ok": True}
return {
"timestamp": timestamp,
"title": "",
"ok": False,
"error": "HTTPSConnectionPool(host='web.archive.org', port=443): Max retries exceeded",
}
detector._fetch_snapshot_title = fake_fetch
detector._normalize_title = lambda title: title
detector._find_sensitive_word = lambda title, words: None
detector._log_info = lambda message: None
detector._handle_exception = lambda exc, domain: None
trip_calls = []
detector._trip_transient_backoff = lambda: trip_calls.append("trip")
result = detector.scan_snapshots("example.com", sensitive_words=[], recent_years=5)
self.assertEqual(1, len(trip_calls))
self.assertEqual(3, result["checked_snapshot_count"])
self.assertEqual(2, result["failed_snapshot_count"])
def test_build_session_disables_env_proxy(self):
detector = WaybackDetector.__new__(WaybackDetector)
session = detector._build_session()
self.assertFalse(session.trust_env)
if __name__ == "__main__":
unittest.main()