Skip to content

Commit a1ce448

Browse files
authored
Merge pull request #13 from andylim-duo/feature/multi-tenant-session-manager
feat(session-manager): tenant validation on session access
2 parents e6d7fc7 + 57dd394 commit a1ce448

2 files changed

Lines changed: 242 additions & 0 deletions

File tree

src/mcp/server/streamable_http_manager.py

Lines changed: 29 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,29 @@ 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("Tenant mismatch for session %s", request_mcp_session_id[:64])
205+
logger.debug(
206+
"Tenant mismatch detail: session bound to '%s', request from '%s'",
207+
session_tenant,
208+
request_tenant,
209+
)
210+
error_response = JSONRPCError(
211+
jsonrpc="2.0",
212+
id=None,
213+
error=ErrorData(code=INVALID_REQUEST, message="Session not found"),
214+
)
215+
response = Response(
216+
content=error_response.model_dump_json(by_alias=True, exclude_unset=True),
217+
status_code=HTTPStatus.NOT_FOUND,
218+
media_type="application/json",
219+
)
220+
await response(scope, receive, send)
221+
return
222+
197223
transport = self._server_instances[request_mcp_session_id]
198224
logger.debug("Session already exists, handling request directly")
199225
# Push back idle deadline on activity
@@ -217,6 +243,7 @@ async def _handle_stateful_request(self, scope: Scope, receive: Receive, send: S
217243

218244
assert http_transport.mcp_session_id is not None
219245
self._server_instances[http_transport.mcp_session_id] = http_transport
246+
self._session_tenants[http_transport.mcp_session_id] = tenant_id_var.get()
220247
logger.info(f"Created new transport with session ID: {new_session_id}")
221248

222249
# Define the server runner
@@ -246,6 +273,7 @@ async def run_server(*, task_status: TaskStatus[None] = anyio.TASK_STATUS_IGNORE
246273
assert http_transport.mcp_session_id is not None
247274
logger.info(f"Session {http_transport.mcp_session_id} idle timeout")
248275
self._server_instances.pop(http_transport.mcp_session_id, None)
276+
self._session_tenants.pop(http_transport.mcp_session_id, None)
249277
await http_transport.terminate()
250278
except Exception:
251279
logger.exception(f"Session {http_transport.mcp_session_id} crashed")
@@ -260,6 +288,7 @@ async def run_server(*, task_status: TaskStatus[None] = anyio.TASK_STATUS_IGNORE
260288
f"{http_transport.mcp_session_id} from active instances."
261289
)
262290
del self._server_instances[http_transport.mcp_session_id]
291+
self._session_tenants.pop(http_transport.mcp_session_id, None)
263292

264293
# Assert task group is not None for type checking
265294
assert self._task_group is not None

tests/server/test_streamable_http_manager.py

Lines changed: 213 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -413,3 +413,216 @@ 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
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
437+
438+
439+
def test_extract_session_id_skips_non_start_messages():
440+
"""_extract_session_id skips non-start messages and returns None when no ID found."""
441+
body_msg: Message = {"type": "http.response.body", "body": b"data"}
442+
start_no_header: Message = {"type": "http.response.start", "status": 200, "headers": []}
443+
444+
# Only body messages → None
445+
assert _extract_session_id([body_msg]) is None
446+
# Start message without session header → None
447+
assert _extract_session_id([body_msg, start_no_header]) is None
448+
449+
450+
def test_extract_status_skips_non_start_messages():
451+
"""_extract_status skips non-start messages and returns None when empty."""
452+
body_msg: Message = {"type": "http.response.body", "body": b"data"}
453+
start_msg: Message = {"type": "http.response.start", "status": 200, "headers": []}
454+
455+
# Only body messages → None
456+
assert _extract_status([body_msg]) is None
457+
# Body then start → returns status from start
458+
assert _extract_status([body_msg, start_msg]) == 200
459+
# Empty list → None
460+
assert _extract_status([]) is None
461+
462+
463+
def _make_scope(session_id: str | None = None) -> dict[str, Any]:
464+
"""Build a minimal ASGI scope for testing, optionally with a session ID."""
465+
headers: list[tuple[bytes, bytes]] = [(b"content-type", b"application/json")]
466+
if session_id is not None:
467+
headers.append((b"mcp-session-id", session_id.encode()))
468+
return {"type": "http", "method": "POST", "path": "/mcp", "headers": headers}
469+
470+
471+
async def _mock_send(messages: list[Message], message: Message) -> None:
472+
"""Async send that collects messages."""
473+
messages.append(message)
474+
475+
476+
async def _mock_receive() -> dict[str, Any]: # pragma: no cover
477+
return {"type": "http.request", "body": b"", "more_body": False}
478+
479+
480+
def _set_tenant(tenant: str | None) -> Any:
481+
"""Set tenant_id_var and return the reset token."""
482+
from mcp.shared._context import tenant_id_var
483+
484+
return tenant_id_var.set(tenant)
485+
486+
487+
def _reset_tenant(token: Any) -> None:
488+
"""Reset tenant_id_var to its previous value."""
489+
from mcp.shared._context import tenant_id_var
490+
491+
tenant_id_var.reset(token)
492+
493+
494+
async def _create_session_blocking(
495+
manager: StreamableHTTPSessionManager,
496+
app: Server[Any],
497+
stop_event: anyio.Event,
498+
tenant: str | None = None,
499+
) -> str:
500+
"""Create a session whose server stays alive until stop_event is set."""
501+
502+
async def blocking_run(*args: Any, **kwargs: Any) -> None:
503+
await stop_event.wait()
504+
505+
app.run = AsyncMock(side_effect=blocking_run)
506+
507+
messages: list[Message] = []
508+
token = _set_tenant(tenant)
509+
try:
510+
await manager.handle_request(_make_scope(), _mock_receive, lambda msg, _msgs=messages: _mock_send(_msgs, msg))
511+
finally:
512+
_reset_tenant(token)
513+
514+
session_id = _extract_session_id(messages)
515+
assert session_id is not None
516+
return session_id
517+
518+
519+
async def _access_session(
520+
manager: StreamableHTTPSessionManager,
521+
session_id: str,
522+
tenant: str | None = None,
523+
) -> int | None:
524+
"""Access an existing session and return the HTTP status code."""
525+
messages: list[Message] = []
526+
token = _set_tenant(tenant)
527+
try:
528+
await manager.handle_request(
529+
_make_scope(session_id), _mock_receive, lambda msg, _msgs=messages: _mock_send(_msgs, msg)
530+
)
531+
finally:
532+
_reset_tenant(token)
533+
534+
return _extract_status(messages)
535+
536+
537+
@pytest.mark.anyio
538+
async def test_tenant_mismatch_returns_404(running_manager: tuple[StreamableHTTPSessionManager, Server]):
539+
"""A request from tenant-b cannot access a session created by tenant-a."""
540+
manager, app = running_manager
541+
stop = anyio.Event()
542+
try:
543+
session_id = await _create_session_blocking(manager, app, stop, tenant="tenant-a")
544+
assert await _access_session(manager, session_id, tenant="tenant-b") == 404
545+
finally:
546+
stop.set()
547+
548+
549+
@pytest.mark.anyio
550+
async def test_two_tenants_cannot_access_each_others_sessions(
551+
running_manager: tuple[StreamableHTTPSessionManager, Server],
552+
):
553+
"""Two tenants each create a session; neither can access the other's."""
554+
manager, app = running_manager
555+
stop = anyio.Event()
556+
try:
557+
session_a = await _create_session_blocking(manager, app, stop, tenant="tenant-a")
558+
session_b = await _create_session_blocking(manager, app, stop, tenant="tenant-b")
559+
assert session_a != session_b
560+
561+
# Tenant-a tries to access tenant-b's session → 404
562+
assert await _access_session(manager, session_b, tenant="tenant-a") == 404
563+
# Tenant-b tries to access tenant-a's session → 404
564+
assert await _access_session(manager, session_a, tenant="tenant-b") == 404
565+
finally:
566+
stop.set()
567+
568+
569+
@pytest.mark.anyio
570+
async def test_same_tenant_can_reuse_session(running_manager: tuple[StreamableHTTPSessionManager, Server]):
571+
"""A request from the same tenant can access its own session."""
572+
manager, app = running_manager
573+
stop = anyio.Event()
574+
try:
575+
session_id = await _create_session_blocking(manager, app, stop, tenant="tenant-a")
576+
status = await _access_session(manager, session_id, tenant="tenant-a")
577+
assert status != 404, "Same tenant should be able to reuse its own session"
578+
finally:
579+
stop.set()
580+
581+
582+
@pytest.mark.anyio
583+
async def test_no_tenant_session_allows_any_access(running_manager: tuple[StreamableHTTPSessionManager, Server]):
584+
"""Sessions created without a tenant (no auth) allow access from any request."""
585+
manager, app = running_manager
586+
stop = anyio.Event()
587+
try:
588+
session_id = await _create_session_blocking(manager, app, stop, tenant=None)
589+
status = await _access_session(manager, session_id, tenant="tenant-a")
590+
assert status != 404, "Session without tenant binding should allow access from any tenant"
591+
finally:
592+
stop.set()
593+
594+
595+
@pytest.mark.anyio
596+
async def test_unauthenticated_request_cannot_access_tenant_session(
597+
running_manager: tuple[StreamableHTTPSessionManager, Server],
598+
):
599+
"""A request with no tenant cannot access a session bound to a tenant."""
600+
manager, app = running_manager
601+
stop = anyio.Event()
602+
try:
603+
session_id = await _create_session_blocking(manager, app, stop, tenant="tenant-a")
604+
assert await _access_session(manager, session_id, tenant=None) == 404
605+
finally:
606+
stop.set()
607+
608+
609+
@pytest.mark.anyio
610+
async def test_session_tenant_cleanup_on_exit(running_manager: tuple[StreamableHTTPSessionManager, Server]):
611+
"""Tenant mapping is cleaned up when a session exits."""
612+
manager, app = running_manager
613+
app.run = AsyncMock(return_value=None)
614+
615+
messages: list[Message] = []
616+
token = _set_tenant("tenant-a")
617+
try:
618+
await manager.handle_request(_make_scope(), _mock_receive, lambda msg, _msgs=messages: _mock_send(_msgs, msg))
619+
finally:
620+
_reset_tenant(token)
621+
622+
session_id = _extract_session_id(messages)
623+
assert session_id is not None
624+
625+
# Wait for the mock server to complete and cleanup to run
626+
await anyio.sleep(0.01)
627+
628+
assert session_id not in manager._session_tenants, "Tenant mapping should be cleaned up after session exits"

0 commit comments

Comments
 (0)