mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
refactor(identity): consolidate principal classifier and service-account names
This commit is contained in:
parent
d450935a9d
commit
9c9c8932aa
6 changed files with 66 additions and 49 deletions
|
|
@ -13,11 +13,7 @@ Invariants:
|
|||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from litellm.constants import (
|
||||
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
|
||||
LITTELM_CLI_SERVICE_ACCOUNT_NAME,
|
||||
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
|
||||
)
|
||||
from litellm.constants import LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME
|
||||
from litellm.identity.context import AuditInfo, ClientInfo, IdentityContext, RequestIds
|
||||
from litellm.identity.principal import (
|
||||
AnonymousPrincipal,
|
||||
|
|
@ -25,27 +21,22 @@ from litellm.identity.principal import (
|
|||
JWTPrincipal,
|
||||
Principal,
|
||||
ServiceAccountPrincipal,
|
||||
classify_principal_kind,
|
||||
)
|
||||
from litellm.identity.service_accounts import SERVICE_ACCOUNT_NAMES
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
_SERVICE_ACCOUNT_NAMES = frozenset(
|
||||
{
|
||||
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
|
||||
LITTELM_CLI_SERVICE_ACCOUNT_NAME,
|
||||
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _principal_from_uak(uak: "UserAPIKeyAuth") -> Principal:
|
||||
if uak.api_key in _SERVICE_ACCOUNT_NAMES or uak.key_alias in _SERVICE_ACCOUNT_NAMES:
|
||||
return ServiceAccountPrincipal(
|
||||
name=uak.api_key if uak.api_key in _SERVICE_ACCOUNT_NAMES else uak.key_alias # type: ignore[arg-type]
|
||||
)
|
||||
kind = classify_principal_kind(uak)
|
||||
|
||||
if uak.jwt_claims:
|
||||
if kind == "service_account":
|
||||
name = uak.api_key if uak.api_key in SERVICE_ACCOUNT_NAMES else uak.key_alias
|
||||
return ServiceAccountPrincipal(name=name) # type: ignore[arg-type]
|
||||
|
||||
if kind == "jwt":
|
||||
claims = uak.jwt_claims
|
||||
scope_claim = claims.get("scope") or claims.get("scp") or ""
|
||||
if isinstance(scope_claim, list):
|
||||
|
|
@ -66,9 +57,9 @@ def _principal_from_uak(uak: "UserAPIKeyAuth") -> Principal:
|
|||
mapped_org_id=uak.org_id,
|
||||
)
|
||||
|
||||
if uak.token:
|
||||
if kind == "api_key":
|
||||
return ApiKeyPrincipal(
|
||||
token_hash=uak.token,
|
||||
token_hash=uak.token, # type: ignore[arg-type]
|
||||
key_alias=uak.key_alias,
|
||||
user_id=uak.user_id,
|
||||
team_id=uak.team_id,
|
||||
|
|
|
|||
|
|
@ -10,7 +10,12 @@ for ``match``-style dispatch.
|
|||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict, Literal, Optional, Tuple, Union
|
||||
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple, Union
|
||||
|
||||
from litellm.identity.service_accounts import SERVICE_ACCOUNT_NAMES
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
|
|
@ -64,3 +69,15 @@ Principal = Union[
|
|||
ServiceAccountPrincipal,
|
||||
AnonymousPrincipal,
|
||||
]
|
||||
|
||||
PrincipalKind = Literal["service_account", "jwt", "api_key", "anonymous"]
|
||||
|
||||
|
||||
def classify_principal_kind(uak: "UserAPIKeyAuth") -> PrincipalKind:
|
||||
if uak.api_key in SERVICE_ACCOUNT_NAMES or uak.key_alias in SERVICE_ACCOUNT_NAMES:
|
||||
return "service_account"
|
||||
if uak.jwt_claims:
|
||||
return "jwt"
|
||||
if uak.token:
|
||||
return "api_key"
|
||||
return "anonymous"
|
||||
|
|
|
|||
|
|
@ -16,11 +16,6 @@ from __future__ import annotations
|
|||
|
||||
from typing import TYPE_CHECKING, Any, Dict, Optional
|
||||
|
||||
from litellm.constants import (
|
||||
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
|
||||
LITTELM_CLI_SERVICE_ACCOUNT_NAME,
|
||||
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
|
||||
)
|
||||
from litellm.identity.context import AuditInfo, ClientInfo, IdentityContext
|
||||
from litellm.identity.extractors.api_key import (
|
||||
extract_api_key_principal,
|
||||
|
|
@ -35,6 +30,7 @@ from litellm.identity.principal import (
|
|||
Principal,
|
||||
ServiceAccountPrincipal,
|
||||
)
|
||||
from litellm.identity.service_accounts import SERVICE_ACCOUNT_NAMES
|
||||
from litellm.integrations.otel.model.spans import SpanRole
|
||||
from litellm.integrations.otel.runtime import traced
|
||||
|
||||
|
|
@ -45,17 +41,9 @@ if TYPE_CHECKING:
|
|||
|
||||
from litellm.identity.cache import IdentityCache
|
||||
|
||||
_SERVICE_ACCOUNT_API_KEYS = frozenset(
|
||||
{
|
||||
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
|
||||
LITTELM_CLI_SERVICE_ACCOUNT_NAME,
|
||||
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _principal_from_raw_key(api_key: Optional[str]) -> Principal:
|
||||
if api_key and api_key in _SERVICE_ACCOUNT_API_KEYS:
|
||||
if api_key and api_key in SERVICE_ACCOUNT_NAMES:
|
||||
return ServiceAccountPrincipal(name=api_key)
|
||||
|
||||
jwt_principal = extract_jwt_principal(api_key)
|
||||
|
|
@ -96,9 +84,7 @@ async def resolve_identity(
|
|||
audit = AuditInfo(changed_by=extract_audit_changed_by(headers))
|
||||
client: ClientInfo
|
||||
if request is not None:
|
||||
client = extract_client_info(
|
||||
request=request, general_settings=general_settings
|
||||
)
|
||||
client = extract_client_info(request=request, general_settings=general_settings)
|
||||
else:
|
||||
client = ClientInfo()
|
||||
|
||||
|
|
|
|||
15
litellm/identity/service_accounts.py
Normal file
15
litellm/identity/service_accounts.py
Normal file
|
|
@ -0,0 +1,15 @@
|
|||
"""Sentinel names for litellm's internal service-account principals."""
|
||||
|
||||
from litellm.constants import (
|
||||
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
|
||||
LITTELM_CLI_SERVICE_ACCOUNT_NAME,
|
||||
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
|
||||
)
|
||||
|
||||
SERVICE_ACCOUNT_NAMES = frozenset(
|
||||
{
|
||||
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
|
||||
LITTELM_CLI_SERVICE_ACCOUNT_NAME,
|
||||
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
|
||||
}
|
||||
)
|
||||
|
|
@ -17,6 +17,7 @@ from typing import TYPE_CHECKING, Optional
|
|||
|
||||
from fastapi import status
|
||||
|
||||
from litellm.identity.principal import classify_principal_kind
|
||||
from litellm.integrations.otel.model.spans import SpanRole
|
||||
from litellm.integrations.otel.runtime import traced
|
||||
from litellm.proxy._types import ProxyErrorTypes, ProxyException
|
||||
|
|
@ -96,21 +97,11 @@ async def _fetch_from_db(
|
|||
return uak
|
||||
|
||||
|
||||
def _principal_kind(uak: "UserAPIKeyAuth") -> str:
|
||||
if uak.api_key and uak.api_key == uak.team_id:
|
||||
return "service_account"
|
||||
if uak.jwt_claims:
|
||||
return "jwt"
|
||||
if uak.token:
|
||||
return "api_key"
|
||||
return "anonymous"
|
||||
|
||||
|
||||
@traced(
|
||||
"identity.load",
|
||||
role=SpanRole.SERVICE,
|
||||
attrs=lambda result: {
|
||||
"identity.principal.kind": _principal_kind(result),
|
||||
"identity.principal.kind": classify_principal_kind(result),
|
||||
"identity.principal.has_team": bool(result.team_id),
|
||||
"identity.principal.has_org": bool(result.org_id),
|
||||
"identity.principal.has_project": bool(result.project_id),
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from litellm.identity.principal import (
|
|||
JWTPrincipal,
|
||||
SSOPrincipal,
|
||||
ServiceAccountPrincipal,
|
||||
classify_principal_kind,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -63,3 +64,19 @@ def test_jwt_principal_hash_ignores_claims():
|
|||
def test_kind_is_not_constructor_arg():
|
||||
with pytest.raises(TypeError):
|
||||
ApiKeyPrincipal(kind="api_key", token_hash="x") # type: ignore[call-arg]
|
||||
|
||||
|
||||
def test_classify_principal_kind_uses_name_set_not_team_id_equality():
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
collision = UserAPIKeyAuth(api_key="myteam", team_id="myteam")
|
||||
assert collision.api_key == collision.team_id
|
||||
assert classify_principal_kind(collision) == "api_key"
|
||||
|
||||
service_account = UserAPIKeyAuth.get_litellm_cli_user_api_key_auth()
|
||||
assert classify_principal_kind(service_account) == "service_account"
|
||||
|
||||
jwt_principal = UserAPIKeyAuth(jwt_claims={"sub": "u"})
|
||||
assert classify_principal_kind(jwt_principal) == "jwt"
|
||||
|
||||
assert classify_principal_kind(UserAPIKeyAuth()) == "anonymous"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue