mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-24 06:24:18 +02:00
fix: HTTP request headers not accessible in background task workers (#3631)
Co-authored-by: Jeremiah Lowin <153965+jlowin@users.noreply.github.com>
This commit is contained in:
parent
5683c0f495
commit
5879119de5
4 changed files with 156 additions and 6 deletions
|
|
@ -160,14 +160,20 @@ def get_client_ip() -> str:
|
|||
```
|
||||
|
||||
<Note>
|
||||
Both raise `RuntimeError` when called outside an HTTP context (e.g., STDIO transport). Use HTTP Headers if you need graceful fallback.
|
||||
Both raise `RuntimeError` when called outside an HTTP context (e.g., STDIO transport).
|
||||
For background tasks created from an HTTP request, FastMCP restores a minimal request
|
||||
backed by the originating request's snapshotted headers. Use HTTP Headers if you need
|
||||
graceful fallback.
|
||||
</Note>
|
||||
|
||||
### HTTP Headers
|
||||
|
||||
<VersionBadge version="2.2.11" />
|
||||
|
||||
Access HTTP headers with graceful fallback—returns an empty dictionary when no HTTP request is available, making it safe for code that might run over any transport.
|
||||
Access HTTP headers with graceful fallback. When a background task originates from an
|
||||
HTTP request, FastMCP restores the originating headers inside the worker. When no HTTP
|
||||
request is available, this returns an empty dictionary, making it safe for code that
|
||||
might run over any transport.
|
||||
|
||||
**Dependency injection:** Use `CurrentHeaders()`:
|
||||
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from __future__ import annotations
|
|||
|
||||
import contextlib
|
||||
import inspect
|
||||
import json
|
||||
import logging
|
||||
import weakref
|
||||
from collections import OrderedDict
|
||||
|
|
@ -207,6 +208,9 @@ _current_worker: ContextVar[Worker | None] = ContextVar("worker", default=None)
|
|||
_task_access_token: ContextVar[AccessToken | None] = ContextVar(
|
||||
"task_access_token", default=None
|
||||
)
|
||||
_task_http_headers: ContextVar[dict[str, str] | None] = ContextVar(
|
||||
"task_http_headers", default=None
|
||||
)
|
||||
|
||||
|
||||
# --- Docket availability check ---
|
||||
|
|
@ -441,6 +445,8 @@ def get_http_request() -> Request:
|
|||
"""Get the current HTTP request.
|
||||
|
||||
Tries MCP SDK's request_ctx first, then falls back to FastMCP's HTTP context.
|
||||
In background tasks, returns a synthetic request populated with the
|
||||
snapshotted headers from the originating HTTP request.
|
||||
"""
|
||||
# Try MCP SDK's request_ctx first (set during normal MCP request handling)
|
||||
request = None
|
||||
|
|
@ -452,6 +458,29 @@ def get_http_request() -> Request:
|
|||
if request is None:
|
||||
request = _current_http_request.get()
|
||||
|
||||
# In Docket workers, restore a minimal request from the snapshotted headers.
|
||||
if request is None:
|
||||
task_headers = _task_http_headers.get()
|
||||
if task_headers:
|
||||
request = Request(
|
||||
{
|
||||
"type": "http",
|
||||
"http_version": "1.1",
|
||||
"method": "POST",
|
||||
"scheme": "http",
|
||||
"path": "/",
|
||||
"raw_path": b"/",
|
||||
"query_string": b"",
|
||||
"headers": [
|
||||
(name.encode("latin-1"), value.encode("latin-1"))
|
||||
for name, value in task_headers.items()
|
||||
],
|
||||
"client": None,
|
||||
"server": None,
|
||||
"root_path": "",
|
||||
}
|
||||
)
|
||||
|
||||
if request is None:
|
||||
raise RuntimeError("No active HTTP request found.")
|
||||
return request
|
||||
|
|
@ -815,6 +844,38 @@ async def _restore_task_access_token(
|
|||
return None
|
||||
|
||||
|
||||
async def _restore_task_http_headers(
|
||||
session_id: str, task_id: str
|
||||
) -> Token[dict[str, str] | None] | None:
|
||||
"""Restore the HTTP header snapshot from Redis into a ContextVar."""
|
||||
docket = _current_docket.get()
|
||||
if docket is None:
|
||||
return None
|
||||
|
||||
headers_key = docket.key(f"fastmcp:task:{session_id}:{task_id}:http_headers")
|
||||
try:
|
||||
async with docket.redis() as redis:
|
||||
headers_data = await redis.get(headers_key)
|
||||
if headers_data is None:
|
||||
return None
|
||||
if isinstance(headers_data, bytes):
|
||||
headers_data = headers_data.decode()
|
||||
restored = json.loads(str(headers_data))
|
||||
if not isinstance(restored, dict):
|
||||
return None
|
||||
return _task_http_headers.set(
|
||||
{str(name).lower(): str(value) for name, value in restored.items()}
|
||||
)
|
||||
except Exception:
|
||||
_logger.warning(
|
||||
"Failed to restore HTTP headers for task %s:%s",
|
||||
session_id,
|
||||
task_id,
|
||||
exc_info=True,
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
async def _restore_task_origin_request_id(session_id: str, task_id: str) -> str | None:
|
||||
"""Restore the origin request ID snapshot for a background task.
|
||||
|
||||
|
|
@ -1120,8 +1181,20 @@ def CurrentFastMCP() -> FastMCP:
|
|||
class _CurrentRequest(Dependency[Request]):
|
||||
"""Async context manager for HTTP Request dependency."""
|
||||
|
||||
_task_http_headers_cv_token: Token[dict[str, str] | None] | None = None
|
||||
|
||||
async def __aenter__(self) -> Request:
|
||||
return get_http_request()
|
||||
try:
|
||||
return get_http_request()
|
||||
except RuntimeError:
|
||||
task_info = get_task_context()
|
||||
if task_info is None:
|
||||
raise
|
||||
if _task_http_headers.get() is None:
|
||||
self._task_http_headers_cv_token = await _restore_task_http_headers(
|
||||
task_info.session_id, task_info.task_id
|
||||
)
|
||||
return get_http_request()
|
||||
|
||||
async def __aexit__(
|
||||
self,
|
||||
|
|
@ -1129,7 +1202,9 @@ class _CurrentRequest(Dependency[Request]):
|
|||
exc_value: BaseException | None,
|
||||
traceback: TracebackType | None,
|
||||
) -> None:
|
||||
pass
|
||||
if self._task_http_headers_cv_token is not None:
|
||||
_task_http_headers.reset(self._task_http_headers_cv_token)
|
||||
self._task_http_headers_cv_token = None
|
||||
|
||||
|
||||
def CurrentRequest() -> Request:
|
||||
|
|
@ -1161,7 +1236,15 @@ def CurrentRequest() -> Request:
|
|||
class _CurrentHeaders(Dependency[dict[str, str]]):
|
||||
"""Async context manager for HTTP Headers dependency."""
|
||||
|
||||
_task_http_headers_cv_token: Token[dict[str, str] | None] | None = None
|
||||
|
||||
async def __aenter__(self) -> dict[str, str]:
|
||||
if _task_http_headers.get() is None:
|
||||
task_info = get_task_context()
|
||||
if task_info is not None:
|
||||
self._task_http_headers_cv_token = await _restore_task_http_headers(
|
||||
task_info.session_id, task_info.task_id
|
||||
)
|
||||
return get_http_headers(include={"authorization"})
|
||||
|
||||
async def __aexit__(
|
||||
|
|
@ -1170,7 +1253,9 @@ class _CurrentHeaders(Dependency[dict[str, str]]):
|
|||
exc_value: BaseException | None,
|
||||
traceback: TracebackType | None,
|
||||
) -> None:
|
||||
pass
|
||||
if self._task_http_headers_cv_token is not None:
|
||||
_task_http_headers.reset(self._task_http_headers_cv_token)
|
||||
self._task_http_headers_cv_token = None
|
||||
|
||||
|
||||
def CurrentHeaders() -> dict[str, str]:
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ Handles queuing tool/prompt/resource executions to Docket as background tasks.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from contextlib import suppress
|
||||
from datetime import datetime, timezone
|
||||
|
|
@ -18,6 +19,7 @@ from fastmcp.server.dependencies import (
|
|||
_current_docket,
|
||||
get_access_token,
|
||||
get_context,
|
||||
get_http_headers,
|
||||
register_task_server,
|
||||
)
|
||||
from fastmcp.server.tasks.config import TaskMeta
|
||||
|
|
@ -122,6 +124,10 @@ async def submit_to_docket(
|
|||
access_token_key = docket.key(
|
||||
f"fastmcp:task:{session_id}:{server_task_id}:access_token"
|
||||
)
|
||||
http_headers = get_http_headers(include_all=True)
|
||||
http_headers_key = docket.key(
|
||||
f"fastmcp:task:{session_id}:{server_task_id}:http_headers"
|
||||
)
|
||||
|
||||
async with docket.redis() as redis:
|
||||
await redis.set(task_meta_key, task_key, ex=ttl_seconds)
|
||||
|
|
@ -133,6 +139,8 @@ async def submit_to_docket(
|
|||
await redis.set(
|
||||
access_token_key, access_token.model_dump_json(), ex=ttl_seconds
|
||||
)
|
||||
if http_headers:
|
||||
await redis.set(http_headers_key, json.dumps(http_headers), ex=ttl_seconds)
|
||||
|
||||
# Register session for Context access in background workers (SEP-1686)
|
||||
# This enables elicitation/sampling from background tasks via weakref
|
||||
|
|
|
|||
|
|
@ -2,10 +2,11 @@ import json
|
|||
|
||||
import pytest
|
||||
from mcp.types import TextContent, TextResourceContents
|
||||
from starlette.requests import Request
|
||||
|
||||
from fastmcp.client import Client
|
||||
from fastmcp.client.transports import SSETransport, StreamableHttpTransport
|
||||
from fastmcp.server.dependencies import get_http_request
|
||||
from fastmcp.server.dependencies import CurrentHeaders, CurrentRequest, get_http_request
|
||||
from fastmcp.server.server import FastMCP
|
||||
from fastmcp.utilities.tests import run_server_async
|
||||
|
||||
|
|
@ -166,3 +167,53 @@ async def test_get_http_headers_excludes_content_type(sse_server: str):
|
|||
# Custom headers should be included
|
||||
assert "x-custom-header" in headers
|
||||
assert headers["x-custom-header"] == "should-be-included"
|
||||
|
||||
|
||||
async def test_background_task_can_read_snapshotted_request_headers():
|
||||
"""Background tools can still access request headers via get_http_request()."""
|
||||
server = FastMCP()
|
||||
|
||||
@server.tool(task=True)
|
||||
async def check_request_header() -> str:
|
||||
request = get_http_request()
|
||||
return request.headers.get("x-tenant-id", "missing")
|
||||
|
||||
async with run_server_async(server, transport="sse") as url:
|
||||
async with Client(
|
||||
transport=SSETransport(url, headers={"X-Tenant-ID": "tenant-123"})
|
||||
) as client:
|
||||
task = await client.call_tool("check_request_header", task=True)
|
||||
result = await task.result()
|
||||
assert result.data == "tenant-123"
|
||||
|
||||
|
||||
async def test_background_task_current_http_dependencies_restore_headers():
|
||||
"""CurrentHeaders/CurrentRequest work in task workers without explicit Context."""
|
||||
server = FastMCP()
|
||||
|
||||
@server.tool(task=True)
|
||||
async def check_headers(
|
||||
headers: dict[str, str] = CurrentHeaders(),
|
||||
request: Request = CurrentRequest(),
|
||||
) -> dict[str, str]:
|
||||
return {
|
||||
"authorization": headers.get("authorization", "missing"),
|
||||
"tenant": request.headers.get("x-tenant-id", "missing"),
|
||||
}
|
||||
|
||||
async with run_server_async(server, transport="sse") as url:
|
||||
async with Client(
|
||||
transport=SSETransport(
|
||||
url,
|
||||
headers={
|
||||
"Authorization": "Bearer tenant-token",
|
||||
"X-Tenant-ID": "tenant-456",
|
||||
},
|
||||
)
|
||||
) as client:
|
||||
task = await client.call_tool("check_headers", task=True)
|
||||
result = await task.result()
|
||||
assert result.data == {
|
||||
"authorization": "Bearer tenant-token",
|
||||
"tenant": "tenant-456",
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue