mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-22 05:24:18 +02:00
Support constrained choice
This commit is contained in:
parent
3805ac9b44
commit
a34fb9e3bb
2 changed files with 236 additions and 97 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue