mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-21 00:24:55 +00:00
fix mypy error
This commit is contained in:
@@ -13,7 +13,7 @@ However, this is under the enterprise license
|
|||||||
import asyncio
|
import asyncio
|
||||||
import os
|
import os
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
from typing import List, Literal, Optional, Union
|
from typing import List, Optional, Union
|
||||||
|
|
||||||
from litellm._logging import verbose_logger
|
from litellm._logging import verbose_logger
|
||||||
from litellm.caching import DualCache
|
from litellm.caching import DualCache
|
||||||
@@ -115,6 +115,7 @@ class PagerDutyAlerting(SlackAlerting):
|
|||||||
user_api_key_budget_reset_at=_meta.get("user_api_key_budget_reset_at"),
|
user_api_key_budget_reset_at=_meta.get("user_api_key_budget_reset_at"),
|
||||||
user_api_key_org_id=_meta.get("user_api_key_org_id"),
|
user_api_key_org_id=_meta.get("user_api_key_org_id"),
|
||||||
user_api_key_team_id=_meta.get("user_api_key_team_id"),
|
user_api_key_team_id=_meta.get("user_api_key_team_id"),
|
||||||
|
user_api_key_project_id=_meta.get("user_api_key_project_id"),
|
||||||
user_api_key_user_id=_meta.get("user_api_key_user_id"),
|
user_api_key_user_id=_meta.get("user_api_key_user_id"),
|
||||||
user_api_key_team_alias=_meta.get("user_api_key_team_alias"),
|
user_api_key_team_alias=_meta.get("user_api_key_team_alias"),
|
||||||
user_api_key_end_user_id=_meta.get("user_api_key_end_user_id"),
|
user_api_key_end_user_id=_meta.get("user_api_key_end_user_id"),
|
||||||
@@ -195,6 +196,7 @@ class PagerDutyAlerting(SlackAlerting):
|
|||||||
),
|
),
|
||||||
user_api_key_org_id=user_api_key_dict.org_id,
|
user_api_key_org_id=user_api_key_dict.org_id,
|
||||||
user_api_key_team_id=user_api_key_dict.team_id,
|
user_api_key_team_id=user_api_key_dict.team_id,
|
||||||
|
user_api_key_project_id=user_api_key_dict.project_id,
|
||||||
user_api_key_user_id=user_api_key_dict.user_id,
|
user_api_key_user_id=user_api_key_dict.user_id,
|
||||||
user_api_key_team_alias=user_api_key_dict.team_alias,
|
user_api_key_team_alias=user_api_key_dict.team_alias,
|
||||||
user_api_key_end_user_id=user_api_key_dict.end_user_id,
|
user_api_key_end_user_id=user_api_key_dict.end_user_id,
|
||||||
|
|||||||
@@ -4506,6 +4506,7 @@ class StandardLoggingPayloadSetup:
|
|||||||
user_api_key_budget_reset_at=None,
|
user_api_key_budget_reset_at=None,
|
||||||
user_api_key_team_id=None,
|
user_api_key_team_id=None,
|
||||||
user_api_key_org_id=None,
|
user_api_key_org_id=None,
|
||||||
|
user_api_key_project_id=None,
|
||||||
user_api_key_user_id=None,
|
user_api_key_user_id=None,
|
||||||
user_api_key_team_alias=None,
|
user_api_key_team_alias=None,
|
||||||
user_api_key_user_email=None,
|
user_api_key_user_email=None,
|
||||||
@@ -5272,6 +5273,7 @@ def get_standard_logging_metadata(
|
|||||||
user_api_key_budget_reset_at=None,
|
user_api_key_budget_reset_at=None,
|
||||||
user_api_key_team_id=None,
|
user_api_key_team_id=None,
|
||||||
user_api_key_org_id=None,
|
user_api_key_org_id=None,
|
||||||
|
user_api_key_project_id=None,
|
||||||
user_api_key_user_id=None,
|
user_api_key_user_id=None,
|
||||||
user_api_key_user_email=None,
|
user_api_key_user_email=None,
|
||||||
user_api_key_team_alias=None,
|
user_api_key_team_alias=None,
|
||||||
|
|||||||
@@ -60,6 +60,7 @@ class _ProxyDBLogger(CustomLogger):
|
|||||||
user_api_key_user_id=user_api_key_dict.user_id,
|
user_api_key_user_id=user_api_key_dict.user_id,
|
||||||
user_api_key_team_id=user_api_key_dict.team_id,
|
user_api_key_team_id=user_api_key_dict.team_id,
|
||||||
user_api_key_org_id=user_api_key_dict.org_id,
|
user_api_key_org_id=user_api_key_dict.org_id,
|
||||||
|
user_api_key_project_id=user_api_key_dict.project_id,
|
||||||
user_api_key_team_alias=user_api_key_dict.team_alias,
|
user_api_key_team_alias=user_api_key_dict.team_alias,
|
||||||
user_api_key_end_user_id=user_api_key_dict.end_user_id,
|
user_api_key_end_user_id=user_api_key_dict.end_user_id,
|
||||||
user_api_key_request_route=user_api_key_dict.request_route,
|
user_api_key_request_route=user_api_key_dict.request_route,
|
||||||
@@ -68,11 +69,11 @@ class _ProxyDBLogger(CustomLogger):
|
|||||||
)
|
)
|
||||||
_metadata["user_api_key"] = user_api_key_dict.api_key
|
_metadata["user_api_key"] = user_api_key_dict.api_key
|
||||||
_metadata["status"] = "failure"
|
_metadata["status"] = "failure"
|
||||||
_metadata["error_information"] = (
|
_metadata[
|
||||||
StandardLoggingPayloadSetup.get_error_information(
|
"error_information"
|
||||||
original_exception=original_exception,
|
] = StandardLoggingPayloadSetup.get_error_information(
|
||||||
traceback_str=traceback_str,
|
original_exception=original_exception,
|
||||||
)
|
traceback_str=traceback_str,
|
||||||
)
|
)
|
||||||
|
|
||||||
existing_metadata: dict = request_data.get("metadata", None) or {}
|
existing_metadata: dict = request_data.get("metadata", None) or {}
|
||||||
@@ -89,15 +90,21 @@ class _ProxyDBLogger(CustomLogger):
|
|||||||
existing_metadata["tags"] = existing_litellm_metadata.get("tags")
|
existing_metadata["tags"] = existing_litellm_metadata.get("tags")
|
||||||
|
|
||||||
request_data["litellm_params"]["proxy_server_request"] = (
|
request_data["litellm_params"]["proxy_server_request"] = (
|
||||||
request_data.get("proxy_server_request") or existing_litellm_params.get("proxy_server_request") or {}
|
request_data.get("proxy_server_request")
|
||||||
|
or existing_litellm_params.get("proxy_server_request")
|
||||||
|
or {}
|
||||||
)
|
)
|
||||||
request_data["litellm_params"]["metadata"] = existing_metadata
|
request_data["litellm_params"]["metadata"] = existing_metadata
|
||||||
|
|
||||||
# Preserve model name and custom_llm_provider
|
# Preserve model name and custom_llm_provider
|
||||||
if "model" not in request_data:
|
if "model" not in request_data:
|
||||||
request_data["model"] = existing_litellm_params.get("model") or request_data.get("model", "")
|
request_data["model"] = existing_litellm_params.get(
|
||||||
|
"model"
|
||||||
|
) or request_data.get("model", "")
|
||||||
if "custom_llm_provider" not in request_data:
|
if "custom_llm_provider" not in request_data:
|
||||||
request_data["custom_llm_provider"] = existing_litellm_params.get("custom_llm_provider") or request_data.get("custom_llm_provider", "")
|
request_data["custom_llm_provider"] = existing_litellm_params.get(
|
||||||
|
"custom_llm_provider"
|
||||||
|
) or request_data.get("custom_llm_provider", "")
|
||||||
|
|
||||||
await proxy_logging_obj.db_spend_update_writer.update_database(
|
await proxy_logging_obj.db_spend_update_writer.update_database(
|
||||||
token=user_api_key_dict.api_key,
|
token=user_api_key_dict.api_key,
|
||||||
@@ -211,7 +218,8 @@ class _ProxyDBLogger(CustomLogger):
|
|||||||
)
|
)
|
||||||
return
|
return
|
||||||
if kwargs.get("stream") is not True or (
|
if kwargs.get("stream") is not True or (
|
||||||
kwargs.get("stream") is True and "complete_streaming_response" in kwargs
|
kwargs.get("stream") is True
|
||||||
|
and "complete_streaming_response" in kwargs
|
||||||
):
|
):
|
||||||
if sl_object is not None:
|
if sl_object is not None:
|
||||||
cost_tracking_failure_debug_info: Union[dict, str] = (
|
cost_tracking_failure_debug_info: Union[dict, str] = (
|
||||||
|
|||||||
@@ -498,6 +498,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
|||||||
user_api_key_user_id=user_api_key_dict.user_id,
|
user_api_key_user_id=user_api_key_dict.user_id,
|
||||||
user_api_key_team_id=user_api_key_dict.team_id,
|
user_api_key_team_id=user_api_key_dict.team_id,
|
||||||
user_api_key_org_id=user_api_key_dict.org_id,
|
user_api_key_org_id=user_api_key_dict.org_id,
|
||||||
|
user_api_key_project_id=user_api_key_dict.project_id,
|
||||||
user_api_key_team_alias=user_api_key_dict.team_alias,
|
user_api_key_team_alias=user_api_key_dict.team_alias,
|
||||||
user_api_key_end_user_id=user_api_key_dict.end_user_id,
|
user_api_key_end_user_id=user_api_key_dict.end_user_id,
|
||||||
user_api_key_request_route=user_api_key_dict.request_route,
|
user_api_key_request_route=user_api_key_dict.request_route,
|
||||||
@@ -1099,7 +1100,9 @@ def create_pass_through_route(
|
|||||||
fastapi_response: Response,
|
fastapi_response: Response,
|
||||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||||
subpath: str = "", # captures sub-paths when include_subpath=True
|
subpath: str = "", # captures sub-paths when include_subpath=True
|
||||||
custom_body: Optional[dict] = None, # accepted for signature compatibility with URL-based path; not forwarded because chat_completion_pass_through_endpoint does not support it
|
custom_body: Optional[
|
||||||
|
dict
|
||||||
|
] = None, # accepted for signature compatibility with URL-based path; not forwarded because chat_completion_pass_through_endpoint does not support it
|
||||||
):
|
):
|
||||||
return await chat_completion_pass_through_endpoint(
|
return await chat_completion_pass_through_endpoint(
|
||||||
fastapi_response=fastapi_response,
|
fastapi_response=fastapi_response,
|
||||||
@@ -2029,8 +2032,8 @@ class InitPassThroughEndpointHelpers:
|
|||||||
parts = key.split(":", 2) # Split into [endpoint_id, type, path]
|
parts = key.split(":", 2) # Split into [endpoint_id, type, path]
|
||||||
if len(parts) == 3:
|
if len(parts) == 3:
|
||||||
route_type = parts[1]
|
route_type = parts[1]
|
||||||
registered_path = InitPassThroughEndpointHelpers._build_full_path_with_root(
|
registered_path = (
|
||||||
parts[2]
|
InitPassThroughEndpointHelpers._build_full_path_with_root(parts[2])
|
||||||
)
|
)
|
||||||
if route_type == "exact" and route == registered_path:
|
if route_type == "exact" and route == registered_path:
|
||||||
return True
|
return True
|
||||||
@@ -2049,8 +2052,8 @@ class InitPassThroughEndpointHelpers:
|
|||||||
parts = key.split(":", 2) # Split into [endpoint_id, type, path]
|
parts = key.split(":", 2) # Split into [endpoint_id, type, path]
|
||||||
if len(parts) == 3:
|
if len(parts) == 3:
|
||||||
route_type = parts[1]
|
route_type = parts[1]
|
||||||
registered_path = InitPassThroughEndpointHelpers._build_full_path_with_root(
|
registered_path = (
|
||||||
parts[2]
|
InitPassThroughEndpointHelpers._build_full_path_with_root(parts[2])
|
||||||
)
|
)
|
||||||
|
|
||||||
if route_type == "exact" and route == registered_path:
|
if route_type == "exact" and route == registered_path:
|
||||||
@@ -2371,9 +2374,7 @@ async def get_pass_through_endpoints(
|
|||||||
|
|
||||||
# Merge: config endpoints not in DB + all DB endpoints (DB overrides config for same path)
|
# Merge: config endpoints not in DB + all DB endpoints (DB overrides config for same path)
|
||||||
db_paths = {ep.path for ep in db_endpoints}
|
db_paths = {ep.path for ep in db_endpoints}
|
||||||
config_only_endpoints = [
|
config_only_endpoints = [ep for ep in config_endpoints if ep.path not in db_paths]
|
||||||
ep for ep in config_endpoints if ep.path not in db_paths
|
|
||||||
]
|
|
||||||
if endpoint_id is not None:
|
if endpoint_id is not None:
|
||||||
# When filtering by endpoint_id, only return if found in DB (config endpoints use generated IDs)
|
# When filtering by endpoint_id, only return if found in DB (config endpoints use generated IDs)
|
||||||
pass_through_endpoints = db_endpoints
|
pass_through_endpoints = db_endpoints
|
||||||
|
|||||||
Reference in New Issue
Block a user