Skip to content

Commit 9e3ded2

Browse files
committed
feat(auth): populate session.tenant_id from auth context on first request
Wire up session.tenant_id so it is set automatically from the auth contextvar on the first authenticated request (set-once semantics). This connects RequestContext.tenant_id and ServerSession.tenant_id, ensuring the session is bound to a tenant for its lifetime.
1 parent 9f4b679 commit 9e3ded2

2 files changed

Lines changed: 68 additions & 2 deletions

File tree

src/mcp/server/lowlevel/server.py

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -451,12 +451,15 @@ async def _handle_request(
451451
task_metadata = None
452452
if hasattr(req, "params") and req.params is not None:
453453
task_metadata = getattr(req.params, "task", None)
454+
tenant_id = get_tenant_id()
455+
if tenant_id is not None and session.tenant_id is None:
456+
session.tenant_id = tenant_id
454457
ctx = ServerRequestContext(
455458
request_id=message.request_id,
456459
meta=message.request_meta,
457460
session=session,
458461
lifespan_context=lifespan_context,
459-
tenant_id=get_tenant_id(),
462+
tenant_id=tenant_id,
460463
experimental=Experimental(
461464
task_metadata=task_metadata,
462465
_client_capabilities=client_capabilities,
@@ -496,10 +499,13 @@ async def _handle_notification(
496499
try:
497500
client_capabilities = session.client_params.capabilities if session.client_params else None
498501
task_support = self._experimental_handlers.task_support if self._experimental_handlers else None
502+
tenant_id = get_tenant_id()
503+
if tenant_id is not None and session.tenant_id is None:
504+
session.tenant_id = tenant_id
499505
ctx = ServerRequestContext(
500506
session=session,
501507
lifespan_context=lifespan_context,
502-
tenant_id=get_tenant_id(),
508+
tenant_id=tenant_id,
503509
experimental=Experimental(
504510
task_metadata=None,
505511
_client_capabilities=client_capabilities,

tests/server/test_multi_tenancy_session.py

Lines changed: 60 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -157,6 +157,66 @@ def test_get_tenant_id_from_auth_context():
157157
auth_context_var.reset(token)
158158

159159

160+
@pytest.mark.anyio
161+
async def test_session_tenant_id_set_from_auth_context_on_first_request(init_options: InitializationOptions):
162+
"""Verify session.tenant_id is populated from auth context on the first request.
163+
164+
The lowlevel server sets session.tenant_id from get_tenant_id() on the
165+
first request that has a tenant. This test simulates that behavior directly.
166+
"""
167+
server_to_client_send, server_to_client_recv = anyio.create_memory_object_stream[SessionMessage](1)
168+
client_to_server_send, client_to_server_recv = anyio.create_memory_object_stream[SessionMessage | Exception](1)
169+
170+
async with server_to_client_send, server_to_client_recv, client_to_server_send, client_to_server_recv:
171+
async with ServerSession(
172+
client_to_server_recv,
173+
server_to_client_send,
174+
init_options,
175+
) as session:
176+
assert session.tenant_id is None
177+
178+
# Simulate what lowlevel/server.py does: set session.tenant_id
179+
# from auth context on first request
180+
access_token = AccessToken(
181+
token="token-first",
182+
client_id="client",
183+
scopes=["read"],
184+
expires_at=int(time.time()) + 3600,
185+
tenant_id="tenant-first",
186+
)
187+
user = AuthenticatedUser(access_token)
188+
context_token = auth_context_var.set(user)
189+
try:
190+
tenant_id = get_tenant_id()
191+
if tenant_id is not None and session.tenant_id is None:
192+
session.tenant_id = tenant_id
193+
finally:
194+
auth_context_var.reset(context_token)
195+
196+
assert session.tenant_id == "tenant-first"
197+
198+
# Simulate a second request with a different tenant —
199+
# session.tenant_id should NOT change (set-once on first request)
200+
access_token2 = AccessToken(
201+
token="token-second",
202+
client_id="client",
203+
scopes=["read"],
204+
expires_at=int(time.time()) + 3600,
205+
tenant_id="tenant-second",
206+
)
207+
user2 = AuthenticatedUser(access_token2)
208+
context_token2 = auth_context_var.set(user2)
209+
try:
210+
tenant_id = get_tenant_id()
211+
if tenant_id is not None and session.tenant_id is None:
212+
session.tenant_id = tenant_id
213+
finally:
214+
auth_context_var.reset(context_token2)
215+
216+
# Still the first tenant — not overwritten
217+
assert session.tenant_id == "tenant-first"
218+
219+
160220
@pytest.mark.anyio
161221
async def test_tenant_context_isolation_between_concurrent_requests():
162222
"""Verify tenant_id doesn't leak between concurrent async contexts.

0 commit comments

Comments
 (0)