mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(auth_utils.py): ensure extracted end user id is always a str
prevents db cost tracking errors
This commit is contained in:
parent
59d23d1d84
commit
c89995ff7a
3 changed files with 16 additions and 11 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue