|
1 | 1 | import base64 |
2 | 2 | import logging |
| 3 | +from collections.abc import AsyncIterator |
| 4 | +from contextlib import asynccontextmanager |
| 5 | +from dataclasses import dataclass |
3 | 6 | from pathlib import Path |
4 | 7 | from types import SimpleNamespace |
5 | 8 | from typing import Annotated, Any |
@@ -1337,6 +1340,45 @@ def prompt_no_context(text: str) -> str: |
1337 | 1340 | assert content.text == "Prompt 'test' works" |
1338 | 1341 |
|
1339 | 1342 |
|
| 1343 | +async def test_parameterized_context_carries_the_request_into_templates_and_prompts(): |
| 1344 | + """`ctx: Context[AppState]` on a resource template or a sync or async prompt is the request's own |
| 1345 | + context, as it is on a tool: lifespan state and the negotiated protocol version are readable.""" |
| 1346 | + |
| 1347 | + @dataclass |
| 1348 | + class AppState: |
| 1349 | + greeting: str |
| 1350 | + |
| 1351 | + @asynccontextmanager |
| 1352 | + async def lifespan(_: MCPServer[AppState]) -> AsyncIterator[AppState]: |
| 1353 | + yield AppState(greeting="Hello") |
| 1354 | + |
| 1355 | + mcp = MCPServer(lifespan=lifespan) |
| 1356 | + |
| 1357 | + @mcp.resource("greeting://{name}") |
| 1358 | + def greeting(name: str, ctx: Context[AppState]) -> str: |
| 1359 | + return f"{ctx.request_context.lifespan_context.greeting}, {name} ({ctx.protocol_version})" |
| 1360 | + |
| 1361 | + @mcp.prompt() |
| 1362 | + def greet_sync(name: str, ctx: Context[AppState]) -> str: |
| 1363 | + return f"{ctx.request_context.lifespan_context.greeting}, {name} ({ctx.protocol_version})" |
| 1364 | + |
| 1365 | + @mcp.prompt() |
| 1366 | + async def greet_async(name: str, ctx: Context[AppState]) -> str: |
| 1367 | + return f"{ctx.request_context.lifespan_context.greeting}, {name} ({ctx.protocol_version})" |
| 1368 | + |
| 1369 | + async with Client(mcp, mode="2026-07-28") as client: |
| 1370 | + resource = await client.read_resource("greeting://Alice") |
| 1371 | + sync_prompt = await client.get_prompt("greet_sync", {"name": "Alice"}) |
| 1372 | + async_prompt = await client.get_prompt("greet_async", {"name": "Alice"}) |
| 1373 | + |
| 1374 | + assert resource.contents == [ |
| 1375 | + TextResourceContents(uri="greeting://Alice", mime_type="text/plain", text="Hello, Alice (2026-07-28)") |
| 1376 | + ] |
| 1377 | + expected = [PromptMessage(role="user", content=TextContent(type="text", text="Hello, Alice (2026-07-28)"))] |
| 1378 | + assert sync_prompt.messages == expected |
| 1379 | + assert async_prompt.messages == expected |
| 1380 | + |
| 1381 | + |
1340 | 1382 | class TestServerPrompts: |
1341 | 1383 | """Test prompt functionality in MCPServer server.""" |
1342 | 1384 |
|
@@ -2151,6 +2193,35 @@ async def briefing(ctx: Context) -> list[UserMessage] | InputRequiredResult: |
2151 | 2193 | assert block.text == "Brief Alice (state=r1)" |
2152 | 2194 |
|
2153 | 2195 |
|
| 2196 | +async def test_prompt_with_parameterized_context_reads_input_responses_on_retry(): |
| 2197 | + """A prompt annotated `ctx: Context[T]` sees the retry's input_responses and request_state, so the |
| 2198 | + multi-round-trip flow completes instead of asking the same question again.""" |
| 2199 | + mcp = MCPServer() |
| 2200 | + |
| 2201 | + @mcp.prompt() |
| 2202 | + async def briefing(ctx: Context[dict[str, Any]]) -> list[UserMessage] | InputRequiredResult: |
| 2203 | + responses = ctx.input_responses |
| 2204 | + if responses and "who" in responses: |
| 2205 | + who = responses["who"] |
| 2206 | + assert isinstance(who, ElicitResult) and who.content is not None |
| 2207 | + return [UserMessage(content=f"Brief {who.content['name']} (state={ctx.request_state})")] |
| 2208 | + return InputRequiredResult(input_requests={"who": _ask_who()}, request_state="r1") |
| 2209 | + |
| 2210 | + with anyio.fail_after(5): |
| 2211 | + async with Client(mcp, mode="2026-07-28") as client: |
| 2212 | + r1 = await client.session.get_prompt("briefing", allow_input_required=True) |
| 2213 | + assert isinstance(r1, InputRequiredResult) |
| 2214 | + |
| 2215 | + r2 = await client.session.get_prompt( |
| 2216 | + "briefing", |
| 2217 | + input_responses={"who": ElicitResult(action="accept", content={"name": "Alice"})}, |
| 2218 | + request_state=r1.request_state, |
| 2219 | + allow_input_required=True, |
| 2220 | + ) |
| 2221 | + assert isinstance(r2, GetPromptResult) |
| 2222 | + assert r2.messages == [PromptMessage(role="user", content=TextContent(type="text", text="Brief Alice (state=r1)"))] |
| 2223 | + |
| 2224 | + |
2154 | 2225 | async def test_prompt_input_required_result_on_legacy_session_is_a_serialization_error(): |
2155 | 2226 | """Pins the shared era gate: a pre-2026 session has no input_required vocabulary, so |
2156 | 2227 | the runner rejects the frame with -32603 — the same posture the tools path has.""" |
@@ -3114,6 +3185,11 @@ def test_context_mcp_server_outside_request_raises() -> None: |
3114 | 3185 | _ = Context().mcp_server |
3115 | 3186 |
|
3116 | 3187 |
|
| 3188 | +def test_context_request_context_outside_request_raises() -> None: |
| 3189 | + with pytest.raises(ValueError, match="outside of a request"): |
| 3190 | + _ = Context().request_context |
| 3191 | + |
| 3192 | + |
3117 | 3193 | async def test_context_notify_outside_a_request_raises() -> None: |
3118 | 3194 | with pytest.raises(ValueError, match="outside of a request"): |
3119 | 3195 | await Context().notify_tools_changed() |
|
0 commit comments