mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-09 07:09:11 +02:00
Add session-scoped state persistence (#2873)
This commit is contained in:
parent
266abef49e
commit
c8c84ff911
13 changed files with 744 additions and 167 deletions
|
|
@ -274,6 +274,50 @@ The environment variable for controlling the server banner has been renamed:
|
||||||
|
|
||||||
This change reflects that the setting now applies to all server startup methods, not just the CLI. The banner is now suppressed when running `python server.py` directly, not just when using `fastmcp run`.
|
This change reflects that the setting now applies to all server startup methods, not just the CLI. The banner is now suppressed when running `python server.py` directly, not just when using `fastmcp run`.
|
||||||
|
|
||||||
|
### Context State Methods Are Async
|
||||||
|
|
||||||
|
<Warning>
|
||||||
|
**Breaking Change:** `ctx.set_state()` and `ctx.get_state()` are now async methods. Synchronous calls will fail.
|
||||||
|
</Warning>
|
||||||
|
|
||||||
|
Context state has changed from request-scoped to session-scoped, persisting across multiple tool calls within the same MCP session. The methods are now async because they interact with a pluggable storage backend.
|
||||||
|
|
||||||
|
<CodeGroup>
|
||||||
|
```python Before
|
||||||
|
@mcp.tool
|
||||||
|
def my_tool(ctx: Context) -> str:
|
||||||
|
ctx.set_state("key", "value")
|
||||||
|
value = ctx.get_state("key")
|
||||||
|
return value
|
||||||
|
```
|
||||||
|
|
||||||
|
```python After
|
||||||
|
@mcp.tool
|
||||||
|
async def my_tool(ctx: Context) -> str:
|
||||||
|
await ctx.set_state("key", "value")
|
||||||
|
value = await ctx.get_state("key")
|
||||||
|
return value
|
||||||
|
```
|
||||||
|
</CodeGroup>
|
||||||
|
|
||||||
|
**What changed:**
|
||||||
|
- State now persists across requests within a session (not just within a single request)
|
||||||
|
- Different clients have isolated state (keyed by session ID)
|
||||||
|
- State expires after 1 day to prevent unbounded memory growth
|
||||||
|
- New method: `await ctx.delete_state(key)`
|
||||||
|
|
||||||
|
**Custom storage backends:**
|
||||||
|
|
||||||
|
By default, state uses an in-memory store. For distributed deployments, provide a custom backend:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from key_value.aio.stores.redis import RedisStore
|
||||||
|
|
||||||
|
mcp = FastMCP("server", session_state_store=RedisStore(...))
|
||||||
|
```
|
||||||
|
|
||||||
|
See [Session State](/servers/context#session-state) for full documentation.
|
||||||
|
|
||||||
## v2.14.0
|
## v2.14.0
|
||||||
|
|
||||||
### OpenAPI Parser Promotion
|
### OpenAPI Parser Promotion
|
||||||
|
|
|
||||||
|
|
@ -150,6 +150,37 @@ Documentation: `docs/servers/providers/transforms.mdx`, `docs/servers/visibility
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
## Session-Scoped State
|
||||||
|
|
||||||
|
v3.0 changes context state from request-scoped to session-scoped. State now persists across multiple tool calls within the same MCP session.
|
||||||
|
|
||||||
|
```python
|
||||||
|
@mcp.tool
|
||||||
|
async def increment_counter(ctx: Context) -> int:
|
||||||
|
count = await ctx.get_state("counter") or 0
|
||||||
|
await ctx.set_state("counter", count + 1)
|
||||||
|
return count + 1
|
||||||
|
```
|
||||||
|
|
||||||
|
State is automatically keyed by session ID, ensuring isolation between different clients. The implementation uses [pykeyvalue](https://github.com/strawgate/py-key-value) for pluggable storage backends:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from key_value.aio.stores.redis import RedisStore
|
||||||
|
|
||||||
|
# Use Redis for distributed deployments
|
||||||
|
mcp = FastMCP("server", session_state_store=RedisStore(...))
|
||||||
|
```
|
||||||
|
|
||||||
|
**Key details:**
|
||||||
|
- Methods are now async: `await ctx.get_state()`, `await ctx.set_state()`, `await ctx.delete_state()`
|
||||||
|
- State expires after 1 day (TTL) to prevent unbounded memory growth
|
||||||
|
- Works during `on_initialize` middleware when using the same session object
|
||||||
|
- For distributed HTTP, session identity comes from the `mcp-session-id` header
|
||||||
|
|
||||||
|
Documentation: `docs/servers/context.mdx`
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
## Visibility System
|
## Visibility System
|
||||||
|
|
||||||
Components can be dynamically enabled/disabled at runtime using the visibility system ([#2708](https://github.com/jlowin/fastmcp/pull/2708)).
|
Components can be dynamically enabled/disabled at runtime using the visibility system ([#2708](https://github.com/jlowin/fastmcp/pull/2708)).
|
||||||
|
|
@ -727,3 +758,19 @@ See `docs/development/v3-notes/auth-provider-env-vars.mdx` for rationale.
|
||||||
`FASTMCP_SHOW_CLI_BANNER` → `FASTMCP_SHOW_SERVER_BANNER` ([#2771](https://github.com/jlowin/fastmcp/pull/2771))
|
`FASTMCP_SHOW_CLI_BANNER` → `FASTMCP_SHOW_SERVER_BANNER` ([#2771](https://github.com/jlowin/fastmcp/pull/2771))
|
||||||
|
|
||||||
Now applies to all server startup methods, not just the CLI.
|
Now applies to all server startup methods, not just the CLI.
|
||||||
|
|
||||||
|
### Context State Methods Are Async
|
||||||
|
|
||||||
|
`ctx.set_state()` and `ctx.get_state()` are now async and session-scoped:
|
||||||
|
|
||||||
|
```python
|
||||||
|
# v2.x
|
||||||
|
ctx.set_state("key", "value")
|
||||||
|
value = ctx.get_state("key")
|
||||||
|
|
||||||
|
# v3.0
|
||||||
|
await ctx.set_state("key", "value")
|
||||||
|
value = await ctx.get_state("key")
|
||||||
|
```
|
||||||
|
|
||||||
|
State now persists across requests within a session. See "Session-Scoped State" above.
|
||||||
|
|
|
||||||
|
|
@ -22,7 +22,7 @@ The `Context` object provides a clean interface to access MCP features within yo
|
||||||
- **Prompt Access**: List and retrieve prompts registered with the server
|
- **Prompt Access**: List and retrieve prompts registered with the server
|
||||||
- **LLM Sampling**: Request the client's LLM to generate text based on provided messages
|
- **LLM Sampling**: Request the client's LLM to generate text based on provided messages
|
||||||
- **User Elicitation**: Request structured input from users during tool execution
|
- **User Elicitation**: Request structured input from users during tool execution
|
||||||
- **State Management**: Store and share data between middleware and the handler within a single request
|
- **Session State**: Store data that persists across requests within an MCP session
|
||||||
- **Request Information**: Access metadata about the current request
|
- **Request Information**: Access metadata about the current request
|
||||||
- **Server Access**: When needed, access the underlying FastMCP server instance
|
- **Server Access**: When needed, access the underlying FastMCP server instance
|
||||||
|
|
||||||
|
|
@ -209,55 +209,57 @@ messages = result.messages
|
||||||
- **`ctx.list_prompts() -> list[MCPPrompt]`**: Returns list of all available prompts
|
- **`ctx.list_prompts() -> list[MCPPrompt]`**: Returns list of all available prompts
|
||||||
- **`ctx.get_prompt(name: str, arguments: dict[str, Any] | None = None) -> GetPromptResult`**: Get a specific prompt with optional arguments
|
- **`ctx.get_prompt(name: str, arguments: dict[str, Any] | None = None) -> GetPromptResult`**: Get a specific prompt with optional arguments
|
||||||
|
|
||||||
### State Management
|
### Session State
|
||||||
|
|
||||||
<VersionBadge version="2.11.0" />
|
<VersionBadge version="3.0.0" />
|
||||||
|
|
||||||
Store and share data between middleware and handlers within a single MCP request. Each MCP request (such as calling a tool, reading a resource, listing tools, or listing resources) receives its own context object with isolated state. Context state is particularly useful for passing information from [middleware](/servers/middleware) to your handlers.
|
Store data that persists across multiple requests within the same MCP session. Session state is automatically keyed by the client's session, ensuring isolation between different clients.
|
||||||
|
|
||||||
To store a value in the context state, use `ctx.set_state(key, value)`. To retrieve a value, use `ctx.get_state(key)`.
|
```python
|
||||||
|
from fastmcp import FastMCP, Context
|
||||||
|
|
||||||
<Warning>
|
mcp = FastMCP("stateful-app")
|
||||||
Context state is scoped to a single MCP request. Each operation (tool call, resource read, list operation, etc.) receives a new context object. State set during one request will not be available in subsequent requests. For persistent data storage across requests, use external storage mechanisms like databases, files, or in-memory caches.
|
|
||||||
</Warning>
|
|
||||||
|
|
||||||
This simplified example shows how to use MCP middleware to store user info in the context state, and how to access that state in a tool:
|
|
||||||
|
|
||||||
```python {7-8, 16-17}
|
|
||||||
from fastmcp.server.middleware import Middleware, MiddlewareContext
|
|
||||||
|
|
||||||
class UserAuthMiddleware(Middleware):
|
|
||||||
async def on_call_tool(self, context: MiddlewareContext, call_next):
|
|
||||||
|
|
||||||
# Middleware stores user info in context state
|
|
||||||
context.fastmcp_context.set_state("user_id", "user_123")
|
|
||||||
context.fastmcp_context.set_state("permissions", ["read", "write"])
|
|
||||||
|
|
||||||
return await call_next(context)
|
|
||||||
|
|
||||||
@mcp.tool
|
@mcp.tool
|
||||||
async def secure_operation(data: str, ctx: Context) -> str:
|
async def increment_counter(ctx: Context) -> int:
|
||||||
"""Tool can access state set by middleware."""
|
"""Increment a counter that persists across tool calls."""
|
||||||
|
count = await ctx.get_state("counter") or 0
|
||||||
|
await ctx.set_state("counter", count + 1)
|
||||||
|
return count + 1
|
||||||
|
|
||||||
user_id = ctx.get_state("user_id") # "user_123"
|
@mcp.tool
|
||||||
permissions = ctx.get_state("permissions") # ["read", "write"]
|
async def get_counter(ctx: Context) -> int:
|
||||||
|
"""Get the current counter value."""
|
||||||
if "write" not in permissions:
|
return await ctx.get_state("counter") or 0
|
||||||
return "Access denied"
|
|
||||||
|
|
||||||
return f"Processing {data} for user {user_id}"
|
|
||||||
```
|
```
|
||||||
|
|
||||||
|
Each client session has its own isolated state—two different clients calling `increment_counter` will each have their own counter.
|
||||||
|
|
||||||
**Method signatures:**
|
**Method signatures:**
|
||||||
- **`ctx.set_state(key: str, value: Any) -> None`**: Store a value in the context state
|
- **`await ctx.set_state(key: str, value: Any) -> None`**: Store a value in session state
|
||||||
- **`ctx.get_state(key: str) -> Any`**: Retrieve a value from the context state (returns None if not found)
|
- **`await ctx.get_state(key: str) -> Any`**: Retrieve a value (returns None if not found)
|
||||||
|
- **`await ctx.delete_state(key: str) -> None`**: Remove a value from session state
|
||||||
|
|
||||||
**State Inheritance:**
|
<Note>
|
||||||
When a new context is created (nested contexts), it inherits a copy of its parent's state. This ensures that:
|
State methods are async and require `await`. State expires after 1 day to prevent unbounded memory growth.
|
||||||
- State set on a child context never affects the parent context
|
</Note>
|
||||||
- State set on a parent context after the child context is initialized is not propagated to the child context
|
|
||||||
|
|
||||||
This makes state management predictable and prevents unexpected side effects between nested operations.
|
#### Custom Storage Backends
|
||||||
|
|
||||||
|
By default, session state uses an in-memory store suitable for single-server deployments. For distributed or serverless deployments, provide a custom storage backend:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from key_value.aio.stores.redis import RedisStore
|
||||||
|
|
||||||
|
# Use Redis for distributed state
|
||||||
|
mcp = FastMCP("distributed-app", session_state_store=RedisStore(...))
|
||||||
|
```
|
||||||
|
|
||||||
|
Any backend compatible with the [py-key-value-aio](https://github.com/strawgate/py-key-value) `AsyncKeyValue` protocol works. See [Storage Backends](/servers/storage-backends) for more options including Redis, DynamoDB, and MongoDB.
|
||||||
|
|
||||||
|
#### State During Initialization
|
||||||
|
|
||||||
|
State set during `on_initialize` middleware persists to subsequent tool calls when using the same session object (STDIO, SSE, single-server HTTP). For distributed/serverless HTTP deployments where different machines handle init and tool calls, state is isolated by the `mcp-session-id` header.
|
||||||
|
|
||||||
### Change Notifications
|
### Change Notifications
|
||||||
|
|
||||||
|
|
@ -345,7 +347,7 @@ async def request_info(ctx: Context) -> dict:
|
||||||
|
|
||||||
- **`ctx.request_id -> str`**: Get the unique ID for the current MCP request
|
- **`ctx.request_id -> str`**: Get the unique ID for the current MCP request
|
||||||
- **`ctx.client_id -> str | None`**: Get the ID of the client making the request, if provided during initialization
|
- **`ctx.client_id -> str | None`**: Get the ID of the client making the request, if provided during initialization
|
||||||
- **`ctx.session_id -> str | None`**: Get the MCP session ID for session-based data sharing (HTTP transports only)
|
- **`ctx.session_id -> str`**: Get the MCP session ID for session-based data sharing. Raises `RuntimeError` if the MCP session is not yet established.
|
||||||
|
|
||||||
#### Request Context Availability
|
#### Request Context Availability
|
||||||
|
|
||||||
|
|
|
||||||
51
examples/persistent_state/README.md
Normal file
51
examples/persistent_state/README.md
Normal file
|
|
@ -0,0 +1,51 @@
|
||||||
|
# Persistent Session State
|
||||||
|
|
||||||
|
This example demonstrates session-scoped state that persists across tool calls within the same MCP session.
|
||||||
|
|
||||||
|
## What it shows
|
||||||
|
|
||||||
|
- State set in one tool call is readable in subsequent calls
|
||||||
|
- Different clients have isolated state (same keys, different values)
|
||||||
|
- Reconnecting creates a new session with fresh state
|
||||||
|
|
||||||
|
## Running
|
||||||
|
|
||||||
|
**HTTP transport:**
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Terminal 1: Start the server
|
||||||
|
uv run python server.py
|
||||||
|
|
||||||
|
# Terminal 2: Run the client
|
||||||
|
uv run python client.py
|
||||||
|
```
|
||||||
|
|
||||||
|
**STDIO transport (in-process):**
|
||||||
|
|
||||||
|
```bash
|
||||||
|
uv run python client_stdio.py
|
||||||
|
```
|
||||||
|
|
||||||
|
## Example output
|
||||||
|
|
||||||
|
```text
|
||||||
|
Each line below is a separate tool call
|
||||||
|
|
||||||
|
Alice connects
|
||||||
|
session a9f6eaa3
|
||||||
|
set user = Alice
|
||||||
|
set secret = alice-password
|
||||||
|
get user → Alice
|
||||||
|
get secret → alice-password
|
||||||
|
|
||||||
|
Bob connects (different session)
|
||||||
|
session 0c3bffc5
|
||||||
|
get user → not found
|
||||||
|
get secret → not found
|
||||||
|
set user = Bob
|
||||||
|
get user → Bob
|
||||||
|
|
||||||
|
Alice reconnects (new session)
|
||||||
|
session e39640e3
|
||||||
|
get user → not found
|
||||||
|
```
|
||||||
85
examples/persistent_state/client.py
Normal file
85
examples/persistent_state/client.py
Normal file
|
|
@ -0,0 +1,85 @@
|
||||||
|
"""Client for testing persistent state.
|
||||||
|
|
||||||
|
Run the server first:
|
||||||
|
uv run python examples/persistent_state/server.py
|
||||||
|
|
||||||
|
Then run this client:
|
||||||
|
uv run python examples/persistent_state/client.py
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
from rich.console import Console
|
||||||
|
|
||||||
|
from fastmcp import Client
|
||||||
|
from fastmcp.client.transports import StreamableHttpTransport
|
||||||
|
|
||||||
|
URL = "http://127.0.0.1:8000/mcp"
|
||||||
|
console = Console()
|
||||||
|
|
||||||
|
|
||||||
|
async def main() -> None:
|
||||||
|
console.print()
|
||||||
|
console.print("[dim italic]Each line below is a separate tool call[/dim italic]")
|
||||||
|
console.print()
|
||||||
|
|
||||||
|
# --- Alice's session ---
|
||||||
|
console.print("[dim]Alice connects[/dim]")
|
||||||
|
|
||||||
|
transport1 = StreamableHttpTransport(url=URL)
|
||||||
|
async with Client(transport=transport1) as alice:
|
||||||
|
result = await alice.call_tool("list_session_info", {})
|
||||||
|
console.print(f" session [cyan]{result.data['session_id'][:8]}[/cyan]")
|
||||||
|
|
||||||
|
await alice.call_tool("set_value", {"key": "user", "value": "Alice"})
|
||||||
|
console.print(" set [white]user[/white] = [green]Alice[/green]")
|
||||||
|
|
||||||
|
await alice.call_tool("set_value", {"key": "secret", "value": "alice-password"})
|
||||||
|
console.print(" set [white]secret[/white] = [green]alice-password[/green]")
|
||||||
|
|
||||||
|
result = await alice.call_tool("get_value", {"key": "user"})
|
||||||
|
console.print(" get [white]user[/white] → [green]Alice[/green]")
|
||||||
|
|
||||||
|
result = await alice.call_tool("get_value", {"key": "secret"})
|
||||||
|
console.print(" get [white]secret[/white] → [green]alice-password[/green]")
|
||||||
|
|
||||||
|
console.print()
|
||||||
|
|
||||||
|
# --- Bob's session ---
|
||||||
|
console.print("[dim]Bob connects (different session)[/dim]")
|
||||||
|
|
||||||
|
transport2 = StreamableHttpTransport(url=URL)
|
||||||
|
async with Client(transport=transport2) as bob:
|
||||||
|
result = await bob.call_tool("list_session_info", {})
|
||||||
|
console.print(f" session [cyan]{result.data['session_id'][:8]}[/cyan]")
|
||||||
|
|
||||||
|
await bob.call_tool("get_value", {"key": "user"})
|
||||||
|
console.print(" get [white]user[/white] → [dim]not found[/dim]")
|
||||||
|
|
||||||
|
await bob.call_tool("get_value", {"key": "secret"})
|
||||||
|
console.print(" get [white]secret[/white] → [dim]not found[/dim]")
|
||||||
|
|
||||||
|
await bob.call_tool("set_value", {"key": "user", "value": "Bob"})
|
||||||
|
console.print(" set [white]user[/white] = [green]Bob[/green]")
|
||||||
|
|
||||||
|
await bob.call_tool("get_value", {"key": "user"})
|
||||||
|
console.print(" get [white]user[/white] → [green]Bob[/green]")
|
||||||
|
|
||||||
|
console.print()
|
||||||
|
|
||||||
|
# --- Alice reconnects ---
|
||||||
|
console.print("[dim]Alice reconnects (new session)[/dim]")
|
||||||
|
|
||||||
|
transport3 = StreamableHttpTransport(url=URL)
|
||||||
|
async with Client(transport=transport3) as alice_again:
|
||||||
|
result = await alice_again.call_tool("list_session_info", {})
|
||||||
|
console.print(f" session [cyan]{result.data['session_id'][:8]}[/cyan]")
|
||||||
|
|
||||||
|
await alice_again.call_tool("get_value", {"key": "user"})
|
||||||
|
console.print(" get [white]user[/white] → [dim]not found[/dim]")
|
||||||
|
|
||||||
|
console.print()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
asyncio.run(main())
|
||||||
88
examples/persistent_state/client_stdio.py
Normal file
88
examples/persistent_state/client_stdio.py
Normal file
|
|
@ -0,0 +1,88 @@
|
||||||
|
"""Client for testing persistent state over STDIO.
|
||||||
|
|
||||||
|
Run directly:
|
||||||
|
uv run python examples/persistent_state/client_stdio.py
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from rich.console import Console
|
||||||
|
|
||||||
|
from fastmcp import Client, FastMCP
|
||||||
|
|
||||||
|
# Add parent directory to path for importing the server module
|
||||||
|
examples_dir: Path = Path(__file__).parent.parent.parent
|
||||||
|
if str(examples_dir) not in sys.path:
|
||||||
|
sys.path.insert(0, str(examples_dir))
|
||||||
|
|
||||||
|
import examples.persistent_state.server as server_module # noqa: E402
|
||||||
|
|
||||||
|
server: FastMCP = server_module.server
|
||||||
|
|
||||||
|
console: Console = Console()
|
||||||
|
|
||||||
|
|
||||||
|
async def main() -> None:
|
||||||
|
console.print()
|
||||||
|
console.print("[dim italic]Each line below is a separate tool call[/dim italic]")
|
||||||
|
console.print()
|
||||||
|
|
||||||
|
# --- Alice's session ---
|
||||||
|
console.print("[dim]Alice connects[/dim]")
|
||||||
|
|
||||||
|
async with Client(server) as alice:
|
||||||
|
result = await alice.call_tool("list_session_info", {})
|
||||||
|
console.print(f" session [cyan]{result.data['session_id'][:8]}[/cyan]")
|
||||||
|
|
||||||
|
await alice.call_tool("set_value", {"key": "user", "value": "Alice"})
|
||||||
|
console.print(" set [white]user[/white] = [green]Alice[/green]")
|
||||||
|
|
||||||
|
await alice.call_tool("set_value", {"key": "secret", "value": "alice-password"})
|
||||||
|
console.print(" set [white]secret[/white] = [green]alice-password[/green]")
|
||||||
|
|
||||||
|
await alice.call_tool("get_value", {"key": "user"})
|
||||||
|
console.print(" get [white]user[/white] → [green]Alice[/green]")
|
||||||
|
|
||||||
|
await alice.call_tool("get_value", {"key": "secret"})
|
||||||
|
console.print(" get [white]secret[/white] → [green]alice-password[/green]")
|
||||||
|
|
||||||
|
console.print()
|
||||||
|
|
||||||
|
# --- Bob's session ---
|
||||||
|
console.print("[dim]Bob connects (different session)[/dim]")
|
||||||
|
|
||||||
|
async with Client(server) as bob:
|
||||||
|
result = await bob.call_tool("list_session_info", {})
|
||||||
|
console.print(f" session [cyan]{result.data['session_id'][:8]}[/cyan]")
|
||||||
|
|
||||||
|
await bob.call_tool("get_value", {"key": "user"})
|
||||||
|
console.print(" get [white]user[/white] → [dim]not found[/dim]")
|
||||||
|
|
||||||
|
await bob.call_tool("get_value", {"key": "secret"})
|
||||||
|
console.print(" get [white]secret[/white] → [dim]not found[/dim]")
|
||||||
|
|
||||||
|
await bob.call_tool("set_value", {"key": "user", "value": "Bob"})
|
||||||
|
console.print(" set [white]user[/white] = [green]Bob[/green]")
|
||||||
|
|
||||||
|
await bob.call_tool("get_value", {"key": "user"})
|
||||||
|
console.print(" get [white]user[/white] → [green]Bob[/green]")
|
||||||
|
|
||||||
|
console.print()
|
||||||
|
|
||||||
|
# --- Alice reconnects ---
|
||||||
|
console.print("[dim]Alice reconnects (new session)[/dim]")
|
||||||
|
|
||||||
|
async with Client(server) as alice_again:
|
||||||
|
result = await alice_again.call_tool("list_session_info", {})
|
||||||
|
console.print(f" session [cyan]{result.data['session_id'][:8]}[/cyan]")
|
||||||
|
|
||||||
|
await alice_again.call_tool("get_value", {"key": "user"})
|
||||||
|
console.print(" get [white]user[/white] → [dim]not found[/dim]")
|
||||||
|
|
||||||
|
console.print()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
asyncio.run(main())
|
||||||
42
examples/persistent_state/server.py
Normal file
42
examples/persistent_state/server.py
Normal file
|
|
@ -0,0 +1,42 @@
|
||||||
|
"""Example: Persistent session-scoped state.
|
||||||
|
|
||||||
|
This demonstrates using Context.get_state() and set_state() to store
|
||||||
|
data that persists across tool calls within the same MCP session.
|
||||||
|
|
||||||
|
Run with:
|
||||||
|
uv run python examples/persistent_state/server.py
|
||||||
|
"""
|
||||||
|
|
||||||
|
from fastmcp import FastMCP
|
||||||
|
from fastmcp.server.context import Context
|
||||||
|
|
||||||
|
server = FastMCP("StateExample")
|
||||||
|
|
||||||
|
|
||||||
|
@server.tool
|
||||||
|
async def set_value(key: str, value: str, ctx: Context) -> str:
|
||||||
|
"""Store a value in session state."""
|
||||||
|
await ctx.set_state(key, value)
|
||||||
|
return f"Stored '{key}' = '{value}'"
|
||||||
|
|
||||||
|
|
||||||
|
@server.tool
|
||||||
|
async def get_value(key: str, ctx: Context) -> str:
|
||||||
|
"""Retrieve a value from session state."""
|
||||||
|
value = await ctx.get_state(key)
|
||||||
|
if value is None:
|
||||||
|
return f"Key '{key}' not found"
|
||||||
|
return f"'{key}' = '{value}'"
|
||||||
|
|
||||||
|
|
||||||
|
@server.tool
|
||||||
|
async def list_session_info(ctx: Context) -> dict[str, str | None]:
|
||||||
|
"""Get information about the current session."""
|
||||||
|
return {
|
||||||
|
"session_id": ctx.session_id,
|
||||||
|
"transport": ctx.transport,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
server.run(transport="streamable-http")
|
||||||
|
|
@ -172,6 +172,8 @@ extend-select = [
|
||||||
"UP", # flake8-unused-imports: Catches unused imports
|
"UP", # flake8-unused-imports: Catches unused imports
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[tool.ruff.lint.isort]
|
||||||
|
known-first-party = ["fastmcp"]
|
||||||
|
|
||||||
[tool.ruff.lint.per-file-ignores]
|
[tool.ruff.lint.per-file-ignores]
|
||||||
"__init__.py" = ["F401", "I001", "RUF013"]
|
"__init__.py" = ["F401", "I001", "RUF013"]
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,5 @@
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import copy
|
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import weakref
|
import weakref
|
||||||
|
|
@ -55,7 +54,7 @@ from fastmcp.server.sampling.run import (
|
||||||
from fastmcp.server.sampling.run import (
|
from fastmcp.server.sampling.run import (
|
||||||
execute_tools as run_sampling_tools,
|
execute_tools as run_sampling_tools,
|
||||||
)
|
)
|
||||||
from fastmcp.server.server import FastMCP
|
from fastmcp.server.server import FastMCP, StateValue
|
||||||
from fastmcp.utilities.json_schema import compress_schema
|
from fastmcp.utilities.json_schema import compress_schema
|
||||||
from fastmcp.utilities.logging import _clamp_logger, get_logger
|
from fastmcp.utilities.logging import _clamp_logger, get_logger
|
||||||
from fastmcp.utilities.types import get_cached_typeadapter
|
from fastmcp.utilities.types import get_cached_typeadapter
|
||||||
|
|
@ -156,29 +155,36 @@ class Context:
|
||||||
request_id = ctx.request_id
|
request_id = ctx.request_id
|
||||||
client_id = ctx.client_id
|
client_id = ctx.client_id
|
||||||
|
|
||||||
# Manage state across the request
|
# Manage state across the session (persists across requests)
|
||||||
ctx.set_state("key", "value")
|
await ctx.set_state("key", "value")
|
||||||
value = ctx.get_state("key")
|
value = await ctx.get_state("key")
|
||||||
|
|
||||||
return str(x)
|
return str(x)
|
||||||
```
|
```
|
||||||
|
|
||||||
State Management:
|
State Management:
|
||||||
Context objects maintain a state dictionary that can be used to store and share
|
Context provides session-scoped state that persists across requests within
|
||||||
data across middleware and tool calls within a request. When a new context
|
the same MCP session. State is automatically keyed by session, ensuring
|
||||||
is created (nested contexts), it inherits a copy of its parent's state, ensuring
|
isolation between different clients.
|
||||||
that modifications in child contexts don't affect parent contexts.
|
|
||||||
|
State set during `on_initialize` middleware will persist to subsequent tool
|
||||||
|
calls when using the same session object (STDIO, SSE, single-server HTTP).
|
||||||
|
For distributed/serverless HTTP deployments where different machines handle
|
||||||
|
the init and tool calls, state is isolated by the mcp-session-id header.
|
||||||
|
|
||||||
The context parameter name can be anything as long as it's annotated with Context.
|
The context parameter name can be anything as long as it's annotated with Context.
|
||||||
The context is optional - tools that don't need it can omit the parameter.
|
The context is optional - tools that don't need it can omit the parameter.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, fastmcp: FastMCP):
|
# Default TTL for session state: 1 day in seconds
|
||||||
|
_STATE_TTL_SECONDS: int = 86400
|
||||||
|
|
||||||
|
def __init__(self, fastmcp: FastMCP, session: ServerSession | None = None):
|
||||||
self._fastmcp: weakref.ref[FastMCP] = weakref.ref(fastmcp)
|
self._fastmcp: weakref.ref[FastMCP] = weakref.ref(fastmcp)
|
||||||
|
self._session: ServerSession | None = session # For state ops during init
|
||||||
self._tokens: list[Token] = []
|
self._tokens: list[Token] = []
|
||||||
self._notification_queue: list[mcp.types.ServerNotificationType] = []
|
self._notification_queue: list[mcp.types.ServerNotificationType] = []
|
||||||
self._state: dict[str, Any] = {}
|
|
||||||
self._exit_stack: AsyncExitStack | None = None
|
self._exit_stack: AsyncExitStack | None = None
|
||||||
self._cancel_scope: anyio.CancelScope | None = None
|
self._cancel_scope: anyio.CancelScope | None = None
|
||||||
|
|
||||||
|
|
@ -192,11 +198,6 @@ class Context:
|
||||||
|
|
||||||
async def __aenter__(self) -> Context:
|
async def __aenter__(self) -> Context:
|
||||||
"""Enter the context manager and set this context as the current context."""
|
"""Enter the context manager and set this context as the current context."""
|
||||||
parent_context = _current_context.get(None)
|
|
||||||
if parent_context is not None:
|
|
||||||
# Inherit state from parent context
|
|
||||||
self._state = copy.deepcopy(parent_context._state)
|
|
||||||
|
|
||||||
# Always set this context and save the token
|
# Always set this context and save the token
|
||||||
token = _current_context.set(self)
|
token = _current_context.set(self)
|
||||||
self._tokens.append(token)
|
self._tokens.append(token)
|
||||||
|
|
@ -462,7 +463,7 @@ class Context:
|
||||||
for other transports.
|
for other transports.
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
RuntimeError if MCP request context is not available.
|
RuntimeError if no session is available.
|
||||||
|
|
||||||
Example:
|
Example:
|
||||||
```python
|
```python
|
||||||
|
|
@ -473,32 +474,37 @@ class Context:
|
||||||
return f"Data stored for session {session_id}"
|
return f"Data stored for session {session_id}"
|
||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
request_ctx = self.request_context
|
from uuid import uuid4
|
||||||
if request_ctx is None:
|
|
||||||
raise RuntimeError(
|
|
||||||
"session_id is not available because the MCP session has not been established yet. "
|
|
||||||
"Check `context.request_context` for None before accessing this attribute."
|
|
||||||
)
|
|
||||||
session = request_ctx.session
|
|
||||||
|
|
||||||
# Try to get the session ID from the session attributes
|
# Get session from request context or _session (for on_initialize)
|
||||||
session_id = getattr(session, "_fastmcp_id", None)
|
request_ctx = self.request_context
|
||||||
|
if request_ctx is not None:
|
||||||
|
session = request_ctx.session
|
||||||
|
elif self._session is not None:
|
||||||
|
session = self._session
|
||||||
|
else:
|
||||||
|
raise RuntimeError(
|
||||||
|
"session_id is not available because no session exists. "
|
||||||
|
"This typically means you're outside a request context."
|
||||||
|
)
|
||||||
|
|
||||||
|
# Check for cached session ID
|
||||||
|
session_id = getattr(session, "_fastmcp_state_prefix", None)
|
||||||
if session_id is not None:
|
if session_id is not None:
|
||||||
return session_id
|
return session_id
|
||||||
|
|
||||||
# Try to get the session ID from the http request headers
|
# For HTTP, try to get from header
|
||||||
request = request_ctx.request
|
if request_ctx is not None:
|
||||||
if request:
|
request = request_ctx.request
|
||||||
session_id = request.headers.get("mcp-session-id")
|
if request:
|
||||||
|
session_id = request.headers.get("mcp-session-id")
|
||||||
|
|
||||||
# Generate a session ID if it doesn't exist.
|
# For STDIO/SSE/in-memory, generate a UUID
|
||||||
if session_id is None:
|
if session_id is None:
|
||||||
from uuid import uuid4
|
|
||||||
|
|
||||||
session_id = str(uuid4())
|
session_id = str(uuid4())
|
||||||
|
|
||||||
# Save the session id to the session attributes
|
# Cache on session for consistency
|
||||||
session._fastmcp_id = session_id # type: ignore[attr-defined]
|
session._fastmcp_state_prefix = session_id # type: ignore[attr-defined]
|
||||||
return session_id
|
return session_id
|
||||||
|
|
||||||
@property
|
@property
|
||||||
|
|
@ -1112,13 +1118,37 @@ class Context:
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unexpected elicitation action: {result.action}")
|
raise ValueError(f"Unexpected elicitation action: {result.action}")
|
||||||
|
|
||||||
def set_state(self, key: str, value: Any) -> None:
|
def _make_state_key(self, key: str) -> str:
|
||||||
"""Set a value in the context state."""
|
"""Create session-prefixed key for state storage."""
|
||||||
self._state[key] = value
|
return f"{self.session_id}:{key}"
|
||||||
|
|
||||||
def get_state(self, key: str) -> Any:
|
async def set_state(self, key: str, value: Any) -> None:
|
||||||
"""Get a value from the context state. Returns None if the key is not found."""
|
"""Set a value in the session-scoped state store.
|
||||||
return self._state.get(key)
|
|
||||||
|
Values persist across requests within the same MCP session.
|
||||||
|
The key is automatically prefixed with the session identifier.
|
||||||
|
State expires after 1 day to prevent unbounded memory growth.
|
||||||
|
"""
|
||||||
|
prefixed_key = self._make_state_key(key)
|
||||||
|
await self.fastmcp._state_store.put(
|
||||||
|
key=prefixed_key,
|
||||||
|
value=StateValue(value=value),
|
||||||
|
ttl=self._STATE_TTL_SECONDS,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def get_state(self, key: str) -> Any:
|
||||||
|
"""Get a value from the session-scoped state store.
|
||||||
|
|
||||||
|
Returns None if the key is not found.
|
||||||
|
"""
|
||||||
|
prefixed_key = self._make_state_key(key)
|
||||||
|
result = await self.fastmcp._state_store.get(key=prefixed_key)
|
||||||
|
return result.value if result is not None else None
|
||||||
|
|
||||||
|
async def delete_state(self, key: str) -> None:
|
||||||
|
"""Delete a value from the session-scoped state store."""
|
||||||
|
prefixed_key = self._make_state_key(key)
|
||||||
|
await self.fastmcp._state_store.delete(key=prefixed_key)
|
||||||
|
|
||||||
async def _periodic_flush(self) -> None:
|
async def _periodic_flush(self) -> None:
|
||||||
"""Background task that flushes the notification queue every second."""
|
"""Background task that flushes the notification queue every second."""
|
||||||
|
|
|
||||||
|
|
@ -96,7 +96,7 @@ class MiddlewareServerSession(ServerSession):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async with fastmcp.server.context.Context(
|
async with fastmcp.server.context.Context(
|
||||||
fastmcp=self.fastmcp
|
fastmcp=self.fastmcp, session=self
|
||||||
) as fastmcp_ctx:
|
) as fastmcp_ctx:
|
||||||
# Create the middleware context.
|
# Create the middleware context.
|
||||||
mw_context = MiddlewareContext(
|
mw_context = MiddlewareContext(
|
||||||
|
|
|
||||||
|
|
@ -30,6 +30,9 @@ import anyio
|
||||||
import httpx
|
import httpx
|
||||||
import mcp.types
|
import mcp.types
|
||||||
import uvicorn
|
import uvicorn
|
||||||
|
from key_value.aio.adapters.pydantic import PydanticAdapter
|
||||||
|
from key_value.aio.protocols import AsyncKeyValue
|
||||||
|
from key_value.aio.stores.memory import MemoryStore
|
||||||
from mcp.server.lowlevel.server import LifespanResultT, NotificationOptions
|
from mcp.server.lowlevel.server import LifespanResultT, NotificationOptions
|
||||||
from mcp.server.stdio import stdio_server
|
from mcp.server.stdio import stdio_server
|
||||||
from mcp.shared.exceptions import McpError
|
from mcp.shared.exceptions import McpError
|
||||||
|
|
@ -99,7 +102,7 @@ from fastmcp.tools.tool_transform import ToolTransformConfig
|
||||||
from fastmcp.utilities.cli import log_server_banner
|
from fastmcp.utilities.cli import log_server_banner
|
||||||
from fastmcp.utilities.components import FastMCPComponent
|
from fastmcp.utilities.components import FastMCPComponent
|
||||||
from fastmcp.utilities.logging import get_logger, temporary_log_level
|
from fastmcp.utilities.logging import get_logger, temporary_log_level
|
||||||
from fastmcp.utilities.types import NotSet, NotSetT
|
from fastmcp.utilities.types import FastMCPBaseModel, NotSet, NotSetT
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from docket import Docket
|
from docket import Docket
|
||||||
|
|
@ -220,6 +223,12 @@ def _lifespan_proxy(
|
||||||
return wrap
|
return wrap
|
||||||
|
|
||||||
|
|
||||||
|
class StateValue(FastMCPBaseModel):
|
||||||
|
"""Wrapper for stored context state values."""
|
||||||
|
|
||||||
|
value: Any
|
||||||
|
|
||||||
|
|
||||||
class FastMCP(Generic[LifespanResultT]):
|
class FastMCP(Generic[LifespanResultT]):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
|
|
@ -242,6 +251,7 @@ class FastMCP(Generic[LifespanResultT]):
|
||||||
on_duplicate: DuplicateBehavior | None = None,
|
on_duplicate: DuplicateBehavior | None = None,
|
||||||
strict_input_validation: bool | None = None,
|
strict_input_validation: bool | None = None,
|
||||||
tasks: bool | None = None,
|
tasks: bool | None = None,
|
||||||
|
session_state_store: AsyncKeyValue | None = None,
|
||||||
# ---
|
# ---
|
||||||
# --- DEPRECATED parameters ---
|
# --- DEPRECATED parameters ---
|
||||||
# ---
|
# ---
|
||||||
|
|
@ -278,6 +288,14 @@ class FastMCP(Generic[LifespanResultT]):
|
||||||
|
|
||||||
self._additional_http_routes: list[BaseRoute] = []
|
self._additional_http_routes: list[BaseRoute] = []
|
||||||
|
|
||||||
|
# Session-scoped state store (shared across all requests)
|
||||||
|
self._state_storage: AsyncKeyValue = session_state_store or MemoryStore()
|
||||||
|
self._state_store: PydanticAdapter[StateValue] = PydanticAdapter[StateValue](
|
||||||
|
key_value=self._state_storage,
|
||||||
|
pydantic_model=StateValue,
|
||||||
|
default_collection="fastmcp_state",
|
||||||
|
)
|
||||||
|
|
||||||
# Create LocalProvider for local components
|
# Create LocalProvider for local components
|
||||||
self._local_provider: LocalProvider = LocalProvider(
|
self._local_provider: LocalProvider = LocalProvider(
|
||||||
on_duplicate=self._on_duplicate
|
on_duplicate=self._on_duplicate
|
||||||
|
|
|
||||||
|
|
@ -12,7 +12,12 @@ from fastmcp.server.middleware import CallNext, Middleware, MiddlewareContext
|
||||||
|
|
||||||
|
|
||||||
class InitializationMiddleware(Middleware):
|
class InitializationMiddleware(Middleware):
|
||||||
"""Middleware that captures initialization details."""
|
"""Middleware that captures initialization details.
|
||||||
|
|
||||||
|
Note: Session state is NOT available during on_initialize because
|
||||||
|
the MCP session has not been established yet. Use instance variables
|
||||||
|
to store data that needs to persist across the session.
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
@ -25,7 +30,7 @@ class InitializationMiddleware(Middleware):
|
||||||
context: MiddlewareContext[mt.InitializeRequest],
|
context: MiddlewareContext[mt.InitializeRequest],
|
||||||
call_next: CallNext[mt.InitializeRequest, None],
|
call_next: CallNext[mt.InitializeRequest, None],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Capture initialization details and store session data."""
|
"""Capture initialization details."""
|
||||||
self.initialized = True
|
self.initialized = True
|
||||||
|
|
||||||
# Extract client info from the initialize params
|
# Extract client info from the initialize params
|
||||||
|
|
@ -34,13 +39,13 @@ class InitializationMiddleware(Middleware):
|
||||||
):
|
):
|
||||||
self.client_info = context.message.params.clientInfo
|
self.client_info = context.message.params.clientInfo
|
||||||
|
|
||||||
# Store data in the context state for cross-request access
|
# Store in instance for cross-request access
|
||||||
if context.fastmcp_context:
|
# (session state is not available during on_initialize)
|
||||||
context.fastmcp_context.set_state("client_initialized", True)
|
self.session_data["client_initialized"] = True
|
||||||
if self.client_info:
|
if self.client_info:
|
||||||
context.fastmcp_context.set_state(
|
self.session_data["client_name"] = getattr(
|
||||||
"client_name", getattr(self.client_info, "name", "unknown")
|
self.client_info, "name", "unknown"
|
||||||
)
|
)
|
||||||
|
|
||||||
return await call_next(context)
|
return await call_next(context)
|
||||||
|
|
||||||
|
|
@ -194,41 +199,36 @@ async def test_multiple_middleware_initialization():
|
||||||
assert detect_mw.tools_modified is True
|
assert detect_mw.tools_modified is True
|
||||||
|
|
||||||
|
|
||||||
async def test_initialization_middleware_with_state_sharing():
|
async def test_session_state_persists_across_tool_calls():
|
||||||
"""Test that state set during initialization is available in later requests."""
|
"""Test that session-scoped state persists across multiple tool calls.
|
||||||
|
|
||||||
|
Session state is only available after the session is established,
|
||||||
|
so it can't be set during on_initialize. This test shows state set
|
||||||
|
during one tool call is accessible in subsequent tool calls.
|
||||||
|
"""
|
||||||
server = FastMCP("TestServer")
|
server = FastMCP("TestServer")
|
||||||
|
|
||||||
class StateTrackingMiddleware(Middleware):
|
class StateTrackingMiddleware(Middleware):
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.init_state = {}
|
self.call_count = 0
|
||||||
self.tool_state = {}
|
self.state_values = []
|
||||||
|
|
||||||
async def on_initialize(
|
|
||||||
self,
|
|
||||||
context: MiddlewareContext[mt.InitializeRequest],
|
|
||||||
call_next: CallNext[mt.InitializeRequest, None],
|
|
||||||
) -> None:
|
|
||||||
# Store some state during initialization
|
|
||||||
if context.fastmcp_context:
|
|
||||||
context.fastmcp_context.set_state("init_timestamp", "2024-01-01")
|
|
||||||
context.fastmcp_context.set_state("client_id", "test-123")
|
|
||||||
self.init_state["timestamp"] = "2024-01-01"
|
|
||||||
self.init_state["client_id"] = "test-123"
|
|
||||||
|
|
||||||
return await call_next(context)
|
|
||||||
|
|
||||||
async def on_call_tool(
|
async def on_call_tool(
|
||||||
self,
|
self,
|
||||||
context: MiddlewareContext[mt.CallToolRequestParams],
|
context: MiddlewareContext[mt.CallToolRequestParams],
|
||||||
call_next: CallNext[mt.CallToolRequestParams, Any],
|
call_next: CallNext[mt.CallToolRequestParams, Any],
|
||||||
) -> Any:
|
) -> Any:
|
||||||
# Try to access state from initialization
|
self.call_count += 1
|
||||||
|
|
||||||
if context.fastmcp_context:
|
if context.fastmcp_context:
|
||||||
timestamp = context.fastmcp_context.get_state("init_timestamp")
|
# Read existing state
|
||||||
client_id = context.fastmcp_context.get_state("client_id")
|
counter = await context.fastmcp_context.get_state("call_counter")
|
||||||
self.tool_state["timestamp"] = timestamp
|
self.state_values.append(counter)
|
||||||
self.tool_state["client_id"] = client_id
|
|
||||||
|
# Increment and save
|
||||||
|
new_counter = (counter or 0) + 1
|
||||||
|
await context.fastmcp_context.set_state("call_counter", new_counter)
|
||||||
|
|
||||||
return await call_next(context)
|
return await call_next(context)
|
||||||
|
|
||||||
|
|
@ -240,20 +240,23 @@ async def test_initialization_middleware_with_state_sharing():
|
||||||
return "success"
|
return "success"
|
||||||
|
|
||||||
async with Client(server) as client:
|
async with Client(server) as client:
|
||||||
# Initialization should have set state
|
# First call - state should be None initially
|
||||||
assert middleware.init_state["timestamp"] == "2024-01-01"
|
|
||||||
assert middleware.init_state["client_id"] == "test-123"
|
|
||||||
|
|
||||||
# Call a tool - state should be accessible
|
|
||||||
result = await client.call_tool("test_tool", {})
|
result = await client.call_tool("test_tool", {})
|
||||||
assert isinstance(result.content[0], TextContent)
|
assert isinstance(result.content[0], TextContent)
|
||||||
assert result.content[0].text == "success"
|
assert result.content[0].text == "success"
|
||||||
|
|
||||||
# State should have been accessible during tool call
|
# Second call - state should show previous value (1)
|
||||||
# Note: State is request-scoped, so it won't persist across requests
|
result = await client.call_tool("test_tool", {})
|
||||||
# This test shows the pattern, but actual cross-request state would need
|
assert isinstance(result.content[0], TextContent)
|
||||||
# external storage (Redis, DB, etc.)
|
|
||||||
# The middleware.tool_state might be None if state doesn't persist
|
# Third call - state should show previous value (2)
|
||||||
|
result = await client.call_tool("test_tool", {})
|
||||||
|
assert isinstance(result.content[0], TextContent)
|
||||||
|
|
||||||
|
# Verify state persisted across calls within the session
|
||||||
|
assert middleware.call_count == 3
|
||||||
|
# First call saw None, second saw 1, third saw 2
|
||||||
|
assert middleware.state_values == [None, 1, 2]
|
||||||
|
|
||||||
|
|
||||||
async def test_middleware_can_access_initialize_result():
|
async def test_middleware_can_access_initialize_result():
|
||||||
|
|
@ -375,3 +378,55 @@ async def test_middleware_mcp_error_after_call_next():
|
||||||
pass
|
pass
|
||||||
|
|
||||||
assert middleware.error_raised is True
|
assert middleware.error_raised is True
|
||||||
|
|
||||||
|
|
||||||
|
async def test_state_isolation_between_streamable_http_clients():
|
||||||
|
"""Test that different HTTP clients have isolated session state.
|
||||||
|
|
||||||
|
Each client should have its own session ID and isolated state.
|
||||||
|
"""
|
||||||
|
from fastmcp.client.transports import StreamableHttpTransport
|
||||||
|
from fastmcp.server.context import Context
|
||||||
|
from fastmcp.utilities.tests import run_server_async
|
||||||
|
|
||||||
|
server = FastMCP("TestServer")
|
||||||
|
|
||||||
|
@server.tool
|
||||||
|
async def store_and_read(value: str, ctx: Context) -> dict:
|
||||||
|
"""Store a value and return session info."""
|
||||||
|
existing = await ctx.get_state("client_value")
|
||||||
|
await ctx.set_state("client_value", value)
|
||||||
|
return {
|
||||||
|
"existing": existing,
|
||||||
|
"stored": value,
|
||||||
|
"session_id": ctx.session_id,
|
||||||
|
}
|
||||||
|
|
||||||
|
async with run_server_async(server, transport="streamable-http") as url:
|
||||||
|
import json
|
||||||
|
|
||||||
|
# Client 1 stores its value
|
||||||
|
transport1 = StreamableHttpTransport(url=url)
|
||||||
|
async with Client(transport=transport1) as client1:
|
||||||
|
result1 = await client1.call_tool(
|
||||||
|
"store_and_read", {"value": "client1-value"}
|
||||||
|
)
|
||||||
|
data1 = json.loads(result1.content[0].text)
|
||||||
|
assert data1["existing"] is None
|
||||||
|
assert data1["stored"] == "client1-value"
|
||||||
|
session_id_1 = data1["session_id"]
|
||||||
|
|
||||||
|
# Client 2 should have completely isolated state
|
||||||
|
transport2 = StreamableHttpTransport(url=url)
|
||||||
|
async with Client(transport=transport2) as client2:
|
||||||
|
result2 = await client2.call_tool(
|
||||||
|
"store_and_read", {"value": "client2-value"}
|
||||||
|
)
|
||||||
|
data2 = json.loads(result2.content[0].text)
|
||||||
|
# Should NOT see client1's value
|
||||||
|
assert data2["existing"] is None
|
||||||
|
assert data2["stored"] == "client2-value"
|
||||||
|
session_id_2 = data2["session_id"]
|
||||||
|
|
||||||
|
# Session IDs should be different
|
||||||
|
assert session_id_1 != session_id_2
|
||||||
|
|
|
||||||
|
|
@ -58,75 +58,188 @@ class TestSessionId:
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
assert context.session_id == "test-session-123"
|
try:
|
||||||
|
assert context.session_id == "test-session-123"
|
||||||
request_ctx.reset(token)
|
finally:
|
||||||
|
request_ctx.reset(token)
|
||||||
|
|
||||||
def test_session_id_without_http_headers(self, context):
|
def test_session_id_without_http_headers(self, context):
|
||||||
"""Test that session_id returns a UUID string when no HTTP headers are available."""
|
"""Test that session_id returns a UUID when no HTTP headers are available.
|
||||||
|
|
||||||
|
For STDIO/SSE/in-memory transports, we generate a UUID and cache it
|
||||||
|
on the session for consistency with state operations.
|
||||||
|
"""
|
||||||
import uuid
|
import uuid
|
||||||
|
|
||||||
from mcp.server.lowlevel.server import request_ctx
|
from mcp.server.lowlevel.server import request_ctx
|
||||||
from mcp.shared.context import RequestContext
|
from mcp.shared.context import RequestContext
|
||||||
|
|
||||||
|
mock_session = MagicMock(wraps={})
|
||||||
token = request_ctx.set(
|
token = request_ctx.set(
|
||||||
RequestContext(
|
RequestContext(
|
||||||
request_id=0,
|
request_id=0,
|
||||||
meta=None,
|
meta=None,
|
||||||
session=MagicMock(wraps={}),
|
session=mock_session,
|
||||||
lifespan_context=MagicMock(),
|
lifespan_context=MagicMock(),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
assert uuid.UUID(context.session_id)
|
try:
|
||||||
|
# session_id should be a valid UUID for non-HTTP transports
|
||||||
request_ctx.reset(token)
|
session_id = context.session_id
|
||||||
|
assert uuid.UUID(session_id) # Valid UUID format
|
||||||
|
# Should be cached on session
|
||||||
|
assert mock_session._fastmcp_state_prefix == session_id
|
||||||
|
finally:
|
||||||
|
request_ctx.reset(token)
|
||||||
|
|
||||||
|
|
||||||
class TestContextState:
|
class TestContextState:
|
||||||
"""Test suite for Context state functionality."""
|
"""Test suite for Context state functionality."""
|
||||||
|
|
||||||
async def test_context_state(self):
|
async def test_context_state_basic(self):
|
||||||
"""Test that state modifications in child contexts don't affect parent."""
|
"""Test basic get/set/delete state operations."""
|
||||||
mock_fastmcp = MagicMock()
|
server = FastMCP("test")
|
||||||
|
mock_session = MagicMock() # Use same session for consistent id()
|
||||||
|
|
||||||
async with Context(fastmcp=mock_fastmcp) as context:
|
async with Context(fastmcp=server, session=mock_session) as context:
|
||||||
assert context.get_state("test1") is None
|
# Initially empty
|
||||||
assert context.get_state("test2") is None
|
assert await context.get_state("test1") is None
|
||||||
context.set_state("test1", "value")
|
assert await context.get_state("test2") is None
|
||||||
context.set_state("test2", 2)
|
|
||||||
assert context.get_state("test1") == "value"
|
|
||||||
assert context.get_state("test2") == 2
|
|
||||||
context.set_state("test1", "new_value")
|
|
||||||
assert context.get_state("test1") == "new_value"
|
|
||||||
|
|
||||||
async def test_context_state_inheritance(self):
|
# Set values
|
||||||
"""Test that child contexts inherit parent state."""
|
await context.set_state("test1", "value")
|
||||||
mock_fastmcp = MagicMock()
|
await context.set_state("test2", 2)
|
||||||
|
|
||||||
async with Context(fastmcp=mock_fastmcp) as context1:
|
# Retrieve values
|
||||||
context1.set_state("key1", "key1-context1")
|
assert await context.get_state("test1") == "value"
|
||||||
context1.set_state("key2", "key2-context1")
|
assert await context.get_state("test2") == 2
|
||||||
async with Context(fastmcp=mock_fastmcp) as context2:
|
|
||||||
# Override one key
|
|
||||||
context2.set_state("key1", "key1-context2")
|
|
||||||
assert context2.get_state("key1") == "key1-context2"
|
|
||||||
assert context1.get_state("key1") == "key1-context1"
|
|
||||||
assert context2.get_state("key2") == "key2-context1"
|
|
||||||
|
|
||||||
async with Context(fastmcp=mock_fastmcp) as context3:
|
# Update value
|
||||||
# Verify state was inherited
|
await context.set_state("test1", "new_value")
|
||||||
assert context3.get_state("key1") == "key1-context2"
|
assert await context.get_state("test1") == "new_value"
|
||||||
assert context3.get_state("key2") == "key2-context1"
|
|
||||||
|
|
||||||
# Add a new key and verify parents were not affected
|
# Delete value
|
||||||
context3.set_state("key-context3-only", 1)
|
await context.delete_state("test1")
|
||||||
assert context1.get_state("key-context3-only") is None
|
assert await context.get_state("test1") is None
|
||||||
assert context2.get_state("key-context3-only") is None
|
|
||||||
assert context3.get_state("key-context3-only") == 1
|
|
||||||
|
|
||||||
assert context1.get_state("key1") == "key1-context1"
|
async def test_context_state_session_isolation(self):
|
||||||
assert context1.get_state("key-context3-only") is None
|
"""Test that different sessions have isolated state."""
|
||||||
|
server = FastMCP("test")
|
||||||
|
session_a = MagicMock()
|
||||||
|
session_b = MagicMock()
|
||||||
|
|
||||||
|
async with Context(fastmcp=server, session=session_a) as context1:
|
||||||
|
await context1.set_state("key", "value-from-A")
|
||||||
|
|
||||||
|
async with Context(fastmcp=server, session=session_b) as context2:
|
||||||
|
# Session B should not see session A's state
|
||||||
|
assert await context2.get_state("key") is None
|
||||||
|
await context2.set_state("key", "value-from-B")
|
||||||
|
assert await context2.get_state("key") == "value-from-B"
|
||||||
|
|
||||||
|
# Verify session A's state is still intact
|
||||||
|
async with Context(fastmcp=server, session=session_a) as context3:
|
||||||
|
assert await context3.get_state("key") == "value-from-A"
|
||||||
|
|
||||||
|
async def test_context_state_persists_across_requests(self):
|
||||||
|
"""Test that state persists across multiple context instances (requests)."""
|
||||||
|
server = FastMCP("test")
|
||||||
|
mock_session = MagicMock() # Same session = same id()
|
||||||
|
|
||||||
|
# First request sets state
|
||||||
|
async with Context(fastmcp=server, session=mock_session) as context1:
|
||||||
|
await context1.set_state("counter", 1)
|
||||||
|
|
||||||
|
# Second request in same session sees the state
|
||||||
|
async with Context(fastmcp=server, session=mock_session) as context2:
|
||||||
|
counter = await context2.get_state("counter")
|
||||||
|
assert counter == 1
|
||||||
|
await context2.set_state("counter", counter + 1)
|
||||||
|
|
||||||
|
# Third request sees updated state
|
||||||
|
async with Context(fastmcp=server, session=mock_session) as context3:
|
||||||
|
assert await context3.get_state("counter") == 2
|
||||||
|
|
||||||
|
async def test_context_state_nested_contexts_share_state(self):
|
||||||
|
"""Test that nested contexts within the same session share state."""
|
||||||
|
server = FastMCP("test")
|
||||||
|
mock_session = MagicMock()
|
||||||
|
|
||||||
|
async with Context(fastmcp=server, session=mock_session) as context1:
|
||||||
|
await context1.set_state("key", "outer-value")
|
||||||
|
|
||||||
|
async with Context(fastmcp=server, session=mock_session) as context2:
|
||||||
|
# Nested context sees same state (same session)
|
||||||
|
assert await context2.get_state("key") == "outer-value"
|
||||||
|
|
||||||
|
# Nested context can modify shared state
|
||||||
|
await context2.set_state("key", "inner-value")
|
||||||
|
|
||||||
|
# Outer context sees the modification
|
||||||
|
assert await context1.get_state("key") == "inner-value"
|
||||||
|
|
||||||
|
async def test_two_clients_same_key_isolated_by_session(self):
|
||||||
|
"""Test that two different clients can store the same key independently.
|
||||||
|
|
||||||
|
Each client gets an auto-generated session ID, and their state is isolated.
|
||||||
|
"""
|
||||||
|
import json
|
||||||
|
|
||||||
|
from fastmcp import Client
|
||||||
|
|
||||||
|
server = FastMCP("test")
|
||||||
|
stored_session_ids: list[str] = []
|
||||||
|
|
||||||
|
@server.tool
|
||||||
|
async def store_and_read(value: str, ctx: Context) -> dict:
|
||||||
|
"""Store a value and return all state info."""
|
||||||
|
stored_session_ids.append(ctx.session_id)
|
||||||
|
existing = await ctx.get_state("shared_key")
|
||||||
|
await ctx.set_state("shared_key", value)
|
||||||
|
new_value = await ctx.get_state("shared_key")
|
||||||
|
return {
|
||||||
|
"session_id": ctx.session_id,
|
||||||
|
"existing_value": existing,
|
||||||
|
"new_value": new_value,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Client 1 stores "value-from-client-1"
|
||||||
|
async with Client(server) as client1:
|
||||||
|
result1 = await client1.call_tool(
|
||||||
|
"store_and_read", {"value": "value-from-client-1"}
|
||||||
|
)
|
||||||
|
data1 = json.loads(result1.content[0].text)
|
||||||
|
assert data1["existing_value"] is None # First write
|
||||||
|
assert data1["new_value"] == "value-from-client-1"
|
||||||
|
session_id_1 = data1["session_id"]
|
||||||
|
|
||||||
|
# Client 2 stores "value-from-client-2" with the SAME key
|
||||||
|
async with Client(server) as client2:
|
||||||
|
result2 = await client2.call_tool(
|
||||||
|
"store_and_read", {"value": "value-from-client-2"}
|
||||||
|
)
|
||||||
|
data2 = json.loads(result2.content[0].text)
|
||||||
|
# Client 2 should NOT see client 1's value (different session)
|
||||||
|
assert data2["existing_value"] is None
|
||||||
|
assert data2["new_value"] == "value-from-client-2"
|
||||||
|
session_id_2 = data2["session_id"]
|
||||||
|
|
||||||
|
# Verify session IDs were auto-generated and are different
|
||||||
|
assert session_id_1 is not None
|
||||||
|
assert session_id_2 is not None
|
||||||
|
assert session_id_1 != session_id_2
|
||||||
|
|
||||||
|
# Client 1 reconnects and should still see their value
|
||||||
|
async with Client(server) as client1_again:
|
||||||
|
# But this is a NEW session (new connection = new session ID)
|
||||||
|
result3 = await client1_again.call_tool(
|
||||||
|
"store_and_read", {"value": "value-from-client-1-again"}
|
||||||
|
)
|
||||||
|
data3 = json.loads(result3.content[0].text)
|
||||||
|
# New session, so existing value is None
|
||||||
|
assert data3["existing_value"] is None
|
||||||
|
assert data3["session_id"] != session_id_1 # Different session
|
||||||
|
|
||||||
|
|
||||||
class TestContextMeta:
|
class TestContextMeta:
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue