From 8d2dba8cac11573c123378e993c71e78935f5830 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 26 Nov 2025 11:34:22 -0800 Subject: [PATCH] fix code qa checks --- .../llms/vertex_ai/videos/transformation.py | 11 +- .../vertex_passthrough_logging_handler.py | 165 ++++++++++-------- 2 files changed, 97 insertions(+), 79 deletions(-) diff --git a/litellm/llms/vertex_ai/videos/transformation.py b/litellm/llms/vertex_ai/videos/transformation.py index aad29d25122..f6678510493 100644 --- a/litellm/llms/vertex_ai/videos/transformation.py +++ b/litellm/llms/vertex_ai/videos/transformation.py @@ -163,7 +163,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): headers: dict, model: str, api_key: Optional[str] = None, - litellm_params: Optional[GenericLiteLLMParams] = None, + litellm_params: Optional[Union[GenericLiteLLMParams, dict]] = None, ) -> dict: """ Validate environment and return headers for Vertex AI OCR. @@ -172,9 +172,12 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): """ # Extract Vertex AI parameters using safe helpers from VertexBase # Use safe_get_* methods that don't mutate litellm_params dict - litellm_params_dict: Dict[str, Any] = ( - litellm_params.model_dump() if litellm_params else {} - ) + if litellm_params is None: + litellm_params_dict: Dict[str, Any] = {} + elif isinstance(litellm_params, dict): + litellm_params_dict = litellm_params + else: + litellm_params_dict = litellm_params.model_dump() vertex_project = VertexBase.safe_get_vertex_ai_project( litellm_params=litellm_params_dict diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py index 407cbd2f3d0..b34a6f455c3 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -89,82 +89,12 @@ class VertexPassthroughLoggingHandler: } elif "predict" in url_route: - from litellm.llms.vertex_ai.image_generation.image_generation_handler import ( - VertexImageGeneration, + return VertexPassthroughLoggingHandler._handle_predict_response( + httpx_response=httpx_response, + logging_obj=logging_obj, + url_route=url_route, + kwargs=kwargs, ) - from litellm.llms.vertex_ai.multimodal_embeddings.transformation import ( - VertexAIMultimodalEmbeddingConfig, - ) - from litellm.types.utils import PassthroughCallTypes - - vertex_image_generation_class = VertexImageGeneration() - - model = VertexPassthroughLoggingHandler.extract_model_from_url(url_route) - - _json_response = httpx_response.json() - - litellm_prediction_response: Union[ - ModelResponse, EmbeddingResponse, ImageResponse - ] = ModelResponse() - if vertex_image_generation_class.is_image_generation_response( - _json_response - ): - litellm_prediction_response = ( - vertex_image_generation_class.process_image_generation_response( - _json_response, - model_response=litellm.ImageResponse(), - model=model, - ) - ) - - logging_obj.call_type = ( - PassthroughCallTypes.passthrough_image_generation.value - ) - elif VertexPassthroughLoggingHandler._is_multimodal_embedding_response( - json_response=_json_response, - ): - # Use multimodal embedding transformation - vertex_multimodal_config = VertexAIMultimodalEmbeddingConfig() - litellm_prediction_response = ( - vertex_multimodal_config.transform_embedding_response( - model=model, - raw_response=httpx_response, - model_response=litellm.EmbeddingResponse(), - logging_obj=logging_obj, - api_key="", - request_data={}, - optional_params={}, - litellm_params={}, - ) - ) - else: - litellm_prediction_response = litellm.vertexAITextEmbeddingConfig.transform_vertex_response_to_openai( - response=_json_response, - model=model, - model_response=litellm.EmbeddingResponse(), - ) - if isinstance(litellm_prediction_response, litellm.EmbeddingResponse): - litellm_prediction_response.model = model - - logging_obj.model = model - logging_obj.model_call_details["model"] = logging_obj.model - logging_obj.model_call_details["custom_llm_provider"] = "vertex_ai" - logging_obj.custom_llm_provider = "vertex_ai" - response_cost = litellm.completion_cost( - completion_response=litellm_prediction_response, - model=model, - custom_llm_provider="vertex_ai", - ) - - kwargs["response_cost"] = response_cost - kwargs["model"] = model - kwargs["custom_llm_provider"] = "vertex_ai" - logging_obj.model_call_details["response_cost"] = response_cost - - return { - "result": litellm_prediction_response, - "kwargs": kwargs, - } elif "rawPredict" in url_route or "streamRawPredict" in url_route: from litellm.llms.vertex_ai.vertex_ai_partner_models import ( get_vertex_ai_partner_model_config, @@ -266,6 +196,91 @@ class VertexPassthroughLoggingHandler: "kwargs": kwargs, } + @staticmethod + def _handle_predict_response( + httpx_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + url_route: str, + kwargs: dict, + ) -> PassThroughEndpointLoggingTypedDict: + """Handle predict endpoint responses (embeddings, image generation).""" + from litellm.llms.vertex_ai.image_generation.image_generation_handler import ( + VertexImageGeneration, + ) + from litellm.llms.vertex_ai.multimodal_embeddings.transformation import ( + VertexAIMultimodalEmbeddingConfig, + ) + from litellm.types.utils import PassthroughCallTypes + + vertex_image_generation_class = VertexImageGeneration() + + model = VertexPassthroughLoggingHandler.extract_model_from_url(url_route) + + _json_response = httpx_response.json() + + litellm_prediction_response: Union[ + ModelResponse, EmbeddingResponse, ImageResponse + ] = ModelResponse() + if vertex_image_generation_class.is_image_generation_response( + _json_response + ): + litellm_prediction_response = ( + vertex_image_generation_class.process_image_generation_response( + _json_response, + model_response=litellm.ImageResponse(), + model=model, + ) + ) + + logging_obj.call_type = ( + PassthroughCallTypes.passthrough_image_generation.value + ) + elif VertexPassthroughLoggingHandler._is_multimodal_embedding_response( + json_response=_json_response, + ): + # Use multimodal embedding transformation + vertex_multimodal_config = VertexAIMultimodalEmbeddingConfig() + litellm_prediction_response = ( + vertex_multimodal_config.transform_embedding_response( + model=model, + raw_response=httpx_response, + model_response=litellm.EmbeddingResponse(), + logging_obj=logging_obj, + api_key="", + request_data={}, + optional_params={}, + litellm_params={}, + ) + ) + else: + litellm_prediction_response = litellm.vertexAITextEmbeddingConfig.transform_vertex_response_to_openai( + response=_json_response, + model=model, + model_response=litellm.EmbeddingResponse(), + ) + if isinstance(litellm_prediction_response, litellm.EmbeddingResponse): + litellm_prediction_response.model = model + + logging_obj.model = model + logging_obj.model_call_details["model"] = logging_obj.model + logging_obj.model_call_details["custom_llm_provider"] = "vertex_ai" + logging_obj.custom_llm_provider = "vertex_ai" + response_cost = litellm.completion_cost( + completion_response=litellm_prediction_response, + model=model, + custom_llm_provider="vertex_ai", + ) + + kwargs["response_cost"] = response_cost + kwargs["model"] = model + kwargs["custom_llm_provider"] = "vertex_ai" + logging_obj.model_call_details["response_cost"] = response_cost + + return { + "result": litellm_prediction_response, + "kwargs": kwargs, + } + @staticmethod def _handle_logging_vertex_collected_chunks( litellm_logging_obj: LiteLLMLoggingObj,