Merge branch 'BerriAI:litellm_internal_staging' into fix/custom-pricing-anthropic-messages-azure-ai

This commit is contained in:
khoapmdx-oss 2026-08-12 08:49:41 +07:00 • committed by GitHub
commit a510a4f812
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
444 changed files with 37584 additions and 6182 deletions

View file

@ -23,30 +23,56 @@ body:
label: What happened?
description: Also tell us, what did you expect to happen?
placeholder: Tell us what you see!
value: "A bug happened!"
validations:
required: true
- type: textarea
id: steps-to-reproduce
id: user-flow
attributes:
label: Steps to Reproduce
description: Please provide a numbered list of the exact steps to reproduce this bug (include a curl/python snippet to reproduce it). Number each step (1., 2., 3., ...) in the order you performed them.
label: User Flow
description: |
Two ordered lists, "Before a (hypothetical) fix" and "After a (hypothetical) fix", walking the same end user through the same task, written strictly from that user's seat. Every rule below applies.
- Describe the real application and the routes its users actually hit, not a generic scenario
- Lead each list with one plain sentence saying where the flow fails (before) or would succeed (after), then number the steps
- Every step is something the user does or observes: the HTTP method and full URL they hit, what they sent, and what visibly came back (status code, error text, the shape of an ID). UI steps name the page URL and what is on screen
- No LiteLLM internals: never name functions, files, DB tables, config classes, hooks, callbacks, or code paths. "The upload hands back an ID that looks like OpenAI's own `file-abc123` instead of the scrambled one the gateway returned" is right, "no managed-file row was registered" is wrong
- Keep the two lists step-for-step identical until they diverge, so the broken step is obvious
- If the bug has a security or authorization consequence, end each list with what another user can do that they shouldn't be able to, and what they could no longer do after a fix
placeholder: |
1. config.yaml file/ .env file/ etc.
2. Run the following code...
3. Observe the error...
value: |
1.
2.
3.
Before a (hypothetical) fix: a developer whose app streams chat completions gets no token counts back, so their cost dashboard reads zero
1. They send POST https://litellm-domain/v1/chat/completions with "stream": true and no stream_options
2. The last SSE chunk arrives with "usage": null, so their app records 0 prompt and 0 completion tokens
3. They open https://litellm-domain/ui/?page=logs and see the request logged at $0 spend
After a (hypothetical) fix: the same request comes back with real token counts, so the dashboard shows real spend
1. The proxy admin sets always_include_stream_usage: true and restarts the proxy
2. The developer sends the same POST https://litellm-domain/v1/chat/completions with "stream": true and no stream_options
3. The last SSE chunk now carries a usage object with real prompt and completion token counts
4. https://litellm-domain/ui/?page=logs shows that request at non-zero spend
validations:
required: true
- type: textarea
id: logs
id: proof-of-bug
attributes:
label: Relevant log output
description: Please copy and paste any relevant log output. This will be automatically formatted into code, so no need for backticks.
render: shell
label: Proof the bug occurs
description: |
The commands (e.g., curl) and their full output, screenshots, or a screen recording demonstrating that the bug happens. Every rule below applies.
- The proof must be completely e2e with no mocks, against a live proxy you ran yourself (e.g., `litellm --config config.yaml --detailed_debug` on localhost:4000), hitting real LLM provider APIs, costing real $ if needed, where the bug involves a provider call. `pytest` commands are not enough
- Show exactly what the end user sees or does, matching the User Flow above step for step
- Start with the config.yaml (or SDK setup) and any env vars the proxy ran with, then the exact version or commit hash the proof was captured at, so a maintainer can stand up the same proxy before running your commands. Keep the real values for env vars that aren't sensitive, they are often the reason the bug happens, and redact only the secrets: never paste a real API key, virtual key, database URL, or other credential, here or anywhere else in the issue
- If the bug applies to more than one of the LLM endpoints (/v1/responses, /v1/chat/completions, /v1/messages), include proof for every one of them, not just one
- For UI bugs: include screenshots and the page URLs you were on. Scrub keys and tokens out of screenshots too (for example, the virtual key is briefly shown in the panel right after you create a virtual key)
placeholder: |
Config / setup the proxy ran with:
Version or commit:
Commands and their full output:
validations:
required: true
- type: dropdown
id: component
attributes:

View file

@ -24,10 +24,53 @@ body:
validations:
required: true
- type: textarea
id: motivation
id: user-flow
attributes:
label: Motivation, pitch
description: Please outline the motivation for the proposal. Is your feature request related to a specific problem? e.g., "I'm working on X and would like Y to be possible". If this is related to another GitHub issue, please link here too.
label: User Flow
description: |
Two ordered lists, "Before this feature (today)" and "After this feature (ideal user flow)", walking the same end user through the same task, written strictly from that user's seat. Every rule below applies.
- Describe the real application and the routes its users actually hit, not a generic scenario. Link any related GitHub issue or provider API docs
- Lead each list with one plain sentence saying where the flow dead-ends today and what it would let them do instead, then number the steps
- Every step is something the user does or observes: the HTTP method and full URL they hit, what they sent, and what visibly came back (status code, error text, the shape of an ID). UI steps name the page URL and what is on screen
- No LiteLLM internals: never name functions, files, DB tables, config classes, hooks, callbacks, or code paths. Ask for the behavior you need, not the implementation you imagine
- Keep the two lists step-for-step identical until they diverge, so the missing capability is obvious
- "Before this feature" is also where you show the workaround you're living with, which is what tells us how badly this is needed
placeholder: |
Before this feature (today): a developer batching nightly summaries has no way to mark those calls as low priority, so they compete with live traffic for the same rate limit
1. They send POST https://litellm-domain/v1/chat/completions for 500 documents in a loop
2. Around document 120 they start getting 429s naming the rpm limit, and their user-facing chat app starts getting them too
3. Their workaround is a hand-rolled sleep between calls, which stretches the batch to 3 hours and still collides at peak
After this feature (ideal user flow): the same batch runs as background work that yields to live traffic
1. The developer sends the same POST with "service_tier": "flex"
2. Batch calls queue behind interactive ones instead of 429ing, and the response comes back with the tier it was served at
3. The live chat app keeps returning 200s throughout the batch
4. https://litellm-domain/ui/?page=logs shows the batch requests tagged with that tier
validations:
required: true
- type: textarea
id: how-far-you-got
attributes:
label: How far you got
description: |
Run as many steps of the "After this feature (ideal user flow)" list as you can against a live proxy you ran yourself (e.g., `litellm --config config.yaml --detailed_debug` on localhost:4000), then paste the commands (e.g., curl) and their full output, ending at the step that dead-ends. Every rule below applies.
- Say plainly what stopped you there, in user terms: the option you passed came back ignored, the response 400'd naming an unsupported field, there is no button on the page for it. This is what proves the feature is genuinely missing rather than undocumented
- No mocks. Where the flow involves a provider call, hit the real provider API, even if it costs real $. `pytest` commands are not enough
- Include the config.yaml (or SDK setup) and env vars the proxy ran with, plus the version or commit you were on. Keep the real values for env vars that aren't sensitive, and redact only the secrets: never paste a real API key, virtual key, database URL, or other credential, here or anywhere else in the issue
- If the provider already supports this, link their API docs and paste a direct call to them succeeding, so we can see the shape LiteLLM should be sending
- For UI asks: include screenshots of the page you got stuck on and its URL. Scrub keys and tokens out of screenshots too (for example, the virtual key is briefly shown in the panel right after you create a virtual key)
placeholder: |
Config / setup the proxy ran with:
Version or commit:
Commands and their full output, up to the step that dead-ends:
What stopped me there:
validations:
required: true
- type: dropdown

View file

@ -0,0 +1,40 @@
name: "Cache Prisma binaries"
description: >-
Cache the Prisma CLI and engine binaries that `prisma generate` downloads, so
only the first job on a given prisma-client-py version pays for the download.
prisma-client-py shells out to `npm install prisma@<version>` whenever its
binary cache directory has no CLI entrypoint, which pulls ~85 MB of query and
schema engines over the network. That normally takes a few seconds, but it is
unbounded: one shard of a proxy-db run took 5m18s on that single step versus
3.8s on its eleven siblings, which pushed the job past its timeout and got a
fully passing test run cancelled.
Callers must not set PRISMA_BINARY_CACHE_DIR. The prisma-client-py default
(~/.cache/prisma-python/binaries/<prisma-version>/<engine-version>) is already
keyed by both versions, so a cache entry can never be served to a run that
expects different binaries.
runs:
using: composite
steps:
- name: Resolve prisma-client-py version
id: version
shell: bash
run: |
version="$(grep -A1 '^name = "prisma"$' uv.lock | sed -n 's/^version = "\(.*\)"$/\1/p' | head -1)"
if [ -z "${version}" ]; then
echo "could not resolve the prisma package version from uv.lock" >&2
exit 1
fi
echo "version=${version}" >> "$GITHUB_OUTPUT"
- name: Restore Prisma binaries
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
# ~/.cache/prisma-python holds the npm install tree prisma-client-py
# drives; ~/.cache/prisma is where @prisma/engines stages its downloads.
path: |
~/.cache/prisma-python
~/.cache/prisma
key: ${{ runner.os }}-prisma-binaries-${{ steps.version.outputs.version }}

View file

@ -83,7 +83,11 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac
🚄 Infrastructure
✅ Test
## Changes
## Caveats (if any)
<!-- Short bullet points, just like the TLDR: one line per bullet, roughly 10 words max
Call out known limitations, follow-up work, or anything a reviewer should watch out for
Leave this section empty if there are none -->
## QA runbook

View file

@ -582,7 +582,9 @@ def build_issue_prompt(*, title: str, body: str) -> str:
Commands whose external dependencies (LLM provider, DB,
network) are mocked or stubbed do NOT count.
Prose-only "steps to reproduce" with no run output, video, or
screenshot do NOT satisfy (1).
screenshot do NOT satisfy (1). An unfilled template scaffold
(bare headings such as "Version or commit:" with nothing under
them, empty numbered lists) counts as absent, not as evidence.
(2) Expected vs. actual behavior (`has_expected_vs_actual`).
FAIL the bug report if either (1) or (2) is missing. Do not bias
@ -595,6 +597,13 @@ def build_issue_prompt(*, title: str, body: str) -> str:
that it does not today).
- Motivation / use case with a concrete example (config, API call,
UI flow, or scenario showing what's blocked today).
- END-TO-END EVIDENCE OF THE DEAD-END (set
`has_dead_end_evidence=true` only when this is present): a video,
a screenshot, or the exact command(s) actually run paired with
their real output, showing the point where the flow stops today.
Mocked or stubbed dependencies do NOT count, and an unfilled
template scaffold (bare headings, empty numbered lists) counts as
absent.
For an issue that is neither a bug report nor a feature request (a
question, support request, or discussion), PASS as long as it has a
@ -608,6 +617,7 @@ def build_issue_prompt(*, title: str, body: str) -> str:
"has_repro": boolean,
"has_expected_vs_actual": boolean,
"has_motivation_example": boolean,
"has_dead_end_evidence": boolean,
"missing": ["plain-english strings naming what is missing"],
"explanation": "1-2 sentence reasoning for the team to skim"
}}
@ -705,6 +715,10 @@ _ISSUE_BUG_LABELS: tuple[tuple[str, str], ...] = (
)
_ISSUE_FEATURE_LABELS: tuple[tuple[str, str], ...] = (
("has_motivation_example", "Motivation and concrete example"),
(
"has_dead_end_evidence",
"End-to-end evidence of the dead-end (video, screenshot, or command + real output)",
),
)
@ -836,8 +850,11 @@ def format_issue_close_comment(verdict: dict) -> str:
"video, a screenshot, or the exact commands you ran with their real output / "
"traceback) plus expected vs. actual behavior. Written steps with no run output, "
"video, or screenshot don't count, and mocked or stubbed runs don't count.\n"
" - For **feature requests**: a concrete description of what should change, plus a "
"use case and example (config / API call / UI flow).\n"
" - For **feature requests**: a concrete description of what should change, a "
"use case and example (config / API call / UI flow), plus end-to-end evidence of "
"the dead-end (a video, a screenshot, or the exact commands you ran with their "
"real output showing where the flow stops today). Mocked or stubbed runs don't "
"count.\n"
"2. Comment `@agent-shin reconsider`. I'll re-run triage and reopen the issue if it "
"now meets the bar. (GitHub doesn't let external authors reopen an issue a maintainer "
"or bot closed, so the comment-based reconsider is the reliable path.)\n"
@ -943,8 +960,10 @@ def format_grace_warning_issue_comment(verdict: dict) -> str:
"screenshot, or the exact commands you ran with their real output / traceback) plus "
"expected vs. actual behavior. Written steps with no run output don't count, and "
"mocked or stubbed runs don't count.\n"
"- For **feature requests**: a concrete description of what should change, plus a use "
"case and example (config / API call / UI flow).\n"
"- For **feature requests**: a concrete description of what should change, a use "
"case and example (config / API call / UI flow), plus end-to-end evidence of the "
"dead-end (a video, a screenshot, or the exact commands you ran with their real "
"output showing where the flow stops today). Mocked or stubbed runs don't count.\n"
"\n"
"**If the issue does get auto-closed in 2 hours**, comment `@agent-shin reconsider` "
"and I'll re-evaluate. If it now meets the bar, I'll reopen the issue.\n"

View file

@ -18,10 +18,25 @@ on:
type: number
default: 2
timeout-minutes:
description: "Job timeout in minutes"
description: >-
Timeout for the test step alone. Setup (checkout, dependency install,
Prisma client generation) gets its own allowance on top, so a slow
runner or a cold binary download can never cancel passing tests.
required: false
type: number
default: 20
job-timeout-minutes:
description: >-
Backstop for the whole job. Keep it >= `timeout-minutes` plus 35: 30 for
the per-step ceilings on the setup steps below, and 5 for the runner
overhead the job clock charges but no step owns (job init, step
transitions, post-job cleanup). That headroom is what makes the test
budget a floor rather than a hope, since setup cannot overrun into it
without failing its own step first. GitHub expressions have no
arithmetic, so the sum is passed in rather than computed.
required: false
type: number
default: 55
max-failures:
description: "Stop after this many failures"
required: false
@ -44,30 +59,35 @@ jobs:
run:
name: Run tests
runs-on: ubuntu-latest
timeout-minutes: ${{ inputs.timeout-minutes }}
timeout-minutes: ${{ inputs.job-timeout-minutes }}
outputs:
decision: ${{ steps.changes.outputs.decision }}
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
timeout-minutes: 3
with:
persist-credentials: false
- name: Detect backend-relevant changes
id: changes
timeout-minutes: 2
uses: ./.github/actions/detect-backend-changes
- name: Set up Python
timeout-minutes: 3
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
timeout-minutes: 3
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache uv dependencies
timeout-minutes: 5
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
@ -79,18 +99,24 @@ jobs:
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 8
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml
- name: Cache Prisma binaries
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 3
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
if: steps.changes.outputs.decision != 'skip'
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
timeout-minutes: 3
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Run tests
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: ${{ inputs.timeout-minutes }}
env:
TEST_PATH: ${{ inputs.test-path }}
MAX_FAILURES: ${{ inputs.max-failures }}

View file

@ -71,10 +71,12 @@ jobs:
if: steps.changes.outputs.relevant == 'true'
run: .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Cache Prisma binaries
if: steps.changes.outputs.relevant == 'true'
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
if: steps.changes.outputs.relevant == 'true'
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Set up Node.js

View file

@ -57,9 +57,10 @@ jobs:
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml
- name: Cache Prisma binaries
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma

View file

@ -43,12 +43,13 @@ jobs:
with:
version: "0.10.9"
- name: Cache Prisma binaries
uses: ./.github/actions/cache-prisma-binaries
# The gate provisions its own measurement env (.venv-typecheck: a frozen
# uv sync of its canonical dependency groups plus a generated Prisma
# client), so no install step here can drift from what local runs measure.
- name: Emit basedpyright counts for HEAD
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
python scripts/type_check_gate.py --emit-counts-dir "$RUNNER_TEMP/basedpyright-counts"
counts_file=$(ls "$RUNNER_TEMP"/basedpyright-counts/basedpyright-counts-*.json)

View file

@ -65,6 +65,12 @@ jobs:
- name: check_provider_folders_documented
run: uv run --no-sync python ./tests/code_coverage_tests/check_provider_folders_documented.py
- name: check_prisma_binary_cache
run: uv run --no-sync python ./tests/code_coverage_tests/check_prisma_binary_cache.py
- name: check_workflow_startup_safety
run: uv run --no-sync python ./tests/code_coverage_tests/check_workflow_startup_safety.py
- name: router_code_coverage
run: uv run --no-sync python ./tests/code_coverage_tests/router_code_coverage.py

View file

@ -43,9 +43,10 @@ jobs:
BASE_SHA: ${{ github.event.pull_request.base.sha }}
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
run: |
MERGE_BASE=$(gh api "repos/${{ github.repository }}/compare/${BASE_SHA}...${HEAD_SHA}?per_page=1" --jq '.merge_base_commit.sha')
retry() { "$@" || { sleep 15; "$@"; } || { sleep 30; "$@"; }; }
MERGE_BASE=$(retry gh api "repos/${{ github.repository }}/compare/${BASE_SHA}...${HEAD_SHA}?per_page=1" --jq '.merge_base_commit.sha')
test -n "$MERGE_BASE"
git fetch --no-tags --depth=1 origin "$MERGE_BASE"
retry git fetch --no-tags --depth=1 origin "$MERGE_BASE"
echo "GATE_BASE_SHA=$MERGE_BASE" >> "$GITHUB_ENV"
- name: Set up Python
@ -71,12 +72,13 @@ jobs:
run: |
uv sync --frozen --group proxy-dev --group e2e-dev
- name: Cache Prisma binaries
uses: ./.github/actions/cache-prisma-binaries
# basedpyright resolves Prisma's generated client (litellm/proxy/schema.prisma)
# only after `prisma generate` writes prisma/client.py et al. Without this the
# DB wrappers typed against the generated client would degrade to Unknown.
- name: Generate Prisma client
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
@ -119,7 +121,6 @@ jobs:
- name: Check basedpyright budget (delta vs base)
env:
GH_TOKEN: ${{ github.token }}
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync python scripts/type_check_gate.py --base "$GATE_BASE_SHA"
@ -161,7 +162,8 @@ jobs:
env:
BASE_SHA: ${{ github.event.pull_request.base.sha }}
run: |
git fetch --no-tags --depth=1 origin "$BASE_SHA"
retry() { "$@" || { sleep 15; "$@"; } || { sleep 30; "$@"; }; }
retry git fetch --no-tags --depth=1 origin "$BASE_SHA"
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
@ -205,7 +207,8 @@ jobs:
GITGUARDIAN_API_KEY: ${{ secrets.GITGUARDIAN_API_KEY }}
run: |
if [ -n "$GITGUARDIAN_API_KEY" ]; then
git fetch --no-tags --unshallow origin
retry() { "$@" || { sleep 15; "$@"; } || { sleep 30; "$@"; }; }
retry git fetch --no-tags --unshallow origin
uv tool run --from 'ggshield==1.48.0' ggshield secret scan repo .
else
echo "GITGUARDIAN_API_KEY not set, skipping ggshield scan"

View file

@ -42,6 +42,11 @@ jobs:
- name: Install dependencies
run: npm ci
- name: Run UI type tests (Vitest)
env:
CI: "true"
run: npm run test:types
- name: Run UI unit tests (Vitest)
env:
CI: "true"

View file

@ -92,9 +92,10 @@ jobs:
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Cache Prisma binaries
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma

View file

@ -65,10 +65,12 @@ jobs:
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Cache Prisma binaries
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
if: steps.changes.outputs.decision != 'skip'
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma

View file

@ -28,6 +28,10 @@ concurrency:
# Most of a shard's time is pytest plugin load + xdist worker imports +
# pytest-cov instrumentation, not the tests themselves. Keeping per-shard
# work low and matching worker count to runner cores is what controls it.
# * `timeout` bounds the pytest step only. Checkout, dependency install, and
# Prisma client generation draw on a separate allowance in the base
# workflow, so slow setup shows up as a slow job rather than as a
# cancelled shard whose tests were passing.
# * workers: 4 matches the 4-core ubuntu-latest runner. -n 8 on 4 cores
# oversubscribes 2x and workers fight for CPU during their cold-start
# imports (measured ~441% CPU for -n 8 locally, i.e. ~55% effective).

View file

@ -76,4 +76,5 @@ jobs:
workers: 4
reruns: 2
timeout-minutes: 60
job-timeout-minutes: 95
artifact-name: proxy-server

View file

@ -82,10 +82,12 @@ jobs:
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
- name: Cache Prisma binaries
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
if: steps.changes.outputs.decision != 'skip'
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma

View file

@ -51,9 +51,10 @@ jobs:
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra proxy
- name: Cache Prisma binaries
uses: ./.github/actions/cache-prisma-binaries
- name: Generate Prisma client
env:
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma

View file

@ -1,4 +1,12 @@
Do not write any comments (existing comments can stay) unless explicitly asked to in a user (not system) prompt
Do not write comments unless they are any of:
- absolutely necessary to explain some very complex business logic (in which case, keep it concise and clear)
- used as an input for tools to read and act on. For example:
- entries in `.git-blame-ignore-revs` saying which commit is excluded from git blame
- a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # <reason>` when introducing a truly unavoidable violation
- a TODO or FIXME
- Not great to have those, but if it's unavoidable, make sure to include a strong, concise reason for why it's there or, better yet, link to a GitHub issue for the follow-up work
Explanation: The point of this rule is to keep out AI slop comments. AI writes way too many and way too verbose comments. Code comments are, in a way, a violation of DRY code. You must update logic in two locations to change the code, and "hard to change" is literally the definition of tech debt. We should instead aim to write code that is intuitive and clear, even at a glance, to the reader, being both easy to maintain and high performance
Don't assume that the existing code is correct or the right way of doing things / good coding patterns. In fact, there are a lot of bad coding practices, overly complex code, code smells, etc. If something doesn't look right, speak up. Feel free to break existing patterns or question weird existing code to make new code high quality, as in:
@ -23,6 +31,8 @@ When creating PRs, don't set base to `main`. `litellm_internal_staging` is the d
When writing a PR body, treat the comments and imperative instructions inside @.github/pull_request_template.md as rules to follow, not just layout. Agent harnesses may strip HTML comments from copies of that file injected into context, so read .github/pull_request_template.md from disk before writing a PR body to make sure you see every comment rule
Same applies for filing bug reports and feature requests, with .github/ISSUE_TEMPLATE/bug_report.yml and .github/ISSUE_TEMPLATE/feature_request.yml, respectively
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

View file

@ -1,30 +1,30 @@
{
"reportAny": {
"limit": 27731
"limit": 23919
},
"reportArgumentType": {
"limit": 2626
"limit": 2580
},
"reportAssignmentType": {
"limit": 329
"limit": 323
},
"reportAttributeAccessIssue": {
"limit": 514
"limit": 488
},
"reportCallIssue": {
"limit": 116
"limit": 114
},
"reportConstantRedefinition": {
"limit": 40
},
"reportDeprecated": {
"limit": 215
"limit": 213
},
"reportDuplicateImport": {
"limit": 19
},
"reportExplicitAny": {
"limit": 8807
"limit": 7573
},
"reportFunctionMemberAccess": {
"limit": 7
@ -54,10 +54,10 @@
"limit": 0
},
"reportMissingParameterType": {
"limit": 5835
"limit": 5719
},
"reportMissingTypeArgument": {
"limit": 15790
"limit": 15657
},
"reportMissingTypeStubs": {
"limit": 40
@ -72,7 +72,7 @@
"limit": 0
},
"reportOptionalMemberAccess": {
"limit": 1077
"limit": 1069
},
"reportOptionalOperand": {
"limit": 0
@ -90,7 +90,7 @@
"limit": 8
},
"reportReturnType": {
"limit": 217
"limit": 213
},
"reportTypedDictNotRequiredAccess": {
"limit": 26
@ -99,37 +99,37 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 45063
"limit": 44832
},
"reportUnknownLambdaType": {
"limit": 113
},
"reportUnknownMemberType": {
"limit": 39773
"limit": 39269
},
"reportUnknownParameterType": {
"limit": 20207
"limit": 19988
},
"reportUnknownVariableType": {
"limit": 31281
"limit": 30923
},
"reportUnnecessaryCast": {
"limit": 122
"limit": 118
},
"reportUnnecessaryComparison": {
"limit": 701
"limit": 699
},
"reportUnnecessaryContains": {
"limit": 5
},
"reportUnnecessaryIsInstance": {
"limit": 862
"limit": 853
},
"reportUntypedBaseClass": {
"limit": 0
},
"reportUntypedFunctionDecorator": {
"limit": 33
"limit": 27
},
"reportUnusedClass": {
"limit": 23
@ -138,7 +138,7 @@
"limit": 139
},
"reportUnusedImport": {
"limit": 555
"limit": 545
},
"reportUnusedVariable": {
"limit": 146

View file

@ -99,6 +99,7 @@ class BaseEmailLogger(CustomLogger):
email_html_content = USER_INVITATION_EMAIL_TEMPLATE.format(
email_logo_url=email_params.logo_url,
recipient_email=email_params.recipient_email,
invitation_link=email_params.base_url,
base_url=email_params.base_url,
email_support_contact=email_params.support_contact,
email_footer=email_params.signature,
@ -826,10 +827,15 @@ class BaseEmailLogger(CustomLogger):
"""
# Early validation
if not user_id:
verbose_proxy_logger.debug("No user_id provided for invitation link")
verbose_proxy_logger.warning(
"No user_id provided for invitation link. Email will link to base URL instead of onboarding page"
)
return base_url
if not await self._is_prisma_client_available():
verbose_proxy_logger.warning(
"Prisma client not available. Email will link to base URL instead of onboarding page"
)
return base_url
# Wait for any concurrent invitation creation to complete
@ -839,11 +845,15 @@ class BaseEmailLogger(CustomLogger):
invitation = await self._get_or_create_invitation(user_id)
if not invitation:
verbose_proxy_logger.warning(
f"Failed to get/create invitation for user_id: {user_id}"
f"Failed to get/create invitation for user_id: {user_id}. Email will link to base URL instead of onboarding page"
)
return base_url
return self._construct_invitation_link(invitation.id, base_url)
invitation_link = self._construct_invitation_link(invitation.id, base_url)
verbose_proxy_logger.info(
f"Successfully created invitation link for user_id: {user_id}"
)
return invitation_link
async def _is_prisma_client_available(self) -> bool:
"""Check if Prisma client is available"""
@ -921,7 +931,9 @@ class BaseEmailLogger(CustomLogger):
# http://localhost:4000/ui/onboarding?invitation_id=7a096b3a-37c6-440f-9dd1-ba22e8043f6b
"""
return f"{base_url}/ui/onboarding?invitation_id={invitation_id}"
base_url = base_url.rstrip("/")
invitation_link = f"{base_url}/ui/onboarding?invitation_id={invitation_id}"
return invitation_link
async def send_email(
self,

View file

@ -3,7 +3,7 @@ Polls LiteLLM_ManagedObjectTable to check if the batch job is complete, and if t
"""
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, List, Optional, Tuple
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
@ -43,11 +43,15 @@ class CheckBatchCost:
# the guaranteed-failing primary query on every subsequent cycle.
self._has_batch_processed_column: bool = True
async def _get_user_info(self, batch_id, user_id) -> dict:
async def _get_user_info(self, batch_id: str, user_id: Optional[str]) -> Dict[str, Any]:
"""
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).
Returns an empty dict when user_id is None: batches created by a team or service
account key carry no user id, and find_unique(where={"user_id": None}) raises.
"""
if not user_id:
return {}
try:
user_row = await self.prisma_client.db.litellm_usertable.find_unique(
where={"user_id": user_id}
@ -62,6 +66,66 @@ class CheckBatchCost:
verbose_proxy_logger.error(f"CheckBatchCost: could not look up user {user_id} for batch {batch_id}: {e}")
return {}
async def _get_key_alias(self, batch_id: str, api_key: str | None) -> str | None:
"""Resolve the creating virtual key's alias from its hashed token."""
if not api_key:
return None
try:
key_row = 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:
verbose_proxy_logger.error(f"CheckBatchCost: could not look up key alias for batch {batch_id}: {e}")
return None
async def _get_team_alias(self, team_id: str | None) -> str | None:
"""Resolve a team's alias from its id."""
if not team_id:
return None
try:
team_row = 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:
verbose_proxy_logger.error(f"CheckBatchCost: could not look up team alias for team {team_id}: {e}")
return None
async def _build_creator_attribution_metadata(
self, job: "LiteLLM_ManagedObjectTable", batch_id: str
) -> Dict[str, Any]:
"""
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
is. Rows created before api_key and request_tags were persisted carry only
created_by and team_id, and fall back to those. A named creating key owns
user_api_key_alias; when it has no alias, or the key has since been rotated or
deleted, the field keeps the creating user's alias that _get_user_info filled in,
because a resolvable name is more useful on the spend row than a null.
"""
api_key = getattr(job, "api_key", None)
team_id = getattr(job, "team_id", None)
request_tags = getattr(job, "request_tags", None)
metadata: Dict[str, Any] = {
"user_api_key_user_id": job.created_by,
"user_api_key": api_key,
"user_api_key_team_id": team_id,
**(await self._get_user_info(batch_id, job.created_by)),
}
key_alias = await self._get_key_alias(batch_id, api_key)
if key_alias is not None:
metadata["user_api_key_alias"] = key_alias
team_alias = await self._get_team_alias(team_id)
if team_alias is not None:
metadata["user_api_key_team_alias"] = team_alias
if isinstance(request_tags, list) and request_tags:
metadata["tags"] = [tag for tag in request_tags if isinstance(tag, str)]
return metadata
async def _cleanup_stale_managed_objects(self) -> None:
"""
Mark managed objects older than MANAGED_OBJECT_STALENESS_CUTOFF_DAYS days
@ -485,9 +549,6 @@ class CheckBatchCost:
function_id=str(uuid.uuid4()),
)
creator_user_id = job.created_by
user_info = await self._get_user_info(batch_id, job.created_by)
logging_obj.update_environment_variables(
litellm_params={
# set the user-agent header so that S3 callback consumers can easily identify CheckBatchCost callbacks
@ -496,11 +557,7 @@ class CheckBatchCost:
"user-agent": CHECK_BATCH_COST_USER_AGENT,
}
},
"metadata": {
"user_api_key_user_id": creator_user_id,
"user_api_key_team_id": getattr(job, "team_id", None),
**user_info,
},
"metadata": await self._build_creator_attribution_metadata(job, batch_id),
},
optional_params={},
)

View file

@ -163,6 +163,8 @@ class _ManagedObjectTableActions(Protocol):
self, where: Mapping[str, str], data: Mapping[str, Mapping[str, object]]
) -> "PrismaManagedObjectRow": ...
async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ...
class _CursorPageArgs(TypedDict, total=False):
cursor: Mapping[str, str]
@ -263,7 +265,24 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
model_object_id: str,
file_purpose: Literal["batch", "fine-tune", "response"],
user_api_key_dict: UserAPIKeyAuth,
request_tags: Sequence[str] | None = None,
persist_attribution: bool = False,
create_if_missing: bool = True,
) -> None:
"""Persist a managed object row, caching it and upserting it in the DB.
persist_attribution is set only by the batch create, which is the one caller
that can speak for the creator; it gates the api_key and request_tags columns
that CheckBatchCost bills against, so a later poll or retrieve of the same
batch cannot record itself as the paying key. Like created_by and team_id,
both are written only in the upsert create branch, never on update.
create_if_missing is cleared by callers that observe a batch they did not
create, such as a poll. They still refresh status and file_object, but a
row absent from the table is left absent rather than created with the
observer as its creator, because created_by and team_id are written from
whoever calls the create branch.
"""
verbose_logger.info(f"Storing LiteLLM Managed {file_purpose} object with id={unified_object_id} in cache")
litellm_managed_object = LiteLLM_ManagedObjectTable(
unified_object_id=unified_object_id,
@ -277,6 +296,29 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
litellm_parent_otel_span=litellm_parent_otel_span,
)
from prisma import Json
api_key = user_api_key_dict.api_key or None
attribution_columns = (
{
**({"api_key": api_key} if api_key is not None else {}),
**({"request_tags": Json(list(request_tags))} if request_tags else {}),
}
if persist_attribution
else {}
)
# FIX: Update status and file_object on every operation to keep state in sync
update_columns: Final = {
"file_object": file_object.model_dump_json(),
"status": file_object.status,
"updated_by": user_api_key_dict.user_id,
}
if not create_if_missing:
await _managed_object_table(self.prisma_client).update_many(
where={"unified_object_id": unified_object_id},
data=update_columns,
)
return
await _managed_object_table(self.prisma_client).upsert(
where={"unified_object_id": unified_object_id},
data={
@ -289,12 +331,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
"team_id": user_api_key_dict.team_id,
"updated_by": user_api_key_dict.user_id,
"status": file_object.status,
**attribution_columns,
},
"update": {
"file_object": file_object.model_dump_json(),
"status": file_object.status,
"updated_by": user_api_key_dict.user_id,
}, # FIX: Update status and file_object on every operation to keep state in sync
"update": update_columns,
},
)
@ -1248,10 +1287,27 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
user_created_file_ids = await self.get_user_created_file_ids(user_api_key_dict, file_ids)
## Filter the response to only include the files created by the user
response.data = user_created_file_ids # type: ignore
self._scope_list_page_cursors(response, user_created_file_ids)
return response
return response
return response
@staticmethod
def _scope_list_page_cursors(response: AsyncCursorPage, data: List[OpenAIFileObject]) -> None:
"""Rebuild ``first_id`` / ``last_id`` from the caller-scoped page.
The upstream cursors point at rows that were just filtered out, so
leaving them in place discloses other callers' file ids. ``has_more``
is always cleared because ``after`` is never forwarded upstream, so
no further page is reachable through the proxy.
"""
if hasattr(response, "first_id"):
response.first_id = data[0].id if data else None
if hasattr(response, "last_id"):
response.last_id = data[-1].id if data else None
if hasattr(response, "has_more"):
response.has_more = False
async def afile_retrieve(
self, file_id: str, litellm_parent_otel_span: Optional[Span], llm_router: Optional[Router] = None
) -> OpenAIFileObject:

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-enterprise"
version = "0.1.54"
version = "0.1.55"
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.54"
version = "0.1.55"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-enterprise==",

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_DailyTeamSpend" ADD COLUMN IF NOT EXISTS "ptu_flat_cost" DOUBLE PRECISION NOT NULL DEFAULT 0.0;

View file

@ -0,0 +1,5 @@
-- Add api_key and request_tags columns to LiteLLM_ManagedObjectTable
-- Captured at batch-create time so CheckBatchCost can attribute batch-cost spend
-- back to the creating virtual key (and its tags) even when created_by is null.
ALTER TABLE "LiteLLM_ManagedObjectTable" ADD COLUMN IF NOT EXISTS "api_key" TEXT;
ALTER TABLE "LiteLLM_ManagedObjectTable" ADD COLUMN IF NOT EXISTS "request_tags" JSONB DEFAULT '[]';

View file

@ -0,0 +1,5 @@
-- AlterTable
ALTER TABLE "LiteLLM_VerificationToken" ADD COLUMN IF NOT EXISTS "settings_updated_at" TIMESTAMP(3);
-- AlterTable
ALTER TABLE "LiteLLM_DeletedVerificationToken" ADD COLUMN IF NOT EXISTS "settings_updated_at" TIMESTAMP(3);

View file

@ -30,7 +30,7 @@ model LiteLLM_BudgetTable {
end_users LiteLLM_EndUserTable[] // multiple end-users can have the same budget
tags LiteLLM_TagTable[] // multiple tags can have the same budget
team_membership LiteLLM_TeamMembership[] // budgets of Users within a Team
organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization
organization_membership LiteLLM_OrganizationMembership[] // budgets of Users within a Organization
}
// Models on proxy
@ -452,6 +452,7 @@ model LiteLLM_VerificationToken {
created_by String?
updated_at DateTime? @default(now()) @updatedAt @map("updated_at")
updated_by String?
settings_updated_at DateTime? @map("settings_updated_at")
last_active DateTime? // When this key was last used
rotation_count Int? @default(0) // Number of times key has been rotated
auto_rotate Boolean? @default(false) // Whether this key should be auto-rotated
@ -548,6 +549,7 @@ model LiteLLM_DeletedVerificationToken {
created_by String? // Original creator
updated_at DateTime? // Last update timestamp before deletion
updated_by String? // Last user who updated before deletion
settings_updated_at DateTime? // Last configuration change before deletion
last_active DateTime? // When this key was last used before deletion
rotation_count Int? @default(0)
auto_rotate Boolean? @default(false)
@ -893,6 +895,7 @@ model LiteLLM_DailyTeamSpend {
api_requests BigInt @default(0)
successful_requests BigInt @default(0)
failed_requests BigInt @default(0)
ptu_flat_cost Float @default(0.0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@ -985,6 +988,8 @@ model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use t
created_at DateTime @default(now())
created_by String?
team_id String?
api_key String?
request_tags Json? @default("[]")
updated_at DateTime @updatedAt
updated_by String?

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.84"
version = "0.4.85"
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.84"
version = "0.4.85"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

@ -197,6 +197,7 @@ standard_logging_payload_excluded_fields: Optional[List[str]] = (
None # Fields to exclude from StandardLoggingPayload before callbacks receive it
)
log_raw_request_response: bool = False
request_correlation_in_logs: bool = False
redact_messages_in_exceptions: Optional[bool] = False
redact_user_api_key_info: Optional[bool] = False
# When True (default — preserves historical behavior), the Router appends

View file

@ -1,4 +1,5 @@
import ast
import contextvars
import logging
import os
import sys
@ -6,12 +7,44 @@ from datetime import datetime
from logging import Formatter
from typing import Any, Final
import litellm
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.litellm_core_utils.secret_redaction import redact_string
set_verbose = False
session_id_var: Final[contextvars.ContextVar[str]] = contextvars.ContextVar("session_id", default="")
trace_id_var: Final[contextvars.ContextVar[str]] = contextvars.ContextVar("trace_id", default="")
_MAX_CORRELATION_ID_LENGTH: Final = 256
def _sanitize_correlation_id(value: str) -> str:
"""Strip control characters, bound length, and redact credential-shaped
content before a caller-controlled trace_id/session_id (e.g.
litellm_session_id, x-litellm-trace-id) is stamped into log lines.
Without the first two, a caller could embed \\r/\\n or terminal escape
sequences to forge fake log entries, or submit an oversized value repeated
across every log line for the request. Without the redaction, a caller
could smuggle a real credential (e.g. an sk-... key) through this field:
CorrelationContextFilter stamps trace_id/session_id onto the record after
SecretRedactionFilter has already run, so those two fields never otherwise
pass through credential redaction.
"""
stripped: Final = "".join(ch for ch in value if ch.isprintable())
return _redact_string(stripped[:_MAX_CORRELATION_ID_LENGTH])
def set_session_id(session_id: str) -> "contextvars.Token[str]":
return session_id_var.set(_sanitize_correlation_id(session_id))
def set_trace_id(trace_id: str) -> "contextvars.Token[str]":
return trace_id_var.set(_sanitize_correlation_id(trace_id))
if set_verbose is True:
logging.warning(
"`litellm.set_verbose` is deprecated. Please set `os.environ['LITELLM_LOG'] = 'DEBUG'` for debug logs."
@ -77,6 +110,28 @@ class SecretRedactionFilter(logging.Filter):
_secret_filter: Final = SecretRedactionFilter()
class CorrelationContextFilter(logging.Filter):
"""Stamps each log record with the current request's trace_id and session_id from contextvars.
Works in tandem with JsonFormatter: the formatter's record.__dict__ loop picks up these
attributes as first-class JSON fields without any formatter-level code.
"""
def filter(self, record: logging.LogRecord) -> bool:
if not litellm.request_correlation_in_logs:
return True
trace_id: Final = trace_id_var.get()
if trace_id:
record.trace_id = trace_id # rebind-ok: stamping the LogRecord is the Filter interface's contract
session_id: Final = session_id_var.get()
if session_id:
record.session_id = session_id # rebind-ok: stamping the LogRecord is the Filter interface's contract
return True
_correlation_filter: Final = CorrelationContextFilter()
json_logs = bool(os.getenv("JSON_LOGS", False))
# Create a handler for the logger (you may need to adapt this based on your needs)
log_level: Final = os.getenv("LITELLM_LOG", "DEBUG")
@ -84,6 +139,7 @@ numeric_level: Final[str] = getattr(logging, log_level.upper())
handler: Final = logging.StreamHandler()
handler.setLevel(numeric_level)
handler.addFilter(_secret_filter)
handler.addFilter(_correlation_filter)
def _try_parse_json_message(message: str) -> dict[str, Any] | None:
@ -146,6 +202,11 @@ def _get_standard_record_attrs() -> frozenset:
_STANDARD_RECORD_ATTRS: Final = _get_standard_record_attrs()
# CorrelationContextFilter is the only legitimate source for these two JSON fields;
# see JsonFormatter.format() for why they're excluded from the generic message-content
# and extra-attribute promotion paths.
_RESERVED_CORRELATION_FIELDS: Final = frozenset(("trace_id", "session_id"))
class JsonFormatter(Formatter):
def __init__(self):
@ -164,13 +225,18 @@ class JsonFormatter(Formatter):
"timestamp": self.formatTime(record),
}
# Parse embedded JSON or Python dict repr in message so sub-fields become first-class properties
# Parse embedded JSON or Python dict repr in message so sub-fields become first-class properties.
# trace_id/session_id are excluded here unconditionally (not just "if not already
# set") - CorrelationContextFilter is the only legitimate source for these two
# fields, and a message that merely happens to parse as JSON/dict (e.g. a proxy
# log line dumping raw request headers) must never be able to claim them, even on
# a record the filter hasn't stamped yet (no correlation context active for it).
parsed = _try_parse_json_message(message_str)
if parsed is None:
parsed = _try_parse_embedded_python_dict(message_str)
if parsed is not None:
for key, value in parsed.items():
if key not in json_record:
if key not in json_record and key not in _RESERVED_CORRELATION_FIELDS:
json_record[key] = value
# Include extra attributes passed via logger.debug("msg", extra={...})
@ -178,6 +244,18 @@ class JsonFormatter(Formatter):
if key not in _STANDARD_RECORD_ATTRS and key not in json_record:
json_record[key] = value
# trace_id/session_id are reserved: CorrelationContextFilter is the only
# legitimate source for these two fields. Without this, a message string
# that happens to parse as JSON/dict (e.g. a proxy log line dumping raw
# request headers) with a "trace_id"/"session_id" key would have already
# claimed the key at the parsed-message step above, and the extra-attributes
# loop's "key not in json_record" guard would then skip the real value -
# letting a caller-supplied header spoof another request's correlation ids.
for reserved_key in _RESERVED_CORRELATION_FIELDS:
value = getattr(record, reserved_key, None)
if value:
json_record[reserved_key] = value
# Set component/logger only if not already supplied via extra={...}
if "component" not in json_record:
json_record["component"] = record.name
@ -190,12 +268,34 @@ class JsonFormatter(Formatter):
return safe_dumps(json_record)
class CorrelationPlainFormatter(logging.Formatter):
"""Appends trace_id/session_id to plain-text log lines stamped by CorrelationContextFilter.
Mirrors JsonFormatter's handling of these two fields so request_correlation_in_logs
behaves the same whether or not json_logs is enabled.
"""
def format(self, record: logging.LogRecord) -> str:
formatted: Final = super().format(record)
trace_id: Final = getattr(record, "trace_id", None)
session_id: Final = getattr(record, "session_id", None)
if not trace_id and not session_id:
return formatted
parts: Final = tuple(
p
for p in (f"trace_id={trace_id}" if trace_id else None, f"session_id={session_id}" if session_id else None)
if p
)
return f"{formatted} [{' '.join(parts)}]"
# Function to set up exception handlers for JSON logging
def _setup_json_exception_handlers(formatter):
# Create a handler with JSON formatting for exceptions
error_handler: Final = logging.StreamHandler()
error_handler.setFormatter(formatter)
error_handler.addFilter(_secret_filter)
error_handler.addFilter(_correlation_filter)
# Setup excepthook for uncaught exceptions
def json_excepthook(exc_type, exc_value, exc_traceback):
@ -243,7 +343,7 @@ if json_logs:
handler.setFormatter(JsonFormatter())
_setup_json_exception_handlers(JsonFormatter())
else:
formatter: Final = logging.Formatter(
formatter: Final = CorrelationPlainFormatter(
"\033[92m%(asctime)s - %(name)s:%(levelname)s\033[0m: %(filename)s:%(lineno)s - %(message)s",
datefmt="%H:%M:%S",
)
@ -346,6 +446,7 @@ def _initialize_loggers_with_handler(handler: logging.Handler):
- Prevents bubbling to parent/root (critical to prevent duplicate JSON logs)
"""
handler.addFilter(_secret_filter)
handler.addFilter(_correlation_filter)
for lg in _get_loggers_to_initialize():
lg.handlers.clear() # remove any existing handlers
lg.addHandler(handler) # add JSON formatter handler

View file

@ -22,6 +22,7 @@ class ResponsesToCompletionBridgeHandlerInputKwargs(TypedDict):
model_response: "ModelResponse"
logging_obj: "LiteLLMLoggingObj"
custom_llm_provider: str
encoding: object
class ResponsesToCompletionBridgeHandler:
@ -102,35 +103,37 @@ class ResponsesToCompletionBridgeHandler:
from litellm import LiteLLMLoggingObj
from litellm.types.utils import ModelResponse
model: Final = kwargs.get("model")
typed_kwargs: Final[dict[str, object]] = kwargs
model: Final = typed_kwargs.get("model")
if model is None or not isinstance(model, str):
raise ValueError("model is required")
custom_llm_provider: Final = kwargs.get("custom_llm_provider")
custom_llm_provider: Final = typed_kwargs.get("custom_llm_provider")
if custom_llm_provider is None or not isinstance(custom_llm_provider, str):
raise ValueError("custom_llm_provider is required")
messages: Final = kwargs.get("messages")
messages: Final = typed_kwargs.get("messages")
if messages is None or not isinstance(messages, list):
raise ValueError("messages is required")
optional_params: Final = kwargs.get("optional_params")
optional_params: Final = typed_kwargs.get("optional_params")
if optional_params is None or not isinstance(optional_params, dict):
raise ValueError("optional_params is required")
litellm_params: Final = kwargs.get("litellm_params")
litellm_params: Final = typed_kwargs.get("litellm_params")
if litellm_params is None or not isinstance(litellm_params, dict):
raise ValueError("litellm_params is required")
headers: Final = kwargs.get("headers")
headers: Final = typed_kwargs.get("headers")
if headers is None or not isinstance(headers, dict):
raise ValueError("headers is required")
model_response: Final = kwargs.get("model_response")
model_response: Final = typed_kwargs.get("model_response")
if model_response is None or not isinstance(model_response, ModelResponse):
raise ValueError("model_response is required")
logging_obj: Final = kwargs.get("logging_obj")
logging_obj: Final = typed_kwargs.get("logging_obj")
if logging_obj is None or not isinstance(logging_obj, LiteLLMLoggingObj):
raise ValueError("logging_obj is required")
@ -143,6 +146,7 @@ class ResponsesToCompletionBridgeHandler:
model_response=model_response,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
encoding=typed_kwargs.get("encoding"),
)
def completion(
@ -205,7 +209,7 @@ class ResponsesToCompletionBridgeHandler:
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
encoding=kwargs.get("encoding"),
encoding=validated_kwargs["encoding"],
api_key=kwargs.get("api_key"),
json_mode=kwargs.get("json_mode"),
)
@ -230,7 +234,7 @@ class ResponsesToCompletionBridgeHandler:
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
encoding=kwargs.get("encoding"),
encoding=validated_kwargs["encoding"],
api_key=kwargs.get("api_key"),
json_mode=kwargs.get("json_mode"),
)
@ -303,7 +307,7 @@ class ResponsesToCompletionBridgeHandler:
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
encoding=kwargs.get("encoding"),
encoding=validated_kwargs["encoding"],
api_key=kwargs.get("api_key"),
json_mode=kwargs.get("json_mode"),
)
@ -328,7 +332,7 @@ class ResponsesToCompletionBridgeHandler:
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
encoding=kwargs.get("encoding"),
encoding=validated_kwargs["encoding"],
api_key=kwargs.get("api_key"),
json_mode=kwargs.get("json_mode"),
)

View file

@ -4,8 +4,8 @@ Handler for transforming /chat/completions api requests to litellm.responses req
import json
import os
from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping
from typing import TYPE_CHECKING, Any, Final, Literal, Union, cast
from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, Union, cast
from openai.types.responses.custom_tool_param import CustomToolParam
from openai.types.responses.response_input_param import (
@ -45,6 +45,9 @@ from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
if TYPE_CHECKING:
from openai.types.responses import ResponseInputImageParam
from openai.types.responses.response_text_config_param import (
ResponseTextConfigParam as ResponseText,
)
from pydantic import BaseModel
from litellm import LiteLLMLoggingObj, ModelResponse
@ -57,6 +60,19 @@ if TYPE_CHECKING:
ChatCompletionThinkingBlock,
OpenAIMessageContentListBlock,
)
from litellm.types.utils import Choices
class _ReasoningSummaryText(TypedDict):
type: str
text: str
class _BuiltReasoningItem(TypedDict):
type: Literal["reasoning"]
id: str
encrypted_content: str | None
summary: Sequence[_ReasoningSummaryText]
def _get_reasoning_items(
@ -72,13 +88,13 @@ def _get_reasoning_items(
def _build_reasoning_item(
item_id: str,
encrypted_content: str | None,
summary_raw: Any,
) -> dict[str, Any]:
summary_raw: Iterable[object] | None,
) -> _BuiltReasoningItem:
"""Build a ChatCompletionReasoningItem-shaped dict from raw response data.
Handles both pydantic objects (attribute access) and plain dicts.
"""
summary: Final[list[dict[str, Any]]] = []
summary: Final[list[_ReasoningSummaryText]] = []
for s in summary_raw or []:
if isinstance(s, dict):
summary.append({"type": s.get("type", "summary_text"), "text": s.get("text", "")})
@ -98,7 +114,7 @@ def _build_reasoning_item(
class _ChatToolCallDict(ChatCompletionToolCallChunk, total=False):
provider_specific_fields: Mapping[str, Any]
provider_specific_fields: Mapping[str, object]
def _tool_call_dict_from_output_item(item: Mapping[str, Any], index: int) -> _ChatToolCallDict:
@ -142,10 +158,10 @@ def _flat_responses_tool_choice(choice_type: str, name: str) -> ToolChoiceFuncti
def _reasoning_item_to_response_input(
r_item: ChatCompletionReasoningItem | dict[str, Any],
) -> dict[str, Any]:
r_item: ChatCompletionReasoningItem,
) -> dict[str, object]:
"""Convert a stored ChatCompletionReasoningItem back to a Responses API input item."""
r_input: Final[dict[str, Any]] = {
r_input: Final[dict[str, object]] = {
"type": "reasoning",
"id": r_item.get("id") or f"rs_{id(r_item)}",
# summary is always required by the Responses API, even when empty
@ -181,7 +197,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
return _flat_responses_tool_choice(choice_type, nested_name)
return tool_choice
def _handle_raw_dict_response_item(self, item: dict[str, Any], index: int) -> tuple[Any | None, int]:
def _handle_raw_dict_response_item(self, item: dict[str, Any], index: int) -> tuple["Choices | None", int]:
"""
Handle raw dict response items from Responses API (e.g., GPT-5 Codex format).
@ -228,8 +244,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
def convert_chat_completion_messages_to_responses_api(
self, messages: list["AllMessageValues"]
) -> tuple[list[Any], str | None]:
input_items: Final[list[Any]] = []
) -> tuple[list[object], str | None]:
input_items: Final[list[object]] = []
instructions: str | None = None
custom_tool_call_ids: Final = frozenset(
tool_call["id"]
@ -270,7 +286,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
# Convert tool message to function call output format
# The Responses API expects 'output' to be a list with input_text/input_image types
# Using list format for consistency across text and multimodal content
tool_output: list[dict[str, Any]]
tool_output: list[dict[str, object]]
if content is None:
tool_output = []
elif isinstance(content, str):
@ -308,7 +324,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
function = tool_call.get("function")
custom = tool_call.get("custom")
if function:
input_tool_call: dict[str, Any] = {
input_tool_call: dict[str, object] = {
"type": "function_call",
"call_id": tool_call["id"],
}
@ -376,15 +392,15 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
elif key == "web_search_options":
self._add_web_search_tool(responses_api_request, value)
def _build_sanitized_litellm_params(self, litellm_params: dict) -> dict[str, Any]:
def _build_sanitized_litellm_params(self, litellm_params: dict) -> dict[str, object]:
"""Build sanitized litellm_params with merged metadata."""
responses_optional_param_keys: Final = set(ResponsesAPIOptionalRequestParams.__annotations__.keys())
sanitized: Final[dict[str, Any]] = {
sanitized: Final[dict[str, object]] = {
key: value for key, value in litellm_params.items() if key not in responses_optional_param_keys
}
legacy_metadata: Final = litellm_params.get("metadata")
existing_litellm_metadata: Final = litellm_params.get("litellm_metadata")
merged_litellm_metadata: Final[dict[str, Any]] = {}
merged_litellm_metadata: Final[dict[str, object]] = {}
if isinstance(legacy_metadata, dict):
merged_litellm_metadata.update(legacy_metadata)
if isinstance(existing_litellm_metadata, dict):
@ -424,7 +440,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
litellm_params: dict,
headers: dict,
litellm_logging_obj: "LiteLLMLoggingObj",
client: Any | None = None,
client: object | None = None,
) -> dict:
(
input_items,
@ -498,9 +514,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
@staticmethod
def _convert_response_output_to_choices(
output_items: list[Any],
handle_raw_dict_callback: Callable | None = None,
) -> list[Any]:
output_items: Sequence[object],
handle_raw_dict_callback: Callable[..., tuple["Choices | None", int]] | None = None,
) -> list["Choices"]:
"""
Convert Responses API output items to chat completion choices.
@ -529,11 +545,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
choices: Final[list[Choices]] = []
index = 0
reasoning_content: str | None = None
pending_reasoning_item: dict[str, Any] | None = None
pending_reasoning_item: _BuiltReasoningItem | None = None
# Collect all tool calls to put them in a single choice
# (Chat Completions API expects all tool calls in one message)
accumulated_tool_calls: Final[list[dict[str, Any]]] = []
accumulated_tool_calls: Final[list[Mapping[str, object]]] = []
tool_call_index = 0
for item in output_items:
@ -640,7 +656,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
return choices
@classmethod
def _extract_output_from_completed_event(cls, parsed_chunk: dict[str, Any]) -> list[dict[str, Any]] | None:
def _extract_output_from_completed_event(cls, parsed_chunk: Mapping[str, object]) -> list[dict[str, object]] | None:
response_payload: Final = parsed_chunk.get("response")
if not isinstance(response_payload, dict):
return None
@ -650,12 +666,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
return cast(list[dict[str, Any]], response_output)
@classmethod
def _recover_output_items_from_raw_sse(cls, raw_sse: str | None) -> list[dict[str, Any]]:
def _recover_output_items_from_raw_sse(cls, raw_sse: str | None) -> list[dict[str, object]]:
if not raw_sse or not isinstance(raw_sse, str):
return []
recovered_output_items: Final[dict[int, dict[str, Any]]] = {}
recovered_text_only_items: Final[dict[int, dict[str, Any]]] = {}
recovered_output_items: Final[dict[int, dict[str, object]]] = {}
recovered_text_only_items: Final[dict[int, dict[str, object]]] = {}
for chunk in raw_sse.splitlines():
parsed_chunk = parse_sse_json_chunk(chunk)
@ -690,7 +706,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
# but text-only items at indices without a matching OUTPUT_ITEM_DONE
# must still be preserved (e.g. multi-output responses where some
# indices only emitted OUTPUT_TEXT_DONE).
merged_items: Final[dict[int, dict[str, Any]]] = {**recovered_text_only_items}
merged_items: Final[dict[int, dict[str, object]]] = {**recovered_text_only_items}
merged_items.update(recovered_output_items)
if merged_items:
@ -699,7 +715,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
return []
@classmethod
def _recover_output_items_from_logging(cls, logging_obj: "LiteLLMLoggingObj") -> list[dict[str, Any]]:
def _recover_output_items_from_logging(cls, logging_obj: "LiteLLMLoggingObj") -> list[dict[str, object]]:
model_call_details: Final = getattr(logging_obj, "model_call_details", {}) or {}
original_response: Final = model_call_details.get("original_response")
return cls._recover_output_items_from_raw_sse(original_response)
@ -714,7 +730,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
messages: list["AllMessageValues"],
optional_params: dict,
litellm_params: dict,
encoding: Any,
encoding: object,
api_key: str | None = None,
json_mode: bool | None = None,
) -> "ModelResponse":
@ -788,7 +804,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
) -> BaseModelResponseIterator:
return OpenAiResponsesToChatCompletionStreamIterator(streaming_response, sync_stream, json_mode)
def _convert_content_str_to_input_text(self, content: str, role: str) -> dict[str, Any]:
def _convert_content_str_to_input_text(self, content: str, role: str) -> dict[str, object]:
if role == "user" or role == "system" or role == "tool":
return {"type": "input_text", "text": content}
else:
@ -825,13 +841,13 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
def _convert_content_to_responses_format(
self,
content: str
| list[Any]
| list[object]
| Iterable[
Union["OpenAIMessageContentListBlock", "ChatCompletionThinkingBlock", "ChatCompletionRedactedThinkingBlock"]
]
| None,
role: str,
) -> list[dict[str, Any]]:
) -> list[dict[str, object]]:
"""Convert chat completion content to responses API format"""
from litellm.types.llms.openai import ChatCompletionImageObject
@ -973,7 +989,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
return optional_params
def _map_reasoning_effort(self, reasoning_effort: str | dict[str, Any]) -> Reasoning | None:
def _map_reasoning_effort(self, reasoning_effort: str | Reasoning) -> Reasoning | None:
# If dict is passed, convert it directly to Reasoning object
if isinstance(reasoning_effort, dict):
return Reasoning(**reasoning_effort)
@ -1006,7 +1022,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
def _add_web_search_tool(
self,
responses_api_request: ResponsesAPIOptionalRequestParams,
web_search_options: Any,
web_search_options: object,
) -> None:
"""
Add web search tool to responses API request.
@ -1024,14 +1040,14 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
tools = []
responses_api_request["tools"] = tools
web_search_tool: Final[dict[str, Any]] = {"type": "web_search"}
web_search_tool: Final[dict[str, object]] = {"type": "web_search"}
if isinstance(web_search_options, dict):
web_search_tool.update(web_search_options)
# Cast to Any to match the expected union type for tools list items
tools.append(cast(Any, web_search_tool))
def _transform_response_format_to_text_format(self, response_format: dict[str, Any] | Any) -> dict[str, Any] | None:
def _transform_response_format_to_text_format(self, response_format: object) -> "ResponseText | None":
"""
Transform Chat Completion response_format parameter to Responses API text.format parameter.
@ -1130,7 +1146,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
def __init__(self, streaming_response, sync_stream: bool, json_mode: bool | None = False):
def __init__(
self,
streaming_response: Union[Iterator[str], AsyncIterator[str], "ModelResponse", "BaseModel"],
sync_stream: bool,
json_mode: bool | None = False,
):
super().__init__(streaming_response, sync_stream, json_mode)
self._chat_completion_id: str | None = None
self._tool_call_index_map: dict[int, int] = {} # mutable-ok: per-stream accumulator state
@ -1387,7 +1408,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
finish_reason: Final = "tool_calls" if has_function_calls else "stop"
# Extract reasoning items with encrypted_content for round-tripping
completed_reasoning_items: list[dict[str, Any]] | None = None
completed_reasoning_items: list[_BuiltReasoningItem] | None = None
for item in output_items:
if not isinstance(item, dict) or item.get("type") != "reasoning":
continue
@ -1439,7 +1460,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
]
)
def chunk_parser(self, chunk: dict) -> "ModelResponseStream":
def chunk_parser(self, chunk: dict[str, object]) -> "ModelResponseStream":
"""
Parse a Responses API streaming chunk and convert to OpenAI format.

View file

@ -1478,6 +1478,10 @@ CLOUDZERO_MAX_FETCHED_DATA_RECORDS: Final = int(os.getenv("CLOUDZERO_MAX_FETCHED
SPEND_LOG_CLEANUP_JOB_NAME: Final = "spend_log_cleanup"
KEY_ROTATION_JOB_NAME: Final = "litellm_key_rotation_job"
EXPIRED_UI_SESSION_KEY_CLEANUP_JOB_NAME: Final = "litellm_expired_ui_session_key_cleanup_job"
WEEKLY_SPEND_REPORT_JOB_ID: Final = "weekly_spend_report_job"
MONTHLY_SPEND_REPORT_JOB_ID: Final = "monthly_spend_report_job"
PROMETHEUS_FALLBACK_STATS_JOB_ID: Final = "prometheus_fallback_stats_job"
SLACK_DAILY_REPORT_LOCK_ID: Final = "slack_daily_report"
SPEND_LOG_RUN_LOOPS: Final = int(os.getenv("SPEND_LOG_RUN_LOOPS", 500))
SPEND_LOG_CLEANUP_BATCH_SIZE: Final = int(os.getenv("SPEND_LOG_CLEANUP_BATCH_SIZE", 1000))
SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES = int(os.getenv("SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES", 3))
@ -1493,6 +1497,8 @@ SPEND_LOG_QUEUE_POLL_INTERVAL: Final = float(os.getenv("SPEND_LOG_QUEUE_POLL_INT
SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE: Final = int(os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000))
DEFAULT_CRON_JOB_LOCK_TTL_SECONDS: Final = int(os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)) # 1 minute
PROXY_BUDGET_RESCHEDULER_MIN_TIME: Final = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597))
RESET_BUDGET_JOB_BATCH_SIZE: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_BATCH_SIZE", "500")))
RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN", "100")))
PROXY_BATCH_POLLING_INTERVAL: Final = int(os.getenv("PROXY_BATCH_POLLING_INTERVAL", 3600))
MAX_OBJECTS_PER_POLL_CYCLE: Final = max(1, int(os.getenv("MAX_OBJECTS_PER_POLL_CYCLE", 50)))
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS: Final = max(1, int(os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7)))
@ -1719,3 +1725,18 @@ BROWSER_SECURITY_HEADERS: Final[frozenset[str]] = frozenset(
)
UNSAFE_PROXY_RESPONSE_HEADERS: Final[frozenset[str]] = HTTP_FRAMING_HEADERS | BROWSER_SECURITY_HEADERS
# PTU reservation rollup writes rows to LiteLLM_DailyTeamSpend with this
# sentinel api_key so PTU flat cost stays distinguishable from real per-request
# spend under the table's composite unique constraint.
PTU_SENTINEL_API_KEY: Final[str] = "__ptu_flat_cost__"
PTU_ROLLUP_JOB_ID: Final[str] = "ptu_flat_cost_rollup_job"
PTU_ROLLUP_LOCK_TTL_SECONDS: Final[int] = 900
# Furthest back the catch-up pass looks for unpriced PTU days when a deployment
# declares no ptu_effective_from, bounding the scan for an open-ended window.
PTU_ROLLUP_MAX_BACKFILL_DAYS: Final[int] = 90
# Slack allowed when deciding a sentinel row is stale. The row's updated_at and the
# run's cutoff are stamped by different hosts, so clock skew between them must not let
# one run delete a charge another just wrote. A stale row is hours old and a concurrent
# one is seconds old, so a few minutes separates them.
PTU_PRUNE_SKEW_GRACE_SECONDS: Final[int] = 300

View file

@ -7,7 +7,7 @@ import asyncio
import contextvars
from collections.abc import Coroutine
from functools import partial
from typing import Any, Final
from typing import Final
import httpx
@ -21,8 +21,10 @@ from litellm.types.llms.openai_evals import (
CancelRunResponse,
CreateEvalRequest,
CreateRunRequest,
DataSourceConfig,
DeleteEvalResponse,
Eval,
GraderConfig,
ListEvalsParams,
ListEvalsResponse,
ListRunsParams,
@ -41,13 +43,13 @@ DEFAULT_OPENAI_API_BASE: Final = "https://api.openai.com"
@client
async def acreate_eval(
data_source_config: dict[str, Any],
testing_criteria: list[dict[str, Any]],
data_source_config: DataSourceConfig,
testing_criteria: list[GraderConfig],
name: str | None = None,
metadata: dict[str, Any] | None = None,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
metadata: dict[str, object] | 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,
custom_llm_provider: str | None = None,
**kwargs,
@ -110,17 +112,17 @@ async def acreate_eval(
@client
def create_eval(
data_source_config: dict[str, Any],
testing_criteria: list[dict[str, Any]],
data_source_config: DataSourceConfig,
testing_criteria: list[GraderConfig],
name: str | None = None,
metadata: dict[str, Any] | None = None,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
metadata: dict[str, object] | 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,
custom_llm_provider: str | None = None,
**kwargs,
) -> Eval | Coroutine[Any, Any, Eval]:
) -> Eval | Coroutine[object, object, Eval]:
"""
Create a new evaluation
@ -231,8 +233,8 @@ async def alist_evals(
before: str | None = None,
order: str | None = None,
order_by: str | None = None,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
**kwargs,
@ -300,12 +302,12 @@ def list_evals(
before: str | None = None,
order: str | None = None,
order_by: str | None = None,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
**kwargs,
) -> ListEvalsResponse | Coroutine[Any, Any, ListEvalsResponse]:
) -> ListEvalsResponse | Coroutine[object, object, ListEvalsResponse]:
"""
List all evaluations
@ -413,8 +415,8 @@ def list_evals(
@client
async def aget_eval(
eval_id: str,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
**kwargs,
@ -470,12 +472,12 @@ async def aget_eval(
@client
def get_eval(
eval_id: str,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
**kwargs,
) -> Eval | Coroutine[Any, Any, Eval]:
) -> Eval | Coroutine[object, object, Eval]:
"""
Get an evaluation by ID
@ -564,10 +566,10 @@ def get_eval(
async def aupdate_eval(
eval_id: str,
name: str | None = None,
metadata: dict[str, Any] | None = None,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
metadata: dict[str, object] | 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,
custom_llm_provider: str | None = None,
**kwargs,
@ -630,14 +632,14 @@ async def aupdate_eval(
def update_eval(
eval_id: str,
name: str | None = None,
metadata: dict[str, Any] | None = None,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
metadata: dict[str, object] | 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,
custom_llm_provider: str | None = None,
**kwargs,
) -> Eval | Coroutine[Any, Any, Eval]:
) -> Eval | Coroutine[object, object, Eval]:
"""
Update an evaluation
@ -783,8 +785,8 @@ def update_eval(
@client
async def adelete_eval(
eval_id: str,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
**kwargs,
@ -840,12 +842,12 @@ async def adelete_eval(
@client
def delete_eval(
eval_id: str,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
**kwargs,
) -> DeleteEvalResponse | Coroutine[Any, Any, DeleteEvalResponse]:
) -> DeleteEvalResponse | Coroutine[object, object, DeleteEvalResponse]:
"""
Delete an evaluation
@ -933,8 +935,8 @@ def delete_eval(
@client
async def acancel_eval(
eval_id: str,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
**kwargs,
@ -990,12 +992,12 @@ async def acancel_eval(
@client
def cancel_eval(
eval_id: str,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
**kwargs,
) -> CancelEvalResponse | Coroutine[Any, Any, CancelEvalResponse]:
) -> CancelEvalResponse | Coroutine[object, object, CancelEvalResponse]:
"""
Cancel a running evaluation
@ -1092,12 +1094,12 @@ def cancel_eval(
@client
async def acreate_run(
eval_id: str,
data_source: dict[str, Any],
data_source: dict[str, object],
name: str | None = None,
metadata: dict[str, Any] | None = None,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
metadata: dict[str, object] | 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,
custom_llm_provider: str | None = None,
**kwargs,
@ -1161,16 +1163,16 @@ async def acreate_run(
@client
def create_run(
eval_id: str,
data_source: dict[str, Any],
data_source: dict[str, object],
name: str | None = None,
metadata: dict[str, Any] | None = None,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_body: dict[str, Any] | None = None,
metadata: dict[str, object] | 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,
custom_llm_provider: str | None = None,
**kwargs,
) -> Run | Coroutine[Any, Any, Run]:
) -> Run | Coroutine[object, object, Run]:
"""
Create a new run for an evaluation
@ -1280,8 +1282,8 @@ async def alist_runs(
after: str | None = None,
before: str | None = None,
order: str | None = None,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
**kwargs,
@ -1349,12 +1351,12 @@ def list_runs(
after: str | None = None,
before: str | None = None,
order: str | None = None,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
**kwargs,
) -> ListRunsResponse | Coroutine[Any, Any, ListRunsResponse]:
) -> ListRunsResponse | Coroutine[object, object, ListRunsResponse]:
"""
List all runs for an evaluation
@ -1462,8 +1464,8 @@ def list_runs(
async def aget_run(
eval_id: str,
run_id: str,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
**kwargs,
@ -1522,12 +1524,12 @@ async def aget_run(
def get_run(
eval_id: str,
run_id: str,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
**kwargs,
) -> Run | Coroutine[Any, Any, Run]:
) -> Run | Coroutine[object, object, Run]:
"""
Get a specific run
@ -1618,8 +1620,8 @@ def get_run(
async def acancel_run(
eval_id: str,
run_id: str,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
**kwargs,
@ -1678,12 +1680,12 @@ async def acancel_run(
def cancel_run(
eval_id: str,
run_id: str,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
**kwargs,
) -> CancelRunResponse | Coroutine[Any, Any, CancelRunResponse]:
) -> CancelRunResponse | Coroutine[object, object, CancelRunResponse]:
"""
Cancel a running run
@ -1783,8 +1785,8 @@ def cancel_run(
async def adelete_run(
eval_id: str,
run_id: str,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
**kwargs,
@ -1843,12 +1845,12 @@ async def adelete_run(
def delete_run(
eval_id: str,
run_id: str,
extra_headers: dict[str, Any] | None = None,
extra_query: dict[str, Any] | None = None,
extra_headers: dict[str, object] | None = None,
extra_query: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
**kwargs,
) -> RunDeleteResponse | Coroutine[Any, Any, RunDeleteResponse]:
) -> RunDeleteResponse | Coroutine[object, object, RunDeleteResponse]:
"""
Delete a run

View file

@ -315,7 +315,12 @@ def image_generation(
or get_secret_str("AZURE_API_KEY")
)
azure_ad_token: Final = optional_params.pop("azure_ad_token", None) or get_secret_str("AZURE_AD_TOKEN")
azure_ad_token_param: Final = optional_params.pop("azure_ad_token", None)
azure_ad_token: Final = (
azure_ad_token_param
if isinstance(azure_ad_token_param, str) and azure_ad_token_param
else get_secret_str("AZURE_AD_TOKEN")
)
# Create azure_ad_token_provider from tenant_id, client_id, client_secret if not already provided
if azure_ad_token_provider is None:

View file

@ -9,6 +9,7 @@ from datetime import timedelta
from typing import TYPE_CHECKING, Any, Final, Literal
from openai import APIError
from pydantic import TypeAdapter
import litellm
import litellm.litellm_core_utils
@ -16,7 +17,7 @@ import litellm.litellm_core_utils.litellm_logging
import litellm.types
from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.constants import HOURS_IN_A_DAY
from litellm.constants import HOURS_IN_A_DAY, SLACK_DAILY_REPORT_LOCK_ID
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.integrations.SlackAlerting.budget_alert_types import get_budget_alert_type
from litellm.integrations.SlackAlerting.hanging_request_check import (
@ -33,10 +34,14 @@ from litellm.llms.custom_httpx.http_handler import (
from litellm.proxy._types import (
AlertType,
CallInfo,
InvitationModel,
InvitationNew,
Litellm_EntityType,
UserAPIKeyAuth,
VirtualKeyEvent,
WebhookEvent,
)
from litellm.repositories.table_repositories import InvitationLinkRepository
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.user_repository import UserRepository
from litellm.types.integrations.slack_alerting import *
@ -46,6 +51,7 @@ from .batching_handler import send_to_webhook, squash_payloads
from .utils import process_slack_alerting_variables
if TYPE_CHECKING:
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
from litellm.router import Router as _Router
Router = _Router
@ -1081,6 +1087,44 @@ Model Info:
if email_logo_url is not None or email_support_contact is not None:
raise ValueError(f"Trying to Customize Email Alerting\n {CommonProxyErrors.not_premium_user.value}")
async def _construct_user_invitation_link(self, recipient_user_id: str | None, base_url: str) -> str:
from litellm.proxy.management_helpers.user_invitation import (
create_invitation_for_user,
)
from litellm.proxy.proxy_server import prisma_client
if recipient_user_id is None or prisma_client is None:
return base_url
try:
existing_invitations: Final = TypeAdapter(list[InvitationModel]).validate_python(
await InvitationLinkRepository(prisma_client).table.find_many( # pyright: ignore[reportAny] # untyped prisma boundary (any-ok), result validated by TypeAdapter
where={"user_id": recipient_user_id}, # mutable-ok: prisma find_many requires a dict where filter
order={"created_at": "desc"}, # mutable-ok: prisma find_many requires a dict order arg
),
from_attributes=True,
)
invitation: Final = (
existing_invitations[0]
if existing_invitations
else TypeAdapter(InvitationModel).validate_python(
await create_invitation_for_user(
data=InvitationNew(user_id=recipient_user_id),
user_api_key_dict=UserAPIKeyAuth(user_id=recipient_user_id),
),
from_attributes=True,
)
)
except Exception as e: # noqa: BLE001 # best-effort link build; any DB/creation failure falls back to base_url
verbose_proxy_logger.error(
"Error creating invitation link for user_id %s: %s",
recipient_user_id,
str(e),
)
return base_url
return f"{base_url.rstrip('/')}/ui/onboarding?invitation_id={invitation.id}"
async def send_key_created_or_user_invited_email(self, webhook_event: WebhookEvent) -> bool:
try:
from litellm.proxy.utils import send_email
@ -1139,11 +1183,14 @@ Model Info:
team_row: Final = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
if team_row is not None:
team_name = team_row.team_alias or "-"
invitation_link: Final = await self._construct_user_invitation_link(
recipient_user_id=recipient_user_id, base_url=base_url
)
email_html_content = USER_INVITED_EMAIL_TEMPLATE.format(
email_logo_url=email_logo_url,
recipient_email=recipient_email,
team_name=team_name,
base_url=base_url,
base_url=invitation_link,
email_support_contact=email_support_contact,
)
else:
@ -1530,7 +1577,11 @@ Model Info:
except Exception:
pass
async def _run_scheduler_helper(self, llm_router) -> bool:
async def _run_scheduler_helper(
self,
llm_router,
pod_lock_manager: "PodLockManager | None" = None,
) -> bool:
"""
Returns:
- True -> report sent
@ -1555,6 +1606,16 @@ Model Info:
interval_seconds: Final = self.alerting_args.daily_report_frequency
if current_time - report_sent >= interval_seconds:
if (
pod_lock_manager is not None
and (
await pod_lock_manager.acquire_lock(
cronjob_id=SLACK_DAILY_REPORT_LOCK_ID, ttl=interval_seconds, allow_reentrant=False
)
)
is False
):
return False
# Sneak in the reporting logic here
await self.send_daily_reports(router=llm_router)
# Also, don't forget to update the report_sent time after sending the report!
@ -1566,7 +1627,11 @@ Model Info:
return report_sent_bool
async def _run_scheduled_daily_report(self, llm_router: Any | None = None):
async def _run_scheduled_daily_report(
self,
llm_router: Any | None = None,
pod_lock_manager: "PodLockManager | None" = None,
):
"""
If 'daily_reports' enabled
@ -1579,7 +1644,7 @@ Model Info:
if "daily_reports" in self.alert_types:
while True:
await self._run_scheduler_helper(llm_router=llm_router)
await self._run_scheduler_helper(llm_router=llm_router, pod_lock_manager=pod_lock_manager)
interval = random.randint(
self.alerting_args.report_check_interval - 3,
self.alerting_args.report_check_interval + 3,

View file

@ -382,19 +382,23 @@ class AnthropicCacheControlHook(CustomPromptManagement):
model: str,
custom_llm_provider: str | None,
tools: list | None = None,
enable_prompt_caching: bool | None = None,
) -> list[CacheControlInjectionPoint]:
"""Default breakpoints when ``litellm.enable_anthropic_prompt_caching`` is on.
Caches the system prompt and the trailing turn, so the stable prefix
(system + tools + history) is reused while the breakpoint advances with
the conversation. Returns [] (stand down) when the flag is off, the
provider does not consume cache_control breakpoints (only anthropic /
bedrock do), the model lacks prompt-caching support, or the request
already carries client-supplied cache_control.
``enable_prompt_caching`` is the per-request override (stamped from key
metadata by the proxy); True turns auto-injection on for this request
even when the global flag is off. Caches the system prompt and the
trailing turn, so the stable prefix (system + tools + history) is
reused while the breakpoint advances with the conversation. Returns []
(stand down) when neither flag is on, the provider does not consume
cache_control breakpoints (only anthropic / bedrock do), the model
lacks prompt-caching support, or the request already carries
client-supplied cache_control.
"""
import litellm
if litellm.enable_anthropic_prompt_caching is not True:
if litellm.enable_anthropic_prompt_caching is not True and enable_prompt_caching is not True:
return []
provider = custom_llm_provider
@ -433,6 +437,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
model: str,
custom_llm_provider: str | None,
tools: list | None = None,
enable_prompt_caching: bool | None = None,
) -> None:
"""For /chat/completions: resolve the injection points the request should carry.
@ -458,6 +463,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
model=model,
custom_llm_provider=custom_llm_provider,
tools=tools,
enable_prompt_caching=enable_prompt_caching,
)
if points:
non_default_params["cache_control_injection_points"] = points
@ -478,12 +484,17 @@ class AnthropicCacheControlHook(CustomPromptManagement):
judgment happens once per request; points a prior pass wrote back
carry the judged stamp and are never re-judged (see
``_should_stand_down``). When none are configured but
``litellm.enable_anthropic_prompt_caching`` is on, synthesize default
breakpoints for the native /v1/messages path. Pops the key from kwargs;
``litellm.enable_anthropic_prompt_caching`` or the per-request
``enable_prompt_caching`` kwarg (stamped from key metadata) is on,
synthesize default breakpoints for the native /v1/messages path. Pops
both keys from kwargs;
if remaining (non-message) points exist they are written back so
downstream transforms can handle them.
"""
typed_messages = cast(list[AllMessageValues], messages) # cast-ok: Anthropic-shaped dicts from v1/messages
enable_prompt_caching: Final = cast( # cast-ok: kwargs is untyped; key stamped as bool by the proxy
bool | None, kwargs.pop("enable_prompt_caching", None)
)
configured: Final = cast( # cast-ok: kwargs is untyped; this key only holds the documented injection-point list
list[CacheControlInjectionPoint] | None, kwargs.pop("cache_control_injection_points", None)
)
@ -497,6 +508,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
tools=tools,
model=model,
custom_llm_provider=custom_llm_provider,
enable_prompt_caching=enable_prompt_caching,
)
if not injection_points:
return messages, system

View file

@ -1,4 +1,5 @@
import json
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final
from typing_extensions import override
@ -12,7 +13,7 @@ from litellm.litellm_core_utils.redact_messages import (
should_redact_message_logging,
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.types.utils import StandardLoggingPayload
from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall, StandardLoggingPayload
if TYPE_CHECKING:
from opentelemetry.trace import Span
@ -22,6 +23,7 @@ from litellm.integrations._types.open_inference import (
ImageAttributes,
MessageAttributes,
MessageContentAttributes,
OpenInferenceMimeTypeValues,
OpenInferenceSpanKindValues,
SpanAttributes,
ToolCallAttributes,
@ -480,6 +482,7 @@ def set_attributes(span: "Span", kwargs, response_obj, attributes: type[BaseLLMO
response_obj_for_attrs,
slp,
)
_safe_emit("mcp tool attrs", _maybe_set_mcp_tool_attrs, span, kwargs, slp, response_obj_for_attrs)
def _sanitize_optional_params(optional_params: dict | None) -> dict:
@ -538,9 +541,12 @@ def _set_request_attributes(
if optional_params.get("user"):
safe_set_attribute(span, "llm.user", optional_params.get("user"))
if response_obj and response_obj.get("id"):
if not hasattr(response_obj, "get"):
return
if response_obj.get("id"):
safe_set_attribute(span, "llm.response.id", response_obj.get("id"))
if response_obj and response_obj.get("model"):
if response_obj.get("model"):
safe_set_attribute(span, "llm.response.model", response_obj.get("model"))
@ -588,6 +594,8 @@ def _coerce_response_obj_for_attrs(response_obj):
- dicts and Pydantic models that already expose `.get` are returned
unchanged (preserves all current behavior, including the Responses API
flow which relies on Pydantic attribute access).
- Pydantic models without `.get` (e.g. the MCP SDK's `CallToolResult`,
logged for `call_mcp_tool` spans) are dumped to a dict.
- `httpx.Response` and other text-only responses (passthrough routes)
are JSON-decoded so the standard extraction paths can read fields like
`id`, `model`, and `usage`. On failure the original object is returned
@ -595,6 +603,9 @@ def _coerce_response_obj_for_attrs(response_obj):
"""
if response_obj is None or hasattr(response_obj, "get"):
return response_obj
dumped: Final = _to_plain_dict(response_obj)
if isinstance(dumped, dict):
return dumped
text: Final = getattr(response_obj, "text", None)
if isinstance(text, str) and text:
try:
@ -1058,3 +1069,65 @@ def _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs):
except Exception:
return None
return None
def _maybe_set_mcp_tool_attrs(
span: "Span",
kwargs: Mapping[str, object],
standard_logging_payload: StandardLoggingPayload | None,
coerced_response_obj: object,
) -> None:
"""Render `call_mcp_tool` spans as OpenInference TOOL spans.
MCP tool calls carry neither `messages` nor `choices`, so the generic
extraction paths leave Input/Output blank. The tool name and arguments live
in `metadata.mcp_tool_call_metadata`; the result is an MCP `CallToolResult`
whose `content` is a list of typed parts.
"""
if standard_logging_payload is None:
return
if standard_logging_payload.get("call_type") != CallTypes.call_mcp_tool.value:
return
metadata: Final = standard_logging_payload.get("metadata")
mcp_meta: Final[StandardLoggingMCPToolCall | None] = metadata.get("mcp_tool_call_metadata") if metadata else None
if mcp_meta is None:
return
tool_name: Final = mcp_meta.get("name") or mcp_meta.get("namespaced_tool_name")
if tool_name:
safe_set_attribute(span, SpanAttributes.TOOL_NAME, tool_name)
if should_redact_message_logging(kwargs): # pyright: ignore[reportArgumentType] # reads, never mutates
return
arguments: Final[object] = mcp_meta.get("arguments")
if arguments is not None:
safe_set_attribute(span, SpanAttributes.INPUT_VALUE, safe_dumps(arguments))
safe_set_attribute(span, SpanAttributes.INPUT_MIME_TYPE, OpenInferenceMimeTypeValues.JSON.value)
_set_mcp_tool_output(span, coerced_response_obj)
def _has_only_text_parts(content: object) -> bool:
return not isinstance(content, list) or all(_coerce_text([part]) is not None for part in content)
def _set_mcp_tool_output(span: "Span", coerced_response_obj: object) -> None:
if not isinstance(coerced_response_obj, Mapping):
return
content: Final[object] = coerced_response_obj.get("content")
text: Final[str | None] = _coerce_text(content)
if text and _has_only_text_parts(content):
safe_set_attribute(span, SpanAttributes.OUTPUT_VALUE, text)
safe_set_attribute(span, SpanAttributes.OUTPUT_MIME_TYPE, OpenInferenceMimeTypeValues.TEXT.value)
return
structured: Final[object] = coerced_response_obj.get("structuredContent")
payload: Final[object] = content if content else structured if structured is not None else content
if payload is None:
return
safe_set_attribute(span, SpanAttributes.OUTPUT_VALUE, safe_dumps(payload))
safe_set_attribute(span, SpanAttributes.OUTPUT_MIME_TYPE, OpenInferenceMimeTypeValues.JSON.value)

View file

@ -10,6 +10,7 @@ Usage:
import os
import time
from collections.abc import AsyncIterable, Iterable
from typing import Final
from urllib.parse import urlparse
@ -84,7 +85,7 @@ def _mock_http_handler_post(
timeout=None,
stream=False,
files=None,
content=None,
content: str | bytes | Iterable[bytes] | AsyncIterable[bytes] | None = None,
logging_obj=None,
):
"""Monkey-patched HTTPHandler.post that intercepts Braintrust calls with endpoint-specific responses."""

View file

@ -54,7 +54,7 @@ USER_INVITED_EMAIL_TEMPLATE: Final = """
You were invited to use OpenAI Proxy API for team {team_name} <br /> <br />
<a href="{base_url}" style="display: inline-block; padding: 10px 20px; background-color: #87ceeb; color: #fff; text-decoration: none; border-radius: 20px;">Get Started here</a> <br /> <br />
<a href="{base_url}" style="display: inline-block; padding: 10px 20px; background-color: #87ceeb; color: #fff; text-decoration: none; border-radius: 20px;">Accept Invitation</a> <br /> <br />
If you have any questions, please send an email to {email_support_contact} <br /> <br />

View file

@ -131,7 +131,7 @@ USER_INVITATION_EMAIL_TEMPLATE: Final = """
</div>
<div class="btn-container">
<a href="{base_url}" class="btn">Accept Invitation</a>
<a href="{invitation_link}" class="btn">Accept Invitation</a>
</div>
<div class="quickstart">

View file

@ -9,6 +9,7 @@ Usage:
"""
import asyncio
from collections.abc import AsyncIterable, Iterable
from typing import Final
from litellm._logging import verbose_logger
@ -113,7 +114,7 @@ async def _mock_async_handler_delete(
headers=None,
timeout=None,
stream=False,
content=None,
content: str | bytes | Iterable[bytes] | AsyncIterable[bytes] | None = None,
):
"""Monkey-patched AsyncHTTPHandler.delete that intercepts GCS calls."""
# Only mock GCS API calls

View file

@ -11,7 +11,7 @@ import json
import os
import re
import traceback
from typing import Any, Final, Literal
from typing import Final, Literal
import httpx
@ -158,7 +158,7 @@ class GenericAPILogger(CustomBatchLogger):
"endpoint not set for GenericAPILogger, GENERIC_LOGGER_ENDPOINT not found in environment variables"
)
self.headers: dict = self._get_headers(headers)
self.headers: dict[str, str] = self._get_headers(headers)
self.endpoint: str = endpoint
self.event_types: list[API_EVENT_TYPES] | None = event_types
self.callback_name: str | None = callback_name
@ -248,18 +248,15 @@ class GenericAPILogger(CustomBatchLogger):
await asyncio.sleep(delay)
async def _post_with_retries(self, data: str) -> httpx.Response:
post_kwargs: Final[dict[str, Any]] = {
"url": self.endpoint,
"headers": self.headers,
"data": data,
}
if self.timeout is not None:
post_kwargs["timeout"] = self.timeout
total_attempts: Final = self.max_retries + 1
for attempt in range(total_attempts):
try:
return await self.async_httpx_client.post(**post_kwargs)
return await self.async_httpx_client.post(
url=self.endpoint,
headers=self.headers,
data=data,
timeout=self.timeout,
)
except Exception as e:
is_last_attempt = attempt == self.max_retries
should_retry = self._should_retry_exception(e)

View file

@ -8,6 +8,7 @@ making actual network calls.
import asyncio
import json
from collections.abc import AsyncIterable, Iterable
from dataclasses import dataclass
from datetime import timedelta
from typing import Final, cast
@ -140,7 +141,7 @@ def create_mock_client_factory(config: MockClientConfig):
stream=False,
logging_obj=None,
files=None,
content=None,
content: str | bytes | Iterable[bytes] | AsyncIterable[bytes] | None = None,
):
"""Monkey-patched AsyncHTTPHandler.post that intercepts API calls."""
if isinstance(url, str) and _is_mock_url(url):
@ -193,7 +194,7 @@ def create_mock_client_factory(config: MockClientConfig):
timeout=None,
stream=False,
files=None,
content=None,
content: str | bytes | Iterable[bytes] | AsyncIterable[bytes] | None = None,
logging_obj=None,
):
"""Monkey-patched HTTPHandler.post that intercepts API calls."""

View file

@ -1,7 +1,8 @@
import os
from collections.abc import Mapping
from dataclasses import dataclass, field
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, cast
from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
import litellm
from litellm._logging import verbose_logger
@ -37,9 +38,11 @@ from litellm.types.utils import (
# OpenTelemetry imports moved to individual functions to avoid import errors when not installed
if TYPE_CHECKING:
from opentelemetry.sdk.trace import TracerProvider as _SDKTracerProvider
from opentelemetry.sdk.trace.export import SpanExporter as _SpanExporter
from opentelemetry.trace import Context as _Context
from opentelemetry.trace import Span as _Span
from opentelemetry.trace import SpanKind as _SpanKind
from opentelemetry.trace import Tracer as _Tracer
from litellm.proxy._types import (
@ -61,6 +64,25 @@ else:
ManagementEndpointLoggingPayload = Any
Context = Any
class _StartSpanRequiredKwargs(TypedDict):
name: str
start_time: int
context: "Context | None"
class _StartSpanKwargs(_StartSpanRequiredKwargs, total=False):
kind: "_SpanKind"
class _UsageCompletionTokensView(TypedDict, total=False):
completion_tokens: int
class _ResponseWithUsageView(TypedDict, total=False):
usage: "_UsageCompletionTokensView | None"
LITELLM_TRACER_NAME: Final = os.getenv("OTEL_TRACER_NAME", "litellm")
LITELLM_METER_NAME: Final = os.getenv("LITELLM_METER_NAME", "litellm")
LITELLM_LOGGER_NAME: Final = os.getenv("LITELLM_LOGGER_NAME", "litellm")
@ -297,9 +319,9 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
config: OpenTelemetryConfig | None = None,
callback_name: str | None = None,
# injection points for testing
tracer_provider: Any | None = None,
logger_provider: Any | None = None,
meter_provider: Any | None = None,
tracer_provider: object | None = None,
logger_provider: object | None = None,
meter_provider: object | None = None,
**kwargs,
):
team_metadata_keys_override: Final = kwargs.pop("baggage_team_metadata_keys", None)
@ -325,7 +347,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
self.OTEL_EXPORTER = self.config.exporter
self.OTEL_ENDPOINT = self.config.endpoint
self.OTEL_HEADERS = self.config.headers
self._tracer_provider_cache: dict[str, Any] = {}
self._tracer_provider_cache: dict[str, _SDKTracerProvider] = {}
self._init_tracing(tracer_provider)
_debug_otel: Final = str(os.getenv("DEBUG_OTEL", "False")).lower()
@ -870,7 +892,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
def _emit_guardrail_spans_from_request_data(
self,
request_data: dict,
parent_span: Any | None,
parent_span: "Span | None",
) -> None:
"""Emit ``guardrail`` spans from the request's proxy-internal metadata bucket
(``standard_logging_guardrail_information``).
@ -896,7 +918,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
# kwargs["litellm_params"]["metadata"]["_otel_internal"]. Pass the
# SAME metadata dict the proxy populated so _handle_failure and
# this hook see the same dedupe markers.
kwargs: Final[dict[str, Any]] = {
kwargs: Final[dict[str, object]] = {
"litellm_params": {"metadata": metadata},
"standard_logging_object": {
"guardrail_information": guardrail_information,
@ -1257,13 +1279,13 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
response_obj,
start_time,
end_time,
context,
context: "Context | None",
):
from opentelemetry.trace import Status, StatusCode
otel_tracer: Final[Tracer] = self.get_tracer_to_use_for_request(kwargs)
span_kwargs: Final[dict[str, Any]] = {
span_kwargs: Final[_StartSpanKwargs] = {
"name": self._get_span_name(kwargs),
"start_time": self._to_ns(start_time),
"context": context,
@ -1454,7 +1476,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
) = _resolve_metric_attribute_filter(attributes)
self._metric_attr_filter_resolved = True
def _filter_metric_attributes(self, attrs: dict[str, Any]) -> dict[str, Any]:
def _filter_metric_attributes(self, attrs: dict[str, str]) -> dict[str, str]:
if not self._metric_attr_filter_resolved:
self._ensure_metric_attribute_filter()
if self._metric_attr_include is not None:
@ -1559,7 +1581,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
def _record_time_per_output_token_metric(
self,
kwargs: dict,
response_obj: Any | None,
response_obj: "_ResponseWithUsageView | None",
end_time: datetime,
duration_s: float,
common_attrs: dict,
@ -1775,10 +1797,10 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
@staticmethod
def _resolve_guardrail_context(
span: Any | None,
parent_span: Any | None,
fallback_ctx: Any | None,
) -> Any | None:
span: "Span | None",
parent_span: "Span | None",
fallback_ctx: "Context | None",
) -> "Context | None":
"""
Return a valid OTEL context for guardrail child spans so they are
never orphaned (Issue #5). Priority:
@ -1945,7 +1967,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
if should_create_primary_span:
# Span 1: Request sent to litellm SDK
otel_tracer: Final[Tracer] = self.get_tracer_to_use_for_request(kwargs)
span_kwargs: Final[dict[str, Any]] = {
span_kwargs: Final[_StartSpanKwargs] = {
"name": self._get_span_name(kwargs),
"start_time": self._to_ns(start_time),
"context": _parent_context,
@ -2131,10 +2153,10 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
@staticmethod
def _tool_calls_kv_pair(
tool_calls: list[ChatCompletionMessageToolCall],
) -> dict[str, Any]:
) -> dict[str, object]:
from litellm.proxy._types import SpanAttributes
kv_pairs: Final[dict[str, Any]] = {}
kv_pairs: Final[dict[str, object]] = {}
for idx, tool_call in enumerate(tool_calls):
_function = tool_call.get("function")
if not _function:
@ -2691,8 +2713,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
import json
try:
_raw_response = json.loads(_raw_response)
for param, val in _raw_response.items():
_parsed: Final[Mapping[str, object]] = json.loads(_raw_response)
for param, val in _parsed.items():
self.safe_set_attribute(
span=span,
key=f"llm.{custom_llm_provider}.{param}",
@ -2722,7 +2744,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
return int(dt * 1e9)
return int(dt.timestamp() * 1e9)
def _get_span_name(self, kwargs):
def _get_span_name(self, kwargs) -> str:
litellm_params: Final = kwargs.get("litellm_params", {})
metadata: Final = litellm_params.get("metadata") or {}
generation_name: Final = metadata.get("generation_name")

View file

@ -10,7 +10,7 @@ import asyncio
import math
import uuid
from collections.abc import AsyncIterator, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, cast
from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
import litellm
from litellm._logging import verbose_logger
@ -73,6 +73,23 @@ WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY: Final = "_websearch_interception_emit_native_b
WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY: Final = "websearch_native_blocks"
class _PlanMetadataView(TypedDict):
websearch_native_blocks: Sequence[Mapping[str, object]] | None
class _AgenticLoopParamsView(TypedDict):
agentic_loop_params: AgenticLoopParams
class _WebSearchSettingsView(TypedDict):
websearch_interception_params: WebSearchInterceptionConfig
class _SearchToolConfig(TypedDict, total=False):
search_tool_name: str
litellm_params: Mapping[str, object] | None
class WebSearchInterceptionLogger(CustomLogger):
"""
CustomLogger that intercepts WebSearch tool calls for models that don't
@ -394,7 +411,7 @@ class WebSearchInterceptionLogger(CustomLogger):
return tool.get("name")
@classmethod
def _sync_forced_tool_choice(cls, tool_choice: Any, converted_tools: list[dict[str, object]]) -> object:
def _sync_forced_tool_choice(cls, tool_choice: object, converted_tools: Sequence[Mapping[str, object]]) -> object:
"""Repoint a forced ``tool_choice`` at ``litellm_web_search`` when it
names a web-search tool that was just converted away.
@ -462,7 +479,7 @@ class WebSearchInterceptionLogger(CustomLogger):
kwargs[WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY] = True
# Convert native web search tools to LiteLLM standard
converted_tools: Final = []
converted_tools: Final[list[dict[str, object]]] = []
for tool in tools:
if is_web_search_tool(tool):
standard_tool = get_litellm_web_search_tool()
@ -833,7 +850,10 @@ class WebSearchInterceptionLogger(CustomLogger):
Anthropic-native clients (Claude Desktop, the Anthropic SDK) can
render citations / sources alongside the model's textual reply.
"""
native_blocks: Final = plan.metadata.get(WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY)
metadata_view: Final[_PlanMetadataView] = {
"websearch_native_blocks": plan.metadata.get(WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY)
}
native_blocks: Final = metadata_view["websearch_native_blocks"]
if not native_blocks:
return response
return self._inject_native_blocks(response, native_blocks)
@ -1278,8 +1298,10 @@ class WebSearchInterceptionLogger(CustomLogger):
kwargs_for_followup: Final = self._prepare_followup_kwargs(kwargs)
if logging_obj is not None:
agentic_params: Final[AgenticLoopParams] = logging_obj.model_call_details.get("agentic_loop_params", {})
full_model_name = agentic_params.get("model", model)
agentic_view: Final[_AgenticLoopParamsView] = {
"agentic_loop_params": logging_obj.model_call_details.get("agentic_loop_params", {})
}
full_model_name = agentic_view["agentic_loop_params"].get("model", model)
verbose_logger.debug(
"WebSearchInterception: Built anthropic request patch [call_id=%s model=%s messages=%d searches=%d]",
_call_id,
@ -1470,7 +1492,7 @@ class WebSearchInterceptionLogger(CustomLogger):
return None
def _select_search_tool_from_router(self, llm_router: object) -> dict[str, Any] | None:
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 [])
@ -1478,9 +1500,9 @@ class WebSearchInterceptionLogger(CustomLogger):
def _select_search_tool_from_list(
self,
search_tools: list[dict[str, Any]],
search_tools: list[_SearchToolConfig],
source: str,
) -> dict[str, Any] | None:
) -> "_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]
if matching_tools:
@ -1675,8 +1697,8 @@ class WebSearchInterceptionLogger(CustomLogger):
@staticmethod
def initialize_from_proxy_config(
litellm_settings: dict[str, Any],
callback_specific_params: dict[str, Any],
litellm_settings: Mapping[str, WebSearchInterceptionConfig],
callback_specific_params: Mapping[str, object],
) -> "WebSearchInterceptionLogger":
"""
Static method to initialize WebSearchInterceptionLogger from proxy config.
@ -1700,7 +1722,10 @@ class WebSearchInterceptionLogger(CustomLogger):
# Get websearch_interception_params from litellm_settings or callback_specific_params
websearch_params: WebSearchInterceptionConfig = {}
if "websearch_interception_params" in litellm_settings:
websearch_params = litellm_settings["websearch_interception_params"]
settings_view: Final[_WebSearchSettingsView] = {
"websearch_interception_params": litellm_settings["websearch_interception_params"]
}
websearch_params = settings_view["websearch_interception_params"]
elif "websearch_interception" in callback_specific_params and isinstance(
callback_specific_params["websearch_interception"], dict
):

View file

@ -8,6 +8,7 @@ This module has no dependencies on proxy code and can be safely imported at the
import json
import os
import time
from collections.abc import Mapping
from pathlib import Path
from typing import Final
@ -71,7 +72,7 @@ def get_litellm_gateway_api_key(
return token_data["key"]
def is_cli_token_fresh(token_data: dict, buffer_hours: float = 0.1) -> bool:
def is_cli_token_fresh(token_data: Mapping[str, object], buffer_hours: float = 0.1) -> bool:
"""Check whether a cached CLI token (as stored in token.json) is still
within its expiration window. Used by `lite auth print-token` to fail
fast, without a network round trip, once the cached token is past

View file

@ -10,7 +10,7 @@ import subprocess
import sys
import time
import traceback
from collections.abc import Callable
from collections.abc import Callable, Mapping, Sequence
from datetime import datetime as dt_object
from functools import lru_cache
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast
@ -25,7 +25,15 @@ from litellm import (
log_raw_request_response,
turn_off_message_logging,
)
from litellm._logging import _is_debugging_on, _redact_string, verbose_logger
from litellm._logging import (
_is_debugging_on,
_redact_string,
session_id_var,
set_session_id,
set_trace_id,
trace_id_var,
verbose_logger,
)
from litellm._uuid import uuid
from litellm.batches.batch_utils import _handle_completed_batch
from litellm.caching.caching import DualCache, InMemoryCache
@ -168,6 +176,9 @@ from .initialize_dynamic_callback_params import (
from .specialty_caches.dynamic_logging_cache import DynamicLoggingCache
if TYPE_CHECKING:
from mcp.types import EmbeddedResource, ImageContent, TextContent
from litellm.integrations.otel.logger import OpenTelemetryV2
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
try:
from litellm_enterprise.enterprise_callbacks.callback_controls import (
@ -203,14 +214,30 @@ except Exception as e:
PagerDutyAlerting = CustomLogger
EnterpriseCallbackControls = None
EnterpriseStandardLoggingPayloadSetupVAR = None
_in_memory_loggers: Final[list[Any]] = []
if TYPE_CHECKING:
from litellm.integrations.generic_api.generic_api_callback import (
GenericAPILogger as _GenericAPILoggerCls,
)
_STANDARD_LOGGING_METADATA_KEYS: Final[frozenset] = frozenset(StandardLoggingMetadata.__annotations__.keys())
_GENERIC_API_LOGGER_CLS: Final = _GenericAPILoggerCls
_RESEND_EMAIL_LOGGER_FACTORY: Final = CustomLogger
_SENDGRID_EMAIL_LOGGER_FACTORY: Final = CustomLogger
_SMTP_EMAIL_LOGGER_FACTORY: Final = CustomLogger
_PAGERDUTY_ALERTING_FACTORY: Final = CustomLogger
else:
_GENERIC_API_LOGGER_CLS: Final = GenericAPILogger
_RESEND_EMAIL_LOGGER_FACTORY: Final = ResendEmailLogger
_SENDGRID_EMAIL_LOGGER_FACTORY: Final = SendGridEmailLogger
_SMTP_EMAIL_LOGGER_FACTORY: Final = SMTPEmailLogger
_PAGERDUTY_ALERTING_FACTORY: Final = PagerDutyAlerting
_in_memory_loggers: Final[list[CustomLogger]] = []
_STANDARD_LOGGING_METADATA_KEYS: Final[frozenset[str]] = frozenset(StandardLoggingMetadata.__annotations__.keys())
### GLOBAL VARIABLES ###
# Cache custom pricing keys as frozenset for O(1) lookups instead of looping through 49 keys
_CUSTOM_PRICING_KEYS: Final[frozenset] = frozenset(CustomPricingLiteLLMParams.model_fields.keys())
_CUSTOM_PRICING_KEYS: Final[frozenset[str]] = frozenset(CustomPricingLiteLLMParams.model_fields.keys())
sentry_sdk_instance = None
capture_exception = None
@ -313,6 +340,7 @@ class Logging(LiteLLMLoggingBaseClass):
applied_guardrails: list[str] | None = None,
kwargs: dict | None = None,
log_raw_request_response: bool = False,
supports_correlation_logging: bool = True,
):
_input: Final[str | None] = messages # save original value of messages
if messages is not None:
@ -338,6 +366,36 @@ class Logging(LiteLLMLoggingBaseClass):
self.call_type = call_type
self.litellm_call_id = litellm_call_id
self.litellm_trace_id: str = litellm_trace_id if litellm_trace_id else str(uuid.uuid4())
# Capture the pre-call *value* (not a contextvars.Token) so restoration works
# even if this attempt's own logging ends up dispatched onto a different
# asyncio Task/context (e.g. via asyncio.create_task or the logging worker) -
# a Token can only be reset in the exact Context where it was created.
self._pre_call_trace_id: str = trace_id_var.get()
self._pre_call_session_id: str = session_id_var.get()
_sid: Final = kwargs.get("litellm_session_id") if kwargs else None
self.litellm_session_id: str = str(_sid) if _sid else ""
# supports_correlation_logging is False for calls originating from the
# sync client entry point (wrapper() in utils.py): a plain OS thread
# has no per-call context isolation the way an asyncio Task does, and
# a thread pool's worker threads are recycled across unrelated
# requests, so stamping trace_id/session_id there risks one request's
# ids leaking into a different, later request on the same thread. Sync
# support is deferred to a follow-up PR with its own safe-restore
# mechanism; async calls (the proxy's only call path) are unaffected.
if supports_correlation_logging:
set_trace_id(self.litellm_trace_id)
set_session_id(self.litellm_session_id)
# set_trace_id()/set_session_id() sanitize (strip control chars, bound
# length) before storing, so the contextvar's actual value can differ
# from self.litellm_trace_id/litellm_session_id. Capture what was
# really stored - _restore_correlation_context_if_unclaimed() must
# compare against this, not the raw ids, or a caller-supplied id
# containing control characters/oversized input would never match
# and cleanup would be skipped forever.
self._own_trace_id: str = trace_id_var.get()
self._own_session_id: str = session_id_var.get()
self.function_id = function_id
self.streaming_chunks: list[Any] = [] # for generating complete stream response
self.sync_streaming_chunks: list[Any] = [] # for generating complete stream response
@ -1246,7 +1304,9 @@ class Logging(LiteLLMLoggingBaseClass):
verbose_logger.exception("LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging %s", e)
return response_obj
def _parse_post_mcp_call_hook_response(self, response: MCPPostCallResponseObject | None) -> Any:
def _parse_post_mcp_call_hook_response(
self, response: MCPPostCallResponseObject | None
) -> "Sequence[TextContent | ImageContent | EmbeddedResource] | None":
"""
Parse the response from the post_mcp_tool_call_hook
@ -1690,7 +1750,7 @@ class Logging(LiteLLMLoggingBaseClass):
self.completion_start_time = completion_start_time
self.model_call_details["completion_start_time"] = self.completion_start_time
def normalize_logging_result(self, result: Any) -> Any:
def normalize_logging_result(self, result: Any) -> object:
"""
Some endpoints return a different type of result than what is expected by the logging system.
This function is used to normalize the result to the expected type.
@ -1726,7 +1786,7 @@ class Logging(LiteLLMLoggingBaseClass):
)
return logging_result
def _merge_hidden_params_from_response_into_metadata(self, logging_result: Any) -> None:
def _merge_hidden_params_from_response_into_metadata(self, logging_result: object) -> None:
"""
Copy response._hidden_params into litellm_params.metadata['hidden_params'].
@ -1787,7 +1847,9 @@ class Logging(LiteLLMLoggingBaseClass):
if (standard_logging_payload := self.model_call_details.get("standard_logging_object")) is not None:
emit_standard_logging_payload(standard_logging_payload)
def _build_standard_logging_payload(self, init_response_obj: Any, start_time: Any, end_time: Any) -> Any:
def _build_standard_logging_payload(
self, init_response_obj: object, start_time: Any, end_time: Any
) -> StandardLoggingPayload | None:
"""Build StandardLoggingPayload and accumulate its construction time."""
_start: Final = time.time()
payload: Final = get_standard_logging_object_payload(
@ -1908,7 +1970,7 @@ class Logging(LiteLLMLoggingBaseClass):
def _is_recognized_call_type_for_logging(
self,
logging_result: Any,
logging_result: object,
):
"""
Returns True if the call type is recognized for logging (eg. ModelResponse, ModelResponseStream, etc.)
@ -1992,7 +2054,67 @@ class Logging(LiteLLMLoggingBaseClass):
if complete_streaming_response is not None:
await self.async_success_handler(result=complete_streaming_response)
def success_handler(self, result=None, start_time=None, end_time=None, cache_hit=None, **kwargs):
def _restore_correlation_context(self) -> None:
"""Restore trace_id/session_id contextvars to their pre-call value.
Without this, a nested LiteLLM call sharing the same asyncio Task as an
outer request (e.g. a guardrail's own LLM-as-judge call, an MCP sampling
call) would leave the outer request's subsequent log lines stamped with
the nested call's trace_id/session_id instead of its own.
Uses a plain set() of the captured pre-call value rather than
contextvars.Token-based reset(), since this can end up called from a
different asyncio Task/context than __init__ ran in (e.g. the request
task's own wrapper() finally block, plus async_success_handler
dispatched separately via asyncio.create_task/the logging worker) -
reset() only works in the exact Context a Token was created in and
raises otherwise. Deliberately NOT idempotent/guarded: each distinct
Task that calls this needs its own restore to actually take effect in
that Task's view of the contextvars, so calling it multiple times
(once per Task involved in this attempt) is required, not just safe.
"""
set_trace_id(self._pre_call_trace_id)
set_session_id(self._pre_call_session_id)
def _restore_correlation_context_if_unclaimed(self) -> None:
"""Guarded variant for __del__-triggered cleanup only.
__del__ can fire arbitrarily late (delayed by cyclic GC, possibly
after the consuming Task/thread has already moved on to a different,
still-active call). Unconditionally restoring in that case would
stomp the active call's trace_id/session_id with this abandoned
stream's stale pre-call snapshot. Only restore if the contextvars
still hold the ids *this* call set - i.e. nothing has claimed them
since - so an unrelated active call is never overwritten.
"""
if trace_id_var.get() == self._own_trace_id and session_id_var.get() == self._own_session_id:
self._restore_correlation_context()
def success_handler(
self,
result: Any = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml)
start_time: datetime.datetime | None = None,
end_time: datetime.datetime | None = None,
cache_hit: bool | None = None,
**kwargs: Any, # kwargs-ok: forwarded to _success_handler_body
) -> None:
"""Restores trace_id/session_id contextvars once this attempt's own success
logging (including any nested calls its callbacks trigger) is fully done."""
try:
return self._success_handler_body(
result=result, start_time=start_time, end_time=end_time, cache_hit=cache_hit, **kwargs
)
finally:
self._restore_correlation_context()
def _success_handler_body(
self,
result: Any = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml)
start_time: datetime.datetime | None = None,
end_time: datetime.datetime | None = None,
cache_hit: bool | None = None,
**kwargs: Any, # kwargs-ok: forwarded from success_handler
) -> None:
verbose_logger.debug("Logging Details LiteLLM-Success Call: Cache_hit=%s", cache_hit)
if not self.should_run_logging(event_type="sync_success"): # prevent double logging
return
@ -2399,7 +2521,31 @@ class Logging(LiteLLMLoggingBaseClass):
e,
)
async def async_success_handler(self, result=None, start_time=None, end_time=None, cache_hit=None, **kwargs):
async def async_success_handler(
self,
result: Any = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml)
start_time: datetime.datetime | None = None,
end_time: datetime.datetime | None = None,
cache_hit: bool | None = None,
**kwargs: Any, # kwargs-ok: forwarded to _async_success_handler_body
) -> None:
"""Restores trace_id/session_id contextvars once this attempt's own success
logging (including any nested calls its callbacks trigger) is fully done."""
try:
return await self._async_success_handler_body(
result=result, start_time=start_time, end_time=end_time, cache_hit=cache_hit, **kwargs
)
finally:
self._restore_correlation_context()
async def _async_success_handler_body(
self,
result: Any = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml)
start_time: datetime.datetime | None = None,
end_time: datetime.datetime | None = None,
cache_hit: bool | None = None,
**kwargs: Any, # kwargs-ok: forwarded from async_success_handler
) -> None:
"""
Implementing async callbacks, to handle asyncio event loop issues when custom integrations need to use async functions.
"""
@ -2791,7 +2937,32 @@ class Logging(LiteLLMLoggingBaseClass):
kwargs=self.model_call_details,
)
def failure_handler(self, exception, traceback_exception, start_time=None, end_time=None):
def failure_handler(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime | None = None,
end_time: datetime.datetime | None = None,
) -> None:
"""Restores trace_id/session_id contextvars once this attempt's own failure
logging (including any nested calls its callbacks trigger) is fully done."""
try:
return self._failure_handler_body(
exception=exception,
traceback_exception=traceback_exception,
start_time=start_time,
end_time=end_time,
)
finally:
self._restore_correlation_context()
def _failure_handler_body(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime | None = None,
end_time: datetime.datetime | None = None,
) -> None:
verbose_logger.debug("Logging Details LiteLLM-Failure Call: %s", litellm.failure_callback)
if not self.should_run_logging(event_type="sync_failure"): # prevent double logging
return
@ -2960,7 +3131,32 @@ class Logging(LiteLLMLoggingBaseClass):
"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while failure logging %s", e
)
async def async_failure_handler(self, exception, traceback_exception, start_time=None, end_time=None):
async def async_failure_handler(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime | None = None,
end_time: datetime.datetime | None = None,
) -> None:
"""Restores trace_id/session_id contextvars once this attempt's own failure
logging (including any nested calls its callbacks trigger) is fully done."""
try:
return await self._async_failure_handler_body(
exception=exception,
traceback_exception=traceback_exception,
start_time=start_time,
end_time=end_time,
)
finally:
self._restore_correlation_context()
async def _async_failure_handler_body(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime | None = None,
end_time: datetime.datetime | None = None,
) -> None:
"""
Implementing async callbacks, to handle asyncio event loop issues when custom integrations need to use async functions.
"""
@ -4043,7 +4239,7 @@ def _init_custom_logger_compatible_class(
for callback in _in_memory_loggers:
if isinstance(callback, PagerDutyAlerting):
return callback
pagerduty_logger: Final = PagerDutyAlerting(**custom_logger_init_args)
pagerduty_logger: Final = _PAGERDUTY_ALERTING_FACTORY(**custom_logger_init_args)
_in_memory_loggers.append(pagerduty_logger)
return pagerduty_logger
elif logging_integration == "anthropic_cache_control_hook":
@ -4073,7 +4269,7 @@ def _init_custom_logger_compatible_class(
return _gcs_pubsub_logger
elif logging_integration == "generic_api":
for callback in _in_memory_loggers:
if isinstance(callback, GenericAPILogger):
if isinstance(callback, _GENERIC_API_LOGGER_CLS):
return callback
generic_api_logger: Final = GenericAPILogger()
_in_memory_loggers.append(generic_api_logger)
@ -4082,21 +4278,21 @@ def _init_custom_logger_compatible_class(
for callback in _in_memory_loggers:
if isinstance(callback, ResendEmailLogger):
return callback
resend_email_logger: Final = ResendEmailLogger()
resend_email_logger: Final = _RESEND_EMAIL_LOGGER_FACTORY()
_in_memory_loggers.append(resend_email_logger)
return resend_email_logger
elif logging_integration == "sendgrid_email":
for callback in _in_memory_loggers:
if isinstance(callback, SendGridEmailLogger):
return callback
sendgrid_email_logger: Final = SendGridEmailLogger()
sendgrid_email_logger: Final = _SENDGRID_EMAIL_LOGGER_FACTORY()
_in_memory_loggers.append(sendgrid_email_logger)
return sendgrid_email_logger
elif logging_integration == "smtp_email":
for callback in _in_memory_loggers:
if isinstance(callback, SMTPEmailLogger):
return callback
smtp_email_logger: Final = SMTPEmailLogger()
smtp_email_logger: Final = _SMTP_EMAIL_LOGGER_FACTORY()
_in_memory_loggers.append(smtp_email_logger)
return smtp_email_logger
elif logging_integration == "humanloop":
@ -4163,7 +4359,7 @@ def _init_custom_logger_compatible_class(
return None
def _maybe_construct_otel_v2(callback_name: str, _in_memory_loggers: list) -> Any | None:
def _maybe_construct_otel_v2(callback_name: str, _in_memory_loggers: list[CustomLogger]) -> "OpenTelemetryV2 | None":
"""If ``LITELLM_OTEL_V2`` is on, build (or reuse) a single ``OpenTelemetryV2``
instance configured via the preset for ``callback_name``.
@ -4194,7 +4390,7 @@ def _maybe_construct_otel_v2(callback_name: str, _in_memory_loggers: list) -> An
return v2_logger
def _maybe_auto_initialize_arize_phoenix(_in_memory_loggers: list) -> None:
def _maybe_auto_initialize_arize_phoenix(_in_memory_loggers: list[CustomLogger]) -> None:
"""
Auto-initialize ArizePhoenixLogger when Phoenix env vars are detected.
@ -4421,7 +4617,7 @@ def get_custom_logger_compatible_class(
return callback
elif logging_integration == "generic_api":
for callback in _in_memory_loggers:
if isinstance(callback, GenericAPILogger):
if isinstance(callback, _GENERIC_API_LOGGER_CLS):
return callback
elif logging_integration == "resend_email":
for callback in _in_memory_loggers:
@ -5061,33 +5257,61 @@ class StandardLoggingPayloadSetup:
return end_time_float - start_time_float
@staticmethod
def _get_standard_logging_payload_trace_id(
def get_standard_logging_payload_trace_id(
logging_obj: Logging,
litellm_params: dict,
litellm_params: Mapping[str, Any],
) -> str:
"""
Returns the `litellm_trace_id` for this request
This helps link sessions when multiple requests are made in a single session
Gated behind `litellm.request_correlation_in_logs`:
- Off (default): legacy behavior, preserved for backward compatibility -
`litellm_session_id` takes priority over `litellm_trace_id` since historically
this field doubled as the session-grouping field.
- On: `litellm_trace_id` takes priority - trace_id and session_id are independent,
see `get_standard_logging_payload_session_id` for session tracking.
"""
dynamic_litellm_session_id: Final = litellm_params.get("litellm_session_id")
dynamic_litellm_trace_id: Final = litellm_params.get("litellm_trace_id")
metadata: Final = litellm_params.get("metadata")
metadata_session_id: Final = metadata.get("session_id") if metadata else None
metadata_trace_id: Final = metadata.get("trace_id") if metadata else None
# Note: we recommend using `litellm_session_id` for session tracking
# `litellm_trace_id` is an internal litellm param
ordered_candidates: Final[tuple[Any, Any, Any, Any]] = (
(dynamic_litellm_trace_id, dynamic_litellm_session_id, metadata_trace_id, metadata_session_id)
if litellm.request_correlation_in_logs
else (dynamic_litellm_session_id, dynamic_litellm_trace_id, metadata_session_id, metadata_trace_id)
)
for candidate in ordered_candidates:
if candidate:
return str(candidate)
return logging_obj.litellm_trace_id
@staticmethod
def get_standard_logging_payload_session_id(
logging_obj: Logging,
litellm_params: Mapping[str, Any],
) -> str:
"""
Returns the end-user/conversation `litellm_session_id` for this request, independent of trace_id.
Only populated when `litellm.request_correlation_in_logs` is enabled - off by default
to avoid changing existing StandardLoggingPayload shape for callers who haven't opted in.
Unlike `get_standard_logging_payload_trace_id`, this never falls back to a generated
per-call trace id: it's empty when the caller never supplied a session id.
"""
if not litellm.request_correlation_in_logs:
return ""
dynamic_litellm_session_id: Final = litellm_params.get("litellm_session_id")
if dynamic_litellm_session_id:
return str(dynamic_litellm_session_id)
elif dynamic_litellm_trace_id:
return str(dynamic_litellm_trace_id)
# Fallback: use metadata.session_id or metadata.trace_id for call chaining
metadata: Final = litellm_params.get("metadata") or {}
metadata_session_id: Final = metadata.get("session_id")
metadata_trace_id: Final = metadata.get("trace_id")
metadata: Final = litellm_params.get("metadata")
metadata_session_id: Final = metadata.get("session_id") if metadata else None
if metadata_session_id:
return str(metadata_session_id)
if metadata_trace_id:
return str(metadata_trace_id)
return logging_obj.litellm_trace_id
return logging_obj.litellm_session_id
@staticmethod
def _get_user_agent_tags(proxy_server_request: dict) -> list[str] | None:
@ -5392,7 +5616,11 @@ def get_standard_logging_object_payload(
payload: Final[StandardLoggingPayload] = StandardLoggingPayload(
id=str(id),
litellm_call_id=kwargs.get("litellm_call_id") or litellm_params.get("litellm_call_id"),
trace_id=StandardLoggingPayloadSetup._get_standard_logging_payload_trace_id(
trace_id=StandardLoggingPayloadSetup.get_standard_logging_payload_trace_id(
logging_obj=logging_obj,
litellm_params=litellm_params,
),
session_id=StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
logging_obj=logging_obj,
litellm_params=litellm_params,
),

View file

@ -49,6 +49,12 @@ _SERVICE_TIER_TO_COST_KEY_SUFFIX: Final[Mapping[str, str]] = MappingProxyType(
}
)
_INCLUSIVE_THRESHOLD_PROVIDERS: Final = frozenset({"xai"})
def _uses_inclusive_token_thresholds(custom_llm_provider: str | None) -> bool:
return custom_llm_provider in _INCLUSIVE_THRESHOLD_PROVIDERS
def _get_token_detail_value(details: object, key: str) -> int | None:
if isinstance(details, dict):
@ -202,7 +208,11 @@ def _parse_above_token_threshold(key: str) -> float:
def _get_token_base_cost(
model_info: ModelInfo, usage: Usage, service_tier: str | None = None
model_info: ModelInfo,
usage: Usage,
service_tier: str | None = None,
*,
threshold_is_inclusive: bool = False,
) -> tuple[float, float, float, float, float]:
"""
Return prompt cost, completion cost, and cache costs for a given model and usage.
@ -210,6 +220,9 @@ def _get_token_base_cost(
If input_tokens > threshold and `input_cost_per_token_above_[x]k_tokens` or `input_cost_per_token_above_[x]_tokens` is set,
then we use the corresponding threshold cost for all token types.
`threshold_is_inclusive` switches that comparison to >=, for providers such as xAI
that bill the higher tier once the prompt reaches the threshold.
Returns:
Tuple[float, float, float, float] - (prompt_cost, completion_cost, cache_creation_cost, cache_read_cost)
"""
@ -262,7 +275,7 @@ def _get_token_base_cost(
# Handle both formats: _above_128k_tokens and _above_128_tokens
threshold_str = key.split("_above_")[1].split("_tokens")[0]
threshold = _parse_above_token_threshold(key)
if usage.prompt_tokens > threshold:
if usage.prompt_tokens > threshold or (threshold_is_inclusive and usage.prompt_tokens == threshold):
# Prefer a service_tier-specific above-threshold key when available,
# e.g. input_cost_per_token_priority_above_200k_tokens for Gemini
# ON_DEMAND_PRIORITY. Falls back to the standard key automatically
@ -777,7 +790,12 @@ def generic_cost_per_token(
cache_creation_cost,
cache_creation_cost_above_1hr,
cache_read_cost,
) = _get_token_base_cost(model_info=model_info, usage=usage, service_tier=service_tier)
) = _get_token_base_cost(
model_info=model_info,
usage=usage,
service_tier=service_tier,
threshold_is_inclusive=_uses_inclusive_token_thresholds(custom_llm_provider),
)
prompt_cost = _calculate_input_cost(
prompt_tokens_details=prompt_tokens_details,
@ -909,7 +927,12 @@ def get_token_type_cost_breakdown(
cache_creation_cost_rate,
cache_creation_cost_above_1hr_rate,
cache_read_cost_rate,
) = _get_token_base_cost(model_info=model_info, usage=usage, service_tier=service_tier)
) = _get_token_base_cost(
model_info=model_info,
usage=usage,
service_tier=service_tier,
threshold_is_inclusive=_uses_inclusive_token_thresholds(custom_llm_provider),
)
reasoning_tokens = (
_parse_completion_tokens_details(usage)["reasoning_tokens"]
@ -996,9 +1019,13 @@ def calculate_image_response_cost_from_usage(
input_tokens_details: Final = getattr(usage, "input_tokens_details", None)
prompt_tokens_details: PromptTokensDetailsWrapper | None = None
if input_tokens_details is not None:
# input_tokens_details may be a dict (e.g. OpenAI image edit responses)
# or an object; read it tolerantly like the output side below, so image
# input tokens are priced at input_cost_per_image_token instead of
# silently falling back to the text rate.
prompt_tokens_details = PromptTokensDetailsWrapper(
text_tokens=getattr(input_tokens_details, "text_tokens", None),
image_tokens=getattr(input_tokens_details, "image_tokens", None),
text_tokens=_get_token_detail_value(input_tokens_details, "text_tokens"),
image_tokens=_get_token_detail_value(input_tokens_details, "image_tokens"),
cached_tokens=0,
)

View file

@ -549,7 +549,7 @@ def _fetch_and_extract_template(
return chat_template, bos_token, eos_token
async def ahf_chat_template(model: str, messages: list, chat_template: Any | None = None):
async def ahf_chat_template(model: str, messages: list, chat_template: str | None = None):
"""HuggingFace chat template (async version)"""
from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import (
_aget_chat_template_file,
@ -576,7 +576,7 @@ async def ahf_chat_template(model: str, messages: list, chat_template: Any | Non
)
def hf_chat_template(model: str, messages: list, chat_template: Any | None = None):
def hf_chat_template(model: str, messages: list, chat_template: str | None = None):
"""HuggingFace chat template (sync version)"""
from litellm.litellm_core_utils.prompt_templates.huggingface_template_handler import (
_get_chat_template_file,
@ -1130,7 +1130,7 @@ def convert_to_azure_openai_messages(
def infer_protocol_value(
value: Any,
value: object,
) -> Literal[
"string_value",
"number_value",
@ -1702,7 +1702,9 @@ def convert_function_to_anthropic_tool_invoke(
_name: Final = get_attribute_or_key(function_call, "name") or ""
_arguments: Final = get_attribute_or_key(function_call, "arguments")
tool_input = parse_tool_call_arguments(_arguments, tool_name=_name, context="Anthropic function to tool invoke")
tool_input: Final = parse_tool_call_arguments(
_arguments, tool_name=_name, context="Anthropic function to tool invoke"
)
anthropic_tool_invoke: Final = [
AnthropicMessagesToolUseParam(
@ -1764,7 +1766,7 @@ def convert_to_anthropic_tool_invoke(
Fixes: https://github.com/BerriAI/litellm/issues/17737
"""
anthropic_tool_invoke: Final[list[AnthropicMessagesToolUseParam | dict[str, Any]]] = []
anthropic_tool_invoke: Final[list[AnthropicMessagesToolUseParam | dict[str, object]]] = []
for tool in tool_calls:
if not get_attribute_or_key(tool, "type") == "function":
@ -1785,7 +1787,7 @@ def convert_to_anthropic_tool_invoke(
# Server tool IDs start with "srvtoolu_"
if tool_id.startswith("srvtoolu_"):
# Create server_tool_use block instead of tool_use
_anthropic_server_tool_use: dict[str, Any] = {
_anthropic_server_tool_use: dict[str, object] = {
"type": "server_tool_use",
"id": tool_id,
"name": tool_name,
@ -2177,7 +2179,7 @@ def _is_orphaned_tool_result(
return False
def _declared_tool_call_ids(message: Mapping[str, Any]) -> frozenset[str]:
def _declared_tool_call_ids(message: Mapping[str, object]) -> frozenset[str]:
tool_calls: Final = message.get("tool_calls")
if not isinstance(tool_calls, list):
return frozenset()
@ -2186,7 +2188,7 @@ def _declared_tool_call_ids(message: Mapping[str, Any]) -> frozenset[str]:
)
def group_tool_exchanges(messages: Sequence[Mapping[str, Any]]) -> tuple[tuple[int, ...], ...]:
def group_tool_exchanges(messages: Sequence[Mapping[str, object]]) -> tuple[tuple[int, ...], ...]:
"""Group message indices into tool exchanges: an assistant row that made
tool calls, together with the tool rows answering the ids it declared.
@ -2204,7 +2206,7 @@ def group_tool_exchanges(messages: Sequence[Mapping[str, Any]]) -> tuple[tuple[i
return tuple(_iter_tool_exchange_groups(messages))
def _iter_tool_exchange_groups(messages: Sequence[Mapping[str, Any]]) -> Iterator[tuple[int, ...]]:
def _iter_tool_exchange_groups(messages: Sequence[Mapping[str, object]]) -> Iterator[tuple[int, ...]]:
index = 0
while index < len(messages):
declared = _declared_tool_call_ids(messages[index])
@ -2409,7 +2411,7 @@ def anthropic_messages_pt(
# Convert ChatCompletionImageUrlObject to dict if needed
image_url_value = m["image_url"]
if isinstance(image_url_value, str):
image_url_input: str | dict[str, Any] = image_url_value
image_url_input: str | dict[str, object] = image_url_value
else:
# ChatCompletionImageUrlObject or dict case - convert to dict
image_url_input = {
@ -3179,7 +3181,7 @@ def _load_image_from_url(image_url):
try:
# Send a GET request to the image URL
client: Final = HTTPHandler(concurrent_limit=1)
response: Final = safe_get(client, image_url)
response: Final[httpx.Response] = safe_get(client, image_url)
response.raise_for_status() # Raise an exception for HTTP errors
# Check the response's content type to ensure it is an image
@ -3382,7 +3384,7 @@ class BedrockImageProcessor:
@staticmethod
def _post_call_image_processing(response: httpx.Response, image_url: str = "") -> tuple[str, str]:
# Check the response's content type to ensure it is an image
content_type = response.headers.get("content-type")
content_type: str | None = response.headers.get("content-type")
# Use helper function to infer content type with fallback logic
content_type = infer_content_type_from_url_and_content(
@ -3406,7 +3408,7 @@ class BedrockImageProcessor:
params={"concurrent_limit": 1},
)
# Send a GET request to the image URL
response: Final = await async_safe_get(client, image_url)
response: Final[httpx.Response] = await async_safe_get(client, image_url)
response.raise_for_status() # Raise an exception for HTTP errors
return BedrockImageProcessor._post_call_image_processing(response, image_url)
@ -3419,7 +3421,7 @@ class BedrockImageProcessor:
try:
client: Final = HTTPHandler(concurrent_limit=1)
# Send a GET request to the image URL
response: Final = safe_get(client, image_url)
response: Final[httpx.Response] = safe_get(client, image_url)
response.raise_for_status() # Raise an exception for HTTP errors
return BedrockImageProcessor._post_call_image_processing(response, image_url)
@ -3967,6 +3969,36 @@ def _rename_duplicate_bedrock_document_names(
return contents
BEDROCK_DOCUMENT_PLACEHOLDER_TEXT: Final = "."
def _with_text_when_document_only(message: BedrockMessageBlock) -> BedrockMessageBlock:
blocks: Final = message["content"]
needs_text: Final = (
message["role"] == "user"
and any("document" in block for block in blocks)
and all("text" not in block for block in blocks)
)
if not needs_text:
return message
placeholder: Final = BedrockContentBlock(text=BEDROCK_DOCUMENT_PLACEHOLDER_TEXT)
cut: Final = len(blocks) - 1 if "cachePoint" in blocks[-1] else len(blocks)
return BedrockMessageBlock(role="user", content=[*blocks[:cut], placeholder, *blocks[cut:]])
def _ensure_document_messages_have_text(
contents: list[BedrockMessageBlock],
) -> list[BedrockMessageBlock]:
"""
Bedrock Converse rejects any user message that carries a document block
without a sibling text block ("A text block must be included when using
documents"), e.g. Claude Code sends the PDF as a document-only user turn.
Inject a placeholder text block, kept ahead of a trailing cachePoint so
the caller's cache boundary stays the final block.
"""
return [_with_text_when_document_only(message) for message in contents]
def _sort_bedrock_assistant_content_blocks(
blocks: list[BedrockContentBlock],
) -> list[BedrockContentBlock]:
@ -4535,7 +4567,7 @@ class BedrockConverseMessagesProcessor:
llm_provider=llm_provider,
)
return _rename_duplicate_bedrock_document_names(contents)
return _ensure_document_messages_have_text(_rename_duplicate_bedrock_document_names(contents))
@staticmethod
def translate_thinking_blocks_to_reasoning_content_blocks(
@ -4911,7 +4943,7 @@ def _bedrock_converse_messages_pt(
llm_provider=llm_provider,
)
return _rename_duplicate_bedrock_document_names(contents)
return _ensure_document_messages_have_text(_rename_duplicate_bedrock_document_names(contents))
def make_valid_bedrock_tool_name(input_tool_name: str) -> str:
@ -5328,10 +5360,10 @@ def get_attribute_or_key(tool_or_function, attribute, default=None):
class NormalizedToolCall(TypedDict):
id: str | None
name: str | None
arguments: dict[str, Any]
arguments: dict[str, object]
def _parse_tool_call_arguments(raw: Any, tool_name: str | None, context: str) -> dict[str, Any]:
def _parse_tool_call_arguments(raw: Any, tool_name: str | None, context: str) -> dict[str, object]:
# Anthropic's tool_use blocks already carry a parsed dict in "input";
# chat completions and the Responses API carry a JSON string that may be
# truncated by the model, so route those through the repair-aware parser.
@ -5352,12 +5384,12 @@ def _parse_tool_call_arguments(raw: Any, tool_name: str | None, context: str) ->
def _tool_calls_from_chat_completion_response(
response: Any, include_all_choices: bool = False
response: object, include_all_choices: bool = False
) -> list[NormalizedToolCall]:
choices: Final = get_attribute_or_key(response, "choices", None)
if not (isinstance(choices, list) and choices):
return []
tool_calls: Final[list[Any]] = []
tool_calls: Final[list[object]] = []
for choice in choices if include_all_choices else choices[:1]:
message = get_attribute_or_key(choice, "message", None)
choice_tool_calls = get_attribute_or_key(message, "tool_calls", None) if message else None
@ -5383,7 +5415,7 @@ def _tool_calls_from_chat_completion_response(
return result
def _tool_calls_from_responses_api_response(response: Any) -> list[NormalizedToolCall]:
def _tool_calls_from_responses_api_response(response: object) -> list[NormalizedToolCall]:
output: Final = get_attribute_or_key(response, "output", None)
if not isinstance(output, list):
return []
@ -5406,7 +5438,7 @@ def _tool_calls_from_responses_api_response(response: Any) -> list[NormalizedToo
return result
def _tool_calls_from_anthropic_messages_response(response: Any) -> list[NormalizedToolCall]:
def _tool_calls_from_anthropic_messages_response(response: object) -> list[NormalizedToolCall]:
content: Final = get_attribute_or_key(response, "content", None)
if not isinstance(content, list):
return []
@ -5425,7 +5457,7 @@ def _tool_calls_from_anthropic_messages_response(response: Any) -> list[Normaliz
return result
def get_tool_calls_from_response(response: Any, include_all_choices: bool = False) -> list[NormalizedToolCall]:
def get_tool_calls_from_response(response: object, include_all_choices: bool = False) -> list[NormalizedToolCall]:
"""
Extract tool/function calls from a response object into a normalized
``{"id", "name", "arguments"}`` shape, regardless of which API surface
@ -5456,7 +5488,7 @@ def get_tool_calls_from_response(response: Any, include_all_choices: bool = Fals
return []
def has_tool_with_name(tools: Any, tool_name: str) -> bool:
def has_tool_with_name(tools: object, tool_name: str) -> bool:
"""
Check whether a tools list (as sent to an LLM) includes a tool with the
given name, regardless of shape: OpenAI-style function tools
@ -5482,9 +5514,9 @@ def has_tool_with_name(tools: Any, tool_name: str) -> bool:
def resolve_structured_messages(
messages: list[dict[str, Any]] | None,
messages: list[dict[str, object]] | None,
request_kwargs: dict[str, Any],
) -> list[dict[str, Any]] | None:
) -> list[dict[str, object]] | None:
"""
Normalize a request's messages to OpenAI-spec chat-completions shape,
regardless of which API surface produced them (chat completions,

View file

@ -145,7 +145,7 @@ class ChunkProcessor:
if first_hidden_params.get("created_at"):
def _created_at(chunk: Any) -> int | float:
def _created_at(chunk: object) -> int | float:
if isinstance(chunk, dict):
params = chunk.get("_hidden_params", {})
else:
@ -158,7 +158,7 @@ class ChunkProcessor:
return chunks
def update_model_response_with_hidden_params(
self, model_response: ModelResponse, chunk: dict[str, Any] | None = None
self, model_response: ModelResponse, chunk: Mapping[str, dict[str, object]] | None = None
) -> ModelResponse:
if chunk is None:
return model_response
@ -176,7 +176,7 @@ class ChunkProcessor:
if not chunks:
return
model: Final = getattr(response, "model", None)
model: Final[str | None] = getattr(response, "model", None)
if not model:
return
@ -214,7 +214,7 @@ class ChunkProcessor:
)
@staticmethod
def _get_chunk_id(chunks: list[dict[str, Any]]) -> str:
def _get_chunk_id(chunks: Sequence[Mapping[str, str]]) -> str:
"""
Chunks:
[{"id": ""}, {"id": "1"}, {"id": "1"}]
@ -225,7 +225,7 @@ class ChunkProcessor:
return ""
@staticmethod
def _get_model_from_chunks(chunks: list[dict[str, Any]], first_chunk_model: str) -> str:
def _get_model_from_chunks(chunks: Sequence[Mapping[str, str]], first_chunk_model: str) -> str:
"""
Get the actual model from chunks, preferring a model that differs from the first chunk.

View file

@ -6,13 +6,14 @@ import logging
import threading
import time
import traceback
from collections.abc import AsyncIterator, Callable, Iterator
from collections.abc import AsyncIterator, Callable, Iterator, Mapping, Sequence
from dataclasses import dataclass
from typing import Any, Final, NoReturn, TypeVar, cast
from typing import Any, Final, NoReturn, Protocol, TypeVar, cast
import anyio
import httpx
from pydantic import BaseModel
from typing_extensions import NotRequired, TypedDict
import litellm
from litellm import verbose_logger
@ -54,7 +55,7 @@ _SYNC_ITER_EXHAUSTED: Final = object()
_GCHUNK_FIELDS: Final[frozenset] = frozenset(GChunk.__annotations__)
def _next_sync_or_exhausted(it: Any) -> Any:
def _next_sync_or_exhausted(it: Any) -> object:
"""
Call next(it) from a thread and return _SYNC_ITER_EXHAUSTED on StopIteration.
@ -68,7 +69,7 @@ def _next_sync_or_exhausted(it: Any) -> Any:
return _SYNC_ITER_EXHAUSTED
def is_async_iterable(obj: Any) -> bool:
def is_async_iterable(obj: object) -> bool:
"""
Check if an object is an async iterable (can be used with 'async for').
@ -81,7 +82,7 @@ def is_async_iterable(obj: Any) -> bool:
return isinstance(obj, collections.abc.AsyncIterable)
def print_verbose(print_statement):
def print_verbose(print_statement: object):
try:
if litellm.set_verbose:
print(print_statement) # noqa: T201
@ -96,18 +97,70 @@ class _ProviderChunkParsed:
@dataclass(frozen=True, slots=True)
class _ProviderChunkEarlyReturn:
value: Any
value: "ModelResponseStream | None"
_ProviderChunkResult = _ProviderChunkParsed | _ProviderChunkEarlyReturn
class _PredibaseStreamData(TypedDict):
token: NotRequired[Mapping[str, str]]
details: Mapping[str, str]
generated_text: str | None
error: str | None
class _Ai21StreamData(TypedDict):
completions: Sequence[Mapping[str, Mapping[str, str]]]
class _MaritalkStreamData(TypedDict):
answer: str
class _NlpCloudStreamData(TypedDict):
generated_text: str
class _AlephAlphaStreamData(TypedDict):
completions: Sequence[Mapping[str, str]]
class _AzureStreamChoice(TypedDict):
delta: Mapping[str, str] | None
finish_reason: str | None
class _AzureStreamData(TypedDict):
choices: Sequence[_AzureStreamChoice]
class _BasetenModelOutput(TypedDict):
data: NotRequired[Sequence[str]]
class _BasetenStreamData(TypedDict):
token: NotRequired[Mapping[str, str]]
model_output: NotRequired["_BasetenModelOutput | str"]
completion: NotRequired[object]
class _DeltaDumpDict(TypedDict):
role: NotRequired[str | None]
tool_calls: NotRequired[Sequence[Mapping[str, object]]]
class _TextCompletionChoiceLike(Protocol):
text: str
finish_reason: str | None
class CustomStreamWrapper:
def __init__(
self,
completion_stream,
model,
logging_obj: Any,
logging_obj: LiteLLMLoggingObject,
custom_llm_provider: str | None = None,
stream_options=None,
make_call: Callable | None = None,
@ -186,7 +239,7 @@ class CustomStreamWrapper:
# Snapshot assumes self._hidden_params is populated from litellm_params
# at init and never mutated during the stream. If that ever changes,
# this cache must be removed.
self._base_hidden_params: dict[str, Any] = {
self._base_hidden_params: dict[str, object] = {
**self._hidden_params,
"response_cost": None,
}
@ -213,7 +266,75 @@ class CustomStreamWrapper:
def __aiter__(self) -> AsyncIterator["ModelResponseStream"]:
return self
def _restore_consumer_correlation_context(self, *, guarded: bool = False) -> None:
"""Restore trace_id/session_id in the *consuming* thread/task/context.
wrapper_async() deliberately skips restoring correlation context when
it returns a stream, so log lines emitted while the caller iterates it
still carry this call's ids (see request_correlation_in_logs).
wrapper() (the sync path) never stamps anything in the first place -
see Logging.__init__'s supports_correlation_logging - so this method
is an inert no-op for sync-created streams, harmless to call anyway
since the class is shared between __next__ and __anext__.
But the terminal success/failure handlers this stream dispatches to
finish the job run on a *different* Task/thread (asyncio.create_task,
threading.Thread, or the shared executor) - restoring there fixes up
that detached context, not the one actually running the caller's
`for`/`async for` loop. Call this at every point control genuinely
returns to that consuming context: natural exhaustion (StopIteration/
StopAsyncIteration), a raised failure, or explicit aclose(). Never let
this raise - it must not break the caller's actual stream handling.
guarded=True (only __del__ uses this) skips the restore unless the
contextvars still hold the ids this stream's own call set, so a
delayed finalizer never overwrites a different, still-active call
that has since taken over the same Task/thread's context.
"""
try:
logging_obj: Final = getattr(self, "logging_obj", None)
if logging_obj is None:
return
method_name: Final = (
"_restore_correlation_context_if_unclaimed" if guarded else "_restore_correlation_context"
)
restore: Final = getattr(logging_obj, method_name, None)
if restore is not None:
restore()
except Exception as restore_error: # noqa: BLE001 # best-effort cleanup; must not raise into the caller
verbose_logger.debug("could not restore correlation context: %s", restore_error)
def __del__(self) -> None:
"""Best-effort correlation-context cleanup for an abandoned async stream.
Only meaningfully applies to streams created by wrapper_async(): it
leaves contextvars "open" across the caller's iteration, so if the
caller never fully consumes the stream - stops early, drops the
reference, cancels it - none of the exit points
_restore_consumer_correlation_context() is called from ever run. For a
sync stream (wrapper()), this is a no-op in practice: wrapper() never
stamps trace_id/session_id for sync calls in the first place (see
Logging.__init__'s supports_correlation_logging), so there is nothing
for this to clean up.
This is a best-effort fallback, not a guarantee: __del__ timing is
unpredictable (delayed by cyclic GC, not guaranteed at interpreter
shutdown, and may run on a different thread), so this can only reduce
how long the leak persists, not eliminate it. That's an acceptable
trade specifically because its blast radius is bounded to the one
asyncio Task this stream's own call ran in - each async call has its
own copy of the contextvars, and Tasks (unlike a thread pool's worker
threads) are never recycled across requests, so a delayed or missed
cleanup here can never misattribute a *different* request's logs.
guarded=True additionally ensures it never clobbers a different,
still-active call's context within that same Task if this fires late.
"""
self._restore_consumer_correlation_context(guarded=True)
async def aclose(self):
# Restore the consumer's outer context only after the underlying
# provider stream's own close (and its diagnostic logging below, if
# closing fails) completes - not before - so those log lines still
# carry this closing stream's own trace_id/session_id.
if self.completion_stream is not None:
stream_to_close: Final = self.completion_stream
self.completion_stream = None
@ -233,6 +354,7 @@ class CustomStreamWrapper:
"CustomStreamWrapper.aclose: error closing completion_stream: %s",
e,
)
self._restore_consumer_correlation_context()
def check_send_stream_usage(self, stream_options: dict | None):
return stream_options is not None and stream_options.get("include_usage", False) is True
@ -347,7 +469,7 @@ class CustomStreamWrapper:
finish_reason = ""
print_verbose(f"chunk: {chunk}")
if chunk.startswith("data:"):
data_json: Final = json.loads(chunk[5:])
data_json: Final[_PredibaseStreamData] = json.loads(chunk[5:])
print_verbose(f"data json: {data_json}")
if "token" in data_json and "text" in data_json["token"]:
text = data_json["token"]["text"]
@ -377,7 +499,7 @@ class CustomStreamWrapper:
def handle_ai21_chunk(self, chunk): # fake streaming
chunk = chunk.decode("utf-8")
data_json: Final = json.loads(chunk)
data_json: Final[_Ai21StreamData] = json.loads(chunk)
try:
text: Final = data_json["completions"][0]["data"]["text"]
is_finished: Final = True
@ -392,7 +514,7 @@ class CustomStreamWrapper:
def handle_maritalk_chunk(self, chunk): # fake streaming
chunk = chunk.decode("utf-8")
data_json: Final = json.loads(chunk)
data_json: Final[_MaritalkStreamData] = json.loads(chunk)
try:
text: Final = data_json["answer"]
is_finished: Final = True
@ -413,7 +535,7 @@ class CustomStreamWrapper:
if self.model and "dolphin" in self.model:
chunk = self.process_chunk(chunk=chunk)
else:
data_json: Final = json.loads(chunk)
data_json: Final[_NlpCloudStreamData] = json.loads(chunk)
chunk = data_json["generated_text"]
text = chunk
if "[DONE]" in text:
@ -430,7 +552,7 @@ class CustomStreamWrapper:
def handle_aleph_alpha_chunk(self, chunk):
chunk = chunk.decode("utf-8")
data_json: Final = json.loads(chunk)
data_json: Final[_AlephAlphaStreamData] = json.loads(chunk)
try:
text: Final = data_json["completions"][0]["completion"]
is_finished: Final = True
@ -458,7 +580,7 @@ class CustomStreamWrapper:
"finish_reason": finish_reason,
}
elif chunk.startswith("data:"):
data_json: Final = json.loads(chunk[5:]) # chunk.startswith("data:"):
data_json: Final[_AzureStreamData] = json.loads(chunk[5:]) # chunk.startswith("data:"):
try:
if len(data_json["choices"]) > 0:
delta: Final = data_json["choices"][0]["delta"]
@ -547,7 +669,7 @@ class CustomStreamWrapper:
text = ""
is_finished = False
finish_reason = None
choices: Final = getattr(chunk, "choices", [])
choices: Final[Sequence[_TextCompletionChoiceLike]] = getattr(chunk, "choices", [])
if len(choices) > 0:
text = choices[0].text
if choices[0].finish_reason is not None:
@ -568,7 +690,7 @@ class CustomStreamWrapper:
is_finished = False
finish_reason = None
usage = None
choices: Final = getattr(chunk, "choices", [])
choices: Final[Sequence[_TextCompletionChoiceLike]] = getattr(chunk, "choices", [])
if len(choices) > 0:
text = choices[0].text
if choices[0].finish_reason is not None:
@ -585,12 +707,12 @@ class CustomStreamWrapper:
except Exception as e:
raise e
def handle_baseten_chunk(self, chunk):
def handle_baseten_chunk(self, chunk) -> str:
try:
chunk = chunk.decode("utf-8")
if len(chunk) > 0:
if chunk.startswith("data:"):
data_json = json.loads(chunk[5:])
data_json: _BasetenStreamData = json.loads(chunk[5:])
if "token" in data_json and "text" in data_json["token"]:
return data_json["token"]["text"]
else:
@ -1256,13 +1378,14 @@ class CustomStreamWrapper:
if response_obj["is_finished"]:
self.received_finish_reason = response_obj["finish_reason"]
if "usage" in response_obj is not None:
_codestral_usage: Final[Usage] = response_obj["usage"]
setattr(
model_response,
"usage",
litellm.Usage(
prompt_tokens=response_obj["usage"].prompt_tokens,
completion_tokens=response_obj["usage"].completion_tokens,
total_tokens=response_obj["usage"].total_tokens,
prompt_tokens=_codestral_usage.prompt_tokens,
completion_tokens=_codestral_usage.completion_tokens,
total_tokens=_codestral_usage.total_tokens,
),
)
elif self.custom_llm_provider == "azure_text":
@ -1405,7 +1528,7 @@ class CustomStreamWrapper:
is None
):
t.function.arguments = ""
_json_delta: Final = delta.model_dump()
_json_delta: Final[_DeltaDumpDict] = delta.model_dump()
if "role" not in _json_delta or _json_delta["role"] is None:
_json_delta["role"] = "assistant" # mistral's api returns role as None
if "tool_calls" in _json_delta and isinstance(_json_delta["tool_calls"], list):
@ -1675,7 +1798,7 @@ class CustomStreamWrapper:
usage.cost, copy it into _hidden_params so litellm's cost
calculator uses it instead of a token-based estimate.
"""
_usage: Final = getattr(response, "usage", None)
_usage: Final[Usage | None] = getattr(response, "usage", None)
if _usage is not None and hasattr(_usage, "cost") and _usage.cost is not None:
if "additional_headers" not in response._hidden_params:
response._hidden_params["additional_headers"] = {}
@ -1839,6 +1962,7 @@ class CustomStreamWrapper:
if self.sent_stream_usage is False and self.send_stream_usage is True:
self.sent_stream_usage = True
return response
self._restore_consumer_correlation_context()
raise # Re-raise StopIteration
else:
self.sent_last_chunk = True
@ -1852,6 +1976,19 @@ class CustomStreamWrapper:
processed_chunk,
cache_hit,
) # log response
# Deliberately do NOT restore context here even though
# completion_stream is already exhausted: this chunk is still
# real data belonging to this call, and the caller's own
# (application-level) log statements processing it run
# immediately after this return, in this same synchronous
# frame - restoring first would make those lines carry the
# wrong ids, which is exactly what leaving context open during
# iteration is meant to prevent (see
# _restore_consumer_correlation_context's docstring). A caller
# that keeps iterating gets cleaned up on its next __next__()
# call (immediate StopIteration, handled above); one that
# stops right here relies on aclose() or the best-effort
# __del__ guard instead.
return processed_chunk
except Exception as e:
traceback_exception: Final = traceback.format_exc()
@ -1879,8 +2016,12 @@ class CustomStreamWrapper:
cache_hit = False
if self.custom_llm_provider is not None and self.custom_llm_provider == "cached_response":
cache_hit = True
self._check_max_streaming_duration()
try:
# Inside the try (not before it) so a raised litellm.Timeout flows
# through the same except Exception -> _handle_stream_fallback_error
# path as every other failure, restoring the consumer's correlation
# context - a check before the try would bypass that entirely.
self._check_max_streaming_duration()
if self.completion_stream is None:
await self.fetch_stream()
@ -2083,10 +2224,17 @@ class CustomStreamWrapper:
)
)
self._restore_consumer_correlation_context()
raise StopAsyncIteration # Re-raise StopIteration
else:
self.sent_last_chunk = True
processed_chunk: Final = self.finish_reason_handler()
# see sync __next__'s sibling branch: deliberately do NOT restore
# here - this chunk is still this call's own data, and restoring
# before returning it would corrupt the caller's own log
# statements processing it. A caller that keeps iterating gets
# cleaned up on the next __anext__() call; one that stops here
# relies on aclose() or the best-effort __del__ guard.
return processed_chunk
def _log_stream_failure_and_raise(self, e: Exception) -> NoReturn:
@ -2138,7 +2286,12 @@ class CustomStreamWrapper:
"""
from litellm.exceptions import MidStreamFallbackError
# Map to OpenAI exception format
# Map to OpenAI exception format. Some providers' mappers (e.g.
# _map_anthropic_exception, _map_aleph_alpha_exception) synchronously
# log a debug diagnostic (the raw status code) as part of mapping -
# restore the consumer's outer context only after this completes, so
# that diagnostic log line still carries the failing stream's own
# trace_id/session_id instead of the consumer's (or an empty one).
if isinstance(e, OpenAIError):
mapped_exception: Exception = e
else:
@ -2152,6 +2305,7 @@ class CustomStreamWrapper:
)
except Exception as mapping_error:
mapped_exception = mapping_error
self._restore_consumer_correlation_context()
def _normalize_status_code(exc: Exception) -> int | None:
"""Best-effort status_code extraction."""

View file

@ -14,6 +14,7 @@ Pattern Overview:
import json
from collections.abc import Mapping
from copy import deepcopy
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Final, cast
@ -110,14 +111,10 @@ EMPTY_EXTRACTED_INPUT: Final = ExtractedInput(scanned=(), images=())
class AnthropicMessagesHandler(BaseTranslation):
"""
Handler for processing Anthropic messages with guardrails.
"""Process Anthropic messages with guardrails.
This class provides methods to:
1. Process input messages (pre-call hook)
2. Process output responses (post-call hook)
Methods can be overridden to customize behavior for different message formats.
In-sequence system entries are untrusted client input. This handler scans and preserves
them through guardrail rewrites; downstream provider handling is out of scope.
"""
def __init__(self):
@ -331,16 +328,30 @@ class AnthropicMessagesHandler(BaseTranslation):
skip_tool: Final = effective_skip_tool_message_for_guardrail(guardrail_to_apply)
scan_only_tool_results: Final = effective_scan_only_tool_results_for_guardrail(guardrail_to_apply)
chat_completion_compatible_request: Final = self._translate_to_openai(data)
# Exclude only the trusted top-level prompt. In-sequence system entries are untrusted
# and must stay aligned with texts_to_check for positional masking. When the top-level
# prompt is included, the pre-existing count mismatch disables positional masking.
translation_source: Final = { # mutable-ok: API message payload
key: value for key, value in data.items() if key != "system"
}
chat_completion_compatible_request: Final = self._translate_to_openai(translation_source)
full_structured_messages: Final = cast(
list[AllMessageValues],
chat_completion_compatible_request.get("messages", []),
)
has_midturn_system_message: Final = any(
str(message.get("role") or "").lower() == "system" for message in full_structured_messages
)
hoisted_system_message: Final = None if skip_system else self._hoisted_top_level_system_message(data)
if hoisted_system_message is not None:
full_structured_messages.insert(0, hoisted_system_message)
# skip_system already excluded the trusted top-level prompt (it is simply not hoisted);
# in-sequence system entries are untrusted and always stay in scope.
scoped_message_indices: Final = scoped_structured_message_indices(
full_structured_messages,
scan_only_tool_results=scan_only_tool_results,
skip_system=skip_system,
skip_system=False,
skip_tool=skip_tool,
)
structured_messages: Final = [full_structured_messages[index] for index in scoped_message_indices]
@ -422,6 +433,8 @@ class AnthropicMessagesHandler(BaseTranslation):
scoped_indices=scoped_message_indices,
guardrailed_scoped=guardrailed_structured_messages,
),
hoisted_system_message=hoisted_system_message,
preserve_system_messages=has_midturn_system_message,
)
else:
# Step 3: Map guardrail responses back to original message structure
@ -435,36 +448,150 @@ class AnthropicMessagesHandler(BaseTranslation):
return data
@staticmethod
def _write_back_structured_messages(data: dict, structured_messages: list) -> None:
"""Convert compressed structured_messages back to Anthropic format and write to data.
def _hoisted_top_level_system_message(
self, data: dict
) -> AllMessageValues | None: # mutable-ok: API message payload
"""Return the system message produced by translating the top-level prompt."""
system: Final = data.get("system")
if not system:
return None
probe: Final = self._translate_to_openai(
{ # mutable-ok: API message payload
"model": data.get("model") or "",
"messages": [], # mutable-ok: API message payload
"system": system,
}
)
hoisted: Final = probe.get("messages") or [] # mutable-ok: API message payload
return hoisted[0] if hoisted else None
``anthropic_messages_pt`` merges every run of consecutive user/tool rows
into a single message, so a turn carrying only tool results and the user
turn that follows it come back fused, and the request the model sees no
longer has the boundaries the client sent. Converting a row at a time
would keep them apart but breaks tool pairing: an assistant row whose
tool results sit outside its own call reads as an orphaned tool call,
and under ``modify_params`` the sanitizer answers it with a synthetic
"tool execution skipped" result and drops the real one. Converting each
assistant row together with the tool rows that answer it, and every
other row on its own, satisfies both.
"""
@staticmethod
def _openai_system_message_to_anthropic(
message: dict[str, Any],
) -> dict[str, Any] | None: # mutable-ok: API message payload
"""Convert an OpenAI system message to the client's Anthropic-shaped entry."""
content: Final = message.get("content")
if isinstance(content, str):
return (
{"role": "system", "content": content} if content else None # mutable-ok: API message payload
) # mutable-ok: API message payload
if not isinstance(content, list):
return None
blocks: Final[list[dict[str, Any]]] = [] # mutable-ok: API message payload
for block in content:
if not isinstance(block, dict) or block.get("type") != "text":
continue
text = block.get("text")
if not isinstance(text, str) or not text:
continue
anthropic_block: dict[str, Any] = { # mutable-ok: API message payload
"type": "text",
"text": text,
} # mutable-ok: API message payload
cache_control = block.get("cache_control")
if cache_control:
anthropic_block["cache_control"] = deepcopy(cache_control)
blocks.append(anthropic_block)
return (
{"role": "system", "content": blocks} if blocks else None # mutable-ok: API message payload
) # mutable-ok: API message payload
@staticmethod
def _is_hoisted_top_level_system(message: object, hoisted_system_message: object) -> bool:
"""Match the hoisted prompt by identity, or by value after serialization."""
if hoisted_system_message is None:
return False
if message is hoisted_system_message:
return True
return (
isinstance(message, dict) and isinstance(hoisted_system_message, dict) and message == hoisted_system_message
)
@staticmethod
def _is_system(message: object) -> bool:
"""Whether the row is an in-sequence system message."""
return isinstance(message, dict) and str(message.get("role") or "").lower() == "system"
@staticmethod
def _defer_systems_inside_tool_exchanges(
structured_messages: list, # mutable-ok: API message payload
) -> list:
"""Hold a system row until the tool exchange around it completes so the call/result pair converts together."""
from litellm.litellm_core_utils.prompt_templates.factory import group_tool_exchanges
non_system_positions: Final[list[int]] = [
index
for index, message in enumerate(structured_messages)
if not AnthropicMessagesHandler._is_system(message)
]
exchange_end_for_start: Final[dict[int, int]] = {
non_system_positions[group[0]]: non_system_positions[group[-1]]
for group in group_tool_exchanges([structured_messages[index] for index in non_system_positions])
if len(group) > 1
}
ordered: Final[list] = [] # mutable-ok: API message payload
deferred_systems: Final[list] = [] # mutable-ok: API message payload
open_exchange_end = -1 # rebind-ok: advances to the enclosing exchange's last index
for index, message in enumerate(structured_messages):
if AnthropicMessagesHandler._is_system(message) and index < open_exchange_end:
deferred_systems.append(message)
continue
open_exchange_end = exchange_end_for_start.get(index, open_exchange_end)
ordered.append(message)
if index >= open_exchange_end and deferred_systems:
ordered.extend(deferred_systems)
deferred_systems.clear()
ordered.extend(deferred_systems)
return ordered
@staticmethod
def _write_back_structured_messages(
data: dict, # mutable-ok: API message payload
structured_messages: list, # mutable-ok: API message payload
hoisted_system_message: object = None,
preserve_system_messages: bool = False,
) -> None:
"""Write a guardrail's structured-message rewrite back without losing corrections."""
from litellm.litellm_core_utils.prompt_templates.factory import (
anthropic_messages_pt,
group_tool_exchanges,
)
_is_system: Final = AnthropicMessagesHandler._is_system
model: Final = str(data.get("model") or "")
non_system: Final = [m for m in structured_messages if m.get("role") != "system"]
groups: Final = tuple([non_system[index] for index in group] for group in group_tool_exchanges(non_system)) or (
non_system,
)
converted: Final = [
message
for group in groups
for message in anthropic_messages_pt(messages=group, model=model, llm_provider="anthropic")
]
converted: Final[list] = [] # mutable-ok: API message payload
def _convert_run(run: list) -> None: # mutable-ok: API message payload
for group in group_tool_exchanges(run):
converted.extend(
anthropic_messages_pt(
messages=[run[index] for index in group], # mutable-ok: API message payload
model=model,
llm_provider="anthropic",
)
)
ordered: Final = AnthropicMessagesHandler._defer_systems_inside_tool_exchanges(structured_messages)
run: Final[list] = [] # mutable-ok: API message payload
hoisted_dropped = False # rebind-ok: flips once the hoisted prompt is dropped
for message in ordered:
if not _is_system(message):
run.append(message)
continue
_convert_run(run)
run.clear()
if not hoisted_dropped and AnthropicMessagesHandler._is_hoisted_top_level_system(
message, hoisted_system_message
):
hoisted_dropped = True
continue
if preserve_system_messages:
anthropic_system = AnthropicMessagesHandler._openai_system_message_to_anthropic(message)
if anthropic_system is not None:
converted.append(anthropic_system)
_convert_run(run)
if not any(not _is_system(message) for message in converted):
converted.extend(anthropic_messages_pt(messages=[], model=model, llm_provider="anthropic"))
for msg in converted:
content = msg.get("content")
if isinstance(content, list):
@ -473,6 +600,31 @@ class AnthropicMessagesHandler(BaseTranslation):
block.pop("cache_control", None)
data["messages"] = converted
@staticmethod
def _extract_midturn_system_text(
message: dict[str, Any], # mutable-ok: API message payload
msg_idx: int,
) -> ExtractedInput:
"""Match the adapter's filtering so positional guardrail write-back stays aligned."""
content: Final = message.get("content")
if isinstance(content, str):
if not content:
return EMPTY_EXTRACTED_INPUT
return ExtractedInput(scanned=(ScannedText(content, MessageContentTarget(msg_idx)),), images=())
if not isinstance(content, list):
return EMPTY_EXTRACTED_INPUT
return ExtractedInput(
scanned=tuple(
ScannedText(text_str, ContentBlockTextTarget(msg_idx, content_idx))
for content_idx, content_item in enumerate(content)
if isinstance(content_item, dict)
and content_item.get("type") == "text"
and isinstance(text_str := content_item.get("text"), str)
and text_str
),
images=(),
)
def extract_request_tool_names(self, data: dict) -> list[str]:
"""Extract tool names from Anthropic messages request (tools[].name)."""
names: Final[list[str]] = []
@ -490,11 +642,17 @@ class AnthropicMessagesHandler(BaseTranslation):
skip_tool_message: bool = False,
scan_only_tool_results: bool = False,
) -> ExtractedInput:
"""Extract text content and images from a message.
In-sequence system entries are scanned even when ``skip_system_message`` is set:
that flag covers only the trusted top-level prompt, which never appears here.
"""
Extract text content and images from a message.
"""
role: Final = str(message.get("role") or "").lower()
if (skip_system_message and role == "system") or (skip_tool_message and role == "tool"):
role: Final = str(message.get("role") or "")
if role == "system":
if scan_only_tool_results:
return EMPTY_EXTRACTED_INPUT
return cls._extract_midturn_system_text(message=message, msg_idx=msg_idx)
if skip_tool_message and role.lower() == "tool":
return EMPTY_EXTRACTED_INPUT
content: Final = message.get("content", None)

View file

@ -58,6 +58,7 @@ from ..common_utils import AnthropicError, process_anthropic_headers
from .transformation import ANTHROPIC_TOOL_NAME_REVERSE_MAP_KEY, AnthropicConfig
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.llms.base_llm.chat.transformation import BaseConfig
@ -206,7 +207,7 @@ class AnthropicChatCompletion(BaseLLM):
client: AsyncHTTPHandler | None,
encoding,
api_key,
logging_obj,
logging_obj: "LiteLLMLoggingObj",
stream,
_is_function_call,
data: dict,
@ -324,7 +325,7 @@ class AnthropicChatCompletion(BaseLLM):
print_verbose: Callable,
encoding,
api_key,
logging_obj,
logging_obj: "LiteLLMLoggingObj",
optional_params: dict,
timeout: float | httpx.Timeout,
litellm_params: dict,

View file

@ -1,12 +1,13 @@
from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping
from typing import (
TYPE_CHECKING,
Any,
Final,
TypeAlias,
cast,
)
from typing_extensions import TypedDict
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.asyncify import run_async_function
@ -39,6 +40,11 @@ _AnthropicSystem: TypeAlias = "str | list[dict[str, object]] | None"
_ContextManagementSpec: TypeAlias = "dict[str, object] | list[dict[str, object]] | None"
class _CompletionKwargs(TypedDict, total=False, extra_items=object):
model: str
custom_llm_provider: str
def _messages_have_compaction_block(messages: _AnthropicMessages) -> bool:
"""Return True when any message carries a ``compaction`` content block."""
for msg in messages:
@ -312,7 +318,7 @@ ANTHROPIC_ADAPTER: Final = AnthropicAdapter()
class LiteLLMMessagesToCompletionTransformationHandler:
@staticmethod
def _route_openai_thinking_to_responses_api_if_needed(
completion_kwargs: dict[str, Any],
completion_kwargs: _CompletionKwargs,
*,
thinking: Mapping[str, object] | None,
) -> None:
@ -377,7 +383,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
@staticmethod
def _normalize_reasoning_effort(
completion_kwargs: dict[str, Any],
completion_kwargs: _CompletionKwargs,
) -> None:
"""
Normalize reasoning_effort values based on target model capabilities.
@ -393,7 +399,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
if reasoning_effort is None:
return
model: Final = cast(str, completion_kwargs.get("model", ""))
model: Final = completion_kwargs.get("model", "")
custom_llm_provider: Final = completion_kwargs.get("custom_llm_provider")
if isinstance(reasoning_effort, str):
@ -417,19 +423,19 @@ class LiteLLMMessagesToCompletionTransformationHandler:
max_tokens: int,
messages: _AnthropicMessages,
model: str,
metadata: dict | None = None,
metadata: dict[str, object] | None = None,
stop_sequences: list[str] | None = None,
stream: bool | None = False,
system: _AnthropicSystem = None,
temperature: float | None = None,
thinking: dict | None = None,
tool_choice: dict | None = None,
tools: list[dict] | None = None,
thinking: dict[str, object] | None = None,
tool_choice: dict[str, object] | None = None,
tools: list[dict[str, object]] | None = None,
top_k: int | None = None,
top_p: float | None = None,
output_format: dict | None = None,
output_format: dict[str, object] | None = None,
extra_kwargs: Mapping[str, object] | None = None,
) -> tuple[dict[str, Any], dict[str, str]]:
) -> tuple[_CompletionKwargs, dict[str, str]]:
"""Prepare kwargs for litellm.completion/acompletion.
Returns:
@ -486,7 +492,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
if openai_request is None:
raise ValueError("Failed to translate request to OpenAI format")
completion_kwargs: Final[dict[str, Any]] = dict(openai_request)
completion_kwargs: Final[_CompletionKwargs] = {**openai_request}
if stream:
completion_kwargs["stream"] = stream
@ -538,17 +544,17 @@ class LiteLLMMessagesToCompletionTransformationHandler:
max_tokens: int,
messages: _AnthropicMessages,
model: str,
metadata: dict | None = None,
metadata: dict[str, object] | None = None,
stop_sequences: list[str] | None = None,
stream: bool | None = False,
system: str | None = None,
temperature: float | None = None,
thinking: dict | None = None,
tool_choice: dict | None = None,
thinking: dict[str, object] | None = None,
tool_choice: dict[str, object] | None = None,
tools: list[dict[str, object]] | None = None,
top_k: int | None = None,
top_p: float | None = None,
output_format: dict | None = None,
output_format: dict[str, object] | None = None,
**kwargs,
) -> AnthropicMessagesResponse | AsyncIterator[bytes] | Iterator[bytes]:
"""Handle non-Anthropic models asynchronously using the adapter"""
@ -625,17 +631,17 @@ class LiteLLMMessagesToCompletionTransformationHandler:
max_tokens: int,
messages: _AnthropicMessages,
model: str,
metadata: dict | None = None,
metadata: dict[str, object] | None = None,
stop_sequences: list[str] | None = None,
stream: bool | None = False,
system: str | None = None,
temperature: float | None = None,
thinking: dict | None = None,
tool_choice: dict | None = None,
thinking: dict[str, object] | None = None,
tool_choice: dict[str, object] | None = None,
tools: list[dict[str, object]] | None = None,
top_k: int | None = None,
top_p: float | None = None,
output_format: dict | None = None,
output_format: dict[str, object] | None = None,
_is_async: bool = False,
**kwargs,
) -> (

View file

@ -74,12 +74,12 @@ from litellm.llms.anthropic.experimental_pass_through.context_management import
)
from litellm.types.llms.anthropic import (
ANTHROPIC_HOSTED_TOOLS,
AllAnthropicPassThroughMessageValues,
AllAnthropicToolsValues,
AnthopicMessagesAssistantMessageParam,
AnthropicFinishReason,
AnthropicMessagesRequest,
AnthropicMessagesSystemMessageParam,
AnthropicMessagesToolChoice,
AnthropicMessagesUserMessageParam,
AnthropicResponseContentBlockRedactedThinking,
AnthropicResponseContentBlockText,
AnthropicResponseContentBlockThinking,
@ -343,7 +343,7 @@ class LiteLLMAnthropicMessagesAdapter:
def translate_anthropic_messages_to_openai(
self,
messages: list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam],
messages: list[AllAnthropicPassThroughMessageValues],
model: str | None = None,
) -> list:
new_messages: Final[list[AllMessageValues]] = []
@ -351,6 +351,11 @@ class LiteLLMAnthropicMessagesAdapter:
user_message: ChatCompletionUserMessage | None = None
tool_message_list: list[ChatCompletionToolMessage] = []
new_user_content_list: list[ChatCompletionTextObject | ChatCompletionImageObject] = []
if m["role"] == "system":
system_message = self._translate_midturn_system_message_to_openai(m, model)
if system_message is not None:
new_messages.append(system_message)
continue
## USER MESSAGE ##
if m["role"] == "user":
## translate user message
@ -848,6 +853,29 @@ class LiteLLMAnthropicMessagesAdapter:
for def_schema in schema[key].values():
LiteLLMAnthropicMessagesAdapter._add_additional_properties_false(def_schema)
def _translate_midturn_system_message_to_openai(
self,
message: AnthropicMessagesSystemMessageParam,
model: str | None,
) -> ChatCompletionSystemMessage | None:
"""Translate an in-sequence system entry without changing its role or position."""
content: Final = message.get("content")
if isinstance(content, str):
return ChatCompletionSystemMessage(role="system", content=content) if content else None
if not isinstance(content, list):
return None
text_parts: Final[list[ChatCompletionTextObject]] = [] # mutable-ok: API message payload
for block in content:
if not isinstance(block, dict) or block.get("type") != "text": # pyright: ignore[reportUnnecessaryIsInstance] # untrusted client payload
continue
text = block.get("text")
if not text:
continue
text_obj = ChatCompletionTextObject(type="text", text=text)
self._add_cache_control_if_applicable(block, text_obj, model)
text_parts.append(text_obj)
return ChatCompletionSystemMessage(role="system", content=text_parts) if text_parts else None
def _add_system_message_to_messages(
self,
new_messages: list[AllMessageValues],
@ -976,6 +1004,17 @@ class LiteLLMAnthropicMessagesAdapter:
model: Final = new_kwargs.get("model", "")
if self.is_anthropic_claude_model(model) or self.is_bedrock_arn_model(model):
new_kwargs["thinking"] = thinking
# Adaptive thinking without its effort tier makes Bedrock Converse
# return zero reasoning blocks, so forward output_config (minus
# `format`, already translated to response_format) for Bedrock
# targets only: other bridged providers reject the raw param, and
# get_llm_provider strips the `bedrock/` prefix before this runs.
if model.startswith(("bedrock/", "converse/", "invoke/")) or self.is_bedrock_arn_model(model):
claude_output_config: Final = anthropic_message_request.get("output_config")
if isinstance(claude_output_config, dict):
effort_config: Final = {k: v for k, v in claude_output_config.items() if k != "format"}
if effort_config:
new_kwargs["output_config"] = effort_config # rebind-ok: out-param store like thinking above
return
reasoning_effort = self.translate_anthropic_thinking_to_reasoning_effort(cast(dict[str, Any], thinking))
@ -1049,8 +1088,8 @@ class LiteLLMAnthropicMessagesAdapter:
tool_name_mapping: dict[str, str] = {}
## CONVERT ANTHROPIC MESSAGES TO OPENAI
messages_list: Final[list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam]] = cast(
list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam],
messages_list: Final[list[AllAnthropicPassThroughMessageValues]] = cast(
list[AllAnthropicPassThroughMessageValues],
anthropic_message_request["messages"],
)
new_messages = self.translate_anthropic_messages_to_openai(

View file

@ -7,8 +7,8 @@ tool through a ``tool_use`` content block, and results are fed back as
``tool_result`` blocks in a user message.
"""
from collections.abc import AsyncIterator, Mapping, Sequence
from typing import Any, Final
from collections.abc import AsyncIterator, Awaitable, Callable, Iterator, Mapping, Sequence
from typing import Any, Final, NamedTuple
from litellm._logging import verbose_logger
from litellm.responses.mcp.request_context import MCPRequestContext
@ -24,14 +24,18 @@ from litellm.types.llms.anthropic_messages.anthropic_response import (
MAX_MCP_TOOL_USE_ITERATIONS: Final = 10
def _get_response_content(response: AnthropicMessagesResponse) -> Sequence[Mapping[str, Any]]:
class _AnthropicMessagesCall(NamedTuple):
fn: Callable[..., Awaitable[AnthropicMessagesResponse | Iterator[bytes] | AsyncIterator[object]]]
def _get_response_content(response: AnthropicMessagesResponse) -> Sequence[Mapping[str, object]]:
content: Final = response.get("content")
if not isinstance(content, list):
return ()
return tuple(block for block in content if isinstance(block, dict))
def _extract_tool_use_blocks(response: AnthropicMessagesResponse) -> Sequence[Mapping[str, Any]]:
def _extract_tool_use_blocks(response: AnthropicMessagesResponse) -> Sequence[Mapping[str, object]]:
"""Return the ``tool_use`` content blocks the model emitted."""
return tuple(block for block in _get_response_content(response) if block.get("type") == "tool_use")
@ -41,7 +45,7 @@ def _get_stop_reason(response: AnthropicMessagesResponse) -> str | None:
return stop_reason if isinstance(stop_reason, str) else None
def _build_tool_result_message(tool_results: Sequence[Mapping[str, Any]]) -> AnthropicMessagesUserMessageParam:
def _build_tool_result_message(tool_results: Sequence[Mapping[str, object]]) -> AnthropicMessagesUserMessageParam:
"""Turn executed tool results into the user message Anthropic expects."""
return AnthropicMessagesUserMessageParam(
role="user",
@ -58,11 +62,11 @@ def _build_tool_result_message(tool_results: Sequence[Mapping[str, Any]]) -> Ant
async def anthropic_messages_with_mcp(
max_tokens: int,
messages: Sequence[Mapping[str, Any]],
messages: Sequence[Mapping[str, object]],
model: str,
tools: Sequence[Mapping[str, Any]] | None = None,
tools: Sequence[Mapping[str, object]] | None = None,
**kwargs: Any, # kwargs-ok: forwarded verbatim to litellm.anthropic_messages, which owns the param contract
) -> AnthropicMessagesResponse | AsyncIterator[Any]:
) -> AnthropicMessagesResponse | Iterator[bytes] | AsyncIterator[object]:
"""
Expand litellm_proxy MCP references for `/v1/messages` and run the tool loop.
@ -81,7 +85,7 @@ async def anthropic_messages_with_mcp(
mcp_references, other_tools = LiteLLM_Proxy_MCP_Handler._parse_mcp_tools(tools)
if not mcp_references:
return await litellm.anthropic_messages(
return await _AnthropicMessagesCall(fn=litellm.anthropic_messages).fn(
max_tokens=max_tokens,
messages=list(messages),
model=model,
@ -114,7 +118,7 @@ async def anthropic_messages_with_mcp(
)
stream: Final = bool(kwargs.pop("stream", False))
base_call_args: Final[Mapping[str, Any]] = {
base_call_args: Final[Mapping[str, object]] = {
"max_tokens": max_tokens,
"model": model,
"tools": all_tools or None,
@ -123,10 +127,12 @@ async def anthropic_messages_with_mcp(
}
if not should_auto_execute:
return await litellm.anthropic_messages(messages=list(messages), stream=stream, **base_call_args)
return await _AnthropicMessagesCall(fn=litellm.anthropic_messages).fn(
messages=list(messages), stream=stream, **base_call_args
)
working_messages: Sequence[Mapping[str, Any]] = tuple(messages)
response: AnthropicMessagesResponse = await litellm.anthropic_messages(
working_messages: Sequence[Mapping[str, object]] = tuple(messages)
response: AnthropicMessagesResponse = await _AnthropicMessagesCall(fn=litellm.anthropic_messages).fn(
messages=list(working_messages), stream=False, **base_call_args
)
@ -161,7 +167,9 @@ async def anthropic_messages_with_mcp(
{"role": "assistant", "content": list(_get_response_content(response))},
_build_tool_result_message(tool_results),
)
response = await litellm.anthropic_messages(messages=list(working_messages), stream=False, **base_call_args)
response = await _AnthropicMessagesCall(fn=litellm.anthropic_messages).fn(
messages=list(working_messages), stream=False, **base_call_args
)
else:
verbose_logger.warning(
"MCP tool loop hit its %s iteration cap for model %s; returning the last response",

View file

@ -8,7 +8,12 @@ from collections.abc import AsyncIterator, Coroutine
from typing import Any, Final
import litellm
from litellm.types.llms.anthropic import AnthropicMessagesRequest
from litellm.types.llms.anthropic import (
AllAnthropicToolsValues,
AnthropicMessagesRequest,
AnthropicOutputConfig,
AnthropicOutputSchema,
)
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
@ -27,24 +32,24 @@ def _build_responses_kwargs(
model: str,
context_management: dict | None = None,
metadata: dict | None = None,
output_config: dict | None = None,
output_config: AnthropicOutputConfig | None = None,
stop_sequences: list[str] | None = None,
stream: bool | None = False,
system: str | None = None,
temperature: float | None = None,
thinking: dict | None = None,
tool_choice: dict | None = None,
tools: list[dict] | None = None,
tools: list[AllAnthropicToolsValues | dict] | None = None,
top_k: int | None = None,
top_p: float | None = None,
output_format: dict | None = None,
output_format: AnthropicOutputSchema | None = None,
extra_kwargs: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""
Build the kwargs dict to pass directly to litellm.responses() / litellm.aresponses().
"""
# Build a typed AnthropicMessagesRequest for the adapter
request_data: Final[dict[str, Any]] = {
request_data: Final[AnthropicMessagesRequest] = {
"model": model,
"messages": messages,
"max_tokens": max_tokens,
@ -128,19 +133,19 @@ class LiteLLMMessagesToResponsesAPIHandler:
model: str,
context_management: dict | None = None,
metadata: dict | None = None,
output_config: dict | None = None,
output_config: AnthropicOutputConfig | None = None,
stop_sequences: list[str] | None = None,
stream: bool | None = False,
system: str | None = None,
temperature: float | None = None,
thinking: dict | None = None,
tool_choice: dict | None = None,
tools: list[dict] | None = None,
tools: list[AllAnthropicToolsValues | dict] | None = None,
top_k: int | None = None,
top_p: float | None = None,
output_format: dict | None = None,
output_format: AnthropicOutputSchema | None = None,
**kwargs,
) -> AnthropicMessagesResponse | AsyncIterator:
) -> AnthropicMessagesResponse | AsyncIterator[bytes]:
responses_kwargs: Final = _build_responses_kwargs(
max_tokens=max_tokens,
messages=messages,
@ -179,23 +184,23 @@ class LiteLLMMessagesToResponsesAPIHandler:
model: str,
context_management: dict | None = None,
metadata: dict | None = None,
output_config: dict | None = None,
output_config: AnthropicOutputConfig | None = None,
stop_sequences: list[str] | None = None,
stream: bool | None = False,
system: str | None = None,
temperature: float | None = None,
thinking: dict | None = None,
tool_choice: dict | None = None,
tools: list[dict] | None = None,
tools: list[AllAnthropicToolsValues | dict] | None = None,
top_k: int | None = None,
top_p: float | None = None,
output_format: dict | None = None,
output_format: AnthropicOutputSchema | None = None,
_is_async: bool = False,
**kwargs,
) -> (
AnthropicMessagesResponse
| AsyncIterator[Any]
| Coroutine[Any, Any, AnthropicMessagesResponse | AsyncIterator[Any]]
| AsyncIterator[bytes]
| Coroutine[None, None, AnthropicMessagesResponse | AsyncIterator[bytes]]
):
if _is_async:
return LiteLLMMessagesToResponsesAPIHandler.async_anthropic_messages_handler(

View file

@ -6,6 +6,7 @@ path used for OpenAI and Azure models.
"""
import json
from collections.abc import Iterable
from typing import Any, Final, cast
from litellm.litellm_core_utils.reasoning_effort_utils import (
@ -15,15 +16,15 @@ from litellm.llms.anthropic.experimental_pass_through.utils import (
is_reasoning_auto_summary_enabled,
)
from litellm.types.llms.anthropic import (
AllAnthropicPassThroughMessageValues,
AllAnthropicToolsValues,
AnthopicMessagesAssistantMessageParam,
AnthropicFinishReason,
AnthropicMessagesRequest,
AnthropicMessagesToolChoice,
AnthropicMessagesUserMessageParam,
AnthropicResponseContentBlockText,
AnthropicResponseContentBlockThinking,
AnthropicResponseContentBlockToolUse,
AnthropicSystemMessageContent,
)
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
@ -72,14 +73,32 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
return source.get("url")
return None
@staticmethod
def _translate_midturn_system_content_to_responses(
content: str | Iterable[AnthropicSystemMessageContent],
) -> list[dict[str, str]]: # mutable-ok: API message payload
"""Convert in-sequence system content to Responses input-text parts."""
if isinstance(content, str):
return (
[{"type": "input_text", "text": content}] if content else [] # mutable-ok: API message payload
) # mutable-ok: API message payload
if not isinstance(content, list):
return [] # mutable-ok: API message payload
return [ # mutable-ok: API message payload
{"type": "input_text", "text": text} # mutable-ok: API message payload
for block in content
if isinstance(block, dict) and block.get("type") == "text" and (text := block.get("text")) # pyright: ignore[reportUnnecessaryIsInstance] # untrusted client payload
]
def translate_messages_to_responses_input(
self,
messages: list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam],
messages: list[AllAnthropicPassThroughMessageValues],
) -> list[dict[str, Any]]:
"""
Convert Anthropic messages list to Responses API `input` items.
Mapping:
system text -> message(role=system, input_text)
user text -> message(role=user, input_text)
user image -> message(role=user, input_image)
user tool_result -> function_call_output
@ -89,6 +108,18 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
input_items: Final[list[dict[str, Any]]] = []
for m in messages:
if m["role"] == "system":
system_parts = self._translate_midturn_system_content_to_responses(m.get("content"))
if system_parts:
input_items.append(
{ # mutable-ok: API message payload
"type": "message",
"role": "system",
"content": system_parts,
}
)
continue
role = m["role"]
content = m.get("content")
@ -300,7 +331,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
"""
model: Final[str] = anthropic_request["model"]
messages_list: Final = cast(
list[AnthropicMessagesUserMessageParam | AnthopicMessagesAssistantMessageParam],
list[AllAnthropicPassThroughMessageValues],
anthropic_request["messages"],
)

View file

@ -228,7 +228,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
litellm_params=litellm_params,
)
data = {"model": None, "messages": messages, **optional_params}
data: dict[str, object] = {"model": None, "messages": messages, **optional_params}
elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model=litellm_params.get("base_model") or model):
data = litellm.AzureOpenAIGPT5Config().transform_request(
model=model,
@ -482,12 +482,12 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
def streaming(
self,
logging_obj,
logging_obj: LiteLLMLoggingObj,
api_base: str,
api_key: str | None,
api_version: str,
dynamic_params: bool,
data: dict,
data: dict[str, object],
model: str,
timeout: Any,
max_retries: int,

View file

@ -5,10 +5,11 @@ Written separately to handle faking streaming for o1 and o3 models.
"""
from collections.abc import Callable
from typing import TYPE_CHECKING, Any, Optional
from typing import TYPE_CHECKING, Optional
import httpx
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.utils import ModelResponse
from ...openai.openai import OpenAIChatCompletion
@ -25,7 +26,7 @@ class AzureOpenAIO1ChatCompletion(BaseAzureLLM, OpenAIChatCompletion):
timeout: float | httpx.Timeout,
optional_params: dict,
litellm_params: dict,
logging_obj: Any,
logging_obj: LiteLLMLoggingObj,
model: str | None = None,
messages: list | None = None,
print_verbose: Callable | None = None,

View file

@ -3,6 +3,7 @@ from typing import Any, Final
from openai import AsyncAzureOpenAI, AzureOpenAI
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.prompt_templates.factory import prompt_factory
from litellm.utils import CustomStreamWrapper, ModelResponse, TextCompletionResponse
@ -39,9 +40,9 @@ class AzureTextCompletion(BaseAzureLLM):
azure_ad_token_provider: Callable | None,
print_verbose: Callable,
timeout,
logging_obj,
logging_obj: LiteLLMLoggingObj,
optional_params,
litellm_params,
litellm_params: dict[str, object],
logger_fn,
acompletion: bool = False,
headers: dict | None = None,
@ -246,7 +247,7 @@ class AzureTextCompletion(BaseAzureLLM):
def streaming(
self,
logging_obj,
logging_obj: LiteLLMLoggingObj,
api_base: str,
api_key: str | None,
api_version: str,
@ -299,7 +300,7 @@ class AzureTextCompletion(BaseAzureLLM):
async def async_streaming(
self,
logging_obj,
logging_obj: LiteLLMLoggingObj,
api_base: str,
api_key: str | None,
api_version: str,

View file

@ -9,6 +9,7 @@ from typing import Final
import httpx
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
@ -40,7 +41,7 @@ class AzureAnthropicChatCompletion(AnthropicChatCompletion):
print_verbose: Callable,
encoding,
api_key,
logging_obj,
logging_obj: LiteLLMLoggingObj,
optional_params: dict,
timeout: float | httpx.Timeout,
litellm_params: dict,

View file

@ -248,7 +248,7 @@ class AzureAIStudioConfig(OpenAIConfig):
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
encoding=encoding,
encoding=encoding if encoding is not None else None,
api_key=api_key,
json_mode=json_mode,
)

View file

@ -28,7 +28,11 @@ from litellm.types.llms.openai import (
from litellm.types.utils import LiteLLMBatch, LlmProviders
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import CommonBatchFilesUtils, resolve_s3_encryption_key_id
from ..common_utils import (
CommonBatchFilesUtils,
merge_bedrock_aws_request_params,
resolve_s3_encryption_key_id,
)
# Bedrock batch input files are uploaded as
# s3://bucket/litellm-bedrock-files-{model, ":" -> "-"}-{uuid4}.jsonl (see
@ -130,7 +134,8 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
Get the complete URL for Bedrock batch creation.
Bedrock batch jobs are created via the model invocation job API.
"""
aws_region_name: Final = self._get_aws_region_name(optional_params, model)
request_params: Final = merge_bedrock_aws_request_params(litellm_params, optional_params)
aws_region_name: Final = self._get_aws_region_name(request_params, model)
# Bedrock model invocation job endpoint
# Format: https://bedrock.{region}.amazonaws.com/model-invocation-job
@ -232,14 +237,15 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
# For Bedrock, we need to return a pre-signed request with AWS auth headers
# Use common utility for AWS signing
request_params: Final = merge_bedrock_aws_request_params(litellm_params, optional_params)
endpoint_url: Final = (
f"https://bedrock.{self._get_aws_region_name(optional_params, model)}.amazonaws.com/model-invocation-job"
f"https://bedrock.{self._get_aws_region_name(request_params, model)}.amazonaws.com/model-invocation-job"
)
signed_headers, signed_data = self.common_utils.sign_aws_request(
service_name="bedrock",
data=bedrock_request,
endpoint_url=endpoint_url,
optional_params=optional_params,
optional_params=request_params,
method="POST",
)
@ -387,11 +393,12 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
endpoint_url: Final = f"https://bedrock.{region}.amazonaws.com/model-invocation-job/{encoded_arn}"
# Use common utility for AWS signing
request_params: Final = merge_bedrock_aws_request_params(litellm_params, optional_params)
signed_headers, _ = self.common_utils.sign_aws_request(
service_name="bedrock",
data={}, # GET request has no body
endpoint_url=endpoint_url,
optional_params=optional_params,
optional_params=request_params,
method="GET",
)

View file

@ -89,7 +89,7 @@ class BedrockConverseLLM(BaseAWSLLM):
model_response: ModelResponse,
timeout: float | httpx.Timeout | None,
encoding,
logging_obj,
logging_obj: LiteLLMLoggingObject,
stream,
optional_params: dict,
litellm_params: dict,

View file

@ -80,6 +80,7 @@ from ..common_utils import (
bedrock_converse_supports_parallel_tool_use_config,
get_anthropic_beta_from_headers,
get_bedrock_tool_name,
is_bedrock_application_inference_profile_arn,
is_claude_4_5_on_bedrock,
normalize_bedrock_opus_output_config_effort,
)
@ -514,6 +515,7 @@ class AmazonConverseConfig(BaseConfig):
supported_params.append("tool_choice")
supported_params.append("thinking")
supported_params.append("reasoning_effort")
supported_params.append("output_config")
# For nova imported models, also add web_search_options
if "nova" in model.lower():
supported_params.append("web_search_options")
@ -564,6 +566,7 @@ class AmazonConverseConfig(BaseConfig):
):
supported_params.append("thinking")
supported_params.append("reasoning_effort")
supported_params.append("output_config")
if base_model.startswith("anthropic"):
supported_params.append("context_management")
@ -919,6 +922,10 @@ class AmazonConverseConfig(BaseConfig):
self._handle_reasoning_effort_parameter(
model=model, reasoning_effort=value, optional_params=optional_params
)
elif param == "output_config" and isinstance(value, dict):
mapped_output_config = dict(value)
normalize_bedrock_opus_output_config_effort(model=model, output_config=mapped_output_config)
optional_params["output_config"] = mapped_output_config # rebind-ok: out-param store like siblings
elif param == "context_management" and isinstance(value, (dict, list)):
self._map_context_management_param(value, optional_params)
if param == "requestMetadata":
@ -1312,7 +1319,12 @@ class AmazonConverseConfig(BaseConfig):
additional_request_params = filter_exceptions_from_params(additional_request_params)
if anthropic_output_config is not None and isinstance(anthropic_output_config, dict):
if base_model.startswith("anthropic"):
# Application inference profile ARNs hide the underlying model, so the
# effort ceiling and capability gates below cannot run; forward
# verbatim (like ``thinking``) and let Bedrock enforce.
if is_bedrock_application_inference_profile_arn(model):
additional_request_params["output_config"] = anthropic_output_config
elif base_model.startswith("anthropic"):
if litellm.drop_params is True and not AnthropicConfig._model_supports_effort_param(model, "bedrock"):
litellm.verbose_logger.warning(
DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING,

View file

@ -36,6 +36,44 @@ class BedrockError(BaseLLMException):
pass
_BEDROCK_AWS_AUTH_PARAMETER_KEYS: Final[tuple[str, ...]] = (
"aws_access_key_id",
"aws_secret_access_key",
"aws_session_token",
"aws_region_name",
"aws_session_name",
"aws_profile_name",
"aws_role_name",
"aws_web_identity_token",
"aws_sts_endpoint",
"aws_external_id",
)
def merge_bedrock_aws_request_params(
litellm_params: Mapping[str, Any],
optional_params: Mapping[str, Any],
) -> dict[str, Any]:
"""Merge deployment and request parameters without allowing auth escalation.
Deployment configuration is authoritative for AWS authentication. When a
deployment supplies static credentials, caller-supplied profile/role/token
selectors must not redirect signing to another identity available on the
server. Requests may still provide AWS credentials when the deployment has
no static credentials configured.
"""
request_params: Final = {**optional_params, **litellm_params} # mutable-ok: AWS helpers require a plain dict
has_static_deployment_credentials: Final = all(
isinstance(litellm_params.get(key), str) and bool(litellm_params.get(key))
for key in ("aws_access_key_id", "aws_secret_access_key", "aws_region_name")
)
if has_static_deployment_credentials:
for key in _BEDROCK_AWS_AUTH_PARAMETER_KEYS:
if key not in litellm_params:
request_params.pop(key, None)
return request_params
# Lazy import cache to avoid circular imports and performance impact
_get_model_info = None

View file

@ -54,7 +54,7 @@ from litellm.types.utils import ExtractedFileData, LlmProviders, SpecialEnums
from litellm.utils import get_llm_provider
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import BedrockError, resolve_s3_encryption_key_id
from ..common_utils import BedrockError, merge_bedrock_aws_request_params, resolve_s3_encryption_key_id
# litellm_params key used to hand the SigV4-signed GET headers from
# `transform_file_content_request` to `validate_environment` (the only hook
@ -285,6 +285,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
"""
Get the complete S3 URL for the file upload request
"""
request_params: Final = merge_bedrock_aws_request_params(litellm_params, optional_params)
bucket_name = litellm_params.get("s3_bucket_name") or os.getenv("AWS_S3_BUCKET_NAME")
if not bucket_name:
raise ValueError(
@ -293,7 +294,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
bucket_name, object_prefix = split_configured_cloud_bucket_name(bucket_name)
s3_region_name: Final = litellm_params.get("s3_region_name") or optional_params.get("s3_region_name")
aws_region_name: Final = s3_region_name or self._get_aws_region_name(optional_params, model)
aws_region_name: Final = s3_region_name or self._get_aws_region_name(request_params, model)
file_data: Final = data.get("file")
purpose: Final = data.get("purpose")
@ -309,7 +310,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
# S3 endpoint URL format
s3_endpoint_url: Final = (
optional_params.get("s3_endpoint_url") or f"https://s3.{aws_region_name}.amazonaws.com"
request_params.get("s3_endpoint_url") or f"https://s3.{aws_region_name}.amazonaws.com"
).rstrip("/")
return f"{s3_endpoint_url}/{bucket_name}/{encoded_object_name}"
@ -843,20 +844,23 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
)
# s3_region_name always wins for S3 operations (same priority as in
# get_complete_file_url above). Overwrite aws_region_name unconditionally
# so the SigV4 region matches the URL region, avoiding SignatureDoesNotMatch.
# get_complete_file_url above). Overwrite aws_region_name unconditionally,
# after the deployment-credential merge, so the SigV4 region matches the
# URL region, avoiding SignatureDoesNotMatch.
merged_params: Final = merge_bedrock_aws_request_params(litellm_params, optional_params)
s3_region_name: Final = litellm_params.get("s3_region_name") or optional_params.get("s3_region_name")
if s3_region_name:
optional_params = {**optional_params, "aws_region_name": s3_region_name}
request_params: Final = (
{**merged_params, "aws_region_name": s3_region_name} if s3_region_name else merged_params
)
# Sign the request and return a pre-signed request object
signed_headers, signed_body = self._sign_s3_request(
content=file_content,
api_base=api_base,
optional_params=optional_params,
optional_params=request_params,
s3_encryption_key_id=resolve_s3_encryption_key_id(
litellm_params=litellm_params,
optional_params=optional_params,
optional_params=request_params,
),
)

View file

@ -67,6 +67,8 @@ class AmazonAnthropicClaudeMessagesConfig(
DEFAULT_BEDROCK_ANTHROPIC_API_VERSION = "bedrock-2023-05-31"
WEBSEARCH_INTERCEPTION_DOCS_URL = "https://docs.litellm.ai/docs/integrations/websearch_interception"
@property
def custom_llm_provider(self) -> str | None:
return "bedrock"
@ -370,8 +372,9 @@ class AmazonAnthropicClaudeMessagesConfig(
"""
Check if the model supports tool search on Bedrock.
On Amazon Bedrock, server-side tool search is supported on Claude Opus 4.5
and Claude Sonnet 4.5 with the tool-search-tool-2025-10-19 beta header.
The model map's ``supports_tool_search`` flag is authoritative when
``model`` resolves to an entry that sets it; the name patterns below
cover ids the map cannot resolve (ARNs, unlisted regional variants).
Ref: https://platform.claude.com/docs/en/agents-and-tools/tool-use/tool-search-tool
@ -381,9 +384,12 @@ class AmazonAnthropicClaudeMessagesConfig(
Returns:
True if the model supports tool search on Bedrock
"""
catalog: Final = AnthropicModelInfo._get_provider_resolved_capability(model, "supports_tool_search", "bedrock")
if catalog is not None:
return catalog
model_lower: Final = model.lower()
# Supported models for tool search on Bedrock
supported_patterns: Final = [
# Opus 4.5
"opus-4.5",
@ -405,10 +411,16 @@ class AmazonAnthropicClaudeMessagesConfig(
"sonnet_4.6",
"sonnet-4-6",
"sonnet_4_6",
# NOTE: Opus 4.7 on Bedrock does not support server-side tool search
# as of launch (2026-04-16). Bedrock rejects the tool type with:
# "tool type 'tool_search_tool_..._20251119' is not supported for this model".
# Re-add the opus-4.7 patterns here once AWS announces support.
# Opus 4.7
"opus-4.7",
"opus_4.7",
"opus-4-7",
"opus_4_7",
# Haiku 4.5
"haiku-4.5",
"haiku_4.5",
"haiku-4-5",
"haiku_4_5",
]
return any(pattern in model_lower for pattern in supported_patterns)
@ -424,11 +436,10 @@ class AmazonAnthropicClaudeMessagesConfig(
"""
Adjust tool search beta header for Bedrock.
Bedrock requires a different beta header for tool search on Opus 4 models
when tool search is used without programmatic tool calling or input examples.
Note: On Amazon Bedrock, server-side tool search is only supported on Claude Opus 4
with the `tool-search-tool-2025-10-19` beta header.
Bedrock requires a different beta header for tool search than the
Anthropic API when tool search is used without programmatic tool
calling or input examples: `tool-search-tool-2025-10-19`, and only on
the models listed in `_supports_tool_search_on_bedrock`.
Ref: https://platform.claude.com/docs/en/agents-and-tools/tool-use/tool-search-tool
@ -572,6 +583,45 @@ class AmazonAnthropicClaudeMessagesConfig(
return filtered_betas
@staticmethod
def _reject_unsupported_web_search_tools(anthropic_messages_request: dict[str, object], model: str) -> None:
"""
Bedrock's Anthropic endpoints cannot execute Anthropic's server-side
``web_search_*`` tool; forwarding it returns an opaque
"The provided request is not valid" 400 from Bedrock. Fail fast with an
error that names the problem and the fix instead.
When web search interception is enabled
(``litellm_settings.callbacks: ["websearch_interception"]``), the tool
is converted to a regular function tool before this transform runs, so
this guard never fires.
"""
from litellm.integrations.websearch_interception.tools import (
is_anthropic_native_web_search_tool,
)
tools: Final = anthropic_messages_request.get("tools")
if not isinstance(tools, list):
return
web_search_tool: Final = next(
(t for t in tools if isinstance(t, dict) and is_anthropic_native_web_search_tool(t)),
None,
)
if web_search_tool is None:
return
raise litellm.BadRequestError(
message=(
f"Bedrock does not support Anthropic's server-side web search tool "
f"(tool type '{web_search_tool.get('type')}', model '{model}'). "
"To use web search with this model, enable LiteLLM's web search interception "
"so the proxy executes the search instead: "
f"{AmazonAnthropicClaudeMessagesConfig.WEBSEARCH_INTERCEPTION_DOCS_URL}. "
"Alternatively, remove the web_search tool from the request."
),
model=model,
llm_provider="bedrock",
)
def _strip_unsupported_bedrock_invoke_fields(
self,
anthropic_messages_request: dict,
@ -630,6 +680,8 @@ class AmazonAnthropicClaudeMessagesConfig(
############## BEDROCK Invoke SPECIFIC TRANSFORMATION ###
#########################################################
self._reject_unsupported_web_search_tools(anthropic_messages_request=anthropic_messages_request, model=model)
# 1. anthropic_version is required for all claude models
if "anthropic_version" not in anthropic_messages_request:
anthropic_messages_request["anthropic_version"] = self.DEFAULT_BEDROCK_ANTHROPIC_API_VERSION

View file

@ -195,7 +195,7 @@ class CodestralTextCompletion:
print_verbose: Callable,
encoding,
api_key: str,
logging_obj,
logging_obj: LiteLLMLogging,
optional_params: dict,
timeout: float | httpx.Timeout,
acompletion=None,
@ -383,7 +383,7 @@ class CodestralTextCompletion:
print_verbose: Callable,
encoding,
api_key,
logging_obj,
logging_obj: LiteLLMLogging,
data: dict,
timeout: float | httpx.Timeout,
optional_params=None,

View file

@ -221,7 +221,7 @@ class BaseLLMAIOHTTPHandler:
timeout=timeout,
stream=stream,
files=files,
content=content,
content=content if content is not None else None,
params=params,
)
except httpx.HTTPStatusError as e:

View file

@ -7,9 +7,9 @@ import ssl
import sys
import threading
import time
from collections.abc import Callable, Mapping
from collections.abc import AsyncIterable, Callable, Iterable, Mapping
from http.cookiejar import CookieJar, DefaultCookiePolicy
from typing import TYPE_CHECKING, Any, Final, Optional
from typing import TYPE_CHECKING, Any, Final, Optional, TypeAlias, TypedDict
import certifi
import httpx
@ -62,8 +62,23 @@ except Exception:
# https://docs.aiohttp.org/en/stable/client_reference.html#aiohttp.TCPConnector
_AIOHTTP_SUPPORTS_SOCKET_FACTORY: Final = "socket_factory" in inspect.signature(TCPConnector.__init__).parameters
_AddrInfo: TypeAlias = tuple[int | socket.AddressFamily, int | socket.SocketKind, int, str, tuple[object, ...]]
def _build_aiohttp_keepalive_socket_factory() -> Callable[[tuple[Any, ...]], socket.socket] | None:
_RequestContent: TypeAlias = str | bytes | Iterable[bytes] | AsyncIterable[bytes]
class _TCPConnectorKwargs(TypedDict, total=False):
local_addr: tuple[str, int] | None
ssl: "ssl.SSLContext | bool"
keepalive_timeout: float
ttl_dns_cache: int
enable_cleanup_closed: bool
limit: int
limit_per_host: int
socket_factory: Callable[[_AddrInfo], socket.socket]
def _build_aiohttp_keepalive_socket_factory() -> Callable[[_AddrInfo], socket.socket] | None:
"""
Build a socket_factory that enables SO_KEEPALIVE on aiohttp TCP sockets.
@ -78,7 +93,7 @@ def _build_aiohttp_keepalive_socket_factory() -> Callable[[tuple[Any, ...]], soc
if not AIOHTTP_SO_KEEPALIVE or not _AIOHTTP_SUPPORTS_SOCKET_FACTORY:
return None
def factory(addr_info: tuple[Any, ...]) -> socket.socket:
def factory(addr_info: _AddrInfo) -> socket.socket:
family, type_, proto = addr_info[0], addr_info[1], addr_info[2]
sock: Final = socket.socket(family=family, type=type_, proto=proto)
sock.setblocking(False)
@ -163,8 +178,8 @@ _STREAMING_ERROR_BODY_READ_EXECUTOR: Final = concurrent.futures.ThreadPoolExecut
def _prepare_request_data_and_content(
data: dict | str | bytes | None = None,
content: Any = None,
) -> tuple[dict | Mapping | None, Any]:
content: _RequestContent | None = None,
) -> tuple[dict | Mapping | None, _RequestContent | None]:
"""
Helper function to route data/content parameters correctly for httpx requests
@ -528,7 +543,7 @@ class AsyncHTTPHandler:
def __init__(
self,
timeout: float | httpx.Timeout | None = None,
event_hooks: Mapping[str, list[Callable[..., Any]]] | None = None,
event_hooks: Mapping[str, list[Callable[..., object]]] | None = None,
concurrent_limit=None, # Kept for backward compatibility, but ignored (no limits)
client_alias: str | None = None, # name for client in logs
ssl_verify: VerifyTypes | None = None,
@ -566,7 +581,7 @@ class AsyncHTTPHandler:
def create_client(
self,
timeout: float | httpx.Timeout | None,
event_hooks: Mapping[str, list[Callable[..., Any]]] | None,
event_hooks: Mapping[str, list[Callable[..., object]]] | None,
ssl_verify: VerifyTypes | None = None,
shared_session: Optional["ClientSession"] = None,
) -> httpx.AsyncClient:
@ -648,7 +663,7 @@ class AsyncHTTPHandler:
stream: bool = False,
logging_obj: LiteLLMLoggingObject | None = None,
files: RequestFiles | None = None,
content: Any = None,
content: _RequestContent | None = None,
):
start_time: Final = time.time()
try:
@ -691,7 +706,7 @@ class AsyncHTTPHandler:
end_time: Final = time.time()
time_delta: Final = round(end_time - start_time, 3)
headers = {}
error_response: Final = getattr(e, "response", None)
error_response: Final[httpx.Response | None] = getattr(e, "response", None)
if error_response is not None:
for key, value in error_response.headers.items():
headers[f"response_headers-{key}"] = value
@ -716,7 +731,7 @@ class AsyncHTTPHandler:
headers: dict | None = None,
timeout: float | httpx.Timeout | None = None,
stream: bool = False,
content: Any = None,
content: _RequestContent | None = None,
):
try:
if timeout is None:
@ -755,7 +770,7 @@ class AsyncHTTPHandler:
await new_client.aclose()
except httpx.TimeoutException as e:
headers = {}
error_response: Final = getattr(e, "response", None)
error_response: Final[httpx.Response | None] = getattr(e, "response", None)
if error_response is not None:
for key, value in error_response.headers.items():
headers[f"response_headers-{key}"] = value
@ -780,7 +795,7 @@ class AsyncHTTPHandler:
headers: dict | None = None,
timeout: float | httpx.Timeout | None = None,
stream: bool = False,
content: Any = None,
content: _RequestContent | None = None,
):
try:
if timeout is None:
@ -819,7 +834,7 @@ class AsyncHTTPHandler:
await new_client.aclose()
except httpx.TimeoutException as e:
headers = {}
error_response: Final = getattr(e, "response", None)
error_response: Final[httpx.Response | None] = getattr(e, "response", None)
if error_response is not None:
for key, value in error_response.headers.items():
headers[f"response_headers-{key}"] = value
@ -844,7 +859,7 @@ class AsyncHTTPHandler:
headers: dict | None = None,
timeout: float | httpx.Timeout | None = None,
stream: bool = False,
content: Any = None,
content: _RequestContent | None = None,
):
try:
if timeout is None:
@ -895,7 +910,7 @@ class AsyncHTTPHandler:
params: dict | None = None,
headers: dict | None = None,
stream: bool = False,
content: Any = None,
content: _RequestContent | None = None,
):
"""
Making POST request for a single connection client.
@ -993,7 +1008,7 @@ class AsyncHTTPHandler:
def _get_ssl_connector_kwargs(
ssl_verify: bool | None = None,
ssl_context: ssl.SSLContext | None = None,
) -> dict[str, Any]:
) -> _TCPConnectorKwargs:
"""
Helper method to get SSL connector initialization arguments for aiohttp TCPConnector.
@ -1004,7 +1019,7 @@ class AsyncHTTPHandler:
Returns:
Dict with appropriate SSL configuration for TCPConnector
"""
connector_kwargs: Final[dict[str, Any]] = {
connector_kwargs: Final[_TCPConnectorKwargs] = {
"local_addr": ("0.0.0.0", 0) if litellm.force_ipv4 else None,
}
@ -1054,7 +1069,7 @@ class AsyncHTTPHandler:
verbose_logger.debug("Creating AiohttpTransport...")
transport_connector_kwargs: Final = {
transport_connector_kwargs: Final[_TCPConnectorKwargs] = {
"keepalive_timeout": AIOHTTP_KEEPALIVE_TIMEOUT,
"ttl_dns_cache": AIOHTTP_TTL_DNS_CACHE,
**connector_kwargs,
@ -1212,7 +1227,7 @@ class HTTPHandler:
stream: bool = False,
timeout: float | httpx.Timeout | None = None,
files: dict | RequestFiles | None = None,
content: Any = None,
content: _RequestContent | None = None,
logging_obj: LiteLLMLoggingObject | None = None,
):
try:
@ -1265,7 +1280,7 @@ class HTTPHandler:
headers: dict | None = None,
stream: bool = False,
timeout: float | httpx.Timeout | None = None,
content: Any = None,
content: _RequestContent | None = None,
):
try:
# Prepare data/content parameters to prevent httpx DeprecationWarning (memory leak fix)
@ -1315,7 +1330,7 @@ class HTTPHandler:
headers: dict | None = None,
stream: bool = False,
timeout: float | httpx.Timeout | None = None,
content: Any = None,
content: _RequestContent | None = None,
):
try:
# Prepare data/content parameters to prevent httpx DeprecationWarning (memory leak fix)
@ -1364,7 +1379,7 @@ class HTTPHandler:
headers: dict | None = None,
timeout: float | httpx.Timeout | None = None,
stream: bool = False,
content: Any = None,
content: _RequestContent | None = None,
):
try:
# Prepare data/content parameters to prevent httpx DeprecationWarning (memory leak fix)

View file

@ -5,7 +5,8 @@ import ssl
from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping
from contextlib import asynccontextmanager
from functools import lru_cache
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeVar, Union, cast, get_type_hints
from types import ModuleType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypedDict, TypeVar, Union, cast, get_type_hints
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
import httpx
@ -148,6 +149,7 @@ from .http_handler import get_shared_realtime_ssl_context
if TYPE_CHECKING:
from aiohttp import ClientSession
from websockets.asyncio.client import ClientConnection
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
@ -176,6 +178,19 @@ else:
_ResponseT = TypeVar("_ResponseT")
class _DeleteRequestKwargs(TypedDict, total=False):
url: str
headers: dict[str, str]
timeout: float | httpx.Timeout | None
json: dict[str, object]
class _MediaUploadKwargs(TypedDict, total=False):
headers: dict[str, str]
content: Iterator[bytes] | AsyncIterator[bytes]
timeout: float | httpx.Timeout
def _google_genai_streaming_hidden_params(
*,
api_base: str,
@ -1413,7 +1428,7 @@ class BaseLLMHTTPHandler:
headers: dict[str, object] | None,
provider_config: BaseOCRConfig,
litellm_params: dict,
) -> tuple[dict[str, Any], str, dict[str, Any], None]:
) -> tuple[dict[str, object], str, dict[str, object], None]:
"""
Shared logic for preparing OCR requests.
Returns: (headers, complete_url, data, files)
@ -1479,7 +1494,7 @@ class BaseLLMHTTPHandler:
headers: dict[str, object] | None,
provider_config: BaseOCRConfig,
litellm_params: dict,
) -> tuple[dict[str, Any], str, dict[str, Any], None]:
) -> tuple[dict[str, object], str, dict[str, object], None]:
"""
Async version of _prepare_ocr_request for providers that need async transforms.
Returns: (headers, complete_url, data, files)
@ -2361,14 +2376,14 @@ class BaseLLMHTTPHandler:
model: str,
input: str | ResponseInputParam,
custom_llm_provider: str,
response_api_optional_request_params: dict[str, Any],
response_api_optional_request_params: dict[str, object],
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
) -> tuple[
str,
str | ResponseInputParam,
str,
dict[str, Any],
dict[str, object],
GenericLiteLLMParams,
]:
if not _has_pre_call_deployment_hook(logging_obj):
@ -2894,7 +2909,7 @@ class BaseLLMHTTPHandler:
},
)
delete_kwargs: Final[dict[str, Any]] = {
delete_kwargs: Final[_DeleteRequestKwargs] = {
"url": url,
"headers": headers,
"timeout": timeout,
@ -2984,7 +2999,7 @@ class BaseLLMHTTPHandler:
},
)
delete_kwargs: Final[dict[str, Any]] = {
delete_kwargs: Final[_DeleteRequestKwargs] = {
"url": url,
"headers": headers,
"timeout": timeout,
@ -3725,7 +3740,7 @@ class BaseLLMHTTPHandler:
timeout: float | httpx.Timeout | None,
) -> httpx.Response:
headers: Final = {**base_headers, "Content-Type": content_type}
kwargs: Final[dict[str, Any]] = {
kwargs: Final[_MediaUploadKwargs] = {
"headers": headers,
"content": self._iter_in_blocks(body_stream.iter_bytes(), self._MEDIA_UPLOAD_BLOCK_SIZE),
}
@ -3762,7 +3777,7 @@ class BaseLLMHTTPHandler:
break
yield cast(bytes, block)
kwargs: Final[dict[str, Any]] = {"headers": headers, "content": _abody()}
kwargs: Final[_MediaUploadKwargs] = {"headers": headers, "content": _abody()}
if timeout is not None:
kwargs["timeout"] = timeout
resp: Final = await client.client.post(url, **kwargs)
@ -5242,7 +5257,7 @@ class BaseLLMHTTPHandler:
def _wrap_responses_response_as_fake_stream(
self,
result: Any,
result: ResponsesAPIResponse,
model: str,
responses_api_provider_config: BaseResponsesAPIConfig,
logging_obj: "LiteLLMLoggingObj",
@ -5365,7 +5380,7 @@ class BaseLLMHTTPHandler:
async def _call_agentic_completion_hooks(
self,
response: Any,
response: object,
model: str,
messages: list[dict],
anthropic_messages_provider_config: "BaseAnthropicMessagesConfig",
@ -5536,7 +5551,7 @@ class BaseLLMHTTPHandler:
async def _call_agentic_chat_completion_hooks(
self,
response: Any,
response: ModelResponse,
model: str,
messages: list[dict],
optional_params: dict,
@ -5760,14 +5775,14 @@ class BaseLLMHTTPHandler:
@staticmethod
async def _open_realtime_backend_ws(
websockets_module: Any,
websockets_module: ModuleType,
url: str,
headers: dict,
ssl_context: Any,
ssl_context: bool | str | ssl.SSLContext,
*,
open_timeout: float = 8.0,
max_attempts: int = 3,
) -> Any:
) -> "ClientConnection":
"""Open the backend realtime websocket, retrying a hung open handshake.
The upstream Live handshake (e.g. Gemini Live) intermittently hangs on
@ -5826,7 +5841,6 @@ class BaseLLMHTTPHandler:
query_params: RealtimeQueryParams | None = None,
):
import websockets
from websockets.asyncio.client import ClientConnection
url: Final = provider_config.get_complete_url(api_base, model, api_key)
headers = provider_config.validate_environment(
@ -5844,12 +5858,12 @@ class BaseLLMHTTPHandler:
ssl_context.verify_mode = ssl.CERT_NONE
backend_ws: Final = await self._open_realtime_backend_ws(websockets, url, headers, ssl_context)
async with backend_ws:
_request_data: Final[dict[str, Any]] = {}
_request_data: Final[dict[str, object]] = {}
if litellm_metadata:
_request_data["litellm_metadata"] = litellm_metadata
realtime_streaming: Final = RealTimeStreaming(
websocket,
cast(ClientConnection, backend_ws),
backend_ws,
logging_obj,
provider_config,
model,
@ -6008,7 +6022,7 @@ class BaseLLMHTTPHandler:
)
else:
url = provider_config.get_complete_url(api_base=api_base, model=model or "", api_version=api_version)
headers: dict[str, Any] = provider_config.validate_environment(
headers: dict[str, object] = provider_config.validate_environment(
headers={}, model=model or "", api_key=api_key
)
else:
@ -6079,7 +6093,7 @@ class BaseLLMHTTPHandler:
if provider_config is not None:
url = provider_config.get_realtime_calls_url(api_base=api_base, model=model or "", api_version=api_version)
headers: dict[str, Any] = provider_config.get_realtime_calls_headers(ephemeral_key=openai_ephemeral_key)
headers: dict[str, object] = provider_config.get_realtime_calls_headers(ephemeral_key=openai_ephemeral_key)
else:
url = f"{api_base.rstrip('/')}/v1/realtime/calls"
headers = {
@ -6247,7 +6261,7 @@ class BaseLLMHTTPHandler:
yield backend
async with _backend_connection() as backend_ws:
_request_data: Final[dict[str, Any]] = {}
_request_data: Final[dict[str, object]] = {}
if litellm_metadata:
_request_data["litellm_metadata"] = litellm_metadata
@ -9444,7 +9458,7 @@ class BaseLLMHTTPHandler:
litellm_params=dict(litellm_params),
extra_body=extra_body,
)
all_optional_params: Final[dict[str, Any]] = dict(litellm_params)
all_optional_params: Final[dict[str, object]] = dict(litellm_params)
all_optional_params.update(vector_store_search_optional_params or {})
headers, signed_json_body = vector_store_provider_config.sign_request(
headers=headers,
@ -9540,7 +9554,7 @@ class BaseLLMHTTPHandler:
extra_body=extra_body,
)
all_optional_params: Final[dict[str, Any]] = dict(litellm_params)
all_optional_params: Final[dict[str, object]] = dict(litellm_params)
all_optional_params.update(vector_store_search_optional_params or {})
headers, signed_json_body = vector_store_provider_config.sign_request(
@ -9860,7 +9874,7 @@ class BaseLLMHTTPHandler:
url: Final = api_base
params: Final[dict[str, Any]] = {}
params: Final[dict[str, object]] = {}
if after is not None:
params["after"] = after
if before is not None:
@ -9938,7 +9952,7 @@ class BaseLLMHTTPHandler:
url: Final = api_base
params: Final[dict[str, Any]] = {}
params: Final[dict[str, object]] = {}
if after is not None:
params["after"] = after
if before is not None:

View file

@ -277,7 +277,7 @@ class GithubCopilotConfig(OpenAIConfig):
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: Any,
encoding: object,
api_key: str | None = None,
json_mode: bool | None = None,
) -> "ModelResponse":

View file

@ -10,7 +10,7 @@ implement the LiteLLM BaseConfig interface. Heavy-lifting lives in:
"""
import json
from collections.abc import AsyncIterator, Iterator
from collections.abc import AsyncIterator, Callable, Iterator
from typing import TYPE_CHECKING, Any, Final
import httpx
@ -713,8 +713,25 @@ class OCIChatConfig(BaseConfig):
class OCIStreamWrapper(CustomStreamWrapper):
"""Custom stream wrapper that dispatches OCI SSE chunks to the correct handler."""
def __init__(self, **kwargs: Any):
super().__init__(**kwargs)
def __init__(
self,
completion_stream: object,
model: str,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: str | None = None,
stream_options: object = None,
make_call: Callable[..., object] | None = None,
_response_headers: dict[str, object] | None = None,
) -> None:
super().__init__(
completion_stream=completion_stream,
model=model,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
stream_options=stream_options,
make_call=make_call,
_response_headers=_response_headers,
)
# Tracks whether any prior Cohere chunk in this stream has emitted
# tool calls. The Cohere handler uses this to decide whether the
# terminal consolidation chunk's tool calls are duplicates (suppress)

View file

@ -217,7 +217,7 @@ class OpenAITextCompletion(BaseLLM):
def streaming(
self,
logging_obj,
logging_obj: LiteLLMLoggingObj,
api_key: str,
data: dict,
headers: dict,
@ -274,7 +274,7 @@ class OpenAITextCompletion(BaseLLM):
async def async_streaming(
self,
logging_obj,
logging_obj: LiteLLMLoggingObj,
api_key: str,
data: dict,
headers: dict,

View file

@ -1,6 +1,6 @@
import time
import types
from collections.abc import AsyncIterator, Callable, Coroutine, Iterable, Iterator
from collections.abc import AsyncIterator, Callable, Coroutine, Iterable, Iterator, Mapping
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast
from urllib.parse import urlparse
@ -61,16 +61,17 @@ class MistralEmbeddingConfig:
def __init__(
self,
) -> None:
locals_: Final = locals().copy()
locals_: Final[Mapping[str, object]] = locals().copy()
for key, value in locals_.items():
if key != "self" and value is not None:
setattr(self.__class__, key, value)
@classmethod
def get_config(cls):
config_attrs: Final[Mapping[str, object]] = cls.__dict__
return {
k: v
for k, v in cls.__dict__.items()
for k, v in config_attrs.items()
if not k.startswith("__")
and not isinstance(
v,
@ -153,7 +154,7 @@ class OpenAIConfig(BaseConfig):
top_p: int | None = None,
response_format: dict | None = None,
) -> None:
locals_: Final = locals().copy()
locals_: Final[Mapping[str, object]] = locals().copy()
for key, value in locals_.items():
if key != "self" and value is not None:
setattr(self.__class__, key, value)
@ -261,7 +262,7 @@ class OpenAIConfig(BaseConfig):
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: Any,
encoding: object,
api_key: str | None = None,
json_mode: bool | None = None,
) -> ModelResponse:
@ -299,7 +300,7 @@ class OpenAIConfig(BaseConfig):
streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse,
sync_stream: bool,
json_mode: bool | None = False,
) -> Any:
) -> "OpenAIChatCompletionResponseIterator":
return OpenAIChatCompletionResponseIterator(
streaming_response=streaming_response,
sync_stream=sync_stream,
@ -478,14 +479,14 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
async def _call_agentic_completion_hooks_openai(
self,
response: Any,
response: object,
model: str,
messages: list[dict],
optional_params: dict,
logging_obj: LiteLLMLoggingObj,
stream: bool,
litellm_params: dict,
) -> Any | None:
) -> object | None:
"""
Call agentic completion hooks for all custom loggers (OpenAI Chat Completions API).
@ -536,7 +537,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
kwargs_with_provider["custom_llm_provider"] = custom_llm_provider
# For OpenAI Chat Completions, use the chat completion agentic loop method
agentic_response = await callback.async_run_chat_completion_agentic_loop(
agentic_response: object = await callback.async_run_chat_completion_agentic_loop(
tools=tool_calls,
model=model,
messages=messages,
@ -580,7 +581,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
timeout: float | httpx.Timeout,
optional_params: dict,
litellm_params: dict,
logging_obj: Any,
logging_obj: LiteLLMLoggingObj,
model: str | None = None,
messages: list | None = None,
print_verbose: Callable | None = None,
@ -1590,7 +1591,7 @@ class OpenAIFilesAPI(BaseLLM):
client: OpenAI | AsyncOpenAI | None = None,
_is_async: bool = False,
) -> OpenAI | AsyncOpenAI | None:
received_args: Final = locals()
received_args: Final[Mapping[str, object]] = locals()
openai_client: OpenAI | AsyncOpenAI | None = None
if client is None:
data: Final = {}
@ -1628,7 +1629,7 @@ class OpenAIFilesAPI(BaseLLM):
max_retries: int | None,
organization: str | None,
client: OpenAI | AsyncOpenAI | None = None,
) -> OpenAIFileObject | Coroutine[Any, Any, OpenAIFileObject]:
) -> OpenAIFileObject | Coroutine[None, None, OpenAIFileObject]:
openai_client: Final[OpenAI | AsyncOpenAI | None] = self.get_openai_client(
api_key=api_key,
api_base=api_base,
@ -1670,7 +1671,7 @@ class OpenAIFilesAPI(BaseLLM):
max_retries: int | None,
organization: str | None,
client: OpenAI | AsyncOpenAI | None = None,
) -> HttpxBinaryResponseContent | Coroutine[Any, Any, HttpxBinaryResponseContent]:
) -> HttpxBinaryResponseContent | Coroutine[None, None, HttpxBinaryResponseContent]:
openai_client: Final[OpenAI | AsyncOpenAI | None] = self.get_openai_client(
api_key=api_key,
api_base=api_base,
@ -1948,7 +1949,7 @@ class OpenAIBatchesAPI(BaseLLM):
client: OpenAI | AsyncOpenAI | None = None,
_is_async: bool = False,
) -> OpenAI | AsyncOpenAI | None:
received_args: Final = locals()
received_args: Final[Mapping[str, object]] = locals()
openai_client: OpenAI | AsyncOpenAI | None = None
if client is None:
data: Final = {}
@ -1986,7 +1987,7 @@ class OpenAIBatchesAPI(BaseLLM):
max_retries: int | None,
organization: str | None,
client: OpenAI | AsyncOpenAI | None = None,
) -> LiteLLMBatch | Coroutine[Any, Any, LiteLLMBatch]:
) -> LiteLLMBatch | Coroutine[None, None, LiteLLMBatch]:
openai_client: Final[OpenAI | AsyncOpenAI | None] = self.get_openai_client(
api_key=api_key,
api_base=api_base,
@ -2160,7 +2161,7 @@ class OpenAIAssistantsAPI(BaseLLM):
organization: str | None,
client: OpenAI | None = None,
) -> OpenAI:
received_args: Final = locals()
received_args: Final[Mapping[str, object]] = locals()
if client is None:
data: Final = {}
for k, v in received_args.items():
@ -2185,7 +2186,7 @@ class OpenAIAssistantsAPI(BaseLLM):
organization: str | None,
client: AsyncOpenAI | None = None,
) -> AsyncOpenAI:
received_args: Final = locals()
received_args: Final[Mapping[str, object]] = locals()
if client is None:
data: Final = {}
for k, v in received_args.items():
@ -2848,7 +2849,7 @@ class OpenAIAssistantsAPI(BaseLLM):
assistant_id: str,
additional_instructions: str | None,
instructions: str | None,
metadata: dict | None,
metadata: dict[str, str] | None,
model: str | None,
stream: bool | None,
tools: Iterable[AssistantToolParam] | None,
@ -2912,23 +2913,32 @@ class OpenAIAssistantsAPI(BaseLLM):
assistant_id: str,
additional_instructions: str | None,
instructions: str | None,
metadata: dict | None,
metadata: dict[str, str] | None,
model: str | None,
tools: Iterable[AssistantToolParam] | None,
event_handler: AssistantEventHandler | None,
) -> AssistantStreamManager[AssistantEventHandler]:
data: Final[dict[str, Any]] = {
"thread_id": thread_id,
"assistant_id": assistant_id,
"additional_instructions": additional_instructions,
"instructions": instructions,
"metadata": metadata,
"model": model,
"tools": tools,
}
runs_stream: Final = client.beta.threads.runs.stream
if event_handler is not None:
data["event_handler"] = event_handler
return client.beta.threads.runs.stream(**data)
return runs_stream(
thread_id=thread_id,
assistant_id=assistant_id,
additional_instructions=additional_instructions,
instructions=instructions,
metadata=metadata,
model=model,
tools=tools,
event_handler=event_handler,
)
return runs_stream(
thread_id=thread_id,
assistant_id=assistant_id,
additional_instructions=additional_instructions,
instructions=instructions,
metadata=metadata,
model=model,
tools=tools,
)
# fmt: off
@ -2984,7 +2994,7 @@ class OpenAIAssistantsAPI(BaseLLM):
assistant_id: str,
additional_instructions: str | None,
instructions: str | None,
metadata: dict | None,
metadata: dict[str, str] | None,
model: str | None,
stream: bool | None,
tools: Iterable[AssistantToolParam] | None,

View file

@ -12,6 +12,7 @@ import httpx
import litellm
from litellm import LlmProviders
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.bedrock.chat.invoke_handler import MockResponseIterator
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.databricks.streaming_utils import ModelResponseIterator
@ -112,7 +113,7 @@ class OpenAILikeChatHandler(OpenAILikeBase):
print_verbose: Callable,
encoding,
api_key,
logging_obj,
logging_obj: LiteLLMLoggingObj,
stream,
data: dict,
optional_params=None,
@ -214,7 +215,7 @@ class OpenAILikeChatHandler(OpenAILikeBase):
print_verbose: Callable,
encoding,
api_key: str | None,
logging_obj,
logging_obj: LiteLLMLoggingObj,
optional_params: dict,
acompletion=None,
litellm_params: dict = {},

View file

@ -9,6 +9,7 @@ from typing import Final
import httpx
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_async_httpx_client,
@ -59,7 +60,7 @@ class PredibaseChatCompletion:
print_verbose: Callable,
encoding,
api_key: str,
logging_obj,
logging_obj: LiteLLMLoggingObj,
optional_params: dict,
litellm_params: dict,
tenant_id: str,
@ -250,7 +251,7 @@ class PredibaseChatCompletion:
print_verbose: Callable,
encoding,
api_key,
logging_obj,
logging_obj: LiteLLMLoggingObj,
data: dict,
timeout: float | httpx.Timeout,
optional_params=None,

View file

@ -6,6 +6,7 @@ from typing import Final
import litellm
from litellm.constants import REPLICATE_POLLING_DELAY_SECONDS
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
@ -128,7 +129,7 @@ def completion(
print_verbose: Callable,
optional_params: dict,
litellm_params: dict,
logging_obj,
logging_obj: LiteLLMLoggingObj,
api_key,
encoding,
custom_prompt_dict={},
@ -246,7 +247,7 @@ async def async_completion(
input_data,
api_key,
api_base,
logging_obj,
logging_obj: LiteLLMLoggingObj,
print_verbose,
headers: dict,
) -> ModelResponse | CustomStreamWrapper:

View file

@ -8,6 +8,7 @@ import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.asyncify import asyncify
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
@ -138,7 +139,7 @@ class SagemakerLLM(BaseAWSLLM):
model_response: ModelResponse,
print_verbose: Callable,
encoding,
logging_obj,
logging_obj: LiteLLMLoggingObj,
optional_params: dict,
litellm_params: dict,
timeout: float | httpx.Timeout | None = None,
@ -431,17 +432,18 @@ class SagemakerLLM(BaseAWSLLM):
if not prepared_request.body:
raise ValueError("Prepared request body is empty")
stream_logging_obj: Final[LiteLLMLoggingObj] = logging_obj
completion_stream: Final = await self.make_async_call(
api_base=prepared_request.url,
headers=prepared_request.headers,
data=cast(str, prepared_request.body),
logging_obj=logging_obj,
logging_obj=stream_logging_obj,
)
streaming_response: Final = CustomStreamWrapper(
completion_stream=completion_stream,
model=model,
custom_llm_provider="sagemaker",
logging_obj=logging_obj,
logging_obj=stream_logging_obj,
)
# LOGGING

View file

@ -3,7 +3,7 @@
## Initial implementation - covers gemini + image gen calls
import json
import time
from collections.abc import Callable, Mapping
from collections.abc import Callable, Mapping, Sequence
from copy import deepcopy
from functools import partial
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast
@ -208,7 +208,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
presence_penalty: float | None = None,
seed: int | None = None,
) -> None:
locals_: Final = locals().copy()
locals_: Final[Mapping[str, object]] = locals().copy()
for key, value in locals_.items():
if key != "self" and value is not None:
setattr(self.__class__, key, value)
@ -1427,7 +1427,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _extract_server_side_tool_invocations(
parts: list[HttpxPartType],
) -> list[dict[str, Any]] | None:
) -> list[dict[str, object]] | None:
"""Extract server-side tool invocations (toolCall/toolResponse) from parts.
These are returned by Gemini when context circulation is enabled
@ -1438,15 +1438,15 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
Returns:
List of server-side invocation dicts if any found, None otherwise.
"""
invocations: Final[list[dict[str, Any]]] = []
invocations: Final[list[dict[str, object]]] = []
# Index toolCalls by id so we can pair them with responses
tool_calls_by_id: Final[dict[str, dict[str, Any]]] = {}
tool_responses_by_id: Final[dict[str, dict[str, Any]]] = {}
tool_calls_by_id: Final[dict[str, dict[str, object]]] = {}
tool_responses_by_id: Final[dict[str, dict[str, object]]] = {}
for part in parts:
if "toolCall" in part:
tc = part["toolCall"]
entry: dict[str, Any] = {
entry: dict[str, object] = {
"tool_type": tc.get("toolType"),
"id": tc.get("id"),
"args": tc.get("args"),
@ -1753,7 +1753,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
response_tokens_details: CompletionTokensDetailsWrapper | None = None
usage_metadata: Final = completion_response["usageMetadata"]
def _get_token_count(detail: Mapping[str, Any]) -> int:
def _get_token_count(detail: Mapping[str, object]) -> int:
raw_token_count: Final = detail.get("tokenCount", detail.get("token_count", 0))
return raw_token_count if isinstance(raw_token_count, int) else 0
@ -2068,7 +2068,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
)
@staticmethod
def _get_stream_chunk_attr(chunk: Any, field_name: str) -> Any:
def _get_stream_chunk_attr(chunk: object, field_name: str) -> object:
if isinstance(chunk, dict):
value = chunk.get(field_name)
if value is not None:
@ -2110,10 +2110,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
def apply_assembled_streaming_response_metadata(
self,
response: ModelResponse,
chunks: list[Any],
chunks: list[object],
) -> None:
for field_name in VERTEX_AI_PROVIDER_METADATA_FIELDS:
merged: list[Any] = []
merged: list[object] = []
for chunk in chunks:
value = VertexGeminiConfig._get_stream_chunk_attr(chunk, field_name)
if not value:
@ -2214,8 +2214,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
functions: ChatCompletionToolCallFunctionChunk | None = None
thinking_blocks: list[ChatCompletionThinkingBlock] | None = None
reasoning_content: str | None = None
thought_signatures: Any | None = None
server_side_tool_invocations: list[dict[str, Any]] | None = None
thought_signatures: Sequence[str] | None = None
server_side_tool_invocations: list[dict[str, object]] | None = None
for idx, candidate in enumerate(_candidates):
if "content" not in candidate:
@ -2370,7 +2370,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: Any,
encoding: object,
api_key: str | None = None,
json_mode: bool | None = None,
) -> ModelResponse:
@ -2486,7 +2486,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
## ADD SERVICE TIER ##
if getattr(raw_response, "headers", None):
if service_tier := raw_response.headers.get("x-gemini-service-tier"):
service_tier: Final[str | None] = raw_response.headers.get("x-gemini-service-tier")
if service_tier:
if service_tier.lower() == "standard":
setattr(model_response, "service_tier", "default")
else:
@ -2660,7 +2661,7 @@ class VertexLLM(VertexBase):
print_verbose: Callable,
data: dict,
timeout: float | httpx.Timeout | None,
encoding,
encoding: object,
logging_obj,
stream,
optional_params: dict,
@ -2756,7 +2757,7 @@ class VertexLLM(VertexBase):
"vertex_ai", "vertex_ai_beta", "gemini"
], # if it's vertex_ai or gemini (google ai studio)
timeout: float | httpx.Timeout | None,
encoding,
encoding: object,
logging_obj,
stream,
optional_params: dict,
@ -2873,7 +2874,7 @@ class VertexLLM(VertexBase):
custom_llm_provider: Literal[
"vertex_ai", "vertex_ai_beta", "gemini"
], # if it's vertex_ai or gemini (google ai studio)
encoding,
encoding: object,
logging_obj,
optional_params: dict,
acompletion: bool,
@ -3122,7 +3123,7 @@ class ModelResponseIterator:
def _apply_stream_candidates(
self,
_candidates: list[Candidates],
model_response: Any,
model_response: "ModelResponseStream",
) -> tuple[list[dict], list[dict], list[dict], list[dict]]:
(
grounding_metadata,
@ -3200,7 +3201,7 @@ class ModelResponseIterator:
def _apply_stream_usage_metadata(
self,
processed_chunk: Any,
processed_chunk: GenerateContentResponseBody,
model_response: Any,
grounding_metadata: list[dict],
) -> Usage | None:

View file

@ -28,7 +28,7 @@ class TextStreamer:
Fake streaming iterator for Vertex AI Model Garden calls
"""
def __init__(self, text):
def __init__(self, text: str):
self.text = text.split() # let's assume words as a streaming unit
self.index = 0

View file

@ -19,12 +19,12 @@ import random
import sys
import time
import traceback
from collections.abc import AsyncIterator, Coroutine, Iterable, Mapping
from collections.abc import AsyncIterator, Coroutine, Iterable, Mapping, Sequence
from concurrent import futures
from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait
from copy import deepcopy
from functools import partial
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast, get_args
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, Union, cast, get_args
from litellm._logging import _redact_string
from litellm._uuid import uuid
@ -504,6 +504,7 @@ async def acompletion(
model=model,
custom_llm_provider=cast(str | None, custom_llm_provider), # cast-ok: read from untyped kwargs
tools=tools,
enable_prompt_caching=cast(bool | None, kwargs.get("enable_prompt_caching")), # cast-ok: untyped kwargs
)
if isinstance(litellm_logging_obj, LiteLLMLoggingObj) and (
@ -596,7 +597,7 @@ async def acompletion(
_, custom_llm_provider, _, _ = get_llm_provider(
model=model,
custom_llm_provider=custom_llm_provider,
api_base=completion_kwargs.get("base_url", None),
api_base=base_url,
)
fallbacks = fallbacks or litellm.model_fallbacks
@ -633,10 +634,10 @@ async def acompletion(
init_response: Final = await loop.run_in_executor(None, func_with_context)
if isinstance(init_response, dict) or isinstance(init_response, ModelResponse): ## CACHING SCENARIO
if isinstance(init_response, dict):
response = ModelResponse(**init_response)
response = _model_response_from_cached_dict(init_response)
response = init_response
elif asyncio.iscoroutine(init_response):
response = await init_response
response = await _resolve_dispatched_chat_response(init_response)
else:
response = init_response
@ -698,6 +699,20 @@ async def acompletion(
)
async def _resolve_dispatched_chat_response(
pending: Coroutine[object, object, ModelResponse | CustomStreamWrapper],
) -> ModelResponse | CustomStreamWrapper:
return await pending
def _model_response_from_cached_dict(cached_response_dict: Mapping[str, object]) -> ModelResponse:
return ModelResponse(**cached_response_dict)
def _transcription_response_from_cached_dict(cached_response_dict: Mapping[str, object]) -> TranscriptionResponse:
return TranscriptionResponse(**cached_response_dict)
async def _async_streaming(response, model, custom_llm_provider, args):
try:
print_verbose(f"received response in _async_streaming: {response}")
@ -983,12 +998,12 @@ def responses_api_bridge_check(
model: str,
custom_llm_provider: str,
web_search_options: OpenAIWebSearchOptions | None = None,
tools: list[Any] | None = None,
reasoning_effort: Any | None = None,
reasoning_summary: Any | None = None,
tools: Sequence[Mapping[str, object]] | None = None,
reasoning_effort: str | Mapping[str, object] | None = None,
reasoning_summary: object | None = None,
api_base: str | None = None,
) -> tuple[dict, str]:
model_info: dict[str, Any] = {}
model_info: dict[str, object] = {}
# Global flag: route ALL OpenAI chat completions through Responses API.
# Returns early with minimal model_info; callers only inspect the "mode" key.
@ -1110,6 +1125,22 @@ def _drop_input_examples_from_tools(
return cleaned_tools
class _ProxyAuthHeadersProvider(Protocol):
def get_auth_headers(self) -> Mapping[str, str]: ...
def _proxy_auth_headers(proxy_auth: _ProxyAuthHeadersProvider) -> Mapping[str, str]:
return proxy_auth.get_auth_headers()
def _provider_config_items(config: Mapping[str, object]) -> Iterable[tuple[str, object]]:
return config.items()
def _locals_snapshot(values: Mapping[str, object]) -> Mapping[str, object]:
return values
def _build_custom_pricing_entry(
custom_llm_provider: str,
kwargs: dict,
@ -1185,13 +1216,31 @@ def _register_custom_pricing_for_request(
)
def _dispatch_metadata(ctx: _CompletionDispatchContext) -> Mapping[str, object] | None:
return ctx.metadata
def _dispatch_client_http(ctx: _CompletionDispatchContext) -> HTTPHandler | AsyncHTTPHandler | None:
return ctx.client
def _dispatch_client_azure(
ctx: _CompletionDispatchContext,
) -> openai.AzureOpenAI | openai.AsyncAzureOpenAI | HTTPHandler | AsyncHTTPHandler | None:
return ctx.client
def _dispatch_client_openai(ctx: _CompletionDispatchContext) -> openai.OpenAI | openai.AsyncOpenAI | None:
return ctx.client
def _complete_azure(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
_azure_detection_model: Final = ctx._azure_detection_model
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
api_version = ctx.api_version
client: Final = ctx.client
client: Final = _dispatch_client_azure(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
extra_headers: Final = ctx.extra_headers
headers = ctx.headers
@ -1232,7 +1281,8 @@ def _complete_azure(ctx: _CompletionDispatchContext) -> _CompletionDispatchResul
"AZURE_AD_TOKEN"
)
azure_ad_token_provider: Final = litellm_params.get("azure_ad_token_provider", None)
azure_ad_token_provider_value: Final = litellm_params.get("azure_ad_token_provider", None)
azure_ad_token_provider: Final = azure_ad_token_provider_value if callable(azure_ad_token_provider_value) else None
headers = headers or litellm.headers
@ -1244,7 +1294,7 @@ def _complete_azure(ctx: _CompletionDispatchContext) -> _CompletionDispatchResul
if litellm.AzureOpenAIO1Config().is_o_series_model(model=_azure_detection_model):
## LOAD CONFIG - if set
config = litellm.AzureOpenAIO1Config.get_config()
for k, v in config.items():
for k, v in _provider_config_items(config):
if (
k not in optional_params
): # completion(top_k=3) > azure_config(top_k=3) <- allows for dynamic variables to be passed in
@ -1273,7 +1323,7 @@ def _complete_azure(ctx: _CompletionDispatchContext) -> _CompletionDispatchResul
else:
## LOAD CONFIG - if set
config = litellm.AzureOpenAIConfig.get_config()
for k, v in config.items():
for k, v in _provider_config_items(config):
if (
k not in optional_params
): # completion(top_k=3) > azure_config(top_k=3) <- allows for dynamic variables to be passed in
@ -1323,7 +1373,7 @@ def _complete_azure_text(ctx: _CompletionDispatchContext) -> _CompletionDispatch
api_base = ctx.api_base
api_key = ctx.api_key
api_version = ctx.api_version
client: Final = ctx.client
client: Final = _dispatch_client_azure(ctx)
extra_headers: Final = ctx.extra_headers
headers = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -1358,7 +1408,8 @@ def _complete_azure_text(ctx: _CompletionDispatchContext) -> _CompletionDispatch
"AZURE_AD_TOKEN"
)
azure_ad_token_provider: Final = litellm_params.get("azure_ad_token_provider", None)
azure_ad_token_provider_value: Final = litellm_params.get("azure_ad_token_provider", None)
azure_ad_token_provider: Final = azure_ad_token_provider_value if callable(azure_ad_token_provider_value) else None
headers = headers or litellm.headers
@ -1367,7 +1418,7 @@ def _complete_azure_text(ctx: _CompletionDispatchContext) -> _CompletionDispatch
## LOAD CONFIG - if set
config: Final = litellm.AzureOpenAIConfig.get_config()
for k, v in config.items():
for k, v in _provider_config_items(config):
if (
k not in optional_params
): # completion(top_k=3) > azure_config(top_k=3) <- allows for dynamic variables to be passed in
@ -1415,7 +1466,7 @@ def _complete_deepseek(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe
acompletion: Final = ctx.acompletion
api_base: Final = ctx.api_base
api_key: Final = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -1466,7 +1517,7 @@ def _complete_azure_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
extra_headers: Final = ctx.extra_headers
headers = ctx.headers
@ -1622,7 +1673,7 @@ def _complete_text_completion_openai(
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_openai(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -1654,7 +1705,7 @@ def _complete_text_completion_openai(
## LOAD CONFIG - if set
config: Final = litellm.OpenAITextCompletionConfig.get_config()
for k, v in config.items():
for k, v in _provider_config_items(config):
if (
k not in optional_params
): # completion(top_k=3) > openai_text_config(top_k=3) <- allows for dynamic variables to be passed in
@ -1704,7 +1755,7 @@ def _complete_fireworks_ai(
acompletion: Final = ctx.acompletion
api_base: Final = ctx.api_base
api_key: Final = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -1755,7 +1806,7 @@ def _complete_heroku(ctx: _CompletionDispatchContext) -> _CompletionDispatchResu
acompletion: Final = ctx.acompletion
api_base: Final = ctx.api_base
api_key: Final = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -1805,7 +1856,7 @@ def _complete_ragflow(ctx: _CompletionDispatchContext) -> _CompletionDispatchRes
acompletion: Final = ctx.acompletion
api_base: Final = ctx.api_base
api_key: Final = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -1855,7 +1906,7 @@ def _complete_xai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
acompletion: Final = ctx.acompletion
api_base: Final = ctx.api_base
api_key: Final = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -1906,7 +1957,7 @@ def _complete_groq(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -1938,7 +1989,7 @@ def _complete_groq(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult
## LOAD CONFIG - if set
config: Final = litellm.GroqChatConfig.get_config()
for k, v in config.items():
for k, v in _provider_config_items(config):
if (
k not in optional_params
): # completion(top_k=3) > openai_config(top_k=3) <- allows for dynamic variables to be passed in
@ -1970,7 +2021,7 @@ def _complete_bedrock_mantle(
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -1987,7 +2038,7 @@ def _complete_bedrock_mantle(
api_key = api_key or litellm.api_key or get_secret("BEDROCK_MANTLE_API_KEY")
headers = headers or litellm.headers
config: Final = litellm.BedrockMantleChatConfig.get_config()
for k, v in config.items():
for k, v in _provider_config_items(config):
if k not in optional_params:
optional_params[k] = v
return base_llm_http_handler.completion(
@ -2014,7 +2065,7 @@ def _complete_a2a(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -2077,7 +2128,7 @@ def _complete_gigachat(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe
acompletion: Final = ctx.acompletion
api_base: Final = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -2139,7 +2190,7 @@ def _complete_sap(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
acompletion: Final = ctx.acompletion
api_base: Final = ctx.api_base
api_key: Final = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -2155,7 +2206,7 @@ def _complete_sap(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
headers = headers or litellm.headers
## LOAD CONFIG - if set
config: Final = litellm.GenAIHubOrchestrationConfig.get_config()
for k, v in config.items():
for k, v in _provider_config_items(config):
if (
k not in optional_params
): # completion(top_k=3) > openai_config(top_k=3) <- allows for dynamic variables to be passed in
@ -2187,7 +2238,7 @@ def _complete_aiohttp_openai(
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
extra_headers: Final = ctx.extra_headers
headers = ctx.headers
@ -2242,7 +2293,7 @@ def _complete_cometapi(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -2291,7 +2342,7 @@ def _complete_minimax(ctx: _CompletionDispatchContext) -> _CompletionDispatchRes
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -2337,7 +2388,7 @@ def _complete_hosted_vllm(ctx: _CompletionDispatchContext) -> _CompletionDispatc
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key: Final = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -2383,7 +2434,7 @@ def _complete_custom_openai(
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
custom_prompt_dict: Final = ctx.custom_prompt_dict
extra_headers = ctx.extra_headers
@ -2392,7 +2443,7 @@ def _complete_custom_openai(
logger_fn: Final = ctx.logger_fn
logging: Final = ctx.logging
messages: Final = ctx.messages
metadata: Final = ctx.metadata
metadata: Final = _dispatch_metadata(ctx)
model: Final = ctx.model
model_response: Final = ctx.model_response
optional_params: Final = ctx.optional_params
@ -2445,7 +2496,7 @@ def _complete_custom_openai(
## LOAD CONFIG - if set
config: Final = litellm.OpenAIConfig.get_config()
for k, v in config.items():
for k, v in _provider_config_items(config):
if (
k not in optional_params
): # completion(top_k=3) > openai_config(top_k=3) <- allows for dynamic variables to be passed in
@ -2522,7 +2573,7 @@ def _complete_mistral(ctx: _CompletionDispatchContext) -> _CompletionDispatchRes
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -2673,7 +2724,7 @@ def _complete_anthropic(ctx: _CompletionDispatchContext) -> _CompletionDispatchR
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
custom_prompt_dict = ctx.custom_prompt_dict
headers: Final = ctx.headers
@ -2972,7 +3023,7 @@ def _complete_huggingface(ctx: _CompletionDispatchContext) -> _CompletionDispatc
acompletion: Final = ctx.acompletion
api_base: Final = ctx.api_base
api_key: Final = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -3015,7 +3066,7 @@ def _complete_oci(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
acompletion: Final = ctx.acompletion
api_base: Final = ctx.api_base
api_key: Final = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -3050,7 +3101,7 @@ def _complete_compactifai(ctx: _CompletionDispatchContext) -> _CompletionDispatc
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -3126,7 +3177,7 @@ def _complete_databricks(ctx: _CompletionDispatchContext) -> _CompletionDispatch
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
headers = ctx.headers
litellm_params: Final = ctx.litellm_params
logging: Final = ctx.logging
@ -3198,7 +3249,7 @@ def _complete_datarobot(ctx: _CompletionDispatchContext) -> _CompletionDispatchR
acompletion: Final = ctx.acompletion
api_base: Final = ctx.api_base
api_key: Final = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -3235,7 +3286,7 @@ def _complete_openrouter(ctx: _CompletionDispatchContext) -> _CompletionDispatch
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
headers = ctx.headers
litellm_params: Final = ctx.litellm_params
logging: Final = ctx.logging
@ -3273,7 +3324,7 @@ def _complete_openrouter(ctx: _CompletionDispatchContext) -> _CompletionDispatch
## Load Config
config: Final = litellm.OpenrouterConfig.get_config()
for k, v in config.items():
for k, v in _provider_config_items(config):
if k == "extra_body":
# we use openai 'extra_body' to pass openrouter specific params - transforms, route, models
if "extra_body" in optional_params:
@ -3314,7 +3365,7 @@ def _complete_vercel_ai_gateway(
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
headers = ctx.headers
litellm_params: Final = ctx.litellm_params
logging: Final = ctx.logging
@ -3351,7 +3402,7 @@ def _complete_vercel_ai_gateway(
## Load Config
config: Final = litellm.VercelAIGatewayConfig.get_config()
for k, v in config.items():
for k, v in _provider_config_items(config):
if k == "extra_body":
# we use openai 'extra_body' to pass vercel specific params - providerOptions
if "extra_body" in optional_params:
@ -3392,7 +3443,7 @@ def _complete_vertex_ai_beta(
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key: Final = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -3457,7 +3508,7 @@ def _complete_vertex_ai_beta(
def _complete_vertex_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
custom_prompt_dict: Final = ctx.custom_prompt_dict
headers: Final = ctx.headers
@ -3754,7 +3805,7 @@ def _complete_text_completion_inception(
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_openai(ctx)
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
logger_fn: Final = ctx.logger_fn
@ -3818,7 +3869,7 @@ def _complete_sagemaker_chat(
acompletion: Final = ctx.acompletion
api_base: Final = ctx.api_base
api_key: Final = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -3881,7 +3932,7 @@ def _complete_bedrock(ctx: _CompletionDispatchContext) -> _CompletionDispatchRes
acompletion: Final = ctx.acompletion
api_base: Final = ctx.api_base
api_key: Final = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_prompt_dict = ctx.custom_prompt_dict
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -4005,7 +4056,7 @@ def _complete_watsonx(ctx: _CompletionDispatchContext) -> _CompletionDispatchRes
acompletion: Final = ctx.acompletion
api_base: Final = ctx.api_base
api_key: Final = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_prompt_dict: Final = ctx.custom_prompt_dict
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -4044,7 +4095,7 @@ def _complete_watsonx_text(
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
logging: Final = ctx.logging
@ -4156,7 +4207,7 @@ def _complete_ollama(ctx: _CompletionDispatchContext) -> _CompletionDispatchResu
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key: Final = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
logging: Final = ctx.logging
@ -4196,7 +4247,7 @@ def _complete_ollama_chat(ctx: _CompletionDispatchContext) -> _CompletionDispatc
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
logging: Final = ctx.logging
@ -4311,7 +4362,7 @@ def _complete_cloudflare(ctx: _CompletionDispatchContext) -> _CompletionDispatch
def _complete_petals(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
api_base = ctx.api_base
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
litellm_params: Final = ctx.litellm_params
logger_fn: Final = ctx.logger_fn
logging: Final = ctx.logging
@ -4353,7 +4404,7 @@ def _complete_snowflake(ctx: _CompletionDispatchContext) -> _CompletionDispatchR
acompletion: Final = ctx.acompletion
api_base: Final = ctx.api_base
api_key: Final = ctx.api_key
client = ctx.client
client = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -4441,7 +4492,7 @@ def _complete_gdc(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -4480,7 +4531,7 @@ def _complete_bytez(ctx: _CompletionDispatchContext) -> _CompletionDispatchResul
acompletion: Final = ctx.acompletion
api_base: Final = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -4520,7 +4571,7 @@ def _complete_lemonade(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe
acompletion: Final = ctx.acompletion
api_base: Final = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -4560,7 +4611,7 @@ def _complete_ovhcloud(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers: Final = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -4603,6 +4654,10 @@ def _complete_ovhcloud(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe
return response
def _custom_api_first_output(resp: httpx.Response | None) -> str:
return resp.json()["data"][0]["output"][0]
def _complete_custom(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
api_base: Final = ctx.api_base
headers: Final = ctx.headers
@ -4651,7 +4706,6 @@ def _complete_custom(ctx: _CompletionDispatchContext) -> _CompletionDispatchResu
**kwargs.get("extra_body", {}),
},
)
response_json: Final = resp.json()
"""
assume all responses from custom api_bases of this format:
{
@ -4665,7 +4719,7 @@ def _complete_custom(ctx: _CompletionDispatchContext) -> _CompletionDispatchResu
]
}
"""
string_response: Final = response_json["data"][0]["output"][0]
string_response: Final = _custom_api_first_output(resp)
## RESPONSE OBJECT
model_response.choices[0].message.content = string_response
model_response.created = int(time.time())
@ -4740,7 +4794,7 @@ def _complete_langgraph(ctx: _CompletionDispatchContext) -> _CompletionDispatchR
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -4789,7 +4843,7 @@ def _complete_langflow(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
client: Final = ctx.client
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
headers = ctx.headers
litellm_params: Final = ctx.litellm_params
@ -4947,7 +5001,7 @@ def completion(
thinking = validate_and_fix_thinking_param(thinking=thinking)
######### unpacking kwargs #####################
args: Final = locals()
args: Final = _locals_snapshot(locals())
# Set by the responses->completion fallback so completion() does not bridge
# back to the Responses API: that round-trip mutually recurses forever for a
@ -5038,7 +5092,7 @@ def completion(
# Inject proxy auth headers if configured
if litellm.proxy_auth is not None:
try:
proxy_headers: Final = litellm.proxy_auth.get_auth_headers()
proxy_headers: Final = _proxy_auth_headers(litellm.proxy_auth)
headers.update(proxy_headers)
except Exception as e:
verbose_logger.warning("Failed to get proxy auth headers: %s", e)
@ -5091,7 +5145,7 @@ def completion(
)
######## end of unpacking kwargs ###########
non_default_params: Final = get_non_default_completion_params(kwargs=kwargs)
litellm_params = {} # used to prevent unbound var errors
litellm_params: dict[str, object] = {} # used to prevent unbound var errors
## PROMPT MANAGEMENT HOOKS ##
from litellm.integrations.anthropic_cache_control_hook import (
@ -5105,6 +5159,7 @@ def completion(
model=model,
custom_llm_provider=cast(str | None, kwargs.get("custom_llm_provider")), # cast-ok: untyped kwargs
tools=tools,
enable_prompt_caching=cast(bool | None, kwargs.get("enable_prompt_caching")), # cast-ok: untyped kwargs
)
if isinstance(litellm_logging_obj, LiteLLMLoggingObj) and (
@ -5913,7 +5968,7 @@ def embedding(
*,
aembedding: Literal[True],
**kwargs,
) -> Coroutine[Any, Any, EmbeddingResponse]:
) -> Coroutine[object, object, EmbeddingResponse]:
...
@ -5964,7 +6019,7 @@ def embedding(
litellm_call_id=None,
logger_fn=None,
**kwargs,
) -> EmbeddingResponse | Coroutine[Any, Any, EmbeddingResponse]:
) -> EmbeddingResponse | Coroutine[object, object, EmbeddingResponse]:
"""
Embedding function that calls an API to generate embeddings for the given input.
@ -6007,7 +6062,7 @@ def embedding(
# Inject proxy auth headers if configured
if litellm.proxy_auth is not None:
try:
proxy_headers: Final = litellm.proxy_auth.get_auth_headers()
proxy_headers: Final = _proxy_auth_headers(litellm.proxy_auth)
headers.update(proxy_headers)
except Exception as e:
verbose_logger.warning("Failed to get proxy auth headers: %s", e)
@ -6084,7 +6139,7 @@ def embedding(
if mock_response is not None:
return mock_embedding(model=model, mock_response=mock_response)
try:
response: EmbeddingResponse | Coroutine[Any, Any, EmbeddingResponse] | None = None
response: EmbeddingResponse | Coroutine[object, object, EmbeddingResponse] | None = None
if azure is True or custom_llm_provider == "azure":
# azure configs
@ -6387,7 +6442,7 @@ def embedding(
response = huggingface_embed.embedding(
model=model,
input=input,
encoding=_get_encoding(),
encoding=sys.modules[__name__].encoding,
api_key=api_key,
api_base=api_base,
logging_obj=logging,
@ -6990,6 +7045,20 @@ def embedding(
###### Text Completion ################
async def _resolve_dispatched_text_completion_response(
pending: Coroutine[
object,
object,
TextCompletionResponse | ModelResponse | CustomStreamWrapper | TextCompletionStreamWrapper,
],
) -> TextCompletionResponse | ModelResponse | CustomStreamWrapper | TextCompletionStreamWrapper:
return await pending
async def _resolve_pending_chat_response(pending: Coroutine[object, object, ModelResponse]) -> ModelResponse:
return await pending
@client
async def atext_completion(*args, **kwargs) -> TextCompletionResponse | TextCompletionStreamWrapper:
"""
@ -7015,7 +7084,7 @@ async def atext_completion(*args, **kwargs) -> TextCompletionResponse | TextComp
else:
response = init_response
elif asyncio.iscoroutine(init_response):
response = await init_response
response = await _resolve_dispatched_text_completion_response(init_response)
else:
response = init_response
@ -7040,7 +7109,7 @@ async def atext_completion(*args, **kwargs) -> TextCompletionResponse | TextComp
if isinstance(response, TextCompletionResponse):
return response
elif asyncio.iscoroutine(response):
response = await response
response = await _resolve_pending_chat_response(response)
text_completion_response = TextCompletionResponse()
text_completion_response = litellm.utils.LiteLLMResponseObjectHandler.convert_chat_to_text_completion(
@ -7330,11 +7399,11 @@ async def aadapter_completion(*, adapter_id: str, **kwargs) -> BaseModel | Adapt
async def aadapter_generate_content(
**kwargs,
) -> dict[str, Any] | AsyncIterator[bytes]:
) -> dict[str, object] | AsyncIterator[bytes]:
from litellm.google_genai.adapters.handler import GenerateContentToCompletionHandler
coro: Final = cast(
Coroutine[Any, Any, dict[str, Any] | AsyncIterator[bytes]],
Coroutine[object, object, dict[str, object] | AsyncIterator[bytes]],
GenerateContentToCompletionHandler.generate_content_handler(**kwargs, _is_async=True),
)
return await coro
@ -7486,7 +7555,7 @@ async def atranscription(*args, **kwargs) -> TranscriptionResponse:
# Await normally
init_response: Final = await loop.run_in_executor(None, func_with_context)
if isinstance(init_response, dict):
response = TranscriptionResponse(**init_response)
response = _transcription_response_from_cached_dict(init_response)
elif isinstance(init_response, TranscriptionResponse): ## CACHING SCENARIO
response = init_response
elif asyncio.iscoroutine(init_response):
@ -7541,7 +7610,7 @@ def transcription(
max_retries: int | None = None,
custom_llm_provider=None,
**kwargs,
) -> TranscriptionResponse | Coroutine[Any, Any, TranscriptionResponse]:
) -> TranscriptionResponse | Coroutine[object, object, TranscriptionResponse]:
"""
Calls openai + azure whisper endpoints.
@ -7608,7 +7677,7 @@ def transcription(
custom_llm_provider=custom_llm_provider,
)
response: TranscriptionResponse | Coroutine[Any, Any, TranscriptionResponse] | None = None
response: TranscriptionResponse | Coroutine[object, object, TranscriptionResponse] | None = None
provider_config: Final = ProviderConfigManager.get_provider_audio_transcription_config(
model=model,
@ -7842,7 +7911,7 @@ def speech(
custom_llm_provider: str | None = None,
aspeech: bool | None = None,
**kwargs,
) -> HttpxBinaryResponseContent | Coroutine[Any, Any, HttpxBinaryResponseContent]:
) -> HttpxBinaryResponseContent | Coroutine[object, object, HttpxBinaryResponseContent]:
user: Final = kwargs.get("user", None)
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
proxy_server_request: Final = kwargs.get("proxy_server_request", None)
@ -7901,7 +7970,7 @@ def speech(
},
custom_llm_provider=custom_llm_provider,
)
response: HttpxBinaryResponseContent | Coroutine[Any, Any, HttpxBinaryResponseContent] | None = None
response: HttpxBinaryResponseContent | Coroutine[object, object, HttpxBinaryResponseContent] | None = None
if custom_llm_provider == "openai" or custom_llm_provider in litellm.openai_compatible_providers:
if voice is None or not (isinstance(voice, str)):
raise litellm.BadRequestError(
@ -8663,7 +8732,7 @@ def stream_chunk_builder(
]
if len(provider_specific_chunks) > 0:
combined_provider_fields: Final[dict[str, Any]] = {}
combined_provider_fields: Final[dict[str, object]] = {}
for chunk in provider_specific_chunks:
fields = chunk["choices"][0]["delta"]["provider_specific_fields"]
if isinstance(fields, dict):
@ -8728,7 +8797,7 @@ def stream_chunk_builder(
async def acount_tokens(
model: str,
messages: list[dict[str, Any]] | None = None,
messages: list[dict[str, object]] | None = None,
tools: list[dict[str, Any]] | None = None,
system: str | None = None,
api_key: str | None = None,
@ -8774,7 +8843,7 @@ async def acount_tokens(
api_base = dynamic_api_base
# Build deployment dict for the token counter
deployment: Final[dict[str, Any]] = {
deployment: Final[dict[str, object]] = {
"litellm_params": {
"model": model,
"api_key": api_key,
@ -8825,29 +8894,37 @@ async def acount_tokens(
# Cache for encoding to avoid repeated __getattr__ calls
_encoding_cache: Any | None = None
_encoding_cache: tiktoken.Encoding | None = None
def _get_encoding():
def _load_module_encoding() -> tiktoken.Encoding:
import sys
return sys.modules[__name__].encoding
def _get_encoding() -> tiktoken.Encoding:
"""Get encoding, loading it lazily if needed."""
global _encoding_cache
if _encoding_cache is None:
import sys
# Access via module to trigger __getattr__ if not cached
_encoding_cache = sys.modules[__name__].encoding
_encoding_cache = _load_module_encoding()
return _encoding_cache
def __getattr__(name: str) -> Any:
def _load_default_encoding() -> tiktoken.Encoding:
from litellm._lazy_imports import _get_default_encoding
return _get_default_encoding()
def __getattr__(name: str) -> tiktoken.Encoding:
"""Lazy import handler for main module"""
if name == "encoding":
# Use _get_default_encoding which properly sets TIKTOKEN_CACHE_DIR
# before loading tiktoken, ensuring the local cache is used
# instead of downloading from the internet
from litellm._lazy_imports import _get_default_encoding
_encoding: Final = _get_default_encoding()
_encoding: Final = _load_default_encoding()
# Cache it in the module's __dict__ for subsequent accesses
import sys

File diff suppressed because it is too large Load diff

View file

@ -49,6 +49,7 @@ class LiteLLM_VerificationToken(LiteLLMPydanticObjectBase):
created_by: str | None = None
updated_at: datetime | None = None
updated_by: str | None = None
settings_updated_at: datetime | None = None
last_active: datetime | None = None
object_permission_id: str | None = None
object_permission: LiteLLM_ObjectPermissionTable | None = None

View file

@ -145,9 +145,8 @@ class MCPOAuth2TokenCache(InMemoryCache):
server.server_id,
)
post_kwargs: Final = {"data": data, **({"headers": token_request.headers} if token_request.headers else {})}
try:
response: Final = await client.post(server.token_url, **post_kwargs)
response: Final = await client.post(server.token_url, data=data, headers=token_request.headers or None)
response.raise_for_status()
except httpx.HTTPStatusError as exc:
raise ValueError(

View file

@ -15,6 +15,8 @@ from collections.abc import Mapping, Sequence
from typing import Any, Final, NamedTuple, Optional, Protocol, Union, runtime_checkable
if typing.TYPE_CHECKING:
from collections.abc import Awaitable, Callable
from fastapi import Request
from mcp.client.session import ClientSession
from mcp.shared.context import RequestContext
@ -28,8 +30,9 @@ if typing.TYPE_CHECKING:
ToolUseContent,
)
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
from litellm.types.utils import ModelResponse
from fastapi import HTTPException
from pydantic import TypeAdapter
@ -1016,7 +1019,7 @@ async def _run_budget_checks(
general_settings=general_settings or {},
route="/chat/completions",
llm_router=_llm_router,
proxy_logging_obj=typing.cast("ProxyLogging", _proxy_logging_obj),
proxy_logging_obj=_proxy_logging_obj,
valid_token=user_api_key_auth,
request=dummy_request,
)
@ -1176,15 +1179,19 @@ async def _build_completion_kwargs(
)
class _AcompletionCall(NamedTuple):
fn: "Callable[..., Awaitable[ModelResponse | CustomStreamWrapper]]"
async def _run_guardrails_and_call_llm(
completion_kwargs: dict[str, Any],
completion_kwargs: dict[str, object],
user_api_key_auth: "UserAPIKeyAuth",
) -> Any:
try:
from litellm.proxy.proxy_server import proxy_logging_obj as _plo
if _plo is not None:
completion_kwargs = await typing.cast("ProxyLogging", _plo).pre_call_hook(
completion_kwargs = await _plo.pre_call_hook(
user_api_key_dict=user_api_key_auth,
data=completion_kwargs,
call_type="acompletion",
@ -1204,10 +1211,10 @@ async def _run_guardrails_and_call_llm(
from litellm.proxy.proxy_server import llm_router
if llm_router is not None:
return await llm_router.acompletion(**completion_kwargs)
return await litellm.acompletion(**completion_kwargs)
return await _AcompletionCall(fn=llm_router.acompletion).fn(**completion_kwargs)
return await _AcompletionCall(fn=litellm.acompletion).fn(**completion_kwargs)
except ImportError:
return await litellm.acompletion(**completion_kwargs)
return await _AcompletionCall(fn=litellm.acompletion).fn(**completion_kwargs)
async def handle_sampling_create_message(

View file

@ -3,8 +3,11 @@ Per-feature OpenAPI snapshot for lazy-loaded routers.
The committed JSON is generated by `python -m litellm.proxy._lazy_openapi_snapshot`
and consumed at runtime so /openapi.json can show full route info for unloaded
features without importing them. CI verifies the file is current and surfaces
any drift as a neutral check.
features without importing them. No CI job regenerates this file; drift surfaces
only indirectly through check-ui-api-types.yml, which rebuilds schema.d.ts from
app.openapi() with the committed snapshot injected. After changing any lazily
loaded route or this generator, rerun the module and commit the JSON, then run
`npm run gen:api` in ui/litellm-dashboard and commit schema.d.ts.
"""
import json
@ -89,8 +92,6 @@ def generate_snapshot() -> dict[str, dict]:
from litellm.proxy.proxy_server import app, ensure_unique_openapi_operation_ids
for feat in LAZY_FEATURES:
if feat.module_path in sys.modules:
continue
try:
module = importlib.import_module(feat.module_path)
feat.register_fn(app, module)
@ -100,7 +101,7 @@ def generate_snapshot() -> dict[str, dict]:
fragments: Final[dict[str, dict]] = {}
used_operation_ids: Final[set[str]] = set()
for feat in LAZY_FEATURES:
feat_routes = [r for r in app.routes if any(getattr(r, "path", "").startswith(p) for p in feat.path_prefixes)]
feat_routes = [r for r in app.routes if feat.matches(getattr(r, "path", ""))]
if not feat_routes:
continue
_stabilize_multi_method_route_ids(feat_routes)

View file

@ -1,7 +1,7 @@
import enum
import json
import os
from collections.abc import Callable
from collections.abc import Callable, Mapping
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal
@ -11,6 +11,7 @@ from pydantic import (
ConfigDict,
Field,
Json,
PositiveInt,
field_validator,
model_validator,
)
@ -1102,9 +1103,12 @@ class AllowedVectorStoreIndexItem(LiteLLMPydanticObjectBase):
class KeyRequestBase(GenerateRequestBase):
key: str | None = None
default_estimated_output_tokens: PositiveInt | None = None
default_estimated_output_tokens_per_model: Mapping[str, PositiveInt] | None = None
budget_id: str | None = None
tags: list[str] | None = None
disable_global_guardrails: bool | None = None
enable_prompt_caching: bool | None = None
throttle_on_budget_exceeded: bool | None = None
enforced_params: list[str] | None = None
allowed_routes: list | None = []
@ -1819,6 +1823,8 @@ class NewTeamRequest(TeamBase):
)
model_tpm_limit: dict[str, int] | None = None
default_estimated_output_tokens: PositiveInt | None = None
default_estimated_output_tokens_per_model: Mapping[str, PositiveInt] | None = None
mcp_rpm_limit: dict[str, int] | None = None
team_member_budget: float | None = None # allow user to set a budget for all team members
team_member_rpm_limit: int | None = None # allow user to set RPM limit for all team members
@ -1883,6 +1889,8 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
prompts: list[str] | None = None
model_rpm_limit: dict[str, int] | None = None
model_tpm_limit: dict[str, int] | None = None
default_estimated_output_tokens: PositiveInt | None = None
default_estimated_output_tokens_per_model: Mapping[str, PositiveInt] | None = None
mcp_rpm_limit: dict[str, int] | None = None
allowed_vector_store_indexes: list[AllowedVectorStoreIndexItem] | None = None
enforced_batch_output_expires_after: dict | None = None
@ -4018,6 +4026,13 @@ class LitellmDataForBackendLLMCall(TypedDict, total=False):
stream_timeout: float | None
user: str | None
num_retries: int | None
# True when the effective timeout came from a caller-controlled source (the
# `x-litellm-timeout`/`x-litellm-stream-timeout` headers, or a `timeout`/`request_timeout`/
# `stream_timeout` field in the request body) rather than deployment config, so a
# deliberately tiny value isn't treated as a deployment health signal (see
# cooldown_handlers._trigger_cooldown_for_failed_deployment).
client_side_timeout: bool
keepalive_seconds: float | None
class LitellmMetadataFromRequestHeaders(TypedDict, total=False):
@ -4097,6 +4112,8 @@ class PassThroughEndpointLoggingTypedDict(TypedDict):
LiteLLM_ManagementEndpoint_MetadataFields: Final = [
"model_rpm_limit",
"model_tpm_limit",
"default_estimated_output_tokens",
"default_estimated_output_tokens_per_model",
"mcp_rpm_limit",
"tag_rpm_limit",
"rpm_limit_type",
@ -4108,6 +4125,7 @@ LiteLLM_ManagementEndpoint_MetadataFields: Final = [
"enforced_batch_output_expires_after",
"enforced_file_expires_after",
"throttle_on_budget_exceeded",
"enable_prompt_caching",
]
LiteLLM_ManagementEndpoint_MetadataFields_Premium: Final = [

View file

@ -11,7 +11,7 @@ The A2A SDK can point to LiteLLM's URL and invoke agents registered with LiteLLM
"""
import json
from collections.abc import AsyncGenerator
from collections.abc import AsyncGenerator, Mapping
from copy import deepcopy
from typing import TYPE_CHECKING, Any, Final
from urllib.parse import urlparse
@ -36,7 +36,7 @@ from litellm.proxy.agent_endpoints.databricks_oauth import (
)
from litellm.proxy.agent_endpoints.utils import merge_agent_headers
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.utils import get_custom_url
from litellm.proxy.utils import ProxyLogging, get_custom_url
from litellm.types.utils import all_litellm_params
if TYPE_CHECKING:
@ -46,7 +46,7 @@ if TYPE_CHECKING:
router: Final = APIRouter()
_PASCAL_TO_WIRE: Final[dict[str, str]] = {
_PASCAL_TO_WIRE: Final[Mapping[str, str]] = {
"SendMessage": "message/send",
"SendStreamingMessage": "message/stream",
"GetTask": "tasks/get",
@ -118,9 +118,9 @@ def _caller_identity_headers(user_api_key_dict: UserAPIKeyAuth) -> dict[str, str
def _forwarding_headers(
user_api_key_dict: UserAPIKeyAuth,
request_data: dict[str, Any],
agent_extra_headers: dict[str, str] | None,
) -> dict[str, str] | None:
request_data: Mapping[str, object],
agent_extra_headers: Mapping[str, str] | None,
) -> Mapping[str, str] | None:
sanitized: Final = (
{k: v for k, v in agent_extra_headers.items() if not k.lower().startswith("x-litellm-")}
if agent_extra_headers
@ -136,7 +136,7 @@ def _forwarding_headers(
def _jsonrpc_error(
request_id: Any | None,
request_id: object,
code: int,
message: str,
status_code: int = 400,
@ -162,7 +162,7 @@ def _get_agent(agent_id: str):
return agent
def _enforce_inbound_trace_id(agent: Any, request: Request) -> None:
def _enforce_inbound_trace_id(agent: "AgentResponse", request: Request) -> None:
"""Raise 400 if agent requires x-litellm-trace-id on inbound calls and it is missing."""
agent_litellm_params: Final = agent.litellm_params or {}
if not agent_litellm_params.get("require_trace_id_on_calls_to_agent"):
@ -181,8 +181,8 @@ def _enforce_inbound_trace_id(agent: Any, request: Request) -> None:
async def _forward_jsonrpc(
agent_url: str,
body: dict[str, Any],
extra_headers: dict[str, str] | None = None,
body: dict[str, object],
extra_headers: Mapping[str, str] | None = None,
) -> dict[str, Any]:
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
@ -205,11 +205,11 @@ async def _forward_jsonrpc(
async def _a2a_sse_event_source(
agent_url: str,
body: dict[str, Any],
request_id: Any | None = None,
extra_headers: dict[str, str] | None = None,
body: Mapping[str, object],
request_id: str | int | None = None,
extra_headers: Mapping[str, str] | None = None,
served_version: A2AVersion = "0.3",
) -> AsyncGenerator[dict, None]:
) -> AsyncGenerator[Mapping[str, object], None]:
"""Stream an upstream A2A SSE response as parsed JSON-RPC event dicts.
Upstream HTTP/JSON-RPC errors are surfaced as a single JSON-RPC error event
@ -234,7 +234,7 @@ async def _a2a_sse_event_source(
try:
if not resp.is_success:
error_body: Final = await resp.aread()
error_event: dict[str, Any] | None = None
error_event: Mapping[str, object] | None = None
try:
parsed: Final = json.loads(error_body)
if isinstance(parsed, dict) and "error" in parsed:
@ -267,12 +267,12 @@ async def _a2a_sse_event_source(
async def _forward_jsonrpc_sse(
agent_url: str,
body: dict[str, Any],
request_id: Any | None = None,
extra_headers: dict[str, str] | None = None,
proxy_logging_obj: Any | None = None,
user_api_key_dict: Any | None = None,
request_data: dict[str, Any] | None = None,
body: Mapping[str, object],
request_id: str | int | None = None,
extra_headers: Mapping[str, str] | None = None,
proxy_logging_obj: ProxyLogging | None = None,
user_api_key_dict: UserAPIKeyAuth | None = None,
request_data: dict[str, object] | None = None,
served_version: A2AVersion = "0.3",
) -> StreamingResponse:
event_source: Final = _a2a_sse_event_source(
@ -283,10 +283,10 @@ async def _forward_jsonrpc_sse(
served_version=served_version,
)
def _serialize_chunk(chunk: Any) -> str:
def _serialize_chunk(chunk: object) -> str:
return f"data: {json.dumps(chunk)}\n\n"
def _serialize_error(proxy_exc: Any) -> str:
def _serialize_error(proxy_exc: object) -> str:
return (
"data: "
+ json.dumps(
@ -331,17 +331,17 @@ async def _forward_jsonrpc_sse(
async def _handle_stream_message(
api_base: str | None,
request_id: Any,
params: dict[str, Any],
litellm_params: dict[str, Any] | None = None,
request_id: str | int,
params: dict[str, object],
litellm_params: dict[str, object] | None = None,
agent_id: str | None = None,
metadata: dict[str, Any] | None = None,
proxy_server_request: dict[str, Any] | None = None,
metadata: dict[str, object] | None = None,
proxy_server_request: dict[str, object] | None = None,
*,
agent_extra_headers: dict[str, str] | None = None,
user_api_key_dict: UserAPIKeyAuth | None = None,
request_data: dict[str, Any] | None = None,
proxy_logging_obj: Any | None = None,
request_data: dict[str, object] | None = None,
proxy_logging_obj: ProxyLogging | None = None,
served_version: A2AVersion = "0.3",
) -> StreamingResponse:
"""Handle message/stream method via SDK functions.
@ -430,7 +430,7 @@ async def _handle_stream_message(
obj = normalize_stream_event(obj, served_version, request_id=request_id)
return json.dumps(obj) + "\n"
def _ndjson_error(proxy_exc: Any) -> str:
def _ndjson_error(proxy_exc: object) -> str:
return (
json.dumps(
{
@ -669,7 +669,7 @@ async def invoke_agent_a2a(
agent_name: Final = agent_card_params.get("name", agent_id)
# Get litellm_params (may include custom_llm_provider for completion bridge)
litellm_params = agent.litellm_params or {}
litellm_params: dict[str, object] = agent.litellm_params or {}
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
# Hand the authenticated key hash to the completion bridge so provider
@ -725,7 +725,7 @@ async def invoke_agent_a2a(
request_data = data
# Build merged headers for the backend agent
static_headers: Final[dict[str, str]] = dict(agent.static_headers or {})
static_headers: Final[Mapping[str, str]] = dict(agent.static_headers or {})
raw_headers: Final = dict(request.headers)
normalized: Final = {k.lower(): v for k, v in raw_headers.items()}
@ -893,7 +893,7 @@ async def invoke_agent_a2a(
detail="Push notification URL must be a string",
)
_validate_push_notification_url(callback_url)
forward_body = {
forward_body: dict[str, object] = {
"jsonrpc": "2.0",
"id": request_id,
"method": method,

View file

@ -13,8 +13,8 @@ import asyncio
import math
import re
import time
from collections.abc import Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast
from collections.abc import Iterator, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, cast
from fastapi import HTTPException, Request, status
from pydantic import BaseModel
@ -87,6 +87,7 @@ from litellm.proxy.utils import PrismaClient, ProxyLogging, log_db_metrics
from litellm.repositories.budget_repository import BudgetRepository
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
from litellm.repositories.organization_repository import OrganizationRepository
from litellm.repositories.prisma_protocols import RowT_co
from litellm.repositories.project_repository import ProjectRepository
from litellm.repositories.table_repositories import (
AccessGroupRepository,
@ -110,11 +111,144 @@ from .auth_utils import get_model_from_request, get_request_route_template
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
Span = _Span | Any
Span = _Span
else:
Span = Any
class _PrismaDictableRow(Protocol):
def dict(self) -> Mapping[str, object]: ...
class _PrismaJWTKeyMappingRow(Protocol):
token: str
class _PrismaModelDumpRow(Protocol):
def model_dump(self) -> Mapping[str, object]: ...
class _PrismaTeamRow(Protocol):
def dict(self) -> Mapping[str, object]: ...
def model_dump(self) -> Mapping[str, object]: ...
class _PrismaVectorStoreRow(Protocol):
def dict(self) -> Mapping[str, object]: ...
def model_dump(self) -> Mapping[str, object]: ...
def __iter__(self) -> Iterator[tuple[str, object]]: ...
class _PrismaUserRow(Protocol):
user_id: str
organization_memberships: Sequence[LiteLLM_OrganizationMembershipTable | None] | None
def __iter__(self) -> Iterator[tuple[str, object]]: ...
class _PrismaAuthTable(Protocol[RowT_co]):
async def find_unique(
self, *, where: Mapping[str, object], include: Mapping[str, object] | None = None
) -> RowT_co | None: ...
async def find_first(
self, *, where: Mapping[str, object], include: Mapping[str, object] | None = None
) -> RowT_co | None: ...
async def find_many(
self,
*,
where: Mapping[str, object],
include: Mapping[str, object] | None = None,
take: int | None = None,
) -> Sequence[RowT_co]: ...
async def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> RowT_co | None: ...
async def create(self, *, data: Mapping[str, object], include: Mapping[str, object] | None = None) -> RowT_co: ...
class _PrismaTableHolder(Protocol[RowT_co]):
@property
def table(self) -> _PrismaAuthTable[RowT_co]: ...
def _dictable_table(repo: _PrismaTableHolder[_PrismaDictableRow]) -> _PrismaAuthTable[_PrismaDictableRow]:
return repo.table
def _jwt_key_mapping_table(
repo: _PrismaTableHolder[_PrismaJWTKeyMappingRow],
) -> _PrismaAuthTable[_PrismaJWTKeyMappingRow]:
return repo.table
def _model_dump_table(repo: _PrismaTableHolder[_PrismaModelDumpRow]) -> _PrismaAuthTable[_PrismaModelDumpRow]:
return repo.table
def _team_table(repo: _PrismaTableHolder[_PrismaTeamRow]) -> _PrismaAuthTable[_PrismaTeamRow]:
return repo.table
def _vector_store_table(repo: _PrismaTableHolder[_PrismaVectorStoreRow]) -> _PrismaAuthTable[_PrismaVectorStoreRow]:
return repo.table
def _user_table(repo: _PrismaTableHolder[_PrismaUserRow]) -> _PrismaAuthTable[_PrismaUserRow]:
return repo.table
def _object_permission_table(
repo: _PrismaTableHolder[LiteLLM_ObjectPermissionTable],
) -> _PrismaAuthTable[LiteLLM_ObjectPermissionTable]:
return repo.table
class _PrismaTagRow(Protocol):
tag_name: str
def dict(self) -> Mapping[str, object]: ...
def _tag_table(repo: _PrismaTableHolder[_PrismaTagRow]) -> _PrismaAuthTable[_PrismaTagRow]:
return repo.table
class _RawCacheRead(Protocol):
async def async_get_cache(self, *, key: str) -> object: ...
def _raw_cache(cache: _RawCacheRead) -> _RawCacheRead:
return cache
class _BudgetCacheRead(Protocol):
async def async_get_cache(self, *, key: str) -> "LiteLLM_BudgetTable | Mapping[str, object] | None": ...
def _budget_cache(cache: _BudgetCacheRead) -> _BudgetCacheRead:
return cache
def _typed_request_body(request_body: dict) -> Mapping[str, object]:
return request_body
class _JsonLoadsObj(Protocol):
def __call__(self, data: str) -> object: ...
def _typed_json_loads(fn: _JsonLoadsObj) -> _JsonLoadsObj:
return fn
_safe_json_loads_obj: Final = _typed_json_loads(safe_json_loads)
last_db_access_time: Final = LimitedSizeOrderedDict(max_size=100)
db_cache_expiry: Final = DEFAULT_IN_MEMORY_TTL # refresh every 5s
@ -384,7 +518,7 @@ _GUARDRAIL_MODIFICATION_KEYS: Final[tuple] = (
)
def _guardrail_modification_check(request_body: dict, team_object: LiteLLM_TeamTable | None) -> None:
def _guardrail_modification_check(request_body: Mapping[str, object], team_object: LiteLLM_TeamTable | None) -> None:
"""
Reject user-supplied metadata flags that would modify guardrail behavior
unless the team has explicit permission. Checked keys include the plural
@ -399,7 +533,7 @@ def _guardrail_modification_check(request_body: dict, team_object: LiteLLM_TeamT
"""
from litellm.proxy.guardrails.guardrail_helpers import can_modify_guardrails
def _coerce_to_dict(container: Any) -> dict | None:
def _coerce_to_dict(container: object) -> dict | None:
"""Accept dict or JSON-string (from multipart/form-data or extra_body).
Without this, an attacker can smuggle guardrail keys past the check by
@ -411,11 +545,11 @@ def _guardrail_modification_check(request_body: dict, team_object: LiteLLM_TeamT
if isinstance(container, dict):
return container
if isinstance(container, str):
parsed: Final = safe_json_loads(container)
parsed: Final = _safe_json_loads_obj(container)
return parsed if isinstance(parsed, dict) else None
return None
def _user_requested_modification(container: Any) -> bool:
def _user_requested_modification(container: object) -> bool:
coerced: Final = _coerce_to_dict(container)
if coerced is None:
return False
@ -731,7 +865,7 @@ async def common_checks(
_enforce_user_param_check(general_settings, request, request_body, route)
_global_proxy_budget_check(global_proxy_spend, skip_all_budget_checks, route)
_guardrail_modification_check(request_body, team_object)
_guardrail_modification_check(_typed_request_body(request_body), team_object)
# 10 [OPTIONAL] Organization RBAC checks
organization_role_based_access_check(user_object=user_object, route=route, request_body=request_body)
@ -955,7 +1089,7 @@ async def get_default_end_user_budget(
# Fetch from database
try:
budget_record: Final = await BudgetRepository(prisma_client).table.find_unique(
budget_record: Final = await _dictable_table(BudgetRepository(prisma_client)).find_unique(
where={"budget_id": litellm.max_end_user_budget_id}
)
@ -1007,14 +1141,16 @@ async def get_team_member_default_budget(
cache_key: Final = f"team_member_default_budget:{budget_id}"
cached_budget: Final = await user_api_key_cache.async_get_cache(key=cache_key)
cached_budget: Final = await _budget_cache(user_api_key_cache).async_get_cache(key=cache_key)
if isinstance(cached_budget, LiteLLM_BudgetTable):
return cached_budget
if isinstance(cached_budget, dict):
return LiteLLM_BudgetTable.model_validate(cached_budget)
try:
budget_record: Final = await BudgetRepository(prisma_client).table.find_unique(where={"budget_id": budget_id})
budget_record: Final = await _dictable_table(BudgetRepository(prisma_client)).find_unique(
where={"budget_id": budget_id}
)
if budget_record is None:
verbose_proxy_logger.warning("Team-default member budget not found in database: %s", budget_id)
@ -1171,7 +1307,7 @@ async def get_end_user_object(
# Fetch from database
try:
response: Final = await EndUserRepository(prisma_client).table.find_unique(
response: Final = await _dictable_table(EndUserRepository(prisma_client)).find_unique(
where={"user_id": end_user_id},
include={"litellm_budget_table": True, "object_permission": True},
)
@ -1243,7 +1379,7 @@ async def resolve_and_validate_end_user_id(
return raw_end_user_id
cache_key: Final = f"end_user_validation:{raw_end_user_id}"
cached: Final = await user_api_key_cache.async_get_cache(key=cache_key)
cached: Final = await _raw_cache(user_api_key_cache).async_get_cache(key=cache_key)
if cached == "valid":
return raw_end_user_id
if cached == "invalid":
@ -1345,8 +1481,8 @@ async def get_tag_objects_batch(
if not tag_names:
return {}
tag_objects: Final = {}
uncached_tags: Final = []
tag_objects: Final = dict[str, LiteLLM_TagTable]()
uncached_tags: Final = list[str]()
# Try to get all tags from cache first
for tag_name in tag_names:
@ -1363,7 +1499,7 @@ async def get_tag_objects_batch(
# Batch fetch uncached tags from DB in one query
if uncached_tags:
try:
db_tags: Final = await TagRepository(prisma_client).table.find_many(
db_tags: Final = await _tag_table(TagRepository(prisma_client)).find_many(
where={"tag_name": {"in": uncached_tags}},
include={"litellm_budget_table": True},
)
@ -1457,7 +1593,7 @@ async def get_team_membership(
# else, check db
try:
response: Final = await TeamMembershipRepository(prisma_client).table.find_unique(
response: Final = await _dictable_table(TeamMembershipRepository(prisma_client)).find_unique(
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}},
include={"litellm_budget_table": True},
)
@ -1524,7 +1660,7 @@ def _should_check_db(key: str, last_db_access_time: LimitedSizeOrderedDict, db_c
return False
def _update_last_db_access_time(key: str, value: Any | None, last_db_access_time: LimitedSizeOrderedDict):
def _update_last_db_access_time(key: str, value: object | None, last_db_access_time: LimitedSizeOrderedDict):
last_db_access_time[key] = (value, time.time())
@ -1545,7 +1681,7 @@ def _get_role_based_permissions(
for role_based_permission in role_based_permissions:
if role_based_permission.role == rbac_role:
return getattr(role_based_permission, key)
return role_based_permission.models if key == "models" else role_based_permission.routes
return None
@ -1586,7 +1722,7 @@ async def _get_fuzzy_user_object(
prisma_client: PrismaClient,
sso_user_id: str | None = None,
user_email: str | None = None,
) -> LiteLLM_UserTable | None:
) -> "_PrismaUserRow | None":
"""
Checks if sso user is in db.
@ -1600,7 +1736,7 @@ async def _get_fuzzy_user_object(
response = None
if sso_user_id is not None:
response = await UserRepository(prisma_client).table.find_unique(
response = await _user_table(UserRepository(prisma_client)).find_unique(
where={"sso_user_id": sso_user_id},
include={"organization_memberships": True},
)
@ -1608,14 +1744,14 @@ async def _get_fuzzy_user_object(
if response is None and user_email is not None:
# Use case-insensitive query to handle emails with different casing
# This matches the pattern used in _check_duplicate_user_email
response = await UserRepository(prisma_client).table.find_first(
response = await _user_table(UserRepository(prisma_client)).find_first(
where={"user_email": {"equals": user_email, "mode": "insensitive"}},
include={"organization_memberships": True},
)
if response is not None and sso_user_id is not None: # update sso_user_id
asyncio.create_task( # background task to update user with sso id
UserRepository(prisma_client).table.update(
_user_table(UserRepository(prisma_client)).update(
where={"user_id": response.user_id},
data={"sso_user_id": sso_user_id},
)
@ -1698,7 +1834,7 @@ async def get_user_object(
)
if should_check_db:
response = await UserRepository(prisma_client).table.find_unique(
response = await _user_table(UserRepository(prisma_client)).find_unique(
where={"user_id": user_id}, include={"organization_memberships": True}
)
@ -1736,7 +1872,7 @@ async def get_user_object(
budget_duration=new_user_params["budget_duration"]
)
response = await UserRepository(prisma_client).table.create(
response = await _user_table(UserRepository(prisma_client)).create(
data=new_user_params,
include={"organization_memberships": True},
)
@ -1802,7 +1938,7 @@ async def get_user_object(
async def _cache_management_object(
key: str,
value: BaseModel | dict[str, Any],
value: BaseModel | Mapping[str, object],
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging | None,
*,
@ -1916,8 +2052,10 @@ async def _delete_cache_key_object(
@log_db_metrics
async def _get_team_db_check(team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None):
response = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
async def _get_team_db_check(
team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None
) -> "_PrismaTeamRow | None":
response = await _team_table(TeamRepository(prisma_client)).find_unique(where={"team_id": team_id})
if response is None and team_id_upsert:
from litellm.proxy.management_endpoints.team_endpoints import new_team
@ -1936,8 +2074,8 @@ async def _get_team_db_check(team_id: str, prisma_client: PrismaClient, team_id_
return response
async def _get_team_object_from_db(team_id: str, prisma_client: PrismaClient):
return await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
async def _get_team_object_from_db(team_id: str, prisma_client: PrismaClient) -> "_PrismaTeamRow | None":
return await _team_table(TeamRepository(prisma_client)).find_unique(where={"team_id": team_id})
async def _get_team_object_from_user_api_key_cache(
@ -2148,7 +2286,7 @@ async def get_access_object(
# Not in cache - fetch from DB
try:
response: Final = await AccessGroupRepository(prisma_client).table.find_unique(
response: Final = await _dictable_table(AccessGroupRepository(prisma_client)).find_unique(
where={"access_group_id": access_group_id}
)
@ -2224,7 +2362,7 @@ async def get_team_object_by_alias(
# Query database by team_alias
try:
teams: Final = await TeamRepository(prisma_client).table.find_many(where={"team_alias": team_alias})
teams: Final = await _team_table(TeamRepository(prisma_client)).find_many(where={"team_alias": team_alias})
if not teams:
raise HTTPException(
@ -2329,7 +2467,9 @@ async def get_org_object_by_alias(
# Query database by organization_alias
try:
orgs = await OrganizationRepository(prisma_client).table.find_many(where={"organization_alias": org_alias})
orgs = await _model_dump_table(OrganizationRepository(prisma_client)).find_many(
where={"organization_alias": org_alias}
)
if not orgs:
raise HTTPException(
@ -2546,7 +2686,7 @@ async def get_jwt_key_mapping_object(
Returns the hashed token (str) if a matching active mapping is found, else None.
"""
mapping: Final = await JWTKeyMappingRepository(prisma_client).table.find_first(
mapping: Final = await _jwt_key_mapping_table(JWTKeyMappingRepository(prisma_client)).find_first(
where={
"jwt_claim_name": jwt_claim_name,
"jwt_claim_value": jwt_claim_value,
@ -2674,7 +2814,7 @@ async def get_object_permission(
# else, check db
try:
response: Final = await ObjectPermissionRepository(prisma_client).table.find_unique(
response: Final = await _dictable_table(ObjectPermissionRepository(prisma_client)).find_unique(
where={"object_permission_id": object_permission_id}
)
@ -2730,7 +2870,7 @@ async def get_managed_vector_store_rows_by_uuids(
if not cache_misses:
return result
rows: Final = await ManagedVectorStoresRepository(prisma_client).table.find_many(
rows: Final = await _vector_store_table(ManagedVectorStoresRepository(prisma_client)).find_many(
where={"vector_store_id": {"in": cache_misses}},
take=len(cache_misses),
)
@ -2804,11 +2944,11 @@ async def get_org_object(
return deserialized_org
# else, check db
try:
query_kwargs: Final[dict[str, Any]] = {"where": {"organization_id": org_id}}
query_kwargs: Final[dict[str, Mapping[str, object]]] = {"where": {"organization_id": org_id}}
if include_budget_table:
query_kwargs["include"] = {"litellm_budget_table": True}
response: Final = await OrganizationRepository(prisma_client).table.find_unique(**query_kwargs)
response: Final = await _model_dump_table(OrganizationRepository(prisma_client)).find_unique(**query_kwargs)
except Exception:
# An operational failure (DB down, timeout, cache fault) is NOT the same fact as a confirmed
# missing row, and relabelling it as "doesn't exist" made every caller unable to tell them
@ -3763,7 +3903,7 @@ async def _virtual_key_soft_budget_check(
)
def _parse_email_list(raw: Any) -> list[str]:
def _parse_email_list(raw: str | Sequence[object] | None) -> list[str]:
"""Parse emails from a list or comma-separated string."""
if isinstance(raw, list):
return [e.strip() for e in raw if isinstance(e, str) and e.strip()]
@ -3773,7 +3913,7 @@ def _parse_email_list(raw: Any) -> list[str]:
def _normalize_alert_emails(
cfg: dict[str, Any] | None,
cfg: Mapping[str, str | Sequence[object] | None] | None,
) -> dict[str, list[str]]:
"""Coerce user-supplied threshold→recipients mapping to Dict[str, List[str]].
@ -3786,8 +3926,8 @@ def _normalize_alert_emails(
def _merge_budget_alert_email_configs(
global_cfg: dict[str, Any] | None,
per_key_cfg: dict[str, Any] | None,
global_cfg: Mapping[str, str | Sequence[object] | None] | None,
per_key_cfg: Mapping[str, str | Sequence[object] | None] | None,
) -> dict[str, list[str]] | None:
"""
Per-threshold additive merge: each threshold's recipient list is the union
@ -4294,7 +4434,7 @@ async def get_project_object(
return deserialized_project
# Fetch from DB
project_row: Final = await ProjectRepository(prisma_client).table.find_unique(
project_row: Final = await _model_dump_table(ProjectRepository(prisma_client)).find_unique(
where={"project_id": project_id},
include={"litellm_budget_table": True},
)
@ -4621,7 +4761,9 @@ async def vector_store_access_check(
#########################################################
# Check if the key can access the vector store
if valid_token is not None and valid_token.object_permission_id is not None:
key_object_permission: Final = await ObjectPermissionRepository(prisma_client).table.find_unique(
key_object_permission: Final = await _object_permission_table(
ObjectPermissionRepository(prisma_client)
).find_unique(
where={"object_permission_id": valid_token.object_permission_id},
)
if key_object_permission is not None:
@ -4633,7 +4775,9 @@ async def vector_store_access_check(
# Check if the team can access the vector store
if team_object is not None and team_object.object_permission_id is not None:
team_object_permission: Final = await ObjectPermissionRepository(prisma_client).table.find_unique(
team_object_permission: Final = await _object_permission_table(
ObjectPermissionRepository(prisma_client)
).find_unique(
where={"object_permission_id": team_object.object_permission_id},
)
if team_object_permission is not None:

View file

@ -1,12 +1,13 @@
import os
import re
import sys
from collections.abc import Iterator, Mapping
from collections.abc import Collection, Iterator, Mapping
from functools import lru_cache
from logging import Logger
from typing import Any, Final
from typing import Any, Final, Protocol
from fastapi import HTTPException, Request, status
from pydantic import PositiveInt, TypeAdapter, ValidationError
import litellm
from litellm import Router, provider_list
@ -261,6 +262,14 @@ _BANNED_REQUEST_BODY_PARAMS: Final[tuple[str, ...]] = (
"aws_sts_endpoint",
"aws_web_identity_token",
"aws_role_name",
# Remaining AWS identity selectors. ``get_credentials`` prefers a named
# profile over the deployment's static keys, so a caller-supplied
# ``aws_profile_name`` signs Bedrock and S3 requests as any profile
# present on the proxy host; the two AssumeRole knobs are banned with it
# so the whole identity-selection family lives behind the same opt-in.
"aws_profile_name",
"aws_session_name",
"aws_external_id",
"vertex_credentials",
# Azure managed-identity / federated-auth token. The Azure provider
# transformer reads ``azure_ad_token`` (top-level or via
@ -999,6 +1008,167 @@ def get_key_model_tpm_limit(
return None
ESTIMATED_OUTPUT_TOKENS_FIELD: Final = "default_estimated_output_tokens"
ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD: Final = "default_estimated_output_tokens_per_model"
ESTIMATED_OUTPUT_TOKENS_METADATA_FIELDS: Final = frozenset(
{ESTIMATED_OUTPUT_TOKENS_FIELD, ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD}
)
_ESTIMATED_OUTPUT_TOKENS_ADAPTER: Final = TypeAdapter(PositiveInt)
_ESTIMATED_OUTPUT_TOKENS_PER_MODEL_ADAPTER: Final = TypeAdapter(Mapping[str, PositiveInt])
def _validated_output_token_estimate(raw: object) -> int | None:
"""Coerce one declared estimate to a positive int, or ignore it."""
if raw is None:
return None
try:
return _ESTIMATED_OUTPUT_TOKENS_ADAPTER.validate_python(raw)
except ValidationError as validation_error:
verbose_proxy_logger.warning(
"Ignoring malformed %s in metadata: %s",
ESTIMATED_OUTPUT_TOKENS_FIELD,
validation_error,
)
return None
def _validated_output_token_estimates_per_model(raw: object) -> Mapping[str, int] | None:
"""Coerce a declared per-model estimate map, or ignore it."""
if raw is None:
return None
try:
return _ESTIMATED_OUTPUT_TOKENS_PER_MODEL_ADAPTER.validate_python(raw)
except ValidationError as validation_error:
verbose_proxy_logger.warning(
"Ignoring malformed %s in metadata: %s",
ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD,
validation_error,
)
return None
def _estimated_output_tokens_from_metadata(
metadata: Mapping[str, Any] | None,
model_name: str | None,
) -> int | None:
"""Resolve the per-model, then global, estimate out of one metadata blob.
The two fields are validated independently so a malformed per-model map
cannot discard a valid global estimate, or the other way round.
"""
if not metadata or ESTIMATED_OUTPUT_TOKENS_METADATA_FIELDS.isdisjoint(metadata):
return None
if model_name is not None:
per_model: Final = _validated_output_token_estimates_per_model(
metadata.get(ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD)
)
per_model_estimate: Final = per_model.get(model_name) if per_model is not None else None
if per_model_estimate is not None:
return per_model_estimate
return _validated_output_token_estimate(metadata.get(ESTIMATED_OUTPUT_TOKENS_FIELD))
def get_estimated_output_tokens(
user_api_key_dict: UserAPIKeyAuth,
model_name: str | None = None,
) -> int | None:
"""Resolve the operator-declared output-token estimate for TPM reservation.
Priority order (returns first found):
1. Key metadata ``default_estimated_output_tokens_per_model[model_name]``
2. Key metadata ``default_estimated_output_tokens``
3. Team metadata ``default_estimated_output_tokens_per_model[model_name]``
4. Team metadata ``default_estimated_output_tokens``
Returns ``None`` when nothing is configured, which leaves the static
heuristic floor in place.
"""
key_estimate: Final = _estimated_output_tokens_from_metadata(user_api_key_dict.metadata, model_name)
if key_estimate is not None:
return key_estimate
return _estimated_output_tokens_from_metadata(user_api_key_dict.team_metadata, model_name)
class OutputTokenEstimateRequest(Protocol):
"""The shape of any management request that can carry an output-token estimate.
Read-only members: the gate inspects a request, it never writes one back.
"""
@property
def metadata(self) -> Mapping[str, object] | None: ...
@property
def default_estimated_output_tokens(self) -> int | None: ...
@property
def default_estimated_output_tokens_per_model(self) -> Mapping[str, int] | None: ...
@property
def model_fields_set(self) -> Collection[str]: ...
def _requested_output_token_estimates(
data: OutputTokenEstimateRequest,
existing_metadata: Mapping[str, object],
) -> tuple[object, object]:
"""The output-token estimates this request would leave stored on the entity.
Mirrors how the management endpoints merge metadata: a supplied ``metadata``
replaces the stored blob wholesale, an omitted one preserves it, and the
dedicated top-level fields overlay whatever survives. Both sources are read
because the same declaration reaches the same stored field either way.
"""
base: Final[Mapping[str, object]] = (
(data.metadata or {}) if "metadata" in data.model_fields_set else existing_metadata
)
return (
data.default_estimated_output_tokens
if data.default_estimated_output_tokens is not None
else base.get(ESTIMATED_OUTPUT_TOKENS_FIELD),
data.default_estimated_output_tokens_per_model
if data.default_estimated_output_tokens_per_model is not None
else base.get(ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD),
)
def enforce_output_token_estimates_are_admin_only(
data: OutputTokenEstimateRequest,
existing_metadata: Mapping[str, object] | None,
user_api_key_dict: UserAPIKeyAuth,
entity: Literal["key", "team"],
) -> None:
"""Only a proxy admin may change what a key or team declares its models emit.
That declaration is what the TPM limiter reserves for a request omitting
``max_tokens``, so lowering or clearing it under-reserves against every
window the request is charged against, including the team and organization
ones the writer may not own. A key's metadata is writable by its holder and
a team's by its team admin, so neither is a trustworthy source for a value
that weakens a limit set above them. Gated on the resulting value rather
than on presence, so a form resending the stored declaration stays a no-op.
"""
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value:
return
stored: Final[Mapping[str, object]] = existing_metadata or {}
if _requested_output_token_estimates(data, stored) == (
stored.get(ESTIMATED_OUTPUT_TOKENS_FIELD),
stored.get(ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD),
):
return
raise HTTPException(
status_code=403,
detail={
"error": f"Only proxy admins can set {ESTIMATED_OUTPUT_TOKENS_FIELD} or "
f"{ESTIMATED_OUTPUT_TOKENS_PER_MODEL_FIELD} on a {entity}. They decide how many output tokens "
"the rate limiter reserves for a request that omits max_tokens."
},
)
def get_model_rate_limit_from_metadata(
user_api_key_dict: UserAPIKeyAuth,
metadata_accessor_key: Literal["team_metadata", "organization_metadata", "project_metadata"],
@ -1141,18 +1311,27 @@ def is_pass_through_provider_route(route: str) -> bool:
return False
def _has_user_setup_sso():
def _has_user_setup_sso() -> bool:
"""
Check if the user has set up single sign-on (SSO) by verifying the presence of Microsoft client ID, Google client ID or generic client ID and UI username environment variables.
Returns a boolean indicating whether SSO has been set up.
Check if the user has set up single sign-on (SSO).
Covers OAuth providers (Microsoft, Google, generic) and SAML IdP metadata.
Used by UI discovery (``sso_configured``) so the login button enables when
any supported SSO path is configured — including SAML-only setups.
"""
microsoft_client_id: Final = os.getenv("MICROSOFT_CLIENT_ID", None)
google_client_id: Final = os.getenv("GOOGLE_CLIENT_ID", None)
generic_client_id: Final = os.getenv("GENERIC_CLIENT_ID", None)
saml_idp_metadata_url: Final = os.getenv("SAML_IDP_METADATA_URL", None)
saml_idp_metadata_xml: Final = os.getenv("SAML_IDP_METADATA_XML", None)
sso_setup = (microsoft_client_id is not None) or (google_client_id is not None) or (generic_client_id is not None)
return sso_setup
return (
microsoft_client_id is not None
or google_client_id is not None
or generic_client_id is not None
or bool(saml_idp_metadata_url)
or bool(saml_idp_metadata_xml)
)
def get_customer_user_header_from_mapping(user_id_mapping) -> list | None:

View file

@ -11,6 +11,7 @@ import click
import requests
from rich.console import Console
from rich.table import Table
from typing_extensions import NotRequired, TypedDict
from litellm.constants import CLI_JWT_EXPIRATION_HOURS
from litellm.litellm_core_utils.cli_token_utils import is_cli_token_fresh
@ -18,6 +19,57 @@ from litellm.litellm_core_utils.cli_token_utils import is_cli_token_fresh
from .private_json import write_private_json
class CliTokenData(TypedDict):
base_url: str
key: str
user_id: str
user_email: str
user_role: str
auth_header_name: str
jwt_token: str
timestamp: float
class CliTeam(TypedDict, total=False):
team_id: str | None
team_alias: str | None
models: list[str]
max_budget: float | None
class CliContextObj(TypedDict):
base_url: str
base_url_explicit: NotRequired[bool]
class CliPollData(TypedDict, total=False):
status: str
key: str
user_id: str
teams: list[str]
team_details: object
requires_team_selection: bool
team_id: str
class CliPollRequestKwargs(TypedDict, total=False):
timeout: int
headers: dict[str, str]
class CliSsoStartData(TypedDict):
login_id: str
poll_secret: str
user_code: str
class CliAuthResult(TypedDict):
api_key: str
user_id: str | None
teams: list[str]
team_id: str | None
# Token storage utilities
def get_token_file_path() -> str:
"""Get the path to store the authentication token"""
@ -27,12 +79,12 @@ def get_token_file_path() -> str:
return str(config_dir / "token.json")
def save_token(token_data: dict[str, Any]) -> None:
def save_token(token_data: CliTokenData) -> None:
"""Save token data to file"""
write_private_json(get_token_file_path(), token_data)
def load_token() -> dict[str, Any] | None:
def load_token() -> CliTokenData | None:
"""Load token data from file"""
token_file: Final = get_token_file_path()
if not os.path.exists(token_file):
@ -65,7 +117,7 @@ def get_stored_api_key(expected_base_url: str | None = None) -> str | None:
# Team selection utilities
def display_teams_table(teams: list[dict[str, Any]]) -> None:
def display_teams_table(teams: list[CliTeam]) -> None:
"""Display teams in a formatted table"""
console: Final = Console()
@ -165,7 +217,7 @@ def display_interactive_team_selection(teams: list[dict[str, Any]], selected_ind
for i, team in enumerate(teams):
team_alias = team.get("team_alias") or "N/A"
team_id = team.get("team_id", "N/A")
models = team.get("models", [])
models: list[str] = team.get("models", [])
max_budget = team.get("max_budget")
# Format models list
@ -249,10 +301,11 @@ def prompt_team_selection_fallback(
while True:
try:
choice = click.prompt(
prompt_response: str = click.prompt(
"\nSelect a team by entering the index number (or 'skip' to continue without a team)",
type=str,
).strip()
)
choice = prompt_response.strip()
if choice.lower() == "skip":
return None
@ -275,7 +328,7 @@ def prompt_team_selection_fallback(
def _response_error_detail(response: requests.Response) -> str | None:
try:
body: Final = response.json()
body: Final[dict[str, object] | list[object] | str | int | float | bool | None] = response.json()
except ValueError:
return None
detail: Final = body.get("detail") if isinstance(body, dict) else None
@ -309,15 +362,15 @@ def _poll_for_ready_data(
other_status_log_every: int = 10,
http_error_log_every: int = 10,
connection_error_log_every: int = 10,
) -> dict[str, Any] | None:
) -> CliPollData | None:
for attempt in range(total_timeout // poll_interval):
try:
request_kwargs: dict[str, Any] = {"timeout": request_timeout}
request_kwargs: CliPollRequestKwargs = {"timeout": request_timeout}
if headers is not None:
request_kwargs["headers"] = headers
response = requests.get(url, **request_kwargs)
if response.status_code == 200:
data = response.json()
data: CliPollData = response.json()
status = data.get("status")
if status == "ready":
return data
@ -341,7 +394,7 @@ def _poll_for_ready_data(
return None
def _normalize_teams(teams, team_details):
def _normalize_teams(teams: object, team_details: object) -> list[CliTeam]:
"""If team_details are a
Args:
@ -365,7 +418,7 @@ def _normalize_teams(teams, team_details):
return []
def _start_cli_sso_flow(base_url: str) -> dict[str, Any]:
def _start_cli_sso_flow(base_url: str) -> CliSsoStartData:
start_url: Final = f"{base_url}/sso/cli/start"
try:
response: Final = requests.post(start_url, timeout=10)
@ -389,7 +442,7 @@ def _start_cli_sso_flow(base_url: str) -> dict[str, Any]:
)
try:
data: Final = response.json()
data: Final[CliSsoStartData] = response.json()
except ValueError:
content_type: Final = response.headers.get("content-type", "unknown")
raise ValueError(
@ -398,7 +451,7 @@ def _start_cli_sso_flow(base_url: str) -> dict[str, Any]:
f"Response starts with: {response.text[:200]!r}"
)
required_fields: Final = ("login_id", "poll_secret", "user_code")
required_fields: Final[tuple[str, ...]] = ("login_id", "poll_secret", "user_code")
missing_fields: Final = tuple(field for field in required_fields if not isinstance(data.get(field), str))
if missing_fields:
raise ValueError(
@ -412,7 +465,7 @@ def _get_cli_sso_poll_headers(poll_secret: str) -> dict[str, str]:
return {"x-litellm-cli-poll-secret": poll_secret}
def _poll_for_authentication(base_url: str, key_id: str, poll_secret: str) -> dict | None:
def _poll_for_authentication(base_url: str, key_id: str, poll_secret: str) -> CliAuthResult | None:
"""
Poll the server for authentication completion and handle team selection.
@ -431,7 +484,7 @@ def _poll_for_authentication(base_url: str, key_id: str, poll_secret: str) -> di
teams = data.get("teams", [])
team_details: Final = data.get("team_details")
user_id = data.get("user_id")
normalized_teams: Final[list[dict[str, Any]]] = _normalize_teams(teams, team_details)
normalized_teams: Final[list[CliTeam]] = _normalize_teams(teams, team_details)
if not normalized_teams:
click.echo("Warning: No teams available for selection.")
return None
@ -478,7 +531,7 @@ def _poll_for_authentication(base_url: str, key_id: str, poll_secret: str) -> di
def _handle_team_selection_during_polling(
base_url: str, key_id: str, poll_secret: str, teams: list[dict[str, Any]]
base_url: str, key_id: str, poll_secret: str, teams: list[CliTeam]
) -> str | None:
"""
Handle team selection and re-poll with selected team_id.
@ -522,7 +575,7 @@ def _handle_team_selection_during_polling(
return None
def _render_and_prompt_for_team_selection(teams: list[dict[str, Any]]) -> str | None:
def _render_and_prompt_for_team_selection(teams: list[CliTeam]) -> str | None:
"""Render teams table and prompt user for a team selection.
Returns the selected team_id as a string, or None if selection was
@ -546,10 +599,11 @@ def _render_and_prompt_for_team_selection(teams: list[dict[str, Any]]) -> str |
# Simple selection
while True:
try:
choice = click.prompt(
prompt_response: str = click.prompt(
"\nSelect a team by entering the index number (or 'skip' to use first team)",
type=str,
).strip()
)
choice = prompt_response.strip()
if choice.lower() == "skip":
# Default to the first team's ID if the user skips an
@ -582,7 +636,8 @@ def login(ctx: click.Context):
from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER
from litellm.proxy.client.cli.interface import show_commands
base_url: Final = ctx.obj["base_url"]
ctx_obj: Final[CliContextObj] = ctx.obj
base_url: Final = ctx_obj["base_url"]
try:
cli_sso_flow: Final = _start_cli_sso_flow(base_url=base_url)
@ -675,8 +730,9 @@ def print_token(ctx: click.Context):
# explicitly pointed us at a server, trust whichever one `lite login`
# actually issued this token for -- that's the whole point of not
# needing a wrapper command.
if ctx.obj.get("base_url_explicit"):
base_url: Final = ctx.obj["base_url"]
ctx_obj: Final[CliContextObj] = ctx.obj
if ctx_obj.get("base_url_explicit"):
base_url: Final = ctx_obj["base_url"]
if token_data.get("base_url") != base_url.rstrip("/"):
click.echo("Not authenticated for this server. Run 'lite login'.", err=True)
sys.exit(1)

View file

@ -7,7 +7,8 @@ import traceback
from collections.abc import AsyncGenerator, Callable, Mapping
from datetime import datetime
from functools import lru_cache
from typing import TYPE_CHECKING, Any, Final, Literal
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, overload
import anyio
import httpx
@ -311,8 +312,49 @@ def _stream_usage_tracking_updates(
}
def _getattr_object(value: object, name: str, default: object = None) -> object:
return getattr(value, name, default)
class _UpstreamHttpResponse(Protocol):
@property
def status_code(self) -> int: ...
@property
def headers(self) -> httpx.Headers: ...
async def aread(self) -> bytes: ...
def _as_upstream_response(response: _UpstreamHttpResponse) -> _UpstreamHttpResponse:
return response
class _ReadsHeaderValues(Protocol):
def get(self, key: str, default: str = "") -> str: ...
def _as_header_reader(headers: _ReadsHeaderValues) -> _ReadsHeaderValues:
return headers
class _DispatchesSuccessHandlers(Protocol):
async def dispatch_success_handlers(
self,
result: object = None,
start_time: object = None,
end_time: object = None,
cache_hit: object = None,
prefer_async_handlers: bool = False,
) -> None: ...
def _as_success_dispatcher(logging_obj: _DispatchesSuccessHandlers) -> _DispatchesSuccessHandlers:
return logging_obj
def _serialize_http_exception_detail(
detail: Any,
detail: object,
) -> tuple[str, dict | None]:
"""
Convert an HTTPException.detail value into (message, structured_fields)
@ -342,7 +384,7 @@ def _serialize_http_exception_detail(
return str(detail), None
def _collect_response_file_search_vector_store_ids(data: dict[str, Any]) -> set[str]:
def _collect_response_file_search_vector_store_ids(data: Mapping[str, object]) -> set[str]:
vector_store_ids: Final[set[str]] = set()
tools: Final = data.get("tools")
if not isinstance(tools, list):
@ -369,7 +411,7 @@ def _collect_response_file_search_vector_store_ids(data: dict[str, Any]) -> set[
async def _authorize_response_file_search_vector_stores(
data: dict[str, Any],
data: Mapping[str, object],
user_api_key_dict: UserAPIKeyAuth,
) -> None:
vector_store_ids: Final = _collect_response_file_search_vector_store_ids(data)
@ -700,7 +742,7 @@ async def create_response(
# Preserve status code from HTTPException (e.g., guardrail blocks)
error_status: Final = getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR)
raw_detail: Final = getattr(e, "detail", "Error processing stream start")
raw_detail: Final = _getattr_object(e, "detail", "Error processing stream start")
message, structured_fields = _serialize_http_exception_detail(raw_detail)
existing_fields: Final = getattr(e, "provider_specific_fields", None) or {}
@ -711,7 +753,7 @@ async def create_response(
# Match ProxyException.to_dict() shape so streaming and non-streaming
# error frames are byte-identical.
error_obj: Final[dict[str, Any]] = {
error_obj: Final[dict[str, object]] = {
"message": message,
"type": getattr(e, "type", "None"),
"param": getattr(e, "param", "None"),
@ -777,7 +819,7 @@ def _is_azure_model_router_request(model: str) -> bool:
def _override_openai_response_model(
*,
response_obj: Any,
response_obj: object,
requested_model: str,
log_context: str,
return_raw_model_name: bool = False,
@ -972,7 +1014,7 @@ def _log_llm_api_exception(e: Exception) -> None:
async def _cancel_llm_call_on_client_disconnect(
request: Request,
llm_api_call: "asyncio.Future[Any]",
llm_api_call: "asyncio.Future[object]",
disconnect_event: asyncio.Event,
) -> None:
try:
@ -1023,7 +1065,7 @@ class ProxyBaseLLMRequestProcessing:
version: str | None = None,
model_region: str | None = None,
response_cost: float | str | None = None,
hidden_params: dict | None = None,
hidden_params: Mapping[str, object] | None = None,
fastest_response_batch_completion: bool | None = None,
request_data: dict | None = {},
timeout: float | httpx.Timeout | None = None,
@ -1115,7 +1157,7 @@ class ProxyBaseLLMRequestProcessing:
@staticmethod
async def build_litellm_proxy_success_headers_from_llm_response(
*,
response: Any,
response: object,
request_data: dict,
request: Request,
user_api_key_dict: UserAPIKeyAuth,
@ -1906,7 +1948,7 @@ class ProxyBaseLLMRequestProcessing:
_captured_user_api_key_dict: Final = user_api_key_dict
_captured_logging_obj: Final = logging_obj
async def _on_deferred_stream_complete(assembled_response, cache_hit):
async def _on_deferred_stream_complete(assembled_response: object, cache_hit: object) -> None:
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data=_captured_data,
captured_user_api_key_dict=_captured_user_api_key_dict,
@ -2157,7 +2199,7 @@ class ProxyBaseLLMRequestProcessing:
@staticmethod
async def _record_container_owners_from_responses_if_needed(
response: Any,
response: object,
user_api_key_dict: UserAPIKeyAuth,
) -> None:
"""Register code-interpreter containers so follow-up file APIs pass ownership checks."""
@ -2180,7 +2222,7 @@ class ProxyBaseLLMRequestProcessing:
)
@staticmethod
def _extract_completed_responses_response(stream_response: Any) -> Any:
def _extract_completed_responses_response(stream_response: object) -> object:
"""Pull the assembled ``ResponsesAPIResponse`` off a streaming iterator.
``ResponsesAPIStreamingIterator`` stores the terminal stream event
@ -2190,17 +2232,17 @@ class ProxyBaseLLMRequestProcessing:
``ResponsesAPIResponse`` directly. Handle both shapes so the
container-ownership recording path can walk ``.output`` either way.
"""
completed: Final = getattr(stream_response, "completed_response", None)
completed: Final = _getattr_object(stream_response, "completed_response")
if completed is None:
return None
response_obj: Final = getattr(completed, "response", None)
response_obj: Final = _getattr_object(completed, "response")
if response_obj is not None:
return response_obj
return completed
@staticmethod
async def _wrap_responses_stream_for_container_ownership(
original_stream_response: Any,
original_stream_response: object,
wrapped_generator: Any,
user_api_key_dict: UserAPIKeyAuth,
):
@ -2299,12 +2341,13 @@ class ProxyBaseLLMRequestProcessing:
if isinstance(result, Response):
return result
content: Final = await result.aread()
upstream: Final = _as_upstream_response(result)
content: Final = await upstream.aread()
return Response(
content=content,
status_code=result.status_code,
status_code=upstream.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(
headers=result.headers,
headers=upstream.headers,
custom_headers=dict(fastapi_response.headers),
),
)
@ -2435,9 +2478,10 @@ class ProxyBaseLLMRequestProcessing:
HttpPassThroughEndpointHelpers,
)
upstream: Final = _as_upstream_response(response)
try:
response_status: Final[int] = response.status_code
content_type: Final[str] = response.headers.get("content-type", "")
response_status: Final[int] = upstream.status_code
content_type: Final[str] = _as_header_reader(upstream.headers).get("content-type", "")
except AttributeError:
return None
@ -2451,20 +2495,20 @@ class ProxyBaseLLMRequestProcessing:
return None
response_headers: Final = HttpPassThroughEndpointHelpers.get_response_headers(
headers=response.headers,
headers=upstream.headers,
custom_headers=custom_headers,
)
callback_headers: Final = await proxy_logging_obj.post_call_response_headers_hook(
data=self.data,
user_api_key_dict=user_api_key_dict,
response=response,
response=upstream,
request_headers=request_headers,
)
if callback_headers:
response_headers.update(callback_headers)
if is_event_stream:
body_bytes = await response.aread()
body_bytes = await upstream.aread()
modified_bytes: Final = await self._handle_event_stream_allm_passthrough_route(
body_bytes=body_bytes,
proxy_logging_obj=proxy_logging_obj,
@ -2477,7 +2521,7 @@ class ProxyBaseLLMRequestProcessing:
headers=response_headers,
)
body_bytes = await response.aread()
body_bytes = await upstream.aread()
try:
parsed: Final = _json.loads(body_bytes)
except (_json.JSONDecodeError, UnicodeDecodeError):
@ -2566,9 +2610,9 @@ class ProxyBaseLLMRequestProcessing:
async def _run_deferred_stream_guardrails(
captured_data: dict,
captured_user_api_key_dict: "UserAPIKeyAuth",
captured_logging_obj: Any,
captured_logging_obj: LiteLLMLoggingObj,
assembled_response: Any,
cache_hit: Any,
cache_hit: object,
) -> None:
"""
Run non-streaming post-call guardrail hooks on an assembled streaming
@ -2646,7 +2690,7 @@ class ProxyBaseLLMRequestProcessing:
# _is_sync_litellm_request (which only recognizes a subset of
# async markers stored in litellm_params).
asyncio.create_task(
captured_logging_obj.dispatch_success_handlers(
_as_success_dispatcher(captured_logging_obj).dispatch_success_handlers(
_response,
cache_hit=cache_hit,
start_time=None,
@ -2717,7 +2761,7 @@ class ProxyBaseLLMRequestProcessing:
headers = getattr(e, "headers", None) or {}
if not headers:
# Try to get headers from e.response.headers (httpx.Response)
_response: Final = getattr(e, "response", None)
_response: Final = _getattr_object(e, "response")
if _response is not None:
_response_headers: Final = getattr(_response, "headers", None)
if _response_headers:
@ -2749,7 +2793,7 @@ class ProxyBaseLLMRequestProcessing:
raise e
if isinstance(e, HTTPException):
raw_detail: Final = getattr(e, "detail", str(e))
raw_detail: Final = _getattr_object(e, "detail", str(e))
message, structured_fields = _serialize_http_exception_detail(raw_detail)
existing_fields: Final = getattr(e, "provider_specific_fields", None) or {}
if structured_fields:
@ -3042,8 +3086,16 @@ class ProxyBaseLLMRequestProcessing:
request=request,
)
@overload
@staticmethod
def _process_chunk_with_cost_injection(chunk: Any, model_name: str) -> Any:
def _process_chunk_with_cost_injection(chunk: bytes, model_name: str) -> bytes: ...
@overload
@staticmethod
def _process_chunk_with_cost_injection(chunk: object, model_name: str) -> object: ...
@staticmethod
def _process_chunk_with_cost_injection(chunk: object, model_name: str) -> object:
"""
Process a streaming chunk and inject cost information if enabled.
@ -3063,12 +3115,12 @@ class ProxyBaseLLMRequestProcessing:
if maybe_modified is not None:
return maybe_modified
elif isinstance(chunk, (bytes, bytearray)):
# Decode to str, inject, and rebuild as bytes
try:
s: Final = chunk.decode("utf-8", errors="ignore")
maybe_mod = ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str(s, model_name)
if maybe_mod is not None:
return (maybe_mod + ("" if maybe_mod.endswith("\n\n") else "\n\n")).encode("utf-8")
s: Final = chunk.decode("utf-8")
if s.endswith(("\n\n", "\r\n\r\n")):
maybe_mod = ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str(s, model_name)
if maybe_mod is not None:
return maybe_mod.encode("utf-8")
except Exception:
pass
elif isinstance(chunk, str):
@ -3106,17 +3158,85 @@ class ProxyBaseLLMRequestProcessing:
obj = json.loads(json_part)
maybe_modified = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(obj, model_name)
if maybe_modified is not None:
# Replace just this line with updated JSON using safe_dumps
lines[idx] = f"data: {safe_dumps(maybe_modified)}"
lines[idx] = "data: " + safe_dumps(maybe_modified) + ("\r" if ln.endswith("\r") else "")
return "\n".join(lines)
return None
except Exception:
return None
@staticmethod
def _anthropic_stream_usage_kwargs(usage: Mapping[str, Any]) -> Mapping[str, Any]:
prompt_tokens: Final = int(usage.get("input_tokens", 0) or 0)
completion_tokens: Final = int(usage.get("output_tokens", 0) or 0)
total_tokens: Final = int(
usage.get("total_tokens", prompt_tokens + completion_tokens) or (prompt_tokens + completion_tokens)
)
web_search_requests: Final = usage.get("web_search_requests")
server_tool_use: Final = (
ServerToolUse(web_search_requests=web_search_requests) if web_search_requests is not None else None
)
return MappingProxyType(
{
key: value
for key, value in (
("prompt_tokens", prompt_tokens),
("completion_tokens", completion_tokens),
("total_tokens", total_tokens),
("completion_tokens_details", usage.get("completion_tokens_details")),
("prompt_tokens_details", usage.get("prompt_tokens_details")),
("cache_creation_input_tokens", usage.get("cache_creation_input_tokens")),
("cache_read_input_tokens", usage.get("cache_read_input_tokens")),
("server_tool_use", server_tool_use),
)
if value is not None
}
)
@staticmethod
def _openai_stream_usage_kwargs(usage: Mapping[str, Any]) -> Mapping[str, Any]:
prompt_tokens: Final = int(usage.get("prompt_tokens", 0) or 0)
completion_tokens: Final = int(usage.get("completion_tokens", 0) or 0)
total_tokens: Final = int(
usage.get("total_tokens", prompt_tokens + completion_tokens) or (prompt_tokens + completion_tokens)
)
return MappingProxyType(
{
key: value
for key, value in (
("prompt_tokens", prompt_tokens),
("completion_tokens", completion_tokens),
("total_tokens", total_tokens),
("completion_tokens_details", usage.get("completion_tokens_details")),
("prompt_tokens_details", usage.get("prompt_tokens_details")),
)
if value is not None
}
)
@staticmethod
def _stream_usage_kwargs_for_event(obj: Mapping[str, object], usage: Mapping[str, Any]) -> Mapping[str, Any] | None:
if obj.get("type") == "message_delta":
return ProxyBaseLLMRequestProcessing._anthropic_stream_usage_kwargs(usage)
if obj.get("object") == "chat.completion.chunk":
return ProxyBaseLLMRequestProcessing._openai_stream_usage_kwargs(usage)
return None
@staticmethod
def _completion_cost_or_none(
model_response: ModelResponse, model_name: str, service_tier: str | None
) -> float | None:
try:
return litellm.completion_cost(
completion_response=model_response, model=model_name, service_tier=service_tier
)
except Exception:
return None
@staticmethod
def _inject_cost_into_usage_dict(obj: dict, model_name: str) -> dict | None:
"""
Inject cost information into a usage dictionary for message_delta events.
Inject cost information into the usage object of a streamed usage event
(Anthropic ``message_delta`` or OpenAI ``chat.completion.chunk``).
Args:
obj: Dictionary containing the SSE event data
@ -3125,57 +3245,21 @@ class ProxyBaseLLMRequestProcessing:
Returns:
Modified dictionary with cost injected, or None if no modification needed
"""
if obj.get("type") == "message_delta" and isinstance(obj.get("usage"), dict):
_usage: Final = obj["usage"]
prompt_tokens: Final = int(_usage.get("input_tokens", 0) or 0)
completion_tokens: Final = int(_usage.get("output_tokens", 0) or 0)
total_tokens: Final = int(
_usage.get("total_tokens", prompt_tokens + completion_tokens) or (prompt_tokens + completion_tokens)
)
# Extract additional usage fields
cache_creation_input_tokens: Final = _usage.get("cache_creation_input_tokens")
cache_read_input_tokens: Final = _usage.get("cache_read_input_tokens")
web_search_requests: Final = _usage.get("web_search_requests")
completion_tokens_details: Final = _usage.get("completion_tokens_details")
prompt_tokens_details: Final = _usage.get("prompt_tokens_details")
usage_kwargs: Final[dict[str, Any]] = {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": total_tokens,
}
# Add optional named parameters
if completion_tokens_details is not None:
usage_kwargs["completion_tokens_details"] = completion_tokens_details
if prompt_tokens_details is not None:
usage_kwargs["prompt_tokens_details"] = prompt_tokens_details
# Handle web_search_requests by wrapping in ServerToolUse
if web_search_requests is not None:
usage_kwargs["server_tool_use"] = ServerToolUse(web_search_requests=web_search_requests)
# Add cache-related fields to **params (handled by Usage.__init__)
if cache_creation_input_tokens is not None:
usage_kwargs["cache_creation_input_tokens"] = cache_creation_input_tokens
if cache_read_input_tokens is not None:
usage_kwargs["cache_read_input_tokens"] = cache_read_input_tokens
_mr: Final = ModelResponse(usage=Usage(**usage_kwargs))
try:
cost_val = litellm.completion_cost(
completion_response=_mr,
model=model_name,
)
except Exception:
cost_val = None
if cost_val is not None:
obj.setdefault("usage", {})["cost"] = cost_val
return obj
return None
usage: Final = obj.get("usage")
if not isinstance(usage, dict):
return None
usage_kwargs: Final = ProxyBaseLLMRequestProcessing._stream_usage_kwargs_for_event(obj, usage)
if usage_kwargs is None:
return None
service_tier: Final = obj.get("service_tier")
cost_val: Final = ProxyBaseLLMRequestProcessing._completion_cost_or_none(
ModelResponse(usage=Usage(**usage_kwargs)),
model_name,
service_tier if isinstance(service_tier, str) else None,
)
if cost_val is None:
return None
return {**obj, "usage": {**usage, "cost": cost_val}}
def maybe_get_model_id(self, _logging_obj: LiteLLMLoggingObj | None) -> str | None:
"""

View file

@ -1,14 +1,21 @@
import asyncio
import json
import time
from collections.abc import Callable, Sequence
from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Final, Literal, Protocol, TypeVar
from types import MappingProxyType
from typing import Final, Literal, Protocol, TypeVar, assert_never
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.caching.dual_cache import DualCache
from litellm.constants import GLOBAL_PROXY_SPEND_CACHE_KEY, LITELLM_PROXY_BUDGET_NAME
from litellm.constants import (
GLOBAL_PROXY_SPEND_CACHE_KEY,
LITELLM_PROXY_BUDGET_NAME,
RESET_BUDGET_JOB_BATCH_SIZE,
RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN,
)
from litellm.proxy._types import (
LiteLLM_BudgetTableFull,
LiteLLM_EndUserTable,
@ -30,7 +37,10 @@ from litellm.repositories.table_repositories import (
TeamMembershipRepository,
)
from litellm.repositories.team_repository import TeamRepository
from litellm.repositories.unit_of_work import spend_reset_unit_of_work
from litellm.repositories.unit_of_work import (
budget_cascade_unit_of_work,
spend_reset_unit_of_work,
)
from litellm.repositories.verification_token_repository import (
VerificationTokenRepository,
)
@ -38,6 +48,9 @@ from litellm.types.services import ServiceTypes
_RowT = TypeVar("_RowT")
_LINKED_KEYS_WHERE: Final[Mapping[str, object]] = MappingProxyType({"budget_duration": None, "spend": {"gt": 0}})
_SPENT_ROWS_WHERE: Final[Mapping[str, object]] = MappingProxyType({"spend": {"gt": 0}})
class _TeamMembershipRow(Protocol):
@property
@ -62,39 +75,130 @@ class _TagRow(Protocol):
def tag_name(self) -> str: ...
class _EndUserRow(Protocol):
@property
def user_id(self) -> str: ...
def _team_membership_counter_key(row: _TeamMembershipRow) -> str:
return f"spend:team_member:{row.user_id}:{row.team_id}"
def _team_membership_cache_key(row: _TeamMembershipRow) -> str:
return f"{row.team_id}_{row.user_id}"
def _team_membership_cache_keys(row: _TeamMembershipRow) -> tuple[str, ...]:
return (f"{row.team_id}_{row.user_id}",)
def _key_counter_key(row: _KeyRow) -> str:
return f"spend:key:{row.token}"
def _key_cache_key(row: _KeyRow) -> str:
return row.token
def _key_cache_keys(row: _KeyRow) -> tuple[str, ...]:
return (row.token,)
def _org_counter_key(row: _OrgRow) -> str:
return f"spend:org:{row.organization_id}"
def _org_cache_keys(row: _OrgRow) -> Sequence[str]:
return [
def _org_cache_keys(row: _OrgRow) -> tuple[str, ...]:
return (
f"org_id:{row.organization_id}",
f"org_id:{row.organization_id}:with_budget",
]
)
def _tag_counter_key(row: _TagRow) -> str:
return f"spend:tag:{row.tag_name}"
def _tag_cache_key(row: _TagRow) -> str:
return f"tag:{row.tag_name}"
def _tag_cache_keys(row: _TagRow) -> tuple[str, ...]:
return (f"tag:{row.tag_name}",)
def _budget_link_where(
budget_ids: Sequence[str],
extra: Mapping[str, object] = MappingProxyType({}),
) -> dict[str, object]:
return {"budget_id": {"in": list(budget_ids)}, **extra}
@dataclass(frozen=True, slots=True)
class _BudgetCascade:
"""Everything one budget-tier reset touches, resolved before any write."""
budgets: tuple[LiteLLM_BudgetTableFull, ...] = ()
budget_ids: tuple[str, ...] = ()
budget_resets: tuple[tuple[str, datetime], ...] = ()
endusers: tuple[_EndUserRow, ...] = ()
counter_keys: tuple[str, ...] = ()
cache_keys: tuple[str, ...] = ()
@dataclass(frozen=True, slots=True)
class _BudgetCascadeCommitted:
cascade: _BudgetCascade
advanced: int
@dataclass(frozen=True, slots=True)
class _BudgetCascadeFailed:
cascade: _BudgetCascade
error: Exception
_EMPTY_CASCADE: Final = _BudgetCascade()
@dataclass(frozen=True, slots=True)
class _ChunkOutcome:
"""One chunk of a reset phase: rows read, and rows whose new budget_reset_at
cleared the due cutoff. Anything else is still due and would come straight
back on the next fetch, so it is not progress."""
fetched: int
advanced: int
_NO_PROGRESS: Final = _ChunkOutcome(fetched=0, advanced=0)
def _as_utc(moment: datetime) -> datetime:
return moment if moment.tzinfo is not None else moment.replace(tzinfo=timezone.utc)
def _count_advanced(reset_ats: Iterable[object], cutoff: datetime) -> int:
"""How many rows the write actually moved past the due cutoff.
A budget_duration of "0s" (or one the parser cannot read) resolves to the
current time, so the row is written and stays due. Counting it as progress
would re-read the same chunk until the per-run cap on every tick.
"""
utc_cutoff: Final = _as_utc(cutoff)
return sum(1 for reset_at in reset_ats if isinstance(reset_at, datetime) and _as_utc(reset_at) > utc_cutoff)
def _phase_is_drained(outcome: _ChunkOutcome) -> bool:
"""A short chunk means the due rows ran out. A full chunk that advanced
nothing would be re-read unchanged forever, so it ends the phase too and
those rows wait for the next tick."""
return outcome.fetched < RESET_BUDGET_JOB_BATCH_SIZE or outcome.advanced == 0
async def _run_phase_in_chunks(process_chunk: Callable[[], Awaitable[_ChunkOutcome]]) -> None:
"""Drive one reset phase a chunk at a time, capped so a single run cannot
spin unbounded: leftovers are picked up by the next tick."""
for _ in range(RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN):
if _phase_is_drained(await process_chunk()):
return
def _budget_cascade_event_metadata(cascade: _BudgetCascade) -> dict[str, object]:
return {
"num_budgets_found": len(cascade.budgets),
"budgets_found": json.dumps(cascade.budgets, indent=4, default=str),
"num_endusers_found": len(cascade.endusers),
"endusers_found": json.dumps(cascade.endusers, indent=4, default=str),
}
class ResetBudgetJob:
@ -122,21 +226,14 @@ class ResetBudgetJob:
Updates db
"""
if self.prisma_client is not None:
### RESET KEY BUDGET ###
await self.reset_budget_for_litellm_keys()
if self.prisma_client is None:
return
### RESET USER BUDGET ###
await self.reset_budget_for_litellm_users()
## Reset Team Budget
await self.reset_budget_for_litellm_teams()
### RESET ENDUSER (Customer) BUDGET and corresponding Budget duration ###
await self.reset_budget_for_litellm_budget_table()
### RESET MULTI-WINDOW BUDGETS ###
await self.reset_budget_windows()
await self.reset_budget_for_litellm_keys()
await self.reset_budget_for_litellm_users()
await self.reset_budget_for_litellm_teams()
await self.reset_budget_for_litellm_budget_table()
await self.reset_budget_windows()
@staticmethod
async def _invalidate_spend_counter(counter_key: str) -> None:
@ -194,238 +291,195 @@ class ResetBudgetJob:
e,
)
async def _cascade_reset_spend_for_budget_link(
async def _fetch_linked_rows(
self,
budgets_to_reset: list[LiteLLM_BudgetTableFull],
table: SpendLinkedTable[_RowT],
counter_key_fn: Callable[[_RowT], str],
where: Mapping[str, object],
log_subject: str,
extra_where: dict[str, object] | None = None,
cache_key_fn: Callable[[_RowT], str | Sequence[str]] | None = None,
):
"""
Generic cascade: zero spend on rows whose budget_id is in the reset set.
) -> tuple[_RowT, ...]:
"""Read the rows the cascade will zero, so their counters can be
invalidated once the transaction commits."""
try:
return tuple(await table.find_many(where=where))
except Exception as e:
verbose_proxy_logger.warning("Failed to fetch %s for counter invalidation: %s", log_subject, e)
return ()
``cache_key_fn`` is optional: when provided, after the DB update each
matching row's entry or entries in ``user_api_key_cache`` are dropped so
cached spend cannot stay pinned above the zeroed DB row after a reset.
async def _collect_endusers_to_reset(self, budget_ids: Sequence[str]) -> tuple[_EndUserRow, ...]:
linked: Final[Sequence[_EndUserRow] | None] = await self.prisma_client.get_data(
table_name="enduser",
query_type="find_all",
budget_id_list=list(budget_ids),
)
if litellm.max_end_user_budget_id is None or litellm.max_end_user_budget_id not in budget_ids:
return tuple(linked or ())
return (*(linked or ()), *await self._get_endusers_with_no_budget_id())
async def _collect_budget_cascade(self, budgets_to_reset: Sequence[LiteLLM_BudgetTableFull]) -> _BudgetCascade:
"""Resolve every row the expiring budget tiers gate, before any write.
Keys carrying their own budget_duration are left out: they run on their
own schedule via reset_budget_for_litellm_keys(), so sweeping them here
would reset them twice.
"""
budget_ids: Final = [b.budget_id for b in budgets_to_reset if b.budget_id is not None]
budget_ids: Final = tuple(b.budget_id for b in budgets_to_reset if b.budget_id is not None)
if not budget_ids:
return _EMPTY_CASCADE
team_memberships: Final[tuple[_TeamMembershipRow, ...]] = await self._fetch_linked_rows(
table=TeamMembershipRepository(self.prisma_client).table,
where=_budget_link_where(budget_ids),
log_subject="team memberships",
)
keys: Final[tuple[_KeyRow, ...]] = await self._fetch_linked_rows(
table=VerificationTokenRepository(self.prisma_client).table,
where=_budget_link_where(budget_ids, _LINKED_KEYS_WHERE),
log_subject="keys",
)
orgs: Final[tuple[_OrgRow, ...]] = await self._fetch_linked_rows(
table=OrganizationRepository(self.prisma_client).table,
where=_budget_link_where(budget_ids, _SPENT_ROWS_WHERE),
log_subject="orgs",
)
tags: Final[tuple[_TagRow, ...]] = await self._fetch_linked_rows(
table=TagRepository(self.prisma_client).table,
where=_budget_link_where(budget_ids, _SPENT_ROWS_WHERE),
log_subject="tags",
)
return _BudgetCascade(
budgets=tuple(budgets_to_reset),
budget_ids=budget_ids,
budget_resets=tuple(
(
b.budget_id,
compute_budget_reset_at(budget_duration=b.budget_duration, settings=self.reset_settings),
)
for b in budgets_to_reset
if b.budget_id is not None and b.budget_duration is not None
),
endusers=await self._collect_endusers_to_reset(budget_ids),
counter_keys=(
*(_team_membership_counter_key(row) for row in team_memberships),
*(_key_counter_key(row) for row in keys),
*(_org_counter_key(row) for row in orgs),
*(_tag_counter_key(row) for row in tags),
),
cache_keys=(
*(key for row in team_memberships for key in _team_membership_cache_keys(row)),
*(key for row in keys for key in _key_cache_keys(row)),
*(key for row in orgs for key in _org_cache_keys(row)),
*(key for row in tags for key in _tag_cache_keys(row)),
),
)
async def _commit_budget_cascade(self, cascade: _BudgetCascade) -> None:
"""Zero the gated spend and advance ``budget_reset_at`` in one transaction.
Advancing the window on its own hides the tier from every later tick
while its dependents stay pinned at the cap for the whole window;
batching both means a mid-cascade failure persists nothing and the rows
stay due for the next run.
"""
if not cascade.budget_ids:
return
where: Final[dict[str, object]] = {"budget_id": {"in": budget_ids}}
if extra_where:
where.update(extra_where)
enduser_ids: Final = tuple(row.user_id for row in cascade.endusers)
async with budget_cascade_unit_of_work(self.prisma_client.db.batch_) as uow:
uow.team_memberships.queue_spend_zero(where=_budget_link_where(cascade.budget_ids))
uow.keys.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _LINKED_KEYS_WHERE))
uow.organizations.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _SPENT_ROWS_WHERE))
uow.tags.queue_spend_zero(where=_budget_link_where(cascade.budget_ids, _SPENT_ROWS_WHERE))
if enduser_ids:
uow.endusers.queue_spend_zero(where={"user_id": {"in": list(enduser_ids)}})
for budget_id, budget_reset_at in cascade.budget_resets:
uow.budgets.queue_window_advance(budget_id=budget_id, budget_reset_at=budget_reset_at)
try:
rows: Sequence[_RowT] = await table.find_many(where=where)
except Exception as e:
rows = ()
verbose_proxy_logger.warning("Failed to fetch %s for counter invalidation: %s", log_subject, e)
update_result: Final = await table.update_many(where=where, data={"spend": 0})
for row in rows:
await self._invalidate_spend_counter(counter_key_fn(row))
if cache_key_fn is not None:
cache_keys = cache_key_fn(row)
if isinstance(cache_keys, str):
cache_keys = [cache_keys]
for cache_key in cache_keys:
await self._invalidate_user_api_key_cache_entry(cache_key)
return update_result
async def reset_budget_for_litellm_team_members(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]):
"""
Resets the budget for all LiteLLM Team Members if their budget has expired
"""
return await self._cascade_reset_spend_for_budget_link(
budgets_to_reset=budgets_to_reset,
table=TeamMembershipRepository(self.prisma_client).table,
counter_key_fn=_team_membership_counter_key,
log_subject="team memberships",
cache_key_fn=_team_membership_cache_key,
)
async def reset_budget_for_keys_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]):
"""
Resets the spend for keys linked to budget tiers that are being reset.
Excludes keys with their own budget_duration; those are reset by
reset_budget_for_litellm_keys() to avoid double-resetting.
"""
return await self._cascade_reset_spend_for_budget_link(
budgets_to_reset=budgets_to_reset,
table=VerificationTokenRepository(self.prisma_client).table,
counter_key_fn=_key_counter_key,
log_subject="keys",
extra_where={"budget_duration": None, "spend": {"gt": 0}},
cache_key_fn=_key_cache_key,
)
async def reset_budget_for_orgs_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]):
"""
Resets the spend for orgs linked to budget tiers that are being reset.
"""
return await self._cascade_reset_spend_for_budget_link(
budgets_to_reset=budgets_to_reset,
table=OrganizationRepository(self.prisma_client).table,
counter_key_fn=_org_counter_key,
log_subject="orgs",
extra_where={"spend": {"gt": 0}},
cache_key_fn=_org_cache_keys,
)
async def reset_budget_for_tags_linked_to_budgets(self, budgets_to_reset: list[LiteLLM_BudgetTableFull]):
"""
Resets the spend for tags linked to budget tiers that are being reset.
Also drops each tag's ``user_api_key_cache`` entry so the next
``_tag_max_budget_check`` reloads the zeroed row from the DB.
``SpendCounterReseed.from_db`` intentionally returns ``None`` for
tags, so the budget check falls back to the cached
``LiteLLM_TagTable.spend`` once the spend counter expires; without
this invalidation, that stale ``.spend`` keeps the tag over-budget
indefinitely.
"""
return await self._cascade_reset_spend_for_budget_link(
budgets_to_reset=budgets_to_reset,
table=TagRepository(self.prisma_client).table,
counter_key_fn=_tag_counter_key,
log_subject="tags",
extra_where={"spend": {"gt": 0}},
cache_key_fn=_tag_cache_key,
)
async def reset_budget_for_litellm_budget_table(self):
"""
Resets the budget for all LiteLLM End-Users (Customers), and Team Members if their budget has expired
The corresponding Budget duration is also updated.
"""
async def _invalidate_budget_cascade_caches(self, cascade: _BudgetCascade) -> None:
for counter_key in cascade.counter_keys:
await self._invalidate_spend_counter(counter_key)
for cache_key in cascade.cache_keys:
await self._invalidate_user_api_key_cache_entry(cache_key)
async def _reset_expired_budget_cascade(self) -> _BudgetCascadeCommitted | _BudgetCascadeFailed:
now: Final = datetime.now(timezone.utc)
start_time: Final = time.time()
endusers_to_reset: list[LiteLLM_EndUserTable] | None = None
budgets_to_reset: list[LiteLLM_BudgetTableFull] | None = None
updated_endusers: Final[list[LiteLLM_EndUserTable]] = []
failed_endusers: Final = []
try:
budgets_to_reset = await self.prisma_client.get_data(
table_name="budget", query_type="find_all", reset_at=now
)
if budgets_to_reset is not None and len(budgets_to_reset) > 0:
for budget in budgets_to_reset:
budget = await ResetBudgetJob._reset_budget_reset_at_date(budget, now, self.reset_settings)
await self.prisma_client.update_data(
query_type="update_many",
data_list=budgets_to_reset,
table_name="budget",
)
budget_ids_to_reset = [budget.budget_id for budget in budgets_to_reset if budget.budget_id is not None]
endusers_to_reset = await self.prisma_client.get_data(
table_name="enduser",
query_type="find_all",
budget_id_list=budget_ids_to_reset,
)
# Also reset end users with no budget_id (NULL) who use the
# default budget via litellm.max_end_user_budget_id. These
# users are enforced in-memory but never had budget_id
# persisted, so the query above misses them.
if litellm.max_end_user_budget_id is not None and litellm.max_end_user_budget_id in budget_ids_to_reset:
default_budget_endusers: Final = await self._get_endusers_with_no_budget_id()
if default_budget_endusers:
if endusers_to_reset is None:
endusers_to_reset = default_budget_endusers
else:
endusers_to_reset.extend(default_budget_endusers)
await self.reset_budget_for_litellm_team_members(budgets_to_reset=budgets_to_reset)
await self.reset_budget_for_keys_linked_to_budgets(budgets_to_reset=budgets_to_reset)
await self.reset_budget_for_orgs_linked_to_budgets(budgets_to_reset=budgets_to_reset)
await self.reset_budget_for_tags_linked_to_budgets(budgets_to_reset=budgets_to_reset)
if endusers_to_reset is not None and len(endusers_to_reset) > 0:
for enduser in endusers_to_reset:
try:
updated_enduser = await ResetBudgetJob._reset_budget_for_enduser(enduser=enduser)
if updated_enduser is not None:
updated_endusers.append(updated_enduser)
else:
failed_endusers.append(
{
"enduser": enduser,
"error": "Returned None without exception",
}
)
except Exception as e:
failed_endusers.append({"enduser": enduser, "error": str(e)})
verbose_proxy_logger.exception("Failed to reset budget for enduser: %s", enduser)
verbose_proxy_logger.debug(
"Updated users %s",
json.dumps(updated_endusers, indent=4, default=str),
)
await self.prisma_client.update_data(
query_type="update_many",
data_list=updated_endusers,
table_name="enduser",
)
end_time = time.time()
if len(failed_endusers) > 0: # If any endusers failed to reset
raise Exception(
f"Failed to reset {len(failed_endusers)} endusers: {json.dumps(failed_endusers, default=str)}"
)
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_success_hook(
service=ServiceTypes.RESET_BUDGET_JOB,
duration=end_time - start_time,
call_type="reset_budget_budget_table",
start_time=start_time,
end_time=end_time,
event_metadata={
"num_budgets_found": (len(budgets_to_reset) if budgets_to_reset else 0),
"budgets_found": json.dumps(budgets_to_reset, indent=4, default=str),
"num_endusers_found": (len(endusers_to_reset) if endusers_to_reset else 0),
"endusers_found": json.dumps(endusers_to_reset, indent=4, default=str),
"num_endusers_updated": len(updated_endusers),
"endusers_updated": json.dumps(updated_endusers, indent=4, default=str),
"num_endusers_failed": len(failed_endusers),
"endusers_failed": json.dumps(failed_endusers, indent=4, default=str),
},
)
budgets_to_reset: Final[Sequence[LiteLLM_BudgetTableFull] | None] = await self.prisma_client.get_data(
table_name="budget",
query_type="find_all",
reset_at=now,
limit=RESET_BUDGET_JOB_BATCH_SIZE,
)
cascade: Final = await self._collect_budget_cascade(budgets_to_reset or ())
except Exception as e:
end_time = time.time()
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_failure_hook(
service=ServiceTypes.RESET_BUDGET_JOB,
duration=end_time - start_time,
error=e,
call_type="reset_budget_endusers",
start_time=start_time,
end_time=end_time,
event_metadata={
"num_budgets_found": (len(budgets_to_reset) if budgets_to_reset else 0),
"budgets_found": json.dumps(budgets_to_reset, indent=4, default=str),
"num_endusers_found": (len(endusers_to_reset) if endusers_to_reset else 0),
"endusers_found": json.dumps(endusers_to_reset, indent=4, default=str),
},
return _BudgetCascadeFailed(cascade=_EMPTY_CASCADE, error=e)
try:
await self._commit_budget_cascade(cascade)
except Exception as e:
return _BudgetCascadeFailed(cascade=cascade, error=e)
await self._invalidate_budget_cascade_caches(cascade)
return _BudgetCascadeCommitted(
cascade=cascade,
advanced=_count_advanced(
(reset_at for _, reset_at in cascade.budget_resets),
cutoff=datetime.now(timezone.utc),
),
)
async def reset_budget_for_litellm_budget_table(self) -> None:
"""
Resets the spend a budget tier gates (end users, team members, keys,
orgs, tags) and advances the tier's budget_reset_at, atomically.
Caches are invalidated only after the transaction commits, so a failed
run cannot leave a zeroed counter in front of an un-reset DB row.
"""
await _run_phase_in_chunks(self._reset_budget_for_litellm_budget_table_chunk)
async def _reset_budget_for_litellm_budget_table_chunk(self) -> _ChunkOutcome:
start_time: Final = time.time()
outcome: Final = await self._reset_expired_budget_cascade()
end_time: Final = time.time()
match outcome:
case _BudgetCascadeCommitted(cascade=cascade, advanced=advanced):
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_success_hook(
service=ServiceTypes.RESET_BUDGET_JOB,
duration=end_time - start_time,
call_type="reset_budget_budget_table",
start_time=start_time,
end_time=end_time,
event_metadata={
**_budget_cascade_event_metadata(cascade),
"num_endusers_updated": len(cascade.endusers),
"num_endusers_failed": 0,
},
)
)
)
verbose_proxy_logger.exception("Failed to reset budget for endusers: %s", e)
return _ChunkOutcome(fetched=len(cascade.budgets), advanced=advanced)
case _BudgetCascadeFailed(cascade=cascade, error=error):
verbose_proxy_logger.exception(
"Failed to reset the budget table cascade (team member, enduser, org and tag spend, plus "
"budget_reset_at); nothing was committed and the budgets stay due for the next run: %s",
error,
exc_info=error,
)
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_failure_hook(
service=ServiceTypes.RESET_BUDGET_JOB,
duration=end_time - start_time,
error=error,
call_type="reset_budget_endusers",
start_time=start_time,
end_time=end_time,
event_metadata=_budget_cascade_event_metadata(cascade),
)
)
return _NO_PROGRESS
case _:
assert_never(outcome)
async def _get_endusers_with_no_budget_id(
self,
@ -486,18 +540,50 @@ class ResetBudgetJob:
for t in updated_teams:
uow.teams.queue_spend_reset(team_id=t.team_id, budget_reset_at=t.budget_reset_at)
async def reset_budget_for_litellm_keys(self):
def _emit_phase_failure(
self,
call_type: str,
error: Exception,
start_time: float,
end_time: float,
event_metadata: dict[str, object],
) -> None:
"""Report rows that could not be reset without failing the chunk: the
rows that did reset are already committed, and raising here would cost
the phase every remaining chunk this tick.
"""
verbose_proxy_logger.error("%s: %s", call_type, error)
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_failure_hook(
service=ServiceTypes.RESET_BUDGET_JOB,
duration=end_time - start_time,
error=error,
call_type=call_type,
start_time=start_time,
end_time=end_time,
event_metadata=event_metadata,
)
)
async def reset_budget_for_litellm_keys(self) -> None:
"""
Resets the budget for all the litellm keys
Catches Exceptions and logs them
"""
await _run_phase_in_chunks(self._reset_budget_for_litellm_keys_chunk)
async def _reset_budget_for_litellm_keys_chunk(self) -> _ChunkOutcome:
now: Final = datetime.utcnow()
start_time: Final = time.time()
keys_to_reset: list[LiteLLM_VerificationToken] | None = None
try:
keys_to_reset = await self.prisma_client.get_data(
table_name="key", query_type="find_all", expires=now, reset_at=now
table_name="key",
query_type="find_all",
expires=now,
reset_at=now,
limit=RESET_BUDGET_JOB_BATCH_SIZE,
)
verbose_proxy_logger.debug("Keys to reset %s", json.dumps(keys_to_reset, indent=4, default=str))
updated_keys: Final[list[LiteLLM_VerificationToken]] = []
@ -528,8 +614,25 @@ class ResetBudgetJob:
await self._invalidate_spend_counter(f"spend:key:{token}")
end_time = time.time()
if len(failed_keys) > 0: # If any keys failed to reset
raise Exception(f"Failed to reset {len(failed_keys)} keys: {json.dumps(failed_keys, default=str)}")
outcome: Final = _ChunkOutcome(
fetched=len(keys_to_reset) if keys_to_reset else 0,
advanced=_count_advanced(
(k.budget_reset_at for k in updated_keys),
cutoff=datetime.now(timezone.utc),
),
)
if len(failed_keys) > 0:
self._emit_phase_failure(
call_type="reset_budget_keys",
error=Exception(f"Failed to reset {len(failed_keys)} keys: {json.dumps(failed_keys, default=str)}"),
start_time=start_time,
end_time=end_time,
event_metadata={
"num_keys_found": len(keys_to_reset) if keys_to_reset else 0,
"keys_found": json.dumps(keys_to_reset, indent=4, default=str),
},
)
return outcome
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_success_hook(
@ -565,16 +668,27 @@ class ResetBudgetJob:
)
)
verbose_proxy_logger.exception("Failed to reset budget for keys: %s", e)
return _NO_PROGRESS
else:
return outcome
async def reset_budget_for_litellm_users(self):
async def reset_budget_for_litellm_users(self) -> None:
"""
Resets the budget for all LiteLLM Internal Users if their budget has expired
"""
await _run_phase_in_chunks(self._reset_budget_for_litellm_users_chunk)
async def _reset_budget_for_litellm_users_chunk(self) -> _ChunkOutcome:
now: Final = datetime.utcnow()
start_time: Final = time.time()
users_to_reset: list[LiteLLM_UserTable] | None = None
try:
users_to_reset = await self.prisma_client.get_data(table_name="user", query_type="find_all", reset_at=now)
users_to_reset = await self.prisma_client.get_data(
table_name="user",
query_type="find_all",
reset_at=now,
limit=RESET_BUDGET_JOB_BATCH_SIZE,
)
updated_users: Final[list[LiteLLM_UserTable]] = []
failed_users: Final = []
if users_to_reset is not None and len(users_to_reset) > 0:
@ -609,8 +723,27 @@ class ResetBudgetJob:
await self._invalidate_global_proxy_spend_cache()
end_time = time.time()
if len(failed_users) > 0: # If any users failed to reset
raise Exception(f"Failed to reset {len(failed_users)} users: {json.dumps(failed_users, default=str)}")
outcome: Final = _ChunkOutcome(
fetched=len(users_to_reset) if users_to_reset else 0,
advanced=_count_advanced(
(u.budget_reset_at for u in updated_users),
cutoff=datetime.now(timezone.utc),
),
)
if len(failed_users) > 0:
self._emit_phase_failure(
call_type="reset_budget_users",
error=Exception(
f"Failed to reset {len(failed_users)} users: {json.dumps(failed_users, default=str)}"
),
start_time=start_time,
end_time=end_time,
event_metadata={
"num_users_found": len(users_to_reset) if users_to_reset else 0,
"users_found": json.dumps(users_to_reset, indent=4, default=str),
},
)
return outcome
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_success_hook(
@ -646,16 +779,27 @@ class ResetBudgetJob:
)
)
verbose_proxy_logger.exception("Failed to reset budget for users: %s", e)
return _NO_PROGRESS
else:
return outcome
async def reset_budget_for_litellm_teams(self):
async def reset_budget_for_litellm_teams(self) -> None:
"""
Resets the budget for all LiteLLM Internal Teams if their budget has expired
"""
await _run_phase_in_chunks(self._reset_budget_for_litellm_teams_chunk)
async def _reset_budget_for_litellm_teams_chunk(self) -> _ChunkOutcome:
now: Final = datetime.utcnow()
start_time: Final = time.time()
teams_to_reset: list[LiteLLM_TeamTable] | None = None
try:
teams_to_reset = await self.prisma_client.get_data(table_name="team", query_type="find_all", reset_at=now)
teams_to_reset = await self.prisma_client.get_data(
table_name="team",
query_type="find_all",
reset_at=now,
limit=RESET_BUDGET_JOB_BATCH_SIZE,
)
updated_teams: Final[list[LiteLLM_TeamTable]] = []
failed_teams: Final = []
if teams_to_reset is not None and len(teams_to_reset) > 0:
@ -688,8 +832,27 @@ class ResetBudgetJob:
await self._invalidate_spend_counter(f"spend:team:{team_id}")
end_time = time.time()
if len(failed_teams) > 0: # If any teams failed to reset
raise Exception(f"Failed to reset {len(failed_teams)} teams: {json.dumps(failed_teams, default=str)}")
outcome: Final = _ChunkOutcome(
fetched=len(teams_to_reset) if teams_to_reset else 0,
advanced=_count_advanced(
(t.budget_reset_at for t in updated_teams),
cutoff=datetime.now(timezone.utc),
),
)
if len(failed_teams) > 0:
self._emit_phase_failure(
call_type="reset_budget_teams",
error=Exception(
f"Failed to reset {len(failed_teams)} teams: {json.dumps(failed_teams, default=str)}"
),
start_time=start_time,
end_time=end_time,
event_metadata={
"num_teams_found": len(teams_to_reset) if teams_to_reset else 0,
"teams_found": json.dumps(teams_to_reset, indent=4, default=str),
},
)
return outcome
asyncio.create_task(
self.proxy_logging_obj.service_logging_obj.async_service_success_hook(
@ -725,6 +888,9 @@ class ResetBudgetJob:
)
)
verbose_proxy_logger.exception("Failed to reset budget for teams: %s", e)
return _NO_PROGRESS
else:
return outcome
@staticmethod
async def _reset_expired_window(
@ -882,33 +1048,6 @@ class ResetBudgetJob:
)
return user
@staticmethod
async def _reset_budget_for_enduser(
enduser: LiteLLM_EndUserTable,
) -> LiteLLM_EndUserTable | None:
try:
enduser.spend = 0.0
except Exception as e:
verbose_proxy_logger.exception("Error resetting budget for enduser: %s. Item: %s", e, enduser)
raise e
return enduser
@staticmethod
async def _reset_budget_reset_at_date(
budget: LiteLLM_BudgetTableFull,
current_time: datetime,
reset_settings: BudgetResetSettings,
) -> LiteLLM_BudgetTableFull:
try:
if budget.budget_duration is not None:
budget.budget_reset_at = compute_budget_reset_at(
budget_duration=budget.budget_duration, settings=reset_settings
)
except Exception as e:
verbose_proxy_logger.exception("Error resetting budget_reset_at for budget: %s. Item: %s", e, budget)
raise e
return budget
@staticmethod
async def _reset_budget_for_key(
key: LiteLLM_VerificationToken,

View file

@ -0,0 +1,185 @@
"""One multi-row ``INSERT ... ON CONFLICT DO UPDATE`` per batch of daily spend rows.
Emitting a statement per aggregated key put every replica's flush on the database as
hundreds of separate statements against the same handful of hot rows, each holding its
row locks for the rest of the enclosing batch transaction. Folding a batch into a single
statement keeps the aggregation identical while collapsing both the statement count and
the window in which those locks are held.
"""
import uuid
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from itertools import groupby
from types import MappingProxyType
from typing import Final, Literal
DailySpendEntity = Literal["user", "team", "org", "tag", "end_user", "agent"]
SqlValue = str | int | float | None
# A queued daily spend transaction, read by column name because the columns are data
# here rather than literals. The concrete TypedDicts in _types.py all satisfy this.
SpendRow = Mapping[str, object]
@dataclass(frozen=True, slots=True)
class DailySpendTable:
"""The physical table behind one entity's daily rollup."""
name: str
entity_id_column: str
carries_request_id: bool = False
DAILY_SPEND_TABLES: Final[Mapping[DailySpendEntity, DailySpendTable]] = MappingProxyType(
{
"user": DailySpendTable(name="LiteLLM_DailyUserSpend", entity_id_column="user_id"),
"team": DailySpendTable(name="LiteLLM_DailyTeamSpend", entity_id_column="team_id"),
"org": DailySpendTable(name="LiteLLM_DailyOrganizationSpend", entity_id_column="organization_id"),
"end_user": DailySpendTable(name="LiteLLM_DailyEndUserSpend", entity_id_column="end_user_id"),
"agent": DailySpendTable(name="LiteLLM_DailyAgentSpend", entity_id_column="agent_id"),
"tag": DailySpendTable(name="LiteLLM_DailyTagSpend", entity_id_column="tag", carries_request_id=True),
}
)
# The unique constraint's columns after the entity id, in constraint order. A NULL can
# never match itself in a unique index, so every one of these is normalized to '': the
# conflict target has to be NULL-free or the row is re-inserted on every single flush.
_KEY_COLUMNS: Final = ("date", "api_key", "model", "custom_llm_provider", "mcp_namespaced_tool_name", "endpoint")
_COUNTER_COLUMNS: Final = (
"prompt_tokens",
"completion_tokens",
"api_requests",
"successful_requests",
"failed_requests",
"cache_read_input_tokens",
"cache_creation_input_tokens",
"compression_saved_tokens",
)
_SPEND_COLUMNS: Final = (
"spend",
"compression_savings_spend",
"prompt_caching_savings_spend",
"autorouter_savings_spend",
)
_CASTS: Final[Mapping[str, str]] = MappingProxyType(
{
**{column: "bigint" for column in _COUNTER_COLUMNS},
**{column: "double precision" for column in _SPEND_COLUMNS},
}
)
def _quoted(columns: Sequence[str]) -> str:
return ", ".join(f'"{column}"' for column in columns)
def _as_text(value: object) -> str:
return "" if value is None else str(value)
def _as_int(value: object) -> int:
return int(value) if isinstance(value, (int, float)) else 0
def _as_float(value: object) -> float:
return float(value) if isinstance(value, (int, float)) else 0.0
def conflict_key(table: DailySpendTable, transaction: SpendRow) -> tuple[str, ...]:
"""The tuple the database arbitrates the upsert on, normalized free of NULLs."""
return tuple(_as_text(transaction.get(column)) for column in (table.entity_id_column, *_KEY_COLUMNS))
def _merge(group: Sequence[SpendRow]) -> SpendRow:
if len(group) == 1:
return group[0]
return {
**group[0],
**{column: sum(_as_int(row.get(column)) for row in group) for column in _COUNTER_COLUMNS},
**{column: sum(_as_float(row.get(column)) for row in group) for column in _SPEND_COLUMNS},
}
def merge_by_conflict_key(
table: DailySpendTable,
transactions: Sequence[SpendRow],
) -> tuple[tuple[tuple[str, ...], SpendRow], ...]:
"""Batch entries keyed by the conflict tuple, in a deterministic order.
The queue keys transactions by their raw field values, so two entries differing only
in a NULL versus an empty member reach the writer separately while arbitrating to the
same row. Postgres rejects a statement whose values touch one row twice, so they are
summed here into the single row they were always destined to become. Ordering by the
key keeps concurrent writers taking row locks in the same sequence.
"""
ordered: Final = sorted(transactions, key=lambda transaction: conflict_key(table, transaction))
return tuple((key, _merge(tuple(group))) for key, group in groupby(ordered, key=lambda t: conflict_key(table, t)))
def _row_params(
table: DailySpendTable,
key: tuple[str, ...],
transaction: SpendRow,
) -> tuple[SqlValue, ...]:
request_id: Final = transaction.get("request_id")
return (
str(uuid.uuid4()),
*key,
None if transaction.get("model_group") is None else _as_text(transaction.get("model_group")),
*(_as_int(transaction.get(column)) for column in _COUNTER_COLUMNS),
*(_as_float(transaction.get(column)) for column in _SPEND_COLUMNS),
*((None if request_id is None else _as_text(request_id),) if table.carries_request_id else ()),
)
def _insert_columns(table: DailySpendTable) -> tuple[str, ...]:
return (
"id",
table.entity_id_column,
*_KEY_COLUMNS,
"model_group",
*_COUNTER_COLUMNS,
*_SPEND_COLUMNS,
*(("request_id",) if table.carries_request_id else ()),
)
def build_bulk_upsert(
table: DailySpendTable,
batch: Sequence[tuple[tuple[str, ...], SpendRow]],
) -> tuple[str, tuple[SqlValue, ...]]:
"""The single statement writing one merged batch, plus its positional arguments."""
columns: Final = _insert_columns(table)
quoted_table: Final = f'"{table.name}"'
rows: Final = ", ".join(
"("
+ ", ".join(
f"${row_index * len(columns) + offset + 1}::{_CASTS.get(column, 'text')}"
for offset, column in enumerate(columns)
)
+ ", (NOW() AT TIME ZONE 'UTC'))"
for row_index in range(len(batch))
)
increments: Final = ", ".join(
f'"{column}" = {quoted_table}."{column}" + EXCLUDED."{column}"'
for column in (*_COUNTER_COLUMNS, *_SPEND_COLUMNS)
)
# request_id names one arbitrary contributing request, so an entry carrying none must
# not blank out the one already recorded.
request_id_update: Final = (
f', "request_id" = COALESCE(EXCLUDED."request_id", {quoted_table}."request_id")'
if table.carries_request_id
else ""
)
sql: Final = (
f'INSERT INTO {quoted_table} ({_quoted(columns)}, "updated_at")\n'
f"VALUES {rows}\n"
f"ON CONFLICT ({_quoted((table.entity_id_column, *_KEY_COLUMNS))}) DO UPDATE SET\n"
f" {increments}{request_id_update},\n"
f" \"updated_at\" = (NOW() AT TIME ZONE 'UTC')"
)
return sql, tuple(value for key, transaction in batch for value in _row_params(table, key, transaction))

View file

@ -12,9 +12,7 @@ import os
import random
import time
import traceback
from collections.abc import Mapping
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, cast, overload
import litellm
@ -41,6 +39,11 @@ from litellm.proxy._types import (
SpendUpdateQueueItem,
ToolDiscoveryQueueItem,
)
from litellm.proxy.db.daily_spend_bulk_upsert import (
DAILY_SPEND_TABLES,
build_bulk_upsert,
merge_by_conflict_key,
)
from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import (
DailySpendUpdateQueue,
)
@ -68,12 +71,6 @@ else:
ProxyLogging = Any
# Only tag rows carry a request_id, so the other entity types spread nothing. Built
# once here rather than as an empty literal per transaction, and read-only so it cannot
# be filled in by accident from one of the call sites that spreads it.
_NO_TAG_REQUEST_ID: Final[Mapping[str, Any]] = MappingProxyType({})
def _get_llm_router():
"""The proxy's router, or None outside a running proxy.
@ -1437,8 +1434,6 @@ class DBSpendUpdateWriter:
daily_spend_transactions: dict[str, DailyUserSpendTransaction],
entity_type: Literal["user"],
entity_id_field: str,
table_name: str,
unique_constraint_name: str,
) -> None:
...
@ -1451,8 +1446,6 @@ class DBSpendUpdateWriter:
daily_spend_transactions: dict[str, DailyTeamSpendTransaction],
entity_type: Literal["team"],
entity_id_field: str,
table_name: str,
unique_constraint_name: str,
) -> None:
...
@ -1465,8 +1458,6 @@ class DBSpendUpdateWriter:
daily_spend_transactions: dict[str, DailyOrganizationSpendTransaction],
entity_type: Literal["org"],
entity_id_field: str,
table_name: str,
unique_constraint_name: str,
) -> None:
...
@ -1479,8 +1470,6 @@ class DBSpendUpdateWriter:
daily_spend_transactions: dict[str, DailyEndUserSpendTransaction],
entity_type: Literal["end_user"],
entity_id_field: str,
table_name: str,
unique_constraint_name: str,
) -> None:
...
@ -1493,8 +1482,6 @@ class DBSpendUpdateWriter:
daily_spend_transactions: dict[str, DailyAgentSpendTransaction],
entity_type: Literal["agent"],
entity_id_field: str,
table_name: str,
unique_constraint_name: str,
) -> None:
...
@ -1507,8 +1494,6 @@ class DBSpendUpdateWriter:
daily_spend_transactions: dict[str, DailyTagSpendTransaction],
entity_type: Literal["tag"],
entity_id_field: str,
table_name: str,
unique_constraint_name: str,
) -> None:
...
# fmt: on
@ -1526,8 +1511,6 @@ class DBSpendUpdateWriter:
| dict[str, DailyAgentSpendTransaction],
entity_type: Literal["user", "team", "org", "tag", "end_user", "agent"],
entity_id_field: str,
table_name: str,
unique_constraint_name: str,
) -> None:
"""
Generic function to update daily spend for any entity type (user, team, org, tag, end_user, agent)
@ -1573,111 +1556,23 @@ class DBSpendUpdateWriter:
)
return
table = DAILY_SPEND_TABLES[entity_type]
try:
async with prisma_client.db.batch_() as batcher:
for _, transaction in transactions_to_process.items():
entity_id = transaction.get(entity_id_field)
# Construct the where clause dynamically
where_clause = {
unique_constraint_name: {
entity_id_field: entity_id,
"date": transaction["date"],
"api_key": transaction["api_key"],
"model": transaction["model"],
"custom_llm_provider": transaction.get("custom_llm_provider") or "",
"mcp_namespaced_tool_name": transaction.get("mcp_namespaced_tool_name")
or "",
"endpoint": transaction.get("endpoint") or "",
}
}
# Get the table dynamically
table = getattr(batcher, table_name)
# Additive metrics that older queued rows may omit; one
# enumeration feeds both the create and the increment below
optional_metrics = {
field: value
for field, value in (
("cache_read_input_tokens", transaction.get("cache_read_input_tokens")),
(
"cache_creation_input_tokens",
transaction.get("cache_creation_input_tokens"),
),
("compression_saved_tokens", transaction.get("compression_saved_tokens")),
(
"compression_savings_spend",
transaction.get("compression_savings_spend"),
),
(
"prompt_caching_savings_spend",
transaction.get("prompt_caching_savings_spend"),
),
("autorouter_savings_spend", transaction.get("autorouter_savings_spend")),
)
if value is not None
}
# Only tag rows carry a request_id. Resolved to a spreadable
# value here so both payloads are built in one shot: a dict
# appended to after construction is one nobody can reason about
# by reading its literal.
tag_request_id: Mapping[str, Any] = (
MappingProxyType({"request_id": transaction["request_id"]})
if entity_type == "tag" and "request_id" in transaction
else _NO_TAG_REQUEST_ID
)
# Common data structure for both create and update
common_data = {
entity_id_field: entity_id,
"date": transaction["date"],
"api_key": transaction["api_key"],
"model": transaction.get("model"),
"model_group": transaction.get("model_group"),
"mcp_namespaced_tool_name": transaction.get("mcp_namespaced_tool_name") or "",
"custom_llm_provider": transaction.get("custom_llm_provider"),
"endpoint": transaction.get("endpoint") or "",
"prompt_tokens": transaction["prompt_tokens"],
"completion_tokens": transaction["completion_tokens"],
"spend": transaction["spend"],
"api_requests": transaction["api_requests"],
"successful_requests": transaction["successful_requests"],
"failed_requests": transaction["failed_requests"],
**optional_metrics,
**tag_request_id,
}
update_data = {
"prompt_tokens": {"increment": transaction["prompt_tokens"]},
"completion_tokens": {"increment": transaction["completion_tokens"]},
"spend": {"increment": transaction["spend"]},
"api_requests": {"increment": transaction["api_requests"]},
"successful_requests": {"increment": transaction["successful_requests"]},
"failed_requests": {"increment": transaction["failed_requests"]},
**{field: {"increment": value} for field, value in optional_metrics.items()},
# An existing row predating the endpoint column gets it filled in here
"endpoint": transaction.get("endpoint") or "",
**tag_request_id,
}
table.upsert(
where=where_clause,
data={
"create": common_data,
"update": update_data,
},
)
# One statement per batch rather than per key: the same rows are
# aggregated, but concurrent writers no longer hold a batch's worth
# of row locks across a hundred round trips.
merged_batch = merge_by_conflict_key(
table=table, transactions=tuple(transactions_to_process.values())
)
sql, params = build_bulk_upsert(table=table, batch=merged_batch)
await prisma_client.db.execute_raw(sql, *params)
except Exception as batch_error:
# Log detailed error information for debugging batch upsert failures
# This helps diagnose issues like unique constraint violations
spend_log_error(
"Daily %s spend batch upsert failed. "
"Table: %s, Constraint: %s, Batch size: %d, Error: %s",
"Daily %s spend batch upsert failed. Table: %s, Rows: %d, Error: %s",
entity_type,
table_name,
unique_constraint_name,
table.name,
len(transactions_to_process),
str(batch_error),
exc=batch_error,
@ -1733,8 +1628,6 @@ class DBSpendUpdateWriter:
daily_spend_transactions=daily_spend_transactions,
entity_type="user",
entity_id_field="user_id",
table_name="litellm_dailyuserspend",
unique_constraint_name="user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint",
)
@staticmethod
@ -1754,8 +1647,6 @@ class DBSpendUpdateWriter:
daily_spend_transactions=daily_spend_transactions,
entity_type="team",
entity_id_field="team_id",
table_name="litellm_dailyteamspend",
unique_constraint_name="team_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint",
)
@staticmethod
@ -1775,8 +1666,6 @@ class DBSpendUpdateWriter:
daily_spend_transactions=daily_spend_transactions,
entity_type="org",
entity_id_field="organization_id",
table_name="litellm_dailyorganizationspend",
unique_constraint_name="organization_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint",
)
@staticmethod
@ -1796,8 +1685,6 @@ class DBSpendUpdateWriter:
daily_spend_transactions=daily_spend_transactions,
entity_type="end_user",
entity_id_field="end_user_id",
table_name="litellm_dailyenduserspend",
unique_constraint_name="end_user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint",
)
@staticmethod
@ -1817,8 +1704,6 @@ class DBSpendUpdateWriter:
daily_spend_transactions=daily_spend_transactions,
entity_type="agent",
entity_id_field="agent_id",
table_name="litellm_dailyagentspend",
unique_constraint_name="agent_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint",
)
@staticmethod
@ -1838,8 +1723,6 @@ class DBSpendUpdateWriter:
daily_spend_transactions=daily_spend_transactions,
entity_type="tag",
entity_id_field="tag",
table_name="litellm_dailytagspend",
unique_constraint_name="tag_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint",
)
async def _common_add_spend_log_transaction_to_daily_transaction(

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