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