Files
getDomain/domainCheck/tests/test_database_job_status_refresh.py

201 lines
7.2 KiB
Python

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