style(auth_v2): apply black formatting to satisfy lint

This commit is contained in:
Claude 2026-06-13 21:47:18 +00:00
parent 2e999ee1ad
commit 6e51b3fc90
No known key found for this signature in database
8 changed files with 66 additions and 21 deletions

View file

@ -21,9 +21,15 @@ def build_authenticators(
by_scheme: Dict[SecuritySchemeType, Authenticator] = {}
if config.api_key is not None:
by_scheme[SecuritySchemeType.API_KEY] = APIKeyAuthenticator(config.api_key)
by_scheme[SecuritySchemeType.HTTP] = HttpAuthenticator(config.http_basic, verifiers, basic_verifier)
by_scheme[SecuritySchemeType.HTTP] = HttpAuthenticator(
config.http_basic, verifiers, basic_verifier
)
by_scheme[SecuritySchemeType.OPENID_CONNECT] = OIDCAuthenticator(verifiers)
by_scheme[SecuritySchemeType.OAUTH2] = OAuth2Authenticator(verifiers, config.oauth2_introspection)
by_scheme[SecuritySchemeType.OAUTH2] = OAuth2Authenticator(
verifiers, config.oauth2_introspection
)
if config.mutual_tls.enabled:
by_scheme[SecuritySchemeType.MUTUAL_TLS] = MutualTLSAuthenticator(config.mutual_tls, config.network)
by_scheme[SecuritySchemeType.MUTUAL_TLS] = MutualTLSAuthenticator(
config.mutual_tls, config.network
)
return [by_scheme[scheme] for scheme in config.scheme_order if scheme in by_scheme]

View file

@ -11,7 +11,12 @@ from jwt import decode as jwt_decode
from starlette.concurrency import run_in_threadpool
from litellm.proxy.auth_v2 import errors
from litellm.proxy.auth_v2.models import AuthMethod, Credential, CredentialRef, SecuritySchemeType
from litellm.proxy.auth_v2.models import (
AuthMethod,
Credential,
CredentialRef,
SecuritySchemeType,
)
from litellm.proxy.auth_v2.config import OIDCProviderConfig
from litellm.proxy.auth_v2.authorization import filter_claim_roles
from litellm.proxy.auth_v2.authenticators.types import Claims
@ -20,7 +25,9 @@ AT_JWT_TYPES = {"at+jwt", "application/at+jwt"}
def apply_role_policy(claims: Claims, provider: OIDCProviderConfig) -> None:
claims["roles"] = filter_claim_roles(claims.get("roles"), provider.allowed_roles, provider.allow_platform_roles)
claims["roles"] = filter_claim_roles(
claims.get("roles"), provider.allowed_roles, provider.allow_platform_roles
)
def extract_bearer(request: Request) -> Optional[str]:
@ -65,7 +72,9 @@ def credential_from_claims(
audience=normalize_audience(claims.get("aud")),
scopes=split_scope(claims.get("scope")),
claims=claims,
credential_ref=CredentialRef(key_id=header.get("kid"), token_id=claims.get("jti")),
credential_ref=CredentialRef(
key_id=header.get("kid"), token_id=claims.get("jti")
),
subject_token=token,
)
@ -80,7 +89,9 @@ class JWTVerifier:
if jwks_client is not None:
self._jwks_client = jwks_client
return
jwks_uri = str(provider.jwks_uri) if provider.jwks_uri else self._discover_jwks()
jwks_uri = (
str(provider.jwks_uri) if provider.jwks_uri else self._discover_jwks()
)
self._jwks_client = PyJWKClient(
jwks_uri,
cache_keys=True,
@ -99,7 +110,9 @@ class JWTVerifier:
return str(jwks_uri)
def verify(self, token: str, *, require_at_jwt: Optional[bool] = None) -> Claims:
enforce = self.provider.require_at_jwt if require_at_jwt is None else require_at_jwt
enforce = (
self.provider.require_at_jwt if require_at_jwt is None else require_at_jwt
)
if enforce:
header = jwt.get_unverified_header(token)
if str(header.get("typ", "")).lower() not in AT_JWT_TYPES:
@ -118,8 +131,12 @@ class JWTVerifier:
raise errors.invalid_token("token verification failed") from exc
async def _verify_jwt_off_loop(verifier: JWTVerifier, token: str, *, require_at_jwt: Optional[bool] = None) -> Claims:
return await run_in_threadpool(functools.partial(verifier.verify, token, require_at_jwt=require_at_jwt))
async def _verify_jwt_off_loop(
verifier: JWTVerifier, token: str, *, require_at_jwt: Optional[bool] = None
) -> Claims:
return await run_in_threadpool(
functools.partial(verifier.verify, token, require_at_jwt=require_at_jwt)
)
def _select_verifier(token: str, verifiers: List[JWTVerifier]) -> Optional[JWTVerifier]:

View file

@ -56,7 +56,9 @@ class RBACEngine(Authorizer):
self._enforcer.add_policy(*rule)
def enforce(self, principal: "Principal", obj: str, act: str) -> bool:
return any(self._enforcer.enforce(role.value, obj, act) for role in principal.roles)
return any(
self._enforcer.enforce(role.value, obj, act) for role in principal.roles
)
def has_any_role(self, principal: "Principal", allowed: Tuple[Role, ...]) -> bool:
allowed_values = {role.value for role in allowed}

View file

@ -16,7 +16,9 @@ class Role(str, Enum):
_PLATFORM_ROLE_VALUES = {Role.PLATFORM_ADMIN.value, Role.PLATFORM_VIEWER.value}
def filter_claim_roles(roles: Any, allowed_roles: List[str], allow_platform_roles: bool) -> List[str]:
def filter_claim_roles(
roles: Any, allowed_roles: List[str], allow_platform_roles: bool
) -> List[str]:
if not isinstance(roles, list):
return []
allowed = set(allowed_roles)

View file

@ -6,12 +6,16 @@ from fastapi import HTTPException
class AuthError(HTTPException):
def __init__(self, status_code: int, detail: str, challenge: Optional[str] = None) -> None:
def __init__(
self, status_code: int, detail: str, challenge: Optional[str] = None
) -> None:
headers = {"WWW-Authenticate": challenge} if challenge else None
super().__init__(status_code=status_code, detail=detail, headers=headers)
def bearer_challenge(error: Optional[str] = None, description: Optional[str] = None) -> str:
def bearer_challenge(
error: Optional[str] = None, description: Optional[str] = None
) -> str:
parts = ['Bearer realm="litellm"']
if error:
parts.append(f'error="{error}"')
@ -29,7 +33,9 @@ def unauthenticated(challenge: str) -> AuthError:
def invalid_token(description: Optional[str] = None) -> AuthError:
return AuthError(401, "Invalid token", bearer_challenge("invalid_token", description))
return AuthError(
401, "Invalid token", bearer_challenge("invalid_token", description)
)
def insufficient_scope() -> AuthError:

View file

@ -16,7 +16,9 @@ class SessionStore(Protocol[SessionValue]):
async def get(self, key: str) -> Optional[SessionValue]: ...
async def set(self, key: str, value: SessionValue, ttl_seconds: Optional[int] = None) -> None: ...
async def set(
self, key: str, value: SessionValue, ttl_seconds: Optional[int] = None
) -> None: ...
async def pop(self, key: str) -> Optional[SessionValue]: ...

View file

@ -21,7 +21,9 @@ class InMemorySessionStore(Generic[SessionValue]):
def _expiry(self, ttl_seconds: Optional[int]) -> float:
return time.time() + (self._default_ttl if ttl_seconds is None else ttl_seconds)
def _live(self, key: str, now: float) -> Optional[Tuple[float, Optional[SessionValue]]]:
def _live(
self, key: str, now: float
) -> Optional[Tuple[float, Optional[SessionValue]]]:
entry = self._entries.get(key)
if entry is None:
return None
@ -34,7 +36,9 @@ class InMemorySessionStore(Generic[SessionValue]):
entry = self._live(self._key(key), time.time())
return entry[1] if entry is not None else None
async def set(self, key: str, value: SessionValue, ttl_seconds: Optional[int] = None) -> None:
async def set(
self, key: str, value: SessionValue, ttl_seconds: Optional[int] = None
) -> None:
self._evict(time.time())
self._entries[self._key(key)] = (self._expiry(ttl_seconds), value)

View file

@ -27,8 +27,12 @@ class RedisSessionStore(Generic[SessionValue]):
raw = await self._client.get(self._key(key))
return cast(SessionValue, json.loads(raw)) if raw is not None else None
async def set(self, key: str, value: SessionValue, ttl_seconds: Optional[int] = None) -> None:
await self._client.set(self._key(key), json.dumps(value), ex=self._ttl(ttl_seconds))
async def set(
self, key: str, value: SessionValue, ttl_seconds: Optional[int] = None
) -> None:
await self._client.set(
self._key(key), json.dumps(value), ex=self._ttl(ttl_seconds)
)
async def pop(self, key: str) -> Optional[SessionValue]:
raw = await self._client.getdel(self._key(key))
@ -38,5 +42,7 @@ class RedisSessionStore(Generic[SessionValue]):
await self._client.delete(self._key(key))
async def add_if_absent(self, key: str, ttl_seconds: Optional[int] = None) -> bool:
added = await self._client.set(self._key(key), "1", nx=True, ex=self._ttl(ttl_seconds))
added = await self._client.set(
self._key(key), "1", nx=True, ex=self._ttl(ttl_seconds)
)
return bool(added)