Merge branch 'simp/r2-gwrun' into simp/integration2
This commit is contained in:
24093
gateway/run.py
24093
gateway/run.py
File diff suppressed because it is too large
Load Diff
1984
gateway/run_adapters.py
Normal file
1984
gateway/run_adapters.py
Normal file
File diff suppressed because it is too large
Load Diff
1162
gateway/run_agent_cache.py
Normal file
1162
gateway/run_agent_cache.py
Normal file
File diff suppressed because it is too large
Load Diff
1523
gateway/run_busy.py
Normal file
1523
gateway/run_busy.py
Normal file
File diff suppressed because it is too large
Load Diff
8
gateway/run_common.py
Normal file
8
gateway/run_common.py
Normal file
@@ -0,0 +1,8 @@
|
||||
"""Leaf constants shared by ``gateway/run.py`` and its ``run_*`` mixin modules.
|
||||
|
||||
Kept import-cycle free (imports nothing from ``gateway.run``) because these values
|
||||
are used as default-argument sentinels, which must resolve at ``def`` time.
|
||||
"""
|
||||
|
||||
# Sentinel for "caller did not pass metadata" vs "caller passed None".
|
||||
_UNSET = object()
|
||||
620
gateway/run_config_loaders.py
Normal file
620
gateway/run_config_loaders.py
Normal file
@@ -0,0 +1,620 @@
|
||||
"""Config/env loaders for runtime knobs (busy modes, reasoning, service tier, timeouts, fallback) for GatewayRunner.
|
||||
|
||||
Split out of ``gateway/run.py``; bound onto ``GatewayRunner`` via the MRO.
|
||||
``gateway.run`` internals are imported lazily inside method bodies (import cycle),
|
||||
so ``patch("gateway.run.X")`` keeps intercepting them at call time.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from gateway.config import Platform
|
||||
from gateway.restart import (
|
||||
DEFAULT_GATEWAY_CRON_DRAIN_TIMEOUT,
|
||||
DEFAULT_GATEWAY_POST_INTERRUPT_GRACE_TIMEOUT,
|
||||
DEFAULT_GATEWAY_RESTART_AFTER_TURN_TIMEOUT,
|
||||
DEFAULT_GATEWAY_RESTART_DRAIN_TIMEOUT,
|
||||
DEFAULT_GATEWAY_SIGNAL_INTERRUPT_GRACE_TIMEOUT,
|
||||
parse_cron_drain_timeout,
|
||||
parse_restart_after_turn_timeout,
|
||||
parse_restart_drain_timeout,
|
||||
parse_signal_interrupt_grace_timeout,
|
||||
)
|
||||
from gateway.session import SessionSource
|
||||
from gateway.session_state import SERVICE_TIER_UNSET as _SERVICE_TIER_UNSET
|
||||
from hermes_cli.config import cfg_get
|
||||
from hermes_cli.fallback_config import get_fallback_chain
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional
|
||||
from utils import is_truthy_value
|
||||
|
||||
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")
|
||||
|
||||
|
||||
class GatewayConfigLoadersMixin:
|
||||
"""Config/env loaders for runtime knobs (busy modes, reasoning, service tier, timeouts, fallback) for GatewayRunner."""
|
||||
|
||||
@staticmethod
|
||||
def _load_prefill_messages() -> List[Dict[str, Any]]:
|
||||
"""Load ephemeral prefill messages from config or env var.
|
||||
|
||||
HERMES_PREFILL_MESSAGES_FILE env wins, then top-level prefill_messages_file in config.yaml,
|
||||
then legacy agent.prefill_messages_file. Relative paths resolve from ~/.hermes/.
|
||||
"""
|
||||
from gateway.run import _hermes_home, _load_gateway_runtime_config
|
||||
file_path = os.getenv("HERMES_PREFILL_MESSAGES_FILE", "")
|
||||
if not file_path:
|
||||
cfg = _load_gateway_runtime_config()
|
||||
file_path = str(cfg.get("prefill_messages_file", "") or "")
|
||||
if not file_path:
|
||||
file_path = str(cfg_get(cfg, "agent", "prefill_messages_file", default="") or "")
|
||||
if not file_path:
|
||||
return []
|
||||
path = Path(file_path).expanduser()
|
||||
if not path.is_absolute():
|
||||
path = _hermes_home / path
|
||||
if not path.exists():
|
||||
logger.warning("Prefill messages file not found: %s", path)
|
||||
return []
|
||||
try:
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
if not isinstance(data, list):
|
||||
logger.warning("Prefill messages file must contain a JSON array: %s", path)
|
||||
return []
|
||||
return data
|
||||
except Exception as e:
|
||||
logger.warning("Failed to load prefill messages from %s: %s", path, e)
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def _load_ephemeral_system_prompt() -> str:
|
||||
"""Load ephemeral system prompt: HERMES_EPHEMERAL_SYSTEM_PROMPT env var first, then
|
||||
``display.personality`` / ``agent.system_prompt`` in config.yaml.
|
||||
"""
|
||||
from gateway.run import _load_gateway_runtime_config
|
||||
from hermes_cli.config import resolve_ephemeral_system_prompt_from_config
|
||||
|
||||
prompt = os.getenv("HERMES_EPHEMERAL_SYSTEM_PROMPT", "")
|
||||
if prompt:
|
||||
return prompt
|
||||
cfg = _load_gateway_runtime_config()
|
||||
return resolve_ephemeral_system_prompt_from_config(cfg)
|
||||
|
||||
def _resolve_model_for_channel(
|
||||
self,
|
||||
platform: Platform,
|
||||
chat_id: str,
|
||||
*,
|
||||
user_config: Optional[dict] = None,
|
||||
thread_id: Optional[str] = None,
|
||||
parent_id: Optional[str] = None,
|
||||
) -> str:
|
||||
"""Resolve model for this channel: channel_overrides else global default.
|
||||
|
||||
Precedence lives in :func:`hermes_cli.model_switch.resolve_effective_model` (shared with the
|
||||
API server so the surfaces cannot diverge). No session tier here: session /model overrides
|
||||
are applied later by ``_apply_session_model_override``.
|
||||
"""
|
||||
from gateway.run import _get_channel_override, _resolve_gateway_model
|
||||
from hermes_cli.model_switch import resolve_effective_model
|
||||
|
||||
override = None
|
||||
config = getattr(self, "config", None)
|
||||
if config:
|
||||
override = _get_channel_override(
|
||||
config,
|
||||
platform,
|
||||
chat_id,
|
||||
thread_id=thread_id,
|
||||
parent_id=parent_id,
|
||||
)
|
||||
return resolve_effective_model(
|
||||
None, # session tier applied downstream (_apply_session_model_override)
|
||||
override,
|
||||
_resolve_gateway_model(user_config),
|
||||
)
|
||||
|
||||
def _get_system_prompt_for_channel(
|
||||
self,
|
||||
platform: Platform,
|
||||
chat_id: str,
|
||||
*,
|
||||
thread_id: Optional[str] = None,
|
||||
parent_id: Optional[str] = None,
|
||||
) -> str:
|
||||
"""Ephemeral system prompt for this channel/thread.
|
||||
|
||||
``channel_overrides`` when set, else the gateway prompt resolved from the CURRENT profile's
|
||||
config on every call (callers run inside ``_profile_runtime_scope``, so routed multiplex
|
||||
profiles get their own personality/system_prompt and ``/personality`` edits apply next turn).
|
||||
Legacy ``channel_prompts`` are applied separately via ``event.channel_prompt`` in ``run_sync``.
|
||||
"""
|
||||
from gateway.run import _get_channel_override
|
||||
config = getattr(self, "config", None)
|
||||
if config:
|
||||
override = _get_channel_override(
|
||||
config,
|
||||
platform,
|
||||
chat_id,
|
||||
thread_id=thread_id,
|
||||
parent_id=parent_id,
|
||||
)
|
||||
if override and override.system_prompt:
|
||||
return (override.system_prompt or "").strip()
|
||||
return self._load_ephemeral_system_prompt()
|
||||
|
||||
@staticmethod
|
||||
def _load_reasoning_config(model: str = "") -> dict | None:
|
||||
"""Load reasoning effort from config.yaml, respecting per-model overrides.
|
||||
|
||||
Thin wrapper over :func:`hermes_constants.resolve_reasoning_config` (per-model override >
|
||||
global ``agent.reasoning_effort``; YAML False = disabled). Empty ``model`` uses ``model.default``.
|
||||
"""
|
||||
from gateway.run import _load_gateway_runtime_config
|
||||
from hermes_constants import resolve_reasoning_config
|
||||
cfg = _load_gateway_runtime_config()
|
||||
return resolve_reasoning_config(cfg, model)
|
||||
|
||||
@staticmethod
|
||||
def _parse_reasoning_command_args(raw_args: str) -> tuple[str, bool]:
|
||||
"""Parse `/reasoning` args into `(value, persist_global)`.
|
||||
|
||||
Session-scoped by default; `--global` in any position persists the change to config.yaml.
|
||||
"""
|
||||
import shlex
|
||||
|
||||
text = str(raw_args or "").strip().replace("—", "--")
|
||||
if not text:
|
||||
return "", False
|
||||
try:
|
||||
tokens = shlex.split(text)
|
||||
except ValueError:
|
||||
tokens = text.split()
|
||||
|
||||
persist_global = False
|
||||
value_tokens = []
|
||||
for token in tokens:
|
||||
if token == "--global":
|
||||
persist_global = True
|
||||
else:
|
||||
value_tokens.append(token)
|
||||
return " ".join(value_tokens).strip().lower(), persist_global
|
||||
|
||||
def _resolve_session_reasoning_config(
|
||||
self,
|
||||
*,
|
||||
source: Optional[SessionSource] = None,
|
||||
session_key: Optional[str] = None,
|
||||
model: str = "",
|
||||
) -> dict | None:
|
||||
"""Resolve reasoning effort for a session, honoring session overrides.
|
||||
|
||||
Priority: session ``/reasoning --session`` > per-model ``agent.reasoning_overrides`` > global
|
||||
``agent.reasoning_effort``. ``model`` must be the session's *effective* model (session
|
||||
``/model`` override included); empty uses ``model.default``.
|
||||
"""
|
||||
resolved_session_key = self._resolve_session_key_or_none(source, session_key)
|
||||
|
||||
if resolved_session_key:
|
||||
_r_state = self._peek_session_state(resolved_session_key)
|
||||
if _r_state is not None and _r_state.conversation.reasoning_override is not None:
|
||||
return _r_state.conversation.reasoning_override
|
||||
return self._load_reasoning_config(model)
|
||||
|
||||
def _set_session_reasoning_override(
|
||||
self,
|
||||
session_key: str,
|
||||
reasoning_config: Optional[dict],
|
||||
) -> None:
|
||||
"""Set or clear the session-scoped reasoning override."""
|
||||
if not session_key:
|
||||
return
|
||||
# Per-session field write: a lazy ``_session_reasoning_overrides = {}`` init replaced the
|
||||
# WHOLE dict, racing concurrent sessions; a SessionState field reset cannot cross sessions.
|
||||
self._session_state(session_key).conversation.reasoning_override = (
|
||||
None if reasoning_config is None else dict(reasoning_config)
|
||||
)
|
||||
|
||||
def _resolve_session_service_tier(
|
||||
self,
|
||||
source=None,
|
||||
session_key: Optional[str] = None,
|
||||
) -> Optional[str]:
|
||||
"""Resolve the effective service tier for a session.
|
||||
|
||||
A session-scoped /fast override beats the config default; the override dict stores
|
||||
"priority" or None (explicit normal), so key presence — not truthiness — decides.
|
||||
"""
|
||||
resolved_session_key = self._resolve_session_key_or_none(source, session_key)
|
||||
|
||||
if resolved_session_key:
|
||||
_t_state = self._peek_session_state(resolved_session_key)
|
||||
if (
|
||||
_t_state is not None
|
||||
and _t_state.conversation.service_tier_override
|
||||
is not _SERVICE_TIER_UNSET
|
||||
):
|
||||
return _t_state.conversation.service_tier_override
|
||||
return self._load_service_tier()
|
||||
|
||||
def _set_session_service_tier_override(
|
||||
self,
|
||||
session_key: str,
|
||||
service_tier,
|
||||
clear: bool = False,
|
||||
) -> None:
|
||||
"""Set or clear the session-scoped /fast override.
|
||||
|
||||
``service_tier`` is "priority" or None (explicit normal). Pass
|
||||
``clear=True`` to remove the override entirely (fall back to config).
|
||||
"""
|
||||
if not session_key:
|
||||
return
|
||||
# Presence-sensitive: "priority" or None (explicit normal) both count as an override; the
|
||||
# sentinel means "no override". Per-session field write: a lazy dict replace races sessions.
|
||||
self._session_state(session_key).conversation.service_tier_override = (
|
||||
_SERVICE_TIER_UNSET if clear else service_tier
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _load_service_tier() -> str | None:
|
||||
"""Load Priority Processing (agent.service_tier) from config.yaml: "fast"/"priority"/"on" =>
|
||||
"priority"; "normal"/"off" disable; None when unset/unsupported.
|
||||
"""
|
||||
from gateway.run import _load_gateway_runtime_config
|
||||
cfg = _load_gateway_runtime_config()
|
||||
raw = str(cfg_get(cfg, "agent", "service_tier", default="") or "").strip()
|
||||
|
||||
value = raw.lower()
|
||||
if not value or value in {"normal", "default", "standard", "off", "none"}:
|
||||
return None
|
||||
if value in {"fast", "priority", "on"}:
|
||||
return "priority"
|
||||
if value in {"auto", "cold"}:
|
||||
return value
|
||||
logger.warning("Unknown service_tier '%s', ignoring", raw)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _load_show_reasoning() -> bool:
|
||||
"""Load show_reasoning toggle from config.yaml display section."""
|
||||
from gateway.run import _load_gateway_runtime_config
|
||||
cfg = _load_gateway_runtime_config()
|
||||
return is_truthy_value(
|
||||
cfg_get(cfg, "display", "show_reasoning"),
|
||||
default=False,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _load_busy_input_mode() -> str:
|
||||
"""Load gateway drain-time busy-input behavior from config/env."""
|
||||
from gateway.run import _load_gateway_runtime_config
|
||||
mode = os.getenv("HERMES_GATEWAY_BUSY_INPUT_MODE", "").strip().lower()
|
||||
if not mode:
|
||||
cfg = _load_gateway_runtime_config()
|
||||
mode = str(cfg_get(cfg, "display", "busy_input_mode", default="") or "").strip().lower()
|
||||
if mode == "queue":
|
||||
return "queue"
|
||||
if mode == "steer":
|
||||
return "steer"
|
||||
return "interrupt"
|
||||
|
||||
@staticmethod
|
||||
def _load_busy_text_mode() -> str:
|
||||
"""Resolve normal busy TEXT follow-up behavior.
|
||||
|
||||
``busy_input_mode`` is the source of truth (default ``interrupt``); legacy ``busy_text_mode``
|
||||
is honored only when explicitly set so existing queue setups keep working.
|
||||
"""
|
||||
from gateway.run import GatewayRunner, _load_gateway_runtime_config
|
||||
# Legacy explicit override wins for backward compat.
|
||||
legacy = os.getenv("HERMES_GATEWAY_BUSY_TEXT_MODE", "").strip().lower()
|
||||
if not legacy:
|
||||
cfg = _load_gateway_runtime_config()
|
||||
legacy = str(cfg_get(cfg, "display", "busy_text_mode", default="") or "").strip().lower()
|
||||
if legacy == "interrupt":
|
||||
return "interrupt"
|
||||
if legacy == "queue":
|
||||
return "queue"
|
||||
# No explicit legacy knob → follow busy_input_mode.
|
||||
input_mode = GatewayRunner._load_busy_input_mode()
|
||||
return "queue" if input_mode == "queue" else "interrupt"
|
||||
|
||||
@staticmethod
|
||||
def _busy_modes_from_config(
|
||||
config: dict,
|
||||
*,
|
||||
fallback_input: str,
|
||||
fallback_text: str,
|
||||
) -> tuple[str, str]:
|
||||
"""Resolve one profile's busy modes without consulting process env."""
|
||||
raw_input = str(
|
||||
cfg_get(config, "display", "busy_input_mode", default="") or ""
|
||||
).strip().lower()
|
||||
input_mode = (
|
||||
raw_input
|
||||
if raw_input in {"interrupt", "queue", "steer"}
|
||||
else fallback_input
|
||||
)
|
||||
|
||||
raw_text = str(
|
||||
cfg_get(config, "display", "busy_text_mode", default="") or ""
|
||||
).strip().lower()
|
||||
if raw_text in {"interrupt", "queue"}:
|
||||
text_mode = raw_text
|
||||
elif raw_input in {"interrupt", "queue", "steer"}:
|
||||
text_mode = "queue" if input_mode == "queue" else "interrupt"
|
||||
else:
|
||||
text_mode = fallback_text
|
||||
return input_mode, text_mode
|
||||
|
||||
def _snapshot_profile_busy_modes(self, profile_name: str, config: dict) -> None:
|
||||
"""Cache a routed profile's busy policy for this gateway lifetime."""
|
||||
input_mode, text_mode = self._busy_modes_from_config(
|
||||
config,
|
||||
fallback_input=getattr(self, "_busy_input_mode", "interrupt"),
|
||||
fallback_text=getattr(self, "_busy_text_mode", "interrupt"),
|
||||
)
|
||||
input_modes = self.__dict__.setdefault("_busy_input_modes_by_profile", {})
|
||||
text_modes = self.__dict__.setdefault("_busy_text_modes_by_profile", {})
|
||||
input_modes[profile_name] = input_mode
|
||||
text_modes[profile_name] = text_mode
|
||||
|
||||
def _busy_profile_name_for_source(self, source: SessionSource) -> Optional[str]:
|
||||
"""Return the routed profile whose busy policy applies, if any."""
|
||||
if not getattr(getattr(self, "config", None), "multiplex_profiles", False):
|
||||
return None
|
||||
name = str(getattr(source, "profile", "") or "").strip()
|
||||
if not name:
|
||||
try:
|
||||
name = str(self._profile_name_for_source(source) or "").strip()
|
||||
except Exception:
|
||||
name = ""
|
||||
return name or None
|
||||
|
||||
def _effective_busy_input_mode(self, source: SessionSource) -> str:
|
||||
"""Resolve busy input mode from the routed profile startup snapshot."""
|
||||
fallback = getattr(self, "_busy_input_mode", "interrupt")
|
||||
profile_name = self._busy_profile_name_for_source(source)
|
||||
if not profile_name:
|
||||
return fallback
|
||||
modes = getattr(self, "_busy_input_modes_by_profile", None)
|
||||
return modes.get(profile_name, fallback) if isinstance(modes, dict) else fallback
|
||||
|
||||
def _effective_busy_text_mode(self, source: SessionSource) -> str:
|
||||
"""Resolve legacy busy text mode from the routed profile snapshot."""
|
||||
fallback = getattr(self, "_busy_text_mode", "interrupt")
|
||||
profile_name = self._busy_profile_name_for_source(source)
|
||||
if not profile_name:
|
||||
return fallback
|
||||
modes = getattr(self, "_busy_text_modes_by_profile", None)
|
||||
return modes.get(profile_name, fallback) if isinstance(modes, dict) else fallback
|
||||
|
||||
@staticmethod
|
||||
def _load_restart_drain_timeout() -> float:
|
||||
"""Load graceful gateway restart/stop drain timeout in seconds."""
|
||||
from gateway.run import _load_gateway_runtime_config
|
||||
raw = os.getenv("HERMES_RESTART_DRAIN_TIMEOUT", "").strip()
|
||||
if not raw:
|
||||
cfg = _load_gateway_runtime_config()
|
||||
raw = str(cfg_get(cfg, "agent", "restart_drain_timeout", default="") or "").strip()
|
||||
value = parse_restart_drain_timeout(raw)
|
||||
if raw and value == DEFAULT_GATEWAY_RESTART_DRAIN_TIMEOUT:
|
||||
try:
|
||||
float(raw)
|
||||
except (TypeError, ValueError):
|
||||
logger.warning(
|
||||
"Invalid restart_drain_timeout '%s', using default %.0fs",
|
||||
raw,
|
||||
DEFAULT_GATEWAY_RESTART_DRAIN_TIMEOUT,
|
||||
)
|
||||
return value
|
||||
|
||||
@staticmethod
|
||||
def _load_env_or_agent_cfg_timeout(env_var: str, cfg_key: str, parse, default: float) -> float:
|
||||
"""Env var (non-empty) else ``agent.<cfg_key>``; warn once when a supplied value fails to parse.
|
||||
|
||||
``0`` is a valid value; the parser falls back to ``default`` on garbage."""
|
||||
from gateway.run import _load_gateway_runtime_config
|
||||
env_raw = os.getenv(env_var)
|
||||
if env_raw is not None and str(env_raw).strip() != "":
|
||||
raw: object = env_raw
|
||||
else:
|
||||
cfg = _load_gateway_runtime_config()
|
||||
raw = cfg_get(cfg, "agent", cfg_key, default=None)
|
||||
value = parse(raw)
|
||||
if raw is not None and str(raw).strip() != "":
|
||||
try:
|
||||
float(raw)
|
||||
except (TypeError, ValueError):
|
||||
logger.warning("Invalid %s '%s', using default %.0fs", cfg_key, raw, default)
|
||||
return value
|
||||
|
||||
@classmethod
|
||||
def _load_restart_after_turn_timeout(cls) -> float:
|
||||
"""Load in-band restart wait-for-idle timeout in seconds."""
|
||||
return cls._load_env_or_agent_cfg_timeout(
|
||||
"HERMES_RESTART_AFTER_TURN_TIMEOUT", "restart_after_turn_timeout",
|
||||
parse_restart_after_turn_timeout, DEFAULT_GATEWAY_RESTART_AFTER_TURN_TIMEOUT,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _load_cron_drain_timeout(cls) -> float:
|
||||
"""Load the cron-only floor under the stop()/drain wait."""
|
||||
return cls._load_env_or_agent_cfg_timeout(
|
||||
"HERMES_CRON_DRAIN_TIMEOUT", "cron_drain_timeout",
|
||||
parse_cron_drain_timeout, DEFAULT_GATEWAY_CRON_DRAIN_TIMEOUT,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _load_signal_interrupt_grace_timeout() -> float:
|
||||
"""Load the unexpected-signal post-interrupt grace in seconds."""
|
||||
from gateway.run import _load_gateway_runtime_config
|
||||
cfg = _load_gateway_runtime_config()
|
||||
raw = cfg_get(
|
||||
cfg,
|
||||
"gateway",
|
||||
"signal_interrupt_grace_timeout",
|
||||
default=None,
|
||||
)
|
||||
value = parse_signal_interrupt_grace_timeout(raw)
|
||||
if raw is not None and raw != "":
|
||||
try:
|
||||
float(raw)
|
||||
except (TypeError, ValueError):
|
||||
logger.warning(
|
||||
"Invalid signal_interrupt_grace_timeout '%s', using default %.0fs",
|
||||
raw,
|
||||
DEFAULT_GATEWAY_SIGNAL_INTERRUPT_GRACE_TIMEOUT,
|
||||
)
|
||||
return value
|
||||
|
||||
def _post_interrupt_grace_timeout(self) -> float:
|
||||
"""Return the grace before teardown after forcibly interrupting agents."""
|
||||
if (
|
||||
getattr(self, "_signal_initiated_shutdown", False)
|
||||
and not getattr(self, "_restart_requested", False)
|
||||
):
|
||||
return max(
|
||||
0.0,
|
||||
float(
|
||||
getattr(
|
||||
self,
|
||||
"_signal_interrupt_grace_timeout",
|
||||
DEFAULT_GATEWAY_SIGNAL_INTERRUPT_GRACE_TIMEOUT,
|
||||
)
|
||||
),
|
||||
)
|
||||
return DEFAULT_GATEWAY_POST_INTERRUPT_GRACE_TIMEOUT
|
||||
|
||||
@staticmethod
|
||||
def _load_background_notifications_mode() -> str:
|
||||
"""Load background process notification mode from config or env var."""
|
||||
from gateway.run import _load_gateway_runtime_config
|
||||
mode = os.getenv("HERMES_BACKGROUND_NOTIFICATIONS", "")
|
||||
if not mode:
|
||||
cfg = _load_gateway_runtime_config()
|
||||
raw = cfg_get(cfg, "display", "background_process_notifications")
|
||||
if raw is False:
|
||||
mode = "off"
|
||||
elif raw not in {None, ""}:
|
||||
mode = str(raw)
|
||||
mode = (mode or "concise").strip().lower()
|
||||
valid = {"concise", "all", "result", "error", "off"}
|
||||
if mode not in valid:
|
||||
logger.warning(
|
||||
"Unknown background_process_notifications '%s', defaulting to 'concise'",
|
||||
mode,
|
||||
)
|
||||
return "concise"
|
||||
return mode
|
||||
|
||||
@staticmethod
|
||||
def _load_provider_routing() -> dict:
|
||||
"""Load OpenRouter provider routing preferences from config.yaml."""
|
||||
from gateway.run import _load_gateway_runtime_config
|
||||
try:
|
||||
# Canonical gateway loader (fail-open): managed overlay + ${VAR}
|
||||
# expansion now apply to provider_routing too.
|
||||
cfg = _load_gateway_runtime_config()
|
||||
return cfg.get("provider_routing", {}) or {}
|
||||
except Exception:
|
||||
pass
|
||||
return {}
|
||||
|
||||
@staticmethod
|
||||
def _load_fallback_model() -> list | None:
|
||||
"""Load fallback provider chain from config.yaml.
|
||||
|
||||
Merges ``fallback_providers`` (kept first) with legacy ``fallback_model`` entries.
|
||||
"""
|
||||
from gateway.run import _load_gateway_runtime_config
|
||||
try:
|
||||
# Canonical gateway loader (fail-open): managed overlay + ${VAR}
|
||||
# expansion now apply to the fallback chain too.
|
||||
cfg = _load_gateway_runtime_config()
|
||||
fb = get_fallback_chain(cfg)
|
||||
if fb:
|
||||
return fb
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
|
||||
def _refresh_fallback_model(self) -> list | None:
|
||||
"""Re-read fallback_providers from disk for the next agent create/reuse.
|
||||
|
||||
Lets a chain edited after startup reach messaging sessions (cron already re-reads per job).
|
||||
A TRANSIENT read/parse failure (user mid-edit, non-atomic write) keeps the last known-good
|
||||
chain; only a successful read that genuinely lacks the key clears it.
|
||||
"""
|
||||
from gateway.run import _hermes_home
|
||||
try:
|
||||
from hermes_cli.config import read_user_config_raw
|
||||
cfg_path = _hermes_home / "config.yaml"
|
||||
if not cfg_path.exists():
|
||||
self._fallback_model = None
|
||||
return self._fallback_model
|
||||
# Raw primitive (raises on parse failure) is required here: the canonical fail-open
|
||||
# loader would return {} on a torn mid-edit write and WIPE the last known-good chain.
|
||||
# The overlay/expansion below fixes the managed-scope/${VAR} drift without losing that.
|
||||
cfg = read_user_config_raw(cfg_path)
|
||||
try:
|
||||
from hermes_cli import managed_scope
|
||||
cfg = managed_scope.apply_managed_overlay(cfg)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
from hermes_cli.config import _expand_env_vars
|
||||
expanded = _expand_env_vars(cfg)
|
||||
if isinstance(expanded, dict):
|
||||
cfg = expanded
|
||||
except Exception:
|
||||
pass
|
||||
except Exception:
|
||||
# Transient failure — keep last known-good chain.
|
||||
logger.debug(
|
||||
"fallback_providers refresh: config.yaml read failed; "
|
||||
"keeping last known-good chain", exc_info=True,
|
||||
)
|
||||
return self._fallback_model
|
||||
self._fallback_model = get_fallback_chain(cfg) or None
|
||||
return self._fallback_model
|
||||
|
||||
@staticmethod
|
||||
def _apply_fallback_chain_to_agent(agent: Any, chain: list | None) -> None:
|
||||
"""Keep a cached agent's fallback chain aligned with current config.
|
||||
|
||||
Skips the rewrite while a cooldown holds the agent on an activated fallback provider
|
||||
(``restore_primary_runtime`` owns that lifecycle); otherwise replaces the chain so
|
||||
mid-uptime ``fallback_providers`` edits apply without a restart.
|
||||
"""
|
||||
if agent is None:
|
||||
return
|
||||
new_chain = list(chain or [])
|
||||
rate_limited_until = getattr(agent, "_rate_limited_until", 0) or 0
|
||||
if (
|
||||
getattr(agent, "_fallback_activated", False)
|
||||
and rate_limited_until > time.monotonic()
|
||||
):
|
||||
return
|
||||
old_chain = list(getattr(agent, "_fallback_chain", []) or [])
|
||||
agent._fallback_chain = new_chain
|
||||
agent._fallback_model = new_chain[0] if new_chain else None
|
||||
if not getattr(agent, "_fallback_activated", False):
|
||||
agent._fallback_index = 0
|
||||
# A config edit means the user changed something — drop the session-scoped unavailability
|
||||
# memo so re-configured entries (e.g. credentials added mid-uptime) get retried. Only on real
|
||||
# content change, so the per-message no-op refresh keeps the memo's rate-limiting benefit.
|
||||
if new_chain != old_chain:
|
||||
unavailable = getattr(agent, "_unavailable_fallback_keys", None)
|
||||
if unavailable:
|
||||
unavailable.clear()
|
||||
532
gateway/run_goals.py
Normal file
532
gateway/run_goals.py
Normal file
@@ -0,0 +1,532 @@
|
||||
"""Goal/heartbeat continuation, post-turn hooks and loop-wakeup watcher methods for GatewayRunner.
|
||||
|
||||
Split out of ``gateway/run.py``; bound onto ``GatewayRunner`` via the MRO.
|
||||
``gateway.run`` internals are imported lazily inside method bodies (import cycle),
|
||||
so ``patch("gateway.run.X")`` keeps intercepting them at call time.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
import asyncio
|
||||
import time
|
||||
from contextlib import suppress
|
||||
from gateway.platforms.base import MessageEvent, MessageType
|
||||
from typing import Any
|
||||
|
||||
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")
|
||||
|
||||
|
||||
class GatewayGoalsMixin:
|
||||
"""Goal/heartbeat continuation, post-turn hooks and loop-wakeup watcher methods for GatewayRunner."""
|
||||
|
||||
# ────────────────────────────────────────────────────────────────
|
||||
# /goal — persistent cross-turn goals (Ralph-style loop)
|
||||
# ────────────────────────────────────────────────────────────────
|
||||
def _goal_max_turns_from_config(self) -> int:
|
||||
"""Resolve the configured /goal turn budget for gateway sessions.
|
||||
|
||||
GatewayRunner.config is a GatewayConfig dataclass, not the full user config mapping, so
|
||||
top-level blocks such as ``goals`` are only reachable via hermes_cli.config.load_config().
|
||||
"""
|
||||
try:
|
||||
goals_cfg = (
|
||||
(self.config or {}).get("goals", {})
|
||||
if isinstance(self.config, dict)
|
||||
else getattr(self.config, "goals", {}) or {}
|
||||
)
|
||||
if not goals_cfg:
|
||||
from hermes_cli.config import load_config
|
||||
|
||||
goals_cfg = (load_config() or {}).get("goals") or {}
|
||||
return int(goals_cfg.get("max_turns", 20) or 20)
|
||||
except Exception:
|
||||
return 20
|
||||
|
||||
async def _warm_goals_session_db(self, label: str) -> None:
|
||||
"""Warm the goals SessionDB cache off-loop (best-effort).
|
||||
|
||||
A cold cache runs the state.db init on the loop thread and freezes the loop for the init
|
||||
duration. The executor hop keeps the profile home override alive under multiplex, so the
|
||||
warm cache belongs to the caller's profile. On failure the caller falls back to the
|
||||
bootstrap windows, so a dropped warm-up is a bounded stall, never a crash.
|
||||
"""
|
||||
try:
|
||||
from hermes_cli.goals import _get_session_db as _warm_goals_db
|
||||
|
||||
await self._run_in_executor_with_context(_warm_goals_db)
|
||||
except Exception as exc:
|
||||
logger.warning("%s: session DB warm-up failed: %s", label, exc)
|
||||
|
||||
async def _session_entry_for_manager(self, event: "MessageEvent", label: str):
|
||||
"""Session entry for a /goal or /heartbeat manager, or None when lookup fails.
|
||||
|
||||
Warms the SessionDB cache off-loop first: a cold cache freezes the loop for the init
|
||||
duration and drops the first write while the reply claims it was set. Internal events look
|
||||
the session up WITHOUT touching activity so they never advance the idle/daily reset clock.
|
||||
"""
|
||||
await self._warm_goals_session_db(label)
|
||||
try:
|
||||
session_entry = await self.async_session_store.get_or_create_session(
|
||||
event.source,
|
||||
touch_activity=not bool(getattr(event, "internal", False)),
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug("%s: session lookup failed: %s", label, exc)
|
||||
return None
|
||||
if not (getattr(session_entry, "session_id", None) or ""):
|
||||
return None
|
||||
return session_entry
|
||||
|
||||
async def _get_goal_manager_for_event(self, event: "MessageEvent"):
|
||||
"""Return ``(GoalManager, session_entry)`` for this event, or ``(None, None)``."""
|
||||
try:
|
||||
from hermes_cli.goals import GoalManager
|
||||
except Exception as exc:
|
||||
logger.debug("goal manager unavailable: %s", exc)
|
||||
return None, None
|
||||
session_entry = await self._session_entry_for_manager(event, "goal manager")
|
||||
if session_entry is None:
|
||||
return None, None
|
||||
max_turns = self._goal_max_turns_from_config()
|
||||
return GoalManager(session_id=session_entry.session_id, default_max_turns=max_turns), session_entry
|
||||
|
||||
async def _get_heartbeat_manager_for_event(self, event: "MessageEvent"):
|
||||
"""Return ``(HeartbeatManager, session_entry)`` for this event, or ``(None, None)``."""
|
||||
try:
|
||||
from hermes_cli.heartbeat import HeartbeatManager
|
||||
except Exception as exc:
|
||||
logger.debug("heartbeat manager unavailable: %s", exc)
|
||||
return None, None
|
||||
session_entry = await self._session_entry_for_manager(event, "heartbeat manager")
|
||||
if session_entry is None:
|
||||
return None, None
|
||||
return HeartbeatManager(session_id=session_entry.session_id), session_entry
|
||||
|
||||
def _register_heartbeat_watch(self, quick_key: str, source: Any, session_id: str) -> None:
|
||||
"""Track a session with an active heartbeat and start the poller.
|
||||
|
||||
The registry maps ``quick_key`` → ``(source, session_id)`` so the poller can rebuild a
|
||||
MessageEvent and enqueue via the adapter FIFO. In-memory by design: heartbeat STATE
|
||||
survives restarts in SessionDB, but firing resumes only when the user touches /heartbeat
|
||||
again (durable schedules belong to cron).
|
||||
"""
|
||||
watch = getattr(self, "_heartbeat_watch", None)
|
||||
if watch is None:
|
||||
watch = {}
|
||||
self._heartbeat_watch = watch
|
||||
watch[quick_key] = (source, session_id)
|
||||
self._start_heartbeat_poller()
|
||||
|
||||
def _unregister_heartbeat_watch(self, quick_key: str) -> None:
|
||||
watch = getattr(self, "_heartbeat_watch", None)
|
||||
if watch:
|
||||
watch.pop(quick_key, None)
|
||||
|
||||
def _start_heartbeat_poller(self) -> None:
|
||||
"""Start the single gateway-wide heartbeat poll task (idempotent)."""
|
||||
existing = getattr(self, "_heartbeat_poll_task", None)
|
||||
if existing is not None and not existing.done():
|
||||
return
|
||||
|
||||
from hermes_cli.heartbeat import POLL_SECONDS
|
||||
|
||||
async def _poll_loop():
|
||||
while True:
|
||||
await asyncio.sleep(POLL_SECONDS)
|
||||
watch = getattr(self, "_heartbeat_watch", None)
|
||||
if not watch:
|
||||
continue
|
||||
# Warm the cache off-loop once per poll. A watch can only be registered through the
|
||||
# warmed /heartbeat command, so this covers only the degraded path where that warm-
|
||||
# up failed.
|
||||
await self._warm_goals_session_db("heartbeat poll")
|
||||
for quick_key, (source, session_id) in list(watch.items()):
|
||||
try:
|
||||
# Busy sessions coalesce their tick to the next idle poll.
|
||||
if quick_key in self._running_agents:
|
||||
continue
|
||||
from hermes_cli.heartbeat import HeartbeatManager
|
||||
|
||||
mgr = HeartbeatManager(session_id=session_id)
|
||||
if not mgr.has_heartbeat():
|
||||
watch.pop(quick_key, None)
|
||||
continue
|
||||
prompt = mgr.due_prompt()
|
||||
if not prompt:
|
||||
continue
|
||||
adapter = self._adapter_for_source(source)
|
||||
if adapter is None:
|
||||
continue
|
||||
hb_event = MessageEvent(
|
||||
text=prompt,
|
||||
message_type=MessageType.TEXT,
|
||||
source=source,
|
||||
message_id=None,
|
||||
channel_prompt=None,
|
||||
)
|
||||
self._enqueue_fifo(quick_key, hb_event, adapter)
|
||||
except Exception as exc:
|
||||
logger.debug("heartbeat poll for %s failed: %s", quick_key, exc)
|
||||
|
||||
try:
|
||||
task = asyncio.create_task(_poll_loop())
|
||||
self._heartbeat_poll_task = task
|
||||
# PERMANENT once started (an infinite while-True loop, no exit condition) — same as a
|
||||
# _spawn_supervised watcher. Tag it so _scale_to_zero_has_live_background_work() doesn't
|
||||
# treat a gateway with an active heartbeat watch as busy forever.
|
||||
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)
|
||||
except Exception:
|
||||
logger.debug("Failed to start heartbeat poller", exc_info=True)
|
||||
|
||||
async def _send_goal_status_notice(self, source: Any, message: str) -> None:
|
||||
"""Send a /goal judge status line back to the originating chat/thread."""
|
||||
adapter = self._adapter_for_source(source)
|
||||
if not adapter:
|
||||
logger.debug("goal continuation: no adapter for %s", getattr(source, "platform", None))
|
||||
return
|
||||
|
||||
try:
|
||||
metadata = self._thread_metadata_for_source(source)
|
||||
except Exception:
|
||||
metadata = None
|
||||
|
||||
result = await adapter.send(source.chat_id, message, metadata=metadata)
|
||||
if result is not None and not getattr(result, "success", True):
|
||||
logger.warning(
|
||||
"goal continuation: status send failed: %s",
|
||||
getattr(result, "error", "unknown error"),
|
||||
)
|
||||
|
||||
async def _defer_goal_status_notice_after_delivery(self, source: Any, message: str) -> None:
|
||||
"""Send a /goal status line after the main response is delivered.
|
||||
|
||||
The adapter sends the agent response after this caller returns, so for reading order the
|
||||
status must follow that send: use the adapter's one-shot post-delivery callback when
|
||||
available, else fall back to direct awaited delivery rather than dropping the notice.
|
||||
"""
|
||||
adapter = self._adapter_for_source(source)
|
||||
if not adapter:
|
||||
logger.debug("goal continuation: no adapter for %s", getattr(source, "platform", None))
|
||||
return
|
||||
|
||||
async def _deliver() -> None:
|
||||
try:
|
||||
await self._send_goal_status_notice(source, message)
|
||||
except Exception as exc:
|
||||
logger.warning("goal continuation: status send failed: %s", exc, exc_info=True)
|
||||
|
||||
try:
|
||||
session_key = self._session_key_for_source(source)
|
||||
except Exception:
|
||||
session_key = None
|
||||
|
||||
if session_key and hasattr(adapter, "register_post_delivery_callback"):
|
||||
try:
|
||||
generation = None
|
||||
active = getattr(adapter, "_active_sessions", {}).get(session_key)
|
||||
if active is not None:
|
||||
generation = getattr(active, "_hermes_run_generation", None)
|
||||
adapter.register_post_delivery_callback(
|
||||
session_key,
|
||||
_deliver,
|
||||
generation=generation,
|
||||
)
|
||||
return
|
||||
except Exception as exc:
|
||||
logger.debug("goal continuation: post-delivery callback registration failed: %s", exc)
|
||||
|
||||
await _deliver()
|
||||
|
||||
async def _post_turn_goal_continuation(
|
||||
self,
|
||||
*,
|
||||
session_entry: Any,
|
||||
source: Any,
|
||||
final_response: str,
|
||||
) -> None:
|
||||
"""Run the goal judge after a gateway turn and, if still active, enqueue a continuation
|
||||
prompt for the same session.
|
||||
|
||||
Called at turn boundary AFTER delivery. Uses the adapter's pending-message/FIFO machinery
|
||||
so a simultaneous real user message is handled by the same queue and takes priority.
|
||||
"""
|
||||
try:
|
||||
from hermes_cli.goals import GoalManager
|
||||
except Exception as exc:
|
||||
logger.debug("goal continuation: goals module unavailable: %s", exc)
|
||||
return
|
||||
|
||||
sid = getattr(session_entry, "session_id", None) or ""
|
||||
if not sid:
|
||||
return
|
||||
|
||||
max_turns = self._goal_max_turns_from_config()
|
||||
|
||||
# Warm the SessionDB cache off-loop: a cold cache runs the state.db init on the loop thread
|
||||
# at the turn boundary; a slow init can drop the goal read and silently end the goal loop.
|
||||
await self._warm_goals_session_db("goal continuation")
|
||||
|
||||
mgr = GoalManager(session_id=sid, default_max_turns=max_turns)
|
||||
if not mgr.is_active():
|
||||
return
|
||||
|
||||
try:
|
||||
from hermes_cli.goals import gather_background_processes as _gather_bg
|
||||
_bg_procs = _gather_bg()
|
||||
except Exception:
|
||||
_bg_procs = None
|
||||
|
||||
# evaluate_after_turn calls judge_goal(), a synchronous HTTP request to the auxiliary LLM;
|
||||
# on the event-loop thread it blocks Discord heartbeats 10-40 s and flaps connections, so it
|
||||
# is offloaded to a thread-pool executor. _run_in_executor_with_context (not bare
|
||||
# run_in_executor): the profile secret scope and aux runtime context are contextvars; a
|
||||
# default-executor hop drops them and aux credential resolution fails under multiplexing.
|
||||
decision = await self._run_in_executor_with_context(
|
||||
lambda: mgr.evaluate_after_turn(
|
||||
final_response or "",
|
||||
user_initiated=True,
|
||||
background_processes=_bg_procs,
|
||||
),
|
||||
)
|
||||
msg = decision.get("message") or ""
|
||||
|
||||
# Defer the status line until after the adapter has delivered the agent's visible final
|
||||
# response. The judge runs after the response is produced but before BasePlatformAdapter
|
||||
# sends it, so sending here would show "✓ Goal achieved" before the answer itself.
|
||||
if msg and source is not None:
|
||||
await self._defer_goal_status_notice_after_delivery(source, msg)
|
||||
|
||||
if not decision.get("should_continue"):
|
||||
return
|
||||
|
||||
prompt = decision.get("continuation_prompt") or ""
|
||||
if not prompt or source is None:
|
||||
return
|
||||
|
||||
# Enqueue via the adapter's FIFO so a user message already in
|
||||
# flight preempts the continuation naturally.
|
||||
try:
|
||||
adapter = self._adapter_for_source(source)
|
||||
_quick_key = self._session_key_for_source(source)
|
||||
if adapter and _quick_key:
|
||||
cont_event = MessageEvent(
|
||||
text=prompt,
|
||||
message_type=MessageType.TEXT,
|
||||
source=source,
|
||||
message_id=None,
|
||||
channel_prompt=None,
|
||||
)
|
||||
self._enqueue_fifo(_quick_key, cont_event, adapter)
|
||||
except Exception as exc:
|
||||
logger.debug("goal continuation: enqueue failed: %s", exc)
|
||||
|
||||
async def _run_post_turn_hooks(
|
||||
self,
|
||||
*,
|
||||
agent_result: Any,
|
||||
source: Any,
|
||||
is_internal: bool,
|
||||
event: Any = None,
|
||||
) -> None:
|
||||
"""Run goal and loop bookkeeping after an agent turn returns."""
|
||||
final_text = self._final_text_for_post_turn_hooks(agent_result, event)
|
||||
|
||||
try:
|
||||
session_entry = await self.async_session_store.get_or_create_session(
|
||||
source,
|
||||
touch_activity=not is_internal,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug("post-turn session resolution failed: %s", exc)
|
||||
return
|
||||
|
||||
# Empty interrupted/errored responses must not drive /goal, but an
|
||||
# in-flight /loop tick still needs to be released and rescheduled.
|
||||
if final_text.strip():
|
||||
try:
|
||||
await self._post_turn_goal_continuation(
|
||||
session_entry=session_entry,
|
||||
source=source,
|
||||
final_response=final_text,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug("goal continuation hook failed: %s", exc)
|
||||
try:
|
||||
await self._post_turn_loop_completion(
|
||||
session_entry=session_entry,
|
||||
source=source,
|
||||
final_response=final_text,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug("loop completion hook failed: %s", exc)
|
||||
|
||||
@staticmethod
|
||||
def _final_text_for_post_turn_hooks(agent_result, event=None) -> str:
|
||||
"""Text for /goal and /loop after a gateway turn.
|
||||
|
||||
Streamed turns return None from _handle_message_with_agent (already_sent). The delivered
|
||||
reply is stashed on the event so those hooks still see it.
|
||||
"""
|
||||
text = ""
|
||||
if isinstance(agent_result, dict):
|
||||
text = str(agent_result.get("final_response") or "")
|
||||
elif isinstance(agent_result, str):
|
||||
text = agent_result
|
||||
if text.strip():
|
||||
return text
|
||||
streamed = getattr(event, "_streamed_final_response", None)
|
||||
if isinstance(streamed, str) and streamed.strip():
|
||||
return streamed
|
||||
return text
|
||||
|
||||
async def _post_turn_loop_completion(
|
||||
self,
|
||||
*,
|
||||
session_entry: Any,
|
||||
source: Any,
|
||||
final_response: str,
|
||||
) -> None:
|
||||
"""Complete a /loop wakeup tick after a gateway turn.
|
||||
|
||||
No-op unless the session has a loop whose tick is in flight (``awaiting_response`` — set
|
||||
when the wakeup was injected). Applies the LOOP_COMPLETE marker / --until judge / caps
|
||||
and schedules the next tick; the idle wakeup watcher fires it when due.
|
||||
"""
|
||||
try:
|
||||
from hermes_cli.loops import LoopManager
|
||||
except Exception as exc:
|
||||
logger.debug("loop completion: loops module unavailable: %s", exc)
|
||||
return
|
||||
|
||||
sid = getattr(session_entry, "session_id", None) or ""
|
||||
if not sid:
|
||||
return
|
||||
|
||||
# Warm the SessionDB cache off-loop: a cold cache at the turn boundary stalls the loop for
|
||||
# the init duration and can drop the tick-completion write (the /goal continuation seam).
|
||||
await self._warm_goals_session_db("loop completion")
|
||||
|
||||
mgr = LoopManager(session_id=sid)
|
||||
state = mgr.state
|
||||
if state is None or not state.awaiting_response:
|
||||
return
|
||||
|
||||
# The --until judge is a sync aux-LLM call — keep it off the event loop.
|
||||
decision = await asyncio.get_running_loop().run_in_executor(
|
||||
None, mgr.complete_tick, final_response or ""
|
||||
)
|
||||
msg = decision.get("message") or ""
|
||||
if msg and source is not None:
|
||||
await self._defer_goal_status_notice_after_delivery(source, msg)
|
||||
|
||||
async def _loop_wakeup_watcher(self, interval: float = 15.0) -> None:
|
||||
"""Fire due /loop wakeups for idle gateway sessions.
|
||||
|
||||
The gateway has no per-session scheduler thread, so a coarse ticker scans persisted loops
|
||||
(SessionDB ``loop:*`` rows) and injects the wakeup prompt into each due session's chat
|
||||
via the same synthetic-message path used by watch notifications. Deferrals: session
|
||||
currently running a turn → skip (the FIFO would race the live turn); active non-parked
|
||||
/goal → skip (goal owns the idle boundary); no routing metadata → skip with a one-time
|
||||
warning (CLI/TUI loops carry no route).
|
||||
"""
|
||||
await asyncio.sleep(5) # let platforms finish connecting
|
||||
warned_no_route: set = set()
|
||||
while self._running:
|
||||
try:
|
||||
from hermes_cli.loops import (
|
||||
LoopManager,
|
||||
goal_blocks_loop_tick,
|
||||
list_active_loops,
|
||||
)
|
||||
|
||||
# Warm the cache off-loop once per scan: the scan reads every persisted loop, so a
|
||||
# cold cache would run the state.db init on the loop thread before the first read.
|
||||
await self._warm_goals_session_db("loop wakeup")
|
||||
|
||||
now = time.time()
|
||||
for sid, state in list_active_loops():
|
||||
if state.awaiting_response or now < state.next_due_at:
|
||||
continue
|
||||
route = state.route or {}
|
||||
platform_name = route.get("platform", "")
|
||||
chat_id = route.get("chat_id", "")
|
||||
if not platform_name or not chat_id:
|
||||
# CLI / TUI-owned loop — their own schedulers drive it.
|
||||
continue
|
||||
adapter = None
|
||||
for p, a in self.adapters.items():
|
||||
if p.value == platform_name:
|
||||
adapter = a
|
||||
break
|
||||
if adapter is None:
|
||||
if sid not in warned_no_route:
|
||||
warned_no_route.add(sid)
|
||||
logger.debug(
|
||||
"loop wakeup: no adapter for platform %r (session %s)",
|
||||
platform_name, sid,
|
||||
)
|
||||
continue
|
||||
|
||||
# Build the source + session key to check business.
|
||||
evt_stub = {
|
||||
"session_key": "",
|
||||
"platform": platform_name,
|
||||
"chat_id": chat_id,
|
||||
"chat_type": route.get("chat_type", ""),
|
||||
"thread_id": route.get("thread_id", ""),
|
||||
"user_id": route.get("user_id", ""),
|
||||
"user_name": route.get("user_name", ""),
|
||||
}
|
||||
source = self._build_process_event_source(evt_stub)
|
||||
if source is None:
|
||||
continue
|
||||
try:
|
||||
session_key = self._session_key_for_source(source)
|
||||
except Exception:
|
||||
session_key = None
|
||||
if session_key and session_key in self._running_agents:
|
||||
continue # busy — stays due, next scan retries
|
||||
if goal_blocks_loop_tick(sid):
|
||||
continue
|
||||
|
||||
mgr = LoopManager(session_id=sid)
|
||||
if not mgr.is_due(now):
|
||||
continue
|
||||
wakeup = mgr.fire_tick()
|
||||
if not wakeup:
|
||||
continue
|
||||
try:
|
||||
synth_event = MessageEvent(
|
||||
text=wakeup,
|
||||
message_type=MessageType.TEXT,
|
||||
source=source,
|
||||
internal=True,
|
||||
)
|
||||
logger.info(
|
||||
"loop wakeup #%s — injecting for %s chat=%s thread=%s",
|
||||
mgr.state.ticks_fired if mgr.state else "?",
|
||||
platform_name, source.chat_id, source.thread_id,
|
||||
)
|
||||
await adapter.handle_message(synth_event)
|
||||
# Slash-command loops dispatch through the command
|
||||
# path and never hit the post-turn completion hook —
|
||||
# complete the tick immediately (caps + scheduling).
|
||||
if wakeup.lstrip().startswith("/"):
|
||||
mgr.complete_tick("")
|
||||
except Exception as exc:
|
||||
logger.warning("loop wakeup injection failed for %s: %s", sid, exc)
|
||||
with suppress(Exception):
|
||||
mgr.abandon_tick()
|
||||
except Exception as exc:
|
||||
logger.debug("loop wakeup watcher error: %s", exc)
|
||||
await asyncio.sleep(interval)
|
||||
2634
gateway/run_inbound.py
Normal file
2634
gateway/run_inbound.py
Normal file
File diff suppressed because it is too large
Load Diff
2098
gateway/run_notifications.py
Normal file
2098
gateway/run_notifications.py
Normal file
File diff suppressed because it is too large
Load Diff
2348
gateway/run_shutdown.py
Normal file
2348
gateway/run_shutdown.py
Normal file
File diff suppressed because it is too large
Load Diff
2071
gateway/run_startup.py
Normal file
2071
gateway/run_startup.py
Normal file
File diff suppressed because it is too large
Load Diff
878
gateway/run_topics.py
Normal file
878
gateway/run_topics.py
Normal file
@@ -0,0 +1,878 @@
|
||||
"""Telegram forum-topic and Discord auto-thread binding/rename methods for GatewayRunner.
|
||||
|
||||
Split out of ``gateway/run.py``; bound onto ``GatewayRunner`` via the MRO.
|
||||
``gateway.run`` internals are imported lazily inside method bodies (import cycle),
|
||||
so ``patch("gateway.run.X")`` keeps intercepting them at call time.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
import asyncio
|
||||
import dataclasses
|
||||
import re
|
||||
from agent.compaction_display import project_compaction_message_for_display
|
||||
from agent.i18n import t
|
||||
from gateway.config import Platform
|
||||
from gateway.platforms.base import MessageEvent, _prefix_within_utf16_limit, utf16_len
|
||||
from gateway.session import SessionSource
|
||||
from pathlib import Path
|
||||
from typing import Optional, Tuple
|
||||
|
||||
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")
|
||||
|
||||
|
||||
class GatewayTopicThreadsMixin:
|
||||
"""Telegram forum-topic and Discord auto-thread binding/rename methods for GatewayRunner."""
|
||||
|
||||
@staticmethod
|
||||
def _telegram_topic_profile_name(source: SessionSource) -> str:
|
||||
"""Profile namespace for Telegram topic-mode rows.
|
||||
|
||||
Use the profile stamped on the routed event (``source.profile``), never the process-global
|
||||
active profile — under multiplex that mis-attributes topic state across bots sharing state.db.
|
||||
"""
|
||||
name = str(getattr(source, "profile", None) or "").strip()
|
||||
return name if name else "default"
|
||||
|
||||
def _telegram_topic_mode_enabled(self, source: SessionSource) -> bool:
|
||||
"""Return whether Telegram DM topic mode is active for this chat."""
|
||||
if source.platform != Platform.TELEGRAM or source.chat_type != "dm":
|
||||
return False
|
||||
session_db = getattr(self, "_session_db", None)
|
||||
if session_db is None:
|
||||
return False
|
||||
# Runs off-loop (always via asyncio.to_thread); use the sync handle.
|
||||
session_db = getattr(session_db, "_db", session_db)
|
||||
try:
|
||||
raw = session_db.is_telegram_topic_mode_enabled(
|
||||
chat_id=str(source.chat_id),
|
||||
user_id=str(source.user_id),
|
||||
profile_name=self._telegram_topic_profile_name(source),
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Failed to read Telegram topic mode state", exc_info=True)
|
||||
return False
|
||||
# Only a real True from the SessionDB enables topic mode; anything else (including MagicMock
|
||||
# from test fixtures that didn't opt in) means off for this chat.
|
||||
return raw is True
|
||||
|
||||
def _is_telegram_topic_root_lobby(self, source: SessionSource) -> bool:
|
||||
"""True for the main Telegram DM (or General topic) when topic mode has made it a lobby."""
|
||||
if source.platform != Platform.TELEGRAM or source.chat_type != "dm":
|
||||
return False
|
||||
if not self._telegram_topic_mode_enabled(source):
|
||||
return False
|
||||
tid = str(source.thread_id or "")
|
||||
return tid in self._TELEGRAM_GENERAL_TOPIC_IDS
|
||||
|
||||
def _is_telegram_topic_lane(self, source: SessionSource) -> bool:
|
||||
"""True for a user-created Telegram private-chat topic lane."""
|
||||
if source.platform != Platform.TELEGRAM or source.chat_type != "dm":
|
||||
return False
|
||||
if not self._telegram_topic_mode_enabled(source):
|
||||
return False
|
||||
tid = str(source.thread_id or "")
|
||||
return bool(tid) and tid not in self._TELEGRAM_GENERAL_TOPIC_IDS
|
||||
|
||||
def _telegram_topic_cooldown_key(self, source: SessionSource) -> Optional[str]:
|
||||
"""Cooldown key for topic-mode cooldowns: (profile, chat_id).
|
||||
|
||||
Profiles sharing a Telegram private chat_id under multiplex must not
|
||||
suppress each other's lobby reminders / capability hints (#76423).
|
||||
"""
|
||||
chat_id = str(source.chat_id or "")
|
||||
if not chat_id:
|
||||
return None
|
||||
return f"{self._telegram_topic_profile_name(source)}:{chat_id}"
|
||||
|
||||
def _should_send_telegram_lobby_reminder(self, source: SessionSource) -> bool:
|
||||
"""Rate-limit root-DM lobby reminders to one per cooldown window, not one per prompt typed."""
|
||||
if not hasattr(self, "_telegram_lobby_reminder_ts"):
|
||||
self._telegram_lobby_reminder_ts = {}
|
||||
key = self._telegram_topic_cooldown_key(source)
|
||||
if not key:
|
||||
return True
|
||||
import time as _time
|
||||
now = _time.monotonic()
|
||||
last = self._telegram_lobby_reminder_ts.get(key, 0.0)
|
||||
if now - last < self._TELEGRAM_LOBBY_REMINDER_COOLDOWN_S:
|
||||
return False
|
||||
self._telegram_lobby_reminder_ts[key] = now
|
||||
return True
|
||||
|
||||
def _telegram_topic_root_lobby_message(self) -> str:
|
||||
return (
|
||||
"This main chat is reserved for system commands.\n\n"
|
||||
"To start a new Hermes chat, open the All Messages topic at the top "
|
||||
"of this bot interface and send any message there. Telegram will "
|
||||
"create a new topic for that message; each topic works as an "
|
||||
"independent Hermes session."
|
||||
)
|
||||
|
||||
def _telegram_topic_root_new_message(self) -> str:
|
||||
return (
|
||||
"To start a new parallel Hermes chat, open the All Messages topic "
|
||||
"at the top of this bot interface and send any message there. "
|
||||
"Telegram will create a new topic for it.\n\n"
|
||||
"Each topic is an independent Hermes session. Use /new inside an "
|
||||
"existing topic only if you want to replace that topic's current session."
|
||||
)
|
||||
|
||||
def _telegram_topic_new_header(self, source: SessionSource) -> Optional[str]:
|
||||
if not self._is_telegram_topic_lane(source):
|
||||
return None
|
||||
return (
|
||||
"Started a new Hermes session in this topic.\n\n"
|
||||
"Tip: for parallel work, open All Messages and send a message there "
|
||||
"to create a separate topic instead of using /new here. /new replaces "
|
||||
"the session attached to the current topic."
|
||||
)
|
||||
|
||||
def _record_telegram_topic_binding(
|
||||
self,
|
||||
source: SessionSource,
|
||||
session_entry,
|
||||
) -> None:
|
||||
"""Persist the Telegram topic -> Hermes session binding for topic lanes."""
|
||||
session_db = getattr(self, "_session_db", None)
|
||||
if session_db is None or not source.chat_id or not source.thread_id:
|
||||
return
|
||||
# Runs off-loop (always via asyncio.to_thread); use the sync handle.
|
||||
session_db = getattr(session_db, "_db", session_db)
|
||||
session_db.bind_telegram_topic(
|
||||
chat_id=str(source.chat_id),
|
||||
thread_id=str(source.thread_id),
|
||||
user_id=str(source.user_id or ""),
|
||||
session_key=session_entry.session_key,
|
||||
session_id=session_entry.session_id,
|
||||
profile_name=self._telegram_topic_profile_name(source),
|
||||
)
|
||||
|
||||
def _sync_telegram_topic_binding(
|
||||
self,
|
||||
source: SessionSource,
|
||||
session_entry,
|
||||
*,
|
||||
reason: str,
|
||||
) -> None:
|
||||
"""Update the topic binding to point at ``session_entry.session_id``.
|
||||
|
||||
Topic lanes persist (chat_id, thread_id) -> session_id so reopening a topic resumes the
|
||||
right session. When compression rotates the id mid-turn a stale binding reloads the
|
||||
oversized parent next message, retriggering preflight compression — sometimes in a loop.
|
||||
"""
|
||||
if not self._is_telegram_topic_lane(source):
|
||||
return
|
||||
try:
|
||||
self._record_telegram_topic_binding(source, session_entry)
|
||||
except Exception:
|
||||
logger.debug(
|
||||
"telegram topic binding refresh failed (%s)", reason, exc_info=True,
|
||||
)
|
||||
|
||||
def _recover_telegram_topic_thread_id(
|
||||
self,
|
||||
source: SessionSource,
|
||||
) -> Optional[str]:
|
||||
"""Pin DM-topic routing to the user's last-active topic.
|
||||
|
||||
Telegram can omit ``message_thread_id`` or surface General (``1``) for topic-mode DM
|
||||
replies; in those lobby-shaped cases keep the conversation on the user's most-recent bound
|
||||
topic. Do not rewrite a non-lobby, previously-unbound thread id: a brand-new DM topic is
|
||||
also "unknown" until its first inbound message is recorded, and rewriting would send its
|
||||
answer into an older lane. Returns None to leave the source alone.
|
||||
"""
|
||||
if (
|
||||
source.platform != Platform.TELEGRAM
|
||||
or source.chat_type != "dm"
|
||||
or not source.chat_id
|
||||
or not source.user_id
|
||||
or not self._telegram_topic_mode_enabled(source)
|
||||
):
|
||||
return None
|
||||
inbound = str(source.thread_id or "")
|
||||
is_lobby = not inbound or inbound in self._TELEGRAM_GENERAL_TOPIC_IDS
|
||||
if not is_lobby:
|
||||
# A non-lobby, unknown thread_id is likely the first message of a new Telegram DM topic:
|
||||
# preserve it to be recorded as a new lane below rather than hijack the latest binding.
|
||||
return None
|
||||
session_db = getattr(self, "_session_db", None)
|
||||
if session_db is None:
|
||||
return None
|
||||
# Runs off-loop (always via asyncio.to_thread); use the sync handle.
|
||||
session_db = getattr(session_db, "_db", session_db)
|
||||
try:
|
||||
bindings = session_db.list_telegram_topic_bindings_for_chat(
|
||||
chat_id=str(source.chat_id),
|
||||
profile_name=self._telegram_topic_profile_name(source),
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("topic-recover: read failed", exc_info=True)
|
||||
return None
|
||||
if not bindings:
|
||||
return None
|
||||
user_id = str(source.user_id)
|
||||
for b in bindings: # newest-first
|
||||
if str(b.get("user_id") or "") == user_id:
|
||||
recovered = str(b.get("thread_id") or "")
|
||||
if recovered and recovered != inbound:
|
||||
return recovered
|
||||
return None
|
||||
return None
|
||||
|
||||
async def _get_telegram_topic_capabilities(self, source: SessionSource) -> dict:
|
||||
"""Read Telegram private-topic capability flags via Bot API getMe."""
|
||||
adapter = self._adapter_for_source(source)
|
||||
bot = getattr(adapter, "_bot", None)
|
||||
if bot is None or not hasattr(bot, "get_me"):
|
||||
return {"checked": False}
|
||||
try:
|
||||
me = await bot.get_me()
|
||||
except Exception:
|
||||
logger.debug("Failed to fetch Telegram getMe topic capabilities", exc_info=True)
|
||||
return {"checked": False}
|
||||
|
||||
def _field(name: str):
|
||||
if hasattr(me, name):
|
||||
return getattr(me, name)
|
||||
api_kwargs = getattr(me, "api_kwargs", None)
|
||||
if isinstance(api_kwargs, dict) and name in api_kwargs:
|
||||
return api_kwargs.get(name)
|
||||
if isinstance(me, dict):
|
||||
return me.get(name)
|
||||
return None
|
||||
|
||||
return {
|
||||
"checked": True,
|
||||
"has_topics_enabled": _field("has_topics_enabled"),
|
||||
"allows_users_to_create_topics": _field("allows_users_to_create_topics"),
|
||||
}
|
||||
|
||||
async def _ensure_telegram_system_topic(self, source: SessionSource) -> None:
|
||||
"""Create/pin the managed System topic after /topic activation when possible."""
|
||||
adapter = self._adapter_for_source(source)
|
||||
if adapter is None or not source.chat_id:
|
||||
return
|
||||
|
||||
thread_id = None
|
||||
create_topic = getattr(adapter, "_create_dm_topic", None)
|
||||
if callable(create_topic):
|
||||
try:
|
||||
thread_id = await create_topic(int(source.chat_id), "System")
|
||||
except Exception:
|
||||
logger.debug("Failed to create Telegram System topic", exc_info=True)
|
||||
if not thread_id:
|
||||
return
|
||||
|
||||
message_id = None
|
||||
try:
|
||||
send_result = await adapter.send(
|
||||
source.chat_id,
|
||||
"System topic for Hermes commands and status.",
|
||||
metadata={"thread_id": str(thread_id)},
|
||||
)
|
||||
message_id = getattr(send_result, "message_id", None)
|
||||
except Exception:
|
||||
logger.debug("Failed to send Telegram System topic intro", exc_info=True)
|
||||
if not message_id:
|
||||
return
|
||||
|
||||
bot = getattr(adapter, "_bot", None)
|
||||
if bot is None or not hasattr(bot, "pin_chat_message"):
|
||||
return
|
||||
try:
|
||||
await bot.pin_chat_message(
|
||||
chat_id=int(source.chat_id),
|
||||
message_id=int(message_id),
|
||||
disable_notification=True,
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Failed to pin Telegram System topic intro", exc_info=True)
|
||||
|
||||
async def _send_telegram_topic_setup_image(self, source: SessionSource) -> None:
|
||||
"""Send the bundled BotFather Threads Settings screenshot when available."""
|
||||
adapter = self._adapter_for_source(source)
|
||||
if adapter is None or not source.chat_id or not hasattr(adapter, "send_image_file"):
|
||||
return
|
||||
image_path = Path(__file__).resolve().parent / "assets" / "telegram-botfather-threads-settings.jpg"
|
||||
if not image_path.exists():
|
||||
return
|
||||
try:
|
||||
await adapter.send_image_file(
|
||||
chat_id=source.chat_id,
|
||||
image_path=str(image_path),
|
||||
caption="BotFather → Bot Settings → Threads Settings",
|
||||
metadata={"thread_id": str(source.thread_id)} if source.thread_id else None,
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Failed to send Telegram topic setup image", exc_info=True)
|
||||
|
||||
def _sanitize_telegram_topic_title(self, title: str) -> str:
|
||||
"""Return a Bot API-safe forum topic name from a generated session title."""
|
||||
cleaned = re.sub(r"\s+", " ", str(title or "")).strip()
|
||||
if not cleaned:
|
||||
return "Hermes Chat"
|
||||
# Telegram forum topic names are short (currently 1-128 chars). Keep
|
||||
# extra room for multi-byte titles and avoid trailing ellipsis churn.
|
||||
if len(cleaned) > 120:
|
||||
cleaned = cleaned[:117].rstrip() + "..."
|
||||
return cleaned
|
||||
|
||||
def _is_discord_auto_thread_lane(self, source: SessionSource) -> bool:
|
||||
"""Return True only for Discord threads Hermes just auto-created."""
|
||||
return (
|
||||
source.platform == Platform.DISCORD
|
||||
and source.chat_type == "thread"
|
||||
and bool(getattr(source, "auto_thread_created", False))
|
||||
and bool(source.thread_id)
|
||||
and bool(getattr(source, "auto_thread_initial_name", None))
|
||||
)
|
||||
|
||||
def _is_relay_discord_channel_lane(self, source: SessionSource) -> bool:
|
||||
"""Shape-only check: a relay-delivered Discord CHANNEL event whose
|
||||
reply the connector MAY auto-thread (title-turn registration gate).
|
||||
|
||||
Deliberately does NOT consult the send-result cache: at registration
|
||||
time (before delivery) the feedback can't exist yet. The rename lane
|
||||
polls the cache at fire time instead."""
|
||||
return (
|
||||
source.platform == Platform.DISCORD
|
||||
and bool(source.chat_id)
|
||||
and not source.thread_id
|
||||
and source.chat_type in ("group", "channel")
|
||||
and getattr(source, "delivered_via_upstream_relay", False) is True
|
||||
)
|
||||
|
||||
def _relay_auto_thread_info(
|
||||
self, source: SessionSource
|
||||
) -> Optional[Tuple[str, str]]:
|
||||
"""(thread_id, initial_name) when the RELAY connector auto-threaded our reply to this
|
||||
source's chat — the title-turn sibling of _is_discord_auto_thread_lane.
|
||||
|
||||
The marker check only matches events ARRIVING IN an auto-created thread (turn 2+); the
|
||||
auto-title fires on the FIRST exchange, whose source is the PARENT channel event with no
|
||||
markers. Preferred: the connector's ``prospective_thread_id`` stamp (anchor message id ==
|
||||
the thread it will create) — per-message, so it names the EXACT thread even when several
|
||||
auto-threads spawn from one channel; the connector's created-name guard enforces
|
||||
no-clobber. Fallback: the per-chat send-result thread_id/auto_thread_name cache (older
|
||||
connectors), which only ever renamed the FIRST thread.
|
||||
"""
|
||||
from gateway.run import _as_thread_info
|
||||
if source.platform != Platform.DISCORD or not source.chat_id:
|
||||
return None
|
||||
if not getattr(source, "delivered_via_upstream_relay", False):
|
||||
return None
|
||||
prospective = getattr(source, "prospective_thread_id", None)
|
||||
if prospective:
|
||||
# Deterministic per-thread identity; the empty initial-name marker
|
||||
# signals the caller to rely on the connector-side no-clobber guard.
|
||||
return (str(prospective), "")
|
||||
adapter = self._adapter_for_source(source)
|
||||
info_fn = getattr(adapter, "auto_thread_info_for_chat", None)
|
||||
if not callable(info_fn):
|
||||
return None
|
||||
try:
|
||||
return _as_thread_info(info_fn(str(source.chat_id)))
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
async def _await_relay_auto_thread_info(
|
||||
self, source: SessionSource
|
||||
) -> Optional[Tuple[str, str]]:
|
||||
"""``_relay_auto_thread_info``, waited out until this turn delivers.
|
||||
|
||||
The legacy send-result path can only answer once the reply is sent, and the caller asks
|
||||
at title time — one turn early. The adapter answers on the send either way, so the
|
||||
timeout is only a backstop for a turn that never sends at all; the turn's own inactivity
|
||||
limit is exactly how long that turn could still be alive.
|
||||
"""
|
||||
from gateway.run import _as_thread_info, _float_env
|
||||
# The connector-stamped prospective id is known at ingest, so most
|
||||
# sessions answer here and never wait at all.
|
||||
known = self._relay_auto_thread_info(source)
|
||||
if known is not None:
|
||||
return known
|
||||
adapter = self._adapter_for_source(source)
|
||||
wait_fn = getattr(adapter, "wait_for_auto_thread_info", None)
|
||||
if not callable(wait_fn) or not source.chat_id:
|
||||
return None
|
||||
# 0 means the operator disabled the turn limit; the backstop still needs one.
|
||||
timeout = _float_env("HERMES_AGENT_TIMEOUT", 1800) or 1800
|
||||
try:
|
||||
return _as_thread_info(await wait_fn(str(source.chat_id), timeout))
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def _sanitize_discord_thread_title(self, title: str) -> str:
|
||||
"""Return a Discord-safe semantic thread title from a session title.
|
||||
|
||||
Discord thread names are capped at 100 characters measured in UTF-16 code units (emoji
|
||||
count double), so truncate with the UTF-16 helpers rather than Python code-point slices.
|
||||
"""
|
||||
cleaned = re.sub(r"\s+", " ", str(title or "")).strip()
|
||||
if not cleaned:
|
||||
return "Hermes Chat"
|
||||
if utf16_len(cleaned) > 80:
|
||||
cleaned = _prefix_within_utf16_limit(cleaned, 77).rstrip() + "..."
|
||||
return cleaned
|
||||
|
||||
async def _rename_discord_auto_thread_for_session_title(
|
||||
self,
|
||||
source: SessionSource,
|
||||
session_id: str,
|
||||
title: str,
|
||||
relay_info: Optional[Tuple[str, str]] = None,
|
||||
) -> None:
|
||||
"""Best-effort semantic rename of a newly auto-created Discord thread.
|
||||
|
||||
``relay_info`` is the (thread_id, initial_name) pair from the relay connector's send-
|
||||
result feedback — supplied on the title turn, where the source is the parent-channel
|
||||
event and carries no auto-thread markers (see _relay_auto_thread_info).
|
||||
"""
|
||||
if relay_info is None and not await asyncio.to_thread(
|
||||
self._is_discord_auto_thread_lane, source
|
||||
):
|
||||
# Relay title turn with no feedback captured at schedule time: the title comes off the
|
||||
# user's opening message, so it beats the delivery that produces the connector's send-
|
||||
# result feedback (thread_id + initial name) by the whole length of the turn.
|
||||
if not self._is_relay_discord_channel_lane(source):
|
||||
return
|
||||
relay_info = await self._await_relay_auto_thread_info(source)
|
||||
if relay_info is None:
|
||||
# True miss: the connector did not auto-thread this reply
|
||||
# (policy off, DM, already-threaded, or send failed).
|
||||
return
|
||||
adapter = self._adapter_for_source(source) if getattr(self, "adapters", None) else None
|
||||
if adapter is None:
|
||||
return
|
||||
rename_thread = getattr(adapter, "rename_thread", None)
|
||||
if rename_thread is None:
|
||||
return
|
||||
target_thread_id = relay_info[0] if relay_info else str(source.thread_id)
|
||||
# Relay lane (relay_info present): ask the CONNECTOR to enforce the no-clobber guard from
|
||||
# its own created-name memory — the gateway can't reliably reproduce the thread's initial
|
||||
# name byte-for-byte (normalization drift silently declined every rename before this).
|
||||
use_connector_guard = relay_info is not None
|
||||
guard_name = (
|
||||
None
|
||||
if use_connector_guard
|
||||
else getattr(source, "auto_thread_initial_name", None)
|
||||
)
|
||||
thread_name = self._sanitize_discord_thread_title(title)
|
||||
# Relay lane only: the connector's egress guard resolves the owning tenant from the
|
||||
# outbound scope_id/user_id caches, keyed by the PARENT channel chat_id (learned at
|
||||
# inbound), not the thread id. rename_thread defaults chat_id to the thread id, so the
|
||||
# lookup misses and the connector declines; pass the parent channel id (the relay source's
|
||||
# chat_id). Native lane needs nothing: its source IS the thread, direct Discord API.
|
||||
parent_chat_id = (
|
||||
str(source.chat_id) if use_connector_guard and source.chat_id else None
|
||||
)
|
||||
logger.info(
|
||||
"discord auto-thread rename: thread=%s lane=%s new_title=%r",
|
||||
target_thread_id,
|
||||
"relay" if use_connector_guard else "native",
|
||||
thread_name,
|
||||
)
|
||||
rename_kwargs = (
|
||||
{
|
||||
"prefer_connector_created": True,
|
||||
"parent_chat_id": parent_chat_id,
|
||||
}
|
||||
if use_connector_guard
|
||||
else {"only_if_current_name": guard_name}
|
||||
)
|
||||
try:
|
||||
renamed = await rename_thread(
|
||||
target_thread_id,
|
||||
thread_name,
|
||||
**rename_kwargs,
|
||||
)
|
||||
logger.info(
|
||||
"discord auto-thread rename result: thread=%s applied=%s",
|
||||
target_thread_id,
|
||||
bool(renamed),
|
||||
)
|
||||
except TypeError:
|
||||
logger.warning(
|
||||
"Discord semantic thread rename raised TypeError (adapter=%s)",
|
||||
type(adapter).__name__,
|
||||
exc_info=True,
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Failed to rename Discord auto-thread for generated session title", exc_info=True)
|
||||
|
||||
def _schedule_rename_from_title_thread(self, source: SessionSource, make_coro, label: str) -> None:
|
||||
"""Schedule a best-effort rename coroutine onto the gateway loop from the auto-title thread.
|
||||
|
||||
The source is copied so the background thread never shares the live dataclass with the
|
||||
loop; failures are logged at debug and never propagate."""
|
||||
from gateway.run import safe_schedule_threadsafe
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
loop = getattr(self, "_gateway_loop", None)
|
||||
if loop is None or loop.is_closed():
|
||||
return
|
||||
try:
|
||||
copied_source = dataclasses.replace(source)
|
||||
except Exception:
|
||||
copied_source = source
|
||||
future = safe_schedule_threadsafe(
|
||||
make_coro(copied_source),
|
||||
loop,
|
||||
logger=logger,
|
||||
log_message=f"{label} failed to schedule",
|
||||
)
|
||||
if future is None:
|
||||
return
|
||||
|
||||
def _log_rename_failure(fut) -> None:
|
||||
try:
|
||||
fut.result()
|
||||
except Exception:
|
||||
logger.debug("%s failed", label, exc_info=True)
|
||||
|
||||
future.add_done_callback(_log_rename_failure)
|
||||
|
||||
def _schedule_discord_semantic_thread_rename(
|
||||
self,
|
||||
source: SessionSource,
|
||||
session_id: str,
|
||||
title: str,
|
||||
) -> None:
|
||||
"""Schedule Discord auto-thread rename from the auto-title background thread."""
|
||||
relay_info = None
|
||||
if not title:
|
||||
return
|
||||
if not self._is_discord_auto_thread_lane(source):
|
||||
# Relay title turn: the source is the PARENT channel event (thread didn't exist at
|
||||
# ingest, no auto-thread markers). The connector's send-result feedback says where the
|
||||
# reply landed, but the auto-title races that delivery, so a cache miss HERE is not a
|
||||
# verdict. Schedule whenever the SHAPE matches; the async rename lane polls the cache
|
||||
# (bounded wait) and no-ops on a true miss.
|
||||
relay_info = self._relay_auto_thread_info(source)
|
||||
if relay_info is None and not self._is_relay_discord_channel_lane(
|
||||
source
|
||||
):
|
||||
return
|
||||
self._schedule_rename_from_title_thread(
|
||||
source,
|
||||
lambda copied: self._rename_discord_auto_thread_for_session_title(
|
||||
copied, session_id, title, relay_info=relay_info
|
||||
),
|
||||
"Discord semantic thread rename",
|
||||
)
|
||||
|
||||
async def _rename_telegram_topic_for_session_title(
|
||||
self,
|
||||
source: SessionSource,
|
||||
session_id: str,
|
||||
title: str,
|
||||
) -> None:
|
||||
"""Best-effort rename of a Telegram DM topic when Hermes auto-titles a session."""
|
||||
if not await asyncio.to_thread(self._is_telegram_topic_lane, source) or not source.chat_id or not source.thread_id:
|
||||
return
|
||||
|
||||
# extra.disable_topic_auto_rename lets the operator disable per-topic auto-rename entirely,
|
||||
# e.g. user-managed topics (ad-hoc Threaded Mode) that auto-rename would keep overwriting.
|
||||
if self._telegram_topic_auto_rename_disabled(source):
|
||||
return
|
||||
|
||||
# Skip rename when the topic is operator-declared via extra.dm_topics. Those topics have
|
||||
# fixed names chosen by the operator (plus optional skill binding); auto-renaming would
|
||||
# silently mutate operator config. Check the class, not the instance — getattr() on a
|
||||
# MagicMock auto-creates attributes, so an instance hasattr() is True for every test double.
|
||||
adapter = self._adapter_for_source(source)
|
||||
if adapter is not None:
|
||||
get_info = getattr(type(adapter), "_get_dm_topic_info", None)
|
||||
if callable(get_info):
|
||||
try:
|
||||
operator_topic = get_info(adapter, str(source.chat_id), str(source.thread_id))
|
||||
except Exception:
|
||||
operator_topic = None
|
||||
# Only treat dict-shaped returns as operator-declared; a
|
||||
# bare MagicMock or other sentinel shouldn't count.
|
||||
if isinstance(operator_topic, dict):
|
||||
return
|
||||
|
||||
session_db = getattr(self, "_session_db", None)
|
||||
if session_db is not None:
|
||||
try:
|
||||
binding = await session_db.get_telegram_topic_binding(
|
||||
chat_id=str(source.chat_id),
|
||||
thread_id=str(source.thread_id),
|
||||
profile_name=self._telegram_topic_profile_name(source),
|
||||
)
|
||||
if binding and str(binding.get("session_id") or "") != str(session_id):
|
||||
return
|
||||
except Exception:
|
||||
logger.debug("Failed to verify Telegram topic binding before rename", exc_info=True)
|
||||
return
|
||||
|
||||
if adapter is None:
|
||||
return
|
||||
topic_name = self._sanitize_telegram_topic_title(title)
|
||||
try:
|
||||
rename_topic = getattr(adapter, "rename_dm_topic", None)
|
||||
if rename_topic is not None:
|
||||
await rename_topic(
|
||||
chat_id=str(source.chat_id),
|
||||
thread_id=str(source.thread_id),
|
||||
name=topic_name,
|
||||
)
|
||||
return
|
||||
|
||||
bot = getattr(adapter, "_bot", None)
|
||||
edit_forum_topic = getattr(bot, "edit_forum_topic", None) if bot is not None else None
|
||||
if edit_forum_topic is None:
|
||||
edit_forum_topic = getattr(bot, "editForumTopic", None) if bot is not None else None
|
||||
if edit_forum_topic is None:
|
||||
return
|
||||
try:
|
||||
await edit_forum_topic(
|
||||
chat_id=int(source.chat_id),
|
||||
message_thread_id=int(source.thread_id),
|
||||
name=topic_name,
|
||||
)
|
||||
except (TypeError, ValueError):
|
||||
await edit_forum_topic(
|
||||
chat_id=source.chat_id,
|
||||
message_thread_id=source.thread_id,
|
||||
name=topic_name,
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Failed to rename Telegram topic for auto-generated title", exc_info=True)
|
||||
|
||||
def _telegram_topic_auto_rename_disabled(self, source: SessionSource) -> bool:
|
||||
"""Return True when operator disabled per-topic auto-rename for this Telegram chat.
|
||||
|
||||
``gateway.platforms.telegram.extra.disable_topic_auto_rename``; default False (auto-rename on).
|
||||
"""
|
||||
platform_cfg = (
|
||||
self.config.platforms.get(source.platform)
|
||||
if getattr(self, "config", None) and getattr(self.config, "platforms", None)
|
||||
else None
|
||||
)
|
||||
if platform_cfg is None:
|
||||
return False
|
||||
extra = getattr(platform_cfg, "extra", None) or {}
|
||||
value = extra.get("disable_topic_auto_rename")
|
||||
if value is None:
|
||||
return False
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
return value.strip().lower() in {"1", "true", "yes", "on"}
|
||||
return bool(value)
|
||||
|
||||
def _schedule_telegram_topic_title_rename(
|
||||
self,
|
||||
source: SessionSource,
|
||||
session_id: str,
|
||||
title: str,
|
||||
) -> None:
|
||||
"""Schedule a topic rename from the auto-title background thread."""
|
||||
if not title or not self._is_telegram_topic_lane(source):
|
||||
return
|
||||
if self._telegram_topic_auto_rename_disabled(source):
|
||||
return
|
||||
self._schedule_rename_from_title_thread(
|
||||
source,
|
||||
lambda copied: self._rename_telegram_topic_for_session_title(copied, session_id, title),
|
||||
"Telegram topic title rename",
|
||||
)
|
||||
|
||||
def _should_send_telegram_capability_hint(self, source: SessionSource) -> bool:
|
||||
"""Rate-limit the BotFather Threads Settings screenshot.
|
||||
|
||||
Repeated /topic while Threads Settings are still off must not re-upload it every time.
|
||||
"""
|
||||
if not hasattr(self, "_telegram_capability_hint_ts"):
|
||||
self._telegram_capability_hint_ts = {}
|
||||
key = self._telegram_topic_cooldown_key(source)
|
||||
if not key:
|
||||
return True
|
||||
import time as _time
|
||||
now = _time.monotonic()
|
||||
last = self._telegram_capability_hint_ts.get(key, 0.0)
|
||||
if now - last < self._TELEGRAM_CAPABILITY_HINT_COOLDOWN_S:
|
||||
return False
|
||||
self._telegram_capability_hint_ts[key] = now
|
||||
return True
|
||||
|
||||
def _telegram_topic_help_text(self) -> str:
|
||||
return (
|
||||
"/topic — enable multi-session DM mode (one bot, many parallel chats)\n"
|
||||
"\n"
|
||||
"Usage:\n"
|
||||
" /topic Enable topic mode, or show status if already on\n"
|
||||
" /topic help Show this message\n"
|
||||
" /topic off Disable topic mode and clear topic bindings\n"
|
||||
" /topic <id> Inside a topic: restore a previous session by ID\n"
|
||||
"\n"
|
||||
"How it works:\n"
|
||||
"1. Run /topic once in this DM — Hermes checks BotFather Threads\n"
|
||||
" Settings are enabled and flips on multi-session mode.\n"
|
||||
"2. Tap All Messages at the top of the bot and send any message.\n"
|
||||
" Telegram creates a new topic for that message; each topic is\n"
|
||||
" an independent Hermes session (fresh history, fresh context).\n"
|
||||
"3. The root DM becomes a system lobby — send /topic, /status,\n"
|
||||
" /help, /usage there. Normal prompts go in a topic.\n"
|
||||
"4. /new inside a topic resets just that topic's session.\n"
|
||||
"5. /topic <id> inside a topic restores an old session into it."
|
||||
)
|
||||
|
||||
async def _disable_telegram_topic_mode_for_chat(self, source: SessionSource) -> str:
|
||||
"""Cleanly disable topic mode for a chat via /topic off."""
|
||||
if not self._session_db:
|
||||
from hermes_state import format_session_db_unavailable
|
||||
return format_session_db_unavailable(prefix=t("gateway.shared.session_db_unavailable_prefix"))
|
||||
chat_id = str(source.chat_id or "")
|
||||
if not chat_id:
|
||||
return "Could not determine chat ID."
|
||||
# No-op if never enabled.
|
||||
try:
|
||||
currently_enabled = await self._session_db.is_telegram_topic_mode_enabled(
|
||||
chat_id=chat_id,
|
||||
user_id=str(source.user_id or ""),
|
||||
profile_name=self._telegram_topic_profile_name(source),
|
||||
)
|
||||
except Exception:
|
||||
currently_enabled = False
|
||||
if not currently_enabled:
|
||||
return "Multi-session topic mode is not currently enabled for this chat."
|
||||
try:
|
||||
await self._session_db.disable_telegram_topic_mode(
|
||||
chat_id=chat_id,
|
||||
profile_name=self._telegram_topic_profile_name(source),
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.exception("Failed to disable Telegram topic mode")
|
||||
return f"Failed to disable topic mode: {exc}"
|
||||
# Reset per-profile+chat debounce state so the user doesn't see a
|
||||
# stale cooldown on the next activation (issue #76423).
|
||||
cooldown_key = self._telegram_topic_cooldown_key(source)
|
||||
if cooldown_key:
|
||||
for attr in ("_telegram_lobby_reminder_ts", "_telegram_capability_hint_ts"):
|
||||
store = getattr(self, attr, None)
|
||||
if isinstance(store, dict):
|
||||
store.pop(cooldown_key, None)
|
||||
return (
|
||||
"Multi-session topic mode is now OFF for this chat.\n\n"
|
||||
"Existing topics in Telegram aren't removed — they'll just stop "
|
||||
"being gated as independent sessions. The root DM works as a "
|
||||
"normal Hermes chat again. Run /topic to re-enable later."
|
||||
)
|
||||
|
||||
async def _telegram_topic_root_status_message(self, source: SessionSource) -> str:
|
||||
lines = [
|
||||
"Telegram multi-session topics are enabled.",
|
||||
"",
|
||||
"To create a new Hermes chat, open All Messages at the top of this "
|
||||
"bot interface and send any message there. Telegram will create a "
|
||||
"new topic for it.",
|
||||
"",
|
||||
]
|
||||
try:
|
||||
sessions = await self._session_db.list_unlinked_telegram_sessions_for_user(
|
||||
chat_id=str(source.chat_id),
|
||||
user_id=str(source.user_id),
|
||||
profile_name=self._telegram_topic_profile_name(source),
|
||||
limit=10,
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Failed to list unlinked Telegram sessions", exc_info=True)
|
||||
sessions = []
|
||||
|
||||
if sessions:
|
||||
lines.append("Previous unlinked sessions:")
|
||||
for session in sessions:
|
||||
session_id = str(session.get("id") or "")
|
||||
title = str(session.get("title") or "Untitled session")
|
||||
preview = str(session.get("preview") or "").strip()
|
||||
line = f"- {title} — `{session_id}`"
|
||||
if preview:
|
||||
line += f" — {preview}"
|
||||
lines.append(line)
|
||||
lines.extend([
|
||||
"",
|
||||
"To restore one:",
|
||||
"1. Create or open a topic. To create a new one, open All Messages and send any message there.",
|
||||
"2. Send /topic <session-id> inside that topic.",
|
||||
f"Example: Send /topic {sessions[0].get('id')} inside a topic.",
|
||||
])
|
||||
else:
|
||||
lines.extend([
|
||||
"No previous unlinked Telegram sessions found.",
|
||||
"",
|
||||
"To restore a previous session later:",
|
||||
"1. Create or open a topic. To create a new one, open All Messages and send any message there.",
|
||||
"2. Send /topic <session-id> inside that topic.",
|
||||
])
|
||||
return "\n".join(lines)
|
||||
|
||||
async def _restore_telegram_topic_session(self, event: MessageEvent, raw_session_id: str) -> str:
|
||||
"""Restore an existing Telegram-owned Hermes session into this topic."""
|
||||
source = event.source
|
||||
session_id = await self._session_db.resolve_session_id(raw_session_id.strip())
|
||||
if not session_id:
|
||||
return f"Session not found: {raw_session_id.strip()}"
|
||||
|
||||
session = await self._session_db.get_session(session_id)
|
||||
if not session:
|
||||
return f"Session not found: {raw_session_id.strip()}"
|
||||
if str(session.get("source") or "") != "telegram":
|
||||
return "That session is not a Telegram session and cannot be restored into this topic."
|
||||
if str(session.get("user_id") or "") != str(source.user_id):
|
||||
return "That session does not belong to this Telegram user."
|
||||
|
||||
linked = await self._session_db.is_telegram_session_linked_to_topic(session_id=session_id)
|
||||
topic_profile = self._telegram_topic_profile_name(source)
|
||||
current_binding = await self._session_db.get_telegram_topic_binding(
|
||||
chat_id=str(source.chat_id),
|
||||
thread_id=str(source.thread_id),
|
||||
profile_name=topic_profile,
|
||||
)
|
||||
if linked:
|
||||
if not current_binding or current_binding.get("session_id") != session_id:
|
||||
return "That session is already linked to another Telegram topic."
|
||||
|
||||
session_key = self._session_key_for_source(source)
|
||||
try:
|
||||
await self._session_db.bind_telegram_topic(
|
||||
chat_id=str(source.chat_id),
|
||||
thread_id=str(source.thread_id),
|
||||
user_id=str(source.user_id),
|
||||
session_key=session_key,
|
||||
session_id=session_id,
|
||||
managed_mode="restored",
|
||||
profile_name=topic_profile,
|
||||
)
|
||||
except ValueError as exc:
|
||||
if "already linked" in str(exc):
|
||||
return "That session is already linked to another Telegram topic."
|
||||
raise
|
||||
|
||||
title = await self._session_db.get_session_title(session_id) or session_id
|
||||
last_assistant = None
|
||||
try:
|
||||
for message in reversed(await self._session_db.get_messages(session_id)):
|
||||
if message.get("role") != "assistant":
|
||||
continue
|
||||
projected = project_compaction_message_for_display(message)
|
||||
if projected is not None and projected.get("content"):
|
||||
last_assistant = str(projected.get("content"))
|
||||
break
|
||||
except Exception:
|
||||
last_assistant = None
|
||||
|
||||
response = f"Session restored: {title}"
|
||||
if last_assistant:
|
||||
response += f"\n\nLast Hermes message:\n{last_assistant}"
|
||||
return response
|
||||
5937
gateway/run_turn.py
Normal file
5937
gateway/run_turn.py
Normal file
File diff suppressed because it is too large
Load Diff
2402
gateway/run_turn_runner.py
Normal file
2402
gateway/run_turn_runner.py
Normal file
File diff suppressed because it is too large
Load Diff
557
gateway/run_voice.py
Normal file
557
gateway/run_voice.py
Normal file
@@ -0,0 +1,557 @@
|
||||
"""Voice-channel / auto-TTS methods for GatewayRunner.
|
||||
|
||||
Split out of ``gateway/run.py``; bound onto ``GatewayRunner`` via the MRO.
|
||||
``gateway.run`` internals are imported lazily inside method bodies (import cycle),
|
||||
so ``patch("gateway.run.X")`` keeps intercepting them at call time.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
import asyncio
|
||||
import functools
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
import time
|
||||
from contextlib import suppress
|
||||
from gateway.config import Platform
|
||||
from gateway.platforms.base import MessageEvent, MessageType, build_auto_tts_output_path
|
||||
from gateway.session import SessionSource
|
||||
from typing import Any, Awaitable, Callable, Dict, List, Optional, cast
|
||||
|
||||
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")
|
||||
|
||||
|
||||
class GatewayVoiceMixin:
|
||||
"""Voice-channel / auto-TTS methods for GatewayRunner."""
|
||||
|
||||
def _voice_key(
|
||||
self, platform: Platform, chat_id: str, profile: Optional[str] = None
|
||||
) -> str:
|
||||
"""Return a platform-namespaced key for voice mode state.
|
||||
|
||||
Under multiplexing the key is ``<profile>:<platform>:<chat_id>`` (profile whose bot speaks);
|
||||
the default profile keeps ``<platform>:<chat_id>`` so persisted state stays valid. Otherwise
|
||||
two bots in one Discord channel share a key and one profile's ``/voice`` flips the other's.
|
||||
"""
|
||||
base = f"{platform.value}:{chat_id}"
|
||||
profile = profile.strip() if isinstance(profile, str) else ""
|
||||
if not profile or profile == "default":
|
||||
return base
|
||||
return f"{profile}:{base}"
|
||||
|
||||
def _voice_key_for_source(self, source: SessionSource) -> str:
|
||||
"""Voice-state key for an inbound source, namespaced by its transport owner.
|
||||
|
||||
Voice mode belongs to the (bot, chat) pair, so the namespace is the profile that OWNS the
|
||||
receiving adapter (matching ``_sync_voice_mode_state_to_adapter``), not the routed profile.
|
||||
"""
|
||||
return self._voice_key(
|
||||
source.platform,
|
||||
source.chat_id,
|
||||
profile=self._adapter_profile_for_source(source),
|
||||
)
|
||||
|
||||
def _bind_voice_input_callback(self, adapter) -> None:
|
||||
"""Route voice transcripts back through the adapter that captured them."""
|
||||
if hasattr(adapter, "_voice_input_callback"):
|
||||
adapter._voice_input_callback = functools.partial(
|
||||
self._handle_voice_channel_input, adapter=adapter
|
||||
)
|
||||
|
||||
def _load_voice_modes(self) -> Dict[str, str]:
|
||||
try:
|
||||
data = json.loads(self._VOICE_MODE_PATH.read_text(encoding="utf-8"))
|
||||
except (FileNotFoundError, json.JSONDecodeError, OSError):
|
||||
return {}
|
||||
|
||||
if not isinstance(data, dict):
|
||||
return {}
|
||||
|
||||
valid_modes = {"off", "voice_only", "all"}
|
||||
result = {}
|
||||
for chat_id, mode in data.items():
|
||||
if mode not in valid_modes:
|
||||
continue
|
||||
key = str(chat_id)
|
||||
# Skip legacy unprefixed keys (warn and skip)
|
||||
if ":" not in key:
|
||||
logger.warning(
|
||||
"Skipping legacy unprefixed voice mode key %r during migration. "
|
||||
"Re-enable voice mode on that chat to rebuild the prefixed key.",
|
||||
key,
|
||||
)
|
||||
continue
|
||||
result[key] = mode
|
||||
return result
|
||||
|
||||
def _save_voice_modes(self) -> None:
|
||||
try:
|
||||
self._VOICE_MODE_PATH.parent.mkdir(parents=True, exist_ok=True)
|
||||
self._VOICE_MODE_PATH.write_text(
|
||||
json.dumps(self._voice_mode, indent=2), encoding="utf-8"
|
||||
)
|
||||
except OSError as e:
|
||||
logger.warning("Failed to save voice modes: %s", e)
|
||||
|
||||
@staticmethod
|
||||
def _toggle_adapter_auto_tts_set(adapter, chat_id: str, on: bool, *, add_to: str, clear_from: str) -> None:
|
||||
"""Add/discard ``chat_id`` in the adapter's ``add_to`` set; adding also clears it from ``clear_from``.
|
||||
|
||||
``/voice off`` and an explicit ``/voice on``/``/voice tts`` are hard overrides of each other."""
|
||||
target = getattr(adapter, add_to, None)
|
||||
if not isinstance(target, set):
|
||||
return
|
||||
if on:
|
||||
target.add(chat_id)
|
||||
other = getattr(adapter, clear_from, None)
|
||||
if isinstance(other, set):
|
||||
other.discard(chat_id)
|
||||
else:
|
||||
target.discard(chat_id)
|
||||
|
||||
def _set_adapter_auto_tts_disabled(self, adapter, chat_id: str, disabled: bool) -> None:
|
||||
"""Update an adapter's in-memory auto-TTS suppression set if present."""
|
||||
self._toggle_adapter_auto_tts_set(
|
||||
adapter, chat_id, disabled, add_to="_auto_tts_disabled_chats", clear_from="_auto_tts_enabled_chats"
|
||||
)
|
||||
|
||||
def _set_adapter_auto_tts_enabled(self, adapter, chat_id: str, enabled: bool) -> None:
|
||||
"""Update an adapter's per-chat auto-TTS opt-in set (auto-TTS even when ``voice.auto_tts`` is False)."""
|
||||
self._toggle_adapter_auto_tts_set(
|
||||
adapter, chat_id, enabled, add_to="_auto_tts_enabled_chats", clear_from="_auto_tts_disabled_chats"
|
||||
)
|
||||
|
||||
def _sync_voice_mode_state_to_adapter(self, adapter) -> None:
|
||||
"""Restore persisted /voice state into a live platform adapter.
|
||||
|
||||
Sets ``_auto_tts_default`` (from ``voice.auto_tts``) and, from ``self._voice_mode``,
|
||||
``_auto_tts_enabled_chats`` (modes ``voice_only``/``all``) and ``_auto_tts_disabled_chats``
|
||||
(mode ``off``).
|
||||
"""
|
||||
platform = getattr(adapter, "platform", None)
|
||||
if not isinstance(platform, Platform):
|
||||
return
|
||||
|
||||
disabled_chats = getattr(adapter, "_auto_tts_disabled_chats", None)
|
||||
enabled_chats = getattr(adapter, "_auto_tts_enabled_chats", None)
|
||||
if not isinstance(disabled_chats, set) and not isinstance(enabled_chats, set):
|
||||
return
|
||||
|
||||
# Push the global voice.auto_tts default (config.yaml) onto the adapter.
|
||||
# Lazy import to avoid adding a module-level dep from gateway → hermes_cli.
|
||||
try:
|
||||
from hermes_cli.config import load_config as _load_full_config
|
||||
_full_cfg = _load_full_config()
|
||||
_auto_tts_default = bool(
|
||||
(_full_cfg.get("voice") or {}).get("auto_tts", False)
|
||||
)
|
||||
except Exception:
|
||||
_auto_tts_default = False
|
||||
if hasattr(adapter, "_auto_tts_default"):
|
||||
adapter._auto_tts_default = _auto_tts_default
|
||||
|
||||
prefix = self._voice_key(platform, "", profile=getattr(adapter, "_owner_profile", None))
|
||||
if isinstance(disabled_chats, set):
|
||||
disabled_chats.clear()
|
||||
disabled_chats.update(
|
||||
key[len(prefix):] for key, mode in self._voice_mode.items()
|
||||
if mode == "off" and key.startswith(prefix)
|
||||
)
|
||||
if isinstance(enabled_chats, set):
|
||||
enabled_chats.clear()
|
||||
enabled_chats.update(
|
||||
key[len(prefix):] for key, mode in self._voice_mode.items()
|
||||
if mode in {"voice_only", "all"} and key.startswith(prefix)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_guild_id(event: MessageEvent) -> Optional[int]:
|
||||
"""Extract Discord guild_id from the raw message object."""
|
||||
raw = getattr(event, "raw_message", None)
|
||||
if raw is None:
|
||||
return None
|
||||
# Slash command interaction
|
||||
if hasattr(raw, "guild_id") and raw.guild_id:
|
||||
return int(raw.guild_id)
|
||||
# Regular message
|
||||
if hasattr(raw, "guild") and raw.guild:
|
||||
return raw.guild.id
|
||||
return None
|
||||
|
||||
async def _handle_voice_channel_join(self, event: MessageEvent) -> str:
|
||||
"""Join the user's current Discord voice channel."""
|
||||
adapter = self._adapter_for_source(event.source)
|
||||
if not hasattr(adapter, "join_voice_channel"):
|
||||
return "Voice channels are not supported on this platform."
|
||||
|
||||
guild_id = self._get_guild_id(event)
|
||||
if not guild_id:
|
||||
return "This command only works in a Discord server."
|
||||
|
||||
voice_channel = await adapter.get_user_voice_channel(
|
||||
guild_id, event.source.user_id
|
||||
)
|
||||
if not voice_channel:
|
||||
return "You need to be in a voice channel first."
|
||||
|
||||
# Wire callbacks BEFORE join so voice input arriving immediately
|
||||
# after connection is not lost.
|
||||
self._bind_voice_input_callback(adapter)
|
||||
voice_profile = self._adapter_profile_for_source(event.source)
|
||||
if hasattr(adapter, "_on_voice_disconnect"):
|
||||
adapter._on_voice_disconnect = functools.partial(
|
||||
self._handle_voice_timeout_cleanup, adapter=adapter
|
||||
)
|
||||
# Let the adapter's inactivity timer see the live voice-reply mode so it
|
||||
# doesn't disconnect a deliberately text-only (/voice off) session.
|
||||
if hasattr(adapter, "_voice_mode_getter"):
|
||||
adapter._voice_mode_getter = lambda chat_id: self._voice_mode.get(
|
||||
self._voice_key(Platform.DISCORD, str(chat_id), profile=voice_profile),
|
||||
"off",
|
||||
)
|
||||
|
||||
try:
|
||||
success = await adapter.join_voice_channel(voice_channel)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to join voice channel: %s", e)
|
||||
adapter._voice_input_callback = None
|
||||
err_lower = str(e).lower()
|
||||
if "pynacl" in err_lower or "nacl" in err_lower or "davey" in err_lower:
|
||||
return (
|
||||
"Voice dependencies are missing (PyNaCl / davey). "
|
||||
f"Install with: `{sys.executable} -m pip install PyNaCl`"
|
||||
)
|
||||
return f"Failed to join voice channel: {e}"
|
||||
|
||||
if success:
|
||||
adapter._voice_text_channels[guild_id] = int(event.source.chat_id)
|
||||
if hasattr(adapter, "_voice_sources"):
|
||||
adapter._voice_sources[guild_id] = event.source.to_dict()
|
||||
self._voice_mode[self._voice_key_for_source(event.source)] = "all"
|
||||
self._save_voice_modes()
|
||||
self._set_adapter_auto_tts_enabled(adapter, event.source.chat_id, enabled=True)
|
||||
return (
|
||||
f"Joined voice channel **{voice_channel.name}**.\n"
|
||||
f"I'll speak my replies and listen to you. Use /voice leave to disconnect."
|
||||
)
|
||||
# Join failed — clear callback
|
||||
adapter._voice_input_callback = None
|
||||
return "Failed to join voice channel. Check bot permissions (Connect + Speak)."
|
||||
|
||||
async def _handle_voice_channel_leave(self, event: MessageEvent) -> str:
|
||||
"""Leave the Discord voice channel."""
|
||||
adapter = self._adapter_for_source(event.source)
|
||||
guild_id = self._get_guild_id(event)
|
||||
|
||||
if not guild_id or not hasattr(adapter, "leave_voice_channel"):
|
||||
return "Not in a voice channel."
|
||||
|
||||
if not hasattr(adapter, "is_in_voice_channel") or not adapter.is_in_voice_channel(guild_id):
|
||||
return "Not in a voice channel."
|
||||
|
||||
try:
|
||||
await adapter.leave_voice_channel(guild_id)
|
||||
except Exception as e:
|
||||
logger.warning("Error leaving voice channel: %s", e)
|
||||
# Always clean up state even if leave raised an exception
|
||||
self._voice_mode[self._voice_key_for_source(event.source)] = "off"
|
||||
self._save_voice_modes()
|
||||
self._set_adapter_auto_tts_disabled(adapter, event.source.chat_id, disabled=True)
|
||||
if hasattr(adapter, "_voice_input_callback"):
|
||||
adapter._voice_input_callback = None
|
||||
return "Left voice channel."
|
||||
|
||||
def _handle_voice_timeout_cleanup(self, chat_id: str, *, adapter=None) -> None:
|
||||
"""Called by the adapter when a voice channel times out.
|
||||
|
||||
Cleans up runner-side voice_mode state that the adapter cannot reach. ``adapter`` is the
|
||||
Discord adapter that timed out (bound at join time); under multiplexing that is a
|
||||
specific profile's bot, not necessarily ``self.adapters[DISCORD]``.
|
||||
"""
|
||||
if adapter is None:
|
||||
adapter = self.adapters.get(Platform.DISCORD)
|
||||
profile = getattr(adapter, "_owner_profile", None)
|
||||
self._voice_mode[self._voice_key(Platform.DISCORD, chat_id, profile=profile)] = "off"
|
||||
self._save_voice_modes()
|
||||
self._set_adapter_auto_tts_disabled(adapter, chat_id, disabled=True)
|
||||
|
||||
def _is_duplicate_voice_transcript(self, guild_id: int, user_id: int, transcript: str) -> bool:
|
||||
"""Suppress repeated STT outputs for the same recent utterance.
|
||||
|
||||
Voice capture can occasionally emit the same utterance twice a few seconds apart, which
|
||||
creates a second queued agent run and overlapping spoken replies.
|
||||
"""
|
||||
from difflib import SequenceMatcher
|
||||
|
||||
normalized = re.sub(r"\s+", " ", transcript).strip().lower()
|
||||
normalized = re.sub(r"[^\w\s]", "", normalized)
|
||||
if not normalized:
|
||||
return False
|
||||
|
||||
now = time.monotonic()
|
||||
window_seconds = 12.0
|
||||
key = (guild_id, user_id)
|
||||
recent_store = getattr(self, "_recent_voice_transcripts", None)
|
||||
if not isinstance(recent_store, dict):
|
||||
recent_store = {}
|
||||
self._recent_voice_transcripts = recent_store
|
||||
recent = [
|
||||
(ts, txt)
|
||||
for ts, txt in recent_store.get(key, [])
|
||||
if now - ts <= window_seconds
|
||||
]
|
||||
|
||||
for _, prior in recent:
|
||||
if prior == normalized:
|
||||
recent_store[key] = recent
|
||||
return True
|
||||
if len(prior) >= 16 and len(normalized) >= 16:
|
||||
if SequenceMatcher(None, prior, normalized).ratio() >= 0.95:
|
||||
recent_store[key] = recent
|
||||
return True
|
||||
|
||||
recent.append((now, normalized))
|
||||
recent_store[key] = recent[-5:]
|
||||
return False
|
||||
|
||||
async def _handle_voice_channel_input(
|
||||
self, guild_id: int, user_id: int, transcript: str, *, adapter=None
|
||||
):
|
||||
"""Handle transcribed voice from a user in a voice channel.
|
||||
|
||||
``adapter`` is the Discord adapter that captured the audio (bound via
|
||||
``_bind_voice_input_callback``); under multiplexing each profile's bot must dispatch
|
||||
through its own adapter, never the default profile's.
|
||||
"""
|
||||
if adapter is None:
|
||||
adapter = self.adapters.get(Platform.DISCORD)
|
||||
if not adapter:
|
||||
return
|
||||
|
||||
text_ch_id = adapter._voice_text_channels.get(guild_id)
|
||||
if not text_ch_id:
|
||||
return
|
||||
|
||||
# Build source — reuse the linked text channel's metadata when available
|
||||
# so voice input shares the same session as the bound text conversation.
|
||||
source_data = getattr(adapter, "_voice_sources", {}).get(guild_id)
|
||||
if source_data:
|
||||
source = SessionSource.from_dict(source_data)
|
||||
source.user_id = str(user_id)
|
||||
source.user_name = str(user_id)
|
||||
else:
|
||||
source = SessionSource(
|
||||
platform=Platform.DISCORD,
|
||||
chat_id=str(text_ch_id),
|
||||
user_id=str(user_id),
|
||||
user_name=str(user_id),
|
||||
chat_type="channel",
|
||||
profile=getattr(adapter, "_owner_profile", None),
|
||||
)
|
||||
|
||||
# Check authorization before processing voice input
|
||||
if not self._is_user_authorized(source):
|
||||
logger.debug("Unauthorized voice input from user %d, ignoring", user_id)
|
||||
return
|
||||
|
||||
if self._is_duplicate_voice_transcript(guild_id, user_id, transcript):
|
||||
logger.info(
|
||||
"Suppressing duplicate voice transcript for guild=%s user=%s: %s",
|
||||
guild_id,
|
||||
user_id,
|
||||
transcript[:100],
|
||||
)
|
||||
return
|
||||
|
||||
# Show transcript in text channel (after auth, with mention sanitization)
|
||||
try:
|
||||
channel = adapter._client.get_channel(text_ch_id)
|
||||
if channel:
|
||||
safe_text = transcript[:2000].replace("@everyone", "@\u200beveryone").replace("@here", "@\u200bhere")
|
||||
await channel.send(f"**[Voice]** <@{user_id}>: {safe_text}")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Build a synthetic MessageEvent for the normal pipeline; SimpleNamespace raw_message lets
|
||||
# _get_guild_id() extract guild_id and _send_voice_reply() play audio in the voice channel.
|
||||
from types import SimpleNamespace
|
||||
# Resolve the bound text channel's channel_prompt so voice input gets
|
||||
# the same per-channel context as typed messages (#50149).
|
||||
channel_prompt: Optional[str] = None
|
||||
resolver = getattr(adapter, "_resolve_channel_prompt", None)
|
||||
if callable(resolver):
|
||||
try:
|
||||
resolved = resolver(str(text_ch_id))
|
||||
channel_prompt = resolved if isinstance(resolved, str) else None
|
||||
except Exception:
|
||||
channel_prompt = None
|
||||
event = MessageEvent(
|
||||
source=source,
|
||||
text=transcript,
|
||||
message_type=MessageType.VOICE,
|
||||
raw_message=SimpleNamespace(guild_id=guild_id, guild=None),
|
||||
channel_prompt=channel_prompt,
|
||||
)
|
||||
|
||||
await adapter.handle_message(event)
|
||||
|
||||
def _should_send_voice_reply(
|
||||
self,
|
||||
event: MessageEvent,
|
||||
response: str,
|
||||
agent_messages: list,
|
||||
already_sent: bool = False,
|
||||
) -> bool:
|
||||
"""Decide whether the runner should send a TTS voice reply.
|
||||
|
||||
False when voice_mode is off for this chat, the response is empty/an error, the agent
|
||||
already called text_to_speech (dedup), or voice input + base adapter auto-TTS already
|
||||
handled it (skip_double) — UNLESS streaming consumed the response (already_sent=True),
|
||||
since then the base adapter has no text for auto-TTS and the runner must handle it.
|
||||
"""
|
||||
if not response or response.startswith("Error:"):
|
||||
return False
|
||||
|
||||
chat_id = event.source.chat_id
|
||||
voice_key = self._voice_key_for_source(event.source)
|
||||
voice_mode = self._voice_mode.get(voice_key)
|
||||
is_voice_input = (event.message_type == MessageType.VOICE)
|
||||
|
||||
adapter = self._adapter_for_source(event.source)
|
||||
adapter_auto_tts = False
|
||||
if adapter and hasattr(adapter, "_should_auto_tts_for_chat"):
|
||||
try:
|
||||
adapter_auto_tts = bool(adapter._should_auto_tts_for_chat(chat_id))
|
||||
except Exception:
|
||||
adapter_auto_tts = False
|
||||
|
||||
should = (
|
||||
(voice_mode == "all")
|
||||
or (voice_mode == "voice_only" and is_voice_input)
|
||||
# ``voice.auto_tts`` (synced into the adapter at startup) is the fallback only when the
|
||||
# chat has no explicit mode; the chat-level all/voice_only/off choice takes precedence.
|
||||
or (voice_mode is None and adapter_auto_tts)
|
||||
)
|
||||
if not should:
|
||||
logger.debug(
|
||||
"Auto voice reply skipped: mode=%s adapter_auto_tts=%s chat=%s platform=%s",
|
||||
voice_mode, adapter_auto_tts, chat_id, event.source.platform.value,
|
||||
)
|
||||
return False
|
||||
|
||||
# Dedup: agent already called TTS tool in THIS turn only
|
||||
last_user_idx = None
|
||||
for i, msg in enumerate(reversed(agent_messages)):
|
||||
if msg.get("role") == "user":
|
||||
last_user_idx = len(agent_messages) - 1 - i; break
|
||||
turn_messages = agent_messages[last_user_idx:] if last_user_idx is not None else agent_messages
|
||||
has_agent_tts = any(
|
||||
msg.get("role") == "assistant"
|
||||
and any(
|
||||
(tc.get("function") or {}).get("name") == "text_to_speech"
|
||||
for tc in (msg.get("tool_calls") or [])
|
||||
)
|
||||
for msg in turn_messages
|
||||
)
|
||||
if has_agent_tts:
|
||||
return False
|
||||
|
||||
# Dedup: base adapter auto-TTS already handles voice input (play_tts plays in VC when
|
||||
# connected), so the runner can skip — unless streaming already delivered the text
|
||||
# (already_sent): then the base adapter gets None, can't run auto-TTS, and the runner must.
|
||||
return not (is_voice_input and not already_sent)
|
||||
|
||||
def _should_echo_stt_transcripts(self) -> bool:
|
||||
"""Return whether inbound voice/STT transcripts should be echoed to chat."""
|
||||
return bool(getattr(self.config, "stt_echo_transcripts", True))
|
||||
|
||||
async def _send_voice_reply(self, event: MessageEvent, text: str) -> None:
|
||||
"""Generate TTS audio and send as a voice message before the text reply."""
|
||||
audio_path = None
|
||||
actual_paths: List[str] = []
|
||||
try:
|
||||
from tools.tts_tool import text_to_speech_tool, _strip_markdown_for_tts
|
||||
|
||||
tts_text = _strip_markdown_for_tts(text)
|
||||
if not tts_text:
|
||||
return
|
||||
|
||||
# Platforms whose native voice bubbles require Ogg/Opus (OPUS_VOICE_PLATFORMS —
|
||||
# Telegram, Matrix, Feishu, WhatsApp, Signal) get an explicit .ogg path; the TTS tool's
|
||||
# central container repair guarantees real Ogg/Opus bytes for every provider.
|
||||
audio_path = build_auto_tts_output_path(event.source.platform)
|
||||
|
||||
result_json = await asyncio.to_thread(
|
||||
text_to_speech_tool, text=tts_text, output_path=audio_path
|
||||
)
|
||||
try:
|
||||
result = json.loads(result_json)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
logger.warning("Auto voice reply TTS returned invalid JSON: %s", result_json[:200] if result_json else result_json)
|
||||
return
|
||||
|
||||
# Delivery may be one combined file or several separately valid files (combination
|
||||
# unavailable or over a platform limit); preserve legacy single-file results.
|
||||
actual_paths = result.get("file_paths") or [
|
||||
result.get("file_path", audio_path)
|
||||
]
|
||||
actual_paths = [
|
||||
str(path) for path in actual_paths
|
||||
if path and os.path.isfile(path)
|
||||
]
|
||||
if not result.get("success") or not actual_paths:
|
||||
logger.warning("Auto voice reply TTS failed: %s", result.get("error"))
|
||||
return
|
||||
|
||||
adapter = self._adapter_for_source(event.source)
|
||||
|
||||
# If connected to a voice channel, play there instead of sending a file
|
||||
guild_id = self._get_guild_id(event)
|
||||
play_in_voice_channel = getattr(adapter, "play_in_voice_channel", None)
|
||||
is_in_voice_channel = getattr(adapter, "is_in_voice_channel", None)
|
||||
send_voice = getattr(adapter, "send_voice", None)
|
||||
in_voice_channel = bool(
|
||||
guild_id
|
||||
and callable(play_in_voice_channel)
|
||||
and callable(is_in_voice_channel)
|
||||
and is_in_voice_channel(guild_id)
|
||||
)
|
||||
reply_anchor = self._reply_anchor_for_event(event)
|
||||
thread_meta = self._thread_metadata_for_source(event.source, reply_anchor)
|
||||
if not in_voice_channel and callable(send_voice):
|
||||
# Mark the auto voice reply as notify-worthy (mirrors the final-text path in
|
||||
# platforms/base.py) so adapters that gate push notifications (Telegram "important"
|
||||
# mode) deliver it as a normal notification, not a silent message. Clone first so
|
||||
# we don't mutate metadata shared with concurrent typing-indicator state.
|
||||
if thread_meta is not None:
|
||||
thread_meta = dict(thread_meta)
|
||||
thread_meta["notify"] = True
|
||||
else:
|
||||
thread_meta = {"notify": True}
|
||||
for actual_path in actual_paths:
|
||||
if in_voice_channel:
|
||||
play_voice = cast(Callable[..., Awaitable[Any]], play_in_voice_channel)
|
||||
await play_voice(guild_id, actual_path)
|
||||
elif callable(send_voice):
|
||||
send_voice_call = cast(Callable[..., Awaitable[Any]], send_voice)
|
||||
send_kwargs: Dict[str, Any] = {
|
||||
"chat_id": event.source.chat_id,
|
||||
"audio_path": actual_path,
|
||||
"reply_to": reply_anchor,
|
||||
"metadata": thread_meta,
|
||||
}
|
||||
await send_voice_call(**send_kwargs)
|
||||
except Exception as e:
|
||||
logger.warning("Auto voice reply failed: %s", e, exc_info=True)
|
||||
finally:
|
||||
for p in ({audio_path, *actual_paths} - {None}):
|
||||
with suppress(OSError):
|
||||
os.unlink(p)
|
||||
455
gateway/run_watchers.py
Normal file
455
gateway/run_watchers.py
Normal file
@@ -0,0 +1,455 @@
|
||||
"""Session expiry / stall / catalog-refresh watcher loops for GatewayRunner.
|
||||
|
||||
Split out of ``gateway/run.py``; bound onto ``GatewayRunner`` via the MRO.
|
||||
``gateway.run`` internals are imported lazily inside method bodies (import cycle),
|
||||
so ``patch("gateway.run.X")`` keeps intercepting them at call time.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
import asyncio
|
||||
import time
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
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")
|
||||
|
||||
|
||||
class GatewaySessionWatchersMixin:
|
||||
"""Session expiry / stall / catalog-refresh watcher loops for GatewayRunner."""
|
||||
|
||||
async def _session_expiry_watcher(self, interval: int = 300):
|
||||
"""Background task that finalizes expired sessions: runs ``on_session_finalize`` hooks,
|
||||
cleans up the cached agent's tool resources, evicts the cache entry, and marks the session
|
||||
finalized so it is not finalized again.
|
||||
"""
|
||||
from gateway.run import _AGENT_PENDING_SENTINEL
|
||||
await asyncio.sleep(60) # initial delay — let the gateway fully start
|
||||
_finalize_failures: dict[str, int] = {} # session_id -> consecutive failure count
|
||||
_MAX_FINALIZE_RETRIES = 3
|
||||
while self._running:
|
||||
try:
|
||||
await self.async_session_store._ensure_loaded()
|
||||
# Collect expired sessions first, then log a single summary.
|
||||
_expired_entries = []
|
||||
for key, entry in list(self.session_store._entries.items()):
|
||||
if entry.expiry_finalized:
|
||||
continue
|
||||
if not await self.async_session_store._is_session_expired(entry):
|
||||
continue
|
||||
_expired_entries.append((key, entry))
|
||||
|
||||
if _expired_entries:
|
||||
# Extract platform names from session keys for a compact summary.
|
||||
# Keys look like "agent:main:telegram:dm:12345" — platform is field [2].
|
||||
_platforms: dict[str, int] = {}
|
||||
for _k, _e in _expired_entries:
|
||||
_parts = _k.split(":")
|
||||
_plat = _parts[2] if len(_parts) > 2 else "unknown"
|
||||
_platforms[_plat] = _platforms.get(_plat, 0) + 1
|
||||
_plat_summary = ", ".join(
|
||||
f"{p}:{c}" for p, c in sorted(_platforms.items())
|
||||
)
|
||||
logger.info(
|
||||
"Session expiry: %d sessions to finalize (%s)",
|
||||
len(_expired_entries), _plat_summary,
|
||||
)
|
||||
|
||||
for key, entry in _expired_entries:
|
||||
try:
|
||||
try:
|
||||
_parts = key.split(":")
|
||||
_platform = _parts[2] if len(_parts) > 2 else ""
|
||||
# Off-loop + bounded: plugin finalize hooks can block arbitrarily, and
|
||||
# this watcher runs on the gateway event loop.
|
||||
await self._finalize_session_off_loop(
|
||||
session_id=entry.session_id,
|
||||
platform=_platform,
|
||||
reason="session_expired",
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
# Close the cached agent's memory provider and tool resources. Idle agents
|
||||
# live in _agent_cache (not _running_agents), so look there.
|
||||
_cached_agent = None
|
||||
_cache_lock = getattr(self, "_agent_cache_lock", None)
|
||||
if _cache_lock is not None:
|
||||
with _cache_lock:
|
||||
_cached = self._agent_cache.get(key)
|
||||
_cached_agent = _cached[0] if isinstance(_cached, tuple) else _cached if _cached else None
|
||||
# Fall back to _running_agents in case the agent is
|
||||
# still mid-turn when the expiry fires.
|
||||
if _cached_agent is None:
|
||||
_exp_state = self._peek_session_state(key)
|
||||
_cached_agent = _exp_state.turn.agent if _exp_state else None
|
||||
if _cached_agent and _cached_agent is not _AGENT_PENDING_SENTINEL:
|
||||
await self._cleanup_agent_resources_off_loop(
|
||||
_cached_agent, context="session expiry"
|
||||
)
|
||||
# Drop the cache entry so the AIAgent (LLM clients, tool schemas, memory
|
||||
# provider refs) can be GC'd; otherwise the cache grows unbounded.
|
||||
self._evict_cached_agent(key)
|
||||
# Permanent finalization: one funnel call drops every conversation-scoped
|
||||
# dict AND boundary security state so they don't grow unbounded. Idle
|
||||
# agent-cache eviction must NOT do this — that session is still alive and a
|
||||
# resumed turn rebuilds from these overrides. Only finalize, /new, /reset clear.
|
||||
self._clear_conversation_scope(
|
||||
key, reason="expiry_finalized"
|
||||
)
|
||||
# Persist finalized flag (sessions.json AND state.db, single write-path);
|
||||
# also drops the /model override — finalization is a conversation boundary.
|
||||
await self.async_session_store.set_expiry_finalized(entry)
|
||||
logger.debug(
|
||||
"Session expiry finalized for %s",
|
||||
entry.session_id,
|
||||
)
|
||||
_finalize_failures.pop(entry.session_id, None)
|
||||
except Exception as e:
|
||||
failures = _finalize_failures.get(entry.session_id, 0) + 1
|
||||
_finalize_failures[entry.session_id] = failures
|
||||
if failures >= _MAX_FINALIZE_RETRIES:
|
||||
logger.warning(
|
||||
"Session finalize gave up after %d attempts for %s: %s. "
|
||||
"Marking as finalized to prevent infinite retry loop.",
|
||||
failures, entry.session_id, e,
|
||||
)
|
||||
await self.async_session_store.set_expiry_finalized(
|
||||
entry, clear_model_override=False
|
||||
)
|
||||
_finalize_failures.pop(entry.session_id, None)
|
||||
else:
|
||||
logger.debug(
|
||||
"Session finalize failed (%d/%d) for %s: %s",
|
||||
failures, _MAX_FINALIZE_RETRIES, entry.session_id, e,
|
||||
)
|
||||
|
||||
if _expired_entries:
|
||||
_done = sum(
|
||||
1 for _, e in _expired_entries if e.expiry_finalized
|
||||
)
|
||||
_failed = len(_expired_entries) - _done
|
||||
if _failed:
|
||||
logger.info(
|
||||
"Session expiry done: %d finalized, %d pending retry",
|
||||
_done, _failed,
|
||||
)
|
||||
else:
|
||||
logger.info(
|
||||
"Session expiry done: %d finalized", _done,
|
||||
)
|
||||
|
||||
# Sweep agents idle beyond the TTL regardless of session reset policy: sessions with
|
||||
# long / "never" reset windows would otherwise pin memory for the gateway's life.
|
||||
try:
|
||||
_idle_evicted = self._sweep_idle_cached_agents()
|
||||
if _idle_evicted:
|
||||
logger.info(
|
||||
"Agent cache idle sweep: evicted %d agent(s)",
|
||||
_idle_evicted,
|
||||
)
|
||||
except Exception as _e:
|
||||
logger.debug("Idle agent sweep failed: %s", _e)
|
||||
|
||||
# Neither LRU cap nor idle TTL knows what a cached transcript costs in memory, so a
|
||||
# busy gateway keeps every warm session's tool output resident until the RSS limit.
|
||||
try:
|
||||
self._sweep_agent_cache_under_pressure()
|
||||
except Exception as _e:
|
||||
logger.debug("Agent cache pressure sweep failed: %s", _e)
|
||||
|
||||
# Prune stale SessionStore entries; the in-memory dict (and sessions.json) would
|
||||
# otherwise grow unbounded with many rotating chats / threads / users.
|
||||
_last_prune_ts = getattr(self, "_last_session_store_prune_ts", 0.0)
|
||||
_prune_interval = 3600.0 # once per hour
|
||||
if time.time() - _last_prune_ts > _prune_interval:
|
||||
try:
|
||||
_max_age = int(
|
||||
getattr(self.config, "session_store_max_age_days", 0) or 0
|
||||
)
|
||||
if _max_age > 0:
|
||||
_pruned = await self.async_session_store.prune_old_entries(_max_age)
|
||||
if _pruned:
|
||||
logger.info(
|
||||
"SessionStore prune: dropped %d stale entries",
|
||||
_pruned,
|
||||
)
|
||||
except Exception as _e:
|
||||
logger.debug("SessionStore prune failed: %s", _e)
|
||||
self._last_session_store_prune_ts = time.time()
|
||||
except Exception as e:
|
||||
logger.debug("Session expiry watcher error: %s", e)
|
||||
# Sleep in small increments so we can stop quickly
|
||||
for _ in range(interval):
|
||||
if not self._running:
|
||||
break
|
||||
await asyncio.sleep(1)
|
||||
|
||||
def _session_stall_timeout_seconds(self) -> float:
|
||||
"""Return configured stall timeout (seconds); 0 disables the watchdog."""
|
||||
from gateway.run import _float_env
|
||||
return _float_env("HERMES_SESSION_STALL_TIMEOUT", 300)
|
||||
|
||||
def _iter_gateway_adapters(self):
|
||||
"""Yield every live platform adapter (default + multiplex profiles)."""
|
||||
seen: set[int] = set()
|
||||
for adapter in list(getattr(self, "adapters", {}).values()):
|
||||
if adapter is None:
|
||||
continue
|
||||
aid = id(adapter)
|
||||
if aid in seen:
|
||||
continue
|
||||
seen.add(aid)
|
||||
yield adapter
|
||||
for amap in list(getattr(self, "_profile_adapters", {}).values()):
|
||||
for adapter in list(amap.values()):
|
||||
if adapter is None:
|
||||
continue
|
||||
aid = id(adapter)
|
||||
if aid in seen:
|
||||
continue
|
||||
seen.add(aid)
|
||||
yield adapter
|
||||
|
||||
def _session_activity_for_stall(self, session_key: str) -> Optional[dict]:
|
||||
"""Return the shared activity snapshot for stall progress: the single source is
|
||||
``AIAgent.get_activity_summary()`` / ``agent.session_activity``; no turn-start or
|
||||
pending-inbound clocks.
|
||||
"""
|
||||
from gateway.run import _AGENT_PENDING_SENTINEL
|
||||
agent = (getattr(self, "_running_agents", None) or {}).get(session_key)
|
||||
if agent is None or agent is _AGENT_PENDING_SENTINEL:
|
||||
return None
|
||||
if not hasattr(agent, "get_activity_summary"):
|
||||
return None
|
||||
try:
|
||||
summary = agent.get_activity_summary()
|
||||
except Exception:
|
||||
return None
|
||||
return summary if isinstance(summary, dict) else None
|
||||
|
||||
async def _check_session_stalls(self, timeout_seconds: float) -> int:
|
||||
"""Scan pending inbound sessions and notify once per stall episode; returns the number of
|
||||
notifications sent this pass (for tests).
|
||||
"""
|
||||
from gateway.run import _STALL_NOTIFY_SEND_TIMEOUT_SECONDS
|
||||
from gateway.session_stall import (
|
||||
format_session_stall_notification,
|
||||
resolve_session_idle_seconds_from_activity,
|
||||
should_clear_session_stall_notification,
|
||||
should_emit_session_stall_notification,
|
||||
)
|
||||
|
||||
notified_map = getattr(self, "_session_stall_notified", None)
|
||||
if notified_map is None:
|
||||
notified_map = {}
|
||||
self._session_stall_notified = notified_map
|
||||
|
||||
sent = 0
|
||||
now = time.time()
|
||||
candidates: Dict[str, tuple[Any, Any]] = {}
|
||||
|
||||
for adapter in self._iter_gateway_adapters():
|
||||
pending_slot = getattr(adapter, "_pending_messages", None) or {}
|
||||
for session_key, event in list(pending_slot.items()):
|
||||
if session_key and session_key not in candidates and event is not None:
|
||||
candidates[session_key] = (adapter, event)
|
||||
|
||||
for session_key, overflow in list(
|
||||
(getattr(self, "_queued_events", None) or {}).items()
|
||||
):
|
||||
if not session_key or session_key in candidates or not overflow:
|
||||
continue
|
||||
event = overflow[0]
|
||||
source = getattr(event, "source", None)
|
||||
adapter = (
|
||||
self._adapter_for_source(source) if source is not None else None
|
||||
)
|
||||
if adapter is None:
|
||||
continue
|
||||
candidates[session_key] = (adapter, event)
|
||||
|
||||
for session_key, (adapter, pending_event) in list(candidates.items()):
|
||||
has_pending = pending_event is not None
|
||||
activity = (
|
||||
self._session_activity_for_stall(session_key) if has_pending else None
|
||||
)
|
||||
idle_seconds = (
|
||||
resolve_session_idle_seconds_from_activity(activity, now=now)
|
||||
if has_pending
|
||||
else None
|
||||
)
|
||||
already = bool(notified_map.get(session_key))
|
||||
if should_clear_session_stall_notification(
|
||||
timeout_seconds=timeout_seconds,
|
||||
idle_seconds=idle_seconds,
|
||||
has_pending_inbound=has_pending,
|
||||
):
|
||||
notified_map.pop(session_key, None)
|
||||
already = False
|
||||
if not should_emit_session_stall_notification(
|
||||
timeout_seconds=timeout_seconds,
|
||||
idle_seconds=idle_seconds,
|
||||
has_pending_inbound=has_pending,
|
||||
already_notified=already,
|
||||
):
|
||||
continue
|
||||
|
||||
if idle_seconds is None:
|
||||
continue
|
||||
mins = max(1, int(idle_seconds // 60))
|
||||
activity = activity or {}
|
||||
logger.warning(
|
||||
"Session stall detected: session=%s idle=%.0fs "
|
||||
"(timeout=%.0fs, ~%d min); pending inbound present "
|
||||
"| last_activity=%s | provenance=%s "
|
||||
"(agent.session_stall_timeout)",
|
||||
session_key,
|
||||
idle_seconds,
|
||||
timeout_seconds,
|
||||
mins,
|
||||
activity.get("last_activity_desc")
|
||||
or activity.get("last_activity_description")
|
||||
or "unknown",
|
||||
activity.get("provenance")
|
||||
or activity.get("last_activity_provenance")
|
||||
or "unknown",
|
||||
)
|
||||
source = getattr(pending_event, "source", None)
|
||||
chat_id = getattr(source, "chat_id", None) if source is not None else None
|
||||
if not chat_id:
|
||||
logger.warning(
|
||||
"Session stall notify skipped (no chat_id): session=%s",
|
||||
session_key,
|
||||
)
|
||||
# Cannot deliver; latch to avoid log spam every tick.
|
||||
notified_map[session_key] = True
|
||||
continue
|
||||
# Re-read pending state + activity IMMEDIATELY before delivery: the snapshot above ages
|
||||
# while earlier candidates await sends; an agent that progressed (or drained its queue)
|
||||
# must not get a false stall notice. Abort, latch un-set, so the next tick re-evaluates.
|
||||
still_pending = (
|
||||
(getattr(adapter, "_pending_messages", None) or {}).get(
|
||||
session_key
|
||||
)
|
||||
is not None
|
||||
or bool(
|
||||
(getattr(self, "_queued_events", None) or {}).get(
|
||||
session_key
|
||||
)
|
||||
)
|
||||
)
|
||||
fresh_idle = resolve_session_idle_seconds_from_activity(
|
||||
self._session_activity_for_stall(session_key),
|
||||
now=time.time(),
|
||||
)
|
||||
if not still_pending or (
|
||||
fresh_idle is not None and fresh_idle < timeout_seconds
|
||||
):
|
||||
logger.info(
|
||||
"Session stall notify aborted (no longer stale): "
|
||||
"session=%s pending=%s fresh_idle=%s",
|
||||
session_key,
|
||||
still_pending,
|
||||
fresh_idle,
|
||||
)
|
||||
# Re-arm: drop any stale latch so a FUTURE genuine stall
|
||||
# episode notifies again.
|
||||
notified_map.pop(session_key, None)
|
||||
continue
|
||||
try:
|
||||
metadata = (
|
||||
self._thread_metadata_for_source(source)
|
||||
if source is not None and hasattr(self, "_thread_metadata_for_source")
|
||||
else None
|
||||
)
|
||||
# Bound the send: a wedged adapter transport (network hang, dead websocket) must not
|
||||
# block the watcher pass — siblings would go unevaluated and the watcher stop.
|
||||
try:
|
||||
result = await asyncio.wait_for(
|
||||
adapter.send(
|
||||
str(chat_id),
|
||||
format_session_stall_notification(idle_seconds),
|
||||
metadata=metadata,
|
||||
),
|
||||
timeout=_STALL_NOTIFY_SEND_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.warning(
|
||||
"Session stall notify send timed out after %.0fs "
|
||||
"for %s; will retry next tick",
|
||||
_STALL_NOTIFY_SEND_TIMEOUT_SECONDS,
|
||||
session_key,
|
||||
)
|
||||
continue # do not latch; retry next tick
|
||||
# Adapters often return SendResult(success=False) instead of raising.
|
||||
if result is not None and getattr(result, "success", True) is False:
|
||||
logger.warning(
|
||||
"Session stall notify failed for %s: %s",
|
||||
session_key,
|
||||
getattr(result, "error", "send returned success=False"),
|
||||
)
|
||||
continue # do not latch; retry next tick
|
||||
sent += 1
|
||||
notified_map[session_key] = True
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Session stall notify failed for %s: %s",
|
||||
session_key,
|
||||
exc,
|
||||
)
|
||||
# Do not latch — retry next watcher tick until delivery or episode clear.
|
||||
|
||||
# Drop latches for sessions that no longer appear in any pending map.
|
||||
for key in list(notified_map.keys()):
|
||||
if key not in candidates:
|
||||
notified_map.pop(key, None)
|
||||
|
||||
return sent
|
||||
|
||||
async def _model_catalog_refresh_watcher(self) -> None:
|
||||
"""Refresh the /model picker's remote catalogs every TTL window. The picker itself only
|
||||
refreshes on a cold/stale open, so if nobody opens ``/model`` the cache never updates.
|
||||
"""
|
||||
from hermes_cli.model_catalog import refresh_catalogs, refresh_interval_seconds
|
||||
|
||||
await asyncio.sleep(30) # let startup settle
|
||||
while self._running:
|
||||
try:
|
||||
await asyncio.to_thread(refresh_catalogs)
|
||||
except Exception as exc:
|
||||
logger.debug("Model catalog refresh failed: %s", exc)
|
||||
try:
|
||||
interval = refresh_interval_seconds()
|
||||
except Exception:
|
||||
interval = 1200.0
|
||||
deadline = time.monotonic() + interval
|
||||
while self._running and time.monotonic() < deadline:
|
||||
await asyncio.sleep(min(30.0, max(0.0, deadline - time.monotonic())))
|
||||
|
||||
async def _session_stall_watcher(self, interval: float = 30.0):
|
||||
"""Periodic pending-inbound + stale-activity stall watchdog.
|
||||
|
||||
Progress comes only from ``get_activity_summary()``. Pending inbound is a notify policy
|
||||
gate, not a progress clock. Notify-only: does not kill the turn (contrast
|
||||
``gateway_timeout`` / ``shutdown_watchdog``).
|
||||
"""
|
||||
# Short initial delay so startup reconnect noise does not false-fire.
|
||||
await asyncio.sleep(min(30.0, max(1.0, float(interval))))
|
||||
while self._running:
|
||||
try:
|
||||
timeout = self._session_stall_timeout_seconds()
|
||||
if timeout > 0:
|
||||
await self._check_session_stalls(timeout)
|
||||
except Exception as exc:
|
||||
logger.debug("Session stall watcher error: %s", exc)
|
||||
# Interruptible sleep
|
||||
steps = max(1, int(float(interval)))
|
||||
for _ in range(steps):
|
||||
if not self._running:
|
||||
break
|
||||
await asyncio.sleep(1)
|
||||
@@ -21,6 +21,8 @@ import ast
|
||||
import inspect
|
||||
|
||||
from gateway import run as gateway_run
|
||||
from gateway import run_turn as gateway_run_turn
|
||||
from gateway import run_turn as gateway_run_turn
|
||||
|
||||
|
||||
def _calls(node: ast.AST) -> set[str]:
|
||||
@@ -53,7 +55,7 @@ def test_auto_reset_cleanup_evicts_cached_agent():
|
||||
conversation's cached agent (and its leaked
|
||||
``context_compressor._previous_summary``) — the cache is keyed on the
|
||||
stable ``session_key`` (#10710)."""
|
||||
tree = ast.parse(inspect.getsource(gateway_run))
|
||||
tree = ast.parse(inspect.getsource(gateway_run_turn))
|
||||
|
||||
# Fingerprint the cleanup branch: the `if <was_auto_reset>:` block that
|
||||
# clears the conversation scope via the funnel (post-#64934 refactor:
|
||||
|
||||
@@ -37,6 +37,8 @@ import ast
|
||||
import inspect
|
||||
|
||||
from gateway import run as gateway_run
|
||||
from gateway import run_turn as gateway_run_turn
|
||||
from gateway import run_turn as gateway_run_turn
|
||||
from gateway.config import GatewayConfig, Platform
|
||||
from gateway.session import SessionSource, SessionStore
|
||||
from hermes_state import SessionDB
|
||||
@@ -47,7 +49,7 @@ from hermes_state import SessionDB
|
||||
# ---------------------------------------------------------------------------
|
||||
def _find_compression_exhausted_reset_block() -> ast.If:
|
||||
"""Return the ``if agent_result.get('compression_exhausted') ...`` block."""
|
||||
tree = ast.parse(inspect.getsource(gateway_run))
|
||||
tree = ast.parse(inspect.getsource(gateway_run_turn))
|
||||
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.If):
|
||||
|
||||
@@ -24,6 +24,8 @@ import ast
|
||||
import inspect
|
||||
|
||||
from gateway import run as gateway_run
|
||||
from gateway import run_turn as gateway_run_turn
|
||||
from gateway import run_turn as gateway_run_turn
|
||||
from gateway import slash_commands as gateway_slash
|
||||
|
||||
|
||||
@@ -47,7 +49,7 @@ def test_run_consumes_was_auto_reset_in_cleanup_block():
|
||||
`session_entry.was_auto_reset = False` so the cleanup (which pops the
|
||||
session model/reasoning overrides) cannot re-fire on the next message and
|
||||
wipe an override stored between turns (#48031)."""
|
||||
tree = ast.parse(inspect.getsource(gateway_run))
|
||||
tree = ast.parse(inspect.getsource(gateway_run_turn))
|
||||
|
||||
# Find the cleanup branch: an `if <flag>:` block that clears the
|
||||
# conversation scope (post-funnel: one _clear_conversation_scope call
|
||||
|
||||
@@ -112,7 +112,7 @@ class TestApprovalCommandWiring:
|
||||
)
|
||||
|
||||
def test_chat_platform_path_redacts_before_send(self):
|
||||
import gateway.run as run
|
||||
import gateway.run_turn_runner as run
|
||||
|
||||
self._assert_redacts_then_uses(run, "_approval_notify_sync", "send_exec_approval")
|
||||
|
||||
|
||||
@@ -22,6 +22,8 @@ import ast
|
||||
import inspect
|
||||
|
||||
from gateway import run as gateway_run
|
||||
from gateway import run_turn as gateway_run_turn
|
||||
from gateway import run_turn as gateway_run_turn
|
||||
|
||||
|
||||
def _calls(node: ast.AST) -> set[str]:
|
||||
@@ -35,7 +37,7 @@ def _calls(node: ast.AST) -> set[str]:
|
||||
def _find_deferred_guarded_reset_chain() -> ast.If:
|
||||
"""Return the ``if agent_result.get('compression_deferred') ... elif
|
||||
agent_result.get('compression_exhausted') ... reset_session`` chain."""
|
||||
tree = ast.parse(inspect.getsource(gateway_run))
|
||||
tree = ast.parse(inspect.getsource(gateway_run_turn))
|
||||
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.If):
|
||||
|
||||
@@ -25,6 +25,10 @@ import textwrap
|
||||
from unittest.mock import MagicMock, call
|
||||
|
||||
from gateway import run as gateway_run
|
||||
from gateway import run_turn as gateway_run_turn
|
||||
from gateway import run_turn_runner as gateway_run_turn_runner
|
||||
from gateway import run_turn as gateway_run_turn
|
||||
from gateway import run_turn_runner as gateway_run_turn_runner
|
||||
from gateway.session_context import set_current_session_id, get_session_env
|
||||
|
||||
|
||||
@@ -108,8 +112,9 @@ def test_every_post_compression_session_id_assignment_persists():
|
||||
would compress correctly, the gateway would update its in-memory
|
||||
session_id, then drop it on next gateway restart.
|
||||
"""
|
||||
source = inspect.getsource(gateway_run)
|
||||
assignments = _session_id_assignments_followed_by_save(source)
|
||||
assignments = []
|
||||
for mod in (gateway_run, gateway_run_turn, gateway_run_turn_runner):
|
||||
assignments += _session_id_assignments_followed_by_save(inspect.getsource(mod))
|
||||
assert assignments, (
|
||||
"No ``session_entry.session_id = ...`` assignments found in gateway/run.py — "
|
||||
"either the structure changed or the AST walker is broken."
|
||||
|
||||
@@ -77,9 +77,8 @@ def test_background_and_main_agent_paths_call_refresh():
|
||||
"""
|
||||
from pathlib import Path
|
||||
|
||||
source = (
|
||||
Path(__file__).resolve().parent.parent.parent / "gateway" / "run.py"
|
||||
).read_text(encoding="utf-8")
|
||||
_gw = Path(__file__).resolve().parent.parent.parent / "gateway"
|
||||
source = "\n".join(p.read_text(encoding="utf-8") for p in sorted(_gw.glob("run*.py")))
|
||||
# The agent-construction site inside TurnRunner.run_sync (extracted from
|
||||
# the old _run_agent_inner closure) references the runner as
|
||||
# ``self._runner``; the background-agent site still uses bare ``self``.
|
||||
|
||||
@@ -21,6 +21,8 @@ def mock_runner():
|
||||
# Bind the actual methods to the mock
|
||||
runner._profile_name_for_source = GatewayRunner._profile_name_for_source.__get__(runner)
|
||||
runner._resolve_profile_home_for_source = GatewayRunner._resolve_profile_home_for_source.__get__(runner)
|
||||
# _handle_message's ingress gates (profile route rejection) live in this helper.
|
||||
runner._hm_admit_event = GatewayRunner._hm_admit_event.__get__(runner)
|
||||
return runner
|
||||
|
||||
|
||||
|
||||
@@ -170,7 +170,7 @@ def test_gateway_run_agent_threads_the_event_message_id_into_the_turn():
|
||||
import ast
|
||||
import inspect
|
||||
|
||||
import gateway.run as gateway_run
|
||||
import gateway.run_turn_runner as gateway_run
|
||||
|
||||
source = inspect.getsource(gateway_run)
|
||||
tree = ast.parse(source)
|
||||
|
||||
@@ -75,7 +75,7 @@ class TestGateWiring:
|
||||
completed checks — a source-level pin so the contract test above
|
||||
cannot drift green while the call site regresses."""
|
||||
import inspect
|
||||
import gateway.run as run_mod
|
||||
import gateway.run_turn_runner as run_mod
|
||||
|
||||
src = inspect.getsource(run_mod)
|
||||
anchor = src.index("_final_for_stream = None")
|
||||
|
||||
@@ -77,7 +77,7 @@ class TestLegacyKeyMigration:
|
||||
voice_path.write_text(json.dumps(legacy_data))
|
||||
|
||||
with patch.object(runner, "_VOICE_MODE_PATH", voice_path):
|
||||
with patch("gateway.run.logger") as mock_logger:
|
||||
with patch("gateway.run_voice.logger") as mock_logger:
|
||||
result = runner._load_voice_modes()
|
||||
|
||||
# Legacy keys without ':' should be skipped
|
||||
|
||||
Reference in New Issue
Block a user