From c600a168997fc91842db022bba1763e86bb51978 Mon Sep 17 00:00:00 2001 From: William Easton Date: Wed, 17 Sep 2025 19:36:14 -0500 Subject: [PATCH] More progress --- src/fastmcp/server/middleware/caching.py | 124 +++++++++++++++-------- tests/server/middleware/test_caching.py | 94 +++++++++++++---- 2 files changed, 158 insertions(+), 60 deletions(-) diff --git a/src/fastmcp/server/middleware/caching.py b/src/fastmcp/server/middleware/caching.py index 0b30fcdcf..7cb285393 100644 --- a/src/fastmcp/server/middleware/caching.py +++ b/src/fastmcp/server/middleware/caching.py @@ -1,6 +1,7 @@ import hashlib import json from collections import defaultdict +from collections.abc import Sequence from datetime import datetime, timedelta, timezone from typing import Any, ClassVar, Generic, Protocol, TypedDict, TypeVar, cast @@ -37,10 +38,10 @@ GLOBAL_KEY = "__global__" CachableTypes = ( ToolResult - | list[Tool] - | list[Resource] - | list[Prompt] - | list[ReadResourceContents] + | Sequence[Tool] + | Sequence[Resource] + | Sequence[Prompt] + | Sequence[ReadResourceContents] | mcp.types.GetPromptResult ) @@ -71,6 +72,13 @@ class CachedResource(Resource): ) +class CachedTool(Tool): + """A cached tool.""" + + 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.""" @@ -328,6 +336,12 @@ 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", @@ -401,9 +415,7 @@ class ResponseCachingMiddleware(Middleware): return await call_next(context=context) if cached_value := await self._get_cache( - context=context, - call_next=call_next, - key=None, + context=context, call_next=call_next, key=None ): return cached_value @@ -411,7 +423,7 @@ class ResponseCachingMiddleware(Middleware): # Convert tool subclasses to Tool objects result = [ - Tool( + CachedTool( name=tool.name, title=tool.title, description=tool.description, @@ -508,7 +520,7 @@ class ResponseCachingMiddleware(Middleware): return await self._cached_call_next( context=context, call_next=call_next, - key=self._make_cache_key(msg=context.message), + key=_make_call_tool_cache_key(msg=context.message), ) async def on_read_resource( @@ -524,6 +536,7 @@ class ResponseCachingMiddleware(Middleware): return await self._cached_call_next( context=context, call_next=call_next, + key=_make_read_resource_cache_key(msg=context.message), ) async def on_get_prompt( @@ -539,7 +552,7 @@ class ResponseCachingMiddleware(Middleware): return await self._cached_call_next( context=context, call_next=call_next, - key=None, + key=_make_get_prompt_cache_key(msg=context.message), ) async def on_notification( @@ -547,8 +560,8 @@ class ResponseCachingMiddleware(Middleware): context: MiddlewareContext[mcp.types.Notification], call_next: CallNext[mcp.types.Notification, Any], ) -> Any: - if isinstance(context.message, mcp.types.ToolListChangedNotification): - await self._backend.delete(collection="tools/list", key=GLOBAL_KEY) + if collection := NOTIFICATION_TO_COLLECTION.get(context.message): + await self._backend.delete(collection=collection, key=GLOBAL_KEY) return await call_next(context) @@ -613,12 +626,7 @@ class ResponseCachingMiddleware(Middleware): return value if self._max_item_size is not None: - size = 0 - - for item in dump_if_base_model(value): - size += len(item.encode("utf-8")) - - if size > self._max_item_size: + if get_size_of_value(value=value) > self._max_item_size: self._stats.mark_too_big(collection=collection) return value @@ -692,33 +700,63 @@ class ResponseCachingMiddleware(Middleware): return False - def _make_cache_key(self, msg: mcp.types.CallToolRequestParams) -> str: - raw = f"{self._get_tool_key(msg)}:{self._get_tool_arguments_str(msg)}" - return hashlib.sha256(raw.encode("utf-8")).hexdigest() - def _get_tool_key(self, msg: mcp.types.CallToolRequestParams) -> str: - return msg.name - - def _get_tool_arguments_str(self, msg: mcp.types.CallToolRequestParams) -> str: - if msg.arguments is None: - return "null" - - try: - return json.dumps(msg.arguments, sort_keys=True, separators=(",", ":")) - - except TypeError: - return repr(msg.arguments) +def _make_call_tool_cache_key(msg: mcp.types.CallToolRequestParams) -> str: + raw = f"{msg.name}:{_get_arguments_str(msg.arguments)}" + return hashlib.sha256(raw.encode("utf-8")).hexdigest() -def dump_if_base_model(value: Any) -> list[str]: - if isinstance(value, BaseModel): - return [value.model_dump_json()] +def _make_read_resource_cache_key(msg: mcp.types.ReadResourceRequestParams) -> str: + raw = f"{msg.uri}" + return hashlib.sha256(raw.encode("utf-8")).hexdigest() - if isinstance(value, list): - return [ - item - for sublist in [dump_if_base_model(val) for val in value] - for item in sublist - ] - return [json.dumps(value, sort_keys=True, separators=(",", ":"))] +def _make_get_prompt_cache_key(msg: mcp.types.GetPromptRequestParams) -> str: + 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: + if arguments is None: + return "null" + + try: + return json.dumps(arguments, sort_keys=True, separators=(",", ":")) + + except TypeError: + return repr(arguments) + + +def get_size_of_content_blocks( + value: mcp.types.ContentBlock | Sequence[mcp.types.ContentBlock], +) -> int: + if isinstance(value, mcp.types.ContentBlock): + value = [value] + + return sum([len(item.model_dump_json()) for item in value]) + + +def get_size_of_tool_result(value: ToolResult) -> int: + content_size = get_size_of_content_blocks(value.content) + structured_content_size = len( + json.dumps( + value.structured_content, sort_keys=True, separators=(",", ":") + ).encode("utf-8") + ) + + return content_size + structured_content_size + + +def get_size_of_one_value(value: BaseModel | ToolResult | ReadResourceContents) -> int: + if isinstance(value, ToolResult): + return get_size_of_tool_result(value) + if isinstance(value, ReadResourceContents): + return len(value.content) + return len(value.model_dump_json()) + + +def get_size_of_value(value: CachableTypes) -> int: + if isinstance(value, BaseModel | ToolResult | ReadResourceContents): + return get_size_of_one_value(value) + + return sum(get_size_of_one_value(item) for item in value) diff --git a/tests/server/middleware/test_caching.py b/tests/server/middleware/test_caching.py index 37bd333d5..4da0067bc 100644 --- a/tests/server/middleware/test_caching.py +++ b/tests/server/middleware/test_caching.py @@ -21,6 +21,7 @@ from fastmcp.client import Client from fastmcp.client.client import CallToolResult from fastmcp.client.transports import FastMCPTransport from fastmcp.prompts.prompt import FunctionPrompt +from fastmcp.resources.resource import Resource from fastmcp.server.middleware.caching import ( CachableTypes, CachedPrompt, @@ -58,6 +59,10 @@ SAMPLE_TOOL_RESULT = ToolResult( content=[TextContent(type="text", text="test_text")], structured_content={"result": "test_result"}, ) +SAMPLE_TOOL_RESULT_LARGE = ToolResult( + content=[TextContent(type="text", text="test_text" * 100)], + structured_content={"result": "test_result"}, +) class CrazyModel(BaseModel): @@ -99,7 +104,7 @@ def dump_mcp_types( | ToolResult | Sequence[BaseModel] | Sequence[ToolResult] - | list[ReadResourceContents], + | Sequence[ReadResourceContents], ) -> list[dict[str, Any]]: if isinstance(model, Sequence): return [dump_mcp_type(model=m) for m in model] @@ -171,18 +176,26 @@ class TrackingCalculator: ) def add_resources(self, fastmcp: FastMCP, prefix: str = ""): - fastmcp.add_resource_fn( - fn=self.get_add_calls, uri="resource://add_calls", name=f"{prefix}add_calls" + fastmcp.add_resource( + resource=Resource.from_function( + fn=self.get_add_calls, + uri="resource://add_calls", + name=f"{prefix}add_calls", + ) ) - fastmcp.add_resource_fn( - fn=self.get_multiply_calls, - uri="resource://multiply_calls", - name=f"{prefix}multiply_calls", + fastmcp.add_resource( + resource=Resource.from_function( + fn=self.get_multiply_calls, + uri="resource://multiply_calls", + name=f"{prefix}multiply_calls", + ) ) - fastmcp.add_resource_fn( - fn=self.get_crazy_calls, - uri="resource://crazy_calls", - name=f"{prefix}crazy_calls", + fastmcp.add_resource( + resource=Resource.from_function( + fn=self.get_crazy_calls, + uri="resource://crazy_calls", + name=f"{prefix}crazy_calls", + ) ) @@ -357,7 +370,9 @@ class TestCacheImplementations: value=value, ttl=3600, ) - result = await cache.get_value(collection="test_collection", key="test_key") + result: CachableTypes | None = await cache.get_value( + collection="test_collection", key="test_key" + ) assert result is not None @@ -373,7 +388,9 @@ class TestCacheImplementations: value=SAMPLE_TOOL_RESULT, ttl=3600, ) - result = await cache.get_value(collection="test_collection", key="test_key") + result: CachableTypes | None = await cache.get_value( + collection="test_collection", key="test_key" + ) assert result is not None assert dump_mcp_types(model=result) == dump_mcp_types(model=SAMPLE_TOOL_RESULT) @@ -471,20 +488,63 @@ class TestResponseCachingMiddleware: is result ) + async def test_large_value(self): + """Test that we can set and get a large value.""" + cache = InMemoryCache() + middleware = ResponseCachingMiddleware(cache, max_item_size=100) + + result = await middleware._store_in_cache_and_return( + context=MiddlewareContext( + method="tools/call", + message=mcp.types.CallToolRequestParams(name="test_tool"), + ), + key="test_key", + value=SAMPLE_TOOL_RESULT_LARGE, + ) + + assert middleware._stats.get_too_big("tools/call") == 1 + def test_cache_key_generation(self): """Test cache key generation.""" - cache = InMemoryCache() - middleware = ResponseCachingMiddleware(cache) + from fastmcp.server.middleware.caching import ( + _make_call_tool_cache_key, + _make_get_prompt_cache_key, + _make_read_resource_cache_key, + ) msg = mcp.types.CallToolRequestParams( name="test_tool", arguments={"param1": "value1", "param2": 42} ) - key = middleware._make_cache_key(msg) + key = _make_call_tool_cache_key(msg) # Should be a SHA256 hash assert len(key) == 64 - assert all(c in "0123456789abcdef" for c in key) + assert key == snapshot( + "7fa3d5c7967a202457eeca0731709fd87ec98546ddaee829ee86ca54b1858c59" + ) + + msg = mcp.types.ReadResourceRequestParams( + uri=AnyUrl("https://test_uri"), + ) + + key = _make_read_resource_cache_key(msg) + + assert len(key) == 64 + assert key == snapshot( + "e34cc47c03ed1ad54f02501d95ecc463b65646568961c97ca4b730cb274e9d42" + ) + + msg = mcp.types.GetPromptRequestParams( + name="test_prompt", arguments={"param1": "value1"} + ) + + key = _make_get_prompt_cache_key(msg) + + assert len(key) == 64 + assert key == snapshot( + "6306ff84fd3ff247a4bd91271e9d727d7f051bba53fb2e3bf80958988c4baf57" + ) async def test_cache_miss_and_hit( self,