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