refactor(gateway): cross-mixin task tracking, retry-entry/hook-registration/notice-key helpers, update-target send

This commit is contained in:
Teknium
2026-09-02 22:24:39 -07:00
parent 8d3c6b6c7d
commit 5cd6d5e752
4 changed files with 170 additions and 206 deletions
+47 -48
View File
@@ -188,13 +188,23 @@ class GatewayAdapterLifecycleMixin:
tasks = getattr(self, "_fatal_handler_tasks", None)
if tasks is None:
tasks = self._fatal_handler_tasks = set()
task = asyncio.create_task(self._handle_adapter_fatal_error_detached(adapter))
tasks.add(task)
task.add_done_callback(tasks.discard)
task = self._track_task_in(tasks, asyncio.create_task(self._handle_adapter_fatal_error_detached(adapter)))
# shield(): a plain `await task` would tunnel the caller's cancellation into the detached
# task; with shield the caller sees CancelledError and the handler runs to completion.
await asyncio.shield(task)
def _reconnect_queue_entry(
self, platform, adapter, platform_config, *, attempts: int, delay: float, queued: bool = True
) -> dict:
"""Build a ``_failed_platforms`` entry (startup failures and runtime fatals share the shape)."""
now = time.monotonic()
return {
"config": platform_config, "attempts": attempts, "next_retry": now + delay,
**({"queued_at": now} if queued else {}),
"credential_claim": self._adapter_credential_claim(platform, adapter),
"listener_claim": self._adapter_listener_claim(platform, adapter),
}
def _queue_retryable_fatal_platform(self, adapter: BasePlatformAdapter) -> bool:
"""Queue a retryable fatal adapter for background reconnection (True when newly queued).
@@ -211,14 +221,9 @@ class GatewayAdapterLifecycleMixin:
# permanent outage — nothing retries and the stranded check treats "queued" as safe.
self._ensure_reconnect_watcher_running()
return False
self._failed_platforms[adapter.platform] = {
"config": platform_config,
"attempts": 0,
"next_retry": time.monotonic(),
"queued_at": time.monotonic(),
"credential_claim": self._adapter_credential_claim(adapter.platform, adapter),
"listener_claim": self._adapter_listener_claim(adapter.platform, adapter),
}
self._failed_platforms[adapter.platform] = self._reconnect_queue_entry(
adapter.platform, adapter, platform_config, attempts=0, delay=0.0,
)
logger.info("%s queued for background reconnection", adapter.platform.value)
self._ensure_reconnect_watcher_running()
return True
@@ -341,6 +346,13 @@ class GatewayAdapterLifecycleMixin:
task.add_done_callback(tasks.discard)
return task
@staticmethod
def _track_task_in(tasks: set, task: "asyncio.Task") -> "asyncio.Task":
"""Register ``task`` in an arbitrary lifecycle set with self-removal on completion."""
tasks.add(task)
task.add_done_callback(tasks.discard)
return task
def _request_clean_exit(self, reason: str) -> None:
self._exit_cleanly = True
self._exit_reason = reason
@@ -640,6 +652,13 @@ class GatewayAdapterLifecycleMixin:
retrying_since=retrying_since_iso,
)
def _mark_platform_fatal(self, status_key: str, adapter) -> None:
"""Record an adapter's fatal error code/message as ``fatal`` runtime status."""
self._update_platform_runtime_status(
status_key, platform_state="fatal", error_code=adapter.fatal_error_code,
error_message=adapter.fatal_error_message,
)
def _bump_reconnect_backoff(
self, platform, info: dict, attempt: int, error_code, error_message: str
) -> int:
@@ -687,10 +706,7 @@ class GatewayAdapterLifecycleMixin:
if success:
await self._install_reconnected_adapter(platform, adapter)
elif adapter.has_fatal_error and not adapter.fatal_error_retryable:
self._update_platform_runtime_status(
platform.value, platform_state="fatal", error_code=adapter.fatal_error_code,
error_message=adapter.fatal_error_message,
)
self._mark_platform_fatal(platform.value, adapter)
logger.warning(
"Reconnect %s: non-retryable error (%s), removing from retry queue",
platform.value, adapter.fatal_error_message,
@@ -779,8 +795,7 @@ class GatewayAdapterLifecycleMixin:
_done, unfinished = await asyncio.wait(tasks, timeout=timeout)
if unfinished:
logger.warning(
"Timed out waiting for %d secondary profile reconnect task(s) during shutdown",
len(unfinished),
"Timed out waiting for %d secondary profile reconnect task(s) during shutdown", len(unfinished),
)
pending.clear()
@@ -880,18 +895,9 @@ class GatewayAdapterLifecycleMixin:
# Register this profile's shell hooks / outbound webhooks: start() registers before any
# profile scope exists, so a secondary profile's `hooks:` block would be silently inert.
try:
from hermes_cli.config import load_config as _load_profile_config
from agent.shell_hooks import register_from_config as _register_shell_hooks
from agent.outbound_webhooks import (register_from_config as _register_outbound_webhooks)
_profile_hooks_cfg = _load_profile_config()
_register_shell_hooks(_profile_hooks_cfg, accept_hooks=False)
_register_outbound_webhooks(_profile_hooks_cfg)
except Exception:
logger.warning(
"shell-hook/webhook registration failed for profile '%s'", profile_name, exc_info=True,
)
self._register_config_hooks(
"shell-hook/webhook registration failed for profile '%s'", profile_name, level=logging.WARNING,
)
profile_cfg = load_gateway_config()
violation = _own_policy_open_startup_violation(profile_cfg)
@@ -929,31 +935,27 @@ class GatewayAdapterLifecycleMixin:
owner = claimed.get(claim) if claim is not None else None
if owner is None:
return False
pv = platform.value
if kind == "credential":
message = (
f"Profile '{owner}' and '{profile_name}' both configure "
f"{platform.value} with the same credential. Give each "
f"profile its own {platform.value} credential."
f"Profile '{owner}' and '{profile_name}' both configure {pv} with the same credential. "
f"Give each profile its own {pv} credential."
)
logger.error(
"Profile '%s' and '%s' both configure %s with the same "
"credential — refusing to start the duplicate (one "
"credential cannot be consumed twice). Give each profile "
"its own %s credential.", owner, profile_name, platform.value, platform.value,
"Profile '%s' and '%s' both configure %s with the same credential — refusing to start the "
"duplicate (one credential cannot be consumed twice). Give each profile its own %s credential.",
owner, profile_name, pv, pv,
)
else:
bind, port = claim[-2:]
message = (
f"Profile '{owner}' and '{profile_name}' both configure "
f"{platform.value} sidecars on the same listener. Configure "
f"a distinct listener for profile '{profile_name}'."
f"Profile '{owner}' and '{profile_name}' both configure {pv} sidecars on the same listener. "
f"Configure a distinct listener for profile '{profile_name}'."
)
logger.error(
"Profile '%s' and '%s' both configure %s sidecars on "
"%s:%s — refusing to start the duplicate listener. "
"Set platforms.%s.extra.sidecar_port to a distinct port "
"for profile '%s'.",
owner, profile_name, platform.value, bind, port, platform.value, profile_name,
"Profile '%s' and '%s' both configure %s sidecars on %s:%s — refusing to start the duplicate "
"listener. Set platforms.%s.extra.sidecar_port to a distinct port for profile '%s'.",
owner, profile_name, pv, bind, port, pv, profile_name,
)
self._update_platform_runtime_status(
f"{profile_name}:{platform.value}", platform_state="fatal",
@@ -1204,10 +1206,7 @@ class GatewayAdapterLifecycleMixin:
"gateway (%s) — parked, not retried. %s", profile_name, platform.value,
adapter.fatal_error_code, adapter.fatal_error_message or "",
)
self._update_platform_runtime_status(
f"{profile_name}:{platform.value}", platform_state="fatal",
error_code=adapter.fatal_error_code, error_message=adapter.fatal_error_message,
)
self._mark_platform_fatal(f"{profile_name}:{platform.value}", adapter)
return
def _handoff() -> None:
+59 -78
View File
@@ -14,14 +14,12 @@ import logging
import time
from contextlib import suppress
from pathlib import Path
from typing import TYPE_CHECKING, Any, Dict, Optional, cast
from typing import Any, Dict, Optional, cast
from gateway.config import Platform, _BUILTIN_PLATFORM_VALUES
from gateway.platforms.base import MessageEvent, MessageType
from gateway.session import SessionEntry, SessionSource
if TYPE_CHECKING: # string annotations only; never imported at runtime (cycle)
from gateway.run import GatewayRunner, TurnRunner # noqa: F401
from gateway.run_shutdown import _notice_target_key, _send_error, _send_failed
# Log-record parity with the origin module.
logger = logging.getLogger("gateway.run")
@@ -78,6 +76,9 @@ class GatewayNotificationsMixin:
from gateway.run import _non_conversational_metadata
return _non_conversational_metadata(self.metadata, platform=self.platform)
async def send(self, text: str):
return await self.adapter.send(self.chat_id, text, metadata=self.send_metadata())
@dataclasses.dataclass
class _CompletionClaim:
"""Pre-flight outcome for one completion delivery."""
@@ -99,11 +100,10 @@ class GatewayNotificationsMixin:
if config and getattr(source, "platform", None) == Platform.SLACK and _is_slack_ignored_channel(config, chat_id):
logger.info("Skipping Slack platform notice for configured ignored channel %s", chat_id)
return
notice_delivery = "public"
if config and hasattr(config, "get_notice_delivery"):
notice_delivery = config.get_notice_delivery(source.platform)
notice_delivery = (
config.get_notice_delivery(source.platform) if config and hasattr(config, "get_notice_delivery")
else "public"
)
metadata = self._thread_metadata_for_source(source)
if notice_delivery == "private" and getattr(source, "user_id", None):
try:
@@ -400,12 +400,9 @@ class GatewayNotificationsMixin:
return None
platform = Platform(platform_str)
adapter = self.adapters.get(platform)
metadata = self._thread_metadata_for_target(
platform, chat_id, pending.get("thread_id"), chat_type=pending.get("chat_type"),
reply_to_message_id=pending.get("message_id"), adapter=adapter,
)
if not adapter:
return None
metadata = self._pending_marker_metadata(platform, chat_id, pending, adapter)
# Fallback session key if not stored (old pending files)
return self._UpdateTarget(
adapter, chat_id, session_key or f"{platform_str}:{chat_id}", metadata, platform,
@@ -414,6 +411,13 @@ class GatewayNotificationsMixin:
pass
return None
def _pending_marker_metadata(self, platform, chat_id, data: dict, adapter):
"""Thread metadata for a persisted update/restart marker (thread_id/chat_type/message_id keys)."""
return self._thread_metadata_for_target(
platform, chat_id, data.get("thread_id"), chat_type=data.get("chat_type"),
reply_to_message_id=data.get("message_id"), adapter=adapter,
)
async def _watch_update_completion_only(self, paths: "_UpdatePaths", deadline: float, poll_interval: float) -> None:
"""Fallback when no adapter/chat can be resolved: wait for the exit code, then notify."""
logger.warning("Update watcher: cannot resolve adapter/chat_id, falling back to completion-only")
@@ -448,9 +452,7 @@ class GatewayNotificationsMixin:
max_chunk = 3500
for i in range(0, len(clean), max_chunk):
try:
await target.adapter.send(
target.chat_id, f"```\n{clean[i:i + max_chunk]}\n```", metadata=target.send_metadata(),
)
await target.send(f"```\n{clean[i:i + max_chunk]}\n```")
except Exception as e:
logger.debug("Update stream send failed: %s", e)
@@ -470,11 +472,9 @@ class GatewayNotificationsMixin:
if not sent_buttons:
default_hint = f" (default: {default})" if default else ""
_p = getattr(adapter, "typed_command_prefix", "/")
await adapter.send(
target.chat_id,
await target.send(
f"⚕ **Update needs your input:**\n\n{prompt_text}{default_hint}\n\n"
f"Reply `{_p}approve` (yes) or `{_p}deny` (no), or type your answer directly.",
metadata=target.send_metadata(),
f"Reply `{_p}approve` (yes) or `{_p}deny` (no), or type your answer directly."
)
# Keep the prompt marker on disk until the user answers so a watcher after a mid-prompt
# gateway restart can recover by re-forwarding it.
@@ -531,11 +531,9 @@ class GatewayNotificationsMixin:
await _flush_buffer()
try:
exit_code = int(paths.exit_code.read_text(encoding="utf-8").strip() or "1")
await target.adapter.send(
target.chat_id,
await target.send(
"✅ Hermes update finished." if exit_code == 0
else "❌ Hermes update failed (exit code {}).".format(exit_code),
metadata=target.send_metadata(),
else "❌ Hermes update failed (exit code {}).".format(exit_code)
)
logger.info("Update finished (exit=%s), notified %s", exit_code, session_key)
except Exception as e:
@@ -570,9 +568,7 @@ class GatewayNotificationsMixin:
paths.exit_code.write_text("124", encoding="utf-8")
await _flush_buffer()
with suppress(Exception):
await target.adapter.send(
target.chat_id, "❌ Hermes update timed out after 30 minutes.", metadata=target.send_metadata(),
)
await target.send("❌ Hermes update timed out after 30 minutes.")
self._clear_update_markers(paths, session_key)
async def _send_update_notification(self) -> bool:
@@ -626,10 +622,7 @@ class GatewayNotificationsMixin:
return _defer("Update notification deferred: %s adapter not connected yet", platform_str)
if adapter and chat_id:
metadata = self._thread_metadata_for_target(
platform, chat_id, pending.get("thread_id"), chat_type=pending.get("chat_type"),
reply_to_message_id=pending.get("message_id"), adapter=adapter,
)
metadata = self._pending_marker_metadata(platform, chat_id, pending, adapter)
from tools.ansi_strip import strip_ansi
output = strip_ansi(output).strip()
if output:
@@ -682,10 +675,7 @@ class GatewayNotificationsMixin:
)
return None
metadata = self._thread_metadata_for_target(
platform, chat_id, thread_id, chat_type=data.get("chat_type"),
reply_to_message_id=data.get("message_id"), adapter=transport.adapter,
)
metadata = self._pending_marker_metadata(platform, chat_id, data, transport.adapter)
if data.get("delivered_via_upstream_relay") is True:
metadata = dict(metadata or {})
for field in ("user_id", "scope_id"):
@@ -697,10 +687,9 @@ class GatewayNotificationsMixin:
)
# adapter.send() catches provider errors (e.g. "Chat not found") and returns
# SendResult(success=False) rather than raising, so inspect the result before claiming success.
if result is not None and getattr(result, "success", True) is False:
if _send_failed(result):
logger.warning(
"Restart notification to %s:%s was not delivered: %s",
platform_str, chat_id, getattr(result, "error", "send returned success=False"),
"Restart notification to %s:%s was not delivered: %s", platform_str, chat_id, _send_error(result),
)
return None
@@ -740,10 +729,8 @@ class GatewayNotificationsMixin:
result = await transport.send(platform, str(home.chat_id), message, metadata=send_metadata)
else:
result = await transport.adapter.send(str(home.chat_id), message)
if result is not None and getattr(result, "success", True) is False:
logger.warning(
failure_fmt, platform.value, home.chat_id, getattr(result, "error", "send returned success=False"),
)
if _send_failed(result):
logger.warning(failure_fmt, platform.value, home.chat_id, _send_error(result))
return False
return True
except Exception as exc:
@@ -770,7 +757,7 @@ class GatewayNotificationsMixin:
)
continue
target = (platform.value, str(home.chat_id), str(home.thread_id) if home.thread_id else None)
target = _notice_target_key(platform.value, home.chat_id, home.thread_id)
if target in skipped or target in delivered:
continue
@@ -921,22 +908,19 @@ class GatewayNotificationsMixin:
"""
from gateway.wake import deliver_wake, persist_delegation_delivery
if evt.get("type") == "async_delegation":
try:
logger.info(
"Async delegation completion — persisting delivery row for api_server session %s (no wake turn)",
raw_sid,
)
await persist_delegation_delivery(adapter, text=synth_text, session_id=raw_sid, evt=evt)
return True
except Exception as e:
logger.warning("Async delegation delivery persist failed for session %s: %s", raw_sid, e)
return False
info = "Async delegation completion — persisting delivery row for api_server session %s (no wake turn)"
fail = "Async delegation delivery persist failed for session %s: %s"
deliver = lambda: persist_delegation_delivery(adapter, text=synth_text, session_id=raw_sid, evt=evt) # noqa: E731
else:
info = "Watch pattern notification — waking api_server session %s via self-post"
fail = "Watch notification self-post wake failed for session %s: %s"
deliver = lambda: deliver_wake(adapter, text=synth_text, session_id=raw_sid) # noqa: E731
try:
logger.info("Watch pattern notification — waking api_server session %s via self-post", raw_sid)
await deliver_wake(adapter, text=synth_text, session_id=raw_sid)
logger.info(info, raw_sid)
await deliver()
return True
except Exception as e:
logger.warning("Watch notification self-post wake failed for session %s: %s", raw_sid, e)
logger.warning(fail, raw_sid, e)
return False
def _resolve_injection_adapter(self, platform_name: str):
@@ -972,10 +956,9 @@ class GatewayNotificationsMixin:
# API-server sessions bind the RAW X-Hermes-Session-Id key (_bind_api_server_session), not a
# structured ``agent:main:...`` key, so _build_process_event_source returned None above.
raw_sid = str(evt.get("origin_session_id") or "").strip()
if not raw_sid:
_sk = str(evt.get("session_key") or "").strip()
if _sk and _parse_session_key(_sk) is None:
raw_sid = _sk
_sk = str(evt.get("session_key") or "").strip()
if not raw_sid and _sk and _parse_session_key(_sk) is None:
raw_sid = _sk
if raw_sid:
adapter = self.adapters.get(Platform.API_SERVER)
if adapter is not None and not adapter_supports_push(adapter):
@@ -1297,13 +1280,17 @@ class GatewayNotificationsMixin:
finally:
# Never strand watcher futures when formatting, delivery, or cancellation interrupts a batch:
# False follows the existing watcher retry path; None remains the ordinary dedupe result.
for _text, _evt, future in entries:
if not future.done():
future.set_result(delivered)
self._settle_batch_waiters(entries, delivered)
# Do not remove a newer flush task that reused the same route key.
if self._completion_notification_batch_tasks.get(key) is current_task:
self._completion_notification_batch_tasks.pop(key, None)
@staticmethod
def _settle_batch_waiters(entries, result) -> None:
for _text, _evt, future in entries:
if not future.done():
future.set_result(result)
async def _cancel_process_completion_batch_tasks(self) -> None:
"""Settle pending completion batches before adapter teardown."""
self._completion_notification_batches_stopping = True
@@ -1320,9 +1307,7 @@ class GatewayNotificationsMixin:
# Defensive cleanup for an orphaned queue with no live flush task.
batches = getattr(self, "_completion_notification_batches", {})
for entries in batches.values():
for _text, _evt, future in entries:
if not future.done():
future.set_result(False)
self._settle_batch_waiters(entries, False)
batches.clear()
getattr(self, "_completion_notification_batch_tasks", {}).clear()
getattr(self, "_completion_notification_batch_flush_tasks", set()).clear()
@@ -1352,10 +1337,8 @@ class GatewayNotificationsMixin:
task = asyncio.create_task(self._flush_process_completion_batch(key))
self._completion_notification_batch_tasks[key] = task
# Keep the flush alive under the gateway's normal lifecycle accounting.
self._background_tasks.add(task)
self._completion_notification_batch_flush_tasks.add(task)
task.add_done_callback(self._background_tasks.discard)
task.add_done_callback(self._completion_notification_batch_flush_tasks.discard)
self._retain_background_task(task)
self._track_task_in(self._completion_notification_batch_flush_tasks, task)
return await future
def _enrich_async_delegation_routing(self, evt: dict) -> None:
@@ -1642,11 +1625,9 @@ class GatewayNotificationsMixin:
"via wait/log — skipping raw notification (#65379)", session_id,
)
break
should_notify = (
notify_mode in {"concise", "all", "result"}
or (notify_mode == "error" and session.exit_code not in {0, None})
)
if should_notify:
if notify_mode in {"concise", "all", "result"} or (
notify_mode == "error" and session.exit_code not in {0, None}
):
message_text = self._format_process_final_message(session_id, session, notify_mode)
await self._send_watcher_message(platform_name, chat_id, thread_id, message_text)
break
@@ -1655,9 +1636,9 @@ class GatewayNotificationsMixin:
# New output — deliver a status update (only in "all" mode; agent_notify watchers
# only care about completion).
new_output = self._redacted_output_tail(session, 500)
message_text = (
f"[Background process {session_id} is still running~ New output:\n{new_output}]"
await self._send_watcher_message(
platform_name, chat_id, thread_id,
f"[Background process {session_id} is still running~ New output:\n{new_output}]",
)
await self._send_watcher_message(platform_name, chat_id, thread_id, message_text)
logger.debug("Process watcher ended: %s", session_id)
+45 -54
View File
@@ -18,7 +18,7 @@ import threading
import time
from contextlib import suppress
from pathlib import Path
from typing import TYPE_CHECKING, Any, Callable, Dict, Optional
from typing import Any, Callable, Dict, Optional
from gateway.config import Platform
from gateway.restart import (
@@ -27,9 +27,6 @@ from gateway.restart import (
from gateway.run_common import _UNSET
from gateway.shutdown_watchdog import arm_shutdown_watchdog, resolve_shutdown_watchdog_delay
if TYPE_CHECKING: # string annotations only; never imported at runtime (cycle)
from gateway.run import GatewayRunner, TurnRunner # noqa: F401
# Log-record parity with the origin module.
logger = logging.getLogger("gateway.run")
@@ -88,6 +85,16 @@ def _send_failed(result: Any) -> bool:
return result is not None and getattr(result, "success", True) is False
def _send_error(result: Any) -> str:
"""Error text of a failed ``send()`` result (adapters may omit it)."""
return getattr(result, "error", "send returned success=False")
def _notice_target_key(platform_value: str, chat_id, thread_id) -> tuple:
"""Dedup key for one notice destination: thread/topic platforms share a chat but route apart."""
return (platform_value, str(chat_id), str(thread_id) if thread_id else None)
class GatewayShutdownMixin:
"""Stop/drain/restart, scale-to-zero and active-work accounting methods for GatewayRunner."""
@@ -321,17 +328,14 @@ class GatewayShutdownMixin:
from cron.scheduler import get_running_job_ids
return len(get_running_job_ids())
def _dashboard_seen():
from gateway.scale_to_zero import dashboard_client_last_seen
return dashboard_client_last_seen()
cron_count = _read_or_awake("cron work count", _cron_count, 1)
api_count = _read_or_awake("api work count", lambda: self._api_server_hook("active_agent_work_count"), 1)
# An attached dashboard/desktop/TUI client (heartbeat file mtime from the dashboard process) is
# inbound activity too — folded into the inbound clock, not a conjunct, so disconnect gets the
# same idle_timeout grace as a message and a lingering marker cannot pin the box.
last_inbound = self._last_inbound_at
seen = _read_or_awake("dashboard heartbeat", _dashboard_seen, time.time())
from gateway.scale_to_zero import dashboard_client_last_seen
seen = _read_or_awake("dashboard heartbeat", dashboard_client_last_seen, time.time())
if seen is not None and seen > last_inbound:
last_inbound = seen
return is_idle(
@@ -524,10 +528,9 @@ class GatewayShutdownMixin:
info["pause_reason"] = reason or "auto-paused after repeated failures"
# Push next_retry far enough out that a stale code path missing "paused" still never fires.
info["next_retry"] = float("inf")
with suppress(Exception):
self._update_platform_runtime_status(
platform.value, platform_state="paused", error_code=None, error_message=info["pause_reason"],
)
self._update_platform_runtime_status(
platform.value, platform_state="paused", error_code=None, error_message=info["pause_reason"],
)
logger.warning(
"%s paused after %d consecutive failures (%s) — "
"fix the underlying issue then run `/platform resume %s` "
@@ -544,10 +547,7 @@ class GatewayShutdownMixin:
info.pop("pause_reason", None)
info["attempts"] = 0
info["next_retry"] = time.monotonic() # retry on next watcher tick
with suppress(Exception):
self._update_platform_runtime_status(
platform.value, platform_state="retrying", error_code=None, error_message=None,
)
self._update_platform_runtime_status(platform.value, platform_state="retrying")
logger.info("%s resumed — retrying on next watcher tick", platform.value)
return True
@@ -687,7 +687,7 @@ class GatewayShutdownMixin:
continue
chat_id = str(target.get("chat_id"))
thread_id = target.get("thread_id")
dedup_key = (job_id, platform.value, chat_id, str(thread_id) if thread_id else None)
dedup_key = (job_id, *_notice_target_key(platform.value, chat_id, thread_id))
if dedup_key in notified:
continue
try:
@@ -695,8 +695,7 @@ class GatewayShutdownMixin:
result = await adapter.send(chat_id, msg, metadata=metadata)
if _send_failed(result):
logger.debug(
"Cron interrupt notice to %s:%s failed: %s", platform.value, chat_id,
getattr(result, "error", "send returned success=False"),
"Cron interrupt notice to %s:%s failed: %s", platform.value, chat_id, _send_error(result),
)
continue
notified.add(dedup_key)
@@ -735,8 +734,8 @@ class GatewayShutdownMixin:
result = await adapter.send(chat_id, msg, **send_kwargs)
if _send_failed(result):
logger.debug(
"Failed to send shutdown notification to %s%s:%s: %s",
where, platform_str, chat_id, getattr(result, "error", "send returned success=False"),
"Failed to send shutdown notification to %s%s:%s: %s", where, platform_str, chat_id,
_send_error(result),
)
return False
logger.info("Sent shutdown notification to %s %s:%s", kind, platform_str, chat_id)
@@ -762,9 +761,8 @@ class GatewayShutdownMixin:
restart_key = None
if restart_source is not None:
with suppress(Exception):
restart_key = (
restart_source.platform.value, str(restart_source.chat_id),
str(restart_source.thread_id) if restart_source.thread_id else None,
restart_key = _notice_target_key(
restart_source.platform.value, restart_source.chat_id, restart_source.thread_id
)
notified: set[tuple[str, str, Optional[str]]] = set()
for session_key in self._snapshot_running_agents():
@@ -772,9 +770,7 @@ class GatewayShutdownMixin:
if target is None:
continue
source, platform_str, chat_id, thread_id = target
# Dedupe only identical targets: thread/topic platforms share a parent chat yet route to
# distinct destinations via metadata.
dedup_key = (platform_str, chat_id, str(thread_id) if thread_id else None)
dedup_key = _notice_target_key(platform_str, chat_id, thread_id)
if dedup_key in notified:
continue
try:
@@ -830,7 +826,7 @@ class GatewayShutdownMixin:
platform.value,
)
continue
dedup_key = (platform.value, str(home.chat_id), str(home.thread_id) if home.thread_id else None)
dedup_key = _notice_target_key(platform.value, home.chat_id, home.thread_id)
if dedup_key in notified:
continue
try:
@@ -921,12 +917,10 @@ class GatewayShutdownMixin:
await self._cleanup_agent_resources_off_loop(agent, context=context)
self._track_deferred_agent_worker(future, agent)
task = asyncio.create_task(_cleanup_when_done())
tasks = getattr(self, "_deferred_agent_cleanup_tasks", None)
if tasks is None:
tasks = self._deferred_agent_cleanup_tasks = set()
tasks.add(task)
task.add_done_callback(tasks.discard)
self._track_task_in(tasks, asyncio.create_task(_cleanup_when_done()))
async def _finalize_session_off_loop(self, *, session_id: Any, platform: str, reason: str, **extra: Any) -> None:
"""Run hermes_cli.lifecycle.finalize_session off-loop, bounded; on timeout the worker is left alone."""
@@ -1213,8 +1207,7 @@ class GatewayShutdownMixin:
"deferring stop() until they finish (cap=%.0fs) so in-flight "
"turns are not amputated (#77184)", active, timeout,
)
with suppress(Exception):
self._update_runtime_status("draining")
self._scale_to_zero_status("draining", "restart wait: status mark failed")
loop = asyncio.get_running_loop()
deadline = loop.time() + timeout
last_status_at = 0.0
@@ -1233,8 +1226,7 @@ class GatewayShutdownMixin:
"(%d wedged and excluded; %.0fs remaining before force drain)",
self._awaitable_work_count(), self._wedged_agent_count(), deadline - now,
)
with suppress(Exception):
self._update_runtime_status("draining")
self._scale_to_zero_status("draining", "restart wait: status mark failed")
last_status_at = now
await asyncio.sleep(0.1)
if self._active_work_count() > 0:
@@ -1298,6 +1290,15 @@ class GatewayShutdownMixin:
# stop() phases. Invoked as ``GatewayRunner._stop_<phase>(self, ctx)`` so shutdown-path tests
# can drive them from bare doubles that are not GatewayRunner instances.
@staticmethod
def _quiet_step(label: str, fn: Callable[[], Any]) -> Any:
"""Run one best-effort teardown step; a failure is debug-logged as ``"<label>: <exc>"``."""
try:
return fn()
except Exception as _e:
logger.debug("%s: %s", label, _e)
return None
@staticmethod
def _stop_kill_tool_subprocesses(phase: str) -> list:
"""Kill tool subprocesses + terminal envs + browsers; returns cron job IDs marked interrupted.
@@ -1307,11 +1308,7 @@ class GatewayShutdownMixin:
"""
def _step(label: str, fn: Callable[[], Any]) -> Any:
try:
return fn()
except Exception as _e:
logger.debug("%s (%s) error: %s", label, phase, _e)
return None
return GatewayShutdownMixin._quiet_step(f"{label} ({phase}) error", fn)
def _kill_processes() -> None:
from tools.process_registry import process_registry
@@ -1510,12 +1507,12 @@ class GatewayShutdownMixin:
for platform, adapter in list(self.adapters.items()):
await self._bounded_adapter_teardown(adapter, platform)
# Disconnect secondary-profile adapters (multiplex mode).
for _prof, _amap in list(getattr(self, "_profile_adapters", {}).items()):
_profile_adapters = getattr(self, "_profile_adapters", {})
for _prof, _amap in list(_profile_adapters.items()):
for platform, adapter in list(_amap.items()):
await self._bounded_adapter_teardown(adapter, platform, profile=_prof)
_amap.clear()
if hasattr(self, "_profile_adapters"):
self._profile_adapters.clear()
_profile_adapters.clear()
logger.info("Shutdown phase: all adapters disconnected at +%.2fs", ctx.elapsed())
def _stop_release_runtime_state(self, ctx: "GatewayShutdownMixin._StopContext") -> None:
@@ -1561,11 +1558,11 @@ class GatewayShutdownMixin:
logger.info("Shutdown phase: final-cleanup tool kill done at +%.2fs", ctx.elapsed())
# Reap the process-global auxiliary-client cache: per-turn cleanup misses clients bound to
# worker-thread loops that died with their executor (cron ticks) → httpx transports leak to EMFILE.
try:
def _reap_aux_clients() -> None:
from agent.auxiliary_client import shutdown_cached_clients
shutdown_cached_clients()
except Exception as _e:
logger.debug("shutdown_cached_clients error: %s", _e)
GatewayShutdownMixin._quiet_step("shutdown_cached_clients error", _reap_aux_clients)
def _stop_quiesce_and_close_session_dbs(self, timeout: float, ctx: "GatewayShutdownMixin._StopContext") -> None:
"""Quiesce the executor, then close SessionDB handles only if no worker is still live."""
@@ -1591,13 +1588,7 @@ class GatewayShutdownMixin:
)
return
logger.info("Shutdown phase: executor quiesced at +%.2fs", ctx.elapsed())
def _step(label: str, fn: Callable[[], Any]) -> None:
try:
fn()
except Exception as _e:
logger.debug("%s: %s", label, _e)
_step = GatewayShutdownMixin._quiet_step
# Close SQLite session DBs so the WAL lock is released; otherwise --replace leaves the old
# connection holding it until exit and the new gateway gets 'database is locked'.
# ``_session_db`` is an AsyncSessionDB facade — unwrap; ``session_store`` holds ``_db``.
+19 -26
View File
@@ -149,11 +149,7 @@ class GatewayStartupMixin:
for task in pending:
task.add_done_callback(late)
if track:
bg = getattr(self, "_background_tasks", None)
if bg is None:
bg = self._background_tasks = set()
bg.add(task)
task.add_done_callback(bg.discard)
self._retain_background_task(task)
return done
async def _finish_startup_restore(self) -> None:
@@ -521,11 +517,9 @@ class GatewayStartupMixin:
# Empty-text internal event: the _is_resume_pending branch in _handle_message_with_agent
# prepends the reason-aware system note before the turn runs.
event = MessageEvent(text="", message_type=MessageType.TEXT, source=source, internal=True)
task = asyncio.create_task(
self._run_startup_resume_event(adapter, event, entry.session_key)
task = self._retain_background_task(
asyncio.create_task(self._run_startup_resume_event(adapter, event, entry.session_key))
)
self._background_tasks.add(task)
task.add_done_callback(self._background_tasks.discard)
if getattr(self, "_startup_restore_in_progress", False):
tasks = getattr(self, "_startup_restore_tasks", None)
if tasks is None:
@@ -548,13 +542,8 @@ class GatewayStartupMixin:
await self._safe_adapter_disconnect(adapter, platform)
def _startup_retry_entry(self, platform, adapter, platform_config, *, queued: bool = True) -> dict:
"""Build a ``_failed_platforms`` entry for a platform that failed at startup."""
return {
"config": platform_config, "attempts": 1, "next_retry": time.monotonic() + 30,
**({"queued_at": time.monotonic()} if queued else {}),
"credential_claim": self._adapter_credential_claim(platform, adapter),
"listener_claim": self._adapter_listener_claim(platform, adapter),
}
"""``_failed_platforms`` entry for a platform that failed at startup (first retry in 30s)."""
return self._reconnect_queue_entry(platform, adapter, platform_config, attempts=1, delay=30, queued=queued)
async def _abort_startup_if_shutdown_requested(
self, adapter: Optional[BasePlatformAdapter] = None, platform: Optional[Platform] = None
@@ -698,8 +687,7 @@ class GatewayStartupMixin:
task._hermes_supervised_watcher = True # type: ignore[attr-defined]
_bg = getattr(self, "_background_tasks", None)
if _bg is not None:
_bg.add(task)
task.add_done_callback(_bg.discard)
self._track_task_in(_bg, task)
except Exception:
logger.debug("Failed to start gateway loop heartbeat", exc_info=True)
@@ -917,18 +905,25 @@ class GatewayStartupMixin:
except Exception:
logger.warning("relay adapter registration failed at gateway startup", exc_info=True)
# Declarative shell hooks from cli-config.yaml. Gateway has no TTY, so consent must come
# from --accept-hooks, HERMES_ACCEPT_HOOKS, or hooks_auto_accept: true; pass
# accept_hooks=False and let register_from_config resolve env + config.
GatewayStartupMixin._register_config_hooks("shell-hook registration failed at gateway startup")
@staticmethod
def _register_config_hooks(fail_fmt: str, *fail_args, level: int = logging.DEBUG) -> None:
"""Register declarative shell hooks + outbound webhooks from the CURRENT scope's config.
Gateway has no TTY, so consent must come from --accept-hooks, HERMES_ACCEPT_HOOKS, or
hooks_auto_accept: true; ``accept_hooks=False`` lets register_from_config resolve env + config.
Never raises (logged at ``level``).
"""
try:
from hermes_cli.config import load_config
from agent.shell_hooks import register_from_config
from agent.outbound_webhooks import register_from_config as register_outbound_webhooks
_hooks_cfg = load_config()
register_from_config(_hooks_cfg, accept_hooks=False)
from agent.outbound_webhooks import register_from_config as register_outbound_webhooks
register_outbound_webhooks(_hooks_cfg)
except Exception:
logger.debug("shell-hook registration failed at gateway startup", exc_info=True)
logger.log(level, fail_fmt, *fail_args, exc_info=True)
async def _start_recover_previous_run(self) -> None:
"""Plugins, relay, hooks, then crash/clean-exit recovery of processes and sessions."""
@@ -1135,9 +1130,7 @@ class GatewayStartupMixin:
# blip — ``_acquire_platform_lock`` emits it retryable only so a MID-RUN reconnect can
# recover. At startup route it non-retryable: with nothing connected the gateway exits
# 78 instead of sitting alive and deaf in the retry queue.
_retryable = adapter.fatal_error_retryable and not (
is_global_startup_conflict(adapter.fatal_error_code)
)
_retryable = adapter.fatal_error_retryable and not is_global_startup_conflict(adapter.fatal_error_code)
self._update_platform_runtime_status(
platform.value, platform_state="retrying" if _retryable else "fatal",
error_code=adapter.fatal_error_code, error_message=adapter.fatal_error_message,