refactor(identity): consolidate principal classifier and service-account names

This commit is contained in:
Yassin Kortam 2026-06-08 14:34:26 -07:00
parent d450935a9d
commit 9c9c8932aa
6 changed files with 66 additions and 49 deletions

View file

@ -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,

View file

@ -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"

View file

@ -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()

View 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,
}
)

View file

@ -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),

View file

@ -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"