From 9c9c8932aa30eec62d4e8e0442049e12b7e085ca Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Mon, 8 Jun 2026 14:34:26 -0700 Subject: [PATCH] refactor(identity): consolidate principal classifier and service-account names --- litellm/identity/adapter.py | 31 +++++++------------ litellm/identity/principal.py | 19 +++++++++++- litellm/identity/resolver.py | 20 ++---------- litellm/identity/service_accounts.py | 15 +++++++++ litellm/identity/store.py | 13 ++------ tests/test_litellm/identity/test_principal.py | 17 ++++++++++ 6 files changed, 66 insertions(+), 49 deletions(-) create mode 100644 litellm/identity/service_accounts.py diff --git a/litellm/identity/adapter.py b/litellm/identity/adapter.py index 2b8ee777c8d..b27902eab49 100644 --- a/litellm/identity/adapter.py +++ b/litellm/identity/adapter.py @@ -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, diff --git a/litellm/identity/principal.py b/litellm/identity/principal.py index 28377f5f281..eba2ff70309 100644 --- a/litellm/identity/principal.py +++ b/litellm/identity/principal.py @@ -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" diff --git a/litellm/identity/resolver.py b/litellm/identity/resolver.py index 30281447ee8..43da2038e16 100644 --- a/litellm/identity/resolver.py +++ b/litellm/identity/resolver.py @@ -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() diff --git a/litellm/identity/service_accounts.py b/litellm/identity/service_accounts.py new file mode 100644 index 00000000000..9da2f08034c --- /dev/null +++ b/litellm/identity/service_accounts.py @@ -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, + } +) diff --git a/litellm/identity/store.py b/litellm/identity/store.py index b4652161002..6847277f0ad 100644 --- a/litellm/identity/store.py +++ b/litellm/identity/store.py @@ -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), diff --git a/tests/test_litellm/identity/test_principal.py b/tests/test_litellm/identity/test_principal.py index f2e663ec895..be206d9333f 100644 --- a/tests/test_litellm/identity/test_principal.py +++ b/tests/test_litellm/identity/test_principal.py @@ -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"