fix code qa checks

This commit is contained in:
Ishaan Jaffer 2025-11-26 11:34:22 -08:00
parent 0f59e5fa3a
commit 8d2dba8cac
2 changed files with 97 additions and 79 deletions

View file

@ -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

View file

@ -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,