mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
feat(proxy): enhance user mapping to support email resolution and async handling
This commit is contained in:
parent
5833d3eadd
commit
ab2f2a2a4c
2 changed files with 118 additions and 9 deletions
|
|
@ -26,6 +26,10 @@ from litellm.proxy._types import (
|
||||||
UserAPIKeyAuth,
|
UserAPIKeyAuth,
|
||||||
)
|
)
|
||||||
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
|
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
|
||||||
|
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
|
||||||
|
from litellm.proxy.auth.auth_checks import get_user_object
|
||||||
|
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
|
||||||
|
from litellm.proxy.auth.auth_checks import get_user_object
|
||||||
|
|
||||||
# Cache special headers as a frozenset for O(1) lookup performance
|
# Cache special headers as a frozenset for O(1) lookup performance
|
||||||
_SPECIAL_HEADERS_CACHE = frozenset(
|
_SPECIAL_HEADERS_CACHE = frozenset(
|
||||||
|
|
@ -707,11 +711,18 @@ class LiteLLMProxyRequestSetup:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def add_internal_user_from_user_mapping(
|
async def add_internal_user_from_user_mapping(
|
||||||
general_settings: Optional[Dict],
|
general_settings: Optional[Dict],
|
||||||
user_api_key_dict: UserAPIKeyAuth,
|
user_api_key_dict: UserAPIKeyAuth,
|
||||||
headers: dict,
|
headers: dict,
|
||||||
) -> UserAPIKeyAuth:
|
) -> UserAPIKeyAuth:
|
||||||
|
"""If configured, map a header to an internal user.
|
||||||
|
|
||||||
|
If the header value looks like an email address, try to resolve an
|
||||||
|
internal user by email. If no user exists, create an internal-only
|
||||||
|
user for that email. In all cases set `user_api_key_dict.user_id` so
|
||||||
|
downstream spend attribution uses the resolved/created user.
|
||||||
|
"""
|
||||||
if general_settings is None:
|
if general_settings is None:
|
||||||
return user_api_key_dict
|
return user_api_key_dict
|
||||||
user_header_mapping = general_settings.get("user_header_mappings")
|
user_header_mapping = general_settings.get("user_header_mappings")
|
||||||
|
|
@ -722,12 +733,50 @@ class LiteLLMProxyRequestSetup:
|
||||||
)
|
)
|
||||||
if not header_name:
|
if not header_name:
|
||||||
return user_api_key_dict
|
return user_api_key_dict
|
||||||
|
|
||||||
header_value = LiteLLMProxyRequestSetup._get_case_insensitive_header(
|
header_value = LiteLLMProxyRequestSetup._get_case_insensitive_header(
|
||||||
headers, header_name
|
headers, header_name
|
||||||
)
|
)
|
||||||
if header_value:
|
if not header_value:
|
||||||
user_api_key_dict.user_id = header_value
|
|
||||||
return user_api_key_dict
|
return user_api_key_dict
|
||||||
|
|
||||||
|
# Quick email-ish heuristic
|
||||||
|
if isinstance(header_value, str) and re.match(r"^[^@\s]+@[^@\s]+\.[^@\s]+$", header_value):
|
||||||
|
try:
|
||||||
|
# Import at runtime to avoid circular imports
|
||||||
|
if prisma_client is not None:
|
||||||
|
try:
|
||||||
|
user_obj = await get_user_object(
|
||||||
|
user_id=str(header_value),
|
||||||
|
prisma_client=prisma_client,
|
||||||
|
user_api_key_cache=user_api_key_cache,
|
||||||
|
user_id_upsert=True,
|
||||||
|
user_email=str(header_value),
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
user_obj = None
|
||||||
|
|
||||||
|
if user_obj is not None:
|
||||||
|
# Ensure role is internal_user (best-effort)
|
||||||
|
try:
|
||||||
|
await prisma_client.db.litellm_usertable.update(
|
||||||
|
where={"user_id": user_obj.user_id},
|
||||||
|
data={"user_role": str(LitellmUserRoles.INTERNAL_USER)},
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
# ignore role-update failures
|
||||||
|
pass
|
||||||
|
|
||||||
|
user_api_key_dict.user_id = user_obj.user_id
|
||||||
|
user_api_key_dict.user_email = getattr(user_obj, "user_email", None)
|
||||||
|
user_api_key_dict.user_role = getattr(user_obj, "user_role", None) or LitellmUserRoles.INTERNAL_USER
|
||||||
|
return user_api_key_dict
|
||||||
|
except Exception:
|
||||||
|
# Fall back to using header value if DB unavailable
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Default: use the raw header value as the user identifier
|
||||||
|
user_api_key_dict.user_id = header_value
|
||||||
return user_api_key_dict
|
return user_api_key_dict
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|
@ -1335,7 +1384,7 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
||||||
data=data, headers=_headers, user_api_key_dict=user_api_key_dict
|
data=data, headers=_headers, user_api_key_dict=user_api_key_dict
|
||||||
)
|
)
|
||||||
|
|
||||||
user_api_key_dict = LiteLLMProxyRequestSetup.add_internal_user_from_user_mapping(
|
user_api_key_dict = await LiteLLMProxyRequestSetup.add_internal_user_from_user_mapping(
|
||||||
general_settings, user_api_key_dict, _headers
|
general_settings, user_api_key_dict, _headers
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -2517,7 +2517,8 @@ def test_get_internal_user_header_from_mapping_none_when_absent():
|
||||||
assert header_name is None
|
assert header_name is None
|
||||||
|
|
||||||
|
|
||||||
def test_add_internal_user_from_user_mapping_sets_user_id_when_header_present():
|
@pytest.mark.asyncio
|
||||||
|
async def test_add_internal_user_from_user_mapping_sets_user_id_when_header_present():
|
||||||
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
|
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
|
||||||
headers = {"X-OpenWebUI-User-Id": "internal-user-123"}
|
headers = {"X-OpenWebUI-User-Id": "internal-user-123"}
|
||||||
general_settings = {
|
general_settings = {
|
||||||
|
|
@ -2530,7 +2531,7 @@ def test_add_internal_user_from_user_mapping_sets_user_id_when_header_present():
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
|
|
||||||
result = LiteLLMProxyRequestSetup.add_internal_user_from_user_mapping(
|
result = await LiteLLMProxyRequestSetup.add_internal_user_from_user_mapping(
|
||||||
general_settings, user_api_key_dict, headers
|
general_settings, user_api_key_dict, headers
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -2538,10 +2539,11 @@ def test_add_internal_user_from_user_mapping_sets_user_id_when_header_present():
|
||||||
assert user_api_key_dict.user_id == "internal-user-123"
|
assert user_api_key_dict.user_id == "internal-user-123"
|
||||||
|
|
||||||
|
|
||||||
def test_add_internal_user_from_user_mapping_no_header_or_mapping_returns_unchanged():
|
@pytest.mark.asyncio
|
||||||
|
async def test_add_internal_user_from_user_mapping_no_header_or_mapping_returns_unchanged():
|
||||||
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
|
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
|
||||||
|
|
||||||
result = LiteLLMProxyRequestSetup.add_internal_user_from_user_mapping(
|
result = await LiteLLMProxyRequestSetup.add_internal_user_from_user_mapping(
|
||||||
None, user_api_key_dict, {"X-OpenWebUI-User-Id": "abc"}
|
None, user_api_key_dict, {"X-OpenWebUI-User-Id": "abc"}
|
||||||
)
|
)
|
||||||
assert result is user_api_key_dict
|
assert result is user_api_key_dict
|
||||||
|
|
@ -2552,13 +2554,71 @@ def test_add_internal_user_from_user_mapping_no_header_or_mapping_returns_unchan
|
||||||
{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"}
|
{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"}
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
result = LiteLLMProxyRequestSetup.add_internal_user_from_user_mapping(
|
result = await LiteLLMProxyRequestSetup.add_internal_user_from_user_mapping(
|
||||||
general_settings, user_api_key_dict, {"Other": "value"}
|
general_settings, user_api_key_dict, {"Other": "value"}
|
||||||
)
|
)
|
||||||
assert result is user_api_key_dict
|
assert result is user_api_key_dict
|
||||||
assert user_api_key_dict.user_id is None
|
assert user_api_key_dict.user_id is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_add_internal_user_from_user_mapping_resolves_email_header_to_internal_user(
|
||||||
|
monkeypatch,
|
||||||
|
):
|
||||||
|
from litellm import proxy as proxy_pkg
|
||||||
|
|
||||||
|
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
|
||||||
|
headers = {"X-OpenWebUI-User-Email": "internal@example.com"}
|
||||||
|
general_settings = {
|
||||||
|
"user_header_mappings": [
|
||||||
|
{
|
||||||
|
"header_name": "X-OpenWebUI-User-Email",
|
||||||
|
"litellm_user_role": "internal_user",
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
fake_prisma_client = object()
|
||||||
|
fake_cache = object()
|
||||||
|
fake_user = MagicMock()
|
||||||
|
fake_user.user_id = "internal-user-db-id"
|
||||||
|
fake_user.user_email = "internal@example.com"
|
||||||
|
fake_user.user_role = "internal_user"
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
proxy_pkg.litellm_pre_call_utils,
|
||||||
|
"prisma_client",
|
||||||
|
fake_prisma_client,
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
proxy_pkg.litellm_pre_call_utils,
|
||||||
|
"user_api_key_cache",
|
||||||
|
fake_cache,
|
||||||
|
)
|
||||||
|
get_user_object_mock = AsyncMock(return_value=fake_user)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
proxy_pkg.litellm_pre_call_utils,
|
||||||
|
"get_user_object",
|
||||||
|
get_user_object_mock,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await LiteLLMProxyRequestSetup.add_internal_user_from_user_mapping(
|
||||||
|
general_settings, user_api_key_dict, headers
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result is user_api_key_dict
|
||||||
|
assert user_api_key_dict.user_id == "internal-user-db-id"
|
||||||
|
assert user_api_key_dict.user_email == "internal@example.com"
|
||||||
|
assert user_api_key_dict.user_role == "internal_user"
|
||||||
|
get_user_object_mock.assert_awaited_once_with(
|
||||||
|
user_id="internal@example.com",
|
||||||
|
prisma_client=fake_prisma_client,
|
||||||
|
user_api_key_cache=fake_cache,
|
||||||
|
user_id_upsert=True,
|
||||||
|
user_email="internal@example.com",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_get_sanitized_user_information_from_key_includes_guardrails_metadata():
|
def test_get_sanitized_user_information_from_key_includes_guardrails_metadata():
|
||||||
"""
|
"""
|
||||||
Test that get_sanitized_user_information_from_key includes guardrails field from key metadata in the returned payload
|
Test that get_sanitized_user_information_from_key includes guardrails field from key metadata in the returned payload
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue