Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py: 15%
378 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1import asyncio
2import re
3from collections.abc import Mapping
4from datetime import datetime
5from typing import TYPE_CHECKING, Any, Final, Literal, cast
6from urllib.parse import urlparse
8import httpx
9from pydantic import TypeAdapter
11import litellm
12from litellm._logging import verbose_proxy_logger
13from litellm.constants import VERTEX_BATCH_PREDICTION_JOBS_ROUTE
14from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
15from litellm.litellm_core_utils.llm_cost_calc.usage_object_transformation import (
16 InteractionsUsageObjectTransformation,
17)
18from litellm.llms.vertex_ai.common_utils import (
19 get_vertex_ai_lyria_generation_cost,
20 get_vertex_location_from_url,
21)
22from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
23 ModelResponseIterator as VertexModelResponseIterator,
24)
25from litellm.llms.vertex_ai.vector_stores.search_api.transformation import (
26 VertexSearchAPIVectorStoreConfig,
27)
28from litellm.llms.vertex_ai.videos.transformation import VertexAIVideoConfig
29from litellm.proxy._types import PassThroughEndpointLoggingTypedDict
30from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution import (
31 is_collection_route,
32 log_batch_registration_result,
33 optional_str,
34 request_tags_from_metadata,
35)
36from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType
37from litellm.types.utils import (
38 Choices,
39 EmbeddingResponse,
40 ImageResponse,
41 ModelResponse,
42 SpecialEnums,
43 StandardPassThroughResponseObject,
44 TextCompletionResponse,
45)
47vertex_search_api_config: Final = VertexSearchAPIVectorStoreConfig()
48if TYPE_CHECKING: 48 ↛ 49line 48 didn't jump to line 49 because the condition on line 48 was never true
49 from litellm.types.utils import LiteLLMBatch
51 from ..success_handler import PassThroughEndpointLogging
52else:
53 PassThroughEndpointLogging = Any
54 LiteLLMBatch = Any
56_VERTEX_INTERACTIONS_PATH: Final = re.compile(r"/projects/[^/]+/locations/[^/]+/interactions/?$")
57_INTERACTIONS_RESPONSE_BODY: Final = TypeAdapter(dict[str, object])
60def _interactions_model(
61 response_body: Mapping[str, object],
62 request_body: Mapping[str, object] | None,
63) -> str | None:
64 response_model: Final = response_body.get("model")
65 if isinstance(response_model, str) and response_model:
66 return response_model
67 request_model: Final = (request_body or {}).get("model")
68 if isinstance(request_model, str) and request_model:
69 return request_model
70 return None
73class VertexPassthroughLoggingHandler:
74 @staticmethod
75 def is_interactions_route(url_route: str) -> bool:
76 return urlparse(url_route).path.rstrip("/").endswith("/interactions")
78 @staticmethod
79 def is_vertex_interactions_route(url_route: str) -> bool:
80 return _VERTEX_INTERACTIONS_PATH.search(urlparse(url_route).path) is not None
82 @staticmethod
83 def interactions_passthrough_handler(
84 httpx_response: httpx.Response,
85 request_body: Mapping[str, object] | None,
86 logging_obj: LiteLLMLoggingObj,
87 kwargs: dict[str, object],
88 start_time: datetime,
89 end_time: datetime,
90 custom_llm_provider: Literal["vertex_ai", "gemini"],
91 vertex_location: str | None,
92 ) -> PassThroughEndpointLoggingTypedDict:
93 response_body: Final = _INTERACTIONS_RESPONSE_BODY.validate_python(httpx_response.json())
94 usage_object: Final = response_body.get("usage")
95 model: Final = _interactions_model(response_body, request_body)
96 if model is None or not InteractionsUsageObjectTransformation.is_interactions_usage_object(usage_object):
97 return {"result": None, "kwargs": kwargs}
99 litellm_model_response: Final = ModelResponse(
100 model=model,
101 usage=InteractionsUsageObjectTransformation.transform_interactions_usage_object(
102 cast(Mapping[str, Any], usage_object)
103 ),
104 )
105 logging_obj.custom_llm_provider = custom_llm_provider
106 logging_kwargs: Final = (
107 VertexPassthroughLoggingHandler._create_vertex_response_logging_payload_for_generate_content(
108 litellm_model_response=litellm_model_response,
109 model=model,
110 kwargs=kwargs,
111 start_time=start_time,
112 end_time=end_time,
113 logging_obj=logging_obj,
114 custom_llm_provider=custom_llm_provider,
115 vertex_location=vertex_location,
116 )
117 )
118 return {
119 "result": litellm_model_response,
120 "kwargs": {**logging_kwargs, "custom_llm_provider": custom_llm_provider},
121 }
123 @staticmethod
124 def vertex_passthrough_handler(
125 httpx_response: httpx.Response,
126 logging_obj: LiteLLMLoggingObj,
127 url_route: str,
128 result: str,
129 start_time: datetime,
130 end_time: datetime,
131 cache_hit: bool,
132 request_body: dict | None = None,
133 **kwargs,
134 ) -> PassThroughEndpointLoggingTypedDict:
135 vertex_location: Final = get_vertex_location_from_url(url_route)
136 if vertex_location is not None:
137 logging_obj.optional_params["vertex_location"] = vertex_location
138 if VertexPassthroughLoggingHandler.is_interactions_route(url_route):
139 return VertexPassthroughLoggingHandler.interactions_passthrough_handler(
140 httpx_response=httpx_response,
141 request_body=request_body,
142 logging_obj=logging_obj,
143 kwargs=kwargs,
144 start_time=start_time,
145 end_time=end_time,
146 custom_llm_provider="vertex_ai",
147 vertex_location=vertex_location,
148 )
149 if "predictLongRunning" in url_route:
150 model = VertexPassthroughLoggingHandler.extract_model_from_url(url_route)
152 vertex_video_config: Final = VertexAIVideoConfig()
153 litellm_video_response: Final = vertex_video_config.transform_video_create_response(
154 model=model,
155 raw_response=httpx_response,
156 logging_obj=logging_obj,
157 custom_llm_provider="vertex_ai",
158 request_data=request_body,
159 )
161 logging_obj.model = model
162 logging_obj.model_call_details["model"] = model
163 logging_obj.model_call_details["custom_llm_provider"] = "vertex_ai"
164 logging_obj.custom_llm_provider = "vertex_ai"
166 response_cost = litellm.completion_cost(
167 completion_response=litellm_video_response,
168 model=model,
169 custom_llm_provider="vertex_ai",
170 call_type="create_video",
171 vertex_location=vertex_location,
172 )
174 # Set response_cost in _hidden_params to prevent recalculation
175 if not hasattr(litellm_video_response, "_hidden_params"):
176 litellm_video_response._hidden_params = {}
177 litellm_video_response._hidden_params["response_cost"] = response_cost
179 kwargs["response_cost"] = response_cost
180 kwargs["model"] = model
181 kwargs["custom_llm_provider"] = "vertex_ai"
182 logging_obj.model_call_details["response_cost"] = response_cost
184 return {
185 "result": litellm_video_response,
186 "kwargs": kwargs,
187 }
189 elif "generateContent" in url_route:
190 model = VertexPassthroughLoggingHandler.extract_model_from_url(url_route)
192 instance_of_vertex_llm: Final = litellm.VertexGeminiConfig()
193 litellm_model_response: Final[ModelResponse] = instance_of_vertex_llm.transform_response(
194 model=model,
195 messages=[{"role": "user", "content": "no-message-pass-through-endpoint"}],
196 raw_response=httpx_response,
197 model_response=litellm.ModelResponse(),
198 logging_obj=logging_obj,
199 optional_params={},
200 litellm_params={},
201 api_key="",
202 request_data={},
203 encoding=getattr(litellm, "encoding", None),
204 )
205 kwargs = VertexPassthroughLoggingHandler._create_vertex_response_logging_payload_for_generate_content(
206 litellm_model_response=litellm_model_response,
207 model=model,
208 kwargs=kwargs,
209 start_time=start_time,
210 end_time=end_time,
211 logging_obj=logging_obj,
212 custom_llm_provider=VertexPassthroughLoggingHandler._get_custom_llm_provider_from_url(url_route),
213 vertex_location=vertex_location,
214 )
216 return {
217 "result": litellm_model_response,
218 "kwargs": kwargs,
219 }
221 elif "embedContent" in url_route or "batchEmbedContents" in url_route:
222 return VertexPassthroughLoggingHandler._handle_embed_content_response(
223 httpx_response=httpx_response,
224 logging_obj=logging_obj,
225 url_route=url_route,
226 kwargs=kwargs,
227 request_body=request_body,
228 )
229 elif "predict" in url_route:
230 return VertexPassthroughLoggingHandler._handle_predict_response(
231 httpx_response=httpx_response,
232 logging_obj=logging_obj,
233 url_route=url_route,
234 kwargs=kwargs,
235 )
236 elif "rawPredict" in url_route or "streamRawPredict" in url_route:
237 from litellm.llms.vertex_ai.vertex_ai_partner_models import (
238 get_vertex_ai_partner_model_config,
239 )
241 model = VertexPassthroughLoggingHandler.extract_model_from_url(url_route)
242 vertex_publisher_or_api_spec = VertexPassthroughLoggingHandler._get_vertex_publisher_or_api_spec_from_url(
243 url_route
244 )
246 _json_response: Final = httpx_response.json()
248 litellm_prediction_response = ModelResponse()
250 if vertex_publisher_or_api_spec is not None:
251 vertex_ai_partner_model_config: Final = get_vertex_ai_partner_model_config(
252 model=model,
253 vertex_publisher_or_api_spec=vertex_publisher_or_api_spec,
254 )
255 litellm_prediction_response = vertex_ai_partner_model_config.transform_response(
256 model=model,
257 raw_response=httpx_response,
258 model_response=litellm_prediction_response,
259 logging_obj=logging_obj,
260 request_data={},
261 encoding=litellm.encoding,
262 optional_params={},
263 litellm_params={},
264 api_key="",
265 messages=[
266 {
267 "role": "user",
268 "content": "no-message-pass-through-endpoint",
269 }
270 ],
271 )
273 kwargs = VertexPassthroughLoggingHandler._create_vertex_response_logging_payload_for_generate_content(
274 litellm_model_response=litellm_prediction_response,
275 model="vertex_ai/" + model,
276 kwargs=kwargs,
277 start_time=start_time,
278 end_time=end_time,
279 logging_obj=logging_obj,
280 custom_llm_provider="vertex_ai",
281 vertex_location=vertex_location,
282 )
284 return {
285 "result": litellm_prediction_response,
286 "kwargs": kwargs,
287 }
288 elif "search" in url_route:
289 litellm_vs_response: Final = vertex_search_api_config.transform_search_vector_store_response(
290 response=httpx_response,
291 litellm_logging_obj=logging_obj,
292 )
293 response_cost = litellm.completion_cost(
294 completion_response=litellm_vs_response,
295 model="vertex_ai/search_api",
296 custom_llm_provider="vertex_ai",
297 call_type="vector_store_search",
298 vertex_location=vertex_location,
299 )
301 standard_pass_through_response_object: Final[StandardPassThroughResponseObject] = {
302 "response": cast(dict, litellm_vs_response),
303 }
305 kwargs["response_cost"] = response_cost
306 kwargs["model"] = "vertex_ai/search_api"
307 logging_obj.model_call_details.setdefault("litellm_params", {})
308 logging_obj.model_call_details["litellm_params"]["base_model"] = "vertex_ai/search_api"
309 logging_obj.model_call_details["response_cost"] = response_cost
311 return {
312 "result": standard_pass_through_response_object,
313 "kwargs": kwargs,
314 }
315 elif "batchPredictionJobs" in url_route:
316 return VertexPassthroughLoggingHandler.batch_prediction_jobs_handler(
317 httpx_response=httpx_response,
318 logging_obj=logging_obj,
319 url_route=url_route,
320 result=result,
321 start_time=start_time,
322 end_time=end_time,
323 cache_hit=cache_hit,
324 **kwargs,
325 )
326 else:
327 return {
328 "result": None,
329 "kwargs": kwargs,
330 }
332 @staticmethod
333 def _handle_predict_response(
334 httpx_response: httpx.Response,
335 logging_obj: LiteLLMLoggingObj,
336 url_route: str,
337 kwargs: dict,
338 ) -> PassThroughEndpointLoggingTypedDict:
339 """Handle predict endpoint responses (embeddings, image generation)."""
340 from litellm.llms.vertex_ai.image_generation.image_generation_handler import (
341 VertexImageGeneration,
342 )
343 from litellm.llms.vertex_ai.multimodal_embeddings.transformation import (
344 VertexAIMultimodalEmbeddingConfig,
345 )
346 from litellm.types.utils import PassthroughCallTypes
348 vertex_image_generation_class: Final = VertexImageGeneration()
350 model: Final = VertexPassthroughLoggingHandler.extract_model_from_url(url_route)
352 _json_response: Final[dict[str, object]] = httpx_response.json()
354 litellm_prediction_response: ModelResponse | EmbeddingResponse | ImageResponse = ModelResponse()
355 if VertexPassthroughLoggingHandler._is_audio_predict_response(
356 model=model,
357 json_response=_json_response,
358 ):
359 return VertexPassthroughLoggingHandler._handle_audio_predict_response(
360 json_response=_json_response,
361 logging_obj=logging_obj,
362 model=model,
363 kwargs=kwargs,
364 )
365 if vertex_image_generation_class.is_image_generation_response(_json_response):
366 litellm_prediction_response = vertex_image_generation_class.process_image_generation_response(
367 _json_response,
368 model_response=litellm.ImageResponse(),
369 model=model,
370 )
372 logging_obj.call_type = PassthroughCallTypes.passthrough_image_generation.value
373 elif VertexPassthroughLoggingHandler._is_multimodal_embedding_response(
374 json_response=_json_response,
375 ):
376 # Use multimodal embedding transformation
377 vertex_multimodal_config: Final = VertexAIMultimodalEmbeddingConfig()
378 litellm_prediction_response = vertex_multimodal_config.transform_embedding_response(
379 model=model,
380 raw_response=httpx_response,
381 model_response=litellm.EmbeddingResponse(),
382 logging_obj=logging_obj,
383 api_key="",
384 request_data={},
385 optional_params={},
386 litellm_params={},
387 )
388 else:
389 litellm_prediction_response = litellm.vertexAITextEmbeddingConfig.transform_vertex_response_to_openai(
390 response=_json_response,
391 model=model,
392 model_response=litellm.EmbeddingResponse(),
393 )
394 if isinstance(litellm_prediction_response, litellm.EmbeddingResponse):
395 litellm_prediction_response.model = model
397 logging_obj.model = model
398 logging_obj.model_call_details["model"] = logging_obj.model
399 logging_obj.model_call_details["custom_llm_provider"] = "vertex_ai"
400 logging_obj.custom_llm_provider = "vertex_ai"
401 response_cost: Final = litellm.completion_cost(
402 completion_response=litellm_prediction_response,
403 model=model,
404 custom_llm_provider="vertex_ai",
405 vertex_location=get_vertex_location_from_url(url_route),
406 )
408 kwargs["response_cost"] = response_cost
409 kwargs["model"] = model
410 kwargs["custom_llm_provider"] = "vertex_ai"
411 logging_obj.model_call_details["response_cost"] = response_cost
413 return {
414 "result": litellm_prediction_response,
415 "kwargs": kwargs,
416 }
418 @staticmethod
419 def _handle_audio_predict_response(
420 json_response: dict, # mutable-ok: passthrough logging receives the decoded provider response dictionary
421 logging_obj: LiteLLMLoggingObj,
422 model: str,
423 kwargs: dict, # mutable-ok: passthrough logging enriches the shared callback metadata dictionary
424 ) -> PassThroughEndpointLoggingTypedDict:
425 prediction_count: Final = VertexPassthroughLoggingHandler._get_audio_prediction_count(
426 json_response=json_response
427 )
428 response_cost: Final = (get_vertex_ai_lyria_generation_cost(model=model) or 0.0) * prediction_count
430 logging_obj.model = model # rebind-ok: passthrough attribution records the resolved Vertex model
431 logging_obj.model_call_details[ # rebind-ok: passthrough attribution enriches callback metadata
432 "model"
433 ] = model
434 logging_obj.model_call_details[ # rebind-ok: passthrough attribution enriches callback metadata
435 "custom_llm_provider"
436 ] = "vertex_ai"
437 logging_obj.custom_llm_provider = ( # rebind-ok: attribution records the resolved provider
438 "vertex_ai"
439 )
440 logging_obj.model_call_details[ # rebind-ok: passthrough attribution enriches callback metadata
441 "response_cost"
442 ] = response_cost
444 kwargs[ # rebind-ok: callback metadata is enriched for downstream hooks
445 "response_cost"
446 ] = response_cost
447 kwargs["model"] = model # rebind-ok: callback metadata records the resolved model
448 kwargs["custom_llm_provider"] = "vertex_ai" # rebind-ok: callback metadata records the resolved provider
450 standard_pass_through_response_object: Final[StandardPassThroughResponseObject] = {
451 "response": json_response,
452 }
453 return { # mutable-ok: passthrough logging contract requires a concrete result dictionary
454 "result": standard_pass_through_response_object,
455 "kwargs": kwargs,
456 }
458 @staticmethod
459 def _is_audio_predict_response(
460 model: str,
461 json_response: Mapping[str, object],
462 ) -> bool:
463 return (
464 VertexPassthroughLoggingHandler._get_audio_prediction_count(json_response=json_response) > 0
465 and get_vertex_ai_lyria_generation_cost(model=model) is not None
466 )
468 @staticmethod
469 def _get_audio_prediction_count(
470 json_response: Mapping[str, object],
471 ) -> int:
472 predictions: Final = json_response.get("predictions")
473 if not isinstance(predictions, list):
474 return 0
475 return sum(
476 1
477 for prediction in predictions
478 if isinstance(prediction, dict) and (prediction.get("audioContent") or prediction.get("bytesBase64Encoded"))
479 )
481 @staticmethod
482 def _extract_embed_content_input(request_body: dict | None, batch: bool) -> str:
483 """Extract raw input text from an :embedContent or :batchEmbedContents request body for token counting."""
484 if not request_body:
485 return ""
486 if batch:
487 texts: Final = []
488 for req in request_body.get("requests", []):
489 for part in req.get("content", {}).get("parts", []):
490 texts.append(part.get("text", ""))
491 return " ".join(texts)
492 else:
493 parts: Final = request_body.get("content", {}).get("parts", [])
494 return " ".join(part.get("text", "") for part in parts)
496 @staticmethod
497 def _handle_embed_content_response(
498 httpx_response: httpx.Response,
499 logging_obj: LiteLLMLoggingObj,
500 url_route: str,
501 kwargs: dict,
502 request_body: dict | None = None,
503 ) -> PassThroughEndpointLoggingTypedDict:
504 """Handle Vertex :embedContent and :batchEmbedContents endpoint responses."""
505 from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import (
506 process_embed_content_response,
507 )
508 from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import (
509 process_response as process_batch_embed_response,
510 )
512 model: Final = VertexPassthroughLoggingHandler.extract_model_from_url(url_route)
513 response_json: Final = httpx_response.json()
514 is_batch: Final = "batchEmbedContents" in url_route
516 input_text: Final = VertexPassthroughLoggingHandler._extract_embed_content_input(
517 request_body=request_body, batch=is_batch
518 )
520 model_response: Final = litellm.EmbeddingResponse()
521 if is_batch:
522 litellm_embedding_response = process_batch_embed_response(
523 input=input_text,
524 model_response=model_response,
525 model=model,
526 _predictions=response_json,
527 )
528 else:
529 litellm_embedding_response = process_embed_content_response(
530 input=input_text,
531 model_response=model_response,
532 model=model,
533 response_json=response_json,
534 )
536 custom_llm_provider: Final = VertexPassthroughLoggingHandler._get_custom_llm_provider_from_url(url_route)
538 litellm_embedding_response.model = model
539 logging_obj.model = model
540 logging_obj.model_call_details["model"] = model
541 logging_obj.model_call_details["custom_llm_provider"] = custom_llm_provider
542 logging_obj.custom_llm_provider = custom_llm_provider
544 response_cost: Final = litellm.completion_cost(
545 completion_response=litellm_embedding_response,
546 model=model,
547 custom_llm_provider=custom_llm_provider,
548 vertex_location=get_vertex_location_from_url(url_route),
549 )
551 kwargs["response_cost"] = response_cost
552 kwargs["model"] = model
553 kwargs["custom_llm_provider"] = custom_llm_provider
554 logging_obj.model_call_details["response_cost"] = response_cost
556 return {
557 "result": litellm_embedding_response,
558 "kwargs": kwargs,
559 }
561 @staticmethod
562 def _handle_logging_vertex_collected_chunks(
563 litellm_logging_obj: LiteLLMLoggingObj,
564 passthrough_success_handler_obj: PassThroughEndpointLogging,
565 url_route: str,
566 request_body: dict,
567 endpoint_type: EndpointType,
568 start_time: datetime,
569 all_chunks: list[str],
570 model: str | None,
571 end_time: datetime,
572 ) -> PassThroughEndpointLoggingTypedDict:
573 """
574 Takes raw chunks from Vertex passthrough endpoint and logs them in litellm callbacks
576 - Builds complete response from chunks
577 - Creates standard logging object
578 - Logs in litellm callbacks
579 """
580 kwargs: dict[str, object] = {}
581 vertex_location: Final = get_vertex_location_from_url(url_route)
582 if vertex_location is not None:
583 litellm_logging_obj.optional_params["vertex_location"] = vertex_location
584 model = model or VertexPassthroughLoggingHandler.extract_model_from_url(url_route)
585 complete_streaming_response: Final = VertexPassthroughLoggingHandler._build_complete_streaming_response(
586 all_chunks=all_chunks,
587 litellm_logging_obj=litellm_logging_obj,
588 model=model,
589 url_route=url_route,
590 )
592 if complete_streaming_response is None:
593 verbose_proxy_logger.error(
594 "Unable to build complete streaming response for Vertex passthrough endpoint, not logging..."
595 )
596 return {
597 "result": None,
598 "kwargs": kwargs,
599 }
601 kwargs = VertexPassthroughLoggingHandler._create_vertex_response_logging_payload_for_generate_content(
602 litellm_model_response=complete_streaming_response,
603 model=model,
604 kwargs=kwargs,
605 start_time=start_time,
606 end_time=end_time,
607 logging_obj=litellm_logging_obj,
608 custom_llm_provider=VertexPassthroughLoggingHandler._get_custom_llm_provider_from_url(url_route),
609 vertex_location=vertex_location,
610 )
612 return {
613 "result": complete_streaming_response,
614 "kwargs": kwargs,
615 }
617 @staticmethod
618 def _build_complete_streaming_response(
619 all_chunks: list[str],
620 litellm_logging_obj: LiteLLMLoggingObj,
621 model: str,
622 url_route: str,
623 ) -> ModelResponse | TextCompletionResponse | None:
624 parsed_chunks = []
625 if "generateContent" in url_route or "streamGenerateContent" in url_route:
626 vertex_iterator: Any = VertexModelResponseIterator(
627 streaming_response=None,
628 sync_stream=False,
629 logging_obj=litellm_logging_obj,
630 )
631 chunk_parsing_logic: Any = vertex_iterator._common_chunk_parsing_logic
632 parsed_chunks = [chunk_parsing_logic(chunk) for chunk in all_chunks]
633 elif "rawPredict" in url_route or "streamRawPredict" in url_route:
634 from litellm.llms.anthropic.chat.handler import ModelResponseIterator
635 from litellm.llms.base_llm.base_model_iterator import (
636 BaseModelResponseIterator,
637 )
639 vertex_iterator = ModelResponseIterator(
640 streaming_response=None,
641 sync_stream=False,
642 )
643 chunk_parsing_logic = vertex_iterator.chunk_parser
644 for chunk in all_chunks:
645 dict_chunk = BaseModelResponseIterator._string_to_dict_parser(chunk)
646 if dict_chunk is None:
647 continue
648 parsed_chunks.append(chunk_parsing_logic(dict_chunk))
649 else:
650 return None
651 if len(parsed_chunks) == 0:
652 return None
653 all_openai_chunks: Final = []
654 for parsed_chunk in parsed_chunks:
655 if parsed_chunk is None:
656 continue
657 all_openai_chunks.append(parsed_chunk)
659 complete_streaming_response: Final = litellm.stream_chunk_builder(chunks=all_openai_chunks)
661 return complete_streaming_response
663 @staticmethod
664 def extract_model_from_url(url: str) -> str:
665 pattern: Final = r"/models/([^:]+)"
666 match: Final = re.search(pattern, url)
667 if match:
668 return match.group(1)
669 return "unknown"
671 @staticmethod
672 def extract_model_name_from_vertex_path(vertex_model_path: str) -> str:
673 """
674 Extract the actual model name from a Vertex AI model path.
676 Examples:
677 - publishers/google/models/gemini-2.5-flash -> gemini-2.5-flash
678 - projects/PROJECT_ID/locations/LOCATION/models/MODEL_ID -> MODEL_ID
680 Args:
681 vertex_model_path: The full Vertex AI model path
683 Returns:
684 The extracted model name for use with LiteLLM
685 """
686 # Handle publishers/google/models/ format
687 if (
688 "publishers/" in vertex_model_path
689 and "models/" in vertex_model_path
690 or "projects/" in vertex_model_path
691 and "models/" in vertex_model_path
692 ):
693 # Extract everything after the last models/
694 parts: Final = vertex_model_path.split("models/")
695 if len(parts) > 1:
696 return parts[-1]
698 # If no recognized pattern, return the original path
699 return vertex_model_path
701 @staticmethod
702 def _get_vertex_publisher_or_api_spec_from_url(url: str) -> str | None:
703 # Check for specific Vertex AI partner publishers
704 if "/publishers/mistralai/" in url:
705 return "mistralai"
706 elif "/publishers/anthropic/" in url:
707 return "anthropic"
708 elif "/publishers/ai21/" in url:
709 return "ai21"
710 elif "/endpoints/openapi/" in url:
711 return "openapi"
712 return None
714 @staticmethod
715 def _get_custom_llm_provider_from_url(url: str) -> str:
716 parsed_url: Final = urlparse(url)
717 if parsed_url.hostname and parsed_url.hostname.endswith("generativelanguage.googleapis.com"):
718 return litellm.LlmProviders.GEMINI.value
719 return litellm.LlmProviders.VERTEX_AI.value
721 @staticmethod
722 def _is_multimodal_embedding_response(json_response: dict) -> bool:
723 """
724 Detect if the response is from a multimodal embedding request.
726 Check if the response contains multimodal embedding fields:
727 - Docs: https://cloud.google.com/vertex-ai/generative-ai/docs/model-reference/multimodal-embeddings-api#response-body
730 Args:
731 json_response: The JSON response from Vertex AI
733 Returns:
734 bool: True if this is a multimodal embedding response
735 """
736 # Check if response contains multimodal embedding fields
737 if "predictions" in json_response:
738 predictions: Final = json_response["predictions"]
739 for prediction in predictions:
740 if isinstance(prediction, dict):
741 # Check for multimodal embedding response fields
742 if any(
743 key in prediction
744 for key in [
745 "textEmbedding",
746 "imageEmbedding",
747 "videoEmbeddings",
748 ]
749 ):
750 return True
752 return False
754 @staticmethod
755 def _create_vertex_response_logging_payload_for_generate_content(
756 litellm_model_response: ModelResponse | TextCompletionResponse,
757 model: str,
758 kwargs: dict,
759 start_time: datetime,
760 end_time: datetime,
761 logging_obj: LiteLLMLoggingObj,
762 custom_llm_provider: str,
763 vertex_location: str | None,
764 ) -> dict:
765 """
766 Create the standard logging object for Vertex passthrough generateContent (streaming and non-streaming)
768 """
770 response_cost: Final = litellm.completion_cost(
771 completion_response=litellm_model_response,
772 model=model,
773 custom_llm_provider=custom_llm_provider,
774 vertex_location=vertex_location,
775 )
777 kwargs["response_cost"] = response_cost
778 kwargs["model"] = model
780 # pretty print standard logging object
781 verbose_proxy_logger.debug("kwargs= %s", kwargs)
783 # set litellm_call_id to logging response object
784 litellm_model_response.id = logging_obj.litellm_call_id
785 logging_obj.model = litellm_model_response.model or model
786 logging_obj.model_call_details["model"] = logging_obj.model
787 logging_obj.model_call_details["custom_llm_provider"] = custom_llm_provider
788 return kwargs
790 @staticmethod
791 def batch_prediction_jobs_handler(
792 httpx_response: httpx.Response,
793 logging_obj: LiteLLMLoggingObj,
794 url_route: str,
795 result: str,
796 start_time: datetime,
797 end_time: datetime,
798 cache_hit: bool,
799 **kwargs,
800 ) -> PassThroughEndpointLoggingTypedDict:
801 """
802 Handle batch prediction jobs passthrough logging.
803 Creates a managed object for cost tracking when batch job is successfully created.
804 """
805 import base64
807 from litellm._uuid import uuid
808 from litellm.llms.vertex_ai.batches.transformation import (
809 VertexAIBatchTransformation,
810 )
812 try:
813 _json_response: Final = httpx_response.json()
815 # Only handle successful batch job creation (POST requests)
816 if httpx_response.status_code == 200 and "name" in _json_response:
817 # Transform Vertex AI response to LiteLLM batch format
818 litellm_batch_response: Final = (
819 VertexAIBatchTransformation.transform_vertex_ai_batch_response_to_openai_batch_response(
820 response=_json_response
821 )
822 )
824 # Extract batch ID and model from the response
825 batch_id = VertexAIBatchTransformation._get_batch_id_from_vertex_ai_batch_response(_json_response)
826 model_name: Final = _json_response.get("model", "unknown")
828 # Create unified object ID for tracking
829 # Format: base64(litellm_proxy;model_id:{};llm_batch_id:{})
830 actual_model_id: Final = VertexPassthroughLoggingHandler.get_actual_model_id_from_router(model_name)
832 unified_id_string: Final = SpecialEnums.LITELLM_MANAGED_BATCH_COMPLETE_STR.value.format(
833 actual_model_id, batch_id
834 )
835 unified_object_id: Final = base64.urlsafe_b64encode(unified_id_string.encode()).decode().rstrip("=")
837 # Store the managed object for cost tracking
838 # This will be picked up by check_batch_cost polling mechanism
839 is_batch_create: Final = is_collection_route(url_route, VERTEX_BATCH_PREDICTION_JOBS_ROUTE)
840 VertexPassthroughLoggingHandler._store_batch_managed_object(
841 unified_object_id=unified_object_id,
842 batch_object=litellm_batch_response,
843 model_object_id=batch_id,
844 logging_obj=logging_obj,
845 is_batch_create=is_batch_create,
846 **kwargs,
847 )
849 # Create a batch job response for logging
850 litellm_model_response = ModelResponse()
851 litellm_model_response.id = str(uuid.uuid4())
852 litellm_model_response.model = model_name
853 litellm_model_response.object = "batch_prediction_job"
854 litellm_model_response.created = int(start_time.timestamp())
856 # Add batch-specific metadata to indicate this is a pending batch job
857 litellm_model_response.choices = [
858 Choices(
859 finish_reason="stop",
860 index=0,
861 message={
862 "role": "assistant",
863 "content": f"Batch prediction job {batch_id} created and is pending. Status will be updated when the batch completes.",
864 "tool_calls": None,
865 "function_call": None,
866 "provider_specific_fields": {
867 "batch_job_id": batch_id,
868 "batch_job_state": "JOB_STATE_PENDING",
869 "unified_object_id": unified_object_id,
870 },
871 },
872 )
873 ]
875 # Set response cost to 0 initially (will be updated when batch completes)
876 response_cost: Final = 0.0
877 kwargs["response_cost"] = response_cost
878 kwargs["model"] = model_name
879 kwargs["batch_id"] = batch_id
880 kwargs["unified_object_id"] = unified_object_id
881 kwargs["batch_job_state"] = "JOB_STATE_PENDING"
883 logging_obj.model = model_name
884 logging_obj.model_call_details["model"] = logging_obj.model
885 logging_obj.model_call_details["response_cost"] = response_cost
886 logging_obj.model_call_details["batch_id"] = batch_id
888 return {
889 "result": litellm_model_response,
890 "kwargs": kwargs,
891 }
892 else:
893 # Handle non-successful responses
894 litellm_model_response = ModelResponse()
895 litellm_model_response.id = str(uuid.uuid4())
896 litellm_model_response.model = "vertex_ai_batch"
897 litellm_model_response.object = "batch_prediction_job"
898 litellm_model_response.created = int(start_time.timestamp())
900 # Add error-specific metadata
901 litellm_model_response.choices = [
902 Choices(
903 finish_reason="stop",
904 index=0,
905 message={
906 "role": "assistant",
907 "content": f"Batch prediction job creation failed. Status: {httpx_response.status_code}",
908 "tool_calls": None,
909 "function_call": None,
910 "provider_specific_fields": {
911 "batch_job_state": "JOB_STATE_FAILED",
912 "status_code": httpx_response.status_code,
913 },
914 },
915 )
916 ]
918 kwargs["response_cost"] = 0.0
919 kwargs["model"] = "vertex_ai_batch"
920 kwargs["batch_job_state"] = "JOB_STATE_FAILED"
922 return {
923 "result": litellm_model_response,
924 "kwargs": kwargs,
925 }
927 except Exception as e:
928 verbose_proxy_logger.error("Error in batch_prediction_jobs_handler: %s", e)
929 # Return basic response on error
930 litellm_model_response = ModelResponse()
931 litellm_model_response.id = str(uuid.uuid4())
932 litellm_model_response.model = "vertex_ai_batch"
933 litellm_model_response.object = "batch_prediction_job"
934 litellm_model_response.created = int(start_time.timestamp())
936 # Add error-specific metadata
937 litellm_model_response.choices = [
938 Choices(
939 finish_reason="stop",
940 index=0,
941 message={
942 "role": "assistant",
943 "content": f"Error creating batch prediction job: {e}",
944 "tool_calls": None,
945 "function_call": None,
946 "provider_specific_fields": {
947 "batch_job_state": "JOB_STATE_FAILED",
948 "error": str(e),
949 },
950 },
951 )
952 ]
954 kwargs["response_cost"] = 0.0
955 kwargs["model"] = "vertex_ai_batch"
956 kwargs["batch_job_state"] = "JOB_STATE_FAILED"
958 return {
959 "result": litellm_model_response,
960 "kwargs": kwargs,
961 }
963 @staticmethod
964 def _store_batch_managed_object(
965 unified_object_id: str,
966 batch_object: LiteLLMBatch,
967 model_object_id: str,
968 logging_obj: LiteLLMLoggingObj,
969 is_batch_create: bool,
970 **kwargs,
971 ) -> None:
972 """
973 Store batch managed object for cost tracking.
974 This will be picked up by the check_batch_cost polling mechanism.
976 A poll refreshes the batch status and file object but neither creates the row
977 nor writes attribution, so the creating key and its tags are persisted from
978 the create alone.
979 """
980 try:
981 # Get the managed files hook from the logging object
982 # This is a bit of a hack, but we need access to the proxy logging system
983 from litellm.proxy.proxy_server import proxy_logging_obj
985 managed_files_hook: Final = proxy_logging_obj.get_proxy_hook("managed_files")
986 if managed_files_hook is not None and hasattr(managed_files_hook, "store_unified_object_id"):
987 # Create a mock user API key dict for the managed object storage
988 from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
990 _request_metadata: Final = (kwargs.get("litellm_params", {}) or {}).get("metadata", {}) or {}
992 user_api_key_dict: Final = UserAPIKeyAuth(
993 user_id=_request_metadata.get("user_api_key_user_id", "default-user"),
994 api_key=optional_str(_request_metadata.get("user_api_key")),
995 team_id=_request_metadata.get("user_api_key_team_id"),
996 team_alias=None,
997 user_role=LitellmUserRoles.CUSTOMER, # Use proper enum value
998 user_email=None,
999 max_budget=None,
1000 spend=0.0, # Set to 0.0 instead of None
1001 models=[], # Set to empty list instead of None
1002 tpm_limit=None,
1003 rpm_limit=None,
1004 budget_duration=None,
1005 budget_reset_at=None,
1006 max_parallel_requests=None,
1007 allowed_model_region=None,
1008 metadata={}, # Set to empty dict instead of None
1009 key_alias=None,
1010 permissions={}, # Set to empty dict instead of None
1011 model_max_budget={}, # Set to empty dict instead of None
1012 model_spend={}, # Set to empty dict instead of None
1013 )
1015 # Store the unified object for batch cost tracking
1016 task: Final = asyncio.create_task(
1017 managed_files_hook.store_unified_object_id(
1018 unified_object_id=unified_object_id,
1019 file_object=batch_object,
1020 litellm_parent_otel_span=None,
1021 model_object_id=model_object_id,
1022 file_purpose="batch",
1023 user_api_key_dict=user_api_key_dict,
1024 request_tags=request_tags_from_metadata(_request_metadata),
1025 persist_attribution=is_batch_create,
1026 create_if_missing=is_batch_create,
1027 )
1028 )
1029 task.add_done_callback(
1030 lambda finished: log_batch_registration_result(
1031 finished, "Vertex AI", unified_object_id, model_object_id, is_batch_create
1032 )
1033 )
1034 else:
1035 verbose_proxy_logger.warning(
1036 "Managed files hook not available, cannot store batch object for cost tracking"
1037 )
1039 except Exception as e:
1040 verbose_proxy_logger.error("Error storing batch managed object: %s", e)
1042 @staticmethod
1043 def get_actual_model_id_from_router(model_name: str) -> str:
1044 from litellm.proxy.proxy_server import llm_router
1046 if llm_router is not None:
1047 # Try to find the model in the router by the extracted model name
1048 extracted_model_name = VertexPassthroughLoggingHandler.extract_model_name_from_vertex_path(model_name)
1050 # Use the existing get_model_ids method from router
1051 model_ids: Final = llm_router.get_model_ids(model_name=extracted_model_name)
1052 if model_ids and len(model_ids) > 0:
1053 # Use the first model ID found
1054 actual_model_id = model_ids[0]
1055 verbose_proxy_logger.info("Found model ID in router: %s", actual_model_id)
1056 return actual_model_id
1057 else:
1058 # Fallback to constructed model name
1059 actual_model_id = extracted_model_name
1060 verbose_proxy_logger.warning("Model not found in router, using constructed name: %s", actual_model_id)
1061 return actual_model_id
1062 else:
1063 # Fallback if router is not available
1064 extracted_model_name = VertexPassthroughLoggingHandler.extract_model_name_from_vertex_path(model_name)
1065 verbose_proxy_logger.warning("Router not available, using constructed model name: %s", extracted_model_name)
1066 return extracted_model_name