mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge remote-tracking branch 'upstream/litellm_internal_staging' into deepkeep-as-internal
This commit is contained in:
commit
e1f3ca0c27
563 changed files with 38035 additions and 8147 deletions
49
.github/workflows/osv-scan.yml
vendored
Normal file
49
.github/workflows/osv-scan.yml
vendored
Normal 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
|
||||
6
.github/workflows/test-linting.yml
vendored
6
.github/workflows/test-linting.yml
vendored
|
|
@ -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__}')"
|
||||
|
|
|
|||
10
CLAUDE.md
10
CLAUDE.md
|
|
@ -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
|
||||
|
||||
|
|
|
|||
26
Dockerfile
26
Dockerfile
|
|
@ -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
|
||||
|
|
|
|||
11
Makefile
11
Makefile
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
)
|
||||
|
|
|
|||
99
db_scripts/partition_spend_logs.sql
Normal file
99
db_scripts/partition_spend_logs.sql
Normal 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";
|
||||
69
db_scripts/unpartition_spend_logs.sql
Normal file
69
db_scripts/unpartition_spend_logs.sql
Normal 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;
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
124
litellm/integrations/datadog/datadog_team_handler.py
Normal file
124
litellm/integrations/datadog/datadog_team_handler.py
Normal 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
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
345
litellm/integrations/focus/destinations/mavvrik_destination.py
Normal file
345
litellm/integrations/focus/destinations/mavvrik_destination.py
Normal 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
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
0
litellm/integrations/mavvrik_focus/__init__.py
Normal file
0
litellm/integrations/mavvrik_focus/__init__.py
Normal file
272
litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py
Normal file
272
litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py
Normal 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
|
||||
)
|
||||
10
litellm/integrations/newrelic/__init__.py
Normal file
10
litellm/integrations/newrelic/__init__.py
Normal 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"]
|
||||
926
litellm/integrations/newrelic/newrelic.py
Normal file
926
litellm/integrations/newrelic/newrelic.py
Normal 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")
|
||||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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")),
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 ##############
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 ##########
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,5 @@
|
|||
from litellm.llms.bedrock.passthrough.guardrail_translation.handler import (
|
||||
BedrockPassthroughGuardrailHandler,
|
||||
)
|
||||
|
||||
__all__ = ["BedrockPassthroughGuardrailHandler"]
|
||||
|
|
@ -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
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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", [])
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
183
litellm/llms/voyage/embedding/transformation_multimodal.py
Normal file
183
litellm/llms/voyage/embedding/transformation_multimodal.py
Normal 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
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
58
litellm/passthrough/timeout_utils.py
Normal file
58
litellm/passthrough/timeout_utils.py
Normal 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()
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue