已合并
[bugfix] NodeManager判断register成功逻辑有误 #543
liu创建于 7月17日
[bugfix] NodeManager判断register成功逻辑有误 #543
已合并
共 5 个文件变更+49-55
| @@ -75,23 +75,7 @@ class ControllerApiClient: | |||
| 75 | nodemanager_config = NodeManagerConfig.from_json() | 75 | nodemanager_config = NodeManagerConfig.from_json() |
| 76 | 76 | ||
| 77 | 77 | ||
| 78 | - def register(register_msg: RegisterMsg): | 78 | + def register(register_msg: RegisterMsg) -> bool: |
| 79 | - # Read config values under lock protection | ||
| 80 | - client_args = {} | ||
| 81 | - try: | ||
| 82 | - client_args = ControllerApiClient._generate_client_args() | ||
| 83 | - with SafeHTTPSClient(timeout=15, **client_args) as client: | ||
| 84 | - _ = client.post("/controller/register", register_msg.model_dump()) | ||
| 85 | - logger.info("Register success!") | ||
| 86 | - return True | ||
| 87 | - except Exception as e: | ||
| 88 | - logger.error( | ||
| 89 | - "Exception occurred while register to controller at %s: %s", client_args.get("address", "unknown"), e | ||
| 90 | - ) | ||
| 91 | - return False | ||
| 92 | - | ||
| 93 | - | ||
| 94 | - def register_after_restore(register_msg: RegisterMsg) -> bool: | ||
| 95 | client_args = {} | 79 | client_args = {} |
| 96 | try: | 80 | try: |
| 97 | client_args = ControllerApiClient._generate_client_args() | 81 | client_args = ControllerApiClient._generate_client_args() |
| @@ -104,13 +88,13 @@ class ControllerApiClient: | |||
| 104 | return False | 88 | return False |
| 105 | 89 | ||
| 106 | if not isinstance(response, dict): | 90 | if not isinstance(response, dict): |
| 107 | - logger.error("Invalid register response from controller after restore: %s", response) | 91 | + logger.error("Invalid register response from controller: %s", response) |
| 108 | return False | 92 | return False |
| 109 | if error := response.get("error"): | 93 | if error := response.get("error"): |
| 110 | - logger.warning("Register rejected by controller after restore: %s", error) | 94 | + logger.warning("Register rejected by controller: %s", error) |
| 111 | return False | 95 | return False |
| 112 | 96 | ||
| 113 | - logger.info("Register after restore success!") | 97 | + logger.info("Register success!") |
| 114 | return True | 98 | return True |
| 115 | 99 | ||
| 116 | 100 | ||
| @@ -213,14 +213,6 @@ class EngineManager(ThreadSafeSingleton): | |||
| 213 | 213 | ||
| 214 | return ControllerApiClient.register(register_msg) | 214 | return ControllerApiClient.register(register_msg) |
| 215 | 215 | ||
| 216 | - def post_register_msg_after_restore(self) -> bool | None: | ||
| 217 | - register_msg = self._gen_register_msg() | ||
| 218 | - if register_msg is None: | ||
| 219 | - return False | ||
| 220 | - logger.debug("register_msg is %s", register_msg) | ||
| 221 | - | ||
| 222 | - return ControllerApiClient.register_after_restore(register_msg) | ||
| 223 | - | ||
| 224 | def post_reregister_msg(self) -> bool | None: | 216 | def post_reregister_msg(self) -> bool | None: |
| 225 | reregister_msg = self._gen_reregister_msg() | 217 | reregister_msg = self._gen_reregister_msg() |
| 226 | if reregister_msg is None: | 218 | if reregister_msg is None: |
| @@ -401,7 +401,7 @@ class HeartbeatManager(ThreadSafeSingleton): | |||
| 401 | # Register for post-snapshot brandnew job name | 401 | # Register for post-snapshot brandnew job name |
| 402 | # Do not consider retry | 402 | # Do not consider retry |
| 403 | # If current register failed, next register will be triggered by next heartbeat report exception | 403 | # If current register failed, next register will be triggered by next heartbeat report exception |
| 404 | - ret = EngineManager().post_register_msg_after_restore() | 404 | + ret = EngineManager().post_register_msg() |
| 405 | self._is_registered_after_restore = ret is True | 405 | self._is_registered_after_restore = ret is True |
| 406 | 406 | ||
| 407 | def _reregister(self) -> None: | 407 | def _reregister(self) -> None: |
| @@ -232,9 +232,7 @@ class TestEngineManager: | |||
| 232 | with pytest.raises(TypeError): | 232 | with pytest.raises(TypeError): |
| 233 | engine_manager._gen_reregister_msg() | 233 | engine_manager._gen_reregister_msg() |
| 234 | 234 | ||
| 235 | - @patch("motor.node_manager.core.engine_manager.ControllerApiClient.register") | 235 | + def _prepare_post_register_config(self, engine_manager): |
| 236 | - def test_post_register_msg_success(self, mock_register, engine_manager): | ||
| 237 | - """Test post_register_msg with successful response""" | ||
| 238 | engine_manager._config.basic_config.job_name = "test_job" | 236 | engine_manager._config.basic_config.job_name = "test_job" |
| 239 | engine_manager._config.basic_config.model_name = "test_model" | 237 | engine_manager._config.basic_config.model_name = "test_model" |
| 240 | engine_manager._config.basic_config.role = PDRole.ROLE_U | 238 | engine_manager._config.basic_config.role = PDRole.ROLE_U |
| @@ -246,29 +244,49 @@ class TestEngineManager: | |||
| 246 | engine_manager._config.api_config.coordinator_api_mgmt_port = 8080 | 244 | engine_manager._config.api_config.coordinator_api_mgmt_port = 8080 |
| 247 | engine_manager._config.basic_config.parallel_config = ParallelConfig(tp_size=2, pp_size=1) | 245 | engine_manager._config.basic_config.parallel_config = ParallelConfig(tp_size=2, pp_size=1) |
| 248 | 246 | ||
| 249 | - mock_register.return_value = True | 247 | + @patch("motor.node_manager.api_client.controller_api_client.ControllerApiClient._generate_client_args") |
| 248 | + | ||
| 249 | + def test_post_register_msg_success(self, mock_http, mock_client_args, engine_manager): | ||
| 250 | + self._prepare_post_register_config(engine_manager) | ||
| 251 | + mock_client_args.return_value = {"address": "controller:8080", "tls_config": None} | ||
| 252 | + mock_http.return_value.__enter__.return_value.post.return_value = {"status": "ok"} | ||
| 250 | 253 | ||
| 251 | result = engine_manager.post_register_msg() | 254 | result = engine_manager.post_register_msg() |
| 255 | + | ||
| 252 | assert result is True | 256 | assert result is True |
| 253 | - mock_register.assert_called_once() | 257 | + mock_http.return_value.__enter__.return_value.post.assert_called_once() |
| 254 | 258 | ||
| 255 | - @patch("motor.node_manager.core.engine_manager.ControllerApiClient.register") | 259 | + @patch("motor.node_manager.api_client.controller_api_client.ControllerApiClient._generate_client_args") |
| 256 | - def test_post_register_msg_failure(self, mock_register, engine_manager): | 260 | + @patch("motor.node_manager.api_client.controller_api_client.SafeHTTPSClient") |
| 257 | - """Test post_register_msg with exception""" | 261 | + def test_post_register_msg_failure_on_exception(self, mock_http, mock_client_args, engine_manager): |
| 258 | - engine_manager._config.basic_config.job_name = "test_job" | 262 | + self._prepare_post_register_config(engine_manager) |
| 259 | - engine_manager._config.basic_config.model_name = "test_model" | 263 | + mock_client_args.return_value = {"address": "controller:8080", "tls_config": None} |
| 260 | - engine_manager._config.basic_config.role = PDRole.ROLE_U | 264 | + mock_http.return_value.__enter__.return_value.post.side_effect = RuntimeError("connection refused") |
| 261 | - engine_manager._config.api_config.pod_ip = "192.168.1.100" | ||
| 262 | - engine_manager._config.api_config.host_ip = "192.168.1.200" | ||
| 263 | - engine_manager._config.endpoint_config.service_ports = ["8080"] | ||
| 264 | - engine_manager._config.api_config.node_manager_port = 8080 | ||
| 265 | - engine_manager._config.api_config.coordinator_api_dns = "localhost" | ||
| 266 | - engine_manager._config.api_config.coordinator_api_mgmt_port = 8080 | ||
| 267 | - engine_manager._config.basic_config.parallel_config = ParallelConfig(tp_size=2, pp_size=1) | ||
| 268 | - | ||
| 269 | - mock_register.return_value = False | ||
| 270 | 265 | ||
| 271 | result = engine_manager.post_register_msg() | 266 | result = engine_manager.post_register_msg() |
| 267 | + | ||
| 268 | + assert result is False | ||
| 269 | + | ||
| 270 | + | ||
| 271 | + | ||
| 272 | + def test_post_register_msg_failure_on_rejected(self, mock_http, mock_client_args, engine_manager): | ||
| 273 | + self._prepare_post_register_config(engine_manager) | ||
| 274 | + mock_client_args.return_value = {"address": "controller:8080", "tls_config": None} | ||
| 275 | + mock_http.return_value.__enter__.return_value.post.return_value = {"error": "already registered"} | ||
| 276 | + | ||
| 277 | + result = engine_manager.post_register_msg() | ||
| 278 | + | ||
| 279 | + assert result is False | ||
| 280 | + | ||
| 281 | + | ||
| 282 | + | ||
| 283 | + def test_post_register_msg_failure_on_invalid_response(self, mock_http, mock_client_args, engine_manager): | ||
| 284 | + self._prepare_post_register_config(engine_manager) | ||
| 285 | + mock_client_args.return_value = {"address": "controller:8080", "tls_config": None} | ||
| 286 | + mock_http.return_value.__enter__.return_value.post.return_value = "not-a-dict" | ||
| 287 | + | ||
| 288 | + result = engine_manager.post_register_msg() | ||
| 289 | + | ||
| 272 | assert result is False | 290 | assert result is False |
| 273 | 291 | ||
| 274 | 292 | ||
| @@ -710,13 +710,13 @@ class TestHeartBeatManager: | |||
| 710 | 710 | ||
| 711 | def test_register_after_restore_success(self, mock_engine_manager_class, heart_beat_manager): | 711 | def test_register_after_restore_success(self, mock_engine_manager_class, heart_beat_manager): |
| 712 | mock_engine_manager = MagicMock() | 712 | mock_engine_manager = MagicMock() |
| 713 | - mock_engine_manager.post_register_msg_after_restore.return_value = True | 713 | + mock_engine_manager.post_register_msg.return_value = True |
| 714 | mock_engine_manager_class.return_value = mock_engine_manager | 714 | mock_engine_manager_class.return_value = mock_engine_manager |
| 715 | 715 | ||
| 716 | heart_beat_manager._register_after_restore() | 716 | heart_beat_manager._register_after_restore() |
| 717 | 717 | ||
| 718 | mock_engine_manager.register_prepare_after_restore.assert_called_once() | 718 | mock_engine_manager.register_prepare_after_restore.assert_called_once() |
| 719 | - mock_engine_manager.post_register_msg_after_restore.assert_called_once() | 719 | + mock_engine_manager.post_register_msg.assert_called_once() |
| 720 | assert heart_beat_manager._is_registered_after_restore is True | 720 | assert heart_beat_manager._is_registered_after_restore is True |
| 721 | 721 | ||
| 722 | 722 | ||
| @@ -727,7 +727,7 @@ class TestHeartBeatManager: | |||
| 727 | 727 | ||
| 728 | heart_beat_manager._register_after_restore() | 728 | heart_beat_manager._register_after_restore() |
| 729 | 729 | ||
| 730 | - mock_engine_manager.post_register_msg_after_restore.assert_not_called() | 730 | + mock_engine_manager.post_register_msg.assert_not_called() |
| 731 | assert heart_beat_manager._is_registered_after_restore is False | 731 | assert heart_beat_manager._is_registered_after_restore is False |
| 732 | assert heart_beat_manager._register_after_restore_retry_count == 1 | 732 | assert heart_beat_manager._register_after_restore_retry_count == 1 |
| 733 | 733 | ||
| @@ -747,14 +747,14 @@ class TestHeartBeatManager: | |||
| 747 | 747 | ||
| 748 | def mock_stop_sleep(_seconds): | 748 | def mock_stop_sleep(_seconds): |
| 749 | call_count["count"] += 1 | 749 | call_count["count"] += 1 |
| 750 | - # First sleep happens after restore register; heartbeat is sent on the next loop. | 750 | + # First sleep happens after register; heartbeat is sent on the next loop. |
| 751 | if call_count["count"] >= 2: | 751 | if call_count["count"] >= 2: |
| 752 | heart_beat_manager.stop_event.set() | 752 | heart_beat_manager.stop_event.set() |
| 753 | 753 | ||
| 754 | mock_engine_manager = MagicMock() | 754 | mock_engine_manager = MagicMock() |
| 755 | mock_engine_manager.is_engine_checkpoint_done.return_value = True | 755 | mock_engine_manager.is_engine_checkpoint_done.return_value = True |
| 756 | mock_engine_manager.register_prepare_after_restore.return_value = None | 756 | mock_engine_manager.register_prepare_after_restore.return_value = None |
| 757 | - mock_engine_manager.post_register_msg_after_restore.return_value = True | 757 | + mock_engine_manager.post_register_msg.return_value = True |
| 758 | mock_engine_manager_class.return_value = mock_engine_manager | 758 | mock_engine_manager_class.return_value = mock_engine_manager |
| 759 | mock_report_heartbeat.return_value = None | 759 | mock_report_heartbeat.return_value = None |
| 760 | mock_sleep.side_effect = mock_stop_sleep | 760 | mock_sleep.side_effect = mock_stop_sleep |
| @@ -770,7 +770,7 @@ class TestHeartBeatManager: | |||
| 770 | heart_beat_manager._report_heartbeat_loop() | 770 | heart_beat_manager._report_heartbeat_loop() |
| 771 | 771 | ||
| 772 | mock_engine_manager.register_prepare_after_restore.assert_called_once() | 772 | mock_engine_manager.register_prepare_after_restore.assert_called_once() |
| 773 | - mock_engine_manager.post_register_msg_after_restore.assert_called_once() | 773 | + mock_engine_manager.post_register_msg.assert_called_once() |
| 774 | mock_report_heartbeat.assert_called_once() | 774 | mock_report_heartbeat.assert_called_once() |
| 775 | 775 | ||
| 776 | 776 | ||