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

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 

7 

8import httpx 

9from pydantic import TypeAdapter 

10 

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) 

46 

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 

50 

51 from ..success_handler import PassThroughEndpointLogging 

52else: 

53 PassThroughEndpointLogging = Any 

54 LiteLLMBatch = Any 

55 

56_VERTEX_INTERACTIONS_PATH: Final = re.compile(r"/projects/[^/]+/locations/[^/]+/interactions/?$") 

57_INTERACTIONS_RESPONSE_BODY: Final = TypeAdapter(dict[str, object]) 

58 

59 

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 

71 

72 

73class VertexPassthroughLoggingHandler: 

74 @staticmethod 

75 def is_interactions_route(url_route: str) -> bool: 

76 return urlparse(url_route).path.rstrip("/").endswith("/interactions") 

77 

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 

81 

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} 

98 

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 } 

122 

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) 

151 

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 ) 

160 

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" 

165 

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 ) 

173 

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 

178 

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 

183 

184 return { 

185 "result": litellm_video_response, 

186 "kwargs": kwargs, 

187 } 

188 

189 elif "generateContent" in url_route: 

190 model = VertexPassthroughLoggingHandler.extract_model_from_url(url_route) 

191 

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 ) 

215 

216 return { 

217 "result": litellm_model_response, 

218 "kwargs": kwargs, 

219 } 

220 

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 ) 

240 

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 ) 

245 

246 _json_response: Final = httpx_response.json() 

247 

248 litellm_prediction_response = ModelResponse() 

249 

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 ) 

272 

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 ) 

283 

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 ) 

300 

301 standard_pass_through_response_object: Final[StandardPassThroughResponseObject] = { 

302 "response": cast(dict, litellm_vs_response), 

303 } 

304 

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 

310 

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 } 

331 

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 

347 

348 vertex_image_generation_class: Final = VertexImageGeneration() 

349 

350 model: Final = VertexPassthroughLoggingHandler.extract_model_from_url(url_route) 

351 

352 _json_response: Final[dict[str, object]] = httpx_response.json() 

353 

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 ) 

371 

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 

396 

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 ) 

407 

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 

412 

413 return { 

414 "result": litellm_prediction_response, 

415 "kwargs": kwargs, 

416 } 

417 

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 

429 

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 

443 

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 

449 

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 } 

457 

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 ) 

467 

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 ) 

480 

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) 

495 

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 ) 

511 

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 

515 

516 input_text: Final = VertexPassthroughLoggingHandler._extract_embed_content_input( 

517 request_body=request_body, batch=is_batch 

518 ) 

519 

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 ) 

535 

536 custom_llm_provider: Final = VertexPassthroughLoggingHandler._get_custom_llm_provider_from_url(url_route) 

537 

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 

543 

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 ) 

550 

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 

555 

556 return { 

557 "result": litellm_embedding_response, 

558 "kwargs": kwargs, 

559 } 

560 

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 

575 

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 ) 

591 

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 } 

600 

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 ) 

611 

612 return { 

613 "result": complete_streaming_response, 

614 "kwargs": kwargs, 

615 } 

616 

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 ) 

638 

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) 

658 

659 complete_streaming_response: Final = litellm.stream_chunk_builder(chunks=all_openai_chunks) 

660 

661 return complete_streaming_response 

662 

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" 

670 

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. 

675 

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 

679 

680 Args: 

681 vertex_model_path: The full Vertex AI model path 

682 

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] 

697 

698 # If no recognized pattern, return the original path 

699 return vertex_model_path 

700 

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 

713 

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 

720 

721 @staticmethod 

722 def _is_multimodal_embedding_response(json_response: dict) -> bool: 

723 """ 

724 Detect if the response is from a multimodal embedding request. 

725 

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 

728 

729 

730 Args: 

731 json_response: The JSON response from Vertex AI 

732 

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 

751 

752 return False 

753 

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) 

767 

768 """ 

769 

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 ) 

776 

777 kwargs["response_cost"] = response_cost 

778 kwargs["model"] = model 

779 

780 # pretty print standard logging object 

781 verbose_proxy_logger.debug("kwargs= %s", kwargs) 

782 

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 

789 

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 

806 

807 from litellm._uuid import uuid 

808 from litellm.llms.vertex_ai.batches.transformation import ( 

809 VertexAIBatchTransformation, 

810 ) 

811 

812 try: 

813 _json_response: Final = httpx_response.json() 

814 

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 ) 

823 

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") 

827 

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) 

831 

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("=") 

836 

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 ) 

848 

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()) 

855 

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 ] 

874 

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" 

882 

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 

887 

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()) 

899 

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 ] 

917 

918 kwargs["response_cost"] = 0.0 

919 kwargs["model"] = "vertex_ai_batch" 

920 kwargs["batch_job_state"] = "JOB_STATE_FAILED" 

921 

922 return { 

923 "result": litellm_model_response, 

924 "kwargs": kwargs, 

925 } 

926 

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()) 

935 

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 ] 

953 

954 kwargs["response_cost"] = 0.0 

955 kwargs["model"] = "vertex_ai_batch" 

956 kwargs["batch_job_state"] = "JOB_STATE_FAILED" 

957 

958 return { 

959 "result": litellm_model_response, 

960 "kwargs": kwargs, 

961 } 

962 

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. 

975 

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 

984 

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 

989 

990 _request_metadata: Final = (kwargs.get("litellm_params", {}) or {}).get("metadata", {}) or {} 

991 

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 ) 

1014 

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 ) 

1038 

1039 except Exception as e: 

1040 verbose_proxy_logger.error("Error storing batch managed object: %s", e) 

1041 

1042 @staticmethod 

1043 def get_actual_model_id_from_router(model_name: str) -> str: 

1044 from litellm.proxy.proxy_server import llm_router 

1045 

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) 

1049 

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