feat(desktop): openExternalFileForIpc opens files via OS handler
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,32 @@
|
||||
"""Shared fixtures for tests/acp_adapter.
|
||||
|
||||
Keeps the ACP server tests offline: ``HermesACPAgent._build_model_state``
|
||||
calls ``hermes_cli.inventory.build_models_payload``, which (without this
|
||||
fixture) performs live network fetches — models.dev registry, GitHub model
|
||||
catalog, Copilot token exchange, Anthropic model list — adding ~3s of real
|
||||
SSL/socket time to every test that creates or loads a session (~147s total
|
||||
for test_server.py alone).
|
||||
|
||||
Tests that assert model-state behavior re-patch these same attributes with
|
||||
``unittest.mock.patch`` / ``monkeypatch``; inner patches win, so this
|
||||
default is transparent to them.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _offline_model_inventory(monkeypatch):
|
||||
"""Stub the shared model inventory so ACP tests never hit the network."""
|
||||
import hermes_cli.inventory as inventory
|
||||
|
||||
class _StubPickerContext:
|
||||
def with_overrides(self, **_kwargs):
|
||||
return self
|
||||
|
||||
monkeypatch.setattr(inventory, "load_picker_context", lambda: _StubPickerContext())
|
||||
monkeypatch.setattr(
|
||||
inventory,
|
||||
"build_models_payload",
|
||||
lambda *_args, **_kwargs: {"providers": []},
|
||||
)
|
||||
@@ -0,0 +1,98 @@
|
||||
"""ACP ``session/set_model`` and the dashboard main slot validate through ``switch_model``.
|
||||
|
||||
Both surfaces used to accept any string (``parse_model_input`` + ``detect_provider_for_model``
|
||||
for ACP; bare provider/model normalization for ``POST /api/model/set``), so a model no catalog
|
||||
knew — or a provider with no credentials — was handed to the session / written to config.yaml
|
||||
and only failed at inference time. They now share the CLI/gateway/TUI ``/model`` pipeline: a
|
||||
rejection from ``switch_model`` is a rejection on these surfaces too, and an acceptance carries
|
||||
the resolved (provider, model) — an explicit ``provider:model`` prefix is honoured as
|
||||
``--provider`` (#59089), never re-detected.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import types
|
||||
|
||||
import pytest
|
||||
|
||||
from hermes_cli.model_switch import ModelSwitchResult
|
||||
|
||||
|
||||
def _acp_agent():
|
||||
from acp_adapter.server import HermesACPAgent
|
||||
made: dict = {}
|
||||
|
||||
class _SM:
|
||||
def _make_agent(self, **kw):
|
||||
made.update(kw)
|
||||
return types.SimpleNamespace(provider=kw.get("requested_provider"), model=kw.get("model"))
|
||||
|
||||
def save_session(self, sid):
|
||||
pass
|
||||
|
||||
return HermesACPAgent(session_manager=_SM()), made
|
||||
|
||||
|
||||
def _state():
|
||||
return types.SimpleNamespace(
|
||||
session_id="s1", cwd=".", model="claude-sonnet-5",
|
||||
agent=types.SimpleNamespace(provider="anthropic", base_url="https://api.anthropic.com", api_key="k"))
|
||||
|
||||
|
||||
def test_acp_and_dashboard_reject_what_switch_model_rejects(monkeypatch):
|
||||
rejected = ModelSwitchResult(success=False, error_message="Unknown provider 'notaprovider'.")
|
||||
monkeypatch.setattr("hermes_cli.model_switch.switch_model", lambda **_kw: rejected)
|
||||
|
||||
agent, made = _acp_agent()
|
||||
state = _state()
|
||||
with pytest.raises(ValueError, match="Unknown provider"):
|
||||
agent._switch_model(state, "notaprovider:whatever")
|
||||
assert made == {} and state.model == "claude-sonnet-5" # session untouched
|
||||
|
||||
from fastapi import HTTPException
|
||||
from hermes_cli.web_server_config import _apply_model_assignment_sync
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
_apply_model_assignment_sync("main", "notaprovider", "whatever", "", "")
|
||||
assert exc.value.status_code == 400 and "Unknown provider" in exc.value.detail
|
||||
|
||||
|
||||
def test_acp_explicit_provider_prefix_becomes_explicit_provider(monkeypatch):
|
||||
seen: dict = {}
|
||||
|
||||
def _switch(**kw):
|
||||
seen.update(kw)
|
||||
return ModelSwitchResult(success=True, new_model=kw["raw_input"], target_provider=kw["explicit_provider"])
|
||||
|
||||
monkeypatch.setattr("hermes_cli.model_switch.switch_model", _switch)
|
||||
agent, made = _acp_agent()
|
||||
old, new_provider, model = agent._switch_model(_state(), "anthropic:claude-sonnet-5", keep_endpoint=True)
|
||||
assert (seen["explicit_provider"], seen["raw_input"]) == ("anthropic", "claude-sonnet-5")
|
||||
assert (old, new_provider, model) == ("anthropic", "anthropic", "claude-sonnet-5")
|
||||
assert made["requested_provider"] == "anthropic" and made["base_url"] == "https://api.anthropic.com"
|
||||
|
||||
|
||||
def test_acp_set_session_model_runs_switch_model_off_the_event_loop(monkeypatch):
|
||||
"""``switch_model`` does ~10 s of sync network I/O on a cold cache; ACP must run it on a
|
||||
worker thread (like the gateway) or every session in the process stalls."""
|
||||
import asyncio
|
||||
import threading
|
||||
|
||||
seen: dict = {}
|
||||
|
||||
def _switch(**kw):
|
||||
seen["thread"] = threading.current_thread()
|
||||
return ModelSwitchResult(success=True, new_model=kw["raw_input"], target_provider="anthropic")
|
||||
|
||||
monkeypatch.setattr("hermes_cli.model_switch.switch_model", _switch)
|
||||
agent, _made = _acp_agent()
|
||||
state = _state()
|
||||
agent.session_manager.get_session = lambda sid: state
|
||||
|
||||
async def _run():
|
||||
loop_thread = threading.current_thread()
|
||||
resp = await agent.set_session_model("anthropic:claude-sonnet-5", "s1")
|
||||
return resp, loop_thread
|
||||
|
||||
resp, loop_thread = asyncio.run(_run())
|
||||
assert resp is not None and state.model == "claude-sonnet-5"
|
||||
assert seen["thread"] is not loop_thread
|
||||
@@ -0,0 +1,197 @@
|
||||
"""Tests for GHSA-96vc-wcxf-jjff and GHSA-qg5c-hvr5-hjgr.
|
||||
|
||||
Two related ACP approval-flow issues:
|
||||
- 96vc: ACP didn't set HERMES_EXEC_ASK, so `check_all_command_guards`
|
||||
took the non-interactive auto-approve path and never consulted the
|
||||
ACP-supplied callback.
|
||||
- qg5c: `_approval_callback` was a module-global in terminal_tool;
|
||||
overlapping ACP sessions overwrote each other's callback slot.
|
||||
|
||||
Both fixed together by:
|
||||
1. Setting HERMES_EXEC_ASK inside _run_agent (wraps the agent call).
|
||||
2. Storing the callback in thread-local state so concurrent executor
|
||||
threads don't collide.
|
||||
"""
|
||||
|
||||
import threading
|
||||
|
||||
import pytest
|
||||
from tools import approval_context
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _isolate_approval_state(monkeypatch):
|
||||
"""Keep these security regression tests hermetic.
|
||||
|
||||
Earlier tests (e.g. tests/acp_adapter/test_permissions.py) lazily load the
|
||||
developer's real ``~/.hermes/config.yaml`` command allowlist into
|
||||
``tools.approval._permanent_approved``. If that allowlist contains a
|
||||
pattern like "recursive delete", ``rm -rf …`` is auto-approved before
|
||||
the interactive callback fires and the GHSA regression assertions fail
|
||||
for reasons unrelated to the code under test.
|
||||
"""
|
||||
import tools.approval as _approval
|
||||
from tools import approval_context
|
||||
|
||||
monkeypatch.setattr(_approval, "_permanent_approved", set())
|
||||
monkeypatch.setattr(_approval, "_session_approved", {})
|
||||
# These tests assert the *manual* interactive-callback path. The default
|
||||
# config is approvals.mode=smart, whose guardian LLM can auto-approve the
|
||||
# command before the callback is consulted (test-order dependent, since
|
||||
# load_config() caching decides which config file is in effect). Pin the
|
||||
# mode so the GHSA regression path is what actually runs.
|
||||
monkeypatch.setattr(approval_context, "_get_approval_mode", lambda: "manual")
|
||||
|
||||
|
||||
class TestThreadLocalApprovalCallback:
|
||||
"""GHSA-qg5c-hvr5-hjgr: set_approval_callback must be per-thread so
|
||||
concurrent ACP sessions don't stomp on each other's handlers."""
|
||||
|
||||
def test_set_and_get_in_same_thread(self):
|
||||
from tools.terminal_tool import (
|
||||
set_approval_callback,
|
||||
_get_approval_callback,
|
||||
)
|
||||
|
||||
cb1 = lambda cmd, desc: "once" # noqa: E731
|
||||
set_approval_callback(cb1)
|
||||
assert _get_approval_callback() is cb1
|
||||
|
||||
def test_callback_not_visible_in_different_thread(self):
|
||||
"""Thread A's callback is NOT visible to Thread B."""
|
||||
from tools.terminal_tool import (
|
||||
set_approval_callback,
|
||||
_get_approval_callback,
|
||||
)
|
||||
|
||||
cb_a = lambda cmd, desc: "thread_a" # noqa: E731
|
||||
cb_b = lambda cmd, desc: "thread_b" # noqa: E731
|
||||
|
||||
seen_in_a = []
|
||||
seen_in_b = []
|
||||
|
||||
def thread_a():
|
||||
set_approval_callback(cb_a)
|
||||
# Pause so thread B has time to set its own callback
|
||||
import time
|
||||
time.sleep(0.05)
|
||||
seen_in_a.append(_get_approval_callback())
|
||||
|
||||
def thread_b():
|
||||
set_approval_callback(cb_b)
|
||||
import time
|
||||
time.sleep(0.05)
|
||||
seen_in_b.append(_get_approval_callback())
|
||||
|
||||
ta = threading.Thread(target=thread_a)
|
||||
tb = threading.Thread(target=thread_b)
|
||||
ta.start()
|
||||
tb.start()
|
||||
ta.join()
|
||||
tb.join()
|
||||
|
||||
# Each thread must see ONLY its own callback — not the other's
|
||||
assert seen_in_a == [cb_a]
|
||||
assert seen_in_b == [cb_b]
|
||||
|
||||
def test_main_thread_callback_not_leaked_to_worker(self):
|
||||
"""A callback set in the main thread does NOT leak into a
|
||||
freshly-spawned worker thread."""
|
||||
from tools.terminal_tool import (
|
||||
set_approval_callback,
|
||||
_get_approval_callback,
|
||||
)
|
||||
|
||||
cb_main = lambda cmd, desc: "main" # noqa: E731
|
||||
set_approval_callback(cb_main)
|
||||
|
||||
worker_saw = []
|
||||
|
||||
def worker():
|
||||
worker_saw.append(_get_approval_callback())
|
||||
|
||||
t = threading.Thread(target=worker)
|
||||
t.start()
|
||||
t.join()
|
||||
|
||||
# Worker thread has no callback set — TLS is empty for it
|
||||
assert worker_saw == [None]
|
||||
# Main thread still has its callback
|
||||
assert _get_approval_callback() is cb_main
|
||||
|
||||
|
||||
def test_sudo_password_cache_does_not_leak_across_threads(self):
|
||||
"""Interactive sudo cache must not bleed into another executor thread."""
|
||||
from tools.terminal_tool_sudo import (
|
||||
_get_cached_sudo_password,
|
||||
_reset_cached_sudo_passwords,
|
||||
_set_cached_sudo_password,
|
||||
)
|
||||
|
||||
_reset_cached_sudo_passwords()
|
||||
_set_cached_sudo_password("main-thread-password")
|
||||
|
||||
worker_saw = []
|
||||
|
||||
def worker():
|
||||
worker_saw.append(_get_cached_sudo_password())
|
||||
|
||||
t = threading.Thread(target=worker)
|
||||
t.start()
|
||||
t.join()
|
||||
|
||||
assert worker_saw == [""]
|
||||
assert _get_cached_sudo_password() == "main-thread-password"
|
||||
|
||||
|
||||
|
||||
class TestAcpExecAskGate:
|
||||
"""GHSA-96vc-wcxf-jjff: ACP's _run_agent must set HERMES_INTERACTIVE so
|
||||
that tools.approval.check_all_command_guards takes the CLI-interactive
|
||||
path (consults the registered callback via prompt_dangerous_approval)
|
||||
instead of the non-interactive auto-approve shortcut.
|
||||
|
||||
(HERMES_EXEC_ASK takes the gateway-queue path which requires a
|
||||
notify_cb registered in _gateway_notify_cbs — not applicable to ACP,
|
||||
which uses a direct callback shape.)"""
|
||||
|
||||
def test_interactive_env_var_routes_to_callback(self, monkeypatch):
|
||||
"""When HERMES_INTERACTIVE is set and an approval callback is
|
||||
registered, a dangerous command must route through the callback."""
|
||||
# Clean env
|
||||
monkeypatch.delenv("HERMES_INTERACTIVE", raising=False)
|
||||
monkeypatch.delenv("HERMES_GATEWAY_SESSION", raising=False)
|
||||
monkeypatch.delenv("HERMES_EXEC_ASK", raising=False)
|
||||
monkeypatch.delenv("HERMES_YOLO_MODE", raising=False)
|
||||
|
||||
from tools.approval import check_all_command_guards
|
||||
|
||||
called_with = []
|
||||
|
||||
def fake_cb(command, description, *, allow_permanent=True):
|
||||
called_with.append((command, description))
|
||||
return "once"
|
||||
|
||||
# Without HERMES_INTERACTIVE: takes auto-approve path, callback NOT called
|
||||
result = check_all_command_guards(
|
||||
"rm -rf /tmp/test-exec-ask", "local", approval_callback=fake_cb,
|
||||
)
|
||||
assert result["approved"] is True
|
||||
assert called_with == [], (
|
||||
"without HERMES_INTERACTIVE the non-interactive auto-approve "
|
||||
"path should fire without consulting the callback"
|
||||
)
|
||||
|
||||
# With HERMES_INTERACTIVE: callback IS called, approval flows through it
|
||||
monkeypatch.setenv("HERMES_INTERACTIVE", "1")
|
||||
called_with.clear()
|
||||
result = check_all_command_guards(
|
||||
"rm -rf /tmp/test-exec-ask", "local", approval_callback=fake_cb,
|
||||
)
|
||||
assert called_with, (
|
||||
"with HERMES_INTERACTIVE the approval path should consult the "
|
||||
"registered callback — this was the ACP bypass in "
|
||||
"GHSA-96vc-wcxf-jjff"
|
||||
)
|
||||
assert result["approved"] is True
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
"""Tests for acp_adapter.auth — provider detection."""
|
||||
|
||||
from acp_adapter.auth import (
|
||||
TERMINAL_SETUP_AUTH_METHOD_ID,
|
||||
build_auth_methods,
|
||||
detect_provider,
|
||||
)
|
||||
|
||||
|
||||
class TestDetectProviderPresence:
|
||||
def test_has_provider_with_resolved_runtime(self, monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"hermes_cli.runtime_provider.resolve_runtime_provider",
|
||||
lambda: {"provider": "openrouter", "api_key": "sk-or-test"},
|
||||
)
|
||||
assert detect_provider() is not None
|
||||
|
||||
def test_has_provider_false_without_credentials(self, monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"hermes_cli.runtime_provider.resolve_runtime_provider",
|
||||
lambda: {"provider": "openrouter", "api_key": ""},
|
||||
)
|
||||
assert detect_provider() is None
|
||||
|
||||
|
||||
class TestDetectProvider:
|
||||
def test_detect_openrouter(self, monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"hermes_cli.runtime_provider.resolve_runtime_provider",
|
||||
lambda: {"provider": "openrouter", "api_key": "sk-or-test"},
|
||||
)
|
||||
assert detect_provider() == "openrouter"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
class TestBuildAuthMethods:
|
||||
def test_build_auth_methods_returns_provider_and_terminal_when_configured(self, monkeypatch):
|
||||
monkeypatch.setattr("acp_adapter.auth.detect_provider", lambda: "openrouter")
|
||||
|
||||
methods = build_auth_methods()
|
||||
payloads = [method.model_dump(by_alias=True, exclude_none=True) for method in methods]
|
||||
|
||||
assert payloads[0]["id"] == "openrouter"
|
||||
assert payloads[0]["name"] == "openrouter runtime credentials"
|
||||
assert any(payload["id"] == TERMINAL_SETUP_AUTH_METHOD_ID for payload in payloads)
|
||||
terminal = next(payload for payload in payloads if payload["id"] == TERMINAL_SETUP_AUTH_METHOD_ID)
|
||||
assert terminal["type"] == "terminal"
|
||||
assert terminal["args"] == ["--setup"]
|
||||
|
||||
@@ -0,0 +1,123 @@
|
||||
"""Tests for ACP pre-edit approval gating."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
from acp_adapter.edit_approval import (
|
||||
EditProposal,
|
||||
build_acp_edit_tool_call,
|
||||
set_edit_approval_requester,
|
||||
should_auto_approve_edit,
|
||||
)
|
||||
from model_tools import handle_function_call
|
||||
|
||||
|
||||
def teardown_function() -> None:
|
||||
set_edit_approval_requester(None)
|
||||
|
||||
|
||||
def test_acp_permission_tool_call_uses_edit_kind_and_diff_content():
|
||||
proposal = EditProposal(
|
||||
tool_name="write_file",
|
||||
path="demo.txt",
|
||||
old_text="old\n",
|
||||
new_text="new\n",
|
||||
arguments={"path": "demo.txt", "content": "new\n"},
|
||||
)
|
||||
|
||||
tool_call = build_acp_edit_tool_call(proposal)
|
||||
|
||||
assert tool_call.kind == "edit"
|
||||
assert tool_call.status == "pending"
|
||||
assert tool_call.rawInput == {"tool": "write_file", "arguments": proposal.arguments}
|
||||
assert len(tool_call.content) == 1
|
||||
diff = tool_call.content[0]
|
||||
assert diff.path == "demo.txt"
|
||||
assert diff.oldText == "old\n"
|
||||
assert diff.newText == "new\n"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_requester_exception_denies_and_does_not_mutate(tmp_path):
|
||||
target = tmp_path / "sample.txt"
|
||||
target.write_text("before\n", encoding="utf-8")
|
||||
|
||||
def boom(_proposal):
|
||||
raise RuntimeError("zed disconnected")
|
||||
|
||||
set_edit_approval_requester(boom)
|
||||
|
||||
result = json.loads(
|
||||
handle_function_call(
|
||||
"write_file",
|
||||
{"path": str(target), "content": "after\n"},
|
||||
task_id="acp-edit-exception",
|
||||
)
|
||||
)
|
||||
|
||||
assert "error" in result
|
||||
assert "Edit approval denied" in result["error"]
|
||||
assert target.read_text(encoding="utf-8") == "before\n"
|
||||
|
||||
|
||||
def test_patch_replace_rejection_does_not_mutate(tmp_path):
|
||||
target = tmp_path / "sample.txt"
|
||||
target.write_text("alpha\nbeta\n", encoding="utf-8")
|
||||
|
||||
set_edit_approval_requester(lambda _proposal: False)
|
||||
|
||||
result = json.loads(
|
||||
handle_function_call(
|
||||
"patch",
|
||||
{
|
||||
"mode": "replace",
|
||||
"path": str(target),
|
||||
"old_string": "beta\n",
|
||||
"new_string": "gamma\n",
|
||||
},
|
||||
task_id="acp-patch-reject",
|
||||
)
|
||||
)
|
||||
|
||||
assert "error" in result
|
||||
assert "Edit approval denied" in result["error"]
|
||||
assert target.read_text(encoding="utf-8") == "alpha\nbeta\n"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_workspace_auto_approval_allows_workspace_and_tmp_but_not_sensitive(tmp_path):
|
||||
workspace_file = tmp_path / "src.py"
|
||||
# Use tempfile.gettempdir() so this test exercises the same code path on
|
||||
# Linux (`/tmp`), macOS (`/private/var/folders/...`) and Windows
|
||||
# (`%LOCALAPPDATA%\Temp`). Before the fix this branch only worked on Linux.
|
||||
tmp_file = Path(tempfile.gettempdir()) / "hermes-acp-auto-approve-test.txt"
|
||||
env_file = tmp_path / ".env"
|
||||
|
||||
assert should_auto_approve_edit(
|
||||
EditProposal("write_file", str(workspace_file), None, "x", {}),
|
||||
"workspace_session",
|
||||
str(tmp_path),
|
||||
)
|
||||
assert should_auto_approve_edit(
|
||||
EditProposal("write_file", str(tmp_file), None, "x", {}),
|
||||
"workspace_session",
|
||||
str(tmp_path),
|
||||
)
|
||||
assert not should_auto_approve_edit(
|
||||
EditProposal("write_file", str(env_file), None, "SECRET=x", {}),
|
||||
"session",
|
||||
str(tmp_path),
|
||||
)
|
||||
@@ -0,0 +1,43 @@
|
||||
"""Only new, contentless ACP sessions are ephemeral; existing state remains durable."""
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
|
||||
from acp_adapter.session import SessionManager
|
||||
from hermes_state import SessionDB
|
||||
|
||||
|
||||
def test_new_session_persists_only_when_content_exists(tmp_path):
|
||||
db = SessionDB(tmp_path / "state.db")
|
||||
manager = SessionManager(db=db, agent_factory=lambda: SimpleNamespace(model="fixture"))
|
||||
state = manager.create_session(cwd=str(tmp_path))
|
||||
assert db.get_session(state.session_id) is None
|
||||
manager.update_cwd(state.session_id, str(tmp_path / "moved"))
|
||||
manager.save_session(state.session_id)
|
||||
empty_fork = manager.fork_session(state.session_id)
|
||||
assert db.get_session(state.session_id) is None
|
||||
assert db.get_session(empty_fork.session_id) is None
|
||||
state.history.append({"role": "user", "content": "kept content"})
|
||||
manager.save_session(state.session_id)
|
||||
fork = manager.fork_session(state.session_id)
|
||||
for sid in (state.session_id, fork.session_id):
|
||||
assert db.get_session(sid)["source"] == "acp"
|
||||
assert db.get_messages_as_conversation(sid)[0]["content"] == "kept content"
|
||||
db.close()
|
||||
|
||||
|
||||
def test_existing_empty_history_still_updates_metadata(tmp_path):
|
||||
db = SessionDB(tmp_path / "state.db")
|
||||
manager = SessionManager(db=db, agent_factory=lambda: SimpleNamespace(model="fixture"))
|
||||
# An old, unprompted ACP client may still own its row: source is not liveness.
|
||||
db.create_session(session_id="existing", source="acp", model="original")
|
||||
state = manager.get_session("existing")
|
||||
assert state is not None
|
||||
assert not state.history
|
||||
state.model = "selected-model"
|
||||
manager.update_cwd(state.session_id, str(tmp_path / "selected"))
|
||||
row = db.get_session(state.session_id)
|
||||
assert row["model"] == "selected-model"
|
||||
assert json.loads(row["model_config"])["cwd"] == state.cwd
|
||||
assert manager.get_session(state.session_id) is state
|
||||
assert db.list_never_active_keyed_sessions(older_than_days=0) == []
|
||||
db.close()
|
||||
@@ -0,0 +1,90 @@
|
||||
"""Tests for acp_adapter.entry startup wiring."""
|
||||
|
||||
import sys
|
||||
|
||||
import acp
|
||||
import pytest
|
||||
|
||||
from acp_adapter import entry
|
||||
|
||||
|
||||
def test_main_enables_unstable_protocol(monkeypatch):
|
||||
calls = {}
|
||||
|
||||
async def fake_run_agent(agent, **kwargs):
|
||||
calls["kwargs"] = kwargs
|
||||
|
||||
monkeypatch.setattr(entry, "_setup_logging", lambda: None)
|
||||
monkeypatch.setattr(entry, "_load_env", lambda: None)
|
||||
monkeypatch.setattr(acp, "run_agent", fake_run_agent)
|
||||
|
||||
entry.main([])
|
||||
|
||||
assert calls["kwargs"]["use_unstable_protocol"] is True
|
||||
|
||||
|
||||
def test_main_skips_configured_mcp_discovery_when_requested(monkeypatch):
|
||||
discovery_calls = []
|
||||
|
||||
async def fake_run_agent(agent, **kwargs):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr(entry, "_setup_logging", lambda: None)
|
||||
monkeypatch.setattr(entry, "_load_env", lambda: None)
|
||||
monkeypatch.setenv("HERMES_ACP_SKIP_CONFIGURED_MCP", "1")
|
||||
monkeypatch.setattr(
|
||||
"tools.mcp_tool_discovery.discover_mcp_tools",
|
||||
lambda: discovery_calls.append(True),
|
||||
)
|
||||
monkeypatch.setattr(acp, "run_agent", fake_run_agent)
|
||||
|
||||
entry.main([])
|
||||
|
||||
assert discovery_calls == []
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_main_setup_offers_browser_install_when_tty(monkeypatch):
|
||||
"""When stdin is a TTY and the user answers yes, model setup is followed
|
||||
by a browser-tools bootstrap call."""
|
||||
monkeypatch.setattr("hermes_cli.main.main", lambda: None)
|
||||
monkeypatch.setattr("sys.stdin.isatty", lambda: True)
|
||||
monkeypatch.setattr("builtins.input", lambda *_args, **_kwargs: "y")
|
||||
|
||||
bootstrap_calls = []
|
||||
monkeypatch.setattr(
|
||||
entry,
|
||||
"_run_setup_browser",
|
||||
lambda assume_yes=False: bootstrap_calls.append(assume_yes) or 0,
|
||||
)
|
||||
|
||||
entry.main(["--setup"])
|
||||
|
||||
assert bootstrap_calls == [False]
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_main_setup_browser_propagates_browser_failure(monkeypatch):
|
||||
"""If browser install fails, exit code is 1."""
|
||||
def fake_ensure(dep, interactive=True):
|
||||
return dep != "browser" # browser fails
|
||||
|
||||
monkeypatch.setattr("hermes_cli.dep_ensure.ensure_dependency", fake_ensure)
|
||||
|
||||
with pytest.raises(SystemExit) as excinfo:
|
||||
entry.main(["--setup-browser"])
|
||||
assert excinfo.value.code == 1
|
||||
@@ -0,0 +1,312 @@
|
||||
"""Tests for acp_adapter.events — callback factories for ACP notifications."""
|
||||
|
||||
import asyncio
|
||||
import gc
|
||||
import uuid
|
||||
import warnings
|
||||
from concurrent.futures import Future
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import acp
|
||||
from acp.schema import AgentPlanUpdate
|
||||
|
||||
from acp_adapter.events import (
|
||||
_build_plan_update_from_todo_result,
|
||||
_send_update,
|
||||
make_message_cb,
|
||||
make_step_cb,
|
||||
make_thinking_cb,
|
||||
make_tool_progress_cb,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def mock_conn():
|
||||
"""Mock ACP Client connection."""
|
||||
conn = MagicMock(spec=acp.Client)
|
||||
conn.session_update = AsyncMock()
|
||||
return conn
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def event_loop_fixture():
|
||||
"""Create a real event loop for testing threadsafe coroutine submission."""
|
||||
loop = asyncio.new_event_loop()
|
||||
yield loop
|
||||
loop.close()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tool progress callback
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestToolProgressCallback:
|
||||
def test_emits_tool_call_start(self, mock_conn, event_loop_fixture):
|
||||
"""Tool progress should emit a ToolCallStart update."""
|
||||
tool_call_ids = {}
|
||||
tool_call_meta = {}
|
||||
loop = event_loop_fixture
|
||||
|
||||
cb = make_tool_progress_cb(mock_conn, "session-1", loop, tool_call_ids, tool_call_meta)
|
||||
|
||||
# Run callback in the event loop context
|
||||
with patch("acp_adapter.events.asyncio.run_coroutine_threadsafe") as mock_rcts:
|
||||
future = MagicMock(spec=Future)
|
||||
future.result.return_value = None
|
||||
mock_rcts.return_value = future
|
||||
|
||||
cb("tool.started", "terminal", "$ ls -la", {"command": "ls -la"})
|
||||
|
||||
# Should have tracked the tool call ID
|
||||
assert "terminal" in tool_call_ids
|
||||
|
||||
# Should have called run_coroutine_threadsafe
|
||||
mock_rcts.assert_called_once()
|
||||
coro = mock_rcts.call_args[0][0]
|
||||
# The coroutine should be conn.session_update
|
||||
assert mock_conn.session_update.called or coro is not None
|
||||
|
||||
|
||||
|
||||
def test_duplicate_same_name_tool_calls_use_fifo_ids(self, mock_conn, event_loop_fixture):
|
||||
"""Multiple same-name tool calls should be tracked independently in order."""
|
||||
tool_call_ids = {}
|
||||
tool_call_meta = {}
|
||||
loop = event_loop_fixture
|
||||
|
||||
progress_cb = make_tool_progress_cb(mock_conn, "session-1", loop, tool_call_ids, tool_call_meta)
|
||||
step_cb = make_step_cb(mock_conn, "session-1", loop, tool_call_ids, tool_call_meta)
|
||||
|
||||
with patch("acp_adapter.events.asyncio.run_coroutine_threadsafe") as mock_rcts:
|
||||
future = MagicMock(spec=Future)
|
||||
future.result.return_value = None
|
||||
mock_rcts.return_value = future
|
||||
|
||||
progress_cb("tool.started", "terminal", "$ ls", {"command": "ls"})
|
||||
progress_cb("tool.started", "terminal", "$ pwd", {"command": "pwd"})
|
||||
assert len(tool_call_ids["terminal"]) == 2
|
||||
|
||||
step_cb(1, [{"name": "terminal", "result": "ok-1"}])
|
||||
assert len(tool_call_ids["terminal"]) == 1
|
||||
|
||||
step_cb(2, [{"name": "terminal", "result": "ok-2"}])
|
||||
assert "terminal" not in tool_call_ids
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Thinking callback
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Step callback
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestStepCallback:
|
||||
def test_completes_tracked_tool_calls(self, mock_conn, event_loop_fixture):
|
||||
"""Step callback should mark tracked tools as completed."""
|
||||
tool_call_ids = {"terminal": "tc-abc123"}
|
||||
loop = event_loop_fixture
|
||||
|
||||
cb = make_step_cb(mock_conn, "session-1", loop, tool_call_ids, {})
|
||||
|
||||
with patch("acp_adapter.events.asyncio.run_coroutine_threadsafe") as mock_rcts:
|
||||
future = MagicMock(spec=Future)
|
||||
future.result.return_value = None
|
||||
mock_rcts.return_value = future
|
||||
|
||||
cb(1, [{"name": "terminal", "result": "success"}])
|
||||
|
||||
# Tool should have been removed from tracking
|
||||
assert "terminal" not in tool_call_ids
|
||||
mock_rcts.assert_called_once()
|
||||
|
||||
|
||||
|
||||
@pytest.mark.parametrize("raw, expected", [("", ""), (0, "0"), (False, "False")])
|
||||
def test_falsey_result_reaches_client_unchanged(self, mock_conn, event_loop_fixture, raw, expected):
|
||||
"""A present-but-falsey ``result`` is the tool's real output, not a missing key (#10845)."""
|
||||
from collections import deque
|
||||
|
||||
cb = make_step_cb(mock_conn, "session-1", event_loop_fixture, {"terminal": deque(["tc-f"])}, {})
|
||||
with patch("acp_adapter.events.asyncio.run_coroutine_threadsafe") as mock_rcts, \
|
||||
patch("acp_adapter.events.build_tool_complete") as mock_btc:
|
||||
mock_rcts.return_value = MagicMock(spec=Future)
|
||||
cb(1, [{"name": "terminal", "result": raw}])
|
||||
mock_btc.assert_called_once_with("tc-f", "terminal", result=expected, function_args=None, snapshot=None)
|
||||
|
||||
def test_result_passed_to_build_tool_complete(self, mock_conn, event_loop_fixture):
|
||||
"""Tool result from prev_tools dict is forwarded to build_tool_complete."""
|
||||
from collections import deque
|
||||
|
||||
tool_call_ids = {"terminal": deque(["tc-xyz789"])}
|
||||
loop = event_loop_fixture
|
||||
|
||||
cb = make_step_cb(mock_conn, "session-1", loop, tool_call_ids, {})
|
||||
|
||||
with patch("acp_adapter.events.asyncio.run_coroutine_threadsafe") as mock_rcts, \
|
||||
patch("acp_adapter.events.build_tool_complete") as mock_btc:
|
||||
future = MagicMock(spec=Future)
|
||||
future.result.return_value = None
|
||||
mock_rcts.return_value = future
|
||||
|
||||
# Provide a result string in the tool info dict
|
||||
cb(1, [{"name": "terminal", "result": '{"output": "hello"}'}])
|
||||
|
||||
mock_btc.assert_called_once_with(
|
||||
"tc-xyz789", "terminal", result='{"output": "hello"}', function_args=None, snapshot=None
|
||||
)
|
||||
|
||||
|
||||
|
||||
def test_tool_progress_captures_snapshot_metadata(self, mock_conn, event_loop_fixture):
|
||||
tool_call_ids = {}
|
||||
tool_call_meta = {}
|
||||
loop = event_loop_fixture
|
||||
|
||||
with patch("acp_adapter.events.make_tool_call_id", return_value="tc-meta"), \
|
||||
patch("acp_adapter.events._send_update") as mock_send, \
|
||||
patch("agent.display.capture_local_edit_snapshot", return_value="snapshot"):
|
||||
cb = make_tool_progress_cb(mock_conn, "session-1", loop, tool_call_ids, tool_call_meta)
|
||||
cb("tool.started", "write_file", None, {"path": "diff-test.txt", "content": "hello"})
|
||||
|
||||
assert list(tool_call_ids["write_file"]) == ["tc-meta"]
|
||||
assert tool_call_meta["tc-meta"] == {
|
||||
"args": {"path": "diff-test.txt", "content": "hello"},
|
||||
"snapshot": "snapshot",
|
||||
}
|
||||
mock_send.assert_called_once()
|
||||
|
||||
def test_todo_completion_emits_native_plan_update_after_tool_completion(self, mock_conn, event_loop_fixture):
|
||||
from collections import deque
|
||||
|
||||
tool_call_ids = {"todo": deque(["tc-todo"])}
|
||||
loop = event_loop_fixture
|
||||
cb = make_step_cb(mock_conn, "session-1", loop, tool_call_ids, {})
|
||||
todo_result = (
|
||||
'{"todos":['
|
||||
'{"id":"inspect","content":"Inspect ACP","status":"completed"},'
|
||||
'{"id":"patch","content":"Patch renderer","status":"in_progress"},'
|
||||
'{"id":"old","content":"Drop stale task","status":"cancelled"}'
|
||||
'],"summary":{"total":3}}'
|
||||
)
|
||||
|
||||
with patch("acp_adapter.events._send_update") as mock_send:
|
||||
cb(1, [{"name": "todo", "result": todo_result}])
|
||||
|
||||
updates = [call.args[3] for call in mock_send.call_args_list]
|
||||
assert [getattr(update, "session_update", None) for update in updates] == [
|
||||
"tool_call_update",
|
||||
"plan",
|
||||
]
|
||||
plan = updates[1]
|
||||
assert isinstance(plan, AgentPlanUpdate)
|
||||
assert [entry.content for entry in plan.entries] == [
|
||||
"Inspect ACP",
|
||||
"Patch renderer",
|
||||
"[cancelled] Drop stale task",
|
||||
]
|
||||
assert [entry.status for entry in plan.entries] == ["completed", "in_progress", "completed"]
|
||||
assert [entry.priority for entry in plan.entries] == ["medium", "medium", "medium"]
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Message callback
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Scheduler-failure regression
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestSendUpdate:
|
||||
def test_scheduler_failure_closes_update_coroutine(self, event_loop_fixture):
|
||||
"""If run_coroutine_threadsafe raises, _send_update must close the coro."""
|
||||
created = {"coro": None}
|
||||
|
||||
async def _session_update(session_id, update):
|
||||
return None
|
||||
|
||||
conn = MagicMock()
|
||||
|
||||
def _capture_update(session_id, update):
|
||||
created["coro"] = _session_update(session_id, update)
|
||||
return created["coro"]
|
||||
|
||||
conn.session_update = _capture_update
|
||||
|
||||
with warnings.catch_warnings(record=True) as caught:
|
||||
warnings.simplefilter("always")
|
||||
with patch(
|
||||
"agent.async_utils.asyncio.run_coroutine_threadsafe",
|
||||
side_effect=RuntimeError("scheduler down"),
|
||||
):
|
||||
_send_update(conn, "session-1", event_loop_fixture, {"type": "noop"})
|
||||
gc.collect()
|
||||
|
||||
assert created["coro"] is not None
|
||||
assert created["coro"].cr_frame is None
|
||||
# Only count warnings about THIS test's coroutine; other tests
|
||||
# may emit unrelated
|
||||
# "coroutine was never awaited" warnings that bleed through.
|
||||
runtime_warnings = [
|
||||
w for w in caught
|
||||
if issubclass(w.category, RuntimeWarning)
|
||||
and "was never awaited" in str(w.message)
|
||||
and "_session_update" in str(w.message)
|
||||
]
|
||||
assert runtime_warnings == []
|
||||
|
||||
|
||||
class TestAssistantMessageIds:
|
||||
"""Streamed chunks carry a per-message ACP messageId; the None flush sentinel starts a new one."""
|
||||
|
||||
def test_deltas_share_one_uuid_until_flush(self, mock_conn, event_loop_fixture):
|
||||
from acp_adapter.events import AssistantMessageIdAllocator
|
||||
|
||||
ids = AssistantMessageIdAllocator()
|
||||
cb = make_message_cb(mock_conn, "s", event_loop_fixture, ids)
|
||||
sent = []
|
||||
with patch("acp_adapter.events._send_update",
|
||||
side_effect=lambda c, s, l, u: sent.append(u)):
|
||||
cb("Hello ")
|
||||
cb("") # empty delta is ignored, not a flush
|
||||
cb("world")
|
||||
cb(None) # flush sentinel — closes the message
|
||||
cb("next turn")
|
||||
assert sent[0].message_id == sent[1].message_id
|
||||
assert sent[2].message_id != sent[0].message_id
|
||||
# ACP requires UUID-format message ids.
|
||||
assert uuid.UUID(sent[0].message_id) and uuid.UUID(sent[2].message_id)
|
||||
|
||||
def test_thought_chunks_carry_id(self, mock_conn, event_loop_fixture):
|
||||
from acp_adapter.events import AssistantMessageIdAllocator
|
||||
|
||||
ids = AssistantMessageIdAllocator()
|
||||
think = make_thinking_cb(mock_conn, "s", event_loop_fixture, ids)
|
||||
msg = make_message_cb(mock_conn, "s", event_loop_fixture, ids)
|
||||
sent = []
|
||||
with patch("acp_adapter.events._send_update",
|
||||
side_effect=lambda c, s, l, u: sent.append(u)):
|
||||
think("pondering")
|
||||
msg("answer")
|
||||
# Reasoning and answer of the same reply share one message id.
|
||||
assert sent[0].message_id == sent[1].message_id
|
||||
|
||||
def test_no_allocator_keeps_legacy_shape(self, mock_conn, event_loop_fixture):
|
||||
cb = make_message_cb(mock_conn, "s", event_loop_fixture)
|
||||
sent = []
|
||||
with patch("acp_adapter.events._send_update",
|
||||
side_effect=lambda c, s, l, u: sent.append(u)):
|
||||
cb("text")
|
||||
assert sent[0].message_id is None
|
||||
@@ -0,0 +1,313 @@
|
||||
"""End-to-end tests for ACP MCP server registration and tool-result reporting.
|
||||
|
||||
Exercises the full flow through the ACP server layer:
|
||||
new_session(mcpServers) → MCP tools registered → prompt() →
|
||||
tool_progress_callback (ToolCallStart) →
|
||||
step_callback with results (ToolCallUpdate with rawOutput) →
|
||||
session_update events arrive at the mock client
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import acp
|
||||
from acp.schema import (
|
||||
EnvVariable,
|
||||
HttpHeader,
|
||||
McpServerHttp,
|
||||
McpServerStdio,
|
||||
NewSessionResponse,
|
||||
PromptResponse,
|
||||
TextContentBlock,
|
||||
ToolCallProgress,
|
||||
ToolCallStart,
|
||||
)
|
||||
|
||||
from acp_adapter.server import HermesACPAgent
|
||||
from acp_adapter.session import SessionManager
|
||||
from acp_adapter.tools import build_tool_start
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def mock_manager():
|
||||
return SessionManager(agent_factory=lambda: MagicMock(name="MockAIAgent"))
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def acp_agent(mock_manager):
|
||||
return HermesACPAgent(session_manager=mock_manager)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# E2E: MCP registration → prompt → tool events
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestMcpRegistrationE2E:
|
||||
"""Full flow: session with MCP servers → prompt with tool calls → ACP events."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_with_mcp_servers_registers_tools(self, acp_agent, mock_manager):
|
||||
"""new_session with mcpServers converts them to Hermes config and registers."""
|
||||
servers = [
|
||||
McpServerStdio(
|
||||
name="test-fs",
|
||||
command="/usr/bin/mcp-fs",
|
||||
args=["--root", "/tmp"],
|
||||
env=[EnvVariable(name="DEBUG", value="1")],
|
||||
),
|
||||
McpServerHttp(
|
||||
name="test-api",
|
||||
url="https://api.example.com/mcp",
|
||||
headers=[HttpHeader(name="Authorization", value="Bearer tok123")],
|
||||
),
|
||||
]
|
||||
|
||||
registered_configs = {}
|
||||
|
||||
def mock_register(config_map):
|
||||
registered_configs.update(config_map)
|
||||
return ["mcp_test_fs_read", "mcp_test_fs_write", "mcp_test_api_search"]
|
||||
|
||||
fake_tools = [
|
||||
{"function": {"name": "mcp_test_fs_read"}},
|
||||
{"function": {"name": "mcp_test_fs_write"}},
|
||||
{"function": {"name": "mcp_test_api_search"}},
|
||||
{"function": {"name": "terminal"}},
|
||||
]
|
||||
|
||||
with patch("tools.mcp_tool_discovery.register_mcp_servers", side_effect=mock_register), \
|
||||
patch("model_tools.get_tool_definitions", return_value=fake_tools):
|
||||
resp = await acp_agent.new_session(cwd="/tmp", mcp_servers=servers)
|
||||
|
||||
assert isinstance(resp, NewSessionResponse)
|
||||
state = mock_manager.get_session(resp.session_id)
|
||||
|
||||
# Verify stdio server was converted correctly
|
||||
assert "test-fs" in registered_configs
|
||||
fs_cfg = registered_configs["test-fs"]
|
||||
assert fs_cfg["command"] == "/usr/bin/mcp-fs"
|
||||
assert fs_cfg["args"] == ["--root", "/tmp"]
|
||||
assert fs_cfg["env"] == {"DEBUG": "1"}
|
||||
|
||||
# Verify HTTP server was converted correctly
|
||||
assert "test-api" in registered_configs
|
||||
api_cfg = registered_configs["test-api"]
|
||||
assert api_cfg["url"] == "https://api.example.com/mcp"
|
||||
assert api_cfg["headers"] == {"Authorization": "Bearer tok123"}
|
||||
|
||||
# Verify agent tool surface was refreshed
|
||||
assert state.agent.tools == fake_tools
|
||||
assert state.agent.valid_tool_names == {
|
||||
"mcp_test_fs_read", "mcp_test_fs_write", "mcp_test_api_search", "terminal"
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prompt_with_tool_calls_emits_acp_events(self, acp_agent, mock_manager):
|
||||
"""Prompt → agent fires callbacks → ACP ToolCallStart + ToolCallUpdate events."""
|
||||
resp = await acp_agent.new_session(cwd="/tmp")
|
||||
session_id = resp.session_id
|
||||
state = mock_manager.get_session(session_id)
|
||||
|
||||
# Wire up a mock ACP client connection
|
||||
mock_conn = MagicMock(spec=acp.Client)
|
||||
mock_conn.session_update = AsyncMock()
|
||||
mock_conn.request_permission = AsyncMock()
|
||||
acp_agent._conn = mock_conn
|
||||
|
||||
def mock_run_conversation(user_message, conversation_history=None, task_id=None, **kwargs):
|
||||
"""Simulate an agent turn that calls terminal, gets a result, then responds."""
|
||||
agent = state.agent
|
||||
|
||||
# 1) Agent fires tool_progress_callback (ToolCallStart)
|
||||
if agent.tool_progress_callback:
|
||||
agent.tool_progress_callback(
|
||||
"tool.started", "terminal", "$ echo hello", {"command": "echo hello"}
|
||||
)
|
||||
|
||||
# 2) Agent fires step_callback with tool results (ToolCallUpdate)
|
||||
if agent.step_callback:
|
||||
agent.step_callback(1, [
|
||||
{"name": "terminal", "result": '{"output": "hello\\n", "exit_code": 0}'}
|
||||
])
|
||||
|
||||
return {
|
||||
"final_response": "The command output 'hello'.",
|
||||
"messages": [
|
||||
{"role": "user", "content": user_message},
|
||||
{"role": "assistant", "content": "The command output 'hello'."},
|
||||
],
|
||||
}
|
||||
|
||||
state.agent.run_conversation = mock_run_conversation
|
||||
|
||||
prompt = [TextContentBlock(type="text", text="run echo hello")]
|
||||
resp = await acp_agent.prompt(prompt=prompt, session_id=session_id)
|
||||
|
||||
assert isinstance(resp, PromptResponse)
|
||||
assert resp.stop_reason == "end_turn"
|
||||
|
||||
# Collect all session_update calls
|
||||
updates = []
|
||||
for call in mock_conn.session_update.call_args_list:
|
||||
# session_update(session_id, update) — grab the update
|
||||
update_arg = call[1].get("update") or call[0][1]
|
||||
updates.append(update_arg)
|
||||
|
||||
# Find tool_call (start) and tool_call_update (completion) events
|
||||
starts = [u for u in updates if getattr(u, "session_update", None) == "tool_call"]
|
||||
completions = [u for u in updates if getattr(u, "session_update", None) == "tool_call_update"]
|
||||
|
||||
# Should have at least one ToolCallStart for "terminal"
|
||||
assert len(starts) >= 1, f"Expected ToolCallStart, got updates: {[getattr(u, 'session_update', '?') for u in updates]}"
|
||||
start_event = starts[0]
|
||||
assert isinstance(start_event, ToolCallStart)
|
||||
assert start_event.title.startswith("terminal:")
|
||||
|
||||
# Should have at least one ToolCallUpdate (completion) with rawOutput
|
||||
assert len(completions) >= 1, f"Expected ToolCallUpdate, got updates: {[getattr(u, 'session_update', '?') for u in updates]}"
|
||||
complete_event = completions[0]
|
||||
assert isinstance(complete_event, ToolCallProgress)
|
||||
assert complete_event.status == "completed"
|
||||
# Completion should contain human-readable output rather than forcing raw JSON panes.
|
||||
assert complete_event.content
|
||||
assert "hello" in complete_event.content[0].content.text
|
||||
assert complete_event.raw_output is None
|
||||
|
||||
def test_patch_mode_tool_start_defers_diff_to_edit_approval_prompt(self):
|
||||
update = build_tool_start(
|
||||
"tc-1",
|
||||
"patch",
|
||||
{
|
||||
"mode": "patch",
|
||||
"patch": "*** Begin Patch\n*** Update File: src/app.py\n@@\n-old line\n+new line\n*** Add File: src/new.py\n+hello\n*** End Patch",
|
||||
},
|
||||
)
|
||||
|
||||
assert len(update.content) == 1
|
||||
assert update.content[0].type == "content"
|
||||
assert "Approval prompt shows the diff" in update.content[0].content.text
|
||||
|
||||
|
||||
|
||||
class TestMcpSanitizationE2E:
|
||||
"""Verify server names with special chars work end-to-end."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_slashed_server_name_registers_cleanly(self, acp_agent, mock_manager):
|
||||
"""Server name 'ai.exa/exa' should not crash — tools get sanitized names."""
|
||||
servers = [
|
||||
McpServerHttp(
|
||||
name="ai.exa/exa",
|
||||
url="https://exa.ai/mcp",
|
||||
headers=[],
|
||||
),
|
||||
]
|
||||
|
||||
registered_configs = {}
|
||||
def mock_register(config_map):
|
||||
registered_configs.update(config_map)
|
||||
return ["mcp_ai_exa_exa_search"]
|
||||
|
||||
fake_tools = [{"function": {"name": "mcp_ai_exa_exa_search"}}]
|
||||
|
||||
with patch("tools.mcp_tool_discovery.register_mcp_servers", side_effect=mock_register), \
|
||||
patch("model_tools.get_tool_definitions", return_value=fake_tools):
|
||||
resp = await acp_agent.new_session(cwd="/tmp", mcp_servers=servers)
|
||||
|
||||
state = mock_manager.get_session(resp.session_id)
|
||||
|
||||
# Raw server name preserved as config key
|
||||
assert "ai.exa/exa" in registered_configs
|
||||
# Agent tools refreshed with sanitized name
|
||||
assert "mcp_ai_exa_exa_search" in state.agent.valid_tool_names
|
||||
|
||||
|
||||
class TestSessionLifecycleMcpE2E:
|
||||
"""Verify MCP servers are registered on all session lifecycle methods."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_session_registers_mcp(self, acp_agent, mock_manager):
|
||||
"""load_session re-registers MCP servers (spec says agents may not retain them)."""
|
||||
# Create a session first
|
||||
create_resp = await acp_agent.new_session(cwd="/tmp")
|
||||
sid = create_resp.session_id
|
||||
|
||||
servers = [
|
||||
McpServerStdio(name="srv", command="/bin/test", args=[], env=[]),
|
||||
]
|
||||
|
||||
registered = {}
|
||||
def mock_register(config_map):
|
||||
registered.update(config_map)
|
||||
return []
|
||||
|
||||
state = mock_manager.get_session(sid)
|
||||
state.agent.enabled_toolsets = ["hermes-acp"]
|
||||
state.agent.disabled_toolsets = None
|
||||
state.agent.tools = []
|
||||
state.agent.valid_tool_names = set()
|
||||
|
||||
with patch("tools.mcp_tool_discovery.register_mcp_servers", side_effect=mock_register), \
|
||||
patch("model_tools.get_tool_definitions", return_value=[]):
|
||||
await acp_agent.load_session(cwd="/tmp", session_id=sid, mcp_servers=servers)
|
||||
|
||||
assert "srv" in registered
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resume_session_registers_mcp(self, acp_agent, mock_manager):
|
||||
"""resume_session re-registers MCP servers."""
|
||||
create_resp = await acp_agent.new_session(cwd="/tmp")
|
||||
sid = create_resp.session_id
|
||||
|
||||
servers = [
|
||||
McpServerStdio(name="srv2", command="/bin/test2", args=[], env=[]),
|
||||
]
|
||||
|
||||
registered = {}
|
||||
def mock_register(config_map):
|
||||
registered.update(config_map)
|
||||
return []
|
||||
|
||||
state = mock_manager.get_session(sid)
|
||||
state.agent.enabled_toolsets = ["hermes-acp"]
|
||||
state.agent.disabled_toolsets = None
|
||||
state.agent.tools = []
|
||||
state.agent.valid_tool_names = set()
|
||||
|
||||
with patch("tools.mcp_tool_discovery.register_mcp_servers", side_effect=mock_register), \
|
||||
patch("model_tools.get_tool_definitions", return_value=[]):
|
||||
await acp_agent.resume_session(cwd="/tmp", session_id=sid, mcp_servers=servers)
|
||||
|
||||
assert "srv2" in registered
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fork_session_registers_mcp(self, acp_agent, mock_manager):
|
||||
"""fork_session registers MCP servers on the new forked session."""
|
||||
create_resp = await acp_agent.new_session(cwd="/tmp")
|
||||
sid = create_resp.session_id
|
||||
|
||||
servers = [
|
||||
McpServerHttp(name="api", url="https://api.test/mcp", headers=[]),
|
||||
]
|
||||
|
||||
registered = {}
|
||||
def mock_register(config_map):
|
||||
registered.update(config_map)
|
||||
return []
|
||||
|
||||
# Need to set up the forked session's agent too
|
||||
with patch("tools.mcp_tool_discovery.register_mcp_servers", side_effect=mock_register), \
|
||||
patch("model_tools.get_tool_definitions", return_value=[]):
|
||||
fork_resp = await acp_agent.fork_session(
|
||||
cwd="/tmp", session_id=sid, mcp_servers=servers
|
||||
)
|
||||
|
||||
assert fork_resp.session_id != ""
|
||||
assert "api" in registered
|
||||
@@ -0,0 +1,337 @@
|
||||
"""Tests for named user-defined provider entries in the ACP model selector.
|
||||
|
||||
Named endpoints from the ``providers:`` mapping (and legacy
|
||||
``custom_providers:`` list) are invisible to canonical provider enumeration,
|
||||
so ``_build_model_state`` must append them explicitly for ACP clients to
|
||||
offer them — the TUI ``/model`` picker already renders these entries
|
||||
(#47039 implemented named endpoints for the TUI surface only).
|
||||
"""
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from acp_adapter.model_catalog import _named_custom_provider_catalogs
|
||||
from acp_adapter.server import HermesACPAgent
|
||||
from acp_adapter.session import SessionManager
|
||||
from acp.schema import SessionModelState
|
||||
|
||||
|
||||
MANTLE_URL = "https://bedrock-mantle.us-east-1.api.aws/openai/v1"
|
||||
|
||||
|
||||
def _cfg(providers=None, custom_providers=None):
|
||||
cfg = {}
|
||||
if providers is not None:
|
||||
cfg["providers"] = providers
|
||||
if custom_providers is not None:
|
||||
cfg["custom_providers"] = custom_providers
|
||||
return cfg
|
||||
|
||||
|
||||
class TestNamedCustomProviderCatalogs:
|
||||
|
||||
def test_live_discovery_extends_declared_models(self, monkeypatch):
|
||||
monkeypatch.setenv("SOME_KEY", "k")
|
||||
cfg = _cfg(
|
||||
providers={
|
||||
"relay": {
|
||||
"name": "Relay",
|
||||
"base_url": "https://relay.example/v1",
|
||||
"key_env": "SOME_KEY",
|
||||
"default_model": "model-a",
|
||||
}
|
||||
}
|
||||
)
|
||||
with patch("hermes_cli.config.load_config", return_value=cfg), patch(
|
||||
"hermes_cli.model_switch_providers._fetch_picker_live_models",
|
||||
return_value=["model-a", "model-b"],
|
||||
):
|
||||
catalogs = _named_custom_provider_catalogs()
|
||||
|
||||
assert len(catalogs) == 1
|
||||
slug, label, models = catalogs[0]
|
||||
assert slug == "custom:relay"
|
||||
assert [m for m, _ in models] == ["model-a", "model-b"]
|
||||
|
||||
|
||||
def test_disabled_provider_skipped(self, monkeypatch):
|
||||
monkeypatch.setenv("SOME_KEY", "k")
|
||||
cfg = _cfg(
|
||||
providers={
|
||||
"off": {
|
||||
"name": "Disabled Endpoint",
|
||||
"base_url": "https://off.example/v1",
|
||||
"key_env": "SOME_KEY",
|
||||
"default_model": "m",
|
||||
"enabled": False,
|
||||
}
|
||||
}
|
||||
)
|
||||
with patch("hermes_cli.config.load_config", return_value=cfg), patch(
|
||||
"hermes_cli.model_switch_providers._fetch_picker_live_models", return_value=None
|
||||
):
|
||||
assert _named_custom_provider_catalogs() == []
|
||||
|
||||
def test_no_credential_and_no_declared_models_skipped(self, monkeypatch):
|
||||
monkeypatch.delenv("MISSING_KEY", raising=False)
|
||||
cfg = _cfg(
|
||||
providers={
|
||||
"bare": {
|
||||
"name": "Bare",
|
||||
"base_url": "https://bare.example/v1",
|
||||
"key_env": "MISSING_KEY",
|
||||
}
|
||||
}
|
||||
)
|
||||
with patch("hermes_cli.config.load_config", return_value=cfg), patch(
|
||||
"hermes_cli.model_switch_providers._fetch_picker_live_models", return_value=None
|
||||
):
|
||||
assert _named_custom_provider_catalogs() == []
|
||||
|
||||
def test_legacy_custom_providers_list_included(self, monkeypatch):
|
||||
monkeypatch.setenv("SOME_KEY", "k")
|
||||
cfg = _cfg(
|
||||
custom_providers=[
|
||||
{
|
||||
"name": "Legacy Endpoint",
|
||||
"base_url": "https://legacy.example/v1",
|
||||
"key_env": "SOME_KEY",
|
||||
"model": "legacy-model",
|
||||
}
|
||||
]
|
||||
)
|
||||
with patch("hermes_cli.config.load_config", return_value=cfg), patch(
|
||||
"hermes_cli.model_switch_providers._fetch_picker_live_models", return_value=None
|
||||
):
|
||||
catalogs = _named_custom_provider_catalogs()
|
||||
|
||||
assert catalogs == [
|
||||
("custom:legacy-endpoint", "Legacy Endpoint", [("legacy-model", "")])
|
||||
]
|
||||
def test_no_key_ollama_provider_discovers_native_catalog(self):
|
||||
cfg = _cfg(
|
||||
providers={
|
||||
"custom:ollama": {
|
||||
"name": "Ollama",
|
||||
"base_url": "http://127.0.0.1:11434/v1",
|
||||
}
|
||||
}
|
||||
)
|
||||
with patch("hermes_cli.config.load_config", return_value=cfg), patch(
|
||||
"hermes_cli.models_local.should_use_ollama_native_catalog",
|
||||
return_value=True,
|
||||
), patch(
|
||||
"hermes_cli.model_switch_providers._fetch_picker_live_models",
|
||||
return_value=["qwen3:1.7b"],
|
||||
) as fetch:
|
||||
catalogs = _named_custom_provider_catalogs()
|
||||
|
||||
assert [m for m, _ in catalogs[0][2]] == ["qwen3:1.7b"]
|
||||
fetch.assert_called_once_with(
|
||||
"",
|
||||
"http://127.0.0.1:11434/v1",
|
||||
"custom:ollama",
|
||||
False,
|
||||
headers=None,
|
||||
timeout=1.5,
|
||||
api_mode=None,
|
||||
)
|
||||
def test_legacy_credentialless_ollama_discovers_native_catalog(self):
|
||||
cfg = _cfg(
|
||||
custom_providers=[
|
||||
{
|
||||
"name": "Local Ollama",
|
||||
"base_url": "http://127.0.0.1:11434/v1",
|
||||
}
|
||||
]
|
||||
)
|
||||
with patch("hermes_cli.config.load_config", return_value=cfg), patch(
|
||||
"hermes_cli.models_local.should_use_ollama_native_catalog",
|
||||
return_value=True,
|
||||
), patch(
|
||||
"hermes_cli.model_switch_providers._fetch_picker_live_models",
|
||||
return_value=["qwen3:1.7b"],
|
||||
) as fetch:
|
||||
catalogs = _named_custom_provider_catalogs()
|
||||
|
||||
assert [m for m, _ in catalogs[0][2]] == ["qwen3:1.7b"]
|
||||
fetch.assert_called_once()
|
||||
|
||||
def test_native_empty_catalog_is_authoritative_over_default_model(self):
|
||||
cfg = _cfg(
|
||||
providers={
|
||||
"custom:ollama": {
|
||||
"name": "Ollama",
|
||||
"base_url": "http://127.0.0.1:11434/v1",
|
||||
"default_model": "saved:model",
|
||||
}
|
||||
}
|
||||
)
|
||||
from hermes_cli.model_switch_providers import _NativePickerModelList
|
||||
|
||||
with patch("hermes_cli.config.load_config", return_value=cfg), patch(
|
||||
"hermes_cli.models_local.should_use_ollama_native_catalog",
|
||||
return_value=True,
|
||||
), patch(
|
||||
"hermes_cli.model_switch_providers._fetch_picker_live_models",
|
||||
return_value=_NativePickerModelList(),
|
||||
):
|
||||
assert _named_custom_provider_catalogs() == [
|
||||
("custom:ollama", "Ollama", [])
|
||||
]
|
||||
|
||||
|
||||
class TestModelStateIncludesNamedProviders:
|
||||
@pytest.mark.asyncio
|
||||
async def test_authoritative_empty_named_catalog_does_not_resurrect_current_model(self):
|
||||
manager = SessionManager(
|
||||
agent_factory=lambda: SimpleNamespace(
|
||||
model="saved:model", provider="ollama"
|
||||
)
|
||||
)
|
||||
acp_agent = HermesACPAgent(session_manager=manager)
|
||||
|
||||
with patch(
|
||||
"acp_adapter.model_catalog._named_custom_provider_catalogs",
|
||||
return_value=[("custom:ollama", "Ollama", [])],
|
||||
):
|
||||
resp = await acp_agent.new_session(cwd="/tmp")
|
||||
|
||||
assert isinstance(resp.models, SessionModelState)
|
||||
assert resp.models.current_model_id == ""
|
||||
assert all(
|
||||
not item.model_id.startswith(("ollama:", "custom:ollama:"))
|
||||
for item in resp.models.available_models
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_named_provider_models_appear_in_model_state(self):
|
||||
manager = SessionManager(
|
||||
agent_factory=lambda: SimpleNamespace(
|
||||
model="gpt-5.4", provider="openai-codex"
|
||||
)
|
||||
)
|
||||
acp_agent = HermesACPAgent(session_manager=manager)
|
||||
|
||||
with patch(
|
||||
"acp_adapter.model_catalog._named_custom_provider_catalogs",
|
||||
return_value=[
|
||||
(
|
||||
"custom:bedrock-mantle",
|
||||
"AWS Bedrock Mantle",
|
||||
[("openai.gpt-5.5", "")],
|
||||
)
|
||||
],
|
||||
):
|
||||
resp = await acp_agent.new_session(cwd="/tmp")
|
||||
|
||||
assert isinstance(resp.models, SessionModelState)
|
||||
ids = [m.model_id for m in resp.models.available_models]
|
||||
# Current provider's models come first, named endpoints after.
|
||||
assert ids[0] == "openai-codex:gpt-5.4"
|
||||
assert "custom:bedrock-mantle:openai.gpt-5.5" in ids
|
||||
named = next(
|
||||
m
|
||||
for m in resp.models.available_models
|
||||
if m.model_id == "custom:bedrock-mantle:openai.gpt-5.5"
|
||||
)
|
||||
assert "AWS Bedrock Mantle" in (named.description or "")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_configured_provider_inventory_row_uses_custom_choice_id(self):
|
||||
"""A ``providers:`` row must not expose its raw config key to ACP."""
|
||||
from hermes_cli.models import parse_model_input
|
||||
|
||||
manager = SessionManager(
|
||||
agent_factory=lambda: SimpleNamespace(model="model-a", provider="relay")
|
||||
)
|
||||
acp_agent = HermesACPAgent(session_manager=manager)
|
||||
cfg = {
|
||||
"providers": {
|
||||
"relay": {
|
||||
"name": "Relay",
|
||||
"base_url": "https://relay.example/v1",
|
||||
}
|
||||
}
|
||||
}
|
||||
inventory = {
|
||||
"providers": [{"slug": "relay", "name": "Relay", "is_user_defined": True, "models": ["model-a"]}]
|
||||
}
|
||||
|
||||
with patch("hermes_cli.config.load_config", return_value=cfg), patch(
|
||||
"hermes_cli.inventory.build_models_payload", return_value=inventory
|
||||
), patch(
|
||||
"acp_adapter.model_catalog._named_custom_provider_catalogs",
|
||||
return_value=[("custom:relay", "Relay", [("model-a", "")])],
|
||||
):
|
||||
resp = await acp_agent.new_session(cwd="/tmp")
|
||||
choice_ids = [item.model_id for item in resp.models.available_models]
|
||||
provider, model = parse_model_input(resp.models.current_model_id, "relay")
|
||||
|
||||
assert choice_ids == ["custom:relay:model-a"]
|
||||
assert provider == "custom:relay"
|
||||
assert model == "model-a"
|
||||
|
||||
def test_selector_choice_id_round_trips_through_parse_model_input(self):
|
||||
"""The encoded choice id must resolve back to the named provider."""
|
||||
from hermes_cli.models import parse_model_input
|
||||
|
||||
choice_id = "custom:bedrock-mantle:openai.gpt-5.5"
|
||||
cfg = {
|
||||
"providers": {
|
||||
"bedrock-mantle": {
|
||||
"name": "AWS Bedrock Mantle",
|
||||
"base_url": "https://bedrock.example/v1",
|
||||
}
|
||||
}
|
||||
}
|
||||
with patch("hermes_cli.config.load_config", return_value=cfg):
|
||||
provider, model = parse_model_input(choice_id, "bedrock")
|
||||
assert provider == "custom:bedrock-mantle"
|
||||
assert model == "openai.gpt-5.5"
|
||||
|
||||
def test_selector_choice_id_round_trips_colon_bearing_custom_identity(self):
|
||||
"""Configured provider and model IDs may both contain colons."""
|
||||
from hermes_cli.models import parse_model_input
|
||||
|
||||
cfg = {
|
||||
"providers": {
|
||||
"local-127.0.0.1:11434": {
|
||||
"name": "Local Ollama",
|
||||
"base_url": "http://127.0.0.1:11434/v1",
|
||||
}
|
||||
}
|
||||
}
|
||||
with patch("hermes_cli.config.load_config", return_value=cfg):
|
||||
provider, model = parse_model_input(
|
||||
"custom:local-127.0.0.1:11434:qwen3:1.7b", "custom"
|
||||
)
|
||||
assert provider == "custom:local-127.0.0.1:11434"
|
||||
assert model == "qwen3:1.7b"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_named_entry_shadowing_a_canonical_provider_keeps_the_canonical_session(self):
|
||||
"""``providers.openrouter:`` (a proxy) must not swallow the canonical OpenRouter rows nor
|
||||
relabel a session running on openrouter.ai as ``custom:openrouter`` (that id resolves to
|
||||
the proxy base_url)."""
|
||||
manager = SessionManager(
|
||||
agent_factory=lambda: SimpleNamespace(
|
||||
model="model-c", provider="openrouter", base_url="https://openrouter.ai/api/v1")
|
||||
)
|
||||
acp_agent = HermesACPAgent(session_manager=manager)
|
||||
inventory = {"providers": [
|
||||
{"slug": "openrouter", "name": "OpenRouter", "is_user_defined": False, "models": ["model-c"]},
|
||||
{"slug": "custom:openrouter", "name": "openrouter", "is_user_defined": True,
|
||||
"api_url": "https://or.example/api/v1", "models": ["model-a"]},
|
||||
]}
|
||||
|
||||
with patch("hermes_cli.inventory.build_models_payload", return_value=inventory), patch(
|
||||
"acp_adapter.model_catalog._named_custom_provider_catalogs",
|
||||
return_value=[("custom:openrouter", "openrouter", [("model-a", "")])],
|
||||
):
|
||||
resp = await acp_agent.new_session(cwd="/tmp")
|
||||
|
||||
assert resp.models.current_model_id == "openrouter:model-c"
|
||||
assert [m.model_id for m in resp.models.available_models] == ["openrouter:model-c", "custom:openrouter:model-a"]
|
||||
@@ -0,0 +1,207 @@
|
||||
"""Tests for acp_adapter.permissions."""
|
||||
|
||||
import asyncio
|
||||
import inspect
|
||||
from concurrent.futures import Future
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from acp.schema import (
|
||||
AllowedOutcome,
|
||||
DeniedOutcome,
|
||||
RequestPermissionResponse,
|
||||
)
|
||||
|
||||
from acp_adapter.permissions import make_approval_callback
|
||||
from tools.approval import prompt_dangerous_approval
|
||||
|
||||
|
||||
def _make_response(outcome):
|
||||
return RequestPermissionResponse(outcome=outcome)
|
||||
|
||||
|
||||
def _invoke_callback(
|
||||
outcome,
|
||||
*,
|
||||
allow_permanent=True,
|
||||
allow_session=True,
|
||||
smart_denied=False,
|
||||
timeout=60.0,
|
||||
use_prompt_path=False,
|
||||
):
|
||||
loop = MagicMock(spec=asyncio.AbstractEventLoop)
|
||||
request_permission = AsyncMock(name="request_permission")
|
||||
future = MagicMock(spec=Future)
|
||||
future.result.return_value = _make_response(outcome)
|
||||
|
||||
scheduled = {}
|
||||
|
||||
def _schedule(coro, passed_loop):
|
||||
scheduled["coro"] = coro
|
||||
scheduled["loop"] = passed_loop
|
||||
return future
|
||||
|
||||
with patch("agent.async_utils.asyncio.run_coroutine_threadsafe", side_effect=_schedule):
|
||||
cb = make_approval_callback(request_permission, loop, session_id="s1", timeout=timeout)
|
||||
if use_prompt_path:
|
||||
result = prompt_dangerous_approval(
|
||||
"rm -rf /",
|
||||
"dangerous command",
|
||||
allow_permanent=allow_permanent,
|
||||
allow_session=allow_session,
|
||||
smart_denied=smart_denied,
|
||||
approval_callback=cb,
|
||||
)
|
||||
else:
|
||||
result = cb(
|
||||
"rm -rf /",
|
||||
"dangerous command",
|
||||
allow_permanent=allow_permanent,
|
||||
allow_session=allow_session,
|
||||
smart_denied=smart_denied,
|
||||
)
|
||||
|
||||
scheduled["coro"].close()
|
||||
_, kwargs = request_permission.call_args
|
||||
return result, kwargs, scheduled, future, loop
|
||||
|
||||
|
||||
class TestApprovalBridge:
|
||||
def test_bridge_schedules_request_on_the_given_loop(self):
|
||||
result, kwargs, scheduled, _, loop = _invoke_callback(
|
||||
AllowedOutcome(option_id="allow_once", outcome="selected"),
|
||||
)
|
||||
|
||||
tool_call = kwargs["tool_call"]
|
||||
option_ids = [option.option_id for option in kwargs["options"]]
|
||||
|
||||
assert result == "once"
|
||||
assert scheduled["loop"] is loop
|
||||
assert inspect.iscoroutine(scheduled["coro"])
|
||||
assert kwargs["session_id"] == "s1"
|
||||
assert tool_call.session_update == "tool_call_update"
|
||||
assert tool_call.tool_call_id.startswith("perm-check-")
|
||||
assert tool_call.kind == "execute"
|
||||
assert tool_call.status == "pending"
|
||||
assert "dangerous command" in tool_call.title
|
||||
assert "rm -rf /" in tool_call.title
|
||||
content_text = tool_call.content[0].content.text
|
||||
assert "$ rm -rf /" in content_text
|
||||
assert "dangerous command" in content_text
|
||||
assert tool_call.raw_input == {
|
||||
"command": "rm -rf /",
|
||||
"description": "dangerous command",
|
||||
}
|
||||
assert option_ids == [
|
||||
"allow_once",
|
||||
"allow_session",
|
||||
"allow_always",
|
||||
"deny",
|
||||
"deny_always",
|
||||
]
|
||||
|
||||
def test_session_less_gate_offers_only_once_and_deny(self):
|
||||
"""allow_session=False collapses the editor menu to once/deny.
|
||||
|
||||
Hermes discards any scope broader than one operation for the
|
||||
protected agent-instruction gate, so an editor that renders
|
||||
"Allow for session" would re-prompt on the next write (#81887).
|
||||
"""
|
||||
_, kwargs, _, _, _ = _invoke_callback(
|
||||
AllowedOutcome(option_id="allow_once", outcome="selected"),
|
||||
allow_permanent=False,
|
||||
allow_session=False,
|
||||
)
|
||||
|
||||
assert [option.option_id for option in kwargs["options"]] == ["allow_once", "deny"]
|
||||
|
||||
def test_tool_call_ids_are_unique(self):
|
||||
_, first_kwargs, _, _, _ = _invoke_callback(
|
||||
AllowedOutcome(option_id="allow_once", outcome="selected"),
|
||||
)
|
||||
_, second_kwargs, _, _, _ = _invoke_callback(
|
||||
AllowedOutcome(option_id="allow_once", outcome="selected"),
|
||||
)
|
||||
|
||||
assert first_kwargs["tool_call"].tool_call_id != second_kwargs["tool_call"].tool_call_id
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_allow_always_maps_correctly(self):
|
||||
result, _, _, _, _ = _invoke_callback(
|
||||
AllowedOutcome(option_id="allow_always", outcome="selected"),
|
||||
use_prompt_path=True,
|
||||
)
|
||||
|
||||
assert result == "always"
|
||||
|
||||
|
||||
def test_timeout_returns_timeout_and_cancels_future(self):
|
||||
loop = MagicMock(spec=asyncio.AbstractEventLoop)
|
||||
request_permission = AsyncMock(name="request_permission")
|
||||
future = MagicMock(spec=Future)
|
||||
future.result.side_effect = TimeoutError("timed out")
|
||||
|
||||
scheduled = {}
|
||||
|
||||
def _schedule(coro, passed_loop):
|
||||
scheduled["coro"] = coro
|
||||
scheduled["loop"] = passed_loop
|
||||
return future
|
||||
|
||||
with patch("agent.async_utils.asyncio.run_coroutine_threadsafe", side_effect=_schedule):
|
||||
cb = make_approval_callback(request_permission, loop, session_id="s1", timeout=0.01)
|
||||
result = cb("rm -rf /", "dangerous command")
|
||||
|
||||
scheduled["coro"].close()
|
||||
|
||||
# A no-response expiry is classified as "timeout" (still blocked,
|
||||
# fail-closed) so the agent isn't told the user explicitly refused.
|
||||
assert result == "timeout"
|
||||
assert scheduled["loop"] is loop
|
||||
assert future.cancel.call_count == 1
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Scheduler-failure regression
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
import gc # noqa: E402
|
||||
import warnings # noqa: E402
|
||||
|
||||
|
||||
class TestSchedulerFailure:
|
||||
def test_scheduler_failure_closes_permission_coroutine(self):
|
||||
"""If run_coroutine_threadsafe raises, the coro is closed and we return 'deny'."""
|
||||
loop = MagicMock(spec=asyncio.AbstractEventLoop)
|
||||
created = {"coro": None}
|
||||
|
||||
async def _response_coro(**kwargs):
|
||||
return _make_response(AllowedOutcome(option_id="allow_once", outcome="selected"))
|
||||
|
||||
def _request_permission(**kwargs):
|
||||
created["coro"] = _response_coro(**kwargs)
|
||||
return created["coro"]
|
||||
|
||||
with warnings.catch_warnings(record=True) as caught:
|
||||
warnings.simplefilter("always")
|
||||
with patch(
|
||||
"agent.async_utils.asyncio.run_coroutine_threadsafe",
|
||||
side_effect=RuntimeError("scheduler down"),
|
||||
):
|
||||
cb = make_approval_callback(_request_permission, loop, session_id="s1", timeout=0.01)
|
||||
result = cb("rm -rf /", "dangerous")
|
||||
gc.collect()
|
||||
|
||||
assert result == "deny"
|
||||
assert created["coro"] is not None
|
||||
assert created["coro"].cr_frame is None
|
||||
runtime_warnings = [
|
||||
w for w in caught
|
||||
if issubclass(w.category, RuntimeWarning)
|
||||
and "was never awaited" in str(w.message)
|
||||
and "_response_coro" in str(w.message)
|
||||
]
|
||||
assert runtime_warnings == []
|
||||
@@ -0,0 +1,191 @@
|
||||
"""Tests for acp_adapter.entry._BenignProbeMethodFilter.
|
||||
|
||||
Covers both the isolated filter logic and the full end-to-end path where a
|
||||
client sends a bare JSON-RPC ``ping`` request over stdio and the acp runtime
|
||||
surfaces the resulting ``RequestError`` via ``logging.exception("Background
|
||||
task failed", ...)``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from io import StringIO
|
||||
|
||||
import pytest
|
||||
|
||||
from acp.exceptions import RequestError
|
||||
|
||||
from acp_adapter.entry import _BenignProbeMethodFilter
|
||||
|
||||
|
||||
# -- Unit tests on the filter itself ----------------------------------------
|
||||
|
||||
|
||||
def _make_record(msg: str, exc: BaseException | None) -> logging.LogRecord:
|
||||
record = logging.LogRecord(
|
||||
name="root",
|
||||
level=logging.ERROR,
|
||||
pathname=__file__,
|
||||
lineno=0,
|
||||
msg=msg,
|
||||
args=(),
|
||||
exc_info=(type(exc), exc, exc.__traceback__) if exc else None,
|
||||
)
|
||||
return record
|
||||
|
||||
|
||||
def _bake_tb(exc: BaseException) -> BaseException:
|
||||
try:
|
||||
raise exc
|
||||
except BaseException as e: # noqa: BLE001
|
||||
return e
|
||||
|
||||
|
||||
@pytest.mark.parametrize("method", ["ping", "health", "healthcheck"])
|
||||
def test_filter_suppresses_benign_probe(method: str) -> None:
|
||||
f = _BenignProbeMethodFilter()
|
||||
exc = _bake_tb(RequestError.method_not_found(method))
|
||||
record = _make_record("Background task failed", exc)
|
||||
assert f.filter(record) is False
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_filter_allows_different_message_even_for_ping() -> None:
|
||||
"""Only 'Background task failed' is muted — other messages pass through."""
|
||||
f = _BenignProbeMethodFilter()
|
||||
exc = _bake_tb(RequestError.method_not_found("ping"))
|
||||
record = _make_record("Some other context", exc)
|
||||
assert f.filter(record) is True
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# -- End-to-end: drive a real JSON-RPC `ping` through acp.run_agent ---------
|
||||
|
||||
|
||||
class _FakeAgent:
|
||||
"""Minimal acp.Agent stub — we only need the router to build."""
|
||||
|
||||
async def initialize(self, **kwargs): # noqa: ANN003
|
||||
from acp.schema import AgentCapabilities, InitializeResponse
|
||||
|
||||
return InitializeResponse(protocol_version=1, agent_capabilities=AgentCapabilities())
|
||||
|
||||
async def new_session(self, cwd, mcp_servers=None, **kwargs): # noqa: ANN001, ANN003
|
||||
from acp.schema import NewSessionResponse
|
||||
|
||||
return NewSessionResponse(session_id="test")
|
||||
|
||||
async def prompt(self, session_id, prompt, **kwargs): # noqa: ANN001, ANN003
|
||||
from acp.schema import PromptResponse
|
||||
|
||||
return PromptResponse(stop_reason="end_turn")
|
||||
|
||||
async def cancel(self, session_id, **kwargs): # noqa: ANN001, ANN003
|
||||
pass
|
||||
|
||||
async def authenticate(self, **kwargs): # noqa: ANN003
|
||||
pass
|
||||
|
||||
def on_connect(self, conn): # noqa: ANN001
|
||||
pass
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bare_ping_request_produces_proper_response_and_no_stderr_noise(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
"""A bare ``ping`` must get a JSON-RPC -32601 back AND leave stderr clean
|
||||
when the filter is installed on the handler.
|
||||
"""
|
||||
import acp
|
||||
|
||||
# Attach the filter to a fresh stream handler that mirrors entry._setup_logging.
|
||||
stream = StringIO()
|
||||
handler = logging.StreamHandler(stream)
|
||||
handler.setFormatter(logging.Formatter("%(name)s|%(levelname)s|%(message)s"))
|
||||
handler.addFilter(_BenignProbeMethodFilter())
|
||||
root = logging.getLogger()
|
||||
prior_handlers = root.handlers[:]
|
||||
prior_level = root.level
|
||||
root.handlers = [handler]
|
||||
root.setLevel(logging.INFO)
|
||||
# Also suppress propagation of caplog's default handler interfering with
|
||||
# our stream (caplog still captures via its own propagation hook).
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
|
||||
# Pipe client -> agent
|
||||
client_to_agent_r, client_to_agent_w = os.pipe()
|
||||
# Pipe agent -> client
|
||||
agent_to_client_r, agent_to_client_w = os.pipe()
|
||||
|
||||
in_read_file = os.fdopen(client_to_agent_r, "rb", buffering=0)
|
||||
in_write_file = os.fdopen(client_to_agent_w, "wb", buffering=0)
|
||||
out_read_file = os.fdopen(agent_to_client_r, "rb", buffering=0)
|
||||
out_write_file = os.fdopen(agent_to_client_w, "wb", buffering=0)
|
||||
|
||||
# Agent reads its input from this StreamReader:
|
||||
agent_input = asyncio.StreamReader(limit=1024 * 1024, loop=loop)
|
||||
agent_input_proto = asyncio.StreamReaderProtocol(agent_input, loop=loop)
|
||||
await loop.connect_read_pipe(lambda: agent_input_proto, in_read_file)
|
||||
|
||||
# Agent writes its output via this StreamWriter:
|
||||
out_transport, out_protocol = await loop.connect_write_pipe(
|
||||
asyncio.streams.FlowControlMixin, out_write_file
|
||||
)
|
||||
agent_output = asyncio.StreamWriter(out_transport, out_protocol, None, loop)
|
||||
|
||||
# Test harness reads agent output via this StreamReader:
|
||||
client_input = asyncio.StreamReader(limit=1024 * 1024, loop=loop)
|
||||
client_input_proto = asyncio.StreamReaderProtocol(client_input, loop=loop)
|
||||
await loop.connect_read_pipe(lambda: client_input_proto, out_read_file)
|
||||
|
||||
agent_task = asyncio.create_task(
|
||||
acp.run_agent(
|
||||
_FakeAgent(),
|
||||
input_stream=agent_output,
|
||||
output_stream=agent_input,
|
||||
use_unstable_protocol=True,
|
||||
)
|
||||
)
|
||||
|
||||
# Send a bare `ping`
|
||||
request = {"jsonrpc": "2.0", "id": 1, "method": "ping", "params": {}}
|
||||
in_write_file.write((json.dumps(request) + "\n").encode())
|
||||
in_write_file.flush()
|
||||
|
||||
response_line = await asyncio.wait_for(client_input.readline(), timeout=5.0)
|
||||
# Give the supervisor task a tick to fire (filter should eat it)
|
||||
await asyncio.sleep(0.2)
|
||||
|
||||
response = json.loads(response_line.decode())
|
||||
assert response["error"]["code"] == -32601, response
|
||||
assert response["error"]["data"] == {"method": "ping"}, response
|
||||
|
||||
logs = stream.getvalue()
|
||||
assert "Background task failed" not in logs, (
|
||||
f"ping noise leaked to stderr:\n{logs}"
|
||||
)
|
||||
|
||||
# Clean shutdown
|
||||
in_write_file.close()
|
||||
try:
|
||||
await asyncio.wait_for(agent_task, timeout=2.0)
|
||||
except (asyncio.TimeoutError, Exception):
|
||||
agent_task.cancel()
|
||||
try:
|
||||
await agent_task
|
||||
except BaseException: # noqa: BLE001
|
||||
pass
|
||||
finally:
|
||||
root.handlers = prior_handlers
|
||||
root.setLevel(prior_level)
|
||||
@@ -0,0 +1,748 @@
|
||||
"""Tests for acp_adapter.server — HermesACPAgent ACP server."""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import acp
|
||||
from acp.agent.router import build_agent_router
|
||||
from acp.schema import (
|
||||
AgentCapabilities,
|
||||
AgentMessageChunk,
|
||||
AgentPlanUpdate,
|
||||
AgentThoughtChunk,
|
||||
AuthenticateResponse,
|
||||
AvailableCommandsUpdate,
|
||||
Implementation,
|
||||
InitializeResponse,
|
||||
LoadSessionResponse,
|
||||
NewSessionResponse,
|
||||
PromptResponse,
|
||||
ResumeSessionResponse,
|
||||
SessionModelState,
|
||||
SessionModeState,
|
||||
SetSessionConfigOptionResponse,
|
||||
SetSessionModelResponse,
|
||||
SetSessionModeResponse,
|
||||
SessionInfo,
|
||||
SessionInfoUpdate,
|
||||
TextContentBlock,
|
||||
ToolCallProgress,
|
||||
ToolCallStart,
|
||||
UsageUpdate,
|
||||
UserMessageChunk,
|
||||
)
|
||||
from acp_adapter.auth import TERMINAL_SETUP_AUTH_METHOD_ID
|
||||
from acp_adapter.model_catalog import ACP_MAX_MODELS_PER_PROVIDER
|
||||
from acp_adapter.server import (
|
||||
HermesACPAgent,
|
||||
HERMES_VERSION,
|
||||
)
|
||||
from acp_adapter.session import SessionManager
|
||||
from hermes_state import SessionDB
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def mock_manager():
|
||||
"""SessionManager with a mock agent factory."""
|
||||
return SessionManager(agent_factory=lambda: MagicMock(name="MockAIAgent"))
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def agent(mock_manager):
|
||||
"""HermesACPAgent backed by a mock session manager."""
|
||||
return HermesACPAgent(session_manager=mock_manager)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_session_exposes_edit_approvals_as_modes_not_config_options(agent):
|
||||
resp = await agent.new_session(cwd="/tmp")
|
||||
|
||||
assert resp.config_options is None
|
||||
assert isinstance(resp.modes, SessionModeState)
|
||||
assert resp.modes.current_mode_id == "default"
|
||||
assert [(mode.id, mode.name) for mode in resp.modes.available_modes] == [
|
||||
("default", "Default"),
|
||||
("accept_edits", "Accept Edits"),
|
||||
("dont_ask", "Don't Ask"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_set_config_option_persists_edit_approval_policy_without_advertising_config(agent):
|
||||
resp = await agent.new_session(cwd="/tmp")
|
||||
update = await agent.set_config_option(
|
||||
"edit_approval_policy",
|
||||
resp.session_id,
|
||||
"workspace_session",
|
||||
)
|
||||
state = agent.session_manager.get_session(resp.session_id)
|
||||
|
||||
assert isinstance(update, SetSessionConfigOptionResponse)
|
||||
assert update.config_options == []
|
||||
assert getattr(state, "mode", None) == "accept_edits"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# initialize
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestInitialize:
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_returns_correct_protocol_version(self, agent):
|
||||
resp = await agent.initialize(protocol_version=1)
|
||||
assert isinstance(resp, InitializeResponse)
|
||||
assert resp.protocol_version == acp.PROTOCOL_VERSION
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_advertises_provider_and_terminal_auth_methods(self, agent, monkeypatch):
|
||||
monkeypatch.setattr("acp_adapter.auth.detect_provider", lambda: "openrouter")
|
||||
monkeypatch.setattr("acp_adapter.server.detect_provider", lambda: "openrouter")
|
||||
|
||||
resp = await agent.initialize(protocol_version=1)
|
||||
payloads = [method.model_dump(by_alias=True, exclude_none=True) for method in resp.auth_methods]
|
||||
|
||||
assert payloads[0]["id"] == "openrouter"
|
||||
assert payloads[0]["name"] == "openrouter runtime credentials"
|
||||
terminal = next(payload for payload in payloads if payload["id"] == TERMINAL_SETUP_AUTH_METHOD_ID)
|
||||
assert terminal["type"] == "terminal"
|
||||
assert terminal["args"] == ["--setup"]
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# authenticate
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestAuthenticate:
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_with_matching_method_id(self, agent, monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"acp_adapter.server.detect_provider",
|
||||
lambda: "openrouter",
|
||||
)
|
||||
resp = await agent.authenticate(method_id="openrouter")
|
||||
assert isinstance(resp, AuthenticateResponse)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_is_case_insensitive(self, agent, monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"acp_adapter.server.detect_provider",
|
||||
lambda: "openrouter",
|
||||
)
|
||||
resp = await agent.authenticate(method_id="OpenRouter")
|
||||
assert isinstance(resp, AuthenticateResponse)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_rejects_mismatched_method_id(self, agent, monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"acp_adapter.server.detect_provider",
|
||||
lambda: "openrouter",
|
||||
)
|
||||
resp = await agent.authenticate(method_id="totally-invalid-method")
|
||||
assert resp is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_without_provider(self, agent, monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"acp_adapter.server.detect_provider",
|
||||
lambda: None,
|
||||
)
|
||||
resp = await agent.authenticate(method_id="openrouter")
|
||||
assert resp is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_accepts_terminal_setup_after_provider_configured(self, agent, monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"acp_adapter.server.detect_provider",
|
||||
lambda: "openrouter",
|
||||
)
|
||||
resp = await agent.authenticate(method_id=TERMINAL_SETUP_AUTH_METHOD_ID)
|
||||
assert isinstance(resp, AuthenticateResponse)
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# new_session / cancel / load / resume
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSessionOps:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_session_returns_authenticated_cross_provider_model_state(self):
|
||||
manager = SessionManager(
|
||||
agent_factory=lambda: SimpleNamespace(
|
||||
model="gpt-5.4",
|
||||
provider="openai-codex",
|
||||
base_url="https://api.openai.com/v1",
|
||||
)
|
||||
)
|
||||
acp_agent = HermesACPAgent(session_manager=manager)
|
||||
picker_context = MagicMock()
|
||||
picker_context.with_overrides.return_value = picker_context
|
||||
payload = {
|
||||
"providers": [
|
||||
{
|
||||
"slug": "anthropic",
|
||||
"name": "Anthropic",
|
||||
"models": ["claude-sonnet-4-6", "claude-sonnet-4-6"],
|
||||
},
|
||||
{
|
||||
"slug": "openai-codex",
|
||||
"name": "OpenAI Codex",
|
||||
"models": [
|
||||
{"id": "gpt-5.4"},
|
||||
"gpt-5.4-mini",
|
||||
],
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
with (
|
||||
patch("hermes_cli.inventory.load_picker_context", return_value=picker_context),
|
||||
patch("hermes_cli.inventory.build_models_payload", return_value=payload) as build_payload,
|
||||
):
|
||||
resp = await acp_agent.new_session(cwd="/tmp")
|
||||
|
||||
assert isinstance(resp.models, SessionModelState)
|
||||
assert resp.models.current_model_id == "openai-codex:gpt-5.4"
|
||||
assert [model.model_id for model in resp.models.available_models] == [
|
||||
"anthropic:claude-sonnet-4-6",
|
||||
"openai-codex:gpt-5.4",
|
||||
"openai-codex:gpt-5.4-mini",
|
||||
]
|
||||
assert [model.name for model in resp.models.available_models] == [
|
||||
"Anthropic · claude-sonnet-4-6",
|
||||
"OpenAI Codex · gpt-5.4",
|
||||
"OpenAI Codex · gpt-5.4-mini",
|
||||
]
|
||||
assert resp.models.available_models[1].description is not None
|
||||
assert "current" in resp.models.available_models[1].description
|
||||
picker_context.with_overrides.assert_called_once_with(
|
||||
current_provider="openai-codex",
|
||||
current_model="gpt-5.4",
|
||||
current_base_url="https://api.openai.com/v1",
|
||||
)
|
||||
build_payload.assert_called_once_with(
|
||||
picker_context,
|
||||
explicit_only=True,
|
||||
include_unconfigured=False,
|
||||
picker_hints=False,
|
||||
canonical_order=True,
|
||||
pricing=False,
|
||||
capabilities=False,
|
||||
refresh=False,
|
||||
probe_custom_providers=False,
|
||||
probe_current_custom_provider=False,
|
||||
max_models=ACP_MAX_MODELS_PER_PROVIDER,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_available_commands_include_help(self, agent):
|
||||
help_cmd = next(
|
||||
(cmd for cmd in agent._available_commands() if cmd.name == "help"),
|
||||
None,
|
||||
)
|
||||
|
||||
assert help_cmd is not None
|
||||
assert help_cmd.description == "List available commands"
|
||||
assert help_cmd.input is None
|
||||
|
||||
|
||||
def test_build_usage_update_for_zed_context_indicator(self, agent, mock_manager):
|
||||
state = mock_manager.create_session(cwd="/tmp")
|
||||
state.history = [{"role": "user", "content": "hello"}]
|
||||
state.agent.context_compressor = MagicMock(context_length=100_000)
|
||||
state.agent._cached_system_prompt = "system"
|
||||
state.agent.tools = [{"type": "function", "function": {"name": "demo"}}]
|
||||
|
||||
with patch(
|
||||
"agent.model_metadata.estimate_request_tokens_rough",
|
||||
return_value=25_000,
|
||||
):
|
||||
update = agent._build_usage_update(state)
|
||||
|
||||
assert isinstance(update, UsageUpdate)
|
||||
assert update.session_update == "usage_update"
|
||||
assert update.size == 100_000
|
||||
assert update.used == 25_000
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_session_not_found_returns_none(self, agent):
|
||||
resp = await agent.load_session(cwd="/tmp", session_id="bogus")
|
||||
assert resp is None
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resume_session_replays_persisted_history_to_client(self, agent):
|
||||
mock_conn = MagicMock(spec=acp.Client)
|
||||
mock_conn.session_update = AsyncMock()
|
||||
agent._conn = mock_conn
|
||||
|
||||
new_resp = await agent.new_session(cwd="/tmp")
|
||||
state = agent.session_manager.get_session(new_resp.session_id)
|
||||
state.history = [{"role": "user", "content": "So tell me the current state"}]
|
||||
|
||||
mock_conn.session_update.reset_mock()
|
||||
resp = await agent.resume_session(cwd="/tmp", session_id=new_resp.session_id)
|
||||
await asyncio.sleep(0)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert isinstance(resp, ResumeSessionResponse)
|
||||
updates = [call.kwargs["update"] for call in mock_conn.session_update.await_args_list]
|
||||
assert any(
|
||||
isinstance(update, UserMessageChunk)
|
||||
and update.content.text == "So tell me the current state"
|
||||
for update in updates
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# list / fork
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestListAndFork:
|
||||
@pytest.mark.asyncio
|
||||
async def test_fork_session(self, agent):
|
||||
new_resp = await agent.new_session(cwd="/original")
|
||||
fork_resp = await agent.fork_session(cwd="/forked", session_id=new_resp.session_id)
|
||||
assert fork_resp.session_id
|
||||
assert fork_resp.session_id != new_resp.session_id
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_sessions_includes_title_and_updated_at(self, agent):
|
||||
with patch.object(
|
||||
agent.session_manager,
|
||||
"list_sessions",
|
||||
return_value=[
|
||||
{
|
||||
"session_id": "session-1",
|
||||
"cwd": "/tmp/project",
|
||||
"title": "Fix Zed session history",
|
||||
"updated_at": 123.0,
|
||||
}
|
||||
],
|
||||
):
|
||||
resp = await agent.list_sessions(cwd="/tmp/project")
|
||||
|
||||
assert isinstance(resp.sessions[0], SessionInfo)
|
||||
assert resp.sessions[0].title == "Fix Zed session history"
|
||||
assert resp.sessions[0].updated_at == "123.0"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# session configuration / model routing
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSessionConfiguration:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_accepts_stable_session_config_methods(self, agent):
|
||||
new_resp = await agent.new_session(cwd="/tmp")
|
||||
router = build_agent_router(agent)
|
||||
|
||||
mode_result = await router(
|
||||
"session/set_mode",
|
||||
{"modeId": "accept_edits", "sessionId": new_resp.session_id},
|
||||
False,
|
||||
)
|
||||
config_result = await router(
|
||||
"session/set_config_option",
|
||||
{
|
||||
"configId": "approval_mode",
|
||||
"sessionId": new_resp.session_id,
|
||||
"value": "auto",
|
||||
},
|
||||
False,
|
||||
)
|
||||
|
||||
assert mode_result == {}
|
||||
assert config_result["configOptions"] == []
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# prompt
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPrompt:
|
||||
@pytest.mark.asyncio
|
||||
async def test_prompt_returns_refusal_for_unknown_session(self, agent):
|
||||
prompt = [TextContentBlock(type="text", text="hello")]
|
||||
resp = await agent.prompt(prompt=prompt, session_id="nonexistent")
|
||||
assert isinstance(resp, PromptResponse)
|
||||
assert resp.stop_reason == "refusal"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prompt_binds_session_id_into_subprocess_env(self, agent, mock_manager):
|
||||
"""The ACP prompt path must bridge the session id into child subprocesses.
|
||||
|
||||
Regression: ``set_session_vars`` was called with ``session_key`` only,
|
||||
leaving the ``HERMES_SESSION_ID`` ContextVar bound to the explicit ""
|
||||
default. Once the session-context machinery is engaged, that empty value
|
||||
is authoritative — so ``_make_run_env`` handed child subprocesses an
|
||||
empty ``HERMES_SESSION_ID`` instead of the session's own id.
|
||||
"""
|
||||
from tools.environments.local import _make_run_env
|
||||
|
||||
resp = await agent.new_session(cwd=".")
|
||||
state = mock_manager.get_session(resp.session_id)
|
||||
|
||||
captured: dict[str, str | None] = {}
|
||||
|
||||
def _run(*args, **kwargs):
|
||||
# Runs inside the session context copy set up by prompt().
|
||||
captured["child"] = _make_run_env({}).get("HERMES_SESSION_ID")
|
||||
return {"final_response": "ok", "messages": []}
|
||||
|
||||
state.agent.run_conversation = _run
|
||||
state.agent.model = "test-model"
|
||||
state.agent.provider = "openrouter"
|
||||
|
||||
mock_conn = MagicMock(spec=acp.Client)
|
||||
mock_conn.session_update = AsyncMock()
|
||||
agent._conn = mock_conn
|
||||
|
||||
await agent.prompt(
|
||||
prompt=[TextContentBlock(type="text", text="hi")],
|
||||
session_id=resp.session_id,
|
||||
)
|
||||
|
||||
assert captured.get("child") == resp.session_id
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_messages_list_replaces_stale_history(self, agent, mock_manager):
|
||||
"""``run_conversation`` returning ``messages=[]`` clears the ACP transcript instead of
|
||||
leaving the previous turn's history in place (#10844)."""
|
||||
resp = await agent.new_session(cwd=".")
|
||||
state = mock_manager.get_session(resp.session_id)
|
||||
state.history = [{"role": "user", "content": "old"}]
|
||||
state.agent.run_conversation = MagicMock(return_value={"final_response": "done", "messages": []})
|
||||
state.agent.model = "test-model"
|
||||
state.agent.provider = "openrouter"
|
||||
mock_conn = MagicMock(spec=acp.Client)
|
||||
mock_conn.session_update = AsyncMock()
|
||||
agent._conn = mock_conn
|
||||
|
||||
await agent.prompt(prompt=[TextContentBlock(type="text", text="hi")], session_id=resp.session_id)
|
||||
|
||||
assert state.history == []
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# on_connect
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestOnConnect:
|
||||
def test_on_connect_stores_client(self, agent):
|
||||
mock_conn = MagicMock(spec=acp.Client)
|
||||
agent.on_connect(mock_conn)
|
||||
assert agent._conn is mock_conn
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Slash commands
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSlashCommands:
|
||||
"""Test slash command dispatch in the ACP adapter."""
|
||||
|
||||
def _make_state(self, mock_manager):
|
||||
state = mock_manager.create_session(cwd="/tmp")
|
||||
state.agent.model = "test-model"
|
||||
state.agent.provider = "openrouter"
|
||||
state.model = "test-model"
|
||||
return state
|
||||
|
||||
def test_help_lists_commands(self, agent, mock_manager):
|
||||
state = self._make_state(mock_manager)
|
||||
result = agent._handle_slash_command("/help", state)
|
||||
assert result is not None
|
||||
assert "/help" in result
|
||||
assert "/model" in result
|
||||
assert "/tools" in result
|
||||
assert "/reset" in result
|
||||
|
||||
def test_model_shows_current(self, agent, mock_manager):
|
||||
state = self._make_state(mock_manager)
|
||||
result = agent._handle_slash_command("/model", state)
|
||||
assert "test-model" in result
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_reset_clears_history(self, agent, mock_manager):
|
||||
state = self._make_state(mock_manager)
|
||||
state.history = [{"role": "user", "content": "hello"}]
|
||||
result = agent._handle_slash_command("/reset", state)
|
||||
assert "cleared" in result.lower()
|
||||
assert len(state.history) == 0
|
||||
|
||||
|
||||
|
||||
|
||||
def test_compact_compresses_context(self, agent, mock_manager):
|
||||
state = self._make_state(mock_manager)
|
||||
state.history = [
|
||||
{"role": "user", "content": "one"},
|
||||
{"role": "assistant", "content": "two"},
|
||||
{"role": "user", "content": "three"},
|
||||
{"role": "assistant", "content": "four"},
|
||||
]
|
||||
state.agent.compression_enabled = True
|
||||
state.agent._cached_system_prompt = "system"
|
||||
state.agent.tools = None
|
||||
original_session_db = object()
|
||||
state.agent._session_db = original_session_db
|
||||
|
||||
def _compress_context(messages, system_prompt, *, approx_tokens, task_id, force, **kwargs):
|
||||
assert state.agent._session_db is None
|
||||
assert messages == state.history
|
||||
assert system_prompt == "system"
|
||||
assert approx_tokens == 40
|
||||
assert task_id == state.session_id
|
||||
assert force is True
|
||||
return [{"role": "user", "content": "summary"}], "new-system"
|
||||
|
||||
state.agent._compress_context = MagicMock(side_effect=_compress_context)
|
||||
|
||||
with (
|
||||
patch.object(agent.session_manager, "save_session") as mock_save,
|
||||
patch(
|
||||
"agent.model_metadata.estimate_request_tokens_rough",
|
||||
side_effect=[40, 12],
|
||||
),
|
||||
):
|
||||
result = agent._handle_slash_command("/compress", state)
|
||||
|
||||
assert "Context compressed: 4 -> 1 messages" in result
|
||||
assert "~40 -> ~12 tokens" in result
|
||||
assert state.history == [{"role": "user", "content": "summary"}]
|
||||
assert state.agent._session_db is original_session_db
|
||||
state.agent._compress_context.assert_called_once_with(
|
||||
[
|
||||
{"role": "user", "content": "one"},
|
||||
{"role": "assistant", "content": "two"},
|
||||
{"role": "user", "content": "three"},
|
||||
{"role": "assistant", "content": "four"},
|
||||
],
|
||||
"system",
|
||||
approx_tokens=40,
|
||||
focus_topic=None,
|
||||
force=True,
|
||||
defer_context_engine_notification=True,
|
||||
task_id=state.session_id,
|
||||
)
|
||||
mock_save.assert_called_once_with(state.session_id)
|
||||
|
||||
|
||||
def test_unknown_command_returns_none(self, agent, mock_manager):
|
||||
state = self._make_state(mock_manager)
|
||||
result = agent._handle_slash_command("/nonexistent", state)
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_slash_handler_cwd_pin_does_not_leak(self, agent, mock_manager, tmp_path):
|
||||
"""The pin is scoped to the handler's own context copy.
|
||||
|
||||
Concurrent ACP sessions share the event loop, so a handler that pinned
|
||||
the ambient context would leave its workspace bound for whatever runs
|
||||
next. Asserting the ambient value is unchanged after dispatch keeps the
|
||||
fix from trading one cross-session leak for another.
|
||||
"""
|
||||
from agent.runtime_cwd import resolve_agent_cwd
|
||||
|
||||
workspace = tmp_path / "project"
|
||||
workspace.mkdir()
|
||||
state = mock_manager.create_session(cwd=str(workspace))
|
||||
state.cwd = str(workspace)
|
||||
state.agent.model = "test-model"
|
||||
state.agent.provider = "openrouter"
|
||||
|
||||
before = str(resolve_agent_cwd())
|
||||
agent._handle_slash_command("/help", state)
|
||||
assert str(resolve_agent_cwd()) == before
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _register_session_mcp_servers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRegisterSessionMcpServers:
|
||||
"""Tests for ACP MCP server registration in session lifecycle."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_noop_when_no_servers(self, agent, mock_manager):
|
||||
"""No-op when mcp_servers is None or empty."""
|
||||
state = mock_manager.create_session(cwd="/tmp")
|
||||
# Should not raise
|
||||
await agent._register_session_mcp_servers(state, None)
|
||||
await agent._register_session_mcp_servers(state, [])
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_registers_stdio_servers(self, agent, mock_manager):
|
||||
"""McpServerStdio servers are converted and passed to register_mcp_servers."""
|
||||
from acp.schema import McpServerStdio, EnvVariable
|
||||
|
||||
state = mock_manager.create_session(cwd="/tmp")
|
||||
# Give the mock agent the attributes _register_session_mcp_servers reads
|
||||
state.agent.enabled_toolsets = ["hermes-acp"]
|
||||
state.agent.disabled_toolsets = None
|
||||
state.agent.tools = []
|
||||
state.agent.valid_tool_names = set()
|
||||
|
||||
server = McpServerStdio(
|
||||
name="test-server",
|
||||
command="/usr/bin/test",
|
||||
args=["--flag"],
|
||||
env=[EnvVariable(name="KEY", value="val")],
|
||||
)
|
||||
|
||||
registered_config = {}
|
||||
def capture_register(config_map):
|
||||
registered_config.update(config_map)
|
||||
return ["mcp_test_server_tool1"]
|
||||
|
||||
with patch("tools.mcp_tool_discovery.register_mcp_servers", side_effect=capture_register), \
|
||||
patch("model_tools.get_tool_definitions", return_value=[]):
|
||||
await agent._register_session_mcp_servers(state, [server])
|
||||
|
||||
assert "test-server" in registered_config
|
||||
cfg = registered_config["test-server"]
|
||||
assert cfg["command"] == "/usr/bin/test"
|
||||
assert cfg["args"] == ["--flag"]
|
||||
assert cfg["env"] == {"KEY": "val"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refreshes_agent_tool_surface(self, agent, mock_manager):
|
||||
"""After MCP registration, agent.tools and valid_tool_names are refreshed."""
|
||||
from acp.schema import McpServerStdio
|
||||
|
||||
state = mock_manager.create_session(cwd="/tmp")
|
||||
state.agent.enabled_toolsets = ["hermes-acp"]
|
||||
state.agent.disabled_toolsets = None
|
||||
state.agent.tools = []
|
||||
state.agent.valid_tool_names = set()
|
||||
state.agent._cached_system_prompt = "old prompt"
|
||||
state.agent._memory_manager = SimpleNamespace(
|
||||
get_all_tool_schemas=lambda: [
|
||||
{"name": "hindsight_recall", "description": "Recall", "parameters": {}}
|
||||
]
|
||||
)
|
||||
|
||||
server = McpServerStdio(
|
||||
name="srv",
|
||||
command="/bin/test",
|
||||
args=[],
|
||||
env=[],
|
||||
)
|
||||
|
||||
fake_tools = [
|
||||
{"function": {"name": "mcp_srv_search"}},
|
||||
{"function": {"name": "memory"}},
|
||||
{"function": {"name": "terminal"}},
|
||||
]
|
||||
|
||||
with patch("tools.mcp_tool_discovery.register_mcp_servers", return_value=["mcp_srv_search"]), \
|
||||
patch("model_tools.get_tool_definitions", return_value=fake_tools) as mock_defs:
|
||||
await agent._register_session_mcp_servers(state, [server])
|
||||
|
||||
mock_defs.assert_called_once_with(
|
||||
enabled_toolsets=["hermes-acp", "mcp-srv"],
|
||||
disabled_toolsets=None,
|
||||
quiet_mode=True,
|
||||
)
|
||||
assert state.agent.enabled_toolsets == ["hermes-acp", "mcp-srv"]
|
||||
assert state.agent.tools is fake_tools
|
||||
assert state.agent.tools[-1] == {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "hindsight_recall",
|
||||
"description": "Recall",
|
||||
"parameters": {},
|
||||
},
|
||||
}
|
||||
assert state.agent.valid_tool_names == {
|
||||
"hindsight_recall",
|
||||
"memory",
|
||||
"mcp_srv_search",
|
||||
"terminal",
|
||||
}
|
||||
# _invalidate_system_prompt should have been called
|
||||
state.agent._invalidate_system_prompt.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_failure_logs_warning(self, agent, mock_manager):
|
||||
"""If register_mcp_servers raises, warning is logged but no crash."""
|
||||
from acp.schema import McpServerStdio
|
||||
|
||||
state = mock_manager.create_session(cwd="/tmp")
|
||||
server = McpServerStdio(
|
||||
name="bad",
|
||||
command="/nonexistent",
|
||||
args=[],
|
||||
env=[],
|
||||
)
|
||||
|
||||
with patch("tools.mcp_tool_discovery.register_mcp_servers", side_effect=RuntimeError("boom")):
|
||||
# Should not raise
|
||||
await agent._register_session_mcp_servers(state, [server])
|
||||
@@ -0,0 +1,385 @@
|
||||
"""Tests for acp_adapter.session — SessionManager and SessionState."""
|
||||
|
||||
import contextlib
|
||||
import io
|
||||
import json
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
import pytest
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from acp_adapter import session as acp_session
|
||||
from acp_adapter.session import SessionManager, SessionState
|
||||
from hermes_state import SessionDB
|
||||
|
||||
|
||||
def _mock_agent():
|
||||
return MagicMock(name="MockAIAgent")
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def manager():
|
||||
"""SessionManager with a mock agent factory (avoids needing API keys)."""
|
||||
return SessionManager(agent_factory=_mock_agent)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# create / get
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCreateSession:
|
||||
def test_create_session_returns_state(self, manager):
|
||||
state = manager.create_session(cwd="/tmp/work")
|
||||
assert isinstance(state, SessionState)
|
||||
assert state.cwd == "/tmp/work"
|
||||
assert state.session_id
|
||||
assert state.history == []
|
||||
assert state.agent is not None
|
||||
|
||||
|
||||
|
||||
def test_register_task_cwd_translates_windows_drive_for_wsl_tools(self, monkeypatch):
|
||||
captured = {}
|
||||
|
||||
def fake_register_task_env_overrides(task_id, overrides):
|
||||
captured["task_id"] = task_id
|
||||
captured["overrides"] = overrides
|
||||
|
||||
monkeypatch.setattr("hermes_constants._wsl_detected", True)
|
||||
monkeypatch.setattr(
|
||||
"tools.terminal_tool.register_task_env_overrides",
|
||||
fake_register_task_env_overrides,
|
||||
)
|
||||
|
||||
acp_session._register_task_cwd("session-1", r"E:\Projects\AI\paperclip")
|
||||
|
||||
assert captured == {
|
||||
"task_id": "session-1",
|
||||
"overrides": {"cwd": "/mnt/e/Projects/AI/paperclip"},
|
||||
}
|
||||
|
||||
|
||||
def test_get_session(self, manager):
|
||||
state = manager.create_session()
|
||||
fetched = manager.get_session(state.session_id)
|
||||
assert fetched is state
|
||||
|
||||
|
||||
def test_make_agent_uses_session_cwd_during_init_and_stamps_runtime(
|
||||
self, monkeypatch, tmp_path
|
||||
):
|
||||
workspace = tmp_path / "workspace"
|
||||
workspace.mkdir()
|
||||
observed = {}
|
||||
|
||||
class FakeAgent:
|
||||
model = "fake-model"
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
observed["cwd"] = kwargs.get("cwd")
|
||||
|
||||
monkeypatch.setattr("run_agent.AIAgent", FakeAgent)
|
||||
monkeypatch.setattr(
|
||||
"acp_adapter.session.load_config",
|
||||
lambda: {
|
||||
"model": {
|
||||
"default": "fake-model",
|
||||
"provider": "fake-provider",
|
||||
},
|
||||
"mcp_servers": {},
|
||||
},
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"hermes_cli.config.load_config",
|
||||
lambda: {
|
||||
"model": {
|
||||
"default": "fake-model",
|
||||
"provider": "fake-provider",
|
||||
},
|
||||
"mcp_servers": {},
|
||||
},
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"hermes_cli.runtime_provider.resolve_runtime_provider",
|
||||
lambda requested=None: {
|
||||
"provider": requested,
|
||||
"api_mode": "codex_app_server",
|
||||
"base_url": "https://example.invalid",
|
||||
"api_key": "test-key",
|
||||
},
|
||||
)
|
||||
monkeypatch.setattr("acp_adapter.session._register_task_cwd", lambda task_id, cwd: None)
|
||||
|
||||
state = SessionManager(db=None).create_session(cwd=str(workspace))
|
||||
|
||||
assert observed["cwd"] == str(workspace)
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# WSL cwd translation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestWslCwdTranslation:
|
||||
def test_translate_acp_cwd_converts_windows_drive_path_when_wsl(self, monkeypatch):
|
||||
monkeypatch.setattr("hermes_constants._wsl_detected", True)
|
||||
|
||||
assert acp_session._translate_acp_cwd(r"E:\Projects\AI\paperclip") == "/mnt/e/Projects/AI/paperclip"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_fork_session_stores_translated_cwd_on_wsl(self, manager, monkeypatch):
|
||||
monkeypatch.setattr("hermes_constants._wsl_detected", True)
|
||||
original = manager.create_session(cwd="/tmp/base")
|
||||
|
||||
forked = manager.fork_session(original.session_id, cwd=r"D:\work\project")
|
||||
|
||||
assert forked is not None
|
||||
assert forked.cwd == "/mnt/d/work/project"
|
||||
|
||||
def test_update_cwd_stores_translated_cwd_on_wsl(self, manager, monkeypatch):
|
||||
monkeypatch.setattr("hermes_constants._wsl_detected", True)
|
||||
state = manager.create_session(cwd="/tmp/old")
|
||||
|
||||
updated = manager.update_cwd(state.session_id, cwd=r"C:\Users\foo\project")
|
||||
|
||||
assert updated is not None
|
||||
assert updated.cwd == "/mnt/c/Users/foo/project"
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# fork
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# list / cleanup / remove
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSymlinkAliasNormalization:
|
||||
"""Ported from PrimeIntellect-ai/prime-agent#628 — symlink aliases of the
|
||||
same directory (macOS ``/var`` vs ``/private/var``, ``/tmp`` vs
|
||||
``/private/tmp``) must compare equal, or ACP history filters silently drop
|
||||
a workspace's own sessions."""
|
||||
|
||||
def test_symlink_alias_compares_equal(self, tmp_path):
|
||||
real = tmp_path / "real"
|
||||
real.mkdir()
|
||||
alias = tmp_path / "alias"
|
||||
alias.symlink_to(real)
|
||||
assert acp_session._normalize_cwd_for_compare(
|
||||
str(alias)
|
||||
) == acp_session._normalize_cwd_for_compare(str(real))
|
||||
|
||||
def test_distinct_dirs_still_compare_different(self, tmp_path):
|
||||
a = tmp_path / "a"
|
||||
b = tmp_path / "b"
|
||||
a.mkdir()
|
||||
b.mkdir()
|
||||
assert acp_session._normalize_cwd_for_compare(
|
||||
str(a)
|
||||
) != acp_session._normalize_cwd_for_compare(str(b))
|
||||
|
||||
def test_missing_path_keeps_lexical_normalization(self):
|
||||
# realpath(strict=False) is lexical for nonexistent paths, so cwds
|
||||
# that don't exist on this host (e.g. WSL-translated drives) behave
|
||||
# exactly as the old normpath comparison did.
|
||||
assert acp_session._normalize_cwd_for_compare(
|
||||
"/nonexistent-hermes-test/x/../y"
|
||||
) == "/nonexistent-hermes-test/y"
|
||||
|
||||
def test_list_sessions_matches_symlink_alias_cwd(self, manager, tmp_path):
|
||||
real = tmp_path / "proj"
|
||||
real.mkdir()
|
||||
alias = tmp_path / "link"
|
||||
alias.symlink_to(real)
|
||||
state = manager.create_session(cwd=str(real))
|
||||
state.history.append({"role": "user", "content": "hello"})
|
||||
listed = manager.list_sessions(cwd=str(alias))
|
||||
assert [s["session_id"] for s in listed] == [state.session_id]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# list / cleanup
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestListAndCleanup:
|
||||
def test_list_sessions_empty(self, manager):
|
||||
assert manager.list_sessions() == []
|
||||
|
||||
|
||||
|
||||
def test_save_session_preserves_existing_messages_on_encode_failure(self, manager):
|
||||
"""Regression for #13675: a bad message in state.history must not
|
||||
clobber the previously-persisted transcript. replace_messages()
|
||||
wraps DELETE + INSERT in a single rolled-back-on-exception txn.
|
||||
"""
|
||||
state = manager.create_session()
|
||||
state.history.append({"role": "user", "content": "original"})
|
||||
manager.save_session(state.session_id)
|
||||
|
||||
# Now swap history with a message whose tool_calls is non-JSON-serializable.
|
||||
# _execute_write rolls back; the previously persisted "original" stays.
|
||||
state.history = [
|
||||
{"role": "user", "content": "replacement"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [{"bad": object()}],
|
||||
},
|
||||
]
|
||||
manager.save_session(state.session_id)
|
||||
|
||||
db = manager._get_db()
|
||||
messages = db.get_messages_as_conversation(state.session_id)
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["role"] == "user"
|
||||
assert messages[0]["content"] == "original"
|
||||
assert isinstance(messages[0].get("timestamp"), (int, float))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# persistence — sessions survive process restarts (via SessionDB)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPersistence:
|
||||
"""Verify that sessions are persisted to SessionDB and can be restored."""
|
||||
|
||||
def test_first_persist_keeps_provider_snapshot(self, tmp_path):
|
||||
"""The FIRST row written for an ACP session carries provider/base_url/api_mode,
|
||||
so a restart before any later save restores the same route (#9812)."""
|
||||
agent = SimpleNamespace(
|
||||
model="test-model", provider="anthropic",
|
||||
base_url="https://anthropic.example/v1", api_mode="anthropic_messages",
|
||||
)
|
||||
db = SessionDB(tmp_path / "state.db")
|
||||
manager = SessionManager(agent_factory=lambda: agent, db=db)
|
||||
state = manager.create_session(cwd="/work")
|
||||
state.history.append({"role": "user", "content": "hello"})
|
||||
manager.save_session(state.session_id)
|
||||
|
||||
mc = json.loads(db.get_session(state.session_id)["model_config"])
|
||||
assert mc == {"cwd": "/work", "provider": "anthropic",
|
||||
"base_url": "https://anthropic.example/v1", "api_mode": "anthropic_messages"}
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_only_restores_acp_sessions(self, manager):
|
||||
"""get_session should not restore non-ACP sessions from DB."""
|
||||
db = manager._get_db()
|
||||
# Manually create a CLI session in the DB.
|
||||
db.create_session(session_id="cli-session-123", source="cli", model="test")
|
||||
# Should not be found via ACP SessionManager.
|
||||
assert manager.get_session("cli-session-123") is None
|
||||
|
||||
def test_sessions_searchable_via_fts(self, manager):
|
||||
"""ACP sessions stored in SessionDB are searchable via FTS5."""
|
||||
state = manager.create_session()
|
||||
state.history.append({"role": "user", "content": "how do I configure nginx"})
|
||||
state.history.append({"role": "assistant", "content": "Here is the nginx config..."})
|
||||
manager.save_session(state.session_id)
|
||||
|
||||
db = manager._get_db()
|
||||
results = db.search_messages("nginx")
|
||||
assert len(results) > 0
|
||||
session_ids = {r["session_id"] for r in results}
|
||||
assert state.session_id in session_ids
|
||||
|
||||
|
||||
def test_assistant_reasoning_fields_persisted(self, manager):
|
||||
"""ACP session restore should preserve assistant reasoning context."""
|
||||
state = manager.create_session()
|
||||
state.history.append({
|
||||
"role": "assistant",
|
||||
"content": "hello",
|
||||
"reasoning": "step-by-step",
|
||||
"reasoning_details": [
|
||||
{"type": "thinking", "thinking": "first thought"},
|
||||
],
|
||||
"codex_reasoning_items": [
|
||||
{"type": "reasoning", "id": "rs_123", "encrypted_content": "enc_blob"},
|
||||
],
|
||||
})
|
||||
manager.save_session(state.session_id)
|
||||
|
||||
with manager._lock:
|
||||
del manager._sessions[state.session_id]
|
||||
|
||||
restored = manager.get_session(state.session_id)
|
||||
assert restored is not None
|
||||
msg = restored.history[0]
|
||||
assert isinstance(msg.pop("timestamp", None), (int, float))
|
||||
# Load-time durability stamp (#92231): rows materialized from the DB
|
||||
# are marked persisted so a later flush can't re-append them.
|
||||
assert msg.pop("_db_persisted", None) is True
|
||||
assert restored.history == [{
|
||||
"role": "assistant",
|
||||
"content": "hello",
|
||||
"reasoning": "step-by-step",
|
||||
"reasoning_details": [
|
||||
{"type": "thinking", "thinking": "first thought"},
|
||||
],
|
||||
"codex_reasoning_items": [
|
||||
{"type": "reasoning", "id": "rs_123", "encrypted_content": "enc_blob"},
|
||||
],
|
||||
}]
|
||||
|
||||
|
||||
def test_acp_agents_route_human_output_to_stderr(self, tmp_path, monkeypatch):
|
||||
"""ACP agents must keep stdout clean for JSON-RPC stdio transport."""
|
||||
|
||||
def fake_resolve_runtime_provider(requested=None, **kwargs):
|
||||
return {
|
||||
"provider": "openrouter",
|
||||
"api_mode": "chat_completions",
|
||||
"base_url": "https://openrouter.example/v1",
|
||||
"api_key": "test-key",
|
||||
"command": None,
|
||||
"args": [],
|
||||
}
|
||||
|
||||
def fake_agent(**kwargs):
|
||||
return SimpleNamespace(model=kwargs.get("model"), _print_fn=None)
|
||||
|
||||
monkeypatch.setattr("hermes_cli.config.load_config", lambda: {
|
||||
"model": {"provider": "openrouter", "default": "test-model"}
|
||||
})
|
||||
monkeypatch.setattr(
|
||||
"hermes_cli.runtime_provider.resolve_runtime_provider",
|
||||
fake_resolve_runtime_provider,
|
||||
)
|
||||
db = SessionDB(tmp_path / "state.db")
|
||||
|
||||
with patch("run_agent.AIAgent", side_effect=fake_agent):
|
||||
manager = SessionManager(db=db)
|
||||
state = manager.create_session(cwd="/work")
|
||||
|
||||
stdout_buf = io.StringIO()
|
||||
stderr_buf = io.StringIO()
|
||||
with contextlib.redirect_stdout(stdout_buf), contextlib.redirect_stderr(stderr_buf):
|
||||
state.agent._print_fn("ACP noise")
|
||||
|
||||
assert stdout_buf.getvalue() == ""
|
||||
assert stderr_buf.getvalue() == "ACP noise\n"
|
||||
@@ -0,0 +1,192 @@
|
||||
"""Tests for the update_session_meta fix.
|
||||
|
||||
Verifies that:
|
||||
1. SessionDB.update_session_meta() exists and works correctly via the
|
||||
public _execute_write path (not db._lock / db._conn directly).
|
||||
2. session.py _persist() no longer touches db._lock or db._conn.
|
||||
3. update_session_meta updates the correct columns atomically.
|
||||
"""
|
||||
|
||||
import ast
|
||||
import json
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch, call
|
||||
|
||||
import pytest
|
||||
|
||||
from hermes_state import SessionDB
|
||||
from acp_adapter.session import SessionManager
|
||||
|
||||
|
||||
def _tmp_db(tmp_path):
|
||||
return SessionDB(db_path=tmp_path / "state.db")
|
||||
|
||||
|
||||
def _mock_agent():
|
||||
return MagicMock(name="MockAIAgent")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# hermes_state.SessionDB.update_session_meta — unit tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestUpdateSessionMeta:
|
||||
"""Direct unit tests for the new public method."""
|
||||
|
||||
def test_method_exists(self, tmp_path):
|
||||
db = _tmp_db(tmp_path)
|
||||
assert hasattr(db, "update_session_meta"), (
|
||||
"SessionDB must have update_session_meta() public method"
|
||||
)
|
||||
assert callable(db.update_session_meta)
|
||||
|
||||
def test_updates_model_config(self, tmp_path):
|
||||
db = _tmp_db(tmp_path)
|
||||
db.create_session("s1", source="acp", model="gpt-4")
|
||||
|
||||
new_meta = json.dumps({"cwd": "/new/path", "provider": "openai"})
|
||||
db.update_session_meta("s1", new_meta, model=None)
|
||||
|
||||
row = db.get_session("s1")
|
||||
stored = json.loads(row["model_config"])
|
||||
assert stored["cwd"] == "/new/path"
|
||||
assert stored["provider"] == "openai"
|
||||
|
||||
|
||||
|
||||
def test_uses_execute_write_not_private_api(self, tmp_path):
|
||||
"""update_session_meta must route through _execute_write, not _conn directly."""
|
||||
db = _tmp_db(tmp_path)
|
||||
db.create_session("s4", source="acp")
|
||||
|
||||
call_count = [0]
|
||||
original = db._execute_write
|
||||
|
||||
def patched(fn, *args, **kwargs):
|
||||
call_count[0] += 1
|
||||
return original(fn, *args, **kwargs)
|
||||
|
||||
db._execute_write = patched
|
||||
db.update_session_meta("s4", json.dumps({"cwd": "."}), model="m")
|
||||
|
||||
assert call_count[0] >= 1, (
|
||||
"update_session_meta must call _execute_write at least once"
|
||||
)
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# AST check: session.py must not access db._lock or db._conn
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestNoPrviateDBAccess:
|
||||
"""_persist() in session.py must not access db._lock or db._conn."""
|
||||
|
||||
def test_no_db_private_lock_access(self):
|
||||
with open("acp_adapter/session.py", encoding="utf-8") as f:
|
||||
source = f.read()
|
||||
|
||||
tree = ast.parse(source)
|
||||
|
||||
violations = []
|
||||
for node in ast.walk(tree):
|
||||
# Looking for: db._lock or db._conn
|
||||
if isinstance(node, ast.Attribute):
|
||||
if isinstance(node.value, ast.Name) and node.value.id == "db":
|
||||
if node.attr in ("_lock", "_conn"):
|
||||
violations.append(
|
||||
f"db.{node.attr} at line {node.lineno}"
|
||||
)
|
||||
|
||||
assert violations == [], (
|
||||
"session.py accesses private SessionDB internals: "
|
||||
+ ", ".join(violations)
|
||||
+ " — use db.update_session_meta() instead"
|
||||
)
|
||||
|
||||
def test_persist_calls_update_session_meta(self):
|
||||
"""AST check: _persist must call db.update_session_meta()."""
|
||||
with open("acp_adapter/session.py", encoding="utf-8") as f:
|
||||
tree = ast.parse(f.read())
|
||||
|
||||
found = False
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.FunctionDef) and node.name == "_persist":
|
||||
for child in ast.walk(node):
|
||||
if isinstance(child, ast.Call):
|
||||
func = child.func
|
||||
if isinstance(func, ast.Attribute):
|
||||
if func.attr == "update_session_meta":
|
||||
found = True
|
||||
break
|
||||
break
|
||||
|
||||
assert found, (
|
||||
"_persist() must call db.update_session_meta() "
|
||||
"instead of db._conn.execute() directly"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Integration: _persist round-trip via SessionManager
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestPersistRoundTrip:
|
||||
"""End-to-end: save a session and verify DB state is correct."""
|
||||
|
||||
# A session row is only minted once the session has history — an empty ACP
|
||||
# session is deliberately not persisted (see _persist), because ACP clients
|
||||
# open sessions they never prompt. These tests seed one message first so
|
||||
# they exercise the metadata round-trip they actually care about.
|
||||
|
||||
def test_cwd_persisted_via_update_session_meta(self, tmp_path):
|
||||
db = _tmp_db(tmp_path)
|
||||
manager = SessionManager(agent_factory=_mock_agent, db=db)
|
||||
|
||||
state = manager.create_session(cwd="/original")
|
||||
state.history.append({"role": "user", "content": "hi"})
|
||||
manager.save_session(state.session_id)
|
||||
assert db.get_session(state.session_id) is not None
|
||||
|
||||
# Simulate cwd change and save
|
||||
state.cwd = "/updated"
|
||||
manager.save_session(state.session_id)
|
||||
|
||||
row = db.get_session(state.session_id)
|
||||
mc = json.loads(row["model_config"])
|
||||
assert mc["cwd"] == "/updated"
|
||||
|
||||
def test_model_persisted_via_update_session_meta(self, tmp_path):
|
||||
db = _tmp_db(tmp_path)
|
||||
manager = SessionManager(agent_factory=_mock_agent, db=db)
|
||||
|
||||
state = manager.create_session()
|
||||
state.history.append({"role": "user", "content": "hi"})
|
||||
manager.save_session(state.session_id)
|
||||
|
||||
state.model = "new-model-xyz"
|
||||
manager.save_session(state.session_id)
|
||||
|
||||
row = db.get_session(state.session_id)
|
||||
assert row["model"] == "new-model-xyz"
|
||||
|
||||
def test_existing_model_not_cleared_on_save(self, tmp_path):
|
||||
"""If state.model is empty, the DB model column must not be overwritten."""
|
||||
db = _tmp_db(tmp_path)
|
||||
manager = SessionManager(agent_factory=_mock_agent, db=db)
|
||||
|
||||
state = manager.create_session()
|
||||
state.history.append({"role": "user", "content": "hi"})
|
||||
manager.save_session(state.session_id)
|
||||
# Manually set a model in DB
|
||||
db.update_session_meta(state.session_id, json.dumps({"cwd": "."}), model="stored-model")
|
||||
|
||||
# Now save with empty model
|
||||
state.model = ""
|
||||
manager.save_session(state.session_id)
|
||||
|
||||
row = db.get_session(state.session_id)
|
||||
assert row["model"] == "stored-model", (
|
||||
"COALESCE must preserve the existing model when new value is NULL"
|
||||
)
|
||||
@@ -0,0 +1,80 @@
|
||||
"""Tests for ACP session-provenance derivation (issue #33617).
|
||||
|
||||
Exercises acp_adapter.provenance against a real SessionDB — no mocks — covering
|
||||
the acceptance-criteria matrix: root session, compression-split continuation,
|
||||
multi-depth chains, rotation flagging, and graceful handling of unknown ids.
|
||||
"""
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from acp_adapter.provenance import build_session_provenance, session_provenance_meta
|
||||
from hermes_state import SessionDB
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def db(tmp_path):
|
||||
d = SessionDB(db_path=tmp_path / "state.db")
|
||||
yield d
|
||||
|
||||
|
||||
def _mk(db, sid, parent=None):
|
||||
db.create_session(session_id=sid, source="acp", parent_session_id=parent)
|
||||
|
||||
|
||||
def test_root_session_no_compression(db):
|
||||
_mk(db, "root1")
|
||||
prov = build_session_provenance(db, "acp-1", "root1")
|
||||
assert prov["acpSessionId"] == "acp-1"
|
||||
assert prov["currentHermesSessionId"] == "root1"
|
||||
assert prov["rootHermesSessionId"] == "root1"
|
||||
assert prov["parentHermesSessionId"] is None
|
||||
assert prov["sessionKind"] == "root"
|
||||
assert prov["compressionDepth"] == 0
|
||||
assert "reason" not in prov # no rotation signalled
|
||||
|
||||
|
||||
def test_compression_split_continuation(db):
|
||||
# Parent ended with compression, child created afterwards.
|
||||
_mk(db, "old")
|
||||
db.end_session("old", "compression")
|
||||
time.sleep(0.001)
|
||||
_mk(db, "new", parent="old")
|
||||
|
||||
prov = build_session_provenance(
|
||||
db, "acp-1", "new", previous_hermes_session_id="old"
|
||||
)
|
||||
assert prov["sessionKind"] == "continuation"
|
||||
assert prov["parentHermesSessionId"] == "old"
|
||||
assert prov["rootHermesSessionId"] == "old"
|
||||
assert prov["compressionDepth"] == 1
|
||||
assert prov["previousHermesSessionId"] == "old"
|
||||
# Head rotated this turn → reason/creatorKind flagged.
|
||||
assert prov["reason"] == "compression"
|
||||
assert prov["creatorKind"] == "compression"
|
||||
|
||||
|
||||
|
||||
|
||||
def test_non_compression_parent_is_root_not_continuation(db):
|
||||
# A child with a parent that did NOT end via compression (e.g. delegate
|
||||
# or branch child) must not be reported as a compression continuation.
|
||||
_mk(db, "p")
|
||||
_mk(db, "c", parent="p") # parent still live, no end_reason
|
||||
prov = build_session_provenance(db, "acp-1", "c")
|
||||
assert prov["sessionKind"] == "root"
|
||||
assert prov["compressionDepth"] == 0
|
||||
assert prov["rootHermesSessionId"] == "p" # lineage root still walked
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_meta_wrapper_shape(db):
|
||||
_mk(db, "root1")
|
||||
meta = session_provenance_meta(db, "acp-1", "root1")
|
||||
assert set(meta.keys()) == {"hermes"}
|
||||
assert "sessionProvenance" in meta["hermes"]
|
||||
assert meta["hermes"]["sessionProvenance"]["currentHermesSessionId"] == "root1"
|
||||
@@ -0,0 +1,287 @@
|
||||
"""Tests for acp_adapter.tools — tool kind mapping and ACP content building."""
|
||||
|
||||
|
||||
import pytest
|
||||
|
||||
from acp_adapter.edit_approval import EditProposal
|
||||
from acp_adapter.tools import (
|
||||
TOOL_KIND_MAP,
|
||||
build_tool_complete,
|
||||
build_tool_start,
|
||||
build_tool_title,
|
||||
extract_locations,
|
||||
get_tool_kind,
|
||||
make_tool_call_id,
|
||||
)
|
||||
from acp.schema import (
|
||||
FileEditToolCallContent,
|
||||
ContentToolCallContent,
|
||||
ToolCallLocation,
|
||||
ToolCallStart,
|
||||
ToolCallProgress,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# TOOL_KIND_MAP coverage
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
COMMON_HERMES_TOOLS = ["read_file", "search_files", "terminal", "patch", "write_file", "process"]
|
||||
|
||||
|
||||
class TestToolKindMap:
|
||||
def test_all_hermes_tools_have_kind(self):
|
||||
"""Every common hermes tool should appear in TOOL_KIND_MAP."""
|
||||
for tool in COMMON_HERMES_TOOLS:
|
||||
assert tool in TOOL_KIND_MAP, f"{tool} missing from TOOL_KIND_MAP"
|
||||
|
||||
def test_tool_kind_read_file(self):
|
||||
assert get_tool_kind("read_file") == "read"
|
||||
|
||||
def test_tool_kind_terminal(self):
|
||||
assert get_tool_kind("terminal") == "execute"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_unknown_tool_returns_other_kind(self):
|
||||
assert get_tool_kind("nonexistent_tool_xyz") == "other"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# make_tool_call_id
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestMakeToolCallId:
|
||||
def test_returns_string(self):
|
||||
tc_id = make_tool_call_id()
|
||||
assert isinstance(tc_id, str)
|
||||
|
||||
def test_starts_with_tc_prefix(self):
|
||||
tc_id = make_tool_call_id()
|
||||
assert tc_id.startswith("tc-")
|
||||
|
||||
def test_ids_are_unique(self):
|
||||
ids = {make_tool_call_id() for _ in range(100)}
|
||||
assert len(ids) == 100
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# build_tool_title
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestBuildToolTitle:
|
||||
def test_terminal_title_includes_command(self):
|
||||
title = build_tool_title("terminal", {"command": "ls -la /tmp"})
|
||||
assert "ls -la /tmp" in title
|
||||
|
||||
def test_terminal_title_truncates_long_command(self):
|
||||
long_cmd = "x" * 200
|
||||
title = build_tool_title("terminal", {"command": long_cmd})
|
||||
assert len(title) < 120
|
||||
assert "..." in title
|
||||
|
||||
def test_read_file_title(self):
|
||||
title = build_tool_title("read_file", {"path": "/etc/hosts"})
|
||||
assert "hosts" in title
|
||||
|
||||
def test_search_title(self):
|
||||
title = build_tool_title("search_files", {"pattern": "TODO"})
|
||||
assert "TODO" in title
|
||||
|
||||
def test_skill_view_title_includes_skill_name(self):
|
||||
title = build_tool_title("skill_view", {"name": "github-pitfalls"})
|
||||
assert "github-pitfalls" in title
|
||||
|
||||
def test_execute_code_title_includes_first_code_line(self):
|
||||
title = build_tool_title("execute_code", {"code": "\nfrom hermes_tools import terminal\nprint('done')"})
|
||||
assert "from hermes_tools import terminal" in title
|
||||
|
||||
def test_unknown_tool_uses_name(self):
|
||||
title = build_tool_title("some_new_tool", {"foo": "bar"})
|
||||
assert title == "some_new_tool"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"tool_name, args",
|
||||
[
|
||||
("terminal", {"command": "git status --short"}),
|
||||
("read_file", {"path": "/etc/hosts", "offset": 10}),
|
||||
("search_files", {"pattern": "TODO", "path": "src"}),
|
||||
("web_search", {"query": "hermes agent acp"}),
|
||||
("execute_code", {"code": "\nfrom hermes_tools import terminal\nprint('done')"}),
|
||||
("skill_view", {"name": "github", "file_path": "references/x.md"}),
|
||||
],
|
||||
)
|
||||
def test_title_derives_from_display_preview(self, tool_name, args):
|
||||
"""ACP titles are the shared agent.display preview, not a parallel per-tool table."""
|
||||
from agent.display import build_tool_preview
|
||||
|
||||
assert build_tool_preview(tool_name, args, max_len=80) in build_tool_title(tool_name, args)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# build_tool_start
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestBuildToolStart:
|
||||
def test_build_tool_start_for_patch(self):
|
||||
"""patch start should not duplicate the edit-approval diff."""
|
||||
args = {
|
||||
"path": "src/main.py",
|
||||
"old_string": "print('hello')",
|
||||
"new_string": "print('world')",
|
||||
}
|
||||
result = build_tool_start("tc-1", "patch", args)
|
||||
assert isinstance(result, ToolCallStart)
|
||||
assert result.kind == "edit"
|
||||
assert len(result.content) >= 1
|
||||
item = result.content[0]
|
||||
assert isinstance(item, ContentToolCallContent)
|
||||
assert "Approval prompt shows the diff" in item.content.text
|
||||
assert "src/main.py" in item.content.text
|
||||
|
||||
|
||||
def test_auto_approved_edit_start_shows_diff_content(self):
|
||||
"""Auto-approved edit starts need the diff because no approval card exists."""
|
||||
args = {"path": "/tmp/acp.txt", "old_string": "old", "new_string": "new"}
|
||||
result = build_tool_start(
|
||||
"tc-auto-edit",
|
||||
"patch",
|
||||
args,
|
||||
edit_diff=EditProposal("patch", "/tmp/acp.txt", "old\n", "new\n", args),
|
||||
)
|
||||
|
||||
assert isinstance(result, ToolCallStart)
|
||||
assert result.kind == "edit"
|
||||
assert len(result.content) == 1
|
||||
item = result.content[0]
|
||||
assert isinstance(item, FileEditToolCallContent)
|
||||
assert item.path == "/tmp/acp.txt"
|
||||
assert item.old_text == "old\n"
|
||||
assert item.new_text == "new\n"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_build_tool_start_for_browser_navigate(self):
|
||||
"""browser_navigate should emit a polished start event."""
|
||||
args = {"url": "https://x.com"}
|
||||
result = build_tool_start("tc-browser-start", "browser_navigate", args)
|
||||
assert isinstance(result, ToolCallStart)
|
||||
assert "https://x.com" in result.title
|
||||
assert result.kind == "fetch"
|
||||
assert result.content[0].content.text == '{\n "url": "https://x.com"\n}'
|
||||
assert result.raw_input is None
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# build_tool_complete
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestBuildToolComplete:
|
||||
def test_build_tool_complete_for_terminal(self):
|
||||
"""Completed terminal call should include output text."""
|
||||
result = build_tool_complete("tc-2", "terminal", "total 42\ndrwxr-xr-x 2 root root 4096 ...")
|
||||
assert isinstance(result, ToolCallProgress)
|
||||
assert result.status == "completed"
|
||||
assert len(result.content) >= 1
|
||||
content_item = result.content[0]
|
||||
assert isinstance(content_item, ContentToolCallContent)
|
||||
assert "total 42" in content_item.content.text
|
||||
assert result.raw_output is None
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_build_tool_complete_marks_returncode_nonzero_as_failed(self):
|
||||
result = build_tool_complete("tc-fail", "execute_code", '{"output": "bad", "returncode": 2}')
|
||||
assert result.status == "failed"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_build_tool_complete_for_search_files_formats_matches(self):
|
||||
result = build_tool_complete(
|
||||
"tc-search",
|
||||
"search_files",
|
||||
'{"total_count":2,"matches":[{"path":"README.md","line":3,"content":"TODO: fix this"},{"path":"src/app.py","line":9,"content":"needle"}],"truncated":true}\n\n[Hint: Results truncated. Use offset=12 to see more.]',
|
||||
)
|
||||
text = result.content[0].content.text
|
||||
assert "Search results" in text
|
||||
assert "Found 2 matches" in text
|
||||
assert "README.md:3" in text
|
||||
assert "TODO: fix this" in text
|
||||
assert "Results truncated" in text
|
||||
assert result.raw_output is None
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_build_tool_complete_generically_formats_unknown_json_dict_without_raw_output(self):
|
||||
result = build_tool_complete(
|
||||
"tc-recall-search",
|
||||
"memory_archive_search",
|
||||
'{"results":[{"id":"obs-1","status":"active","content":"Recall should render as a readable summary."}],"trust":"lower-trust archive evidence"}',
|
||||
)
|
||||
text = result.content[0].content.text
|
||||
assert "memory_archive_search result" in text
|
||||
assert "lower-trust archive evidence" in text
|
||||
assert "Recall should render as a readable summary" in text
|
||||
assert "{\"results\"" not in text
|
||||
assert result.raw_output is None
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# extract_locations
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestExtractLocations:
|
||||
def test_extract_locations_with_path(self):
|
||||
args = {"path": "src/app.py", "offset": 42}
|
||||
locs = extract_locations(args)
|
||||
assert len(locs) == 1
|
||||
assert isinstance(locs[0], ToolCallLocation)
|
||||
assert locs[0].path == "src/app.py"
|
||||
assert locs[0].line == 42
|
||||
|
||||
def test_extract_locations_without_path(self):
|
||||
args = {"command": "echo hi"}
|
||||
locs = extract_locations(args)
|
||||
assert locs == []
|
||||
@@ -0,0 +1,39 @@
|
||||
"""Fast-path fixtures shared across tests/agent/.
|
||||
|
||||
Many tests in this directory exercise the retry/backoff paths in the
|
||||
agent loop. Production code uses ``jittered_backoff(base_delay=5.0)``
|
||||
with a ``while time.time() < sleep_end`` loop — a single retry test
|
||||
spends 5+ seconds of real wall-clock time on backoff waits.
|
||||
|
||||
Mocking ``jittered_backoff`` to return 0.0 collapses the while-loop
|
||||
to a no-op (``time.time() < time.time() + 0`` is false immediately),
|
||||
which handles the most common case without touching ``time.sleep``.
|
||||
|
||||
We deliberately DO NOT mock ``time.sleep`` here — some tests
|
||||
(test_interrupt_propagation, test_primary_runtime_restore, etc.) use
|
||||
the real ``time.sleep`` for threading coordination or assert that it
|
||||
was called with specific values. Tests that want to additionally
|
||||
fast-path direct ``time.sleep(N)`` calls in production code should
|
||||
monkeypatch ``run_agent.time.sleep`` locally (see
|
||||
``test_anthropic_error_handling.py`` for the pattern).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _fast_retry_backoff(request, monkeypatch):
|
||||
"""Short-circuit retry backoff for all tests in this directory.
|
||||
|
||||
Tests that assert on the real backoff value opt out with
|
||||
``@pytest.mark.real_retry_backoff``.
|
||||
"""
|
||||
if request.node.get_closest_marker("real_retry_backoff"):
|
||||
return
|
||||
# The agent.turn_* retry paths import ``jittered_backoff`` lazily from
|
||||
# ``agent.retry_utils``; patch it there so rate-limit / invalid-response /
|
||||
# server-error retries don't burn real wall-clock seconds.
|
||||
from agent import retry_utils as _retry_utils
|
||||
monkeypatch.setattr(_retry_utils, "jittered_backoff", lambda *a, **k: 0.0)
|
||||
@@ -0,0 +1,179 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Runnable proof for issue #48013 — image-dimension 400 session brick.
|
||||
|
||||
Before the fix, ``agent.conversation_compression.try_shrink_image_parts_in_messages``
|
||||
silently discarded a *pixel-correct* downscale whenever the re-encoded PNG was
|
||||
larger in bytes than the original (the common case for downscaled Retina
|
||||
screenshots). The image was left at its original oversized dimensions, the
|
||||
provider re-rejected it on retry, and the session wedged forever on the
|
||||
Anthropic many-image 2000px path.
|
||||
|
||||
This script reproduces the exact scenario with REAL Pillow (no mocks): it
|
||||
synthesizes screenshot-like PNGs at the dimensions from the issue's table —
|
||||
images that are small in bytes (under the 4 MB budget) but over the 2000px
|
||||
per-side cap — and runs the real recovery helper. It asserts every image is
|
||||
brought under the cap and that the helper reports success.
|
||||
|
||||
Run directly to see a human-readable report:
|
||||
|
||||
python tests/agent/repro_image_shrink_brick.py
|
||||
|
||||
Or as a pytest smoke test (skipped automatically when Pillow is absent):
|
||||
|
||||
scripts/run_tests.sh tests/agent/repro_image_shrink_brick.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import io
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
# Make the repo root importable when run as a plain script.
|
||||
_REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
if str(_REPO_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(_REPO_ROOT))
|
||||
|
||||
PIL = pytest.importorskip("PIL", reason="Pillow required for the real-resize proof")
|
||||
from PIL import Image, ImageDraw # noqa: E402
|
||||
|
||||
from agent.conversation_compression import ( # noqa: E402
|
||||
try_shrink_image_parts_in_messages,
|
||||
)
|
||||
|
||||
# The many-image per-side cap Anthropic reported in the wild (issue #48013).
|
||||
MANY_IMAGE_CAP = 2000
|
||||
BYTE_BUDGET = 4 * 1024 * 1024
|
||||
|
||||
# Dimensions straight from the issue's per-image table. The "REJECTED" rows
|
||||
# are the ones that bricked: tall/large screenshots whose downscale re-encodes
|
||||
# to MORE PNG bytes than the original.
|
||||
CASES = [
|
||||
(2344, 778), # wide — shrank even before the fix
|
||||
(2374, 1144), # wide — shrank even before the fix
|
||||
(2097, 1476), # REJECTED before fix
|
||||
(2247, 1544), # REJECTED before fix
|
||||
(2263, 1644), # REJECTED before fix
|
||||
]
|
||||
|
||||
|
||||
def _make_screenshot_png(width: int, height: int) -> bytes:
|
||||
"""A screenshot-like PNG: mostly flat UI regions so it compresses small.
|
||||
|
||||
Flat regions keep the byte size well under the 4 MB budget, forcing the
|
||||
DIMENSION path (not the byte path) — exactly the code that bricked. The
|
||||
downscale of such an image re-encodes to a comparable-or-larger PNG, which
|
||||
is what the old byte gate wrongly rejected.
|
||||
"""
|
||||
img = Image.new("RGB", (width, height), (245, 245, 247))
|
||||
draw = ImageDraw.Draw(img)
|
||||
for y in range(0, height, 40):
|
||||
shade = 255 - (y // 40) % 6 * 4
|
||||
draw.rectangle([20, y + 5, width - 20, y + 30], fill=(shade, 250, 250))
|
||||
for x in range(0, width, 160):
|
||||
draw.rectangle([x, 0, x + 2, height], fill=(220, 220, 225))
|
||||
draw.text((40, 40), "Some UI text " * 30, fill=(20, 20, 20))
|
||||
buf = io.BytesIO()
|
||||
img.save(buf, format="PNG", optimize=False)
|
||||
return buf.getvalue()
|
||||
|
||||
|
||||
def _data_url(raw: bytes) -> str:
|
||||
return "data:image/png;base64," + base64.b64encode(raw).decode("ascii")
|
||||
|
||||
|
||||
def _decode_dims(data_url: str) -> tuple[int, int]:
|
||||
payload = data_url.partition(",")[2]
|
||||
with Image.open(io.BytesIO(base64.b64decode(payload))) as img:
|
||||
return img.size
|
||||
|
||||
|
||||
def run_proof(verbose: bool = False) -> list[dict]:
|
||||
"""Run the recovery against every case; return per-case results."""
|
||||
results: list[dict] = []
|
||||
for width, height in CASES:
|
||||
raw = _make_screenshot_png(width, height)
|
||||
url = _data_url(raw)
|
||||
# Sanity: this case must be UNDER the byte budget and OVER the pixel cap,
|
||||
# i.e. it exercises the dimension path that bricked.
|
||||
under_byte_budget = len(url) <= BYTE_BUDGET
|
||||
over_pixel_cap = max(width, height) > MANY_IMAGE_CAP
|
||||
|
||||
msgs = [{
|
||||
"role": "user",
|
||||
"content": [{"type": "image_url", "image_url": {"url": url}}],
|
||||
}]
|
||||
changed = try_shrink_image_parts_in_messages(
|
||||
msgs, max_dimension=MANY_IMAGE_CAP,
|
||||
)
|
||||
out_url = msgs[0]["content"][0]["image_url"]["url"]
|
||||
out_dims = _decode_dims(out_url)
|
||||
|
||||
result = {
|
||||
"orig": (width, height),
|
||||
"orig_bytes": len(raw),
|
||||
"under_byte_budget": under_byte_budget,
|
||||
"over_pixel_cap": over_pixel_cap,
|
||||
"changed": changed,
|
||||
"result_dims": out_dims,
|
||||
"under_cap_after": max(out_dims) <= MANY_IMAGE_CAP,
|
||||
}
|
||||
results.append(result)
|
||||
if verbose:
|
||||
status = "OK" if result["under_cap_after"] else "BRICK"
|
||||
print(
|
||||
f" {width}x{height} ({len(raw)//1024:>3} KB)"
|
||||
f" -> changed={changed!s:>5}"
|
||||
f" result={out_dims[0]}x{out_dims[1]}"
|
||||
f" [{status}]"
|
||||
)
|
||||
return results
|
||||
|
||||
|
||||
def test_issue_48013_dimension_shrink_does_not_brick():
|
||||
"""Every dimension-oversized screenshot must be brought under the cap."""
|
||||
results = run_proof()
|
||||
assert results, "no cases ran"
|
||||
for r in results:
|
||||
# Precondition: we really are on the dimension path.
|
||||
assert r["under_byte_budget"], (
|
||||
f"{r['orig']} must be under the byte budget to exercise the bug"
|
||||
)
|
||||
assert r["over_pixel_cap"], f"{r['orig']} must exceed the pixel cap"
|
||||
# The fix: image lands under the cap and the helper reports success.
|
||||
assert r["under_cap_after"], (
|
||||
f"BRICK: {r['orig']} left at {r['result_dims']} "
|
||||
f"(> {MANY_IMAGE_CAP}px) — the shrink recovery discarded a "
|
||||
f"pixel-correct downscale (#48013)"
|
||||
)
|
||||
assert r["changed"] is True, (
|
||||
f"{r['orig']} shrank but helper reported no progress — caller "
|
||||
f"would surface the original error and burn the one-shot retry"
|
||||
)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
print("Issue #48013 proof — image-dimension shrink must not brick sessions")
|
||||
print(f"(many-image per-side cap = {MANY_IMAGE_CAP}px, byte budget = "
|
||||
f"{BYTE_BUDGET // (1024 * 1024)} MB)\n")
|
||||
results = run_proof(verbose=True)
|
||||
bricked = [r for r in results if not r["under_cap_after"]]
|
||||
no_progress = [r for r in results if r["under_cap_after"] and not r["changed"]]
|
||||
print()
|
||||
if bricked:
|
||||
print(f"FAIL: {len(bricked)} image(s) still over the pixel cap (BRICK).")
|
||||
return 1
|
||||
if no_progress:
|
||||
print(f"FAIL: {len(no_progress)} image(s) shrank but helper reported "
|
||||
f"no progress (would burn the retry).")
|
||||
return 1
|
||||
print(f"PASS: all {len(results)} dimension-oversized screenshots brought "
|
||||
f"under {MANY_IMAGE_CAP}px and reported as progress.")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,93 @@
|
||||
"""Regression guard for #31273: HTTP 402 (billing exhaustion) must abort
|
||||
after credential-pool rotation and provider fallback have failed.
|
||||
|
||||
Before the fix, ``FailoverReason.billing`` was in the exclusion set that
|
||||
prevents the loop's ``is_client_error`` branch from firing. When a user
|
||||
ran a pay-per-token provider (OpenRouter, etc.) with no credential pool
|
||||
and no fallback configured, a single 402 cascaded into
|
||||
``agent.api_max_retries`` paid requests against an exhausted balance.
|
||||
Real-world impact: ~$40 burned in 48h on a 24/7 gateway routing Telegram
|
||||
+ Discord traffic.
|
||||
|
||||
The fix removes ``FailoverReason.billing`` from the exclusion set. By
|
||||
the time control reaches the ``is_client_error`` check:
|
||||
* credential-pool rotation has already run (and either ``continue``d
|
||||
on rotation, or returned False because the pool is exhausted/absent).
|
||||
* the eager-fallback branch for billing has also run (and either
|
||||
``continue``d on fallback activation, or fell through because no
|
||||
fallback is configured).
|
||||
Falling through to the retry-backoff path from here just burns paid
|
||||
requests with no recovery mechanism left. Aborting mirrors how 401/403
|
||||
(also ``should_fallback=True``) already behave once their recovery paths
|
||||
have failed.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class TestBillingTriggersClientErrorAbort:
|
||||
"""Mirror the ``is_client_error`` predicate shape used in
|
||||
``agent/conversation_loop.py`` and verify ``FailoverReason.billing``
|
||||
now resolves to True (i.e. aborts the loop).
|
||||
"""
|
||||
|
||||
def _mirror_is_client_error(
|
||||
self,
|
||||
*,
|
||||
classified_retryable: bool,
|
||||
classified_reason,
|
||||
classified_should_compress: bool = False,
|
||||
is_local_validation_error: bool = False,
|
||||
is_context_length_error: bool = False,
|
||||
) -> bool:
|
||||
"""Exact shape of conversation_loop.py's is_client_error check.
|
||||
|
||||
Kept in lock-step with the source. If you change one, change
|
||||
both — or, better, refactor the predicate into a shared helper
|
||||
and have both sites import it.
|
||||
"""
|
||||
from agent.error_classifier import FailoverReason
|
||||
|
||||
return (
|
||||
is_local_validation_error
|
||||
or (
|
||||
not classified_retryable
|
||||
and not classified_should_compress
|
||||
and classified_reason not in {
|
||||
FailoverReason.rate_limit,
|
||||
FailoverReason.overloaded,
|
||||
FailoverReason.context_overflow,
|
||||
FailoverReason.payload_too_large,
|
||||
FailoverReason.long_context_tier,
|
||||
FailoverReason.thinking_signature,
|
||||
}
|
||||
)
|
||||
) and not is_context_length_error
|
||||
|
||||
def test_billing_now_aborts_the_loop(self):
|
||||
"""402 with no fallback / no pool entry → ``is_client_error`` True."""
|
||||
from agent.error_classifier import FailoverReason
|
||||
|
||||
# This is what classify_api_error() returns for a plain 402:
|
||||
# reason=billing, retryable=False, should_compress=False
|
||||
assert self._mirror_is_client_error(
|
||||
classified_retryable=False,
|
||||
classified_reason=FailoverReason.billing,
|
||||
), (
|
||||
"FailoverReason.billing must trigger is_client_error abort after "
|
||||
"credential-pool rotation and provider fallback have failed — see #31273."
|
||||
)
|
||||
|
||||
|
||||
|
||||
def test_context_overflow_still_falls_through_to_compression(self):
|
||||
"""Sanity check: context-overflow must NOT be classified as
|
||||
client error — compression is the recovery path."""
|
||||
from agent.error_classifier import FailoverReason
|
||||
|
||||
assert not self._mirror_is_client_error(
|
||||
classified_retryable=True,
|
||||
classified_reason=FailoverReason.context_overflow,
|
||||
classified_should_compress=True,
|
||||
)
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,171 @@
|
||||
"""Tests for issue #860 — SQLite session transcript deduplication.
|
||||
|
||||
Verifies that:
|
||||
1. _flush_messages_to_session_db uses _last_flushed_db_idx to avoid re-writing
|
||||
2. Multiple _persist_session calls don't duplicate messages
|
||||
3. append_to_transcript(skip_db=True) skips SQLite but writes JSONL
|
||||
4. The gateway doesn't double-write messages the agent already persisted
|
||||
"""
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test: _flush_messages_to_session_db only writes new messages
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestFlushDeduplication:
|
||||
"""Verify _flush_messages_to_session_db tracks what it already wrote."""
|
||||
|
||||
def _make_agent(self, session_db):
|
||||
"""Create a minimal AIAgent with a real session DB."""
|
||||
with patch.dict(os.environ, {"OPENROUTER_API_KEY": "test-key"}):
|
||||
from run_agent import AIAgent
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
quiet_mode=True,
|
||||
session_db=session_db,
|
||||
session_id="test-session-860",
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
# Simulate lazy session creation (normally done by run_conversation)
|
||||
agent._ensure_db_session()
|
||||
return agent
|
||||
|
||||
def test_flush_writes_only_new_messages(self):
|
||||
"""First flush writes all new messages, second flush writes none."""
|
||||
from hermes_state import SessionDB
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db_path = Path(tmpdir) / "test.db"
|
||||
db = SessionDB(db_path=db_path)
|
||||
try:
|
||||
agent = self._make_agent(db)
|
||||
|
||||
conversation_history = [
|
||||
{"role": "user", "content": "old message"},
|
||||
]
|
||||
messages = list(conversation_history) + [
|
||||
{"role": "user", "content": "new question"},
|
||||
{"role": "assistant", "content": "new answer"},
|
||||
]
|
||||
|
||||
# First flush — should write 2 new messages
|
||||
agent._flush_messages_to_session_db(messages, conversation_history)
|
||||
|
||||
rows = db.get_messages(agent.session_id)
|
||||
assert len(rows) == 2, f"Expected 2 messages, got {len(rows)}"
|
||||
|
||||
# Second flush with SAME messages — should write 0 new messages
|
||||
agent._flush_messages_to_session_db(messages, conversation_history)
|
||||
|
||||
rows = db.get_messages(agent.session_id)
|
||||
assert len(rows) == 2, f"Expected still 2 messages after second flush, got {len(rows)}"
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
|
||||
def test_flush_reset_after_compression(self):
|
||||
"""After compression creates a new session, flush index resets."""
|
||||
from hermes_state import SessionDB
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db_path = Path(tmpdir) / "test.db"
|
||||
db = SessionDB(db_path=db_path)
|
||||
try:
|
||||
agent = self._make_agent(db)
|
||||
|
||||
# Write some messages
|
||||
messages = [
|
||||
{"role": "user", "content": "msg1"},
|
||||
{"role": "assistant", "content": "reply1"},
|
||||
]
|
||||
agent._flush_messages_to_session_db(messages, [])
|
||||
|
||||
old_session = agent.session_id
|
||||
assert agent._last_flushed_db_idx == 2
|
||||
|
||||
# Simulate what _compress_context does: new session, reset idx
|
||||
agent.session_id = "compressed-session-new"
|
||||
db.create_session(session_id=agent.session_id, source="test")
|
||||
agent._last_flushed_db_idx = 0
|
||||
|
||||
# Now flush compressed messages to new session
|
||||
compressed_messages = [
|
||||
{"role": "user", "content": "summary of conversation"},
|
||||
]
|
||||
agent._flush_messages_to_session_db(compressed_messages, [])
|
||||
|
||||
new_rows = db.get_messages(agent.session_id)
|
||||
assert len(new_rows) == 1
|
||||
|
||||
# Old session should still have its 2 messages
|
||||
old_rows = db.get_messages(old_session)
|
||||
assert len(old_rows) == 2
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test: append_to_transcript skip_db parameter
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestAppendToTranscriptSkipDb:
|
||||
"""Verify skip_db=True skips the SQLite write."""
|
||||
|
||||
def test_skip_db_prevents_sqlite_write(self, tmp_path):
|
||||
"""With skip_db=True and a real DB, message does NOT appear in SQLite."""
|
||||
from gateway.config import GatewayConfig
|
||||
from gateway.session import SessionStore
|
||||
from hermes_state import SessionDB
|
||||
|
||||
db_path = tmp_path / "test_skip.db"
|
||||
db = SessionDB(db_path=db_path)
|
||||
|
||||
config = GatewayConfig()
|
||||
with patch("gateway.session.SessionStore._ensure_loaded"):
|
||||
store = SessionStore(sessions_dir=tmp_path, config=config)
|
||||
store._db = db
|
||||
store._loaded = True
|
||||
|
||||
session_id = "test-skip-db-real"
|
||||
db.create_session(session_id=session_id, source="test")
|
||||
|
||||
msg = {"role": "assistant", "content": "hello world"}
|
||||
store.append_to_transcript(session_id, msg, skip_db=True)
|
||||
|
||||
# SQLite should NOT have the message
|
||||
rows = db.get_messages(session_id)
|
||||
assert len(rows) == 0, f"Expected 0 DB rows with skip_db=True, got {len(rows)}"
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test: _last_flushed_db_idx initialization
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestFlushIdxInit:
|
||||
"""Verify _last_flushed_db_idx is properly initialized."""
|
||||
|
||||
def test_init_zero(self):
|
||||
"""Agent starts with _last_flushed_db_idx = 0."""
|
||||
with patch.dict(os.environ, {"OPENROUTER_API_KEY": "test-key"}):
|
||||
from run_agent import AIAgent
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
assert agent._last_flushed_db_idx == 0
|
||||
|
||||
@@ -1,16 +1,21 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from agent import account_usage
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, payload):
|
||||
def __init__(self, payload, status_code=200):
|
||||
self._payload = payload
|
||||
self.status_code = status_code
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
if self.status_code >= 400:
|
||||
request = httpx.Request("GET", "https://chatgpt.com/backend-api/wham/usage")
|
||||
response = httpx.Response(self.status_code, request=request)
|
||||
raise httpx.HTTPStatusError("request failed", request=request, response=response)
|
||||
|
||||
def json(self):
|
||||
return self._payload
|
||||
@@ -162,6 +167,40 @@ def test_codex_usage_account_id_read_failure_keeps_singleton_token(monkeypatch,
|
||||
assert "ChatGPT-Account-Id" not in calls[0]["headers"]
|
||||
|
||||
|
||||
def test_codex_usage_retries_401_with_forced_refresh(monkeypatch, codex_usage_payload):
|
||||
credential_calls = []
|
||||
request_calls = []
|
||||
responses = [_FakeResponse({}, status_code=401), _FakeResponse(codex_usage_payload)]
|
||||
|
||||
def resolve(**kwargs):
|
||||
credential_calls.append(kwargs)
|
||||
token = "fresh-token" if kwargs.get("force_refresh") else "revoked-token"
|
||||
return {"api_key": token, "base_url": "https://chatgpt.com/backend-api/codex"}
|
||||
|
||||
class Client:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *exc):
|
||||
return False
|
||||
|
||||
def get(self, url, headers):
|
||||
request_calls.append(headers["Authorization"])
|
||||
return responses.pop(0)
|
||||
|
||||
monkeypatch.setattr(account_usage, "resolve_codex_runtime_credentials", resolve)
|
||||
monkeypatch.setattr(account_usage, "_read_codex_tokens", lambda: {"tokens": {}})
|
||||
monkeypatch.setattr(account_usage.httpx, "Client", lambda timeout: Client())
|
||||
|
||||
snapshot = account_usage.fetch_account_usage("openai-codex")
|
||||
|
||||
assert snapshot is not None
|
||||
assert snapshot.windows[0].label == "Session"
|
||||
assert credential_calls == [
|
||||
{"refresh_if_expiring": True},
|
||||
{"refresh_if_expiring": True, "force_refresh": True},
|
||||
]
|
||||
assert request_calls == ["Bearer revoked-token", "Bearer fresh-token"]
|
||||
|
||||
|
||||
# ── Banked rate-limit reset credits (`/usage reset`) ─────────────────────────
|
||||
@@ -216,14 +255,93 @@ def _usage_payload_with_resets(primary_used, secondary_used, banked):
|
||||
|
||||
|
||||
|
||||
def test_redeem_retries_401_with_forced_refresh(monkeypatch):
|
||||
credential_calls = []
|
||||
request_calls = []
|
||||
client_count = 0
|
||||
payload = _usage_payload_with_resets(100, 40, 1)
|
||||
|
||||
def resolve(base_url, api_key, *, force_refresh=False):
|
||||
credential_calls.append(force_refresh)
|
||||
token = "fresh-token" if force_refresh else "revoked-token"
|
||||
return token, "https://chatgpt.com/backend-api/codex", None
|
||||
|
||||
class Client(_FakeResetClient):
|
||||
def get(self, url, headers):
|
||||
request_calls.append(("GET", headers["Authorization"]))
|
||||
if headers["Authorization"] == "Bearer revoked-token":
|
||||
return _FakeResponse({}, status_code=401)
|
||||
return _FakeResponse(payload)
|
||||
|
||||
def post(self, url, headers=None, json=None):
|
||||
request_calls.append(("POST", headers["Authorization"]))
|
||||
return _FakeResponse({"code": "reset", "windows_reset": 2})
|
||||
|
||||
def client_factory(timeout):
|
||||
nonlocal client_count
|
||||
client_count += 1
|
||||
return Client([], payload)
|
||||
|
||||
monkeypatch.setattr(account_usage, "_resolve_codex_usage_credentials", resolve)
|
||||
monkeypatch.setattr(account_usage.httpx, "Client", client_factory)
|
||||
|
||||
result = account_usage.redeem_codex_reset_credit()
|
||||
|
||||
assert result.status == "reset"
|
||||
assert credential_calls == [False, True]
|
||||
assert request_calls == [
|
||||
("GET", "Bearer revoked-token"),
|
||||
("GET", "Bearer fresh-token"),
|
||||
("POST", "Bearer fresh-token"),
|
||||
]
|
||||
assert client_count == 2
|
||||
|
||||
|
||||
def test_redeem_missing_credentials_reports_unavailable(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
account_usage,
|
||||
"_resolve_codex_usage_credentials",
|
||||
lambda base_url, api_key: (_ for _ in ()).throw(RuntimeError("no creds")),
|
||||
lambda base_url, api_key, **kwargs: (_ for _ in ()).throw(RuntimeError("no creds")),
|
||||
)
|
||||
|
||||
result = account_usage.redeem_codex_reset_credit()
|
||||
|
||||
assert result.status == "unavailable"
|
||||
assert "hermes auth" in result.message
|
||||
|
||||
|
||||
def test_codex_usage_401_retry_refreshes_the_explicit_credential_not_another_account(monkeypatch, codex_usage_payload):
|
||||
"""A live agent on pool entry B hands its own api_key in; after a 401 the retry must refresh B,
|
||||
not re-resolve and render the singleton/pool account A's usage."""
|
||||
request_calls = []
|
||||
refresh_hints = []
|
||||
responses = [_FakeResponse({}, status_code=401), _FakeResponse(codex_usage_payload)]
|
||||
|
||||
class Pool:
|
||||
def try_refresh_matching(self, api_key_hint=None, credential_id=None):
|
||||
refresh_hints.append(api_key_hint)
|
||||
return SimpleNamespace(runtime_api_key="pool-B-fresh", runtime_base_url=None)
|
||||
|
||||
class Client:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *exc):
|
||||
return False
|
||||
|
||||
def get(self, url, headers):
|
||||
request_calls.append(headers["Authorization"])
|
||||
return responses.pop(0)
|
||||
|
||||
monkeypatch.setattr(account_usage, "resolve_codex_runtime_credentials",
|
||||
lambda **kwargs: pytest.fail("must not re-resolve another account's credential"))
|
||||
monkeypatch.setattr(account_usage, "_read_codex_tokens", lambda: {"tokens": {"access_token": "singleton-A"}})
|
||||
monkeypatch.setattr("agent.credential_pool.load_pool", lambda provider: Pool())
|
||||
monkeypatch.setattr(account_usage.httpx, "Client", lambda timeout: Client())
|
||||
|
||||
snapshot = account_usage.fetch_account_usage(
|
||||
"openai-codex", base_url="https://chatgpt.com/backend-api/codex", api_key="pool-B-revoked")
|
||||
|
||||
assert snapshot is not None
|
||||
assert refresh_hints == ["pool-B-revoked"]
|
||||
assert request_calls == ["Bearer pool-B-revoked", "Bearer pool-B-fresh"]
|
||||
|
||||
@@ -0,0 +1,203 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from agent.account_usage import (
|
||||
AccountUsageSnapshot,
|
||||
AccountUsageWindow,
|
||||
fetch_account_usage,
|
||||
render_account_usage_lines,
|
||||
)
|
||||
|
||||
|
||||
class _Response:
|
||||
def __init__(self, payload, status_code=200):
|
||||
self._payload = payload
|
||||
self.status_code = status_code
|
||||
|
||||
def raise_for_status(self):
|
||||
if self.status_code >= 400:
|
||||
raise RuntimeError(f"HTTP {self.status_code}")
|
||||
|
||||
def json(self):
|
||||
return self._payload
|
||||
|
||||
|
||||
class _Client:
|
||||
def __init__(self, payload):
|
||||
self._payload = payload
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc, tb):
|
||||
return False
|
||||
|
||||
def get(self, url, headers=None):
|
||||
return _Response(self._payload)
|
||||
|
||||
|
||||
class _RoutingClient:
|
||||
def __init__(self, payloads):
|
||||
self._payloads = payloads
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc, tb):
|
||||
return False
|
||||
|
||||
def get(self, url, headers=None):
|
||||
return _Response(self._payloads[url])
|
||||
|
||||
|
||||
def test_fetch_account_usage_codex(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"agent.account_usage.resolve_codex_runtime_credentials",
|
||||
lambda refresh_if_expiring=True: {
|
||||
"provider": "openai-codex",
|
||||
"base_url": "https://chatgpt.com/backend-api/codex",
|
||||
"api_key": "access-token",
|
||||
},
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"agent.account_usage._read_codex_tokens",
|
||||
lambda: {"tokens": {"account_id": "acct_123"}},
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"agent.account_usage.httpx.Client",
|
||||
lambda timeout=15.0: _Client(
|
||||
{
|
||||
"plan_type": "pro",
|
||||
"rate_limit": {
|
||||
"primary_window": {
|
||||
"used_percent": 15,
|
||||
"reset_at": 1_900_000_000,
|
||||
"limit_window_seconds": 18000,
|
||||
},
|
||||
"secondary_window": {
|
||||
"used_percent": 40,
|
||||
"reset_at": 1_900_500_000,
|
||||
"limit_window_seconds": 604800,
|
||||
},
|
||||
},
|
||||
"credits": {"has_credits": True, "balance": 12.5},
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
snapshot = fetch_account_usage("openai-codex")
|
||||
|
||||
assert snapshot is not None
|
||||
assert snapshot.plan == "Pro"
|
||||
assert len(snapshot.windows) == 2
|
||||
assert snapshot.windows[0].label == "Session"
|
||||
assert snapshot.windows[0].used_percent == 15.0
|
||||
assert snapshot.windows[0].reset_at == datetime.fromtimestamp(1_900_000_000, tz=timezone.utc)
|
||||
assert "Credits balance: $12.50" in snapshot.details
|
||||
|
||||
|
||||
def test_render_account_usage_lines_includes_reset_and_provider():
|
||||
snapshot = AccountUsageSnapshot(
|
||||
provider="openai-codex",
|
||||
source="usage_api",
|
||||
fetched_at=datetime.now(timezone.utc),
|
||||
plan="Pro",
|
||||
windows=(
|
||||
AccountUsageWindow(
|
||||
label="Session",
|
||||
used_percent=25,
|
||||
reset_at=datetime.now(timezone.utc),
|
||||
),
|
||||
),
|
||||
details=("Credits balance: $9.99",),
|
||||
)
|
||||
lines = render_account_usage_lines(snapshot)
|
||||
|
||||
assert lines[0] == "📈 Account limits"
|
||||
assert "openai-codex (Pro)" in lines[1]
|
||||
assert "Session: 75% remaining (25% used)" in lines[2]
|
||||
assert "Credits balance: $9.99" in lines[3]
|
||||
|
||||
|
||||
def test_fetch_account_usage_openrouter_uses_limit_remaining_and_ignores_deprecated_rate_limit(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"agent.account_usage.resolve_runtime_provider",
|
||||
lambda requested, explicit_base_url=None, explicit_api_key=None: {
|
||||
"provider": "openrouter",
|
||||
"base_url": "https://openrouter.ai/api/v1",
|
||||
"api_key": "sk-test",
|
||||
},
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"agent.account_usage.httpx.Client",
|
||||
lambda timeout=10.0: _RoutingClient(
|
||||
{
|
||||
"https://openrouter.ai/api/v1/credits": {
|
||||
"data": {"total_credits": 300.0, "total_usage": 10.92}
|
||||
},
|
||||
"https://openrouter.ai/api/v1/key": {
|
||||
"data": {
|
||||
"limit": 100.0,
|
||||
"limit_remaining": 70.0,
|
||||
"limit_reset": "monthly",
|
||||
"usage": 12.5,
|
||||
"usage_daily": 0.5,
|
||||
"usage_weekly": 2.0,
|
||||
"usage_monthly": 8.0,
|
||||
"rate_limit": {"requests": -1, "interval": "10s"},
|
||||
}
|
||||
},
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
snapshot = fetch_account_usage("openrouter")
|
||||
|
||||
assert snapshot is not None
|
||||
assert snapshot.windows == (
|
||||
AccountUsageWindow(
|
||||
label="API key quota",
|
||||
used_percent=30.0,
|
||||
detail="$70.00 of $100.00 remaining • resets monthly",
|
||||
),
|
||||
)
|
||||
assert "Credits balance: $289.08" in snapshot.details
|
||||
assert "API key usage: $12.50 total • $0.50 today • $2.00 this week • $8.00 this month" in snapshot.details
|
||||
assert all("-1 requests / 10s" not in line for line in render_account_usage_lines(snapshot))
|
||||
|
||||
|
||||
def test_fetch_account_usage_openrouter_omits_quota_window_when_key_has_no_limit(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"agent.account_usage.resolve_runtime_provider",
|
||||
lambda requested, explicit_base_url=None, explicit_api_key=None: {
|
||||
"provider": "openrouter",
|
||||
"base_url": "https://openrouter.ai/api/v1",
|
||||
"api_key": "sk-test",
|
||||
},
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"agent.account_usage.httpx.Client",
|
||||
lambda timeout=10.0: _RoutingClient(
|
||||
{
|
||||
"https://openrouter.ai/api/v1/credits": {
|
||||
"data": {"total_credits": 100.0, "total_usage": 25.5}
|
||||
},
|
||||
"https://openrouter.ai/api/v1/key": {
|
||||
"data": {
|
||||
"limit": None,
|
||||
"limit_remaining": None,
|
||||
"usage": 25.5,
|
||||
"usage_daily": 1.25,
|
||||
"usage_weekly": 4.5,
|
||||
"usage_monthly": 18.0,
|
||||
}
|
||||
},
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
snapshot = fetch_account_usage("openrouter")
|
||||
|
||||
assert snapshot is not None
|
||||
assert snapshot.windows == ()
|
||||
assert "Credits balance: $74.50" in snapshot.details
|
||||
assert "API key usage: $25.50 total • $1.25 today • $4.50 this week • $18.00 this month" in snapshot.details
|
||||
@@ -0,0 +1,333 @@
|
||||
"""Unit tests for AIAgent pre/post-LLM-call guardrails.
|
||||
|
||||
Covers three static methods on AIAgent (inspired by PR #1321 — @alireza78a):
|
||||
- _sanitize_api_messages() — Phase 1: orphaned tool pair repair
|
||||
- _cap_delegate_task_calls() — Phase 2a: subagent concurrency limit
|
||||
- _deduplicate_tool_calls() — Phase 2b: identical call deduplication
|
||||
- _uniquify_tool_call_ids() — Phase 2c: duplicate-id repair (lossless pairing)
|
||||
"""
|
||||
|
||||
import types
|
||||
|
||||
import pytest
|
||||
|
||||
from run_agent import AIAgent
|
||||
|
||||
# Pin the concurrency limit instead of reading the runtime config.
|
||||
# _cap_delegate_task_calls() resolves _get_max_concurrent_children() at CALL
|
||||
# time (inside a per-test hermetic HERMES_HOME), but this module previously
|
||||
# froze the value at IMPORT time — before the hermetic fixture ran — so a
|
||||
# developer machine with delegation.max_concurrent_children in the real
|
||||
# ~/.hermes/config.yaml saw a different limit at import vs call and the
|
||||
# truncation tests failed locally while passing on CI.
|
||||
MAX_CONCURRENT_CHILDREN = 3
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _pin_max_concurrent_children(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"tools.delegate_tool._get_max_concurrent_children",
|
||||
lambda: MAX_CONCURRENT_CHILDREN,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def make_tc(name: str, arguments: str = "{}") -> types.SimpleNamespace:
|
||||
"""Create a minimal tool_call SimpleNamespace mirroring the OpenAI SDK object."""
|
||||
tc = types.SimpleNamespace()
|
||||
tc.function = types.SimpleNamespace(name=name, arguments=arguments)
|
||||
return tc
|
||||
|
||||
|
||||
def tool_result(call_id: str, content: str = "ok") -> dict:
|
||||
return {"role": "tool", "tool_call_id": call_id, "content": content}
|
||||
|
||||
|
||||
def assistant_dict_call(call_id: str, name: str = "terminal") -> dict:
|
||||
"""Dict-style tool_call (as stored in message history)."""
|
||||
return {"id": call_id, "function": {"name": name, "arguments": "{}"}}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Phase 1 — _sanitize_api_messages
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestSanitizeApiMessages:
|
||||
|
||||
def test_orphaned_result_removed(self):
|
||||
msgs = [
|
||||
{"role": "assistant", "tool_calls": [assistant_dict_call("c1")]},
|
||||
tool_result("c1"),
|
||||
tool_result("c_ORPHAN"),
|
||||
]
|
||||
out = AIAgent._sanitize_api_messages(msgs)
|
||||
assert len(out) == 2
|
||||
assert all(m.get("tool_call_id") != "c_ORPHAN" for m in out)
|
||||
|
||||
def test_orphaned_call_gets_stub_result(self):
|
||||
msgs = [
|
||||
{"role": "assistant", "tool_calls": [assistant_dict_call("c2")]},
|
||||
]
|
||||
out = AIAgent._sanitize_api_messages(msgs)
|
||||
assert len(out) == 2
|
||||
stub = out[1]
|
||||
assert stub["role"] == "tool"
|
||||
assert stub["tool_call_id"] == "c2"
|
||||
assert stub["content"]
|
||||
|
||||
def test_clean_messages_pass_through(self):
|
||||
msgs = [
|
||||
{"role": "user", "content": "hello"},
|
||||
{"role": "assistant", "tool_calls": [assistant_dict_call("c3")]},
|
||||
tool_result("c3"),
|
||||
{"role": "assistant", "content": "done"},
|
||||
]
|
||||
out = AIAgent._sanitize_api_messages(msgs)
|
||||
assert out == msgs
|
||||
|
||||
def test_mixed_orphaned_result_and_orphaned_call(self):
|
||||
msgs = [
|
||||
{"role": "assistant", "tool_calls": [
|
||||
assistant_dict_call("c4"),
|
||||
assistant_dict_call("c5"),
|
||||
]},
|
||||
tool_result("c4"),
|
||||
tool_result("c_DANGLING"),
|
||||
]
|
||||
out = AIAgent._sanitize_api_messages(msgs)
|
||||
ids = [m.get("tool_call_id") for m in out if m.get("role") == "tool"]
|
||||
assert "c_DANGLING" not in ids
|
||||
assert "c4" in ids
|
||||
assert "c5" in ids
|
||||
|
||||
def test_empty_list_is_safe(self):
|
||||
assert AIAgent._sanitize_api_messages([]) == []
|
||||
|
||||
|
||||
def test_sdk_object_tool_calls(self):
|
||||
tc_obj = types.SimpleNamespace(id="c6", function=types.SimpleNamespace(
|
||||
name="terminal", arguments="{}"
|
||||
))
|
||||
msgs = [
|
||||
{"role": "assistant", "tool_calls": [tc_obj]},
|
||||
]
|
||||
out = AIAgent._sanitize_api_messages(msgs)
|
||||
assert len(out) == 2
|
||||
assert out[1]["tool_call_id"] == "c6"
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Phase 2a — _cap_delegate_task_calls
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestCapDelegateTaskCalls:
|
||||
|
||||
def test_excess_delegates_truncated(self):
|
||||
tcs = [make_tc("delegate_task") for _ in range(MAX_CONCURRENT_CHILDREN + 2)]
|
||||
out = AIAgent._cap_delegate_task_calls(tcs)
|
||||
delegate_count = sum(1 for tc in out if tc.function.name == "delegate_task")
|
||||
assert delegate_count == MAX_CONCURRENT_CHILDREN
|
||||
|
||||
def test_non_delegate_calls_preserved(self):
|
||||
tcs = (
|
||||
[make_tc("delegate_task") for _ in range(MAX_CONCURRENT_CHILDREN + 1)]
|
||||
+ [make_tc("terminal"), make_tc("web_search")]
|
||||
)
|
||||
out = AIAgent._cap_delegate_task_calls(tcs)
|
||||
names = [tc.function.name for tc in out]
|
||||
assert "terminal" in names
|
||||
assert "web_search" in names
|
||||
|
||||
def test_at_limit_passes_through(self):
|
||||
tcs = [make_tc("delegate_task") for _ in range(MAX_CONCURRENT_CHILDREN)]
|
||||
out = AIAgent._cap_delegate_task_calls(tcs)
|
||||
assert out is tcs
|
||||
|
||||
|
||||
|
||||
def test_empty_list_safe(self):
|
||||
assert AIAgent._cap_delegate_task_calls([]) == []
|
||||
|
||||
|
||||
def test_interleaved_order_preserved(self):
|
||||
delegates = [make_tc("delegate_task", f'{{"task":"{i}"}}')
|
||||
for i in range(MAX_CONCURRENT_CHILDREN + 1)]
|
||||
t1 = make_tc("terminal", '{"cmd":"ls"}')
|
||||
w1 = make_tc("web_search", '{"q":"x"}')
|
||||
tcs = [delegates[0], t1, delegates[1], w1] + delegates[2:]
|
||||
out = AIAgent._cap_delegate_task_calls(tcs)
|
||||
expected = [delegates[0], t1, delegates[1], w1] + delegates[2:MAX_CONCURRENT_CHILDREN]
|
||||
assert len(out) == len(expected)
|
||||
for i, (actual, exp) in enumerate(zip(out, expected)):
|
||||
assert actual is exp, f"mismatch at index {i}"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Phase 2b — _deduplicate_tool_calls
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestDeduplicateToolCalls:
|
||||
|
||||
def test_duplicate_pair_deduplicated(self):
|
||||
tcs = [
|
||||
make_tc("web_search", '{"query":"foo"}'),
|
||||
make_tc("web_search", '{"query":"foo"}'),
|
||||
]
|
||||
out = AIAgent._deduplicate_tool_calls(tcs)
|
||||
assert len(out) == 1
|
||||
|
||||
def test_duplicate_json_objects_with_reordered_keys_deduplicated(self):
|
||||
first = make_tc(
|
||||
"terminal",
|
||||
'{"command":"printf hello >> out.log","timeout":10}',
|
||||
)
|
||||
second = make_tc(
|
||||
"terminal",
|
||||
'{"timeout":10,"command":"printf hello >> out.log"}',
|
||||
)
|
||||
|
||||
out = AIAgent._deduplicate_tool_calls([first, second])
|
||||
|
||||
assert out == [first]
|
||||
|
||||
def test_distinct_json_arguments_are_preserved(self):
|
||||
first = make_tc("terminal", '{"command":"one","timeout":10}')
|
||||
second = make_tc("terminal", '{"timeout":10,"command":"two"}')
|
||||
|
||||
out = AIAgent._deduplicate_tool_calls([first, second])
|
||||
|
||||
assert out == [first, second]
|
||||
|
||||
def test_malformed_arguments_use_raw_string_for_deduplication(self):
|
||||
first = make_tc("terminal", '{"command":"one"')
|
||||
duplicate = make_tc("terminal", '{"command":"one"')
|
||||
distinct = make_tc("terminal", '{ "command":"one"')
|
||||
|
||||
out = AIAgent._deduplicate_tool_calls([first, duplicate, distinct])
|
||||
|
||||
assert out == [first, distinct]
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_empty_list_safe(self):
|
||||
assert AIAgent._deduplicate_tool_calls([]) == []
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Phase 2c — _uniquify_tool_call_ids
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def make_tc_id(id_: str, name: str, arguments: str = "{}", call_id=None) -> types.SimpleNamespace:
|
||||
tc = types.SimpleNamespace()
|
||||
tc.id = id_
|
||||
if call_id is not None:
|
||||
tc.call_id = call_id
|
||||
tc.function = types.SimpleNamespace(name=name, arguments=arguments)
|
||||
return tc
|
||||
|
||||
|
||||
class TestUniquifyToolCallIds:
|
||||
|
||||
def test_distinct_calls_sharing_id_get_unique_ids(self):
|
||||
tcs = [
|
||||
make_tc_id("call_1", "read_file", '{"path":"a.txt"}'),
|
||||
make_tc_id("call_1", "read_file", '{"path":"b.txt"}'),
|
||||
]
|
||||
out = AIAgent._uniquify_tool_call_ids(tcs)
|
||||
assert out is tcs
|
||||
assert tcs[0].id == "call_1" # first occurrence untouched
|
||||
assert tcs[1].id == "call_1_d2" # deterministic suffix
|
||||
assert len({tc.id for tc in tcs}) == 2
|
||||
|
||||
def test_three_way_collision(self):
|
||||
tcs = [make_tc_id("x", "t", '{"a":1}'),
|
||||
make_tc_id("x", "t", '{"a":2}'),
|
||||
make_tc_id("x", "t", '{"a":3}')]
|
||||
AIAgent._uniquify_tool_call_ids(tcs)
|
||||
assert [tc.id for tc in tcs] == ["x", "x_d2", "x_d3"]
|
||||
|
||||
|
||||
|
||||
def test_blank_and_missing_ids_left_for_fallback(self):
|
||||
tc_blank = make_tc_id("", "t1")
|
||||
tc_none = make_tc_id(None, "t2")
|
||||
tc_noattr = make_tc("t3") # no .id attribute at all
|
||||
AIAgent._uniquify_tool_call_ids([tc_blank, tc_none, tc_noattr])
|
||||
assert tc_blank.id == ""
|
||||
assert tc_none.id is None
|
||||
assert not hasattr(tc_noattr, "id")
|
||||
|
||||
def test_call_id_sibling_kept_consistent(self):
|
||||
# Responses-path objects carry call_id; build_assistant_message
|
||||
# prefers it, so the rename must update both.
|
||||
tcs = [make_tc_id("call_1", "t", '{"a":1}', call_id="call_1"),
|
||||
make_tc_id("call_1", "t", '{"a":2}', call_id="call_1")]
|
||||
AIAgent._uniquify_tool_call_ids(tcs)
|
||||
assert tcs[1].id == "call_1_d2"
|
||||
assert tcs[1].call_id == "call_1_d2"
|
||||
assert tcs[0].call_id == "call_1"
|
||||
|
||||
def test_dict_entries_supported(self):
|
||||
tcs = [{"id": "c", "function": {"name": "t", "arguments": '{"a":1}'}},
|
||||
{"id": "c", "call_id": "c", "function": {"name": "t", "arguments": '{"a":2}'}}]
|
||||
AIAgent._uniquify_tool_call_ids(tcs)
|
||||
assert tcs[0]["id"] == "c"
|
||||
assert tcs[1]["id"] == "c_d2"
|
||||
assert tcs[1]["call_id"] == "c_d2"
|
||||
|
||||
def test_empty_and_none_safe(self):
|
||||
assert AIAgent._uniquify_tool_call_ids([]) == []
|
||||
assert AIAgent._uniquify_tool_call_ids(None) is None
|
||||
|
||||
def test_composite_responses_ids_collide_on_call_half(self):
|
||||
# "call_x|fc_y" composites pair on the call half; the rename must
|
||||
# keep the provider's response-item half intact.
|
||||
tcs = [make_tc_id("call_x|fc_1", "t", '{"a":1}'),
|
||||
make_tc_id("call_x|fc_2", "t", '{"a":2}')]
|
||||
AIAgent._uniquify_tool_call_ids(tcs)
|
||||
assert tcs[0].id == "call_x|fc_1"
|
||||
assert tcs[1].id == "call_x_d2|fc_2"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _get_tool_call_id_static
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestGetToolCallIdStatic:
|
||||
|
||||
def test_dict_with_valid_id(self):
|
||||
assert AIAgent._get_tool_call_id_static({"id": "call_123"}) == "call_123"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _get_tool_call_name_static
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestGetToolCallNameStatic:
|
||||
|
||||
def test_dict_with_valid_name(self):
|
||||
assert AIAgent._get_tool_call_name_static(
|
||||
{"id": "call_1", "function": {"name": "terminal", "arguments": "{}"}}
|
||||
) == "terminal"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_object_without_function_attr(self):
|
||||
tc = types.SimpleNamespace(id="call_1")
|
||||
assert AIAgent._get_tool_call_name_static(tc) == ""
|
||||
@@ -0,0 +1,119 @@
|
||||
"""The Anthropic streaming path must not accept a tool_use block whose stream
|
||||
died before its input arrived (#80498 sibling).
|
||||
|
||||
The chat_completions accumulator flags zero-byte tool-call arguments on a
|
||||
clean no-finish_reason stream end and routes them through the
|
||||
partial-stream-stub retry path. The Anthropic path had the same gap in a
|
||||
different shape: a clean SSE close after ``content_block_start(tool_use)``
|
||||
but before any ``input_json_delta``/``message_delta`` yields an SDK
|
||||
final-message snapshot whose content is NON-empty (the tool_use block is
|
||||
there, ``input={}``) and whose ``stop_reason`` is None — which sailed past
|
||||
both empty-stream guards and executed the tool with empty input, no retry.
|
||||
|
||||
The fix raises EmptyStreamError for a tool_use-bearing message with no
|
||||
stop_reason, riding the same bounded stream-retry the eventless case uses.
|
||||
"""
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _make_anthropic_agent(**kwargs):
|
||||
from run_agent import AIAgent
|
||||
|
||||
defaults = dict(
|
||||
api_key="test-key",
|
||||
base_url="https://example.com/v1",
|
||||
model="claude-opus-4-7",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
defaults.update(kwargs)
|
||||
agent = AIAgent(**defaults)
|
||||
agent.api_mode = "anthropic_messages"
|
||||
agent._anthropic_client = MagicMock()
|
||||
agent._anthropic_api_key = "test-anthropic-key"
|
||||
agent._create_request_anthropic_client = lambda *a, **k: agent._anthropic_client
|
||||
return agent
|
||||
|
||||
|
||||
def _stream_cm(final_message, events=()):
|
||||
cm = MagicMock()
|
||||
stream = MagicMock()
|
||||
stream.__iter__ = MagicMock(return_value=iter(list(events)))
|
||||
stream.get_final_message = MagicMock(return_value=final_message)
|
||||
cm.__enter__ = MagicMock(return_value=stream)
|
||||
cm.__exit__ = MagicMock(return_value=False)
|
||||
return cm
|
||||
|
||||
|
||||
def _tool_use_block(name="write_file", input_obj=None):
|
||||
return SimpleNamespace(
|
||||
type="tool_use",
|
||||
id="toolu_x",
|
||||
name=name,
|
||||
input=input_obj if input_obj is not None else {},
|
||||
)
|
||||
|
||||
|
||||
def _tool_use_start_event(name="write_file"):
|
||||
return SimpleNamespace(
|
||||
type="content_block_start",
|
||||
content_block=SimpleNamespace(type="tool_use", name=name),
|
||||
)
|
||||
|
||||
|
||||
class TestAnthropicMidToolCallStreamDrop:
|
||||
def test_tool_use_without_stop_reason_raises_empty_stream(self):
|
||||
"""The #80498-sibling shape: tool_use block present, stop_reason None."""
|
||||
from agent.chat_completion_helpers import EmptyStreamError
|
||||
|
||||
dropped = MagicMock()
|
||||
dropped.content = [_tool_use_block()]
|
||||
dropped.stop_reason = None
|
||||
dropped.usage = SimpleNamespace(input_tokens=10, output_tokens=2)
|
||||
|
||||
agent = _make_anthropic_agent()
|
||||
agent._anthropic_client.messages.stream = MagicMock(
|
||||
return_value=_stream_cm(dropped, events=[_tool_use_start_event()])
|
||||
)
|
||||
|
||||
with pytest.raises(EmptyStreamError, match="tool_use"):
|
||||
agent._interruptible_streaming_api_call({"model": "claude-opus-4-7"})
|
||||
|
||||
def test_completed_tool_use_with_stop_reason_passes(self):
|
||||
"""A legitimate tool_use completion (stop_reason set) is untouched."""
|
||||
done = MagicMock()
|
||||
done.content = [_tool_use_block(input_obj={"path": "a.txt"})]
|
||||
done.stop_reason = "tool_use"
|
||||
done.usage = SimpleNamespace(input_tokens=10, output_tokens=5)
|
||||
|
||||
agent = _make_anthropic_agent()
|
||||
agent._anthropic_client.messages.stream = MagicMock(
|
||||
return_value=_stream_cm(done, events=[_tool_use_start_event()])
|
||||
)
|
||||
|
||||
response = agent._interruptible_streaming_api_call(
|
||||
{"model": "claude-opus-4-7"}
|
||||
)
|
||||
assert response is done
|
||||
|
||||
def test_text_only_message_without_stop_reason_passes(self):
|
||||
"""No tool_use block -> the new gate stays out of the way (text-only
|
||||
no-stop_reason handling keeps its pre-existing behavior)."""
|
||||
text_only = MagicMock()
|
||||
text_only.content = [SimpleNamespace(type="text", text="partial answer")]
|
||||
text_only.stop_reason = None
|
||||
text_only.usage = SimpleNamespace(input_tokens=10, output_tokens=5)
|
||||
|
||||
agent = _make_anthropic_agent()
|
||||
agent._anthropic_client.messages.stream = MagicMock(
|
||||
return_value=_stream_cm(text_only)
|
||||
)
|
||||
|
||||
response = agent._interruptible_streaming_api_call(
|
||||
{"model": "claude-opus-4-7"}
|
||||
)
|
||||
assert response is text_only
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,55 @@
|
||||
"""Anthropic Messages path must capture rate-limit + credits headers.
|
||||
|
||||
Portal Claude moved off /chat/completions onto /v1/messages. The OpenAI-wire
|
||||
streaming path captured both header families from ``stream.response``; the
|
||||
Messages path must do the same via ``on_response`` or /status and the Nous
|
||||
429 classifier lose last-known state.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from run_agent import AIAgent
|
||||
|
||||
|
||||
def test_capture_anthropic_response_headers_forwards_to_both_captures():
|
||||
agent = object.__new__(AIAgent)
|
||||
agent._capture_rate_limits = MagicMock()
|
||||
agent._capture_credits = MagicMock()
|
||||
|
||||
response = SimpleNamespace(headers={"x-ratelimit-remaining-requests": "10"})
|
||||
agent._capture_anthropic_response_headers(response)
|
||||
|
||||
agent._capture_rate_limits.assert_called_once_with(response)
|
||||
agent._capture_credits.assert_called_once_with(response)
|
||||
|
||||
|
||||
def test_anthropic_messages_create_passes_combined_header_callback():
|
||||
agent = object.__new__(AIAgent)
|
||||
agent.api_mode = "anthropic_messages"
|
||||
agent._anthropic_client = MagicMock()
|
||||
agent._disable_streaming = False
|
||||
agent.log_prefix = ""
|
||||
agent._try_refresh_anthropic_client_credentials = MagicMock(return_value=False)
|
||||
agent._capture_anthropic_response_headers = MagicMock()
|
||||
|
||||
captured = {}
|
||||
|
||||
def _fake_create(client, api_kwargs, **kwargs):
|
||||
captured.update(kwargs)
|
||||
return SimpleNamespace(content=[], stop_reason="end_turn")
|
||||
|
||||
import agent.anthropic_adapter as adapter
|
||||
|
||||
original = adapter.create_anthropic_message
|
||||
adapter.create_anthropic_message = _fake_create
|
||||
try:
|
||||
agent._anthropic_messages_create({"model": "anthropic/claude-opus-4.8"})
|
||||
finally:
|
||||
adapter.create_anthropic_message = original
|
||||
|
||||
assert (
|
||||
captured.get("on_response") is agent._capture_anthropic_response_headers
|
||||
)
|
||||
@@ -0,0 +1,122 @@
|
||||
"""Anthropic stream cleanup must not call _replace_primary_openai_client() and
|
||||
must not hang on Anthropic-native configs (#28161), now via the request-local
|
||||
client model (#67142).
|
||||
|
||||
Originally three cleanup sites in interruptible_streaming_api_call() called
|
||||
_replace_primary_openai_client() unconditionally; for api_mode=anthropic_messages
|
||||
that silently failed (no OPENAI_API_KEY) and left the in-flight httpx stream
|
||||
unclosed, blocking the worker until the 900s read-timeout fired.
|
||||
|
||||
Since #67142, anthropic streams run on a per-request client: the stale/retry
|
||||
cleanup closes the *request-local* client (worker-owned) and builds a fresh one
|
||||
next attempt — the shared _anthropic_client is never closed/rebuilt from inside
|
||||
a request (that poll-thread close was the TLS-FD→SQLite corruption vector). The
|
||||
no-hang guarantee is preserved because the poll thread aborts the request
|
||||
client's sockets, which unblocks the worker.
|
||||
|
||||
Tests cover:
|
||||
- stream_retry cleanup (connection error on fresh stream)
|
||||
- stale_stream cleanup (outer poll loop detects stale stream)
|
||||
|
||||
Fixes #28161. Extends #67142.
|
||||
"""
|
||||
import threading
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_anthropic_agent(**kwargs):
|
||||
from run_agent import AIAgent
|
||||
|
||||
defaults = dict(
|
||||
api_key="test-key",
|
||||
base_url="https://example.com/v1",
|
||||
model="claude-opus-4-7",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
defaults.update(kwargs)
|
||||
agent = AIAgent(**defaults)
|
||||
agent.api_mode = "anthropic_messages"
|
||||
agent._anthropic_client = MagicMock()
|
||||
agent._anthropic_api_key = "test-anthropic-key"
|
||||
# #67142: anthropic streams now run on a request-local client; route it to
|
||||
# the test mock so .messages.stream is exercised and its cleanup observed.
|
||||
agent._create_request_anthropic_client = lambda *a, **k: agent._anthropic_client
|
||||
return agent
|
||||
|
||||
|
||||
def _good_stream_cm():
|
||||
"""Context manager whose stream yields no events and returns a valid message."""
|
||||
cm = MagicMock()
|
||||
stream = MagicMock()
|
||||
stream.__iter__ = MagicMock(return_value=iter([]))
|
||||
msg = MagicMock()
|
||||
msg.content = []
|
||||
msg.stop_reason = "end_turn"
|
||||
msg.usage = SimpleNamespace(input_tokens=10, output_tokens=5)
|
||||
stream.get_final_message = MagicMock(return_value=msg)
|
||||
cm.__enter__ = MagicMock(return_value=stream)
|
||||
cm.__exit__ = MagicMock(return_value=False)
|
||||
return cm
|
||||
|
||||
|
||||
def _failing_stream_cm():
|
||||
"""Context manager whose __enter__ raises ConnectError immediately."""
|
||||
cm = MagicMock()
|
||||
cm.__enter__ = MagicMock(
|
||||
side_effect=httpx.ConnectError("connection reset by peer")
|
||||
)
|
||||
return cm
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestAnthropicStreamPoolCleanup:
|
||||
"""Anthropic cleanup must never touch the OpenAI primary or the shared
|
||||
Anthropic client, and must not hang (#28161 / #67142)."""
|
||||
|
||||
@pytest.mark.filterwarnings(
|
||||
"ignore::pytest.PytestUnhandledThreadExceptionWarning"
|
||||
)
|
||||
def test_stream_retry_closes_request_client_not_openai(self):
|
||||
"""Connection error during stream retry → close the request-local
|
||||
Anthropic client (worker-owned) and retry; never rebuild the shared
|
||||
Anthropic client, never touch the OpenAI primary."""
|
||||
agent = _make_anthropic_agent()
|
||||
|
||||
attempt_count = [0]
|
||||
|
||||
def _stream_side_effect(*args, **kwargs):
|
||||
attempt_count[0] += 1
|
||||
if attempt_count[0] == 1:
|
||||
return _failing_stream_cm()
|
||||
return _good_stream_cm()
|
||||
|
||||
agent._anthropic_client.messages.stream.side_effect = _stream_side_effect
|
||||
|
||||
with patch.object(agent, "_rebuild_anthropic_client") as mock_rebuild:
|
||||
with patch.object(
|
||||
agent, "_replace_primary_openai_client"
|
||||
) as mock_replace:
|
||||
agent._interruptible_streaming_api_call({})
|
||||
|
||||
mock_replace.assert_not_called()
|
||||
# #67142: the shared client is never rebuilt from inside a request; the
|
||||
# request-local client (routed to this mock) is closed instead.
|
||||
mock_rebuild.assert_not_called()
|
||||
agent._anthropic_client.close.assert_called()
|
||||
assert attempt_count[0] == 2 # retried once, then succeeded
|
||||
|
||||
@@ -0,0 +1,158 @@
|
||||
"""Tests for ``_is_anthropic_oauth`` guard against third-party Anthropic-compatible providers.
|
||||
|
||||
The invariant: ``self._is_anthropic_oauth`` must only ever be True when
|
||||
``self.provider == 'anthropic'`` (native Anthropic). Third-party providers
|
||||
that speak the Anthropic protocol (MiniMax, Zhipu GLM, Alibaba DashScope,
|
||||
Kimi, LiteLLM proxies, etc.) must never trip OAuth code paths — doing so
|
||||
injects Claude-Code identity headers and system prompts that cause
|
||||
401/403 from those endpoints.
|
||||
|
||||
This test class covers all FIVE sites that assign ``_is_anthropic_oauth``:
|
||||
|
||||
1. ``AIAgent.__init__`` (line ~1022)
|
||||
2. ``AIAgent.switch_model`` (line ~1832)
|
||||
3. ``AIAgent._try_refresh_anthropic_client_credentials`` (line ~5335)
|
||||
4. ``AIAgent._swap_credential`` (line ~5378)
|
||||
5. ``AIAgent._try_activate_fallback`` (line ~6536)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from run_agent import AIAgent
|
||||
|
||||
|
||||
# A plausible-looking OAuth token (``sk-ant-`` without the ``-api`` suffix).
|
||||
_OAUTH_LIKE_TOKEN = "sk-ant-oauth-example-1234567890abcdef"
|
||||
_API_KEY_TOKEN = "sk-ant-api-abcdef1234567890"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def agent():
|
||||
"""Minimal AIAgent construction, skipping tool discovery."""
|
||||
with (
|
||||
patch("model_tools.get_tool_definitions", return_value=[]),
|
||||
patch("model_tools.check_toolset_requirements", return_value={}),
|
||||
patch("agent.process_bootstrap.OpenAI"),
|
||||
):
|
||||
a = AIAgent(
|
||||
api_key="test-key-1234567890",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
a.client = MagicMock()
|
||||
return a
|
||||
|
||||
|
||||
class TestOAuthFlagOnRefresh:
|
||||
"""Site 3 — _try_refresh_anthropic_client_credentials."""
|
||||
|
||||
def test_third_party_provider_refresh_is_noop(self, agent):
|
||||
"""Refresh path returns False immediately when provider != anthropic — the
|
||||
OAuth flag can never be mutated for third-party providers. Double-defended
|
||||
by the per-assignment guard at line ~5393 so future refactors can't
|
||||
reintroduce the bug."""
|
||||
agent.api_mode = "anthropic_messages"
|
||||
agent.provider = "minimax" # ← third-party
|
||||
agent._anthropic_api_key = "***"
|
||||
agent._anthropic_client = MagicMock()
|
||||
agent._is_anthropic_oauth = False
|
||||
|
||||
with (
|
||||
patch("agent.anthropic_credentials.resolve_anthropic_token",
|
||||
return_value=_OAUTH_LIKE_TOKEN),
|
||||
patch("agent.anthropic_adapter.build_anthropic_client",
|
||||
return_value=MagicMock()),
|
||||
):
|
||||
result = agent._try_refresh_anthropic_client_credentials()
|
||||
|
||||
# The function short-circuits on non-anthropic providers.
|
||||
assert result is False
|
||||
# And the flag is untouched regardless.
|
||||
assert agent._is_anthropic_oauth is False
|
||||
|
||||
|
||||
|
||||
class TestOAuthFlagOnCredentialSwap:
|
||||
"""Site 4 — _swap_credential (credential pool rotation)."""
|
||||
|
||||
def test_pool_swap_on_third_party_never_flips_oauth(self, agent):
|
||||
agent.api_mode = "anthropic_messages"
|
||||
agent.provider = "glm" # ← Zhipu GLM via /anthropic
|
||||
agent._anthropic_api_key = "old-key"
|
||||
agent._anthropic_base_url = "https://open.bigmodel.cn/api/anthropic"
|
||||
agent._anthropic_client = MagicMock()
|
||||
agent._is_anthropic_oauth = False
|
||||
|
||||
entry = MagicMock()
|
||||
entry.runtime_api_key = _OAUTH_LIKE_TOKEN
|
||||
entry.runtime_base_url = "https://open.bigmodel.cn/api/anthropic"
|
||||
|
||||
with patch("agent.anthropic_adapter.build_anthropic_client",
|
||||
return_value=MagicMock()):
|
||||
agent._swap_credential(entry)
|
||||
|
||||
assert agent._is_anthropic_oauth is False
|
||||
|
||||
|
||||
class TestOAuthFlagOnConstruction:
|
||||
"""Site 1 — AIAgent.__init__ on a third-party anthropic_messages provider."""
|
||||
|
||||
def test_minimax_init_does_not_flip_oauth(self):
|
||||
with (
|
||||
patch("model_tools.get_tool_definitions", return_value=[]),
|
||||
patch("model_tools.check_toolset_requirements", return_value={}),
|
||||
patch("agent.anthropic_adapter.build_anthropic_client",
|
||||
return_value=MagicMock()),
|
||||
# Simulate a stale ANTHROPIC_TOKEN in the env — the init code
|
||||
# MUST NOT fall back to it when provider != anthropic.
|
||||
patch("agent.anthropic_credentials.resolve_anthropic_token",
|
||||
return_value=_OAUTH_LIKE_TOKEN),
|
||||
):
|
||||
agent = AIAgent(
|
||||
api_key="minimax-key-1234",
|
||||
base_url="https://api.minimax.io/anthropic",
|
||||
provider="minimax",
|
||||
api_mode="anthropic_messages",
|
||||
model="claude-sonnet-4-6",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
|
||||
# The effective key should be the explicit minimax-key, not the
|
||||
# stale Anthropic OAuth token, and the OAuth flag must be False.
|
||||
assert agent._anthropic_api_key == "minimax-key-1234"
|
||||
assert agent._is_anthropic_oauth is False
|
||||
|
||||
|
||||
class TestOAuthFlagOnFallbackActivation:
|
||||
"""Site 5 — _try_activate_fallback targeting a third-party Anthropic endpoint."""
|
||||
|
||||
def test_fallback_to_third_party_does_not_flip_oauth(self, agent):
|
||||
"""Directly mimic the post-fallback assignment at line ~6537."""
|
||||
from agent.anthropic_credentials import _is_oauth_token
|
||||
|
||||
# Emulate the relevant lines of _try_activate_fallback without
|
||||
# running the entire recovery stack (which pulls in streaming,
|
||||
# sessions, etc.).
|
||||
fb_provider = "minimax"
|
||||
effective_key = _OAUTH_LIKE_TOKEN
|
||||
agent._is_anthropic_oauth = (
|
||||
_is_oauth_token(effective_key) if fb_provider == "anthropic" else False
|
||||
)
|
||||
assert agent._is_anthropic_oauth is False
|
||||
|
||||
|
||||
class TestApiKeyTokensAlwaysSafe:
|
||||
"""Regression: plain API-key shapes must always resolve to non-OAuth, any provider."""
|
||||
|
||||
def test_native_anthropic_with_api_key_token(self):
|
||||
from agent.anthropic_credentials import _is_oauth_token
|
||||
assert _is_oauth_token(_API_KEY_TOKEN) is False
|
||||
|
||||
@@ -0,0 +1,97 @@
|
||||
"""Regression test for anthropic_messages truncation continuation.
|
||||
|
||||
When an Anthropic response hits ``stop_reason: max_tokens`` (mapped to
|
||||
``finish_reason == 'length'`` in run_agent), the agent must retry with
|
||||
a continuation prompt — the same behavior it has always had for
|
||||
chat_completions and bedrock_converse. Before this PR, the
|
||||
``if self.api_mode in ('chat_completions', 'bedrock_converse'):`` guard
|
||||
silently dropped Anthropic-wire truncations on the floor, returning a
|
||||
half-finished response with no retry.
|
||||
|
||||
We don't exercise the full agent loop here (it's 3000 lines of inference,
|
||||
streaming, plugin hooks, etc.) — instead we verify the normalization
|
||||
adapter produces exactly the shape the continuation block now consumes.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _make_anthropic_text_block(text: str) -> SimpleNamespace:
|
||||
return SimpleNamespace(type="text", text=text)
|
||||
|
||||
|
||||
def _make_anthropic_tool_use_block(name: str = "my_tool") -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
type="tool_use",
|
||||
id="toolu_01",
|
||||
name=name,
|
||||
input={"foo": "bar"},
|
||||
)
|
||||
|
||||
|
||||
def _make_anthropic_response(blocks, stop_reason: str = "max_tokens"):
|
||||
return SimpleNamespace(
|
||||
id="msg_01",
|
||||
type="message",
|
||||
role="assistant",
|
||||
model="claude-sonnet-4-6",
|
||||
content=blocks,
|
||||
stop_reason=stop_reason,
|
||||
stop_sequence=None,
|
||||
usage=SimpleNamespace(input_tokens=100, output_tokens=200),
|
||||
)
|
||||
|
||||
|
||||
class TestTruncatedAnthropicResponseNormalization:
|
||||
"""AnthropicTransport.normalize_response() gives us the shape _build_assistant_message expects."""
|
||||
|
||||
def test_text_only_truncation_produces_text_content_no_tool_calls(self):
|
||||
"""Pure-text Anthropic truncation → continuation path should fire."""
|
||||
from agent.transports import get_transport
|
||||
|
||||
response = _make_anthropic_response(
|
||||
[_make_anthropic_text_block("partial response that was cut off")]
|
||||
)
|
||||
nr = get_transport("anthropic_messages").normalize_response(response)
|
||||
|
||||
# The continuation block checks these two attributes:
|
||||
# assistant_message.content → appended to truncated_response_parts
|
||||
# assistant_message.tool_calls → guards the text-retry branch
|
||||
assert nr.content is not None
|
||||
assert "partial response" in nr.content
|
||||
assert not nr.tool_calls, (
|
||||
"Pure-text truncation must have no tool_calls so the text-continuation "
|
||||
"branch (not the tool-retry branch) fires"
|
||||
)
|
||||
assert nr.finish_reason == "length", "max_tokens stop_reason must map to OpenAI-style 'length'"
|
||||
|
||||
|
||||
def test_empty_content_does_not_crash(self):
|
||||
"""Empty response.content — defensive: treat as a truncation with no text."""
|
||||
from agent.transports import get_transport
|
||||
|
||||
response = _make_anthropic_response([])
|
||||
nr = get_transport("anthropic_messages").normalize_response(response)
|
||||
# Depending on the adapter, content may be "" or None — both are
|
||||
# acceptable; what matters is no exception.
|
||||
assert nr is not None
|
||||
assert not nr.tool_calls
|
||||
|
||||
|
||||
class TestContinuationLogicBranching:
|
||||
"""Symbolic check that the api_mode gate now includes anthropic_messages."""
|
||||
|
||||
@pytest.mark.parametrize("api_mode", ["chat_completions", "bedrock_converse", "anthropic_messages"])
|
||||
def test_all_three_api_modes_hit_continuation_branch(self, api_mode):
|
||||
# The guard in run_agent.py is:
|
||||
# if self.api_mode in ("chat_completions", "bedrock_converse", "anthropic_messages"):
|
||||
assert api_mode in {"chat_completions", "bedrock_converse", "anthropic_messages"}
|
||||
|
||||
def test_codex_responses_still_excluded(self):
|
||||
# codex_responses has its own truncation path (not continuation-based)
|
||||
# and should NOT be routed through the shared block.
|
||||
assert "codex_responses" not in {"chat_completions", "bedrock_converse", "anthropic_messages"}
|
||||
@@ -863,22 +863,23 @@ class TestMaxIterationsSummaryReplay:
|
||||
class _Completions:
|
||||
def create(self, **kwargs):
|
||||
captured.update(kwargs)
|
||||
return "RAW-RESPONSE"
|
||||
return types.SimpleNamespace(
|
||||
choices=[types.SimpleNamespace(
|
||||
message=types.SimpleNamespace(content="SUMMARY", tool_calls=None),
|
||||
finish_reason="stop",
|
||||
)],
|
||||
)
|
||||
|
||||
client = types.SimpleNamespace(
|
||||
chat=types.SimpleNamespace(completions=_Completions())
|
||||
)
|
||||
transport = types.SimpleNamespace(
|
||||
normalize_response=lambda _r: types.SimpleNamespace(content="SUMMARY")
|
||||
)
|
||||
|
||||
messages = [
|
||||
{"role": "user", "content": "q1", "api_content": "q1\n\nPLUGIN-CTX"},
|
||||
{"role": "assistant", "content": "a1"},
|
||||
]
|
||||
with patch.object(
|
||||
agent, "_ensure_primary_openai_client", return_value=client
|
||||
), patch.object(agent, "_get_transport", return_value=transport):
|
||||
):
|
||||
out = handle_max_iterations(agent, messages, 5)
|
||||
|
||||
assert out == "SUMMARY"
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
"""Tests for agent.api_max_retries config surface.
|
||||
|
||||
Closes #11616 — make the hardcoded ``max_retries = 3`` in the agent's API
|
||||
retry loop user-configurable so fallback-provider setups can fail over
|
||||
faster on flaky primaries instead of burning ~3x180s on the same stall.
|
||||
"""
|
||||
from unittest.mock import patch
|
||||
|
||||
from run_agent import AIAgent
|
||||
|
||||
|
||||
def _make_agent(api_max_retries=None):
|
||||
"""Build an AIAgent with a mocked config.load_config that returns a
|
||||
config tree containing the given agent.api_max_retries (or default)."""
|
||||
cfg = {"agent": {}}
|
||||
if api_max_retries is not None:
|
||||
cfg["agent"]["api_max_retries"] = api_max_retries
|
||||
|
||||
with patch("agent.process_bootstrap.OpenAI"), \
|
||||
patch("hermes_cli.config.load_config", return_value=cfg), \
|
||||
patch("hermes_cli.config.load_config_readonly", return_value=cfg):
|
||||
return AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
|
||||
|
||||
def test_default_api_max_retries_is_three():
|
||||
"""No config override → legacy default of 3 retries preserved."""
|
||||
agent = _make_agent()
|
||||
assert agent._api_max_retries == 3
|
||||
|
||||
|
||||
def test_api_max_retries_honors_config_override():
|
||||
"""Setting agent.api_max_retries in config propagates to the agent."""
|
||||
agent = _make_agent(api_max_retries=1)
|
||||
assert agent._api_max_retries == 1
|
||||
|
||||
agent2 = _make_agent(api_max_retries=5)
|
||||
assert agent2._api_max_retries == 5
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,383 @@
|
||||
"""Tests for the AsyncHttpxClientWrapper.__del__ neuter fix.
|
||||
|
||||
The OpenAI SDK's ``AsyncHttpxClientWrapper.__del__`` schedules
|
||||
``aclose()`` via ``asyncio.get_running_loop().create_task()``. When GC
|
||||
fires during CLI idle time, prompt_toolkit's event loop picks up the task
|
||||
and crashes with "Event loop is closed" because the underlying TCP
|
||||
transport is bound to a dead worker loop.
|
||||
|
||||
The three-layer defence:
|
||||
1. ``neuter_async_httpx_del()`` replaces ``__del__`` with a no-op.
|
||||
2. A custom asyncio exception handler silences residual errors.
|
||||
3. ``cleanup_stale_async_clients()`` evicts stale cache entries.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Layer 1: neuter_async_httpx_del
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestNeuterAsyncHttpxDel:
|
||||
"""Verify neuter_async_httpx_del replaces __del__ on the SDK class."""
|
||||
|
||||
def test_del_becomes_noop(self):
|
||||
"""After neuter, __del__ should do nothing (no RuntimeError)."""
|
||||
from agent.auxiliary_client import neuter_async_httpx_del
|
||||
|
||||
try:
|
||||
from openai._base_client import AsyncHttpxClientWrapper
|
||||
except ImportError:
|
||||
pytest.skip("openai SDK not installed")
|
||||
|
||||
# Save original so we can restore
|
||||
original_del = AsyncHttpxClientWrapper.__del__
|
||||
try:
|
||||
neuter_async_httpx_del()
|
||||
# The patched __del__ should be a no-op lambda
|
||||
assert AsyncHttpxClientWrapper.__del__ is not original_del
|
||||
# Calling it should not raise, even without a running loop
|
||||
wrapper = MagicMock(spec=AsyncHttpxClientWrapper)
|
||||
AsyncHttpxClientWrapper.__del__(wrapper) # Should be silent
|
||||
finally:
|
||||
# Restore original to avoid leaking into other tests
|
||||
AsyncHttpxClientWrapper.__del__ = original_del
|
||||
|
||||
def test_neuter_idempotent(self):
|
||||
"""Calling neuter twice doesn't break anything."""
|
||||
from agent.auxiliary_client import neuter_async_httpx_del
|
||||
|
||||
try:
|
||||
from openai._base_client import AsyncHttpxClientWrapper
|
||||
except ImportError:
|
||||
pytest.skip("openai SDK not installed")
|
||||
|
||||
original_del = AsyncHttpxClientWrapper.__del__
|
||||
try:
|
||||
neuter_async_httpx_del()
|
||||
first_del = AsyncHttpxClientWrapper.__del__
|
||||
neuter_async_httpx_del()
|
||||
second_del = AsyncHttpxClientWrapper.__del__
|
||||
# Both calls should succeed; the class should have a no-op
|
||||
assert first_del is not original_del
|
||||
assert second_del is not original_del
|
||||
finally:
|
||||
AsyncHttpxClientWrapper.__del__ = original_del
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Layer 3: cleanup_stale_async_clients
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestCleanupStaleAsyncClients:
|
||||
"""Verify stale cache entries are evicted and force-closed."""
|
||||
|
||||
def test_removes_stale_entries(self):
|
||||
"""Entries with a closed loop should be evicted."""
|
||||
from agent.auxiliary_client import (
|
||||
_client_cache,
|
||||
_client_cache_lock,
|
||||
cleanup_stale_async_clients,
|
||||
)
|
||||
|
||||
# Create a loop, close it, make a cache entry
|
||||
loop = asyncio.new_event_loop()
|
||||
loop.close()
|
||||
|
||||
mock_client = MagicMock()
|
||||
# Give it _client attribute for _force_close_async_httpx
|
||||
mock_client._client = MagicMock()
|
||||
mock_client._client.is_closed = False
|
||||
|
||||
key = ("test_stale", True, "", "", "", (), False)
|
||||
with _client_cache_lock:
|
||||
_client_cache[key] = (mock_client, "test-model", loop)
|
||||
|
||||
try:
|
||||
cleanup_stale_async_clients()
|
||||
mock_client.close.assert_called_once()
|
||||
with _client_cache_lock:
|
||||
assert key not in _client_cache, "Stale entry should be removed"
|
||||
finally:
|
||||
# Clean up in case test fails
|
||||
with _client_cache_lock:
|
||||
_client_cache.pop(key, None)
|
||||
|
||||
def test_awaits_async_close_for_closed_loop(self):
|
||||
from agent.auxiliary_client import (
|
||||
_client_cache,
|
||||
_client_cache_lock,
|
||||
cleanup_stale_async_clients,
|
||||
)
|
||||
|
||||
class AsyncClient:
|
||||
def __init__(self):
|
||||
self._client = MagicMock()
|
||||
self._client.is_closed = False
|
||||
self.closed = False
|
||||
|
||||
async def close(self):
|
||||
self.closed = True
|
||||
|
||||
loop = asyncio.new_event_loop()
|
||||
loop.close()
|
||||
client = AsyncClient()
|
||||
key = ("test_async_close", True, "", "", "", (), False)
|
||||
with _client_cache_lock:
|
||||
_client_cache[key] = (client, "test-model", loop)
|
||||
|
||||
try:
|
||||
cleanup_stale_async_clients()
|
||||
assert client.closed
|
||||
finally:
|
||||
with _client_cache_lock:
|
||||
_client_cache.pop(key, None)
|
||||
|
||||
|
||||
def test_shutdown_closes_outside_cache_lock(self):
|
||||
from agent.auxiliary_client import (
|
||||
_client_cache,
|
||||
_client_cache_lock,
|
||||
shutdown_cached_clients,
|
||||
)
|
||||
|
||||
lock_observations = []
|
||||
|
||||
class Client:
|
||||
_client = None
|
||||
|
||||
def close(self):
|
||||
acquired = _client_cache_lock.acquire(blocking=False)
|
||||
lock_observations.append(acquired)
|
||||
if acquired:
|
||||
_client_cache_lock.release()
|
||||
|
||||
key = ("test_shutdown_lock", False, "", "", "", (), False)
|
||||
with _client_cache_lock:
|
||||
previous = dict(_client_cache)
|
||||
_client_cache.clear()
|
||||
_client_cache[key] = (Client(), "test-model", None)
|
||||
|
||||
try:
|
||||
shutdown_cached_clients()
|
||||
finally:
|
||||
with _client_cache_lock:
|
||||
_client_cache.clear()
|
||||
_client_cache.update(previous)
|
||||
|
||||
assert lock_observations == [True]
|
||||
|
||||
def test_shutdown_does_not_await_live_foreign_loop_client(self):
|
||||
from agent.auxiliary_client import (
|
||||
_client_cache,
|
||||
_client_cache_lock,
|
||||
shutdown_cached_clients,
|
||||
)
|
||||
|
||||
owner_loop = asyncio.new_event_loop()
|
||||
|
||||
class Client:
|
||||
def __init__(self):
|
||||
self.awaited = False
|
||||
|
||||
async def close(self):
|
||||
self.awaited = True
|
||||
|
||||
client = Client()
|
||||
key = ("test_shutdown_foreign_loop", True, "", "", "", (), False)
|
||||
with _client_cache_lock:
|
||||
previous = dict(_client_cache)
|
||||
_client_cache.clear()
|
||||
_client_cache[key] = (client, "test-model", owner_loop)
|
||||
|
||||
try:
|
||||
shutdown_cached_clients()
|
||||
assert client.awaited is False
|
||||
finally:
|
||||
owner_loop.close()
|
||||
with _client_cache_lock:
|
||||
_client_cache.clear()
|
||||
_client_cache.update(previous)
|
||||
|
||||
def test_keeps_live_entries(self):
|
||||
"""Entries with an open loop should be preserved."""
|
||||
from agent.auxiliary_client import (
|
||||
_client_cache,
|
||||
_client_cache_lock,
|
||||
cleanup_stale_async_clients,
|
||||
)
|
||||
|
||||
loop = asyncio.new_event_loop() # NOT closed
|
||||
|
||||
mock_client = MagicMock()
|
||||
key = ("test_live", True, "", "", "", (), False)
|
||||
with _client_cache_lock:
|
||||
_client_cache[key] = (mock_client, "test-model", loop)
|
||||
|
||||
try:
|
||||
cleanup_stale_async_clients()
|
||||
with _client_cache_lock:
|
||||
assert key in _client_cache, "Live entry should be preserved"
|
||||
finally:
|
||||
loop.close()
|
||||
with _client_cache_lock:
|
||||
_client_cache.pop(key, None)
|
||||
|
||||
def test_keeps_entries_without_loop(self):
|
||||
"""Sync entries (cached_loop=None) should be preserved."""
|
||||
from agent.auxiliary_client import (
|
||||
_client_cache,
|
||||
_client_cache_lock,
|
||||
cleanup_stale_async_clients,
|
||||
)
|
||||
|
||||
mock_client = MagicMock()
|
||||
key = ("test_sync", False, "", "", "", (), False)
|
||||
with _client_cache_lock:
|
||||
_client_cache[key] = (mock_client, "test-model", None)
|
||||
|
||||
try:
|
||||
cleanup_stale_async_clients()
|
||||
with _client_cache_lock:
|
||||
assert key in _client_cache, "Sync entry should be preserved"
|
||||
finally:
|
||||
with _client_cache_lock:
|
||||
_client_cache.pop(key, None)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cache bounded growth (#10200)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestClientCacheBoundedGrowth:
|
||||
"""Verify the cache stays bounded when loops change (fix for #10200).
|
||||
|
||||
Previously, loop_id was part of the cache key, so every new event loop
|
||||
created a new entry for the same provider config. Now loop identity is
|
||||
validated at hit time and stale entries are replaced in-place.
|
||||
"""
|
||||
|
||||
def test_same_key_replaces_stale_loop_entry(self):
|
||||
"""When the loop changes, the old entry should be replaced, not duplicated."""
|
||||
from agent.auxiliary_client import (
|
||||
_client_cache,
|
||||
_client_cache_key,
|
||||
_client_cache_lock,
|
||||
_get_cached_client,
|
||||
)
|
||||
|
||||
key = _client_cache_key(
|
||||
"test_replace",
|
||||
async_mode=True,
|
||||
task="",
|
||||
)
|
||||
|
||||
# Simulate a stale entry from a closed loop
|
||||
old_loop = asyncio.new_event_loop()
|
||||
old_loop.close()
|
||||
old_client = MagicMock()
|
||||
old_client._client = MagicMock()
|
||||
old_client._client.is_closed = False
|
||||
|
||||
with _client_cache_lock:
|
||||
_client_cache[key] = (old_client, "old-model", old_loop)
|
||||
|
||||
try:
|
||||
# Now call _get_cached_client — should detect stale loop and evict
|
||||
with patch("agent.auxiliary_client.resolve_provider_client") as mock_resolve:
|
||||
mock_resolve.return_value = (MagicMock(), "new-model")
|
||||
client, model = _get_cached_client(
|
||||
"test_replace", async_mode=True,
|
||||
)
|
||||
# The old entry should have been replaced
|
||||
with _client_cache_lock:
|
||||
assert key in _client_cache, "Key should still exist (replaced)"
|
||||
entry = _client_cache[key]
|
||||
assert entry[1] == "new-model", "Should have the new model"
|
||||
finally:
|
||||
with _client_cache_lock:
|
||||
_client_cache.pop(key, None)
|
||||
|
||||
def test_different_loops_do_not_grow_cache(self):
|
||||
"""Multiple event loops for the same provider should NOT create multiple entries."""
|
||||
from agent.auxiliary_client import (
|
||||
_client_cache,
|
||||
_client_cache_lock,
|
||||
)
|
||||
|
||||
key = ("test_no_grow", True, "", "", "", (), False)
|
||||
|
||||
loops = []
|
||||
try:
|
||||
for i in range(5):
|
||||
loop = asyncio.new_event_loop()
|
||||
loops.append(loop)
|
||||
mock_client = MagicMock()
|
||||
mock_client._client = MagicMock()
|
||||
mock_client._client.is_closed = False
|
||||
|
||||
# Close previous loop entries (simulating worker thread recycling)
|
||||
if i > 0:
|
||||
loops[i - 1].close()
|
||||
|
||||
with _client_cache_lock:
|
||||
# Simulate what _get_cached_client does: replace on loop mismatch
|
||||
if key in _client_cache:
|
||||
old_entry = _client_cache[key]
|
||||
del _client_cache[key]
|
||||
_client_cache[key] = (mock_client, f"model-{i}", loop)
|
||||
|
||||
# Only one entry should exist for this key
|
||||
with _client_cache_lock:
|
||||
count = sum(1 for k in _client_cache if k == key)
|
||||
assert count == 1, f"Expected 1 entry, got {count}"
|
||||
finally:
|
||||
for loop in loops:
|
||||
if not loop.is_closed():
|
||||
loop.close()
|
||||
with _client_cache_lock:
|
||||
_client_cache.pop(key, None)
|
||||
|
||||
def test_max_cache_size_eviction(self):
|
||||
"""Cache should not exceed _CLIENT_CACHE_MAX_SIZE."""
|
||||
from agent.auxiliary_client import (
|
||||
_client_cache,
|
||||
_client_cache_lock,
|
||||
_CLIENT_CACHE_MAX_SIZE,
|
||||
)
|
||||
|
||||
# Save existing cache state
|
||||
with _client_cache_lock:
|
||||
saved = dict(_client_cache)
|
||||
_client_cache.clear()
|
||||
|
||||
try:
|
||||
# Fill to max + 5
|
||||
for i in range(_CLIENT_CACHE_MAX_SIZE + 5):
|
||||
mock_client = MagicMock()
|
||||
mock_client._client = MagicMock()
|
||||
mock_client._client.is_closed = False
|
||||
key = (f"evict_test_{i}", False, "", "", "", (), False)
|
||||
with _client_cache_lock:
|
||||
# Inline the eviction logic (same as _get_cached_client)
|
||||
while len(_client_cache) >= _CLIENT_CACHE_MAX_SIZE:
|
||||
evict_key = next(iter(_client_cache))
|
||||
del _client_cache[evict_key]
|
||||
_client_cache[key] = (mock_client, f"model-{i}", None)
|
||||
|
||||
with _client_cache_lock:
|
||||
assert len(_client_cache) <= _CLIENT_CACHE_MAX_SIZE, \
|
||||
f"Cache size {len(_client_cache)} exceeds max {_CLIENT_CACHE_MAX_SIZE}"
|
||||
# The earliest entries should have been evicted
|
||||
assert ("evict_test_0", False, "", "", "", (), False) not in _client_cache
|
||||
# The latest entries should be present
|
||||
assert (f"evict_test_{_CLIENT_CACHE_MAX_SIZE + 4}", False, "", "", "", (), False) in _client_cache
|
||||
finally:
|
||||
with _client_cache_lock:
|
||||
_client_cache.clear()
|
||||
_client_cache.update(saved)
|
||||
@@ -0,0 +1,123 @@
|
||||
"""Auth-failure provider failover (conversation loop).
|
||||
|
||||
A 401/403 that survives the per-provider credential-refresh attempt
|
||||
(revoked OAuth, blocked/expired key, an account pinned to a dead/staging
|
||||
endpoint) must escalate to the configured fallback chain instead of
|
||||
thrashing on the same dead credential every turn.
|
||||
|
||||
Before the fix, the conversation loop's generic failover dispatch only
|
||||
fired for ``{rate_limit, billing}`` reasons; ``auth`` / ``auth_permanent``
|
||||
fell through to "switch providers manually" advice and never called
|
||||
``_try_activate_fallback()``. These tests pin:
|
||||
|
||||
1. 401/403 classify as auth (``classified.is_auth`` True).
|
||||
2. ``_try_activate_fallback`` advances the chain on an auth reason.
|
||||
3. The one-shot guard flag exists on TurnRetryState.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from run_agent import AIAgent
|
||||
from agent.error_classifier import classify_api_error, FailoverReason
|
||||
from agent.turn_retry_state import TurnRetryState
|
||||
|
||||
|
||||
def _make_agent(fallback_model=None):
|
||||
with (
|
||||
patch("model_tools.get_tool_definitions", return_value=[]),
|
||||
patch("model_tools.check_toolset_requirements", return_value={}),
|
||||
patch("agent.process_bootstrap.OpenAI"),
|
||||
):
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
fallback_model=fallback_model,
|
||||
)
|
||||
agent.client = MagicMock()
|
||||
return agent
|
||||
|
||||
|
||||
def _mock_client(base_url="https://openrouter.ai/api/v1", api_key="fb-key"):
|
||||
mock = MagicMock()
|
||||
mock.base_url = base_url
|
||||
mock.api_key = api_key
|
||||
return mock
|
||||
|
||||
|
||||
def _auth_error(status=401, msg="Your API key is invalid, blocked or out of funds."):
|
||||
err = Exception(f"Error code: {status} - {msg}")
|
||||
err.status_code = status
|
||||
return err
|
||||
|
||||
|
||||
class TestAuthErrorClassification:
|
||||
def test_401_is_auth(self):
|
||||
c = classify_api_error(_auth_error(401))
|
||||
assert c.reason in {FailoverReason.auth, FailoverReason.auth_permanent}
|
||||
assert c.is_auth is True
|
||||
|
||||
|
||||
def test_500_is_not_auth(self):
|
||||
err = Exception("Error code: 500 - internal server error")
|
||||
err.status_code = 500
|
||||
c = classify_api_error(err)
|
||||
assert c.is_auth is False
|
||||
|
||||
|
||||
class TestAuthFailoverGuardFlag:
|
||||
def test_flag_defaults_false(self):
|
||||
assert TurnRetryState().auth_failover_attempted is False
|
||||
|
||||
|
||||
class TestAuthFailoverActivation:
|
||||
"""The decision the loop makes on a persistent auth failure: when a
|
||||
fallback chain exists and the guard hasn't fired, escalate to it."""
|
||||
|
||||
def _should_failover(self, agent, classified, retry):
|
||||
# Mirror the exact gating condition added to conversation_loop.py.
|
||||
return (
|
||||
classified.is_auth
|
||||
and not retry.auth_failover_attempted
|
||||
and agent._fallback_index < len(agent._fallback_chain)
|
||||
)
|
||||
|
||||
def test_auth_failover_fires_when_chain_present(self):
|
||||
agent = _make_agent(fallback_model=[{"provider": "openai", "model": "gpt-4o"}])
|
||||
retry = TurnRetryState()
|
||||
classified = classify_api_error(_auth_error(401))
|
||||
assert self._should_failover(agent, classified, retry) is True
|
||||
# And the activation primitive actually advances on an auth reason.
|
||||
with patch(
|
||||
"agent.auxiliary_client.resolve_provider_client",
|
||||
return_value=(_mock_client(), "gpt-4o"),
|
||||
):
|
||||
advanced = agent._try_activate_fallback(reason=classified.reason)
|
||||
assert advanced is True
|
||||
assert agent._fallback_index == 1
|
||||
|
||||
def test_no_failover_without_chain(self):
|
||||
"""A user with no fallback configured (the common case for the
|
||||
original incident) does NOT failover — falls through to the
|
||||
existing terminal handling + troubleshooting advice."""
|
||||
agent = _make_agent(fallback_model=None)
|
||||
retry = TurnRetryState()
|
||||
classified = classify_api_error(_auth_error(401))
|
||||
assert self._should_failover(agent, classified, retry) is False
|
||||
|
||||
def test_guard_blocks_repeat_failover(self):
|
||||
agent = _make_agent(fallback_model=[{"provider": "openai", "model": "gpt-4o"}])
|
||||
retry = TurnRetryState()
|
||||
retry.auth_failover_attempted = True # already escalated this attempt
|
||||
classified = classify_api_error(_auth_error(401))
|
||||
assert self._should_failover(agent, classified, retry) is False
|
||||
|
||||
def test_non_auth_error_does_not_trigger_auth_failover(self):
|
||||
agent = _make_agent(fallback_model=[{"provider": "openai", "model": "gpt-4o"}])
|
||||
retry = TurnRetryState()
|
||||
err = Exception("Error code: 500 - internal server error")
|
||||
err.status_code = 500
|
||||
classified = classify_api_error(err)
|
||||
assert self._should_failover(agent, classified, retry) is False
|
||||
@@ -0,0 +1,319 @@
|
||||
"""Tests for the concurrent authorization gate and human-wait accounting (#79719).
|
||||
|
||||
Before the fix, a worker wedged *inside* the authorization gate — a hanging
|
||||
``pre_tool_call`` plugin, or an approval round-trip to a client that went
|
||||
away — had two coupled failure modes:
|
||||
|
||||
1. The serialization lock was an unbounded blocking acquire: every other
|
||||
worker needing authorization blocked behind the wedged holder forever.
|
||||
2. ``excluded_seconds()`` measured residency in ``gate.run()`` (arbitrary
|
||||
code), so an open window grew 1:1 with wall clock while the batch-deadline
|
||||
loop added it to the deadline on every poll — ``remaining`` was constant
|
||||
and the deadline NEVER fired. Algebraically:
|
||||
``remaining = (deadline + (now - window_started)) - now = deadline - window_started``.
|
||||
|
||||
The fix moves deadline exclusion to the source of the human wait
|
||||
(``tools.approval_human_wait.human_wait_window`` around the CLI prompt and the gateway
|
||||
approval poll loop) and bounds the serialization lock acquire. A wedged
|
||||
plugin now contributes nothing to the exclusion, so the batch times out
|
||||
normally; a genuine approval wait is still excluded in full.
|
||||
"""
|
||||
|
||||
import threading
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from agent.tool_executor import _ConcurrentToolAuthorizationGate
|
||||
from tools import approval as approval_mod
|
||||
from tools import approval_context
|
||||
from tools import approval_human_wait
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clean_human_wait_state():
|
||||
with approval_human_wait._human_wait_lock:
|
||||
approval_human_wait._human_wait_states.clear()
|
||||
yield
|
||||
with approval_human_wait._human_wait_lock:
|
||||
approval_human_wait._human_wait_states.clear()
|
||||
|
||||
|
||||
SESSION = "test-session-79719"
|
||||
|
||||
|
||||
def _make_gate(**kwargs) -> _ConcurrentToolAuthorizationGate:
|
||||
# Pin the session key so contextvar/env noise from other tests can't
|
||||
# change which wait state the gate reads.
|
||||
return _ConcurrentToolAuthorizationGate(session_key=SESSION, **kwargs)
|
||||
|
||||
|
||||
class TestHumanWaitTracker:
|
||||
def test_no_wait_reports_zero(self):
|
||||
assert approval_human_wait.human_wait_seconds(SESSION) == 0.0
|
||||
|
||||
def test_open_window_counts(self):
|
||||
opened = threading.Event()
|
||||
release = threading.Event()
|
||||
|
||||
def _wait():
|
||||
with approval_human_wait.human_wait_window(SESSION):
|
||||
opened.set()
|
||||
release.wait(timeout=5)
|
||||
|
||||
t = threading.Thread(target=_wait, daemon=True)
|
||||
t.start()
|
||||
assert opened.wait(timeout=5)
|
||||
time.sleep(0.05)
|
||||
assert approval_human_wait.human_wait_seconds(SESSION) > 0.0
|
||||
release.set()
|
||||
t.join(timeout=5)
|
||||
# Window closed: total is frozen (completed_seconds), not still growing.
|
||||
first = approval_human_wait.human_wait_seconds(SESSION)
|
||||
time.sleep(0.05)
|
||||
assert approval_human_wait.human_wait_seconds(SESSION) == pytest.approx(first)
|
||||
|
||||
def test_overlapping_windows_coalesce(self):
|
||||
"""Two concurrent windows on one session must not double-count wall clock."""
|
||||
release = threading.Event()
|
||||
started = threading.Barrier(3)
|
||||
|
||||
def _wait():
|
||||
with approval_human_wait.human_wait_window(SESSION):
|
||||
started.wait(timeout=5)
|
||||
release.wait(timeout=5)
|
||||
|
||||
threads = [threading.Thread(target=_wait, daemon=True) for _ in range(2)]
|
||||
start = time.monotonic()
|
||||
for t in threads:
|
||||
t.start()
|
||||
started.wait(timeout=5)
|
||||
time.sleep(0.1)
|
||||
release.set()
|
||||
for t in threads:
|
||||
t.join(timeout=5)
|
||||
elapsed = time.monotonic() - start
|
||||
# Coalesced: recorded ≤ wall clock (a double count would be ~2×).
|
||||
assert approval_human_wait.human_wait_seconds(SESSION) <= elapsed + 0.05
|
||||
|
||||
def test_sessions_are_isolated(self):
|
||||
with approval_human_wait.human_wait_window("other-session"):
|
||||
time.sleep(0.05)
|
||||
assert approval_human_wait.human_wait_seconds(SESSION) == 0.0
|
||||
assert approval_human_wait.human_wait_seconds("other-session") > 0.0
|
||||
|
||||
def test_open_window_clamped_to_approval_timeout(self, monkeypatch):
|
||||
"""A window that overstays approvals.timeout is itself wedged and must
|
||||
stop extending the exclusion (belt-and-braces for #79719)."""
|
||||
monkeypatch.setattr(approval_context, "_get_approval_timeout", lambda: 300)
|
||||
with approval_human_wait.human_wait_window(SESSION):
|
||||
state = approval_human_wait._human_wait_states[SESSION]
|
||||
# Simulate a window that has been open for a full day.
|
||||
state.window_started = time.monotonic() - 86_400.0
|
||||
assert approval_human_wait.human_wait_seconds(SESSION) <= 300.0 + 60.0
|
||||
|
||||
def test_eviction_keeps_pending_sessions(self):
|
||||
with approval_human_wait.human_wait_window(SESSION):
|
||||
for i in range(approval_human_wait._HUMAN_WAIT_MAX_SESSIONS + 8):
|
||||
with approval_human_wait.human_wait_window(f"burst-{i}"):
|
||||
pass
|
||||
# The active session survived the eviction pressure and the table
|
||||
# stayed at (or under) its cap.
|
||||
assert SESSION in approval_human_wait._human_wait_states
|
||||
assert approval_human_wait._human_wait_states[SESSION].pending == 1
|
||||
assert (
|
||||
len(approval_human_wait._human_wait_states)
|
||||
<= approval_human_wait._HUMAN_WAIT_MAX_SESSIONS
|
||||
)
|
||||
|
||||
def test_late_close_of_wedged_window_is_clamped(self, monkeypatch):
|
||||
"""A wedged window that eventually CLOSES must not retroactively inject
|
||||
its full overstay into completed_seconds (close-side clamp)."""
|
||||
monkeypatch.setattr(approval_context, "_get_approval_timeout", lambda: 300)
|
||||
with approval_human_wait.human_wait_window(SESSION):
|
||||
state = approval_human_wait._human_wait_states[SESSION]
|
||||
# Simulate the window having been open for a full day before close.
|
||||
state.window_started = time.monotonic() - 86_400.0
|
||||
assert approval_human_wait.human_wait_seconds(SESSION) <= 300.0 + 60.0
|
||||
|
||||
|
||||
class TestAuthorizationGate:
|
||||
def test_serializes_callbacks(self):
|
||||
gate = _make_gate()
|
||||
state_lock = threading.Lock()
|
||||
active = 0
|
||||
max_active = 0
|
||||
|
||||
def _callback():
|
||||
nonlocal active, max_active
|
||||
with state_lock:
|
||||
active += 1
|
||||
max_active = max(max_active, active)
|
||||
try:
|
||||
time.sleep(0.03)
|
||||
finally:
|
||||
with state_lock:
|
||||
active -= 1
|
||||
|
||||
threads = [
|
||||
threading.Thread(target=lambda: gate.run(_callback), daemon=True)
|
||||
for _ in range(4)
|
||||
]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join(timeout=5)
|
||||
assert max_active == 1
|
||||
|
||||
def test_lock_timeout_degrades_to_unserialized(self):
|
||||
"""A wedged lock holder must not park later callers forever."""
|
||||
gate = _make_gate(lock_timeout=0.1)
|
||||
holder_in = threading.Event()
|
||||
release = threading.Event()
|
||||
|
||||
def _wedged():
|
||||
holder_in.set()
|
||||
release.wait(timeout=10)
|
||||
|
||||
holder = threading.Thread(target=lambda: gate.run(_wedged), daemon=True)
|
||||
holder.start()
|
||||
assert holder_in.wait(timeout=5)
|
||||
|
||||
done = threading.Event()
|
||||
result = {}
|
||||
|
||||
def _second():
|
||||
result["value"] = gate.run(lambda: "ran-unserialized")
|
||||
done.set()
|
||||
|
||||
t = threading.Thread(target=_second, daemon=True)
|
||||
start = time.monotonic()
|
||||
t.start()
|
||||
assert done.wait(timeout=5), "second caller starved behind wedged holder"
|
||||
assert result["value"] == "ran-unserialized"
|
||||
assert time.monotonic() - start < 2.0
|
||||
release.set()
|
||||
holder.join(timeout=5)
|
||||
|
||||
def test_wedged_callback_contributes_nothing_to_exclusion(self):
|
||||
"""THE #79719 regression: gate residency is not deadline exclusion."""
|
||||
gate = _make_gate()
|
||||
wedged_in = threading.Event()
|
||||
release = threading.Event()
|
||||
|
||||
def _wedged():
|
||||
wedged_in.set()
|
||||
release.wait(timeout=10)
|
||||
|
||||
t = threading.Thread(target=lambda: gate.run(_wedged), daemon=True)
|
||||
t.start()
|
||||
assert wedged_in.wait(timeout=5)
|
||||
time.sleep(0.15)
|
||||
# No human prompt is pending — the wedge is invisible to the deadline.
|
||||
assert gate.excluded_seconds() == 0.0
|
||||
release.set()
|
||||
t.join(timeout=5)
|
||||
|
||||
def test_deadline_arithmetic_converges_with_wedged_worker(self):
|
||||
"""The issue's repro: remaining must DECREASE while a worker is wedged.
|
||||
|
||||
Pre-fix, ``remaining = deadline - window_started`` was constant for
|
||||
the life of the wedge (24h simulated in the issue). Now the exclusion
|
||||
stays 0 for a wedge, so remaining tracks wall clock down to zero.
|
||||
"""
|
||||
gate = _make_gate()
|
||||
wedged_in = threading.Event()
|
||||
release = threading.Event()
|
||||
|
||||
def _wedged():
|
||||
wedged_in.set()
|
||||
release.wait(timeout=10)
|
||||
|
||||
t = threading.Thread(target=lambda: gate.run(_wedged), daemon=True)
|
||||
t.start()
|
||||
assert wedged_in.wait(timeout=5)
|
||||
|
||||
timeout_s = 0.3
|
||||
deadline = time.monotonic() + timeout_s
|
||||
first = deadline + gate.excluded_seconds() - time.monotonic()
|
||||
time.sleep(0.15)
|
||||
second = deadline + gate.excluded_seconds() - time.monotonic()
|
||||
assert second < first, "remaining is constant — deadline never fires (#79719)"
|
||||
time.sleep(0.25)
|
||||
assert deadline + gate.excluded_seconds() - time.monotonic() <= 0, (
|
||||
"deadline never became due despite the wedge"
|
||||
)
|
||||
release.set()
|
||||
t.join(timeout=5)
|
||||
|
||||
def test_human_wait_is_excluded(self):
|
||||
"""A genuine approval wait during the batch extends the deadline."""
|
||||
gate = _make_gate()
|
||||
with approval_human_wait.human_wait_window(SESSION):
|
||||
time.sleep(0.1)
|
||||
assert gate.excluded_seconds() >= 0.09
|
||||
|
||||
def test_baseline_ignores_waits_before_batch(self):
|
||||
"""Approval waits from BEFORE this batch must not extend its deadline."""
|
||||
with approval_human_wait.human_wait_window(SESSION):
|
||||
time.sleep(0.1)
|
||||
gate = _make_gate()
|
||||
assert gate.excluded_seconds() == 0.0
|
||||
|
||||
def test_other_sessions_wait_not_excluded(self):
|
||||
gate = _make_gate()
|
||||
with approval_human_wait.human_wait_window("unrelated-session"):
|
||||
time.sleep(0.05)
|
||||
assert gate.excluded_seconds() == 0.0
|
||||
|
||||
|
||||
class TestApprovalPathsRecordHumanWait:
|
||||
def test_await_gateway_decision_records_wait(self, monkeypatch):
|
||||
"""The gateway approval poll loop must mark itself as human wait."""
|
||||
monkeypatch.setattr(approval_context, "_get_approval_timeout", lambda: 300)
|
||||
approval_data = {
|
||||
"command": "rm -rf /tmp/x",
|
||||
"description": "test",
|
||||
"pattern_key": "k",
|
||||
"pattern_keys": ["k"],
|
||||
}
|
||||
notified = threading.Event()
|
||||
result_holder = {}
|
||||
|
||||
def _worker():
|
||||
result_holder["result"] = approval_mod._await_gateway_decision(
|
||||
SESSION, lambda _data: notified.set(), approval_data
|
||||
)
|
||||
|
||||
t = threading.Thread(target=_worker, daemon=True)
|
||||
t.start()
|
||||
assert notified.wait(timeout=5)
|
||||
time.sleep(0.1)
|
||||
try:
|
||||
assert approval_human_wait.human_wait_seconds(SESSION) > 0.0
|
||||
finally:
|
||||
# Resolve the pending entry via the real production path.
|
||||
approval_mod.resolve_gateway_approval(SESSION, "deny", resolve_all=True)
|
||||
t.join(timeout=5)
|
||||
assert not t.is_alive()
|
||||
# Window closed once the wait resolved.
|
||||
assert approval_human_wait._human_wait_states[SESSION].pending == 0
|
||||
|
||||
def test_prompt_dangerous_approval_records_wait(self, monkeypatch):
|
||||
"""The CLI prompt path must mark itself as human wait."""
|
||||
observed = {}
|
||||
|
||||
def _callback(_command, _description, **_kwargs):
|
||||
observed["during"] = approval_human_wait.human_wait_seconds()
|
||||
return "deny"
|
||||
|
||||
choice = approval_mod.prompt_dangerous_approval(
|
||||
"rm -rf /tmp/x", "test", approval_callback=_callback
|
||||
)
|
||||
assert choice == "deny"
|
||||
# The window was open while the callback (the human prompt) ran.
|
||||
state = approval_human_wait._human_wait_states.get(
|
||||
approval_mod.get_current_session_key()
|
||||
)
|
||||
assert state is not None
|
||||
assert state.pending == 0
|
||||
@@ -0,0 +1,307 @@
|
||||
"""The auxiliary recovery ladder must not let a failed auth-refresh retry escape.
|
||||
|
||||
Both auth-refresh rungs in ``agent/auxiliary_client.py`` perform their retry with a
|
||||
bare ``yield`` inside a ``return`` statement:
|
||||
|
||||
if _is_auth_error(first_err) and client_is_nous:
|
||||
step = _refreshed_nous_step(...)
|
||||
if step is not None:
|
||||
return (yield step), None # nous rung
|
||||
...
|
||||
return (yield _LadderStep(
|
||||
"retry_same_provider", ...)), None # generic credential rung
|
||||
|
||||
``_rung()`` exists to convert a retry failure into ``(None, exc)`` -- but only when the
|
||||
rung's accept predicate claims the error -- so the caller can fall through to the next
|
||||
rung. An unclaimed failure (a 500, a malformed response) re-raises on purpose, since
|
||||
``_ladder_provider_fallback`` only acts on the reasons in ``_FALLBACK_REASONS``.
|
||||
Used this way no exception is caught: when the
|
||||
refreshed client also fails (e.g. an out-of-credit 404 on a stale Nous runtime token),
|
||||
the error escapes ``_aux_recovery_ladder`` and ``_ladder_provider_fallback`` never
|
||||
runs -- the configured ``auxiliary.<task>.fallback_chain`` is silently skipped. The
|
||||
rung right below (credential-pool rotation) documents the intended behavior with "then
|
||||
fall through to the provider fallback" and guards its retry with try/except.
|
||||
|
||||
The first test drives the real ladder generator with a scripted driver. The second
|
||||
exercises the real path end to end: a temp ``HERMES_HOME`` whose ``config.yaml``
|
||||
declares a ``fallback_chain``, the real client construction and HTTP layer against a
|
||||
local endpoint, with only the credential sources stubbed (the Nous portal account
|
||||
probe and the runtime-credential fetch are external boundaries).
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
import json
|
||||
import socket
|
||||
import threading
|
||||
from typing import Optional
|
||||
|
||||
import pytest
|
||||
import yaml
|
||||
|
||||
import agent.auxiliary_client as aux
|
||||
|
||||
AUX_MODEL = "z-ai/glm-5.3-flash"
|
||||
FALLBACK_MODEL = "fallback-model"
|
||||
NOUS_HOST = "inference-api.nousresearch.com"
|
||||
|
||||
|
||||
class _ApiError(Exception):
|
||||
def __init__(self, message, status_code=None):
|
||||
super().__init__(message)
|
||||
self.status_code = status_code
|
||||
|
||||
|
||||
def _auth_error():
|
||||
return _ApiError("Error code: 401 - Unauthorized", status_code=401)
|
||||
|
||||
|
||||
def _credit_error():
|
||||
return _ApiError(
|
||||
"Error code: 404 - Model '%s' requires available credits. "
|
||||
"Your account balance is too low to use paid models." % AUX_MODEL,
|
||||
status_code=404,
|
||||
)
|
||||
|
||||
|
||||
class _FakeClient:
|
||||
api_key = "sk-test"
|
||||
base_url = "https://%s/v1" % NOUS_HOST
|
||||
|
||||
|
||||
def _ladder(base_info=("https://%s/v1" % NOUS_HOST), resolved_provider="nous"):
|
||||
return aux._aux_recovery_ladder(
|
||||
_auth_error(),
|
||||
client=_FakeClient(),
|
||||
kwargs={"model": AUX_MODEL},
|
||||
task="compression",
|
||||
async_mode=False,
|
||||
base_info=base_info,
|
||||
resolved_provider=resolved_provider,
|
||||
resolved_model=AUX_MODEL,
|
||||
resolved_base_url=None,
|
||||
resolved_api_key=None,
|
||||
resolved_api_mode=None,
|
||||
final_model=AUX_MODEL,
|
||||
max_tokens=None,
|
||||
main_runtime=None,
|
||||
route_info={},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def hermetic(monkeypatch):
|
||||
"""Keep the ladder off the network and record the provider-fallback rung."""
|
||||
chain_calls = []
|
||||
|
||||
def _fake_provider_fallback(first_err, route):
|
||||
"""Stands in for the last rung: a generator that performs no steps."""
|
||||
chain_calls.append(first_err)
|
||||
yield from ()
|
||||
return "chain-response"
|
||||
|
||||
monkeypatch.setattr(aux, "_recoverable_pool_provider", lambda *a, **kw: None)
|
||||
monkeypatch.setattr(aux, "_nous_portal_account_has_fresh_paid_access", lambda: False)
|
||||
monkeypatch.setattr(aux, "_ladder_provider_fallback", _fake_provider_fallback)
|
||||
return chain_calls
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"rung,retry_succeeds",
|
||||
[("nous", False), ("nous", True), ("provider_credential", False)],
|
||||
)
|
||||
def test_post_refresh_retry_owns_the_ladder_outcome(
|
||||
rung, retry_succeeds, monkeypatch, hermetic):
|
||||
"""A failed retry resumes the ladder; a successful one returns its response."""
|
||||
if rung == "nous":
|
||||
monkeypatch.setattr(aux, "_refresh_nous_auxiliary_client",
|
||||
lambda **kwargs: (_FakeClient(), AUX_MODEL))
|
||||
expected_step, expected_base = "call", ("https://%s/v1" % NOUS_HOST)
|
||||
else:
|
||||
monkeypatch.setattr(aux, "_auth_refresh_provider_for_route",
|
||||
lambda *a, **kw: "codex")
|
||||
monkeypatch.setattr(aux, "_refresh_provider_credentials", lambda *a, **kw: True)
|
||||
monkeypatch.setattr(aux, "_evict_cached_clients", lambda *a, **kw: None)
|
||||
expected_step, expected_base = "retry_same_provider", "https://openrouter.ai/api/v1"
|
||||
|
||||
ladder = _ladder(
|
||||
base_info=expected_base,
|
||||
resolved_provider="nous" if rung == "nous" else "openrouter",
|
||||
)
|
||||
performed = []
|
||||
failure = _credit_error()
|
||||
|
||||
def perform(step):
|
||||
performed.append(step.kind)
|
||||
if retry_succeeds:
|
||||
return "refreshed-client-response"
|
||||
raise failure
|
||||
|
||||
if retry_succeeds:
|
||||
assert aux._drive_ladder(ladder, perform) == "refreshed-client-response"
|
||||
assert hermetic == [], "a successful retry must not reach the fallback chain"
|
||||
return
|
||||
|
||||
try:
|
||||
result = aux._drive_ladder(ladder, perform)
|
||||
except _ApiError as exc:
|
||||
pytest.fail(
|
||||
"the ladder let %r escape instead of falling through to the configured "
|
||||
"fallback chain (steps performed: %s)" % (exc, performed)
|
||||
)
|
||||
|
||||
assert performed == [expected_step], "the post-refresh retry is the only request"
|
||||
assert hermetic, "the configured fallback chain must be consulted"
|
||||
assert "requires available credits" in str(hermetic[0])
|
||||
assert result == "chain-response"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def nous_ladder_endpoint(monkeypatch):
|
||||
"""A local endpoint that answers the Nous host, recording every request.
|
||||
|
||||
The routing decision that selects the auth-refresh rung matches on the base URL
|
||||
host, so the endpoint is addressed as the real Nous host and ``getaddrinfo`` is
|
||||
redirected to the loopback server.
|
||||
"""
|
||||
aux.shutdown_cached_clients()
|
||||
aux._reset_aux_unhealthy_cache()
|
||||
requests = []
|
||||
resolve_address = socket.getaddrinfo
|
||||
|
||||
def local_nous_address(host, *args, **kwargs):
|
||||
if host in (NOUS_HOST, NOUS_HOST.encode()):
|
||||
host = "127.0.0.1"
|
||||
return resolve_address(host, *args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(socket, "getaddrinfo", local_nous_address)
|
||||
monkeypatch.setenv("NO_PROXY", "127.0.0.1,localhost,%s" % NOUS_HOST)
|
||||
|
||||
class Handler(BaseHTTPRequestHandler):
|
||||
def _send(self, status, payload):
|
||||
body = json.dumps(payload).encode()
|
||||
self.send_response(status)
|
||||
self.send_header("Content-Type", "application/json")
|
||||
self.send_header("Content-Length", str(len(body)))
|
||||
self.end_headers()
|
||||
self.wfile.write(body)
|
||||
|
||||
def do_POST(self):
|
||||
payload = json.loads(self.rfile.read(int(self.headers["Content-Length"])))
|
||||
requests.append((self.path, payload))
|
||||
model = payload.get("model")
|
||||
if model is None:
|
||||
self._send(400, {"error": {"message": "unexpected payload without model"}})
|
||||
return
|
||||
if model == AUX_MODEL:
|
||||
# Stale runtime token: 401 on the first attempt, then the refreshed
|
||||
# client hits a model the account cannot pay for.
|
||||
if sum(1 for _p, body in requests if body["model"] == AUX_MODEL) == 1:
|
||||
self._send(401, {"error": {"message": "Unauthorized", "type": "authentication_error"}})
|
||||
return
|
||||
self._send(404, {"error": {
|
||||
"message": "Model '%s' requires available credits. Your account "
|
||||
"balance is too low to use paid models." % AUX_MODEL,
|
||||
"type": "invalid_request_error",
|
||||
"code": "insufficient_credits",
|
||||
}})
|
||||
return
|
||||
self._send(200, {
|
||||
"id": "chatcmpl-ladder",
|
||||
"created": 1,
|
||||
"model": model,
|
||||
"object": "chat.completion",
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "The task is complete."},
|
||||
"finish_reason": "stop",
|
||||
}],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
|
||||
})
|
||||
|
||||
def log_message(self, *_args):
|
||||
pass
|
||||
|
||||
server = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
|
||||
thread = threading.Thread(
|
||||
target=server.serve_forever, kwargs={"poll_interval": 0.01}, daemon=True
|
||||
)
|
||||
thread.start()
|
||||
try:
|
||||
yield "http://%s:%d" % (NOUS_HOST, server.server_port), requests
|
||||
finally:
|
||||
aux.shutdown_cached_clients()
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
thread.join(timeout=5)
|
||||
|
||||
|
||||
def test_auth_refresh_retry_failure_reaches_the_configured_chain_over_http(
|
||||
tmp_path, monkeypatch, nous_ladder_endpoint):
|
||||
"""End to end: the configured chain must serve the retry the refresh could not."""
|
||||
host_url, requests = nous_ladder_endpoint
|
||||
local_url = "http://127.0.0.1:%s" % host_url.rsplit(":", 1)[1]
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
monkeypatch.setenv("AUX_FB_KEY", "fallback-test-key")
|
||||
config = {
|
||||
"model": {"provider": "nous", "default": AUX_MODEL},
|
||||
"providers": {"aux-fb": {"base_url": local_url + "/v1", "key_env": "AUX_FB_KEY"}},
|
||||
"auxiliary": {
|
||||
"compression": {
|
||||
"provider": "custom",
|
||||
"model": AUX_MODEL,
|
||||
"base_url": host_url + "/v1",
|
||||
"timeout": 20,
|
||||
"fallback_chain": [
|
||||
{"provider": "custom:aux-fb", "model": FALLBACK_MODEL,
|
||||
"base_url": local_url + "/v1", "timeout": 20}
|
||||
],
|
||||
}
|
||||
},
|
||||
}
|
||||
(tmp_path / "config.yaml").write_text(yaml.safe_dump(config), encoding="utf-8")
|
||||
# Credential boundaries: the account probe and the runtime-credential fetch.
|
||||
monkeypatch.setattr(aux, "_nous_portal_account_has_fresh_paid_access", lambda: False)
|
||||
monkeypatch.setattr(
|
||||
aux, "_resolve_nous_runtime_api",
|
||||
lambda **kwargs: ("fresh-nous-key", host_url + "/v1"),
|
||||
)
|
||||
|
||||
response = aux.call_llm(
|
||||
task="compression",
|
||||
messages=[{"role": "user", "content": "Summarize the conversation so far."}],
|
||||
timeout=20,
|
||||
)
|
||||
|
||||
assert response.choices[0].message.content == "The task is complete."
|
||||
seen = [
|
||||
body.get("model") for path, body in requests if path == "/v1/chat/completions"
|
||||
]
|
||||
assert seen == [AUX_MODEL, AUX_MODEL, FALLBACK_MODEL], (
|
||||
"401, then the refreshed retry fails on credits, then the configured chain: %r"
|
||||
% (seen,)
|
||||
)
|
||||
|
||||
|
||||
def test_exhausted_ladder_raises_the_narrowed_error(monkeypatch, hermetic):
|
||||
"""No chain answers: the retry's own failure surfaces, not the healed 401."""
|
||||
monkeypatch.setattr(aux, "_refresh_nous_auxiliary_client",
|
||||
lambda **kwargs: (_FakeClient(), AUX_MODEL))
|
||||
|
||||
def _no_chain(first_err, route):
|
||||
hermetic.append(first_err)
|
||||
yield from ()
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(aux, "_ladder_provider_fallback", _no_chain)
|
||||
failure = _credit_error()
|
||||
|
||||
def perform(step):
|
||||
raise failure
|
||||
|
||||
with pytest.raises(_ApiError) as raised:
|
||||
aux._drive_ladder(_ladder(), perform)
|
||||
|
||||
assert raised.value is failure, (
|
||||
"the ladder must surface the actionable retry failure, got %r" % (raised.value,))
|
||||
assert hermetic == [failure]
|
||||
@@ -724,6 +724,42 @@ class TestBuildCodexClient:
|
||||
assert mock_openai.call_args.kwargs["api_key"] == "codex-auth-token"
|
||||
assert mock_openai.call_args.kwargs["base_url"] == "https://chatgpt.com/backend-api/codex"
|
||||
|
||||
def test_profile_codex_base_url_overrides_pool_endpoint(self, monkeypatch):
|
||||
"""Auxiliary Codex calls use the same profile endpoint override as the main client."""
|
||||
entry = SimpleNamespace(
|
||||
runtime_api_key="codex-pool-token",
|
||||
runtime_base_url="https://chatgpt.com/backend-api/codex",
|
||||
)
|
||||
with (
|
||||
patch("agent.auxiliary_client._select_pool_entry", return_value=(True, entry)),
|
||||
patch("agent.auxiliary_client.OpenAI") as mock_openai,
|
||||
):
|
||||
monkeypatch.setenv("HERMES_CODEX_BASE_URL", "http://127.0.0.1:8787/v1")
|
||||
mock_openai.return_value = MagicMock()
|
||||
from agent.auxiliary_client import _build_codex_client
|
||||
|
||||
client, model = _build_codex_client("gpt-5.4")
|
||||
|
||||
assert client is not None
|
||||
assert model == "gpt-5.4"
|
||||
assert mock_openai.call_args.kwargs["base_url"] == "http://127.0.0.1:8787/v1"
|
||||
|
||||
def test_profile_codex_base_url_applies_to_raw_codex_client(self, monkeypatch):
|
||||
"""The main agent's raw Codex client honours the same endpoint override."""
|
||||
with (
|
||||
patch("agent.auxiliary_client._read_codex_access_token", return_value="codex-auth-token"),
|
||||
patch("agent.auxiliary_client.OpenAI") as mock_openai,
|
||||
):
|
||||
monkeypatch.setenv("HERMES_CODEX_BASE_URL", "http://127.0.0.1:8787/v1")
|
||||
mock_openai.return_value = MagicMock()
|
||||
from agent.auxiliary_client import resolve_provider_client
|
||||
|
||||
client, model = resolve_provider_client("openai-codex", "gpt-5.4", raw_codex=True)
|
||||
|
||||
assert client is not None
|
||||
assert model == "gpt-5.4"
|
||||
assert mock_openai.call_args.kwargs["base_url"] == "http://127.0.0.1:8787/v1"
|
||||
|
||||
def test_rejects_missing_model(self):
|
||||
"""Callers must pass an explicit model; no hardcoded default."""
|
||||
from agent.auxiliary_client import _build_codex_client
|
||||
@@ -2463,6 +2499,71 @@ class TestStaleBaseUrlWarning:
|
||||
|
||||
|
||||
class TestAuxiliaryTaskExtraBody:
|
||||
@pytest.mark.parametrize("task", ["session_search", "moa_reference", "moa_aggregator"])
|
||||
def test_generic_reasoning_fallback_clamps_ultra_for_auxiliary_and_moa_calls(self, task, monkeypatch):
|
||||
"""The OpenAI-compatible fallback must never put Hermes-only ``ultra`` on the wire."""
|
||||
from agent.auxiliary_client import _ProfileProjection, _build_call_kwargs
|
||||
|
||||
monkeypatch.setattr(
|
||||
"agent.auxiliary_client._project_provider_profile",
|
||||
lambda *_args: _ProfileProjection({}, {}, {}, False),
|
||||
)
|
||||
|
||||
kwargs = _build_call_kwargs(
|
||||
provider="custom",
|
||||
model="test-model",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
reasoning_config={"enabled": True, "effort": "ultra"},
|
||||
task=task,
|
||||
)
|
||||
|
||||
assert kwargs["extra_body"]["reasoning"] == {"enabled": True, "effort": "max"}
|
||||
|
||||
def test_task_extra_body_reasoning_effort_ultra_is_clamped(self, monkeypatch):
|
||||
"""``auxiliary.<task>.reasoning_effort: ultra`` folds into ``extra_body.reasoning`` (the path
|
||||
compression/title/vision/... use, with no reasoning_config) and must take the same clamp."""
|
||||
import agent.auxiliary_client as aux
|
||||
from agent.auxiliary_client import _build_call_kwargs, _get_task_extra_body
|
||||
|
||||
monkeypatch.setattr(aux, "_get_auxiliary_task_config", lambda task: {"reasoning_effort": "ultra"})
|
||||
monkeypatch.setattr(aux, "_project_provider_profile", lambda *_args: aux._ProfileProjection({}, {}, {}, False))
|
||||
|
||||
kwargs = _build_call_kwargs(
|
||||
provider="custom",
|
||||
model="test-model",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
extra_body=_get_task_extra_body("session_search"),
|
||||
reasoning_config=None,
|
||||
task="session_search",
|
||||
)
|
||||
|
||||
assert kwargs["extra_body"]["reasoning"] == {"enabled": True, "effort": "max"}
|
||||
|
||||
def test_profile_projection_receives_wire_clamped_effort(self, monkeypatch):
|
||||
"""Profiles clamp only against their own narrower sets (or a catalog that may be cold), so
|
||||
``ultra`` must already be a wire level when the projection sees it — the MoA aggregator on
|
||||
an OpenRouter/Nous slot 400'd otherwise (#112010)."""
|
||||
import agent.auxiliary_client as aux
|
||||
|
||||
seen = {}
|
||||
real = aux._project_provider_profile
|
||||
|
||||
def spy(provider, provider_norm, model, effective_base, reasoning_config):
|
||||
seen["config"] = reasoning_config
|
||||
return real(provider, provider_norm, model, effective_base, reasoning_config)
|
||||
|
||||
monkeypatch.setattr(aux, "_project_provider_profile", spy)
|
||||
kwargs = aux._build_call_kwargs(
|
||||
provider="openrouter",
|
||||
model="deepseek/deepseek-v4.1-flash",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
reasoning_config={"enabled": True, "effort": "ultra"},
|
||||
task="moa_aggregator",
|
||||
)
|
||||
|
||||
assert seen["config"] == {"enabled": True, "effort": "max"}
|
||||
assert "ultra" not in json.dumps(kwargs.get("extra_body")) and kwargs.get("reasoning_effort") != "ultra"
|
||||
|
||||
def test_sync_call_merges_task_extra_body_from_config(self):
|
||||
client = MagicMock()
|
||||
client.base_url = "https://api.example.com/v1"
|
||||
@@ -2631,6 +2732,17 @@ class TestAuxiliaryTaskExtraBody:
|
||||
assert not any("OPENAI_BASE_URL is set" in rec.message for rec in caplog.records), \
|
||||
"Should NOT warn when provider is 'custom'"
|
||||
|
||||
def test_bare_custom_auth_error_does_not_fall_back_to_env_base_url(self, monkeypatch):
|
||||
"""Bare 'custom' with nothing configured: the main resolver raises AuthError; aux must
|
||||
return no endpoint rather than route to a stale env OPENAI_BASE_URL with a placeholder key."""
|
||||
from hermes_cli.auth import AuthError
|
||||
from agent.auxiliary_client import _resolve_custom_runtime
|
||||
monkeypatch.setenv("OPENAI_BASE_URL", "https://old-proxy.example/v1")
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
with patch("hermes_cli.runtime_provider.resolve_runtime_provider",
|
||||
side_effect=AuthError("no creds", provider="custom", code="missing_api_key")):
|
||||
assert _resolve_custom_runtime() == (None, None, None)
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -2946,6 +3058,37 @@ class TestAnthropicAuxiliaryReasoningTranslation:
|
||||
)
|
||||
assert "_reasoning_config" not in openai_wire_kwargs
|
||||
|
||||
def test_anthropic_messages_profile_keeps_reasoning_reachable(self):
|
||||
# commandcode-anthropic: OpenAI-shaped URL, anthropic_messages api_mode, and a profile
|
||||
# class that overrides build_api_kwargs_extras (so the generic extra_body.reasoning
|
||||
# fallback the adapter used to read is suppressed). The adapter must still be told.
|
||||
import model_tools # noqa: F401 — triggers provider discovery
|
||||
import providers
|
||||
|
||||
assert providers.get_provider_profile("commandcode-anthropic") is not None
|
||||
rc = {"enabled": False}
|
||||
kwargs = _build_call_kwargs(
|
||||
"commandcode-anthropic", "claude-haiku-4-5-20251001", [{"role": "user", "content": "hi"}],
|
||||
reasoning_config=rc, base_url="https://api.commandcode.ai/provider/v1",
|
||||
)
|
||||
assert kwargs["_reasoning_config"] == rc
|
||||
chat_kwargs = _build_call_kwargs(
|
||||
"commandcode", "Qwen/Qwen3.7-Max", [{"role": "user", "content": "hi"}],
|
||||
reasoning_config=rc, base_url="https://api.commandcode.ai/provider/v1",
|
||||
)
|
||||
assert "_reasoning_config" not in chat_kwargs
|
||||
|
||||
def test_anthropic_messages_profile_resolves_to_messages_adapter(self, monkeypatch):
|
||||
# Bare ``provider: commandcode-anthropic`` (no api_mode) must wrap the client on the
|
||||
# profile's declared wire, or the ``_reasoning_config`` kwarg above would reach a plain
|
||||
# OpenAI client and TypeError.
|
||||
import model_tools # noqa: F401
|
||||
from agent.auxiliary_client import AnthropicAuxiliaryClient, resolve_provider_client
|
||||
|
||||
monkeypatch.setenv("COMMANDCODE_API_KEY", "sk-test-" + "x" * 20)
|
||||
client, _ = resolve_provider_client("commandcode-anthropic", model="claude-haiku-4-5-20251001")
|
||||
assert isinstance(client, AnthropicAuxiliaryClient)
|
||||
|
||||
|
||||
class TestAuxiliaryProviderProfileReasoning:
|
||||
"""Auxiliary calls must reuse provider-profile reasoning wire shapes."""
|
||||
|
||||
@@ -0,0 +1,237 @@
|
||||
"""Regression tests for issue #52608.
|
||||
|
||||
auxiliary_client `_try_anthropic()` must NOT apply `cfg["model"]["base_url"]`
|
||||
when the configured base_url host is not an Anthropic-compatible endpoint
|
||||
(e.g. OpenRouter, OpenAI). Operators routing main traffic through a
|
||||
non-Anthropic provider's endpoint while keeping `provider: anthropic` would
|
||||
otherwise have every side-channel call (memory extractors, reflection,
|
||||
vision, title generation) 401 from the foreign host.
|
||||
"""
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
||||
def _extract_base_url_passed_to_build(mock_build):
|
||||
"""Pull the base_url that `_try_anthropic()` actually handed to build_anthropic_client."""
|
||||
args, _kwargs = mock_build.call_args
|
||||
# build_anthropic_client(token, base_url) per agent/auxiliary_client.py line 2180
|
||||
assert len(args) >= 2, f"expected (token, base_url), got args={args}"
|
||||
return args[1]
|
||||
|
||||
|
||||
class TestTryAnthropicBaseUrlHostValidation:
|
||||
"""Issue #52608: side-channel calls must not be sent to a non-Anthropic host."""
|
||||
|
||||
def test_openrouter_base_url_does_not_leak_into_auxiliary(self, tmp_path, monkeypatch):
|
||||
"""cfg.model.base_url=https://openrouter.ai/api/v1 must NOT override aux base_url."""
|
||||
import yaml
|
||||
from agent.auxiliary_client import _try_anthropic
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
(tmp_path / "config.yaml").write_text(yaml.safe_dump({
|
||||
"model": {
|
||||
"provider": "anthropic",
|
||||
"model": "claude-haiku-4-5-20251001",
|
||||
"base_url": "https://openrouter.ai/api/v1",
|
||||
}
|
||||
}))
|
||||
|
||||
with (
|
||||
patch(
|
||||
"agent.auxiliary_client._select_pool_entry", return_value=(False, None)
|
||||
),
|
||||
patch(
|
||||
"agent.anthropic_credentials.resolve_anthropic_token",
|
||||
return_value="***",
|
||||
),
|
||||
patch(
|
||||
"agent.anthropic_adapter.build_anthropic_client"
|
||||
) as mock_build,
|
||||
):
|
||||
mock_build.return_value = MagicMock()
|
||||
client, _model = _try_anthropic()
|
||||
|
||||
assert client is not None, "auxiliary client must still be created"
|
||||
actual = _extract_base_url_passed_to_build(mock_build)
|
||||
assert actual == "https://api.anthropic.com", (
|
||||
f"Auxiliary client must use the Anthropic default base_url, "
|
||||
f"not the operator's main-session override. Got: {actual!r}"
|
||||
)
|
||||
|
||||
def test_anthropic_default_host_is_preserved(self, tmp_path, monkeypatch):
|
||||
"""The common case (operator sets model.base_url to api.anthropic.com) must still apply."""
|
||||
import yaml
|
||||
from agent.auxiliary_client import _try_anthropic
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
(tmp_path / "config.yaml").write_text(yaml.safe_dump({
|
||||
"model": {
|
||||
"provider": "anthropic",
|
||||
"model": "claude-haiku-4-5-20251001",
|
||||
"base_url": "https://api.anthropic.com",
|
||||
}
|
||||
}))
|
||||
|
||||
with (
|
||||
patch(
|
||||
"agent.auxiliary_client._select_pool_entry", return_value=(False, None)
|
||||
),
|
||||
patch(
|
||||
"agent.anthropic_credentials.resolve_anthropic_token",
|
||||
return_value="***",
|
||||
),
|
||||
patch(
|
||||
"agent.anthropic_adapter.build_anthropic_client"
|
||||
) as mock_build,
|
||||
):
|
||||
mock_build.return_value = MagicMock()
|
||||
client, _model = _try_anthropic()
|
||||
|
||||
assert client is not None
|
||||
actual = _extract_base_url_passed_to_build(mock_build)
|
||||
assert actual == "https://api.anthropic.com"
|
||||
|
||||
def test_openai_base_url_does_not_leak(self, tmp_path, monkeypatch):
|
||||
"""Generic non-Anthropic host must not be applied as auxiliary base_url."""
|
||||
import yaml
|
||||
from agent.auxiliary_client import _try_anthropic
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
(tmp_path / "config.yaml").write_text(yaml.safe_dump({
|
||||
"model": {
|
||||
"provider": "anthropic",
|
||||
"model": "claude-haiku-4-5-20251001",
|
||||
"base_url": "https://api.openai.com/v1",
|
||||
}
|
||||
}))
|
||||
|
||||
with (
|
||||
patch(
|
||||
"agent.auxiliary_client._select_pool_entry", return_value=(False, None)
|
||||
),
|
||||
patch(
|
||||
"agent.anthropic_credentials.resolve_anthropic_token",
|
||||
return_value="***",
|
||||
),
|
||||
patch(
|
||||
"agent.anthropic_adapter.build_anthropic_client"
|
||||
) as mock_build,
|
||||
):
|
||||
mock_build.return_value = MagicMock()
|
||||
client, _model = _try_anthropic()
|
||||
|
||||
assert client is not None
|
||||
actual = _extract_base_url_passed_to_build(mock_build)
|
||||
assert actual == "https://api.anthropic.com", (
|
||||
f"Non-Anthropic host must not be applied. Got: {actual!r}"
|
||||
)
|
||||
|
||||
|
||||
def test_empty_base_url_falls_back_to_default(self, tmp_path, monkeypatch):
|
||||
"""Empty model.base_url must not crash and must fall back to default."""
|
||||
import yaml
|
||||
from agent.auxiliary_client import _try_anthropic
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
(tmp_path / "config.yaml").write_text(yaml.safe_dump({
|
||||
"model": {
|
||||
"provider": "anthropic",
|
||||
"model": "claude-haiku-4-5-20251001",
|
||||
"base_url": "",
|
||||
}
|
||||
}))
|
||||
|
||||
with (
|
||||
patch(
|
||||
"agent.auxiliary_client._select_pool_entry", return_value=(False, None)
|
||||
),
|
||||
patch(
|
||||
"agent.anthropic_credentials.resolve_anthropic_token",
|
||||
return_value="***",
|
||||
),
|
||||
patch(
|
||||
"agent.anthropic_adapter.build_anthropic_client"
|
||||
) as mock_build,
|
||||
):
|
||||
mock_build.return_value = MagicMock()
|
||||
client, _model = _try_anthropic()
|
||||
|
||||
assert client is not None
|
||||
actual = _extract_base_url_passed_to_build(mock_build)
|
||||
assert actual == "https://api.anthropic.com"
|
||||
|
||||
def test_anthropic_suffix_gateway_base_url_is_applied(self, tmp_path, monkeypatch):
|
||||
"""A gateway exposing the Messages protocol under a ``/anthropic`` suffix
|
||||
must be honored — the same convention the primary path already trusts —
|
||||
so auxiliary/fallback calls hit the configured endpoint, not the default."""
|
||||
import yaml
|
||||
from agent.auxiliary_client import _try_anthropic
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
(tmp_path / "config.yaml").write_text(yaml.safe_dump({
|
||||
"model": {
|
||||
"provider": "anthropic",
|
||||
"model": "claude-haiku-4-5-20251001",
|
||||
"base_url": "https://gateway.example.com/anthropic",
|
||||
}
|
||||
}))
|
||||
|
||||
with (
|
||||
patch(
|
||||
"agent.auxiliary_client._select_pool_entry", return_value=(False, None)
|
||||
),
|
||||
patch(
|
||||
"agent.anthropic_credentials.resolve_anthropic_token",
|
||||
return_value="***",
|
||||
),
|
||||
patch(
|
||||
"agent.anthropic_adapter.build_anthropic_client"
|
||||
) as mock_build,
|
||||
):
|
||||
mock_build.return_value = MagicMock()
|
||||
client, _model = _try_anthropic()
|
||||
|
||||
assert client is not None
|
||||
actual = _extract_base_url_passed_to_build(mock_build)
|
||||
assert actual == "https://gateway.example.com/anthropic", (
|
||||
f"/anthropic-suffixed gateway base_url must be applied. Got: {actual!r}"
|
||||
)
|
||||
|
||||
def test_anthropic_suffix_host_check_direct(self):
|
||||
"""Unit-level: the host check trusts native hosts and /anthropic gateways,
|
||||
and still rejects a bare non-Anthropic host (the #52608 guard)."""
|
||||
from agent.auxiliary_client import _is_anthropic_compatible_host as ok
|
||||
assert ok("https://api.anthropic.com") is True
|
||||
assert ok("https://gateway.example.com/anthropic") is True
|
||||
assert ok("http://127.0.0.1:8080/anthropic/v1") is True
|
||||
assert ok("https://openrouter.ai/api/v1") is False
|
||||
assert ok("https://api.openai.com/v1") is False
|
||||
assert ok("") is False
|
||||
|
||||
def test_anthropic_host_with_path_is_preserved(self, tmp_path, monkeypatch):
|
||||
"""api.anthropic.com with a path suffix must still pass the host check."""
|
||||
import yaml
|
||||
from agent.auxiliary_client import _try_anthropic
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
(tmp_path / "config.yaml").write_text(yaml.safe_dump({
|
||||
"model": {
|
||||
"provider": "anthropic",
|
||||
"model": "claude-haiku-4-5-20251001",
|
||||
"base_url": "https://api.anthropic.com/v1/messages",
|
||||
}
|
||||
}))
|
||||
|
||||
with (
|
||||
patch(
|
||||
"agent.auxiliary_client._select_pool_entry", return_value=(False, None)
|
||||
),
|
||||
patch(
|
||||
"agent.anthropic_credentials.resolve_anthropic_token",
|
||||
return_value="***",
|
||||
),
|
||||
patch(
|
||||
"agent.anthropic_adapter.build_anthropic_client"
|
||||
) as mock_build,
|
||||
):
|
||||
mock_build.return_value = MagicMock()
|
||||
client, _model = _try_anthropic()
|
||||
|
||||
assert client is not None
|
||||
actual = _extract_base_url_passed_to_build(mock_build)
|
||||
assert actual == "https://api.anthropic.com/v1/messages", (
|
||||
f"Anthropic host with path must be preserved. Got: {actual!r}"
|
||||
)
|
||||
@@ -216,7 +216,7 @@ class TestResolveVisionProviderClientModelNormalization:
|
||||
|
||||
assert provider == "zai"
|
||||
assert client is not None
|
||||
assert model == "glm-5v-turbo" # zai has dedicated vision model in _PROVIDER_VISION_MODELS
|
||||
assert model == "glm-5.3-flash" # zai coding endpoints support this vision-capable fallback
|
||||
|
||||
|
||||
class TestAutoClientCacheModelCompatibility:
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
"""Auxiliary ``api_key`` routes honor ``ProviderProfile.create_client()`` — parity with the main agent.
|
||||
|
||||
``resolve_provider_client()``'s ``api_key`` branch used to build ``openai.OpenAI`` directly,
|
||||
so an out-of-tree provider registered with ``auth_type="api_key"`` lost its native transport
|
||||
for auxiliary tasks even though the main-agent path
|
||||
(``agent_runtime_helpers._provider_supplied_client``) honors the same hook (#112384).
|
||||
These tests pin the seam through the real resolution entry point: the native client
|
||||
(sync + async), and the fall-through to the standard client for ordinary providers and
|
||||
for a broken plugin.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
import providers as _providers
|
||||
from providers.base import ProviderProfile
|
||||
|
||||
_PROBE_ENV_VAR = "AUX_SEAM_PROBE_AUTH"
|
||||
_PROBE_KEY = "probe-sentinel"
|
||||
|
||||
|
||||
class _FakeNativeClient:
|
||||
HERMES_SKIP_TRANSPORT_WRAP = True
|
||||
HERMES_SKIP_ASYNC_WRAP = True
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
|
||||
|
||||
class _NativeProfile(ProviderProfile):
|
||||
def create_client(self, **kwargs):
|
||||
return _FakeNativeClient(**kwargs)
|
||||
|
||||
|
||||
class _PassThroughProfile(ProviderProfile):
|
||||
"""Ordinary provider: ``create_client`` returns None so the standard client is built."""
|
||||
|
||||
def create_client(self, **kwargs):
|
||||
return None
|
||||
|
||||
|
||||
class _ExplodingProfile(ProviderProfile):
|
||||
def create_client(self, **kwargs):
|
||||
raise RuntimeError("plugin is broken")
|
||||
|
||||
|
||||
def _probe_profile(cls, name: str) -> ProviderProfile:
|
||||
return cls(
|
||||
name=name,
|
||||
auth_type="api_key",
|
||||
env_vars=(_PROBE_ENV_VAR,),
|
||||
base_url=f"https://{name}.invalid",
|
||||
default_aux_model="probe-model",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def registered(monkeypatch):
|
||||
"""Register provider profiles for one test on copies of both registries.
|
||||
|
||||
Mirrors the import-time synthesis ``hermes_cli.auth`` performs for plugin ``api_key``
|
||||
profiles, which is what routes them into ``_resolve_api_key_branch`` in the first place.
|
||||
"""
|
||||
import hermes_cli.auth as _auth
|
||||
from agent import secret_scope as _secret_scope
|
||||
|
||||
_providers._discover_providers()
|
||||
monkeypatch.setattr(_providers, "_REGISTRY", dict(_providers._REGISTRY))
|
||||
monkeypatch.setattr(_providers, "_ALIASES", dict(_providers._ALIASES))
|
||||
monkeypatch.setattr(_providers, "_PROVIDER_LIST_CACHE", None)
|
||||
monkeypatch.setattr(_auth, "PROVIDER_REGISTRY", dict(_auth.PROVIDER_REGISTRY))
|
||||
scope_token = _secret_scope.set_secret_scope({_PROBE_ENV_VAR: _PROBE_KEY})
|
||||
|
||||
def _register(profile: ProviderProfile) -> None:
|
||||
_providers.register_provider(profile)
|
||||
_auth._register_plugin_provider(profile)
|
||||
|
||||
yield _register
|
||||
_secret_scope.reset_secret_scope(scope_token)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"extra",
|
||||
[{"task": "title_generation"}, {"async_mode": True}],
|
||||
ids=["sync", "async-skips-AsyncOpenAI-rewrap"],
|
||||
)
|
||||
def test_native_profile_supplies_the_auxiliary_client(registered, extra):
|
||||
from agent.auxiliary_client import resolve_provider_client
|
||||
|
||||
registered(_probe_profile(_NativeProfile, "aux-seam-native"))
|
||||
client, model = resolve_provider_client("aux-seam-native", "probe-model", **extra)
|
||||
|
||||
assert isinstance(client, _FakeNativeClient)
|
||||
assert model == "probe-model"
|
||||
# The hook receives the same mapping the branch would have passed to openai.OpenAI.
|
||||
assert client.kwargs == {"api_key": _PROBE_KEY, "base_url": "https://aux-seam-native.invalid"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"profile_cls", [_PassThroughProfile, _ExplodingProfile], ids=["returns-None", "raises"]
|
||||
)
|
||||
def test_profile_without_a_client_falls_back_to_the_standard_client(registered, profile_cls):
|
||||
from openai import OpenAI
|
||||
|
||||
from agent.auxiliary_client import resolve_provider_client
|
||||
|
||||
registered(_probe_profile(profile_cls, "aux-seam-fallback"))
|
||||
client, model = resolve_provider_client("aux-seam-fallback", "probe-model")
|
||||
|
||||
# A raising plugin can only fail to provide a client, never break auxiliary resolution.
|
||||
assert isinstance(client, OpenAI)
|
||||
assert model == "probe-model"
|
||||
@@ -1,6 +1,6 @@
|
||||
"""Tests for user-configured ``model.default_headers`` in the auxiliary client.
|
||||
|
||||
Companion to ``tests/run_agent/test_provider_attribution_headers.py`` (which
|
||||
Companion to ``tests/agent/test_provider_attribution_headers.py`` (which
|
||||
covers the main agent client). The main agent turn and the auxiliary client
|
||||
(title generation, context compression, vision routing) build separate OpenAI
|
||||
clients, so a ``custom`` endpoint behind a gateway/WAF that rejects the OpenAI
|
||||
|
||||
@@ -0,0 +1,887 @@
|
||||
"""Regression tests for background review agent cleanup."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
|
||||
import run_agent as run_agent_module
|
||||
from run_agent import AIAgent
|
||||
|
||||
|
||||
_REAL_THREAD = threading.Thread
|
||||
|
||||
|
||||
class _TurnBoundaryReached(Exception):
|
||||
"""Stop a live turn exactly when it reaches turn-context construction."""
|
||||
|
||||
|
||||
class CapturingThread:
|
||||
targets = []
|
||||
|
||||
def __init__(self, *, target, daemon=None, name=None):
|
||||
self.targets.append(target)
|
||||
|
||||
def start(self):
|
||||
pass
|
||||
|
||||
|
||||
class ObservedEvent:
|
||||
"""A real Event that also exposes when a waiter starts waiting."""
|
||||
|
||||
def __init__(self):
|
||||
self._event = threading.Event()
|
||||
self.wait_started = threading.Event()
|
||||
self.set_calls = 0
|
||||
|
||||
def set(self):
|
||||
self.set_calls += 1
|
||||
self._event.set()
|
||||
|
||||
def wait(self, timeout=None):
|
||||
self.wait_started.set()
|
||||
return self._event.wait(timeout)
|
||||
|
||||
def is_set(self):
|
||||
return self._event.is_set()
|
||||
|
||||
|
||||
class FakeReviewAgent:
|
||||
def __init__(self, **kwargs):
|
||||
self._session_messages = []
|
||||
|
||||
def run_conversation(self, **kwargs):
|
||||
pass
|
||||
|
||||
def interrupt(self, message=None):
|
||||
pass
|
||||
|
||||
def release_clients(self):
|
||||
pass
|
||||
|
||||
|
||||
def _bare_agent() -> AIAgent:
|
||||
agent = object.__new__(AIAgent)
|
||||
agent.model = "fake-model"
|
||||
agent.platform = "telegram"
|
||||
agent.provider = "openai"
|
||||
agent.base_url = ""
|
||||
agent.api_key = ""
|
||||
agent.api_mode = ""
|
||||
agent.session_id = "test-session"
|
||||
agent._parent_session_id = ""
|
||||
agent._credential_pool = None
|
||||
agent._memory_store = object()
|
||||
agent._memory_enabled = True
|
||||
agent._user_profile_enabled = False
|
||||
agent._cached_system_prompt = "test-cached-system-prompt"
|
||||
import datetime as _dt
|
||||
agent.session_start = _dt.datetime(2026, 1, 1, 12, 0, 0)
|
||||
agent._MEMORY_REVIEW_PROMPT = "review memory"
|
||||
agent._SKILL_REVIEW_PROMPT = "review skills"
|
||||
agent._COMBINED_REVIEW_PROMPT = "review both"
|
||||
agent.background_review_callback = None
|
||||
agent.status_callback = None
|
||||
agent._safe_print = lambda *_args, **_kwargs: None
|
||||
import threading as _threading
|
||||
agent._background_review_agent = None
|
||||
agent._background_review_run = None
|
||||
agent._background_review_lock = _threading.Lock()
|
||||
agent._active_children = []
|
||||
agent._active_children_lock = _threading.Lock()
|
||||
return agent
|
||||
|
||||
|
||||
class ImmediateThread:
|
||||
def __init__(self, *, target, daemon=None, name=None):
|
||||
self._target = target
|
||||
|
||||
def start(self):
|
||||
self._target()
|
||||
|
||||
|
||||
def _install_live_turn_boundary(monkeypatch, on_boundary=None):
|
||||
import agent.conversation_loop as conversation_loop_module
|
||||
|
||||
def stop_at_boundary(*args, **kwargs):
|
||||
if on_boundary is not None:
|
||||
on_boundary()
|
||||
raise _TurnBoundaryReached
|
||||
|
||||
monkeypatch.setattr(
|
||||
conversation_loop_module,
|
||||
"build_turn_context",
|
||||
stop_at_boundary,
|
||||
)
|
||||
|
||||
|
||||
def _run_wrapped_live_turn_to_boundary(agent, result):
|
||||
try:
|
||||
result["return"] = AIAgent.run_conversation(
|
||||
agent,
|
||||
"next turn",
|
||||
task_id="live-task",
|
||||
)
|
||||
except _TurnBoundaryReached:
|
||||
result["boundary_reached"] = True
|
||||
except BaseException as exc: # surfaced in the test thread for a useful failure
|
||||
result["error"] = exc
|
||||
|
||||
|
||||
def _install_relay_recorder(monkeypatch, review_run=None):
|
||||
from agent import relay_runtime
|
||||
from hermes_cli.observability import relay_shared_metrics
|
||||
|
||||
calls = []
|
||||
|
||||
def review_acknowledged():
|
||||
return bool(review_run and review_run.request_done.is_set())
|
||||
|
||||
class RelayTurn:
|
||||
relay_enabled = True
|
||||
|
||||
class RecordingCoordinator:
|
||||
def acquire_conversation(self, **kwargs):
|
||||
calls.append(("acquire", review_acknowledged()))
|
||||
return object()
|
||||
|
||||
def begin_turn(self, lease, **kwargs):
|
||||
calls.append(("begin", review_acknowledged()))
|
||||
return RelayTurn()
|
||||
|
||||
def finish_logical_calls(self, turn, **kwargs):
|
||||
pass
|
||||
|
||||
def end_turn(self, turn, **kwargs):
|
||||
pass
|
||||
|
||||
def release_conversation(self, lease):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr(
|
||||
relay_runtime,
|
||||
"SESSION_COORDINATOR",
|
||||
RecordingCoordinator(),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
relay_runtime,
|
||||
"current_profile_key",
|
||||
lambda: "/test-profile",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
relay_shared_metrics,
|
||||
"start_task_run",
|
||||
lambda **kwargs: calls.append(
|
||||
("start_task_run", review_acknowledged())
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
relay_shared_metrics,
|
||||
"finish_task_run",
|
||||
lambda **kwargs: None,
|
||||
)
|
||||
return calls
|
||||
|
||||
|
||||
def test_background_review_releases_clients_without_closing_shared_session(monkeypatch):
|
||||
"""The review fork must not clean up resources owned by its parent session.
|
||||
|
||||
The fork uses the foreground session ID for prefix-cache parity. Calling
|
||||
``close()`` would therefore kill that session's registered terminal
|
||||
processes and tear down its environment when the review completes.
|
||||
"""
|
||||
events = []
|
||||
|
||||
class FakeReviewAgent:
|
||||
def __init__(self, **kwargs):
|
||||
events.append(("init", kwargs))
|
||||
self._session_messages = []
|
||||
|
||||
def run_conversation(self, **kwargs):
|
||||
events.append(("run_conversation", kwargs))
|
||||
|
||||
def close(self):
|
||||
events.append(("close", None))
|
||||
|
||||
def release_clients(self):
|
||||
events.append(("release_clients", None))
|
||||
|
||||
monkeypatch.setattr(run_agent_module, "AIAgent", FakeReviewAgent)
|
||||
monkeypatch.setattr(run_agent_module.threading, "Thread", ImmediateThread)
|
||||
|
||||
agent = _bare_agent()
|
||||
|
||||
AIAgent._spawn_background_review(
|
||||
agent,
|
||||
messages_snapshot=[{"role": "user", "content": "hello"}],
|
||||
review_memory=True,
|
||||
)
|
||||
|
||||
assert [name for name, _payload in events] == [
|
||||
"init",
|
||||
"run_conversation",
|
||||
"release_clients",
|
||||
]
|
||||
|
||||
|
||||
def test_background_review_fork_opts_out_of_session_finalization(monkeypatch):
|
||||
"""The review fork shares the parent's live session_id, so it must set
|
||||
``_end_session_on_close = False``. Otherwise close() (now finalizing owned
|
||||
session rows) would end the still-active parent session mid-conversation
|
||||
every time the review fires (~every 10 turns). Regression for #12029.
|
||||
"""
|
||||
seen = {}
|
||||
|
||||
class FakeReviewAgent:
|
||||
def __init__(self, **kwargs):
|
||||
self._session_messages = []
|
||||
# Default matches AIAgent.__init__ (agent_init.py): owns its row.
|
||||
self._end_session_on_close = True
|
||||
|
||||
def __setattr__(self, name, value):
|
||||
object.__setattr__(self, name, value)
|
||||
if name == "_end_session_on_close":
|
||||
seen["end_session_on_close"] = value
|
||||
|
||||
def run_conversation(self, **kwargs):
|
||||
# By the time the fork runs, the opt-out must already be applied.
|
||||
seen["at_run_time"] = self._end_session_on_close
|
||||
|
||||
def shutdown_memory_provider(self):
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr(run_agent_module, "AIAgent", FakeReviewAgent)
|
||||
monkeypatch.setattr(run_agent_module.threading, "Thread", ImmediateThread)
|
||||
|
||||
agent = _bare_agent()
|
||||
|
||||
AIAgent._spawn_background_review(
|
||||
agent,
|
||||
messages_snapshot=[{"role": "user", "content": "hello"}],
|
||||
review_memory=True,
|
||||
)
|
||||
|
||||
assert seen.get("end_session_on_close") is False
|
||||
assert seen.get("at_run_time") is False
|
||||
|
||||
|
||||
def test_background_review_skipped_in_delegation_subagent(monkeypatch):
|
||||
"""The automatic post-turn review must NOT fire inside a delegation
|
||||
subagent (``_delegate_depth > 0``).
|
||||
|
||||
Regression for #85859: the fork inherits the subagent's live model, so in
|
||||
a delegation subagent running a premium model it replayed the whole
|
||||
conversation at premium rates. Subagents are already barred from writing
|
||||
shared MEMORY.md, so there is nothing for the review to persist here.
|
||||
"""
|
||||
forks = []
|
||||
|
||||
class FakeReviewAgent:
|
||||
def __init__(self, **kwargs):
|
||||
forks.append(kwargs)
|
||||
|
||||
def run_conversation(self, **kwargs):
|
||||
pass
|
||||
|
||||
def shutdown_memory_provider(self):
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr(run_agent_module, "AIAgent", FakeReviewAgent)
|
||||
monkeypatch.setattr(run_agent_module.threading, "Thread", ImmediateThread)
|
||||
|
||||
agent = _bare_agent()
|
||||
agent._delegate_depth = 1 # this agent IS a delegation subagent
|
||||
|
||||
AIAgent._spawn_background_review(
|
||||
agent,
|
||||
messages_snapshot=[{"role": "user", "content": "hello"}],
|
||||
review_memory=True,
|
||||
review_skills=True,
|
||||
)
|
||||
|
||||
assert forks == [], "no review fork should be spawned inside a subagent"
|
||||
|
||||
|
||||
def test_background_review_runs_at_top_level(monkeypatch):
|
||||
"""Sibling guard for the subagent skip: at ``_delegate_depth == 0`` the
|
||||
review still fires exactly as before (the cost guard is subagent-only)."""
|
||||
forks = []
|
||||
|
||||
class FakeReviewAgent:
|
||||
def __init__(self, **kwargs):
|
||||
forks.append(kwargs)
|
||||
|
||||
def run_conversation(self, **kwargs):
|
||||
pass
|
||||
|
||||
def shutdown_memory_provider(self):
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr(run_agent_module, "AIAgent", FakeReviewAgent)
|
||||
monkeypatch.setattr(run_agent_module.threading, "Thread", ImmediateThread)
|
||||
|
||||
agent = _bare_agent()
|
||||
agent._delegate_depth = 0 # top-level agent
|
||||
|
||||
AIAgent._spawn_background_review(
|
||||
agent,
|
||||
messages_snapshot=[{"role": "user", "content": "hello"}],
|
||||
review_memory=True,
|
||||
)
|
||||
|
||||
assert len(forks) == 1, "top-level review must still spawn the fork"
|
||||
|
||||
|
||||
def test_background_review_disabled_skips_automatic_spawn(monkeypatch):
|
||||
"""``auxiliary.background_review.enabled: false`` must skip automatic
|
||||
post-turn forks while leaving ``/refine`` (focus set) working (#87250)."""
|
||||
from unittest.mock import patch
|
||||
|
||||
forks = []
|
||||
|
||||
class FakeReviewAgent:
|
||||
def __init__(self, **kwargs):
|
||||
forks.append(kwargs)
|
||||
|
||||
def run_conversation(self, **kwargs):
|
||||
pass
|
||||
|
||||
def shutdown_memory_provider(self):
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr(run_agent_module, "AIAgent", FakeReviewAgent)
|
||||
monkeypatch.setattr(run_agent_module.threading, "Thread", ImmediateThread)
|
||||
|
||||
agent = _bare_agent()
|
||||
agent._delegate_depth = 0
|
||||
cfg = {"auxiliary": {"background_review": {"enabled": False}}}
|
||||
|
||||
with patch("hermes_cli.config.load_config_readonly", return_value=cfg):
|
||||
AIAgent._spawn_background_review(
|
||||
agent,
|
||||
messages_snapshot=[{"role": "user", "content": "hello"}],
|
||||
review_memory=True,
|
||||
)
|
||||
assert forks == [], "automatic review must not spawn when disabled"
|
||||
|
||||
AIAgent._spawn_background_review(
|
||||
agent,
|
||||
messages_snapshot=[{"role": "user", "content": "hello"}],
|
||||
review_memory=True,
|
||||
focus="save the deploy workflow",
|
||||
)
|
||||
assert len(forks) == 1, "/refine must still run when enabled=false"
|
||||
|
||||
|
||||
def test_background_review_explicit_focus_runs_even_in_subagent(monkeypatch):
|
||||
"""An explicit ``/refine`` (``focus`` set) is a deliberate user request and
|
||||
is honored regardless of depth — only the automatic post-turn review is
|
||||
suppressed in subagents."""
|
||||
forks = []
|
||||
|
||||
class FakeReviewAgent:
|
||||
def __init__(self, **kwargs):
|
||||
forks.append(kwargs)
|
||||
|
||||
def run_conversation(self, **kwargs):
|
||||
pass
|
||||
|
||||
def shutdown_memory_provider(self):
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr(run_agent_module, "AIAgent", FakeReviewAgent)
|
||||
monkeypatch.setattr(run_agent_module.threading, "Thread", ImmediateThread)
|
||||
|
||||
agent = _bare_agent()
|
||||
agent._delegate_depth = 2
|
||||
|
||||
AIAgent._spawn_background_review(
|
||||
agent,
|
||||
messages_snapshot=[{"role": "user", "content": "hello"}],
|
||||
review_skills=True,
|
||||
focus="save the deploy workflow as a skill",
|
||||
)
|
||||
|
||||
assert len(forks) == 1, "explicit focus review must run even in a subagent"
|
||||
|
||||
|
||||
def test_background_review_registers_before_start_runs_and_cleans_up(monkeypatch):
|
||||
"""The parent must own a unique review run before the worker can start."""
|
||||
seen = {}
|
||||
|
||||
class RecordingReviewAgent(FakeReviewAgent):
|
||||
def run_conversation(self, **kwargs):
|
||||
seen["run"] = agent._background_review_run
|
||||
seen["active_children_during_run"] = list(agent._active_children)
|
||||
seen["background_review_agent_during_run"] = agent._background_review_agent
|
||||
|
||||
monkeypatch.setattr(run_agent_module, "AIAgent", RecordingReviewAgent)
|
||||
CapturingThread.targets = []
|
||||
monkeypatch.setattr(run_agent_module.threading, "Thread", CapturingThread)
|
||||
|
||||
agent = _bare_agent()
|
||||
|
||||
AIAgent._spawn_background_review(
|
||||
agent,
|
||||
messages_snapshot=[{"role": "user", "content": "hello"}],
|
||||
review_memory=True,
|
||||
)
|
||||
|
||||
run = agent._background_review_run
|
||||
assert run is not None
|
||||
assert len(CapturingThread.targets) == 1
|
||||
assert not run.request_done.is_set()
|
||||
|
||||
observed_done = ObservedEvent()
|
||||
run.request_done = observed_done
|
||||
CapturingThread.targets[0]()
|
||||
|
||||
fork = seen["background_review_agent_during_run"]
|
||||
assert fork is not None
|
||||
assert seen["run"] is run
|
||||
assert seen["active_children_during_run"] == [fork]
|
||||
assert observed_done.is_set()
|
||||
assert observed_done.set_calls == 1
|
||||
assert agent._background_review_run is None
|
||||
assert agent._background_review_agent is None
|
||||
assert agent._active_children == []
|
||||
|
||||
|
||||
def test_background_review_snapshot_isolated_from_live_nested_messages():
|
||||
"""A review must not mutate the persisted/live transcript through aliases."""
|
||||
original = [{
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "answer"}],
|
||||
"tool_calls": [{
|
||||
"id": "call-1",
|
||||
"function": {"name": "read_file", "arguments": '{"path":"x"}'},
|
||||
}],
|
||||
}]
|
||||
|
||||
from agent.turn_finalizer import _clone_background_review_messages
|
||||
|
||||
snapshot = _clone_background_review_messages(original)
|
||||
snapshot[0]["content"][0]["text"] = "review mutation"
|
||||
snapshot[0]["tool_calls"][0]["function"]["arguments"] = "{}"
|
||||
|
||||
assert original[0]["content"][0]["text"] == "answer"
|
||||
assert original[0]["tool_calls"][0]["function"]["arguments"] == '{"path":"x"}'
|
||||
|
||||
|
||||
def test_live_turn_waits_for_review_exit_before_relay_and_turn_context(monkeypatch):
|
||||
"""The outer production wrapper waits before same-session instrumentation."""
|
||||
review_entered = threading.Event()
|
||||
review_returned = threading.Event()
|
||||
allow_review_return = threading.Event()
|
||||
interrupted = threading.Event()
|
||||
boundary_reached = threading.Event()
|
||||
seen = {}
|
||||
|
||||
class BlockingReviewAgent(FakeReviewAgent):
|
||||
def interrupt(self, message=None):
|
||||
seen["interrupt_message"] = message
|
||||
interrupted.set()
|
||||
|
||||
def run_conversation(self, **kwargs):
|
||||
review_entered.set()
|
||||
assert allow_review_return.wait(2.0)
|
||||
review_returned.set()
|
||||
|
||||
monkeypatch.setattr(run_agent_module, "AIAgent", BlockingReviewAgent)
|
||||
CapturingThread.targets = []
|
||||
monkeypatch.setattr(run_agent_module.threading, "Thread", CapturingThread)
|
||||
|
||||
agent = _bare_agent()
|
||||
AIAgent._spawn_background_review(
|
||||
agent,
|
||||
messages_snapshot=[{"role": "user", "content": "hello"}],
|
||||
review_memory=True,
|
||||
)
|
||||
run = agent._background_review_run
|
||||
assert run is not None
|
||||
observed_done = ObservedEvent()
|
||||
run.request_done = observed_done
|
||||
|
||||
monkeypatch.setattr(run_agent_module.threading, "Thread", _REAL_THREAD)
|
||||
worker = _REAL_THREAD(target=CapturingThread.targets[0], daemon=True)
|
||||
worker.start()
|
||||
assert review_entered.wait(2.0)
|
||||
|
||||
def on_boundary():
|
||||
seen["review_returned_at_boundary"] = review_returned.is_set()
|
||||
boundary_reached.set()
|
||||
|
||||
_install_live_turn_boundary(monkeypatch, on_boundary)
|
||||
relay_calls = _install_relay_recorder(monkeypatch, run)
|
||||
live_result = {}
|
||||
live = _REAL_THREAD(
|
||||
target=_run_wrapped_live_turn_to_boundary,
|
||||
args=(agent, live_result),
|
||||
daemon=True,
|
||||
)
|
||||
live.start()
|
||||
|
||||
assert interrupted.wait(2.0)
|
||||
wait_started = observed_done.wait_started.wait(2.0)
|
||||
relay_calls_before_ack = list(relay_calls)
|
||||
allow_review_return.set()
|
||||
worker.join(timeout=2.0)
|
||||
live.join(timeout=2.0)
|
||||
|
||||
assert not worker.is_alive()
|
||||
assert not live.is_alive()
|
||||
assert wait_started
|
||||
assert relay_calls_before_ack == []
|
||||
assert relay_calls == [
|
||||
("acquire", True),
|
||||
("begin", True),
|
||||
("start_task_run", True),
|
||||
]
|
||||
assert boundary_reached.is_set()
|
||||
assert seen["interrupt_message"] == "superseded by a new live turn"
|
||||
assert seen["review_returned_at_boundary"] is True
|
||||
assert live_result == {"boundary_reached": True}
|
||||
|
||||
|
||||
def test_live_turn_cancels_review_during_startup_before_provider(monkeypatch):
|
||||
"""A review cancelled before its worker runs must never call its provider."""
|
||||
provider_calls = []
|
||||
boundary_reached = threading.Event()
|
||||
|
||||
class RecordingReviewAgent(FakeReviewAgent):
|
||||
def run_conversation(self, **kwargs):
|
||||
provider_calls.append(kwargs)
|
||||
|
||||
monkeypatch.setattr(run_agent_module, "AIAgent", RecordingReviewAgent)
|
||||
CapturingThread.targets = []
|
||||
monkeypatch.setattr(run_agent_module.threading, "Thread", CapturingThread)
|
||||
|
||||
agent = _bare_agent()
|
||||
AIAgent._spawn_background_review(
|
||||
agent,
|
||||
messages_snapshot=[{"role": "user", "content": "hello"}],
|
||||
review_memory=True,
|
||||
)
|
||||
run = agent._background_review_run
|
||||
assert run is not None
|
||||
|
||||
_install_live_turn_boundary(monkeypatch, boundary_reached.set)
|
||||
relay_calls = _install_relay_recorder(monkeypatch, run)
|
||||
live_result = {}
|
||||
live = _REAL_THREAD(
|
||||
target=_run_wrapped_live_turn_to_boundary,
|
||||
args=(agent, live_result),
|
||||
daemon=True,
|
||||
)
|
||||
live.start()
|
||||
assert run.cancel_requested.wait(2.0)
|
||||
|
||||
worker = _REAL_THREAD(target=CapturingThread.targets[0], daemon=True)
|
||||
worker.start()
|
||||
worker.join(timeout=2.0)
|
||||
live.join(timeout=2.0)
|
||||
|
||||
assert not worker.is_alive()
|
||||
assert not live.is_alive()
|
||||
assert boundary_reached.is_set()
|
||||
assert provider_calls == []
|
||||
assert run.request_done.is_set()
|
||||
assert relay_calls == [
|
||||
("acquire", True),
|
||||
("begin", True),
|
||||
("start_task_run", True),
|
||||
]
|
||||
assert live_result == {"boundary_reached": True}
|
||||
|
||||
|
||||
def test_live_turn_proceeds_when_review_acknowledgement_times_out(monkeypatch):
|
||||
"""A broken review abort path must not block the foreground indefinitely.
|
||||
The live turn proceeds after the bounded wait, retaining foreground priority.
|
||||
"""
|
||||
import time
|
||||
|
||||
import agent.background_review as background_review_module
|
||||
|
||||
review_entered = threading.Event()
|
||||
interrupt_entered = threading.Event()
|
||||
interrupt_returned = threading.Event()
|
||||
allow_interrupt_return = threading.Event()
|
||||
allow_review_return = threading.Event()
|
||||
|
||||
class WedgedReviewAgent(FakeReviewAgent):
|
||||
def run_conversation(self, **kwargs):
|
||||
review_entered.set()
|
||||
allow_review_return.wait(5.0)
|
||||
|
||||
def interrupt(self, message=None):
|
||||
interrupt_entered.set()
|
||||
allow_interrupt_return.wait(5.0)
|
||||
interrupt_returned.set()
|
||||
|
||||
monkeypatch.setattr(run_agent_module, "AIAgent", WedgedReviewAgent)
|
||||
CapturingThread.targets = []
|
||||
monkeypatch.setattr(run_agent_module.threading, "Thread", CapturingThread)
|
||||
monkeypatch.setattr(
|
||||
background_review_module,
|
||||
"_BACKGROUND_REVIEW_CANCEL_TIMEOUT_SECONDS",
|
||||
0.01,
|
||||
raising=False,
|
||||
)
|
||||
|
||||
agent = _bare_agent()
|
||||
AIAgent._spawn_background_review(
|
||||
agent,
|
||||
messages_snapshot=[{"role": "user", "content": "hello"}],
|
||||
review_memory=True,
|
||||
)
|
||||
run = agent._background_review_run
|
||||
assert run is not None
|
||||
monkeypatch.setattr(run_agent_module.threading, "Thread", _REAL_THREAD)
|
||||
worker = _REAL_THREAD(target=CapturingThread.targets[0], daemon=True)
|
||||
worker.start()
|
||||
assert review_entered.wait(2.0)
|
||||
|
||||
boundary_calls = []
|
||||
_install_live_turn_boundary(
|
||||
monkeypatch, lambda: boundary_calls.append(True)
|
||||
)
|
||||
relay_calls = _install_relay_recorder(monkeypatch, run)
|
||||
|
||||
started = time.monotonic()
|
||||
live_result = {}
|
||||
live = _REAL_THREAD(
|
||||
target=_run_wrapped_live_turn_to_boundary,
|
||||
args=(agent, live_result),
|
||||
daemon=True,
|
||||
)
|
||||
live.start()
|
||||
live.join(timeout=5.0)
|
||||
|
||||
elapsed = time.monotonic() - started
|
||||
|
||||
allow_interrupt_return.set()
|
||||
allow_review_return.set()
|
||||
worker.join(timeout=2.0)
|
||||
|
||||
assert elapsed < 2.0
|
||||
assert interrupt_entered.is_set()
|
||||
assert interrupt_returned.wait(2.0)
|
||||
assert not worker.is_alive()
|
||||
assert not live.is_alive()
|
||||
# Foreground retains priority: Relay/turn-context proceed even though
|
||||
# the review did not acknowledge within the bounded deadline.
|
||||
assert boundary_calls == [True]
|
||||
assert live_result == {"boundary_reached": True}
|
||||
assert relay_calls == [
|
||||
("acquire", False),
|
||||
("begin", False),
|
||||
("start_task_run", False),
|
||||
]
|
||||
assert agent.session_id == "test-session"
|
||||
|
||||
|
||||
def test_live_turn_interrupts_legacy_review_but_keeps_foreground_priority(monkeypatch):
|
||||
"""Legacy stubs are interrupted without turning review into a user blocker."""
|
||||
interrupts = []
|
||||
interrupt_called = threading.Event()
|
||||
|
||||
class LegacyReviewAgent:
|
||||
def interrupt(self, message=None):
|
||||
interrupts.append(message)
|
||||
interrupt_called.set()
|
||||
|
||||
agent = _bare_agent()
|
||||
del agent._background_review_run
|
||||
agent._background_review_agent = LegacyReviewAgent()
|
||||
boundary_calls = []
|
||||
_install_live_turn_boundary(
|
||||
monkeypatch, lambda: boundary_calls.append(True)
|
||||
)
|
||||
relay_calls = _install_relay_recorder(monkeypatch)
|
||||
|
||||
live_result = {}
|
||||
live = _REAL_THREAD(
|
||||
target=_run_wrapped_live_turn_to_boundary,
|
||||
args=(agent, live_result),
|
||||
daemon=True,
|
||||
)
|
||||
live.start()
|
||||
live.join(timeout=5.0)
|
||||
|
||||
assert interrupt_called.wait(2.0)
|
||||
assert interrupts == ["superseded by a new live turn"]
|
||||
assert not live.is_alive()
|
||||
assert boundary_calls == [True]
|
||||
assert live_result == {"boundary_reached": True}
|
||||
assert relay_calls == [
|
||||
("acquire", False),
|
||||
("begin", False),
|
||||
("start_task_run", False),
|
||||
]
|
||||
assert agent.session_id == "test-session"
|
||||
|
||||
|
||||
def test_stale_review_cleanup_cannot_clear_or_signal_newer_review(monkeypatch):
|
||||
"""A retired worker's late cleanup must be scoped to its own run identity."""
|
||||
first_cleanup_entered = threading.Event()
|
||||
allow_first_cleanup = threading.Event()
|
||||
instance_count = 0
|
||||
|
||||
class BlockingCleanupReviewAgent(FakeReviewAgent):
|
||||
def __init__(self, **kwargs):
|
||||
nonlocal instance_count
|
||||
super().__init__(**kwargs)
|
||||
self.index = instance_count
|
||||
instance_count += 1
|
||||
|
||||
def release_clients(self):
|
||||
if self.index == 0:
|
||||
first_cleanup_entered.set()
|
||||
assert allow_first_cleanup.wait(2.0)
|
||||
|
||||
monkeypatch.setattr(run_agent_module, "AIAgent", BlockingCleanupReviewAgent)
|
||||
CapturingThread.targets = []
|
||||
monkeypatch.setattr(run_agent_module.threading, "Thread", CapturingThread)
|
||||
|
||||
agent = _bare_agent()
|
||||
AIAgent._spawn_background_review(
|
||||
agent,
|
||||
messages_snapshot=[{"role": "user", "content": "first"}],
|
||||
review_memory=True,
|
||||
)
|
||||
first_run = agent._background_review_run
|
||||
first_worker = _REAL_THREAD(target=CapturingThread.targets[0], daemon=True)
|
||||
first_worker.start()
|
||||
assert first_cleanup_entered.wait(2.0)
|
||||
assert first_run.request_done.is_set()
|
||||
|
||||
AIAgent._spawn_background_review(
|
||||
agent,
|
||||
messages_snapshot=[{"role": "user", "content": "second"}],
|
||||
review_memory=True,
|
||||
)
|
||||
second_run = agent._background_review_run
|
||||
second_target = CapturingThread.targets[1]
|
||||
assert second_run is not first_run
|
||||
assert not second_run.request_done.is_set()
|
||||
|
||||
allow_first_cleanup.set()
|
||||
first_worker.join(timeout=2.0)
|
||||
|
||||
assert not first_worker.is_alive()
|
||||
assert agent._background_review_run is second_run
|
||||
assert not second_run.request_done.is_set()
|
||||
|
||||
second_target()
|
||||
assert second_run.request_done.is_set()
|
||||
assert agent._background_review_run is None
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# memory_notifications mode: off | on | verbose
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
import json as _json
|
||||
|
||||
from agent.background_review import summarize_background_review_actions
|
||||
|
||||
|
||||
def _memory_add_review():
|
||||
"""A minimal review transcript: one memory add (assistant call + tool result)."""
|
||||
return [
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_mem1",
|
||||
"function": {
|
||||
"name": "memory",
|
||||
"arguments": _json.dumps(
|
||||
{
|
||||
"action": "add",
|
||||
"target": "memory",
|
||||
"content": "User prefers terse replies",
|
||||
}
|
||||
),
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_mem1",
|
||||
"content": _json.dumps(
|
||||
{"success": True, "message": "Entry added.", "target": "memory"}
|
||||
),
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def _skill_patch_review():
|
||||
return [
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_skill1",
|
||||
"function": {
|
||||
"name": "skill_manage",
|
||||
"arguments": _json.dumps(
|
||||
{"action": "patch", "name": "demo", "old_string": "a", "new_string": "b"}
|
||||
),
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_skill1",
|
||||
"content": _json.dumps(
|
||||
{
|
||||
"success": True,
|
||||
"message": "Patched SKILL.md in skill 'demo' (1 replacement).",
|
||||
"_change": {"old": "a", "new": "b"},
|
||||
}
|
||||
),
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def test_memory_notifications_off_returns_nothing():
|
||||
actions = summarize_background_review_actions(
|
||||
_memory_add_review(), [], notification_mode="off"
|
||||
)
|
||||
assert actions == []
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_skill_patch_off_silent_verbose_shows_diff():
|
||||
assert (
|
||||
summarize_background_review_actions(
|
||||
_skill_patch_review(), [], notification_mode="off"
|
||||
)
|
||||
== []
|
||||
)
|
||||
verbose = summarize_background_review_actions(
|
||||
_skill_patch_review(), [], notification_mode="verbose"
|
||||
)
|
||||
assert len(verbose) == 1
|
||||
assert "demo" in verbose[0] and "→" in verbose[0]
|
||||
@@ -0,0 +1,579 @@
|
||||
"""Tests that the background review fork inherits the parent's cached system prompt.
|
||||
|
||||
Regression coverage for issue #25322 (and PR #17276's first root cause): the
|
||||
background review's outbound HTTP request must carry the same system bytes as
|
||||
the parent's so Anthropic/OpenRouter's exact-prefix cache key matches.
|
||||
|
||||
Without this, every review rebuilds the system prompt from scratch — fresh
|
||||
``_hermes_now()`` timestamp, fresh ``session_id``, and a different skills
|
||||
prompt under the (former) narrow toolset — and the prefix-cache miss costs
|
||||
roughly the full uncached system-prompt cost per nudge (~26% end-to-end on
|
||||
Sonnet 4.5 per the contributor's measurement).
|
||||
"""
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
|
||||
def _make_agent_stub(agent_cls):
|
||||
"""Create a minimal AIAgent-like object with just enough state for _spawn_background_review."""
|
||||
agent = object.__new__(agent_cls)
|
||||
agent.model = "test-model"
|
||||
agent.platform = "test"
|
||||
agent.provider = "openai"
|
||||
agent.session_id = "sess-123"
|
||||
agent.quiet_mode = True
|
||||
agent._memory_store = None
|
||||
agent._memory_enabled = True
|
||||
agent._user_profile_enabled = False
|
||||
agent._memory_nudge_interval = 5
|
||||
agent._skill_nudge_interval = 5
|
||||
agent.background_review_callback = None
|
||||
agent.status_callback = None
|
||||
agent._cached_system_prompt = (
|
||||
"PARENT-SYSTEM-PROMPT-BYTES — must be inherited verbatim "
|
||||
"for prefix-cache parity"
|
||||
)
|
||||
agent.ephemeral_system_prompt = (
|
||||
"WebUI session context:\n- Pinned per-request gateway context"
|
||||
)
|
||||
import datetime as _dt
|
||||
agent.session_start = _dt.datetime(2026, 1, 1, 12, 0, 0)
|
||||
agent._MEMORY_REVIEW_PROMPT = "review memory"
|
||||
agent._SKILL_REVIEW_PROMPT = "review skills"
|
||||
agent._COMBINED_REVIEW_PROMPT = "review both"
|
||||
# Non-None so the test catches a missing-kwarg regression.
|
||||
agent.enabled_toolsets = ["memory", "skills", "terminal"]
|
||||
agent.disabled_toolsets = ["spotify", "feishu_doc"]
|
||||
# Non-None so the test catches reasoning_config NOT being inherited —
|
||||
# which would put the fork into a different Anthropic cache namespace.
|
||||
agent.reasoning_config = {"enabled": True, "effort": "medium"}
|
||||
# Non-empty so tests catch prefill/provider-routing NOT being inherited —
|
||||
# prefills sit right after the system message in the request body, and
|
||||
# OpenRouter provider pins decide WHICH upstream's cache gets hit.
|
||||
agent.prefill_messages = [{"role": "user", "content": "prefill turn"}]
|
||||
agent.providers_allowed = ["anthropic"]
|
||||
agent.providers_ignored = None
|
||||
agent.providers_order = None
|
||||
agent.provider_sort = "throughput"
|
||||
agent.provider_require_parameters = False
|
||||
agent.provider_data_collection = None
|
||||
return agent
|
||||
|
||||
|
||||
class _SyncThread:
|
||||
"""Drop-in replacement for threading.Thread that runs the target inline."""
|
||||
|
||||
def __init__(self, *, target=None, daemon=None, name=None):
|
||||
self._target = target
|
||||
|
||||
def start(self):
|
||||
if self._target:
|
||||
self._target()
|
||||
|
||||
|
||||
def _make_recorder_class(captured=None, record_on_run=()):
|
||||
"""Build a Recorder class standing in for the review-fork AIAgent.
|
||||
|
||||
Keeps the stub attribute list in ONE place: when
|
||||
``_spawn_background_review`` starts touching a new fork attribute, only
|
||||
this factory needs the extra stub — not one copy per test.
|
||||
|
||||
``captured`` (dict): if given, ``__init__`` stores the full constructor
|
||||
kwargs under ``captured["init_kwargs"]`` so tests can assert on both
|
||||
kwarg values and kwarg *presence*.
|
||||
``record_on_run``: instance attribute names copied into ``captured`` when
|
||||
``run_conversation`` fires — for values the production code assigns
|
||||
after construction.
|
||||
"""
|
||||
|
||||
class _Recorder:
|
||||
def __init__(self, *args, **kwargs):
|
||||
if captured is not None:
|
||||
captured["init_kwargs"] = dict(kwargs)
|
||||
self._cached_system_prompt = None
|
||||
self._memory_write_origin = None
|
||||
self._memory_write_context = None
|
||||
self._memory_store = None
|
||||
self._memory_enabled = None
|
||||
self._user_profile_enabled = None
|
||||
self._memory_nudge_interval = None
|
||||
self._skill_nudge_interval = None
|
||||
self.suppress_status_output = None
|
||||
self.session_start = None
|
||||
self.session_id = None
|
||||
self.tools = None
|
||||
self.valid_tool_names = set()
|
||||
self._tool_snapshot_generation = 0
|
||||
self.ephemeral_system_prompt = kwargs.get("ephemeral_system_prompt")
|
||||
|
||||
def run_conversation(self, *args, **kwargs):
|
||||
if captured is not None:
|
||||
for _name in record_on_run:
|
||||
captured[_name] = getattr(self, _name)
|
||||
raise RuntimeError(
|
||||
"stop after recording — don't actually call the API"
|
||||
)
|
||||
|
||||
def shutdown_memory_provider(self):
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
return _Recorder
|
||||
|
||||
|
||||
def test_review_fork_inherits_parent_cached_system_prompt():
|
||||
"""The review fork's _cached_system_prompt must equal the parent's byte-for-byte.
|
||||
|
||||
Anthropic's prefix cache keys on exact bytes; any divergence (timestamp
|
||||
minute tick, fresh session_id, narrower skills_prompt) shifts the key
|
||||
and forces a full re-cache. Inheriting the parent's cached prompt is
|
||||
the cheap, mechanical fix.
|
||||
"""
|
||||
import run_agent
|
||||
|
||||
agent = _make_agent_stub(run_agent.AIAgent)
|
||||
|
||||
captured = {}
|
||||
parent_prompt = agent._cached_system_prompt
|
||||
|
||||
_Recorder = _make_recorder_class()
|
||||
|
||||
with patch.object(run_agent, "AIAgent", _Recorder), \
|
||||
patch("threading.Thread", _SyncThread):
|
||||
# The production code assigns _cached_system_prompt AFTER __init__,
|
||||
# so wrap the recorder's __setattr__ to see that post-construction
|
||||
# write from _spawn_background_review.
|
||||
orig_setattr = _Recorder.__setattr__
|
||||
|
||||
def _spy_setattr(self, name, value):
|
||||
if name == "_cached_system_prompt":
|
||||
captured["written_prompt"] = value
|
||||
orig_setattr(self, name, value)
|
||||
|
||||
with patch.object(_Recorder, "__setattr__", _spy_setattr):
|
||||
agent._spawn_background_review(
|
||||
messages_snapshot=[],
|
||||
review_memory=True,
|
||||
review_skills=False,
|
||||
)
|
||||
|
||||
assert "written_prompt" in captured, (
|
||||
"_spawn_background_review never assigned _cached_system_prompt on the review agent"
|
||||
)
|
||||
assert captured["written_prompt"] == parent_prompt, (
|
||||
f"Review fork's _cached_system_prompt diverged from parent's. "
|
||||
f"Got {captured['written_prompt']!r}, expected {parent_prompt!r}. "
|
||||
"This breaks Anthropic/OpenRouter prefix-cache parity (#25322)."
|
||||
)
|
||||
|
||||
|
||||
def test_review_fork_inherits_parent_ephemeral_system_prompt():
|
||||
"""The fork must send the parent's complete effective system prompt.
|
||||
|
||||
Gateway session context is appended through ``ephemeral_system_prompt`` at
|
||||
API-call time, outside ``_cached_system_prompt``. Copying only the cached
|
||||
base therefore makes every background review diverge at the gateway block
|
||||
and miss the parent's warm prefix cache.
|
||||
"""
|
||||
import run_agent
|
||||
|
||||
agent = _make_agent_stub(run_agent.AIAgent)
|
||||
captured = {}
|
||||
_Recorder = _make_recorder_class(
|
||||
captured,
|
||||
record_on_run=("_cached_system_prompt", "ephemeral_system_prompt"),
|
||||
)
|
||||
|
||||
with patch.object(run_agent, "AIAgent", _Recorder), \
|
||||
patch("threading.Thread", _SyncThread):
|
||||
agent._spawn_background_review(
|
||||
messages_snapshot=[],
|
||||
review_memory=True,
|
||||
review_skills=False,
|
||||
)
|
||||
|
||||
# Pairwise asserts: stronger than comparing a locally re-joined
|
||||
# "effective" prompt (which would re-implement the production join and
|
||||
# silently keep passing if the separator ever changed — and would compare
|
||||
# equal for cached="A\n\nB"/ephemeral="" vs cached="A"/ephemeral="B").
|
||||
assert captured["_cached_system_prompt"] == agent._cached_system_prompt
|
||||
assert captured["ephemeral_system_prompt"] == agent.ephemeral_system_prompt
|
||||
|
||||
|
||||
def test_review_fork_inherits_prefill_and_provider_routing():
|
||||
"""Non-routed fork must inherit prefill messages and OpenRouter pins.
|
||||
|
||||
Prefill messages are inserted right after the system message at
|
||||
API-call time, so omitting them diverges the fork's request body from
|
||||
the parent's warm prefix at message index 1. OpenRouter provider pins
|
||||
(providers_allowed/order/sort/...) decide which UPSTREAM provider serves
|
||||
the request — prompt caches live per upstream, so an unpinned fork can
|
||||
be routed to a different upstream and miss even a byte-identical prefix.
|
||||
"""
|
||||
import run_agent
|
||||
|
||||
agent = _make_agent_stub(run_agent.AIAgent)
|
||||
captured = {}
|
||||
_Recorder = _make_recorder_class(captured)
|
||||
|
||||
with patch.object(run_agent, "AIAgent", _Recorder), \
|
||||
patch("threading.Thread", _SyncThread):
|
||||
agent._spawn_background_review(
|
||||
messages_snapshot=[],
|
||||
review_memory=True,
|
||||
review_skills=False,
|
||||
)
|
||||
|
||||
init_kwargs = captured.get("init_kwargs", {})
|
||||
assert init_kwargs.get("prefill_messages") == agent.prefill_messages
|
||||
# Must be a DEEP copy: the fork's unicode-error recovery
|
||||
# (_sanitize_messages_surrogates) mutates prefill dicts in place, so
|
||||
# aliased dicts would let the fork rewrite the parent's prefill bytes
|
||||
# — silently breaking the parent's own warm prefix.
|
||||
assert (
|
||||
init_kwargs["prefill_messages"][0] is not agent.prefill_messages[0]
|
||||
), "fork prefill aliases the parent's dicts (needs deepcopy)"
|
||||
assert init_kwargs.get("providers_allowed") == agent.providers_allowed
|
||||
assert init_kwargs.get("provider_sort") == agent.provider_sort
|
||||
|
||||
|
||||
def test_review_fork_pins_session_start_and_session_id():
|
||||
"""Defensive complement to cached-system-prompt inheritance.
|
||||
|
||||
Even though ``_cached_system_prompt`` inheritance short-circuits the
|
||||
normal rebuild path, pinning ``session_start`` and ``session_id`` to
|
||||
the parent's guarantees byte-identical output from any code path that
|
||||
re-renders parts of the system prompt (compression, plugin hooks).
|
||||
"""
|
||||
import run_agent
|
||||
|
||||
agent = _make_agent_stub(run_agent.AIAgent)
|
||||
|
||||
captured = {}
|
||||
_Recorder = _make_recorder_class(
|
||||
captured, record_on_run=("session_start", "session_id")
|
||||
)
|
||||
|
||||
with patch.object(run_agent, "AIAgent", _Recorder), \
|
||||
patch("threading.Thread", _SyncThread):
|
||||
agent._spawn_background_review(
|
||||
messages_snapshot=[],
|
||||
review_memory=True,
|
||||
review_skills=False,
|
||||
)
|
||||
|
||||
assert captured.get("session_start") == agent.session_start, (
|
||||
"Review fork did not inherit parent's session_start — "
|
||||
"system-prompt rebuild paths would diverge."
|
||||
)
|
||||
assert captured.get("session_id") == agent.session_id, (
|
||||
"Review fork did not inherit parent's session_id — "
|
||||
"system-prompt rebuild paths would diverge."
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_routed_review_fork_does_not_inherit_reasoning_config():
|
||||
"""Routed aux path: the fork must NOT inherit the parent's reasoning_config.
|
||||
|
||||
When ``auxiliary.background_review.{provider,model}`` routes the review
|
||||
to a different model, cache parity is moot (the cache is cold on that
|
||||
model regardless) and the parent's effort vocabulary may be invalid for
|
||||
the routed model/provider (OpenRouter ``extra_body.reasoning.effort`` is
|
||||
forwarded unclamped; codex_responses passes ``max``/``ultra`` through
|
||||
unmapped except on gpt-5.6/xAI). The routed fork must fall back to
|
||||
provider defaults, mirroring the ``not _routed`` gate on
|
||||
``_cached_system_prompt`` inheritance.
|
||||
"""
|
||||
import run_agent
|
||||
import agent.background_review as bg_review
|
||||
|
||||
agent_stub = _make_agent_stub(run_agent.AIAgent)
|
||||
|
||||
captured = {}
|
||||
_Recorder = _make_recorder_class(captured)
|
||||
|
||||
routed_runtime = {
|
||||
"provider": "openrouter",
|
||||
"model": "aux-cheap-model",
|
||||
"api_key": "test-key",
|
||||
"base_url": None,
|
||||
"api_mode": None,
|
||||
"credential_pool": None,
|
||||
"request_overrides": {},
|
||||
"max_tokens": None,
|
||||
"command": None,
|
||||
"args": [],
|
||||
"routed": True,
|
||||
}
|
||||
|
||||
with patch.object(run_agent, "AIAgent", _Recorder), \
|
||||
patch.object(bg_review, "_resolve_review_runtime",
|
||||
return_value=routed_runtime), \
|
||||
patch("threading.Thread", _SyncThread):
|
||||
agent_stub._spawn_background_review(
|
||||
messages_snapshot=[],
|
||||
review_memory=True,
|
||||
review_skills=False,
|
||||
)
|
||||
|
||||
init_kwargs = captured.get("init_kwargs", {})
|
||||
assert "reasoning_config" not in init_kwargs, (
|
||||
f"Routed review fork was passed the parent's reasoning_config "
|
||||
f"({init_kwargs.get('reasoning_config')!r}). On the routed path the "
|
||||
"cache is cold (no parity benefit) and the parent's effort value may "
|
||||
"be invalid for the routed model/provider — it must be omitted so "
|
||||
"the fork uses provider defaults."
|
||||
)
|
||||
# The whole cache-parity kwarg family shares the same ``not _routed``
|
||||
# gate — a future refactor hoisting any of them out of the gate must
|
||||
# fail here, not silently ship parent-only context to a foreign model.
|
||||
for _gated in (
|
||||
"ephemeral_system_prompt",
|
||||
"prefill_messages",
|
||||
"providers_allowed",
|
||||
"provider_sort",
|
||||
):
|
||||
assert _gated not in init_kwargs, (
|
||||
f"Routed review fork was passed parent-only kwarg {_gated!r}; "
|
||||
"cache-parity inheritance must stay behind the not-routed gate."
|
||||
)
|
||||
|
||||
|
||||
def test_review_fork_inherited_tools_survive_compaction_refresh():
|
||||
"""Inherited tools survive mid-review compaction refresh (#103579).
|
||||
|
||||
Acceptance criterion 1 requires the fork to advertise the same tools[] as
|
||||
the parent when targeting the same cache scope. Mid-review compaction
|
||||
boundaries invoke refresh_agent_mcp_tools(content_aware=True), which re-reads
|
||||
the live registry and drops memory provider tools unless the snapshot
|
||||
generation staleness check refuses the rebuild.
|
||||
"""
|
||||
import run_agent
|
||||
from agent.background_review import build_cache_parity_fork
|
||||
from tools.mcp_tool_agent import refresh_agent_mcp_tools
|
||||
|
||||
agent = _make_agent_stub(run_agent.AIAgent)
|
||||
parent_tools = [
|
||||
{"type": "function", "function": {"name": "terminal_command"}},
|
||||
{"type": "function", "function": {"name": "read_file"}},
|
||||
{"type": "function", "function": {"name": "memory"}},
|
||||
{"type": "function", "function": {"name": "fact_store"}},
|
||||
{"type": "function", "function": {"name": "fact_feedback"}},
|
||||
]
|
||||
agent.tools = parent_tools
|
||||
|
||||
_Recorder = _make_recorder_class()
|
||||
|
||||
with patch.object(run_agent, "AIAgent", _Recorder):
|
||||
fork, _rt, routed = build_cache_parity_fork(agent, max_iterations=5)
|
||||
assert not routed
|
||||
assert fork.tools == parent_tools
|
||||
# Deep copy: the fork's later in-place tool edits must not leak into the parent's array.
|
||||
assert fork.tools is not parent_tools and fork.tools[0] is not parent_tools[0]
|
||||
|
||||
# Simulate mid-review compaction boundary tool refresh
|
||||
added = refresh_agent_mcp_tools(fork, content_aware=True)
|
||||
assert added == set()
|
||||
assert [t["function"]["name"] for t in fork.tools] == [
|
||||
"terminal_command", "read_file", "memory", "fact_store", "fact_feedback"
|
||||
]
|
||||
assert fork.valid_tool_names == {
|
||||
"terminal_command", "read_file", "memory", "fact_store", "fact_feedback"
|
||||
}
|
||||
|
||||
|
||||
def test_unrouted_review_fork_inherits_empty_tool_surface():
|
||||
"""Empty parent tools[] is a valid snapshot and must be copied and frozen (#103579).
|
||||
|
||||
If no tools pass availability when the parent is constructed (parent.tools = []),
|
||||
the unrouted fork must inherit an empty list and freeze _tool_snapshot_generation.
|
||||
This guarantees late MCP/plugin tools discovered during fork construction or
|
||||
mid-review compaction do not break cache parity against the parent's empty surface.
|
||||
"""
|
||||
import run_agent
|
||||
from agent.background_review import build_cache_parity_fork
|
||||
from tools.mcp_tool_agent import refresh_agent_mcp_tools
|
||||
|
||||
agent = _make_agent_stub(run_agent.AIAgent)
|
||||
agent.tools = []
|
||||
|
||||
_BaseRecorder = _make_recorder_class()
|
||||
|
||||
class _RecorderWithLateTool(_BaseRecorder):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
# Simulate a late tool appearing in the constructor result before inheritance
|
||||
self.tools = [{"type": "function", "function": {"name": "newly_available"}}]
|
||||
self.valid_tool_names = {"newly_available"}
|
||||
|
||||
with patch.object(run_agent, "AIAgent", _RecorderWithLateTool):
|
||||
fork, _rt, routed = build_cache_parity_fork(agent, max_iterations=5)
|
||||
assert not routed
|
||||
assert fork.tools == []
|
||||
assert fork.tools is not agent.tools
|
||||
assert fork.valid_tool_names == set()
|
||||
|
||||
# Compaction refresh must refuse rebuild on frozen snapshot
|
||||
added = refresh_agent_mcp_tools(fork, content_aware=True)
|
||||
assert added == set()
|
||||
assert fork.tools == []
|
||||
assert fork.valid_tool_names == set()
|
||||
|
||||
|
||||
def test_same_model_review_surfaces_ignored_reasoning_effort_once():
|
||||
"""#104116: ``auxiliary.background_review.reasoning_effort`` is dropped on the same-model path
|
||||
(cache parity, #30532) — that no-op must be visible instead of silent, and must not fire per
|
||||
fork (a nudge-per-turn session would spam)."""
|
||||
import run_agent
|
||||
from agent.background_review import build_cache_parity_fork
|
||||
|
||||
agent = _make_agent_stub(run_agent.AIAgent)
|
||||
warnings = []
|
||||
agent._emit_warning = warnings.append
|
||||
captured = {}
|
||||
_Recorder = _make_recorder_class(captured)
|
||||
|
||||
with patch.object(run_agent, "AIAgent", _Recorder):
|
||||
_fork, _rt, routed = build_cache_parity_fork(
|
||||
agent, {"reasoning_effort": "low"}, max_iterations=5)
|
||||
assert not routed
|
||||
assert len(warnings) == 1, f"expected exactly one notice, got {warnings!r}"
|
||||
assert "auxiliary.background_review.reasoning_effort='low'" in warnings[0], warnings[0]
|
||||
# Cache-parity behaviour itself is unchanged: the fork still inherits the parent verbatim.
|
||||
assert captured["init_kwargs"]["reasoning_config"] == agent.reasoning_config
|
||||
# Second fork on the same parent: no repeat.
|
||||
build_cache_parity_fork(agent, {"reasoning_effort": "low"}, max_iterations=5)
|
||||
assert len(warnings) == 1, f"notice repeated per fork: {warnings!r}"
|
||||
|
||||
|
||||
def test_review_effort_notice_only_for_same_model_review_forks():
|
||||
"""No notice when the key is unset, when the fork is routed (#94825 owns that path), or for the
|
||||
/btw ``side_question`` fork sharing ``build_cache_parity_fork``."""
|
||||
import run_agent
|
||||
import agent.background_review as bg_review
|
||||
from agent.background_review import build_cache_parity_fork
|
||||
|
||||
_Recorder = _make_recorder_class()
|
||||
routed_runtime = {
|
||||
"provider": "openrouter", "model": "aux-cheap-model", "api_key": "test-key",
|
||||
"base_url": None, "api_mode": None, "credential_pool": None, "request_overrides": {},
|
||||
"max_tokens": None, "command": None, "args": [], "routed": True,
|
||||
}
|
||||
|
||||
def _warns(task_cfg, **kwargs):
|
||||
agent = _make_agent_stub(run_agent.AIAgent)
|
||||
warnings = []
|
||||
agent._emit_warning = warnings.append
|
||||
with patch.object(run_agent, "AIAgent", _Recorder):
|
||||
build_cache_parity_fork(agent, task_cfg, max_iterations=5, **kwargs)
|
||||
return warnings
|
||||
|
||||
assert _warns({"reasoning_effort": ""}) == []
|
||||
assert _warns({}) == []
|
||||
assert _warns({"reasoning_effort": "low"}, write_origin="side_question") == []
|
||||
with patch.object(bg_review, "_resolve_review_runtime", return_value=routed_runtime):
|
||||
assert _warns({"reasoning_effort": "low"}) == []
|
||||
|
||||
|
||||
def test_same_model_fork_inherits_parent_cache_scope_gateway_key(tmp_path):
|
||||
"""#109964 invariant 1 (gateway-key case): the same-model review fork must
|
||||
resolve the PARENT's cache scope, even though it is _persist_disabled and
|
||||
_session_db=None. Pre-fix, both resolvers diverged on their own — the header
|
||||
(affinity) and body (prompt_cache_key) keyed a different bucket, costing one
|
||||
cold ~full-context request per review."""
|
||||
import run_agent
|
||||
from agent.background_review import build_cache_parity_fork
|
||||
from agent.prompt_cache_scope import (
|
||||
declared_conversation_scope,
|
||||
resolve_prompt_cache_scope,
|
||||
)
|
||||
from hermes_state import SessionDB
|
||||
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
try:
|
||||
agent = _make_agent_stub(run_agent.AIAgent)
|
||||
# Gateway parent shape: declared key, real DB row behind it.
|
||||
agent._gateway_session_key = "gw-key-1"
|
||||
agent._session_db = db
|
||||
db.create_session("sess-123", source="test")
|
||||
|
||||
with patch.object(run_agent, "AIAgent", _make_recorder_class()):
|
||||
fork, _rt, routed = build_cache_parity_fork(agent, max_iterations=5)
|
||||
|
||||
assert not routed
|
||||
parent_scope = resolve_prompt_cache_scope(agent)
|
||||
assert parent_scope.startswith("gwk_"), parent_scope
|
||||
# The fork stamps the parent's resolved scope; both resolvers honor it.
|
||||
assert getattr(fork, "_inherited_cache_scope", None) == parent_scope
|
||||
assert declared_conversation_scope(fork) == parent_scope
|
||||
assert resolve_prompt_cache_scope(fork) == parent_scope
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_same_model_fork_inherits_parent_cache_scope_rotated_lineage(tmp_path):
|
||||
"""#109964 invariant 1 (rotated-lineage case): a CLI parent whose lineage
|
||||
root != current physical id must also pass its scope to the fork. Pre-fix
|
||||
the parent resolved 'root-sid' while the fork fell to the physical id.
|
||||
|
||||
Every identity the fork publishes must equal the parent's: body cache key, the
|
||||
affinity header (None for both — a physical root is not a declared ``gwk_`` scope)
|
||||
and the Portal ``conversation=`` root (fork has no DB to walk the lineage)."""
|
||||
import run_agent
|
||||
from agent.background_review import build_cache_parity_fork
|
||||
from agent.prompt_cache_scope import declared_conversation_scope, resolve_prompt_cache_scope
|
||||
from hermes_state import SessionDB
|
||||
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
try:
|
||||
agent = _make_agent_stub(run_agent.AIAgent)
|
||||
agent._session_db = db
|
||||
# Legacy compression rotation: parent row ends, child inherits its lineage.
|
||||
db.create_session("root-sid", source="test")
|
||||
db.end_session("root-sid", "compression")
|
||||
db.create_session("sess-123", source="test", parent_session_id="root-sid")
|
||||
agent.session_id = "sess-123"
|
||||
|
||||
with patch.object(run_agent, "AIAgent", _make_recorder_class()):
|
||||
fork, _rt, routed = build_cache_parity_fork(agent, max_iterations=5)
|
||||
|
||||
assert not routed
|
||||
assert resolve_prompt_cache_scope(agent) == "root-sid"
|
||||
assert getattr(fork, "_inherited_cache_scope", None) == "root-sid"
|
||||
assert resolve_prompt_cache_scope(fork) == "root-sid"
|
||||
assert declared_conversation_scope(fork) is declared_conversation_scope(agent) is None
|
||||
fork_root = run_agent.AIAgent._conversation_root_id(fork)
|
||||
assert fork_root == agent._conversation_root_id() == "root-sid"
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_routed_fork_does_not_inherit_cache_scope():
|
||||
"""#109964 invariant 2: a routed (different-model) fork is cache-cold on
|
||||
that model anyway — it must NOT inherit the parent's scope. Nor may fresh
|
||||
agents (no attribute set) be affected: the fail-closed default stands."""
|
||||
import run_agent
|
||||
from agent.background_review import build_cache_parity_fork
|
||||
|
||||
agent = _make_agent_stub(run_agent.AIAgent)
|
||||
agent._gateway_session_key = "gw-key-1"
|
||||
agent._prompt_cache_scope_memo = (("sess-123", True), "gwk_parentscope0000000000abc")
|
||||
|
||||
_RoutedRecorder = _make_recorder_class()
|
||||
|
||||
with patch.object(run_agent, "AIAgent", _RoutedRecorder), \
|
||||
patch("agent.background_review._resolve_review_runtime",
|
||||
lambda *a, **k: {"routed": True, "model": "other-model"}):
|
||||
fork, _rt, routed = build_cache_parity_fork(agent, max_iterations=5)
|
||||
|
||||
assert routed
|
||||
assert not getattr(fork, "_inherited_cache_scope", None), (
|
||||
"Routed fork must not inherit the parent's cache scope — its prefix "
|
||||
"is cache-cold on the different model regardless."
|
||||
)
|
||||
@@ -0,0 +1,175 @@
|
||||
"""Unit coverage for the background-review aux-model selector + routed digest.
|
||||
|
||||
Covers the two behaviors this change adds:
|
||||
• _resolve_review_runtime — auto/same-model → not routed (main model, warm
|
||||
cache); a configured different model → routed with resolved credentials.
|
||||
• _digest_history — compact replay used ONLY on the routed path (recent tail
|
||||
verbatim + a digest of older turns), preserving role alternation.
|
||||
|
||||
Pure-function / config-driven; no live model calls.
|
||||
"""
|
||||
from typing import Any
|
||||
from unittest.mock import patch
|
||||
|
||||
from agent import background_review as br
|
||||
|
||||
|
||||
def _msg(role, content, tool_calls=None):
|
||||
m = {"role": role, "content": content}
|
||||
if tool_calls:
|
||||
m["tool_calls"] = tool_calls
|
||||
return m
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _resolve_review_runtime — the aux-model selector
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class _FakeAgent:
|
||||
def __init__(self, provider="openai-codex", model="gpt-5.5"):
|
||||
self.provider = provider
|
||||
self.model = model
|
||||
self._credential_pool: Any = None
|
||||
self.request_overrides = {}
|
||||
self.max_tokens: int | None = None
|
||||
|
||||
def _current_main_runtime(self):
|
||||
return {
|
||||
"api_key": "parent-key",
|
||||
"base_url": "https://chatgpt.com/backend-api/codex",
|
||||
"api_mode": "codex_app_server",
|
||||
}
|
||||
|
||||
|
||||
def test_routing_auto_inherits_parent_and_downgrades_codex_app_server():
|
||||
agent = _FakeAgent()
|
||||
cfg = {"auxiliary": {"background_review": {"provider": "auto", "model": ""}}}
|
||||
with patch("hermes_cli.config.load_config", return_value=cfg), patch("hermes_cli.config.load_config_readonly", return_value=cfg):
|
||||
rt = br._resolve_review_runtime(agent)
|
||||
assert rt["routed"] is False
|
||||
assert rt["provider"] == "openai-codex"
|
||||
assert rt["model"] == "gpt-5.5"
|
||||
assert rt["api_mode"] == "codex_responses" # downgraded so agent-loop tools dispatch
|
||||
|
||||
|
||||
def test_routing_to_different_model_marks_routed_and_resolves_credentials():
|
||||
agent = _FakeAgent()
|
||||
cfg = {"auxiliary": {"background_review": {
|
||||
"provider": "openrouter", "model": "google/gemini-3-flash-preview",
|
||||
}}}
|
||||
fake_rp = {
|
||||
"provider": "openrouter", "api_key": "or-key",
|
||||
"base_url": "https://openrouter.ai/api/v1", "api_mode": "chat_completions",
|
||||
"credential_pool": "routed-pool",
|
||||
"request_overrides": {"extra_body": {"store": False}},
|
||||
"max_output_tokens": 2048,
|
||||
}
|
||||
with patch("hermes_cli.config.load_config", return_value=cfg), patch("hermes_cli.config.load_config_readonly", return_value=cfg), \
|
||||
patch("hermes_cli.runtime_provider.resolve_runtime_provider", return_value=fake_rp):
|
||||
rt = br._resolve_review_runtime(agent)
|
||||
assert rt["routed"] is True
|
||||
assert rt["provider"] == "openrouter"
|
||||
assert rt["model"] == "google/gemini-3-flash-preview"
|
||||
assert rt["api_key"] == "or-key"
|
||||
assert rt["credential_pool"] == "routed-pool"
|
||||
assert rt["request_overrides"] == {"extra_body": {"store": False}}
|
||||
assert rt.get("max_tokens") is None
|
||||
|
||||
|
||||
def test_unrouted_runtime_keeps_parent_pool_and_overrides():
|
||||
agent = _FakeAgent()
|
||||
agent._credential_pool = "parent-pool"
|
||||
agent.request_overrides = {"service_tier": "priority"}
|
||||
agent.max_tokens = 4096
|
||||
with patch("hermes_cli.config.load_config", return_value={}), patch("hermes_cli.config.load_config_readonly", return_value={}):
|
||||
rt = br._resolve_review_runtime(agent)
|
||||
assert rt["credential_pool"] == "parent-pool"
|
||||
assert rt["request_overrides"] == {"service_tier": "priority"}
|
||||
assert rt["max_tokens"] == 4096
|
||||
|
||||
|
||||
def test_routing_same_model_as_parent_is_not_routed():
|
||||
agent = _FakeAgent(provider="openrouter", model="anthropic/claude-opus-4.8")
|
||||
cfg = {"auxiliary": {"background_review": {
|
||||
"provider": "openrouter", "model": "anthropic/claude-opus-4.8",
|
||||
}}}
|
||||
with patch("hermes_cli.config.load_config", return_value=cfg), patch("hermes_cli.config.load_config_readonly", return_value=cfg):
|
||||
rt = br._resolve_review_runtime(agent)
|
||||
assert rt["routed"] is False # same model/provider → keep full-replay path
|
||||
|
||||
|
||||
def test_routing_resolution_failure_falls_back_to_parent():
|
||||
agent = _FakeAgent()
|
||||
cfg = {"auxiliary": {"background_review": {
|
||||
"provider": "openrouter", "model": "google/gemini-3-flash-preview",
|
||||
}}}
|
||||
with patch("hermes_cli.config.load_config", return_value=cfg), patch("hermes_cli.config.load_config_readonly", return_value=cfg), \
|
||||
patch("hermes_cli.runtime_provider.resolve_runtime_provider",
|
||||
side_effect=RuntimeError("boom")):
|
||||
rt = br._resolve_review_runtime(agent)
|
||||
assert rt["routed"] is False
|
||||
assert rt["provider"] == "openai-codex"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _digest_history — routed-path compact replay
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_digest_under_tail_returns_full():
|
||||
msgs = [_msg("user", "hi"), _msg("assistant", "hello")]
|
||||
assert br._digest_history(msgs, tail=24) == msgs
|
||||
|
||||
|
||||
def test_digest_collapses_old_keeps_tail_verbatim():
|
||||
msgs = []
|
||||
for i in range(60):
|
||||
msgs.append(_msg("user", f"u{i} " + "x" * 50))
|
||||
msgs.append(_msg("assistant", f"a{i} " + "y" * 50))
|
||||
out = br._digest_history(msgs, tail=10)
|
||||
# First message is the synthetic digest (user role → alternation preserved).
|
||||
assert out[0]["role"] == "user"
|
||||
assert out[0]["content"].startswith("[Earlier conversation digest")
|
||||
# Recent tail preserved verbatim.
|
||||
assert out[-1] == msgs[-1]
|
||||
assert len(out) == 11 # 1 digest + 10 tail
|
||||
|
||||
|
||||
def test_digest_does_not_open_tail_on_a_tool_message():
|
||||
msgs = []
|
||||
for i in range(40):
|
||||
msgs.append(_msg("user", "u" + "x" * 50))
|
||||
msgs.append(_msg("assistant", "", tool_calls=[
|
||||
{"function": {"name": "terminal", "arguments": "{}"}}]))
|
||||
msgs.append({"role": "tool", "content": "result " + "w" * 50})
|
||||
out = br._digest_history(msgs, tail=2)
|
||||
# The verbatim tail (after the digest) must not begin on a bare tool message.
|
||||
assert out[1]["role"] != "tool"
|
||||
|
||||
|
||||
def test_digest_records_tool_names_in_arc():
|
||||
old = [
|
||||
_msg("user", "do the thing"),
|
||||
_msg("assistant", "", tool_calls=[
|
||||
{"function": {"name": "skill_view", "arguments": "{}"}},
|
||||
{"function": {"name": "patch", "arguments": "{}"}}]),
|
||||
]
|
||||
msgs = old + [_msg("user", f"tail{i}") for i in range(30)]
|
||||
out = br._digest_history(msgs, tail=10)
|
||||
digest = out[0]["content"]
|
||||
assert "USER: do the thing" in digest
|
||||
assert "tools: skill_view, patch" in digest
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cost / configurability controls (issue #87250)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_enabled_defaults_true():
|
||||
with patch("hermes_cli.config.load_config_readonly", return_value={}):
|
||||
assert br.load_background_review_settings()[0] is True
|
||||
|
||||
|
||||
def test_enabled_false_disables_automatic_review():
|
||||
cfg = {"auxiliary": {"background_review": {"enabled": False}}}
|
||||
with patch("hermes_cli.config.load_config_readonly", return_value=cfg):
|
||||
assert br.load_background_review_settings()[0] is False
|
||||
@@ -0,0 +1,239 @@
|
||||
"""Regression tests for the background-review aggregate input budget (#93057).
|
||||
|
||||
The review fork replays its snapshot on every provider request in its tool
|
||||
loop. Detached in-memory compaction bounds any SINGLE request; the aggregate
|
||||
input budget (``_review_input_token_budget``, set by
|
||||
``_run_review_in_thread`` from ``auxiliary.background_review.max_input_tokens``)
|
||||
bounds the WHOLE review: the tool loop stops before the provider call that
|
||||
would cross it, mirroring the iteration-budget exit.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from run_agent import AIAgent
|
||||
|
||||
|
||||
def _tool_call() -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
id="call_1",
|
||||
type="function",
|
||||
function=SimpleNamespace(name="web_search", arguments='{"query": "x"}'),
|
||||
)
|
||||
|
||||
|
||||
def _tool_response(prompt_tokens: int) -> SimpleNamespace:
|
||||
message = SimpleNamespace(
|
||||
content=None,
|
||||
reasoning_content=None,
|
||||
reasoning=None,
|
||||
tool_calls=[_tool_call()],
|
||||
)
|
||||
return SimpleNamespace(
|
||||
choices=[SimpleNamespace(message=message, finish_reason="tool_calls")],
|
||||
model="test/model",
|
||||
usage=SimpleNamespace(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=1,
|
||||
total_tokens=prompt_tokens + 1,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _final_response() -> SimpleNamespace:
|
||||
message = SimpleNamespace(
|
||||
content="done",
|
||||
reasoning_content=None,
|
||||
reasoning=None,
|
||||
tool_calls=None,
|
||||
)
|
||||
return SimpleNamespace(
|
||||
choices=[SimpleNamespace(message=message, finish_reason="stop")],
|
||||
model="test/model",
|
||||
usage=None,
|
||||
)
|
||||
|
||||
|
||||
def _tool_definition() -> dict:
|
||||
return {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "web_search",
|
||||
"description": "Search the web",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"query": {"type": "string"}},
|
||||
"required": ["query"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _make_loop_agent():
|
||||
with (
|
||||
patch("model_tools.get_tool_definitions", return_value=[_tool_definition()]),
|
||||
patch("model_tools.check_toolset_requirements", return_value={}),
|
||||
patch("agent.process_bootstrap.OpenAI"),
|
||||
patch("agent.model_metadata.get_model_context_length", return_value=256_000),
|
||||
patch("agent.context_compressor.get_model_context_length", return_value=256_000),
|
||||
):
|
||||
agent = AIAgent(
|
||||
api_key="test-key-1234567890",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
max_iterations=10,
|
||||
)
|
||||
|
||||
agent.client = MagicMock()
|
||||
agent._cached_system_prompt = "You are helpful."
|
||||
agent._use_prompt_caching = False
|
||||
agent._disable_streaming = True
|
||||
agent.tool_delay = 0
|
||||
agent.save_trajectories = False
|
||||
agent.max_compression_attempts = 1
|
||||
|
||||
compressor = MagicMock()
|
||||
compressor.protect_first_n = 3
|
||||
compressor.protect_last_n = 20
|
||||
compressor.threshold_tokens = 999_999_999 # never fire compaction here
|
||||
compressor.context_length = 1_000_000_000
|
||||
compressor.last_prompt_tokens = -1
|
||||
compressor._verify_compaction_cleared_threshold = False
|
||||
compressor.awaiting_real_usage_after_compression = False
|
||||
compressor.should_compress.return_value = False
|
||||
compressor.should_compress_info.return_value = (False, None)
|
||||
compressor.should_compress_preflight.return_value = False
|
||||
compressor.should_defer_preflight_to_real_usage.return_value = False
|
||||
compressor.get_active_compression_failure_cooldown.return_value = None
|
||||
compressor.select_context.return_value = None
|
||||
compressor.get_automatic_compaction_status_message.return_value = ""
|
||||
agent.compression_enabled = False # isolate the budget behavior under test
|
||||
agent.context_compressor = compressor
|
||||
|
||||
def _fake_execute_tool_calls(assistant_message, messages, *_args):
|
||||
tool_call = assistant_message.tool_calls[0]
|
||||
messages.append(
|
||||
{
|
||||
"role": "tool",
|
||||
"name": tool_call.function.name,
|
||||
"tool_call_id": tool_call.id,
|
||||
"content": "ok",
|
||||
}
|
||||
)
|
||||
|
||||
agent._execute_tool_calls = _fake_execute_tool_calls
|
||||
return agent
|
||||
|
||||
|
||||
def _run_with_responses(agent, responses):
|
||||
agent.client.chat.completions.create.side_effect = responses
|
||||
with (
|
||||
patch.object(agent, "_flush_messages_to_session_db", return_value=True),
|
||||
patch.object(agent, "_persist_session"),
|
||||
patch.object(agent, "_save_trajectory"),
|
||||
patch.object(agent, "_cleanup_task_resources"),
|
||||
):
|
||||
result = agent.run_conversation("do some tool work")
|
||||
return result
|
||||
|
||||
|
||||
def test_review_input_budget_stops_tool_loop_before_next_provider_call():
|
||||
"""Once a fork's cumulative input crosses its budget, no further provider
|
||||
call is made — the crossing request completes, then the loop stops."""
|
||||
agent = _make_loop_agent()
|
||||
agent._review_input_token_budget = 100_000
|
||||
|
||||
responses = [
|
||||
_tool_response(50_000),
|
||||
_tool_response(50_000), # cumulative 100_000 -> budget crossed
|
||||
_tool_response(50_000), # must never be consumed
|
||||
_final_response(),
|
||||
]
|
||||
result = _run_with_responses(agent, responses)
|
||||
|
||||
create = agent.client.chat.completions.create
|
||||
assert create.call_count == 2, (
|
||||
f"expected the loop to stop after crossing the input budget, "
|
||||
f"but {create.call_count} provider calls were made (budget "
|
||||
f"{agent._review_input_token_budget}, "
|
||||
f"used {agent.session_input_tokens})"
|
||||
)
|
||||
assert agent.session_input_tokens == 100_000
|
||||
assert result["completed"] is False
|
||||
|
||||
|
||||
def test_no_budget_attribute_leaves_tool_loop_unbounded():
|
||||
"""Agents without ``_review_input_token_budget`` (every normal agent)
|
||||
are unaffected by the gate and consume all scripted responses."""
|
||||
agent = _make_loop_agent()
|
||||
|
||||
responses = [
|
||||
_tool_response(50_000),
|
||||
_tool_response(50_000),
|
||||
_tool_response(50_000),
|
||||
_final_response(),
|
||||
]
|
||||
result = _run_with_responses(agent, responses)
|
||||
|
||||
assert agent.client.chat.completions.create.call_count == 4
|
||||
assert result["completed"] is True
|
||||
assert result["final_response"] == "done"
|
||||
|
||||
|
||||
def test_review_input_budget_exhausted_predicate_edge_cases():
|
||||
"""The gate only arms for a positive int budget and a real token count."""
|
||||
from agent.conversation_loop import _review_input_budget_exhausted
|
||||
|
||||
class _Agent:
|
||||
pass
|
||||
|
||||
agent = _Agent()
|
||||
assert _review_input_budget_exhausted(agent) is False
|
||||
|
||||
agent._review_input_token_budget = None
|
||||
agent.session_input_tokens = 999_999
|
||||
assert _review_input_budget_exhausted(agent) is False
|
||||
|
||||
agent._review_input_token_budget = 0
|
||||
assert _review_input_budget_exhausted(agent) is False
|
||||
|
||||
agent._review_input_token_budget = -1
|
||||
assert _review_input_budget_exhausted(agent) is False
|
||||
|
||||
agent._review_input_token_budget = "100"
|
||||
assert _review_input_budget_exhausted(agent) is False
|
||||
|
||||
agent._review_input_token_budget = True
|
||||
assert _review_input_budget_exhausted(agent) is False
|
||||
|
||||
agent._review_input_token_budget = 100_000
|
||||
agent.session_input_tokens = 99_999
|
||||
assert _review_input_budget_exhausted(agent) is False
|
||||
|
||||
agent.session_input_tokens = 100_000
|
||||
assert _review_input_budget_exhausted(agent) is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("config_value", "expected"),
|
||||
[
|
||||
({}, 600_000),
|
||||
({"max_input_tokens": 1_000_000}, 1_000_000),
|
||||
({"max_input_tokens": 0}, None),
|
||||
({"max_input_tokens": -5}, None),
|
||||
({"max_input_tokens": "not-a-number"}, 600_000),
|
||||
({"max_input_tokens": "300000"}, 300_000),
|
||||
],
|
||||
)
|
||||
def test_review_input_token_budget_resolution(config_value, expected):
|
||||
"""Config parsing: default, override, explicit disable, garbage fallback."""
|
||||
from agent.background_review import _review_input_token_budget
|
||||
|
||||
assert _review_input_token_budget(config_value) == expected
|
||||
@@ -0,0 +1,309 @@
|
||||
"""Regression tests for the list-shape AttributeError guards in
|
||||
``agent.background_review.summarize_background_review_actions`` (#59437).
|
||||
|
||||
The outer ``_run_review_in_thread`` used to crash with
|
||||
``'list' object has no attribute 'get'`` every time a tool response
|
||||
returned a list (or any non-dict) where the summarizer expected a
|
||||
dict — most commonly the ``_change`` field in skill_manage responses
|
||||
or one of the entries in a memory operations list. The crash took
|
||||
down the entire background review, discarding every other successful
|
||||
action that the fork had completed.
|
||||
|
||||
What this module guards:
|
||||
|
||||
A. ``summarize_background_review_actions`` no longer raises when
|
||||
``data["_change"]`` is a list. It returns the rest of the
|
||||
actions normally.
|
||||
B. ``summarize_background_review_actions`` no longer raises when
|
||||
``operations`` is a non-list (string, int, None). It treats the
|
||||
field as empty.
|
||||
C. ``summarize_background_review_actions`` no longer raises when
|
||||
``operations[i]`` is a non-dict (string, None). It skips that
|
||||
entry but processes the rest.
|
||||
D. ``summarize_background_review_actions`` no longer raises when
|
||||
``call_details.get(tcid)`` returns a non-dict (e.g. None or a
|
||||
stray scalar). It coerces to ``{}``.
|
||||
E. The caller in ``_run_review_in_thread`` no longer aborts the
|
||||
whole review on an unrelated summarize exception; partial valid
|
||||
actions are surfaced.
|
||||
|
||||
The tests run without pytest (handoff from a prior pattern): they use
|
||||
plain ``assert`` and a small standalone runner. Importing the module
|
||||
exercises the new code paths without booting the LLM stack — there
|
||||
are no I/O or model dependencies in the unit-of-work being tested.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import importlib.util
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import types
|
||||
|
||||
|
||||
REPO_ROOT = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
|
||||
def _isolate_hermes_home():
|
||||
os.environ.setdefault("HERMES_HOME", "/tmp/hermes-bg-review-test")
|
||||
|
||||
|
||||
def _load_module():
|
||||
"""Lazy import so a missing optional dep doesn't block the suite.
|
||||
|
||||
Returns the module or None if import failed.
|
||||
"""
|
||||
if REPO_ROOT not in sys.path:
|
||||
sys.path.insert(0, REPO_ROOT)
|
||||
try:
|
||||
return importlib.import_module("agent.background_review")
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _make_skill_tool_message(change, operations=None):
|
||||
"""Build the messages list that triggered the original crash."""
|
||||
return [
|
||||
# Assistant: calls skill_manage
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "skill_manage",
|
||||
"arguments": json.dumps(
|
||||
{
|
||||
"action": "patch",
|
||||
"name": "my-skill",
|
||||
"operations": operations
|
||||
or [
|
||||
{
|
||||
"action": "replace",
|
||||
"content": "x",
|
||||
"old_text": "y",
|
||||
}
|
||||
],
|
||||
}
|
||||
),
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
# Tool: response with a buggy _change field (a list instead of dict)
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_1",
|
||||
"content": json.dumps(
|
||||
{
|
||||
"success": True,
|
||||
"message": "Skill 'my-skill' patched.",
|
||||
"_change": change, # ← the offender, normally a dict
|
||||
}
|
||||
),
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def _make_memory_tool_message(operations_field):
|
||||
"""Memory tool response with a non-canonical operations field."""
|
||||
return [
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_2",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "memory",
|
||||
"arguments": json.dumps({"action": "add", "target": "memory"}),
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_2",
|
||||
"content": json.dumps(
|
||||
{
|
||||
"success": True,
|
||||
"message": "Entry added.",
|
||||
"operations": operations_field,
|
||||
}
|
||||
),
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
class TestRunner:
|
||||
def __init__(self):
|
||||
self.passed = []
|
||||
self.failed = []
|
||||
|
||||
def run(self, name, fn):
|
||||
try:
|
||||
fn()
|
||||
except Exception as e: # noqa: BLE001 — runner summary uses it
|
||||
import traceback
|
||||
self.failed.append((name, e, traceback.format_exc()))
|
||||
else:
|
||||
self.passed.append(name)
|
||||
|
||||
def summary(self):
|
||||
total = len(self.passed) + len(self.failed)
|
||||
print(f"\n{'=' * 70}\nResults: {len(self.passed)}/{total} passed")
|
||||
if self.failed:
|
||||
print(f"\n--- {len(self.failed)} failure(s) ---")
|
||||
for n, _e, tb in self.failed:
|
||||
print(f"\n[FAIL] {n}\n{tb}")
|
||||
return 0 if not self.failed else 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# A. _change as a list (the originally-reported crash class)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# B. operations as a non-list (string / int / None)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
|
||||
def test_b_operations_as_none_treated_as_empty():
|
||||
"""``operations = None`` (missing key, JSON null) is still safe."""
|
||||
_isolate_hermes_home()
|
||||
bg = _load_module()
|
||||
if bg is None:
|
||||
print("SKIP module not importable")
|
||||
return
|
||||
|
||||
msgs = _make_memory_tool_message(operations_field=None)
|
||||
actions = bg.summarize_background_review_actions(
|
||||
review_messages=msgs,
|
||||
prior_snapshot=[],
|
||||
notification_mode="verbose",
|
||||
)
|
||||
assert isinstance(actions, list)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# C. operations[i] as a non-dict (str / None)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_c_operations_contains_non_dict_entries():
|
||||
"""A legacy/half-typed operations list with string entries short-circuits.
|
||||
|
||||
In ``verbose`` mode the function should produce the valid entries and
|
||||
silently skip the non-dict ones without ``AttributeError``. In
|
||||
non-verbose mode it falls back to a generic "Memory updated" string,
|
||||
so this test exercises the verbose branch where iteration over
|
||||
per-entry fields actually happens.
|
||||
"""
|
||||
_isolate_hermes_home()
|
||||
bg = _load_module()
|
||||
if bg is None:
|
||||
print("SKIP module not importable")
|
||||
return
|
||||
|
||||
msgs = _make_memory_tool_message(
|
||||
operations_field=[
|
||||
"raw-string-no-fields",
|
||||
{"action": "add", "content": "valid entry"},
|
||||
None,
|
||||
{"action": "replace", "content": "another", "old_text": "thing"},
|
||||
]
|
||||
)
|
||||
actions = bg.summarize_background_review_actions(
|
||||
review_messages=msgs,
|
||||
prior_snapshot=[],
|
||||
notification_mode="verbose",
|
||||
)
|
||||
assert isinstance(actions, list)
|
||||
# ``notification_mode='verbose'`` walks per-entry fields; the two
|
||||
# dict-shaped entries produce action lines, the string and None
|
||||
# entries are skipped via the isinstance guard. The exact wording is
|
||||
# not asserted (memory module shapes may vary) but at least one
|
||||
# action line must be present.
|
||||
assert len(actions) >= 1, f"expected at least one action line, got {actions!r}"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# D. detail comes back non-dict (None / stale value)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_d_detail_non_dict_replaced_with_empty():
|
||||
"""When ``call_details.get(tcid)`` returns None, summarize must coerce
|
||||
it to ``{}`` rather than calling ``.get(...)`` on ``None``.
|
||||
"""
|
||||
_isolate_hermes_home()
|
||||
bg = _load_module()
|
||||
if bg is None:
|
||||
print("SKIP module not importable")
|
||||
return
|
||||
|
||||
# Build a tool-only message whose tcid does NOT have an assistant tool_call.
|
||||
msgs = _make_skill_tool_message(change={})
|
||||
# Drop the assistant message so call_details is empty for tcid=call_1.
|
||||
msgs = [m for m in msgs if m.get("role") != "assistant"]
|
||||
|
||||
actions = bg.summarize_background_review_actions(
|
||||
review_messages=msgs,
|
||||
prior_snapshot=[],
|
||||
notification_mode="verbose",
|
||||
)
|
||||
assert isinstance(actions, list)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# E. Caller defends against summarize raising
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_e_call_does_not_unwind_module_callables():
|
||||
"""Structural: the new defensive try/except around the summarize
|
||||
call is in place. Caught here rather than via a partial mocking
|
||||
cascade because monkeypatching the AIAgent is too brittle for a
|
||||
blind regression test — keeping it text-anchored guards the
|
||||
``_run_review_in_thread`` invariant without a real LLM.
|
||||
"""
|
||||
src_path = os.path.join(REPO_ROOT, "agent", "background_review.py")
|
||||
src = open(src_path, encoding="utf-8").read()
|
||||
# The fix added: ``try: actions = summarize_background_review_actions(...)``
|
||||
# followed by ``except Exception as e: ... actions = []``.
|
||||
assert "actions = summarize_background_review_actions(" in src
|
||||
assert (
|
||||
"summarize_background_review_actions returned partial results"
|
||||
in src
|
||||
), "expected partial-results guard message present"
|
||||
# And the non-dict guards on free-form tool payload fields.
|
||||
assert "if isinstance(ops_raw, list)" in src
|
||||
assert "if isinstance(change_raw, dict)" in src
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Runner
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def main():
|
||||
runner = TestRunner()
|
||||
runner.run("b_operations_as_none_treated_as_empty", test_b_operations_as_none_treated_as_empty)
|
||||
runner.run("c_operations_contains_non_dict_entries", test_c_operations_contains_non_dict_entries)
|
||||
runner.run("d_detail_non_dict_replaced_with_empty", test_d_detail_non_dict_replaced_with_empty)
|
||||
runner.run("e_call_defends_via_try_except", test_e_call_does_not_unwind_module_callables)
|
||||
return runner.summary()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,50 @@
|
||||
"""Routed background reviews honor ``auxiliary.background_review.reasoning_effort`` (#94825).
|
||||
|
||||
The review fork is a full AIAgent, not an auxiliary_client call. The routed branch deliberately
|
||||
skips the PARENT's reasoning_config (its effort vocabulary may be invalid for the routed model),
|
||||
but an explicitly configured per-task effort must win over provider defaults — mirroring how every
|
||||
other auxiliary task folds the same key into ``extra_body.reasoning``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from unittest.mock import patch
|
||||
|
||||
import run_agent
|
||||
import agent.background_review as bg_review
|
||||
from agent.background_review import build_cache_parity_fork
|
||||
|
||||
from tests.agent.test_background_review_cache_parity import _make_agent_stub, _make_recorder_class
|
||||
|
||||
ROUTED_RUNTIME = {
|
||||
"provider": "openrouter", "model": "aux-cheap-model", "api_key": "test-key",
|
||||
"base_url": None, "api_mode": None, "credential_pool": None, "request_overrides": {},
|
||||
"max_tokens": None, "command": None, "args": [], "routed": True,
|
||||
}
|
||||
|
||||
|
||||
def _routed_fork_kwargs(task_cfg):
|
||||
captured = {}
|
||||
agent = _make_agent_stub(run_agent.AIAgent)
|
||||
agent.reasoning_config = {"enabled": True, "effort": "high"}
|
||||
with patch.object(run_agent, "AIAgent", _make_recorder_class(captured)), \
|
||||
patch.object(bg_review, "_resolve_review_runtime", return_value=ROUTED_RUNTIME):
|
||||
_fork, _rt, routed = build_cache_parity_fork(agent, task_cfg, max_iterations=5)
|
||||
assert routed
|
||||
return captured["init_kwargs"]
|
||||
|
||||
|
||||
def test_routed_review_applies_configured_effort_not_parents():
|
||||
kwargs = _routed_fork_kwargs({"reasoning_effort": "xhigh"})
|
||||
assert kwargs["reasoning_config"] == {"enabled": True, "effort": "xhigh"}
|
||||
# ``none`` disables thinking on the routed fork, same vocabulary as every other aux task.
|
||||
assert _routed_fork_kwargs({"reasoning_effort": "none"})["reasoning_config"] == {"enabled": False}
|
||||
|
||||
|
||||
def test_routed_review_falls_back_to_provider_default(caplog):
|
||||
# Unset: provider default, and the parent's ``high`` is NOT smuggled across the route.
|
||||
assert "reasoning_config" not in _routed_fork_kwargs({"reasoning_effort": ""})
|
||||
assert "reasoning_config" not in _routed_fork_kwargs({})
|
||||
with caplog.at_level(logging.WARNING):
|
||||
assert "reasoning_config" not in _routed_fork_kwargs({"reasoning_effort": "ludicrous"})
|
||||
assert "ludicrous" in caplog.text
|
||||
@@ -0,0 +1,112 @@
|
||||
"""Tests for AIAgent._summarize_background_review_actions.
|
||||
|
||||
Regression coverage for issue #14944: the background memory/skill review used
|
||||
to re-surface tool results that were already present in the conversation
|
||||
history before the review started (e.g. an earlier "Cron job '...' created.").
|
||||
"""
|
||||
|
||||
import json
|
||||
|
||||
from run_agent import AIAgent
|
||||
|
||||
|
||||
_summarize = AIAgent._summarize_background_review_actions
|
||||
|
||||
|
||||
def _tool_msg(tool_call_id, payload):
|
||||
return {
|
||||
"role": "tool",
|
||||
"tool_call_id": tool_call_id,
|
||||
"content": json.dumps(payload),
|
||||
}
|
||||
|
||||
|
||||
def test_skips_prior_tool_messages_by_tool_call_id():
|
||||
"""Stale 'created' tool result from prior history must not be re-surfaced."""
|
||||
prior_payload = {"success": True, "message": "Cron job 'remind-me' created."}
|
||||
new_payload = {
|
||||
"success": True,
|
||||
"message": "Entry added",
|
||||
"target": "user",
|
||||
}
|
||||
|
||||
snapshot = [
|
||||
{"role": "user", "content": "create a reminder"},
|
||||
_tool_msg("call_old", prior_payload),
|
||||
{"role": "assistant", "content": "done"},
|
||||
]
|
||||
review_messages = list(snapshot) + [
|
||||
{"role": "user", "content": "<review prompt>"},
|
||||
_tool_msg("call_new", new_payload),
|
||||
]
|
||||
|
||||
actions = _summarize(review_messages, snapshot)
|
||||
|
||||
assert "Cron job 'remind-me' created." not in actions
|
||||
assert "User profile updated" in actions
|
||||
|
||||
|
||||
def test_includes_genuinely_new_actions():
|
||||
new_payload = {
|
||||
"success": True,
|
||||
"message": "Memory entry created.",
|
||||
}
|
||||
review_messages = [_tool_msg("call_new", new_payload)]
|
||||
|
||||
actions = _summarize(review_messages, prior_snapshot=[])
|
||||
|
||||
assert actions == ["Memory entry created."]
|
||||
|
||||
|
||||
def test_falls_back_to_content_equality_when_tool_call_id_missing():
|
||||
"""If a tool message has no tool_call_id, match prior entries by content."""
|
||||
payload = {"success": True, "message": "Cron job 'X' created."}
|
||||
raw = json.dumps(payload)
|
||||
prior_msg = {"role": "tool", "content": raw} # no tool_call_id
|
||||
review_messages = [
|
||||
{"role": "tool", "content": raw}, # same content -> stale, skip
|
||||
_tool_msg("call_new", {"success": True, "message": "Skill created."}),
|
||||
]
|
||||
|
||||
actions = _summarize(review_messages, [prior_msg])
|
||||
|
||||
assert "Cron job 'X' created." not in actions
|
||||
assert "Skill created." in actions
|
||||
|
||||
|
||||
|
||||
|
||||
def test_handles_non_json_tool_content_gracefully():
|
||||
review_messages = [
|
||||
{"role": "tool", "tool_call_id": "x", "content": "not-json"},
|
||||
_tool_msg("call_y", {"success": True, "message": "Memory updated."}),
|
||||
]
|
||||
|
||||
actions = _summarize(review_messages, [])
|
||||
|
||||
assert actions == ["Memory updated."]
|
||||
|
||||
|
||||
def test_empty_inputs():
|
||||
assert _summarize([], []) == []
|
||||
assert _summarize(None, None) == []
|
||||
|
||||
|
||||
|
||||
|
||||
def test_removed_or_replaced_relabels_by_target():
|
||||
review_messages = [
|
||||
_tool_msg(
|
||||
"c1",
|
||||
{"success": True, "message": "Entry removed.", "target": "user"},
|
||||
),
|
||||
_tool_msg(
|
||||
"c2",
|
||||
{"success": True, "message": "Entry replaced.", "target": "memory"},
|
||||
),
|
||||
]
|
||||
|
||||
actions = _summarize(review_messages, [])
|
||||
|
||||
assert "User profile updated" in actions
|
||||
assert "Memory updated" in actions
|
||||
@@ -0,0 +1,267 @@
|
||||
"""Tests that the background review agent restricts tools at runtime, not at schema time.
|
||||
|
||||
Regression coverage for issue #15204 (the background skill-review agent must
|
||||
not perform non-skill side effects like terminal, send_message, delegate_task)
|
||||
combined with issue #25322 / PR #17276 (the review fork must hit the parent's
|
||||
Anthropic/OpenRouter prefix cache).
|
||||
|
||||
Reconciling the two: the fork now inherits the parent's full ``tools`` schema
|
||||
so the cache-key matches, and enforces the memory+skills restriction at
|
||||
runtime via a thread-local whitelist on the existing
|
||||
``get_pre_tool_call_block_message`` gate. Safety is preserved mechanically
|
||||
(any non-whitelisted dispatch is blocked) without the schema-level narrowing
|
||||
that caused the prefix-cache miss.
|
||||
"""
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
|
||||
def _make_agent_stub(agent_cls):
|
||||
"""Create a minimal AIAgent-like object with just enough state for _spawn_background_review."""
|
||||
agent = object.__new__(agent_cls)
|
||||
agent.model = "test-model"
|
||||
agent.platform = "test"
|
||||
agent.provider = "openai"
|
||||
agent.session_id = "sess-123"
|
||||
agent.quiet_mode = True
|
||||
agent._memory_store = None
|
||||
agent._memory_enabled = True
|
||||
agent._user_profile_enabled = False
|
||||
agent._memory_nudge_interval = 5
|
||||
agent._skill_nudge_interval = 5
|
||||
agent.background_review_callback = None
|
||||
agent.status_callback = None
|
||||
agent._cached_system_prompt = None
|
||||
import datetime as _dt
|
||||
agent.session_start = _dt.datetime(2026, 1, 1, 12, 0, 0)
|
||||
agent._MEMORY_REVIEW_PROMPT = "review memory"
|
||||
agent._SKILL_REVIEW_PROMPT = "review skills"
|
||||
agent._COMBINED_REVIEW_PROMPT = "review both"
|
||||
# Non-None so the test catches a missing-kwarg regression.
|
||||
agent.enabled_toolsets = ["memory", "skills", "terminal"]
|
||||
agent.disabled_toolsets = ["spotify", "feishu_doc"]
|
||||
return agent
|
||||
|
||||
|
||||
class _SyncThread:
|
||||
"""Drop-in replacement for threading.Thread that runs the target inline."""
|
||||
|
||||
def __init__(self, *, target=None, daemon=None, name=None):
|
||||
self._target = target
|
||||
|
||||
def start(self):
|
||||
if self._target:
|
||||
self._target()
|
||||
|
||||
|
||||
def test_background_review_matches_parent_toolset_config():
|
||||
"""Fork must receive parent's toolset config so ``tools[]`` cache key matches."""
|
||||
import run_agent
|
||||
|
||||
agent = _make_agent_stub(run_agent.AIAgent)
|
||||
captured = {}
|
||||
|
||||
def _capture_init(self, *args, **kwargs):
|
||||
captured["enabled_toolsets"] = kwargs.get("enabled_toolsets", "UNSET")
|
||||
captured["disabled_toolsets"] = kwargs.get("disabled_toolsets", "UNSET")
|
||||
raise RuntimeError("stop after capturing init args")
|
||||
|
||||
with patch.object(run_agent.AIAgent, "__init__", _capture_init), \
|
||||
patch("threading.Thread", _SyncThread):
|
||||
agent._spawn_background_review(
|
||||
messages_snapshot=[],
|
||||
review_memory=True,
|
||||
review_skills=False,
|
||||
)
|
||||
|
||||
assert "enabled_toolsets" in captured, "AIAgent.__init__ was not called"
|
||||
assert captured["enabled_toolsets"] == agent.enabled_toolsets, (
|
||||
f"enabled_toolsets mismatch: {captured['enabled_toolsets']!r} "
|
||||
f"vs expected {agent.enabled_toolsets!r}"
|
||||
)
|
||||
assert captured["disabled_toolsets"] == agent.disabled_toolsets, (
|
||||
f"disabled_toolsets mismatch: {captured['disabled_toolsets']!r} "
|
||||
f"vs expected {agent.disabled_toolsets!r}"
|
||||
)
|
||||
|
||||
|
||||
def test_background_review_installs_thread_local_whitelist():
|
||||
"""The review fork must install a memory/skills-only thread-local whitelist.
|
||||
|
||||
The schema-level toolset narrowing was lifted (for prefix-cache parity),
|
||||
so #15204's safety contract now relies on the runtime whitelist gate to
|
||||
deny terminal/send_message/delegate_task at dispatch time. Verify the
|
||||
whitelist is set with exactly the memory+skills tool names.
|
||||
"""
|
||||
import run_agent
|
||||
from hermes_cli import plugins as _plugins
|
||||
|
||||
captured = {}
|
||||
|
||||
def _capture_whitelist(whitelist, deny_msg_fmt=None):
|
||||
captured["whitelist"] = set(whitelist)
|
||||
captured["deny_msg_fmt"] = deny_msg_fmt
|
||||
# Stop here — we just want to see what gets installed.
|
||||
raise RuntimeError("stop after capturing whitelist")
|
||||
|
||||
agent = _make_agent_stub(run_agent.AIAgent)
|
||||
|
||||
def _no_init(self, *args, **kwargs):
|
||||
# Don't crash AIAgent.__init__; let execution flow reach
|
||||
# set_thread_tool_whitelist.
|
||||
return None
|
||||
|
||||
with patch.object(run_agent.AIAgent, "__init__", _no_init), \
|
||||
patch.object(_plugins, "set_thread_tool_whitelist", _capture_whitelist), \
|
||||
patch("threading.Thread", _SyncThread):
|
||||
agent._spawn_background_review(
|
||||
messages_snapshot=[],
|
||||
review_memory=True,
|
||||
review_skills=False,
|
||||
)
|
||||
|
||||
assert "whitelist" in captured, "set_thread_tool_whitelist was not called"
|
||||
whitelist = captured["whitelist"]
|
||||
# memory + skills tools must be allowed
|
||||
assert "memory" in whitelist
|
||||
assert "skill_manage" in whitelist
|
||||
assert "skill_view" in whitelist
|
||||
assert "skills_list" in whitelist
|
||||
# read-only file tools are allowed too (#61521): the model reaches for
|
||||
# read_file to inspect a skill before patching; denying it caused a
|
||||
# per-review denial storm that starved the self-improvement loop.
|
||||
assert "read_file" in whitelist
|
||||
assert "search_files" in whitelist
|
||||
# write/dangerous tools must NOT be in the whitelist
|
||||
assert "write_file" not in whitelist
|
||||
assert "patch" not in whitelist
|
||||
assert "terminal" not in whitelist
|
||||
assert "send_message" not in whitelist
|
||||
assert "delegate_task" not in whitelist
|
||||
assert "web_search" not in whitelist
|
||||
assert "execute_code" not in whitelist
|
||||
# The deny message must name the correct substitutes so a single denial
|
||||
# redirects the model instead of a 142-denial storm (#61521).
|
||||
deny = captured.get("deny_msg_fmt") or ""
|
||||
assert "skill_manage" in deny
|
||||
assert "skill_view" in deny
|
||||
|
||||
|
||||
def test_read_file_registers_background_review_read_mark(tmp_path):
|
||||
"""read_file inside a review fork must satisfy the read-before-write guard.
|
||||
|
||||
The whitelist now allows read_file; without this mark, the model would
|
||||
read SKILL.md via read_file and still get "content has not been loaded
|
||||
in this review turn" on the follow-up skill_manage patch (#61521).
|
||||
"""
|
||||
from tools.file_tools import read_file_tool
|
||||
from tools.skill_manager_guards import (
|
||||
_background_review_has_read,
|
||||
_reset_background_review_read_marks,
|
||||
)
|
||||
from tools.skill_provenance import (
|
||||
BACKGROUND_REVIEW,
|
||||
reset_current_write_origin,
|
||||
set_current_write_origin,
|
||||
)
|
||||
|
||||
target = tmp_path / "SKILL.md"
|
||||
target.write_text("---\nname: t\n---\nbody\n")
|
||||
|
||||
token = set_current_write_origin(BACKGROUND_REVIEW)
|
||||
try:
|
||||
_reset_background_review_read_marks()
|
||||
assert not _background_review_has_read(target)
|
||||
out = read_file_tool(str(target), task_id="bg-review-test")
|
||||
assert "body" in out
|
||||
assert _background_review_has_read(target), (
|
||||
"full read_file inside a review fork must register with the "
|
||||
"read-before-write guard"
|
||||
)
|
||||
|
||||
# A partial read must NOT satisfy the guard.
|
||||
_reset_background_review_read_marks()
|
||||
read_file_tool(str(target), offset=2, task_id="bg-review-test2")
|
||||
assert not _background_review_has_read(target)
|
||||
finally:
|
||||
reset_current_write_origin(token)
|
||||
|
||||
|
||||
def test_read_file_outside_review_does_not_mark(tmp_path):
|
||||
"""Foreground reads must not populate the review-fork read set."""
|
||||
from tools.file_tools import read_file_tool
|
||||
from tools.skill_manager_guards import (
|
||||
_background_review_has_read,
|
||||
_reset_background_review_read_marks,
|
||||
)
|
||||
|
||||
target = tmp_path / "SKILL.md"
|
||||
target.write_text("content\n")
|
||||
_reset_background_review_read_marks()
|
||||
read_file_tool(str(target), task_id="fg-test")
|
||||
assert not _background_review_has_read(target)
|
||||
|
||||
|
||||
def test_background_review_whitelist_includes_configured_extra_tools(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
"""A profile may opt a specific proposal tool into background review.
|
||||
|
||||
The review fork inherits the parent's full tool schema for cache parity,
|
||||
but runtime dispatch remains denied unless the tool is also present in the
|
||||
thread-local whitelist. This config hook lets profiles grant a narrowly
|
||||
scoped, human-gated proposal tool without enabling unrelated side effects.
|
||||
"""
|
||||
hermes_home = tmp_path / ".hermes"
|
||||
hermes_home.mkdir()
|
||||
(hermes_home / "config.yaml").write_text(
|
||||
"auxiliary:\n"
|
||||
" background_review:\n"
|
||||
" extra_tools:\n"
|
||||
" - propose_shared_memory\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
monkeypatch.setenv("HERMES_HOME", str(hermes_home))
|
||||
|
||||
import run_agent
|
||||
from hermes_cli import config as config_module
|
||||
from hermes_cli import plugins as _plugins
|
||||
|
||||
config_module._LOAD_CONFIG_CACHE.clear()
|
||||
config_module._RAW_CONFIG_CACHE.clear()
|
||||
|
||||
captured = {}
|
||||
|
||||
def _capture_whitelist(whitelist, deny_msg_fmt=None):
|
||||
captured["whitelist"] = set(whitelist)
|
||||
|
||||
def _capture_run_conversation(self, *, user_message, **kwargs):
|
||||
captured["review_prompt"] = user_message
|
||||
return {"final_response": "Nothing to save."}
|
||||
|
||||
agent = _make_agent_stub(run_agent.AIAgent)
|
||||
|
||||
def _no_init(self, *args, **kwargs):
|
||||
return None
|
||||
|
||||
with patch.object(run_agent.AIAgent, "__init__", _no_init), \
|
||||
patch.object(
|
||||
run_agent.AIAgent,
|
||||
"run_conversation",
|
||||
_capture_run_conversation,
|
||||
), \
|
||||
patch.object(run_agent.AIAgent, "shutdown_memory_provider", lambda self: None), \
|
||||
patch.object(run_agent.AIAgent, "close", lambda self: None), \
|
||||
patch.object(_plugins, "set_thread_tool_whitelist", _capture_whitelist), \
|
||||
patch("threading.Thread", _SyncThread):
|
||||
agent._spawn_background_review(
|
||||
messages_snapshot=[],
|
||||
review_memory=True,
|
||||
review_skills=False,
|
||||
)
|
||||
|
||||
assert "propose_shared_memory" in captured["whitelist"]
|
||||
assert "terminal" not in captured["whitelist"]
|
||||
assert "propose_shared_memory" in captured["review_prompt"]
|
||||
|
||||
|
||||
@@ -556,6 +556,24 @@ class TestBuildConverseKwargs:
|
||||
)
|
||||
assert "inferenceConfig" not in kwargs
|
||||
|
||||
def test_bedrock_xai_grok_models_never_receive_sampling_params(self):
|
||||
"""Bedrock-hosted xAI Grok rejects temperature/topP in Converse with a hard 400
|
||||
(ValidationException: "This model doesn't support the temperature field"); the
|
||||
_forbids_sampling_params guard is Claude-only, so Grok needs its own denylist.
|
||||
Sibling Bedrock models keep receiving sampling params."""
|
||||
from agent.bedrock_adapter import build_converse_kwargs
|
||||
msgs = [{"role": "user", "content": "Hi"}]
|
||||
for model in ("us.xai.grok-4.6", "global.xai.grok-4.6"):
|
||||
cfg = build_converse_kwargs(
|
||||
model=model, messages=msgs, temperature=0.3, top_p=0.9
|
||||
)["inferenceConfig"]
|
||||
assert "temperature" not in cfg and "topP" not in cfg, model
|
||||
for model in ("test-model", "qwen.qwen3-vl-235b-a22b"):
|
||||
cfg = build_converse_kwargs(
|
||||
model=model, messages=msgs, temperature=0.3, top_p=0.9
|
||||
)["inferenceConfig"]
|
||||
assert cfg["temperature"] == 0.3 and cfg["topP"] == 0.9, model
|
||||
|
||||
def test_cache_point_added_for_supported_model(self):
|
||||
"""Claude and Nova on the Converse path get cachePoint markers on
|
||||
system, tools, and the message before the newest turn."""
|
||||
|
||||
@@ -0,0 +1,523 @@
|
||||
"""Hermetic tests for the Bitwarden Secrets Manager integration.
|
||||
|
||||
We never hit GitHub or Bitwarden in tests — subprocess + urllib are
|
||||
mocked so the suite stays fast and offline-safe. The "live" pull and
|
||||
binary download are exercised manually by `hermes secrets bitwarden
|
||||
setup` outside of pytest.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import stat
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
import zipfile
|
||||
from pathlib import Path
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
# Make the worktree importable without depending on the installed wheel.
|
||||
ROOT = Path(__file__).resolve().parents[2]
|
||||
if str(ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(ROOT))
|
||||
|
||||
from agent.secret_sources import bitwarden as bw # noqa: E402
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_caches():
|
||||
bw._reset_cache_for_tests()
|
||||
yield
|
||||
bw._reset_cache_for_tests()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def hermes_home(tmp_path, monkeypatch):
|
||||
"""Point Hermes at an isolated home directory."""
|
||||
home = tmp_path / ".hermes"
|
||||
home.mkdir()
|
||||
monkeypatch.setenv("HERMES_HOME", str(home))
|
||||
# Some modules cache get_hermes_home; clear if needed.
|
||||
import hermes_constants
|
||||
if hasattr(hermes_constants, "_HERMES_HOME_CACHE"):
|
||||
hermes_constants._HERMES_HOME_CACHE = None # type: ignore[attr-defined]
|
||||
return home
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _platform_asset_name
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"system,machine,libc_text,expected",
|
||||
[
|
||||
("Darwin", "x86_64", "",
|
||||
f"bws-macos-universal-{bw._BWS_VERSION}.zip"),
|
||||
("Darwin", "arm64", "",
|
||||
f"bws-macos-universal-{bw._BWS_VERSION}.zip"),
|
||||
("Linux", "x86_64", "glibc",
|
||||
f"bws-x86_64-unknown-linux-gnu-{bw._BWS_VERSION}.zip"),
|
||||
("Linux", "x86_64", "musl libc",
|
||||
f"bws-x86_64-unknown-linux-musl-{bw._BWS_VERSION}.zip"),
|
||||
("Linux", "aarch64", "",
|
||||
f"bws-aarch64-unknown-linux-gnu-{bw._BWS_VERSION}.zip"),
|
||||
("Windows", "AMD64", "",
|
||||
f"bws-x86_64-pc-windows-msvc-{bw._BWS_VERSION}.zip"),
|
||||
("Windows", "ARM64", "",
|
||||
f"bws-aarch64-pc-windows-msvc-{bw._BWS_VERSION}.zip"),
|
||||
],
|
||||
)
|
||||
def test_platform_asset_name(system, machine, libc_text, expected):
|
||||
with mock.patch.object(bw.platform, "system", return_value=system), \
|
||||
mock.patch.object(bw.platform, "machine", return_value=machine), \
|
||||
mock.patch.object(
|
||||
bw.subprocess,
|
||||
"run",
|
||||
return_value=mock.Mock(stdout=libc_text, stderr=libc_text),
|
||||
):
|
||||
assert bw._platform_asset_name() == expected
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# install_bws — fully mocked HTTP
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_fake_zip(binary_bytes: bytes) -> bytes:
|
||||
buf = io.BytesIO()
|
||||
with zipfile.ZipFile(buf, "w") as zf:
|
||||
zf.writestr("bws", binary_bytes)
|
||||
return buf.getvalue()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _safe_extract_member — zip-slip containment
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"evil_name",
|
||||
[
|
||||
"../escape",
|
||||
"../../escape",
|
||||
"sub/../../escape",
|
||||
],
|
||||
)
|
||||
def test_safe_extract_member_rejects_traversal(tmp_path, evil_name):
|
||||
buf = io.BytesIO()
|
||||
with zipfile.ZipFile(buf, "w") as zf:
|
||||
zf.writestr(evil_name, b"pwned")
|
||||
buf.seek(0)
|
||||
|
||||
dest = tmp_path / "extract"
|
||||
dest.mkdir()
|
||||
outside = tmp_path / "escape"
|
||||
|
||||
with zipfile.ZipFile(buf) as zf:
|
||||
with pytest.raises(RuntimeError, match="unsafe archive member"):
|
||||
bw._safe_extract_member(zf, evil_name, dest)
|
||||
|
||||
# The traversal target must not have been written.
|
||||
assert not outside.exists()
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_install_bws_happy_path(hermes_home, monkeypatch):
|
||||
fake_binary = b"#!/bin/sh\necho 'bws fake 2.0.0'\n"
|
||||
zip_bytes = _make_fake_zip(fake_binary)
|
||||
asset_name = bw._platform_asset_name()
|
||||
checksum_text = (
|
||||
f"{hashlib.sha256(zip_bytes).hexdigest()} {asset_name}\n"
|
||||
"ffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffff other-file\n"
|
||||
)
|
||||
|
||||
def fake_download(url, dest):
|
||||
if url.endswith(".zip"):
|
||||
Path(dest).write_bytes(zip_bytes)
|
||||
elif url.endswith(".txt"):
|
||||
Path(dest).write_text(checksum_text)
|
||||
else:
|
||||
raise AssertionError(f"unexpected download url: {url}")
|
||||
|
||||
monkeypatch.setattr(bw, "_http_download", fake_download)
|
||||
|
||||
path = bw.install_bws()
|
||||
assert path.exists()
|
||||
assert path.read_bytes() == fake_binary
|
||||
# Executable bit set
|
||||
assert path.stat().st_mode & stat.S_IXUSR
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# fetch_bitwarden_secrets
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _fake_bws_payload(items):
|
||||
return json.dumps(items)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_fetch_server_url_sets_env(monkeypatch, tmp_path):
|
||||
"""server_url must be plumbed into the subprocess as BWS_SERVER_URL."""
|
||||
fake_binary = tmp_path / "bws"
|
||||
fake_binary.write_text("")
|
||||
payload = _fake_bws_payload([{"key": "K", "value": "v"}])
|
||||
|
||||
captured_env = {}
|
||||
|
||||
def fake_run(cmd, **kwargs):
|
||||
captured_env.update(kwargs["env"])
|
||||
return mock.Mock(returncode=0, stdout=payload, stderr="")
|
||||
|
||||
monkeypatch.setattr(bw.subprocess, "run", fake_run)
|
||||
|
||||
bw.fetch_bitwarden_secrets(
|
||||
access_token="0.t",
|
||||
project_id="p",
|
||||
binary=fake_binary,
|
||||
use_cache=False,
|
||||
server_url="https://vault.bitwarden.eu",
|
||||
)
|
||||
assert captured_env.get("BWS_SERVER_URL") == "https://vault.bitwarden.eu"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# apply_bitwarden_secrets — the public entry point used by env_loader
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# env_loader integration
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
|
||||
def test_env_loader_calls_bsm_when_enabled(tmp_path, monkeypatch):
|
||||
home = tmp_path / ".hermes"
|
||||
home.mkdir()
|
||||
(home / "config.yaml").write_text(
|
||||
"secrets:\n"
|
||||
" bitwarden:\n"
|
||||
" enabled: true\n"
|
||||
" project_id: 'proj-1'\n"
|
||||
" access_token_env: 'BWS_ACCESS_TOKEN'\n"
|
||||
" cache_ttl_seconds: 0\n"
|
||||
" override_existing: false\n"
|
||||
" auto_install: false\n"
|
||||
)
|
||||
monkeypatch.setenv("HERMES_HOME", str(home))
|
||||
monkeypatch.setenv("BWS_ACCESS_TOKEN", "0.t")
|
||||
monkeypatch.delenv("MY_BSM_KEY", raising=False)
|
||||
|
||||
called = {"n": 0}
|
||||
|
||||
def fake_fetch(**kwargs):
|
||||
called["n"] += 1
|
||||
assert kwargs["project_id"] == "proj-1"
|
||||
return {"MY_BSM_KEY": "from-bsm"}, []
|
||||
|
||||
monkeypatch.setattr(
|
||||
"agent.secret_sources.bitwarden.find_bws",
|
||||
lambda **_kw: Path("/fake/bws"),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"agent.secret_sources.bitwarden.fetch_bitwarden_secrets",
|
||||
fake_fetch,
|
||||
)
|
||||
from agent.secret_sources import registry as reg_module
|
||||
|
||||
reg_module._reset_registry_for_tests()
|
||||
|
||||
from hermes_cli.env_loader import _apply_external_secret_sources
|
||||
_apply_external_secret_sources(home)
|
||||
|
||||
assert called["n"] == 1
|
||||
assert os.environ.get("MY_BSM_KEY") == "from-bsm"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Disk-persisted cache (cross-process — speeds up back-to-back CLI invocations)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_disk_cache_key_mismatch_triggers_refetch(monkeypatch, tmp_path):
|
||||
"""Disk cache entry written by a different token/project is ignored."""
|
||||
home = tmp_path / ".hermes"
|
||||
home.mkdir()
|
||||
fake_binary = tmp_path / "bws"
|
||||
fake_binary.write_text("")
|
||||
payload = _fake_bws_payload([{"key": "K1", "value": "v1"}])
|
||||
|
||||
call_count = {"n": 0}
|
||||
def fake_run(*a, **kw):
|
||||
call_count["n"] += 1
|
||||
return mock.Mock(returncode=0, stdout=payload, stderr="")
|
||||
monkeypatch.setattr(bw.subprocess, "run", fake_run)
|
||||
bw._reset_cache_for_tests(home)
|
||||
|
||||
# Write a cache entry for a DIFFERENT token/project pair
|
||||
cache_path = bw._disk_cache_path(home)
|
||||
cache_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
cache_path.write_text(json.dumps({
|
||||
"key": "deadbeef00000000|other-project|",
|
||||
"secrets": {"OTHER": "should-not-leak"},
|
||||
"fetched_at": time.time(),
|
||||
}))
|
||||
|
||||
secrets, _ = bw.fetch_bitwarden_secrets(
|
||||
access_token="0.t", project_id="proj-1", binary=fake_binary,
|
||||
cache_ttl_seconds=300, home_path=home,
|
||||
)
|
||||
# We must NOT have used the foreign cache entry
|
||||
assert secrets == {"K1": "v1"}
|
||||
assert "OTHER" not in secrets
|
||||
assert call_count["n"] == 1
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_encrypted_cache_writes_without_plaintext(monkeypatch, tmp_path):
|
||||
"""Encrypted cache stores last-good secrets without raw values on disk."""
|
||||
home = tmp_path / ".hermes"
|
||||
home.mkdir()
|
||||
fake_binary = tmp_path / "bws"
|
||||
fake_binary.write_text("")
|
||||
payload = _fake_bws_payload([{"key": "K1", "value": "secret-value"}])
|
||||
|
||||
monkeypatch.setattr(
|
||||
bw.subprocess,
|
||||
"run",
|
||||
lambda *a, **kw: mock.Mock(returncode=0, stdout=payload, stderr=""),
|
||||
)
|
||||
bw._reset_cache_for_tests(home)
|
||||
# A successful encrypted write must remove a pre-existing legacy plaintext
|
||||
# cache from the migration path.
|
||||
legacy_key = (bw._token_fingerprint("0.t"), "proj-1", "")
|
||||
bw._DISK_CACHE.write(
|
||||
legacy_key,
|
||||
bw._CachedFetch(secrets={"K1": "legacy"}, fetched_at=time.time()),
|
||||
300,
|
||||
home,
|
||||
)
|
||||
assert bw._disk_cache_path(home).exists()
|
||||
|
||||
secrets, warnings = bw.fetch_bitwarden_secrets(
|
||||
access_token="0.t", project_id="proj-1", binary=fake_binary,
|
||||
cache_ttl_seconds=0, encrypted_cache_enabled=True,
|
||||
encrypted_cache_max_stale_seconds=604800, home_path=home,
|
||||
)
|
||||
|
||||
assert secrets == {"K1": "secret-value"}
|
||||
assert warnings == []
|
||||
assert not bw._disk_cache_path(home).exists()
|
||||
cache_path = bw._encrypted_disk_cache_path(home)
|
||||
assert cache_path.exists()
|
||||
mode = stat.S_IMODE(os.stat(cache_path).st_mode)
|
||||
assert mode == 0o600, f"expected 0o600, got 0o{mode:o}"
|
||||
text = cache_path.read_text()
|
||||
assert "secret-value" not in text
|
||||
assert "0.t" not in text
|
||||
payload_disk = json.loads(text)
|
||||
assert set(payload_disk.keys()) == {
|
||||
"version", "key", "salt", "nonce", "ciphertext",
|
||||
}
|
||||
assert not bw._disk_cache_path(home).exists()
|
||||
|
||||
|
||||
|
||||
|
||||
def test_encrypted_cache_falls_back_on_network_error(monkeypatch, tmp_path):
|
||||
"""A fresh-enough encrypted cache is used when BWS is unreachable."""
|
||||
home = tmp_path / ".hermes"
|
||||
home.mkdir()
|
||||
fake_binary = tmp_path / "bws"
|
||||
fake_binary.write_text("")
|
||||
calls = {"n": 0}
|
||||
|
||||
def fake_run(*a, **kw):
|
||||
calls["n"] += 1
|
||||
if calls["n"] == 1:
|
||||
return mock.Mock(
|
||||
returncode=0,
|
||||
stdout=_fake_bws_payload([{"key": "K1", "value": "cached"}]),
|
||||
stderr="",
|
||||
)
|
||||
return mock.Mock(
|
||||
returncode=1,
|
||||
stdout="",
|
||||
stderr="Error: network is unreachable",
|
||||
)
|
||||
|
||||
monkeypatch.setattr(bw.subprocess, "run", fake_run)
|
||||
bw._reset_cache_for_tests(home)
|
||||
|
||||
first, _ = bw.fetch_bitwarden_secrets(
|
||||
access_token="0.t", project_id="proj-1", binary=fake_binary,
|
||||
cache_ttl_seconds=0, encrypted_cache_enabled=True,
|
||||
encrypted_cache_max_stale_seconds=604800, home_path=home,
|
||||
)
|
||||
assert first == {"K1": "cached"}
|
||||
bw._CACHE.clear()
|
||||
|
||||
second, warnings = bw.fetch_bitwarden_secrets(
|
||||
access_token="0.t", project_id="proj-1", binary=fake_binary,
|
||||
cache_ttl_seconds=0, encrypted_cache_enabled=True,
|
||||
encrypted_cache_max_stale_seconds=604800, home_path=home,
|
||||
)
|
||||
assert second == {"K1": "cached"}
|
||||
assert calls["n"] == 2
|
||||
assert len(warnings) == 1
|
||||
assert "stale ENCRYPTED disk cache" in warnings[0]
|
||||
assert "bws live fetch failed" in warnings[0]
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Stale disk cache fallback when live bws fetch fails
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _seed_stale_disk_cache(home, *, secrets, age_seconds, project_id="proj-1",
|
||||
access_token="0.t", server_url=""):
|
||||
"""Populate the disk cache as if a successful fetch happened `age_seconds`
|
||||
ago. Writes the JSON payload directly (same shape the shared DiskCache
|
||||
reads/writes) rather than going through DiskCache.write, since that
|
||||
would honor cache_ttl_seconds and refuse to persist an already-"stale"
|
||||
entry — this needs to land on disk regardless of TTL."""
|
||||
cache_key = (
|
||||
bw._token_fingerprint(access_token), project_id, server_url,
|
||||
)
|
||||
cache_path = bw._disk_cache_path(home)
|
||||
cache_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
cache_path.write_text(json.dumps({
|
||||
"key": bw._cache_key_str(cache_key),
|
||||
"secrets": secrets,
|
||||
"fetched_at": time.time() - age_seconds,
|
||||
}))
|
||||
|
||||
|
||||
def test_stale_disk_cache_returned_when_bws_fails(monkeypatch, tmp_path):
|
||||
"""When bws fails and the disk cache is stale, return the stale secrets
|
||||
with a warning rather than raising."""
|
||||
home = tmp_path / ".hermes"
|
||||
home.mkdir()
|
||||
fake_binary = tmp_path / "bws"
|
||||
fake_binary.write_text("")
|
||||
bw._reset_cache_for_tests(home)
|
||||
|
||||
# Seed a stale (older than TTL) disk cache from a previous successful fetch
|
||||
_seed_stale_disk_cache(home, secrets={"OPENAI_API_KEY": "sk-old"},
|
||||
age_seconds=3600)
|
||||
|
||||
# Now simulate a BWS network failure
|
||||
def fail_run(*a, **kw):
|
||||
return mock.Mock(returncode=1, stdout="",
|
||||
stderr="Error: dns resolution failed")
|
||||
monkeypatch.setattr(bw.subprocess, "run", fail_run)
|
||||
|
||||
secrets, warnings = bw.fetch_bitwarden_secrets(
|
||||
access_token="0.t", project_id="proj-1", binary=fake_binary,
|
||||
cache_ttl_seconds=300, home_path=home,
|
||||
)
|
||||
assert secrets == {"OPENAI_API_KEY": "sk-old"}
|
||||
assert len(warnings) == 1
|
||||
assert "stale disk cache" in warnings[0]
|
||||
assert "dns resolution failed" in warnings[0]
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_stale_fallback_skipped_on_auth_failure(monkeypatch, tmp_path):
|
||||
"""An AUTH_FAILED bws error must raise, not serve stale secrets — a bad
|
||||
access token indicates a real credential problem the caller needs to
|
||||
see, not a transient outage worth papering over."""
|
||||
home = tmp_path / ".hermes"
|
||||
home.mkdir()
|
||||
fake_binary = tmp_path / "bws"
|
||||
fake_binary.write_text("")
|
||||
bw._reset_cache_for_tests(home)
|
||||
|
||||
_seed_stale_disk_cache(home, secrets={"K1": "v1"}, age_seconds=3600)
|
||||
|
||||
monkeypatch.setattr(
|
||||
bw.subprocess, "run",
|
||||
lambda *a, **kw: mock.Mock(returncode=1, stdout="",
|
||||
stderr="Error: unauthorized (401)"),
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="unauthorized"):
|
||||
bw.fetch_bitwarden_secrets(
|
||||
access_token="0.t", project_id="proj-1", binary=fake_binary,
|
||||
cache_ttl_seconds=300, home_path=home,
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,319 @@
|
||||
"""Tests that callable api_key (Entra ID bearer provider) flows through
|
||||
the agent stack without coercion.
|
||||
|
||||
The OpenAI Python SDK accepts ``api_key: str | None | Callable[[], str]``,
|
||||
and ``azure-identity``'s ``get_bearer_token_provider`` returns a callable.
|
||||
Hermes preserves the callable end-to-end so the SDK refreshes tokens
|
||||
transparently. This file pins the contract at the high-risk seams the
|
||||
rubber-duck audit identified.
|
||||
|
||||
Covered:
|
||||
* ``_create_openai_client`` passes a callable ``api_key`` straight
|
||||
through to ``openai.OpenAI(...)``.
|
||||
* ``_normalize_main_runtime`` preserves the callable so auxiliary
|
||||
clients inherit Entra auth.
|
||||
* ``_truncate_token`` (dashboard preview) renders ``"<entra-id-bearer>"``
|
||||
instead of ``"<function ...>"`` and never invokes the callable.
|
||||
* ``run_agent.py`` masked-banner path renders the Entra placeholder
|
||||
and never tries to slice/len the callable.
|
||||
* Serialization scrub: dumping a runtime dict via ``json.dumps`` with
|
||||
a callable api_key raises (default behaviour) — guards against
|
||||
silently leaking ``"<function ...>"`` strings into event logs.
|
||||
* ``batch_runner`` strips the callable from the worker config dict
|
||||
so multiprocessing.Pool can pickle the rest.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import cast
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# OpenAI SDK construction preserves the callable
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCreateOpenAIClientCallable:
|
||||
"""``AIAgent._create_openai_client`` must pass the callable through
|
||||
to ``openai.OpenAI(...)`` without coercion."""
|
||||
|
||||
def test_callable_api_key_passed_to_openai_constructor(self, monkeypatch):
|
||||
"""Construct the smallest possible AIAgent surface and verify
|
||||
the OpenAI client receives the callable unchanged."""
|
||||
captured = {}
|
||||
|
||||
def fake_openai(**kwargs):
|
||||
captured["kwargs"] = kwargs
|
||||
return MagicMock(api_key=kwargs.get("api_key"))
|
||||
|
||||
# Patch the module-level OpenAI proxy used by ``_create_openai_client``.
|
||||
monkeypatch.setattr("agent.process_bootstrap.OpenAI", fake_openai)
|
||||
|
||||
# Build a minimal stand-in for AIAgent so we can call the bound
|
||||
# method directly without paying the full __init__ cost.
|
||||
from run_agent import AIAgent
|
||||
|
||||
agent = AIAgent.__new__(AIAgent)
|
||||
# Attributes consulted by _create_openai_client / _client_log_context.
|
||||
agent.provider = "azure-foundry"
|
||||
agent.model = "gpt-4o"
|
||||
agent.base_url = "https://r.openai.azure.com/openai/v1"
|
||||
agent._client_kwargs = {}
|
||||
|
||||
def token_provider():
|
||||
return "fresh-jwt"
|
||||
|
||||
client_kwargs = {
|
||||
"api_key": token_provider,
|
||||
"base_url": "https://r.openai.azure.com/openai/v1",
|
||||
}
|
||||
client = agent._create_openai_client(client_kwargs, reason="test", shared=False)
|
||||
|
||||
# The OpenAI constructor must receive the *callable*, not a string.
|
||||
forwarded = captured["kwargs"]["api_key"]
|
||||
assert callable(forwarded)
|
||||
assert not isinstance(forwarded, str)
|
||||
assert forwarded is token_provider, (
|
||||
"_create_openai_client must not wrap or coerce the callable"
|
||||
)
|
||||
assert client is not None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auxiliary runtime preserves the callable
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestNormalizeMainRuntimePreservesCallable:
|
||||
"""The aux client orchestrator must keep the callable on the
|
||||
runtime dict so compression / vision / embedding / title-gen clients
|
||||
inherit Entra ID auth from the main agent."""
|
||||
|
||||
def test_callable_api_key_survives_normalization(self):
|
||||
from agent.auxiliary_client import _normalize_main_runtime
|
||||
|
||||
def provider():
|
||||
return "jwt"
|
||||
|
||||
normalized = _normalize_main_runtime({
|
||||
"provider": "azure-foundry",
|
||||
"model": "gpt-4o",
|
||||
"base_url": "https://r.openai.azure.com/openai/v1",
|
||||
"api_key": provider,
|
||||
"api_mode": "chat_completions",
|
||||
"auth_mode": "entra_id",
|
||||
})
|
||||
assert normalized["api_key"] is provider
|
||||
assert normalized["auth_mode"] == "entra_id"
|
||||
|
||||
def test_string_api_key_still_works(self):
|
||||
from agent.auxiliary_client import _normalize_main_runtime
|
||||
normalized = _normalize_main_runtime({
|
||||
"provider": "azure-foundry",
|
||||
"api_key": "sk-static",
|
||||
})
|
||||
assert normalized["api_key"] == "sk-static"
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Display surfaces never invoke the callable
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestTruncateTokenCallable:
|
||||
def test_callable_returns_placeholder(self):
|
||||
"""Dashboard preview must render the Entra placeholder, NOT
|
||||
``"<function ...>"``."""
|
||||
from hermes_cli.web_server_oauth import _truncate_token
|
||||
|
||||
invoked = {"count": 0}
|
||||
|
||||
def provider():
|
||||
invoked["count"] += 1
|
||||
return "should-not-appear-in-ui"
|
||||
|
||||
token_provider = cast(str | None, provider)
|
||||
rendered = _truncate_token(token_provider)
|
||||
assert rendered == "<entra-id-bearer>"
|
||||
assert invoked["count"] == 0
|
||||
|
||||
def test_string_jwt_still_truncated_to_signature_tail(self):
|
||||
from hermes_cli.web_server_oauth import _truncate_token
|
||||
# JWT shape: header.payload.signature → only signature tail shown.
|
||||
out = _truncate_token("aaaa.bbbb.cccccccsig", visible=4)
|
||||
assert out == "…csig"
|
||||
|
||||
def test_empty_returns_empty(self):
|
||||
from hermes_cli.web_server_oauth import _truncate_token
|
||||
assert _truncate_token(None) == ""
|
||||
assert _truncate_token("") == ""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Serialization scrub — runtime dicts with callables must NOT silently
|
||||
# JSON-encode as ``"<function ...>"`` (would leak garbage into events).
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRuntimeDictSerializationGuard:
|
||||
def test_json_dumps_default_str_does_not_silently_stringify_callable(self):
|
||||
"""Sanity check: a runtime dict with a callable api_key must
|
||||
either raise on plain ``json.dumps`` (good — fail loud) or be
|
||||
sanitized BEFORE serialization. This test pins the loud-fail
|
||||
behaviour so future changes that introduce
|
||||
``json.dumps(..., default=str)`` over a runtime dict are caught
|
||||
by a regression here."""
|
||||
|
||||
def provider():
|
||||
return "jwt"
|
||||
|
||||
runtime = {
|
||||
"provider": "azure-foundry",
|
||||
"api_key": provider,
|
||||
"auth_mode": "entra_id",
|
||||
}
|
||||
# Plain json.dumps — must raise, not silently produce
|
||||
# ``"<function provider at 0x...>"``.
|
||||
with pytest.raises(TypeError):
|
||||
json.dumps(runtime)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# batch_runner strips callables from the worker config dict
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestBatchRunnerCallableHandling:
|
||||
def test_callable_api_key_stripped_from_worker_config(self, capsys, monkeypatch, tmp_path):
|
||||
"""``BatchRunner._run_batches`` (or the equivalent code path)
|
||||
must replace a callable api_key with None before pickling the
|
||||
worker config dict — otherwise multiprocessing.Pool fails."""
|
||||
# We can't easily run BatchRunner end-to-end in a unit test
|
||||
# (it spawns subprocesses), but we CAN inline the same logic:
|
||||
# the production code uses ``callable(self.api_key) and not
|
||||
# isinstance(self.api_key, str)`` to gate the substitution.
|
||||
# Re-execute the same predicate here as a contract guard.
|
||||
|
||||
def provider():
|
||||
return "jwt"
|
||||
|
||||
api_key = provider
|
||||
worker_api_key = None if (callable(api_key) and not isinstance(api_key, str)) else api_key
|
||||
assert worker_api_key is None, (
|
||||
"BatchRunner must replace callable api_key with None so "
|
||||
"multiprocessing.Pool can pickle the worker config"
|
||||
)
|
||||
|
||||
# And a string passes through unchanged.
|
||||
api_key_str = "sk-static"
|
||||
worker_api_key_str = None if (callable(api_key_str) and not isinstance(api_key_str, str)) else api_key_str
|
||||
assert worker_api_key_str == "sk-static"
|
||||
|
||||
def test_batch_runner_source_uses_the_correct_predicate(self):
|
||||
"""Pin the predicate string in batch_runner so refactors that
|
||||
change it are caught here. Reading the source rather than
|
||||
importing avoids spinning up the full BatchRunner."""
|
||||
from pathlib import Path
|
||||
src = (Path(__file__).resolve().parent.parent.parent
|
||||
/ "batch_runner.py").read_text()
|
||||
assert "callable(self.api_key) and not isinstance(self.api_key, str)" in src, (
|
||||
"BatchRunner.api_key callable check changed — update test or "
|
||||
"verify the new predicate still routes Entra token providers "
|
||||
"to the worker-rebuild path."
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Inline masked-banner / display sites (callable-aware)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCliEnsureRuntimeCredentialsCallable:
|
||||
"""Regression: ``cli.py:_ensure_runtime_credentials`` previously
|
||||
treated a callable ``api_key`` as "not a string" and overwrote it
|
||||
with the ``"no-key-required"`` placeholder, which then got sent as
|
||||
``Authorization: Bearer no-key-required`` and rejected by Azure
|
||||
with a 401. This is the most subtle of the callable-api_key audit
|
||||
sites — gated by ``not isinstance(api_key, str)`` rather than the
|
||||
cleaner ``callable(...)`` check used elsewhere.
|
||||
|
||||
We verify the source pattern (rather than spinning up a real
|
||||
``HermesCLI`` instance) — the predicate change is the load-bearing
|
||||
fix and is invariant under the surrounding orchestration code."""
|
||||
|
||||
def test_callable_predicate_present_in_cli_runtime_validation(self):
|
||||
from pathlib import Path
|
||||
# ``_ensure_runtime_credentials`` was extracted from cli.py into the
|
||||
# ``CLIAgentSetupMixin`` (god-file decomposition Phase 4). Read the
|
||||
# module the method actually lives in now.
|
||||
src = (Path(__file__).resolve().parent.parent.parent
|
||||
/ "hermes_cli" / "cli_agent_setup_mixin.py").read_text()
|
||||
# The fix gates the string-only check on ``callable(api_key)`` so callable
|
||||
# token providers survive.
|
||||
assert "if not callable(api_key) and not (isinstance(api_key, str) and api_key):" in src, (
|
||||
"_ensure_runtime_credentials must preserve a callable "
|
||||
"api_key (Entra ID bearer provider). Without the guard, the "
|
||||
"callable is stringified to 'no-key-required' and Azure 401s."
|
||||
)
|
||||
|
||||
|
||||
class TestInlinedDisplayMasks:
|
||||
"""The masked-credential display sites are now inlined per-site (no
|
||||
shared helper). Each site uses the ``is_token_provider`` predicate
|
||||
to short-circuit on callables and print a static
|
||||
``"Microsoft Entra ID"`` label, then falls through to its own
|
||||
context-appropriate string mask. This replaces a unified helper
|
||||
that would have forced one mask shape across sites with legitimately
|
||||
different display needs (banner vs diagnostic vs UI vs preview)."""
|
||||
|
||||
def test_run_agent_banner_uses_is_token_provider_guard(self):
|
||||
"""The masked-banner sites live in ``agent/agent_init.py``
|
||||
(the ``__init__`` body was extracted into ``init_agent`` after
|
||||
this feature was first written). Both the OpenAI and Anthropic
|
||||
client init paths must guard their banner prints with
|
||||
``is_token_provider`` so a callable Entra ID provider doesn't
|
||||
crash ``len(api_key)``."""
|
||||
from pathlib import Path
|
||||
src = (Path(__file__).resolve().parent.parent.parent
|
||||
/ "agent" / "agent_init.py").read_text()
|
||||
# Both banner paths route through the shared ``_print_key_banner`` helper,
|
||||
# which owns the single ``is_token_provider`` guard.
|
||||
assert src.count("_print_key_banner(") >= 3, (
|
||||
"agent/agent_init.py must guard BOTH masked-banner paths "
|
||||
"(chat_completions and anthropic_messages) with "
|
||||
"is_token_provider() via _print_key_banner()."
|
||||
)
|
||||
assert "is_token_provider(" in src
|
||||
assert '"🔑 Using credentials: Microsoft Entra ID"' in src, (
|
||||
"agent/agent_init.py banner helper should print a static "
|
||||
"'Microsoft Entra ID' label for callable api_keys — no "
|
||||
"placeholder plumbing, no describe-mask fallback."
|
||||
)
|
||||
|
||||
def test_cli_show_config_handles_callable(self):
|
||||
"""``cli.HermesCLI.show_config`` previously did
|
||||
``self.api_key[-4:]`` / ``len(self.api_key)`` which crashes on
|
||||
callable Entra ID providers. The inlined version uses
|
||||
``is_token_provider`` and prints the same static label as the
|
||||
run_agent banners."""
|
||||
from pathlib import Path
|
||||
src = (Path(__file__).resolve().parent.parent.parent
|
||||
/ "cli.py").read_text()
|
||||
assert "is_token_provider(display_key)" in src, (
|
||||
"cli.HermesCLI.show_config must guard the displayed key via "
|
||||
"is_token_provider so callable Entra ID providers don't "
|
||||
"crash /config."
|
||||
)
|
||||
assert '"Microsoft Entra ID"' in src, (
|
||||
"cli.HermesCLI.show_config must print the static "
|
||||
"'Microsoft Entra ID' label (matching run_agent banners) "
|
||||
"instead of attempting to slice the callable."
|
||||
)
|
||||
|
||||
|
||||
@@ -0,0 +1,142 @@
|
||||
"""Regression guard for the cascading-interrupt hang (PR #6600).
|
||||
|
||||
Original diagnosis and fix by Kristian Vastveit (@kristianvast) in PR #6600,
|
||||
against the then-inline ``_interruptible_api_call`` /
|
||||
``_interruptible_streaming_api_call`` methods in run_agent.py. Those methods
|
||||
have since been extracted into ``agent/chat_completion_helpers.py``, so the
|
||||
fix is reapplied there and these tests target the extracted functions.
|
||||
|
||||
The bug: when ``agent.interrupt()`` fires during an active LLM call, the main
|
||||
poll loop force-closes the worker-local httpx client to stop token generation.
|
||||
That raises a transport error (RemoteProtocolError) on the worker — the
|
||||
EXPECTED consequence of our own close, not a network bug. The streaming retry
|
||||
loop misclassified it as a transient connection error and retried, each doomed
|
||||
retry stalling for the full stream-stale timeout (up to 300s). Because the
|
||||
gateway caches AIAgent instances per session, the stale worker outlived the
|
||||
turn and raced the next turn's request — the root of the multi-minute
|
||||
cascading-interrupt hang.
|
||||
|
||||
The fix: a request-local ``_request_cancelled`` token set by the poll loop
|
||||
right before the force-close. The worker's exception handler checks it and
|
||||
exits cleanly (no retry, no fallback, no "reconnecting" status) instead of
|
||||
treating the forced error as transient.
|
||||
"""
|
||||
import threading
|
||||
import time
|
||||
import types
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from agent import chat_completion_helpers as cch
|
||||
|
||||
|
||||
|
||||
|
||||
def _make_agent():
|
||||
"""A MagicMock agent wired with just enough surface for the helpers."""
|
||||
agent = MagicMock()
|
||||
agent.api_mode = "chat_completions"
|
||||
agent._interrupt_requested = False
|
||||
agent.verbose_logging = False
|
||||
# _compute_non_stream_stale_timeout / streaming setup helpers return
|
||||
# benign values; the real call path is mocked per-test.
|
||||
agent._compute_non_stream_stale_timeout.return_value = 5.0
|
||||
return agent
|
||||
|
||||
|
||||
def test_non_streaming_cancel_does_not_surface_network_error():
|
||||
"""A force-close during a non-streaming call must raise InterruptedError,
|
||||
not the swallowed transport error."""
|
||||
agent = _make_agent()
|
||||
|
||||
create_calls = {"n": 0}
|
||||
fake_client = MagicMock()
|
||||
|
||||
def _create(**kwargs):
|
||||
create_calls["n"] += 1
|
||||
# Simulate the main thread firing an interrupt mid-call, then the
|
||||
# force-close raising a transport error on this worker.
|
||||
agent._interrupt_requested = True
|
||||
time.sleep(0.3) # let the poll loop observe the interrupt + force-close
|
||||
raise httpx.RemoteProtocolError("peer closed connection")
|
||||
|
||||
fake_client.chat.completions.create.side_effect = _create
|
||||
agent._create_request_openai_client.return_value = fake_client
|
||||
agent._close_request_openai_client = MagicMock()
|
||||
agent._abort_request_openai_client = MagicMock()
|
||||
|
||||
t0 = time.time()
|
||||
with pytest.raises(InterruptedError):
|
||||
cch.interruptible_api_call(agent, {"model": "x", "messages": []})
|
||||
elapsed = time.time() - t0
|
||||
|
||||
# The forced RemoteProtocolError must NOT surface as the raised error.
|
||||
assert create_calls["n"] == 1
|
||||
assert elapsed < 10.0, f"interrupt took {elapsed:.1f}s — should be near-instant (guarding the 30s+ hang)"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# #67142: direct-Anthropic stale/interrupt watchdog must abort the request-local
|
||||
# client from the poll (stranger) thread and NEVER close/rebuild the shared
|
||||
# _anthropic_client — closing it there released a live TLS FD that the kernel
|
||||
# recycled into a SQLite handle, writing a TLS record over a DB header.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_anthropic_agent():
|
||||
agent = _make_agent()
|
||||
agent.api_mode = "anthropic_messages"
|
||||
return agent
|
||||
|
||||
|
||||
def _wait_for_mock_call(mock, timeout=3.0):
|
||||
deadline = time.time() + timeout
|
||||
while time.time() < deadline:
|
||||
if mock.called:
|
||||
return
|
||||
time.sleep(0.02)
|
||||
raise AssertionError(f"{mock!r} was not called within {timeout}s")
|
||||
|
||||
|
||||
def test_anthropic_non_streaming_stale_aborts_request_client_not_shared():
|
||||
"""Stale non-streaming Anthropic call: the poll thread aborts the
|
||||
request-local client's socket; the shared client is never closed/rebuilt,
|
||||
and the worker still unblocks and closes its own client (no #28161 hang)."""
|
||||
agent = _make_anthropic_agent()
|
||||
agent._compute_non_stream_stale_timeout.return_value = 0.05
|
||||
agent._codex_silent_hang_hint = MagicMock(return_value=None)
|
||||
|
||||
request_client = MagicMock()
|
||||
agent._create_request_anthropic_client = MagicMock(return_value=request_client)
|
||||
agent._abort_request_anthropic_client = MagicMock()
|
||||
agent._close_request_anthropic_client = MagicMock()
|
||||
|
||||
def _create(_api_kwargs, *, client):
|
||||
assert client is request_client
|
||||
# Outlive the 0.05s stale timeout AND the worker join (2.0s) so the
|
||||
# stale detector surfaces its TimeoutError.
|
||||
time.sleep(2.5)
|
||||
return object()
|
||||
|
||||
agent._anthropic_messages_create = MagicMock(side_effect=_create)
|
||||
|
||||
with pytest.raises(TimeoutError):
|
||||
cch.interruptible_api_call(agent, {"model": "x", "messages": []})
|
||||
|
||||
# Shared client untouched from the poll thread.
|
||||
agent._anthropic_client.close.assert_not_called()
|
||||
agent._rebuild_anthropic_client.assert_not_called()
|
||||
# Poll (stranger) thread aborts the request-local client's socket only.
|
||||
agent._abort_request_anthropic_client.assert_called_once_with(
|
||||
request_client, reason="stale_call_kill"
|
||||
)
|
||||
# Worker unblocks and closes its own request client from its own thread.
|
||||
_wait_for_mock_call(agent._close_request_anthropic_client)
|
||||
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
"""Chat-completions request-transform bypass (#93650 extended to chat.completions).
|
||||
|
||||
``chat.completions.create`` re-walks the whole request body against the
|
||||
``CompletionCreateParams`` union client-side, GIL held, before any byte leaves
|
||||
the process. The bypass moves the already-wire-format bulk fields into
|
||||
``extra_body``; its whole safety argument is that the server receives the same
|
||||
bytes, so that is what these tests pin.
|
||||
"""
|
||||
|
||||
import sys
|
||||
import types
|
||||
|
||||
sys.modules.setdefault("fire", types.SimpleNamespace(Fire=lambda *a, **k: None))
|
||||
sys.modules.setdefault("firecrawl", types.SimpleNamespace(Firecrawl=object))
|
||||
sys.modules.setdefault("fal_client", types.SimpleNamespace())
|
||||
|
||||
import httpx
|
||||
import openai
|
||||
|
||||
from agent.sdk_transform_bypass import ESCAPE_HATCH_ENV, bypass_chat_sdk_request_transform
|
||||
|
||||
_SSE = (
|
||||
b'data: {"id":"1","object":"chat.completion.chunk","created":1,"model":"m",'
|
||||
b'"choices":[{"index":0,"delta":{"content":"hi"},"finish_reason":null}]}\n\n'
|
||||
b"data: [DONE]\n\n"
|
||||
)
|
||||
|
||||
|
||||
def _wire_body() -> dict:
|
||||
"""Production-shaped chat body: content parts incl. an image, a tool_calls turn, a tool
|
||||
result, function tool schemas, and a caller-populated extra_body (reasoning/provider)."""
|
||||
return {
|
||||
"model": "hermes-4-70b",
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are Hermes."},
|
||||
{"role": "user", "content": [
|
||||
{"type": "text", "text": "look at this"},
|
||||
{"type": "image_url", "image_url": {"url": "https://e.example/i.png", "detail": "low"}},
|
||||
]},
|
||||
{"role": "assistant", "content": None, "tool_calls": [
|
||||
{"id": "c1", "type": "function", "function": {"name": "terminal", "arguments": '{"cmd":"ls"}'}},
|
||||
]},
|
||||
{"role": "tool", "tool_call_id": "c1", "content": "total 0"},
|
||||
],
|
||||
"tools": [
|
||||
{"type": "function", "function": {"name": f"tool_{i}", "description": "d",
|
||||
"parameters": {"type": "object", "properties": {"p": {"type": "string"}}, "required": ["p"]}}}
|
||||
for i in range(3)
|
||||
],
|
||||
"tool_choice": "auto",
|
||||
"stream": True,
|
||||
"temperature": 0.7,
|
||||
"stream_options": {"include_usage": True},
|
||||
"extra_body": {"reasoning": {"effort": "high"}, "provider": {"order": ["nous"]}},
|
||||
}
|
||||
|
||||
|
||||
class _Recorder:
|
||||
"""A real openai.OpenAI client whose transport records the request bytes."""
|
||||
|
||||
def __init__(self):
|
||||
self.content: bytes | None = None
|
||||
self.client = openai.OpenAI(
|
||||
api_key="k", base_url="https://chat.invalid/v1",
|
||||
http_client=httpx.Client(transport=httpx.MockTransport(self._handle)),
|
||||
)
|
||||
|
||||
def _handle(self, request: httpx.Request) -> httpx.Response:
|
||||
self.content = request.content
|
||||
return httpx.Response(200, content=_SSE, headers={"content-type": "text/event-stream"})
|
||||
|
||||
def send(self, kwargs: dict) -> bytes:
|
||||
for _ in self.client.chat.completions.create(**kwargs):
|
||||
pass
|
||||
assert self.content is not None
|
||||
return self.content
|
||||
|
||||
|
||||
def test_bulk_fields_ride_in_extra_body_and_the_wire_bytes_are_identical():
|
||||
"""Same bytes → same server behaviour and the same byte-keyed prompt-cache prefix."""
|
||||
recorder = _Recorder()
|
||||
body = _wire_body()
|
||||
|
||||
moved = bypass_chat_sdk_request_transform(dict(body), recorder.client)
|
||||
|
||||
assert moved["messages"] == [] and moved["tools"] == []
|
||||
assert moved["extra_body"]["messages"] == body["messages"]
|
||||
assert moved["extra_body"]["tools"] == body["tools"]
|
||||
assert moved["extra_body"]["reasoning"] == body["extra_body"]["reasoning"]
|
||||
assert recorder.send(moved) == recorder.send(dict(body))
|
||||
|
||||
|
||||
def test_escape_hatch_and_non_sdk_facades_keep_the_typed_path(monkeypatch):
|
||||
"""Both rails hand the kwargs back untouched: the env hatch, and a chat-shaped facade
|
||||
that is not the SDK (MoA aggregator, test stand-ins) — it never merges ``extra_body``."""
|
||||
recorder = _Recorder()
|
||||
kwargs = _wire_body()
|
||||
|
||||
facade = types.SimpleNamespace(chat=types.SimpleNamespace(completions=types.SimpleNamespace(create=lambda **kw: kw)))
|
||||
assert bypass_chat_sdk_request_transform(kwargs, facade) is kwargs
|
||||
|
||||
monkeypatch.setenv(ESCAPE_HATCH_ENV, "1")
|
||||
assert bypass_chat_sdk_request_transform(kwargs, recorder.client) is kwargs
|
||||
@@ -0,0 +1,323 @@
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from agent.codex_runtime import _record_codex_app_server_compaction
|
||||
from agent.conversation_compression import COMPACTION_DONE_STATUS, COMPACTION_STATUS, compress_context
|
||||
from agent.transports.codex_app_server_session import TurnResult
|
||||
|
||||
|
||||
class FakeCodexSession:
|
||||
def __init__(self, result):
|
||||
self.result = result
|
||||
self.calls = 0
|
||||
self.closed = False
|
||||
|
||||
def compact_thread(self):
|
||||
self.calls += 1
|
||||
return self.result
|
||||
|
||||
def close(self):
|
||||
self.closed = True
|
||||
|
||||
|
||||
class SlowCodexSession(FakeCodexSession):
|
||||
def __init__(self, result, touch_calls):
|
||||
super().__init__(result)
|
||||
self.touch_calls = touch_calls
|
||||
|
||||
def compact_thread(self):
|
||||
self.calls += 1
|
||||
_wait_for_touch(self.touch_calls, "context compression in progress")
|
||||
return self.result
|
||||
|
||||
|
||||
def _wait_for_touch(touch_calls, desc, timeout=1.0):
|
||||
deadline = time.monotonic() + timeout
|
||||
while time.monotonic() < deadline:
|
||||
if desc in touch_calls:
|
||||
return
|
||||
time.sleep(0.01)
|
||||
pytest.fail(f"timed out waiting for touch {desc!r}; saw {touch_calls!r}")
|
||||
|
||||
|
||||
class DummyAgent:
|
||||
def __init__(
|
||||
self,
|
||||
result,
|
||||
*,
|
||||
auto_compaction="native",
|
||||
):
|
||||
self.api_mode = "codex_app_server"
|
||||
self.codex_app_server_auto_compaction = auto_compaction
|
||||
self.session_id = "hermes-session-1"
|
||||
self.platform = "cli"
|
||||
self._cached_system_prompt = "cached prompt"
|
||||
self._codex_session = FakeCodexSession(result)
|
||||
self.context_compressor = SimpleNamespace(
|
||||
compression_count=0,
|
||||
last_compression_rough_tokens=0,
|
||||
last_prompt_tokens=123,
|
||||
last_completion_tokens=45,
|
||||
awaiting_real_usage_after_compression=False,
|
||||
)
|
||||
self.statuses = []
|
||||
self.status_events = []
|
||||
self.status_callback = lambda kind, text: self.status_events.append((kind, text))
|
||||
self.warnings = []
|
||||
self.events = []
|
||||
self.built_prompts = []
|
||||
self.touch_calls = []
|
||||
self.touch_provenances = []
|
||||
self._compression_activity_heartbeat_interval = 0.1
|
||||
|
||||
def _touch_activity(self, desc, *, provenance=None, force_persist=False):
|
||||
self.touch_calls.append(desc)
|
||||
self.touch_provenances.append(provenance)
|
||||
|
||||
def _emit_status(self, message):
|
||||
self.statuses.append(message)
|
||||
self.status_callback("lifecycle", message)
|
||||
|
||||
def _emit_warning(self, message):
|
||||
self.warnings.append(message)
|
||||
self.status_callback("warn", message)
|
||||
|
||||
def _build_system_prompt(self, system_message):
|
||||
self.built_prompts.append(system_message)
|
||||
return "built prompt"
|
||||
|
||||
def event_callback(self, name, payload):
|
||||
self.events.append((name, payload))
|
||||
|
||||
|
||||
def test_codex_app_server_native_auto_mode_leaves_thread_compaction_to_codex():
|
||||
agent = DummyAgent(
|
||||
TurnResult(thread_id="thread-1", turn_id="compact-turn-1")
|
||||
)
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
|
||||
returned, prompt = compress_context(
|
||||
agent,
|
||||
messages,
|
||||
"system",
|
||||
approx_tokens=100000,
|
||||
task_id="test",
|
||||
)
|
||||
|
||||
assert returned is messages
|
||||
assert prompt == "cached prompt"
|
||||
assert agent._codex_session.calls == 0
|
||||
assert agent.context_compressor.compression_count == 0
|
||||
assert agent.events == []
|
||||
|
||||
|
||||
def test_codex_app_server_compaction_heartbeat_refreshes_activity_while_waiting():
|
||||
agent = DummyAgent(
|
||||
TurnResult(thread_id="thread-1", turn_id="compact-turn-1")
|
||||
)
|
||||
agent._codex_session = SlowCodexSession(
|
||||
agent._codex_session.result,
|
||||
agent.touch_calls,
|
||||
)
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
|
||||
returned, prompt = compress_context(
|
||||
agent,
|
||||
messages,
|
||||
"system",
|
||||
approx_tokens=100000,
|
||||
task_id="test",
|
||||
force=True,
|
||||
)
|
||||
|
||||
assert returned is messages
|
||||
assert prompt == "cached prompt"
|
||||
assert agent._codex_session.calls == 1
|
||||
assert "context compression started" in agent.touch_calls
|
||||
assert "context compression in progress" in agent.touch_calls
|
||||
assert agent.touch_calls[-1] == "context compression completed"
|
||||
from agent.session_activity import ActivityProvenance
|
||||
|
||||
assert agent.touch_provenances
|
||||
assert all(
|
||||
p is ActivityProvenance.AGENT_COMPRESSION for p in agent.touch_provenances
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_codex_app_server_compression_failure_preserves_bookkeeping():
|
||||
agent = DummyAgent(TurnResult(error="compact failed"))
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
|
||||
returned, prompt = compress_context(
|
||||
agent,
|
||||
messages,
|
||||
"system",
|
||||
approx_tokens=100000,
|
||||
force=True,
|
||||
)
|
||||
|
||||
assert returned is messages
|
||||
assert prompt == "cached prompt"
|
||||
assert agent._codex_session.calls == 1
|
||||
assert agent.context_compressor.compression_count == 0
|
||||
assert agent.context_compressor.last_prompt_tokens == 123
|
||||
assert agent.warnings
|
||||
assert agent.touch_calls[0] == "context compression started"
|
||||
assert agent.touch_calls[-1] == "context compression failed"
|
||||
assert agent.status_events == [
|
||||
("lifecycle", COMPACTION_STATUS),
|
||||
("warn", "⚠ Codex app-server compaction failed: compact failed"),
|
||||
]
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_codex_native_boundary_clears_stale_hermes_fallback_streak():
|
||||
from unittest.mock import patch
|
||||
|
||||
from agent.context_compressor import ContextCompressor
|
||||
|
||||
with patch(
|
||||
"agent.context_compressor.get_model_context_length",
|
||||
return_value=100_000,
|
||||
):
|
||||
compressor = ContextCompressor(model="test-model", quiet_mode=True)
|
||||
compressor._fallback_compression_streak = 1
|
||||
compressor._last_summary_fallback_used = True
|
||||
|
||||
agent = DummyAgent(
|
||||
TurnResult(thread_id="thread-1", turn_id="normal-turn-1")
|
||||
)
|
||||
agent.context_compressor = compressor
|
||||
turn = TurnResult(
|
||||
thread_id="thread-1",
|
||||
turn_id="normal-turn-1",
|
||||
compacted=True,
|
||||
)
|
||||
|
||||
assert _record_codex_app_server_compaction(agent, turn) is True
|
||||
assert compressor._fallback_compression_streak == 0
|
||||
assert compressor._verify_compaction_cleared_threshold is True
|
||||
|
||||
|
||||
class RecordingCooldownCompressor(SimpleNamespace):
|
||||
"""Compressor stub exposing the real cooldown API surface."""
|
||||
|
||||
def __init__(self, remaining=0.0):
|
||||
super().__init__(
|
||||
compression_count=0,
|
||||
last_compression_rough_tokens=0,
|
||||
last_prompt_tokens=123,
|
||||
last_completion_tokens=45,
|
||||
awaiting_real_usage_after_compression=False,
|
||||
)
|
||||
self.remaining = remaining
|
||||
self.recorded = []
|
||||
|
||||
def get_active_compression_failure_cooldown(self, *, refresh=False):
|
||||
if self.remaining <= 0:
|
||||
return None
|
||||
return {"remaining_seconds": self.remaining, "error": "prior failure"}
|
||||
|
||||
def _record_compression_failure_cooldown(self, seconds, error):
|
||||
self.recorded.append((seconds, error))
|
||||
self.remaining = float(seconds)
|
||||
|
||||
|
||||
def test_interrupted_codex_compaction_arms_the_failure_cooldown():
|
||||
"""Regression: the codex path returned unchanged with no brake, so the
|
||||
session stayed above threshold and the next turn retried immediately."""
|
||||
from agent.context_compressor import _SUMMARY_FAILURE_COOLDOWN_SECONDS
|
||||
|
||||
agent = DummyAgent(
|
||||
TurnResult(
|
||||
thread_id="thread-1",
|
||||
turn_id="compact-turn-1",
|
||||
interrupted=True,
|
||||
error="compact turn interrupted",
|
||||
),
|
||||
auto_compaction="hermes",
|
||||
)
|
||||
agent.context_compressor = RecordingCooldownCompressor()
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
|
||||
returned, prompt = compress_context(
|
||||
agent, messages, "system", approx_tokens=100000, task_id="test"
|
||||
)
|
||||
|
||||
assert returned is messages
|
||||
assert prompt == "cached prompt"
|
||||
assert agent.context_compressor.recorded == [
|
||||
(_SUMMARY_FAILURE_COOLDOWN_SECONDS, "compact turn interrupted")
|
||||
]
|
||||
|
||||
|
||||
def test_codex_compaction_error_without_interrupt_also_arms_cooldown():
|
||||
agent = DummyAgent(
|
||||
TurnResult(thread_id="thread-1", turn_id="compact-turn-1", error="boom"),
|
||||
auto_compaction="hermes",
|
||||
)
|
||||
agent.context_compressor = RecordingCooldownCompressor()
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
|
||||
compress_context(
|
||||
agent, messages, "system", approx_tokens=100000, task_id="test"
|
||||
)
|
||||
|
||||
assert len(agent.context_compressor.recorded) == 1
|
||||
assert agent.context_compressor.recorded[0][1] == "boom"
|
||||
|
||||
|
||||
def test_active_cooldown_blocks_automatic_codex_compaction():
|
||||
agent = DummyAgent(
|
||||
TurnResult(thread_id="thread-1", turn_id="compact-turn-1"),
|
||||
auto_compaction="hermes",
|
||||
)
|
||||
agent.context_compressor = RecordingCooldownCompressor(remaining=120.0)
|
||||
session = agent._codex_session
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
|
||||
returned, prompt = compress_context(
|
||||
agent, messages, "system", approx_tokens=100000, task_id="test"
|
||||
)
|
||||
|
||||
assert returned is messages
|
||||
assert prompt == "cached prompt"
|
||||
assert session.calls == 0, "compaction ran despite an active cooldown"
|
||||
|
||||
|
||||
def test_force_bypasses_the_codex_compaction_cooldown():
|
||||
"""An explicit /compress is a user decision and must not be braked by a
|
||||
failure it did not cause."""
|
||||
agent = DummyAgent(TurnResult(thread_id="thread-1", turn_id="compact-turn-1"))
|
||||
agent.context_compressor = RecordingCooldownCompressor(remaining=120.0)
|
||||
session = agent._codex_session
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
|
||||
compress_context(
|
||||
agent, messages, "system", approx_tokens=100000, task_id="test", force=True
|
||||
)
|
||||
|
||||
assert session.calls == 1
|
||||
|
||||
|
||||
def test_successful_codex_compaction_arms_no_cooldown():
|
||||
agent = DummyAgent(
|
||||
TurnResult(thread_id="thread-1", turn_id="compact-turn-1"),
|
||||
auto_compaction="hermes",
|
||||
)
|
||||
agent.context_compressor = RecordingCooldownCompressor()
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
|
||||
compress_context(
|
||||
agent, messages, "system", approx_tokens=100000, task_id="test"
|
||||
)
|
||||
|
||||
assert agent.context_compressor.recorded == []
|
||||
@@ -0,0 +1,788 @@
|
||||
"""Integration test for the codex_app_server runtime path through AIAgent.
|
||||
|
||||
Verifies that:
|
||||
- api_mode='codex_app_server' is accepted on AIAgent construction
|
||||
- run_conversation() takes the early-return path and never enters the
|
||||
chat completions loop
|
||||
- Projected messages from a fake Codex session land in the messages list
|
||||
- tool_iterations from the codex session tick the skill nudge counter
|
||||
- Memory nudge counter ticks once per turn
|
||||
- The returned dict has the same shape as the chat_completions path
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import run_agent
|
||||
from agent.transports.codex_app_server_session import CodexAppServerSession, TurnResult
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fake_session(monkeypatch):
|
||||
"""Replace CodexAppServerSession with a stub that returns a fixed
|
||||
TurnResult, so we can drive AIAgent without spawning real codex."""
|
||||
|
||||
def fake_run_turn(self, user_input: str, **kwargs):
|
||||
return TurnResult(
|
||||
final_text=f"echo: {user_input}",
|
||||
projected_messages=[
|
||||
{"role": "assistant", "content": None,
|
||||
"tool_calls": [{"id": "exec_1", "type": "function",
|
||||
"function": {"name": "exec_command",
|
||||
"arguments": "{}"}}]},
|
||||
{"role": "tool", "tool_call_id": "exec_1", "content": "ok"},
|
||||
{"role": "assistant", "content": f"echo: {user_input}"},
|
||||
],
|
||||
tool_iterations=1,
|
||||
interrupted=False,
|
||||
error=None,
|
||||
turn_id="turn-stub-1",
|
||||
thread_id="thread-stub-1",
|
||||
)
|
||||
|
||||
monkeypatch.setattr(CodexAppServerSession, "run_turn", fake_run_turn)
|
||||
monkeypatch.setattr(
|
||||
CodexAppServerSession, "ensure_started", lambda self: "thread-stub-1"
|
||||
)
|
||||
|
||||
|
||||
def _make_codex_agent(**kwargs):
|
||||
"""Construct an AIAgent in codex_app_server mode without contacting any
|
||||
real provider. We pass api_mode explicitly so the constructor takes the
|
||||
fast path for direct credentials."""
|
||||
return run_agent.AIAgent(
|
||||
api_key="stub",
|
||||
base_url="https://stub.invalid",
|
||||
provider="openai",
|
||||
api_mode="codex_app_server",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
class TestApiModeAccepted:
|
||||
def test_api_mode_is_codex_app_server(self):
|
||||
agent = _make_codex_agent()
|
||||
assert agent.api_mode == "codex_app_server"
|
||||
|
||||
|
||||
class TestRunConversationCodexPath:
|
||||
def test_run_conversation_returns_codex_shape(self, fake_session):
|
||||
agent = _make_codex_agent()
|
||||
# No background review fork during tests
|
||||
with patch.object(agent, "_spawn_background_review", return_value=None):
|
||||
result = agent.run_conversation("hello there")
|
||||
assert result["final_response"] == "echo: hello there"
|
||||
assert result["completed"] is True
|
||||
assert result["partial"] is False
|
||||
assert result["error"] is None
|
||||
assert result["api_calls"] == 1
|
||||
assert result["codex_thread_id"] == "thread-stub-1"
|
||||
assert result["codex_turn_id"] == "turn-stub-1"
|
||||
|
||||
def test_codex_app_server_token_usage_updates_session_accounting(self, monkeypatch):
|
||||
def fake_run_turn(self, user_input: str, **kwargs):
|
||||
return TurnResult(
|
||||
final_text="done",
|
||||
projected_messages=[{"role": "assistant", "content": "done"}],
|
||||
turn_id="turn-usage-1",
|
||||
thread_id="thread-usage-1",
|
||||
token_usage_last={
|
||||
"totalTokens": 130,
|
||||
"inputTokens": 80,
|
||||
"cachedInputTokens": 20,
|
||||
"outputTokens": 25,
|
||||
"reasoningOutputTokens": 5,
|
||||
},
|
||||
model_context_window=200000,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(CodexAppServerSession, "run_turn", fake_run_turn)
|
||||
monkeypatch.setattr(
|
||||
CodexAppServerSession, "ensure_started", lambda self: "thread-usage-1"
|
||||
)
|
||||
agent = _make_codex_agent()
|
||||
with patch.object(agent, "_spawn_background_review", return_value=None):
|
||||
result = agent.run_conversation("hello")
|
||||
|
||||
assert result["api_calls"] == 1
|
||||
assert result["prompt_tokens"] == 100
|
||||
assert result["completion_tokens"] == 25
|
||||
assert result["total_tokens"] == 130
|
||||
assert result["input_tokens"] == 80
|
||||
assert result["output_tokens"] == 25
|
||||
assert result["cache_read_tokens"] == 20
|
||||
assert result["cache_write_tokens"] == 0
|
||||
assert result["reasoning_tokens"] == 5
|
||||
assert result["last_prompt_tokens"] == 100
|
||||
|
||||
assert agent.session_api_calls == 1
|
||||
assert agent.session_prompt_tokens == 100
|
||||
assert agent.session_completion_tokens == 25
|
||||
assert agent.session_total_tokens == 130
|
||||
assert agent.session_input_tokens == 80
|
||||
assert agent.session_output_tokens == 25
|
||||
assert agent.session_cache_read_tokens == 20
|
||||
assert agent.session_cache_write_tokens == 0
|
||||
assert agent.session_reasoning_tokens == 5
|
||||
assert agent.context_compressor.last_prompt_tokens == 100
|
||||
assert agent.context_compressor.last_completion_tokens == 25
|
||||
assert agent.context_compressor.last_total_tokens == 130
|
||||
assert agent.context_compressor.context_length == 200000
|
||||
|
||||
def test_native_codex_compaction_updates_bookkeeping(self, monkeypatch):
|
||||
def fake_run_turn(self, user_input: str, **kwargs):
|
||||
return TurnResult(
|
||||
final_text="done",
|
||||
projected_messages=[{"role": "assistant", "content": "done"}],
|
||||
turn_id="turn-compact-1",
|
||||
thread_id="thread-compact-1",
|
||||
compacted=True,
|
||||
token_usage_last={
|
||||
"totalTokens": 300_000,
|
||||
"inputTokens": 300_000,
|
||||
"cachedInputTokens": 0,
|
||||
"outputTokens": 0,
|
||||
"reasoningOutputTokens": 0,
|
||||
},
|
||||
)
|
||||
|
||||
monkeypatch.setattr(CodexAppServerSession, "run_turn", fake_run_turn)
|
||||
monkeypatch.setattr(
|
||||
CodexAppServerSession, "ensure_started", lambda self: "thread-compact-1"
|
||||
)
|
||||
events = []
|
||||
agent = _make_codex_agent(event_callback=lambda name, payload: events.append((name, payload)))
|
||||
|
||||
with patch.object(agent, "_spawn_background_review", return_value=None):
|
||||
result = agent.run_conversation("hello")
|
||||
|
||||
assert result["completed"] is True
|
||||
assert agent.context_compressor.compression_count == 1
|
||||
# A compacted turn with real usage is judged against that same real
|
||||
# prompt count, exactly like a normal completed compression boundary.
|
||||
assert agent.context_compressor.last_prompt_tokens == 300_000
|
||||
assert agent.context_compressor.awaiting_real_usage_after_compression is False
|
||||
assert agent.context_compressor._ineffective_compression_count == 1
|
||||
assert events == [
|
||||
(
|
||||
"session:compress",
|
||||
{
|
||||
"platform": "",
|
||||
"session_id": agent.session_id,
|
||||
"old_session_id": "",
|
||||
"in_place": False,
|
||||
"compression_count": 1,
|
||||
"runtime": "codex_app_server",
|
||||
"thread_id": "thread-compact-1",
|
||||
"turn_id": "turn-compact-1",
|
||||
},
|
||||
)
|
||||
]
|
||||
|
||||
def test_projected_messages_are_spliced(self, fake_session):
|
||||
agent = _make_codex_agent()
|
||||
with patch.object(agent, "_spawn_background_review", return_value=None):
|
||||
result = agent.run_conversation("hello")
|
||||
msgs = result["messages"]
|
||||
# User message + 3 projected (assistant tool_call + tool + assistant text)
|
||||
assert len(msgs) >= 4
|
||||
assert msgs[0]["role"] == "user"
|
||||
assert msgs[0]["content"] == "hello"
|
||||
# Last assistant message has the final text
|
||||
final = [m for m in msgs if m.get("role") == "assistant"
|
||||
and m.get("content") == "echo: hello"]
|
||||
assert final, f"expected final assistant message in {msgs}"
|
||||
|
||||
def test_projected_messages_are_synced_to_external_memory(self, fake_session):
|
||||
agent = _make_codex_agent()
|
||||
agent._memory_manager = MagicMock()
|
||||
agent._memory_manager.build_system_prompt.return_value = ""
|
||||
|
||||
with patch.object(agent, "_spawn_background_review", return_value=None):
|
||||
result = agent.run_conversation("hello")
|
||||
|
||||
agent._memory_manager.sync_all.assert_called_once()
|
||||
assert agent._memory_manager.sync_all.call_args.kwargs["messages"] == result["messages"]
|
||||
|
||||
def test_nudge_counters_tick(self, fake_session):
|
||||
"""The skill nudge counter must accumulate tool_iterations across
|
||||
turns. The memory nudge counter is gated on memory being configured
|
||||
(which we skip via skip_memory=True), so we don't assert on it here —
|
||||
a separate test below covers that path explicitly."""
|
||||
agent = _make_codex_agent()
|
||||
agent._iters_since_skill = 0
|
||||
agent._user_turn_count = 0
|
||||
with patch.object(agent, "_spawn_background_review", return_value=None):
|
||||
agent.run_conversation("first")
|
||||
assert agent._iters_since_skill == 1 # one tool_iteration in fake turn
|
||||
# _user_turn_count is incremented by run_conversation pre-loop, not
|
||||
# by the codex helper — confirms we delegate that to the standard flow.
|
||||
assert agent._user_turn_count == 1
|
||||
with patch.object(agent, "_spawn_background_review", return_value=None):
|
||||
agent.run_conversation("second")
|
||||
assert agent._iters_since_skill == 2
|
||||
assert agent._user_turn_count == 2
|
||||
|
||||
def test_user_message_not_duplicated(self, fake_session):
|
||||
"""Regression guard: the user message must appear exactly once in
|
||||
the messages list. The standard run_conversation pre-loop appends
|
||||
it, and the codex helper must NOT append again."""
|
||||
agent = _make_codex_agent()
|
||||
with patch.object(agent, "_spawn_background_review", return_value=None):
|
||||
result = agent.run_conversation("ping unique 12345")
|
||||
user_count = sum(
|
||||
1 for m in result["messages"]
|
||||
if m.get("role") == "user" and m.get("content") == "ping unique 12345"
|
||||
)
|
||||
assert user_count == 1, f"user message appeared {user_count}× in {result['messages']}"
|
||||
|
||||
def test_background_review_NOT_invoked_below_threshold(self, fake_session):
|
||||
"""A single turn shouldn't trigger background review — counters
|
||||
haven't reached the nudge interval (default 10)."""
|
||||
agent = _make_codex_agent()
|
||||
agent._memory_nudge_interval = 10
|
||||
agent._skill_nudge_interval = 10
|
||||
agent._iters_since_skill = 0
|
||||
with patch.object(agent, "_spawn_background_review",
|
||||
return_value=None) as spawn:
|
||||
agent.run_conversation("ping")
|
||||
# Below threshold → review should NOT fire (was a real bug:
|
||||
# the helper was calling _spawn_background_review() with no
|
||||
# args after every turn, which would crash with TypeError).
|
||||
assert not spawn.called
|
||||
|
||||
def test_background_review_skill_trigger_fires_above_threshold(
|
||||
self, monkeypatch
|
||||
):
|
||||
"""When tool iterations cross the skill nudge interval, the
|
||||
background review fires with review_skills=True and the right
|
||||
messages_snapshot signature."""
|
||||
from agent.transports.codex_app_server_session import (
|
||||
CodexAppServerSession, TurnResult,
|
||||
)
|
||||
# Make the fake session report 10 tool iterations in one turn
|
||||
# (matching the default skill threshold).
|
||||
def fake_run_turn(self, user_input: str, **kwargs):
|
||||
return TurnResult(
|
||||
final_text=f"echo: {user_input}",
|
||||
projected_messages=[
|
||||
{"role": "assistant", "content": f"echo: {user_input}"},
|
||||
],
|
||||
tool_iterations=10,
|
||||
turn_id="t1", thread_id="th1",
|
||||
)
|
||||
monkeypatch.setattr(CodexAppServerSession, "run_turn", fake_run_turn)
|
||||
monkeypatch.setattr(
|
||||
CodexAppServerSession, "ensure_started", lambda self: "th1"
|
||||
)
|
||||
|
||||
agent = _make_codex_agent()
|
||||
agent._skill_nudge_interval = 10
|
||||
agent._iters_since_skill = 0
|
||||
# Make valid_tool_names include 'skill_manage' so the gate passes
|
||||
agent.valid_tool_names = set(getattr(agent, "valid_tool_names", set()))
|
||||
agent.valid_tool_names.add("skill_manage")
|
||||
|
||||
with patch.object(agent, "_spawn_background_review",
|
||||
return_value=None) as spawn:
|
||||
agent.run_conversation("do tool work")
|
||||
|
||||
assert spawn.called, "skill threshold tripped but review didn't fire"
|
||||
# Verify the call signature matches what _spawn_background_review
|
||||
# actually expects — this is the regression guard for the original
|
||||
# bug where the codex path called it with no args at all.
|
||||
call = spawn.call_args
|
||||
assert "messages_snapshot" in call.kwargs
|
||||
assert isinstance(call.kwargs["messages_snapshot"], list)
|
||||
assert call.kwargs["review_skills"] is True
|
||||
# Counter should be reset after the review fires
|
||||
assert agent._iters_since_skill == 0
|
||||
|
||||
def test_background_review_signature_never_breaks(self, fake_session):
|
||||
"""Even when no trigger fires, the helper must never call
|
||||
_spawn_background_review with the wrong signature. Run a turn,
|
||||
then run another turn after manually tripping the skill counter
|
||||
and confirm the call shape is the kwargs-only form the function
|
||||
actually accepts."""
|
||||
agent = _make_codex_agent()
|
||||
agent._skill_nudge_interval = 1 # very low so any iter trips it
|
||||
agent._iters_since_skill = 0
|
||||
agent.valid_tool_names = set(getattr(agent, "valid_tool_names", set()))
|
||||
agent.valid_tool_names.add("skill_manage")
|
||||
|
||||
with patch.object(agent, "_spawn_background_review",
|
||||
return_value=None) as spawn:
|
||||
agent.run_conversation("first")
|
||||
# The fake session reports tool_iterations=1, which trips
|
||||
# _skill_nudge_interval=1. So review should fire.
|
||||
assert spawn.called
|
||||
# Critical invariant: positional args must be empty, all real
|
||||
# args must be kwargs (matching _spawn_background_review's
|
||||
# actual signature).
|
||||
call = spawn.call_args
|
||||
assert call.args == (), (
|
||||
f"expected no positional args, got {call.args!r} — "
|
||||
"would crash _spawn_background_review at runtime"
|
||||
)
|
||||
assert "messages_snapshot" in call.kwargs
|
||||
|
||||
def test_chat_completions_loop_is_not_entered(self, fake_session):
|
||||
"""The early-return must bypass the regular API call loop entirely.
|
||||
We confirm by patching the SDK call and asserting it's never invoked."""
|
||||
agent = _make_codex_agent()
|
||||
# The chat_completions loop calls self.client.chat.completions.create(...)
|
||||
# If our early-return works, that path is dead.
|
||||
with patch.object(agent, "client") as client_mock, patch.object(
|
||||
agent, "_spawn_background_review", return_value=None
|
||||
):
|
||||
agent.run_conversation("hi")
|
||||
assert not client_mock.chat.completions.create.called
|
||||
|
||||
def test_gateway_terminal_cwd_seeds_codex_thread_cwd(self, monkeypatch, tmp_path):
|
||||
"""Gateway sessions set TERMINAL_CWD without pinning agent.session_cwd.
|
||||
Codex app-server must still start in that configured workspace instead
|
||||
of falling back to the Hermes daemon process cwd."""
|
||||
from agent.transports.codex_app_server_session import (
|
||||
CodexAppServerSession, TurnResult,
|
||||
)
|
||||
|
||||
captured: dict[str, str] = {}
|
||||
|
||||
def fake_init(self, **kwargs):
|
||||
captured["cwd"] = kwargs["cwd"]
|
||||
self._thread_id = "thread-stub-1"
|
||||
|
||||
def fake_run_turn(self, user_input: str, **kwargs):
|
||||
return TurnResult(
|
||||
final_text="ok",
|
||||
projected_messages=[{"role": "assistant", "content": "ok"}],
|
||||
turn_id="turn-stub-1",
|
||||
thread_id="thread-stub-1",
|
||||
)
|
||||
|
||||
monkeypatch.setenv("TERMINAL_CWD", str(tmp_path))
|
||||
monkeypatch.setattr(CodexAppServerSession, "__init__", fake_init)
|
||||
monkeypatch.setattr(CodexAppServerSession, "run_turn", fake_run_turn)
|
||||
|
||||
agent = _make_codex_agent()
|
||||
assert agent.session_cwd is None
|
||||
with patch.object(agent, "_spawn_background_review", return_value=None):
|
||||
agent.run_conversation("hi")
|
||||
|
||||
assert captured["cwd"] == str(tmp_path)
|
||||
|
||||
def _capture_routing_agent(self, monkeypatch):
|
||||
"""Build a codex agent with a CodexAppServerSession stub that captures
|
||||
the request_routing passed at construction time, so we can assert how
|
||||
the gateway-context approval routing was resolved."""
|
||||
captured: dict = {}
|
||||
|
||||
def fake_init(self, **kwargs):
|
||||
captured.update(kwargs)
|
||||
self._thread_id = "thread-stub-1"
|
||||
|
||||
def fake_run_turn(self, user_input: str, **kwargs):
|
||||
return TurnResult(
|
||||
final_text="ok",
|
||||
projected_messages=[{"role": "assistant", "content": "ok"}],
|
||||
turn_id="turn-stub-1",
|
||||
thread_id="thread-stub-1",
|
||||
)
|
||||
|
||||
monkeypatch.setattr(CodexAppServerSession, "__init__", fake_init)
|
||||
monkeypatch.setattr(CodexAppServerSession, "run_turn", fake_run_turn)
|
||||
monkeypatch.setattr(
|
||||
CodexAppServerSession, "ensure_started", lambda self: "thread-stub-1"
|
||||
)
|
||||
return captured
|
||||
|
||||
def test_approvals_mode_off_auto_approves_codex_server_requests(
|
||||
self, monkeypatch
|
||||
):
|
||||
"""When the user disables Hermes approvals, codex app-server approval
|
||||
requests should not fail closed just because no interactive callback is
|
||||
wired (the typical gateway path). Codex's own sandbox permission
|
||||
profile remains the filesystem boundary."""
|
||||
captured = self._capture_routing_agent(monkeypatch)
|
||||
with patch(
|
||||
"hermes_cli.config.load_config_readonly",
|
||||
return_value={"approvals": {"mode": "off"}},
|
||||
):
|
||||
agent = _make_codex_agent()
|
||||
with patch.object(
|
||||
agent, "_spawn_background_review", return_value=None
|
||||
):
|
||||
agent.run_conversation("write something")
|
||||
routing = captured["request_routing"]
|
||||
assert routing.auto_approve_exec is True
|
||||
assert routing.auto_approve_apply_patch is True
|
||||
|
||||
def test_yaml_boolean_false_approval_mode_also_auto_approves(
|
||||
self, monkeypatch
|
||||
):
|
||||
"""YAML 1.1 parses unquoted `off` as False; match the normal approval
|
||||
subsystem's compatibility behavior for codex app-server routing too."""
|
||||
captured = self._capture_routing_agent(monkeypatch)
|
||||
with patch(
|
||||
"hermes_cli.config.load_config_readonly",
|
||||
return_value={"approvals": {"mode": False}},
|
||||
):
|
||||
agent = _make_codex_agent()
|
||||
with patch.object(
|
||||
agent, "_spawn_background_review", return_value=None
|
||||
):
|
||||
agent.run_conversation("write something")
|
||||
routing = captured["request_routing"]
|
||||
assert routing.auto_approve_exec is True
|
||||
assert routing.auto_approve_apply_patch is True
|
||||
|
||||
def test_manual_approvals_keep_codex_server_requests_fail_closed(
|
||||
self, monkeypatch
|
||||
):
|
||||
"""Default (manual) approvals must preserve the fail-closed behavior —
|
||||
this fix is a no-op for users who haven't opted out."""
|
||||
captured = self._capture_routing_agent(monkeypatch)
|
||||
with patch(
|
||||
"hermes_cli.config.load_config",
|
||||
return_value={"approvals": {"mode": "manual"}},
|
||||
):
|
||||
agent = _make_codex_agent()
|
||||
with patch.object(
|
||||
agent, "_spawn_background_review", return_value=None
|
||||
):
|
||||
agent.run_conversation("write something")
|
||||
routing = captured["request_routing"]
|
||||
assert routing.auto_approve_exec is False
|
||||
assert routing.auto_approve_apply_patch is False
|
||||
|
||||
def test_frozen_yolo_env_auto_approves_codex_server_requests(
|
||||
self, monkeypatch
|
||||
):
|
||||
"""--yolo / HERMES_YOLO_MODE (frozen into _YOLO_MODE_FROZEN at import
|
||||
time — a prompt-injection-safe process-scoped bypass) should flow
|
||||
through to codex app-server routing so gateway/cron contexts do not
|
||||
fail closed when the user launched with yolo mode."""
|
||||
import tools.approval as _approval
|
||||
|
||||
captured = self._capture_routing_agent(monkeypatch)
|
||||
monkeypatch.setattr(_approval, "_YOLO_MODE_FROZEN", True)
|
||||
with patch(
|
||||
"hermes_cli.config.load_config",
|
||||
return_value={"approvals": {"mode": "manual"}},
|
||||
):
|
||||
agent = _make_codex_agent()
|
||||
with patch.object(
|
||||
agent, "_spawn_background_review", return_value=None
|
||||
):
|
||||
agent.run_conversation("write something")
|
||||
routing = captured["request_routing"]
|
||||
assert routing.auto_approve_exec is True
|
||||
assert routing.auto_approve_apply_patch is True
|
||||
|
||||
def test_session_yolo_auto_approves_codex_server_requests(
|
||||
self, monkeypatch
|
||||
):
|
||||
"""The /yolo session toggle should be honored at Codex session creation
|
||||
time, independent of the startup-time approvals config."""
|
||||
captured = self._capture_routing_agent(monkeypatch)
|
||||
with patch(
|
||||
"hermes_cli.config.load_config",
|
||||
return_value={"approvals": {"mode": "manual"}},
|
||||
):
|
||||
agent = _make_codex_agent()
|
||||
with patch(
|
||||
"tools.approval.is_approval_bypass_active_for_session",
|
||||
return_value=True,
|
||||
), patch.object(
|
||||
agent, "_spawn_background_review", return_value=None
|
||||
):
|
||||
agent.run_conversation("write something")
|
||||
routing = captured["request_routing"]
|
||||
assert routing.auto_approve_exec is True
|
||||
assert routing.auto_approve_apply_patch is True
|
||||
|
||||
|
||||
class TestReviewForkApiModeDowngrade:
|
||||
"""When the parent agent runs on codex_app_server, the background
|
||||
review fork must downgrade to codex_responses — otherwise the fork
|
||||
can't dispatch agent-loop tools (memory, skill_manage) which is the
|
||||
whole point of the review."""
|
||||
|
||||
def test_codex_app_server_parent_downgrades_review_fork(self):
|
||||
"""Live test against the real _spawn_background_review code path:
|
||||
verify the review_agent gets api_mode=codex_responses when the
|
||||
parent is codex_app_server."""
|
||||
from unittest.mock import MagicMock, patch as _patch
|
||||
agent = _make_codex_agent()
|
||||
# Pretend memory + skills are configured so the review fork
|
||||
# reaches the AIAgent constructor.
|
||||
agent._memory_store = MagicMock()
|
||||
agent._memory_enabled = True
|
||||
agent._user_profile_enabled = True
|
||||
# Mock _current_main_runtime to return the parent's codex_app_server
|
||||
# state so we can confirm the helper detects + downgrades it.
|
||||
agent._current_main_runtime = lambda: {
|
||||
"api_mode": "codex_app_server",
|
||||
"base_url": "https://chatgpt.com/backend-api/codex",
|
||||
"api_key": "stub-token",
|
||||
}
|
||||
# Capture what AIAgent gets constructed with inside the helper.
|
||||
captured = {}
|
||||
|
||||
def _capture_init(self, **kwargs):
|
||||
captured.update(kwargs)
|
||||
# Set bare attributes the rest of the spawn function reads
|
||||
# so it can finish without exploding.
|
||||
self.api_mode = kwargs.get("api_mode")
|
||||
self.provider = kwargs.get("provider")
|
||||
self.model = kwargs.get("model")
|
||||
self._memory_write_origin = None
|
||||
self._memory_write_context = None
|
||||
self._memory_store = None
|
||||
self._memory_enabled = False
|
||||
self._user_profile_enabled = False
|
||||
self._memory_nudge_interval = 0
|
||||
self._skill_nudge_interval = 0
|
||||
self.suppress_status_output = False
|
||||
self._session_messages = []
|
||||
|
||||
def _no_op_run_conv(*a, **kw):
|
||||
return {"final_response": "", "messages": []}
|
||||
self.run_conversation = _no_op_run_conv
|
||||
|
||||
def _no_op_close(*a, **kw):
|
||||
return None
|
||||
self.close = _no_op_close
|
||||
|
||||
with _patch("run_agent.AIAgent.__init__", _capture_init):
|
||||
agent._spawn_background_review(
|
||||
messages_snapshot=[{"role": "user", "content": "x"}],
|
||||
review_memory=True,
|
||||
review_skills=False,
|
||||
)
|
||||
# Wait for the spawned thread to actually execute
|
||||
import time
|
||||
for _ in range(30):
|
||||
if "api_mode" in captured:
|
||||
break
|
||||
time.sleep(0.1)
|
||||
|
||||
assert captured.get("api_mode") == "codex_responses", (
|
||||
f"review fork should be downgraded to codex_responses when "
|
||||
f"parent is codex_app_server; got {captured.get('api_mode')!r}"
|
||||
)
|
||||
|
||||
|
||||
class TestErrorHandling:
|
||||
def test_session_exception_returns_partial_with_error(self, monkeypatch):
|
||||
def boom_run_turn(self, user_input, **kwargs):
|
||||
raise RuntimeError("subprocess died")
|
||||
|
||||
monkeypatch.setattr(CodexAppServerSession, "ensure_started",
|
||||
lambda self: "t1")
|
||||
monkeypatch.setattr(CodexAppServerSession, "run_turn", boom_run_turn)
|
||||
|
||||
agent = _make_codex_agent()
|
||||
with patch.object(agent, "_spawn_background_review", return_value=None):
|
||||
result = agent.run_conversation("hi")
|
||||
assert result["completed"] is False
|
||||
assert result["partial"] is True
|
||||
assert "subprocess died" in result["error"]
|
||||
assert "codex-runtime auto" in result["final_response"]
|
||||
|
||||
def test_interrupted_turn_marked_partial(self, monkeypatch):
|
||||
def interrupted_turn(self, user_input, **kwargs):
|
||||
return TurnResult(
|
||||
final_text="",
|
||||
projected_messages=[],
|
||||
tool_iterations=0,
|
||||
interrupted=True,
|
||||
error="user interrupted",
|
||||
turn_id="t",
|
||||
thread_id="th",
|
||||
)
|
||||
monkeypatch.setattr(CodexAppServerSession, "ensure_started",
|
||||
lambda self: "th")
|
||||
monkeypatch.setattr(CodexAppServerSession, "run_turn", interrupted_turn)
|
||||
|
||||
agent = _make_codex_agent()
|
||||
with patch.object(agent, "_spawn_background_review", return_value=None):
|
||||
result = agent.run_conversation("hi")
|
||||
assert result["completed"] is False
|
||||
assert result["partial"] is True
|
||||
assert result["error"] == "user interrupted"
|
||||
|
||||
|
||||
class TestSessionRetirementOnRunAgent:
|
||||
"""run_agent.py side: when run_turn returns should_retire=True, the
|
||||
AIAgent must close + null _codex_session so the next turn respawns."""
|
||||
|
||||
def test_should_retire_drops_session(self, monkeypatch):
|
||||
closes = {"count": 0}
|
||||
|
||||
def fake_run_turn(self, user_input, **kwargs):
|
||||
return TurnResult(
|
||||
final_text="",
|
||||
projected_messages=[],
|
||||
tool_iterations=0,
|
||||
interrupted=True,
|
||||
error="turn timed out after 600.0s",
|
||||
turn_id="tu1",
|
||||
thread_id="th1",
|
||||
should_retire=True,
|
||||
)
|
||||
|
||||
def fake_close(self):
|
||||
closes["count"] += 1
|
||||
|
||||
monkeypatch.setattr(CodexAppServerSession, "ensure_started",
|
||||
lambda self: "th1")
|
||||
monkeypatch.setattr(CodexAppServerSession, "run_turn", fake_run_turn)
|
||||
monkeypatch.setattr(CodexAppServerSession, "close", fake_close)
|
||||
|
||||
agent = _make_codex_agent()
|
||||
with patch.object(agent, "_spawn_background_review", return_value=None):
|
||||
result = agent.run_conversation("hi")
|
||||
|
||||
# The session was closed and cleared
|
||||
assert closes["count"] == 1
|
||||
assert getattr(agent, "_codex_session", "MISSING") is None
|
||||
# Partial result was still returned (caller still sees the error)
|
||||
assert result["partial"] is True
|
||||
assert result["error"] == "turn timed out after 600.0s"
|
||||
|
||||
def test_normal_turn_keeps_session(self, fake_session):
|
||||
"""fake_session fixture returns should_retire=False (default).
|
||||
The session must stay attached for the next turn to reuse."""
|
||||
agent = _make_codex_agent()
|
||||
with patch.object(agent, "_spawn_background_review", return_value=None):
|
||||
agent.run_conversation("hi")
|
||||
# Session was lazily created and still attached.
|
||||
assert getattr(agent, "_codex_session", None) is not None
|
||||
|
||||
def test_exception_path_also_drops_session(self, monkeypatch):
|
||||
"""Even if run_turn raises (not just sets should_retire), we must
|
||||
drop the session — a thrown exception is the strongest possible
|
||||
signal the process is dead."""
|
||||
closes = {"count": 0}
|
||||
|
||||
def boom_run_turn(self, user_input, **kwargs):
|
||||
raise RuntimeError("codex segfaulted")
|
||||
|
||||
def fake_close(self):
|
||||
closes["count"] += 1
|
||||
|
||||
monkeypatch.setattr(CodexAppServerSession, "ensure_started",
|
||||
lambda self: "th1")
|
||||
monkeypatch.setattr(CodexAppServerSession, "run_turn", boom_run_turn)
|
||||
monkeypatch.setattr(CodexAppServerSession, "close", fake_close)
|
||||
|
||||
agent = _make_codex_agent()
|
||||
with patch.object(agent, "_spawn_background_review", return_value=None):
|
||||
result = agent.run_conversation("hi")
|
||||
|
||||
assert closes["count"] == 1
|
||||
assert agent._codex_session is None
|
||||
assert result["completed"] is False
|
||||
assert "codex segfaulted" in result["error"]
|
||||
|
||||
|
||||
class TestCodexToolProgressBridge:
|
||||
"""#38835 / #33200: Codex app-server item notifications must surface as
|
||||
Hermes tool-progress so gateways show verbose breadcrumbs on this route.
|
||||
The original item/started-only mapper was superseded by the full event
|
||||
bridge (make_codex_app_server_event_bridge); these tests pin the same
|
||||
mapping contract against the bridge helpers."""
|
||||
|
||||
def test_mapper_command_execution(self):
|
||||
from agent.codex_runtime import (
|
||||
_codex_item_to_args,
|
||||
_codex_item_to_preview,
|
||||
_codex_item_to_tool_name,
|
||||
)
|
||||
item = {"type": "commandExecution", "command": "ls -la", "cwd": "/tmp"}
|
||||
assert _codex_item_to_tool_name(item) == "exec_command"
|
||||
assert _codex_item_to_preview(item) == "ls -la"
|
||||
assert _codex_item_to_args(item) == {"command": "ls -la", "cwd": "/tmp"}
|
||||
|
||||
def test_mapper_file_change(self):
|
||||
from agent.codex_runtime import (
|
||||
_codex_item_to_preview,
|
||||
_codex_item_to_tool_name,
|
||||
)
|
||||
item = {
|
||||
"type": "fileChange",
|
||||
"changes": [{"path": "a.py"}, {"path": "b.py"}],
|
||||
}
|
||||
assert _codex_item_to_tool_name(item) == "apply_patch"
|
||||
assert _codex_item_to_preview(item) == "a.py, b.py"
|
||||
|
||||
def test_mapper_mcp_and_dynamic_tool_calls(self):
|
||||
from agent.codex_runtime import (
|
||||
_codex_item_to_args,
|
||||
_codex_item_to_tool_name,
|
||||
)
|
||||
mcp = {"type": "mcpToolCall", "server": "fs", "tool": "read", "arguments": {"p": 1}}
|
||||
assert _codex_item_to_tool_name(mcp) == "mcp.fs.read"
|
||||
assert _codex_item_to_args(mcp) == {"p": 1}
|
||||
|
||||
dyn = {"type": "dynamicToolCall", "tool": "web_search", "arguments": {"q": "x"}}
|
||||
assert _codex_item_to_tool_name(dyn) == "web_search"
|
||||
|
||||
def test_bridge_ignores_non_tool_items_and_other_methods(self):
|
||||
from agent.codex_runtime import make_codex_app_server_event_bridge
|
||||
events = []
|
||||
agent = SimpleNamespace(
|
||||
tool_progress_callback=lambda *a, **kw: events.append(a),
|
||||
_fire_stream_delta=None,
|
||||
_fire_reasoning_delta=None,
|
||||
_emit_interim_assistant_message=None,
|
||||
)
|
||||
on_event = make_codex_app_server_event_bridge(agent)
|
||||
# agentMessage started items are not tool-shaped
|
||||
on_event({"method": "item/started", "params": {
|
||||
"item": {"type": "agentMessage", "text": "hi"}}})
|
||||
# malformed / empty notes
|
||||
on_event({"method": "item/completed", "params": {}})
|
||||
on_event({})
|
||||
assert events == []
|
||||
|
||||
def test_session_wired_with_on_event_that_fires_tool_progress(self, monkeypatch):
|
||||
"""The session is constructed with an on_event hook that, when fed an
|
||||
item/started note, calls the agent's tool_progress_callback."""
|
||||
captured_init = {}
|
||||
events = []
|
||||
|
||||
def fake_init(self, **kwargs):
|
||||
captured_init.update(kwargs)
|
||||
# minimal attrs so the rest of run_turn stubs work
|
||||
self._client = None
|
||||
|
||||
def fake_run_turn(self, user_input, **kwargs):
|
||||
# Exercise the wired on_event hook with a real item/started note.
|
||||
on_event = captured_init.get("on_event")
|
||||
if on_event:
|
||||
on_event({"method": "item/started", "params": {"item": {
|
||||
"type": "commandExecution", "command": "pytest", "cwd": "/repo"}}})
|
||||
return TurnResult(final_text="done", projected_messages=[
|
||||
{"role": "assistant", "content": "done"}], turn_id="t1", thread_id="th1")
|
||||
|
||||
monkeypatch.setattr(CodexAppServerSession, "__init__", fake_init)
|
||||
monkeypatch.setattr(CodexAppServerSession, "ensure_started", lambda self: "th1")
|
||||
monkeypatch.setattr(CodexAppServerSession, "run_turn", fake_run_turn)
|
||||
|
||||
agent = _make_codex_agent()
|
||||
agent.tool_progress_callback = lambda kind, name, preview, args: events.append(
|
||||
(kind, name, preview))
|
||||
with patch.object(agent, "_spawn_background_review", return_value=None):
|
||||
agent.run_conversation("run the tests")
|
||||
|
||||
assert "on_event" in captured_init and captured_init["on_event"] is not None
|
||||
assert ("tool.started", "exec_command", "pytest") in events
|
||||
@@ -0,0 +1,82 @@
|
||||
"""Codex app-server session lifecycle on hard agent teardown (#65260).
|
||||
|
||||
The Codex runtime drops ``agent._codex_session`` on turn crash and on
|
||||
retirement (agent/codex_runtime.py), but ``AIAgent.close()`` — the hard
|
||||
teardown for /new, /reset, and session expiry — had no owner for it, so
|
||||
the app-server child process survived until interpreter exit.
|
||||
"""
|
||||
|
||||
import threading
|
||||
|
||||
from run_agent import AIAgent
|
||||
|
||||
|
||||
class _FakeCodexSession:
|
||||
def __init__(self, raises: bool = False):
|
||||
self.close_calls = 0
|
||||
self._raises = raises
|
||||
|
||||
def close(self):
|
||||
self.close_calls += 1
|
||||
if self._raises:
|
||||
raise RuntimeError("app-server already dead")
|
||||
|
||||
|
||||
def _bare_agent(session_id: str) -> AIAgent:
|
||||
"""Minimal agent shell exercising close() without a real build."""
|
||||
agent = AIAgent.__new__(AIAgent)
|
||||
agent.session_id = session_id
|
||||
agent.client = None
|
||||
agent._active_children_lock = threading.Lock()
|
||||
agent._active_children = set()
|
||||
agent._end_session_on_close = False
|
||||
agent._session_messages = ["retained"]
|
||||
return agent
|
||||
|
||||
|
||||
def test_agent_close_releases_codex_app_server_session(monkeypatch):
|
||||
agent = _bare_agent("test-codex-lifecycle")
|
||||
codex_session = _FakeCodexSession()
|
||||
agent._codex_session = codex_session
|
||||
|
||||
monkeypatch.setattr("run_agent.cleanup_vm", lambda _task_id: None)
|
||||
monkeypatch.setattr("run_agent.cleanup_browser", lambda _task_id: None)
|
||||
|
||||
agent.close()
|
||||
agent.close()
|
||||
|
||||
# Idempotent: the second close must not re-close a released session.
|
||||
assert codex_session.close_calls == 1
|
||||
assert agent._codex_session is None
|
||||
assert agent._session_messages == []
|
||||
|
||||
|
||||
def test_close_clears_reference_even_when_session_close_raises(monkeypatch):
|
||||
"""A wedged app-server must not strand a stale session reference.
|
||||
|
||||
The attribute is cleared BEFORE close() precisely so a raising close
|
||||
can't leave a dead session attached to the agent.
|
||||
"""
|
||||
agent = _bare_agent("test-codex-lifecycle-raises")
|
||||
codex_session = _FakeCodexSession(raises=True)
|
||||
agent._codex_session = codex_session
|
||||
|
||||
monkeypatch.setattr("run_agent.cleanup_vm", lambda _task_id: None)
|
||||
monkeypatch.setattr("run_agent.cleanup_browser", lambda _task_id: None)
|
||||
|
||||
agent.close()
|
||||
|
||||
assert codex_session.close_calls == 1
|
||||
assert agent._codex_session is None
|
||||
|
||||
|
||||
def test_close_without_codex_session_is_a_noop(monkeypatch):
|
||||
"""Non-Codex sessions (the common case) must be unaffected."""
|
||||
agent = _bare_agent("test-no-codex")
|
||||
|
||||
monkeypatch.setattr("run_agent.cleanup_vm", lambda _task_id: None)
|
||||
monkeypatch.setattr("run_agent.cleanup_browser", lambda _task_id: None)
|
||||
|
||||
agent.close()
|
||||
|
||||
assert getattr(agent, "_codex_session", None) is None
|
||||
@@ -71,13 +71,15 @@ class TestCodexAuxiliaryTimeoutFdOwnership:
|
||||
shutdown(); the real close() must land on the owning thread in the
|
||||
adapter's ``finally``."""
|
||||
|
||||
def _stalled():
|
||||
deadline = time.monotonic() + 30.0
|
||||
while time.monotonic() < deadline:
|
||||
time.sleep(0.02)
|
||||
yield SimpleNamespace(type="response.in_progress")
|
||||
def _one_keepalive_then_block():
|
||||
# Let the owner process one keepalive, then keep it inside the
|
||||
# stream past the watchdog window. The Timer is consequently
|
||||
# the only deadline observer that can win this timeout.
|
||||
yield SimpleNamespace(type="response.in_progress")
|
||||
time.sleep(1.0)
|
||||
yield SimpleNamespace(type="response.in_progress")
|
||||
|
||||
adapter, events = _adapter_with_recording_client(_stalled())
|
||||
adapter, events = _adapter_with_recording_client(_one_keepalive_then_block())
|
||||
owner_tid = threading.get_ident()
|
||||
|
||||
def _consume(stream, *, model, on_event):
|
||||
|
||||
@@ -0,0 +1,142 @@
|
||||
"""Tests for codex_responses_adapter multimodal tool-result handling.
|
||||
|
||||
Tool messages can contain a list of OpenAI-style content parts
|
||||
(``[{type:"text"...}, {type:"image_url"...}]``) when the
|
||||
``vision_analyze`` native fast path returns image bytes for the main model.
|
||||
This file verifies the Codex Responses adapter:
|
||||
|
||||
1. Converts that list into ``function_call_output.output`` as an array of
|
||||
``input_text``/``input_image`` items (not a stringified blob).
|
||||
2. Preserves array-shaped output through the preflight validator.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from agent.codex_responses_adapter import (
|
||||
_chat_messages_to_responses_input,
|
||||
_preflight_codex_input_items,
|
||||
)
|
||||
|
||||
|
||||
def _build_messages_with_multimodal_tool_result():
|
||||
return [
|
||||
{"role": "user", "content": "What's in /tmp/foo.png?"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [{
|
||||
"id": "call_abc",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "vision_analyze",
|
||||
"arguments": '{"image_url": "/tmp/foo.png", "question": "describe"}',
|
||||
},
|
||||
}],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"name": "vision_analyze",
|
||||
"tool_call_id": "call_abc",
|
||||
"content": [
|
||||
{"type": "text", "text": "Image loaded."},
|
||||
{"type": "image_url", "image_url": {"url": "data:image/png;base64,XYZ"}},
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
class TestMultimodalToolResultConversion:
|
||||
def test_list_content_becomes_output_array(self):
|
||||
items = _chat_messages_to_responses_input(
|
||||
_build_messages_with_multimodal_tool_result()
|
||||
)
|
||||
# Find the function_call_output item
|
||||
outputs = [it for it in items if it.get("type") == "function_call_output"]
|
||||
assert len(outputs) == 1
|
||||
out = outputs[0]
|
||||
assert out["call_id"] == "call_abc"
|
||||
# Output should be a LIST (array form), not a string
|
||||
assert isinstance(out["output"], list), \
|
||||
f"Expected array output for multimodal tool result, got {type(out['output']).__name__}: {out['output']!r}"
|
||||
types = [p.get("type") for p in out["output"]]
|
||||
assert "input_text" in types
|
||||
assert "input_image" in types
|
||||
|
||||
def test_input_image_preserves_data_url(self):
|
||||
items = _chat_messages_to_responses_input(
|
||||
_build_messages_with_multimodal_tool_result()
|
||||
)
|
||||
out = next(it for it in items if it.get("type") == "function_call_output")
|
||||
image_parts = [p for p in out["output"] if p.get("type") == "input_image"]
|
||||
assert len(image_parts) == 1
|
||||
assert image_parts[0]["image_url"] == "data:image/png;base64,XYZ"
|
||||
|
||||
def test_string_tool_content_still_string_output(self):
|
||||
msgs = [
|
||||
{"role": "user", "content": "hi"},
|
||||
{
|
||||
"role": "assistant", "content": "",
|
||||
"tool_calls": [{
|
||||
"id": "call_x", "type": "function",
|
||||
"function": {"name": "terminal", "arguments": "{}"},
|
||||
}],
|
||||
},
|
||||
{
|
||||
"role": "tool", "name": "terminal", "tool_call_id": "call_x",
|
||||
"content": "ls output here",
|
||||
},
|
||||
]
|
||||
items = _chat_messages_to_responses_input(msgs)
|
||||
out = next(it for it in items if it.get("type") == "function_call_output")
|
||||
assert isinstance(out["output"], str)
|
||||
assert out["output"] == "ls output here"
|
||||
|
||||
|
||||
class TestPreflightAcceptsArrayOutput:
|
||||
def test_preflight_passes_array_through(self):
|
||||
raw = [
|
||||
{
|
||||
"type": "function_call",
|
||||
"call_id": "call_abc",
|
||||
"name": "vision_analyze",
|
||||
"arguments": "{}",
|
||||
},
|
||||
{
|
||||
"type": "function_call_output",
|
||||
"call_id": "call_abc",
|
||||
"output": [
|
||||
{"type": "input_text", "text": "Image loaded."},
|
||||
{"type": "input_image", "image_url": "data:image/png;base64,ABC"},
|
||||
],
|
||||
},
|
||||
]
|
||||
normalized = _preflight_codex_input_items(raw)
|
||||
out = [it for it in normalized if it.get("type") == "function_call_output"][0]
|
||||
assert isinstance(out["output"], list)
|
||||
assert len(out["output"]) == 2
|
||||
assert out["output"][1]["type"] == "input_image"
|
||||
assert out["output"][1]["image_url"] == "data:image/png;base64,ABC"
|
||||
|
||||
def test_preflight_drops_unknown_part_types(self):
|
||||
raw = [
|
||||
{
|
||||
"type": "function_call",
|
||||
"call_id": "call_abc", "name": "vision_analyze", "arguments": "{}",
|
||||
},
|
||||
{
|
||||
"type": "function_call_output",
|
||||
"call_id": "call_abc",
|
||||
"output": [
|
||||
{"type": "input_text", "text": "ok"},
|
||||
{"type": "garbage", "data": "nope"}, # unknown — should be dropped
|
||||
{"type": "input_image", "image_url": "data:image/png;base64,ZZ"},
|
||||
],
|
||||
},
|
||||
]
|
||||
normalized = _preflight_codex_input_items(raw)
|
||||
out = [it for it in normalized if it.get("type") == "function_call_output"][0]
|
||||
# The "garbage" part is dropped; valid parts remain
|
||||
types = [p.get("type") for p in out["output"]]
|
||||
assert types == ["input_text", "input_image"]
|
||||
|
||||
|
||||
@@ -0,0 +1,132 @@
|
||||
"""Regression coverage for #32892.
|
||||
|
||||
The openai SDK's ``responses.stream()`` / ``responses.parse()`` eagerly
|
||||
call ``_make_tools(tools)``, which iterates ``tools`` *without* a None
|
||||
guard. Passing ``tools=None`` therefore raises::
|
||||
|
||||
TypeError: 'NoneType' object is not iterable
|
||||
|
||||
…before any HTTP request is issued. This trips the
|
||||
``openai-codex`` / ``gpt-5.5`` combo on ``chatgpt.com/backend-api/codex``
|
||||
whenever the user runs Hermes without external tools registered: the
|
||||
agent loop catches the TypeError, sees no HTTP status, classifies it as
|
||||
non-retryable, and aborts (#32892).
|
||||
|
||||
These tests pin the defence:
|
||||
:func:`agent.transports.codex.ResponsesApiTransport.build_kwargs` must
|
||||
never emit ``tools=None`` — only add the ``tools`` key when there are
|
||||
function tools to expose. When there are no tools, the entire ``tools``
|
||||
key (plus ``tool_choice`` and ``parallel_tool_calls`` which are
|
||||
meaningless without it) is omitted from the kwargs.
|
||||
|
||||
Note: #33042 separately removed the SDK's ``responses.stream()`` helper
|
||||
from our own Codex call paths, so the specific iteration crash inside
|
||||
``_make_tools`` is also structurally avoided in normal operation. This
|
||||
test class additionally pins the SDK's ``_make_tools(None)`` contract so
|
||||
we notice if upstream ever changes it.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
import types
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
# Stub optional deps the parent module imports at top level — keeps this
|
||||
# test file runnable in the same environment as the existing Codex tests.
|
||||
sys.modules.setdefault("fire", types.SimpleNamespace(Fire=lambda *a, **k: None))
|
||||
sys.modules.setdefault("firecrawl", types.SimpleNamespace(Firecrawl=object))
|
||||
sys.modules.setdefault("fal_client", types.SimpleNamespace())
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def transport():
|
||||
"""Fresh ``ResponsesApiTransport`` per test (it is stateless but
|
||||
the import has side-effects on a global transport registry)."""
|
||||
from agent.transports.codex import ResponsesApiTransport
|
||||
|
||||
return ResponsesApiTransport()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def codex_messages() -> List[Dict[str, Any]]:
|
||||
"""Minimal Codex-shaped chat history mirroring the #32892 reproducer:
|
||||
one system + one short user message, with no tool calls in history."""
|
||||
return [
|
||||
{"role": "system", "content": "You are Hermes."},
|
||||
{"role": "user", "content": "Hey! What can I help you with?"},
|
||||
]
|
||||
|
||||
|
||||
def _build_kwargs_no_tools(transport, messages) -> Dict[str, Any]:
|
||||
"""Exercise the real ``build_kwargs`` for the codex backend with no tools."""
|
||||
return transport.build_kwargs(
|
||||
model="gpt-5.5",
|
||||
messages=messages,
|
||||
tools=None,
|
||||
is_codex_backend=True,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# build_kwargs: the "tools=None" key must never appear
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_build_kwargs_omits_tools_key_when_no_tools(transport, codex_messages):
|
||||
"""``build_kwargs`` must not place ``tools=None`` in the outgoing dict.
|
||||
|
||||
Putting ``tools=None`` reaches ``responses.stream()`` which calls
|
||||
``_make_tools(None)`` and crashes with the #32892 TypeError before any
|
||||
request is sent.
|
||||
"""
|
||||
kwargs = _build_kwargs_no_tools(transport, codex_messages)
|
||||
|
||||
assert "tools" not in kwargs, (
|
||||
f"tools key must be omitted entirely when no tools are registered, "
|
||||
f"got kwargs={sorted(kwargs)}"
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
def test_build_kwargs_keeps_required_codex_fields_without_tools(transport, codex_messages):
|
||||
"""The toolless build must still emit the non-negotiable Codex fields
|
||||
(model / instructions / input / store) — otherwise we'd just be moving
|
||||
the bug from the SDK to preflight."""
|
||||
kwargs = _build_kwargs_no_tools(transport, codex_messages)
|
||||
|
||||
assert kwargs["model"] == "gpt-5.5"
|
||||
assert kwargs["instructions"] == "You are Hermes."
|
||||
assert kwargs["store"] is False
|
||||
assert isinstance(kwargs["input"], list)
|
||||
assert kwargs["input"] and kwargs["input"][0]["role"] == "user"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_openai_sdk_raises_typeerror_on_tools_none():
|
||||
"""Document the upstream behaviour the two defences guard against.
|
||||
|
||||
If the SDK ever fixes ``_make_tools(None)`` to return ``omit``
|
||||
gracefully, this test will start failing — at which point the agent
|
||||
defences become belt-only and this test should be flipped to an
|
||||
``xfail`` so we notice the upstream change.
|
||||
"""
|
||||
from openai.resources.responses.responses import _make_tools
|
||||
|
||||
with pytest.raises(TypeError, match="NoneType.*not iterable"):
|
||||
_make_tools(None)
|
||||
@@ -5,6 +5,7 @@ import pytest
|
||||
from agent.codex_responses_adapter import (
|
||||
_chat_content_to_responses_parts,
|
||||
_chat_messages_to_responses_input,
|
||||
_classify_responses_issuer,
|
||||
_sanitize_replayed_fn_name,
|
||||
_format_responses_error,
|
||||
_normalize_codex_response,
|
||||
@@ -310,6 +311,7 @@ def test_normalize_codex_response_treats_summary_only_reasoning_as_incomplete():
|
||||
|
||||
_OVERSIZED_ITEM_ID = "x" * 408
|
||||
_VALID_ITEM_ID = "msg_abc123"
|
||||
_FOREIGN_ITEM_ID = "123e4567-e89b-12d3-a456-426614174000"
|
||||
|
||||
|
||||
# The codex app-server overflows the Responses 64-char call_id limit for
|
||||
@@ -462,6 +464,62 @@ def test_chat_messages_to_responses_input_canonicalizes_fc_only_pair():
|
||||
assert len(call["call_id"]) <= 64
|
||||
|
||||
|
||||
def test_chat_messages_to_responses_input_uniquifies_call_id_reused_across_turns():
|
||||
"""A stored call_id (e.g. a short-lived id like "terminal:0") can recur
|
||||
on a later, unrelated turn. Replayed verbatim, both function_call items
|
||||
and both function_call_output items would carry the same call_id, and
|
||||
the Responses API rejects the whole request with 400 "Duplicate
|
||||
function_call_output" (#102629). Each occurrence must get a unique
|
||||
call_id, still correctly paired with its own output."""
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"call_id": "terminal:0",
|
||||
"function": {"name": "terminal", "arguments": '{"command":"first"}'},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "terminal:0",
|
||||
"content": "first result",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"call_id": "terminal:0",
|
||||
"function": {"name": "terminal", "arguments": '{"command":"second"}'},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "terminal:0",
|
||||
"content": "second result",
|
||||
},
|
||||
]
|
||||
|
||||
items = _chat_messages_to_responses_input(messages)
|
||||
|
||||
calls = [i for i in items if i.get("type") == "function_call"]
|
||||
outputs = [i for i in items if i.get("type") == "function_call_output"]
|
||||
assert len(calls) == 2
|
||||
assert len(outputs) == 2
|
||||
|
||||
call_ids = [c["call_id"] for c in calls]
|
||||
assert len(set(call_ids)) == 2, "duplicate call_ids would 400 the whole request"
|
||||
|
||||
assert calls[0]["call_id"] == outputs[0]["call_id"]
|
||||
assert calls[1]["call_id"] == outputs[1]["call_id"]
|
||||
assert outputs[0]["output"] == "first result"
|
||||
assert outputs[1]["output"] == "second result"
|
||||
|
||||
|
||||
def test_preflight_codex_input_items_sanitizes_replayed_fn_name():
|
||||
"""The preflight choke-point also coerces invalid replayed names
|
||||
(covers callers that build input items without the chat converter)."""
|
||||
@@ -522,6 +580,116 @@ def test_preflight_codex_input_items_drops_short_id_for_github_responses():
|
||||
assert items[0]["content"] == [{"type": "output_text", "text": "pong"}]
|
||||
|
||||
|
||||
def test_chat_messages_to_responses_input_drops_foreign_id_for_codex_backend():
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "pong",
|
||||
"codex_message_items": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": "pong"}],
|
||||
"id": _FOREIGN_ITEM_ID,
|
||||
"phase": "final_answer",
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
codex_items = _chat_messages_to_responses_input(
|
||||
messages, current_issuer_kind="codex_backend"
|
||||
)
|
||||
xai_items = _chat_messages_to_responses_input(
|
||||
messages, current_issuer_kind="xai_responses"
|
||||
)
|
||||
|
||||
codex_message = next(item for item in codex_items if item.get("type") == "message")
|
||||
xai_message = next(item for item in xai_items if item.get("type") == "message")
|
||||
assert "id" not in codex_message
|
||||
assert codex_message["phase"] == "final_answer"
|
||||
assert xai_message["id"] == _FOREIGN_ITEM_ID
|
||||
|
||||
|
||||
def _reasoning_history(item):
|
||||
return [
|
||||
{"role": "assistant", "content": "done", "codex_reasoning_items": [item]},
|
||||
{"role": "user", "content": "next"},
|
||||
]
|
||||
|
||||
|
||||
def test_reasoning_replay_requires_matching_issuer_model_on_same_endpoint():
|
||||
# Blobs are sealed to the minting model, not just the endpoint: same endpoint + other model must drop.
|
||||
issuer = "other:https://responses.example.com/v1"
|
||||
normalized, _ = _normalize_codex_response(
|
||||
SimpleNamespace(
|
||||
status="completed",
|
||||
output=[
|
||||
SimpleNamespace(type="reasoning", id="rs_a", encrypted_content="model-a-blob", summary=[]),
|
||||
SimpleNamespace(
|
||||
type="message", role="assistant", status="completed", id="msg_a",
|
||||
content=[SimpleNamespace(type="output_text", text="done")],
|
||||
),
|
||||
],
|
||||
),
|
||||
issuer_kind=issuer, issuer_model="gpt-5.6-sol",
|
||||
)
|
||||
captured = normalized.codex_reasoning_items[0]
|
||||
assert captured["_issuer_model"] == "gpt-5.6-sol"
|
||||
|
||||
same = _chat_messages_to_responses_input(
|
||||
_reasoning_history(captured), current_issuer_kind=issuer, current_issuer_model="gpt-5.6-sol"
|
||||
)
|
||||
other = _chat_messages_to_responses_input(
|
||||
_reasoning_history(captured), current_issuer_kind=issuer, current_issuer_model="gpt-5.7-sol"
|
||||
)
|
||||
replayed = [i for i in same if i.get("type") == "reasoning"]
|
||||
assert [i["encrypted_content"] for i in replayed] == ["model-a-blob"]
|
||||
assert "_issuer_model" not in replayed[0] and "_issuer_kind" not in replayed[0]
|
||||
assert not any(i.get("type") == "reasoning" for i in other)
|
||||
|
||||
|
||||
def test_legacy_endpoint_stamped_item_without_model_replays_on_same_issuer():
|
||||
# WHY: native compaction checkpoints and reasoning persisted before model stamping carry only the
|
||||
# endpoint stamp; dropping them would erase every existing session's context once after upgrade.
|
||||
issuer = "other:https://responses.example.com/v1"
|
||||
legacy = {"type": "reasoning", "encrypted_content": "legacy-blob", "_issuer_kind": issuer}
|
||||
items = _chat_messages_to_responses_input(
|
||||
_reasoning_history(legacy), current_issuer_kind=issuer, current_issuer_model="gpt-5.6-sol"
|
||||
)
|
||||
replayed = [i for i in items if i.get("type") == "reasoning"]
|
||||
assert [i["encrypted_content"] for i in replayed] == ["legacy-blob"]
|
||||
# A different endpoint stamp still drops.
|
||||
foreign = _chat_messages_to_responses_input(
|
||||
_reasoning_history(legacy), current_issuer_kind="codex_backend", current_issuer_model="gpt-5.6-sol"
|
||||
)
|
||||
assert not any(i.get("type") == "reasoning" for i in foreign)
|
||||
|
||||
|
||||
def test_issuer_kind_is_canonical_across_trailing_slash_and_host_case():
|
||||
# The openai SDK stores ``client.base_url`` with a trailing slash; the aux adapter and the main
|
||||
# transport must agree on one issuer kind or aux calls drop every main-minted blob.
|
||||
canonical = _classify_responses_issuer(base_url="https://h/v1")
|
||||
assert _classify_responses_issuer(base_url="https://h/v1/") == canonical
|
||||
assert _classify_responses_issuer(base_url=" HTTPS://H/v1 ") == canonical
|
||||
assert _classify_responses_issuer(base_url="https://other/v1") != canonical
|
||||
|
||||
|
||||
def test_legacy_raw_endpoint_stamp_replays_on_canonical_issuer():
|
||||
# WHY: items persisted before issuer canonicalisation carry the raw ``agent.base_url`` (trailing slash,
|
||||
# host case); they must still replay on the same endpoint instead of being dropped as foreign.
|
||||
legacy = {"type": "reasoning", "encrypted_content": "legacy-blob", "_issuer_kind": "other:https://H/v1/"}
|
||||
items = _chat_messages_to_responses_input(
|
||||
_reasoning_history(legacy), current_issuer_kind="other:https://h/v1", current_issuer_model="gpt-5.6-sol"
|
||||
)
|
||||
assert [i["encrypted_content"] for i in items if i.get("type") == "reasoning"] == ["legacy-blob"]
|
||||
foreign = _chat_messages_to_responses_input(
|
||||
_reasoning_history(legacy), current_issuer_kind="other:https://other/v1", current_issuer_model="gpt-5.6-sol"
|
||||
)
|
||||
assert not any(i.get("type") == "reasoning" for i in foreign)
|
||||
|
||||
|
||||
def test_preflight_codex_api_kwargs_drops_oversized_message_id_end_to_end():
|
||||
kwargs = _preflight_codex_api_kwargs(
|
||||
{
|
||||
|
||||
@@ -0,0 +1,157 @@
|
||||
"""Regression tests for the SDK request-transform bypass (#93650).
|
||||
|
||||
``responses.create`` re-walks the whole request body against the
|
||||
``ResponseCreateParams`` union graph client-side, holding the GIL. #93650
|
||||
documents that walk wedging for 12+ hours on a ~1.4 MB conversation and
|
||||
freezing the entire agent — no in-process watchdog can fire while the GIL
|
||||
is held, and no socket kill helps a pre-network hang. Bulk wire-format
|
||||
fields are therefore routed through ``extra_body``, which the SDK merges
|
||||
into the JSON body *after* the transform.
|
||||
"""
|
||||
|
||||
import sys
|
||||
import types
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
sys.modules.setdefault("fire", types.SimpleNamespace(Fire=lambda *a, **k: None))
|
||||
sys.modules.setdefault("firecrawl", types.SimpleNamespace(Firecrawl=object))
|
||||
sys.modules.setdefault("fal_client", types.SimpleNamespace())
|
||||
|
||||
from agent.sdk_transform_bypass import (
|
||||
_is_plain_json_data,
|
||||
bypass_sdk_request_transform as _bypass_sdk_request_transform,
|
||||
)
|
||||
|
||||
|
||||
def _wire_kwargs():
|
||||
return {
|
||||
"model": "gpt-5.6-sol",
|
||||
"instructions": "You are Hermes.",
|
||||
"input": [
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "Ping"}]},
|
||||
{"type": "function_call_output", "call_id": "c1", "output": "ok"},
|
||||
],
|
||||
"tools": [{"type": "function", "name": "terminal", "parameters": {}}],
|
||||
"store": False,
|
||||
"stream": True,
|
||||
"timeout": 1800.0,
|
||||
}
|
||||
|
||||
|
||||
class TestIsPlainJsonData:
|
||||
def test_accepts_nested_wire_payloads(self):
|
||||
assert _is_plain_json_data(_wire_kwargs()["input"])
|
||||
|
||||
def test_rejects_non_json_leaves(self):
|
||||
assert not _is_plain_json_data([{"role": "user", "content": object()}])
|
||||
|
||||
def test_rejects_non_string_dict_keys(self):
|
||||
assert not _is_plain_json_data({1: "a"})
|
||||
|
||||
def test_rejects_generators(self):
|
||||
assert not _is_plain_json_data((item for item in ()))
|
||||
|
||||
|
||||
class TestBypassSdkRequestTransform:
|
||||
def test_moves_bulk_fields_to_extra_body(self):
|
||||
kwargs = _wire_kwargs()
|
||||
original_input = kwargs["input"]
|
||||
|
||||
bypassed = _bypass_sdk_request_transform(kwargs)
|
||||
|
||||
assert "input" not in bypassed
|
||||
assert "tools" not in bypassed
|
||||
assert bypassed["extra_body"]["input"] is original_input
|
||||
assert bypassed["extra_body"]["tools"] == kwargs["tools"]
|
||||
# Scalar configuration stays on the typed path.
|
||||
assert bypassed["model"] == "gpt-5.6-sol"
|
||||
assert bypassed["stream"] is True
|
||||
assert bypassed["timeout"] == 1800.0
|
||||
# The caller's mapping is untouched.
|
||||
assert kwargs["input"] is original_input
|
||||
assert "extra_body" not in kwargs
|
||||
|
||||
def test_merges_with_existing_extra_body_and_keeps_caller_precedence(self):
|
||||
kwargs = _wire_kwargs()
|
||||
caller_extra = {"prompt_cache_retention": "24h", "input": "explicit-wins"}
|
||||
kwargs["extra_body"] = caller_extra
|
||||
|
||||
bypassed = _bypass_sdk_request_transform(kwargs)
|
||||
|
||||
# An explicit extra_body entry wins, exactly as the SDK's
|
||||
# post-transform merge would have resolved the collision.
|
||||
assert bypassed["extra_body"]["input"] == "explicit-wins"
|
||||
assert bypassed["extra_body"]["prompt_cache_retention"] == "24h"
|
||||
assert bypassed["extra_body"]["tools"] == kwargs["tools"]
|
||||
assert caller_extra == {
|
||||
"prompt_cache_retention": "24h",
|
||||
"input": "explicit-wins",
|
||||
}
|
||||
|
||||
def test_non_json_field_stays_on_typed_sdk_path(self):
|
||||
kwargs = _wire_kwargs()
|
||||
kwargs["input"] = [{"role": "user", "content": object()}]
|
||||
|
||||
bypassed = _bypass_sdk_request_transform(kwargs)
|
||||
|
||||
assert bypassed["input"] == kwargs["input"]
|
||||
assert bypassed["extra_body"] == {"tools": kwargs["tools"]}
|
||||
|
||||
def test_string_input_stays_in_place(self):
|
||||
kwargs = _wire_kwargs()
|
||||
kwargs["input"] = "plain prompt"
|
||||
kwargs.pop("tools")
|
||||
|
||||
bypassed = _bypass_sdk_request_transform(kwargs)
|
||||
|
||||
assert bypassed is kwargs
|
||||
|
||||
def test_env_escape_hatch_restores_passthrough(self, monkeypatch):
|
||||
monkeypatch.setenv("HERMES_CODEX_SDK_TRANSFORM", "1")
|
||||
kwargs = _wire_kwargs()
|
||||
|
||||
assert _bypass_sdk_request_transform(kwargs) is kwargs
|
||||
|
||||
|
||||
class TestRunCodexStreamRoutesPayloadViaExtraBody:
|
||||
def _make_agent(self):
|
||||
from run_agent import AIAgent
|
||||
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://chatgpt.com/backend-api/codex",
|
||||
model="gpt-5.6-sol",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
agent._interrupt_requested = False
|
||||
return agent
|
||||
|
||||
def test_create_receives_input_via_extra_body(self):
|
||||
from agent.codex_runtime import run_codex_stream
|
||||
|
||||
agent = self._make_agent()
|
||||
events = [
|
||||
SimpleNamespace(
|
||||
type="response.completed",
|
||||
response=SimpleNamespace(
|
||||
id="r1",
|
||||
status="completed",
|
||||
output=[],
|
||||
usage=None,
|
||||
),
|
||||
)
|
||||
]
|
||||
mock_client = MagicMock()
|
||||
mock_client.responses.create.return_value = iter(events)
|
||||
|
||||
run_codex_stream(agent, _wire_kwargs(), client=mock_client)
|
||||
|
||||
create_kwargs = mock_client.responses.create.call_args.kwargs
|
||||
assert "input" not in create_kwargs
|
||||
assert "tools" not in create_kwargs
|
||||
assert create_kwargs["stream"] is True
|
||||
assert create_kwargs["extra_body"]["input"][0]["role"] == "user"
|
||||
assert create_kwargs["extra_body"]["tools"][0]["name"] == "terminal"
|
||||
@@ -0,0 +1,76 @@
|
||||
"""Tests for the ``_codex_silent_hang_hint`` heuristic.
|
||||
|
||||
The helper substitutes an actionable hint into the stale-call timeout
|
||||
warning when the request matches a known Codex silent-reject pattern
|
||||
(gpt-5.5 family on the ChatGPT Codex backend). See issue #21444 for
|
||||
symptom history. The recommended workaround for ChatGPT Codex OAuth
|
||||
accounts is `gpt-5.4` / `gpt-5.3-codex`, not `gpt-5.4-codex`.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _make_agent(tmp_path: Path, **overrides):
|
||||
from run_agent import AIAgent
|
||||
kwargs = dict(
|
||||
model="gpt-5.5",
|
||||
provider="openai-codex",
|
||||
api_key="sk-dummy",
|
||||
base_url="https://chatgpt.com/backend-api/codex",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
platform="cli",
|
||||
)
|
||||
kwargs.update(overrides)
|
||||
return AIAgent(**kwargs)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _isolate_hermes_home(monkeypatch, tmp_path):
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
(tmp_path / ".env").write_text("", encoding="utf-8")
|
||||
|
||||
|
||||
# ── positive cases: hint fires ─────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_hint_fires_for_bare_gpt_5_5_on_codex(tmp_path):
|
||||
agent = _make_agent(tmp_path)
|
||||
agent.api_mode = "codex_responses"
|
||||
hint = agent._codex_silent_hang_hint(model="gpt-5.5")
|
||||
assert hint is not None
|
||||
assert "gpt-5.4" in hint
|
||||
assert "gpt-5.3-codex" in hint
|
||||
assert "gpt-5.4-codex" in hint
|
||||
assert "fallback chain" in hint
|
||||
|
||||
|
||||
def test_hint_fires_for_vendor_prefixed_gpt_5_5(tmp_path):
|
||||
agent = _make_agent(tmp_path, model="openai/gpt-5.5")
|
||||
agent.api_mode = "codex_responses"
|
||||
hint = agent._codex_silent_hang_hint(model="openai/gpt-5.5")
|
||||
assert hint is not None
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ── negative cases: hint stays None ────────────────────────────────────────
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,700 @@
|
||||
"""Regression tests for the May 2026 xAI OAuth (SuperGrok / X Premium) bugs.
|
||||
|
||||
Three distinct failure modes the user community hit during rollout:
|
||||
|
||||
1. ``RuntimeError("Expected to have received `response.created` before
|
||||
`error`")`` on multi-turn xAI OAuth conversations. The OpenAI SDK's
|
||||
Responses streaming state machine collapses an upstream ``error`` SSE
|
||||
frame into a generic stream-ordering error. ``_run_codex_stream``
|
||||
now treats this the same way it already treats the missing
|
||||
``response.completed`` postlude — fall back to a non-stream
|
||||
``responses.create(stream=True)`` which surfaces the real provider
|
||||
error. Also closes #8133 (``response.in_progress`` prelude on custom
|
||||
relays) and #14634 (``codex.rate_limits`` prelude on codex-lb).
|
||||
|
||||
2. The HTTP 403 entitlement error xAI returns when an OAuth token lacks
|
||||
SuperGrok / X Premium ("You have either run out of available
|
||||
resources or do not have an active Grok subscription") used to read
|
||||
as a confusing wall of JSON. ``_summarize_api_error`` now appends a
|
||||
one-line hint pointing the user at https://grok.com and ``/model``.
|
||||
|
||||
3. Multi-turn replay of ``codex_reasoning_items`` (with
|
||||
``encrypted_content``) was briefly suppressed for ``is_xai_responses``
|
||||
in PR #26644 on the theory that xAI's OAuth/SuperGrok surface
|
||||
rejected replayed encrypted reasoning items. That suppression was
|
||||
reverted shortly after: xAI confirmed they explicitly want Hermes to
|
||||
thread encrypted reasoning back across turns, and the original
|
||||
multi-turn failure mode was actually the prelude-SSE issue closed by
|
||||
Fix A above. The remaining tests here lock in that xAI receives
|
||||
replayed reasoning AND that we ask xAI to echo it back in the
|
||||
``include`` array.
|
||||
"""
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fix A: prelude error surfacing via wire `error` events
|
||||
#
|
||||
# With the migration to ``responses.create(stream=True)`` raw event iteration,
|
||||
# the SDK's high-level state-machine RuntimeError no longer mediates between
|
||||
# the wire and us — we read the wire directly. When the chatgpt.com Codex
|
||||
# backend (or xAI, codex-lb, custom relays) emits a ``type=error`` frame as
|
||||
# its first event, our consumer raises ``_StreamErrorEvent`` straight from
|
||||
# the wire payload, which carries the real provider message in ``.body`` /
|
||||
# ``.message`` shape for ``_summarize_api_error`` to consume. This is
|
||||
# strictly better than the old "SDK raises RuntimeError → we retry → fall
|
||||
# back to a second non-stream call" two-phase dance, because the error
|
||||
# surfaces on the first event instead of after one wasted round trip.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_codex_agent():
|
||||
"""Build a minimal AIAgent wired for codex_responses streaming tests."""
|
||||
from run_agent import AIAgent
|
||||
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://api.x.ai/v1",
|
||||
model="grok-4.3",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
agent.api_mode = "codex_responses"
|
||||
agent.provider = "xai-oauth"
|
||||
agent._interrupt_requested = False
|
||||
return agent
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"provider_message",
|
||||
[
|
||||
"You do not have an active Grok subscription",
|
||||
"rate limit exceeded",
|
||||
"model not available",
|
||||
],
|
||||
)
|
||||
def test_codex_stream_wire_error_event_surfaces_stream_error_event(provider_message):
|
||||
"""A wire ``type=error`` SSE frame raises ``_StreamErrorEvent`` with the
|
||||
provider's real message in the body."""
|
||||
from run_agent import _StreamErrorEvent
|
||||
|
||||
agent = _make_codex_agent()
|
||||
|
||||
class _ErrorCreateStream:
|
||||
def __iter__(self_inner):
|
||||
yield SimpleNamespace(type="error", message=provider_message, code="forbidden")
|
||||
|
||||
def close(self_inner):
|
||||
pass
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.responses.create.return_value = _ErrorCreateStream()
|
||||
|
||||
with pytest.raises(_StreamErrorEvent) as excinfo:
|
||||
agent._run_codex_stream({}, client=mock_client)
|
||||
|
||||
assert provider_message in str(excinfo.value)
|
||||
assert excinfo.value.body["error"]["message"] == provider_message
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Nested error envelope on ``type=error`` SSE frames (opencode#36130 port)
|
||||
#
|
||||
# The Responses spec carries error details at the top level of the frame,
|
||||
# but the official OpenAI SDK and several OpenAI-compatible proxies wrap
|
||||
# them in an HTTP-style nested envelope:
|
||||
# {"type": "error", "error": {"code": ..., "message": ..., "param": ...}}
|
||||
# Before the fix, _raise_stream_error only read top-level fields, so these
|
||||
# frames collapsed to the generic "stream emitted error event" placeholder
|
||||
# and the error classifier never saw the provider's real code/message.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
|
||||
def test_codex_stream_wire_error_event_nested_envelope_attr_style():
|
||||
"""Details nested under ``error`` (SDK attr-object shape) are surfaced."""
|
||||
from run_agent import _StreamErrorEvent
|
||||
|
||||
agent = _make_codex_agent()
|
||||
|
||||
class _ErrorCreateStream:
|
||||
def __iter__(self_inner):
|
||||
yield SimpleNamespace(
|
||||
type="error",
|
||||
message=None,
|
||||
code=None,
|
||||
param=None,
|
||||
error=SimpleNamespace(
|
||||
type="rate_limit_error",
|
||||
code="rate_limit_exceeded",
|
||||
message="Slow down",
|
||||
param=None,
|
||||
),
|
||||
)
|
||||
|
||||
def close(self_inner):
|
||||
pass
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.responses.create.return_value = _ErrorCreateStream()
|
||||
|
||||
with pytest.raises(_StreamErrorEvent) as excinfo:
|
||||
agent._run_codex_stream({}, client=mock_client)
|
||||
|
||||
assert "Slow down" in str(excinfo.value)
|
||||
assert excinfo.value.code == "rate_limit_exceeded"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fix B: friendly entitlement message
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_summarize_api_error_decorates_xai_entitlement_403():
|
||||
"""xAI's OAuth 403 must surface the X Premium+ gotcha + neutral causes.
|
||||
|
||||
Wording deliberately leads with the X Premium+ gotcha because that's
|
||||
the #1 confusing case: people see Grok in their X app, assume it
|
||||
works here too, and hit this 403 with no idea API access is a
|
||||
separate SKU. Other causes (no subscription, wrong tier, exhausted
|
||||
quota) follow.
|
||||
"""
|
||||
from run_agent import AIAgent
|
||||
|
||||
error = RuntimeError(
|
||||
"HTTP 403: Error code: 403 - {'code': 'The caller does not have permission "
|
||||
"to execute the specified operation', 'error': 'You have either run out of "
|
||||
"available resources or do not have an active Grok subscription. Manage "
|
||||
"subscriptions at https://grok.com'}"
|
||||
)
|
||||
summary = AIAgent._summarize_api_error(error)
|
||||
# The original xAI text must survive — it's still useful diagnostic info.
|
||||
assert "do not have an active Grok subscription" in summary
|
||||
# The hint MUST lead with the X Premium+ gotcha (most likely cause
|
||||
# for users who think they're subscribed).
|
||||
assert "X Premium+ does NOT include" in summary
|
||||
assert "standalone SuperGrok subscribers" in summary
|
||||
# Other causes still listed.
|
||||
assert "no Grok subscription" in summary
|
||||
assert "tier doesn't include this model" in summary
|
||||
assert "quota is exhausted" in summary
|
||||
# The hint must point at the usage page where the user can verify.
|
||||
assert "https://grok.com/?_s=usage" in summary
|
||||
# Switching providers is still a valid escape hatch.
|
||||
assert "/model" in summary
|
||||
|
||||
|
||||
def test_summarize_api_error_does_not_accuse_subscribers():
|
||||
"""Hint must not confidently say the user has no subscription.
|
||||
|
||||
Don Piedro reported his subscription is active. The hint must not
|
||||
contradict him — leading with the X Premium+ gotcha gives subscribers
|
||||
a plausible reason ("oh, I'm on Premium+ not pure SuperGrok") instead
|
||||
of accusing them of lying about having a subscription.
|
||||
"""
|
||||
from run_agent import AIAgent
|
||||
|
||||
error = RuntimeError(
|
||||
"HTTP 403: do not have an active Grok subscription"
|
||||
)
|
||||
summary = AIAgent._summarize_api_error(error)
|
||||
# MUST NOT contain language that flatly assumes the user is unsubscribed.
|
||||
assert "lacks SuperGrok" not in summary
|
||||
assert "you are not subscribed" not in summary.lower()
|
||||
# MUST lead with the most-likely-but-non-accusatory cause.
|
||||
assert "X Premium+ does NOT include" in summary
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fix D: _StreamErrorEvent xAI entitlement classified as auth, not retryable
|
||||
#
|
||||
# run_codex_stream raises _StreamErrorEvent (status_code=None)
|
||||
# when the Responses stream emits a ``type=error`` SSE frame. Before this
|
||||
# fix, classify_api_error had no match for "grok subscription" in its pattern
|
||||
# lists, so it returned FailoverReason.unknown (retryable=True) — burning
|
||||
# max_retries before the agent stopped. _is_entitlement_failure was never
|
||||
# called because it only runs when FailoverReason.auth is returned.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_classify_api_error_stream_event_grok_subscription_is_auth():
|
||||
"""_StreamErrorEvent with xAI subscription message classifies as auth/non-retryable.
|
||||
|
||||
The SSE error path has status_code=None, so _classify_by_status is
|
||||
skipped. The explicit pattern added at step 1 must fire first and
|
||||
return auth/non-retryable so _is_entitlement_failure can stop the loop.
|
||||
"""
|
||||
from run_agent import _StreamErrorEvent
|
||||
from agent.error_classifier import classify_api_error, FailoverReason
|
||||
|
||||
err = _StreamErrorEvent(
|
||||
"You have either run out of available resources or do not have an "
|
||||
"active Grok subscription. Manage subscriptions at https://grok.com",
|
||||
code="The caller does not have permission to execute the specified operation",
|
||||
)
|
||||
result = classify_api_error(err, provider="xai-oauth", model="grok-4.3")
|
||||
assert result.reason == FailoverReason.auth
|
||||
assert result.retryable is False
|
||||
assert result.should_fallback is True
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fix C: reasoning replay gating for xai-oauth
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _assistant_msg_with_encrypted_reasoning(text="hi from grok", encrypted="enc_blob"):
|
||||
return {
|
||||
"role": "assistant",
|
||||
"content": text,
|
||||
"codex_reasoning_items": [
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": "rs_xai_001",
|
||||
"encrypted_content": encrypted,
|
||||
"summary": [],
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def test_codex_reasoning_replay_default_includes_encrypted_content():
|
||||
"""Native Codex backend (default) must still replay encrypted reasoning."""
|
||||
from agent.codex_responses_adapter import _chat_messages_to_responses_input
|
||||
|
||||
msgs = [
|
||||
{"role": "user", "content": "hi"},
|
||||
_assistant_msg_with_encrypted_reasoning(),
|
||||
{"role": "user", "content": "what's your name?"},
|
||||
]
|
||||
|
||||
items = _chat_messages_to_responses_input(msgs)
|
||||
reasoning = [it for it in items if it.get("type") == "reasoning"]
|
||||
assert len(reasoning) == 1
|
||||
assert reasoning[0]["encrypted_content"] == "enc_blob"
|
||||
|
||||
|
||||
def test_codex_reasoning_replay_includes_encrypted_content_for_xai():
|
||||
"""xAI must receive replayed encrypted reasoning items (May 2026 reversal).
|
||||
|
||||
Earlier we stripped these on the theory that the OAuth/SuperGrok
|
||||
surface rejected them. xAI subsequently confirmed they explicitly
|
||||
want Hermes to thread encrypted reasoning back across turns for
|
||||
cross-turn coherence — that's the whole point of the partnership
|
||||
integration.
|
||||
"""
|
||||
from agent.codex_responses_adapter import _chat_messages_to_responses_input
|
||||
|
||||
msgs = [
|
||||
{"role": "user", "content": "hi"},
|
||||
_assistant_msg_with_encrypted_reasoning(),
|
||||
{"role": "user", "content": "what's your name?"},
|
||||
]
|
||||
|
||||
items = _chat_messages_to_responses_input(msgs, is_xai_responses=True)
|
||||
reasoning = [it for it in items if it.get("type") == "reasoning"]
|
||||
assert len(reasoning) == 1, (
|
||||
"xAI must receive replayed reasoning items — see docstring for the "
|
||||
"May 2026 reversal of the earlier suppression gate."
|
||||
)
|
||||
assert reasoning[0]["encrypted_content"] == "enc_blob"
|
||||
|
||||
# And the assistant's visible text must still be present alongside it.
|
||||
assistant_items = [
|
||||
it for it in items
|
||||
if it.get("role") == "assistant" or it.get("type") == "message"
|
||||
]
|
||||
assert assistant_items, "assistant message must still be present"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fix D: entitlement 403 must NOT trigger credential-pool refresh loop
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"message",
|
||||
[
|
||||
# The exact wire text RaidenTyler and Don Piedro captured.
|
||||
"You have either run out of available resources or do not have an "
|
||||
"active Grok subscription. Manage at https://grok.com",
|
||||
# Permission-style variant from the same 403 body.
|
||||
"The caller does not have permission to execute the specified "
|
||||
"operation for grok-4.3",
|
||||
],
|
||||
)
|
||||
def test_is_entitlement_failure_matches_real_xai_bodies(message):
|
||||
from run_agent import AIAgent
|
||||
|
||||
assert AIAgent._is_entitlement_failure(
|
||||
{"message": message, "reason": "permission_denied"},
|
||||
403,
|
||||
)
|
||||
|
||||
|
||||
def test_is_entitlement_failure_false_for_status_other_than_401_403():
|
||||
"""200/429/500 must never be classified as entitlement, even if body matches."""
|
||||
from run_agent import AIAgent
|
||||
|
||||
body = {
|
||||
"message": "do not have an active Grok subscription",
|
||||
}
|
||||
assert not AIAgent._is_entitlement_failure(body, 500)
|
||||
assert not AIAgent._is_entitlement_failure(body, 429)
|
||||
assert not AIAgent._is_entitlement_failure(body, 200)
|
||||
|
||||
|
||||
|
||||
|
||||
def test_recover_with_credential_pool_skips_refresh_on_entitlement_403():
|
||||
"""The recovery path must NOT call pool.try_refresh_current() on entitlement 403.
|
||||
|
||||
Before the fix, an unsubscribed xAI OAuth account would burn the agent
|
||||
loop indefinitely: refresh → 403 → refresh → 403, infinitely. With
|
||||
the entitlement guard, recovery returns False so the error surfaces
|
||||
normally with the friendly hint from _summarize_api_error.
|
||||
"""
|
||||
from agent.error_classifier import FailoverReason
|
||||
|
||||
agent = _make_codex_agent()
|
||||
|
||||
# Wire a fake credential pool that records refresh attempts.
|
||||
refresh_calls = {"n": 0}
|
||||
|
||||
class _FakePool:
|
||||
def try_refresh_matching(self, api_key_hint=None):
|
||||
refresh_calls["n"] += 1
|
||||
return MagicMock(id="should_not_be_called")
|
||||
|
||||
def mark_exhausted_and_rotate(self, **_kwargs):
|
||||
return None
|
||||
|
||||
def has_available(self):
|
||||
return False
|
||||
|
||||
agent._credential_pool = _FakePool()
|
||||
|
||||
error_context = {
|
||||
"reason": "The caller does not have permission to execute the specified operation",
|
||||
"message": "You have either run out of available resources or do not have an "
|
||||
"active Grok subscription. Manage at https://grok.com",
|
||||
}
|
||||
|
||||
recovered, _retried_429 = agent._recover_with_credential_pool(
|
||||
status_code=403,
|
||||
has_retried_429=False,
|
||||
classified_reason=FailoverReason.auth,
|
||||
error_context=error_context,
|
||||
)
|
||||
|
||||
assert recovered is False, "Entitlement 403 must surface, not silently recover"
|
||||
assert refresh_calls["n"] == 0, "try_refresh_current must NOT be called on entitlement 403"
|
||||
|
||||
|
||||
def test_recover_with_credential_pool_rotates_on_xai_spending_limit_403():
|
||||
"""xAI's explicit spending-limit 403 must rotate, not hit the entitlement guard."""
|
||||
from agent.error_classifier import FailoverReason, classify_api_error
|
||||
|
||||
agent = _make_codex_agent()
|
||||
next_entry = MagicMock(id="healthy-account")
|
||||
refresh_calls = {"n": 0}
|
||||
|
||||
class _SpendingLimitError(Exception):
|
||||
status_code = 403
|
||||
body = {
|
||||
"code": "personal-team-blocked:spending-limit",
|
||||
"error": (
|
||||
"You have run out of credits or need a Grok subscription. "
|
||||
"Add credits at Grok or upgrade at Grok."
|
||||
),
|
||||
}
|
||||
|
||||
class _FakePool:
|
||||
provider = "xai-oauth"
|
||||
|
||||
def try_refresh_matching(self, api_key_hint=None):
|
||||
refresh_calls["n"] += 1
|
||||
return MagicMock(id="should_not_be_called")
|
||||
|
||||
def mark_exhausted_and_rotate(
|
||||
self,
|
||||
*,
|
||||
status_code,
|
||||
error_context=None,
|
||||
api_key_hint=None,
|
||||
failure_reason=None,
|
||||
):
|
||||
assert status_code == 403
|
||||
assert api_key_hint == "test-key"
|
||||
# An xAI spending-limit 403 classifies as billing, and the pool
|
||||
# must be told so — otherwise a sole-credential pool gives a spent
|
||||
# account the transient 60s cooldown instead of the full bench.
|
||||
assert failure_reason == "billing"
|
||||
assert error_context == {
|
||||
"reason": "personal-team-blocked:spending-limit",
|
||||
"message": (
|
||||
"You have run out of credits or need a Grok subscription. "
|
||||
"Add credits at Grok or upgrade at Grok."
|
||||
),
|
||||
}
|
||||
return next_entry
|
||||
|
||||
error = _SpendingLimitError("Error code: 403")
|
||||
classified = classify_api_error(error, provider="xai-oauth", model="grok-4.5")
|
||||
error_context = agent._extract_api_error_context(error)
|
||||
setattr(agent, "_credential_pool", _FakePool())
|
||||
agent._swap_credential = MagicMock()
|
||||
|
||||
recovered, retried_429 = agent._recover_with_credential_pool(
|
||||
status_code=error.status_code,
|
||||
has_retried_429=False,
|
||||
classified_reason=classified.reason,
|
||||
error_context=error_context,
|
||||
)
|
||||
|
||||
assert classified.reason == FailoverReason.billing
|
||||
assert recovered is True
|
||||
assert retried_429 is False
|
||||
assert refresh_calls["n"] == 0
|
||||
agent._swap_credential.assert_called_once_with(next_entry)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fix D-bis: bad-credentials 403 must NOT be classified as entitlement (#29344)
|
||||
#
|
||||
# xAI returns the same permission-denied ``code`` text for two distinct
|
||||
# conditions: unsubscribed account vs. stale OAuth access token. The
|
||||
# ``error`` field's ``[WKE=unauthenticated:...]`` suffix (and the
|
||||
# accompanying "OAuth2 access token could not be validated" phrasing) is
|
||||
# xAI's authoritative disambiguator — when present, the body is an auth
|
||||
# failure, not entitlement, and the credential-pool refresh path must
|
||||
# run. Pre-fix, long-running TUI sessions stuck on a stale token
|
||||
# surfaced as a non-retryable client error; the workaround was to exit
|
||||
# and reopen the TUI so the startup-resolve path refreshed.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fix E: grok-4.3 context length must be 1M, not 256K
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_grok_4_3_context_length_is_1m():
|
||||
"""grok-4.3 ships with 1M context per docs.x.ai/developers/models/grok-4.3.
|
||||
|
||||
Hermes' substring-match fallback used to return 256k (from the
|
||||
"grok-4" catch-all) which under-reported the model's real capacity.
|
||||
"""
|
||||
from agent.model_metadata import DEFAULT_CONTEXT_LENGTHS
|
||||
|
||||
# The entry exists with the expected value.
|
||||
assert DEFAULT_CONTEXT_LENGTHS["grok-4.3"] == 1_000_000
|
||||
|
||||
# And longest-first substring matching resolves grok-4.3 and
|
||||
# grok-4.3-latest to the new value, NOT the grok-4 catch-all.
|
||||
for slug in ("grok-4.3", "grok-4.3-latest"):
|
||||
matched_key = max(
|
||||
(k for k in DEFAULT_CONTEXT_LENGTHS if k in slug.lower()),
|
||||
key=len,
|
||||
)
|
||||
assert matched_key == "grok-4.3", (
|
||||
f"Expected longest-first match to land on grok-4.3 for {slug}, "
|
||||
f"got {matched_key}"
|
||||
)
|
||||
assert DEFAULT_CONTEXT_LENGTHS[matched_key] == 1_000_000
|
||||
|
||||
|
||||
def test_grok_4_still_resolves_to_256k():
|
||||
"""Regression guard: grok-4 (non-.3) must still resolve to 256k."""
|
||||
from agent.model_metadata import DEFAULT_CONTEXT_LENGTHS
|
||||
|
||||
for slug in ("grok-4", "grok-4-0709"):
|
||||
matched_key = max(
|
||||
(k for k in DEFAULT_CONTEXT_LENGTHS if k in slug.lower()),
|
||||
key=len,
|
||||
)
|
||||
# grok-4-0709 contains "grok-4" but not "grok-4.3"; matched key
|
||||
# must be "grok-4" (or a more specific variant family if one is
|
||||
# ever added). The 256k contract must hold.
|
||||
assert DEFAULT_CONTEXT_LENGTHS[matched_key] == 256_000
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cross-issuer reasoning replay guard
|
||||
#
|
||||
# When a session switches model providers mid-conversation (e.g. user runs
|
||||
# /model gpt-5.5 after several turns on grok-4.3), the persisted reasoning
|
||||
# items carry encrypted_content that only the issuing endpoint can decrypt.
|
||||
# Replaying them against the new endpoint deterministically returns HTTP 400
|
||||
# invalid_encrypted_content and breaks every subsequent turn. The cross-issuer
|
||||
# guard stamps each reasoning item with its issuer on normalize and drops
|
||||
# foreign-issuer items on replay.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _stamped_assistant_msg(issuer_kind, *, text="hi", encrypted="enc_blob", rs_id="rs_001"):
|
||||
return {
|
||||
"role": "assistant",
|
||||
"content": text,
|
||||
"codex_reasoning_items": [
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": rs_id,
|
||||
"encrypted_content": encrypted,
|
||||
"summary": [],
|
||||
"_issuer_kind": issuer_kind,
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_unstamped_reasoning_is_replayed_for_backwards_compat():
|
||||
"""Reasoning items persisted before this patch don't carry _issuer_kind.
|
||||
They must still be replayed (legacy-compatible behaviour).
|
||||
"""
|
||||
from agent.codex_responses_adapter import _chat_messages_to_responses_input
|
||||
|
||||
msgs = [
|
||||
{"role": "user", "content": "hi"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "hello",
|
||||
"codex_reasoning_items": [
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": "rs_legacy",
|
||||
"encrypted_content": "legacy_blob",
|
||||
"summary": [],
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "user", "content": "next"},
|
||||
]
|
||||
|
||||
items = _chat_messages_to_responses_input(
|
||||
msgs, current_issuer_kind="codex_backend"
|
||||
)
|
||||
reasoning = [it for it in items if it.get("type") == "reasoning"]
|
||||
assert len(reasoning) == 1
|
||||
assert reasoning[0]["encrypted_content"] == "legacy_blob"
|
||||
|
||||
|
||||
def test_normalize_codex_response_stamps_issuer_on_reasoning():
|
||||
"""Reasoning captured from a response must be stamped with the issuer so
|
||||
a later replay against a different endpoint can drop it.
|
||||
"""
|
||||
from types import SimpleNamespace
|
||||
|
||||
from agent.codex_responses_adapter import _normalize_codex_response
|
||||
|
||||
reasoning_item = SimpleNamespace(
|
||||
type="reasoning",
|
||||
id="rs_new",
|
||||
encrypted_content="fresh_blob",
|
||||
summary=[],
|
||||
)
|
||||
message_item = SimpleNamespace(
|
||||
type="message",
|
||||
role="assistant",
|
||||
status="completed",
|
||||
content=[SimpleNamespace(type="output_text", text="ok")],
|
||||
id="msg_1",
|
||||
)
|
||||
response = SimpleNamespace(output=[reasoning_item, message_item], status="completed")
|
||||
|
||||
msg, _ = _normalize_codex_response(response, issuer_kind="xai_responses")
|
||||
assert msg.codex_reasoning_items and len(msg.codex_reasoning_items) == 1
|
||||
assert msg.codex_reasoning_items[0]["_issuer_kind"] == "xai_responses"
|
||||
assert msg.codex_reasoning_items[0]["encrypted_content"] == "fresh_blob"
|
||||
|
||||
|
||||
def test_transport_round_trip_drops_foreign_reasoning():
|
||||
"""Full transport flow: build_kwargs against codex_backend after grok turns
|
||||
must produce an `input` array that contains zero foreign reasoning items.
|
||||
"""
|
||||
from agent.transports.codex import ResponsesApiTransport
|
||||
|
||||
transport = ResponsesApiTransport()
|
||||
messages = [
|
||||
{"role": "system", "content": "you are hermes"},
|
||||
{"role": "user", "content": "hi"},
|
||||
_stamped_assistant_msg("xai_responses", encrypted="grok_blob"),
|
||||
{"role": "user", "content": "엑스다임 프로젝트 파악, 스킬로 정리."},
|
||||
]
|
||||
|
||||
kwargs = transport.build_kwargs(
|
||||
model="gpt-5.5",
|
||||
messages=messages,
|
||||
tools=None,
|
||||
is_codex_backend=True,
|
||||
is_xai_responses=False,
|
||||
is_github_responses=False,
|
||||
base_url="https://chatgpt.com/backend-api/codex",
|
||||
instructions="you are hermes",
|
||||
)
|
||||
|
||||
reasoning = [it for it in kwargs["input"] if it.get("type") == "reasoning"]
|
||||
assert reasoning == [], (
|
||||
"Cross-issuer reasoning leaked through build_kwargs — this is the "
|
||||
"exact regression that broke session 40de1ae0 on 2026-05-25 01:09."
|
||||
)
|
||||
@@ -0,0 +1,223 @@
|
||||
"""E2E tests for the ``command`` secret source.
|
||||
|
||||
These exercise the REAL resolution path: real helper shell scripts written
|
||||
to a temp dir (chmod +x), real ``/bin/sh -c`` subprocesses, and a real temp
|
||||
HERMES_HOME with a config.yaml routing ``secrets.provider: command`` through
|
||||
``hermes_cli.env_loader._apply_external_secret_sources``.
|
||||
|
||||
Security invariants under test (ported from the desktop TS provider):
|
||||
|
||||
* the requested key travels ONLY via the ``HERMES_SECRET_KEY`` env var —
|
||||
never interpolated into the shell string (hostile key names are inert);
|
||||
* hard timeout + degrade-to-empty on every failure mode, never raise;
|
||||
* failure logging carries structured fields only — never the command
|
||||
string or any secret value.
|
||||
|
||||
NOTE: tests assert on key NAMES, lengths, and presence — never log secret
|
||||
values themselves.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import stat
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[2]
|
||||
if str(ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(ROOT))
|
||||
|
||||
from agent.secret_sources.command import ( # noqa: E402
|
||||
CommandSource,
|
||||
_run_helper,
|
||||
unquote_dotenv_value,
|
||||
)
|
||||
from agent.secret_sources.base import ( # noqa: E402
|
||||
reset_source_environment,
|
||||
set_source_environment,
|
||||
)
|
||||
from hermes_cli import env_loader # noqa: E402
|
||||
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
os.name == "nt", reason="the command secret provider is POSIX-only"
|
||||
)
|
||||
|
||||
|
||||
def _write_helper(tmp_path: Path, body: str, name: str = "helper.sh") -> Path:
|
||||
"""Write a real executable helper script and return its path."""
|
||||
script = tmp_path / name
|
||||
script.write_text("#!/bin/sh\n" + body + "\n", encoding="utf-8")
|
||||
script.chmod(script.stat().st_mode | stat.S_IXUSR | stat.S_IXGRP | stat.S_IXOTH)
|
||||
return script
|
||||
|
||||
|
||||
def test_profile_helper_does_not_inherit_process_secret(monkeypatch):
|
||||
monkeypatch.setenv("LEAK_CANARY", "global-secret")
|
||||
token = set_source_environment({"PROFILE_ONLY": "profile-value"})
|
||||
try:
|
||||
output = _run_helper(
|
||||
'printf "%s|%s" "${LEAK_CANARY-unset}" "$PROFILE_ONLY"',
|
||||
"",
|
||||
1.0,
|
||||
1024,
|
||||
)
|
||||
finally:
|
||||
reset_source_environment(token)
|
||||
|
||||
assert output == "unset|profile-value"
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clean_env(monkeypatch):
|
||||
"""Each test starts with a clean source map, applied-home guard, and no
|
||||
leftover test keys in os.environ."""
|
||||
env_loader._SECRET_SOURCES.clear()
|
||||
env_loader.reset_secret_source_cache()
|
||||
for key in ("CMDTEST_API_KEY", "CMDTEST_TOKEN", "CMDTEST_OTHER_KEY"):
|
||||
monkeypatch.delenv(key, raising=False)
|
||||
yield
|
||||
env_loader._SECRET_SOURCES.clear()
|
||||
env_loader.reset_secret_source_cache()
|
||||
for key in ("CMDTEST_API_KEY", "CMDTEST_TOKEN", "CMDTEST_OTHER_KEY"):
|
||||
os.environ.pop(key, None)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Parsing semantics (pure functions, mirroring the TS parseSecretOutput)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_unquote_strips_one_layer_of_matching_quotes():
|
||||
assert unquote_dotenv_value('"abc"') == "abc"
|
||||
assert unquote_dotenv_value("'abc'") == "abc"
|
||||
assert unquote_dotenv_value('""') == ""
|
||||
assert unquote_dotenv_value('"') == '"' # lone quote left intact
|
||||
assert unquote_dotenv_value(" plain ") == "plain"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Real-subprocess resolution
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_helper_stdout_is_returned_verbatim(tmp_path):
|
||||
helper = _write_helper(tmp_path, "printf 'sk-test-bare-12345'")
|
||||
value = _run_helper(str(helper), "CMDTEST_API_KEY", 3.0, 1024)
|
||||
assert value == "sk-test-bare-12345"
|
||||
|
||||
|
||||
def test_timeout_kills_hung_helper_and_degrades_to_empty(tmp_path):
|
||||
helper = _write_helper(tmp_path, "sleep 30")
|
||||
start = time.monotonic()
|
||||
value = _run_helper(str(helper), "CMDTEST_API_KEY", 2.0, 1024)
|
||||
elapsed = time.monotonic() - start
|
||||
assert value is None
|
||||
assert elapsed < 6.0, f"helper not killed within the bound (took {elapsed:.1f}s)"
|
||||
|
||||
|
||||
def test_failure_logging_never_leaks_command_or_secret(tmp_path, capfd):
|
||||
secret_value = "sk-super-secret-value-do-not-log"
|
||||
helper = _write_helper(
|
||||
tmp_path,
|
||||
f"echo '{secret_value}' >&2\nexit 7",
|
||||
name="my-distinctive-helper-name.sh",
|
||||
)
|
||||
value = _run_helper(str(helper), "CMDTEST_API_KEY", 3.0, 1024)
|
||||
assert value is None
|
||||
captured = capfd.readouterr()
|
||||
combined = captured.out + captured.err
|
||||
# Structured fields only — never the command string, the helper's
|
||||
# stderr, or any secret value.
|
||||
assert "my-distinctive-helper-name" not in combined
|
||||
assert secret_value not in combined
|
||||
assert "code=7" in combined # the structured field IS logged
|
||||
|
||||
|
||||
def test_fetch_parses_dotenv_blob(tmp_path):
|
||||
helper = _write_helper(
|
||||
tmp_path,
|
||||
"printf 'CMDTEST_API_KEY=sk-applied\\nCMDTEST_TOKEN=tok-applied\\n'",
|
||||
)
|
||||
result = CommandSource().fetch({"enabled": True, "command": str(helper)}, tmp_path)
|
||||
assert result.error is None
|
||||
assert result.secrets == {"CMDTEST_API_KEY": "sk-applied", "CMDTEST_TOKEN": "tok-applied"}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Dispatch E2E through env_loader against a real temp HERMES_HOME
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clean_registry():
|
||||
from agent.secret_sources import registry
|
||||
registry._reset_registry_for_tests()
|
||||
yield
|
||||
registry._reset_registry_for_tests()
|
||||
|
||||
|
||||
def test_registry_command_source_applies_and_records_source(tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
helper = _write_helper(
|
||||
tmp_path, "printf 'CMDTEST_API_KEY=sk-dispatch\\nCMDTEST_TOKEN=tok-dispatch\\n'"
|
||||
)
|
||||
(tmp_path / "config.yaml").write_text(
|
||||
"secrets:\n"
|
||||
" command:\n"
|
||||
" enabled: true\n"
|
||||
f" command: {helper}\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
env_loader._apply_external_secret_sources(tmp_path)
|
||||
|
||||
assert os.environ.get("CMDTEST_API_KEY") == "sk-dispatch"
|
||||
assert env_loader.get_secret_source("CMDTEST_API_KEY") == "command"
|
||||
assert env_loader.get_secret_source("CMDTEST_TOKEN") == "command"
|
||||
assert (
|
||||
env_loader.format_secret_source_suffix("CMDTEST_API_KEY")
|
||||
== " (from Command helper)"
|
||||
)
|
||||
|
||||
|
||||
def test_registry_status_line_printed_once_per_home(tmp_path, monkeypatch, capsys):
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
helper = _write_helper(tmp_path, "printf 'CMDTEST_API_KEY=sk-once\\n'")
|
||||
(tmp_path / "config.yaml").write_text(
|
||||
"secrets:\n command:\n enabled: true\n"
|
||||
f" command: {helper}\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
for _ in range(3): # idempotency guard: only the first call does work
|
||||
env_loader._apply_external_secret_sources(tmp_path)
|
||||
|
||||
err = capsys.readouterr().err
|
||||
assert err.count("Command helper: applied 1 secret") == 1
|
||||
|
||||
|
||||
def test_registry_failing_helper_does_not_block_startup(tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
(tmp_path / "config.yaml").write_text(
|
||||
"secrets:\n command:\n enabled: true\n command: exit 9\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
# Must not raise — config/helper errors never block startup.
|
||||
env_loader._apply_external_secret_sources(tmp_path)
|
||||
assert env_loader.get_secret_source("CMDTEST_API_KEY") is None
|
||||
|
||||
|
||||
def test_registry_helper_error_prints_remediation(tmp_path, monkeypatch, capsys):
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
(tmp_path / "config.yaml").write_text(
|
||||
"secrets:\n command:\n enabled: true\n command: ''\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
env_loader._apply_external_secret_sources(tmp_path)
|
||||
err = capsys.readouterr().err
|
||||
assert "secrets.command.command" in err
|
||||
@@ -0,0 +1,54 @@
|
||||
"""Regression tests for AIAgent.commit_memory_session.
|
||||
|
||||
Issue #22394: commit_memory_session was calling MemoryManager.on_session_end
|
||||
but never ContextEngine.on_session_end. Context engines that accumulate
|
||||
per-session state (LCM-style DAGs, summary stores) leaked that state from a
|
||||
rotated-out session into whatever continued under the same compressor
|
||||
instance.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
|
||||
def _make_minimal_agent(memory_manager, context_compressor, session_id="abc"):
|
||||
"""Build an object with just enough surface for commit_memory_session to run.
|
||||
|
||||
AIAgent.__init__ is too heavy for a focused unit test — bind the method
|
||||
to a SimpleNamespace-style object that has the attributes the method
|
||||
actually touches.
|
||||
"""
|
||||
from run_agent import AIAgent
|
||||
|
||||
obj = SimpleNamespace(
|
||||
_memory_manager=memory_manager,
|
||||
context_compressor=context_compressor,
|
||||
session_id=session_id,
|
||||
)
|
||||
obj.commit_memory_session = AIAgent.commit_memory_session.__get__(obj)
|
||||
return obj
|
||||
|
||||
|
||||
def test_commit_memory_session_notifies_context_engine():
|
||||
"""Both the memory manager AND the context engine receive on_session_end."""
|
||||
mm = MagicMock()
|
||||
ctx = MagicMock()
|
||||
agent = _make_minimal_agent(mm, ctx, session_id="sess-42")
|
||||
|
||||
msgs = [{"role": "user", "content": "hi"}, {"role": "assistant", "content": "yo"}]
|
||||
agent.commit_memory_session(msgs)
|
||||
|
||||
mm.on_session_end.assert_called_once_with(msgs)
|
||||
ctx.on_session_end.assert_called_once_with("sess-42", msgs)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,249 @@
|
||||
"""Compaction ALWAYS rebuilds the system prompt from the live builder (#95681).
|
||||
|
||||
The old keep-prompt containment branch restored the stored bytes whenever the
|
||||
reloaded memory blocks were embedded — so prompt-builder changes (guidance
|
||||
diets, renames, new blocks) never reached long-lived sessions (Bot Mode
|
||||
forever-chats, gateway channels). New contract:
|
||||
|
||||
1. builder output byte-equal -> keep the ORIGINAL string object (identity
|
||||
preserved for KV/prefix caches keyed on it)
|
||||
2. builder output differs -> the rebuilt prompt wins, logged
|
||||
3. plugin sections re-render at the same boundary; a RAISING plugin falls
|
||||
back to its last good bytes (fail-open), never silently vanishes
|
||||
"""
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from agent.system_prompt import invalidate_system_prompt
|
||||
|
||||
|
||||
def _agent(**over):
|
||||
base = dict(
|
||||
_cached_system_prompt="OLD PROMPT",
|
||||
_cached_system_prompt_static="OLD",
|
||||
_memory_store=None,
|
||||
_memory_manager=None,
|
||||
provider="",
|
||||
model="",
|
||||
platform="",
|
||||
_memory_enabled=False,
|
||||
_user_profile_enabled=False,
|
||||
)
|
||||
base.update(over)
|
||||
return SimpleNamespace(**base)
|
||||
|
||||
|
||||
class TestInvalidateClearsPluginFreeze(unittest.TestCase):
|
||||
def test_invalidate_stashes_and_clears_plugin_snapshot(self):
|
||||
agent = _agent()
|
||||
agent._plugin_system_prompt_sections_snapshot = ("frozen-section",)
|
||||
invalidate_system_prompt(agent)
|
||||
self.assertFalse(hasattr(agent, "_plugin_system_prompt_sections_snapshot"))
|
||||
self.assertEqual(agent._plugin_system_prompt_sections_previous, ("frozen-section",))
|
||||
self.assertIsNone(agent._cached_system_prompt)
|
||||
|
||||
def test_invalidate_without_snapshot_is_noop_for_plugins(self):
|
||||
agent = _agent()
|
||||
invalidate_system_prompt(agent)
|
||||
self.assertFalse(hasattr(agent, "_plugin_system_prompt_sections_snapshot"))
|
||||
|
||||
|
||||
class TestPluginRerenderFailOpen(unittest.TestCase):
|
||||
def test_raising_plugin_render_falls_back_to_previous_bytes(self):
|
||||
from agent.system_prompt import _frozen_plugin_prompt_sections
|
||||
|
||||
agent = _agent(_cached_system_prompt=None)
|
||||
agent._plugin_system_prompt_sections_previous = ("last-good",)
|
||||
with patch("hermes_cli.plugins.render_system_prompt_sections",
|
||||
side_effect=RuntimeError("plugin exploded")):
|
||||
rendered = _frozen_plugin_prompt_sections(agent)
|
||||
self.assertEqual(rendered, ("last-good",))
|
||||
|
||||
def test_raising_plugin_render_without_previous_is_empty(self):
|
||||
from agent.system_prompt import _frozen_plugin_prompt_sections
|
||||
|
||||
agent = _agent(_cached_system_prompt=None)
|
||||
with patch("hermes_cli.plugins.render_system_prompt_sections",
|
||||
side_effect=RuntimeError("plugin exploded")):
|
||||
rendered = _frozen_plugin_prompt_sections(agent)
|
||||
self.assertEqual(rendered, ())
|
||||
|
||||
|
||||
def _init_repo(path, first_commit):
|
||||
import subprocess
|
||||
path.mkdir()
|
||||
for cmd in (
|
||||
["git", "init", "-q", "-b", "main"],
|
||||
["git", "config", "user.email", "t@t"],
|
||||
["git", "config", "user.name", "t"],
|
||||
["git", "config", "core.autocrlf", "false"],
|
||||
):
|
||||
subprocess.run(cmd, cwd=path, check=True)
|
||||
(path / "main.py").write_text("print(1)\n")
|
||||
subprocess.run(["git", "add", "-A"], cwd=path, check=True)
|
||||
subprocess.run(["git", "commit", "-qm", first_commit], cwd=path, check=True)
|
||||
return path
|
||||
|
||||
|
||||
class TestCommitAlwaysRebuilds(unittest.TestCase):
|
||||
"""Source-level contract pins for the commit-site semantics."""
|
||||
|
||||
def _src(self):
|
||||
import inspect
|
||||
from agent import conversation_compression as cc
|
||||
return inspect.getsource(cc)
|
||||
|
||||
def test_keep_prompt_branch_requires_byte_equality(self):
|
||||
src = self._src()
|
||||
i = src.find("rebuilt_system_prompt = agent._build_system_prompt(")
|
||||
self.assertGreater(i, 0, "commit site must always run the live builder")
|
||||
window = src[i:i + 900]
|
||||
self.assertIn("rebuilt_system_prompt == cached_system_prompt", window,
|
||||
"keep-prompt must be gated on BYTE EQUALITY of the "
|
||||
"rebuilt output, not on memory containment")
|
||||
self.assertNotIn("_cached_prompt_reflects_builtin_memory(agent, cached_system_prompt)",
|
||||
window,
|
||||
"the containment keep-prompt gate must not return")
|
||||
|
||||
def test_drift_rebuild_is_logged(self):
|
||||
src = self._src()
|
||||
self.assertIn("Compaction rebuilt a drifted system prompt", src)
|
||||
|
||||
|
||||
class TestWorkspaceSnapshotPinnedAcrossCompaction(unittest.TestCase):
|
||||
"""Compaction rebuilds must not invalidate the prefix at the workspace snapshot (#103326)."""
|
||||
|
||||
def test_workspace_snapshot_replayed_across_rebuilds_when_repo_mutates(self):
|
||||
import tempfile, shutil, subprocess
|
||||
from pathlib import Path
|
||||
from agent.system_prompt import build_system_prompt, invalidate_system_prompt
|
||||
|
||||
tmp = Path(tempfile.mkdtemp(prefix="test-pinned-ws-"))
|
||||
try:
|
||||
repo = _init_repo(tmp / "proj", "init commit")
|
||||
|
||||
agent = _agent(
|
||||
load_soul_identity=False,
|
||||
skip_context_files=True,
|
||||
valid_tool_names={"terminal", "file_write"},
|
||||
platform="cli",
|
||||
model="gpt-4o",
|
||||
_memory_enabled=False,
|
||||
_user_profile_enabled=False,
|
||||
_task_completion_guidance=False,
|
||||
_parallel_tool_call_guidance=False,
|
||||
_tool_use_enforcement=False,
|
||||
_execution_guidance=False,
|
||||
_environment_probe=False,
|
||||
_bot_mode_protocol=False,
|
||||
_kanban_worker_guidance="",
|
||||
pass_session_id=False,
|
||||
session_id="s1",
|
||||
_emit_status=lambda *a, **k: None,
|
||||
)
|
||||
|
||||
with patch("agent.prompt_builder.load_soul_md", return_value=""), \
|
||||
patch("agent.prompt_builder.build_environment_hints", return_value="ENV HINTS"), \
|
||||
patch("agent.system_prompt.resolve_context_cwd", return_value=repo):
|
||||
|
||||
# First build: pins the snapshot
|
||||
p1 = build_system_prompt(agent)
|
||||
self.assertIn("Workspace (snapshot at session start", p1)
|
||||
|
||||
# Now repo mutates (agent touched and committed new files)
|
||||
(repo / "new_file.py").write_text("print(2)\n")
|
||||
subprocess.run(["git", "add", "-A"], cwd=repo, check=True)
|
||||
subprocess.run(["git", "commit", "-qm", "second commit"], cwd=repo, check=True)
|
||||
(repo / "untracked.txt").write_text("wip\n")
|
||||
|
||||
# Invalidate prompt (as happens during context compression)
|
||||
invalidate_system_prompt(agent)
|
||||
|
||||
# Second build: must replay pinned snapshot without re-probing git
|
||||
p2 = build_system_prompt(agent)
|
||||
self.assertEqual(p1, p2, "Prompt must remain byte-identical despite repo mutations")
|
||||
self.assertNotIn("second commit", p2, "Rebuilt prompt must not leak mutated git log")
|
||||
|
||||
finally:
|
||||
shutil.rmtree(tmp, ignore_errors=True)
|
||||
|
||||
def test_workspace_snapshot_reprobes_when_cwd_changes(self):
|
||||
import tempfile, shutil, subprocess
|
||||
from pathlib import Path
|
||||
from agent.system_prompt import build_system_prompt
|
||||
|
||||
tmp = Path(tempfile.mkdtemp(prefix="test-pinned-cwd-"))
|
||||
try:
|
||||
repo1 = _init_repo(tmp / "r1", "init r1")
|
||||
repo2 = _init_repo(tmp / "r2", "init r2")
|
||||
|
||||
agent = _agent(
|
||||
load_soul_identity=False,
|
||||
skip_context_files=True,
|
||||
valid_tool_names={"terminal"},
|
||||
platform="cli",
|
||||
model="gpt-4o",
|
||||
_memory_enabled=False,
|
||||
_user_profile_enabled=False,
|
||||
_task_completion_guidance=False,
|
||||
_parallel_tool_call_guidance=False,
|
||||
_tool_use_enforcement=False,
|
||||
_execution_guidance=False,
|
||||
_environment_probe=False,
|
||||
_bot_mode_protocol=False,
|
||||
_kanban_worker_guidance="",
|
||||
pass_session_id=False,
|
||||
session_id="s1",
|
||||
_emit_status=lambda *a, **k: None,
|
||||
)
|
||||
|
||||
with patch("agent.prompt_builder.load_soul_md", return_value=""), \
|
||||
patch("agent.prompt_builder.build_environment_hints", return_value="ENV HINTS"), \
|
||||
patch("agent.system_prompt.resolve_context_cwd", return_value=repo1):
|
||||
p1 = build_system_prompt(agent)
|
||||
self.assertIn(f"init {repo1.name}", p1)
|
||||
|
||||
with patch("agent.prompt_builder.load_soul_md", return_value=""), \
|
||||
patch("agent.prompt_builder.build_environment_hints", return_value="ENV HINTS"), \
|
||||
patch("agent.system_prompt.resolve_context_cwd", return_value=repo2):
|
||||
p2 = build_system_prompt(agent)
|
||||
self.assertIn(f"init {repo2.name}", p2)
|
||||
|
||||
finally:
|
||||
shutil.rmtree(tmp, ignore_errors=True)
|
||||
|
||||
def test_session_boundary_drops_the_pin_so_a_new_session_resnapshots(self):
|
||||
"""A /new, /resume or /branch reuses the AIAgent; the next session must see the live repo."""
|
||||
import tempfile, shutil, subprocess
|
||||
from pathlib import Path
|
||||
from agent.system_prompt import build_system_prompt, invalidate_system_prompt
|
||||
from run_agent import AIAgent
|
||||
|
||||
tmp = Path(tempfile.mkdtemp(prefix="test-pinned-boundary-"))
|
||||
try:
|
||||
repo = _init_repo(tmp / "proj", "init commit")
|
||||
agent = _agent(
|
||||
load_soul_identity=False, skip_context_files=True, valid_tool_names={"terminal"},
|
||||
platform="cli", model="gpt-4o", _task_completion_guidance=False,
|
||||
_parallel_tool_call_guidance=False, _tool_use_enforcement=False, _execution_guidance=False,
|
||||
_environment_probe=False, _bot_mode_protocol=False, _kanban_worker_guidance="",
|
||||
pass_session_id=False, session_id="s1", _emit_status=lambda *a, **k: None,
|
||||
_frozen_workspace_snapshot=None, context_compressor=None, _session_db=None,
|
||||
_transition_context_engine_session=lambda **kw: None,
|
||||
)
|
||||
with patch("agent.prompt_builder.load_soul_md", return_value=""), \
|
||||
patch("agent.prompt_builder.build_environment_hints", return_value="ENV HINTS"), \
|
||||
patch("agent.system_prompt.resolve_context_cwd", return_value=repo):
|
||||
build_system_prompt(agent)
|
||||
subprocess.run(["git", "commit", "-qm", "second commit", "--allow-empty"], cwd=repo, check=True)
|
||||
# The CLI session boundary (cli_session_mixin.new_session) on the same agent object.
|
||||
AIAgent.reset_session_state(agent)
|
||||
invalidate_system_prompt(agent)
|
||||
self.assertIn("second commit", build_system_prompt(agent))
|
||||
finally:
|
||||
shutil.rmtree(tmp, ignore_errors=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,102 @@
|
||||
"""Item 1 regression — the run_agent._compress_context fallback shim must be loud.
|
||||
|
||||
Before the fix, _compress_context wrapped the imports of _DB_PERSISTED_MARKER
|
||||
(agent.context_compressor) and the scoped-identity stamping helper
|
||||
(agent.conversation_compression._stamp_scoped_twins) in a try/except that silently defined local
|
||||
fallbacks (a hard-coded ``"_db_persisted"`` literal and a local copy of the
|
||||
identity helper) with NO logging. If the canonical constant/helper is renamed
|
||||
or removed upstream, the import raises, the fallback silently keeps stamping
|
||||
with the stale literal, and the stamping key splits from the flush's — the
|
||||
duplicate-row bug this PR fixes returns with no error anywhere.
|
||||
|
||||
The fix imports both symbols UNCONDITIONALLY (no fallback), so a
|
||||
renamed/removed symbol must fail the wrapper loudly with ImportError before
|
||||
any stamping happens.
|
||||
|
||||
The ``already_present`` outcome is load-bearing: it keeps compress_context
|
||||
from touching the deleted module-global name (which would raise NameError on
|
||||
BOTH pre- and post-fix code and make the test non-discriminating), because the
|
||||
stamp block at conversation_compression.py:3834-3859 — the ONLY in-module use
|
||||
of _stamp_scoped_twins — is skipped for already_present.
|
||||
"""
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import agent.conversation_compression as conversation_compression
|
||||
from agent.conversation_compression import CompressionCommitFence
|
||||
from hermes_state import SessionDB
|
||||
|
||||
|
||||
def _build_agent_with_db(db: SessionDB, session_id: str, platform: str = "telegram"):
|
||||
with patch.dict(os.environ, {"OPENROUTER_API_KEY": "test-key"}):
|
||||
from run_agent import AIAgent
|
||||
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
platform=platform,
|
||||
quiet_mode=True,
|
||||
session_db=db,
|
||||
session_id=session_id,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
|
||||
compressor = MagicMock()
|
||||
# A real user row in the stub return makes _ensure_compressed_has_user_turn
|
||||
# return `already_present`, so the in-module stamp block (the only user of
|
||||
# _stamp_scoped_twins inside compress_context) is skipped and
|
||||
# the deleted name is referenced ONLY by the run_agent shim import.
|
||||
compressor.compress.return_value = [
|
||||
{"role": "user", "content": "real user row"},
|
||||
]
|
||||
compressor.compression_count = 1
|
||||
compressor.last_prompt_tokens = 0
|
||||
compressor.last_completion_tokens = 0
|
||||
compressor._last_summary_error = None
|
||||
compressor._last_compress_aborted = False
|
||||
compressor._last_summary_auth_failure = False
|
||||
compressor._last_aux_model_failure_model = None
|
||||
compressor._last_aux_model_failure_error = None
|
||||
agent.context_compressor = compressor
|
||||
# ROTATION fallback path — pin in_place=False so the fork-rotation path is
|
||||
# exercised regardless of the global default (flipped to True in #38763).
|
||||
agent.compression_in_place = False
|
||||
return agent
|
||||
|
||||
|
||||
class TestCompressContextFallbackShim:
|
||||
def test_compress_context_shim_import_failure_is_loud(
|
||||
self, tmp_path: Path, monkeypatch
|
||||
):
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
parent = "PARENT_ROT_SHIM_LOUD"
|
||||
db.create_session(parent, source="cli")
|
||||
db.append_message(parent, "user", "persisted question")
|
||||
db.append_message(parent, "assistant", "persisted answer")
|
||||
|
||||
loaded = db.get_messages_as_conversation(parent)
|
||||
messages = [*loaded, {"role": "user", "content": "live question"}]
|
||||
|
||||
agent = _build_agent_with_db(db, parent)
|
||||
agent._persist_user_message_idx = len(messages) - 1
|
||||
|
||||
# Delete the canonical helper from its defining module: only the shim
|
||||
# import can still reference it (already_present skips the in-module
|
||||
# stamp block). Pre-fix the except branch silently defines a fallback;
|
||||
# post-fix the unconditional import must raise ImportError.
|
||||
monkeypatch.delattr(
|
||||
conversation_compression, "_stamp_scoped_twins"
|
||||
)
|
||||
with pytest.raises(ImportError):
|
||||
agent._compress_context(
|
||||
messages,
|
||||
"sys",
|
||||
approx_tokens=120_000,
|
||||
commit_fence=CompressionCommitFence(),
|
||||
)
|
||||
@@ -0,0 +1,74 @@
|
||||
"""Regression test: _compress_context tolerates plugin engines with strict signatures.
|
||||
|
||||
Added to ``ContextEngine.compress`` ABC signature (Apr 2026) allows passing
|
||||
``focus_topic`` to all engines. Older plugins written against the prior ABC
|
||||
(no focus_topic kwarg) would raise TypeError. _compress_context retries
|
||||
without focus_topic on TypeError so manual /compress <focus> doesn't crash
|
||||
on older plugins.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
|
||||
from run_agent import AIAgent
|
||||
|
||||
|
||||
def _make_agent_with_engine(engine):
|
||||
agent = object.__new__(AIAgent)
|
||||
agent.context_compressor = engine
|
||||
agent.session_id = "sess-1"
|
||||
agent.model = "test-model"
|
||||
agent.platform = "cli"
|
||||
agent.logs_dir = MagicMock()
|
||||
agent.quiet_mode = True
|
||||
agent._todo_store = MagicMock()
|
||||
agent._todo_store.format_for_injection.return_value = ""
|
||||
agent._memory_manager = None
|
||||
agent._session_db = None
|
||||
agent._cached_system_prompt = None
|
||||
agent.log_prefix = ""
|
||||
agent._vprint = lambda *a, **kw: None
|
||||
agent._last_flushed_db_idx = 0
|
||||
# Stub the few AIAgent methods _compress_context uses.
|
||||
agent._invalidate_system_prompt = lambda *a, **kw: None
|
||||
agent._build_system_prompt = lambda *a, **kw: "new-system-prompt"
|
||||
agent.commit_memory_session = lambda *a, **kw: None
|
||||
return agent
|
||||
|
||||
|
||||
def test_compress_context_falls_back_when_engine_rejects_focus_topic():
|
||||
"""Older plugins without focus_topic in compress() signature don't crash."""
|
||||
captured_kwargs = []
|
||||
|
||||
class _StrictOldPluginEngine:
|
||||
"""Mimics a plugin written against the pre-focus_topic ABC."""
|
||||
|
||||
compression_count = 0
|
||||
|
||||
def compress(self, messages, current_tokens=None):
|
||||
# NOTE: no focus_topic kwarg — TypeError if caller passes one.
|
||||
captured_kwargs.append({"current_tokens": current_tokens})
|
||||
return [messages[0], messages[-1]]
|
||||
|
||||
engine = _StrictOldPluginEngine()
|
||||
agent = _make_agent_with_engine(engine)
|
||||
|
||||
messages = [
|
||||
{"role": "user", "content": "one"},
|
||||
{"role": "assistant", "content": "two"},
|
||||
{"role": "user", "content": "three"},
|
||||
{"role": "assistant", "content": "four"},
|
||||
]
|
||||
|
||||
# Directly invoke the compression call site — this is the line that
|
||||
# used to blow up with TypeError under focus_topic+strict plugin.
|
||||
try:
|
||||
compressed = engine.compress(messages, current_tokens=100, focus_topic="foo")
|
||||
except TypeError:
|
||||
compressed = engine.compress(messages, current_tokens=100)
|
||||
|
||||
# Fallback succeeded: engine was called once without focus_topic.
|
||||
assert compressed == [messages[0], messages[-1]]
|
||||
assert captured_kwargs == [{"current_tokens": 100}]
|
||||
# Silence unused-var warning on agent.
|
||||
assert agent.context_compressor is engine
|
||||
@@ -0,0 +1,135 @@
|
||||
"""Regression tests for #58630: every compression abort path must reset
|
||||
per-attempt in-place compaction state.
|
||||
|
||||
After a successful in-place compaction sets ``_last_compaction_in_place=True``
|
||||
(run-level gateway signal), a later attempt that aborts or skips through ANY
|
||||
early-return path in ``compress_context`` must NOT reuse that stale flag as
|
||||
the flush baseline: ``conversation_history_after_compression()`` would then
|
||||
treat all current messages (including unflushed new turns) as persisted, and a
|
||||
restart would lose them.
|
||||
|
||||
The fix records a per-attempt outcome (``_last_compression_attempt_recorded``/
|
||||
``_last_compression_attempt_in_place``) at the very top of
|
||||
``compress_context`` — before the codex-app-server route, breaker gates, lock
|
||||
acquisition, rotated-parent skips, compressor-abort, no-progress and
|
||||
empty-transcript returns — so every abort path leaves the attempt outcome
|
||||
``None`` and callers retain the previous flush baseline.
|
||||
"""
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
|
||||
def _make_agent(session_db):
|
||||
with patch.dict(os.environ, {"OPENROUTER_API_KEY": "test-key"}):
|
||||
from run_agent import AIAgent
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
quiet_mode=True,
|
||||
session_db=session_db,
|
||||
session_id="abort-state-session",
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
return agent
|
||||
|
||||
|
||||
class _InPlaceSuccessCompressor:
|
||||
_last_compress_aborted = False
|
||||
_last_summary_error = None
|
||||
compression_count = 1
|
||||
_last_compression_made_progress = True
|
||||
_last_summary_fallback_used = False
|
||||
last_compression_rough_tokens = 0
|
||||
last_prompt_tokens = 0
|
||||
last_completion_tokens = 0
|
||||
awaiting_real_usage_after_compression = False
|
||||
|
||||
def compress(self, _messages, **_kwargs):
|
||||
return [
|
||||
{"role": "user", "content": "[summary] earlier state"},
|
||||
{"role": "assistant", "content": "retained tail"},
|
||||
]
|
||||
|
||||
|
||||
class _BreakerBlockedCompressor(_InPlaceSuccessCompressor):
|
||||
"""Trips the pre-lock automatic-compression breaker gate."""
|
||||
|
||||
def _automatic_compression_blocked(self):
|
||||
return True
|
||||
|
||||
def compress(self, messages, **_kwargs): # pragma: no cover - must not run
|
||||
raise AssertionError("compress() must not be reached when blocked")
|
||||
|
||||
|
||||
class _NoProgressCompressor(_InPlaceSuccessCompressor):
|
||||
"""Returns a semantically-equal transcript (no-op attempt)."""
|
||||
|
||||
_last_compression_made_progress = False
|
||||
|
||||
def compress(self, messages, **_kwargs):
|
||||
return [dict(m) for m in messages]
|
||||
|
||||
|
||||
class TestAbortPathsResetPerAttemptState:
|
||||
def _in_place_success(self, agent, messages):
|
||||
from agent.conversation_compression import (
|
||||
compress_context,
|
||||
conversation_history_after_compression,
|
||||
)
|
||||
agent.context_compressor = _InPlaceSuccessCompressor()
|
||||
compacted, _ = compress_context(
|
||||
agent, messages, "system", approx_tokens=100_000
|
||||
)
|
||||
assert agent._last_compaction_in_place is True
|
||||
assert agent._last_compression_attempt_in_place is True
|
||||
return compacted, conversation_history_after_compression(
|
||||
agent, compacted, None
|
||||
)
|
||||
|
||||
def test_breaker_blocked_skip_retains_previous_baseline(self):
|
||||
"""Pre-lock breaker skip after an in-place success must keep baseline."""
|
||||
from agent.conversation_compression import (
|
||||
compress_context,
|
||||
conversation_history_after_compression,
|
||||
)
|
||||
from hermes_state import SessionDB
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db = SessionDB(db_path=Path(tmpdir) / "test.db")
|
||||
agent = _make_agent(db)
|
||||
agent.compression_in_place = True
|
||||
original = [
|
||||
{"role": "user", "content": "old question"},
|
||||
{"role": "assistant", "content": "old answer"},
|
||||
]
|
||||
agent._flush_messages_to_session_db(original, [])
|
||||
compacted, history = self._in_place_success(agent, original)
|
||||
|
||||
messages = compacted + [
|
||||
{"role": "user", "content": "new request"},
|
||||
{"role": "assistant", "content": "new answer"},
|
||||
]
|
||||
agent.context_compressor = _BreakerBlockedCompressor()
|
||||
returned, _ = compress_context(
|
||||
agent, messages, "system", approx_tokens=100_000
|
||||
)
|
||||
assert returned is messages
|
||||
# Per-attempt outcome must be reset even though this attempt
|
||||
# returned before acquiring the lock.
|
||||
assert agent._last_compression_attempt_recorded is True
|
||||
assert agent._last_compression_attempt_in_place is None
|
||||
new_history = conversation_history_after_compression(
|
||||
agent, returned, history
|
||||
)
|
||||
# Skip = previous baseline stays authoritative: not all-persisted
|
||||
# (would drop the new pair on restart), not None (would re-append
|
||||
# the compacted rows).
|
||||
assert new_history is history
|
||||
db.close()
|
||||
|
||||
|
||||
@@ -0,0 +1,145 @@
|
||||
"""Tests for context compression boundary alignment.
|
||||
|
||||
Verifies that _align_boundary_backward correctly handles tool result groups
|
||||
so that parallel tool calls are never split during compression.
|
||||
"""
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
from agent.context_compressor import ContextCompressor
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _tc(call_id: str) -> dict:
|
||||
"""Create a minimal tool_call dict."""
|
||||
return {"id": call_id, "type": "function", "function": {"name": "test", "arguments": "{}"}}
|
||||
|
||||
|
||||
def _tool_result(call_id: str, content: str = "result") -> dict:
|
||||
"""Create a tool result message."""
|
||||
return {"role": "tool", "tool_call_id": call_id, "content": content}
|
||||
|
||||
|
||||
def _assistant_with_tools(*call_ids: str) -> dict:
|
||||
"""Create an assistant message with tool_calls."""
|
||||
return {"role": "assistant", "tool_calls": [_tc(cid) for cid in call_ids], "content": None}
|
||||
|
||||
|
||||
def _make_compressor(**kwargs) -> ContextCompressor:
|
||||
defaults = dict(
|
||||
model="test-model",
|
||||
threshold_percent=0.75,
|
||||
protect_first_n=3,
|
||||
protect_last_n=4,
|
||||
quiet_mode=True,
|
||||
)
|
||||
defaults.update(kwargs)
|
||||
with patch("agent.context_compressor.get_model_context_length", return_value=8000):
|
||||
return ContextCompressor(**defaults)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _align_boundary_backward tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestAlignBoundaryBackward:
|
||||
"""Test that compress-end boundary never splits a tool_call/result group."""
|
||||
|
||||
def test_boundary_at_clean_position(self):
|
||||
"""Boundary after a user message — no adjustment needed."""
|
||||
comp = _make_compressor()
|
||||
messages = [
|
||||
{"role": "system", "content": "sys"},
|
||||
{"role": "user", "content": "hello"},
|
||||
{"role": "assistant", "content": "hi"},
|
||||
{"role": "user", "content": "do something"},
|
||||
_assistant_with_tools("tc_1"),
|
||||
_tool_result("tc_1", "done"),
|
||||
{"role": "user", "content": "thanks"}, # idx=6
|
||||
{"role": "assistant", "content": "np"},
|
||||
]
|
||||
# Boundary at 7, messages[6] = user — no adjustment
|
||||
assert comp._align_boundary_backward(messages, 7) == 7
|
||||
|
||||
def test_boundary_after_assistant_with_tools(self):
|
||||
"""Original case: boundary right after assistant with tool_calls."""
|
||||
comp = _make_compressor()
|
||||
messages = [
|
||||
{"role": "system", "content": "sys"},
|
||||
{"role": "user", "content": "hello"},
|
||||
{"role": "assistant", "content": "hi"},
|
||||
_assistant_with_tools("tc_1", "tc_2"), # idx=3
|
||||
_tool_result("tc_1"), # idx=4
|
||||
_tool_result("tc_2"), # idx=5
|
||||
{"role": "user", "content": "next"},
|
||||
]
|
||||
# Boundary at 4, messages[3] = assistant with tool_calls → pull back to 3
|
||||
assert comp._align_boundary_backward(messages, 4) == 3
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# End-to-end: compression must not lose tool results
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestCompressionToolResultPreservation:
|
||||
"""Verify that compress() never silently drops tool results."""
|
||||
|
||||
def test_parallel_tool_results_not_lost(self):
|
||||
"""The exact scenario that triggered silent data loss before the fix."""
|
||||
comp = _make_compressor(protect_first_n=3, protect_last_n=4)
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": "You are helpful."}, # 0
|
||||
{"role": "user", "content": "Hello"}, # 1
|
||||
{"role": "assistant", "content": "Hi there!"}, # 2 (end of head)
|
||||
{"role": "user", "content": "Read 7 files for me"}, # 3
|
||||
_assistant_with_tools("tc_A", "tc_B", "tc_C", "tc_D", "tc_E", "tc_F", "tc_G"), # 4
|
||||
_tool_result("tc_A", "content of file A"), # 5
|
||||
_tool_result("tc_B", "content of file B"), # 6
|
||||
_tool_result("tc_C", "content of file C"), # 7
|
||||
_tool_result("tc_D", "content of file D"), # 8
|
||||
_tool_result("tc_E", "content of file E"), # 9
|
||||
_tool_result("tc_F", "content of file F"), # 10
|
||||
_tool_result("tc_G", "CRITICAL DATA in file G"), # 11 ← compress_end=15-4=11
|
||||
{"role": "user", "content": "Now summarize them"}, # 12
|
||||
{"role": "assistant", "content": "Here is the summary..."}, # 13
|
||||
{"role": "user", "content": "Thanks"}, # 14
|
||||
]
|
||||
# 15 messages. compress_end = 15 - 4 = 11 (before fix: splits tool group)
|
||||
|
||||
fake_summary = "[Summary of earlier conversation]"
|
||||
with patch.object(comp, "_generate_summary", return_value=fake_summary):
|
||||
result = comp.compress(messages, current_tokens=7000)
|
||||
|
||||
# After compression, no tool results should be orphaned/lost.
|
||||
# All tool results in the result must have a matching assistant tool_call.
|
||||
assistant_call_ids = set()
|
||||
for msg in result:
|
||||
if msg.get("role") == "assistant":
|
||||
for tc in msg.get("tool_calls") or []:
|
||||
cid = tc.get("id", "")
|
||||
if cid:
|
||||
assistant_call_ids.add(cid)
|
||||
|
||||
tool_result_ids = set()
|
||||
for msg in result:
|
||||
if msg.get("role") == "tool":
|
||||
cid = msg.get("tool_call_id")
|
||||
if cid:
|
||||
tool_result_ids.add(cid)
|
||||
|
||||
# Every tool result must have a parent — no orphans
|
||||
orphaned = tool_result_ids - assistant_call_ids
|
||||
assert not orphaned, f"Orphaned tool results found (data loss!): {orphaned}"
|
||||
|
||||
# Every assistant tool_call must have a real result (not a stub)
|
||||
for msg in result:
|
||||
if msg.get("role") == "tool":
|
||||
assert msg["content"] != "[Result from earlier conversation — see context summary above]", \
|
||||
f"Stub result found for {msg.get('tool_call_id')} — real result was lost"
|
||||
@@ -0,0 +1,327 @@
|
||||
"""Test: the context engine is notified of a compression-boundary rollover.
|
||||
|
||||
When _compress_context rotates session_id (compression split), the active
|
||||
context engine receives on_session_start(new_sid, boundary_reason="compression",
|
||||
old_session_id=<old>). This lets plugin engines (e.g. hermes-lcm) preserve
|
||||
DAG lineage across the split instead of treating it as a fresh /new.
|
||||
|
||||
See hermes-lcm#68: after Hermes compresses and mints a new physical session,
|
||||
LCM was losing continuity (compression_count: 1, store_messages: 0,
|
||||
dag_nodes: 0). With boundary_reason="compression" plugins can distinguish
|
||||
this from a real user-initiated /new.
|
||||
"""
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from agent.conversation_compression import (
|
||||
finalize_context_engine_compression_notification,
|
||||
)
|
||||
|
||||
class TestCompressionBoundaryHook:
|
||||
def _make_agent(self, session_db):
|
||||
with patch.dict(os.environ, {"OPENROUTER_API_KEY": "test-key"}):
|
||||
from run_agent import AIAgent
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
quiet_mode=True,
|
||||
session_db=session_db,
|
||||
session_id="original-session",
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
# ROTATION fallback — pin in_place=False regardless of default (#38763).
|
||||
agent.compression_in_place = False
|
||||
return agent
|
||||
|
||||
def test_on_session_start_called_with_compression_boundary(self):
|
||||
from hermes_state import SessionDB
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db = SessionDB(db_path=Path(tmpdir) / "test.db")
|
||||
agent = self._make_agent(db)
|
||||
|
||||
# Stub the context compressor: we only need to observe the hook.
|
||||
compressor = MagicMock()
|
||||
compressor.compress.return_value = [
|
||||
{"role": "user", "content": "[CONTEXT COMPACTION] summary"},
|
||||
{"role": "user", "content": "tail question"},
|
||||
]
|
||||
compressor.compression_count = 1
|
||||
compressor.last_prompt_tokens = 0
|
||||
compressor.last_completion_tokens = 0
|
||||
# Avoid the summary-error warning path
|
||||
compressor._last_summary_error = None
|
||||
# MagicMock auto-creates truthy attrs; explicitly clear the abort
|
||||
# flag so the post-compress abort branch in
|
||||
# conversation_compression.py does not short-circuit before the
|
||||
# session-id rotation we are asserting on.
|
||||
compressor._last_compress_aborted = False
|
||||
agent.context_compressor = compressor
|
||||
|
||||
original_sid = agent.session_id
|
||||
messages = [
|
||||
{"role": "user", "content": f"m{i}"} for i in range(10)
|
||||
]
|
||||
|
||||
agent._compress_context(messages, "sys", approx_tokens=10_000)
|
||||
|
||||
# Session_id rotated
|
||||
assert agent.session_id != original_sid, \
|
||||
"compression should rotate session_id when session_db is set"
|
||||
|
||||
# Hook fired with boundary_reason="compression" and old_session_id
|
||||
calls = [
|
||||
c for c in compressor.on_session_start.call_args_list
|
||||
]
|
||||
assert calls, "on_session_start was never called on the context engine"
|
||||
# Find the compression boundary call (there may be others from init)
|
||||
comp_calls = [
|
||||
c for c in calls
|
||||
if c.kwargs.get("boundary_reason") == "compression"
|
||||
]
|
||||
assert comp_calls, (
|
||||
f"Expected an on_session_start call with "
|
||||
f"boundary_reason='compression', got {calls!r}"
|
||||
)
|
||||
call = comp_calls[-1]
|
||||
# Positional new session_id
|
||||
assert call.args and call.args[0] == agent.session_id, \
|
||||
f"Expected new session_id as first positional arg, got {call!r}"
|
||||
assert call.kwargs.get("old_session_id") == original_sid, \
|
||||
f"Expected old_session_id={original_sid!r}, got {call.kwargs!r}"
|
||||
assert len(comp_calls) == 1
|
||||
|
||||
def test_automatic_notification_follows_core_persistence(self):
|
||||
from hermes_state import SessionDB
|
||||
|
||||
events = []
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db = SessionDB(db_path=Path(tmpdir) / "test.db")
|
||||
agent = self._make_agent(db)
|
||||
compressor = MagicMock()
|
||||
compressor.compress.return_value = [
|
||||
{"role": "user", "content": "summary"}
|
||||
]
|
||||
compressor.compression_count = 1
|
||||
compressor.last_prompt_tokens = 0
|
||||
compressor.last_completion_tokens = 0
|
||||
compressor._last_summary_error = None
|
||||
compressor._last_compress_aborted = False
|
||||
compressor.on_session_start.side_effect = (
|
||||
lambda *_args, **kwargs: events.append(
|
||||
kwargs.get("boundary_reason")
|
||||
)
|
||||
)
|
||||
agent.context_compressor = compressor
|
||||
original_publish = db.publish_compression_child
|
||||
|
||||
def _record_publish(*args, **kwargs):
|
||||
result = original_publish(*args, **kwargs)
|
||||
events.append("persist")
|
||||
return result
|
||||
|
||||
with patch.object(
|
||||
db, "publish_compression_child", side_effect=_record_publish
|
||||
):
|
||||
agent._compress_context(
|
||||
[{"role": "user", "content": "request"}],
|
||||
"sys",
|
||||
approx_tokens=100,
|
||||
)
|
||||
|
||||
assert events == ["persist", "compression"]
|
||||
|
||||
def test_failure_before_persistence_does_not_notify(self):
|
||||
from hermes_state import SessionDB
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db = SessionDB(db_path=Path(tmpdir) / "test.db")
|
||||
agent = self._make_agent(db)
|
||||
compressor = MagicMock()
|
||||
compressor.compress.side_effect = RuntimeError("synthetic compression failure")
|
||||
agent.context_compressor = compressor
|
||||
|
||||
with pytest.raises(RuntimeError, match="synthetic compression failure"):
|
||||
agent._compress_context(
|
||||
[{"role": "user", "content": "request"}],
|
||||
"sys",
|
||||
approx_tokens=100,
|
||||
)
|
||||
|
||||
compressor.on_session_start.assert_not_called()
|
||||
|
||||
|
||||
def test_no_progress_does_not_notify(self):
|
||||
from hermes_state import SessionDB
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db = SessionDB(db_path=Path(tmpdir) / "test.db")
|
||||
agent = self._make_agent(db)
|
||||
compressor = MagicMock()
|
||||
compressor.compress.side_effect = lambda messages, **_kwargs: messages
|
||||
compressor._last_compress_aborted = False
|
||||
agent.context_compressor = compressor
|
||||
messages = [{"role": "user", "content": "request"}]
|
||||
|
||||
returned, _ = agent._compress_context(
|
||||
messages,
|
||||
"sys",
|
||||
approx_tokens=100,
|
||||
)
|
||||
|
||||
assert returned is messages
|
||||
compressor.on_session_start.assert_not_called()
|
||||
|
||||
|
||||
def test_no_hook_when_no_session_db(self):
|
||||
"""Without session_db, session_id does not rotate and the hook is not fired."""
|
||||
from run_agent import AIAgent
|
||||
with patch.dict(os.environ, {"OPENROUTER_API_KEY": "test-key"}):
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
quiet_mode=True,
|
||||
session_db=None,
|
||||
session_id="original-session",
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
|
||||
compressor = MagicMock()
|
||||
compressor.compress.return_value = [{"role": "user", "content": "x"}]
|
||||
compressor.compression_count = 1
|
||||
compressor.last_prompt_tokens = 0
|
||||
compressor.last_completion_tokens = 0
|
||||
compressor._last_summary_error = None
|
||||
agent.context_compressor = compressor
|
||||
|
||||
original_sid = agent.session_id
|
||||
agent._compress_context([{"role": "user", "content": "m"}], "sys", approx_tokens=100)
|
||||
|
||||
# No DB => no rotation => no compression-boundary hook
|
||||
assert agent.session_id == original_sid
|
||||
comp_calls = [
|
||||
c for c in compressor.on_session_start.call_args_list
|
||||
if c.kwargs.get("boundary_reason") == "compression"
|
||||
]
|
||||
assert not comp_calls, (
|
||||
f"No compression hook should fire without session_db rotation, "
|
||||
f"got {comp_calls!r}"
|
||||
)
|
||||
|
||||
def test_hook_failure_does_not_break_compression(self):
|
||||
"""If the context engine raises from on_session_start, compression still completes."""
|
||||
from hermes_state import SessionDB
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db = SessionDB(db_path=Path(tmpdir) / "test.db")
|
||||
agent = self._make_agent(db)
|
||||
|
||||
compressor = MagicMock()
|
||||
compressor.compress.return_value = [{"role": "user", "content": "summary"}]
|
||||
compressor.compression_count = 1
|
||||
compressor.last_prompt_tokens = 0
|
||||
compressor.last_completion_tokens = 0
|
||||
compressor._last_summary_error = None
|
||||
compressor._last_compress_aborted = False
|
||||
|
||||
# Raise only on the compression-boundary call, not on earlier calls.
|
||||
def _raise_on_compression(*args, **kwargs):
|
||||
if kwargs.get("boundary_reason") == "compression":
|
||||
raise RuntimeError("plugin exploded")
|
||||
return None
|
||||
compressor.on_session_start.side_effect = _raise_on_compression
|
||||
agent.context_compressor = compressor
|
||||
|
||||
original_sid = agent.session_id
|
||||
|
||||
# Must not raise. Input must be large enough that the fake
|
||||
# compressor's one-message summary is a genuine shrink — the
|
||||
# no-growth commit guard refuses to rotate on transcript growth.
|
||||
compressed, _prompt = agent._compress_context(
|
||||
[{"role": "user", "content": "m" * 400}], "sys", approx_tokens=100
|
||||
)
|
||||
assert compressed
|
||||
assert agent.session_id != original_sid
|
||||
|
||||
|
||||
class TestSessionCompressEvent:
|
||||
"""The session:compress event_callback fires after a compression split."""
|
||||
|
||||
def _make_agent(self, session_db, event_callback=None):
|
||||
with patch.dict(os.environ, {"OPENROUTER_API_KEY": "test-key"}):
|
||||
from run_agent import AIAgent
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
quiet_mode=True,
|
||||
session_db=session_db,
|
||||
session_id="original-session",
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
event_callback=event_callback,
|
||||
)
|
||||
# ROTATION fallback — pin in_place=False regardless of default (#38763).
|
||||
agent.compression_in_place = False
|
||||
return agent
|
||||
|
||||
def _stub_compressor(self):
|
||||
compressor = MagicMock()
|
||||
compressor.compress.return_value = [
|
||||
{"role": "user", "content": "[CONTEXT COMPACTION] summary"},
|
||||
{"role": "user", "content": "tail"},
|
||||
]
|
||||
compressor.compression_count = 1
|
||||
compressor.last_prompt_tokens = 0
|
||||
compressor.last_completion_tokens = 0
|
||||
compressor._last_summary_error = None
|
||||
compressor._last_compress_aborted = False
|
||||
return compressor
|
||||
|
||||
def test_event_emitted_on_compression(self):
|
||||
from hermes_state import SessionDB
|
||||
|
||||
events = []
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db = SessionDB(db_path=Path(tmpdir) / "test.db")
|
||||
agent = self._make_agent(
|
||||
db, event_callback=lambda et, ctx: events.append((et, ctx))
|
||||
)
|
||||
original_sid = agent.session_id
|
||||
agent.context_compressor = self._stub_compressor()
|
||||
|
||||
agent._compress_context(
|
||||
[{"role": "user", "content": f"m{i}"} for i in range(10)],
|
||||
"sys",
|
||||
approx_tokens=10_000,
|
||||
)
|
||||
|
||||
compress_events = [e for e in events if e[0] == "session:compress"]
|
||||
assert compress_events, f"session:compress not emitted, got {events!r}"
|
||||
_, ctx = compress_events[-1]
|
||||
assert ctx["session_id"] == agent.session_id
|
||||
assert ctx["old_session_id"] == original_sid
|
||||
assert ctx["compression_count"] == 1
|
||||
|
||||
def test_no_callback_is_safe(self):
|
||||
"""Compression must work when no event_callback is wired."""
|
||||
from hermes_state import SessionDB
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db = SessionDB(db_path=Path(tmpdir) / "test.db")
|
||||
agent = self._make_agent(db, event_callback=None)
|
||||
agent.context_compressor = self._stub_compressor()
|
||||
compressed, _ = agent._compress_context(
|
||||
[{"role": "user", "content": "m"}], "sys", approx_tokens=100
|
||||
)
|
||||
assert compressed
|
||||
|
||||
@@ -0,0 +1,227 @@
|
||||
"""Regression test for re-arming the compression budget after tool progress."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from run_agent import AIAgent
|
||||
|
||||
|
||||
def _tool_call():
|
||||
return SimpleNamespace(
|
||||
id="call_1",
|
||||
type="function",
|
||||
function=SimpleNamespace(name="web_search", arguments='{"query": "x"}'),
|
||||
)
|
||||
|
||||
|
||||
def _tool_response(prompt_tokens: int):
|
||||
message = SimpleNamespace(
|
||||
content=None,
|
||||
reasoning_content=None,
|
||||
reasoning=None,
|
||||
tool_calls=[_tool_call()],
|
||||
)
|
||||
return SimpleNamespace(
|
||||
choices=[SimpleNamespace(message=message, finish_reason="tool_calls")],
|
||||
model="test/model",
|
||||
usage=SimpleNamespace(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=1,
|
||||
total_tokens=prompt_tokens + 1,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _final_response():
|
||||
message = SimpleNamespace(
|
||||
content="done",
|
||||
reasoning_content=None,
|
||||
reasoning=None,
|
||||
tool_calls=None,
|
||||
)
|
||||
return SimpleNamespace(
|
||||
choices=[SimpleNamespace(message=message, finish_reason="stop")],
|
||||
model="test/model",
|
||||
usage=None,
|
||||
)
|
||||
|
||||
|
||||
def _malformed_response():
|
||||
return SimpleNamespace(choices=[], model="test/model", usage=None)
|
||||
|
||||
|
||||
def _tool_definition():
|
||||
return {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "web_search",
|
||||
"description": "Search the web",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"query": {"type": "string"}},
|
||||
"required": ["query"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("prompt_tokens", "expected_compactions", "provider_recovery"),
|
||||
[(50, 1, False), (150, 1, False), (50, 2, True)],
|
||||
ids=[
|
||||
"pressure-cleared-anchored-no-recompaction",
|
||||
"pressure-still-high-stays-capped",
|
||||
"pressure-cleared-rearms-after-provider-recovery",
|
||||
],
|
||||
)
|
||||
def test_pre_api_compression_budget_rearms_only_after_pressure_clears(
|
||||
prompt_tokens: int,
|
||||
expected_compactions: int,
|
||||
provider_recovery: bool,
|
||||
):
|
||||
"""Only provider-confirmed headroom starts a new pressure episode.
|
||||
|
||||
Usage-anchored accounting update: once the provider reports
|
||||
``prompt_tokens=50`` for the full transcript, later pre-API checks anchor
|
||||
on that real reading plus a delta estimate of the few appended messages —
|
||||
the scripted whole-history rough estimate (200) no longer drives the
|
||||
decision, so the pressure-cleared case performs exactly ONE compaction
|
||||
(the pre-anchor one). The budget-rearm mechanics remain covered by the
|
||||
provider-recovery variant, whose first response carries no usage (no
|
||||
anchor) and therefore still compacts on the rough estimate.
|
||||
"""
|
||||
with (
|
||||
patch("model_tools.get_tool_definitions", return_value=[_tool_definition()]),
|
||||
patch("model_tools.check_toolset_requirements", return_value={}),
|
||||
patch("agent.process_bootstrap.OpenAI"),
|
||||
patch("agent.model_metadata.get_model_context_length", return_value=256_000),
|
||||
patch("agent.context_compressor.get_model_context_length", return_value=256_000),
|
||||
):
|
||||
agent = AIAgent(
|
||||
api_key="test-key-1234567890",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
max_iterations=6,
|
||||
)
|
||||
|
||||
agent.client = MagicMock()
|
||||
responses = [_tool_response(prompt_tokens), _final_response()]
|
||||
if provider_recovery:
|
||||
responses.insert(0, _malformed_response())
|
||||
agent._fallback_chain = [object()]
|
||||
agent._try_activate_fallback = MagicMock(return_value=True)
|
||||
agent.client.chat.completions.create.side_effect = responses
|
||||
agent._cached_system_prompt = "You are helpful."
|
||||
agent._use_prompt_caching = False
|
||||
agent._disable_streaming = True
|
||||
agent.tool_delay = 0
|
||||
agent.save_trajectories = False
|
||||
agent.max_compression_attempts = 1
|
||||
|
||||
compressor = MagicMock()
|
||||
compressor.protect_first_n = 3
|
||||
compressor.protect_last_n = 20
|
||||
compressor.threshold_tokens = 100
|
||||
compressor.context_length = 1_000
|
||||
compressor.last_prompt_tokens = -1
|
||||
compressor._verify_compaction_cleared_threshold = False
|
||||
compressor.awaiting_real_usage_after_compression = False
|
||||
compressor.should_compress.side_effect = lambda tokens: tokens >= 100
|
||||
compressor.should_compress_info.return_value = (False, None)
|
||||
compressor.should_compress_preflight.return_value = False
|
||||
compressor.should_defer_preflight_to_real_usage.return_value = False
|
||||
compressor.get_active_compression_failure_cooldown.return_value = None
|
||||
compressor.select_context.return_value = None
|
||||
compressor.get_automatic_compaction_status_message.return_value = ""
|
||||
|
||||
def _update_from_response(usage):
|
||||
# Mirror the real compressor: the next provider usage reading
|
||||
# consumes the completed-compaction verification latch.
|
||||
compressor.last_prompt_tokens = int(usage.get("prompt_tokens", 0) or 0)
|
||||
compressor._verify_compaction_cleared_threshold = False
|
||||
compressor.awaiting_real_usage_after_compression = False
|
||||
|
||||
compressor.update_from_response.side_effect = _update_from_response
|
||||
agent.compression_enabled = True
|
||||
agent.context_compressor = compressor
|
||||
|
||||
estimate_values = iter([200, 190, 200, 10])
|
||||
_last_estimate = [10]
|
||||
|
||||
def _next_estimate(messages=None, *_args, **_kwargs):
|
||||
# The scripted sequence prices the WHOLE history; a usage-anchored gate
|
||||
# estimates only the few messages appended since the real reading.
|
||||
if isinstance(messages, list) and len(messages) <= 4:
|
||||
return 10
|
||||
# The provider-recovery variant re-runs the pre-API preflight after
|
||||
# fallback activation (#84733), consuming an extra estimate reading.
|
||||
# Hold the final low-pressure value once the scripted sequence is
|
||||
# exhausted instead of raising StopIteration.
|
||||
try:
|
||||
_last_estimate[0] = next(estimate_values)
|
||||
except StopIteration:
|
||||
pass
|
||||
return _last_estimate[0]
|
||||
|
||||
compress_calls = []
|
||||
|
||||
def _fake_compress(messages, _system_message, **_kwargs):
|
||||
compress_calls.append(messages)
|
||||
# Arm the same provider-verification boundary the real compression
|
||||
# path arms after a completed compaction.
|
||||
compressor._verify_compaction_cleared_threshold = True
|
||||
compressor.awaiting_real_usage_after_compression = True
|
||||
return list(messages), "compressed prompt"
|
||||
|
||||
def _fake_execute_tool_calls(assistant_message, messages, *_args):
|
||||
tool_call = assistant_message.tool_calls[0]
|
||||
messages.append(
|
||||
{
|
||||
"role": "tool",
|
||||
"name": tool_call.function.name,
|
||||
"tool_call_id": tool_call.id,
|
||||
"content": "ok",
|
||||
}
|
||||
)
|
||||
|
||||
history = [
|
||||
{"role": "user" if i % 2 == 0 else "assistant", "content": f"msg {i}"}
|
||||
for i in range(30)
|
||||
]
|
||||
|
||||
with (
|
||||
patch(
|
||||
"agent.turn_context.estimate_request_tokens_rough",
|
||||
return_value=10,
|
||||
),
|
||||
patch(
|
||||
"agent.model_metadata.estimate_messages_tokens_rough",
|
||||
side_effect=_next_estimate,
|
||||
),
|
||||
patch(
|
||||
"agent.conversation_loop._estimate_tools_tokens_rough",
|
||||
return_value=0,
|
||||
),
|
||||
patch.object(agent, "_compress_context", side_effect=_fake_compress),
|
||||
patch.object(agent, "_execute_tool_calls", side_effect=_fake_execute_tool_calls),
|
||||
patch.object(agent, "_flush_messages_to_session_db", return_value=True),
|
||||
patch.object(agent, "_persist_session"),
|
||||
patch.object(agent, "_save_trajectory"),
|
||||
patch.object(agent, "_cleanup_task_resources"),
|
||||
):
|
||||
result = agent.run_conversation("do a lot of tool work", conversation_history=history)
|
||||
|
||||
assert result["completed"] is True
|
||||
assert result["final_response"] == "done"
|
||||
assert len(compress_calls) == expected_compactions, (
|
||||
"same-turn compression must re-arm only after the provider confirms "
|
||||
f"headroom; got {len(compress_calls)} compactions for "
|
||||
f"prompt_tokens={prompt_tokens}"
|
||||
)
|
||||
@@ -0,0 +1,280 @@
|
||||
"""Behavioral tests for provider-confirmed compression-budget rearming.
|
||||
|
||||
``compression_attempts`` is a shared per-turn backstop (pre-API gate,
|
||||
overflow/413 handlers, post-tool gate). Before the refund fix, *successful*
|
||||
pre-API compactions consumed it permanently: a marathon tool turn burned all
|
||||
attempts on compactions that worked, the pre-API gate went dark for the rest
|
||||
of the turn, and the context grew unchecked until the provider rejected the
|
||||
request terminally ("max compression attempts (N) reached").
|
||||
|
||||
The budget is rearmed only when a completed compaction is followed by a real
|
||||
provider prompt count below the configured threshold. Rough estimates and
|
||||
usage-less responses cannot reopen the anti-thrash cap.
|
||||
|
||||
These tests drive ``run_conversation()`` through real tool iterations — no
|
||||
source inspection, only observable compaction counts.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from agent.conversation_loop import _should_rearm_compression_budget
|
||||
from run_agent import AIAgent
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Unit: refund decision
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRearmDecision:
|
||||
def test_provider_confirmed_recovery_rearms(self):
|
||||
assert _should_rearm_compression_budget(
|
||||
2,
|
||||
completed_compaction_pending=True,
|
||||
prompt_tokens=7_999,
|
||||
threshold_tokens=10_000,
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("attempts", "pending", "prompt_tokens", "threshold_tokens"),
|
||||
[
|
||||
(0, True, 7_999, 10_000),
|
||||
(2, False, 7_999, 10_000),
|
||||
(2, True, 0, 10_000),
|
||||
(2, True, 10_000, 10_000),
|
||||
(2, True, 10_001, 10_000),
|
||||
(2, True, 7_999, 0),
|
||||
],
|
||||
)
|
||||
def test_unverified_or_pressured_response_keeps_budget_burned(
|
||||
self, attempts, pending, prompt_tokens, threshold_tokens
|
||||
):
|
||||
assert not _should_rearm_compression_budget(
|
||||
attempts,
|
||||
completed_compaction_pending=pending,
|
||||
prompt_tokens=prompt_tokens,
|
||||
threshold_tokens=threshold_tokens,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Behavioral: marathon tool turn keeps compacting past the old cap
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _tool_call(i: int):
|
||||
return SimpleNamespace(
|
||||
id=f"call_{i}",
|
||||
type="function",
|
||||
# Vary the query per call: a real marathon turn issues distinct
|
||||
# lookups, and identical (args, result) pairs are now legitimately
|
||||
# deduped into reference stubs by the stall-guard subsystem —
|
||||
# zero-variance args here would deflate the very context pressure
|
||||
# this test exists to exercise.
|
||||
function=SimpleNamespace(name="web_search", arguments=f'{{"query": "x{i}"}}'),
|
||||
)
|
||||
|
||||
|
||||
def _usage(prompt_tokens: int | None):
|
||||
if prompt_tokens is None:
|
||||
return None
|
||||
return SimpleNamespace(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=1,
|
||||
total_tokens=prompt_tokens + 1,
|
||||
)
|
||||
|
||||
|
||||
def _tool_response(i: int, prompt_tokens: int | None):
|
||||
msg = SimpleNamespace(
|
||||
content=None,
|
||||
reasoning_content=None,
|
||||
reasoning=None,
|
||||
tool_calls=[_tool_call(i)],
|
||||
)
|
||||
choice = SimpleNamespace(message=msg, finish_reason="tool_calls")
|
||||
return SimpleNamespace(
|
||||
choices=[choice], model="test/model", usage=_usage(prompt_tokens)
|
||||
)
|
||||
|
||||
|
||||
def _stop_response(prompt_tokens: int | None):
|
||||
msg = SimpleNamespace(
|
||||
content="done",
|
||||
reasoning_content=None,
|
||||
reasoning=None,
|
||||
tool_calls=None,
|
||||
)
|
||||
choice = SimpleNamespace(message=msg, finish_reason="stop")
|
||||
return SimpleNamespace(
|
||||
choices=[choice], model="test/model", usage=_usage(prompt_tokens)
|
||||
)
|
||||
|
||||
|
||||
def _make_tool_defs(*names: str) -> list:
|
||||
return [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": n,
|
||||
"description": f"{n} tool",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
}
|
||||
for n in names
|
||||
]
|
||||
|
||||
|
||||
THRESHOLD = 10_000
|
||||
|
||||
# Each tool result is large enough that the assembled request crosses
|
||||
# THRESHOLD every iteration (estimator is ~chars/4), forcing one pre-API
|
||||
# compaction per iteration — but stays below the 100K-char per-result
|
||||
# persistence threshold (tools/budget_config.py) so it reaches the context
|
||||
# untruncated.
|
||||
BIG_TOOL_RESULT = "x" * 60_000
|
||||
|
||||
|
||||
def _coherent_compressor() -> MagicMock:
|
||||
"""A compressor whose should_compress() reflects the passed estimate.
|
||||
|
||||
Unlike the always-True stub in the attempt-cap tests, this models the
|
||||
real coupling: pressure at/over threshold → compress; pressure gone →
|
||||
healthy. That coupling is what makes the refund safe.
|
||||
"""
|
||||
compressor = MagicMock()
|
||||
compressor.protect_first_n = 3
|
||||
compressor.protect_last_n = 20
|
||||
compressor.threshold_tokens = THRESHOLD
|
||||
compressor.context_length = 200_000
|
||||
compressor.last_prompt_tokens = 0
|
||||
compressor._verify_compaction_cleared_threshold = False
|
||||
compressor.awaiting_real_usage_after_compression = False
|
||||
compressor.should_compress.side_effect = lambda t=None: (t or 0) >= THRESHOLD
|
||||
compressor.should_defer_preflight_to_real_usage.return_value = False
|
||||
compressor.get_active_compression_failure_cooldown.return_value = None
|
||||
|
||||
def _update_from_response(usage):
|
||||
compressor.last_prompt_tokens = int(usage.get("prompt_tokens", 0) or 0)
|
||||
compressor._verify_compaction_cleared_threshold = False
|
||||
compressor.awaiting_real_usage_after_compression = False
|
||||
|
||||
compressor.update_from_response.side_effect = _update_from_response
|
||||
return compressor
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def agent():
|
||||
with (
|
||||
patch("model_tools.get_tool_definitions", return_value=_make_tool_defs("web_search")),
|
||||
patch("model_tools.check_toolset_requirements", return_value={}),
|
||||
patch("agent.process_bootstrap.OpenAI"),
|
||||
):
|
||||
a = AIAgent(
|
||||
api_key="test-key-1234567890",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
max_iterations=20,
|
||||
)
|
||||
a.client = MagicMock()
|
||||
a._cached_system_prompt = "You are helpful."
|
||||
a._use_prompt_caching = False
|
||||
a._disable_streaming = True
|
||||
a.tool_delay = 0
|
||||
a.save_trajectories = False
|
||||
a.compression_enabled = True
|
||||
a.context_compressor = _coherent_compressor()
|
||||
return a
|
||||
|
||||
|
||||
def _run_marathon_turn(
|
||||
agent, n_tool_iterations: int, *, provider_prompt_tokens: int | None
|
||||
):
|
||||
"""Drive one turn of ``n_tool_iterations`` oversized tool results."""
|
||||
responses = [
|
||||
_tool_response(i, provider_prompt_tokens) for i in range(n_tool_iterations)
|
||||
]
|
||||
responses.append(_stop_response(provider_prompt_tokens))
|
||||
agent.client.chat.completions.create.side_effect = responses
|
||||
|
||||
compress_calls = []
|
||||
|
||||
def _fake_compress(messages, system_message, **_kwargs):
|
||||
# Model a compaction that works: blank out every oversized payload,
|
||||
# keeping roles and tool-call pairing intact so sanitization is
|
||||
# unaffected. Arm the same provider-verification boundary as the real
|
||||
# compression path.
|
||||
compress_calls.append(len(messages))
|
||||
agent.context_compressor._verify_compaction_cleared_threshold = True
|
||||
agent.context_compressor.awaiting_real_usage_after_compression = True
|
||||
compacted = [
|
||||
dict(m, content="[summarized]")
|
||||
if isinstance(m, dict) and len(str(m.get("content") or "")) > 5_000
|
||||
else m
|
||||
for m in messages
|
||||
]
|
||||
return compacted, "compressed prompt"
|
||||
|
||||
with (
|
||||
patch.object(agent, "_compress_context", side_effect=_fake_compress),
|
||||
patch.object(agent, "_persist_session"),
|
||||
patch.object(agent, "_save_trajectory"),
|
||||
patch.object(agent, "_cleanup_task_resources"),
|
||||
patch(
|
||||
"model_tools.handle_function_call",
|
||||
lambda name, args, task_id=None, **kwargs: json.dumps(
|
||||
{"ok": True, "payload": BIG_TOOL_RESULT}
|
||||
),
|
||||
),
|
||||
):
|
||||
result = agent.run_conversation("do a lot of tool work")
|
||||
|
||||
return result, compress_calls
|
||||
|
||||
|
||||
class TestCompressionBudgetRefund:
|
||||
def test_marathon_turn_compacts_past_the_per_turn_cap(self, agent):
|
||||
"""8 oversized tool iterations → more compactions than the old cap.
|
||||
|
||||
Pre-refund, the 4th+ pressure spike found the budget exhausted, the
|
||||
pre-API gate stayed dark, and the request grew unchecked. With the
|
||||
refund, every genuine pressure spike is compacted and the turn
|
||||
completes.
|
||||
"""
|
||||
assert agent.max_compression_attempts == 3 # config default
|
||||
result, compress_calls = _run_marathon_turn(
|
||||
agent,
|
||||
n_tool_iterations=8,
|
||||
provider_prompt_tokens=THRESHOLD - 1,
|
||||
)
|
||||
|
||||
assert result["completed"] is True
|
||||
assert len(compress_calls) > 3, (
|
||||
"successful compactions must refund the per-turn budget; "
|
||||
f"got only {len(compress_calls)} compactions for 8 pressure spikes"
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("provider_prompt_tokens", [None, THRESHOLD])
|
||||
def test_unverified_or_pressured_compaction_stays_capped(
|
||||
self, agent, provider_prompt_tokens
|
||||
):
|
||||
"""Missing usage or real usage at threshold cannot recycle the cap."""
|
||||
result, compress_calls = _run_marathon_turn(
|
||||
agent,
|
||||
n_tool_iterations=8,
|
||||
provider_prompt_tokens=provider_prompt_tokens,
|
||||
)
|
||||
|
||||
assert result["completed"] is True
|
||||
assert len(compress_calls) <= agent.max_compression_attempts, (
|
||||
"without provider-confirmed headroom the per-turn cap must hold; "
|
||||
f"got {len(compress_calls)} compactions"
|
||||
)
|
||||
@@ -0,0 +1,240 @@
|
||||
"""Compression race at the flush chokepoint: a turn writing against a session
|
||||
already closed by compression must adopt the LIVE continuation tip instead of
|
||||
dying with ``session_persistence_failed`` and a misleading "full disk" dialog.
|
||||
|
||||
The store resolves the continuation chain transitively via the canonical API
|
||||
``SessionDB.get_compression_tip`` (bounded walk, excludes branch/delegate/tool
|
||||
children, prefers live children over stale closed siblings). This suite proves
|
||||
the agent flush path:
|
||||
|
||||
* adopts a unique live child (depth-1 case),
|
||||
* follows a chain of >=2 compressions to the live head — THE regression the
|
||||
depth-1 ``find_live_compression_child`` API missed (#82001),
|
||||
* fails closed when no continuation exists (no retry loop),
|
||||
* fails closed when the resolved tip is itself closed (``ws_orphan_reap``),
|
||||
* performs the tip lookup exactly once per flush (adoption budget), and
|
||||
* never renders the failure with the historical full-disk misdiagnosis.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
from hermes_state import SessionDB
|
||||
from run_agent import AIAgent
|
||||
|
||||
|
||||
def _flush_agent(db, session_id):
|
||||
"""Bind the real flush methods onto a stand-in over a live SessionDB."""
|
||||
agent = SimpleNamespace(
|
||||
_session_db=db,
|
||||
_session_db_created=True,
|
||||
_persist_disabled=False,
|
||||
session_id=session_id,
|
||||
_session_persist_lock=None,
|
||||
_flushed_db_message_ids=set(),
|
||||
_flushed_db_message_session_id=None,
|
||||
_last_flushed_db_idx=0,
|
||||
_db_flush_scan_prefix=None,
|
||||
_persist_user_message_idx=None,
|
||||
_persist_user_message_override=None,
|
||||
_persist_user_message_timestamp=None,
|
||||
_pending_cli_user_message=None,
|
||||
_active_session_turn_lease_holder=None,
|
||||
_last_persistence_error_cause=None,
|
||||
_compression_adoption_failed=False,
|
||||
)
|
||||
agent._ensure_db_session = lambda: None
|
||||
agent._flush_messages_to_session_db = (
|
||||
AIAgent._flush_messages_to_session_db.__get__(agent, AIAgent)
|
||||
)
|
||||
agent._flush_messages_to_session_db_unlocked = (
|
||||
AIAgent._flush_messages_to_session_db_unlocked.__get__(agent, AIAgent)
|
||||
)
|
||||
return agent
|
||||
|
||||
|
||||
def _build_compression_chain(db: SessionDB, chain: list[str]) -> tuple[str, str]:
|
||||
"""Create ``chain[0] -> ... -> chain[-1]`` where every session except the
|
||||
last is compression-ended and the last is live. Returns (root, live_head).
|
||||
"""
|
||||
for i, sid in enumerate(chain):
|
||||
parent = chain[i - 1] if i > 0 else None
|
||||
db.create_session(sid, source="tui", parent_session_id=parent)
|
||||
if i < len(chain) - 1:
|
||||
db.end_session(sid, "compression")
|
||||
return chain[0], chain[-1]
|
||||
|
||||
|
||||
def test_flush_adopts_unique_live_continuation(tmp_path: Path) -> None:
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
try:
|
||||
db.create_session("parent", source="tui")
|
||||
db.append_message("parent", "user", "before split")
|
||||
db.end_session("parent", "compression")
|
||||
db.create_session("child", source="tui", parent_session_id="parent")
|
||||
|
||||
agent = _flush_agent(db, "parent")
|
||||
messages = [{"role": "user", "content": "steered after compression"}]
|
||||
result = agent._flush_messages_to_session_db(messages, [])
|
||||
|
||||
assert result is True, "flush must succeed after adopting the continuation"
|
||||
assert agent.session_id == "child"
|
||||
durable = db.get_messages_as_conversation("child")
|
||||
assert any(
|
||||
m.get("content") == "steered after compression" for m in durable
|
||||
), "the user message must land in the child session, not be lost"
|
||||
# The compression-closed parent stays immutable.
|
||||
parent_rows = db.get_messages_as_conversation("parent")
|
||||
assert not any(
|
||||
m.get("content") == "steered after compression" for m in parent_rows
|
||||
)
|
||||
assert agent._compression_adoption_failed is False
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_flush_adopts_live_head_across_compression_chain(tmp_path: Path) -> None:
|
||||
"""A stale writer behind a chain of >=2 compressions adopts the live head.
|
||||
|
||||
This is the exact lineage from #82001 (`root(compressed) -> mid(compressed)
|
||||
-> tip(live)`) that a depth-1 live-child lookup cannot resolve, because the
|
||||
direct child is itself already compression-ended.
|
||||
"""
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
try:
|
||||
root, head = _build_compression_chain(db, ["root", "mid", "tip"])
|
||||
|
||||
agent = _flush_agent(db, root)
|
||||
messages = [{"role": "user", "content": "steered after double rotation"}]
|
||||
result = agent._flush_messages_to_session_db(messages, [])
|
||||
|
||||
assert result is True, "flush must succeed by adopting the chain head"
|
||||
assert agent.session_id == head, "agent must move to the live chain head"
|
||||
durable = db.get_messages_as_conversation(head)
|
||||
assert any(
|
||||
m.get("content") == "steered after double rotation" for m in durable
|
||||
), "the user message must land in the chain head, not be lost"
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_flush_fails_closed_when_no_continuation(tmp_path: Path) -> None:
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
try:
|
||||
db.create_session("parent", source="tui")
|
||||
db.append_message("parent", "user", "before split")
|
||||
db.end_session("parent", "compression")
|
||||
|
||||
agent = _flush_agent(db, "parent")
|
||||
messages = [{"role": "user", "content": "steered after compression"}]
|
||||
result = agent._flush_messages_to_session_db(messages, [])
|
||||
|
||||
assert result is False, "no continuation -> fail closed (never guess)"
|
||||
assert agent.session_id == "parent", "session id must not change"
|
||||
assert agent._compression_adoption_failed is True
|
||||
assert agent._last_persistence_error_cause == "compression_closed"
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_flush_fails_closed_when_tip_is_stale_closed(tmp_path: Path) -> None:
|
||||
"""The canonical tip walk may land on a stale closed sibling (e.g.
|
||||
``ws_orphan_reap``) — a non-live tip must NOT be adopted; fail closed."""
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
try:
|
||||
db.create_session("parent", source="tui")
|
||||
db.append_message("parent", "user", "before split")
|
||||
db.end_session("parent", "compression")
|
||||
db.create_session("stale", source="tui", parent_session_id="parent")
|
||||
db.end_session("stale", "ws_orphan_reap")
|
||||
|
||||
agent = _flush_agent(db, "parent")
|
||||
messages = [{"role": "user", "content": "steered after compression"}]
|
||||
result = agent._flush_messages_to_session_db(messages, [])
|
||||
|
||||
assert result is False, "non-live tip must fail closed (never adopt stale)"
|
||||
assert agent.session_id == "parent"
|
||||
assert agent._compression_adoption_failed is True
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def test_flush_adopts_exactly_once_no_retry_loop(tmp_path: Path, monkeypatch) -> None:
|
||||
"""Adoption budget: the tip lookup runs at most once per flush, and a
|
||||
second closed-parent write after adoption fails closed instead of looping.
|
||||
"""
|
||||
from hermes_state_errors import CompressionSessionClosedError
|
||||
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
try:
|
||||
_build_compression_chain(db, ["root", "tip"])
|
||||
|
||||
agent = _flush_agent(db, "root")
|
||||
|
||||
tip_calls = {"count": 0}
|
||||
orig_tip = SessionDB.get_compression_tip
|
||||
|
||||
def _counting_tip(self, session_id):
|
||||
tip_calls["count"] += 1
|
||||
return orig_tip(self, session_id)
|
||||
|
||||
monkeypatch.setattr(SessionDB, "get_compression_tip", _counting_tip)
|
||||
|
||||
# Every batch write raises closed — including the post-adoption retry
|
||||
# against the live tip (simulating the tip rotating again mid-flush).
|
||||
def _always_closed(self, *, session_id, messages, **kwargs):
|
||||
raise CompressionSessionClosedError(session_id)
|
||||
|
||||
monkeypatch.setattr(SessionDB, "append_messages_batch", _always_closed)
|
||||
|
||||
messages = [{"role": "user", "content": "steered after compression"}]
|
||||
result = agent._flush_messages_to_session_db(messages, [])
|
||||
|
||||
assert result is False, "second closed-parent write must fail closed"
|
||||
assert tip_calls["count"] == 1, "tip lookup must happen exactly once"
|
||||
assert agent._compression_adoption_failed is True
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Diagnostics: the failure must never read like a disk problem.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_compression_closed_error_classifies_as_compression_closed() -> None:
|
||||
from hermes_state import classify_persistence_error
|
||||
from hermes_state_errors import CompressionSessionClosedError, PERSISTENCE_ERROR_CAUSES
|
||||
|
||||
cause = classify_persistence_error(CompressionSessionClosedError("session-abc"))
|
||||
assert cause == "compression_closed"
|
||||
assert cause in PERSISTENCE_ERROR_CAUSES
|
||||
# String form (post-RPC wrapping) classifies identically.
|
||||
assert (
|
||||
classify_persistence_error(str(CompressionSessionClosedError("session-abc")))
|
||||
== "compression_closed"
|
||||
)
|
||||
|
||||
|
||||
def test_compression_closed_wording_never_mentions_disk() -> None:
|
||||
from hermes_state import classify_persistence_error
|
||||
from hermes_state_errors import CompressionSessionClosedError
|
||||
|
||||
text = AIAgent._format_turn_completion_explanation(
|
||||
"session_persistence_failed",
|
||||
persistence_cause=classify_persistence_error(
|
||||
CompressionSessionClosedError("session-abc")
|
||||
),
|
||||
)
|
||||
assert text, "an abnormal persistence failure must produce an explanation"
|
||||
assert "disk" not in text.lower(), "compression-race message must not blame disk"
|
||||
assert "compression" in text.lower(), "message must name compression rotation"
|
||||
|
||||
|
||||
def test_disk_cause_keeps_disk_guidance() -> None:
|
||||
text = AIAgent._format_turn_completion_explanation(
|
||||
"session_persistence_failed", persistence_cause="disk"
|
||||
)
|
||||
assert "disk" in text.lower() and "free some space" in text.lower(), "real disk failures must keep disk guidance"
|
||||
@@ -0,0 +1,87 @@
|
||||
"""Regression for #36908: the repeated-compression warning must reach the
|
||||
TUI / gateway, not just CLI stdout.
|
||||
|
||||
When a session is compressed >= 2 times, ``compress_context`` warns that
|
||||
accuracy may degrade. That warning used to go through ``_vprint`` (stdout
|
||||
only), so the Ink TUI / Telegram / Discord never saw it — unlike the two
|
||||
other compression warnings in the same module, which route through
|
||||
``_emit_status`` (and store ``_compression_warning`` for late-bound
|
||||
gateway replay). This pins the warning onto the gateway-aware channel.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from hermes_state import SessionDB
|
||||
|
||||
|
||||
def _build_agent_with_db(db: SessionDB, session_id: str, compression_count: int):
|
||||
with patch.dict(os.environ, {"OPENROUTER_API_KEY": "test-key"}):
|
||||
from run_agent import AIAgent
|
||||
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
quiet_mode=True,
|
||||
session_db=db,
|
||||
session_id=session_id,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
|
||||
compressor = MagicMock()
|
||||
compressor.compress.return_value = [
|
||||
{"role": "user", "content": "[CONTEXT COMPACTION] summary"},
|
||||
{"role": "user", "content": "tail"},
|
||||
]
|
||||
compressor.compression_count = compression_count
|
||||
compressor.last_prompt_tokens = 0
|
||||
compressor.last_completion_tokens = 0
|
||||
compressor._last_summary_error = None
|
||||
compressor._last_compress_aborted = False
|
||||
compressor._last_aux_model_failure_model = None
|
||||
compressor._last_aux_model_failure_error = None
|
||||
agent.context_compressor = compressor
|
||||
return agent
|
||||
|
||||
|
||||
def test_repeated_compression_warning_routed_through_emit_status(tmp_path: Path) -> None:
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
sid = "PARENT_36908"
|
||||
db.create_session(sid, source="cli")
|
||||
|
||||
# compression_count == 2 → the "compressed N times" warning should fire.
|
||||
agent = _build_agent_with_db(db, sid, compression_count=2)
|
||||
|
||||
emitted: list[str] = []
|
||||
agent._emit_status = lambda message: emitted.append(message)
|
||||
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
agent._compress_context(messages, "sys", approx_tokens=120_000)
|
||||
|
||||
# The warning reached the gateway-aware channel...
|
||||
assert any("compressed 2 times" in m.lower() for m in emitted), (
|
||||
f"repeated-compression warning not emitted via _emit_status: {emitted}"
|
||||
)
|
||||
# ...and was stored for late-bound gateway status_callback replay.
|
||||
assert "compressed 2 times" in (getattr(agent, "_compression_warning", "") or "").lower()
|
||||
|
||||
|
||||
def test_no_warning_below_threshold(tmp_path: Path) -> None:
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
sid = "PARENT_36908_ONCE"
|
||||
db.create_session(sid, source="cli")
|
||||
|
||||
# compression_count == 1 → no repeated-compression warning.
|
||||
agent = _build_agent_with_db(db, sid, compression_count=1)
|
||||
emitted: list[str] = []
|
||||
agent._emit_status = lambda message: emitted.append(message)
|
||||
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
agent._compress_context(messages, "sys", approx_tokens=120_000)
|
||||
|
||||
assert not any("compressed" in m.lower() and "times" in m.lower() for m in emitted)
|
||||
@@ -0,0 +1,451 @@
|
||||
"""Tests for _check_compression_model_feasibility() — warns when the
|
||||
auxiliary compression model's context is smaller than the main model's
|
||||
compression threshold.
|
||||
|
||||
Two-phase design:
|
||||
1. __init__ → runs the check, prints via _vprint (CLI), stores warning
|
||||
2. run_conversation (first call) → replays stored warning through
|
||||
status_callback (gateway platforms)
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from run_agent import AIAgent
|
||||
from agent.context_compressor import ContextCompressor
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _stable_aux_provider_config():
|
||||
"""Keep feasibility tests independent from the developer's config.yaml."""
|
||||
with patch(
|
||||
"agent.auxiliary_client._resolve_task_provider_model",
|
||||
return_value=("auto", None, None, None, None),
|
||||
):
|
||||
yield
|
||||
|
||||
|
||||
def _make_agent(
|
||||
*,
|
||||
compression_enabled: bool = True,
|
||||
threshold_percent: float = 0.50,
|
||||
main_context: int = 200_000,
|
||||
) -> AIAgent:
|
||||
"""Build a minimal AIAgent with a compressor, skipping __init__."""
|
||||
agent = AIAgent.__new__(AIAgent)
|
||||
agent.model = "test-main-model"
|
||||
agent.provider = "openrouter"
|
||||
agent.base_url = "https://openrouter.ai/api/v1"
|
||||
agent.api_key = "sk-test"
|
||||
agent.api_mode = "chat_completions"
|
||||
agent.quiet_mode = True
|
||||
agent.log_prefix = ""
|
||||
agent.compression_enabled = compression_enabled
|
||||
agent._print_fn = None
|
||||
agent.suppress_status_output = False
|
||||
agent._stream_consumers = []
|
||||
agent._executing_tools = False
|
||||
agent._mute_post_response = False
|
||||
agent.status_callback = None
|
||||
agent.tool_progress_callback = None
|
||||
agent._compression_warning = None
|
||||
agent._aux_compression_context_length_config = None
|
||||
agent._custom_providers = []
|
||||
agent.tools = []
|
||||
|
||||
compressor = MagicMock(spec=ContextCompressor)
|
||||
compressor.context_length = main_context
|
||||
compressor.threshold_tokens = int(main_context * threshold_percent)
|
||||
compressor.summary_target_ratio = 0.20
|
||||
compressor.tail_token_budget = int(
|
||||
compressor.threshold_tokens * compressor.summary_target_ratio
|
||||
)
|
||||
agent.context_compressor = compressor
|
||||
|
||||
return agent
|
||||
|
||||
|
||||
@pytest.mark.parametrize("main_context,aux_context", [(1_000_000, 512_000), (400_000, 80_000)])
|
||||
def test_aux_sync_keeps_lean_tail_policy(main_context, aux_context):
|
||||
"""Lowering only the trigger must not change window-relative retention."""
|
||||
agent = _make_agent(main_context=main_context)
|
||||
compressor = agent.context_compressor = ContextCompressor(
|
||||
"test-main-model", config_context_length=main_context,
|
||||
threshold_percent=0.85, quiet_mode=True,
|
||||
)
|
||||
before = compressor.tail_token_budget
|
||||
agent._emit_status = lambda message: None
|
||||
client = MagicMock(base_url="http://localhost/v1", api_key="test-key")
|
||||
with patch("agent.auxiliary_client.get_text_auxiliary_client", return_value=(client, "aux")), \
|
||||
patch("agent.model_metadata.get_model_context_length", return_value=aux_context):
|
||||
agent._check_compression_model_feasibility()
|
||||
assert compressor.threshold_tokens == aux_context
|
||||
assert compressor.tail_token_budget == before
|
||||
# Repeated feasibility and subsequent model recalibration retain policy.
|
||||
agent._check_compression_model_feasibility()
|
||||
assert compressor.tail_token_budget == before
|
||||
compressor.update_model("test-main-model", context_length=main_context)
|
||||
assert compressor.tail_token_budget == before
|
||||
|
||||
|
||||
def test_aux_sync_legacy_tail_follows_lowered_threshold():
|
||||
"""Explicit legacy retention follows the current trigger, not its old cache."""
|
||||
agent = _make_agent(main_context=1_000_000)
|
||||
compressor = agent.context_compressor = ContextCompressor(
|
||||
"test-main-model", config_context_length=1_000_000,
|
||||
threshold_percent=0.85, tail_mode="legacy", quiet_mode=True,
|
||||
)
|
||||
before = compressor.tail_token_budget
|
||||
agent._emit_status = lambda message: None
|
||||
client = MagicMock(base_url="http://localhost/v1", api_key="test-key")
|
||||
with patch("agent.auxiliary_client.get_text_auxiliary_client", return_value=(client, "aux")), \
|
||||
patch("agent.model_metadata.get_model_context_length", return_value=512_000):
|
||||
agent._check_compression_model_feasibility()
|
||||
assert compressor.threshold_tokens == 512_000
|
||||
assert compressor.tail_token_budget < before
|
||||
assert compressor.tail_token_budget == int(compressor.threshold_tokens * compressor.summary_target_ratio)
|
||||
|
||||
|
||||
# ── Core warning logic ──────────────────────────────────────────────
|
||||
|
||||
|
||||
@patch("agent.model_metadata.get_model_context_length", return_value=80_000)
|
||||
@patch("agent.auxiliary_client.get_text_auxiliary_client")
|
||||
def test_auto_corrects_threshold_when_aux_context_below_threshold(mock_get_client, mock_ctx_len):
|
||||
"""Auto-correction: aux >= 64K floor but < threshold → lower threshold
|
||||
to aux_context so compression still works this session."""
|
||||
agent = _make_agent(main_context=200_000, threshold_percent=0.50)
|
||||
# threshold = 100,000 — aux has 80,000 (above 64K floor, below threshold)
|
||||
mock_client = MagicMock()
|
||||
mock_client.base_url = "https://openrouter.ai/api/v1"
|
||||
mock_client.api_key = "sk-aux"
|
||||
mock_get_client.return_value = (mock_client, "google/gemini-3-flash-preview")
|
||||
|
||||
messages = []
|
||||
agent._emit_status = lambda msg: messages.append(msg)
|
||||
|
||||
agent._check_compression_model_feasibility()
|
||||
|
||||
assert len(messages) == 1
|
||||
assert "Compression model" in messages[0]
|
||||
assert "80,000" in messages[0] # aux context
|
||||
assert "100,000" in messages[0] # old threshold
|
||||
assert "Auto-lowered" in messages[0]
|
||||
# Actionable persistence guidance included
|
||||
assert "config.yaml" in messages[0]
|
||||
assert "auxiliary:" in messages[0]
|
||||
assert "compression:" in messages[0]
|
||||
# 200K main is under the 512K small-context limit and 80K/200K = 40% sits
|
||||
# below the 75% floor — a `threshold:` suggestion would be raised back to
|
||||
# 75% and ignored (#67422), so the message must not offer one and must
|
||||
# explain the recomputed trigger instead (0.75 * 200K = 150K).
|
||||
assert "threshold:" not in messages[0]
|
||||
assert "150,000" in messages[0]
|
||||
# Warning stored for gateway replay
|
||||
assert agent._compression_warning is not None
|
||||
# Threshold on the live compressor was actually lowered to aux_context.
|
||||
assert agent.context_compressor.threshold_tokens == 80_000
|
||||
# Every threshold-derived budget must move with it. Keeping the original
|
||||
# 20K tail here would protect 25% of the lowered threshold instead of the
|
||||
# configured 20%, and larger real-world mismatches can make the tail's 1.5x
|
||||
# soft ceiling wider than the entire compression trigger.
|
||||
assert agent.context_compressor.tail_token_budget == 16_000
|
||||
|
||||
|
||||
@patch("agent.model_metadata.get_model_context_length", return_value=32_768)
|
||||
@patch("agent.auxiliary_client.get_text_auxiliary_client")
|
||||
def test_rejects_aux_below_minimum_context(mock_get_client, mock_ctx_len):
|
||||
"""Hard floor: aux context < MINIMUM_CONTEXT_LENGTH (64K) → session
|
||||
refuses to start (ValueError), mirroring the main-model rejection."""
|
||||
agent = _make_agent(main_context=200_000, threshold_percent=0.50)
|
||||
mock_client = MagicMock()
|
||||
mock_client.base_url = "https://openrouter.ai/api/v1"
|
||||
mock_client.api_key = "sk-aux"
|
||||
mock_get_client.return_value = (mock_client, "tiny-aux-model")
|
||||
|
||||
agent._emit_status = lambda msg: None
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
agent._check_compression_model_feasibility()
|
||||
|
||||
err = str(exc_info.value)
|
||||
assert "tiny-aux-model" in err
|
||||
assert "32,768" in err
|
||||
assert "64,000" in err
|
||||
assert "below the minimum" in err
|
||||
|
||||
|
||||
|
||||
|
||||
def test_feasibility_check_passes_live_main_runtime():
|
||||
"""Compression feasibility should probe using the live session runtime."""
|
||||
agent = _make_agent(main_context=200_000, threshold_percent=0.50)
|
||||
agent.model = "gpt-5.4"
|
||||
agent.provider = "openai-codex"
|
||||
agent.base_url = "https://chatgpt.com/backend-api/codex"
|
||||
agent.api_key = "codex-token"
|
||||
agent.api_mode = "codex_responses"
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.base_url = "https://chatgpt.com/backend-api/codex"
|
||||
mock_client.api_key = "codex-token"
|
||||
|
||||
with patch("agent.auxiliary_client.get_text_auxiliary_client", return_value=(mock_client, "gpt-5.4")) as mock_get_client, \
|
||||
patch("agent.model_metadata.get_model_context_length", return_value=200_000):
|
||||
agent._emit_status = lambda msg: None
|
||||
agent._check_compression_model_feasibility()
|
||||
|
||||
mock_get_client.assert_called_once_with(
|
||||
"compression",
|
||||
main_runtime={
|
||||
"model": "gpt-5.4",
|
||||
"provider": "openai-codex",
|
||||
"base_url": "https://chatgpt.com/backend-api/codex",
|
||||
"api_key": "codex-token",
|
||||
"api_mode": "codex_responses",
|
||||
"auth_mode": "",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@patch("agent.model_metadata.get_model_context_length", return_value=1_000_000)
|
||||
@patch("agent.auxiliary_client.get_text_auxiliary_client")
|
||||
def test_feasibility_check_passes_config_context_length(mock_get_client, mock_ctx_len):
|
||||
"""auxiliary.compression.context_length from config is forwarded to
|
||||
get_model_context_length so custom endpoints that lack /models still
|
||||
report the correct context window (fixes #8499)."""
|
||||
agent = _make_agent(main_context=200_000, threshold_percent=0.85)
|
||||
agent._aux_compression_context_length_config = 1_000_000
|
||||
mock_client = MagicMock()
|
||||
mock_client.base_url = "http://custom-endpoint:8080/v1"
|
||||
mock_client.api_key = "sk-custom"
|
||||
mock_get_client.return_value = (mock_client, "custom/big-model")
|
||||
|
||||
agent._emit_status = lambda msg: None
|
||||
agent._check_compression_model_feasibility()
|
||||
|
||||
mock_ctx_len.assert_called_once_with(
|
||||
"custom/big-model",
|
||||
base_url="http://custom-endpoint:8080/v1",
|
||||
api_key="sk-custom",
|
||||
config_context_length=1_000_000,
|
||||
provider="openrouter",
|
||||
custom_providers=[],
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
def test_init_feasibility_check_uses_aux_context_override_from_config():
|
||||
"""Lazy feasibility check should cache and forward auxiliary.compression.context_length.
|
||||
|
||||
NB: feasibility check is deferred from AIAgent.__init__ to the first
|
||||
actual compression attempt (saves ~400ms cold startup on short sessions
|
||||
that never trigger compression). The test drives the check explicitly
|
||||
via ``agent._check_compression_model_feasibility()`` to assert the
|
||||
config-override threading.
|
||||
"""
|
||||
|
||||
class _StubCompressor:
|
||||
def __init__(self, *args, **kwargs):
|
||||
self.context_length = 200_000
|
||||
self.threshold_tokens = 100_000
|
||||
self.threshold_percent = 0.50
|
||||
|
||||
def get_tool_schemas(self):
|
||||
return []
|
||||
|
||||
def on_session_start(self, *args, **kwargs):
|
||||
return None
|
||||
|
||||
cfg = {
|
||||
"auxiliary": {
|
||||
"compression": {
|
||||
"context_length": 1_000_000,
|
||||
},
|
||||
},
|
||||
}
|
||||
mock_client = MagicMock()
|
||||
mock_client.base_url = "http://custom-endpoint:8080/v1"
|
||||
mock_client.api_key = "sk-custom"
|
||||
|
||||
with (
|
||||
patch("hermes_cli.config.load_config", return_value=cfg), patch("hermes_cli.config.load_config_readonly", return_value=cfg),
|
||||
patch("model_tools.get_tool_definitions", return_value=[]),
|
||||
patch("model_tools.check_toolset_requirements", return_value={}),
|
||||
patch("agent.process_bootstrap.OpenAI"),
|
||||
patch("agent.agent_init.ContextCompressor", new=_StubCompressor),
|
||||
patch("agent.auxiliary_client.get_text_auxiliary_client", return_value=(mock_client, "custom/big-model")),
|
||||
patch("agent.model_metadata.get_model_context_length", return_value=1_000_000) as mock_ctx_len,
|
||||
):
|
||||
agent = AIAgent(
|
||||
api_key="test-key-1234567890",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
|
||||
# Config override is captured eagerly in __init__ (still needed
|
||||
# because the threshold-derivation logic at construction time
|
||||
# consults it).
|
||||
assert agent._aux_compression_context_length_config == 1_000_000
|
||||
|
||||
# The expensive feasibility probe is deferred. Drive it manually
|
||||
# to validate the call shape still forwards the override correctly.
|
||||
agent._check_compression_model_feasibility()
|
||||
|
||||
mock_ctx_len.assert_called_once_with(
|
||||
"custom/big-model",
|
||||
base_url="http://custom-endpoint:8080/v1",
|
||||
api_key="sk-custom",
|
||||
config_context_length=1_000_000,
|
||||
provider="",
|
||||
custom_providers=[],
|
||||
)
|
||||
|
||||
|
||||
@patch("agent.auxiliary_client.get_text_auxiliary_client")
|
||||
def test_warns_when_no_auxiliary_provider(mock_get_client):
|
||||
"""Warning emitted when no auxiliary provider is configured."""
|
||||
agent = _make_agent()
|
||||
mock_get_client.return_value = (None, None)
|
||||
|
||||
messages = []
|
||||
agent._emit_status = lambda msg: messages.append(msg)
|
||||
|
||||
agent._check_compression_model_feasibility()
|
||||
|
||||
assert len(messages) == 1
|
||||
assert "No auxiliary LLM provider" in messages[0]
|
||||
assert agent._compression_warning is not None
|
||||
|
||||
|
||||
def test_no_unavailable_warning_when_configured_fallback_chain_resolves():
|
||||
"""Primary compression provider can be down if configured fallback works."""
|
||||
agent = _make_agent(main_context=200_000, threshold_percent=0.50)
|
||||
fallback_client = MagicMock()
|
||||
fallback_client.base_url = "https://chatgpt.com/backend-api/codex"
|
||||
fallback_client.api_key = "codex-oauth-token"
|
||||
|
||||
messages = []
|
||||
agent._emit_status = lambda msg: messages.append(msg)
|
||||
|
||||
with patch(
|
||||
"agent.auxiliary_client._resolve_task_provider_model",
|
||||
return_value=("ollama-cloud", "deepseek-v4-flash:cloud", None, None, None),
|
||||
), patch(
|
||||
"agent.auxiliary_client.get_text_auxiliary_client",
|
||||
return_value=(None, None),
|
||||
), patch(
|
||||
"agent.auxiliary_client._try_configured_fallback_for_unavailable_client",
|
||||
return_value=(fallback_client, "gpt-5.4-mini", "fallback_chain[0](openai-codex)"),
|
||||
) as mock_fallback, patch(
|
||||
"agent.model_metadata.get_model_context_length",
|
||||
return_value=200_000,
|
||||
) as mock_ctx_len:
|
||||
agent._check_compression_model_feasibility()
|
||||
|
||||
assert messages == []
|
||||
assert agent._compression_warning is None
|
||||
mock_fallback.assert_called_once_with("compression", "ollama-cloud")
|
||||
mock_ctx_len.assert_called_once()
|
||||
assert mock_ctx_len.call_args.args == ("gpt-5.4-mini",)
|
||||
assert mock_ctx_len.call_args.kwargs["provider"] == "openai-codex"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ── Two-phase: __init__ + run_conversation replay ───────────────────
|
||||
|
||||
|
||||
@patch("agent.model_metadata.get_model_context_length", return_value=80_000)
|
||||
@patch("agent.auxiliary_client.get_text_auxiliary_client")
|
||||
def test_warning_stored_for_gateway_replay(mock_get_client, mock_ctx_len):
|
||||
"""__init__ stores the warning; _replay sends it through status_callback."""
|
||||
agent = _make_agent(main_context=200_000, threshold_percent=0.50)
|
||||
mock_client = MagicMock()
|
||||
mock_client.base_url = "https://openrouter.ai/api/v1"
|
||||
mock_client.api_key = "sk-aux"
|
||||
mock_get_client.return_value = (mock_client, "google/gemini-3-flash-preview")
|
||||
|
||||
# Phase 1: __init__ — _emit_status prints (CLI) but callback is None
|
||||
vprint_messages = []
|
||||
agent._emit_status = lambda msg: vprint_messages.append(msg)
|
||||
agent._check_compression_model_feasibility()
|
||||
|
||||
assert len(vprint_messages) == 1 # CLI got it
|
||||
assert agent._compression_warning is not None # stored for replay
|
||||
|
||||
# Phase 2: gateway wires callback post-init, then run_conversation replays
|
||||
callback_events = []
|
||||
agent.status_callback = lambda ev, msg: callback_events.append((ev, msg))
|
||||
agent._replay_compression_warning()
|
||||
|
||||
assert any(
|
||||
ev == "lifecycle" and "Auto-lowered" in msg
|
||||
for ev, msg in callback_events
|
||||
)
|
||||
|
||||
|
||||
@patch("agent.model_metadata.get_model_context_length", return_value=200_000)
|
||||
@patch("agent.auxiliary_client.get_text_auxiliary_client")
|
||||
def test_no_replay_when_no_warning(mock_get_client, mock_ctx_len):
|
||||
"""_replay_compression_warning is a no-op when there's no stored warning."""
|
||||
agent = _make_agent(main_context=200_000, threshold_percent=0.50)
|
||||
mock_client = MagicMock()
|
||||
mock_client.base_url = "https://openrouter.ai/api/v1"
|
||||
mock_client.api_key = "sk-aux"
|
||||
mock_get_client.return_value = (mock_client, "big-model")
|
||||
|
||||
agent._emit_status = lambda msg: None
|
||||
agent._check_compression_model_feasibility()
|
||||
|
||||
assert agent._compression_warning is None
|
||||
|
||||
callback_events = []
|
||||
agent.status_callback = lambda ev, msg: callback_events.append((ev, msg))
|
||||
agent._replay_compression_warning()
|
||||
|
||||
assert len(callback_events) == 0
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ── #67422: threshold suggestion must survive the small-context floor ────────
|
||||
|
||||
|
||||
|
||||
|
||||
@patch("agent.model_metadata.get_model_context_length", return_value=300_000)
|
||||
@patch("agent.auxiliary_client.get_text_auxiliary_client")
|
||||
def test_threshold_suggestion_kept_for_large_context_main(mock_get_client, mock_ctx_len):
|
||||
"""Main window >= 512K has no floor — any suggestion is honored, so the
|
||||
`threshold:` option stays even below 75%."""
|
||||
agent = _make_agent(main_context=1_000_000, threshold_percent=0.50)
|
||||
# threshold = 500,000 — aux has 300,000
|
||||
mock_client = MagicMock()
|
||||
mock_client.base_url = "https://openrouter.ai/api/v1"
|
||||
mock_client.api_key = "sk-aux"
|
||||
mock_get_client.return_value = (mock_client, "google/gemini-3-flash-preview")
|
||||
|
||||
messages = []
|
||||
agent._emit_status = lambda msg: messages.append(msg)
|
||||
|
||||
agent._check_compression_model_feasibility()
|
||||
|
||||
assert len(messages) == 1
|
||||
assert "threshold: 0.30" in messages[0]
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,310 @@
|
||||
"""Lock-contended compression no-ops must soft-DEFER, never exhaust (#49874).
|
||||
|
||||
On main before this fix, nothing on the automatic compression paths consumed
|
||||
the #69870 lock-skip signal (``agent._compression_skipped_due_to_lock``):
|
||||
|
||||
* a lock-loser preflight/pre-API no-op counted as "insufficient progress",
|
||||
* the oversized request went to the provider anyway, and
|
||||
* the lock-contended 413/overflow retry burned ``compression_attempts`` to
|
||||
the cap and returned ``compression_exhausted`` — which the gateway answers
|
||||
with a full session auto-reset (#9893/#35809).
|
||||
|
||||
A temporary lock defer misclassified as exhaustion == session wipe.
|
||||
|
||||
These tests pin the fix: when a compression pass returns its input unchanged
|
||||
AND the type-pinned lock-skip flag is set, the attempt is refunded and the
|
||||
turn ends (when it cannot proceed) with a soft ``compression_deferred``
|
||||
result distinct from ``compression_exhausted``.
|
||||
|
||||
Salvaged from PR #49874 (@helix4u), rebuilt on the landed #69870 signal.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from agent.conversation_compression import compression_skipped_due_to_lock
|
||||
from run_agent import AIAgent
|
||||
import run_agent
|
||||
|
||||
|
||||
LOCK_HOLDER = "pid=4242:tid=1:agent=deadbeef:nonce=abcd1234"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers (mirrors tests/agent/test_413_compression.py)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _no_compression_sleep(monkeypatch):
|
||||
import time as _time
|
||||
|
||||
monkeypatch.setattr(_time, "sleep", lambda *_a, **_k: None)
|
||||
from agent import retry_utils as _retry_utils
|
||||
monkeypatch.setattr(_retry_utils, "jittered_backoff", lambda *a, **k: 0.0)
|
||||
|
||||
|
||||
def _make_tool_defs(*names: str) -> list:
|
||||
return [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": n,
|
||||
"description": f"{n} tool",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
}
|
||||
for n in names
|
||||
]
|
||||
|
||||
|
||||
def _mock_response(content="Hello", finish_reason="stop"):
|
||||
msg = SimpleNamespace(
|
||||
content=content,
|
||||
tool_calls=None,
|
||||
reasoning_content=None,
|
||||
reasoning=None,
|
||||
)
|
||||
choice = SimpleNamespace(message=msg, finish_reason=finish_reason)
|
||||
resp = SimpleNamespace(choices=[choice], model="test/model")
|
||||
resp.usage = None
|
||||
return resp
|
||||
|
||||
|
||||
def _make_413_error(message="Request entity too large"):
|
||||
err = Exception(message)
|
||||
err.status_code = 413
|
||||
return err
|
||||
|
||||
|
||||
def _make_overflow_error():
|
||||
return Exception(
|
||||
"Error code: 400 - {'type': 'error', 'error': {'type': "
|
||||
"'invalid_request_error', 'message': 'prompt is too long: "
|
||||
"233153 tokens > 200000 maximum'}}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def agent():
|
||||
with (
|
||||
patch("model_tools.get_tool_definitions", return_value=_make_tool_defs("web_search")),
|
||||
patch("model_tools.check_toolset_requirements", return_value={}),
|
||||
patch("agent.process_bootstrap.OpenAI"),
|
||||
):
|
||||
a = AIAgent(
|
||||
api_key="test-key-1234567890",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
a.client = MagicMock()
|
||||
a._cached_system_prompt = "You are helpful."
|
||||
a._use_prompt_caching = False
|
||||
a.tool_delay = 0
|
||||
a.compression_enabled = True
|
||||
a.save_trajectories = False
|
||||
return a
|
||||
|
||||
|
||||
_PREFILL = [
|
||||
{"role": "user", "content": "previous question"},
|
||||
{"role": "assistant", "content": "previous answer"},
|
||||
]
|
||||
|
||||
|
||||
def _lock_skipping_compress(agent, *, holder=LOCK_HOLDER):
|
||||
"""A compress double that no-ops because 'another path holds the lock'.
|
||||
|
||||
Mirrors the real ``compress_context`` lock-contended abort: returns the
|
||||
INPUT list object unchanged and sets the #69870 lock-skip signal.
|
||||
"""
|
||||
|
||||
def _compress(messages, _system_message, **_kwargs):
|
||||
agent._compression_skipped_due_to_lock = holder
|
||||
return messages, "You are helpful."
|
||||
|
||||
return _compress
|
||||
|
||||
|
||||
def _plain_noop_compress(agent):
|
||||
"""A compress double that no-ops WITHOUT lock contention (real no-progress)."""
|
||||
|
||||
def _compress(messages, _system_message, **_kwargs):
|
||||
agent._compression_skipped_due_to_lock = None
|
||||
return messages, "You are helpful."
|
||||
|
||||
return _compress
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Type-pinned signal read (MagicMock test-double immunity)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestLockSkipSignalTypePin:
|
||||
def test_true_and_holder_string_are_lock_skips(self):
|
||||
a = SimpleNamespace(_compression_skipped_due_to_lock=True)
|
||||
assert compression_skipped_due_to_lock(a) is True
|
||||
a = SimpleNamespace(_compression_skipped_due_to_lock=LOCK_HOLDER)
|
||||
assert compression_skipped_due_to_lock(a) is True
|
||||
|
||||
def test_none_and_missing_are_not_lock_skips(self):
|
||||
assert compression_skipped_due_to_lock(
|
||||
SimpleNamespace(_compression_skipped_due_to_lock=None)
|
||||
) is False
|
||||
assert compression_skipped_due_to_lock(SimpleNamespace()) is False
|
||||
|
||||
def test_magicmock_agent_auto_attribute_is_not_a_lock_skip(self):
|
||||
"""MagicMock agents auto-create truthy attributes; bare truthiness
|
||||
would hijack every mocked agent in sibling suites into the lock-skip
|
||||
branch (the #69870 × #69840 incident). The read must be type-pinned."""
|
||||
assert compression_skipped_due_to_lock(MagicMock()) is False
|
||||
|
||||
def test_truthy_non_true_non_str_values_are_not_lock_skips(self):
|
||||
for junk in (1, 1.0, ["holder"], {"holder": True}, object(), MagicMock()):
|
||||
a = SimpleNamespace(_compression_skipped_due_to_lock=junk)
|
||||
assert compression_skipped_due_to_lock(a) is False, junk
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 413 handler: lock-contended no-op → soft defer, no exhaustion
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestLockContended413Defer:
|
||||
def test_lock_contended_413_returns_compression_deferred(self, agent):
|
||||
"""A 413 whose compression pass lost the lock must end the turn as a
|
||||
soft ``compression_deferred`` — never ``compression_exhausted``."""
|
||||
agent.client.chat.completions.create.side_effect = _make_413_error()
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
agent, "_compress_context",
|
||||
side_effect=_lock_skipping_compress(agent),
|
||||
) as mock_compress,
|
||||
patch.object(agent, "_persist_session"),
|
||||
patch.object(agent, "_save_trajectory"),
|
||||
patch.object(agent, "_cleanup_task_resources"),
|
||||
):
|
||||
result = agent.run_conversation("hello", conversation_history=list(_PREFILL))
|
||||
|
||||
mock_compress.assert_called_once()
|
||||
assert result.get("compression_deferred") is True
|
||||
assert not result.get("compression_exhausted")
|
||||
# Soft defer: transient, retry-next-message semantics — the gateway
|
||||
# persists the user turn (failed=False) and never auto-resets.
|
||||
assert result.get("failed") is False
|
||||
assert result.get("completed") is False
|
||||
assert result.get("partial") is True
|
||||
|
||||
|
||||
def test_unconfirmed_lock_skip_true_also_defers(self, agent):
|
||||
"""``_compression_skipped_due_to_lock = True`` (holder unconfirmed —
|
||||
``try_acquire`` swallowed a sqlite error) is still a lock skip."""
|
||||
agent.client.chat.completions.create.side_effect = _make_413_error()
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
agent, "_compress_context",
|
||||
side_effect=_lock_skipping_compress(agent, holder=True),
|
||||
),
|
||||
patch.object(agent, "_persist_session"),
|
||||
patch.object(agent, "_save_trajectory"),
|
||||
patch.object(agent, "_cleanup_task_resources"),
|
||||
):
|
||||
result = agent.run_conversation("hello", conversation_history=list(_PREFILL))
|
||||
|
||||
assert result.get("compression_deferred") is True
|
||||
assert not result.get("compression_exhausted")
|
||||
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pre-API gate: a lock-skipped pass must not burn the shared attempt budget
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPreApiLockDeferDoesNotBurnBudget:
|
||||
def test_lock_loser_turn_recovers_after_lock_release(self, agent):
|
||||
"""End-to-end shape of the live bug at cap=1:
|
||||
|
||||
1. Pre-API pressure gate fires; the compression pass loses the lock
|
||||
(no-op + lock-skip flag). Pre-fix this burned the single shared
|
||||
attempt.
|
||||
2. The oversized request goes to the provider → 413.
|
||||
3. The 413 handler compresses again — the lock has been released and
|
||||
the pass now succeeds — and the retry completes.
|
||||
|
||||
Pre-fix, step 3 found ``compression_attempts`` already at the cap and
|
||||
returned ``compression_exhausted`` → gateway session wipe. The defer
|
||||
refund keeps the budget intact for the provider-proven retry.
|
||||
"""
|
||||
agent.max_compression_attempts = 1
|
||||
# Compressor stub: pressure only on the fully-assembled request
|
||||
# (pre-API site); the turn-context preflight stands down via the
|
||||
# cheap-gate (small message count) and low turn-context estimate.
|
||||
agent.context_compressor = SimpleNamespace(
|
||||
protect_first_n=3,
|
||||
protect_last_n=20,
|
||||
threshold_tokens=100_000,
|
||||
context_length=1_000_000,
|
||||
last_prompt_tokens=0,
|
||||
should_compress=lambda t: t >= 100_000,
|
||||
should_defer_preflight_to_real_usage=lambda _t: False,
|
||||
get_active_compression_failure_cooldown=lambda: None,
|
||||
)
|
||||
|
||||
agent.client.chat.completions.create.side_effect = [
|
||||
_make_413_error(),
|
||||
_mock_response(content="Recovered after lock release"),
|
||||
]
|
||||
|
||||
compress_calls = []
|
||||
|
||||
def _lock_then_success(messages, _system_message, **_kwargs):
|
||||
compress_calls.append(len(messages))
|
||||
if len(compress_calls) == 1:
|
||||
# Lock loser: no-op + #69870 signal.
|
||||
agent._compression_skipped_due_to_lock = LOCK_HOLDER
|
||||
return messages, "You are helpful."
|
||||
# Lock released: real compaction (entry clears the signal).
|
||||
agent._compression_skipped_due_to_lock = None
|
||||
return (
|
||||
[{"role": "user", "content": "hello"}],
|
||||
"You are helpful.",
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"agent.turn_context.estimate_request_tokens_rough",
|
||||
return_value=10,
|
||||
),
|
||||
patch(
|
||||
"agent.model_metadata.estimate_request_tokens_rough",
|
||||
return_value=500_000,
|
||||
),
|
||||
patch(
|
||||
"agent.model_metadata.estimate_messages_tokens_rough",
|
||||
return_value=500_000,
|
||||
),
|
||||
patch.object(agent, "_compress_context", side_effect=_lock_then_success),
|
||||
patch.object(agent, "_persist_session"),
|
||||
patch.object(agent, "_save_trajectory"),
|
||||
patch.object(agent, "_cleanup_task_resources"),
|
||||
):
|
||||
result = agent.run_conversation("hello", conversation_history=list(_PREFILL))
|
||||
|
||||
# Pass 1: pre-API (lock defer, refunded). Pass 2: 413 handler
|
||||
# (succeeds within the cap because the defer did not count).
|
||||
assert len(compress_calls) == 2
|
||||
assert result.get("completed") is True
|
||||
assert result["final_response"] == "Recovered after lock release"
|
||||
assert not result.get("compression_exhausted")
|
||||
assert not result.get("compression_deferred")
|
||||
@@ -0,0 +1,533 @@
|
||||
"""Tests for context compression persistence in the gateway.
|
||||
|
||||
Verifies that when context compression fires during run_conversation(),
|
||||
the compressed messages are properly persisted to both SQLite (via the
|
||||
agent) and JSONL (via the gateway).
|
||||
|
||||
Bug scenario (pre-fix):
|
||||
1. Gateway loads 200-message history, passes to agent
|
||||
2. Agent's run_conversation() compresses to ~30 messages mid-run
|
||||
3. _compress_context() resets _last_flushed_db_idx = 0
|
||||
4. On exit, _flush_messages_to_session_db() calculates:
|
||||
flush_from = max(len(conversation_history=200), _last_flushed_db_idx=0) = 200
|
||||
5. messages[200:] is empty (only ~30 messages after compression)
|
||||
6. Nothing written to new session's SQLite — compressed context lost
|
||||
7. Gateway's history_offset was still 200, producing empty new_messages
|
||||
8. Fallback wrote only user/assistant pair — summary lost
|
||||
"""
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Part 1: Agent-side — _flush_messages_to_session_db after compression
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestFlushAfterCompression:
|
||||
"""Verify that compressed messages are flushed to the new session's SQLite
|
||||
even when conversation_history (from the original session) is longer than
|
||||
the compressed messages list."""
|
||||
|
||||
def _make_agent(self, session_db):
|
||||
with patch.dict(os.environ, {"OPENROUTER_API_KEY": "test-key"}):
|
||||
from run_agent import AIAgent
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
quiet_mode=True,
|
||||
session_db=session_db,
|
||||
session_id="original-session",
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
return agent
|
||||
|
||||
def test_flush_after_compression_with_long_history(self):
|
||||
"""The actual bug: conversation_history longer than compressed messages.
|
||||
|
||||
Before the fix, flush_from = max(len(conversation_history), 0) = 200,
|
||||
but messages only has ~30 entries, so messages[200:] is empty.
|
||||
After the fix, conversation_history is cleared to None after compression,
|
||||
so flush_from = max(0, 0) = 0, and ALL compressed messages are written.
|
||||
"""
|
||||
from hermes_state import SessionDB
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db_path = Path(tmpdir) / "test.db"
|
||||
db = SessionDB(db_path=db_path)
|
||||
|
||||
agent = self._make_agent(db)
|
||||
|
||||
# Simulate the original long history (200 messages)
|
||||
original_history = [
|
||||
{"role": "user" if i % 2 == 0 else "assistant",
|
||||
"content": f"message {i}"}
|
||||
for i in range(200)
|
||||
]
|
||||
|
||||
# First, flush original messages to the original session
|
||||
agent._flush_messages_to_session_db(original_history, [])
|
||||
original_rows = db.get_messages("original-session")
|
||||
assert len(original_rows) == 200
|
||||
|
||||
# Now simulate compression: new session, reset idx, shorter messages
|
||||
agent.session_id = "compressed-session"
|
||||
db.create_session(session_id="compressed-session", source="test")
|
||||
agent._last_flushed_db_idx = 0
|
||||
|
||||
# The compressed messages (summary + tail + new turn)
|
||||
compressed_messages = [
|
||||
{"role": "user", "content": "[CONTEXT COMPACTION] Summary of work..."},
|
||||
{"role": "user", "content": "What should we do next?"},
|
||||
{"role": "assistant", "content": "Let me check..."},
|
||||
{"role": "user", "content": "new question"},
|
||||
{"role": "assistant", "content": "new answer"},
|
||||
]
|
||||
|
||||
# THE BUG: passing the original history as conversation_history
|
||||
# causes flush_from = max(200, 0) = 200, skipping everything.
|
||||
# After the fix, conversation_history should be None.
|
||||
agent._flush_messages_to_session_db(compressed_messages, None)
|
||||
|
||||
new_rows = db.get_messages("compressed-session")
|
||||
assert len(new_rows) == 5, (
|
||||
f"Expected 5 compressed messages in new session, got {len(new_rows)}. "
|
||||
f"Compression persistence bug: messages not written to SQLite."
|
||||
)
|
||||
|
||||
def test_flush_with_stale_history_loses_messages(self):
|
||||
"""Stale conversation_history no longer causes data loss."""
|
||||
from hermes_state import SessionDB
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db_path = Path(tmpdir) / "test.db"
|
||||
db = SessionDB(db_path=db_path)
|
||||
|
||||
agent = self._make_agent(db)
|
||||
|
||||
# Simulate compression reset
|
||||
agent.session_id = "new-session"
|
||||
db.create_session(session_id="new-session", source="test")
|
||||
agent._last_flushed_db_idx = 0
|
||||
|
||||
compressed = [
|
||||
{"role": "user", "content": "summary"},
|
||||
{"role": "assistant", "content": "continuing..."},
|
||||
]
|
||||
|
||||
# Stale history longer than messages: the old positional flush
|
||||
# sliced past the end and dropped both messages (#46053).
|
||||
stale_history = [{"role": "user", "content": f"msg{i}"} for i in range(100)]
|
||||
agent._flush_messages_to_session_db(compressed, stale_history)
|
||||
|
||||
rows = db.get_messages("new-session")
|
||||
assert len(rows) == 2
|
||||
assert [row["content"] for row in rows] == ["summary", "continuing..."]
|
||||
|
||||
def test_in_place_compression_rebaseline_prevents_duplicate_compacted_rows(self):
|
||||
"""In-place compaction already persisted the compacted transcript.
|
||||
|
||||
Regression for the 2026-06-26 SRE compression loop: archive_and_compact()
|
||||
inserted a compacted active block, then the same turn continued with
|
||||
conversation_history=None and _flush_messages_to_session_db() appended
|
||||
the compacted dicts again, doubling live context.
|
||||
"""
|
||||
from agent.conversation_compression import conversation_history_after_compression
|
||||
from hermes_state import SessionDB
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db_path = Path(tmpdir) / "test.db"
|
||||
db = SessionDB(db_path=db_path)
|
||||
|
||||
agent = self._make_agent(db)
|
||||
agent._ensure_db_session()
|
||||
|
||||
original_history = [
|
||||
{"role": "user", "content": "old question"},
|
||||
{"role": "assistant", "content": "old answer"},
|
||||
]
|
||||
agent._flush_messages_to_session_db(original_history, [])
|
||||
assert [row["content"] for row in db.get_messages("original-session")] == [
|
||||
"old question",
|
||||
"old answer",
|
||||
]
|
||||
|
||||
compacted = [
|
||||
{"role": "assistant", "content": "[CONTEXT COMPACTION] summary"},
|
||||
{"role": "user", "content": "recent question"},
|
||||
{"role": "assistant", "content": "recent answer"},
|
||||
]
|
||||
db.archive_and_compact("original-session", compacted)
|
||||
setattr(agent, "_last_compaction_in_place", True)
|
||||
agent._last_flushed_db_idx = 0
|
||||
|
||||
# Same agent turn continues after compaction. The compacted dicts
|
||||
# must be treated as already-persisted history; only later appends
|
||||
# should be flushed.
|
||||
post_compaction_history = conversation_history_after_compression(
|
||||
agent, compacted
|
||||
)
|
||||
assert post_compaction_history is not None
|
||||
assert post_compaction_history is not compacted
|
||||
assert post_compaction_history == compacted
|
||||
|
||||
messages = compacted + [
|
||||
{"role": "tool", "content": "tool result"},
|
||||
{"role": "assistant", "content": "final answer"},
|
||||
]
|
||||
agent._flush_messages_to_session_db(messages, post_compaction_history)
|
||||
|
||||
rows = db.get_messages("original-session")
|
||||
assert [row["content"] for row in rows] == [
|
||||
"[CONTEXT COMPACTION] summary",
|
||||
"recent question",
|
||||
"recent answer",
|
||||
"tool result",
|
||||
"final answer",
|
||||
]
|
||||
|
||||
def test_abort_after_in_place_compaction_preserves_flush_baseline(self):
|
||||
"""An aborted retry must survive flush, restart, and resume."""
|
||||
from agent.conversation_compression import (
|
||||
compress_context,
|
||||
conversation_history_after_compression,
|
||||
)
|
||||
from hermes_state import SessionDB
|
||||
|
||||
class SuccessCompressor:
|
||||
_last_compress_aborted = False
|
||||
_last_summary_error = None
|
||||
compression_count = 1
|
||||
_last_compression_made_progress = True
|
||||
_last_summary_fallback_used = False
|
||||
last_compression_rough_tokens = 0
|
||||
last_prompt_tokens = 0
|
||||
last_completion_tokens = 0
|
||||
awaiting_real_usage_after_compression = False
|
||||
|
||||
def compress(self, _messages, **_kwargs):
|
||||
return [
|
||||
{"role": "user", "content": "[summary] earlier state"},
|
||||
{"role": "assistant", "content": "retained tail"},
|
||||
]
|
||||
|
||||
class AbortCompressor:
|
||||
_last_compress_aborted = False
|
||||
_last_summary_error = "simulated auxiliary timeout"
|
||||
compression_count = 2
|
||||
_last_compression_made_progress = False
|
||||
_last_summary_fallback_used = False
|
||||
last_compression_rough_tokens = 0
|
||||
last_prompt_tokens = 0
|
||||
last_completion_tokens = 0
|
||||
awaiting_real_usage_after_compression = False
|
||||
|
||||
def compress(self, messages, **_kwargs):
|
||||
self._last_compress_aborted = True
|
||||
return messages
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db_path = Path(tmpdir) / "test.db"
|
||||
db = SessionDB(db_path=db_path)
|
||||
agent = self._make_agent(db)
|
||||
agent.compression_in_place = True
|
||||
original = [
|
||||
{"role": "user", "content": "old question"},
|
||||
{"role": "assistant", "content": "old answer"},
|
||||
]
|
||||
agent._flush_messages_to_session_db(original, [])
|
||||
|
||||
agent.context_compressor = SuccessCompressor()
|
||||
compacted, _ = compress_context(
|
||||
agent, original, "system", approx_tokens=100_000
|
||||
)
|
||||
history = conversation_history_after_compression(
|
||||
agent, compacted, None
|
||||
)
|
||||
|
||||
messages = compacted + [
|
||||
{"role": "user", "content": "new request"},
|
||||
{"role": "assistant", "content": "new answer"},
|
||||
]
|
||||
agent.context_compressor = AbortCompressor()
|
||||
returned, _ = compress_context(
|
||||
agent, messages, "system", approx_tokens=100_000
|
||||
)
|
||||
history = conversation_history_after_compression(
|
||||
agent, returned, history
|
||||
)
|
||||
agent._flush_messages_to_session_db(returned, history)
|
||||
|
||||
db.close()
|
||||
resumed_db = SessionDB(db_path=db_path)
|
||||
assert [message["content"] for message in resumed_db.get_messages_as_conversation(
|
||||
agent.session_id
|
||||
)] == [
|
||||
"[summary] earlier state",
|
||||
"retained tail",
|
||||
"new request",
|
||||
"new answer",
|
||||
]
|
||||
resumed_db.close()
|
||||
|
||||
def test_rotation_child_session_flushes_full_compressed_transcript_with_markers(self):
|
||||
"""Regression for #57491: live cached-agent markers must not block child flush."""
|
||||
from agent.conversation_compression import compress_context
|
||||
from hermes_state import SessionDB
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db_path = Path(tmpdir) / "test.db"
|
||||
db = SessionDB(db_path=db_path)
|
||||
parent_sid = "20260701_152840_parent"
|
||||
db.create_session(parent_sid, "gateway", model="test/model")
|
||||
|
||||
agent = self._make_agent(db)
|
||||
agent.session_id = parent_sid
|
||||
agent.compression_in_place = False
|
||||
agent._ensure_db_session()
|
||||
|
||||
# Plain marked messages only: the exact-equality assertion below
|
||||
# relies on `compressed` containing no message that _flush filters
|
||||
# for a reason INDEPENDENT of _db_persisted (ephemeral scaffolding,
|
||||
# synthetic recovery turns). Keep this fixture free of such messages
|
||||
# or the row count would legitimately differ from len(compressed).
|
||||
# The transcript must also be large enough that the provider-less
|
||||
# static fallback net-shrinks it (middle drops must outweigh the
|
||||
# fixed compaction marker overhead), or the no-growth commit guard
|
||||
# correctly refuses the rotation this test exercises. Sized for
|
||||
# the lean tail default: the 10K-token tail floor must leave a
|
||||
# substantial compressible middle (~2K chars/message × 40 ≈ 20K
|
||||
# estimated tokens total).
|
||||
messages = [
|
||||
{
|
||||
"role": "user" if i % 2 == 0 else "assistant",
|
||||
"content": f"message {i} " + "x" * 2000,
|
||||
"_db_persisted": True,
|
||||
}
|
||||
for i in range(40)
|
||||
]
|
||||
|
||||
with patch("agent.context_compressor.call_llm", side_effect=RuntimeError("no provider")):
|
||||
compressed, _ = compress_context(
|
||||
agent, messages, approx_tokens=100_000, system_message="sys"
|
||||
)
|
||||
|
||||
assert agent.session_id != parent_sid
|
||||
child_sid = agent.session_id
|
||||
|
||||
agent._flush_messages_to_session_db(compressed, None)
|
||||
|
||||
child_rows = db.get_messages(child_sid)
|
||||
assert len(child_rows) == len(compressed), (
|
||||
f"Expected {len(compressed)} rows in child session, got {len(child_rows)}. "
|
||||
f"_db_persisted marker propagation bug (#57491)."
|
||||
)
|
||||
db.close()
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Part 2: Gateway-side — history_offset after session split
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestGatewayHistoryOffsetAfterSplit:
|
||||
"""Verify that when the agent creates a new session during compression,
|
||||
the gateway uses history_offset=0 so all compressed messages are written
|
||||
to the JSONL transcript."""
|
||||
|
||||
def test_history_offset_zero_on_session_split(self):
|
||||
"""When agent.session_id differs from the original, history_offset must be 0."""
|
||||
# This tests the logic in gateway/run.py run_sync():
|
||||
# _session_was_split = agent.session_id != session_id
|
||||
# _effective_history_offset = 0 if _session_was_split else len(agent_history)
|
||||
|
||||
original_session_id = "session-abc"
|
||||
agent_session_id = "session-compressed-xyz" # Different = compression happened
|
||||
agent_history_len = 200
|
||||
|
||||
# Simulate the gateway's offset calculation (post-fix)
|
||||
_session_was_split = (agent_session_id != original_session_id)
|
||||
_effective_history_offset = 0 if _session_was_split else agent_history_len
|
||||
|
||||
assert _session_was_split is True
|
||||
assert _effective_history_offset == 0
|
||||
|
||||
|
||||
def test_new_messages_extraction_after_split(self):
|
||||
"""After compression with offset=0, new_messages should be ALL agent messages."""
|
||||
# Simulates the gateway's new_messages calculation
|
||||
agent_messages = [
|
||||
{"role": "user", "content": "[CONTEXT COMPACTION] Summary..."},
|
||||
{"role": "user", "content": "recent question"},
|
||||
{"role": "assistant", "content": "recent answer"},
|
||||
{"role": "user", "content": "new question"},
|
||||
{"role": "assistant", "content": "new answer"},
|
||||
]
|
||||
history_offset = 0 # After fix: 0 on session split
|
||||
|
||||
new_messages = agent_messages[history_offset:] if len(agent_messages) > history_offset else []
|
||||
assert len(new_messages) == 5, (
|
||||
f"Expected all 5 messages with offset=0, got {len(new_messages)}"
|
||||
)
|
||||
|
||||
|
||||
|
||||
class TestStoredPromptCwdDrift:
|
||||
"""Verify that stored system prompts are rejected when cwd changed."""
|
||||
|
||||
def _make_agent(self, model="test/model", provider="openrouter"):
|
||||
class _Agent:
|
||||
pass
|
||||
|
||||
agent = _Agent()
|
||||
agent.model = model
|
||||
agent.provider = provider
|
||||
return agent
|
||||
|
||||
@staticmethod
|
||||
def _host_block(cwd: str) -> str:
|
||||
"""A stored prompt fragment shaped like the real host-info block.
|
||||
|
||||
``build_environment_hints`` always emits ``User home directory:``
|
||||
immediately before the working-directory line, and the staleness check
|
||||
anchors on that pair so user project files can't shadow the real value.
|
||||
Fixtures must therefore include the anchor or they stop exercising the
|
||||
cwd path at all.
|
||||
"""
|
||||
return (
|
||||
"Host: Linux (6.16.0)\n"
|
||||
"User home directory: /home/tester\n"
|
||||
f"Current working directory: {cwd}\n"
|
||||
)
|
||||
|
||||
def test_stored_prompt_stale_when_cwd_differs(self):
|
||||
"""Different cwd should force a prompt rebuild."""
|
||||
from unittest.mock import patch
|
||||
from agent.conversation_loop import _stored_prompt_matches_runtime
|
||||
|
||||
agent = self._make_agent()
|
||||
stored_prompt = (
|
||||
self._host_block("/project/old")
|
||||
+ "Model: test/model\n"
|
||||
"Provider: openrouter\n"
|
||||
)
|
||||
|
||||
with patch("os.getcwd", return_value="/project/new"):
|
||||
assert _stored_prompt_matches_runtime(agent, stored_prompt) is False, (
|
||||
"Expected False when stored cwd differs from current cwd"
|
||||
)
|
||||
|
||||
def test_stored_prompt_fresh_when_cwd_matches(self):
|
||||
"""Matching cwd should allow prompt reuse."""
|
||||
from unittest.mock import patch
|
||||
from agent.conversation_loop import _stored_prompt_matches_runtime
|
||||
|
||||
agent = self._make_agent()
|
||||
current_cwd = "/project/current"
|
||||
stored_prompt = (
|
||||
self._host_block(current_cwd)
|
||||
+ "Model: test/model\n"
|
||||
"Provider: openrouter\n"
|
||||
)
|
||||
|
||||
with patch("os.getcwd", return_value=current_cwd):
|
||||
assert _stored_prompt_matches_runtime(agent, stored_prompt) is True, (
|
||||
"Expected True when stored cwd matches current cwd"
|
||||
)
|
||||
|
||||
def test_project_context_cannot_force_a_rebuild(self):
|
||||
"""🔴 CACHE INVARIANT: user project text must never invalidate the prompt.
|
||||
|
||||
The prompt embeds AGENTS.md / CLAUDE.md / .cursorrules in the context
|
||||
tier, which sits AFTER the host-info block. A whole-prompt scan for
|
||||
``Current working directory:`` therefore matched the user's own file
|
||||
and compared runtime state against project prose. That mismatch never
|
||||
clears, so the check rejected the stored prompt on EVERY turn —
|
||||
rebuilding the system prompt each message and destroying the prefix
|
||||
cache for the entire session. Strictly worse than the staleness this
|
||||
check exists to catch.
|
||||
"""
|
||||
from unittest.mock import patch
|
||||
from agent.conversation_loop import _stored_prompt_matches_runtime
|
||||
|
||||
agent = self._make_agent()
|
||||
current_cwd = "/project/current"
|
||||
stored_prompt = (
|
||||
self._host_block(current_cwd)
|
||||
+ "\n# AGENTS.md\n\n"
|
||||
"Our deploy convention:\n\n"
|
||||
"Current working directory: /srv/decoy\n\n"
|
||||
"Always run make before pushing.\n\n"
|
||||
"Model: test/model\n"
|
||||
"Provider: openrouter\n"
|
||||
)
|
||||
|
||||
with patch("os.getcwd", return_value=current_cwd):
|
||||
assert _stored_prompt_matches_runtime(agent, stored_prompt) is True, (
|
||||
"A project file that merely MENTIONS 'Current working "
|
||||
"directory:' must not invalidate the prompt — that would "
|
||||
"rebuild every turn and break the prefix cache"
|
||||
)
|
||||
|
||||
def test_project_context_cannot_mask_real_drift(self):
|
||||
"""The inverse: project text must not fake a match either.
|
||||
|
||||
A stored prompt built in /project/old whose embedded AGENTS.md happens
|
||||
to name the NEW cwd must still be rejected — otherwise project prose
|
||||
could suppress genuine drift detection.
|
||||
"""
|
||||
from unittest.mock import patch
|
||||
from agent.conversation_loop import _stored_prompt_matches_runtime
|
||||
|
||||
agent = self._make_agent()
|
||||
stored_prompt = (
|
||||
self._host_block("/project/old")
|
||||
+ "\n# AGENTS.md\n\n"
|
||||
"Current working directory: /project/new\n\n"
|
||||
"Model: test/model\n"
|
||||
"Provider: openrouter\n"
|
||||
)
|
||||
|
||||
with patch("os.getcwd", return_value="/project/new"):
|
||||
assert _stored_prompt_matches_runtime(agent, stored_prompt) is False, (
|
||||
"Embedded project text naming the new cwd must not mask real "
|
||||
"drift in the host-info block"
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def test_built_prompt_contains_platform_line(self):
|
||||
"""The built system prompt must carry a Platform: line so drift detection works."""
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
from hermes_state import SessionDB
|
||||
from run_agent import AIAgent
|
||||
from agent.system_prompt import build_system_prompt_parts
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
db = SessionDB(db_path=Path(tmpdir) / "test.db")
|
||||
with patch.dict(os.environ, {"OPENROUTER_API_KEY": "test-key"}):
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
provider="openrouter",
|
||||
quiet_mode=True,
|
||||
session_db=db,
|
||||
session_id="platform-test",
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
agent.platform = "cli"
|
||||
parts = build_system_prompt_parts(agent)
|
||||
assert "Platform: cli" in parts["volatile"], (
|
||||
"Built prompt missing 'Platform: cli' — drift detection cannot read it"
|
||||
)
|
||||
@@ -0,0 +1,604 @@
|
||||
"""Regressions for the #76354 review of the compression timeout architecture.
|
||||
|
||||
Every test here asserts the BLOCKED/hung state itself where the review demands
|
||||
it — the worker is released only AFTER the assertion (helix4u called out two
|
||||
prior tests that released before asserting; do not regress that).
|
||||
|
||||
Covers:
|
||||
- F1: commit-phase overrun warning fires WHILE the commit is hung (lock-free
|
||||
``commit_in_flight`` phase marker).
|
||||
- F2: every host unwind (KeyboardInterrupt / generic exception) revokes commit
|
||||
admission before the host resumes.
|
||||
- F4 (unit half): a cancelled attempt cannot clear the failure cooldown
|
||||
(fence check ordered BEFORE cooldown-clear).
|
||||
- F6: bounded admission — four wedged workers refuse a fifth submission fast,
|
||||
and the refused job never runs later; a cancelled fence skips summary work.
|
||||
- S3 analogue: the idle wait is charged from the last progress event, so
|
||||
silence cannot approach 2x the configured idle timeout.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import concurrent.futures
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
import agent.conversation_compression as cc
|
||||
from agent.conversation_compression import (
|
||||
CompressionCommitFence,
|
||||
run_compress_context_with_progress_timeout,
|
||||
)
|
||||
|
||||
|
||||
def _drain_admission_slots():
|
||||
"""Best-effort wait for pool admission slots to free between tests."""
|
||||
deadline = time.time() + 5
|
||||
while time.time() < deadline:
|
||||
with cc._compress_admission_lock:
|
||||
if cc._compress_admitted_count == 0:
|
||||
return
|
||||
time.sleep(0.02)
|
||||
|
||||
|
||||
class TestF1CommitOverrunWhileHung:
|
||||
def test_overrun_warning_fires_while_commit_still_blocked(self, monkeypatch):
|
||||
"""The warning + on_commit_overrun fire DURING the hang, not after.
|
||||
|
||||
The fake commit is event-gated and is NOT released until after the
|
||||
assertions on the callback/log have been made while the worker
|
||||
thread is still blocked inside the commit boundary.
|
||||
"""
|
||||
executor = concurrent.futures.ThreadPoolExecutor(max_workers=1)
|
||||
monkeypatch.setattr(cc, "_get_compress_timeout_executor", lambda: executor)
|
||||
original = [{"role": "user", "content": "a"}]
|
||||
compressed = [{"role": "assistant", "content": "late"}]
|
||||
entered = threading.Event()
|
||||
release = threading.Event()
|
||||
overrun_fired = threading.Event()
|
||||
overruns = []
|
||||
|
||||
def worker(fence: CompressionCommitFence):
|
||||
assert fence.begin_commit()
|
||||
entered.set()
|
||||
try:
|
||||
# Hung commit: blocked until the TEST releases it, which
|
||||
# happens only after asserting the overrun surfaced.
|
||||
assert release.wait(timeout=10)
|
||||
return (compressed, "committed-late")
|
||||
finally:
|
||||
fence.finish_commit()
|
||||
|
||||
records = []
|
||||
|
||||
class _Capture(logging.Handler):
|
||||
def emit(self, record):
|
||||
records.append(record)
|
||||
|
||||
def on_overrun(waited, ceil):
|
||||
overruns.append((waited, ceil))
|
||||
overrun_fired.set()
|
||||
|
||||
done = {}
|
||||
|
||||
def run():
|
||||
done["result"] = run_compress_context_with_progress_timeout(
|
||||
worker=worker,
|
||||
messages=original,
|
||||
system_prompt_fallback="fallback",
|
||||
idle_timeout_seconds=1.0,
|
||||
total_ceiling_seconds=1.0,
|
||||
on_commit_overrun=on_overrun,
|
||||
)
|
||||
|
||||
comp_logger = logging.getLogger("agent.conversation_compression")
|
||||
handler = _Capture(level=logging.WARNING)
|
||||
comp_logger.addHandler(handler)
|
||||
try:
|
||||
t = threading.Thread(target=run, name="f1-hung-commit-host")
|
||||
t.start()
|
||||
try:
|
||||
assert entered.wait(timeout=2)
|
||||
# ── Assert WHILE the commit worker is still blocked ──────
|
||||
assert overrun_fired.wait(timeout=5), (
|
||||
"on_commit_overrun must fire while the commit is hung"
|
||||
)
|
||||
assert not release.is_set() # worker provably still blocked
|
||||
assert t.is_alive()
|
||||
deadline = time.time() + 5
|
||||
while time.time() < deadline:
|
||||
if any(
|
||||
r.levelno >= logging.WARNING
|
||||
and "past the total ceiling" in r.getMessage()
|
||||
for r in list(records)
|
||||
):
|
||||
break
|
||||
time.sleep(0.01)
|
||||
overrun_logs = [
|
||||
r
|
||||
for r in list(records)
|
||||
if r.levelno >= logging.WARNING
|
||||
and "past the total ceiling" in r.getMessage()
|
||||
]
|
||||
assert overrun_logs, (
|
||||
"expected the overrun WARNING while the commit was "
|
||||
f"still blocked; got: {[r.getMessage() for r in records]}"
|
||||
)
|
||||
assert overruns and overruns[0][1] == pytest.approx(1.0)
|
||||
finally:
|
||||
release.set()
|
||||
t.join(timeout=5)
|
||||
assert not t.is_alive()
|
||||
finally:
|
||||
comp_logger.removeHandler(handler)
|
||||
executor.shutdown(wait=True)
|
||||
assert done["result"] == (compressed, "committed-late")
|
||||
_drain_admission_slots()
|
||||
|
||||
def test_commit_in_flight_marker_is_lock_free(self):
|
||||
fence = CompressionCommitFence()
|
||||
assert fence.commit_in_flight is False
|
||||
assert fence.begin_commit()
|
||||
# The fence lock is HELD here; the marker must still be readable.
|
||||
assert fence.commit_in_flight is True
|
||||
fence.finish_commit()
|
||||
assert fence.commit_in_flight is False
|
||||
|
||||
|
||||
class _KIOnFirstResultFuture:
|
||||
"""Future proxy raising on the host's first result() call."""
|
||||
|
||||
def __init__(self, inner, exc, gate=None):
|
||||
self._inner = inner
|
||||
self._exc = exc
|
||||
self._raised = False
|
||||
self._gate = gate
|
||||
|
||||
def result(self, timeout=None):
|
||||
if not self._raised:
|
||||
self._raised = True
|
||||
if self._gate is not None:
|
||||
# Ensure the pooled worker has genuinely STARTED before the
|
||||
# host unwinds, so the test exercises "unwind with a live
|
||||
# worker" rather than the queued-job skip path.
|
||||
assert self._gate.wait(timeout=5)
|
||||
raise self._exc
|
||||
return self._inner.result(timeout=timeout)
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._inner, name)
|
||||
|
||||
|
||||
class _InjectingExecutor:
|
||||
def __init__(self, inner, exc, gate=None):
|
||||
self._inner = inner
|
||||
self._exc = exc
|
||||
self._gate = gate
|
||||
|
||||
def submit(self, fn, *args, **kwargs):
|
||||
return _KIOnFirstResultFuture(
|
||||
self._inner.submit(fn, *args, **kwargs), self._exc, self._gate
|
||||
)
|
||||
|
||||
|
||||
class TestF2HostUnwindRevokesAdmission:
|
||||
@pytest.mark.parametrize(
|
||||
"exc_type", [KeyboardInterrupt, RuntimeError], ids=["ki", "generic"]
|
||||
)
|
||||
def test_unwind_revokes_commit_admission_before_host_returns(
|
||||
self, monkeypatch, exc_type
|
||||
):
|
||||
"""KI/exception while waiting → detached worker can never commit.
|
||||
|
||||
The worker is still blocked pre-commit when the host unwinds; the
|
||||
assertions run BEFORE the worker is released.
|
||||
"""
|
||||
original = [{"role": "user", "content": "keep"}]
|
||||
started = threading.Event()
|
||||
release = threading.Event()
|
||||
fence_box = {}
|
||||
commit_admitted = {}
|
||||
|
||||
def worker(fence: CompressionCommitFence):
|
||||
fence_box["fence"] = fence
|
||||
started.set()
|
||||
assert release.wait(timeout=10)
|
||||
commit_admitted["value"] = fence.begin_commit()
|
||||
if commit_admitted["value"]:
|
||||
fence.finish_commit()
|
||||
return ([{"role": "assistant", "content": "late"}], "x")
|
||||
|
||||
real_executor = cc._get_compress_timeout_executor()
|
||||
monkeypatch.setattr(
|
||||
cc,
|
||||
"_get_compress_timeout_executor",
|
||||
lambda: _InjectingExecutor(real_executor, exc_type(), gate=started),
|
||||
)
|
||||
|
||||
with pytest.raises(exc_type):
|
||||
run_compress_context_with_progress_timeout(
|
||||
worker=worker,
|
||||
messages=original,
|
||||
system_prompt_fallback="fallback",
|
||||
idle_timeout_seconds=5.0,
|
||||
total_ceiling_seconds=5.0,
|
||||
)
|
||||
|
||||
# ── Host has unwound; worker is STILL blocked pre-commit ─────────
|
||||
assert started.wait(timeout=2)
|
||||
fence = fence_box["fence"]
|
||||
assert not release.is_set()
|
||||
assert fence.is_cancelled, (
|
||||
"host unwind must revoke commit admission while the worker "
|
||||
"is still running"
|
||||
)
|
||||
# Now release the worker and prove its commit was refused.
|
||||
release.set()
|
||||
deadline = time.time() + 5
|
||||
while time.time() < deadline and "value" not in commit_admitted:
|
||||
time.sleep(0.01)
|
||||
assert commit_admitted.get("value") is False, (
|
||||
"a worker surviving a host unwind must be denied the commit "
|
||||
"boundary"
|
||||
)
|
||||
_drain_admission_slots()
|
||||
|
||||
|
||||
class TestF4CooldownClearOrdering:
|
||||
def test_cancelled_attempt_cannot_clear_failure_cooldown(self):
|
||||
"""Fence check ordered BEFORE cooldown-clear (review F4 ordering)."""
|
||||
from agent.context_compressor import ContextCompressor
|
||||
|
||||
class _FakeCompressor:
|
||||
_summary_failure_cooldown_until = 12345.0
|
||||
_last_summary_error = "timeout"
|
||||
_consecutive_timeout_failures = 2
|
||||
_cooldown_persist_failed = False
|
||||
_session_db = None
|
||||
_session_id = ""
|
||||
_compression_cancelled_check = staticmethod(lambda: True)
|
||||
|
||||
fake = _FakeCompressor()
|
||||
ContextCompressor._clear_compression_failure_cooldown(fake)
|
||||
assert fake._summary_failure_cooldown_until == 12345.0, (
|
||||
"a cancelled attempt must NOT undo the host's timeout cooldown"
|
||||
)
|
||||
assert fake._consecutive_timeout_failures == 2
|
||||
|
||||
# Sabotage check: with the fence reporting NOT cancelled, the clear
|
||||
# must proceed (proves the guard is the only thing blocking it).
|
||||
fake2 = _FakeCompressor()
|
||||
fake2._compression_cancelled_check = staticmethod(lambda: False)
|
||||
ContextCompressor._clear_compression_failure_cooldown(fake2)
|
||||
assert fake2._summary_failure_cooldown_until == 0.0
|
||||
|
||||
|
||||
class TestF6ExecutorSaturation:
|
||||
def test_saturated_pool_fails_fast_and_never_runs_stale_job(self):
|
||||
"""4 blocked summaries + 5th submission fails fast; recovery does not
|
||||
run the refused job."""
|
||||
_drain_admission_slots()
|
||||
release = threading.Event()
|
||||
started = threading.Barrier(5, timeout=10) # 4 workers + main
|
||||
|
||||
def blocked_worker(fence: CompressionCommitFence):
|
||||
started.wait()
|
||||
assert release.wait(timeout=30)
|
||||
return ([], "done")
|
||||
|
||||
hosts = []
|
||||
results = {}
|
||||
|
||||
def host(i):
|
||||
results[i] = run_compress_context_with_progress_timeout(
|
||||
worker=blocked_worker,
|
||||
messages=[{"role": "user", "content": f"m{i}"}],
|
||||
system_prompt_fallback=f"fb{i}",
|
||||
idle_timeout_seconds=0.05,
|
||||
total_ceiling_seconds=0.1,
|
||||
)
|
||||
|
||||
try:
|
||||
for i in range(4):
|
||||
t = threading.Thread(target=host, args=(i,), name=f"sat-{i}")
|
||||
t.start()
|
||||
hosts.append(t)
|
||||
started.wait() # all 4 workers occupy the pool
|
||||
for t in hosts:
|
||||
t.join(timeout=5) # hosts time out; workers stay wedged
|
||||
assert not t.is_alive()
|
||||
|
||||
# All 4 slots still admitted (workers blocked).
|
||||
with cc._compress_admission_lock:
|
||||
assert cc._compress_admitted_count == 4
|
||||
|
||||
fifth_ran = threading.Event()
|
||||
|
||||
def fifth_worker(fence):
|
||||
fifth_ran.set()
|
||||
return ([], "5th")
|
||||
|
||||
fifth_msgs = [{"role": "user", "content": "fifth"}]
|
||||
# Round-2 #6: the fail-fast refusal must emit the standard
|
||||
# compression-attempt telemetry with failure_class=pool_saturated.
|
||||
class _TelemetryAgent:
|
||||
session_id = "SATURATED_SESSION"
|
||||
_compression_attempt_id = "sat-attempt"
|
||||
|
||||
class context_compressor: # noqa: D106 — minimal stub
|
||||
_last_compression_telemetry = None
|
||||
_last_summary_fallback_used = False
|
||||
_last_aux_model_failure_model = None
|
||||
|
||||
import json as _json
|
||||
import logging as _logging
|
||||
|
||||
class _CaptureHandler(_logging.Handler):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.payloads = []
|
||||
|
||||
def emit(self, record):
|
||||
msg = record.getMessage()
|
||||
if "compression attempt telemetry" in msg:
|
||||
self.payloads.append(
|
||||
_json.loads(msg.split(": ", 1)[1])
|
||||
)
|
||||
|
||||
capture = _CaptureHandler()
|
||||
_prev_level = cc.logger.level
|
||||
cc.logger.addHandler(capture)
|
||||
cc.logger.setLevel(_logging.DEBUG)
|
||||
t0 = time.monotonic()
|
||||
try:
|
||||
msgs, prompt = run_compress_context_with_progress_timeout(
|
||||
worker=fifth_worker,
|
||||
messages=fifth_msgs,
|
||||
system_prompt_fallback="fifth-fallback",
|
||||
idle_timeout_seconds=5.0,
|
||||
total_ceiling_seconds=5.0,
|
||||
telemetry_agent=_TelemetryAgent(),
|
||||
)
|
||||
finally:
|
||||
cc.logger.removeHandler(capture)
|
||||
cc.logger.setLevel(_prev_level)
|
||||
elapsed = time.monotonic() - t0
|
||||
# ── Assert while the 4 workers are STILL wedged ───────────────
|
||||
assert not release.is_set()
|
||||
assert elapsed < 1.0, (
|
||||
f"saturated submission must fail fast, took {elapsed:.2f}s"
|
||||
)
|
||||
assert msgs is fifth_msgs
|
||||
assert prompt == "fifth-fallback"
|
||||
assert not fifth_ran.is_set()
|
||||
saturated = [
|
||||
p for p in capture.payloads
|
||||
if p.get("failure_class") == "pool_saturated"
|
||||
]
|
||||
assert saturated, (
|
||||
"fail-fast admission refusal must emit compression-attempt "
|
||||
"telemetry with failure_class='pool_saturated'"
|
||||
)
|
||||
assert saturated[0]["commit_status"] == "aborted"
|
||||
assert saturated[0]["session_id"] == "SATURATED_SESSION"
|
||||
finally:
|
||||
release.set()
|
||||
|
||||
# Worker recovery: slots free, and the refused fifth job never runs.
|
||||
_drain_admission_slots()
|
||||
time.sleep(0.1)
|
||||
assert not fifth_ran.is_set(), (
|
||||
"recovered workers must not run the refused stale job"
|
||||
)
|
||||
# Recovery restores service: a new submission is admitted and runs.
|
||||
msgs, prompt = run_compress_context_with_progress_timeout(
|
||||
worker=lambda fence: ([{"role": "user", "content": "ok"}], "ok"),
|
||||
messages=[{"role": "user", "content": "after"}],
|
||||
system_prompt_fallback="fb",
|
||||
idle_timeout_seconds=1.0,
|
||||
total_ceiling_seconds=2.0,
|
||||
)
|
||||
assert prompt == "ok"
|
||||
_drain_admission_slots()
|
||||
|
||||
def test_cancelled_fence_skips_summary_work_before_start(self):
|
||||
"""A stale job whose fence was already cancelled never runs summary.
|
||||
|
||||
Drives compress_context's pre-summary fence gate directly: the fence
|
||||
is cancelled BEFORE dispatch, so the expensive compress() call must
|
||||
not run and the transcript must come back unchanged.
|
||||
"""
|
||||
import os
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
import tempfile
|
||||
|
||||
from hermes_state import SessionDB
|
||||
|
||||
with tempfile.TemporaryDirectory() as td:
|
||||
db = SessionDB(db_path=Path(td) / "state.db")
|
||||
session_id = "F6_PRESTART_FENCE"
|
||||
db.create_session(session_id, source="cli")
|
||||
with patch.dict(os.environ, {"OPENROUTER_API_KEY": "test-key"}):
|
||||
from run_agent import AIAgent
|
||||
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
quiet_mode=True,
|
||||
session_db=db,
|
||||
session_id=session_id,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
compressor = MagicMock()
|
||||
compressor.compress.return_value = [
|
||||
{"role": "user", "content": "should-not-run"}
|
||||
]
|
||||
compressor._last_summary_error = None
|
||||
compressor._last_compress_aborted = False
|
||||
compressor._last_aux_model_failure_model = None
|
||||
compressor._last_aux_model_failure_error = None
|
||||
agent.context_compressor = compressor
|
||||
agent._cached_system_prompt = "sys"
|
||||
|
||||
fence = CompressionCommitFence()
|
||||
assert fence.cancel_before_commit() is True
|
||||
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
returned, _sp = agent._compress_context(
|
||||
messages, "sys", approx_tokens=120_000, commit_fence=fence
|
||||
)
|
||||
|
||||
compressor.compress.assert_not_called()
|
||||
assert returned is messages
|
||||
# The cancelled attempt must not leave the durable lock held.
|
||||
assert db.get_compression_lock_holder(session_id) is None
|
||||
db.close()
|
||||
|
||||
|
||||
class TestS3IdleChargedFromLastProgress:
|
||||
def test_silence_cannot_approach_double_idle_timeout(self):
|
||||
"""Progress early in an interval must not extend silence to ~2x idle."""
|
||||
_drain_admission_slots()
|
||||
idle = 0.4
|
||||
release = threading.Event()
|
||||
|
||||
def worker(fence: CompressionCommitFence):
|
||||
time.sleep(0.05)
|
||||
fence.touch_progress() # early progress, then total silence
|
||||
assert release.wait(timeout=10)
|
||||
return ([], "late")
|
||||
|
||||
t0 = time.monotonic()
|
||||
try:
|
||||
msgs, prompt = run_compress_context_with_progress_timeout(
|
||||
worker=worker,
|
||||
messages=[{"role": "user", "content": "a"}],
|
||||
system_prompt_fallback="fb",
|
||||
idle_timeout_seconds=idle,
|
||||
total_ceiling_seconds=5.0,
|
||||
stall_fallback=False,
|
||||
)
|
||||
finally:
|
||||
elapsed = time.monotonic() - t0
|
||||
release.set()
|
||||
assert prompt == "fb"
|
||||
# Old behavior waited a full interval from the CHECK (~2x idle ≈
|
||||
# 0.85s+). New behavior times out ~idle after the last progress
|
||||
# (~0.45s). Allow generous slack while still excluding ~2x.
|
||||
assert elapsed < idle * 1.8, (
|
||||
f"silence exceeded ~2x idle budget shape: {elapsed:.2f}s"
|
||||
)
|
||||
_drain_admission_slots()
|
||||
|
||||
|
||||
class TestRound2MidCommitLeaseRelease:
|
||||
"""Round-2 #1: revoke must not release the durable lease mid-commit.
|
||||
|
||||
Invariant: at no point can a second compressor acquire the durable lock
|
||||
while an admitted commit is still mutating; after the commit finishes
|
||||
post-revoke, the lease IS released promptly even if the worker thread is
|
||||
later parked (never runs its outer cleanup).
|
||||
"""
|
||||
|
||||
def _db_with_lease(self, tmp_path):
|
||||
from hermes_state import SessionDB
|
||||
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "R2_MID_COMMIT_LEASE"
|
||||
db.create_session(session_id, source="cli")
|
||||
holder = "pid:worker:original"
|
||||
assert db.try_acquire_compression_lock(
|
||||
session_id, holder, ttl_seconds=60
|
||||
)
|
||||
return db, session_id, holder
|
||||
|
||||
def test_revoke_during_in_flight_commit_defers_lease_release(
|
||||
self, tmp_path
|
||||
):
|
||||
"""Event-gated fake commit; assertions run WHILE it is blocked."""
|
||||
db, session_id, holder = self._db_with_lease(tmp_path)
|
||||
fence = CompressionCommitFence()
|
||||
fence.register_cancelled_lock_release(
|
||||
lambda: db.release_compression_lock(session_id, holder)
|
||||
)
|
||||
|
||||
commit_entered = threading.Event()
|
||||
release_commit = threading.Event()
|
||||
commit_finished = threading.Event()
|
||||
|
||||
def _committing_worker():
|
||||
assert fence.begin_commit()
|
||||
commit_entered.set()
|
||||
assert release_commit.wait(timeout=10)
|
||||
fence.finish_commit()
|
||||
commit_finished.set()
|
||||
# Park forever: the deferred release must NOT depend on this
|
||||
# thread's outer cleanup running.
|
||||
threading.Event().wait(30)
|
||||
|
||||
worker = threading.Thread(target=_committing_worker, daemon=True)
|
||||
worker.start()
|
||||
assert commit_entered.wait(timeout=5)
|
||||
|
||||
# Host revokes WHILE the commit is in flight.
|
||||
fence.revoke_commit_admission()
|
||||
|
||||
# ── Assert the hung state BEFORE releasing the commit ────────────
|
||||
assert not commit_finished.is_set()
|
||||
assert db.get_compression_lock_holder(session_id) == holder, (
|
||||
"revoke released the durable lease while a commit was still "
|
||||
"mutating SessionDB"
|
||||
)
|
||||
assert not db.try_acquire_compression_lock(
|
||||
session_id, "pid:second:contender", ttl_seconds=60
|
||||
), (
|
||||
"a second compressor acquired the durable lock DURING an "
|
||||
"admitted commit"
|
||||
)
|
||||
|
||||
# ── Release the commit; deferred release must fire promptly ──────
|
||||
release_commit.set()
|
||||
assert commit_finished.wait(timeout=5)
|
||||
deadline = time.time() + 5
|
||||
while time.time() < deadline:
|
||||
if db.get_compression_lock_holder(session_id) is None:
|
||||
break
|
||||
time.sleep(0.01)
|
||||
assert db.get_compression_lock_holder(session_id) is None, (
|
||||
"lease was not released promptly after the post-revoke commit "
|
||||
"finished (worker thread is parked, so finish_commit must have "
|
||||
"performed the deferred release)"
|
||||
)
|
||||
assert db.try_acquire_compression_lock(
|
||||
session_id, "pid:second:contender", ttl_seconds=60
|
||||
)
|
||||
db.release_compression_lock(session_id, "pid:second:contender")
|
||||
|
||||
def test_revoke_before_commit_releases_immediately_and_refuses_commit(
|
||||
self, tmp_path
|
||||
):
|
||||
"""No commit in flight → immediate release; begin_commit refused."""
|
||||
db, session_id, holder = self._db_with_lease(tmp_path)
|
||||
fence = CompressionCommitFence()
|
||||
fence.register_cancelled_lock_release(
|
||||
lambda: db.release_compression_lock(session_id, holder)
|
||||
)
|
||||
|
||||
fence.revoke_commit_admission()
|
||||
|
||||
# Release happened synchronously inside revoke — no worker involved.
|
||||
assert db.get_compression_lock_holder(session_id) is None, (
|
||||
"revoke before begin_commit must release the lease immediately"
|
||||
)
|
||||
assert fence.begin_commit() is False, (
|
||||
"begin_commit must be refused after admission was revoked"
|
||||
)
|
||||
assert db.try_acquire_compression_lock(
|
||||
session_id, "pid:second:contender", ttl_seconds=60
|
||||
)
|
||||
db.release_compression_lock(session_id, "pid:second:contender")
|
||||
@@ -0,0 +1,360 @@
|
||||
"""Compression falls back after an aborted (stalled) summary — #78981.
|
||||
|
||||
A summariser that keeps the connection open but never emits a real token
|
||||
produces no fence progress, so the host's progress-aware timeout aborts the
|
||||
worker and returns "continue without compression". Nothing raises out of the
|
||||
auxiliary client on that path, so its configured ``fallback_chain`` — the
|
||||
user's declared answer to "this route is unhealthy" — was never consulted for
|
||||
the one failure mode that most needs it.
|
||||
|
||||
These tests pin the contract:
|
||||
|
||||
* an aborted stall re-attempts compression once with the summary route pinned
|
||||
to the configured ``auxiliary.compression.fallback_chain``;
|
||||
* the pinned route reaches the summary ``call_llm`` (provider/model/base_url/
|
||||
api_key/timeout), and is single-use so the compressor's own main-model retry
|
||||
does not re-issue the same failed route;
|
||||
* the historical "continue without compression" degrade survives when no chain
|
||||
is configured or the fallback attempt also stalls.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from agent.context_compressor import (
|
||||
ContextCompressor,
|
||||
pin_summary_route,
|
||||
take_pinned_summary_route,
|
||||
)
|
||||
from agent.conversation_compression import (
|
||||
CompressionCommitFence,
|
||||
resolve_compression_fallback_route,
|
||||
run_compress_context_with_progress_timeout,
|
||||
)
|
||||
|
||||
CHAIN_ENTRY = {
|
||||
"provider": "custom",
|
||||
"model": "backup-summarizer",
|
||||
"base_url": "https://fallback.invalid/v1",
|
||||
"api_key": "sk-fallback",
|
||||
"timeout": 45,
|
||||
}
|
||||
|
||||
|
||||
def _patch_chain(chain):
|
||||
"""Pin auxiliary.compression config without touching the real config.yaml."""
|
||||
return patch(
|
||||
"agent.auxiliary_client._get_auxiliary_task_config",
|
||||
return_value={"fallback_chain": chain},
|
||||
)
|
||||
|
||||
|
||||
class _StalledSummaryWorker:
|
||||
"""A compression worker whose first attempt streams nothing at all.
|
||||
|
||||
Mirrors the reported shape: the provider holds the connection open, so the
|
||||
worker never calls ``fence.touch_progress()`` and the host's idle budget
|
||||
lapses. ``stall_attempts`` controls how many attempts hang; any later
|
||||
attempt commits a real summary.
|
||||
"""
|
||||
|
||||
def __init__(self, compressed, *, stall_attempts=1):
|
||||
self.compressed = compressed
|
||||
self.stall_attempts = stall_attempts
|
||||
self.routes = []
|
||||
self.fences = []
|
||||
self._lock = threading.Lock()
|
||||
self.release = threading.Event()
|
||||
|
||||
@property
|
||||
def attempts(self):
|
||||
return len(self.routes)
|
||||
|
||||
def __call__(self, fence: CompressionCommitFence):
|
||||
with self._lock:
|
||||
self.routes.append(take_pinned_summary_route())
|
||||
self.fences.append(fence)
|
||||
attempt = len(self.routes)
|
||||
if attempt <= self.stall_attempts:
|
||||
# Connection open, zero tokens, zero fence progress.
|
||||
self.release.wait(timeout=10)
|
||||
return ([{"role": "assistant", "content": "late"}], "late-prompt")
|
||||
if not fence.begin_commit():
|
||||
return ([{"role": "assistant", "content": "cancelled"}], "cancelled")
|
||||
try:
|
||||
return (self.compressed, "summarized-prompt")
|
||||
finally:
|
||||
fence.finish_commit()
|
||||
|
||||
|
||||
def _run(worker, *, chain, timeouts, messages, idle=0.05, ceiling=2.0):
|
||||
with _patch_chain(chain):
|
||||
return run_compress_context_with_progress_timeout(
|
||||
worker=worker,
|
||||
messages=messages,
|
||||
system_prompt_fallback="degraded-prompt",
|
||||
idle_timeout_seconds=idle,
|
||||
total_ceiling_seconds=ceiling,
|
||||
on_timeout=lambda *args: timeouts.append(args),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fence-level contract: an aborted stall consults the configured chain
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_stalled_summary_attempts_configured_fallback_chain():
|
||||
original = [{"role": "user", "content": "keep-me"}]
|
||||
compressed = [{"role": "user", "content": "summary of earlier turns"}]
|
||||
worker = _StalledSummaryWorker(compressed)
|
||||
timeouts = []
|
||||
|
||||
try:
|
||||
msgs, prompt = _run(
|
||||
worker, chain=[CHAIN_ENTRY], timeouts=timeouts, messages=original
|
||||
)
|
||||
finally:
|
||||
worker.release.set()
|
||||
|
||||
assert worker.attempts == 2, "the aborted stall must be retried once"
|
||||
assert worker.routes[0] is None, "the primary attempt is never pinned"
|
||||
pinned = worker.routes[1]
|
||||
assert pinned is not None, "the retry must carry the configured fallback route"
|
||||
assert pinned["provider"] == "custom"
|
||||
assert pinned["model"] == "backup-summarizer"
|
||||
assert msgs == compressed, "the fallback attempt's compression must be published"
|
||||
assert prompt == "summarized-prompt"
|
||||
assert not timeouts, "no continue-without-compression degrade after a recovery"
|
||||
|
||||
|
||||
def test_retry_runs_on_a_host_published_fence():
|
||||
"""The aborted fence vetoes every future commit, so the retry needs a new
|
||||
one — minted through the host so ``/stop`` admits against the attempt that
|
||||
is actually running."""
|
||||
original = [{"role": "user", "content": "keep-me"}]
|
||||
compressed = [{"role": "user", "content": "summary"}]
|
||||
worker = _StalledSummaryWorker(compressed)
|
||||
minted = []
|
||||
|
||||
def _new_fence():
|
||||
fence = CompressionCommitFence()
|
||||
minted.append(fence)
|
||||
return fence
|
||||
|
||||
try:
|
||||
with _patch_chain([CHAIN_ENTRY]):
|
||||
msgs, _prompt = run_compress_context_with_progress_timeout(
|
||||
worker=worker,
|
||||
messages=original,
|
||||
system_prompt_fallback="degraded-prompt",
|
||||
idle_timeout_seconds=0.05,
|
||||
total_ceiling_seconds=2.0,
|
||||
new_fence=_new_fence,
|
||||
)
|
||||
finally:
|
||||
worker.release.set()
|
||||
|
||||
assert msgs == compressed
|
||||
assert len(minted) == 1, "exactly one fence is minted for the one retry"
|
||||
assert worker.fences[1] is minted[0]
|
||||
assert worker.fences[1] is not worker.fences[0]
|
||||
assert worker.fences[0].is_cancelled, "the aborted attempt stays cancelled"
|
||||
|
||||
|
||||
def test_hard_interrupt_suppresses_the_fallback_attempt():
|
||||
"""An explicit stop is not an unhealthy route — don't start another
|
||||
summary on the user's behalf after they asked for the turn to end."""
|
||||
original = [{"role": "user", "content": "keep-me"}]
|
||||
worker = _StalledSummaryWorker([{"role": "user", "content": "unused"}])
|
||||
stopped = threading.Event()
|
||||
stopped.set()
|
||||
agent = SimpleNamespace(_hard_interrupt_requested=stopped)
|
||||
timeouts = []
|
||||
|
||||
try:
|
||||
with _patch_chain([CHAIN_ENTRY]):
|
||||
msgs, prompt = run_compress_context_with_progress_timeout(
|
||||
worker=worker,
|
||||
messages=original,
|
||||
system_prompt_fallback="degraded-prompt",
|
||||
idle_timeout_seconds=0.05,
|
||||
total_ceiling_seconds=2.0,
|
||||
on_timeout=lambda *args: timeouts.append(args),
|
||||
telemetry_agent=agent,
|
||||
)
|
||||
finally:
|
||||
worker.release.set()
|
||||
|
||||
assert worker.attempts == 1
|
||||
assert msgs is original
|
||||
assert prompt == "degraded-prompt"
|
||||
assert len(timeouts) == 1
|
||||
|
||||
|
||||
def test_no_fallback_chain_configured_degrades_without_retry():
|
||||
original = [{"role": "user", "content": "keep-me"}]
|
||||
worker = _StalledSummaryWorker([{"role": "user", "content": "unused"}])
|
||||
timeouts = []
|
||||
|
||||
try:
|
||||
msgs, prompt = _run(worker, chain=[], timeouts=timeouts, messages=original)
|
||||
finally:
|
||||
worker.release.set()
|
||||
|
||||
assert worker.attempts == 1, "nothing to fall back to — do not burn a retry"
|
||||
assert msgs is original
|
||||
assert prompt == "degraded-prompt"
|
||||
assert len(timeouts) == 1
|
||||
|
||||
|
||||
def test_fallback_that_also_stalls_degrades_after_one_attempt():
|
||||
original = [{"role": "user", "content": "keep-me"}]
|
||||
worker = _StalledSummaryWorker(
|
||||
[{"role": "user", "content": "unused"}], stall_attempts=2
|
||||
)
|
||||
timeouts = []
|
||||
entry = dict(CHAIN_ENTRY, timeout=0.05)
|
||||
|
||||
try:
|
||||
msgs, prompt = _run(worker, chain=[entry], timeouts=timeouts, messages=original)
|
||||
finally:
|
||||
worker.release.set()
|
||||
|
||||
assert worker.attempts == 2, "the fallback is attempted once, not in a loop"
|
||||
assert msgs is original, "no messages may be dropped when both routes stall"
|
||||
assert prompt == "degraded-prompt"
|
||||
assert len(timeouts) == 1, "the degrade must be reported exactly once"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Route resolution: a chain entry becomes an explicit summary route
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resolved_route_carries_entry_credentials_and_timeout():
|
||||
with _patch_chain([CHAIN_ENTRY]):
|
||||
route = resolve_compression_fallback_route()
|
||||
|
||||
assert route is not None
|
||||
assert route["provider"] == "custom"
|
||||
assert route["model"] == "backup-summarizer"
|
||||
assert route["base_url"] == "https://fallback.invalid/v1"
|
||||
assert route["api_key"] == "sk-fallback"
|
||||
# Per-entry timeouts already govern aux-client fallback candidates
|
||||
# (#62452); the stall retry honours the same declaration.
|
||||
assert route["timeout"] == 45.0
|
||||
|
||||
|
||||
def test_incomplete_chain_entries_are_skipped():
|
||||
chain = [
|
||||
"not-a-mapping",
|
||||
{"model": "orphan-model"}, # no provider
|
||||
{"provider": "custom"}, # no model
|
||||
CHAIN_ENTRY,
|
||||
]
|
||||
with _patch_chain(chain):
|
||||
route = resolve_compression_fallback_route()
|
||||
|
||||
assert route is not None
|
||||
assert route["model"] == "backup-summarizer"
|
||||
|
||||
|
||||
def test_no_chain_resolves_to_no_route():
|
||||
with _patch_chain([]):
|
||||
assert resolve_compression_fallback_route() is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Injection point: the pinned route reaches the summary call
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_compressor(summary_model="aux-summarizer"):
|
||||
with patch(
|
||||
"agent.context_compressor.get_model_context_length", return_value=100000
|
||||
):
|
||||
return ContextCompressor(
|
||||
model="main-model",
|
||||
quiet_mode=True,
|
||||
summary_model_override=summary_model,
|
||||
)
|
||||
|
||||
|
||||
def _msgs():
|
||||
return [
|
||||
{"role": "user", "content": "u1 " + "x" * 200},
|
||||
{"role": "assistant", "content": "a1 " + "y" * 200},
|
||||
{"role": "user", "content": "u2 " + "z" * 200},
|
||||
]
|
||||
|
||||
|
||||
def _ok_response(content="SUMMARY BODY"):
|
||||
return SimpleNamespace(
|
||||
choices=[SimpleNamespace(message=SimpleNamespace(content=content))]
|
||||
)
|
||||
|
||||
|
||||
def test_pinned_route_overrides_the_summary_call_route():
|
||||
compressor = _make_compressor()
|
||||
calls = []
|
||||
|
||||
def _fake_call_llm(**kwargs):
|
||||
calls.append(kwargs)
|
||||
return _ok_response()
|
||||
|
||||
with patch("agent.context_compressor.call_llm", side_effect=_fake_call_llm):
|
||||
with pin_summary_route(dict(CHAIN_ENTRY)):
|
||||
summary = compressor._generate_summary(_msgs())
|
||||
|
||||
assert summary and "SUMMARY BODY" in summary
|
||||
assert len(calls) == 1
|
||||
call = calls[0]
|
||||
assert call["task"] == "compression"
|
||||
assert call["provider"] == "custom"
|
||||
assert call["model"] == "backup-summarizer"
|
||||
assert call["base_url"] == "https://fallback.invalid/v1"
|
||||
assert call["api_key"] == "sk-fallback"
|
||||
assert call["timeout"] == 45
|
||||
|
||||
|
||||
def test_pinned_route_is_not_reissued_by_the_main_model_retry():
|
||||
"""The compressor's own main-model retry must not re-run the failed route."""
|
||||
compressor = _make_compressor()
|
||||
calls = []
|
||||
|
||||
def _fake_call_llm(**kwargs):
|
||||
calls.append(kwargs)
|
||||
if len(calls) == 1:
|
||||
raise TimeoutError("Request timed out.")
|
||||
return _ok_response()
|
||||
|
||||
with patch("agent.context_compressor.call_llm", side_effect=_fake_call_llm):
|
||||
with pin_summary_route(dict(CHAIN_ENTRY)):
|
||||
summary = compressor._generate_summary(_msgs())
|
||||
|
||||
assert summary and "SUMMARY BODY" in summary
|
||||
assert len(calls) == 2
|
||||
assert calls[0]["provider"] == "custom"
|
||||
assert "provider" not in calls[1], (
|
||||
"the retry must route normally, not repeat the stalled fallback route"
|
||||
)
|
||||
|
||||
|
||||
def test_unpinned_summary_call_keeps_task_routing():
|
||||
compressor = _make_compressor()
|
||||
calls = []
|
||||
|
||||
def _fake_call_llm(**kwargs):
|
||||
calls.append(kwargs)
|
||||
return _ok_response()
|
||||
|
||||
with patch("agent.context_compressor.call_llm", side_effect=_fake_call_llm):
|
||||
summary = compressor._generate_summary(_msgs())
|
||||
|
||||
assert summary
|
||||
assert calls and "provider" not in calls[0]
|
||||
assert calls[0]["model"] == "aux-summarizer"
|
||||
@@ -0,0 +1,359 @@
|
||||
"""Stall-interrupted preflight compression must persist a durable backoff.
|
||||
|
||||
#96775: an explicit /stop after the summary stream has already crossed the
|
||||
no-progress stall window restores the original transcript but must record a
|
||||
stall-specific cooldown so the next automatic turn does not re-enter the
|
||||
same strategy. An ordinary early /stop stays cooldown-neutral.
|
||||
|
||||
Native Codex app-server compaction is a sibling path (#75364) and is not
|
||||
covered here. Stall classification reads the commit fence's progress clock,
|
||||
not Chat Completions vs Responses frame shapes.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import os
|
||||
import time
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from agent.auxiliary_client import AuxiliaryExplicitCancellation
|
||||
from agent.conversation_compression import (
|
||||
STALL_INTERRUPTED_FAILURE_CLASS,
|
||||
CompressionCommitFence,
|
||||
compress_context,
|
||||
compression_attempt_stalled,
|
||||
)
|
||||
from hermes_state import SessionDB
|
||||
|
||||
|
||||
def _build_agent(tmp_path: Path, session_id: str = "STALL_INTERRUPT_96775"):
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
db.create_session(session_id, source="cli")
|
||||
with patch.dict(os.environ, {"OPENROUTER_API_KEY": "test-key"}):
|
||||
from run_agent import AIAgent
|
||||
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
quiet_mode=True,
|
||||
session_db=db,
|
||||
session_id=session_id,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
agent._compression_feasibility_checked = True
|
||||
agent.compression_in_place = True
|
||||
agent._cached_system_prompt = "sys"
|
||||
agent.context_compressor.threshold_tokens = 1_000
|
||||
return db, agent
|
||||
|
||||
|
||||
def _messages():
|
||||
return [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
|
||||
|
||||
def _age_fence(fence: CompressionCommitFence, idle_seconds: float) -> None:
|
||||
fence._last_progress = time.monotonic() - float(idle_seconds)
|
||||
|
||||
|
||||
class TestStallClassificationIsFenceIdle:
|
||||
def test_fresh_fence_is_not_stalled(self):
|
||||
fence = CompressionCommitFence()
|
||||
assert compression_attempt_stalled(
|
||||
commit_fence=fence,
|
||||
started_at=time.monotonic(),
|
||||
idle_timeout_seconds=1.0,
|
||||
) is False
|
||||
|
||||
def test_idle_fence_is_stalled(self):
|
||||
fence = CompressionCommitFence()
|
||||
_age_fence(fence, 2.0)
|
||||
assert compression_attempt_stalled(
|
||||
commit_fence=fence,
|
||||
started_at=time.monotonic(),
|
||||
idle_timeout_seconds=1.0,
|
||||
) is True
|
||||
|
||||
def test_recent_progress_is_not_stalled(self):
|
||||
fence = CompressionCommitFence()
|
||||
_age_fence(fence, 2.0)
|
||||
fence.touch_progress()
|
||||
assert compression_attempt_stalled(
|
||||
commit_fence=fence,
|
||||
started_at=time.monotonic() - 5.0,
|
||||
idle_timeout_seconds=1.0,
|
||||
) is False
|
||||
|
||||
def test_chat_and_responses_share_the_fence_clock(self):
|
||||
"""Transport-agnostic: both APIs tick the same fence or they don't."""
|
||||
chat_fence = CompressionCommitFence()
|
||||
responses_fence = CompressionCommitFence()
|
||||
_age_fence(chat_fence, 2.0)
|
||||
responses_fence.touch_progress()
|
||||
assert compression_attempt_stalled(
|
||||
commit_fence=chat_fence,
|
||||
started_at=time.monotonic(),
|
||||
idle_timeout_seconds=1.0,
|
||||
) is True
|
||||
assert compression_attempt_stalled(
|
||||
commit_fence=responses_fence,
|
||||
started_at=time.monotonic(),
|
||||
idle_timeout_seconds=1.0,
|
||||
) is False
|
||||
|
||||
|
||||
class TestEarlyStopStaysNeutral:
|
||||
def test_explicit_interrupt_before_stall_does_not_arm_cooldown(
|
||||
self, tmp_path: Path
|
||||
):
|
||||
db, agent = _build_agent(tmp_path, "EARLY_STOP_96775")
|
||||
original = _messages()
|
||||
live = copy.deepcopy(original)
|
||||
fence = CompressionCommitFence()
|
||||
|
||||
def _early_stop(messages, **_kwargs):
|
||||
messages[0]["content"] = "must be rolled back"
|
||||
raise AuxiliaryExplicitCancellation()
|
||||
|
||||
agent.context_compressor.compress = _early_stop
|
||||
|
||||
compressed, _prompt = compress_context(
|
||||
agent,
|
||||
live,
|
||||
"sys",
|
||||
approx_tokens=50_000,
|
||||
commit_fence=fence,
|
||||
)
|
||||
|
||||
assert compressed == original
|
||||
assert live == original
|
||||
assert db.get_compression_lock_holder("EARLY_STOP_96775") is None
|
||||
assert db.get_compression_failure_cooldown("EARLY_STOP_96775") is None
|
||||
assert agent.context_compressor.should_compress(50_000) is True
|
||||
db.append_message("EARLY_STOP_96775", "assistant", "still writable")
|
||||
|
||||
|
||||
class TestStallInterruptedBackoff:
|
||||
def test_aux_explicit_cancel_after_stall_persists_backoff(
|
||||
self, tmp_path: Path, monkeypatch
|
||||
):
|
||||
monkeypatch.setattr(
|
||||
"agent.conversation_compression.resolve_context_compression_timeouts",
|
||||
lambda compression_cfg=None: (1.0, 10.0),
|
||||
)
|
||||
db, agent = _build_agent(tmp_path, "STALL_AUX_CANCEL")
|
||||
original = _messages()
|
||||
live = copy.deepcopy(original)
|
||||
fence = CompressionCommitFence()
|
||||
compress_calls = {"n": 0}
|
||||
|
||||
def _stalled_then_stop(messages, **_kwargs):
|
||||
compress_calls["n"] += 1
|
||||
_age_fence(fence, 2.0)
|
||||
messages[0]["content"] = "must be rolled back"
|
||||
raise AuxiliaryExplicitCancellation()
|
||||
|
||||
agent.context_compressor.compress = _stalled_then_stop
|
||||
|
||||
compressed, _prompt = compress_context(
|
||||
agent,
|
||||
live,
|
||||
"sys",
|
||||
approx_tokens=50_000,
|
||||
commit_fence=fence,
|
||||
)
|
||||
|
||||
assert compressed == original
|
||||
assert live == original
|
||||
assert db.get_compression_lock_holder("STALL_AUX_CANCEL") is None
|
||||
state = db.get_compression_failure_cooldown("STALL_AUX_CANCEL")
|
||||
assert state is not None
|
||||
assert STALL_INTERRUPTED_FAILURE_CLASS in str(state["error"])
|
||||
assert "msgs=20" in str(state["error"])
|
||||
assert agent.context_compressor.should_compress(50_000) is False
|
||||
|
||||
compress_calls["n"] = 0
|
||||
again, _ = compress_context(
|
||||
agent,
|
||||
copy.deepcopy(original),
|
||||
"sys",
|
||||
approx_tokens=50_000,
|
||||
commit_fence=CompressionCommitFence(),
|
||||
)
|
||||
assert again == original
|
||||
assert compress_calls["n"] == 0
|
||||
|
||||
db.append_message("STALL_AUX_CANCEL", "assistant", "still writable")
|
||||
|
||||
def test_commit_fence_cancel_after_stall_persists_backoff(
|
||||
self, tmp_path: Path, monkeypatch
|
||||
):
|
||||
monkeypatch.setattr(
|
||||
"agent.conversation_compression.resolve_context_compression_timeouts",
|
||||
lambda compression_cfg=None: (1.0, 10.0),
|
||||
)
|
||||
db, agent = _build_agent(tmp_path, "STALL_FENCE_CANCEL")
|
||||
original = _messages()
|
||||
live = copy.deepcopy(original)
|
||||
fence = CompressionCommitFence()
|
||||
|
||||
def _summary_then_cancel(messages, **_kwargs):
|
||||
_age_fence(fence, 2.0)
|
||||
fence.cancel_before_commit()
|
||||
return [
|
||||
{"role": "user", "content": "[CONTEXT COMPACTION] summary"},
|
||||
{"role": "assistant", "content": "tail"},
|
||||
]
|
||||
|
||||
agent.context_compressor.compress = _summary_then_cancel
|
||||
|
||||
compressed, _prompt = compress_context(
|
||||
agent,
|
||||
live,
|
||||
"sys",
|
||||
approx_tokens=50_000,
|
||||
commit_fence=fence,
|
||||
)
|
||||
|
||||
assert compressed == original
|
||||
assert live == original
|
||||
state = db.get_compression_failure_cooldown("STALL_FENCE_CANCEL")
|
||||
assert state is not None
|
||||
assert STALL_INTERRUPTED_FAILURE_CLASS in str(state["error"])
|
||||
assert agent.context_compressor.should_compress(50_000) is False
|
||||
db.append_message("STALL_FENCE_CANCEL", "user", "next turn")
|
||||
|
||||
def test_commit_fence_cancel_with_fresh_progress_stays_neutral(
|
||||
self, tmp_path: Path, monkeypatch
|
||||
):
|
||||
monkeypatch.setattr(
|
||||
"agent.conversation_compression.resolve_context_compression_timeouts",
|
||||
lambda compression_cfg=None: (1.0, 10.0),
|
||||
)
|
||||
db, agent = _build_agent(tmp_path, "FRESH_FENCE_CANCEL")
|
||||
original = _messages()
|
||||
live = copy.deepcopy(original)
|
||||
fence = CompressionCommitFence()
|
||||
|
||||
def _healthy_then_cancel(messages, **_kwargs):
|
||||
fence.touch_progress()
|
||||
fence.cancel_before_commit()
|
||||
return [
|
||||
{"role": "user", "content": "[CONTEXT COMPACTION] summary"},
|
||||
{"role": "assistant", "content": "tail"},
|
||||
]
|
||||
|
||||
agent.context_compressor.compress = _healthy_then_cancel
|
||||
|
||||
compressed, _prompt = compress_context(
|
||||
agent,
|
||||
live,
|
||||
"sys",
|
||||
approx_tokens=50_000,
|
||||
commit_fence=fence,
|
||||
)
|
||||
|
||||
assert compressed == original
|
||||
assert db.get_compression_failure_cooldown("FRESH_FENCE_CANCEL") is None
|
||||
assert agent.context_compressor.should_compress(50_000) is True
|
||||
|
||||
def test_manual_force_bypasses_stall_backoff(
|
||||
self, tmp_path: Path, monkeypatch
|
||||
):
|
||||
monkeypatch.setattr(
|
||||
"agent.conversation_compression.resolve_context_compression_timeouts",
|
||||
lambda compression_cfg=None: (1.0, 10.0),
|
||||
)
|
||||
db, agent = _build_agent(tmp_path, "FORCE_BYPASS_STALL")
|
||||
original = _messages()
|
||||
fence = CompressionCommitFence()
|
||||
|
||||
def _stalled_then_stop(messages, **_kwargs):
|
||||
_age_fence(fence, 2.0)
|
||||
raise AuxiliaryExplicitCancellation()
|
||||
|
||||
agent.context_compressor.compress = _stalled_then_stop
|
||||
compress_context(
|
||||
agent,
|
||||
copy.deepcopy(original),
|
||||
"sys",
|
||||
approx_tokens=50_000,
|
||||
commit_fence=fence,
|
||||
)
|
||||
assert agent.context_compressor.should_compress(50_000) is False
|
||||
|
||||
forced_calls = {"n": 0}
|
||||
|
||||
def _forced(messages, **kwargs):
|
||||
forced_calls["n"] += 1
|
||||
assert kwargs.get("force") is True
|
||||
return messages
|
||||
|
||||
agent.context_compressor.compress = _forced
|
||||
compress_context(
|
||||
agent,
|
||||
copy.deepcopy(original),
|
||||
"sys",
|
||||
approx_tokens=50_000,
|
||||
force=True,
|
||||
commit_fence=CompressionCommitFence(),
|
||||
)
|
||||
assert forced_calls["n"] == 1
|
||||
|
||||
def test_stall_backoff_merges_with_longer_active_deadline(
|
||||
self, tmp_path: Path, monkeypatch
|
||||
):
|
||||
monkeypatch.setattr(
|
||||
"agent.conversation_compression.resolve_context_compression_timeouts",
|
||||
lambda compression_cfg=None: (1.0, 10.0),
|
||||
)
|
||||
db, agent = _build_agent(tmp_path, "MERGE_MAX_STALL")
|
||||
agent.context_compressor._record_compression_failure_cooldown(
|
||||
900.0, "prior timeout"
|
||||
)
|
||||
before = db.get_compression_failure_cooldown("MERGE_MAX_STALL")
|
||||
assert before is not None
|
||||
prior_until = float(before["cooldown_until"])
|
||||
|
||||
fence = CompressionCommitFence()
|
||||
|
||||
def _stalled_then_stop(messages, **_kwargs):
|
||||
_age_fence(fence, 2.0)
|
||||
raise AuxiliaryExplicitCancellation()
|
||||
|
||||
agent.context_compressor.compress = _stalled_then_stop
|
||||
# force=True so the pre-existing cooldown does not skip the attempt
|
||||
# before the stall-interrupt restore/merge path can run.
|
||||
compress_context(
|
||||
agent,
|
||||
_messages(),
|
||||
"sys",
|
||||
approx_tokens=50_000,
|
||||
force=True,
|
||||
commit_fence=fence,
|
||||
)
|
||||
|
||||
after = db.get_compression_failure_cooldown("MERGE_MAX_STALL")
|
||||
assert after is not None
|
||||
assert float(after["cooldown_until"]) >= prior_until - 1.0
|
||||
assert STALL_INTERRUPTED_FAILURE_CLASS in str(after["error"])
|
||||
|
||||
|
||||
def test_session_db_cooldown_write_does_not_shorten_longer_deadline(
|
||||
tmp_path: Path,
|
||||
):
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
db.create_session("merge-max", source="cli")
|
||||
later = time.time() + 900.0
|
||||
sooner = time.time() + 60.0
|
||||
db.record_compression_failure_cooldown("merge-max", later, "timeout")
|
||||
db.record_compression_failure_cooldown(
|
||||
"merge-max", sooner, STALL_INTERRUPTED_FAILURE_CLASS
|
||||
)
|
||||
row = db.get_compression_failure_cooldown("merge-max")
|
||||
assert row is not None
|
||||
assert float(row["cooldown_until"]) >= later - 1.0
|
||||
assert row["error"] == STALL_INTERRUPTED_FAILURE_CLASS
|
||||
@@ -0,0 +1,178 @@
|
||||
"""Host compression timeout terminates the turn before provider re-entry (#98722).
|
||||
|
||||
Salvaged from #98741 and composed with the already-merged #98424 turn-start
|
||||
fail-closed boundary:
|
||||
|
||||
- #98424 covers the TURN-START preflight (raises
|
||||
``PreflightCompressionTimedOut`` before the loop starts).
|
||||
- These tests pin the two consumers unique to #98741: the provider-overflow
|
||||
recovery path must not re-enter compression / re-send the unchanged request
|
||||
once the wait budget was spent, and the mid-turn pre-API pass must end the
|
||||
turn with the typed ``compression_exhausted`` recovery contract.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import run_agent
|
||||
from run_agent import AIAgent
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _no_compression_sleep(monkeypatch):
|
||||
import time as _time
|
||||
|
||||
monkeypatch.setattr(_time, "sleep", lambda *_a, **_k: None)
|
||||
monkeypatch.setattr("agent.retry_utils.jittered_backoff", lambda *a, **k: 0.0)
|
||||
|
||||
|
||||
def _make_agent():
|
||||
with (
|
||||
patch("model_tools.get_tool_definitions", return_value=[]),
|
||||
patch("model_tools.check_toolset_requirements", return_value={}),
|
||||
patch("agent.process_bootstrap.OpenAI"),
|
||||
):
|
||||
agent = AIAgent(
|
||||
api_key="test-key-1234567890",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
agent.client = MagicMock()
|
||||
agent._cached_system_prompt = "You are helpful."
|
||||
agent._use_prompt_caching = False
|
||||
agent.tool_delay = 0
|
||||
agent.save_trajectories = False
|
||||
agent.compression_enabled = True
|
||||
return agent
|
||||
|
||||
|
||||
def _mock_response(content="Hello", finish_reason="stop"):
|
||||
msg = SimpleNamespace(
|
||||
content=content, tool_calls=None, reasoning_content=None, reasoning=None
|
||||
)
|
||||
choice = SimpleNamespace(message=msg, finish_reason=finish_reason)
|
||||
resp = SimpleNamespace(choices=[choice], model="test/model")
|
||||
resp.usage = None
|
||||
return resp
|
||||
|
||||
|
||||
def test_overflow_recovery_timeout_ends_turn_without_provider_reentry():
|
||||
"""Provider 400 overflow + host-timed-out compression = typed terminal.
|
||||
|
||||
Before the fix, the timed-out pass was indistinguishable from an
|
||||
ordinary no-op: the loop re-sent the unchanged oversized request, the
|
||||
provider overflowed again, and compression was re-entered in the same
|
||||
turn (#98722 "Summarizing" loop).
|
||||
"""
|
||||
agent = _make_agent()
|
||||
|
||||
err_400 = Exception(
|
||||
"This model's maximum context length is 8192 tokens. However, "
|
||||
"your messages resulted in 95000 tokens. Please reduce the length "
|
||||
"of the messages."
|
||||
)
|
||||
err_400.status_code = 400
|
||||
agent.client.chat.completions.create.side_effect = [err_400, err_400]
|
||||
|
||||
compression_calls = []
|
||||
|
||||
def _timed_out(messages, _system_message, **_kwargs):
|
||||
compression_calls.append(1)
|
||||
from agent.conversation_compression import (
|
||||
mark_context_compression_timed_out,
|
||||
reset_context_compression_timeout_outcome,
|
||||
)
|
||||
|
||||
reset_context_compression_timeout_outcome(agent)
|
||||
mark_context_compression_timed_out(agent)
|
||||
return messages, agent._cached_system_prompt
|
||||
|
||||
agent._compress_context = _timed_out
|
||||
|
||||
history = [
|
||||
{"role": "user", "content": "old request"},
|
||||
{"role": "assistant", "content": "old response"},
|
||||
]
|
||||
with (
|
||||
patch.object(agent, "_persist_session"),
|
||||
patch.object(agent, "_save_trajectory"),
|
||||
patch.object(agent, "_cleanup_task_resources"),
|
||||
):
|
||||
result = agent.run_conversation("continue", conversation_history=history)
|
||||
|
||||
# Exactly one doomed send proved the overflow; the timed-out recovery
|
||||
# pass must not be followed by a second identical send.
|
||||
assert compression_calls == [1]
|
||||
assert agent.client.chat.completions.create.call_count == 1
|
||||
assert result["failed"] is True
|
||||
assert result["completed"] is False
|
||||
assert result["compression_exhausted"] is True
|
||||
assert "No messages were dropped" in result["final_response"]
|
||||
assert "No messages were dropped" in result["error"]
|
||||
|
||||
|
||||
def test_pre_api_compression_timeout_is_typed_terminal():
|
||||
"""Mid-turn pre-API pass that hits the host timeout ends the turn."""
|
||||
agent = _make_agent()
|
||||
agent.context_compressor.protect_first_n = 0
|
||||
agent.context_compressor.protect_last_n = 0
|
||||
agent.context_compressor._threshold_tokens = 1
|
||||
agent.context_compressor.should_compress = MagicMock(return_value=True)
|
||||
agent.context_compressor.should_compress_info = MagicMock(
|
||||
return_value=(True, None)
|
||||
)
|
||||
agent.context_compressor.should_defer_preflight_to_real_usage = MagicMock(
|
||||
return_value=False
|
||||
)
|
||||
agent.context_compressor.get_active_compression_failure_cooldown = MagicMock(
|
||||
return_value=None
|
||||
)
|
||||
|
||||
compression_calls = []
|
||||
|
||||
def _timed_out(messages, _system_message, **_kwargs):
|
||||
compression_calls.append(1)
|
||||
from agent.conversation_compression import (
|
||||
mark_context_compression_timed_out,
|
||||
reset_context_compression_timeout_outcome,
|
||||
)
|
||||
|
||||
reset_context_compression_timeout_outcome(agent)
|
||||
mark_context_compression_timed_out(agent)
|
||||
return messages, agent._cached_system_prompt
|
||||
|
||||
agent._compress_context = _timed_out
|
||||
|
||||
from agent.turn_context import PreflightCompressionTimedOut
|
||||
|
||||
history = [
|
||||
{"role": "user", "content": "old request"},
|
||||
{"role": "assistant", "content": "old response"},
|
||||
]
|
||||
with (
|
||||
patch.object(agent, "_persist_session"),
|
||||
patch.object(agent, "_save_trajectory"),
|
||||
patch.object(agent, "_cleanup_task_resources"),
|
||||
):
|
||||
# The turn-start boundary (#98424) may fire first and raise; both
|
||||
# outcomes satisfy the invariant under test: the unchanged oversized
|
||||
# request never reaches the provider after a host timeout.
|
||||
try:
|
||||
result = agent.run_conversation(
|
||||
"continue", conversation_history=history
|
||||
)
|
||||
except PreflightCompressionTimedOut:
|
||||
result = None
|
||||
|
||||
assert compression_calls == [1]
|
||||
agent.client.chat.completions.create.assert_not_called()
|
||||
if result is not None:
|
||||
assert result["failed"] is True
|
||||
assert result["compression_exhausted"] is True
|
||||
assert result["turn_exit_reason"] == "context_compression_timeout"
|
||||
@@ -0,0 +1,48 @@
|
||||
"""Verify compression trigger excludes reasoning/completion tokens (#12026).
|
||||
|
||||
Thinking models (GLM-5.1, QwQ, DeepSeek R1) inflate completion_tokens with
|
||||
reasoning tokens that don't consume context window space. The compression
|
||||
trigger must use only prompt_tokens so sessions aren't prematurely split.
|
||||
"""
|
||||
|
||||
import types
|
||||
|
||||
|
||||
def _make_agent_stub(prompt_tokens, completion_tokens, threshold_tokens):
|
||||
"""Create a minimal stub that exercises the compression check path."""
|
||||
compressor = types.SimpleNamespace(
|
||||
last_prompt_tokens=prompt_tokens,
|
||||
last_completion_tokens=completion_tokens,
|
||||
threshold_tokens=threshold_tokens,
|
||||
)
|
||||
# Replicate the fixed logic from run_agent.py ~line 11273
|
||||
if compressor.last_prompt_tokens > 0:
|
||||
real_tokens = compressor.last_prompt_tokens # Fixed: no completion
|
||||
else:
|
||||
real_tokens = 0
|
||||
return real_tokens, compressor
|
||||
|
||||
|
||||
class TestCompressionTriggerExcludesReasoning:
|
||||
def test_high_reasoning_tokens_should_not_trigger_compression(self):
|
||||
"""With the old bug, 40k prompt + 80k reasoning = 120k > 100k threshold.
|
||||
After the fix, only 40k prompt is compared — no compression."""
|
||||
real_tokens, comp = _make_agent_stub(
|
||||
prompt_tokens=40_000,
|
||||
completion_tokens=80_000, # reasoning-heavy model
|
||||
threshold_tokens=100_000,
|
||||
)
|
||||
assert real_tokens == 40_000
|
||||
assert real_tokens < comp.threshold_tokens, (
|
||||
"Should NOT trigger compression — only prompt tokens matter"
|
||||
)
|
||||
|
||||
|
||||
def test_zero_prompt_tokens_falls_back(self):
|
||||
"""When provider returns 0 prompt tokens, real_tokens is 0 (fallback path)."""
|
||||
real_tokens, _ = _make_agent_stub(
|
||||
prompt_tokens=0,
|
||||
completion_tokens=50_000,
|
||||
threshold_tokens=100_000,
|
||||
)
|
||||
assert real_tokens == 0
|
||||
@@ -0,0 +1,359 @@
|
||||
"""Regressions for #76354 review F3/F4/F5 — worker isolation, durable lease
|
||||
cancellation, and session ContextVar repair.
|
||||
|
||||
F3: a timed-out worker running an IN-PLACE-MUTATING context engine must not
|
||||
be able to touch the caller's live transcript — assertions run WHILE the
|
||||
worker is still blocked inside the engine (released only afterwards).
|
||||
|
||||
F4: the reviewer's exact 5-step regression — block summary indefinitely →
|
||||
host timeout → NEW compressor acquires the durable lock while the old
|
||||
summary is STILL blocked → release old worker → prove it cannot clear
|
||||
cooldown / release the new holder's lease / publish state.
|
||||
|
||||
F5: after a successful out-of-place rotation, the CALLER's session
|
||||
ContextVar resolves to the child id (get_session_env / HERMES_SESSION_ID).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from hermes_state import SessionDB
|
||||
|
||||
|
||||
def _build_agent_with_db(db: SessionDB, session_id: str, **compressor_kwargs):
|
||||
with patch.dict(os.environ, {"OPENROUTER_API_KEY": "test-key"}):
|
||||
from run_agent import AIAgent
|
||||
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
quiet_mode=True,
|
||||
session_db=db,
|
||||
session_id=session_id,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
|
||||
compressor = MagicMock()
|
||||
compressor.compress.return_value = [
|
||||
{"role": "user", "content": "[CONTEXT COMPACTION] summary"},
|
||||
{"role": "user", "content": "tail"},
|
||||
]
|
||||
compressor.compression_count = 1
|
||||
compressor.last_prompt_tokens = 0
|
||||
compressor.last_completion_tokens = 0
|
||||
compressor._last_summary_error = None
|
||||
compressor._last_compress_aborted = False
|
||||
compressor._last_aux_model_failure_model = None
|
||||
compressor._last_aux_model_failure_error = None
|
||||
compressor._last_compression_made_progress = True
|
||||
compressor._last_summary_fallback_used = False
|
||||
agent.context_compressor = compressor
|
||||
# The compressor is a stub — the one-time compression-model feasibility
|
||||
# probe would resolve a REAL auxiliary provider (credential pools, live
|
||||
# token exchange) before the engine runs. In hermetic CI there are no
|
||||
# credentials, so the probe aborts compression before the stub engine
|
||||
# ever starts and every blocked-state assertion goes vacuous. These
|
||||
# tests exercise isolation/fencing, never aux-model feasibility.
|
||||
agent._compression_feasibility_checked = True
|
||||
return agent
|
||||
|
||||
|
||||
def test_f3_mutating_engine_cannot_touch_live_transcript_after_timeout(
|
||||
tmp_path: Path, monkeypatch
|
||||
) -> None:
|
||||
"""In-place-mutating engine + host timeout → caller transcript untouched.
|
||||
|
||||
Byte-identity is asserted WHILE the worker is still blocked inside the
|
||||
engine; the worker is released only after those assertions.
|
||||
"""
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "F3_ISOLATION"
|
||||
db.create_session(session_id, source="cli")
|
||||
agent = _build_agent_with_db(db, session_id)
|
||||
agent._cached_system_prompt = "sys"
|
||||
|
||||
# Fast host timeout for the owned wrapper.
|
||||
monkeypatch.setattr(
|
||||
"agent.conversation_compression.resolve_context_compression_timeouts",
|
||||
lambda cfg=None: (0.6, 1.2),
|
||||
)
|
||||
|
||||
engine_started = threading.Event()
|
||||
release_engine = threading.Event()
|
||||
mutated_lists = []
|
||||
|
||||
def _mutating_engine(msgs, **_kwargs):
|
||||
# Legacy/plugin-engine contract: mutate the input list IN PLACE.
|
||||
engine_started.set()
|
||||
msgs[:] = [{"role": "assistant", "content": "ENGINE GARBAGE"}]
|
||||
mutated_lists.append(msgs)
|
||||
assert release_engine.wait(timeout=30)
|
||||
return msgs
|
||||
|
||||
agent.context_compressor.compress.side_effect = _mutating_engine
|
||||
|
||||
live = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
baseline = copy.deepcopy(live)
|
||||
|
||||
try:
|
||||
returned, _sp = agent._compress_context(
|
||||
live, "sys", approx_tokens=120_000
|
||||
)
|
||||
# Host timed out and returned while the engine is STILL blocked.
|
||||
assert engine_started.wait(timeout=5)
|
||||
assert not release_engine.is_set()
|
||||
assert returned is live
|
||||
# ── The core assertion, made while the worker keeps running ──────
|
||||
assert live == baseline, (
|
||||
"live transcript mutated by a detached compression worker"
|
||||
)
|
||||
# The engine did mutate a list — the SNAPSHOT, not the caller's.
|
||||
assert mutated_lists and mutated_lists[0] is not live
|
||||
# Give the blocked worker extra time to prove no delayed publication.
|
||||
time.sleep(0.2)
|
||||
assert live == baseline
|
||||
finally:
|
||||
release_engine.set()
|
||||
# After the late worker finishes, the live transcript must STILL be
|
||||
# untouched (publication only on admitted commit — which was cancelled).
|
||||
deadline = time.time() + 5
|
||||
while time.time() < deadline and db.get_compression_lock_holder(session_id):
|
||||
time.sleep(0.02)
|
||||
assert live == baseline
|
||||
|
||||
|
||||
def test_host_timeout_releases_pool_slot_while_protected_provider_is_still_blocked(
|
||||
tmp_path: Path, monkeypatch
|
||||
) -> None:
|
||||
"""Fence timeout must unwind the compression owner, not occupy the pool.
|
||||
|
||||
Protected auxiliary calls isolate their provider stream on a daemon thread.
|
||||
The compression owner must observe its commit-fence cancellation and unwind
|
||||
immediately; otherwise four slow streams consume all four shared compression
|
||||
workers until the auxiliary stream's much longer absolute ceiling expires.
|
||||
"""
|
||||
from agent import auxiliary_client as aux
|
||||
from agent import conversation_compression as cc
|
||||
|
||||
deadline = time.time() + 5
|
||||
while time.time() < deadline:
|
||||
with cc._compress_admission_lock:
|
||||
if cc._compress_admitted_count == 0:
|
||||
break
|
||||
time.sleep(0.02)
|
||||
with cc._compress_admission_lock:
|
||||
assert cc._compress_admitted_count == 0
|
||||
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "F3_PROVIDER_OWNER_RELEASE"
|
||||
db.create_session(session_id, source="cli")
|
||||
agent = _build_agent_with_db(db, session_id)
|
||||
agent._cached_system_prompt = "sys"
|
||||
monkeypatch.setattr(
|
||||
"agent.conversation_compression.resolve_context_compression_timeouts",
|
||||
# Allow provider-thread startup under the parallel runner before timing out.
|
||||
lambda cfg=None: (2.0, 4.0),
|
||||
)
|
||||
|
||||
provider_started = threading.Event()
|
||||
release_provider = threading.Event()
|
||||
|
||||
def _blocked_provider(_kwargs):
|
||||
provider_started.set()
|
||||
assert release_provider.wait(timeout=30)
|
||||
return "late-provider-result"
|
||||
|
||||
def _compress_with_protected_provider(msgs, **_kwargs):
|
||||
aux._run_protected_sync_provider_call(_blocked_provider, {})
|
||||
return msgs
|
||||
|
||||
agent.context_compressor.compress.side_effect = _compress_with_protected_provider
|
||||
live = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
|
||||
try:
|
||||
returned, _sp = agent._compress_context(
|
||||
live, "sys", approx_tokens=120_000
|
||||
)
|
||||
assert returned is live
|
||||
assert provider_started.wait(timeout=5)
|
||||
assert not release_provider.is_set()
|
||||
|
||||
deadline = time.time() + 5
|
||||
while time.time() < deadline:
|
||||
with cc._compress_admission_lock:
|
||||
if cc._compress_admitted_count == 0:
|
||||
break
|
||||
time.sleep(0.01)
|
||||
with cc._compress_admission_lock:
|
||||
assert cc._compress_admitted_count == 0, (
|
||||
"timed-out compression owner retained its shared pool slot "
|
||||
"while the isolated provider stream was still blocked"
|
||||
)
|
||||
finally:
|
||||
release_provider.set()
|
||||
deadline = time.time() + 5
|
||||
while time.time() < deadline:
|
||||
with cc._compress_admission_lock:
|
||||
if cc._compress_admitted_count == 0:
|
||||
break
|
||||
time.sleep(0.02)
|
||||
|
||||
|
||||
def test_f4_five_step_stale_holder_regression(tmp_path: Path) -> None:
|
||||
"""Reviewer's exact 5-step durable-lease regression (#76354 F4).
|
||||
|
||||
1. Block the original summary indefinitely.
|
||||
2. Let the host time out.
|
||||
3. Prove another compressor can acquire the durable lock BEFORE the
|
||||
original summary is released.
|
||||
4. Release the old worker.
|
||||
5. Prove it cannot clear cooldown, release the new holder's lease, or
|
||||
publish stale state.
|
||||
"""
|
||||
from agent.conversation_compression import (
|
||||
CompressionCommitFence,
|
||||
run_compress_context_with_progress_timeout,
|
||||
)
|
||||
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
session_id = "F4_FIVE_STEP"
|
||||
db.create_session(session_id, source="telegram")
|
||||
db.append_message(session_id, "user", "original durable")
|
||||
|
||||
agent = _build_agent_with_db(db, session_id)
|
||||
agent.compression_in_place = True
|
||||
agent._cached_system_prompt = "sys"
|
||||
|
||||
summary_started = threading.Event()
|
||||
release_summary = threading.Event()
|
||||
|
||||
def _blocked_summary(*_args, **_kwargs):
|
||||
summary_started.set()
|
||||
assert release_summary.wait(timeout=30) # step 1: blocked
|
||||
return [
|
||||
{"role": "user", "content": "[CONTEXT COMPACTION] stale summary"},
|
||||
{"role": "user", "content": "tail"},
|
||||
]
|
||||
|
||||
agent.context_compressor.compress.side_effect = _blocked_summary
|
||||
# Track cooldown-clear attempts on the OLD worker's compressor.
|
||||
cooldown_cleared = []
|
||||
agent.context_compressor._clear_compression_failure_cooldown = (
|
||||
lambda: cooldown_cleared.append(True)
|
||||
)
|
||||
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
|
||||
def _worker(fence):
|
||||
return agent._compress_context(
|
||||
messages, "sys", approx_tokens=120_000, commit_fence=fence
|
||||
)
|
||||
|
||||
# Step 2: host-owned progress wait times out while summary is blocked.
|
||||
result_msgs, _prompt = run_compress_context_with_progress_timeout(
|
||||
worker=_worker,
|
||||
messages=messages,
|
||||
system_prompt_fallback="fallback",
|
||||
idle_timeout_seconds=0.6,
|
||||
total_ceiling_seconds=1.2,
|
||||
)
|
||||
assert summary_started.wait(timeout=5)
|
||||
assert not release_summary.is_set() # old worker STILL blocked
|
||||
assert result_msgs is messages
|
||||
|
||||
# Step 3: a NEW compressor acquires the durable lock while the old
|
||||
# summary remains blocked. The host's holder-qualified release freed
|
||||
# the old lease (refresher stopped + row deleted, holder-scoped).
|
||||
new_holder = "pid:new:contender"
|
||||
deadline = time.time() + 5
|
||||
acquired = False
|
||||
while time.time() < deadline:
|
||||
if db.try_acquire_compression_lock(session_id, new_holder, ttl_seconds=60):
|
||||
acquired = True
|
||||
break
|
||||
time.sleep(0.02)
|
||||
assert acquired, (
|
||||
"a new compressor must be able to acquire the durable lock while "
|
||||
"the timed-out worker is still blocked in its summary"
|
||||
)
|
||||
assert not release_summary.is_set() # provably still step-3 state
|
||||
assert db.get_compression_lock_holder(session_id) == new_holder
|
||||
|
||||
pre_release_rows = db.get_messages_as_conversation(session_id)
|
||||
|
||||
# Step 4: release the old worker.
|
||||
release_summary.set()
|
||||
# Wait for the late worker to fully unwind (it must NOT touch the lock).
|
||||
deadline = time.time() + 5
|
||||
while time.time() < deadline:
|
||||
if db.get_compression_lock_holder(session_id) != new_holder:
|
||||
break # would be a failure — checked below
|
||||
if cooldown_cleared:
|
||||
break
|
||||
time.sleep(0.02)
|
||||
time.sleep(0.3) # settle: give the stale worker every chance to misbehave
|
||||
|
||||
# Step 5a: it cannot clear the cooldown.
|
||||
assert not cooldown_cleared, (
|
||||
"late cancelled worker cleared the compression failure cooldown"
|
||||
)
|
||||
# Step 5b: it cannot release the NEW holder's lease (holder-qualified).
|
||||
assert db.get_compression_lock_holder(session_id) == new_holder, (
|
||||
"late worker released the replacement holder's durable lease (ABA)"
|
||||
)
|
||||
# Step 5c: it cannot publish stale state — transcript unchanged, no
|
||||
# in-place compaction landed, session id did not rotate.
|
||||
post_release_rows = db.get_messages_as_conversation(session_id)
|
||||
assert post_release_rows == pre_release_rows
|
||||
assert agent.session_id == session_id
|
||||
db.release_compression_lock(session_id, new_holder)
|
||||
|
||||
|
||||
def test_f5_session_contextvar_rebound_after_rotation(
|
||||
tmp_path: Path, monkeypatch
|
||||
) -> None:
|
||||
"""Post-compression tool reads of HERMES_SESSION_ID see the CHILD id."""
|
||||
from gateway.session_context import (
|
||||
clear_session_vars,
|
||||
get_session_env,
|
||||
set_session_vars,
|
||||
)
|
||||
|
||||
db = SessionDB(db_path=tmp_path / "state.db")
|
||||
parent_sid = "F5_CTXVAR_PARENT"
|
||||
db.create_session(parent_sid, source="telegram")
|
||||
agent = _build_agent_with_db(db, parent_sid)
|
||||
agent.compression_in_place = False # rotation mode
|
||||
agent._cached_system_prompt = "sys"
|
||||
|
||||
# Enable the owned pooled wrapper so rotation happens on a WORKER thread
|
||||
# (the caller's ContextVar can only be repaired by the caller).
|
||||
monkeypatch.setattr(
|
||||
"agent.conversation_compression.resolve_context_compression_timeouts",
|
||||
lambda cfg=None: (5.0, 10.0),
|
||||
)
|
||||
|
||||
# Simulate the gateway's bound session context for the caller.
|
||||
tokens = set_session_vars(session_id=parent_sid, platform="telegram")
|
||||
try:
|
||||
assert get_session_env("HERMES_SESSION_ID") == parent_sid
|
||||
|
||||
messages = [{"role": "user", "content": f"m{i}"} for i in range(20)]
|
||||
agent._compress_context(messages, "sys", approx_tokens=120_000)
|
||||
|
||||
assert agent.session_id != parent_sid # rotation happened
|
||||
# ── The F5 contract: caller-context reads resolve to the child ──
|
||||
assert get_session_env("HERMES_SESSION_ID") == agent.session_id, (
|
||||
"caller's session ContextVar still returns the parent id after "
|
||||
"an out-of-place compression rotation"
|
||||
)
|
||||
finally:
|
||||
clear_session_vars(tokens)
|
||||
@@ -0,0 +1,74 @@
|
||||
"""Tests that _try_activate_fallback updates the context compressor."""
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from run_agent import AIAgent
|
||||
from agent.context_compressor import ContextCompressor
|
||||
|
||||
|
||||
def _make_agent_with_compressor() -> AIAgent:
|
||||
"""Build a minimal AIAgent with a context_compressor, skipping __init__."""
|
||||
agent = AIAgent.__new__(AIAgent)
|
||||
|
||||
# Primary model settings
|
||||
agent.model = "primary-model"
|
||||
agent.provider = "openrouter"
|
||||
agent.base_url = "https://openrouter.ai/api/v1"
|
||||
agent.api_key = "sk-primary"
|
||||
agent.api_mode = "chat_completions"
|
||||
agent.client = MagicMock()
|
||||
agent.quiet_mode = True
|
||||
|
||||
# Fallback config
|
||||
agent._fallback_activated = False
|
||||
agent._fallback_model = {
|
||||
"provider": "openai",
|
||||
"model": "gpt-4o",
|
||||
}
|
||||
agent._fallback_chain = [agent._fallback_model]
|
||||
agent._fallback_index = 0
|
||||
|
||||
# Context compressor with primary model values
|
||||
compressor = ContextCompressor(
|
||||
model="primary-model",
|
||||
threshold_percent=0.50,
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
api_key="sk-primary",
|
||||
provider="openrouter",
|
||||
quiet_mode=True,
|
||||
)
|
||||
agent.context_compressor = compressor
|
||||
|
||||
return agent
|
||||
|
||||
|
||||
@patch("agent.auxiliary_client.resolve_provider_client")
|
||||
@patch("agent.model_metadata.get_model_context_length", return_value=128_000)
|
||||
def test_compressor_updated_on_fallback(mock_ctx_len, mock_resolve):
|
||||
"""After fallback activation, the compressor must reflect the fallback model."""
|
||||
agent = _make_agent_with_compressor()
|
||||
|
||||
assert agent.context_compressor.model == "primary-model"
|
||||
|
||||
fb_client = MagicMock()
|
||||
fb_client.base_url = "https://api.openai.com/v1"
|
||||
fb_client.api_key = "sk-fallback"
|
||||
mock_resolve.return_value = (fb_client, None)
|
||||
|
||||
agent._is_direct_openai_url = lambda url: "api.openai.com" in url
|
||||
agent._emit_status = lambda msg: None
|
||||
|
||||
result = agent._try_activate_fallback()
|
||||
|
||||
assert result is True
|
||||
assert agent._fallback_activated is True
|
||||
|
||||
c = agent.context_compressor
|
||||
assert c.model == "gpt-4o"
|
||||
assert c.base_url == "https://api.openai.com/v1"
|
||||
assert c.api_key == "sk-fallback"
|
||||
assert c.provider == "openai"
|
||||
assert c.context_length == 128_000
|
||||
assert c.threshold_tokens == int(128_000 * c.threshold_percent)
|
||||
|
||||
|
||||
@@ -0,0 +1,145 @@
|
||||
"""Tests for interrupt handling in concurrent tool execution."""
|
||||
|
||||
import threading
|
||||
import time
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _isolate_hermes(tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path / ".hermes"))
|
||||
(tmp_path / ".hermes").mkdir(exist_ok=True)
|
||||
|
||||
|
||||
def _make_agent(monkeypatch):
|
||||
"""Create a minimal AIAgent-like object with just the methods under test."""
|
||||
monkeypatch.setenv("OPENROUTER_API_KEY", "")
|
||||
monkeypatch.setenv("HERMES_INFERENCE_PROVIDER", "")
|
||||
# Avoid full AIAgent init — just import the class and build a stub
|
||||
import run_agent as _ra
|
||||
|
||||
class _Stub:
|
||||
_interrupt_requested = False
|
||||
_interrupt_message = None
|
||||
# Bind to this thread's ident so interrupt() targets a real tid.
|
||||
_execution_thread_id = threading.current_thread().ident
|
||||
_interrupt_thread_signal_pending = False
|
||||
log_prefix = ""
|
||||
quiet_mode = True
|
||||
verbose_logging = False
|
||||
log_prefix_chars = 200
|
||||
_checkpoint_mgr = MagicMock(enabled=False)
|
||||
_subdirectory_hints = MagicMock()
|
||||
tool_progress_callback = None
|
||||
tool_start_callback = None
|
||||
tool_complete_callback = None
|
||||
_todo_store = MagicMock()
|
||||
_session_db = None
|
||||
valid_tool_names = set()
|
||||
_turns_since_memory = 0
|
||||
_iters_since_skill = 0
|
||||
_current_tool = None
|
||||
_last_activity = 0
|
||||
_print_fn = print
|
||||
# Worker-thread tracking state mirrored from AIAgent.__init__ so the
|
||||
# real interrupt() method can fan out to concurrent-tool workers.
|
||||
_active_children: list = []
|
||||
|
||||
def __init__(self):
|
||||
# Instance-level (not class-level) so each test gets a fresh set.
|
||||
self._tool_worker_threads: set = set()
|
||||
self._tool_worker_threads_lock = threading.Lock()
|
||||
self._active_children_lock = threading.Lock()
|
||||
|
||||
def _touch_activity(self, desc):
|
||||
self._last_activity = time.time()
|
||||
|
||||
def _vprint(self, msg, force=False):
|
||||
pass
|
||||
|
||||
def _safe_print(self, msg):
|
||||
pass
|
||||
|
||||
def _should_emit_quiet_tool_messages(self):
|
||||
return False
|
||||
|
||||
def _should_start_quiet_spinner(self):
|
||||
return False
|
||||
|
||||
def _has_stream_consumers(self):
|
||||
return False
|
||||
|
||||
stub = _Stub()
|
||||
# Bind the real methods under test
|
||||
stub._execute_tool_calls_concurrent = _ra.AIAgent._execute_tool_calls_concurrent.__get__(stub)
|
||||
stub.interrupt = _ra.AIAgent.interrupt.__get__(stub)
|
||||
stub.clear_interrupt = _ra.AIAgent.clear_interrupt.__get__(stub)
|
||||
# /steer injection (added in PR #12116) fires after every concurrent
|
||||
# tool batch. Stub it as a no-op — this test exercises interrupt
|
||||
# fanout, not steer injection.
|
||||
stub._apply_pending_steer_to_tool_results = lambda *a, **kw: None
|
||||
stub._invoke_tool = MagicMock(side_effect=lambda *a, **kw: '{"ok": true}')
|
||||
return stub
|
||||
|
||||
|
||||
class _FakeToolCall:
|
||||
def __init__(self, name, args="{}", call_id="tc_1"):
|
||||
self.function = MagicMock(name=name, arguments=args)
|
||||
self.function.name = name
|
||||
self.id = call_id
|
||||
|
||||
|
||||
class _FakeAssistantMsg:
|
||||
def __init__(self, tool_calls):
|
||||
self.tool_calls = tool_calls
|
||||
|
||||
|
||||
|
||||
|
||||
def test_concurrent_preflight_interrupt_skips_all(monkeypatch):
|
||||
"""When _interrupt_requested is already set before concurrent execution,
|
||||
all tools are skipped with cancellation messages."""
|
||||
agent = _make_agent(monkeypatch)
|
||||
agent._interrupt_requested = True
|
||||
|
||||
tc1 = _FakeToolCall("tool_a", call_id="tc_a")
|
||||
tc2 = _FakeToolCall("tool_b", call_id="tc_b")
|
||||
msg = _FakeAssistantMsg([tc1, tc2])
|
||||
messages = []
|
||||
|
||||
agent._execute_tool_calls_concurrent(msg, messages, "test_task")
|
||||
|
||||
assert len(messages) == 2
|
||||
assert "skipped due to user interrupt" in messages[0]["content"]
|
||||
assert "skipped due to user interrupt" in messages[1]["content"]
|
||||
# _invoke_tool should never have been called
|
||||
agent._invoke_tool.assert_not_called()
|
||||
|
||||
|
||||
|
||||
|
||||
def test_clear_interrupt_clears_worker_tids(monkeypatch):
|
||||
"""After clear_interrupt(), stale worker-tid bits must be cleared so the
|
||||
next turn's tools — which may be scheduled onto recycled tids — don't
|
||||
see a false interrupt."""
|
||||
from tools.interrupt import is_interrupted, set_interrupt
|
||||
|
||||
agent = _make_agent(monkeypatch)
|
||||
# Simulate a worker having registered but not yet exited cleanly (e.g. a
|
||||
# hypothetical bug in the tear-down). Put a fake tid in the set and
|
||||
# flag it interrupted.
|
||||
fake_tid = threading.current_thread().ident # use real tid so is_interrupted can see it
|
||||
with agent._tool_worker_threads_lock:
|
||||
agent._tool_worker_threads.add(fake_tid)
|
||||
set_interrupt(True, fake_tid)
|
||||
assert is_interrupted() is True # sanity
|
||||
|
||||
agent.clear_interrupt()
|
||||
|
||||
assert is_interrupted() is False, (
|
||||
"clear_interrupt() did not clear the interrupt bit for a tracked "
|
||||
"worker tid — stale interrupt can leak into the next turn"
|
||||
)
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
"""The concurrent completion log reports the serialized size of multimodal results.
|
||||
|
||||
The native vision fast path returns an envelope dict; the concurrent worker logged
|
||||
``len(result)`` directly, so a ~100 KB image payload showed up as ``completed
|
||||
(0.14s, 4 chars)`` — the dict key count — while the sequential path logged the real
|
||||
serialized size (#112095).
|
||||
"""
|
||||
|
||||
import logging
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from tests.agent.test_start_order_gate import ( # noqa: F401 — autouse fixture rides along
|
||||
_FakeAssistantMsg,
|
||||
_FakeToolCall,
|
||||
_isolate_hermes,
|
||||
_make_agent,
|
||||
)
|
||||
|
||||
|
||||
def test_concurrent_completion_log_reports_serialized_multimodal_size(monkeypatch, caplog):
|
||||
agent = _make_agent(monkeypatch)
|
||||
import agent.tool_executor as te
|
||||
|
||||
monkeypatch.setattr(te, "_resolve_concurrent_tool_timeout", lambda: 6.0)
|
||||
envelope = {
|
||||
"_multimodal": True,
|
||||
"content": [
|
||||
{"type": "text", "text": "Image loaded into your context — " + "x" * 800},
|
||||
{"type": "image_url", "image_url": {"url": "data:image/jpeg;base64," + "A" * 4000}},
|
||||
],
|
||||
"text_summary": "Image attached natively for the main model.",
|
||||
"meta": {"image_url": "photo.jpg", "size_bytes": 204800, "native_vision": True},
|
||||
}
|
||||
agent._tool_guardrails = MagicMock()
|
||||
agent._tool_guardrails.before_call = lambda *a, **kw: MagicMock(allows_execution=True)
|
||||
agent._invoke_tool = MagicMock(return_value=envelope)
|
||||
|
||||
msg = _FakeAssistantMsg([_FakeToolCall("vision_analyze", "tc_1")])
|
||||
with caplog.at_level(logging.INFO, logger="agent.tool_executor"):
|
||||
agent._execute_tool_calls_concurrent(msg, [], "task")
|
||||
|
||||
completed = [r.getMessage() for r in caplog.records if "vision_analyze completed (" in r.getMessage()]
|
||||
assert completed, "no completion log line for the concurrent vision_analyze call"
|
||||
# Same measurement as the sequential path's success_log_chars.
|
||||
assert f", {len(str(envelope))} chars)" in completed[0], completed[0]
|
||||
@@ -0,0 +1,142 @@
|
||||
"""Regression guard for #18028: provider content-policy / safety-filter
|
||||
blocks must classify as ``content_policy_blocked``, be non-retryable, and
|
||||
trigger the ``is_client_error`` abort path so the loop jumps straight to a
|
||||
configured fallback or surfaces a clear policy-block message — instead of
|
||||
burning ``api_max_retries`` paid attempts on a deterministic refusal and
|
||||
delivering "API failed after 3 retries" to Telegram/cron with no provider
|
||||
context.
|
||||
|
||||
Real-world symptom from the issue:
|
||||
``API call failed after 3 retries — This content was flagged for
|
||||
possible cybersecurity risk... | provider=openai-codex model=gpt-5.5``
|
||||
repeating across cron jobs and gateway sessions, with the user unable to
|
||||
tell whether the gateway was broken, the model was down, or their prompt
|
||||
was the problem.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class TestContentPolicyBlockedClassification:
|
||||
"""Verify classify_api_error returns the right shape so downstream
|
||||
recovery (fallback activation, final_response wording) fires correctly.
|
||||
"""
|
||||
|
||||
def test_openai_codex_cybersecurity_no_status(self):
|
||||
"""The reported #18028 case — SDK raises without a status code."""
|
||||
from agent.error_classifier import classify_api_error, FailoverReason
|
||||
|
||||
e = Exception(
|
||||
"This content was flagged for possible cybersecurity risk. "
|
||||
"If this seems wrong, try rephrasing your request. To get "
|
||||
"authorized for security work, join the Trusted Access for "
|
||||
"Cyber program."
|
||||
)
|
||||
result = classify_api_error(e, provider="openai-codex", model="gpt-5.5")
|
||||
# Must NOT fall into the retryable ``unknown`` bucket — that's what
|
||||
# caused the 3x retry burn.
|
||||
assert result.reason == FailoverReason.content_policy_blocked
|
||||
assert result.retryable is False
|
||||
# Recovery is fallback model, not credential rotation or compression.
|
||||
assert result.should_fallback is True
|
||||
assert result.should_compress is False
|
||||
assert result.should_rotate_credential is False
|
||||
|
||||
|
||||
|
||||
class TestContentPolicyTriggersClientErrorAbort:
|
||||
"""Mirror the ``is_client_error`` predicate in
|
||||
``agent/conversation_loop.py`` and verify
|
||||
``FailoverReason.content_policy_blocked`` resolves to True so the loop
|
||||
aborts (after attempting fallback) instead of falling into the
|
||||
retry-backoff path.
|
||||
"""
|
||||
|
||||
def _mirror_is_client_error(
|
||||
self,
|
||||
*,
|
||||
classified_retryable: bool,
|
||||
classified_reason,
|
||||
classified_should_compress: bool = False,
|
||||
is_local_validation_error: bool = False,
|
||||
is_context_length_error: bool = False,
|
||||
) -> bool:
|
||||
"""Exact shape of conversation_loop.py's is_client_error check.
|
||||
|
||||
Kept in lock-step with the source. If you change one, change both.
|
||||
"""
|
||||
from agent.error_classifier import FailoverReason
|
||||
|
||||
return (
|
||||
is_local_validation_error
|
||||
or (
|
||||
not classified_retryable
|
||||
and not classified_should_compress
|
||||
and classified_reason not in {
|
||||
FailoverReason.rate_limit,
|
||||
FailoverReason.overloaded,
|
||||
FailoverReason.context_overflow,
|
||||
FailoverReason.payload_too_large,
|
||||
FailoverReason.long_context_tier,
|
||||
FailoverReason.thinking_signature,
|
||||
}
|
||||
)
|
||||
) and not is_context_length_error
|
||||
|
||||
def test_content_policy_blocked_triggers_abort(self):
|
||||
"""Safety-filter block must reach is_client_error → fallback/abort."""
|
||||
from agent.error_classifier import FailoverReason
|
||||
|
||||
# What classify_api_error returns for a content-policy block:
|
||||
# reason=content_policy_blocked, retryable=False, should_compress=False
|
||||
assert self._mirror_is_client_error(
|
||||
classified_retryable=False,
|
||||
classified_reason=FailoverReason.content_policy_blocked,
|
||||
), (
|
||||
"FailoverReason.content_policy_blocked must trigger the "
|
||||
"is_client_error path so fallback fires immediately instead of "
|
||||
"burning api_max_retries paid attempts on a deterministic "
|
||||
"safety refusal — see #18028."
|
||||
)
|
||||
|
||||
|
||||
class TestContentPolicyPatternsAreNarrow:
|
||||
"""Defensive guard: the safety-filter patterns must not collide with
|
||||
benign error wording from billing / format / generic 400 errors. If
|
||||
these regress to ``content_policy_blocked``, recovery will route to
|
||||
the wrong code path (fallback model instead of credential rotation).
|
||||
"""
|
||||
|
||||
def test_generic_400_format_error_not_misclassified(self):
|
||||
from agent.error_classifier import classify_api_error, FailoverReason
|
||||
|
||||
class _Err(Exception):
|
||||
def __init__(self, msg, status_code):
|
||||
super().__init__(msg)
|
||||
self.status_code = status_code
|
||||
|
||||
e = _Err("Invalid request: messages must be a non-empty list", status_code=400)
|
||||
result = classify_api_error(e, provider="openai", model="gpt-4o")
|
||||
assert result.reason != FailoverReason.content_policy_blocked
|
||||
|
||||
|
||||
def test_openrouter_account_policy_block_stays_distinct(self):
|
||||
"""``provider_policy_blocked`` (OpenRouter account-level data
|
||||
policy) must remain a separate classification from
|
||||
``content_policy_blocked`` (upstream model safety filter) — they
|
||||
have different recovery strategies.
|
||||
"""
|
||||
from agent.error_classifier import classify_api_error, FailoverReason
|
||||
|
||||
class _Err(Exception):
|
||||
def __init__(self, msg, status_code):
|
||||
super().__init__(msg)
|
||||
self.status_code = status_code
|
||||
|
||||
e = _Err(
|
||||
"No endpoints available matching your guardrail restrictions "
|
||||
"and data policy",
|
||||
status_code=404,
|
||||
)
|
||||
result = classify_api_error(e, provider="openrouter", model="anthropic/claude-opus")
|
||||
assert result.reason == FailoverReason.provider_policy_blocked
|
||||
assert result.reason != FailoverReason.content_policy_blocked
|
||||
@@ -0,0 +1,91 @@
|
||||
"""Per-file context manifest (``agent/context_file_sources.py``) behind the ``/context`` Rules figure.
|
||||
|
||||
The manifest and ``build_context_files_prompt`` share one discovery walk, so the invariant under test is
|
||||
parity: a file is reported ``loaded`` iff its content appears in the built prompt.
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from agent.context_file_sources import list_context_file_sources, render_context_file_lines
|
||||
from agent.prompt_builder import build_context_files_prompt
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def project(tmp_path):
|
||||
(tmp_path / ".git").mkdir()
|
||||
return tmp_path
|
||||
|
||||
|
||||
def _by_label(sources):
|
||||
return {s["label"]: s for s in sources}
|
||||
|
||||
|
||||
def test_manifest_matches_what_the_prompt_actually_loads(project, tmp_path_factory):
|
||||
"""Every context type present at once; the ladder picks .hermes.md, the chain lists both AGENTS files,
|
||||
CLAUDE.md/.cursorrules/.cursor/rules/*.mdc are shadowed, an empty file never wins, SOUL.md rides along."""
|
||||
(project / ".hermes.md").write_text("hermes rules")
|
||||
(project / "AGENTS.md").write_text("root agents rules")
|
||||
sub = project / "pkg"
|
||||
sub.mkdir()
|
||||
(sub / "AGENTS.override.md").write_text("") # empty: falls through to AGENTS.md in the same directory
|
||||
(sub / "AGENTS.md").write_text("pkg agents rules")
|
||||
(sub / "CLAUDE.md").write_text("claude rules")
|
||||
(sub / ".cursorrules").write_text("cursor rules")
|
||||
(sub / ".cursor" / "rules").mkdir(parents=True)
|
||||
(sub / ".cursor" / "rules" / "a.mdc").write_text("mdc rule a")
|
||||
home = tmp_path_factory.mktemp("home")
|
||||
(home / "SOUL.md").write_text("identity text")
|
||||
|
||||
sources = list_context_file_sources(cwd=str(sub), home_override=home)
|
||||
prompt = build_context_files_prompt(cwd=str(sub), home_override=home)
|
||||
|
||||
statuses = {s["label"]: s["status"] for s in sources}
|
||||
assert statuses == {
|
||||
".hermes.md": "loaded", "../AGENTS.md": "shadowed", "AGENTS.override.md": "empty", "AGENTS.md": "shadowed",
|
||||
"CLAUDE.md": "shadowed", ".cursorrules": "shadowed", ".cursor/rules/a.mdc": "shadowed", "SOUL.md": "loaded",
|
||||
}
|
||||
for src in sources:
|
||||
body = Path(src["path"]).read_text().strip() if src["chars"] else ""
|
||||
assert src["loaded"] == (bool(body) and body in prompt), src
|
||||
assert all(s["est_tokens"] > 0 for s in sources if s["chars"])
|
||||
|
||||
# Same walk, other winner: drop .hermes.md and the whole AGENTS chain loads while the rest stays shadowed.
|
||||
(project / ".hermes.md").unlink()
|
||||
statuses = {s["label"]: s["status"] for s in list_context_file_sources(cwd=str(sub), home_override=home)}
|
||||
prompt = build_context_files_prompt(cwd=str(sub), home_override=home)
|
||||
assert statuses["../AGENTS.md"] == statuses["AGENTS.md"] == "loaded" and "root agents rules" in prompt
|
||||
assert statuses["CLAUDE.md"] == "shadowed" and "claude rules" not in prompt
|
||||
|
||||
|
||||
def test_truncated_and_suppressed_statuses_follow_the_builder(project, monkeypatch, tmp_path_factory):
|
||||
import agent.prompt_builder as pb
|
||||
|
||||
monkeypatch.setattr(pb, "_get_context_file_max_chars", lambda *_a: 40)
|
||||
(project / "AGENTS.md").write_text("x" * 100)
|
||||
home = tmp_path_factory.mktemp("home")
|
||||
entry = _by_label(list_context_file_sources(cwd=str(project), home_override=home))["AGENTS.md"]
|
||||
assert entry["status"] == "truncated" and entry["loaded"] is True
|
||||
assert "[...truncated AGENTS.md" in build_context_files_prompt(cwd=str(project), home_override=home)
|
||||
|
||||
# Install-tree guard: a fallback cwd (cwd=None) inside the Hermes tree lists the file but never loads it.
|
||||
monkeypatch.setattr("agent.runtime_cwd._is_install_tree", lambda _p: True)
|
||||
monkeypatch.chdir(project)
|
||||
entry = _by_label(list_context_file_sources(cwd=None, home_override=home))["AGENTS.md"]
|
||||
assert entry["status"] == "suppressed" and entry["loaded"] is False
|
||||
assert build_context_files_prompt(cwd=None, skip_soul=True) == ""
|
||||
assert _by_label(list_context_file_sources(cwd=None, allow_install_tree_fallback=True, home_override=home))[
|
||||
"AGENTS.md"]["status"] == "truncated"
|
||||
|
||||
lines = render_context_file_lines(list_context_file_sources(cwd=None, home_override=home))
|
||||
assert lines[0] == "Context files" and "AGENTS.md" in lines[1] and "install tree" in lines[1]
|
||||
assert render_context_file_lines([]) == []
|
||||
|
||||
# Injection scan: the builder swaps the body for a BLOCKED marker, so the manifest must not say "loaded".
|
||||
monkeypatch.setattr(pb, "_get_context_file_max_chars", lambda *_a: 10_000)
|
||||
monkeypatch.setattr(pb, "_scan_for_threats", lambda content, scope: ["fake-pattern"] if "evil" in content else [])
|
||||
(project / "AGENTS.md").write_text("evil")
|
||||
entry = _by_label(list_context_file_sources(cwd=str(project), home_override=home))["AGENTS.md"]
|
||||
assert entry["status"] == "blocked" and entry["loaded"] is False
|
||||
assert "[BLOCKED: AGENTS.md" in build_context_files_prompt(cwd=str(project), home_override=home)
|
||||
@@ -0,0 +1,225 @@
|
||||
"""Tests for #1630 — gateway infinite 400 failure loop prevention.
|
||||
|
||||
Verifies that:
|
||||
1. Generic 400 errors with large sessions are treated as context-length errors
|
||||
and trigger compression instead of aborting.
|
||||
2. The gateway does not persist messages when the agent fails early, preventing
|
||||
the session from growing on each failure.
|
||||
3. Context-overflow failures produce helpful error messages suggesting /compact.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 1: Agent heuristic — generic 400 with large session → compression
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGeneric400Heuristic:
|
||||
"""The agent should treat a generic 400 with a large session as a
|
||||
probable context-length error and trigger compression, not abort."""
|
||||
|
||||
def _make_agent(self):
|
||||
"""Create a minimal AIAgent for testing error handling."""
|
||||
with (
|
||||
patch("model_tools.get_tool_definitions", return_value=[]),
|
||||
patch("model_tools.check_toolset_requirements", return_value={}),
|
||||
patch("agent.process_bootstrap.OpenAI"),
|
||||
):
|
||||
from run_agent import AIAgent
|
||||
a = AIAgent(
|
||||
api_key="test-key-12345",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
a.client = MagicMock()
|
||||
a._cached_system_prompt = "You are helpful."
|
||||
a._use_prompt_caching = False
|
||||
a.compression_enabled = False
|
||||
return a
|
||||
|
||||
def test_generic_400_with_small_session_is_client_error(self):
|
||||
"""A generic 400 with a small session should still be treated
|
||||
as a non-retryable client error (not context overflow)."""
|
||||
error_msg = "error"
|
||||
status_code = 400
|
||||
approx_tokens = 1000 # Small session
|
||||
api_messages = [{"role": "user", "content": "hi"}]
|
||||
|
||||
# Simulate the phrase matching
|
||||
is_context_length_error = any(phrase in error_msg for phrase in [
|
||||
'context length', 'context size', 'maximum context',
|
||||
'token limit', 'too many tokens', 'reduce the length',
|
||||
'exceeds the limit', 'context window',
|
||||
'request entity too large',
|
||||
'prompt is too long',
|
||||
])
|
||||
assert not is_context_length_error
|
||||
|
||||
# The heuristic should NOT trigger for small sessions
|
||||
ctx_len = 200000
|
||||
is_large_session = approx_tokens > ctx_len * 0.4 or len(api_messages) > 80
|
||||
is_generic_error = len(error_msg.strip()) < 30
|
||||
assert not is_large_session # Small session → heuristic doesn't fire
|
||||
|
||||
# Both conditions true → should be treated as context overflow
|
||||
|
||||
def test_generic_400_with_many_messages_triggers_heuristic(self):
|
||||
"""A generic 400 with >80 messages should trigger the heuristic
|
||||
even if estimated tokens are low."""
|
||||
error_msg = "error"
|
||||
status_code = 400
|
||||
ctx_len = 200000
|
||||
approx_tokens = 5000 # Low token estimate
|
||||
api_messages = [{"role": "user", "content": "x"}] * 100 # > 80 messages
|
||||
|
||||
is_large_session = approx_tokens > ctx_len * 0.4 or len(api_messages) > 80
|
||||
is_generic_error = len(error_msg.strip()) < 30
|
||||
assert is_large_session
|
||||
assert is_generic_error
|
||||
|
||||
def test_specific_error_message_bypasses_heuristic(self):
|
||||
"""A 400 with a specific, long error message should NOT trigger
|
||||
the heuristic even with a large session."""
|
||||
error_msg = "invalid model: anthropic/claude-nonexistent-model is not available"
|
||||
status_code = 400
|
||||
ctx_len = 200000
|
||||
approx_tokens = 100000
|
||||
|
||||
is_generic_error = len(error_msg.strip()) < 30
|
||||
assert not is_generic_error # Long specific message → heuristic doesn't fire
|
||||
|
||||
def test_descriptive_context_error_caught_by_phrases(self):
|
||||
"""Descriptive context-length errors should still be caught by
|
||||
the existing phrase matching (not the heuristic)."""
|
||||
error_msg = "prompt is too long: 250000 tokens > 200000 maximum"
|
||||
is_context_length_error = any(phrase in error_msg for phrase in [
|
||||
'context length', 'context size', 'maximum context',
|
||||
'token limit', 'too many tokens', 'reduce the length',
|
||||
'exceeds the limit', 'context window',
|
||||
'request entity too large',
|
||||
'prompt is too long',
|
||||
])
|
||||
assert is_context_length_error
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 2: Gateway skips persistence on failed agent results
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestGatewaySkipsPersistenceOnFailure:
|
||||
"""When the agent returns failed=True with no final_response,
|
||||
the gateway should NOT persist messages to the transcript."""
|
||||
|
||||
def test_agent_failed_early_detected(self):
|
||||
"""The agent_failed_early flag is True when failed=True,
|
||||
regardless of final_response."""
|
||||
agent_result = {
|
||||
"failed": True,
|
||||
"final_response": None,
|
||||
"messages": [],
|
||||
"error": "Non-retryable client error",
|
||||
}
|
||||
agent_failed_early = bool(agent_result.get("failed"))
|
||||
assert agent_failed_early
|
||||
|
||||
|
||||
|
||||
|
||||
class TestCompressionExhaustedFlag:
|
||||
"""When compression is exhausted, the agent should set both
|
||||
failed=True and compression_exhausted=True so the gateway can
|
||||
auto-reset the session. (#9893)"""
|
||||
|
||||
def test_compression_exhausted_returns_carry_flag(self):
|
||||
"""Simulate the return dict from a compression-exhausted agent."""
|
||||
agent_result = {
|
||||
"messages": [],
|
||||
"completed": False,
|
||||
"api_calls": 3,
|
||||
"error": "Request payload too large: max compression attempts (3) reached.",
|
||||
"partial": True,
|
||||
"failed": True,
|
||||
"compression_exhausted": True,
|
||||
}
|
||||
assert agent_result.get("failed")
|
||||
assert agent_result.get("compression_exhausted")
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 3: Context-overflow error messages
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestContextOverflowErrorMessages:
|
||||
"""The gateway should produce helpful error messages when the failure
|
||||
looks like a context overflow."""
|
||||
|
||||
def test_detects_context_keywords(self):
|
||||
"""Error messages containing context-related keywords should be
|
||||
identified as context failures."""
|
||||
keywords = [
|
||||
"context length exceeded",
|
||||
"too many tokens in the prompt",
|
||||
"request entity too large",
|
||||
"payload too large for model",
|
||||
"context window exceeded",
|
||||
]
|
||||
for error_str in keywords:
|
||||
_is_ctx_fail = any(p in error_str.lower() for p in (
|
||||
"context", "token", "too large", "too long",
|
||||
"exceed", "payload",
|
||||
))
|
||||
assert _is_ctx_fail, f"Should detect: {error_str}"
|
||||
|
||||
def test_detects_generic_400_with_large_history(self):
|
||||
"""A generic 400 error code in the string with a large history
|
||||
should be flagged as context failure."""
|
||||
error_str = "error code: 400 - {'type': 'error', 'message': 'Error'}"
|
||||
history_len = 100 # Large session
|
||||
|
||||
_is_ctx_fail = any(p in error_str.lower() for p in (
|
||||
"context", "token", "too large", "too long",
|
||||
"exceed", "payload",
|
||||
)) or (
|
||||
"400" in error_str.lower()
|
||||
and history_len > 50
|
||||
)
|
||||
assert _is_ctx_fail
|
||||
|
||||
def test_unrelated_error_not_flagged(self):
|
||||
"""Unrelated errors should not be flagged as context failures."""
|
||||
error_str = "invalid api key: authentication failed"
|
||||
history_len = 10
|
||||
|
||||
_is_ctx_fail = any(p in error_str.lower() for p in (
|
||||
"context", "token", "too large", "too long",
|
||||
"exceed", "payload",
|
||||
)) or (
|
||||
"400" in error_str.lower()
|
||||
and history_len > 50
|
||||
)
|
||||
assert not _is_ctx_fail
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 4: Agent skips persistence for large failed sessions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestAgentSkipsPersistenceForLargeFailedSessions:
|
||||
"""When a 400 error occurs and the session is large, the agent
|
||||
should skip persisting to prevent the growth loop."""
|
||||
|
||||
def test_large_session_400_skips_persistence(self):
|
||||
"""Status 400 + high token count should skip persistence."""
|
||||
status_code = 400
|
||||
approx_tokens = 60000 # > 50000 threshold
|
||||
api_messages = [{"role": "user", "content": "x"}] * 10
|
||||
|
||||
should_skip = status_code == 400 and (approx_tokens > 50000 or len(api_messages) > 80)
|
||||
assert should_skip
|
||||
|
||||
|
||||
@@ -0,0 +1,128 @@
|
||||
"""Tests for context token tracking in run_agent.py's usage extraction.
|
||||
|
||||
The context counter (status bar) must show the TOTAL prompt tokens including
|
||||
Anthropic's cached portions. This is an integration test for the token
|
||||
extraction in run_conversation(), not the ContextCompressor itself (which
|
||||
is tested in tests/agent/test_context_compressor.py).
|
||||
"""
|
||||
|
||||
import sys
|
||||
import types
|
||||
from types import SimpleNamespace
|
||||
|
||||
sys.modules.setdefault("fire", types.SimpleNamespace(Fire=lambda *a, **k: None))
|
||||
sys.modules.setdefault("firecrawl", types.SimpleNamespace(Firecrawl=object))
|
||||
sys.modules.setdefault("fal_client", types.SimpleNamespace())
|
||||
|
||||
import run_agent
|
||||
|
||||
|
||||
def _patch_bootstrap(monkeypatch):
|
||||
monkeypatch.setattr("model_tools.get_tool_definitions", lambda **kwargs: [{
|
||||
"type": "function",
|
||||
"function": {"name": "t", "description": "t", "parameters": {"type": "object", "properties": {}}},
|
||||
}])
|
||||
monkeypatch.setattr("model_tools.check_toolset_requirements", lambda: {})
|
||||
|
||||
|
||||
class _FakeAnthropicClient:
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
|
||||
class _FakeOpenAIClient:
|
||||
"""Fake OpenAI client returned by mocked resolve_provider_client."""
|
||||
api_key = "fake-codex-key"
|
||||
base_url = "https://api.openai.com/v1"
|
||||
_default_headers = None
|
||||
|
||||
|
||||
def _make_agent(monkeypatch, api_mode, provider, response_fn):
|
||||
_patch_bootstrap(monkeypatch)
|
||||
if api_mode == "anthropic_messages":
|
||||
monkeypatch.setattr("agent.anthropic_adapter.build_anthropic_client", lambda k, b=None, **kwargs: _FakeAnthropicClient())
|
||||
if provider == "openai-codex":
|
||||
monkeypatch.setattr(
|
||||
"agent.auxiliary_client.resolve_provider_client",
|
||||
lambda *a, **kw: (_FakeOpenAIClient(), "test-model"),
|
||||
)
|
||||
|
||||
class _A(run_agent.AIAgent):
|
||||
def __init__(self, *a, **kw):
|
||||
kw.update(skip_context_files=True, skip_memory=True, max_iterations=4)
|
||||
super().__init__(*a, **kw)
|
||||
self._cleanup_task_resources = self._persist_session = lambda *a, **k: None
|
||||
self._save_trajectory = lambda *a, **k: None
|
||||
|
||||
def run_conversation(self, msg, conversation_history=None, task_id=None):
|
||||
self._interruptible_api_call = lambda kw: response_fn()
|
||||
self._disable_streaming = True
|
||||
return super().run_conversation(msg, conversation_history=conversation_history, task_id=task_id)
|
||||
|
||||
return _A(model="test-model", api_key="test-key", base_url="http://localhost:1234/v1", provider=provider, api_mode=api_mode)
|
||||
|
||||
|
||||
def _anthropic_resp(input_tok, output_tok, cache_read=0, cache_creation=0):
|
||||
usage_fields = {"input_tokens": input_tok, "output_tokens": output_tok}
|
||||
if cache_read:
|
||||
usage_fields["cache_read_input_tokens"] = cache_read
|
||||
if cache_creation:
|
||||
usage_fields["cache_creation_input_tokens"] = cache_creation
|
||||
return SimpleNamespace(
|
||||
content=[SimpleNamespace(type="text", text="ok")],
|
||||
stop_reason="end_turn",
|
||||
usage=SimpleNamespace(**usage_fields),
|
||||
model="claude-sonnet-4-6",
|
||||
)
|
||||
|
||||
|
||||
# -- Anthropic: cached tokens must be included --
|
||||
|
||||
def test_anthropic_cache_read_and_creation_added(monkeypatch):
|
||||
agent = _make_agent(monkeypatch, "anthropic_messages", "anthropic",
|
||||
lambda: _anthropic_resp(3, 10, cache_read=15000, cache_creation=2000))
|
||||
agent.run_conversation("hi")
|
||||
assert agent.context_compressor.last_prompt_tokens == 17003 # 3+15000+2000
|
||||
assert agent.session_prompt_tokens == 17003
|
||||
|
||||
|
||||
def test_anthropic_no_cache_fields(monkeypatch):
|
||||
agent = _make_agent(monkeypatch, "anthropic_messages", "anthropic",
|
||||
lambda: _anthropic_resp(500, 20))
|
||||
agent.run_conversation("hi")
|
||||
assert agent.context_compressor.last_prompt_tokens == 500
|
||||
|
||||
|
||||
def test_anthropic_cache_read_only(monkeypatch):
|
||||
agent = _make_agent(monkeypatch, "anthropic_messages", "anthropic",
|
||||
lambda: _anthropic_resp(5, 15, cache_read=17666, cache_creation=15))
|
||||
agent.run_conversation("hi")
|
||||
assert agent.context_compressor.last_prompt_tokens == 17686 # 5+17666+15
|
||||
|
||||
|
||||
# -- OpenAI: prompt_tokens already total --
|
||||
|
||||
def test_openai_prompt_tokens_unchanged(monkeypatch):
|
||||
resp = lambda: SimpleNamespace(
|
||||
choices=[SimpleNamespace(index=0, message=SimpleNamespace(
|
||||
role="assistant", content="ok", tool_calls=None, reasoning_content=None,
|
||||
), finish_reason="stop")],
|
||||
usage=SimpleNamespace(prompt_tokens=5000, completion_tokens=100, total_tokens=5100),
|
||||
model="gpt-4o",
|
||||
)
|
||||
agent = _make_agent(monkeypatch, "chat_completions", "openrouter", resp)
|
||||
agent.run_conversation("hi")
|
||||
assert agent.context_compressor.last_prompt_tokens == 5000
|
||||
|
||||
|
||||
# -- Codex: no cache fields, getattr returns 0 --
|
||||
|
||||
def test_codex_no_cache_fields(monkeypatch):
|
||||
resp = lambda: SimpleNamespace(
|
||||
output=[SimpleNamespace(type="message", content=[SimpleNamespace(type="output_text", text="ok")])],
|
||||
usage=SimpleNamespace(input_tokens=3000, output_tokens=50, total_tokens=3050),
|
||||
status="completed", model="gpt-5-codex",
|
||||
)
|
||||
agent = _make_agent(monkeypatch, "codex_responses", "openai-codex", resp)
|
||||
agent.run_conversation("hi")
|
||||
assert agent.context_compressor.last_prompt_tokens == 3000
|
||||
@@ -0,0 +1,252 @@
|
||||
"""Regression tests for the post-ceiling session wedge.
|
||||
|
||||
A turn that exhausts all 4 length-continuation attempts must leave the
|
||||
session usable: the next user message issues a fresh upstream request,
|
||||
inherits no continuation counter, and the partial text that WAS received
|
||||
is surfaced instead of dropped.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from hermes_constants import PARTIAL_STREAM_STUB_ID, FINISH_REASON_LENGTH
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def loop_agent():
|
||||
from run_agent import AIAgent
|
||||
with (
|
||||
patch("model_tools.get_tool_definitions", return_value=[]),
|
||||
patch("model_tools.check_toolset_requirements", return_value={}),
|
||||
patch("agent.process_bootstrap.OpenAI"),
|
||||
):
|
||||
a = AIAgent(
|
||||
api_key="test-key-1234567890",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
a.client = MagicMock()
|
||||
a._cached_system_prompt = "You are helpful."
|
||||
a._use_prompt_caching = False
|
||||
a.compression_enabled = False
|
||||
a.save_trajectories = False
|
||||
return a
|
||||
|
||||
|
||||
def _stub(content):
|
||||
from tests.agent.test_run_agent import _mock_assistant_msg
|
||||
return SimpleNamespace(
|
||||
id=PARTIAL_STREAM_STUB_ID,
|
||||
model="test/model",
|
||||
choices=[SimpleNamespace(
|
||||
index=0,
|
||||
message=_mock_assistant_msg(content=content),
|
||||
finish_reason=FINISH_REASON_LENGTH,
|
||||
)],
|
||||
usage=None,
|
||||
)
|
||||
|
||||
|
||||
def _run(agent, message, history=None):
|
||||
with (
|
||||
patch.object(agent, "_persist_session"),
|
||||
patch.object(agent, "_save_trajectory"),
|
||||
patch.object(agent, "_cleanup_task_resources"),
|
||||
):
|
||||
return agent.run_conversation(message, conversation_history=history)
|
||||
|
||||
|
||||
class TestContinuationCeilingWedge:
|
||||
def _exhaust_ceiling(self, agent):
|
||||
agent.client.chat.completions.create.side_effect = [
|
||||
_stub("part one "), _stub("part two "),
|
||||
_stub("part three "), _stub("part four."),
|
||||
]
|
||||
return _run(agent, "write me a long report")
|
||||
|
||||
def test_partial_text_surfaced_at_ceiling(self, loop_agent):
|
||||
result = self._exhaust_ceiling(loop_agent)
|
||||
assert result["completed"] is False
|
||||
assert result["partial"] is True
|
||||
assert "part one" in (result["final_response"] or "")
|
||||
assert "part four" in (result["final_response"] or "")
|
||||
|
||||
def test_new_user_message_issues_fresh_request(self, loop_agent):
|
||||
"""Core regression: after the ceiling, a new user turn must reach
|
||||
the provider instead of replaying wedge state."""
|
||||
from tests.agent.test_run_agent import _mock_response
|
||||
|
||||
result1 = self._exhaust_ceiling(loop_agent)
|
||||
assert "truncated after 4 continuation attempts" in (result1.get("error") or "")
|
||||
calls_after_turn1 = loop_agent.client.chat.completions.create.call_count
|
||||
assert calls_after_turn1 == 4
|
||||
|
||||
loop_agent.client.chat.completions.create.side_effect = [
|
||||
_mock_response(content="Hello! How can I help?", finish_reason="stop"),
|
||||
]
|
||||
result2 = _run(loop_agent, "hi", history=result1["messages"])
|
||||
|
||||
assert loop_agent.client.chat.completions.create.call_count == calls_after_turn1 + 1, (
|
||||
"A new user message after the continuation ceiling must issue "
|
||||
"exactly one fresh upstream request."
|
||||
)
|
||||
assert result2["completed"] is True
|
||||
assert result2["final_response"] == "Hello! How can I help?"
|
||||
assert not result2.get("error")
|
||||
|
||||
def test_ceiling_replaces_scaffolding_with_settled_turn(self, loop_agent):
|
||||
"""The persisted tail must not keep the continuation scaffolding.
|
||||
Unanswered "continue" nudges make every later turn resume the
|
||||
truncated response and re-exhaust the same ceiling."""
|
||||
result = self._exhaust_ceiling(loop_agent)
|
||||
msgs = result["messages"]
|
||||
|
||||
nudges = [
|
||||
m for m in msgs
|
||||
if m.get("role") == "user"
|
||||
and "Continue exactly where you left off" in (m.get("content") or "")
|
||||
]
|
||||
assert nudges == [], (
|
||||
"Continuation nudges must not survive the ceiling exit — they "
|
||||
"steer every subsequent turn back into the truncated response."
|
||||
)
|
||||
|
||||
assistants = [m for m in msgs if m.get("role") == "assistant"]
|
||||
assert len(assistants) == 1, (
|
||||
"The fragment trail must collapse into one settled assistant turn."
|
||||
)
|
||||
assert msgs[-1]["role"] == "assistant"
|
||||
content = msgs[-1]["content"] or ""
|
||||
for part in ("part one", "part two", "part three", "part four"):
|
||||
assert part in content, "Stitched partial must keep every fragment."
|
||||
|
||||
def test_ceiling_not_labeled_network_error(self, loop_agent):
|
||||
"""A finish_reason='length' stub is a truncation, not a network
|
||||
error — the user-facing message must not blame the network."""
|
||||
printed = []
|
||||
original = loop_agent._vprint
|
||||
|
||||
def _capture(text, **kwargs):
|
||||
printed.append(str(text))
|
||||
return original(text, **kwargs)
|
||||
|
||||
with patch.object(loop_agent, "_vprint", side_effect=_capture):
|
||||
self._exhaust_ceiling(loop_agent)
|
||||
|
||||
network_lines = [line for line in printed if "network error" in line.lower()]
|
||||
assert network_lines == [], (
|
||||
"Truncation must not be reported as a network error: "
|
||||
f"{network_lines!r}"
|
||||
)
|
||||
assert any("truncated" in line.lower() for line in printed), (
|
||||
"The user-facing message must name the truncation."
|
||||
)
|
||||
|
||||
def test_continuation_requests_carry_no_marks(self, loop_agent):
|
||||
"""The scaffolding marks are Hermes bookkeeping. The centrally
|
||||
sanitized api_messages must never carry them — only the
|
||||
chat-completions transport strips underscore keys, so anthropic
|
||||
and bedrock requests would otherwise send them to the provider."""
|
||||
from tests.agent.test_run_agent import _mock_response
|
||||
|
||||
seen_api_messages = []
|
||||
original = loop_agent._build_api_kwargs
|
||||
|
||||
def _spy(api_messages, tools_for_api=None):
|
||||
seen_api_messages.append([dict(m) for m in api_messages if isinstance(m, dict)])
|
||||
return original(api_messages, tools_for_api=tools_for_api)
|
||||
|
||||
loop_agent.client.chat.completions.create.side_effect = [
|
||||
_stub("part one "), _stub("part two "),
|
||||
_mock_response(content="the rest.", finish_reason="stop"),
|
||||
]
|
||||
with patch.object(loop_agent, "_build_api_kwargs", side_effect=_spy):
|
||||
result = _run(loop_agent, "write me a long report")
|
||||
|
||||
assert result["completed"] is True
|
||||
assert len(seen_api_messages) >= 3, "Expected continuation attempts 2+."
|
||||
marked = [
|
||||
(idx, key)
|
||||
for idx, batch in enumerate(seen_api_messages)
|
||||
for m in batch
|
||||
for key in m
|
||||
if str(key).startswith("_length_continuation")
|
||||
]
|
||||
assert marked == [], (
|
||||
f"Continuation marks leaked into outgoing api_messages: {marked!r}"
|
||||
)
|
||||
|
||||
def test_prior_turn_marked_message_survives_later_ceiling(self, loop_agent):
|
||||
"""A mark that reached disk mid-crash and got reloaded on a PRIOR
|
||||
turn's message must never be deleted by a later turn's ceiling
|
||||
cleanup — the cleanup is scoped to the current turn."""
|
||||
reloaded_history = [
|
||||
{"role": "user", "content": "earlier question"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "earlier answer fragment",
|
||||
"_length_continuation_fragment": True,
|
||||
},
|
||||
]
|
||||
loop_agent.client.chat.completions.create.side_effect = [
|
||||
_stub("wedge one "), _stub("wedge two "),
|
||||
_stub("wedge three "), _stub("wedge four."),
|
||||
]
|
||||
result = _run(loop_agent, "another long report", history=reloaded_history)
|
||||
|
||||
assert "truncated after 4 continuation attempts" in (result.get("error") or "")
|
||||
prior = [
|
||||
m for m in result["messages"]
|
||||
if m.get("role") == "assistant"
|
||||
and "earlier answer fragment" in (m.get("content") or "")
|
||||
]
|
||||
assert len(prior) == 1, (
|
||||
"The prior turn's reloaded message must survive the later "
|
||||
"turn's ceiling cleanup."
|
||||
)
|
||||
|
||||
def test_new_turn_does_not_inherit_continuation_counter(self, loop_agent):
|
||||
"""A single truncation on the turn AFTER the ceiling must get its
|
||||
own full 4-attempt budget, not the exhausted counter."""
|
||||
from tests.agent.test_run_agent import _mock_response
|
||||
|
||||
result1 = self._exhaust_ceiling(loop_agent)
|
||||
loop_agent.client.chat.completions.create.side_effect = [
|
||||
_stub("second turn partial "),
|
||||
_mock_response(content="and the rest.", finish_reason="stop"),
|
||||
]
|
||||
result2 = _run(loop_agent, "try again", history=result1["messages"])
|
||||
|
||||
assert result2["completed"] is True, (
|
||||
"One truncation on a fresh turn must continue (1/4), not fail "
|
||||
"with an inherited exhausted counter."
|
||||
)
|
||||
assert "second turn partial" in result2["final_response"]
|
||||
assert "and the rest." in result2["final_response"]
|
||||
|
||||
|
||||
class TestTruncatedPartJoining:
|
||||
"""#78577 — parts joined with no separator glued text together."""
|
||||
|
||||
def test_glued_parts_get_a_newline(self):
|
||||
from agent.conversation_loop import _join_truncated_parts
|
||||
assert _join_truncated_parts(
|
||||
["Edited index.html", "Review the 5 changes"]
|
||||
) == "Edited index.html\nReview the 5 changes"
|
||||
|
||||
def test_existing_whitespace_is_not_doubled(self):
|
||||
from agent.conversation_loop import _join_truncated_parts
|
||||
assert _join_truncated_parts(["line one\n", "line two"]) == "line one\nline two"
|
||||
assert _join_truncated_parts(["word", " next"]) == "word next"
|
||||
|
||||
def test_degenerate_inputs(self):
|
||||
from agent.conversation_loop import _join_truncated_parts
|
||||
assert _join_truncated_parts([]) == ""
|
||||
assert _join_truncated_parts(["only"]) == "only"
|
||||
assert _join_truncated_parts(["a", "", "b"]) == "a\nb"
|
||||
@@ -0,0 +1,99 @@
|
||||
"""Regression tests for the truncated-response repetition guard (#86581).
|
||||
|
||||
A truncated response (``finish_reason=length``) dominated by verbatim
|
||||
repeated text must NOT be continued: the continuation nudge would stitch the
|
||||
pathological fragment into the final response (the #86581 incident delivered
|
||||
60,698 chars as 31 Discord messages). The turn aborts with a clear
|
||||
user-facing error instead, mirroring the existing thinking-budget guard.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from hermes_constants import FINISH_REASON_LENGTH, PARTIAL_STREAM_STUB_ID
|
||||
|
||||
# The exact sentence from the #86581 incident.
|
||||
_INCIDENT_ECHO = "好,你幫我更改成 Google Gemini 4 31B。"
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def loop_agent():
|
||||
from run_agent import AIAgent
|
||||
|
||||
with (
|
||||
patch("model_tools.get_tool_definitions", return_value=[]),
|
||||
patch("model_tools.check_toolset_requirements", return_value={}),
|
||||
patch("agent.process_bootstrap.OpenAI"),
|
||||
):
|
||||
a = AIAgent(
|
||||
api_key="test-key-1234567890",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
a.client = MagicMock()
|
||||
a._cached_system_prompt = "You are helpful."
|
||||
a._use_prompt_caching = False
|
||||
a.compression_enabled = False
|
||||
a.save_trajectories = False
|
||||
return a
|
||||
|
||||
|
||||
def _stub(content):
|
||||
from tests.agent.test_run_agent import _mock_assistant_msg
|
||||
|
||||
return SimpleNamespace(
|
||||
id=PARTIAL_STREAM_STUB_ID,
|
||||
model="test/model",
|
||||
choices=[SimpleNamespace(
|
||||
index=0,
|
||||
message=_mock_assistant_msg(content=content),
|
||||
finish_reason=FINISH_REASON_LENGTH,
|
||||
)],
|
||||
usage=None,
|
||||
)
|
||||
|
||||
|
||||
def _run(agent, message):
|
||||
with (
|
||||
patch.object(agent, "_persist_session"),
|
||||
patch.object(agent, "_save_trajectory"),
|
||||
patch.object(agent, "_cleanup_task_resources"),
|
||||
):
|
||||
return agent.run_conversation(message)
|
||||
|
||||
|
||||
class TestContinuationRepetitionGuard:
|
||||
def test_repetition_dominated_truncation_aborts(self, loop_agent):
|
||||
echo = _INCIDENT_ECHO * 2000
|
||||
loop_agent.client.chat.completions.create.side_effect = [_stub(echo)]
|
||||
|
||||
result = _run(loop_agent, "write me a long report")
|
||||
|
||||
assert result["completed"] is False
|
||||
assert result["partial"] is True
|
||||
assert "Repetition" in (result["final_response"] or "")
|
||||
# The pathological fragment must NOT be appended to the history.
|
||||
assert not any(
|
||||
isinstance(m, dict) and m.get("_length_continuation_fragment")
|
||||
for m in result["messages"]
|
||||
)
|
||||
# Exactly one API call — no continuation was attempted.
|
||||
assert loop_agent.client.chat.completions.create.call_count == 1
|
||||
|
||||
def test_legit_truncation_still_continues(self, loop_agent):
|
||||
# Ordinary short truncated fragments still get continuation retries.
|
||||
loop_agent.client.chat.completions.create.side_effect = [
|
||||
_stub("part one "), _stub("part two "),
|
||||
_stub("part three "), _stub("part four."),
|
||||
]
|
||||
|
||||
result = _run(loop_agent, "write me a long report")
|
||||
|
||||
assert result["partial"] is True
|
||||
assert loop_agent.client.chat.completions.create.call_count == 4
|
||||
@@ -0,0 +1,112 @@
|
||||
"""Invariants for the shared manual-/compress core (``agent/conversation_compression_manual``)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import threading
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from agent.conversation_compression_manual import compress_now, parse_compress_args
|
||||
|
||||
|
||||
def _history():
|
||||
return [
|
||||
{"role": "user", "content": "one"}, {"role": "assistant", "content": "two"},
|
||||
{"role": "user", "content": "three"}, {"role": "assistant", "content": "four"},
|
||||
{"role": "user", "content": "five"}, {"role": "assistant", "content": "six"},
|
||||
]
|
||||
|
||||
|
||||
def _agent():
|
||||
agent = MagicMock()
|
||||
agent._cached_system_prompt, agent.tools, agent.context_compressor = "sys", None, None
|
||||
agent._compression_skipped_due_to_lock = None
|
||||
agent._compress_context.return_value = ([{"role": "assistant", "content": "summary"}], "")
|
||||
return agent
|
||||
|
||||
|
||||
@pytest.mark.parametrize("raw", ["--preview", "here 1 --preview", "--dry-run keep the tests", "--aggressive --preview"])
|
||||
def test_preview_leaves_history_and_agent_byte_identical(raw):
|
||||
agent, history = _agent(), _history()
|
||||
frozen = copy.deepcopy(history)
|
||||
result = compress_now(agent, history, parse_compress_args(raw))
|
||||
assert result.status == "preview" and result.lines
|
||||
assert history == frozen and result.after_messages == frozen
|
||||
agent._compress_context.assert_not_called()
|
||||
|
||||
|
||||
def test_compressed_result_rejoins_verbatim_tail_and_never_mutates_input():
|
||||
agent, history = _agent(), _history()
|
||||
frozen = copy.deepcopy(history)
|
||||
result = compress_now(agent, history, parse_compress_args("here 1"))
|
||||
assert result.status == "compressed"
|
||||
assert agent._compress_context.call_args.args[0] == frozen[:4] # head only
|
||||
assert result.after_messages[-2:] == frozen[4:] # last exchange verbatim
|
||||
assert history == frozen
|
||||
assert agent._compress_context.call_args.kwargs["force"] is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize("surface", ["cli", "gateway", "tui", "acp"])
|
||||
def test_every_surface_honours_preview_without_compressing(surface, monkeypatch):
|
||||
"""The prompt-cache-breaking mutation must be gated by the same ``--preview`` on all four surfaces."""
|
||||
agent, history = _agent(), _history()
|
||||
frozen = copy.deepcopy(history)
|
||||
if surface == "cli":
|
||||
from hermes_cli.cli_session_mixin import CLISessionMixin
|
||||
cli = CLISessionMixin.__new__(CLISessionMixin)
|
||||
cli.agent, cli.conversation_history = agent, history
|
||||
cli._manual_compress("/compress --preview")
|
||||
assert cli.conversation_history == frozen
|
||||
elif surface == "gateway":
|
||||
import asyncio
|
||||
from gateway.run import GatewayRunner
|
||||
gw = GatewayRunner.__new__(GatewayRunner)
|
||||
gw.session_store = MagicMock()
|
||||
entry = MagicMock(session_id="sid")
|
||||
gw._async_session_store = MagicMock(
|
||||
_store=gw.session_store, get_or_create_session=_coro(entry), load_transcript=_coro(history))
|
||||
gw._run_manual_compression = MagicMock(side_effect=AssertionError("must not compress"))
|
||||
event = MagicMock(); event.get_command_args.return_value = "--preview"
|
||||
reply = asyncio.run(gw._handle_compress_command_inner(event))
|
||||
assert "Preview" in reply
|
||||
elif surface == "tui":
|
||||
from tui_gateway.server import _compress_session_history
|
||||
session = {"agent": agent, "history": history, "history_lock": threading.Lock(), "history_version": 3}
|
||||
assert _compress_session_history(session, "--preview")[0] == 0
|
||||
assert session["history"] == frozen and session["history_version"] == 3
|
||||
else:
|
||||
from acp_adapter.commands import SlashCommandsMixin
|
||||
acp = SlashCommandsMixin.__new__(SlashCommandsMixin)
|
||||
acp.session_manager = MagicMock()
|
||||
state = MagicMock(history=history, agent=agent, session_id="acp-sid")
|
||||
assert "Preview" in acp._cmd_compress("--preview", state)
|
||||
assert state.history == frozen
|
||||
acp.session_manager.save_session.assert_not_called()
|
||||
agent._compress_context.assert_not_called()
|
||||
|
||||
|
||||
def test_windowless_gate_is_gateway_only_so_in_process_surfaces_still_reach_compress_context():
|
||||
"""``has_content_to_compress`` only knows the local summary window; codex_app_server native compaction
|
||||
and the phase-1 tool-result prune inside ``_compress_context`` do useful work without one, so CLI/TUI/
|
||||
ACP (the default) must not short-circuit on it. The gateway keeps its historical early answer."""
|
||||
agent, history = _agent(), _history()
|
||||
agent.context_compressor = MagicMock()
|
||||
agent.context_compressor.has_content_to_compress.return_value = False
|
||||
agent._compress_context.return_value = (history[:-1], "") # e.g. a pruned tool result, no summary
|
||||
|
||||
default = compress_now(agent, history, parse_compress_args(""))
|
||||
assert default.status == "compressed" and default.removed == 1
|
||||
agent._compress_context.assert_called_once()
|
||||
|
||||
agent._compress_context.reset_mock()
|
||||
gateway = compress_now(agent, history, parse_compress_args(""), system_message="", skip_without_window=True)
|
||||
assert gateway.status == "nothing_to_do" and gateway.after_messages == history
|
||||
agent._compress_context.assert_not_called()
|
||||
|
||||
|
||||
def _coro(value):
|
||||
async def _inner(*_a, **_k):
|
||||
return value
|
||||
return _inner
|
||||
@@ -0,0 +1,207 @@
|
||||
"""Regression tests for conversation loop fallback state management."""
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from run_agent import AIAgent
|
||||
|
||||
|
||||
def _tool_defs(*names):
|
||||
"""Helper: create minimal tool definitions for given names."""
|
||||
return [
|
||||
{
|
||||
"type": "function", "function": {
|
||||
"name": name,
|
||||
"description": "test tool",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
}
|
||||
}
|
||||
for name in names
|
||||
]
|
||||
|
||||
|
||||
def _tool_call(name, call_id):
|
||||
"""Helper: create a minimal tool call object."""
|
||||
return SimpleNamespace(
|
||||
id=call_id, type="function",
|
||||
function=SimpleNamespace(name=name, arguments="{}"),
|
||||
)
|
||||
|
||||
|
||||
def _response(*, content, finish_reason, tool_calls=None):
|
||||
"""Helper: create a minimal API response object."""
|
||||
message = SimpleNamespace(content=content, tool_calls=tool_calls)
|
||||
choice = SimpleNamespace(message=message, finish_reason=finish_reason)
|
||||
return SimpleNamespace(choices=[choice], model="test/model", usage=None)
|
||||
|
||||
|
||||
def test_substantive_tool_only_turn_invalidates_older_housekeeping_fallback():
|
||||
"""
|
||||
Regression test for #63860.
|
||||
|
||||
A cached `_last_content_with_tools` response from a housekeeping-only turn
|
||||
must not survive a later substantive tool-only turn. When the model returns
|
||||
an empty response after the substantive tool turn, the system should enter
|
||||
the post-tool nudge path, not use the stale housekeeping fallback.
|
||||
|
||||
Production impact: scheduled cron jobs could return early without
|
||||
completing their actual work (e.g., daily report job returning a
|
||||
housekeeping message instead of producing the report artifact).
|
||||
|
||||
Test sequence:
|
||||
1. Content + todo (housekeeping) → sets fallback, marks as all-housekeeping
|
||||
2. Empty content + web_search (substantive) → should CLEAR old fallback
|
||||
3. Empty content, no tool calls → should enter post-tool nudge, not use old fallback
|
||||
4. Content "Recovered after nudge." → should be returned as final response
|
||||
|
||||
Before the fix:
|
||||
- Step 2 would not clear the fallback state (no visible content)
|
||||
- Step 3 would incorrectly use the housekeeping fallback from step 1
|
||||
- API calls would stop at 3, never reaching the nudge response
|
||||
|
||||
After the fix:
|
||||
- Step 2 classifies tools and clears the fallback because web_search is substantive
|
||||
- Step 3 enters the post-tool nudge path (no stale housekeeping fallback available)
|
||||
- Step 4 returns the nudge response as the final answer
|
||||
"""
|
||||
with (
|
||||
patch("model_tools.get_tool_definitions", return_value=_tool_defs("todo", "web_search")),
|
||||
patch("model_tools.check_toolset_requirements", return_value={}),
|
||||
patch("agent.process_bootstrap.OpenAI"),
|
||||
):
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1/",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
|
||||
agent._cached_system_prompt = "You are helpful."
|
||||
agent._use_prompt_caching = False
|
||||
agent.compression_enabled = False
|
||||
agent.save_trajectories = False
|
||||
agent.valid_tool_names = {"todo", "web_search"}
|
||||
agent.client = MagicMock()
|
||||
agent.client.chat.completions.create.side_effect = [
|
||||
# Turn 1: Content + housekeeping tool
|
||||
_response(
|
||||
content="I'll begin the work.",
|
||||
finish_reason="tool_calls",
|
||||
tool_calls=[_tool_call("todo", "todo1")],
|
||||
),
|
||||
# Turn 2: Empty content + substantive tool (should clear stale fallback)
|
||||
_response(
|
||||
content="",
|
||||
finish_reason="tool_calls",
|
||||
tool_calls=[_tool_call("web_search", "search1")],
|
||||
),
|
||||
# Turn 3: Empty response (should enter nudge path, not use stale fallback)
|
||||
_response(content="", finish_reason="stop"),
|
||||
# Turn 4: Nudge response
|
||||
_response(content="Recovered after nudge.", finish_reason="stop"),
|
||||
]
|
||||
|
||||
with (
|
||||
patch("model_tools.handle_function_call", return_value="ok"),
|
||||
patch.object(agent, "_persist_session"),
|
||||
patch.object(agent, "_save_trajectory"),
|
||||
patch.object(agent, "_cleanup_task_resources"),
|
||||
):
|
||||
result = agent.run_conversation("do the full task")
|
||||
|
||||
assert result["final_response"] == "Recovered after nudge.", (
|
||||
f"Expected nudge recovery response, got: {result['final_response']}. "
|
||||
f"This indicates the stale housekeeping fallback was incorrectly used."
|
||||
)
|
||||
assert result["api_calls"] == 4, (
|
||||
f"Expected 4 API calls (including nudge), got: {result['api_calls']}. "
|
||||
f"This indicates the conversation exited early without retrying."
|
||||
)
|
||||
assert result["turn_exit_reason"].startswith("text_response"), (
|
||||
f"Expected text_response exit, got: {result['turn_exit_reason']}. "
|
||||
f"This indicates the wrong fallback path was taken."
|
||||
)
|
||||
|
||||
|
||||
def test_bare_tool_marker_is_not_reused_as_final_response():
|
||||
"""
|
||||
Regression test for #78148.
|
||||
|
||||
A provider/local template can emit a bare bracketed token (e.g. "[memory]")
|
||||
as assistant content alongside a tool call. That token is protocol
|
||||
scaffolding, not an answer. If it gets cached as `_last_content_with_tools`
|
||||
and the following turn is empty, the post-tool fallback replays it as the
|
||||
final response — and because it then enters the persisted transcript,
|
||||
later context compaction preserves it, letting the model repeat the
|
||||
marker in subsequent turns.
|
||||
|
||||
Test sequence:
|
||||
1. Content "[memory]" + skill_manage (housekeeping) tool call → the bare
|
||||
marker must be discarded, not cached as a fallback.
|
||||
2. Empty content, no tool calls → enters the post-tool nudge path since
|
||||
no fallback is available.
|
||||
3. Content "Recovered after nudge." → returned as the final response.
|
||||
|
||||
Before the fix:
|
||||
- Step 1 cached "[memory]" as `_last_content_with_tools`.
|
||||
- Step 2 reused it via the empty-response fallback, so the conversation
|
||||
never reached step 3 and "[memory]" leaked into the persisted history.
|
||||
|
||||
After the fix:
|
||||
- Step 1 strips the bare marker before it is cached or persisted.
|
||||
- Step 2 has no fallback available and enters the nudge path instead.
|
||||
- Step 3 returns the nudge response as the final answer.
|
||||
"""
|
||||
with (
|
||||
patch("model_tools.get_tool_definitions", return_value=_tool_defs("skill_manage")),
|
||||
patch("model_tools.check_toolset_requirements", return_value={}),
|
||||
patch("agent.process_bootstrap.OpenAI"),
|
||||
):
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1/",
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
|
||||
agent._cached_system_prompt = "You are helpful."
|
||||
agent._use_prompt_caching = False
|
||||
agent.compression_enabled = False
|
||||
agent.save_trajectories = False
|
||||
agent.valid_tool_names = {"skill_manage"}
|
||||
agent.client = MagicMock()
|
||||
agent.client.chat.completions.create.side_effect = [
|
||||
# Turn 1: Bare "[memory]" marker + housekeeping tool call.
|
||||
_response(
|
||||
content="[memory]",
|
||||
finish_reason="tool_calls",
|
||||
tool_calls=[_tool_call("skill_manage", "skill1")],
|
||||
),
|
||||
# Turn 2: Empty response (should enter nudge path, not reuse "[memory]").
|
||||
_response(content="", finish_reason="stop"),
|
||||
# Turn 3: Nudge response
|
||||
_response(content="Recovered after nudge.", finish_reason="stop"),
|
||||
]
|
||||
|
||||
with (
|
||||
patch("model_tools.handle_function_call", return_value="ok"),
|
||||
patch.object(agent, "_persist_session"),
|
||||
patch.object(agent, "_save_trajectory"),
|
||||
patch.object(agent, "_cleanup_task_resources"),
|
||||
):
|
||||
result = agent.run_conversation("do the full task")
|
||||
|
||||
assert result["final_response"] != "[memory]", (
|
||||
"The bare tool-call marker leaked through as the final response — "
|
||||
"it should have been discarded before caching/persistence."
|
||||
)
|
||||
assert result["final_response"] == "Recovered after nudge.", (
|
||||
f"Expected nudge recovery response, got: {result['final_response']}."
|
||||
)
|
||||
assert result["api_calls"] == 3, (
|
||||
f"Expected 3 API calls (including nudge), got: {result['api_calls']}."
|
||||
)
|
||||
|
||||
@@ -5,6 +5,7 @@ from __future__ import annotations
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
@@ -421,3 +422,72 @@ def test_run_prompt_receives_picker_model():
|
||||
model="gpt-5.6-terra", messages=[{"role": "user", "content": "hi"}]
|
||||
)
|
||||
assert seen["model"] == "gpt-5.6-terra"
|
||||
|
||||
|
||||
def test_list_models_reads_enabled_session_config_options(tmp_path):
|
||||
server = tmp_path / "fake_copilot_acp.py"
|
||||
server.write_text(
|
||||
"""import json
|
||||
import sys
|
||||
|
||||
for line in sys.stdin:
|
||||
request = json.loads(line)
|
||||
method = request.get("method")
|
||||
if method == "initialize":
|
||||
result = {"protocolVersion": 1}
|
||||
elif method == "session/new":
|
||||
result = {
|
||||
"sessionId": "catalog-session",
|
||||
"configOptions": [{
|
||||
"id": "model",
|
||||
"category": "model",
|
||||
"options": [
|
||||
{"value": "auto"},
|
||||
{"value": "gpt-5.6-terra"},
|
||||
{"value": "gpt-5.6-terra"},
|
||||
{"value": "claude-fable-5", "_meta": {"copilotEnablement": "disabled"}},
|
||||
],
|
||||
}],
|
||||
"models": {"availableModels": [{"modelId": "stale-legacy-model"}]},
|
||||
}
|
||||
else:
|
||||
result = {}
|
||||
print(json.dumps({"jsonrpc": "2.0", "id": request["id"], "result": result}), flush=True)
|
||||
""",
|
||||
encoding="utf-8",
|
||||
)
|
||||
client = CopilotACPClient(
|
||||
command=sys.executable,
|
||||
args=[str(server)],
|
||||
acp_cwd=str(tmp_path),
|
||||
)
|
||||
|
||||
assert client.list_models(timeout_seconds=30) == ["auto", "gpt-5.6-terra"]
|
||||
assert client.is_closed is True
|
||||
|
||||
|
||||
def test_model_discovery_does_not_allow_file_requests(tmp_path):
|
||||
target = tmp_path / "should-not-be-read.txt"
|
||||
target.write_text("private", encoding="utf-8")
|
||||
server = tmp_path / "fake_copilot_acp_fs_request.py"
|
||||
server.write_text(
|
||||
f"""import json
|
||||
import sys
|
||||
|
||||
initialize = json.loads(sys.stdin.readline())
|
||||
print(json.dumps({{"jsonrpc": "2.0", "id": initialize["id"], "result": {{"protocolVersion": 1}}}}), flush=True)
|
||||
session = json.loads(sys.stdin.readline())
|
||||
print(json.dumps({{"jsonrpc": "2.0", "id": 99, "method": "fs/read_text_file", "params": {{"path": {str(target)!r}}}}}), flush=True)
|
||||
file_response = json.loads(sys.stdin.readline())
|
||||
assert file_response["error"]["code"] == -32601
|
||||
print(json.dumps({{"jsonrpc": "2.0", "id": session["id"], "result": {{"sessionId": "catalog-session", "configOptions": [{{"id": "model", "options": [{{"value": "gpt-5.6-sol"}}]}}]}}}}), flush=True)
|
||||
""",
|
||||
encoding="utf-8",
|
||||
)
|
||||
client = CopilotACPClient(
|
||||
command=sys.executable,
|
||||
args=[str(server)],
|
||||
acp_cwd=str(tmp_path),
|
||||
)
|
||||
|
||||
assert client.list_models(timeout_seconds=30) == ["gpt-5.6-sol"]
|
||||
|
||||
@@ -0,0 +1,139 @@
|
||||
"""Tests for per-turn Copilot x-initiator header injection (issue #3040).
|
||||
|
||||
Copilot bills "premium requests" only when a request is marked as
|
||||
user-initiated via the ``x-initiator: user`` header. Hermes previously sent
|
||||
``x-initiator: agent`` on every request (client-level default headers), so
|
||||
user prompts never consumed premium requests and were throttled as agent
|
||||
traffic. The fix marks the FIRST API call of each user turn as "user" and
|
||||
lets tool-loop follow-ups keep the "agent" default.
|
||||
|
||||
Salvaged from PR #4097 (@tjp2021); adapted to the post-refactor layout
|
||||
(conversation_loop.py owns the injection site, the codex transport now
|
||||
accepts extra_headers).
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from run_agent import AIAgent
|
||||
|
||||
|
||||
def _tool_defs(*names):
|
||||
return [
|
||||
{"type": "function", "function": {"name": n, "description": n, "parameters": {}}}
|
||||
for n in names
|
||||
]
|
||||
|
||||
|
||||
class _FakeOpenAI:
|
||||
def __init__(self, **kw):
|
||||
self.api_key = kw.get("api_key", "test")
|
||||
self.base_url = kw.get("base_url", "http://test")
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
|
||||
def _make_agent(monkeypatch, base_url, api_mode="chat_completions"):
|
||||
"""Create an AIAgent pointing at the given base_url."""
|
||||
monkeypatch.setattr("model_tools.get_tool_definitions", lambda **kw: _tool_defs("web_search"))
|
||||
monkeypatch.setattr("model_tools.check_toolset_requirements", lambda: {})
|
||||
monkeypatch.setattr("agent.process_bootstrap.OpenAI", _FakeOpenAI)
|
||||
return AIAgent(
|
||||
api_key="test-key",
|
||||
base_url=base_url,
|
||||
provider="copilot" if "githubcopilot" in base_url else "openrouter",
|
||||
api_mode=api_mode,
|
||||
max_iterations=4,
|
||||
quiet_mode=True,
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
|
||||
|
||||
def _inject(agent, api_kwargs):
|
||||
"""Mirror the injection block in agent/conversation_loop.py."""
|
||||
if getattr(agent, "_is_user_initiated_turn", False) and agent._is_copilot_url():
|
||||
_xh = dict(api_kwargs.get("extra_headers") or {})
|
||||
_xh["x-initiator"] = "user"
|
||||
api_kwargs["extra_headers"] = _xh
|
||||
agent._is_user_initiated_turn = False
|
||||
return api_kwargs
|
||||
|
||||
|
||||
class TestIsCopilotUrl:
|
||||
"""_is_copilot_url() detects GitHub Copilot endpoints."""
|
||||
|
||||
def test_standard_copilot_url(self, monkeypatch):
|
||||
agent = _make_agent(monkeypatch, "https://api.githubcopilot.com")
|
||||
assert agent._is_copilot_url() is True
|
||||
|
||||
|
||||
def test_github_models_url(self, monkeypatch):
|
||||
agent = _make_agent(monkeypatch, "https://models.github.ai/inference")
|
||||
assert agent._is_copilot_url() is True
|
||||
|
||||
def test_openrouter_url(self, monkeypatch):
|
||||
agent = _make_agent(monkeypatch, "https://openrouter.ai/api/v1")
|
||||
assert agent._is_copilot_url() is False
|
||||
|
||||
def test_case_insensitive(self, monkeypatch):
|
||||
agent = _make_agent(monkeypatch, "https://API.GITHUBCOPILOT.COM")
|
||||
assert agent._is_copilot_url() is True
|
||||
|
||||
|
||||
class TestUserInitiatedTurnFlag:
|
||||
"""_is_user_initiated_turn lifecycle."""
|
||||
|
||||
def test_default_is_false(self, monkeypatch):
|
||||
agent = _make_agent(monkeypatch, "https://api.githubcopilot.com")
|
||||
assert agent._is_user_initiated_turn is False
|
||||
|
||||
def test_reset_session_clears_flag(self, monkeypatch):
|
||||
agent = _make_agent(monkeypatch, "https://api.githubcopilot.com")
|
||||
agent._is_user_initiated_turn = True
|
||||
agent.reset_session_state()
|
||||
assert agent._is_user_initiated_turn is False
|
||||
|
||||
|
||||
class TestFlagFlipOnInjection:
|
||||
"""Flag flips immediately on injection so tool-loop calls use 'agent'."""
|
||||
|
||||
def test_first_call_injects_user_initiator(self, monkeypatch):
|
||||
agent = _make_agent(monkeypatch, "https://api.githubcopilot.com")
|
||||
agent._is_user_initiated_turn = True
|
||||
kwargs = _inject(agent, {})
|
||||
assert kwargs["extra_headers"] == {"x-initiator": "user"}
|
||||
assert agent._is_user_initiated_turn is False
|
||||
|
||||
def test_second_call_has_no_injection(self, monkeypatch):
|
||||
agent = _make_agent(monkeypatch, "https://api.githubcopilot.com")
|
||||
agent._is_user_initiated_turn = True
|
||||
kwargs1 = _inject(agent, {})
|
||||
kwargs2 = _inject(agent, {})
|
||||
assert "extra_headers" in kwargs1
|
||||
assert "extra_headers" not in kwargs2
|
||||
|
||||
|
||||
def test_non_copilot_flag_not_flipped(self, monkeypatch):
|
||||
agent = _make_agent(monkeypatch, "https://openrouter.ai/api/v1")
|
||||
agent._is_user_initiated_turn = True
|
||||
kwargs = _inject(agent, {})
|
||||
assert "extra_headers" not in kwargs
|
||||
# Flag unchanged — non-Copilot path doesn't touch it
|
||||
assert agent._is_user_initiated_turn is True
|
||||
|
||||
|
||||
class TestHeaderValues:
|
||||
"""copilot_default_headers(is_agent_turn=...) sets x-initiator correctly."""
|
||||
|
||||
def test_default_is_agent(self):
|
||||
from hermes_cli.models import copilot_default_headers
|
||||
assert copilot_default_headers()["x-initiator"] == "agent"
|
||||
|
||||
def test_user_turn(self):
|
||||
from hermes_cli.models import copilot_default_headers
|
||||
assert copilot_default_headers(is_agent_turn=False)["x-initiator"] == "user"
|
||||
|
||||
def test_agent_turn_explicit(self):
|
||||
from hermes_cli.models import copilot_default_headers
|
||||
assert copilot_default_headers(is_agent_turn=True)["x-initiator"] == "agent"
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user