fix mypy error

This commit is contained in:
Harshit Jain
2026-02-19 14:04:00 +05:30
parent 31752c7b78
commit bdf01fa283
4 changed files with 8797 additions and 8784 deletions
@@ -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