diff --git a/src/fastmcp/server/context.py b/src/fastmcp/server/context.py index a3302c818..a37e8c927 100644 --- a/src/fastmcp/server/context.py +++ b/src/fastmcp/server/context.py @@ -6,7 +6,8 @@ from collections.abc import Generator from contextlib import contextmanager from contextvars import ContextVar, Token from dataclasses import dataclass -from typing import TypeVar, cast +from enum import Enum +from typing import Literal, TypeVar, cast, get_origin from mcp import LoggingLevel, ServerSession from mcp.server.lowlevel.helper_types import ReadResourceContents @@ -314,7 +315,7 @@ class Context: async def elicit( self, message: str, - response_type: type[T] | None = None, + response_type: type[T] | list[str] | None = None, ) -> AcceptedElicitation[T] | DeclinedElicitation | CancelledElicitation: """ Send an elicitation request to the client and await the response. @@ -338,10 +339,29 @@ class Context: if response_type is None: response_type = str # type: ignore - if response_type in {bool, int, float, str}: + # if the user provided a list of strings, treat it as a Literal + if isinstance(response_type, list): + if not all(isinstance(item, str) for item in response_type): + raise ValueError( + "List of options must be a list of strings. Received: " + f"{response_type}" + ) + # Convert list of options to Literal type and wrap + choice_literal = Literal[*tuple(response_type)] # type: ignore + response_type = ScalarElicitationType[choice_literal] # type: ignore + # if the user provided a primitive scalar, wrap it in an object schema + elif response_type in {bool, int, float, str}: + response_type = ScalarElicitationType[response_type] # type: ignore + # if the user provided a Literal type, wrap it in an object schema + elif get_origin(response_type) is Literal: + response_type = ScalarElicitationType[response_type] # type: ignore + # if the user provided an Enum type, wrap it in an object schema + elif isinstance(response_type, type) and issubclass(response_type, Enum): response_type = ScalarElicitationType[response_type] # type: ignore - requested_schema = get_elicitation_schema(response_type) # type: ignore + response_type = cast(type[T], response_type) + + requested_schema = get_elicitation_schema(response_type) result = await self.session.elicit( message=message, diff --git a/tests/client/test_elicitation.py b/tests/client/test_elicitation.py index d79cfbdca..70998625f 100644 --- a/tests/client/test_elicitation.py +++ b/tests/client/test_elicitation.py @@ -3,6 +3,8 @@ from enum import Enum from typing import Literal import pytest +from pydantic import BaseModel +from typing_extensions import TypedDict from fastmcp import Context, FastMCP from fastmcp.client.client import Client @@ -159,23 +161,122 @@ async def test_elicitation_cancel_action(): assert result.data == "Request was canceled" -async def test_elicitation_number_schema(): - """Test elicitation with number schema.""" - mcp = FastMCP("TestServer") +class TestScalarResponseTypes: + async def test_elicitation_str_response(self): + """Test elicitation with string schema.""" + mcp = FastMCP("TestServer") - @mcp.tool - async def get_age(context: Context) -> str: - result = await context.elicit(message="How old are you?", response_type=int) - if result.action == "accept": - return f"You are {result.data} years old" - return "No age provided" + @mcp.tool + async def my_tool(context: Context) -> str: + result = await context.elicit(message="", response_type=str) + return result.data # type: ignore[attr-defined] - async def elicitation_handler(message, response_type, params, ctx): - return ElicitResult(action="accept", content=response_type(value=25)) + async def elicitation_handler(message, response_type, params, ctx): + return ElicitResult(action="accept", content={"value": "hello"}) - async with Client(mcp, elicitation_handler=elicitation_handler) as client: - result = await client.call_tool("get_age", {}) - assert result.data == "You are 25 years old" + async with Client(mcp, elicitation_handler=elicitation_handler) as client: + result = await client.call_tool("my_tool", {}) + assert result.data == "hello" + + async def test_elicitation_int_response(self): + """Test elicitation with number schema.""" + mcp = FastMCP("TestServer") + + @mcp.tool + async def my_tool(context: Context) -> int: + result = await context.elicit(message="", response_type=int) + return result.data # type: ignore[attr-defined] + + async def elicitation_handler(message, response_type, params, ctx): + return ElicitResult(action="accept", content={"value": 42}) + + async with Client(mcp, elicitation_handler=elicitation_handler) as client: + result = await client.call_tool("my_tool", {}) + assert result.data == 42 + + async def test_elicitation_float_response(self): + """Test elicitation with number schema.""" + mcp = FastMCP("TestServer") + + @mcp.tool + async def my_tool(context: Context) -> float: + result = await context.elicit(message="", response_type=float) + return result.data # type: ignore[attr-defined] + + async def elicitation_handler(message, response_type, params, ctx): + return ElicitResult(action="accept", content={"value": 3.14}) + + async with Client(mcp, elicitation_handler=elicitation_handler) as client: + result = await client.call_tool("my_tool", {}) + assert result.data == 3.14 + + async def test_elicitation_bool_response(self): + """Test elicitation with boolean schema.""" + mcp = FastMCP("TestServer") + + @mcp.tool + async def my_tool(context: Context) -> bool: + result = await context.elicit(message="", response_type=bool) + return result.data # type: ignore[attr-defined] + + async def elicitation_handler(message, response_type, params, ctx): + return ElicitResult(action="accept", content={"value": True}) + + async with Client(mcp, elicitation_handler=elicitation_handler) as client: + result = await client.call_tool("my_tool", {}) + assert result.data is True + + async def test_elicitation_literal_response(self): + """Test elicitation with literal schema.""" + mcp = FastMCP("TestServer") + + @mcp.tool + async def my_tool(context: Context) -> Literal["x", "y"]: + result = await context.elicit(message="", response_type=Literal["x", "y"]) # type: ignore + return result.data # type: ignore[attr-defined] + + async def elicitation_handler(message, response_type, params, ctx): + return ElicitResult(action="accept", content={"value": "x"}) + + async with Client(mcp, elicitation_handler=elicitation_handler) as client: + result = await client.call_tool("my_tool", {}) + assert result.data == "x" + + async def test_elicitation_enum_response(self): + """Test elicitation with enum schema.""" + mcp = FastMCP("TestServer") + + class ResponseEnum(Enum): + X = "x" + Y = "y" + + @mcp.tool + async def my_tool(context: Context) -> ResponseEnum: + result = await context.elicit(message="", response_type=ResponseEnum) + return result.data # type: ignore[attr-defined] + + async def elicitation_handler(message, response_type, params, ctx): + return ElicitResult(action="accept", content={"value": "x"}) + + async with Client(mcp, elicitation_handler=elicitation_handler) as client: + result = await client.call_tool("my_tool", {}) + assert result.data == "x" + + async def test_elicitation_list_response(self): + """Test elicitation with list schema.""" + mcp = FastMCP("TestServer") + + @mcp.tool + async def my_tool(context: Context) -> str: + result = await context.elicit(message="", response_type=["x", "y"]) + return result.data # type: ignore[attr-defined] + + async def elicitation_handler(message, response_type, params, ctx): + return ElicitResult(action="accept", content={"value": "x"}) + + async with Client(mcp, elicitation_handler=elicitation_handler) as client: + result = await client.call_tool("my_tool", {}) + assert result.data == "x" async def test_elicitation_handler_error(): @@ -237,29 +338,48 @@ async def test_elicitation_multiple_calls(): assert call_count == 2 -async def test_dataclass_response_type(): +@dataclass +class UserInfo: + name: str + age: int + + +class UserInfoTypedDict(TypedDict): + name: str + age: int + + +class UserInfoPydantic(BaseModel): + name: str + age: int + + +@pytest.mark.parametrize( + "structured_type", [UserInfo, UserInfoTypedDict, UserInfoPydantic] +) +async def test_structured_response_type( + structured_type: type[UserInfo | UserInfoTypedDict | UserInfoPydantic], +): """Test elicitation with dataclass response type.""" mcp = FastMCP("TestServer") - @dataclass - class UserInfo: - name: str - age: int - @mcp.tool async def get_user_info(context: Context) -> str: result = await context.elicit( - message="Please provide your information", response_type=UserInfo + message="Please provide your information", response_type=structured_type ) if result.action == "accept": - return f"User: {result.data.name}, age: {result.data.age}" + if isinstance(result.data, dict): + return f"User: {result.data['name']}, age: {result.data['age']}" + else: + return f"User: {result.data.name}, age: {result.data.age}" return "No user info provided" async def elicitation_handler(message, response_type, params, ctx): # Verify we get the dataclass type assert ( TypeAdapter(response_type).json_schema() - == TypeAdapter(UserInfo).json_schema() + == TypeAdapter(structured_type).json_schema() ) # Verify the schema has the dataclass fields (available in params) @@ -387,73 +507,72 @@ class TestValidation: ) -async def test_pattern_matching_accept(): - """Test pattern matching with AcceptedElicitation.""" - mcp = FastMCP("TestServer") +class TestPatternMatching: + async def test_pattern_matching_accept(self): + """Test pattern matching with AcceptedElicitation.""" + mcp = FastMCP("TestServer") - @mcp.tool - async def pattern_match_tool(context: Context) -> str: - result = await context.elicit("Enter your name:", response_type=str) + @mcp.tool + async def pattern_match_tool(context: Context) -> str: + result = await context.elicit("Enter your name:", response_type=str) - match result: - case AcceptedElicitation(data=name): - return f"Hello {name}!" - case DeclinedElicitation(): - return "You declined" - case CancelledElicitation(): - return "Cancelled" + match result: + case AcceptedElicitation(data=name): + return f"Hello {name}!" + case DeclinedElicitation(): + return "You declined" + case CancelledElicitation(): + return "Cancelled" - async def elicitation_handler(message, response_type, params, ctx): - return ElicitResult(action="accept", content={"value": "Alice"}) + async def elicitation_handler(message, response_type, params, ctx): + return ElicitResult(action="accept", content={"value": "Alice"}) - async with Client(mcp, elicitation_handler=elicitation_handler) as client: - result = await client.call_tool("pattern_match_tool", {}) - assert result.data == "Hello Alice!" + async with Client(mcp, elicitation_handler=elicitation_handler) as client: + result = await client.call_tool("pattern_match_tool", {}) + assert result.data == "Hello Alice!" + async def test_pattern_matching_decline(self): + """Test pattern matching with DeclinedElicitation.""" + mcp = FastMCP("TestServer") -async def test_pattern_matching_decline(): - """Test pattern matching with DeclinedElicitation.""" - mcp = FastMCP("TestServer") + @mcp.tool + async def pattern_match_tool(context: Context) -> str: + result = await context.elicit("Enter your name:", response_type=str) - @mcp.tool - async def pattern_match_tool(context: Context) -> str: - result = await context.elicit("Enter your name:", response_type=str) + match result: + case AcceptedElicitation(data=name): + return f"Hello {name}!" + case DeclinedElicitation(): + return "You declined" + case CancelledElicitation(): + return "Cancelled" - match result: - case AcceptedElicitation(data=name): - return f"Hello {name}!" - case DeclinedElicitation(): - return "You declined" - case CancelledElicitation(): - return "Cancelled" + async def elicitation_handler(message, response_type, params, ctx): + return ElicitResult(action="decline") - async def elicitation_handler(message, response_type, params, ctx): - return ElicitResult(action="decline") + async with Client(mcp, elicitation_handler=elicitation_handler) as client: + result = await client.call_tool("pattern_match_tool", {}) + assert result.data == "You declined" - async with Client(mcp, elicitation_handler=elicitation_handler) as client: - result = await client.call_tool("pattern_match_tool", {}) - assert result.data == "You declined" + async def test_pattern_matching_cancel(self): + """Test pattern matching with CancelledElicitation.""" + mcp = FastMCP("TestServer") + @mcp.tool + async def pattern_match_tool(context: Context) -> str: + result = await context.elicit("Enter your name:", response_type=str) -async def test_pattern_matching_cancel(): - """Test pattern matching with CancelledElicitation.""" - mcp = FastMCP("TestServer") + match result: + case AcceptedElicitation(data=name): + return f"Hello {name}!" + case DeclinedElicitation(): + return "You declined" + case CancelledElicitation(): + return "Cancelled" - @mcp.tool - async def pattern_match_tool(context: Context) -> str: - result = await context.elicit("Enter your name:", response_type=str) + async def elicitation_handler(message, response_type, params, ctx): + return ElicitResult(action="cancel") - match result: - case AcceptedElicitation(data=name): - return f"Hello {name}!" - case DeclinedElicitation(): - return "You declined" - case CancelledElicitation(): - return "Cancelled" - - async def elicitation_handler(message, response_type, params, ctx): - return ElicitResult(action="cancel") - - async with Client(mcp, elicitation_handler=elicitation_handler) as client: - result = await client.call_tool("pattern_match_tool", {}) - assert result.data == "Cancelled" + async with Client(mcp, elicitation_handler=elicitation_handler) as client: + result = await client.call_tool("pattern_match_tool", {}) + assert result.data == "Cancelled"