mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
Merge origin/litellm_internal_staging into litellm_anthropic_stream_model_alias
This commit is contained in:
commit
a0cf82b359
452 changed files with 28559 additions and 5525 deletions
|
|
@ -1483,7 +1483,7 @@ jobs:
|
|||
- run:
|
||||
name: Run tests
|
||||
command: |
|
||||
uv run --no-sync python -m pytest -vv tests/local_testing/test_basic_python_version.py -k "not v2_resolver"
|
||||
uv run --no-sync python -m pytest -vv tests/local_testing/test_basic_python_version.py -k "not legacy_resolver"
|
||||
|
||||
installing_litellm_on_python_3_13:
|
||||
docker:
|
||||
|
|
@ -1507,9 +1507,9 @@ jobs:
|
|||
- run:
|
||||
name: Run tests
|
||||
command: |
|
||||
uv run --no-sync python -m pytest -v tests/local_testing/test_basic_python_version.py -k "not v2_resolver"
|
||||
uv run --no-sync python -m pytest -v tests/local_testing/test_basic_python_version.py -k "not legacy_resolver"
|
||||
|
||||
installing_litellm_on_python_v2_migration_resolver:
|
||||
installing_litellm_on_python_legacy_migration_resolver:
|
||||
docker:
|
||||
- *python312_image
|
||||
- image: cimg/postgres:16.0@sha256:b125148bc76e8e8eee5eb3ad6020a3a14110a14e8192f1c645128afebe2e2f84
|
||||
|
|
@ -1536,10 +1536,10 @@ jobs:
|
|||
url: tcp://localhost:5432
|
||||
timeout: "60"
|
||||
- run:
|
||||
name: Run v2 migration resolver proxy smoke test
|
||||
name: Run legacy migration resolver proxy smoke test
|
||||
command: |
|
||||
uv run --no-sync python -m pytest -vv \
|
||||
tests/local_testing/test_basic_python_version.py::test_litellm_proxy_server_config_no_general_settings_v2_resolver
|
||||
tests/local_testing/test_basic_python_version.py::test_litellm_proxy_server_config_no_general_settings_legacy_resolver
|
||||
|
||||
helm_chart_testing:
|
||||
machine:
|
||||
|
|
@ -2879,7 +2879,8 @@ jobs:
|
|||
command: |
|
||||
if grep -q "Error: P1001: Can't reach database server at" docker_output.log && \
|
||||
(grep -q "Database setup failed after multiple retries" docker_output.log || \
|
||||
grep -q "ERROR: Application startup failed. Exiting." docker_output.log); then
|
||||
grep -q "ERROR: Application startup failed. Exiting." docker_output.log || \
|
||||
grep -q "Database migration cannot proceed" docker_output.log); then
|
||||
echo "Expected error found. Test passed."
|
||||
else
|
||||
echo "Expected error not found. Test failed."
|
||||
|
|
@ -3011,7 +3012,7 @@ workflows:
|
|||
filters: *main_branches
|
||||
- installing_litellm_on_python_3_13:
|
||||
filters: *main_branches
|
||||
- installing_litellm_on_python_v2_migration_resolver:
|
||||
- installing_litellm_on_python_legacy_migration_resolver:
|
||||
filters: *main_branches
|
||||
- helm_chart_testing:
|
||||
requires:
|
||||
|
|
|
|||
23
.github/actions/cache-cargo-build/action.yml
vendored
23
.github/actions/cache-cargo-build/action.yml
vendored
|
|
@ -4,17 +4,16 @@ description: >-
|
|||
so only the first job on a given Cargo.lock compiles the bridge from scratch.
|
||||
|
||||
litellm builds through maturin, which compiles litellm-rust/crates/python-bridge
|
||||
in release mode before it can produce a wheel. `uv sync` therefore pays a full
|
||||
build in every job that installs the workspace: measured at 2m40s per unit shard
|
||||
on 2026-08-21, more than the whole unit tier spends running tests. Nothing caught
|
||||
it, because the uv cache holds wheels uv downloads rather than wheels it builds,
|
||||
and a path dependency whose source moves every commit could never hit that cache
|
||||
anyway. Cargo rebuilds only what changed when its target directory survives, so a
|
||||
warm job pays for the bridge crate alone.
|
||||
in the dev profile for editable installs. `uv sync` therefore pays a full build
|
||||
in every job that installs the workspace. Nothing caught it, because the uv cache
|
||||
holds wheels uv downloads rather than wheels it builds, and a path dependency
|
||||
whose source moves every commit could never hit that cache anyway. Cargo rebuilds
|
||||
only what changed when its target directory survives, so a warm job pays for the
|
||||
bridge crate alone.
|
||||
|
||||
The key namespace is separate from test-rust.yml's. Both cache the same directory,
|
||||
but that workflow fills it with debug and clippy artifacts, which a release build
|
||||
cannot reuse, and a shared key would let whichever ran first deny the other a save.
|
||||
The key namespace is separate from test-rust.yml's check and release caches. They
|
||||
cache the same directory for different workloads, and a shared key would let
|
||||
whichever ran first deny the others a save.
|
||||
|
||||
runs:
|
||||
using: composite
|
||||
|
|
@ -26,6 +25,6 @@ runs:
|
|||
~/.cargo/registry
|
||||
~/.cargo/git
|
||||
litellm-rust/target
|
||||
key: ${{ runner.os }}-cargo-release-${{ hashFiles('litellm-rust/Cargo.lock') }}
|
||||
key: ${{ runner.os }}-maturin-dev-${{ hashFiles('litellm-rust/Cargo.lock') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-cargo-release-
|
||||
${{ runner.os }}-maturin-dev-
|
||||
|
|
|
|||
77
.github/workflows/test-redis-compat.yml
vendored
Normal file
77
.github/workflows/test-redis-compat.yml
vendored
Normal file
|
|
@ -0,0 +1,77 @@
|
|||
name: "Unit Tests: Redis Client Version Compatibility"
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths:
|
||||
- "litellm/_redis.py"
|
||||
- "litellm/_redis_credential_provider.py"
|
||||
- "tests/test_litellm/test_redis.py"
|
||||
- "tests/test_litellm/caching/test_redis_connection_pool.py"
|
||||
- ".github/workflows/test-redis-compat.yml"
|
||||
- "pyproject.toml"
|
||||
- "uv.lock"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
redis-compat:
|
||||
name: "redis-py ${{ matrix.redis-version }}"
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
# 5.3.1 is the version pinned in uv.lock (redisvl caps it below 6); the
|
||||
# newer legs prove the inspect.signature introspection in litellm/_redis.py
|
||||
# keeps extracting kwargs on the redis-py releases people actually run now.
|
||||
# Only the exact release 6.0.0 is skipped: rq (pulled by the proxy extra)
|
||||
# specifies `redis != 6`, which excludes 6.0.0 alone, so 6.4.0 stands in
|
||||
# for the 6.x line.
|
||||
redis-version: ["5.3.1", "6.4.0", "7.4.1", "8.0.1"]
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
|
||||
|
||||
- name: Pin redis-py to the matrix version
|
||||
env:
|
||||
REDIS_VERSION: ${{ matrix.redis-version }}
|
||||
run: |
|
||||
uv pip install "redis==${REDIS_VERSION:?}"
|
||||
uv run --no-sync python -c "import redis; assert redis.__version__ == '${REDIS_VERSION:?}', redis.__version__; print('redis-py', redis.__version__)"
|
||||
|
||||
- name: Run redis unit tests
|
||||
run: |
|
||||
uv run --no-sync pytest \
|
||||
tests/test_litellm/test_redis.py \
|
||||
tests/test_litellm/caching/test_redis_connection_pool.py \
|
||||
--tb=short -vv \
|
||||
--reruns 2 \
|
||||
--reruns-delay 1 \
|
||||
--durations=20
|
||||
1
.github/workflows/test-unit.yml
vendored
1
.github/workflows/test-unit.yml
vendored
|
|
@ -103,6 +103,7 @@ jobs:
|
|||
tests/test_litellm/completion_extras
|
||||
tests/test_litellm/compression
|
||||
tests/test_litellm/containers
|
||||
tests/test_litellm/endpoints
|
||||
tests/test_litellm/experimental_mcp_client
|
||||
tests/test_litellm/models
|
||||
tests/test_litellm/repositories
|
||||
|
|
|
|||
|
|
@ -37,7 +37,7 @@ Same applies for filing bug reports and feature requests, with .github/ISSUE_TEM
|
|||
|
||||
If you're resolving a linear ticket, in the "## Linear ticket" section of the PR, say "Resolves LIT-1234", replacing "LIT-1234" with the actual ticket id that you're resolving. If you don't have the ticket id, don't make one up or search for it. Just leave the section blank
|
||||
|
||||
Never use `pytest` commands or the like as "Screenshots / Proof of Fix". We prefer curl'ing a live proxy instance running on localhost:4000 (I like to run it with `python litellm/proxy/proxy_cli.py --config litellm/proxy/dev_config.yaml --detailed_debug --reload --use_v2_migration_resolver 2>&1 | tee litellm.log`; the Admin UI dev server is `npm run dev` in `ui/litellm-dashboard`, served on port 3000) and showing both the command run and the output. Also, it should hit real LLM provider APIs, not mocks, and cost real $$$ because that is the most realistic test. The proof of fix should be exactly what the end user / customer would see / do. The run logs in PR #27703 is a prime example of how to do it (not a huge fan of using a python test script that future me and the team will have no visibility into; I prefer just curl commands or a short list of bash commands (e.g., using `for`)). If it's a UI thing, just tell me which URLs to go to (e.g., http://localhost:4000/ui/?page=logs), where to click, what fields to fill out, etc. along with the other commands to run in an ordered list, and I'll do it myself and post the screenshots after you make the PR
|
||||
Never use `pytest` commands or the like as "Screenshots / Proof of Fix". We prefer curl'ing a live proxy instance running on localhost:4000 (I like to run it with `python litellm/proxy/proxy_cli.py --config litellm/proxy/dev_config.yaml --detailed_debug --reload 2>&1 | tee litellm.log`; the Admin UI dev server is `npm run dev` in `ui/litellm-dashboard`, served on port 3000) and showing both the command run and the output. Also, it should hit real LLM provider APIs, not mocks, and cost real $$$ because that is the most realistic test. The proof of fix should be exactly what the end user / customer would see / do. The run logs in PR #27703 is a prime example of how to do it (not a huge fan of using a python test script that future me and the team will have no visibility into; I prefer just curl commands or a short list of bash commands (e.g., using `for`)). If it's a UI thing, just tell me which URLs to go to (e.g., http://localhost:4000/ui/?page=logs), where to click, what fields to fill out, etc. along with the other commands to run in an ordered list, and I'll do it myself and post the screenshots after you make the PR
|
||||
|
||||
If you ever write any human-facing text (pull requests, issues, commit messages, discussion posts, github comments, release notes, docs, etc.), always follow these guidelines to sound less AI-y:
|
||||
- don't use emojis
|
||||
|
|
|
|||
15
Dockerfile
15
Dockerfile
|
|
@ -1,10 +1,10 @@
|
|||
# syntax=docker/dockerfile:1.7
|
||||
|
||||
# Base image for building
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
|
||||
# Runtime image
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
|
||||
# Pinned by digest like the other base images; bump explicitly on Node upgrades.
|
||||
ARG UI_BUILD_IMAGE=node:24.19-alpine3.24@sha256:d32cdf619f63fe0471182d08996dd516c6275bb5fd31ae06e55a570bd9e1ad43
|
||||
|
|
@ -40,8 +40,8 @@ COPY --from=uvbin /uvx /usr/local/bin/uvx
|
|||
RUN apk add --no-cache \
|
||||
bash \
|
||||
gcc \
|
||||
python3 \
|
||||
python3-dev \
|
||||
python-3.13 \
|
||||
python-3.13-dev \
|
||||
rust \
|
||||
openssl \
|
||||
openssl-dev \
|
||||
|
|
@ -51,6 +51,7 @@ RUN apk add --no-cache \
|
|||
|
||||
ENV UV_PROJECT_ENVIRONMENT=/app/.venv \
|
||||
UV_LINK_MODE=copy \
|
||||
UV_PYTHON_DOWNLOADS=0 \
|
||||
PATH="/app/.venv/bin:${PATH}"
|
||||
|
||||
# Copy dependency metadata first for layer caching
|
||||
|
|
@ -65,7 +66,7 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
|
|||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra saml \
|
||||
--python python3
|
||||
--python python3.13
|
||||
|
||||
# Copy full source tree
|
||||
COPY . .
|
||||
|
|
@ -86,7 +87,7 @@ RUN uv sync --frozen --no-default-groups --no-editable \
|
|||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra saml \
|
||||
--python python3
|
||||
--python python3.13
|
||||
|
||||
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
npm_config_cache=/root/.npm \
|
||||
|
|
@ -101,7 +102,7 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
|
|||
USER root
|
||||
|
||||
# node (without npm) is required by the prisma CLI at runtime
|
||||
RUN apk add --no-cache bash openssl tzdata nodejs python3 libsndfile
|
||||
RUN apk add --no-cache bash openssl tzdata nodejs python-3.13 libsndfile
|
||||
|
||||
WORKDIR /app
|
||||
ENV PATH="/app/.venv/bin:${PATH}" \
|
||||
|
|
|
|||
|
|
@ -354,6 +354,8 @@ curl -X POST 'http://0.0.0.0:4000/v1/chat/completions' \
|
|||
| [Petals (`petals`)](https://docs.litellm.ai/docs/providers/petals) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Pinstripes (`pinstripes`)](https://docs.litellm.ai/docs/providers/pinstripes) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Predibase (`predibase`)](https://docs.litellm.ai/docs/providers/predibase) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Qwen AI Platform (`qwen_ai_platform`)](https://docs.litellm.ai/docs/providers/qwencloud) | ✅ | ✅ | ✅ | ✅ | ✅ | | | | | ✅ |
|
||||
| [QwenCloud (`qwencloud`)](https://docs.litellm.ai/docs/providers/qwencloud) | ✅ | ✅ | ✅ | ✅ | ✅ | | | | | ✅ |
|
||||
| [Recraft (`recraft`)](https://docs.litellm.ai/docs/providers/recraft) | | | | | ✅ | | | | | |
|
||||
| [Replicate (`replicate`)](https://docs.litellm.ai/docs/providers/replicate) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Sagemaker Chat (`sagemaker_chat`)](https://docs.litellm.ai/docs/providers/aws_sagemaker) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
|
||||
|
||||
FROM $UV_IMAGE AS uvbin
|
||||
|
|
@ -16,7 +16,7 @@ COPY --from=uvbin /uv /uvx /usr/local/bin/
|
|||
# instead of nodeenv downloading one whose dynamic deps may not be in Wolfi
|
||||
# (e.g. Node 26.2.0 needs libatomic). Retry for transient apk.cgr.dev flakes.
|
||||
RUN for i in 1 2 3; do \
|
||||
apk add --no-cache bash gcc python3 python3-dev openssl openssl-dev libsndfile nodejs npm && break; \
|
||||
apk add --no-cache bash gcc python-3.13 python-3.13-dev openssl openssl-dev libsndfile nodejs npm && break; \
|
||||
[ $i = 3 ] && { echo "apk add failed after 3 retries" >&2; exit 1; }; \
|
||||
sleep 5; \
|
||||
done
|
||||
|
|
@ -46,7 +46,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
|||
--extra proxy-runtime \
|
||||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--python python3
|
||||
--python python3.13
|
||||
|
||||
# Stage 2 — copy source and install the project + workspace members.
|
||||
COPY . .
|
||||
|
|
@ -57,7 +57,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
|||
--extra proxy-runtime \
|
||||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--python python3
|
||||
--python python3.13
|
||||
|
||||
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
npm_config_cache=/root/.npm \
|
||||
|
|
@ -71,7 +71,7 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
|
|||
USER root
|
||||
|
||||
RUN for i in 1 2 3; do \
|
||||
apk add --no-cache bash openssl tzdata python3 libsndfile libatomic && break; \
|
||||
apk add --no-cache bash openssl tzdata python-3.13 libsndfile libatomic && break; \
|
||||
[ $i = 3 ] && { echo "apk add failed after 3 retries" >&2; exit 1; }; \
|
||||
sleep 5; \
|
||||
done
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"limit": 16171
|
||||
"limit": 14076
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2226
|
||||
"limit": 2216
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"limit": 319
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 19
|
||||
},
|
||||
"reportExplicitAny": {
|
||||
"limit": 5199
|
||||
"limit": 4128
|
||||
},
|
||||
"reportFunctionMemberAccess": {
|
||||
"limit": 7
|
||||
|
|
@ -42,7 +42,7 @@
|
|||
"limit": 12
|
||||
},
|
||||
"reportIndexIssue": {
|
||||
"limit": 35
|
||||
"limit": 25
|
||||
},
|
||||
"reportInvalidTypeForm": {
|
||||
"limit": 34
|
||||
|
|
@ -54,10 +54,10 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportMissingParameterType": {
|
||||
"limit": 5611
|
||||
"limit": 5601
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15350
|
||||
"limit": 15306
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 40
|
||||
|
|
@ -99,19 +99,19 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 44368
|
||||
"limit": 44364
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 109
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 38468
|
||||
"limit": 38350
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19665
|
||||
"limit": 19626
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 30066
|
||||
"limit": 29890
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 111
|
||||
|
|
@ -123,7 +123,7 @@
|
|||
"limit": 5
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"limit": 828
|
||||
"limit": 826
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"limit": 0
|
||||
|
|
@ -141,6 +141,6 @@
|
|||
"limit": 543
|
||||
},
|
||||
"reportUnusedVariable": {
|
||||
"limit": 139
|
||||
"limit": 137
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,10 +1,10 @@
|
|||
# syntax=docker/dockerfile:1.7
|
||||
|
||||
# Base image for building
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
|
||||
# Runtime image
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
|
||||
# Pinned by digest like the other base images; bump explicitly on Node upgrades.
|
||||
ARG UI_BUILD_IMAGE=node:24.19-alpine3.24@sha256:d32cdf619f63fe0471182d08996dd516c6275bb5fd31ae06e55a570bd9e1ad43
|
||||
|
|
@ -39,8 +39,8 @@ COPY --from=uvbin /uvx /usr/local/bin/uvx
|
|||
RUN apk add --no-cache \
|
||||
bash \
|
||||
gcc \
|
||||
python3 \
|
||||
python3-dev \
|
||||
python-3.13 \
|
||||
python-3.13-dev \
|
||||
openssl \
|
||||
openssl-dev \
|
||||
nodejs \
|
||||
|
|
@ -49,6 +49,7 @@ RUN apk add --no-cache \
|
|||
|
||||
ENV UV_PROJECT_ENVIRONMENT=/app/.venv \
|
||||
UV_LINK_MODE=copy \
|
||||
UV_PYTHON_DOWNLOADS=0 \
|
||||
PATH="/app/.venv/bin:${PATH}"
|
||||
|
||||
# Copy dependency metadata first for layer caching
|
||||
|
|
@ -63,7 +64,7 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
|
|||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra saml \
|
||||
--python python3
|
||||
--python python3.13
|
||||
|
||||
# Copy full source tree
|
||||
COPY . .
|
||||
|
|
@ -84,7 +85,7 @@ RUN uv sync --frozen --no-default-groups --no-editable \
|
|||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra saml \
|
||||
--python python3
|
||||
--python python3.13
|
||||
|
||||
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
npm_config_cache=/root/.npm \
|
||||
|
|
@ -98,7 +99,7 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
|
|||
USER root
|
||||
|
||||
# node (without npm) is required by the prisma CLI at runtime
|
||||
RUN apk add --no-cache bash openssl tzdata nodejs python3 libsndfile
|
||||
RUN apk add --no-cache bash openssl tzdata nodejs python-3.13 libsndfile
|
||||
|
||||
WORKDIR /app
|
||||
ENV PATH="/app/.venv/bin:${PATH}" \
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
# syntax=docker/dockerfile:1.7
|
||||
|
||||
# Base images
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG PROXY_EXTRAS_SOURCE=published
|
||||
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
|
||||
# Pinned by digest like the other base images; bump explicitly on Node upgrades.
|
||||
|
|
@ -37,8 +37,8 @@ COPY --from=uvbin /uvx /usr/local/bin/uvx
|
|||
|
||||
RUN for i in 1 2 3; do \
|
||||
apk add --no-cache \
|
||||
python3 \
|
||||
python3-dev \
|
||||
python-3.13 \
|
||||
python-3.13-dev \
|
||||
gcc \
|
||||
rust \
|
||||
bash \
|
||||
|
|
@ -52,6 +52,7 @@ RUN for i in 1 2 3; do \
|
|||
|
||||
ENV UV_PROJECT_ENVIRONMENT=/app/.venv \
|
||||
UV_LINK_MODE=copy \
|
||||
UV_PYTHON_DOWNLOADS=0 \
|
||||
PATH="/app/.venv/bin:${PATH}" \
|
||||
LITELLM_NON_ROOT=true \
|
||||
XDG_CACHE_HOME=/app/.cache
|
||||
|
|
@ -69,7 +70,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
|
|||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra saml \
|
||||
--python python3
|
||||
--python python3.13
|
||||
|
||||
# Copy full source tree
|
||||
COPY . .
|
||||
|
|
@ -96,7 +97,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
|
|||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra saml \
|
||||
--python python3 \
|
||||
--python python3.13 \
|
||||
--no-sources-package litellm-proxy-extras; \
|
||||
else \
|
||||
uv sync --frozen --no-default-groups --no-editable \
|
||||
|
|
@ -105,7 +106,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
|
|||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra saml \
|
||||
--python python3; \
|
||||
--python python3.13; \
|
||||
fi
|
||||
|
||||
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
|
|
@ -124,7 +125,7 @@ RUN for i in 1 2 3; do \
|
|||
apk upgrade --no-cache && break || sleep 5; \
|
||||
done && \
|
||||
for i in 1 2 3; do \
|
||||
apk add --no-cache python3 bash openssl tzdata libsndfile nodejs && break || sleep 5; \
|
||||
apk add --no-cache python-3.13 bash openssl tzdata libsndfile nodejs && break || sleep 5; \
|
||||
done
|
||||
|
||||
# Copy only what runtime needs. The application is installed inside the venv;
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ Polls LiteLLM_ManagedObjectTable to check if the batch job is complete, and if t
|
|||
from dataclasses import replace as dataclasses_replace
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Dict, Final, List, Literal, Optional, Tuple, cast
|
||||
from typing import TYPE_CHECKING, Final, List, Literal, Optional, Tuple, cast
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
|
|
@ -87,7 +87,7 @@ class CheckBatchCost:
|
|||
return
|
||||
self.batch_processed_support_confirmed = True
|
||||
|
||||
async def _get_user_info(self, batch_id: str, user_id: Optional[str]) -> Dict[str, Any]:
|
||||
async def _get_user_info(self, batch_id: str, user_id: Optional[str]) -> dict[str, str | None]:
|
||||
"""
|
||||
Look up user email and key alias by user_id for enriching the S3 callback metadata.
|
||||
Returns a dict with user_api_key_user_email and user_api_key_alias (both may be None).
|
||||
|
|
@ -97,8 +97,10 @@ class CheckBatchCost:
|
|||
if not user_id:
|
||||
return {}
|
||||
try:
|
||||
user_row = await self.prisma_client.db.litellm_usertable.find_unique(
|
||||
where={"user_id": user_id}
|
||||
user_row: prisma_models.LiteLLM_UserTable | None = (
|
||||
await self.prisma_client.db.litellm_usertable.find_unique(
|
||||
where={"user_id": user_id}
|
||||
)
|
||||
)
|
||||
if user_row is None:
|
||||
return {}
|
||||
|
|
@ -115,8 +117,10 @@ class CheckBatchCost:
|
|||
if not api_key:
|
||||
return None
|
||||
try:
|
||||
key_row = await self.prisma_client.db.litellm_verificationtoken.find_unique(
|
||||
where={"token": api_key}
|
||||
key_row: prisma_models.LiteLLM_VerificationToken | None = (
|
||||
await self.prisma_client.db.litellm_verificationtoken.find_unique(
|
||||
where={"token": api_key}
|
||||
)
|
||||
)
|
||||
return getattr(key_row, "key_alias", None) if key_row is not None else None
|
||||
except Exception as e:
|
||||
|
|
@ -128,8 +132,10 @@ class CheckBatchCost:
|
|||
if not team_id:
|
||||
return None
|
||||
try:
|
||||
team_row = await self.prisma_client.db.litellm_teamtable.find_unique(
|
||||
where={"team_id": team_id}
|
||||
team_row: prisma_models.LiteLLM_TeamTable | None = (
|
||||
await self.prisma_client.db.litellm_teamtable.find_unique(
|
||||
where={"team_id": team_id}
|
||||
)
|
||||
)
|
||||
return getattr(team_row, "team_alias", None) if team_row is not None else None
|
||||
except Exception as e:
|
||||
|
|
@ -138,7 +144,7 @@ class CheckBatchCost:
|
|||
|
||||
async def _build_creator_attribution_metadata(
|
||||
self, job: "LiteLLM_ManagedObjectTable", batch_id: str
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Rebuild the spend-tracking metadata for the key, team, and tags that created the
|
||||
batch so the batch-cost spend log is attributed the same way a non-batch request
|
||||
|
|
@ -152,7 +158,7 @@ class CheckBatchCost:
|
|||
team_id = getattr(job, "team_id", None)
|
||||
request_tags = getattr(job, "request_tags", None)
|
||||
|
||||
metadata: Dict[str, Any] = {
|
||||
metadata: dict[str, object] = {
|
||||
"user_api_key_user_id": job.created_by,
|
||||
"user_api_key": api_key,
|
||||
"user_api_key_team_id": team_id,
|
||||
|
|
|
|||
|
|
@ -182,6 +182,10 @@ class _ManagedObjectTableActions(Protocol):
|
|||
async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ...
|
||||
|
||||
|
||||
class _SchedulerWithJobLookup(Protocol):
|
||||
def get_job(self, job_id: str) -> object: ...
|
||||
|
||||
|
||||
class _CursorPageArgs(TypedDict, total=False):
|
||||
cursor: Mapping[str, str]
|
||||
skip: int
|
||||
|
|
@ -853,7 +857,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
file_ids.append(file_id)
|
||||
return file_ids
|
||||
|
||||
def get_file_ids_from_responses_input(self, input: Union[str, List[Dict[str, Any]]]) -> List[str]:
|
||||
def get_file_ids_from_responses_input(self, input: Union[str, List[Dict[str, object]]]) -> List[str]:
|
||||
"""
|
||||
Gets file ids from responses API input.
|
||||
|
||||
|
|
@ -878,7 +882,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
# Check for direct input_file type
|
||||
if item.get("type") == "input_file":
|
||||
file_id = item.get("file_id")
|
||||
if file_id:
|
||||
if isinstance(file_id, str) and file_id:
|
||||
file_ids.append(file_id)
|
||||
|
||||
# Check for input_file in content array
|
||||
|
|
@ -887,7 +891,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
for content_item in content:
|
||||
if isinstance(content_item, dict) and content_item.get("type") == "input_file":
|
||||
file_id = content_item.get("file_id")
|
||||
if file_id:
|
||||
if isinstance(file_id, str) and file_id:
|
||||
file_ids.append(file_id)
|
||||
|
||||
return file_ids
|
||||
|
|
@ -1227,7 +1231,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
|
||||
# Handle both output_file_id and error_file_id
|
||||
for file_attr in ["output_file_id", "error_file_id"]:
|
||||
file_id_value = getattr(response, file_attr, None)
|
||||
file_id_value: str | None = getattr(response, file_attr, None)
|
||||
if file_id_value and model_id:
|
||||
decoded_output_file_id = _is_base64_encoded_unified_file_id(file_id_value)
|
||||
if decoded_output_file_id and "llm_output_file_id," in decoded_output_file_id:
|
||||
|
|
@ -1496,7 +1500,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
import litellm.proxy.proxy_server as proxy_server_module
|
||||
|
||||
# Check if the scheduler has the batch cost checking job registered
|
||||
scheduler = getattr(proxy_server_module, "scheduler", None)
|
||||
scheduler: Final[_SchedulerWithJobLookup | None] = getattr(proxy_server_module, "scheduler", None)
|
||||
if scheduler is None:
|
||||
return False
|
||||
|
||||
|
|
@ -1542,7 +1546,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
)
|
||||
MAX_MATCHES_TO_RETURN = 10
|
||||
|
||||
batches = await self.prisma_client.db.litellm_managedobjecttable.find_many(
|
||||
batches = await _managed_object_table(self.prisma_client).find_many(
|
||||
where={
|
||||
"file_purpose": "batch",
|
||||
"batch_processed": False,
|
||||
|
|
@ -1552,11 +1556,14 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
order={"created_at": "desc"},
|
||||
)
|
||||
|
||||
referencing_batches = []
|
||||
referencing_batches: Final[list[dict[str, object]]] = []
|
||||
for batch in batches:
|
||||
try:
|
||||
# Parse the batch file_object to check for file references
|
||||
batch_data = json.loads(batch.file_object) if isinstance(batch.file_object, str) else batch.file_object
|
||||
decoded_file_object = _decode_json_blob(batch.file_object)
|
||||
batch_data: Mapping[str, object] = (
|
||||
decoded_file_object if isinstance(decoded_file_object, Mapping) else {}
|
||||
)
|
||||
|
||||
# Extract file IDs from batch
|
||||
# Batches typically reference the unified file ID in input_file_id
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.62"
|
||||
version = "0.1.63"
|
||||
description = "Package for LiteLLM Enterprise features"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.1.62"
|
||||
version = "0.1.63"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
|
||||
|
||||
FROM $UV_IMAGE AS uvbin
|
||||
|
|
@ -16,7 +16,7 @@ COPY --from=uvbin /uv /uvx /usr/local/bin/
|
|||
# instead of nodeenv downloading one whose dynamic deps may not be in Wolfi
|
||||
# (e.g. Node 26.2.0 needs libatomic). Retry for transient apk.cgr.dev flakes.
|
||||
RUN for i in 1 2 3; do \
|
||||
apk add --no-cache bash gcc python3 python3-dev openssl openssl-dev libsndfile nodejs npm && break; \
|
||||
apk add --no-cache bash gcc python-3.13 python-3.13-dev openssl openssl-dev libsndfile nodejs npm && break; \
|
||||
[ $i = 3 ] && { echo "apk add failed after 3 retries" >&2; exit 1; }; \
|
||||
sleep 5; \
|
||||
done
|
||||
|
|
@ -47,7 +47,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
|||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra bedrock-realtime \
|
||||
--python python3
|
||||
--python python3.13
|
||||
|
||||
# Stage 2 — copy source and install the project + workspace members.
|
||||
COPY . .
|
||||
|
|
@ -59,7 +59,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
|||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra bedrock-realtime \
|
||||
--python python3
|
||||
--python python3.13
|
||||
|
||||
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
npm_config_cache=/root/.npm \
|
||||
|
|
@ -73,7 +73,7 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
|
|||
USER root
|
||||
|
||||
RUN for i in 1 2 3; do \
|
||||
apk add --no-cache bash openssl tzdata python3 libsndfile libatomic && break; \
|
||||
apk add --no-cache bash openssl tzdata python-3.13 libsndfile libatomic && break; \
|
||||
[ $i = 3 ] && { echo "apk add failed after 3 retries" >&2; exit 1; }; \
|
||||
sleep 5; \
|
||||
done
|
||||
|
|
|
|||
|
|
@ -86,6 +86,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
|
|||
"/comprehendmedical",
|
||||
"/cohere/",
|
||||
"/gemini/",
|
||||
"/gigachat/",
|
||||
"/google/",
|
||||
"/vertex_ai/",
|
||||
"/vertex-ai/",
|
||||
|
|
|
|||
|
|
@ -7,6 +7,10 @@ metadata:
|
|||
{{- include "litellm.commonLabels" . | nindent 4 }}
|
||||
app.kubernetes.io/component: backend
|
||||
spec:
|
||||
{{- with .Values.backend.strategy }}
|
||||
strategy:
|
||||
{{- toYaml . | nindent 4 }}
|
||||
{{- end }}
|
||||
selector:
|
||||
matchLabels:
|
||||
{{- include "litellm.backend.selectorLabels" . | nindent 6 }}
|
||||
|
|
|
|||
|
|
@ -7,6 +7,10 @@ metadata:
|
|||
{{- include "litellm.commonLabels" . | nindent 4 }}
|
||||
app.kubernetes.io/component: gateway
|
||||
spec:
|
||||
{{- with .Values.gateway.strategy }}
|
||||
strategy:
|
||||
{{- toYaml . | nindent 4 }}
|
||||
{{- end }}
|
||||
selector:
|
||||
matchLabels:
|
||||
{{- include "litellm.gateway.selectorLabels" . | nindent 6 }}
|
||||
|
|
|
|||
|
|
@ -7,6 +7,8 @@
|
|||
#
|
||||
# Running this pre-upgrade closes the window where new application pods would
|
||||
# otherwise serve traffic against the previous release's unmigrated schema.
|
||||
# Argo CD users can swap the Helm hook for a PreSync hook through
|
||||
# `migrationJob.hooks`, which re-runs the Job on every sync.
|
||||
apiVersion: batch/v1
|
||||
kind: Job
|
||||
metadata:
|
||||
|
|
@ -14,10 +16,18 @@ metadata:
|
|||
labels:
|
||||
{{- include "litellm.commonLabels" . | nindent 4 }}
|
||||
app.kubernetes.io/component: migrations
|
||||
{{- if or .Values.migrationJob.hooks.helm.enabled .Values.migrationJob.hooks.argocd.enabled }}
|
||||
annotations:
|
||||
{{- if .Values.migrationJob.hooks.helm.enabled }}
|
||||
helm.sh/hook: pre-install,pre-upgrade
|
||||
helm.sh/hook-delete-policy: before-hook-creation
|
||||
helm.sh/hook-weight: "0"
|
||||
helm.sh/hook-weight: {{ .Values.migrationJob.hooks.helm.weight | default "0" | quote }}
|
||||
{{- end }}
|
||||
{{- if .Values.migrationJob.hooks.argocd.enabled }}
|
||||
argocd.argoproj.io/hook: PreSync
|
||||
argocd.argoproj.io/hook-delete-policy: BeforeHookCreation
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
spec:
|
||||
backoffLimit: {{ .Values.migrationJob.backoffLimit }}
|
||||
ttlSecondsAfterFinished: {{ .Values.migrationJob.ttlSecondsAfterFinished }}
|
||||
|
|
|
|||
|
|
@ -7,6 +7,10 @@ metadata:
|
|||
{{- include "litellm.commonLabels" . | nindent 4 }}
|
||||
app.kubernetes.io/component: ui
|
||||
spec:
|
||||
{{- with .Values.ui.strategy }}
|
||||
strategy:
|
||||
{{- toYaml . | nindent 4 }}
|
||||
{{- end }}
|
||||
selector:
|
||||
matchLabels:
|
||||
{{- include "litellm.ui.selectorLabels" . | nindent 6 }}
|
||||
|
|
|
|||
63
helm/litellm/tests/migration_job_hooks_tests.yaml
Normal file
63
helm/litellm/tests/migration_job_hooks_tests.yaml
Normal file
|
|
@ -0,0 +1,63 @@
|
|||
suite: test migrations Job hook annotations
|
||||
templates:
|
||||
- migrations-job.yaml
|
||||
values:
|
||||
- ./values/required.yaml
|
||||
tests:
|
||||
- it: runs as a Helm pre-install / pre-upgrade hook by default
|
||||
asserts:
|
||||
- equal:
|
||||
path: metadata.annotations["helm.sh/hook"]
|
||||
value: pre-install,pre-upgrade
|
||||
- equal:
|
||||
path: metadata.annotations["helm.sh/hook-delete-policy"]
|
||||
value: before-hook-creation
|
||||
- equal:
|
||||
path: metadata.annotations["helm.sh/hook-weight"]
|
||||
value: "0"
|
||||
- notExists:
|
||||
path: metadata.annotations["argocd.argoproj.io/hook"]
|
||||
|
||||
- it: adds the Argo CD PreSync hook when asked
|
||||
set:
|
||||
migrationJob.hooks.argocd.enabled: true
|
||||
asserts:
|
||||
- equal:
|
||||
path: metadata.annotations["argocd.argoproj.io/hook"]
|
||||
value: PreSync
|
||||
- equal:
|
||||
path: metadata.annotations["argocd.argoproj.io/hook-delete-policy"]
|
||||
value: BeforeHookCreation
|
||||
|
||||
- it: drops the Helm hook so Argo CD owns the Job
|
||||
set:
|
||||
migrationJob.hooks.argocd.enabled: true
|
||||
migrationJob.hooks.helm.enabled: false
|
||||
asserts:
|
||||
- equal:
|
||||
path: metadata.annotations["argocd.argoproj.io/hook"]
|
||||
value: PreSync
|
||||
- notExists:
|
||||
path: metadata.annotations["helm.sh/hook"]
|
||||
- notExists:
|
||||
path: metadata.annotations["helm.sh/hook-delete-policy"]
|
||||
- notExists:
|
||||
path: metadata.annotations["helm.sh/hook-weight"]
|
||||
|
||||
- it: renders an ordinary Job when both hooks are disabled
|
||||
set:
|
||||
migrationJob.hooks.helm.enabled: false
|
||||
asserts:
|
||||
- notExists:
|
||||
path: metadata.annotations
|
||||
- equal:
|
||||
path: kind
|
||||
value: Job
|
||||
|
||||
- it: honours a custom Helm hook weight
|
||||
set:
|
||||
migrationJob.hooks.helm.weight: "-5"
|
||||
asserts:
|
||||
- equal:
|
||||
path: metadata.annotations["helm.sh/hook-weight"]
|
||||
value: "-5"
|
||||
66
helm/litellm/tests/rollout_strategy_tests.yaml
Normal file
66
helm/litellm/tests/rollout_strategy_tests.yaml
Normal file
|
|
@ -0,0 +1,66 @@
|
|||
suite: test rolling update strategy on the component deployments
|
||||
templates:
|
||||
- gateway/deployment.yaml
|
||||
- gateway/configmap.yaml
|
||||
- backend/deployment.yaml
|
||||
- ui/deployment.yaml
|
||||
values:
|
||||
- ./values/required.yaml
|
||||
tests:
|
||||
- it: leaves the strategy to Kubernetes defaults when unset
|
||||
asserts:
|
||||
- notExists:
|
||||
path: spec.strategy
|
||||
|
||||
- it: renders the configured strategy on each deployment
|
||||
set:
|
||||
gateway.strategy:
|
||||
type: RollingUpdate
|
||||
rollingUpdate:
|
||||
maxUnavailable: 0
|
||||
maxSurge: 1
|
||||
backend.strategy:
|
||||
type: RollingUpdate
|
||||
rollingUpdate:
|
||||
maxUnavailable: "25%"
|
||||
maxSurge: 2
|
||||
ui.strategy:
|
||||
type: Recreate
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.strategy
|
||||
value:
|
||||
type: RollingUpdate
|
||||
rollingUpdate:
|
||||
maxUnavailable: 0
|
||||
maxSurge: 1
|
||||
template: gateway/deployment.yaml
|
||||
- equal:
|
||||
path: spec.strategy
|
||||
value:
|
||||
type: RollingUpdate
|
||||
rollingUpdate:
|
||||
maxUnavailable: 25%
|
||||
maxSurge: 2
|
||||
template: backend/deployment.yaml
|
||||
- equal:
|
||||
path: spec.strategy
|
||||
value:
|
||||
type: Recreate
|
||||
template: ui/deployment.yaml
|
||||
|
||||
- it: keeps a component on the cluster default when only another one sets a strategy
|
||||
set:
|
||||
gateway.strategy:
|
||||
type: Recreate
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.strategy.type
|
||||
value: Recreate
|
||||
template: gateway/deployment.yaml
|
||||
- notExists:
|
||||
path: spec.strategy
|
||||
template: backend/deployment.yaml
|
||||
- notExists:
|
||||
path: spec.strategy
|
||||
template: ui/deployment.yaml
|
||||
|
|
@ -75,6 +75,22 @@ serviceAccounts:
|
|||
# generate` — the migration engine doesn't need the generated client.
|
||||
migrationJob:
|
||||
enabled: true
|
||||
# Which controller is responsible for running the Job.
|
||||
#
|
||||
# `helm.enabled` renders the Helm pre-install / pre-upgrade hook, so the Job
|
||||
# runs whenever `helm upgrade` sees a change to apply. `argocd.enabled`
|
||||
# renders an Argo CD PreSync hook instead, which runs the Job on every sync
|
||||
# even when the rendered manifests are unchanged: the way to re-run
|
||||
# migrations on demand from a GitOps pipeline. Turning the Helm hook off
|
||||
# while the Argo CD hook is on leaves the Job out of Helm's own upgrade
|
||||
# path, which is what Argo CD users want since Argo, not Helm, applies the
|
||||
# manifests.
|
||||
hooks:
|
||||
helm:
|
||||
enabled: true
|
||||
weight: "0"
|
||||
argocd:
|
||||
enabled: false
|
||||
backoffLimit: 4
|
||||
ttlSecondsAfterFinished: 120
|
||||
# Wall-clock budget for the whole Job, shared across every `backoffLimit`
|
||||
|
|
@ -257,6 +273,15 @@ gateway:
|
|||
initialDelaySeconds: 5
|
||||
periodSeconds: 10
|
||||
timeoutSeconds: 10
|
||||
# Rolling update tuning for the gateway Deployment. Empty by default, so
|
||||
# Kubernetes applies its own RollingUpdate defaults (25% maxSurge /
|
||||
# 25% maxUnavailable). Example, for a surge-only rollout behind a load
|
||||
# balancer that must never lose capacity:
|
||||
# type: RollingUpdate
|
||||
# rollingUpdate:
|
||||
# maxUnavailable: 0
|
||||
# maxSurge: 1
|
||||
strategy: {}
|
||||
# Optional startupProbe. Empty by default, so existing installs are unchanged
|
||||
# and liveness/readiness apply from container start. Set it to gate
|
||||
# liveness/readiness until a slow cold start finishes — a high failureThreshold
|
||||
|
|
@ -369,6 +394,8 @@ backend:
|
|||
initialDelaySeconds: 5
|
||||
periodSeconds: 10
|
||||
timeoutSeconds: 10
|
||||
# Same shape as gateway.strategy.
|
||||
strategy: {}
|
||||
# Optional startupProbe; same shape as gateway.startupProbe. Empty by default.
|
||||
startupProbe: {}
|
||||
hpa:
|
||||
|
|
@ -433,6 +460,8 @@ ui:
|
|||
httpGet: { path: /, port: http }
|
||||
initialDelaySeconds: 2
|
||||
periodSeconds: 10
|
||||
# Same shape as gateway.strategy.
|
||||
strategy: {}
|
||||
# Optional startupProbe; same shape as gateway.startupProbe. Empty by default.
|
||||
startupProbe: {}
|
||||
hpa:
|
||||
|
|
|
|||
|
|
@ -0,0 +1,21 @@
|
|||
DO $$
|
||||
BEGIN
|
||||
IF EXISTS (
|
||||
SELECT 1 FROM information_schema.columns
|
||||
WHERE table_name = 'LiteLLM_ShadowEvalJob' AND column_name = 'api_key_id'
|
||||
) THEN
|
||||
ALTER TABLE "LiteLLM_ShadowEvalJob" RENAME COLUMN "api_key_id" TO "target_id";
|
||||
END IF;
|
||||
END $$;
|
||||
|
||||
ALTER TABLE "LiteLLM_ShadowEvalJob" ADD COLUMN IF NOT EXISTS "target_type" TEXT NOT NULL DEFAULT 'key';
|
||||
|
||||
DROP INDEX IF EXISTS "LiteLLM_ShadowEvalJob_one_active_per_key_direction";
|
||||
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_ShadowEvalJob_one_active_per_target_direction"
|
||||
ON "LiteLLM_ShadowEvalJob"("target_type", "target_id", "direction") WHERE "stopped_at" IS NULL;
|
||||
|
||||
DROP INDEX IF EXISTS "LiteLLM_ShadowEvalJob_api_key_id_idx";
|
||||
|
||||
CREATE INDEX IF NOT EXISTS "LiteLLM_ShadowEvalJob_target_type_target_id_idx"
|
||||
ON "LiteLLM_ShadowEvalJob"("target_type", "target_id");
|
||||
|
|
@ -0,0 +1,3 @@
|
|||
ALTER TABLE "LiteLLM_ShadowEvalJob" ADD COLUMN IF NOT EXISTS "router_names" TEXT[] NOT NULL DEFAULT ARRAY[]::TEXT[];
|
||||
|
||||
ALTER TABLE "LiteLLM_ShadowEvalAttempt" ADD COLUMN IF NOT EXISTS "router_name" TEXT;
|
||||
|
|
@ -1529,14 +1529,16 @@ model LiteLLM_AutoRouterSession {
|
|||
model LiteLLM_ShadowEvalJob {
|
||||
id String @id @default(cuid())
|
||||
group_id String // legs of one job share this; the API's job id
|
||||
api_key_id String // hashed virtual key whose traffic this leg shadows
|
||||
router_name String // the auto-router under evaluation, in either direction
|
||||
target_type String @default("key") // key | team | user
|
||||
target_id String // hashed virtual key, team_id, or user_id whose traffic this leg shadows
|
||||
router_name String // first (often only) auto-router under evaluation; router_names is the full set
|
||||
router_names String[] @default([]) // all routers this job runs as shadow arms; empty on legacy rows, whose set is (router_name)
|
||||
direction String @default("forward") // forward | reverse
|
||||
baseline_model String? // reverse only: the fixed model the router is judged against
|
||||
judge_model String
|
||||
shadow_percentage Float
|
||||
max_turns Int // sample-count ceiling: the whole budget on pre-max_budget jobs, the error-loop valve otherwise
|
||||
max_budget Float? // per-key USD cap on the eval's own shadow + judge spend; null on jobs from before spend budgets
|
||||
max_budget Float? // per-target USD cap on the eval's own shadow + judge spend; null on jobs from before spend budgets
|
||||
created_at DateTime @default(now())
|
||||
created_by String?
|
||||
ends_at DateTime
|
||||
|
|
@ -1544,7 +1546,7 @@ model LiteLLM_ShadowEvalJob {
|
|||
stopped_by String? // operator who stopped it early; null when it ended on its own
|
||||
|
||||
@@index([group_id])
|
||||
@@index([api_key_id])
|
||||
@@index([target_type, target_id])
|
||||
@@index([created_at])
|
||||
}
|
||||
|
||||
|
|
@ -1554,6 +1556,7 @@ model LiteLLM_ShadowEvalAttempt {
|
|||
job_id String
|
||||
request_id String // the judged real request
|
||||
outcome String // real | shadow | tie | error
|
||||
router_name String? // the arm this verdict scores; NULL on legacy rows, meaning the job's own router
|
||||
tier String? // router's tier for the prompt, when classified
|
||||
real_model String?
|
||||
shadow_model String?
|
||||
|
|
|
|||
|
|
@ -7,7 +7,8 @@ import subprocess
|
|||
import tempfile
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Optional
|
||||
|
||||
from litellm_proxy_extras._logging import logger
|
||||
from litellm_proxy_extras.replica_identity import (
|
||||
|
|
@ -50,6 +51,17 @@ _SPEND_LOGS_PK_CLAUSE_RE = re.compile(
|
|||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
_PRISMA_ATTEMPTS: Final = 4
|
||||
|
||||
_TRANSIENT_PRISMA_FAILURES: Final = MappingProxyType(
|
||||
{
|
||||
"deadlock detected": "a deadlock on the migration advisory lock (a concurrent migrate deploy)",
|
||||
"P1001": "an unreachable database server",
|
||||
"P1002": "a database server that timed out",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
PARTITIONED_SPEND_LOGS_PUSH_ERROR = (
|
||||
"LiteLLM_SpendLogs is a partitioned table (see db_scripts/partition_spend_logs.sql), "
|
||||
"so its primary key must include the partition key (\"startTime\"). `prisma db push` "
|
||||
|
|
@ -274,6 +286,23 @@ class ProxyExtrasDBManager:
|
|||
env=prisma_env,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _transient_prisma_failure(stderr: str) -> str | None:
|
||||
"""Why a failed prisma command is worth retrying, or None.
|
||||
|
||||
v1 retried every failure, so it absorbed a database that was not up yet
|
||||
or another instance holding the migration lock. v2 fails fast, which is
|
||||
right for a broken migration and wrong for these.
|
||||
"""
|
||||
return next(
|
||||
(
|
||||
reason
|
||||
for marker, reason in _TRANSIENT_PRISMA_FAILURES.items()
|
||||
if marker in stderr
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _is_permission_error(error_message: str) -> bool:
|
||||
"""
|
||||
|
|
@ -512,6 +541,13 @@ class ProxyExtrasDBManager:
|
|||
try:
|
||||
import psycopg
|
||||
except ImportError:
|
||||
logger.warning(
|
||||
"psycopg is not installed; skipping the LiteLLM_SpendLogs "
|
||||
"partition check. If this table is partitioned (see "
|
||||
"db_scripts/partition_spend_logs.sql), schema reconciliation "
|
||||
"will try to rewrite its primary key and fail. Install the "
|
||||
"litellm[extra_proxy] extra, which now includes psycopg."
|
||||
)
|
||||
return False
|
||||
|
||||
cleaned_url = ProxyExtrasDBManager._strip_prisma_query_params(database_url)
|
||||
|
|
@ -648,7 +684,7 @@ class ProxyExtrasDBManager:
|
|||
@staticmethod
|
||||
def _setup_database_v2(use_migrate: bool) -> bool:
|
||||
"""
|
||||
v2 migration resolver (opt-in via --use_v2_migration_resolver).
|
||||
v2 migration resolver (what the proxy CLI selects by default).
|
||||
|
||||
Runs `prisma migrate deploy` and handles standard recovery paths
|
||||
(P3005 baseline, P3009/P3018 idempotent errors). Critically, it does
|
||||
|
|
@ -669,20 +705,46 @@ class ProxyExtrasDBManager:
|
|||
original_dir = os.getcwd()
|
||||
os.chdir(migrations_dir)
|
||||
try:
|
||||
subprocess.run(
|
||||
[_get_prisma_command(), "db", "push", "--accept-data-loss"],
|
||||
timeout=prisma_command_timeout(),
|
||||
check=True,
|
||||
env=_get_prisma_env(),
|
||||
for attempt in range(_PRISMA_ATTEMPTS):
|
||||
try:
|
||||
subprocess.run(
|
||||
[_get_prisma_command(), "db", "push", "--accept-data-loss"],
|
||||
timeout=prisma_command_timeout(),
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
env=_get_prisma_env(),
|
||||
)
|
||||
return True
|
||||
except subprocess.TimeoutExpired:
|
||||
logger.info(
|
||||
"prisma db push attempt %s timed out, retrying",
|
||||
attempt + 1,
|
||||
)
|
||||
time.sleep(random.randrange(5, 15))
|
||||
except subprocess.CalledProcessError as e:
|
||||
stderr = e.stderr or ""
|
||||
transient = ProxyExtrasDBManager._transient_prisma_failure(
|
||||
stderr
|
||||
)
|
||||
# Re-raise as RuntimeError so proxy_cli.py's
|
||||
# `except RuntimeError` catches it and exits cleanly.
|
||||
if transient is None or attempt == _PRISMA_ATTEMPTS - 1:
|
||||
raise RuntimeError(
|
||||
f"prisma db push failed.\n\nDetail: {e}"
|
||||
f"\n\nPrisma error:\n{stderr}"
|
||||
) from e
|
||||
logger.info(
|
||||
"prisma db push attempt %s failed on %s, retrying. "
|
||||
"Prisma error:\n%s",
|
||||
attempt + 1,
|
||||
transient,
|
||||
stderr,
|
||||
)
|
||||
time.sleep(random.randrange(5, 15))
|
||||
raise RuntimeError(
|
||||
f"prisma db push failed after {_PRISMA_ATTEMPTS} attempts."
|
||||
)
|
||||
return True
|
||||
except (
|
||||
subprocess.CalledProcessError,
|
||||
subprocess.TimeoutExpired,
|
||||
) as e:
|
||||
# Re-raise as RuntimeError so proxy_cli.py's
|
||||
# `except RuntimeError` catches it and exits cleanly.
|
||||
raise RuntimeError(f"prisma db push failed.\n\nDetail: {e}") from e
|
||||
finally:
|
||||
os.chdir(original_dir)
|
||||
|
||||
|
|
@ -692,7 +754,7 @@ class ProxyExtrasDBManager:
|
|||
original_dir = os.getcwd()
|
||||
os.chdir(migrations_dir)
|
||||
try:
|
||||
for attempt in range(4):
|
||||
for attempt in range(_PRISMA_ATTEMPTS):
|
||||
try:
|
||||
result = subprocess.run(
|
||||
[_get_prisma_command(), "migrate", "deploy"],
|
||||
|
|
@ -807,16 +869,36 @@ class ProxyExtrasDBManager:
|
|||
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
|
||||
) from e
|
||||
|
||||
raise RuntimeError(
|
||||
"Database migration failed and cannot be auto-recovered. "
|
||||
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
|
||||
) from e
|
||||
transient = ProxyExtrasDBManager._transient_prisma_failure(stderr)
|
||||
if transient is None:
|
||||
raise RuntimeError(
|
||||
"Database migration failed and cannot be auto-recovered. "
|
||||
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
|
||||
) from e
|
||||
|
||||
if attempt == _PRISMA_ATTEMPTS - 1:
|
||||
raise RuntimeError(
|
||||
f"Database migration failed after "
|
||||
f"{_PRISMA_ATTEMPTS} attempts on {transient}. "
|
||||
"Check database connectivity and load."
|
||||
f"\n\nPrisma error:\n{stderr}"
|
||||
) from e
|
||||
|
||||
logger.info(
|
||||
"prisma migrate deploy attempt %s failed on %s, retrying. "
|
||||
"Prisma error:\n%s",
|
||||
attempt + 1,
|
||||
transient,
|
||||
stderr,
|
||||
)
|
||||
time.sleep(random.randrange(5, 15))
|
||||
continue
|
||||
|
||||
raise RuntimeError(
|
||||
"Database migration failed after 4 attempts (retry loop "
|
||||
"exhausted by timeouts or repeated idempotent-recovery "
|
||||
"continues). Check database connectivity, load, and "
|
||||
"_prisma_migrations ledger state."
|
||||
f"Database migration failed after {_PRISMA_ATTEMPTS} "
|
||||
"attempts (retry loop exhausted by timeouts or repeated "
|
||||
"idempotent-recovery continues). Check database connectivity, "
|
||||
"load, and _prisma_migrations ledger state."
|
||||
)
|
||||
finally:
|
||||
os.chdir(original_dir)
|
||||
|
|
@ -864,10 +946,11 @@ class ProxyExtrasDBManager:
|
|||
|
||||
Args:
|
||||
use_migrate: Whether to use prisma migrate instead of db push
|
||||
use_v2_resolver: Opt into the v2 migration resolver (safer during
|
||||
use_v2_resolver: Run the v2 migration resolver (safer during
|
||||
rolling deploys; does not run the diff-and-force recovery
|
||||
that causes schema thrashing). Defaults to False for
|
||||
backwards compatibility.
|
||||
that causes schema thrashing). Defaults to False here so
|
||||
direct callers keep the old behavior; the proxy CLI passes
|
||||
True, so the proxy's runtime default is v2.
|
||||
|
||||
Returns:
|
||||
bool: True if setup was successful, False otherwise
|
||||
|
|
@ -885,7 +968,7 @@ class ProxyExtrasDBManager:
|
|||
@staticmethod
|
||||
def _run_migrations(use_migrate: bool, use_v2_resolver: bool) -> bool:
|
||||
if use_v2_resolver:
|
||||
logger.info("Using v2 migration resolver (--use_v2_migration_resolver)")
|
||||
logger.info("Using v2 migration resolver")
|
||||
return ProxyExtrasDBManager._setup_database_v2(use_migrate=use_migrate)
|
||||
|
||||
schema_path = ProxyExtrasDBManager._get_prisma_dir() + "/schema.prisma"
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.91"
|
||||
version = "0.4.92"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.4.91"
|
||||
version = "0.4.92"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -1,242 +0,0 @@
|
|||
"""Regression tests for ProxyExtrasDBManager v2 migration resolver.
|
||||
|
||||
The v2 resolver is opt-in via `--use_v2_migration_resolver` / the
|
||||
`use_v2_resolver=True` kwarg. These tests exercise the v2 path; the v1
|
||||
(default) behavior is unchanged from pre-fix.
|
||||
"""
|
||||
|
||||
import subprocess
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm_proxy_extras.utils import (
|
||||
ProxyExtrasDBManager,
|
||||
_max_migration_timestamp,
|
||||
_migration_timestamp,
|
||||
)
|
||||
|
||||
|
||||
def _fake_migrate_deploy_failure(returncode: int, stderr: str):
|
||||
def _run(*args, **kwargs):
|
||||
raise subprocess.CalledProcessError(
|
||||
returncode=returncode,
|
||||
cmd=args[0],
|
||||
stderr=stderr,
|
||||
output="",
|
||||
)
|
||||
|
||||
return _run
|
||||
|
||||
|
||||
def test_v2_p3018_permission_error_raises_runtime_error(monkeypatch, tmp_path):
|
||||
"""v2: a permission failure during migrate deploy raises RuntimeError."""
|
||||
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@localhost:9/x")
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager, "_warn_if_db_ahead_of_head", lambda _: None
|
||||
)
|
||||
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
|
||||
(tmp_path / "schema.prisma").write_text("// stub")
|
||||
|
||||
stderr = (
|
||||
"Error: P3018\nMigration name: 20250326162113_baseline\n"
|
||||
"Database error code: 42501\npermission denied for schema public"
|
||||
)
|
||||
with patch("subprocess.run", side_effect=_fake_migrate_deploy_failure(1, stderr)):
|
||||
with pytest.raises(RuntimeError, match="permission"):
|
||||
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
|
||||
|
||||
|
||||
def test_v2_non_idempotent_p3009_raises_runtime_error(monkeypatch, tmp_path):
|
||||
"""v2: a non-idempotent migration failure raises (no silent recovery)."""
|
||||
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@localhost:9/x")
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager, "_warn_if_db_ahead_of_head", lambda _: None
|
||||
)
|
||||
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
|
||||
(tmp_path / "schema.prisma").write_text("// stub")
|
||||
|
||||
stderr = (
|
||||
"Error: P3009\nMigration `20260101000000_genuinely_broken` failed\n"
|
||||
'Reason: syntax error at or near "BRKN" LINE 42'
|
||||
)
|
||||
with patch("subprocess.run", side_effect=_fake_migrate_deploy_failure(1, stderr)):
|
||||
with pytest.raises(RuntimeError, match="cannot be auto-recovered"):
|
||||
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
|
||||
|
||||
|
||||
def test_strip_prisma_query_params_removes_connection_limit():
|
||||
"""DATABASE_URLs with Prisma-specific params should be parseable by psycopg."""
|
||||
url = "postgresql://u:p@h:5432/db?connection_limit=100&pool_timeout=60&sslmode=require"
|
||||
stripped = ProxyExtrasDBManager._strip_prisma_query_params(url)
|
||||
assert "connection_limit" not in stripped
|
||||
assert "pool_timeout" not in stripped
|
||||
assert "sslmode=require" in stripped
|
||||
|
||||
|
||||
def test_strip_prisma_query_params_passthrough_no_query():
|
||||
"""URLs without query strings are returned unchanged."""
|
||||
url = "postgresql://u:p@h:5432/db"
|
||||
assert ProxyExtrasDBManager._strip_prisma_query_params(url) == url
|
||||
|
||||
|
||||
def test_migration_timestamp_extracts_leading_digits():
|
||||
assert _migration_timestamp("20260101000000_add_foo") == 20260101000000
|
||||
assert _migration_timestamp("20250326162113_baseline") == 20250326162113
|
||||
|
||||
|
||||
def test_migration_timestamp_returns_zero_on_malformed():
|
||||
assert _migration_timestamp("0_init") == 0
|
||||
assert _migration_timestamp("not_a_migration") == 0
|
||||
|
||||
|
||||
def test_max_migration_timestamp():
|
||||
names = {"20250326000000_a", "20260415000000_b", "20251115000000_c"}
|
||||
assert _max_migration_timestamp(names) == 20260415000000
|
||||
|
||||
|
||||
def test_max_migration_timestamp_empty_set():
|
||||
assert _max_migration_timestamp(set()) == 0
|
||||
|
||||
|
||||
def test_v1_default_still_calls_resolve_all_migrations(monkeypatch, tmp_path):
|
||||
"""v1 (default) continues to call _resolve_all_migrations on the happy path.
|
||||
|
||||
This is the existing buggy behavior — we're not fixing it in v1, only
|
||||
offering v2 as opt-in. This test pins the default so that a future
|
||||
inadvertent default flip is caught.
|
||||
"""
|
||||
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
|
||||
(tmp_path / "schema.prisma").write_text("// stub")
|
||||
|
||||
# Stub `prisma migrate deploy` to claim success with pending migrations
|
||||
# applied, which is the code path that triggers the legacy post-migration
|
||||
# sanity check (a call to _resolve_all_migrations).
|
||||
class FakeResult:
|
||||
stdout = "Applied migration.\n"
|
||||
stderr = ""
|
||||
|
||||
def fake_run(cmd, *args, **kwargs):
|
||||
return FakeResult()
|
||||
|
||||
resolve_called = {"n": 0}
|
||||
|
||||
def fake_resolve(*args, **kwargs):
|
||||
resolve_called["n"] += 1
|
||||
|
||||
monkeypatch.setattr("subprocess.run", fake_run)
|
||||
monkeypatch.setattr(ProxyExtrasDBManager, "_resolve_all_migrations", fake_resolve)
|
||||
|
||||
ok = ProxyExtrasDBManager.setup_database(use_migrate=True) # v2 flag NOT set
|
||||
assert ok is True
|
||||
assert resolve_called["n"] == 1, "v1 default should still invoke the legacy path"
|
||||
|
||||
|
||||
def test_v2_db_push_wraps_subprocess_error_as_runtime_error(monkeypatch, tmp_path):
|
||||
"""v2: a failing `prisma db push` must raise RuntimeError, not leak
|
||||
CalledProcessError past proxy_cli.py's `except RuntimeError`."""
|
||||
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
|
||||
(tmp_path / "schema.prisma").write_text("// stub")
|
||||
|
||||
stderr = "db push error"
|
||||
with patch("subprocess.run", side_effect=_fake_migrate_deploy_failure(1, stderr)):
|
||||
with pytest.raises(RuntimeError, match="prisma db push failed"):
|
||||
ProxyExtrasDBManager.setup_database(use_migrate=False, use_v2_resolver=True)
|
||||
|
||||
|
||||
def test_v2_warn_ahead_of_head_swallows_db_errors(monkeypatch, tmp_path):
|
||||
"""_warn_if_db_ahead_of_head must never raise — it's informational.
|
||||
|
||||
Non-connection DB errors (e.g. InsufficientPrivilege from a user
|
||||
without SELECT on _prisma_migrations) must be caught, not propagated.
|
||||
"""
|
||||
import psycopg
|
||||
|
||||
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@localhost:9/x")
|
||||
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
|
||||
(tmp_path / "schema.prisma").write_text("// stub")
|
||||
|
||||
class _FakeConn:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def execute(self, *a, **kw):
|
||||
# Simulate an InsufficientPrivilege (subclass of DatabaseError).
|
||||
raise psycopg.errors.InsufficientPrivilege("permission denied")
|
||||
|
||||
def _fake_connect(*a, **kw):
|
||||
return _FakeConn()
|
||||
|
||||
monkeypatch.setattr("psycopg.connect", _fake_connect)
|
||||
|
||||
# Must not raise.
|
||||
ProxyExtrasDBManager._warn_if_db_ahead_of_head(str(tmp_path))
|
||||
|
||||
|
||||
def test_v2_resolve_specific_migration_failure_raises_runtime_error(
|
||||
monkeypatch, tmp_path
|
||||
):
|
||||
"""If marking a migration as applied fails inside P3009 idempotent
|
||||
recovery, the subprocess error must be re-raised as RuntimeError so
|
||||
proxy_cli.py catches it cleanly (instead of leaking CalledProcessError)."""
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager, "_warn_if_db_ahead_of_head", lambda _: None
|
||||
)
|
||||
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
|
||||
(tmp_path / "schema.prisma").write_text("// stub")
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager, "_roll_back_migration", lambda *a, **kw: None
|
||||
)
|
||||
|
||||
# First call: migrate deploy -> P3009 idempotent error.
|
||||
# Recovery path tries _resolve_specific_migration; that also raises.
|
||||
def _failing_resolve(*a, **kw):
|
||||
raise subprocess.CalledProcessError(
|
||||
returncode=1,
|
||||
cmd="prisma migrate resolve --applied",
|
||||
stderr="resolve failed",
|
||||
output="",
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager, "_resolve_specific_migration", _failing_resolve
|
||||
)
|
||||
|
||||
stderr = (
|
||||
"Error: P3009\nMigration `20260101000000_some_migration` failed\n"
|
||||
"relation already exists"
|
||||
)
|
||||
with patch("subprocess.run", side_effect=_fake_migrate_deploy_failure(1, stderr)):
|
||||
with pytest.raises(
|
||||
RuntimeError, match="Failed to mark migration .* as applied"
|
||||
):
|
||||
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
|
||||
|
||||
|
||||
def test_v2_does_not_call_resolve_all_migrations(monkeypatch, tmp_path):
|
||||
"""v2 must never call _resolve_all_migrations — that's the bug it fixes."""
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager, "_warn_if_db_ahead_of_head", lambda _: None
|
||||
)
|
||||
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
|
||||
(tmp_path / "schema.prisma").write_text("// stub")
|
||||
|
||||
class FakeResult:
|
||||
stdout = "Applied migration.\n"
|
||||
stderr = ""
|
||||
|
||||
monkeypatch.setattr("subprocess.run", lambda *a, **kw: FakeResult())
|
||||
|
||||
resolve_called = {"n": 0}
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager,
|
||||
"_resolve_all_migrations",
|
||||
lambda *a, **kw: resolve_called.__setitem__("n", resolve_called["n"] + 1),
|
||||
)
|
||||
|
||||
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
|
||||
assert ok is True
|
||||
assert resolve_called["n"] == 0, "v2 must not invoke the diff-and-force recovery"
|
||||
|
|
@ -30,3 +30,12 @@ tokio = { version = "1", features = ["rt-multi-thread", "macros", "time", "net"]
|
|||
tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] }
|
||||
futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] }
|
||||
base64 = "0.22"
|
||||
|
||||
[profile.release]
|
||||
opt-level = 3
|
||||
lto = "thin"
|
||||
codegen-units = 1
|
||||
panic = "unwind"
|
||||
debug = false
|
||||
incremental = false
|
||||
strip = "symbols"
|
||||
|
|
|
|||
|
|
@ -10,7 +10,8 @@ name = "_native"
|
|||
crate-type = ["cdylib"]
|
||||
|
||||
[features]
|
||||
default = ["extension-module"]
|
||||
default = ["abi3"]
|
||||
abi3 = ["pyo3/abi3-py310"]
|
||||
extension-module = ["pyo3/extension-module"]
|
||||
|
||||
[dependencies]
|
||||
|
|
|
|||
|
|
@ -7,6 +7,9 @@ warnings.filterwarnings("ignore", message=".*conflict with protected namespace.*
|
|||
# Suppress Pydantic 2.11+ deprecation warning about accessing model_fields on instances
|
||||
# This warning can accumulate during streaming and cause memory leaks
|
||||
warnings.filterwarnings("ignore", message=".*Accessing the.*attribute on the instance is deprecated.*")
|
||||
# ReadOnly on TypedDict fields is repo-wide static discipline (LIT012); pydantic warns it
|
||||
# cannot enforce it at runtime, which floods proxy boot once such a type is schema-walked
|
||||
warnings.filterwarnings("ignore", message=".*`ReadOnly` qualifier.*")
|
||||
### INIT VARIABLES #########################
|
||||
import threading
|
||||
import os
|
||||
|
|
@ -656,6 +659,8 @@ aiml_models: Set = set()
|
|||
deepgram_models: Set = set()
|
||||
elevenlabs_models: Set = set()
|
||||
dashscope_models: Set = set()
|
||||
qwencloud_models: Set = set()
|
||||
qwen_ai_platform_models: Set = set()
|
||||
moonshot_models: Set = set()
|
||||
publicai_models: Set = set()
|
||||
darkbloom_models: Set = set()
|
||||
|
|
@ -906,6 +911,10 @@ def _populate_provider_model_sets(model_cost_map: Dict) -> None:
|
|||
heroku_models.add(key)
|
||||
elif value.get("litellm_provider") == "dashscope":
|
||||
dashscope_models.add(key)
|
||||
elif value.get("litellm_provider") == "qwencloud":
|
||||
qwencloud_models.add(key)
|
||||
elif value.get("litellm_provider") == "qwen_ai_platform":
|
||||
qwen_ai_platform_models.add(key)
|
||||
elif value.get("litellm_provider") == "modelscope":
|
||||
modelscope_models.add(key)
|
||||
elif value.get("litellm_provider") == "moonshot":
|
||||
|
|
@ -1069,6 +1078,8 @@ model_list = list(
|
|||
| deepgram_models
|
||||
| elevenlabs_models
|
||||
| dashscope_models
|
||||
| qwencloud_models
|
||||
| qwen_ai_platform_models
|
||||
| moonshot_models
|
||||
| publicai_models
|
||||
| darkbloom_models
|
||||
|
|
@ -1175,6 +1186,8 @@ def _build_models_by_provider() -> dict:
|
|||
"elevenlabs": elevenlabs_models,
|
||||
"heroku": heroku_models,
|
||||
"dashscope": dashscope_models,
|
||||
"qwencloud": qwencloud_models,
|
||||
"qwen_ai_platform": qwen_ai_platform_models,
|
||||
"modelscope": modelscope_models,
|
||||
"moonshot": moonshot_models,
|
||||
"publicai": publicai_models,
|
||||
|
|
@ -2011,6 +2024,24 @@ if TYPE_CHECKING:
|
|||
from .llms.dashscope.rerank.transformation import (
|
||||
DashScopeRerankConfig as DashScopeRerankConfig,
|
||||
)
|
||||
from .llms.dashscope.qwencloud import (
|
||||
QwenCloudChatConfig as QwenCloudChatConfig,
|
||||
)
|
||||
from .llms.dashscope.qwencloud import (
|
||||
QwenCloudEmbeddingConfig as QwenCloudEmbeddingConfig,
|
||||
)
|
||||
from .llms.dashscope.qwencloud import (
|
||||
QwenCloudRerankConfig as QwenCloudRerankConfig,
|
||||
)
|
||||
from .llms.dashscope.qwen_ai_platform import (
|
||||
QwenAIPlatformChatConfig as QwenAIPlatformChatConfig,
|
||||
)
|
||||
from .llms.dashscope.qwen_ai_platform import (
|
||||
QwenAIPlatformEmbeddingConfig as QwenAIPlatformEmbeddingConfig,
|
||||
)
|
||||
from .llms.dashscope.qwen_ai_platform import (
|
||||
QwenAIPlatformRerankConfig as QwenAIPlatformRerankConfig,
|
||||
)
|
||||
from .llms.modelscope.chat.transformation import (
|
||||
ModelScopeChatConfig as ModelScopeChatConfig,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -310,6 +310,8 @@ LLM_CONFIG_NAMES: Final = (
|
|||
"GigaChatConfig",
|
||||
"GigaChatEmbeddingConfig",
|
||||
"DashScopeChatConfig",
|
||||
"QwenCloudChatConfig",
|
||||
"QwenAIPlatformChatConfig",
|
||||
"ModelScopeChatConfig",
|
||||
"MoonshotChatConfig",
|
||||
"DockerModelRunnerChatConfig",
|
||||
|
|
@ -1172,6 +1174,14 @@ _LLM_CONFIGS_IMPORT_MAP: Final = {
|
|||
".llms.dashscope.chat.transformation",
|
||||
"DashScopeChatConfig",
|
||||
),
|
||||
"QwenCloudChatConfig": (
|
||||
".llms.dashscope.qwencloud",
|
||||
"QwenCloudChatConfig",
|
||||
),
|
||||
"QwenAIPlatformChatConfig": (
|
||||
".llms.dashscope.qwen_ai_platform",
|
||||
"QwenAIPlatformChatConfig",
|
||||
),
|
||||
"GDCGeminiConfig": (
|
||||
".llms.gdc.chat.transformation",
|
||||
"GDCGeminiConfig",
|
||||
|
|
|
|||
|
|
@ -264,13 +264,17 @@ def _plain_log_format(stdout: TextIO | None, stderr: TextIO | None) -> str:
|
|||
|
||||
|
||||
class LevelRoutingStreamHandler(logging.StreamHandler):
|
||||
"""Writes records below WARNING to stdout and WARNING and above to stderr.
|
||||
"""Writes records below WARNING and invalid-key warnings to stdout, others to stderr.
|
||||
|
||||
Collectors that derive severity from the stream report every stderr line as an error.
|
||||
Invalid-key warnings route to stdout so LITELLM_LOG=ERROR can suppress them.
|
||||
"""
|
||||
|
||||
def emit(self, record: logging.LogRecord) -> None:
|
||||
preferred: Final = sys.stdout if record.levelno < logging.WARNING else sys.stderr
|
||||
is_stdout_record: Final = record.levelno < logging.WARNING or (
|
||||
record.levelno == logging.WARNING and record.name == verbose_proxy_stdout_logger.name
|
||||
)
|
||||
preferred: Final = sys.stdout if is_stdout_record else sys.stderr
|
||||
if preferred is None or getattr(preferred, "closed", False):
|
||||
self.stream = sys.stderr # rebind-ok: fall back to the pre-fix stream rather than raising per record
|
||||
else:
|
||||
|
|
@ -508,6 +512,9 @@ else:
|
|||
handler.setFormatter(formatter)
|
||||
|
||||
verbose_proxy_logger = logging.getLogger("LiteLLM Proxy")
|
||||
# Malformed virtual key rejections log through this child; LevelRoutingStreamHandler
|
||||
# writes its WARNING records to stdout. It has no handler or level of its own.
|
||||
verbose_proxy_stdout_logger: Final = verbose_proxy_logger.getChild("stdout")
|
||||
verbose_router_logger = logging.getLogger("LiteLLM Router")
|
||||
verbose_logger = logging.getLogger("LiteLLM")
|
||||
|
||||
|
|
@ -520,6 +527,7 @@ verbose_logger.addHandler(handler)
|
|||
# handlers (JSON mode, uvicorn log config, a host app's root handler).
|
||||
verbose_router_logger.addFilter(_stdout_truncation_filter)
|
||||
verbose_proxy_logger.addFilter(_stdout_truncation_filter)
|
||||
verbose_proxy_stdout_logger.addFilter(_stdout_truncation_filter)
|
||||
verbose_logger.addFilter(_stdout_truncation_filter)
|
||||
|
||||
|
||||
|
|
@ -683,6 +691,7 @@ def _turn_on_json():
|
|||
- Adds a JSON formatter to all loggers
|
||||
"""
|
||||
handler: Final = LevelRoutingStreamHandler()
|
||||
handler.setLevel(numeric_level)
|
||||
handler.setFormatter(JsonFormatter())
|
||||
_initialize_loggers_with_handler(handler)
|
||||
# Set up exception handlers
|
||||
|
|
@ -700,12 +709,14 @@ def _disable_debugging():
|
|||
verbose_logger.disabled = True
|
||||
verbose_router_logger.disabled = True
|
||||
verbose_proxy_logger.disabled = True
|
||||
verbose_proxy_stdout_logger.disabled = True
|
||||
|
||||
|
||||
def _enable_debugging():
|
||||
verbose_logger.disabled = False
|
||||
verbose_router_logger.disabled = False
|
||||
verbose_proxy_logger.disabled = False
|
||||
verbose_proxy_stdout_logger.disabled = False
|
||||
|
||||
|
||||
def print_verbose(print_statement):
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ import json
|
|||
# s/o [@Frank Colson](https://www.linkedin.com/in/frank-colson-422b9b183/) for this redis implementation
|
||||
import os
|
||||
from collections.abc import Callable, Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from urllib.parse import urlsplit, urlunsplit
|
||||
|
||||
|
|
@ -38,9 +39,25 @@ from ._logging import verbose_logger
|
|||
AZURE_REDIS_SCOPE: Final = "https://redis.azure.com/.default"
|
||||
|
||||
|
||||
def _get_redis_kwargs():
|
||||
arg_spec: Final = inspect.getfullargspec(redis.Redis)
|
||||
def _unwrapped_init_args(cls: type) -> frozenset[str]:
|
||||
"""Every parameter on a single class's own ``__init__``, decorator-unwrapped.
|
||||
|
||||
Unlike ``_init_arg_names`` below, this does not walk the MRO: ``redis.Redis``
|
||||
and ``redis.RedisCluster`` (sync and async) each declare every real
|
||||
constructor parameter directly on their own ``__init__``, so MRO-walking is
|
||||
unnecessary — and it actively breaks the several tests here that mock the
|
||||
class with ``patch(..., autospec=True)``, since ``inspect.getmro`` needs a
|
||||
real ``__mro__`` that an autospec'd stand-in for a class does not provide.
|
||||
|
||||
Still unwraps first: redis-py >= 7.4 decorates these ``__init__``s with
|
||||
``@deprecated_args`` too, which the same class of bug as ``_init_arg_names``
|
||||
would otherwise silently empty this allowlist through (see its docstring).
|
||||
"""
|
||||
spec: Final = inspect.getfullargspec(inspect.unwrap(cls.__init__))
|
||||
return frozenset(spec.args + spec.kwonlyargs)
|
||||
|
||||
|
||||
def _get_redis_kwargs():
|
||||
# Only allow primitive arguments
|
||||
exclude_args: Final = {
|
||||
"self",
|
||||
|
|
@ -60,7 +77,7 @@ def _get_redis_kwargs():
|
|||
"azure_client_secret",
|
||||
}
|
||||
|
||||
available_args: Final = {x for x in arg_spec.args if x not in exclude_args} | include_args
|
||||
available_args: Final = {x for x in _unwrapped_init_args(redis.Redis) if x not in exclude_args} | include_args
|
||||
|
||||
return available_args
|
||||
|
||||
|
|
@ -120,15 +137,23 @@ def _get_redis_url_kwargs(client: type | None = None) -> tuple[str, ...]:
|
|||
return tuple(x for x in _init_arg_names(connection_cls) if x not in exclude_args) + include_args
|
||||
|
||||
|
||||
def _get_redis_cluster_kwargs(client=None):
|
||||
def _get_redis_cluster_kwargs(client: type | None = None):
|
||||
"""Config kwargs the target cluster client's constructor actually accepts.
|
||||
|
||||
Defaults to the sync ``redis.RedisCluster``, but the async cluster client
|
||||
(``redis.asyncio.cluster.RedisCluster``) declares connection settings such as
|
||||
``decode_responses`` on its own constructor, where the sync class takes them
|
||||
through ``**kwargs`` and so never names them in its signature. Introspecting
|
||||
only the sync class regardless of which client is actually built silently
|
||||
drops those for every async cluster caller.
|
||||
"""
|
||||
if client is None:
|
||||
client = redis.Redis.from_url
|
||||
arg_spec: Final = inspect.getfullargspec(redis.RedisCluster)
|
||||
client = redis.RedisCluster
|
||||
|
||||
# Only allow primitive arguments
|
||||
exclude_args: Final = {"self", "connection_pool", "retry", "host", "port", "startup_nodes"}
|
||||
|
||||
available_args = {x for x in arg_spec.args if x not in exclude_args}
|
||||
available_args = {x for x in _unwrapped_init_args(client) if x not in exclude_args}
|
||||
available_args |= {
|
||||
"password",
|
||||
"username",
|
||||
|
|
@ -161,6 +186,79 @@ def _get_redis_env_kwarg_mapping():
|
|||
return {f"{PREFIX}{x.upper()}": x for x in _get_redis_kwargs() if x not in exclude_from_environment}
|
||||
|
||||
|
||||
def _str_to_bool(value: str) -> bool:
|
||||
return value.lower() in ("true", "1", "yes")
|
||||
|
||||
|
||||
def _coerce_redis_kwargs_types(
|
||||
redis_kwargs: Mapping[str, object],
|
||||
client: type | tuple[type, ...] = redis.Redis,
|
||||
) -> dict[str, object]: # mutable-ok: a caller mutates the returned kwargs before constructing its client
|
||||
"""Coerces string values to the numeric/boolean type ``client``'s constructor
|
||||
declares for that parameter. ``client`` may be a tuple of client classes; a
|
||||
parameter's type is taken from the first signature that declares it, which
|
||||
lets cluster callers coerce cluster-only kwargs such as
|
||||
``cluster_error_retry_attempts`` alongside the shared connection kwargs.
|
||||
|
||||
Environment variables are always strings, and Helm ``--set`` stringifies values
|
||||
too, so a config value like ``health_check_interval`` or ``socket_timeout``
|
||||
can arrive as ``"30"``/``"5.5"`` rather than a real number. redis-py's own
|
||||
connection-health-check arithmetic (``loop.time() + self.health_check_interval``)
|
||||
then raises ``TypeError`` on every Redis operation instead of connecting.
|
||||
|
||||
``max_connections``, ``socket_timeout``, and ``socket_connect_timeout`` use an
|
||||
explicit target type rather than the parameter's own signature default: redis-py
|
||||
8.x changed the timeout defaults from ``None`` to int ``5``, so inferring the
|
||||
type from the default would make a fractional ``"5.5"`` fail ``int()`` and get
|
||||
silently dropped on 8.x while working on older versions. ``socket_keepalive``
|
||||
is explicit too: its signature default is ``None``, which carries no type to
|
||||
infer from, and leaving it a string makes ``"false"`` truthy.
|
||||
"""
|
||||
signatures: Final = tuple(inspect.signature(c) for c in (client if isinstance(client, tuple) else (client,)))
|
||||
explicit_param_types: Final = MappingProxyType(
|
||||
{
|
||||
"max_connections": int,
|
||||
"socket_timeout": float,
|
||||
"socket_connect_timeout": float,
|
||||
"socket_keepalive": bool,
|
||||
}
|
||||
)
|
||||
result: Final = dict(redis_kwargs) # mutable-ok: per-key try/except coercion below needs to drop individual keys
|
||||
for key, value in redis_kwargs.items():
|
||||
if not isinstance(value, str):
|
||||
continue
|
||||
param = next((sig.parameters[key] for sig in signatures if key in sig.parameters), None)
|
||||
if param is None:
|
||||
continue
|
||||
explicit_type = explicit_param_types.get(key)
|
||||
if explicit_type is bool:
|
||||
result[key] = _str_to_bool(value)
|
||||
continue
|
||||
if explicit_type is not None:
|
||||
try:
|
||||
result[key] = explicit_type(value)
|
||||
except (ValueError, TypeError):
|
||||
del result[key]
|
||||
continue
|
||||
default: object = param.default # pyright: ignore[reportAny] # inspect.Parameter.default is stubbed as Any
|
||||
if default is inspect.Parameter.empty:
|
||||
continue
|
||||
# bool must be checked before int, since bool subclasses int
|
||||
if isinstance(default, bool):
|
||||
result[key] = _str_to_bool(value)
|
||||
elif isinstance(default, int):
|
||||
try:
|
||||
result[key] = int(value)
|
||||
except (ValueError, TypeError):
|
||||
del result[key]
|
||||
elif isinstance(default, float):
|
||||
try:
|
||||
result[key] = float(value)
|
||||
except (ValueError, TypeError):
|
||||
del result[key]
|
||||
return result
|
||||
|
||||
|
||||
def _redis_kwargs_from_environment():
|
||||
mapping: Final = _get_redis_env_kwarg_mapping()
|
||||
|
||||
|
|
@ -505,7 +603,12 @@ def _get_redis_client_logic(**env_overrides):
|
|||
raise ValueError("Either 'host' or 'url' must be specified for redis.")
|
||||
|
||||
# litellm.print_verbose(f"redis_kwargs: {redis_kwargs}")
|
||||
return redis_kwargs
|
||||
coercion_client: Final = (
|
||||
(redis.Redis, redis.RedisCluster, async_redis.RedisCluster)
|
||||
if redis_kwargs.get("startup_nodes")
|
||||
else redis.Redis
|
||||
)
|
||||
return _coerce_redis_kwargs_types(redis_kwargs, client=coercion_client)
|
||||
|
||||
|
||||
def init_redis_cluster(redis_kwargs) -> redis.RedisCluster:
|
||||
|
|
@ -657,7 +760,9 @@ def get_redis_client(**env_overrides):
|
|||
if "sentinel_nodes" in redis_kwargs and "service_name" in redis_kwargs:
|
||||
return _init_redis_sentinel(redis_kwargs)
|
||||
|
||||
return redis.Redis(**redis_kwargs)
|
||||
return redis.Redis( # pyright: ignore[reportCallIssue] # object-valued kwargs match no overload statically
|
||||
**redis_kwargs, # pyright: ignore[reportArgumentType] # allow-listed and coerced against this signature
|
||||
)
|
||||
|
||||
|
||||
def get_redis_async_client(
|
||||
|
|
@ -669,7 +774,7 @@ def get_redis_async_client(
|
|||
if "startup_nodes" in redis_kwargs:
|
||||
from redis.cluster import ClusterNode
|
||||
|
||||
args = _get_redis_cluster_kwargs()
|
||||
args = _get_redis_cluster_kwargs(async_redis.RedisCluster)
|
||||
cluster_kwargs: Final = {}
|
||||
for arg in redis_kwargs:
|
||||
if arg in args:
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ import hashlib
|
|||
import json
|
||||
import time
|
||||
import traceback
|
||||
from collections.abc import Mapping
|
||||
from enum import Enum
|
||||
from typing import Any, Final
|
||||
|
||||
|
|
@ -506,7 +507,7 @@ class Cache:
|
|||
|
||||
def _get_cache_logic(
|
||||
self,
|
||||
cached_result: Any | None,
|
||||
cached_result: object | None,
|
||||
max_age: float | None,
|
||||
):
|
||||
"""
|
||||
|
|
@ -538,8 +539,8 @@ class Cache:
|
|||
return cached_result
|
||||
|
||||
@staticmethod
|
||||
def _get_safe_cache_lookup_kwargs(kwargs: dict[str, Any]) -> dict[str, Any]:
|
||||
cache_lookup_kwargs: Final[dict[str, Any]] = {}
|
||||
def _get_safe_cache_lookup_kwargs(kwargs: Mapping[str, object]) -> dict[str, object]:
|
||||
cache_lookup_kwargs: Final[dict[str, object]] = {}
|
||||
for prompt_kwarg in ("messages", "input"):
|
||||
if prompt_kwarg in kwargs:
|
||||
cache_lookup_kwargs[prompt_kwarg] = kwargs[prompt_kwarg]
|
||||
|
|
@ -552,7 +553,7 @@ class Cache:
|
|||
|
||||
@staticmethod
|
||||
def _update_metadata_from_cache_lookup_kwargs(
|
||||
original_kwargs: dict[str, Any], cache_lookup_kwargs: dict[str, Any]
|
||||
original_kwargs: Mapping[str, object], cache_lookup_kwargs: Mapping[str, object]
|
||||
) -> None:
|
||||
original_metadata: Final = original_kwargs.get("metadata")
|
||||
cache_lookup_metadata: Final = cache_lookup_kwargs.get("metadata")
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ import ast
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose
|
||||
|
|
@ -39,6 +39,12 @@ if TYPE_CHECKING:
|
|||
from litellm.router import Router
|
||||
|
||||
|
||||
class _QdrantCollectionDetailsResponse(Protocol):
|
||||
"""The qdrant `/collections/{name}` response, whose body is kept as an opaque JSON object."""
|
||||
|
||||
def json(self) -> dict[str, object]: ...
|
||||
|
||||
|
||||
class QdrantSemanticCache(BaseCache):
|
||||
CACHE_KEY_FIELD_NAME = "litellm_cache_key"
|
||||
embedding_max_input_tokens: int | None = None
|
||||
|
|
@ -115,15 +121,15 @@ class QdrantSemanticCache(BaseCache):
|
|||
raise ValueError(f"Error from qdrant checking if /collections exist {collection_exists.text}")
|
||||
|
||||
if collection_exists.json()["result"]["exists"]:
|
||||
collection_details = self.sync_client.get(
|
||||
collection_details: _QdrantCollectionDetailsResponse = self.sync_client.get(
|
||||
url=f"{self.qdrant_api_base}/collections/{self.collection_name}",
|
||||
headers=self.headers,
|
||||
)
|
||||
self.collection_info = collection_details.json()
|
||||
self.collection_info: dict[str, object] = collection_details.json()
|
||||
print_verbose(f"Collection already exists.\nCollection details:{self.collection_info}")
|
||||
self._ensure_cache_key_payload_index()
|
||||
else:
|
||||
quantization_params: dict[str, Any]
|
||||
quantization_params: dict[str, dict[str, object]]
|
||||
if quantization_config is None or quantization_config == "binary":
|
||||
quantization_params = {
|
||||
"binary": {
|
||||
|
|
@ -214,7 +220,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
resolve_embedding_max_input_tokens(self.embedding_max_input_tokens, self.embedding_model, router),
|
||||
)
|
||||
|
||||
def _get_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> EmbeddingResponse:
|
||||
def _get_embedding(self, prompt: str, metadata: dict[str, object] | None = None) -> EmbeddingResponse:
|
||||
"""Embed via the proxy Router when it serves the model, else direct."""
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_model_list, llm_router
|
||||
|
|
@ -241,7 +247,7 @@ class QdrantSemanticCache(BaseCache):
|
|||
num_retries=0,
|
||||
)
|
||||
|
||||
async def _get_async_embedding(self, prompt: str, metadata: dict[str, Any] | None = None) -> EmbeddingResponse:
|
||||
async def _get_async_embedding(self, prompt: str, metadata: dict[str, object] | None = None) -> EmbeddingResponse:
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_model_list, llm_router
|
||||
except ImportError:
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ import time
|
|||
from collections.abc import Awaitable, Callable, Sequence
|
||||
from contextvars import ContextVar
|
||||
from datetime import timedelta
|
||||
from typing import TYPE_CHECKING, Any, Final, TypeVar, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, TypeVar, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
|
|
@ -58,6 +58,26 @@ else:
|
|||
Span = Any
|
||||
|
||||
|
||||
class _AsyncRedisCommands(Protocol):
|
||||
"""Async redis commands this cache issues.
|
||||
|
||||
redis-py's type stubs omit these methods on RedisCluster, so the union returned by
|
||||
init_async_client() is untyped at every call site without this protocol.
|
||||
"""
|
||||
|
||||
def ping(self) -> Awaitable[bool]: ...
|
||||
|
||||
def delete(self, *names: str) -> Awaitable[int]: ...
|
||||
|
||||
def ttl(self, name: str) -> Awaitable[int]: ...
|
||||
|
||||
def rpush(self, name: str, *values: str | bytes | float) -> Awaitable[int]: ...
|
||||
|
||||
def lpop(self, name: str, count: int | None = None) -> Awaitable[object]: ...
|
||||
|
||||
def pipeline(self, transaction: bool = True) -> "Pipeline[bytes]": ...
|
||||
|
||||
|
||||
def _get_call_stack_info(num_frames: int = 2) -> str:
|
||||
"""
|
||||
Get the function names from the previous 1-2 functions in the call stack.
|
||||
|
|
@ -429,6 +449,9 @@ class RedisCache(BaseCache):
|
|||
self.redis_async_client = redis_async_client
|
||||
return redis_async_client
|
||||
|
||||
def _async_commands(self) -> _AsyncRedisCommands:
|
||||
return self.init_async_client()
|
||||
|
||||
def check_and_fix_namespace(self, key: str) -> str:
|
||||
"""
|
||||
Make sure each key starts with the given namespace
|
||||
|
|
@ -1055,19 +1078,17 @@ class RedisCache(BaseCache):
|
|||
await self.async_set_cache_pipeline(self.redis_batch_writing_buffer)
|
||||
self.redis_batch_writing_buffer = []
|
||||
|
||||
def _get_cache_logic(self, cached_response: Any):
|
||||
def _get_cache_logic(self, cached_response: bytes | str | None):
|
||||
"""
|
||||
Common 'get_cache_logic' across sync + async redis client implementations
|
||||
"""
|
||||
if cached_response is None:
|
||||
return cached_response
|
||||
# cached_response is in `b{} convert it to ModelResponse
|
||||
cached_response = cached_response.decode("utf-8") # Convert bytes to string
|
||||
return None
|
||||
decoded: Final = cached_response.decode("utf-8") if isinstance(cached_response, bytes) else cached_response
|
||||
try:
|
||||
cached_response = json.loads(cached_response) # Convert string to dictionary
|
||||
return json.loads(decoded)
|
||||
except Exception:
|
||||
cached_response = ast.literal_eval(cached_response)
|
||||
return cached_response
|
||||
return ast.literal_eval(decoded)
|
||||
|
||||
def get_cache(self, key, parent_otel_span: Span | None = None, **kwargs):
|
||||
try:
|
||||
|
|
@ -1314,8 +1335,7 @@ class RedisCache(BaseCache):
|
|||
raise e
|
||||
|
||||
async def ping(self) -> bool:
|
||||
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `ping`
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
_redis_client: Final = self._async_commands()
|
||||
start_time: Final = time.time()
|
||||
print_verbose("Pinging Async Redis Cache")
|
||||
try:
|
||||
|
|
@ -1349,8 +1369,7 @@ class RedisCache(BaseCache):
|
|||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def delete_cache_keys(self, keys):
|
||||
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `delete`
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
_redis_client: Final = self._async_commands()
|
||||
keys = [self.check_and_fix_namespace(key=key) for key in keys]
|
||||
# keys is a list, unpack it so it gets passed as individual elements to delete
|
||||
await _redis_client.delete(*keys)
|
||||
|
|
@ -1415,8 +1434,7 @@ class RedisCache(BaseCache):
|
|||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_delete_cache(self, key: str):
|
||||
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `delete`
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
_redis_client: Final = self._async_commands()
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
# keys is str
|
||||
return await _redis_client.delete(key)
|
||||
|
|
@ -1523,8 +1541,7 @@ class RedisCache(BaseCache):
|
|||
Redis ref: https://redis.io/docs/latest/commands/ttl/
|
||||
"""
|
||||
try:
|
||||
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `ttl`
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
_redis_client: Final = self._async_commands()
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
ttl: Final = await _redis_client.ttl(key)
|
||||
if ttl <= -1: # -1 means the key does not exist, -2 key does not exist
|
||||
|
|
@ -1554,7 +1571,7 @@ class RedisCache(BaseCache):
|
|||
Returns:
|
||||
int: The length of the list after the push operation
|
||||
"""
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
_redis_client: Final = self._async_commands()
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
start_time: Final = time.time()
|
||||
try:
|
||||
|
|
@ -1621,7 +1638,7 @@ class RedisCache(BaseCache):
|
|||
if len(rpush_list) == 0:
|
||||
return []
|
||||
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
_redis_client: Final = self._async_commands()
|
||||
start_time: Final = time.time()
|
||||
|
||||
try:
|
||||
|
|
@ -1678,7 +1695,7 @@ class RedisCache(BaseCache):
|
|||
parent_otel_span: Span | None = None,
|
||||
**kwargs,
|
||||
) -> Any | list[Any]:
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
_redis_client: Final = self._async_commands()
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
start_time: Final = time.time()
|
||||
print_verbose(f"LPOP from Redis list: key: {key}, count: {count}")
|
||||
|
|
@ -1810,7 +1827,7 @@ class RedisCache(BaseCache):
|
|||
if len(lpop_list) == 0:
|
||||
return []
|
||||
|
||||
_redis_client: Final[Any] = self.init_async_client()
|
||||
_redis_client: Final = self._async_commands()
|
||||
start_time: Final = time.time()
|
||||
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -45,14 +45,14 @@ class ResponsesToCompletionBridgeHandler:
|
|||
return bool(stream)
|
||||
|
||||
@staticmethod
|
||||
def _is_preformatted_cached_chat_stream(result: Any) -> bool:
|
||||
def _is_preformatted_cached_chat_stream(result: object) -> bool:
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
return isinstance(result, CustomStreamWrapper) and result.custom_llm_provider == "cached_response"
|
||||
|
||||
@staticmethod
|
||||
def _coerce_response_object(
|
||||
response_obj: Any,
|
||||
response_obj: object,
|
||||
hidden_params: dict | None,
|
||||
) -> "ResponsesAPIResponse":
|
||||
if isinstance(response_obj, ResponsesAPIResponse):
|
||||
|
|
@ -78,8 +78,8 @@ class ResponsesToCompletionBridgeHandler:
|
|||
for _ in stream_iter:
|
||||
pass
|
||||
|
||||
completed: Final = getattr(stream_iter, "completed_response", None)
|
||||
response_obj: Final = getattr(completed, "response", None) if completed else None
|
||||
completed: Final[object] = getattr(stream_iter, "completed_response", None)
|
||||
response_obj: Final[object] = getattr(completed, "response", None) if completed else None
|
||||
if response_obj is None:
|
||||
raise ValueError("Stream ended without a completed response")
|
||||
|
||||
|
|
@ -93,8 +93,8 @@ class ResponsesToCompletionBridgeHandler:
|
|||
async for _ in stream_iter:
|
||||
pass
|
||||
|
||||
completed: Final = getattr(stream_iter, "completed_response", None)
|
||||
response_obj: Final = getattr(completed, "response", None) if completed else None
|
||||
completed: Final[object] = getattr(stream_iter, "completed_response", None)
|
||||
response_obj: Final[object] = getattr(completed, "response", None) if completed else None
|
||||
if response_obj is None:
|
||||
raise ValueError("Stream ended without a completed response")
|
||||
|
||||
|
|
@ -157,7 +157,7 @@ class ResponsesToCompletionBridgeHandler:
|
|||
def completion(
|
||||
self, *args, **kwargs
|
||||
) -> Union[
|
||||
Coroutine[Any, Any, Union["ModelResponse", "CustomStreamWrapper"]],
|
||||
Coroutine[None, None, Union["ModelResponse", "CustomStreamWrapper"]],
|
||||
"ModelResponse",
|
||||
"CustomStreamWrapper",
|
||||
]:
|
||||
|
|
|
|||
|
|
@ -212,7 +212,8 @@ def _tool_call_dict_from_output_item(item: Mapping[str, Any], index: int) -> _Ch
|
|||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
|
||||
is_custom: Final = item.get("type") == "custom_tool_call"
|
||||
item_type: Final[object] = item.get("type")
|
||||
is_custom: Final = item_type == "custom_tool_call"
|
||||
arguments: Final = (item.get("input") if is_custom else item.get("arguments")) or ""
|
||||
name: Final = item.get("name") or ("custom_tool" if is_custom else "")
|
||||
function_chunk: Final = ChatCompletionToolCallFunctionChunk(name=name, arguments=arguments)
|
||||
|
|
@ -222,7 +223,7 @@ def _tool_call_dict_from_output_item(item: Mapping[str, Any], index: int) -> _Ch
|
|||
function=function_chunk,
|
||||
index=index,
|
||||
)
|
||||
raw_provider_fields: Final = item.get("provider_specific_fields")
|
||||
raw_provider_fields: Final[object] = item.get("provider_specific_fields")
|
||||
if isinstance(raw_provider_fields, dict):
|
||||
provider_specific_fields = raw_provider_fields
|
||||
elif raw_provider_fields and hasattr(raw_provider_fields, "__dict__"):
|
||||
|
|
@ -507,7 +508,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
def _merge_responses_api_request_into_request_data(
|
||||
self,
|
||||
request_data: dict[str, Any],
|
||||
request_data: dict[str, object],
|
||||
responses_api_request: "ResponsesAPIOptionalRequestParams",
|
||||
instructions: str | None,
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -484,6 +484,22 @@ FIREWORKS_AI_80_B: Final = int(os.getenv("FIREWORKS_AI_80_B", 80))
|
|||
#### Logging callback constants ####
|
||||
REDACTED_BY_LITELM_STRING: Final = "REDACTED_BY_LITELM"
|
||||
MAX_LANGFUSE_INITIALIZED_CLIENTS: Final = int(os.getenv("MAX_LANGFUSE_INITIALIZED_CLIENTS", 50))
|
||||
# Backpressure + lifetime bounds for the /v1/messages streaming relay (see
|
||||
# BaseAnthropicMessagesStreamingIterator.async_sse_wrapper). The relay queue is
|
||||
# bounded so a slow client throttles the upstream pump instead of letting it
|
||||
# buffer the whole response in memory; the detached-drain cap bounds how many
|
||||
# post-disconnect drains may run concurrently so client behavior can't create
|
||||
# unbounded worker state.
|
||||
ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE: Final = int(
|
||||
os.getenv("ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE", "1024")
|
||||
)
|
||||
# Setting this to 0 disables detached draining entirely: every post-disconnect
|
||||
# pump bills whatever partial output it has already collected and aborts the
|
||||
# upstream stream immediately, instead of continuing to drain for the real
|
||||
# terminal usage.
|
||||
ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS: Final = int(
|
||||
os.getenv("ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS", "100")
|
||||
)
|
||||
LOGGING_WORKER_CONCURRENCY: Final = int(os.getenv("LOGGING_WORKER_CONCURRENCY", 100)) # Must be above 0
|
||||
LOGGING_WORKER_MAX_QUEUE_SIZE: Final = int(os.getenv("LOGGING_WORKER_MAX_QUEUE_SIZE", 50_000))
|
||||
LOGGING_WORKER_MAX_TIME_PER_COROUTINE: Final = float(os.getenv("LOGGING_WORKER_MAX_TIME_PER_COROUTINE", 20.0))
|
||||
|
|
@ -614,6 +630,8 @@ LITELLM_CHAT_PROVIDERS: Final = [
|
|||
"nscale",
|
||||
"nebius",
|
||||
"dashscope",
|
||||
"qwencloud",
|
||||
"qwen_ai_platform",
|
||||
"modelscope",
|
||||
"moonshot",
|
||||
"publicai",
|
||||
|
|
@ -783,6 +801,7 @@ openai_compatible_endpoints: Final[list] = [
|
|||
"inference.api.nscale.com/v1",
|
||||
"api.studio.nebius.ai/v1",
|
||||
"https://dashscope-intl.aliyuncs.com/compatible-mode/v1",
|
||||
"https://dashscope.aliyuncs.com/compatible-mode/v1",
|
||||
"https://api-inference.modelscope.cn/v1",
|
||||
"https://api.moonshot.ai/v1",
|
||||
"https://api.publicai.co/v1",
|
||||
|
|
@ -806,6 +825,7 @@ openai_compatible_endpoints: Final[list] = [
|
|||
"https://api.meta.ai/v1",
|
||||
"https://api.cognition.ai/v1",
|
||||
"https://api.scx.ai/v1",
|
||||
"https://gigachat.devices.sberbank.ru/api/v1",
|
||||
]
|
||||
|
||||
|
||||
|
|
@ -855,6 +875,8 @@ openai_compatible_providers: Final[list] = [
|
|||
"nscale",
|
||||
"nebius",
|
||||
"dashscope",
|
||||
"qwencloud",
|
||||
"qwen_ai_platform",
|
||||
"modelscope",
|
||||
"moonshot",
|
||||
"v0",
|
||||
|
|
@ -885,6 +907,8 @@ openai_text_completion_compatible_providers: Final[list] = [ # providers that s
|
|||
"featherless_ai",
|
||||
"nebius",
|
||||
"dashscope",
|
||||
"qwencloud",
|
||||
"qwen_ai_platform",
|
||||
"modelscope",
|
||||
"moonshot",
|
||||
"publicai",
|
||||
|
|
@ -1092,7 +1116,7 @@ nebius_models: Final[set] = set(
|
|||
]
|
||||
)
|
||||
|
||||
dashscope_models: Final[set] = set(
|
||||
dashscope_models: Final[frozenset] = frozenset(
|
||||
[
|
||||
"qwen-turbo",
|
||||
"qwen-plus",
|
||||
|
|
@ -1107,6 +1131,10 @@ dashscope_models: Final[set] = set(
|
|||
]
|
||||
)
|
||||
|
||||
qwencloud_models: Final[frozenset] = frozenset(dashscope_models)
|
||||
|
||||
qwen_ai_platform_models: Final[frozenset] = frozenset(dashscope_models)
|
||||
|
||||
nebius_embedding_models: Final[set] = set(
|
||||
[
|
||||
"BAAI/bge-en-icl",
|
||||
|
|
@ -1410,6 +1438,12 @@ DEFAULT_SOFT_BUDGET: Final = float(
|
|||
) # by default all litellm proxy keys have a soft budget of 50.0
|
||||
# makes it clear this is a rate limit error for a litellm virtual key
|
||||
RATE_LIMIT_ERROR_MESSAGE_FOR_VIRTUAL_KEY: Final = "LiteLLM Virtual Key user_api_key_hash"
|
||||
# Prefix of the 401 raised when a submitted virtual key is not shaped like one.
|
||||
INVALID_VIRTUAL_KEY_ERROR_MESSAGE: Final = "LiteLLM Virtual Key expected"
|
||||
# Attribute stamped on that 401 at its raise site so log routing recognises it by
|
||||
# provenance. Message text is caller-influenceable on other 401s, so it must not
|
||||
# be used to classify.
|
||||
INVALID_VIRTUAL_KEY_ERROR_MARKER: Final = "_litellm_invalid_virtual_key_error"
|
||||
|
||||
# Python garbage collection threshold configuration
|
||||
# Format: "gen0,gen1,gen2" e.g., "1000,50,50"
|
||||
|
|
|
|||
|
|
@ -641,12 +641,12 @@ def cost_per_token(
|
|||
return xai_cost_per_token(model=model, usage=usage_block)
|
||||
elif custom_llm_provider == "lemonade":
|
||||
return lemonade_cost_per_token(model=model, usage=usage_block)
|
||||
elif custom_llm_provider == "dashscope":
|
||||
elif custom_llm_provider in ("dashscope", "qwencloud", "qwen_ai_platform"):
|
||||
from litellm.llms.dashscope.cost_calculator import (
|
||||
cost_per_token as dashscope_cost_per_token,
|
||||
)
|
||||
|
||||
return dashscope_cost_per_token(model=model, usage=usage_block)
|
||||
return dashscope_cost_per_token(model=model, usage=usage_block, custom_llm_provider=custom_llm_provider)
|
||||
elif custom_llm_provider == "azure_ai":
|
||||
return azure_ai_cost_per_token(
|
||||
model=model,
|
||||
|
|
@ -1910,12 +1910,15 @@ def ocr_cost(
|
|||
if credits is not None and cost_per_credit is not None:
|
||||
return cost_per_credit * credits, 0.0
|
||||
|
||||
ocr_cost_per_page: float | None = None
|
||||
if model_info is not None:
|
||||
ocr_cost_per_page = model_info.get("ocr_cost_per_page")
|
||||
ocr_cost_per_page: Final = model_info.get("ocr_cost_per_page") if model_info is not None else None
|
||||
annotation_cost_per_page: Final = model_info.get("annotation_cost_per_page") if model_info is not None else None
|
||||
annotation_rate: Final = annotation_cost_per_page if annotation_cost_per_page is not None else ocr_cost_per_page
|
||||
|
||||
pages_processed: Final = response.usage_info.pages_processed
|
||||
if pages_processed is None:
|
||||
annotation_pages: Final = response.usage_info.pages_processed_annotation or 0
|
||||
has_billable_annotation_pages: Final = annotation_rate is not None and annotation_pages > 0
|
||||
|
||||
if pages_processed is None and not has_billable_annotation_pages:
|
||||
if cost_per_credit is not None or ocr_cost_per_page is None:
|
||||
# Surface missing usage data instead of silently under-reporting
|
||||
# cost. The previous behavior raised ValueError; we now return 0.0
|
||||
|
|
@ -1931,7 +1934,7 @@ def ocr_cost(
|
|||
return 0.0, 0.0
|
||||
raise ValueError("OCR response pages_processed is None")
|
||||
|
||||
if ocr_cost_per_page is None:
|
||||
if ocr_cost_per_page is None and not has_billable_annotation_pages:
|
||||
# No per-page pricing configured. Either the model is on credit-based
|
||||
# pricing (and credits weren't returned, so the credit branch above did
|
||||
# not match) or the model has no OCR pricing entry at all. Surface a
|
||||
|
|
@ -1947,8 +1950,9 @@ def ocr_cost(
|
|||
)
|
||||
return 0.0, 0.0
|
||||
|
||||
total_ocr_processing_cost: Final[float] = ocr_cost_per_page * pages_processed
|
||||
return total_ocr_processing_cost, 0.0
|
||||
ocr_pages_cost: Final = (ocr_cost_per_page or 0.0) * (pages_processed or 0)
|
||||
annotation_pages_cost: Final = (annotation_rate or 0.0) * annotation_pages
|
||||
return ocr_pages_cost + annotation_pages_cost, 0.0
|
||||
|
||||
|
||||
def vector_store_search_cost(
|
||||
|
|
@ -2268,6 +2272,10 @@ def batch_cost_calculator(
|
|||
return total_prompt_cost, total_completion_cost
|
||||
|
||||
|
||||
def _attribute_value(obj: object, name: str) -> object:
|
||||
return getattr(obj, name)
|
||||
|
||||
|
||||
def _summable_prompt_token_fields(prompt_tokens_details: BaseModel) -> list[str]:
|
||||
field_names: Final = list(type(prompt_tokens_details).model_fields)
|
||||
if getattr(prompt_tokens_details, "cache_write_tokens", None) is None:
|
||||
|
|
@ -2293,7 +2301,7 @@ class BaseTokenUsageProcessor:
|
|||
for usage in usage_objects:
|
||||
# Handle direct attributes by checking what exists in the model
|
||||
for attr in dir(usage):
|
||||
if not attr.startswith("_") and not callable(getattr(usage, attr)):
|
||||
if not attr.startswith("_") and not callable(_attribute_value(usage, attr)):
|
||||
current_val = getattr(combined, attr, 0)
|
||||
new_val = getattr(usage, attr, 0)
|
||||
if (
|
||||
|
|
@ -2313,7 +2321,7 @@ class BaseTokenUsageProcessor:
|
|||
if (
|
||||
hasattr(usage.prompt_tokens_details, attr)
|
||||
and not attr.startswith("_")
|
||||
and not callable(getattr(usage.prompt_tokens_details, attr))
|
||||
and not callable(_attribute_value(usage.prompt_tokens_details, attr))
|
||||
):
|
||||
current_val = getattr(combined.prompt_tokens_details, attr, 0) or 0
|
||||
new_val = getattr(usage.prompt_tokens_details, attr, 0) or 0
|
||||
|
|
@ -2332,7 +2340,9 @@ class BaseTokenUsageProcessor:
|
|||
# Check what keys exist in the model's completion_tokens_details
|
||||
# Access model_fields on the class, not the instance, to avoid Pydantic 2.11+ deprecation warnings
|
||||
for attr in type(usage.completion_tokens_details).model_fields:
|
||||
if not attr.startswith("_") and not callable(getattr(usage.completion_tokens_details, attr)):
|
||||
if not attr.startswith("_") and not callable(
|
||||
_attribute_value(usage.completion_tokens_details, attr)
|
||||
):
|
||||
current_val = getattr(combined.completion_tokens_details, attr, 0) or 0
|
||||
new_val = getattr(usage.completion_tokens_details, attr, 0) or 0
|
||||
if isinstance(new_val, (int, float)):
|
||||
|
|
|
|||
|
|
@ -115,9 +115,11 @@ class SpeechToCompletionBridgeHandler:
|
|||
**request_data,
|
||||
)
|
||||
|
||||
requested_response_format: Final = optional_params.get("response_format")
|
||||
if isinstance(result, ModelResponse):
|
||||
return self.transformation_handler.transform_response(
|
||||
model_response=result,
|
||||
response_format=requested_response_format if isinstance(requested_response_format, str) else None,
|
||||
)
|
||||
else:
|
||||
raise Exception(f"Unmapped response type. Got type: {type(result)}")
|
||||
|
|
|
|||
|
|
@ -1,10 +1,14 @@
|
|||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, cast
|
||||
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
from litellm.constants import OPENAI_CHAT_COMPLETION_PARAMS
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
from litellm.types.llms.openai import ChatCompletionUserMessage, HttpxBinaryResponseContent
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
|
|
@ -16,7 +20,64 @@ def _completion_response_cost(model_response: "ModelResponse") -> float | None:
|
|||
return response_cost if isinstance(response_cost, float) else None
|
||||
|
||||
|
||||
GEMINI_TTS_CHAT_AUDIO_FORMAT: Final = "pcm16"
|
||||
GEMINI_TTS_RAW_RESPONSE_FORMAT: Final = "pcm"
|
||||
GEMINI_TTS_SUPPORTED_RESPONSE_FORMATS: Final = frozenset({"wav", GEMINI_TTS_RAW_RESPONSE_FORMAT})
|
||||
|
||||
|
||||
class ChatAudioParam(TypedDict):
|
||||
voice: ReadOnly[str]
|
||||
format: ReadOnly[NotRequired[str]]
|
||||
|
||||
|
||||
class SpeechToCompletionBridgeTransformationHandler:
|
||||
def _validate_response_format(
|
||||
self, model: str, custom_llm_provider: str, optional_params: Mapping[str, object]
|
||||
) -> None:
|
||||
if not self._is_gemini_tts_model(model):
|
||||
return
|
||||
response_format: Final = optional_params.get("response_format")
|
||||
if not isinstance(response_format, str) or response_format in GEMINI_TTS_SUPPORTED_RESPONSE_FORMATS:
|
||||
return
|
||||
from litellm.exceptions import BadRequestError
|
||||
|
||||
supported: Final = ", ".join(sorted(GEMINI_TTS_SUPPORTED_RESPONSE_FORMATS))
|
||||
raise BadRequestError(
|
||||
message=(
|
||||
f"Gemini TTS only produces raw PCM16 audio, so response_format='{response_format}'"
|
||||
f" is not supported. Supported response formats: {supported}."
|
||||
),
|
||||
model=model,
|
||||
llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
def _chat_completion_params(self, optional_params: Mapping[str, object]) -> Mapping[str, object]:
|
||||
return MappingProxyType(
|
||||
{
|
||||
param: value
|
||||
for param, value in optional_params.items()
|
||||
if param in OPENAI_CHAT_COMPLETION_PARAMS and param != "response_format"
|
||||
}
|
||||
)
|
||||
|
||||
def _chat_audio_format(self, model: str, optional_params: Mapping[str, object]) -> str | None:
|
||||
if self._is_gemini_tts_model(model):
|
||||
return GEMINI_TTS_CHAT_AUDIO_FORMAT
|
||||
response_format: Final = optional_params.get("response_format")
|
||||
return response_format if isinstance(response_format, str) else None
|
||||
|
||||
def _chat_audio_param(
|
||||
self, model: str, voice: str | Mapping[str, object] | None, optional_params: Mapping[str, object]
|
||||
) -> ChatAudioParam | None:
|
||||
if not isinstance(voice, str):
|
||||
return None
|
||||
audio_format: Final = self._chat_audio_format(model, optional_params)
|
||||
if audio_format is None:
|
||||
voice_only: Final[ChatAudioParam] = {"voice": voice}
|
||||
return voice_only
|
||||
audio: Final[ChatAudioParam] = {"voice": voice, "format": audio_format}
|
||||
return audio
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -28,36 +89,20 @@ class SpeechToCompletionBridgeTransformationHandler:
|
|||
litellm_logging_obj: "LiteLLMLoggingObj",
|
||||
custom_llm_provider: str,
|
||||
) -> dict:
|
||||
passed_optional_params: Final = {}
|
||||
for op in optional_params:
|
||||
if op in OPENAI_CHAT_COMPLETION_PARAMS:
|
||||
passed_optional_params[op] = optional_params[op]
|
||||
|
||||
if voice is not None:
|
||||
if isinstance(voice, str):
|
||||
passed_optional_params["audio"] = {"voice": voice}
|
||||
if "response_format" in optional_params:
|
||||
passed_optional_params["audio"]["format"] = optional_params["response_format"]
|
||||
|
||||
return_kwargs = {
|
||||
self._validate_response_format(model, custom_llm_provider, optional_params)
|
||||
user_message: Final[ChatCompletionUserMessage] = {"role": "user", "content": input}
|
||||
return_kwargs: Final = {
|
||||
"model": model,
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": input,
|
||||
}
|
||||
],
|
||||
"messages": [user_message],
|
||||
"modalities": ["audio"],
|
||||
**passed_optional_params,
|
||||
**self._chat_completion_params(optional_params),
|
||||
"audio": self._chat_audio_param(model, voice, optional_params),
|
||||
**litellm_params,
|
||||
"headers": headers,
|
||||
"litellm_logging_obj": litellm_logging_obj,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
}
|
||||
|
||||
# filter out None values
|
||||
return_kwargs = {k: v for k, v in return_kwargs.items() if v is not None}
|
||||
return return_kwargs
|
||||
return {k: v for k, v in return_kwargs.items() if v is not None}
|
||||
|
||||
def _convert_pcm16_to_wav(self, pcm_data: bytes, sample_rate: int = 24000, channels: int = 1) -> bytes:
|
||||
"""
|
||||
|
|
@ -103,7 +148,14 @@ class SpeechToCompletionBridgeTransformationHandler:
|
|||
"""Check if the model is a Gemini TTS model that returns PCM16 data."""
|
||||
return "gemini" in model.lower() and ("tts" in model.lower() or "preview-tts" in model.lower())
|
||||
|
||||
def transform_response(self, model_response: "ModelResponse") -> "HttpxBinaryResponseContent":
|
||||
def _gemini_tts_response_body(self, decoded_audio: bytes, response_format: str | None) -> tuple[bytes, str]:
|
||||
if response_format == GEMINI_TTS_RAW_RESPONSE_FORMAT:
|
||||
return decoded_audio, "audio/pcm"
|
||||
return self._convert_pcm16_to_wav(decoded_audio), "audio/wav"
|
||||
|
||||
def transform_response(
|
||||
self, model_response: "ModelResponse", response_format: str | None
|
||||
) -> "HttpxBinaryResponseContent":
|
||||
import base64
|
||||
|
||||
import httpx
|
||||
|
|
@ -114,23 +166,17 @@ class SpeechToCompletionBridgeTransformationHandler:
|
|||
audio_part: Final = cast(Choices, model_response.choices[0]).message.audio
|
||||
if audio_part is None:
|
||||
raise ValueError("No audio part found in the response")
|
||||
audio_content: Final = audio_part.data
|
||||
decoded_audio: Final = base64.b64decode(audio_part.data)
|
||||
|
||||
# Decode base64 to get binary content
|
||||
binary_data = base64.b64decode(audio_content)
|
||||
|
||||
# Check if this is a Gemini TTS model that returns raw PCM16 data
|
||||
model: Final = getattr(model_response, "model", "")
|
||||
headers: Final = {}
|
||||
if self._is_gemini_tts_model(model):
|
||||
# Convert PCM16 to WAV format for proper audio file playback
|
||||
binary_data = self._convert_pcm16_to_wav(binary_data)
|
||||
headers["Content-Type"] = "audio/wav"
|
||||
else:
|
||||
headers["Content-Type"] = "audio/mpeg"
|
||||
|
||||
# Create an httpx.Response object
|
||||
response: Final = httpx.Response(status_code=200, content=binary_data, headers=headers)
|
||||
content, content_type = (
|
||||
self._gemini_tts_response_body(decoded_audio, response_format)
|
||||
if self._is_gemini_tts_model(model)
|
||||
else (decoded_audio, "audio/mpeg")
|
||||
)
|
||||
response: Final = httpx.Response(
|
||||
status_code=200, content=content, headers=MappingProxyType({"Content-Type": content_type})
|
||||
)
|
||||
binary_response: Final = HttpxBinaryResponseContent(response)
|
||||
binary_response.set_response_cost(_completion_response_cost(model_response))
|
||||
return binary_response
|
||||
|
|
|
|||
|
|
@ -108,7 +108,6 @@ class GoogleGenAIStreamWrapper(AdapterCompletionStreamWrapper):
|
|||
|
||||
def __init__(self, completion_stream: object):
|
||||
self.sent_first_chunk = False
|
||||
# State tracking for accumulating partial tool calls
|
||||
self.accumulated_tool_calls = dict[int, _ToolCallAccumulator]()
|
||||
self._returned_response = False
|
||||
super().__init__(completion_stream)
|
||||
|
|
@ -723,7 +722,7 @@ class GoogleGenAIAdapter:
|
|||
)
|
||||
|
||||
for tool_call in tool_calls:
|
||||
if not hasattr(tool_call, "function"):
|
||||
if not hasattr(tool_call, "function") or isinstance(tool_call, ChatCompletionDeltaCustomToolCall):
|
||||
continue
|
||||
|
||||
# 3. Use `index` as the primary key for accumulation
|
||||
|
|
|
|||
|
|
@ -52,10 +52,10 @@ class GenerateContentSetupResult(BaseModel):
|
|||
model_config: ClassVar[ConfigDict] = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
model: str
|
||||
request_body: dict[str, Any]
|
||||
request_body: dict[str, object]
|
||||
custom_llm_provider: str
|
||||
generate_content_provider_config: BaseGoogleGenAIGenerateContentConfig | None
|
||||
generate_content_config_dict: dict[str, Any]
|
||||
generate_content_config_dict: dict[str, object]
|
||||
native_request_fields: dict[str, object]
|
||||
litellm_params: GenericLiteLLMParams
|
||||
litellm_logging_obj: LiteLLMLoggingObj
|
||||
|
|
@ -68,7 +68,7 @@ class GenerateContentHelper:
|
|||
@staticmethod
|
||||
def mock_generate_content_response(
|
||||
mock_response: str = "This is a mock response from Google GenAI generate_content.",
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""Mock response for generate_content for testing purposes"""
|
||||
return {
|
||||
"text": mock_response,
|
||||
|
|
@ -239,9 +239,9 @@ async def agenerate_content(
|
|||
tools: ToolConfigDict | None = None,
|
||||
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
|
||||
# The extra values given here take precedence over values defined on the client or passed to this method.
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
# LiteLLM specific params,
|
||||
custom_llm_provider: str | None = None,
|
||||
|
|
@ -307,9 +307,9 @@ def generate_content(
|
|||
tools: ToolConfigDict | None = None,
|
||||
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
|
||||
# The extra values given here take precedence over values defined on the client or passed to this method.
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
# LiteLLM specific params,
|
||||
custom_llm_provider: str | None = None,
|
||||
|
|
@ -397,9 +397,9 @@ async def agenerate_content_stream(
|
|||
tools: ToolConfigDict | None = None,
|
||||
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
|
||||
# The extra values given here take precedence over values defined on the client or passed to this method.
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
# LiteLLM specific params,
|
||||
custom_llm_provider: str | None = None,
|
||||
|
|
@ -492,9 +492,9 @@ def generate_content_stream(
|
|||
tools: ToolConfigDict | None = None,
|
||||
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
|
||||
# The extra values given here take precedence over values defined on the client or passed to this method.
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
# LiteLLM specific params,
|
||||
custom_llm_provider: str | None = None,
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import contextvars
|
|||
import importlib
|
||||
from collections.abc import Coroutine
|
||||
from functools import partial
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast, overload
|
||||
from typing import TYPE_CHECKING, Final, Literal, Optional, cast, overload
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.images.utils import ImageEditRequestUtils
|
||||
|
|
@ -151,7 +151,7 @@ def image_generation(
|
|||
*,
|
||||
aimg_generation: Literal[True],
|
||||
**kwargs,
|
||||
) -> Coroutine[Any, Any, ImageResponse]:
|
||||
) -> Coroutine[object, object, ImageResponse]:
|
||||
...
|
||||
|
||||
|
||||
|
|
@ -197,7 +197,7 @@ def image_generation(
|
|||
api_version: str | None = None,
|
||||
custom_llm_provider=None,
|
||||
**kwargs,
|
||||
) -> ImageResponse | Coroutine[Any, Any, ImageResponse]:
|
||||
) -> ImageResponse | Coroutine[object, object, ImageResponse]:
|
||||
"""
|
||||
Maps the https://api.openai.com/v1/images/generations endpoint.
|
||||
|
||||
|
|
@ -386,6 +386,8 @@ def image_generation(
|
|||
litellm.LlmProviders.VERTEX_AI,
|
||||
litellm.LlmProviders.OPENROUTER,
|
||||
litellm.LlmProviders.DASHSCOPE,
|
||||
litellm.LlmProviders.QWENCLOUD,
|
||||
litellm.LlmProviders.QWEN_AI_PLATFORM,
|
||||
):
|
||||
if image_generation_config is None:
|
||||
raise ValueError(f"image generation config is not supported for {custom_llm_provider}")
|
||||
|
|
@ -723,14 +725,14 @@ def image_edit(
|
|||
user: str | None = None,
|
||||
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
|
||||
# The extra values given here take precedence over values defined on the client or passed to this method.
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
# LiteLLM specific params,
|
||||
custom_llm_provider: str | None = None,
|
||||
**kwargs,
|
||||
) -> ImageResponse | Coroutine[Any, Any, ImageResponse]:
|
||||
) -> ImageResponse | Coroutine[object, object, ImageResponse]:
|
||||
"""
|
||||
Maps the image edit functionality, similar to OpenAI's images/edits endpoint.
|
||||
"""
|
||||
|
|
@ -769,7 +771,7 @@ def image_edit(
|
|||
images: Final = image if isinstance(image, list) else ([image] if image is not None else [])
|
||||
|
||||
headers_from_kwargs: Final = kwargs.get("headers")
|
||||
merged_extra_headers: Final[dict[str, Any]] = {}
|
||||
merged_extra_headers: Final[dict[str, object]] = {}
|
||||
if isinstance(headers_from_kwargs, dict):
|
||||
merged_extra_headers.update(headers_from_kwargs)
|
||||
if isinstance(extra_headers, dict):
|
||||
|
|
@ -974,9 +976,9 @@ async def aimage_edit(
|
|||
user: str | None = None,
|
||||
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
|
||||
# The extra values given here take precedence over values defined on the client or passed to this method.
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
# LiteLLM specific params,
|
||||
custom_llm_provider: str | None = None,
|
||||
|
|
@ -1044,7 +1046,7 @@ async def aimage_edit(
|
|||
)
|
||||
|
||||
|
||||
def __getattr__(name: str) -> Any:
|
||||
def __getattr__(name: str) -> type["ImageEditRequestUtils"]:
|
||||
"""Lazy import handler for images.main module"""
|
||||
if name == "ImageEditRequestUtils":
|
||||
# Lazy load ImageEditRequestUtils to avoid heavy import from images.utils at module load time
|
||||
|
|
|
|||
|
|
@ -545,7 +545,6 @@ class SlackAlerting(CustomBatchLogger):
|
|||
# Get the appropriate budget alert type handler
|
||||
budget_alert_class: Final = get_budget_alert_type(type)
|
||||
_id: Final = budget_alert_class.get_id(user_info)
|
||||
user_info_json: Final = user_info.model_dump(exclude_none=True)
|
||||
user_info_str: Final = self._get_user_info_str(user_info)
|
||||
event_message = budget_alert_class.get_event_message()
|
||||
|
||||
|
|
@ -575,7 +574,22 @@ class SlackAlerting(CustomBatchLogger):
|
|||
webhook_event = WebhookEvent(
|
||||
event=event,
|
||||
event_message=event_message,
|
||||
**user_info_json,
|
||||
spend=user_info.spend,
|
||||
max_budget=user_info.max_budget,
|
||||
soft_budget=user_info.soft_budget,
|
||||
token=user_info.token,
|
||||
customer_id=user_info.customer_id,
|
||||
user_id=user_info.user_id,
|
||||
team_id=user_info.team_id,
|
||||
team_alias=user_info.team_alias,
|
||||
organization_id=user_info.organization_id,
|
||||
user_email=user_info.user_email,
|
||||
key_alias=user_info.key_alias,
|
||||
projected_exceeded_date=user_info.projected_exceeded_date,
|
||||
projected_spend=user_info.projected_spend,
|
||||
event_group=user_info.event_group,
|
||||
alert_emails=user_info.alert_emails,
|
||||
max_budget_alert_emails=user_info.max_budget_alert_emails,
|
||||
)
|
||||
await self.send_alert(
|
||||
message=event_message + "\n\n" + user_info_str,
|
||||
|
|
@ -657,7 +671,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
"""
|
||||
Create a standard message for a budget alert
|
||||
"""
|
||||
_all_fields_as_dict: Final = user_info.model_dump(exclude_none=True)
|
||||
_all_fields_as_dict: Final[dict[str, object]] = user_info.model_dump(exclude_none=True)
|
||||
_all_fields_as_dict.pop("token")
|
||||
msg = ""
|
||||
for k, v in _all_fields_as_dict.items():
|
||||
|
|
@ -1006,7 +1020,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
except Exception:
|
||||
pass
|
||||
|
||||
async def model_added_alert(self, model_name: str, litellm_model_name: str, passed_model_info: Any):
|
||||
async def model_added_alert(self, model_name: str, litellm_model_name: str, passed_model_info: object):
|
||||
base_model_from_user: Final = getattr(passed_model_info, "base_model", None)
|
||||
model_info = {}
|
||||
base_model = ""
|
||||
|
|
@ -1973,7 +1987,7 @@ Model Info:
|
|||
try:
|
||||
message = f"`{event_name}`\n"
|
||||
|
||||
key_event_dict: Final = key_event.model_dump()
|
||||
key_event_dict: Final[dict[str, object]] = key_event.model_dump()
|
||||
|
||||
# Add Created by information first
|
||||
message += "*Action Done by:*\n"
|
||||
|
|
|
|||
|
|
@ -3,10 +3,12 @@ Arize Phoenix prompt manager that integrates with LiteLLM's prompt management sy
|
|||
Fetches prompt versions from Arize Phoenix and provides workspace-based access control.
|
||||
"""
|
||||
|
||||
from typing import Any, Final
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Final, cast
|
||||
|
||||
from jinja2 import DictLoader, select_autoescape
|
||||
from jinja2.sandbox import ImmutableSandboxedEnvironment
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.integrations.custom_prompt_management import CustomPromptManagement
|
||||
from litellm.integrations.prompt_management_base import (
|
||||
|
|
@ -20,6 +22,31 @@ from litellm.types.utils import StandardCallbackDynamicParams
|
|||
from .arize_phoenix_client import ArizePhoenixClient
|
||||
|
||||
|
||||
class ArizePhoenixContentPart(TypedDict, total=False):
|
||||
type: ReadOnly[str]
|
||||
text: ReadOnly[str]
|
||||
|
||||
|
||||
class ArizePhoenixTemplateMessage(TypedDict, total=False):
|
||||
role: ReadOnly[str]
|
||||
content: ReadOnly[Sequence[ArizePhoenixContentPart]]
|
||||
|
||||
|
||||
class ArizePhoenixTemplateBody(TypedDict, total=False):
|
||||
messages: ReadOnly[Sequence[ArizePhoenixTemplateMessage]]
|
||||
|
||||
|
||||
class ArizePhoenixPromptMetadata(TypedDict):
|
||||
model_name: ReadOnly[str | None]
|
||||
model_provider: ReadOnly[str | None]
|
||||
description: ReadOnly[str]
|
||||
template_type: ReadOnly[str | None]
|
||||
template_format: ReadOnly[str]
|
||||
invocation_parameters: ReadOnly[Mapping[str, Mapping[str, object]]]
|
||||
temperature: ReadOnly[float | None]
|
||||
max_tokens: ReadOnly[int | None]
|
||||
|
||||
|
||||
class ArizePhoenixPromptTemplate:
|
||||
"""
|
||||
Represents a prompt template loaded from Arize Phoenix.
|
||||
|
|
@ -28,10 +55,10 @@ class ArizePhoenixPromptTemplate:
|
|||
def __init__(
|
||||
self,
|
||||
template_id: str,
|
||||
messages: list[dict[str, Any]],
|
||||
metadata: dict[str, Any],
|
||||
messages: Sequence[ArizePhoenixTemplateMessage],
|
||||
metadata: ArizePhoenixPromptMetadata,
|
||||
model: str | None = None,
|
||||
):
|
||||
) -> None:
|
||||
self.template_id = template_id
|
||||
self.messages = messages
|
||||
self.metadata = metadata
|
||||
|
|
@ -43,7 +70,7 @@ class ArizePhoenixPromptTemplate:
|
|||
self.description = metadata.get("description", "")
|
||||
self.template_format = metadata.get("template_format", "MUSTACHE")
|
||||
|
||||
def __repr__(self):
|
||||
def __repr__(self) -> str:
|
||||
return f"ArizePhoenixPromptTemplate(id='{self.template_id}', model='{self.model}')"
|
||||
|
||||
|
||||
|
|
@ -109,7 +136,7 @@ class ArizePhoenixTemplateManager:
|
|||
|
||||
def _parse_prompt_data(self, data: dict[str, Any], prompt_version_id: str) -> ArizePhoenixPromptTemplate:
|
||||
"""Parse Arize Phoenix prompt data and extract messages and metadata."""
|
||||
template_data: Final = data.get("template", {})
|
||||
template_data: Final[ArizePhoenixTemplateBody] = data.get("template", {})
|
||||
messages: Final = template_data.get("messages", [])
|
||||
|
||||
# Extract invocation parameters
|
||||
|
|
@ -129,7 +156,7 @@ class ArizePhoenixTemplateManager:
|
|||
break
|
||||
|
||||
# Build metadata dictionary
|
||||
metadata: Final = {
|
||||
metadata: Final[ArizePhoenixPromptMetadata] = {
|
||||
"model_name": data.get("model_name"),
|
||||
"model_provider": data.get("model_provider"),
|
||||
"description": data.get("description", ""),
|
||||
|
|
@ -146,7 +173,9 @@ class ArizePhoenixTemplateManager:
|
|||
metadata=metadata,
|
||||
)
|
||||
|
||||
def render_template(self, template_id: str, variables: dict[str, Any] | None = None) -> list[AllMessageValues]:
|
||||
def render_template(
|
||||
self, template_id: str, variables: Mapping[str, object] | None = None
|
||||
) -> list[AllMessageValues]:
|
||||
"""Render a template with the given variables and return formatted messages."""
|
||||
if template_id not in self.prompts:
|
||||
raise ValueError(f"Template '{template_id}' not found")
|
||||
|
|
@ -174,7 +203,9 @@ class ArizePhoenixTemplateManager:
|
|||
# Combine rendered content
|
||||
final_content = " ".join(rendered_content_parts)
|
||||
|
||||
rendered_messages.append({"role": role, "content": final_content})
|
||||
rendered_messages.append(
|
||||
cast("AllMessageValues", {"role": role, "content": final_content}) # cast-ok: Phoenix roles are OpenAI
|
||||
)
|
||||
|
||||
return rendered_messages
|
||||
|
||||
|
|
@ -243,8 +274,8 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
|
|||
def get_prompt_template(
|
||||
self,
|
||||
prompt_id: str,
|
||||
prompt_variables: dict[str, Any] | None = None,
|
||||
) -> tuple[list[AllMessageValues], dict[str, Any]]:
|
||||
prompt_variables: Mapping[str, object] | None = None,
|
||||
) -> tuple[list[AllMessageValues], dict[str, object]]:
|
||||
"""
|
||||
Get a prompt template and render it with variables.
|
||||
|
||||
|
|
@ -263,7 +294,7 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
|
|||
rendered_messages: Final = self.prompt_manager.render_template(prompt_id, prompt_variables or {})
|
||||
|
||||
# Extract metadata
|
||||
metadata: Final = {
|
||||
metadata: Final[dict[str, object]] = {
|
||||
"model": template.model,
|
||||
"temperature": template.temperature,
|
||||
"max_tokens": template.max_tokens,
|
||||
|
|
@ -271,7 +302,7 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
|
|||
|
||||
# Add additional invocation parameters
|
||||
invocation_params: Final = template.invocation_parameters
|
||||
provider_params = {}
|
||||
provider_params: Mapping[str, object] = {}
|
||||
|
||||
if "openai" in invocation_params:
|
||||
provider_params = invocation_params["openai"]
|
||||
|
|
@ -289,12 +320,12 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
|
|||
self,
|
||||
user_id: str | None,
|
||||
messages: list[AllMessageValues],
|
||||
function_call: dict[str, Any] | str | None = None,
|
||||
litellm_params: dict[str, Any] | None = None,
|
||||
function_call: dict[str, object] | str | None = None,
|
||||
litellm_params: dict[str, object] | None = None,
|
||||
prompt_id: str | None = None,
|
||||
prompt_variables: dict[str, Any] | None = None,
|
||||
prompt_variables: dict[str, object] | None = None,
|
||||
**kwargs,
|
||||
) -> tuple[list[AllMessageValues], dict[str, Any] | None]:
|
||||
) -> tuple[list[AllMessageValues], dict[str, object] | None]:
|
||||
"""
|
||||
Pre-call hook that processes the prompt template before making the LLM call.
|
||||
"""
|
||||
|
|
@ -335,9 +366,9 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
|
|||
|
||||
except Exception as e:
|
||||
# Log error but don't fail the call
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
litellm._logging.verbose_proxy_logger.error("Error in Arize Phoenix prompt pre_call_hook: %s", e)
|
||||
verbose_proxy_logger.error("Error in Arize Phoenix prompt pre_call_hook: %s", e)
|
||||
return messages, litellm_params
|
||||
|
||||
def get_available_prompts(self) -> list[str]:
|
||||
|
|
@ -393,7 +424,8 @@ class ArizePhoenixPromptManager(CustomPromptManagement):
|
|||
rendered_messages, prompt_metadata = self.get_prompt_template(prompt_id, prompt_variables)
|
||||
|
||||
# Extract model from metadata (if specified)
|
||||
template_model: Final = prompt_metadata.get("model")
|
||||
raw_template_model: Final = prompt_metadata.get("model")
|
||||
template_model: Final = raw_template_model if isinstance(raw_template_model, str) else None
|
||||
|
||||
# Extract optional parameters from metadata
|
||||
optional_params: Final = {}
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ BitBucket prompt manager that integrates with LiteLLM's prompt management system
|
|||
Fetches .prompt files from BitBucket repositories and provides team-based access control.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from jinja2 import DictLoader, select_autoescape
|
||||
|
|
@ -65,7 +66,7 @@ class BitBucketTemplateManager:
|
|||
|
||||
def __init__(
|
||||
self,
|
||||
bitbucket_config: dict[str, Any],
|
||||
bitbucket_config: Mapping[str, object],
|
||||
prompt_id: str | None = None,
|
||||
):
|
||||
self.bitbucket_config = bitbucket_config
|
||||
|
|
@ -123,7 +124,7 @@ class BitBucketTemplateManager:
|
|||
template_content = content
|
||||
|
||||
# Parse YAML frontmatter
|
||||
metadata: dict[str, Any] = {}
|
||||
metadata: dict[str, object] = {}
|
||||
if frontmatter_str:
|
||||
try:
|
||||
import yaml
|
||||
|
|
@ -141,9 +142,9 @@ class BitBucketTemplateManager:
|
|||
metadata=metadata,
|
||||
)
|
||||
|
||||
def _parse_yaml_basic(self, yaml_str: str) -> dict[str, Any]:
|
||||
def _parse_yaml_basic(self, yaml_str: str) -> dict[str, object]:
|
||||
"""Basic YAML parser for simple cases when PyYAML is not available."""
|
||||
result: Final[dict[str, Any]] = {}
|
||||
result: Final[dict[str, object]] = {}
|
||||
for line in yaml_str.split("\n"):
|
||||
line = line.strip()
|
||||
if ":" in line and not line.startswith("#"):
|
||||
|
|
@ -162,7 +163,7 @@ class BitBucketTemplateManager:
|
|||
result[key] = value.strip("\"'")
|
||||
return result
|
||||
|
||||
def render_template(self, template_id: str, variables: dict[str, Any] | None = None) -> str:
|
||||
def render_template(self, template_id: str, variables: Mapping[str, object] | None = None) -> str:
|
||||
"""Render a template with the given variables."""
|
||||
if template_id not in self.prompts:
|
||||
raise ValueError(f"Template '{template_id}' not found")
|
||||
|
|
@ -209,7 +210,7 @@ class BitBucketPromptManager(CustomPromptManagement):
|
|||
|
||||
def __init__(
|
||||
self,
|
||||
bitbucket_config: dict[str, Any],
|
||||
bitbucket_config: Mapping[str, object],
|
||||
prompt_id: str | None = None,
|
||||
):
|
||||
self.bitbucket_config = bitbucket_config
|
||||
|
|
@ -234,7 +235,7 @@ class BitBucketPromptManager(CustomPromptManagement):
|
|||
def get_prompt_template(
|
||||
self,
|
||||
prompt_id: str,
|
||||
prompt_variables: dict[str, Any] | None = None,
|
||||
prompt_variables: Mapping[str, object] | None = None,
|
||||
) -> tuple[str, dict[str, Any]]:
|
||||
"""
|
||||
Get a prompt template and render it with variables.
|
||||
|
|
@ -267,12 +268,12 @@ class BitBucketPromptManager(CustomPromptManagement):
|
|||
self,
|
||||
user_id: str | None,
|
||||
messages: list[AllMessageValues],
|
||||
function_call: dict[str, Any] | str | None = None,
|
||||
litellm_params: dict[str, Any] | None = None,
|
||||
function_call: Mapping[str, object] | str | None = None,
|
||||
litellm_params: dict[str, object] | None = None,
|
||||
prompt_id: str | None = None,
|
||||
prompt_variables: dict[str, Any] | None = None,
|
||||
prompt_variables: Mapping[str, object] | None = None,
|
||||
**kwargs,
|
||||
) -> tuple[list[AllMessageValues], dict[str, Any] | None]:
|
||||
) -> tuple[list[AllMessageValues], dict[str, object] | None]:
|
||||
"""
|
||||
Pre-call hook that processes the prompt template before making the LLM call.
|
||||
"""
|
||||
|
|
@ -316,9 +317,9 @@ class BitBucketPromptManager(CustomPromptManagement):
|
|||
|
||||
except Exception as e:
|
||||
# Log error but don't fail the call
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
litellm._logging.verbose_proxy_logger.error("Error in BitBucket prompt pre_call_hook: %s", e)
|
||||
verbose_proxy_logger.error("Error in BitBucket prompt pre_call_hook: %s", e)
|
||||
return messages, litellm_params
|
||||
|
||||
def _parse_prompt_to_messages(self, prompt_content: str) -> list[AllMessageValues]:
|
||||
|
|
@ -384,14 +385,14 @@ class BitBucketPromptManager(CustomPromptManagement):
|
|||
def post_call_hook(
|
||||
self,
|
||||
user_id: str | None,
|
||||
response: Any,
|
||||
response: object,
|
||||
input_messages: list[AllMessageValues],
|
||||
function_call: dict[str, Any] | str | None = None,
|
||||
litellm_params: dict[str, Any] | None = None,
|
||||
function_call: Mapping[str, object] | str | None = None,
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
prompt_id: str | None = None,
|
||||
prompt_variables: dict[str, Any] | None = None,
|
||||
prompt_variables: Mapping[str, object] | None = None,
|
||||
**kwargs,
|
||||
) -> Any:
|
||||
) -> object:
|
||||
"""
|
||||
Post-call hook for any post-processing after the LLM call.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -19,14 +19,29 @@
|
|||
"""Transform LiteLLM data to CloudZero AnyCost CBF format."""
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any, Final
|
||||
from typing import Final, SupportsFloat, SupportsIndex, SupportsInt
|
||||
|
||||
import polars as pl
|
||||
from typing_extensions import Buffer
|
||||
|
||||
from ...types.integrations.cloudzero import CBFRecord
|
||||
from .cz_resource_names import CZEntityType, CZRNGenerator
|
||||
|
||||
|
||||
def _as_int(value: object) -> int:
|
||||
"""The integer form of a spend table cell, computed the way :func:`int` computes it."""
|
||||
if isinstance(value, (str, Buffer, SupportsInt, SupportsIndex)):
|
||||
return int(value)
|
||||
raise TypeError(f"int() argument must be a string or a number, not {type(value).__name__!r}")
|
||||
|
||||
|
||||
def _as_float(value: object) -> float:
|
||||
"""The floating point form of a spend table cell, computed the way :func:`float` computes it."""
|
||||
if isinstance(value, (str, Buffer, SupportsFloat, SupportsIndex)):
|
||||
return float(value)
|
||||
raise TypeError(f"float() argument must be a string or a number, not {type(value).__name__!r}")
|
||||
|
||||
|
||||
class CBFTransformer:
|
||||
"""Transform LiteLLM usage data to CloudZero Billing Format (CBF)."""
|
||||
|
||||
|
|
@ -82,15 +97,15 @@ class CBFTransformer:
|
|||
|
||||
return pl.DataFrame(cbf_data)
|
||||
|
||||
def _create_cbf_record(self, row: dict[str, Any]) -> CBFRecord:
|
||||
def _create_cbf_record(self, row: dict[str, object]) -> CBFRecord:
|
||||
"""Create a single CBF record from LiteLLM daily spend row."""
|
||||
|
||||
# Parse date (daily spend tables use date strings like '2025-04-19')
|
||||
usage_date: Final = self._parse_date(row.get("date"))
|
||||
|
||||
# Calculate total tokens
|
||||
prompt_tokens: Final = int(row.get("prompt_tokens", 0))
|
||||
completion_tokens: Final = int(row.get("completion_tokens", 0))
|
||||
prompt_tokens: Final = _as_int(row.get("prompt_tokens", 0))
|
||||
completion_tokens: Final = _as_int(row.get("completion_tokens", 0))
|
||||
total_tokens: Final = prompt_tokens + completion_tokens
|
||||
|
||||
# Create CloudZero Resource Name (CZRN) as resource_id
|
||||
|
|
@ -154,7 +169,7 @@ class CBFTransformer:
|
|||
"time/usage_start": (
|
||||
usage_date.isoformat() if usage_date else None
|
||||
), # Required: ISO-formatted UTC datetime
|
||||
"cost/cost": float(row.get("spend", 0.0)), # Required: billed cost
|
||||
"cost/cost": _as_float(row.get("spend", 0.0)), # Required: billed cost
|
||||
"resource/id": resource_id, # CZRN (CloudZero Resource Name)
|
||||
# Usage metrics for token consumption
|
||||
"usage/amount": total_tokens, # Numeric value of tokens consumed
|
||||
|
|
@ -187,7 +202,7 @@ class CBFTransformer:
|
|||
|
||||
return CBFRecord(cbf_record)
|
||||
|
||||
def _parse_date(self, date_str) -> datetime | None:
|
||||
def _parse_date(self, date_str: object) -> datetime | None:
|
||||
"""Parse date string from daily spend tables (e.g., '2025-04-19')."""
|
||||
if date_str is None:
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import contextvars
|
|||
import hashlib
|
||||
import os
|
||||
import secrets
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, get_args
|
||||
|
||||
|
|
@ -227,13 +228,13 @@ class CustomGuardrail(CustomLogger):
|
|||
)
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def render_violation_message(self, default: str, context: dict[str, Any] | None = None) -> str:
|
||||
def render_violation_message(self, default: str, context: Mapping[str, object] | None = None) -> str:
|
||||
"""Return a custom violation message if template is configured."""
|
||||
|
||||
if not self.violation_message_template:
|
||||
return default
|
||||
|
||||
format_context: Final[dict[str, Any]] = {"default_message": default}
|
||||
format_context: Final[dict[str, object]] = {"default_message": default}
|
||||
if context:
|
||||
format_context.update(context)
|
||||
try:
|
||||
|
|
@ -661,7 +662,7 @@ class CustomGuardrail(CustomLogger):
|
|||
value: Final = self._get_admin_metadata(data).get("opted_out_global_guardrails")
|
||||
return value if isinstance(value, list) else []
|
||||
|
||||
def _is_valid_response_type(self, result: Any) -> bool:
|
||||
def _is_valid_response_type(self, result: object) -> bool:
|
||||
"""
|
||||
Check if result is a valid LLMResponseTypes instance.
|
||||
|
||||
|
|
@ -722,7 +723,7 @@ class CustomGuardrail(CustomLogger):
|
|||
return None
|
||||
return f"{_PRE_CALL_EXECUTED_TOKEN}:{name}"
|
||||
|
||||
def mark_pre_call_hook_ran(self, data: dict[str, Any]) -> None:
|
||||
def mark_pre_call_hook_ran(self, data: dict[str, object]) -> None:
|
||||
"""
|
||||
Record that this guardrail's ``async_pre_call_hook`` already ran for this
|
||||
request, so the deployment-level hook does not run it a second time.
|
||||
|
|
@ -747,7 +748,7 @@ class CustomGuardrail(CustomLogger):
|
|||
return
|
||||
data["metadata"] = {PRE_CALL_EXECUTED_GUARDRAILS_KEY: [marker]}
|
||||
|
||||
def _pre_call_hook_already_ran(self, data: dict[str, Any]) -> bool:
|
||||
def _pre_call_hook_already_ran(self, data: dict[str, object]) -> bool:
|
||||
marker: Final = self._pre_call_marker()
|
||||
if marker is None:
|
||||
return False
|
||||
|
|
@ -1170,7 +1171,7 @@ class CustomGuardrail(CustomLogger):
|
|||
This gets logged on downsteam Langfuse, DataDog, etc.
|
||||
"""
|
||||
# Convert None to empty dict to satisfy type requirements
|
||||
guardrail_response: dict[str, Any] | str = {} if response is None else response
|
||||
guardrail_response: dict[str, object] | str = {} if response is None else response
|
||||
|
||||
# For apply_guardrail functions in custom_code_guardrail scenario,
|
||||
# simplify the logged response to "allow", "deny", or "mask"
|
||||
|
|
|
|||
|
|
@ -31,6 +31,9 @@ if TYPE_CHECKING:
|
|||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.anthropic_messages.transformation import (
|
||||
BaseAnthropicMessagesConfig,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.mcp import (
|
||||
MCPPostCallResponseObject,
|
||||
|
|
@ -39,7 +42,7 @@ if TYPE_CHECKING:
|
|||
)
|
||||
from litellm.types.router import PreRoutingHookResponse
|
||||
|
||||
Span = _Span | Any
|
||||
Span = _Span
|
||||
else:
|
||||
Span = Any
|
||||
LiteLLMLoggingObj = Any
|
||||
|
|
@ -268,7 +271,9 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
) -> list[dict]:
|
||||
return healthy_deployments
|
||||
|
||||
async def async_pre_call_deployment_hook(self, kwargs: dict[str, Any], call_type: CallTypes | None) -> dict | None:
|
||||
async def async_pre_call_deployment_hook(
|
||||
self, kwargs: dict[str, object], call_type: CallTypes | None
|
||||
) -> dict | None:
|
||||
"""
|
||||
Allow modifying the request just before it's sent to the deployment.
|
||||
|
||||
|
|
@ -344,9 +349,9 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
async def async_post_call_streaming_deployment_hook(
|
||||
self,
|
||||
request_data: dict,
|
||||
response_chunk: Any,
|
||||
response_chunk: object,
|
||||
call_type: CallTypes | None,
|
||||
) -> Any | None:
|
||||
) -> object | None:
|
||||
"""
|
||||
Allow modifying streaming chunks just before they're returned to the user.
|
||||
|
||||
|
|
@ -378,7 +383,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
"""
|
||||
|
||||
def translate_completion_output_params_streaming(
|
||||
self, completion_stream: Any
|
||||
self, completion_stream: object
|
||||
) -> AdapterCompletionStreamWrapper | None:
|
||||
"""
|
||||
Translates the streaming chunk, from the OpenAI format to the custom format.
|
||||
|
|
@ -418,9 +423,9 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
self,
|
||||
data: dict,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: Any,
|
||||
response: object,
|
||||
request_headers: dict[str, str] | None = None,
|
||||
litellm_call_info: dict[str, Any] | None = None,
|
||||
litellm_call_info: dict[str, object] | None = None,
|
||||
) -> dict[str, str] | None:
|
||||
"""
|
||||
Called after an LLM API call (success or failure) to allow injecting custom HTTP response headers.
|
||||
|
|
@ -471,11 +476,11 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
) -> Any:
|
||||
pass
|
||||
|
||||
async def async_logging_hook(self, kwargs: dict, result: Any, call_type: str) -> tuple[dict, Any]:
|
||||
async def async_logging_hook(self, kwargs: dict, result: object, call_type: str) -> tuple[dict, object]:
|
||||
"""For masking logged request/response. Return a modified version of the request/result."""
|
||||
return kwargs, result
|
||||
|
||||
def logging_hook(self, kwargs: dict, result: Any, call_type: str) -> tuple[dict, Any]:
|
||||
def logging_hook(self, kwargs: dict, result: object, call_type: str) -> tuple[dict, object]:
|
||||
"""For masking logged request/response. Return a modified version of the request/result."""
|
||||
return kwargs, result
|
||||
|
||||
|
|
@ -581,7 +586,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
|
||||
async def async_should_run_agentic_loop(
|
||||
self,
|
||||
response: Any,
|
||||
response: object,
|
||||
model: str,
|
||||
messages: list[dict],
|
||||
tools: list[dict] | None,
|
||||
|
|
@ -642,8 +647,8 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
tools: dict,
|
||||
model: str,
|
||||
messages: list[dict],
|
||||
response: Any,
|
||||
anthropic_messages_provider_config: Any,
|
||||
response: object,
|
||||
anthropic_messages_provider_config: "BaseAnthropicMessagesConfig | None",
|
||||
anthropic_messages_optional_request_params: dict,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
stream: bool,
|
||||
|
|
@ -711,8 +716,8 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
tools: dict,
|
||||
model: str,
|
||||
messages: list[dict],
|
||||
response: Any,
|
||||
anthropic_messages_provider_config: Any,
|
||||
response: object,
|
||||
anthropic_messages_provider_config: "BaseAnthropicMessagesConfig | None",
|
||||
anthropic_messages_optional_request_params: dict,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
stream: bool,
|
||||
|
|
@ -728,7 +733,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
|
||||
async def async_post_agentic_loop_response_hook(
|
||||
self,
|
||||
response: Any,
|
||||
response: object,
|
||||
plan: AgenticLoopPlan,
|
||||
kwargs: dict,
|
||||
) -> Any:
|
||||
|
|
@ -767,7 +772,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
|
||||
async def async_should_run_chat_completion_agentic_loop(
|
||||
self,
|
||||
response: Any,
|
||||
response: object,
|
||||
model: str,
|
||||
messages: list[dict],
|
||||
tools: list[dict] | None,
|
||||
|
|
@ -785,12 +790,12 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
tools: dict,
|
||||
model: str,
|
||||
messages: list[dict],
|
||||
response: Any,
|
||||
response: object,
|
||||
optional_params: dict,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
stream: bool,
|
||||
kwargs: dict,
|
||||
) -> Any:
|
||||
) -> object:
|
||||
"""
|
||||
Hook to execute chat completion agentic loop based on context from should_run hook.
|
||||
"""
|
||||
|
|
@ -800,7 +805,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
tools: dict,
|
||||
model: str,
|
||||
messages: list[dict],
|
||||
response: Any,
|
||||
response: object,
|
||||
optional_params: dict,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
stream: bool,
|
||||
|
|
@ -1056,7 +1061,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
|
||||
def _redact_base64(
|
||||
self,
|
||||
value: Any,
|
||||
value: object,
|
||||
depth: int = 0,
|
||||
max_depth: int = DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER,
|
||||
) -> object:
|
||||
|
|
@ -1079,7 +1084,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
|
||||
return value
|
||||
|
||||
def _should_keep_content(self, content: Any) -> bool:
|
||||
def _should_keep_content(self, content: object) -> bool:
|
||||
"""Return True if this content item should be retained."""
|
||||
if not isinstance(content, dict):
|
||||
return True
|
||||
|
|
|
|||
|
|
@ -20,10 +20,11 @@ import time
|
|||
import traceback
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime as datetimeObj
|
||||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import httpx
|
||||
from httpx import Response
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -62,6 +63,18 @@ from litellm.types.utils import StandardLoggingPayload
|
|||
|
||||
from ..additional_logging_utils import AdditionalLoggingUtils
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
||||
class _DatadogLoggingKwargs(TypedDict, total=False):
|
||||
"""The subset of logging ``kwargs`` that the Datadog payload builder reads."""
|
||||
|
||||
standard_logging_object: ReadOnly[StandardLoggingPayload | None]
|
||||
|
||||
|
||||
# max number of logs DD API can accept
|
||||
|
||||
|
||||
|
|
@ -87,6 +100,11 @@ def _resolve_dd_batch_size() -> int:
|
|||
return max(1, min(value, DD_MAX_BATCH_SIZE))
|
||||
|
||||
|
||||
def _span_attribute(span: object, name: str) -> object:
|
||||
"""Read an optional attribute off whatever span object the active tracer hands back."""
|
||||
return getattr(span, name, None)
|
||||
|
||||
|
||||
class DataDogLogger(
|
||||
CustomBatchLogger,
|
||||
AdditionalLoggingUtils,
|
||||
|
|
@ -271,9 +289,9 @@ class DataDogLogger(
|
|||
self,
|
||||
request_data: dict,
|
||||
original_exception: Exception,
|
||||
user_api_key_dict: Any,
|
||||
user_api_key_dict: "UserAPIKeyAuth",
|
||||
traceback_str: str | None = None,
|
||||
) -> Any | None:
|
||||
) -> "HTTPException | None":
|
||||
"""
|
||||
Log proxy-level failures (e.g. 401 auth, DB connection errors) to Datadog.
|
||||
|
||||
|
|
@ -297,7 +315,7 @@ class DataDogLogger(
|
|||
status_code = int(_code)
|
||||
|
||||
# Use project-standard sanitized user context when running in proxy
|
||||
user_context: dict[str, Any] = {}
|
||||
user_context: dict[str, object] = {}
|
||||
try:
|
||||
from litellm.proxy.litellm_pre_call_utils import (
|
||||
LiteLLMProxyRequestSetup,
|
||||
|
|
@ -553,8 +571,8 @@ class DataDogLogger(
|
|||
|
||||
def create_datadog_logging_payload(
|
||||
self,
|
||||
kwargs: dict | Any,
|
||||
response_obj: Any,
|
||||
kwargs: _DatadogLoggingKwargs,
|
||||
response_obj: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> DatadogPayload:
|
||||
|
|
@ -562,8 +580,8 @@ class DataDogLogger(
|
|||
Helper function to create a datadog payload for logging
|
||||
|
||||
Args:
|
||||
kwargs (Union[dict, Any]): request kwargs
|
||||
response_obj (Any): llm api response
|
||||
kwargs: request kwargs, read for its standard logging object
|
||||
response_obj: llm api response
|
||||
start_time (datetime.datetime): start time of request
|
||||
end_time (datetime.datetime): end time of request
|
||||
|
||||
|
|
@ -625,7 +643,7 @@ class DataDogLogger(
|
|||
self,
|
||||
payload: ServiceLoggerPayload,
|
||||
error: str | None = "",
|
||||
parent_otel_span: Any | None = None,
|
||||
parent_otel_span: object = None,
|
||||
start_time: datetimeObj | float | None = None,
|
||||
end_time: float | datetimeObj | None = None,
|
||||
event_metadata: dict | None = None,
|
||||
|
|
@ -659,7 +677,7 @@ class DataDogLogger(
|
|||
self,
|
||||
payload: ServiceLoggerPayload,
|
||||
error: str | None = "",
|
||||
parent_otel_span: Any | None = None,
|
||||
parent_otel_span: object = None,
|
||||
start_time: datetimeObj | float | None = None,
|
||||
end_time: float | datetimeObj | None = None,
|
||||
event_metadata: dict | None = None,
|
||||
|
|
@ -696,7 +714,7 @@ class DataDogLogger(
|
|||
|
||||
def _create_v0_logging_payload(
|
||||
self,
|
||||
kwargs: dict | Any,
|
||||
kwargs: dict,
|
||||
response_obj: Any,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
|
|
@ -810,11 +828,11 @@ class DataDogLogger(
|
|||
if current_span is None:
|
||||
return None
|
||||
|
||||
trace_id: Final = getattr(current_span, "trace_id", None)
|
||||
trace_id: Final = _span_attribute(current_span, "trace_id")
|
||||
if trace_id is None:
|
||||
return None
|
||||
|
||||
span_id: Final = getattr(current_span, "span_id", None)
|
||||
span_id: Final = _span_attribute(current_span, "span_id")
|
||||
trace_context: Final[dict[str, str]] = {"trace_id": str(trace_id)}
|
||||
if span_id is not None:
|
||||
trace_context["span_id"] = str(span_id)
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ API Reference: https://docs.datadoghq.com/llm_observability/setup/api/?tab=examp
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import Any, Final, Literal
|
||||
|
||||
|
|
@ -334,7 +335,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
|
|||
|
||||
def _get_response_messages(
|
||||
self, standard_logging_payload: StandardLoggingPayload, call_type: str | None
|
||||
) -> list[Any]:
|
||||
) -> list[object]:
|
||||
"""
|
||||
Get the messages from the response object
|
||||
|
||||
|
|
@ -484,7 +485,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
|
|||
# Default fallback for unknown or passthrough operations
|
||||
return "llm"
|
||||
|
||||
def _ensure_string_content(self, messages: str | list[Any] | dict[Any, Any] | None) -> list[Any]:
|
||||
def _ensure_string_content(self, messages: str | Sequence[object] | Mapping[object, object] | None) -> list[object]:
|
||||
if messages is None:
|
||||
return []
|
||||
if isinstance(messages, str):
|
||||
|
|
@ -495,11 +496,11 @@ class DataDogLLMObsLogger(CustomBatchLogger):
|
|||
return [str(messages.get("content", ""))]
|
||||
return []
|
||||
|
||||
def _get_dd_llm_obs_payload_metadata(self, standard_logging_payload: StandardLoggingPayload) -> dict[str, Any]:
|
||||
def _get_dd_llm_obs_payload_metadata(self, standard_logging_payload: StandardLoggingPayload) -> dict[str, object]:
|
||||
"""
|
||||
Fields to track in DD LLM Observability metadata from litellm standard logging payload
|
||||
"""
|
||||
_metadata: Final[dict[str, Any]] = {
|
||||
_metadata: Final[dict[str, object]] = {
|
||||
"model_name": standard_logging_payload.get("model", "unknown"),
|
||||
"model_provider": standard_logging_payload.get("custom_llm_provider", "unknown"),
|
||||
"id": standard_logging_payload.get("id", "unknown"),
|
||||
|
|
@ -647,7 +648,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
|
|||
|
||||
return spend_metrics
|
||||
|
||||
def _process_input_messages_preserving_tool_calls(self, messages: list[Any]) -> list[dict[str, Any]]:
|
||||
def _process_input_messages_preserving_tool_calls(self, messages: Sequence[object]) -> list[dict[str, object]]:
|
||||
"""
|
||||
Process input messages while preserving tool_calls and tool message types.
|
||||
|
||||
|
|
@ -671,13 +672,13 @@ class DataDogLLMObsLogger(CustomBatchLogger):
|
|||
return processed
|
||||
|
||||
@staticmethod
|
||||
def _tool_calls_kv_pair(tool_calls: list[dict[str, Any]]) -> dict[str, Any]:
|
||||
def _tool_calls_kv_pair(tool_calls: list[dict[str, Any]]) -> dict[str, object]:
|
||||
"""
|
||||
Extract tool call information into key-value pairs for Datadog metadata.
|
||||
|
||||
Similar to OpenTelemetry's implementation but adapted for Datadog's format.
|
||||
"""
|
||||
kv_pairs: Final[dict[str, Any]] = {}
|
||||
kv_pairs: Final[dict[str, object]] = {}
|
||||
for idx, tool_call in enumerate(tool_calls):
|
||||
try:
|
||||
# Extract tool call ID
|
||||
|
|
@ -712,11 +713,11 @@ class DataDogLLMObsLogger(CustomBatchLogger):
|
|||
|
||||
return kv_pairs
|
||||
|
||||
def _extract_tool_call_metadata(self, standard_logging_payload: StandardLoggingPayload) -> dict[str, Any]:
|
||||
def _extract_tool_call_metadata(self, standard_logging_payload: StandardLoggingPayload) -> dict[str, object]:
|
||||
"""
|
||||
Extract tool call information from both input messages and response for Datadog metadata.
|
||||
"""
|
||||
tool_call_metadata: Final[dict[str, Any]] = {}
|
||||
tool_call_metadata: Final[dict[str, object]] = {}
|
||||
|
||||
try:
|
||||
# Extract tool calls from input messages
|
||||
|
|
|
|||
|
|
@ -3,12 +3,21 @@ Based on Google's GenAI Kit dotprompt implementation: https://google.github.io/d
|
|||
"""
|
||||
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from pathlib import Path
|
||||
from typing import Any, Final
|
||||
|
||||
import yaml
|
||||
from jinja2 import DictLoader, select_autoescape
|
||||
from jinja2.sandbox import ImmutableSandboxedEnvironment
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
|
||||
class _PromptFileJson(TypedDict):
|
||||
"""JSON form of a .prompt file: rendered template text plus its frontmatter."""
|
||||
|
||||
content: ReadOnly[NotRequired[str]]
|
||||
metadata: ReadOnly[NotRequired[dict[str, object]]]
|
||||
|
||||
|
||||
def strip_version_suffix(prompt_id: str) -> str | None:
|
||||
|
|
@ -167,7 +176,7 @@ class PromptManager:
|
|||
template_id=prompt_id,
|
||||
)
|
||||
|
||||
def _parse_frontmatter(self, content: str) -> tuple[dict[str, Any], str]:
|
||||
def _parse_frontmatter(self, content: str) -> tuple[dict[str, object], str]:
|
||||
"""Parse YAML frontmatter from prompt content."""
|
||||
# Match YAML frontmatter between --- delimiters
|
||||
frontmatter_pattern: Final = r"^---\s*\n(.*?)\n---\s*\n(.*)$"
|
||||
|
|
@ -178,7 +187,7 @@ class PromptManager:
|
|||
template_content = match.group(2)
|
||||
|
||||
try:
|
||||
frontmatter = yaml.safe_load(frontmatter_yaml) or {}
|
||||
frontmatter: dict[str, object] = yaml.safe_load(frontmatter_yaml) or {}
|
||||
except yaml.YAMLError as e:
|
||||
raise ValueError(f"Invalid YAML frontmatter: {e}")
|
||||
else:
|
||||
|
|
@ -191,7 +200,7 @@ class PromptManager:
|
|||
def render(
|
||||
self,
|
||||
prompt_id: str,
|
||||
prompt_variables: dict[str, Any] | None = None,
|
||||
prompt_variables: Mapping[str, object] | None = None,
|
||||
version: int | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
|
|
@ -231,7 +240,7 @@ class PromptManager:
|
|||
except Exception as e:
|
||||
raise ValueError(f"Error rendering template '{prompt_id}': {e}")
|
||||
|
||||
def _validate_input(self, variables: dict[str, Any], schema: dict[str, Any]) -> None:
|
||||
def _validate_input(self, variables: Mapping[str, object], schema: Mapping[str, str]) -> None:
|
||||
"""Basic validation of input variables against schema."""
|
||||
for field_name, field_type in schema.items():
|
||||
if field_name in variables:
|
||||
|
|
@ -291,7 +300,7 @@ class PromptManager:
|
|||
"""Get a list of all available prompt IDs."""
|
||||
return list(self.prompts.keys())
|
||||
|
||||
def get_prompt_metadata(self, prompt_id: str) -> dict[str, Any] | None:
|
||||
def get_prompt_metadata(self, prompt_id: str) -> dict[str, object] | None:
|
||||
"""Get metadata for a specific prompt."""
|
||||
template: Final = self.prompts.get(prompt_id)
|
||||
return template.metadata if template else None
|
||||
|
|
@ -302,12 +311,12 @@ class PromptManager:
|
|||
if self.prompt_directory:
|
||||
self._load_prompts()
|
||||
|
||||
def add_prompt(self, prompt_id: str, content: str, metadata: dict[str, Any] | None = None) -> None:
|
||||
def add_prompt(self, prompt_id: str, content: str, metadata: dict[str, object] | None = None) -> None:
|
||||
"""Add a prompt template programmatically."""
|
||||
template: Final = PromptTemplate(content=content, metadata=metadata or {}, template_id=prompt_id)
|
||||
self.prompts[prompt_id] = template
|
||||
|
||||
def prompt_file_to_json(self, file_path: str | Path) -> dict[str, Any]:
|
||||
def prompt_file_to_json(self, file_path: str | Path) -> _PromptFileJson:
|
||||
"""Convert a .prompt file to JSON format.
|
||||
|
||||
Args:
|
||||
|
|
@ -324,7 +333,7 @@ class PromptManager:
|
|||
|
||||
return {"content": template_content.strip(), "metadata": frontmatter}
|
||||
|
||||
def json_to_prompt_file(self, prompt_data: dict[str, Any]) -> str:
|
||||
def json_to_prompt_file(self, prompt_data: _PromptFileJson) -> str:
|
||||
"""Convert JSON prompt data to .prompt file format.
|
||||
|
||||
Args:
|
||||
|
|
|
|||
|
|
@ -6,10 +6,11 @@ import re
|
|||
import uuid
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime, timezone, tzinfo
|
||||
from typing import Any, Final, TypedDict, cast
|
||||
from typing import Any, Final, Protocol, cast
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, Field
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -35,6 +36,34 @@ GALILEO_CLOUD_API_BASE_URL: Final = "https://api.galileo.ai"
|
|||
GALILEO_MAX_IN_MEMORY_RECORDS: Final = 1000
|
||||
|
||||
|
||||
class _GalileoLoginBody(TypedDict):
|
||||
"""Decoded body of the Galileo login response."""
|
||||
|
||||
access_token: ReadOnly[str]
|
||||
|
||||
|
||||
class _GalileoLoginResponse(Protocol):
|
||||
"""The login call's HTTP response, read for the access token it carries."""
|
||||
|
||||
def json(self) -> _GalileoLoginBody: ...
|
||||
|
||||
|
||||
class _JsonResponse(Protocol):
|
||||
"""An HTTP response read only for whatever JSON body it decodes to."""
|
||||
|
||||
def json(self) -> object: ...
|
||||
|
||||
|
||||
def _login_access_token(response: _GalileoLoginResponse) -> str:
|
||||
"""Read the bearer token out of a Galileo login response body."""
|
||||
return response.json()["access_token"]
|
||||
|
||||
|
||||
def _decoded_body(response: _JsonResponse) -> object:
|
||||
"""Decode a response body without asserting anything about its shape."""
|
||||
return response.json()
|
||||
|
||||
|
||||
class GalileoStandardLoggingFields(TypedDict, total=False):
|
||||
call_type: str
|
||||
model: str
|
||||
|
|
@ -156,7 +185,7 @@ class GalileoObserve(CustomLogger):
|
|||
},
|
||||
)
|
||||
galileo_login_response.raise_for_status()
|
||||
access_token: Final = galileo_login_response.json()["access_token"]
|
||||
access_token: Final = _login_access_token(galileo_login_response)
|
||||
self.headers = {
|
||||
"accept": "application/json",
|
||||
"Content-Type": "application/json",
|
||||
|
|
@ -421,7 +450,7 @@ class GalileoObserve(CustomLogger):
|
|||
try:
|
||||
verbose_logger.debug(
|
||||
"Galileo Logger HTTP error response json: %s",
|
||||
response.json(),
|
||||
_decoded_body(response),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -4,12 +4,80 @@ Now supports selecting a tag via `config["tag"]`; falls back to branch ("main").
|
|||
"""
|
||||
|
||||
import base64
|
||||
from typing import Any, Final
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Final, Protocol, TypedDict
|
||||
from urllib.parse import quote
|
||||
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
|
||||
class GitLabFilePayload(TypedDict, total=False):
|
||||
"""A repository-files API entry."""
|
||||
|
||||
content: ReadOnly[str]
|
||||
encoding: ReadOnly[str]
|
||||
|
||||
|
||||
class GitLabTreeEntry(TypedDict, total=False):
|
||||
"""A repository-tree API entry."""
|
||||
|
||||
path: ReadOnly[str]
|
||||
type: ReadOnly[str]
|
||||
|
||||
|
||||
class GitLabBranch(TypedDict, total=False):
|
||||
"""A repository-branches API entry."""
|
||||
|
||||
name: ReadOnly[str]
|
||||
type: ReadOnly[str]
|
||||
|
||||
|
||||
class GitLabFileMetadata(TypedDict):
|
||||
"""The response headers a raw file request exposes as metadata."""
|
||||
|
||||
content_type: ReadOnly[str | None]
|
||||
content_length: ReadOnly[str | None]
|
||||
last_modified: ReadOnly[str | None]
|
||||
|
||||
|
||||
class _FileJsonResponse(Protocol):
|
||||
def json(self) -> GitLabFilePayload: ...
|
||||
|
||||
|
||||
class _TreeJsonResponse(Protocol):
|
||||
def json(self) -> Sequence[GitLabTreeEntry] | None: ...
|
||||
|
||||
|
||||
class _ProjectJsonResponse(Protocol):
|
||||
def json(self) -> Mapping[str, object]: ...
|
||||
|
||||
|
||||
class _BranchesJsonResponse(Protocol):
|
||||
def json(self) -> Sequence[GitLabBranch] | None: ...
|
||||
|
||||
|
||||
def _file_payload(resp: _FileJsonResponse) -> GitLabFilePayload:
|
||||
"""The JSON body of a repository-files response."""
|
||||
return resp.json()
|
||||
|
||||
|
||||
def _tree_entries(resp: _TreeJsonResponse) -> Sequence[GitLabTreeEntry]:
|
||||
"""The entries of a repository-tree response."""
|
||||
return resp.json() or []
|
||||
|
||||
|
||||
def _project_info(resp: _ProjectJsonResponse) -> Mapping[str, object]:
|
||||
"""The JSON body of a project response."""
|
||||
return resp.json()
|
||||
|
||||
|
||||
def _branch_entries(resp: _BranchesJsonResponse) -> Sequence[GitLabBranch] | None:
|
||||
"""The JSON body of a repository-branches response."""
|
||||
return resp.json()
|
||||
|
||||
|
||||
class GitLabClient:
|
||||
"""
|
||||
Client for interacting with the GitLab API to fetch files.
|
||||
|
|
@ -42,12 +110,12 @@ class GitLabClient:
|
|||
|
||||
self.project: str | int = project
|
||||
self.access_token: str = str(access_token)
|
||||
self.auth_method = config.get("auth_method", "token") # 'token' or 'oauth'
|
||||
self.auth_method: str = config.get("auth_method", "token") # 'token' or 'oauth'
|
||||
self.branch = config.get("branch", None)
|
||||
if not self.branch:
|
||||
self.branch = "main"
|
||||
self.tag = config.get("tag")
|
||||
self.base_url = config.get("base_url", "https://gitlab.com/api/v4")
|
||||
self.base_url: str = config.get("base_url", "https://gitlab.com/api/v4")
|
||||
|
||||
if not all([self.project, self.access_token]):
|
||||
raise ValueError("project and access_token are required")
|
||||
|
|
@ -159,7 +227,7 @@ class GitLabClient:
|
|||
if resp.status_code == 404:
|
||||
return None
|
||||
resp.raise_for_status()
|
||||
data: Final = resp.json()
|
||||
data: Final = _file_payload(resp)
|
||||
content: Final = data.get("content")
|
||||
encoding: Final = data.get("encoding", "")
|
||||
if content and encoding == "base64":
|
||||
|
|
@ -208,7 +276,7 @@ class GitLabClient:
|
|||
return []
|
||||
resp.raise_for_status()
|
||||
|
||||
data: Final = resp.json() or []
|
||||
data: Final = _tree_entries(resp)
|
||||
files: Final[list[str]] = []
|
||||
for item in data:
|
||||
if item.get("type") == "blob":
|
||||
|
|
@ -229,13 +297,13 @@ class GitLabClient:
|
|||
raise Exception("Authentication failed. Check your GitLab token and auth_method.")
|
||||
raise Exception(f"Failed to list files in '{directory_path}': {e}")
|
||||
|
||||
def get_repository_info(self) -> dict[str, Any]:
|
||||
def get_repository_info(self) -> Mapping[str, object]:
|
||||
"""Get information about the project/repository."""
|
||||
url: Final = f"{self.base_url}/projects/{self._project_enc}"
|
||||
try:
|
||||
resp: Final = self.http_handler.get(url, headers=self.headers)
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
return _project_info(resp)
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to get repository info: {e}")
|
||||
|
||||
|
|
@ -247,18 +315,18 @@ class GitLabClient:
|
|||
except Exception:
|
||||
return False
|
||||
|
||||
def get_branches(self) -> list[dict[str, Any]]:
|
||||
def get_branches(self) -> list[GitLabBranch]:
|
||||
"""Get list of branches in the repository."""
|
||||
url: Final = f"{self.base_url}/projects/{self._project_enc}/repository/branches"
|
||||
try:
|
||||
resp: Final = self.http_handler.get(url, headers=self.headers)
|
||||
resp.raise_for_status()
|
||||
data: Final = resp.json()
|
||||
data: Final = _branch_entries(resp)
|
||||
return data if isinstance(data, list) else []
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to get branches: {e}")
|
||||
|
||||
def get_file_metadata(self, file_path: str, *, ref: str | None = None) -> dict[str, Any] | None:
|
||||
def get_file_metadata(self, file_path: str, *, ref: str | None = None) -> GitLabFileMetadata | None:
|
||||
"""
|
||||
Get minimal metadata about a file via RAW endpoint headers at a given ref.
|
||||
|
||||
|
|
|
|||
|
|
@ -2,10 +2,12 @@
|
|||
GitLab prompt manager with configurable prompts folder.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, TypeVar
|
||||
|
||||
from jinja2 import DictLoader, select_autoescape
|
||||
from jinja2.sandbox import ImmutableSandboxedEnvironment
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.integrations.custom_prompt_management import CustomPromptManagement
|
||||
|
||||
|
|
@ -24,6 +26,19 @@ from litellm.types.utils import StandardCallbackDynamicParams
|
|||
|
||||
GITLAB_PREFIX: Final = "gitlab::"
|
||||
|
||||
_ResponseT = TypeVar("_ResponseT")
|
||||
|
||||
|
||||
class GitLabCachedPrompt(TypedDict):
|
||||
id: ReadOnly[str]
|
||||
path: ReadOnly[str]
|
||||
content: ReadOnly[str]
|
||||
metadata: ReadOnly[Mapping[str, object]]
|
||||
model: ReadOnly[str | None]
|
||||
temperature: ReadOnly[float | None]
|
||||
max_tokens: ReadOnly[int | None]
|
||||
optional_params: ReadOnly[Mapping[str, object]]
|
||||
|
||||
|
||||
def encode_prompt_id(raw_id: str) -> str:
|
||||
"""Convert GitLab path IDs like 'invoice/extract' → 'gitlab::invoice::extract'"""
|
||||
|
|
@ -206,7 +221,7 @@ class GitLabTemplateManager:
|
|||
result[key] = value.strip("\"'")
|
||||
return result
|
||||
|
||||
def render_template(self, template_id: str, variables: dict[str, Any] | None = None) -> str:
|
||||
def render_template(self, template_id: str, variables: Mapping[str, object] | None = None) -> str:
|
||||
if template_id not in self.prompts:
|
||||
raise ValueError(f"Template '{template_id}' not found")
|
||||
template: Final = self.prompts[template_id]
|
||||
|
|
@ -313,7 +328,7 @@ class GitLabPromptManager(CustomPromptManagement):
|
|||
def get_prompt_template(
|
||||
self,
|
||||
prompt_id: str,
|
||||
prompt_variables: dict[str, Any] | None = None,
|
||||
prompt_variables: Mapping[str, object] | None = None,
|
||||
*,
|
||||
ref: str | None = None,
|
||||
) -> tuple[str, dict[str, Any]]:
|
||||
|
|
@ -338,13 +353,13 @@ class GitLabPromptManager(CustomPromptManagement):
|
|||
self,
|
||||
user_id: str | None,
|
||||
messages: list[AllMessageValues],
|
||||
function_call: dict[str, Any] | str | None = None,
|
||||
litellm_params: dict[str, Any] | None = None,
|
||||
function_call: Mapping[str, object] | str | None = None,
|
||||
litellm_params: dict[str, object] | None = None,
|
||||
prompt_id: str | None = None,
|
||||
prompt_variables: dict[str, Any] | None = None,
|
||||
prompt_variables: Mapping[str, object] | None = None,
|
||||
prompt_version: str | None = None,
|
||||
**kwargs,
|
||||
) -> tuple[list[AllMessageValues], dict[str, Any] | None]:
|
||||
) -> tuple[list[AllMessageValues], dict[str, object] | None]:
|
||||
if not prompt_id:
|
||||
return messages, litellm_params
|
||||
try:
|
||||
|
|
@ -377,9 +392,9 @@ class GitLabPromptManager(CustomPromptManagement):
|
|||
|
||||
return final_messages, litellm_params
|
||||
except Exception as e:
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
litellm._logging.verbose_proxy_logger.error("Error in GitLab prompt pre_call_hook: %s", e)
|
||||
verbose_proxy_logger.error("Error in GitLab prompt pre_call_hook: %s", e)
|
||||
return messages, litellm_params
|
||||
|
||||
def _parse_prompt_to_messages(self, prompt_content: str) -> list[AllMessageValues]:
|
||||
|
|
@ -435,14 +450,14 @@ class GitLabPromptManager(CustomPromptManagement):
|
|||
def post_call_hook(
|
||||
self,
|
||||
user_id: str | None,
|
||||
response: Any,
|
||||
response: _ResponseT,
|
||||
input_messages: list[AllMessageValues],
|
||||
function_call: dict[str, Any] | str | None = None,
|
||||
litellm_params: dict[str, Any] | None = None,
|
||||
function_call: Mapping[str, object] | str | None = None,
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
prompt_id: str | None = None,
|
||||
prompt_variables: dict[str, Any] | None = None,
|
||||
prompt_variables: Mapping[str, object] | None = None,
|
||||
**kwargs,
|
||||
) -> Any:
|
||||
) -> _ResponseT:
|
||||
return response
|
||||
|
||||
def get_available_prompts(self) -> list[str]:
|
||||
|
|
@ -498,7 +513,7 @@ class GitLabPromptManager(CustomPromptManagement):
|
|||
messages: Final = self._parse_prompt_to_messages(rendered_prompt)
|
||||
template_model: Final = prompt_metadata.get("model")
|
||||
|
||||
optional_params: Final[dict[str, Any]] = {}
|
||||
optional_params: Final[dict[str, object]] = {}
|
||||
for param in [
|
||||
"temperature",
|
||||
"max_tokens",
|
||||
|
|
@ -658,14 +673,14 @@ class GitLabPromptCache:
|
|||
self.template_manager: GitLabTemplateManager = self.prompt_manager.prompt_manager
|
||||
|
||||
# In-memory stores
|
||||
self._by_file: dict[str, dict[str, Any]] = {}
|
||||
self._by_id: dict[str, dict[str, Any]] = {}
|
||||
self._by_file: dict[str, GitLabCachedPrompt] = {}
|
||||
self._by_id: dict[str, GitLabCachedPrompt] = {}
|
||||
|
||||
# -------------------------
|
||||
# Public API
|
||||
# -------------------------
|
||||
|
||||
def load_all(self, *, recursive: bool = True) -> dict[str, dict[str, Any]]:
|
||||
def load_all(self, *, recursive: bool = True) -> dict[str, GitLabCachedPrompt]:
|
||||
"""
|
||||
Scan GitLab for all .prompt files under prompts_path, load and parse each,
|
||||
and return the mapping of repo file path -> JSON-like dict.
|
||||
|
|
@ -695,7 +710,7 @@ class GitLabPromptCache:
|
|||
|
||||
return self._by_id
|
||||
|
||||
def reload(self, *, recursive: bool = True) -> dict[str, dict[str, Any]]:
|
||||
def reload(self, *, recursive: bool = True) -> dict[str, GitLabCachedPrompt]:
|
||||
"""Clear the cache and re-load from GitLab."""
|
||||
self._by_file.clear()
|
||||
self._by_id.clear()
|
||||
|
|
@ -709,11 +724,11 @@ class GitLabPromptCache:
|
|||
"""Return the template IDs (relative to prompts_path, without extension) currently cached."""
|
||||
return list(self._by_id.keys())
|
||||
|
||||
def get_by_file(self, file_path: str) -> dict[str, Any] | None:
|
||||
def get_by_file(self, file_path: str) -> GitLabCachedPrompt | None:
|
||||
"""Get a cached prompt JSON by repo file path."""
|
||||
return self._by_file.get(file_path)
|
||||
|
||||
def get_by_id(self, prompt_id: str) -> dict[str, Any] | None:
|
||||
def get_by_id(self, prompt_id: str) -> GitLabCachedPrompt | None:
|
||||
"""Get a cached prompt JSON by prompt ID (relative to prompts_path)."""
|
||||
if prompt_id in self._by_id:
|
||||
return self._by_id[prompt_id]
|
||||
|
|
@ -728,7 +743,7 @@ class GitLabPromptCache:
|
|||
# Internals
|
||||
# -------------------------
|
||||
|
||||
def _template_to_json(self, prompt_id: str, tmpl: GitLabPromptTemplate) -> dict[str, Any]:
|
||||
def _template_to_json(self, prompt_id: str, tmpl: GitLabPromptTemplate) -> GitLabCachedPrompt:
|
||||
"""
|
||||
Normalize a GitLabPromptTemplate into a JSON-like dict that is easy to serialize.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -89,7 +89,7 @@ def _extract_cache_read_input_tokens(usage_obj) -> int:
|
|||
|
||||
# Check prompt_tokens_details.cached_tokens (used by Gemini and other providers)
|
||||
if hasattr(usage_obj, "prompt_tokens_details"):
|
||||
prompt_tokens_details: Final = getattr(usage_obj, "prompt_tokens_details", None)
|
||||
prompt_tokens_details: Final[object] = getattr(usage_obj, "prompt_tokens_details", None)
|
||||
if prompt_tokens_details is not None and hasattr(prompt_tokens_details, "cached_tokens"):
|
||||
cached_tokens: Final = getattr(prompt_tokens_details, "cached_tokens", None)
|
||||
if cached_tokens is not None and isinstance(cached_tokens, (int, float)) and cached_tokens > 0:
|
||||
|
|
@ -623,9 +623,16 @@ class LangFuseLogger:
|
|||
)
|
||||
|
||||
# Apply custom masking function if provided
|
||||
if masking_function is not None and callable(masking_function):
|
||||
input = self._apply_masking_function(input, masking_function)
|
||||
output = self._apply_masking_function(output, masking_function)
|
||||
masked_input: Final[object] = (
|
||||
self._apply_masking_function(input, masking_function)
|
||||
if masking_function is not None and callable(masking_function)
|
||||
else input
|
||||
)
|
||||
masked_output: Final[object] = (
|
||||
self._apply_masking_function(output, masking_function)
|
||||
if masking_function is not None and callable(masking_function)
|
||||
else output
|
||||
)
|
||||
|
||||
clean_metadata = redact_user_api_key_info(metadata=clean_metadata)
|
||||
|
||||
|
|
@ -651,15 +658,15 @@ class LangFuseLogger:
|
|||
|
||||
# Special keys that are found in the function arguments and not the metadata
|
||||
if "input" in update_trace_keys:
|
||||
trace_params["input"] = input if not mask_input else "redacted-by-litellm"
|
||||
trace_params["input"] = masked_input if not mask_input else "redacted-by-litellm"
|
||||
if "output" in update_trace_keys:
|
||||
trace_params["output"] = output if not mask_output else "redacted-by-litellm"
|
||||
trace_params["output"] = masked_output if not mask_output else "redacted-by-litellm"
|
||||
else: # don't overwrite an existing trace
|
||||
trace_params = {
|
||||
"id": trace_id,
|
||||
"name": trace_name,
|
||||
"session_id": session_id,
|
||||
"input": input if not mask_input else "redacted-by-litellm",
|
||||
"input": masked_input if not mask_input else "redacted-by-litellm",
|
||||
"version": clean_metadata.pop(
|
||||
"trace_version", clean_metadata.get("version", None)
|
||||
), # If provided just version, it will applied to the trace as well, if applied a trace version it will take precedence
|
||||
|
|
@ -669,9 +676,9 @@ class LangFuseLogger:
|
|||
trace_params[key.replace("trace_", "")] = clean_metadata.pop(key, None)
|
||||
|
||||
if level == "ERROR":
|
||||
trace_params["status_message"] = output
|
||||
trace_params["status_message"] = masked_output
|
||||
else:
|
||||
trace_params["output"] = output if not mask_output else "redacted-by-litellm"
|
||||
trace_params["output"] = masked_output if not mask_output else "redacted-by-litellm"
|
||||
|
||||
if debug is True or (isinstance(debug, str) and debug.lower() == "true"):
|
||||
debug_metadata: Final = {
|
||||
|
|
@ -708,7 +715,7 @@ class LangFuseLogger:
|
|||
("aws_region_name", aws_region_name, bool(aws_region_name)),
|
||||
("cache_hit", kwargs.get("cache_hit") or False, self._supports_tags() and "cache_hit" in kwargs),
|
||||
)
|
||||
enrichments: Final[Mapping[str, Any]] = {
|
||||
enrichments: Final[Mapping[str, object]] = {
|
||||
key: value for key, value, include in candidate_enrichments if include
|
||||
}
|
||||
|
||||
|
|
@ -802,8 +809,8 @@ class LangFuseLogger:
|
|||
"end_time": end_time,
|
||||
"model": model_name,
|
||||
"model_parameters": optional_params,
|
||||
"input": input if not mask_input else "redacted-by-litellm",
|
||||
"output": output if not mask_output else "redacted-by-litellm",
|
||||
"input": masked_input if not mask_input else "redacted-by-litellm",
|
||||
"output": masked_output if not mask_output else "redacted-by-litellm",
|
||||
"usage": usage,
|
||||
"usage_details": usage_details,
|
||||
"metadata": {
|
||||
|
|
@ -825,8 +832,8 @@ class LangFuseLogger:
|
|||
prompt_management_metadata=prompt_management_metadata,
|
||||
langfuse_client=self.Langfuse,
|
||||
)
|
||||
if output is not None and isinstance(output, str) and level == "ERROR":
|
||||
generation_params["status_message"] = output
|
||||
if masked_output is not None and isinstance(masked_output, str) and level == "ERROR":
|
||||
generation_params["status_message"] = masked_output
|
||||
|
||||
if self._supports_completion_start_time():
|
||||
generation_params["completion_start_time"] = kwargs.get("completion_start_time", None)
|
||||
|
|
@ -935,7 +942,7 @@ class LangFuseLogger:
|
|||
return Version(self.langfuse_sdk_version) >= Version("2.7.3")
|
||||
|
||||
@staticmethod
|
||||
def _apply_masking_function(data: Any, masking_function: Callable[[Any], Any]) -> Any:
|
||||
def _apply_masking_function(data: object, masking_function: Callable[[object], object]) -> object:
|
||||
"""
|
||||
Apply a masking function to data, handling different data types.
|
||||
|
||||
|
|
@ -1049,7 +1056,7 @@ def _add_prompt_to_generation_params(
|
|||
generation_params: dict,
|
||||
clean_metadata: dict,
|
||||
prompt_management_metadata: StandardLoggingPromptManagementMetadata | None,
|
||||
langfuse_client: Any,
|
||||
langfuse_client: object,
|
||||
) -> dict:
|
||||
from langfuse import Langfuse
|
||||
from langfuse.model import (
|
||||
|
|
|
|||
|
|
@ -4,9 +4,12 @@ Opik Logger that logs LLM events to an Opik server
|
|||
|
||||
import asyncio
|
||||
import traceback
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
from typing import Any, Final
|
||||
|
||||
from typing_extensions import ReadOnly, TypedDict, Unpack
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_batch_logger import CustomBatchLogger
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
|
|
@ -23,7 +26,7 @@ except Exception:
|
|||
opik_client = None
|
||||
|
||||
|
||||
def _should_skip_event(kwargs: dict[str, Any]) -> bool:
|
||||
def _should_skip_event(kwargs: Mapping[str, object]) -> bool:
|
||||
"""Check if event should be skipped due to missing standard_logging_object."""
|
||||
if kwargs.get("standard_logging_object") is None:
|
||||
verbose_logger.debug("OpikLogger skipping event; no standard_logging_object found")
|
||||
|
|
@ -31,12 +34,24 @@ def _should_skip_event(kwargs: dict[str, Any]) -> bool:
|
|||
return False
|
||||
|
||||
|
||||
class _OpikLoggerKwargs(TypedDict, total=False):
|
||||
"""Constructor options accepted by ``OpikLogger``."""
|
||||
|
||||
project_name: ReadOnly[str | None]
|
||||
url: ReadOnly[str | None]
|
||||
api_key: ReadOnly[str | None]
|
||||
workspace: ReadOnly[str | None]
|
||||
batch_size: ReadOnly[int | None]
|
||||
flush_interval: ReadOnly[int | None]
|
||||
max_queue_size: ReadOnly[int | None]
|
||||
|
||||
|
||||
class OpikLogger(CustomBatchLogger):
|
||||
"""
|
||||
Opik Logger for logging events to an Opik Server
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
def __init__(self, **kwargs: Unpack[_OpikLoggerKwargs]) -> None:
|
||||
self.async_httpx_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback)
|
||||
self.sync_httpx_client = _get_httpx_client()
|
||||
|
||||
|
|
@ -95,7 +110,7 @@ class OpikLogger(CustomBatchLogger):
|
|||
|
||||
async def async_log_success_event(
|
||||
self,
|
||||
kwargs: dict[str, Any],
|
||||
kwargs: dict[str, object],
|
||||
response_obj: Any,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
|
|
@ -163,7 +178,7 @@ class OpikLogger(CustomBatchLogger):
|
|||
except Exception as e:
|
||||
verbose_logger.exception("OpikLogger failed to log success event - %s\n%s", e, traceback.format_exc())
|
||||
|
||||
def _sync_send(self, url: str, headers: dict[str, str], batch: dict[str, Any]) -> None:
|
||||
def _sync_send(self, url: str, headers: dict[str, str], batch: dict[str, object]) -> None:
|
||||
try:
|
||||
response: Final = self.sync_httpx_client.post(
|
||||
url=url,
|
||||
|
|
@ -178,7 +193,7 @@ class OpikLogger(CustomBatchLogger):
|
|||
|
||||
def log_success_event(
|
||||
self,
|
||||
kwargs: dict[str, Any],
|
||||
kwargs: dict[str, object],
|
||||
response_obj: Any,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
|
|
@ -247,7 +262,7 @@ class OpikLogger(CustomBatchLogger):
|
|||
except Exception as e:
|
||||
verbose_logger.exception("OpikLogger failed to log success event - %s\n%s", e, traceback.format_exc())
|
||||
|
||||
async def _submit_batch(self, url: str, headers: dict[str, str], batch: dict[str, Any]) -> None:
|
||||
async def _submit_batch(self, url: str, headers: dict[str, str], batch: dict[str, object]) -> None:
|
||||
try:
|
||||
response: Final = await self.async_httpx_client.post(
|
||||
url=url,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
"""Data extraction functions for Opik payload building."""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm import _logging
|
||||
|
|
@ -35,8 +36,8 @@ def normalize_provider_name(provider: str | None) -> str | None:
|
|||
|
||||
|
||||
def extract_opik_metadata(
|
||||
litellm_metadata: dict[str, Any],
|
||||
standard_logging_metadata: dict[str, Any],
|
||||
litellm_metadata: Mapping[str, Any],
|
||||
standard_logging_metadata: Mapping[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Merge Opik metadata from three sources in increasing priority order:
|
||||
|
|
@ -97,7 +98,7 @@ def extract_span_identifiers(
|
|||
|
||||
|
||||
def extract_tags(
|
||||
opik_metadata: dict[str, Any],
|
||||
opik_metadata: Mapping[str, Any],
|
||||
custom_llm_provider: str | None,
|
||||
) -> list[str]:
|
||||
"""
|
||||
|
|
@ -122,7 +123,7 @@ def apply_proxy_header_overrides(
|
|||
project_name: str,
|
||||
tags: list[str],
|
||||
thread_id: str | None,
|
||||
proxy_headers: dict[str, Any],
|
||||
proxy_headers: Mapping[str, str],
|
||||
) -> tuple[str, list[str], str | None]:
|
||||
"""
|
||||
Apply overrides from proxy request headers (opik_* prefix).
|
||||
|
|
@ -148,7 +149,7 @@ def apply_proxy_header_overrides(
|
|||
thread_id = value
|
||||
elif param_key == "tags":
|
||||
try:
|
||||
parsed_tags = json.loads(value)
|
||||
parsed_tags: object = json.loads(value)
|
||||
if isinstance(parsed_tags, list):
|
||||
tags.extend(parsed_tags)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
|
|
@ -158,11 +159,11 @@ def apply_proxy_header_overrides(
|
|||
|
||||
|
||||
def extract_and_build_metadata(
|
||||
opik_metadata: dict[str, Any],
|
||||
standard_logging_metadata: dict[str, Any],
|
||||
standard_logging_object: dict[str, Any],
|
||||
litellm_kwargs: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
opik_metadata: Mapping[str, object],
|
||||
standard_logging_metadata: Mapping[str, object],
|
||||
standard_logging_object: Mapping[str, object],
|
||||
litellm_kwargs: Mapping[str, object],
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Build the complete metadata dictionary from all available sources.
|
||||
|
||||
|
|
|
|||
|
|
@ -62,6 +62,8 @@ class GenAIMapper:
|
|||
GenAI.RESPONSE_TIME_TO_FIRST_CHUNK: lambda d: d.time_to_first_chunk_seconds,
|
||||
GenAI.USAGE_INPUT_TOKENS: lambda d: d.usage.input_tokens,
|
||||
GenAI.USAGE_OUTPUT_TOKENS: lambda d: d.usage.output_tokens,
|
||||
GenAI.USAGE_CACHE_CREATION_INPUT_TOKENS: lambda d: d.usage.cache_creation_input_tokens,
|
||||
GenAI.USAGE_CACHE_READ_INPUT_TOKENS: lambda d: d.usage.cache_read_input_tokens,
|
||||
Error.TYPE: lambda d: d.error.error_type if d.error else None,
|
||||
Server.ADDRESS: lambda d: d.server.address if d.server else None,
|
||||
Server.PORT: lambda d: d.server.port if d.server else None,
|
||||
|
|
|
|||
|
|
@ -95,6 +95,22 @@ class LLMUsage:
|
|||
input_tokens: int | None = None
|
||||
output_tokens: int | None = None
|
||||
total_tokens: int | None = None
|
||||
cache_creation_input_tokens: int | None = None
|
||||
cache_read_input_tokens: int | None = None
|
||||
|
||||
@classmethod
|
||||
def from_standard_logging_payload(cls, payload: StandardLoggingPayload) -> LLMUsage:
|
||||
# Cache token counts only exist on the raw provider usage object under metadata
|
||||
metadata: Final[Mapping[str, object]] = payload.get("metadata") or {}
|
||||
raw_usage: Final = metadata.get("usage_object")
|
||||
usage_object: Final[Mapping[str, object]] = raw_usage if isinstance(raw_usage, Mapping) else {}
|
||||
return cls(
|
||||
input_tokens=as_int(payload.get("prompt_tokens")),
|
||||
output_tokens=as_int(payload.get("completion_tokens")),
|
||||
total_tokens=as_int(payload.get("total_tokens")),
|
||||
cache_creation_input_tokens=as_int(usage_object.get("cache_creation_input_tokens")),
|
||||
cache_read_input_tokens=as_int(usage_object.get("cache_read_input_tokens")),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
|
|
@ -363,11 +379,7 @@ class LLMCallSpanData:
|
|||
response_model=context.response_model,
|
||||
response_id=as_str(response.get("id")),
|
||||
request_params=LLMRequestParams.from_model_parameters(params),
|
||||
usage=LLMUsage(
|
||||
input_tokens=as_int(payload.get("prompt_tokens")),
|
||||
output_tokens=as_int(payload.get("completion_tokens")),
|
||||
total_tokens=as_int(payload.get("total_tokens")),
|
||||
),
|
||||
usage=LLMUsage.from_standard_logging_payload(payload),
|
||||
finish_reasons=finish_reasons,
|
||||
error=_parse_error(payload),
|
||||
response_cost=as_float(payload.get("response_cost")),
|
||||
|
|
|
|||
|
|
@ -110,6 +110,8 @@ class GenAI:
|
|||
# usage
|
||||
USAGE_INPUT_TOKENS: Final = "gen_ai.usage.input_tokens"
|
||||
USAGE_OUTPUT_TOKENS: Final = "gen_ai.usage.output_tokens"
|
||||
USAGE_CACHE_CREATION_INPUT_TOKENS: Final = "gen_ai.usage.cache_creation.input_tokens"
|
||||
USAGE_CACHE_READ_INPUT_TOKENS: Final = "gen_ai.usage.cache_read.input_tokens"
|
||||
# content (opt-in, gated by capture mode)
|
||||
INPUT_MESSAGES: Final = "gen_ai.input.messages"
|
||||
OUTPUT_MESSAGES: Final = "gen_ai.output.messages"
|
||||
|
|
|
|||
|
|
@ -11,9 +11,10 @@ identical metrics. The attribute cardinality filter is reused from v1 by import
|
|||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Any, Final, TypeAlias
|
||||
from typing import Any, Final, Literal, Protocol, TypeAlias
|
||||
|
||||
from opentelemetry.metrics import Histogram, Meter
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -151,6 +152,29 @@ METRIC_ATTRIBUTE_CEILING: Final[frozenset[str]] = frozenset(
|
|||
BOUNDED_HIDDEN_PARAM_KEYS: Final[tuple[str, ...]] = ("model_id",)
|
||||
|
||||
|
||||
class _TokenUsage(TypedDict, total=False):
|
||||
"""The token counts a response's ``usage`` carries, as the recorder reads them."""
|
||||
|
||||
prompt_tokens: ReadOnly[int]
|
||||
completion_tokens: ReadOnly[int]
|
||||
|
||||
|
||||
class _ResponseView(Protocol):
|
||||
"""The one read the recorder makes on a litellm response object."""
|
||||
|
||||
def get(self, key: Literal["usage"], /) -> _TokenUsage | None: ...
|
||||
|
||||
|
||||
class _MetricKwargs(TypedDict, total=False):
|
||||
"""The logging kwargs the recorder reads directly."""
|
||||
|
||||
call_type: ReadOnly[str | None]
|
||||
litellm_params: ReadOnly[Mapping[str, object] | None]
|
||||
response_cost: ReadOnly[float | None]
|
||||
completion_start_time: ReadOnly[datetime | float | str | None]
|
||||
api_call_start_time: ReadOnly[datetime | float | str | None]
|
||||
|
||||
|
||||
def resolve_error_type(kwargs: Mapping[str, Any]) -> str:
|
||||
"""The ``error.type`` value for a failed request.
|
||||
|
||||
|
|
@ -192,8 +216,8 @@ class GenAIMetricRecorder:
|
|||
|
||||
def record(
|
||||
self,
|
||||
kwargs: Mapping[str, Any],
|
||||
response_obj: Any,
|
||||
kwargs: _MetricKwargs,
|
||||
response_obj: _ResponseView | None,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
) -> None:
|
||||
|
|
@ -218,7 +242,7 @@ class GenAIMetricRecorder:
|
|||
|
||||
def record_failure(
|
||||
self,
|
||||
kwargs: Mapping[str, Any],
|
||||
kwargs: _MetricKwargs,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
) -> None:
|
||||
|
|
@ -342,7 +366,7 @@ class GenAIMetricRecorder:
|
|||
# Per-metric recording
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
def _record_token_usage(self, response_obj: Any, common_attrs: dict) -> None:
|
||||
def _record_token_usage(self, response_obj: _ResponseView | None, common_attrs: dict) -> None:
|
||||
if not response_obj:
|
||||
return
|
||||
usage: Final = response_obj.get("usage")
|
||||
|
|
@ -353,7 +377,7 @@ class GenAIMetricRecorder:
|
|||
self._metrics.token_usage.record(usage.get("prompt_tokens", 0), attributes=in_attrs)
|
||||
self._metrics.token_usage.record(usage.get("completion_tokens", 0), attributes=out_attrs)
|
||||
|
||||
def _record_time_to_first_token(self, kwargs: Mapping[str, Any], common_attrs: dict) -> None:
|
||||
def _record_time_to_first_token(self, kwargs: _MetricKwargs, common_attrs: dict) -> None:
|
||||
time_to_first_chunk: Final = time_to_first_chunk_seconds(kwargs)
|
||||
if time_to_first_chunk is None:
|
||||
return
|
||||
|
|
@ -361,15 +385,14 @@ class GenAIMetricRecorder:
|
|||
|
||||
def _record_time_per_output_token(
|
||||
self,
|
||||
kwargs: Mapping[str, Any],
|
||||
response_obj: Any,
|
||||
kwargs: _MetricKwargs,
|
||||
response_obj: _ResponseView | None,
|
||||
end_time: datetime,
|
||||
duration_s: float,
|
||||
common_attrs: dict,
|
||||
) -> None:
|
||||
completion_tokens = None
|
||||
if response_obj and (usage := response_obj.get("usage")):
|
||||
completion_tokens = usage.get("completion_tokens")
|
||||
usage: Final = response_obj.get("usage") if response_obj else None
|
||||
completion_tokens: Final = usage.get("completion_tokens") if usage else None
|
||||
if completion_tokens is None or completion_tokens <= 0:
|
||||
return
|
||||
|
||||
|
|
|
|||
|
|
@ -12,7 +12,10 @@ For batching specific details see CustomBatchLogger class
|
|||
import asyncio
|
||||
import atexit
|
||||
import os
|
||||
from typing import Any, Final
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Final
|
||||
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
|
|
@ -34,6 +37,21 @@ from litellm.types.integrations.posthog import (
|
|||
from litellm.types.utils import StandardCallbackDynamicParams, StandardLoggingPayload
|
||||
|
||||
|
||||
class PostHogBatchPayload(TypedDict):
|
||||
api_key: ReadOnly[str]
|
||||
batch: ReadOnly[Sequence[PostHogEventPayload]]
|
||||
|
||||
|
||||
class PostHogLiteLLMParams(TypedDict, total=False):
|
||||
metadata: ReadOnly[Mapping[str, object]]
|
||||
|
||||
|
||||
class PostHogLogKwargs(TypedDict, total=False):
|
||||
standard_logging_object: ReadOnly[StandardLoggingPayload]
|
||||
standard_callback_dynamic_params: ReadOnly[StandardCallbackDynamicParams]
|
||||
litellm_params: ReadOnly[PostHogLiteLLMParams]
|
||||
|
||||
|
||||
class PostHogLogger(CustomBatchLogger):
|
||||
def __init__(self, **kwargs):
|
||||
"""
|
||||
|
|
@ -137,7 +155,7 @@ class PostHogLogger(CustomBatchLogger):
|
|||
if len(self.log_queue) >= self.batch_size:
|
||||
await self.flush_queue()
|
||||
|
||||
def create_posthog_event_payload(self, kwargs: dict[str, Any]) -> PostHogEventPayload:
|
||||
def create_posthog_event_payload(self, kwargs: PostHogLogKwargs) -> PostHogEventPayload:
|
||||
"""
|
||||
Helper function to create a PostHog event payload for logging
|
||||
|
||||
|
|
@ -171,11 +189,11 @@ class PostHogLogger(CustomBatchLogger):
|
|||
def _create_posthog_properties(
|
||||
self,
|
||||
standard_logging_object: StandardLoggingPayload,
|
||||
kwargs: dict[str, Any],
|
||||
kwargs: PostHogLogKwargs,
|
||||
event_name: str,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""Create PostHog properties following LLM Analytics spec"""
|
||||
properties: Final = {}
|
||||
properties: Final[dict[str, object]] = {}
|
||||
|
||||
# Core model information
|
||||
properties["$ai_model"] = self._safe_get(standard_logging_object, "model", "")
|
||||
|
|
@ -211,16 +229,19 @@ class PostHogLogger(CustomBatchLogger):
|
|||
properties["$ai_error"] = error_str
|
||||
|
||||
# Add trace properties
|
||||
self._add_trace_properties(properties, kwargs)
|
||||
self._add_trace_properties(properties, standard_logging_object, kwargs)
|
||||
|
||||
# Add custom metadata fields
|
||||
self._add_custom_metadata_properties(properties, kwargs)
|
||||
|
||||
return properties
|
||||
|
||||
def _add_trace_properties(self, properties: dict[str, Any], kwargs: dict[str, Any]):
|
||||
standard_logging_object: Final = self._safe_get(kwargs, "standard_logging_object", {})
|
||||
|
||||
def _add_trace_properties(
|
||||
self,
|
||||
properties: dict[str, object],
|
||||
standard_logging_object: StandardLoggingPayload,
|
||||
kwargs: PostHogLogKwargs,
|
||||
) -> None:
|
||||
trace_id: Final = self._safe_get(standard_logging_object, "trace_id", self._safe_uuid())
|
||||
properties["$ai_trace_id"] = trace_id
|
||||
|
||||
|
|
@ -232,7 +253,7 @@ class PostHogLogger(CustomBatchLogger):
|
|||
if parent_id:
|
||||
properties["$ai_parent_id"] = parent_id
|
||||
|
||||
def _add_custom_metadata_properties(self, properties: dict[str, Any], kwargs: dict[str, Any]):
|
||||
def _add_custom_metadata_properties(self, properties: dict[str, object], kwargs: PostHogLogKwargs) -> None:
|
||||
"""Add custom metadata fields to PostHog properties"""
|
||||
metadata: Final = self._extract_metadata(kwargs)
|
||||
if not isinstance(metadata, dict):
|
||||
|
|
@ -277,7 +298,7 @@ class PostHogLogger(CustomBatchLogger):
|
|||
if key not in litellm_internal_fields:
|
||||
properties[key] = value
|
||||
|
||||
def _get_distinct_id(self, standard_logging_object: StandardLoggingPayload, kwargs: dict[str, Any]) -> str:
|
||||
def _get_distinct_id(self, standard_logging_object: StandardLoggingPayload, kwargs: PostHogLogKwargs) -> str:
|
||||
metadata: Final = self._extract_metadata(kwargs)
|
||||
user_id: Final = self._safe_get(metadata, "user_id")
|
||||
if user_id:
|
||||
|
|
@ -291,7 +312,7 @@ class PostHogLogger(CustomBatchLogger):
|
|||
|
||||
return self._safe_uuid()
|
||||
|
||||
def _get_credentials_for_request(self, kwargs: dict[str, Any]) -> tuple[str | None, str | None]:
|
||||
def _get_credentials_for_request(self, kwargs: PostHogLogKwargs) -> tuple[str | None, str | None]:
|
||||
"""
|
||||
Get PostHog credentials for this request.
|
||||
|
||||
|
|
@ -334,7 +355,7 @@ class PostHogLogger(CustomBatchLogger):
|
|||
verbose_logger.debug("[POSTHOG MOCK] Mock mode enabled - API calls will be intercepted")
|
||||
|
||||
# Group events by credentials for batch sending
|
||||
batches_by_credentials: Final[dict[tuple[str, str], list]] = {}
|
||||
batches_by_credentials: Final[dict[tuple[str, str], list[PostHogEventPayload]]] = {}
|
||||
for item in self.log_queue:
|
||||
key = (item["api_key"], item["api_url"])
|
||||
if key not in batches_by_credentials:
|
||||
|
|
@ -380,18 +401,19 @@ class PostHogLogger(CustomBatchLogger):
|
|||
verbose_logger.error("PostHog: Failed to initialize async components: %s", e)
|
||||
raise
|
||||
|
||||
def _extract_metadata(self, kwargs: dict[str, Any]) -> dict[str, Any]:
|
||||
litellm_params: Final = kwargs.get("litellm_params", {}) or {}
|
||||
return litellm_params.get("metadata", {}) or {}
|
||||
def _extract_metadata(self, kwargs: PostHogLogKwargs) -> Mapping[str, object]:
|
||||
litellm_params: Final[PostHogLiteLLMParams] = kwargs.get("litellm_params", {}) or {}
|
||||
metadata: Final[Mapping[str, object]] = litellm_params.get("metadata", {}) or {}
|
||||
return metadata
|
||||
|
||||
def _safe_uuid(self) -> str:
|
||||
return str(uuid.uuid4())
|
||||
|
||||
def _create_posthog_payload(self, events: list, api_key: str) -> dict[str, Any]:
|
||||
def _create_posthog_payload(self, events: Sequence[PostHogEventPayload], api_key: str) -> PostHogBatchPayload:
|
||||
return {"api_key": api_key, "batch": events}
|
||||
|
||||
def _safe_get(self, obj: Any, key: str, default: Any = None) -> Any:
|
||||
if obj is None or not hasattr(obj, "get"):
|
||||
def _safe_get(self, obj: Mapping[str, object] | None, key: str, default: object = None) -> object:
|
||||
if not isinstance(obj, Mapping):
|
||||
return default
|
||||
return obj.get(key, default)
|
||||
|
||||
|
|
@ -412,7 +434,7 @@ class PostHogLogger(CustomBatchLogger):
|
|||
|
||||
try:
|
||||
# Group events by credentials (same logic as async_send_batch)
|
||||
batches_by_credentials: Final[dict[tuple[str, str], list]] = {}
|
||||
batches_by_credentials: Final[dict[tuple[str, str], list[PostHogEventPayload]]] = {}
|
||||
for item in self.log_queue:
|
||||
key = (item["api_key"], item["api_url"])
|
||||
if key not in batches_by_credentials:
|
||||
|
|
|
|||
|
|
@ -1,8 +1,11 @@
|
|||
"""Shadow Eval Logger: samples a shadowed key's successful LLM requests (chat completions,
|
||||
Anthropic Messages, and Responses API surfaces, each normalized to chat shape), duplicates
|
||||
each against the job's other arm in a detached task (the auto-router for a forward job, the
|
||||
fixed baseline model for a reverse one), blind-judges real vs shadow, and appends one
|
||||
``LiteLLM_ShadowEvalAttempt`` row (verdict or error) as the feature's only hot-path write.
|
||||
each through every shadow arm in one detached task (each candidate auto-router for a
|
||||
forward job, the fixed baseline model for a reverse one), blind-judges real vs each arm,
|
||||
and appends one ``LiteLLM_ShadowEvalAttempt`` row per arm (verdict or error) as the
|
||||
feature's only hot-path write. A multi-router job's arms therefore score the identical
|
||||
sampled requests against the identical real responses, which is what makes their win
|
||||
rates comparable head-to-head.
|
||||
Counts, status, and spend derive from those rows at read time, so nothing can disagree
|
||||
across pods or stop races; the hook reads active jobs through a short-TTL cache."""
|
||||
|
||||
|
|
@ -498,12 +501,16 @@ def _decision_classifier_cost(metadata: Mapping[str, object]) -> float:
|
|||
return float(raw) if isinstance(raw, (int, float)) else 0.0
|
||||
|
||||
|
||||
def _request_was_routed_by(request_metadata: Mapping[str, object], router_name: str) -> bool:
|
||||
"""Whether the router under evaluation served this request, which is what decides
|
||||
the direction it belongs to. A forward job skips its own router's traffic, since
|
||||
duplicating it would compare the router to itself: guaranteed ties, judge spend for
|
||||
zero information. A reverse job samples exactly that traffic and nothing else."""
|
||||
return _routing_decision(request_metadata).get("router_model_name") == router_name
|
||||
def _direction_admits(request_metadata: Mapping[str, object], job: "ActiveShadowEvalJob") -> bool:
|
||||
"""Whether this request belongs to the job's direction. A forward job skips traffic
|
||||
any of its candidate routers served: duplicating a router's own request compares it
|
||||
to itself (guaranteed ties), and judging a sibling against another candidate's live
|
||||
response would score candidates against each other instead of against the incumbent.
|
||||
A reverse job samples exactly its one router's traffic and nothing else."""
|
||||
routed_by: Final = _routing_decision(request_metadata).get("router_model_name")
|
||||
if job.direction == "reverse":
|
||||
return routed_by == job.router_name
|
||||
return routed_by not in job.arm_router_names
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -546,6 +553,7 @@ class ActiveShadowEvalJob(BaseModel):
|
|||
|
||||
id: str
|
||||
router_name: str
|
||||
router_names: tuple[str, ...] = ()
|
||||
direction: ShadowEvalDirection = "forward"
|
||||
baseline_model: str | None = None
|
||||
shadow_percentage: float
|
||||
|
|
@ -567,12 +575,25 @@ class ActiveShadowEvalJob(BaseModel):
|
|||
raise ValueError("baseline_model is set for exactly the reverse jobs")
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _reverse_evaluates_one_router(self) -> "ActiveShadowEvalJob":
|
||||
"""A reverse row naming several routers is unsamplable (there is no one traffic
|
||||
slice they share) and fails closed."""
|
||||
if self.direction == "reverse" and len(self.arm_router_names) > 1:
|
||||
raise ValueError("a reverse job evaluates exactly one router")
|
||||
return self
|
||||
|
||||
@property
|
||||
def shadow_target(self) -> str:
|
||||
"""The model the duplicated arm calls: the router itself for a forward job, the
|
||||
fixed baseline for a reverse one. Total because the validator above pins
|
||||
def arm_router_names(self) -> tuple[str, ...]:
|
||||
"""The job's full router set; rows from before router_names existed hold it in
|
||||
router_name alone. The one place that reading lives on the sampling side."""
|
||||
return self.router_names or (self.router_name,)
|
||||
|
||||
def arm_target(self, arm_router: str) -> str:
|
||||
"""The model one duplicated arm calls: the candidate router itself for a forward
|
||||
job, the fixed baseline for a reverse one. Total because the validator above pins
|
||||
baseline_model to reverse jobs and only those."""
|
||||
return self.baseline_model or self.router_name
|
||||
return self.baseline_model or arm_router
|
||||
|
||||
|
||||
def _as_active_job(record: object, attempts: int, spend: float) -> ActiveShadowEvalJob | None:
|
||||
|
|
@ -592,7 +613,12 @@ _JOBS_CACHE_KEY: Final = "shadow_eval:active_jobs"
|
|||
|
||||
|
||||
class ShadowEvalLogger(CustomLogger):
|
||||
"""Fires blind pairwise shadow evaluations for keys with an active shadow-eval job."""
|
||||
"""Fires blind pairwise shadow evaluations for targets with an active shadow-eval job.
|
||||
|
||||
A job targets a virtual key, a team, or a user; a request qualifies for a job when
|
||||
any of its resolved identities (key hash, team id, user id) matches the job's
|
||||
target, so team and user jobs cover JWT-authenticated traffic, which carries no
|
||||
key hash at all."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -617,10 +643,10 @@ class ShadowEvalLogger(CustomLogger):
|
|||
# generation; the refill absorbs written rows and resets.
|
||||
self._job_starts: dict[str, int] = {} # mutable-ok: per-generation counter
|
||||
|
||||
async def _active_jobs(self) -> Mapping[str, tuple[ActiveShadowEvalJob, ...]]:
|
||||
"""Active jobs by api_key_id, cache-first. A key holds at most one job per
|
||||
direction, so the value is a collection. A DB fault returns empty without
|
||||
caching, so sampling pauses for that request and the next one retries."""
|
||||
async def _active_jobs(self) -> Mapping[tuple[str, str], tuple[ActiveShadowEvalJob, ...]]:
|
||||
"""Active jobs by (target_type, target_id), cache-first. A target holds at most
|
||||
one job per direction, so the value is a collection. A DB fault returns empty
|
||||
without caching, so sampling pauses for that request and the next one retries."""
|
||||
cached: Final = await self._jobs_cache.async_get_cache(_JOBS_CACHE_KEY)
|
||||
if cached is not None:
|
||||
return cached # pyright: ignore[reportReturnType] # cache stores exactly this mapping shape
|
||||
|
|
@ -652,10 +678,10 @@ class ShadowEvalLogger(CustomLogger):
|
|||
)
|
||||
for row in grouped or []
|
||||
}
|
||||
by_key: Final = tuple(
|
||||
by_target: Final = tuple(
|
||||
sorted(
|
||||
(
|
||||
(str(record.api_key_id), job)
|
||||
((str(record.target_type), str(record.target_id)), job)
|
||||
for record in records or []
|
||||
if (job := _as_active_job(record, *attempt_stats.get(str(record.id), (0, 0.0)))) is not None
|
||||
),
|
||||
|
|
@ -663,7 +689,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
)
|
||||
)
|
||||
jobs: Final = MappingProxyType(
|
||||
{key: tuple(job for _, job in group) for key, group in groupby(by_key, key=itemgetter(0))}
|
||||
{target: tuple(job for _, job in group) for target, group in groupby(by_target, key=itemgetter(0))}
|
||||
)
|
||||
await self._jobs_cache.async_set_cache(_JOBS_CACHE_KEY, jobs)
|
||||
self._job_starts = {} # rebind-ok: new generation, counts absorbed into the fill
|
||||
|
|
@ -691,7 +717,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
now >= job.ends_at
|
||||
or job.attempts + self._job_starts.get(job.id, 0) >= job.max_turns
|
||||
or (job.max_budget is not None and job.spend >= job.max_budget)
|
||||
or _request_was_routed_by(request_metadata, job.router_name) != (job.direction == "reverse")
|
||||
or not _direction_admits(request_metadata, job)
|
||||
):
|
||||
continue
|
||||
if not _sample_hits(request_id, job.id, job.shadow_percentage):
|
||||
|
|
@ -720,8 +746,18 @@ class ShadowEvalLogger(CustomLogger):
|
|||
if should_redact_message_logging(dict(kwargs)): # mutable-ok: predicate takes a plain dict
|
||||
return
|
||||
metadata: Final = payload.get("metadata") or _EMPTY_METADATA
|
||||
api_key_hash: Final = metadata.get("user_api_key_hash")
|
||||
if not api_key_hash:
|
||||
# Each identity the request resolved to is a candidate target; JWT-auth
|
||||
# requests carry no key hash but do carry a team and user.
|
||||
targets: Final = tuple(
|
||||
(target_type, str(value))
|
||||
for target_type, value in (
|
||||
("key", metadata.get("user_api_key_hash")),
|
||||
("team", metadata.get("user_api_key_team_id")),
|
||||
("user", metadata.get("user_api_key_user_id")),
|
||||
)
|
||||
if value
|
||||
)
|
||||
if not targets:
|
||||
return
|
||||
request_id: Final = payload.get("id") or ""
|
||||
if not request_id:
|
||||
|
|
@ -731,8 +767,11 @@ class ShadowEvalLogger(CustomLogger):
|
|||
return # only surfaces this table can normalize are comparable; unknown types fail closed
|
||||
if ops.wire_params and _request_mutating_guardrail_ran(request_metadata):
|
||||
return # the wire-body snapshot predates the rewrite; replaying it would resurrect stripped content
|
||||
active_jobs: Final = await self._active_jobs()
|
||||
eligible: Final = self._sampled_jobs(
|
||||
(await self._active_jobs()).get(str(api_key_hash), ()), request_metadata, request_id
|
||||
tuple(job for target in targets for job in active_jobs.get(target, ())),
|
||||
request_metadata,
|
||||
request_id,
|
||||
)
|
||||
if not eligible:
|
||||
return
|
||||
|
|
@ -755,7 +794,10 @@ class ShadowEvalLogger(CustomLogger):
|
|||
if self._inflight_shadow_tasks >= _MAX_CONCURRENT_SHADOW_TASKS:
|
||||
self._record_funnel(job.id, "shed")
|
||||
continue
|
||||
self._job_starts[job.id] = self._job_starts.get(job.id, 0) + 1
|
||||
# One start writes one attempt row per arm, and max_turns is a row
|
||||
# ceiling, so admission must pre-count every arm or a multi-router
|
||||
# job overshoots the valve N-fold within a cache generation.
|
||||
self._job_starts[job.id] = self._job_starts.get(job.id, 0) + len(job.arm_router_names)
|
||||
self._inflight_shadow_tasks += 1
|
||||
asyncio.create_task(
|
||||
self._run_shadow_eval(
|
||||
|
|
@ -794,32 +836,74 @@ class ShadowEvalLogger(CustomLogger):
|
|||
shadow_params: Mapping[str, object],
|
||||
parent_metadata: Mapping[str, object],
|
||||
) -> None:
|
||||
"""Budget gate -> shadow call -> blind judge -> one attempt row, and every exit
|
||||
in exactly one coverage bucket: the gates that decline to spend on an admitted
|
||||
sample (no DB to record into, an over-budget key, an unverifiable or exhausted
|
||||
eval budget) count it withheld, so eligible traffic still reconciles as
|
||||
not_sampled + unjudgeable + shed + withheld + attempt rows. The prisma gate sits
|
||||
above the dispatch so no provider spend happens without a place to record the
|
||||
outcome, and the budget read lives here rather than in the success hook."""
|
||||
"""Budget gates once per sampled request, then every router arm in turn: shadow
|
||||
call -> blind judge -> one attempt row stamped with the arm. The gates that
|
||||
decline to spend on an admitted sample (no DB to record into, an over-budget key,
|
||||
an unverifiable or exhausted eval budget) count the REQUEST withheld before any
|
||||
arm runs, so funnel counters stay per-request and a leg's eligible traffic still
|
||||
reconciles as not_sampled + unjudgeable + shed + withheld + sampled requests,
|
||||
where each sampled request writes one attempt row per arm. A budget crossed
|
||||
mid-loop lets the remaining arms overshoot by one round, the same class of
|
||||
overshoot as the samples already in flight when the cap is crossed. The prisma
|
||||
gate sits above the dispatch so no provider spend happens without a place to
|
||||
record the outcome, and the budget read lives here rather than in the success
|
||||
hook."""
|
||||
prisma: Final = self._prisma_provider()
|
||||
if prisma is None:
|
||||
self._record_funnel(job.id, "withheld")
|
||||
return
|
||||
if await _key_or_team_is_over_budget(parent_metadata):
|
||||
self._record_funnel(job.id, "withheld")
|
||||
return
|
||||
if job.max_budget is not None:
|
||||
try:
|
||||
spend: Final = await self._read_job_spend(_job_spend_counter_key(job.id), job.spend, job.max_budget)
|
||||
except Exception as e: # noqa: BLE001 # unverifiable budget: skip the sample rather than spend on it
|
||||
verbose_logger.warning("shadow_eval: budget unverifiable for %s, sample skipped: %s", job.id, e)
|
||||
self._record_funnel(job.id, "withheld")
|
||||
return
|
||||
if spend >= job.max_budget:
|
||||
self._record_funnel(job.id, "withheld")
|
||||
return
|
||||
for arm_router in job.arm_router_names:
|
||||
await self._run_shadow_arm(
|
||||
prisma=prisma,
|
||||
job=job,
|
||||
arm_router=arm_router,
|
||||
request_id=request_id,
|
||||
messages=messages,
|
||||
real_text=real_text,
|
||||
real_model=real_model,
|
||||
real_cost=real_cost,
|
||||
real_classifier_cost=real_classifier_cost,
|
||||
real_cache_hit=real_cache_hit,
|
||||
control_tier=control_tier,
|
||||
shadow_params=shadow_params,
|
||||
parent_metadata=parent_metadata,
|
||||
)
|
||||
|
||||
async def _run_shadow_arm(
|
||||
self,
|
||||
prisma: "PrismaClient",
|
||||
job: ActiveShadowEvalJob,
|
||||
arm_router: str,
|
||||
request_id: str,
|
||||
messages: Sequence[Mapping[str, object]],
|
||||
real_text: str,
|
||||
real_model: str,
|
||||
real_cost: float,
|
||||
real_classifier_cost: float,
|
||||
real_cache_hit: bool,
|
||||
control_tier: str | None,
|
||||
shadow_params: Mapping[str, object],
|
||||
parent_metadata: Mapping[str, object],
|
||||
) -> None:
|
||||
"""One arm's pipeline: shadow call -> blind judge -> one attempt row, every exit
|
||||
recording this arm's outcome, so one arm's fault never silences a sibling arm."""
|
||||
try:
|
||||
if prisma is None:
|
||||
self._record_funnel(job.id, "withheld")
|
||||
return
|
||||
if await _key_or_team_is_over_budget(parent_metadata):
|
||||
self._record_funnel(job.id, "withheld")
|
||||
return
|
||||
if job.max_budget is not None:
|
||||
try:
|
||||
spend: Final = await self._read_job_spend(_job_spend_counter_key(job.id), job.spend, job.max_budget)
|
||||
except Exception as e: # noqa: BLE001 # unverifiable budget: skip the sample rather than spend on it
|
||||
verbose_logger.warning("shadow_eval: budget unverifiable for %s, sample skipped: %s", job.id, e)
|
||||
self._record_funnel(job.id, "withheld")
|
||||
return
|
||||
if spend >= job.max_budget:
|
||||
self._record_funnel(job.id, "withheld")
|
||||
return
|
||||
shadow: Final = await self._call_router_shadow(job.shadow_target, messages, shadow_params, parent_metadata)
|
||||
shadow: Final = await self._call_router_shadow(
|
||||
job.arm_target(arm_router), messages, shadow_params, parent_metadata
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # detached task: nothing billed yet, record and never raise
|
||||
verbose_logger.debug("shadow_eval: pipeline failed for %s: %s", request_id, e)
|
||||
await self._record_attempt(
|
||||
|
|
@ -827,6 +911,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
job,
|
||||
request_id,
|
||||
control_tier,
|
||||
router_name=arm_router,
|
||||
outcome="error",
|
||||
error=f"pipeline error: {e}",
|
||||
real_cost=real_cost,
|
||||
|
|
@ -840,6 +925,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
job,
|
||||
request_id,
|
||||
control_tier,
|
||||
router_name=arm_router,
|
||||
outcome="error",
|
||||
error=shadow.error,
|
||||
shadow_cost=shadow.cost,
|
||||
|
|
@ -864,6 +950,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
job,
|
||||
request_id,
|
||||
control_tier,
|
||||
router_name=arm_router,
|
||||
outcome="error",
|
||||
error=verdict.error,
|
||||
shadow=shadow,
|
||||
|
|
@ -880,6 +967,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
job,
|
||||
request_id,
|
||||
control_tier,
|
||||
router_name=arm_router,
|
||||
outcome=verdict.preference,
|
||||
shadow=shadow,
|
||||
real_model=real_model,
|
||||
|
|
@ -898,6 +986,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
job,
|
||||
request_id,
|
||||
control_tier,
|
||||
router_name=arm_router,
|
||||
outcome="error",
|
||||
error=f"pipeline error: {e}",
|
||||
shadow=shadow,
|
||||
|
|
@ -915,6 +1004,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
request_id: str,
|
||||
control_tier: str | None,
|
||||
*,
|
||||
router_name: str,
|
||||
outcome: str,
|
||||
real_cost: float,
|
||||
real_classifier_cost: float,
|
||||
|
|
@ -937,6 +1027,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
data={ # mutable-ok: Prisma payload
|
||||
"job_id": job.id,
|
||||
"request_id": request_id,
|
||||
"router_name": router_name,
|
||||
"outcome": outcome,
|
||||
"tier": control_tier if job.direction == "reverse" else (shadow.tier if shadow else None),
|
||||
"real_model": real_model or None,
|
||||
|
|
@ -1056,7 +1147,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
)
|
||||
|
||||
|
||||
_EMPTY_JOBS: Final[Mapping[str, tuple[ActiveShadowEvalJob, ...]]] = MappingProxyType({})
|
||||
_EMPTY_JOBS: Final[Mapping[tuple[str, str], tuple[ActiveShadowEvalJob, ...]]] = MappingProxyType({})
|
||||
|
||||
|
||||
def _default_prisma_provider() -> "PrismaClient | None":
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ from litellm._logging import verbose_logger
|
|||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage
|
||||
from litellm.types.prompts.init_prompts import PromptSpec
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
from litellm.types.utils import CallTypes, StandardCallbackDynamicParams
|
||||
from litellm.types.vector_stores import (
|
||||
LiteLLM_ManagedVectorStore,
|
||||
VectorStoreResultContent,
|
||||
|
|
@ -226,7 +226,7 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
self,
|
||||
request_data: dict,
|
||||
response: Any,
|
||||
call_type: Any | None,
|
||||
call_type: CallTypes | None,
|
||||
) -> Any | None:
|
||||
"""
|
||||
Add search results to the response after successful LLM call.
|
||||
|
|
@ -283,7 +283,7 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
self,
|
||||
request_data: dict,
|
||||
response_chunk: Any,
|
||||
call_type: Any | None,
|
||||
call_type: CallTypes | None,
|
||||
) -> Any | None:
|
||||
"""
|
||||
Add search results to the final streaming chunk.
|
||||
|
|
|
|||
|
|
@ -1633,16 +1633,18 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
def _select_search_tool_from_router(self, llm_router: object) -> "_SearchToolConfig | None":
|
||||
if llm_router is None or not hasattr(llm_router, "search_tools"):
|
||||
return None
|
||||
search_tools: Final = list(getattr(llm_router, "search_tools") or [])
|
||||
search_tools: Final = tuple(getattr(llm_router, "search_tools", None) or ())
|
||||
return self._select_search_tool_from_list(search_tools=search_tools, source="router")
|
||||
|
||||
def _select_search_tool_from_list(
|
||||
self,
|
||||
search_tools: list[_SearchToolConfig],
|
||||
search_tools: Sequence[_SearchToolConfig],
|
||||
source: str,
|
||||
) -> "_SearchToolConfig | None":
|
||||
if self.search_tool_name:
|
||||
matching_tools = [tool for tool in search_tools if tool.get("search_tool_name") == self.search_tool_name]
|
||||
matching_tools: Final = tuple(
|
||||
tool for tool in search_tools if tool.get("search_tool_name") == self.search_tool_name
|
||||
)
|
||||
if matching_tools:
|
||||
search_provider = (matching_tools[0].get("litellm_params", {}) or {}).get("search_provider")
|
||||
verbose_logger.debug(
|
||||
|
|
|
|||
|
|
@ -7,7 +7,13 @@ import os
|
|||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
from litellm.types.files import get_file_mime_type_from_extension
|
||||
from litellm.types.files import (
|
||||
AUDIO_FILE_TYPES,
|
||||
FILE_EXTENSIONS,
|
||||
FILE_MIME_TYPES,
|
||||
FileType,
|
||||
get_file_mime_type_from_extension,
|
||||
)
|
||||
from litellm.types.utils import FileTypes
|
||||
|
||||
|
||||
|
|
@ -323,3 +329,75 @@ def calculate_request_duration(file: FileTypes) -> float | None:
|
|||
except Exception:
|
||||
# Silently fail if duration extraction fails
|
||||
return None
|
||||
|
||||
|
||||
DEFAULT_SPEECH_MEDIA_TYPE: Final = "audio/mpeg"
|
||||
|
||||
|
||||
def _speech_media_type_for_response_format(response_format: str) -> str | None:
|
||||
file_type: Final = next(
|
||||
(candidate for candidate, extensions in FILE_EXTENSIONS.items() if response_format.lower() in extensions),
|
||||
None,
|
||||
)
|
||||
if file_type is None or file_type not in AUDIO_FILE_TYPES:
|
||||
return None
|
||||
return FILE_MIME_TYPES[file_type]
|
||||
|
||||
|
||||
def resolve_speech_media_type(upstream_content_type: str | None, response_format: str | None) -> str:
|
||||
upstream_media_type: Final = (upstream_content_type or "").split(";", 1)[0].strip().lower()
|
||||
if upstream_media_type.startswith("audio/"):
|
||||
return upstream_media_type
|
||||
requested_media_type: Final = (
|
||||
None if response_format is None else _speech_media_type_for_response_format(response_format)
|
||||
)
|
||||
return requested_media_type or DEFAULT_SPEECH_MEDIA_TYPE
|
||||
|
||||
|
||||
_OGG_OPUS_HEAD_WINDOW: Final = 64
|
||||
_ADTS_SYNC_AND_LAYER_MASK: Final = 0xF6
|
||||
_ADTS_SYNC_AND_LAYER: Final = 0xF0
|
||||
_ADTS_SAMPLE_RATE_INDEX_LIMIT: Final = 13
|
||||
_MPEG_SYNC_MASK: Final = 0xE0
|
||||
_MPEG_LAYER_MASK: Final = 0x06
|
||||
_MPEG_RESERVED_VERSION: Final = 0x01
|
||||
_MPEG_INVALID_BITRATE_INDEX: Final = 0x0F
|
||||
_MPEG_RESERVED_SAMPLE_RATE_INDEX: Final = 0x03
|
||||
|
||||
|
||||
def _adts_aac_frame_media_type(header: bytes) -> str | None:
|
||||
sample_rate_index: Final = (header[2] >> 2) & 0x0F
|
||||
return FILE_MIME_TYPES[FileType.AAC] if sample_rate_index < _ADTS_SAMPLE_RATE_INDEX_LIMIT else None
|
||||
|
||||
|
||||
def _mpeg_audio_frame_media_type(header: bytes) -> str | None:
|
||||
version: Final = (header[1] >> 3) & 0x03
|
||||
layer: Final = header[1] & _MPEG_LAYER_MASK
|
||||
bitrate_index: Final = header[2] >> 4
|
||||
sample_rate_index: Final = (header[2] >> 2) & 0x03
|
||||
if (
|
||||
(header[1] & _MPEG_SYNC_MASK) != _MPEG_SYNC_MASK
|
||||
or version == _MPEG_RESERVED_VERSION
|
||||
or layer == 0
|
||||
or bitrate_index == _MPEG_INVALID_BITRATE_INDEX
|
||||
or sample_rate_index == _MPEG_RESERVED_SAMPLE_RATE_INDEX
|
||||
):
|
||||
return None
|
||||
return FILE_MIME_TYPES[FileType.MP3]
|
||||
|
||||
|
||||
def speech_media_type_from_audio_bytes(audio: bytes) -> str | None:
|
||||
if audio[:4] == b"RIFF" and audio[8:12] == b"WAVE":
|
||||
return FILE_MIME_TYPES[FileType.WAV]
|
||||
if audio[:4] == b"fLaC":
|
||||
return FILE_MIME_TYPES[FileType.FLAC]
|
||||
if audio[:4] == b"OggS":
|
||||
is_opus: Final = b"OpusHead" in audio[:_OGG_OPUS_HEAD_WINDOW]
|
||||
return FILE_MIME_TYPES[FileType.OPUS if is_opus else FileType.OGG]
|
||||
if audio[:3] == b"ID3":
|
||||
return FILE_MIME_TYPES[FileType.MP3]
|
||||
if len(audio) < 3 or audio[0] != 0xFF:
|
||||
return None
|
||||
if (audio[1] & _ADTS_SYNC_AND_LAYER_MASK) == _ADTS_SYNC_AND_LAYER:
|
||||
return _adts_aac_frame_media_type(audio)
|
||||
return _mpeg_audio_frame_media_type(audio)
|
||||
|
|
|
|||
|
|
@ -50,6 +50,9 @@ OPTIONAL_KWARGS_KEYS: Final = (
|
|||
"vertex_ai_project",
|
||||
"vertex_ai_location",
|
||||
"vertex_ai_credentials",
|
||||
"gigachat_scope",
|
||||
"gigachat_auth_url",
|
||||
"gigachat_access_token",
|
||||
"tpm",
|
||||
"rpm",
|
||||
"itpm",
|
||||
|
|
|
|||
|
|
@ -369,6 +369,9 @@ def get_llm_provider(
|
|||
elif endpoint == "https://api.meta.ai/v1":
|
||||
custom_llm_provider = "meta"
|
||||
dynamic_api_key = get_secret_str("META_API_KEY")
|
||||
elif endpoint == "https://gigachat.devices.sberbank.ru/api/v1":
|
||||
custom_llm_provider = "gigachat"
|
||||
dynamic_api_key = get_secret_str("GIGACHAT_API_KEY")
|
||||
elif (json_provider := JSONProviderRegistry.get_by_base_url(endpoint)) is not None:
|
||||
custom_llm_provider = json_provider.slug
|
||||
dynamic_api_key = api_key if api_key is not None else get_secret_str(json_provider.api_key_env)
|
||||
|
|
@ -533,6 +536,14 @@ def get_llm_provider(
|
|||
)
|
||||
|
||||
|
||||
def _dashscope_family_chat_config(custom_llm_provider: str) -> "litellm.DashScopeChatConfig":
|
||||
if custom_llm_provider == "qwencloud":
|
||||
return litellm.QwenCloudChatConfig()
|
||||
if custom_llm_provider == "qwen_ai_platform":
|
||||
return litellm.QwenAIPlatformChatConfig()
|
||||
return litellm.DashScopeChatConfig()
|
||||
|
||||
|
||||
def _get_openai_compatible_provider_info(
|
||||
model: str,
|
||||
api_base: str | None,
|
||||
|
|
@ -782,11 +793,11 @@ def _get_openai_compatible_provider_info(
|
|||
api_base,
|
||||
dynamic_api_key,
|
||||
) = litellm.HerokuChatConfig()._get_openai_compatible_provider_info(api_base, api_key)
|
||||
elif custom_llm_provider == "dashscope":
|
||||
elif custom_llm_provider in ("dashscope", "qwencloud", "qwen_ai_platform"):
|
||||
(
|
||||
api_base,
|
||||
dynamic_api_key,
|
||||
) = litellm.DashScopeChatConfig()._get_openai_compatible_provider_info(api_base, api_key)
|
||||
) = _dashscope_family_chat_config(custom_llm_provider)._get_openai_compatible_provider_info(api_base, api_key)
|
||||
elif custom_llm_provider == "modelscope":
|
||||
(
|
||||
api_base,
|
||||
|
|
@ -867,6 +878,9 @@ def _get_openai_compatible_provider_info(
|
|||
# Manus is OpenAI compatible for responses API
|
||||
api_base = api_base or get_secret_str("MANUS_API_BASE") or "https://api.manus.im"
|
||||
dynamic_api_key = api_key or get_secret_str("MANUS_API_KEY")
|
||||
elif custom_llm_provider == "gigachat":
|
||||
api_base = api_base or get_secret_str("GIGACHAT_API_BASE") or "https://gigachat.devices.sberbank.ru/api/v1"
|
||||
dynamic_api_key = api_key or get_secret_str("GIGACHAT_API_KEY")
|
||||
|
||||
if api_base is not None and not isinstance(api_base, str):
|
||||
raise Exception(f"api base needs to be a string. api_base={api_base}")
|
||||
|
|
|
|||
|
|
@ -2141,6 +2141,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
|
||||
logging_result: Final = self.normalize_logging_result(result=result)
|
||||
|
||||
if isinstance(result, Response) and isinstance(logging_result, (ModelResponse, EmbeddingResponse)):
|
||||
result = logging_result
|
||||
|
||||
if standard_logging_object is None and result is not None and self.stream is not True:
|
||||
if self._is_recognized_call_type_for_logging(logging_result=logging_result) or isinstance(
|
||||
logging_result, (dict, list)
|
||||
|
|
@ -6152,7 +6155,10 @@ def get_standard_logging_object_payload(
|
|||
|
||||
def emit_standard_logging_payload(payload: StandardLoggingPayload):
|
||||
if os.getenv("LITELLM_PRINT_STANDARD_LOGGING_PAYLOAD"):
|
||||
print(json.dumps(payload, indent=4), flush=True) # noqa: T201
|
||||
try:
|
||||
print(json.dumps(payload, indent=4, default=str), flush=True) # noqa: T201
|
||||
except Exception as e: # noqa: BLE001 # Safe catch-all for verbose logging
|
||||
verbose_logger.exception("Error serializing standard logging payload for debug output: %s", e)
|
||||
|
||||
|
||||
def get_standard_logging_metadata(
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ Helper utilities for tracking the cost of built-in tools.
|
|||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Final, Literal
|
||||
from typing import Final, Literal
|
||||
|
||||
import litellm
|
||||
from litellm.constants import OPENAI_FILE_SEARCH_COST_PER_1K_CALLS
|
||||
|
|
@ -16,6 +16,7 @@ from litellm.types.llms.openai import (
|
|||
WebSearchOptions,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
ChatCompletionAnnotation,
|
||||
Message,
|
||||
ModelInfo,
|
||||
ModelResponse,
|
||||
|
|
@ -49,7 +50,7 @@ class StandardBuiltInToolCostTracking:
|
|||
@staticmethod
|
||||
def get_cost_for_built_in_tools(
|
||||
model: str,
|
||||
response_object: Any,
|
||||
response_object: object,
|
||||
usage: Usage | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
standard_built_in_tools_params: StandardBuiltInToolsParams | None = None,
|
||||
|
|
@ -201,8 +202,7 @@ class StandardBuiltInToolCostTracking:
|
|||
model_info: Final = StandardBuiltInToolCostTracking._safe_get_model_info(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
file_search_raw: Final[Any] = standard_built_in_tools_params.get("file_search", {})
|
||||
file_search_usage: Final[FileSearchTool | None] = FileSearchTool(**file_search_raw) if file_search_raw else None
|
||||
file_search_usage: Final[FileSearchTool | None] = standard_built_in_tools_params.get("file_search") or None
|
||||
|
||||
# Convert model_info to dict and extract usage parameters
|
||||
model_info_dict: Final = dict(model_info) if model_info is not None else None
|
||||
|
|
@ -245,7 +245,7 @@ class StandardBuiltInToolCostTracking:
|
|||
|
||||
@staticmethod
|
||||
def _extract_file_search_params(
|
||||
file_search_usage: Any,
|
||||
file_search_usage: object,
|
||||
) -> tuple[float | None, float | None]:
|
||||
"""Extract and convert file search parameters safely."""
|
||||
storage_gb = None
|
||||
|
|
@ -335,7 +335,7 @@ class StandardBuiltInToolCostTracking:
|
|||
|
||||
@staticmethod
|
||||
def _extract_token_counts(
|
||||
computer_use_usage: Any,
|
||||
computer_use_usage: object,
|
||||
) -> tuple[int | None, int | None]:
|
||||
"""Extract and convert token counts safely."""
|
||||
input_tokens = None
|
||||
|
|
@ -351,9 +351,9 @@ class StandardBuiltInToolCostTracking:
|
|||
return input_tokens, output_tokens
|
||||
|
||||
@staticmethod
|
||||
def _safe_convert_to_int(value: Any) -> int | None:
|
||||
def _safe_convert_to_int(value: object) -> int | None:
|
||||
"""Safely convert a value to int."""
|
||||
if value is not None:
|
||||
if isinstance(value, (int, float, str)):
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError):
|
||||
|
|
@ -381,7 +381,7 @@ class StandardBuiltInToolCostTracking:
|
|||
return usage.model_copy(update={"server_tool_use": server_tool_use})
|
||||
|
||||
@staticmethod
|
||||
def response_object_includes_web_search_call(response_object: Any, usage: Usage | None = None) -> bool:
|
||||
def response_object_includes_web_search_call(response_object: object, usage: Usage | None = None) -> bool:
|
||||
"""
|
||||
Check if the response object includes a web search call.
|
||||
|
||||
|
|
@ -446,7 +446,7 @@ class StandardBuiltInToolCostTracking:
|
|||
|
||||
@staticmethod
|
||||
def response_object_includes_file_search_call(
|
||||
response_object: Any,
|
||||
response_object: object,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if the response object includes a file search call.
|
||||
|
|
@ -477,11 +477,11 @@ class StandardBuiltInToolCostTracking:
|
|||
message: Message | None = getattr(choice, "message", None)
|
||||
if message is None:
|
||||
continue
|
||||
if annotations := getattr(message, "annotations", None):
|
||||
if len(annotations) > 0:
|
||||
for annotation in annotations:
|
||||
if annotation.get("type", None) == annotation_type:
|
||||
return True
|
||||
annotations: list[ChatCompletionAnnotation] | None = getattr(message, "annotations", None)
|
||||
if annotations:
|
||||
for annotation in annotations:
|
||||
if annotation.get("type", None) == annotation_type:
|
||||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -522,10 +522,8 @@ class StandardBuiltInToolCostTracking:
|
|||
if model_info is None:
|
||||
return 0.0
|
||||
|
||||
search_context_raw: Final[Any] = model_info.get("search_context_cost_per_query", {})
|
||||
search_context_pricing: Final[SearchContextCostPerQuery] = (
|
||||
SearchContextCostPerQuery(**search_context_raw) if search_context_raw else SearchContextCostPerQuery()
|
||||
)
|
||||
search_context_raw: Final = model_info.get("search_context_cost_per_query")
|
||||
search_context_pricing: Final[SearchContextCostPerQuery] = search_context_raw or SearchContextCostPerQuery()
|
||||
if web_search_options.get("search_context_size", None) == "low":
|
||||
return search_context_pricing.get("search_context_size_low", 0.0)
|
||||
elif web_search_options.get("search_context_size", None) == "medium":
|
||||
|
|
@ -545,10 +543,8 @@ class StandardBuiltInToolCostTracking:
|
|||
"""
|
||||
if model_info is None:
|
||||
return 0.0
|
||||
search_context_raw: Final[Any] = model_info.get("search_context_cost_per_query", {}) or {}
|
||||
search_context_pricing: Final[SearchContextCostPerQuery] = (
|
||||
SearchContextCostPerQuery(**search_context_raw) if search_context_raw else SearchContextCostPerQuery()
|
||||
)
|
||||
search_context_raw: Final = model_info.get("search_context_cost_per_query")
|
||||
search_context_pricing: Final[SearchContextCostPerQuery] = search_context_raw or SearchContextCostPerQuery()
|
||||
return search_context_pricing.get("search_context_size_medium", 0.0)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -714,7 +710,7 @@ class StandardBuiltInToolCostTracking:
|
|||
response_object: ModelResponse,
|
||||
) -> bool:
|
||||
for _choice in response_object.choices:
|
||||
message = getattr(_choice, "message", None)
|
||||
message: Message | None = getattr(_choice, "message", None)
|
||||
if (
|
||||
message is not None
|
||||
and hasattr(message, "annotations")
|
||||
|
|
|
|||
|
|
@ -144,7 +144,6 @@ def _is_choice_non_empty(choice: StreamingChoices) -> bool:
|
|||
# Check model_extra for dynamically added fields on the choice
|
||||
choice_extra_fields: Final[Mapping[str, object]] = choice.model_extra or {}
|
||||
for extra_field_name, extra_field_value in choice_extra_fields.items():
|
||||
# Skip certain structural fields that are just default/None placeholders
|
||||
if extra_field_name == "index" and extra_field_value == 0:
|
||||
continue
|
||||
if extra_field_name in {"finish_reason", "logprobs"} and extra_field_value is None:
|
||||
|
|
@ -192,7 +191,6 @@ def _is_delta_non_empty(delta: Delta) -> bool:
|
|||
# Check model_extra for dynamically added fields (this is where Pydantic stores them)
|
||||
delta_extra_fields: Final[Mapping[str, object]] = delta.model_extra or {}
|
||||
for extra_field_value in delta_extra_fields.values():
|
||||
# Even structural fields are meaningful if they have actual content
|
||||
if _has_meaningful_content(extra_field_value):
|
||||
return True
|
||||
|
||||
|
|
|
|||
|
|
@ -205,6 +205,41 @@ def is_non_content_values_set(message: AllMessageValues) -> bool:
|
|||
return any(message.get(key, None) is not None for key in message if key not in ignore_keys)
|
||||
|
||||
|
||||
_IMAGE_CONTENT_PART_TYPES: Final = frozenset({"image_url", "input_image", "image"})
|
||||
_IMAGE_SCAN_MAX_DEPTH: Final = 4
|
||||
|
||||
|
||||
def _content_parts_contain_image(parts: Sequence[object]) -> bool:
|
||||
"""Depth-bounded frontier walk over nested content lists, iterative because the repo bans
|
||||
recursion; an Anthropic tool_result nests its image parts exactly one level down."""
|
||||
frontier = parts # rebind-ok: depth-bounded frontier walk
|
||||
for _ in range(_IMAGE_SCAN_MAX_DEPTH):
|
||||
if any(isinstance(part, Mapping) and part.get("type") in _IMAGE_CONTENT_PART_TYPES for part in frontier):
|
||||
return True
|
||||
frontier = tuple( # rebind-ok: depth-bounded frontier walk
|
||||
nested
|
||||
for part in frontier
|
||||
if isinstance(part, Mapping)
|
||||
for content in (part.get("content"),)
|
||||
if isinstance(content, list)
|
||||
for nested in content
|
||||
)
|
||||
if not frontier:
|
||||
return False
|
||||
return False
|
||||
|
||||
|
||||
def request_contains_image_content(messages: Sequence[Mapping[str, object]]) -> bool:
|
||||
"""Whether any message carries an image content part, across the dialects that reach
|
||||
pre-routing hooks untranslated: chat-completions ``image_url``, Responses ``input_image``,
|
||||
and Anthropic Messages ``image``, including images nested inside ``tool_result`` blocks."""
|
||||
return any(
|
||||
isinstance(content, list) and _content_parts_contain_image(content)
|
||||
for message in messages
|
||||
for content in (message.get("content"),)
|
||||
)
|
||||
|
||||
|
||||
def _audio_or_image_in_message_content(message: AllMessageValues) -> bool:
|
||||
"""
|
||||
Checks if message content contains an image or audio
|
||||
|
|
@ -520,10 +555,10 @@ def update_messages_with_model_file_ids(
|
|||
|
||||
|
||||
def update_responses_input_with_model_file_ids(
|
||||
input: Any,
|
||||
input: object,
|
||||
model_id: str | None = None,
|
||||
model_file_id_mapping: dict[str, dict[str, str]] | None = None,
|
||||
) -> str | list[dict[str, Any]]:
|
||||
) -> object:
|
||||
"""
|
||||
Updates responses API input with provider-specific file IDs.
|
||||
File IDs are always inside the content array, not as direct input_file items.
|
||||
|
|
@ -604,8 +639,8 @@ def update_responses_input_with_model_file_ids(
|
|||
|
||||
|
||||
def _decode_vector_store_ids_in_tools(
|
||||
tools: list[dict[str, Any]] | None,
|
||||
) -> list[dict[str, Any]] | None:
|
||||
tools: list[dict[str, object]] | None,
|
||||
) -> list[dict[str, object]] | None:
|
||||
"""
|
||||
Decodes unified (LiteLLM-managed) vector_store_ids in file_search tools to
|
||||
provider-native IDs. Non-unified IDs are passed through unchanged.
|
||||
|
|
@ -657,10 +692,10 @@ def _decode_vector_store_ids_in_tools(
|
|||
|
||||
|
||||
def update_responses_tools_with_model_file_ids(
|
||||
tools: list[dict[str, Any]] | None,
|
||||
tools: list[dict[str, object]] | None,
|
||||
model_id: str | None = None,
|
||||
model_file_id_mapping: dict[str, dict[str, str]] | None = None,
|
||||
) -> list[dict[str, Any]] | None:
|
||||
) -> list[dict[str, object]] | None:
|
||||
"""
|
||||
Updates responses API tools with provider-specific file IDs.
|
||||
|
||||
|
|
@ -853,7 +888,7 @@ def extract_file_data(file_data: FileTypes) -> ExtractedFileData:
|
|||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _estimate_json_bytes(obj: Any) -> int:
|
||||
def _estimate_json_bytes(obj: object) -> int:
|
||||
"""Estimate the JSON-serialised byte size of ``obj`` without materialising
|
||||
JSON. Walks iteratively (no recursion stack risk).
|
||||
|
||||
|
|
@ -1944,7 +1979,7 @@ def drop_tool_reference_parts_from_tool_messages(
|
|||
return [_drop_tool_reference_parts(message) for message in messages] # mutable-ok: pipelines mutate message lists
|
||||
|
||||
|
||||
def _attempt_json_repair(s: str) -> Any | None:
|
||||
def _attempt_json_repair(s: str) -> object | None:
|
||||
"""
|
||||
Attempt to repair truncated JSON produced by LLM tool calls.
|
||||
|
||||
|
|
@ -2060,7 +2095,7 @@ def parse_tool_call_arguments(
|
|||
raise ValueError(error_message) from original_error
|
||||
|
||||
|
||||
def split_concatenated_json_objects(raw: str) -> list[dict[str, Any]]:
|
||||
def split_concatenated_json_objects(raw: str) -> list[dict[str, object]]:
|
||||
"""
|
||||
Split a string that contains one or more concatenated JSON objects into
|
||||
a list of parsed dicts.
|
||||
|
|
@ -2096,7 +2131,7 @@ def split_concatenated_json_objects(raw: str) -> list[dict[str, Any]]:
|
|||
return []
|
||||
|
||||
decoder: Final = json.JSONDecoder()
|
||||
results: Final[list[dict[str, Any]]] = []
|
||||
results: Final[list[dict[str, object]]] = []
|
||||
idx = 0
|
||||
length: Final = len(raw)
|
||||
|
||||
|
|
|
|||
|
|
@ -1500,6 +1500,6 @@ class RealTimeStreaming:
|
|||
pass
|
||||
|
||||
|
||||
def client_sent_openai_beta_realtime_header(websocket: Any) -> bool:
|
||||
def client_sent_openai_beta_realtime_header(websocket: _ScopedWebSocket) -> bool:
|
||||
"""True when the client WebSocket includes ``OpenAI-Beta: realtime=v1``."""
|
||||
return RealTimeStreaming._detect_beta_header(websocket)
|
||||
|
|
|
|||
|
|
@ -73,6 +73,18 @@ class _ContentChunk(TypedDict):
|
|||
choices: Sequence[_ContentChoice]
|
||||
|
||||
|
||||
class _FunctionCallDelta(TypedDict):
|
||||
function_call: ReadOnly[FunctionCall]
|
||||
|
||||
|
||||
class _FunctionCallChoice(TypedDict):
|
||||
delta: ReadOnly[_FunctionCallDelta]
|
||||
|
||||
|
||||
class _FunctionCallChunk(TypedDict):
|
||||
choices: ReadOnly[Sequence[_FunctionCallChoice]]
|
||||
|
||||
|
||||
class _AudioDelta(TypedDict, total=False):
|
||||
audio: ChatCompletionAudioDelta | None
|
||||
|
||||
|
|
@ -588,7 +600,7 @@ class ChunkProcessor:
|
|||
|
||||
return tool_calls_list
|
||||
|
||||
def get_combined_function_call_content(self, function_call_chunks: list[dict[str, Any]]) -> FunctionCall:
|
||||
def get_combined_function_call_content(self, function_call_chunks: Sequence["_FunctionCallChunk"]) -> FunctionCall:
|
||||
argument_list: Final = []
|
||||
delta = function_call_chunks[0]["choices"][0]["delta"]
|
||||
function_call = delta.get("function_call", "")
|
||||
|
|
|
|||
|
|
@ -862,6 +862,8 @@ class CustomStreamWrapper:
|
|||
model_response: Final = ModelResponseStream(**args)
|
||||
if self.response_id is not None:
|
||||
model_response.id = self.response_id
|
||||
elif model_response.id:
|
||||
self.response_id = model_response.id
|
||||
if self.system_fingerprint is not None:
|
||||
model_response.system_fingerprint = self.system_fingerprint
|
||||
|
||||
|
|
|
|||
|
|
@ -4,8 +4,9 @@ import base64
|
|||
import io
|
||||
import struct
|
||||
from collections.abc import Callable, Iterable, Mapping, Sequence
|
||||
from typing import Any, Final, Literal, cast
|
||||
from typing import Final, Literal, cast
|
||||
|
||||
import httpx
|
||||
import tiktoken
|
||||
|
||||
import litellm
|
||||
|
|
@ -171,6 +172,10 @@ def calculate_tiles_needed(
|
|||
return total_tiles
|
||||
|
||||
|
||||
def _unpack_ints(fmt: str, buffer: bytes) -> tuple[int, ...]:
|
||||
return struct.unpack(fmt, buffer)
|
||||
|
||||
|
||||
def get_image_type(image_data: bytes) -> str | None:
|
||||
"""take an image (really only the first ~100 bytes max are needed)
|
||||
and return 'png' 'gif' 'jpeg' 'webp' 'heic' or None. method added to
|
||||
|
|
@ -210,9 +215,9 @@ def get_image_dimensions(
|
|||
if data.startswith(("http://", "https://")):
|
||||
try:
|
||||
client: Final = _get_httpx_client()
|
||||
response: Final = safe_get(client, data)
|
||||
response: Final[httpx.Response] = safe_get(client, data)
|
||||
max_bytes: Final = int(MAX_IMAGE_URL_DOWNLOAD_SIZE_MB * 1024 * 1024)
|
||||
content_length: Final = response.headers.get("Content-Length")
|
||||
content_length: Final[str | None] = response.headers.get("Content-Length")
|
||||
if content_length is not None and int(content_length) > max_bytes:
|
||||
pass # skip download; img_data stays None
|
||||
else:
|
||||
|
|
@ -229,10 +234,10 @@ def get_image_dimensions(
|
|||
img_type: Final = get_image_type(img_data)
|
||||
|
||||
if img_type == "png":
|
||||
w, h = struct.unpack(">LL", img_data[16:24])
|
||||
w, h = _unpack_ints(">LL", img_data[16:24])
|
||||
return w, h
|
||||
elif img_type == "gif":
|
||||
w, h = struct.unpack("<HH", img_data[6:10])
|
||||
w, h = _unpack_ints("<HH", img_data[6:10])
|
||||
return w, h
|
||||
elif img_type == "jpeg":
|
||||
with io.BytesIO(img_data) as fhandle:
|
||||
|
|
@ -245,25 +250,25 @@ def get_image_dimensions(
|
|||
while ord(byte) == 0xFF:
|
||||
byte = fhandle.read(1)
|
||||
ftype = ord(byte)
|
||||
size = struct.unpack(">H", fhandle.read(2))[0] - 2
|
||||
size = _unpack_ints(">H", fhandle.read(2))[0] - 2
|
||||
fhandle.seek(1, 1)
|
||||
h, w = struct.unpack(">HH", fhandle.read(4))
|
||||
h, w = _unpack_ints(">HH", fhandle.read(4))
|
||||
return w, h
|
||||
elif img_type == "webp":
|
||||
# For WebP, the dimensions are stored at different offsets depending on the format
|
||||
# Check for VP8X (extended format)
|
||||
if img_data[12:16] == b"VP8X":
|
||||
w = struct.unpack("<I", img_data[24:27] + b"\x00")[0] + 1
|
||||
h = struct.unpack("<I", img_data[27:30] + b"\x00")[0] + 1
|
||||
w = _unpack_ints("<I", img_data[24:27] + b"\x00")[0] + 1
|
||||
h = _unpack_ints("<I", img_data[27:30] + b"\x00")[0] + 1
|
||||
return w, h
|
||||
# Check for VP8 (lossy format)
|
||||
elif img_data[12:16] == b"VP8 ":
|
||||
w = struct.unpack("<H", img_data[26:28])[0] & 0x3FFF
|
||||
h = struct.unpack("<H", img_data[28:30])[0] & 0x3FFF
|
||||
w = _unpack_ints("<H", img_data[26:28])[0] & 0x3FFF
|
||||
h = _unpack_ints("<H", img_data[28:30])[0] & 0x3FFF
|
||||
return w, h
|
||||
# Check for VP8L (lossless format)
|
||||
elif img_data[12:16] == b"VP8L":
|
||||
bits: Final = struct.unpack("<I", img_data[21:25])[0]
|
||||
bits: Final = _unpack_ints("<I", img_data[21:25])[0]
|
||||
w = (bits & 0x3FFF) + 1
|
||||
h = ((bits >> 14) & 0x3FFF) + 1
|
||||
return w, h
|
||||
|
|
@ -420,8 +425,8 @@ def token_counter(
|
|||
|
||||
def _count_function_call_tokens(
|
||||
key: str,
|
||||
value: Any,
|
||||
message: Mapping[str, Any],
|
||||
value: object,
|
||||
message: Mapping[str, object],
|
||||
count_function: TokenCounterFunction,
|
||||
) -> int:
|
||||
"""
|
||||
|
|
@ -587,7 +592,7 @@ def _fix_model_name(model: str) -> str:
|
|||
|
||||
|
||||
def _count_image_tokens(
|
||||
image_url: Any,
|
||||
image_url: object,
|
||||
use_default_image_token_count: bool,
|
||||
) -> int:
|
||||
"""
|
||||
|
|
@ -627,7 +632,7 @@ def _count_image_tokens(
|
|||
raise ValueError(f"Invalid image_url type: {type(image_url).__name__}. Expected str or dict with 'url' field.")
|
||||
|
||||
|
||||
def _validate_anthropic_content(content: Mapping[str, Any]) -> type:
|
||||
def _validate_anthropic_content(content: Mapping[str, object]) -> type:
|
||||
"""
|
||||
Validate and determine which Anthropic TypedDict applies.
|
||||
|
||||
|
|
@ -642,7 +647,7 @@ def _validate_anthropic_content(content: Mapping[str, Any]) -> type:
|
|||
"tool_result": AnthropicMessagesToolResultParam,
|
||||
}
|
||||
|
||||
expected_cls: Final = mapping.get(content_type)
|
||||
expected_cls: Final = mapping.get(content_type) if isinstance(content_type, str) else None
|
||||
if expected_cls is None:
|
||||
raise ValueError(f"Unknown Anthropic content type: '{content_type}'")
|
||||
|
||||
|
|
@ -693,8 +698,28 @@ def _count_document_tokens(
|
|||
)
|
||||
|
||||
|
||||
def _count_file_tokens(
|
||||
file_value: object,
|
||||
count_function: TokenCounterFunction,
|
||||
use_default_image_token_count: bool,
|
||||
) -> int:
|
||||
"""An OpenAI `file` block is the chat-completions spelling of a document, so it prices like one."""
|
||||
if not isinstance(file_value, Mapping):
|
||||
return 0
|
||||
filename: Final = file_value.get("filename")
|
||||
file_data: Final = file_value.get("file_data")
|
||||
name_tokens: Final = count_function(filename) if isinstance(filename, str) and filename else 0
|
||||
if not isinstance(file_data, str) or not file_data:
|
||||
return name_tokens
|
||||
return name_tokens + calculate_img_tokens(
|
||||
data=file_data,
|
||||
mode="auto",
|
||||
use_default_image_token_count=use_default_image_token_count,
|
||||
)
|
||||
|
||||
|
||||
def _count_anthropic_content(
|
||||
content: Mapping[str, Any],
|
||||
content: Mapping[str, object],
|
||||
count_function: TokenCounterFunction,
|
||||
use_default_image_token_count: bool,
|
||||
default_token_count: int | None,
|
||||
|
|
@ -709,7 +734,7 @@ def _count_anthropic_content(
|
|||
avoiding hardcoded field names.
|
||||
"""
|
||||
typeddict_cls: Final = _validate_anthropic_content(content)
|
||||
type_hints: Final = getattr(typeddict_cls, "__annotations__", {})
|
||||
type_hints: Final[Mapping[str, object]] = getattr(typeddict_cls, "__annotations__", {})
|
||||
tokens = 0
|
||||
|
||||
# Fields to skip (metadata/identifiers that don't contribute to prompt tokens)
|
||||
|
|
@ -778,6 +803,12 @@ def _count_content_list(
|
|||
use_default_image_token_count,
|
||||
default_token_count,
|
||||
)
|
||||
elif c["type"] == "file":
|
||||
num_tokens += _count_file_tokens(
|
||||
c.get("file"),
|
||||
count_function,
|
||||
use_default_image_token_count,
|
||||
)
|
||||
elif c["type"] in ("tool_use", "tool_result"):
|
||||
num_tokens += _count_anthropic_content(
|
||||
c,
|
||||
|
|
@ -807,7 +838,7 @@ def _count_content_list(
|
|||
raise ValueError(
|
||||
f"Invalid content item type: {content_type}. "
|
||||
f"Expected str or dict with 'type' field "
|
||||
f"(text, image_url, image, document, tool_use, tool_result, thinking, tool_reference)."
|
||||
f"(text, image_url, image, document, file, tool_use, tool_result, thinking, tool_reference)."
|
||||
)
|
||||
return num_tokens
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -11,8 +11,11 @@ A2A Protocol Format:
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional
|
||||
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
|
@ -23,6 +26,13 @@ if TYPE_CHECKING:
|
|||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
||||
class _A2ATextPart(TypedDict, total=False):
|
||||
"""The subset of an A2A message part this handler reads text from."""
|
||||
|
||||
kind: ReadOnly[str]
|
||||
text: ReadOnly[str]
|
||||
|
||||
|
||||
class A2AGuardrailHandler(BaseTranslation):
|
||||
"""
|
||||
Handler for processing A2A Protocol messages with guardrails.
|
||||
|
|
@ -41,7 +51,7 @@ class A2AGuardrailHandler(BaseTranslation):
|
|||
data: dict,
|
||||
guardrail_to_apply: "CustomGuardrail",
|
||||
litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> Any:
|
||||
) -> dict:
|
||||
"""
|
||||
Process A2A input messages by applying guardrails to text content.
|
||||
|
||||
|
|
@ -214,12 +224,12 @@ class A2AGuardrailHandler(BaseTranslation):
|
|||
|
||||
async def process_output_streaming_response(
|
||||
self,
|
||||
responses_so_far: list[Any],
|
||||
responses_so_far: list[object],
|
||||
guardrail_to_apply: "CustomGuardrail",
|
||||
litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
|
||||
request_data: dict | None = None,
|
||||
) -> list[Any]:
|
||||
) -> list[object]:
|
||||
"""
|
||||
Process A2A streaming output by applying guardrails to accumulated text.
|
||||
|
||||
|
|
@ -305,11 +315,12 @@ class A2AGuardrailHandler(BaseTranslation):
|
|||
|
||||
def _parse_streaming_responses(
|
||||
self,
|
||||
responses_so_far: list[Any],
|
||||
) -> tuple[list[dict[str, Any] | None], list[tuple[int, dict[str, Any]]]]:
|
||||
responses_so_far: list[object],
|
||||
) -> tuple[list[dict[str, object] | None], list[tuple[int, dict[str, object]]]]:
|
||||
"""Parse JSON-RPC items, returning aligned parsed list and valid entries."""
|
||||
parsed: Final[list[dict[str, Any] | None]] = [None] * len(responses_so_far)
|
||||
parsed: Final[list[dict[str, object] | None]] = [None] * len(responses_so_far)
|
||||
for i, item in enumerate(responses_so_far):
|
||||
obj: dict[str, object]
|
||||
if isinstance(item, dict):
|
||||
obj = item
|
||||
elif isinstance(item, str):
|
||||
|
|
@ -326,7 +337,7 @@ class A2AGuardrailHandler(BaseTranslation):
|
|||
|
||||
def _collect_text_from_parsed_chunks(
|
||||
self,
|
||||
valid_parsed: list[tuple[int, dict[str, Any]]],
|
||||
valid_parsed: list[tuple[int, dict[str, object]]],
|
||||
) -> tuple[str, list[int]]:
|
||||
"""Collect text from parsed chunks, returning combined text and indices."""
|
||||
from litellm.llms.a2a.common_utils import extract_text_from_a2a_response
|
||||
|
|
@ -411,7 +422,7 @@ class A2AGuardrailHandler(BaseTranslation):
|
|||
|
||||
def _extract_texts_from_parts(
|
||||
self,
|
||||
parts: list[dict[str, Any]],
|
||||
parts: Sequence[_A2ATextPart],
|
||||
path: tuple[str, ...],
|
||||
texts_to_check: list[str],
|
||||
task_mappings: list[tuple[tuple[str, ...], int]],
|
||||
|
|
|
|||
|
|
@ -100,16 +100,6 @@ InputWriteBackTarget = (
|
|||
)
|
||||
|
||||
|
||||
class _SSEDelta(TypedDict, total=False):
|
||||
type: ReadOnly[str]
|
||||
text: ReadOnly[str]
|
||||
stop_reason: ReadOnly[str | None]
|
||||
|
||||
|
||||
class _SSEEventData(TypedDict, total=False):
|
||||
delta: ReadOnly[_SSEDelta]
|
||||
|
||||
|
||||
def _as_str_mapping(value: Mapping[str, object]) -> Mapping[str, object]:
|
||||
return value
|
||||
|
||||
|
|
@ -157,6 +147,16 @@ class ExtractedInput:
|
|||
EMPTY_EXTRACTED_INPUT: Final = ExtractedInput(scanned=(), images=())
|
||||
|
||||
|
||||
class _AnthropicSSEDelta(TypedDict, total=False):
|
||||
type: ReadOnly[str]
|
||||
text: ReadOnly[str]
|
||||
stop_reason: ReadOnly[str | None]
|
||||
|
||||
|
||||
class _AnthropicSSEEvent(TypedDict, total=False):
|
||||
delta: ReadOnly[_AnthropicSSEDelta]
|
||||
|
||||
|
||||
class AnthropicMessagesHandler(BaseTranslation):
|
||||
"""Process Anthropic messages with guardrails.
|
||||
|
||||
|
|
@ -859,12 +859,28 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
@staticmethod
|
||||
def _image_sources(block: Mapping[str, object]) -> tuple[str, ...]:
|
||||
"""Normalize an Anthropic image block into strings a guardrail can read.
|
||||
|
||||
base64 becomes a data URI so the format travels with the payload, which is what
|
||||
the OpenAI path already puts in this field. A file source yields nothing: those
|
||||
bytes live behind the Files API and this extractor has no client to fetch them.
|
||||
"""
|
||||
source: Final = block.get("source")
|
||||
if not isinstance(source, Mapping):
|
||||
return ()
|
||||
# Could be base64 or url
|
||||
|
||||
source_type: Final = source.get("type")
|
||||
if source_type == "url":
|
||||
url: Final = source.get("url")
|
||||
return (url,) if isinstance(url, str) and url else ()
|
||||
|
||||
data: Final = source.get("data")
|
||||
return (data,) if data else ()
|
||||
if not isinstance(data, str) or not data:
|
||||
return ()
|
||||
media_type: Final = source.get("media_type")
|
||||
if isinstance(media_type, str) and media_type:
|
||||
return (f"data:{media_type};base64,{data}",)
|
||||
return (data,)
|
||||
|
||||
async def _apply_guardrail_responses_to_input(
|
||||
self,
|
||||
|
|
@ -1231,8 +1247,8 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
# Only process content_block_delta events
|
||||
if event_type == "content_block_delta" and data_line:
|
||||
try:
|
||||
data: _SSEEventData = json.loads(data_line)
|
||||
delta = data.get("delta", {})
|
||||
data: _AnthropicSSEEvent = json.loads(data_line)
|
||||
delta: _AnthropicSSEDelta = data.get("delta", {})
|
||||
if delta.get("type") == "text_delta":
|
||||
text += delta.get("text", "")
|
||||
except json.JSONDecodeError:
|
||||
|
|
@ -1294,9 +1310,9 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
# Check for message_delta event with stop_reason
|
||||
if event_type == "message_delta" and data_line:
|
||||
try:
|
||||
data: _SSEEventData = json.loads(data_line)
|
||||
delta = data.get("delta", {})
|
||||
stop_reason = delta.get("stop_reason")
|
||||
data: _AnthropicSSEEvent = json.loads(data_line)
|
||||
delta: _AnthropicSSEDelta = data.get("delta", {})
|
||||
stop_reason: str | None = delta.get("stop_reason")
|
||||
if stop_reason is not None:
|
||||
return True
|
||||
except json.JSONDecodeError:
|
||||
|
|
|
|||
|
|
@ -66,6 +66,10 @@ if TYPE_CHECKING:
|
|||
from litellm.llms.base_llm.chat.transformation import BaseConfig
|
||||
|
||||
|
||||
def _loads_stream_chunk(payload: str) -> dict[str, object]:
|
||||
return json.loads(payload)
|
||||
|
||||
|
||||
async def make_call(
|
||||
client: AsyncHTTPHandler | None,
|
||||
api_base: str,
|
||||
|
|
@ -78,7 +82,7 @@ async def make_call(
|
|||
json_mode: bool,
|
||||
speed: str | None = None,
|
||||
tool_name_reverse_map: dict[str, str] | None = None,
|
||||
) -> tuple[Any, httpx.Headers]:
|
||||
) -> tuple["ModelResponseIterator", httpx.Headers]:
|
||||
if client is None:
|
||||
client = litellm.module_level_aclient
|
||||
|
||||
|
|
@ -93,7 +97,7 @@ async def make_call(
|
|||
)
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_headers = getattr(e, "headers", None)
|
||||
error_response: Final = getattr(e, "response", None)
|
||||
error_response: Final[object] = getattr(e, "response", None)
|
||||
if error_headers is None and error_response:
|
||||
error_headers = getattr(error_response, "headers", None)
|
||||
raise AnthropicError(
|
||||
|
|
@ -138,7 +142,7 @@ def make_sync_call(
|
|||
json_mode: bool,
|
||||
speed: str | None = None,
|
||||
tool_name_reverse_map: dict[str, str] | None = None,
|
||||
) -> tuple[Any, httpx.Headers]:
|
||||
) -> tuple["ModelResponseIterator", httpx.Headers]:
|
||||
if client is None:
|
||||
client = litellm.module_level_client # re-use a module level client
|
||||
|
||||
|
|
@ -153,7 +157,7 @@ def make_sync_call(
|
|||
)
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_headers = getattr(e, "headers", None)
|
||||
error_response: Final = getattr(e, "response", None)
|
||||
error_response: Final[object] = getattr(e, "response", None)
|
||||
if error_headers is None and error_response:
|
||||
error_headers = getattr(error_response, "headers", None)
|
||||
raise AnthropicError(
|
||||
|
|
@ -292,7 +296,7 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
status_code: Final = getattr(e, "status_code", 500)
|
||||
error_headers = getattr(e, "headers", None)
|
||||
error_text = getattr(e, "text", str(e))
|
||||
error_response: Final = getattr(e, "response", None)
|
||||
error_response: Final[object] = getattr(e, "response", None)
|
||||
if error_headers is None and error_response:
|
||||
error_headers = getattr(error_response, "headers", None)
|
||||
if error_response and hasattr(error_response, "text"):
|
||||
|
|
@ -593,7 +597,7 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
status_code: Final = getattr(e, "status_code", 500)
|
||||
error_headers = getattr(e, "headers", None)
|
||||
error_text = getattr(e, "text", str(e))
|
||||
error_response: Final = getattr(e, "response", None)
|
||||
error_response: Final[object] = getattr(e, "response", None)
|
||||
if error_headers is None and error_response:
|
||||
error_headers = getattr(error_response, "headers", None)
|
||||
if error_response and hasattr(error_response, "text"):
|
||||
|
|
@ -664,10 +668,10 @@ class ModelResponseIterator:
|
|||
|
||||
# Accumulate web_search_tool_result blocks for multi-turn reconstruction
|
||||
# See: https://github.com/BerriAI/litellm/issues/17737
|
||||
self.web_search_results: list[dict[str, Any]] = []
|
||||
self.web_search_results: list[dict[str, object]] = []
|
||||
|
||||
# Accumulate compaction blocks for multi-turn reconstruction
|
||||
self.compaction_blocks: list[dict[str, Any]] = []
|
||||
self.compaction_blocks: list[dict[str, object]] = []
|
||||
|
||||
# Accumulate streamed thinking text so final usage can split reasoning
|
||||
# tokens from regular output tokens.
|
||||
|
|
@ -727,7 +731,7 @@ class ModelResponseIterator:
|
|||
str,
|
||||
ChatCompletionToolCallChunk | None,
|
||||
list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock],
|
||||
dict[str, Any],
|
||||
dict[str, object],
|
||||
str | None,
|
||||
]:
|
||||
"""
|
||||
|
|
@ -735,7 +739,7 @@ class ModelResponseIterator:
|
|||
"""
|
||||
text = ""
|
||||
tool_use: ChatCompletionToolCallChunk | None = None
|
||||
provider_specific_fields: Final = {}
|
||||
provider_specific_fields: Final[dict[str, object]] = {}
|
||||
reasoning_content: str | None = None
|
||||
content_block: Final = ContentBlockDelta(**chunk)
|
||||
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] = []
|
||||
|
|
@ -809,8 +813,8 @@ class ModelResponseIterator:
|
|||
def _handle_redacted_thinking_content(
|
||||
self,
|
||||
content_block_start: ContentBlockStart,
|
||||
provider_specific_fields: dict[str, Any],
|
||||
) -> tuple[list[ChatCompletionRedactedThinkingBlock], dict[str, Any]]:
|
||||
provider_specific_fields: dict[str, object],
|
||||
) -> tuple[list[ChatCompletionRedactedThinkingBlock], dict[str, object]]:
|
||||
"""
|
||||
Handle the redacted thinking content
|
||||
"""
|
||||
|
|
@ -878,7 +882,7 @@ class ModelResponseIterator:
|
|||
tool_use: ChatCompletionToolCallChunk | None = None
|
||||
finish_reason = ""
|
||||
usage: Usage | None = None
|
||||
provider_specific_fields: dict[str, Any] = {}
|
||||
provider_specific_fields: dict[str, object] = {}
|
||||
reasoning_content: str | None = None
|
||||
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = None
|
||||
|
||||
|
|
@ -1212,7 +1216,7 @@ class ModelResponseIterator:
|
|||
|
||||
# Try to parse as valid JSON first
|
||||
try:
|
||||
data_json: Final = json.loads(data_str)
|
||||
data_json: Final = _loads_stream_chunk(data_str)
|
||||
return self.chunk_parser(chunk=data_json)
|
||||
except json.JSONDecodeError:
|
||||
# Switch to accumulation mode and start accumulating
|
||||
|
|
@ -1330,7 +1334,7 @@ class ModelResponseIterator:
|
|||
str_line = str_line[index:]
|
||||
|
||||
if str_line.startswith("data:"):
|
||||
data_json: Final = json.loads(str_line[5:])
|
||||
data_json: Final = _loads_stream_chunk(str_line[5:])
|
||||
return self.chunk_parser(chunk=data_json)
|
||||
else:
|
||||
return ModelResponseStream(id=self.response_id)
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Any, Final, NoReturn, cast
|
|||
|
||||
import httpx
|
||||
from pydantic import ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm.constants import (
|
||||
|
|
@ -125,7 +126,25 @@ else:
|
|||
_ANTHROPIC_TOOL_NAME_INVALID_CHARS: Final = re.compile(r"[^a-zA-Z0-9_-]")
|
||||
_ANTHROPIC_TOOL_NAME_MAX_LEN: Final = 128
|
||||
|
||||
_ENUM_TYPE_CHECKS: Final[Mapping[str, Callable[[Any], bool]]] = MappingProxyType(
|
||||
|
||||
class _AnthropicUsageIteration(TypedDict, total=False):
|
||||
"""One entry of the ``usage.iterations`` array on an Anthropic response."""
|
||||
|
||||
input_tokens: ReadOnly[int | None]
|
||||
output_tokens: ReadOnly[int | None]
|
||||
cache_creation_input_tokens: ReadOnly[int | None]
|
||||
cache_read_input_tokens: ReadOnly[int | None]
|
||||
|
||||
|
||||
class _AnthropicToolResultBlock(TypedDict, total=False):
|
||||
"""A ``*_tool_result`` content block on an Anthropic response."""
|
||||
|
||||
type: ReadOnly[str]
|
||||
tool_use_id: ReadOnly[str]
|
||||
content: ReadOnly[object]
|
||||
|
||||
|
||||
_ENUM_TYPE_CHECKS: Final[Mapping[str, Callable[[object], bool]]] = MappingProxyType(
|
||||
{
|
||||
"null": lambda v: v is None,
|
||||
"boolean": lambda v: isinstance(v, bool),
|
||||
|
|
@ -440,7 +459,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
optional_params.pop("speed", None)
|
||||
|
||||
@staticmethod
|
||||
def _raise_invalid_reasoning_effort(model: str, value: Any, llm_provider: str) -> NoReturn:
|
||||
def _raise_invalid_reasoning_effort(model: str, value: object, llm_provider: str) -> NoReturn:
|
||||
"""Raise a ``BadRequestError`` for an unrecognised ``reasoning_effort``.
|
||||
|
||||
Args:
|
||||
|
|
@ -1992,19 +2011,35 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return data
|
||||
|
||||
def _apply_output_config(self, data: dict, model: str, optional_params: dict) -> None:
|
||||
"""Validate and apply output_config to the request data."""
|
||||
"""Validate and apply output_config to the request data.
|
||||
|
||||
The ``drop_params`` gate here is an effort gate: ``format`` is a
|
||||
structured-output field, not an effort field, so it survives the drop
|
||||
and is vetted where it is consumed (the map's
|
||||
``supports_native_structured_output`` flag on emission paths).
|
||||
"""
|
||||
if "output_config" not in optional_params:
|
||||
return
|
||||
output_config: Final = optional_params.get("output_config")
|
||||
if not output_config or not isinstance(output_config, dict):
|
||||
return
|
||||
if litellm.drop_params is True and not self._model_supports_effort_param(model, self._resolved_provider):
|
||||
if (
|
||||
litellm.drop_params is True
|
||||
and any(key != "format" for key in output_config)
|
||||
and not self._model_supports_effort_param(model, self._resolved_provider)
|
||||
):
|
||||
litellm.verbose_logger.warning(
|
||||
DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING,
|
||||
model,
|
||||
)
|
||||
optional_params.pop("output_config", None)
|
||||
data.pop("output_config", None)
|
||||
preserved_format: Final = output_config.get("format")
|
||||
if preserved_format is None:
|
||||
optional_params.pop("output_config", None)
|
||||
data.pop("output_config", None)
|
||||
return
|
||||
format_only: Final = {"format": preserved_format} # mutable-ok: json body
|
||||
optional_params["output_config"] = format_only # rebind-ok: out-param store
|
||||
data["output_config"] = format_only # rebind-ok: out-param store
|
||||
return
|
||||
effort: Final = output_config.get("effort")
|
||||
valid_efforts: Final = ["high", "medium", "low", "xhigh", "max"]
|
||||
|
|
@ -2059,22 +2094,22 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
self, completion_response: dict
|
||||
) -> tuple[
|
||||
str,
|
||||
list[Any] | None,
|
||||
list[object] | None,
|
||||
list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None,
|
||||
str | None,
|
||||
list[ChatCompletionToolCallChunk],
|
||||
list[Any] | None,
|
||||
list[Any] | None,
|
||||
list[Any] | None,
|
||||
list[object] | None,
|
||||
list[_AnthropicToolResultBlock] | None,
|
||||
list[object] | None,
|
||||
]:
|
||||
text_content = ""
|
||||
citations: list[Any] | None = None
|
||||
citations: list[object] | None = None
|
||||
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = None
|
||||
reasoning_content: str | None = None
|
||||
tool_calls: Final[list[ChatCompletionToolCallChunk]] = []
|
||||
web_search_results: list[Any] | None = None
|
||||
tool_results: list[Any] | None = None
|
||||
compaction_blocks: list[Any] | None = None
|
||||
web_search_results: list[object] | None = None
|
||||
tool_results: list[_AnthropicToolResultBlock] | None = None
|
||||
compaction_blocks: list[object] | None = None
|
||||
for idx, content in enumerate(completion_response["content"]):
|
||||
if content["type"] == "text":
|
||||
text_content += content["text"]
|
||||
|
|
@ -2284,7 +2319,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
raw_speed: Final = _usage.get("speed")
|
||||
resolved_speed: Final = raw_speed if isinstance(raw_speed, str) else speed
|
||||
|
||||
iterations: Final[list[Any] | None] = _usage.get("iterations")
|
||||
iterations: Final[Sequence[_AnthropicUsageIteration] | None] = _usage.get("iterations")
|
||||
if iterations:
|
||||
prompt_tokens = sum(it.get("input_tokens", 0) or 0 for it in iterations)
|
||||
completion_tokens = sum(it.get("output_tokens", 0) or 0 for it in iterations)
|
||||
|
|
@ -2377,7 +2412,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
def _build_code_interpreter_results(
|
||||
self,
|
||||
tool_results: list[Any],
|
||||
tool_results: Sequence[_AnthropicToolResultBlock],
|
||||
code_by_id: dict[str, str],
|
||||
container_id: str | None,
|
||||
) -> list[OutputCodeInterpreterCall]:
|
||||
|
|
@ -2403,11 +2438,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
def _build_provider_specific_fields(
|
||||
self,
|
||||
completion_response: dict,
|
||||
citations: list[Any] | None,
|
||||
citations: Sequence[object] | None,
|
||||
thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None,
|
||||
web_search_results: list[Any] | None,
|
||||
tool_results: list[Any] | None,
|
||||
compaction_blocks: list[Any] | None,
|
||||
web_search_results: Sequence[object] | None,
|
||||
tool_results: Sequence[_AnthropicToolResultBlock] | None,
|
||||
compaction_blocks: Sequence[object] | None,
|
||||
tool_calls: list[ChatCompletionToolCallChunk],
|
||||
) -> dict[str, Any]:
|
||||
provider_specific_fields: Final[dict[str, Any]] = {
|
||||
|
|
|
|||
|
|
@ -865,13 +865,9 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
f"Failed to fetch models from Anthropic. Status code: {response.status_code}, Response: {response.text}"
|
||||
)
|
||||
|
||||
models: Final = response.json()["data"]
|
||||
models: Final[Sequence[Mapping[str, str]]] = response.json()["data"]
|
||||
|
||||
litellm_model_names: Final = []
|
||||
for model in models:
|
||||
stripped_model_name = model["id"]
|
||||
litellm_model_name = "anthropic/" + stripped_model_name
|
||||
litellm_model_names.append(litellm_model_name)
|
||||
litellm_model_names: Final = ["anthropic/" + model["id"] for model in models]
|
||||
return litellm_model_names
|
||||
|
||||
def get_token_counter(self) -> BaseTokenCounter | None:
|
||||
|
|
@ -1077,7 +1073,7 @@ def strip_empty_content_blocks_from_anthropic_messages(
|
|||
return out
|
||||
|
||||
|
||||
def _is_empty_text_block(block: Any) -> bool:
|
||||
def _is_empty_text_block(block: object) -> bool:
|
||||
if not isinstance(block, dict) or block.get("type") != "text":
|
||||
return False
|
||||
text: Final = block.get("text")
|
||||
|
|
@ -1131,7 +1127,7 @@ def normalize_anthropic_tool_use_id(raw_id: str) -> str:
|
|||
return sanitized or "tool_use_id"
|
||||
|
||||
|
||||
def _sanitize_tool_use_id_content_block(block: Any) -> Any:
|
||||
def _sanitize_tool_use_id_content_block(block: object) -> object:
|
||||
if not isinstance(block, dict):
|
||||
return block
|
||||
block_type: Final = block.get("type")
|
||||
|
|
|
|||
|
|
@ -18,6 +18,24 @@ TOOL_NAME_PREFIX_LENGTH: Final = OPENAI_MAX_TOOL_NAME_LENGTH - TOOL_NAME_HASH_LE
|
|||
PROVIDERS_PROXYING_AN_UNKNOWN_BACKEND: Final = frozenset({"litellm_proxy"})
|
||||
|
||||
|
||||
def _optional_attr(source: object, name: str) -> object:
|
||||
return getattr(source, name, None)
|
||||
|
||||
|
||||
def _as_string_mapping(value: object) -> Mapping[str, object] | None:
|
||||
if isinstance(value, Mapping):
|
||||
return value
|
||||
return None
|
||||
|
||||
|
||||
def _thought_signature(provider_specific_fields: object) -> str | None:
|
||||
fields: Final = _as_string_mapping(provider_specific_fields)
|
||||
if fields is None:
|
||||
return None
|
||||
signature: Final = fields.get("thought_signature")
|
||||
return signature if isinstance(signature, str) else None
|
||||
|
||||
|
||||
_ANTHROPIC_TOOL_SCHEMA_KEYS: Final = frozenset(
|
||||
{"name", "type", "input_schema", "description", "cache_control", "strict"}
|
||||
)
|
||||
|
|
@ -56,7 +74,7 @@ def truncate_tool_name(name: str) -> str:
|
|||
|
||||
|
||||
def create_tool_name_mapping(
|
||||
tools: list[dict[str, Any]],
|
||||
tools: Sequence[Mapping[str, object]],
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
Create a mapping of truncated tool names to original names.
|
||||
|
|
@ -70,6 +88,8 @@ def create_tool_name_mapping(
|
|||
mapping: Final[dict[str, str]] = {}
|
||||
for tool in tools:
|
||||
original_name = tool.get("name", "")
|
||||
if not isinstance(original_name, str):
|
||||
continue
|
||||
truncated_name = truncate_tool_name(original_name)
|
||||
if truncated_name != original_name:
|
||||
mapping[truncated_name] = original_name
|
||||
|
|
@ -286,44 +306,44 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
|
||||
### FOR [BETA] `/v1/messages` endpoint support
|
||||
|
||||
def _extract_signature_from_tool_call(self, tool_call: Any) -> str | None:
|
||||
def _extract_signature_from_tool_call(self, tool_call: object) -> str | None:
|
||||
"""
|
||||
Extract signature from a tool call's provider_specific_fields.
|
||||
Only checks provider_specific_fields, not thinking blocks.
|
||||
"""
|
||||
signature = None
|
||||
fields: Final = _optional_attr(tool_call, "provider_specific_fields")
|
||||
if fields:
|
||||
return _thought_signature(fields)
|
||||
|
||||
if hasattr(tool_call, "provider_specific_fields") and tool_call.provider_specific_fields:
|
||||
if "thought_signature" in tool_call.provider_specific_fields:
|
||||
signature = tool_call.provider_specific_fields["thought_signature"]
|
||||
elif hasattr(tool_call.function, "provider_specific_fields") and tool_call.function.provider_specific_fields:
|
||||
if "thought_signature" in tool_call.function.provider_specific_fields:
|
||||
signature = tool_call.function.provider_specific_fields["thought_signature"]
|
||||
function_fields: Final = _optional_attr(_optional_attr(tool_call, "function"), "provider_specific_fields")
|
||||
if function_fields:
|
||||
return _thought_signature(function_fields)
|
||||
|
||||
return signature
|
||||
return None
|
||||
|
||||
def _extract_signature_from_tool_use_content(self, content: dict[str, Any]) -> str | None:
|
||||
def _extract_signature_from_tool_use_content(self, content: Mapping[str, object]) -> str | None:
|
||||
"""
|
||||
Extract signature from a tool_use content block's provider_specific_fields.
|
||||
"""
|
||||
provider_specific_fields: Final = content.get("provider_specific_fields", {})
|
||||
provider_specific_fields: Final = _as_string_mapping(content.get("provider_specific_fields", {}))
|
||||
if provider_specific_fields:
|
||||
return provider_specific_fields.get("signature")
|
||||
signature: Final = provider_specific_fields.get("signature")
|
||||
return signature if isinstance(signature, str) else None
|
||||
return None
|
||||
|
||||
def _add_cache_control_if_applicable(
|
||||
self,
|
||||
source: Any,
|
||||
target: Any,
|
||||
source: object,
|
||||
target: object,
|
||||
model: str | None,
|
||||
) -> None:
|
||||
"""
|
||||
Extract cache_control from source and add to target if it should be preserved.
|
||||
|
||||
This method accepts Any type to support both regular dicts and TypedDict objects.
|
||||
TypedDict objects (like ChatCompletionTextObject, ChatCompletionImageObject, etc.)
|
||||
are dicts at runtime but have specific types at type-check time. Using Any allows
|
||||
this method to work with both while maintaining runtime correctness.
|
||||
This method accepts an unconstrained type to support both regular dicts and
|
||||
TypedDict objects. TypedDict objects (like ChatCompletionTextObject,
|
||||
ChatCompletionImageObject, etc.) are dicts at runtime but have specific types at
|
||||
type-check time, so the widest parameter type works with both.
|
||||
|
||||
Args:
|
||||
source: Dict or TypedDict containing potential cache_control field
|
||||
|
|
@ -751,7 +771,7 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
|
||||
return new_tools, tool_name_mapping
|
||||
|
||||
def translate_anthropic_output_format_to_openai(self, output_format: Any) -> dict[str, object] | None:
|
||||
def translate_anthropic_output_format_to_openai(self, output_format: object) -> dict[str, object] | None:
|
||||
"""
|
||||
Translate Anthropic's output_format to OpenAI's response_format.
|
||||
|
||||
|
|
@ -1366,7 +1386,7 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
|
||||
@classmethod
|
||||
def _first_positive_prompt_tokens_detail_value(cls, usage: Usage, field_names: tuple[str, ...]) -> int:
|
||||
prompt_tokens_details: Final = getattr(usage, "prompt_tokens_details", None)
|
||||
prompt_tokens_details: Final = _optional_attr(usage, "prompt_tokens_details")
|
||||
if prompt_tokens_details is None:
|
||||
return 0
|
||||
|
||||
|
|
@ -1374,7 +1394,7 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
if isinstance(prompt_tokens_details, dict):
|
||||
value = cls._positive_int(prompt_tokens_details.get(field_name))
|
||||
else:
|
||||
value = cls._positive_int(getattr(prompt_tokens_details, field_name, None))
|
||||
value = cls._positive_int(_optional_attr(prompt_tokens_details, field_name))
|
||||
if value > 0:
|
||||
return value
|
||||
return 0
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@ Mirrors Anthropic's native ``compact_20260112`` for non-Anthropic providers:
|
|||
|
||||
import re
|
||||
from collections.abc import Awaitable, Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypeVar, Union, cast
|
||||
from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, TypeVar, Union, cast
|
||||
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict, Unpack
|
||||
|
||||
|
|
@ -232,7 +232,7 @@ async def _check_summary_model_access(
|
|||
|
||||
key_models: Final = list(getattr(user_api_key_auth, "models", None) or [])
|
||||
team_id: Final[str | None] = getattr(user_api_key_auth, "team_id", None)
|
||||
team_model_aliases: Final = getattr(user_api_key_auth, "team_model_aliases", None)
|
||||
team_model_aliases: Final[dict[str, str] | None] = getattr(user_api_key_auth, "team_model_aliases", None)
|
||||
team_models: Final = list(getattr(user_api_key_auth, "team_models", None) or [])
|
||||
user_id: Final[str | None] = getattr(user_api_key_auth, "user_id", None)
|
||||
project_id: Final[str | None] = getattr(user_api_key_auth, "project_id", None)
|
||||
|
|
@ -443,7 +443,9 @@ async def _check_summary_model_budget(
|
|||
)
|
||||
return False
|
||||
|
||||
end_user_model_max_budget: Final = getattr(user_api_key_auth, "end_user_model_max_budget", None)
|
||||
end_user_model_max_budget: Final[dict[str, object] | None] = getattr(
|
||||
user_api_key_auth, "end_user_model_max_budget", None
|
||||
)
|
||||
end_user_id: Final[str | None] = getattr(user_api_key_auth, "end_user_id", None)
|
||||
if isinstance(end_user_model_max_budget, dict) and end_user_model_max_budget and end_user_id is not None:
|
||||
try:
|
||||
|
|
@ -854,8 +856,8 @@ def _extract_summary_text(raw: str | None) -> str | None:
|
|||
|
||||
|
||||
def _system_to_openai_message(
|
||||
system: str | list[dict[str, Any]] | None,
|
||||
) -> Mapping[str, object] | None:
|
||||
system: str | list[dict[str, object]] | None,
|
||||
) -> dict[str, object] | None:
|
||||
"""Translate Anthropic-shaped ``system`` to an OpenAI system message.
|
||||
|
||||
Accepts a bare string or a list of Anthropic content blocks; returns
|
||||
|
|
@ -866,10 +868,10 @@ def _system_to_openai_message(
|
|||
if isinstance(system, str):
|
||||
return {"role": "system", "content": system} if system else None
|
||||
if isinstance(system, list):
|
||||
parts: Final[tuple[str, ...]] = tuple(
|
||||
parts: Final[list[object]] = [
|
||||
block.get("text", "") for block in system if isinstance(block, dict) and block.get("type") == "text"
|
||||
)
|
||||
joined: Final = "\n\n".join(part for part in parts if part)
|
||||
]
|
||||
joined: Final = "\n\n".join(part for part in parts if isinstance(part, str) and part)
|
||||
return {"role": "system", "content": joined} if joined else None
|
||||
return None
|
||||
|
||||
|
|
@ -951,7 +953,7 @@ async def _call_summary_model(
|
|||
summary_model: str,
|
||||
summary_messages: Sequence[Mapping[str, object]],
|
||||
metadata: Mapping[str, object],
|
||||
llm_router: object,
|
||||
llm_router: Optional["Router"],
|
||||
allowed_model_region: str | None = None,
|
||||
max_tokens: int = COMPACT_SUMMARY_MAX_TOKENS,
|
||||
) -> Union["ModelResponse", "CustomStreamWrapper"]:
|
||||
|
|
@ -1036,10 +1038,9 @@ def _extract_usage(response: object) -> tuple[int, int]:
|
|||
usage: Final[object] = getattr(response, "usage", None)
|
||||
if usage is None:
|
||||
return 0, 0
|
||||
return (
|
||||
int(getattr(usage, "prompt_tokens", 0) or 0),
|
||||
int(getattr(usage, "completion_tokens", 0) or 0),
|
||||
)
|
||||
prompt_tokens: Final[int | None] = getattr(usage, "prompt_tokens", 0)
|
||||
completion_tokens: Final[int | None] = getattr(usage, "completion_tokens", 0)
|
||||
return int(prompt_tokens or 0), int(completion_tokens or 0)
|
||||
|
||||
|
||||
def apply_client_compaction_block_history(
|
||||
|
|
|
|||
|
|
@ -8,6 +8,10 @@ import httpx
|
|||
from pydantic import TypeAdapter
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from litellm.constants import (
|
||||
ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS,
|
||||
ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import process_response_headers
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
|
|
@ -21,6 +25,9 @@ from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
|
|||
|
||||
GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ: Final = PassThroughEndpointLogging()
|
||||
|
||||
_UPSTREAM_PUMP_TASKS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: stdlib strong-ref set for pump tasks
|
||||
_DETACHED_STREAM_DRAINS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: bounded strong-ref set, detached drains
|
||||
|
||||
INCOMPLETE_STREAM_ERROR_MESSAGE: Final = (
|
||||
"Provider stream ended before emitting a message_stop event; "
|
||||
"the response is incomplete and any partial content (e.g. tool_use input JSON) may be truncated."
|
||||
|
|
@ -133,6 +140,34 @@ def _is_terminal_stream_chunk(chunk: object) -> bool:
|
|||
return _is_message_stop_chunk(chunk) or _is_provider_error_chunk(chunk)
|
||||
|
||||
|
||||
def _try_claim_detached_drain_slot() -> bool:
|
||||
"""Claim a detached-drain slot for the current task, bounding concurrency.
|
||||
|
||||
Returns True if a slot was claimed (the caller may keep draining upstream
|
||||
for billing) or False if the cap is already reached (the caller should stop
|
||||
and bill what it has). Only touched from the event loop, so the check +
|
||||
insert need no lock.
|
||||
"""
|
||||
if len(_DETACHED_STREAM_DRAINS) >= ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS:
|
||||
return False
|
||||
current_task: Final = asyncio.current_task()
|
||||
if current_task is not None:
|
||||
_DETACHED_STREAM_DRAINS.add(current_task)
|
||||
current_task.add_done_callback(_DETACHED_STREAM_DRAINS.discard)
|
||||
return True
|
||||
|
||||
|
||||
def _exception_left_unconsumed(queue: "asyncio.Queue[bytes | None | BaseException]", exc: BaseException) -> bool:
|
||||
"""After client detach the relay never reads the queue again, so drain it here.
|
||||
|
||||
The forwarded exception still sitting in the queue means the relay tore
|
||||
down before re-raising it, so the proxy's failure handling never ran and
|
||||
the caller must salvage spend itself.
|
||||
"""
|
||||
remaining: Final = tuple(queue.get_nowait() for _ in range(queue.qsize()))
|
||||
return any(item is exc for item in remaining)
|
||||
|
||||
|
||||
def _sse_event(event_type: str, payload: Mapping[str, object]) -> bytes:
|
||||
return f"event: {event_type}\ndata: {json.dumps(payload)}\n\n".encode()
|
||||
|
||||
|
|
@ -414,17 +449,167 @@ class BaseAnthropicMessagesStreamingIterator:
|
|||
|
||||
async def async_sse_wrapper(
|
||||
self,
|
||||
completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | dict],
|
||||
completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | Mapping[str, object]],
|
||||
) -> AsyncIterator[bytes]:
|
||||
"""
|
||||
Generic async SSE wrapper that converts streaming chunks to SSE format
|
||||
and handles logging.
|
||||
|
||||
The upstream read runs in a detached background task (``_pump_upstream``)
|
||||
so that a client disconnect tears down only this client-facing generator,
|
||||
never the upstream drain + billing. The provider (e.g. Bedrock) keeps
|
||||
generating and billing the full response regardless of the client, so
|
||||
draining it to completion is what lets spend tracking see the real
|
||||
terminal ``message_delta`` / ``message_stop`` usage instead of a
|
||||
truncated placeholder count.
|
||||
|
||||
Chunks reach the client through a bounded queue. While the client is
|
||||
connected the pump blocks on a full queue (racing the disconnect
|
||||
signal), so a slow reader throttles the upstream read exactly as the old
|
||||
direct ``yield`` did instead of letting the whole response buffer in
|
||||
memory. Once the client goes away the pump stops enqueueing and only
|
||||
keeps a single ``collected_chunks`` copy for billing, and the number of
|
||||
such post-disconnect drains running at once is capped so client behavior
|
||||
can't create unbounded worker state; over the cap the pump bills what it
|
||||
has rather than draining further. Detached-drain lifetime is otherwise
|
||||
bounded by the upstream stream/read timeout.
|
||||
|
||||
An upstream failure (Bedrock read / decode / chunk-conversion error)
|
||||
that happens while the client is still connected is forwarded through
|
||||
the queue and re-raised here, so the original provider exception (and
|
||||
its status) reaches the proxy's failure handling unchanged rather than
|
||||
being masked by a generic incomplete-stream event.
|
||||
|
||||
This method provides the common logic for both Anthropic and Bedrock implementations.
|
||||
"""
|
||||
collected_chunks: Final = []
|
||||
saw_terminal_event = False
|
||||
queue: Final[asyncio.Queue[bytes | None | BaseException]] = asyncio.Queue(
|
||||
maxsize=ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE
|
||||
)
|
||||
client_detached: Final = asyncio.Event()
|
||||
|
||||
pump_task: Final = asyncio.create_task(self._pump_upstream_to_queue(completion_stream, queue, client_detached))
|
||||
_UPSTREAM_PUMP_TASKS.add(pump_task)
|
||||
pump_task.add_done_callback(_UPSTREAM_PUMP_TASKS.discard)
|
||||
|
||||
reached_end = False # rebind-ok: flipped once the relay consumes the end-of-stream sentinel
|
||||
try:
|
||||
while True:
|
||||
item = await queue.get()
|
||||
if item is None:
|
||||
reached_end = True
|
||||
break
|
||||
if isinstance(item, BaseException):
|
||||
raise item
|
||||
yield item
|
||||
finally:
|
||||
client_detached.set()
|
||||
if not reached_end:
|
||||
self._dispatch_pending_deferred_logging()
|
||||
|
||||
def _dispatch_pending_deferred_logging(self) -> None:
|
||||
"""Fire deferred billing that a torn-down response would otherwise drop.
|
||||
|
||||
When the pump finishes draining while the client is still connected it
|
||||
stores the logging coroutine for ProxyLogging._fire_deferred_stream_logging,
|
||||
which the proxy only fires on a normally completed response: a client
|
||||
disconnect (GeneratorExit / CancelledError) re-raises past it. Without
|
||||
this dispatch that window loses the spend row entirely.
|
||||
"""
|
||||
deferred_cb: Final = getattr(self.litellm_logging_obj, "_on_deferred_stream_complete", None)
|
||||
deferred_args: Final = getattr(self.litellm_logging_obj, "_deferred_stream_complete_args", None)
|
||||
if deferred_cb is None or deferred_args is None:
|
||||
return
|
||||
self.litellm_logging_obj._on_deferred_stream_complete = None
|
||||
self.litellm_logging_obj._deferred_stream_complete_args = None
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(async_coroutine=deferred_cb(*deferred_args))
|
||||
|
||||
async def _bill_collected_chunks(
|
||||
self,
|
||||
collected_chunks: list[bytes], # mutable-ok: SSE buffer forwarded to list-typed _handle_streaming_logging
|
||||
*,
|
||||
stream_teardown: bool,
|
||||
) -> None:
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
try:
|
||||
await self._handle_streaming_logging(collected_chunks, stream_teardown=stream_teardown)
|
||||
except Exception as exc: # noqa: BLE001 # billing is best-effort; never crash the pump
|
||||
verbose_proxy_logger.warning(
|
||||
"async_sse_wrapper billing failed after %d chunks: %s(%s)",
|
||||
len(collected_chunks),
|
||||
type(exc).__name__,
|
||||
exc,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _abort_upstream(
|
||||
completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | Mapping[str, object]],
|
||||
) -> None:
|
||||
"""Close the upstream provider stream so it stops generating and billing."""
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
try:
|
||||
await aclose_if_supported(completion_stream)
|
||||
except Exception as exc: # noqa: BLE001 # abort is best-effort; log and continue
|
||||
verbose_proxy_logger.warning(
|
||||
"async_sse_wrapper failed to abort upstream stream: %s(%s)",
|
||||
type(exc).__name__,
|
||||
exc,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _enqueue_for_client(
|
||||
queue: "asyncio.Queue[bytes | None | BaseException]",
|
||||
client_detached: "asyncio.Event",
|
||||
item: bytes | None | BaseException,
|
||||
) -> bool:
|
||||
"""Deliver one item to the client, applying backpressure.
|
||||
|
||||
Returns True if the item was queued, False if the client disconnected
|
||||
before there was room (the item is then dropped, since a gone client
|
||||
can't receive it). Never blocks once the client has detached.
|
||||
"""
|
||||
if client_detached.is_set():
|
||||
return False
|
||||
try:
|
||||
queue.put_nowait(item)
|
||||
except asyncio.QueueFull:
|
||||
pass
|
||||
else:
|
||||
return True
|
||||
put_task: Final = asyncio.ensure_future(queue.put(item))
|
||||
detached_task: Final = asyncio.ensure_future(client_detached.wait())
|
||||
try:
|
||||
await asyncio.wait(frozenset((put_task, detached_task)), return_when=asyncio.FIRST_COMPLETED)
|
||||
finally:
|
||||
if not detached_task.done():
|
||||
detached_task.cancel()
|
||||
if put_task.done() and not put_task.cancelled():
|
||||
return True
|
||||
put_task.cancel()
|
||||
return False
|
||||
|
||||
async def _pump_upstream_to_queue(
|
||||
self,
|
||||
completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | Mapping[str, object]],
|
||||
queue: "asyncio.Queue[bytes | None | BaseException]",
|
||||
client_detached: "asyncio.Event",
|
||||
) -> None:
|
||||
"""Drain the whole upstream into ``queue`` (backpressured) and bill once.
|
||||
|
||||
Runs detached so a client disconnect can't interrupt the upstream read;
|
||||
see ``async_sse_wrapper`` for the full rationale. On a completed drain
|
||||
the success billing (or deferred park) happens before the end-of-stream
|
||||
sentinel is enqueued: the relay can only tear down after consuming the
|
||||
sentinel, so its teardown can never outrun the park and get mistaken
|
||||
for a client disconnect, and a sentinel the client never consumes falls
|
||||
back to dispatching the parked billing here.
|
||||
"""
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
collected_chunks: Final[list[bytes]] = [] # mutable-ok: SSE billing buffer appended to across the drain
|
||||
saw_terminal_event = False # rebind-ok: accumulates across the upstream loop
|
||||
draining_detached = False # rebind-ok: set once this pump claims a detached-drain slot
|
||||
try:
|
||||
async for chunk in completion_stream:
|
||||
if self.completion_start_time is None:
|
||||
|
|
@ -432,17 +617,62 @@ class BaseAnthropicMessagesStreamingIterator:
|
|||
saw_terminal_event = saw_terminal_event or _is_terminal_stream_chunk(chunk)
|
||||
encoded_chunk = self._convert_chunk_to_sse_format(chunk)
|
||||
collected_chunks.append(encoded_chunk)
|
||||
yield encoded_chunk
|
||||
except (GeneratorExit, asyncio.CancelledError):
|
||||
# A client disconnect tears the generator down at the yield, so the
|
||||
# post-loop logging below never runs and the tokens already streamed
|
||||
# (and billed by the provider) would never reach spend tracking. See LIT-5839.
|
||||
if collected_chunks:
|
||||
await self._handle_streaming_logging(collected_chunks, stream_teardown=True)
|
||||
raise
|
||||
if not client_detached.is_set():
|
||||
await self._enqueue_for_client(queue, client_detached, encoded_chunk)
|
||||
continue
|
||||
if not draining_detached:
|
||||
if not _try_claim_detached_drain_slot():
|
||||
verbose_proxy_logger.warning(
|
||||
"async_sse_wrapper: detached-drain cap (%d) reached; billing %d partial "
|
||||
"chunks and aborting the upstream stream to stop provider billing",
|
||||
ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS,
|
||||
len(collected_chunks),
|
||||
)
|
||||
await self._bill_collected_chunks(collected_chunks, stream_teardown=True)
|
||||
await self._abort_upstream(completion_stream)
|
||||
return
|
||||
draining_detached = True
|
||||
except Exception as exc: # noqa: BLE001 # upstream errors are handled/forwarded by _handle_pump_upstream_error
|
||||
await self._handle_pump_upstream_error(queue, client_detached, collected_chunks, exc)
|
||||
return
|
||||
|
||||
if not saw_terminal_event:
|
||||
yield _incomplete_stream_error_sse_event()
|
||||
if client_detached.is_set():
|
||||
await self._bill_collected_chunks(collected_chunks, stream_teardown=True)
|
||||
return
|
||||
if not saw_terminal_event and not await self._enqueue_for_client(
|
||||
queue, client_detached, _incomplete_stream_error_sse_event()
|
||||
):
|
||||
await self._bill_collected_chunks(collected_chunks, stream_teardown=True)
|
||||
return
|
||||
await self._bill_collected_chunks(collected_chunks, stream_teardown=False)
|
||||
if not await self._enqueue_for_client(queue, client_detached, None):
|
||||
self._dispatch_pending_deferred_logging()
|
||||
|
||||
# Handle logging after all chunks are processed
|
||||
await self._handle_streaming_logging(collected_chunks)
|
||||
async def _handle_pump_upstream_error(
|
||||
self,
|
||||
queue: "asyncio.Queue[bytes | None | BaseException]",
|
||||
client_detached: "asyncio.Event",
|
||||
collected_chunks: list[bytes], # mutable-ok: SSE buffer forwarded to list-typed _bill_collected_chunks
|
||||
exc: BaseException,
|
||||
) -> None:
|
||||
"""Forward a provider error to a still-connected client, else salvage partial spend.
|
||||
|
||||
Handing the original exception to the client-facing generator lets it
|
||||
re-raise so the proxy's failure handling keeps the provider status and
|
||||
owns logging (no success-bill). If the client already went away, or
|
||||
disconnects before ever consuming the queued exception, no failure hook
|
||||
runs, so bill the partial instead of dropping the request.
|
||||
"""
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
if not client_detached.is_set() and await self._enqueue_for_client(queue, client_detached, exc):
|
||||
await client_detached.wait()
|
||||
if not _exception_left_unconsumed(queue, exc):
|
||||
return
|
||||
verbose_proxy_logger.warning(
|
||||
"async_sse_wrapper upstream pump failed after client disconnect (%d chunks): %s(%s)",
|
||||
len(collected_chunks),
|
||||
type(exc).__name__,
|
||||
exc,
|
||||
)
|
||||
await self._bill_collected_chunks(collected_chunks, stream_teardown=True)
|
||||
|
|
|
|||
|
|
@ -179,14 +179,14 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def _assistant_block_group_key(indexed_block: tuple[int, Mapping[str, Any]]) -> str:
|
||||
def _assistant_block_group_key(indexed_block: tuple[int, Mapping[str, object]]) -> str:
|
||||
"""Group a run of consecutive thinking blocks together; keep every other block alone."""
|
||||
index, block = indexed_block
|
||||
return "thinking" if block.get("type") == "thinking" else f"block:{index}"
|
||||
|
||||
@classmethod
|
||||
def _assistant_group_to_input_item(
|
||||
cls, group: tuple[Mapping[str, Any], ...]
|
||||
cls, group: tuple[Mapping[str, object], ...]
|
||||
) -> dict[str, Any] | None: # mutable-ok: API message payload
|
||||
first: Final = group[0]
|
||||
btype: Final = first.get("type")
|
||||
|
|
@ -206,7 +206,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
def translate_messages_to_responses_input(
|
||||
self,
|
||||
messages: list[AllAnthropicPassThroughMessageValues],
|
||||
) -> list[dict[str, Any]]:
|
||||
) -> list[dict[str, object]]:
|
||||
"""
|
||||
Convert Anthropic messages list to Responses API `input` items.
|
||||
|
||||
|
|
@ -220,7 +220,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
assistant thinking -> reasoning
|
||||
assistant tool_use -> function_call
|
||||
"""
|
||||
input_items: Final[list[dict[str, Any]]] = []
|
||||
input_items: Final[list[dict[str, object]]] = []
|
||||
|
||||
for m in messages:
|
||||
if m["role"] == "system":
|
||||
|
|
@ -248,7 +248,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
}
|
||||
)
|
||||
elif isinstance(content, list):
|
||||
user_parts: list[dict[str, Any]] = []
|
||||
user_parts: list[Mapping[str, object]] = []
|
||||
tool_image_parts: list[dict[str, Any]] = [] # mutable-ok: json content parts
|
||||
for block in content:
|
||||
if not isinstance(block, dict):
|
||||
|
|
@ -379,9 +379,9 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
def translate_tools_to_responses_api(
|
||||
self,
|
||||
tools: list[AllAnthropicToolsValues],
|
||||
) -> list[dict[str, Any]]:
|
||||
) -> list[dict[str, object]]:
|
||||
"""Convert Anthropic tool definitions to Responses API function tools."""
|
||||
result: Final[list[dict[str, Any]]] = []
|
||||
result: Final[list[dict[str, object]]] = []
|
||||
for tool in tools:
|
||||
tool_dict = cast(dict[str, Any], tool)
|
||||
tool_type = tool_dict.get("type", "")
|
||||
|
|
@ -392,7 +392,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
continue
|
||||
# Responses turns strict mode on when `strict` is omitted, silently rewriting
|
||||
# `required` to every property. Anthropic tools are non-strict unless asked.
|
||||
func_tool: dict[str, Any] = {
|
||||
func_tool: dict[str, object] = {
|
||||
"type": "function",
|
||||
"name": tool_name,
|
||||
"strict": bool(tool_dict.get("strict")),
|
||||
|
|
@ -407,7 +407,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
@staticmethod
|
||||
def translate_tool_choice_to_responses_api(
|
||||
tool_choice: AnthropicMessagesToolChoice,
|
||||
) -> str | dict[str, Any]:
|
||||
) -> str | dict[str, object]:
|
||||
"""Convert Anthropic tool_choice to Responses API tool_choice."""
|
||||
tc_type: Final = tool_choice.get("type")
|
||||
if tc_type == "any":
|
||||
|
|
@ -420,8 +420,8 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
|
||||
@staticmethod
|
||||
def translate_context_management_to_responses_api(
|
||||
context_management: dict[str, Any],
|
||||
) -> list[dict[str, Any]] | None:
|
||||
context_management: dict[str, object],
|
||||
) -> list[dict[str, object]] | None:
|
||||
"""
|
||||
Convert Anthropic context_management dict to OpenAI Responses API array format.
|
||||
|
||||
|
|
@ -435,13 +435,13 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
if not isinstance(edits, list):
|
||||
return None
|
||||
|
||||
result: Final[list[dict[str, Any]]] = []
|
||||
result: Final[list[dict[str, object]]] = []
|
||||
for edit in edits:
|
||||
if not isinstance(edit, dict):
|
||||
continue
|
||||
edit_type = edit.get("type", "")
|
||||
if edit_type == "compact_20260112":
|
||||
entry: dict[str, Any] = {"type": "compaction"}
|
||||
entry: dict[str, object] = {"type": "compaction"}
|
||||
trigger = edit.get("trigger")
|
||||
if isinstance(trigger, dict) and trigger.get("value") is not None:
|
||||
entry["compact_threshold"] = int(trigger["value"])
|
||||
|
|
@ -451,9 +451,9 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
|
||||
@staticmethod
|
||||
def translate_thinking_to_reasoning(
|
||||
thinking: dict[str, Any],
|
||||
output_config: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any] | None:
|
||||
thinking: dict[str, object],
|
||||
output_config: dict[str, object] | None = None,
|
||||
) -> dict[str, object] | None:
|
||||
"""
|
||||
Convert Anthropic thinking param to Responses API reasoning param.
|
||||
|
||||
|
|
@ -473,12 +473,14 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
if isinstance(output_config, dict) and output_config.get("effort"):
|
||||
effort = output_config["effort"]
|
||||
elif thinking_type == "enabled":
|
||||
effort = reasoning_effort_from_thinking_budget(thinking.get("budget_tokens", 0))
|
||||
raw_budget: Final = thinking.get("budget_tokens", 0)
|
||||
budget_tokens: Final = int(raw_budget) if isinstance(raw_budget, (int, float)) else 0
|
||||
effort = reasoning_effort_from_thinking_budget(budget_tokens)
|
||||
else:
|
||||
return None
|
||||
|
||||
auto_summary: Final = is_reasoning_auto_summary_enabled()
|
||||
result: Final[dict[str, Any]] = {"effort": effort}
|
||||
result: Final[dict[str, object]] = {"effort": effort}
|
||||
summary: Final = thinking.get("summary")
|
||||
if summary:
|
||||
result["summary"] = summary
|
||||
|
|
@ -570,7 +572,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
# output_format / output_config.format -> text format
|
||||
# output_format: {"type": "json_schema", "schema": {...}}
|
||||
# output_config: {"format": {"type": "json_schema", "schema": {...}}}
|
||||
output_format: Any = anthropic_request.get("output_format")
|
||||
output_format: object = anthropic_request.get("output_format")
|
||||
output_config = anthropic_request.get("output_config")
|
||||
if not isinstance(output_format, dict) and isinstance(output_config, dict):
|
||||
output_format = output_config.get("format")
|
||||
|
|
@ -620,7 +622,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
ResponseReasoningItem,
|
||||
)
|
||||
|
||||
content: Final[list[dict[str, Any]]] = []
|
||||
content: Final[list[dict[str, object]]] = []
|
||||
stop_reason: AnthropicFinishReason = "end_turn"
|
||||
|
||||
for item in response.output:
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ import asyncio
|
|||
import json
|
||||
import time
|
||||
from collections.abc import Coroutine
|
||||
from typing import Any, Final
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -116,7 +116,7 @@ class AnthropicFilesHandler:
|
|||
api_key: str | None = None,
|
||||
timeout: float | httpx.Timeout = 600.0,
|
||||
max_retries: int | None = None,
|
||||
) -> HttpxBinaryResponseContent | Coroutine[Any, Any, HttpxBinaryResponseContent]:
|
||||
) -> HttpxBinaryResponseContent | Coroutine[object, object, HttpxBinaryResponseContent]:
|
||||
"""
|
||||
Retrieve file content from Anthropic.
|
||||
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ import asyncio
|
|||
import json
|
||||
import time
|
||||
from collections.abc import Callable, Coroutine
|
||||
from typing import Any, Final
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from openai import (
|
||||
|
|
@ -374,7 +374,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
except Exception as e:
|
||||
status_code: Final = getattr(e, "status_code", 500)
|
||||
error_headers = getattr(e, "headers", None)
|
||||
error_response: Final = getattr(e, "response", None)
|
||||
error_response: Final[object] = getattr(e, "response", None)
|
||||
error_body: Final = getattr(e, "body", None)
|
||||
if error_headers is None and error_response:
|
||||
error_headers = getattr(error_response, "headers", None)
|
||||
|
|
@ -392,7 +392,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
model: str,
|
||||
api_base: str,
|
||||
data: dict,
|
||||
timeout: Any,
|
||||
timeout: float | httpx.Timeout,
|
||||
dynamic_params: bool,
|
||||
model_response: ModelResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
|
|
@ -502,7 +502,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
dynamic_params: bool,
|
||||
data: dict[str, object],
|
||||
model: str,
|
||||
timeout: Any,
|
||||
timeout: float | httpx.Timeout,
|
||||
max_retries: int,
|
||||
azure_ad_token: str | None = None,
|
||||
azure_ad_token_provider: Callable | None = None,
|
||||
|
|
@ -578,7 +578,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
dynamic_params: bool,
|
||||
data: dict,
|
||||
model: str,
|
||||
timeout: Any,
|
||||
timeout: float | httpx.Timeout,
|
||||
max_retries: int,
|
||||
azure_ad_token: str | None = None,
|
||||
azure_ad_token_provider: Callable | None = None,
|
||||
|
|
@ -634,7 +634,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
except Exception as e:
|
||||
status_code: Final = getattr(e, "status_code", 500)
|
||||
error_headers = getattr(e, "headers", None)
|
||||
error_response: Final = getattr(e, "response", None)
|
||||
error_response: Final[object] = getattr(e, "response", None)
|
||||
message: Final = getattr(e, "message", str(e))
|
||||
error_body: Final = getattr(e, "body", None)
|
||||
if error_headers is None and error_response:
|
||||
|
|
@ -754,7 +754,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
aembedding=None,
|
||||
headers: dict | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
) -> EmbeddingResponse | Coroutine[Any, Any, EmbeddingResponse]:
|
||||
) -> EmbeddingResponse | Coroutine[object, object, EmbeddingResponse]:
|
||||
if headers:
|
||||
optional_params["extra_headers"] = headers
|
||||
if self._client_session is None:
|
||||
|
|
@ -1268,7 +1268,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
headers["Authorization"] = f"Bearer {azure_ad_token}"
|
||||
|
||||
# init AzureOpenAI Client
|
||||
azure_client_params: Final[dict[str, Any]] = self.initialize_azure_sdk_client(
|
||||
azure_client_params: Final[dict[str, object]] = self.initialize_azure_sdk_client(
|
||||
litellm_params=litellm_params or {},
|
||||
api_key=api_key,
|
||||
model_name=model or "",
|
||||
|
|
|
|||
|
|
@ -51,15 +51,13 @@ else:
|
|||
AsyncHTTPHandler = Any
|
||||
|
||||
|
||||
class _AzureRawAnnotation(TypedDict, total=False):
|
||||
type: ReadOnly[str]
|
||||
class _AzureRawAnnotation(ChatCompletionAnnotation, total=False):
|
||||
text: ReadOnly[str]
|
||||
start_index: ReadOnly[int]
|
||||
end_index: ReadOnly[int]
|
||||
url_citation: ReadOnly[ChatCompletionAnnotationURLCitation]
|
||||
|
||||
|
||||
_TransformedAnnotation: TypeAlias = ChatCompletionAnnotation | _AzureRawAnnotation
|
||||
_TransformedAnnotation: TypeAlias = ChatCompletionAnnotation
|
||||
|
||||
|
||||
class _AzureText(TypedDict, total=False):
|
||||
|
|
@ -223,18 +221,11 @@ class AzureAIAgentsHandler:
|
|||
"""Build the ModelResponse from agent output."""
|
||||
from litellm.types.utils import Choices, Message, Usage
|
||||
|
||||
message_kwargs: Final[dict[str, Any]] = {
|
||||
"content": content,
|
||||
"role": "assistant",
|
||||
}
|
||||
if annotations:
|
||||
message_kwargs["annotations"] = annotations
|
||||
|
||||
model_response.choices = [
|
||||
Choices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
message=Message(**message_kwargs),
|
||||
message=Message(content=content, role="assistant", annotations=annotations or None),
|
||||
)
|
||||
]
|
||||
model_response.model = model
|
||||
|
|
@ -655,9 +646,6 @@ class AzureAIAgentsHandler:
|
|||
|
||||
if data_str == "[DONE]":
|
||||
# Send final chunk with finish_reason
|
||||
final_delta_kwargs: dict[str, Any] = {"content": None}
|
||||
if collected_annotations:
|
||||
final_delta_kwargs["annotations"] = collected_annotations
|
||||
final_chunk = ModelResponseStream(
|
||||
id=response_id,
|
||||
created=created,
|
||||
|
|
@ -667,7 +655,7 @@ class AzureAIAgentsHandler:
|
|||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=Delta(**final_delta_kwargs),
|
||||
delta=Delta(content=None, annotations=collected_annotations or None),
|
||||
)
|
||||
],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -5,7 +5,8 @@
|
|||
import base64
|
||||
import json
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import TYPE_CHECKING, Any, Final, Generic, TypeVar, cast
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, Generic, Protocol, TypeVar, cast, runtime_checkable
|
||||
|
||||
from litellm import verbose_logger
|
||||
from litellm.llms.base_llm.managed_resources.isolation import (
|
||||
|
|
@ -38,6 +39,30 @@ else:
|
|||
ResourceObjectType = TypeVar("ResourceObjectType")
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _HasIdentifier(Protocol):
|
||||
id: str
|
||||
|
||||
|
||||
class _ManagedResourceRecord(Protocol[ResourceObjectType]):
|
||||
unified_resource_id: str
|
||||
resource_object: ResourceObjectType
|
||||
|
||||
def model_dump(self) -> dict[str, object]: ...
|
||||
|
||||
|
||||
class _ManagedResourceTable(Protocol[ResourceObjectType]):
|
||||
async def create(self, *, data: Mapping[str, object]) -> object: ...
|
||||
|
||||
async def find_first(self, *, where: Mapping[str, object]) -> _ManagedResourceRecord[ResourceObjectType] | None: ...
|
||||
|
||||
async def find_many(
|
||||
self, *, where: Mapping[str, object], take: int, order: Mapping[str, str]
|
||||
) -> list[_ManagedResourceRecord[ResourceObjectType]]: ...
|
||||
|
||||
async def delete(self, *, where: Mapping[str, object]) -> object: ...
|
||||
|
||||
|
||||
class BaseManagedResource(ABC, Generic[ResourceObjectType]):
|
||||
"""
|
||||
Base class for managing resources with target_model_names support.
|
||||
|
|
@ -64,6 +89,9 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
|
|||
self.internal_usage_cache = internal_usage_cache
|
||||
self.prisma_client = prisma_client
|
||||
|
||||
def _resource_table(self) -> _ManagedResourceTable[ResourceObjectType]:
|
||||
return getattr(self.prisma_client.db, self.table_name)
|
||||
|
||||
# ============================================================================
|
||||
# ABSTRACT METHODS
|
||||
# ============================================================================
|
||||
|
|
@ -137,7 +165,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
|
|||
litellm_parent_otel_span: Span | None,
|
||||
model_mappings: dict[str, str],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
additional_db_fields: dict[str, Any] | None = None,
|
||||
additional_db_fields: Mapping[str, object] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Store unified resource ID with model mappings in cache and database.
|
||||
|
|
@ -153,7 +181,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
|
|||
verbose_logger.info("Storing LiteLLM Managed %s with id=%s in cache", self.resource_type, unified_resource_id)
|
||||
|
||||
# Prepare cache data
|
||||
cache_data: Final = {
|
||||
cache_data: Final[dict[str, object]] = {
|
||||
"unified_resource_id": unified_resource_id,
|
||||
"resource_object": resource_object,
|
||||
"model_mappings": model_mappings,
|
||||
|
|
@ -176,7 +204,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
|
|||
)
|
||||
|
||||
# Prepare database data
|
||||
db_data: Final = {
|
||||
db_data: Final[dict[str, object]] = {
|
||||
"unified_resource_id": unified_resource_id,
|
||||
"model_mappings": json.dumps(model_mappings),
|
||||
"flat_model_resource_ids": list(model_mappings.values()),
|
||||
|
|
@ -205,7 +233,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
|
|||
db_data.update(additional_db_fields)
|
||||
|
||||
# Store in database
|
||||
table: Final = getattr(self.prisma_client.db, self.table_name)
|
||||
table: Final = self._resource_table()
|
||||
result: Final = await table.create(data=db_data)
|
||||
|
||||
verbose_logger.debug(
|
||||
|
|
@ -240,7 +268,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
|
|||
return result
|
||||
|
||||
# Check database
|
||||
table: Final = getattr(self.prisma_client.db, self.table_name)
|
||||
table: Final = self._resource_table()
|
||||
db_object: Final = await table.find_first(where={"unified_resource_id": unified_resource_id})
|
||||
|
||||
if db_object:
|
||||
|
|
@ -264,7 +292,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
|
|||
The deleted resource object or None if not found
|
||||
"""
|
||||
# Get old value from database
|
||||
table: Final = getattr(self.prisma_client.db, self.table_name)
|
||||
table: Final = self._resource_table()
|
||||
initial_value: Final = await table.find_first(where={"unified_resource_id": unified_resource_id})
|
||||
|
||||
if initial_value is None:
|
||||
|
|
@ -515,7 +543,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
|
|||
user_api_key_dict: UserAPIKeyAuth,
|
||||
limit: int | None = None,
|
||||
after: str | None = None,
|
||||
additional_filters: dict[str, Any] | None = None,
|
||||
additional_filters: Mapping[str, object] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
List resources created by a user.
|
||||
|
|
@ -533,7 +561,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
|
|||
if owner_filter is None:
|
||||
return build_list_page([])
|
||||
|
||||
where_clause: Final[dict[str, Any]] = {**owner_filter}
|
||||
where_clause: Final[dict[str, object]] = {**owner_filter}
|
||||
|
||||
if after:
|
||||
where_clause["id"] = {"gt": after}
|
||||
|
|
@ -544,14 +572,14 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
|
|||
|
||||
# Fetch resources
|
||||
fetch_limit: Final = limit or 20
|
||||
table: Final = getattr(self.prisma_client.db, self.table_name)
|
||||
table: Final = self._resource_table()
|
||||
resources: Final = await table.find_many(
|
||||
where=where_clause,
|
||||
take=fetch_limit,
|
||||
order={"created_at": "desc"},
|
||||
)
|
||||
|
||||
resource_objects: Final[list[Any]] = []
|
||||
resource_objects: Final[list[object]] = []
|
||||
for resource in resources:
|
||||
try:
|
||||
# Stop once we have enough
|
||||
|
|
@ -559,12 +587,13 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
|
|||
break
|
||||
|
||||
# Parse resource object
|
||||
resource_data = resource.resource_object
|
||||
if isinstance(resource_data, str):
|
||||
resource_data = json.loads(resource_data)
|
||||
stored_resource = resource.resource_object
|
||||
resource_data: object = (
|
||||
json.loads(stored_resource) if isinstance(stored_resource, str) else stored_resource
|
||||
)
|
||||
|
||||
# Set unified ID
|
||||
if hasattr(resource_data, "id"):
|
||||
if isinstance(resource_data, _HasIdentifier):
|
||||
resource_data.id = resource.unified_resource_id
|
||||
elif isinstance(resource_data, dict):
|
||||
resource_data["id"] = resource.unified_resource_id
|
||||
|
|
|
|||
|
|
@ -75,6 +75,7 @@ class OCRUsageInfo(LiteLLMPydanticObjectBase):
|
|||
"""Usage information from OCR response."""
|
||||
|
||||
pages_processed: int | None = None
|
||||
pages_processed_annotation: int | None = None
|
||||
credits: float | None = None
|
||||
doc_size_bytes: int | None = None
|
||||
|
||||
|
|
|
|||
|
|
@ -3,7 +3,6 @@ from typing import TYPE_CHECKING, Any, Final
|
|||
import httpx
|
||||
|
||||
from litellm.anthropic_beta_headers_manager import filter_and_transform_beta_headers
|
||||
from litellm.litellm_core_utils.litellm_logging import verbose_logger
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
convert_to_anthropic_image_obj,
|
||||
)
|
||||
|
|
@ -16,17 +15,16 @@ from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation
|
|||
AmazonInvokeConfig,
|
||||
)
|
||||
from litellm.llms.bedrock.common_utils import (
|
||||
convert_bedrock_invoke_output_format_to_inline_schema,
|
||||
apply_bedrock_invoke_structured_output,
|
||||
get_anthropic_beta_from_headers,
|
||||
normalize_bedrock_opus_output_config_effort,
|
||||
normalize_custom_field_on_tools,
|
||||
normalize_tool_input_schema_types_for_bedrock_invoke,
|
||||
pop_bedrock_invoke_output_config_format,
|
||||
strip_unsupported_bedrock_invoke_output_config_keys,
|
||||
)
|
||||
from litellm.types.llms.anthropic import ANTHROPIC_TOOL_SEARCH_BETA_HEADER
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse
|
||||
from litellm.utils import _supports_factory
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
|
@ -212,36 +210,14 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
|
|||
anthropic_request.pop("model", None)
|
||||
anthropic_request.pop("stream", None)
|
||||
anthropic_request.pop("stream_chunk_size", None)
|
||||
output_format: Final = anthropic_request.pop("output_format", None)
|
||||
output_config_format: Final = pop_bedrock_invoke_output_config_format(anthropic_request)
|
||||
if output_format:
|
||||
convert_bedrock_invoke_output_format_to_inline_schema(
|
||||
output_format=output_format,
|
||||
request_body=anthropic_request,
|
||||
)
|
||||
elif output_config_format:
|
||||
convert_bedrock_invoke_output_format_to_inline_schema(
|
||||
output_format=output_config_format,
|
||||
request_body=anthropic_request,
|
||||
)
|
||||
if not (
|
||||
_supports_factory(
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
key="supports_output_config",
|
||||
)
|
||||
or AnthropicConfig._model_supports_effort_param(model, "bedrock")
|
||||
):
|
||||
if anthropic_request.pop("output_config", None) is not None:
|
||||
verbose_logger.warning(
|
||||
"Bedrock Invoke: stripping unsupported `output_config` for "
|
||||
"model=%s — neither `supports_output_config` nor any "
|
||||
"`supports_*_reasoning_effort` flag is set in "
|
||||
"model_prices_and_context_window.json. Add the capability "
|
||||
"flag to the model JSON entry if this model accepts "
|
||||
"`output_config`.",
|
||||
model,
|
||||
)
|
||||
apply_bedrock_invoke_structured_output(
|
||||
model=model,
|
||||
request_body=anthropic_request,
|
||||
)
|
||||
strip_unsupported_bedrock_invoke_output_config_keys(
|
||||
model=model,
|
||||
request_body=anthropic_request,
|
||||
)
|
||||
if "anthropic_version" not in anthropic_request:
|
||||
anthropic_request["anthropic_version"] = self.anthropic_version
|
||||
|
||||
|
|
|
|||
|
|
@ -177,6 +177,95 @@ def convert_bedrock_invoke_output_format_to_inline_schema(
|
|||
request_body["messages"] = new_messages
|
||||
|
||||
|
||||
def _bedrock_model_supports(model: str, key: str) -> bool:
|
||||
from litellm.utils import _supports_factory
|
||||
|
||||
return _supports_factory(model=model, custom_llm_provider="bedrock", key=key)
|
||||
|
||||
|
||||
def apply_bedrock_invoke_structured_output(
|
||||
model: str,
|
||||
request_body: dict[str, object], # mutable-ok: edited in place like siblings
|
||||
) -> None:
|
||||
"""
|
||||
Route Anthropic structured-output params to what the Bedrock model supports.
|
||||
|
||||
Consumes the legacy top-level ``output_format`` and the newer
|
||||
``output_config.format``, keeping the pre-existing precedence of the legacy
|
||||
field when a request carries both. Models flagged
|
||||
``supports_native_structured_output`` in the model map get the schema
|
||||
forwarded as ``output_config.format``, which Bedrock relays to the model for
|
||||
enforced structured output. For every other model the schema is inlined into
|
||||
the last user message as best-effort text, with a warning because nothing
|
||||
enforces it.
|
||||
"""
|
||||
legacy_output_format: Final = request_body.pop("output_format", None)
|
||||
output_config_format: Final = pop_bedrock_invoke_output_config_format(request_body)
|
||||
schema_format: Final = legacy_output_format if isinstance(legacy_output_format, dict) else output_config_format
|
||||
if schema_format is None:
|
||||
return
|
||||
|
||||
if _bedrock_model_supports(model, "supports_native_structured_output"):
|
||||
existing_output_config: Final = request_body.get("output_config")
|
||||
if isinstance(existing_output_config, dict):
|
||||
existing_output_config["format"] = schema_format
|
||||
else:
|
||||
request_body["output_config"] = {"format": schema_format} # rebind-ok: out-param # mutable-ok: json
|
||||
return
|
||||
|
||||
verbose_logger.warning(
|
||||
"Bedrock Invoke: model=%s does not advertise `supports_native_structured_output` "
|
||||
"in model_prices_and_context_window.json, so the JSON schema was inlined into "
|
||||
"the last user message and is NOT enforced by the model.",
|
||||
model,
|
||||
)
|
||||
convert_bedrock_invoke_output_format_to_inline_schema(
|
||||
output_format=schema_format,
|
||||
request_body=request_body,
|
||||
)
|
||||
|
||||
|
||||
def strip_unsupported_bedrock_invoke_output_config_keys(
|
||||
model: str,
|
||||
request_body: dict[str, object], # mutable-ok: edited in place like siblings
|
||||
) -> None:
|
||||
"""
|
||||
Drop ``output_config`` keys the Bedrock model does not accept.
|
||||
|
||||
``format`` survives unconditionally: it is only attached for models whose map
|
||||
entry advertises ``supports_native_structured_output``. Effort-bearing keys
|
||||
survive only when the map flags ``supports_output_config`` or a
|
||||
``supports_*_reasoning_effort`` tier; otherwise they are dropped with a
|
||||
warning so Bedrock does not reject the request.
|
||||
"""
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
|
||||
output_config: Final = request_body.get("output_config")
|
||||
if not isinstance(output_config, dict):
|
||||
return
|
||||
if all(key == "format" for key in output_config):
|
||||
return
|
||||
if _bedrock_model_supports(model, "supports_output_config") or AnthropicConfig._model_supports_effort_param(
|
||||
model, "bedrock"
|
||||
):
|
||||
return
|
||||
|
||||
verbose_logger.warning(
|
||||
"Bedrock Invoke: stripping unsupported `output_config` keys for "
|
||||
"model=%s: neither `supports_output_config` nor any "
|
||||
"`supports_*_reasoning_effort` flag is set in "
|
||||
"model_prices_and_context_window.json. Add the capability "
|
||||
"flag to the model JSON entry if this model accepts "
|
||||
"`output_config`.",
|
||||
model,
|
||||
)
|
||||
preserved_format: Final = output_config.get("format")
|
||||
if preserved_format is None:
|
||||
request_body.pop("output_config", None)
|
||||
else:
|
||||
request_body["output_config"] = {"format": preserved_format} # rebind-ok: out-param # mutable-ok: json
|
||||
|
||||
|
||||
def normalize_custom_field_on_tools(request_body: dict) -> None:
|
||||
"""
|
||||
Drop the ``custom`` field from each tool, first hoisting a boolean
|
||||
|
|
@ -1487,6 +1576,7 @@ class CommonBatchFilesUtils:
|
|||
aws_role_name=optional_params.get("aws_role_name"),
|
||||
aws_web_identity_token=optional_params.get("aws_web_identity_token"),
|
||||
aws_sts_endpoint=optional_params.get("aws_sts_endpoint"),
|
||||
aws_external_id=optional_params.get("aws_external_id"),
|
||||
)
|
||||
|
||||
# Prepare the request data
|
||||
|
|
|
|||
|
|
@ -113,6 +113,7 @@ class BedrockFilesHandler(BaseAWSLLM):
|
|||
aws_role_name=optional_params.get("aws_role_name"),
|
||||
aws_web_identity_token=optional_params.get("aws_web_identity_token"),
|
||||
aws_sts_endpoint=optional_params.get("aws_sts_endpoint"),
|
||||
aws_external_id=optional_params.get("aws_external_id"),
|
||||
)
|
||||
|
||||
# Create S3 client
|
||||
|
|
|
|||
|
|
@ -146,6 +146,7 @@ class _BedrockS3RequestParams(BaseModel):
|
|||
aws_role_name: str | None = None
|
||||
aws_web_identity_token: str | None = None
|
||||
aws_sts_endpoint: str | None = None
|
||||
aws_external_id: str | None = None
|
||||
s3_region_name: str | None = None
|
||||
s3_endpoint_url: str | None = None
|
||||
|
||||
|
|
@ -1029,6 +1030,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
aws_role_name=optional_params.get("aws_role_name"),
|
||||
aws_web_identity_token=optional_params.get("aws_web_identity_token"),
|
||||
aws_sts_endpoint=optional_params.get("aws_sts_endpoint"),
|
||||
aws_external_id=optional_params.get("aws_external_id"),
|
||||
)
|
||||
|
||||
# Calculate SHA256 hash of the content (REQUIRED for S3)
|
||||
|
|
@ -1290,6 +1292,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
aws_role_name=request_params.aws_role_name,
|
||||
aws_web_identity_token=request_params.aws_web_identity_token,
|
||||
aws_sts_endpoint=request_params.aws_sts_endpoint,
|
||||
aws_external_id=request_params.aws_external_id,
|
||||
)
|
||||
|
||||
empty_body_hash: Final = hashlib.sha256(b"").hexdigest()
|
||||
|
|
|
|||
|
|
@ -29,14 +29,14 @@ from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation
|
|||
AmazonInvokeConfig,
|
||||
)
|
||||
from litellm.llms.bedrock.common_utils import (
|
||||
convert_bedrock_invoke_output_format_to_inline_schema,
|
||||
apply_bedrock_invoke_structured_output,
|
||||
ensure_bedrock_anthropic_messages_tool_names,
|
||||
get_anthropic_beta_from_headers,
|
||||
is_claude_4_5_on_bedrock,
|
||||
normalize_bedrock_opus_output_config_effort,
|
||||
normalize_custom_field_on_tools,
|
||||
normalize_tool_input_schema_types_for_bedrock_invoke,
|
||||
pop_bedrock_invoke_output_config_format,
|
||||
strip_unsupported_bedrock_invoke_output_config_keys,
|
||||
)
|
||||
from litellm.llms.bedrock.request_metadata import (
|
||||
bedrock_request_metadata_headers,
|
||||
|
|
@ -51,7 +51,6 @@ from litellm.types.llms.openai import AllMessageValues
|
|||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
|
||||
from litellm.types.utils import GenericStreamingChunk as GChunk
|
||||
from litellm.utils import _supports_factory
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
|
@ -708,52 +707,25 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
# 4. Remove `ttl` field from cache_control in messages (Bedrock doesn't support it for older models)
|
||||
self._remove_ttl_from_cache_control(anthropic_messages_request=anthropic_messages_request, model=model)
|
||||
|
||||
# 5. Convert structured-output params to inline schema.
|
||||
# Bedrock Invoke doesn't support top-level `output_format`; its
|
||||
# accepted `output_config` subset is also narrower than Anthropic's, so
|
||||
# consume the newer `output_config.format` shape here instead of
|
||||
# forwarding it as an unknown nested key.
|
||||
# 5. Route structured-output params (`output_format` /
|
||||
# `output_config.format`) to native enforcement or the inline-schema
|
||||
# fallback, then strip `output_config` keys the model does not accept.
|
||||
# Ref: https://github.com/BerriAI/litellm/issues/22797
|
||||
existing_output_config: Final = anthropic_messages_request.get("output_config")
|
||||
if isinstance(existing_output_config, dict):
|
||||
anthropic_messages_request["output_config"] = dict(existing_output_config)
|
||||
output_format: Final = anthropic_messages_request.pop("output_format", None)
|
||||
output_config_format: Final = pop_bedrock_invoke_output_config_format(anthropic_messages_request)
|
||||
if output_format:
|
||||
convert_bedrock_invoke_output_format_to_inline_schema(
|
||||
output_format=output_format,
|
||||
request_body=anthropic_messages_request,
|
||||
)
|
||||
elif output_config_format:
|
||||
convert_bedrock_invoke_output_format_to_inline_schema(
|
||||
output_format=output_config_format,
|
||||
request_body=anthropic_messages_request,
|
||||
)
|
||||
apply_bedrock_invoke_structured_output(
|
||||
model=model,
|
||||
request_body=anthropic_messages_request,
|
||||
)
|
||||
normalize_bedrock_opus_output_config_effort(
|
||||
model=model,
|
||||
output_config=anthropic_messages_request.get("output_config"),
|
||||
)
|
||||
|
||||
# 5a. Bedrock Invoke supports output_config (effort) for Claude 4.6+ models,
|
||||
# but older models do not — strip it to avoid request rejection.
|
||||
# Ref: https://github.com/BerriAI/litellm/issues/22797
|
||||
if not (
|
||||
_supports_factory(
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
key="supports_output_config",
|
||||
)
|
||||
or AnthropicConfig._model_supports_effort_param(model, "bedrock")
|
||||
):
|
||||
if anthropic_messages_request.pop("output_config", None) is not None:
|
||||
verbose_logger.warning(
|
||||
"Bedrock Invoke: stripping unsupported `output_config` for "
|
||||
"model=%s — neither `supports_output_config` nor any "
|
||||
"`supports_*_reasoning_effort` flag is set in "
|
||||
"model_prices_and_context_window.json. Add the capability "
|
||||
"flag to the model JSON entry if this model accepts "
|
||||
"`output_config`.",
|
||||
model,
|
||||
)
|
||||
strip_unsupported_bedrock_invoke_output_config_keys(
|
||||
model=model,
|
||||
request_body=anthropic_messages_request,
|
||||
)
|
||||
|
||||
# 5b. Hoist `custom.defer_loading` then drop `custom` (Bedrock doesn't support it)
|
||||
# Ref: https://github.com/BerriAI/litellm/issues/22847
|
||||
|
|
@ -774,9 +746,11 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
if filtered_betas:
|
||||
anthropic_messages_request["anthropic_beta"] = filtered_betas
|
||||
|
||||
remaining_output_config: Final = anthropic_messages_request.get("output_config")
|
||||
if (
|
||||
litellm.drop_params is True
|
||||
and "output_config" in anthropic_messages_request
|
||||
and isinstance(remaining_output_config, dict)
|
||||
and any(key != "format" for key in remaining_output_config)
|
||||
and not AnthropicConfig._model_supports_effort_param(model, "bedrock")
|
||||
):
|
||||
verbose_logger.warning(
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue