mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-20 04:24:17 +02:00
PR Clean-up
This commit is contained in:
parent
d5f01dae15
commit
00fbc97f99
1 changed files with 97 additions and 30 deletions
|
|
@ -1,3 +1,5 @@
|
|||
"""A middleware for response caching."""
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from collections import defaultdict
|
||||
|
|
@ -21,7 +23,7 @@ try:
|
|||
from diskcache import Cache as DiskCacheClient
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"fastmcp[caching] is required to use the caching middleware. Please install it with `pip install fastmcp[caching] or `uv add fastmcp[caching]`"
|
||||
"fastmcp[caching] is required to use the caching middleware. Please install it with `pip install fastmcp[caching]` or `uv add fastmcp[caching]`"
|
||||
)
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
|
@ -49,11 +51,12 @@ CachableTypeVar = TypeVar("CachableTypeVar", bound=CachableTypes)
|
|||
|
||||
|
||||
def make_collection_key(collection: str, key: str) -> str:
|
||||
"""For cache backends that dont support collections, we combine the collection name and key into a single string."""
|
||||
return f"{collection}:{key}"
|
||||
|
||||
|
||||
class CachedPrompt(Prompt):
|
||||
"""A cached prompt."""
|
||||
"""A no-op prompt that can be cached/pickled and provided during list calls."""
|
||||
|
||||
def render(
|
||||
self, arguments: dict[str, Any] | None = None
|
||||
|
|
@ -64,7 +67,7 @@ class CachedPrompt(Prompt):
|
|||
|
||||
|
||||
class CachedResource(Resource):
|
||||
"""A cached resource."""
|
||||
"""A no-op resource that can be cached/pickled and provided during list calls."""
|
||||
|
||||
def read(self) -> str | bytes:
|
||||
raise NotImplementedError(
|
||||
|
|
@ -73,14 +76,14 @@ class CachedResource(Resource):
|
|||
|
||||
|
||||
class CachedTool(Tool):
|
||||
"""A cached tool."""
|
||||
"""A no-op tool that can be cached/pickled and provided during list calls."""
|
||||
|
||||
def run(self, arguments: dict[str, Any]) -> ToolResult:
|
||||
raise NotImplementedError("Run called on CachedTool, this should never happen")
|
||||
|
||||
|
||||
class CacheEntry(BaseModel, Generic[CachableTypeVar]):
|
||||
"""A cache entry."""
|
||||
"""A cache entry helper that can be stored in a cache backend."""
|
||||
|
||||
model_config: ClassVar[ConfigDict] = ConfigDict(
|
||||
frozen=True, arbitrary_types_allowed=True
|
||||
|
|
@ -129,7 +132,7 @@ class CacheProtocol(Protocol):
|
|||
collection: str,
|
||||
key: str,
|
||||
) -> CachableTypes | None:
|
||||
"""Get a value from the cache."""
|
||||
"""Get a value from the cache using the collection and key."""
|
||||
|
||||
if not (cache_entry := await self.get_entry(collection=collection, key=key)):
|
||||
return None
|
||||
|
|
@ -140,7 +143,7 @@ class CacheProtocol(Protocol):
|
|||
self,
|
||||
cache_entry: CacheEntry[CachableTypes],
|
||||
) -> None:
|
||||
"""Set a value in the cache."""
|
||||
"""Set a value in the cache using the collection and key."""
|
||||
|
||||
async def set_value(
|
||||
self,
|
||||
|
|
@ -149,7 +152,7 @@ class CacheProtocol(Protocol):
|
|||
value: CachableTypes,
|
||||
ttl: int,
|
||||
) -> None:
|
||||
"""Set a value in the cache."""
|
||||
"""Set a value in the cache using the collection and key."""
|
||||
|
||||
await self.set_entry(
|
||||
cache_entry=CacheEntry.from_value(
|
||||
|
|
@ -162,7 +165,7 @@ class CacheProtocol(Protocol):
|
|||
collection: str,
|
||||
key: str,
|
||||
) -> None:
|
||||
"""Delete a value from the cache."""
|
||||
"""Delete a value from the cache using the collection and key."""
|
||||
|
||||
def make_collection_key(self, collection: str, key: str) -> str:
|
||||
return f"{collection}:{key}"
|
||||
|
|
@ -222,10 +225,12 @@ DEFAULT_MEMORY_CACHE_MAX_ENTRIES = 1000
|
|||
|
||||
|
||||
def _memory_cache_ttu(_key: Any, value: CacheEntry[CachableTypes], now: float) -> float:
|
||||
"""TTU function for the memory cache. Determines the TTL of the cache entry."""
|
||||
return now + value.ttl
|
||||
|
||||
|
||||
def _memory_cache_getsizeof(value: CacheEntry[CachableTypes]) -> int:
|
||||
"""Getsizeof function for the memory cache. Currently measures how many entries are in the cache."""
|
||||
return 1
|
||||
|
||||
|
||||
|
|
@ -277,34 +282,46 @@ class InMemoryCache(CacheProtocol):
|
|||
class CacheMethodStats(BaseModel):
|
||||
"""Stats for a cache method."""
|
||||
|
||||
hits: int = Field(default=0)
|
||||
misses: int = Field(default=0)
|
||||
too_big: int = Field(default=0)
|
||||
hits: int = Field(default=0, description="The number of hits for the cache method.")
|
||||
misses: int = Field(
|
||||
default=0, description="The number of misses for the cache method."
|
||||
)
|
||||
too_big: int = Field(
|
||||
default=0,
|
||||
description="The number of items that exceeded the size limit for cache entries.",
|
||||
)
|
||||
|
||||
|
||||
class CacheStats(BaseModel):
|
||||
"""Stats for the cache."""
|
||||
|
||||
collections: dict[str, CacheMethodStats] = Field(
|
||||
default_factory=lambda: defaultdict[str, CacheMethodStats](CacheMethodStats)
|
||||
default_factory=lambda: defaultdict(CacheMethodStats),
|
||||
description="Stats are organized by collection (method).",
|
||||
)
|
||||
|
||||
def get_misses(self, collection: str) -> int:
|
||||
"""Get the number of misses for a collection."""
|
||||
return self.collections[collection].misses
|
||||
|
||||
def get_hits(self, collection: str) -> int:
|
||||
"""Get the number of hits for a collection."""
|
||||
return self.collections[collection].hits
|
||||
|
||||
def get_too_big(self, collection: str) -> int:
|
||||
"""Get the number of items that exceeded the size limit for a collection."""
|
||||
return self.collections[collection].too_big
|
||||
|
||||
def mark_miss(self, collection: str) -> None:
|
||||
"""Mark a miss for a collection."""
|
||||
self.collections[collection].misses += 1
|
||||
|
||||
def mark_hit(self, collection: str) -> None:
|
||||
"""Mark a hit for a collection."""
|
||||
self.collections[collection].hits += 1
|
||||
|
||||
def mark_too_big(self, collection: str) -> None:
|
||||
"""Mark a too big for a collection."""
|
||||
self.collections[collection].too_big += 1
|
||||
|
||||
|
||||
|
|
@ -315,34 +332,34 @@ class SharedMethodSettings(TypedDict):
|
|||
|
||||
|
||||
class ListToolsSettings(SharedMethodSettings):
|
||||
pass
|
||||
"""Configuration options for Tool-related caching."""
|
||||
|
||||
|
||||
class ListResourcesSettings(SharedMethodSettings):
|
||||
pass
|
||||
"""Configuration options for Resource-related caching."""
|
||||
|
||||
|
||||
class ListPromptsSettings(SharedMethodSettings):
|
||||
pass
|
||||
"""Configuration options for Prompt-related caching."""
|
||||
|
||||
|
||||
class CallToolSettings(SharedMethodSettings):
|
||||
"""Extra configuration options for Tool-related caching."""
|
||||
"""Configuration options for Tool-related caching."""
|
||||
|
||||
included_tools: NotRequired[list[str]]
|
||||
excluded_tools: NotRequired[list[str]]
|
||||
|
||||
|
||||
class ReadResourceSettings(SharedMethodSettings):
|
||||
pass
|
||||
"""Configuration options for Resource-related caching."""
|
||||
|
||||
|
||||
class GetPromptSettings(SharedMethodSettings):
|
||||
pass
|
||||
"""Configuration options for Prompt-related caching."""
|
||||
|
||||
|
||||
class MethodSettings(TypedDict):
|
||||
"""Config for the response caching middleware methods."""
|
||||
"""Configuration options for mcp "methods" in the response caching middleware."""
|
||||
|
||||
list_tools: NotRequired[ListToolsSettings]
|
||||
call_tool: NotRequired[CallToolSettings]
|
||||
|
|
@ -356,12 +373,6 @@ class MethodSettings(TypedDict):
|
|||
|
||||
MethodSettingsType = TypeVar("MethodSettingsType", bound=SharedMethodSettings)
|
||||
|
||||
NOTIFICATION_TO_COLLECTION = {
|
||||
mcp.types.ToolListChangedNotification: "tools/list",
|
||||
mcp.types.ResourceListChangedNotification: "resources/list",
|
||||
mcp.types.PromptListChangedNotification: "prompts/list",
|
||||
}
|
||||
|
||||
MCP_METHOD_TO_METHOD_SETTINGS_KEY = {
|
||||
"tools/list": "list_tools",
|
||||
"tools/call": "call_tool",
|
||||
|
|
@ -394,11 +405,16 @@ DEFAULT_METHOD_SETTINGS: MethodSettings = MethodSettings(
|
|||
|
||||
|
||||
class ResponseCachingMiddleware(Middleware):
|
||||
"""Caches tool call responses based on method name and params.
|
||||
"""The response caching middleware offers a simple way to cache responses to mcp methods. The Middleware
|
||||
supports cache invalidation via notifications from the server. The Middleware implements TTL-based caching
|
||||
but cache implementations may offer additional features like LRU eviction, size limits, and more.
|
||||
|
||||
When items are retrieved from the cache they will no longer be the original objects, but rather no-op objects
|
||||
this means that response caching may not be compatible with other middleware that expects original subclasses.
|
||||
|
||||
Notes:
|
||||
- Only caches `tools/call` requests.
|
||||
- Cache key derived from tool name and arguments.
|
||||
- Caches `tools/call`, `resources/read`, `prompts/get`, `tools/list`, `resources/list`, and `prompts/list` requests.
|
||||
- Cache keys are derived from method name and arguments.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
|
|
@ -467,6 +483,8 @@ class ResponseCachingMiddleware(Middleware):
|
|||
context: MiddlewareContext[mcp.types.ListResourcesRequest],
|
||||
call_next: CallNext[mcp.types.ListResourcesRequest, list[Resource]],
|
||||
) -> list[Resource]:
|
||||
"""List resources from the cache, if caching is enabled, and the result is in the cache. Otherwise,
|
||||
otherwise call the next middleware and store the result in the cache if caching is enabled."""
|
||||
if self._should_bypass_caching(context=context):
|
||||
return await call_next(context)
|
||||
|
||||
|
|
@ -497,6 +515,8 @@ class ResponseCachingMiddleware(Middleware):
|
|||
context: MiddlewareContext[mcp.types.ListPromptsRequest],
|
||||
call_next: CallNext[mcp.types.ListPromptsRequest, list[Prompt]],
|
||||
) -> list[Prompt]:
|
||||
"""List prompts from the cache, if caching is enabled, and the result is in the cache. Otherwise,
|
||||
otherwise call the next middleware and store the result in the cache if caching is enabled."""
|
||||
if self._should_bypass_caching(context=context):
|
||||
return await call_next(context)
|
||||
|
||||
|
|
@ -531,6 +551,8 @@ class ResponseCachingMiddleware(Middleware):
|
|||
context: MiddlewareContext[mcp.types.CallToolRequestParams],
|
||||
call_next: CallNext[mcp.types.CallToolRequestParams, ToolResult],
|
||||
) -> Any:
|
||||
"""Call a tool from the cache, if caching is enabled, and the result is in the cache. Otherwise,
|
||||
otherwise call the next middleware and store the result in the cache if caching is enabled."""
|
||||
if self._should_bypass_caching(context=context):
|
||||
return await call_next(context=context)
|
||||
|
||||
|
|
@ -550,6 +572,8 @@ class ResponseCachingMiddleware(Middleware):
|
|||
mcp.types.ReadResourceRequestParams, list[ReadResourceContents]
|
||||
],
|
||||
) -> list[ReadResourceContents]:
|
||||
"""Read a resource from the cache, if caching is enabled, and the result is in the cache. Otherwise,
|
||||
otherwise call the next middleware and store the result in the cache if caching is enabled."""
|
||||
if self._should_bypass_caching(context=context):
|
||||
return await call_next(context=context)
|
||||
|
||||
|
|
@ -566,6 +590,8 @@ class ResponseCachingMiddleware(Middleware):
|
|||
mcp.types.GetPromptRequestParams, mcp.types.GetPromptResult
|
||||
],
|
||||
) -> mcp.types.GetPromptResult:
|
||||
"""Get a prompt from the cache, if caching is enabled, and the result is in the cache. Otherwise,
|
||||
otherwise call the next middleware and store the result in the cache if caching is enabled."""
|
||||
if self._should_bypass_caching(context=context):
|
||||
return await call_next(context)
|
||||
|
||||
|
|
@ -580,7 +606,18 @@ class ResponseCachingMiddleware(Middleware):
|
|||
context: MiddlewareContext[mcp.types.Notification],
|
||||
call_next: CallNext[mcp.types.Notification, Any],
|
||||
) -> Any:
|
||||
if collection := NOTIFICATION_TO_COLLECTION.get(context.message):
|
||||
"""Handle a notification from the server. If the notification is a tool/resource/prompt list changed
|
||||
notification, delete the cache for the affected method."""
|
||||
if isinstance(context.message, mcp.types.ToolListChangedNotification):
|
||||
collection = "tools/list"
|
||||
elif isinstance(context.message, mcp.types.ResourceListChangedNotification):
|
||||
collection = "resources/list"
|
||||
elif isinstance(context.message, mcp.types.PromptListChangedNotification):
|
||||
collection = "prompts/list"
|
||||
else:
|
||||
collection = None
|
||||
|
||||
if collection:
|
||||
await self._backend.delete(collection=collection, key=GLOBAL_KEY)
|
||||
|
||||
return await call_next(context)
|
||||
|
|
@ -591,6 +628,9 @@ class ResponseCachingMiddleware(Middleware):
|
|||
call_next: CallNext[Any, CachableTypeVar],
|
||||
key: str | None = None,
|
||||
) -> CachableTypeVar:
|
||||
"""Perform the cached lookup, if the result is not in the cache, call the next middleware and return
|
||||
the result."""
|
||||
|
||||
if key is None:
|
||||
key = GLOBAL_KEY
|
||||
|
||||
|
|
@ -615,6 +655,8 @@ class ResponseCachingMiddleware(Middleware):
|
|||
call_next: CallNext[Any, CachableTypeVar],
|
||||
key: str | None = None,
|
||||
) -> CachableTypeVar | None:
|
||||
"""Get a value from the cache and update the cache stats."""
|
||||
|
||||
if key is None:
|
||||
key = GLOBAL_KEY
|
||||
|
||||
|
|
@ -638,6 +680,8 @@ class ResponseCachingMiddleware(Middleware):
|
|||
key: str | None,
|
||||
value: CachableTypeVar,
|
||||
) -> CachableTypeVar:
|
||||
"""Store a value in the cache (if it's not too big) with the appropriate TTL."""
|
||||
|
||||
if key is None:
|
||||
key = GLOBAL_KEY
|
||||
|
||||
|
|
@ -664,6 +708,8 @@ class ResponseCachingMiddleware(Middleware):
|
|||
def _matches_tool_cache_settings(
|
||||
self, context: MiddlewareContext[mcp.types.CallToolRequestParams]
|
||||
) -> bool:
|
||||
"""Check if the tool matches the cache settings for tool calls."""
|
||||
|
||||
tool_name = context.message.name
|
||||
|
||||
tool_call_cache_settings: CallToolSettings | None = self._get_cache_settings(
|
||||
|
|
@ -689,6 +735,8 @@ class ResponseCachingMiddleware(Middleware):
|
|||
context: MiddlewareContext[Any],
|
||||
settings_type: type[MethodSettingsType] = SharedMethodSettings,
|
||||
) -> MethodSettingsType | None:
|
||||
"""Get the cache settings for a method."""
|
||||
|
||||
if not context.method:
|
||||
return None
|
||||
|
||||
|
|
@ -705,6 +753,8 @@ class ResponseCachingMiddleware(Middleware):
|
|||
return cast(MethodSettingsType, self.method_settings[method_settings_key])
|
||||
|
||||
def _get_cache_ttl(self, context: MiddlewareContext[Any]) -> int:
|
||||
"""Get the cache TTL for a method."""
|
||||
|
||||
settings: SharedMethodSettings | None = self._get_cache_settings(
|
||||
context=context
|
||||
)
|
||||
|
|
@ -715,6 +765,8 @@ class ResponseCachingMiddleware(Middleware):
|
|||
return settings["ttl"]
|
||||
|
||||
def _should_bypass_caching(self, context: MiddlewareContext[Any]) -> bool:
|
||||
"""Check if the method should bypass caching."""
|
||||
|
||||
if not self._get_cache_settings(context=context):
|
||||
return True
|
||||
|
||||
|
|
@ -722,21 +774,29 @@ class ResponseCachingMiddleware(Middleware):
|
|||
|
||||
|
||||
def _make_call_tool_cache_key(msg: mcp.types.CallToolRequestParams) -> str:
|
||||
"""Make a cache key for a tool call by hashing the tool name and its arguments."""
|
||||
|
||||
raw = f"{msg.name}:{_get_arguments_str(msg.arguments)}"
|
||||
return hashlib.sha256(raw.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def _make_read_resource_cache_key(msg: mcp.types.ReadResourceRequestParams) -> str:
|
||||
"""Make a cache key for a resource read by hashing the resource URI."""
|
||||
|
||||
raw = f"{msg.uri}"
|
||||
return hashlib.sha256(raw.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def _make_get_prompt_cache_key(msg: mcp.types.GetPromptRequestParams) -> str:
|
||||
"""Make a cache key for a prompt get by hashing the prompt name and its arguments."""
|
||||
|
||||
raw = f"{msg.name}:{_get_arguments_str(msg.arguments)}"
|
||||
return hashlib.sha256(raw.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def _get_arguments_str(arguments: dict[str, Any] | None) -> str:
|
||||
"""Get a string representation of the arguments."""
|
||||
|
||||
if arguments is None:
|
||||
return "null"
|
||||
|
||||
|
|
@ -750,6 +810,8 @@ def _get_arguments_str(arguments: dict[str, Any] | None) -> str:
|
|||
def get_size_of_content_blocks(
|
||||
value: mcp.types.ContentBlock | Sequence[mcp.types.ContentBlock],
|
||||
) -> int:
|
||||
"""Get the size of a series of content blocks by summing the size of the JSON representation of each block."""
|
||||
|
||||
if isinstance(value, mcp.types.ContentBlock):
|
||||
value = [value]
|
||||
|
||||
|
|
@ -757,6 +819,8 @@ def get_size_of_content_blocks(
|
|||
|
||||
|
||||
def get_size_of_tool_result(value: ToolResult) -> int:
|
||||
"""Get the size of a tool result by summing the size of the content blocks and the size of the structured content."""
|
||||
|
||||
content_size = get_size_of_content_blocks(value.content)
|
||||
structured_content_size = len(
|
||||
json.dumps(
|
||||
|
|
@ -768,6 +832,8 @@ def get_size_of_tool_result(value: ToolResult) -> int:
|
|||
|
||||
|
||||
def get_size_of_one_value(value: BaseModel | ToolResult | ReadResourceContents) -> int:
|
||||
"""Get the size of an mcp type."""
|
||||
|
||||
if isinstance(value, ToolResult):
|
||||
return get_size_of_tool_result(value)
|
||||
if isinstance(value, ReadResourceContents):
|
||||
|
|
@ -776,6 +842,7 @@ def get_size_of_one_value(value: BaseModel | ToolResult | ReadResourceContents)
|
|||
|
||||
|
||||
def get_size_of_value(value: CachableTypes) -> int:
|
||||
"""Get the size of a cache entry."""
|
||||
if isinstance(value, BaseModel | ToolResult | ReadResourceContents):
|
||||
return get_size_of_one_value(value)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue