fix(auth_utils.py): ensure extracted end user id is always a str

prevents db cost tracking errors
This commit is contained in:
Krrish Dholakia 2025-01-22 17:59:35 -08:00
parent 59d23d1d84
commit c89995ff7a
3 changed files with 16 additions and 11 deletions

View file

@ -498,13 +498,13 @@ def _has_user_setup_sso():
def get_end_user_id_from_request_body(request_body: dict) -> Optional[str]:
# openai - check 'user'
if "user" in request_body:
return request_body["user"]
if "user" in request_body and request_body["user"] is not None:
return str(request_body["user"])
# anthropic - check 'litellm_metadata'
end_user_id = request_body.get("litellm_metadata", {}).get("user", None)
if end_user_id:
return end_user_id
return str(end_user_id)
metadata = request_body.get("metadata")
if metadata and "user_id" in metadata:
return metadata["user_id"]
if metadata and "user_id" in metadata and metadata["user_id"] is not None:
return str(metadata["user_id"])
return None

View file

@ -1049,7 +1049,6 @@ async def update_database( # noqa: PLR0915
response_obj=completion_response,
start_time=start_time,
end_time=end_time,
end_user_id=end_user_id,
)
payload["spend"] = response_cost

View file

@ -29,9 +29,7 @@ def _is_master_key(api_key: str, _master_key: Optional[str]) -> bool:
return False
def get_logging_payload(
kwargs, response_obj, start_time, end_time, end_user_id: Optional[str]
) -> SpendLogsPayload:
def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogsPayload:
from litellm.proxy.proxy_server import general_settings, master_key
@ -58,9 +56,17 @@ def get_logging_payload(
usage = dict(usage)
id = cast(dict, response_obj).get("id") or kwargs.get("litellm_call_id")
api_key = metadata.get("user_api_key", "")
standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get(
"standard_logging_object", None
standard_logging_payload = cast(
Optional[StandardLoggingPayload], kwargs.get("standard_logging_object", None)
)
if standard_logging_payload is not None:
end_user_id = standard_logging_payload["metadata"].get(
"user_api_key_end_user_id"
)
else:
end_user_id = None
if api_key is not None and isinstance(api_key, str):
if api_key.startswith("sk-"):
# hash the api_key