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:
parent
3e6a081d60
commit
1ce6d95f20
|
|
@ -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:
|
||||
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Reference in New Issue