mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(mcp): make the EMA retention gate authoritative across pods and tighten the assertion store
This commit is contained in:
parent
83e147b553
commit
cfa074d970
5 changed files with 93 additions and 39 deletions
|
|
@ -404,7 +404,7 @@ model LiteLLM_MCPServerOAuthClient {
|
|||
}
|
||||
|
||||
// The enterprise IdP identity assertion captured at SSO login, one row per user.
|
||||
// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?, connected_at}.
|
||||
// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}.
|
||||
model LiteLLM_SSOIdentityAssertion {
|
||||
user_id String @id
|
||||
assertion_b64 String
|
||||
|
|
|
|||
|
|
@ -71,8 +71,10 @@ def assertion_from_sso_login(id_token: object, refresh_token: object) -> SSOIden
|
|||
try:
|
||||
claims = _IdTokenClaims.model_validate(jwt.decode(raw_id_token, options={"verify_signature": False}))
|
||||
expires_at = datetime.fromtimestamp(claims.exp, tz=timezone.utc) if claims.exp is not None else None
|
||||
except Exception: # noqa: BLE001 # any decode failure means the token is not EMA-exchangeable; never raise into login
|
||||
verbose_proxy_logger.warning("SSO id_token could not be decoded as a JWT; not retaining it for EMA egress.")
|
||||
except Exception: # noqa: BLE001 # decode failure = not retainable; never raise into login
|
||||
verbose_proxy_logger.warning(
|
||||
"SSO id_token could not be decoded or its claims were unusable; not retaining it for EMA egress."
|
||||
)
|
||||
return None
|
||||
return SSOIdentityAssertion(
|
||||
id_token=SecretStr(raw_id_token),
|
||||
|
|
@ -82,21 +84,31 @@ def assertion_from_sso_login(id_token: object, refresh_token: object) -> SSOIden
|
|||
)
|
||||
|
||||
|
||||
def ema_assertion_retention_enabled() -> bool:
|
||||
"""Whether any registered MCP server uses ``oauth2_id_jag``, evaluated per login so the
|
||||
gateway only retains bearer material while an EMA upstream exists to spend it on."""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # avoids a module-load circular import
|
||||
async def ema_assertion_retention_enabled() -> bool:
|
||||
"""Whether any MCP server uses ``oauth2_id_jag``, evaluated per login so the gateway only
|
||||
retains bearer material while an EMA upstream exists to spend it on. The local registry is
|
||||
the fast path; when it has none, the DB is consulted as the authoritative source, because
|
||||
the registry is a per-process snapshot while the assertion write targets the shared DB. A
|
||||
server added on another pod (or not yet loaded during startup) must still enable retention
|
||||
here, or the drop only surfaces later as an unexplained challenge at the EMA upstream."""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # avoids import cycle
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth # noqa: PLC0415 # runtime global, not wired at import time
|
||||
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global
|
||||
from litellm.types.mcp import MCPAuth # noqa: PLC0415 # runtime global
|
||||
|
||||
servers = global_mcp_server_manager.get_registry().values()
|
||||
return any(server.auth_type == MCPAuth.oauth2_id_jag for server in servers)
|
||||
if any(server.auth_type == MCPAuth.oauth2_id_jag for server in servers):
|
||||
return True
|
||||
if prisma_client is None:
|
||||
return False
|
||||
row = await prisma_client.db.litellm_mcpservertable.find_first(where={"auth_type": MCPAuth.oauth2_id_jag.value})
|
||||
return row is not None
|
||||
|
||||
|
||||
async def persist_sso_identity_assertion(user_id: str, assertion: SSOIdentityAssertion) -> None:
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper # noqa: PLC0415 # runtime global, not wired at import time
|
||||
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global, not wired at import time
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper # noqa: PLC0415 # runtime global
|
||||
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global
|
||||
|
||||
if prisma_client is None:
|
||||
return
|
||||
|
|
@ -119,8 +131,8 @@ async def persist_sso_identity_assertion(user_id: str, assertion: SSOIdentityAss
|
|||
async def fetch_sso_identity_assertion(user_id: str) -> SSOIdentityAssertion | None:
|
||||
"""The stored assertion for ``user_id``, or ``None`` when absent, undecryptable (salt-key
|
||||
rotation), or unparseable. Expiry is not judged here; the reader owns that policy."""
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper # noqa: PLC0415 # runtime global, not wired at import time
|
||||
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global, not wired at import time
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper # noqa: PLC0415 # runtime global
|
||||
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global
|
||||
|
||||
if prisma_client is None:
|
||||
return None
|
||||
|
|
@ -150,39 +162,38 @@ async def fetch_sso_identity_assertion(user_id: str) -> SSOIdentityAssertion | N
|
|||
async def rotate_sso_identity_assertions_master_key(prisma_client: PrismaClient, new_master_key: str) -> None:
|
||||
"""Re-encrypt every stored assertion under ``new_master_key`` during a salt-key rotation,
|
||||
mirroring the sibling per-user credential tables; an unreadable row is skipped so one
|
||||
corrupt row does not abort the rotation."""
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: PLC0415 # runtime global, not wired at import time
|
||||
corrupt row does not abort the rotation. Rows are decrypted one at a time inside the loop
|
||||
so the whole table's plaintext is never held in memory at once."""
|
||||
from prisma.models import LiteLLM_SSOIdentityAssertion as AssertionRow # noqa: PLC0415 # generated at runtime
|
||||
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: PLC0415 # runtime global
|
||||
decrypt_value_helper,
|
||||
encrypt_value_helper,
|
||||
)
|
||||
|
||||
rows = await prisma_client.db.litellm_ssoidentityassertion.find_many()
|
||||
decoded = [
|
||||
(
|
||||
row,
|
||||
_MAYBE_STR_ADAPTER.validate_python(
|
||||
decrypt_value_helper(row.assertion_b64, _ASSERTION_DECRYPT_LOG_KEY, exception_type="debug")
|
||||
),
|
||||
async def _rotate_row(row: AssertionRow) -> bool:
|
||||
plaintext = _MAYBE_STR_ADAPTER.validate_python(
|
||||
decrypt_value_helper(row.assertion_b64, _ASSERTION_DECRYPT_LOG_KEY, exception_type="debug")
|
||||
)
|
||||
for row in rows
|
||||
]
|
||||
for row, plaintext in decoded:
|
||||
if plaintext is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"rotate_sso_identity_assertions_master_key: could not decrypt assertion for user_id=%s, skipping",
|
||||
row.user_id,
|
||||
)
|
||||
continue
|
||||
return False
|
||||
re_encrypted = _STR_ADAPTER.validate_python(encrypt_value_helper(plaintext, new_encryption_key=new_master_key))
|
||||
await prisma_client.db.litellm_ssoidentityassertion.update(
|
||||
where={"user_id": row.user_id},
|
||||
data={"assertion_b64": re_encrypted},
|
||||
)
|
||||
rotated = sum(1 for _, plaintext in decoded if plaintext is not None)
|
||||
return True
|
||||
|
||||
rows = await prisma_client.db.litellm_ssoidentityassertion.find_many()
|
||||
outcomes = [await _rotate_row(row) for row in rows]
|
||||
verbose_proxy_logger.info(
|
||||
"rotate_sso_identity_assertions_master_key: rotated %d row(s), skipped %d",
|
||||
rotated,
|
||||
len(decoded) - rotated,
|
||||
sum(outcomes),
|
||||
len(outcomes) - sum(outcomes),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -193,7 +204,7 @@ async def retain_sso_identity_assertion_for_ema(user_id: str, assertion: SSOIden
|
|||
if assertion is None:
|
||||
return
|
||||
try:
|
||||
if not ema_assertion_retention_enabled():
|
||||
if not await ema_assertion_retention_enabled():
|
||||
return
|
||||
await persist_sso_identity_assertion(user_id, assertion)
|
||||
except Exception as exc: # noqa: BLE001 # the login itself must not fail on an egress-side write
|
||||
|
|
|
|||
|
|
@ -404,7 +404,7 @@ model LiteLLM_MCPServerOAuthClient {
|
|||
}
|
||||
|
||||
// The enterprise IdP identity assertion captured at SSO login, one row per user.
|
||||
// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?, connected_at}.
|
||||
// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}.
|
||||
model LiteLLM_SSOIdentityAssertion {
|
||||
user_id String @id
|
||||
assertion_b64 String
|
||||
|
|
|
|||
|
|
@ -404,7 +404,7 @@ model LiteLLM_MCPServerOAuthClient {
|
|||
}
|
||||
|
||||
// The enterprise IdP identity assertion captured at SSO login, one row per user.
|
||||
// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?, connected_at}.
|
||||
// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}.
|
||||
model LiteLLM_SSOIdentityAssertion {
|
||||
user_id String @id
|
||||
assertion_b64 String
|
||||
|
|
|
|||
|
|
@ -43,10 +43,15 @@ def _make_id_token(exp_offset: int = 3600, iss: str = ISSUER) -> str:
|
|||
)
|
||||
|
||||
|
||||
def _make_prisma(stored: dict):
|
||||
def _make_prisma(stored: dict, db_has_id_jag_server: bool = False):
|
||||
"""A fake prisma client whose sso-assertion table reads and writes ``stored``
|
||||
(user_id -> assertion_b64), covering upsert, find_unique, find_many, and update."""
|
||||
(user_id -> assertion_b64), covering upsert, find_unique, find_many, and update.
|
||||
``db_has_id_jag_server`` drives the retention gate's authoritative DB fallback;
|
||||
it is wired explicitly so the gate never reads a truthy bare MagicMock."""
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpservertable.find_first = AsyncMock(
|
||||
return_value=MagicMock() if db_has_id_jag_server else None
|
||||
)
|
||||
|
||||
async def _upsert(where, data):
|
||||
stored[where["user_id"]] = data["update"]["assertion_b64"]
|
||||
|
|
@ -125,18 +130,54 @@ def test_assertion_without_exp_or_iss_still_retained():
|
|||
assert assertion.issuer is None
|
||||
|
||||
|
||||
def test_retention_gate_requires_an_id_jag_server():
|
||||
with patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager:
|
||||
@pytest.mark.asyncio
|
||||
async def test_retention_gate_requires_an_id_jag_server():
|
||||
with (
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager,
|
||||
patch("litellm.proxy.proxy_server.prisma_client", _make_prisma({}, db_has_id_jag_server=False)),
|
||||
):
|
||||
manager.get_registry.return_value = {
|
||||
"s1": _server_with_auth(MCPAuth.oauth2),
|
||||
"s2": _server_with_auth(None),
|
||||
}
|
||||
assert ema_assertion_retention_enabled() is False
|
||||
assert await ema_assertion_retention_enabled() is False
|
||||
manager.get_registry.return_value = {
|
||||
"s1": _server_with_auth(MCPAuth.oauth2),
|
||||
"s2": _server_with_auth(MCPAuth.oauth2_id_jag),
|
||||
}
|
||||
assert ema_assertion_retention_enabled() is True
|
||||
assert await ema_assertion_retention_enabled() is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retention_gate_falls_back_to_db_when_registry_is_cold():
|
||||
"""A server added on another pod (or before this pod's DB load) is invisible to the local
|
||||
registry; the gate must still enable retention off the authoritative DB row, and read
|
||||
False only when neither the registry nor the DB knows an id_jag server."""
|
||||
with patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager:
|
||||
manager.get_registry.return_value = {"s1": _server_with_auth(MCPAuth.oauth2)}
|
||||
db_backed = _make_prisma({}, db_has_id_jag_server=True)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", db_backed):
|
||||
assert await ema_assertion_retention_enabled() is True
|
||||
db_backed.db.litellm_mcpservertable.find_first.assert_awaited_once_with(
|
||||
where={"auth_type": MCPAuth.oauth2_id_jag.value}
|
||||
)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", None):
|
||||
assert await ema_assertion_retention_enabled() is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retain_persists_when_only_the_db_knows_the_id_jag_server():
|
||||
stored = {}
|
||||
prisma = _make_prisma(stored, db_has_id_jag_server=True)
|
||||
with (
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager,
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma),
|
||||
):
|
||||
manager.get_registry.return_value = {}
|
||||
await retain_sso_identity_assertion_for_ema(
|
||||
user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None)
|
||||
)
|
||||
assert "user-a" in stored
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -274,11 +315,13 @@ async def test_rotation_reencrypts_under_new_key(monkeypatch):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rotation_skips_unreadable_rows():
|
||||
async def test_rotation_skips_unreadable_rows_but_rotates_readable_ones():
|
||||
stored = {"good": None, "bad": "garbage-blob"}
|
||||
prisma = _make_prisma(stored)
|
||||
token = _make_id_token()
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
|
||||
await persist_sso_identity_assertion("good", assertion_from_sso_login(token, None))
|
||||
good_blob_before = stored["good"]
|
||||
await rotate_sso_identity_assertions_master_key(prisma_client=prisma, new_master_key="another-new-salt-key-0000")
|
||||
assert stored["bad"] == "garbage-blob"
|
||||
assert stored["good"] != good_blob_before
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue