Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/utils.py: 42%

3509 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1import asyncio 

2import contextlib 

3import copy 

4import hashlib 

5import inspect 

6import json 

7import math 

8import os 

9import smtplib 

10import ssl 

11import sys 

12import threading 

13import time 

14import traceback 

15from collections.abc import ( 

16 AsyncGenerator, 

17 AsyncIterable, 

18 AsyncIterator, 

19 Awaitable, 

20 Callable, 

21 Coroutine, 

22 Mapping, 

23 Sequence, 

24) 

25from dataclasses import dataclass, field 

26from datetime import date, datetime, timedelta, timezone 

27from email.mime.multipart import MIMEMultipart 

28from email.mime.text import MIMEText 

29from functools import partial 

30from itertools import takewhile 

31from types import MappingProxyType 

32from typing import ( 

33 TYPE_CHECKING, 

34 Any, 

35 ClassVar, 

36 Final, 

37 Generic, 

38 Literal, 

39 NoReturn, 

40 Optional, 

41 Protocol, 

42 TypeAlias, 

43 TypeVar, 

44 Union, 

45 cast, 

46 overload, 

47) 

48 

49from typing_extensions import ReadOnly, TypedDict 

50 

51from litellm import _custom_logger_compatible_callbacks_literal 

52from litellm.constants import ( 

53 DEFAULT_MODEL_CREATED_AT_TIME, 

54 LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL, 

55 MAX_TEAM_LIST_LIMIT, 

56 PROXY_REJECTED_BEFORE_ROUTING_KEY, 

57 REDIS_SPEND_LOGS_BUFFER_DEQUEUE_COUNT, 

58 SPEND_LOG_QUEUE_MAX_BYTES, 

59 SPEND_LOG_WRITE_BATCH_MAX_BYTES, 

60 SPEND_LOG_WRITE_BATCH_MAX_ROWS, 

61) 

62from litellm.litellm_core_utils.bug_report import ( 

63 bug_report_notice, 

64 should_report_bug, 

65 strip_bug_report_notice, 

66) 

67from litellm.proxy._types import ( 

68 CommonProxyErrors, 

69 ProxyErrorTypes, 

70 ProxyException, 

71 SpendLogsMetadata, 

72 SpendLogsPayload, 

73) 

74from litellm.proxy.bug_report_config import build_proxy_bug_report 

75from litellm.proxy.common_utils.openai_error_payload import ( 

76 litellm_call_id_headers, 

77 openai_error_param, 

78 with_litellm_call_id, 

79) 

80from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error 

81from litellm.types.guardrails import GuardrailEventHooks 

82from litellm.types.proxy.model_listing import ModelInfoResponse 

83from litellm.types.utils import CallTypes, CallTypesLiteral, ModelInfo, Usage 

84 

85try: 

86 from litellm_enterprise.enterprise_callbacks.send_emails.base_email import ( 

87 BaseEmailLogger, 

88 ) 

89 from litellm_enterprise.enterprise_callbacks.send_emails.resend_email import ( 

90 ResendEmailLogger, 

91 ) 

92 from litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email import ( 

93 SendGridEmailLogger, 

94 ) 

95 from litellm_enterprise.enterprise_callbacks.send_emails.smtp_email import ( 

96 SMTPEmailLogger, 

97 ) 

98except ImportError: 

99 BaseEmailLogger = None 

100 SendGridEmailLogger = None 

101 SMTPEmailLogger = None 

102 ResendEmailLogger = None 

103 

104try: 

105 import backoff 

106except ImportError: 

107 raise ImportError("backoff is not installed. Please install it via 'pip install backoff'") 

108 

109from fastapi import HTTPException, status 

110from pydantic import TypeAdapter, ValidationError 

111 

112import litellm 

113import litellm.litellm_core_utils 

114import litellm.litellm_core_utils.litellm_logging 

115from litellm import ( 

116 EmbeddingResponse, 

117 ImageResponse, 

118 ModelResponse, 

119 ModelResponseStream, 

120 Router, 

121) 

122from litellm._logging import _redact_string, verbose_proxy_logger 

123from litellm._service_logger import ServiceLogging, ServiceTypes 

124from litellm.caching.caching import DualCache, RedisCache 

125from litellm.caching.dual_cache import LimitedSizeOrderedDict 

126from litellm.exceptions import ( 

127 GuardrailRaisedException, 

128 RejectedRequestError, 

129 SensitiveDataRouteException, 

130) 

131from litellm.integrations.custom_guardrail import ( 

132 CustomGuardrail, 

133 ModifyResponseException, 

134) 

135from litellm.integrations.custom_logger import CustomLogger 

136from litellm.integrations.prometheus import PrometheusLogger 

137from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting 

138from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert 

139from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route 

140from litellm.litellm_core_utils.core_helpers import ( 

141 coerce_token_limit, 

142 get_or_create_metadata_bucket, 

143 independent_snapshot, 

144 is_expected_client_error, 

145) 

146from litellm.litellm_core_utils.litellm_logging import Logging 

147from litellm.litellm_core_utils.safe_json_dumps import safe_dumps 

148from litellm.litellm_core_utils.safe_json_loads import safe_json_loads 

149from litellm.litellm_core_utils.served_output_texts import ( 

150 record_served_output_texts, 

151 served_stream_output_texts, 

152) 

153from litellm.litellm_core_utils.token_counter import offload_token_count 

154from litellm.llms import load_guardrail_translation_mappings 

155from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler 

156from litellm.proxy._types import ( 

157 AlertType, 

158 CallInfo, 

159 LiteLLM_VerificationTokenView, 

160 Member, 

161 UserAPIKeyAuth, 

162) 

163from litellm.proxy.agent_endpoints.auth.agent_access_groups import CeilingResolver, resolve_agent_access_group_ceiling 

164from litellm.proxy.auth.route_checks import RouteChecks 

165from litellm.proxy.common_utils.callback_utils import add_guardrail_to_applied_guardrails_header 

166from litellm.proxy.common_utils.config_sync_pubsub import publish_config_param_change 

167from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache 

168from litellm.proxy.db.create_views import ( 

169 create_missing_views, 

170 create_view_tolerating_race, 

171 should_create_missing_views, 

172) 

173from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter 

174from litellm.proxy.db.db_url_settings import ( 

175 DatabaseURLSettings, 

176 add_missing_query_params, 

177 token_refresh_params_from_url, 

178) 

179from litellm.proxy.db.exception_handler import ( 

180 PrismaDBExceptionHandler, 

181 call_with_db_reconnect_retry, 

182) 

183from litellm.proxy.db.health_check_latest import ( 

184 LatestHealthCheckRow, 

185 fetch_latest_health_checks, 

186 fetch_latest_health_checks_for_models, 

187) 

188from litellm.proxy.db.log_db_metrics import log_db_metrics 

189from litellm.proxy.db.pgbouncer import database_url_is_pooled 

190from litellm.proxy.db.prisma_client import ( 

191 PrismaWrapper, 

192 parse_iam_endpoint_from_url, 

193) 

194from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper 

195from litellm.proxy.db.spend_log_batching import ( 

196 spend_log_queue_within_budget, 

197 spend_log_row_bytes, 

198 spend_log_write_batches, 

199) 

200from litellm.proxy.db.token_auth import ( 

201 DatabaseTokenAuth, 

202 mint_database_token, 

203 resolve_database_token_auth, 

204) 

205from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( 

206 UnifiedLLMGuardrails, 

207 resolve_endpoint_translation, 

208) 

209from litellm.proxy.hooks import PROXY_HOOKS, get_proxy_hook 

210from litellm.proxy.hooks.cache_control_check import _PROXY_CacheControlCheck 

211from litellm.proxy.hooks.parallel_request_limiter import ( 

212 _PROXY_MaxParallelRequestsHandler, 

213) 

214from litellm.proxy.hooks.parallel_request_limiter_v3 import ( 

215 _PROXY_MaxParallelRequestsHandler_v3, 

216) 

217from litellm.proxy.hooks.sensitive_data_routing import ( 

218 _PROXY_SensitiveDataRoutingHandler, 

219) 

220from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_guardrails_from_auth_metadata 

221from litellm.proxy.management_helpers.key_settings_audit import with_settings_updated_at 

222from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor 

223from litellm.proxy.policy_engine.policy_registry import get_policy_registry 

224from litellm.proxy.policy_engine.policy_resolver import PolicyResolver 

225from litellm.repositories.budget_repository import BudgetRepository 

226from litellm.repositories.config_repository import ConfigRepository 

227from litellm.repositories.table_repositories import ( 

228 EndUserRepository, 

229 HealthCheckRepository, 

230 SpendLogsRepository, 

231 UserNotificationsRepository, 

232) 

233from litellm.repositories.team_repository import TeamRepository 

234from litellm.repositories.user_repository import UserRepository 

235from litellm.repositories.verification_token_repository import ( 

236 VerificationTokenRepository, 

237) 

238from litellm.router_utils.common_utils import resolve_model_group_alias 

239from litellm.secret_managers.main import str_to_bool 

240from litellm.types.integrations.slack_alerting import DEFAULT_ALERT_TYPES 

241from litellm.types.llms.openai import ResponsesAPIResponse 

242from litellm.types.mcp import ( 

243 MCPDuringCallResponseObject, 

244 MCPPreCallRequestObject, 

245 MCPPreCallResponseObject, 

246) 

247from litellm.types.proxy.policy_engine.pipeline_types import PipelineExecutionResult 

248from litellm.types.utils import LLMResponseTypes, LoggedLiteLLMParams 

249from litellm.utils import ( 

250 _add_custom_logger_callback_to_specific_event, # pyright: ignore[reportPrivateUsage] # only string-to-logger helper 

251) 

252 

253if TYPE_CHECKING: 253 ↛ 254line 253 didn't jump to line 254 because the condition on line 253 was never true

254 from mcp.types import CallToolResult 

255 from opentelemetry.trace import Span as _Span 

256 from prisma import models as prisma_models 

257 from prisma.actions import LiteLLM_DeprecatedVerificationTokenActions 

258 from prisma.client import TransactionManager 

259 from prisma.models import LiteLLM_DeprecatedVerificationToken 

260 from prisma.types import HttpConfig 

261 

262 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj 

263 from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation 

264 from litellm.models.team import LiteLLM_TeamTableCachedObj 

265 from litellm.proxy.db.autorouter_session_rollup import AutoRouterTurnTransaction 

266 from litellm.proxy.db.baseline_accounting import BaselineAccountingRecord 

267 from litellm.proxy.db.spend_log_tool_index import ToolUsageTransaction 

268 from litellm.repositories.prisma_protocols import TableActions 

269 from litellm.types.proxy.policy_engine.pipeline_types import GuardrailPipeline 

270 

271 Span = _Span | object 

272else: 

273 Span = Any 

274 

275_T: Final = TypeVar("_T") 

276 

277 

278class _ViewCountRow(TypedDict): 

279 view_count: ReadOnly[int] 

280 view_names: ReadOnly[Sequence[str] | None] 

281 

282 

283class _RelTuplesRow(TypedDict): 

284 reltuples: ReadOnly[int] 

285 

286 

287_VIEW_SETUP_POLL_INTERVAL_SECONDS: Final = 5.0 

288_VIEW_SETUP_DEADLINE_SECONDS: Final = 15 * 60.0 

289_VIEW_SETUP_GATE_TABLE: Final = '"LiteLLM_SpendLogs"' 

290_VIEW_SETUP_GATE_PROBE_ROWS: Final = TypeAdapter(tuple[Mapping[str, bool], ...]) 

291 

292_ViewSetupOutcome: TypeAlias = Literal["ready", "timed_out"] 

293_ViewSetupAttempt: TypeAlias = Literal["ready", "table_missing"] | Exception 

294 

295 

296class _EndUserBatchTable(Protocol): 

297 def upsert(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ... 297 ↛ exitline 297 didn't return from function 'upsert' because

298 

299 

300class _EndUserSpendBatch(Protocol): 

301 @property 

302 def litellm_endusertable(self) -> _EndUserBatchTable: ... 302 ↛ exitline 302 didn't return from function 'litellm_endusertable' because

303 

304 

305unified_guardrail: Final = UnifiedLLMGuardrails() 

306 

307NON_OPENAI_STREAM_GUARDRAIL_TRANSLATION_CALL_TYPES: "frozenset[CallTypes]" = frozenset({CallTypes.anthropic_messages}) 

308 

309 

310def print_verbose(print_statement: object): 

311 """ 

312 Prints the given `print_statement` to the console if `litellm.set_verbose` is True. 

313 Also logs the `print_statement` at the debug level using `verbose_proxy_logger`. 

314 

315 :param print_statement: The statement to be printed and logged. 

316 :type print_statement: Any 

317 """ 

318 import traceback 

319 

320 verbose_proxy_logger.debug("%s\n%s", print_statement, traceback.format_exc()) 

321 if litellm.set_verbose: 321 ↛ 322line 321 didn't jump to line 322 because the condition on line 321 was never true

322 print(f"LiteLLM Proxy: {_redact_string(str(print_statement))}") # noqa: T201 

323 

324 

325def _get_email_logger_class(): 

326 """ 

327 Determine which email logger class to use based on environment variables. 

328 Priority: SendGrid > Resend > SMTP > BaseEmailLogger (fallback) 

329 

330 Returns: 

331 The email logger class to use, or None if BaseEmailLogger is not available 

332 """ 

333 if BaseEmailLogger is None: 333 ↛ 334line 333 didn't jump to line 334 because the condition on line 333 was never true

334 return None 

335 

336 # Check for SendGrid API key 

337 if SendGridEmailLogger is not None and os.getenv("SENDGRID_API_KEY"): 337 ↛ 338line 337 didn't jump to line 338 because the condition on line 337 was never true

338 return SendGridEmailLogger 

339 

340 # Check for Resend API key 

341 if ResendEmailLogger is not None and os.getenv("RESEND_API_KEY"): 341 ↛ 342line 341 didn't jump to line 342 because the condition on line 341 was never true

342 return ResendEmailLogger 

343 

344 # Check for SMTP configuration 

345 if SMTPEmailLogger is not None and os.getenv("SMTP_HOST"): 345 ↛ 346line 345 didn't jump to line 346 because the condition on line 345 was never true

346 return SMTPEmailLogger 

347 

348 # Fallback to BaseEmailLogger (though it won't actually send emails) 

349 return BaseEmailLogger 

350 

351 

352class InternalUsageCache: 

353 def __init__(self, dual_cache: DualCache): 

354 self.dual_cache: DualCache = dual_cache 

355 

356 async def async_get_cache( 

357 self, 

358 key: str, 

359 litellm_parent_otel_span: Span | None, 

360 local_only: bool = False, 

361 **kwargs: object, 

362 ) -> Any: 

363 return await self.dual_cache.async_get_cache( 

364 key=key, 

365 local_only=local_only, 

366 parent_otel_span=litellm_parent_otel_span, 

367 **kwargs, 

368 ) 

369 

370 async def async_set_cache( 

371 self, 

372 key: str, 

373 value: object, 

374 litellm_parent_otel_span: Span | None, 

375 local_only: bool = False, 

376 **kwargs: object, 

377 ) -> None: 

378 return await self.dual_cache.async_set_cache( 

379 key=key, 

380 value=value, 

381 local_only=local_only, 

382 litellm_parent_otel_span=litellm_parent_otel_span, 

383 **kwargs, 

384 ) 

385 

386 async def async_batch_set_cache( 

387 self, 

388 cache_list: list[tuple[str, object]], 

389 litellm_parent_otel_span: Span | None, 

390 local_only: bool = False, 

391 **kwargs: object, 

392 ) -> None: 

393 return await self.dual_cache.async_set_cache_pipeline( 

394 cache_list=cache_list, 

395 local_only=local_only, 

396 litellm_parent_otel_span=litellm_parent_otel_span, 

397 **kwargs, 

398 ) 

399 

400 async def async_batch_get_cache( 

401 self, 

402 keys: Sequence[str | None], 

403 parent_otel_span: Span | None = None, 

404 local_only: bool = False, 

405 ): 

406 return await self.dual_cache.async_batch_get_cache( 

407 keys=list(keys), 

408 parent_otel_span=parent_otel_span, 

409 local_only=local_only, 

410 ) 

411 

412 async def async_increment_cache( 

413 self, 

414 key: str, 

415 value: float, 

416 litellm_parent_otel_span: Span | None, 

417 local_only: bool = False, 

418 **kwargs, 

419 ): 

420 return await self.dual_cache.async_increment_cache( 

421 key=key, 

422 value=value, 

423 local_only=local_only, 

424 parent_otel_span=litellm_parent_otel_span, 

425 **kwargs, 

426 ) 

427 

428 def set_cache( 

429 self, 

430 key: str, 

431 value: object, 

432 local_only: bool = False, 

433 **kwargs: object, 

434 ) -> None: 

435 return self.dual_cache.set_cache( 

436 key=key, 

437 value=value, 

438 local_only=local_only, 

439 **kwargs, 

440 ) 

441 

442 def get_cache( 

443 self, 

444 key: str, 

445 local_only: bool = False, 

446 **kwargs: object, 

447 ) -> Any: 

448 return self.dual_cache.get_cache( 

449 key=key, 

450 local_only=local_only, 

451 **kwargs, 

452 ) 

453 

454 

455### LOGGING ### 

456 

457# Cache for inspect.signature checks — avoids repeated introspection per request 

458_CALLBACK_ACCEPTS_CALL_INFO: Final[dict[int, bool]] = {} 

459 

460 

461def _accepts_litellm_call_info(cb: CustomLogger) -> bool: 

462 key: Final = id(type(cb)) 

463 if key not in _CALLBACK_ACCEPTS_CALL_INFO: 

464 sig: Final = inspect.signature(cb.async_post_call_response_headers_hook) 

465 _CALLBACK_ACCEPTS_CALL_INFO[key] = "litellm_call_info" in sig.parameters 

466 return _CALLBACK_ACCEPTS_CALL_INFO[key] 

467 

468 

469def _enrich_http_exception_with_guardrail_context(exc: BaseException, callback: object) -> None: 

470 """ 

471 If `exc` is an HTTPException with a dict `detail`, mutate it in place to 

472 add `guardrail_name` and `guardrail_mode` taken from the callback instance. 

473 

474 Uses setdefault so guardrails that already populate these fields explicitly 

475 win over the inferred defaults. No-op for non-HTTPException, non-dict-detail, 

476 or callbacks without `guardrail_name`. Never raises. 

477 """ 

478 if not isinstance(exc, HTTPException): 

479 return 

480 detail: Final = getattr(exc, "detail", None) 

481 if not isinstance(detail, dict): 

482 return 

483 guardrail_name: Final[object] = getattr(callback, "guardrail_name", None) 

484 if guardrail_name: 

485 detail.setdefault("guardrail_name", guardrail_name) 

486 event_hook: Final[object] = getattr(callback, "event_hook", None) 

487 if event_hook: 

488 detail.setdefault("guardrail_mode", event_hook) 

489 

490 

491def _record_raising_guardrail(request_data: Mapping[str, object], callback: object) -> None: 

492 guardrail_name: Final[object] = getattr(callback, "guardrail_name", None) 

493 if isinstance(request_data, dict) and isinstance(guardrail_name, str): 

494 add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=guardrail_name) 

495 

496 

497class _UpstreamStreamBoundary(Generic[_T]): 

498 __slots__ = ("_source", "_upstream", "failure") 

499 

500 def __init__(self, upstream: AsyncIterable[_T]) -> None: 

501 self._source: Final = upstream 

502 self._upstream: Final = upstream.__aiter__() 

503 self.failure: BaseException | None = None 

504 

505 def __getattr__(self, name: str) -> object: 

506 return getattr(self._source, name) 

507 

508 def __aiter__(self) -> "_UpstreamStreamBoundary[_T]": 

509 return self 

510 

511 async def __anext__(self) -> _T: 

512 try: 

513 return await self._upstream.__anext__() 

514 except StopAsyncIteration: 

515 raise 

516 except Exception as e: 

517 self.failure = e 

518 raise 

519 

520 

521class _StreamIteratorHook(Protocol[_T]): 

522 def __call__(self, *, response: AsyncIterator[_T]) -> AsyncGenerator[_T, None]: ... 522 ↛ exitline 522 didn't return from function '__call__' because

523 

524 

525def _is_client_error_exception(exc: Exception) -> bool: 

526 if isinstance(exc, HTTPException): 

527 return exc.status_code < 500 

528 if isinstance(exc, ProxyException): 

529 return not (exc.code.isdigit() and int(exc.code) >= 500) 

530 return False 

531 

532 

533def _exception_changes_request_flow(exc: BaseException) -> bool: 

534 """ 

535 True for guardrail exceptions the proxy turns into an alternate request flow 

536 (a reroute or a passthrough response) rather than a block. A pipeline step 

537 configured to block must honor that block, so these are surfaced as the 

538 generic pipeline block instead of being re-raised verbatim. 

539 """ 

540 return isinstance(exc, (SensitiveDataRouteException, ModifyResponseException)) 

541 

542 

543def _policy_state_metadata(data: Mapping[str, object]) -> Mapping[str, object]: 

544 """ 

545 Return the metadata bucket the policy engine wrote its pipeline state into. 

546 

547 The route decides the bucket (``litellm_metadata`` for ``/v1/messages``, 

548 responses, batches, files and bedrock, ``metadata`` everywhere else), and both 

549 buckets can be present at once because callers send their own provider-facing 

550 ``metadata`` (Claude Code sends ``metadata.user_id``) or their own 

551 ``litellm_metadata``. Pipeline slots are stripped from caller input before the 

552 policy engine runs, so whichever bucket carries them is the proxy's own write. 

553 """ 

554 return next( 

555 ( 

556 bucket 

557 for bucket in (data.get("metadata"), data.get("litellm_metadata")) 

558 if isinstance(bucket, dict) 

559 and ("_guardrail_pipelines" in bucket or "_pipeline_managed_guardrails" in bucket) 

560 ), 

561 {}, 

562 ) 

563 

564 

565def _policy_pipelines(data: Mapping[str, object]) -> tuple[tuple[str, "GuardrailPipeline"], ...]: 

566 pipelines: Final = _policy_state_metadata(data).get("_guardrail_pipelines") 

567 return ( 

568 tuple(cast("Sequence[tuple[str, GuardrailPipeline]]", pipelines)) # cast-ok: the policy engine wrote the slot 

569 if pipelines 

570 else () 

571 ) 

572 

573 

574def _pipeline_step_guardrail_names(pipelines: Sequence[tuple[str, "GuardrailPipeline"]]) -> frozenset[str]: 

575 return frozenset(step.guardrail for _policy_name, pipeline in pipelines for step in pipeline.steps) 

576 

577 

578def pipeline_managed_guardrail_names( 

579 data: Mapping[str, object], mode: Literal["pre_call", "post_call"] 

580) -> frozenset[str]: 

581 return _pipeline_step_guardrail_names( 

582 tuple((policy_name, pipeline) for policy_name, pipeline in _policy_pipelines(data) if pipeline.mode == mode) 

583 ) 

584 

585 

586def _partition_post_call_callbacks() -> tuple[tuple[CustomGuardrail, ...], tuple[CustomLogger, ...]]: 

587 resolved: Final = tuple( 

588 litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class( 

589 cast( # cast-ok: the resolver returns None for unknown names, filtered below 

590 _custom_logger_compatible_callbacks_literal, callback 

591 ) 

592 ) 

593 if isinstance(callback, str) 

594 else callback 

595 for callback in litellm.callbacks 

596 ) 

597 present: Final = tuple(callback for callback in resolved if callback is not None) 

598 guardrails: Final = tuple(callback for callback in present if isinstance(callback, CustomGuardrail)) 

599 others: Final = cast( # cast-ok: mirrors the legacy loop, which treated every non-guardrail entry as a CustomLogger 

600 "tuple[CustomLogger, ...]", 

601 tuple(callback for callback in present if not isinstance(callback, CustomGuardrail)), 

602 ) 

603 return (guardrails, others) 

604 

605 

606def _merge_pipeline_metadata_bucket(data: dict, bucket_key: str, modified_bucket_value: object) -> None: 

607 if not isinstance(modified_bucket_value, dict): 

608 return 

609 modified_bucket: Final = cast("dict[str, object]", modified_bucket_value) # cast-ok: metadata buckets are str-keyed 

610 surviving_writes: Final = {key: value for key, value in modified_bucket.items() if key != "guardrails"} 

611 existing_bucket: Final = data.get(bucket_key) 

612 if isinstance(existing_bucket, dict): 

613 cast("dict[str, object]", existing_bucket).update(surviving_writes) # cast-ok: metadata buckets are str-keyed 

614 else: 

615 data[bucket_key] = surviving_writes 

616 

617 

618def _merge_pipeline_metadata_writes(data: dict, modified_data: Mapping[str, object]) -> None: 

619 """ 

620 Copy metadata-bucket writes from a pipeline's working copy back onto the request. 

621 

622 Post_call pipelines run step hooks against a copied request dict so the payload 

623 already sent upstream stays untouched, but hooks record proxy-internal logging 

624 state in the metadata buckets (``applied_guardrails`` for response headers, 

625 ``standard_logging_guardrail_information`` for spend logs), and those writes 

626 must reach the request dict the proxy keeps reading after the pipeline returns. 

627 

628 The ``guardrails`` key is the executor's per-step activation flag for 

629 ``should_run_guardrail``, not a hook write, so it stays in the working copy. 

630 """ 

631 for bucket_key in ("metadata", "litellm_metadata"): 

632 _merge_pipeline_metadata_bucket(data, bucket_key, modified_data.get(bucket_key)) 

633 

634 

635def _pipeline_step_supports_streaming(guardrail_name: str, translation: "BaseTranslation | None") -> bool: 

636 callback: Final = PipelineExecutor.find_guardrail_callback(guardrail_name) 

637 if callback is None: 

638 return False 

639 if PipelineExecutor.supports_unified_execution(callback): 

640 return True 

641 return ( 

642 translation is not None 

643 and type(translation).assembles_streamed_response 

644 and PipelineExecutor.supports_streaming_execution(callback) 

645 ) 

646 

647 

648def _post_call_pipelines(data: Mapping[str, object]) -> tuple[tuple[str, "GuardrailPipeline"], ...]: 

649 return tuple( 

650 (policy_name, pipeline) for policy_name, pipeline in _policy_pipelines(data) if pipeline.mode == "post_call" 

651 ) 

652 

653 

654_PENDING_BACKGROUND_RESPONSE_STATUSES: Final = frozenset(("queued", "in_progress")) 

655 

656 

657def _is_pending_background_response(response: LLMResponseTypes) -> bool: 

658 return isinstance(response, ResponsesAPIResponse) and response.status in _PENDING_BACKGROUND_RESPONSE_STATUSES 

659 

660 

661def _guardrails_outside_pipeline(policy_name: str, pipeline: "GuardrailPipeline") -> frozenset[str]: 

662 resolved: Final = PolicyResolver.resolve_policy_guardrails( 

663 policy_name=policy_name, policies=get_policy_registry().get_all_policies() 

664 ) 

665 return frozenset(resolved.guardrails) - frozenset(step.guardrail for step in pipeline.steps) 

666 

667 

668def _guardrails_run_standalone_pre_call(data: Mapping[str, object]) -> frozenset[str]: 

669 return frozenset( 

670 callback.guardrail_name 

671 for callback in litellm.callbacks 

672 if isinstance(callback, CustomGuardrail) 

673 and callback.guardrail_name is not None 

674 and callback.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) 

675 ) 

676 

677 

678def _without_names( 

679 bucket: dict[str, object], # mutable-ok: the applied_* header slots live in the request-state dict hooks write 

680 slot: str, 

681 names: frozenset[str], 

682) -> None: 

683 claimed: Final = bucket.get(slot) 

684 if not isinstance(claimed, list): 

685 return 

686 remaining: Final = [ # mutable-ok: the slot stays a list, the shape every applied_* header writer appends to 

687 name for name in claimed if name not in names 

688 ] 

689 if remaining: 

690 bucket[slot] = remaining # rebind-ok: the slot lives in the shared request-state dict, rewritten in place 

691 else: 

692 bucket.pop(slot) 

693 

694 

695def _withdraw_deferred_claims( 

696 data: dict[str, object], # mutable-ok: same request-payload shape as post_call_success_hook's data 

697 deferred: Sequence[tuple[str, "GuardrailPipeline"]], 

698) -> None: 

699 outside_by_policy: Final = MappingProxyType( 

700 {policy_name: _guardrails_outside_pipeline(policy_name, pipeline) for policy_name, pipeline in deferred} 

701 ) 

702 running_elsewhere: Final = pipeline_managed_guardrail_names(data, "pre_call").union( 

703 _guardrails_run_standalone_pre_call(data), *outside_by_policy.values() 

704 ) 

705 withdrawn_policies: Final = frozenset(name for name, outside in outside_by_policy.items() if not outside) 

706 withdrawn_guardrails: Final = _pipeline_step_guardrail_names(deferred) - running_elsewhere 

707 _, bucket = get_or_create_metadata_bucket(data) 

708 _without_names(bucket, "applied_policies", withdrawn_policies) 

709 _without_names(bucket, "applied_guardrails", withdrawn_guardrails) 

710 sources: Final = bucket.get("policy_sources") 

711 if not isinstance(sources, dict): 

712 return 

713 remaining_sources: Final = { # mutable-ok: policy_sources stays a dict, the shape its writer updates in place 

714 name: reason for name, reason in sources.items() if name not in withdrawn_policies 

715 } 

716 if remaining_sources: 

717 bucket["policy_sources"] = remaining_sources 

718 else: 

719 bucket.pop("policy_sources") 

720 

721 

722def _defer_post_call_pipelines( 

723 data: dict[str, object], # mutable-ok: same request-payload shape as post_call_success_hook's data 

724 response: ResponsesAPIResponse, 

725) -> None: 

726 deferred: Final = _post_call_pipelines(data) 

727 if not deferred: 

728 return 

729 verbose_proxy_logger.debug( 

730 "Post_call guardrail pipelines wait for background response %s (status=%s) to be retrieved complete: %s", 

731 response.id, 

732 response.status, 

733 ", ".join(policy_name for policy_name, _pipeline in deferred), 

734 ) 

735 tag_matched: Final = _tag_matched_deferrals(data, deferred) 

736 if tag_matched: 

737 verbose_proxy_logger.warning( 

738 "Policy engine: background response %s matched post_call policies through a request tag at submit; " 

739 "retrieval re-matches only the key, team, and model scopes, so a tag carried in the request body " 

740 "does not govern the completed response: %s", 

741 response.id, 

742 ", ".join(tag_matched), 

743 ) 

744 body_selected: Final = _body_selected_deferrals(data, deferred) 

745 if body_selected: 

746 verbose_proxy_logger.warning( 

747 "Policy engine: background response %s matched post_call policies through the request body's policies " 

748 "list at submit; retrieval carries no request body, so those policies do not govern the completed " 

749 "response: %s", 

750 response.id, 

751 ", ".join(body_selected), 

752 ) 

753 _withdraw_deferred_claims(data, deferred) 

754 

755 

756def _tag_matched_deferrals( 

757 data: Mapping[str, object], deferred: Sequence[tuple[str, "GuardrailPipeline"]] 

758) -> tuple[str, ...]: 

759 sources: Final = _policy_state_metadata(data).get("policy_sources") 

760 if not isinstance(sources, dict): 

761 return () 

762 return tuple( 

763 policy_name 

764 for policy_name, _pipeline in deferred 

765 if policy_name in sources and "tag:" in str(sources[policy_name]) 

766 ) 

767 

768 

769def _body_selected_deferrals( 

770 data: Mapping[str, object], deferred: Sequence[tuple[str, "GuardrailPipeline"]] 

771) -> tuple[str, ...]: 

772 sources: Final = _policy_state_metadata(data).get("policy_sources") 

773 attributed: Final = frozenset(sources) if isinstance(sources, dict) else frozenset() 

774 return tuple(policy_name for policy_name, _pipeline in deferred if policy_name not in attributed) 

775 

776 

777def _pipeline_unsupported_streaming_guardrails( 

778 pipeline: "GuardrailPipeline", translation: "BaseTranslation | None" 

779) -> tuple[str, ...]: 

780 return tuple( 

781 dict.fromkeys( 

782 step.guardrail 

783 for step in pipeline.steps 

784 if not _pipeline_step_supports_streaming(step.guardrail, translation) 

785 ) 

786 ) 

787 

788 

789def _pipeline_is_streamable( 

790 policy_name: str, pipeline: "GuardrailPipeline", translation: "BaseTranslation | None" 

791) -> bool: 

792 unsupported: Final = _pipeline_unsupported_streaming_guardrails(pipeline, translation) 

793 if not unsupported: 

794 return True 

795 verbose_proxy_logger.warning( 

796 "Policy '%s' has post_call pipeline guardrails a streaming pipeline cannot run on this route yet; they " 

797 "need the unified apply_guardrail interface, or a post-call hook without a streaming iterator hook on a " 

798 "route whose translation assembles the streamed response. The stream skips the pipeline and its " 

799 "guardrails run on their own: %s", 

800 policy_name, 

801 ", ".join(unsupported), 

802 ) 

803 return False 

804 

805 

806def _streaming_pipeline_translation(user_api_key_dict: UserAPIKeyAuth) -> "BaseTranslation | None": 

807 resolved: Final = resolve_endpoint_translation(user_api_key_dict, None) 

808 return None if resolved is None else resolved[1] 

809 

810 

811def stream_gated_guardrail_names( 

812 request_data: Mapping[str, object], user_api_key_dict: UserAPIKeyAuth 

813) -> frozenset[str]: 

814 translation: Final = _streaming_pipeline_translation(user_api_key_dict) 

815 if translation is None: 

816 return frozenset() 

817 return _pipeline_step_guardrail_names( 

818 tuple( 

819 (policy_name, pipeline) 

820 for policy_name, pipeline in _post_call_pipelines(request_data) 

821 if not _pipeline_unsupported_streaming_guardrails(pipeline, translation) 

822 ) 

823 ) 

824 

825 

826def _streamable_post_call_pipelines( 

827 request_data: Mapping[str, object], user_api_key_dict: UserAPIKeyAuth 

828) -> tuple[tuple[str, "GuardrailPipeline"], ...]: 

829 """ 

830 The post_call pipelines a streaming response can be gated through. 

831 

832 Streaming pipelines scan the buffered stream through the endpoint guardrail 

833 translation of the request route, so every step's guardrail needs either the 

834 unified apply_guardrail interface or, on a route whose translation assembles 

835 the streamed response, a post-call hook that is its only streaming path, and 

836 the route needs a translation. A pipeline that 

837 cannot be run that way yet is left out and its guardrails run on the stream 

838 on their own, the way they did before pipelines ran on streams at all, with 

839 a warning naming the pipeline. 

840 """ 

841 post_call_pipelines: Final = _post_call_pipelines(request_data) 

842 if not post_call_pipelines: 

843 return () 

844 translation: Final = _streaming_pipeline_translation(user_api_key_dict) 

845 if translation is None: 

846 verbose_proxy_logger.warning( 

847 "Policies with post_call guardrail pipelines cannot scan streaming responses on route %s yet " 

848 "(no endpoint guardrail translation); the stream skips the pipelines and their guardrails run " 

849 "on their own: %s", 

850 user_api_key_dict.request_route, 

851 ", ".join(policy_name for policy_name, _pipeline in post_call_pipelines), 

852 ) 

853 return () 

854 return tuple( 

855 (policy_name, pipeline) 

856 for policy_name, pipeline in post_call_pipelines 

857 if _pipeline_is_streamable(policy_name, pipeline, translation) 

858 ) 

859 

860 

861def _prompt_block_text(block: object) -> str: 

862 if isinstance(block, str): 

863 return block 

864 if not isinstance(block, dict): 

865 return "" 

866 block_text: Final = block.get("text") 

867 return block_text if isinstance(block_text, str) else "" 

868 

869 

870def _system_prompt_text(system_input: object) -> str: 

871 if isinstance(system_input, str): 

872 return system_input 

873 if not isinstance(system_input, list): 

874 return "" 

875 return "".join(_prompt_block_text(block) for block in system_input) 

876 

877 

878def _count_request_input_tokens(model: str, request_input: object, system_input: object) -> int: 

879 system_text: Final = _system_prompt_text(system_input) 

880 system_tokens: Final = litellm.token_counter(model=model, text=system_text) if system_text else 0 

881 if isinstance(request_input, str): 

882 return system_tokens + litellm.token_counter(model=model, text=request_input) 

883 if not isinstance(request_input, list) or not request_input: 

884 return system_tokens 

885 text_entries: Final = tuple(entry for entry in request_input if isinstance(entry, str)) 

886 if len(text_entries) == len(request_input): 

887 return system_tokens + litellm.token_counter(model=model, text="".join(text_entries)) 

888 return system_tokens + litellm.token_counter( 

889 model=model, messages=request_input, use_default_image_token_count=True 

890 ) 

891 

892 

893def _estimate_dispatched_failure_usage(model: str, request_input: object, system_input: object) -> Usage | None: 

894 """A request that failed after dispatch consumed provider-billed input 

895 tokens, but no provider usage ever came back. Estimate the input side with 

896 the same tokenizer fallback interrupted streams use, so the spend log's 

897 failure row records what was sent instead of zero.""" 

898 try: 

899 input_tokens: Final = _count_request_input_tokens( 

900 model=model, request_input=request_input, system_input=system_input 

901 ) 

902 except Exception: 

903 return None 

904 if input_tokens <= 0: 

905 return None 

906 return Usage(prompt_tokens=input_tokens, completion_tokens=0, total_tokens=input_tokens) 

907 

908 

909_INPUT_ESTIMABLE_CALL_TYPES: Final = frozenset( 

910 call_type.value 

911 for call_type in ( 

912 CallTypes.completion, 

913 CallTypes.acompletion, 

914 CallTypes.text_completion, 

915 CallTypes.atext_completion, 

916 CallTypes.anthropic_messages, 

917 CallTypes.aanthropic_messages, 

918 CallTypes.responses, 

919 CallTypes.aresponses, 

920 CallTypes.embedding, 

921 CallTypes.aembedding, 

922 CallTypes.moderation, 

923 CallTypes.amoderation, 

924 CallTypes.image_generation, 

925 CallTypes.aimage_generation, 

926 CallTypes.speech, 

927 CallTypes.aspeech, 

928 CallTypes.rerank, 

929 CallTypes.arerank, 

930 CallTypes.generate_content, 

931 CallTypes.agenerate_content, 

932 CallTypes.generate_content_stream, 

933 CallTypes.agenerate_content_stream, 

934 ) 

935) 

936 

937 

938def _failure_usage_to_lift( 

939 model_call_details: Mapping[str, object], 

940 request_body: Mapping[str, object], 

941 dispatched: bool, 

942) -> tuple[object, object] | None: 

943 """A stream that broke mid-flight still billed the provider for the chunks 

944 already delivered; the streaming handler stashes that recovered usage and 

945 cost in model_call_details, so prefer it. Otherwise a request that was 

946 dispatched to a provider and failed without upstream usage gets an 

947 estimated input-side Usage with zero cost. The raw request body backfills 

948 the system prompt when the SDK bridges an endpoint (e.g. /v1/messages on a 

949 chat-completions provider) without filling optional_params. Returns the 

950 (combined_usage_object, response_cost) pair to lift, or None.""" 

951 recovered_usage: Final = model_call_details.get("combined_usage_object") 

952 if recovered_usage is not None: 952 ↛ 953line 952 didn't jump to line 953 because the condition on line 952 was never true

953 return recovered_usage, model_call_details.get("response_cost") 

954 if not dispatched or model_call_details.get(LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL): 

955 return None 

956 if str(model_call_details.get("call_type")) not in _INPUT_ESTIMABLE_CALL_TYPES: 956 ↛ 958line 956 didn't jump to line 958 because the condition on line 956 was always true

957 return None 

958 optional_params: Final = model_call_details.get("optional_params") 

959 dispatched_system: Final = ( 

960 (optional_params.get("system") or optional_params.get("instructions")) 

961 if isinstance(optional_params, dict) 

962 else None 

963 ) 

964 system_input: Final = dispatched_system or request_body.get("system") or request_body.get("instructions") 

965 estimated_usage: Final = _estimate_dispatched_failure_usage( 

966 model=str(model_call_details.get("model") or ""), 

967 request_input=model_call_details.get("messages"), 

968 system_input=system_input, 

969 ) 

970 if estimated_usage is None: 

971 return None 

972 return estimated_usage, 0.0 

973 

974 

975_EMPTY_LIFT: Final = MappingProxyType({}) 

976 

977 

978def _reached_deployment(litellm_logging_obj: Logging) -> bool: 

979 """A provider handoff or a cached response both mean the router selected a deployment.""" 

980 caching_details: Final = litellm_logging_obj.caching_details 

981 return litellm_logging_obj.model_call_details.get("first_api_call_start_time") is not None or ( 

982 caching_details is not None and caching_details.get("cache_hit") is True 

983 ) 

984 

985 

986def _stamp_deployment_attribution( 

987 litellm_params: dict[str, object], model_group: str | None, team_id: str | None, dispatched: bool 

988) -> Mapping[str, object]: 

989 """Stamp provider and logging-metadata attribution onto ``litellm_params`` and return it. 

990 ``litellm_params["model_info"]`` stays unset: the router's cooldown and per-deployment rpm 

991 callbacks key off it and must not count a proxy-side reject against the deployment. A failure 

992 after the provider handoff keeps the metadata the router stamped; a request that never reached a 

993 provider is flagged ``PROXY_REJECTED_BEFORE_ROUTING_KEY`` (deployment metrics key off it) whatever 

994 its metadata says, since ``metadata.model_info`` can be caller supplied.""" 

995 attribution: Final = _deployment_attribution_for_model_group(model_group, team_id) 

996 if "custom_llm_provider" in attribution: 996 ↛ 997line 996 didn't jump to line 997 because the condition on line 996 was never true

997 litellm_params["custom_llm_provider"] = attribution["custom_llm_provider"] 

998 if dispatched: 

999 return attribution 

1000 litellm_params[PROXY_REJECTED_BEFORE_ROUTING_KEY] = True 

1001 if "model_info" not in attribution: 1001 ↛ 1003line 1001 didn't jump to line 1003 because the condition on line 1001 was always true

1002 return attribution 

1003 if litellm_params.get("metadata") is None: 

1004 litellm_params["metadata"] = {} # mutable-ok: legacy logging payload is populated in place 

1005 metadata: Final = litellm_params["metadata"] 

1006 if not isinstance(metadata, dict): 

1007 return attribution 

1008 metadata.setdefault("model_info", attribution["model_info"]) 

1009 metadata.setdefault("deployment", attribution["deployment"]) 

1010 if isinstance(model_group, str): 

1011 metadata.setdefault("model_group", model_group) 

1012 return attribution 

1013 

1014 

1015def _deployment_attribution_for_model_group(model_group: object, team_id: str | None) -> Mapping[str, object]: 

1016 """Provider fields the router would have stamped had it reached a deployment: 

1017 ``custom_llm_provider`` when every deployment in the group resolves to the same 

1018 provider, plus ``model_info`` and ``deployment`` when the group has exactly one. 

1019 ``team_id`` picks the key's team deployments over a global group of the same public name.""" 

1020 if not isinstance(model_group, str): 

1021 return _EMPTY_LIFT 

1022 

1023 from litellm.proxy.proxy_server import llm_router 

1024 

1025 if llm_router is None: 1025 ↛ 1026line 1025 didn't jump to line 1026 because the condition on line 1025 was never true

1026 return _EMPTY_LIFT 

1027 deployments: Final = llm_router.get_model_list(model_name=model_group, team_id=team_id) 

1028 if not deployments: 1028 ↛ 1031line 1028 didn't jump to line 1031 because the condition on line 1028 was always true

1029 return _EMPTY_LIFT 

1030 

1031 def _provider_for_deployment(deployment: Mapping[str, object]) -> str | None: 

1032 litellm_params: Final = cast( # cast-ok: router deployment parameters are mapping-shaped 

1033 Mapping[str, object], deployment["litellm_params"] 

1034 ) 

1035 try: 

1036 provider: Final = litellm.get_llm_provider( 

1037 model=cast(str, litellm_params["model"]), # cast-ok: router deployment model is a string 

1038 custom_llm_provider=cast( # cast-ok: router deployment provider is optional 

1039 str | None, litellm_params.get("custom_llm_provider") 

1040 ), 

1041 )[1] 

1042 return cast(str | None, provider) # cast-ok: provider resolver returns an optional provider string 

1043 except Exception: # noqa: BLE001 # get_llm_provider raises for unmapped models 

1044 return None 

1045 

1046 providers: Final = frozenset(_provider_for_deployment(deployment) for deployment in deployments) 

1047 shared_provider: Final = next(iter(providers)) if len(providers) == 1 else None 

1048 single_deployment: Final = deployments[0] if len(deployments) == 1 else None 

1049 single_deployment_params: Final = ( 

1050 cast( # cast-ok: router deployment parameters are mapping-shaped 

1051 Mapping[str, object], single_deployment["litellm_params"] 

1052 ) 

1053 if single_deployment is not None 

1054 else None 

1055 ) 

1056 return MappingProxyType( 

1057 { 

1058 **({"custom_llm_provider": shared_provider} if shared_provider is not None else {}), 

1059 **( 

1060 { # mutable-ok: frozen immediately by the outer MappingProxyType 

1061 "model_info": dict( # mutable-ok: preserve the router's mutable model-info payload 

1062 single_deployment.get("model_info") or {} 

1063 ), 

1064 "deployment": single_deployment_params["model"], 

1065 } 

1066 if single_deployment is not None and single_deployment_params is not None 

1067 else {} # mutable-ok: frozen immediately by the outer MappingProxyType 

1068 ), 

1069 } 

1070 ) 

1071 

1072 

1073def _call_type_for_route(route: str | None) -> str | None: 

1074 """The route's call type when it maps to a single operation (its async and sync variants); 

1075 None for routes shared by several operations, since the method is not known here.""" 

1076 if route is None: 

1077 return None 

1078 call_types: Final = get_call_types_for_route(route) 

1079 if not call_types: 

1080 return None 

1081 operations: Final = frozenset(call_type.value.removeprefix("a") for call_type in call_types) 

1082 return call_types[0].value if len(operations) == 1 else None 

1083 

1084 

1085_PROXY_ONLY_LLM_API_ERRORS: Final = (HTTPException, ProxyException, GuardrailRaisedException) 

1086 

1087 

1088def _failure_fields_to_lift(request_data: Mapping[str, object]) -> Mapping[str, object]: 

1089 """Failure-path callbacks run after ``litellm_logging_obj`` is popped from 

1090 request_data (it is not serialisable), so the caller merges these fields 

1091 onto request_data first: the request start and first-handoff instants for 

1092 duration and preprocessing latency, the call type, recovered or estimated 

1093 usage for token counts, and the standard logging object for deployment 

1094 attribution on failed-request spend logs.""" 

1095 _logging_obj: Final = request_data.get("litellm_logging_obj") 

1096 if _logging_obj is None: 

1097 return _EMPTY_LIFT 

1098 _model_call_details: Final = getattr(_logging_obj, "model_call_details", {}) 

1099 _first_handoff: Final = _model_call_details.get("first_api_call_start_time") 

1100 _usage_to_lift: Final = _failure_usage_to_lift( 

1101 model_call_details=_model_call_details, 

1102 request_body=request_data, 

1103 dispatched=_first_handoff is not None, 

1104 ) 

1105 _entries: Final = ( 

1106 ("start_time", _model_call_details.get("start_time")), 

1107 ("first_api_call_start_time", _first_handoff), 

1108 ("call_type", _model_call_details.get("call_type")), 

1109 ("combined_usage_object", None if _usage_to_lift is None else _usage_to_lift[0]), 

1110 ("response_cost", None if _usage_to_lift is None else (_usage_to_lift[1] or 0.0)), 

1111 ("standard_logging_object", _model_call_details.get("standard_logging_object")), 

1112 ) 

1113 return MappingProxyType({key: value for key, value in _entries if value is not None}) 

1114 

1115 

1116@dataclass(frozen=True) 

1117class _CallbackCapabilities: 

1118 """Cached per-hook capability flags derived from ``litellm.callbacks``. 

1119 

1120 Recomputing this per request walked the callback list and resolved every 

1121 string entry via ``get_custom_logger_compatible_class`` — a measurable 

1122 chunk of overhead on streaming and non-streaming chat completions. 

1123 """ 

1124 

1125 has_post_call_response_headers: bool = False 

1126 has_iterator_override: bool = False 

1127 has_streaming_chunk_override: bool = False 

1128 has_guardrail: bool = False 

1129 has_pre_call_override: bool = False 

1130 has_content_enforcer: bool = False 

1131 has_moderation_override: bool = False 

1132 # Tuple[(resolved_callback, "override" | "apply_guardrail"), ...] 

1133 # Ordered the same as ``litellm.callbacks``; used to build the streaming 

1134 # iterator chain without re-scanning per request. 

1135 iterator_overrides: tuple[tuple[Any, str], ...] = field(default_factory=tuple) 

1136 # Resolved CustomLogger callbacks in original order. Pre-resolving once 

1137 # avoids the per-request ``get_custom_logger_compatible_class`` walk for 

1138 # every string entry in ``litellm.callbacks``. 

1139 resolved_callbacks: tuple[object, ...] = field(default_factory=tuple) 

1140 listed_models_filters: tuple[CustomLogger, ...] = field(default_factory=tuple) 

1141 

1142 

1143def _overrides_hook(callback: CustomLogger, hook_name: str) -> bool: 

1144 leaf_to_base: Final = takewhile(lambda klass: klass is not CustomLogger, type(callback).__mro__) 

1145 return any(hook_name in klass.__dict__ for klass in leaf_to_base) 

1146 

1147 

1148def _overrides_moderation_hook(callback: CustomLogger) -> bool: 

1149 return _overrides_hook(callback, "async_moderation_hook") 

1150 

1151 

1152_LISTED_MODEL_NAMES: Final = TypeAdapter(tuple[str, ...]) 

1153 

1154 

1155@dataclass(frozen=True, slots=True) 

1156class MalformedListingFilterReturn: 

1157 callback: str 

1158 tag: Literal["malformed_listing_filter_return"] = "malformed_listing_filter_return" 

1159 

1160 

1161async def _names_kept_by_listing_callbacks( 

1162 callbacks: Sequence[CustomLogger], 

1163 user_api_key_dict: UserAPIKeyAuth, 

1164 model_names: tuple[str, ...], 

1165) -> tuple[str, ...] | MalformedListingFilterReturn: 

1166 if not callbacks or not model_names: 

1167 return model_names 

1168 returned: Final = await callbacks[0].async_filter_listed_models(user_api_key_dict, model_names) 

1169 try: 

1170 kept: Final = frozenset(_LISTED_MODEL_NAMES.validate_python(returned)) 

1171 except ValidationError: 

1172 return MalformedListingFilterReturn(callback=type(callbacks[0]).__name__) 

1173 return await _names_kept_by_listing_callbacks( 

1174 callbacks[1:], user_api_key_dict, tuple(name for name in model_names if name in kept) 

1175 ) 

1176 

1177 

1178def _raise_malformed_listing_filter_return(error: MalformedListingFilterReturn) -> NoReturn: 

1179 raise ProxyException( 

1180 message=f"{error.callback}.async_filter_listed_models must return a sequence of model names", 

1181 type=ProxyErrorTypes.internal_server_error, 

1182 param=None, 

1183 code=500, 

1184 ) 

1185 

1186 

1187class ProxyLogging: 

1188 """ 

1189 Logging/Custom Handlers for proxy. 

1190 

1191 Implemented mainly to: 

1192 - log successful/failed db read/writes 

1193 - support the max parallel request integration 

1194 """ 

1195 

1196 def __init__( 

1197 self, 

1198 user_api_key_cache: UserApiKeyCache, 

1199 premium_user: bool = False, 

1200 ): 

1201 ## INITIALIZE LITELLM CALLBACKS ## 

1202 self.call_details: dict = {} 

1203 self.call_details["user_api_key_cache"] = user_api_key_cache 

1204 self.internal_usage_cache: InternalUsageCache = InternalUsageCache( 

1205 dual_cache=DualCache(default_in_memory_ttl=1) # ping redis cache every 1s 

1206 ) 

1207 self.max_parallel_request_limiter = _PROXY_MaxParallelRequestsHandler(self.internal_usage_cache) 

1208 self.cache_control_check = _PROXY_CacheControlCheck() 

1209 self.alerting: list[str] | None = None 

1210 self.alerting_threshold: float = 300 # default to 5 min. threshold 

1211 self.alert_types: list[AlertType] = DEFAULT_ALERT_TYPES 

1212 self.alert_to_webhook_url: dict | None = None 

1213 self.slack_alerting_instance: SlackAlerting = SlackAlerting( 

1214 alerting_threshold=self.alerting_threshold, 

1215 alerting=self.alerting, 

1216 internal_usage_cache=self.internal_usage_cache.dual_cache, 

1217 ) 

1218 self.email_logging_instance: Any | None = None 

1219 if BaseEmailLogger is not None: 1219 ↛ 1226line 1219 didn't jump to line 1226 because the condition on line 1219 was always true

1220 email_logger_class: Final = _get_email_logger_class() 

1221 if email_logger_class is not None: 1221 ↛ 1226line 1221 didn't jump to line 1226 because the condition on line 1221 was always true

1222 # All email logger classes now accept internal_usage_cache 

1223 self.email_logging_instance = email_logger_class( 

1224 internal_usage_cache=self.internal_usage_cache.dual_cache, 

1225 ) 

1226 self.premium_user = premium_user 

1227 self.service_logging_obj = ServiceLogging() 

1228 self.db_spend_update_writer = DBSpendUpdateWriter() 

1229 self.proxy_hook_mapping: dict[str, CustomLogger] = {} 

1230 

1231 # Guard flags to prevent duplicate background tasks 

1232 self.daily_report_started: bool = False 

1233 self.hanging_requests_check_started: bool = False 

1234 self.deprecation_check_started: bool = False 

1235 

1236 def startup_event( 

1237 self, 

1238 llm_router: Router | None, 

1239 redis_usage_cache: RedisCache | None, 

1240 ): 

1241 """Initialize logging and alerting on proxy startup""" 

1242 ## UPDATE SLACK ALERTING ## 

1243 self.slack_alerting_instance.update_values(llm_router=llm_router) 

1244 

1245 ## UPDATE INTERNAL USAGE CACHE ## 

1246 self.update_values( 

1247 redis_cache=redis_usage_cache 

1248 ) # used by parallel request limiter for rate limiting keys across instances 

1249 

1250 self._init_litellm_callbacks( 

1251 llm_router=llm_router 

1252 ) # INITIALIZE LITELLM CALLBACKS ON SERVER STARTUP <- do this to catch any logging errors on startup, not when calls are being made 

1253 

1254 if ( 1254 ↛ 1267line 1254 didn't jump to line 1267 because the condition on line 1254 was always true

1255 self.slack_alerting_instance is not None 

1256 and "daily_reports" in self.slack_alerting_instance.alert_types 

1257 and not self.daily_report_started 

1258 ): 

1259 asyncio.create_task( 

1260 self.slack_alerting_instance._run_scheduled_daily_report( 

1261 llm_router=llm_router, 

1262 pod_lock_manager=self.db_spend_update_writer.pod_lock_manager, 

1263 ) 

1264 ) # RUN DAILY REPORT (if scheduled) 

1265 self.daily_report_started = True 

1266 

1267 if ( 1267 ↛ 1277line 1267 didn't jump to line 1277 because the condition on line 1267 was always true

1268 self.slack_alerting_instance is not None 

1269 and AlertType.llm_requests_hanging in self.slack_alerting_instance.alert_types 

1270 and not self.hanging_requests_check_started 

1271 ): 

1272 asyncio.create_task( 

1273 self.slack_alerting_instance.hanging_request_check.check_for_hanging_requests() 

1274 ) # RUN HANGING REQUEST CHECK (if user wants to alert on hanging requests) 

1275 self.hanging_requests_check_started = True 

1276 

1277 self._ensure_deprecation_check_scheduled() 

1278 

1279 def _ensure_deprecation_check_scheduled(self) -> None: 

1280 """Alerting can be configured at startup or by a later config reload, so schedule from either path""" 

1281 if self.alerting is None or self.deprecation_check_started: 1281 ↛ 1284line 1281 didn't jump to line 1284 because the condition on line 1281 was always true

1282 return 

1283 

1284 try: 

1285 asyncio.get_running_loop() 

1286 except RuntimeError: 

1287 return 

1288 

1289 asyncio.create_task( 

1290 self.slack_alerting_instance.run_scheduled_deprecation_check( 

1291 pod_lock_manager=self.db_spend_update_writer.pod_lock_manager 

1292 ) 

1293 ) 

1294 self.deprecation_check_started = True 

1295 

1296 def update_values( 

1297 self, 

1298 alerting: list | None = None, 

1299 alerting_threshold: float | None = None, 

1300 redis_cache: RedisCache | None = None, 

1301 alert_types: list[AlertType] | None = None, 

1302 alerting_args: dict | None = None, 

1303 alert_to_webhook_url: dict | None = None, 

1304 alert_type_config: dict | None = None, 

1305 ): 

1306 updated_slack_alerting: bool = False 

1307 if alerting is not None: 1307 ↛ 1308line 1307 didn't jump to line 1308 because the condition on line 1307 was never true

1308 self.alerting = alerting 

1309 updated_slack_alerting = True 

1310 if alerting_threshold is not None: 1310 ↛ 1311line 1310 didn't jump to line 1311 because the condition on line 1310 was never true

1311 self.alerting_threshold = alerting_threshold 

1312 updated_slack_alerting = True 

1313 if alert_types is not None: 1313 ↛ 1314line 1313 didn't jump to line 1314 because the condition on line 1313 was never true

1314 self.alert_types = alert_types 

1315 updated_slack_alerting = True 

1316 if alert_to_webhook_url is not None: 1316 ↛ 1317line 1316 didn't jump to line 1317 because the condition on line 1316 was never true

1317 self.alert_to_webhook_url = alert_to_webhook_url 

1318 updated_slack_alerting = True 

1319 if alert_type_config is not None: 1319 ↛ 1320line 1319 didn't jump to line 1320 because the condition on line 1319 was never true

1320 updated_slack_alerting = True 

1321 

1322 if updated_slack_alerting is True: 1322 ↛ 1323line 1322 didn't jump to line 1323 because the condition on line 1322 was never true

1323 self._ensure_deprecation_check_scheduled() 

1324 self.slack_alerting_instance.update_values( 

1325 alerting=self.alerting, 

1326 alerting_threshold=self.alerting_threshold, 

1327 alert_types=self.alert_types, 

1328 alerting_args=alerting_args, 

1329 alert_to_webhook_url=self.alert_to_webhook_url, 

1330 alert_type_config=alert_type_config, 

1331 ) 

1332 

1333 if self.alerting is not None and ("slack" in self.alerting or "ms_teams" in self.alerting): 

1334 # NOTE: ENSURE we only add callbacks when alerting is on 

1335 # We should NOT add callbacks when alerting is off 

1336 if ( 

1337 "daily_reports" in self.alert_types 

1338 or "outage_alerts" in self.alert_types 

1339 or "region_outage_alerts" in self.alert_types 

1340 ): 

1341 litellm.logging_callback_manager.add_litellm_callback(self.slack_alerting_instance) 

1342 litellm.logging_callback_manager.add_litellm_success_callback( 

1343 self.slack_alerting_instance.response_taking_too_long_callback 

1344 ) 

1345 

1346 if redis_cache is not None: 1346 ↛ 1347line 1346 didn't jump to line 1347 because the condition on line 1346 was never true

1347 self.internal_usage_cache.dual_cache.redis_cache = redis_cache 

1348 self.db_spend_update_writer.redis_update_buffer.redis_cache = redis_cache 

1349 self.db_spend_update_writer.pod_lock_manager.redis_cache = redis_cache 

1350 

1351 def _add_proxy_hooks(self, llm_router: Router | None = None): 

1352 """ 

1353 Add proxy hooks to litellm.callbacks 

1354 """ 

1355 from litellm.proxy.proxy_server import prisma_client 

1356 

1357 for hook in PROXY_HOOKS: 

1358 proxy_hook = get_proxy_hook(hook) 

1359 expected_args = inspect.getfullargspec(proxy_hook).args 

1360 if "prisma_client" in expected_args and prisma_client is None: 1360 ↛ 1361line 1360 didn't jump to line 1361 because the condition on line 1360 was never true

1361 verbose_proxy_logger.debug( 

1362 "Skipping proxy hook %s: it requires a database and no prisma client is configured", hook 

1363 ) 

1364 continue 

1365 passed_in_args: dict[str, Any] = {} 

1366 if "internal_usage_cache" in expected_args: 

1367 passed_in_args["internal_usage_cache"] = self.internal_usage_cache 

1368 if "prisma_client" in expected_args: 

1369 passed_in_args["prisma_client"] = prisma_client 

1370 proxy_hook_obj = cast(CustomLogger, proxy_hook(**passed_in_args)) 

1371 litellm.logging_callback_manager.add_litellm_callback(proxy_hook_obj) 

1372 

1373 self.proxy_hook_mapping[hook] = proxy_hook_obj 

1374 

1375 def get_proxy_hook(self, hook: str) -> CustomLogger | None: 

1376 """ 

1377 Get a proxy hook from the proxy_hook_mapping 

1378 """ 

1379 return self.proxy_hook_mapping.get(hook) 

1380 

1381 def _init_litellm_callbacks(self, llm_router: Router | None = None): 

1382 self._add_proxy_hooks(llm_router) 

1383 litellm.logging_callback_manager.add_litellm_callback(self.service_logging_obj) 

1384 

1385 # Track string callbacks and their initialized instances so we can 

1386 # replace them in-place, preventing duplicates (string + instance) in 

1387 # litellm.callbacks which caused double-counting of metrics. 

1388 string_callbacks_to_replace: Final[dict[int, CustomLogger]] = {} 

1389 

1390 for idx, callback in enumerate(litellm.callbacks): 

1391 if isinstance(callback, str): 1391 ↛ 1392line 1391 didn't jump to line 1392 because the condition on line 1391 was never true

1392 initialized_callback = litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class( 

1393 cast(_custom_logger_compatible_callbacks_literal, callback), 

1394 internal_usage_cache=self.internal_usage_cache.dual_cache, 

1395 llm_router=llm_router, 

1396 ) 

1397 

1398 if initialized_callback is not None: 

1399 string_callbacks_to_replace[idx] = initialized_callback 

1400 

1401 # Replace string entries in litellm.callbacks with initialized instances 

1402 for idx, initialized_callback in string_callbacks_to_replace.items(): 1402 ↛ 1403line 1402 didn't jump to line 1403 because the loop on line 1402 never started

1403 litellm.callbacks[idx] = initialized_callback 

1404 

1405 # Fan ``litellm.callbacks`` (the "all events" registry) out into the 

1406 # success/failure event lists eagerly, at startup. ``completion()`` does 

1407 # this lazily in ``function_setup`` on the first call, but request paths 

1408 # that build their own logging object and never run ``function_setup`` — 

1409 # notably pass-through endpoints — read ``litellm._async_success_callback`` 

1410 # directly. Without this, a config-registered logger (e.g. ``otel``) is 

1411 # invisible to pass-through traffic until some other request warms the 

1412 # global lists. The manager dedupes, so this is idempotent with 

1413 # ``function_setup``. 

1414 for callback in litellm.callbacks: 

1415 if isinstance(callback, CustomLogger): 1415 ↛ 1414line 1415 didn't jump to line 1414 because the condition on line 1415 was always true

1416 litellm.logging_callback_manager.add_litellm_success_callback(callback) 

1417 litellm.logging_callback_manager.add_litellm_failure_callback(callback) 

1418 litellm.logging_callback_manager.add_litellm_async_success_callback(callback) 

1419 litellm.logging_callback_manager.add_litellm_async_failure_callback(callback) 

1420 

1421 # Runs after load_config applied every litellm_settings key: logger __init__s read e.g. s3_callback_params 

1422 success_callbacks: Final = tuple(cb for cb in litellm.success_callback if isinstance(cb, str)) 

1423 failure_callbacks: Final = tuple(cb for cb in litellm.failure_callback if isinstance(cb, str)) 

1424 for callback in success_callbacks: 1424 ↛ 1425line 1424 didn't jump to line 1425 because the loop on line 1424 never started

1425 _add_custom_logger_callback_to_specific_event(callback, "success") 

1426 for callback in failure_callbacks: 1426 ↛ 1427line 1426 didn't jump to line 1427 because the loop on line 1426 never started

1427 _add_custom_logger_callback_to_specific_event(callback, "failure") 

1428 

1429 async def update_request_status(self, litellm_call_id: str, status: Literal["success", "fail"]): 

1430 # only use this if slack alerting is being used 

1431 if self.alerting is None: 1431 ↛ 1435line 1431 didn't jump to line 1435 because the condition on line 1431 was always true

1432 return 

1433 

1434 # current alerting threshold 

1435 alerting_threshold: float = self.alerting_threshold 

1436 

1437 # add a 100 second buffer to the alerting threshold 

1438 # ensures we don't send errant hanging request slack alerts 

1439 alerting_threshold += 100 

1440 

1441 await self.internal_usage_cache.async_set_cache( 

1442 key=f"request_status:{litellm_call_id}", 

1443 value=status, 

1444 local_only=True, 

1445 ttl=alerting_threshold, 

1446 litellm_parent_otel_span=None, 

1447 ) 

1448 

1449 def _convert_user_api_key_auth_to_dict(self, user_api_key_auth_obj): 

1450 """ 

1451 Helper function to convert UserAPIKeyAuth object to dictionary. 

1452 Handles both Pydantic models and regular objects. 

1453 """ 

1454 if user_api_key_auth_obj is not None: 

1455 if hasattr(user_api_key_auth_obj, "model_dump"): 

1456 # If it's a Pydantic model, convert to dict 

1457 return user_api_key_auth_obj.model_dump() 

1458 elif hasattr(user_api_key_auth_obj, "__dict__"): 

1459 # If it's a regular object, convert to dict 

1460 return user_api_key_auth_obj.__dict__ 

1461 return {} 

1462 

1463 def _convert_mcp_to_llm_format(self, request_obj, kwargs: dict) -> dict: 

1464 """ 

1465 Convert MCP tool call to LLM message format for existing guardrail validation. 

1466 """ 

1467 from litellm.types.llms.openai import ChatCompletionUserMessage 

1468 

1469 guardrail_context: Final = TypeAdapter(Mapping[str, object]).validate_python( 

1470 kwargs.get("guardrail_context") or MappingProxyType({}) 

1471 ) 

1472 

1473 parent_metadata: Final = copy.deepcopy( 

1474 TypeAdapter(dict[str, object]).validate_python(guardrail_context.get("metadata") or MappingProxyType({})) 

1475 ) 

1476 

1477 # Create a synthetic message that represents the tool call 

1478 tool_call_content: Final = f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}" 

1479 

1480 synthetic_message: Final = ChatCompletionUserMessage(role="user", content=tool_call_content) 

1481 

1482 synthetic_metadata: Final[dict[str, object]] = { # mutable-ok: existing guardrail hooks mutate request metadata 

1483 **MappingProxyType({key: value for key, value in parent_metadata.items() if key != "guardrails"}), 

1484 "headers": kwargs.get("headers") or {}, 

1485 "user_api_key_user_id": kwargs.get("user_api_key_user_id"), 

1486 "user_api_key_team_id": kwargs.get("user_api_key_team_id"), 

1487 "user_api_key_end_user_id": kwargs.get("user_api_key_end_user_id"), 

1488 } 

1489 

1490 # Create synthetic LLM data that guardrails can process 

1491 synthetic_data: Final = { 

1492 "messages": [synthetic_message], 

1493 "model": guardrail_context.get("model", kwargs.get("model", "mcp-tool-call")), 

1494 "user_api_key_user_id": kwargs.get("user_api_key_user_id"), 

1495 "user_api_key_team_id": kwargs.get("user_api_key_team_id"), 

1496 "user_api_key_end_user_id": kwargs.get("user_api_key_end_user_id"), 

1497 "user_api_key_hash": kwargs.get("user_api_key_hash"), 

1498 "user_api_key_request_route": kwargs.get("user_api_key_request_route"), 

1499 "mcp_tool_name": request_obj.tool_name, # Keep original for reference 

1500 "mcp_arguments": request_obj.arguments, # Keep original for reference 

1501 # Surface the per-MCP-server rate-limit identity so the 

1502 # ParallelRequestLimiterV3 hook can apply mcp_rpm_limit on the 

1503 # synthetic call_mcp_tool payload (otherwise a key with 

1504 # mcp_rpm_limit could exceed it via the MCP path). 

1505 "mcp_server_name": kwargs.get("mcp_rate_limit_server_name"), 

1506 # Raw Bearer token from the original HTTP request — allows guardrails 

1507 # (e.g. MCPJWTSigner) to independently verify the caller's identity 

1508 # before re-signing an outbound token (FR-5 verify+re-sign). 

1509 "incoming_bearer_token": kwargs.get("incoming_bearer_token"), 

1510 "metadata": synthetic_metadata, 

1511 } 

1512 user_api_key_auth: Final = kwargs.get("user_api_key_auth") 

1513 if isinstance(user_api_key_auth, UserAPIKeyAuth): 

1514 add_guardrails_from_auth_metadata( 

1515 user_api_key_dict=user_api_key_auth, 

1516 data=synthetic_data, 

1517 metadata_variable_name="metadata", 

1518 ) 

1519 synthetic_metadata["user_api_key_metadata"] = copy.deepcopy(user_api_key_auth.metadata) 

1520 synthetic_metadata["user_api_key_team_metadata"] = copy.deepcopy(user_api_key_auth.team_metadata) 

1521 merged_guardrails: Final = ( 

1522 *TypeAdapter(tuple[object, ...]).validate_python(synthetic_metadata.get("guardrails") or ()), 

1523 *TypeAdapter(tuple[object, ...]).validate_python(parent_metadata.get("guardrails") or ()), 

1524 ) 

1525 synthetic_metadata["guardrails"] = [ # mutable-ok: existing guardrail selection and policy hooks require a list 

1526 selection for index, selection in enumerate(merged_guardrails) if selection not in merged_guardrails[:index] 

1527 ] 

1528 return synthetic_data 

1529 

1530 def _convert_llm_result_to_mcp_response(self, llm_result, request_obj) -> MCPPreCallResponseObject | None: 

1531 """ 

1532 Convert LLM guardrail result back to MCP response format. 

1533 """ 

1534 from litellm.types.mcp import MCPPreCallResponseObject 

1535 

1536 # If result is an exception, it means the guardrail blocked the request 

1537 if isinstance(llm_result, Exception): 

1538 return MCPPreCallResponseObject( 

1539 should_proceed=False, 

1540 error_message=str(llm_result), 

1541 modified_arguments=None, 

1542 ) 

1543 

1544 # If result is a dict with modified messages, check for content filtering 

1545 if isinstance(llm_result, dict): 

1546 modified_messages: Final = llm_result.get("messages") 

1547 if modified_messages: 

1548 # Check if content was blocked/modified 

1549 original_content: Final = f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}" 

1550 new_content: Final = modified_messages[0].get("content", "") if modified_messages else "" 

1551 

1552 if new_content != original_content: 

1553 # Content was modified - could be masking, redaction, or blocking 

1554 if not new_content or "blocked" in new_content.lower() or "violation" in new_content.lower(): 

1555 # Content was blocked completely 

1556 return MCPPreCallResponseObject( 

1557 should_proceed=False, 

1558 error_message="Content blocked by guardrail", 

1559 modified_arguments=None, 

1560 ) 

1561 else: 

1562 # Content was masked/redacted - extract the modified arguments 

1563 try: 

1564 # Try to parse the modified arguments from the masked content 

1565 modified_args = self._extract_modified_arguments_from_content(new_content, request_obj) 

1566 if modified_args is not None: 

1567 # Return the masked/redacted arguments for the MCP call to use 

1568 return MCPPreCallResponseObject( 

1569 should_proceed=True, 

1570 error_message=None, 

1571 modified_arguments=modified_args, 

1572 ) 

1573 else: 

1574 # Could not parse modified arguments, allow original call but warn 

1575 verbose_proxy_logger.warning( 

1576 "Could not parse modified arguments from guardrail response: %s", new_content 

1577 ) 

1578 return None 

1579 except Exception as e: 

1580 verbose_proxy_logger.error("Error parsing modified arguments: %s", e) 

1581 # Fallback: allow original call 

1582 return None 

1583 

1584 # If result is a string, it's likely an error message 

1585 if isinstance(llm_result, str): 

1586 return MCPPreCallResponseObject(should_proceed=False, error_message=llm_result, modified_arguments=None) 

1587 

1588 return None 

1589 

1590 def _extract_modified_arguments_from_content(self, masked_content: str, request_obj) -> dict | None: 

1591 """ 

1592 Extract modified/masked arguments from the guardrail response content. 

1593 """ 

1594 import json 

1595 

1596 verbose_proxy_logger.debug("Extracting modified args from content: %s", masked_content) 

1597 

1598 try: 

1599 # The format should be: "Tool: <tool_name>\nArguments: <json_arguments>" 

1600 # Parse the arguments section 

1601 lines: Final = masked_content.strip().split("\n") 

1602 for i, line in enumerate(lines): 

1603 if line.startswith("Arguments:"): 

1604 # Get the arguments part - everything after "Arguments: " 

1605 args_text = line[len("Arguments:") :].strip() 

1606 

1607 verbose_proxy_logger.debug("Found arguments text: %s", args_text) 

1608 

1609 # Try to parse as JSON first 

1610 try: 

1611 modified_args = json.loads(args_text) 

1612 verbose_proxy_logger.debug("Successfully parsed JSON args: %s", modified_args) 

1613 return modified_args 

1614 except json.JSONDecodeError as e: 

1615 # If JSON parsing fails, try to extract key-value pairs manually 

1616 verbose_proxy_logger.debug("Failed to parse JSON arguments: %s, error: %s", args_text, e) 

1617 return self._parse_arguments_manually(args_text, request_obj.arguments) 

1618 

1619 # If we can't find the Arguments: line, return None 

1620 verbose_proxy_logger.warning("Could not find 'Arguments:' line in masked content") 

1621 return None 

1622 

1623 except Exception as e: 

1624 verbose_proxy_logger.error("Error extracting modified arguments: %s", e) 

1625 return None 

1626 

1627 def _parse_arguments_manually(self, args_text: str, original_args: dict) -> dict | None: 

1628 """ 

1629 Try to manually parse arguments when JSON parsing fails. 

1630 This is a fallback for cases where the guardrail modifies the format. 

1631 """ 

1632 import re 

1633 

1634 try: 

1635 # Start with original arguments and try to apply modifications 

1636 modified_args: Final = original_args.copy() 

1637 

1638 # Look for simple key-value patterns 

1639 # This is a basic implementation - can be enhanced based on specific guardrail formats 

1640 for key, original_value in original_args.items(): 

1641 if isinstance(original_value, str): 

1642 # Look for the key in the masked content and try to extract its value 

1643 pattern = rf"['\"]?{re.escape(key)}['\"]?\s*:\s*['\"]?([^,'\"]*)['\"]?" 

1644 match = re.search(pattern, args_text, re.IGNORECASE) 

1645 if match: 

1646 new_value = match.group(1).strip() 

1647 if new_value: 

1648 modified_args[key] = new_value 

1649 

1650 return modified_args 

1651 

1652 except Exception as e: 

1653 verbose_proxy_logger.error("Error in manual argument parsing: %s", e) 

1654 return None 

1655 

1656 def _convert_llm_result_to_mcp_during_response(self, llm_result, request_obj) -> MCPDuringCallResponseObject | None: 

1657 """ 

1658 Convert LLM guardrail result back to MCP during call response format. 

1659 """ 

1660 # If result is an exception, it means the guardrail wants to stop execution 

1661 if isinstance(llm_result, Exception): 

1662 return MCPDuringCallResponseObject(should_continue=False, error_message=str(llm_result)) 

1663 

1664 # If result is a dict with modified messages, check for content filtering 

1665 if isinstance(llm_result, dict): 

1666 modified_messages: Final = llm_result.get("messages") 

1667 if modified_messages: 

1668 # Check if content was blocked/modified 

1669 original_content: Final = f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}" 

1670 new_content: Final = modified_messages[0].get("content", "") if modified_messages else "" 

1671 

1672 if new_content != original_content: 

1673 # Content was modified, could be masking or blocking 

1674 if not new_content or "blocked" in new_content.lower(): 

1675 # Content was blocked 

1676 return MCPDuringCallResponseObject( 

1677 should_continue=False, 

1678 error_message="Content blocked by guardrail during execution", 

1679 ) 

1680 else: 

1681 # Content was masked/modified - for now, stop execution 

1682 return MCPDuringCallResponseObject( 

1683 should_continue=False, 

1684 error_message="Content modified by guardrail during execution", 

1685 ) 

1686 

1687 # If result is a string, it's likely an error message 

1688 if isinstance(llm_result, str): 

1689 return MCPDuringCallResponseObject(should_continue=False, error_message=llm_result) 

1690 

1691 return None 

1692 

1693 def get_combined_callback_list(self, dynamic_success_callbacks: list | None, global_callbacks: list) -> list: 

1694 if dynamic_success_callbacks is None: 

1695 return list(global_callbacks) 

1696 return list(dict.fromkeys(dynamic_success_callbacks + global_callbacks)) 

1697 

1698 def _parse_pre_mcp_call_hook_response( 

1699 self, 

1700 response: MCPPreCallResponseObject, 

1701 original_request: MCPPreCallRequestObject, 

1702 ) -> Mapping[str, object]: 

1703 """ 

1704 Parse the response from the pre_mcp_tool_call_hook 

1705 

1706 1. Check if the call should proceed 

1707 2. Apply any argument modifications 

1708 3. Handle validation errors 

1709 """ 

1710 result: Final = { 

1711 "should_proceed": response.should_proceed, 

1712 "modified_arguments": response.modified_arguments or original_request.arguments, 

1713 "error_message": response.error_message, 

1714 "hidden_params": response.hidden_params, 

1715 } 

1716 return result 

1717 

1718 def _create_mcp_request_object_from_kwargs(self, kwargs: dict) -> "MCPPreCallRequestObject": 

1719 """ 

1720 Helper function to create MCPPreCallRequestObject from kwargs for standard pre_call_hook. 

1721 """ 

1722 from litellm.types.llms.base import HiddenParams 

1723 from litellm.types.mcp import MCPPreCallRequestObject 

1724 

1725 user_api_key_auth_dict: Final = self._convert_user_api_key_auth_to_dict(kwargs.get("user_api_key_auth")) 

1726 

1727 return MCPPreCallRequestObject( 

1728 tool_name=kwargs.get("name", ""), 

1729 arguments=kwargs.get("arguments", {}), 

1730 server_name=kwargs.get("server_name"), 

1731 user_api_key_auth=user_api_key_auth_dict, 

1732 hidden_params=HiddenParams(), 

1733 ) 

1734 

1735 def _convert_mcp_hook_response_to_kwargs(self, response_data: dict | None, original_kwargs: dict) -> dict: 

1736 """ 

1737 Helper function to convert pre_call_hook response back to kwargs for MCP usage. 

1738 

1739 Supports: 

1740 - modified_arguments: Override tool call arguments 

1741 - extra_headers: Inject custom headers into the outbound MCP request 

1742 """ 

1743 if not response_data: 

1744 return original_kwargs 

1745 

1746 modified_kwargs: Final = original_kwargs.copy() 

1747 

1748 if response_data.get("modified_arguments"): 

1749 modified_kwargs["arguments"] = response_data["modified_arguments"] 

1750 

1751 if response_data.get("extra_headers"): 

1752 # Merge rather than replace — a prior guardrail in the chain may have 

1753 # already injected headers (e.g. tracing IDs). Later guardrails win on 

1754 # key collisions so that the most-specific guardrail (e.g. JWT signer) 

1755 # takes precedence over earlier ones. 

1756 existing: Final = modified_kwargs.get("extra_headers") or {} 

1757 modified_kwargs["extra_headers"] = { 

1758 **existing, 

1759 **response_data["extra_headers"], 

1760 } 

1761 

1762 return modified_kwargs 

1763 

1764 async def process_pre_call_hook_response(self, response, data, call_type): 

1765 if isinstance(response, Exception): 1765 ↛ 1766line 1765 didn't jump to line 1766 because the condition on line 1765 was never true

1766 raise response 

1767 if isinstance(response, dict): 1767 ↛ 1769line 1767 didn't jump to line 1769 because the condition on line 1767 was always true

1768 return response 

1769 if isinstance(response, str): 

1770 if call_type in ["completion", "text_completion"]: 

1771 raise RejectedRequestError( 

1772 message=response, 

1773 model=data.get("model", ""), 

1774 llm_provider="", 

1775 request_data=data, 

1776 ) 

1777 else: 

1778 raise HTTPException(status_code=400, detail={"error": response}) 

1779 return data 

1780 

1781 def _should_use_guardrail_load_balancing( 

1782 self, 

1783 guardrail_name: str, 

1784 ) -> bool: 

1785 """ 

1786 Check if load balancing should be used for this guardrail. 

1787 

1788 Returns True if the router has multiple deployments for this guardrail name. 

1789 """ 

1790 from litellm.proxy.proxy_server import llm_router 

1791 

1792 if llm_router is None or not hasattr(llm_router, "guardrail_list"): 

1793 return False 

1794 

1795 matching: Final = [g for g in llm_router.guardrail_list if g.get("guardrail_name") == guardrail_name] 

1796 return len(matching) > 1 

1797 

1798 async def _execute_guardrail_hook( 

1799 self, 

1800 callback: "CustomGuardrail", 

1801 hook_type: str, 

1802 data: dict, 

1803 user_api_key_dict: UserAPIKeyAuth | None, 

1804 call_type: CallTypesLiteral, 

1805 response: LLMResponseTypes | None = None, 

1806 ) -> object: 

1807 """ 

1808 Execute a single guardrail's hook. 

1809 

1810 Args: 

1811 callback: The guardrail callback to execute 

1812 hook_type: One of "pre_call", "during_call", "post_call" 

1813 data: Request data 

1814 user_api_key_dict: User API key auth 

1815 call_type: Type of call 

1816 response: Response object (for post_call hooks) 

1817 

1818 Returns: 

1819 Result from the guardrail execution 

1820 """ 

1821 # Use unified_guardrail if callback has apply_guardrail method 

1822 has_apply_guardrail: Final = "apply_guardrail" in type(callback).__dict__ and not getattr( 

1823 callback, "use_native_lifecycle_hooks", False 

1824 ) 

1825 use_unified: Final = has_apply_guardrail and not ( 

1826 hook_type == "during_call" and getattr(callback, "use_native_during_call_hook", False) 

1827 ) 

1828 if use_unified: 

1829 data["guardrail_to_apply"] = callback 

1830 

1831 target: Final = unified_guardrail if use_unified else callback 

1832 

1833 if hook_type == "pre_call": 

1834 return await target.async_pre_call_hook( 

1835 user_api_key_dict=user_api_key_dict, 

1836 cache=self.call_details["user_api_key_cache"], 

1837 data=data, 

1838 call_type=call_type, 

1839 ) 

1840 elif hook_type == "during_call": 

1841 return await target.async_moderation_hook( 

1842 data=data, 

1843 user_api_key_dict=user_api_key_dict, 

1844 call_type=call_type, 

1845 ) 

1846 elif hook_type == "post_call": 

1847 return await target.async_post_call_success_hook( 

1848 user_api_key_dict=user_api_key_dict, 

1849 data=data, 

1850 response=response, 

1851 ) 

1852 else: 

1853 raise ValueError(f"Unknown hook_type: {hook_type}") 

1854 

1855 async def _execute_guardrail_with_load_balancing( 

1856 self, 

1857 guardrail_name: str, 

1858 hook_type: str, 

1859 data: dict, 

1860 user_api_key_dict: UserAPIKeyAuth | None, 

1861 call_type: CallTypesLiteral, 

1862 response: LLMResponseTypes | None = None, 

1863 ) -> object: 

1864 """ 

1865 Execute a guardrail using the router's load balancing. 

1866 

1867 Args: 

1868 guardrail_name: Name of the guardrail 

1869 hook_type: One of "pre_call", "during_call", "post_call" 

1870 data: Request data 

1871 user_api_key_dict: User API key auth 

1872 call_type: Type of call 

1873 response: Response object (for post_call hooks) 

1874 

1875 Returns: 

1876 Result from the guardrail execution 

1877 """ 

1878 from litellm.proxy.proxy_server import llm_router 

1879 

1880 if llm_router is None: 

1881 raise ValueError("Router not initialized") 

1882 

1883 # Select guardrail using router's load balancing 

1884 selected_guardrail: Final = llm_router.get_available_guardrail(guardrail_name=guardrail_name) 

1885 

1886 callback: Final[CustomGuardrail | None] = selected_guardrail.get("callback") 

1887 if callback is None: 

1888 raise ValueError(f"No callback found for guardrail: {guardrail_name}") 

1889 

1890 return await self._execute_guardrail_hook( 

1891 callback=callback, 

1892 hook_type=hook_type, 

1893 data=data, 

1894 user_api_key_dict=user_api_key_dict, 

1895 call_type=call_type, 

1896 response=response, 

1897 ) 

1898 

1899 async def _process_guardrail_callback( 

1900 self, 

1901 callback: CustomGuardrail, 

1902 data: dict, 

1903 user_api_key_dict: UserAPIKeyAuth | None, 

1904 call_type: CallTypesLiteral, 

1905 event_type: GuardrailEventHooks, 

1906 ) -> dict | None: 

1907 """ 

1908 Process a guardrail callback during pre-call hook. 

1909 

1910 Supports load balancing when multiple guardrail deployments exist. 

1911 

1912 Args: 

1913 callback: The CustomGuardrail callback to process 

1914 data: The request data dictionary 

1915 user_api_key_dict: User API key authentication details 

1916 call_type: The type of API call being made 

1917 

1918 Returns: 

1919 Updated data dictionary if guardrail passes, None if guardrail should be skipped 

1920 """ 

1921 from litellm.types.guardrails import GuardrailEventHooks 

1922 

1923 # Determine the event type based on call type 

1924 if event_type is GuardrailEventHooks.pre_call and call_type == CallTypes.call_mcp_tool.value: 

1925 event_type = GuardrailEventHooks.pre_mcp_call 

1926 

1927 # Check if the guardrail should run for this request 

1928 if callback.should_run_guardrail(data=data, event_type=event_type) is not True: 

1929 return None 

1930 

1931 guardrail_name: Final = callback.guardrail_name 

1932 

1933 # Track timing and errors for prometheus metrics 

1934 # Use time.perf_counter() for more accurate duration measurements 

1935 guardrail_start_time: Final = time.perf_counter() 

1936 status = "success" 

1937 error_type = None 

1938 

1939 try: 

1940 # Check if load balancing should be used 

1941 if guardrail_name and self._should_use_guardrail_load_balancing(guardrail_name): 

1942 response = await self._execute_guardrail_with_load_balancing( 

1943 guardrail_name=guardrail_name, 

1944 hook_type="pre_call", 

1945 data=data, 

1946 user_api_key_dict=user_api_key_dict, 

1947 call_type=call_type, 

1948 ) 

1949 else: 

1950 # Single guardrail - execute directly 

1951 response = await self._execute_guardrail_hook( 

1952 callback=callback, 

1953 hook_type="pre_call", 

1954 data=data, 

1955 user_api_key_dict=user_api_key_dict, 

1956 call_type=call_type, 

1957 ) 

1958 

1959 # Process the response if one was returned 

1960 if response is not None: 

1961 data = await self.process_pre_call_hook_response(response=response, data=data, call_type=call_type) 

1962 

1963 callback.mark_pre_call_hook_ran(data) 

1964 

1965 except SensitiveDataRouteException: 

1966 status = "intervened" 

1967 raise 

1968 except Exception as e: 

1969 status = "error" 

1970 error_type = type(e).__name__ 

1971 _enrich_http_exception_with_guardrail_context(e, callback) 

1972 # Re-raise the exception to maintain existing behavior 

1973 raise 

1974 finally: 

1975 # Record prometheus metrics 

1976 guardrail_end_time: Final = time.perf_counter() 

1977 latency_seconds: Final = guardrail_end_time - guardrail_start_time 

1978 

1979 # Get guardrail name for metrics (fallback if not set) 

1980 metrics_guardrail_name: Final = ( 

1981 guardrail_name or getattr(callback, "guardrail_name", callback.__class__.__name__) or "unknown" 

1982 ) 

1983 

1984 self._emit_guardrail_metrics( 

1985 guardrail_name=metrics_guardrail_name, 

1986 latency_seconds=latency_seconds, 

1987 status=status, 

1988 error_type=error_type, 

1989 hook_type="pre_call", 

1990 ) 

1991 

1992 return data 

1993 

1994 async def _run_sequential_guardrail_callback( 

1995 self, 

1996 callback: CustomGuardrail, 

1997 data: dict, # mutable-ok: matches _process_guardrail_callback's own request-payload typing 

1998 raw_request_snapshot: dict | None, # mutable-ok: same request-payload shape as data 

1999 user_api_key_dict: UserAPIKeyAuth, 

2000 call_type: CallTypesLiteral, 

2001 ) -> dict: # mutable-ok: callers reassign the loop's own data from this return value 

2002 """ 

2003 Run one guardrail from the sequential pre_call loop and return what the 

2004 rest of the loop should carry forward. 

2005 

2006 A guardrail opted into ``scan_raw_request`` always evaluates a fresh 

2007 copy of ``raw_request_snapshot`` (taken before any guardrail in this 

2008 hook ran) instead of ``data`` (the live, possibly already-mutated 

2009 payload), so its block/pass decision can never depend on where it's 

2010 declared relative to a guardrail that masks or rewrites content. It's 

2011 declared block-only, same contract as ``run_in_parallel``: any data it 

2012 returns is discarded, since applying its view on top of a stale 

2013 snapshot would silently undo whatever a later guardrail already did to 

2014 the live request. A guardrail that mutates content (e.g. PII masking) 

2015 should never set this flag -- if one does anyway, its returned 

2016 mutation is discarded and a warning is logged so the misconfiguration 

2017 is visible instead of silently forwarding unredacted content. 

2018 """ 

2019 scans_raw_request: Final = callback.scan_raw_request 

2020 should_use_raw_snapshot: Final = scans_raw_request and raw_request_snapshot is not None 

2021 input_data: Final = independent_snapshot(raw_request_snapshot) if should_use_raw_snapshot else data 

2022 # _process_guardrail_callback always calls mark_pre_call_hook_ran on a 

2023 # successful run, which unconditionally stamps bookkeeping metadata onto 

2024 # the dict regardless of whether the guardrail's own hook mutated 

2025 # anything -- so comparing `result` straight against `input_data` would 

2026 # warn on every single scan_raw_request call. Apply that same stamp to a 

2027 # throwaway, guaranteed-independent copy first (never the live request or 

2028 # raw_request_snapshot itself) so the comparison isolates the guardrail's 

2029 # own content mutation from this bookkeeping noise without risking a 

2030 # premature marker write into shared state. 

2031 expected_if_unmutated: Final[dict | None] = ( # mutable-ok: same request-payload shape as data 

2032 independent_snapshot(input_data) if scans_raw_request else None 

2033 ) 

2034 if expected_if_unmutated is not None: 

2035 callback.mark_pre_call_hook_ran(expected_if_unmutated) 

2036 try: 

2037 result: Final = await self._process_guardrail_callback( 

2038 callback=callback, 

2039 data=input_data, 

2040 user_api_key_dict=user_api_key_dict, 

2041 call_type=call_type, 

2042 event_type=GuardrailEventHooks.pre_call, 

2043 ) 

2044 except SensitiveDataRouteException: 

2045 raise 

2046 except Exception: 

2047 _record_raising_guardrail(data, callback) 

2048 raise 

2049 if ( 

2050 scans_raw_request 

2051 and expected_if_unmutated is not None 

2052 and result is not None 

2053 and result != expected_if_unmutated 

2054 ): 

2055 verbose_proxy_logger.warning( 

2056 "Guardrail '%s' has scan_raw_request=True but returned a modified payload; " 

2057 "scan_raw_request is for block-only guardrails and this mutation is being " 

2058 "discarded. Remove scan_raw_request from this guardrail's config if it needs " 

2059 "to mask/rewrite content.", 

2060 callback.guardrail_name or callback.__class__.__name__, 

2061 ) 

2062 if scans_raw_request: 

2063 if result is not None: 

2064 # _process_guardrail_callback only stamped input_data (a throwaway 

2065 # snapshot copy), never the live data returned here -- without this, 

2066 # a deployment-level guardrail sharing this name would see no marker 

2067 # via _pre_call_hook_already_ran and re-run the same guardrail a 

2068 # second time on live kwargs. 

2069 callback.mark_pre_call_hook_ran(data) 

2070 return data 

2071 if result is None: 

2072 return data 

2073 return result 

2074 

2075 async def _process_prompt_template( 

2076 self, 

2077 data: dict, 

2078 litellm_logging_obj: "LiteLLMLoggingObj", 

2079 prompt_id: str, 

2080 prompt_version: int | None, 

2081 call_type: CallTypesLiteral, 

2082 ) -> None: 

2083 """Process prompt template if applicable.""" 

2084 

2085 from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY 

2086 from litellm.responses.utils import ResponsesAPIRequestUtils 

2087 from litellm.utils import get_non_default_completion_params 

2088 

2089 raw_prompt_environment: Final = data.get("prompt_environment", None) 

2090 prompt_environment: Final = raw_prompt_environment if isinstance(raw_prompt_environment, str) else None 

2091 prompt_spec: Final = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec( 

2092 prompt_id, 

2093 version=prompt_version, 

2094 environment=prompt_environment, 

2095 ) 

2096 custom_logger: Final = ( 

2097 IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=prompt_spec) 

2098 if prompt_spec is not None 

2099 else None 

2100 ) 

2101 litellm_prompt_id: str | None = None 

2102 if prompt_spec is not None: 

2103 litellm_prompt_id = prompt_spec.litellm_params.prompt_id 

2104 data.pop("prompt_id", None) 

2105 data.pop("prompt_environment", None) 

2106 

2107 if custom_logger and prompt_spec is not None: 

2108 is_responses_call: Final = call_type == "aresponses" 

2109 original_responses_input: Final = data.get("input", "") if is_responses_call else "" 

2110 client_messages: Final = ( 

2111 ResponsesAPIRequestUtils.responses_input_to_chat_messages(original_responses_input) 

2112 if is_responses_call 

2113 else data.get("messages", []) 

2114 ) 

2115 ( 

2116 model, 

2117 messages, 

2118 optional_params, 

2119 ) = await litellm_logging_obj.async_get_chat_completion_prompt( 

2120 model=data.get("model", ""), 

2121 messages=client_messages, 

2122 non_default_params=get_non_default_completion_params(kwargs=data) or {}, 

2123 prompt_id=litellm_prompt_id, 

2124 prompt_spec=prompt_spec, 

2125 prompt_management_logger=custom_logger, 

2126 prompt_variables=data.pop("prompt_variables", None) or {}, 

2127 prompt_label=data.pop("prompt_label", None) or {}, 

2128 prompt_version=data.pop("prompt_version", None) or {}, 

2129 request_kwargs=data, 

2130 injected_for_every_deployment=True, 

2131 ) 

2132 

2133 data.update(optional_params) 

2134 data["model"] = model 

2135 if is_responses_call: 

2136 data["input"] = ResponsesAPIRequestUtils.merge_prompt_management_input( 

2137 original_input=original_responses_input, 

2138 client_input=client_messages, 

2139 merged_input=messages, 

2140 ) 

2141 else: 

2142 data["messages"] = messages 

2143 # prevent re-processing the prompt template 

2144 data.pop("prompt_id", None) 

2145 data.pop("prompt_variables", None) 

2146 data.pop("prompt_label", None) 

2147 data.pop("prompt_version", None) 

2148 data.pop("prompt_environment", None) 

2149 

2150 def _process_guardrail_metadata(self, data: dict) -> None: 

2151 """Process guardrails from metadata and add to applied_guardrails.""" 

2152 from litellm.proxy.common_utils.callback_utils import ( 

2153 add_guardrail_to_applied_guardrails_header, 

2154 ) 

2155 

2156 metadata_standard: Final = data.get("metadata") or {} 

2157 metadata_litellm: Final = data.get("litellm_metadata") or {} 

2158 

2159 guardrails_in_metadata = [] 

2160 if isinstance(metadata_standard, dict) and "guardrails" in metadata_standard: 

2161 guardrails_in_metadata = metadata_standard.get("guardrails", []) 

2162 elif isinstance(metadata_litellm, dict) and "guardrails" in metadata_litellm: 

2163 guardrails_in_metadata = metadata_litellm.get("guardrails", []) 

2164 

2165 if guardrails_in_metadata and isinstance(guardrails_in_metadata, list): 

2166 applied_guardrails = [] 

2167 if isinstance(metadata_standard, dict) and "applied_guardrails" in metadata_standard: 2167 ↛ 2168line 2167 didn't jump to line 2168 because the condition on line 2167 was never true

2168 applied_guardrails = metadata_standard.get("applied_guardrails", []) 

2169 elif isinstance(metadata_litellm, dict) and "applied_guardrails" in metadata_litellm: 2169 ↛ 2170line 2169 didn't jump to line 2170 because the condition on line 2169 was never true

2170 applied_guardrails = metadata_litellm.get("applied_guardrails", []) 

2171 

2172 if not isinstance(applied_guardrails, list): 2172 ↛ 2173line 2172 didn't jump to line 2173 because the condition on line 2172 was never true

2173 applied_guardrails = [] 

2174 

2175 for guardrail_name in guardrails_in_metadata: 

2176 if isinstance(guardrail_name, str) and guardrail_name not in applied_guardrails: 

2177 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=guardrail_name) 

2178 

2179 async def _maybe_execute_pipelines( 

2180 self, 

2181 data: dict, 

2182 user_api_key_dict: UserAPIKeyAuth, 

2183 call_type: str, 

2184 event_hook: str, 

2185 raw_request_snapshot: dict | None = None, # mutable-ok: same request-payload shape as data 

2186 response: LLMResponseTypes | None = None, 

2187 ) -> tuple[dict, LLMResponseTypes | None]: # mutable-ok: returns the request-payload dict onward 

2188 """ 

2189 Execute guardrail pipelines if any are configured for this request. 

2190 

2191 Checks metadata for pipelines resolved by the policy engine 

2192 and executes them. Handles the result (allow/block/modify_response). 

2193 

2194 ``raw_request_snapshot`` (taken before any guardrail or pipeline ran) 

2195 is forwarded so a pipeline step whose guardrail opted into 

2196 ``scan_raw_request`` evaluates the pristine request, not whatever an 

2197 earlier ``pass_data`` step in the same pipeline already rewrote. 

2198 

2199 Returns the (possibly modified) data dict, plus the replacement 

2200 response when a post_call pipeline step returned one (None when the 

2201 response is unchanged), matching the flat callback-loop contract. 

2202 """ 

2203 pipelines: Final = _policy_pipelines(data) 

2204 if not pipelines: 2204 ↛ 2207line 2204 didn't jump to line 2207 because the condition on line 2204 was always true

2205 return data, None 

2206 

2207 current_response = response # rebind-ok: chains each pipeline's replacement response into the next 

2208 for policy_name, pipeline in pipelines: 

2209 if pipeline.mode != event_hook: 

2210 continue 

2211 

2212 step_input: dict = {**data, "response": current_response} if current_response is not None else data 

2213 

2214 result: PipelineExecutionResult = await PipelineExecutor.execute_steps( 

2215 steps=pipeline.steps, 

2216 mode=pipeline.mode, 

2217 data=step_input, 

2218 user_api_key_dict=user_api_key_dict, 

2219 call_type=call_type, 

2220 policy_name=policy_name, 

2221 raw_request_snapshot=raw_request_snapshot, 

2222 ) 

2223 

2224 data = self._handle_pipeline_result( 

2225 result=result, 

2226 data=data, 

2227 policy_name=policy_name, 

2228 original_response=current_response, 

2229 ) 

2230 

2231 if current_response is not None and result.modified_data is not None: 

2232 current_response = result.modified_data.get("response", current_response) 

2233 

2234 return data, current_response if current_response is not response else None 

2235 

2236 @staticmethod 

2237 def _handle_pipeline_result( 

2238 result: PipelineExecutionResult, 

2239 data: dict, 

2240 policy_name: str, 

2241 original_response: "LLMResponseTypes | Sequence[object] | None" = None, 

2242 ) -> dict: 

2243 """ 

2244 Handle a PipelineExecutionResult — allow, block, or modify_response. 

2245 

2246 Returns data dict if allowed, raises on block/modify_response. 

2247 ``original_response`` is set on the post_call path, where the request 

2248 payload (already sent upstream) must stay untouched; a replacement 

2249 response carried in ``modified_data`` is adopted by the caller, and 

2250 metadata-bucket writes (applied guardrails, guardrail logging info) 

2251 are merged back so headers and spend logs still see them, on block 

2252 and modify_response too, so failure spend records keep guardrail 

2253 cost and status. On the 

2254 streaming path it is the buffered chunk list, carried into 

2255 ``ModifyResponseException.original_response`` for usage reporting. 

2256 """ 

2257 if result.terminal_action == "allow": 

2258 if result.modified_data is not None: 

2259 if original_response is None: 

2260 data.update(result.modified_data) 

2261 else: 

2262 _merge_pipeline_metadata_writes(data, result.modified_data) 

2263 return data 

2264 

2265 if result.modified_data is not None: 

2266 _merge_pipeline_metadata_writes(data, result.modified_data) 

2267 

2268 if result.terminal_action == "block": 

2269 blocking_step: Final = result.step_results[-1] if result.step_results else None 

2270 callback: Final = ( 

2271 PipelineExecutor.find_guardrail_callback(blocking_step.guardrail_name) 

2272 if blocking_step is not None 

2273 else None 

2274 ) 

2275 if callback is not None: 

2276 _record_raising_guardrail(data, callback) 

2277 original_exception: Final = result.original_exception 

2278 if original_exception is not None and not _exception_changes_request_flow(original_exception): 

2279 if callback is not None: 

2280 _enrich_http_exception_with_guardrail_context(original_exception, callback) 

2281 raise original_exception 

2282 

2283 step_results_serializable: Final = [ 

2284 { 

2285 "guardrail": sr.guardrail_name, 

2286 "outcome": sr.outcome, 

2287 "action": sr.action_taken, 

2288 } 

2289 for sr in result.step_results 

2290 ] 

2291 error_detail: Final = { 

2292 "error": { 

2293 "message": f"Content blocked by guardrail pipeline '{policy_name}'", 

2294 "type": "guardrail_pipeline_error", 

2295 "pipeline_context": { 

2296 "policy": policy_name, 

2297 "step_results": step_results_serializable, 

2298 }, 

2299 } 

2300 } 

2301 raise HTTPException(status_code=400, detail=error_detail) 

2302 

2303 if result.terminal_action == "modify_response": 

2304 raise ModifyResponseException( 

2305 message=result.modify_response_message or "Response modified by pipeline", 

2306 model=data.get("model", "unknown"), 

2307 request_data=data, 

2308 guardrail_name=f"pipeline:{policy_name}", 

2309 detection_info=None, 

2310 original_response=original_response, 

2311 ) 

2312 

2313 return data 

2314 

2315 def has_pre_call_guardrails(self, request_metadata: Mapping[str, object]) -> bool: 

2316 """ 

2317 Whether anything configured would inspect the content of a request carrying this metadata. 

2318 

2319 Evaluated with the same predicate the pre-call loop uses, so a proxy configured only with 

2320 post-call guardrails answers False. Callers that must pay a real cost to build the hook's 

2321 input, such as streaming a batch input file off disk, use this to skip that work. 

2322 

2323 A content-enforcing ``CustomLogger`` counts too. It is not a guardrail and has no event 

2324 hook to consult, but it judges the payload the same way, so a proxy configured only with 

2325 one of those still has something to say about every record. 

2326 """ 

2327 if request_metadata.get("_guardrail_pipelines"): 

2328 return True 

2329 caps: Final = ProxyLogging._callback_capabilities() 

2330 if caps.has_content_enforcer: 

2331 return True 

2332 probe: Final = {"metadata": dict(request_metadata)} # mutable-ok: should_run_guardrail takes a dict 

2333 return any( 

2334 isinstance(callback, CustomGuardrail) 

2335 and callback.should_run_guardrail(data=probe, event_type=GuardrailEventHooks.pre_call) 

2336 for callback in caps.resolved_callbacks 

2337 ) 

2338 

2339 # The actual implementation of the function 

2340 @overload 

2341 async def pre_call_hook( 

2342 self, 

2343 user_api_key_dict: UserAPIKeyAuth, 

2344 data: None, 

2345 call_type: CallTypesLiteral, 

2346 guardrails_only: bool = False, 

2347 skip_guardrails: bool = False, 

2348 ) -> None: 

2349 pass 

2350 

2351 @overload 

2352 async def pre_call_hook( 

2353 self, 

2354 user_api_key_dict: UserAPIKeyAuth, 

2355 data: dict, 

2356 call_type: CallTypesLiteral, 

2357 guardrails_only: bool = False, 

2358 skip_guardrails: bool = False, 

2359 ) -> dict: 

2360 pass 

2361 

2362 async def pre_call_hook( 

2363 self, 

2364 user_api_key_dict: UserAPIKeyAuth, 

2365 data: dict | None, 

2366 call_type: CallTypesLiteral, 

2367 guardrails_only: bool = False, 

2368 skip_guardrails: bool = False, 

2369 ) -> dict | None: 

2370 """ 

2371 Allows users to modify/reject the incoming request to the proxy, without having to deal with parsing Request body. 

2372 

2373 Covers: 

2374 1. /chat/completions 

2375 2. /embeddings 

2376 3. /image/generation 

2377 

2378 With ``guardrails_only`` the walk is limited to guardrails and guardrail pipelines: rate 

2379 limiting, budget accounting, prompt templates and hanging-request alerting are skipped. 

2380 Use it to scan a payload that is not itself a request, such as one record of a batch file. 

2381 """ 

2382 verbose_proxy_logger.debug("Inside Proxy Logging Pre-call hook!") 

2383 

2384 if guardrails_only and skip_guardrails: 2384 ↛ 2385line 2384 didn't jump to line 2385 because the condition on line 2384 was never true

2385 raise ValueError("guardrails_only and skip_guardrails are mutually exclusive") 

2386 

2387 if not guardrails_only: 2387 ↛ 2390line 2387 didn't jump to line 2390 because the condition on line 2387 was always true

2388 self._init_response_taking_too_long_task(data=data) 

2389 

2390 if data is None: 2390 ↛ 2391line 2390 didn't jump to line 2391 because the condition on line 2390 was never true

2391 return None 

2392 

2393 litellm_logging_obj: Final = cast(Optional["LiteLLMLoggingObj"], data.get("litellm_logging_obj", None)) 

2394 prompt_id: Final[str | None] = data.get("prompt_id", None) 

2395 

2396 ## PROMPT TEMPLATE CHECK ## 

2397 

2398 if ( 2398 ↛ 2404line 2398 didn't jump to line 2404 because the condition on line 2398 was never true

2399 not guardrails_only 

2400 and litellm_logging_obj is not None 

2401 and prompt_id is not None 

2402 and (call_type == "completion" or call_type == "acompletion" or call_type == "aresponses") 

2403 ): 

2404 from litellm.proxy.prompts.prompt_registry import parse_prompt_version 

2405 

2406 await self._process_prompt_template( 

2407 data=data, 

2408 litellm_logging_obj=litellm_logging_obj, 

2409 prompt_id=prompt_id, 

2410 prompt_version=parse_prompt_version(data.get("prompt_version", None)), 

2411 call_type=call_type, 

2412 ) 

2413 

2414 # Snapshotted here, before _maybe_execute_pipelines or any guardrail in 

2415 # this hook has run, so a scan_raw_request guardrail's block/pass 

2416 # decision never depends on its position in the guardrails list or on 

2417 # a pipeline that runs ahead of it: an earlier guardrail (pipelined or 

2418 # not) that masks/rewrites content can't hide a violation from a later 

2419 # one that opted into scanning the original request. Only computed 

2420 # when at least one registered guardrail actually opted in, and via 

2421 # independent_snapshot (not safe_deep_copy) since this isolation 

2422 # guarantee must hold even under litellm.safe_memory_mode, which 

2423 # otherwise makes deep copies return the original object. 

2424 needs_raw_request_snapshot: Final = any( 

2425 isinstance(cb, CustomGuardrail) and cb.scan_raw_request 

2426 for cb in ProxyLogging._callback_capabilities().resolved_callbacks 

2427 ) 

2428 raw_request_snapshot: Final[dict | None] = ( # mutable-ok: same request-payload shape as data 

2429 independent_snapshot(data) if needs_raw_request_snapshot else None 

2430 ) 

2431 

2432 try: 

2433 if not skip_guardrails: 2433 ↛ 2443line 2433 didn't jump to line 2443 because the condition on line 2433 was always true

2434 data, _ = await self._maybe_execute_pipelines( # rebind-ok: pipeline edits feed the callback loop below 

2435 data=data, 

2436 user_api_key_dict=user_api_key_dict, 

2437 call_type=call_type, 

2438 event_hook="pre_call", 

2439 raw_request_snapshot=raw_request_snapshot, 

2440 ) 

2441 

2442 # Get pipeline-managed guardrails to skip in normal loop 

2443 pipeline_managed: Final[frozenset[str]] = ( 

2444 frozenset() if skip_guardrails else pipeline_managed_guardrail_names(data, "pre_call") 

2445 ) 

2446 

2447 caps: Final = ProxyLogging._callback_capabilities() 

2448 # Skip the per-request callback walk entirely when nothing in 

2449 # ``litellm.callbacks`` overrides ``async_pre_call_hook`` and no 

2450 # CustomGuardrail is configured. Saves the loop overhead + 

2451 # ``time.time()`` x2 per registered callback for the common 

2452 # "callbacks=[]" case on small / dev deployments. 

2453 if ( 2453 ↛ 2458line 2453 didn't jump to line 2458 because the condition on line 2453 was never true

2454 (skip_guardrails or not caps.has_guardrail) 

2455 and not caps.has_content_enforcer 

2456 and (guardrails_only or not caps.has_pre_call_override) 

2457 ): 

2458 if data is not None: 

2459 self._process_guardrail_metadata(data) 

2460 return data 

2461 

2462 parallel_guardrails: Final[tuple[CustomGuardrail, ...]] = ( 

2463 () 

2464 if skip_guardrails 

2465 else tuple( 

2466 cb 

2467 for cb in caps.resolved_callbacks 

2468 if isinstance(cb, CustomGuardrail) 

2469 and getattr(cb, "run_in_parallel", False) 

2470 and not (cb.guardrail_name and cb.guardrail_name in pipeline_managed) 

2471 ) 

2472 ) 

2473 

2474 deferred_route_exc: SensitiveDataRouteException | None = None 

2475 for _callback in caps.resolved_callbacks: 

2476 start_time = time.time() 

2477 try: 

2478 if isinstance(_callback, CustomGuardrail) and data is not None: 2478 ↛ 2479line 2478 didn't jump to line 2479 because the condition on line 2478 was never true

2479 if skip_guardrails: 

2480 continue 

2481 

2482 # Skip guardrails managed by a pipeline 

2483 if _callback.guardrail_name and _callback.guardrail_name in pipeline_managed: 

2484 continue 

2485 

2486 if getattr(_callback, "run_in_parallel", False): 

2487 continue 

2488 

2489 data = await self._run_sequential_guardrail_callback( 

2490 callback=_callback, 

2491 data=data, 

2492 raw_request_snapshot=raw_request_snapshot, 

2493 user_api_key_dict=user_api_key_dict, 

2494 call_type=call_type, 

2495 ) 

2496 

2497 elif ( 

2498 _callback is not None 

2499 and isinstance(_callback, CustomLogger) 

2500 and (not guardrails_only or _callback.enforces_request_content) 

2501 and "async_pre_call_hook" in vars(_callback.__class__) 

2502 and _callback.__class__.async_pre_call_hook != CustomLogger.async_pre_call_hook 

2503 ): 

2504 if call_type == "call_mcp_tool" and user_api_key_dict is None: 2504 ↛ 2505line 2504 didn't jump to line 2505 because the condition on line 2504 was never true

2505 continue 

2506 

2507 response: Exception | str | Mapping[str, object] | None = await _callback.async_pre_call_hook( 

2508 user_api_key_dict=user_api_key_dict, 

2509 cache=self.call_details["user_api_key_cache"], 

2510 data=data, 

2511 call_type=call_type, 

2512 ) 

2513 if response is not None: 

2514 data = await self.process_pre_call_hook_response( 

2515 response=response, data=data, call_type=call_type 

2516 ) 

2517 except SensitiveDataRouteException as e: 

2518 # Defer the reroute until remaining guardrails have run so later 

2519 # security checks are not skipped; the first reroute wins and a 

2520 # later guardrail that blocks still propagates. Fall through to the 

2521 # service-span recording below so the triggering guardrail is still 

2522 # timed like every other callback. 

2523 if deferred_route_exc is None: 

2524 deferred_route_exc = e 

2525 

2526 end_time = time.time() 

2527 duration = end_time - start_time 

2528 if ( 

2529 hasattr(self, "service_logging_obj") and duration > 0.01 

2530 ): # only if duration is non-negligible - don't spam the logs 

2531 await self.service_logging_obj.async_service_success_hook( 

2532 service=ServiceTypes.PROXY_PRE_CALL, 

2533 duration=duration, 

2534 call_type=f"{_callback.__class__.__name__}", 

2535 parent_otel_span=user_api_key_dict.parent_otel_span, 

2536 start_time=start_time, 

2537 end_time=end_time, 

2538 ) 

2539 

2540 if deferred_route_exc is not None and data is not None: 2540 ↛ 2541line 2540 didn't jump to line 2541 because the condition on line 2540 was never true

2541 data = await self._handle_sensitive_data_route_exception(deferred_route_exc, data, user_api_key_dict) 

2542 

2543 if parallel_guardrails and data is not None: 2543 ↛ 2544line 2543 didn't jump to line 2544 because the condition on line 2543 was never true

2544 await self._run_parallel_pre_call_guardrails( 

2545 guardrails=parallel_guardrails, 

2546 data=data, 

2547 raw_request_snapshot=raw_request_snapshot, 

2548 user_api_key_dict=user_api_key_dict, 

2549 call_type=call_type, 

2550 ) 

2551 

2552 if data is not None: 2552 ↛ 2555line 2552 didn't jump to line 2555 because the condition on line 2552 was always true

2553 self._process_guardrail_metadata(data) 

2554 

2555 return data 

2556 except SensitiveDataRouteException as e: 

2557 data = await self._handle_sensitive_data_route_exception(e, data, user_api_key_dict) 

2558 if data is not None: 

2559 self._process_guardrail_metadata(data) 

2560 return data 

2561 except Exception: 

2562 if data is not None: 2562 ↛ 2564line 2562 didn't jump to line 2564 because the condition on line 2562 was always true

2563 self._process_guardrail_metadata(data) 

2564 raise 

2565 

2566 async def _run_parallel_pre_call_guardrails( 

2567 self, 

2568 guardrails: tuple[CustomGuardrail, ...], 

2569 data: dict, 

2570 raw_request_snapshot: dict | None, # mutable-ok: same request-payload shape as data 

2571 user_api_key_dict: UserAPIKeyAuth, 

2572 call_type: CallTypesLiteral, 

2573 ) -> None: 

2574 """ 

2575 Run opted-in pre_call guardrails concurrently against one shared payload 

2576 snapshot. These guardrails are declared block-only, so any modified data 

2577 they return is discarded; they run for their blocking side effect (raising 

2578 to reject the request before it reaches the LLM). Every guardrail is 

2579 awaited to completion (``return_exceptions=True``) so a raise by one never 

2580 leaves the others running as unobserved background tasks. A guardrail that 

2581 blocks (any exception other than a reroute or passthrough) takes precedence 

2582 over one that only changes the request flow, so a fast reroute can never 

2583 let a slower block be bypassed; the request is rejected before it reaches 

2584 the LLM, preserving the pre-call barrier that ``during_call`` guardrails 

2585 cannot provide. Per-guardrail latency is recorded by 

2586 ``_process_guardrail_callback``'s own metrics. 

2587 

2588 A guardrail that also opted into ``scan_raw_request`` evaluates 

2589 ``raw_request_snapshot`` (taken before the sequential loop ran) instead 

2590 of ``data`` (the sequential loop's output), for the same reason the 

2591 sequential branch does: its block decision must not depend on what a 

2592 sequential guardrail already masked or rewrote. 

2593 """ 

2594 

2595 def _input_for(callback: CustomGuardrail) -> dict: # mutable-ok: same request-payload shape as data 

2596 if not callback.scan_raw_request or raw_request_snapshot is None: 

2597 return data 

2598 return independent_snapshot(raw_request_snapshot) 

2599 

2600 results: Final = await asyncio.gather( 

2601 *( 

2602 self._process_guardrail_callback( 

2603 callback=callback, 

2604 data=_input_for(callback), 

2605 user_api_key_dict=user_api_key_dict, 

2606 call_type=call_type, 

2607 event_type=GuardrailEventHooks.pre_call, 

2608 ) 

2609 for callback in guardrails 

2610 ), 

2611 return_exceptions=True, 

2612 ) 

2613 for callback, result in zip(guardrails, results, strict=True): 

2614 # _process_guardrail_callback stamped mark_pre_call_hook_ran on 

2615 # _input_for's throwaway snapshot copy for a scan_raw_request 

2616 # guardrail, never on the live, shared `data` -- without this, a 

2617 # deployment-level guardrail sharing this name would see no marker 

2618 # via _pre_call_hook_already_ran and re-run it a second time on 

2619 # live kwargs. 

2620 if callback.scan_raw_request and not isinstance(result, BaseException) and result is not None: 

2621 callback.mark_pre_call_hook_ran(data) 

2622 if isinstance(result, BaseException) and not isinstance(result, SensitiveDataRouteException): 

2623 _record_raising_guardrail(data, callback) 

2624 raised: Final = tuple(result for result in results if isinstance(result, BaseException)) 

2625 blocking: Final = next((exc for exc in raised if not _exception_changes_request_flow(exc)), None) 

2626 if blocking is not None: 

2627 raise blocking 

2628 if raised: 

2629 raise raised[0] 

2630 

2631 async def _handle_sensitive_data_route_exception( 

2632 self, 

2633 exc: SensitiveDataRouteException, 

2634 data: dict | None, 

2635 user_api_key_dict: UserAPIKeyAuth | None, 

2636 ) -> dict | None: 

2637 """ 

2638 Handle SensitiveDataRouteException by rerouting the current request to 

2639 the target model and, when sticky_session_routing is enabled, persisting 

2640 the session override so subsequent requests reuse the same model. 

2641 """ 

2642 if data is None: 

2643 return None 

2644 

2645 verbose_proxy_logger.info( 

2646 "SensitiveDataRouteException caught: session_id=%s route_to_model=%s guardrail=%s sticky=%s", 

2647 exc.session_id, 

2648 exc.route_to_model, 

2649 exc.guardrail_name, 

2650 exc.sticky_session_routing, 

2651 ) 

2652 

2653 if exc.sticky_session_routing: 

2654 sensitive_routing_hook: Final = self.get_proxy_hook("sensitive_data_routing") 

2655 if isinstance(sensitive_routing_hook, _PROXY_SensitiveDataRoutingHandler): 

2656 await sensitive_routing_hook.set_session_routing( 

2657 session_id=exc.session_id, 

2658 model=exc.route_to_model, 

2659 user_api_key_dict=user_api_key_dict, 

2660 guardrail_name=exc.guardrail_name, 

2661 ) 

2662 else: 

2663 verbose_proxy_logger.warning( 

2664 "SensitiveDataRouteException requested sticky routing for session_id=%s " 

2665 "but the 'sensitive_data_routing' hook is not registered. Only this request " 

2666 "will be rerouted; subsequent requests will not be sticky.", 

2667 exc.session_id, 

2668 ) 

2669 

2670 original_model: Final = data.get("model") 

2671 data["model"] = exc.route_to_model 

2672 

2673 metadata: Final = data.get("metadata") or {} 

2674 metadata["sensitive_data_routing_applied"] = True 

2675 metadata["sensitive_data_routing_original_model"] = original_model 

2676 metadata["sensitive_data_routing_guardrail"] = exc.guardrail_name 

2677 metadata["sensitive_data_routing_detection_info"] = exc.detection_info 

2678 data["metadata"] = metadata 

2679 

2680 return data 

2681 

2682 @staticmethod 

2683 def _emit_guardrail_metrics( 

2684 guardrail_name: str, 

2685 latency_seconds: float, 

2686 status: str, 

2687 error_type: str | None, 

2688 hook_type: str, 

2689 ) -> None: 

2690 for prom_callback in litellm.callbacks: 

2691 if isinstance(prom_callback, PrometheusLogger): 

2692 prom_callback._record_guardrail_metrics( 

2693 guardrail_name=guardrail_name, 

2694 latency_seconds=latency_seconds, 

2695 status=status, 

2696 error_type=error_type, 

2697 hook_type=hook_type, 

2698 ) 

2699 break 

2700 

2701 @staticmethod 

2702 async def _run_guardrail_with_metrics( 

2703 callback: object, 

2704 coro: Awaitable[_T], 

2705 hook_type: str, 

2706 request_data: Mapping[str, object], 

2707 ) -> _T: 

2708 """ 

2709 Await `coro`, recording its latency and status to the 

2710 `litellm_guardrail_latency_seconds` metric under `hook_type`, and 

2711 enriching any raised HTTPException with the originating callback's 

2712 `guardrail_name`/`guardrail_mode` before re-raising. 

2713 """ 

2714 guardrail_name: Final = getattr(callback, "guardrail_name", None) or type(callback).__name__ 

2715 start_time: Final = time.perf_counter() 

2716 status = "success" 

2717 error_type: str | None = None 

2718 try: 

2719 return await coro 

2720 except SensitiveDataRouteException: 

2721 status = "intervened" 

2722 raise 

2723 except Exception as e: 

2724 status = "error" 

2725 error_type = type(e).__name__ 

2726 _enrich_http_exception_with_guardrail_context(e, callback) 

2727 _record_raising_guardrail(request_data, callback) 

2728 raise 

2729 finally: 

2730 ProxyLogging._emit_guardrail_metrics( 

2731 guardrail_name=guardrail_name, 

2732 latency_seconds=time.perf_counter() - start_time, 

2733 status=status, 

2734 error_type=error_type, 

2735 hook_type=hook_type, 

2736 ) 

2737 

2738 @staticmethod 

2739 async def _wrap_streaming_iterator_with_enrichment( 

2740 callback: object, 

2741 response: AsyncIterable[_T], 

2742 hook: _StreamIteratorHook[_T], 

2743 request_data: Mapping[str, object], 

2744 ) -> AsyncGenerator[_T, None]: 

2745 upstream: Final = _UpstreamStreamBoundary(response) 

2746 try: 

2747 async for chunk in hook(response=upstream): 

2748 yield chunk 

2749 except Exception as e: 

2750 if e is not upstream.failure: 

2751 _enrich_http_exception_with_guardrail_context(e, callback) 

2752 _record_raising_guardrail(request_data, callback) 

2753 raise 

2754 

2755 # Cache for callback-capability detection. Keyed on a signature of 

2756 # litellm.callbacks (length + each item's id) so we recompute when the 

2757 # callback list mutates (add/remove) without iterating every request. 

2758 _callback_capabilities_cache: ClassVar[dict[tuple[int, tuple[int, ...]], "_CallbackCapabilities"]] = {} 

2759 

2760 @staticmethod 

2761 def _callback_capabilities() -> "_CallbackCapabilities": 

2762 """ 

2763 Inspect ``litellm.callbacks`` once and answer the per-hook capability 

2764 questions used to short-circuit no-op work on the chat-completions hot 

2765 path. Per-request callers iterated ``litellm.callbacks`` and called 

2766 ``get_custom_logger_compatible_class`` for every string entry — that 

2767 scanning cost dominated the proxy overhead on low-config deployments. 

2768 

2769 Cache invalidates whenever the list length or member identities change. 

2770 """ 

2771 callbacks: Final = litellm.callbacks 

2772 sig: Final = (len(callbacks), tuple(id(c) for c in callbacks)) 

2773 cache: Final = ProxyLogging._callback_capabilities_cache 

2774 cached: Final = cache.get(sig) 

2775 if cached is not None: 

2776 return cached 

2777 

2778 has_post_call_response_headers = False 

2779 has_iterator_override = False 

2780 has_streaming_chunk_override = False 

2781 has_guardrail = False 

2782 has_pre_call_override = False 

2783 has_content_enforcer = False 

2784 has_moderation_override = False 

2785 iterator_overrides: Final[list[tuple[Any, str]]] = [] # (callback, kind) 

2786 resolved_callbacks: Final[list[CustomLogger]] = [] 

2787 

2788 for callback in callbacks: 

2789 if isinstance(callback, str): 2789 ↛ 2790line 2789 didn't jump to line 2790 because the condition on line 2789 was never true

2790 resolved = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class( 

2791 cast(_custom_logger_compatible_callbacks_literal, callback) 

2792 ) 

2793 else: 

2794 resolved = callback 

2795 if resolved is None or not isinstance(resolved, CustomLogger): 2795 ↛ 2796line 2795 didn't jump to line 2796 because the condition on line 2795 was never true

2796 continue 

2797 resolved_callbacks.append(resolved) 

2798 cls = type(resolved) 

2799 if cls is CustomLogger: 2799 ↛ 2800line 2799 didn't jump to line 2800 because the condition on line 2799 was never true

2800 continue 

2801 if isinstance(resolved, CustomGuardrail): 2801 ↛ 2802line 2801 didn't jump to line 2802 because the condition on line 2801 was never true

2802 has_guardrail = True 

2803 elif _overrides_moderation_hook(resolved): 2803 ↛ 2804line 2803 didn't jump to line 2804 because the condition on line 2803 was never true

2804 has_moderation_override = True 

2805 # Use the same leaf-class ``__dict__`` check as the other hook 

2806 # capabilities: only callbacks that actually override the hook 

2807 # contribute to the flag. Setting this for every ``CustomLogger`` 

2808 # instance (the prior behaviour) forced the full 

2809 # ``post_call_response_headers_hook`` body to run on every request 

2810 # even when no registered callback customized response headers. 

2811 cls_attrs = cls.__dict__ 

2812 if "async_post_call_response_headers_hook" in cls_attrs: 

2813 has_post_call_response_headers = True 

2814 if "async_post_call_streaming_iterator_hook" in cls_attrs: 

2815 has_iterator_override = True 

2816 iterator_overrides.append((resolved, "override")) 

2817 elif "apply_guardrail" in cls_attrs and not getattr(resolved, "use_native_lifecycle_hooks", False): 2817 ↛ 2818line 2817 didn't jump to line 2818 because the condition on line 2817 was never true

2818 iterator_overrides.append((resolved, "apply_guardrail")) 

2819 # Walk the MRO for ``async_post_call_streaming_hook`` rather than 

2820 # using the leaf-class ``__dict__`` check used by the other flags: 

2821 # before this PR the hook was unconditionally invoked, so a 

2822 # callback that inherits an override from an intermediate parent 

2823 # (e.g. a vendor base class providing the override, with the 

2824 # registered class adding nothing else) MUST still be detected. 

2825 # A leaf-class miss here would silently drop the inherited hook. 

2826 base_streaming_hook = CustomLogger.async_post_call_streaming_hook 

2827 cls_streaming_hook = getattr( 

2828 cls, 

2829 "async_post_call_streaming_hook", 

2830 base_streaming_hook, 

2831 ) 

2832 if getattr(cls_streaming_hook, "__func__", cls_streaming_hook) is not getattr( 2832 ↛ 2835line 2832 didn't jump to line 2835 because the condition on line 2832 was never true

2833 base_streaming_hook, "__func__", base_streaming_hook 

2834 ): 

2835 has_streaming_chunk_override = True 

2836 if "async_pre_call_hook" in cls_attrs: 

2837 has_pre_call_override = True 

2838 if resolved.enforces_request_content is True: 2838 ↛ 2839line 2838 didn't jump to line 2839 because the condition on line 2838 was never true

2839 has_content_enforcer = True 

2840 

2841 caps: Final = _CallbackCapabilities( 

2842 has_post_call_response_headers=has_post_call_response_headers, 

2843 has_iterator_override=has_iterator_override 

2844 or any(kind == "apply_guardrail" for _, kind in iterator_overrides), 

2845 has_streaming_chunk_override=has_streaming_chunk_override, 

2846 has_guardrail=has_guardrail, 

2847 has_pre_call_override=has_pre_call_override, 

2848 has_content_enforcer=has_content_enforcer, 

2849 has_moderation_override=has_moderation_override, 

2850 iterator_overrides=tuple(iterator_overrides), 

2851 resolved_callbacks=tuple(resolved_callbacks), 

2852 listed_models_filters=tuple( 

2853 callback for callback in resolved_callbacks if _overrides_hook(callback, "async_filter_listed_models") 

2854 ), 

2855 ) 

2856 # Limit cache to handle test churn without leaking; production 

2857 # callback lists are stable so this rarely grows past 1 entry. 

2858 if len(cache) >= 32: 

2859 cache.clear() 

2860 cache[sig] = caps 

2861 return caps 

2862 

2863 @staticmethod 

2864 def _stream_requires_guardrail_translation(user_api_key_dict: UserAPIKeyAuth) -> bool: 

2865 route: Final = user_api_key_dict.request_route 

2866 if not route: 

2867 return False 

2868 call_types: Final = get_call_types_for_route(route) 

2869 if not call_types: 

2870 return False 

2871 return call_types[0] in NON_OPENAI_STREAM_GUARDRAIL_TRANSLATION_CALL_TYPES 

2872 

2873 @staticmethod 

2874 def has_post_call_response_headers_callbacks() -> bool: 

2875 return ProxyLogging._callback_capabilities().has_post_call_response_headers 

2876 

2877 @staticmethod 

2878 def has_streaming_callbacks() -> bool: 

2879 caps: Final = ProxyLogging._callback_capabilities() 

2880 return caps.has_iterator_override or caps.has_streaming_chunk_override or caps.has_guardrail 

2881 

2882 @staticmethod 

2883 def has_streaming_chunk_hook_overrides() -> bool: 

2884 """True iff any callback overrides ``async_post_call_streaming_hook`` 

2885 (the per-chunk hook, distinct from the iterator wrapper).""" 

2886 caps: Final = ProxyLogging._callback_capabilities() 

2887 return caps.has_streaming_chunk_override or caps.has_guardrail 

2888 

2889 def needs_iterator_wrap(self) -> bool: 

2890 """Whether ``async_data_generator`` needs to wrap the upstream stream 

2891 through ``async_post_call_streaming_iterator_hook``. Instance method 

2892 so tests can override the gate via ``MagicMock(spec=ProxyLogging)``. 

2893 """ 

2894 return ProxyLogging._callback_capabilities().has_iterator_override 

2895 

2896 def needs_per_chunk_streaming_hook(self) -> bool: 

2897 """Whether ``async_data_generator`` needs to call the per-chunk 

2898 ``_apply_streaming_chunk_hooks`` for every emitted chunk. Instance 

2899 method for the same reason as :py:meth:`needs_iterator_wrap`. 

2900 """ 

2901 caps: Final = ProxyLogging._callback_capabilities() 

2902 return caps.has_streaming_chunk_override or caps.has_guardrail 

2903 

2904 @staticmethod 

2905 def has_during_call_guardrails() -> bool: 

2906 return ProxyLogging._callback_capabilities().has_guardrail 

2907 

2908 async def during_call_hook( 

2909 self, 

2910 data: dict, 

2911 user_api_key_dict: UserAPIKeyAuth | None, 

2912 call_type: CallTypesLiteral, 

2913 ): 

2914 caps: Final = ProxyLogging._callback_capabilities() 

2915 if not caps.has_guardrail and not caps.has_moderation_override: 2915 ↛ 2918line 2915 didn't jump to line 2918 because the condition on line 2915 was always true

2916 return data 

2917 # Step 1: Collect all guardrail tasks to run in parallel 

2918 guardrail_tasks: Final = [] 

2919 

2920 for callback in litellm.callbacks: 

2921 if ( 

2922 isinstance(callback, CustomLogger) 

2923 and not isinstance(callback, CustomGuardrail) 

2924 and _overrides_moderation_hook(callback) 

2925 and user_api_key_dict is not None 

2926 ): 

2927 guardrail_tasks.append( 

2928 callback.async_moderation_hook( 

2929 data=data, 

2930 user_api_key_dict=user_api_key_dict, 

2931 call_type=call_type, 

2932 ) 

2933 ) 

2934 elif isinstance(callback, CustomGuardrail): 

2935 ################################################################ 

2936 # Check if guardrail should be run for GuardrailEventHooks.during_call hook 

2937 ################################################################ 

2938 

2939 # V1 implementation - backwards compatibility 

2940 if callback.event_hook is None and hasattr(callback, "moderation_check"): 

2941 if callback.moderation_check == "pre_call": 

2942 continue 

2943 else: 

2944 # Main - V2 Guardrails implementation 

2945 from litellm.types.guardrails import GuardrailEventHooks 

2946 

2947 event_type = GuardrailEventHooks.during_call 

2948 if call_type == CallTypes.call_mcp_tool.value: 

2949 event_type = GuardrailEventHooks.during_mcp_call 

2950 

2951 if callback.should_run_guardrail(data=data, event_type=event_type) is not True: 

2952 continue 

2953 # Convert user_api_key_dict to proper format for async_moderation_hook 

2954 if call_type == CallTypes.call_mcp_tool.value: 

2955 user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(user_api_key_dict) 

2956 else: 

2957 user_api_key_auth_dict = user_api_key_dict 

2958 guardrail_tasks.append( 

2959 self._run_during_call_guardrail( 

2960 callback=callback, 

2961 data=data, 

2962 user_api_key_dict=user_api_key_dict, 

2963 user_api_key_auth_dict=user_api_key_auth_dict, 

2964 call_type=call_type, 

2965 ) 

2966 ) 

2967 

2968 # Step 2: Run all guardrail tasks in parallel 

2969 if guardrail_tasks: 

2970 try: 

2971 await asyncio.gather(*guardrail_tasks) 

2972 except Exception as e: 

2973 # If any guardrail raises an exception, it will propagate here 

2974 raise e 

2975 

2976 return data 

2977 

2978 async def _run_during_call_guardrail( 

2979 self, 

2980 callback: CustomGuardrail, 

2981 data: dict[str, object], # mutable-ok: request payload dict, guardrail_to_apply is written in place 

2982 user_api_key_dict: UserAPIKeyAuth | None, 

2983 user_api_key_auth_dict: UserAPIKeyAuth | dict[str, object] | None, 

2984 call_type: CallTypesLiteral, 

2985 ) -> None: 

2986 if ( 

2987 "apply_guardrail" in type(callback).__dict__ 

2988 and not callback.use_native_lifecycle_hooks 

2989 and user_api_key_dict is not None 

2990 and not callback.use_native_during_call_hook 

2991 ): 

2992 data["guardrail_to_apply"] = callback 

2993 await self._run_guardrail_with_metrics( 

2994 callback, 

2995 unified_guardrail.async_moderation_hook( 

2996 user_api_key_dict=user_api_key_dict, 

2997 data=data, 

2998 call_type=call_type, 

2999 ), 

3000 "during_call", 

3001 request_data=data, 

3002 ) 

3003 return 

3004 await self._run_guardrail_with_metrics( 

3005 callback, 

3006 callback.async_moderation_hook( 

3007 data=data, 

3008 user_api_key_dict=user_api_key_auth_dict, 

3009 call_type=call_type, 

3010 ), 

3011 "during_call", 

3012 request_data=data, 

3013 ) 

3014 

3015 async def failed_tracking_alert( 

3016 self, 

3017 error_message: str, 

3018 failing_model: str, 

3019 ): 

3020 if self.alerting is None: 

3021 return 

3022 

3023 if self.slack_alerting_instance: 

3024 await self.slack_alerting_instance.failed_tracking_alert( 

3025 error_message=error_message, 

3026 failing_model=failing_model, 

3027 ) 

3028 

3029 async def budget_alerts( 

3030 self, 

3031 type: Literal[ 

3032 "token_budget", 

3033 "user_budget", 

3034 "soft_budget", 

3035 "max_budget_alert", 

3036 "team_budget", 

3037 "organization_budget", 

3038 "proxy_budget", 

3039 "projected_limit_exceeded", 

3040 "project_budget", 

3041 ], 

3042 user_info: CallInfo, 

3043 ): 

3044 # For soft_budget alerts with alert_emails set, allow email sending even if alerting is None 

3045 # This enables team-specific soft budget email alerts via metadata.soft_budget_alerting_emails 

3046 # Note: user_info is a CallInfo that can represent user/team/org level info. For team budgets, 

3047 # alert_emails is populated from team_object.metadata.soft_budget_alerting_emails (see auth_checks.py) 

3048 is_soft_budget_with_alert_emails: Final = ( 

3049 type == "soft_budget" and user_info.alert_emails is not None and len(user_info.alert_emails) > 0 

3050 ) 

3051 

3052 if self.alerting is None and not is_soft_budget_with_alert_emails: 

3053 # do nothing if alerting is not switched on (unless it's a soft_budget alert with team-specific emails) 

3054 return 

3055 

3056 if self.alerting is not None and ( 

3057 "slack" in self.alerting or "ms_teams" in self.alerting or "webhook" in self.alerting 

3058 ): 

3059 if self.slack_alerting_instance is not None: 

3060 await self.slack_alerting_instance.budget_alerts( 

3061 type=type, 

3062 user_info=user_info, 

3063 ) 

3064 

3065 # Call email_logging_instance if: 

3066 # 1. "email" is in alerting config, OR 

3067 # 2. It's a soft_budget alert with team-specific alert_emails (bypasses global alerting config) 

3068 should_send_email = (self.alerting is not None and "email" in self.alerting) or is_soft_budget_with_alert_emails 

3069 

3070 if should_send_email and self.email_logging_instance is not None: 

3071 await self.email_logging_instance.budget_alerts( 

3072 type=type, 

3073 user_info=user_info, 

3074 ) 

3075 

3076 async def alerting_handler( 

3077 self, 

3078 message: str, 

3079 level: Literal["Low", "Medium", "High"], 

3080 alert_type: AlertType, 

3081 request_data: dict | None = None, 

3082 ): 

3083 """ 

3084 Alerting based on thresholds: - https://github.com/BerriAI/litellm/issues/1298 

3085 

3086 - Responses taking too long 

3087 - Requests are hanging 

3088 - Calls are failing 

3089 - DB Read/Writes are failing 

3090 - Proxy Close to max budget 

3091 - Key Close to max budget 

3092 

3093 Parameters: 

3094 level: str - Low|Medium|High - if calls might fail (Medium) or are failing (High); Currently, no alerts would be 'Low'. 

3095 message: str - what is the alert about 

3096 """ 

3097 if self.alerting is None: 3097 ↛ 3100line 3097 didn't jump to line 3100 because the condition on line 3097 was always true

3098 return 

3099 

3100 from datetime import datetime 

3101 

3102 # Get the current timestamp 

3103 current_time: Final = datetime.now().strftime("%H:%M:%S") 

3104 _proxy_base_url: Final = os.getenv("PROXY_BASE_URL", None) 

3105 formatted_message = f"Level: `{level}`\nTimestamp: `{current_time}`\n\nMessage: {message}" 

3106 if _proxy_base_url is not None: 

3107 formatted_message += f"\n\nProxy URL: `{_proxy_base_url}`" 

3108 

3109 extra_kwargs: Final = {} 

3110 alerting_metadata = {} 

3111 if request_data is not None: 

3112 _url: Final = await _add_langfuse_trace_id_to_alert(request_data=request_data) 

3113 

3114 if _url is not None: 

3115 extra_kwargs["🪢 Langfuse Trace"] = _url 

3116 formatted_message += f"\n\n🪢 Langfuse Trace: {_url}" 

3117 if ( 

3118 "metadata" in request_data 

3119 and request_data["metadata"].get("alerting_metadata", None) is not None 

3120 and isinstance(request_data["metadata"]["alerting_metadata"], dict) 

3121 ): 

3122 alerting_metadata = request_data["metadata"]["alerting_metadata"] 

3123 if "slack" in self.alerting or "ms_teams" in self.alerting: 

3124 await self.slack_alerting_instance.send_alert( 

3125 message=message, 

3126 level=level, 

3127 alert_type=alert_type, 

3128 user_info=None, 

3129 alerting_metadata=alerting_metadata, 

3130 **extra_kwargs, 

3131 ) 

3132 for client in self.alerting: 

3133 if client == "sentry": 

3134 if litellm.utils.sentry_sdk_instance is not None: 

3135 litellm.utils.sentry_sdk_instance.capture_message(formatted_message) 

3136 else: 

3137 raise Exception("Missing SENTRY_DSN from environment") 

3138 

3139 async def failure_handler(self, original_exception, duration: float, call_type: str, traceback_str=""): 

3140 """ 

3141 Log failed db read/writes 

3142 

3143 Currently only logs exceptions to sentry 

3144 """ 

3145 ### ALERTING ### 

3146 if AlertType.db_exceptions not in self.alert_types: 3146 ↛ 3147line 3146 didn't jump to line 3147 because the condition on line 3146 was never true

3147 return 

3148 if isinstance(original_exception, HTTPException): 

3149 if isinstance(original_exception.detail, str): 3149 ↛ 3151line 3149 didn't jump to line 3151 because the condition on line 3149 was always true

3150 error_message = original_exception.detail 

3151 elif isinstance(original_exception.detail, dict): 

3152 error_message = json.dumps(original_exception.detail) 

3153 else: 

3154 error_message = str(original_exception) 

3155 else: 

3156 error_message = str(original_exception) 

3157 if isinstance(traceback_str, str): 3157 ↛ 3159line 3157 didn't jump to line 3159 because the condition on line 3157 was always true

3158 error_message += traceback_str[:1000] 

3159 error_message = _redact_string(error_message) 

3160 asyncio.create_task( 

3161 self.alerting_handler( 

3162 message=f"DB read/write call failed: {error_message}", 

3163 level="High", 

3164 alert_type=AlertType.db_exceptions, 

3165 request_data={}, 

3166 ) 

3167 ) 

3168 

3169 if hasattr(self, "service_logging_obj"): 3169 ↛ 3177line 3169 didn't jump to line 3177 because the condition on line 3169 was always true

3170 await self.service_logging_obj.async_service_failure_hook( 

3171 service=ServiceTypes.DB, 

3172 duration=duration, 

3173 error=error_message, 

3174 call_type=call_type, 

3175 ) 

3176 

3177 if litellm.utils.capture_exception: 3177 ↛ 3178line 3177 didn't jump to line 3178 because the condition on line 3177 was never true

3178 litellm.utils.capture_exception(error=original_exception) 

3179 

3180 async def post_call_failure_hook( 

3181 self, 

3182 request_data: dict, 

3183 original_exception: Exception, 

3184 user_api_key_dict: UserAPIKeyAuth, 

3185 error_type: ProxyErrorTypes | None = None, 

3186 route: str | None = None, 

3187 traceback_str: str | None = None, 

3188 ) -> HTTPException | None: 

3189 """ 

3190 Allows users to raise custom exceptions/log when a call fails, without having to deal with parsing Request body. 

3191 Callbacks can return or raise HTTPException to transform error responses sent to clients. 

3192 

3193 Covers: 

3194 1. /chat/completions 

3195 2. /embeddings 

3196 3. /image/generation 

3197 

3198 Args: 

3199 - request_data: dict - The request data. 

3200 - original_exception: Exception - The original exception. 

3201 - user_api_key_dict: UserAPIKeyAuth - The user api key dict. 

3202 - error_type: Optional[ProxyErrorTypes] - The error type. 

3203 - route: Optional[str] - The route. 

3204 - traceback_str: Optional[str] - The traceback string, sometimes upstream endpoints might need to send the upstream traceback. In which case we use this 

3205 

3206 Returns: 

3207 - Optional[HTTPException]: If any callback returns or raises an HTTPException, the first one found is returned. 

3208 Otherwise, returns None and the original exception is used. 

3209 """ 

3210 

3211 logging_obj: Final[object] = request_data.get("litellm_logging_obj") # pyright: ignore[reportUnknownVariableType, reportUnknownMemberType] # legacy request data is narrowed to Logging below 

3212 if isinstance(logging_obj, Logging) and logging_obj.baseline_cache_context is not None: 3212 ↛ 3213line 3212 didn't jump to line 3213 because the condition on line 3212 was never true

3213 await logging_obj.invalidate_baseline_cache_estimate("failed_request", completed=True) 

3214 

3215 ### ALERTING ### 

3216 await self.update_request_status(litellm_call_id=request_data.get("litellm_call_id", ""), status="fail") 

3217 if AlertType.llm_exceptions in self.alert_types and not _is_client_error_exception(original_exception): 

3218 """ 

3219 Just alert on LLM API exceptions. Do not alert on user errors 

3220 

3221 Related issue - https://github.com/BerriAI/litellm/issues/3395 

3222 """ 

3223 litellm_debug_info: Final[str | None] = getattr(original_exception, "litellm_debug_info", None) 

3224 exception_str = str(original_exception) 

3225 if litellm_debug_info is not None: 

3226 exception_str += litellm_debug_info 

3227 

3228 asyncio.create_task( 

3229 self.alerting_handler( 

3230 message=_redact_string(f"LLM API call failed: `{exception_str}`"), 

3231 level="High", 

3232 alert_type=AlertType.llm_exceptions, 

3233 request_data=request_data, 

3234 ) 

3235 ) 

3236 

3237 # Auth and pass-through failure bodies are unstripped client input, and 

3238 # the logging handler below flattens body keys into model_call_details, 

3239 # so drop the key before it can masquerade as the built payload. 

3240 request_data.pop("standard_logging_object", None) 

3241 

3242 ### LOGGING ### 

3243 if self._is_proxy_only_llm_api_error( 

3244 original_exception=original_exception, 

3245 error_type=error_type, 

3246 route=user_api_key_dict.request_route, 

3247 ): 

3248 await self._handle_logging_proxy_only_error( 

3249 request_data=request_data, 

3250 user_api_key_dict=user_api_key_dict, 

3251 route=route, 

3252 original_exception=original_exception, 

3253 ) 

3254 

3255 request_data.update(await offload_token_count(_failure_fields_to_lift)(request_data)) 

3256 

3257 # Remove before callbacks iterate — not serialisable 

3258 request_data.pop("litellm_logging_obj", None) 

3259 

3260 redacted_traceback_str: Final = _redact_string(traceback_str) if traceback_str is not None else None 

3261 

3262 # Track the first HTTPException returned or raised by any callback 

3263 transformed_exception: HTTPException | None = None 

3264 

3265 for callback in litellm.callbacks: 

3266 try: 

3267 _callback: CustomLogger | None = None 

3268 if isinstance(callback, str): 3268 ↛ 3269line 3268 didn't jump to line 3269 because the condition on line 3268 was never true

3269 _callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class( 

3270 cast(_custom_logger_compatible_callbacks_literal, callback) 

3271 ) 

3272 else: 

3273 _callback = callback 

3274 if _callback is not None and isinstance(_callback, CustomLogger): 3274 ↛ 3265line 3274 didn't jump to line 3265 because the condition on line 3274 was always true

3275 try: 

3276 hook_result = await _callback.async_post_call_failure_hook( 

3277 request_data=request_data, 

3278 user_api_key_dict=user_api_key_dict, 

3279 original_exception=original_exception, 

3280 traceback_str=redacted_traceback_str, 

3281 ) 

3282 # If callback returned an HTTPException, use it (first one wins) 

3283 if isinstance(hook_result, HTTPException) and transformed_exception is None: 3283 ↛ 3284line 3283 didn't jump to line 3284 because the condition on line 3283 was never true

3284 transformed_exception = hook_result 

3285 except HTTPException as e: 

3286 # If callback raised an HTTPException, use it (first one wins) 

3287 if transformed_exception is None: 

3288 transformed_exception = e 

3289 except Exception as e: 

3290 # Log non-HTTPException errors from callbacks but don't break the flow 

3291 verbose_proxy_logger.exception( 

3292 "[Non-Blocking] Error in async_post_call_failure_hook callback: %s", e 

3293 ) 

3294 except Exception as e: 

3295 verbose_proxy_logger.exception("[Non-Blocking] Error setting up post_call_failure_hook callback: %s", e) 

3296 

3297 return transformed_exception 

3298 

3299 def _is_proxy_only_llm_api_error( 

3300 self, 

3301 original_exception: Exception, 

3302 error_type: ProxyErrorTypes | None = None, 

3303 route: str | None = None, 

3304 ) -> bool: 

3305 """ 

3306 Return True if the error is a Proxy Only LLM API Error 

3307 

3308 Prevents double logging of LLM API exceptions 

3309 

3310 e.g should only return True for: 

3311 - Authentication Errors from user_api_key_auth 

3312 - HTTP HTTPException (rate limit errors) 

3313 - ProxyException (guardrail blocks, budget / rate-limit errors) 

3314 - GuardrailRaisedException (guardrail blocks / guardrail failures) 

3315 """ 

3316 

3317 ######################################################### 

3318 # Only log LLM API and info route errors for proxy level hooks 

3319 # eg. Authentication errors, rate limit errors, etc. 

3320 # Note: This fixes a security issue where we 

3321 # would log temporary keys/auth info 

3322 # from management endpoints 

3323 ######################################################### 

3324 if route is None: 3324 ↛ 3325line 3324 didn't jump to line 3325 because the condition on line 3324 was never true

3325 return False 

3326 if not (RouteChecks.is_llm_api_route(route) or RouteChecks.is_info_route(route)): 

3327 return False 

3328 

3329 return isinstance(original_exception, _PROXY_ONLY_LLM_API_ERRORS) or (error_type == ProxyErrorTypes.auth_error) 

3330 

3331 async def _handle_logging_proxy_only_error( 

3332 self, 

3333 request_data: dict, 

3334 user_api_key_dict: UserAPIKeyAuth, 

3335 route: str | None = None, 

3336 original_exception: Exception | None = None, 

3337 ): 

3338 """ 

3339 Handle logging for proxy only errors by calling `litellm_logging_obj.async_failure_handler` 

3340 

3341 Is triggered when self._is_proxy_only_error() returns True 

3342 """ 

3343 litellm_logging_obj: Logging | None = request_data.get("litellm_logging_obj", None) 

3344 if litellm_logging_obj is None: 

3345 from litellm._uuid import uuid 

3346 

3347 request_data.setdefault("litellm_call_id", str(uuid.uuid4())) 

3348 user_api_key_logged_metadata: Final = LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key( 

3349 user_api_key_dict=user_api_key_dict 

3350 ) 

3351 

3352 litellm_logging_obj, data = litellm.utils.function_setup( 

3353 original_function=route or "IGNORE_THIS", 

3354 rules_obj=litellm.utils.Rules(), 

3355 start_time=datetime.now(), 

3356 **request_data, 

3357 ) 

3358 request_data["litellm_logging_obj"] = litellm_logging_obj # rebind-ok: lifted then popped by the caller 

3359 if "metadata" not in request_data: 

3360 request_data["metadata"] = {} 

3361 request_data["metadata"].update(user_api_key_logged_metadata) 

3362 

3363 if litellm_logging_obj is not None: 3363 ↛ exitline 3363 didn't return from function '_handle_logging_proxy_only_error' because the condition on line 3363 was always true

3364 ## UPDATE LOGGING INPUT 

3365 _optional_params: Final = {} 

3366 _litellm_params: Final = {} 

3367 

3368 litellm_param_keys: Final = LoggedLiteLLMParams.__annotations__.keys() 

3369 for k, v in request_data.items(): 

3370 if k in litellm_param_keys: 

3371 _litellm_params[k] = v 

3372 elif k not in ("model", "user", "litellm_logging_obj"): 

3373 _optional_params[k] = v 

3374 

3375 attribution: Final = _stamp_deployment_attribution( 

3376 _litellm_params, 

3377 request_data.get("model"), 

3378 user_api_key_dict.team_id, 

3379 dispatched=_reached_deployment(litellm_logging_obj), 

3380 ) 

3381 

3382 litellm_logging_obj.update_environment_variables( 

3383 model=request_data.get("model", ""), 

3384 user=request_data.get("user", ""), 

3385 optional_params=_optional_params, 

3386 litellm_params=_litellm_params, 

3387 **( 

3388 { # mutable-ok: frozen immediately by keyword expansion 

3389 "custom_llm_provider": attribution["custom_llm_provider"] 

3390 } 

3391 if "custom_llm_provider" in attribution 

3392 else {} # mutable-ok: frozen immediately by keyword expansion 

3393 ), 

3394 ) 

3395 

3396 input: list | str | dict = "" 

3397 body_shape_call_type: str | None = None 

3398 if "messages" in request_data and isinstance(request_data["messages"], list): 

3399 input = request_data["messages"] 

3400 litellm_logging_obj.model_call_details["messages"] = input 

3401 body_shape_call_type = CallTypes.acompletion.value 

3402 elif "prompt" in request_data and isinstance(request_data["prompt"], str): 

3403 input = request_data["prompt"] 

3404 litellm_logging_obj.model_call_details["prompt"] = input 

3405 body_shape_call_type = CallTypes.atext_completion.value 

3406 elif "input" in request_data and isinstance(request_data["input"], list): 

3407 input = request_data["input"] 

3408 litellm_logging_obj.model_call_details["input"] = input 

3409 body_shape_call_type = CallTypes.aembedding.value 

3410 resolved_call_type: Final = _call_type_for_route(route) or body_shape_call_type 

3411 if resolved_call_type is not None and litellm_logging_obj.call_type != CallTypes.pass_through.value: 

3412 litellm_logging_obj.call_type = resolved_call_type 

3413 litellm_logging_obj.model_call_details["call_type"] = resolved_call_type 

3414 # Pass-through endpoints are logged via the callback loop's 

3415 # async_post_call_failure_hook — skip pre_call and failure handlers. 

3416 if litellm_logging_obj.call_type == CallTypes.pass_through.value: 

3417 return 

3418 # This is a proxy-gate error (auth/rate-limit) for a request that never 

3419 # reached a provider. ``pre_call`` below still fires every callback's 

3420 # input hook so the failure is logged — but tracing callbacks must not 

3421 # fabricate an LLM-call span for a call that did not happen (and, since 

3422 # this runs inside the live ``auth`` phase span, would otherwise nest it 

3423 # under auth). The marker tells them to skip span creation. 

3424 litellm_logging_obj.model_call_details[LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL] = True 

3425 litellm_logging_obj.pre_call( 

3426 input=input, 

3427 api_key="", 

3428 ) 

3429 

3430 await self._dispatch_proxy_only_failure_handlers( 

3431 litellm_logging_obj=litellm_logging_obj, 

3432 original_exception=original_exception, 

3433 ) 

3434 

3435 @staticmethod 

3436 async def _dispatch_proxy_only_failure_handlers( 

3437 litellm_logging_obj: Logging, 

3438 original_exception: Exception | None, 

3439 ) -> None: 

3440 """Runs the async failure handler plus the threaded sync handler. Expected 

3441 client (4xx) errors skip traceback formatting unless 

3442 litellm.log_client_error_tracebacks is set.""" 

3443 include_traceback: Final = litellm.log_client_error_tracebacks or not is_expected_client_error( 

3444 original_exception 

3445 ) 

3446 traceback_str: Final = traceback.format_exc() if include_traceback else "" 

3447 await litellm_logging_obj.async_failure_handler( 

3448 exception=original_exception, 

3449 traceback_exception=traceback_str, 

3450 ) 

3451 

3452 threading.Thread( 

3453 target=litellm_logging_obj.failure_handler, 

3454 args=( 

3455 original_exception, 

3456 traceback_str, 

3457 ), 

3458 daemon=True, 

3459 ).start() 

3460 

3461 async def _run_post_call_pipelines( 

3462 self, 

3463 data: dict[str, object], # mutable-ok: same request-payload shape as post_call_success_hook's data 

3464 user_api_key_dict: UserAPIKeyAuth, 

3465 response: LLMResponseTypes, 

3466 ) -> LLMResponseTypes | None: 

3467 if _is_pending_background_response(response): 3467 ↛ 3468line 3467 didn't jump to line 3468 because the condition on line 3467 was never true

3468 _defer_post_call_pipelines(data, response) 

3469 return None 

3470 _, pipeline_response = await self._maybe_execute_pipelines( 

3471 data=data, 

3472 user_api_key_dict=user_api_key_dict, 

3473 call_type=getattr(data.get("litellm_logging_obj"), "call_type", None) or "acompletion", 

3474 event_hook="post_call", 

3475 response=response, 

3476 ) 

3477 return pipeline_response 

3478 

3479 async def post_call_success_hook( 

3480 self, 

3481 data: dict, 

3482 response: LLMResponseTypes, 

3483 user_api_key_dict: UserAPIKeyAuth, 

3484 ): 

3485 """ 

3486 Allow user to modify outgoing data 

3487 

3488 Covers: 

3489 1. /chat/completions 

3490 2. /embeddings 

3491 3. /image/generation 

3492 4. /files 

3493 """ 

3494 

3495 from litellm.proxy.proxy_server import llm_router 

3496 from litellm.types.guardrails import GuardrailEventHooks 

3497 

3498 pipeline_response: Final = await self._run_post_call_pipelines( 

3499 data=data, 

3500 user_api_key_dict=user_api_key_dict, 

3501 response=response, 

3502 ) 

3503 if pipeline_response is not None: 3503 ↛ 3504line 3503 didn't jump to line 3504 because the condition on line 3503 was never true

3504 response = pipeline_response # rebind-ok: adopt the pipeline's replacement response, same contract as the callback loops below 

3505 

3506 pipeline_managed: Final = pipeline_managed_guardrail_names(data, "post_call") 

3507 guardrail_callbacks, other_callbacks = _partition_post_call_callbacks() 

3508 try: 

3509 # Merge model-level guardrails before checking which guardrails to run 

3510 guardrail_data: Final = _check_and_merge_model_level_guardrails(data=data, llm_router=llm_router) 

3511 

3512 parallel_guardrails: Final[tuple[CustomGuardrail, ...]] = tuple( 

3513 callback 

3514 for callback in guardrail_callbacks 

3515 if getattr(callback, "run_in_parallel", False) 

3516 and not (callback.guardrail_name and callback.guardrail_name in pipeline_managed) 

3517 ) 

3518 

3519 for callback in guardrail_callbacks: 3519 ↛ 3522line 3519 didn't jump to line 3522 because the loop on line 3519 never started

3520 # Main - V2 Guardrails implementation 

3521 

3522 if callback.guardrail_name and callback.guardrail_name in pipeline_managed: 

3523 continue 

3524 

3525 if getattr(callback, "run_in_parallel", False): 

3526 continue 

3527 

3528 if ( 

3529 callback.should_run_guardrail( 

3530 data=guardrail_data, 

3531 event_type=GuardrailEventHooks.post_call, 

3532 ) 

3533 is not True 

3534 ): 

3535 continue 

3536 

3537 guardrail_response: Any | None = None 

3538 

3539 if "apply_guardrail" in type(callback).__dict__ and not callback.use_native_lifecycle_hooks: 

3540 data["guardrail_to_apply"] = callback 

3541 guardrail_response = await self._run_guardrail_with_metrics( 

3542 callback, 

3543 unified_guardrail.async_post_call_success_hook( 

3544 user_api_key_dict=user_api_key_dict, 

3545 data=data, 

3546 response=response, 

3547 ), 

3548 "post_call", 

3549 request_data=data, 

3550 ) 

3551 else: 

3552 guardrail_response = await self._run_guardrail_with_metrics( 

3553 callback, 

3554 callback.async_post_call_success_hook( 

3555 user_api_key_dict=user_api_key_dict, 

3556 data=data, 

3557 response=response, 

3558 ), 

3559 "post_call", 

3560 request_data=data, 

3561 ) 

3562 

3563 if guardrail_response is not None: 

3564 response = guardrail_response 

3565 

3566 if parallel_guardrails: 3566 ↛ 3567line 3566 didn't jump to line 3567 because the condition on line 3566 was never true

3567 await self._run_parallel_post_call_guardrails( 

3568 guardrails=parallel_guardrails, 

3569 data=data, 

3570 guardrail_data=guardrail_data, 

3571 response=response, 

3572 user_api_key_dict=user_api_key_dict, 

3573 ) 

3574 

3575 ############ Handle CustomLogger ############################### 

3576 ################################################################# 

3577 

3578 for callback in other_callbacks: 

3579 callback_response: LLMResponseTypes | None = await callback.async_post_call_success_hook( 

3580 user_api_key_dict=user_api_key_dict, data=data, response=response 

3581 ) 

3582 if callback_response is not None: 

3583 response = callback_response 

3584 except Exception as e: 

3585 raise e 

3586 return response 

3587 

3588 async def _run_parallel_post_call_guardrails( 

3589 self, 

3590 guardrails: tuple[CustomGuardrail, ...], 

3591 data: dict, 

3592 guardrail_data: dict, 

3593 response: LLMResponseTypes, 

3594 user_api_key_dict: UserAPIKeyAuth, 

3595 ) -> None: 

3596 """ 

3597 Run opted-in post_call guardrails concurrently against the response 

3598 produced by the sequential guardrails. These guardrails are declared 

3599 block-only, so any modified response they return is discarded; they run 

3600 for their blocking side effect (raising to reject the response before it 

3601 reaches the client). Every guardrail is awaited to completion 

3602 (``return_exceptions=True``) so a raise by one never leaves the others 

3603 running as unobserved background tasks. A guardrail that blocks (any 

3604 exception other than a passthrough) takes precedence over one that only 

3605 changes the response flow, so a fast passthrough can never let a slower 

3606 block be bypassed. Each per-guardrail coroutine sets ``guardrail_to_apply`` 

3607 immediately before awaiting, and the unified hook pops it before its first 

3608 suspension point, so concurrent guardrails never race on that key. 

3609 """ 

3610 

3611 async def _run_one(callback: CustomGuardrail) -> None: 

3612 if callback.should_run_guardrail(data=guardrail_data, event_type=GuardrailEventHooks.post_call) is not True: 

3613 return 

3614 if "apply_guardrail" in type(callback).__dict__ and not callback.use_native_lifecycle_hooks: 

3615 data["guardrail_to_apply"] = callback 

3616 await self._run_guardrail_with_metrics( 

3617 callback, 

3618 unified_guardrail.async_post_call_success_hook( 

3619 user_api_key_dict=user_api_key_dict, 

3620 data=data, 

3621 response=response, 

3622 ), 

3623 "post_call", 

3624 request_data=data, 

3625 ) 

3626 else: 

3627 await self._run_guardrail_with_metrics( 

3628 callback, 

3629 callback.async_post_call_success_hook( 

3630 user_api_key_dict=user_api_key_dict, 

3631 data=data, 

3632 response=response, 

3633 ), 

3634 "post_call", 

3635 request_data=data, 

3636 ) 

3637 

3638 results: Final = await asyncio.gather( 

3639 *(_run_one(callback) for callback in guardrails), 

3640 return_exceptions=True, 

3641 ) 

3642 raised: Final = tuple(result for result in results if isinstance(result, BaseException)) 

3643 blocking: Final = next((exc for exc in raised if not _exception_changes_request_flow(exc)), None) 

3644 if blocking is not None: 

3645 raise blocking 

3646 if raised: 

3647 raise raised[0] 

3648 

3649 async def post_mcp_call_hook( 

3650 self, 

3651 response: "CallToolResult", 

3652 request_data: Mapping[str, Any], 

3653 user_api_key_dict: UserAPIKeyAuth | None = None, 

3654 ) -> "CallToolResult": 

3655 """ 

3656 Run guardrails configured for ``post_mcp_call`` against an MCP tool result. 

3657 

3658 The MCP counterpart of ``post_call_success_hook``: guardrails that 

3659 implement ``apply_guardrail`` see the tool result's text through the 

3660 unified guardrail seam (``MCPGuardrailTranslationHandler``), so a text 

3661 guardrail can mask sensitive values in the result without any MCP-specific 

3662 code of its own. Guardrails that instead implement 

3663 ``async_post_mcp_tool_call_hook`` are dispatched by 

3664 ``Logging.async_post_mcp_tool_call_hook`` and are not run here. 

3665 

3666 A guardrail that rejects the result raises, and the exception propagates 

3667 (matching the inbound ``pre_mcp_call`` behavior) rather than being 

3668 swallowed into an unguarded result. 

3669 """ 

3670 caps: Final = ProxyLogging._callback_capabilities() 

3671 if not caps.has_guardrail: 

3672 return response 

3673 

3674 handler_cls: Final = load_guardrail_translation_mappings().get(CallTypes.call_mcp_tool) 

3675 if handler_cls is None: 

3676 verbose_proxy_logger.debug("MCP guardrail translation handler unavailable; skipping post_mcp_call hook") 

3677 return response 

3678 

3679 for callback in caps.resolved_callbacks: 

3680 if not isinstance(callback, CustomGuardrail): 

3681 continue 

3682 if "apply_guardrail" not in type(callback).__dict__ or callback.use_native_lifecycle_hooks: 

3683 continue 

3684 if ( 

3685 callback.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_mcp_call) 

3686 is not True 

3687 ): 

3688 continue 

3689 response = await self._run_guardrail_with_metrics( 

3690 callback, 

3691 handler_cls().process_output_response( 

3692 response=response, 

3693 guardrail_to_apply=callback, 

3694 litellm_logging_obj=request_data.get("litellm_logging_obj"), 

3695 user_api_key_dict=user_api_key_dict, 

3696 request_data=request_data, 

3697 ), 

3698 "post_mcp_call", 

3699 request_data=request_data, 

3700 ) 

3701 return response 

3702 

3703 async def post_call_response_headers_hook( 

3704 self, 

3705 data: dict, 

3706 user_api_key_dict: UserAPIKeyAuth, 

3707 response: object, 

3708 request_headers: dict[str, str] | None = None, 

3709 ) -> dict[str, str]: 

3710 """ 

3711 Calls async_post_call_response_headers_hook on all CustomLogger callbacks. 

3712 Merges all returned header dicts (later callbacks override earlier ones). 

3713 

3714 Returns: 

3715 Dict[str, str]: Merged headers from all callbacks. 

3716 """ 

3717 merged_headers: Final[dict[str, str]] = {} 

3718 # Outer call sites in common_request_processing.py already gate this 

3719 # call with ``has_post_call_response_headers_callbacks()``. The 

3720 # cached detection makes the redundant interior guard cheap, but the 

3721 # guard would still iterate every code path through this function so 

3722 # keep it cheap and rely on the cached capability lookup. 

3723 if not ProxyLogging._callback_capabilities().has_post_call_response_headers: 

3724 return merged_headers 

3725 

3726 try: 

3727 # Build litellm_call_info — normalized routing metadata for callbacks 

3728 litellm_call_info: Final = self._build_litellm_call_info(data=data, response=response) 

3729 

3730 for callback in litellm.callbacks: 

3731 _callback: CustomLogger | None = None 

3732 if isinstance(callback, str): 3732 ↛ 3733line 3732 didn't jump to line 3733 because the condition on line 3732 was never true

3733 _callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class( 

3734 cast(_custom_logger_compatible_callbacks_literal, callback) 

3735 ) 

3736 else: 

3737 _callback = callback 

3738 

3739 if _callback is not None and isinstance(_callback, CustomLogger): 3739 ↛ 3730line 3739 didn't jump to line 3730 because the condition on line 3739 was always true

3740 if _accepts_litellm_call_info(_callback): 3740 ↛ 3750line 3740 didn't jump to line 3750 because the condition on line 3740 was always true

3741 result = await _callback.async_post_call_response_headers_hook( 

3742 data=data, 

3743 user_api_key_dict=user_api_key_dict, 

3744 response=response, 

3745 request_headers=request_headers, 

3746 litellm_call_info=litellm_call_info, 

3747 ) 

3748 else: 

3749 # Backwards compat: callback doesn't accept litellm_call_info 

3750 result = await _callback.async_post_call_response_headers_hook( 

3751 data=data, 

3752 user_api_key_dict=user_api_key_dict, 

3753 response=response, 

3754 request_headers=request_headers, 

3755 ) 

3756 if result is not None: 3756 ↛ 3757line 3756 didn't jump to line 3757 because the condition on line 3756 was never true

3757 merged_headers.update(result) 

3758 except Exception as e: 

3759 verbose_proxy_logger.exception("Error in post_call_response_headers_hook: %s", str(e)) 

3760 return merged_headers 

3761 

3762 async def hidden_by_listing_callbacks( 

3763 self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str] 

3764 ) -> frozenset[str]: 

3765 filters: Final = ProxyLogging._callback_capabilities().listed_models_filters 

3766 if not filters: 3766 ↛ 3768line 3766 didn't jump to line 3768 because the condition on line 3766 was always true

3767 return frozenset() 

3768 candidates: Final = tuple(model_names) 

3769 kept: Final = await _names_kept_by_listing_callbacks(filters, user_api_key_dict, candidates) 

3770 if isinstance(kept, MalformedListingFilterReturn): 

3771 _raise_malformed_listing_filter_return(kept) 

3772 return frozenset(candidates).difference(kept) 

3773 

3774 @staticmethod 

3775 def _build_litellm_call_info(data: dict, response: object) -> dict[str, object]: 

3776 """ 

3777 Build a normalized dict of routing metadata from response._hidden_params 

3778 and data, abstracting away the metadata vs litellm_metadata split. 

3779 """ 

3780 hidden_params: Final = getattr(response, "_hidden_params", {}) or {} 

3781 

3782 # model_info: check both metadata keys (chat uses "metadata", responses uses "litellm_metadata") 

3783 model_info: Final = ( 

3784 (data.get("metadata") or {}).get("model_info") 

3785 or (data.get("litellm_metadata") or {}).get("model_info") 

3786 or {} 

3787 ) 

3788 

3789 return { 

3790 "custom_llm_provider": hidden_params.get("custom_llm_provider") 

3791 or getattr(response, "custom_llm_provider", None), 

3792 "model_info": model_info, 

3793 "api_base": hidden_params.get("api_base"), 

3794 "model_id": hidden_params.get("model_id"), 

3795 } 

3796 

3797 def is_a2a_streaming_response(self, response: dict) -> bool: 

3798 expected_keys: Final = ["jsonrpc", "id", "result"] 

3799 return all(key in response for key in expected_keys) 

3800 

3801 async def async_post_call_streaming_hook( 

3802 self, 

3803 data: dict, 

3804 response: ModelResponse | EmbeddingResponse | ImageResponse | ModelResponseStream, 

3805 user_api_key_dict: UserAPIKeyAuth, 

3806 str_so_far: str | None = None, 

3807 ): 

3808 """ 

3809 Allow user to modify outgoing streaming data -> per chunk 

3810 

3811 Covers: 

3812 1. /chat/completions 

3813 """ 

3814 # Per-chunk fast path: skip the response-string materialization and 

3815 # callback scan when no configured callback overrides 

3816 # ``async_post_call_streaming_hook`` AND no CustomGuardrail is 

3817 # active. ``get_response_string`` walks every choice/delta on the 

3818 # chunk so paying it per chunk for no-op callbacks dominated stream 

3819 # CPU time even after the iterator-chain fix. 

3820 caps: Final = ProxyLogging._callback_capabilities() 

3821 if not caps.has_streaming_chunk_override and not caps.has_guardrail: 

3822 return response 

3823 

3824 from litellm.proxy.proxy_server import llm_router 

3825 

3826 response_str: str | None = None 

3827 if isinstance(response, (ModelResponse, ModelResponseStream)): 

3828 response_str = litellm.get_response_string(response_obj=response) 

3829 elif isinstance(response, dict) and self.is_a2a_streaming_response(response): 

3830 from litellm.llms.a2a.common_utils import extract_text_from_a2a_response 

3831 

3832 response_str = extract_text_from_a2a_response(response) 

3833 if response_str is not None: 

3834 # Cache model-level guardrails check per-request to avoid repeated 

3835 # dict lookups + llm_router.get_deployment() per callback per chunk. 

3836 _cached_guardrail_data: dict | None = None 

3837 _guardrail_data_computed = False 

3838 pipeline_gated: Final = ( 

3839 stream_gated_guardrail_names(data, user_api_key_dict) if caps.has_guardrail else frozenset() 

3840 ) 

3841 

3842 for callback in litellm.callbacks: 

3843 try: 

3844 _callback: CustomLogger | None = None 

3845 if isinstance(callback, CustomGuardrail): 

3846 if callback.guardrail_name in pipeline_gated: 

3847 continue 

3848 # Main - V2 Guardrails implementation 

3849 from litellm.types.guardrails import GuardrailEventHooks 

3850 

3851 ## CHECK FOR MODEL-LEVEL GUARDRAILS (cached per-request) 

3852 if not _guardrail_data_computed: 

3853 _cached_guardrail_data = _check_and_merge_model_level_guardrails( 

3854 data=data, llm_router=llm_router 

3855 ) 

3856 _guardrail_data_computed = True 

3857 

3858 if ( 

3859 callback.should_run_guardrail( 

3860 data=_cached_guardrail_data, 

3861 event_type=GuardrailEventHooks.post_call, 

3862 ) 

3863 is not True 

3864 ): 

3865 continue 

3866 if isinstance(callback, str): 

3867 _callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class( 

3868 cast(_custom_logger_compatible_callbacks_literal, callback) 

3869 ) 

3870 else: 

3871 _callback = callback 

3872 if _callback is not None and isinstance(_callback, CustomLogger): 

3873 if str_so_far is not None: 

3874 complete_response = str_so_far + response_str 

3875 else: 

3876 complete_response = response_str 

3877 callback_response: ( 

3878 ModelResponse | EmbeddingResponse | ImageResponse | ModelResponseStream | None 

3879 ) 

3880 callback_response = await _callback.async_post_call_streaming_hook( 

3881 user_api_key_dict=user_api_key_dict, 

3882 response=complete_response, 

3883 ) 

3884 if callback_response is not None: 

3885 response = callback_response 

3886 except Exception as e: 

3887 raise e 

3888 return response 

3889 

3890 async def async_post_call_streaming_iterator_hook( 

3891 self, 

3892 response, 

3893 user_api_key_dict: UserAPIKeyAuth, 

3894 request_data: dict, 

3895 ): 

3896 """ 

3897 Allow user to modify outgoing streaming data -> Given a whole response iterator. 

3898 This hook is best used when you need to modify multiple chunks of the response at once. 

3899 

3900 Covers: 

3901 1. /chat/completions 

3902 """ 

3903 caps: Final = ProxyLogging._callback_capabilities() 

3904 post_call_pipelines: Final = _streamable_post_call_pipelines(request_data, user_api_key_dict) 

3905 # Fast path: no real overrides. Internal proxy CustomLogger callbacks 

3906 # (e.g. _PROXY_CacheControlCheck, ManagedFiles) inherit the default 

3907 # ``async for chunk: yield chunk`` body, so wrapping the iterator 

3908 # through each of them adds N pass-through trampolines per chunk for 

3909 # zero behavior change. Skip the chain entirely and stream through. 

3910 if not caps.iterator_overrides and not post_call_pipelines: 

3911 try: 

3912 async for chunk in response: 

3913 yield chunk 

3914 except (GeneratorExit, asyncio.CancelledError): 

3915 raise 

3916 except Exception as e: 

3917 if not ProxyLogging._discard_deferred_stream_logging_for_failure(request_data, e): 

3918 ProxyLogging._fire_deferred_stream_logging(request_data) 

3919 raise 

3920 ProxyLogging._fire_deferred_stream_logging(request_data) 

3921 return 

3922 

3923 from litellm.proxy.proxy_server import llm_router 

3924 

3925 # Merge model-level guardrails before checking which guardrails to run 

3926 request_data = _check_and_merge_model_level_guardrails(data=request_data, llm_router=llm_router) 

3927 

3928 current_response = response 

3929 stream_needs_translation: Final = ProxyLogging._stream_requires_guardrail_translation(user_api_key_dict) 

3930 

3931 pipeline_gated_names: Final = _pipeline_step_guardrail_names(post_call_pipelines) 

3932 for resolved_callback, kind in caps.iterator_overrides: 

3933 if isinstance(resolved_callback, CustomGuardrail): 

3934 if resolved_callback.guardrail_name in pipeline_gated_names: 

3935 continue 

3936 if ( 

3937 resolved_callback.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) 

3938 is not True 

3939 ): 

3940 continue 

3941 effective_kind = ( 

3942 "apply_guardrail" 

3943 if ( 

3944 kind == "override" 

3945 and stream_needs_translation 

3946 and isinstance(resolved_callback, CustomGuardrail) 

3947 and resolved_callback.uses_apply_guardrail_interface() 

3948 and getattr(resolved_callback, "use_native_lifecycle_hooks", False) is not True 

3949 and not resolved_callback.mask_response_content 

3950 ) 

3951 else kind 

3952 ) 

3953 hook: _StreamIteratorHook[object] = ( 

3954 partial( 

3955 resolved_callback.async_post_call_streaming_iterator_hook, 

3956 user_api_key_dict=user_api_key_dict, 

3957 request_data=request_data, 

3958 ) 

3959 if effective_kind == "override" 

3960 else partial( 

3961 unified_guardrail.async_post_call_streaming_iterator_hook, 

3962 user_api_key_dict=user_api_key_dict, 

3963 request_data=request_data, 

3964 guardrail_to_apply=resolved_callback, 

3965 buffer_until_moderated_default=(kind == "override"), 

3966 ) 

3967 ) 

3968 current_response = self._wrap_streaming_iterator_with_enrichment( 

3969 resolved_callback, 

3970 current_response, 

3971 hook, 

3972 request_data=request_data, 

3973 ) 

3974 

3975 pipeline_translation: Final = ( 

3976 resolve_endpoint_translation(user_api_key_dict, None) if post_call_pipelines else None 

3977 ) 

3978 if pipeline_translation is not None: 

3979 current_response = self._pipeline_gated_stream( 

3980 response=current_response, 

3981 user_api_key_dict=user_api_key_dict, 

3982 request_data=request_data, 

3983 pipelines=post_call_pipelines, 

3984 translation=pipeline_translation, 

3985 ) 

3986 

3987 served_chunks: Final[list[object]] = [] # mutable-ok: accumulates while yielding to the client 

3988 try: 

3989 async for chunk in current_response: 

3990 served_chunks.append(chunk) 

3991 yield chunk 

3992 except (GeneratorExit, asyncio.CancelledError): 

3993 ProxyLogging._record_served_stream_output(request_data, served_chunks) 

3994 raise 

3995 except Exception as e: 

3996 ProxyLogging._record_served_stream_output(request_data, served_chunks) 

3997 if not ProxyLogging._discard_deferred_stream_logging_for_failure(request_data, e): 

3998 ProxyLogging._fire_deferred_stream_logging(request_data) 

3999 raise 

4000 

4001 # Fire deferred logging AFTER all guardrail end-of-stream blocks 

4002 # completed. unified_guardrail writes guardrail_information during 

4003 # its end-of-stream block (inside current_response), so by the time 

4004 # we reach this point the metadata is fully populated. 

4005 ProxyLogging._record_served_stream_output(request_data, served_chunks) 

4006 ProxyLogging._fire_deferred_stream_logging(request_data) 

4007 

4008 async def _pipeline_gated_stream( 

4009 self, 

4010 response: "AsyncGenerator[object, None]", 

4011 user_api_key_dict: UserAPIKeyAuth, 

4012 request_data: dict, # mutable-ok: same request-payload shape the hooks mutate 

4013 pipelines: "tuple[tuple[str, GuardrailPipeline], ...]", 

4014 translation: "tuple[str, BaseTranslation]", 

4015 ) -> "AsyncGenerator[object, None]": 

4016 """ 

4017 Execute post_call policy pipelines against a streamed response. 

4018 

4019 Buffers the whole stream (nothing reaches the client until every 

4020 pipeline allows it), then runs each pipeline's steps against the 

4021 assembled output through the endpoint guardrail translation, the same 

4022 machinery flat post_call guardrails use at end of stream. An allow 

4023 releases the buffered chunks: verbatim when no guardrail rewrote the 

4024 output, rewritten in place when one rewrote text or a tool call and the 

4025 translation delivers ended-stream rewrites (later steps then re-scan the 

4026 rewritten chunks, so rewrites chain). A rewrite the translation cannot 

4027 deliver yet (one on a route without write-back, or a shape the route 

4028 refuses) is discarded by the executor and the original chunks are 

4029 released; a block or modify_response terminates with the translation's 

4030 block chunks or the raised error. 

4031 """ 

4032 buffered: Final[list[object]] = [] # mutable-ok: accumulates the stream before the pipeline verdict 

4033 async for item in response: 

4034 buffered.append(item) 

4035 if not buffered: 

4036 return 

4037 

4038 call_type, endpoint_translation = translation 

4039 

4040 for policy_name, pipeline in pipelines: 

4041 result: PipelineExecutionResult = await PipelineExecutor.execute_steps( 

4042 steps=pipeline.steps, 

4043 mode="post_call", 

4044 data=request_data, 

4045 user_api_key_dict=user_api_key_dict, 

4046 call_type=call_type, 

4047 policy_name=policy_name, 

4048 streaming_chunks=buffered, 

4049 endpoint_translation=endpoint_translation, 

4050 ) 

4051 try: 

4052 ProxyLogging._handle_pipeline_result( 

4053 result, data=request_data, policy_name=policy_name, original_response=buffered 

4054 ) 

4055 except ModifyResponseException as e: 

4056 if e.original_response is None: 

4057 e.original_response = buffered 

4058 async for block_chunk in unified_guardrail.handle_streaming_block( 

4059 e, endpoint_translation, stream_started=False, responses_so_far=() 

4060 ): 

4061 yield block_chunk 

4062 return 

4063 except HTTPException as e: 

4064 async for error_chunk in unified_guardrail.emit_streaming_http_error( 

4065 e, call_type, buffered, request_data 

4066 ): 

4067 yield error_chunk 

4068 return 

4069 

4070 for buffered_item in buffered: 

4071 yield buffered_item 

4072 

4073 @staticmethod 

4074 def _record_served_stream_output(request_data: Mapping[str, object], served_chunks: Sequence[object]) -> None: 

4075 logging_obj: Final = request_data.get("litellm_logging_obj") 

4076 if not isinstance(logging_obj, Logging): 

4077 return 

4078 record_served_output_texts(logging_obj.model_call_details, served_stream_output_texts(served_chunks)) 

4079 

4080 @staticmethod 

4081 def _fire_deferred_stream_logging(request_data: dict) -> None: 

4082 """ 

4083 Fire the deferred streaming logging callback after the full streaming 

4084 pipeline (including guardrail end-of-stream blocks) has completed. 

4085 

4086 CSW.__anext__ stores the callback and args on logging_obj instead of 

4087 scheduling via create_task (which would race with unified_guardrail's 

4088 end-of-stream block). This method retrieves and fires them. 

4089 """ 

4090 logging_obj: Final = request_data.get("litellm_logging_obj") 

4091 if logging_obj is None: 

4092 return 

4093 _deferred_cb: Final[Callable[..., Coroutine[object, object, object]] | None] = getattr( 

4094 logging_obj, "_on_deferred_stream_complete", None 

4095 ) 

4096 _args: Final[tuple[object, ...] | None] = getattr(logging_obj, "_deferred_stream_complete_args", None) 

4097 if _deferred_cb is not None and _args is not None: 

4098 logging_obj._on_deferred_stream_complete = None 

4099 logging_obj._deferred_stream_complete_args = None 

4100 asyncio.create_task(_deferred_cb(*_args)) 

4101 

4102 @staticmethod 

4103 def _discard_deferred_stream_logging_for_failure(request_data: Mapping[str, object], error: Exception) -> bool: 

4104 """Drop the parked success dispatch for an assembled chat stream that ends in an error 

4105 ``post_call_failure_hook`` logs as a failure, billing its usage on the failure row instead. 

4106 Returns False when the parked dispatch should still be flushed by the caller.""" 

4107 logging_obj: Final = request_data.get("litellm_logging_obj") 

4108 if not isinstance(logging_obj, Logging): 

4109 return False 

4110 _args: Final[tuple[object, ...] | None] = getattr(logging_obj, "_deferred_stream_complete_args", None) 

4111 assembled: Final = _args[0] if _args else None 

4112 if not isinstance(error, _PROXY_ONLY_LLM_API_ERRORS) or not isinstance(assembled, ModelResponse): 

4113 return False 

4114 logging_obj._on_deferred_stream_complete = None 

4115 logging_obj._deferred_stream_complete_args = None 

4116 logging_obj.record_assembled_response_for_failure(assembled) 

4117 return True 

4118 

4119 async def _arelease_max_parallel_requests_on_disconnect( 

4120 self, 

4121 user_api_key_dict: UserAPIKeyAuth, 

4122 ) -> None: 

4123 """ 

4124 Release the api-key max_parallel_requests slot when a streaming 

4125 response is cancelled mid-flight (client disconnect) and no logging 

4126 callback fired for it. Neither the success nor failure callback runs on 

4127 the resulting CancelledError / GeneratorExit, so the pre-call +1 would 

4128 otherwise leak. 

4129 

4130 Awaited from the shielded streaming cleanup rather than scheduled 

4131 fire-and-forget, so the caller can make it the single owner of the 

4132 release: when a disconnect-time success event does fire (partial-spend 

4133 billing or a deferred-guardrail flush), that event's own limiter 

4134 callback releases the slot and this is not called at all. Two 

4135 concurrent releases of the same acquisition would otherwise race and 

4136 double-decrement under the limiter's in-memory fallback. 

4137 """ 

4138 limiter: Final = self.get_proxy_hook("parallel_request_limiter") 

4139 if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): 

4140 return 

4141 await limiter.async_release_max_parallel_requests_on_disconnect(user_api_key_dict) 

4142 

4143 def _init_response_taking_too_long_task(self, data: dict | None = None): 

4144 """ 

4145 Initialize the response taking too long task if user is using slack alerting 

4146 

4147 Only run task if user is using slack alerting 

4148 

4149 This handles checking for if a request is hanging for too long 

4150 """ 

4151 ## ALERTING ### 

4152 if self.slack_alerting_instance and self.slack_alerting_instance.alerting is not None: 4152 ↛ 4153line 4152 didn't jump to line 4153 because the condition on line 4152 was never true

4153 asyncio.create_task(self.slack_alerting_instance.response_taking_too_long(request_data=data)) 

4154 

4155 

4156### DB CONNECTOR ### 

4157# Define the retry decorator with backoff strategy 

4158# Function to be called whenever a retry is about to happen 

4159def on_backoff(details): 

4160 # The 'tries' key in the details dictionary contains the number of completed tries 

4161 print_verbose(f"Backing off... this was attempt #{details['tries']}") 

4162 

4163 

4164def jsonify_object(data: dict) -> dict: 

4165 db_data: Final = copy.deepcopy(data) 

4166 

4167 for k, v in db_data.items(): 

4168 if isinstance(v, dict): 

4169 try: 

4170 db_data[k] = json.dumps(v) 

4171 except Exception: 

4172 # This avoids Prisma retrying this 5 times, and making 5 clients 

4173 db_data[k] = "failed-to-serialize-json" 

4174 return db_data 

4175 

4176 

4177# In-memory cache for deprecated key lookups: 

4178# maps old_token_hash -> (active_token_id, cache_expires_at_ts, revoke_at_ts). 

4179# Avoids a DB query on every auth request for non-deprecated keys. 

4180# Bounded to prevent memory leaks from accumulated rotations. 

4181_deprecated_key_cache: Final[LimitedSizeOrderedDict] = LimitedSizeOrderedDict(max_size=1000) 

4182_DEPRECATED_KEY_CACHE_TTL_SECONDS: Final = 60 

4183_PRISMA_DEFAULT_TX_TIMEOUT: Final = timedelta(seconds=5) 

4184 

4185 

4186async def _lookup_deprecated_key( 

4187 db: PrismaWrapper | RoutingPrismaWrapper, 

4188 hashed_token: str, 

4189) -> str | None: 

4190 """ 

4191 Check if a token exists in the deprecated keys table and is still within its grace period. 

4192 

4193 Returns the active_token_id if found and valid, otherwise None. 

4194 Uses an in-memory cache to avoid DB queries on every auth request. 

4195 """ 

4196 now: Final = datetime.now(timezone.utc) 

4197 now_ts: Final = now.timestamp() 

4198 

4199 # Check cache first 

4200 cached: Final = _deprecated_key_cache.get(hashed_token) 

4201 if cached is not None: 4201 ↛ 4202line 4201 didn't jump to line 4202 because the condition on line 4201 was never true

4202 active_token_id, cache_expires_at_ts, revoke_at_ts = cached 

4203 if now_ts < cache_expires_at_ts and now_ts < revoke_at_ts: 

4204 return active_token_id 

4205 _deprecated_key_cache.pop(hashed_token, None) 

4206 

4207 try: 

4208 deprecated_keys_table: Final[ 

4209 LiteLLM_DeprecatedVerificationTokenActions[LiteLLM_DeprecatedVerificationToken] 

4210 ] = db.litellm_deprecatedverificationtoken 

4211 deprecated_row: Final = await deprecated_keys_table.find_first( 

4212 where={ 

4213 "token": hashed_token, 

4214 "revoke_at": {"gt": now}, 

4215 } 

4216 ) 

4217 if deprecated_row and deprecated_row.active_token_id: 4217 ↛ 4218line 4217 didn't jump to line 4218 because the condition on line 4217 was never true

4218 revoke_at: Final = deprecated_row.revoke_at 

4219 _deprecated_key_cache[hashed_token] = ( 

4220 deprecated_row.active_token_id, 

4221 now_ts + _DEPRECATED_KEY_CACHE_TTL_SECONDS, 

4222 revoke_at.timestamp(), 

4223 ) 

4224 return deprecated_row.active_token_id 

4225 # Only cache positive results; negative lookups are fast on indexed columns 

4226 # and caching them risks evicting real deprecated key entries. 

4227 except Exception as e: 

4228 verbose_proxy_logger.debug("Deprecated key lookup skipped: %s", e) 

4229 

4230 return None 

4231 

4232 

4233# DualCache for LiteLLM_Config param_name reads. 

4234# Redis layer is attached in proxy_server._init_cache. 

4235LITELLM_CONFIG_CACHE_TTL_SECONDS: Final[int] = int(os.environ.get("LITELLM_CONFIG_PARAM_CACHE_TTL_SECONDS", "60")) 

4236_CONFIG_CACHE_MISS: Final[str] = "__litellm_config_param_miss__" 

4237 

4238litellm_config_cache: Final[DualCache] = DualCache( 

4239 default_in_memory_ttl=LITELLM_CONFIG_CACHE_TTL_SECONDS, 

4240 default_redis_ttl=LITELLM_CONFIG_CACHE_TTL_SECONDS, 

4241) 

4242 

4243 

4244class _ConfigRow: 

4245 """Mimics the Prisma litellm_config row shape for cached entries.""" 

4246 

4247 __slots__ = ("param_name", "param_value") 

4248 

4249 def __init__(self, param_name: str, param_value: object) -> None: 

4250 self.param_name = param_name 

4251 self.param_value = param_value 

4252 

4253 

4254def _config_cache_key(param_name: str) -> str: 

4255 return f"litellm_config:param:{param_name}" 

4256 

4257 

4258def _pack_config_row(row: Any) -> dict[str, object]: 

4259 return {"param_name": row.param_name, "param_value": row.param_value} 

4260 

4261 

4262def _unpack_config_row(cached: object) -> _ConfigRow | None: 

4263 if cached is None or cached == _CONFIG_CACHE_MISS: 

4264 return None 

4265 if isinstance(cached, dict): 4265 ↛ 4267line 4265 didn't jump to line 4267 because the condition on line 4265 was always true

4266 return _ConfigRow(cached["param_name"], cached["param_value"]) 

4267 return None 

4268 

4269 

4270async def get_config_param(prisma_client: "PrismaClient", param_name: str) -> Any | None: 

4271 """Cached read of a LiteLLM_Config row; returns row, _ConfigRow shim, or None.""" 

4272 cache_key: Final = _config_cache_key(param_name) 

4273 cached: Final = await litellm_config_cache.async_get_cache(cache_key) 

4274 if cached is not None: 

4275 return _unpack_config_row(cached) 

4276 

4277 row: Final = await prisma_client.get_generic_data(key="param_name", value=param_name, table_name="config") 

4278 cache_value: Final[Mapping[str, object] | str] = _pack_config_row(row) if row is not None else _CONFIG_CACHE_MISS 

4279 await litellm_config_cache.async_set_cache(cache_key, cache_value, ttl=LITELLM_CONFIG_CACHE_TTL_SECONDS) 

4280 return row 

4281 

4282 

4283async def evict_config_param(param_name: str) -> None: 

4284 await litellm_config_cache.async_delete_cache(_config_cache_key(param_name)) 

4285 

4286 

4287async def invalidate_config_param(param_name: str) -> None: 

4288 """Evict from both cache layers; call after every LiteLLM_Config write.""" 

4289 await evict_config_param(param_name) 

4290 await publish_config_param_change(param_name) 

4291 

4292 

4293async def prefetch_config_params(prisma_client: "PrismaClient | None", param_names: list[str]) -> None: 

4294 """Batch-load LiteLLM_Config rows into the cache with one find_many.""" 

4295 if not param_names: 4295 ↛ 4296line 4295 didn't jump to line 4296 because the condition on line 4295 was never true

4296 return 

4297 try: 

4298 config_table: Final = cast( # cast-ok: ConfigRepository.table is prisma's litellm_config actions object 

4299 "TableActions[prisma_models.LiteLLM_Config]", ConfigRepository(prisma_client).table 

4300 ) 

4301 rows: Final = await config_table.find_many(where={"param_name": {"in": param_names}}) 

4302 except Exception as e: 

4303 verbose_proxy_logger.debug( 

4304 "prefetch_config_params failed, falling through to per-param queries: %s", 

4305 e, 

4306 ) 

4307 return 

4308 by_name: Final = {row.param_name: row for row in rows} 

4309 for name in param_names: 

4310 row = by_name.get(name) 

4311 cache_value: Mapping[str, object] | str = _pack_config_row(row) if row is not None else _CONFIG_CACHE_MISS 

4312 await litellm_config_cache.async_set_cache( 

4313 _config_cache_key(name), cache_value, ttl=LITELLM_CONFIG_CACHE_TTL_SECONDS 

4314 ) 

4315 

4316 

4317_WRITER_WRITABILITY_PROBE_SQL: Final = "SELECT current_setting('transaction_read_only') AS transaction_read_only" 

4318_WRITER_WRITABILITY_PROBE_ROWS: Final = TypeAdapter(list[dict[str, object]]) 

4319_READ_ONLY_RECREATE_BACKOFF_CAP_SECONDS: Final = 600 

4320 

4321 

4322class _ForcedRecreateDeclined(Exception): 

4323 """A forced recreate was declined by the engine-generation guard. 

4324 

4325 Distinct from a reconnect *failure*: the machinery worked, it just found 

4326 that another path had already replaced the writer, so it left the engines 

4327 alone. The caller's engine may still be poisoned, so the cycle must not 

4328 report success, but it must not count as a failure either, or the record 

4329 of what could not be repaired would gate the retry that recovers. 

4330 """ 

4331 

4332 

4333@dataclass(frozen=True, slots=True) 

4334class _StaleReadEngine: 

4335 """The read engine a query observed, identified rather than only counted. 

4336 

4337 `PrismaClient.read_db` resolves to the reader while it is available and to 

4338 the writer once it is not, and the two carry independent generation 

4339 counters that both start at zero and advance on the same reconnect 

4340 cadence. A bare generation compared across that switch would silently pit 

4341 one engine's counter against another's, so the wrapper is carried with the 

4342 number and a switch counts as the engine having moved. 

4343 

4344 Holding the wrapper itself rather than its `id()` is load-bearing, not 

4345 incidental: the strong reference keeps the wrapper alive, so its address 

4346 cannot be recycled under a stored observation and match an unrelated 

4347 engine later. It is only free because writer and reader both live as long 

4348 as the client does; a replaceable reader would make this a retention leak. 

4349 """ 

4350 

4351 wrapper: PrismaWrapper 

4352 generation: int 

4353 

4354 @classmethod 

4355 def observe(cls, wrapper: PrismaWrapper) -> "_StaleReadEngine": 

4356 return cls(wrapper=wrapper, generation=wrapper.engine_generation) 

4357 

4358 def is_still_live(self, current: PrismaWrapper) -> bool: 

4359 """Whether this exact engine is still serving reads, unreplaced. 

4360 

4361 A True answer must never be the only thing standing between a poisoned 

4362 engine and its repair. The generation moves only after a replacement 

4363 connects, and a recreate whose connect raises leaves it unmoved until 

4364 some later recreate succeeds, so this can report an engine as live 

4365 after it has stopped working. What bounds that is the failed-repair 

4366 record in `_cooldown_applies`, written by a repair attempt that fails 

4367 rather than by whatever broke the engine: the two need not be the same 

4368 recreate, since the synchronous token-refresh fallback in 

4369 `PrismaWrapper.__getattr__` recreates outside the reconnect machinery 

4370 and records nothing. The record is written only for callers that named 

4371 an engine, and it collapses the rest of the burst for up to one 

4372 cooldown window rather than guaranteeing a repair, since the cooldown 

4373 conjunct underneath it still expires and lets a later caller retry. 

4374 """ 

4375 return self.wrapper is current and self.generation == current.engine_generation 

4376 

4377 

4378class PrismaClient: 

4379 spend_log_transactions: list = [] 

4380 _spend_log_transactions_lock = asyncio.Lock() 

4381 spend_log_flush_requested: "asyncio.Event | None" = None 

4382 spend_log_queue_bytes: ClassVar[int] = 0 

4383 spend_logs_queue_monitor_task: "asyncio.Task[None] | None" = None 

4384 spend_log_write_lock = asyncio.Lock() 

4385 tool_usage_transactions: list["ToolUsageTransaction"] = [] 

4386 _tool_usage_transactions_lock = asyncio.Lock() 

4387 autorouter_turn_transactions: ClassVar[ 

4388 list["AutoRouterTurnTransaction"] 

4389 ] = [] # mutable-ok: drained queue, mirrors tool_usage_transactions 

4390 _autorouter_turn_transactions_lock = asyncio.Lock() 

4391 

4392 # How long a health probe failure waits for an in-flight planned engine 

4393 # replacement to settle before deciding whether to report itself. Generous 

4394 # against a replacement that takes well under a second, and far short of the 

4395 # reconnect budget an outage-hung `connect()` runs under, so a real outage 

4396 # is never waited out. 

4397 PLANNED_ENGINE_REPLACEMENT_SETTLE_SECONDS: ClassVar[float] = 5.0 

4398 

4399 def __init__( 

4400 self, 

4401 database_url: str, 

4402 proxy_logging_obj: ProxyLogging, 

4403 http_client: "HttpConfig | None" = None, 

4404 ): 

4405 ## init logging object 

4406 self.baseline_accounting_transactions: list[ 

4407 BaselineAccountingRecord 

4408 ] = [] # mutable-ok: locked background queue 

4409 self.baseline_accounting_lock: Final = asyncio.Lock() 

4410 self.proxy_logging_obj = proxy_logging_obj 

4411 self.token_auth: DatabaseTokenAuth | None = resolve_database_token_auth() 

4412 verbose_proxy_logger.debug("Creating Prisma Client..") 

4413 try: 

4414 from prisma import Prisma 

4415 from prisma.types import DatasourceOverride 

4416 except Exception as e: 

4417 verbose_proxy_logger.error("Failed to import Prisma client: %s", e) 

4418 verbose_proxy_logger.error("This usually means 'prisma generate' hasn't been run yet.") 

4419 verbose_proxy_logger.error("Please run 'prisma generate' to generate the Prisma client.") 

4420 raise Exception("Unable to find Prisma binaries. Please run 'prisma generate' first.") 

4421 token_auth: Final = self.token_auth 

4422 writer_token_auth: Final = None if database_url_is_pooled() else token_auth 

4423 # When read-replica routing is on, tag log lines with [writer]/[reader] 

4424 # so the two wrappers' interleaved token refresh logs can be told apart. 

4425 # Single-DB deployments get an empty prefix (logs unchanged). 

4426 read_replica_url = os.getenv("DATABASE_URL_READ_REPLICA") 

4427 writer_log_prefix: Final = "[writer]" if read_replica_url else "" 

4428 if http_client is not None: 4428 ↛ 4429line 4428 didn't jump to line 4429 because the condition on line 4428 was never true

4429 writer_wrapper = PrismaWrapper( 

4430 original_prisma=Prisma(http=http_client), 

4431 token_auth=writer_token_auth, 

4432 log_prefix=writer_log_prefix, 

4433 ) 

4434 else: 

4435 writer_wrapper = PrismaWrapper( 

4436 original_prisma=Prisma(), 

4437 token_auth=writer_token_auth, 

4438 log_prefix=writer_log_prefix, 

4439 ) 

4440 

4441 # Optional read-replica routing. When DATABASE_URL_READ_REPLICA is set, 

4442 # reads (find_*, count, group_by, query_raw/_first) are routed to the 

4443 # reader endpoint and writes stay on the writer. Falls back to the 

4444 # writer-only wrapper when the env var is unset, preserving existing 

4445 # single-DB deployments. 

4446 self.db: PrismaWrapper | RoutingPrismaWrapper 

4447 if read_replica_url: 4447 ↛ 4448line 4447 didn't jump to line 4448 because the condition on line 4447 was never true

4448 try: 

4449 # If token auth is enabled, the reader refreshes its own token on 

4450 # the same cadence as the writer. We parse the static endpoint 

4451 # pieces (host/port/user/db) once from the reader URL — only 

4452 # the token rotates after that. 

4453 reader_iam_endpoint: Final = ( 

4454 parse_iam_endpoint_from_url(read_replica_url) if token_auth is not None else None 

4455 ) 

4456 # Mint a fresh token for the reader BEFORE constructing the 

4457 # Prisma client. Mirrors what `proxy_cli.py` already does for 

4458 # the writer — without this, the reader Prisma is built with 

4459 # whatever placeholder URL the user supplied (no real token), 

4460 # and the first query falls through to the synchronous fallback 

4461 # path in `PrismaWrapper.__getattr__`, which deadlocks the event 

4462 # loop and times out after 30s. 

4463 if token_auth is not None and reader_iam_endpoint is not None: 

4464 reader_token: Final = mint_database_token(token_auth, reader_iam_endpoint) 

4465 read_replica_url = add_missing_query_params( 

4466 reader_iam_endpoint.build_url(reader_token), 

4467 token_refresh_params_from_url(read_replica_url), 

4468 ) 

4469 os.environ["DATABASE_URL_READ_REPLICA"] = read_replica_url 

4470 reader_datasource: Final = DatasourceOverride(url=read_replica_url) 

4471 if http_client is not None: 

4472 reader_prisma = Prisma(http=http_client, datasource=reader_datasource) 

4473 else: 

4474 reader_prisma = Prisma(datasource=reader_datasource) 

4475 reader_wrapper: Final = PrismaWrapper( 

4476 original_prisma=reader_prisma, 

4477 token_auth=token_auth, 

4478 db_url_env_var="DATABASE_URL_READ_REPLICA", 

4479 iam_endpoint=reader_iam_endpoint, 

4480 recreate_uses_datasource=True, 

4481 log_prefix="[reader]", 

4482 ) 

4483 self.db = RoutingPrismaWrapper(writer=writer_wrapper, reader=reader_wrapper) 

4484 verbose_proxy_logger.info( 

4485 "PrismaClient: read-replica routing enabled via DATABASE_URL_READ_REPLICA" 

4486 + (f" (with {token_auth.label} auto-refresh)" if token_auth is not None else "") 

4487 ) 

4488 except Exception as e: 

4489 # Reader is opt-in; never let its construction fail proxy 

4490 # startup. Mirrors the runtime contract from 

4491 # `RoutingPrismaWrapper.connect`: reader-side failures are 

4492 # logged and we keep serving traffic via the writer alone. 

4493 # This recovers from transient credential-provider hiccups 

4494 # during the reader token mint, malformed DATABASE_URL_READ_REPLICA, 

4495 # and Prisma construction errors. Operator restart is required 

4496 # to retry read-routing once the underlying issue is resolved. 

4497 verbose_proxy_logger.warning( 

4498 "Failed to initialize read replica Prisma client: %s. " 

4499 "Falling back to writer-only mode (no read routing) until proxy restart.", 

4500 e, 

4501 ) 

4502 self.db = writer_wrapper 

4503 else: 

4504 self.db = writer_wrapper # Client to connect to Prisma db 

4505 self._db_reconnect_lock = asyncio.Lock() 

4506 self._db_health_watchdog_task: asyncio.Task | None = None 

4507 self._view_setup_task: asyncio.Task[_ViewSetupOutcome] | None = None 

4508 self._db_last_reconnect_attempt_ts: float = 0.0 

4509 self._db_reconnect_cooldown_seconds: int = max(1, int(os.getenv("PRISMA_RECONNECT_COOLDOWN_SECONDS", "15"))) 

4510 self._db_read_only_recreate_ts: float = 0.0 

4511 self._db_read_only_recreate_streak: int = 0 

4512 self._db_health_watchdog_interval_seconds: int = max( 

4513 5, int(os.getenv("PRISMA_HEALTH_WATCHDOG_INTERVAL_SECONDS", "30")) 

4514 ) 

4515 self._db_health_watchdog_enabled: bool = ( 

4516 str_to_bool(os.getenv("PRISMA_HEALTH_WATCHDOG_ENABLED", "true")) is True 

4517 ) 

4518 self._db_health_watchdog_probe_timeout_seconds: float = max( 

4519 0.5, 

4520 float(os.getenv("PRISMA_HEALTH_WATCHDOG_PROBE_TIMEOUT_SECONDS", "5.0")), 

4521 ) 

4522 self._db_watchdog_reconnect_timeout_seconds: float = max( 

4523 1.0, float(os.getenv("PRISMA_WATCHDOG_RECONNECT_TIMEOUT_SECONDS", "30.0")) 

4524 ) 

4525 self._db_auth_reconnect_timeout_seconds: float = max( 

4526 0.5, float(os.getenv("PRISMA_AUTH_RECONNECT_TIMEOUT_SECONDS", "2.0")) 

4527 ) 

4528 self._db_auth_reconnect_lock_timeout_seconds: float = max( 

4529 0.0, 

4530 float(os.getenv("PRISMA_AUTH_RECONNECT_LOCK_TIMEOUT_SECONDS", "0.1")), 

4531 ) 

4532 self._consecutive_reconnect_failures: int = 0 

4533 # Last generation of each read engine whose repair was attempted and 

4534 # failed. Scoped to the engine rather than counted globally so an 

4535 # unrelated reconnect failure cannot suppress a stale reader's 

4536 # recovery, and keyed per wrapper rather than held in one slot so a 

4537 # writer failure cannot evict the reader's record and hand the waiver 

4538 # back to a caller whose engine is still unrepaired. Bounded at two 

4539 # entries: a client has one writer and at most one reader. 

4540 self._failed_recreate_generations: Mapping[PrismaWrapper, int] = MappingProxyType({}) 

4541 self._reconnect_escalation_threshold: int = max(1, int(os.getenv("PRISMA_RECONNECT_ESCALATION_THRESHOLD", "3"))) 

4542 self._engine_pidfd: int = -1 

4543 self._engine_pid: int = 0 

4544 self._watching_engine: bool = False 

4545 self._engine_confirmed_dead: bool = False 

4546 self._engine_wait_thread: threading.Thread | None = None 

4547 verbose_proxy_logger.debug("Success - Created Prisma Client") 

4548 

4549 @property 

4550 def writer_db(self) -> PrismaWrapper: 

4551 """Underlying writer Prisma wrapper, regardless of read-replica routing.""" 

4552 if isinstance(self.db, RoutingPrismaWrapper): 4552 ↛ 4553line 4552 didn't jump to line 4553 because the condition on line 4552 was never true

4553 return self.db.writer 

4554 return self.db 

4555 

4556 @property 

4557 def read_db(self) -> PrismaWrapper: 

4558 """Underlying wrapper that top-level reads are dispatched to. 

4559 

4560 Identical to `writer_db` without a read replica. With one configured 

4561 it is the reader, which is the engine `query_first` actually runs on, 

4562 so anything reasoning about the state of the connection that served a 

4563 read has to consult this rather than the writer. 

4564 """ 

4565 if isinstance(self.db, RoutingPrismaWrapper): 4565 ↛ 4566line 4565 didn't jump to line 4566 because the condition on line 4565 was never true

4566 return self.db.read_target 

4567 return self.db 

4568 

4569 def tx(self, *, timeout: timedelta = _PRISMA_DEFAULT_TX_TIMEOUT) -> "TransactionManager": 

4570 """Open an interactive transaction on the writer. 

4571 

4572 Callers go through this instead of reaching into ``self.db`` so writer 

4573 selection and read-replica routing stay encapsulated in the wrapper. 

4574 """ 

4575 return cast("TransactionManager", self.db.tx(timeout=timeout)) # cast-ok: untyped __getattr__ delegate 

4576 

4577 def get_request_status(self, payload: dict | SpendLogsPayload) -> Literal["success", "failure"]: 

4578 """ 

4579 Determine if a request was successful or failed based on payload metadata. 

4580 

4581 Args: 

4582 payload (Union[dict, SpendLogsPayload]): Request payload containing metadata 

4583 

4584 Returns: 

4585 Literal["success", "failure"]: Request status 

4586 """ 

4587 try: 

4588 # Get metadata and convert to dict if it's a JSON string 

4589 payload_metadata: Final[dict | SpendLogsMetadata | str] = payload.get("metadata", {}) 

4590 if isinstance(payload_metadata, str): 4590 ↛ 4593line 4590 didn't jump to line 4593 because the condition on line 4590 was always true

4591 payload_metadata_json: dict | SpendLogsMetadata = cast(dict, json.loads(payload_metadata)) 

4592 else: 

4593 payload_metadata_json = payload_metadata 

4594 

4595 # Check status in metadata dict 

4596 return "failure" if payload_metadata_json.get("status") == "failure" else "success" 

4597 

4598 except (json.JSONDecodeError, AttributeError): 

4599 # Default to success if metadata parsing fails 

4600 return "success" 

4601 

4602 def hash_token(self, token: str): 

4603 # Hash the string using SHA-256 

4604 hashed_token: Final = hashlib.sha256(token.encode()).hexdigest() 

4605 

4606 return hashed_token 

4607 

4608 def jsonify_object(self, data: Mapping[str, object]) -> dict[str, object]: 

4609 db_data: Final[dict[str, object]] = copy.deepcopy(dict(data)) 

4610 

4611 for k, v in db_data.items(): 

4612 if isinstance(v, dict): 

4613 try: 

4614 db_data[k] = json.dumps(v) 

4615 except Exception: 

4616 # This avoids Prisma retrying this 5 times, and making 5 clients 

4617 db_data[k] = "failed-to-serialize-json" 

4618 return db_data 

4619 

4620 @backoff.on_exception( 

4621 backoff.expo, 

4622 Exception, # base exception to catch for the backoff 

4623 max_tries=3, # maximum number of retries 

4624 max_time=10, # maximum total time to retry for 

4625 on_backoff=on_backoff, # specifying the function to call on backoff 

4626 ) 

4627 async def check_view_exists(self): 

4628 """ 

4629 Checks if the LiteLLM_VerificationTokenView and MonthlyGlobalSpend exists in the user's db. 

4630 

4631 LiteLLM_VerificationTokenView: This view is used for getting the token + team data in user_api_key_auth 

4632 

4633 MonthlyGlobalSpend: This view is used for the admin view to see global spend for this month 

4634 

4635 If the view doesn't exist, one will be created. 

4636 """ 

4637 

4638 # Check to see if all of the necessary views exist and if they do, simply return 

4639 # This is more efficient because it lets us check for all views in one 

4640 # query instead of multiple queries. 

4641 try: 

4642 expected_views: Final = [ 

4643 "LiteLLM_VerificationTokenView", 

4644 "MonthlyGlobalSpend", 

4645 "Last30dKeysBySpend", 

4646 "Last30dModelsBySpend", 

4647 "MonthlyGlobalSpendPerKey", 

4648 "MonthlyGlobalSpendPerUserPerKey", 

4649 "Last30dTopEndUsersSpend", 

4650 "DailyTagSpend", 

4651 ] 

4652 required_view: Final = "LiteLLM_VerificationTokenView" 

4653 expected_views_str: Final = ", ".join(f"'{view}'" for view in expected_views) 

4654 pg_schema: Final = os.getenv("DATABASE_SCHEMA", "public") 

4655 ret: Final[Sequence[_ViewCountRow]] = await self.db.query_raw(f""" 

4656 WITH existing_views AS ( 

4657 SELECT viewname 

4658 FROM pg_views 

4659 WHERE schemaname = '{pg_schema}' AND viewname IN ( 

4660 {expected_views_str} 

4661 ) 

4662 ) 

4663 SELECT 

4664 (SELECT COUNT(*) FROM existing_views) AS view_count, 

4665 ARRAY_AGG(viewname) AS view_names 

4666 FROM existing_views 

4667 """) 

4668 expected_total_views: Final = len(expected_views) 

4669 if ret[0]["view_count"] == expected_total_views: 4669 ↛ 4670line 4669 didn't jump to line 4670 because the condition on line 4669 was never true

4670 verbose_proxy_logger.info("All necessary views exist!") 

4671 return 

4672 else: 

4673 ## check if required view exists ## 

4674 if ret[0]["view_names"] and required_view not in ret[0]["view_names"]: 4674 ↛ 4675line 4674 didn't jump to line 4675 because the condition on line 4674 was never true

4675 await self.health_check() # make sure we can connect to db 

4676 await create_view_tolerating_race( 

4677 self.db, 

4678 "LiteLLM_VerificationTokenView", 

4679 """ 

4680 CREATE VIEW "LiteLLM_VerificationTokenView" AS 

4681 SELECT 

4682 v.*, 

4683 t.spend AS team_spend, 

4684 t.max_budget AS team_max_budget, 

4685 t.model_max_budget AS team_model_max_budget, 

4686 t.tpm_limit AS team_tpm_limit, 

4687 t.rpm_limit AS team_rpm_limit, 

4688 t.tpd_limit AS team_tpd_limit 

4689 FROM "LiteLLM_VerificationToken" v 

4690 LEFT JOIN "LiteLLM_TeamTable" t ON v.team_id = t.team_id; 

4691 """, 

4692 ) 

4693 else: 

4694 should_create_views: Final = await should_create_missing_views(db=self.db) 

4695 if should_create_views: 4695 ↛ 4700line 4695 didn't jump to line 4700 because the condition on line 4695 was always true

4696 await create_missing_views(db=self.db) 

4697 else: 

4698 # don't block execution if these views are missing 

4699 # Convert lists to sets for efficient difference calculation 

4700 ret_view_names_set: Final = set(ret[0]["view_names"]) if ret[0]["view_names"] else set() 

4701 expected_views_set: Final = set(expected_views) 

4702 # Find missing views 

4703 missing_views: Final = expected_views_set - ret_view_names_set 

4704 

4705 verbose_proxy_logger.warning( 

4706 "\n\n\x1b[93mNot all views exist in db, needed for UI 'Usage' tab. Missing=%s.\nRun 'create_views.py' from https://github.com/BerriAI/litellm/tree/main/db_scripts to create missing views.\x1b[0m\n", 

4707 missing_views, 

4708 ) 

4709 

4710 except Exception: 

4711 raise 

4712 return 

4713 

4714 @log_db_metrics 

4715 @backoff.on_exception( 

4716 backoff.expo, 

4717 Exception, # base exception to catch for the backoff 

4718 max_tries=1, # maximum number of retries 

4719 max_time=2, # maximum total time to retry for 

4720 on_backoff=on_backoff, # specifying the function to call on backoff 

4721 ) 

4722 async def get_generic_data( 

4723 self, 

4724 key: str, 

4725 value: object, 

4726 table_name: Literal["users", "keys", "config", "spend"], 

4727 ): 

4728 """ 

4729 Generic implementation of get data. 

4730 

4731 Self-heals across a single transient transport blip via 

4732 `call_with_db_reconnect_retry`: on `httpx.ReadError` / 

4733 `ClientNotConnectedError` / similar, attempt one DB reconnect and 

4734 retry once before surfacing the failure. Restores the 1.82.6 behavior 

4735 that was lost in 1.83.x — see issue #25143. 

4736 """ 

4737 start_time: Final = time.time() 

4738 

4739 async def _do_query(): 

4740 if table_name == "users": 4740 ↛ 4741line 4740 didn't jump to line 4741 because the condition on line 4740 was never true

4741 return await UserRepository(self).table.find_first(where={key: value}) 

4742 elif table_name == "keys": 4742 ↛ 4743line 4742 didn't jump to line 4743 because the condition on line 4742 was never true

4743 return await VerificationTokenRepository(self).table.find_first(where={key: value}) 

4744 elif table_name == "config": 4744 ↛ 4749line 4744 didn't jump to line 4749 because the condition on line 4744 was always true

4745 config_table: Final = cast( # cast-ok: ConfigRepository.table is prisma's litellm_config actions object 

4746 "TableActions[prisma_models.LiteLLM_Config]", ConfigRepository(self).table 

4747 ) 

4748 return await config_table.find_first(where={key: value}) 

4749 elif table_name == "spend": 

4750 return await self.db.l.find_first(where={key: value}) 

4751 return None 

4752 

4753 try: 

4754 return await call_with_db_reconnect_retry( 

4755 self, 

4756 _do_query, 

4757 reason=f"prisma_get_generic_data_{table_name}_lookup_failure", 

4758 ) 

4759 except Exception as e: 

4760 error_msg = f"LiteLLM Prisma Client Exception get_generic_data: {e}" 

4761 verbose_proxy_logger.error(error_msg) 

4762 error_msg = error_msg + f"\nException Type: {type(e)}" 

4763 error_traceback: Final = error_msg + "\n" + traceback.format_exc() 

4764 end_time: Final = time.time() 

4765 _duration: Final = end_time - start_time 

4766 asyncio.create_task( 

4767 self.proxy_logging_obj.failure_handler( 

4768 original_exception=e, 

4769 duration=_duration, 

4770 traceback_str=error_traceback, 

4771 call_type="get_generic_data", 

4772 ) 

4773 ) 

4774 

4775 raise e 

4776 

4777 async def _query_first_with_cached_plan_fallback(self, sql_query: str, *args) -> dict | None: 

4778 """ 

4779 Execute a query, recovering once from PostgreSQL's "cached plan must not 

4780 change result type" error. 

4781 

4782 That error surfaces during rolling deployments when a schema change 

4783 invalidates the prepared-statement plans that pooled connections still 

4784 hold. Clearing only the server-side plans with DEALLOCATE ALL makes 

4785 things worse: Prisma's query engine keeps a per-connection client-side 

4786 cache of prepared-statement names, so once the server drops a plan the 

4787 engine re-sends a name PostgreSQL no longer recognizes and the 

4788 connection breaks with `prepared statement "sN" does not exist`. With a 

4789 small pool that connection stays poisoned and every auth lookup fails. 

4790 

4791 Recreating the Prisma client kills the engine subprocess and drops the 

4792 server-side plans and the engine's client-side name cache together, so 

4793 the retried query is prepared fresh. We reconnect through 

4794 `attempt_db_reconnect`, which is singleflight: when a schema change 

4795 poisons every pooled connection at once, the first cached-plan error 

4796 recreates the client and the concurrent waiters reuse that single 

4797 recreate instead of racing to kill each other's fresh engine. We pass 

4798 `force_recreate` so the reconnect skips its `SELECT 1` liveness probe: 

4799 the connection is healthy here, it is the prepared statements on it 

4800 that are stale, so a passing probe would otherwise skip the recreate 

4801 and leave the retry to hit the same error. We then retry the identical 

4802 query exactly once. 

4803 

4804 The retry reuses the original query byte-for-byte. Mutating the SQL 

4805 (e.g. injecting a unique comment) would defeat PostgreSQL's plan cache, 

4806 forcing a fresh plan on every request and pegging the database CPU. 

4807 

4808 The reconnect cooldown must not gate the engine this query itself saw 

4809 as stale, or a migration landing within the cooldown of an earlier 

4810 reconnect leaves auth failing until it elapses. The engine observed 

4811 before the query names it, so the reconnect bypasses the cooldown only 

4812 while that same engine is still the live one. 

4813 

4814 It is observed from `read_db`, not `writer_db`: `query_first` is a 

4815 top-level read, so with a read replica configured it runs on the reader 

4816 and it is the reader's prepared statements that went stale. Naming the 

4817 writer here would let an unrelated writer reconnect re-arm the cooldown 

4818 while the reader stayed poisoned. 

4819 """ 

4820 stale_read_engine: Final = _StaleReadEngine.observe(self.read_db) 

4821 try: 

4822 return await self.db.query_first(sql_query, *args) 

4823 except Exception as e: 

4824 if "cached plan must not change result type" not in str(e): 

4825 raise 

4826 verbose_proxy_logger.warning( 

4827 "PostgreSQL cached plan error detected for token lookup; " 

4828 "recreating the database connection and retrying with the same " 

4829 "query. This may occur during rolling deployments when schema " 

4830 "changes are applied." 

4831 ) 

4832 await self.attempt_db_reconnect( 

4833 reason="postgres_cached_plan_error", 

4834 force_recreate=True, 

4835 stale_read_engine=stale_read_engine, 

4836 ) 

4837 return await self.db.query_first(sql_query, *args) 

4838 

4839 @backoff.on_exception( 

4840 backoff.expo, 

4841 Exception, # base exception to catch for the backoff 

4842 max_tries=3, # maximum number of retries 

4843 max_time=10, # maximum total time to retry for 

4844 on_backoff=on_backoff, # specifying the function to call on backoff 

4845 ) 

4846 @log_db_metrics 

4847 async def get_data( 

4848 self, 

4849 token: str | list | None = None, 

4850 user_id: str | None = None, 

4851 user_id_list: Sequence[str] | None = None, 

4852 team_id: str | None = None, 

4853 team_id_list: Sequence[str] | None = None, 

4854 key_val: dict | None = None, 

4855 table_name: Literal[ 

4856 "user", "key", "config", "spend", "enduser", "budget", "team", "user_notification", "combined_view" 

4857 ] 

4858 | None = None, 

4859 query_type: Literal["find_unique", "find_all"] = "find_unique", 

4860 expires: datetime | None = None, 

4861 reset_at: datetime | None = None, 

4862 offset: int | None = None, # pagination, what row number to start from 

4863 limit: int | None = None, # pagination, number of rows to getch when find_all==True 

4864 parent_otel_span: Span | None = None, 

4865 proxy_logging_obj: ProxyLogging | None = None, 

4866 budget_id_list: list[str] | None = None, 

4867 check_deprecated: bool = True, 

4868 ): 

4869 args_passed_in: Final = locals() 

4870 start_time: Final = time.time() 

4871 hashed_token: str | None = None 

4872 try: 

4873 response: Any = None 

4874 if (token is not None and table_name is None) or (table_name is not None and table_name == "key"): 

4875 # check if plain text or hash 

4876 if token is not None: 

4877 if isinstance(token, str): 4877 ↛ 4880line 4877 didn't jump to line 4880 because the condition on line 4877 was always true

4878 hashed_token = _hash_token_if_needed(token=token) 

4879 verbose_proxy_logger.debug("PrismaClient: find_unique for token: %s", hashed_token) 

4880 if query_type == "find_unique" and hashed_token is not None: 

4881 if token is None: 4881 ↛ 4882line 4881 didn't jump to line 4882 because the condition on line 4881 was never true

4882 raise HTTPException( 

4883 status_code=400, 

4884 detail={"error": f"No token passed in. Token={token}"}, 

4885 ) 

4886 response = await VerificationTokenRepository(self).table.find_unique( 

4887 where={"token": hashed_token}, 

4888 include={"litellm_budget_table": True}, 

4889 ) 

4890 if response is not None: 4890 ↛ 4892line 4890 didn't jump to line 4892 because the condition on line 4890 was never true

4891 # for prisma we need to cast the expires time to str 

4892 if response.expires is not None and isinstance(response.expires, datetime): 

4893 response.expires = response.expires.isoformat() 

4894 else: 

4895 # Token does not exist. 

4896 raise HTTPException( 

4897 status_code=status.HTTP_401_UNAUTHORIZED, 

4898 detail=f"Authentication Error: invalid user key - user key does not exist in db. User Key={token}", 

4899 ) 

4900 elif query_type == "find_all" and user_id is not None: 

4901 response = await VerificationTokenRepository(self).table.find_many( 

4902 where={"user_id": user_id}, 

4903 include={"litellm_budget_table": True}, 

4904 ) 

4905 if response is not None and len(response) > 0: 4905 ↛ 4906line 4905 didn't jump to line 4906 because the condition on line 4905 was never true

4906 for r in response: 

4907 if isinstance(r.expires, datetime): 

4908 r.expires = r.expires.isoformat() 

4909 elif query_type == "find_all" and team_id is not None: 

4910 response = await VerificationTokenRepository(self).table.find_many( 

4911 take=limit, 

4912 where={"team_id": team_id}, 

4913 include={"litellm_budget_table": True}, 

4914 ) 

4915 if response is not None and len(response) > 0: 4915 ↛ 4916line 4915 didn't jump to line 4916 because the condition on line 4915 was never true

4916 for r in response: 

4917 if isinstance(r.expires, datetime): 

4918 r.expires = r.expires.isoformat() 

4919 elif query_type == "find_all" and expires is not None and reset_at is not None: 4919 ↛ 4935line 4919 didn't jump to line 4935 because the condition on line 4919 was always true

4920 response = await VerificationTokenRepository(self).table.find_many( 

4921 take=limit, 

4922 where={ 

4923 "OR": [ 

4924 {"expires": None}, 

4925 {"expires": {"gt": expires}}, 

4926 ], 

4927 "budget_reset_at": {"lt": reset_at}, 

4928 "NOT": {"budget_duration": None}, 

4929 }, 

4930 ) 

4931 if response is not None and len(response) > 0: 4931 ↛ 4932line 4931 didn't jump to line 4932 because the condition on line 4931 was never true

4932 for r in response: 

4933 if isinstance(r.expires, datetime): 

4934 r.expires = r.expires.isoformat() 

4935 elif query_type == "find_all": 

4936 where_filter: Final[dict[str, dict[str, Sequence[str]]]] = {} 

4937 if token is not None: 

4938 where_filter["token"] = {} 

4939 if isinstance(token, str): 

4940 token = _hash_token_if_needed(token=token) 

4941 where_filter["token"]["in"] = [token] 

4942 elif isinstance(token, list): 

4943 hashed_tokens: Final[list[str]] = [] 

4944 for t in token: 

4945 assert isinstance(t, str) 

4946 if t.startswith("sk-"): 

4947 new_token = self.hash_token(token=t) 

4948 hashed_tokens.append(new_token) 

4949 else: 

4950 hashed_tokens.append(t) 

4951 where_filter["token"]["in"] = hashed_tokens 

4952 response = await VerificationTokenRepository(self).table.find_many( 

4953 order={"spend": "desc"}, 

4954 where=where_filter, 

4955 include={"litellm_budget_table": True}, 

4956 ) 

4957 if response is not None: 4957 ↛ 4961line 4957 didn't jump to line 4961 because the condition on line 4957 was always true

4958 return response 

4959 else: 

4960 # Token does not exist. 

4961 raise HTTPException( 

4962 status_code=status.HTTP_401_UNAUTHORIZED, 

4963 detail="Authentication Error: invalid user key - token does not exist", 

4964 ) 

4965 elif (user_id is not None and table_name is None) or (table_name is not None and table_name == "user"): 

4966 if query_type == "find_unique": 

4967 if key_val is None: 4967 ↛ 4970line 4967 didn't jump to line 4970 because the condition on line 4967 was always true

4968 key_val = {"user_id": user_id} 

4969 

4970 response = await UserRepository(self).table.find_unique( 

4971 where=key_val, 

4972 include={"organization_memberships": True}, 

4973 ) 

4974 

4975 elif query_type == "find_all" and key_val is not None: 

4976 response = await UserRepository(self).table.find_many(where=key_val) 

4977 elif query_type == "find_all" and reset_at is not None: 4977 ↛ 4997line 4977 didn't jump to line 4997 because the condition on line 4977 was always true

4978 response = await UserRepository(self).table.find_many( 

4979 take=limit, 

4980 where={ 

4981 # A user seeded from default_internal_user_params 

4982 # (or created via /user/new without an explicit 

4983 # budget_reset_at) has budget_duration set but 

4984 # budget_reset_at = NULL. `{"lt": reset_at}` never 

4985 # matches NULL, so such users would never be reset 

4986 # and their spend would accumulate for the lifetime 

4987 # of the row, silently exceeding max_budget. Treat a 

4988 # NULL budget_reset_at with a non-NULL budget_duration 

4989 # as due, matching the budget-table query below. 

4990 "NOT": {"budget_duration": None}, 

4991 "OR": [ 

4992 {"budget_reset_at": None}, 

4993 {"budget_reset_at": {"lt": reset_at}}, 

4994 ], 

4995 }, 

4996 ) 

4997 elif query_type == "find_all" and user_id_list is not None: 

4998 response = await UserRepository(self).table.find_many(where={"user_id": {"in": user_id_list}}) 

4999 elif query_type == "find_all": 

5000 if expires is not None: 

5001 response = await UserRepository(self).table.find_many( 

5002 order={"spend": "desc"}, 

5003 where={ 

5004 "OR": [ 

5005 {"expires": None}, 

5006 {"expires": {"gt": expires}}, 

5007 ], 

5008 }, 

5009 ) 

5010 else: 

5011 # return all users in the table, get their key aliases ordered by spend 

5012 sql_query = """ 

5013 SELECT 

5014 u.*, 

5015 json_agg(v.key_alias) AS key_aliases 

5016 FROM 

5017 "LiteLLM_UserTable" u 

5018 LEFT JOIN "LiteLLM_VerificationToken" v ON u.user_id = v.user_id 

5019 GROUP BY 

5020 u.user_id 

5021 ORDER BY u.spend DESC 

5022 LIMIT $1 

5023 OFFSET $2 

5024 """ 

5025 response = await self.db.query_raw(sql_query, limit, offset) 

5026 return response 

5027 elif table_name == "spend": 5027 ↛ 5028line 5027 didn't jump to line 5028 because the condition on line 5027 was never true

5028 verbose_proxy_logger.debug("PrismaClient: get_data: table_name == 'spend'") 

5029 if key_val is not None: 

5030 if query_type == "find_unique": 

5031 response = await SpendLogsRepository(self).table.find_unique( 

5032 where={ 

5033 key_val["key"]: key_val["value"], 

5034 } 

5035 ) 

5036 elif query_type == "find_all": 

5037 response = await SpendLogsRepository(self).table.find_many( 

5038 where={ 

5039 key_val["key"]: key_val["value"], 

5040 } 

5041 ) 

5042 return response 

5043 else: 

5044 response = await SpendLogsRepository(self).table.find_many( 

5045 order={"startTime": "desc"}, 

5046 ) 

5047 return response 

5048 elif table_name == "budget" and reset_at is not None: 

5049 if query_type == "find_all": 5049 ↛ exitline 5049 didn't return from function 'get_data' because the condition on line 5049 was always true

5050 response = await BudgetRepository(self).table.find_many( 

5051 take=limit, 

5052 where={ 

5053 "NOT": {"budget_duration": None}, 

5054 "OR": [ 

5055 {"budget_reset_at": None}, 

5056 {"budget_reset_at": {"lt": reset_at}}, 

5057 ], 

5058 }, 

5059 ) 

5060 return response 

5061 

5062 elif table_name == "enduser" and budget_id_list is not None: 5062 ↛ 5063line 5062 didn't jump to line 5063 because the condition on line 5062 was never true

5063 if query_type == "find_all": 

5064 response = await EndUserRepository(self).table.find_many( 

5065 where={"budget_id": {"in": budget_id_list}} 

5066 ) 

5067 return response 

5068 elif table_name == "team": 

5069 if query_type == "find_unique": 

5070 response = await TeamRepository(self).table.find_unique( 

5071 where={"team_id": team_id}, 

5072 include={"litellm_model_table": True}, 

5073 ) 

5074 elif query_type == "find_all" and reset_at is not None: 

5075 response = await TeamRepository(self).table.find_many( 

5076 take=limit, 

5077 where={ 

5078 # Same NULL budget_reset_at gap as the user query 

5079 # above: a team with a budget_duration but no 

5080 # initialized budget_reset_at would never be reset. 

5081 "NOT": {"budget_duration": None}, 

5082 "OR": [ 

5083 {"budget_reset_at": None}, 

5084 {"budget_reset_at": {"lt": reset_at}}, 

5085 ], 

5086 }, 

5087 ) 

5088 elif query_type == "find_all" and user_id is not None: 5088 ↛ 5089line 5088 didn't jump to line 5089 because the condition on line 5088 was never true

5089 response = await TeamRepository(self).table.find_many( 

5090 where={ 

5091 "members": {"has": user_id}, 

5092 }, 

5093 include={"litellm_budget_table": True}, 

5094 ) 

5095 elif query_type == "find_all" and team_id_list is not None: 5095 ↛ 5097line 5095 didn't jump to line 5097 because the condition on line 5095 was always true

5096 response = await TeamRepository(self).table.find_many(where={"team_id": {"in": team_id_list}}) 

5097 elif query_type == "find_all" and team_id_list is None: 

5098 response = await TeamRepository(self).table.find_many(take=MAX_TEAM_LIST_LIMIT) 

5099 return response 

5100 elif table_name == "user_notification": 5100 ↛ 5101line 5100 didn't jump to line 5101 because the condition on line 5100 was never true

5101 if query_type == "find_unique": 

5102 response = await UserNotificationsRepository(self).table.find_unique(where={"user_id": user_id}) 

5103 elif query_type == "find_all": 

5104 response = await UserNotificationsRepository(self).table.find_many() 

5105 return response 

5106 elif table_name == "combined_view": 5106 ↛ exitline 5106 didn't return from function 'get_data' because the condition on line 5106 was always true

5107 # check if plain text or hash 

5108 if token is not None: 5108 ↛ 5112line 5108 didn't jump to line 5112 because the condition on line 5108 was always true

5109 if isinstance(token, str): 5109 ↛ 5112line 5109 didn't jump to line 5112 because the condition on line 5109 was always true

5110 hashed_token = _hash_token_if_needed(token=token) 

5111 verbose_proxy_logger.debug("PrismaClient: find_unique for token: %s", hashed_token) 

5112 if query_type == "find_unique": 5112 ↛ exitline 5112 didn't return from function 'get_data' because the condition on line 5112 was always true

5113 if token is None: 5113 ↛ 5114line 5113 didn't jump to line 5114 because the condition on line 5113 was never true

5114 raise HTTPException( 

5115 status_code=400, 

5116 detail={"error": f"No token passed in. Token={token}"}, 

5117 ) 

5118 

5119 sql_query = """ 

5120 SELECT  

5121 v.*, 

5122 t.spend AS team_spend,  

5123 t.max_budget AS team_max_budget, 

5124 t.soft_budget AS team_soft_budget, 

5125 t.model_max_budget AS team_model_max_budget, 

5126 t.tpm_limit AS team_tpm_limit, 

5127 t.rpm_limit AS team_rpm_limit, 

5128 t.tpd_limit AS team_tpd_limit, 

5129 t.models AS team_models, 

5130 t.metadata AS team_metadata, 

5131 t.blocked AS team_blocked, 

5132 t.team_alias AS team_alias, 

5133 t.metadata AS team_metadata, 

5134 t.members_with_roles AS team_members_with_roles, 

5135 t.object_permission_id AS team_object_permission_id, 

5136 t.organization_id as org_id, 

5137 p.project_alias AS project_alias, 

5138 tm.spend AS team_member_spend, 

5139 b_tm.tpm_limit AS team_member_tpm_limit, 

5140 b_tm.rpm_limit AS team_member_rpm_limit, 

5141 m.aliases AS team_model_aliases, 

5142 -- Added comma to separate b.* columns 

5143 b.max_budget AS litellm_budget_table_max_budget, 

5144 b.tpm_limit AS litellm_budget_table_tpm_limit, 

5145 b.rpm_limit AS litellm_budget_table_rpm_limit, 

5146 b.tpd_limit AS litellm_budget_table_tpd_limit, 

5147 b.model_max_budget as litellm_budget_table_model_max_budget, 

5148 b.soft_budget as litellm_budget_table_soft_budget, 

5149 o.metadata as organization_metadata, 

5150 o.organization_alias as organization_alias, 

5151 b2.max_budget as organization_max_budget, 

5152 b2.tpm_limit as organization_tpm_limit, 

5153 b2.rpm_limit as organization_rpm_limit 

5154 FROM "LiteLLM_VerificationToken" AS v 

5155 LEFT JOIN "LiteLLM_TeamTable" AS t ON v.team_id = t.team_id 

5156 LEFT JOIN "LiteLLM_TeamMembership" AS tm ON v.team_id = tm.team_id AND tm.user_id = v.user_id 

5157 LEFT JOIN "LiteLLM_BudgetTable" AS b_tm ON tm.budget_id = b_tm.budget_id 

5158 LEFT JOIN "LiteLLM_ModelTable" m ON t.model_id = m.id 

5159 LEFT JOIN "LiteLLM_BudgetTable" AS b ON v.budget_id = b.budget_id 

5160 LEFT JOIN "LiteLLM_ProjectTable" AS p ON v.project_id = p.project_id 

5161 LEFT JOIN "LiteLLM_OrganizationTable" AS o ON v.organization_id = o.organization_id 

5162 LEFT JOIN "LiteLLM_BudgetTable" AS b2 ON o.budget_id = b2.budget_id 

5163 WHERE v.token = $1 

5164 """ 

5165 

5166 response = await self._query_first_with_cached_plan_fallback(sql_query, hashed_token) 

5167 

5168 # If not found in main table, check deprecated keys (grace period) 

5169 # check_deprecated=False on the recursive call prevents unbounded chaining 

5170 if response is None and hashed_token is not None and check_deprecated: 

5171 active_token_id: Final = await _lookup_deprecated_key(db=self.db, hashed_token=hashed_token) 

5172 if active_token_id: 5172 ↛ 5176line 5172 didn't jump to line 5176 because the condition on line 5172 was never true

5173 # The recursive call returns a finished 

5174 # LiteLLM_VerificationTokenView; the dict 

5175 # normalization below would crash subscripting it. 

5176 deprecated_response: Final = await self.get_data( 

5177 token=active_token_id, 

5178 table_name="combined_view", 

5179 query_type="find_unique", 

5180 parent_otel_span=parent_otel_span, 

5181 proxy_logging_obj=proxy_logging_obj, 

5182 check_deprecated=False, 

5183 ) 

5184 if deprecated_response is not None: 

5185 verbose_proxy_logger.debug("Deprecated key used during grace period") 

5186 return deprecated_response 

5187 

5188 if response is not None: 

5189 if response["team_models"] is None: 5189 ↛ 5191line 5189 didn't jump to line 5191 because the condition on line 5189 was always true

5190 response["team_models"] = [] 

5191 if response["team_blocked"] is None: 5191 ↛ 5194line 5191 didn't jump to line 5194 because the condition on line 5191 was always true

5192 response["team_blocked"] = False 

5193 

5194 team_member: Member | None = None 

5195 if response["team_members_with_roles"] is not None and response["user_id"] is not None: 5195 ↛ 5197line 5195 didn't jump to line 5197 because the condition on line 5195 was never true

5196 ## find the team member corresponding to user id 

5197 """ 

5198 [ 

5199 { 

5200 "role": "admin", 

5201 "user_id": "default_user_id", 

5202 "user_email": null 

5203 }, 

5204 { 

5205 "role": "user", 

5206 "user_id": null, 

5207 "user_email": "test@email.com" 

5208 } 

5209 ] 

5210 """ 

5211 for tm in response["team_members_with_roles"]: 

5212 if tm.get("user_id") is not None and response["user_id"] == tm.get("user_id"): 

5213 team_member = Member(**tm) 

5214 response["team_member"] = team_member 

5215 response = LiteLLM_VerificationTokenView(**response, last_refreshed_at=time.time()) 

5216 # for prisma we need to cast the expires time to str 

5217 if response.expires is not None and isinstance(response.expires, datetime): 5217 ↛ 5218line 5217 didn't jump to line 5218 because the condition on line 5217 was never true

5218 response.expires = response.expires.isoformat() 

5219 return response 

5220 except Exception as e: 

5221 import traceback 

5222 

5223 prisma_query_info: Final = ( 

5224 f"LiteLLM Prisma Client Exception: Error with `get_data`. Args passed in: {args_passed_in}" 

5225 ) 

5226 error_msg: Final = prisma_query_info + str(e) 

5227 print_verbose(error_msg) 

5228 error_traceback: Final = error_msg + "\n" + traceback.format_exc() 

5229 verbose_proxy_logger.debug(error_traceback) 

5230 end_time: Final = time.time() 

5231 _duration: Final = end_time - start_time 

5232 

5233 asyncio.create_task( 

5234 self.proxy_logging_obj.failure_handler( 

5235 original_exception=e, 

5236 duration=_duration, 

5237 call_type="get_data", 

5238 traceback_str=error_traceback, 

5239 ) 

5240 ) 

5241 raise e 

5242 

5243 def jsonify_team_object(self, db_data: Mapping[str, object]) -> dict[str, object]: 

5244 db_data = self.jsonify_object(data=db_data) 

5245 if db_data.get("members_with_roles", None) is not None and isinstance(db_data["members_with_roles"], list): 

5246 db_data["members_with_roles"] = json.dumps(db_data["members_with_roles"]) 

5247 if db_data.get("budget_limits", None) is not None and isinstance(db_data["budget_limits"], list): 

5248 db_data["budget_limits"] = json.dumps(db_data["budget_limits"]) 

5249 return db_data 

5250 

5251 # Define a retrying strategy with exponential backoff 

5252 @backoff.on_exception( 

5253 backoff.expo, 

5254 Exception, # base exception to catch for the backoff 

5255 max_tries=3, # maximum number of retries 

5256 max_time=10, # maximum total time to retry for 

5257 on_backoff=on_backoff, # specifying the function to call on backoff 

5258 ) 

5259 async def insert_data( 

5260 self, 

5261 data: Mapping[str, object], 

5262 table_name: Literal["user", "key", "config", "spend", "team", "user_notification"], 

5263 ): 

5264 """ 

5265 Add a key to the database. If it already exists, do nothing. 

5266 """ 

5267 start_time: Final = time.time() 

5268 try: 

5269 verbose_proxy_logger.debug( 

5270 "PrismaClient: insert_data: %s", 

5271 {**data, "token": self.hash_token(token=cast("str", data["token"]))} # cast-ok: a key token is a str 

5272 if data.get("token") is not None 

5273 else data, 

5274 ) 

5275 if table_name == "key": 

5276 token: Final = cast("str", data["token"]) # cast-ok: the key table's token column is a str 

5277 hashed_token: Final = self.hash_token(token=token) 

5278 db_data = self.jsonify_object(data=data) 

5279 db_data["token"] = hashed_token 

5280 # Prisma rejects nullable JSON fields set to None (no default). 

5281 # Strip them so the DB stores NULL via the column's nullable constraint. 

5282 if db_data.get("budget_limits") is None: 

5283 db_data.pop("budget_limits", None) 

5284 print_verbose("PrismaClient: Before upsert into litellm_verificationtoken") 

5285 new_verification_token: Final = await VerificationTokenRepository(self).table.upsert( 

5286 where={ 

5287 "token": hashed_token, 

5288 }, 

5289 data={ 

5290 "create": {**db_data}, 

5291 "update": {}, # don't do anything if it already exists 

5292 }, 

5293 include={"litellm_budget_table": True}, 

5294 ) 

5295 verbose_proxy_logger.info("Data Inserted into Keys Table") 

5296 return new_verification_token 

5297 elif table_name == "user": 5297 ↛ 5321line 5297 didn't jump to line 5321 because the condition on line 5297 was always true

5298 db_data = self.jsonify_object(data=data) 

5299 try: 

5300 new_user_row: Final = await UserRepository(self).table.upsert( 

5301 where={"user_id": data["user_id"]}, 

5302 data={ 

5303 "create": {**db_data}, 

5304 "update": {}, # don't do anything if it already exists 

5305 }, 

5306 ) 

5307 except Exception as e: 

5308 if ( 5308 ↛ 5312line 5308 didn't jump to line 5312 because the condition on line 5308 was never true

5309 "Foreign key constraint failed on the field: `LiteLLM_UserTable_organization_id_fkey (index)`" 

5310 in str(e) 

5311 ): 

5312 raise HTTPException( 

5313 status_code=400, 

5314 detail={ 

5315 "error": f"Foreign Key Constraint failed. Organization ID={db_data['organization_id']} does not exist in LiteLLM_OrganizationTable. Create via `/organization/new`." 

5316 }, 

5317 ) 

5318 raise e 

5319 verbose_proxy_logger.info("Data Inserted into User Table") 

5320 return new_user_row 

5321 elif table_name == "team": 

5322 db_data = self.jsonify_team_object(db_data=data) 

5323 new_team_row: Final = await TeamRepository(self).table.upsert( 

5324 where={"team_id": data["team_id"]}, 

5325 data={ 

5326 "create": {**db_data}, 

5327 "update": {}, # don't do anything if it already exists 

5328 }, 

5329 ) 

5330 verbose_proxy_logger.info("Data Inserted into Team Table") 

5331 return new_team_row 

5332 elif table_name == "config": 

5333 """ 

5334 For each param, 

5335 get the existing table values 

5336 

5337 Add the new values 

5338 

5339 Update DB 

5340 """ 

5341 tasks: Final = [] 

5342 for k, v in data.items(): 

5343 updated_data = v 

5344 updated_data = json.dumps(updated_data) 

5345 updated_table_row = ConfigRepository(self).table.upsert( 

5346 where={"param_name": k}, 

5347 data={ 

5348 "create": {"param_name": k, "param_value": updated_data}, 

5349 "update": {"param_value": updated_data}, 

5350 }, 

5351 ) 

5352 

5353 tasks.append(updated_table_row) 

5354 await asyncio.gather(*tasks) 

5355 # invalidate cache so other pods see writes from save_config 

5356 for k in data: 

5357 await invalidate_config_param(k) 

5358 verbose_proxy_logger.info("Data Inserted into Config Table") 

5359 elif table_name == "spend": 

5360 db_data = self.jsonify_object(data=data) 

5361 new_spend_row: Final = await SpendLogsRepository(self).table.upsert( 

5362 where={"request_id": data["request_id"]}, 

5363 data={ 

5364 "create": {**db_data}, 

5365 "update": {}, # don't do anything if it already exists 

5366 }, 

5367 ) 

5368 verbose_proxy_logger.info("Data Inserted into Spend Table") 

5369 return new_spend_row 

5370 elif table_name == "user_notification": 

5371 db_data = self.jsonify_object(data=data) 

5372 new_user_notification_row: Final = await UserNotificationsRepository(self).table.upsert( 

5373 where={"request_id": data["request_id"]}, 

5374 data={ 

5375 "create": {**db_data}, 

5376 "update": {}, # don't do anything if it already exists 

5377 }, 

5378 ) 

5379 verbose_proxy_logger.info("Data Inserted into Model Request Table") 

5380 return new_user_notification_row 

5381 

5382 except Exception as e: 

5383 import traceback 

5384 

5385 error_msg: Final = f"LiteLLM Prisma Client Exception in insert_data: {e}" 

5386 print_verbose(error_msg) 

5387 error_traceback: Final = error_msg + "\n" + traceback.format_exc() 

5388 end_time: Final = time.time() 

5389 _duration: Final = end_time - start_time 

5390 asyncio.create_task( 

5391 self.proxy_logging_obj.failure_handler( 

5392 original_exception=e, 

5393 duration=_duration, 

5394 call_type="insert_data", 

5395 traceback_str=error_traceback, 

5396 ) 

5397 ) 

5398 raise e 

5399 

5400 # Define a retrying strategy with exponential backoff 

5401 @backoff.on_exception( 

5402 backoff.expo, 

5403 Exception, # base exception to catch for the backoff 

5404 max_tries=3, # maximum number of retries 

5405 max_time=10, # maximum total time to retry for 

5406 on_backoff=on_backoff, # specifying the function to call on backoff 

5407 ) 

5408 async def update_data( 

5409 self, 

5410 token: str | None = None, 

5411 data: Mapping[str, object] = {}, 

5412 data_list: list | None = None, 

5413 user_id: str | None = None, 

5414 team_id: str | None = None, 

5415 query_type: Literal["update", "update_many"] = "update", 

5416 table_name: Literal["user", "key", "config", "spend", "team", "enduser", "budget"] | None = None, 

5417 update_key_values: dict[str, object] | None = None, 

5418 update_key_values_custom_query: dict[str, object] | None = None, 

5419 ): 

5420 """ 

5421 Update existing data 

5422 """ 

5423 verbose_proxy_logger.debug("PrismaClient: update_data, table_name: %s", table_name) 

5424 start_time: Final = time.time() 

5425 try: 

5426 db_data: Final = self.jsonify_object(data=data) 

5427 if update_key_values is not None: 

5428 update_key_values = self.jsonify_object(data=update_key_values) 

5429 if token is not None: 

5430 print_verbose(f"token: [set={token is not None}]") 

5431 # check if plain text or hash 

5432 token = _hash_token_if_needed(token=token) 

5433 db_data["token"] = token 

5434 response: Final = await VerificationTokenRepository(self).table.update( 

5435 where={"token": token}, 

5436 data=with_settings_updated_at(db_data), 

5437 ) 

5438 verbose_proxy_logger.debug("\033[91m" + f"DB Token Table update succeeded {response}" + "\033[0m") 

5439 _data: dict = {} 

5440 if response is not None: 

5441 try: 

5442 _data = response.model_dump() 

5443 except Exception: 

5444 _data = response.dict() # pyright: ignore[reportDeprecated] # pydantic-v1 row fallback 

5445 return {"token": token, "data": _data} 

5446 elif user_id is not None or (table_name is not None and table_name == "user") and query_type == "update": 

5447 """ 

5448 If data['spend'] + data['user'], update the user table with spend info as well 

5449 """ 

5450 if user_id is None: 

5451 user_id = cast("str", db_data["user_id"]) # cast-ok: the user table's user_id column is a str 

5452 if update_key_values is None: 

5453 if update_key_values_custom_query is not None: 

5454 update_key_values = update_key_values_custom_query 

5455 else: 

5456 update_key_values = db_data 

5457 update_user_row: Final = await UserRepository(self).table.upsert( 

5458 where={"user_id": user_id}, 

5459 data={ 

5460 "create": {**db_data}, 

5461 "update": {**update_key_values}, # just update user-specified values, if it already exists 

5462 }, 

5463 ) 

5464 verbose_proxy_logger.info( 

5465 "\033[91m" + f"DB User Table - update succeeded {update_user_row}" + "\033[0m" 

5466 ) 

5467 return {"user_id": user_id, "data": update_user_row} 

5468 elif team_id is not None or (table_name is not None and table_name == "team") and query_type == "update": 

5469 """ 

5470 If data['spend'] + data['user'], update the user table with spend info as well 

5471 """ 

5472 if team_id is None: 

5473 team_id = cast("str | None", db_data["team_id"]) # cast-ok: team_id column is a nullable str 

5474 if update_key_values is None: 

5475 update_key_values = db_data 

5476 if "team_id" not in db_data and team_id is not None: 

5477 db_data["team_id"] = team_id 

5478 if "members_with_roles" in db_data and isinstance(db_data["members_with_roles"], list): 

5479 db_data["members_with_roles"] = json.dumps(db_data["members_with_roles"]) 

5480 if "members_with_roles" in update_key_values and isinstance( 

5481 update_key_values["members_with_roles"], list 

5482 ): 

5483 update_key_values["members_with_roles"] = json.dumps(update_key_values["members_with_roles"]) 

5484 update_team_row: Final = await TeamRepository(self).table.upsert( 

5485 where={"team_id": team_id}, 

5486 data={ 

5487 "create": {**db_data}, 

5488 "update": {**update_key_values}, # just update user-specified values, if it already exists 

5489 }, 

5490 ) 

5491 verbose_proxy_logger.info( 

5492 "\033[91m" + f"DB Team Table - update succeeded {update_team_row}" + "\033[0m" 

5493 ) 

5494 return {"team_id": team_id, "data": update_team_row} 

5495 elif ( 

5496 table_name is not None 

5497 and table_name == "key" 

5498 and query_type == "update_many" 

5499 and data_list is not None 

5500 and isinstance(data_list, list) 

5501 ): 

5502 """ 

5503 Batch write update queries 

5504 """ 

5505 batcher = self.db.batch_() 

5506 for idx, t in enumerate(data_list): 

5507 # check if plain text or hash 

5508 if t.token.startswith("sk-"): 

5509 t.token = self.hash_token(token=t.token) 

5510 try: 

5511 data_json = self.jsonify_object(data=t.model_dump(exclude_none=True)) 

5512 except Exception: 

5513 data_json = self.jsonify_object(data=t.dict(exclude_none=True)) 

5514 batcher.litellm_verificationtoken.update( 

5515 where={"token": t.token}, 

5516 data={**data_json}, 

5517 ) 

5518 await batcher.commit() 

5519 print_verbose("\033[91m" + "DB Token Table update succeeded" + "\033[0m") 

5520 elif ( 

5521 table_name is not None 

5522 and table_name == "user" 

5523 and query_type == "update_many" 

5524 and data_list is not None 

5525 and isinstance(data_list, list) 

5526 ): 

5527 """ 

5528 Batch write update queries 

5529 """ 

5530 batcher = self.db.batch_() 

5531 for idx, user in enumerate(data_list): 

5532 try: 

5533 data_json = self.jsonify_object(data=user.model_dump(exclude_none=True)) 

5534 except Exception: 

5535 data_json = self.jsonify_object(data=user.dict()) 

5536 batcher.litellm_usertable.upsert( 

5537 where={"user_id": user.user_id}, 

5538 data={ 

5539 "create": {**data_json}, 

5540 "update": {**data_json}, # just update user-specified values, if it already exists 

5541 }, 

5542 ) 

5543 await batcher.commit() 

5544 verbose_proxy_logger.info("\033[91m" + "DB User Table Batch update succeeded" + "\033[0m") 

5545 elif ( 

5546 table_name is not None 

5547 and table_name == "enduser" 

5548 and query_type == "update_many" 

5549 and data_list is not None 

5550 and isinstance(data_list, list) 

5551 ): 

5552 """ 

5553 Batch write update queries 

5554 """ 

5555 batcher = self.db.batch_() 

5556 for enduser in data_list: 

5557 try: 

5558 data_json = self.jsonify_object(data=enduser.model_dump(exclude_none=True)) 

5559 except Exception: 

5560 data_json = self.jsonify_object(data=enduser.dict()) 

5561 batcher.litellm_endusertable.upsert( 

5562 where={"user_id": enduser.user_id}, 

5563 data={ 

5564 "create": {**data_json}, 

5565 "update": {**data_json}, # just update end-user-specified values, if it already exists 

5566 }, 

5567 ) 

5568 await batcher.commit() 

5569 verbose_proxy_logger.info("\033[91m" + "DB End User Table Batch update succeeded" + "\033[0m") 

5570 elif ( 

5571 table_name is not None 

5572 and table_name == "budget" 

5573 and query_type == "update_many" 

5574 and data_list is not None 

5575 and isinstance(data_list, list) 

5576 ): 

5577 """ 

5578 Batch write update queries 

5579 """ 

5580 batcher = self.db.batch_() 

5581 for budget in data_list: 

5582 try: 

5583 data_json = self.jsonify_object(data=budget.model_dump(exclude_none=True)) 

5584 except Exception: 

5585 data_json = self.jsonify_object(data=budget.dict()) 

5586 batcher.litellm_budgettable.upsert( 

5587 where={"budget_id": budget.budget_id}, 

5588 data={ 

5589 "create": {**data_json}, 

5590 "update": {**data_json}, # just update end-user-specified values, if it already exists 

5591 }, 

5592 ) 

5593 await batcher.commit() 

5594 verbose_proxy_logger.info("\033[91m" + "DB Budget Table Batch update succeeded" + "\033[0m") 

5595 elif ( 

5596 table_name is not None 

5597 and table_name == "team" 

5598 and query_type == "update_many" 

5599 and data_list is not None 

5600 and isinstance(data_list, list) 

5601 ): 

5602 # Batch write update queries 

5603 batcher = self.db.batch_() 

5604 for idx, team in enumerate(data_list): 

5605 try: 

5606 data_json = self.jsonify_team_object(db_data=team.model_dump(exclude_none=True)) 

5607 except Exception: 

5608 data_json = self.jsonify_object(data=team.dict(exclude_none=True)) 

5609 batcher.litellm_teamtable.upsert( 

5610 where={"team_id": team.team_id}, 

5611 data={ 

5612 "create": {**data_json}, 

5613 "update": {**data_json}, # just update user-specified values, if it already exists 

5614 }, 

5615 ) 

5616 await batcher.commit() 

5617 verbose_proxy_logger.info("\033[91m" + "DB Team Table Batch update succeeded" + "\033[0m") 

5618 

5619 except Exception as e: 

5620 import traceback 

5621 

5622 error_msg: Final = f"LiteLLM Prisma Client Exception - update_data: {e}" 

5623 print_verbose(error_msg) 

5624 error_traceback: Final = error_msg + "\n" + traceback.format_exc() 

5625 end_time: Final = time.time() 

5626 _duration: Final = end_time - start_time 

5627 asyncio.create_task( 

5628 self.proxy_logging_obj.failure_handler( 

5629 original_exception=e, 

5630 duration=_duration, 

5631 call_type="update_data", 

5632 traceback_str=error_traceback, 

5633 ) 

5634 ) 

5635 raise e 

5636 

5637 # Define a retrying strategy with exponential backoff 

5638 @backoff.on_exception( 

5639 backoff.expo, 

5640 Exception, # base exception to catch for the backoff 

5641 max_tries=3, # maximum number of retries 

5642 max_time=10, # maximum total time to retry for 

5643 on_backoff=on_backoff, # specifying the function to call on backoff 

5644 ) 

5645 async def delete_data( 

5646 self, 

5647 tokens: Sequence[str | None] | None = None, 

5648 team_id_list: Sequence[str] | None = None, 

5649 table_name: Literal["user", "key", "config", "spend", "team"] | None = None, 

5650 user_id: str | None = None, 

5651 ): 

5652 """ 

5653 Allow user to delete a key(s) 

5654 

5655 Ensure user owns that key, unless admin. 

5656 """ 

5657 start_time: Final = time.time() 

5658 try: 

5659 if tokens is not None and isinstance(tokens, list): 5659 ↛ 5660line 5659 didn't jump to line 5660 because the condition on line 5659 was never true

5660 hashed_tokens: Final[list[str | None]] = [] 

5661 for token in tokens: 

5662 if isinstance(token, str) and token.startswith("sk-"): 

5663 hashed_token = self.hash_token(token=token) 

5664 else: 

5665 hashed_token = token 

5666 hashed_tokens.append(hashed_token) 

5667 filter_query: dict[str, object] = {} 

5668 if user_id is not None: 

5669 filter_query = {"AND": [{"token": {"in": hashed_tokens}}, {"user_id": user_id}]} 

5670 else: 

5671 filter_query = {"token": {"in": hashed_tokens}} 

5672 

5673 deleted_tokens: Final[int] = await VerificationTokenRepository(self).table.delete_many( 

5674 where=filter_query 

5675 ) 

5676 verbose_proxy_logger.debug("deleted_tokens: %s", deleted_tokens) 

5677 return {"deleted_keys": deleted_tokens} 

5678 elif table_name == "team" and team_id_list is not None and isinstance(team_id_list, list): 5678 ↛ 5680line 5678 didn't jump to line 5680 because the condition on line 5678 was never true

5679 # admin only endpoint -> `/team/delete` 

5680 await TeamRepository(self).table.delete_many(where={"team_id": {"in": team_id_list}}) 

5681 return {"deleted_teams": team_id_list} 

5682 elif table_name == "key" and team_id_list is not None and isinstance(team_id_list, list): 5682 ↛ exitline 5682 didn't return from function 'delete_data' because the condition on line 5682 was always true

5683 # admin only endpoint -> `/team/delete` 

5684 await VerificationTokenRepository(self).table.delete_many(where={"team_id": {"in": team_id_list}}) 

5685 except Exception as e: 

5686 import traceback 

5687 

5688 error_msg: Final = f"LiteLLM Prisma Client Exception - delete_data: {e}" 

5689 print_verbose(error_msg) 

5690 error_traceback: Final = error_msg + "\n" + traceback.format_exc() 

5691 end_time: Final = time.time() 

5692 _duration: Final = end_time - start_time 

5693 asyncio.create_task( 

5694 self.proxy_logging_obj.failure_handler( 

5695 original_exception=e, 

5696 duration=_duration, 

5697 call_type="delete_data", 

5698 traceback_str=error_traceback, 

5699 ) 

5700 ) 

5701 raise e 

5702 

5703 # Define a retrying strategy with exponential backoff 

5704 @backoff.on_exception( 

5705 backoff.expo, 

5706 Exception, # base exception to catch for the backoff 

5707 max_tries=3, # maximum number of retries 

5708 max_time=10, # maximum total time to retry for 

5709 on_backoff=on_backoff, # specifying the function to call on backoff 

5710 ) 

5711 async def connect(self): 

5712 start_time: Final = time.time() 

5713 try: 

5714 verbose_proxy_logger.debug("PrismaClient: connect() called Attempting to Connect to DB") 

5715 if self.db.is_connected() is False: 5715 ↛ exitline 5715 didn't return from function 'connect' because the condition on line 5715 was always true

5716 verbose_proxy_logger.debug("PrismaClient: DB not connected, Attempting to Connect to DB") 

5717 await self.db.connect() 

5718 except Exception as e: 

5719 import traceback 

5720 

5721 error_msg: Final = f"LiteLLM Prisma Client Exception connect(): {e}" 

5722 verbose_proxy_logger.warning(error_msg) 

5723 error_traceback: Final = error_msg + "\n" + traceback.format_exc() 

5724 end_time: Final = time.time() 

5725 _duration: Final = end_time - start_time 

5726 asyncio.create_task( 

5727 self.proxy_logging_obj.failure_handler( 

5728 original_exception=e, 

5729 duration=_duration, 

5730 call_type="connect", 

5731 traceback_str=error_traceback, 

5732 ) 

5733 ) 

5734 raise e 

5735 

5736 # Define a retrying strategy with exponential backoff 

5737 @backoff.on_exception( 

5738 backoff.expo, 

5739 Exception, # base exception to catch for the backoff 

5740 max_tries=3, # maximum number of retries 

5741 max_time=10, # maximum total time to retry for 

5742 on_backoff=on_backoff, # specifying the function to call on backoff 

5743 ) 

5744 async def disconnect(self): 

5745 start_time: Final = time.time() 

5746 try: 

5747 await self.db.disconnect() 

5748 except Exception as e: 

5749 import traceback 

5750 

5751 error_msg: Final = f"LiteLLM Prisma Client Exception disconnect(): {e}" 

5752 print_verbose(error_msg) 

5753 error_traceback: Final = error_msg + "\n" + traceback.format_exc() 

5754 end_time: Final = time.time() 

5755 _duration: Final = end_time - start_time 

5756 asyncio.create_task( 

5757 self.proxy_logging_obj.failure_handler( 

5758 original_exception=e, 

5759 duration=_duration, 

5760 call_type="disconnect", 

5761 traceback_str=error_traceback, 

5762 ) 

5763 ) 

5764 raise e 

5765 

5766 def _get_engine_pid(self) -> int: 

5767 """Get the PID of the writer's engine subprocess, or 0 if unavailable. 

5768 

5769 Must never raise: prisma's ``_engine`` property raises 

5770 ``ClientNotConnectedError`` on a disconnected client, and an exception 

5771 escaping from the reconnect path would leave it unable to recover. 

5772 """ 

5773 try: 

5774 prisma_obj: Final = self.writer_db._original_prisma 

5775 if prisma_obj.is_connected() is not True: 5775 ↛ 5776line 5775 didn't jump to line 5776 because the condition on line 5775 was never true

5776 return 0 

5777 engine: Final = prisma_obj._engine 

5778 process: Final = getattr(engine, "process", None) if engine is not None else None 

5779 if process is not None: 5779 ↛ 5785line 5779 didn't jump to line 5785 because the condition on line 5779 was always true

5780 pid: Final[object] = process.pid 

5781 if isinstance(pid, int): 5781 ↛ 5785line 5781 didn't jump to line 5785 because the condition on line 5781 was always true

5782 return pid 

5783 except (AttributeError, TypeError): 

5784 pass 

5785 return 0 

5786 

5787 def _is_engine_alive(self) -> bool: 

5788 if self._engine_pid <= 0: 

5789 return True 

5790 try: 

5791 os.kill(self._engine_pid, 0) 

5792 return True 

5793 except ProcessLookupError: 

5794 return False 

5795 except (PermissionError, OSError): 

5796 return True 

5797 

5798 @staticmethod 

5799 def _reap_all_zombies() -> set: 

5800 """Reap ALL zombie child processes via waitpid(-1, WNOHANG). 

5801 

5802 Returns a set of reaped PIDs. As PID 1 in Docker (or any 

5803 process that spawns children), we must reap ALL terminated 

5804 children to prevent zombie accumulation. 

5805 

5806 No-op on Windows: os.waitpid and os.WNOHANG are Unix-only. 

5807 """ 

5808 if sys.platform == "win32": 

5809 return set() 

5810 reaped: Final[set] = set() 

5811 while True: 

5812 try: 

5813 pid, _ = os.waitpid(-1, os.WNOHANG) 

5814 if pid == 0: 

5815 break 

5816 reaped.add(pid) 

5817 except ChildProcessError: 

5818 break 

5819 return reaped 

5820 

5821 def _try_waitpid_watch(self, pid: int) -> bool: 

5822 """Watch engine PID via os.waitpid() in a dedicated thread. 

5823 

5824 The thread blocks on os.waitpid(pid, 0) which is a kernel-level 

5825 wait and with zero CPU overhead, instant detection when the process exits. 

5826 When the process dies, the thread notifies the asyncio event loop 

5827 via call_soon_threadsafe. 

5828 

5829 Returns True if the thread was started, False on failure. 

5830 On Windows, returns False immediately (os.waitpid/WNOHANG are Unix-only); 

5831 caller falls back to os.kill polling. 

5832 """ 

5833 if sys.platform == "win32": 5833 ↛ 5834line 5833 didn't jump to line 5834 because the condition on line 5833 was never true

5834 return False 

5835 try: 

5836 probe_pid, _ = os.waitpid(pid, os.WNOHANG) 

5837 except ChildProcessError: 

5838 verbose_proxy_logger.debug( 

5839 "PID %s is not a child process; skipping waitpid watch.", 

5840 pid, 

5841 ) 

5842 return False 

5843 

5844 if probe_pid == pid: 5844 ↛ 5845line 5844 didn't jump to line 5845 because the condition on line 5844 was never true

5845 verbose_proxy_logger.warning( 

5846 "prisma-query-engine PID %s already dead at watch start.", 

5847 pid, 

5848 ) 

5849 if self._consume_expected_death(pid): 

5850 verbose_proxy_logger.info( 

5851 "PID %s death was planned (engine already replaced); not reconnecting.", 

5852 pid, 

5853 ) 

5854 self._cleanup_engine_watcher() 

5855 return True 

5856 self._engine_confirmed_dead = True 

5857 self._reap_all_zombies() 

5858 self._cleanup_engine_watcher() 

5859 asyncio.create_task( 

5860 self.attempt_db_reconnect( 

5861 reason="engine_process_death", 

5862 force=True, 

5863 ) 

5864 ) 

5865 return True 

5866 

5867 try: 

5868 loop: Final = asyncio.get_running_loop() 

5869 except RuntimeError: 

5870 return False 

5871 

5872 thread: Final = threading.Thread( 

5873 target=self._waitpid_thread_func, 

5874 args=(pid, loop), 

5875 daemon=True, 

5876 name=f"prisma-engine-waitpid-{pid}", 

5877 ) 

5878 thread.start() 

5879 self._engine_wait_thread = thread 

5880 return True 

5881 

5882 def _waitpid_thread_func(self, pid: int, loop: asyncio.AbstractEventLoop) -> None: 

5883 """Thread function: block until engine PID exits, then notify event loop. 

5884 

5885 Note: uvloop/libuv may reap the child first via waitpid(-1, WNOHANG) 

5886 in its SIGCHLD handler. In that case our waitpid raises ChildProcessError. 

5887 we still notify the event loop because the engine is dead either way. 

5888 """ 

5889 try: 

5890 os.waitpid(pid, 0) 

5891 except ChildProcessError: 

5892 pass 

5893 except OSError: 

5894 pass 

5895 try: 

5896 loop.call_soon_threadsafe(self._on_engine_death_from_thread, pid) 

5897 except RuntimeError: 

5898 pass 

5899 

5900 def _consume_expected_death(self, pid: int) -> bool: 

5901 """True iff ``pid`` was killed on purpose by a planned recreate. 

5902 

5903 `PrismaWrapper.recreate_prisma_client` records the old engine PID in 

5904 `_expected_engine_deaths` before SIGTERM-ing it (IAM token refresh, 

5905 guarded reconnect). When the watcher then sees that PID die, this lets 

5906 it recognize the death as planned and skip its own reconnect, which 

5907 would otherwise kill the engine the recreate just spawned (#29176). 

5908 

5909 Consumes (removes) the PID so a later real crash of a reused PID is 

5910 still handled. Tolerant of `self.db` stand-ins (tests / older clients) 

5911 that don't expose a real set. 

5912 """ 

5913 expected: Final = getattr(self.db, "_expected_engine_deaths", None) 

5914 if isinstance(expected, set) and pid in expected: 

5915 expected.discard(pid) 

5916 return True 

5917 return False 

5918 

5919 def _on_engine_death_from_thread(self, dead_pid: int) -> None: 

5920 """Called on the event loop thread when the waitpid thread detects engine death.""" 

5921 if self._engine_confirmed_dead: 5921 ↛ 5922line 5921 didn't jump to line 5922 because the condition on line 5921 was never true

5922 return 

5923 if dead_pid != self._engine_pid: 5923 ↛ 5925line 5923 didn't jump to line 5925 because the condition on line 5923 was always true

5924 return 

5925 if self._consume_expected_death(dead_pid): 

5926 verbose_proxy_logger.info( 

5927 "prisma-query-engine PID %s exited as part of a planned restart; " 

5928 "not reconnecting (engine already replaced).", 

5929 dead_pid, 

5930 ) 

5931 self._cleanup_engine_watcher() 

5932 return 

5933 verbose_proxy_logger.error( 

5934 "prisma-query-engine PID %s exited (waitpid thread); triggering reconnect.", 

5935 dead_pid, 

5936 ) 

5937 self._engine_confirmed_dead = True 

5938 self._reap_all_zombies() 

5939 self._cleanup_engine_watcher() 

5940 asyncio.create_task( 

5941 self.attempt_db_reconnect( 

5942 reason="engine_process_death", 

5943 force=True, 

5944 ) 

5945 ) 

5946 

5947 def _try_pidfd_watch(self, pid: int) -> bool: 

5948 """ 

5949 Watch engine PID via pidfd_open + asyncio event loop reader. 

5950 

5951 Returns True if pidfd watch was set up, False if unavailable or failed. 

5952 Broad OSError catch handles both ENOSYS and SECCOMP-blocked syscalls. 

5953 """ 

5954 if not hasattr(os, "pidfd_open"): 

5955 return False 

5956 fd = -1 

5957 try: 

5958 fd = os.pidfd_open(pid, 0) 

5959 asyncio.get_running_loop().add_reader(fd, self._on_pidfd_readable) 

5960 self._engine_pidfd = fd 

5961 return True 

5962 except OSError: 

5963 if fd >= 0: 

5964 os.close(fd) 

5965 return False 

5966 

5967 def _on_pidfd_readable(self) -> None: 

5968 """pidfd became readable: engine process exited or became zombie. 

5969 

5970 Sets _engine_confirmed_dead BEFORE cleanup so _run_reconnect_cycle 

5971 takes the heavy path (recreate Prisma client + re-arm watcher). 

5972 """ 

5973 if self._engine_confirmed_dead: 

5974 # Already handled -- just clean up pidfd resources. 

5975 if self._engine_pidfd >= 0: 

5976 try: 

5977 asyncio.get_running_loop().remove_reader(self._engine_pidfd) 

5978 except Exception: 

5979 pass 

5980 try: 

5981 os.close(self._engine_pidfd) 

5982 except OSError: 

5983 pass 

5984 self._engine_pidfd = -1 

5985 return 

5986 dead_pid: Final = self._engine_pid 

5987 if self._consume_expected_death(dead_pid): 

5988 verbose_proxy_logger.info( 

5989 "prisma-query-engine PID %s exited (pidfd event) as part of a " 

5990 "planned restart; not reconnecting (engine already replaced).", 

5991 dead_pid, 

5992 ) 

5993 self._cleanup_engine_watcher() 

5994 return 

5995 verbose_proxy_logger.error( 

5996 "prisma-query-engine PID %s exited (pidfd event); triggering reconnect.", 

5997 dead_pid, 

5998 ) 

5999 self._engine_confirmed_dead = True 

6000 self._reap_all_zombies() 

6001 self._cleanup_engine_watcher() 

6002 asyncio.create_task( 

6003 self.attempt_db_reconnect( 

6004 reason="engine_process_death", 

6005 force=True, 

6006 ) 

6007 ) 

6008 

6009 async def _poll_engine_proc(self) -> None: 

6010 """poll via os.kill(pid, 0) every 1s. 

6011 Only used when BOTH waitpid thread and pidfd are unavailable 

6012 (e.g., PID is not our child process and pidfd_open fails) 

6013 """ 

6014 while self._watching_engine and self._engine_pid > 0: 

6015 try: 

6016 os.kill(self._engine_pid, 0) 

6017 except ProcessLookupError: 

6018 dead_pid = self._engine_pid 

6019 if self._consume_expected_death(dead_pid): 

6020 verbose_proxy_logger.info( 

6021 "prisma-query-engine PID %s gone as part of a planned " 

6022 "restart; not reconnecting (engine already replaced).", 

6023 dead_pid, 

6024 ) 

6025 self._cleanup_engine_watcher() 

6026 return 

6027 verbose_proxy_logger.error( 

6028 "prisma-query-engine PID %s gone; triggering reconnect.", 

6029 dead_pid, 

6030 ) 

6031 self._engine_confirmed_dead = True 

6032 self._reap_all_zombies() 

6033 self._cleanup_engine_watcher() 

6034 await self.attempt_db_reconnect( 

6035 reason="engine_process_death", 

6036 force=True, 

6037 ) 

6038 return 

6039 except (PermissionError, OSError): 

6040 verbose_proxy_logger.debug( 

6041 "Cannot signal PID %s; stopping engine poll.", 

6042 self._engine_pid, 

6043 ) 

6044 self._cleanup_engine_watcher() 

6045 return 

6046 await asyncio.sleep(1) 

6047 

6048 def _cleanup_engine_watcher(self) -> None: 

6049 """Clean up pidfd reader, waitpid thread ref, or stop polling and reset state.""" 

6050 self._watching_engine = False 

6051 if self._engine_pidfd >= 0: 6051 ↛ 6052line 6051 didn't jump to line 6052 because the condition on line 6051 was never true

6052 try: 

6053 asyncio.get_running_loop().remove_reader(self._engine_pidfd) 

6054 except Exception: 

6055 pass 

6056 try: 

6057 os.close(self._engine_pidfd) 

6058 except OSError: 

6059 pass 

6060 self._engine_pidfd = -1 

6061 self._engine_wait_thread = None 

6062 self._engine_pid = 0 

6063 

6064 async def _start_engine_watcher(self) -> None: 

6065 """ 

6066 Start watching the Prisma query engine process for death. 

6067 

6068 Detection priority: 

6069 1. os.waitpid() in a dedicated thread, works with all event loops. 

6070 2. pidfd_open kernel fd registered with asyncio. 

6071 3. os.kill(pid, 0) polling (1s), last-resort fallback when neither 

6072 waitpid thread nor pidfd are available. 

6073 

6074 """ 

6075 if self._watching_engine or self._engine_pidfd >= 0 or self._engine_wait_thread is not None: 6075 ↛ 6076line 6075 didn't jump to line 6076 because the condition on line 6075 was never true

6076 return 

6077 pid: Final = self._get_engine_pid() 

6078 if pid == 0: 6078 ↛ 6079line 6078 didn't jump to line 6079 because the condition on line 6078 was never true

6079 verbose_proxy_logger.debug("Could not find prisma-query-engine PID; engine death detection unavailable.") 

6080 return 

6081 self._engine_pid = pid 

6082 self._engine_confirmed_dead = False 

6083 verbose_proxy_logger.info("Found prisma-query-engine at PID %s.", pid) 

6084 waitpid_ok: Final = self._try_waitpid_watch(pid) 

6085 pidfd_ok: Final = False if waitpid_ok else self._try_pidfd_watch(pid) 

6086 if waitpid_ok: 6086 ↛ 6091line 6086 didn't jump to line 6091 because the condition on line 6086 was always true

6087 verbose_proxy_logger.info( 

6088 "Watching engine PID %s via waitpid thread.", 

6089 pid, 

6090 ) 

6091 elif pidfd_ok: 

6092 verbose_proxy_logger.info( 

6093 "Watching engine PID %s via pidfd.", 

6094 pid, 

6095 ) 

6096 else: 

6097 verbose_proxy_logger.info( 

6098 "Watching engine PID %s via os.kill polling.", 

6099 pid, 

6100 ) 

6101 self._watching_engine = True 

6102 asyncio.create_task(self._poll_engine_proc()) 

6103 

6104 def _stop_engine_watcher(self) -> None: 

6105 """Stop watching the engine process and clean up all resources.""" 

6106 self._cleanup_engine_watcher() 

6107 self._engine_confirmed_dead = False 

6108 verbose_proxy_logger.debug("Stopped engine process watcher.") 

6109 

6110 def _handle_writer_engine_replaced(self) -> None: 

6111 """Re-arm the engine watcher after a planned writer-engine restart. 

6112 

6113 Wired as `PrismaWrapper.on_engine_replaced` and invoked from inside 

6114 `recreate_prisma_client` once the new engine is connected (IAM token 

6115 refresh, guarded reconnect). The old watcher was tracking the engine 

6116 we just intentionally killed, so we tear it down and re-arm on the new 

6117 PID. Scheduling `_start_engine_watcher` as a task (rather than awaiting) 

6118 keeps us from blocking the recreate while it still holds the wrapper's 

6119 reconnection lock. Without this re-arm, a planned restart would leave 

6120 the proxy with no engine-death detection until the next reconnect. 

6121 """ 

6122 self._engine_confirmed_dead = False 

6123 self._cleanup_engine_watcher() 

6124 asyncio.create_task(self._start_engine_watcher()) 

6125 

6126 async def _run_reconnect_cycle( 

6127 self, 

6128 timeout_seconds: float | None = None, 

6129 force_recreate: bool = False, 

6130 ) -> None: 

6131 """ 

6132 Run a reconnect cycle with a single overall timeout budget. 

6133 

6134 Uses the _engine_confirmed_dead flag (set by waitpid thread / pidfd / poll 

6135 handlers) to choose between heavy reconnect (engine dead -- recreate 

6136 Prisma client, re-arm watcher) and direct reconnect (network blip -- 

6137 recreate Prisma client, re-arm watcher, SELECT 1). Both paths recreate 

6138 the client via the non-blocking kill-then-construct flow rather than 

6139 calling disconnect(), which blocks the event loop on the synchronous 

6140 subprocess.Popen.wait() inside prisma-client-py (see issue #26191). 

6141 

6142 `force_recreate` skips the direct path's liveness probe, for callers 

6143 whose failure lives in the session state rather than the connection 

6144 (stale prepared statements after a schema change): a reachable writer 

6145 proves nothing about those, so the probe must not veto the recreate. 

6146 """ 

6147 effective_timeout: Final = ( 

6148 timeout_seconds if timeout_seconds is not None else self._db_watchdog_reconnect_timeout_seconds 

6149 ) 

6150 

6151 # Snapshot the writer's engine generation BEFORE any await. Both 

6152 # reconnect branches forward it to recreate_prisma_client as an 

6153 # optimistic-lock token: if a concurrent IAM token refresh replaces the 

6154 # engine after this point, the generation moves and the recreate becomes 

6155 # a no-op instead of killing the engine the refresh just spawned 

6156 # (#29176). Captured here — atomically with the dead-engine decision 

6157 # below — rather than inside the reconnect closures, because those run 

6158 # after an `asyncio.wait_for(...)` yield during which a refresh could 

6159 # otherwise slip in and bump the very generation the closure then reads. 

6160 expected_generation: Final = getattr(self.writer_db, "_engine_generation", None) 

6161 

6162 engine_is_dead: Final = self._engine_confirmed_dead or (self._engine_pid > 0 and not self._is_engine_alive()) 

6163 

6164 if engine_is_dead: 

6165 dead_pid: Final = self._engine_pid 

6166 verbose_proxy_logger.warning( 

6167 "prisma-query-engine PID %s is dead; reconnecting.", 

6168 dead_pid, 

6169 ) 

6170 self._reap_all_zombies() 

6171 self._cleanup_engine_watcher() 

6172 

6173 async def _do_heavy_reconnect() -> None: 

6174 db_url: Final = os.getenv("DATABASE_URL", "") 

6175 if not db_url: 

6176 verbose_proxy_logger.error("DATABASE_URL not set; cannot recreate Prisma client.") 

6177 raise RuntimeError("DATABASE_URL not set") 

6178 # Forward the entry-snapshot generation. The engine was 

6179 # confirmed dead, but a concurrent IAM refresh may have already 

6180 # respawned it; the guard makes this recreate a no-op in that 

6181 # case rather than killing the fresh engine (#29176). Unlike the 

6182 # direct path there is no SELECT 1 probe here, so the generation 

6183 # guard is the only thing standing between a crash-reconnect and 

6184 # a refresh that raced it. 

6185 recreated: Final = await self.db.recreate_prisma_client(db_url, expected_generation=expected_generation) 

6186 await self._start_engine_watcher() 

6187 # Same contract as the direct path below: a forced caller asked 

6188 # for its engine to be replaced, so a decline is not a success. 

6189 # Reachable here because the escalation threshold flips 

6190 # `_engine_confirmed_dead`, which routes the next cycle, forced 

6191 # callers included, down this branch. 

6192 if force_recreate is True and recreated is False: 

6193 # Clear the dead-engine flag first, restoring the policy the 

6194 # non-forced path already has: a decline does not raise for 

6195 # it, so it falls through to the clear below. Only the 

6196 # forced branch would strand the flag, and stranding it 

6197 # routes the next cycle back down this probe-free branch, 

6198 # where the refreshed generation now matches and the 

6199 # recreate kills the healthy engine a refresh just spawned 

6200 # (#29176). This has to stay AFTER `_start_engine_watcher` 

6201 # above: clearing the flag while the watcher is still torn 

6202 # down would be worse than either alone. 

6203 self._engine_confirmed_dead = False 

6204 raise _ForcedRecreateDeclined( 

6205 "Forced Prisma recreate declined by the generation guard; " 

6206 "the engine that failed was not replaced" 

6207 ) 

6208 

6209 await asyncio.wait_for(_do_heavy_reconnect(), timeout=effective_timeout) 

6210 # Only clear the "dead engine" flag after the heavy reconnect 

6211 # actually completed. If `_do_heavy_reconnect()` raises (timeout, 

6212 # missing DATABASE_URL, recreate failure), the flag stays True so 

6213 # the next attempt re-enters the heavy branch instead of silently 

6214 # demoting to the lightweight path. 

6215 self._engine_confirmed_dead = False 

6216 else: 

6217 verbose_proxy_logger.debug("Performing Prisma DB reconnect (engine alive or unknown).") 

6218 

6219 async def _do_direct_reconnect() -> None: 

6220 db_url: Final = os.getenv("DATABASE_URL", "") 

6221 if not db_url: 

6222 verbose_proxy_logger.error("DATABASE_URL not set; cannot reconnect Prisma client.") 

6223 raise RuntimeError("DATABASE_URL not set") 

6224 # Probe the writer BEFORE recreating. A concurrent IAM token 

6225 # refresh may have just replaced the engine (issue #29176); if 

6226 # the writer answers SELECT 1 the connection is already healthy 

6227 # and recreating would needlessly kill that fresh engine. If we 

6228 # do recreate, the entry-snapshot generation lets the wrapper 

6229 # detect a refresh that landed since cycle entry and skip the 

6230 # redundant restart. 

6231 writer: Final = self.writer_db 

6232 if force_recreate is False: 

6233 try: 

6234 if await self._writer_is_read_only(writer): 

6235 verbose_proxy_logger.warning( 

6236 "Writer answers the probe but its session is read-only " 

6237 "(writes fail with SQLSTATE 25006); recreating Prisma client." 

6238 ) 

6239 else: 

6240 verbose_proxy_logger.info( 

6241 "Writer healthy on probe; skipping recreate (engine " 

6242 "likely already replaced by a token refresh)." 

6243 ) 

6244 if isinstance(self.db, RoutingPrismaWrapper): 

6245 self.db.mark_writer_recovered() 

6246 await self._start_engine_watcher() 

6247 return 

6248 except Exception as probe_err: 

6249 verbose_proxy_logger.warning( 

6250 "Writer probe failed (%s); recreating Prisma client.", 

6251 probe_err, 

6252 ) 

6253 # Fresh Prisma client + new engine subprocess. The previous 

6254 # "lightweight" path called `disconnect()` which blocks the 

6255 # event loop on `subprocess.Popen.wait()`; since that call 

6256 # ends up killing the engine anyway, we do it non-blockingly 

6257 # via `_kill_engine_process` inside `recreate_prisma_client`. 

6258 self._cleanup_engine_watcher() 

6259 recreated: Final = await self.db.recreate_prisma_client(db_url, expected_generation=expected_generation) 

6260 await self._start_engine_watcher() 

6261 # Smoke-test the writer specifically; query_raw on the routing 

6262 # wrapper sends to the reader, which would not validate the 

6263 # newly-recreated writer engine. The reader is left to the 

6264 # caller's own retried query, a stronger check than SELECT 1, 

6265 # and a reader that fails to come back sets `_reader_unavailable` 

6266 # so reads fall through to the writer just recreated here. 

6267 await self.writer_db.query_raw("SELECT 1") 

6268 # A recreate can decline: the optimistic-lock guard no-ops when 

6269 # the writer generation moved since cycle entry, and the routing 

6270 # wrapper then leaves the reader untouched as well. Callers that 

6271 # merely suspect a transport blip are happy either way, but a 

6272 # forced caller asked for this engine to be replaced because its 

6273 # session state is poisoned, and it was not. Do not report that 

6274 # as a success: it would reset the consecutive-failure count and 

6275 # log a repair that never happened. 

6276 if force_recreate is True and recreated is False: 

6277 raise _ForcedRecreateDeclined( 

6278 "Forced Prisma recreate declined by the generation guard; " 

6279 "the engine that failed was not replaced" 

6280 ) 

6281 

6282 await asyncio.wait_for(_do_direct_reconnect(), timeout=effective_timeout) 

6283 

6284 def _cooldown_applies(self, stale_read_engine: "_StaleReadEngine | None") -> bool: 

6285 """ 

6286 Whether the reconnect cooldown should still gate this caller. 

6287 

6288 The cooldown collapses a burst of callers onto one recreate, so it 

6289 keeps gating a caller whose named engine has already been replaced: 

6290 that recreate is the one it was waiting for. While that engine is still 

6291 the live one the damage is still being served, so deferring to an 

6292 unrelated reconnect's cooldown would leave it broken until the cooldown 

6293 elapses. 

6294 

6295 A named engine always describes the one that served the failing read 

6296 (see `_query_first_with_cached_plan_fallback`), so it is compared 

6297 against `read_db`, identity included: `read_db` can resolve to a 

6298 different wrapper than it did at observation time. 

6299 

6300 The waiver is withdrawn once a repair of this same engine has been 

6301 tried and failed. A failed recreate leaves the generation where it was, 

6302 so without this every queued caller would still see its own engine live 

6303 and run its own full recreate serially instead of collapsing onto one 

6304 attempt, which is what the cooldown is for. The record is scoped to the 

6305 engine rather than to a global failure count: an unrelated reconnect 

6306 failing somewhere else says nothing about whether this engine can be 

6307 repaired, and gating on it would suppress the recovery this method 

6308 exists to allow. 

6309 

6310 The record is never cleared, and does not need to be. Generations are 

6311 monotonic per wrapper, so once the engine is repaired every later 

6312 caller names a higher one and the entry can never match again. And this 

6313 method is only ever the first half of the gate: the cooldown window 

6314 itself still expires, so an engine that can never be repaired degrades 

6315 to the plain cooldown rather than being suppressed forever. 

6316 """ 

6317 if stale_read_engine is None: 

6318 return True 

6319 if self._failed_recreate_generations.get(stale_read_engine.wrapper) == stale_read_engine.generation: 

6320 return True 

6321 return not stale_read_engine.is_still_live(self.read_db) 

6322 

6323 async def _attempt_reconnect_inside_lock( 

6324 self, 

6325 force: bool, 

6326 reason: str, 

6327 timeout_seconds: float | None, 

6328 force_recreate: bool = False, 

6329 stale_read_engine: "_StaleReadEngine | None" = None, 

6330 ) -> bool: 

6331 now: Final = time.time() 

6332 if ( 

6333 force is False 

6334 and self._cooldown_applies(stale_read_engine) 

6335 and now - self._db_last_reconnect_attempt_ts < self._db_reconnect_cooldown_seconds 

6336 ): 

6337 verbose_proxy_logger.debug( 

6338 "Skipping DB reconnect attempt inside lock due to cooldown. reason=%s", 

6339 reason, 

6340 ) 

6341 return False 

6342 

6343 # Escalate to heavy reconnect after consecutive lightweight failures. 

6344 # When the Prisma engine process is alive but not accepting connections 

6345 # (e.g., startup race condition), lightweight reconnects (disconnect + 

6346 # connect) will never succeed. Force a full Prisma client recreation 

6347 # to recover from this state. 

6348 if self._consecutive_reconnect_failures >= self._reconnect_escalation_threshold: 

6349 verbose_proxy_logger.warning( 

6350 "Escalating to heavy reconnect after %d consecutive failures. reason=%s", 

6351 self._consecutive_reconnect_failures, 

6352 reason, 

6353 ) 

6354 self._engine_confirmed_dead = True 

6355 

6356 verbose_proxy_logger.warning("Attempting Prisma DB reconnect. reason=%s", reason) 

6357 

6358 reconnect_succeeded = False 

6359 try: 

6360 await self._run_reconnect_cycle(timeout_seconds=timeout_seconds, force_recreate=force_recreate) 

6361 reconnect_succeeded = True 

6362 self._consecutive_reconnect_failures = 0 

6363 verbose_proxy_logger.info("Prisma DB reconnect succeeded. reason=%s", reason) 

6364 except _ForcedRecreateDeclined as declined: 

6365 # A decline is raised only when the recreate returns False, which 

6366 # happens only at the generation guard, and the generation moves 

6367 # only after a replacement has connected. So a decline is proof 

6368 # that a replacement SUCCEEDED, and zeroing a consecutive-failure 

6369 # count on that proof is right by definition rather than by 

6370 # analogy to what a reported success used to do. Note what it 

6371 # proves is that the WRITER was replaced, not that this caller's 

6372 # engine was repaired: on a read replica the reader can still be 

6373 # poisoned, since the wrapper returns before touching it. Leaving 

6374 # the count at the threshold would let the escalation check above 

6375 # re-arm the dead-engine flag on the very next attempt and send a 

6376 # healthy replacement back down the probe-free heavy path. 

6377 self._consecutive_reconnect_failures = 0 

6378 verbose_proxy_logger.warning("Prisma DB reconnect declined. reason=%s detail=%s", reason, declined) 

6379 except Exception as reconnect_err: 

6380 self._consecutive_reconnect_failures += 1 

6381 # Remember WHICH engine could not be repaired, so the rest of this 

6382 # caller's burst collapses onto the cooldown instead of each 

6383 # retrying the recreate that just failed. Recorded only for a 

6384 # caller that named a generation: a watchdog or transport-error 

6385 # reconnect failing here is unrelated to any stale read engine and 

6386 # must not suppress its waiver. 

6387 if stale_read_engine is not None: 

6388 # Key off the wrapper the CALLER named, never a freshly resolved 

6389 # `read_db`. A failed reader recreate is itself what marks the 

6390 # reader unavailable, so re-resolving here would file the 

6391 # reader's failure under the writer: the poisoned reader would 

6392 # lose its record and the healthy writer would gain a spurious 

6393 # one, wrong in both directions at once. 

6394 self._failed_recreate_generations = MappingProxyType( 

6395 {**self._failed_recreate_generations, stale_read_engine.wrapper: stale_read_engine.generation} 

6396 ) 

6397 verbose_proxy_logger.error( 

6398 "Prisma DB reconnect failed (%d consecutive). reason=%s error=%s", 

6399 self._consecutive_reconnect_failures, 

6400 reason, 

6401 reconnect_err, 

6402 ) 

6403 finally: 

6404 self._db_last_reconnect_attempt_ts = time.time() 

6405 

6406 return reconnect_succeeded 

6407 

6408 async def attempt_db_reconnect( 

6409 self, 

6410 reason: str, 

6411 force: bool = False, 

6412 timeout_seconds: float | None = None, 

6413 lock_timeout_seconds: float | None = None, 

6414 force_recreate: bool = False, 

6415 stale_read_engine: "_StaleReadEngine | None" = None, 

6416 ) -> bool: 

6417 """ 

6418 Attempt to reconnect the Prisma client in a singleflight manner. 

6419 

6420 `force` bypasses the cooldown unconditionally; `force_recreate` 

6421 bypasses the liveness probe that would otherwise skip recreating a 

6422 reachable engine; `stale_read_engine` bypasses the cooldown only while 

6423 the engine that produced the caller's failure is still the live one 

6424 (see `_cooldown_applies`). 

6425 

6426 A `force_recreate` caller can also get False for a third reason: the 

6427 generation guard declined because another path had already replaced 

6428 the engine, which is a successful outcome reported as False. Callers 

6429 that branch on the return value (`exception_handler` raises on False, 

6430 `auth_checks` retries only on True) would misread that as a dead end, 

6431 and are safe today only because neither passes `force_recreate`. Do 

6432 not add it to one of them without revisiting how it reads the result. 

6433 

6434 Returns: 

6435 bool: True if reconnection succeeded, else False. 

6436 """ 

6437 now: Final = time.time() 

6438 if ( 

6439 force is False 

6440 and self._cooldown_applies(stale_read_engine) 

6441 and now - self._db_last_reconnect_attempt_ts < self._db_reconnect_cooldown_seconds 

6442 ): 

6443 verbose_proxy_logger.debug( 

6444 "Skipping DB reconnect attempt due to cooldown. reason=%s", 

6445 reason, 

6446 ) 

6447 return False 

6448 

6449 if lock_timeout_seconds is None: 

6450 async with self._db_reconnect_lock: 

6451 return await self._attempt_reconnect_inside_lock( 

6452 force, reason, timeout_seconds, force_recreate, stale_read_engine 

6453 ) 

6454 

6455 lock_acquired_by_timeout_task = False 

6456 

6457 async def _acquire_reconnect_lock() -> bool: 

6458 nonlocal lock_acquired_by_timeout_task 

6459 await self._db_reconnect_lock.acquire() 

6460 lock_acquired_by_timeout_task = True 

6461 return True 

6462 

6463 acquire_task: Final = asyncio.create_task(_acquire_reconnect_lock()) 

6464 

6465 async def _abandon_acquire_task() -> None: 

6466 acquire_task.cancel() 

6467 try: 

6468 await acquire_task 

6469 except asyncio.CancelledError: 

6470 pass 

6471 except Exception: 

6472 pass 

6473 

6474 # Defensive cleanup for timeout/cancel race on Python 3.9-3.11. 

6475 if lock_acquired_by_timeout_task: 

6476 try: 

6477 self._db_reconnect_lock.release() 

6478 except RuntimeError: 

6479 pass 

6480 

6481 try: 

6482 done, _pending = await asyncio.wait( 

6483 {acquire_task}, 

6484 timeout=lock_timeout_seconds, 

6485 return_when=asyncio.FIRST_COMPLETED, 

6486 ) 

6487 except asyncio.CancelledError: 

6488 await asyncio.shield(_abandon_acquire_task()) 

6489 raise 

6490 if acquire_task not in done: 

6491 await _abandon_acquire_task() 

6492 verbose_proxy_logger.debug( 

6493 "Skipping DB reconnect attempt due to lock acquisition timeout. reason=%s timeout=%ss", 

6494 reason, 

6495 lock_timeout_seconds, 

6496 ) 

6497 return False 

6498 

6499 try: 

6500 acquire_task.result() 

6501 except Exception as lock_acquire_err: 

6502 verbose_proxy_logger.debug( 

6503 "Skipping DB reconnect attempt due to lock acquisition error. reason=%s error=%s", 

6504 reason, 

6505 lock_acquire_err, 

6506 ) 

6507 return False 

6508 

6509 try: 

6510 return await self._attempt_reconnect_inside_lock( 

6511 force, reason, timeout_seconds, force_recreate, stale_read_engine 

6512 ) 

6513 finally: 

6514 self._db_reconnect_lock.release() 

6515 

6516 async def start_db_health_watchdog_task(self) -> None: 

6517 """Start background tasks that monitor DB health: 

6518 - A periodic SELECT 1 probe that triggers reconnect on network/connection failure. 

6519 - A process-level watcher that detects engine death via waitpid thread, pidfd, or os.kill polling. 

6520 """ 

6521 if self._db_health_watchdog_enabled is not True: 6521 ↛ 6522line 6521 didn't jump to line 6522 because the condition on line 6521 was never true

6522 verbose_proxy_logger.debug("Prisma DB health watchdog disabled via PRISMA_HEALTH_WATCHDOG_ENABLED") 

6523 return 

6524 if self._db_health_watchdog_task is not None: 6524 ↛ 6525line 6524 didn't jump to line 6525 because the condition on line 6524 was never true

6525 return 

6526 # Let planned writer-engine restarts (IAM token refresh, guarded 

6527 # reconnect) re-arm the watcher on the new PID instead of being 

6528 # mistaken for a crash (issue #29176). Set on the writer wrapper since 

6529 # the watcher tracks the writer engine. 

6530 self.writer_db.on_engine_replaced = self._handle_writer_engine_replaced 

6531 self._db_health_watchdog_task = asyncio.create_task(self._db_health_watchdog_loop()) 

6532 verbose_proxy_logger.info( 

6533 "Started Prisma DB health watchdog (interval=%ss, reconnect_cooldown=%ss, probe_timeout=%ss, reconnect_timeout=%ss)", 

6534 self._db_health_watchdog_interval_seconds, 

6535 self._db_reconnect_cooldown_seconds, 

6536 self._db_health_watchdog_probe_timeout_seconds, 

6537 self._db_watchdog_reconnect_timeout_seconds, 

6538 ) 

6539 await self._start_engine_watcher() 

6540 

6541 async def stop_db_health_watchdog_task(self) -> None: 

6542 """Stop DB health watchdog task and engine watcher gracefully.""" 

6543 self._stop_engine_watcher() 

6544 if self._db_health_watchdog_task is None: 6544 ↛ 6545line 6544 didn't jump to line 6545 because the condition on line 6544 was never true

6545 return 

6546 self._db_health_watchdog_task.cancel() 

6547 try: 

6548 await self._db_health_watchdog_task 

6549 except asyncio.CancelledError: 

6550 pass 

6551 self._db_health_watchdog_task = None 

6552 verbose_proxy_logger.info("Stopped Prisma DB health watchdog") 

6553 

6554 def start_view_setup_task(self) -> None: 

6555 if self._view_setup_task is not None: 6555 ↛ 6556line 6555 didn't jump to line 6556 because the condition on line 6555 was never true

6556 return 

6557 self._view_setup_task = asyncio.create_task(self._run_view_setup()) 

6558 

6559 async def stop_view_setup_task(self) -> None: 

6560 if self._view_setup_task is None: 6560 ↛ 6561line 6560 didn't jump to line 6561 because the condition on line 6560 was never true

6561 return 

6562 self._view_setup_task.cancel() 

6563 with contextlib.suppress(asyncio.CancelledError): 

6564 await self._view_setup_task 

6565 self._view_setup_task = None 

6566 

6567 async def _run_view_setup( 

6568 self, 

6569 poll_interval_seconds: float = _VIEW_SETUP_POLL_INTERVAL_SECONDS, 

6570 deadline_seconds: float = _VIEW_SETUP_DEADLINE_SECONDS, 

6571 ) -> _ViewSetupOutcome: 

6572 deadline: Final = time.monotonic() + deadline_seconds 

6573 while True: 

6574 if (attempt := await self._attempt_view_setup()) == "ready": 6574 ↛ 6576line 6574 didn't jump to line 6576 because the condition on line 6574 was always true

6575 return "ready" 

6576 if time.monotonic() >= deadline: 

6577 self._log_view_setup_timeout(attempt, deadline_seconds) 

6578 return "timed_out" 

6579 await asyncio.sleep(poll_interval_seconds) 

6580 

6581 async def _attempt_view_setup(self) -> _ViewSetupAttempt: 

6582 try: 

6583 if not await self._view_setup_gate_table_present(): 6583 ↛ 6584line 6583 didn't jump to line 6584 because the condition on line 6583 was never true

6584 verbose_proxy_logger.debug( 

6585 "Waiting for table %s before creating the spend views", _VIEW_SETUP_GATE_TABLE 

6586 ) 

6587 return "table_missing" 

6588 await self._set_spend_logs_row_count_in_proxy_state() 

6589 await self.check_view_exists() 

6590 return "ready" 

6591 except Exception as e: 

6592 verbose_proxy_logger.warning("Spend view setup attempt failed, retrying until the schema settles: %s", e) 

6593 return e 

6594 

6595 def _log_view_setup_timeout( 

6596 self, last_attempt: Literal["table_missing"] | Exception, deadline_seconds: float 

6597 ) -> None: 

6598 if isinstance(last_attempt, Exception): 

6599 verbose_proxy_logger.error( 

6600 "Gave up creating the spend views after %ss; the last attempt failed with: %s. " 

6601 "Fix that error and restart the proxy.", 

6602 deadline_seconds, 

6603 last_attempt, 

6604 ) 

6605 return 

6606 verbose_proxy_logger.error( 

6607 "Gave up creating the spend views: table %s did not appear within %ss. " 

6608 "Run the database migrations against this database and restart the proxy.", 

6609 _VIEW_SETUP_GATE_TABLE, 

6610 deadline_seconds, 

6611 ) 

6612 

6613 async def _view_setup_gate_table_present(self) -> bool: 

6614 rows: Final = _VIEW_SETUP_GATE_PROBE_ROWS.validate_python( 

6615 await self.db.query_raw("SELECT to_regclass($1) IS NOT NULL AS present", _VIEW_SETUP_GATE_TABLE) 

6616 ) 

6617 return rows[0]["present"] 

6618 

6619 async def _db_health_watchdog_loop(self) -> None: 

6620 while True: 

6621 try: 

6622 await asyncio.sleep(self._db_health_watchdog_interval_seconds) 

6623 await asyncio.wait_for( 

6624 self.db.query_raw("SELECT 1"), 

6625 timeout=self._db_health_watchdog_probe_timeout_seconds, 

6626 ) 

6627 if isinstance(self.db, RoutingPrismaWrapper) and self.db.writer_unavailable: 6627 ↛ 6628line 6627 didn't jump to line 6628 because the condition on line 6627 was never true

6628 await self.attempt_db_reconnect( 

6629 reason="db_health_watchdog_writer_unavailable", 

6630 timeout_seconds=self._db_watchdog_reconnect_timeout_seconds, 

6631 ) 

6632 continue 

6633 if await asyncio.wait_for( 6633 ↛ 6637line 6633 didn't jump to line 6637 because the condition on line 6633 was never true

6634 self._writer_is_read_only(self.writer_db), 

6635 timeout=self._db_health_watchdog_probe_timeout_seconds, 

6636 ): 

6637 await self.recreate_read_only_writer( 

6638 reason="db_health_watchdog_writer_read_only", 

6639 timeout_seconds=self._db_watchdog_reconnect_timeout_seconds, 

6640 ) 

6641 continue 

6642 self._db_read_only_recreate_streak = 0 

6643 self._db_read_only_recreate_ts = 0.0 

6644 except asyncio.CancelledError: 

6645 break 

6646 except Exception as e: 

6647 if isinstance(e, asyncio.TimeoutError) or PrismaDBExceptionHandler.is_database_infrastructure_error(e): 

6648 await self.attempt_db_reconnect( 

6649 reason="db_health_watchdog_connection_error", 

6650 timeout_seconds=self._db_watchdog_reconnect_timeout_seconds, 

6651 ) 

6652 else: 

6653 verbose_proxy_logger.debug("Prisma DB health watchdog observed non-DB error: %s", e) 

6654 

6655 async def recreate_read_only_writer(self, reason: str, timeout_seconds: float | None = None) -> bool: 

6656 """Force-recreate the client behind a writer session that rejects writes 

6657 (SQLSTATE 25006). Each recreate doubles the wait before the next one 

6658 until the watchdog sees a writable session again, so a database that is 

6659 read-only as a whole (replica, failover in progress) does not get its 

6660 engine killed on every watchdog cycle or failed write.""" 

6661 backoff_seconds: Final = min( 

6662 self._db_reconnect_cooldown_seconds * 2 ** min(self._db_read_only_recreate_streak, 10), 

6663 _READ_ONLY_RECREATE_BACKOFF_CAP_SECONDS, 

6664 ) 

6665 if time.time() - self._db_read_only_recreate_ts < backoff_seconds: 

6666 verbose_proxy_logger.debug( 

6667 "Writer session still read-only after %s recreate(s); backing off %ss. reason=%s", 

6668 self._db_read_only_recreate_streak, 

6669 backoff_seconds, 

6670 reason, 

6671 ) 

6672 return False 

6673 verbose_proxy_logger.warning( 

6674 "Writer session is read-only (writes fail with SQLSTATE 25006); recreating Prisma client. reason=%s", 

6675 reason, 

6676 ) 

6677 self._db_read_only_recreate_ts = time.time() 

6678 self._db_read_only_recreate_streak += 1 

6679 return await self.attempt_db_reconnect(reason=reason, timeout_seconds=timeout_seconds, force_recreate=True) 

6680 

6681 async def _writer_is_read_only(self, writer: PrismaWrapper) -> bool: 

6682 """True iff the pooled writer session answers reads but rejects writes (SQLSTATE 25006).""" 

6683 rows: Final = _WRITER_WRITABILITY_PROBE_ROWS.validate_python( 

6684 await writer.query_raw(_WRITER_WRITABILITY_PROBE_SQL) 

6685 ) 

6686 return any(row.get("transaction_read_only") == "on" for row in rows) 

6687 

6688 def _probe_target_wrapper(self) -> PrismaWrapper: 

6689 """The Prisma wrapper a `SELECT 1` health probe actually reaches. 

6690 

6691 `health_check()` issues `query_raw`, which `RoutingPrismaWrapper` sends 

6692 to the reader unless the reader is degraded. The writer's engine state 

6693 therefore says nothing about a probe that failed against the reader, so 

6694 the gate has to follow the same routing rule the probe did. 

6695 """ 

6696 if isinstance(self.db, RoutingPrismaWrapper): 6696 ↛ 6697line 6696 didn't jump to line 6697 because the condition on line 6696 was never true

6697 return self.db.writer if self.db.reader_unavailable else self.db.reader 

6698 return self.db 

6699 

6700 async def _run_health_probe(self, wrapper: PrismaWrapper) -> object: 

6701 """Issue the `SELECT 1` a health check is made of, against `wrapper`. 

6702 

6703 Takes the wrapper rather than re-reading `self.db`, because routing is 

6704 re-resolved on every attribute access: a reader that recovers between 

6705 the caller picking its target and the query going out would send the 

6706 probe to a different engine than the one whose generation the caller is 

6707 about to check, and attribute the failure to the wrong replacement. 

6708 """ 

6709 sql_query: Final = "SELECT 1" 

6710 response: Final[object] = await wrapper.query_raw(sql_query) 

6711 return response 

6712 

6713 async def _probe_answers_now(self, wrapper: PrismaWrapper) -> bool: 

6714 try: 

6715 await self._run_health_probe(wrapper) 

6716 except Exception as probe_error: # noqa: BLE001 # any failure means the database is not answering 

6717 verbose_proxy_logger.debug("Prisma health_check() confirmation probe failed: %s", probe_error) 

6718 return False 

6719 return True 

6720 

6721 async def _planned_engine_replacement_absorbed( 

6722 self, 

6723 e: Exception, 

6724 wrapper: PrismaWrapper, 

6725 generation_before: int, 

6726 ) -> bool: 

6727 """True iff `e` is a connection-class probe failure that a completed 

6728 planned query-engine replacement explains. 

6729 

6730 Planned replacements (RDS IAM token refresh, guarded reconnect) kill the 

6731 running query engine and spawn a new one. A `SELECT 1` probe that races 

6732 that sub-second window fails with a transport error against the engine's 

6733 local HTTP port even though nothing is wrong with the database, and 

6734 reporting it drives a false-positive `db_exceptions` alert on every 

6735 replacement. 

6736 

6737 Two things must both hold, because neither is sufficient alone. The 

6738 engine generation must have moved, which says a replacement completed 

6739 rather than merely being attempted: reconnect attempts during a real 

6740 outage hold the same lock for tens of seconds, so gating on an in-flight 

6741 replacement would swallow most of an outage's alerts. And a fresh probe 

6742 must succeed, because `Prisma.connect()` polls the query engine's own 

6743 `/status` endpoint rather than round-tripping to the database, so a 

6744 future engine that binds before it validates its connection pool would 

6745 let the generation advance with the database still unreachable. 

6746 

6747 Waiting for an in-flight replacement to settle is what makes the 

6748 generation check meaningful, since the generation has not moved yet at 

6749 the instant the probe fails. The wait is generous against a replacement 

6750 that takes well under a second and short enough that an outage-hung 

6751 reconnect is not waited out; a replacement that has not settled by then 

6752 reports rather than stays silent. 

6753 """ 

6754 if not PrismaDBExceptionHandler.is_database_connection_error(e): 

6755 return False 

6756 await wrapper.wait_for_planned_engine_replacement(self.PLANNED_ENGINE_REPLACEMENT_SETTLE_SECONDS) 

6757 if wrapper.engine_generation == generation_before: 

6758 return False 

6759 return await self._probe_answers_now(wrapper) 

6760 

6761 async def _report_health_check_failure( 

6762 self, 

6763 e: Exception, 

6764 duration: float, 

6765 traceback_str: str, 

6766 wrapper: PrismaWrapper, 

6767 generation_before: int, 

6768 ) -> None: 

6769 if await self._planned_engine_replacement_absorbed(e, wrapper, generation_before): 

6770 verbose_proxy_logger.info( 

6771 "Prisma health_check() connection error raced a planned query-engine replacement; " 

6772 "not reporting it as a DB exception: %s", 

6773 e, 

6774 ) 

6775 return 

6776 await self.proxy_logging_obj.failure_handler( 

6777 original_exception=e, 

6778 duration=duration, 

6779 call_type="health_check", 

6780 traceback_str=traceback_str, 

6781 ) 

6782 

6783 @backoff.on_exception( 

6784 backoff.expo, 

6785 Exception, 

6786 max_tries=3, 

6787 max_time=10, 

6788 on_backoff=on_backoff, 

6789 ) 

6790 async def health_check(self): 

6791 """ 

6792 Health check endpoint for the prisma client 

6793 """ 

6794 start_time: Final = time.time() 

6795 probe_wrapper: Final = self._probe_target_wrapper() 

6796 generation_before: Final = probe_wrapper.engine_generation 

6797 try: 

6798 return await self._run_health_probe(probe_wrapper) 

6799 except Exception as e: 

6800 import traceback 

6801 

6802 error_msg: Final = f"LiteLLM Prisma Client Exception health_check(): {e}" 

6803 verbose_proxy_logger.warning(error_msg) 

6804 error_traceback: Final = error_msg + "\n" + traceback.format_exc() 

6805 end_time: Final = time.time() 

6806 _duration: Final = end_time - start_time 

6807 asyncio.create_task( 

6808 self._report_health_check_failure( 

6809 e=e, 

6810 duration=_duration, 

6811 traceback_str=error_traceback, 

6812 wrapper=probe_wrapper, 

6813 generation_before=generation_before, 

6814 ) 

6815 ) 

6816 raise e 

6817 

6818 async def _get_spend_logs_row_count(self) -> int: 

6819 """ 

6820 Get the row count from LiteLLM_SpendLogs table using PostgreSQL system statistics. 

6821 """ 

6822 

6823 @backoff.on_exception( 

6824 backoff.expo, 

6825 Exception, 

6826 max_tries=3, 

6827 max_time=10, 

6828 on_backoff=on_backoff, 

6829 ) 

6830 async def _fetch_row_count() -> int: 

6831 sql_query: Final = """ 

6832 SELECT reltuples::BIGINT 

6833 FROM pg_class 

6834 WHERE oid = '"LiteLLM_SpendLogs"'::regclass; 

6835 """ 

6836 result: Final[Sequence[_RelTuplesRow]] = await self.db.query_raw(query=sql_query) 

6837 return result[0]["reltuples"] 

6838 

6839 try: 

6840 return await _fetch_row_count() 

6841 except Exception as e: 

6842 verbose_proxy_logger.error("Error getting LiteLLM_SpendLogs row count: %s", e) 

6843 return 0 

6844 

6845 @backoff.on_exception( 

6846 backoff.expo, 

6847 Exception, 

6848 max_tries=3, 

6849 max_time=10, 

6850 on_backoff=on_backoff, 

6851 ) 

6852 async def _set_spend_logs_row_count_in_proxy_state(self) -> None: 

6853 """ 

6854 Set the `LiteLLM_SpendLogs`row count in proxy state. 

6855 

6856 This is used later to determine if we should run expensive UI Usage queries. 

6857 """ 

6858 from litellm.proxy.proxy_server import proxy_state 

6859 

6860 _num_spend_logs_rows: Final = await self._get_spend_logs_row_count() 

6861 proxy_state.set_proxy_state_variable( 

6862 variable_name="spend_logs_row_count", 

6863 value=_num_spend_logs_rows, 

6864 ) 

6865 

6866 # Health Check Database Methods 

6867 def _validate_response_time(self, response_time_ms: float | None) -> float | None: 

6868 """Validate and clean response time value""" 

6869 if response_time_ms is None: 6869 ↛ 6870line 6869 didn't jump to line 6870 because the condition on line 6869 was never true

6870 return None 

6871 try: 

6872 value: Final = float(response_time_ms) 

6873 return value if math.isfinite(value) else None 

6874 except (ValueError, TypeError): 

6875 verbose_proxy_logger.warning("Invalid response_time_ms value: %s", response_time_ms) 

6876 return None 

6877 

6878 def _clean_details(self, details: dict | None) -> dict | None: 

6879 """Clean and validate details JSON""" 

6880 if not isinstance(details, dict): 6880 ↛ 6882line 6880 didn't jump to line 6882 because the condition on line 6880 was always true

6881 return None 

6882 try: 

6883 return safe_json_loads(safe_dumps(details)) 

6884 except Exception as e: 

6885 verbose_proxy_logger.warning("Failed to clean details JSON: %s", e) 

6886 return None 

6887 

6888 async def save_health_check_result( 

6889 self, 

6890 model_name: str, 

6891 status: str, 

6892 healthy_count: int = 0, 

6893 unhealthy_count: int = 0, 

6894 error_message: str | None = None, 

6895 response_time_ms: float | None = None, 

6896 details: dict | None = None, 

6897 checked_by: str | None = None, 

6898 model_id: str | None = None, 

6899 ): 

6900 """Save health check result to database""" 

6901 try: 

6902 # Build base data with required fields 

6903 health_check_data: Final = { 

6904 "model_name": str(model_name), 

6905 "status": str(status), 

6906 "healthy_count": int(healthy_count), 

6907 "unhealthy_count": int(unhealthy_count), 

6908 } 

6909 

6910 # Add optional fields using dict comprehension and helper methods 

6911 optional_fields: Final = { 

6912 "error_message": str(error_message)[:500] if error_message else None, 

6913 "response_time_ms": self._validate_response_time(response_time_ms), 

6914 "details": self._clean_details(details), 

6915 "checked_by": str(checked_by) if checked_by else None, 

6916 "model_id": str(model_id) if model_id else None, 

6917 } 

6918 

6919 # Add only non-None optional fields 

6920 health_check_data.update({k: v for k, v in optional_fields.items() if v is not None}) 

6921 

6922 verbose_proxy_logger.debug("Saving health check data: %s", health_check_data) 

6923 return await HealthCheckRepository(self).table.create(data=health_check_data) 

6924 

6925 except Exception as e: 

6926 verbose_proxy_logger.error("Error saving health check result for model %s: %s", model_name, e) 

6927 return None 

6928 

6929 async def get_health_check_history( 

6930 self, 

6931 model_name: str | None = None, 

6932 limit: int = 100, 

6933 offset: int = 0, 

6934 status_filter: str | None = None, 

6935 ) -> "Sequence[prisma_models.LiteLLM_HealthCheckTable]": 

6936 """ 

6937 Get health check history with optional filtering 

6938 """ 

6939 try: 

6940 where_clause: Final[dict[str, str]] = {} 

6941 if model_name: 

6942 where_clause["model_name"] = model_name 

6943 if status_filter: 

6944 where_clause["status"] = status_filter 

6945 

6946 results: Final = await HealthCheckRepository(self).table.find_many( 

6947 where=where_clause, 

6948 order={"checked_at": "desc"}, 

6949 take=limit, 

6950 skip=offset, 

6951 ) 

6952 return results 

6953 except Exception as e: 

6954 verbose_proxy_logger.error("Error getting health check history: %s", e) 

6955 return [] 

6956 

6957 async def get_all_latest_health_checks(self) -> tuple[LatestHealthCheckRow, ...]: 

6958 """Latest health check per (model_id, model_name), deduplicated in Postgres.""" 

6959 return await fetch_latest_health_checks(self) 

6960 

6961 async def get_latest_health_checks_for_models(self, model_names: Sequence[str]) -> tuple[LatestHealthCheckRow, ...]: 

6962 """Same as ``get_all_latest_health_checks``, bounded to the named models.""" 

6963 return await fetch_latest_health_checks_for_models(self, model_names) 

6964 

6965 

6966### HELPER FUNCTIONS ### 

6967 

6968 

6969async def _cache_user_row(user_id: str, cache: DualCache, db: PrismaClient): 

6970 """ 

6971 Check if a user_id exists in cache, 

6972 if not retrieve it. 

6973 """ 

6974 cache_key: Final = f"{user_id}_user_api_key_user_id" 

6975 response: Final = cache.get_cache(key=cache_key) 

6976 if response is None: # Cache miss 

6977 user_row: Final = await db.get_data(user_id=user_id) 

6978 if user_row is not None: 

6979 print_verbose(f"User Row: {user_row}, type = {type(user_row)}") 

6980 if hasattr(user_row, "model_dump_json") and callable(getattr(user_row, "model_dump_json")): 

6981 cache_value: Final[str] = user_row.model_dump_json() 

6982 cache.set_cache(key=cache_key, value=cache_value, ttl=600) # store for 10 minutes 

6983 

6984 

6985def _should_use_smtp_ssl(smtp_port: int) -> bool: 

6986 """ 

6987 Port 465 expects an immediate TLS handshake (implicit SSL), so a plain 

6988 smtplib.SMTP connection hangs waiting for an SMTP banner. Use SMTP_SSL 

6989 there, or when SMTP_USE_SSL is explicitly enabled. 

6990 """ 

6991 return os.getenv("SMTP_USE_SSL", "False") == "True" or smtp_port == 465 

6992 

6993 

6994def _create_smtp_connection(smtp_host: str, smtp_port: int, timeout: float) -> smtplib.SMTP: 

6995 if _should_use_smtp_ssl(smtp_port=smtp_port): 

6996 return smtplib.SMTP_SSL(host=smtp_host, port=smtp_port, context=ssl.create_default_context(), timeout=timeout) 

6997 return smtplib.SMTP(host=smtp_host, port=smtp_port, timeout=timeout) 

6998 

6999 

7000def _send_smtp_message( 

7001 email_message: MIMEMultipart, 

7002 smtp_host: str, 

7003 smtp_port: int, 

7004 smtp_username: str | None, 

7005 smtp_password: str | None, 

7006 sender_email: str, 

7007 receiver_email: str, 

7008 timeout: float, 

7009) -> None: 

7010 using_ssl: Final = _should_use_smtp_ssl(smtp_port=smtp_port) 

7011 with _create_smtp_connection( 

7012 smtp_host=smtp_host, 

7013 smtp_port=smtp_port, 

7014 timeout=timeout, 

7015 ) as server: 

7016 if not using_ssl and os.getenv("SMTP_TLS", "True") != "False": 

7017 server.starttls(context=ssl.create_default_context()) 

7018 

7019 if smtp_username and smtp_password: 

7020 server.login( 

7021 user=smtp_username, 

7022 password=smtp_password, 

7023 ) 

7024 

7025 server.send_message( 

7026 msg=email_message, 

7027 from_addr=sender_email, 

7028 to_addrs=receiver_email, 

7029 ) 

7030 

7031 

7032async def send_email( 

7033 receiver_email: str | None = None, 

7034 subject: str | None = None, 

7035 html: str | None = None, 

7036): 

7037 """ 

7038 smtp_host, 

7039 smtp_port, 

7040 smtp_username, 

7041 smtp_password, 

7042 sender_name, 

7043 sender_email, 

7044 """ 

7045 ## SERVER SETUP ## 

7046 

7047 smtp_host: Final = os.getenv("SMTP_HOST") 

7048 smtp_port: Final = int(os.getenv("SMTP_PORT", "587")) # default to port 587 

7049 smtp_username: Final = os.getenv("SMTP_USERNAME") 

7050 smtp_password: Final = os.getenv("SMTP_PASSWORD") 

7051 sender_email: Final = os.getenv("SMTP_SENDER_EMAIL", None) 

7052 if sender_email is None: 

7053 raise ValueError("Trying to use SMTP, but SMTP_SENDER_EMAIL is not set") 

7054 if receiver_email is None: 

7055 raise ValueError(f"No receiver email provided for SMTP email. {receiver_email}") 

7056 if subject is None: 

7057 raise ValueError(f"No subject provided for SMTP email. {subject}") 

7058 if html is None: 

7059 raise ValueError(f"No HTML body provided for SMTP email. {html}") 

7060 

7061 ## EMAIL SETUP ## 

7062 email_message: Final = MIMEMultipart() 

7063 email_message["From"] = sender_email 

7064 email_message["To"] = receiver_email 

7065 email_message["Subject"] = subject 

7066 verbose_proxy_logger.debug("sending email from %s to %s", sender_email, receiver_email) 

7067 

7068 if smtp_host is None: 

7069 raise ValueError("Trying to use SMTP, but SMTP_HOST is not set") 

7070 

7071 # Attach the body to the email 

7072 email_message.attach(MIMEText(html, "html")) 

7073 

7074 try: 

7075 smtp_timeout: Final = float(os.getenv("SMTP_TIMEOUT", "30")) 

7076 await asyncio.to_thread( 

7077 _send_smtp_message, 

7078 email_message=email_message, 

7079 smtp_host=smtp_host, 

7080 smtp_port=smtp_port, 

7081 smtp_username=smtp_username, 

7082 smtp_password=smtp_password, 

7083 sender_email=sender_email, 

7084 receiver_email=receiver_email, 

7085 timeout=smtp_timeout, 

7086 ) 

7087 

7088 except Exception as e: 

7089 verbose_proxy_logger.exception("An error occurred while sending the email:" + str(e)) 

7090 

7091 

7092def hash_token(token: str): 

7093 import hashlib 

7094 

7095 # Hash the string using SHA-256 

7096 hashed_token: Final = hashlib.sha256(token.encode()).hexdigest() 

7097 

7098 return hashed_token 

7099 

7100 

7101def hash_password(password: str) -> str: 

7102 """Hash a password using scrypt with a random salt.""" 

7103 import base64 

7104 import hashlib 

7105 import os 

7106 

7107 salt: Final = os.urandom(16) 

7108 dk: Final = hashlib.scrypt(password.encode(), salt=salt, n=16384, r=8, p=1, dklen=32) 

7109 return "scrypt:" + base64.b64encode(salt + dk).decode() 

7110 

7111 

7112def verify_password(password: str, stored: str) -> bool: 

7113 """Verify a password against a stored hash. Supports scrypt and SHA256.""" 

7114 import base64 

7115 import hashlib 

7116 import secrets 

7117 

7118 if stored.startswith("scrypt:"): 

7119 try: 

7120 raw: Final = base64.b64decode(stored[7:]) 

7121 salt, dk = raw[:16], raw[16:] 

7122 dk2: Final = hashlib.scrypt(password.encode(), salt=salt, n=16384, r=8, p=1, dklen=32) 

7123 return secrets.compare_digest(dk, dk2) 

7124 except Exception: 

7125 return False 

7126 # SHA256 fallback (not vulnerable to pass-the-hash: checks sha256(input) == stored) 

7127 if len(stored) == 64 and all(c in "0123456789abcdef" for c in stored): 

7128 return secrets.compare_digest(hashlib.sha256(password.encode()).hexdigest().encode(), stored.encode()) 

7129 return False 

7130 

7131 

7132async def migrate_passwords_to_scrypt_async(prisma_client) -> str: 

7133 """ 

7134 Migrate plaintext passwords in the DB to scrypt. SHA256 passwords 

7135 are left alone (they migrate on next login via the SHA256 fallback). 

7136 Skips quickly if no plaintext passwords exist. 

7137 """ 

7138 all_with_pw: Final = await UserRepository(prisma_client).table.find_many( 

7139 where={"password": {"not": None}}, 

7140 ) 

7141 

7142 def _is_sha256_hex(s: str) -> bool: 

7143 return len(s) == 64 and all(c in "0123456789abcdef" for c in s) 

7144 

7145 plaintext_users: Final = [ 

7146 (u.user_id, u.password) 

7147 for u in all_with_pw 

7148 if u.password and not u.password.startswith("scrypt:") and not _is_sha256_hex(u.password) 

7149 ] 

7150 if not plaintext_users: 7150 ↛ 7153line 7150 didn't jump to line 7153 because the condition on line 7150 was always true

7151 return "No plaintext passwords found" 

7152 

7153 for user_id, plaintext_password in plaintext_users: 

7154 await UserRepository(prisma_client).table.update( 

7155 where={"user_id": user_id}, 

7156 data={"password": hash_password(plaintext_password)}, 

7157 ) 

7158 return f"Migrated {len(plaintext_users)} plaintext passwords to scrypt" 

7159 

7160 

7161def _hash_token_if_needed(token: str) -> str: 

7162 """ 

7163 Hash the token if it's a string and starts with "sk-" 

7164 

7165 Else return the token as is 

7166 """ 

7167 if token.startswith("sk-"): 

7168 return hash_token(token=token) 

7169 else: 

7170 return token 

7171 

7172 

7173async def enqueue_spend_logs( 

7174 prisma_client: PrismaClient, 

7175 logs: Sequence[Mapping[str, object]], 

7176 *, 

7177 at_head: bool = False, 

7178 max_bytes: int = SPEND_LOG_QUEUE_MAX_BYTES, 

7179) -> None: 

7180 """Queue spend logs for the next flush, held under ``SPEND_LOG_QUEUE_MAX_BYTES``. 

7181 

7182 ``at_head`` replays a batch the DB refused, so it flushes before the logs 

7183 that piled up during the outage. Past the budget the oldest logs are 

7184 dropped, which keeps a long outage from growing the queue until the pod 

7185 dies. 

7186 """ 

7187 added: Final = sum(spend_log_row_bytes(row) for row in logs) 

7188 async with prisma_client._spend_log_transactions_lock: 

7189 queued: Final = ( 

7190 tuple(logs) + tuple(prisma_client.spend_log_transactions) 

7191 if at_head 

7192 else tuple(prisma_client.spend_log_transactions) + tuple(logs) 

7193 ) 

7194 kept, kept_bytes = spend_log_queue_within_budget(queued, PrismaClient.spend_log_queue_bytes + added, max_bytes) 

7195 prisma_client.spend_log_transactions[:] = kept 

7196 PrismaClient.spend_log_queue_bytes = kept_bytes 

7197 if len(kept) < len(queued): 7197 ↛ 7198line 7197 didn't jump to line 7198 because the condition on line 7197 was never true

7198 verbose_proxy_logger.error( 

7199 "Spend tracking - spend log queue is at its %d byte budget; dropped the %d oldest spend logs", 

7200 max_bytes, 

7201 len(queued) - len(kept), 

7202 ) 

7203 

7204 

7205def request_spend_log_flush(prisma_client: PrismaClient) -> None: 

7206 """Wake this client's queue monitor now rather than leaving the rows for its next poll. 

7207 

7208 The Responses API hands the client an id it can chain from straight away, and that 

7209 lookup reads the DB, so the row cannot sit in this worker's queue for a poll interval. 

7210 Repeated requests coalesce into the monitor's next pass, so the batching holds. 

7211 A request made before the monitor is running is dropped, and loses nothing: the 

7212 monitor reads the queue on its first pass, before it ever waits on a request. 

7213 """ 

7214 flush_requested: Final = prisma_client.spend_log_flush_requested 

7215 if flush_requested is not None: 7215 ↛ exitline 7215 didn't return from function 'request_spend_log_flush' because the condition on line 7215 was always true

7216 flush_requested.set() 

7217 

7218 

7219async def _wait_for_spend_log_flush_request(flush_requested: asyncio.Event, interval: float) -> bool: 

7220 """Wait out ``interval``, returning early and True when a flush was requested.""" 

7221 try: 

7222 await asyncio.wait_for(flush_requested.wait(), timeout=interval) 

7223 except asyncio.TimeoutError: 

7224 return False 

7225 flush_requested.clear() 

7226 return True 

7227 

7228 

7229async def dequeue_spend_logs(prisma_client: PrismaClient, limit: int) -> list[dict[str, object]]: 

7230 """Take up to ``limit`` of the oldest queued spend logs off the queue. 

7231 

7232 Every enqueue and dequeue goes through this pair so the byte total the 

7233 queue is bounded by stays in step with what the queue actually holds. 

7234 """ 

7235 async with prisma_client._spend_log_transactions_lock: 

7236 popped: Final = prisma_client.spend_log_transactions[:limit] 

7237 prisma_client.spend_log_transactions[:] = prisma_client.spend_log_transactions[limit:] 

7238 PrismaClient.spend_log_queue_bytes = max( 

7239 0, PrismaClient.spend_log_queue_bytes - sum(spend_log_row_bytes(row) for row in popped) 

7240 ) 

7241 return popped 

7242 

7243 

7244class ProxyUpdateSpend: 

7245 @staticmethod 

7246 async def update_end_user_spend( 

7247 n_retry_times: int, 

7248 prisma_client: PrismaClient, 

7249 proxy_logging_obj: ProxyLogging, 

7250 end_user_list_transactions: dict[str, float], 

7251 ): 

7252 for i in range(n_retry_times + 1): 7252 ↛ exitline 7252 didn't return from function 'update_end_user_spend' because the loop on line 7252 didn't complete

7253 start_time = time.time() 

7254 try: 

7255 async with prisma_client.db.tx(timeout=timedelta(seconds=60)) as transaction: 

7256 batcher: _EndUserSpendBatch 

7257 async with transaction.batch_() as batcher: 

7258 # Sort by end_user_id for consistent lock ordering across pods to prevent deadlocks. 

7259 for end_user_id, response_cost in sorted(end_user_list_transactions.items()): 

7260 if litellm.max_end_user_budget is not None: 7260 ↛ 7261line 7260 didn't jump to line 7261 because the condition on line 7260 was never true

7261 pass 

7262 batcher.litellm_endusertable.upsert( 

7263 where={"user_id": end_user_id}, 

7264 data={ 

7265 "create": { 

7266 "user_id": end_user_id, 

7267 "spend": response_cost, 

7268 "blocked": False, 

7269 }, 

7270 "update": {"spend": {"increment": response_cost}}, 

7271 }, 

7272 ) 

7273 

7274 break 

7275 except Exception as e: 

7276 await DBSpendUpdateWriter._handle_spend_update_failure( 

7277 e=e, 

7278 attempt=i, 

7279 n_retry_times=n_retry_times, 

7280 start_time=start_time, 

7281 proxy_logging_obj=proxy_logging_obj, 

7282 ) 

7283 

7284 @staticmethod 

7285 async def update_spend_logs( 

7286 n_retry_times: int, 

7287 prisma_client: PrismaClient, 

7288 db_writer_client: AsyncHTTPHandler | None, 

7289 proxy_logging_obj: ProxyLogging, 

7290 logs_to_process: list[dict[str, object]] | None = None, 

7291 ): 

7292 BATCH_SIZE: Final = 1000 # Preferred size of each batch to write to the database 

7293 MAX_LOGS_PER_INTERVAL: Final = 10000 # Maximum number of logs to flush in a single interval 

7294 popped_batch = False 

7295 if logs_to_process is None: 7295 ↛ 7296line 7295 didn't jump to line 7296 because the condition on line 7295 was never true

7296 logs_to_process = await dequeue_spend_logs(prisma_client, MAX_LOGS_PER_INTERVAL) 

7297 popped_batch = True 

7298 if len(logs_to_process) > 0: 

7299 verbose_proxy_logger.info( 

7300 "Spend tracking - processing %d spend logs for DB write", 

7301 len(logs_to_process), 

7302 ) 

7303 start_time: Final = time.time() 

7304 try: 

7305 for i in range(n_retry_times + 1): 7305 ↛ 7373line 7305 didn't jump to line 7373 because the loop on line 7305 didn't complete

7306 try: 

7307 base_url = os.getenv("SPEND_LOGS_URL", None) 

7308 if len(logs_to_process) > 0 and base_url is not None and db_writer_client is not None: 7308 ↛ 7309line 7308 didn't jump to line 7309 because the condition on line 7308 was never true

7309 if not base_url.endswith("/"): 

7310 base_url += "/" 

7311 verbose_proxy_logger.debug("base_url: %s", base_url) 

7312 json_data = json.dumps(logs_to_process) 

7313 response = await db_writer_client.post( 

7314 url=base_url + "spend/update", 

7315 data=json_data, 

7316 headers={"Content-Type": "application/json"}, 

7317 ) 

7318 del json_data 

7319 if response.status_code == 200: 

7320 # Items already removed from queue at start of function 

7321 pass 

7322 else: 

7323 for j in range(0, len(logs_to_process), BATCH_SIZE): 

7324 batch = logs_to_process[j : j + BATCH_SIZE] 

7325 batch_with_dates = [prisma_client.jsonify_object({**entry}) for entry in batch] 

7326 isolation_budget = MAX_SPEND_LOG_ISOLATION_FAILURES_PER_BATCH 

7327 for statement_rows in spend_log_write_batches( 

7328 batch_with_dates, 

7329 SPEND_LOG_WRITE_BATCH_MAX_BYTES, 

7330 SPEND_LOG_WRITE_BATCH_MAX_ROWS, 

7331 ): 

7332 isolation_budget = await _create_spend_logs_with_poison_isolation( 

7333 SpendLogsRepository(prisma_client), 

7334 statement_rows, 

7335 isolation_budget, 

7336 ) 

7337 verbose_proxy_logger.debug("Flushed %s logs to the DB.", len(batch)) 

7338 # Explicitly clear batch memory 

7339 del batch, batch_with_dates 

7340 

7341 # Items already removed from queue at start of function 

7342 async with prisma_client._spend_log_transactions_lock: 

7343 remaining_count = len(prisma_client.spend_log_transactions) 

7344 verbose_proxy_logger.debug( 

7345 "%s logs processed. Remaining in queue: %s", len(logs_to_process), remaining_count 

7346 ) 

7347 break 

7348 except Exception as e: 

7349 if not _is_transient_spend_log_write_error(e): 

7350 if PrismaDBExceptionHandler.is_prisma_error(e): 

7351 await requeue_spend_logs(prisma_client, proxy_logging_obj, logs_to_process) 

7352 verbose_proxy_logger.warning( 

7353 "Spend tracking - DB error writing spend logs, requeued %d rows for the next flush. error=%s", 

7354 len(logs_to_process), 

7355 str(e), 

7356 ) 

7357 raise 

7358 verbose_proxy_logger.warning( 

7359 "Spend tracking - transient DB error writing spend logs, retry %d/%d. logs_count=%d, error=%s", 

7360 i + 1, 

7361 n_retry_times, 

7362 len(logs_to_process), 

7363 str(e), 

7364 ) 

7365 if i >= n_retry_times: 

7366 await requeue_spend_logs(prisma_client, proxy_logging_obj, logs_to_process) 

7367 raise 

7368 await asyncio.sleep(2**i) 

7369 except Exception as e: 

7370 _raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) 

7371 finally: 

7372 # Clean up logs_to_process only if we popped it (caller-owned otherwise) 

7373 if popped_batch: 7373 ↛ 7374line 7373 didn't jump to line 7374 because the condition on line 7373 was never true

7374 del logs_to_process 

7375 

7376 @staticmethod 

7377 def disable_spend_updates() -> bool: 

7378 """ 

7379 returns True if should not update spend in db 

7380 Skips writing spend logs and updates to key, team, user spend to DB 

7381 """ 

7382 from litellm.proxy.proxy_server import general_settings 

7383 

7384 if general_settings.get("disable_spend_updates") is True: 7384 ↛ 7385line 7384 didn't jump to line 7385 because the condition on line 7384 was never true

7385 return True 

7386 return False 

7387 

7388 

7389async def update_spend( 

7390 prisma_client: PrismaClient, 

7391 db_writer_client: AsyncHTTPHandler | None, 

7392 proxy_logging_obj: ProxyLogging, 

7393): 

7394 """ 

7395 Batch write updates to db. 

7396 

7397 Triggered every minute. 

7398 

7399 NOTE: This job now skips tag spend updates, which are handled by a separate 

7400 scheduler job (update_daily_tag_spend) at a longer interval to reduce contention. 

7401 

7402 Requires: 

7403 user_id_list: dict, 

7404 keys_list: list, 

7405 team_list: list, 

7406 spend_logs: list, 

7407 """ 

7408 n_retry_times: Final = 3 

7409 await proxy_logging_obj.db_spend_update_writer.db_update_spend_transaction_handler( 

7410 prisma_client=prisma_client, 

7411 n_retry_times=n_retry_times, 

7412 proxy_logging_obj=proxy_logging_obj, 

7413 ) 

7414 

7415 ### UPDATE SPEND LOGS ### 

7416 await recover_parked_spend_logs(prisma_client, proxy_logging_obj) 

7417 # Check queue size with lock protection 

7418 queue_size: Final = await _total_queued_spend_transactions(prisma_client) 

7419 verbose_proxy_logger.debug("Spend Logs transactions: %s", queue_size) 

7420 

7421 # Process spend log transactions when called directly. 

7422 # This keeps backwards compatibility with the old behavior. 

7423 # See update_spend_logs_job and _monitor_spend_logs_queue for the new behavior. 

7424 # Safe to keep: under high concurrency this can take up to ~30s to run, 

7425 # so it's unlikely to overlap with monitor_spend_logs_queue. 

7426 if queue_size > 0: 

7427 await update_spend_logs_job( 

7428 prisma_client=prisma_client, 

7429 db_writer_client=db_writer_client, 

7430 proxy_logging_obj=proxy_logging_obj, 

7431 ) 

7432 

7433 

7434async def _park_spend_logs_in_redis(proxy_logging_obj: ProxyLogging, rows: Sequence[Mapping[str, object]]) -> bool: 

7435 try: 

7436 return await proxy_logging_obj.db_spend_update_writer.redis_update_buffer.store_spend_logs_in_redis(rows) 

7437 except Exception as e: # noqa: BLE001 # a Redis fault falls back to the in-memory queue, never loses the rows 

7438 verbose_proxy_logger.warning( 

7439 "Spend tracking - could not park spend logs in Redis, keeping them in memory: %s", e 

7440 ) 

7441 return False 

7442 

7443 

7444async def requeue_spend_logs( 

7445 prisma_client: PrismaClient, 

7446 proxy_logging_obj: ProxyLogging, 

7447 rows: Sequence[Mapping[str, object]], 

7448) -> None: 

7449 """Park rows from a failed or cancelled write in Redis, falling back to the head of the in-memory queue.""" 

7450 if await _park_spend_logs_in_redis(proxy_logging_obj, rows): 

7451 return 

7452 await enqueue_spend_logs(prisma_client, rows, at_head=True) 

7453 

7454 

7455async def recover_parked_spend_logs( 

7456 prisma_client: PrismaClient, 

7457 proxy_logging_obj: ProxyLogging, 

7458 limit: int = REDIS_SPEND_LOGS_BUFFER_DEQUEUE_COUNT, 

7459) -> int: 

7460 """Move spend-log rows parked in Redis back to the head of the in-memory queue for the next write.""" 

7461 try: 

7462 rows: Final = ( 

7463 await proxy_logging_obj.db_spend_update_writer.redis_update_buffer.get_spend_logs_from_redis_buffer(limit) 

7464 ) 

7465 except Exception as e: # noqa: BLE001 # Redis being down must not stop the regular in-memory flush 

7466 verbose_proxy_logger.warning("Spend tracking - could not read parked spend logs from Redis: %s", e) 

7467 return 0 

7468 if len(rows) == 0: 7468 ↛ 7470line 7468 didn't jump to line 7470 because the condition on line 7468 was always true

7469 return 0 

7470 try: 

7471 await enqueue_spend_logs(prisma_client, rows, at_head=True) 

7472 except BaseException: 

7473 await _park_spend_logs_in_redis(proxy_logging_obj, rows) 

7474 raise 

7475 verbose_proxy_logger.info("Spend tracking - recovered %d parked spend log rows from Redis", len(rows)) 

7476 return len(rows) 

7477 

7478 

7479async def _total_queued_spend_transactions(prisma_client: PrismaClient) -> int: 

7480 """Pending entries across every request-time spend queue, sized under each queue's 

7481 lock. Every drain trigger reads this one owner, so a queue added later joins the 

7482 direct path, the batch job's emptiness check and the monitor at once.""" 

7483 async with prisma_client._spend_log_transactions_lock: 

7484 spend_queue_size: Final = len(prisma_client.spend_log_transactions) 

7485 async with prisma_client._tool_usage_transactions_lock: 

7486 tool_queue_size: Final = len(prisma_client.tool_usage_transactions) 

7487 async with prisma_client._autorouter_turn_transactions_lock: 

7488 autorouter_queue_size: Final = len(prisma_client.autorouter_turn_transactions) 

7489 from litellm.proxy.db.shadow_eval_funnel import pending_shadow_eval_funnel_events 

7490 

7491 async with prisma_client.baseline_accounting_lock: 

7492 baseline_queue_size: Final = len(prisma_client.baseline_accounting_transactions) 

7493 return ( 

7494 spend_queue_size 

7495 + tool_queue_size 

7496 + autorouter_queue_size 

7497 + baseline_queue_size 

7498 + pending_shadow_eval_funnel_events() 

7499 ) 

7500 

7501 

7502async def update_daily_tag_spend( 

7503 prisma_client: PrismaClient, 

7504 proxy_logging_obj: ProxyLogging, 

7505): 

7506 """ 

7507 Separate scheduler job to commit daily tag spend updates. 

7508 

7509 Runs at a longer interval (2.3x default) than the main update_spend job 

7510 to reduce query contention for DailyTagSpend table. 

7511 

7512 This is called by a dedicated scheduler job and does NOT process: 

7513 - Regular spend updates (user, key, team, org) 

7514 - End-user spend 

7515 - Agent spend 

7516 - Spend logs 

7517 

7518 Only processes tag spend transactions from the daily_tag_spend_update_queue. 

7519 

7520 Args: 

7521 prisma_client: PrismaClient instance 

7522 proxy_logging_obj: ProxyLogging instance for error handling 

7523 """ 

7524 n_retry_times: Final = 3 

7525 try: 

7526 if proxy_logging_obj.db_spend_update_writer.redis_update_buffer._should_commit_spend_updates_to_redis(): 7526 ↛ 7527line 7526 didn't jump to line 7527 because the condition on line 7526 was never true

7527 await proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis( 

7528 prisma_client=prisma_client, 

7529 n_retry_times=n_retry_times, 

7530 proxy_logging_obj=proxy_logging_obj, 

7531 ) 

7532 else: 

7533 await proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db( 

7534 prisma_client=prisma_client, 

7535 n_retry_times=n_retry_times, 

7536 proxy_logging_obj=proxy_logging_obj, 

7537 ) 

7538 except Exception as e: 

7539 # NOTE: keep this as a plain ``error`` (no traceback) to match the 

7540 # historical behavior of this site. ``spend_log_error`` would attach 

7541 # the active exception's traceback whenever the suppression env var 

7542 # is unset, which would be a regression for operators who never saw 

7543 # one here before. 

7544 verbose_proxy_logger.error("Error updating daily tag spend: %s", e) 

7545 

7546 

7547async def update_spend_logs_job( 

7548 prisma_client: PrismaClient, 

7549 db_writer_client: AsyncHTTPHandler | None, 

7550 proxy_logging_obj: ProxyLogging, 

7551): 

7552 """ 

7553 Job to process spend_log_transactions queue. 

7554 

7555 This job is triggered based on queue size rather than time. 

7556 Pops the batch once, writes spend logs, then runs guardrail usage tracking. 

7557 """ 

7558 from litellm.proxy.db.baseline_accounting import flush_baseline_accounting 

7559 

7560 if await _total_queued_spend_transactions(prisma_client) == 0: 7560 ↛ 7561line 7560 didn't jump to line 7561 because the condition on line 7560 was never true

7561 await flush_baseline_accounting(prisma_client) 

7562 return 

7563 async with prisma_client.spend_log_write_lock: 

7564 await _run_spend_logs_job(prisma_client, db_writer_client, proxy_logging_obj) 

7565 

7566 

7567async def _run_spend_logs_job( 

7568 prisma_client: PrismaClient, 

7569 db_writer_client: AsyncHTTPHandler | None, 

7570 proxy_logging_obj: ProxyLogging, 

7571) -> None: 

7572 from litellm.proxy.db.baseline_accounting import flush_baseline_accounting 

7573 

7574 n_retry_times: Final = 3 

7575 MAX_LOGS_PER_INTERVAL: Final = 10000 

7576 

7577 logs_to_process: Final = await dequeue_spend_logs(prisma_client, MAX_LOGS_PER_INTERVAL) 

7578 

7579 try: 

7580 await ProxyUpdateSpend.update_spend_logs( 

7581 n_retry_times=n_retry_times, 

7582 prisma_client=prisma_client, 

7583 proxy_logging_obj=proxy_logging_obj, 

7584 db_writer_client=db_writer_client, 

7585 logs_to_process=logs_to_process, 

7586 ) 

7587 except asyncio.CancelledError: 

7588 await requeue_spend_logs(prisma_client, proxy_logging_obj, logs_to_process) 

7589 verbose_proxy_logger.warning( 

7590 "Spend tracking - spend log write cancelled, requeued %d rows for the next flush", 

7591 len(logs_to_process), 

7592 ) 

7593 raise 

7594 

7595 # Guardrail/policy usage tracking (same batch, outside spend-logs update) 

7596 try: 

7597 from litellm.proxy.guardrails.usage_tracking import ( 

7598 process_spend_logs_guardrail_usage, 

7599 ) 

7600 

7601 await process_spend_logs_guardrail_usage( 

7602 prisma_client=prisma_client, 

7603 logs_to_process=logs_to_process, 

7604 ) 

7605 except Exception as guardrail_tracking_err: 

7606 verbose_proxy_logger.warning( 

7607 "Spend tracking - guardrail usage tracking failed (non-fatal): %s", 

7608 guardrail_tracking_err, 

7609 ) 

7610 

7611 # Tool usage tracking: drain the request-time queue into the tool index and the 

7612 # LiteLLM_DailyToolSpend rollup. Never retried; a dropped batch is permanently 

7613 # absent from the rollup, so failures log at error. 

7614 async with prisma_client._tool_usage_transactions_lock: 

7615 tool_usage_to_process: Final = prisma_client.tool_usage_transactions[:MAX_LOGS_PER_INTERVAL] 

7616 prisma_client.tool_usage_transactions = prisma_client.tool_usage_transactions[len(tool_usage_to_process) :] 

7617 try: 

7618 from litellm.proxy.db.spend_log_tool_index import flush_tool_usage_transactions 

7619 

7620 await flush_tool_usage_transactions( 

7621 prisma_client=prisma_client, 

7622 transactions=tool_usage_to_process, 

7623 ) 

7624 except Exception as tool_tracking_err: 

7625 verbose_proxy_logger.error( 

7626 "Spend tracking - tool usage flush failed; %s tool usage transactions dropped: %s", 

7627 len(tool_usage_to_process), 

7628 tool_tracking_err, 

7629 ) 

7630 

7631 await flush_baseline_accounting(prisma_client) 

7632 

7633 async with prisma_client._autorouter_turn_transactions_lock: 

7634 autorouter_turns_to_process: Final = prisma_client.autorouter_turn_transactions[:MAX_LOGS_PER_INTERVAL] 

7635 remaining_autorouter_turns: Final = prisma_client.autorouter_turn_transactions[ 

7636 len(autorouter_turns_to_process) : 

7637 ] 

7638 prisma_client.autorouter_turn_transactions = remaining_autorouter_turns # rebind-ok: drain under lock 

7639 try: 

7640 from litellm.proxy.db.autorouter_session_rollup import flush_autorouter_turn_transactions 

7641 

7642 await flush_autorouter_turn_transactions( 

7643 prisma_client=prisma_client, 

7644 transactions=autorouter_turns_to_process, 

7645 ) 

7646 except Exception as autorouter_tracking_err: # noqa: BLE001 # a drain bug must not abort the spend job 

7647 verbose_proxy_logger.error( 

7648 "Spend tracking - auto-router session rollup drain failed; %s turn transactions dropped: %s", 

7649 len(autorouter_turns_to_process), 

7650 autorouter_tracking_err, 

7651 ) 

7652 

7653 try: 

7654 from litellm.proxy.db.shadow_eval_funnel import flush_shadow_eval_funnel 

7655 

7656 await flush_shadow_eval_funnel(prisma_client) 

7657 except Exception as funnel_err: # noqa: BLE001 # a drain bug must not abort the spend job 

7658 verbose_proxy_logger.error("Spend tracking - shadow eval funnel drain failed: %s", funnel_err) 

7659 

7660 

7661MAX_SPEND_LOG_DRAIN_ITERATIONS: Final = 20 

7662 

7663 

7664async def drain_spend_logs_queue( 

7665 prisma_client: PrismaClient, 

7666 db_writer_client: "AsyncHTTPHandler | None", 

7667 proxy_logging_obj: ProxyLogging, 

7668) -> None: 

7669 monitor_task: Final = prisma_client.spend_logs_queue_monitor_task 

7670 if monitor_task is not None: 7670 ↛ 7676line 7670 didn't jump to line 7676 because the condition on line 7670 was always true

7671 monitor_task.cancel() 

7672 with contextlib.suppress(asyncio.CancelledError): 

7673 await monitor_task 

7674 prisma_client.spend_logs_queue_monitor_task = None # rebind-ok: the client owns its monitor handle 

7675 

7676 async with prisma_client.spend_log_write_lock: 

7677 try: 

7678 await _drain_spend_logs_queue_to_db(prisma_client, db_writer_client, proxy_logging_obj) 

7679 finally: 

7680 await _park_remaining_spend_logs(prisma_client, proxy_logging_obj) 

7681 

7682 

7683async def _drain_spend_logs_queue_to_db( 

7684 prisma_client: PrismaClient, 

7685 db_writer_client: "AsyncHTTPHandler | None", 

7686 proxy_logging_obj: ProxyLogging, 

7687) -> None: 

7688 for _ in range(MAX_SPEND_LOG_DRAIN_ITERATIONS): 7688 ↛ 7693line 7688 didn't jump to line 7693 because the loop on line 7688 didn't complete

7689 if await _total_queued_spend_transactions(prisma_client) == 0: 7689 ↛ 7691line 7689 didn't jump to line 7691 because the condition on line 7689 was always true

7690 return 

7691 await _run_spend_logs_job(prisma_client, db_writer_client, proxy_logging_obj) 

7692 

7693 remaining: Final = await _total_queued_spend_transactions(prisma_client) 

7694 if remaining > 0: 

7695 spend_log_error( 

7696 "Spend tracking - %d spend log rows still queued after %d drain passes", 

7697 remaining, 

7698 MAX_SPEND_LOG_DRAIN_ITERATIONS, 

7699 ) 

7700 

7701 

7702async def _park_remaining_spend_logs(prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging) -> None: 

7703 rows: Final = await dequeue_spend_logs(prisma_client, sys.maxsize) 

7704 if len(rows) == 0 or await _park_spend_logs_in_redis(proxy_logging_obj, rows): 7704 ↛ 7706line 7704 didn't jump to line 7706 because the condition on line 7704 was always true

7705 return 

7706 await enqueue_spend_logs(prisma_client, rows, at_head=True) 

7707 spend_log_error( 

7708 "Spend tracking - %d spend log rows could not be written or parked in Redis and will be lost on exit", 

7709 len(rows), 

7710 ) 

7711 

7712 

7713async def _monitor_spend_logs_queue( 

7714 prisma_client: PrismaClient, 

7715 db_writer_client: AsyncHTTPHandler | None, 

7716 proxy_logging_obj: ProxyLogging, 

7717): 

7718 """ 

7719 Background task that monitors the spend_log_transactions queue size 

7720 and triggers processing when the threshold is reached. 

7721 

7722 Args: 

7723 prisma_client: Prisma client instance 

7724 db_writer_client: Optional HTTP handler for external spend logs endpoint 

7725 proxy_logging_obj: Proxy logging object 

7726 """ 

7727 from litellm.constants import ( 

7728 SPEND_LOG_QUEUE_POLL_INTERVAL, 

7729 SPEND_LOG_QUEUE_SIZE_THRESHOLD, 

7730 ) 

7731 

7732 threshold: Final = SPEND_LOG_QUEUE_SIZE_THRESHOLD 

7733 base_interval: Final = SPEND_LOG_QUEUE_POLL_INTERVAL 

7734 max_backoff: Final = 30.0 # Maximum backoff interval in seconds 

7735 backoff_multiplier: Final = 1.5 # Exponential backoff multiplier 

7736 current_interval = base_interval 

7737 flush_requested: Final = asyncio.Event() 

7738 prisma_client.spend_log_flush_requested = flush_requested # rebind-ok: the client owns its monitor's flush signal 

7739 

7740 verbose_proxy_logger.info( 

7741 "Starting spend logs queue monitor (threshold: %s, poll_interval: %ss)", threshold, base_interval 

7742 ) 

7743 

7744 while True: 

7745 try: 

7746 await recover_parked_spend_logs(prisma_client, proxy_logging_obj) 

7747 # Check queue sizes with lock protection; the tool usage queue keeps 

7748 # the monitor firing when a prior failed run left it nonempty. 

7749 queue_size = await _total_queued_spend_transactions(prisma_client) 

7750 

7751 if queue_size > 0: 

7752 if queue_size >= threshold: 

7753 verbose_proxy_logger.debug( 

7754 "Spend logs queue size (%s) reached threshold (%s), triggering processing", 

7755 queue_size, 

7756 threshold, 

7757 ) 

7758 # Reset to base interval when threshold is reached 

7759 current_interval = base_interval 

7760 else: 

7761 verbose_proxy_logger.debug( 

7762 "Spend logs queue size (%s) below threshold (%s), processing with backoff", 

7763 queue_size, 

7764 threshold, 

7765 ) 

7766 # Exponential backoff when below threshold but still processing 

7767 current_interval = min(current_interval * backoff_multiplier, max_backoff) 

7768 

7769 await update_spend_logs_job( 

7770 prisma_client=prisma_client, 

7771 db_writer_client=db_writer_client, 

7772 proxy_logging_obj=proxy_logging_obj, 

7773 ) 

7774 else: 

7775 from litellm.proxy.db.baseline_accounting import flush_baseline_accounting 

7776 

7777 await flush_baseline_accounting(prisma_client) 

7778 current_interval = min(current_interval * backoff_multiplier, max_backoff) 

7779 

7780 if await _wait_for_spend_log_flush_request(flush_requested, current_interval): 

7781 current_interval = base_interval 

7782 except Exception as e: 

7783 spend_log_error("Error in spend logs queue monitor: %s", str(e), exc=e) 

7784 # Continue monitoring even if there's an error, with exponential backoff 

7785 current_interval = min(current_interval * backoff_multiplier, max_backoff) 

7786 await asyncio.sleep(current_interval) 

7787 

7788 

7789MAX_SPEND_LOG_ISOLATION_FAILURES_PER_BATCH: Final = 256 

7790 

7791 

7792def _is_transient_spend_log_write_error(e: Exception) -> bool: 

7793 return PrismaDBExceptionHandler.is_database_transport_error(e) or PrismaDBExceptionHandler.is_deadlock_error(e) 

7794 

7795 

7796async def _create_spend_logs_with_poison_isolation( 

7797 repo: SpendLogsRepository, 

7798 rows: Sequence[Mapping[str, object]], 

7799 failure_budget: int, 

7800) -> int: 

7801 """Write spend-log rows, isolating any row Postgres rejects on its data. 

7802 

7803 ``create_many`` writes the whole batch in a single statement, so one row 

7804 carrying bytes Postgres refuses (a residual NUL byte is the canonical case) 

7805 fails the entire insert and drops every good row alongside it. On a genuine 

7806 data-layer rejection the batch is bisected so the good rows still persist 

7807 and only the offending row is dropped and logged. Transport failures, 

7808 including the "can't reach database server" outage that prisma mislabels as 

7809 a ``DataError``, are re-raised unchanged so the caller's connection-retry 

7810 path still runs. 

7811 

7812 ``failure_budget`` caps the *failed* inserts the isolation may issue, which 

7813 is the work an authenticated caller flooding poisoned rows can amplify. The 

7814 one insert a statement needs when nothing is poisoned is not charged, so a 

7815 caller can thread a single budget through every statement of a flush and 

7816 bound the whole flush's failed inserts and log lines by the initial value, 

7817 without a large healthy flush ever running out and losing rows. When the 

7818 budget is spent the still-failing remainder is dropped wholesale (the 

7819 pre-existing drop-the-batch behavior) under one log line, and a statement 

7820 reached afterwards is still attempted, so clean rows behind a poison flood 

7821 persist. Returns the budget left after this subtree. 

7822 """ 

7823 try: 

7824 await repo.table.create_many(data=rows, skip_duplicates=True) 

7825 return failure_budget 

7826 except Exception as e: 

7827 if not PrismaDBExceptionHandler.is_prisma_data_error(e): 7827 ↛ 7828line 7827 didn't jump to line 7828 because the condition on line 7827 was never true

7828 raise 

7829 if PrismaDBExceptionHandler.is_database_service_unavailable_error(e): 7829 ↛ 7830line 7829 didn't jump to line 7830 because the condition on line 7829 was never true

7830 raise 

7831 if PrismaDBExceptionHandler.is_deadlock_error(e): 7831 ↛ 7832line 7831 didn't jump to line 7832 because the condition on line 7831 was never true

7832 raise 

7833 budget_left: Final = max(failure_budget - 1, 0) 

7834 if len(rows) == 1: 

7835 request_id: Final = rows[0].get("request_id") 

7836 spend_log_error( 

7837 "Spend tracking - dropping spend log row Postgres rejected. request_id=%s error=%s", 

7838 request_id, 

7839 str(e), 

7840 exc=e, 

7841 ) 

7842 return budget_left 

7843 if budget_left <= 0: 7843 ↛ 7844line 7843 didn't jump to line 7844 because the condition on line 7843 was never true

7844 spend_log_error( 

7845 "Spend tracking - dropping %d spend log rows without per-row isolation; " 

7846 "isolation failure budget exhausted for this flush", 

7847 len(rows), 

7848 ) 

7849 return 0 

7850 mid: Final = len(rows) // 2 

7851 remaining: Final = await _create_spend_logs_with_poison_isolation(repo, rows[:mid], budget_left) 

7852 if remaining <= 0: 7852 ↛ 7853line 7852 didn't jump to line 7853 because the condition on line 7852 was never true

7853 spend_log_error( 

7854 "Spend tracking - dropping %d spend log rows without per-row isolation; " 

7855 "isolation failure budget exhausted for this flush", 

7856 len(rows) - mid, 

7857 ) 

7858 return 0 

7859 return await _create_spend_logs_with_poison_isolation(repo, rows[mid:], remaining) 

7860 

7861 

7862def _raise_failed_update_spend_exception(e: Exception, start_time: float, proxy_logging_obj: ProxyLogging): 

7863 """ 

7864 Raise an exception for failed update spend logs 

7865 

7866 - Calls proxy_logging_obj.failure_handler to log the error 

7867 - Ensures error messages says "Non-Blocking" 

7868 """ 

7869 import traceback 

7870 

7871 error_msg: Final = f"[Non-Blocking]LiteLLM Prisma Client Exception - update spend logs: {e}" 

7872 error_traceback: Final = error_msg + "\n" + traceback.format_exc() 

7873 end_time: Final = time.time() 

7874 _duration: Final = end_time - start_time 

7875 asyncio.create_task( 

7876 proxy_logging_obj.failure_handler( 

7877 original_exception=e, 

7878 duration=_duration, 

7879 call_type="update_spend", 

7880 traceback_str=error_traceback, 

7881 ) 

7882 ) 

7883 raise e 

7884 

7885 

7886def _get_month_end_date(today: date) -> date: 

7887 if today.month == 12: 

7888 return date(today.year + 1, 1, 1) - timedelta(days=1) 

7889 return date(today.year, today.month + 1, 1) - timedelta(days=1) 

7890 

7891 

7892def _is_projected_spend_over_limit(current_spend: float, soft_budget_limit: float | None): 

7893 if soft_budget_limit is None: 

7894 # If there's no limit, we can't exceed it. 

7895 return False 

7896 

7897 today: Final = date.today() 

7898 

7899 # Finding the first day of the next month, then subtracting one day to get the end of the current month. 

7900 end_month: Final = _get_month_end_date(today) 

7901 

7902 remaining_days: Final = (end_month - today).days 

7903 

7904 # Check for the start of the month to avoid division by zero 

7905 if today.day == 1: 

7906 daily_spend_estimate = current_spend 

7907 else: 

7908 daily_spend_estimate = current_spend / (today.day - 1) 

7909 

7910 # Total projected spend for the month 

7911 projected_spend: Final = current_spend + (daily_spend_estimate * remaining_days) 

7912 

7913 if projected_spend > soft_budget_limit: 

7914 print_verbose("Projected spend exceeds soft budget limit!") 

7915 return True 

7916 return False 

7917 

7918 

7919def _get_projected_spend_over_limit(current_spend: float, soft_budget_limit: float | None) -> tuple | None: 

7920 if soft_budget_limit is None: 

7921 return None 

7922 

7923 today: Final = date.today() 

7924 end_month: Final = _get_month_end_date(today) 

7925 remaining_days: Final = (end_month - today).days 

7926 

7927 # assuming the current spend till today (not including today) 

7928 if today.day == 1: 

7929 daily_spend = current_spend 

7930 else: 

7931 daily_spend = current_spend / (today.day - 1) 

7932 projected_spend: Final = current_spend + (daily_spend * remaining_days) 

7933 

7934 if projected_spend > soft_budget_limit: 

7935 if daily_spend <= 0: 

7936 limit_exceed_date = today 

7937 else: 

7938 remaining_budget: Final = soft_budget_limit - current_spend 

7939 if remaining_budget <= 0: 

7940 limit_exceed_date = today 

7941 else: 

7942 approx_days: Final = remaining_budget / daily_spend 

7943 limit_exceed_date = today + timedelta(days=approx_days) 

7944 

7945 # return the projected spend and the date it will exceeded 

7946 return projected_spend, limit_exceed_date 

7947 

7948 return None 

7949 

7950 

7951def _is_valid_team_configs(team_id=None, team_config=None, request_data=None): 

7952 if team_id is None or team_config is None or request_data is None: 

7953 return 

7954 # check if valid model called for team 

7955 if "models" in team_config: 

7956 valid_models: Final = team_config.pop("models") 

7957 model_in_request: Final = request_data["model"] 

7958 if model_in_request not in valid_models: 

7959 raise Exception( 

7960 f"Invalid model for team {team_id}: {model_in_request}. Valid models for team are: {valid_models}\n" 

7961 ) 

7962 return 

7963 

7964 

7965def _to_ns(dt): 

7966 return int(dt.timestamp() * 1e9) 

7967 

7968 

7969def _check_and_merge_model_level_guardrails( 

7970 data: dict, 

7971 llm_router: Router | None, 

7972 trust_client_model_info: bool = True, 

7973 model_alias: str | None = None, 

7974) -> dict: 

7975 """ 

7976 Check if the model has guardrails defined and merge them with existing guardrails in the request data. 

7977 

7978 Args: 

7979 data: The request data dict 

7980 llm_router: The LLM router instance to get deployment info from 

7981 model_alias: Resolve guardrails for this model group instead of data["model"] 

7982 trust_client_model_info: If False, ignore metadata.model_info.id and 

7983 resolve guardrails by alias-union only. Set to False on the 

7984 pre_call path because add_litellm_data_to_request preserves 

7985 client-supplied model_info when allow_client_pricing_override is 

7986 set, so a caller could spoof an unguarded model_info.id while 

7987 requesting a guarded alias and bypass guardrails (veria-ai HIGH 

7988 on #29654). Defaults to True for post_call paths where the 

7989 router has populated model_info.id itself. 

7990 

7991 Returns: 

7992 Modified data dict with merged guardrails (if any model-level guardrails exist) 

7993 """ 

7994 if llm_router is None: 

7995 return data 

7996 

7997 metadata: Final = data.get("metadata") or {} 

7998 litellm_metadata: Final = data.get("litellm_metadata") or {} 

7999 model_info: Final = metadata.get("model_info") or {} 

8000 model_id: Final = model_info.get("id") if trust_client_model_info else None 

8001 # route_request resolves team-scoped public model names with the 

8002 # server-populated team id; pre_call lookup must do the same so 

8003 # team-scoped guardrails are not silently skipped (greptile/veria-ai 

8004 # Medium on #29654). 

8005 team_id: Final = metadata.get("user_api_key_team_id") or litellm_metadata.get("user_api_key_team_id") 

8006 

8007 model_level_guardrails: list[object] | None = None 

8008 if model_id is not None: 8008 ↛ 8009line 8008 didn't jump to line 8009 because the condition on line 8008 was never true

8009 deployment: Final = llm_router.get_deployment(model_id=model_id) 

8010 if deployment is None: 

8011 return data 

8012 deployment_guardrails: Final = deployment.litellm_params.get("guardrails") 

8013 # Bare-string guardrail names were truthy-accepted before; preserve 

8014 # that contract so post_call callers don't silently lose them. 

8015 if isinstance(deployment_guardrails, list): 

8016 model_level_guardrails = deployment_guardrails 

8017 elif deployment_guardrails: 

8018 model_level_guardrails = [deployment_guardrails] 

8019 else: 

8020 # Pre_call paths run before route_request picks a deployment, so we 

8021 # don't know which deployment's litellm_params.guardrails will apply. 

8022 # Take the UNION across all deployments in the group so a guardrail 

8023 # set on ANY eligible deployment still fires (#29652; addresses 

8024 # veria-ai HIGH on the single-deployment fallback that would skip 

8025 # non-first deployments). 

8026 alias: Final = model_alias if model_alias is not None else data.get("model") 

8027 if not isinstance(alias, str) or not alias: 

8028 return data 

8029 # Pass team_id so team-scoped public model names resolve the same way 

8030 # route_request resolves them; otherwise team-scoped deployments are 

8031 # invisible to this lookup and their guardrails are silently dropped. 

8032 deployments: Final = llm_router.get_model_list(model_name=alias, team_id=team_id) or [] 

8033 seen: Final[set] = set() 

8034 union: Final[list] = [] 

8035 for dep in deployments: 8035 ↛ 8036line 8035 didn't jump to line 8036 because the loop on line 8035 never started

8036 litellm_params_dep = dep.get("litellm_params") or {} 

8037 guardrails = litellm_params_dep.get("guardrails") 

8038 if isinstance(guardrails, str): 

8039 guardrails = [guardrails] 

8040 elif not isinstance(guardrails, list): 

8041 continue 

8042 for g in guardrails: 

8043 key = g if isinstance(g, str) else repr(g) 

8044 if key not in seen: 

8045 seen.add(key) 

8046 union.append(g) 

8047 model_level_guardrails = union or None 

8048 

8049 if model_level_guardrails is None: 8049 ↛ 8053line 8049 didn't jump to line 8053 because the condition on line 8049 was always true

8050 return data 

8051 

8052 # Merge model-level guardrails with existing ones 

8053 return _merge_guardrails_with_existing(data, model_level_guardrails) 

8054 

8055 

8056def _merge_guardrails_with_existing(data: dict, model_level_guardrails: object) -> dict: 

8057 """ 

8058 Merge model-level guardrails with any existing guardrails in the request data. 

8059 

8060 Args: 

8061 data: The request data dict 

8062 model_level_guardrails: Guardrails defined at the model level 

8063 

8064 Returns: 

8065 Modified data dict with merged guardrails in metadata 

8066 """ 

8067 modified_data: Final = data.copy() 

8068 metadata: Final = modified_data.setdefault("metadata", {}) 

8069 existing_guardrails = metadata.get("guardrails", []) 

8070 

8071 # Ensure existing_guardrails is a list 

8072 if not isinstance(existing_guardrails, list): 

8073 existing_guardrails = [existing_guardrails] if existing_guardrails else [] 

8074 

8075 # Ensure model_level_guardrails is a list 

8076 if not isinstance(model_level_guardrails, list): 

8077 model_level_guardrails = [model_level_guardrails] if model_level_guardrails else [] 

8078 

8079 # Combine existing and model-level guardrails 

8080 metadata["guardrails"] = list(set(existing_guardrails + model_level_guardrails)) 

8081 return modified_data 

8082 

8083 

8084def get_error_message_str(e: Exception) -> str: 

8085 error_message = "" 

8086 if isinstance(e, HTTPException): 

8087 if isinstance(e.detail, str): 

8088 error_message = e.detail 

8089 elif isinstance(e.detail, dict): 

8090 error_message = json.dumps(e.detail) 

8091 elif hasattr(e, "message"): 

8092 _error: Final = getattr(e, "message", None) 

8093 if isinstance(_error, str): 

8094 error_message = _error 

8095 elif isinstance(_error, dict): 

8096 error_message = json.dumps(_error) 

8097 else: 

8098 error_message = str(e) 

8099 else: 

8100 error_message = str(e) 

8101 return error_message 

8102 

8103 

8104def _get_redoc_url() -> str | None: 

8105 """ 

8106 Get the Redoc URL from the environment variables. 

8107 

8108 - If REDOC_URL is set, return it. 

8109 - If NO_REDOC is True, return None. 

8110 - Otherwise, default to "/redoc". 

8111 """ 

8112 if redoc_url := os.getenv("REDOC_URL"): 8112 ↛ 8113line 8112 didn't jump to line 8113 because the condition on line 8112 was never true

8113 return redoc_url 

8114 

8115 if str_to_bool(os.getenv("NO_REDOC")) is True: 8115 ↛ 8116line 8115 didn't jump to line 8116 because the condition on line 8115 was never true

8116 return None 

8117 

8118 return "/redoc" 

8119 

8120 

8121def _get_docs_url() -> str | None: 

8122 """ 

8123 Get the docs (Swagger UI) URL from the environment variables. 

8124 

8125 - If DOCS_URL is set, return it. 

8126 - If NO_DOCS is True, return None. 

8127 - Otherwise, default to "/". 

8128 """ 

8129 if docs_url := os.getenv("DOCS_URL"): 8129 ↛ 8130line 8129 didn't jump to line 8130 because the condition on line 8129 was never true

8130 return docs_url 

8131 

8132 if str_to_bool(os.getenv("NO_DOCS")) is True: 8132 ↛ 8133line 8132 didn't jump to line 8133 because the condition on line 8132 was never true

8133 return None 

8134 

8135 return "/" 

8136 

8137 

8138def _get_openapi_url() -> str | None: 

8139 """ 

8140 Get the OpenAPI JSON URL from the environment variables. 

8141 

8142 - If OPENAPI_URL is set, return it. 

8143 - If NO_OPENAPI is True, return None. 

8144 - Otherwise, default to "/openapi.json". 

8145 """ 

8146 if openapi_url := os.getenv("OPENAPI_URL"): 8146 ↛ 8147line 8146 didn't jump to line 8147 because the condition on line 8146 was never true

8147 return openapi_url 

8148 

8149 if str_to_bool(os.getenv("NO_OPENAPI")) is True: 8149 ↛ 8150line 8149 didn't jump to line 8150 because the condition on line 8149 was never true

8150 return None 

8151 

8152 return "/openapi.json" 

8153 

8154 

8155def _recreate_writer_on_read_only_transaction(prisma_client: "PrismaClient | None") -> None: 

8156 if prisma_client is None: 

8157 return 

8158 asyncio.create_task(prisma_client.recreate_read_only_writer(reason="postgres_read_only_transaction")) 

8159 

8160 

8161def handle_exception_on_proxy(e: Exception, litellm_call_id: str | None = None) -> ProxyException: 

8162 """ 

8163 Returns an Exception as ProxyException, this ensures all exceptions are OpenAI API compatible 

8164 """ 

8165 from fastapi import status 

8166 

8167 verbose_proxy_logger.exception("Exception: %s", e) 

8168 if PrismaDBExceptionHandler.is_read_only_transaction_error(e): 8168 ↛ 8169line 8168 didn't jump to line 8169 because the condition on line 8168 was never true

8169 from litellm.proxy.proxy_server import prisma_client 

8170 

8171 _recreate_writer_on_read_only_transaction(prisma_client) 

8172 

8173 headers: Final = litellm_call_id_headers(litellm_call_id) 

8174 if isinstance(e, HTTPException): 

8175 return ProxyException( 

8176 message=getattr(e, "detail", f"error({e})"), 

8177 type=ProxyErrorTypes.internal_server_error, 

8178 param=openai_error_param(e), 

8179 headers=headers, 

8180 code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR), 

8181 ) 

8182 elif isinstance(e, ProxyException): 

8183 return with_litellm_call_id(e, litellm_call_id) 

8184 _status_code: Final = getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR) 

8185 if should_report_bug(e): 

8186 verbose_proxy_logger.error(bug_report_notice(build_proxy_bug_report(e))) 

8187 return ProxyException( 

8188 message=strip_bug_report_notice(str(e)), 

8189 type=ProxyErrorTypes.internal_server_error, 

8190 param=openai_error_param(e), 

8191 headers=headers, 

8192 code=_status_code, 

8193 ) 

8194 

8195 

8196def _premium_user_check(feature: str | None = None): 

8197 """ 

8198 Raises an HTTPException if the user is not a premium user 

8199 """ 

8200 from litellm.proxy.proxy_server import premium_user 

8201 

8202 if feature: 

8203 detail_msg = f"This feature is only available for LiteLLM Enterprise users: {feature}. {CommonProxyErrors.not_premium_user.value}" 

8204 else: 

8205 detail_msg = ( 

8206 f"This feature is only available for LiteLLM Enterprise users. {CommonProxyErrors.not_premium_user.value}" 

8207 ) 

8208 

8209 if not premium_user: 8209 ↛ exitline 8209 didn't return from function '_premium_user_check' because the condition on line 8209 was always true

8210 raise HTTPException( 

8211 status_code=403, 

8212 detail={"error": detail_msg}, 

8213 ) 

8214 

8215 

8216def is_known_model(model: str | None, llm_router: Router | None) -> bool: 

8217 """ 

8218 Returns True if the model is in the llm_router model names 

8219 """ 

8220 if model is None or llm_router is None: 8220 ↛ 8222line 8220 didn't jump to line 8222 because the condition on line 8220 was always true

8221 return False 

8222 model_names: Final = llm_router.get_model_names() 

8223 

8224 model_names_set: Final = set(model_names) 

8225 

8226 is_in_list = False 

8227 if model in model_names_set: 

8228 is_in_list = True 

8229 

8230 return is_in_list 

8231 

8232 

8233def is_known_vector_store_index(index_name: str) -> bool: 

8234 """ 

8235 Returns True if the vector store index is in the llm_router vector store indexes 

8236 """ 

8237 

8238 if litellm.vector_store_index_registry is None: 

8239 return False 

8240 return index_name in litellm.vector_store_index_registry.get_vector_store_indexes() 

8241 

8242 

8243def join_paths(base_path: str, route: str) -> str: 

8244 # Remove trailing slashes from base_path and leading slashes from route 

8245 base_path = base_path.rstrip("/") 

8246 route = route.lstrip("/") 

8247 

8248 # If base_path is empty, return route with leading slash 

8249 if not base_path: 

8250 return f"/{route}" if route else "/" 

8251 

8252 # If route is empty, return just base_path 

8253 if not route: 8253 ↛ 8257line 8253 didn't jump to line 8257 because the condition on line 8253 was always true

8254 return base_path 

8255 

8256 # Check if base_path already ends with the route to avoid duplication 

8257 if base_path.endswith(f"/{route}"): 

8258 final_path = base_path 

8259 else: 

8260 # Join with single slash 

8261 final_path = f"{base_path}/{route}" 

8262 

8263 return final_path 

8264 

8265 

8266def get_custom_url(request_base_url: str, route: str | None = None) -> str: 

8267 # Use environment variable value, otherwise use URL from request 

8268 server_base_url: Final = get_proxy_base_url() 

8269 if server_base_url is not None: 8269 ↛ 8270line 8269 didn't jump to line 8270 because the condition on line 8269 was never true

8270 base_url = server_base_url 

8271 else: 

8272 base_url = request_base_url 

8273 

8274 # get_request_root_path() returns the prefix the router is actually 

8275 # resolving this request under: the matched SERVER_ROOT_PATHS entry when 

8276 # PerRequestRootPathMiddleware ran, otherwise the SERVER_ROOT_PATH scalar. 

8277 # This keeps the emitted URL under one prefix — the one the client called — 

8278 # instead of stacking the scalar onto a request already living under a 

8279 # dynamic prefix (which would produce /tenant-a/legacy/... — a path that 

8280 # doesn't exist). join_paths()'s tail-dedup then collapses the append when 

8281 # base_url (i.e. request.base_url) already ends in the same prefix. 

8282 from litellm.proxy.middleware.per_request_root_path_middleware import ( # noqa: PLC0415 # lazy: middleware imports utils 

8283 get_request_root_path, 

8284 ) 

8285 

8286 server_root_path: Final = get_request_root_path() 

8287 if route is not None: 

8288 if server_root_path != "": 8288 ↛ 8290line 8288 didn't jump to line 8290 because the condition on line 8288 was never true

8289 # First join base_url with server_root_path, then with route 

8290 intermediate_url: Final = join_paths(base_url, server_root_path) 

8291 return join_paths(intermediate_url, route) 

8292 else: 

8293 return join_paths(base_url, route) 

8294 else: 

8295 return join_paths(base_url, server_root_path) 

8296 

8297 

8298def get_proxy_base_url() -> str | None: 

8299 """ 

8300 Get the proxy base url from the environment variables. 

8301 """ 

8302 return os.getenv("PROXY_BASE_URL") 

8303 

8304 

8305def get_server_root_path() -> str: 

8306 """ 

8307 Get the server root path from the environment variables. 

8308 

8309 - If SERVER_ROOT_PATH is set, return it. 

8310 - Otherwise, default to "/". 

8311 """ 

8312 return os.getenv("SERVER_ROOT_PATH", "") 

8313 

8314 

8315def normalize_route_for_root_path(route: str) -> str | None: 

8316 """Strip SERVER_ROOT_PATH prefix. Returns de-prefixed route, or None if route is not under root path.""" 

8317 root_path: Final = get_server_root_path() 

8318 if root_path and root_path != "/": 8318 ↛ 8319line 8318 didn't jump to line 8319 because the condition on line 8318 was never true

8319 if route.startswith(root_path + "/"): 

8320 return route[len(root_path) :] 

8321 return None 

8322 return route 

8323 

8324 

8325def get_prisma_client_or_throw(message: str): 

8326 from litellm.proxy.proxy_server import prisma_client 

8327 

8328 if prisma_client is None: 8328 ↛ 8329line 8328 didn't jump to line 8329 because the condition on line 8328 was never true

8329 raise HTTPException( 

8330 status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, 

8331 detail={"error": message}, 

8332 ) 

8333 return prisma_client 

8334 

8335 

8336def is_valid_api_key(key: str) -> bool: 

8337 """ 

8338 Validates API key format: 

8339 - sk- keys: must match ^sk-[A-Za-z0-9_-]+$ 

8340 - hashed keys: must match ^[a-fA-F0-9]{64}$ 

8341 - Length between 20 and 100 characters 

8342 """ 

8343 import re 

8344 

8345 if not isinstance(key, str): 8345 ↛ 8346line 8345 didn't jump to line 8346 because the condition on line 8345 was never true

8346 return False 

8347 if 3 <= len(key) <= 100: 

8348 if re.match(r"^sk-[A-Za-z0-9_-]+$", key): 

8349 return True 

8350 if re.match(r"^[a-fA-F0-9]{64}$", key): 8350 ↛ 8351line 8350 didn't jump to line 8351 because the condition on line 8350 was never true

8351 return True 

8352 return False 

8353 

8354 

8355def construct_database_url_from_env_vars() -> str | None: 

8356 """ 

8357 Construct a DATABASE_URL from individual environment variables. 

8358 Returns: 

8359 Optional[str]: The constructed DATABASE_URL or None if required variables are missing 

8360 """ 

8361 import urllib.parse 

8362 

8363 # Check if all required variables are provided 

8364 database_host: Final = os.getenv("DATABASE_HOST") 

8365 database_username: Final = os.getenv("DATABASE_USERNAME") 

8366 database_password: Final = os.getenv("DATABASE_PASSWORD") 

8367 database_name: Final = os.getenv("DATABASE_NAME") 

8368 database_schema: Final = os.getenv("DATABASE_SCHEMA") 

8369 

8370 if database_host and database_username and database_name: 

8371 # Handle the problem of special character escaping in the database URL 

8372 database_username_enc: Final = urllib.parse.quote_plus(database_username) 

8373 database_password_enc: Final = urllib.parse.quote_plus(database_password) if database_password else "" 

8374 database_name_enc: Final = urllib.parse.quote_plus(database_name) 

8375 

8376 # Construct DATABASE_URL from the provided variables 

8377 if database_password: 

8378 database_url = ( 

8379 f"postgresql://{database_username_enc}:{database_password_enc}@{database_host}/{database_name_enc}" 

8380 ) 

8381 else: 

8382 database_url = f"postgresql://{database_username_enc}@{database_host}/{database_name_enc}" 

8383 

8384 if database_schema: 

8385 database_url += f"?schema={database_schema}" 

8386 

8387 return add_missing_query_params(database_url, DatabaseURLSettings.from_env().tls_params()) 

8388 

8389 return None 

8390 

8391 

8392async def _get_validated_team_object( 

8393 user_api_key_dict: "UserAPIKeyAuth", 

8394 team_id: str, 

8395 prisma_client: "PrismaClient", 

8396 user_api_key_cache: "UserApiKeyCache", 

8397 proxy_logging_obj: "ProxyLogging", 

8398) -> "LiteLLM_TeamTableCachedObj": 

8399 from litellm.proxy.auth.auth_checks import get_team_object 

8400 from litellm.proxy.management_endpoints.team_endpoints import validate_membership 

8401 

8402 team_object: Final = await get_team_object( 

8403 team_id=team_id, 

8404 prisma_client=prisma_client, 

8405 user_api_key_cache=user_api_key_cache, 

8406 proxy_logging_obj=proxy_logging_obj, 

8407 ) 

8408 await validate_membership(user_api_key_dict=user_api_key_dict, team_table=team_object) 

8409 return team_object 

8410 

8411 

8412async def _get_team_object_for_access_groups( 

8413 team_id: str | None, 

8414 prisma_client: Optional["PrismaClient"], 

8415 user_api_key_cache: Optional["UserApiKeyCache"], 

8416 proxy_logging_obj: Optional["ProxyLogging"], 

8417) -> Optional["LiteLLM_TeamTableCachedObj"]: 

8418 from litellm.proxy.auth.auth_checks import get_team_object 

8419 

8420 if team_id is None or prisma_client is None or user_api_key_cache is None or proxy_logging_obj is None: 

8421 return None 

8422 try: 

8423 return await get_team_object( 

8424 team_id=team_id, 

8425 prisma_client=prisma_client, 

8426 user_api_key_cache=user_api_key_cache, 

8427 proxy_logging_obj=proxy_logging_obj, 

8428 ) 

8429 except HTTPException: 

8430 verbose_proxy_logger.debug("Could not fetch team %s while listing models", team_id) 

8431 return None 

8432 

8433 

8434async def _get_access_group_models( 

8435 user_api_key_dict: "UserAPIKeyAuth", 

8436 team_object: Optional["LiteLLM_TeamTableCachedObj"], 

8437 prisma_client: Optional["PrismaClient"], 

8438 user_api_key_cache: Optional["UserApiKeyCache"], 

8439 proxy_logging_obj: Optional["ProxyLogging"], 

8440) -> tuple[str, ...]: 

8441 from litellm.proxy.auth.auth_checks import ( 

8442 _get_models_from_access_groups, 

8443 get_authorized_resources_from_key_access_groups, 

8444 ) 

8445 

8446 team_group_models: Final = await _get_models_from_access_groups( 

8447 access_group_ids=(team_object.access_group_ids or ()) if team_object is not None else (), 

8448 prisma_client=prisma_client, 

8449 user_api_key_cache=user_api_key_cache, 

8450 proxy_logging_obj=proxy_logging_obj, 

8451 ) 

8452 key_group_models: Final = await get_authorized_resources_from_key_access_groups( 

8453 valid_token=user_api_key_dict, 

8454 team_object=team_object, 

8455 resource_field="access_model_names", 

8456 ) 

8457 return tuple(dict.fromkeys((*team_group_models, *key_group_models))) 

8458 

8459 

8460async def _agent_access_group_visible_models( 

8461 user_api_key_dict: "UserAPIKeyAuth", 

8462 llm_router: "Router | None", 

8463 include_model_access_groups: bool, 

8464 return_wildcard_routes: bool, 

8465 team_id: str | None, 

8466 resolve_agent_ceiling: CeilingResolver, 

8467) -> frozenset[str] | None: 

8468 """Models an agent key may still list once its attached access groups cap it, ``None`` when 

8469 nothing caps it, so ``/v1/models`` never advertises a model the same key would be denied on.""" 

8470 from litellm.proxy.auth.model_checks import get_complete_model_list, get_team_models 

8471 

8472 if not user_api_key_dict.agent_id: 8472 ↛ 8474line 8472 didn't jump to line 8474 because the condition on line 8472 was always true

8473 return None 

8474 ceiling: Final = await resolve_agent_ceiling(user_api_key_dict.agent_id) 

8475 if ceiling is None: 

8476 return None 

8477 if llm_router is None: 

8478 return ceiling.models 

8479 proxy_model_list: Final = llm_router.get_model_names() 

8480 model_access_groups: Final = llm_router.get_model_access_groups() 

8481 granted: Final = get_team_models( 

8482 team_models=sorted(ceiling.models), 

8483 proxy_model_list=proxy_model_list, 

8484 model_access_groups=model_access_groups, 

8485 include_model_access_groups=include_model_access_groups, 

8486 ) 

8487 if not granted: 

8488 return frozenset() 

8489 return frozenset( 

8490 get_complete_model_list( 

8491 key_models=granted, 

8492 team_models=(), 

8493 proxy_model_list=proxy_model_list, 

8494 user_model=None, 

8495 infer_model_from_keys=False, 

8496 return_wildcard_routes=return_wildcard_routes, 

8497 llm_router=llm_router, 

8498 model_access_groups=model_access_groups, 

8499 include_model_access_groups=include_model_access_groups, 

8500 team_id=team_id, 

8501 ) 

8502 ) 

8503 

8504 

8505async def get_available_models_for_user( 

8506 user_api_key_dict: "UserAPIKeyAuth", 

8507 llm_router: Optional["Router"], 

8508 general_settings: dict, 

8509 user_model: str | None, 

8510 prisma_client: Optional["PrismaClient"] = None, 

8511 proxy_logging_obj: Optional["ProxyLogging"] = None, 

8512 team_id: str | None = None, 

8513 include_model_access_groups: bool = False, 

8514 only_model_access_groups: bool = False, 

8515 return_wildcard_routes: bool = False, 

8516 user_api_key_cache: Optional["UserApiKeyCache"] = None, 

8517 resolve_agent_ceiling: CeilingResolver = resolve_agent_access_group_ceiling, 

8518) -> list[str]: 

8519 """ 

8520 Get the list of models available to a user based on their API key and team permissions. 

8521 

8522 Args: 

8523 user_api_key_dict: User API key authentication object 

8524 llm_router: LiteLLM router instance 

8525 general_settings: General settings from config 

8526 user_model: User-specific model 

8527 prisma_client: Prisma client for database operations 

8528 proxy_logging_obj: Proxy logging object 

8529 team_id: Specific team ID to check (optional) 

8530 include_model_access_groups: Whether to include model access groups 

8531 only_model_access_groups: Whether to only return model access groups 

8532 return_wildcard_routes: Whether to return wildcard routes 

8533 

8534 Returns: 

8535 List of model names available to the user 

8536 """ 

8537 from litellm.proxy.auth.model_checks import ( 

8538 get_complete_model_list, 

8539 get_key_models, 

8540 get_team_models, 

8541 ) 

8542 

8543 # Get proxy model list and access groups 

8544 if llm_router is None: 8544 ↛ 8545line 8544 didn't jump to line 8545 because the condition on line 8544 was never true

8545 proxy_model_list = [] 

8546 model_access_groups = {} 

8547 else: 

8548 proxy_model_list = llm_router.get_model_names() 

8549 model_access_groups = llm_router.get_model_access_groups() 

8550 

8551 requested_team_object: Final = ( 

8552 await _get_validated_team_object( 

8553 user_api_key_dict=user_api_key_dict, 

8554 team_id=team_id, 

8555 prisma_client=prisma_client, 

8556 user_api_key_cache=user_api_key_cache, 

8557 proxy_logging_obj=proxy_logging_obj, 

8558 ) 

8559 if team_id and prisma_client and proxy_logging_obj and user_api_key_cache 

8560 else None 

8561 ) 

8562 

8563 key_models: Final[Sequence[str]] = ( 

8564 () 

8565 if requested_team_object is not None 

8566 else get_key_models( 

8567 user_api_key_dict=user_api_key_dict, 

8568 proxy_model_list=proxy_model_list, 

8569 model_access_groups=model_access_groups, 

8570 include_model_access_groups=include_model_access_groups, 

8571 ) 

8572 ) 

8573 

8574 team_models: Final = get_team_models( 

8575 team_models=( 

8576 requested_team_object.models if requested_team_object is not None else user_api_key_dict.team_models 

8577 ), 

8578 proxy_model_list=proxy_model_list, 

8579 model_access_groups=model_access_groups, 

8580 include_model_access_groups=include_model_access_groups, 

8581 ) 

8582 

8583 effective_team_id: Final = team_id or user_api_key_dict.team_id 

8584 

8585 access_group_models: Final = ( 

8586 await _get_access_group_models( 

8587 user_api_key_dict=user_api_key_dict, 

8588 team_object=requested_team_object 

8589 or await _get_team_object_for_access_groups( 

8590 team_id=effective_team_id, 

8591 prisma_client=prisma_client, 

8592 user_api_key_cache=user_api_key_cache, 

8593 proxy_logging_obj=proxy_logging_obj, 

8594 ), 

8595 prisma_client=prisma_client, 

8596 user_api_key_cache=user_api_key_cache, 

8597 proxy_logging_obj=proxy_logging_obj, 

8598 ) 

8599 if key_models or team_models 

8600 else () 

8601 ) 

8602 

8603 granted_key_models: Final = (*key_models, *access_group_models) if key_models else key_models 

8604 granted_team_models: Final = (*team_models, *access_group_models) if team_models else team_models 

8605 

8606 # Get complete model list 

8607 all_models: Final = get_complete_model_list( 

8608 key_models=granted_key_models, 

8609 team_models=granted_team_models, 

8610 proxy_model_list=proxy_model_list, 

8611 user_model=user_model, 

8612 infer_model_from_keys=general_settings.get("infer_model_from_keys", False), 

8613 return_wildcard_routes=return_wildcard_routes, 

8614 llm_router=llm_router, 

8615 model_access_groups=model_access_groups, 

8616 include_model_access_groups=include_model_access_groups, 

8617 only_model_access_groups=only_model_access_groups, 

8618 team_id=effective_team_id, 

8619 ) 

8620 

8621 agent_visible: Final = await _agent_access_group_visible_models( 

8622 user_api_key_dict=user_api_key_dict, 

8623 llm_router=llm_router, 

8624 include_model_access_groups=include_model_access_groups, 

8625 return_wildcard_routes=return_wildcard_routes, 

8626 team_id=effective_team_id, 

8627 resolve_agent_ceiling=resolve_agent_ceiling, 

8628 ) 

8629 if agent_visible is None: 8629 ↛ 8631line 8629 didn't jump to line 8631 because the condition on line 8629 was always true

8630 return all_models 

8631 capped: Final = [m for m in all_models if m in agent_visible] # mutable-ok: callers expect the list all_models is 

8632 return capped 

8633 

8634 

8635def _safe_get_model_info(model: str, get_model_info: Callable[[str], ModelInfo]) -> ModelInfo | None: 

8636 try: 

8637 return get_model_info(model) 

8638 except Exception as e: 

8639 verbose_proxy_logger.debug( 

8640 "create_model_info_response: cost map lookup failed for %s: %s", 

8641 model, 

8642 e, 

8643 ) 

8644 return None 

8645 

8646 

8647def _resolve_listing_model_info( 

8648 deployment_model: str | None, 

8649 listed_model: str, 

8650 listed_info: ModelInfo | None, 

8651 get_model_info: Callable[[str], ModelInfo], 

8652) -> tuple[ModelInfo, ...]: 

8653 """ 

8654 Cost-map entries describing one deployment behind a listed model, best source first. 

8655 

8656 The name a model is listed under is an arbitrary public alias, so it often misses the 

8657 cost map and lands on a fallback-generalization rule that answers with a conservative 

8658 family baseline instead of the real model's limits; the deployment's underlying model 

8659 is what the request actually reaches. Both names are kept because either can 

8660 generalize, and because a deployment's own model is registered into the cost map as a 

8661 stub that carries no limits of its own. Exact entries are consulted before generalized 

8662 ones, and each field is then taken from the first entry that has it. 

8663 

8664 ``listed_info`` is resolved once by the caller, since a group with several distinct 

8665 underlying models resolves the same alias for each of them. 

8666 """ 

8667 # Fast path, and the only one a wildcard-expanded name takes: with a single name 

8668 # there is nothing to order, so skip the generalization test entirely. This keeps 

8669 # the per-model cost of the listing on the hot path #33721 exists to protect. 

8670 if deployment_model is None or deployment_model == listed_model: 

8671 return () if listed_info is None else (listed_info,) 

8672 

8673 deployment_info: Final = _safe_get_model_info(deployment_model, get_model_info) 

8674 if deployment_info is None: 8674 ↛ 8675line 8674 didn't jump to line 8675 because the condition on line 8674 was never true

8675 return () if listed_info is None else (listed_info,) 

8676 if listed_info is None: 8676 ↛ 8679line 8676 didn't jump to line 8679 because the condition on line 8676 was always true

8677 return (deployment_info,) 

8678 

8679 from litellm.utils import is_generalized_model_info 

8680 

8681 # Both names resolved: the deployment's model leads unless it only generalized 

8682 # while the listed name is an exact cost-map entry. 

8683 if is_generalized_model_info(deployment_info) and not is_generalized_model_info(listed_info): 

8684 return (listed_info, deployment_info) 

8685 return (deployment_info, listed_info) 

8686 

8687 

8688def _first_token_limit(candidates: tuple[ModelInfo, ...], field: str) -> int | None: 

8689 return next( 

8690 (limit for limit in (coerce_token_limit(info.get(field)) for info in candidates) if limit is not None), 

8691 None, 

8692 ) 

8693 

8694 

8695def _group_token_limit(candidate_sets: tuple[tuple[ModelInfo, ...], ...], field: str) -> int | None: 

8696 """The widest limit any deployment behind the listed name declares for ``field``. 

8697 

8698 A model group is normally one model behind several interchangeable deployments, so 

8699 there is a single value to report and the choice of aggregate does not arise. 

8700 

8701 When a group genuinely mixes models no single number is right, and the widest is the 

8702 deliberate pick over the narrowest for two reasons. It is what ``/model_group/info`` 

8703 has long reported to the Admin UI, so the two surfaces agree; disagreeing is the very 

8704 complaint this resolution path exists to fix. And of the two ways to be wrong, 

8705 under-advertising is worse: a client that trusts a narrowed window silently refuses 

8706 prompts the group would have served, while an over-long prompt that reaches a smaller 

8707 deployment comes back as a legible context-length error -- and does not reach one at 

8708 all when ``enable_pre_call_checks`` is set, which filters deployments the prompt does 

8709 not fit. 

8710 """ 

8711 limits: Final = tuple( 

8712 limit for limit in (_first_token_limit(candidates, field) for candidates in candidate_sets) if limit is not None 

8713 ) 

8714 return max(limits) if limits else None 

8715 

8716 

8717def create_model_info_response( 

8718 model_id: str, 

8719 provider: str, 

8720 include_metadata: bool = False, 

8721 fallback_type: str | None = None, 

8722 llm_router: Optional["Router"] = None, 

8723 get_model_info: Callable[[str], ModelInfo] = litellm.get_model_info, 

8724) -> ModelInfoResponse: 

8725 """ 

8726 Create a standardized OpenAI-compatible model object. 

8727 

8728 When include_metadata is true, attaches the model's configured fallbacks 

8729 (resolved via the router under fallback_type, defaulting to "general"). 

8730 Raises HTTPException(400) for an unknown fallback_type. 

8731 """ 

8732 from litellm.proxy.auth.model_checks import get_all_fallbacks 

8733 

8734 base: Final[ModelInfoResponse] = { 

8735 "id": model_id, 

8736 "object": "model", 

8737 "created": DEFAULT_MODEL_CREATED_AT_TIME, 

8738 "owned_by": provider, 

8739 } 

8740 

8741 alias_target: Final = ( 

8742 resolve_model_group_alias(llm_router.model_group_alias, model_id) if llm_router is not None else None 

8743 ) 

8744 lookup_model: Final = alias_target if alias_target is not None else model_id 

8745 

8746 listing_info: Final = llm_router.get_model_listing_info(lookup_model) if llm_router is not None else None 

8747 

8748 # One entry per distinct model behind the listed name; (None,) when the router knows 

8749 # nothing about it, so the listed name is resolved on its own as before. 

8750 deployment_models: Final[tuple[str | None, ...]] = ( 

8751 listing_info.cost_map_keys if listing_info is not None and listing_info.cost_map_keys else (None,) 

8752 ) 

8753 listed_info: Final = _safe_get_model_info(lookup_model, get_model_info) 

8754 candidate_sets: Final = tuple( 

8755 _resolve_listing_model_info( 

8756 deployment_model=deployment_model, 

8757 listed_model=lookup_model, 

8758 listed_info=listed_info, 

8759 get_model_info=get_model_info, 

8760 ) 

8761 for deployment_model in deployment_models 

8762 ) 

8763 

8764 max_input_tokens: int | None = _group_token_limit(candidate_sets, "max_input_tokens") 

8765 max_output_tokens: int | None = _group_token_limit(candidate_sets, "max_output_tokens") 

8766 mode: Final = next( 

8767 ( 

8768 m 

8769 for m in ( 

8770 cast("Mapping[str, object]", info).get("mode") # cast-ok: an entry need not carry "mode" 

8771 for candidates in candidate_sets 

8772 for info in candidates 

8773 ) 

8774 if isinstance(m, str) 

8775 ), 

8776 None, 

8777 ) 

8778 if mode is not None: 

8779 base["mode"] = mode 

8780 

8781 if listing_info is not None: 

8782 if listing_info.max_input_tokens is not None: 8782 ↛ 8783line 8782 didn't jump to line 8783 because the condition on line 8782 was never true

8783 max_input_tokens = listing_info.max_input_tokens 

8784 if listing_info.max_output_tokens is not None: 8784 ↛ 8785line 8784 didn't jump to line 8785 because the condition on line 8784 was never true

8785 max_output_tokens = listing_info.max_output_tokens 

8786 

8787 if llm_router is not None: 8787 ↛ 8792line 8787 didn't jump to line 8792 because the condition on line 8787 was always true

8788 configured_mode: Final = llm_router.get_configured_mode(lookup_model) 

8789 if isinstance(configured_mode, str): 8789 ↛ 8790line 8789 didn't jump to line 8790 because the condition on line 8789 was never true

8790 base["mode"] = configured_mode 

8791 

8792 if max_input_tokens is not None: 

8793 base["max_input_tokens"] = max_input_tokens 

8794 if max_output_tokens is not None: 

8795 base["max_output_tokens"] = max_output_tokens 

8796 

8797 if not include_metadata: 

8798 return base 

8799 

8800 effective_fallback_type: Final = fallback_type if fallback_type is not None else "general" 

8801 

8802 valid_fallback_types: Final = ["general", "context_window", "content_policy"] 

8803 if effective_fallback_type not in valid_fallback_types: 

8804 raise HTTPException( 

8805 status_code=400, 

8806 detail=f"Invalid fallback_type. Must be one of: {valid_fallback_types}", 

8807 ) 

8808 

8809 fallbacks: Final = get_all_fallbacks( 

8810 model=model_id, 

8811 llm_router=llm_router, 

8812 fallback_type=effective_fallback_type, 

8813 ) 

8814 return {**base, "metadata": {"fallbacks": fallbacks}} 

8815 

8816 

8817def validate_model_access( 

8818 model_id: str, 

8819 available_models: list[str], 

8820) -> None: 

8821 """ 

8822 Validate that a model is accessible to the user. 

8823 Supports batch requests with comma-separated model IDs. 

8824 

8825 Args: 

8826 model_id: The model ID to validate (can be comma-separated for batch requests) 

8827 available_models: List of models available to the user 

8828 

8829 Raises: 

8830 HTTPException: If the model is not accessible 

8831 """ 

8832 # Handle batch requests with comma-separated models 

8833 if "," in model_id: 

8834 models: Final = [m.strip() for m in model_id.split(",")] 

8835 inaccessible_models: Final = [m for m in models if m not in available_models] 

8836 if inaccessible_models: 8836 ↛ exitline 8836 didn't return from function 'validate_model_access' because the condition on line 8836 was always true

8837 raise HTTPException( 

8838 status_code=404, 

8839 detail="The following model(s) do not exist or are not accessible: {}".format( 

8840 ", ".join(inaccessible_models) 

8841 ), 

8842 ) 

8843 else: 

8844 # Single model validation 

8845 if model_id not in available_models: 8845 ↛ exitline 8845 didn't return from function 'validate_model_access' because the condition on line 8845 was always true

8846 raise HTTPException( 

8847 status_code=404, 

8848 detail=f"The model `{model_id}` does not exist or is not accessible", 

8849 ) 

8850 

8851 

8852_PRESERVED_NONE_FIELDS: Final[list[tuple[str, str]]] = [ 

8853 ("message", "content"), # null when tool_calls present (issue #6677) 

8854 ("message", "role"), # always required by OpenAI spec 

8855 ("delta", "content"), # null in streaming chunks 

8856] 

8857 

8858 

8859def model_dump_with_preserved_fields( 

8860 obj: Any, 

8861 preserve_fields: list[str] | None = None, 

8862 exclude_unset: bool = True, 

8863) -> dict[str, object]: 

8864 """ 

8865 Serialize a Pydantic model to a dictionary while preserving specific fields 

8866 even if they are None. 

8867 

8868 Fields listed in _PRESERVED_NONE_FIELDS are restored after 

8869 model_dump(exclude_none=True) strips them. 

8870 

8871 Args: 

8872 obj: The Pydantic BaseModel instance to serialize 

8873 preserve_fields: Deprecated, kept for backward compatibility. 

8874 exclude_unset: Whether to exclude fields that were not explicitly set 

8875 

8876 Returns: 

8877 Dictionary representation with None values excluded except for preserved fields 

8878 """ 

8879 result: Final = obj.model_dump(exclude_none=True, exclude_unset=exclude_unset) 

8880 

8881 choices: Final = result.get("choices") 

8882 if not choices: 

8883 return result 

8884 

8885 obj_choices: Final = obj.choices 

8886 for choice_obj, choice_dict in zip(obj_choices, choices): 

8887 for sub_object, field_name in _PRESERVED_NONE_FIELDS: 

8888 sub_dict = choice_dict.get(sub_object) 

8889 if sub_dict is None: 

8890 continue 

8891 if field_name not in sub_dict: 

8892 sub_obj = getattr(choice_obj, sub_object, None) 

8893 if sub_obj is not None and hasattr(sub_obj, field_name): 

8894 sub_dict[field_name] = getattr(sub_obj, field_name) 

8895 

8896 return result