PR Clean-up

This commit is contained in:
William Easton 2025-09-18 08:20:20 -05:00
commit 00fbc97f99
No known key found for this signature in database

View file

@ -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)