diff --git a/gateway/run.py b/gateway/run.py index 4aee4c91ef918..a7ba984fe16f9 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -20772,6 +20772,7 @@ async def _watch_update_progress( chat_id = None session_key = None metadata = None + initiating_gateway_pid = None for path in (claimed_path, pending_path): if path.exists(): try: @@ -20780,6 +20781,7 @@ async def _watch_update_progress( chat_id = pending.get("chat_id") chat_type = pending.get("chat_type") session_key = pending.get("session_key") + initiating_gateway_pid = pending.get("initiating_gateway_pid") thread_id = pending.get("thread_id") message_id = pending.get("message_id") if platform_str and chat_id: @@ -20865,35 +20867,81 @@ async def _flush_buffer() -> None: pass await _flush_buffer() - # Send final status + # A successful gateway update writes its exit marker before it + # asks systemd to restart the gateway. The initiating process + # must leave the durable target marker for the next process; + # otherwise a send racing with shutdown can fail and erase the + # only record of who should receive the completion message. try: exit_code_raw = exit_code_path.read_text(encoding="utf-8").strip() or "1" exit_code = int(exit_code_raw) - if exit_code == 0: + except (OSError, ValueError): + exit_code = 1 + + if ( + exit_code == 0 + and initiating_gateway_pid is not None + and str(initiating_gateway_pid) == str(os.getpid()) + ): + try: await adapter.send( + chat_id, + "✅ Hermes update installed.\n♻️ Restarting gateway…", + metadata=_non_conversational_metadata(metadata, platform=platform), + ) + except Exception as e: + # The restart may already be stopping this process. + # Keeping the markers is what guarantees that the new + # gateway can still deliver the definitive result. + logger.info("Pre-restart update notification was not delivered: %s", e) + return + + # Send final status for failed updates and legacy markers that + # predate initiating_gateway_pid. Cleanup only after delivery. + delivered = False + try: + if exit_code == 0: + send_result = await adapter.send( chat_id, "✅ Hermes update finished.", metadata=_non_conversational_metadata(metadata, platform=platform), ) else: - await adapter.send( + send_result = await adapter.send( chat_id, "❌ Hermes update failed (exit code {}).".format(exit_code), metadata=_non_conversational_metadata(metadata, platform=platform), ) - logger.info("Update finished (exit=%s), notified %s", exit_code, session_key) + if ( + send_result is not None + and getattr(send_result, "success", True) is False + ): + logger.warning( + "Update final notification was not delivered: %s", + getattr(send_result, "error", "send returned success=False"), + ) + else: + delivered = True + logger.info( + "Update finished (exit=%s), notified %s", + exit_code, + session_key, + ) except Exception as e: logger.warning("Update final notification failed: %s", e) - # Cleanup - for p in (pending_path, claimed_path, output_path, - exit_code_path, prompt_path): - p.unlink(missing_ok=True) - (_hermes_home / ".update_response").unlink(missing_ok=True) - _up_done = self._peek_session_state(session_key) - if _up_done is not None: - _up_done.persistent.update_prompt_pending = False - return + if delivered: + for p in (pending_path, claimed_path, output_path, + exit_code_path, prompt_path): + p.unlink(missing_ok=True) + (_hermes_home / ".update_response").unlink(missing_ok=True) + _up_done = self._peek_session_state(session_key) + if _up_done is not None: + _up_done.persistent.update_prompt_pending = False + return + + await asyncio.sleep(poll_interval) + continue # Check for new output if output_path.exists(): @@ -21037,6 +21085,22 @@ async def _send_update_notification(self) -> bool: exit_code_raw = exit_code_path.read_text(encoding="utf-8").strip() or "1" exit_code = int(exit_code_raw) + initiating_gateway_pid = pending.get("initiating_gateway_pid") + if ( + exit_code == 0 + and initiating_gateway_pid is not None + and str(initiating_gateway_pid) == str(os.getpid()) + ): + logger.info( + "Update completion deferred until the gateway restarts " + "(initiating pid=%s)", + initiating_gateway_pid, + ) + cleanup = False + active_pending_path = pending_path + claimed_path.replace(pending_path) + return False + # Read the captured update output output = "" if output_path.exists(): @@ -21080,18 +21144,47 @@ async def _send_update_notification(self) -> bool: if len(output) > 3500: output = "…" + output[-3500:] if exit_code == 0: - msg = f"✅ Hermes update finished.\n\n```\n{output}\n```" + msg = ( + "✅ Hermes update finished.\n" + "♻️ Gateway restarted and is back online.\n\n" + f"```\n{output}\n```" + ) else: msg = f"❌ Hermes update failed.\n\n```\n{output}\n```" elif exit_code == 0: - msg = "✅ Hermes update finished successfully." + msg = ( + "✅ Hermes update finished successfully.\n" + "♻️ Gateway restarted and is back online." + ) else: msg = "❌ Hermes update failed. Check the gateway logs or run `hermes update` manually for details." - await adapter.send( - chat_id, - msg, - metadata=_non_conversational_metadata(metadata, platform=platform), - ) + try: + send_result = await adapter.send( + chat_id, + msg, + metadata=_non_conversational_metadata(metadata, platform=platform), + ) + except Exception as send_error: + # A transient platform reconnect or network failure must + # not destroy the only durable completion target. Restore + # the claim so the startup watcher can retry delivery. + logger.warning("Post-update notification send failed: %s", send_error) + cleanup = False + active_pending_path = pending_path + claimed_path.replace(pending_path) + return False + if ( + send_result is not None + and getattr(send_result, "success", True) is False + ): + logger.warning( + "Post-update notification was not delivered: %s", + getattr(send_result, "error", "send returned success=False"), + ) + cleanup = False + active_pending_path = pending_path + claimed_path.replace(pending_path) + return False logger.info( "Sent post-update notification to %s:%s (exit=%s)", platform_str, diff --git a/gateway/slash_commands.py b/gateway/slash_commands.py index c65acc814f5a0..c219225305a99 100644 --- a/gateway/slash_commands.py +++ b/gateway/slash_commands.py @@ -5590,6 +5590,11 @@ async def _handle_update_command(self, event: MessageEvent) -> str: "user_id": event.source.user_id, "session_key": session_key, "timestamp": datetime.now().isoformat(), + # Lets the completion watcher distinguish the gateway that + # initiated the update from the freshly restarted process. The + # old process must not consume the final notification marker just + # before systemd stops it. + "initiating_gateway_pid": os.getpid(), } if event.source.thread_id: pending["thread_id"] = event.source.thread_id diff --git a/tests/gateway/test_update_command.py b/tests/gateway/test_update_command.py index a56dec11d80d6..a0c138b945919 100644 --- a/tests/gateway/test_update_command.py +++ b/tests/gateway/test_update_command.py @@ -5,13 +5,14 @@ """ import json +import os from pathlib import Path from unittest.mock import patch, MagicMock, AsyncMock import pytest from gateway.config import Platform -from gateway.platforms.base import MessageEvent +from gateway.platforms.base import MessageEvent, SendResult from gateway.session import SessionSource @@ -132,6 +133,7 @@ async def test_writes_pending_marker(self, tmp_path): assert data["chat_id"] == "99999" assert data["chat_type"] == "dm" assert data["message_id"] == "m-update" + assert data["initiating_gateway_pid"] == os.getpid() assert "timestamp" in data assert not (hermes_home / ".update_exit_code").exists() @@ -347,8 +349,8 @@ async def test_sends_notification_with_output(self, tmp_path): @pytest.mark.asyncio - async def test_cleans_up_on_error(self, tmp_path): - """Files are cleaned up even if notification fails.""" + async def test_preserves_markers_when_notification_send_fails(self, tmp_path): + """A transient send failure remains retryable after restart.""" runner = _make_runner() hermes_home = tmp_path / "hermes" hermes_home.mkdir() @@ -368,12 +370,95 @@ async def test_cleans_up_on_error(self, tmp_path): runner.adapters = {Platform.TELEGRAM: mock_adapter} with patch("gateway.run._hermes_home", hermes_home): - await runner._send_update_notification() + result = await runner._send_update_notification() + + assert result is False + assert pending_path.exists() + assert output_path.exists() + assert exit_code_path.exists() + assert not (hermes_home / ".update_pending.claimed.json").exists() + + + @pytest.mark.asyncio + async def test_preserves_markers_when_notification_returns_failure(self, tmp_path): + """SendResult(success=False) is a failed delivery, not a completed send.""" + runner = _make_runner() + hermes_home = tmp_path / "hermes" + hermes_home.mkdir() + pending_path = hermes_home / ".update_pending.json" + output_path = hermes_home / ".update_output.txt" + exit_code_path = hermes_home / ".update_exit_code" + pending_path.write_text(json.dumps({ + "platform": "telegram", "chat_id": "111", "user_id": "222", + })) + output_path.write_text("✓ Done") + exit_code_path.write_text("0") + mock_adapter = AsyncMock() + mock_adapter.send.return_value = SendResult( + success=False, + error="Telegram temporarily unavailable", + ) + runner.adapters = {Platform.TELEGRAM: mock_adapter} + + with patch("gateway.run._hermes_home", hermes_home): + result = await runner._send_update_notification() + + assert result is False + assert pending_path.exists() + assert output_path.exists() + assert exit_code_path.exists() + assert not (hermes_home / ".update_pending.claimed.json").exists() - # Files should still be cleaned up (finally block) + + @pytest.mark.asyncio + async def test_initiating_gateway_defers_success_until_restart(self, tmp_path): + """The old process cannot consume the new process's completion marker.""" + runner = _make_runner() + hermes_home = tmp_path / "hermes" + hermes_home.mkdir() + pending_path = hermes_home / ".update_pending.json" + pending_path.write_text(json.dumps({ + "platform": "telegram", + "chat_id": "111", + "initiating_gateway_pid": 101, + })) + (hermes_home / ".update_exit_code").write_text("0") + runner.adapters = {Platform.TELEGRAM: AsyncMock()} + + with patch("gateway.run._hermes_home", hermes_home), \ + patch("gateway.run.os.getpid", return_value=101): + result = await runner._send_update_notification() + + assert result is False + assert pending_path.exists() + runner.adapters[Platform.TELEGRAM].send.assert_not_called() + + + @pytest.mark.asyncio + async def test_restarted_gateway_sends_final_online_notification(self, tmp_path): + """A new PID delivers the definitive completion and then cleans up.""" + runner = _make_runner() + hermes_home = tmp_path / "hermes" + hermes_home.mkdir() + pending_path = hermes_home / ".update_pending.json" + pending_path.write_text(json.dumps({ + "platform": "telegram", + "chat_id": "111", + "initiating_gateway_pid": 101, + })) + (hermes_home / ".update_exit_code").write_text("0") + adapter = AsyncMock() + runner.adapters = {Platform.TELEGRAM: adapter} + + with patch("gateway.run._hermes_home", hermes_home), \ + patch("gateway.run.os.getpid", return_value=202): + result = await runner._send_update_notification() + + assert result is True + message = adapter.send.await_args.args[1] + assert "Gateway restarted and is back online" in message assert not pending_path.exists() - assert not output_path.exists() - assert not exit_code_path.exists() + assert not (hermes_home / ".update_exit_code").exists() @pytest.mark.asyncio diff --git a/tests/gateway/test_update_streaming.py b/tests/gateway/test_update_streaming.py index 1e1134b7ad05c..a3867c0adacf2 100644 --- a/tests/gateway/test_update_streaming.py +++ b/tests/gateway/test_update_streaming.py @@ -16,7 +16,7 @@ import pytest from gateway.config import Platform -from gateway.platforms.base import MessageEvent +from gateway.platforms.base import MessageEvent, SendResult from gateway.session import SessionSource @@ -169,6 +169,72 @@ async def test_spawns_with_gateway_flag(self, tmp_path): class TestWatchUpdateProgress: """Tests for _watch_update_progress() streaming output.""" + @pytest.mark.asyncio + async def test_success_from_initiating_process_preserves_restart_marker(self, tmp_path): + """The watcher reports restart intent but leaves final delivery to the new PID.""" + runner = _make_runner() + hermes_home = tmp_path / "hermes" + hermes_home.mkdir() + pending_path = hermes_home / ".update_pending.json" + pending_path.write_text(json.dumps({ + "platform": "telegram", + "chat_id": "111", + "session_key": "agent:main:telegram:dm:111", + "initiating_gateway_pid": 101, + })) + output_path = hermes_home / ".update_output.txt" + output_path.write_text("✓ Code updated!\n") + exit_path = hermes_home / ".update_exit_code" + exit_path.write_text("0") + adapter = AsyncMock() + runner.adapters = {Platform.TELEGRAM: adapter} + + with patch("gateway.run._hermes_home", hermes_home), \ + patch("gateway.run.os.getpid", return_value=101): + await runner._watch_update_progress( + poll_interval=0.01, + stream_interval=0.01, + timeout=1.0, + ) + + sent = " ".join(str(call) for call in adapter.send.await_args_list) + assert "Restarting gateway" in sent + assert pending_path.exists() + assert output_path.exists() + assert exit_path.exists() + + @pytest.mark.asyncio + async def test_final_notification_retries_send_result_failure(self, tmp_path): + """A soft adapter failure must not consume legacy completion markers.""" + runner = _make_runner() + hermes_home = tmp_path / "hermes" + hermes_home.mkdir() + pending_path = hermes_home / ".update_pending.json" + pending_path.write_text(json.dumps({ + "platform": "telegram", + "chat_id": "111", + "session_key": "agent:main:telegram:dm:111", + })) + exit_path = hermes_home / ".update_exit_code" + exit_path.write_text("0") + adapter = AsyncMock() + adapter.send.side_effect = [ + SendResult(success=False, error="temporary failure"), + SendResult(success=True), + ] + runner.adapters = {Platform.TELEGRAM: adapter} + + with patch("gateway.run._hermes_home", hermes_home): + await runner._watch_update_progress( + poll_interval=0.01, + stream_interval=0.01, + timeout=1.0, + ) + + assert adapter.send.await_count == 2 + assert not pending_path.exists() + assert not exit_path.exists() + @pytest.mark.asyncio async def test_streams_output_to_adapter(self, tmp_path): """New output is sent to the adapter periodically.""" @@ -397,4 +463,3 @@ def fake_input(prompt, default=""): assert len(calls) == 1 assert "Restore" in calls[0] -