mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-21 04:54:17 +02:00
Update response caching PR
This commit is contained in:
parent
28370827dc
commit
57cb83a145
8 changed files with 348 additions and 1275 deletions
|
|
@ -1,155 +0,0 @@
|
|||
# Elasticsearch Cache Backend – Contrib Module for FastMCP
|
||||
|
||||
This backend plugs into `ResponseCachingMiddleware` and stores cache entries in Elasticsearch.
|
||||
|
||||
- **Package**: `fastmcp.contrib.middleware.caching.elasticsearch`
|
||||
- **Class**: `ElasticsearchCache`
|
||||
- **Works with**: `fastmcp.server.middleware.caching.ResponseCachingMiddleware`
|
||||
|
||||
---
|
||||
|
||||
## 📦 Installation
|
||||
|
||||
Install Elasticsearch async client and FastMCP caching extras:
|
||||
|
||||
```bash
|
||||
uv add elasticsearch fastmcp[contrib-middleware-elasticsearch_cache]
|
||||
# or
|
||||
pip install elasticsearch fastmcp[contrib-middleware-elasticsearch_cache]
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 🚀 Quick Start
|
||||
|
||||
```python
|
||||
from elasticsearch import AsyncElasticsearch
|
||||
from fastmcp import FastMCP
|
||||
from fastmcp.server.middleware.caching import ResponseCachingMiddleware
|
||||
from fastmcp.contrib.middleware.caching.elasticsearch import ElasticsearchCache
|
||||
|
||||
mcp = FastMCP(name="Cache Demo")
|
||||
|
||||
es = AsyncElasticsearch(hosts=["http://localhost:9200"]) # configure as needed
|
||||
cache = ElasticsearchCache(elasticsearch_client=es, index="fastmcp-response-cache")
|
||||
|
||||
mcp.add_middleware(ResponseCachingMiddleware(cache_backend=cache))
|
||||
|
||||
# ... define tools/resources/prompts ...
|
||||
|
||||
# When using FastMCP.run(), ensure the event loop closes the ES client when done
|
||||
# await es.close() in your shutdown path if you manage the loop yourself.
|
||||
```
|
||||
|
||||
This enables caching for default MCP methods per `ResponseCachingMiddleware`:
|
||||
- **tools**: `tools/call`, `tools/list`
|
||||
- **resources**: `resources/read`, `resources/list`
|
||||
- **prompts**: `prompts/get`, `prompts/list`
|
||||
|
||||
---
|
||||
|
||||
## ⚙️ Configuration
|
||||
|
||||
Constructor:
|
||||
|
||||
```python
|
||||
ElasticsearchCache(
|
||||
elasticsearch_client: AsyncElasticsearch,
|
||||
index: str | None = None,
|
||||
mapping: dict[str, object] | None = None,
|
||||
)
|
||||
```
|
||||
|
||||
- **elasticsearch_client**: A live `AsyncElasticsearch` instance
|
||||
- **index**: Target index name. Default: `"fastmcp-response-cache"`
|
||||
- **mapping**: Custom index mapping. Default mapping is optimized for cache fields:
|
||||
|
||||
```json
|
||||
{
|
||||
"properties": {
|
||||
"created_at": {"type": "date"},
|
||||
"expires_at": {"type": "date"},
|
||||
"ttl": {"type": "integer"},
|
||||
"collection": {"type": "keyword"},
|
||||
"key": {"type": "keyword"},
|
||||
"value": {"type": "keyword", "index": false}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Notes:
|
||||
- Values are stored as JSON strings under `value` (non-indexed) for size and simplicity.
|
||||
- Each cache entry is keyed by `f"{collection}:{key}"` and saved as the document `_id`.
|
||||
|
||||
---
|
||||
|
||||
## 🔧 Lifecycle and Maintenance
|
||||
|
||||
The backend lazily creates the index on first use via `setup()`; calls are internally guarded by an async lock.
|
||||
|
||||
Available maintenance helpers:
|
||||
|
||||
```python
|
||||
# Drop the entire cache index
|
||||
await cache.clear()
|
||||
|
||||
# Delete only expired entries
|
||||
await cache.cull()
|
||||
```
|
||||
|
||||
- `clear()` deletes the index.
|
||||
- `cull()` issues a `delete_by_query` filtering `expires_at < now`.
|
||||
|
||||
---
|
||||
|
||||
## ✅ Compatibility
|
||||
|
||||
- Works with the core cache entry types defined by `ResponseCachingMiddleware`:
|
||||
- `tools/call` (tool results)
|
||||
- `resources/read` (resource contents)
|
||||
- `prompts/get` (prompt result)
|
||||
- List endpoints for tools, resources, prompts
|
||||
- Uses Pydantic discriminators to serialize/validate cache entries round-trip.
|
||||
|
||||
---
|
||||
|
||||
## 🛡️ Production Notes
|
||||
|
||||
- Ensure your ES cluster has sufficient storage.
|
||||
- Consider a dedicated index per environment or service instance via the `index` parameter.
|
||||
- Close the `AsyncElasticsearch` client on shutdown to avoid warnings and resource leaks.
|
||||
|
||||
---
|
||||
|
||||
## 🧪 Manual Cache Lifecycle
|
||||
|
||||
```python
|
||||
import asyncio
|
||||
from elasticsearch import AsyncElasticsearch
|
||||
from fastmcp.server.middleware.caching import ResponseCachingMiddleware
|
||||
from fastmcp.contrib.middleware.caching.elasticsearch import ElasticsearchCache
|
||||
|
||||
async def main():
|
||||
es = AsyncElasticsearch(hosts=["http://localhost:9200"])
|
||||
cache = ElasticsearchCache(es)
|
||||
await cache.setup() # optional; auto-runs on first access
|
||||
|
||||
# Simulate storing an entry through middleware flow
|
||||
# In practice, the middleware calls set_entry() for you.
|
||||
|
||||
await cache.cull()
|
||||
await es.close()
|
||||
|
||||
asyncio.run(main())
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 📚 Imports
|
||||
|
||||
```python
|
||||
from fastmcp.contrib.middleware.caching.elasticsearch import ElasticsearchCache
|
||||
from fastmcp.server.middleware.caching import ResponseCachingMiddleware
|
||||
```
|
||||
|
||||
This module is part of `fastmcp.contrib`. See `docs/servers/middleware.mdx` and `docs/servers/tools.mdx` for response caching details.
|
||||
|
|
@ -1,5 +0,0 @@
|
|||
from .elasticsearch_cache import ElasticsearchCache
|
||||
|
||||
__all__ = [
|
||||
"ElasticsearchCache",
|
||||
]
|
||||
|
|
@ -1,161 +0,0 @@
|
|||
import asyncio
|
||||
import json
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, Any
|
||||
|
||||
from elasticsearch import AsyncElasticsearch
|
||||
from pydantic import Field, TypeAdapter
|
||||
|
||||
from fastmcp.server.middleware.caching import (
|
||||
CacheEntryTypes,
|
||||
CacheProtocol,
|
||||
)
|
||||
|
||||
DEFAULT_MAPPING = {
|
||||
"properties": {
|
||||
"created_at": {
|
||||
"type": "date",
|
||||
},
|
||||
"expires_at": {
|
||||
"type": "date",
|
||||
},
|
||||
"ttl": {
|
||||
"type": "integer",
|
||||
},
|
||||
"collection": {
|
||||
"type": "keyword",
|
||||
},
|
||||
"key": {
|
||||
"type": "keyword",
|
||||
},
|
||||
"value": {
|
||||
"type": "keyword",
|
||||
"index": False,
|
||||
"doc_values": False,
|
||||
"ignore_above": 256,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class ElasticsearchCache(CacheProtocol):
|
||||
"""A cache client that uses Elasticsearch."""
|
||||
|
||||
setup_called: bool
|
||||
setup_lock: asyncio.Lock
|
||||
cached_entry_typeadapter: TypeAdapter[CacheEntryTypes]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
elasticsearch_client: AsyncElasticsearch,
|
||||
index: str | None = None,
|
||||
mapping: dict[str, Any] | None = None,
|
||||
):
|
||||
"""Initialize the Elasticsearch cache.
|
||||
|
||||
Args:
|
||||
elasticsearch_client: The Elasticsearch client to use.
|
||||
index: The index to use for the cache. Defaults to "fastmcp-response-cache".
|
||||
mapping: The mapping to use for the cache. Defaults to the default mapping.
|
||||
"""
|
||||
self.elasticsearch_client = elasticsearch_client
|
||||
self.index = index or "fastmcp-response-cache"
|
||||
self.mapping = mapping or DEFAULT_MAPPING
|
||||
self.setup_called = False
|
||||
self.setup_lock = asyncio.Lock()
|
||||
self.cached_entry_typeadapter = TypeAdapter(
|
||||
Annotated[CacheEntryTypes, Field(discriminator="collection")],
|
||||
)
|
||||
|
||||
async def get_entry(self, collection: str, key: str) -> CacheEntryTypes | None:
|
||||
if not self.setup_called:
|
||||
await self.setup()
|
||||
|
||||
collection_key = self.make_collection_key(collection=collection, key=key)
|
||||
|
||||
elasticsearch_response = await self.elasticsearch_client.options(
|
||||
ignore_status=404
|
||||
).get(index=self.index, id=collection_key)
|
||||
|
||||
if (
|
||||
elasticsearch_response.body is None
|
||||
or elasticsearch_response.body.get("error")
|
||||
or not elasticsearch_response.body.get("found")
|
||||
):
|
||||
return None
|
||||
|
||||
source = elasticsearch_response.body.get("_source")
|
||||
|
||||
source["value"] = json.loads(source["value"])
|
||||
|
||||
cache_entry = self.cached_entry_typeadapter.validate_python(source)
|
||||
if cache_entry.is_expired():
|
||||
await self.delete(collection=collection, key=key)
|
||||
return None
|
||||
|
||||
return cache_entry
|
||||
|
||||
async def set_entry(
|
||||
self,
|
||||
cache_entry: CacheEntryTypes,
|
||||
) -> None:
|
||||
if not self.setup_called:
|
||||
await self.setup()
|
||||
|
||||
collection_key = self.make_collection_key(
|
||||
collection=cache_entry.collection, key=cache_entry.key
|
||||
)
|
||||
|
||||
document = json.loads(cache_entry.model_dump_json(serialize_as_any=True))
|
||||
|
||||
document["value"] = json.dumps(document["value"])
|
||||
|
||||
await self.elasticsearch_client.index(
|
||||
index=self.index,
|
||||
id=collection_key,
|
||||
body=document,
|
||||
)
|
||||
|
||||
async def delete(self, collection: str, key: str) -> None:
|
||||
collection_key = self.make_collection_key(collection=collection, key=key)
|
||||
|
||||
await self.elasticsearch_client.options(ignore_status=404).delete(
|
||||
index=self.index, id=collection_key
|
||||
)
|
||||
|
||||
async def setup(self) -> None:
|
||||
if self.setup_called:
|
||||
return
|
||||
|
||||
async with self.setup_lock:
|
||||
if self.setup_called:
|
||||
return
|
||||
|
||||
if await self.elasticsearch_client.options(
|
||||
ignore_status=404
|
||||
).indices.exists(index=self.index):
|
||||
return
|
||||
|
||||
await self.elasticsearch_client.options(ignore_status=404).indices.create(
|
||||
index=self.index,
|
||||
mappings=self.mapping,
|
||||
)
|
||||
|
||||
self.setup_called = True
|
||||
|
||||
async def clear(self) -> None:
|
||||
await self.elasticsearch_client.options(ignore_status=404).indices.delete(
|
||||
index=self.index,
|
||||
)
|
||||
|
||||
async def cull(self) -> None:
|
||||
await self.elasticsearch_client.options(ignore_status=404).delete_by_query(
|
||||
index=self.index,
|
||||
body={
|
||||
"query": {
|
||||
"range": {
|
||||
"expires_at": {"lt": datetime.now(tz=timezone.utc).timestamp()},
|
||||
},
|
||||
},
|
||||
},
|
||||
)
|
||||
|
|
@ -2,34 +2,27 @@
|
|||
|
||||
import hashlib
|
||||
import json
|
||||
from abc import ABC
|
||||
from collections import defaultdict
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, ClassVar, Literal, Protocol, TypedDict, TypeVar, cast
|
||||
from typing import Any, TypedDict, TypeVar, cast
|
||||
|
||||
import mcp.types
|
||||
from key_value.aio.adapters.pydantic import PydanticAdapter
|
||||
from key_value.aio.protocols.key_value import AsyncKeyValue
|
||||
from key_value.aio.stores.memory import MemoryStore
|
||||
from key_value.aio.wrappers.statistics import StatisticsWrapper
|
||||
from key_value.aio.wrappers.statistics.wrapper import (
|
||||
KVStoreCollectionStatistics,
|
||||
)
|
||||
from mcp.server.lowlevel.helper_types import ReadResourceContents
|
||||
from mcp.types import GetPromptResult, PromptMessage
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_serializer, field_validator
|
||||
from pydantic.fields import computed_field
|
||||
from pydantic.type_adapter import TypeAdapter
|
||||
from typing_extensions import NotRequired, Self, overload, runtime_checkable
|
||||
from mcp.types import PromptMessage
|
||||
from pydantic import BaseModel, Field
|
||||
from typing_extensions import NotRequired, Self, override
|
||||
|
||||
from fastmcp.prompts.prompt import Prompt
|
||||
from fastmcp.resources.resource import Resource
|
||||
from fastmcp.server.middleware.middleware import CallNext, Middleware, MiddlewareContext
|
||||
from fastmcp.tools.tool import Tool, ToolResult
|
||||
from fastmcp.utilities.logging import get_logger
|
||||
from fastmcp.utilities.types import get_cached_typeadapter
|
||||
|
||||
try:
|
||||
from cachetools import TLRUCache as MemoryCacheClient
|
||||
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]`"
|
||||
)
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
|
@ -42,80 +35,24 @@ ONE_MB_IN_BYTES = 1024 * 1024
|
|||
|
||||
GLOBAL_KEY = "__global__"
|
||||
|
||||
CachableValueTypes = (
|
||||
ToolResult
|
||||
| list[Tool]
|
||||
| list[Resource]
|
||||
| list[Prompt]
|
||||
| list[ReadResourceContents]
|
||||
| GetPromptResult
|
||||
)
|
||||
|
||||
CachableValueTypesVar = TypeVar("CachableValueTypesVar", bound=CachableValueTypes)
|
||||
class CachableToolResult(BaseModel, ToolResult):
|
||||
structured_content: dict[str, Any] | None
|
||||
content: list[mcp.types.ContentBlock]
|
||||
|
||||
|
||||
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 BaseCacheEntry(BaseModel, ABC):
|
||||
model_config: ClassVar[ConfigDict] = ConfigDict(
|
||||
frozen=True, arbitrary_types_allowed=True
|
||||
)
|
||||
|
||||
created_at: datetime = Field(default_factory=lambda: datetime.now(tz=timezone.utc))
|
||||
|
||||
key: str
|
||||
|
||||
ttl: int
|
||||
|
||||
def is_expired(self) -> bool:
|
||||
return datetime.now(tz=timezone.utc) > self.expires_at
|
||||
|
||||
@computed_field
|
||||
@property
|
||||
def expires_at(self) -> datetime:
|
||||
return self.created_at + timedelta(seconds=self.ttl)
|
||||
|
||||
|
||||
class ToolResultCacheEntry(BaseCacheEntry):
|
||||
collection: Literal["tools/call"] = Field(default="tools/call")
|
||||
value: ToolResult
|
||||
|
||||
@field_validator("value", mode="before")
|
||||
@classmethod
|
||||
def validate_value(cls, value: dict[str, Any] | ToolResult) -> ToolResult:
|
||||
if isinstance(value, ToolResult):
|
||||
return value
|
||||
|
||||
content_block_type_adapter: TypeAdapter[list[mcp.types.ContentBlock]] = (
|
||||
get_cached_typeadapter(list[mcp.types.ContentBlock])
|
||||
def from_tool_result(cls, tool_result: ToolResult) -> Self:
|
||||
return cls(
|
||||
structured_content=tool_result.structured_content,
|
||||
content=tool_result.content,
|
||||
)
|
||||
|
||||
content = content_block_type_adapter.validate_python(value.get("content"))
|
||||
|
||||
structured_content = value.get("structured_content")
|
||||
|
||||
return ToolResult(
|
||||
content=content,
|
||||
structured_content=structured_content,
|
||||
)
|
||||
|
||||
@field_serializer("value", when_used="always")
|
||||
def serialize_value(self, value: ToolResult) -> dict[str, Any]:
|
||||
return {
|
||||
"content": [item.model_dump() for item in value.content],
|
||||
"structured_content": value.structured_content,
|
||||
}
|
||||
|
||||
|
||||
class ListToolsCacheEntry(BaseCacheEntry):
|
||||
collection: Literal["tools/list"] = Field(default="tools/list")
|
||||
value: list[Tool]
|
||||
def get_size(self) -> int:
|
||||
return _get_size_of_tool_result(self)
|
||||
|
||||
|
||||
class CachablePrompt(Prompt):
|
||||
@override
|
||||
async def render(
|
||||
self,
|
||||
arguments: dict[str, Any] | None = None,
|
||||
|
|
@ -126,44 +63,28 @@ class CachablePrompt(Prompt):
|
|||
)
|
||||
|
||||
@classmethod
|
||||
def from_prompt(cls, prompt: Prompt) -> Self:
|
||||
return cls(
|
||||
name=prompt.name,
|
||||
title=prompt.title,
|
||||
description=prompt.description,
|
||||
arguments=prompt.arguments,
|
||||
meta=prompt.meta,
|
||||
tags=prompt.tags,
|
||||
enabled=prompt.enabled,
|
||||
)
|
||||
def from_list_prompts(cls, prompts: list[Prompt]) -> list[Self]:
|
||||
cachable_prompts: list[Self] = []
|
||||
for prompt in prompts:
|
||||
cachable_prompts.append(
|
||||
cls(
|
||||
name=prompt.name,
|
||||
title=prompt.title,
|
||||
description=prompt.description,
|
||||
arguments=prompt.arguments,
|
||||
meta=prompt.meta,
|
||||
tags=prompt.tags,
|
||||
enabled=prompt.enabled,
|
||||
)
|
||||
)
|
||||
return cachable_prompts
|
||||
|
||||
|
||||
class ListPromptsCacheEntry(BaseCacheEntry):
|
||||
collection: Literal["prompts/list"] = Field(default="prompts/list")
|
||||
value: Sequence[Prompt]
|
||||
|
||||
@field_validator("value", mode="before")
|
||||
@classmethod
|
||||
def validate_value(
|
||||
cls, value: Sequence[dict[str, Any]] | Sequence[Prompt]
|
||||
) -> Sequence[CachablePrompt]:
|
||||
results = []
|
||||
|
||||
for item in value:
|
||||
if isinstance(item, Prompt):
|
||||
results.append(CachablePrompt.from_prompt(prompt=item))
|
||||
else:
|
||||
results.append(CachablePrompt.model_validate(item))
|
||||
|
||||
return results
|
||||
|
||||
|
||||
class GetPromptCacheEntry(BaseCacheEntry):
|
||||
collection: Literal["prompts/get"] = Field(default="prompts/get")
|
||||
value: mcp.types.GetPromptResult
|
||||
class CachablePromptResult(mcp.types.GetPromptResult): ...
|
||||
|
||||
|
||||
class CachableResource(Resource):
|
||||
@override
|
||||
async def read(self) -> str | bytes:
|
||||
"""Read the resource content."""
|
||||
raise NotImplementedError(
|
||||
|
|
@ -171,274 +92,69 @@ class CachableResource(Resource):
|
|||
)
|
||||
|
||||
@classmethod
|
||||
def from_resource(cls, resource: Resource) -> Self:
|
||||
return cls(
|
||||
name=resource.name,
|
||||
description=resource.description,
|
||||
uri=resource.uri,
|
||||
mime_type=resource.mime_type,
|
||||
annotations=resource.annotations,
|
||||
meta=resource.meta,
|
||||
tags=resource.tags,
|
||||
enabled=resource.enabled,
|
||||
)
|
||||
|
||||
|
||||
class ListResourcesCacheEntry(BaseCacheEntry):
|
||||
collection: Literal["resources/list"] = Field(default="resources/list")
|
||||
value: Sequence[Resource]
|
||||
|
||||
@field_validator("value", mode="before")
|
||||
@classmethod
|
||||
def validate_value(
|
||||
cls, value: Sequence[dict[str, Any]] | Sequence[Resource]
|
||||
) -> Sequence[Resource]:
|
||||
results = []
|
||||
for item in value:
|
||||
if isinstance(item, Resource):
|
||||
results.append(CachableResource.from_resource(resource=item))
|
||||
else:
|
||||
results.append(CachableResource.model_validate(item))
|
||||
|
||||
return results
|
||||
|
||||
|
||||
class ReadResourceCacheEntry(BaseCacheEntry):
|
||||
collection: Literal["resources/read"] = Field(default="resources/read")
|
||||
value: Sequence[ReadResourceContents]
|
||||
|
||||
@field_validator("value", mode="before")
|
||||
@classmethod
|
||||
def validate_value(
|
||||
cls, value: Sequence[dict[str, Any]] | Sequence[ReadResourceContents]
|
||||
) -> Sequence[ReadResourceContents]:
|
||||
resource_contents: list[ReadResourceContents] = []
|
||||
for item in value:
|
||||
if isinstance(item, ReadResourceContents):
|
||||
resource_contents.append(item)
|
||||
continue
|
||||
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
|
||||
if not (content := item.get("content")):
|
||||
continue
|
||||
|
||||
mime_type = item.get("mime_type")
|
||||
|
||||
resource_contents.append(
|
||||
ReadResourceContents(content=content, mime_type=mime_type)
|
||||
def from_list_resources(cls, resources: list[Resource]) -> list[Self]:
|
||||
cachable_resources: list[Self] = []
|
||||
for resource in resources:
|
||||
cachable_resources.append(
|
||||
cls(
|
||||
name=resource.name,
|
||||
description=resource.description,
|
||||
uri=resource.uri,
|
||||
mime_type=resource.mime_type,
|
||||
annotations=resource.annotations,
|
||||
meta=resource.meta,
|
||||
tags=resource.tags,
|
||||
enabled=resource.enabled,
|
||||
)
|
||||
)
|
||||
|
||||
return resource_contents
|
||||
|
||||
@field_serializer("value")
|
||||
def serialize_value(
|
||||
self, value: Sequence[ReadResourceContents]
|
||||
) -> list[dict[str, Any]]:
|
||||
return [
|
||||
{
|
||||
"content": item.content,
|
||||
"mime_type": item.mime_type,
|
||||
}
|
||||
for item in value
|
||||
]
|
||||
return cachable_resources
|
||||
|
||||
|
||||
CacheEntryTypes = (
|
||||
GetPromptCacheEntry
|
||||
| ListPromptsCacheEntry
|
||||
| ReadResourceCacheEntry
|
||||
| ListResourcesCacheEntry
|
||||
| ToolResultCacheEntry
|
||||
| ListToolsCacheEntry
|
||||
)
|
||||
class CachableTool(Tool):
|
||||
@classmethod
|
||||
def from_list_tools(cls, tools: list[Tool]) -> list[Self]:
|
||||
cachable_tools: list[Self] = []
|
||||
for tool in tools:
|
||||
cachable_tools.append(
|
||||
cls(
|
||||
name=tool.name,
|
||||
description=tool.description,
|
||||
parameters=tool.parameters,
|
||||
output_schema=tool.output_schema,
|
||||
annotations=tool.annotations,
|
||||
serializer=tool.serializer,
|
||||
meta=tool.meta,
|
||||
tags=tool.tags,
|
||||
enabled=tool.enabled,
|
||||
)
|
||||
)
|
||||
return cachable_tools
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class CacheProtocol(Protocol):
|
||||
"""A protocol for a cache client."""
|
||||
|
||||
async def get_entry(
|
||||
self,
|
||||
collection: str,
|
||||
key: str,
|
||||
) -> CacheEntryTypes | None:
|
||||
"""Get a cache entry from the cache."""
|
||||
|
||||
async def set_entry(
|
||||
self,
|
||||
cache_entry: CacheEntryTypes,
|
||||
) -> None:
|
||||
"""Set a value in the cache using the collection and key."""
|
||||
|
||||
async def delete(
|
||||
self,
|
||||
collection: str,
|
||||
key: str,
|
||||
) -> None:
|
||||
"""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}"
|
||||
class CachableReadResourceContents(BaseModel, ReadResourceContents): ...
|
||||
|
||||
|
||||
class DiskCache(CacheProtocol):
|
||||
"""A caching client that uses the DiskCache library to cache to disk."""
|
||||
|
||||
@overload
|
||||
def __init__(self, *, disk_cache: DiskCacheClient):
|
||||
"""Initialize the disk cache with a diskcache client."""
|
||||
|
||||
@overload
|
||||
def __init__(self, path: str, *, size_limit: int = ONE_GB_IN_BYTES):
|
||||
"""Initialize a 1GB disk cache at the provided path."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path: str | None = None,
|
||||
*,
|
||||
disk_cache: DiskCacheClient | None = None,
|
||||
size_limit: int = ONE_GB_IN_BYTES,
|
||||
):
|
||||
self._cache = disk_cache or DiskCacheClient(
|
||||
directory=path, size_limit=size_limit
|
||||
)
|
||||
|
||||
async def get_entry(self, collection: str, key: str) -> CacheEntryTypes | None:
|
||||
collection_key: str = self.make_collection_key(collection=collection, key=key)
|
||||
|
||||
cache_entry = self._cache.get(key=collection_key)
|
||||
|
||||
if cache_entry is None:
|
||||
return None
|
||||
|
||||
return cache_entry # pyright: ignore[reportReturnType]
|
||||
|
||||
async def set_entry(
|
||||
self,
|
||||
cache_entry: CacheEntryTypes,
|
||||
) -> None:
|
||||
collection_key: str = self.make_collection_key(
|
||||
collection=cache_entry.collection, key=cache_entry.key
|
||||
)
|
||||
|
||||
self._cache.set(key=collection_key, value=cache_entry, expire=cache_entry.ttl)
|
||||
|
||||
async def delete(self, collection: str, key: str) -> None:
|
||||
collection_key = self.make_collection_key(collection=collection, key=key)
|
||||
|
||||
self._cache.delete(key=collection_key)
|
||||
class CachableToolList(BaseModel):
|
||||
cachable_tools: list[CachableTool]
|
||||
|
||||
|
||||
DEFAULT_MEMORY_CACHE_MAX_ENTRIES = 1000
|
||||
class CachableResourceList(BaseModel):
|
||||
cachable_resources: list[CachableResource]
|
||||
|
||||
|
||||
def _memory_cache_ttu(_key: Any, value: BaseCacheEntry, now: float) -> float:
|
||||
"""TTU function for the memory cache. Determines the TTL of the cache entry."""
|
||||
return now + value.ttl
|
||||
class CachablePromptList(BaseModel):
|
||||
cachable_prompts: list[CachablePrompt]
|
||||
|
||||
|
||||
def _memory_cache_getsizeof(value: BaseCacheEntry) -> int:
|
||||
"""Getsizeof function for the memory cache. Currently measures how many entries are in the cache."""
|
||||
return 1
|
||||
|
||||
|
||||
class InMemoryCache(CacheProtocol):
|
||||
"""A simple in-memory cache."""
|
||||
|
||||
def __init__(self, max_entries: int = DEFAULT_MEMORY_CACHE_MAX_ENTRIES):
|
||||
"""Initialize the in-memory cache.
|
||||
|
||||
Args:
|
||||
max_entries: The maximum number of entries to store in the cache. Defaults to 1000.
|
||||
"""
|
||||
self.max_entries = max_entries
|
||||
self._cache: MemoryCacheClient[Any, CacheEntryTypes] = MemoryCacheClient[
|
||||
Any, CacheEntryTypes
|
||||
](
|
||||
maxsize=max_entries,
|
||||
ttu=_memory_cache_ttu,
|
||||
getsizeof=_memory_cache_getsizeof,
|
||||
)
|
||||
|
||||
async def get_entry(self, collection: str, key: str) -> CacheEntryTypes | None:
|
||||
collection_key: str = self.make_collection_key(collection=collection, key=key)
|
||||
|
||||
return self._cache.get(collection_key)
|
||||
|
||||
async def set_entry(
|
||||
self,
|
||||
cache_entry: CacheEntryTypes,
|
||||
) -> None:
|
||||
collection_key: str = self.make_collection_key(
|
||||
collection=cache_entry.collection, key=cache_entry.key
|
||||
)
|
||||
|
||||
self._cache[collection_key] = cache_entry
|
||||
|
||||
async def delete(self, collection: str, key: str) -> None:
|
||||
collection_key = self.make_collection_key(collection=collection, key=key)
|
||||
|
||||
self._cache.pop(collection_key, None)
|
||||
|
||||
async def setup(self) -> None:
|
||||
return None
|
||||
|
||||
async def clear(self) -> None:
|
||||
self._cache.clear()
|
||||
|
||||
|
||||
class CacheMethodStats(BaseModel):
|
||||
"""Stats for a cache method."""
|
||||
|
||||
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(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
|
||||
class CachableReadResourceContentsList(BaseModel):
|
||||
cachable_read_resource_contents: list[CachableReadResourceContents]
|
||||
|
||||
|
||||
class SharedMethodSettings(TypedDict):
|
||||
"""Shared config for a cache method."""
|
||||
|
||||
ttl: NotRequired[int]
|
||||
enabled: NotRequired[bool]
|
||||
|
||||
|
||||
class ListToolsSettings(SharedMethodSettings):
|
||||
|
|
@ -483,15 +199,6 @@ class MethodSettings(TypedDict):
|
|||
|
||||
MethodSettingsType = TypeVar("MethodSettingsType", bound=SharedMethodSettings)
|
||||
|
||||
MCP_METHOD_TO_CACHE_ENTRY_TYPE: dict[str, type[CacheEntryTypes]] = {
|
||||
"tools/list": ListToolsCacheEntry,
|
||||
"tools/call": ToolResultCacheEntry,
|
||||
"resources/list": ListResourcesCacheEntry,
|
||||
"resources/read": ReadResourceCacheEntry,
|
||||
"prompts/list": ListPromptsCacheEntry,
|
||||
"prompts/get": GetPromptCacheEntry,
|
||||
}
|
||||
|
||||
|
||||
MCP_METHOD_TO_METHOD_SETTINGS_KEY = {
|
||||
"tools/list": "list_tools",
|
||||
|
|
@ -524,6 +231,15 @@ DEFAULT_METHOD_SETTINGS: MethodSettings = MethodSettings(
|
|||
)
|
||||
|
||||
|
||||
class ResponseCachingStatistics(BaseModel):
|
||||
list_tools: KVStoreCollectionStatistics | None = Field(default=None)
|
||||
list_resources: KVStoreCollectionStatistics | None = Field(default=None)
|
||||
list_prompts: KVStoreCollectionStatistics | None = Field(default=None)
|
||||
read_resource: KVStoreCollectionStatistics | None = Field(default=None)
|
||||
get_prompt: KVStoreCollectionStatistics | None = Field(default=None)
|
||||
call_tool: KVStoreCollectionStatistics | None = Field(default=None)
|
||||
|
||||
|
||||
class ResponseCachingMiddleware(Middleware):
|
||||
"""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
|
||||
|
|
@ -539,7 +255,7 @@ class ResponseCachingMiddleware(Middleware):
|
|||
|
||||
def __init__(
|
||||
self,
|
||||
cache_backend: CacheProtocol | None = None,
|
||||
cache_store: AsyncKeyValue | None = None,
|
||||
method_settings: MethodSettings | None = None,
|
||||
default_ttl: int = ONE_HOUR_IN_SECONDS,
|
||||
max_item_size: int | None = None,
|
||||
|
|
@ -553,15 +269,56 @@ class ResponseCachingMiddleware(Middleware):
|
|||
max_item_size: The maximum size of an item to cache. Defaults to no size limit.
|
||||
"""
|
||||
self._default_ttl: int = default_ttl
|
||||
self._backend: CacheProtocol = cache_backend or InMemoryCache()
|
||||
self._max_item_size: int | None = max_item_size
|
||||
self._backend: AsyncKeyValue = cache_store or MemoryStore()
|
||||
self._stats: StatisticsWrapper = StatisticsWrapper(store=self._backend)
|
||||
|
||||
self._stats = CacheStats()
|
||||
self._max_item_size: int | None = max_item_size
|
||||
|
||||
self.method_settings: MethodSettings = (
|
||||
method_settings or DEFAULT_METHOD_SETTINGS
|
||||
)
|
||||
|
||||
self._list_tools_cache: PydanticAdapter[CachableToolList] = PydanticAdapter(
|
||||
key_value=self._stats,
|
||||
pydantic_model=CachableToolList,
|
||||
default_collection="tools/list",
|
||||
)
|
||||
|
||||
self._list_resources_cache: PydanticAdapter[CachableResourceList] = (
|
||||
PydanticAdapter(
|
||||
key_value=self._stats,
|
||||
pydantic_model=CachableResourceList,
|
||||
default_collection="resources/list",
|
||||
)
|
||||
)
|
||||
|
||||
self._list_prompts_cache: PydanticAdapter[CachablePromptList] = PydanticAdapter(
|
||||
key_value=self._stats,
|
||||
pydantic_model=CachablePromptList,
|
||||
default_collection="prompts/list",
|
||||
)
|
||||
|
||||
self._read_resource_cache: PydanticAdapter[CachableReadResourceContentsList] = (
|
||||
PydanticAdapter(
|
||||
key_value=self._stats,
|
||||
pydantic_model=CachableReadResourceContentsList,
|
||||
default_collection="resources/read",
|
||||
)
|
||||
)
|
||||
|
||||
self._get_prompt_cache: PydanticAdapter[CachablePromptResult] = PydanticAdapter(
|
||||
key_value=self._stats,
|
||||
pydantic_model=CachablePromptResult,
|
||||
default_collection="prompts/get",
|
||||
)
|
||||
|
||||
self._call_tool_cache: PydanticAdapter[CachableToolResult] = PydanticAdapter(
|
||||
key_value=self._stats,
|
||||
pydantic_model=CachableToolResult,
|
||||
default_collection="tools/call",
|
||||
)
|
||||
|
||||
@override
|
||||
async def on_list_tools(
|
||||
self,
|
||||
context: MiddlewareContext[mcp.types.ListToolsRequest],
|
||||
|
|
@ -570,8 +327,21 @@ class ResponseCachingMiddleware(Middleware):
|
|||
if self._should_bypass_caching(context=context):
|
||||
return await call_next(context=context)
|
||||
|
||||
return await self._cached_call_next(context=context, call_next=call_next)
|
||||
if cached_value := await self._list_tools_cache.get(key=GLOBAL_KEY):
|
||||
return cached_value.cachable_tools
|
||||
|
||||
value: list[Tool] = await call_next(context=context)
|
||||
|
||||
cachable_tools: list[CachableTool] = CachableTool.from_list_tools(tools=value)
|
||||
|
||||
await self._list_tools_cache.put(
|
||||
key=GLOBAL_KEY,
|
||||
value=CachableToolList(cachable_tools=cachable_tools),
|
||||
)
|
||||
|
||||
return cachable_tools
|
||||
|
||||
@override
|
||||
async def on_list_resources(
|
||||
self,
|
||||
context: MiddlewareContext[mcp.types.ListResourcesRequest],
|
||||
|
|
@ -582,8 +352,23 @@ class ResponseCachingMiddleware(Middleware):
|
|||
if self._should_bypass_caching(context=context):
|
||||
return await call_next(context)
|
||||
|
||||
return await self._cached_call_next(context=context, call_next=call_next)
|
||||
if cached_value := await self._list_resources_cache.get(key=GLOBAL_KEY):
|
||||
return cached_value.cachable_resources
|
||||
|
||||
value: list[Resource] = await call_next(context=context)
|
||||
|
||||
cachable_resources: list[CachableResource] = (
|
||||
CachableResource.from_list_resources(resources=value)
|
||||
)
|
||||
|
||||
await self._list_resources_cache.put(
|
||||
key=GLOBAL_KEY,
|
||||
value=CachableResourceList(cachable_resources=cachable_resources),
|
||||
)
|
||||
|
||||
return cachable_resources
|
||||
|
||||
@override
|
||||
async def on_list_prompts(
|
||||
self,
|
||||
context: MiddlewareContext[mcp.types.ListPromptsRequest],
|
||||
|
|
@ -594,8 +379,23 @@ class ResponseCachingMiddleware(Middleware):
|
|||
if self._should_bypass_caching(context=context):
|
||||
return await call_next(context)
|
||||
|
||||
return await self._cached_call_next(context=context, call_next=call_next)
|
||||
if cached_value := await self._list_prompts_cache.get(key=GLOBAL_KEY):
|
||||
return cached_value.cachable_prompts
|
||||
|
||||
value: list[Prompt] = await call_next(context=context)
|
||||
|
||||
cachable_prompts: list[CachablePrompt] = CachablePrompt.from_list_prompts(
|
||||
prompts=value
|
||||
)
|
||||
|
||||
await self._list_prompts_cache.put(
|
||||
key=GLOBAL_KEY,
|
||||
value=CachablePromptList(cachable_prompts=cachable_prompts),
|
||||
)
|
||||
|
||||
return cachable_prompts
|
||||
|
||||
@override
|
||||
async def on_call_tool(
|
||||
self,
|
||||
context: MiddlewareContext[mcp.types.CallToolRequestParams],
|
||||
|
|
@ -609,12 +409,28 @@ class ResponseCachingMiddleware(Middleware):
|
|||
if not self._matches_tool_cache_settings(context=context):
|
||||
return await call_next(context=context)
|
||||
|
||||
return await self._cached_call_next(
|
||||
context=context,
|
||||
call_next=call_next,
|
||||
key=_make_call_tool_cache_key(msg=context.message),
|
||||
if cached_value := await self._call_tool_cache.get(
|
||||
key=_make_call_tool_cache_key(msg=context.message)
|
||||
):
|
||||
return cached_value
|
||||
|
||||
tool_result: ToolResult = await call_next(context=context)
|
||||
|
||||
cachable_value: CachableToolResult = CachableToolResult.from_tool_result(
|
||||
tool_result=tool_result
|
||||
)
|
||||
|
||||
if self._max_item_size and cachable_value.get_size() > self._max_item_size:
|
||||
return tool_result
|
||||
|
||||
await self._call_tool_cache.put(
|
||||
key=_make_call_tool_cache_key(msg=context.message),
|
||||
value=cachable_value,
|
||||
)
|
||||
|
||||
return cachable_value
|
||||
|
||||
@override
|
||||
async def on_read_resource(
|
||||
self,
|
||||
context: MiddlewareContext[mcp.types.ReadResourceRequestParams],
|
||||
|
|
@ -627,12 +443,28 @@ class ResponseCachingMiddleware(Middleware):
|
|||
if self._should_bypass_caching(context=context):
|
||||
return await call_next(context=context)
|
||||
|
||||
return await self._cached_call_next(
|
||||
context=context,
|
||||
call_next=call_next,
|
||||
if cached_value := await self._read_resource_cache.get(
|
||||
key=_make_read_resource_cache_key(msg=context.message)
|
||||
):
|
||||
return cached_value.cachable_read_resource_contents
|
||||
|
||||
value: list[ReadResourceContents] = await call_next(context=context)
|
||||
|
||||
cachable_read_resource_contents: list[CachableReadResourceContents] = [
|
||||
CachableReadResourceContents(content=item.content, mime_type=item.mime_type)
|
||||
for item in value
|
||||
]
|
||||
|
||||
await self._read_resource_cache.put(
|
||||
key=_make_read_resource_cache_key(msg=context.message),
|
||||
value=CachableReadResourceContentsList(
|
||||
cachable_read_resource_contents=cachable_read_resource_contents
|
||||
),
|
||||
)
|
||||
|
||||
return cachable_read_resource_contents
|
||||
|
||||
@override
|
||||
async def on_get_prompt(
|
||||
self,
|
||||
context: MiddlewareContext[mcp.types.GetPromptRequestParams],
|
||||
|
|
@ -645,119 +477,42 @@ class ResponseCachingMiddleware(Middleware):
|
|||
if self._should_bypass_caching(context=context):
|
||||
return await call_next(context)
|
||||
|
||||
return await self._cached_call_next(
|
||||
context=context,
|
||||
call_next=call_next,
|
||||
key=_make_get_prompt_cache_key(msg=context.message),
|
||||
if cached_value := await self._get_prompt_cache.get(
|
||||
key=_make_get_prompt_cache_key(msg=context.message)
|
||||
):
|
||||
return cached_value
|
||||
|
||||
value: mcp.types.GetPromptResult = await call_next(context=context)
|
||||
|
||||
cachable_value: CachablePromptResult = CachablePromptResult(
|
||||
messages=value.messages, description=value.description, _meta=value.meta
|
||||
)
|
||||
|
||||
await self._get_prompt_cache.put(
|
||||
key=_make_get_prompt_cache_key(msg=context.message),
|
||||
value=cachable_value,
|
||||
)
|
||||
|
||||
return cachable_value
|
||||
|
||||
@override
|
||||
async def on_notification(
|
||||
self,
|
||||
context: MiddlewareContext[mcp.types.Notification],
|
||||
call_next: CallNext[mcp.types.Notification, Any],
|
||||
context: MiddlewareContext[mcp.types.Notification[Any, Any]],
|
||||
call_next: CallNext[mcp.types.Notification[Any, Any], Any],
|
||||
) -> Any:
|
||||
"""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"
|
||||
_ = await self._list_tools_cache.delete(key=GLOBAL_KEY)
|
||||
elif isinstance(context.message, mcp.types.ResourceListChangedNotification):
|
||||
collection = "resources/list"
|
||||
_ = await self._list_resources_cache.delete(key=GLOBAL_KEY)
|
||||
elif isinstance(context.message, mcp.types.PromptListChangedNotification):
|
||||
collection = "prompts/list"
|
||||
_ = await self._list_prompts_cache.delete(key=GLOBAL_KEY)
|
||||
else:
|
||||
collection = None
|
||||
pass
|
||||
|
||||
if collection:
|
||||
await self._backend.delete(collection=collection, key=GLOBAL_KEY)
|
||||
|
||||
return await call_next(context)
|
||||
|
||||
async def _cached_call_next(
|
||||
self,
|
||||
context: MiddlewareContext[Any],
|
||||
call_next: CallNext[Any, CachableValueTypesVar],
|
||||
key: str | None = None,
|
||||
) -> CachableValueTypesVar:
|
||||
"""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
|
||||
|
||||
if cached_value := await self._get_cache(
|
||||
context=context,
|
||||
call_next=call_next,
|
||||
key=key,
|
||||
):
|
||||
return cached_value
|
||||
|
||||
result: CachableValueTypesVar = await call_next(context)
|
||||
|
||||
return await self._store_in_cache_and_return(
|
||||
context=context,
|
||||
key=key,
|
||||
value=result,
|
||||
)
|
||||
|
||||
async def _get_cache(
|
||||
self,
|
||||
context: MiddlewareContext[Any],
|
||||
call_next: CallNext[Any, CachableValueTypesVar],
|
||||
key: str | None = None,
|
||||
) -> CachableValueTypesVar | None:
|
||||
"""Get a value from the cache and update the cache stats."""
|
||||
|
||||
if key is None:
|
||||
key = GLOBAL_KEY
|
||||
|
||||
if not (collection := context.method):
|
||||
logger.warning("No method found on context, skipping cache")
|
||||
return None
|
||||
|
||||
if cached_entry := await self._backend.get_entry(
|
||||
collection=collection, key=key
|
||||
):
|
||||
self._stats.mark_hit(collection=collection)
|
||||
return cast(CachableValueTypesVar, cached_entry.value)
|
||||
|
||||
self._stats.mark_miss(collection=collection)
|
||||
|
||||
return None
|
||||
|
||||
async def _store_in_cache_and_return(
|
||||
self,
|
||||
context: MiddlewareContext[Any],
|
||||
key: str | None,
|
||||
value: CachableValueTypesVar,
|
||||
) -> CachableValueTypesVar:
|
||||
"""Store a value in the cache (if it's not too big) with the appropriate TTL."""
|
||||
|
||||
if key is None:
|
||||
key = GLOBAL_KEY
|
||||
|
||||
if not (collection := context.method):
|
||||
logger.warning("No method found on context, skipping cache")
|
||||
return value
|
||||
|
||||
if self._max_item_size is not None:
|
||||
if get_size_of_value(value=value) > self._max_item_size:
|
||||
self._stats.mark_too_big(collection=collection)
|
||||
return value
|
||||
|
||||
ttl: int = self._get_cache_ttl(context=context)
|
||||
|
||||
cache_entry: CacheEntryTypes = MCP_METHOD_TO_CACHE_ENTRY_TYPE[collection](
|
||||
collection=collection, # pyright: ignore[reportArgumentType]
|
||||
key=key,
|
||||
value=value, # pyright: ignore[reportArgumentType]
|
||||
ttl=ttl,
|
||||
)
|
||||
|
||||
await self._backend.set_entry(
|
||||
cache_entry=cache_entry,
|
||||
)
|
||||
|
||||
return value
|
||||
return await call_next(context=context)
|
||||
|
||||
def _matches_tool_cache_settings(
|
||||
self, context: MiddlewareContext[mcp.types.CallToolRequestParams]
|
||||
|
|
@ -810,7 +565,7 @@ class ResponseCachingMiddleware(Middleware):
|
|||
"""Get the cache TTL for a method."""
|
||||
|
||||
settings: SharedMethodSettings | None = self._get_cache_settings(
|
||||
context=context
|
||||
context=context,
|
||||
)
|
||||
|
||||
if not settings or "ttl" not in settings:
|
||||
|
|
@ -821,11 +576,24 @@ class ResponseCachingMiddleware(Middleware):
|
|||
def _should_bypass_caching(self, context: MiddlewareContext[Any]) -> bool:
|
||||
"""Check if the method should bypass caching."""
|
||||
|
||||
if not self._get_cache_settings(context=context):
|
||||
if not (cache_settings := self._get_cache_settings(context=context)):
|
||||
return True
|
||||
|
||||
if cache_settings.get("enabled") is False:
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
def statistics(self) -> ResponseCachingStatistics:
|
||||
return ResponseCachingStatistics(
|
||||
list_tools=self._stats.statistics.collections.get("tools/list"),
|
||||
list_resources=self._stats.statistics.collections.get("resources/list"),
|
||||
list_prompts=self._stats.statistics.collections.get("prompts/list"),
|
||||
read_resource=self._stats.statistics.collections.get("resources/read"),
|
||||
get_prompt=self._stats.statistics.collections.get("prompts/get"),
|
||||
call_tool=self._stats.statistics.collections.get("tools/call"),
|
||||
)
|
||||
|
||||
|
||||
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."""
|
||||
|
|
@ -861,7 +629,7 @@ def _get_arguments_str(arguments: dict[str, Any] | None) -> str:
|
|||
return repr(arguments)
|
||||
|
||||
|
||||
def get_size_of_content_blocks(
|
||||
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."""
|
||||
|
|
@ -872,10 +640,10 @@ def get_size_of_content_blocks(
|
|||
return sum([len(item.model_dump_json()) for item in value])
|
||||
|
||||
|
||||
def get_size_of_tool_result(value: ToolResult) -> int:
|
||||
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)
|
||||
content_size = _get_size_of_content_blocks(value.content)
|
||||
structured_content_size = len(
|
||||
json.dumps(
|
||||
value.structured_content, sort_keys=True, separators=(",", ":")
|
||||
|
|
@ -885,19 +653,19 @@ def get_size_of_tool_result(value: ToolResult) -> int:
|
|||
return content_size + structured_content_size
|
||||
|
||||
|
||||
def get_size_of_one_value(value: BaseModel | ToolResult | ReadResourceContents) -> int:
|
||||
"""Get the size of an mcp type."""
|
||||
# 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):
|
||||
return len(value.content)
|
||||
return len(value.model_dump_json())
|
||||
# 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: CachableValueTypes) -> int:
|
||||
"""Get the size of a cache entry."""
|
||||
if isinstance(value, BaseModel | ToolResult | ReadResourceContents):
|
||||
return get_size_of_one_value(value)
|
||||
# def get_size_of_value(value: CachableListTypes) -> int:
|
||||
# """Get the size of a cache entry."""
|
||||
# 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)
|
||||
# return sum(get_size_of_one_value(item) for item in value)
|
||||
|
|
|
|||
|
|
@ -8,8 +8,10 @@ from unittest.mock import AsyncMock, MagicMock
|
|||
import mcp.types
|
||||
import pytest
|
||||
from inline_snapshot import snapshot
|
||||
from key_value.aio.stores.disk import DiskStore
|
||||
from key_value.aio.stores.memory import MemoryStore
|
||||
from mcp.server.lowlevel.helper_types import ReadResourceContents
|
||||
from mcp.types import GetPromptResult, PromptMessage, TextContent, TextResourceContents
|
||||
from mcp.types import PromptMessage, TextContent, TextResourceContents
|
||||
from pydantic import AnyUrl, BaseModel
|
||||
|
||||
from fastmcp import FastMCP
|
||||
|
|
@ -18,23 +20,9 @@ from fastmcp.client.transports import FastMCPTransport
|
|||
from fastmcp.prompts.prompt import FunctionPrompt, Prompt
|
||||
from fastmcp.resources.resource import Resource
|
||||
from fastmcp.server.middleware.caching import (
|
||||
BaseCacheEntry,
|
||||
CachablePrompt,
|
||||
CachableResource,
|
||||
CacheMethodStats,
|
||||
CacheProtocol,
|
||||
CacheStats,
|
||||
CallToolSettings,
|
||||
DiskCache,
|
||||
GetPromptCacheEntry,
|
||||
InMemoryCache,
|
||||
ListPromptsCacheEntry,
|
||||
ListResourcesCacheEntry,
|
||||
ListToolsCacheEntry,
|
||||
MethodSettings,
|
||||
ReadResourceCacheEntry,
|
||||
ResponseCachingMiddleware,
|
||||
ToolResultCacheEntry,
|
||||
)
|
||||
from fastmcp.server.middleware.middleware import CallNext, MiddlewareContext
|
||||
from fastmcp.tools.tool import Tool, ToolResult
|
||||
|
|
@ -147,11 +135,13 @@ class TrackingCalculator:
|
|||
add_calls: int
|
||||
multiply_calls: int
|
||||
crazy_calls: int
|
||||
very_large_response_calls: int
|
||||
|
||||
def __init__(self):
|
||||
self.add_calls = 0
|
||||
self.multiply_calls = 0
|
||||
self.crazy_calls = 0
|
||||
self.very_large_response_calls = 0
|
||||
|
||||
def add(self, a: int, b: int) -> int:
|
||||
self.add_calls += 1
|
||||
|
|
@ -161,6 +151,10 @@ class TrackingCalculator:
|
|||
self.multiply_calls += 1
|
||||
return a * b
|
||||
|
||||
def very_large_response(self) -> str:
|
||||
self.very_large_response_calls += 1
|
||||
return "istenchars" * 100000 # 1,000,000 characters, 1mb
|
||||
|
||||
def crazy(self, a: CrazyModel) -> CrazyModel:
|
||||
self.crazy_calls += 1
|
||||
return a
|
||||
|
|
@ -183,6 +177,11 @@ class TrackingCalculator:
|
|||
tool=Tool.from_function(fn=self.multiply, name=f"{prefix}multiply")
|
||||
)
|
||||
fastmcp.add_tool(tool=Tool.from_function(fn=self.crazy, name=f"{prefix}crazy"))
|
||||
fastmcp.add_tool(
|
||||
tool=Tool.from_function(
|
||||
fn=self.very_large_response, name=f"{prefix}very_large_response"
|
||||
)
|
||||
)
|
||||
|
||||
def add_prompts(self, fastmcp: FastMCP, prefix: str = ""):
|
||||
fastmcp.add_prompt(
|
||||
|
|
@ -251,365 +250,6 @@ def sample_tool_result() -> ToolResult:
|
|||
)
|
||||
|
||||
|
||||
class TestCacheEntry:
|
||||
"""Test CacheEntry class functionality."""
|
||||
|
||||
def test_init_and_expiration(self):
|
||||
"""Test cache entry initialization and expiration logic."""
|
||||
entry: ToolResultCacheEntry[ToolResult] = ToolResultCacheEntry(
|
||||
key="test_key",
|
||||
value=ToolResult(
|
||||
content=[{"type": "text", "text": "success"}],
|
||||
structured_content={"result": "success"},
|
||||
),
|
||||
ttl=3600,
|
||||
)
|
||||
|
||||
assert entry.key == "test_key"
|
||||
assert not entry.is_expired()
|
||||
|
||||
# Test expired entry
|
||||
expired_entry: ToolResultCacheEntry[ToolResult] = ToolResultCacheEntry(
|
||||
key="test_key",
|
||||
value=ToolResult(
|
||||
content=[{"type": "text", "text": "success"}],
|
||||
structured_content={"result": "success"},
|
||||
),
|
||||
ttl=-1,
|
||||
)
|
||||
|
||||
assert expired_entry.is_expired()
|
||||
|
||||
def test_tool_result_cache_entry_validation(self):
|
||||
"""Test ToolResultCacheEntry class functionality."""
|
||||
entry: ToolResultCacheEntry[ToolResult] = ToolResultCacheEntry(
|
||||
key="test_key",
|
||||
value=SAMPLE_TOOL_RESULT,
|
||||
ttl=3600,
|
||||
)
|
||||
|
||||
dumped_entry = entry.model_dump()
|
||||
new_entry = ToolResultCacheEntry.model_validate(dumped_entry)
|
||||
|
||||
assert dump_mcp_types(model=new_entry.value.content) == snapshot(
|
||||
[{"type": "text", "text": "test_text", "annotations": None, "meta": None}]
|
||||
)
|
||||
assert new_entry.value.structured_content == snapshot({"result": "test_result"})
|
||||
|
||||
def test_read_resource_cache_entry_validation(self):
|
||||
"""Test ReadResourceCacheEntry class functionality."""
|
||||
entry: ReadResourceCacheEntry[list[ReadResourceContents]] = (
|
||||
ReadResourceCacheEntry(
|
||||
key="test_key",
|
||||
value=[SAMPLE_READ_RESOURCE_CONTENTS],
|
||||
ttl=3600,
|
||||
)
|
||||
)
|
||||
|
||||
dumped_entry = entry.model_dump()
|
||||
new_entry = ReadResourceCacheEntry.model_validate(dumped_entry)
|
||||
|
||||
assert new_entry.value == snapshot(
|
||||
[ReadResourceContents(content="test_text", mime_type="text/plain")]
|
||||
)
|
||||
|
||||
def test_get_prompt_cache_entry_validation(self):
|
||||
"""Test GetPromptCacheEntry class functionality."""
|
||||
entry: GetPromptCacheEntry = GetPromptCacheEntry(
|
||||
key="test_key",
|
||||
value=SAMPLE_GET_PROMPT_RESULT,
|
||||
ttl=3600,
|
||||
)
|
||||
|
||||
dumped_entry = entry.model_dump()
|
||||
new_entry = GetPromptCacheEntry.model_validate(dumped_entry)
|
||||
|
||||
assert new_entry.value == snapshot(
|
||||
GetPromptResult(
|
||||
messages=[
|
||||
PromptMessage(
|
||||
role="user", content=TextContent(type="text", text="test_text")
|
||||
)
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def test_list_tools_cache_entry_validation(self):
|
||||
"""Test ListToolsCacheEntry class functionality."""
|
||||
entry: ListToolsCacheEntry = ListToolsCacheEntry(
|
||||
key="test_key",
|
||||
value=[SAMPLE_TOOL],
|
||||
ttl=3600,
|
||||
)
|
||||
|
||||
dumped_entry = entry.model_dump()
|
||||
new_entry = ListToolsCacheEntry.model_validate(dumped_entry)
|
||||
|
||||
assert new_entry.value == snapshot(
|
||||
[Tool(name="test_tool", parameters={"param1": "value1", "param2": 42})]
|
||||
)
|
||||
|
||||
def test_list_resources_cache_entry_validation(self):
|
||||
"""Test ListResourcesCacheEntry class functionality."""
|
||||
entry: ListResourcesCacheEntry = ListResourcesCacheEntry(
|
||||
key="test_key",
|
||||
value=[SAMPLE_RESOURCE],
|
||||
ttl=3600,
|
||||
)
|
||||
|
||||
dumped_entry = entry.model_dump()
|
||||
new_entry = ListResourcesCacheEntry.model_validate(dumped_entry)
|
||||
|
||||
assert new_entry.value == snapshot(
|
||||
[CachableResource(name="test_resource", uri=AnyUrl("https://test_uri/"))]
|
||||
)
|
||||
|
||||
def test_list_prompts_cache_entry_validation(self):
|
||||
"""Test ListPromptsCacheEntry class functionality."""
|
||||
entry: ListPromptsCacheEntry = ListPromptsCacheEntry(
|
||||
key="test_key",
|
||||
value=[SAMPLE_PROMPT],
|
||||
ttl=3600,
|
||||
)
|
||||
|
||||
dumped_entry = entry.model_dump()
|
||||
new_entry = ListPromptsCacheEntry.model_validate(dumped_entry)
|
||||
|
||||
assert new_entry.value == snapshot(
|
||||
[CachablePrompt(name="test_prompt", arguments=[])]
|
||||
)
|
||||
|
||||
|
||||
class TestMemoryCache:
|
||||
"""Test InMemoryCache implementation."""
|
||||
|
||||
async def test_size_limit(self, sample_tool_result):
|
||||
"""Test cache size limit enforcement."""
|
||||
cache = InMemoryCache(max_entries=2)
|
||||
|
||||
# Fill cache to capacity
|
||||
await cache.set_entry(
|
||||
cache_entry=ToolResultCacheEntry(
|
||||
key="key1", value=sample_tool_result, ttl=3600
|
||||
)
|
||||
)
|
||||
await cache.set_entry(
|
||||
cache_entry=ToolResultCacheEntry(
|
||||
key="key2", value=sample_tool_result, ttl=3600
|
||||
)
|
||||
)
|
||||
|
||||
# Add one more - should evict the first
|
||||
await cache.set_entry(
|
||||
cache_entry=ToolResultCacheEntry(
|
||||
key="key3", value=sample_tool_result, ttl=3600
|
||||
)
|
||||
)
|
||||
|
||||
assert len(cache._cache) == 2
|
||||
assert "tools/call:key1" not in cache._cache
|
||||
assert "tools/call:key2" in cache._cache
|
||||
assert "tools/call:key3" in cache._cache
|
||||
|
||||
|
||||
class TestCacheImplementations:
|
||||
"""Test InMemoryCache implementation."""
|
||||
|
||||
@pytest.fixture(params=["memory", "disk"])
|
||||
async def cache(self, request):
|
||||
if request.param == "memory":
|
||||
return InMemoryCache()
|
||||
else:
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
return DiskCache(path=temp_dir)
|
||||
|
||||
async def test_get_none_if_not_set(self, cache: CacheProtocol):
|
||||
"""Test that we get None if a value is not set."""
|
||||
assert await cache.get_entry(collection="tools/call", key="test_key") is None
|
||||
|
||||
async def test_list_tools(self, cache: CacheProtocol):
|
||||
"""Test that we can list tools from the cache."""
|
||||
await cache.set_entry(
|
||||
cache_entry=ListToolsCacheEntry(
|
||||
key="test_key", value=[SAMPLE_TOOL], ttl=3600
|
||||
)
|
||||
)
|
||||
result = await cache.get_entry(collection="tools/list", key="test_key")
|
||||
|
||||
assert result is not None
|
||||
assert dump_mcp_types(model=result.value) == snapshot(
|
||||
[
|
||||
{
|
||||
"name": "test_tool",
|
||||
"title": None,
|
||||
"description": None,
|
||||
"tags": set(),
|
||||
"meta": None,
|
||||
"enabled": True,
|
||||
"parameters": {"param1": "value1", "param2": 42},
|
||||
"output_schema": None,
|
||||
"annotations": None,
|
||||
"serializer": None,
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
async def test_tool_result(self, cache: CacheProtocol):
|
||||
"""Test that we can get a tool result from the cache."""
|
||||
await cache.set_entry(
|
||||
cache_entry=ToolResultCacheEntry(
|
||||
key="test_key", value=SAMPLE_TOOL_RESULT, ttl=3600
|
||||
)
|
||||
)
|
||||
result = await cache.get_entry(collection="tools/call", key="test_key")
|
||||
|
||||
assert result is not None
|
||||
assert dump_mcp_types(model=result.value) == snapshot(
|
||||
{
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "test_text",
|
||||
"annotations": None,
|
||||
"meta": None,
|
||||
}
|
||||
],
|
||||
"structured_content": {"result": "test_result"},
|
||||
}
|
||||
)
|
||||
|
||||
async def test_list_resources(self, cache: CacheProtocol):
|
||||
"""Test that we can list resources from the cache."""
|
||||
await cache.set_entry(
|
||||
cache_entry=ListResourcesCacheEntry(
|
||||
key="test_key", value=[SAMPLE_RESOURCE], ttl=3600
|
||||
)
|
||||
)
|
||||
result = await cache.get_entry(collection="resources/list", key="test_key")
|
||||
|
||||
assert result is not None
|
||||
assert dump_mcp_types(model=result.value) == snapshot(
|
||||
[
|
||||
{
|
||||
"name": "test_resource",
|
||||
"title": None,
|
||||
"description": None,
|
||||
"tags": set(),
|
||||
"meta": None,
|
||||
"enabled": True,
|
||||
"uri": AnyUrl("https://test_uri/"),
|
||||
"mime_type": "text/plain",
|
||||
"annotations": None,
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
async def test_read_resource(self, cache: CacheProtocol):
|
||||
"""Test that we can read a resource from the cache."""
|
||||
await cache.set_entry(
|
||||
cache_entry=ReadResourceCacheEntry(
|
||||
key="test_key", value=[SAMPLE_READ_RESOURCE_CONTENTS], ttl=3600
|
||||
)
|
||||
)
|
||||
result = await cache.get_entry(collection="resources/read", key="test_key")
|
||||
|
||||
assert result is not None
|
||||
assert dump_mcp_types(model=result.value) == snapshot(
|
||||
[{"content": "test_text", "mime_type": "text/plain"}]
|
||||
)
|
||||
|
||||
async def test_list_prompts(self, cache: CacheProtocol):
|
||||
"""Test that we can list prompts from the cache."""
|
||||
await cache.set_entry(
|
||||
cache_entry=ListPromptsCacheEntry(
|
||||
key="test_key", value=[SAMPLE_PROMPT], ttl=3600
|
||||
)
|
||||
)
|
||||
result = await cache.get_entry(collection="prompts/list", key="test_key")
|
||||
|
||||
assert result is not None
|
||||
assert dump_mcp_types(model=result.value) == snapshot(
|
||||
[
|
||||
{
|
||||
"name": "test_prompt",
|
||||
"title": None,
|
||||
"description": None,
|
||||
"tags": set(),
|
||||
"meta": None,
|
||||
"enabled": True,
|
||||
"arguments": [],
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
async def test_get_prompt(self, cache: CacheProtocol):
|
||||
"""Test that we can get a prompt from the cache."""
|
||||
entry = GetPromptCacheEntry(
|
||||
key="test_key", value=SAMPLE_GET_PROMPT_RESULT, ttl=3600
|
||||
)
|
||||
await cache.set_entry(cache_entry=entry)
|
||||
result = await cache.get_entry(collection=entry.collection, key="test_key")
|
||||
|
||||
assert result is not None
|
||||
assert dump_mcp_types(model=result.value) == snapshot(
|
||||
{
|
||||
"meta": None,
|
||||
"description": None,
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": {
|
||||
"type": "text",
|
||||
"text": "test_text",
|
||||
"annotations": None,
|
||||
"meta": None,
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
async def test_set_get_delete_get_value(self, cache: CacheProtocol):
|
||||
"""Test that we can set, get, delete, and get a value from the cache."""
|
||||
await cache.set_entry(
|
||||
cache_entry=ToolResultCacheEntry(
|
||||
key="test_key",
|
||||
value=SAMPLE_TOOL_RESULT,
|
||||
ttl=3600,
|
||||
)
|
||||
)
|
||||
result: BaseCacheEntry | None = await cache.get_entry(
|
||||
collection="tools/call", key="test_key"
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert dump_mcp_types(model=result.value) == dump_mcp_types(
|
||||
model=SAMPLE_TOOL_RESULT
|
||||
)
|
||||
|
||||
await cache.delete(collection="tools/call", key="test_key")
|
||||
|
||||
assert await cache.get_entry(collection="tools/call", key="test_key") is None
|
||||
|
||||
async def test_expiration_and_cleanup(self, cache: CacheProtocol):
|
||||
"""Test cache expiration and cleanup."""
|
||||
# Create an expired entry
|
||||
await cache.set_entry(
|
||||
cache_entry=ToolResultCacheEntry(
|
||||
key="expired_key",
|
||||
value=SAMPLE_TOOL_RESULT,
|
||||
ttl=-1,
|
||||
)
|
||||
)
|
||||
|
||||
# Should return None and remove expired entry
|
||||
result = await cache.get_entry(collection="tools/call", key="expired_key")
|
||||
|
||||
assert result is None
|
||||
|
||||
assert await cache.get_entry(collection="tools/call", key="expired_key") is None
|
||||
|
||||
|
||||
class TestResponseCachingMiddleware:
|
||||
"""Test ResponseCachingMiddleware functionality."""
|
||||
|
||||
|
|
@ -676,21 +316,42 @@ 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)
|
||||
|
||||
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,
|
||||
def test_method_settings(self):
|
||||
"""Test method TTL."""
|
||||
middleware = ResponseCachingMiddleware(
|
||||
method_settings={
|
||||
"list_tools": {"ttl": 100},
|
||||
"call_tool": {"enabled": False},
|
||||
},
|
||||
default_ttl=1000,
|
||||
)
|
||||
|
||||
assert middleware._stats.get_too_big("tools/call") == 1
|
||||
tool_list_settings = middleware._get_cache_settings(
|
||||
context=MiddlewareContext(method="tools/list", message=MagicMock())
|
||||
)
|
||||
assert tool_list_settings == {"ttl": 100}
|
||||
|
||||
call_tool_settings = middleware._get_cache_settings(
|
||||
context=MiddlewareContext(method="tools/call", message=MagicMock())
|
||||
)
|
||||
assert call_tool_settings == {"enabled": False}
|
||||
|
||||
other_methods = [
|
||||
"resources/list",
|
||||
"prompts/list",
|
||||
"resources/read",
|
||||
"prompts/get",
|
||||
]
|
||||
for method in other_methods:
|
||||
cache_settings = middleware._get_cache_settings(
|
||||
context=MiddlewareContext(method=method, message=MagicMock())
|
||||
)
|
||||
assert cache_settings is None
|
||||
|
||||
should_bypass = middleware._should_bypass_caching(
|
||||
context=MiddlewareContext(method=method, message=MagicMock())
|
||||
)
|
||||
assert should_bypass
|
||||
|
||||
def test_cache_key_generation(self):
|
||||
"""Test cache key generation."""
|
||||
|
|
@ -734,50 +395,11 @@ class TestResponseCachingMiddleware:
|
|||
"6306ff84fd3ff247a4bd91271e9d727d7f051bba53fb2e3bf80958988c4baf57"
|
||||
)
|
||||
|
||||
async def test_cache_miss_and_hit(
|
||||
self,
|
||||
):
|
||||
"""Test cache miss and hit scenarios."""
|
||||
middleware = ResponseCachingMiddleware()
|
||||
|
||||
mock_call_next = AsyncMock(
|
||||
return_value=ToolResult(
|
||||
content=[{"type": "text", "text": "test result"}],
|
||||
structured_content={"result": "success", "value": 123},
|
||||
)
|
||||
)
|
||||
|
||||
mock_context = MagicMock(
|
||||
spec=MiddlewareContext[mcp.types.CallToolRequestParams]
|
||||
)
|
||||
mock_context.message = mcp.types.CallToolRequestParams(
|
||||
name="test_tool", arguments={"param1": "value1", "param2": 42}
|
||||
)
|
||||
mock_context.method = "tools/call"
|
||||
|
||||
# First call - cache miss
|
||||
result1 = await middleware.on_call_tool(
|
||||
context=mock_context, call_next=mock_call_next
|
||||
)
|
||||
assert middleware._stats.get_misses("tools/call") == 1
|
||||
assert middleware._stats.get_hits("tools/call") == 0
|
||||
|
||||
# Second call - cache hit
|
||||
mock_call_next.reset_mock()
|
||||
result2 = await middleware.on_call_tool(
|
||||
context=mock_context, call_next=mock_call_next
|
||||
)
|
||||
|
||||
assert result1.content == result2.content
|
||||
assert not mock_call_next.called # Should not call downstream
|
||||
assert middleware._stats.get_hits("tools/call") == 1
|
||||
assert middleware._stats.get_misses("tools/call") == 1
|
||||
|
||||
|
||||
class TestResponseCachingMiddlewareIntegration:
|
||||
"""Integration tests with real FastMCP server."""
|
||||
|
||||
@pytest.fixture(params=["memory", "disk", "elasticsearch"])
|
||||
@pytest.fixture(params=["memory", "disk"])
|
||||
async def caching_server(
|
||||
self,
|
||||
tracking_calculator: TrackingCalculator,
|
||||
|
|
@ -787,10 +409,10 @@ class TestResponseCachingMiddlewareIntegration:
|
|||
mcp = FastMCP("CachingTestServer")
|
||||
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
disk_store = DiskStore(directory=temp_dir)
|
||||
response_caching_middleware = ResponseCachingMiddleware(
|
||||
cache_backend=DiskCache(path=temp_dir)
|
||||
if request.param == "disk"
|
||||
else InMemoryCache()
|
||||
cache_store=disk_store if request.param == "disk" else MemoryStore(),
|
||||
max_item_size=100000, # 100kb
|
||||
)
|
||||
|
||||
mcp.add_middleware(middleware=response_caching_middleware)
|
||||
|
|
@ -801,6 +423,8 @@ class TestResponseCachingMiddlewareIntegration:
|
|||
|
||||
yield mcp
|
||||
|
||||
await disk_store.close()
|
||||
|
||||
@pytest.fixture
|
||||
def non_caching_server(self, tracking_calculator: TrackingCalculator):
|
||||
"""Create a FastMCP server for non-caching tests."""
|
||||
|
|
@ -815,7 +439,7 @@ class TestResponseCachingMiddlewareIntegration:
|
|||
|
||||
async with Client(caching_server) as client:
|
||||
pre_tool_list: list[mcp.types.Tool] = await client.list_tools()
|
||||
assert len(pre_tool_list) == 3
|
||||
assert len(pre_tool_list) == 4
|
||||
|
||||
# Add a tool and make sure it's missing from the list tool response
|
||||
caching_server.add_tool(
|
||||
|
|
@ -823,7 +447,7 @@ class TestResponseCachingMiddlewareIntegration:
|
|||
)
|
||||
|
||||
post_tool_list: list[mcp.types.Tool] = await client.list_tools()
|
||||
assert len(post_tool_list) == 3
|
||||
assert len(post_tool_list) == 4
|
||||
|
||||
assert pre_tool_list == post_tool_list
|
||||
|
||||
|
|
@ -846,6 +470,26 @@ class TestResponseCachingMiddlewareIntegration:
|
|||
)
|
||||
assert call_tool_result_one == call_tool_result_two
|
||||
|
||||
async def test_call_tool_very_large_value(
|
||||
self,
|
||||
caching_server: FastMCP,
|
||||
tracking_calculator: TrackingCalculator,
|
||||
):
|
||||
"""Test that caching works with a real FastMCP server."""
|
||||
tracking_calculator.add_tools(fastmcp=caching_server)
|
||||
|
||||
async with Client[FastMCPTransport](caching_server) as client:
|
||||
call_tool_result_one: CallToolResult = await client.call_tool(
|
||||
"very_large_response", {}
|
||||
)
|
||||
|
||||
assert tracking_calculator.very_large_response_calls == 1
|
||||
call_tool_result_two: CallToolResult = await client.call_tool(
|
||||
"very_large_response", {}
|
||||
)
|
||||
assert call_tool_result_one == call_tool_result_two
|
||||
assert tracking_calculator.very_large_response_calls == 2
|
||||
|
||||
async def test_list_resources(
|
||||
self, caching_server: FastMCP, tracking_calculator: TrackingCalculator
|
||||
):
|
||||
|
|
@ -918,21 +562,3 @@ class TestResponseCachingMiddlewareIntegration:
|
|||
)
|
||||
|
||||
assert pre_prompt == post_prompt
|
||||
|
||||
|
||||
class TestCacheStats:
|
||||
"""Test CacheStats functionality."""
|
||||
|
||||
def test_stats_initialization(self):
|
||||
"""Test cache stats initialization."""
|
||||
stats = CacheStats(
|
||||
collections={
|
||||
"tools/call": CacheMethodStats(hits=5, misses=10),
|
||||
"tools/list": CacheMethodStats(hits=0, misses=0),
|
||||
}
|
||||
)
|
||||
|
||||
assert stats.get_hits("tools/call") == 5
|
||||
assert stats.get_misses("tools/call") == 10
|
||||
assert stats.get_hits("tools/list") == 0
|
||||
assert stats.get_misses("tools/list") == 0
|
||||
|
|
|
|||
2
uv.lock
generated
2
uv.lock
generated
|
|
@ -762,7 +762,7 @@ requires-dist = [
|
|||
{ name = "rich", specifier = ">=13.9.4" },
|
||||
{ name = "websockets", specifier = ">=15.0.1" },
|
||||
]
|
||||
provides-extras = ["openai"]
|
||||
provides-extras = ["caching", "contrib-middleware-elasticsearch-cache", "openai"]
|
||||
|
||||
[package.metadata.requires-dev]
|
||||
dev = [
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue