A general-plugin context engine is one shared instance; agent init copied it per agent with copy.deepcopy() only. Engines that hold a SQLite connection or lock (hermes-lcm) already expose clone_for_agent() for exactly this, but it was never called, so every init logged "could not be safely copied … falling back to built-in compressor" and the engine was unusable through the plugin system. ContextEngine grows clone_for_agent() (default: deepcopy, the previous behaviour) and _select_context_engine calls it; the failure message now names the hook to override. Docs: context-engine-plugin.md documents the per-agent clone contract. Test change (existing on main): tests/agent/test_context_engine.py:: test_agent_init_source_deepcopies_singleton_not_aliases was a source-reading pin on the literal `copy.deepcopy(_candidate)` line, which this fix intentionally replaces. It is superseded by tests/agent/test_plugin_context_engine_clone.py, which drives the real _select_context_engine seam and asserts the invariant it guarded (child update_model() never mutates the shared singleton) plus the new clone_for_agent() path. Fixes #99640 credit: @stephenschoettler #62374 credit: @686f6c61 #99677
316 lines
11 KiB
Python
316 lines
11 KiB
Python
"""Tests for the ContextEngine ABC and plugin slot."""
|
|
|
|
import json
|
|
import pytest
|
|
from typing import Any, Dict, List
|
|
|
|
from agent.context_engine import ContextEngine
|
|
from agent.context_compressor import ContextCompressor
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# A minimal concrete engine for testing the ABC
|
|
# ---------------------------------------------------------------------------
|
|
|
|
class StubEngine(ContextEngine):
|
|
"""Minimal engine that satisfies the ABC without doing real work."""
|
|
|
|
def __init__(self, context_length=200000, threshold_pct=0.50):
|
|
self.context_length = context_length
|
|
self.threshold_tokens = int(context_length * threshold_pct)
|
|
self._compress_called = False
|
|
self._tools_called = []
|
|
|
|
@property
|
|
def name(self) -> str:
|
|
return "stub"
|
|
|
|
def update_model(self, model="", context_length=0, base_url="", api_key="",
|
|
provider="", api_mode="", **kwargs) -> None:
|
|
"""Mirror ContextCompressor.update_model — recompute threshold from the
|
|
new context_length. This is the mutation that corrupted the shared
|
|
singleton in #42449."""
|
|
self.context_length = context_length
|
|
self.threshold_tokens = int(context_length * 0.20)
|
|
|
|
def update_from_response(self, usage: Dict[str, Any]) -> None:
|
|
self.last_prompt_tokens = usage.get("prompt_tokens", 0)
|
|
self.last_completion_tokens = usage.get("completion_tokens", 0)
|
|
self.last_total_tokens = usage.get("total_tokens", 0)
|
|
|
|
def should_compress(self, prompt_tokens: int = None) -> bool:
|
|
tokens = prompt_tokens if prompt_tokens is not None else self.last_prompt_tokens
|
|
return tokens >= self.threshold_tokens
|
|
|
|
def compress(self, messages: List[Dict[str, Any]], current_tokens: int = None) -> List[Dict[str, Any]]:
|
|
self._compress_called = True
|
|
self.compression_count += 1
|
|
# Trivial: just return as-is
|
|
return messages
|
|
|
|
def get_tool_schemas(self) -> List[Dict[str, Any]]:
|
|
return [
|
|
{
|
|
"name": "stub_search",
|
|
"description": "Search the stub engine",
|
|
"parameters": {"type": "object", "properties": {}},
|
|
}
|
|
]
|
|
|
|
def handle_tool_call(self, name: str, args: Dict[str, Any]) -> str:
|
|
self._tools_called.append(name)
|
|
return json.dumps({"ok": True, "tool": name})
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# ABC contract tests
|
|
# ---------------------------------------------------------------------------
|
|
|
|
class TestContextEngineABC:
|
|
"""Verify the ABC enforces the required interface."""
|
|
|
|
|
|
def test_missing_methods_raises(self):
|
|
"""A subclass missing required methods cannot be instantiated."""
|
|
class Incomplete(ContextEngine):
|
|
@property
|
|
def name(self):
|
|
return "incomplete"
|
|
with pytest.raises(TypeError):
|
|
Incomplete()
|
|
|
|
def test_stub_engine_satisfies_abc(self):
|
|
engine = StubEngine()
|
|
assert isinstance(engine, ContextEngine)
|
|
assert engine.name == "stub"
|
|
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Default method behavior
|
|
# ---------------------------------------------------------------------------
|
|
|
|
class TestDefaults:
|
|
"""Verify ABC default implementations work correctly."""
|
|
|
|
|
|
|
|
def test_default_get_status(self):
|
|
engine = StubEngine()
|
|
engine.last_prompt_tokens = 50000
|
|
status = engine.get_status()
|
|
assert status["last_prompt_tokens"] == 50000
|
|
assert status["context_length"] == 200000
|
|
assert status["threshold_tokens"] == 100000
|
|
assert 0 < status["usage_percent"] <= 100
|
|
|
|
|
|
def test_on_session_reset(self):
|
|
engine = StubEngine()
|
|
engine.last_prompt_tokens = 999
|
|
engine.compression_count = 3
|
|
engine.on_session_reset()
|
|
assert engine.last_prompt_tokens == 0
|
|
assert engine.compression_count == 0
|
|
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# StubEngine behavior
|
|
# ---------------------------------------------------------------------------
|
|
|
|
class TestStubEngine:
|
|
|
|
|
|
|
|
def test_tool_schemas(self):
|
|
engine = StubEngine()
|
|
schemas = engine.get_tool_schemas()
|
|
assert len(schemas) == 1
|
|
assert schemas[0]["name"] == "stub_search"
|
|
|
|
def test_handle_tool_call(self):
|
|
engine = StubEngine()
|
|
result = engine.handle_tool_call("stub_search", {})
|
|
assert json.loads(result)["ok"] is True
|
|
assert "stub_search" in engine._tools_called
|
|
|
|
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# ContextCompressor session reset via ABC
|
|
# ---------------------------------------------------------------------------
|
|
|
|
class TestCompressorSessionReset:
|
|
"""Verify ContextCompressor.on_session_reset() clears all state."""
|
|
|
|
def test_reset_clears_state(self):
|
|
c = ContextCompressor(model="test", quiet_mode=True, config_context_length=200000)
|
|
c.last_prompt_tokens = 50000
|
|
c.compression_count = 3
|
|
c._previous_summary = "some old summary"
|
|
c._context_probed = True
|
|
c._context_probe_persistable = True
|
|
|
|
c.on_session_reset()
|
|
|
|
assert c.last_prompt_tokens == 0
|
|
assert c.last_completion_tokens == 0
|
|
assert c.last_total_tokens == 0
|
|
assert c.compression_count == 0
|
|
assert c._context_probed is False
|
|
assert c._context_probe_persistable is False
|
|
assert c._previous_summary is None
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Plugin slot (PluginManager integration)
|
|
# ---------------------------------------------------------------------------
|
|
|
|
class TestPluginContextEngineSlot:
|
|
"""Test register_context_engine on PluginContext."""
|
|
|
|
def test_register_engine(self):
|
|
from hermes_cli.plugins import PluginManager, PluginContext, PluginManifest
|
|
mgr = PluginManager()
|
|
manifest = PluginManifest(name="test-lcm")
|
|
ctx = PluginContext(manifest, mgr)
|
|
|
|
engine = StubEngine()
|
|
ctx.register_context_engine(engine)
|
|
|
|
assert mgr._context_engine is engine
|
|
assert mgr._context_engine.name == "stub"
|
|
|
|
|
|
|
|
def test_get_plugin_context_engine(self):
|
|
from hermes_cli.plugins import PluginManager, get_plugin_context_engine
|
|
import hermes_cli.plugins as plugins_mod
|
|
|
|
# Inject a test manager
|
|
old_mgr = plugins_mod._plugin_manager
|
|
try:
|
|
mgr = PluginManager()
|
|
plugins_mod._plugin_manager = mgr
|
|
|
|
assert get_plugin_context_engine() is None
|
|
|
|
engine = StubEngine()
|
|
mgr._context_engine = engine
|
|
assert get_plugin_context_engine() is engine
|
|
finally:
|
|
plugins_mod._plugin_manager = old_mgr
|
|
|
|
|
|
|
|
class TestPluginContextEngineDeepCopy:
|
|
"""Verify that the plugin context engine singleton is deep-copied before
|
|
mutation in agent_init — regression test for #42449."""
|
|
|
|
|
|
def test_deepcopy_preserves_engine_name(self):
|
|
"""Deep-copied engine retains its identity (name property)."""
|
|
import copy
|
|
engine = StubEngine(context_length=500000)
|
|
clone = copy.deepcopy(engine)
|
|
assert clone.name == engine.name == "stub"
|
|
|
|
def test_deepcopy_preserves_compressor_state(self):
|
|
"""Deep-copied engine starts with the same token counters."""
|
|
import copy
|
|
engine = StubEngine(context_length=500000)
|
|
engine.last_prompt_tokens = 1000
|
|
engine.last_total_tokens = 1500
|
|
engine.compression_count = 3
|
|
|
|
clone = copy.deepcopy(engine)
|
|
assert clone.last_prompt_tokens == 1000
|
|
assert clone.last_total_tokens == 1500
|
|
assert clone.compression_count == 3
|
|
assert clone is not engine
|
|
|
|
|
|
|
|
class TestInitAgentDoesNotMutatePluginSingleton:
|
|
"""Regression coverage for #42449: a child agent's init must not mutate the
|
|
shared plugin context-engine singleton via update_model().
|
|
|
|
Note: these replicate the init_agent selection-block *pattern*; the production
|
|
seam (``_select_context_engine`` → ``clone_for_agent()``) is driven directly by
|
|
``tests/agent/test_plugin_context_engine_clone.py``.
|
|
"""
|
|
|
|
def test_child_init_does_not_corrupt_parent_singleton(self, monkeypatch):
|
|
import hermes_cli.plugins as plugins_mod
|
|
from hermes_cli.plugins import PluginManager
|
|
|
|
# Register a "parent" engine as the global plugin singleton, sized for
|
|
# a 1M-context model (DeepSeek-style), threshold 20% => 200K.
|
|
singleton = StubEngine(context_length=1_000_000, threshold_pct=0.20)
|
|
old_mgr = plugins_mod._plugin_manager
|
|
try:
|
|
mgr = PluginManager()
|
|
mgr._context_engine = singleton
|
|
plugins_mod._plugin_manager = mgr
|
|
|
|
# Replicate init_agent's fallback selection-block pattern: fetch the
|
|
# singleton, deepcopy it, then mutate the copy via update_model with
|
|
# a SMALLER child context (MiniMax-style 204800).
|
|
import copy
|
|
from hermes_cli.plugins import get_plugin_context_engine
|
|
|
|
_candidate = get_plugin_context_engine()
|
|
assert _candidate is singleton
|
|
_selected_engine = copy.deepcopy(_candidate)
|
|
_selected_engine.update_model(
|
|
model="MiniMax-M2", context_length=204800, provider="minimax",
|
|
)
|
|
|
|
# The child's smaller context must NOT leak back into the parent
|
|
# singleton (the #42449 corruption).
|
|
assert singleton.context_length == 1_000_000, (
|
|
"parent singleton context_length was corrupted by child init"
|
|
)
|
|
assert singleton.threshold_tokens == 200_000
|
|
# And the child's own engine reflects the child model.
|
|
assert _selected_engine.context_length == 204800
|
|
assert _selected_engine is not singleton
|
|
finally:
|
|
plugins_mod._plugin_manager = old_mgr
|
|
|
|
def test_unpicklable_engine_falls_back_gracefully(self, monkeypatch):
|
|
"""Copy-failure path: an engine holding uncopyable state (a lock — the
|
|
plugin docs prescribe locks/DB connections for stateful engines) makes
|
|
copy.deepcopy raise. init_agent must NOT silently drop it with a
|
|
misleading 'not found'; it falls back to the built-in compressor and
|
|
logs an accurate copy-failure warning. Regression for the deepcopy-
|
|
copy-failure path."""
|
|
import threading
|
|
|
|
class _UncopyableEngine(StubEngine):
|
|
def __init__(self):
|
|
super().__init__(context_length=1_000_000, threshold_pct=0.20)
|
|
self._lock = threading.RLock() # RLock can't be deepcopied
|
|
|
|
engine = _UncopyableEngine()
|
|
# Sanity: the engine genuinely defeats deepcopy.
|
|
import copy
|
|
with pytest.raises(Exception):
|
|
copy.deepcopy(engine)
|
|
|
|
# Replicate the init_agent fallback block's copy-failure handling.
|
|
selected = None
|
|
copy_failed = False
|
|
try:
|
|
selected = copy.deepcopy(engine)
|
|
except Exception:
|
|
copy_failed = True
|
|
selected = None
|
|
|
|
assert copy_failed is True
|
|
assert selected is None
|
|
# The original engine is untouched (no partial mutation).
|
|
assert engine.context_length == 1_000_000
|