Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 19 additions & 6 deletions gateway/run_inbound.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,7 @@
from gateway.platforms.event import MessageEvent, MessageType
from gateway.run_common import _UNSET
from gateway.session import (
SessionSource, is_shared_multi_user_session, neutralize_untrusted_inline_text
SessionSource, build_session_context, is_shared_multi_user_session, neutralize_untrusted_inline_text
)
from gateway.turn_lease import TurnLeaseTimeoutError
from typing import Any, Dict, List, Optional, Tuple
Expand Down Expand Up @@ -955,7 +955,8 @@ async def _hm_run_exec_quick_command(self, command: str, exec_cmd: str) -> str:
return f"Quick command error: {e}"

async def _hm_dispatch_quick_and_plugin_commands(
self, event: "MessageEvent", source: SessionSource, command: Optional[str]
self, event: "MessageEvent", source: SessionSource, command: Optional[str],
session_key: Optional[str] = None,
) -> Tuple[bool, Optional[str], Optional[str]]:
"""Drain gate, user-defined quick commands (exec/alias) and plugin slash commands →
``(handled, result, command)``; an alias quick command rewrites ``command``."""
Expand Down Expand Up @@ -993,9 +994,19 @@ async def _hm_dispatch_quick_and_plugin_commands(
from hermes_cli.plugins import get_plugin_command_handler
plugin_handler = get_plugin_command_handler(command.replace("_", "-"))
if plugin_handler:
result = plugin_handler(event.get_command_args().strip())
if asyncio.iscoroutine(result):
result = await result
_session_env_tokens = None
if source is not None and session_key:
context = build_session_context(source, self.config)
context.session_key = session_key
_session_env_tokens = self._set_session_env(context)
try:
result = plugin_handler(event.get_command_args().strip())
if asyncio.iscoroutine(result):
result = await result
finally:
if _session_env_tokens is not None:
from gateway.session_context import restore_session_vars
restore_session_vars(_session_env_tokens)
return True, str(result) if result else None, command
except Exception as e:
logger.warning("Plugin command dispatch failed: %s", e)
Expand Down Expand Up @@ -1146,7 +1157,9 @@ async def _hm_dispatch_idle_commands(
if not _handled:
_handled, _result = await self._hm_dispatch_canonical_command(event, source, _quick_key, canonical)
if not _handled:
_handled, _result, command = await self._hm_dispatch_quick_and_plugin_commands(event, source, command)
_handled, _result, command = await self._hm_dispatch_quick_and_plugin_commands(
event, source, command, _quick_key,
)
if not _handled:
_result = self._hm_skill_slash_rewrite(event, source, _quick_key, command)
_handled = _result is not None
Expand Down
20 changes: 16 additions & 4 deletions gateway/session_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,13 +69,13 @@ def session_context_engaged() -> bool:
)}


def _runtime_cwd(func: str, *args: Any) -> None:
def _runtime_cwd(func: str, *args: Any) -> Any:
"""Best-effort call of ``agent.runtime_cwd.<func>``; import/runtime failures are ignored."""
try:
from agent import runtime_cwd
getattr(runtime_cwd, func)(*args)
return getattr(runtime_cwd, func)(*args)
except Exception:
pass
return None


def set_current_session_id(session_id: str) -> None:
Expand Down Expand Up @@ -140,10 +140,22 @@ def set_session_vars(
tokens = [var.set(value) for var, value in zip(_SESSION_VARS, values)]
tokens.append(_SESSION_ASYNC_DELIVERY.set(bool(async_delivery)))
tokens.append(_SESSION_HISTORY_DELIVERY.set(_UNSET if session_history_delivery is None else session_history_delivery))
_runtime_cwd("set_session_cwd", cwd)
if (cwd_token := _runtime_cwd("set_session_cwd", cwd)) is not None:
tokens.append(cwd_token)
return tokens


def restore_session_vars(tokens: list) -> None:
"""Restore a nested session binding from the tokens returned by ``set_session_vars``.

Unlike ``clear_session_vars``, this is only for a local scope that must preserve an outer
binding. Reset in reverse order so every ContextVar, including runtime cwd, regains its
exact prior value.
"""
for token in reversed(tokens):
token.var.reset(token)


def clear_session_vars(tokens: list) -> None:
"""Mark session context variables as explicitly cleared (``""``, not ``_UNSET``), so
``get_session_env`` returns empty instead of stale ``os.environ`` values. Async-delivery
Expand Down
167 changes: 167 additions & 0 deletions tests/gateway/test_plugin_command_session_context.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,167 @@
"""Regression tests for session context during gateway plugin slash commands."""

import asyncio
from pathlib import Path

import pytest

from gateway.session import build_session_key
from gateway.session_context import (
async_delivery_supported,
clear_session_vars,
get_session_env,
session_history_delivery_supported,
set_session_vars,
)
from tests.gateway.test_unknown_command import _make_event, _make_runner


@pytest.mark.asyncio
@pytest.mark.parametrize("async_handler", [False, True])
async def test_plugin_command_handler_has_inbound_session_context_and_restores_it(
monkeypatch, tmp_path, async_handler,
):
from gateway.run import GatewayRunner
from hermes_cli import plugins as plugins_mod
from agent.runtime_cwd import resolve_context_cwd

runner = _make_runner()
runner._draining = False
runner._set_session_env = GatewayRunner._set_session_env.__get__(runner)
event = _make_event("/observe value")
expected_key = build_session_key(event.source)
observed = []

outer_tokens = set_session_vars(
platform="outer-platform",
source="outer-source",
chat_id="outer-chat",
chat_type="outer-type",
chat_name="outer-name",
thread_id="outer-thread",
user_id="outer-user",
user_id_alt="outer-user-alt",
user_name="outer-user-name",
scope_id="outer-scope",
session_key="outer-key",
session_id="outer-session",
ui_session_id="outer-ui-session",
message_id="outer-message",
profile="outer-profile",
browser_control_principal="outer-principal",
browser_control_transport_family="outer-transport",
cron_session="1",
parent_chat_id="outer-parent-chat",
async_delivery=False,
session_history_delivery="1",
cwd=str(tmp_path),
)

session_names = (
"HERMES_SESSION_PLATFORM",
"HERMES_SESSION_SOURCE",
"HERMES_SESSION_CHAT_ID",
"HERMES_SESSION_CHAT_TYPE",
"HERMES_SESSION_CHAT_NAME",
"HERMES_SESSION_THREAD_ID",
"HERMES_SESSION_USER_ID",
"HERMES_SESSION_USER_ID_ALT",
"HERMES_SESSION_USER_NAME",
"HERMES_SESSION_SCOPE_ID",
"HERMES_SESSION_KEY",
"HERMES_SESSION_ID",
"HERMES_UI_SESSION_ID",
"HERMES_SESSION_MESSAGE_ID",
"HERMES_SESSION_PROFILE",
"HERMES_BROWSER_CONTROL_PRINCIPAL",
"HERMES_BROWSER_CONTROL_TRANSPORT_FAMILY",
"HERMES_CRON_SESSION",
"HERMES_SESSION_PARENT_CHAT_ID",
)
outer_context = {name: get_session_env(name) for name in session_names}

def observe():
observed.append({
"key": get_session_env("HERMES_SESSION_KEY"),
"platform": get_session_env("HERMES_SESSION_PLATFORM"),
"chat_id": get_session_env("HERMES_SESSION_CHAT_ID"),
"user_id": get_session_env("HERMES_SESSION_USER_ID"),
})
return "observed"

async def async_observe(_args):
await asyncio.sleep(0)
return observe()

def sync_observe(_args):
return observe()

monkeypatch.setattr(plugins_mod, "get_plugin_commands", lambda: {"observe": {}})
monkeypatch.setattr(
plugins_mod,
"get_plugin_command_handler",
lambda name: (async_observe if async_handler else sync_observe) if name == "observe" else None,
)

try:
handled, result, command = await runner._hm_dispatch_quick_and_plugin_commands(
event, event.source, "observe", expected_key,
)

assert (handled, result, command) == (True, "observed", "observe")
assert observed == [{
"key": expected_key,
"platform": "telegram",
"chat_id": "c1",
"user_id": "u1",
}]
assert {name: get_session_env(name) for name in session_names} == outer_context
assert async_delivery_supported() is False
assert session_history_delivery_supported() is True
assert resolve_context_cwd() == Path(tmp_path)
finally:
clear_session_vars(outer_tokens)


@pytest.mark.asyncio
async def test_plugin_command_exception_restores_outer_session_context(monkeypatch, tmp_path):
from gateway.run import GatewayRunner
from hermes_cli import plugins as plugins_mod
from agent.runtime_cwd import resolve_context_cwd

runner = _make_runner()
runner._draining = False
runner._set_session_env = GatewayRunner._set_session_env.__get__(runner)
event = _make_event("/explode")
observed = []
outer_tokens = set_session_vars(
platform="outer-platform",
session_key="outer-key",
async_delivery=False,
session_history_delivery="1",
cwd=str(tmp_path),
)

def explode(_args):
observed.append(get_session_env("HERMES_SESSION_KEY"))
raise RuntimeError("plugin boom")

monkeypatch.setattr(plugins_mod, "get_plugin_commands", lambda: {"explode": {}})
monkeypatch.setattr(
plugins_mod, "get_plugin_command_handler", lambda name: explode if name == "explode" else None,
)

try:
result = await runner._hm_dispatch_quick_and_plugin_commands(
event, event.source, "explode", build_session_key(event.source),
)

assert observed == [build_session_key(event.source)]
assert result == (False, None, "explode")
assert get_session_env("HERMES_SESSION_KEY") == "outer-key"
assert get_session_env("HERMES_SESSION_PLATFORM") == "outer-platform"
assert async_delivery_supported() is False
assert session_history_delivery_supported() is True
assert resolve_context_cwd() == Path(tmp_path)
finally:
clear_session_vars(outer_tokens)