mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
resolve
This commit is contained in:
parent
8c59adab98
commit
0dc69ad66b
2 changed files with 191 additions and 16 deletions
|
|
@ -112,15 +112,44 @@ try:
|
|||
)
|
||||
_session_id_auth_storage: Dict[uuid.UUID, "MCPAuthenticatedUser"] = {}
|
||||
|
||||
# Robust auth lookup keyed by id(session_object). Unlike
|
||||
# _session_auth_storage (which relies on the SDK-private _read_stream
|
||||
# attribute), this dict uses the session's Python object identity and
|
||||
# needs no private-attribute access. Populated lazily on the first
|
||||
# successful ContextVar read inside get_or_extract_auth_context.
|
||||
_session_obj_auth_storage: Dict[int, "MCPAuthenticatedUser"] = {}
|
||||
|
||||
_captured_session_id_container_var: contextvars.ContextVar[
|
||||
Optional[Dict[str, uuid.UUID]]
|
||||
] = contextvars.ContextVar("captured_session_id_container", default=None)
|
||||
|
||||
class _SessionIdCapturingDict(dict):
|
||||
"""Intercepts ``__setitem__`` on the SSE transport's ``_read_stream_writers``
|
||||
dict so we can capture the session-ID assigned by the MCP SDK.
|
||||
|
||||
Contract assumption: the MCP SDK's ``connect_sse`` calls
|
||||
``self._read_stream_writers[session_id] = writer`` *synchronously* in
|
||||
the calling task (before spawning sub-tasks). If the SDK ever moves
|
||||
this into a child task, the ``_captured_session_id_container_var``
|
||||
lookup will return ``None`` and the capture will fail silently. The
|
||||
``else`` branch below logs a warning so we can detect this.
|
||||
"""
|
||||
|
||||
def __setitem__(self, key, value):
|
||||
container = _captured_session_id_container_var.get()
|
||||
if container is not None:
|
||||
container["session_id"] = key
|
||||
else:
|
||||
# Defensive: if we're inside a sub-task that doesn't inherit
|
||||
# the ContextVar, the capture silently fails. Log so
|
||||
# operators can diagnose auth-binding issues.
|
||||
verbose_logger.debug(
|
||||
"_SessionIdCapturingDict.__setitem__: "
|
||||
"_captured_session_id_container_var is None — "
|
||||
"session ID %s was NOT captured (expected during "
|
||||
"startup or if the MCP SDK changed internal task structure)",
|
||||
key,
|
||||
)
|
||||
super().__setitem__(key, value)
|
||||
|
||||
active_mcp_session_var: contextvars.ContextVar[Optional[_McpServerSession]] = (
|
||||
|
|
@ -2734,6 +2763,17 @@ if MCP_AVAILABLE:
|
|||
finally:
|
||||
if session_id:
|
||||
_session_id_auth_storage.pop(session_id, None)
|
||||
# Clean up session-object-identity storage.
|
||||
# The auth object stored here is the same reference
|
||||
# as ``auth`` above, so identity comparison is O(n)
|
||||
# but n is tiny (one entry per active SSE session).
|
||||
stale_keys = [
|
||||
k
|
||||
for k, v in _session_obj_auth_storage.items()
|
||||
if v is auth
|
||||
]
|
||||
for k in stale_keys:
|
||||
_session_obj_auth_storage.pop(k, None)
|
||||
finally:
|
||||
_captured_session_id_container_var.reset(_capture_token)
|
||||
|
||||
|
|
@ -2888,15 +2928,29 @@ if MCP_AVAILABLE:
|
|||
return auth
|
||||
elif auth:
|
||||
return cast(MCPAuthenticatedUser, auth)
|
||||
# Fallback to session read_stream (available in SDK request handlers)
|
||||
# Fallback: try session-object-identity lookup (robust, no private attrs)
|
||||
try:
|
||||
from mcp.server.lowlevel.server import request_ctx
|
||||
|
||||
session = request_ctx.get().session
|
||||
stored = _session_obj_auth_storage.get(id(session))
|
||||
if stored is not None:
|
||||
return stored
|
||||
# Last resort: legacy _read_stream lookup (private SDK attribute)
|
||||
read_stream = getattr(session, "_read_stream", None)
|
||||
return _session_auth_storage.get(read_stream) if read_stream else None
|
||||
if read_stream is not None:
|
||||
stored = _session_auth_storage.get(read_stream)
|
||||
if stored is not None:
|
||||
verbose_logger.debug(
|
||||
"get_active_auth_context: recovered auth via legacy "
|
||||
"_read_stream lookup — consider upgrading SDK"
|
||||
)
|
||||
# Promote to robust storage for future lookups
|
||||
_session_obj_auth_storage[id(session)] = stored
|
||||
return stored
|
||||
except Exception:
|
||||
return None
|
||||
pass
|
||||
return None
|
||||
|
||||
app = FastAPI(
|
||||
title=LITELLM_MCP_SERVER_NAME,
|
||||
|
|
@ -2999,8 +3053,16 @@ if MCP_AVAILABLE:
|
|||
Optional[str],
|
||||
]:
|
||||
"""
|
||||
Get auth context from ContextVar first, then fall back to the
|
||||
session read_stream (which survives cross-task boundaries in the MCP SDK).
|
||||
Get auth context from ContextVar first, then fall back to session
|
||||
storage (which survives cross-task boundaries in the MCP SDK).
|
||||
|
||||
Lookup priority:
|
||||
1. ContextVar (set by set_auth_context in the SSE/POST handler)
|
||||
2. _session_obj_auth_storage[id(session)] — robust, no private attrs
|
||||
3. _session_auth_storage[session._read_stream] — legacy, fragile
|
||||
|
||||
When (1) succeeds, we lazily cache the auth in (2) so future requests
|
||||
from SDK sub-tasks can find it without ContextVar.
|
||||
"""
|
||||
(
|
||||
user_api_key_auth,
|
||||
|
|
@ -3011,26 +3073,56 @@ if MCP_AVAILABLE:
|
|||
raw_headers,
|
||||
_client_ip,
|
||||
) = get_auth_context()
|
||||
# Fallback: read from session read_stream if ContextVar was lost
|
||||
if user_api_key_auth is None:
|
||||
|
||||
# Helper: resolve the current MCP session (if available)
|
||||
def _get_current_session():
|
||||
try:
|
||||
from mcp.server.lowlevel.server import request_ctx
|
||||
|
||||
session = request_ctx.get().session
|
||||
read_stream = getattr(session, "_read_stream", None)
|
||||
stored = _session_auth_storage.get(read_stream) if read_stream else None
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
f"get_or_extract_auth_context FALLBACK failed: {e}"
|
||||
return request_ctx.get().session
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
# When ContextVar has auth, lazily cache it by session identity so
|
||||
# subsequent requests from SDK sub-tasks can recover it.
|
||||
if user_api_key_auth is not None:
|
||||
session = _get_current_session()
|
||||
if session is not None and id(session) not in _session_obj_auth_storage:
|
||||
_session_obj_auth_storage[id(session)] = MCPAuthenticatedUser(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=_client_ip,
|
||||
)
|
||||
stored = None
|
||||
else:
|
||||
# Fallback: try session-object-identity lookup first (robust)
|
||||
stored: Optional[MCPAuthenticatedUser] = None
|
||||
session = _get_current_session()
|
||||
if session is not None:
|
||||
stored = _session_obj_auth_storage.get(id(session))
|
||||
if stored is None:
|
||||
# Last resort: legacy _read_stream lookup (private SDK attribute)
|
||||
read_stream = getattr(session, "_read_stream", None)
|
||||
if read_stream is not None:
|
||||
stored = _session_auth_storage.get(read_stream)
|
||||
if stored is not None:
|
||||
verbose_logger.debug(
|
||||
"get_or_extract_auth_context: recovered auth via "
|
||||
"legacy _read_stream lookup — consider upgrading SDK"
|
||||
)
|
||||
# Promote to robust storage for future lookups
|
||||
_session_obj_auth_storage[id(session)] = stored
|
||||
|
||||
verbose_logger.debug(
|
||||
f"get_or_extract_auth_context FALLBACK: stored={stored}, type={type(stored)}"
|
||||
f"get_or_extract_auth_context FALLBACK: stored={stored}, "
|
||||
f"type={type(stored)}"
|
||||
)
|
||||
if stored and isinstance(stored, MCPAuthenticatedUser):
|
||||
verbose_logger.debug(
|
||||
"get_or_extract_auth_context: Recovered auth from session read_stream"
|
||||
"get_or_extract_auth_context: Recovered auth from session storage"
|
||||
)
|
||||
user_api_key_auth = stored.user_api_key_auth
|
||||
mcp_auth_header = stored.mcp_auth_header
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from litellm.proxy._experimental.mcp_server.server import (
|
|||
handle_sse_post_messages,
|
||||
_captured_session_id_container_var,
|
||||
_session_id_auth_storage,
|
||||
_session_obj_auth_storage,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
|
@ -284,3 +285,85 @@ async def test_anonymous_session_still_works():
|
|||
|
||||
finish_run.set()
|
||||
await task
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_obj_auth_storage_lazy_cache_and_fallback():
|
||||
"""
|
||||
Verify that get_or_extract_auth_context:
|
||||
1. Lazily caches auth in _session_obj_auth_storage when ContextVar is set.
|
||||
2. Falls back to _session_obj_auth_storage when ContextVar is lost.
|
||||
|
||||
This proves the robust id(session)-based lookup works without relying
|
||||
on the private _read_stream attribute.
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
get_or_extract_auth_context,
|
||||
set_auth_context,
|
||||
)
|
||||
|
||||
# Create a fake session object (any object with a stable id)
|
||||
fake_session = MagicMock(name="FakeServerSession")
|
||||
session_id_key = id(fake_session)
|
||||
|
||||
# Create a fake request_ctx
|
||||
fake_request_ctx_value = MagicMock()
|
||||
fake_request_ctx_value.session = fake_session
|
||||
|
||||
test_auth = UserAPIKeyAuth(api_key="lazy-cache-key", user_id="lazy-user")
|
||||
|
||||
# Ensure _session_obj_auth_storage starts clean for this session
|
||||
_session_obj_auth_storage.pop(session_id_key, None)
|
||||
|
||||
# --- Step 1: ContextVar has auth, session is available ---
|
||||
# get_or_extract_auth_context should lazily cache auth.
|
||||
with patch(
|
||||
"mcp.server.lowlevel.server.request_ctx",
|
||||
) as mock_ctx_var:
|
||||
mock_ctx_var.get.return_value = fake_request_ctx_value
|
||||
|
||||
# Set the ContextVar-based auth
|
||||
set_auth_context(
|
||||
user_api_key_auth=test_auth,
|
||||
mcp_auth_header="auth-hdr",
|
||||
mcp_servers=["srv1"],
|
||||
mcp_server_auth_headers={},
|
||||
oauth2_headers={},
|
||||
raw_headers={},
|
||||
client_ip="127.0.0.1",
|
||||
)
|
||||
|
||||
result = await get_or_extract_auth_context()
|
||||
assert result[0] is not None
|
||||
assert result[0].api_key == "lazy-cache-key"
|
||||
|
||||
# Verify lazy cache was populated
|
||||
assert session_id_key in _session_obj_auth_storage
|
||||
cached = _session_obj_auth_storage[session_id_key]
|
||||
assert cached.user_api_key_auth.api_key == "lazy-cache-key"
|
||||
|
||||
# --- Step 2: ContextVar is cleared (simulates SDK sub-task) ---
|
||||
# get_or_extract_auth_context should recover from _session_obj_auth_storage.
|
||||
set_auth_context(
|
||||
user_api_key_auth=None,
|
||||
mcp_auth_header=None,
|
||||
mcp_servers=None,
|
||||
mcp_server_auth_headers=None,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
client_ip=None,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"mcp.server.lowlevel.server.request_ctx",
|
||||
) as mock_ctx_var:
|
||||
mock_ctx_var.get.return_value = fake_request_ctx_value
|
||||
|
||||
result = await get_or_extract_auth_context()
|
||||
assert (
|
||||
result[0] is not None
|
||||
), "Fallback to _session_obj_auth_storage failed — auth was lost"
|
||||
assert result[0].api_key == "lazy-cache-key"
|
||||
|
||||
# Cleanup
|
||||
_session_obj_auth_storage.pop(session_id_key, None)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue