Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py: 51%
202 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""Identity-only session tokens for the gateway-level (aggregate ``/mcp``) DCR front door.
3A DCR client that signs in through LiteLLM SSO holds ONE bearer that carries ONLY a
4litellm identity; unlike the :mod:`.envelope` bridge bearer it seals no upstream
5credential, because the custody model vaults every upstream token server-side in
6``LiteLLM_MCPUserCredentials`` and egress resolves them by user at call time. The token
7is therefore a stable REFERENCE, not an authorization: admission reloads the live user
8record and policy on every request, so deactivating the user (or their team) kills
9outstanding sessions immediately without a revocation store.
11Wire shape: ``llm_session_`` (access) / ``llm_srefresh_`` (refresh) + a JWT signed with
12the injected key material: HS256 under the default master-key-derived secret (the same
13signing approach as :mod:`.envelope`), or RS256 under an operator-provided RSA private
14key (:class:`AsymmetricSessionKeys`) so downstream validators hold only the public half.
15Claims are ``iss``/``iat``/``exp``
16plus ``jti`` (per-mint uniqueness, so two tokens minted in the same second never
17collide and a future revocation list has a stable handle), ``kind``, ``user_id``, and
18``client_id``; ``client_id`` binds the refresh token
19to the DCR client it was issued to (RFC 6749 section 6) and is carried on the access
20token for parity and audit. There is no encrypted payload: nothing in a session token
21is secret beyond the signature, and reprs never print the signed value because minted
22tokens are ``SecretStr``.
24This module is pure and unwired: it imports nothing from endpoint or edge code, reads
25no proxy globals, and takes all key material and the clock as explicit parameters.
26Failures are values: :func:`open_session_token` and :func:`open_session_refresh_token`
27are total over hostile, attacker-controlled input and return a
28``SessionTokenOpenError`` variant rather than raising. PyJWT's ``iat``/``nbf``/``exp``
29validators are disabled for the same reasons documented in :mod:`.envelope` (they
30raise on hostile claim types and compare against the wall clock instead of the
31injected ``now``); the strict pydantic claims model is the sole, total type gate.
32"""
34from __future__ import annotations
36import secrets
37from collections import Counter
38from datetime import datetime, timedelta
39from functools import lru_cache
40from typing import Final, Literal, TypeAlias
42import jwt
43from cryptography.exceptions import UnsupportedAlgorithm
44from cryptography.hazmat.primitives import serialization
45from cryptography.hazmat.primitives.asymmetric import rsa
46from pydantic import BaseModel, ConfigDict, Field, SecretStr, ValidationError, field_validator, model_validator
48SESSION_TOKEN_PREFIX: Final = "llm_session_"
49"""Marker prefix on every serialized session ACCESS token so the admission edge can cheaply
50tell a gateway session from a litellm key, JWT, or bridge envelope before doing any
51cryptography. Distinct from the ``llm_env_``/``llm_refresh_`` envelope prefixes."""
53SESSION_REFRESH_PREFIX: Final = "llm_srefresh_"
54"""Marker prefix on every serialized session REFRESH token. A distinct prefix keeps the two
55credentials routable without crypto and, together with the signed ``kind`` claim, stops one
56from being presented where the other is expected: the refresh token is only ever presented
57back to the token endpoint, never at the MCP edge."""
59SESSION_ISSUER: Final = "litellm-mcp-gateway"
60"""``iss`` claim stamped into every session token and required back on open. Distinct from
61the envelope issuer so a token of one family can never validate in the other even under a
62hypothetical shared signing key."""
64SESSION_TTL_SECONDS: Final = 3600
65"""Session ACCESS token lifetime (1h), matching the BYOK session bearer window: a
66client-held credential never outlives a bounded window, and each refresh re-validates
67the live user before re-minting."""
69SESSION_REFRESH_TTL_SECONDS: Final = 1209600
70"""Session REFRESH token lifetime (14 days), matching the refresh-envelope bound. Each
71renewal re-validates the sealed user against the live record (deactivation gates it) and
72rotates the refresh token, so the practical bound is idle time, not a fixed session."""
74MAX_SESSION_TOKEN_BYTES: Final = 4096
75"""Size cap on the serialized token (prefix + JWT, in bytes) and on any candidate accepted
76by the openers. Session claims are small; the only variable-length field is ``client_id``
77(a sealed DCR client record), and 4096 leaves ample headroom under common 8-16KB header
78limits while bounding hostile input before JWT parsing."""
80_SESSION_JWT_ALGORITHM: Final = "HS256"
82_SESSION_RSA_ALGORITHM: Final = "RS256"
84_MIN_RSA_KEY_BITS: Final = 2048
85"""RFC 7518 section 3.3: RS256 requires a key of at least 2048 bits."""
87SessionTokenKind = Literal["session", "session_refresh"]
88"""Which credential a session token is. Stamped into the signed claims and required to match
89on open, so a signature-valid token of one kind cannot be replayed as the other even if its
90wire prefix is swapped (the prefix is not part of the signed payload; this claim is)."""
92SessionAudience = Literal["proxy_api"]
93"""The non-MCP audience a session REFRESH token can be minted for. ``None`` (the default and
94the only value ever on an MCP wire) means the aggregate MCP gateway; ``"proxy_api"`` means the
95refresh grant re-mints the proxy-API CLI credential instead of an MCP session pair. The audience
96is read only from the signed claims, never from the request, so a token of one audience can
97never be redeemed as the other."""
100class SessionPrincipal(BaseModel):
101 """The litellm user a session token identifies and the DCR client it was issued to.
103 ``user_id`` is the SSO-established litellm user subject, never a credential: admission
104 reloads the live user record by it, so current role, team, and revocation state are
105 enforced at use time rather than frozen at mint time. ``client_id`` is the (stateless,
106 gateway-sealed) DCR client identifier the token was issued to; the token endpoint
107 requires it to match on the refresh grant.
109 ``resource_server_id`` is the single MCP server this session was authorized for when
110 the client requested a per-server RFC 8707 resource at authorize time, or ``None`` for
111 the aggregate scope. It is a RESTRICTION carried for admission to intersect against
112 the live grant resolution, never a grant by itself; the refresh grant re-mints from
113 this principal so the restriction survives rotation.
114 """
116 model_config = ConfigDict(frozen=True)
117 user_id: str = Field(min_length=1)
118 client_id: str = Field(min_length=1)
119 resource_server_id: str | None = None
120 audience: SessionAudience | None = None
121 team_id: str | None = None
124class SessionKeys(BaseModel):
125 """Injected key material: the HS256 signing key.
127 ``signing_key`` must be at least 32 bytes: HS256's HMAC-SHA256 has a 256-bit security
128 level, RFC 7518 requires a key of at least that size, and a shorter key makes PyJWT
129 emit ``InsecureKeyLengthWarning``.
130 """
132 model_config = ConfigDict(frozen=True)
133 signing_key: SecretStr = Field(min_length=32)
136class SessionRotatedPublicKey(BaseModel):
137 """The public half of a retired signing key, kept verifiable under its ``kid`` during a
138 rotation window so tokens minted before the rotation stay valid until they expire."""
140 model_config = ConfigDict(frozen=True)
141 kid: str = Field(min_length=1)
142 public_key_pem: str = Field(min_length=1)
144 @field_validator("public_key_pem")
145 @classmethod
146 def _pem_is_an_rsa_public_key(cls, value: str) -> str:
147 try:
148 loaded: Final = serialization.load_pem_public_key(value.encode())
149 except (ValueError, TypeError, UnsupportedAlgorithm) as exc:
150 raise ValueError(f"public_key_pem is not a loadable PEM public key: {exc}") from exc
151 if not isinstance(loaded, rsa.RSAPublicKey):
152 raise ValueError("public_key_pem must be an RSA public key in PEM format") # noqa: TRY004 # pydantic validators must raise ValueError
153 if loaded.key_size < _MIN_RSA_KEY_BITS:
154 raise ValueError(f"public_key_pem must be an RSA key of at least {_MIN_RSA_KEY_BITS} bits")
155 return value
158class AsymmetricSessionKeys(BaseModel):
159 """Injected RS256 key material: the issuer-held RSA private key and the stable ``kid``
160 stamped into every minted token's JOSE header, plus the public halves of previously
161 rotated keys that verification still accepts while their tokens age out. Downstream
162 validators never need the private key: :func:`session_public_key_pem` yields the
163 public half to distribute."""
165 model_config = ConfigDict(frozen=True)
166 private_key_pem: SecretStr
167 kid: str = Field(min_length=1)
168 previous_public_keys: tuple[SessionRotatedPublicKey, ...] = ()
170 @field_validator("private_key_pem")
171 @classmethod
172 def _pem_is_a_strong_rsa_private_key(cls, value: SecretStr) -> SecretStr:
173 try:
174 loaded: Final = serialization.load_pem_private_key(value.get_secret_value().encode(), password=None)
175 except (ValueError, TypeError, UnsupportedAlgorithm) as exc:
176 raise ValueError(f"private_key_pem is not a loadable unencrypted PEM private key: {exc}") from exc
177 if not isinstance(loaded, rsa.RSAPrivateKey):
178 raise ValueError("private_key_pem must be an unencrypted RSA private key in PEM format") # noqa: TRY004 # pydantic validators must raise ValueError
179 if loaded.key_size < _MIN_RSA_KEY_BITS:
180 raise ValueError(f"private_key_pem must be an RSA key of at least {_MIN_RSA_KEY_BITS} bits")
181 return value
183 @model_validator(mode="after")
184 def _kids_are_unique(self) -> AsymmetricSessionKeys:
185 kids: Final = (self.kid, *(previous.kid for previous in self.previous_public_keys))
186 duplicates: Final = tuple(kid for kid, count in Counter(kids).items() if count > 1)
187 if duplicates:
188 raise ValueError(
189 f"every kid must be unique across the current and previous keys; duplicated: {', '.join(duplicates)}"
190 )
191 return self
194SessionSigningKeys: TypeAlias = SessionKeys | AsymmetricSessionKeys
195"""Every key material shape the mints and openers accept: the default master-key-derived
196HS256 secret, or operator-configured RS256 RSA keys."""
199@lru_cache(maxsize=8)
200def _public_key_pem_from_private(private_key_pem: str) -> str:
201 loaded: Final = serialization.load_pem_private_key(private_key_pem.encode(), password=None)
202 return (
203 loaded.public_key()
204 .public_bytes(serialization.Encoding.PEM, serialization.PublicFormat.SubjectPublicKeyInfo)
205 .decode()
206 )
209def session_public_key_pem(keys: AsymmetricSessionKeys) -> str:
210 """The PEM public half of the current RS256 signing key: the only material a downstream
211 validator (an external gateway verifying ``kid``-matched tokens) ever needs."""
212 return _public_key_pem_from_private(keys.private_key_pem.get_secret_value())
215class MintedSessionToken(BaseModel):
216 """A minted session token: the client-held bearer value and when it expires."""
218 model_config = ConfigDict(frozen=True)
219 token: SecretStr
220 expires_at: datetime
223class OpenedSessionToken(BaseModel):
224 """A validated session token of either kind: the principal it was minted for, the
225 ``jti`` so the token endpoint can enforce single-use rotation on a refresh token, and
226 the signed ``kind``/``iat``/``exp`` so an introspection response can report the
227 token's metadata without re-decoding."""
229 model_config = ConfigDict(frozen=True)
230 principal: SessionPrincipal
231 jti: str
232 kind: SessionTokenKind
233 iat: int
234 exp: int
237class SessionTokenTooLarge(BaseModel):
238 """The serialized token exceeded ``MAX_SESSION_TOKEN_BYTES``; carries sizes only. Only
239 reachable through an oversized ``client_id``, which registration should have bounded."""
241 model_config = ConfigDict(frozen=True)
242 tag: Literal["session_token_too_large"] = "session_token_too_large"
243 size_bytes: int
244 max_bytes: int
247SessionTokenMintError: TypeAlias = SessionTokenTooLarge
250class NotASessionToken(BaseModel):
251 """The candidate does not carry the expected session prefix."""
253 model_config = ConfigDict(frozen=True)
254 tag: Literal["not_a_session_token"] = "not_a_session_token"
257class SessionBadSignature(BaseModel):
258 """The JWT signature does not verify under the provided signing key."""
260 model_config = ConfigDict(frozen=True)
261 tag: Literal["session_bad_signature"] = "session_bad_signature"
264class SessionExpired(BaseModel):
265 """The token's ``exp`` is not in the future relative to the provided ``now``."""
267 model_config = ConfigDict(frozen=True)
268 tag: Literal["session_expired"] = "session_expired"
271class SessionMalformed(BaseModel):
272 """The token is not a well-formed session token: undecodable JWT, wrong issuer, wrong
273 ``kind``, or missing/mistyped/extra claims."""
275 model_config = ConfigDict(frozen=True)
276 tag: Literal["session_malformed"] = "session_malformed"
279SessionTokenOpenError: TypeAlias = NotASessionToken | SessionBadSignature | SessionExpired | SessionMalformed
282class _SessionClaims(BaseModel):
283 """Decoded-claims boundary that pins the exact shape the mints emit.
285 ``user_id``/``client_id`` mirror the ``min_length`` constraints of
286 :class:`SessionPrincipal` so any claim set that validates here also constructs a
287 principal, keeping the openers raise-free: a correctly signed JWT with an empty
288 identity claim fails here and maps to ``SessionMalformed``. ``strict`` rejects coerced
289 types (``exp: "123"``) and ``extra="forbid"`` rejects any claim the gateway never
290 mints; PyJWT's own registered-claim validators are disabled at decode (see module
291 docstring), so this model is the sole, total type gate for every claim.
292 """
294 model_config = ConfigDict(frozen=True, strict=True, extra="forbid")
295 iss: str
296 iat: int
297 exp: int
298 jti: str = Field(min_length=1)
299 kind: SessionTokenKind
300 user_id: str = Field(min_length=1)
301 client_id: str = Field(min_length=1)
302 resource_server_id: str | None = None
303 audience: SessionAudience | None = None
304 team_id: str | None = None
307def is_session_token(candidate: str) -> bool:
308 """Cheap prefix check for a session ACCESS token so the admission edge can route gateway
309 sessions vs keys, JWTs, and envelopes without crypto."""
310 return candidate.startswith(SESSION_TOKEN_PREFIX)
313def is_session_refresh_token(candidate: str) -> bool:
314 """Cheap prefix check for a session REFRESH token so the token endpoint can route a
315 refresh grant without crypto."""
316 return candidate.startswith(SESSION_REFRESH_PREFIX)
319def mint_session_token(
320 principal: SessionPrincipal,
321 keys: SessionSigningKeys,
322 now: datetime,
323) -> MintedSessionToken | SessionTokenMintError:
324 """Mint the short-lived session ACCESS token for ``principal``.
326 ``exp`` is ``SESSION_TTL_SECONDS`` from ``now``. Returns ``SessionTokenTooLarge`` when
327 the serialized token exceeds ``MAX_SESSION_TOKEN_BYTES``.
328 """
329 return _mint(
330 kind="session",
331 prefix=SESSION_TOKEN_PREFIX,
332 principal=principal,
333 expires_at=now + timedelta(seconds=SESSION_TTL_SECONDS),
334 keys=keys,
335 now=now,
336 )
339def mint_session_refresh_token(
340 principal: SessionPrincipal,
341 keys: SessionSigningKeys,
342 now: datetime,
343) -> MintedSessionToken | SessionTokenMintError:
344 """Mint the long-lived session REFRESH token for ``principal``.
346 ``exp`` is ``SESSION_REFRESH_TTL_SECONDS`` from ``now``. Minting a distinct
347 ``kind="session_refresh"`` claim is what keeps a refresh token from ever opening as an
348 access credential at the MCP edge.
349 """
350 return _mint(
351 kind="session_refresh",
352 prefix=SESSION_REFRESH_PREFIX,
353 principal=principal,
354 expires_at=now + timedelta(seconds=SESSION_REFRESH_TTL_SECONDS),
355 keys=keys,
356 now=now,
357 )
360def open_session_token(
361 candidate: str,
362 keys: SessionSigningKeys,
363 now: datetime,
364) -> OpenedSessionToken | SessionTokenOpenError:
365 """Validate a session ACCESS ``candidate`` and recover the principal.
367 Never raises for bad input: every invalid, expired, tampered, or wrong-kind candidate
368 maps to a distinct ``SessionTokenOpenError`` variant.
369 """
370 return _open(candidate, prefix=SESSION_TOKEN_PREFIX, expected_kind="session", keys=keys, now=now)
373def open_session_refresh_token(
374 candidate: str,
375 keys: SessionSigningKeys,
376 now: datetime,
377) -> OpenedSessionToken | SessionTokenOpenError:
378 """Validate a session REFRESH ``candidate`` and recover the principal.
380 Total over hostile input exactly like :func:`open_session_token`. The
381 ``kind="session_refresh"`` claim is required, so an access token re-prefixed as a
382 refresh one is rejected as ``SessionMalformed``.
383 """
384 return _open(candidate, prefix=SESSION_REFRESH_PREFIX, expected_kind="session_refresh", keys=keys, now=now)
387def _mint(
388 kind: SessionTokenKind,
389 prefix: str,
390 principal: SessionPrincipal,
391 expires_at: datetime,
392 keys: SessionSigningKeys,
393 now: datetime,
394) -> MintedSessionToken | SessionTokenTooLarge:
395 """Sign the claims for either token kind and enforce the size cap. Shared by both mints
396 so the JWT shape, issuer, and size guard cannot drift between access and refresh."""
397 claims: Final = _SessionClaims(
398 iss=SESSION_ISSUER,
399 iat=int(now.timestamp()),
400 exp=int(expires_at.timestamp()),
401 jti=secrets.token_urlsafe(16),
402 kind=kind,
403 user_id=principal.user_id,
404 client_id=principal.client_id,
405 resource_server_id=principal.resource_server_id,
406 audience=principal.audience,
407 team_id=principal.team_id,
408 )
409 token: Final = prefix + _sign_claims(claims, keys)
410 size_bytes: Final = len(token.encode("utf-8"))
411 if size_bytes > MAX_SESSION_TOKEN_BYTES:
412 return SessionTokenTooLarge(size_bytes=size_bytes, max_bytes=MAX_SESSION_TOKEN_BYTES)
413 return MintedSessionToken(token=SecretStr(token), expires_at=expires_at)
416def _sign_claims(claims: _SessionClaims, keys: SessionSigningKeys) -> str:
417 """Sign the claim set under whichever key material was injected: RS256 with the ``kid``
418 in the JOSE header (so a validator can pick the right public key), or the default
419 HS256 secret with no header extras (byte-compatible with every pre-RS256 token)."""
420 payload: Final = claims.model_dump(exclude_none=True)
421 if isinstance(keys, AsymmetricSessionKeys):
422 return jwt.encode(
423 payload,
424 keys.private_key_pem.get_secret_value(),
425 algorithm=_SESSION_RSA_ALGORITHM,
426 headers={"kid": keys.kid},
427 )
428 return jwt.encode(payload, keys.signing_key.get_secret_value(), algorithm=_SESSION_JWT_ALGORITHM)
431def _open(
432 candidate: str,
433 prefix: str,
434 expected_kind: SessionTokenKind,
435 keys: SessionSigningKeys,
436 now: datetime,
437) -> OpenedSessionToken | SessionTokenOpenError:
438 """Prefix-route, size-bound, signature-verify, kind-check, and expiry-check an
439 attacker-controlled candidate, shared by both openers so the security gate is identical
440 for access and refresh. Returns the opened token or a distinct error; never raises."""
441 if not candidate.startswith(prefix):
442 return NotASessionToken()
443 # UTF-8 byte length is never below character length, so a character count already over
444 # the cap rejects an oversize candidate in O(1) without encoding it; the exact byte
445 # check then runs only on candidates already bounded to the cap in characters.
446 if len(candidate) > MAX_SESSION_TOKEN_BYTES:
447 return SessionMalformed()
448 if len(candidate.encode("utf-8", "surrogatepass")) > MAX_SESSION_TOKEN_BYTES:
449 return SessionMalformed()
450 claims: Final = _decode_claims(candidate.removeprefix(prefix), keys)
451 if not isinstance(claims, _SessionClaims):
452 return claims
453 if claims.kind != expected_kind:
454 return SessionMalformed()
455 if now.timestamp() >= claims.exp:
456 return SessionExpired()
457 return OpenedSessionToken(
458 principal=SessionPrincipal(
459 user_id=claims.user_id,
460 client_id=claims.client_id,
461 resource_server_id=claims.resource_server_id,
462 audience=claims.audience,
463 team_id=claims.team_id,
464 ),
465 jti=claims.jti,
466 kind=claims.kind,
467 iat=claims.iat,
468 exp=claims.exp,
469 )
472class _VerificationMaterial(BaseModel):
473 model_config = ConfigDict(frozen=True)
474 key: SecretStr
475 algorithm: Literal["HS256", "RS256"]
478def _verification_material(
479 compact: str,
480 keys: SessionSigningKeys,
481) -> _VerificationMaterial | SessionBadSignature | SessionMalformed:
482 """Pick the single key and algorithm the candidate is allowed to verify under.
484 HS256 mode has exactly one secret. RS256 mode routes by the JOSE header ``kid``: the
485 current key's derived public half, or a retired key's stored public half during a
486 rotation window. An unknown or missing ``kid`` is ``SessionBadSignature`` (a foreign
487 key), and an undecodable header is ``SessionMalformed``. The algorithm is pinned per
488 key shape, never read from the header, so an HS256 token can never be verified
489 against a public key or vice versa.
490 """
491 if isinstance(keys, SessionKeys):
492 return _VerificationMaterial(key=keys.signing_key, algorithm=_SESSION_JWT_ALGORITHM)
493 try:
494 header: Final = jwt.get_unverified_header(compact)
495 except jwt.InvalidTokenError:
496 return SessionMalformed()
497 kid: Final = header.get("kid")
498 if kid == keys.kid:
499 return _VerificationMaterial(key=SecretStr(session_public_key_pem(keys)), algorithm=_SESSION_RSA_ALGORITHM)
500 for previous in keys.previous_public_keys:
501 if previous.kid == kid:
502 return _VerificationMaterial(key=SecretStr(previous.public_key_pem), algorithm=_SESSION_RSA_ALGORITHM)
503 return SessionBadSignature()
506def _decode_claims(
507 compact: str,
508 keys: SessionSigningKeys,
509) -> _SessionClaims | SessionBadSignature | SessionMalformed:
510 """Verify the signature and shape of an attacker-controlled compact JWT.
512 ``compact`` is fully hostile and bounded to ``MAX_SESSION_TOKEN_BYTES`` by the caller.
513 The accepted algorithm is pinned by :func:`_verification_material` from the injected
514 key shape, so ``alg`` confusion (``none``, or HS256 signed with a public key as the
515 secret) fails before or at signature verification. PyJWT's ``iat``/``nbf``/``exp``
516 validators are disabled: they raise on hostile claim
517 types and, for ``iat``/``nbf``, compare against the wall clock rather than the injected
518 ``now`` (``exp`` is checked by the caller against ``now``). Apart from a signature
519 mismatch, every decode failure is ``SessionMalformed``: a non-UTF-8 candidate surfaces
520 as ``UnicodeEncodeError`` (a ``ValueError``), a non-string registered claim as a
521 ``TypeError`` from PyJWT's claim validators, and a wrong issuer or structurally invalid
522 token as an ``InvalidTokenError``. ``_SessionClaims`` is the total type gate.
523 """
524 material: Final = _verification_material(compact, keys)
525 if not isinstance(material, _VerificationMaterial):
526 return material
527 try:
528 payload: Final = jwt.decode(
529 compact,
530 material.key.get_secret_value(),
531 algorithms=[material.algorithm],
532 issuer=SESSION_ISSUER,
533 options={
534 "verify_exp": False,
535 "verify_iat": False,
536 "verify_nbf": False,
537 "require": ["iss", "iat", "exp"],
538 },
539 )
540 except jwt.InvalidSignatureError:
541 return SessionBadSignature()
542 except (jwt.InvalidTokenError, ValueError, TypeError):
543 return SessionMalformed()
544 try:
545 return _SessionClaims.model_validate(payload)
546 except ValidationError:
547 return SessionMalformed()