diff --git a/docs/servers/dependency-injection.mdx b/docs/servers/dependency-injection.mdx index d13f43952..40fc7b65b 100644 --- a/docs/servers/dependency-injection.mdx +++ b/docs/servers/dependency-injection.mdx @@ -160,14 +160,20 @@ def get_client_ip() -> str: ``` -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. ### HTTP Headers -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()`: diff --git a/src/fastmcp/server/dependencies.py b/src/fastmcp/server/dependencies.py index 4d2c2b1c4..a01e4e06c 100644 --- a/src/fastmcp/server/dependencies.py +++ b/src/fastmcp/server/dependencies.py @@ -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]: diff --git a/src/fastmcp/server/tasks/handlers.py b/src/fastmcp/server/tasks/handlers.py index 034f133dd..051785ddd 100644 --- a/src/fastmcp/server/tasks/handlers.py +++ b/src/fastmcp/server/tasks/handlers.py @@ -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 diff --git a/tests/server/http/test_http_dependencies.py b/tests/server/http/test_http_dependencies.py index e637af269..f1a1e6f47 100644 --- a/tests/server/http/test_http_dependencies.py +++ b/tests/server/http/test_http_dependencies.py @@ -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", + }