feat(middleware): add adaptive execution intercepts
Signed-off-by: Bryan Bednarski <bbednarski@nvidia.com>
This commit is contained in:
+86
-2
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user