Skip to content

Commit 7d574e1

Browse files
committed
feat(session-group): add read_resource and get_prompt routing
ClientSessionGroup could aggregate resources and prompts and call_tool against them, but offered no routed way to actually read a resource or get a prompt: callers had to track the owning session themselves. Add read_resource(name, ...) and get_prompt(name, arguments, ...) mirroring call_tool: they resolve the aggregate key (honoring component_name_hook) to the owning session via new _resource_to_session / _prompt_to_session reverse indexes and forward the resource's wire URI / prompt's wire name. Both carry the same allow_input_required overloads as ClientSession. Reverse indexes are cleaned up on disconnect_from_server. Adds unit + in-memory end-to-end tests and documents the methods in the session-groups guide.
1 parent f1b6589 commit 7d574e1

3 files changed

Lines changed: 273 additions & 1 deletion

File tree

‎docs/client/session-groups.md‎

Lines changed: 15 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -60,6 +60,20 @@ Run it again. `print(sorted(group.tools))` now shows both:
6060
The hook runs on **every** name from **every** server, not only on conflicts: there is no
6161
prefix-on-collision mode. Pick one scheme and let it apply everywhere.
6262

63+
## Reading resources and prompts
64+
65+
`call_tool` is not the only routed call. `group.read_resource(name)` and `group.get_prompt(name, arguments)` work the same way: you pass the aggregate key (the one in `group.resources` / `group.prompts`, prefixed if you use a hook), and the group finds the owning session and forwards the call with the resource's real URI or the prompt's real name.
66+
67+
```python
68+
# `Library.hours` is the group key; `library://hours` goes on the wire.
69+
result = await group.read_resource("Library.hours")
70+
71+
# `arguments` are forwarded to the owning server unchanged.
72+
prompt = await group.get_prompt("Greeter.greet", {"name": "Ada"})
73+
```
74+
75+
Both raise `KeyError` if the key isn't an aggregated resource/prompt, and both accept the same `allow_input_required` flag as `ClientSession`, so a server that needs input mid-call behaves identically through the group.
76+
6377
## Adding and removing servers
6478

6579
`connect_to_server` returns the `ClientSession` it opened. Keep it if you ever want that server gone: `await group.disconnect_from_server(session)` removes its tools, resources, and prompts from the group.
@@ -74,7 +88,7 @@ If you already hold a connected `ClientSession` (`Client.session` is one), hand
7488

7589
* `ClientSessionGroup` holds many server connections and merges their tools, resources, and prompts into one `dict` each.
7690
* `connect_to_server(params)` per server. It takes transport parameters, never the URL or `Transport` a `Client` takes.
77-
* `group.call_tool(name, arguments)` routes to the owning server for you.
91+
* `group.call_tool(name, arguments)` routes to the owning server for you; `group.read_resource(name)` and `group.get_prompt(name, arguments)` route the same way.
7892
* Names must be unique across the whole group; two servers with a `search` tool cannot coexist on their own.
7993
* `component_name_hook=` rewrites every registered name. The dict key changes, the wire name does not.
8094
* `connect_with_session` adds a session you already hold; `disconnect_from_server` removes one.

‎src/mcp/client/session_group.py‎

Lines changed: 122 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -116,6 +116,8 @@ class _ComponentNames(BaseModel):
116116
# Client-server connection management.
117117
_sessions: dict[mcp.ClientSession, _ComponentNames]
118118
_tool_to_session: dict[str, mcp.ClientSession]
119+
_resource_to_session: dict[str, mcp.ClientSession]
120+
_prompt_to_session: dict[str, mcp.ClientSession]
119121
_exit_stack: contextlib.AsyncExitStack
120122
_session_exit_stacks: dict[mcp.ClientSession, contextlib.AsyncExitStack]
121123

@@ -138,6 +140,8 @@ def __init__(
138140

139141
self._sessions = {}
140142
self._tool_to_session = {}
143+
self._resource_to_session = {}
144+
self._prompt_to_session = {}
141145
if exit_stack is None:
142146
self._exit_stack = contextlib.AsyncExitStack()
143147
self._owns_exit_stack = True
@@ -249,6 +253,112 @@ async def call_tool(
249253
allow_input_required=allow_input_required,
250254
)
251255

256+
@overload
257+
async def read_resource(
258+
self,
259+
name: str,
260+
*,
261+
input_responses: types.InputResponses | None = None,
262+
request_state: str | None = None,
263+
meta: types.RequestParamsMeta | None = None,
264+
allow_input_required: Literal[False] = False,
265+
) -> types.ReadResourceResult: ...
266+
267+
@overload
268+
async def read_resource(
269+
self,
270+
name: str,
271+
*,
272+
input_responses: types.InputResponses | None = None,
273+
request_state: str | None = None,
274+
meta: types.RequestParamsMeta | None = None,
275+
allow_input_required: bool,
276+
) -> types.ReadResourceResult | types.InputRequiredResult: ...
277+
278+
async def read_resource(
279+
self,
280+
name: str,
281+
*,
282+
input_responses: types.InputResponses | None = None,
283+
request_state: str | None = None,
284+
meta: types.RequestParamsMeta | None = None,
285+
allow_input_required: bool = False,
286+
) -> types.ReadResourceResult | types.InputRequiredResult:
287+
"""Reads an aggregated resource, routing to the server that owns it.
288+
289+
``name`` is the aggregate key (i.e. the value used in ``resources``,
290+
which is affected by ``component_name_hook``), not necessarily the
291+
resource's wire URI.
292+
293+
Raises:
294+
KeyError: If ``name`` is not an aggregated resource.
295+
RuntimeError: If the server returns an ``InputRequiredResult`` and
296+
``allow_input_required`` is ``False``.
297+
"""
298+
session = self._resource_to_session[name]
299+
return await session.read_resource(
300+
self.resources[name].uri,
301+
input_responses=input_responses,
302+
request_state=request_state,
303+
meta=meta,
304+
allow_input_required=allow_input_required,
305+
)
306+
307+
@overload
308+
async def get_prompt(
309+
self,
310+
name: str,
311+
arguments: dict[str, str] | None = None,
312+
*,
313+
input_responses: types.InputResponses | None = None,
314+
request_state: str | None = None,
315+
meta: types.RequestParamsMeta | None = None,
316+
allow_input_required: Literal[False] = False,
317+
) -> types.GetPromptResult: ...
318+
319+
@overload
320+
async def get_prompt(
321+
self,
322+
name: str,
323+
arguments: dict[str, str] | None = None,
324+
*,
325+
input_responses: types.InputResponses | None = None,
326+
request_state: str | None = None,
327+
meta: types.RequestParamsMeta | None = None,
328+
allow_input_required: bool,
329+
) -> types.GetPromptResult | types.InputRequiredResult: ...
330+
331+
async def get_prompt(
332+
self,
333+
name: str,
334+
arguments: dict[str, str] | None = None,
335+
*,
336+
input_responses: types.InputResponses | None = None,
337+
request_state: str | None = None,
338+
meta: types.RequestParamsMeta | None = None,
339+
allow_input_required: bool = False,
340+
) -> types.GetPromptResult | types.InputRequiredResult:
341+
"""Gets an aggregated prompt, routing to the server that owns it.
342+
343+
``name`` is the aggregate key (i.e. the value used in ``prompts``,
344+
which is affected by ``component_name_hook``), not necessarily the
345+
prompt's wire name.
346+
347+
Raises:
348+
KeyError: If ``name`` is not an aggregated prompt.
349+
RuntimeError: If the server returns an ``InputRequiredResult`` and
350+
``allow_input_required`` is ``False``.
351+
"""
352+
session = self._prompt_to_session[name]
353+
return await session.get_prompt(
354+
self.prompts[name].name,
355+
arguments,
356+
input_responses=input_responses,
357+
request_state=request_state,
358+
meta=meta,
359+
allow_input_required=allow_input_required,
360+
)
361+
252362
async def disconnect_from_server(self, session: mcp.ClientSession) -> None:
253363
"""Disconnects from a single MCP server."""
254364

@@ -272,6 +382,12 @@ async def disconnect_from_server(self, session: mcp.ClientSession) -> None:
272382
for name in component_names.resources:
273383
if name in self._resources: # pragma: no branch
274384
del self._resources[name]
385+
if name in self._resource_to_session: # pragma: no branch
386+
del self._resource_to_session[name]
387+
# Remove prompts' reverse index for this session.
388+
for name in component_names.prompts:
389+
if name in self._prompt_to_session: # pragma: no branch
390+
del self._prompt_to_session[name]
275391
# Remove tools associated with the session.
276392
for name in component_names.tools:
277393
if name in self._tools: # pragma: no branch
@@ -382,13 +498,16 @@ async def _aggregate_components(self, server_info: types.Implementation, session
382498
resources_temp: dict[str, types.Resource] = {}
383499
tools_temp: dict[str, types.Tool] = {}
384500
tool_to_session_temp: dict[str, mcp.ClientSession] = {}
501+
resource_to_session_temp: dict[str, mcp.ClientSession] = {}
502+
prompt_to_session_temp: dict[str, mcp.ClientSession] = {}
385503

386504
# Query the server for its prompts and aggregate to list.
387505
try:
388506
prompts = (await session.list_prompts()).prompts
389507
for prompt in prompts:
390508
name = self._component_name(prompt.name, server_info)
391509
prompts_temp[name] = prompt
510+
prompt_to_session_temp[name] = session
392511
component_names.prompts.add(name)
393512
except MCPError as err: # pragma: no cover
394513
logging.warning(f"Could not fetch prompts: {err}")
@@ -399,6 +518,7 @@ async def _aggregate_components(self, server_info: types.Implementation, session
399518
for resource in resources:
400519
name = self._component_name(resource.name, server_info)
401520
resources_temp[name] = resource
521+
resource_to_session_temp[name] = session
402522
component_names.resources.add(name)
403523
except MCPError as err: # pragma: no cover
404524
logging.warning(f"Could not fetch resources: {err}")
@@ -442,6 +562,8 @@ async def _aggregate_components(self, server_info: types.Implementation, session
442562
self._resources.update(resources_temp)
443563
self._tools.update(tools_temp)
444564
self._tool_to_session.update(tool_to_session_temp)
565+
self._resource_to_session.update(resource_to_session_temp)
566+
self._prompt_to_session.update(prompt_to_session_temp)
445567

446568
def _component_name(self, name: str, server_info: types.Implementation) -> str:
447569
if self._component_name_hook:

‎tests/client/test_session_group.py‎

Lines changed: 136 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,13 +6,15 @@
66
import pytest
77

88
import mcp
9+
from mcp import Client
910
from mcp.client.session_group import (
1011
ClientSessionGroup,
1112
ClientSessionParameters,
1213
SseServerParameters,
1314
StreamableHttpParameters,
1415
)
1516
from mcp.client.stdio import StdioServerParameters
17+
from mcp.server import MCPServer
1618
from mcp.shared.exceptions import MCPError
1719

1820

@@ -402,3 +404,137 @@ async def test_client_session_group_establish_session_parameterized(
402404
# 3. Assert returned values
403405
assert returned_server_info is mock_initialize_result.server_info
404406
assert returned_session is mock_entered_session
407+
408+
409+
@pytest.mark.anyio
410+
async def test_read_resource_routes_to_owning_session():
411+
"""read_resource resolves the aggregate key to the owning session and passes the wire URI."""
412+
mock_session = mock.AsyncMock(spec=mcp.ClientSession)
413+
resource = types.Resource(name="hours", uri="library://hours")
414+
expected = types.ReadResourceResult(contents=[])
415+
mock_session.read_resource.return_value = expected
416+
417+
group = ClientSessionGroup()
418+
group._resources = {"Library.hours": resource}
419+
group._resource_to_session = {"Library.hours": mock_session}
420+
421+
result = await group.read_resource("Library.hours")
422+
423+
assert result is expected
424+
mock_session.read_resource.assert_awaited_once_with(
425+
"library://hours",
426+
input_responses=None,
427+
request_state=None,
428+
meta=None,
429+
allow_input_required=False,
430+
)
431+
432+
433+
@pytest.mark.anyio
434+
async def test_read_resource_unknown_name_raises_key_error():
435+
group = ClientSessionGroup()
436+
with pytest.raises(KeyError):
437+
await group.read_resource("missing")
438+
439+
440+
@pytest.mark.anyio
441+
async def test_get_prompt_routes_to_owning_session():
442+
"""get_prompt resolves the aggregate key to the owning session and passes the wire name."""
443+
mock_session = mock.AsyncMock(spec=mcp.ClientSession)
444+
prompt = types.Prompt(name="greet")
445+
expected = types.GetPromptResult(messages=[])
446+
mock_session.get_prompt.return_value = expected
447+
448+
group = ClientSessionGroup()
449+
group._prompts = {"Web.greet": prompt}
450+
group._prompt_to_session = {"Web.greet": mock_session}
451+
452+
result = await group.get_prompt("Web.greet", {"name": "Ada"})
453+
454+
assert result is expected
455+
mock_session.get_prompt.assert_awaited_once_with(
456+
"greet",
457+
{"name": "Ada"},
458+
input_responses=None,
459+
request_state=None,
460+
meta=None,
461+
allow_input_required=False,
462+
)
463+
464+
465+
@pytest.mark.anyio
466+
async def test_get_prompt_unknown_name_raises_key_error():
467+
group = ClientSessionGroup()
468+
with pytest.raises(KeyError):
469+
await group.get_prompt("missing")
470+
471+
472+
@pytest.mark.anyio
473+
async def test_disconnect_clears_resource_and_prompt_reverse_index():
474+
"""disconnect_from_server drops the session's resource/prompt routing entries."""
475+
session = mock.Mock(spec=mcp.ClientSession)
476+
group = ClientSessionGroup()
477+
group._resources = {"res1": mock.Mock(spec=types.Resource)}
478+
group._prompts = {"prm1": mock.Mock(spec=types.Prompt)}
479+
group._resource_to_session = {"res1": session}
480+
group._prompt_to_session = {"prm1": session}
481+
group._sessions = {
482+
session: ClientSessionGroup._ComponentNames(
483+
prompts={"prm1"},
484+
resources={"res1"},
485+
tools=set(),
486+
)
487+
}
488+
489+
await group.disconnect_from_server(session)
490+
491+
assert "res1" not in group._resource_to_session
492+
assert "prm1" not in group._prompt_to_session
493+
494+
495+
def _server_info(client: Client) -> types.Implementation:
496+
assert client.server_info is not None
497+
return client.server_info
498+
499+
500+
@pytest.mark.anyio
501+
async def test_read_resource_end_to_end_routes_through_the_owning_server():
502+
"""The group reads an aggregated resource against a real in-memory session."""
503+
server = MCPServer("Library")
504+
505+
@server.resource("library://hours")
506+
def hours() -> str:
507+
return "Mon-Fri 09:00-17:00"
508+
509+
async with Client(server) as client:
510+
group = ClientSessionGroup()
511+
await group.connect_with_session(_server_info(client), client.session)
512+
513+
(name,) = group.resources
514+
result = await group.read_resource(name)
515+
516+
assert isinstance(result, types.ReadResourceResult)
517+
(content,) = result.contents
518+
assert isinstance(content, types.TextResourceContents)
519+
assert content.text == "Mon-Fri 09:00-17:00"
520+
521+
522+
@pytest.mark.anyio
523+
async def test_get_prompt_end_to_end_routes_through_the_owning_server():
524+
"""The group gets an aggregated prompt against a real in-memory session."""
525+
server = MCPServer("Greeter")
526+
527+
@server.prompt()
528+
def greet(name: str) -> str:
529+
return f"Hello, {name}!"
530+
531+
async with Client(server) as client:
532+
group = ClientSessionGroup()
533+
await group.connect_with_session(_server_info(client), client.session)
534+
535+
result = await group.get_prompt("greet", {"name": "Ada"})
536+
537+
assert isinstance(result, types.GetPromptResult)
538+
(message,) = result.messages
539+
assert isinstance(message.content, types.TextContent)
540+
assert message.content.text == "Hello, Ada!"

0 commit comments

Comments
 (0)