mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-20 20:44:17 +02:00
436 lines
14 KiB
Python
436 lines
14 KiB
Python
from typing import Annotated
|
|
|
|
import pytest
|
|
|
|
from fastmcp import Context
|
|
from fastmcp.exceptions import NotFoundError, PromptError
|
|
from fastmcp.prompts import Prompt
|
|
from fastmcp.prompts.prompt import FunctionPrompt, PromptMessage, TextContent
|
|
from fastmcp.prompts.prompt_manager import PromptManager
|
|
from fastmcp.utilities.tests import caplog_for_fastmcp
|
|
from tests.conftest import get_fn_name
|
|
|
|
|
|
class TestPromptManager:
|
|
async def test_add_prompt(self):
|
|
"""Test adding a prompt to the manager."""
|
|
|
|
def fn() -> str:
|
|
return "Hello, world!"
|
|
|
|
manager = PromptManager()
|
|
prompt = Prompt.from_function(fn)
|
|
added = manager.add_prompt(prompt)
|
|
assert added == prompt
|
|
assert await manager.get_prompt("fn") == prompt
|
|
|
|
async def test_add_duplicate_prompt(self, caplog):
|
|
"""Test adding the same prompt twice."""
|
|
|
|
def fn() -> str:
|
|
return "Hello, world!"
|
|
|
|
manager = PromptManager(duplicate_behavior="warn")
|
|
prompt = Prompt.from_function(fn)
|
|
first = manager.add_prompt(prompt)
|
|
|
|
with caplog_for_fastmcp(caplog):
|
|
second = manager.add_prompt(prompt)
|
|
|
|
assert first == second
|
|
assert "Prompt already exists" in caplog.text
|
|
|
|
async def test_disable_warn_on_duplicate_prompts(self, caplog):
|
|
"""Test disabling warning on duplicate prompts."""
|
|
|
|
def fn() -> str:
|
|
return "Hello, world!"
|
|
|
|
manager = PromptManager(duplicate_behavior="ignore")
|
|
prompt = Prompt.from_function(fn)
|
|
first = manager.add_prompt(prompt)
|
|
second = manager.add_prompt(prompt)
|
|
assert first == second
|
|
assert "Prompt already exists" not in caplog.text
|
|
|
|
async def test_warn_on_duplicate_prompts(self, caplog):
|
|
"""Test warning on duplicate prompts."""
|
|
manager = PromptManager(duplicate_behavior="warn")
|
|
|
|
def test_fn() -> str:
|
|
return "Test prompt"
|
|
|
|
prompt = Prompt.from_function(test_fn, name="test_prompt")
|
|
|
|
manager.add_prompt(prompt)
|
|
|
|
with caplog_for_fastmcp(caplog):
|
|
manager.add_prompt(prompt)
|
|
|
|
assert "Prompt already exists: test_prompt" in caplog.text
|
|
# Should have the prompt
|
|
assert await manager.get_prompt("test_prompt") is not None
|
|
|
|
async def test_error_on_duplicate_prompts(self):
|
|
"""Test error on duplicate prompts."""
|
|
manager = PromptManager(duplicate_behavior="error")
|
|
|
|
def test_fn() -> str:
|
|
return "Test prompt"
|
|
|
|
prompt = Prompt.from_function(test_fn, name="test_prompt")
|
|
|
|
manager.add_prompt(prompt)
|
|
|
|
with pytest.raises(ValueError, match="Prompt already exists: test_prompt"):
|
|
manager.add_prompt(prompt)
|
|
|
|
async def test_replace_duplicate_prompts(self):
|
|
"""Test replacing duplicate prompts."""
|
|
manager = PromptManager(duplicate_behavior="replace")
|
|
|
|
def original_fn() -> str:
|
|
return "Original prompt"
|
|
|
|
def replacement_fn() -> str:
|
|
return "Replacement prompt"
|
|
|
|
prompt1 = Prompt.from_function(original_fn, name="test_prompt")
|
|
prompt2 = Prompt.from_function(replacement_fn, name="test_prompt")
|
|
|
|
manager.add_prompt(prompt1)
|
|
manager.add_prompt(prompt2)
|
|
|
|
# Should have replaced with the new prompt
|
|
prompt = await manager.get_prompt("test_prompt")
|
|
assert prompt is not None
|
|
assert isinstance(prompt, FunctionPrompt)
|
|
assert get_fn_name(prompt.fn) == "replacement_fn"
|
|
|
|
async def test_ignore_duplicate_prompts(self):
|
|
"""Test ignoring duplicate prompts."""
|
|
manager = PromptManager(duplicate_behavior="ignore")
|
|
|
|
def original_fn() -> str:
|
|
return "Original prompt"
|
|
|
|
def replacement_fn() -> str:
|
|
return "Replacement prompt"
|
|
|
|
prompt1 = Prompt.from_function(original_fn, name="test_prompt")
|
|
prompt2 = Prompt.from_function(replacement_fn, name="test_prompt")
|
|
|
|
manager.add_prompt(prompt1)
|
|
result = manager.add_prompt(prompt2)
|
|
|
|
# Should keep the original
|
|
prompt = await manager.get_prompt("test_prompt")
|
|
assert prompt is not None
|
|
assert isinstance(prompt, FunctionPrompt)
|
|
assert get_fn_name(prompt.fn) == "original_fn"
|
|
# Result should be the original prompt
|
|
assert isinstance(result, FunctionPrompt)
|
|
assert get_fn_name(result.fn) == "original_fn"
|
|
|
|
async def test_get_prompts(self):
|
|
"""Test retrieving all prompts."""
|
|
|
|
def fn1() -> str:
|
|
return "Hello, world!"
|
|
|
|
def fn2() -> str:
|
|
return "Goodbye, world!"
|
|
|
|
manager = PromptManager()
|
|
prompt1 = Prompt.from_function(fn1)
|
|
prompt2 = Prompt.from_function(fn2)
|
|
manager.add_prompt(prompt1)
|
|
manager.add_prompt(prompt2)
|
|
prompts = await manager.get_prompts()
|
|
assert len(prompts) == 2
|
|
assert prompts["fn1"] == prompt1
|
|
assert prompts["fn2"] == prompt2
|
|
|
|
|
|
class TestRenderPrompt:
|
|
async def test_render_prompt(self):
|
|
"""Test rendering a prompt."""
|
|
|
|
def fn() -> str:
|
|
"""An example prompt."""
|
|
return "Hello, world!"
|
|
|
|
manager = PromptManager()
|
|
prompt = Prompt.from_function(fn)
|
|
manager.add_prompt(prompt)
|
|
result = await manager.render_prompt("fn")
|
|
assert result.description == "An example prompt."
|
|
assert result.messages == [
|
|
PromptMessage(
|
|
role="user", content=TextContent(type="text", text="Hello, world!")
|
|
)
|
|
]
|
|
|
|
async def test_render_prompt_with_args(self):
|
|
"""Test rendering a prompt with arguments."""
|
|
|
|
def fn(name: str) -> str:
|
|
"""An example prompt."""
|
|
return f"Hello, {name}!"
|
|
|
|
manager = PromptManager()
|
|
prompt = Prompt.from_function(fn)
|
|
manager.add_prompt(prompt)
|
|
result = await manager.render_prompt("fn", arguments={"name": "World"})
|
|
assert result.description == "An example prompt."
|
|
assert result.messages == [
|
|
PromptMessage(
|
|
role="user", content=TextContent(type="text", text="Hello, World!")
|
|
)
|
|
]
|
|
|
|
async def test_render_prompt_callable_object(self):
|
|
"""Test rendering a prompt with a callable object."""
|
|
|
|
class MyPrompt:
|
|
"""A callable object that can be used as a prompt."""
|
|
|
|
def __call__(self, name: str) -> str:
|
|
"""ignore this"""
|
|
return f"Hello, {name}!"
|
|
|
|
manager = PromptManager()
|
|
prompt = Prompt.from_function(MyPrompt())
|
|
manager.add_prompt(prompt)
|
|
result = await manager.render_prompt("MyPrompt", arguments={"name": "World"})
|
|
assert result.description == "A callable object that can be used as a prompt."
|
|
assert result.messages == [
|
|
PromptMessage(
|
|
role="user", content=TextContent(type="text", text="Hello, World!")
|
|
)
|
|
]
|
|
|
|
async def test_render_prompt_callable_object_async(self):
|
|
"""Test rendering a prompt with a callable object."""
|
|
|
|
class MyPrompt:
|
|
"""A callable object that can be used as a prompt."""
|
|
|
|
async def __call__(self, name: str) -> str:
|
|
"""ignore this"""
|
|
return f"Hello, {name}!"
|
|
|
|
manager = PromptManager()
|
|
prompt = Prompt.from_function(MyPrompt())
|
|
manager.add_prompt(prompt)
|
|
result = await manager.render_prompt("MyPrompt", arguments={"name": "World"})
|
|
assert result.description == "A callable object that can be used as a prompt."
|
|
assert result.messages == [
|
|
PromptMessage(
|
|
role="user", content=TextContent(type="text", text="Hello, World!")
|
|
)
|
|
]
|
|
|
|
async def test_render_unknown_prompt(self):
|
|
"""Test rendering a non-existent prompt."""
|
|
manager = PromptManager()
|
|
with pytest.raises(NotFoundError, match="Unknown prompt: unknown"):
|
|
await manager.render_prompt("unknown")
|
|
|
|
async def test_render_prompt_with_missing_args(self):
|
|
"""Test rendering a prompt with missing required arguments."""
|
|
|
|
def fn(name: str) -> str:
|
|
return f"Hello, {name}!"
|
|
|
|
manager = PromptManager()
|
|
prompt = Prompt.from_function(fn)
|
|
manager.add_prompt(prompt)
|
|
with pytest.raises(PromptError, match="Missing required arguments"):
|
|
await manager.render_prompt("fn")
|
|
|
|
async def test_prompt_with_varargs_not_allowed(self):
|
|
"""Test that a prompt with *args is not allowed."""
|
|
|
|
def fn(*args: int) -> str:
|
|
return f"Hello, {args}!"
|
|
|
|
manager = PromptManager()
|
|
with pytest.raises(
|
|
ValueError, match=r"Functions with \*args are not supported as prompts"
|
|
):
|
|
manager.add_prompt(Prompt.from_function(fn))
|
|
|
|
async def test_prompt_with_varkwargs_not_allowed(self):
|
|
"""Test that a prompt with **kwargs is not allowed."""
|
|
|
|
def fn(**kwargs: int) -> str:
|
|
return f"Hello, {kwargs}!"
|
|
|
|
manager = PromptManager()
|
|
with pytest.raises(
|
|
ValueError, match=r"Functions with \*\*kwargs are not supported as prompts"
|
|
):
|
|
manager.add_prompt(Prompt.from_function(fn))
|
|
|
|
|
|
class TestPromptTags:
|
|
"""Test functionality related to prompt tags."""
|
|
|
|
async def test_add_prompt_with_tags(self):
|
|
"""Test adding a prompt with tags."""
|
|
|
|
def greeting() -> str:
|
|
return "Hello, world!"
|
|
|
|
manager = PromptManager()
|
|
prompt = Prompt.from_function(greeting, tags={"greeting", "simple"})
|
|
manager.add_prompt(prompt)
|
|
|
|
prompt = await manager.get_prompt("greeting")
|
|
assert prompt is not None
|
|
assert prompt.tags == {"greeting", "simple"}
|
|
|
|
async def test_add_prompt_with_empty_tags(self):
|
|
"""Test adding a prompt with empty tags."""
|
|
|
|
def greeting() -> str:
|
|
return "Hello, world!"
|
|
|
|
manager = PromptManager()
|
|
prompt = Prompt.from_function(greeting, tags=set())
|
|
manager.add_prompt(prompt)
|
|
|
|
prompt = await manager.get_prompt("greeting")
|
|
assert prompt is not None
|
|
assert prompt.tags == set()
|
|
|
|
async def test_add_prompt_with_none_tags(self):
|
|
"""Test adding a prompt with None tags."""
|
|
|
|
def greeting() -> str:
|
|
return "Hello, world!"
|
|
|
|
manager = PromptManager()
|
|
prompt = Prompt.from_function(greeting, tags=None)
|
|
manager.add_prompt(prompt)
|
|
|
|
prompt = await manager.get_prompt("greeting")
|
|
assert prompt is not None
|
|
assert prompt.tags == set()
|
|
|
|
async def test_list_prompts_with_tags(self):
|
|
"""Test listing prompts with specific tags."""
|
|
|
|
def greeting() -> str:
|
|
return "Hello, world!"
|
|
|
|
def weather(location: str) -> str:
|
|
return f"Weather for {location}"
|
|
|
|
def summary(text: str) -> str:
|
|
return f"Summary of: {text}"
|
|
|
|
manager = PromptManager()
|
|
manager.add_prompt(Prompt.from_function(greeting, tags={"greeting", "simple"}))
|
|
manager.add_prompt(Prompt.from_function(weather, tags={"weather", "location"}))
|
|
manager.add_prompt(
|
|
Prompt.from_function(summary, tags={"summary", "nlp", "simple"})
|
|
)
|
|
|
|
# Filter prompts by tags
|
|
prompts = await manager.get_prompts()
|
|
simple_prompts = [p for p in prompts.values() if "simple" in p.tags]
|
|
assert len(simple_prompts) == 2
|
|
assert {p.name for p in simple_prompts} == {"greeting", "summary"}
|
|
|
|
nlp_prompts = [p for p in prompts.values() if "nlp" in p.tags]
|
|
assert len(nlp_prompts) == 1
|
|
assert nlp_prompts[0].name == "summary"
|
|
|
|
|
|
class TestContextHandling:
|
|
"""Test context handling in prompts."""
|
|
|
|
def test_context_parameter_detection(self):
|
|
"""Test that context parameters are properly detected in
|
|
Prompt.from_function()."""
|
|
|
|
def prompt_with_context(x: int, ctx: Context) -> str:
|
|
return str(x)
|
|
|
|
Prompt.from_function(prompt_with_context)
|
|
|
|
def prompt_without_context(x: int) -> str:
|
|
return str(x)
|
|
|
|
Prompt.from_function(prompt_without_context)
|
|
|
|
def test_parameterized_context_parameter_detection(self):
|
|
"""Test that parameterized context parameters are properly detected in
|
|
Prompt.from_function()."""
|
|
|
|
def prompt_with_context(x: int, ctx: Context) -> str:
|
|
return str(x)
|
|
|
|
Prompt.from_function(prompt_with_context)
|
|
|
|
def test_parameterized_union_context_parameter_detection(self):
|
|
"""Test that context parameters in a union are properly detected in
|
|
Prompt.from_function()."""
|
|
|
|
def prompt_with_context(x: int, ctx: Context | None) -> str:
|
|
return str(x)
|
|
|
|
Prompt.from_function(prompt_with_context)
|
|
|
|
async def test_context_injection(self):
|
|
"""Test that context is properly injected during prompt rendering."""
|
|
|
|
def prompt_with_context(x: int, ctx: Context) -> str:
|
|
assert isinstance(ctx, Context)
|
|
return str(x)
|
|
|
|
prompt = Prompt.from_function(prompt_with_context)
|
|
|
|
from fastmcp import FastMCP
|
|
|
|
mcp = FastMCP()
|
|
context = Context(fastmcp=mcp)
|
|
|
|
async with context:
|
|
messages = await prompt.render(arguments={"x": 42})
|
|
|
|
assert len(messages) == 1
|
|
assert messages[0].content.text == "42" # type: ignore[attr-defined]
|
|
|
|
async def test_context_optional(self):
|
|
"""Test that context is optional when rendering prompts."""
|
|
|
|
def prompt_with_context(x: int, ctx: Context | None = None) -> str:
|
|
return str(x)
|
|
|
|
prompt = Prompt.from_function(prompt_with_context)
|
|
|
|
# Even for optional context, we need to provide a context
|
|
from fastmcp import FastMCP
|
|
|
|
mcp = FastMCP()
|
|
context = Context(fastmcp=mcp)
|
|
|
|
async with context:
|
|
messages = await prompt.render(
|
|
arguments={"x": 42},
|
|
)
|
|
|
|
assert len(messages) == 1
|
|
assert messages[0].content.text == "42" # type: ignore[attr-defined]
|
|
|
|
async def test_annotated_context_parameter_detection(self):
|
|
"""Test that annotated context parameters are properly detected in
|
|
Prompt.from_function()."""
|
|
|
|
def prompt_with_context(x: int, ctx: Annotated[Context, "ctx"]) -> str:
|
|
return str(x)
|
|
|
|
Prompt.from_function(prompt_with_context)
|