Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
84 changes: 79 additions & 5 deletions gateway/platforms/api_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -501,7 +501,10 @@ def __len__(self) -> int:

_CORS_HEADERS = {
"Access-Control-Allow-Methods": "GET, POST, DELETE, OPTIONS",
"Access-Control-Allow-Headers": "Authorization, Content-Type, Idempotency-Key",
"Access-Control-Allow-Headers": (
"Authorization, Content-Type, Idempotency-Key, "
"X-Portal-Client-Id, X-Hermes-Session-Key, X-Hermes-Thread-Id"
),
}


Expand Down Expand Up @@ -718,6 +721,64 @@ def __init__(self, config: PlatformConfig):
# in-flight run by run_id.
self._run_approval_sessions: Dict[str, str] = {}
self._session_db: Optional[Any] = None # Lazy-init SessionDB for session continuity
self._plugin_manager = self._load_plugin_manager()

def _load_plugin_manager(self):
"""Best-effort plugin manager discovery for API-server plugin hooks."""
try:
from hermes_cli.plugins import discover_plugins, get_plugin_manager

discover_plugins()
return get_plugin_manager()
except Exception as exc:
logger.debug("[%s] Plugin manager unavailable for API server hooks: %s", self.name, exc)
return None

def _mount_plugin_api_routes(self, router: Any) -> None:
"""Mount plugin-contributed API routes without breaking core startup."""
manager = getattr(self, "_plugin_manager", None)
if manager is None or not hasattr(manager, "get_api_server_routes"):
return

existing_routes: set[tuple[str, str]] = set()
for resource in router.resources():
path = getattr(resource, "canonical", None)
if not path:
continue
for route_info in resource:
method = getattr(route_info, "method", None)
if method:
existing_routes.add((str(method).upper(), path))

for route in manager.get_api_server_routes() or []:
try:
method = str(route.get("method", "")).upper()
path = route.get("path")
handler = route.get("handler")
if not method or not path or handler is None:
raise ValueError("missing method/path/handler")
if not isinstance(path, str) or not path.startswith("/"):
raise ValueError("path must start with /")
if (method, path) in existing_routes:
raise ValueError(f"route already registered: {method} {path}")
add_fn = getattr(router, f"add_{method.lower()}", None)
if add_fn is None:
raise ValueError(f"unsupported method: {method}")
route_name = route.get("name")
if route_name:
add_fn(path, handler, name=route_name)
else:
add_fn(path, handler)
existing_routes.add((method, path))
except Exception as exc:
logger.warning(
"[%s] Failed to register plugin route %s %s (plugin=%s): %s",
self.name,
route.get("method"),
route.get("path"),
route.get("plugin"),
exc,
)

@staticmethod
def _parse_cors_origins(value: Any) -> tuple[str, ...]:
Expand Down Expand Up @@ -1077,7 +1138,7 @@ async def _handle_capabilities(self, request: "web.Request") -> "web.Response":
if auth_err:
return auth_err

return web.json_response({
payload = {
"object": "hermes.api_server.capabilities",
"platform": "hermes-agent",
"model": self._model_name,
Expand Down Expand Up @@ -1144,7 +1205,20 @@ async def _handle_capabilities(self, request: "web.Request") -> "web.Response":
"session_chat": {"method": "POST", "path": "/api/sessions/{session_id}/chat"},
"session_chat_stream": {"method": "POST", "path": "/api/sessions/{session_id}/chat/stream"},
},
})
}

manager = self._plugin_manager
if manager is not None and hasattr(manager, "get_api_server_capabilities"):
plugin_caps = manager.get_api_server_capabilities(adapter=self, request=request) or []
if plugin_caps:
payload.setdefault("extensions", {})
payload["extensions"]["plugins"] = {
str(entry.get("plugin")): entry.get("capabilities", {})
for entry in plugin_caps
if isinstance(entry, dict) and entry.get("plugin")
}

return web.json_response(payload)

async def _handle_skills(self, request: "web.Request") -> "web.Response":
"""GET /v1/skills β€” list installed skills visible to the API-server agent.
Expand Down Expand Up @@ -4069,7 +4143,7 @@ async def _sweep_orphaned_runs(self) -> None:
stale_statuses = [
run_id
for run_id, status in list(self._run_statuses.items())
if status.get("status") in {"completed", "failed", "cancelled"}
if status.get("status") in {"completed", "failed", "cancelled", "routed"}
and now - float(status.get("updated_at", 0) or 0) > self._RUN_STATUS_TTL
]
for run_id in stale_statuses:
Expand Down Expand Up @@ -4130,7 +4204,7 @@ async def connect(self) -> bool:
# native routes first lets those shims no-op instead of shadowing the
# upstream session-control handlers.
self._app["api_server_adapter"] = self

self._mount_plugin_api_routes(self._app.router)
# Start background sweep to clean up orphaned (unconsumed) run streams
sweep_task = asyncio.create_task(self._sweep_orphaned_runs())
try:
Expand Down
73 changes: 73 additions & 0 deletions hermes_cli/plugins.py
Original file line number Diff line number Diff line change
Expand Up @@ -952,6 +952,43 @@ def register_hook(self, hook_name: str, callback: Callable) -> None:
self._manager._hooks.setdefault(hook_name, []).append(callback)
logger.debug("Plugin %s registered hook: %s", self.manifest.name, hook_name)

# -- API server registration --------------------------------------------

def register_api_server_route(
self,
method: str,
path: str,
handler: Callable,
*,
name: str | None = None,
) -> None:
"""Register an aiohttp route contribution for the API server adapter."""
self._manager._api_server_routes.append(
{
"method": str(method or "").upper(),
"path": path,
"handler": handler,
"name": name,
"plugin": self.manifest.name,
}
)
logger.debug(
"Plugin %s registered API server route: %s %s",
self.manifest.name,
str(method or "").upper(),
path,
)

def register_api_server_capability(self, provider: Callable) -> None:
"""Register a provider callback for /v1/capabilities extensions."""
self._manager._api_server_capability_providers.append(
{
"plugin": self.manifest.name,
"provider": provider,
}
)
logger.debug("Plugin %s registered API server capability provider", self.manifest.name)

# -- skill registration -------------------------------------------------

def register_skill(
Expand Down Expand Up @@ -1022,6 +1059,8 @@ def __init__(self) -> None:
# Plugin-registered auxiliary tasks: key β†’ {key, display_name,
# description, defaults, plugin}. See PluginContext.register_auxiliary_task.
self._aux_tasks: Dict[str, Dict[str, Any]] = {}
self._api_server_routes: List[Dict[str, Any]] = []
self._api_server_capability_providers: List[Dict[str, Any]] = []

# -----------------------------------------------------------------------
# Public
Expand All @@ -1044,6 +1083,8 @@ def discover_and_load(self, force: bool = False) -> None:
self._plugin_commands.clear()
self._plugin_skills.clear()
self._aux_tasks.clear()
self._api_server_routes.clear()
self._api_server_capability_providers.clear()
self._context_engine = None
self._discovered = True

Expand Down Expand Up @@ -1600,6 +1641,38 @@ def list_plugins(self) -> List[Dict[str, Any]]:
)
return result

def get_api_server_routes(self) -> List[Dict[str, Any]]:
"""Return plugin-contributed API server routes in registration order."""
return list(self._api_server_routes)

def get_api_server_capabilities(self, *, adapter: Any = None, request: Any = None) -> List[Dict[str, Any]]:
"""Resolve plugin capability contributions for /v1/capabilities."""
results: List[Dict[str, Any]] = []
for entry in self._api_server_capability_providers:
plugin_name = entry.get("plugin", "unknown")
provider = entry.get("provider")
if not callable(provider):
continue
try:
payload = provider(adapter=adapter, request=request)
except Exception as exc:
logger.warning(
"Plugin '%s' API server capability provider failed: %s",
plugin_name,
exc,
)
continue
if payload is None:
continue
if not isinstance(payload, dict):
logger.warning(
"Plugin '%s' API server capability provider returned non-dict payload; skipping",
plugin_name,
)
continue
results.append({"plugin": plugin_name, "capabilities": payload})
return results

# -----------------------------------------------------------------------
# Plugin skill lookups
# -----------------------------------------------------------------------
Expand Down
114 changes: 114 additions & 0 deletions tests/gateway/test_api_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -676,6 +676,102 @@ async def test_capabilities_requires_auth_when_key_configured(self, auth_adapter
data = await authed.json()
assert data["auth"]["required"] is True

@pytest.mark.asyncio
async def test_capabilities_include_plugin_extension_payload(self, adapter):
class _PluginManager:
def get_api_server_capabilities(self, *, adapter=None, request=None):
return [
{"plugin": "alpha", "capabilities": {"alpha_flag": True}},
{"plugin": "beta", "capabilities": {"beta_count": 2}},
]

adapter._plugin_manager = _PluginManager()
app = _create_app(adapter)
async with TestClient(TestServer(app)) as cli:
resp = await cli.get("/v1/capabilities")
assert resp.status == 200
data = await resp.json()
assert data["extensions"]["plugins"] == {
"alpha": {"alpha_flag": True},
"beta": {"beta_count": 2},
}


class TestPluginApiServerRoutes:
def test_mount_plugin_routes_after_core_routes(self, adapter):
async def _plugin_handler(request):
return web.json_response({"ok": True})

class _PluginManager:
def get_api_server_routes(self):
return [
{
"method": "GET",
"path": "/v1/plugins/test-ping",
"handler": _plugin_handler,
"name": "test_ping",
"plugin": "plugin-test",
}
]

adapter._plugin_manager = _PluginManager()
app = _create_app(adapter)
adapter._mount_plugin_api_routes(app.router)
paths = [res.canonical for res in app.router.resources()]
assert "/v1/models" in paths
assert "/v1/plugins/test-ping" in paths
assert paths.index("/v1/plugins/test-ping") > paths.index("/v1/models")

def test_plugin_route_registration_failure_is_isolated(self, adapter, caplog):
async def _plugin_handler(request):
return web.json_response({"ok": True})

class _PluginManager:
def get_api_server_routes(self):
return [
{
"method": "NOPE",
"path": "/v1/plugins/bad",
"handler": _plugin_handler,
"name": "bad",
"plugin": "plugin-test",
}
]

adapter._plugin_manager = _PluginManager()
app = _create_app(adapter)
with caplog.at_level("WARNING"):
adapter._mount_plugin_api_routes(app.router)
paths = [res.canonical for res in app.router.resources()]
assert "/v1/models" in paths
assert "/v1/plugins/bad" not in paths
assert "plugin route" in caplog.text.lower()

def test_plugin_route_cannot_shadow_core_route(self, adapter, caplog):
async def _plugin_handler(request):
return web.json_response({"shadow": True})

class _PluginManager:
def get_api_server_routes(self):
return [
{
"method": "GET",
"path": "/v1/models",
"handler": _plugin_handler,
"name": "models_shadow",
"plugin": "plugin-test",
}
]

adapter._plugin_manager = _PluginManager()
app = _create_app(adapter)
with caplog.at_level("WARNING"):
adapter._mount_plugin_api_routes(app.router)

matching = [res for res in app.router.resources() if res.canonical == "/v1/models"]
assert len(matching) == 1
assert "route already registered" in caplog.text.lower()


# ---------------------------------------------------------------------------
# /v1/skills and /v1/toolsets endpoints
Expand Down Expand Up @@ -3007,6 +3103,24 @@ async def test_cors_allows_idempotency_key_header(self):
assert resp.status == 200
assert "Idempotency-Key" in resp.headers.get("Access-Control-Allow-Headers", "")

@pytest.mark.asyncio
async def test_cors_allows_mobile_routing_headers(self):
adapter = _make_adapter(cors_origins=["http://localhost:3000"])
app = _create_app(adapter)
async with TestClient(TestServer(app)) as cli:
resp = await cli.options(
"/v1/runs",
headers={
"Origin": "http://localhost:3000",
"Access-Control-Request-Method": "POST",
"Access-Control-Request-Headers": "X-Portal-Client-Id, X-Hermes-Thread-Id",
},
)
assert resp.status == 200
allow_headers = resp.headers.get("Access-Control-Allow-Headers", "")
assert "X-Portal-Client-Id" in allow_headers
assert "X-Hermes-Thread-Id" in allow_headers

@pytest.mark.asyncio
async def test_cors_sets_vary_origin_header(self):
adapter = _make_adapter(cors_origins=["http://localhost:3000"])
Expand Down
1 change: 0 additions & 1 deletion tests/gateway/test_api_server_runs.py
Original file line number Diff line number Diff line change
Expand Up @@ -189,7 +189,6 @@ async def test_start_with_valid_auth(self, auth_adapter):
)
assert resp.status == 202


# ---------------------------------------------------------------------------
# GET /v1/runs/{run_id} β€” poll run status
# ---------------------------------------------------------------------------
Expand Down
Loading