spamming the house with try catches
This commit is contained in:
@@ -4,7 +4,7 @@ import json
|
|||||||
import os
|
import os
|
||||||
from dataclasses import asdict
|
from dataclasses import asdict
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Literal, cast
|
from typing import Any, Literal
|
||||||
|
|
||||||
from mcp.server.fastmcp import FastMCP
|
from mcp.server.fastmcp import FastMCP
|
||||||
|
|
||||||
@@ -57,7 +57,15 @@ def create_broker_server(service: WfMcpService) -> FastMCP:
|
|||||||
|
|
||||||
@server.tool()
|
@server.tool()
|
||||||
async def refresh_connection_catalog(connection_id: str) -> dict[str, Any]:
|
async def refresh_connection_catalog(connection_id: str) -> dict[str, Any]:
|
||||||
await service.refresh_connection_catalog(connection_id)
|
try:
|
||||||
|
await service.refresh_connection_catalog(connection_id)
|
||||||
|
except Exception as exc:
|
||||||
|
return {
|
||||||
|
"connection_id": connection_id,
|
||||||
|
"refreshed": False,
|
||||||
|
"error_type": type(exc).__name__,
|
||||||
|
"error": str(exc),
|
||||||
|
}
|
||||||
snapshot = service.get_connection_snapshot(connection_id)
|
snapshot = service.get_connection_snapshot(connection_id)
|
||||||
if snapshot is None:
|
if snapshot is None:
|
||||||
return {"connection_id": connection_id, "refreshed": False}
|
return {"connection_id": connection_id, "refreshed": False}
|
||||||
|
|||||||
+43
-6
@@ -52,6 +52,40 @@ def _json_dump(data: Any) -> None:
|
|||||||
print(json.dumps(data, indent=2))
|
print(json.dumps(data, indent=2))
|
||||||
|
|
||||||
|
|
||||||
|
async def _refresh_all(service, connection_id: str | None) -> list[dict[str, Any]]:
|
||||||
|
target_ids = (
|
||||||
|
[connection_id]
|
||||||
|
if connection_id is not None
|
||||||
|
else [connection.id for connection in service.connections.list_enabled()]
|
||||||
|
)
|
||||||
|
results: list[dict[str, Any]] = []
|
||||||
|
for target_id in target_ids:
|
||||||
|
try:
|
||||||
|
await service.refresh_connection_catalog(target_id)
|
||||||
|
snapshot = service.get_connection_snapshot(target_id)
|
||||||
|
results.append(
|
||||||
|
{
|
||||||
|
"connection_id": target_id,
|
||||||
|
"refreshed": snapshot is not None,
|
||||||
|
"node_count": 0 if snapshot is None else len(snapshot.nodes),
|
||||||
|
"resource_count": 0
|
||||||
|
if snapshot is None
|
||||||
|
else len(snapshot.resources),
|
||||||
|
"prompt_count": 0 if snapshot is None else len(snapshot.prompts),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
results.append(
|
||||||
|
{
|
||||||
|
"connection_id": target_id,
|
||||||
|
"refreshed": False,
|
||||||
|
"error_type": type(exc).__name__,
|
||||||
|
"error": str(exc),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return results
|
||||||
|
|
||||||
|
|
||||||
def main(argv: list[str] | None = None) -> int:
|
def main(argv: list[str] | None = None) -> int:
|
||||||
parser = build_parser()
|
parser = build_parser()
|
||||||
args = parser.parse_args(argv)
|
args = parser.parse_args(argv)
|
||||||
@@ -82,12 +116,15 @@ def main(argv: list[str] | None = None) -> int:
|
|||||||
return 0
|
return 0
|
||||||
|
|
||||||
if args.command == "refresh":
|
if args.command == "refresh":
|
||||||
if args.connection_id:
|
results = asyncio.run(_refresh_all(service, args.connection_id))
|
||||||
asyncio.run(service.refresh_connection_catalog(args.connection_id))
|
_json_dump(
|
||||||
else:
|
{
|
||||||
for connection in service.connections.list_enabled():
|
"results": results,
|
||||||
asyncio.run(service.refresh_connection_catalog(connection.id))
|
"catalog": service.get_catalog().as_payload(),
|
||||||
_json_dump(service.get_catalog().as_payload())
|
}
|
||||||
|
)
|
||||||
|
if any(not result["refreshed"] for result in results):
|
||||||
|
return 1
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
parser.error(f"unknown command {args.command!r}")
|
parser.error(f"unknown command {args.command!r}")
|
||||||
|
|||||||
+50
-37
@@ -292,44 +292,57 @@ class WfMcpService:
|
|||||||
payload={"server": connection.server},
|
payload={"server": connection.server},
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
capabilities = await discover_connection_capabilities(
|
try:
|
||||||
connection=connection,
|
capabilities = await discover_connection_capabilities(
|
||||||
auth=auth,
|
connection=connection,
|
||||||
adapter=adapter,
|
auth=auth,
|
||||||
)
|
adapter=adapter,
|
||||||
specs = specs_from_discovered_tools(
|
|
||||||
connection=connection,
|
|
||||||
auth=auth,
|
|
||||||
adapter=adapter,
|
|
||||||
tools=capabilities.tools,
|
|
||||||
emit_event=self._record_event,
|
|
||||||
)
|
|
||||||
self.register_specs(
|
|
||||||
connection_id,
|
|
||||||
*specs,
|
|
||||||
max_age_seconds=max_age_seconds,
|
|
||||||
)
|
|
||||||
snapshot = snapshot_from_specs(
|
|
||||||
connection_id,
|
|
||||||
specs=self.specs_by_connection.get(connection_id, {}),
|
|
||||||
resources=capabilities.resources,
|
|
||||||
prompts=capabilities.prompts,
|
|
||||||
metadata=capabilities.metadata,
|
|
||||||
fetched_at_epoch_ms=int(time.time() * 1000),
|
|
||||||
max_age_seconds=max_age_seconds or self.default_catalog_max_age_seconds,
|
|
||||||
)
|
|
||||||
self.store.save_catalog(snapshot)
|
|
||||||
self._record_event(
|
|
||||||
make_event(
|
|
||||||
"catalog_refresh_completed",
|
|
||||||
connection_id=connection_id,
|
|
||||||
payload={
|
|
||||||
"node_count": len(snapshot.nodes),
|
|
||||||
"resource_count": len(snapshot.resources),
|
|
||||||
"prompt_count": len(snapshot.prompts),
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
)
|
specs = specs_from_discovered_tools(
|
||||||
|
connection=connection,
|
||||||
|
auth=auth,
|
||||||
|
adapter=adapter,
|
||||||
|
tools=capabilities.tools,
|
||||||
|
emit_event=self._record_event,
|
||||||
|
)
|
||||||
|
self.register_specs(
|
||||||
|
connection_id,
|
||||||
|
*specs,
|
||||||
|
max_age_seconds=max_age_seconds,
|
||||||
|
)
|
||||||
|
snapshot = snapshot_from_specs(
|
||||||
|
connection_id,
|
||||||
|
specs=self.specs_by_connection.get(connection_id, {}),
|
||||||
|
resources=capabilities.resources,
|
||||||
|
prompts=capabilities.prompts,
|
||||||
|
metadata=capabilities.metadata,
|
||||||
|
fetched_at_epoch_ms=int(time.time() * 1000),
|
||||||
|
max_age_seconds=max_age_seconds or self.default_catalog_max_age_seconds,
|
||||||
|
)
|
||||||
|
self.store.save_catalog(snapshot)
|
||||||
|
self._record_event(
|
||||||
|
make_event(
|
||||||
|
"catalog_refresh_completed",
|
||||||
|
connection_id=connection_id,
|
||||||
|
payload={
|
||||||
|
"node_count": len(snapshot.nodes),
|
||||||
|
"resource_count": len(snapshot.resources),
|
||||||
|
"prompt_count": len(snapshot.prompts),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
self._record_event(
|
||||||
|
make_event(
|
||||||
|
"catalog_refresh_failed",
|
||||||
|
connection_id=connection_id,
|
||||||
|
payload={
|
||||||
|
"error_type": type(exc).__name__,
|
||||||
|
"error": str(exc),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
|
||||||
def compile_plan(self, plan: RawWorkflowPlan) -> Workflow:
|
def compile_plan(self, plan: RawWorkflowPlan) -> Workflow:
|
||||||
node_defs: dict[str, Any] = {}
|
node_defs: dict[str, Any] = {}
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ from wf_mcp import (
|
|||||||
load_broker_config,
|
load_broker_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
from test_wf_mcp_support import FakeAdapter, local_temp_root
|
from test_wf_mcp_support import FailingDiscoveryAdapter, FakeAdapter, local_temp_root
|
||||||
|
|
||||||
|
|
||||||
def test_load_broker_config_resolves_relative_store_root() -> None:
|
def test_load_broker_config_resolves_relative_store_root() -> None:
|
||||||
@@ -80,3 +80,23 @@ def test_build_service_from_config_registers_connections() -> None:
|
|||||||
|
|
||||||
ids = [connection.id for connection in service.connections.list_all()]
|
ids = [connection.id for connection in service.connections.list_all()]
|
||||||
assert ids == ["demo.personal", "demo.work"]
|
assert ids == ["demo.personal", "demo.work"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_broker_refresh_tool_returns_structured_error() -> None:
|
||||||
|
service = WfMcpService(store=FileStore(local_temp_root() / "broker_fail_store"))
|
||||||
|
service.register_connection(
|
||||||
|
ConnectionConfig(id="demo.personal", server="demo", account="personal")
|
||||||
|
)
|
||||||
|
service.register_adapter("demo", FailingDiscoveryAdapter())
|
||||||
|
|
||||||
|
server = create_broker_server(service)
|
||||||
|
|
||||||
|
_content, structured = asyncio.run(
|
||||||
|
server.call_tool("refresh_connection_catalog", {"connection_id": "demo.personal"})
|
||||||
|
)
|
||||||
|
assert structured == {
|
||||||
|
"connection_id": "demo.personal",
|
||||||
|
"refreshed": False,
|
||||||
|
"error_type": "PermissionError",
|
||||||
|
"error": "Access is denied",
|
||||||
|
}
|
||||||
|
|||||||
@@ -11,7 +11,13 @@ from wf_mcp import (
|
|||||||
WfMcpService,
|
WfMcpService,
|
||||||
)
|
)
|
||||||
|
|
||||||
from test_wf_mcp_support import FakeAdapter, echo_tool, finalize_tool, local_temp_root
|
from test_wf_mcp_support import (
|
||||||
|
FailingDiscoveryAdapter,
|
||||||
|
FakeAdapter,
|
||||||
|
echo_tool,
|
||||||
|
finalize_tool,
|
||||||
|
local_temp_root,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_service_builds_namespaced_catalog() -> None:
|
def test_service_builds_namespaced_catalog() -> None:
|
||||||
@@ -311,3 +317,29 @@ def test_service_can_invoke_raw_method_and_notification() -> None:
|
|||||||
assert "raw_method_completed" in event_kinds
|
assert "raw_method_completed" in event_kinds
|
||||||
assert "raw_notification_started" in event_kinds
|
assert "raw_notification_started" in event_kinds
|
||||||
assert "raw_notification_completed" in event_kinds
|
assert "raw_notification_completed" in event_kinds
|
||||||
|
|
||||||
|
|
||||||
|
def test_service_records_catalog_refresh_failures() -> None:
|
||||||
|
service = WfMcpService(store=FileStore(local_temp_root() / "refresh_fail_store"))
|
||||||
|
service.register_connection(
|
||||||
|
ConnectionConfig(id="demo.personal", server="demo", account="personal")
|
||||||
|
)
|
||||||
|
service.register_adapter("demo", FailingDiscoveryAdapter())
|
||||||
|
|
||||||
|
try:
|
||||||
|
asyncio.run(service.refresh_connection_catalog("demo.personal"))
|
||||||
|
except PermissionError as exc:
|
||||||
|
assert str(exc) == "Access is denied"
|
||||||
|
else:
|
||||||
|
raise AssertionError("expected refresh to fail")
|
||||||
|
|
||||||
|
failure_events = [
|
||||||
|
event
|
||||||
|
for event in service.list_events()
|
||||||
|
if event.kind == "catalog_refresh_failed"
|
||||||
|
]
|
||||||
|
assert len(failure_events) == 1
|
||||||
|
assert failure_events[0].payload == {
|
||||||
|
"error_type": "PermissionError",
|
||||||
|
"error": "Access is denied",
|
||||||
|
}
|
||||||
|
|||||||
@@ -243,7 +243,17 @@ class FakeAdapter:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class FailingDiscoveryAdapter(FakeAdapter):
|
||||||
|
async def list_tools(
|
||||||
|
self,
|
||||||
|
connection: ConnectionConfig,
|
||||||
|
auth: AuthRecord | None,
|
||||||
|
) -> list[DiscoveredTool]:
|
||||||
|
raise PermissionError("Access is denied")
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
"FailingDiscoveryAdapter",
|
||||||
"FakeAdapter",
|
"FakeAdapter",
|
||||||
"echo_tool",
|
"echo_tool",
|
||||||
"everything_server_connection",
|
"everything_server_connection",
|
||||||
|
|||||||
Reference in New Issue
Block a user