hermes-agent/tests/agent/test_context_engine_select_context.py
xue xinglong dec464c351 feat(context-engine): add select_context() per-turn selection hook
Adds an optional, no-op-default select_context() hook to the ContextEngine
ABC, called every turn after the request messages are assembled and before
provider dispatch — independent of should_compress(). Lets an engine select
or replace which context enters the prompt for a single request (retrieval,
topic routing, role/branch switching) without mutating persisted history,
removing the need to abuse should_compress()=True as a per-turn callback.

The host call site (_apply_context_engine_selection) is fail-open: a missing
hook, an exception, or an invalid return value leaves the assembled request
untouched. Additive and non-breaking: the built-in compressor and every
existing engine are unaffected.

Consolidates the per-turn request-assembly surface proposed across #41918,

Related: #36765 #41918 #24949 #47109 #50053 #23837 #25115 #29370
2026-07-23 19:44:35 -07:00

189 lines
6 KiB
Python

"""Tests for the per-turn ``ContextEngine.select_context()`` hook.
``select_context()`` is the *selection / routing* verb — distinct from
compression — that lets an external context engine replace which context
enters the prompt for a single request, every turn, independent of
``should_compress()``. It is additive and no-op by default, and the host
call site (``_apply_context_engine_selection``) is fail-open: a missing hook,
an exception, or an invalid return value must leave the assembled request
untouched and must never mutate persisted history.
This pins the contract that engines such as retrieval-augmented, topic-routed,
and role-switching engines rely on (RFC #36765), consolidating the per-turn
request-assembly surface proposed across #41918, #24949, #47109, and #50053.
"""
from __future__ import annotations
from typing import Any, Dict, List
from unittest.mock import MagicMock
from agent.context_engine import ContextEngine
from agent.conversation_loop import _apply_context_engine_selection
class _MinimalEngine(ContextEngine):
"""Concrete engine implementing only the abstract methods."""
@property
def name(self) -> str:
return "minimal"
def update_from_response(self, usage: Dict[str, Any]) -> None:
pass
def should_compress(self, prompt_tokens: int = None) -> bool:
return False
def compress(
self,
messages: List[Dict[str, Any]],
current_tokens: int = None,
focus_topic: str = None,
) -> List[Dict[str, Any]]:
return messages
def _agent_with(engine) -> Any:
agent = MagicMock()
agent.session_id = "test-session"
agent.context_compressor = engine
return agent
REQUEST = [
{"role": "system", "content": "sys"},
{"role": "user", "content": "hello"},
]
HISTORY = [{"role": "user", "content": "hello"}]
# -- ABC default -----------------------------------------------------------
def test_default_select_context_is_noop():
"""The base implementation returns None (no replacement)."""
engine = _MinimalEngine()
assert (
engine.select_context(
REQUEST,
conversation_messages=HISTORY,
incoming_message=HISTORY[-1],
budget_tokens=0,
)
is None
)
# -- Host call site: _apply_context_engine_selection -----------------------
def test_none_return_leaves_request_unchanged():
"""An engine returning None falls through to the assembled request."""
engine = _MinimalEngine() # default select_context -> None
agent = _agent_with(engine)
out = _apply_context_engine_selection(
agent, REQUEST, HISTORY, HISTORY[-1], logger=MagicMock()
)
assert out is REQUEST
def test_missing_hook_leaves_request_unchanged():
"""An engine without select_context (older/stub base) is a no-op."""
engine = object() # no select_context attribute
agent = _agent_with(engine)
out = _apply_context_engine_selection(
agent, REQUEST, HISTORY, HISTORY[-1], logger=MagicMock()
)
assert out is REQUEST
def test_no_engine_leaves_request_unchanged():
agent = MagicMock()
agent.session_id = "test-session"
agent.context_compressor = None
out = _apply_context_engine_selection(
agent, REQUEST, HISTORY, HISTORY[-1], logger=MagicMock()
)
assert out is REQUEST
def test_valid_list_replaces_request():
"""A valid list of dicts replaces the request messages for this call."""
replacement = [
{"role": "system", "content": "sys"},
{"role": "user", "content": "routed-context"},
]
class _Engine(_MinimalEngine):
def select_context(self, request_messages, **kwargs):
return replacement
agent = _agent_with(_Engine())
out = _apply_context_engine_selection(
agent, REQUEST, HISTORY, HISTORY[-1], logger=MagicMock()
)
assert out is replacement
def test_exception_fails_open():
"""A raising hook is swallowed; the unmodified request is used."""
class _Engine(_MinimalEngine):
def select_context(self, request_messages, **kwargs):
raise RuntimeError("backend offline")
logger = MagicMock()
agent = _agent_with(_Engine())
out = _apply_context_engine_selection(
agent, REQUEST, HISTORY, HISTORY[-1], logger=logger
)
assert out is REQUEST
assert logger.warning.called
def test_non_list_return_is_ignored():
"""A non-list return value is rejected and logged, request unchanged."""
class _Engine(_MinimalEngine):
def select_context(self, request_messages, **kwargs):
return {"role": "user", "content": "oops not a list"}
logger = MagicMock()
agent = _agent_with(_Engine())
out = _apply_context_engine_selection(
agent, REQUEST, HISTORY, HISTORY[-1], logger=logger
)
assert out is REQUEST
assert logger.warning.called
def test_list_of_non_dicts_is_ignored():
"""A list that isn't all dicts is rejected, request unchanged."""
class _Engine(_MinimalEngine):
def select_context(self, request_messages, **kwargs):
return ["not", "dicts"]
agent = _agent_with(_Engine())
out = _apply_context_engine_selection(
agent, REQUEST, HISTORY, HISTORY[-1], logger=MagicMock()
)
assert out is REQUEST
def test_persisted_history_not_mutated():
"""The hook must not mutate the persisted conversation history."""
class _Engine(_MinimalEngine):
def select_context(self, request_messages, *, conversation_messages=None, **kwargs):
# Even a misbehaving engine touching its inputs must not affect
# what the host persists — the host passes the live list, so we
# assert the host contract by checking the engine received it and
# the canonical copy is unchanged after the call.
return list(request_messages)
history_snapshot = [dict(m) for m in HISTORY]
agent = _agent_with(_Engine())
_apply_context_engine_selection(
agent, REQUEST, HISTORY, HISTORY[-1], logger=MagicMock()
)
assert HISTORY == history_snapshot