refactor: add typed mcp source connection seam

This commit is contained in:
lda
2026-06-07 13:50:57 +07:00 Verified
parent f79741eb2a
commit 195a967527
27 changed files with 690 additions and 217 deletions
+17 -5
View File
@@ -9,6 +9,10 @@ from mcp.types import METHOD_NOT_FOUND
from wf_authoring import NodeSpec
from wf_sources_mcp.catalog import DiscoveredPrompt, DiscoveredResource, DiscoveredTool
from wf_sources_mcp.connections import (
McpSourceConnection,
mcp_source_connection_from_connection_config,
)
from wf_sources_mcp.sdk import BackendAdapter, ToolExecutor
from ..auth import AuthRecord
@@ -34,14 +38,18 @@ async def discover_connection_capabilities(
auth: AuthRecord | None,
adapter: BackendAdapter,
) -> DiscoveredConnectionCapabilities:
tools = await adapter.list_tools(connection, auth)
# Compatibility boundary: broker callers still pass ConnectionConfig. Runtime
# internals use McpSourceConnection so the session code can move to
# wf_sources_mcp in a later slice.
source_connection = mcp_source_connection_from_connection_config(connection)
tools = await adapter.list_tools(source_connection, auth)
resources = await _list_optional_capabilities(
lambda: adapter.list_resources(connection, auth)
lambda: adapter.list_resources(source_connection, auth)
)
prompts = await _list_optional_capabilities(
lambda: adapter.list_prompts(connection, auth)
lambda: adapter.list_prompts(source_connection, auth)
)
metadata = await adapter.get_connection_metadata(connection, auth)
metadata = await adapter.get_connection_metadata(source_connection, auth)
return DiscoveredConnectionCapabilities(
tools=tools,
resources=resources,
@@ -77,9 +85,13 @@ def specs_from_discovered_tools(
tools: list[DiscoveredTool],
emit_event: Callable[[McpEvent], None] | None = None,
) -> list[NodeSpec[Any, Any]]:
# Compatibility boundary: broker callers still pass ConnectionConfig. Runtime
# internals use McpSourceConnection so the session code can move to
# wf_sources_mcp in a later slice.
source_connection = mcp_source_connection_from_connection_config(connection)
return [
wrap_discovered_tool(
connection=connection,
connection=source_connection,
auth=auth,
executor=executor,
tool=tool,
+4 -1
View File
@@ -23,6 +23,7 @@ from wf_sources_mcp.catalog import (
CatalogPromptEntry,
CatalogResourceEntry,
)
from wf_sources_mcp.connections import mcp_source_connection_from_connection_config
from wf_sources_mcp.sdk import ToolExecutor
from wf_sources_mcp.storage import CatalogStore
@@ -273,8 +274,10 @@ class SourceCatalogService:
async def invoke_tool(payload: BaseModel) -> NodeReturn[BaseModel]:
connection = self.connection_lookup(entry.connection_id)
auth = self.load_auth(connection)
# Compatibility boundary: broker callers still pass ConnectionConfig.
source_connection = mcp_source_connection_from_connection_config(connection)
result = await self.tool_executor_for(connection).call_tool(
connection,
source_connection,
auth,
entry.local_name,
payload.model_dump(exclude_unset=True),
@@ -5,6 +5,7 @@ from dataclasses import dataclass, field
from typing import Any
from wf_api.source_registry_admin import WorkflowSourceRegistryMutationProvider
from wf_sources_mcp.connections import mcp_source_connection_from_connection_config
from wf_sources_mcp.source_registry import (
McpSourceRegistryEntry,
SourceRegistryFile,
@@ -162,8 +163,12 @@ class SourceRegistryAdminProvider(WorkflowSourceRegistryMutationProvider):
auth_diagnostics = []
if self.load_auth is not None:
for source_id in sorted(after):
# Compatibility boundary: broker callers still pass ConnectionConfig.
source_connection = mcp_source_connection_from_connection_config(
after[source_id]
)
diagnostic = connection_auth_diagnostic(
after[source_id],
source_connection,
load_auth_ref=self.load_auth,
)
if diagnostic is not None:
@@ -27,6 +27,10 @@ from wf_mcp.models import ConnectionConfig
from wf_mcp.shared.errors import error_payload
from wf_sources_mcp.auth import AuthRecord, connection_auth_diagnostic
from wf_sources_mcp.catalog.models import CatalogSnapshot
from wf_sources_mcp.connections import (
McpSourceConnection,
mcp_source_connection_from_connection_config,
)
from wf_sources_mcp.sdk import BackendAdapter, ToolExecutor
from wf_sources_mcp.storage import AuthStore, CatalogStore
@@ -74,6 +78,9 @@ class UpstreamTransportService:
old compatibility surface has no callers.
"""
# Compatibility boundary: broker callers still pass ConnectionConfig.
# Check legacy metadata for auth_ref first to avoid requiring transport
# metadata just for auth resolution.
auth_ref = connection.metadata.get("auth_ref")
if isinstance(auth_ref, str):
return self.load_auth(auth_ref)
@@ -98,6 +105,8 @@ class UpstreamTransportService:
) -> dict[str, Any]:
adapter = require_adapter(connection, self.adapters)
auth = self.load_connection_auth(connection)
# Compatibility boundary: broker callers still pass ConnectionConfig.
source_connection = mcp_source_connection_from_connection_config(connection)
self.event_sink(
make_event(
"resource_read_started",
@@ -106,7 +115,7 @@ class UpstreamTransportService:
payload={"uri": uri},
)
)
result = await adapter.read_resource(connection, auth, uri)
result = await adapter.read_resource(source_connection, auth, uri)
self.event_sink(
make_event(
"resource_read_completed",
@@ -126,6 +135,8 @@ class UpstreamTransportService:
) -> dict[str, Any]:
adapter = require_adapter(connection, self.adapters)
auth = self.load_connection_auth(connection)
# Compatibility boundary: broker callers still pass ConnectionConfig.
source_connection = mcp_source_connection_from_connection_config(connection)
self.event_sink(
make_event(
"prompt_get_started",
@@ -134,7 +145,7 @@ class UpstreamTransportService:
payload={"argument_keys": sorted((arguments or {}).keys())},
)
)
result = await adapter.get_prompt(connection, auth, local_name, arguments)
result = await adapter.get_prompt(source_connection, auth, local_name, arguments)
self.event_sink(
make_event(
"prompt_get_completed",
@@ -154,6 +165,8 @@ class UpstreamTransportService:
) -> dict[str, Any]:
adapter = require_adapter(connection, self.adapters)
auth = self.load_connection_auth(connection)
# Compatibility boundary: broker callers still pass ConnectionConfig.
source_connection = mcp_source_connection_from_connection_config(connection)
self.event_sink(
make_event(
"raw_method_started",
@@ -162,7 +175,7 @@ class UpstreamTransportService:
payload={"params": params or {}},
)
)
result = await adapter.invoke_method(connection, auth, method, params)
result = await adapter.invoke_method(source_connection, auth, method, params)
self.event_sink(
make_event(
"raw_method_completed",
@@ -182,6 +195,8 @@ class UpstreamTransportService:
) -> None:
adapter = require_adapter(connection, self.adapters)
auth = self.load_connection_auth(connection)
# Compatibility boundary: broker callers still pass ConnectionConfig.
source_connection = mcp_source_connection_from_connection_config(connection)
self.event_sink(
make_event(
"raw_notification_started",
@@ -190,7 +205,7 @@ class UpstreamTransportService:
payload={"params": params or {}},
)
)
await adapter.send_notification(connection, auth, method, params)
await adapter.send_notification(source_connection, auth, method, params)
self.event_sink(
make_event(
"raw_notification_completed",
@@ -309,8 +324,10 @@ class UpstreamTransportService:
)
)
continue
# Compatibility boundary: broker callers still pass ConnectionConfig.
source_connection = mcp_source_connection_from_connection_config(connection)
auth_diagnostic = connection_auth_diagnostic(
connection,
source_connection,
# The diagnostic helper passes the explicit auth_ref to this
# loader, matching load_connection_auth's auth_ref-first path.
load_auth_ref=self.load_auth,
@@ -323,7 +340,7 @@ class UpstreamTransportService:
adapter = require_adapter(connection, self.adapters)
auth = self.load_connection_auth(connection)
await asyncio.wait_for(
adapter.list_tools(connection, auth),
adapter.list_tools(source_connection, auth),
timeout=LIVE_SOURCE_CHECK_TIMEOUT_SECONDS,
)
except _LIVE_SOURCE_CHECK_FAILURES as exc:
+2 -20
View File
@@ -1,29 +1,11 @@
from __future__ import annotations
import re
from dataclasses import dataclass, field
from wf_sources_mcp.ids import CONNECTION_ID_PATTERN, parse_connection_id
from .models import ConnectionConfig
CONNECTION_ID_PATTERN = r"^[A-Za-z0-9_][A-Za-z0-9_.-]*$"
def parse_connection_id(connection_id: str) -> tuple[str, str]:
# Connection ids are logical source ids, but they also key persisted auth and
# catalog files. Keep this parser conservative so unsafe ids are rejected
# before they reach either registry or store boundaries.
if not re.fullmatch(CONNECTION_ID_PATTERN, connection_id):
raise ValueError(
"connection id must start with alphanumeric or underscore and contain "
"only [A-Za-z0-9_.-]"
)
if "." not in connection_id:
raise ValueError("connection id must look like '<server>.<account>'")
server, account = connection_id.split(".", 1)
if not server or not account:
raise ValueError("connection id must look like '<server>.<account>'")
return server, account
def qualify_node_name(connection_id: str, local_name: str) -> str:
parse_connection_id(connection_id)
+2 -2
View File
@@ -75,8 +75,8 @@ class McpRuntimePool:
async def call_tool(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
connection,
auth,
tool_name: str,
payload: dict[str, Any],
) -> ToolCallResult:
+23 -28
View File
@@ -19,6 +19,7 @@ from pydantic import AnyUrl
from wf_sources_mcp.auth import AuthRecord, mcp_auth_env, mcp_auth_headers
from wf_sources_mcp.catalog import DiscoveredPrompt, DiscoveredResource, DiscoveredTool
from wf_sources_mcp.connections import McpSourceConnection
from wf_sources_mcp.sdk import BackendAdapter, ToolCallResult
from wf_sources_mcp.sdk.converters import (
prompt_to_discovered,
@@ -26,31 +27,26 @@ from wf_sources_mcp.sdk.converters import (
tool_result_to_call_result,
tool_to_discovered,
)
from ..models import ConnectionConfig
from wf_sources_mcp.transports import HttpSourceTransport, StdioSourceTransport
class McpSdkAdapter(BackendAdapter):
@asynccontextmanager
async def _session(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
):
transport = connection.metadata.get("transport", "stdio")
if transport == "stdio":
command = connection.metadata["command"]
args = list(connection.metadata.get("args", []))
env = connection.metadata.get("env")
cwd = connection.metadata.get("cwd")
transport = connection.transport
if isinstance(transport, StdioSourceTransport):
auth_env = mcp_auth_env(auth)
env = dict(transport.env)
if auth_env:
env = {**(env or {}), **auth_env}
env = {**env, **auth_env}
params = StdioServerParameters(
command=command,
args=args,
command=transport.command,
args=list(transport.args),
env=env,
cwd=cwd,
)
async with stdio_client(params) as (read_stream, write_stream):
async with ClientSession(read_stream, write_stream) as session:
@@ -58,13 +54,12 @@ class McpSdkAdapter(BackendAdapter):
yield session
return
if transport == "streamable_http":
url = connection.metadata["url"]
if isinstance(transport, HttpSourceTransport):
headers = mcp_auth_headers(auth)
http_client = httpx.AsyncClient(headers=headers or None)
async with http_client:
async with streamable_http_client(
url,
str(transport.url),
http_client=http_client,
) as (read_stream, write_stream, _get_session_id):
async with ClientSession(read_stream, write_stream) as session:
@@ -72,11 +67,11 @@ class McpSdkAdapter(BackendAdapter):
yield session
return
raise ValueError(f"unsupported MCP transport {transport!r}")
raise ValueError(f"unsupported MCP transport {transport.kind!r}")
async def list_tools(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
) -> list[DiscoveredTool]:
async with self._session(connection, auth) as session:
@@ -85,7 +80,7 @@ class McpSdkAdapter(BackendAdapter):
async def list_resources(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
) -> list[DiscoveredResource]:
async with self._session(connection, auth) as session:
@@ -94,7 +89,7 @@ class McpSdkAdapter(BackendAdapter):
async def list_prompts(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
) -> list[DiscoveredPrompt]:
async with self._session(connection, auth) as session:
@@ -103,17 +98,17 @@ class McpSdkAdapter(BackendAdapter):
async def get_connection_metadata(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
) -> dict[str, Any]:
return {
"server": connection.server,
"transport": connection.metadata.get("transport", "stdio"),
"server": connection.provider,
"transport": connection.transport.kind,
}
async def read_resource(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
uri: str,
) -> dict[str, Any]:
@@ -123,7 +118,7 @@ class McpSdkAdapter(BackendAdapter):
async def get_prompt(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
prompt_name: str,
arguments: dict[str, str] | None = None,
@@ -134,7 +129,7 @@ class McpSdkAdapter(BackendAdapter):
async def invoke_method(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
method: str,
params: dict[str, Any] | None = None,
@@ -148,7 +143,7 @@ class McpSdkAdapter(BackendAdapter):
async def send_notification(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
method: str,
params: dict[str, Any] | None = None,
@@ -160,7 +155,7 @@ class McpSdkAdapter(BackendAdapter):
async def call_tool(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
tool_name: str,
payload: dict[str, Any],
+2 -2
View File
@@ -21,9 +21,9 @@ if TYPE_CHECKING:
from fastmcp.resources.template import ResourceTemplate
from fastmcp.tools.base import Tool
from wf_sources_mcp.ids import RESERVED_CONNECTION_IDS
ADMIN_NAMESPACE = "wf.admin"
RESERVED_CONNECTION_IDS = frozenset({ADMIN_NAMESPACE, "wf.mcp"})
"""Source ids reserved by wf-mcp system capabilities."""
@dataclass(frozen=True, slots=True)
+3 -4
View File
@@ -9,12 +9,11 @@ from pydantic import BaseModel, ConfigDict, Field, create_model
from wf_authoring import NodeReturn, NodeSpec
from wf_core import RuntimeContext
from wf_mcp.broker.events import McpEvent, make_event
from wf_sources_mcp.auth import AuthRecord
from wf_sources_mcp.catalog import DiscoveredTool
from wf_sources_mcp.connections import McpSourceConnection
from wf_sources_mcp.sdk import ToolExecutor
from ..auth import AuthRecord
from ..models import ConnectionConfig
_JSON_TYPE_MAP: dict[str, object] = {
"string": str,
"integer": int,
@@ -118,7 +117,7 @@ def _model_from_schema(name: str, schema: dict[str, Any]) -> type[BaseModel]:
def wrap_discovered_tool(
*,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
executor: ToolExecutor,
tool: DiscoveredTool,
+29 -6
View File
@@ -21,23 +21,31 @@ from .auth import (
)
if TYPE_CHECKING:
from .connections import (
McpSourceConnection,
mcp_source_connection_from_connection_config,
mcp_source_connection_from_registry_entry,
)
from .source_registry import (
FileSourceRegistryStore,
HttpSourceTransport,
McpSourceRegistryEntry,
SourceRegistryFile,
SourceRegistryStore,
SourceTransport,
StdioSourceTransport,
connection_config_to_registry_entry,
registry_entry_to_connection_config,
workflow_mcp_source_to_connection_config,
)
from .transports import (
HttpSourceTransport,
SourceTransport,
StdioSourceTransport,
)
__all__ = [
"AuthRecord",
"FileSourceRegistryStore",
"HttpSourceTransport",
"McpSourceConnection",
"McpSourceRegistryEntry",
"SourceRegistryFile",
"SourceRegistryStore",
@@ -50,6 +58,8 @@ __all__ = [
"mcp_auth_env",
"mcp_auth_from_neutral",
"mcp_auth_headers",
"mcp_source_connection_from_connection_config",
"mcp_source_connection_from_registry_entry",
"neutral_auth_from_mcp",
"registry_entry_to_connection_config",
"workflow_mcp_source_to_connection_config",
@@ -57,14 +67,19 @@ __all__ = [
def __getattr__(name: str) -> object:
if name in {
"McpSourceConnection",
"mcp_source_connection_from_connection_config",
"mcp_source_connection_from_registry_entry",
}:
from . import connections
return getattr(connections, name)
if name in {
"FileSourceRegistryStore",
"HttpSourceTransport",
"McpSourceRegistryEntry",
"SourceRegistryFile",
"SourceRegistryStore",
"SourceTransport",
"StdioSourceTransport",
"connection_config_to_registry_entry",
"registry_entry_to_connection_config",
"workflow_mcp_source_to_connection_config",
@@ -72,4 +87,12 @@ def __getattr__(name: str) -> object:
from . import source_registry
return getattr(source_registry, name)
if name in {
"HttpSourceTransport",
"SourceTransport",
"StdioSourceTransport",
}:
from . import transports
return getattr(transports, name)
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
+14 -11
View File
@@ -1,22 +1,18 @@
"""MCP upstream-source auth helpers.
This module is canonical for MCP-as-source auth interpretation. The temporary
TYPE_CHECKING dependency on `wf_mcp.broker.models.ConnectionConfig` exists until
connection runtime DTOs move out of the compatibility MCP facade.
This module is canonical for MCP-as-source auth interpretation. Runtime-facing
helpers consume source-connection-like objects instead of broker config DTOs.
"""
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
from typing import Any, Protocol
from wf_api.auth import AuthRecord as NeutralAuthRecord
from wf_artifacts import DependencyDiagnostic, DiagnosticSeverity
if TYPE_CHECKING:
from wf_mcp.broker.models import ConnectionConfig
@dataclass(slots=True)
class AuthRecord:
@@ -85,11 +81,18 @@ def mcp_auth_env(auth: AuthRecord | None) -> dict[str, str]:
}
def auth_ref_for_connection(connection: ConnectionConfig) -> str | None:
class SourceConnectionLike(Protocol):
@property
def id(self) -> str: ...
@property
def auth_ref(self) -> str | None: ...
def auth_ref_for_connection(connection: SourceConnectionLike) -> str | None:
"""Return the explicit auth ref for one source connection, if present."""
auth_ref = connection.metadata.get("auth_ref")
return auth_ref if isinstance(auth_ref, str) else None
return connection.auth_ref
def auth_missing_diagnostic(
@@ -117,7 +120,7 @@ def auth_missing_diagnostic(
def connection_auth_diagnostic(
connection: ConnectionConfig,
connection: SourceConnectionLike,
*,
load_auth_ref: Callable[[str], AuthRecord | None],
logical_ref: str | None = None,
+148
View File
@@ -0,0 +1,148 @@
from __future__ import annotations
from dataclasses import dataclass, field
from typing import TYPE_CHECKING
from wf_sources_mcp.ids import parse_connection_id
from wf_sources_mcp.source_registry import McpSourceRegistryEntry
from wf_sources_mcp.transports import (
HttpSourceTransport,
SourceTransport,
StdioSourceTransport,
)
if TYPE_CHECKING:
from wf_mcp.broker.models import ConnectionConfig
_FLAT_HTTP_TRANSPORTS = {"http", "streamable-http", "streamable_http", "sse"}
_CONNECTION_METADATA_KEYS = {
"transport",
"command",
"args",
"env",
"cwd",
"url",
"headers",
"profile",
"auth_ref",
}
@dataclass(frozen=True, slots=True)
class McpSourceConnection:
"""Typed runtime-facing MCP source connection.
This is the object runtime/session code should consume. Legacy broker
`ConnectionConfig.metadata` remains at the compatibility edge only.
"""
id: str
provider: str
account: str
transport: SourceTransport
enabled: bool = True
profile: str | None = None
auth_ref: str | None = None
metadata: dict[str, object] = field(default_factory=dict)
def __post_init__(self) -> None:
provider, account = parse_connection_id(self.id)
if not self.provider:
raise ValueError("provider must not be empty")
if not self.account:
raise ValueError("account must not be empty")
if provider != self.provider or account != self.account:
raise ValueError(
"MCP source connection id must match provider/account fields"
)
def mcp_source_connection_from_registry_entry(
entry: McpSourceRegistryEntry,
) -> McpSourceConnection:
"""Adapt persisted desired-source registry state to runtime source shape."""
return McpSourceConnection(
id=entry.id,
provider=entry.provider,
account=entry.account,
enabled=entry.enabled,
profile=entry.profile,
transport=entry.transport,
auth_ref=entry.auth_ref,
metadata=dict(entry.metadata),
)
def mcp_source_connection_from_connection_config(
connection: ConnectionConfig,
) -> McpSourceConnection:
"""Adapt legacy broker connection config into typed source shape.
Keep all metadata-bag reads in this compatibility converter. Runtime/session
code should use `McpSourceConnection.transport` directly.
"""
transport = _transport_from_connection_metadata(connection)
profile = connection.metadata.get("profile")
auth_ref = connection.metadata.get("auth_ref")
metadata = {
str(key): value
for key, value in connection.metadata.items()
if key not in _CONNECTION_METADATA_KEYS
}
return McpSourceConnection(
id=connection.id,
provider=connection.server,
account=connection.account,
enabled=connection.enabled,
profile=profile if isinstance(profile, str) else None,
transport=transport,
auth_ref=auth_ref if isinstance(auth_ref, str) else None,
metadata=metadata,
)
def _transport_from_connection_metadata(connection: ConnectionConfig) -> SourceTransport:
transport = connection.metadata.get("transport")
if isinstance(transport, dict):
kind = transport.get("kind")
if kind == "stdio":
return StdioSourceTransport.model_validate(transport)
if kind == "http":
return HttpSourceTransport.model_validate(transport)
raise ValueError(
f"connection {connection.id!r} has unsupported metadata.transport.kind {kind!r}"
)
if isinstance(transport, str):
if transport == "stdio":
return StdioSourceTransport(
command=str(connection.metadata.get("command", "")),
args=tuple(str(arg) for arg in connection.metadata.get("args", ())),
env={
str(key): str(value)
for key, value in dict(connection.metadata.get("env", {})).items()
},
)
if transport in _FLAT_HTTP_TRANSPORTS:
url = connection.metadata.get("url", "")
return HttpSourceTransport(
url=url if isinstance(url, str) else str(url), # type: ignore[arg-type]
headers={
str(key): str(value)
for key, value in dict(
connection.metadata.get("headers", {})
).items()
},
)
raise ValueError(
f"connection {connection.id!r} has unrecognized metadata.transport {transport!r}"
)
raise ValueError(f"connection {connection.id!r} requires metadata.transport")
__all__ = [
"McpSourceConnection",
"mcp_source_connection_from_connection_config",
"mcp_source_connection_from_registry_entry",
]
+35
View File
@@ -0,0 +1,35 @@
from __future__ import annotations
import re
CONNECTION_ID_PATTERN = r"^[A-Za-z0-9_][A-Za-z0-9_.-]*$"
RESERVED_CONNECTION_IDS = frozenset({"wf.admin", "wf.mcp"})
"""Source ids reserved by built-in workflow/MCP control surfaces."""
def parse_connection_id(connection_id: str) -> tuple[str, str]:
"""Validate and split one MCP source id into provider/account parts.
Source ids also key persisted auth, registry, and catalog files. Keep this
conservative so unsafe ids are rejected before reaching store boundaries.
"""
if not re.fullmatch(CONNECTION_ID_PATTERN, connection_id):
raise ValueError(
"connection id must start with alphanumeric or underscore and contain "
"only [A-Za-z0-9_.-]"
)
if "." not in connection_id:
raise ValueError("connection id must look like '<server>.<account>'")
server, account = connection_id.split(".", 1)
if not server or not account:
raise ValueError("connection id must look like '<server>.<account>'")
return server, account
__all__ = [
"CONNECTION_ID_PATTERN",
"RESERVED_CONNECTION_IDS",
"parse_connection_id",
]
+13 -19
View File
@@ -1,19 +1,13 @@
"""Protocol/result contracts for MCP upstream source providers.
The temporary `wf_mcp.broker.models.ConnectionConfig` dependency remains until
broker runtime connection DTOs move to a neutral/source-provider package.
"""
"""Protocol/result contracts for MCP upstream source providers."""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, Protocol
from typing import Any, Protocol
from wf_sources_mcp.auth import AuthRecord
from wf_sources_mcp.catalog import DiscoveredPrompt, DiscoveredResource, DiscoveredTool
if TYPE_CHECKING:
from wf_mcp.broker.models import ConnectionConfig
from wf_sources_mcp.connections import McpSourceConnection
@dataclass(slots=True)
@@ -26,38 +20,38 @@ class ToolCallResult:
class BackendAdapter(Protocol):
async def list_tools(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
) -> list[DiscoveredTool]: ...
async def list_resources(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
) -> list[DiscoveredResource]: ...
async def list_prompts(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
) -> list[DiscoveredPrompt]: ...
async def get_connection_metadata(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
) -> dict[str, Any]: ...
async def read_resource(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
uri: str,
) -> dict[str, Any]: ...
async def get_prompt(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
prompt_name: str,
arguments: dict[str, str] | None = None,
@@ -65,7 +59,7 @@ class BackendAdapter(Protocol):
async def invoke_method(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
method: str,
params: dict[str, Any] | None = None,
@@ -73,7 +67,7 @@ class BackendAdapter(Protocol):
async def send_notification(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
method: str,
params: dict[str, Any] | None = None,
@@ -81,7 +75,7 @@ class BackendAdapter(Protocol):
async def call_tool(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
tool_name: str,
payload: dict[str, Any],
@@ -98,7 +92,7 @@ class ToolExecutor(Protocol):
async def call_tool(
self,
connection: ConnectionConfig,
connection: McpSourceConnection,
auth: AuthRecord | None,
tool_name: str,
payload: dict[str, Any],
+8 -32
View File
@@ -8,14 +8,9 @@ runtime DTOs move out of the compatibility MCP facade.
from __future__ import annotations
from pathlib import Path
from typing import TYPE_CHECKING, Annotated, Literal, Protocol
from typing import TYPE_CHECKING, Literal, Protocol
from pydantic import (
AnyHttpUrl,
Field,
field_validator,
model_validator,
)
from pydantic import Field, field_validator, model_validator
from wf_api.source_registry import (
AtomicJsonRegistryStore,
@@ -25,12 +20,12 @@ from wf_api.source_registry import (
from wf_api.source_registry import (
SourceRegistryStore as GenericSourceRegistryStore,
)
# Temporary low-level compatibility imports. `wf_mcp.shared.names` currently
# pulls in FastMCP transitively; keep this visible until reserved-name parsing
# moves to a neutral/source package.
from wf_mcp.connections import parse_connection_id
from wf_mcp.shared.names import RESERVED_CONNECTION_IDS
from wf_sources_mcp.ids import RESERVED_CONNECTION_IDS, parse_connection_id
from wf_sources_mcp.transports import (
HttpSourceTransport,
SourceTransport,
StdioSourceTransport,
)
if TYPE_CHECKING:
from wf_mcp.models import ConnectionConfig
@@ -50,25 +45,6 @@ _TRANSPORT_METADATA_KEYS = {
}
class StdioSourceTransport(SourceRegistryBaseModel):
kind: Literal["stdio"] = "stdio"
command: str = Field(min_length=1)
args: tuple[str, ...] = ()
env: dict[str, str] = Field(default_factory=dict)
class HttpSourceTransport(SourceRegistryBaseModel):
kind: Literal["http"] = "http"
url: AnyHttpUrl
headers: dict[str, str] = Field(default_factory=dict)
SourceTransport = Annotated[
StdioSourceTransport | HttpSourceTransport,
Field(discriminator="kind"),
]
class McpSourceRegistryEntry(SourceRegistryBaseModel):
"""Desired MCP source configuration persisted by server-owned mutation."""
+33
View File
@@ -0,0 +1,33 @@
from __future__ import annotations
from typing import Annotated, Literal
from pydantic import AnyHttpUrl, Field
from wf_api.source_registry import SourceRegistryBaseModel
class StdioSourceTransport(SourceRegistryBaseModel):
kind: Literal["stdio"] = "stdio"
command: str = Field(min_length=1)
args: tuple[str, ...] = ()
env: dict[str, str] = Field(default_factory=dict)
class HttpSourceTransport(SourceRegistryBaseModel):
kind: Literal["http"] = "http"
url: AnyHttpUrl
headers: dict[str, str] = Field(default_factory=dict)
SourceTransport = Annotated[
StdioSourceTransport | HttpSourceTransport,
Field(discriminator="kind"),
]
__all__ = [
"HttpSourceTransport",
"SourceTransport",
"StdioSourceTransport",
]