441 lines
20 KiB
Python
441 lines
20 KiB
Python
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()
|