mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-23 05:54:19 +02:00
More progress
This commit is contained in:
parent
84b3e0ff60
commit
c600a16899
2 changed files with 155 additions and 57 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue