From 1a83935aa437cca63a10af073c15f5b9dbda18d9 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Mon, 22 Jul 2024 21:31:21 -0700 Subject: [PATCH] fix(proxy/utils.py): add stronger typing for litellm params in failure call logging --- litellm/proxy/utils.py | 14 ++++++++++---- litellm/tests/test_custom_callback_input.py | 1 + litellm/types/utils.py | 19 +++++++++++++++++++ 3 files changed, 30 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 5846dd25a86..cebd79aa29d 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -25,7 +25,7 @@ from typing_extensions import overload import litellm import litellm.litellm_core_utils import litellm.litellm_core_utils.litellm_logging -from litellm import EmbeddingResponse, ImageResponse, ModelResponse +from litellm import EmbeddingResponse, ImageResponse, ModelResponse, get_litellm_params from litellm._logging import verbose_proxy_logger from litellm._service_logger import ServiceLogging, ServiceTypes from litellm.caching import DualCache, RedisCache @@ -50,7 +50,7 @@ from litellm.proxy.hooks.max_budget_limiter import _PROXY_MaxBudgetLimiter from litellm.proxy.hooks.parallel_request_limiter import ( _PROXY_MaxParallelRequestsHandler, ) -from litellm.types.utils import CallTypes +from litellm.types.utils import CallTypes, LoggedLiteLLMParams if TYPE_CHECKING: from opentelemetry.trace import Span as _Span @@ -602,14 +602,20 @@ class ProxyLogging: if litellm_logging_obj is not None: ## UPDATE LOGGING INPUT _optional_params = {} + _litellm_params = {} + + litellm_param_keys = LoggedLiteLLMParams.__annotations__.keys() for k, v in request_data.items(): - if k != "model" and k != "user" and k != "litellm_params": + if k in litellm_param_keys: + _litellm_params[k] = v + elif k != "model" and k != "user": _optional_params[k] = v + litellm_logging_obj.update_environment_variables( model=request_data.get("model", ""), user=request_data.get("user", ""), optional_params=_optional_params, - litellm_params=request_data.get("litellm_params", {}), + litellm_params=_litellm_params, ) input: Union[list, str, dict] = "" diff --git a/litellm/tests/test_custom_callback_input.py b/litellm/tests/test_custom_callback_input.py index eae0412d391..9c18899a577 100644 --- a/litellm/tests/test_custom_callback_input.py +++ b/litellm/tests/test_custom_callback_input.py @@ -234,6 +234,7 @@ class CompletionCustomHandler( ) assert isinstance(kwargs["optional_params"], dict) assert isinstance(kwargs["litellm_params"], dict) + assert isinstance(kwargs["litellm_params"]["metadata"], Optional[dict]) assert isinstance(kwargs["start_time"], (datetime, type(None))) assert isinstance(kwargs["stream"], bool) assert isinstance(kwargs["user"], (str, type(None))) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 6581fea5f8f..88bfa19e90c 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -1029,3 +1029,22 @@ class GenericImageParsingChunk(TypedDict): class ResponseFormatChunk(TypedDict, total=False): type: Required[Literal["json_object", "text"]] response_schema: dict + + +class LoggedLiteLLMParams(TypedDict, total=False): + force_timeout: Optional[float] + custom_llm_provider: Optional[str] + api_base: Optional[str] + litellm_call_id: Optional[str] + model_alias_map: Optional[dict] + metadata: Optional[dict] + model_info: Optional[dict] + proxy_server_request: Optional[dict] + acompletion: Optional[bool] + preset_cache_key: Optional[str] + no_log: Optional[bool] + input_cost_per_second: Optional[float] + input_cost_per_token: Optional[float] + output_cost_per_token: Optional[float] + output_cost_per_second: Optional[float] + cooldown_time: Optional[float]