fix(middleware): chain request rewrites sequentially

Focused salvage of request-middleware composition from PR #73656.

(cherry picked from commit 089f76e821)
This commit is contained in:
cresslank 2026-07-28 21:18:38 +00:00 committed by Teknium
parent 3e6a081d60
commit 1ce6d95f20
No known key found for this signature in database
4 changed files with 245 additions and 65 deletions

View File

@ -62,6 +62,12 @@ return {
Hermes stores those trace entries in later observer hook payloads as
`middleware_trace`.
If multiple plugins register the same request middleware kind, Hermes applies
them in registration order. Each callback receives the effective request or
arguments returned by the previous callback, while `original_request` or
`original_args` remains the pre-middleware snapshot. Callback failures are
fail-open and do not discard rewrites already applied by earlier middleware.
Execution middleware receives a `next_call` callback. Call it to continue the
chain:

View File

@ -10,7 +10,7 @@ from __future__ import annotations
import logging
from copy import deepcopy
from dataclasses import dataclass, field
from typing import Any, Callable, Dict, List
from typing import Any, Callable, Dict, List, Optional
logger = logging.getLogger(__name__)
@ -83,7 +83,8 @@ def apply_llm_request_middleware(
Middleware may return ``{"request": {...}}`` to replace the effective
provider kwargs before Hermes sends them.
"""
if not _has_middleware(LLM_REQUEST_MIDDLEWARE):
callbacks = _get_middleware_callbacks(LLM_REQUEST_MIDDLEWARE)
if not callbacks:
return RequestMiddlewareResult(
payload=request,
original_payload=request,
@ -91,29 +92,13 @@ def apply_llm_request_middleware(
trace=[],
)
original_request = _safe_copy(request)
current_request = _safe_copy(original_request)
trace: List[Dict[str, Any]] = []
for result in _invoke_middleware(
return _run_request_chain(
LLM_REQUEST_MIDDLEWARE,
request=current_request,
original_request=original_request,
callbacks,
request,
payload_key="request",
original_key="original_request",
**context,
):
if not isinstance(result, dict):
continue
next_request = result.get("request")
if not isinstance(next_request, dict):
continue
current_request = _safe_copy(next_request)
trace.append(_trace_entry(result))
return RequestMiddlewareResult(
payload=current_request,
original_payload=original_request,
changed=bool(trace),
trace=trace,
)
@ -145,7 +130,8 @@ def apply_tool_request_middleware(
current_args = _safe_copy(relay_args)
trace.append({"source": "nemo_relay"})
if not _has_middleware(TOOL_REQUEST_MIDDLEWARE):
callbacks = _get_middleware_callbacks(TOOL_REQUEST_MIDDLEWARE)
if not callbacks:
return RequestMiddlewareResult(
payload=args if not trace else current_args,
original_payload=args,
@ -153,26 +139,16 @@ def apply_tool_request_middleware(
trace=trace,
)
for result in _invoke_middleware(
return _run_request_chain(
TOOL_REQUEST_MIDDLEWARE,
tool_name=tool_name,
args=current_args,
original_args=original_args,
**context,
):
if not isinstance(result, dict):
continue
next_args = result.get("args")
if not isinstance(next_args, dict):
continue
current_args = _safe_copy(next_args)
trace.append(_trace_entry(result))
return RequestMiddlewareResult(
payload=current_args,
callbacks,
current_args,
payload_key="args",
original_key="original_args",
original_payload=original_args,
changed=bool(trace),
trace=trace,
initial_trace=trace,
tool_name=tool_name,
**context,
)
@ -233,24 +209,63 @@ def run_api_execution_middleware(
return run_llm_execution_middleware(request, next_call, **context)
def _invoke_middleware(kind: str, **kwargs: Any) -> List[Any]:
from hermes_cli.plugins import invoke_middleware
return invoke_middleware(kind, **middleware_payload(**kwargs))
def _has_middleware(kind: str) -> bool:
from hermes_cli.plugins import has_middleware
return has_middleware(kind)
def _get_middleware_callbacks(kind: str) -> List[Callable]:
from hermes_cli.plugins import get_plugin_manager
return list(get_plugin_manager()._middleware.get(kind, []))
def _run_request_chain(
kind: str,
callbacks: List[Callable],
payload: Dict[str, Any],
*,
payload_key: str,
original_key: str,
original_payload: Optional[Dict[str, Any]] = None,
initial_trace: Optional[List[Dict[str, Any]]] = None,
**context: Any,
) -> RequestMiddlewareResult:
"""Apply request middleware serially to the current effective payload."""
if original_payload is None:
original_payload = _safe_copy(payload)
current_payload = _safe_copy(payload)
trace = list(initial_trace or [])
for callback in callbacks:
callback_payload = middleware_payload(
**context,
**{
payload_key: _safe_copy(current_payload),
original_key: _safe_copy(original_payload),
},
)
try:
result = callback(**callback_payload)
except Exception as exc:
logger.warning(
"Middleware '%s' callback %s raised: %s",
kind,
getattr(callback, "__name__", repr(callback)),
exc,
)
continue
if not isinstance(result, dict):
continue
next_payload = result.get(payload_key)
if not isinstance(next_payload, dict):
continue
current_payload = _safe_copy(next_payload)
trace.append(_trace_entry(result))
return RequestMiddlewareResult(
payload=current_payload,
original_payload=original_payload,
changed=bool(trace),
trace=trace,
)
def _run_execution_chain(
kind: str,
callbacks: List[Callable],

View File

@ -1,7 +1,9 @@
"""Tests for the Hermes plugin system (hermes_cli.plugins)."""
import logging
import json
import logging
import os
import subprocess
import sys
import types
from pathlib import Path
@ -295,6 +297,161 @@ class TestPluginDiscovery:
assert run_tool_execution_middleware("terminal", args, lambda payload: payload) is args
assert has_middleware("tool_request") is False
def test_llm_request_middleware_chains_effective_request_with_original_isolated(
self, monkeypatch
):
seen = []
def first(**kwargs):
kwargs["original_request"]["tampered"] = True
return {
"request": {**kwargs["request"], "first": True},
"source": "first",
}
def second(**kwargs):
seen.append((kwargs["request"], kwargs["original_request"]))
return {
"request": {**kwargs["request"], "second": True},
"source": "second",
}
manager = PluginManager()
manager._middleware = {"llm_request": [first, second]}
monkeypatch.setattr("hermes_cli.plugins.get_plugin_manager", lambda: manager)
request = {"messages": []}
result = apply_llm_request_middleware(request)
assert seen == [({"messages": [], "first": True}, request)]
assert result.payload == {"messages": [], "first": True, "second": True}
assert result.original_payload == request
assert result.trace == [{"source": "first"}, {"source": "second"}]
assert request == {"messages": []}
def test_request_middleware_failure_keeps_prior_rewrite_isolated(
self, monkeypatch, caplog
):
seen = []
def first(**kwargs):
return {
"request": {**kwargs["request"], "first": True},
"source": "first",
}
def failing(**kwargs):
kwargs["request"]["failed_mutation"] = True
raise RuntimeError("broken plugin")
def malformed(**kwargs):
seen.append(kwargs["request"])
return {"request": "not-a-dict", "source": "malformed"}
manager = PluginManager()
manager._middleware = {"llm_request": [first, failing, malformed]}
monkeypatch.setattr("hermes_cli.plugins.get_plugin_manager", lambda: manager)
with caplog.at_level(logging.WARNING):
result = apply_llm_request_middleware({"messages": []})
assert seen == [{"messages": [], "first": True}]
assert result.payload == {"messages": [], "first": True}
assert result.trace == [{"source": "first"}]
assert "Middleware 'llm_request' callback failing raised: broken plugin" in caplog.text
def test_tool_request_middleware_runs_relay_before_sequential_plugins(
self, monkeypatch
):
seen = []
monkeypatch.setattr(
"agent.relay_runtime.apply_tool_request_intercepts",
lambda **kwargs: {**kwargs["args"], "relay": True},
)
def first(**kwargs):
seen.append(("first", kwargs["args"], kwargs["original_args"]))
return {"args": {**kwargs["args"], "first": True}, "source": "first"}
def second(**kwargs):
seen.append(("second", kwargs["args"], kwargs["original_args"]))
return {"args": {**kwargs["args"], "second": True}, "source": "second"}
manager = PluginManager()
manager._middleware = {"tool_request": [first, second]}
monkeypatch.setattr("hermes_cli.plugins.get_plugin_manager", lambda: manager)
args = {"path": "README.md"}
result = apply_tool_request_middleware("read_file", args, session_id="s1")
assert seen == [
("first", {"path": "README.md", "relay": True}, args),
("second", {"path": "README.md", "relay": True, "first": True}, args),
]
assert result.payload == {
"path": "README.md",
"relay": True,
"first": True,
"second": True,
}
assert result.original_payload == args
assert result.trace == [
{"source": "nemo_relay"},
{"source": "first"},
{"source": "second"},
]
def test_request_middleware_composes_discovered_plugins_in_fresh_process(
self, tmp_path
):
home = tmp_path / "home"
plugins = home / "plugins"
_make_plugin_dir(
plugins,
"first",
register_body=(
"ctx.register_middleware('llm_request', "
"lambda **kw: {'request': {**kw['request'], 'first': True}})"
),
auto_enable=False,
)
_make_plugin_dir(
plugins,
"second",
register_body=(
"ctx.register_middleware('llm_request', "
"lambda **kw: {'request': {**kw['request'], "
"'saw_first': kw['request'].get('first')}})"
),
auto_enable=False,
)
(home / "config.yaml").write_text(
yaml.safe_dump({"plugins": {"enabled": ["first", "second"]}})
)
script = """
import json
from hermes_cli.middleware import apply_llm_request_middleware
from hermes_cli.plugins import get_plugin_manager
get_plugin_manager().discover_and_load()
result = apply_llm_request_middleware({"base": True})
print(json.dumps({"payload": result.payload, "original": result.original_payload}))
"""
env = dict(os.environ, HERMES_HOME=str(home))
completed = subprocess.run(
[sys.executable, "-c", script],
cwd=Path(__file__).resolve().parents[2],
env=env,
text=True,
capture_output=True,
check=True,
)
assert json.loads(completed.stdout) == {
"payload": {"base": True, "first": True, "saw_first": True},
"original": {"base": True},
}

View File

@ -100,14 +100,12 @@ class TestHandleFunctionCall:
def test_tool_request_and_execution_middleware_wrap_registry_dispatch(self, monkeypatch):
seen = {}
def fake_invoke_middleware(kind, **kwargs):
if kind == "tool_request":
return [{
"args": {**kwargs["args"], "rewritten": True},
"source": "test-middleware",
"reason": "rewrite",
}]
return []
def request_middleware(**kwargs):
return {
"args": {**kwargs["args"], "rewritten": True},
"source": "test-middleware",
"reason": "rewrite",
}
def execution_middleware(**kwargs):
seen["execution_args"] = kwargs["args"]
@ -120,9 +118,13 @@ class TestHandleFunctionCall:
manager = type(
"Manager",
(),
{"_middleware": {"tool_request": [fake_invoke_middleware], "tool_execution": [execution_middleware]}},
{
"_middleware": {
"tool_request": [request_middleware],
"tool_execution": [execution_middleware],
}
},
)()
monkeypatch.setattr("hermes_cli.plugins.invoke_middleware", fake_invoke_middleware)
monkeypatch.setattr("hermes_cli.plugins.get_plugin_manager", lambda: manager)
hook_calls = []
monkeypatch.setattr(