adapters
This commit is contained in:
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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: ...
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
Reference in New Issue
Block a user