This commit is contained in:
lda
2026-04-29 17:29:48 +07:00 Verified
parent 7e020c1291
commit a0819e8988
5 changed files with 240 additions and 1 deletions
+6
View File
@@ -1,3 +1,4 @@
from .adapters import BackendAdapter, DiscoveredTool, ToolCallResult
from .catalog import CombinedCatalog
from .connections import ConnectionRegistry, parse_connection_id, qualify_node_name
from .models import (
@@ -9,18 +10,23 @@ from .models import (
)
from .service import WfMcpService
from .store import FileStore, Store
from .wrappers import wrap_discovered_tool
__all__ = [
"AuthRecord",
"BackendAdapter",
"CatalogNodeEntry",
"CatalogSnapshot",
"CombinedCatalog",
"ConnectionConfig",
"ConnectionRegistry",
"DiscoveredTool",
"FileStore",
"RawWorkflowPlan",
"Store",
"ToolCallResult",
"WfMcpService",
"parse_connection_id",
"qualify_node_name",
"wrap_discovered_tool",
]
+39
View File
@@ -0,0 +1,39 @@
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Protocol
from .models import AuthRecord, ConnectionConfig
@dataclass(slots=True)
class DiscoveredTool:
name: str
description: str | None
input_schema: dict[str, Any]
output_schema: dict[str, Any]
outcomes: tuple[str, ...] = ("ok",)
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass(slots=True)
class ToolCallResult:
outcome: str
output: dict[str, Any] = field(default_factory=dict)
meta: dict[str, Any] = field(default_factory=dict)
class BackendAdapter(Protocol):
async def list_tools(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
) -> list[DiscoveredTool]: ...
async def call_tool(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
tool_name: str,
payload: dict[str, Any],
) -> ToolCallResult: ...
+34
View File
@@ -7,10 +7,12 @@ from typing import Any
from wf_authoring import NodeSpec, build_async_registry
from wf_core import NodeUse, Workflow, execute_workflow_async
from .adapters import BackendAdapter
from .catalog import CombinedCatalog, snapshot_from_specs
from .connections import ConnectionRegistry, parse_connection_id, qualify_node_name
from .models import AuthRecord, CatalogSnapshot, ConnectionConfig, RawWorkflowPlan
from .store import Store
from .wrappers import wrap_discovered_tool
def _qualify_spec(connection_id: str, spec: NodeSpec[Any, Any]) -> NodeSpec[Any, Any]:
@@ -30,6 +32,7 @@ class WfMcpService:
store: Store
default_catalog_max_age_seconds: int = 300
connections: ConnectionRegistry = field(default_factory=ConnectionRegistry)
adapters: dict[str, BackendAdapter] = field(default_factory=dict)
specs_by_connection: dict[str, dict[str, NodeSpec[Any, Any]]] = field(
default_factory=dict
)
@@ -38,6 +41,9 @@ class WfMcpService:
parse_connection_id(connection.id)
self.connections.register(connection)
def register_adapter(self, server: str, adapter: BackendAdapter) -> None:
self.adapters[server] = adapter
def save_auth(self, record: AuthRecord) -> None:
self.store.save_auth(record)
@@ -72,6 +78,34 @@ class WfMcpService:
snapshots[connection.id] = snapshot
return CombinedCatalog(snapshots=snapshots)
async def refresh_connection_catalog(
self,
connection_id: str,
*,
max_age_seconds: int | None = None,
) -> None:
connection = self.connections.get(connection_id)
adapter = self.adapters.get(connection.server)
if adapter is None:
raise KeyError(f"no adapter registered for server {connection.server!r}")
auth = self.load_auth(connection_id)
tools = await adapter.list_tools(connection, auth)
specs = [
wrap_discovered_tool(
connection=connection,
auth=auth,
adapter=adapter,
tool=tool,
)
for tool in tools
]
self.register_specs(
connection_id,
*specs,
max_age_seconds=max_age_seconds,
)
def compile_plan(self, plan: RawWorkflowPlan) -> Workflow:
node_defs: dict[str, Any] = {}
for step in plan.nodes:
+71
View File
@@ -0,0 +1,71 @@
from __future__ import annotations
from typing import Any, cast
from pydantic import BaseModel, ConfigDict, Field, create_model
from wf_authoring import NodeReturn, NodeSpec
from wf_core import RuntimeContext
from .adapters import BackendAdapter, DiscoveredTool
from .models import AuthRecord, ConnectionConfig
def _model_from_schema(name: str, schema: dict[str, Any]) -> type[BaseModel]:
properties = cast(dict[str, Any], schema.get("properties", {}))
required = set(cast(list[str], schema.get("required", [])))
field_defs: dict[str, tuple[object, object]] = {}
for field_name in properties:
default = ... if field_name in required else None
field_defs[field_name] = (Any, Field(default=default))
raw_field_defs = cast(dict[str, Any], field_defs)
model = create_model(
name,
__config__=ConfigDict(extra="allow"),
**raw_field_defs,
)
return cast(type[BaseModel], model)
def wrap_discovered_tool(
*,
connection: ConnectionConfig,
auth: AuthRecord | None,
adapter: BackendAdapter,
tool: DiscoveredTool,
) -> NodeSpec[BaseModel, BaseModel]:
input_model = _model_from_schema(
f"{connection.id}_{tool.name}_Input",
tool.input_schema,
)
output_model = _model_from_schema(
f"{connection.id}_{tool.name}_Output",
tool.output_schema,
)
async def invoke_tool(
payload: BaseModel,
ctx: RuntimeContext,
) -> NodeReturn[BaseModel]:
result = await adapter.call_tool(
connection=connection,
auth=auth,
tool_name=tool.name,
payload=payload.model_dump(),
)
return NodeReturn(
outcome=result.outcome,
output=output_model.model_validate(result.output),
)
return NodeSpec(
name=tool.name,
input_model=input_model,
output_model=output_model,
outcomes=tool.outcomes,
fn=invoke_tool,
description=tool.description,
is_async=True,
)