Skip to content

Commit d4bd472

Browse files
committed
feat(session-manager): add tenant validation on session access
Prevent cross-tenant session hijacking by validating that the authenticated tenant matches the session's bound tenant on every request. Sessions created without a tenant (no auth) remain accessible to all requests for backward compatibility. Adds a parallel _session_tenants dict that records the tenant_id from tenant_id_var at session creation time. On existing session lookup, a mismatch returns 404 (same as "session not found" to avoid information leakage). Tenant mappings are cleaned up alongside sessions on all exit paths: idle timeout, crash, and shutdown. Includes 7 tests covering bidirectional isolation, same-tenant reuse, backward compatibility, unauthenticated access rejection, and cleanup. 100% branch coverage on streamable_http_manager.py.
1 parent e6d7fc7 commit d4bd472

2 files changed

Lines changed: 216 additions & 0 deletions

File tree

src/mcp/server/streamable_http_manager.py

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121
StreamableHTTPServerTransport,
2222
)
2323
from mcp.server.transport_security import TransportSecuritySettings
24+
from mcp.shared._context import tenant_id_var
2425
from mcp.types import INVALID_REQUEST, ErrorData, JSONRPCError
2526

2627
if TYPE_CHECKING:
@@ -89,6 +90,7 @@ def __init__(
8990
# Session tracking (only used if not stateless)
9091
self._session_creation_lock = anyio.Lock()
9192
self._server_instances: dict[str, StreamableHTTPServerTransport] = {}
93+
self._session_tenants: dict[str, str | None] = {}
9294

9395
# The task group will be set during lifespan
9496
self._task_group = None
@@ -135,6 +137,7 @@ async def lifespan(app: Starlette) -> AsyncIterator[None]:
135137
self._task_group = None
136138
# Clear any remaining server instances
137139
self._server_instances.clear()
140+
self._session_tenants.clear()
138141

139142
async def handle_request(self, scope: Scope, receive: Receive, send: Send) -> None:
140143
"""Process ASGI request with proper session handling and transport setup.
@@ -194,6 +197,27 @@ async def _handle_stateful_request(self, scope: Scope, receive: Receive, send: S
194197

195198
# Existing session case
196199
if request_mcp_session_id is not None and request_mcp_session_id in self._server_instances:
200+
# Validate that the requesting tenant matches the session's tenant
201+
session_tenant = self._session_tenants.get(request_mcp_session_id)
202+
request_tenant = tenant_id_var.get()
203+
if session_tenant is not None and request_tenant != session_tenant:
204+
logger.warning(
205+
f"Tenant mismatch for session {request_mcp_session_id[:64]}: "
206+
f"session bound to '{session_tenant}', request from '{request_tenant}'"
207+
)
208+
error_response = JSONRPCError(
209+
jsonrpc="2.0",
210+
id=None,
211+
error=ErrorData(code=INVALID_REQUEST, message="Session not found"),
212+
)
213+
response = Response(
214+
content=error_response.model_dump_json(by_alias=True, exclude_unset=True),
215+
status_code=HTTPStatus.NOT_FOUND,
216+
media_type="application/json",
217+
)
218+
await response(scope, receive, send)
219+
return
220+
197221
transport = self._server_instances[request_mcp_session_id]
198222
logger.debug("Session already exists, handling request directly")
199223
# Push back idle deadline on activity
@@ -217,6 +241,7 @@ async def _handle_stateful_request(self, scope: Scope, receive: Receive, send: S
217241

218242
assert http_transport.mcp_session_id is not None
219243
self._server_instances[http_transport.mcp_session_id] = http_transport
244+
self._session_tenants[http_transport.mcp_session_id] = tenant_id_var.get()
220245
logger.info(f"Created new transport with session ID: {new_session_id}")
221246

222247
# Define the server runner
@@ -246,6 +271,7 @@ async def run_server(*, task_status: TaskStatus[None] = anyio.TASK_STATUS_IGNORE
246271
assert http_transport.mcp_session_id is not None
247272
logger.info(f"Session {http_transport.mcp_session_id} idle timeout")
248273
self._server_instances.pop(http_transport.mcp_session_id, None)
274+
self._session_tenants.pop(http_transport.mcp_session_id, None)
249275
await http_transport.terminate()
250276
except Exception:
251277
logger.exception(f"Session {http_transport.mcp_session_id} crashed")
@@ -260,6 +286,7 @@ async def run_server(*, task_status: TaskStatus[None] = anyio.TASK_STATUS_IGNORE
260286
f"{http_transport.mcp_session_id} from active instances."
261287
)
262288
del self._server_instances[http_transport.mcp_session_id]
289+
self._session_tenants.pop(http_transport.mcp_session_id, None)
263290

264291
# Assert task group is not None for type checking
265292
assert self._task_group is not None

tests/server/test_streamable_http_manager.py

Lines changed: 189 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -413,3 +413,192 @@ def test_session_idle_timeout_rejects_non_positive():
413413
def test_session_idle_timeout_rejects_stateless():
414414
with pytest.raises(RuntimeError, match="not supported in stateless"):
415415
StreamableHTTPSessionManager(app=Server("test"), session_idle_timeout=30, stateless=True)
416+
417+
418+
# --- Multi-tenancy: session-level tenant isolation ---
419+
420+
421+
def _extract_session_id(messages: list[Message]) -> str | None:
422+
"""Extract the MCP session ID from ASGI response messages."""
423+
for msg in messages:
424+
if msg["type"] == "http.response.start":
425+
for header_name, header_value in msg.get("headers", []):
426+
if header_name.decode().lower() == MCP_SESSION_ID_HEADER.lower():
427+
return header_value.decode()
428+
return None # pragma: no cover
429+
430+
431+
def _extract_status(messages: list[Message]) -> int | None:
432+
"""Extract the HTTP status code from ASGI response messages."""
433+
for msg in messages:
434+
if msg["type"] == "http.response.start":
435+
return msg["status"]
436+
return None # pragma: no cover
437+
438+
439+
def _make_scope(session_id: str | None = None) -> dict[str, Any]:
440+
"""Build a minimal ASGI scope for testing, optionally with a session ID."""
441+
headers: list[tuple[bytes, bytes]] = [(b"content-type", b"application/json")]
442+
if session_id is not None:
443+
headers.append((b"mcp-session-id", session_id.encode()))
444+
return {"type": "http", "method": "POST", "path": "/mcp", "headers": headers}
445+
446+
447+
async def _mock_send(messages: list[Message], message: Message) -> None:
448+
"""Async send that collects messages."""
449+
messages.append(message)
450+
451+
452+
async def _mock_receive() -> dict[str, Any]: # pragma: no cover
453+
return {"type": "http.request", "body": b"", "more_body": False}
454+
455+
456+
def _set_tenant(tenant: str | None) -> Any:
457+
"""Set tenant_id_var if tenant is not None; return the token (or None)."""
458+
from mcp.shared._context import tenant_id_var
459+
460+
return tenant_id_var.set(tenant) if tenant is not None else None
461+
462+
463+
def _reset_tenant(token: Any) -> None:
464+
"""Reset tenant_id_var if a token was set."""
465+
from mcp.shared._context import tenant_id_var
466+
467+
if token is not None:
468+
tenant_id_var.reset(token)
469+
470+
471+
async def _create_session_blocking(
472+
manager: StreamableHTTPSessionManager,
473+
app: Server[Any],
474+
stop_event: anyio.Event,
475+
tenant: str | None = None,
476+
) -> str:
477+
"""Create a session whose server stays alive until stop_event is set."""
478+
479+
async def blocking_run(*args: Any, **kwargs: Any) -> None:
480+
await stop_event.wait()
481+
482+
app.run = AsyncMock(side_effect=blocking_run)
483+
484+
messages: list[Message] = []
485+
token = _set_tenant(tenant)
486+
try:
487+
await manager.handle_request(
488+
_make_scope(), _mock_receive, lambda msg, _msgs=messages: _mock_send(_msgs, msg)
489+
)
490+
finally:
491+
_reset_tenant(token)
492+
493+
session_id = _extract_session_id(messages)
494+
assert session_id is not None
495+
return session_id
496+
497+
498+
async def _access_session(
499+
manager: StreamableHTTPSessionManager,
500+
session_id: str,
501+
tenant: str | None = None,
502+
) -> int | None:
503+
"""Access an existing session and return the HTTP status code."""
504+
messages: list[Message] = []
505+
token = _set_tenant(tenant)
506+
try:
507+
await manager.handle_request(
508+
_make_scope(session_id), _mock_receive, lambda msg, _msgs=messages: _mock_send(_msgs, msg)
509+
)
510+
finally:
511+
_reset_tenant(token)
512+
513+
return _extract_status(messages)
514+
515+
516+
@pytest.mark.anyio
517+
async def test_tenant_mismatch_returns_404(running_manager: tuple[StreamableHTTPSessionManager, Server]):
518+
"""A request from tenant-b cannot access a session created by tenant-a."""
519+
manager, app = running_manager
520+
stop = anyio.Event()
521+
session_id = await _create_session_blocking(manager, app, stop, tenant="tenant-a")
522+
523+
assert await _access_session(manager, session_id, tenant="tenant-b") == 404
524+
stop.set()
525+
526+
527+
@pytest.mark.anyio
528+
async def test_two_tenants_cannot_access_each_others_sessions(
529+
running_manager: tuple[StreamableHTTPSessionManager, Server],
530+
):
531+
"""Two tenants each create a session; neither can access the other's."""
532+
manager, app = running_manager
533+
stop = anyio.Event()
534+
535+
session_a = await _create_session_blocking(manager, app, stop, tenant="tenant-a")
536+
session_b = await _create_session_blocking(manager, app, stop, tenant="tenant-b")
537+
assert session_a != session_b
538+
539+
# Tenant-a tries to access tenant-b's session → 404
540+
assert await _access_session(manager, session_b, tenant="tenant-a") == 404
541+
# Tenant-b tries to access tenant-a's session → 404
542+
assert await _access_session(manager, session_a, tenant="tenant-b") == 404
543+
stop.set()
544+
545+
546+
@pytest.mark.anyio
547+
async def test_same_tenant_can_reuse_session(running_manager: tuple[StreamableHTTPSessionManager, Server]):
548+
"""A request from the same tenant can access its own session."""
549+
manager, app = running_manager
550+
stop = anyio.Event()
551+
session_id = await _create_session_blocking(manager, app, stop, tenant="tenant-a")
552+
553+
status = await _access_session(manager, session_id, tenant="tenant-a")
554+
assert status != 404, "Same tenant should be able to reuse its own session"
555+
stop.set()
556+
557+
558+
@pytest.mark.anyio
559+
async def test_no_tenant_session_allows_any_access(running_manager: tuple[StreamableHTTPSessionManager, Server]):
560+
"""Sessions created without a tenant (no auth) allow access from any request."""
561+
manager, app = running_manager
562+
stop = anyio.Event()
563+
session_id = await _create_session_blocking(manager, app, stop, tenant=None)
564+
565+
status = await _access_session(manager, session_id, tenant="tenant-a")
566+
assert status != 404, "Session without tenant binding should allow access from any tenant"
567+
stop.set()
568+
569+
570+
@pytest.mark.anyio
571+
async def test_unauthenticated_request_cannot_access_tenant_session(
572+
running_manager: tuple[StreamableHTTPSessionManager, Server],
573+
):
574+
"""A request with no tenant cannot access a session bound to a tenant."""
575+
manager, app = running_manager
576+
stop = anyio.Event()
577+
session_id = await _create_session_blocking(manager, app, stop, tenant="tenant-a")
578+
579+
assert await _access_session(manager, session_id, tenant=None) == 404
580+
stop.set()
581+
582+
583+
@pytest.mark.anyio
584+
async def test_session_tenant_cleanup_on_exit(running_manager: tuple[StreamableHTTPSessionManager, Server]):
585+
"""Tenant mapping is cleaned up when a session exits."""
586+
manager, app = running_manager
587+
app.run = AsyncMock(return_value=None)
588+
589+
messages: list[Message] = []
590+
token = _set_tenant("tenant-a")
591+
try:
592+
await manager.handle_request(
593+
_make_scope(), _mock_receive, lambda msg, _msgs=messages: _mock_send(_msgs, msg)
594+
)
595+
finally:
596+
_reset_tenant(token)
597+
598+
session_id = _extract_session_id(messages)
599+
assert session_id is not None
600+
601+
# Wait for the mock server to complete and cleanup to run
602+
await anyio.sleep(0.01)
603+
604+
assert session_id not in manager._session_tenants, "Tenant mapping should be cleaned up after session exits"

0 commit comments

Comments
 (0)