mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-22 21:44:18 +02:00
Unify SamplingHandler and promote OpenAI handler (#2616)
* Unify SamplingHandler and promote OpenAI handler Consolidates ServerSamplingHandler and ClientSamplingHandler into a single SamplingHandler type alias. Moves OpenAISamplingHandler from experimental to fastmcp.client.sampling.handlers.openai as the canonical location. Backwards compatibility maintained for imports from experimental. * Remove unreachable code paths in OpenAI handler * Fix docstring and use elif for mutually exclusive branches
This commit is contained in:
parent
41ec7ee06d
commit
da77cfa73f
20 changed files with 490 additions and 536 deletions
|
|
@ -185,7 +185,7 @@ For full-featured sampling with tool support, use the built-in OpenAI handler. I
|
|||
|
||||
```python
|
||||
from fastmcp import Client
|
||||
from fastmcp.experimental.sampling.handlers.openai import OpenAISamplingHandler
|
||||
from fastmcp.client.sampling.handlers.openai import OpenAISamplingHandler
|
||||
|
||||
client = Client(
|
||||
"my_mcp_server.py",
|
||||
|
|
@ -212,5 +212,5 @@ Tool execution happens on the server side. The client's role is to pass tools to
|
|||
</Note>
|
||||
|
||||
<Tip>
|
||||
To implement a custom sampling handler, see the [OpenAISamplingHandler source code](https://github.com/jlowin/fastmcp/blob/main/src/fastmcp/experimental/sampling/handlers/openai.py) as a reference.
|
||||
To implement a custom sampling handler, see the [OpenAISamplingHandler source code](https://github.com/jlowin/fastmcp/blob/main/src/fastmcp/client/sampling/handlers/openai.py) as a reference.
|
||||
</Tip>
|
||||
|
|
@ -461,9 +461,7 @@ FastMCP provides an OpenAI-compatible handler that works with OpenAI's API and c
|
|||
import os
|
||||
from openai import OpenAI
|
||||
from fastmcp import FastMCP
|
||||
from fastmcp.experimental.sampling.handlers.openai import (
|
||||
OpenAISamplingHandler,
|
||||
)
|
||||
from fastmcp.client.sampling.handlers.openai import OpenAISamplingHandler
|
||||
|
||||
server = FastMCP(
|
||||
name="My Server",
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ import asyncio
|
|||
from pydantic import BaseModel
|
||||
|
||||
from fastmcp import Client, Context, FastMCP
|
||||
from fastmcp.experimental.sampling.handlers.openai import OpenAISamplingHandler
|
||||
from fastmcp.client.sampling.handlers.openai import OpenAISamplingHandler
|
||||
|
||||
# Create the MCP server
|
||||
mcp = FastMCP("Sampling Test Server")
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ import asyncio
|
|||
from pydantic import BaseModel
|
||||
|
||||
from fastmcp import Client, Context, FastMCP
|
||||
from fastmcp.experimental.sampling.handlers.openai import OpenAISamplingHandler
|
||||
from fastmcp.client.sampling.handlers.openai import OpenAISamplingHandler
|
||||
|
||||
|
||||
# Define a structured output model
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ import asyncio
|
|||
from pydantic import BaseModel, Field
|
||||
|
||||
from fastmcp import Client, Context, FastMCP
|
||||
from fastmcp.experimental.sampling.handlers.openai import OpenAISamplingHandler
|
||||
from fastmcp.client.sampling.handlers.openai import OpenAISamplingHandler
|
||||
|
||||
|
||||
# Define tools (available to the LLM during sampling)
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ from mcp.types import ContentBlock
|
|||
from openai import OpenAI
|
||||
|
||||
from fastmcp import FastMCP
|
||||
from fastmcp.experimental.sampling.handlers.openai import OpenAISamplingHandler
|
||||
from fastmcp.client.sampling.handlers.openai import OpenAISamplingHandler
|
||||
from fastmcp.server.context import Context
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -47,7 +47,6 @@ from fastmcp.client.roots import (
|
|||
create_roots_callback,
|
||||
)
|
||||
from fastmcp.client.sampling import (
|
||||
ClientSamplingHandler,
|
||||
SamplingHandler,
|
||||
create_sampling_callback,
|
||||
)
|
||||
|
|
@ -82,7 +81,6 @@ from .transports import (
|
|||
|
||||
__all__ = [
|
||||
"Client",
|
||||
"ClientSamplingHandler",
|
||||
"ElicitationHandler",
|
||||
"LogHandler",
|
||||
"MessageHandler",
|
||||
|
|
@ -248,7 +246,7 @@ class Client(Generic[ClientTransportT]):
|
|||
),
|
||||
name: str | None = None,
|
||||
roots: RootsList | RootsHandler | None = None,
|
||||
sampling_handler: ClientSamplingHandler | None = None,
|
||||
sampling_handler: SamplingHandler | None = None,
|
||||
sampling_capabilities: mcp.types.SamplingCapability | None = None,
|
||||
elicitation_handler: ElicitationHandler | None = None,
|
||||
log_handler: LogHandler | None = None,
|
||||
|
|
@ -368,7 +366,7 @@ class Client(Generic[ClientTransportT]):
|
|||
|
||||
def set_sampling_callback(
|
||||
self,
|
||||
sampling_callback: ClientSamplingHandler,
|
||||
sampling_callback: SamplingHandler,
|
||||
sampling_capabilities: mcp.types.SamplingCapability | None = None,
|
||||
) -> None:
|
||||
"""Set the sampling callback for the client."""
|
||||
|
|
|
|||
|
|
@ -1,56 +0,0 @@
|
|||
import inspect
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import TypeAlias
|
||||
|
||||
import mcp.types
|
||||
from mcp import CreateMessageResult
|
||||
from mcp.client.session import ClientSession, SamplingFnT
|
||||
from mcp.shared.context import LifespanContextT, RequestContext
|
||||
from mcp.types import CreateMessageRequestParams as SamplingParams
|
||||
from mcp.types import SamplingMessage
|
||||
|
||||
from fastmcp.server.sampling.handler import ServerSamplingHandler
|
||||
|
||||
__all__ = ["SamplingHandler", "SamplingMessage", "SamplingParams"]
|
||||
|
||||
|
||||
ClientSamplingHandler: TypeAlias = Callable[
|
||||
[
|
||||
list[SamplingMessage],
|
||||
SamplingParams,
|
||||
RequestContext[ClientSession, LifespanContextT],
|
||||
],
|
||||
str | CreateMessageResult | Awaitable[str | CreateMessageResult],
|
||||
]
|
||||
|
||||
SamplingHandler: TypeAlias = (
|
||||
ClientSamplingHandler[LifespanContextT] | ServerSamplingHandler[LifespanContextT]
|
||||
)
|
||||
|
||||
|
||||
def create_sampling_callback(
|
||||
sampling_handler: ClientSamplingHandler[LifespanContextT],
|
||||
) -> SamplingFnT:
|
||||
async def _sampling_handler(
|
||||
context: RequestContext[ClientSession, LifespanContextT],
|
||||
params: SamplingParams,
|
||||
) -> CreateMessageResult | mcp.types.ErrorData:
|
||||
try:
|
||||
result = sampling_handler(params.messages, params, context)
|
||||
if inspect.isawaitable(result):
|
||||
result = await result
|
||||
|
||||
if isinstance(result, str):
|
||||
result = CreateMessageResult(
|
||||
role="assistant",
|
||||
model="fastmcp-client",
|
||||
content=mcp.types.TextContent(type="text", text=result),
|
||||
)
|
||||
return result
|
||||
except Exception as e:
|
||||
return mcp.types.ErrorData(
|
||||
code=mcp.types.INTERNAL_ERROR,
|
||||
message=str(e),
|
||||
)
|
||||
|
||||
return _sampling_handler
|
||||
69
src/fastmcp/client/sampling/__init__.py
Normal file
69
src/fastmcp/client/sampling/__init__.py
Normal file
|
|
@ -0,0 +1,69 @@
|
|||
import inspect
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import TypeAlias, TypeVar
|
||||
|
||||
import mcp.types
|
||||
from mcp import ClientSession, CreateMessageResult
|
||||
from mcp.client.session import SamplingFnT
|
||||
from mcp.server.session import ServerSession
|
||||
from mcp.shared.context import LifespanContextT, RequestContext
|
||||
from mcp.types import CreateMessageRequestParams as SamplingParams
|
||||
from mcp.types import CreateMessageResultWithTools, SamplingMessage
|
||||
|
||||
# Result type that handlers can return
|
||||
SamplingHandlerResult: TypeAlias = (
|
||||
str | CreateMessageResult | CreateMessageResultWithTools
|
||||
)
|
||||
|
||||
# Session type for sampling handlers - works with both client and server sessions
|
||||
SessionT = TypeVar("SessionT", ClientSession, ServerSession)
|
||||
|
||||
# Unified sampling handler type that works for both clients and servers.
|
||||
# Handlers receive messages and parameters from the MCP sampling flow
|
||||
# and return LLM responses.
|
||||
SamplingHandler: TypeAlias = Callable[
|
||||
[
|
||||
list[SamplingMessage],
|
||||
SamplingParams,
|
||||
RequestContext[SessionT, LifespanContextT],
|
||||
],
|
||||
SamplingHandlerResult | Awaitable[SamplingHandlerResult],
|
||||
]
|
||||
|
||||
|
||||
__all__ = [
|
||||
"RequestContext",
|
||||
"SamplingHandler",
|
||||
"SamplingHandlerResult",
|
||||
"SamplingMessage",
|
||||
"SamplingParams",
|
||||
"create_sampling_callback",
|
||||
]
|
||||
|
||||
|
||||
def create_sampling_callback(
|
||||
sampling_handler: SamplingHandler,
|
||||
) -> SamplingFnT:
|
||||
async def _sampling_handler(
|
||||
context,
|
||||
params: SamplingParams,
|
||||
) -> CreateMessageResult | CreateMessageResultWithTools | mcp.types.ErrorData:
|
||||
try:
|
||||
result = sampling_handler(params.messages, params, context)
|
||||
if inspect.isawaitable(result):
|
||||
result = await result
|
||||
|
||||
if isinstance(result, str):
|
||||
result = CreateMessageResult(
|
||||
role="assistant",
|
||||
model="fastmcp-client",
|
||||
content=mcp.types.TextContent(type="text", text=result),
|
||||
)
|
||||
return result
|
||||
except Exception as e:
|
||||
return mcp.types.ErrorData(
|
||||
code=mcp.types.INTERNAL_ERROR,
|
||||
message=str(e),
|
||||
)
|
||||
|
||||
return _sampling_handler
|
||||
0
src/fastmcp/client/sampling/handlers/__init__.py
Normal file
0
src/fastmcp/client/sampling/handlers/__init__.py
Normal file
399
src/fastmcp/client/sampling/handlers/openai.py
Normal file
399
src/fastmcp/client/sampling/handlers/openai.py
Normal file
|
|
@ -0,0 +1,399 @@
|
|||
"""OpenAI sampling handler for FastMCP."""
|
||||
|
||||
import json
|
||||
from collections.abc import Iterator, Sequence
|
||||
from typing import Any, get_args
|
||||
|
||||
from mcp import ClientSession, ServerSession
|
||||
from mcp.shared.context import LifespanContextT, RequestContext
|
||||
from mcp.types import CreateMessageRequestParams as SamplingParams
|
||||
from mcp.types import (
|
||||
CreateMessageResult,
|
||||
CreateMessageResultWithTools,
|
||||
ModelPreferences,
|
||||
SamplingMessage,
|
||||
StopReason,
|
||||
TextContent,
|
||||
Tool,
|
||||
ToolChoice,
|
||||
ToolResultContent,
|
||||
ToolUseContent,
|
||||
)
|
||||
|
||||
try:
|
||||
from openai import NOT_GIVEN, AsyncOpenAI, NotGiven
|
||||
from openai.types.chat import (
|
||||
ChatCompletion,
|
||||
ChatCompletionAssistantMessageParam,
|
||||
ChatCompletionMessageParam,
|
||||
ChatCompletionMessageToolCallParam,
|
||||
ChatCompletionSystemMessageParam,
|
||||
ChatCompletionToolChoiceOptionParam,
|
||||
ChatCompletionToolMessageParam,
|
||||
ChatCompletionToolParam,
|
||||
ChatCompletionUserMessageParam,
|
||||
)
|
||||
from openai.types.shared.chat_model import ChatModel
|
||||
from openai.types.shared_params import FunctionDefinition
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"The `openai` package is not installed. "
|
||||
"Please install `fastmcp[openai]` or add `openai` to your dependencies manually."
|
||||
) from e
|
||||
|
||||
|
||||
class OpenAISamplingHandler:
|
||||
"""Sampling handler that uses the OpenAI API."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
default_model: ChatModel,
|
||||
client: AsyncOpenAI | None = None,
|
||||
) -> None:
|
||||
self.client: AsyncOpenAI = client or AsyncOpenAI()
|
||||
self.default_model: ChatModel = default_model
|
||||
|
||||
async def __call__(
|
||||
self,
|
||||
messages: list[SamplingMessage],
|
||||
params: SamplingParams,
|
||||
context: RequestContext[ServerSession, LifespanContextT]
|
||||
| RequestContext[ClientSession, LifespanContextT],
|
||||
) -> CreateMessageResult | CreateMessageResultWithTools:
|
||||
openai_messages: list[ChatCompletionMessageParam] = (
|
||||
self._convert_to_openai_messages(
|
||||
system_prompt=params.systemPrompt,
|
||||
messages=messages,
|
||||
)
|
||||
)
|
||||
|
||||
model: ChatModel = self._select_model_from_preferences(params.modelPreferences)
|
||||
|
||||
# Convert MCP tools to OpenAI format
|
||||
openai_tools: list[ChatCompletionToolParam] | NotGiven = NOT_GIVEN
|
||||
if params.tools:
|
||||
openai_tools = self._convert_tools_to_openai(params.tools)
|
||||
|
||||
# Convert tool_choice to OpenAI format
|
||||
openai_tool_choice: ChatCompletionToolChoiceOptionParam | NotGiven = NOT_GIVEN
|
||||
if params.toolChoice:
|
||||
openai_tool_choice = self._convert_tool_choice_to_openai(params.toolChoice)
|
||||
|
||||
response = await self.client.chat.completions.create(
|
||||
model=model,
|
||||
messages=openai_messages,
|
||||
temperature=(
|
||||
params.temperature if params.temperature is not None else NOT_GIVEN
|
||||
),
|
||||
max_tokens=params.maxTokens,
|
||||
stop=params.stopSequences if params.stopSequences else NOT_GIVEN,
|
||||
tools=openai_tools,
|
||||
tool_choice=openai_tool_choice,
|
||||
)
|
||||
|
||||
# Return appropriate result type based on whether tools were provided
|
||||
if params.tools:
|
||||
return self._chat_completion_to_result_with_tools(response)
|
||||
return self._chat_completion_to_create_message_result(response)
|
||||
|
||||
@staticmethod
|
||||
def _iter_models_from_preferences(
|
||||
model_preferences: ModelPreferences | str | list[str] | None,
|
||||
) -> Iterator[str]:
|
||||
if model_preferences is None:
|
||||
return
|
||||
|
||||
if isinstance(model_preferences, str) and model_preferences in get_args(
|
||||
ChatModel
|
||||
):
|
||||
yield model_preferences
|
||||
|
||||
elif isinstance(model_preferences, list):
|
||||
yield from model_preferences
|
||||
|
||||
elif isinstance(model_preferences, ModelPreferences):
|
||||
if not (hints := model_preferences.hints):
|
||||
return
|
||||
|
||||
for hint in hints:
|
||||
if not (name := hint.name):
|
||||
continue
|
||||
|
||||
yield name
|
||||
|
||||
@staticmethod
|
||||
def _convert_to_openai_messages(
|
||||
system_prompt: str | None, messages: Sequence[SamplingMessage]
|
||||
) -> list[ChatCompletionMessageParam]:
|
||||
openai_messages: list[ChatCompletionMessageParam] = []
|
||||
|
||||
if system_prompt:
|
||||
openai_messages.append(
|
||||
ChatCompletionSystemMessageParam(
|
||||
role="system",
|
||||
content=system_prompt,
|
||||
)
|
||||
)
|
||||
|
||||
for message in messages:
|
||||
content = message.content
|
||||
|
||||
# Handle list content (from CreateMessageResultWithTools)
|
||||
if isinstance(content, list):
|
||||
# Collect tool calls and text from the list
|
||||
tool_calls: list[ChatCompletionMessageToolCallParam] = []
|
||||
text_parts: list[str] = []
|
||||
# Collect tool results separately to maintain correct ordering
|
||||
tool_messages: list[ChatCompletionToolMessageParam] = []
|
||||
|
||||
for item in content:
|
||||
if isinstance(item, ToolUseContent):
|
||||
tool_calls.append(
|
||||
ChatCompletionMessageToolCallParam(
|
||||
id=item.id,
|
||||
type="function",
|
||||
function={
|
||||
"name": item.name,
|
||||
"arguments": json.dumps(item.input),
|
||||
},
|
||||
)
|
||||
)
|
||||
elif isinstance(item, TextContent):
|
||||
text_parts.append(item.text)
|
||||
elif isinstance(item, ToolResultContent):
|
||||
# Collect tool results (added after assistant message)
|
||||
content_text = ""
|
||||
if item.content:
|
||||
result_texts = []
|
||||
for sub_item in item.content:
|
||||
if isinstance(sub_item, TextContent):
|
||||
result_texts.append(sub_item.text)
|
||||
content_text = "\n".join(result_texts)
|
||||
tool_messages.append(
|
||||
ChatCompletionToolMessageParam(
|
||||
role="tool",
|
||||
tool_call_id=item.toolUseId,
|
||||
content=content_text,
|
||||
)
|
||||
)
|
||||
|
||||
# Add assistant message with tool calls if present
|
||||
# OpenAI requires: assistant (with tool_calls) -> tool messages
|
||||
if tool_calls or text_parts:
|
||||
msg_content = "\n".join(text_parts) if text_parts else None
|
||||
if tool_calls:
|
||||
openai_messages.append(
|
||||
ChatCompletionAssistantMessageParam(
|
||||
role="assistant",
|
||||
content=msg_content,
|
||||
tool_calls=tool_calls,
|
||||
)
|
||||
)
|
||||
# Add tool messages AFTER assistant message
|
||||
openai_messages.extend(tool_messages)
|
||||
elif msg_content:
|
||||
if message.role == "user":
|
||||
openai_messages.append(
|
||||
ChatCompletionUserMessageParam(
|
||||
role="user",
|
||||
content=msg_content,
|
||||
)
|
||||
)
|
||||
else:
|
||||
openai_messages.append(
|
||||
ChatCompletionAssistantMessageParam(
|
||||
role="assistant",
|
||||
content=msg_content,
|
||||
)
|
||||
)
|
||||
elif tool_messages:
|
||||
# Tool results only (assistant message was in previous message)
|
||||
openai_messages.extend(tool_messages)
|
||||
continue
|
||||
|
||||
# Handle ToolUseContent (assistant's tool calls)
|
||||
if isinstance(content, ToolUseContent):
|
||||
openai_messages.append(
|
||||
ChatCompletionAssistantMessageParam(
|
||||
role="assistant",
|
||||
tool_calls=[
|
||||
ChatCompletionMessageToolCallParam(
|
||||
id=content.id,
|
||||
type="function",
|
||||
function={
|
||||
"name": content.name,
|
||||
"arguments": json.dumps(content.input),
|
||||
},
|
||||
)
|
||||
],
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Handle ToolResultContent (user's tool results)
|
||||
if isinstance(content, ToolResultContent):
|
||||
# Extract text parts from the content list
|
||||
result_texts: list[str] = []
|
||||
if content.content:
|
||||
for item in content.content:
|
||||
if isinstance(item, TextContent):
|
||||
result_texts.append(item.text)
|
||||
openai_messages.append(
|
||||
ChatCompletionToolMessageParam(
|
||||
role="tool",
|
||||
tool_call_id=content.toolUseId,
|
||||
content="\n".join(result_texts),
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Handle TextContent
|
||||
if isinstance(content, TextContent):
|
||||
if message.role == "user":
|
||||
openai_messages.append(
|
||||
ChatCompletionUserMessageParam(
|
||||
role="user",
|
||||
content=content.text,
|
||||
)
|
||||
)
|
||||
else:
|
||||
openai_messages.append(
|
||||
ChatCompletionAssistantMessageParam(
|
||||
role="assistant",
|
||||
content=content.text,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
raise ValueError(f"Unsupported content type: {type(content)}")
|
||||
|
||||
return openai_messages
|
||||
|
||||
@staticmethod
|
||||
def _chat_completion_to_create_message_result(
|
||||
chat_completion: ChatCompletion,
|
||||
) -> CreateMessageResult:
|
||||
if len(chat_completion.choices) == 0:
|
||||
raise ValueError("No response for completion")
|
||||
|
||||
first_choice = chat_completion.choices[0]
|
||||
|
||||
if content := first_choice.message.content:
|
||||
return CreateMessageResult(
|
||||
content=TextContent(type="text", text=content),
|
||||
role="assistant",
|
||||
model=chat_completion.model,
|
||||
)
|
||||
|
||||
raise ValueError("No content in response from completion")
|
||||
|
||||
def _select_model_from_preferences(
|
||||
self, model_preferences: ModelPreferences | str | list[str] | None
|
||||
) -> ChatModel:
|
||||
for model_option in self._iter_models_from_preferences(model_preferences):
|
||||
if model_option in get_args(ChatModel):
|
||||
chosen_model: ChatModel = model_option # type: ignore[assignment]
|
||||
return chosen_model
|
||||
|
||||
return self.default_model
|
||||
|
||||
@staticmethod
|
||||
def _convert_tools_to_openai(tools: list[Tool]) -> list[ChatCompletionToolParam]:
|
||||
"""Convert MCP tools to OpenAI tool format."""
|
||||
openai_tools: list[ChatCompletionToolParam] = []
|
||||
for tool in tools:
|
||||
# Build parameters dict, ensuring required fields
|
||||
parameters: dict[str, Any] = dict(tool.inputSchema)
|
||||
if "type" not in parameters:
|
||||
parameters["type"] = "object"
|
||||
|
||||
openai_tools.append(
|
||||
ChatCompletionToolParam(
|
||||
type="function",
|
||||
function=FunctionDefinition(
|
||||
name=tool.name,
|
||||
description=tool.description or "",
|
||||
parameters=parameters,
|
||||
),
|
||||
)
|
||||
)
|
||||
return openai_tools
|
||||
|
||||
@staticmethod
|
||||
def _convert_tool_choice_to_openai(
|
||||
tool_choice: ToolChoice,
|
||||
) -> ChatCompletionToolChoiceOptionParam:
|
||||
"""Convert MCP tool_choice to OpenAI format."""
|
||||
if tool_choice.mode == "auto":
|
||||
return "auto"
|
||||
elif tool_choice.mode == "required":
|
||||
return "required"
|
||||
elif tool_choice.mode == "none":
|
||||
return "none"
|
||||
else:
|
||||
raise ValueError(f"Unsupported tool_choice mode: {tool_choice.mode!r}")
|
||||
|
||||
@staticmethod
|
||||
def _chat_completion_to_result_with_tools(
|
||||
chat_completion: ChatCompletion,
|
||||
) -> CreateMessageResultWithTools:
|
||||
"""Convert OpenAI response to CreateMessageResultWithTools."""
|
||||
if len(chat_completion.choices) == 0:
|
||||
raise ValueError("No response for completion")
|
||||
|
||||
first_choice = chat_completion.choices[0]
|
||||
message = first_choice.message
|
||||
|
||||
# Determine stop reason
|
||||
stop_reason: StopReason
|
||||
if first_choice.finish_reason == "tool_calls":
|
||||
stop_reason = "toolUse"
|
||||
elif first_choice.finish_reason == "stop":
|
||||
stop_reason = "endTurn"
|
||||
elif first_choice.finish_reason == "length":
|
||||
stop_reason = "maxTokens"
|
||||
else:
|
||||
stop_reason = "endTurn"
|
||||
|
||||
# Build content list
|
||||
content: list[TextContent | ToolUseContent] = []
|
||||
|
||||
# Add text content if present
|
||||
if message.content:
|
||||
content.append(TextContent(type="text", text=message.content))
|
||||
|
||||
# Add tool calls if present
|
||||
if message.tool_calls:
|
||||
for tool_call in message.tool_calls:
|
||||
# Skip non-function tool calls
|
||||
if not hasattr(tool_call, "function"):
|
||||
continue
|
||||
func = tool_call.function # type: ignore[union-attr]
|
||||
# Parse the arguments JSON string
|
||||
try:
|
||||
arguments = json.loads(func.arguments) # type: ignore[union-attr]
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(
|
||||
f"Invalid JSON in tool arguments for "
|
||||
f"'{func.name}': {func.arguments}" # type: ignore[union-attr]
|
||||
) from e
|
||||
|
||||
content.append(
|
||||
ToolUseContent(
|
||||
type="tool_use",
|
||||
id=tool_call.id,
|
||||
name=func.name, # type: ignore[union-attr]
|
||||
input=arguments,
|
||||
)
|
||||
)
|
||||
|
||||
# Must have at least some content
|
||||
if not content:
|
||||
raise ValueError("No content in response from completion")
|
||||
|
||||
return CreateMessageResultWithTools(
|
||||
content=content, # type: ignore[arg-type]
|
||||
role="assistant",
|
||||
model=chat_completion.model,
|
||||
stopReason=stop_reason,
|
||||
)
|
||||
|
|
@ -0,0 +1,5 @@
|
|||
# Re-export for backwards compatibility
|
||||
# The canonical location is now fastmcp.client.sampling.handlers
|
||||
from fastmcp.client.sampling.handlers.openai import OpenAISamplingHandler
|
||||
|
||||
__all__ = ["OpenAISamplingHandler"]
|
||||
|
|
@ -1,21 +0,0 @@
|
|||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Awaitable
|
||||
|
||||
from mcp import ClientSession, CreateMessageResult
|
||||
from mcp.server.session import ServerSession
|
||||
from mcp.shared.context import LifespanContextT, RequestContext
|
||||
from mcp.types import CreateMessageRequestParams as SamplingParams
|
||||
from mcp.types import (
|
||||
SamplingMessage,
|
||||
)
|
||||
|
||||
|
||||
class BaseLLMSamplingHandler(ABC):
|
||||
@abstractmethod
|
||||
def __call__(
|
||||
self,
|
||||
messages: list[SamplingMessage],
|
||||
params: SamplingParams,
|
||||
context: RequestContext[ServerSession, LifespanContextT]
|
||||
| RequestContext[ClientSession, LifespanContextT],
|
||||
) -> str | CreateMessageResult | Awaitable[str | CreateMessageResult]: ...
|
||||
|
|
@ -1,417 +1,5 @@
|
|||
import json
|
||||
from collections.abc import Iterator, Sequence
|
||||
from typing import Any, get_args
|
||||
# Re-export for backwards compatibility
|
||||
# The canonical location is now fastmcp.client.sampling.handlers.openai
|
||||
from fastmcp.client.sampling.handlers.openai import OpenAISamplingHandler
|
||||
|
||||
from mcp import ClientSession, ServerSession
|
||||
from mcp.shared.context import LifespanContextT, RequestContext
|
||||
from mcp.types import CreateMessageRequestParams as SamplingParams
|
||||
from mcp.types import (
|
||||
CreateMessageResult,
|
||||
CreateMessageResultWithTools,
|
||||
ModelPreferences,
|
||||
SamplingMessage,
|
||||
StopReason,
|
||||
TextContent,
|
||||
Tool,
|
||||
ToolChoice,
|
||||
ToolResultContent,
|
||||
ToolUseContent,
|
||||
)
|
||||
|
||||
try:
|
||||
from openai import NOT_GIVEN, AsyncOpenAI, NotGiven
|
||||
from openai.types.chat import (
|
||||
ChatCompletion,
|
||||
ChatCompletionAssistantMessageParam,
|
||||
ChatCompletionMessageParam,
|
||||
ChatCompletionMessageToolCallParam,
|
||||
ChatCompletionSystemMessageParam,
|
||||
ChatCompletionToolChoiceOptionParam,
|
||||
ChatCompletionToolMessageParam,
|
||||
ChatCompletionToolParam,
|
||||
ChatCompletionUserMessageParam,
|
||||
)
|
||||
from openai.types.shared.chat_model import ChatModel
|
||||
from openai.types.shared_params import FunctionDefinition
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"The `openai` package is not installed. Please install `fastmcp[openai]` or add `openai` to your dependencies manually."
|
||||
) from e
|
||||
|
||||
from typing_extensions import override
|
||||
|
||||
from fastmcp.experimental.sampling.handlers.base import BaseLLMSamplingHandler
|
||||
|
||||
|
||||
class OpenAISamplingHandler(BaseLLMSamplingHandler):
|
||||
def __init__(
|
||||
self,
|
||||
default_model: ChatModel,
|
||||
client: AsyncOpenAI | None = None,
|
||||
) -> None:
|
||||
self.client: AsyncOpenAI = client or AsyncOpenAI()
|
||||
self.default_model: ChatModel = default_model
|
||||
|
||||
@override
|
||||
async def __call__(
|
||||
self,
|
||||
messages: list[SamplingMessage],
|
||||
params: SamplingParams,
|
||||
context: RequestContext[ServerSession, LifespanContextT]
|
||||
| RequestContext[ClientSession, LifespanContextT],
|
||||
) -> CreateMessageResult | CreateMessageResultWithTools:
|
||||
openai_messages: list[ChatCompletionMessageParam] = (
|
||||
self._convert_to_openai_messages(
|
||||
system_prompt=params.systemPrompt,
|
||||
messages=messages,
|
||||
)
|
||||
)
|
||||
|
||||
model: ChatModel = self._select_model_from_preferences(params.modelPreferences)
|
||||
|
||||
# Convert MCP tools to OpenAI format
|
||||
openai_tools: list[ChatCompletionToolParam] | NotGiven = NOT_GIVEN
|
||||
if params.tools:
|
||||
openai_tools = self._convert_tools_to_openai(params.tools)
|
||||
|
||||
# Convert tool_choice to OpenAI format
|
||||
openai_tool_choice: ChatCompletionToolChoiceOptionParam | NotGiven = NOT_GIVEN
|
||||
if params.toolChoice:
|
||||
openai_tool_choice = self._convert_tool_choice_to_openai(params.toolChoice)
|
||||
|
||||
response = await self.client.chat.completions.create(
|
||||
model=model,
|
||||
messages=openai_messages,
|
||||
temperature=(
|
||||
params.temperature if params.temperature is not None else NOT_GIVEN
|
||||
),
|
||||
max_tokens=params.maxTokens,
|
||||
stop=params.stopSequences if params.stopSequences else NOT_GIVEN,
|
||||
tools=openai_tools,
|
||||
tool_choice=openai_tool_choice,
|
||||
)
|
||||
|
||||
# Return appropriate result type based on whether tools were provided
|
||||
if params.tools:
|
||||
return self._chat_completion_to_result_with_tools(response)
|
||||
return self._chat_completion_to_create_message_result(response)
|
||||
|
||||
@staticmethod
|
||||
def _iter_models_from_preferences(
|
||||
model_preferences: ModelPreferences | str | list[str] | None,
|
||||
) -> Iterator[str]:
|
||||
if model_preferences is None:
|
||||
return
|
||||
|
||||
if isinstance(model_preferences, str) and model_preferences in get_args(
|
||||
ChatModel
|
||||
):
|
||||
yield model_preferences
|
||||
|
||||
if isinstance(model_preferences, list):
|
||||
yield from model_preferences
|
||||
|
||||
if isinstance(model_preferences, ModelPreferences):
|
||||
if not (hints := model_preferences.hints):
|
||||
return
|
||||
|
||||
for hint in hints:
|
||||
if not (name := hint.name):
|
||||
continue
|
||||
|
||||
yield name
|
||||
|
||||
@staticmethod
|
||||
def _convert_to_openai_messages(
|
||||
system_prompt: str | None, messages: Sequence[SamplingMessage]
|
||||
) -> list[ChatCompletionMessageParam]:
|
||||
openai_messages: list[ChatCompletionMessageParam] = []
|
||||
|
||||
if system_prompt:
|
||||
openai_messages.append(
|
||||
ChatCompletionSystemMessageParam(
|
||||
role="system",
|
||||
content=system_prompt,
|
||||
)
|
||||
)
|
||||
|
||||
if isinstance(messages, str):
|
||||
openai_messages.append(
|
||||
ChatCompletionUserMessageParam(
|
||||
role="user",
|
||||
content=messages,
|
||||
)
|
||||
)
|
||||
|
||||
if isinstance(messages, list):
|
||||
for message in messages:
|
||||
if isinstance(message, str):
|
||||
openai_messages.append(
|
||||
ChatCompletionUserMessageParam(
|
||||
role="user",
|
||||
content=message,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
content = message.content
|
||||
|
||||
# Handle list content (from CreateMessageResultWithTools)
|
||||
if isinstance(content, list):
|
||||
# Collect tool calls and text from the list
|
||||
tool_calls: list[ChatCompletionMessageToolCallParam] = []
|
||||
text_parts: list[str] = []
|
||||
# Collect tool results separately to maintain correct ordering
|
||||
tool_messages: list[ChatCompletionToolMessageParam] = []
|
||||
|
||||
for item in content:
|
||||
if isinstance(item, ToolUseContent):
|
||||
tool_calls.append(
|
||||
ChatCompletionMessageToolCallParam(
|
||||
id=item.id,
|
||||
type="function",
|
||||
function={
|
||||
"name": item.name,
|
||||
"arguments": json.dumps(item.input),
|
||||
},
|
||||
)
|
||||
)
|
||||
elif isinstance(item, TextContent):
|
||||
text_parts.append(item.text)
|
||||
elif isinstance(item, ToolResultContent):
|
||||
# Collect tool results (added after assistant message)
|
||||
content_text = ""
|
||||
if item.content:
|
||||
result_texts = []
|
||||
for sub_item in item.content:
|
||||
if isinstance(sub_item, TextContent):
|
||||
result_texts.append(sub_item.text)
|
||||
content_text = "\n".join(result_texts)
|
||||
tool_messages.append(
|
||||
ChatCompletionToolMessageParam(
|
||||
role="tool",
|
||||
tool_call_id=item.toolUseId,
|
||||
content=content_text,
|
||||
)
|
||||
)
|
||||
|
||||
# Add assistant message with tool calls if present
|
||||
# OpenAI requires: assistant (with tool_calls) -> tool messages
|
||||
if tool_calls or text_parts:
|
||||
msg_content = "\n".join(text_parts) if text_parts else None
|
||||
if tool_calls:
|
||||
openai_messages.append(
|
||||
ChatCompletionAssistantMessageParam(
|
||||
role="assistant",
|
||||
content=msg_content,
|
||||
tool_calls=tool_calls,
|
||||
)
|
||||
)
|
||||
# Add tool messages AFTER assistant message
|
||||
openai_messages.extend(tool_messages)
|
||||
elif msg_content:
|
||||
if message.role == "user":
|
||||
openai_messages.append(
|
||||
ChatCompletionUserMessageParam(
|
||||
role="user",
|
||||
content=msg_content,
|
||||
)
|
||||
)
|
||||
else:
|
||||
openai_messages.append(
|
||||
ChatCompletionAssistantMessageParam(
|
||||
role="assistant",
|
||||
content=msg_content,
|
||||
)
|
||||
)
|
||||
elif tool_messages:
|
||||
# Tool results only (assistant message was in previous message)
|
||||
openai_messages.extend(tool_messages)
|
||||
continue
|
||||
|
||||
# Handle ToolUseContent (assistant's tool calls)
|
||||
if isinstance(content, ToolUseContent):
|
||||
openai_messages.append(
|
||||
ChatCompletionAssistantMessageParam(
|
||||
role="assistant",
|
||||
tool_calls=[
|
||||
ChatCompletionMessageToolCallParam(
|
||||
id=content.id,
|
||||
type="function",
|
||||
function={
|
||||
"name": content.name,
|
||||
"arguments": json.dumps(content.input),
|
||||
},
|
||||
)
|
||||
],
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Handle ToolResultContent (user's tool results)
|
||||
if isinstance(content, ToolResultContent):
|
||||
# Extract text parts from the content list
|
||||
result_texts: list[str] = []
|
||||
if content.content:
|
||||
for item in content.content:
|
||||
if isinstance(item, TextContent):
|
||||
result_texts.append(item.text)
|
||||
openai_messages.append(
|
||||
ChatCompletionToolMessageParam(
|
||||
role="tool",
|
||||
tool_call_id=content.toolUseId,
|
||||
content="\n".join(result_texts),
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Handle TextContent
|
||||
if isinstance(content, TextContent):
|
||||
if message.role == "user":
|
||||
openai_messages.append(
|
||||
ChatCompletionUserMessageParam(
|
||||
role="user",
|
||||
content=content.text,
|
||||
)
|
||||
)
|
||||
else:
|
||||
openai_messages.append(
|
||||
ChatCompletionAssistantMessageParam(
|
||||
role="assistant",
|
||||
content=content.text,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
raise ValueError(f"Unsupported content type: {type(content)}")
|
||||
|
||||
return openai_messages
|
||||
|
||||
@staticmethod
|
||||
def _chat_completion_to_create_message_result(
|
||||
chat_completion: ChatCompletion,
|
||||
) -> CreateMessageResult:
|
||||
if len(chat_completion.choices) == 0:
|
||||
raise ValueError("No response for completion")
|
||||
|
||||
first_choice = chat_completion.choices[0]
|
||||
|
||||
if content := first_choice.message.content:
|
||||
return CreateMessageResult(
|
||||
content=TextContent(type="text", text=content),
|
||||
role="assistant",
|
||||
model=chat_completion.model,
|
||||
)
|
||||
|
||||
raise ValueError("No content in response from completion")
|
||||
|
||||
def _select_model_from_preferences(
|
||||
self, model_preferences: ModelPreferences | str | list[str] | None
|
||||
) -> ChatModel:
|
||||
for model_option in self._iter_models_from_preferences(model_preferences):
|
||||
if model_option in get_args(ChatModel):
|
||||
chosen_model: ChatModel = model_option # type: ignore[assignment]
|
||||
return chosen_model
|
||||
|
||||
return self.default_model
|
||||
|
||||
@staticmethod
|
||||
def _convert_tools_to_openai(tools: list[Tool]) -> list[ChatCompletionToolParam]:
|
||||
"""Convert MCP tools to OpenAI tool format."""
|
||||
openai_tools: list[ChatCompletionToolParam] = []
|
||||
for tool in tools:
|
||||
# Build parameters dict, ensuring required fields
|
||||
parameters: dict[str, Any] = dict(tool.inputSchema)
|
||||
if "type" not in parameters:
|
||||
parameters["type"] = "object"
|
||||
|
||||
openai_tools.append(
|
||||
ChatCompletionToolParam(
|
||||
type="function",
|
||||
function=FunctionDefinition(
|
||||
name=tool.name,
|
||||
description=tool.description or "",
|
||||
parameters=parameters,
|
||||
),
|
||||
)
|
||||
)
|
||||
return openai_tools
|
||||
|
||||
@staticmethod
|
||||
def _convert_tool_choice_to_openai(
|
||||
tool_choice: ToolChoice,
|
||||
) -> ChatCompletionToolChoiceOptionParam:
|
||||
"""Convert MCP tool_choice to OpenAI format."""
|
||||
if tool_choice.mode == "auto":
|
||||
return "auto"
|
||||
elif tool_choice.mode == "required":
|
||||
return "required"
|
||||
elif tool_choice.mode == "none":
|
||||
return "none"
|
||||
else:
|
||||
raise ValueError(f"Unsupported tool_choice mode: {tool_choice.mode!r}")
|
||||
|
||||
@staticmethod
|
||||
def _chat_completion_to_result_with_tools(
|
||||
chat_completion: ChatCompletion,
|
||||
) -> CreateMessageResultWithTools:
|
||||
"""Convert OpenAI response to CreateMessageResultWithTools."""
|
||||
if len(chat_completion.choices) == 0:
|
||||
raise ValueError("No response for completion")
|
||||
|
||||
first_choice = chat_completion.choices[0]
|
||||
message = first_choice.message
|
||||
|
||||
# Determine stop reason
|
||||
stop_reason: StopReason
|
||||
if first_choice.finish_reason == "tool_calls":
|
||||
stop_reason = "toolUse"
|
||||
elif first_choice.finish_reason == "stop":
|
||||
stop_reason = "endTurn"
|
||||
elif first_choice.finish_reason == "length":
|
||||
stop_reason = "maxTokens"
|
||||
else:
|
||||
stop_reason = "endTurn"
|
||||
|
||||
# Build content list
|
||||
content: list[TextContent | ToolUseContent] = []
|
||||
|
||||
# Add text content if present
|
||||
if message.content:
|
||||
content.append(TextContent(type="text", text=message.content))
|
||||
|
||||
# Add tool calls if present
|
||||
if message.tool_calls:
|
||||
for tool_call in message.tool_calls:
|
||||
# Skip non-function tool calls
|
||||
if not hasattr(tool_call, "function"):
|
||||
continue
|
||||
func = tool_call.function # type: ignore[union-attr]
|
||||
# Parse the arguments JSON string
|
||||
try:
|
||||
arguments = json.loads(func.arguments) # type: ignore[union-attr]
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(
|
||||
f"Invalid JSON in tool arguments for "
|
||||
f"'{func.name}': {func.arguments}" # type: ignore[union-attr]
|
||||
) from e
|
||||
|
||||
content.append(
|
||||
ToolUseContent(
|
||||
type="tool_use",
|
||||
id=tool_call.id,
|
||||
name=func.name, # type: ignore[union-attr]
|
||||
input=arguments,
|
||||
)
|
||||
)
|
||||
|
||||
# Must have at least some content
|
||||
if not content:
|
||||
raise ValueError("No content in response from completion")
|
||||
|
||||
return CreateMessageResultWithTools(
|
||||
content=content, # type: ignore[arg-type]
|
||||
role="assistant",
|
||||
model=chat_completion.model,
|
||||
stopReason=stop_reason,
|
||||
)
|
||||
__all__ = ["OpenAISamplingHandler"]
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
"""Sampling module for FastMCP servers."""
|
||||
|
||||
from fastmcp.server.sampling.handler import ServerSamplingHandler
|
||||
from fastmcp.server.sampling.run import SampleStep, SamplingResult
|
||||
from fastmcp.server.sampling.sampling_tool import SamplingTool
|
||||
|
||||
|
|
@ -8,5 +7,4 @@ __all__ = [
|
|||
"SampleStep",
|
||||
"SamplingResult",
|
||||
"SamplingTool",
|
||||
"ServerSamplingHandler",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,22 +0,0 @@
|
|||
from collections.abc import Awaitable, Callable
|
||||
from typing import TypeAlias
|
||||
|
||||
from mcp import CreateMessageResult
|
||||
from mcp.server.session import ServerSession
|
||||
from mcp.shared.context import LifespanContextT, RequestContext
|
||||
from mcp.types import CreateMessageRequestParams as SamplingParams
|
||||
from mcp.types import CreateMessageResultWithTools, SamplingMessage
|
||||
|
||||
# Result type that handlers can return
|
||||
SamplingHandlerResult: TypeAlias = (
|
||||
str | CreateMessageResult | CreateMessageResultWithTools
|
||||
)
|
||||
|
||||
ServerSamplingHandler: TypeAlias = Callable[
|
||||
[
|
||||
list[SamplingMessage],
|
||||
SamplingParams,
|
||||
RequestContext[ServerSession, LifespanContextT],
|
||||
],
|
||||
SamplingHandlerResult | Awaitable[SamplingHandlerResult],
|
||||
]
|
||||
|
|
@ -94,12 +94,12 @@ from fastmcp.utilities.types import NotSet, NotSetT
|
|||
if TYPE_CHECKING:
|
||||
from fastmcp.client import Client
|
||||
from fastmcp.client.client import FastMCP1Server
|
||||
from fastmcp.client.sampling import SamplingHandler
|
||||
from fastmcp.client.transports import ClientTransport, ClientTransportT
|
||||
from fastmcp.server.openapi import ComponentFn as OpenAPIComponentFn
|
||||
from fastmcp.server.openapi import FastMCPOpenAPI, RouteMap
|
||||
from fastmcp.server.openapi import RouteMapFn as OpenAPIRouteMapFn
|
||||
from fastmcp.server.proxy import FastMCPProxy
|
||||
from fastmcp.server.sampling.handler import ServerSamplingHandler
|
||||
from fastmcp.tools.tool import ToolResultSerializerType
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
|
@ -208,7 +208,7 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
streamable_http_path: str | None = None,
|
||||
json_response: bool | None = None,
|
||||
stateless_http: bool | None = None,
|
||||
sampling_handler: ServerSamplingHandler[LifespanResultT] | None = None,
|
||||
sampling_handler: SamplingHandler | None = None,
|
||||
sampling_handler_behavior: Literal["always", "fallback"] | None = None,
|
||||
):
|
||||
# Resolve server default for background task support
|
||||
|
|
@ -288,9 +288,7 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
# Set up MCP protocol handlers
|
||||
self._setup_handlers()
|
||||
|
||||
self.sampling_handler: ServerSamplingHandler[LifespanResultT] | None = (
|
||||
sampling_handler
|
||||
)
|
||||
self.sampling_handler: SamplingHandler | None = sampling_handler
|
||||
self.sampling_handler_behavior: Literal["always", "fallback"] = (
|
||||
sampling_handler_behavior or "fallback"
|
||||
)
|
||||
|
|
|
|||
0
tests/client/sampling/__init__.py
Normal file
0
tests/client/sampling/__init__.py
Normal file
0
tests/client/sampling/handlers/__init__.py
Normal file
0
tests/client/sampling/handlers/__init__.py
Normal file
|
|
@ -18,7 +18,7 @@ from openai.types.chat import (
|
|||
)
|
||||
from openai.types.chat.chat_completion import Choice
|
||||
|
||||
from fastmcp.experimental.sampling.handlers.openai import OpenAISamplingHandler
|
||||
from fastmcp.client.sampling.handlers.openai import OpenAISamplingHandler
|
||||
|
||||
|
||||
def test_convert_sampling_messages_to_openai_messages():
|
||||
Loading…
Add table
Add a link
Reference in a new issue