Skip to content

Commit 74c41a5

Browse files
Resolve: Types
1 parent b5c86a9 commit 74c41a5

2 files changed

Lines changed: 10 additions & 23 deletions

File tree

src/amrita_core/tools/mcp.py

Lines changed: 9 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -4,16 +4,13 @@
44
from collections.abc import Awaitable, Callable, Iterable
55
from contextlib import nullcontext
66
from copy import deepcopy
7-
from typing import Any, Generic, TypeVar, overload
7+
from pathlib import Path
8+
from typing import Any, overload
89

9-
from fastmcp import Client, FastMCP
10+
from fastmcp import Client
1011
from fastmcp.client.client import CallToolResult
11-
from fastmcp.client.transports.base import ClientTransport
12-
from fastmcp.mcp_config import MCPConfig
1312
from mcp.types import TextContent
14-
from pydantic import AnyUrl
1513
from typing_extensions import Self
16-
from zipp import Path
1714

1815
from amrita_core.logging import logger
1916

@@ -27,23 +24,16 @@
2724
cast_mcp_properties_to_amrita,
2825
)
2926

30-
MCP_SERVER_SCRIPT_TYPE = TypeVar(
31-
"MCP_SERVER_SCRIPT_TYPE",
32-
str,
33-
ClientTransport,
34-
AnyUrl,
35-
FastMCP,
36-
MCPConfig,
37-
dict[str, Any],
38-
covariant=True,
27+
MCP_SERVER_SCRIPT_TYPE = (
28+
str | Path # TODO: Support all types of scripts
3929
)
4030

4131

4232
class NOT_GIVEN:
4333
pass
4434

4535

46-
class MCPClient(Generic[MCP_SERVER_SCRIPT_TYPE]):
36+
class MCPClient:
4737
"""Reusable MCP Client"""
4838

4939
mcp_client: Client | None = None
@@ -174,9 +164,7 @@ def __init__(self, tools_manager: MultiToolsManager | None = None) -> None:
174164
self.script_to_clients = {}
175165
self._lock = Lock()
176166

177-
def get_client_by_script(
178-
self, server_script: MCP_SERVER_SCRIPT_TYPE
179-
) -> MCPClient[MCP_SERVER_SCRIPT_TYPE]:
167+
def get_client_by_script(self, server_script: MCP_SERVER_SCRIPT_TYPE) -> MCPClient:
180168
"""Get MCP Client (without operating stored MCP Server)
181169
Args:
182170
server_script (str, optional): MCP Server script path (or URI).
@@ -219,7 +207,7 @@ def register_only(
219207
self,
220208
*,
221209
server_script: MCP_SERVER_SCRIPT_TYPE | None = None,
222-
client: MCPClient[MCP_SERVER_SCRIPT_TYPE] | None = None,
210+
client: MCPClient | None = None,
223211
) -> Self:
224212
"""Register MCP Server only, without initialization"""
225213
if client is not None:
@@ -253,9 +241,7 @@ async def initialize_this(
253241
self, server_script: MCP_SERVER_SCRIPT_TYPE, fail_then_raise: bool = False
254242
) -> Self:
255243
"""Register and initialize single MCP Server"""
256-
client: MCPClient[MCP_SERVER_SCRIPT_TYPE] = self.get_client_by_script(
257-
server_script
258-
)
244+
client: MCPClient = self.get_client_by_script(server_script)
259245
async with self._lock:
260246
try:
261247
await self._load_this(client)

tests/test_functions.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -85,6 +85,7 @@ def test_agent_runtime_init_with_session_string(mock_config, mock_preset):
8585
assert runtime.session_id == session_id
8686
assert runtime.context == mock_session_data.memory
8787

88+
8889
def test_agent_runtime_init_new_session(mock_config, mock_preset):
8990
"""Test AgentRuntime initialization creating new session."""
9091
with (

0 commit comments

Comments
 (0)