392 lines
17 KiB
Python
392 lines
17 KiB
Python
"""Test the elicitation feature using stdio transport."""
|
|
|
|
from typing import Any
|
|
|
|
import pytest
|
|
from pydantic import BaseModel, Field
|
|
|
|
from mcp import Client, types
|
|
from mcp.client.session import ClientSession, ElicitationFnT
|
|
from mcp.server.mcpserver import Context, MCPServer
|
|
from mcp.shared._context import RequestContext
|
|
from mcp.types import ElicitRequestParams, ElicitResult, TextContent
|
|
|
|
|
|
# Shared schema for basic tests
|
|
class AnswerSchema(BaseModel):
|
|
answer: str = Field(description="The user's answer to the question")
|
|
|
|
|
|
def create_ask_user_tool(mcp: MCPServer):
|
|
"""Create a standard ask_user tool that handles all elicitation responses."""
|
|
|
|
@mcp.tool(description="A tool that uses elicitation")
|
|
async def ask_user(prompt: str, ctx: Context) -> str:
|
|
result = await ctx.elicit(message=f"Tool wants to ask: {prompt}", schema=AnswerSchema)
|
|
|
|
if result.action == "accept" and result.data:
|
|
return f"User answered: {result.data.answer}"
|
|
elif result.action == "decline":
|
|
return "User declined to answer"
|
|
else: # pragma: no cover
|
|
return "User cancelled"
|
|
|
|
return ask_user
|
|
|
|
|
|
async def call_tool_and_assert(
|
|
mcp: MCPServer,
|
|
elicitation_callback: ElicitationFnT,
|
|
tool_name: str,
|
|
args: dict[str, Any],
|
|
expected_text: str | None = None,
|
|
text_contains: list[str] | None = None,
|
|
):
|
|
"""Helper to create session, call tool, and assert result."""
|
|
async with Client(mcp, elicitation_callback=elicitation_callback) as client:
|
|
result = await client.call_tool(tool_name, args)
|
|
assert len(result.content) == 1
|
|
assert isinstance(result.content[0], TextContent)
|
|
|
|
if expected_text is not None:
|
|
assert result.content[0].text == expected_text
|
|
elif text_contains is not None: # pragma: no branch
|
|
for substring in text_contains:
|
|
assert substring in result.content[0].text
|
|
|
|
return result
|
|
|
|
|
|
@pytest.mark.anyio
|
|
async def test_stdio_elicitation():
|
|
"""Test the elicitation feature using stdio transport."""
|
|
mcp = MCPServer(name="StdioElicitationServer")
|
|
create_ask_user_tool(mcp)
|
|
|
|
# Create a custom handler for elicitation requests
|
|
async def elicitation_callback(context: RequestContext[ClientSession], params: ElicitRequestParams):
|
|
if params.message == "Tool wants to ask: What is your name?":
|
|
return ElicitResult(action="accept", content={"answer": "Test User"})
|
|
else: # pragma: no cover
|
|
raise ValueError(f"Unexpected elicitation message: {params.message}")
|
|
|
|
await call_tool_and_assert(
|
|
mcp, elicitation_callback, "ask_user", {"prompt": "What is your name?"}, "User answered: Test User"
|
|
)
|
|
|
|
|
|
@pytest.mark.anyio
|
|
async def test_stdio_elicitation_decline():
|
|
"""Test elicitation with user declining."""
|
|
mcp = MCPServer(name="StdioElicitationDeclineServer")
|
|
create_ask_user_tool(mcp)
|
|
|
|
async def elicitation_callback(context: RequestContext[ClientSession], params: ElicitRequestParams):
|
|
return ElicitResult(action="decline")
|
|
|
|
await call_tool_and_assert(
|
|
mcp, elicitation_callback, "ask_user", {"prompt": "What is your name?"}, "User declined to answer"
|
|
)
|
|
|
|
|
|
@pytest.mark.anyio
|
|
async def test_elicitation_schema_validation():
|
|
"""Test that elicitation schemas must only contain primitive types."""
|
|
mcp = MCPServer(name="ValidationTestServer")
|
|
|
|
def create_validation_tool(name: str, schema_class: type[BaseModel]):
|
|
@mcp.tool(name=name, description=f"Tool testing {name}")
|
|
async def tool(ctx: Context) -> str:
|
|
try:
|
|
await ctx.elicit(message="This should fail validation", schema=schema_class)
|
|
return "Should not reach here" # pragma: no cover
|
|
except TypeError as e:
|
|
return f"Validation failed as expected: {str(e)}"
|
|
|
|
return tool
|
|
|
|
# Test cases for invalid schemas
|
|
class InvalidListSchema(BaseModel):
|
|
numbers: list[int] = Field(description="List of numbers")
|
|
|
|
class NestedModel(BaseModel):
|
|
value: str
|
|
|
|
class InvalidNestedSchema(BaseModel):
|
|
nested: NestedModel = Field(description="Nested model")
|
|
|
|
create_validation_tool("invalid_list", InvalidListSchema)
|
|
create_validation_tool("nested_model", InvalidNestedSchema)
|
|
|
|
# Dummy callback (won't be called due to validation failure)
|
|
async def elicitation_callback(
|
|
context: RequestContext[ClientSession], params: ElicitRequestParams
|
|
): # pragma: no cover
|
|
return ElicitResult(action="accept", content={})
|
|
|
|
async with Client(mcp, elicitation_callback=elicitation_callback) as client:
|
|
# Test both invalid schemas
|
|
for tool_name, field_name in [("invalid_list", "numbers"), ("nested_model", "nested")]:
|
|
result = await client.call_tool(tool_name, {})
|
|
assert len(result.content) == 1
|
|
assert isinstance(result.content[0], TextContent)
|
|
assert "Validation failed as expected" in result.content[0].text
|
|
assert field_name in result.content[0].text
|
|
|
|
|
|
@pytest.mark.anyio
|
|
async def test_elicitation_with_optional_fields():
|
|
"""Test that Optional fields work correctly in elicitation schemas."""
|
|
mcp = MCPServer(name="OptionalFieldServer")
|
|
|
|
class OptionalSchema(BaseModel):
|
|
required_name: str = Field(description="Your name (required)")
|
|
optional_age: int | None = Field(default=None, description="Your age (optional)")
|
|
optional_email: str | None = Field(default=None, description="Your email (optional)")
|
|
subscribe: bool | None = Field(default=False, description="Subscribe to newsletter?")
|
|
|
|
@mcp.tool(description="Tool with optional fields")
|
|
async def optional_tool(ctx: Context) -> str:
|
|
result = await ctx.elicit(message="Please provide your information", schema=OptionalSchema)
|
|
|
|
if result.action == "accept" and result.data:
|
|
info = [f"Name: {result.data.required_name}"]
|
|
if result.data.optional_age is not None:
|
|
info.append(f"Age: {result.data.optional_age}")
|
|
if result.data.optional_email is not None:
|
|
info.append(f"Email: {result.data.optional_email}")
|
|
info.append(f"Subscribe: {result.data.subscribe}")
|
|
return ", ".join(info)
|
|
else: # pragma: no cover
|
|
return f"User {result.action}"
|
|
|
|
# Test cases with different field combinations
|
|
test_cases: list[tuple[dict[str, Any], str]] = [
|
|
(
|
|
# All fields provided
|
|
{"required_name": "John Doe", "optional_age": 30, "optional_email": "john@example.com", "subscribe": True},
|
|
"Name: John Doe, Age: 30, Email: john@example.com, Subscribe: True",
|
|
),
|
|
(
|
|
# Only required fields
|
|
{"required_name": "Jane Smith"},
|
|
"Name: Jane Smith, Subscribe: False",
|
|
),
|
|
]
|
|
|
|
for content, expected in test_cases:
|
|
|
|
async def callback(context: RequestContext[ClientSession], params: ElicitRequestParams):
|
|
return ElicitResult(action="accept", content=content)
|
|
|
|
await call_tool_and_assert(mcp, callback, "optional_tool", {}, expected)
|
|
|
|
# Test invalid optional field
|
|
class InvalidOptionalSchema(BaseModel):
|
|
name: str = Field(description="Name")
|
|
optional_list: list[int] | None = Field(default=None, description="Invalid optional list")
|
|
|
|
@mcp.tool(description="Tool with invalid optional field")
|
|
async def invalid_optional_tool(ctx: Context) -> str:
|
|
try:
|
|
await ctx.elicit(message="This should fail", schema=InvalidOptionalSchema)
|
|
return "Should not reach here" # pragma: no cover
|
|
except TypeError as e:
|
|
return f"Validation failed: {str(e)}"
|
|
|
|
async def elicitation_callback(
|
|
context: RequestContext[ClientSession], params: ElicitRequestParams
|
|
): # pragma: no cover
|
|
return ElicitResult(action="accept", content={})
|
|
|
|
await call_tool_and_assert(
|
|
mcp,
|
|
elicitation_callback,
|
|
"invalid_optional_tool",
|
|
{},
|
|
text_contains=["Validation failed:", "optional_list"],
|
|
)
|
|
|
|
# Test valid list[str] for multi-select enum
|
|
class ValidMultiSelectSchema(BaseModel):
|
|
name: str = Field(description="Name")
|
|
tags: list[str] = Field(description="Tags")
|
|
|
|
@mcp.tool(description="Tool with valid list[str] field")
|
|
async def valid_multiselect_tool(ctx: Context) -> str:
|
|
result = await ctx.elicit(message="Please provide tags", schema=ValidMultiSelectSchema)
|
|
if result.action == "accept" and result.data:
|
|
return f"Name: {result.data.name}, Tags: {', '.join(result.data.tags)}"
|
|
return f"User {result.action}" # pragma: no cover
|
|
|
|
async def multiselect_callback(context: RequestContext[ClientSession], params: ElicitRequestParams):
|
|
if "Please provide tags" in params.message:
|
|
return ElicitResult(action="accept", content={"name": "Test", "tags": ["tag1", "tag2"]})
|
|
return ElicitResult(action="decline") # pragma: no cover
|
|
|
|
await call_tool_and_assert(mcp, multiselect_callback, "valid_multiselect_tool", {}, "Name: Test, Tags: tag1, tag2")
|
|
|
|
# Test Optional[list[str]] for optional multi-select enum
|
|
class OptionalMultiSelectSchema(BaseModel):
|
|
name: str = Field(description="Name")
|
|
tags: list[str] | None = Field(default=None, description="Optional tags")
|
|
|
|
@mcp.tool(description="Tool with optional list[str] field")
|
|
async def optional_multiselect_tool(ctx: Context) -> str:
|
|
result = await ctx.elicit(message="Please provide optional tags", schema=OptionalMultiSelectSchema)
|
|
if result.action == "accept" and result.data:
|
|
tags_str = ", ".join(result.data.tags) if result.data.tags else "none"
|
|
return f"Name: {result.data.name}, Tags: {tags_str}"
|
|
return f"User {result.action}" # pragma: no cover
|
|
|
|
async def optional_multiselect_callback(context: RequestContext[ClientSession], params: ElicitRequestParams):
|
|
if "Please provide optional tags" in params.message:
|
|
return ElicitResult(action="accept", content={"name": "Test", "tags": ["tag1", "tag2"]})
|
|
return ElicitResult(action="decline") # pragma: no cover
|
|
|
|
await call_tool_and_assert(
|
|
mcp, optional_multiselect_callback, "optional_multiselect_tool", {}, "Name: Test, Tags: tag1, tag2"
|
|
)
|
|
|
|
|
|
@pytest.mark.anyio
|
|
async def test_elicitation_with_default_values():
|
|
"""Test that default values work correctly in elicitation schemas and are included in JSON."""
|
|
mcp = MCPServer(name="DefaultValuesServer")
|
|
|
|
class DefaultsSchema(BaseModel):
|
|
name: str = Field(default="Guest", description="User name")
|
|
age: int = Field(default=18, description="User age")
|
|
subscribe: bool = Field(default=True, description="Subscribe to newsletter")
|
|
email: str = Field(description="Email address (required)")
|
|
|
|
@mcp.tool(description="Tool with default values")
|
|
async def defaults_tool(ctx: Context) -> str:
|
|
result = await ctx.elicit(message="Please provide your information", schema=DefaultsSchema)
|
|
|
|
if result.action == "accept" and result.data:
|
|
return (
|
|
f"Name: {result.data.name}, Age: {result.data.age}, "
|
|
f"Subscribe: {result.data.subscribe}, Email: {result.data.email}"
|
|
)
|
|
else: # pragma: no cover
|
|
return f"User {result.action}"
|
|
|
|
# First verify that defaults are present in the JSON schema sent to clients
|
|
async def callback_schema_verify(context: RequestContext[ClientSession], params: ElicitRequestParams):
|
|
# Verify the schema includes defaults
|
|
assert isinstance(params, types.ElicitRequestFormParams), "Expected form mode elicitation"
|
|
schema = params.requested_schema
|
|
props = schema["properties"]
|
|
|
|
assert props["name"]["default"] == "Guest"
|
|
assert props["age"]["default"] == 18
|
|
assert props["subscribe"]["default"] is True
|
|
assert "default" not in props["email"] # Required field has no default
|
|
|
|
return ElicitResult(action="accept", content={"email": "test@example.com"})
|
|
|
|
await call_tool_and_assert(
|
|
mcp,
|
|
callback_schema_verify,
|
|
"defaults_tool",
|
|
{},
|
|
"Name: Guest, Age: 18, Subscribe: True, Email: test@example.com",
|
|
)
|
|
|
|
# Test overriding defaults
|
|
async def callback_override(context: RequestContext[ClientSession], params: ElicitRequestParams):
|
|
return ElicitResult(
|
|
action="accept", content={"email": "john@example.com", "name": "John", "age": 25, "subscribe": False}
|
|
)
|
|
|
|
await call_tool_and_assert(
|
|
mcp, callback_override, "defaults_tool", {}, "Name: John, Age: 25, Subscribe: False, Email: john@example.com"
|
|
)
|
|
|
|
|
|
@pytest.mark.anyio
|
|
async def test_elicitation_with_enum_titles():
|
|
"""Test elicitation with enum schemas using oneOf/anyOf for titles."""
|
|
mcp = MCPServer(name="ColorPreferencesApp")
|
|
|
|
# Test single-select with titles using oneOf
|
|
class FavoriteColorSchema(BaseModel):
|
|
user_name: str = Field(description="Your name")
|
|
favorite_color: str = Field(
|
|
description="Select your favorite color",
|
|
json_schema_extra={
|
|
"oneOf": [
|
|
{"const": "red", "title": "Red"},
|
|
{"const": "green", "title": "Green"},
|
|
{"const": "blue", "title": "Blue"},
|
|
{"const": "yellow", "title": "Yellow"},
|
|
]
|
|
},
|
|
)
|
|
|
|
@mcp.tool(description="Single color selection")
|
|
async def select_favorite_color(ctx: Context) -> str:
|
|
result = await ctx.elicit(message="Select your favorite color", schema=FavoriteColorSchema)
|
|
if result.action == "accept" and result.data:
|
|
return f"User: {result.data.user_name}, Favorite: {result.data.favorite_color}"
|
|
return f"User {result.action}" # pragma: no cover
|
|
|
|
# Test multi-select with titles using anyOf
|
|
class FavoriteColorsSchema(BaseModel):
|
|
user_name: str = Field(description="Your name")
|
|
favorite_colors: list[str] = Field(
|
|
description="Select your favorite colors",
|
|
json_schema_extra={
|
|
"items": {
|
|
"anyOf": [
|
|
{"const": "red", "title": "Red"},
|
|
{"const": "green", "title": "Green"},
|
|
{"const": "blue", "title": "Blue"},
|
|
{"const": "yellow", "title": "Yellow"},
|
|
]
|
|
}
|
|
},
|
|
)
|
|
|
|
@mcp.tool(description="Multiple color selection")
|
|
async def select_favorite_colors(ctx: Context) -> str:
|
|
result = await ctx.elicit(message="Select your favorite colors", schema=FavoriteColorsSchema)
|
|
if result.action == "accept" and result.data:
|
|
return f"User: {result.data.user_name}, Colors: {', '.join(result.data.favorite_colors)}"
|
|
return f"User {result.action}" # pragma: no cover
|
|
|
|
# Test legacy enumNames format
|
|
class LegacyColorSchema(BaseModel):
|
|
user_name: str = Field(description="Your name")
|
|
color: str = Field(
|
|
description="Select a color",
|
|
json_schema_extra={"enum": ["red", "green", "blue"], "enumNames": ["Red", "Green", "Blue"]},
|
|
)
|
|
|
|
@mcp.tool(description="Legacy enum format")
|
|
async def select_color_legacy(ctx: Context) -> str:
|
|
result = await ctx.elicit(message="Select a color (legacy format)", schema=LegacyColorSchema)
|
|
if result.action == "accept" and result.data:
|
|
return f"User: {result.data.user_name}, Color: {result.data.color}"
|
|
return f"User {result.action}" # pragma: no cover
|
|
|
|
async def enum_callback(context: RequestContext[ClientSession], params: ElicitRequestParams):
|
|
if "colors" in params.message and "legacy" not in params.message:
|
|
return ElicitResult(action="accept", content={"user_name": "Bob", "favorite_colors": ["red", "green"]})
|
|
elif "color" in params.message:
|
|
if "legacy" in params.message:
|
|
return ElicitResult(action="accept", content={"user_name": "Charlie", "color": "green"})
|
|
else:
|
|
return ElicitResult(action="accept", content={"user_name": "Alice", "favorite_color": "blue"})
|
|
return ElicitResult(action="decline") # pragma: no cover
|
|
|
|
# Test single-select with titles
|
|
await call_tool_and_assert(mcp, enum_callback, "select_favorite_color", {}, "User: Alice, Favorite: blue")
|
|
|
|
# Test multi-select with titles
|
|
await call_tool_and_assert(mcp, enum_callback, "select_favorite_colors", {}, "User: Bob, Colors: red, green")
|
|
|
|
# Test legacy enumNames format
|
|
await call_tool_and_assert(mcp, enum_callback, "select_color_legacy", {}, "User: Charlie, Color: green")
|