From 3b7f1d563dc0e35fb723c15a761cac146ea09e96 Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Tue, 10 Jun 2025 21:26:46 -0700 Subject: [PATCH] Add VertexAI Anthropic passthrough - cost calculation, token tracking (#11611) * fix(vertex_ai/anthropic): extract vertex credentials from correct dict Fixes passing credentials for vertex anthropic via `/v1/messages` * fix(vertex_ai/anthropic): working e2e vertex ai anthropic messages via `/v1/messages` Ensures vertex anthropic models can be called via /v1/messages * fix(vertex_passthrough_logging_handler.py): fix vertex anthropic logging to emit the model name on spend logs enables observability of passthrough route * fix(vertex_passthrough_logging_handler.py): don't redefine variable * fix(vertex_passthrough_logging_handler.py): log cost + token usage for anthropic passthrough requests on vertex ai --- .../vertex_ai_partner_models/__init__.py | 24 ++++++++ .../vertex_passthrough_logging_handler.py | 60 ++++++++++++++++++- 2 files changed, 81 insertions(+), 3 deletions(-) create mode 100644 litellm/llms/vertex_ai/vertex_ai_partner_models/__init__.py diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/__init__.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/__init__.py new file mode 100644 index 00000000000..cc0ecc2e3c6 --- /dev/null +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/__init__.py @@ -0,0 +1,24 @@ +from litellm.llms.base_llm.chat.transformation import BaseConfig + + +def get_vertex_ai_partner_model_config( + model: str, vertex_publisher_or_api_spec: str +) -> BaseConfig: + """Return config for handling response transformation for vertex ai partner models""" + if vertex_publisher_or_api_spec == "anthropic": + from .anthropic.transformation import VertexAIAnthropicConfig + + return VertexAIAnthropicConfig() + elif vertex_publisher_or_api_spec == "ai21": + from .ai21.transformation import VertexAIAi21Config + + return VertexAIAi21Config() + elif ( + vertex_publisher_or_api_spec == "openapi" + or vertex_publisher_or_api_spec == "mistralai" + ): + from .llama3.transformation import VertexAILlama3Config + + return VertexAILlama3Config() + else: + raise ValueError(f"Unsupported model: {model}") 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 131a0031c9c..ab037b0f1e9 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 @@ -123,17 +123,56 @@ class VertexPassthroughLoggingHandler: "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, + ) + model = VertexPassthroughLoggingHandler.extract_model_from_url(url_route) + vertex_publisher_or_api_spec = VertexPassthroughLoggingHandler._get_vertex_publisher_or_api_spec_from_url( + url_route + ) _json_response = httpx_response.json() litellm_prediction_response = ModelResponse() + + if vertex_publisher_or_api_spec is not None: + vertex_ai_partner_model_config = get_vertex_ai_partner_model_config( + model=model, + vertex_publisher_or_api_spec=vertex_publisher_or_api_spec, + ) + litellm_prediction_response = ( + vertex_ai_partner_model_config.transform_response( + model=model, + raw_response=httpx_response, + model_response=litellm_prediction_response, + logging_obj=logging_obj, + request_data={}, + encoding=litellm.encoding, + optional_params={}, + litellm_params={}, + api_key="", + messages=[ + { + "role": "user", + "content": "no-message-pass-through-endpoint", + } + ], + ) + ) - logging_obj.model = model - logging_obj.model_call_details["model"] = logging_obj.model + kwargs = VertexPassthroughLoggingHandler._create_vertex_response_logging_payload_for_generate_content( + litellm_model_response=litellm_prediction_response, + model="vertex_ai/" + model, + kwargs=kwargs, + start_time=start_time, + end_time=end_time, + logging_obj=logging_obj, + custom_llm_provider="vertex_ai", + ) return { - "result": None, + "result": litellm_prediction_response, "kwargs": kwargs, } else: @@ -236,6 +275,19 @@ class VertexPassthroughLoggingHandler: return match.group(1) return "unknown" + @staticmethod + def _get_vertex_publisher_or_api_spec_from_url(url: str) -> Optional[str]: + # Check for specific Vertex AI partner publishers + if "/publishers/mistralai/" in url: + return "mistralai" + elif "/publishers/anthropic/" in url: + return "anthropic" + elif "/publishers/ai21/" in url: + return "ai21" + elif "/endpoints/openapi/" in url: + return "openapi" + return None + @staticmethod def _get_custom_llm_provider_from_url(url: str) -> str: parsed_url = urlparse(url) @@ -259,10 +311,12 @@ class VertexPassthroughLoggingHandler: Create the standard logging object for Vertex passthrough generateContent (streaming and non-streaming) """ + response_cost = litellm.completion_cost( completion_response=litellm_model_response, model=model, ) + kwargs["response_cost"] = response_cost kwargs["model"] = model