fix type issue

This commit is contained in:
lda
2026-05-19 23:19:02 +07:00 Verified
parent c820294f88
commit 53a9e7f85d
2 changed files with 38 additions and 17 deletions
+7 -10
View File
@@ -4,8 +4,11 @@ from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from typing import Any
from mcp.client.session import ClientSession
from ..models import AuthRecord, ConnectionConfig
from ..sdk import ToolCallResult
from ..sdk.converters import tool_result_to_call_result
@dataclass(slots=True)
@@ -14,20 +17,14 @@ class PersistentMcpSession:
connection: ConnectionConfig
auth: AuthRecord | None
client: Any
client: ClientSession
close_callback: Callable[[], Awaitable[None]] | None = None
async def call_tool(self, tool_name: str, payload: dict[str, Any]) -> ToolCallResult:
return await self.client.call_tool(tool_name, payload)
result = await self.client.call_tool(tool_name, payload)
return tool_result_to_call_result(result)
async def close(self) -> None:
"""Close this runtime without assuming a specific SDK client shape."""
"""Close the transport/session stack owned by the runtime factory."""
if self.close_callback is not None:
await self.close_callback()
return
close = getattr(self.client, "close", None)
if close is None:
return
result = close()
if hasattr(result, "__await__"):
await result
+31 -7
View File
@@ -2,8 +2,10 @@ from __future__ import annotations
import asyncio
from dataclasses import dataclass, field
from typing import Any
from typing import Any, cast
from mcp.client.session import ClientSession
from mcp.types import CallToolResult
from wf_authoring import build_async_registry
from wf_core import RuntimeContext
from wf_mcp.capabilities import DiscoveredTool
@@ -46,14 +48,26 @@ class FakeStatefulClient:
closed: bool = False
calls: list[tuple[str, dict[str, Any]]] = field(default_factory=list)
async def call_tool(self, tool_name: str, payload: dict[str, Any]) -> ToolCallResult:
async def call_tool(
self, tool_name: str, payload: dict[str, Any]
) -> CallToolResult:
self.calls.append((tool_name, payload))
if tool_name == "browser_navigate":
self.page_open = True
return ToolCallResult(outcome="ok", output={"content": "opened"})
return CallToolResult(
content=[],
structuredContent={"content": "opened"},
)
if tool_name == "browser_snapshot" and self.page_open:
return ToolCallResult(outcome="ok", output={"content": "snapshot"})
return ToolCallResult(outcome="error", output={"message": "No open page"})
return CallToolResult(
content=[],
structuredContent={"content": "snapshot"},
)
return CallToolResult(
content=[],
structuredContent={"message": "No open page"},
isError=True,
)
async def close(self) -> None:
self.closed = True
@@ -129,7 +143,12 @@ def test_runtime_pool_reuses_stateful_session_for_same_connection() -> None:
) -> PersistentMcpSession:
client = FakeStatefulClient()
created_clients.append(client)
return PersistentMcpSession(connection=connection, auth=auth, client=client)
return PersistentMcpSession(
connection=connection,
auth=auth,
client=cast(ClientSession, client),
close_callback=client.close,
)
async def run_calls() -> ToolCallResult:
pool = McpRuntimePool(factory)
@@ -164,7 +183,12 @@ def test_runtime_pool_replaces_session_when_fingerprint_changes() -> None:
) -> PersistentMcpSession:
client = FakeStatefulClient()
created_clients.append(client)
return PersistentMcpSession(connection=connection, auth=auth, client=client)
return PersistentMcpSession(
connection=connection,
auth=auth,
client=cast(ClientSession, client),
close_callback=client.close,
)
async def run_calls() -> None:
pool = McpRuntimePool(factory)