mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-21 21:14:17 +02:00
Use custom cacheable models in middleware to preserve key field
Instead of adding custom serializers to FastMCPComponent (which affects equality globally), use dedicated cacheable models in the middleware. This approach: - Isolates caching concerns to the middleware - Follows existing patterns (CachableReadResourceContents, CachableToolResult) - Preserves key field through serialization without side effects - Maintains architectural separation Co-authored-by: William Easton <strawgate@users.noreply.github.com>
This commit is contained in:
parent
d970e8a4aa
commit
c4b53ce141
2 changed files with 155 additions and 102 deletions
|
|
@ -74,6 +74,132 @@ class CachableToolResult(BaseModel):
|
|||
)
|
||||
|
||||
|
||||
class CachableTool(BaseModel):
|
||||
"""A wrapper for Tool that can be cached, preserving the key field."""
|
||||
|
||||
name: str
|
||||
key: str
|
||||
title: str | None
|
||||
description: str | None
|
||||
parameters: dict[str, Any]
|
||||
output_schema: dict[str, Any] | None
|
||||
annotations: Any | None
|
||||
meta: dict[str, Any] | None
|
||||
tags: set[str]
|
||||
enabled: bool
|
||||
|
||||
@classmethod
|
||||
def wrap(cls, tool: Tool) -> Self:
|
||||
return cls(
|
||||
name=tool.name,
|
||||
key=tool.key,
|
||||
title=tool.title,
|
||||
description=tool.description,
|
||||
parameters=tool.parameters,
|
||||
output_schema=tool.output_schema,
|
||||
annotations=tool.annotations,
|
||||
meta=tool.meta,
|
||||
tags=tool.tags,
|
||||
enabled=tool.enabled,
|
||||
)
|
||||
|
||||
def unwrap(self) -> Tool:
|
||||
return Tool(
|
||||
name=self.name,
|
||||
key=self.key,
|
||||
title=self.title,
|
||||
description=self.description,
|
||||
parameters=self.parameters,
|
||||
output_schema=self.output_schema,
|
||||
annotations=self.annotations,
|
||||
meta=self.meta,
|
||||
tags=self.tags,
|
||||
enabled=self.enabled,
|
||||
)
|
||||
|
||||
|
||||
class CachableResource(BaseModel):
|
||||
"""A wrapper for Resource that can be cached, preserving the key field."""
|
||||
|
||||
name: str
|
||||
key: str
|
||||
title: str | None
|
||||
description: str | None
|
||||
uri: str
|
||||
mime_type: str
|
||||
annotations: Any | None
|
||||
meta: dict[str, Any] | None
|
||||
tags: set[str]
|
||||
enabled: bool
|
||||
|
||||
@classmethod
|
||||
def wrap(cls, resource: Resource) -> Self:
|
||||
return cls(
|
||||
name=resource.name,
|
||||
key=resource.key,
|
||||
title=resource.title,
|
||||
description=resource.description,
|
||||
uri=str(resource.uri),
|
||||
mime_type=resource.mime_type,
|
||||
annotations=resource.annotations,
|
||||
meta=resource.meta,
|
||||
tags=resource.tags,
|
||||
enabled=resource.enabled,
|
||||
)
|
||||
|
||||
def unwrap(self) -> Resource:
|
||||
return Resource(
|
||||
name=self.name,
|
||||
key=self.key,
|
||||
title=self.title,
|
||||
description=self.description,
|
||||
uri=self.uri,
|
||||
mime_type=self.mime_type,
|
||||
annotations=self.annotations,
|
||||
meta=self.meta,
|
||||
tags=self.tags,
|
||||
enabled=self.enabled,
|
||||
)
|
||||
|
||||
|
||||
class CachablePrompt(BaseModel):
|
||||
"""A wrapper for Prompt that can be cached, preserving the key field."""
|
||||
|
||||
name: str
|
||||
key: str
|
||||
title: str | None
|
||||
description: str | None
|
||||
arguments: list[Any] | None
|
||||
meta: dict[str, Any] | None
|
||||
tags: set[str]
|
||||
enabled: bool
|
||||
|
||||
@classmethod
|
||||
def wrap(cls, prompt: Prompt) -> Self:
|
||||
return cls(
|
||||
name=prompt.name,
|
||||
key=prompt.key,
|
||||
title=prompt.title,
|
||||
description=prompt.description,
|
||||
arguments=prompt.arguments,
|
||||
meta=prompt.meta,
|
||||
tags=prompt.tags,
|
||||
enabled=prompt.enabled,
|
||||
)
|
||||
|
||||
def unwrap(self) -> Prompt:
|
||||
return Prompt(
|
||||
name=self.name,
|
||||
key=self.key,
|
||||
title=self.title,
|
||||
description=self.description,
|
||||
arguments=self.arguments,
|
||||
meta=self.meta,
|
||||
tags=self.tags,
|
||||
enabled=self.enabled,
|
||||
)
|
||||
|
||||
|
||||
class SharedMethodSettings(TypedDict):
|
||||
"""Shared config for a cache method."""
|
||||
|
||||
|
|
@ -182,22 +308,26 @@ class ResponseCachingMiddleware(Middleware):
|
|||
call_tool_settings or CallToolSettings()
|
||||
)
|
||||
|
||||
self._list_tools_cache: PydanticAdapter[list[Tool]] = PydanticAdapter(
|
||||
self._list_tools_cache: PydanticAdapter[list[CachableTool]] = PydanticAdapter(
|
||||
key_value=self._stats,
|
||||
pydantic_model=list[Tool],
|
||||
pydantic_model=list[CachableTool],
|
||||
default_collection="tools/list",
|
||||
)
|
||||
|
||||
self._list_resources_cache: PydanticAdapter[list[Resource]] = PydanticAdapter(
|
||||
key_value=self._stats,
|
||||
pydantic_model=list[Resource],
|
||||
default_collection="resources/list",
|
||||
self._list_resources_cache: PydanticAdapter[list[CachableResource]] = (
|
||||
PydanticAdapter(
|
||||
key_value=self._stats,
|
||||
pydantic_model=list[CachableResource],
|
||||
default_collection="resources/list",
|
||||
)
|
||||
)
|
||||
|
||||
self._list_prompts_cache: PydanticAdapter[list[Prompt]] = PydanticAdapter(
|
||||
key_value=self._stats,
|
||||
pydantic_model=list[Prompt],
|
||||
default_collection="prompts/list",
|
||||
self._list_prompts_cache: PydanticAdapter[list[CachablePrompt]] = (
|
||||
PydanticAdapter(
|
||||
key_value=self._stats,
|
||||
pydantic_model=list[CachablePrompt],
|
||||
default_collection="prompts/list",
|
||||
)
|
||||
)
|
||||
|
||||
self._read_resource_cache: PydanticAdapter[
|
||||
|
|
@ -234,26 +364,12 @@ class ResponseCachingMiddleware(Middleware):
|
|||
return await call_next(context)
|
||||
|
||||
if cached_value := await self._list_tools_cache.get(key=GLOBAL_KEY):
|
||||
return cached_value
|
||||
return [item.unwrap() for item in cached_value]
|
||||
|
||||
tools: Sequence[Tool] = await call_next(context=context)
|
||||
|
||||
# Turn any subclass of Tool into a Tool
|
||||
cachable_tools: list[Tool] = [
|
||||
Tool(
|
||||
name=tool.name,
|
||||
key=tool.key,
|
||||
title=tool.title,
|
||||
description=tool.description,
|
||||
parameters=tool.parameters,
|
||||
output_schema=tool.output_schema,
|
||||
annotations=tool.annotations,
|
||||
meta=tool.meta,
|
||||
tags=tool.tags,
|
||||
enabled=tool.enabled,
|
||||
)
|
||||
for tool in tools
|
||||
]
|
||||
# Wrap tools in cacheable models
|
||||
cachable_tools: list[CachableTool] = [CachableTool.wrap(tool) for tool in tools]
|
||||
|
||||
await self._list_tools_cache.put(
|
||||
key=GLOBAL_KEY,
|
||||
|
|
@ -261,7 +377,7 @@ class ResponseCachingMiddleware(Middleware):
|
|||
ttl=self._list_tools_settings.get("ttl", FIVE_MINUTES_IN_SECONDS),
|
||||
)
|
||||
|
||||
return cachable_tools
|
||||
return [item.unwrap() for item in cachable_tools]
|
||||
|
||||
@override
|
||||
async def on_list_resources(
|
||||
|
|
@ -275,25 +391,13 @@ class ResponseCachingMiddleware(Middleware):
|
|||
return await call_next(context)
|
||||
|
||||
if cached_value := await self._list_resources_cache.get(key=GLOBAL_KEY):
|
||||
return cached_value
|
||||
return [item.unwrap() for item in cached_value]
|
||||
|
||||
resources: Sequence[Resource] = await call_next(context=context)
|
||||
|
||||
# Turn any subclass of Resource into a Resource
|
||||
cachable_resources: list[Resource] = [
|
||||
Resource(
|
||||
name=resource.name,
|
||||
key=resource.key,
|
||||
title=resource.title,
|
||||
description=resource.description,
|
||||
tags=resource.tags,
|
||||
meta=resource.meta,
|
||||
mime_type=resource.mime_type,
|
||||
annotations=resource.annotations,
|
||||
enabled=resource.enabled,
|
||||
uri=resource.uri,
|
||||
)
|
||||
for resource in resources
|
||||
# Wrap resources in cacheable models
|
||||
cachable_resources: list[CachableResource] = [
|
||||
CachableResource.wrap(resource) for resource in resources
|
||||
]
|
||||
|
||||
await self._list_resources_cache.put(
|
||||
|
|
@ -302,7 +406,7 @@ class ResponseCachingMiddleware(Middleware):
|
|||
ttl=self._list_resources_settings.get("ttl", FIVE_MINUTES_IN_SECONDS),
|
||||
)
|
||||
|
||||
return cachable_resources
|
||||
return [item.unwrap() for item in cachable_resources]
|
||||
|
||||
@override
|
||||
async def on_list_prompts(
|
||||
|
|
@ -316,23 +420,13 @@ class ResponseCachingMiddleware(Middleware):
|
|||
return await call_next(context)
|
||||
|
||||
if cached_value := await self._list_prompts_cache.get(key=GLOBAL_KEY):
|
||||
return cached_value
|
||||
return [item.unwrap() for item in cached_value]
|
||||
|
||||
prompts: Sequence[Prompt] = await call_next(context=context)
|
||||
|
||||
# Turn any subclass of Prompt into a Prompt
|
||||
cachable_prompts: list[Prompt] = [
|
||||
Prompt(
|
||||
name=prompt.name,
|
||||
key=prompt.key,
|
||||
title=prompt.title,
|
||||
description=prompt.description,
|
||||
tags=prompt.tags,
|
||||
meta=prompt.meta,
|
||||
enabled=prompt.enabled,
|
||||
arguments=prompt.arguments,
|
||||
)
|
||||
for prompt in prompts
|
||||
# Wrap prompts in cacheable models
|
||||
cachable_prompts: list[CachablePrompt] = [
|
||||
CachablePrompt.wrap(prompt) for prompt in prompts
|
||||
]
|
||||
|
||||
await self._list_prompts_cache.put(
|
||||
|
|
@ -341,7 +435,7 @@ class ResponseCachingMiddleware(Middleware):
|
|||
ttl=self._list_prompts_settings.get("ttl", FIVE_MINUTES_IN_SECONDS),
|
||||
)
|
||||
|
||||
return cachable_prompts
|
||||
return [item.unwrap() for item in cachable_prompts]
|
||||
|
||||
@override
|
||||
async def on_call_tool(
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ from collections.abc import Sequence
|
|||
from typing import Annotated, Any, TypedDict
|
||||
|
||||
from mcp.types import Icon
|
||||
from pydantic import BeforeValidator, Field, PrivateAttr, model_serializer
|
||||
from pydantic import BeforeValidator, Field, PrivateAttr
|
||||
from typing_extensions import Self, TypeVar
|
||||
|
||||
import fastmcp
|
||||
|
|
@ -119,47 +119,6 @@ class FastMCPComponent(FastMCPBaseModel):
|
|||
copy._key = key
|
||||
return copy
|
||||
|
||||
@model_serializer(mode="wrap")
|
||||
def _serialize_model(self, serializer: Any, info: Any) -> dict[str, Any]:
|
||||
"""Custom serializer to include the key field."""
|
||||
data = serializer(self)
|
||||
# Include _key in serialization if it's set
|
||||
if self._key is not None:
|
||||
data["key"] = self._key
|
||||
return data
|
||||
|
||||
@classmethod
|
||||
def model_validate(
|
||||
cls,
|
||||
obj: Any,
|
||||
*,
|
||||
strict: bool | None = None,
|
||||
from_attributes: bool | None = None,
|
||||
context: dict[str, Any] | None = None,
|
||||
) -> Self:
|
||||
"""Validate and create a model instance, handling the key attribute."""
|
||||
if isinstance(obj, dict):
|
||||
# Extract key from dict if present (don't mutate original dict)
|
||||
key = obj.get("key")
|
||||
if key is not None:
|
||||
# Create a copy without the key field
|
||||
obj = {k: v for k, v in obj.items() if k != "key"}
|
||||
instance = super().model_validate(
|
||||
obj,
|
||||
strict=strict,
|
||||
from_attributes=from_attributes,
|
||||
context=context,
|
||||
)
|
||||
if key is not None:
|
||||
instance._key = key
|
||||
return instance
|
||||
return super().model_validate(
|
||||
obj,
|
||||
strict=strict,
|
||||
from_attributes=from_attributes,
|
||||
context=context,
|
||||
)
|
||||
|
||||
def __eq__(self, other: object) -> bool:
|
||||
if type(self) is not type(other):
|
||||
return False
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue