Files
getDomain/domain-api/tests/test_worker_control_service.py

352 lines
14 KiB
Python

import json
import subprocess
import unittest
from unittest.mock import Mock, patch
from app.core.config import settings
from app.services.worker_control_service import (
WORKER_CONTROL_CHANNEL,
WORKER_PENDING_COMMAND_KEY,
detect_worker_runtime,
send_worker_command,
start_worker,
stop_worker,
)
class WorkerControlServiceTests(unittest.TestCase):
@patch("app.services.worker_control_service.get_redis")
def test_send_worker_command_persists_request_id_and_publishes_same_payload(self, mock_get_redis) -> None:
redis_client = Mock()
mock_get_redis.return_value = redis_client
with patch("app.services.worker_control_service._runtime_config", return_value={"worker_mode": "windows-local", "worker_service_name": "domaincheck-worker"}):
ok, message = send_worker_command(
"start_detection",
payload={"job_id": 1, "job_code": "detect-20260419030000-abc123"},
)
self.assertTrue(ok)
self.assertIn("已发送 Worker 控制指令", message)
redis_client.set.assert_called_once()
set_args = redis_client.set.call_args.args
self.assertEqual(WORKER_PENDING_COMMAND_KEY, set_args[0])
serialized = set_args[1]
self.assertEqual(120, redis_client.set.call_args.kwargs["ex"])
payload = json.loads(serialized)
self.assertEqual("start_detection", payload["action"])
self.assertEqual(1, payload["job_id"])
self.assertEqual("detect-20260419030000-abc123", payload["job_code"])
self.assertTrue(payload["request_id"].startswith("workerctl-"))
redis_client.publish.assert_called_once_with(WORKER_CONTROL_CHANNEL, serialized)
@patch("app.services.worker_control_service.get_redis")
def test_send_worker_command_scopes_pending_command_to_target_worker_instances(self, mock_get_redis) -> None:
redis_client = Mock()
mock_get_redis.return_value = redis_client
with patch("app.services.worker_control_service._runtime_config", return_value={"worker_mode": "linux-systemd", "worker_service_name": "domaincheck-worker"}):
ok, message = send_worker_command(
"start_detection",
payload={
"job_id": 2,
"target_node_codes": ["mainland-controller-01-a", "mainland-controller-01-b"],
},
)
self.assertTrue(ok)
self.assertIn("mainland-controller-01-a,mainland-controller-01-b", message)
self.assertEqual(2, redis_client.set.call_count)
set_keys = [call.args[0] for call in redis_client.set.call_args_list]
self.assertEqual(
[
f"{WORKER_PENDING_COMMAND_KEY}:mainland-controller-01-a",
f"{WORKER_PENDING_COMMAND_KEY}:mainland-controller-01-b",
],
set_keys,
)
serialized = redis_client.set.call_args_list[0].args[1]
payload = json.loads(serialized)
self.assertEqual(["mainland-controller-01-a", "mainland-controller-01-b"], payload["target_node_codes"])
redis_client.publish.assert_called_once_with(WORKER_CONTROL_CHANNEL, serialized)
@patch("app.services.worker_control_service._expand_linux_worker_control_units")
@patch("app.services.worker_control_service.get_redis")
def test_send_worker_command_expands_local_linux_worker_instances_when_targets_unspecified(
self,
mock_get_redis,
mock_expand_linux_worker_control_units,
) -> None:
redis_client = Mock()
mock_get_redis.return_value = redis_client
mock_expand_linux_worker_control_units.return_value = [
"domaincheck-worker",
"domaincheck-worker@a.service",
"domaincheck-worker@b.service",
]
with patch("app.services.worker_control_service._runtime_config", return_value={"worker_mode": "linux-systemd", "worker_service_name": "domaincheck-worker"}), \
patch.object(settings, "node_code", "mainland-controller-01"):
ok, message = send_worker_command("start_detection", payload={"job_id": 3})
self.assertTrue(ok)
self.assertIn("mainland-controller-01,mainland-controller-01-a,mainland-controller-01-b", message)
self.assertEqual(3, redis_client.set.call_count)
set_keys = [call.args[0] for call in redis_client.set.call_args_list]
self.assertEqual(
[
f"{WORKER_PENDING_COMMAND_KEY}:mainland-controller-01",
f"{WORKER_PENDING_COMMAND_KEY}:mainland-controller-01-a",
f"{WORKER_PENDING_COMMAND_KEY}:mainland-controller-01-b",
],
set_keys,
)
serialized = redis_client.set.call_args_list[0].args[1]
payload = json.loads(serialized)
self.assertEqual(
["mainland-controller-01", "mainland-controller-01-a", "mainland-controller-01-b"],
payload["target_node_codes"],
)
redis_client.publish.assert_called_once_with(WORKER_CONTROL_CHANNEL, serialized)
@patch("app.services.worker_control_service._build_direct_redis_client")
@patch("app.services.worker_control_service.get_redis")
def test_send_worker_command_falls_back_to_direct_redis_client(
self,
mock_get_redis,
mock_build_direct_redis_client,
) -> None:
mock_get_redis.side_effect = RecursionError("maximum recursion depth exceeded")
direct_client = Mock()
mock_build_direct_redis_client.return_value = direct_client
with patch("app.services.worker_control_service._runtime_config", return_value={"worker_mode": "windows-local", "worker_service_name": "domaincheck-worker"}):
ok, message = send_worker_command("start_detection", payload={"job_id": 9})
self.assertTrue(ok)
self.assertIn("已发送 Worker 控制指令", message)
direct_client.set.assert_called_once()
serialized = direct_client.set.call_args.args[1]
payload = json.loads(serialized)
self.assertEqual("start_detection", payload["action"])
self.assertEqual(9, payload["job_id"])
direct_client.publish.assert_called_once_with(WORKER_CONTROL_CHANNEL, serialized)
direct_client.close.assert_called_once()
@patch("app.services.worker_control_service._probe_linux_worker_instance_count", return_value=0)
@patch("app.services.worker_control_service._run_shell")
@patch("app.services.worker_control_service.probe_systemd_service")
@patch("app.services.worker_control_service._runtime_config")
def test_detect_worker_runtime_prefers_fast_pgrep_probe(
self,
mock_runtime_config,
mock_probe_systemd_service,
mock_run_shell,
_mock_instance_count,
) -> None:
mock_runtime_config.return_value = {
"worker_mode": "linux-systemd",
"worker_service_name": "domaincheck-worker",
}
mock_probe_systemd_service.return_value = {
"mode": "linux-systemd",
"service_name": "domaincheck-worker",
"running": True,
"process_count": 1,
"latest_start_time": "2026-04-22 23:00:00",
"message": "active/running",
}
mock_run_shell.return_value = subprocess.CompletedProcess(
args=["bash", "-lc", "pgrep -fc '[d]etect_worker.py' || true"],
returncode=0,
stdout="80\n",
stderr="",
)
runtime = detect_worker_runtime()
self.assertTrue(runtime["running"])
self.assertEqual(80, runtime["process_count"])
self.assertEqual(1, mock_run_shell.call_count)
@patch("app.services.worker_control_service._probe_linux_worker_instance_count", return_value=0)
@patch("app.services.worker_control_service._run_shell")
@patch("app.services.worker_control_service.probe_systemd_service")
@patch("app.services.worker_control_service._runtime_config")
def test_detect_worker_runtime_falls_back_when_pgrep_probe_is_unavailable(
self,
mock_runtime_config,
mock_probe_systemd_service,
mock_run_shell,
_mock_instance_count,
) -> None:
mock_runtime_config.return_value = {
"worker_mode": "linux-systemd",
"worker_service_name": "domaincheck-worker",
}
mock_probe_systemd_service.return_value = {
"mode": "linux-systemd",
"service_name": "domaincheck-worker",
"running": True,
"process_count": 1,
"latest_start_time": "2026-04-22 23:00:00",
"message": "active/running",
}
mock_run_shell.side_effect = [
subprocess.CompletedProcess(
args=["bash", "-lc", "pgrep -fc '[d]etect_worker.py' || true"],
returncode=0,
stdout="",
stderr="pgrep: command not found\n",
),
subprocess.CompletedProcess(
args=["bash", "-lc", "ps -eo args= | grep '[d]etect_worker.py' | wc -l"],
returncode=0,
stdout="12\n",
stderr="",
),
]
runtime = detect_worker_runtime()
self.assertTrue(runtime["running"])
self.assertEqual(12, runtime["process_count"])
self.assertEqual(2, mock_run_shell.call_count)
@patch("app.services.worker_control_service._probe_linux_worker_process_count", return_value=7)
@patch("app.services.worker_control_service._probe_linux_worker_instance_count", return_value=0)
@patch("app.services.worker_control_service.probe_systemd_service")
@patch("app.services.worker_control_service._runtime_config")
def test_detect_worker_runtime_keeps_service_offline_when_only_unmanaged_processes_exist(
self,
mock_runtime_config,
mock_probe_systemd_service,
_mock_instance_count,
_mock_process_count,
) -> None:
mock_runtime_config.return_value = {
"worker_mode": "linux-systemd",
"worker_service_name": "domaincheck-worker",
}
mock_probe_systemd_service.return_value = {
"mode": "linux-systemd",
"service_name": "domaincheck-worker",
"running": False,
"process_count": 0,
"latest_start_time": "",
"message": "inactive/dead",
}
runtime = detect_worker_runtime()
self.assertFalse(runtime["running"])
self.assertEqual(7, runtime["process_count"])
self.assertIn("unmanaged worker processes", runtime["message"])
@patch("app.services.worker_control_service._probe_linux_worker_process_count", return_value=30)
@patch("app.services.worker_control_service._probe_linux_worker_instance_count", return_value=3)
@patch("app.services.worker_control_service.probe_systemd_service")
@patch("app.services.worker_control_service._runtime_config")
def test_detect_worker_runtime_accepts_active_template_instances_when_base_service_is_inactive(
self,
mock_runtime_config,
mock_probe_systemd_service,
_mock_instance_count,
_mock_process_count,
) -> None:
mock_runtime_config.return_value = {
"worker_mode": "linux-systemd",
"worker_service_name": "domaincheck-worker",
}
mock_probe_systemd_service.return_value = {
"mode": "linux-systemd",
"service_name": "domaincheck-worker",
"running": False,
"process_count": 0,
"latest_start_time": "",
"message": "inactive/dead",
}
runtime = detect_worker_runtime()
self.assertTrue(runtime["running"])
self.assertEqual(30, runtime["process_count"])
self.assertEqual("template instances active (3)", runtime["message"])
@patch("app.services.worker_control_service._expand_linux_worker_control_units")
@patch("app.services.worker_control_service._run_systemctl")
@patch("app.services.worker_control_service._runtime_config")
def test_start_worker_includes_template_instances(
self,
mock_runtime_config,
mock_run_systemctl,
mock_expand_units,
) -> None:
mock_runtime_config.return_value = {
"worker_mode": "linux-systemd",
"worker_service_name": "domaincheck-worker",
}
mock_expand_units.return_value = [
"domaincheck-worker",
"domaincheck-worker@a",
"domaincheck-worker@b",
]
mock_run_systemctl.return_value = subprocess.CompletedProcess(
args=["systemctl", "start", "domaincheck-worker", "domaincheck-worker@a", "domaincheck-worker@b"],
returncode=0,
stdout="",
stderr="",
)
ok, message = start_worker()
self.assertTrue(ok)
self.assertIn("附带 2 个实例", message)
mock_run_systemctl.assert_called_once_with(
["start", "domaincheck-worker", "domaincheck-worker@a", "domaincheck-worker@b"],
timeout=45,
)
@patch("app.services.worker_control_service._expand_linux_worker_control_units")
@patch("app.services.worker_control_service._run_systemctl")
@patch("app.services.worker_control_service._runtime_config")
def test_stop_worker_includes_template_instances(
self,
mock_runtime_config,
mock_run_systemctl,
mock_expand_units,
) -> None:
mock_runtime_config.return_value = {
"worker_mode": "linux-systemd",
"worker_service_name": "domaincheck-worker",
}
mock_expand_units.return_value = [
"domaincheck-worker",
"domaincheck-worker@a",
"domaincheck-worker@b",
]
mock_run_systemctl.return_value = subprocess.CompletedProcess(
args=["systemctl", "stop", "domaincheck-worker", "domaincheck-worker@a", "domaincheck-worker@b"],
returncode=0,
stdout="",
stderr="",
)
ok, message = stop_worker()
self.assertTrue(ok)
self.assertIn("附带 2 个实例", message)
mock_run_systemctl.assert_called_once_with(
["stop", "domaincheck-worker", "domaincheck-worker@a", "domaincheck-worker@b"],
timeout=45,
)
if __name__ == "__main__":
unittest.main()