mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
refactor(proxy): fully type auth_v2 - eliminate every avoidable Any
The decision core used Any where it was over-cautious about coupling. Replace it with real types (under TYPE_CHECKING, so the framework-free core stays import- light) and small Protocols: - identity is UserAPIKeyAuth end to end (authenticators, entry, principal, context, stages); the authenticator chain returns it, not Any - SupportsEnforce types the authorizer's engine param; CasbinRuleRow / PolicyDB / PolicyAdminDB type the dynamic Prisma casbin-rule access (read and write), narrowed with one explicit cast at the adapter boundary - AuthContext deps are PrismaClient / UserApiKeyCache / ProxyLogging / Span - loaders, settings, key sets, role conversion all carry their real types The only Any left is Dict[str, Any] for genuinely arbitrary JSON (decoded JWT claims, RFC 7662 introspection responses, request bodies) - that is the honest type, not a gap. Verified: mypy clean on all 21 files, zero bare Any annotations, 129 tests green, still imports with casbin/authlib absent.
This commit is contained in:
parent
764ffd834e
commit
69a8d4f986
12 changed files with 203 additions and 78 deletions
|
|
@ -1,20 +1,29 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import secrets
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, List, Optional, Protocol, runtime_checkable
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Protocol, runtime_checkable
|
||||
|
||||
from fastapi import HTTPException, status
|
||||
|
||||
from ..context import AuthMethod
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
|
||||
from .jwt_claims import JWTSettings
|
||||
from .oauth2_introspection import IntrospectionSettings
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AuthResult:
|
||||
"""The output of the authenticator chain: the resolved identity and how."""
|
||||
|
||||
identity: "UserAPIKeyAuth"
|
||||
identity: UserAPIKeyAuth
|
||||
method: AuthMethod
|
||||
|
||||
|
||||
|
|
@ -24,7 +33,7 @@ class Authenticator(Protocol):
|
|||
|
||||
def can_handle(self, api_key: Optional[str]) -> bool: ...
|
||||
|
||||
async def authenticate(self, api_key: str, ctx: "AuthContext") -> Any: ...
|
||||
async def authenticate(self, api_key: str, ctx: AuthContext) -> UserAPIKeyAuth: ...
|
||||
|
||||
|
||||
class AuthContext:
|
||||
|
|
@ -32,10 +41,10 @@ class AuthContext:
|
|||
|
||||
def __init__(
|
||||
self,
|
||||
prisma_client: Any,
|
||||
user_api_key_cache: Any,
|
||||
proxy_logging_obj: Any,
|
||||
parent_otel_span: Any = None,
|
||||
prisma_client: Optional[PrismaClient],
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
parent_otel_span: Optional[Span] = None,
|
||||
):
|
||||
self.prisma_client = prisma_client
|
||||
self.user_api_key_cache = user_api_key_cache
|
||||
|
|
@ -63,7 +72,7 @@ class MasterKeyAuthenticator:
|
|||
except Exception:
|
||||
return False
|
||||
|
||||
async def authenticate(self, api_key: str, ctx: AuthContext) -> Any:
|
||||
async def authenticate(self, api_key: str, ctx: AuthContext) -> UserAPIKeyAuth:
|
||||
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.proxy_server import litellm_proxy_admin_name
|
||||
|
|
@ -83,7 +92,7 @@ class VirtualKeyAuthenticator:
|
|||
def can_handle(self, api_key: Optional[str]) -> bool:
|
||||
return isinstance(api_key, str) and api_key.startswith("sk-")
|
||||
|
||||
async def authenticate(self, api_key: str, ctx: AuthContext) -> Any:
|
||||
async def authenticate(self, api_key: str, ctx: AuthContext) -> UserAPIKeyAuth:
|
||||
from litellm.proxy._types import hash_token
|
||||
from litellm.proxy.auth.auth_checks import get_key_object
|
||||
|
||||
|
|
@ -96,7 +105,7 @@ class VirtualKeyAuthenticator:
|
|||
)
|
||||
|
||||
|
||||
def _load_jwt_settings() -> Any:
|
||||
def _load_jwt_settings() -> JWTSettings:
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
from .jwt_claims import JWTSettings
|
||||
|
|
@ -142,7 +151,7 @@ class JWTAuthenticator:
|
|||
return False
|
||||
return _jwt_is_configured()
|
||||
|
||||
async def authenticate(self, api_key: str, ctx: AuthContext) -> Any:
|
||||
async def authenticate(self, api_key: str, ctx: AuthContext) -> UserAPIKeyAuth:
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
|
||||
from .jwt_claims import extract_identity
|
||||
|
|
@ -174,7 +183,7 @@ class JWTAuthenticator:
|
|||
)
|
||||
|
||||
|
||||
def _load_introspection_settings() -> Any:
|
||||
def _load_introspection_settings() -> Optional[IntrospectionSettings]:
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
from .oauth2_introspection import IntrospectionSettings
|
||||
|
|
@ -208,7 +217,7 @@ class OAuth2IntrospectionAuthenticator:
|
|||
return False
|
||||
return _load_introspection_settings() is not None
|
||||
|
||||
async def authenticate(self, api_key: str, ctx: AuthContext) -> Any:
|
||||
async def authenticate(self, api_key: str, ctx: AuthContext) -> UserAPIKeyAuth:
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
|
||||
from .oauth2_introspection import (
|
||||
|
|
@ -217,6 +226,11 @@ class OAuth2IntrospectionAuthenticator:
|
|||
)
|
||||
|
||||
settings = _load_introspection_settings()
|
||||
if settings is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="auth_v2: OAuth2 introspection is not configured",
|
||||
)
|
||||
data = await self._introspect(api_key, settings)
|
||||
try:
|
||||
identity = parse_introspection_response(data, settings)
|
||||
|
|
@ -238,7 +252,9 @@ class OAuth2IntrospectionAuthenticator:
|
|||
user_role=user_role,
|
||||
)
|
||||
|
||||
async def _introspect(self, token: str, settings: Any) -> Any:
|
||||
async def _introspect(
|
||||
self, token: str, settings: IntrospectionSettings
|
||||
) -> Dict[str, Any]:
|
||||
import base64
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
|
|
|
|||
|
|
@ -1,5 +1,10 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from authlib.jose import KeySet
|
||||
|
||||
|
||||
class JWTVerificationError(Exception):
|
||||
|
|
@ -38,7 +43,7 @@ def build_claims_options(
|
|||
|
||||
def verify(
|
||||
token: str,
|
||||
key_set: Any,
|
||||
key_set: KeySet,
|
||||
issuer: Optional[str] = None,
|
||||
audience: Optional[str] = None,
|
||||
algorithms: Optional[List[str]] = None,
|
||||
|
|
@ -74,10 +79,10 @@ class JWKSProvider:
|
|||
def __init__(self, jwks_uri: str, ttl_seconds: float = 600.0):
|
||||
self.jwks_uri = jwks_uri
|
||||
self.ttl_seconds = ttl_seconds
|
||||
self._key_set: Any = None
|
||||
self._key_set: Optional[KeySet] = None
|
||||
self._fetched_at: float = 0.0
|
||||
|
||||
async def get_key_set(self) -> Any:
|
||||
async def get_key_set(self) -> KeySet:
|
||||
now = time.monotonic()
|
||||
if self._key_set is not None and (now - self._fetched_at) < self.ttl_seconds:
|
||||
return self._key_set
|
||||
|
|
|
|||
|
|
@ -25,7 +25,7 @@ class IntrospectionIdentity:
|
|||
role: Optional[str] = None
|
||||
|
||||
|
||||
def _scopes(raw: Any) -> List[str]:
|
||||
def _scopes(raw: object) -> List[str]:
|
||||
if raw is None:
|
||||
return []
|
||||
if isinstance(raw, str):
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import logging
|
|||
from typing import Any, Dict, Optional
|
||||
|
||||
from ..principal import Principal
|
||||
from ..protocols import SupportsEnforce
|
||||
from .route_map import GovernedRoute, match_route
|
||||
|
||||
logger = logging.getLogger("litellm.proxy.auth.v2")
|
||||
|
|
@ -32,7 +33,7 @@ def authorize(
|
|||
principal: Principal,
|
||||
route: str,
|
||||
request_data: Optional[Dict[str, Any]],
|
||||
enforcer: Any,
|
||||
enforcer: SupportsEnforce,
|
||||
method: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Enforce policy for ``route``. No-op (loud) for routes v2 doesn't yet govern.
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
import time
|
||||
from typing import Any, List, Optional, Tuple
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
from ..protocols import CasbinRuleRow, PolicyDB
|
||||
|
||||
# Always-present bootstrap: principals carrying the proxy_admin role keep full
|
||||
# access so enabling auth_v2 never locks admins out. Granular custom roles are
|
||||
|
|
@ -23,27 +25,20 @@ _cache: Optional[
|
|||
] = None
|
||||
|
||||
|
||||
def _row_values(row: Any) -> List[str]:
|
||||
values = [
|
||||
getattr(row, "v0", None),
|
||||
getattr(row, "v1", None),
|
||||
getattr(row, "v2", None),
|
||||
getattr(row, "v3", None),
|
||||
getattr(row, "v4", None),
|
||||
getattr(row, "v5", None),
|
||||
]
|
||||
def _row_values(row: CasbinRuleRow) -> List[str]:
|
||||
values = [row.v0, row.v1, row.v2, row.v3, row.v4, row.v5]
|
||||
return [v for v in values if v is not None and v != ""]
|
||||
|
||||
|
||||
def _split_rules(
|
||||
rows: List[Any],
|
||||
rows: List[CasbinRuleRow],
|
||||
) -> Tuple[List[List[str]], List[List[str]], List[List[str]], List[List[str]]]:
|
||||
policies: List[List[str]] = [list(p) for p in DEFAULT_POLICIES]
|
||||
groupings: List[List[str]] = []
|
||||
resource_groupings: List[List[str]] = []
|
||||
domain_groupings: List[List[str]] = []
|
||||
for row in rows:
|
||||
ptype = getattr(row, "ptype", None)
|
||||
ptype = row.ptype
|
||||
values = _row_values(row)
|
||||
if ptype == "p":
|
||||
policies.append(values)
|
||||
|
|
@ -62,7 +57,7 @@ def reset_cache() -> None:
|
|||
|
||||
|
||||
async def load_policy_snapshot(
|
||||
prisma_client: Any,
|
||||
prisma_client: Optional[PolicyDB],
|
||||
) -> Tuple[List[List[str]], List[List[str]], List[List[str]], List[List[str]]]:
|
||||
"""Load (policies, groupings, resource_groupings, domain_groupings) from
|
||||
LiteLLM_CasbinRule, with a short TTL cache.
|
||||
|
|
@ -75,7 +70,7 @@ async def load_policy_snapshot(
|
|||
if _cache is not None and (now - _cache[0]) < _CACHE_TTL_SECONDS:
|
||||
return _cache[1], _cache[2], _cache[3], _cache[4]
|
||||
|
||||
rows: List[Any] = []
|
||||
rows: List[CasbinRuleRow] = []
|
||||
if prisma_client is not None:
|
||||
rows = await prisma_client.db.litellm_casbinrule.find_many()
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,6 @@
|
|||
from typing import Any, Optional, Tuple
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Optional, Tuple, cast
|
||||
|
||||
from fastapi import HTTPException, Request, status
|
||||
|
||||
|
|
@ -11,29 +13,43 @@ from .authz.policy_store import load_policy_snapshot
|
|||
from .authz.route_map import is_inference_route, match_route
|
||||
from .context import AuthMethod, RequestAuthContext, set_auth_context
|
||||
from .principal import Principal, build_principal
|
||||
from .protocols import PolicyDB
|
||||
from .stages.budgets import enforce_hierarchy_budgets
|
||||
from .stages.end_user import resolve_end_user
|
||||
from .stages.enrichment import enrich_identity
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_TeamTable,
|
||||
LiteLLM_UserTable,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
async def _anonymous_identity(api_key: Optional[str]) -> Any:
|
||||
|
||||
async def _anonymous_identity(api_key: Optional[str]) -> UserAPIKeyAuth:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
return UserAPIKeyAuth(api_key=api_key)
|
||||
|
||||
|
||||
async def _build_enforcer(principal: Principal, prisma_client: Any) -> CasbinEnforcer:
|
||||
async def _build_enforcer(
|
||||
principal: Principal, prisma_client: Optional[PrismaClient]
|
||||
) -> CasbinEnforcer:
|
||||
"""Build the casbin engine for this principal from the current policy snapshot.
|
||||
|
||||
The principal's identity-to-role bridges are added on top of the stored
|
||||
groupings. One engine authorizes both the control plane and model calls.
|
||||
"""
|
||||
# The Prisma client's casbin-rule table is generated dynamically; narrow it to
|
||||
# the read surface the store needs at this single adapter boundary.
|
||||
policy_db = cast(Optional[PolicyDB], prisma_client)
|
||||
(
|
||||
policies,
|
||||
groupings,
|
||||
resource_groupings,
|
||||
domain_groupings,
|
||||
) = await load_policy_snapshot(prisma_client)
|
||||
) = await load_policy_snapshot(policy_db)
|
||||
return CasbinEnforcer(
|
||||
policies,
|
||||
groupings + principal.groupings,
|
||||
|
|
@ -42,13 +58,13 @@ async def _build_enforcer(principal: Principal, prisma_client: Any) -> CasbinEnf
|
|||
)
|
||||
|
||||
|
||||
async def _enrich_for_limits(identity: Any, ctx: AuthContext) -> None:
|
||||
async def _enrich_for_limits(identity: UserAPIKeyAuth, ctx: AuthContext) -> None:
|
||||
"""Fill user/team budget+limit fields for non-key logins (master/JWT/OAuth) so
|
||||
the existing pre-call budget/limit hooks can enforce them. Virtual keys are
|
||||
already populated by get_key_object and skip this."""
|
||||
from litellm.proxy.auth.auth_checks import get_team_object, get_user_object
|
||||
|
||||
async def load_user(user_id: str) -> Any:
|
||||
async def load_user(user_id: str) -> Optional[LiteLLM_UserTable]:
|
||||
return await get_user_object(
|
||||
user_id=user_id,
|
||||
prisma_client=ctx.prisma_client,
|
||||
|
|
@ -58,7 +74,7 @@ async def _enrich_for_limits(identity: Any, ctx: AuthContext) -> None:
|
|||
proxy_logging_obj=ctx.proxy_logging_obj,
|
||||
)
|
||||
|
||||
async def load_team(team_id: str) -> Any:
|
||||
async def load_team(team_id: str) -> Optional[LiteLLM_TeamTable]:
|
||||
return await get_team_object(
|
||||
team_id=team_id,
|
||||
prisma_client=ctx.prisma_client,
|
||||
|
|
@ -70,7 +86,9 @@ async def _enrich_for_limits(identity: Any, ctx: AuthContext) -> None:
|
|||
await enrich_identity(identity, load_user=load_user, load_team=load_team)
|
||||
|
||||
|
||||
async def _enforce_budgets(identity: Any, route: str, ctx: AuthContext) -> None:
|
||||
async def _enforce_budgets(
|
||||
identity: UserAPIKeyAuth, route: str, ctx: AuthContext
|
||||
) -> None:
|
||||
"""Enforce team/org/global budgets (reusing v1's functions) and surface a
|
||||
breach as the same ProxyException v1 raises."""
|
||||
import litellm
|
||||
|
|
@ -101,7 +119,7 @@ async def _best_effort_identity(api_key: Optional[str], ctx: AuthContext) -> Aut
|
|||
|
||||
def _establish_context(
|
||||
request: Request, result: AuthResult, route: str
|
||||
) -> Tuple[Principal, Any]:
|
||||
) -> Tuple[Principal, UserAPIKeyAuth]:
|
||||
"""Derive the principal and publish the typed auth context for downstream
|
||||
stages (budget/limit hooks, end-user resolver, telemetry span)."""
|
||||
principal = build_principal(result.identity)
|
||||
|
|
@ -120,7 +138,7 @@ def _establish_context(
|
|||
async def user_api_key_auth_v2(
|
||||
request: Request,
|
||||
api_key: str = "",
|
||||
) -> Any:
|
||||
) -> UserAPIKeyAuth:
|
||||
"""auth_v2 entry point: authenticator chain establishes identity, casbin
|
||||
authorizes governed routes. Routes v2 doesn't yet own are loud-open."""
|
||||
from litellm.proxy.auth.user_api_key_auth import _get_bearer_token
|
||||
|
|
@ -194,7 +212,7 @@ async def user_api_key_auth_v2(
|
|||
|
||||
|
||||
class _DenyAll:
|
||||
def enforce(self, *_args: Any) -> bool:
|
||||
def enforce(self, *_args: object) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import Any, Dict, List, Optional
|
||||
from typing import Dict, List, Optional, cast
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -12,6 +12,7 @@ from .authz.policy_admin import (
|
|||
make_permission_rule,
|
||||
)
|
||||
from .authz.policy_store import reset_cache
|
||||
from .protocols import CasbinRuleRow, PolicyAdminDB
|
||||
|
||||
router = APIRouter(tags=["auth_v2"])
|
||||
|
||||
|
|
@ -41,10 +42,11 @@ def rule_to_row_data(rule: List[str]) -> Dict[str, str]:
|
|||
return data
|
||||
|
||||
|
||||
def row_to_rule(row: Any) -> List[str]:
|
||||
rule = [row.ptype]
|
||||
for index in range(6):
|
||||
value = getattr(row, f"v{index}", None)
|
||||
def row_to_rule(row: CasbinRuleRow) -> List[str]:
|
||||
rule: List[str] = []
|
||||
if row.ptype is not None:
|
||||
rule.append(row.ptype)
|
||||
for value in (row.v0, row.v1, row.v2, row.v3, row.v4, row.v5):
|
||||
if value is not None and value != "":
|
||||
rule.append(value)
|
||||
return rule
|
||||
|
|
@ -58,7 +60,7 @@ def _require_admin(user_api_key_dict: UserAPIKeyAuth) -> None:
|
|||
)
|
||||
|
||||
|
||||
def _prisma() -> Any:
|
||||
def _prisma() -> PolicyAdminDB:
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
|
|
@ -66,7 +68,7 @@ def _prisma() -> Any:
|
|||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="auth_v2 policy administration requires a connected database",
|
||||
)
|
||||
return prisma_client
|
||||
return cast(PolicyAdminDB, prisma_client)
|
||||
|
||||
|
||||
async def _add_rule(rule: List[str]) -> None:
|
||||
|
|
|
|||
|
|
@ -1,5 +1,11 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, List, Optional
|
||||
from enum import Enum
|
||||
from typing import TYPE_CHECKING, List, Optional
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
|
|
@ -17,21 +23,20 @@ class Principal:
|
|||
groupings: List[List[str]]
|
||||
|
||||
|
||||
def _role_to_str(role: Any) -> Optional[str]:
|
||||
def _role_to_str(role: object) -> Optional[str]:
|
||||
if role is None:
|
||||
return None
|
||||
return getattr(role, "value", role)
|
||||
if isinstance(role, Enum):
|
||||
value = role.value
|
||||
return value if isinstance(value, str) else str(value)
|
||||
return str(role)
|
||||
|
||||
|
||||
def build_principal(identity: Any) -> Principal:
|
||||
"""Derive a :class:`Principal` from an authenticated identity object.
|
||||
|
||||
Duck-typed on ``user_id`` / ``team_id`` / ``token`` / ``user_role`` so it
|
||||
stays decoupled from the full ``UserAPIKeyAuth`` import.
|
||||
"""
|
||||
user_id = getattr(identity, "user_id", None)
|
||||
team_id = getattr(identity, "team_id", None)
|
||||
token = getattr(identity, "token", None)
|
||||
def build_principal(identity: UserAPIKeyAuth) -> Principal:
|
||||
"""Derive a :class:`Principal` from an authenticated identity."""
|
||||
user_id = identity.user_id
|
||||
team_id = identity.team_id
|
||||
token = identity.token
|
||||
|
||||
if user_id:
|
||||
subject = f"user:{user_id}"
|
||||
|
|
@ -43,7 +48,7 @@ def build_principal(identity: Any) -> Principal:
|
|||
domain = f"team:{team_id}" if team_id else "*"
|
||||
|
||||
groupings: List[List[str]] = []
|
||||
role = _role_to_str(getattr(identity, "user_role", None))
|
||||
role = _role_to_str(identity.user_role)
|
||||
if role:
|
||||
groupings.append([subject, f"role:{role}"])
|
||||
|
||||
|
|
|
|||
65
litellm/proxy/auth/v2/protocols.py
Normal file
65
litellm/proxy/auth/v2/protocols.py
Normal file
|
|
@ -0,0 +1,65 @@
|
|||
from typing import Dict, List, Optional, Protocol, runtime_checkable
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class SupportsEnforce(Protocol):
|
||||
"""Anything that can decide a casbin-style authorization check.
|
||||
|
||||
Lets the authorizer depend on the capability, not the concrete
|
||||
``CasbinEnforcer``, so test doubles and the loud-open sentinel type-check too.
|
||||
"""
|
||||
|
||||
def enforce(self, subject: str, domain: str, obj: str, action: str) -> bool: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class CasbinRuleRow(Protocol):
|
||||
"""A persisted casbin rule row (the LiteLLM_CasbinRule shape).
|
||||
|
||||
Typed structurally so the policy store doesn't depend on the generated Prisma
|
||||
model class while still being fully typed over ``ptype`` and ``v0``..``v5``.
|
||||
"""
|
||||
|
||||
ptype: Optional[str]
|
||||
v0: Optional[str]
|
||||
v1: Optional[str]
|
||||
v2: Optional[str]
|
||||
v3: Optional[str]
|
||||
v4: Optional[str]
|
||||
v5: Optional[str]
|
||||
|
||||
|
||||
class CasbinRuleTable(Protocol):
|
||||
"""The LiteLLM_CasbinRule Prisma accessor surface the policy store reads."""
|
||||
|
||||
async def find_many(self) -> List[CasbinRuleRow]: ...
|
||||
|
||||
|
||||
class _PolicyDBNamespace(Protocol):
|
||||
litellm_casbinrule: CasbinRuleTable
|
||||
|
||||
|
||||
class PolicyDB(Protocol):
|
||||
"""A Prisma client narrowed to the casbin-rule table the policy store needs."""
|
||||
|
||||
db: _PolicyDBNamespace
|
||||
|
||||
|
||||
class CasbinRuleWriteTable(Protocol):
|
||||
"""The casbin-rule table surface the policy-admin endpoints read and write."""
|
||||
|
||||
async def find_many(self) -> List[CasbinRuleRow]: ...
|
||||
|
||||
async def create(self, data: Dict[str, str]) -> object: ...
|
||||
|
||||
async def delete_many(self, where: Dict[str, str]) -> int: ...
|
||||
|
||||
|
||||
class _AdminDBNamespace(Protocol):
|
||||
litellm_casbinrule: CasbinRuleWriteTable
|
||||
|
||||
|
||||
class PolicyAdminDB(Protocol):
|
||||
"""A Prisma client narrowed to the casbin-rule writes the admin API performs."""
|
||||
|
||||
db: _AdminDBNamespace
|
||||
|
|
@ -1,7 +1,16 @@
|
|||
from typing import Any, Optional
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import LiteLLM_TeamTable, UserAPIKeyAuth
|
||||
|
||||
from ..authn.authenticators import AuthContext
|
||||
|
||||
|
||||
async def enforce_hierarchy_budgets(identity: Any, route: str, ctx: Any) -> None:
|
||||
async def enforce_hierarchy_budgets(
|
||||
identity: UserAPIKeyAuth, route: str, ctx: AuthContext
|
||||
) -> None:
|
||||
"""Enforce team, organization, and global budgets for an auth_v2 request.
|
||||
|
||||
These hierarchy caps live in v1's ``common_checks`` (not the pre-call hooks),
|
||||
|
|
@ -21,7 +30,7 @@ async def enforce_hierarchy_budgets(identity: Any, route: str, ctx: Any) -> None
|
|||
from litellm.proxy.proxy_server import litellm_proxy_admin_name
|
||||
|
||||
team_id = getattr(identity, "team_id", None)
|
||||
team_object: Optional[Any] = None
|
||||
team_object: Optional[LiteLLM_TeamTable] = None
|
||||
if team_id is not None:
|
||||
team_object = await get_team_object(
|
||||
team_id=team_id,
|
||||
|
|
|
|||
|
|
@ -1,8 +1,14 @@
|
|||
from typing import Any, Awaitable, Callable, Optional
|
||||
from __future__ import annotations
|
||||
|
||||
# Loaders are injected so the mapping is unit-testable without a DB.
|
||||
UserLoader = Callable[[str], Awaitable[Optional[Any]]]
|
||||
TeamLoader = Callable[[str], Awaitable[Optional[Any]]]
|
||||
from typing import TYPE_CHECKING, Awaitable, Callable, Dict, Optional
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
# Loaders are injected so the mapping is unit-testable without a DB. The loaded
|
||||
# row is only read by attribute, so ``object`` is enough and keeps this decoupled.
|
||||
UserLoader = Callable[[str], Awaitable[Optional[object]]]
|
||||
TeamLoader = Callable[[str], Awaitable[Optional[object]]]
|
||||
|
||||
# Source attr on the user/team row -> destination attr on the identity. The
|
||||
# destination user_*/team_* slots are distinct from key-level fields, so this
|
||||
|
|
@ -23,7 +29,9 @@ _TEAM_FIELD_MAP = {
|
|||
}
|
||||
|
||||
|
||||
def _copy_missing(identity: Any, source: Any, field_map: dict) -> None:
|
||||
def _copy_missing(
|
||||
identity: UserAPIKeyAuth, source: object, field_map: Dict[str, str]
|
||||
) -> None:
|
||||
for src_attr, dest_attr in field_map.items():
|
||||
value = getattr(source, src_attr, None)
|
||||
if value is not None and getattr(identity, dest_attr, None) is None:
|
||||
|
|
@ -31,11 +39,11 @@ def _copy_missing(identity: Any, source: Any, field_map: dict) -> None:
|
|||
|
||||
|
||||
async def enrich_identity(
|
||||
identity: Any,
|
||||
identity: UserAPIKeyAuth,
|
||||
*,
|
||||
load_user: Optional[UserLoader] = None,
|
||||
load_team: Optional[TeamLoader] = None,
|
||||
) -> Any:
|
||||
) -> UserAPIKeyAuth:
|
||||
"""Populate the identity's user/team budget+limit fields from the user/team rows.
|
||||
|
||||
Virtual keys arrive fully populated via ``get_key_object``; master/JWT/OAuth
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import Any, Optional
|
||||
|
||||
from litellm.proxy.auth.v2.principal import build_principal
|
||||
|
|
@ -12,9 +13,9 @@ class _Identity:
|
|||
user_role: Any = None
|
||||
|
||||
|
||||
class _Role:
|
||||
def __init__(self, value):
|
||||
self.value = value
|
||||
class _Role(str, Enum):
|
||||
# Mirrors LitellmUserRoles (a str Enum); build_principal reads .value off enums.
|
||||
PROXY_ADMIN = "proxy_admin"
|
||||
|
||||
|
||||
def test_user_id_becomes_subject_and_team_becomes_domain():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue