feat(middleware): add adaptive execution intercepts

Signed-off-by: Bryan Bednarski <bbednarski@nvidia.com>
This commit is contained in:
Bryan Bednarski
2026-06-03 11:22:06 -07:00
parent b4b9a93848
commit 2e0c9083db
14 changed files with 2013 additions and 149 deletions
+86 -2
View File
@@ -49,7 +49,7 @@ from typing import Any, Callable, Dict, List, Optional, Set, Union
from hermes_constants import get_hermes_home
from utils import env_var_enabled
from hermes_cli.config import cfg_get
OBSERVER_SCHEMA_VERSION = "hermes.observer.v1"
from hermes_cli.middleware import OBSERVER_SCHEMA_VERSION, VALID_MIDDLEWARE
def get_bundled_plugins_dir() -> Path:
@@ -277,6 +277,7 @@ class LoadedPlugin:
module: Optional[types.ModuleType] = None
tools_registered: List[str] = field(default_factory=list)
hooks_registered: List[str] = field(default_factory=list)
middleware_registered: List[str] = field(default_factory=list)
commands_registered: List[str] = field(default_factory=list)
enabled: bool = False
error: Optional[str] = None
@@ -952,6 +953,27 @@ class PluginContext:
self._manager._hooks.setdefault(hook_name, []).append(callback)
logger.debug("Plugin %s registered hook: %s", self.manifest.name, hook_name)
# -- middleware registration -------------------------------------------
def register_middleware(self, kind: str, callback: Callable) -> None:
"""Register a behavior-changing middleware callback.
Middleware is separate from observer hooks: request middleware may
rewrite the effective payload, and execution middleware may wrap the
real callback. Unknown kinds are stored for forward compatibility but
warned so plugin authors can catch typos.
"""
if kind not in VALID_MIDDLEWARE:
logger.warning(
"Plugin '%s' registered unknown middleware '%s' "
"(valid: %s)",
self.manifest.name,
kind,
", ".join(sorted(VALID_MIDDLEWARE)),
)
self._manager._middleware.setdefault(kind, []).append(callback)
logger.debug("Plugin %s registered middleware: %s", self.manifest.name, kind)
# -- skill registration -------------------------------------------------
def register_skill(
@@ -1010,6 +1032,7 @@ class PluginManager:
def __init__(self) -> None:
self._plugins: Dict[str, LoadedPlugin] = {}
self._hooks: Dict[str, List[Callable]] = {}
self._middleware: Dict[str, List[Callable]] = {}
self._plugin_tool_names: Set[str] = set()
self._plugin_platform_names: Set[str] = set()
self._cli_commands: Dict[str, dict] = {}
@@ -1039,6 +1062,7 @@ class PluginManager:
if force:
self._plugins.clear()
self._hooks.clear()
self._middleware.clear()
self._plugin_tool_names.clear()
self._cli_commands.clear()
self._plugin_commands.clear()
@@ -1449,15 +1473,28 @@ class PluginManager:
for h in p.hooks_registered
}
)
loaded.middleware_registered = list(
{
kind
for kind, cbs in self._middleware.items()
if cbs
}
- {
kind
for name, p in self._plugins.items()
for kind in p.middleware_registered
}
)
loaded.commands_registered = [
c for c in self._plugin_commands
if self._plugin_commands[c].get("plugin") == manifest.name
]
loaded.enabled = True
logger.debug(
" registered: %d tool(s), %d hook(s), %d slash command(s), %d CLI command(s)",
" registered: %d tool(s), %d hook(s), %d middleware, %d slash command(s), %d CLI command(s)",
len(loaded.tools_registered),
len(loaded.hooks_registered),
len(loaded.middleware_registered),
len(loaded.commands_registered),
sum(
1 for c in self._cli_commands
@@ -1575,6 +1612,33 @@ class PluginManager:
"""Return True when at least one callback is registered for a hook."""
return bool(self._hooks.get(hook_name))
def has_middleware(self, kind: str) -> bool:
"""Return True when at least one callback is registered for middleware."""
return bool(self._middleware.get(kind))
def invoke_middleware(self, kind: str, **kwargs: Any) -> List[Any]:
"""Call registered middleware callbacks for *kind*.
Each callback is isolated so one plugin cannot break the base runtime
path. Middleware that wants to change behavior must return the shape
documented by the caller-specific contract.
"""
callbacks = self._middleware.get(kind, [])
results: List[Any] = []
for cb in callbacks:
try:
ret = cb(**kwargs)
if ret is not None:
results.append(ret)
except Exception as exc:
logger.warning(
"Middleware '%s' callback %s raised: %s",
kind,
getattr(cb, "__name__", repr(cb)),
exc,
)
return results
# -----------------------------------------------------------------------
# Introspection
# -----------------------------------------------------------------------
@@ -1594,6 +1658,7 @@ class PluginManager:
"enabled": loaded.enabled,
"tools": len(loaded.tools_registered),
"hooks": len(loaded.hooks_registered),
"middleware": len(loaded.middleware_registered),
"commands": len(loaded.commands_registered),
"error": loaded.error,
}
@@ -1655,6 +1720,23 @@ def invoke_hook(hook_name: str, **kwargs: Any) -> List[Any]:
return get_plugin_manager().invoke_hook(hook_name, **kwargs)
def invoke_middleware(kind: str, **kwargs: Any) -> List[Any]:
"""Invoke registered middleware callbacks.
Returns a list of non-``None`` return values from middleware callbacks.
"""
return get_plugin_manager().invoke_middleware(kind, **kwargs)
def has_middleware(kind: str) -> bool:
"""Return True when middleware callbacks are registered for ``kind``."""
manager = get_plugin_manager()
method = getattr(manager, "has_middleware", None)
if callable(method):
return bool(method(kind))
return bool(getattr(manager, "_middleware", {}).get(kind))
def has_hook(hook_name: str) -> bool:
"""Return True when a hook has registered callbacks."""
return get_plugin_manager().has_hook(hook_name)
@@ -1683,6 +1765,7 @@ def get_pre_tool_call_block_message(
tool_call_id: str = "",
turn_id: str = "",
api_request_id: str = "",
middleware_trace: Optional[List[Dict[str, Any]]] = None,
) -> Optional[str]:
"""Check ``pre_tool_call`` hooks for a blocking directive.
@@ -1709,6 +1792,7 @@ def get_pre_tool_call_block_message(
tool_call_id=tool_call_id,
turn_id=turn_id,
api_request_id=api_request_id,
middleware_trace=list(middleware_trace or []),
)
for result in hook_results: