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 d709956df5c..d0fe37596f8 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 @@ -130,6 +130,13 @@ class VertexPassthroughLoggingHandler: "kwargs": kwargs, } + elif "embedContent" in url_route or "batchEmbedContents" in url_route: + return VertexPassthroughLoggingHandler._handle_embed_content_response( + httpx_response=httpx_response, + logging_obj=logging_obj, + url_route=url_route, + kwargs=kwargs, + ) elif "predict" in url_route: return VertexPassthroughLoggingHandler._handle_predict_response( httpx_response=httpx_response, @@ -322,6 +329,60 @@ class VertexPassthroughLoggingHandler: "kwargs": kwargs, } + @staticmethod + def _handle_embed_content_response( + httpx_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + url_route: str, + kwargs: dict, + ) -> PassThroughEndpointLoggingTypedDict: + """Handle Vertex :embedContent and :batchEmbedContents endpoint responses.""" + from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import ( + process_embed_content_response, + process_response as process_batch_embed_response, + ) + + model = VertexPassthroughLoggingHandler.extract_model_from_url(url_route) + response_json = httpx_response.json() + + model_response = litellm.EmbeddingResponse() + if "batchEmbedContents" in url_route: + litellm_embedding_response = process_batch_embed_response( + input="", + model_response=model_response, + model=model, + _predictions=response_json, + ) + else: + litellm_embedding_response = process_embed_content_response( + input="", + model_response=model_response, + model=model, + response_json=response_json, + ) + + litellm_embedding_response.model = model + logging_obj.model = model + logging_obj.model_call_details["model"] = 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_embedding_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_embedding_response, + "kwargs": kwargs, + } + @staticmethod def _handle_logging_vertex_collected_chunks( litellm_logging_obj: LiteLLMLoggingObj,