mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
style(auth_v2): apply black formatting to satisfy lint
This commit is contained in:
parent
2e999ee1ad
commit
6e51b3fc90
8 changed files with 66 additions and 21 deletions
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]: ...
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue