Merge remote-tracking branch 'upstream/litellm_internal_staging' into deepkeep-as-internal

This commit is contained in:
Yaniv Israel 2026-06-15 13:02:50 +03:00
commit e1f3ca0c27
563 changed files with 38035 additions and 8147 deletions

49
.github/workflows/osv-scan.yml vendored Normal file
View file

@ -0,0 +1,49 @@
name: OSV Scan
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_branch
- "litellm_**"
paths:
- uv.lock
- ui/litellm-dashboard/package-lock.json
- osv-scanner.toml
- .github/workflows/osv-scan.yml
schedule:
- cron: "23 6 * * *"
workflow_dispatch:
permissions: {}
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
osv-scan:
name: osv-scan
runs-on: ubuntu-latest
timeout-minutes: 10
permissions:
contents: read
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Download osv-scanner v2.3.8
run: |
curl -fsSL --retry 3 -o "$RUNNER_TEMP/osv-scanner" \
https://github.com/google/osv-scanner/releases/download/v2.3.8/osv-scanner_linux_amd64
echo "bc98e15319ed0d515e3f9235287ba53cdc5535d576d24fd573978ecfe9ab92dc $RUNNER_TEMP/osv-scanner" | sha256sum -c -
chmod +x "$RUNNER_TEMP/osv-scanner"
- name: Scan lockfiles
run: |
"$RUNNER_TEMP/osv-scanner" scan source \
--config osv-scanner.toml \
-L uv.lock \
-L ui/litellm-dashboard/package-lock.json

View file

@ -67,6 +67,12 @@ jobs:
uv run --no-sync ruff check .
cd ..
- name: Check strict-rule budget (delta vs base)
env:
BASE_SHA: ${{ github.event.pull_request.base.sha }}
run: |
uv run --no-sync python scripts/ruff_strict_gate.py --base "$BASE_SHA"
- name: Print OpenAI version
run: |
uv run --no-sync python -c "import openai; print(f'OpenAI version: {openai.__version__}')"

View file

@ -36,6 +36,8 @@ Don't hesitate to use values in .env to get needed API keys and other secrets, a
Run tests, format your code, and lint your code before each commit
When you fix strict-rule violations gated by `ruff-strict-budget.json`, run `make lint-strict-budget-update` and commit the lowered baselines so the ceilings ratchet down instead of leaving stale headroom
Ask to commit and push your work when you're done (or if you're confident that your code is good and works, just do it)
When you must use real LLM models to, for example, write e2e tests, write a QA runbook, etc., make sure to use the latest models (doesn't have to be smartest, can also be a modern small, fast one. No strong preference for smart vs fast here, just use something modern) as of the year and month of the current date. Do a web search as necessary to figure that out
@ -57,11 +59,15 @@ Follow these coding conventions for new/updated code (a three-line fix in a lega
- Composition over inheritance
- Never-nester: early returns over deep nesting
- Don't throw; model failures as values (One function (e.g., raise_public) maps error union to existing public exception contracts via exhaustive match + assert_never)
- No mutation; instead of mutable lists and dicts, prefer tuples, NamedTuples, frozen dataclasses, etc.
- No mutation; don't reassign variables, global or local. Instead of mutable lists and dicts, prefer tuples, frozen dataclasses (with slots=True), etc.
- Use dependency injection
- Fully typed; no `Any` or coarse types like dict[str, Any]. Every function parameter must be strongly typed
- Fully typed; no `Any` or coarse types like `dict[str, Any]` or just `dict`. Every function parameter must be strongly typed
- Use tagged unions + match
- No monster files or god objects
- No file sprawl: deliberate file and folder structure
- Standard over hand-rolled: use the official SDK or a library where one exists; where none does, follow industry standards instead of inventing local conventions
if you're trying to create a new function that relies on untyped stuff, instead of adding more Any's and bringing it closer to the max, just validate it in the caller (a simple function that returns the typed thing or raises will do) and then pass the now typed variable in
Follow conventional commits for commit names and PR titles

View file

@ -68,22 +68,24 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
USER root
RUN apk add --no-cache bash openssl tzdata nodejs npm python3 libsndfile && \
npm install -g npm@11.14.0 tar@7.5.11 glob@13.0.6 @isaacs/brace-expansion@5.0.1 brace-expansion@5.0.5 minimatch@10.2.4 diff@8.0.3 picomatch@4.0.4 && \
GLOBAL="$(npm root -g)" && \
for pkg in tar glob @isaacs/brace-expansion brace-expansion minimatch diff picomatch; do \
name="${pkg##*/}"; \
find "$GLOBAL/npm" -type d -name "$name" -path "*/node_modules/$pkg" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/$pkg" "$d"; \
done; \
done && \
npm cache clean --force && \
{ apk del --no-cache npm 2>/dev/null || true; }
# node (without npm) is required by the prisma CLI at runtime
RUN apk add --no-cache bash openssl tzdata nodejs python3 libsndfile
WORKDIR /app
ENV PATH="/app/.venv/bin:${PATH}"
COPY --from=builder /app /app
# Copy only what runtime needs. The application is installed inside the venv;
# the rest of the builder's /app is source and build metadata that must not
# ship (manifest-scanning tools attribute everything in it to this image).
# entrypoint.sh invokes litellm/proxy/prisma_migration.py by source path.
COPY --from=builder /app/.venv /app/.venv
COPY --from=builder /app/docker /app/docker
COPY --from=builder /app/schema.prisma /app/schema.prisma
COPY --from=builder /app/litellm/proxy/prisma_migration.py /app/litellm/proxy/prisma_migration.py
# enterprise/ is imported by source path at runtime (proxy_cli puts the
# working directory on sys.path; litellm/proxy/hooks resolves
# enterprise.enterprise_hooks from it)
COPY --from=builder /app/enterprise /app/enterprise
# Prisma binaries live in $HOME/.cache (default prisma-python location),
# which is /root/.cache here. Copy only the Prisma subdirs — copying the
# whole /root/.cache drags in the uv build cache (~660 MB, includes a

View file

@ -5,6 +5,7 @@
test-unit-integrations test-unit-core-utils test-unit-other test-unit-root \
test-proxy-unit-a test-proxy-unit-b test-integration test-unit-helm \
info lint lint-dev format \
lint-strict-budget lint-strict-budget-update \
install-dev install-proxy-dev install-test-deps install-hooks \
install-helm-unittest check-circular-imports check-import-safety
@ -24,6 +25,8 @@ help:
@echo " make lint-ruff - Run Ruff linting only"
@echo " make lint-mypy - Run MyPy type checking only"
@echo " make lint-black - Check Black formatting (matches CI)"
@echo " make lint-strict-budget - Gate the codebase total of each strict ruff rule against its ceiling"
@echo " make lint-strict-budget-update - Re-capture per-rule baselines in ruff-strict-budget.json (ratchet)"
@echo " make check-circular-imports - Check for circular imports"
@echo " make check-import-safety - Check import safety"
@echo " make test - Run all tests"
@ -122,6 +125,12 @@ lint-mypy: install-dev
lint-black: format-check
lint-strict-budget: install-dev
$(UV_RUN) python scripts/ruff_strict_gate.py
lint-strict-budget-update: install-dev
$(UV_RUN) python scripts/ruff_strict_gate.py --update
check-circular-imports: install-dev
cd litellm && $(UV_RUN) python ../tests/documentation_tests/test_circular_imports.py && cd ..
@ -129,7 +138,7 @@ check-import-safety: install-dev
@$(UV_RUN) python -c "from litellm import *; print('[from litellm import *] OK! no issues!');" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1)
# Combined linting (matches test-linting.yml workflow)
lint: format-check lint-ruff lint-mypy check-circular-imports check-import-safety
lint: format-check lint-ruff lint-mypy check-circular-imports check-import-safety lint-strict-budget
# Faster linting for local development (only checks changed code)
lint-dev: lint-format-changed lint-mypy check-circular-imports check-import-safety

View file

@ -20,7 +20,11 @@ DatabaseURLSettings.from_env().apply_to_env()
from litellm.proxy.proxy_server import app
from backend.routes.allowlist import BACKEND_EXACT_PATHS, BACKEND_PATH_PREFIXES
from backend.routes.allowlist import (
BACKEND_EXACT_PATHS,
BACKEND_MOUNT_PATHS,
BACKEND_PATH_PREFIXES,
)
def _is_backend_route(route) -> bool:
@ -29,8 +33,9 @@ def _is_backend_route(route) -> bool:
if path is None:
return False
if isinstance(route, Mount):
# Static UI mounts are served by the dedicated UI container, not here.
return False
# The dashboard UI static mounts are served by the dedicated UI container.
# Only Mounts in the backend allowlist (e.g. swagger docs) remain on backend.
return path in BACKEND_MOUNT_PATHS
if path in BACKEND_EXACT_PATHS:
return True
return any(path.startswith(prefix) for prefix in BACKEND_PATH_PREFIXES)

View file

@ -135,3 +135,9 @@ BACKEND_EXACT_PATHS: frozenset[str] = frozenset(
"/fallback/login",
}
)
BACKEND_MOUNT_PATHS: frozenset[str] = frozenset(
{
"/swagger", # API documentation static assets belong to the backend
}
)

View file

@ -0,0 +1,99 @@
-- Converts an existing LiteLLM_SpendLogs table into a native Postgres
-- range-partitioned table keyed on "startTime".
--
-- Why: at high request volume, retention via DELETE leaves dead tuples that
-- autovacuum cannot reclaim quickly enough, so the table keeps growing on disk
-- (seen at 450GB+ after ~1 month). With partitioning, retention drops whole
-- partitions, which is instant and returns disk to the OS immediately.
--
-- This is an opt-in, manual operation. The default LiteLLM schema is NOT
-- partitioned, so existing installs are unaffected until you run this.
--
-- IMPORTANT
-- * Test on a staging copy first and take a backup.
-- * Postgres cannot convert a populated table to partitioned in place, so this
-- renames the old table aside and creates a fresh partitioned table.
-- * The partition key ("startTime") must be part of the primary key, so the
-- PK becomes composite ("request_id", "startTime"). LiteLLM's write path uses
-- INSERT ... ON CONFLICT DO NOTHING, which is compatible with this.
-- * Choose a partition granularity ("day" is the recommended default for
-- high-volume tables) and keep it consistent with SPEND_LOG_PARTITION_INTERVAL.
--
-- After running this, enable the feature and set a retention period in
-- proxy_config.yaml:
-- general_settings:
-- use_spend_logs_partitioning: true
-- maximum_spend_logs_retention_period: "30d"
-- The spend-log cleanup job then verifies the table is partitioned and reclaims
-- disk by dropping expired partitions instead of deleting rows. It also
-- pre-creates upcoming partitions on each run. To roll back, see
-- db_scripts/unpartition_spend_logs.sql.
BEGIN;
ALTER TABLE "LiteLLM_SpendLogs" RENAME TO "LiteLLM_SpendLogs_legacy";
-- Renaming a table does NOT rename its indexes, and index names are unique per
-- schema. Move the legacy table's indexes aside so the CREATE INDEX statements
-- below actually create indexes on the new partitioned table instead of being
-- silently skipped by IF NOT EXISTS, and so the new PK keeps the canonical
-- name instead of getting a "_pkey1" suffix.
ALTER INDEX IF EXISTS "LiteLLM_SpendLogs_pkey"
RENAME TO "LiteLLM_SpendLogs_legacy_pkey";
ALTER INDEX IF EXISTS "LiteLLM_SpendLogs_startTime_idx"
RENAME TO "LiteLLM_SpendLogs_legacy_startTime_idx";
ALTER INDEX IF EXISTS "LiteLLM_SpendLogs_startTime_request_id_idx"
RENAME TO "LiteLLM_SpendLogs_legacy_startTime_request_id_idx";
ALTER INDEX IF EXISTS "LiteLLM_SpendLogs_end_user_idx"
RENAME TO "LiteLLM_SpendLogs_legacy_end_user_idx";
ALTER INDEX IF EXISTS "LiteLLM_SpendLogs_session_id_idx"
RENAME TO "LiteLLM_SpendLogs_legacy_session_id_idx";
CREATE TABLE "LiteLLM_SpendLogs" (
LIKE "LiteLLM_SpendLogs_legacy" INCLUDING DEFAULTS INCLUDING GENERATED
) PARTITION BY RANGE ("startTime");
ALTER TABLE "LiteLLM_SpendLogs"
ADD PRIMARY KEY ("request_id", "startTime");
-- Recreate every index Prisma defines on the table. LIKE ... INCLUDING DEFAULTS
-- INCLUDING GENERATED copies columns and defaults but NOT indexes, so without
-- these the admin-UI cost-reporting queries that filter by end_user/session_id
-- fall back to sequential scans. On a partitioned parent these propagate to
-- every current and future partition automatically.
CREATE INDEX IF NOT EXISTS "LiteLLM_SpendLogs_startTime_idx"
ON "LiteLLM_SpendLogs" ("startTime");
CREATE INDEX IF NOT EXISTS "LiteLLM_SpendLogs_startTime_request_id_idx"
ON "LiteLLM_SpendLogs" ("startTime", "request_id");
CREATE INDEX IF NOT EXISTS "LiteLLM_SpendLogs_end_user_idx"
ON "LiteLLM_SpendLogs" ("end_user");
CREATE INDEX IF NOT EXISTS "LiteLLM_SpendLogs_session_id_idx"
ON "LiteLLM_SpendLogs" ("session_id");
-- Safety net: any row whose startTime has no explicit partition lands here so
-- writes never fail. The cleanup job never drops the DEFAULT partition.
CREATE TABLE IF NOT EXISTS "LiteLLM_SpendLogs_pdefault"
PARTITION OF "LiteLLM_SpendLogs" DEFAULT;
COMMIT;
-- Backfill (optional). Rows route to the correct partition automatically.
-- For large legacy tables, copy in time-bounded batches during a low-traffic
-- window instead of one statement, or simply keep "LiteLLM_SpendLogs_legacy"
-- read-only until its data ages past your retention, then DROP it.
--
-- Backfilled rows land in the DEFAULT partition until explicit partitions
-- cover their dates. Postgres refuses to create a partition whose range
-- overlaps rows already in DEFAULT, so the cleanup job may log a warning when
-- pre-creating today's partition right after a backfill; it recovers on its
-- own once those dates age out, and future partitions are unaffected because
-- they are always created ahead of writes.
--
-- INSERT INTO "LiteLLM_SpendLogs"
-- SELECT * FROM "LiteLLM_SpendLogs_legacy"
-- WHERE "startTime" >= now() - interval '30 days';
--
-- DROP TABLE "LiteLLM_SpendLogs_legacy";

View file

@ -0,0 +1,69 @@
-- Rolls back db_scripts/partition_spend_logs.sql: converts the native
-- range-partitioned "LiteLLM_SpendLogs" table back into a plain,
-- non-partitioned table matching the default LiteLLM schema.
--
-- When/why: run this if you want to stop using partition-based retention and
-- return to DELETE-based cleanup, or to restore the original single-column
-- primary key ("request_id") that the partitioned layout had to widen to a
-- composite ("request_id", "startTime").
--
-- IMPORTANT
-- * Test on a staging copy first and take a backup.
-- * Postgres cannot convert a partitioned table back in place, so this
-- renames the partitioned table aside and creates a fresh plain table.
-- * The composite PK could in principle hold the same "request_id" in more
-- than one partition, so rows are copied with ON CONFLICT DO NOTHING to
-- restore the single-column PK without failing on such duplicates.
-- * For large tables the INSERT ... SELECT copies every surviving row and may
-- run long; do it during a low-traffic window.
-- * Also remove use_spend_logs_partitioning from proxy_config.yaml (or set it
-- to false) so the cleanup job returns to DELETE-based retention.
BEGIN;
ALTER TABLE "LiteLLM_SpendLogs" RENAME TO "LiteLLM_SpendLogs_partitioned";
-- Renaming a table does NOT rename its indexes, and index names are unique per
-- schema. Move the partitioned table's indexes aside so the CREATE INDEX
-- statements below actually create indexes on the new plain table instead of
-- being silently skipped by IF NOT EXISTS, and so the new PK keeps the
-- canonical name.
ALTER INDEX IF EXISTS "LiteLLM_SpendLogs_pkey"
RENAME TO "LiteLLM_SpendLogs_partitioned_pkey";
ALTER INDEX IF EXISTS "LiteLLM_SpendLogs_pkey1"
RENAME TO "LiteLLM_SpendLogs_partitioned_pkey1";
ALTER INDEX IF EXISTS "LiteLLM_SpendLogs_startTime_idx"
RENAME TO "LiteLLM_SpendLogs_partitioned_startTime_idx";
ALTER INDEX IF EXISTS "LiteLLM_SpendLogs_startTime_request_id_idx"
RENAME TO "LiteLLM_SpendLogs_partitioned_startTime_request_id_idx";
ALTER INDEX IF EXISTS "LiteLLM_SpendLogs_end_user_idx"
RENAME TO "LiteLLM_SpendLogs_partitioned_end_user_idx";
ALTER INDEX IF EXISTS "LiteLLM_SpendLogs_session_id_idx"
RENAME TO "LiteLLM_SpendLogs_partitioned_session_id_idx";
CREATE TABLE "LiteLLM_SpendLogs" (
LIKE "LiteLLM_SpendLogs_partitioned" INCLUDING DEFAULTS INCLUDING GENERATED
);
ALTER TABLE "LiteLLM_SpendLogs"
ADD PRIMARY KEY ("request_id");
CREATE INDEX IF NOT EXISTS "LiteLLM_SpendLogs_startTime_idx"
ON "LiteLLM_SpendLogs" ("startTime");
CREATE INDEX IF NOT EXISTS "LiteLLM_SpendLogs_startTime_request_id_idx"
ON "LiteLLM_SpendLogs" ("startTime", "request_id");
CREATE INDEX IF NOT EXISTS "LiteLLM_SpendLogs_end_user_idx"
ON "LiteLLM_SpendLogs" ("end_user");
CREATE INDEX IF NOT EXISTS "LiteLLM_SpendLogs_session_id_idx"
ON "LiteLLM_SpendLogs" ("session_id");
INSERT INTO "LiteLLM_SpendLogs"
SELECT * FROM "LiteLLM_SpendLogs_partitioned"
ON CONFLICT ("request_id") DO NOTHING;
DROP TABLE "LiteLLM_SpendLogs_partitioned";
COMMIT;

View file

@ -66,36 +66,31 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
USER root
RUN apk add --no-cache bash openssl tzdata nodejs npm python3 libsndfile && \
npm install -g npm@11.12.1 tar@7.5.11 glob@11.1.0 @isaacs/brace-expansion@5.0.1 minimatch@10.2.4 diff@8.0.3 && \
GLOBAL="$(npm root -g)" && \
find "$GLOBAL/npm" -type d -name "tar" -path "*/node_modules/tar" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/tar" "$d"; \
done && \
find "$GLOBAL/npm" -type d -name "glob" -path "*/node_modules/glob" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/glob" "$d"; \
done && \
find "$GLOBAL/npm" -type d -name "brace-expansion" -path "*/node_modules/@isaacs/brace-expansion" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/@isaacs/brace-expansion" "$d"; \
done && \
find "$GLOBAL/npm" -type d -name "minimatch" -path "*/node_modules/minimatch" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/minimatch" "$d"; \
done && \
find "$GLOBAL/npm" -type d -name "diff" -path "*/node_modules/diff" | while read d; do \
rm -rf "$d" && cp -rL "$GLOBAL/diff" "$d"; \
done && \
npm cache clean --force && \
{ apk del --no-cache npm 2>/dev/null || true; }
# node (without npm) is required by the prisma CLI at runtime
RUN apk add --no-cache bash openssl tzdata nodejs python3 libsndfile
WORKDIR /app
ENV PATH="/app/.venv/bin:${PATH}"
COPY --from=builder /app /app
# Copy only what runtime needs. The application is installed inside the venv;
# the rest of the builder's /app is source and build metadata that must not
# ship (manifest-scanning tools attribute everything in it to this image).
# entrypoint.sh invokes litellm/proxy/prisma_migration.py by source path.
COPY --from=builder /app/.venv /app/.venv
COPY --from=builder /app/docker /app/docker
COPY --from=builder /app/schema.prisma /app/schema.prisma
COPY --from=builder /app/litellm/proxy/prisma_migration.py /app/litellm/proxy/prisma_migration.py
# enterprise/ is imported by source path at runtime (proxy_cli puts the
# working directory on sys.path; litellm/proxy/hooks resolves
# enterprise.enterprise_hooks from it)
COPY --from=builder /app/enterprise /app/enterprise
# Prisma binaries live in $HOME/.cache (default prisma-python location),
# which is /root/.cache here. Copy them from the builder so they survive
# deployments that volume-mount /app/.cache (e.g. readOnlyRootFilesystem
# + emptyDir) — otherwise the mount would shadow the baked-in query engine.
COPY --from=builder /root/.cache /root/.cache
# Only the Prisma subdirs: the whole /root/.cache drags in the uv build cache.
COPY --from=builder /root/.cache/prisma /root/.cache/prisma
COPY --from=builder /root/.cache/prisma-python /root/.cache/prisma-python
RUN find /app/.venv -type f -path "*/tornado/test/*" -delete && \
find /app/.venv -type d -path "*/tornado/test" -delete

View file

@ -95,7 +95,21 @@ RUN for i in 1 2 3; do \
apk add --no-cache python3 bash openssl tzdata libsndfile nodejs && break || sleep 5; \
done
COPY --from=builder /app /app
# Copy only what runtime needs. The application is installed inside the venv;
# the rest of the builder's /app is source and build metadata that must not
# ship (manifest-scanning tools attribute everything in it to this image).
# entrypoint.sh invokes litellm/proxy/prisma_migration.py by source path.
# Prisma caches live under /app/.cache here (XDG_CACHE_HOME /
# PRISMA_BINARY_CACHE_DIR) so the runtime prisma generate finds them.
COPY --from=builder /app/.venv /app/.venv
COPY --from=builder /app/docker /app/docker
COPY --from=builder /app/schema.prisma /app/schema.prisma
COPY --from=builder /app/litellm/proxy/prisma_migration.py /app/litellm/proxy/prisma_migration.py
# enterprise/ is imported by source path at runtime (proxy_cli puts the
# working directory on sys.path; litellm/proxy/hooks resolves
# enterprise.enterprise_hooks from it)
COPY --from=builder /app/enterprise /app/enterprise
COPY --from=builder /app/.cache /app/.cache
COPY --from=builder /var/lib/litellm/ui /var/lib/litellm/ui
COPY --from=builder /var/lib/litellm/assets /var/lib/litellm/assets

View file

@ -504,7 +504,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
if retrieve_file_id
else False
)
if potential_file_id:
if potential_file_id and "llm_output_file_id," in potential_file_id:
model_id = self.get_model_id_from_unified_file_id(potential_file_id)
if model_id:
data["model"] = model_id
@ -1058,7 +1058,12 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
return file_id.split("llm_output_file_model_id,")[1].split(";")[0]
def get_output_file_id_from_unified_file_id(self, file_id: str) -> str:
return file_id.split("llm_output_file_id,")[1].split(";")[0]
marker = "llm_output_file_id,"
if marker not in file_id:
raise ValueError(
f"Unified id does not contain {marker!r}: {file_id[:80]!r}"
)
return file_id.split(marker, 1)[1].split(";")[0]
async def async_post_call_success_hook(
self, data: Dict, user_api_key_dict: UserAPIKeyAuth, response: LLMResponseTypes
@ -1099,13 +1104,33 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
for file_attr in ["output_file_id", "error_file_id"]:
file_id_value = getattr(response, file_attr, None)
if file_id_value and model_id:
original_file_id = file_id_value
unified_file_id = self.get_unified_output_file_id(
output_file_id=original_file_id,
model_id=model_id,
model_name=resolved_model_name,
decoded_output_file_id = _is_base64_encoded_unified_file_id(
file_id_value
)
setattr(response, file_attr, unified_file_id)
if (
decoded_output_file_id
and "llm_output_file_id," in decoded_output_file_id
):
provider_file_id = (
self.get_output_file_id_from_unified_file_id(
decoded_output_file_id
)
)
unified_file_id = file_id_value
elif decoded_output_file_id:
verbose_logger.warning(
f"Skipping {file_attr}={file_id_value!r}: "
"unified id is not a managed file output id"
)
continue
else:
provider_file_id = file_id_value
unified_file_id = self.get_unified_output_file_id(
output_file_id=provider_file_id,
model_id=model_id,
model_name=resolved_model_name,
)
setattr(response, file_attr, unified_file_id)
# Use llm_router credentials when available. Without credentials,
# Azure and other auth-required providers return 500/401.
@ -1125,27 +1150,27 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
or {}
)
file_object = await litellm.afile_retrieve(
file_id=original_file_id,
file_id=provider_file_id,
**_creds,
)
else:
file_object = await litellm.afile_retrieve(
custom_llm_provider=model_name.split("/")[0] if model_name and "/" in model_name else "openai", # type: ignore[arg-type]
file_id=original_file_id,
file_id=provider_file_id,
)
verbose_logger.debug(
f"Successfully retrieved file object for {file_attr}={original_file_id}"
f"Successfully retrieved file object for {file_attr}={provider_file_id}"
)
except Exception as e:
verbose_logger.warning(
f"Failed to retrieve file object for {file_attr}={original_file_id}: {str(e)}. Storing with None and will fetch on-demand."
f"Failed to retrieve file object for {file_attr}={provider_file_id}: {str(e)}. Storing with None and will fetch on-demand."
)
await self.store_unified_file_id(
file_id=unified_file_id,
file_object=file_object,
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
model_mappings={model_id: original_file_id},
model_mappings={model_id: provider_file_id},
user_api_key_dict=user_api_key_dict,
)
await self.store_unified_object_id(

View file

@ -43,6 +43,7 @@ from typing import (
Type,
)
from litellm.types.integrations.datadog import DatadogInitParams
from litellm.types.integrations.newrelic import NewRelicInitParams
from litellm._logging import (
set_verbose,
_turn_on_debug,
@ -154,10 +155,12 @@ _custom_logger_compatible_callbacks_literal = Literal[
"gitlab",
"cloudzero",
"focus",
"mavvrik",
"vantage",
"posthog",
"levo",
"compression_interception",
"newrelic",
]
cold_storage_custom_logger: Optional[_custom_logger_compatible_callbacks_literal] = None
logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None
@ -359,6 +362,9 @@ enable_gemini_default_thinking_level_low: bool = (
####################
logging: bool = True
enable_loadbalancing_on_batch_endpoints: Optional[bool] = None
require_managed_files: bool = (
False # proxy only - require target_model_names on POST /v1/files
)
enable_caching_on_provider_specific_optional_params: bool = (
False # feature-flag for caching on optional params - e.g. 'top_k'
)
@ -412,6 +418,7 @@ s3_callback_params: Optional[Dict] = None
s3_audit_callback_params: Optional[Dict] = None
datadog_llm_observability_params: Optional[Union[DatadogLLMObsInitParams, Dict]] = None
datadog_params: Optional[Union[DatadogInitParams, Dict]] = None
newrelic_params: Optional[Union[NewRelicInitParams, Dict]] = None
aws_sqs_callback_params: Optional[Dict] = None
generic_logger_headers: Optional[Dict] = None
default_key_generate_params: Optional[Dict] = None
@ -1373,6 +1380,7 @@ from .search.main import *
from .realtime_api.main import (
_arealtime,
acreate_realtime_client_secret,
acreate_realtime_transcription_session,
arealtime_calls,
)
from .responses.main import _aresponses_websocket
@ -1723,6 +1731,9 @@ if TYPE_CHECKING:
from .llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig as VoyageContextualEmbeddingConfig,
)
from .llms.voyage.embedding.transformation_multimodal import (
VoyageMultimodalEmbeddingConfig as VoyageMultimodalEmbeddingConfig,
)
from .llms.infinity.embedding.transformation import (
InfinityEmbeddingConfig as InfinityEmbeddingConfig,
)

View file

@ -223,6 +223,7 @@ LLM_CONFIG_NAMES = (
"GenAIHubOrchestrationConfig",
"VoyageEmbeddingConfig",
"VoyageContextualEmbeddingConfig",
"VoyageMultimodalEmbeddingConfig",
"InfinityEmbeddingConfig",
"PerplexityEmbeddingConfig",
"AzureAIStudioConfig",
@ -903,6 +904,10 @@ _LLM_CONFIGS_IMPORT_MAP = {
".llms.voyage.embedding.transformation_contextual",
"VoyageContextualEmbeddingConfig",
),
"VoyageMultimodalEmbeddingConfig": (
".llms.voyage.embedding.transformation_multimodal",
"VoyageMultimodalEmbeddingConfig",
),
"InfinityEmbeddingConfig": (
".llms.infinity.embedding.transformation",
"InfinityEmbeddingConfig",

View file

@ -19,6 +19,7 @@ from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
A2AStreamingContext,
)
from litellm.a2a_protocol.providers.config_manager import A2AProviderConfigManager
from litellm.interactions.agents.utils import merge_agent_headers
# litellm_params key carrying the authenticated principal (hashed virtual key) so
# A2A provider configs can scope provider-side state (e.g. LangFlow session memory)
@ -48,6 +49,7 @@ class A2ACompletionBridgeHandler:
params: Dict[str, Any],
litellm_params: Dict[str, Any],
api_base: Optional[str] = None,
agent_extra_headers: Optional[Dict[str, str]] = None,
*,
_skip_a2a_provider_routing: bool = False,
) -> Dict[str, Any]:
@ -59,6 +61,8 @@ class A2ACompletionBridgeHandler:
params: A2A MessageSendParams containing the message
litellm_params: Agent's litellm_params (custom_llm_provider, model, etc.)
api_base: API base URL from agent_card_params
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
admin extra_headers) to forward on the upstream HTTP call.
Returns:
A2A SendMessageResponse dict
@ -80,6 +84,7 @@ class A2ACompletionBridgeHandler:
params=params,
api_base=api_base,
litellm_params=litellm_params,
agent_extra_headers=agent_extra_headers,
)
# Extract message from params
@ -106,7 +111,7 @@ class A2ACompletionBridgeHandler:
)
# Build completion params dict
completion_params = {
completion_params: Dict[str, Any] = {
"model": full_model,
"messages": openai_messages,
"api_base": api_base,
@ -128,6 +133,12 @@ class A2ACompletionBridgeHandler:
params=params,
)
if agent_extra_headers:
completion_params["extra_headers"] = merge_agent_headers(
dynamic_headers=agent_extra_headers,
static_headers=completion_params.get("extra_headers"),
)
# Call litellm.acompletion
response = await litellm.acompletion(**completion_params)
@ -149,6 +160,7 @@ class A2ACompletionBridgeHandler:
params: Dict[str, Any],
litellm_params: Dict[str, Any],
api_base: Optional[str] = None,
agent_extra_headers: Optional[Dict[str, str]] = None,
*,
_skip_a2a_provider_routing: bool = False,
) -> AsyncIterator[Dict[str, Any]]:
@ -166,6 +178,8 @@ class A2ACompletionBridgeHandler:
params: A2A MessageSendParams containing the message
litellm_params: Agent's litellm_params (custom_llm_provider, model, etc.)
api_base: API base URL from agent_card_params
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
admin extra_headers) to forward on the upstream HTTP call.
Yields:
A2A streaming response events
@ -187,6 +201,7 @@ class A2ACompletionBridgeHandler:
params=params,
api_base=api_base,
litellm_params=litellm_params,
agent_extra_headers=agent_extra_headers,
):
yield chunk
@ -222,7 +237,7 @@ class A2ACompletionBridgeHandler:
)
# Build completion params dict
completion_params = {
completion_params: Dict[str, Any] = {
"model": full_model,
"messages": openai_messages,
"api_base": api_base,
@ -244,6 +259,12 @@ class A2ACompletionBridgeHandler:
params=params,
)
if agent_extra_headers:
completion_params["extra_headers"] = merge_agent_headers(
dynamic_headers=agent_extra_headers,
static_headers=completion_params.get("extra_headers"),
)
# 1. Emit initial task event (kind: "task", status: "submitted")
task_event = A2ACompletionBridgeTransformation.create_task_event(ctx)
yield task_event
@ -305,6 +326,7 @@ async def handle_a2a_completion(
params: Dict[str, Any],
litellm_params: Dict[str, Any],
api_base: Optional[str] = None,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
"""Convenience function for non-streaming A2A completion."""
return await A2ACompletionBridgeHandler.handle_non_streaming(
@ -312,6 +334,7 @@ async def handle_a2a_completion(
params=params,
litellm_params=litellm_params,
api_base=api_base,
agent_extra_headers=agent_extra_headers,
)
@ -320,6 +343,7 @@ async def handle_a2a_completion_streaming(
params: Dict[str, Any],
litellm_params: Dict[str, Any],
api_base: Optional[str] = None,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> AsyncIterator[Dict[str, Any]]:
"""Convenience function for streaming A2A completion."""
async for chunk in A2ACompletionBridgeHandler.handle_streaming(
@ -327,5 +351,6 @@ async def handle_a2a_completion_streaming(
params=params,
litellm_params=litellm_params,
api_base=api_base,
agent_extra_headers=agent_extra_headers,
):
yield chunk

View file

@ -132,6 +132,7 @@ async def _send_message_via_completion_bridge(
custom_llm_provider: str,
api_base: Optional[str],
litellm_params: Dict[str, Any],
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> LiteLLMSendMessageResponse:
"""
Route a send_message through the LiteLLM completion bridge (e.g. LangGraph, Bedrock AgentCore).
@ -157,6 +158,7 @@ async def _send_message_via_completion_bridge(
params=params,
litellm_params=litellm_params,
api_base=api_base,
agent_extra_headers=agent_extra_headers,
)
return LiteLLMSendMessageResponse.from_dict(
@ -283,6 +285,7 @@ async def asend_message(
custom_llm_provider=custom_llm_provider,
api_base=api_base,
litellm_params=litellm_params,
agent_extra_headers=agent_extra_headers,
)
# Standard A2A client flow
@ -509,6 +512,7 @@ async def asend_message_streaming( # noqa: PLR0915
params=params,
litellm_params=litellm_params,
api_base=api_base,
agent_extra_headers=agent_extra_headers,
):
yield chunk
return

View file

@ -37,6 +37,7 @@ class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig):
request_id=request_id,
params=params,
litellm_params=litellm_params,
agent_extra_headers=kwargs.get("agent_extra_headers"),
)
async def handle_streaming(
@ -57,5 +58,6 @@ class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig):
request_id=request_id,
params=params,
litellm_params=litellm_params,
agent_extra_headers=kwargs.get("agent_extra_headers"),
):
yield chunk

View file

@ -6,7 +6,7 @@ completion bridge that would otherwise strip the envelope.
"""
import json
from typing import Any, AsyncIterator, Dict, cast
from typing import Any, AsyncIterator, Dict, Optional, cast
from litellm._logging import verbose_logger
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
@ -29,6 +29,7 @@ class BedrockAgentCoreA2AHandler:
request_id: str,
params: Dict[str, Any],
litellm_params: Dict[str, Any],
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
"""
Handle non-streaming A2A request to AgentCore.
@ -37,6 +38,8 @@ class BedrockAgentCoreA2AHandler:
request_id: A2A JSON-RPC request ID
params: A2A MessageSendParams containing the message
litellm_params: Agent's litellm_params (model, api_key, etc.)
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
admin extra_headers) to forward on the upstream HTTP call.
Returns:
A2A JSON-RPC response dict from the AgentCore agent
@ -47,6 +50,7 @@ class BedrockAgentCoreA2AHandler:
params=params,
litellm_params=litellm_params,
method="message/send",
agent_extra_headers=agent_extra_headers,
)
)
@ -77,6 +81,7 @@ class BedrockAgentCoreA2AHandler:
request_id: str,
params: Dict[str, Any],
litellm_params: Dict[str, Any],
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> AsyncIterator[Dict[str, Any]]:
"""
Handle streaming A2A request to AgentCore.
@ -85,6 +90,8 @@ class BedrockAgentCoreA2AHandler:
request_id: A2A JSON-RPC request ID
params: A2A MessageSendParams containing the message
litellm_params: Agent's litellm_params (model, api_key, etc.)
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
admin extra_headers) to forward on the upstream HTTP call.
Yields:
A2A streaming response events from the AgentCore agent
@ -96,6 +103,7 @@ class BedrockAgentCoreA2AHandler:
litellm_params=litellm_params,
method="message/send",
stream=True,
agent_extra_headers=agent_extra_headers,
)
)

View file

@ -6,11 +6,66 @@ and signs requests via AmazonAgentCoreConfig (SigV4 or JWT).
"""
import json
from typing import Any, AsyncIterator, Dict, Tuple
from typing import Any, AsyncIterator, Dict, Mapping, Optional, Tuple
from litellm._logging import verbose_logger
from litellm.llms.bedrock.chat.agentcore.transformation import AmazonAgentCoreConfig
# Reserved outbound header names that must never be sourced from per-request
# ``agent_extra_headers`` for AgentCore requests. ``agent_extra_headers`` carries
# values rewritten from the client-controlled ``x-a2a-{agent}-*`` convention, so
# allowing these would let any caller with access to the agent spoof the AWS
# request identity / SigV4 metadata by overwriting headers the proxy sets from
# trusted server-side config.
#
# The runtime headers (session / user id) are derived server-side from
# ``runtimeSessionId`` / ``runtimeUserId`` in the agent's ``litellm_params``;
# ``authorization`` is set by the AgentCore signer (JWT or SigV4); ``host`` and
# the ``x-amz-*`` family are owned by SigV4 itself.
_RESERVED_EXACT_HEADERS = frozenset(
{
"authorization",
"host",
}
)
_RESERVED_PREFIX_HEADERS: Tuple[str, ...] = (
"x-amzn-bedrock-agentcore-runtime-",
"x-amz-",
)
def _filter_reserved_headers(
agent_extra_headers: Optional[Mapping[str, str]],
) -> Optional[Dict[str, str]]:
"""
Strip reserved AWS / AgentCore headers from caller-supplied
``agent_extra_headers`` before they are merged into the signed request.
Returns ``None`` if the result is empty.
"""
if not agent_extra_headers:
return None
filtered: Dict[str, str] = {}
dropped: list = []
for k, v in agent_extra_headers.items():
k_lower = k.lower()
if k_lower in _RESERVED_EXACT_HEADERS or any(
k_lower.startswith(prefix) for prefix in _RESERVED_PREFIX_HEADERS
):
dropped.append(k)
continue
filtered[k] = v
if dropped:
verbose_logger.warning(
"BedrockAgentCore A2A: dropping reserved header(s) from "
"agent_extra_headers (not forwarded to AgentCore): %s",
sorted(dropped),
)
return filtered or None
class BedrockAgentCoreA2ATransformation:
"""
@ -27,6 +82,7 @@ class BedrockAgentCoreA2ATransformation:
litellm_params: Dict[str, Any],
method: str = "message/send",
stream: bool = False,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[str, dict, bytes]:
"""
Build the AgentCore URL, construct a JSON-RPC envelope, and sign the request.
@ -37,6 +93,15 @@ class BedrockAgentCoreA2ATransformation:
litellm_params: Agent's litellm_params (model, api_key, etc.)
method: JSON-RPC method name (default: "message/send")
stream: Whether this is a streaming request
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
admin extra_headers) to forward on the upstream HTTP call. Merged into
the headers dict before signing so SigV4 includes them in the signature.
Reserved AWS / AgentCore identity headers (``authorization``, ``host``,
``x-amzn-bedrock-agentcore-runtime-*``, ``x-amz-*``) are filtered out
here to prevent a caller-controlled ``x-a2a-{agent}-*`` header from
spoofing the AgentCore runtime user id or other SigV4 metadata. Use
``api_key`` / ``runtimeUserId`` / ``runtimeSessionId`` in litellm_params
(not ``agent_extra_headers``) to override those values.
Returns:
Tuple of (url, signed_headers, signed_body_bytes)
@ -85,6 +150,13 @@ class BedrockAgentCoreA2ATransformation:
if runtime_user_id:
headers["X-Amzn-Bedrock-AgentCore-Runtime-User-Id"] = runtime_user_id
# Merge per-request agent headers before signing so SigV4 covers them.
# Reserved headers are stripped first to prevent client-controlled values
# from spoofing the AgentCore runtime identity / SigV4 metadata.
safe_extra_headers = _filter_reserved_headers(agent_extra_headers)
if safe_extra_headers:
headers.update(safe_extra_headers)
# Sign the request (SigV4 or JWT depending on api_key presence)
signed_headers, signed_body = agentcore_config.sign_request(
headers=headers,

View file

@ -31,6 +31,7 @@ class PydanticAIProviderConfig(BaseA2AProviderConfig):
params=params,
api_base=api_base,
timeout=kwargs.get("timeout", 60.0),
agent_extra_headers=kwargs.get("agent_extra_headers"),
)
async def handle_streaming(
@ -50,5 +51,6 @@ class PydanticAIProviderConfig(BaseA2AProviderConfig):
timeout=kwargs.get("timeout", 60.0),
chunk_size=kwargs.get("chunk_size", 50),
delay_ms=kwargs.get("delay_ms", 10),
agent_extra_headers=kwargs.get("agent_extra_headers"),
):
yield chunk

View file

@ -28,6 +28,7 @@ class PydanticAIHandler:
params: Dict[str, Any],
api_base: Optional[str] = None,
timeout: float = 60.0,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
"""
Handle non-streaming request to Pydantic AI agent.
@ -37,6 +38,8 @@ class PydanticAIHandler:
params: A2A MessageSendParams containing the message
api_base: Base URL of the Pydantic AI agent
timeout: Request timeout in seconds
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
admin extra_headers) to forward on the upstream HTTP call.
Returns:
A2A SendMessageResponse dict
@ -51,6 +54,7 @@ class PydanticAIHandler:
request_id=request_id,
params=params,
timeout=timeout,
agent_extra_headers=agent_extra_headers,
)
return response_data
@ -63,6 +67,7 @@ class PydanticAIHandler:
timeout: float = 60.0,
chunk_size: int = 50,
delay_ms: int = 10,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> AsyncIterator[Dict[str, Any]]:
"""
Handle streaming request to Pydantic AI agent with fake streaming.
@ -78,6 +83,8 @@ class PydanticAIHandler:
timeout: Request timeout in seconds
chunk_size: Number of characters per chunk
delay_ms: Delay between chunks in milliseconds
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
admin extra_headers) to forward on the upstream HTTP call.
Yields:
A2A streaming response events
@ -94,6 +101,7 @@ class PydanticAIHandler:
request_id=request_id,
params=params,
timeout=timeout,
agent_extra_headers=agent_extra_headers,
)
# Convert raw task response to fake streaming chunks

View file

@ -6,7 +6,7 @@ This module provides fake streaming by converting non-streaming responses into s
"""
import asyncio
from typing import Any, AsyncIterator, Dict, cast
from typing import Any, AsyncIterator, Dict, Optional, cast
from uuid import uuid4
from litellm._logging import verbose_logger
@ -86,6 +86,7 @@ class PydanticAITransformation:
request_id: str,
max_attempts: int = 30,
poll_interval: float = 0.5,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
"""
Poll for task completion using tasks/get method.
@ -112,7 +113,10 @@ class PydanticAITransformation:
response = await client.post(
endpoint,
json=poll_request,
headers={"Content-Type": "application/json"},
headers={
**(agent_extra_headers or {}),
"Content-Type": "application/json",
},
)
response.raise_for_status()
poll_data = response.json()
@ -142,6 +146,7 @@ class PydanticAITransformation:
request_id: str,
params: Any,
timeout: float = 60.0,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
"""
Send a request to Pydantic AI agent and return the raw task response.
@ -189,7 +194,10 @@ class PydanticAITransformation:
response = await client.post(
endpoint,
json=a2a_request,
headers={"Content-Type": "application/json"},
headers={
**(agent_extra_headers or {}),
"Content-Type": "application/json",
},
)
response.raise_for_status()
response_data = response.json()
@ -211,6 +219,7 @@ class PydanticAITransformation:
endpoint=endpoint,
task_id=task_id,
request_id=request_id,
agent_extra_headers=agent_extra_headers,
)
verbose_logger.info(
@ -225,6 +234,7 @@ class PydanticAITransformation:
request_id: str,
params: Any,
timeout: float = 60.0,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
"""
Send a non-streaming A2A request to Pydantic AI agent and wait for completion.
@ -234,6 +244,7 @@ class PydanticAITransformation:
request_id: A2A JSON-RPC request ID
params: A2A MessageSendParams containing the message (dict or Pydantic model)
timeout: Request timeout in seconds
agent_extra_headers: Per-request headers to forward on the upstream HTTP call.
Returns:
Standard A2A non-streaming response format with message
@ -244,6 +255,7 @@ class PydanticAITransformation:
request_id=request_id,
params=params,
timeout=timeout,
agent_extra_headers=agent_extra_headers,
)
# Transform to standard A2A non-streaming format
@ -258,6 +270,7 @@ class PydanticAITransformation:
request_id: str,
params: Any,
timeout: float = 60.0,
agent_extra_headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
"""
Send a request to Pydantic AI agent and return the raw task response.
@ -269,6 +282,7 @@ class PydanticAITransformation:
request_id: A2A JSON-RPC request ID
params: A2A MessageSendParams containing the message
timeout: Request timeout in seconds
agent_extra_headers: Per-request headers to forward on the upstream HTTP call.
Returns:
Raw Pydantic AI task response (with history/artifacts)
@ -278,6 +292,7 @@ class PydanticAITransformation:
request_id=request_id,
params=params,
timeout=timeout,
agent_extra_headers=agent_extra_headers,
)
@staticmethod

View file

@ -75,7 +75,7 @@
"effort-2025-11-24": "effort-2025-11-24",
"fast-mode-2026-02-01": null,
"files-api-2025-04-14": null,
"fine-grained-tool-streaming-2025-05-14": null,
"fine-grained-tool-streaming-2025-05-14": "fine-grained-tool-streaming-2025-05-14",
"interleaved-thinking-2025-05-14": null,
"mcp-client-2025-11-20": null,
"mcp-client-2025-04-04": null,
@ -106,7 +106,7 @@
"effort-2025-11-24": "effort-2025-11-24",
"fast-mode-2026-02-01": null,
"files-api-2025-04-14": null,
"fine-grained-tool-streaming-2025-05-14": null,
"fine-grained-tool-streaming-2025-05-14": "fine-grained-tool-streaming-2025-05-14",
"interleaved-thinking-2025-05-14": null,
"mcp-client-2025-11-20": null,
"mcp-client-2025-04-04": null,

View file

@ -22,6 +22,7 @@ import litellm
from litellm._logging import print_verbose, verbose_logger
from litellm.constants import (
DEFAULT_REDIS_MAJOR_VERSION,
REDIS_CIRCUIT_BREAKER_ENABLED,
REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD,
REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT,
)
@ -114,15 +115,23 @@ class RedisCircuitBreaker:
OPEN = "open"
HALF_OPEN = "half_open"
def __init__(self, failure_threshold: int, recovery_timeout: int) -> None:
def __init__(
self,
failure_threshold: int,
recovery_timeout: int,
enabled: bool = True,
) -> None:
self.failure_threshold = failure_threshold
self.recovery_timeout = recovery_timeout
self.enabled = enabled
self._failure_count = 0
self._opened_at: Optional[float] = None
self._state = self.CLOSED
def is_open(self) -> bool:
"""Returns True if Redis calls should be skipped."""
if not self.enabled:
return False
if self._state == self.HALF_OPEN:
# Probe already in flight — fast-fail all concurrent requests.
# Only the one call that caused the OPEN→HALF_OPEN transition
@ -136,6 +145,8 @@ class RedisCircuitBreaker:
return False
def record_failure(self) -> None:
if not self.enabled:
return
self._failure_count += 1
self._opened_at = time.time()
if self._failure_count >= self.failure_threshold:
@ -149,6 +160,8 @@ class RedisCircuitBreaker:
self._state = self.OPEN
def record_success(self) -> None:
if not self.enabled:
return
if self._state == self.HALF_OPEN:
verbose_logger.info("Redis circuit breaker CLOSED — Redis recovered")
self._failure_count = 0
@ -243,6 +256,7 @@ class RedisCache(BaseCache):
self._circuit_breaker = RedisCircuitBreaker(
failure_threshold=REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD,
recovery_timeout=REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT,
enabled=REDIS_CIRCUIT_BREAKER_ENABLED,
)
self._setup_health_pings()

View file

@ -171,6 +171,8 @@ class ResponsesToCompletionBridgeHandler:
model_response = validated_kwargs["model_response"]
logging_obj = validated_kwargs["logging_obj"]
custom_llm_provider = validated_kwargs["custom_llm_provider"]
if kwargs.get("stream") is True and "stream" not in optional_params:
optional_params = {**optional_params, "stream": True}
request_data = self.transformation_handler.transform_request(
model=model,
@ -263,6 +265,8 @@ class ResponsesToCompletionBridgeHandler:
model_response = validated_kwargs["model_response"]
logging_obj = validated_kwargs["logging_obj"]
custom_llm_provider = validated_kwargs["custom_llm_provider"]
if kwargs.get("stream") is True and "stream" not in optional_params:
optional_params = {**optional_params, "stream": True}
try:
request_data = self.transformation_handler.transform_request(

View file

@ -398,6 +398,9 @@ REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD = int(
REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT = int(
os.getenv("REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT", 60)
)
REDIS_CIRCUIT_BREAKER_ENABLED = (
os.getenv("REDIS_CIRCUIT_BREAKER_ENABLED", "true").lower() == "true"
)
# Default Redis major version to assume when version cannot be determined
# Using 7 as it's the modern version that supports LPOP with count parameter
DEFAULT_REDIS_MAJOR_VERSION = int(os.getenv("DEFAULT_REDIS_MAJOR_VERSION", 7))
@ -418,6 +421,7 @@ REPLICATE_POLLING_DELAY_SECONDS = float(
DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS = int(
os.getenv("DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS", 4096)
)
DEFAULT_OCI_CHAT_MAX_TOKENS = 4096
TOGETHER_AI_4_B = int(os.getenv("TOGETHER_AI_4_B", 4))
TOGETHER_AI_8_B = int(os.getenv("TOGETHER_AI_8_B", 8))
TOGETHER_AI_21_B = int(os.getenv("TOGETHER_AI_21_B", 21))
@ -1480,6 +1484,7 @@ DB_SPEND_UPDATE_JOB_NAME = "db_spend_update_job"
DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME = "db_daily_tag_spend_update_job"
PROMETHEUS_EMIT_BUDGET_METRICS_JOB_NAME = "prometheus_emit_budget_metrics"
CLOUDZERO_EXPORT_USAGE_DATA_JOB_NAME = "cloudzero_export_usage_data"
MAVVRIK_FOCUS_EXPORT_JOB_NAME = "mavvrik_focus_export_usage_data"
CLOUDZERO_MAX_FETCHED_DATA_RECORDS = int(
os.getenv("CLOUDZERO_MAX_FETCHED_DATA_RECORDS", 50000)
)
@ -1494,6 +1499,10 @@ SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES = int(
SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS = float(
os.getenv("SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS", 0.5)
)
SPEND_LOG_PARTITION_INTERVAL = os.getenv("SPEND_LOG_PARTITION_INTERVAL", "day")
SPEND_LOG_PARTITION_PRECREATE_AHEAD = int(
os.getenv("SPEND_LOG_PARTITION_PRECREATE_AHEAD", 7)
)
SPEND_LOG_QUEUE_SIZE_THRESHOLD = int(os.getenv("SPEND_LOG_QUEUE_SIZE_THRESHOLD", 100))
SPEND_LOG_QUEUE_POLL_INTERVAL = float(os.getenv("SPEND_LOG_QUEUE_POLL_INTERVAL", 2.0))
SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE = int(

View file

@ -2488,6 +2488,11 @@ class RealtimeAPITokenUsageProcessor(BaseTokenUsageProcessor):
)
_TRANSCRIPTION_COMPLETED_EVENT_TYPE = (
"conversation.item.input_audio_transcription.completed"
)
def handle_realtime_stream_cost_calculation(
results: OpenAIRealtimeStreamList,
combined_usage_object: Usage,
@ -2533,4 +2538,99 @@ def handle_realtime_stream_cost_calculation(
break # exit if we find a valid model
total_cost = input_cost_per_token + output_cost_per_token
if any(r.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE for r in results):
total_cost += handle_realtime_transcription_cost_calculation(
results=results,
custom_llm_provider=custom_llm_provider,
litellm_model_name=litellm_model_name,
)
return total_cost
def handle_realtime_transcription_cost_calculation(
results: OpenAIRealtimeStreamList,
custom_llm_provider: str,
litellm_model_name: str,
) -> float:
"""
Cost for realtime transcription sessions (e.g. gpt-realtime-whisper).
Transcription sessions emit no `response.done` events; instead each
`conversation.item.input_audio_transcription.completed` event carries a
`usage` object billed by the ASR model. The usage is one of:
- {"type": "duration", "seconds": <float>} → priced via input_cost_per_second
- {"type": "tokens", "input_tokens": ...} → priced via input/audio token cost
"""
completed_events = [
cast(dict, result)
for result in results
if result.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE
]
if not completed_events:
return 0.0
model_name = (
_get_transcription_model_name_from_results(results) or litellm_model_name
)
try:
model_info = litellm.get_model_info(
model=model_name, custom_llm_provider=custom_llm_provider
)
except Exception:
model_info = None
total_cost = 0.0
for event in completed_events:
usage = event.get("usage") or {}
total_cost += _transcription_usage_cost(usage, model_info)
return total_cost
def _get_transcription_model_name_from_results(
results: OpenAIRealtimeStreamList,
) -> Optional[str]:
"""Resolve the ASR model from a transcription_session.* / session.* event."""
for result in results:
if result.get("type") in (
"transcription_session.created",
"transcription_session.updated",
"session.created",
"session.updated",
):
session = cast(dict, result).get("session", {}) or {}
transcription = (
(session.get("audio", {}) or {}).get("input", {}) or {}
).get("transcription", {}) or session.get("input_audio_transcription", {})
model = (transcription or {}).get("model") or session.get("model")
if model:
return model
return None
def _transcription_usage_cost(usage: dict, model_info: Optional[ModelInfo]) -> float:
if model_info is None:
return 0.0
usage_type = usage.get("type")
if usage_type == "duration":
seconds = usage.get("seconds") or 0.0
per_second = model_info.get("input_cost_per_second") or 0.0
return float(seconds) * float(per_second)
if usage_type == "tokens":
input_token_details = usage.get("input_token_details") or {}
audio_tokens = input_token_details.get("audio_tokens") or 0
text_tokens = input_token_details.get("text_tokens") or 0
output_tokens = usage.get("output_tokens") or 0
audio_cost = float(audio_tokens) * float(
model_info.get("input_cost_per_audio_token")
or model_info.get("input_cost_per_token")
or 0.0
)
text_cost = float(text_tokens) * float(
model_info.get("input_cost_per_token") or 0.0
)
output_cost = float(output_tokens) * float(
model_info.get("output_cost_per_token") or 0.0
)
return audio_cost + text_cost + output_cost
return 0.0

View file

@ -18,6 +18,42 @@ else:
GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ = PassThroughEndpointLogging()
def _encode_google_genai_sse_event(event_lines: List[str]) -> bytes:
return ("\n".join(event_lines) + "\n\n").encode("utf-8")
def _next_google_genai_sse_chunk(line_iter) -> bytes:
event_lines: List[str] = []
while True:
try:
line = next(line_iter)
except StopIteration:
if event_lines:
return _encode_google_genai_sse_event(event_lines)
raise
if line == "":
if event_lines:
return _encode_google_genai_sse_event(event_lines)
continue
event_lines.append(line)
async def _anext_google_genai_sse_chunk(line_iter) -> bytes:
event_lines: List[str] = []
while True:
try:
line = await line_iter.__anext__()
except StopAsyncIteration:
if event_lines:
return _encode_google_genai_sse_event(event_lines)
raise
if line == "":
if event_lines:
return _encode_google_genai_sse_event(event_lines)
continue
event_lines.append(line)
class BaseGoogleGenAIGenerateContentStreamingIterator:
"""
Base class for Google GenAI Generate Content streaming iterators that provides common logic
@ -91,18 +127,17 @@ class GoogleGenAIGenerateContentStreamingIterator(
self.generate_content_provider_config = generate_content_provider_config
self.litellm_metadata = litellm_metadata
self.custom_llm_provider = custom_llm_provider
# Store the iterator once to avoid multiple stream consumption
self.stream_iterator = response.iter_bytes()
# Gemini streamGenerateContent uses SSE line framing; iter_lines keeps
# large inlineData payloads (e.g. image/jpeg) intact within one event.
self.stream_iterator = response.iter_lines()
def __iter__(self):
return self
def __next__(self):
try:
# Get the next chunk from the stored iterator
chunk = next(self.stream_iterator)
chunk = _next_google_genai_sse_chunk(self.stream_iterator)
self.collected_chunks.append(chunk)
# Just yield raw bytes
return chunk
except StopIteration:
raise StopIteration
@ -147,18 +182,17 @@ class AsyncGoogleGenAIGenerateContentStreamingIterator(
self.generate_content_provider_config = generate_content_provider_config
self.litellm_metadata = litellm_metadata
self.custom_llm_provider = custom_llm_provider
# Store the async iterator once to avoid multiple stream consumption
self.stream_iterator = response.aiter_bytes()
# Gemini streamGenerateContent uses SSE line framing; aiter_lines keeps
# large inlineData payloads (e.g. image/jpeg) intact within one event.
self.stream_iterator = response.aiter_lines()
def __aiter__(self):
return self
async def __anext__(self):
try:
# Get the next chunk from the stored async iterator
chunk = await self.stream_iterator.__anext__()
chunk = await _anext_google_genai_sse_chunk(self.stream_iterator)
self.collected_chunks.append(chunk)
# Just yield raw bytes
return chunk
except StopAsyncIteration:
await self._handle_async_streaming_logging()

View file

@ -1,7 +1,7 @@
from abc import ABC, abstractmethod
from typing import Literal
from litellm.proxy._types import CallInfo
from litellm.proxy._types import CallInfo, Litellm_EntityType
class BaseBudgetAlertType(ABC):
@ -31,6 +31,8 @@ class SoftBudgetAlert(BaseBudgetAlertType):
return "Soft Budget Crossed: "
def get_id(self, user_info: CallInfo) -> str:
if user_info.event_group == Litellm_EntityType.TEAM:
return user_info.team_id or "default_id"
return user_info.token or "default_id"

View file

@ -8,6 +8,7 @@ Notes:
"""
import asyncio
import time
from typing import TYPE_CHECKING, Any, Optional
import litellm
@ -36,11 +37,15 @@ class AlertingHangingRequestCheck:
slack_alerting_object: SlackAlerting,
):
self.slack_alerting_object = slack_alerting_object
# checks run every alerting_threshold / 2 seconds, so entries must
# stay cached for at least 1.5x the threshold to guarantee a check
# happens after they cross it
self.hanging_request_cache_ttl = int(
self.slack_alerting_object.alerting_threshold * 1.5
+ HANGING_ALERT_BUFFER_TIME_SECONDS
)
self.hanging_request_cache = InMemoryCache(
default_ttl=int(
self.slack_alerting_object.alerting_threshold
+ HANGING_ALERT_BUFFER_TIME_SECONDS
),
default_ttl=self.hanging_request_cache_ttl,
)
async def add_request_to_hanging_request_check(
@ -76,10 +81,7 @@ class AlertingHangingRequestCheck:
await self.hanging_request_cache.async_set_cache(
key=hanging_request_data.request_id,
value=hanging_request_data,
ttl=int(
self.slack_alerting_object.alerting_threshold
+ HANGING_ALERT_BUFFER_TIME_SECONDS
),
ttl=self.hanging_request_cache_ttl,
)
return
@ -111,6 +113,9 @@ class AlertingHangingRequestCheck:
if hanging_request_data is None:
continue
if hanging_request_data.alerted:
continue
request_status = (
await proxy_logging_obj.internal_usage_cache.async_get_cache(
key="request_status:{}".format(hanging_request_data.request_id),
@ -127,12 +132,21 @@ class AlertingHangingRequestCheck:
)
continue
request_age_seconds = time.time() - hanging_request_data.created_at
if request_age_seconds < self.slack_alerting_object.alerting_threshold:
# in-flight but below the alerting threshold; keep it cached
# so a later check can alert if it never completes
continue
################
# Send the Alert on Slack
################
await self.send_hanging_request_alert(
hanging_request_data=hanging_request_data
)
# flag so the entry is skipped on later ticks; one alert per hang,
# with the existing TTL still handling cleanup
hanging_request_data.alerted = True
return

View file

@ -290,6 +290,21 @@
},
"description": "Langsmith Logging Integration"
},
{
"id": "newrelic",
"displayName": "New Relic",
"logo": "newrelic.png",
"supports_key_team_logging": false,
"dynamic_params": {
"NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED": {
"type": "text",
"ui_name": "Record AI Content (default: true)",
"description": "Whether to record AI message content. Set to false to disable.",
"required": false
}
},
"description": "New Relic AI Monitoring Integration"
},
{
"id": "openmeter",
"displayName": "OpenMeter",

View file

@ -92,12 +92,26 @@ class DataDogLogger(
# Class variables or attributes
def __init__(
self,
dd_api_key: Optional[str] = None,
dd_site: Optional[str] = None,
dd_agent_host: Optional[str] = None,
dd_agent_port: Optional[str] = None,
allow_env_credentials: bool = True,
**kwargs,
):
"""
Initializes the datadog logger, checks if the correct env variables are set
Required environment variables (Direct API):
Args:
dd_api_key: Datadog API key. Falls back to DD_API_KEY env var when allow_env_credentials is True.
dd_site: Datadog site (e.g. "us5.datadoghq.com"). Falls back to DD_SITE env var.
dd_agent_host: Hostname or IP of DataDog agent. Falls back to LITELLM_DD_AGENT_HOST env var.
dd_agent_port: Port of DataDog agent (default: 10518). Falls back to LITELLM_DD_AGENT_PORT env var.
allow_env_credentials: When False, the API key is never read from DD_API_KEY env var. Set to
False for team/key-scoped loggers whose destination (dd_agent_host/dd_site) is caller-supplied,
so the proxy's global DD_API_KEY is never sent to an untrusted host.
Required environment variables (Direct API) when kwargs not provided:
`DD_API_KEY` - your datadog api key
`DD_SITE` - your datadog site, example = `"us5.datadoghq.com"`
@ -130,12 +144,21 @@ class DataDogLogger(
)
# Configure DataDog endpoint (Agent or Direct API)
# Use LITELLM_DD_AGENT_HOST to avoid conflicts with ddtrace's DD_AGENT_HOST
dd_agent_host = os.getenv("LITELLM_DD_AGENT_HOST")
if dd_agent_host:
self._configure_dd_agent(dd_agent_host=dd_agent_host)
# Prefer explicit kwargs, then fall back to env vars
resolved_agent_host = dd_agent_host or os.getenv("LITELLM_DD_AGENT_HOST")
if resolved_agent_host:
self._configure_dd_agent(
dd_agent_host=resolved_agent_host,
dd_agent_port=dd_agent_port,
dd_api_key=dd_api_key,
allow_env_credentials=allow_env_credentials,
)
else:
self._configure_dd_direct_api()
self._configure_dd_direct_api(
dd_api_key=dd_api_key,
dd_site=dd_site,
allow_env_credentials=allow_env_credentials,
)
# Optional override for testing
dd_base_url = get_datadog_base_url_from_env()
@ -172,34 +195,60 @@ class DataDogLogger(
).model_dump()
return dict_datadog_params
def _configure_dd_agent(self, dd_agent_host: str) -> None:
def _configure_dd_agent(
self,
dd_agent_host: str,
dd_agent_port: Optional[str] = None,
dd_api_key: Optional[str] = None,
allow_env_credentials: bool = True,
) -> None:
"""
Configure DataDog Agent for log forwarding
Args:
dd_agent_host: Hostname or IP of DataDog agent
dd_agent_port: Port of DataDog agent. Falls back to LITELLM_DD_AGENT_PORT env var (default: 10518).
dd_api_key: Datadog API key. Falls back to DD_API_KEY env var when allow_env_credentials is True. Optional when using agent.
allow_env_credentials: When False, never read the API key from DD_API_KEY env var.
"""
dd_agent_port = os.getenv(
resolved_port = dd_agent_port or os.getenv(
"LITELLM_DD_AGENT_PORT", "10518"
) # default port for logs
self.intake_url = f"http://{dd_agent_host}:{dd_agent_port}/api/v2/logs"
self.DD_API_KEY = os.getenv("DD_API_KEY") # Optional when using agent
self.intake_url = f"http://{dd_agent_host}:{resolved_port}/api/v2/logs"
self.DD_API_KEY = dd_api_key or (
os.getenv("DD_API_KEY") if allow_env_credentials else None
) # Optional when using agent
verbose_logger.debug(f"Datadog: Using DD Agent at {self.intake_url}")
def _configure_dd_direct_api(self) -> None:
def _configure_dd_direct_api(
self,
dd_api_key: Optional[str] = None,
dd_site: Optional[str] = None,
allow_env_credentials: bool = True,
) -> None:
"""
Configure direct DataDog API connection
Args:
dd_api_key: Datadog API key. Falls back to DD_API_KEY env var when allow_env_credentials is True.
dd_site: Datadog site. Falls back to DD_SITE env var.
allow_env_credentials: When False, never read the API key from DD_API_KEY env var.
Raises:
Exception: If required environment variables are not set
Exception: If required credentials are not provided via args or env vars
"""
if os.getenv("DD_API_KEY", None) is None:
resolved_api_key = dd_api_key or (
os.getenv("DD_API_KEY") if allow_env_credentials else None
)
resolved_site = dd_site or os.getenv("DD_SITE")
if resolved_api_key is None:
raise Exception("DD_API_KEY is not set, set 'DD_API_KEY=<>")
if os.getenv("DD_SITE", None) is None:
if resolved_site is None:
raise Exception("DD_SITE is not set in .env, set 'DD_SITE=<>")
self.DD_API_KEY = os.getenv("DD_API_KEY")
self.intake_url = f"https://http-intake.logs.{os.getenv('DD_SITE')}/api/v2/logs"
self.DD_API_KEY = resolved_api_key
self.intake_url = f"https://http-intake.logs.{resolved_site}/api/v2/logs"
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
"""

View file

@ -0,0 +1,124 @@
"""
DataDog Team Handler
Used to get the DataDogLogger for a given request.
Handles Key/Team Based Datadog Logging, following the same pattern as LangFuseHandler.
"""
from typing import TYPE_CHECKING, Any, Dict, Optional, TypedDict
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.litellm_logging import StandardCallbackDynamicParams
from .datadog import DataDogLogger
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import DynamicLoggingCache
else:
DynamicLoggingCache = Any
class DatadogLoggingConfig(TypedDict):
dd_api_key: Optional[str]
dd_site: Optional[str]
dd_agent_host: Optional[str]
dd_agent_port: Optional[str]
class DataDogHandler:
@staticmethod
def get_datadog_logger_for_request(
standard_callback_dynamic_params: StandardCallbackDynamicParams,
in_memory_dynamic_logger_cache: DynamicLoggingCache,
) -> DataDogLogger:
"""
Get a team-scoped DataDogLogger for a given request.
Resolves and caches per-team DataDogLogger instances using DynamicLoggingCache,
keyed by the team's DD credentials. Each unique set of credentials gets its own
logger instance with its own batch/flush loop.
Note: This handler is only called when team-scoped DD credentials are present.
The global (env-var based) DataDogLogger is managed separately by
_init_custom_logger_compatible_class via _in_memory_loggers.
"""
_credentials = DataDogHandler.get_dynamic_datadog_logging_config(
standard_callback_dynamic_params=standard_callback_dynamic_params,
)
credentials_dict = dict(_credentials)
# check if datadog logger is already cached
temp_datadog_logger = in_memory_dynamic_logger_cache.get_cache(
credentials=credentials_dict, service_name="datadog"
)
# if not cached, create a new datadog logger and cache it
if temp_datadog_logger is None:
temp_datadog_logger = (
DataDogHandler._create_datadog_logger_from_credentials(
credentials=credentials_dict,
in_memory_dynamic_logger_cache=in_memory_dynamic_logger_cache,
)
)
return temp_datadog_logger
@staticmethod
def _create_datadog_logger_from_credentials(
credentials: Dict,
in_memory_dynamic_logger_cache: DynamicLoggingCache,
) -> DataDogLogger:
"""
Create a DataDogLogger from the credentials and cache it.
"""
# When the destination is caller-supplied (dd_agent_host/dd_site), never fall back to the
# proxy's DD_API_KEY env var, otherwise it would be sent to a team-controlled host.
allow_env_credentials = (
credentials.get("dd_agent_host") is None
and credentials.get("dd_site") is None
)
datadog_logger = DataDogLogger(
dd_api_key=credentials.get("dd_api_key"),
dd_site=credentials.get("dd_site"),
dd_agent_host=credentials.get("dd_agent_host"),
dd_agent_port=credentials.get("dd_agent_port"),
allow_env_credentials=allow_env_credentials,
)
in_memory_dynamic_logger_cache.set_cache(
credentials=credentials,
service_name="datadog",
logging_obj=datadog_logger,
)
verbose_logger.debug(
"Datadog: Created and cached new DataDogLogger for team-scoped credentials"
)
return datadog_logger
@staticmethod
def get_dynamic_datadog_logging_config(
standard_callback_dynamic_params: StandardCallbackDynamicParams,
) -> DatadogLoggingConfig:
"""
Get the Datadog logging config for a given request from dynamic params.
"""
return DatadogLoggingConfig(
dd_api_key=standard_callback_dynamic_params.get("dd_api_key"),
dd_site=standard_callback_dynamic_params.get("dd_site"),
dd_agent_host=standard_callback_dynamic_params.get("dd_agent_host"),
dd_agent_port=standard_callback_dynamic_params.get("dd_agent_port"),
)
@staticmethod
def _dynamic_datadog_credentials_are_passed(
standard_callback_dynamic_params: StandardCallbackDynamicParams,
) -> bool:
"""
Check if dynamic Datadog credentials are passed in standard_callback_dynamic_params.
"""
if (
standard_callback_dynamic_params.get("dd_api_key") is not None
or standard_callback_dynamic_params.get("dd_site") is not None
or standard_callback_dynamic_params.get("dd_agent_host") is not None
):
return True
return False

View file

@ -4,6 +4,7 @@ from .base import FocusDestination, FocusTimeWindow
from .factory import FocusDestinationFactory
from .gcs_destination import FocusGCSDestination
from .s3_destination import FocusS3Destination
from .mavvrik_destination import FocusMavvrikDestination
from .vantage_destination import FocusVantageDestination
__all__ = [
@ -12,5 +13,6 @@ __all__ = [
"FocusGCSDestination",
"FocusTimeWindow",
"FocusS3Destination",
"FocusMavvrikDestination",
"FocusVantageDestination",
]

View file

@ -8,6 +8,7 @@ from typing import Any, Dict, Optional
from .base import FocusDestination
from .gcs_destination import FocusGCSDestination
from .s3_destination import FocusS3Destination
from .mavvrik_destination import FocusMavvrikDestination
from .vantage_destination import FocusVantageDestination
@ -32,6 +33,8 @@ class FocusDestinationFactory:
return FocusVantageDestination(prefix=prefix, config=normalized_config)
if provider_lower == "gcs":
return FocusGCSDestination(prefix=prefix, config=normalized_config)
if provider_lower == "mavvrik":
return FocusMavvrikDestination(prefix=prefix, config=normalized_config)
raise NotImplementedError(
f"Provider '{provider}' not supported for Focus export"
)
@ -87,6 +90,15 @@ class FocusDestinationFactory:
"FOCUS_GCS_BUCKET_NAME must be provided for GCS exports"
)
return {k: v for k, v in resolved.items() if v is not None}
if provider == "mavvrik":
resolved = {
"api_key": overrides.get("api_key") or os.getenv("MAVVRIK_API_KEY"),
"api_endpoint": overrides.get("api_endpoint")
or os.getenv("MAVVRIK_API_ENDPOINT"),
"connection_id": overrides.get("connection_id")
or os.getenv("MAVVRIK_CONNECTION_ID"),
}
return {k: v for k, v in resolved.items() if v is not None}
raise NotImplementedError(
f"Provider '{provider}' not supported for Focus export configuration"
)

View file

@ -0,0 +1,345 @@
"""Mavvrik GCS destination for FOCUS export.
Flow:
1. GET /metrics/agent/ai/{connection_id}/upload-url → GCS signed URL
2. PUT <signed_url> with CSV content
"""
from __future__ import annotations
import gzip
from typing import Any, Optional
from urllib.parse import urlparse
from litellm._logging import verbose_logger
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_async_httpx_client,
httpxSpecialProvider,
)
from .base import FocusDestination, FocusTimeWindow
_MAVVRIK_ALLOWED_SUFFIXES = (".mavvrik.dev", ".mavvrik.ai", ".mavvrik.app")
# GCS requires intermediate chunks to be a multiple of 256 KB.
# 8 MB gives a good balance between round-trips and memory pressure.
_GCS_CHUNK_SIZE = 8 * 1024 * 1024 # 8 MB
def _validate_api_endpoint(api_endpoint: str) -> None:
if not api_endpoint.startswith("https://"):
raise ValueError("MAVVRIK_API_ENDPOINT must be an HTTPS URL")
hostname = (urlparse(api_endpoint).hostname or "").lower()
if not any(hostname.endswith(suffix) for suffix in _MAVVRIK_ALLOWED_SUFFIXES):
raise ValueError(
"MAVVRIK_API_ENDPOINT host must be a Mavvrik domain "
"(e.g. https://api.mavvrik.dev/<tenant_id>)"
)
def _validate_gcs_url(url: str, label: str) -> None:
parsed = urlparse(url)
if parsed.scheme != "https":
raise ValueError(
f"Mavvrik FOCUS destination: {label} must be HTTPS, got scheme '{parsed.scheme}'"
)
hostname = (parsed.hostname or "").lower()
if not (
hostname == "storage.googleapis.com"
or hostname.endswith(".storage.googleapis.com")
):
raise ValueError(
f"Mavvrik FOCUS destination: {label} must be a GCS endpoint "
f"(storage.googleapis.com), got '{hostname}'"
)
class FocusMavvrikDestination(FocusDestination):
"""Upload FOCUS CSV exports to Mavvrik via GCS signed URL."""
def __init__(
self,
*,
prefix: str,
config: Optional[dict[str, Any]] = None,
) -> None:
config = config or {}
api_key = config.get("api_key")
api_endpoint = config.get("api_endpoint")
connection_id = config.get("connection_id")
if not api_key:
raise ValueError(
"MAVVRIK_API_KEY must be provided for Mavvrik FOCUS destination "
"(set MAVVRIK_API_KEY env var or pass in destination_config)"
)
if not api_endpoint:
raise ValueError(
"MAVVRIK_API_ENDPOINT must be provided for Mavvrik FOCUS destination "
"(set MAVVRIK_API_ENDPOINT env var or pass in destination_config)"
)
if not connection_id:
raise ValueError(
"MAVVRIK_CONNECTION_ID must be provided for Mavvrik FOCUS destination "
"(set MAVVRIK_CONNECTION_ID env var or pass in destination_config)"
)
_validate_api_endpoint(api_endpoint)
self.api_key = api_key
self.api_endpoint = api_endpoint.rstrip("/")
self.connection_id = connection_id
self.prefix = prefix
self._http: AsyncHTTPHandler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.LoggingCallback
)
self._registered = False
@property
def _agent_url(self) -> str:
return f"{self.api_endpoint}/metrics/agent/ai/{self.connection_id}"
@property
def _upload_url_endpoint(self) -> str:
return f"{self.api_endpoint}/metrics/agent/ai/{self.connection_id}/upload-url"
@property
def _auth_headers(self) -> dict[str, str]:
return {"Content-Type": "application/json", "x-api-key": self.api_key}
async def _ensure_registered(self) -> Optional[int]:
"""POST agent endpoint to register/initialize the connector (once per instance).
Returns metricsMarker from the Mavvrik response — the last date index
Mavvrik has successfully processed. Used by the logger to catch up any
dates that were missed due to previous export failures.
Returns None if the connector was already registered (cached).
"""
if self._registered:
return None
resp = await self._http.client.request(
method="POST",
url=self._agent_url,
headers=self._auth_headers,
json={"name": self.connection_id},
timeout=30.0,
)
if resp.status_code == 410:
# Connector has been disconnected in Mavvrik — reset flag so next
# delivery attempt re-registers after it becomes active again.
self._registered = False
raise RuntimeError(
"Mavvrik FOCUS destination: connector is disconnected (410). "
"Re-enable the connection in the Mavvrik dashboard."
)
if resp.status_code >= 400:
raise RuntimeError(
f"Mavvrik FOCUS destination: register failed "
f"({resp.status_code}): {resp.text[:200]}"
)
self._registered = True
metrics_marker = resp.json().get("metricsMarker", 0)
verbose_logger.debug(
"Mavvrik FOCUS destination: connector registered (metricsMarker=%s)",
metrics_marker,
)
return metrics_marker
async def _get_signed_url(self, date_str: str) -> str:
"""GET upload-url endpoint → GCS signed URL for the given date."""
params = {"name": date_str, "type": "metrics", "datetime": date_str}
resp = await self._http.client.request(
method="GET",
url=self._upload_url_endpoint,
headers=self._auth_headers,
params=params,
timeout=30.0,
)
if resp.status_code >= 400:
raise RuntimeError(
f"Mavvrik FOCUS destination: failed to get signed URL "
f"({resp.status_code}): {resp.text[:200]}"
)
signed_url = resp.json().get("url")
if not signed_url:
raise RuntimeError(
f"Mavvrik FOCUS destination: response missing 'url' field: {resp.json()}"
)
_validate_gcs_url(signed_url, "signed URL")
verbose_logger.debug(
"Mavvrik FOCUS destination: got signed URL for date %s", date_str
)
return signed_url
async def _upload_to_gcs(self, signed_url: str, content: bytes) -> None:
"""Upload gzip-compressed CSV to GCS via chunked resumable upload.
The full CSV is gzip-compressed first, then uploaded in _GCS_CHUNK_SIZE
chunks using the GCS resumable upload protocol. GCS assembles the chunks
server-side into a single complete object — the bucket receives one file
regardless of how many chunks were sent.
Intermediate chunks: Content-Range: bytes X-Y/* → expect 308
Final chunk: Content-Range: bytes X-Y/T → expect 200/201
This handles exports larger than available memory for a single PUT while
keeping the destination code self-contained (no changes to the FOCUS
pipeline upstream).
"""
gzip_bytes = gzip.compress(content)
total = len(gzip_bytes)
# Step 1: initiate resumable upload session
metadata = b'{"contentEncoding":"gzip","contentDisposition":"attachment"}'
init_resp = await self._http.client.request(
method="POST",
url=signed_url,
headers={
"Content-Type": "application/gzip",
"x-goog-resumable": "start",
},
content=metadata,
timeout=30.0,
)
if init_resp.status_code not in (200, 201):
raise RuntimeError(
f"Mavvrik FOCUS destination: GCS session init failed "
f"({init_resp.status_code}): {init_resp.text[:400]}"
)
session_uri = init_resp.headers.get("Location")
if not session_uri:
raise RuntimeError(
"Mavvrik FOCUS destination: GCS session init missing Location header"
)
_validate_gcs_url(session_uri, "session URI")
verbose_logger.debug(
"Mavvrik FOCUS destination: GCS session started, uploading %d gzip bytes "
"in %d chunk(s)",
total,
max(1, -(-total // _GCS_CHUNK_SIZE)), # ceiling division
)
# Step 2: upload in chunks; cancel session on any failure to avoid
# lingering GCS sessions (they stay open for ~1 week otherwise).
offset = 0
try:
while offset < total:
chunk = gzip_bytes[offset : offset + _GCS_CHUNK_SIZE]
chunk_end = offset + len(chunk) - 1
is_final = (offset + len(chunk)) >= total
content_range = (
f"bytes {offset}-{chunk_end}/{total}"
if is_final
else f"bytes {offset}-{chunk_end}/*"
)
expected_statuses = {200, 201} if is_final else {308}
resp = await self._http.client.request(
method="PUT",
url=session_uri,
headers={
"Content-Type": "application/gzip",
"Content-Range": content_range,
},
content=chunk,
timeout=120.0,
)
if resp.status_code not in expected_statuses:
raise RuntimeError(
f"Mavvrik FOCUS destination: GCS chunk upload failed "
f"(chunk offset={offset}, expected={expected_statuses}, "
f"got={resp.status_code}): {resp.text[:400]}"
)
offset += len(chunk)
verbose_logger.debug(
"Mavvrik FOCUS destination: uploaded chunk offset=%d/%d",
offset,
total,
)
except Exception:
# Cancel the open GCS session so it doesn't linger for up to 1 week.
try:
await self._http.client.request(
method="DELETE", url=session_uri, timeout=10.0
)
verbose_logger.debug(
"Mavvrik FOCUS destination: cancelled GCS session after error"
)
except Exception:
pass
raise
async def get_metrics_marker(self) -> Optional[int]:
"""Register with Mavvrik and return the current metricsMarker.
The metricsMarker is a Unix timestamp (seconds) representing the last
date Mavvrik has successfully ingested. Called on every scheduled run
so the logger can detect and catch up any dates missed due to previous
export failures.
Always calls the Mavvrik register API — unlike deliver() which skips
registration once _registered is True, catch-up requires a fresh
marker value on every run.
"""
resp = await self._http.client.request(
method="POST",
url=self._agent_url,
headers=self._auth_headers,
json={"name": self.connection_id},
timeout=30.0,
)
if resp.status_code == 410:
self._registered = False
raise RuntimeError(
"Mavvrik FOCUS destination: connector is disconnected (410). "
"Re-enable the connection in the Mavvrik dashboard."
)
if resp.status_code >= 400:
raise RuntimeError(
f"Mavvrik FOCUS destination: register failed "
f"({resp.status_code}): {resp.text[:200]}"
)
self._registered = True
metrics_marker = resp.json().get("metricsMarker", 0)
verbose_logger.debug(
"Mavvrik FOCUS destination: got metricsMarker=%s", metrics_marker
)
return metrics_marker
async def deliver(
self,
*,
content: bytes,
time_window: FocusTimeWindow,
filename: str,
) -> None:
"""Upload FOCUS CSV to Mavvrik via GCS signed URL.
Uses the start date of the time window as the object date key.
"""
if not content:
verbose_logger.debug(
"Mavvrik FOCUS destination: empty content, skipping upload"
)
return
date_str = time_window.start_time.strftime("%Y-%m-%d")
verbose_logger.debug(
"Mavvrik FOCUS destination: uploading %d bytes for date=%s (%s)",
len(content),
date_str,
filename,
)
await self._ensure_registered()
signed_url = await self._get_signed_url(date_str)
await self._upload_to_gcs(signed_url, content)
verbose_logger.debug(
"Mavvrik FOCUS destination: upload complete for date=%s", date_str
)

View file

@ -43,6 +43,7 @@ class LangfuseOtelLogger(OpenTelemetry):
"""
_utils.set_attributes(span, kwargs, response_obj, LangfuseLLMObsOTELAttributes)
span.set_attribute("langfuse.observation.type", "generation")
#########################################################
# Set Langfuse specific attributes

View file

@ -0,0 +1,272 @@
"""MavvrikFocusLogger — FOCUS-based Mavvrik export logger.
Usage in config.yaml:
litellm_settings:
callbacks: ["mavvrik"]
Required env vars:
MAVVRIK_API_KEY
MAVVRIK_API_ENDPOINT
MAVVRIK_CONNECTION_ID
Optional env vars:
MAVVRIK_FOCUS_MAX_ROWS — row cap per export window (default: 500000)
Only daily frequency is supported. The Mavvrik ingestion protocol stores one
file per calendar date (metrics/YYYY-MM-DD). Hourly or interval exports would
overwrite each other within the same day, producing incomplete data.
"""
from __future__ import annotations
import os
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, List, Optional
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import MAVVRIK_FOCUS_EXPORT_JOB_NAME
from litellm.integrations.focus.destinations.base import FocusTimeWindow
from litellm.integrations.focus.focus_logger import FocusLogger
if TYPE_CHECKING:
from apscheduler.schedulers.asyncio import AsyncIOScheduler
else:
AsyncIOScheduler = Any
def _parse_metrics_marker(
marker: Optional[object],
) -> Optional[datetime]:
"""Parse metricsMarker from Mavvrik register response into a UTC datetime.
Handles both formats Mavvrik may return:
- Unix timestamp (int/float): e.g. 1749340800
- ISO date string: e.g. "2026-06-09" or "2026-06-09T00:00:00Z"
Returns None for falsy values (0, None, empty string) which indicate
no data has been ingested yet.
"""
if not marker:
return None
try:
if isinstance(marker, (int, float)):
return datetime.fromtimestamp(float(marker), tz=timezone.utc).replace(
hour=0, minute=0, second=0, microsecond=0
)
if isinstance(marker, str):
marker = marker.strip()
if not marker:
return None
# Try ISO date first (YYYY-MM-DD), then full ISO datetime
for fmt in ("%Y-%m-%d", "%Y-%m-%dT%H:%M:%SZ", "%Y-%m-%dT%H:%M:%S"):
try:
return datetime.strptime(marker, fmt).replace(tzinfo=timezone.utc)
except ValueError:
continue
except Exception:
pass
verbose_proxy_logger.warning(
"Mavvrik FOCUS: could not parse metricsMarker %r — skipping catch-up", marker
)
return None
class MavvrikFocusLogger(FocusLogger):
"""FOCUS-based export logger that routes to the Mavvrik destination."""
def __init__(self, **kwargs: Any) -> None:
frequency = os.getenv("MAVVRIK_FOCUS_FREQUENCY", "daily").lower()
if frequency != "daily":
raise ValueError(
f"MAVVRIK_FOCUS_FREQUENCY='{frequency}' is not supported. "
"Only 'daily' is allowed -- the Mavvrik ingestion protocol stores one "
"file per calendar date (metrics/YYYY-MM-DD). Hourly or interval "
"exports would overwrite each other within the same day."
)
super().__init__(
provider="mavvrik",
export_format="csv",
frequency="daily",
prefix="mavvrik_focus_exports",
destination_config={
"api_key": os.getenv("MAVVRIK_API_KEY"),
"api_endpoint": os.getenv("MAVVRIK_API_ENDPOINT"),
"connection_id": os.getenv("MAVVRIK_CONNECTION_ID"),
},
**kwargs,
)
raw = os.getenv("MAVVRIK_FOCUS_MAX_ROWS")
self._max_rows: Optional[int] = int(raw) if raw else 500_000
async def _export_window(
self,
*,
window: FocusTimeWindow,
limit: Optional[int],
) -> None:
"""Export with Mavvrik row cap applied when no explicit limit is passed."""
effective_limit = limit if limit is not None else self._max_rows
engine = self._ensure_engine()
data = await engine._database.get_usage_data(
limit=effective_limit,
start_time_utc=window.start_time,
end_time_utc=window.end_time,
)
if effective_limit is not None and len(data) >= effective_limit:
verbose_proxy_logger.warning(
"Mavvrik FOCUS export: row cap reached (%d rows). "
"Some data for window %s→%s may be excluded. "
"Increase MAVVRIK_FOCUS_MAX_ROWS to export all rows.",
effective_limit,
window.start_time.date(),
window.end_time.date(),
)
if data.is_empty():
verbose_proxy_logger.debug(
"Mavvrik FOCUS export: no usage data for window %s", window
)
return
normalized = engine._transformer.transform(data)
if normalized.is_empty():
return
payload = engine._serializer.serialize(normalized)
if not payload:
return
await engine._destination.deliver(
content=payload,
time_window=window,
filename=engine._build_filename(window),
)
# Maximum number of days to catch up in a single run. Prevents runaway
# loops if the connector was disabled for a long time, and avoids querying
# data that has likely been cleaned up from LiteLLM_DailyUserSpend.
_MAX_CATCHUP_DAYS = 7
async def _run_scheduled_export(self) -> None:
"""Export today's window, catching up any dates Mavvrik has not yet received.
On each run:
1. Register with Mavvrik → get metricsMarker (last successfully ingested date)
2. If metricsMarker is behind yesterday, catch up missed dates (capped at
_MAX_CATCHUP_DAYS to avoid runaway loops on long outages)
3. Export yesterday (today's daily window)
This ensures a failed export on day N is automatically retried on day N+1
without any manual intervention.
"""
engine = self._ensure_engine()
from litellm.integrations.focus.destinations.mavvrik_destination import ( # noqa: PLC0415
FocusMavvrikDestination,
)
destination = engine._destination
if not isinstance(destination, FocusMavvrikDestination):
await super()._run_scheduled_export()
return
# Register and get the last date Mavvrik has processed.
# metricsMarker may be a Unix timestamp (int/float) or an ISO date string.
marker = await destination.get_metrics_marker()
now = datetime.now(timezone.utc)
yesterday = now.replace(hour=0, minute=0, second=0, microsecond=0) - timedelta(
days=1
)
last_ingested = _parse_metrics_marker(marker)
# Catch up missed dates, capped at _MAX_CATCHUP_DAYS
if last_ingested and last_ingested < yesterday:
# Never go further back than _MAX_CATCHUP_DAYS from yesterday
earliest_catchup = yesterday - timedelta(days=self._MAX_CATCHUP_DAYS - 1)
catch_up_date = max(last_ingested + timedelta(days=1), earliest_catchup)
if last_ingested + timedelta(days=1) < earliest_catchup:
verbose_proxy_logger.warning(
"Mavvrik FOCUS export: metricsMarker is more than %d days behind "
"(%s). Catching up from %s only; earlier data will not be re-exported.",
self._MAX_CATCHUP_DAYS,
last_ingested.date(),
catch_up_date.date(),
)
while catch_up_date < yesterday:
verbose_proxy_logger.info(
"Mavvrik FOCUS export: catching up missed date %s",
catch_up_date.date(),
)
window = FocusTimeWindow(
start_time=catch_up_date,
end_time=catch_up_date + timedelta(days=1),
frequency="daily",
)
await self._export_window(window=window, limit=None)
catch_up_date += timedelta(days=1)
# Export yesterday's window (the normal daily run)
window = FocusTimeWindow(
start_time=yesterday,
end_time=yesterday + timedelta(days=1),
frequency="daily",
)
await self._export_window(window=window, limit=None)
async def initialize_mavvrik_focus_export_job(self) -> None:
"""Scheduler entry point — uses Mavvrik-specific pod-lock key."""
from litellm.proxy.proxy_server import proxy_logging_obj # noqa: PLC0415
pod_lock_manager = None
if proxy_logging_obj is not None:
writer = getattr(proxy_logging_obj, "db_spend_update_writer", None)
if writer is not None:
pod_lock_manager = getattr(writer, "pod_lock_manager", None)
if pod_lock_manager and pod_lock_manager.redis_cache:
acquired = await pod_lock_manager.acquire_lock(
cronjob_id=MAVVRIK_FOCUS_EXPORT_JOB_NAME
)
if not acquired:
verbose_proxy_logger.debug(
"Mavvrik FOCUS export: unable to acquire pod lock"
)
return
try:
await self._run_scheduled_export()
finally:
await pod_lock_manager.release_lock(
cronjob_id=MAVVRIK_FOCUS_EXPORT_JOB_NAME
)
else:
await self._run_scheduled_export()
@staticmethod
async def init_mavvrik_focus_background_job(
scheduler: AsyncIOScheduler,
) -> None:
"""Register the Mavvrik FOCUS export job on the provided scheduler."""
loggers: List[MavvrikFocusLogger] = [
cb
for cb in litellm.logging_callback_manager.get_custom_loggers_for_type(
callback_type=MavvrikFocusLogger
)
if type(cb) is MavvrikFocusLogger
]
if not loggers:
verbose_proxy_logger.debug(
"No MavvrikFocusLogger registered; skipping scheduler"
)
return
logger = loggers[0]
trigger_kwargs = logger._build_scheduler_trigger()
scheduler.add_job( # type: ignore[attr-defined]
logger.initialize_mavvrik_focus_export_job,
id=MAVVRIK_FOCUS_EXPORT_JOB_NAME,
replace_existing=True,
**trigger_kwargs,
)
verbose_proxy_logger.info(
"mavvrik_focus: background export job scheduled (%s)", trigger_kwargs
)

View file

@ -0,0 +1,10 @@
"""
New Relic AI Monitoring Integration for LiteLLM
This module provides integration with New Relic's AI Monitoring feature to track
LLM requests, responses, and usage metrics.
"""
from litellm.integrations.newrelic.newrelic import NewRelicLogger
__all__ = ["NewRelicLogger"]

View file

@ -0,0 +1,926 @@
"""
New Relic AI Monitoring Integration for LiteLLM
This module provides integration with New Relic's AI Monitoring feature to track
LLM requests, responses, and usage metrics.
Environment Variables (consumed by the New Relic agent at process bootstrap -
set via container env, or before invoking `newrelic-admin run-program`):
NEW_RELIC_LICENSE_KEY: Your New Relic license key (required)
NEW_RELIC_APP_NAME: Your application name (required)
UI- and runtime-toggleable:
NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED: Whether to record message
content (optional, default: true)
Configuration:
Message logging can be controlled via (both must agree to record):
1. turn_off_message_logging parameter - pass via callback initialization or config YAML
2. NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED env var
Default behavior: Messages ARE recorded unless explicitly disabled by either method
Either method can disable recording - both must enable for recording to occur
Usage - Python SDK:
import litellm
litellm.callbacks = ["newrelic"]
# Or with explicit configuration:
from litellm.integrations.newrelic import NewRelicLogger
litellm.callbacks = [NewRelicLogger(turn_off_message_logging=True)]
Usage - Proxy Server (config.yaml):
litellm_settings:
callbacks: ["newrelic"]
newrelic_params:
turn_off_message_logging: true # Disable message content recording
# Or disable via environment variable:
# export NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED=false
# Ensure New Relic agent is initialized (use newrelic-admin or initialize manually)
# newrelic-admin run-program python your_app.py
"""
import json
import os
import threading
import time
import uuid
from typing import Any, Dict, List, Optional, Tuple, Union
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.redact_messages import should_redact_message_logging
from litellm.types.integrations.newrelic import NewRelicInitParams
from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus
from litellm.types.utils import ModelResponse, Message, StandardLoggingPayload
try:
import newrelic.agent as _newrelic_agent
except ImportError:
_newrelic_agent = None # type: ignore
class NewRelicLogger(CustomLogger):
"""
New Relic logger for LiteLLM to send AI monitoring events.
This logger creates two types of New Relic custom events:
1. LlmChatCompletionSummary - One per completion request
2. LlmChatCompletionMessage - One per message (request and response)
"""
# Class-level state for supportability metric emission, shared across all instances.
# Protected by _metric_lock to ensure thread-safe access.
_last_metric_emission_time: float = 0.0
_metric_lock = threading.Lock()
def __init__(self, **kwargs):
#########################################################
# Handle newrelic_params set as litellm.newrelic_params
#########################################################
dict_newrelic_params = self._get_newrelic_params()
# Use setdefault so constructor kwargs take priority over global params.
# model_dump() always returns all fields (including defaults), so update()
# would silently overwrite explicit constructor args like turn_off_message_logging=True.
for k, v in dict_newrelic_params.items():
kwargs.setdefault(k, v)
# CustomLogger.__init__ will set self.turn_off_message_logging from kwargs
super().__init__(**kwargs)
# Check for required environment variables
self.license_key = os.getenv("NEW_RELIC_LICENSE_KEY")
self.app_name = os.getenv("NEW_RELIC_APP_NAME")
# Validate configuration
if not self.license_key or not self.app_name:
verbose_logger.warning(
"New Relic integration requires NEW_RELIC_LICENSE_KEY and "
"NEW_RELIC_APP_NAME environment variables. Integration will be disabled."
)
self.enabled = False
elif _newrelic_agent is None:
verbose_logger.error(
"New Relic Python agent not installed. Review the New Relic integration documentation at https://docs.litellm.ai/docs/observability/newrelic."
)
self.enabled = False
else:
try:
# timeout=0 forces non-blocking startup: the agent connects in a
# background thread regardless of newrelic.ini / NEW_RELIC_STARTUP_TIMEOUT.
_newrelic_agent.register_application(timeout=0)
self.enabled = True
verbose_logger.info(
f"New Relic AI Monitoring initialized for app: {self.app_name}, "
f"content recording: {self.record_content}"
)
except Exception as e:
verbose_logger.error(
f"Failed to initialize New Relic agent: {e}. "
"Integration will be disabled."
)
self.enabled = False
def _get_newrelic_params(self) -> Dict:
"""
Get the newrelic_params from litellm.newrelic_params
These are params specific to initializing the NewRelicLogger e.g. turn_off_message_logging
"""
dict_newrelic_params: Dict = {}
if litellm.newrelic_params is not None:
if isinstance(litellm.newrelic_params, NewRelicInitParams):
dict_newrelic_params = litellm.newrelic_params.model_dump()
elif isinstance(litellm.newrelic_params, Dict):
# only allow params that are of NewRelicInitParams
dict_newrelic_params = NewRelicInitParams(
**litellm.newrelic_params
).model_dump()
return dict_newrelic_params
@property
def record_content(self) -> bool:
"""Whether to record message content in New Relic.
Both turn_off_message_logging param AND NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED
env var must agree to record content. If either disables recording, content will not
be recorded. Read at call time so UI config changes take effect without a restart.
Default: True (record content) unless explicitly disabled by either method.
"""
return (not self.turn_off_message_logging) and self._parse_bool_env(
"NEW_RELIC_AI_MONITORING_RECORD_CONTENT_ENABLED", True
)
def _parse_bool_env(self, var_name: str, default: bool = False) -> bool:
"""Parse a boolean environment variable.
Accepts true/false, 1/0, yes/no, on/off (case-insensitive,
whitespace-tolerant) — matching the convention used in
``litellm/__init__.py`` and the standard library's
``configparser.BOOLEAN_STATES``. Unrecognised values log a
warning and fall back to ``default`` rather than silently
flipping user intent.
"""
raw = os.getenv(var_name)
if not raw:
return default
value = raw.strip().lower()
if value in ("1", "true", "yes", "on"):
return True
if value in ("0", "false", "no", "off"):
return False
verbose_logger.warning(
f"{var_name}={raw!r} is not a recognised boolean "
f"(accepts true/false, 1/0, yes/no, on/off). "
f"Falling back to default ({default})."
)
return default
def _get_litellm_version(self) -> str:
"""
Get litellm version for supportability metrics.
Returns:
Version string (e.g., "1.80.0") or "unknown" if unable to determine
"""
try:
from importlib.metadata import version
return version("litellm")
except Exception as e:
verbose_logger.warning(f"Unable to determine litellm version: {e}")
return "unknown"
def _emit_supportability_metric(self):
"""
Emit New Relic supportability metric for LiteLLM usage.
Per spec, this metric should be emitted at least once every 27 hours
to indicate the library is in use. Format:
Supportability/Python/ML/LiteLLM/{version}
This method updates _last_metric_emission_time and should
be called within a lock when checking periodic emission.
"""
try:
litellm_version = self._get_litellm_version()
metric_name = f"Supportability/Python/ML/LiteLLM/{litellm_version}"
# Record metric with value of 1 (will be aggregated by New Relic)
app = _newrelic_agent.application()
# Always update the timestamp so the 27-hour back-off applies
# regardless of whether the app is ready, preventing lock contention
# on every request when the agent is slow to register or never starts.
NewRelicLogger._last_metric_emission_time = time.time()
if app and app.enabled:
app.record_custom_metric(metric_name, 1)
verbose_logger.info(
f"Emitted New Relic supportability metric: {metric_name}"
)
else:
verbose_logger.info(
"New Relic application is not enabled; skipping metric recording."
)
except Exception as e:
verbose_logger.warning(f"Failed to emit supportability metric: {e}")
def _check_and_emit_periodic_metric(self):
"""
Check if 27 hours have passed since last metric emission and re-emit if needed.
Uses a mutex to ensure only one thread emits the metric even if multiple
requests are being processed concurrently.
"""
# Quick check without lock to avoid unnecessary locking
current_time = time.time()
time_since_last_emission = (
current_time - NewRelicLogger._last_metric_emission_time
)
if time_since_last_emission >= 97200: # 27 hours = 97200 seconds
# Acquire lock to ensure only one thread emits
with NewRelicLogger._metric_lock:
# Double-check inside lock in case another thread just emitted
current_time = time.time()
time_since_last_emission = (
current_time - NewRelicLogger._last_metric_emission_time
)
if time_since_last_emission >= 97200:
self._emit_supportability_metric()
def _get_trace_context(
self,
kwargs: Dict,
standard_logging_object: Optional[StandardLoggingPayload] = None,
) -> str:
"""
Get the New Relic trace ID for AI monitoring events.
This integration runs in LiteLLM's async logging worker, outside the
New Relic agent's current transaction. Because we can't call
`newrelic.agent.current_trace_id()` to let the agent populate the
trace_id on AIM custom events, we manually simulate what the agent
would do. An AIM event without a trace_id is malformed per the NR
schema, so this method always returns a valid string.
Resolution order:
1. W3C traceparent header (litellm_params.metadata.headers.traceparent) -
what the agent would link to if we were in-transaction.
2. StandardLoggingPayload.trace_id - LiteLLM's internal trace for
retry/fallback grouping.
3. Generated UUID - synthetic grouping key when upstream context is
absent or parsing it fails.
Span IDs are intentionally not emitted: any span ID recoverable from
the inbound traceparent is the caller's parent span, not ours.
Returns:
trace_id: always a non-empty string.
"""
trace_id: Optional[str] = None
try:
litellm_params = kwargs.get("litellm_params") or {}
metadata = litellm_params.get("metadata") or {}
headers = metadata.get("headers") or {}
# Normalize header key lookup to be case-insensitive per W3C spec
traceparent = next(
(v for k, v in headers.items() if k.lower() == "traceparent"), None
)
if traceparent:
# Extract trace_id from traceparent header if available
# traceparent format: "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-00"
parts = traceparent.split("-")
if len(parts) == 4:
trace_id = parts[1]
if not trace_id and standard_logging_object:
slo_trace_id = standard_logging_object.get("trace_id")
if slo_trace_id:
trace_id = slo_trace_id
except Exception as e:
verbose_logger.warning(
f"Unable to parse New Relic trace context from upstream sources: {e}"
)
if not trace_id:
trace_id = uuid.uuid4().hex
verbose_logger.debug(
f"New Relic trace_id not available from distributed tracing headers or "
f"StandardLoggingPayload. Generated trace_id={trace_id} for AI monitoring "
f"event grouping."
)
return trace_id
def _extract_completion_id(self, kwargs: Dict, response_obj: ModelResponse) -> str:
"""
Extract completion ID from kwargs or response_obj, or generate one.
"""
completion_id = None
if response_obj:
completion_id = response_obj.get("id")
if not completion_id:
completion_id = kwargs.get("litellm_call_id")
# If still not found, generate UUID and log warning per spec
if not completion_id:
completion_id = str(uuid.uuid4())
return completion_id
def _get_vendor(
self,
kwargs: Dict,
standard_logging_object: Optional[StandardLoggingPayload] = None,
) -> str:
"""Extract vendor/provider, preferring StandardLoggingPayload."""
if standard_logging_object:
vendor = standard_logging_object.get("custom_llm_provider")
if vendor:
return vendor
litellm_params = kwargs.get("litellm_params", {}) or {}
return litellm_params.get("custom_llm_provider") or "litellm"
def _get_model_names(
self,
kwargs: Dict,
response_obj: ModelResponse,
standard_logging_object: Optional[StandardLoggingPayload] = None,
) -> Tuple[str, str]:
"""
Extract request and response model names, preferring StandardLoggingPayload
for the request model.
Returns:
Tuple of (request_model, response_model)
"""
request_model = None
if standard_logging_object:
slo_model = standard_logging_object.get("model")
if slo_model:
request_model = str(slo_model)
if not request_model:
request_model = str(kwargs.get("model") or "unknown")
response_model: str = str(response_obj.get("model") or request_model)
return request_model, response_model
def _extract_usage(
self,
response_obj: ModelResponse,
standard_logging_object: Optional[StandardLoggingPayload] = None,
) -> Dict[str, int]:
"""Extract usage statistics, preferring StandardLoggingPayload."""
if standard_logging_object:
prompt = standard_logging_object.get("prompt_tokens")
completion = standard_logging_object.get("completion_tokens")
total = standard_logging_object.get("total_tokens")
if any(x is not None for x in [prompt, completion, total]):
return {
"prompt_tokens": prompt or 0,
"completion_tokens": completion or 0,
"total_tokens": total or 0,
}
usage = response_obj.get("usage", None)
if not usage:
return {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
return {
"prompt_tokens": usage.get("prompt_tokens") or 0,
"completion_tokens": usage.get("completion_tokens") or 0,
"total_tokens": usage.get("total_tokens") or 0,
}
def _get_finish_reason(self, response_obj: ModelResponse) -> str:
"""
Extract finish reason from first choice in the response.
Returns "unknown" if choices are not present or finish_reason is not found.
"""
choices = response_obj.get("choices") or []
if choices and len(choices) > 0:
return choices[0].get("finish_reason") or "unknown"
return "unknown"
def _to_epoch_ms(self, t: Any) -> float:
"""Convert a datetime or float timestamp to epoch milliseconds."""
if hasattr(t, "timestamp"):
return t.timestamp() * 1000.0
return float(t) * 1000.0
def _get_duration(
self,
kwargs: Dict,
start_time: Any,
end_time: Any,
standard_logging_object: Optional[StandardLoggingPayload] = None,
) -> Optional[float]:
"""
Extract duration in milliseconds.
Resolution order:
1. StandardLoggingPayload.response_time (already computed by LiteLLM)
2. llm_api_duration_ms from kwargs
3. Calculated from start_time and end_time
"""
if standard_logging_object:
response_time = standard_logging_object.get("response_time")
if response_time is not None:
return (
float(response_time) * 1000.0
) # SLO stores seconds; convert to ms
duration_ms = kwargs.get("llm_api_duration_ms")
if duration_ms is not None:
return float(duration_ms)
if start_time is not None and end_time is not None:
return self._to_epoch_ms(end_time) - self._to_epoch_ms(start_time)
return None
def _get_request_params(
self,
kwargs: Dict,
standard_logging_object: Optional[StandardLoggingPayload] = None,
) -> Dict[str, Any]:
"""
Extract request parameters like temperature and max_tokens, preferring
StandardLoggingPayload.model_parameters.
Returns dict with available parameters, omitting those not present.
"""
if standard_logging_object:
source_params = standard_logging_object.get("model_parameters") or {}
else:
source_params = kwargs.get("optional_params") or {}
params = {}
temperature = source_params.get("temperature")
if temperature is not None:
params["temperature"] = temperature
max_tokens = source_params.get("max_tokens")
if max_tokens is not None:
params["max_tokens"] = max_tokens
return params
def _extract_message_content(self, message: Union[Message, Dict]) -> str:
"""
Extract content from a message, handling various formats.
Handles tool calls, multimodal content (as JSON), and standard text content.
Returns empty string if content is None or missing.
"""
content = message.get("content")
# Handle tool calls
if message.get("tool_calls"):
try:
return json.dumps(message["tool_calls"])
except Exception:
return str(message["tool_calls"])
# Handle None or missing content
if content is None:
return ""
# Handle list content (multimodal)
if isinstance(content, list):
try:
return json.dumps(content)
except Exception:
return str(content)
# Handle non-string content
if not isinstance(content, str):
return str(content)
return content
def _extract_all_messages(
self,
kwargs: Dict,
response_obj: ModelResponse,
response_model: str,
vendor: str,
standard_logging_object: Optional[StandardLoggingPayload] = None,
) -> List[Dict[str, Any]]:
"""
Extract all messages (request + response) with sequence numbers and timestamps.
Processes request messages from StandardLoggingPayload.messages (preferred) or
kwargs["messages"] (fallback), and response messages from response_obj["choices"].
Assigns sequential numbers starting at 0.
Adds timestamps from StandardLoggingPayload (preferred) or kwargs if available
(converted to epoch milliseconds).
"""
messages = []
sequence = 0
# Extract timestamps, preferring StandardLoggingPayload
start_time = None
if standard_logging_object:
start_time = standard_logging_object.get("startTime")
if not start_time:
start_time = kwargs.get("start_time")
end_time = None
if standard_logging_object:
end_time = standard_logging_object.get("endTime")
if not end_time:
end_time = kwargs.get("end_time")
# Content is recorded only when the NR-specific switches allow it AND
# LiteLLM's wider redaction decision (turn_off_message_logging, dynamic
# params, headers) does not require redaction. Async streaming hands the
# callback an unredacted async_complete_streaming_response, so without
# this gate generated content would still reach NR even when the user
# has globally disabled message logging.
record_content = self.record_content and not should_redact_message_logging(
kwargs
)
# Extract request messages, preferring StandardLoggingPayload.
# SLO messages can be a string (serialized/redacted), so only use it when it's a list.
slo_messages = (
standard_logging_object.get("messages") if standard_logging_object else None
)
if isinstance(slo_messages, list):
request_messages = slo_messages
else:
request_messages = kwargs.get("messages") or []
for msg in request_messages:
message_data = {
"role": msg.get("role") or "user",
"sequence": sequence,
"response.model": response_model,
"vendor": vendor,
}
# Add timestamp for request message if available (convert to milliseconds)
if start_time is not None:
message_data["timestamp"] = int(self._to_epoch_ms(start_time))
if record_content:
message_data["content"] = self._extract_message_content(msg)
messages.append(message_data)
sequence += 1
# Extract response messages from choices
choices = response_obj.get("choices") or []
if choices and len(choices) > 0:
for choice in choices:
# Prefer "message" (non-streaming); fall back to "delta" (streaming-assembled)
message = choice.get("message", None) or choice.get("delta", None)
if message:
message_data = {
"role": message.get("role") or "assistant",
"sequence": sequence,
"response.model": response_model,
"vendor": vendor,
"is_response": True,
}
# Add timestamp for response message if available (convert to milliseconds)
if end_time is not None:
message_data["timestamp"] = int(self._to_epoch_ms(end_time))
if record_content:
message_data["content"] = self._extract_message_content(message)
messages.append(message_data)
sequence += 1
return messages
def _record_summary_event(
self,
request_id: str,
trace_id: Optional[str],
request_model: str,
response_model: str,
vendor: str,
finish_reason: str,
num_messages: int,
usage: Dict[str, int],
duration: Optional[float] = None,
request_params: Optional[Dict[str, Any]] = None,
):
"""Record LlmChatCompletionSummary event to New Relic."""
try:
event_data = {
"id": request_id,
"request_id": request_id,
"request.model": request_model,
"response.model": response_model,
"response.choices.finish_reason": finish_reason,
"response.number_of_messages": num_messages,
"vendor": vendor,
"ingest_source": "litellm",
"response.usage.prompt_tokens": usage["prompt_tokens"],
"response.usage.completion_tokens": usage["completion_tokens"],
"response.usage.total_tokens": usage["total_tokens"],
}
# Add optional attributes if present
if trace_id:
event_data["trace_id"] = trace_id
if duration is not None:
event_data["duration"] = duration
# Add request parameters if present
if request_params:
if "temperature" in request_params:
event_data["request.temperature"] = request_params["temperature"]
if "max_tokens" in request_params:
event_data["request.max_tokens"] = request_params["max_tokens"]
app = _newrelic_agent.application()
if app and app.enabled:
app.record_custom_event("LlmChatCompletionSummary", event_data)
else:
verbose_logger.warning(
"New Relic application is not enabled; skipping summary event recording."
)
except Exception as e:
verbose_logger.warning(f"Failed to record New Relic summary event: {e}")
self.handle_callback_failure("newrelic")
def _record_message_events(
self,
request_id: str,
llm_response_id: str,
trace_id: Optional[str],
messages: List[Dict[str, Any]],
):
"""Record LlmChatCompletionMessage events to New Relic.
Args:
request_id: Agent-generated UUID that links to Summary event's id
llm_response_id: LLM's response ID (e.g., "chatcmpl-...") for message id format
trace_id: Trace ID for distributed tracing (None if not available)
messages: List of message dicts to record
"""
try:
app = _newrelic_agent.application()
if not (app and app.enabled):
verbose_logger.warning(
"New Relic application is not enabled; skipping message event recording."
)
return
for message in messages:
sequence = message["sequence"]
event_data = {
"id": f"{llm_response_id}-{sequence}",
"request_id": request_id,
"completion_id": request_id,
"role": message["role"],
"sequence": sequence,
"response.model": message["response.model"],
"vendor": message["vendor"],
"ingest_source": "litellm",
"token_count": 0, # Per-message token counts are not available from LiteLLM
}
# Add trace context if available
if trace_id:
event_data["trace_id"] = trace_id
# Add content only if it was included in the message data
if "content" in message:
event_data["content"] = message["content"]
# Add is_response only if True (per spec, omit for request messages)
if message.get("is_response"):
event_data["is_response"] = True
# Forward actual request/response timestamp (ms) so NR uses the
# real LLM call window rather than the async-logger fire time.
# Requires newrelic>=11.2.0 which reads params["timestamp"] as
# the intrinsic event timestamp.
if "timestamp" in message:
event_data["timestamp"] = message["timestamp"]
app.record_custom_event("LlmChatCompletionMessage", event_data)
except Exception as e:
verbose_logger.warning(f"Failed to record New Relic message events: {e}")
self.handle_callback_failure("newrelic")
def _record_error_metric(self):
"""Record error metric to New Relic."""
try:
if not self.enabled:
return
self._check_and_emit_periodic_metric()
app = _newrelic_agent.application()
if app and app.enabled:
app.record_custom_metric("LLM/LiteLLM/Error", 1)
except Exception as e:
verbose_logger.warning(f"Failed to record New Relic error metric: {e}")
self.handle_callback_failure("newrelic")
def _process_success(
self,
kwargs: Dict,
response_obj: ModelResponse,
start_time: Optional[float] = None,
end_time: Optional[float] = None,
):
"""
Core logic for processing successful LLM calls.
Used by both sync and async success event handlers.
"""
# Early exit if not enabled
if not self.enabled:
return
# Check and emit periodic supportability metric if 27 hours have passed
self._check_and_emit_periodic_metric()
# Use StandardLoggingPayload where available for normalized, pre-computed values
standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get(
"standard_logging_object"
)
# Get trace context
trace_id = self._get_trace_context(kwargs, standard_logging_object)
# Generate unique request ID for this request (used as Summary event id)
request_id = str(uuid.uuid4())
# Extract data from response
llm_response_id = self._extract_completion_id(kwargs, response_obj)
vendor = self._get_vendor(kwargs, standard_logging_object)
request_model, response_model = self._get_model_names(
kwargs, response_obj, standard_logging_object
)
usage = self._extract_usage(response_obj, standard_logging_object)
finish_reason = self._get_finish_reason(response_obj)
# Extract additional summary event fields
duration = self._get_duration(
kwargs, start_time, end_time, standard_logging_object
)
request_params = self._get_request_params(kwargs, standard_logging_object)
# Extract all messages
messages = self._extract_all_messages(
kwargs, response_obj, response_model, vendor, standard_logging_object
)
# Record summary event
self._record_summary_event(
request_id=request_id,
trace_id=trace_id,
request_model=request_model,
response_model=response_model,
vendor=vendor,
finish_reason=finish_reason,
num_messages=len(messages),
usage=usage,
duration=duration,
request_params=request_params,
)
# Record message events
self._record_message_events(
request_id=request_id,
llm_response_id=llm_response_id,
trace_id=trace_id,
messages=messages,
)
async def async_health_check(self) -> IntegrationHealthCheckStatus:
"""
Check if the New Relic integration is healthy.
Verifies that the integration is enabled and the New Relic agent
has an active, connected application, then records a small
`LiteLLMConnectionTest` custom event so the user can confirm the
end-to-end pipeline in the New Relic UI via NRQL:
`SELECT * FROM LiteLLMConnectionTest SINCE 1 hour ago`.
The `LiteLLMConnectionTest` event type is intentionally outside the
`Llm*` family that AI Monitoring queries, so test events do not
appear in AI Monitoring dashboards.
"""
if not self.enabled:
return IntegrationHealthCheckStatus(
status="unhealthy",
error_message="New Relic integration is disabled. Check that "
"NEW_RELIC_LICENSE_KEY and NEW_RELIC_APP_NAME are set and the "
"newrelic package is installed.",
)
try:
app = _newrelic_agent.application()
if not (app and app.enabled):
return IntegrationHealthCheckStatus(
status="unhealthy",
error_message=(
"New Relic Python agent not installed. Review the New Relic integration documentation at https://docs.litellm.ai/docs/observability/newrelic."
),
)
app.record_custom_event(
"LiteLLMConnectionTest",
{
"is_test_event": True,
"app_name": self.app_name,
"source": "litellm-proxy",
"timestamp": time.time(),
},
)
return IntegrationHealthCheckStatus(status="healthy", error_message=None)
except Exception as e:
return IntegrationHealthCheckStatus(
status="unhealthy",
error_message=str(e),
)
# CustomLogger interface implementation
def log_pre_api_call(self, model, messages, kwargs):
"""Unused per spec."""
pass
def log_post_api_call(self, kwargs, response_obj, start_time, end_time):
"""Unused per spec."""
pass
def log_success_event(self, kwargs, response_obj, start_time, end_time):
"""
Main success path for non-streaming requests.
Note: New Relic's record_custom_event is synchronous but non-blocking
(in-memory operation), so it's safe to call from sync context.
"""
try:
self._process_success(kwargs, response_obj, start_time, end_time)
except Exception as e:
verbose_logger.warning(f"Error in New Relic log_success_event: {e}")
self.handle_callback_failure("newrelic")
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
"""
Main success path for async/streaming requests.
Note: New Relic's SDK is thread-safe and record_custom_event is fast,
so we can call it directly without asyncio.to_thread().
"""
try:
self._process_success(kwargs, response_obj, start_time, end_time)
except Exception as e:
verbose_logger.warning(f"Error in New Relic async_log_success_event: {e}")
self.handle_callback_failure("newrelic")
def log_failure_event(self, kwargs, response_obj, start_time, end_time):
"""
Log error metric for failed LLM calls (sync).
Per spec: Do not send AI events on failure, only record error metric.
"""
try:
self._record_error_metric()
except Exception as e:
verbose_logger.warning(f"Error in New Relic log_failure_event: {e}")
self.handle_callback_failure("newrelic")
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
"""
Log error metric for failed LLM calls (async).
Per spec: Do not send AI events on failure, only record error metric.
"""
try:
self._record_error_metric()
except Exception as e:
verbose_logger.warning(f"Error in New Relic async_log_failure_event: {e}")
self.handle_callback_failure("newrelic")

View file

@ -1,7 +1,18 @@
import os
from dataclasses import dataclass, field
from datetime import datetime
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Union, cast
from typing import (
TYPE_CHECKING,
Any,
Dict,
FrozenSet,
List,
Optional,
Set,
Tuple,
Union,
cast,
)
import litellm
from litellm._logging import verbose_logger
@ -82,6 +93,88 @@ _VALID_CAPTURE_MODES = {
CAPTURE_MODE_SPAN_AND_EVENT,
}
METRIC_METADATA_KEYS: Tuple[str, ...] = (
"user_api_key_hash",
"user_api_key_alias",
"user_api_key_team_id",
"user_api_key_org_id",
"user_api_key_user_id",
"user_api_key_team_alias",
"user_api_key_user_email",
"spend_logs_metadata",
"requester_ip_address",
"requester_metadata",
"user_api_key_end_user_id",
"prompt_management_metadata",
"applied_guardrails",
"mcp_tool_call_metadata",
"vector_store_request_metadata",
)
TOKEN_TYPE_ATTRIBUTE: str = "gen_ai.token.type"
VALID_METRIC_ATTRIBUTE_NAMES: FrozenSet[str] = frozenset(
(
"gen_ai.operation.name",
"gen_ai.system",
"gen_ai.request.model",
"gen_ai.framework",
"hidden_params",
)
+ tuple(f"metadata.{key}" for key in METRIC_METADATA_KEYS)
)
@dataclass(frozen=True)
class OTELMetricAttributeFilter:
include_list: Optional[List[str]] = None
exclude_list: Optional[List[str]] = None
def _build_metric_attribute_filter(value: Any) -> OTELMetricAttributeFilter:
if isinstance(value, OTELMetricAttributeFilter):
return value
if not isinstance(value, dict):
raise ValueError(
"otel.attributes must be a mapping with optional 'include_list' / "
f"'exclude_list', got {type(value).__name__}"
)
return OTELMetricAttributeFilter(
include_list=value.get("include_list"),
exclude_list=value.get("exclude_list"),
)
def _resolve_metric_attribute_filter(
attributes: Optional[OTELMetricAttributeFilter],
) -> Tuple[Optional[FrozenSet[str]], Optional[FrozenSet[str]]]:
if attributes is None:
return None, None
include = attributes.include_list or None
exclude = attributes.exclude_list or None
if include and exclude:
raise ValueError(
"otel.attributes: include_list and exclude_list are mutually exclusive"
)
requested = include or exclude or []
if TOKEN_TYPE_ATTRIBUTE in requested:
raise ValueError(
f"otel.attributes: {TOKEN_TYPE_ATTRIBUTE} is a structural token-usage "
"discriminator and cannot be filtered"
)
unknown = sorted(
name for name in requested if name not in VALID_METRIC_ATTRIBUTE_NAMES
)
if unknown:
raise ValueError(
f"otel.attributes: unknown attribute name(s) {unknown}. "
f"Valid names: {sorted(VALID_METRIC_ATTRIBUTE_NAMES)}"
)
return (
frozenset(include) if include else None,
frozenset(exclude) if exclude else None,
)
def _normalize_team_metadata_keys(value: Any) -> List[str]:
"""Coerce a team-metadata allowlist from a list or comma-separated string.
@ -117,6 +210,9 @@ class OpenTelemetryConfig:
# under ``litellm.team.metadata``. Empty by default so none of a team's
# metadata leaves the process until explicitly allowlisted.
baggage_team_metadata_keys: List[str] = field(default_factory=list)
# Prometheus-style include/exclude control over which attributes are stamped
# on emitted metrics, to cap metric cardinality.
attributes: Optional[OTELMetricAttributeFilter] = None
def __post_init__(self) -> None:
# If endpoint is specified but exporter is still the default "console",
@ -211,15 +307,29 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
**kwargs,
):
team_metadata_keys_override = kwargs.pop("baggage_team_metadata_keys", None)
metric_attributes_override = kwargs.pop("attributes", None)
if config is None:
config = OpenTelemetryConfig.from_env()
if team_metadata_keys_override is not None:
config.baggage_team_metadata_keys = _normalize_team_metadata_keys(
team_metadata_keys_override
)
if metric_attributes_override is not None:
config.attributes = _build_metric_attribute_filter(
metric_attributes_override
)
self.config = config
self.callback_name = callback_name
# Resolved on first metric record, not here: the proxy populates
# callback_settings.otel.attributes after this logger is constructed, so
# reading it now would miss it. An explicit config is validated eagerly so
# a bad config still fails at startup.
self._metric_attr_include: Optional[FrozenSet[str]] = None
self._metric_attr_exclude: Optional[FrozenSet[str]] = None
self._metric_attr_filter_resolved = False
if config.attributes is not None:
self._ensure_metric_attribute_filter()
self.OTEL_EXPORTER = self.config.exporter
self.OTEL_ENDPOINT = self.config.endpoint
self.OTEL_HEADERS = self.config.headers
@ -1318,6 +1428,38 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
return None
return safe_dumps(filtered)
def _ensure_metric_attribute_filter(self) -> None:
"""Resolve the include/exclude filter once, falling back to the proxy's
callback_settings.otel.attributes when no explicit config was passed."""
if self._metric_attr_filter_resolved:
return
attributes = self.config.attributes
if attributes is None and self.callback_name in (None, "otel"):
otel_settings = (litellm.callback_settings or {}).get("otel") or {}
raw = (
otel_settings.get("attributes")
if isinstance(otel_settings, dict)
else None
)
if raw is not None:
attributes = _build_metric_attribute_filter(raw)
(
self._metric_attr_include,
self._metric_attr_exclude,
) = _resolve_metric_attribute_filter(attributes)
self._metric_attr_filter_resolved = True
def _filter_metric_attributes(self, attrs: Dict[str, Any]) -> Dict[str, Any]:
if not self._metric_attr_filter_resolved:
self._ensure_metric_attribute_filter()
if self._metric_attr_include is not None:
return {k: v for k, v in attrs.items() if k in self._metric_attr_include}
if self._metric_attr_exclude is not None:
return {
k: v for k, v in attrs.items() if k not in self._metric_attr_exclude
}
return attrs
def _record_metrics(self, kwargs, response_obj, start_time, end_time):
duration_s = (end_time - start_time).total_seconds()
params = kwargs.get("litellm_params") or {}
@ -1336,23 +1478,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
std_log = kwargs.get("standard_logging_object")
md = getattr(std_log, "metadata", None) or (std_log or {}).get("metadata", {})
for key in [
"user_api_key_hash",
"user_api_key_alias",
"user_api_key_team_id",
"user_api_key_org_id",
"user_api_key_user_id",
"user_api_key_team_alias",
"user_api_key_user_email",
"spend_logs_metadata",
"requester_ip_address",
"requester_metadata",
"user_api_key_end_user_id",
"prompt_management_metadata",
"applied_guardrails",
"mcp_tool_call_metadata",
"vector_store_request_metadata",
]:
for key in METRIC_METADATA_KEYS:
value = md.get(key)
if value is None:
continue
@ -1368,6 +1494,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
if hidden_params:
common_attrs["hidden_params"] = safe_dumps(hidden_params)
common_attrs = self._filter_metric_attributes(common_attrs)
if self._operation_duration_histogram:
self._operation_duration_histogram.record(
duration_s, attributes=common_attrs
@ -1377,8 +1505,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
and (usage := response_obj.get("usage"))
and self._token_usage_histogram
):
in_attrs = {**common_attrs, "gen_ai.token.type": "input"}
out_attrs = {**common_attrs, "gen_ai.token.type": "output"}
in_attrs = {**common_attrs, TOKEN_TYPE_ATTRIBUTE: "input"}
out_attrs = {**common_attrs, TOKEN_TYPE_ATTRIBUTE: "output"}
self._token_usage_histogram.record(
usage.get("prompt_tokens", 0), attributes=in_attrs
)

View file

@ -17,7 +17,7 @@ from litellm.integrations.otel.model.payloads import (
ServiceSpanData,
)
from litellm.integrations.otel.plumbing.providers import to_otel_span_kind
from litellm.integrations.otel.model.semconv import Error
from litellm.integrations.otel.model.semconv import Error, ExceptionEvent
from litellm.integrations.otel.model.spans import (
SPAN_REGISTRY,
SpanRole,
@ -179,9 +179,17 @@ class SpanEmitter:
else None
)
if error and (error.error_type or error.message):
span.set_attribute(Error.TYPE, error.error_type or "error")
span.set_status(
Status(StatusCode.ERROR, error.message or error.error_type or "error")
error_type = error.error_type or "error"
message = error.message or error.error_type or "error"
span.set_attribute(Error.TYPE, error_type)
span.set_status(Status(StatusCode.ERROR, message))
# Carry the full message on the standard ``exception`` event so backends
# map it as full text under ``exception.message``. Setting it as a bare
# string attribute instead lets backends like Elasticsearch dynamic-map
# it to a ``keyword`` capped at 1024 chars, truncating the message.
span.add_event(
ExceptionEvent.NAME,
{ExceptionEvent.TYPE: error_type, ExceptionEvent.MESSAGE: message},
)
# On success leave the status UNSET (the semconv default) rather than
# forcing OK — that matches the FastAPI server span and avoids implying a

View file

@ -63,6 +63,20 @@ class GenAIMapper:
# routing) onto the boundary-born LLM span — stamp it directly here.
LiteLLM.PROVIDER_MODEL: lambda d: d.identity.provider_model or None,
f"{LiteLLM.COST_PREFIX}total": lambda d: d.response_cost,
# Per-component cost breakdown (from the StandardLoggingPayload
# ``cost_breakdown``). Each component is omitted when the source didn't
# report it, so spans stay sparse rather than carrying zeros.
f"{LiteLLM.COST_PREFIX}input": lambda d: d.cost.input,
f"{LiteLLM.COST_PREFIX}output": lambda d: d.cost.output,
f"{LiteLLM.COST_PREFIX}cache_read": lambda d: d.cost.cache_read,
f"{LiteLLM.COST_PREFIX}cache_creation": lambda d: d.cost.cache_creation,
f"{LiteLLM.COST_PREFIX}tool_usage": lambda d: d.cost.tool_usage,
f"{LiteLLM.COST_PREFIX}original": lambda d: d.cost.original,
f"{LiteLLM.COST_PREFIX}discount_amount": lambda d: d.cost.discount_amount,
f"{LiteLLM.COST_PREFIX}discount_percent": lambda d: d.cost.discount_percent,
f"{LiteLLM.COST_PREFIX}margin_fixed_amount": lambda d: d.cost.margin_fixed_amount,
f"{LiteLLM.COST_PREFIX}margin_percent": lambda d: d.cost.margin_percent,
f"{LiteLLM.COST_PREFIX}margin_total_amount": lambda d: d.cost.margin_total_amount,
LiteLLM.REQUEST_STREAMING: lambda d: d.is_streaming,
}

View file

@ -34,6 +34,7 @@ __all__ = [
"RequestIdentity",
"GuardrailSpanData",
"LLMCallSpanData",
"LLMCost",
"LLMRequestParams",
"LLMUsage",
"MCPToolCallSpanData",
@ -91,6 +92,49 @@ class LLMUsage:
total_tokens: int | None = None
@dataclass(frozen=True)
class LLMCost:
"""Per-component cost breakdown, from the StandardLoggingPayload
``cost_breakdown`` (``litellm.types.utils.CostBreakdown``).
Each field is the USD cost of one component, or ``None`` when the source did
not report it — so the mapper omits absent components instead of emitting 0.
The final (post-discount/post-margin) total is carried separately on
``LLMCallSpanData.response_cost``. Free-form ``additional_costs`` are not
surfaced here: span attributes are scalar and there is no agreed key shape
for them yet.
"""
input: float | None = None
output: float | None = None
cache_read: float | None = None
cache_creation: float | None = None
tool_usage: float | None = None
original: float | None = None
discount_amount: float | None = None
discount_percent: float | None = None
margin_fixed_amount: float | None = None
margin_percent: float | None = None
margin_total_amount: float | None = None
@classmethod
def from_breakdown(cls, breakdown: Mapping[str, object] | None) -> "LLMCost":
b = breakdown or {}
return cls(
input=as_float(b.get("input_cost")),
output=as_float(b.get("output_cost")),
cache_read=as_float(b.get("cache_read_cost")),
cache_creation=as_float(b.get("cache_creation_cost")),
tool_usage=as_float(b.get("tool_usage_cost")),
original=as_float(b.get("original_cost")),
discount_amount=as_float(b.get("discount_amount")),
discount_percent=as_float(b.get("discount_percent")),
margin_fixed_amount=as_float(b.get("margin_fixed_amount")),
margin_percent=as_float(b.get("margin_percent")),
margin_total_amount=as_float(b.get("margin_total_amount")),
)
@dataclass(frozen=True)
class SpanError:
error_type: str | None = None
@ -255,6 +299,7 @@ class LLMCallSpanData:
server: ServerInfo | None
identity: RequestIdentity
is_streaming: bool | None = None
cost: LLMCost = field(default_factory=LLMCost)
tools: tuple[ToolDefinition, ...] = ()
# Raw messages and response, needed by vendor mappers (OpenInference,
# Langfuse, Weave) that stamp message-level attributes. ``messages_in`` is
@ -302,6 +347,9 @@ class LLMCallSpanData:
finish_reasons=finish_reasons,
error=_parse_error(payload),
response_cost=as_float(payload.get("response_cost")),
cost=LLMCost.from_breakdown(
cast("Mapping[str, object] | None", payload.get("cost_breakdown"))
),
server=ServerInfo.from_api_base(context.api_base),
identity=context.identity,
is_streaming=as_bool(payload.get("stream")),

View file

@ -146,6 +146,21 @@ class Error:
TYPE: Final = "error.type"
class ExceptionEvent:
"""OTel exception-event name and attribute keys (semconv ``exception.*``).
The full error message rides ``exception.message`` on a span event rather than
a custom string attribute. Backends recognise these semantic-convention names
and map them as full text; an unrecognised key (e.g. ``error_message``) falls
into the default dynamic template, which truncates strings to a 1024-char
``keyword``.
"""
NAME: Final = "exception"
TYPE: Final = "exception.type"
MESSAGE: Final = "exception.message"
class Server:
ADDRESS: Final = "server.address"
PORT: Final = "server.port"

View file

@ -17,6 +17,7 @@ from opentelemetry.sdk.trace.export.in_memory_span_exporter import (
)
from opentelemetry.trace import Span, SpanKind, Tracer
from litellm._version import version as litellm_version
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
from litellm.integrations.otel.model.semconv import LiteLLM
from litellm.integrations.otel.model.spans import LiteLLMSpanKind
@ -207,7 +208,10 @@ def build_tracer_provider(
def get_tracer(provider: TracerProvider, name: str = "litellm") -> Tracer:
return provider.get_tracer(name)
# Stamp the instrumentation scope with the LiteLLM package version so every
# emitted span carries a deterministic ``scope.version`` (the standard OTel
# location for the emitting library's version) for downstream consumers.
return provider.get_tracer(name, litellm_version)
def in_memory_provider(

View file

@ -2,11 +2,40 @@
Utility functions for the Agents API SDK.
"""
from typing import Optional
from typing import Dict, Mapping, Optional
from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig
def merge_agent_headers(
*,
dynamic_headers: Optional[Mapping[str, str]] = None,
static_headers: Optional[Mapping[str, str]] = None,
) -> Optional[Dict[str, str]]:
"""Merge outbound HTTP headers for A2A agent calls.
Merge rules:
- Start with ``dynamic_headers`` (values extracted from the incoming client request).
- Overlay ``static_headers`` (admin-configured per agent).
- Comparison is case-insensitive (HTTP headers are case-insensitive), so a
static ``Authorization`` strips any dynamic ``authorization`` before the
static value is written. The static side's casing is preserved.
If both contain the same header (case-insensitively), ``static_headers`` wins.
"""
merged: Dict[str, str] = {}
if dynamic_headers:
merged.update({str(k): str(v) for k, v in dynamic_headers.items()})
if static_headers:
static_lower = {str(k).lower() for k in static_headers}
merged = {k: v for k, v in merged.items() if k.lower() not in static_lower}
merged.update({str(k): str(v) for k, v in static_headers.items()})
return merged or None
def get_provider_agents_api_config(
custom_llm_provider: Optional[str],
) -> Optional[BaseAgentsAPIConfig]:

View file

@ -25,6 +25,7 @@ from litellm.integrations.datadog.datadog_metrics import DatadogMetricsLogger
from litellm.integrations.deepeval import DeepEvalLogger
from litellm.integrations.dotprompt import DotpromptManager
from litellm.integrations.focus.focus_logger import FocusLogger
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import MavvrikFocusLogger
from litellm.integrations.vantage.vantage_logger import VantageLogger
from litellm.integrations.galileo import GalileoObserve
from litellm.integrations.gcs_bucket.gcs_bucket import GCSBucketLogger
@ -39,6 +40,7 @@ from litellm.integrations.langsmith import LangsmithLogger
from litellm.integrations.litellm_agent import LiteLLMAgentModelResolver
from litellm.integrations.literal_ai import LiteralAILogger
from litellm.integrations.mlflow import MlflowLogger
from litellm.integrations.newrelic import NewRelicLogger
from litellm.integrations.openmeter import OpenMeterLogger
from litellm.integrations.opentelemetry import OpenTelemetry
from litellm.integrations.opik.opik import OpikLogger
@ -102,8 +104,10 @@ class CustomLoggerRegistry:
"gitlab": GitLabPromptManager,
"cloudzero": CloudZeroLogger,
"focus": FocusLogger,
"mavvrik": MavvrikFocusLogger,
"vantage": VantageLogger,
"posthog": PostHogLogger,
"newrelic": NewRelicLogger,
}
try:

View file

@ -131,6 +131,8 @@ def get_next_standardized_reset_time(
# Handle different time units
if unit == "d":
return _handle_day_reset(current_time, base_midnight, value, tz)
elif unit == "w":
return _handle_day_reset(current_time, base_midnight, value * 7, tz)
elif unit == "h":
return _handle_hour_reset(current_time, base_midnight, value)
elif unit == "m":

View file

@ -4,7 +4,7 @@ from litellm.llms.openai.data_residency import infer_openai_data_residency
# Pre-define optional kwargs keys as frozenset for O(1) lookups
# These are extracted from kwargs only if present, avoiding unnecessary .get() calls
_OPTIONAL_KWARGS_KEYS = frozenset(
OPTIONAL_KWARGS_KEYS = frozenset(
{
"azure_ad_token",
"tenant_id",
@ -39,6 +39,9 @@ _OPTIONAL_KWARGS_KEYS = frozenset(
}
)
# Backward-compatible alias for existing imports/tests.
_OPTIONAL_KWARGS_KEYS = OPTIONAL_KWARGS_KEYS
def _get_base_model_from_litellm_call_metadata(
metadata: Optional[dict],
@ -166,7 +169,7 @@ def get_litellm_params(
# Sparse extraction: only add kwargs keys that are actually present
if kwargs:
for key in _OPTIONAL_KWARGS_KEYS:
for key in OPTIONAL_KWARGS_KEYS:
if key in kwargs:
litellm_params[key] = kwargs[key]

View file

@ -1,5 +1,5 @@
import re
from typing import Optional, Tuple
from typing import Optional, Tuple, cast
from urllib.parse import urlparse
import litellm
@ -7,7 +7,7 @@ from litellm.constants import REPLICATE_MODEL_NAME_WITH_ID_LENGTH
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
from litellm.secret_managers.main import get_secret, get_secret_str
from ..types.router import LiteLLM_Params
from ..types.router import GenericLiteLLMParams, LiteLLM_Params
def _endpoint_matches_api_base(endpoint: str, api_base: str) -> bool:
@ -159,7 +159,7 @@ def get_llm_provider( # noqa: PLR0915
custom_llm_provider: Optional[str] = None,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
litellm_params: Optional[LiteLLM_Params] = None,
litellm_params: Optional[GenericLiteLLMParams] = None,
) -> Tuple[str, str, Optional[str], Optional[str]]:
"""
Returns the provider for a given model name - e.g. 'azure/chatgpt-v-2' -> 'azure'
@ -178,7 +178,7 @@ def get_llm_provider( # noqa: PLR0915
)
if litellm.LiteLLMProxyChatConfig._should_use_litellm_proxy_by_default(
litellm_params=litellm_params
litellm_params=cast(Optional[LiteLLM_Params], litellm_params)
):
return litellm.LiteLLMProxyChatConfig.litellm_proxy_get_custom_llm_provider_info(
model=model, api_base=api_base, api_key=api_key
@ -186,12 +186,10 @@ def get_llm_provider( # noqa: PLR0915
## IF LITELLM PARAMS GIVEN ##
if litellm_params:
assert (
custom_llm_provider is None and api_base is None and api_key is None
), "Either pass in litellm_params or the custom_llm_provider/api_base/api_key. Otherwise, these values will be overriden."
custom_llm_provider = litellm_params.custom_llm_provider
api_base = litellm_params.api_base
api_key = litellm_params.api_key
if custom_llm_provider is None and api_base is None and api_key is None:
custom_llm_provider = litellm_params.custom_llm_provider
api_base = litellm_params.api_base
api_key = litellm_params.api_key
dynamic_api_key = None
# check if llm provider provided
@ -235,6 +233,7 @@ def get_llm_provider( # noqa: PLR0915
api_base=api_base,
api_key=api_key,
dynamic_api_key=dynamic_api_key,
litellm_params=litellm_params,
)
# check if llm provider part of model name
@ -250,6 +249,7 @@ def get_llm_provider( # noqa: PLR0915
api_base=api_base,
api_key=api_key,
dynamic_api_key=dynamic_api_key,
litellm_params=litellm_params,
)
elif model.split("/", 1)[0] in litellm.provider_list:
custom_llm_provider = model.split("/", 1)[0]
@ -570,6 +570,7 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
api_base: Optional[str],
api_key: Optional[str],
dynamic_api_key: Optional[str],
litellm_params: Optional[GenericLiteLLMParams] = None,
) -> Tuple[str, str, Optional[str], Optional[str]]:
"""
Returns:
@ -637,7 +638,7 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
api_base,
dynamic_api_key,
) = litellm.BedrockMantleChatConfig()._get_openai_compatible_provider_info(
api_base, api_key
api_base, api_key, litellm_params=litellm_params
)
elif custom_llm_provider == "nvidia_nim":
# nvidia_nim is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.endpoints.anyscale.com/v1

View file

@ -295,6 +295,15 @@ def get_supported_openai_params( # noqa: PLR0915
elif custom_llm_provider == "predibase":
return litellm.PredibaseConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "voyage":
if (
request_type == "embeddings"
and litellm.VoyageMultimodalEmbeddingConfig.is_multimodal_embeddings(model)
):
return (
litellm.VoyageMultimodalEmbeddingConfig().get_supported_openai_params(
model=model
)
)
return litellm.VoyageEmbeddingConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "infinity":
return litellm.InfinityEmbeddingConfig().get_supported_openai_params(

View file

@ -53,11 +53,19 @@ _supported_callback_params = [
"braintrust_host",
"slack_webhook_url",
"lunary_public_key",
"dd_api_key",
"dd_site",
"dd_agent_host",
"dd_agent_port",
]
_request_blocked_callback_params = {
"gcs_bucket_name",
"gcs_path_service_account",
"dd_api_key",
"dd_site",
"dd_agent_host",
"dd_agent_port",
}

View file

@ -158,6 +158,7 @@ from ..integrations.litellm_agent import LiteLLMAgentModelResolver
from ..integrations.literal_ai import LiteralAILogger
from ..integrations.logfire_logger import LogfireLevel, LogfireLogger
from ..integrations.lunary import LunaryLogger
from ..integrations.newrelic import NewRelicLogger
from ..integrations.openmeter import OpenMeterLogger
from ..integrations.opik.opik import OpikLogger
from ..integrations.posthog import PostHogLogger
@ -380,13 +381,14 @@ class Logging(LiteLLMLoggingBaseClass):
List[Union[str, Callable, CustomLogger]]
] = dynamic_async_failure_callbacks
# Process dynamic callbacks
self.process_dynamic_callbacks()
## DYNAMIC LANGFUSE / GCS / logging callback KEYS ##
self.standard_callback_dynamic_params: StandardCallbackDynamicParams = (
self.initialize_standard_callback_dynamic_params(kwargs)
)
# Process dynamic callbacks (after standard_callback_dynamic_params is initialized,
# so team-scoped credentials are available for callback initialization)
self.process_dynamic_callbacks()
self.standard_built_in_tools_params: StandardBuiltInToolsParams = (
self.initialize_standard_built_in_tools_params(kwargs)
)
@ -481,8 +483,21 @@ class Logging(LiteLLMLoggingBaseClass):
isinstance(callback, str)
and callback in litellm._known_custom_logger_compatible_callbacks
):
# For callbacks that support team-scoped credentials (e.g. datadog),
# pass only the relevant dynamic params as custom_logger_init_args.
_custom_logger_init_args: Optional[dict] = None
if callback == "datadog":
_custom_logger_init_args = {
k: v
for k, v in self.standard_callback_dynamic_params.items()
if k.startswith("dd_")
}
callback_class = _init_custom_logger_compatible_class(
callback, internal_usage_cache=None, llm_router=None # type: ignore
callback, # type: ignore[arg-type]
internal_usage_cache=None,
llm_router=None, # type: ignore
custom_logger_init_args=_custom_logger_init_args,
)
if callback_class is not None:
processed_list.append(callback_class)
@ -3507,9 +3522,7 @@ class Logging(LiteLLMLoggingBaseClass):
else:
return None
def _handle_anthropic_messages_response_logging(
self, result: Any
) -> Union[ModelResponse, ResponsesAPIResponse]:
def _handle_anthropic_messages_response_logging(self, result: Any) -> ModelResponse:
"""
Handles logging for Anthropic messages responses.
@ -3528,15 +3541,14 @@ class Logging(LiteLLMLoggingBaseClass):
return result
elif isinstance(result, ModelResponse):
return result
elif isinstance(
if isinstance(
result,
(ResponseCompletedEvent, ResponseIncompleteEvent, ResponseFailedEvent),
):
# anthropic_messages() can route to OpenAI Responses API; in that path
# the assembled streaming result is one of these terminal events rather than
# a ModelResponse. Return the inner response so downstream handlers
# (_transform_usage_objects, normalize_logging_result) can process it.
return result.response
result = result.response
if isinstance(result, ResponsesAPIResponse):
return self._translate_responses_api_response_to_model_response(result)
httpx_response = self.model_call_details.get("httpx_response", None)
if httpx_response and isinstance(httpx_response, httpx.Response):
@ -3570,6 +3582,55 @@ class Logging(LiteLLMLoggingBaseClass):
)
return result
def _translate_responses_api_response_to_model_response(
self, result: ResponsesAPIResponse
) -> ModelResponse:
"""
Convert a Responses API response into a ModelResponse for spend_logs.
The proxy UI parses spend_log rows expecting chat-completion shape
(response.choices[0].message); a raw ResponsesAPIResponse dump (output[...])
would render as empty in the Logs tab. Translation also yields full
choices/message detail downstream consumers can rely on.
"""
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
)
try:
return LiteLLMResponsesTransformationHandler().transform_response(
model=self.model,
raw_response=result,
model_response=litellm.ModelResponse(),
logging_obj=self,
request_data={},
messages=[],
optional_params={},
litellm_params={},
encoding=litellm.encoding,
)
except Exception as e:
verbose_logger.debug(
"Responses API -> ModelResponse translation failed for "
"anthropic_messages logging (%s); falling back to minimal "
"usage-only ModelResponse to keep the spend_logs row.",
str(e),
)
model_response = litellm.ModelResponse()
model_response.model = self.model
usage = getattr(result, "usage", None)
if usage is not None and ResponseAPILoggingUtils._is_response_api_usage(
usage
):
setattr(
model_response,
"usage",
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
),
)
return model_response
def _handle_non_streaming_google_genai_generate_content_response_logging(
self, result: Any
) -> ModelResponse:
@ -3894,6 +3955,24 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
_in_memory_loggers.append(_prometheus_logger)
return _prometheus_logger # type: ignore
elif logging_integration == "datadog":
# Check if team-scoped credentials are provided
_dd_api_key = custom_logger_init_args.get("dd_api_key")
_dd_site = custom_logger_init_args.get("dd_site")
_dd_agent_host = custom_logger_init_args.get("dd_agent_host")
_dd_agent_port = custom_logger_init_args.get("dd_agent_port")
if _dd_api_key or _dd_site or _dd_agent_host:
# Team-scoped credentials: use DynamicLoggingCache for per-credential isolation
from litellm.integrations.datadog.datadog_team_handler import (
DataDogHandler,
)
return DataDogHandler.get_datadog_logger_for_request(
standard_callback_dynamic_params=custom_logger_init_args, # type: ignore
in_memory_dynamic_logger_cache=in_memory_dynamic_logger_cache,
)
# Global (env-var based): reuse cached instance
for callback in _in_memory_loggers:
if isinstance(callback, DataDogLogger):
return callback # type: ignore
@ -4124,6 +4203,17 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
focus_logger = FocusLogger()
_in_memory_loggers.append(focus_logger)
return focus_logger # type: ignore
elif logging_integration == "mavvrik":
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
MavvrikFocusLogger,
)
for callback in _in_memory_loggers:
if type(callback) is MavvrikFocusLogger:
return callback # type: ignore
mavvrik_focus_logger = MavvrikFocusLogger()
_in_memory_loggers.append(mavvrik_focus_logger)
return mavvrik_focus_logger # type: ignore
elif logging_integration == "vantage":
from litellm.integrations.vantage.vantage_logger import VantageLogger
@ -4419,6 +4509,13 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
gitlab_logger = GitLabPromptManager(gitlab_config=gitlab_config)
_in_memory_loggers.append(gitlab_logger)
return gitlab_logger # type: ignore
elif logging_integration == "newrelic":
for callback in _in_memory_loggers:
if isinstance(callback, NewRelicLogger):
return callback # type: ignore
newrelic_logger = NewRelicLogger()
_in_memory_loggers.append(newrelic_logger)
return newrelic_logger # type: ignore
return None
except Exception as e:
verbose_logger.exception(
@ -4720,6 +4817,10 @@ def get_custom_logger_compatible_class( # noqa: PLR0915
for callback in _in_memory_loggers:
if isinstance(callback, SMTPEmailLogger):
return callback
elif logging_integration == "newrelic":
for callback in _in_memory_loggers:
if isinstance(callback, NewRelicLogger):
return callback
return None
except Exception as e:

View file

@ -949,6 +949,43 @@ def calculate_image_response_cost_from_usage(
return prompt_cost + completion_cost
def calculate_image_response_web_search_cost(
image_response: ImageResponse,
custom_llm_provider: str,
model_info: ModelInfo,
) -> float:
"""
Cost of Google Search grounding performed during image generation.
The grounding request count is carried on the image usage object by the
provider transformers; it is billed with the same per-request accounting
used for chat completions.
"""
usage = image_response.usage
if usage is None:
return 0.0
web_search_requests = getattr(usage, "web_search_requests", None)
if not web_search_requests:
return 0.0
from litellm.llms import get_cost_for_web_search_request
synthetic_usage = Usage(
prompt_tokens_details=PromptTokensDetailsWrapper(
web_search_requests=web_search_requests
)
)
return (
get_cost_for_web_search_request(
custom_llm_provider=custom_llm_provider,
usage=synthetic_usage,
model_info=model_info,
)
or 0.0
)
class CostCalculatorUtils:
@staticmethod
def _call_type_has_image_response(call_type: str) -> bool:

View file

@ -4741,6 +4741,12 @@ class BedrockConverseMessagesProcessor:
guardContent={"text": {"text": element["text"]}}
)
_parts.append(_part)
elif element["type"] in ("grounding_source", "query"):
# Contextual grounding tags are guardrail metadata; the
# model only needs the underlying text, so render them
# as plain text on the generate path.
_part = BedrockContentBlock(text=element["text"])
_parts.append(_part)
elif element["type"] == "image_url":
format: Optional[str] = None
if isinstance(element["image_url"], dict):
@ -5173,6 +5179,12 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915
guardContent={"text": {"text": element["text"]}}
)
_parts.append(_part)
elif element["type"] in ("grounding_source", "query"):
# Contextual grounding tags are guardrail metadata; the
# model only needs the underlying text, so render them as
# plain text on the generate path.
_part = BedrockContentBlock(text=element["text"])
_parts.append(_part)
elif element["type"] == "image_url":
format: Optional[str] = None
if isinstance(element["image_url"], dict):

View file

@ -47,6 +47,7 @@ class RealTimeStreaming:
user_api_key_dict: Optional[Any] = None,
request_data: Optional[Dict] = None,
backend_uses_beta_protocol: Optional[bool] = None,
force_transcription_model: Optional[str] = None,
):
self.websocket = websocket
self.backend_ws = backend_ws
@ -100,6 +101,11 @@ class RealTimeStreaming:
self._flushing_pending_messages_until_setup: bool = False
self._pending_messages_until_setup: List[str] = []
self._pending_messages_byte_total: int = 0
# Whether this is a transcription-only session (session.type == "transcription",
# e.g. gpt-realtime-whisper). Such sessions must not be sent response.create and
# their input_audio_transcription.completed usage drives duration-based cost.
self._force_transcription_model = force_transcription_model
self._is_transcription_session: bool = force_transcription_model is not None
# Per-connection caps for pre-setup audio frames (message count + total bytes).
_MAX_BUFFERED_MESSAGES: int = 200
@ -111,8 +117,12 @@ class RealTimeStreaming:
"input_audio_buffer.append",
"input_audio_buffer.commit",
"input_audio_buffer.clear",
"input_audio_buffer.end",
]
)
_CLIENT_AUDIO_BUFFER_COMMIT_TYPES = frozenset(
["input_audio_buffer.commit", "input_audio_buffer.end"]
)
_AUDIO_FORMAT_MAP: Dict[str, Dict[str, Any]] = {
"pcm16": {"type": "audio/pcm", "rate": 24000},
"g711_ulaw": {"type": "audio/G711-ulaw", "rate": 8000},
@ -207,6 +217,8 @@ class RealTimeStreaming:
self.session_tools = tools
# GA: session.type is required; log it for traceability but no action needed
verbose_logger.debug(f"Realtime session.type: {session.get('type')}")
if session.get("type") == "transcription":
self._is_transcription_session = True
except (json.JSONDecodeError, AttributeError, TypeError):
pass
@ -223,6 +235,55 @@ class RealTimeStreaming:
except (AttributeError, TypeError):
pass
def _detect_transcription_session_from_backend(
self, event_obj: Union[dict, OpenAIRealtimeEvents]
) -> None:
"""Flag transcription-only sessions from backend session events."""
try:
event_type = event_obj.get("type", "")
if event_type in (
"transcription_session.created",
"transcription_session.updated",
):
self._is_transcription_session = True
elif event_type in ("session.created", "session.updated"):
session = cast(dict, event_obj).get("session", {}) or {}
if session.get("type") == "transcription":
self._is_transcription_session = True
except (AttributeError, TypeError):
pass
def _capture_transcription_usage(
self, event_obj: Union[dict, OpenAIRealtimeEvents]
) -> None:
"""
Append a usage-only transcription completed event to the logged results so
the cost calculator can bill it by audio duration. The default logged event
types exclude this event, so it is captured here directly for transcription
sessions rather than widening logging for every realtime session. Only the
type and usage are kept — the transcript is already captured separately in
input_messages, so it is not duplicated into the response log here.
"""
try:
usage = event_obj.get("usage")
if usage is None:
return
# If this event type is already captured by store_message (e.g. the user
# logs all realtime events), don't append a second copy.
if self._should_store_message(event_obj):
return
self.messages.append(
cast(
OpenAIRealtimeEvents,
{
"type": "conversation.item.input_audio_transcription.completed",
"usage": usage,
},
)
)
except (AttributeError, TypeError):
pass
def _collect_tool_calls_from_response_done(
self, event_obj: Union[dict, OpenAIRealtimeEvents]
) -> None:
@ -283,6 +344,7 @@ class RealTimeStreaming:
backend, False if the provider transformation produced no output and
the message was effectively dropped.
"""
message = self._enforce_transcription_session_model(message)
if self.provider_config:
transformed = self.provider_config.transform_realtime_request(
message, self.model, self.session_configuration_request
@ -302,12 +364,128 @@ class RealTimeStreaming:
await self.backend_ws.send(message) # type: ignore[union-attr, attr-defined]
return True
def _enforce_transcription_session_model(self, message: str) -> str:
"""Force client transcription session updates to the authorized model.
`/v1/realtime?intent=transcription` may intentionally omit `model` from
the upstream URL for Azure compatibility, but the proxy still authorizes
a resolved LiteLLM model before opening the backend websocket. If a
client later sends a transcription `session.update`, any model embedded
in that update must be rewritten to the same authorized model instead of
allowing a post-auth model/deployment switch.
Normal realtime sessions keep their independent nested transcription
model behavior because `_force_transcription_model` is only set for
transcription-intent websocket routes.
"""
if self._force_transcription_model is None:
return message
try:
message_obj = json.loads(message)
except (json.JSONDecodeError, TypeError):
return message
if message_obj.get("type") not in (
"session.update",
"transcription_session.update",
):
return message
session = message_obj.get("session")
if not isinstance(session, dict):
return message
if session.get("type") == "transcription":
self._is_transcription_session = True
authorized_model = self._force_transcription_model
changed = False
transcription = session.get("input_audio_transcription")
if (
isinstance(transcription, dict)
and transcription.get("model") != authorized_model
):
session["input_audio_transcription"] = {
**transcription,
"model": authorized_model,
}
changed = True
audio = session.get("audio")
if isinstance(audio, dict):
audio_input = audio.get("input")
if isinstance(audio_input, dict):
nested_transcription = audio_input.get("transcription")
if (
isinstance(nested_transcription, dict)
and nested_transcription.get("model") != authorized_model
):
session["audio"] = {
**audio,
"input": {
**audio_input,
"transcription": {
**nested_transcription,
"model": authorized_model,
},
},
}
changed = True
if not changed:
return message
return json.dumps(message_obj)
def _uses_deferred_backend_setup(self) -> bool:
"""True when setup is deferred until the client's first session.update."""
if self.provider_config is None:
return False
return not self.provider_config.requires_session_configuration()
@staticmethod
def _collapse_buffered_audio_messages(messages: List[str]) -> List[str]:
"""Apply ``input_audio_buffer.clear`` semantics before replaying buffered frames.
During deferred Gemini Live setup, ``clear`` is buffered alongside appends.
On flush each append becomes a provider ``realtimeInput``; ``clear`` must
drop preceding uncommitted appends instead of being forwarded as a no-op.
"""
collapsed: List[str] = []
pending_appends: List[str] = []
for message in messages:
try:
msg_type = json.loads(message).get("type")
except (json.JSONDecodeError, TypeError):
collapsed.extend(pending_appends)
pending_appends = []
collapsed.append(message)
continue
if msg_type == "input_audio_buffer.append":
pending_appends.append(message)
elif msg_type == "input_audio_buffer.clear":
pending_appends = []
elif msg_type in RealTimeStreaming._CLIENT_AUDIO_BUFFER_COMMIT_TYPES:
collapsed.extend(pending_appends)
pending_appends = []
collapsed.append(message)
else:
collapsed.extend(pending_appends)
pending_appends = []
collapsed.append(message)
collapsed.extend(pending_appends)
return collapsed
def _sync_pending_messages_byte_total(self) -> None:
self._pending_messages_byte_total = sum(
len(message.encode("utf-8"))
for message in self._pending_messages_until_setup
)
def _should_buffer_client_message_until_setup(self, message: str) -> bool:
if not self._uses_deferred_backend_setup():
return False
@ -323,6 +501,18 @@ class RealTimeStreaming:
return msg_obj.get("type") in RealTimeStreaming._CLIENT_AUDIO_BUFFER_TYPES
def _buffer_pending_message_until_setup(self, message: str) -> None:
try:
msg_type = json.loads(message).get("type")
except (json.JSONDecodeError, TypeError):
msg_type = None
if msg_type == "input_audio_buffer.clear":
self._pending_messages_until_setup = self._collapse_buffered_audio_messages(
self._pending_messages_until_setup + [message]
)
self._sync_pending_messages_byte_total()
return
msg_bytes = len(message.encode("utf-8"))
if (
len(self._pending_messages_until_setup)
@ -340,7 +530,9 @@ class RealTimeStreaming:
)
async def _flush_pending_messages_until_setup(self) -> bool:
pending = self._pending_messages_until_setup
pending = self._collapse_buffered_audio_messages(
self._pending_messages_until_setup
)
self._pending_messages_until_setup = []
self._pending_messages_byte_total = 0
for idx, message in enumerate(pending):
@ -425,16 +617,13 @@ class RealTimeStreaming:
if sent:
self._guardrail_turn_detection_update_sent = True
def _has_realtime_guardrails(self) -> bool:
"""Return True if any callback is registered for realtime guardrail event types."""
def _has_realtime_guardrails_for_event_hooks(
self,
event_hooks: List[Any],
) -> bool:
"""Return True if any callback would run for one of ``event_hooks``."""
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import GuardrailEventHooks
_realtime_event_types = [
GuardrailEventHooks.realtime_input_transcription,
GuardrailEventHooks.pre_call,
GuardrailEventHooks.post_call,
]
return any(
isinstance(cb, CustomGuardrail)
and any(
@ -442,31 +631,45 @@ class RealTimeStreaming:
data=self.request_data,
event_type=et,
)
for et in _realtime_event_types
for et in event_hooks
)
for cb in litellm.callbacks
)
def _has_realtime_guardrails(self) -> bool:
"""Return True if any callback is registered for realtime guardrail event types."""
from litellm.types.guardrails import GuardrailEventHooks
return self._has_realtime_guardrails_for_event_hooks(
[
GuardrailEventHooks.realtime_input_transcription,
GuardrailEventHooks.pre_call,
GuardrailEventHooks.post_call,
]
)
def _has_audio_transcription_guardrails(self) -> bool:
"""Return True if any callback needs to run on audio transcriptions (VAD path).
"""Return True when a guardrail is configured for the audio/VAD transcript path.
When this returns True, we inject a session.update to disable the LLM's
auto-response so the guardrail can gate it first.
Must match the same hook criteria as run_realtime_guardrails() so that
any guardrail that would actually check the transcript also disables
auto-response before the transcript arrives.
Only ``realtime_input_transcription`` hooks disable ``server_vad`` auto-response.
``pre_call`` / ``post_call`` guardrails (e.g. Model Armor on chat completions)
must not override ``turn_detection.create_response`` on realtime sessions.
"""
return self._has_realtime_guardrails()
from litellm.types.guardrails import GuardrailEventHooks
return self._has_realtime_guardrails_for_event_hooks(
[GuardrailEventHooks.realtime_input_transcription]
)
async def run_realtime_guardrails(
self,
transcript: str,
item_id: Optional[str] = None,
pre_block_backend_message: Optional[str] = None,
event_hooks: Optional[List[Any]] = None,
) -> bool:
"""
Run registered guardrails on a completed speech transcription.
Run registered guardrails on realtime text (transcript, user message, tool output).
Returns True if blocked (synthetic warning already sent to client).
Returns False if clean (caller should send response.create to the backend).
@ -477,15 +680,17 @@ class RealTimeStreaming:
specific message to be sent first — e.g. Gemini Live requires a
matching ``toolResponse`` immediately after a ``toolCall`` before any
other client messages can be accepted.
``event_hooks`` selects which guardrail modes to evaluate. Audio/VAD
transcript completion uses ``realtime_input_transcription`` only;
typed user messages and tool outputs use ``pre_call``.
"""
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import GuardrailEventHooks
_realtime_event_types = [
GuardrailEventHooks.realtime_input_transcription,
GuardrailEventHooks.pre_call,
GuardrailEventHooks.post_call,
]
if event_hooks is None:
event_hooks = [GuardrailEventHooks.realtime_input_transcription]
_realtime_event_types = event_hooks
_check_data = {**self.request_data, "transcript": transcript}
_already_run: set = set()
@ -719,6 +924,8 @@ class RealTimeStreaming:
"""
event_type = event_obj.get("type")
self._detect_transcription_session_from_backend(event_obj)
# Send session.created to the client FIRST so it stays in sync, then inject
# the disable-auto-response session.update; otherwise a backend error could
# reach the client before it sees session.created.
@ -736,6 +943,14 @@ class RealTimeStreaming:
self._collect_user_input_from_backend_event(event_obj)
self.store_message(event_obj)
await self.websocket.send_text(raw_response)
# Transcription-only sessions (e.g. gpt-realtime-whisper) have no
# assistant turn: capture audio-duration usage for cost and never
# trigger response.create.
if self._is_transcription_session:
self._capture_transcription_usage(event_obj)
return True
blocked = await self.run_realtime_guardrails(
transcript,
item_id=event_obj.get("item_id"),
@ -992,6 +1207,8 @@ class RealTimeStreaming:
guardrail_turn_detection_injected = False
msg_type: Optional[str] = None
try:
from litellm.types.guardrails import GuardrailEventHooks
msg_obj = json.loads(message)
msg_type = msg_obj.get("type")
@ -1044,6 +1261,7 @@ class RealTimeStreaming:
blocked = await self.run_realtime_guardrails(
output_text,
pre_block_backend_message=sanitized_msg,
event_hooks=[GuardrailEventHooks.pre_call],
)
if blocked:
# ``_pending_guardrail_message`` is
@ -1069,7 +1287,8 @@ class RealTimeStreaming:
combined_text = " ".join(texts)
if combined_text:
blocked = await self.run_realtime_guardrails(
combined_text
combined_text,
event_hooks=[GuardrailEventHooks.pre_call],
)
if blocked:
# Store the guardrail reason so the next response.create

View file

@ -469,12 +469,16 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
if should_start_new_block and not self.sent_content_block_finish:
# Queue the sequence: content_block_stop -> content_block_start
# For text blocks the trigger chunk is not emitted as a separate
# delta because content_block_start carries the information.
# For tool_use blocks we must also emit the trigger chunk's delta
# when it carries input_json_delta data, because some providers
# (e.g. xAI, Gemini) include tool arguments in the same streaming
# chunk as the function name/id.
# -> (optionally) the trigger chunk's delta.
#
# The synthesized content_block_start always carries an
# empty body, so the chunk that *triggered* the transition
# also carries the new block's first delta. It must be
# re-emitted or the first token of the new block is lost.
# This applies to text_delta and thinking_delta (the first
# non-empty text/thinking token) as well as input_json_delta
# (providers like xAI/Gemini bundle tool arguments with the
# function name/id in a single chunk).
# 1. Stop current content block
self.chunk_queue.append(
@ -493,14 +497,9 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
}
)
# 3. If the trigger chunk carries tool argument data, queue it
# so the input_json_delta is not silently dropped.
if (
processed_chunk.get("type") == "content_block_delta"
and isinstance(processed_chunk.get("delta"), dict)
and processed_chunk["delta"].get("type") == "input_json_delta"
and processed_chunk["delta"].get("partial_json")
):
# 3. If the trigger chunk carries delta content, queue it
# so the first delta of the new block is not silently dropped.
if self._trigger_delta_has_content(processed_chunk):
self.chunk_queue.append(processed_chunk)
self.sent_content_block_finish = False
@ -711,12 +710,16 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
if not self.queued_usage_chunk:
if should_start_new_block and not self.sent_content_block_finish:
# Queue the sequence: content_block_stop -> content_block_start
# For text blocks the trigger chunk is not emitted as a separate
# delta because content_block_start carries the information.
# For tool_use blocks we must also emit the trigger chunk's delta
# when it carries input_json_delta data, because some providers
# (e.g. xAI, Gemini) include tool arguments in the same streaming
# chunk as the function name/id.
# -> (optionally) the trigger chunk's delta.
#
# The synthesized content_block_start always carries an
# empty body, so the chunk that *triggered* the transition
# also carries the new block's first delta. It must be
# re-emitted or the first token of the new block is lost.
# This applies to text_delta and thinking_delta (the
# first non-empty text/thinking token) as well as
# input_json_delta (providers like xAI/Gemini bundle tool
# arguments with the function name/id in a single chunk).
# 1. Stop current content block
self.chunk_queue.append(
@ -733,15 +736,9 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
}
)
# 3. If the trigger chunk carries tool argument data, queue it
# so the input_json_delta is not silently dropped.
if (
processed_chunk.get("type") == "content_block_delta"
and isinstance(processed_chunk.get("delta"), dict)
and processed_chunk["delta"].get("type")
== "input_json_delta"
and processed_chunk["delta"].get("partial_json")
):
# 3. If the trigger chunk carries delta content, queue it
# so the first delta of the new block is not silently dropped.
if self._trigger_delta_has_content(processed_chunk):
self.chunk_queue.append(processed_chunk)
# Reset state for new block
@ -898,6 +895,38 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
def _increment_content_block_index(self):
self.current_content_block_index += 1
@staticmethod
def _trigger_delta_has_content(processed_chunk: Dict[str, Any]) -> bool:
"""Return True if a translated trigger chunk carries a non-empty
``content_block_delta`` payload that must be re-emitted after a
block transition.
When an upstream chunk both *triggers* a new content block (its type
differs from the active block) and *carries* delta content, that
content belongs to the new block. The synthesized
``content_block_start`` only ever carries an empty body — see
``_translate_streaming_openai_chunk_to_anthropic_content_block``,
which returns an empty ``TextBlock``/``ToolUseBlock``/thinking block —
so the trigger chunk's delta must be re-queued or the first token of
the new block (the first non-empty text/thinking delta, or bundled
tool arguments) is silently dropped.
"""
if processed_chunk.get("type") != "content_block_delta":
return False
delta = processed_chunk.get("delta")
if not isinstance(delta, dict):
return False
delta_type = delta.get("type")
if delta_type == "text_delta":
return bool(delta.get("text"))
if delta_type == "input_json_delta":
return bool(delta.get("partial_json"))
if delta_type == "thinking_delta":
return bool(delta.get("thinking"))
if delta_type == "signature_delta":
return bool(delta.get("signature"))
return False
def _should_start_new_content_block(self, chunk: "ModelResponseStream") -> bool:
"""
Determine if we should start a new content block based on the processed chunk.

View file

@ -102,9 +102,9 @@ def _build_responses_kwargs(
from litellm.types.utils import CallTypes
if isinstance(value, LiteLLMLoggingObject):
# Reclassify as acompletion so the success handler doesn't try to
# validate the Responses API event as an AnthropicResponse.
# (Mirrors the pattern used in LiteLLMMessagesToCompletionTransformationHandler.)
# Keep call_type as anthropic_messages so spend_logs are billed
# against /v1/messages; the success handler translates the
# Responses API result back to a ModelResponse for the row.
setattr(value, "call_type", CallTypes.anthropic_messages.value)
responses_kwargs[key] = value
elif key not in excluded and key not in responses_kwargs and value is not None:

View file

@ -8,6 +8,7 @@ from typing import Any, Optional, cast
from litellm._logging import _redact_string, verbose_proxy_logger
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.types.realtime import RealtimeQueryParams
from ....litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from ....litellm_core_utils.realtime_streaming import RealTimeStreaming
@ -35,6 +36,7 @@ class AzureOpenAIRealtime(AzureChatCompletion):
model: str,
api_version: Optional[str],
realtime_protocol: Optional[str] = None,
query_params: Optional[RealtimeQueryParams] = None,
) -> str:
"""
Construct Azure realtime WebSocket URL.
@ -46,6 +48,7 @@ class AzureOpenAIRealtime(AzureChatCompletion):
realtime_protocol: Protocol version to use:
- "GA" or "v1": Uses /openai/v1/realtime (GA path)
- "beta" or None: Uses /openai/realtime (beta path, default)
query_params: Extra query params to forward (e.g. intent=transcription).
Returns:
WebSocket URL string
@ -54,6 +57,8 @@ class AzureOpenAIRealtime(AzureChatCompletion):
beta/default: "wss://.../openai/realtime?api-version=2024-10-01-preview&deployment=gpt-4o-realtime-preview"
GA/v1: "wss://.../openai/v1/realtime?model=gpt-realtime-deployment"
"""
from urllib.parse import urlencode
api_base = api_base.replace("https://", "wss://")
# Determine path based on realtime_protocol (case-insensitive)
@ -61,13 +66,25 @@ class AzureOpenAIRealtime(AzureChatCompletion):
"GA",
"V1",
)
intent = (query_params or {}).get("intent")
if _is_ga:
path = "/openai/v1/realtime"
return f"{api_base}{path}?model={model}"
query_parts = []
if intent != "transcription" and (
query_params is None or "model" in query_params
):
query_parts.append(urlencode({"model": model}))
else:
# Default to beta path for backwards compatibility
path = "/openai/realtime"
return f"{api_base}{path}?api-version={api_version}&deployment={model}"
query_parts = [urlencode({"api-version": api_version, "deployment": model})]
if intent:
query_parts.append(urlencode({"intent": intent}))
qs = "&".join(query_parts)
return f"{api_base}{path}?{qs}" if qs else f"{api_base}{path}"
async def async_realtime(
self,
@ -81,6 +98,7 @@ class AzureOpenAIRealtime(AzureChatCompletion):
client: Optional[Any] = None,
timeout: Optional[float] = None,
realtime_protocol: Optional[str] = None,
query_params: Optional[RealtimeQueryParams] = None,
user_api_key_dict: Optional[Any] = None,
litellm_metadata: Optional[dict] = None,
):
@ -96,7 +114,11 @@ class AzureOpenAIRealtime(AzureChatCompletion):
raise ValueError("api_version is required for Azure OpenAI calls")
url = self._construct_url(
api_base, model, api_version, realtime_protocol=realtime_protocol
api_base,
model,
api_version,
realtime_protocol=realtime_protocol,
query_params=query_params,
)
try:
@ -113,9 +135,15 @@ class AzureOpenAIRealtime(AzureChatCompletion):
websocket,
cast(ClientConnection, backend_ws),
logging_obj,
model=model,
user_api_key_dict=user_api_key_dict,
request_data={"litellm_metadata": litellm_metadata or {}},
backend_uses_beta_protocol=backend_uses_beta_protocol,
force_transcription_model=(
model
if (query_params or {}).get("intent") == "transcription"
else None
),
)
await realtime_streaming.bidirectional_forward()

View file

@ -40,6 +40,13 @@ class AzureRealtimeHTTPConfig(BaseRealtimeHTTPConfig):
version = api_version or get_secret_str("AZURE_API_VERSION") or "2024-12-17"
return f"{base}/openai/realtime/calls?api-version={version}"
def get_transcription_session_url(
self, api_base: Optional[str], model: str, api_version: Optional[str] = None
) -> str:
base = self.get_api_base(api_base).rstrip("/")
version = api_version or get_secret_str("AZURE_API_VERSION") or "2024-12-17"
return f"{base}/openai/realtime/transcription_sessions?api-version={version}"
def get_realtime_calls_headers(self, ephemeral_key: str) -> dict:
return {
"api-key": ephemeral_key,

View file

@ -185,6 +185,40 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
default_api_version=AZURE_DEFAULT_RESPONSES_API_VERSION,
)
def supports_native_websocket(self) -> bool:
return True
def get_websocket_url(
self,
api_base: Optional[str],
litellm_params: dict,
) -> str:
"""
Azure Responses WebSocket endpoint is at /openai/v1/responses with no
api-version query param. Auth is via Authorization header, model is sent
in the response.create body — not the URL.
"""
if api_base is None:
raise ValueError("api_base is required for Azure WebSocket")
parsed_url = httpx.URL(api_base)
path = parsed_url.path.rstrip("/")
# Strip existing /openai/responses path if the api_base already contains it
for suffix in ("/openai/v1/responses", "/openai/responses"):
if path.endswith(suffix):
path = path[: -len(suffix)]
break
scheme = "wss" if parsed_url.scheme == "https" else "ws"
return str(
parsed_url.copy_with(
scheme=scheme, path=f"{path}/openai/v1/responses", query=None
)
)
def model_in_websocket_url(self) -> bool:
# Azure sends the model in the response.create body, not the URL
return False
#########################################################
########## DELETE RESPONSE API TRANSFORMATION ##############
#########################################################

View file

@ -59,6 +59,15 @@ class BaseRealtimeHTTPConfig(ABC):
) -> str:
"""Return the full URL for POST /realtime/client_secrets."""
def get_transcription_session_url(
self, api_base: Optional[str], model: str, api_version: Optional[str] = None
) -> str:
"""Return the full URL for POST /realtime/transcription_sessions."""
base = (api_base or "").rstrip("/")
if base.endswith("/v1"):
base = base[:-3]
return f"{base}/v1/realtime/transcription_sessions"
@abstractmethod
def validate_environment(
self,

View file

@ -258,6 +258,31 @@ class BaseResponsesAPIConfig(ABC):
"""
return False
def get_websocket_url(
self,
api_base: Optional[str],
litellm_params: dict,
) -> str:
"""
Return the wss:// URL for the provider's native Responses WebSocket endpoint.
Defaults to converting the HTTP URL from get_complete_url. Providers whose
WebSocket path differs from their HTTP path (e.g. Azure uses
/openai/v1/responses without api-version) should override this.
"""
http_url = self.get_complete_url(
api_base=api_base, litellm_params=litellm_params
)
return http_url.replace("https://", "wss://").replace("http://", "ws://")
def model_in_websocket_url(self) -> bool:
"""
Return True if the model should be appended as a ?model= query param to
the WebSocket URL. Providers that identify the model via the request body
(e.g. Azure Responses API) should override this to return False.
"""
return True
#########################################################
########## CANCEL RESPONSE API TRANSFORMATION ##########
#########################################################

View file

@ -861,14 +861,58 @@ class BaseAWSLLM:
with tracer.trace("boto3.client(sts)"):
sts_client = boto3.client("sts", **sts_client_kwargs)
# The session policy is an IAM PERMISSION CEILING — effective
# permissions are the intersection of the role's identity policies
# and this policy. Any action not listed here is silently denied
# even when the IAM role grants it. So every Bedrock route we
# support needs a matching action statement, or it 403s on OIDC
# auth only (static creds + IRSA take other code paths).
# https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRoleWithWebIdentity.html
# https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/sts/client/assume_role_with_web_identity.html
bedrock_session_policy = {
"Version": "2012-10-17",
"Statement": [
{
"Sid": "BedrockLiteLLM",
"Effect": "Allow",
"Action": [
"bedrock:InvokeModel",
"bedrock:InvokeModelWithResponseStream",
"bedrock:ApplyGuardrail",
"bedrock:GetGuardrail",
"bedrock:ListGuardrails",
],
"Resource": "*",
"Condition": {"Bool": {"aws:SecureTransport": "true"}},
},
# Claude Platform on AWS (added by #27678 for the
# ``bedrock/claude_platform/<model>`` route) lives under
# a separate IAM action namespace; without these entries
# the OIDC path 403s on every claude_platform request
# even with a fully permissive identity policy (#30200).
{
"Sid": "ClaudePlatformLiteLLM",
"Effect": "Allow",
"Action": [
"aws-external-anthropic:CreateInference",
"aws-external-anthropic:CreateBatchInference",
"aws-external-anthropic:CancelBatchInference",
"aws-external-anthropic:DeleteBatchInference",
"aws-external-anthropic:CountTokens",
"aws-external-anthropic:Get*",
"aws-external-anthropic:List*",
],
"Resource": "*",
"Condition": {"Bool": {"aws:SecureTransport": "true"}},
},
],
}
assume_role_params = {
"RoleArn": aws_role_name,
"RoleSessionName": aws_session_name,
"WebIdentityToken": oidc_token,
"DurationSeconds": 3600,
"Policy": '{"Version":"2012-10-17","Statement":[{"Sid":"BedrockLiteLLM","Effect":"Allow","Action":["bedrock:InvokeModel","bedrock:InvokeModelWithResponseStream","bedrock:ApplyGuardrail","bedrock:GetGuardrail","bedrock:ListGuardrails"],"Resource":"*","Condition":{"Bool":{"aws:SecureTransport":"true"}}}]}',
"Policy": json.dumps(bedrock_session_policy, separators=(",", ":")),
}
# Add ExternalId parameter if provided

View file

@ -32,7 +32,7 @@ def make_sync_call(
logging_obj: LiteLLMLoggingObject,
json_mode: Optional[bool] = False,
fake_stream: bool = False,
stream_chunk_size: int = 1024,
stream_chunk_size: Optional[int] = None,
):
if client is None:
client = _get_httpx_client() # Create a new client if none provided
@ -108,7 +108,7 @@ class BedrockConverseLLM(BaseAWSLLM):
fake_stream: bool = False,
json_mode: Optional[bool] = False,
api_key: Optional[str] = None,
stream_chunk_size: int = 1024,
stream_chunk_size: Optional[int] = None,
) -> CustomStreamWrapper:
request_data = await litellm.AmazonConverseConfig()._async_transform_request(
model=model,
@ -268,7 +268,7 @@ class BedrockConverseLLM(BaseAWSLLM):
):
## SETUP ##
stream = optional_params.pop("stream", None)
stream_chunk_size = optional_params.pop("stream_chunk_size", 1024)
stream_chunk_size = optional_params.pop("stream_chunk_size", None)
unencoded_model_id = optional_params.pop("model_id", None)
fake_stream = optional_params.pop("fake_stream", False)
json_mode = optional_params.get("json_mode", False)

View file

@ -197,7 +197,7 @@ async def make_call(
fake_stream: bool = False,
json_mode: Optional[bool] = False,
bedrock_invoke_provider: Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL] = None,
stream_chunk_size: int = 1024,
stream_chunk_size: Optional[int] = None,
):
try:
if client is None:
@ -294,7 +294,7 @@ def make_sync_call(
fake_stream: bool = False,
json_mode: Optional[bool] = False,
bedrock_invoke_provider: Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL] = None,
stream_chunk_size: int = 1024,
stream_chunk_size: Optional[int] = None,
):
try:
if client is None:
@ -790,7 +790,7 @@ class BedrockLLM(BaseAWSLLM):
## SETUP ##
stream = optional_params.pop("stream", None)
stream_chunk_size = optional_params.pop("stream_chunk_size", 1024)
stream_chunk_size = optional_params.pop("stream_chunk_size", None)
provider = self.get_bedrock_invoke_provider(model)
modelId = self.get_bedrock_model_id(
@ -1203,7 +1203,7 @@ class BedrockLLM(BaseAWSLLM):
extra_headers: Optional[dict] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[AsyncHTTPHandler] = None,
stream_chunk_size: int = 1024,
stream_chunk_size: Optional[int] = None,
) -> Union[ModelResponse, CustomStreamWrapper]:
transformed_request = (
await litellm.AmazonAnthropicClaudeConfig().async_transform_request(
@ -1350,7 +1350,7 @@ class BedrockLLM(BaseAWSLLM):
logger_fn=None,
headers={},
client: Optional[AsyncHTTPHandler] = None,
stream_chunk_size: int = 1024,
stream_chunk_size: Optional[int] = None,
) -> CustomStreamWrapper:
# The call is not made here; instead, we prepare the necessary objects for the stream.

View file

@ -215,6 +215,7 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
anthropic_request.pop("model", None)
anthropic_request.pop("stream", None)
anthropic_request.pop("stream_chunk_size", None)
output_format = anthropic_request.pop("output_format", None)
output_config_format = pop_bedrock_invoke_output_config_format(
anthropic_request

View file

@ -150,6 +150,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
) -> dict:
## SETUP ##
stream = optional_params.pop("stream", None)
optional_params.pop("stream_chunk_size", None)
custom_prompt_dict: dict = litellm_params.pop("custom_prompt_dict", None) or {}
hf_model_name = litellm_params.get("hf_model_name", None)

View file

@ -0,0 +1,5 @@
from litellm.llms.bedrock.passthrough.guardrail_translation.handler import (
BedrockPassthroughGuardrailHandler,
)
__all__ = ["BedrockPassthroughGuardrailHandler"]

View file

@ -0,0 +1,507 @@
from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union
from litellm._logging import verbose_proxy_logger
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.llms.base_llm.guardrail_translation.utils import (
effective_skip_system_message_for_guardrail,
effective_skip_tool_message_for_guardrail,
)
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
_CONVERSE_ACTIONS = frozenset({"converse", "converse-stream"})
_EVENT_STREAM_CONTENT_TYPE = "vnd.amazon.eventstream"
_EVENT_STREAM_MEDIA_TYPE = "application/vnd.amazon.eventstream"
def _is_converse_endpoint(endpoint: str) -> bool:
parts = endpoint.rstrip("/").split("/")
return bool(parts) and parts[-1] in _CONVERSE_ACTIONS
def _generic_passthrough_handler() -> BaseTranslation:
"""
Fallback for non-Converse Bedrock routes (e.g. invoke). The generic
handler scans the full request/response payload so blocking guardrails
still run, matching how other passthrough providers are guarded.
"""
from litellm.llms.pass_through.guardrail_translation.handler import (
PassThroughEndpointHandler,
)
return PassThroughEndpointHandler()
_StringHolder = Tuple[Any, Union[str, int]]
def _collect_strings(node: Any, holders: List[_StringHolder]) -> None:
"""
Record a (container, key) holder for every non-empty string value nested
under an arbitrary JSON node, so prompt content a caller hides in fields
like ``toolUse.input`` or ``toolResult.content[].json`` is still scanned
and can be written back in place. Iterative to avoid unbounded recursion
on deeply nested payloads.
"""
stack: List[Any] = [node]
while stack:
current = stack.pop()
if isinstance(current, dict):
for key, value in current.items():
if isinstance(value, str):
if value:
holders.append((current, key))
else:
stack.append(value)
elif isinstance(current, list):
for index, value in enumerate(current):
if isinstance(value, str):
if value:
holders.append((current, index))
else:
stack.append(value)
def _collect_block_text(block: dict, holders: List[_StringHolder]) -> None:
text = block.get("text")
if isinstance(text, str) and text:
holders.append((block, "text"))
def _extract_converse_texts(
body: dict,
skip_system: bool,
skip_tool: bool,
) -> Tuple[List[str], List[_StringHolder]]:
"""
Walk a Bedrock Converse request body and collect text content.
Returns (texts, holders) where each holder is the (container, key) pair
that owns the extracted string, so write-back mutates it in place. Besides
top-level ``text`` blocks this scans the arbitrary-JSON fields a caller can
hide prompt content in -- ``toolUse.input`` and
``toolResult.content[].json`` (alongside ``toolResult.content[].text``) --
as well as the request-level fields still forwarded to Bedrock that a caller
can route blocked content through: ``toolConfig.tools`` (tool names,
descriptions and input schemas) and ``additionalModelRequestFields``. Tool
message blocks are skipped when tool messages are excluded, but tool
definitions are always scanned to match the chat-completions guardrail path.
"""
holders: List[_StringHolder] = []
if not skip_system:
for block in body.get("system") or []:
if isinstance(block, dict):
_collect_block_text(block, holders)
for message in body.get("messages") or []:
if not isinstance(message, dict):
continue
for block in message.get("content") or []:
if not isinstance(block, dict):
continue
if skip_tool and ("toolUse" in block or "toolResult" in block):
continue
_collect_block_text(block, holders)
tool_use = block.get("toolUse")
if isinstance(tool_use, dict):
_collect_strings(tool_use.get("input"), holders)
tool_result = block.get("toolResult")
if isinstance(tool_result, dict):
for inner in tool_result.get("content") or []:
if isinstance(inner, dict):
_collect_block_text(inner, holders)
_collect_strings(inner.get("json"), holders)
tool_config = body.get("toolConfig")
if isinstance(tool_config, dict):
_collect_strings(tool_config.get("tools"), holders)
_collect_strings(body.get("additionalModelRequestFields"), holders)
texts = [container[key] for container, key in holders]
return texts, holders
def _extract_converse_output_texts(
content_blocks: List[Any],
) -> Tuple[List[str], List[_StringHolder]]:
"""
Collect user-visible text from Bedrock Converse output content blocks.
Covers ``text`` blocks plus the other content-bearing fields a model can
emit -- ``toolUse.input``, ``reasoningContent.reasoningText.text`` and
``citationsContent.content[].text`` -- while leaving structural values such
as reasoning signatures and citation sources untouched.
"""
holders: List[_StringHolder] = []
for block in content_blocks:
if not isinstance(block, dict):
continue
_collect_block_text(block, holders)
tool_use = block.get("toolUse")
if isinstance(tool_use, dict):
_collect_strings(tool_use.get("input"), holders)
reasoning = block.get("reasoningContent")
if isinstance(reasoning, dict):
reasoning_text = reasoning.get("reasoningText")
if isinstance(reasoning_text, dict):
_collect_block_text(reasoning_text, holders)
citations = block.get("citationsContent")
if isinstance(citations, dict):
for cited in citations.get("content") or []:
if isinstance(cited, dict):
_collect_block_text(cited, holders)
texts = [container[key] for container, key in holders]
return texts, holders
def _write_back_texts(
guardrailed_texts: List[str],
holders: List[_StringHolder],
) -> None:
if len(guardrailed_texts) < len(holders):
verbose_proxy_logger.warning(
"BedrockPassthroughGuardrailHandler: guardrail returned %d texts for %d "
"extracted fields; the unreturned fields keep their original text",
len(guardrailed_texts),
len(holders),
)
for idx, (container, key) in enumerate(holders):
if idx >= len(guardrailed_texts):
break
container[key] = guardrailed_texts[idx]
_DeltaHolder = Tuple[Any, Any, Union[str, int]]
def _collect_stream_delta_text_holders(delta: Any) -> List[_DeltaHolder]:
"""
Collect the user-visible text strings a Bedrock Converse ``contentBlockDelta``
can carry, matching the coverage of the non-streaming output handler.
Each holder is ``(group_key, container, key)`` where ``container[key]`` is the
text. ``group_key`` ties together fragments that belong to the same logical
stream (e.g. a single mask token split across frames) so they are
concatenated before guardrailing and redistributed afterwards. Structural
values such as reasoning signatures, redacted reasoning and citation sources
are left out so they are never rewritten.
"""
holders: List[_DeltaHolder] = []
if not isinstance(delta, dict):
return holders
if isinstance(delta.get("text"), str):
holders.append(("text", delta, "text"))
tool_use = delta.get("toolUse")
if isinstance(tool_use, dict) and isinstance(tool_use.get("input"), str):
holders.append(("tool", tool_use, "input"))
reasoning = delta.get("reasoningContent")
if isinstance(reasoning, dict) and isinstance(reasoning.get("text"), str):
holders.append(("reasoning", reasoning, "text"))
citations = delta.get("citationsContent")
if isinstance(citations, dict):
for index, cited in enumerate(citations.get("content") or []):
if isinstance(cited, dict) and isinstance(cited.get("text"), str):
holders.append((("citation", index), cited, "text"))
return holders
class BedrockPassthroughGuardrailHandler(BaseTranslation):
@staticmethod
def is_event_stream_content_type(content_type: str) -> bool:
return _EVENT_STREAM_CONTENT_TYPE in content_type
@staticmethod
def event_stream_media_type() -> str:
return _EVENT_STREAM_MEDIA_TYPE
@staticmethod
def event_stream_endpoint_is_de_anonymizable(endpoint: str) -> bool:
return _is_converse_endpoint(endpoint)
@staticmethod
async def de_anonymize_event_stream( # noqa: PLR0915
body_bytes: bytes,
proxy_logging_obj: "ProxyLogging",
user_api_key_dict: "UserAPIKeyAuth",
data: dict,
) -> bytes:
import json as _json
import struct
from binascii import crc32 as esm_crc32
from botocore.eventstream import EventStreamBuffer
frames: list[dict] = []
offset = 0
while offset + 16 <= len(body_bytes):
total_length = struct.unpack("!I", body_bytes[offset : offset + 4])[0]
if total_length < 16 or offset + total_length > len(body_bytes):
break
frame_raw = body_bytes[offset : offset + total_length]
offset += total_length
try:
buf = EventStreamBuffer()
buf.add_data(frame_raw)
msg = next(iter(buf))
event_type = msg.headers.get(":event-type")
payload_bytes = msg.payload
except Exception as e:
verbose_proxy_logger.debug(
"BedrockPassthroughGuardrailHandler: could not decode event-stream "
"frame, forwarding it unmodified: %s",
e,
)
frames.append({"raw": frame_raw, "texts": []})
continue
texts: List[Tuple[Any, str]] = []
if event_type == "contentBlockDelta":
try:
payload_dict = _json.loads(payload_bytes)
texts = [
(group_key, container[key])
for group_key, container, key in _collect_stream_delta_text_holders(
payload_dict.get("delta")
)
]
except Exception as e:
verbose_proxy_logger.debug(
"BedrockPassthroughGuardrailHandler: could not parse "
"contentBlockDelta payload, forwarding frame unmodified: %s",
e,
)
frames.append({"raw": frame_raw, "texts": texts})
trailing_bytes = body_bytes[offset:]
group_order: List[Any] = []
group_members: dict[Any, list[Tuple[int, int]]] = {}
group_texts: dict[Any, list[str]] = {}
for frame_idx, frame in enumerate(frames):
for local_idx, (group_key, text) in enumerate(frame["texts"]):
if group_key not in group_members:
group_members[group_key] = []
group_texts[group_key] = []
group_order.append(group_key)
group_members[group_key].append((frame_idx, local_idx))
group_texts[group_key].append(text)
active_groups = [gk for gk in group_order if "".join(group_texts[gk])]
if not active_groups:
return body_bytes
synthetic_response: dict = {
"output": {
"message": {
"role": "assistant",
"content": [
{"text": "".join(group_texts[gk])} for gk in active_groups
],
}
},
"stopReason": "end_turn",
}
processed = await proxy_logging_obj.post_call_success_hook(
data=data,
user_api_key_dict=user_api_key_dict,
response=synthetic_response, # type: ignore[arg-type]
)
if not isinstance(processed, dict):
verbose_proxy_logger.debug(
"BedrockPassthroughGuardrailHandler: post_call_success_hook returned %s, "
"leaving event stream unmodified",
type(processed).__name__,
)
return body_bytes
try:
processed_blocks = processed["output"]["message"]["content"] # type: ignore[index]
de_anonymized_texts = [
processed_blocks[i]["text"] for i in range(len(active_groups))
]
except (KeyError, IndexError, TypeError):
return body_bytes
new_text_map: dict[Tuple[int, int], str] = {}
for group_key, de_anonymized_text in zip(active_groups, de_anonymized_texts):
members = group_members[group_key]
orig_texts = group_texts[group_key]
total_orig = sum(len(t) for t in orig_texts) or 1
de_anon_len = len(de_anonymized_text)
pos = 0
for k, member in enumerate(members):
if k == len(members) - 1:
new_text_map[member] = de_anonymized_text[pos:]
else:
end = pos + round(de_anon_len * len(orig_texts[k]) / total_orig)
new_text_map[member] = de_anonymized_text[pos:end]
pos = end
result_parts: list[bytes] = []
for frame_idx, frame in enumerate(frames):
if not frame["texts"]:
result_parts.append(frame["raw"])
continue
frame_raw = frame["raw"]
orig_total = struct.unpack("!I", frame_raw[0:4])[0]
orig_hdrs_len = struct.unpack("!I", frame_raw[4:8])[0]
headers_bytes = frame_raw[12 : 12 + orig_hdrs_len]
try:
payload_dict = _json.loads(
frame_raw[12 + orig_hdrs_len : orig_total - 4]
)
for local_idx, (_, container, key) in enumerate(
_collect_stream_delta_text_holders(payload_dict.get("delta"))
):
new_text = new_text_map.get((frame_idx, local_idx))
if new_text is not None:
container[key] = new_text
new_payload = _json.dumps(payload_dict, separators=(",", ":")).encode()
except Exception:
result_parts.append(frame_raw)
continue
new_total = 12 + orig_hdrs_len + len(new_payload) + 4
prelude = struct.pack("!II", new_total, orig_hdrs_len)
prelude_crc_val = esm_crc32(prelude) & 0xFFFFFFFF
prelude_crc_b = struct.pack("!I", prelude_crc_val)
part_for_msg_crc = prelude_crc_b + headers_bytes + new_payload
msg_crc_val = esm_crc32(part_for_msg_crc, prelude_crc_val) & 0xFFFFFFFF
msg_crc_b = struct.pack("!I", msg_crc_val)
result_parts.append(
prelude + prelude_crc_b + headers_bytes + new_payload + msg_crc_b
)
result_parts.append(trailing_bytes)
return b"".join(result_parts)
async def process_input_messages(
self,
data: dict,
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> Any:
endpoint = data.get("endpoint", "")
body = data.get("data")
if not _is_converse_endpoint(endpoint):
return await _generic_passthrough_handler().process_input_messages(
data=data,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=litellm_logging_obj,
)
if not isinstance(body, dict) or not isinstance(body.get("messages"), list):
return data
skip_system = effective_skip_system_message_for_guardrail(guardrail_to_apply)
skip_tool = effective_skip_tool_message_for_guardrail(guardrail_to_apply)
texts, holders = _extract_converse_texts(body, skip_system, skip_tool)
if not texts:
return data
inputs = GenericGuardrailAPIInputs(texts=texts)
model = data.get("model")
if model:
inputs["model"] = model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,
)
guardrailed_texts = guardrailed_inputs.get("texts", [])
if guardrailed_texts:
_write_back_texts(guardrailed_texts, holders)
return data
async def process_output_response(
self,
response: Any,
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None,
user_api_key_dict: Optional[Any] = None,
request_data: Optional[dict] = None,
) -> Any:
endpoint = (request_data or {}).get("endpoint", "")
if endpoint and not _is_converse_endpoint(endpoint):
return await _generic_passthrough_handler().process_output_response(
response=response,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=litellm_logging_obj,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
)
if not isinstance(response, dict):
return response
output_message = (
response.get("output", {}).get("message", {})
if isinstance(response.get("output"), dict)
else {}
)
content_blocks = (
output_message.get("content") if isinstance(output_message, dict) else None
)
if not isinstance(content_blocks, list):
return response
texts, holders = _extract_converse_output_texts(content_blocks)
if not texts:
return response
effective_request_data = request_data or {}
if (
"litellm_metadata" not in effective_request_data
and user_api_key_dict is not None
):
user_metadata = self.transform_user_api_key_dict_to_metadata(
user_api_key_dict
)
if user_metadata:
effective_request_data = {
**effective_request_data,
"litellm_metadata": user_metadata,
}
inputs = GenericGuardrailAPIInputs(texts=texts)
model = effective_request_data.get("model") if effective_request_data else None
if model:
inputs["model"] = model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=effective_request_data,
input_type="response",
logging_obj=litellm_logging_obj,
)
guardrailed_texts = guardrailed_inputs.get("texts", [])
if guardrailed_texts:
_write_back_texts(guardrailed_texts, holders)
return response

View file

@ -12,8 +12,10 @@ from typing import Iterator, AsyncIterator, Any, List, Optional, Tuple, Union
import litellm
from litellm._logging import verbose_logger
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.router import GenericLiteLLMParams
from ...openai_like.chat.transformation import OpenAILikeChatConfig
@ -34,13 +36,19 @@ class BedrockMantleChatConfig(OpenAILikeChatConfig):
return super().get_config()
def _get_openai_compatible_provider_info(
self, api_base: Optional[str], api_key: Optional[str]
self,
api_base: Optional[str],
api_key: Optional[str],
litellm_params: Optional[GenericLiteLLMParams] = None,
) -> Tuple[Optional[str], Optional[str]]:
region = (
get_secret_str("BEDROCK_MANTLE_REGION")
(litellm_params.aws_region_name if litellm_params else None)
or get_secret_str("BEDROCK_MANTLE_REGION")
or get_secret_str("AWS_REGION_NAME")
or get_secret_str("AWS_REGION")
or BEDROCK_MANTLE_DEFAULT_REGION
)
BaseAWSLLM._validate_aws_region_name(region)
api_base = (
api_base
or get_secret_str("BEDROCK_MANTLE_API_BASE")

View file

@ -16,7 +16,7 @@ BaseAWSLLM._sign_request after the request body is finalized.
"""
import re
from typing import Optional, Tuple
from typing import Any, Dict, List, Optional, Tuple
from botocore.exceptions import (
CredentialRetrievalError,
@ -25,9 +25,11 @@ from botocore.exceptions import (
ProfileNotFound,
)
from litellm._logging import verbose_logger
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
@ -48,6 +50,11 @@ _MANTLE_HOST_RE = re.compile(
r"^https?://bedrock-mantle\.([^/.]+)\.api\.aws", re.IGNORECASE
)
# Per Bedrock Mantle Responses API validation errors.
_BEDROCK_MANTLE_SUPPORTED_RESPONSE_TOOL_TYPES = frozenset(
{"function", "mcp", "custom", "namespace", "tool_search"}
)
class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig):
def __init__(
@ -67,6 +74,7 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig):
def _resolve_region(params: dict) -> str:
region = params.get("aws_region_name")
if region:
BaseAWSLLM._validate_aws_region_name(region)
return region
base = params.get("api_base") or get_secret_str("BEDROCK_MANTLE_API_BASE")
if base:
@ -125,6 +133,56 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig):
def supports_native_websocket(self) -> bool:
return False
@staticmethod
def _filter_unsupported_tools(tools: List[Any]) -> List[Any]:
"""Keep only tool types Mantle's Responses API accepts."""
kept: List[Any] = []
dropped_types: List[str] = []
for tool in tools:
if not isinstance(tool, dict):
kept.append(tool)
continue
tool_type = tool.get("type")
if tool_type in _BEDROCK_MANTLE_SUPPORTED_RESPONSE_TOOL_TYPES:
kept.append(tool)
else:
dropped_types.append(str(tool_type))
if dropped_types:
verbose_logger.warning(
"Bedrock Mantle Responses API: dropping unsupported tool type(s) "
"%s (supported: %s).",
sorted(set(dropped_types)),
sorted(_BEDROCK_MANTLE_SUPPORTED_RESPONSE_TOOL_TYPES),
)
return kept
def map_openai_params(
self,
response_api_optional_params: ResponsesAPIOptionalRequestParams,
model: str,
drop_params: bool,
) -> Dict:
params = super().map_openai_params(
response_api_optional_params=response_api_optional_params,
model=model,
drop_params=drop_params,
)
tools = params.get("tools")
if not tools:
return params
tools_list = tools if isinstance(tools, list) else [tools]
filtered = self._filter_unsupported_tools(tools_list)
if filtered:
params["tools"] = filtered
else:
params.pop("tools", None)
return params
def sign_request(
self,
headers: dict,

View file

@ -125,6 +125,7 @@ from litellm.types.vector_stores import (
VectorStoreSearchOptionalRequestParams,
VectorStoreSearchResponse,
)
from litellm.types.realtime import RealtimeQueryParams
from litellm.types.videos.main import VideoObject
from litellm.utils import (
CustomStreamWrapper,
@ -2305,6 +2306,7 @@ class BaseLLMHTTPHandler:
if extra_body:
data.update(extra_body)
stream = bool(stream or data.get("stream"))
# Preserve the OpenAI-style request context (not sent to the provider) for streaming
# hooks/metadata; the streaming iterator now consumes this to run deployment hooks
@ -2467,6 +2469,7 @@ class BaseLLMHTTPHandler:
if extra_body:
data.update(extra_body)
stream = bool(stream or data.get("stream"))
# Preserve the OpenAI-style request context (not sent to the provider) for streaming
# hooks/metadata; the streaming iterator now consumes this to run deployment hooks
@ -5315,6 +5318,23 @@ class BaseLLMHTTPHandler:
headers=error_headers,
)
@staticmethod
def _append_query_params(
url: str, query_params: Optional[RealtimeQueryParams]
) -> str:
"""Append query_params to url, skipping keys already present in the URL."""
if not query_params:
return url
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
parsed = urlparse(url)
existing = dict(parse_qsl(parsed.query))
extras = {k: v for k, v in query_params.items() if k not in existing}
if not extras:
return url
new_query = parsed.query + ("&" if parsed.query else "") + urlencode(extras)
return urlunparse(parsed._replace(query=new_query))
async def async_realtime(
self,
model: str,
@ -5328,11 +5348,14 @@ class BaseLLMHTTPHandler:
timeout: Optional[float] = None,
user_api_key_dict: Optional[Any] = None,
litellm_metadata: Optional[Dict[str, Any]] = None,
query_params: Optional[RealtimeQueryParams] = None,
):
import websockets
from websockets.asyncio.client import ClientConnection
url = provider_config.get_complete_url(api_base, model, api_key)
url = self._append_query_params(
provider_config.get_complete_url(api_base, model, api_key), query_params
)
headers = provider_config.validate_environment(
headers=headers,
model=model,
@ -5373,6 +5396,11 @@ class BaseLLMHTTPHandler:
model,
user_api_key_dict=user_api_key_dict,
request_data=_request_data,
force_transcription_model=(
model
if (query_params or {}).get("intent") == "transcription"
else None
),
)
if _session_config:
realtime_streaming.session_configuration_request = _session_config
@ -5437,6 +5465,69 @@ class BaseLLMHTTPHandler:
"""
Forward POST /v1/realtime/client_secrets to upstream provider.
Uses provider_config (BaseRealtimeHTTPConfig) for URL construction and
header auth when available; falls back to the legacy OpenAI-style defaults.
"""
return await self._async_realtime_session_post(
endpoint="client_secrets",
api_base=api_base,
api_key=api_key,
request_data=request_data,
logging_obj=logging_obj,
timeout=timeout,
provider_config=provider_config,
model=model,
extra_headers=extra_headers,
client=client,
api_version=api_version,
)
async def async_realtime_transcription_session_handler(
self,
api_base: str,
api_key: str,
request_data: Dict[str, Any],
logging_obj: LiteLLMLoggingObj,
timeout: Union[float, httpx.Timeout],
provider_config: Optional[Any] = None,
model: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
api_version: Optional[str] = None,
) -> httpx.Response:
"""Forward POST /v1/realtime/transcription_sessions to upstream provider."""
return await self._async_realtime_session_post(
endpoint="transcription_sessions",
api_base=api_base,
api_key=api_key,
request_data=request_data,
logging_obj=logging_obj,
timeout=timeout,
provider_config=provider_config,
model=model,
extra_headers=extra_headers,
client=client,
api_version=api_version,
)
async def _async_realtime_session_post(
self,
endpoint: Literal["client_secrets", "transcription_sessions"],
api_base: str,
api_key: str,
request_data: Dict[str, Any],
logging_obj: LiteLLMLoggingObj,
timeout: Union[float, httpx.Timeout],
provider_config: Optional[Any] = None,
model: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
api_version: Optional[str] = None,
) -> httpx.Response:
"""
Shared POST flow for the realtime HTTP session endpoints
(client_secrets and transcription_sessions).
Uses provider_config (BaseRealtimeHTTPConfig) for URL construction and
header auth when available; falls back to the legacy OpenAI-style defaults.
"""
@ -5448,14 +5539,19 @@ class BaseLLMHTTPHandler:
async_httpx_client = client
if provider_config is not None:
url = provider_config.get_complete_url(
api_base=api_base, model=model or "", api_version=api_version
)
if endpoint == "transcription_sessions":
url = provider_config.get_transcription_session_url(
api_base=api_base, model=model or "", api_version=api_version
)
else:
url = provider_config.get_complete_url(
api_base=api_base, model=model or "", api_version=api_version
)
headers: Dict[str, Any] = provider_config.validate_environment(
headers={}, model=model or "", api_key=api_key
)
else:
url = f"{api_base.rstrip('/')}/v1/realtime/client_secrets"
url = f"{api_base.rstrip('/')}/v1/realtime/{endpoint}"
headers = {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
@ -5575,7 +5671,7 @@ class BaseLLMHTTPHandler:
)
raise
async def async_responses_websocket(
async def async_responses_websocket( # noqa: PLR0915
self,
model: str,
websocket: Any,
@ -5628,7 +5724,11 @@ class BaseLLMHTTPHandler:
import websockets
from websockets.asyncio.client import ClientConnection
litellm_params = GenericLiteLLMParams()
litellm_params = GenericLiteLLMParams(
api_base=api_base,
api_key=api_key,
**kwargs,
)
headers = responses_api_provider_config.validate_environment(
headers={},
model=model,
@ -5637,21 +5737,21 @@ class BaseLLMHTTPHandler:
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
http_url = responses_api_provider_config.get_complete_url(
ws_url = responses_api_provider_config.get_websocket_url(
api_base=api_base,
litellm_params={},
litellm_params=dict(litellm_params),
)
ws_url = http_url.replace("https://", "wss://").replace("http://", "ws://")
# OpenAI's WebSocket responses endpoint requires ?model= in the URL,
# matching the Realtime API convention (wss://.../v1/realtime?model=...).
# Use urllib.parse so existing query params (e.g. api-version) are preserved.
_parsed = urlparse(ws_url)
_qs = parse_qs(_parsed.query)
if "model" not in _qs:
_qs["model"] = [model]
ws_url = urlunparse(
_parsed._replace(query=urlencode({k: v[0] for k, v in _qs.items()}))
)
# Some providers (e.g. OpenAI) require ?model= in the WebSocket URL.
# Providers that send the model in the request body (e.g. Azure) set
# model_in_websocket_url() to False to suppress this append.
if responses_api_provider_config.model_in_websocket_url():
_parsed = urlparse(ws_url)
_qs = parse_qs(_parsed.query)
if "model" not in _qs:
_qs["model"] = [model]
ws_url = urlunparse(
_parsed._replace(query=urlencode({k: v[0] for k, v in _qs.items()}))
)
try:
ssl_context = get_shared_realtime_ssl_context()
@ -5679,6 +5779,41 @@ class BaseLLMHTTPHandler:
_request_data: Dict[str, Any] = {}
if litellm_metadata:
_request_data["litellm_metadata"] = litellm_metadata
_ws_guardrail_callbacks: list = []
_ws_output_guardrail_callbacks: list = []
try:
import litellm as _litellm
# Use duck-typing so any guardrail that exposes the PII
# masking interface works, not just _OPTIONAL_PresidioPIIMasking.
# This avoids a layering violation (SDK importing from proxy).
_ws_guardrail_callbacks = [
cb
for cb in _litellm.callbacks
if callable(getattr(cb, "check_pii", None))
and callable(
getattr(cb, "get_presidio_settings_from_request_data", None)
)
and callable(getattr(cb, "_unmask_pii_text", None))
and getattr(cb, "output_parse_pii", False)
]
_ws_output_guardrail_callbacks = [
cb
for cb in _litellm.callbacks
if callable(getattr(cb, "check_pii", None))
and callable(
getattr(cb, "get_presidio_settings_from_request_data", None)
)
and getattr(cb, "apply_to_output", False)
]
except Exception as _guardrail_exc:
verbose_logger.warning(
"Responses WebSocket: failed to collect guardrail "
"callbacks — PII masking will be skipped. Error: %s",
_guardrail_exc,
)
streaming = ResponsesWebSocketStreaming(
websocket=websocket,
backend_ws=cast(ClientConnection, backend_ws),
@ -5686,6 +5821,9 @@ class BaseLLMHTTPHandler:
user_api_key_dict=user_api_key_dict,
request_data=_request_data,
first_message=first_message,
guardrail_callbacks=_ws_guardrail_callbacks,
output_guardrail_callbacks=_ws_output_guardrail_callbacks,
authorized_model=model,
)
await streaming.bidirectional_forward()

View file

@ -174,12 +174,110 @@ def map_openai_image_params_to_gemini(
mapped_params["imageConfig"] = image_config_param
for key, value in filtered_params.items():
if key not in ("n", "size", "imageConfig") and key not in optional_params:
if (
key not in ("n", "size", "imageConfig", "tools", "web_search_options")
and key not in optional_params
):
mapped_params[key] = value
return mapped_params
def _dedupe_gemini_search_tools(tools: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
search_tool_keys = VertexGeminiConfig._search_tool_keys()
seen_search_keys: set[str] = set()
deduped_tools: List[Dict[str, Any]] = []
for tool in tools:
if not isinstance(tool, dict):
deduped_tools.append(tool)
continue
search_key = next((key for key in search_tool_keys if key in tool), None)
if search_key is None:
deduped_tools.append(tool)
continue
if search_key in seen_search_keys:
continue
seen_search_keys.add(search_key)
deduped_tools.append(tool)
return deduped_tools
def _has_gemini_search_tool(tools: List[Any]) -> bool:
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
search_tool_keys = VertexGeminiConfig._search_tool_keys()
return any(
isinstance(tool, dict) and any(key in tool for key in search_tool_keys)
for tool in tools
)
def map_gemini_image_tools_params(
non_default_params: Dict[str, Any],
mapped_params: Dict[str, Any],
) -> Dict[str, Any]:
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
gemini_config = VertexGeminiConfig()
result = dict(mapped_params)
result.pop("web_search_options", None)
tools_value = non_default_params.get("tools")
if isinstance(tools_value, list) and tools_value:
mapped_tools = gemini_config._map_function(
value=tools_value, optional_params=result
)
result = gemini_config._add_tools_to_optional_params(result, mapped_tools)
web_search_options = non_default_params.get("web_search_options")
existing_tools = result.get("tools")
if isinstance(web_search_options, dict) and not (
isinstance(existing_tools, list) and _has_gemini_search_tool(existing_tools)
):
search_tool = gemini_config._map_web_search_options(web_search_options)
result = gemini_config._add_tools_to_optional_params(result, [search_tool])
gemini_config._drop_search_tools_mixed_with_functions(result)
if isinstance(result.get("tools"), list):
result["tools"] = _dedupe_gemini_search_tools(result["tools"])
return result
def get_gemini_image_web_search_requests(
response_data: Dict[str, Any],
) -> Optional[int]:
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
grounding_metadata: List[Dict[str, Any]] = []
for candidate in response_data.get("candidates", []):
if not isinstance(candidate, dict):
continue
candidate_grounding = candidate.get("groundingMetadata")
if isinstance(candidate_grounding, list):
grounding_metadata.extend(candidate_grounding)
elif isinstance(candidate_grounding, dict):
grounding_metadata.append(candidate_grounding)
return VertexGeminiConfig._calculate_web_search_requests(grounding_metadata)
def get_gemini_image_generation_config(
model: str,
optional_params: Dict[str, Any],

View file

@ -7,6 +7,7 @@ from typing import Any
import litellm
from litellm.litellm_core_utils.llm_cost_calc.utils import (
calculate_image_response_cost_from_usage,
calculate_image_response_web_search_cost,
)
from litellm.types.utils import ImageResponse
@ -23,22 +24,25 @@ def cost_calculator(
custom_llm_provider="gemini",
)
if isinstance(image_response, ImageResponse):
token_based_cost = calculate_image_response_cost_from_usage(
model=model,
image_response=image_response,
custom_llm_provider="gemini",
)
if token_based_cost is not None:
return token_based_cost
output_cost_per_image: float = _model_info.get("output_cost_per_image") or 0.0
num_images: int = 0
if isinstance(image_response, ImageResponse):
if image_response.data:
num_images = len(image_response.data)
return output_cost_per_image * num_images
else:
if not isinstance(image_response, ImageResponse):
raise ValueError(
f"image_response must be of type ImageResponse got type={type(image_response)}"
)
web_search_cost = calculate_image_response_web_search_cost(
image_response=image_response,
custom_llm_provider="gemini",
model_info=_model_info,
)
token_based_cost = calculate_image_response_cost_from_usage(
model=model,
image_response=image_response,
custom_llm_provider="gemini",
)
if token_based_cost is not None:
return token_based_cost + web_search_cost
output_cost_per_image: float = _model_info.get("output_cost_per_image") or 0.0
num_images: int = len(image_response.data) if image_response.data else 0
return output_cost_per_image * num_images + web_search_cost

View file

@ -7,7 +7,9 @@ from litellm.llms.base_llm.image_generation.transformation import (
)
from litellm.llms.gemini.common_utils import (
get_gemini_image_generation_config,
get_gemini_image_web_search_requests,
is_gemini_image_model,
map_gemini_image_tools_params,
map_openai_image_params_to_gemini,
)
from litellm.llms.gemini.image_usage_transformation import (
@ -41,7 +43,7 @@ class GoogleImageGenConfig(BaseImageGenerationConfig):
"""
supported_params = ["n", "size"]
if is_gemini_image_model(model):
supported_params.append("imageConfig")
supported_params.extend(["imageConfig", "tools", "web_search_options"])
return supported_params # type: ignore[return-value]
def map_openai_params(
@ -51,12 +53,17 @@ class GoogleImageGenConfig(BaseImageGenerationConfig):
model: str,
drop_params: bool,
) -> dict:
return map_openai_image_params_to_gemini(
mapped_params = map_openai_image_params_to_gemini(
params=non_default_params,
model=model,
supported_params=self.get_supported_openai_params(model),
optional_params=optional_params,
)
if is_gemini_image_model(model):
mapped_params = map_gemini_image_tools_params(
non_default_params, mapped_params
)
return mapped_params
def get_complete_url(
self,
@ -140,6 +147,10 @@ class GoogleImageGenConfig(BaseImageGenerationConfig):
optional_params=optional_params,
),
}
if tools := optional_params.get("tools"):
request_body["tools"] = tools
if tool_config := optional_params.get("toolConfig"):
request_body["toolConfig"] = tool_config
return request_body
else:
# For other Imagen models, use the original Imagen format
@ -217,6 +228,11 @@ class GoogleImageGenConfig(BaseImageGenerationConfig):
model_response.usage = transform_gemini_image_usage(
response_data["usageMetadata"]
)
web_search_requests = get_gemini_image_web_search_requests(response_data)
if web_search_requests and model_response.usage is not None:
setattr(
model_response.usage, "web_search_requests", web_search_requests
)
else:
# Original Imagen format - predictions with generated images
predictions = response_data.get("predictions", [])

View file

@ -195,13 +195,48 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
def get_audio_mime_type(self, input_audio_format: str = "pcm16"):
mime_types = {
"pcm16": "audio/pcm",
"pcm16": "audio/pcm;rate=24000",
"g711_ulaw": "audio/pcmu",
"g711_alaw": "audio/pcma",
}
return mime_types.get(input_audio_format, "application/octet-stream")
def _manual_turn_detection_enabled(
self, session_configuration_request: Optional[str]
) -> bool:
if not session_configuration_request:
return False
try:
setup = json.loads(session_configuration_request).get("setup", {})
automatic_detection = setup.get("realtimeInputConfig", {}).get(
"automaticActivityDetection", {}
)
return (
isinstance(automatic_detection, dict)
and automatic_detection.get("disabled") is True
)
except (json.JSONDecodeError, TypeError, AttributeError):
return False
def _handle_input_audio_buffer_commit_or_end(
self, session_configuration_request: Optional[str]
) -> List[str]:
"""Map OpenAI buffer commit/end to Gemini Live turn-boundary signals."""
if self._manual_turn_detection_enabled(session_configuration_request):
realtime_input_dict: BidiGenerateContentRealtimeInput = {
"activityEnd": True,
}
verbose_logger.debug(
"Gemini Realtime: Sending activityEnd realtimeInput to backend"
)
else:
realtime_input_dict = {"audioStreamEnd": True}
verbose_logger.debug(
"Gemini Realtime: Sending audioStreamEnd realtimeInput to backend"
)
return [json.dumps({"realtimeInput": realtime_input_dict})]
def map_automatic_turn_detection(
self, value: OpenAIRealtimeTurnDetection
) -> AutomaticActivityDetection:
@ -656,6 +691,19 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
)
messages.append(gemini_msg)
return messages
if msg_type in ("input_audio_buffer.commit", "input_audio_buffer.end"):
return self._handle_input_audio_buffer_commit_or_end(
session_configuration_request
)
if msg_type == "input_audio_buffer.clear":
# Local OpenAI buffer op — nothing to forward to Gemini Live.
verbose_logger.debug(
"Gemini Realtime: input_audio_buffer.clear is a local buffer op"
)
return []
# Unknown/unsupported OpenAI event type — drop silently rather than
# forwarding raw JSON as text input to the model.
return []

View file

@ -101,6 +101,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig):
def __init__(self) -> None:
super().__init__()
self.authenticator = Authenticator()
self._stream_item_ids_by_output_index: Dict[int, str] = {}
@property
def custom_llm_provider(self) -> LlmProviders:
@ -129,6 +130,61 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig):
"""
return dict(response_api_optional_params)
def transform_streaming_response(
self,
model: str,
parsed_chunk: dict,
logging_obj: LiteLLMLoggingObj,
) -> Any:
parsed_chunk = self._normalize_stream_item_id(parsed_chunk)
return super().transform_streaming_response(
model=model,
parsed_chunk=parsed_chunk,
logging_obj=logging_obj,
)
def _normalize_stream_item_id(self, parsed_chunk: dict) -> dict:
"""Rewrite streamed item ids to one stable id per output_index.
GitHub Copilot tags each event of a single output item with a different
item id, so clients that key streaming state by item id (e.g. the Vercel
AI SDK) crash with "reasoning part <id> not found" / "text part <id> not
found". Every sub-event carries a top-level ``item_id`` (whatever the
item type), so its presence is the rewrite signal; output_item.added /
.done instead nest the id under ``item``. The anchor is keyed by
output_index and taken from output_item.added, which the protocol always
emits first, so it is written before any sub-event reads it. Copilot
accepts that id paired with the final encrypted_content next turn, so
multi-turn replay is unaffected.
State is keyed by output_index on this config, which
ProviderConfigManager builds fresh per request, so it is stream-scoped.
"""
output_index = parsed_chunk.get("output_index")
if not isinstance(output_index, int):
return parsed_chunk
if parsed_chunk.get("type") == "response.output_item.added":
item = parsed_chunk.get("item")
if isinstance(item, dict) and isinstance(item.get("id"), str):
self._stream_item_ids_by_output_index[output_index] = item["id"]
return parsed_chunk
stable_id = self._stream_item_ids_by_output_index.get(output_index)
if stable_id is None:
return parsed_chunk
if isinstance(parsed_chunk.get("item_id"), str):
parsed_chunk = dict(parsed_chunk)
parsed_chunk["item_id"] = stable_id
elif parsed_chunk.get("type") == "response.output_item.done":
item = parsed_chunk.get("item")
if isinstance(item, dict):
parsed_chunk = dict(parsed_chunk)
parsed_chunk["item"] = {**item, "id": stable_id}
return parsed_chunk
def validate_environment(
self,
headers: dict,

View file

@ -25,6 +25,7 @@ from typing import (
import httpx
import litellm
from litellm.constants import DEFAULT_OCI_CHAT_MAX_TOKENS
from litellm.litellm_core_utils.logging_utils import track_llm_api_timing
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.llms.custom_httpx.http_handler import (
@ -87,15 +88,20 @@ STREAMING_TIMEOUT = 60 * 5
def _model_uses_max_completion_tokens(model: str) -> bool:
"""Return True for OCI-hosted models that require ``maxCompletionTokens``.
Reasoning models on OCI (e.g. the OpenAI GPT-5 family) reject ``maxTokens``
with HTTP 400 and require ``maxCompletionTokens`` per OpenAI's reasoning-API
convention. Driven by ``supports_reasoning`` in
``model_prices_and_context_window.json`` so new model families are picked
up via a catalog update rather than a code change.
OpenAI commercial models proxied through OCI (``openai.*``) reject
``maxTokens`` with HTTP 400 on the reasoning families (gpt-5.x, o-series)
and accept ``maxCompletionTokens`` everywhere, so route the whole vendor
prefix to it rather than chasing each new release in
``model_prices_and_context_window.json``. The ``openai.gpt-oss-*`` open
weights are served by OCI's own stack and keep ``maxTokens``. Any other
vendor falls back to the catalog's ``supports_reasoning`` flag.
"""
if not model:
return False
name = model[4:] if model.lower().startswith("oci/") else model
lowered = name.lower()
if lowered.startswith("openai."):
return not lowered.startswith("openai.gpt-oss")
return supports_reasoning(model=name, custom_llm_provider="oci")
@ -193,19 +199,49 @@ def _normalize_response_format(selected_params: Dict, vendor: OCIVendors) -> Non
rf = selected_params.get("responseFormat")
if not isinstance(rf, dict) or "type" not in rf:
return
rf_payload = dict(rf)
selected_params["responseFormat"] = rf_payload
response_type = rf_payload["type"]
if "json_schema" in rf_payload:
raw_schema = rf_payload.pop("json_schema")
rf_payload["jsonSchema"] = (
dict(raw_schema) if isinstance(raw_schema, dict) else raw_schema
)
rf_type = str(rf["type"]).lower()
raw_schema = rf.get("json_schema")
json_schema = raw_schema if isinstance(raw_schema, dict) else None
if rf_type == "text":
selected_params["responseFormat"] = {"type": "TEXT"}
return
if vendor == OCIVendors.COHERE:
rf_payload["type"] = response_type
else:
fmt = response_type.upper()
rf_payload["type"] = "JSON_OBJECT" if fmt == "JSON" else fmt
# OCI Cohere has no JSON_SCHEMA type; a schema rides on JSON_OBJECT.
payload: Dict[str, Any] = {"type": "JSON_OBJECT"}
if json_schema is not None and json_schema.get("schema") is not None:
payload["schema"] = json_schema["schema"]
selected_params["responseFormat"] = payload
return
if rf_type == "json_schema":
if json_schema is None:
raise OCIError(
status_code=400,
message="response_format type 'json_schema' requires a 'json_schema' object",
)
# OCI's ResponseJsonSchema accepts only name/description/schema/isStrict.
# OpenAI sends `strict` instead of `isStrict`; forwarding it (or any
# other extra key) makes OCI reject the whole request with HTTP 400.
oci_schema: Dict[str, Any] = {"name": json_schema.get("name") or "response"}
if json_schema.get("description") is not None:
oci_schema["description"] = json_schema["description"]
if json_schema.get("schema") is not None:
oci_schema["schema"] = json_schema["schema"]
if json_schema.get("strict") is not None:
oci_schema["isStrict"] = json_schema["strict"]
selected_params["responseFormat"] = {
"type": "JSON_SCHEMA",
"jsonSchema": oci_schema,
}
return
fmt = rf_type.upper()
selected_params["responseFormat"] = {
"type": "JSON_OBJECT" if fmt == "JSON" else fmt
}
def get_vendor_from_model(model: str) -> OCIVendors:
@ -297,6 +333,11 @@ class OCIChatConfig(BaseConfig):
if get_vendor_from_model(model) == OCIVendors.COHERE
else self.openai_to_oci_generic_param_map
)
# `n` is intentionally not advertised for Cohere even though n=1 is
# tolerated: Cohere has no numGenerations field, so n>1 cannot be
# honoured and advertising it would be misleading. Callers that gate on
# this list strip n=1 (a no-op, matching what map_openai_params does);
# callers that bypass it have n=1 dropped there. Both paths converge.
return [key for key, value in param_map.items() if value]
def map_openai_params(
@ -317,6 +358,19 @@ class OCIChatConfig(BaseConfig):
for key, value in {**non_default_params, **optional_params}.items():
alias = param_map.get(key)
if alias is False:
# max_retries is a litellm-level control param (litellm applies
# retries itself); it is never a generation param OCI accepts, so
# drop it silently. The litellm proxy injects it on every request,
# which otherwise 500s OCI calls unless drop_params is set.
if key == "max_retries":
continue
# n=1 (or None) is the OpenAI default: a single generation, which
# every OCI model produces anyway. Drop it silently so standard
# clients that always send n=1 (e.g. the MLflow gateway) are not
# rejected; only n>1 is genuinely unsupported on Cohere, which
# has no numGenerations field.
if key == "n" and (value is None or value == 1):
continue
if drop_params or litellm.drop_params:
continue
raise OCIError(
@ -451,6 +505,13 @@ class OCIChatConfig(BaseConfig):
elif oci_alias in optional_params:
selected_params[target] = optional_params[oci_alias] # type: ignore[index]
# OCI's server-side default token cap is tiny (~20 tokens), so an
# omitted max_tokens silently truncates the response mid-string. Most
# callers never send a limit (MLflow judges among them), so inject a
# sane default when one is absent, mirroring litellm's Anthropic config.
if max_tokens_key not in selected_params:
selected_params[max_tokens_key] = DEFAULT_OCI_CHAT_MAX_TOKENS
# OCI expects uppercase reasoning levels (LOW/MEDIUM/HIGH/NONE); OpenAI
# clients send lowercase. OpenAI's "disable" maps to OCI's "NONE".
if "reasoningEffort" in selected_params:

View file

@ -157,8 +157,14 @@ class OpenAIRealtime(OpenAIChatCompletion):
websocket,
cast(ClientConnection, backend_ws),
logging_obj,
model=model,
user_api_key_dict=user_api_key_dict,
request_data={"litellm_metadata": litellm_metadata or {}},
force_transcription_model=(
model
if (query_params or {}).get("intent") == "transcription"
else None
),
)
await realtime_streaming.bidirectional_forward()

View file

@ -41,6 +41,14 @@ class OpenAIRealtimeHTTPConfig(BaseRealtimeHTTPConfig):
base = base[:-3]
return f"{base}/v1/realtime/calls"
def get_transcription_session_url(
self, api_base: Optional[str], model: str, api_version: Optional[str] = None
) -> str:
base = self.get_api_base(api_base).rstrip("/")
if base.endswith("/v1"):
base = base[:-3]
return f"{base}/v1/realtime/transcription_sessions"
def validate_environment(
self,
headers: dict,

View file

@ -35,7 +35,6 @@ from pydantic import BaseModel
from litellm._logging import verbose_proxy_logger
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
OpenAiResponsesToChatCompletionStreamIterator,
)
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
@ -479,90 +478,137 @@ class OpenAIResponsesHandler(BaseTranslation):
) -> List[Any]:
"""
Process output streaming response by applying guardrails to text content.
Mirrors the Chat Completions handler pattern: extract text from the final
chunk, apply the guardrail, then write the result back in-place so the
caller sees the modified content (e.g. PII tokens replaced).
For ``response.completed`` events (the normal end-of-stream signal) we
use the same per-item extraction + task-mapping approach as
``process_output_response`` so that unmasking / blocking works correctly
for every output item.
"""
if not responses_so_far:
return responses_so_far
final_chunk = responses_so_far[-1]
# Accept both plain dicts and Pydantic models (BaseLiteLLMOpenAIResponseObject
# exposes a .get() shim, so all the .get() calls below work for both).
if not (isinstance(final_chunk, dict) or hasattr(final_chunk, "get")):
return responses_so_far
# ------------------------------------------------------------------ #
# Case 1: response.completed — full response is available in the #
# final chunk; iterate output items, apply guardrail, write back. #
# ------------------------------------------------------------------ #
if final_chunk.get("type") == "response.completed":
response_obj = final_chunk.get("response") or {}
if not hasattr(response_obj, "get"):
return responses_so_far
outputs: List[Any] = response_obj.get("output") or []
texts_to_check: List[str] = []
tool_calls_to_check: List[ChatCompletionToolCallChunk] = []
task_mappings: List[Tuple[int, int]] = []
for output_idx, output_item in enumerate(outputs):
self._extract_output_text_and_images(
output_item=output_item,
output_idx=output_idx,
texts_to_check=texts_to_check,
images_to_check=[],
task_mappings=task_mappings,
tool_calls_to_check=tool_calls_to_check,
)
if texts_to_check or tool_calls_to_check:
if request_data is None:
request_data = {}
if "response" not in request_data:
request_data["response"] = response_obj
if "litellm_metadata" not in request_data:
user_metadata = self.transform_user_api_key_dict_to_metadata(
user_api_key_dict
)
if user_metadata:
request_data["litellm_metadata"] = user_metadata
inputs = GenericGuardrailAPIInputs(texts=texts_to_check)
if tool_calls_to_check:
inputs["tool_calls"] = cast(
List[ChatCompletionToolCallChunk], tool_calls_to_check
)
response_model = response_obj.get("model")
if response_model:
inputs["model"] = response_model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="response",
logging_obj=litellm_logging_obj,
)
guardrailed_texts = guardrailed_inputs.get("texts", [])
# Write guardrailed texts back into the output items in-place.
# final_chunk is a reference into responses_so_far so this
# mutates the list that the caller holds.
await self._apply_guardrail_responses_to_output(
response=response_obj,
responses=guardrailed_texts,
task_mappings=task_mappings,
)
return responses_so_far
# ------------------------------------------------------------------ #
# Case 2: response.output_item.done — extract tool calls only. #
# ------------------------------------------------------------------ #
if final_chunk.get("type") == "response.output_item.done":
# convert openai response to model response
model_response_stream = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(
final_chunk
)
tool_calls = model_response_stream.choices[0].delta.tool_calls
if tool_calls:
inputs = GenericGuardrailAPIInputs()
inputs["tool_calls"] = cast(
List[ChatCompletionToolCallChunk], tool_calls
)
# Include model information if available
if (
hasattr(model_response_stream, "model")
and model_response_stream.model
):
inputs["model"] = model_response_stream.model
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=request_data if request_data is not None else {},
input_type="response",
logging_obj=litellm_logging_obj,
)
return responses_so_far
elif final_chunk.get("type") == "response.completed":
# convert openai response to model response
outputs = final_chunk.get("response", {}).get("output", [])
return responses_so_far
model_response_choices = LiteLLMResponsesTransformationHandler._convert_response_output_to_choices(
output_items=outputs,
handle_raw_dict_callback=None,
)
if model_response_choices:
tool_calls = model_response_choices[0].message.tool_calls
text = model_response_choices[0].message.content
guardrail_inputs = GenericGuardrailAPIInputs()
if text:
guardrail_inputs["texts"] = [text]
if tool_calls:
guardrail_inputs["tool_calls"] = cast(
List[ChatCompletionToolCallChunk], tool_calls
)
# Include model information from the response if available
response_model = final_chunk.get("response", {}).get("model")
if response_model:
guardrail_inputs["model"] = response_model
if tool_calls or text:
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=guardrail_inputs,
request_data=request_data if request_data is not None else {},
input_type="response",
logging_obj=litellm_logging_obj,
)
return responses_so_far
else:
verbose_proxy_logger.debug(
"Skipping output guardrail - model response has no choices"
)
# model_response_stream = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(final_chunk)
# tool_calls = model_response_stream.choices[0].tool_calls
# convert openai response to model response
# ------------------------------------------------------------------ #
# Fallback: apply guardrail to the accumulated text string. #
# No structured write-back is possible here; guardrails that only #
# need to block/flag (not rewrite) still work correctly. #
# ------------------------------------------------------------------ #
string_so_far = self.get_streaming_string_so_far(responses_so_far)
inputs = GenericGuardrailAPIInputs(texts=[string_so_far])
# Try to get model from the final chunk if available
if isinstance(final_chunk, dict):
if string_so_far:
fallback_inputs = GenericGuardrailAPIInputs(texts=[string_so_far])
response_model = (
final_chunk.get("response", {}).get("model")
if isinstance(final_chunk.get("response"), dict)
else None
)
if response_model:
inputs["model"] = response_model
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=request_data if request_data is not None else {},
input_type="response",
logging_obj=litellm_logging_obj,
)
fallback_inputs["model"] = response_model
await guardrail_to_apply.apply_guardrail(
inputs=fallback_inputs,
request_data=request_data if request_data is not None else {},
input_type="response",
logging_obj=litellm_logging_obj,
)
return responses_so_far
def _check_streaming_has_ended(self, responses_so_far: List[Any]) -> bool:
@ -721,7 +767,7 @@ class OpenAIResponsesHandler(BaseTranslation):
async def _apply_guardrail_responses_to_output(
self,
response: "ResponsesAPIResponse",
response: Union["ResponsesAPIResponse", Dict[Any, Any]],
responses: List[str],
task_mappings: List[Tuple[int, int]],
) -> None:

View file

@ -131,7 +131,8 @@
"base_class": "openai_gpt",
"param_mappings": {
"max_completion_tokens": "max_tokens"
}
},
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"]
},
"parasail": {
"base_url": "https://api.parasail.io/v1",
@ -141,5 +142,14 @@
"special_handling": {
"force_store_false": true
}
},
"empiriolabs": {
"base_url": "https://api.empiriolabs.ai/v1",
"api_key_env": "EMPIRIOLABS_API_KEY",
"api_base_env": "EMPIRIOLABS_API_BASE",
"param_mappings": {
"max_completion_tokens": "max_tokens"
},
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"]
}
}

View file

@ -1,7 +1,7 @@
"""
Calls Parallel AI's /search endpoint to search the web.
Calls Parallel AI's /v1/search endpoint to search the web.
Parallel AI API Reference: https://docs.parallel.ai/api-reference/search-and-extract-api-beta/search
Parallel AI API Reference: https://docs.parallel.ai/api-reference/search/search
"""
from typing import Dict, List, Optional, TypedDict, Union
@ -18,36 +18,43 @@ from litellm.secret_managers.main import get_secret_str
class _ParallelAISourcePolicy(TypedDict, total=False):
"""Source policy for Parallel AI search results."""
allowed_domains: List[str] # Optional - list of allowed domains
disallowed_domains: List[str] # Optional - list of disallowed domains
include_domains: List[str]
exclude_domains: List[str]
after_date: str
class _ParallelAISearchRequestRequired(TypedDict):
"""Required fields for Parallel AI Search API request."""
# Note: At least one of objective or search_queries must be provided
pass
class _ParallelAIExcerptSettings(TypedDict, total=False):
max_chars_per_result: int
class ParallelAISearchRequest(_ParallelAISearchRequestRequired, total=False):
class _ParallelAIAdvancedSettings(TypedDict, total=False):
source_policy: _ParallelAISourcePolicy
excerpt_settings: _ParallelAIExcerptSettings
fetch_policy: Dict
location: str
max_results: int
class ParallelAISearchRequest(TypedDict, total=False):
"""
Parallel AI Search API request format.
Based on: https://docs.parallel.ai/api-reference/search-and-extract-api-beta/search
Parallel AI v1 Search API request format.
Based on: https://docs.parallel.ai/api-reference/search/search
"""
search_queries: List[str] # Required - at least one keyword search query
objective: str # Optional - natural-language description of search goal
search_queries: List[str] # Optional - list of keyword search queries
processor: str # Optional - search processor ('base', 'pro'), default 'base'
max_results: int # Optional - maximum number of results, default 10
max_chars_per_result: int # Optional - max characters per result excerpt
source_policy: _ParallelAISourcePolicy # Optional - source policy for allowed/disallowed domains
mode: str # Optional - 'turbo', 'basic', or 'advanced' (default 'advanced')
max_chars_total: int # Optional - upper bound on total excerpt characters
session_id: str # Optional - tracks calls across search/extract requests
client_model: str # Optional - model consuming the results
advanced_settings: _ParallelAIAdvancedSettings
LEGACY_PROCESSOR_TO_MODE = {"base": "basic", "pro": "advanced"}
class ParallelAISearchConfig(BaseSearchConfig):
PARALLEL_AI_API_BASE = "https://api.parallel.ai"
PARALLEL_HEADER_SEARCH_EXTRACT_VALUE = "search-extract-2025-10-10"
@staticmethod
def ui_friendly_name() -> str:
@ -60,9 +67,6 @@ class ParallelAISearchConfig(BaseSearchConfig):
api_base: Optional[str] = None,
**kwargs,
) -> Dict:
"""
Validate environment and return headers.
"""
api_key = (
api_key
or get_secret_str("PARALLEL_AI_API_KEY")
@ -74,7 +78,6 @@ class ParallelAISearchConfig(BaseSearchConfig):
)
headers["x-api-key"] = api_key
headers["Content-Type"] = "application/json"
headers["parallel-beta"] = self.PARALLEL_HEADER_SEARCH_EXTRACT_VALUE
return headers
def get_complete_url(
@ -84,32 +87,18 @@ class ParallelAISearchConfig(BaseSearchConfig):
data: Optional[Union[Dict, List[Dict]]] = None,
**kwargs,
) -> str:
"""
Get complete URL for Search endpoint.
"""
api_base = (
api_base
or get_secret_str("PARALLEL_AI_API_BASE")
or self.PARALLEL_AI_API_BASE
)
# Parallel AI search endpoint is at /v1beta/search
if not api_base.endswith("/v1beta/search"):
if api_base.endswith("/"):
api_base = f"{api_base}v1beta/search"
else:
api_base = f"{api_base}/v1beta/search"
api_base = api_base.rstrip("/")
if not api_base.endswith("/v1/search"):
api_base = f"{api_base.removesuffix('/v1')}/v1/search"
return api_base
def _transform_query_to_objective(self, query: Union[str, List[str]]) -> str:
"""
Transform query to objective.
"""
if isinstance(query, list):
return " ".join(query)
return query
def transform_search_request(
self,
query: Union[str, List[str]],
@ -117,57 +106,78 @@ class ParallelAISearchConfig(BaseSearchConfig):
**kwargs,
) -> Dict:
"""
Transform Search request to Parallel AI API format.
Transform Search request to Parallel AI v1 API format.
Args:
query: Search query (string or list of strings)
- If string: maps to `objective` (natural language)
- If string: maps to `search_queries` (single item) and `objective`
- If list: maps to `search_queries` (keyword queries)
optional_params: Optional parameters for the request
- max_results: Maximum number of search results (default 10)
- search_domain_filter: List of domains to include -> maps to `source_policy.allowed_domains`
- exclude_domains: List of domains to exclude -> maps to `source_policy.disallowed_domains`
- processor: Search processor ('base', 'pro')
- max_chars_per_result: Max characters per result excerpt
- mode: Search mode ('turbo', 'basic', 'advanced'); defaults to 'basic'
- processor: Legacy v1beta param; 'base' maps to mode 'basic', 'pro' to 'advanced'
- max_results: Maximum number of search results -> `advanced_settings.max_results`
- search_domain_filter: Domains to include -> `advanced_settings.source_policy.include_domains`
- exclude_domains: Domains to exclude -> `advanced_settings.source_policy.exclude_domains`
- country: ISO 3166-1 alpha-2 code -> `advanced_settings.location`
- max_chars_per_result: -> `advanced_settings.excerpt_settings.max_chars_per_result`
- Any other params are passed through to the request body as-is
Returns:
Dict with typed request data following ParallelAISearchRequest spec
Dict with request data following the v1 search request spec
"""
params = dict(optional_params)
request_data: ParallelAISearchRequest = {}
# Map query to objective (string or list both become objective)
if isinstance(query, list):
request_data["objective"] = self._transform_query_to_objective(query)
request_data["search_queries"] = query
else:
request_data["search_queries"] = [query]
request_data["objective"] = query
# Transform Perplexity unified spec parameters to Parallel AI format
if "max_results" in optional_params:
request_data["max_results"] = optional_params["max_results"]
mode = params.pop("mode", None)
processor = params.pop("processor", None)
if mode is None and processor is not None:
mode = LEGACY_PROCESSOR_TO_MODE.get(processor, processor)
# the v1 API defaults to 'advanced' when mode is omitted; default to 'basic'
# instead to keep v1beta's default tier (processor 'base') and litellm's
# $0.004/query cost map entry for `parallel_ai/search` accurate
request_data["mode"] = mode or "basic"
advanced_settings: _ParallelAIAdvancedSettings = {}
if "max_results" in params:
advanced_settings["max_results"] = params.pop("max_results")
if "country" in params:
advanced_settings["location"] = params.pop("country")
if "max_chars_per_result" in params:
advanced_settings["excerpt_settings"] = {
"max_chars_per_result": params.pop("max_chars_per_result")
}
# Map domain filters to source_policy
source_policy: _ParallelAISourcePolicy = {}
if "search_domain_filter" in optional_params:
source_policy["allowed_domains"] = optional_params["search_domain_filter"]
if "search_domain_filter" in params:
source_policy["include_domains"] = params.pop("search_domain_filter")
if "exclude_domains" in optional_params:
source_policy["disallowed_domains"] = optional_params["exclude_domains"]
if "exclude_domains" in params:
source_policy["exclude_domains"] = params.pop("exclude_domains")
if source_policy:
request_data["source_policy"] = source_policy
advanced_settings["source_policy"] = source_policy
# Convert to dict before dynamic key assignments
result_data = dict(request_data)
advanced_settings.update(params.pop("advanced_settings", {}))
# pass through all other parameters as-is
for param, value in optional_params.items():
if (
param not in self.get_supported_perplexity_optional_params()
and param not in result_data
):
result_data[param] = value
if advanced_settings:
request_data["advanced_settings"] = advanced_settings
# unified-spec param with no v1 equivalent
params.pop("max_tokens_per_page", None)
result_data: Dict = dict(request_data)
result_data.update(params)
return result_data
def transform_search_response(
@ -177,36 +187,27 @@ class ParallelAISearchConfig(BaseSearchConfig):
**kwargs,
) -> SearchResponse:
"""
Transform Parallel AI API response to LiteLLM unified SearchResponse format.
Transform Parallel AI v1 API response to LiteLLM unified SearchResponse format.
Parallel AI → LiteLLM mappings:
- results[].title → SearchResult.title
- results[].url → SearchResult.url
- results[].excerpts (array) → SearchResult.snippet (joined string)
- No date/last_updated fields in Parallel AI response (set to None)
Args:
raw_response: Raw httpx response from Parallel AI API
logging_obj: Logging object for tracking
Returns:
SearchResponse with standardized format
Parallel AI -> LiteLLM mappings:
- results[].title -> SearchResult.title
- results[].url -> SearchResult.url
- results[].excerpts (array) -> SearchResult.snippet (joined string)
- results[].publish_date -> SearchResult.date
"""
response_json = raw_response.json()
# Transform results to SearchResult objects
results = []
for result in response_json.get("results", []):
# Join excerpts array into a single snippet string
excerpts = result.get("excerpts", [])
excerpts = result.get("excerpts") or []
snippet = " ... ".join(excerpts) if excerpts else ""
search_result = SearchResult(
title=result.get("title", ""),
url=result.get("url", ""),
title=result.get("title") or "",
url=result.get("url") or "",
snippet=snippet,
date=None, # Parallel AI doesn't provide date in response
last_updated=None, # Parallel AI doesn't provide last_updated in response
date=result.get("publish_date"),
last_updated=None,
)
results.append(search_result)

View file

@ -1,15 +1,18 @@
"""Pass-Through Endpoint guardrail translation handler."""
from litellm.llms.pass_through.guardrail_translation.handler import (
LlmPassthroughRouteHandler,
PassThroughEndpointHandler,
)
from litellm.types.utils import CallTypes
guardrail_translation_mappings = {
CallTypes.pass_through: PassThroughEndpointHandler,
CallTypes.allm_passthrough_route: LlmPassthroughRouteHandler,
}
__all__ = [
"guardrail_translation_mappings",
"LlmPassthroughRouteHandler",
"PassThroughEndpointHandler",
]

View file

@ -6,7 +6,7 @@ It uses the field targeting configuration from litellm_logging_obj
to extract specific fields for guardrail processing.
"""
from typing import TYPE_CHECKING, Any, List, Optional
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type
from litellm._logging import verbose_proxy_logger
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
@ -16,6 +16,8 @@ from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
class PassThroughEndpointHandler(BaseTranslation):
@ -208,3 +210,128 @@ class PassThroughEndpointHandler(BaseTranslation):
)
return response
_PROVIDER_HANDLERS: Dict[str, Type[BaseTranslation]] = {}
def _get_provider_handlers() -> Dict[str, Type[BaseTranslation]]:
global _PROVIDER_HANDLERS
if not _PROVIDER_HANDLERS:
from litellm.llms.bedrock.passthrough.guardrail_translation.handler import (
BedrockPassthroughGuardrailHandler,
)
_PROVIDER_HANDLERS = {"bedrock": BedrockPassthroughGuardrailHandler}
return _PROVIDER_HANDLERS
class LlmPassthroughRouteHandler(BaseTranslation):
"""
Dispatcher for allm_passthrough_route guardrail translation.
Routes to a per-provider handler based on data["custom_llm_provider"].
Unknown providers are skipped with a debug log.
"""
async def process_input_messages(
self,
data: dict,
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> Any:
provider = data.get("custom_llm_provider")
handler_cls = _get_provider_handlers().get(provider or "")
if handler_cls is None:
verbose_proxy_logger.debug(
"LlmPassthroughRouteHandler: no handler for provider=%s, skipping guardrail",
provider,
)
return data
return await handler_cls().process_input_messages(
data=data,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=litellm_logging_obj,
)
async def process_output_response(
self,
response: Any,
guardrail_to_apply: "CustomGuardrail",
litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None,
user_api_key_dict: Optional[Any] = None,
request_data: Optional[dict] = None,
) -> Any:
provider = (request_data or {}).get("custom_llm_provider")
handler_cls = _get_provider_handlers().get(provider or "")
if handler_cls is None:
verbose_proxy_logger.debug(
"LlmPassthroughRouteHandler: no handler for provider=%s, skipping guardrail",
provider,
)
return response
return await handler_cls().process_output_response(
response=response,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=litellm_logging_obj,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
)
@staticmethod
def is_event_stream_response(provider: Optional[str], content_type: str) -> bool:
handler_cls = _get_provider_handlers().get(provider or "")
detector = getattr(handler_cls, "is_event_stream_content_type", None)
if detector is None:
return False
return detector(content_type)
@staticmethod
def event_stream_media_type(provider: Optional[str]) -> Optional[str]:
handler_cls = _get_provider_handlers().get(provider or "")
getter = getattr(handler_cls, "event_stream_media_type", None)
if getter is None:
return None
return getter()
@staticmethod
def _resolve_event_stream_de_anonymizer(provider: Optional[str]):
handler_cls = _get_provider_handlers().get(provider or "")
return getattr(handler_cls, "de_anonymize_event_stream", None)
@staticmethod
def supports_event_stream_de_anonymization(
provider: Optional[str], endpoint: Optional[str]
) -> bool:
handler_cls = _get_provider_handlers().get(provider or "")
endpoint_check = getattr(
handler_cls, "event_stream_endpoint_is_de_anonymizable", None
)
if endpoint_check is None:
return False
return endpoint_check(endpoint or "")
@staticmethod
async def de_anonymize_event_stream(
body_bytes: bytes,
proxy_logging_obj: "ProxyLogging",
user_api_key_dict: "UserAPIKeyAuth",
data: dict,
) -> bytes:
provider = data.get("custom_llm_provider")
de_anonymize = LlmPassthroughRouteHandler._resolve_event_stream_de_anonymizer(
provider
)
if de_anonymize is None:
verbose_proxy_logger.debug(
"LlmPassthroughRouteHandler: no event-stream handler for provider=%s, "
"leaving stream unmodified",
provider,
)
return body_bytes
return await de_anonymize(
body_bytes=body_bytes,
proxy_logging_obj=proxy_logging_obj,
user_api_key_dict=user_api_key_dict,
data=data,
)

View file

@ -1,17 +1,32 @@
"""
Support for Snowflake REST API
Snowflake Cortex REST API — Chat Transformation
Routes to native Cortex REST API endpoints based on model:
- Claude models → POST /api/v2/cortex/v1/messages (Anthropic format)
- All other models → POST /api/v2/cortex/v1/chat/completions (OpenAI format)
Ref: https://docs.snowflake.com/en/user-guide/snowflake-cortex/cortex-rest-api
"""
import json
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
from typing import TYPE_CHECKING, Any, Dict, List, Optional
import httpx
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ChatCompletionMessageToolCall, Function, ModelResponse
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk
from litellm.types.utils import (
ChatCompletionMessageToolCall,
ChatCompletionUsageBlock,
Choices,
Function,
GenericStreamingChunk,
Message,
ModelResponse,
Usage,
)
from ...base_llm.base_model_iterator import BaseModelResponseIterator
from ...openai_like.chat.transformation import OpenAIGPTConfig
from ..utils import SnowflakeBaseConfig
if TYPE_CHECKING:
@ -21,69 +36,343 @@ if TYPE_CHECKING:
else:
LiteLLMLoggingObj = Any
ANTHROPIC_VERSION = "2023-06-01"
_CLAUDE_MODEL_PREFIXES = (
"claude-",
"claude_",
)
def _is_claude_model(model: str) -> bool:
"""Return True if model name (after stripping snowflake/ prefix) is a Claude model."""
name = model.lower().removeprefix("snowflake/")
return any(name.startswith(p) for p in _CLAUDE_MODEL_PREFIXES)
class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
"""
Reference: https://docs.snowflake.com/en/user-guide/snowflake-cortex/cortex-llm-rest-api
Snowflake Cortex REST API — unified provider.
Snowflake Cortex LLM REST API supports function calling with specific models (e.g., Claude 3.5 Sonnet).
This config handles transformation between OpenAI format and Snowflake's tool_spec format.
Auto-routes based on model name:
- Claude models → /api/v2/cortex/v1/messages (Anthropic Messages format)
- All others → /api/v2/cortex/v1/chat/completions (OpenAI format)
Auth:
PAT: api_key="pat/<token>" → X-Snowflake-Authorization-Token-Type: PROGRAMMATIC_ACCESS_TOKEN
JWT: api_key="<jwt>" → X-Snowflake-Authorization-Token-Type: KEYPAIR_JWT
"""
@classmethod
def get_config(cls):
return super().get_config()
def _transform_tool_calls_from_snowflake_to_openai(
self, content_list: List[Dict[str, Any]]
) -> Tuple[str, Optional[List[ChatCompletionMessageToolCall]]]:
def get_supported_openai_params(self, model: str) -> List[str]:
params = [
"temperature",
"max_tokens",
"max_completion_tokens",
"top_p",
"stream",
"tools",
"tool_choice",
]
if _is_claude_model(model):
params.append("thinking")
return params
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
api_base = self._get_api_base(api_base, optional_params)
if _is_claude_model(model):
return f"{api_base}/cortex/v1/messages"
return f"{api_base}/cortex/v1/chat/completions"
def validate_environment(
self,
headers: dict,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> dict:
headers = super().validate_environment(
headers=headers,
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
api_key=api_key,
api_base=api_base,
)
if _is_claude_model(model):
headers["anthropic-version"] = ANTHROPIC_VERSION
return headers
def _transform_tools_to_anthropic(self, tools: List[Dict]) -> List[Dict]:
"""
Transform Snowflake tool calls to OpenAI format.
Convert tools from OpenAI format to Anthropic format.
Args:
content_list: Snowflake's content_list array containing text and tool_use items
OpenAI: {"type": "function", "function": {"name": ..., "parameters": {...}}}
Anthropic: {"name": ..., "description": ..., "input_schema": {...}}
"""
anthropic_tools = []
for tool in tools:
if tool.get("type") == "function" and "function" in tool:
func = tool["function"]
anthropic_tool: Dict[str, Any] = {
"name": func.get("name", ""),
}
if "description" in func:
anthropic_tool["description"] = func["description"]
if "parameters" in func:
anthropic_tool["input_schema"] = func["parameters"]
else:
anthropic_tool["input_schema"] = {
"type": "object",
"properties": {},
}
anthropic_tools.append(anthropic_tool)
else:
anthropic_tools.append(tool)
return anthropic_tools
Returns:
Tuple of (text_content, tool_calls)
def _extract_system_and_messages(
self, messages: List[AllMessageValues]
) -> tuple[Optional[str], List[Dict]]:
"""
Split messages into system prompt and conversation turns for Anthropic format.
Snowflake format in content_list:
{
"type": "tool_use",
"tool_use": {
"tool_use_id": "tooluse_...",
"name": "get_weather",
"input": {"location": "Paris"}
}
- system messages → collected and joined (preserves guardrail prompts)
- assistant messages with tool_calls → tool_use content blocks
- tool role messages → user role with tool_result content blocks
"""
system_parts: List[str] = []
conversation: List[Dict] = []
for msg in messages:
if isinstance(msg, dict):
role = msg.get("role", "")
content: Any = msg.get("content", "")
else:
role = getattr(msg, "role", "")
content = getattr(msg, "content", "")
if role == "system":
if isinstance(content, str) and content:
system_parts.append(content)
elif isinstance(content, list):
system_parts.append(
"\n".join(
b.get("text", "")
for b in content
if b.get("type") == "text"
)
)
elif role == "assistant":
tool_calls = (
msg.get("tool_calls")
if isinstance(msg, dict)
else getattr(msg, "tool_calls", None)
)
if tool_calls: # type: ignore[truthy-bool]
content_blocks: List[Dict[str, Any]] = []
if content:
content_blocks.append({"type": "text", "text": content})
for tc in tool_calls: # type: ignore[attr-defined]
func = (
tc.get("function", {})
if isinstance(tc, dict)
else getattr(tc, "function", {})
)
tc_id = (
tc.get("id", "")
if isinstance(tc, dict)
else getattr(tc, "id", "")
)
func_name = (
func.get("name", "")
if isinstance(func, dict)
else getattr(func, "name", "")
)
func_args = (
func.get("arguments", "{}")
if isinstance(func, dict)
else getattr(func, "arguments", "{}")
)
try:
input_data = (
json.loads(func_args)
if isinstance(func_args, str)
else func_args
)
except (json.JSONDecodeError, TypeError):
input_data = {}
content_blocks.append(
{
"type": "tool_use",
"id": tc_id,
"name": func_name,
"input": input_data,
}
)
conversation.append(
{"role": "assistant", "content": content_blocks}
)
else:
conversation.append({"role": "assistant", "content": content})
elif role == "tool":
tool_call_id = (
msg.get("tool_call_id", "")
if isinstance(msg, dict)
else getattr(msg, "tool_call_id", "")
)
tool_content = (
content if isinstance(content, str) else json.dumps(content)
)
tool_result_block = {
"type": "tool_result",
"tool_use_id": tool_call_id,
"content": tool_content,
}
if (
conversation
and conversation[-1]["role"] == "user"
and isinstance(conversation[-1]["content"], list)
and conversation[-1]["content"]
and conversation[-1]["content"][0].get("type") == "tool_result"
):
conversation[-1]["content"].append(tool_result_block)
else:
conversation.append(
{"role": "user", "content": [tool_result_block]}
)
else:
conversation.append({"role": role, "content": content})
system: Optional[str] = "\n\n".join(system_parts) if system_parts else None
return system, conversation
def transform_request(
self,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
stream: bool = optional_params.pop("stream", False) or False
extra_body = optional_params.pop("extra_body", {})
if _is_claude_model(model):
return self._transform_request_anthropic(
model, messages, optional_params, stream, extra_body
)
return self._transform_request_openai(
model, messages, optional_params, stream, extra_body
)
def _transform_request_openai(
self,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
stream: bool,
extra_body: dict,
) -> dict:
"""OpenAI format for /chat/completions endpoint."""
max_tokens = optional_params.pop("max_tokens", None)
max_completion_tokens = optional_params.pop("max_completion_tokens", None)
resolved_max = max_completion_tokens or max_tokens
body: dict = {
"model": model.removeprefix("snowflake/"),
"messages": messages,
"stream": stream,
**optional_params,
**extra_body,
}
OpenAI format (returned tool_calls):
ChatCompletionMessageToolCall(
id="tooluse_...",
type="function",
function=Function(name="get_weather", arguments='{"location": "Paris"}')
)
if resolved_max is not None:
body["max_completion_tokens"] = resolved_max
return body
def _transform_tool_choice_to_anthropic(self, tool_choice: Any) -> Dict[str, Any]:
"""
text_content = ""
tool_calls: List[ChatCompletionMessageToolCall] = []
Convert tool_choice from OpenAI format to Anthropic format.
for idx, content_item in enumerate(content_list):
if content_item.get("type") == "text":
text_content += content_item.get("text", "")
OpenAI string values: "auto", "required", "none"
OpenAI dict: {"type": "function", "function": {"name": "..."}}
Anthropic: {"type": "auto"}, {"type": "any"}, {"type": "tool", "name": "..."}
"""
if isinstance(tool_choice, str):
mapping = {
"auto": {"type": "auto"},
"required": {"type": "any"},
"none": {"type": "none"},
}
return mapping.get(tool_choice, {"type": "auto"})
elif isinstance(tool_choice, dict):
if tool_choice.get("type") == "function":
func = tool_choice.get("function", {})
return {"type": "tool", "name": func.get("name", "")}
return tool_choice
return {"type": "auto"}
## TOOL CALLING
elif content_item.get("type") == "tool_use":
tool_use_data = content_item.get("tool_use", {})
tool_call = ChatCompletionMessageToolCall(
id=tool_use_data.get("tool_use_id", ""),
type="function",
function=Function(
name=tool_use_data.get("name", ""),
arguments=json.dumps(tool_use_data.get("input", {})),
),
)
tool_calls.append(tool_call)
def _transform_request_anthropic(
self,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
stream: bool,
extra_body: dict,
) -> dict:
"""Anthropic Messages format for /messages endpoint."""
system, conversation = self._extract_system_and_messages(messages)
return text_content, tool_calls if tool_calls else None
if "tools" in optional_params:
optional_params["tools"] = self._transform_tools_to_anthropic(
optional_params["tools"]
)
if "tool_choice" in optional_params:
optional_params["tool_choice"] = self._transform_tool_choice_to_anthropic(
optional_params["tool_choice"]
)
max_completion_tokens = optional_params.pop("max_completion_tokens", None)
if max_completion_tokens and "max_tokens" not in optional_params:
optional_params["max_tokens"] = max_completion_tokens
model_name = model.removeprefix("snowflake/")
body: Dict[str, Any] = {
"model": model_name,
"messages": conversation,
"stream": stream,
**optional_params,
**extra_body,
}
if system is not None:
body["system"] = system
if "max_tokens" not in body:
body["max_tokens"] = (
4096 # reasonable default; Anthropic API max varies by model
)
return body
def transform_response(
self,
@ -99,6 +388,24 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
) -> ModelResponse:
if _is_claude_model(model):
return self._transform_response_anthropic(
model, raw_response, model_response, logging_obj, request_data, messages
)
return self._transform_response_openai(
model, raw_response, model_response, logging_obj, request_data, messages
)
def _transform_response_openai(
self,
model: str,
raw_response: httpx.Response,
model_response: ModelResponse,
logging_obj: LiteLLMLoggingObj,
request_data: dict,
messages: List[AllMessageValues],
) -> ModelResponse:
"""Parse standard OpenAI chat completions response."""
response_json = raw_response.json()
logging_obj.post_call(
@ -108,180 +415,278 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig):
additional_args={"complete_input_dict": request_data},
)
## RESPONSE TRANSFORMATION
# Snowflake returns content_list (not content) with tool_use objects
# We need to transform this to OpenAI's format with content + tool_calls
if "choices" in response_json and len(response_json["choices"]) > 0:
choice = response_json["choices"][0]
if "message" in choice and "content_list" in choice["message"]:
content_list = choice["message"]["content_list"]
(
text_content,
tool_calls,
) = self._transform_tool_calls_from_snowflake_to_openai(content_list)
# Update the choice message with OpenAI format
choice["message"]["content"] = text_content
if tool_calls:
choice["message"]["tool_calls"] = tool_calls
# Remove Snowflake-specific content_list
del choice["message"]["content_list"]
returned_response = ModelResponse(**response_json)
returned_response.model = "snowflake/" + (returned_response.model or "")
if model is not None:
returned_response._hidden_params["model"] = model
return returned_response
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
"""
If api_base is not provided, use the default DeepSeek /chat/completions endpoint.
"""
api_base = self._get_api_base(api_base, optional_params)
return f"{api_base}/cortex/inference:complete"
def _transform_tools(self, tools: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
"""
Transform OpenAI tool format to Snowflake tool format.
Args:
tools: List of tools in OpenAI format
Returns:
List of tools in Snowflake format
OpenAI format:
{
"type": "function",
"function": {
"name": "get_weather",
"description": "...",
"parameters": {...}
}
}
Snowflake format:
{
"tool_spec": {
"type": "generic",
"name": "get_weather",
"description": "...",
"input_schema": {...}
}
}
"""
snowflake_tools: List[Dict[str, Any]] = []
for tool in tools:
if tool.get("type") == "function":
function = tool.get("function", {})
snowflake_tool: Dict[str, Any] = {
"tool_spec": {
"type": "generic",
"name": function.get("name"),
"input_schema": function.get(
"parameters",
{"type": "object", "properties": {}},
),
}
}
# Add description if present
if "description" in function:
snowflake_tool["tool_spec"]["description"] = function["description"]
snowflake_tools.append(snowflake_tool)
return snowflake_tools
def _transform_tool_choice(
self, tool_choice: Union[str, Dict[str, Any]]
) -> Dict[str, Any]:
"""
Transform OpenAI tool_choice format to Snowflake format.
Snowflake requires tool_choice to be an object, not a string.
Ref: https://docs.snowflake.com/en/developer-guide/snowflake-rest-api/reference/cortex-inference#post--api-v2-cortex-inference-complete-req-body-schema
Args:
tool_choice: Tool choice in OpenAI format (str or dict)
Returns:
Tool choice in Snowflake format (always an object, never a string)
OpenAI format (string):
"auto", "required", "none"
OpenAI format (dict):
{"type": "function", "function": {"name": "get_weather"}}
Snowflake format:
{"type": "auto"} / {"type": "any"} / {"type": "none"}
{"type": "tool", "name": ["get_weather"]}
Snowflake's API (like Anthropic) requires tool_choice as an object
with a "type" field, not as a bare string.
"""
if isinstance(tool_choice, str):
# Snowflake requires object format, not string.
# Map OpenAI string values to Snowflake object format.
# "required" maps to "any" (Snowflake/Anthropic convention).
_type_map = {
"auto": "auto",
"required": "any",
"none": "none",
}
mapped_type = _type_map.get(tool_choice, tool_choice)
return {"type": mapped_type}
if isinstance(tool_choice, dict):
if tool_choice.get("type") == "function":
function_name = tool_choice.get("function", {}).get("name")
if function_name:
return {
"type": "tool",
"name": [function_name], # Snowflake expects array
}
return tool_choice
def transform_request(
def _transform_response_anthropic(
self,
model: str,
raw_response: httpx.Response,
model_response: ModelResponse,
logging_obj: LiteLLMLoggingObj,
request_data: dict,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
stream: bool = optional_params.pop("stream", None) or False
extra_body = optional_params.pop("extra_body", {})
) -> ModelResponse:
"""Parse Anthropic Messages response into OpenAI format."""
response_json = raw_response.json()
## TOOL CALLING
# Transform tools from OpenAI format to Snowflake's tool_spec format
tools = optional_params.pop("tools", None)
if tools:
optional_params["tools"] = self._transform_tools(tools)
logging_obj.post_call(
input=messages,
api_key="",
original_response=response_json,
additional_args={"complete_input_dict": request_data},
)
# Transform tool_choice from OpenAI format to Snowflake's tool name array format
tool_choice = optional_params.pop("tool_choice", None)
if tool_choice:
optional_params["tool_choice"] = self._transform_tool_choice(tool_choice)
text_content = ""
tool_calls = []
return {
"model": model,
"messages": messages,
"stream": stream,
**optional_params,
**extra_body,
for block in response_json.get("content", []):
if block.get("type") == "text":
text_content += block.get("text", "")
elif block.get("type") == "tool_use":
tool_calls.append(
ChatCompletionMessageToolCall(
id=block.get("id", ""),
type="function",
function=Function(
name=block.get("name", ""),
arguments=json.dumps(block.get("input", {})),
),
)
)
_stop_reason_map = {
"end_turn": "stop",
"max_tokens": "length",
"tool_use": "tool_calls",
"stop_sequence": "stop",
}
finish_reason = _stop_reason_map.get(
response_json.get("stop_reason", "end_turn"), "stop"
)
message = Message(content=text_content or None, role="assistant")
if tool_calls:
message.tool_calls = tool_calls
choice = Choices(
finish_reason=finish_reason,
index=0,
message=message,
)
usage_data = response_json.get("usage", {})
usage = Usage(
prompt_tokens=usage_data.get("input_tokens", 0),
completion_tokens=usage_data.get("output_tokens", 0),
total_tokens=usage_data.get("input_tokens", 0)
+ usage_data.get("output_tokens", 0),
)
model_response.choices = [choice]
model_response.usage = usage # type: ignore[attr-defined]
model_response.model = "snowflake/" + response_json.get("model", model)
model_response.id = response_json.get("id", "")
if model is not None:
model_response._hidden_params["model"] = model
return model_response
def get_model_response_iterator(
self,
streaming_response: Any,
sync_stream: bool,
json_mode: Optional[bool] = False,
) -> Any:
return SnowflakeStreamingHandler(
streaming_response=streaming_response,
sync_stream=sync_stream,
json_mode=json_mode,
)
class SnowflakeStreamingHandler(BaseModelResponseIterator):
"""
Parse streaming events from both Snowflake endpoints.
- /chat/completions: OpenAI SSE format (has "choices" key)
- /messages: Anthropic SSE format (has "type" key like content_block_delta)
"""
def __init__(
self,
streaming_response: Any,
sync_stream: bool,
json_mode: Optional[bool] = False,
):
super().__init__(streaming_response=streaming_response, sync_stream=sync_stream)
self._tool_index = 0
self._tool_id = ""
self._tool_name = ""
self._input_tokens = 0
def chunk_parser(self, chunk: dict) -> GenericStreamingChunk:
if "choices" in chunk:
return self._parse_openai_chunk(chunk)
return self._parse_anthropic_chunk(chunk)
def _parse_openai_chunk(self, chunk: dict) -> GenericStreamingChunk:
choices = chunk.get("choices", [])
if not choices:
return GenericStreamingChunk(
text="",
is_finished=False,
finish_reason="",
usage=None,
index=0,
tool_use=None,
)
choice = choices[0]
delta = choice.get("delta", {})
finish_reason = choice.get("finish_reason") or ""
text = delta.get("content") or ""
tool_use = None
tool_calls = delta.get("tool_calls")
if tool_calls:
tc = tool_calls[0]
func = tc.get("function", {})
tool_use = ChatCompletionToolCallChunk(
id=tc.get("id", ""),
type="function",
function={
"name": func.get("name", ""),
"arguments": func.get("arguments", ""),
},
index=tc.get("index", 0),
)
return GenericStreamingChunk(
text=text,
is_finished=finish_reason != "",
finish_reason=finish_reason,
usage=None,
index=choice.get("index", 0),
tool_use=tool_use,
)
def _parse_anthropic_chunk(self, chunk: dict) -> GenericStreamingChunk:
event_type = chunk.get("type", "")
if event_type == "message_start":
message = chunk.get("message", {})
usage_data = message.get("usage", {})
self._input_tokens = usage_data.get("input_tokens", 0)
return GenericStreamingChunk(
text="",
is_finished=False,
finish_reason="",
usage=None,
index=0,
tool_use=None,
)
elif event_type == "content_block_delta":
delta = chunk.get("delta", {})
delta_type = delta.get("type", "")
if delta_type == "text_delta":
return GenericStreamingChunk(
text=delta.get("text", ""),
is_finished=False,
finish_reason="",
usage=None,
index=chunk.get("index", 0),
tool_use=None,
)
elif delta_type == "input_json_delta":
return GenericStreamingChunk(
text="",
is_finished=False,
finish_reason="",
usage=None,
index=chunk.get("index", 0),
tool_use=ChatCompletionToolCallChunk(
id=self._tool_id,
type="function",
function={
"name": self._tool_name,
"arguments": delta.get("partial_json", ""),
},
index=self._tool_index,
),
)
elif event_type == "content_block_start":
content_block = chunk.get("content_block", {})
if content_block.get("type") == "tool_use":
self._tool_id = content_block.get("id", "")
self._tool_name = content_block.get("name", "")
self._tool_index = chunk.get("index", 0)
return GenericStreamingChunk(
text="",
is_finished=False,
finish_reason="",
usage=None,
index=chunk.get("index", 0),
tool_use=ChatCompletionToolCallChunk(
id=self._tool_id,
type="function",
function={"name": self._tool_name, "arguments": ""},
index=self._tool_index,
),
)
elif event_type == "message_delta":
delta = chunk.get("delta", {})
stop_reason = delta.get("stop_reason", "")
usage_data = chunk.get("usage", {})
_stop_map = {
"end_turn": "stop",
"max_tokens": "length",
"tool_use": "tool_calls",
"stop_sequence": "stop",
}
usage = None
if usage_data or self._input_tokens:
output_t = usage_data.get("output_tokens", 0)
input_t = self._input_tokens or usage_data.get("input_tokens", 0)
usage = ChatCompletionUsageBlock(
prompt_tokens=input_t,
completion_tokens=output_t,
total_tokens=input_t + output_t,
)
return GenericStreamingChunk(
text="",
is_finished=True,
finish_reason=_stop_map.get(stop_reason, "stop"),
usage=usage,
index=0,
tool_use=None,
)
elif event_type == "message_stop":
return GenericStreamingChunk(
text="",
is_finished=True,
finish_reason="stop",
usage=None,
index=0,
tool_use=None,
)
return GenericStreamingChunk(
text="",
is_finished=False,
finish_reason="",
usage=None,
index=0,
tool_use=None,
)

View file

@ -3422,14 +3422,18 @@ class ModelResponseIterator:
self.has_seen_tool_calls = True
break
# Handle final chunk with finishReason but no content.
# _process_candidates skips candidates without "content",
# so the finish_reason from the final chunk is lost.
# _process_candidates skips candidates without a "content" part, so a
# content-less chunk leaves choices empty and the downstream streaming
# handler hits IndexError on choices[0]. This covers the final chunk
# (finishReason, no content) and mid-stream metadata-only chunks
# (grounding/web-search/thought, no content and no finishReason — seen
# with web_search + reasoning) by emitting an empty-delta choice.
if not model_response.choices and _candidates:
from litellm.types.utils import Delta, StreamingChoices
for candidate in _candidates:
finish_reason_str = candidate.get("finishReason")
mapped_finish_reason = None
if finish_reason_str is not None:
if self.has_seen_tool_calls:
mapped_finish_reason = "tool_calls"
@ -3437,14 +3441,14 @@ class ModelResponseIterator:
mapped_finish_reason = VertexGeminiConfig._check_finish_reason(
None, finish_reason_str
)
choice = StreamingChoices(
finish_reason=mapped_finish_reason,
index=candidate.get("index", 0),
delta=Delta(content=None, role=None),
logprobs=None,
enhancements=None,
)
model_response.choices.append(choice)
choice = StreamingChoices(
finish_reason=mapped_finish_reason,
index=candidate.get("index", 0),
delta=Delta(content=None, role=None),
logprobs=None,
enhancements=None,
)
model_response.choices.append(choice)
# Also handle the case where the final chunk has empty
# content (e.g. text:"") WITH finishReason. In this case

View file

@ -5,6 +5,7 @@ Vertex AI Image Generation Cost Calculator
import litellm
from litellm.litellm_core_utils.llm_cost_calc.utils import (
calculate_image_response_cost_from_usage,
calculate_image_response_web_search_cost,
)
from litellm.types.utils import ImageResponse
@ -21,16 +22,20 @@ def cost_calculator(
custom_llm_provider="vertex_ai",
)
web_search_cost = calculate_image_response_web_search_cost(
image_response=image_response,
custom_llm_provider="vertex_ai",
model_info=_model_info,
)
token_based_cost = calculate_image_response_cost_from_usage(
model=model,
image_response=image_response,
custom_llm_provider="vertex_ai",
)
if token_based_cost is not None:
return token_based_cost
return token_based_cost + web_search_cost
output_cost_per_image: float = _model_info.get("output_cost_per_image") or 0.0
num_images: int = 0
if image_response.data:
num_images = len(image_response.data)
return output_cost_per_image * num_images
num_images: int = len(image_response.data) if image_response.data else 0
return output_cost_per_image * num_images + web_search_cost

View file

@ -7,6 +7,10 @@ import litellm
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from litellm.llms.gemini.common_utils import (
get_gemini_image_web_search_requests,
map_gemini_image_tools_params,
)
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM
from litellm.secret_managers.main import get_secret_str
@ -52,6 +56,8 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM):
"aspect_ratio",
"imageSize",
"image_size",
"tools",
"web_search_options",
]
def map_openai_params(
@ -77,9 +83,10 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM):
mapped_params["aspectRatio"] = v
elif k in ("imageSize", "image_size"):
mapped_params["imageSize"] = v
else:
elif k not in ("tools", "web_search_options"):
mapped_params[k] = v
mapped_params = map_gemini_image_tools_params(non_default_params, mapped_params)
return mapped_params
def _map_size_to_aspect_ratio(self, size: str) -> str:
@ -247,6 +254,11 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM):
"generationConfig": generation_config,
}
if tools := optional_params.get("tools"):
request_body["tools"] = tools
if tool_config := optional_params.get("toolConfig"):
request_body["toolConfig"] = tool_config
return request_body
def _transform_image_usage(self, usage: dict) -> ImageUsage:
@ -324,4 +336,8 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM):
if usage_metadata := response_data.get("usageMetadata", None):
model_response.usage = self._transform_image_usage(usage_metadata)
web_search_requests = get_gemini_image_web_search_requests(response_data)
if web_search_requests and model_response.usage is not None:
setattr(model_response.usage, "web_search_requests", web_search_requests)
return model_response

View file

@ -90,7 +90,8 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig):
def get_audio_mime_type(self, input_audio_format: str = "pcm16") -> str:
mime_types = {
"pcm16": "audio/pcm;rate=16000",
# Gemini Live native audio (OpenAI GA realtime default) is 24kHz PCM.
"pcm16": "audio/pcm;rate=24000",
"g711_ulaw": "audio/pcmu",
"g711_alaw": "audio/pcma",
}

View file

@ -0,0 +1,183 @@
"""
Transform request/response for Voyage multimodal embeddings.
Voyage multimodal models use /v1/multimodalembeddings and accept `inputs`
containing content blocks, unlike standard Voyage embeddings which use
/v1/embeddings and a string/list `input` field.
"""
from typing import Any, Dict, List, Optional, Union
import httpx
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues
from litellm.types.utils import EmbeddingResponse, Usage
class VoyageMultimodalEmbeddingError(BaseLLMException):
def __init__(
self,
status_code: int,
message: str,
headers: Union[dict, httpx.Headers] = {},
):
self.status_code = status_code
self.message = message
self.request = httpx.Request(
method="POST", url="https://api.voyageai.com/v1/multimodalembeddings"
)
self.response = httpx.Response(status_code=status_code, request=self.request)
super().__init__(
status_code=status_code,
message=message,
headers=headers,
)
class VoyageMultimodalEmbeddingConfig(BaseEmbeddingConfig):
"""
Reference: https://docs.voyageai.com/reference/multimodal-embeddings-api
"""
@staticmethod
def is_multimodal_embeddings(model: str) -> bool:
return "multimodal" in model.lower()
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
if api_base:
if not api_base.endswith("/multimodalembeddings"):
api_base = f"{api_base}/multimodalembeddings"
return api_base
return "https://api.voyageai.com/v1/multimodalembeddings"
def get_supported_openai_params(self, model: str) -> list:
return ["dimensions"]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
if "dimensions" in non_default_params:
optional_params["output_dimension"] = non_default_params["dimensions"]
return optional_params
def validate_environment(
self,
headers: dict,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> dict:
if api_key is None:
api_key = (
get_secret_str("VOYAGE_API_KEY")
or get_secret_str("VOYAGE_AI_API_KEY")
or get_secret_str("VOYAGE_AI_TOKEN")
)
if not api_key:
raise ValueError(
"Voyage API key is required for multimodal embeddings. "
"Set VOYAGE_API_KEY / VOYAGE_AI_API_KEY / VOYAGE_AI_TOKEN "
"or pass `api_key` explicitly."
)
return {"Authorization": f"Bearer {api_key}"}
def _normalize_content_item(self, item: Dict[str, Any]) -> Dict[str, Any]:
item_type = item.get("type")
if item_type == "image_url":
image_url = item.get("image_url")
if isinstance(image_url, dict):
image_url = image_url.get("url")
if image_url is None:
raise ValueError(
"Voyage multimodal embeddings require a non-empty `image_url`. "
"Got an image content block without a `url`."
)
if isinstance(image_url, str) and image_url.startswith("data:image/"):
_, _, encoded = image_url.partition(",")
return {"type": "image_base64", "image_base64": encoded}
return {"type": "image_url", "image_url": image_url}
return item
def _normalize_input_item(self, item: Any) -> Dict[str, Any]:
if isinstance(item, str):
return {"content": [{"type": "text", "text": item}]}
if isinstance(item, dict) and "content" in item:
content = item.get("content") or []
return {
**item,
"content": [
self._normalize_content_item(content_item)
for content_item in content
],
}
return item
def transform_embedding_request(
self,
model: str,
input: AllEmbeddingInputValues,
optional_params: dict,
headers: dict,
) -> dict:
inputs = input if isinstance(input, list) else [input]
return {
"inputs": [self._normalize_input_item(item) for item in inputs],
"model": model,
**optional_params,
}
def transform_embedding_response(
self,
model: str,
raw_response: httpx.Response,
model_response: EmbeddingResponse,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str] = None,
request_data: dict = {},
optional_params: dict = {},
litellm_params: dict = {},
) -> EmbeddingResponse:
try:
raw_response_json = raw_response.json()
except Exception:
raise VoyageMultimodalEmbeddingError(
message=raw_response.text, status_code=raw_response.status_code
)
model_response.model = raw_response_json.get("model")
model_response.data = raw_response_json.get("data")
model_response.object = raw_response_json.get("object")
usage_payload = raw_response_json.get("usage", {})
total_tokens = usage_payload.get("total_tokens", 0)
model_response.usage = Usage(
prompt_tokens=total_tokens,
total_tokens=total_tokens,
)
return model_response
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
) -> BaseLLMException:
return VoyageMultimodalEmbeddingError(
message=error_message, status_code=status_code, headers=headers
)

View file

@ -86,6 +86,7 @@ from litellm.litellm_core_utils.audio_utils.utils import (
get_audio_file_for_health_check,
)
from litellm.litellm_core_utils.completion_timeout import CompletionTimeout
from litellm.litellm_core_utils.get_litellm_params import OPTIONAL_KWARGS_KEYS
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.get_provider_specific_headers import (
ProviderSpecificHeaderUtils,
@ -1407,11 +1408,19 @@ def completion( # type: ignore # noqa: PLR0915
if deployment_id is not None: # azure llms
model = deployment_id
custom_llm_provider = "azure"
_supplemental_provider_params = {
k: kwargs[k] for k in OPTIONAL_KWARGS_KEYS if k in kwargs
}
model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider(
model=model,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
api_key=api_key,
litellm_params=(
GenericLiteLLMParams(**_supplemental_provider_params)
if _supplemental_provider_params
else None
),
)
## RESPONSES API BRIDGE LOGIC ## - check early and normalize model name

View file

@ -4409,6 +4409,23 @@
"/v1/audio/transcriptions"
]
},
"azure/gpt-realtime-whisper": {
"input_cost_per_second": 0.0002833333333333333,
"litellm_provider": "azure",
"mode": "audio_transcription",
"source": "https://learn.microsoft.com/en-us/azure/ai-foundry/openai/concepts/gpt-realtime-whisper",
"supported_endpoints": [
"/v1/realtime",
"/v1/realtime/transcription_sessions"
],
"supported_modalities": [
"audio"
],
"supported_output_modalities": [
"text"
],
"supports_audio_input": true
},
"azure/gpt-5.1-2025-11-13": {
"cache_read_input_token_cost": 1.25e-07,
"cache_read_input_token_cost_priority": 2.5e-07,
@ -7557,6 +7574,45 @@
"supports_function_calling": true,
"supports_tool_choice": true
},
"azure_ai/deepseek-v3.1": {
"input_cost_per_token": 1.23e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 131072,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.94e-06,
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/deepseek/",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"azure_ai/deepseek-v4-pro": {
"input_cost_per_token": 1.74e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 1000000,
"max_output_tokens": 384000,
"max_tokens": 384000,
"mode": "chat",
"output_cost_per_token": 3.48e-06,
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/deepseek/",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"azure_ai/deepseek-v4-flash": {
"input_cost_per_token": 1.9e-07,
"litellm_provider": "azure_ai",
"max_input_tokens": 1000000,
"max_output_tokens": 384000,
"max_tokens": 384000,
"mode": "chat",
"output_cost_per_token": 5.1e-07,
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/deepseek/",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"azure_ai/embed-v-4-0": {
"input_cost_per_token": 1.2e-07,
"litellm_provider": "azure_ai",
@ -35796,7 +35852,17 @@
"max_input_tokens": 32000,
"max_tokens": 32000,
"mode": "embedding",
"output_cost_per_token": 0.0
"output_cost_per_token": 0.0,
"supports_vision": true
},
"voyage/voyage-multimodal-3.5": {
"input_cost_per_token": 1.2e-07,
"litellm_provider": "voyage",
"max_input_tokens": 32000,
"max_tokens": 32000,
"mode": "embedding",
"output_cost_per_token": 0.0,
"supports_vision": true
},
"wandb/openai/gpt-oss-120b": {
"max_tokens": 131072,
@ -40916,6 +40982,23 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-realtime-whisper": {
"input_cost_per_second": 0.0002833333333333333,
"litellm_provider": "openai",
"mode": "audio_transcription",
"source": "https://platform.openai.com/docs/models/gpt-realtime-whisper",
"supported_endpoints": [
"/v1/realtime",
"/v1/realtime/transcription_sessions"
],
"supported_modalities": [
"audio"
],
"supported_output_modalities": [
"text"
],
"supports_audio_input": true
},
"sora-2": {
"litellm_provider": "openai",
"mode": "video_generation",
@ -41680,6 +41763,48 @@
"supports_tool_choice": true,
"supports_vision": true
},
"bedrock_mantle/google.gemma-4-31b": {
"input_cost_per_token": 1.4e-07,
"output_cost_per_token": 4e-07,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 256000,
"max_output_tokens": 256000,
"max_tokens": 256000,
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": false,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_vision": true
},
"bedrock_mantle/google.gemma-4-26b-a4b": {
"input_cost_per_token": 1.3e-07,
"output_cost_per_token": 4e-07,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 256000,
"max_output_tokens": 256000,
"max_tokens": 256000,
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": false,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_vision": true
},
"bedrock_mantle/google.gemma-4-e2b": {
"input_cost_per_token": 4e-08,
"output_cost_per_token": 8e-08,
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 128000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": false,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_vision": true
},
"volcengine/doubao-seed-2-0-pro-260215": {
"litellm_provider": "volcengine",
"max_input_tokens": 256000,

View file

@ -20,7 +20,6 @@ from typing import (
import httpx
from httpx._types import CookieTypes, QueryParamTypes, RequestFiles
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
@ -201,12 +200,6 @@ def llm_passthrough_route(
_is_async = allm_passthrough_route
if client is None:
if _is_async:
client = litellm.module_level_aclient
else:
client = litellm.module_level_client
litellm_logging_obj = cast("LiteLLMLoggingObj", kwargs.get("litellm_logging_obj"))
model, custom_llm_provider, api_key, api_base = get_llm_provider(
@ -218,6 +211,26 @@ def llm_passthrough_route(
litellm_params_dict = get_litellm_params(**kwargs)
if client is None:
from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
get_async_httpx_client,
)
from litellm.passthrough.timeout_utils import resolve_llm_passthrough_timeout
from litellm.types.llms.custom_http import httpxSpecialProvider
resolved_timeout = resolve_llm_passthrough_timeout(
kwargs=kwargs,
litellm_params=litellm_params_dict,
)
if _is_async:
client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.PassThroughEndpoint,
params={"timeout": resolved_timeout},
)
else:
client = _get_httpx_client(params={"timeout": resolved_timeout})
# Add model_id to litellm_params if present in kwargs (for Bedrock Application Inference Profiles)
if "model_id" in kwargs:
litellm_params_dict["model_id"] = kwargs["model_id"]

View file

@ -0,0 +1,58 @@
import sys
from typing import Optional
DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS = 600.0
def resolve_pass_through_request_timeout(
endpoint_timeout: Optional[float] = None,
) -> float:
"""
Resolve the upstream httpx timeout for pass_through_request.
Precedence: per-endpoint timeout -> general_settings.pass_through_request_timeout -> 600s default.
Uses sys.modules to read general_settings only when the proxy module is already
loaded, avoiding a fastapi transitive import in pure SDK contexts.
"""
if endpoint_timeout is not None:
return float(endpoint_timeout)
try:
proxy_server = sys.modules.get("litellm.proxy.proxy_server")
if proxy_server is not None:
global_timeout = getattr(proxy_server, "general_settings", {}).get(
"pass_through_request_timeout"
)
if global_timeout is not None:
return float(global_timeout)
except Exception:
pass
return DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS
def resolve_llm_passthrough_timeout(
kwargs: Optional[dict] = None,
litellm_params: Optional[dict] = None,
router_timeout: Optional[float] = None,
) -> float:
"""
Resolve upstream httpx timeout for SDK native passthrough (e.g. Bedrock /converse).
Precedence: kwargs timeout/request_timeout -> litellm_params timeout/request_timeout
-> router_timeout -> general_settings.pass_through_request_timeout -> 600s default.
"""
kwargs = kwargs or {}
litellm_params = litellm_params or {}
for source in (kwargs, litellm_params):
for key in ("timeout", "request_timeout"):
val = source.get(key)
if val is not None:
return float(val)
if router_timeout is not None:
return float(router_timeout)
return resolve_pass_through_request_timeout()

View file

@ -10,8 +10,9 @@ class MCPUpstreamAuthError(Exception):
(typically HTTP 401) and the gateway should surface it transparently to
the client instead of swallowing it.
Only relevant for pass-through MCP servers (see
``MCPServer.is_oauth_passthrough``). The gateway converts this exception
Relevant for MCP servers that delegate OAuth to the upstream server,
including pass-through servers and OAuth2 servers with
``delegate_auth_to_upstream`` enabled. The gateway converts this exception
into an HTTP 401 response on single-server routes, preserving any
``WWW-Authenticate`` challenge emitted by the upstream so standards-
compliant MCP clients can trigger the upstream OAuth flow.

View file

@ -2777,28 +2777,40 @@ class MCPServerManager:
Uses anyio.fail_after() instead of asyncio.wait_for() to avoid conflicts
with the MCP SDK's anyio TaskGroup. See GitHub issue #20715 for details.
For pass-through MCP servers (``MCPServer.is_oauth_passthrough``) an
For OAuth pass-through and upstream-delegated OAuth2 MCP servers, an
upstream HTTP 401 is converted into :class:`MCPUpstreamAuthError`
instead of being swallowed to an empty tool list. That lets the
single-server HTTP routes surface a proper 401 + ``WWW-Authenticate``
challenge so standards-compliant MCP clients trigger the upstream
OAuth flow. Non-pass-through servers keep today's swallow-and-log
behaviour so the multi-server ``/mcp`` aggregator doesn't get
tainted by a single bad server.
OAuth flow. Other servers keep today's swallow-and-log behaviour so
the multi-server ``/mcp`` aggregator doesn't get tainted by a single
bad server.
Args:
client: MCP client instance
server_name: Name of the server for logging
server: Optional MCPServer; when pass-through, auth errors are
re-raised as :class:`MCPUpstreamAuthError`.
server: Optional MCPServer; when upstream auth is delegated, auth
errors are re-raised as :class:`MCPUpstreamAuthError`.
Returns:
List of tools from the server
"""
is_passthrough = bool(server is not None and server.is_oauth_passthrough)
should_surface_upstream_auth = bool(
server is not None
and (
server.is_oauth_passthrough
or (
server.auth_type == MCPAuth.oauth2
and getattr(server, "delegate_auth_to_upstream", False) is True
and not server.has_client_credentials
)
)
)
try:
with anyio.fail_after(MCP_TOOL_LISTING_TIMEOUT):
tools = await client.list_tools(raise_on_error=is_passthrough)
tools = await client.list_tools(
raise_on_error=should_surface_upstream_auth
)
verbose_logger.debug(f"Tools from {server_name}: {tools}")
return tools
except TimeoutError:
@ -2815,12 +2827,12 @@ class MCPServerManager:
)
return []
except Exception as e:
if is_passthrough:
if should_surface_upstream_auth:
auth_info = _extract_upstream_auth_failure(e)
if auth_info is not None:
status_code, www_authenticate = auth_info
verbose_logger.info(
f"Upstream auth failure from pass-through MCP server "
f"Upstream auth failure from MCP server "
f"{server_name}: HTTP {status_code}"
)
raise MCPUpstreamAuthError(

Some files were not shown because too many files have changed in this diff Show more