More progress

This commit is contained in:
William Easton 2025-09-17 19:36:14 -05:00
commit c600a16899
No known key found for this signature in database
2 changed files with 155 additions and 57 deletions

View file

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

View file

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