mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
chore(merge): sync with litellm_internal_staging
This commit is contained in:
commit
acb5ebd34b
287 changed files with 25258 additions and 3771 deletions
40
.github/actions/cache-prisma-binaries/action.yml
vendored
Normal file
40
.github/actions/cache-prisma-binaries/action.yml
vendored
Normal 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 }}
|
||||
6
.github/pull_request_template.md
vendored
6
.github/pull_request_template.md
vendored
|
|
@ -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
|
||||
|
||||
|
|
|
|||
34
.github/workflows/_test-unit-base.yml
vendored
34
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -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 }}
|
||||
|
|
|
|||
6
.github/workflows/check-ui-api-types.yml
vendored
6
.github/workflows/check-ui-api-types.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
5
.github/workflows/mutation-test.yml
vendored
5
.github/workflows/mutation-test.yml
vendored
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
6
.github/workflows/test-code-quality.yml
vendored
6
.github/workflows/test-code-quality.yml
vendored
|
|
@ -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
|
||||
|
||||
|
|
|
|||
6
.github/workflows/test-linting.yml
vendored
6
.github/workflows/test-linting.yml
vendored
|
|
@ -71,12 +71,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 +120,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"
|
||||
|
||||
|
|
|
|||
5
.github/workflows/test-litellm-ui-unit.yml
vendored
5
.github/workflows/test-litellm-ui-unit.yml
vendored
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
4
.github/workflows/test-unit-proxy-db.yml
vendored
4
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -76,4 +76,5 @@ jobs:
|
|||
workers: 4
|
||||
reruns: 2
|
||||
timeout-minutes: 60
|
||||
job-timeout-minutes: 95
|
||||
artifact-name: proxy-server
|
||||
|
|
|
|||
6
.github/workflows/test-unit-proxy-legacy.yml
vendored
6
.github/workflows/test-unit-proxy-legacy.yml
vendored
|
|
@ -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
|
||||
|
||||
|
|
|
|||
5
.github/workflows/weekly_load_anomaly.yml
vendored
5
.github/workflows/weekly_load_anomaly.yml
vendored
|
|
@ -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
|
||||
|
||||
|
|
|
|||
10
CLAUDE.md
10
CLAUDE.md
|
|
@ -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:
|
||||
|
||||
|
|
|
|||
|
|
@ -1,18 +1,18 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"limit": 27731
|
||||
"limit": 26391
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2626
|
||||
"limit": 2614
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"limit": 329
|
||||
"limit": 327
|
||||
},
|
||||
"reportAttributeAccessIssue": {
|
||||
"limit": 514
|
||||
},
|
||||
"reportCallIssue": {
|
||||
"limit": 116
|
||||
"limit": 114
|
||||
},
|
||||
"reportConstantRedefinition": {
|
||||
"limit": 40
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 19
|
||||
},
|
||||
"reportExplicitAny": {
|
||||
"limit": 8807
|
||||
"limit": 8319
|
||||
},
|
||||
"reportFunctionMemberAccess": {
|
||||
"limit": 7
|
||||
|
|
@ -54,10 +54,10 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportMissingParameterType": {
|
||||
"limit": 5835
|
||||
"limit": 5825
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15790
|
||||
"limit": 15695
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 40
|
||||
|
|
@ -90,7 +90,7 @@
|
|||
"limit": 8
|
||||
},
|
||||
"reportReturnType": {
|
||||
"limit": 217
|
||||
"limit": 213
|
||||
},
|
||||
"reportTypedDictNotRequiredAccess": {
|
||||
"limit": 26
|
||||
|
|
@ -99,22 +99,22 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 45063
|
||||
"limit": 44996
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 113
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 39773
|
||||
"limit": 39643
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 20207
|
||||
"limit": 20132
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 31281
|
||||
"limit": 31153
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 122
|
||||
"limit": 118
|
||||
},
|
||||
"reportUnnecessaryComparison": {
|
||||
"limit": 701
|
||||
|
|
@ -123,7 +123,7 @@
|
|||
"limit": 5
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"limit": 862
|
||||
"limit": 857
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"limit": 0
|
||||
|
|
|
|||
|
|
@ -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={},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_DailyTeamSpend" ADD COLUMN IF NOT EXISTS "ptu_flat_cost" DOUBLE PRECISION NOT NULL DEFAULT 0.0;
|
||||
|
|
@ -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 '[]';
|
||||
|
|
@ -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
|
||||
|
|
@ -893,6 +893,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 +986,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?
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1493,6 +1493,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 +1721,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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,18 @@ 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 WebSearchInterceptionLogger(CustomLogger):
|
||||
"""
|
||||
CustomLogger that intercepts WebSearch tool calls for models that don't
|
||||
|
|
@ -394,7 +406,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 +474,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 +845,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 +1293,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,
|
||||
|
|
@ -1676,7 +1693,7 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
@staticmethod
|
||||
def initialize_from_proxy_config(
|
||||
litellm_settings: dict[str, Any],
|
||||
callback_specific_params: dict[str, Any],
|
||||
callback_specific_params: Mapping[str, object],
|
||||
) -> "WebSearchInterceptionLogger":
|
||||
"""
|
||||
Static method to initialize WebSearchInterceptionLogger from proxy config.
|
||||
|
|
@ -1700,7 +1717,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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ import subprocess
|
|||
import sys
|
||||
import time
|
||||
import traceback
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Mapping
|
||||
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
|
||||
|
|
@ -313,6 +321,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 +347,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
|
||||
|
|
@ -1992,7 +2031,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 +2498,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 +2914,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 +3108,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.
|
||||
"""
|
||||
|
|
@ -5074,33 +5247,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:
|
||||
|
|
@ -5405,7 +5606,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,
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -213,7 +213,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 +301,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
|
||||
|
|
@ -1839,6 +1908,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 +1922,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 +1962,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 +2170,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 +2232,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 +2251,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."""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
) -> (
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -572,6 +574,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 +671,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
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,6 +1103,8 @@ 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
|
||||
|
|
@ -1819,6 +1822,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 +1888,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 +4025,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 +4111,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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -999,6 +1000,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 +1303,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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
185
litellm/proxy/db/daily_spend_bulk_upsert.py
Normal file
185
litellm/proxy/db/daily_spend_bulk_upsert.py
Normal 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))
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any, Final
|
||||
from typing import Any, Final, TypeVar
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import (
|
||||
|
|
@ -311,14 +311,17 @@ def _coerce_timeout(value: Any, fallback: float) -> float:
|
|||
return fallback
|
||||
|
||||
|
||||
_ReadResultT: Final = TypeVar("_ReadResultT")
|
||||
|
||||
|
||||
async def call_with_db_reconnect_retry(
|
||||
prisma_client: Any,
|
||||
coro_factory: Callable[[], Awaitable[Any]],
|
||||
coro_factory: Callable[[], Awaitable[_ReadResultT]],
|
||||
*,
|
||||
reason: str,
|
||||
timeout_seconds: float | None = None,
|
||||
lock_timeout_seconds: float | None = None,
|
||||
) -> Any:
|
||||
) -> _ReadResultT:
|
||||
"""Run a Prisma read coroutine with one transport-reconnect-and-retry.
|
||||
|
||||
The canonical "self-heal a transient DB transport blip" wrapper used by
|
||||
|
|
|
|||
|
|
@ -9,10 +9,10 @@ import asyncio
|
|||
import json
|
||||
import os
|
||||
import re
|
||||
from collections.abc import AsyncGenerator
|
||||
from collections.abc import AsyncGenerator, Coroutine, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from re import Pattern
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypedDict, cast
|
||||
|
||||
import yaml
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -28,6 +28,7 @@ from litellm.types.utils import (
|
|||
GenericGuardrailAPIInputs,
|
||||
GuardrailStatus,
|
||||
GuardrailTracingDetail,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
)
|
||||
|
||||
|
|
@ -83,6 +84,46 @@ WORD_NUMBER_SEQUENCE_PATTERN: Final = re.compile(
|
|||
WORD_NUMBER_TOKEN_FINDER: Final = re.compile(rf"(?:{WORD_NUMBER_TOKEN_REGEX})", re.IGNORECASE)
|
||||
|
||||
|
||||
class ConditionalCategoryConfig(TypedDict):
|
||||
identifier_words: Sequence[str]
|
||||
block_words: Sequence[str]
|
||||
action: ContentFilterAction
|
||||
severity: str
|
||||
|
||||
|
||||
class CompiledPatternEntry(TypedDict):
|
||||
regex: Pattern[str]
|
||||
pattern_name: str
|
||||
action: ContentFilterAction
|
||||
keyword_regex: Pattern[str] | None
|
||||
allow_word_numbers: bool
|
||||
|
||||
|
||||
class _PatternExtraLookup(TypedDict):
|
||||
keyword_pattern: str | None
|
||||
allow_word_numbers: bool
|
||||
|
||||
|
||||
class _CategoryConfigView(TypedDict):
|
||||
category: object
|
||||
enabled: object
|
||||
action: object
|
||||
category_file: str | None
|
||||
|
||||
|
||||
class CategoryFileData(TypedDict, total=False):
|
||||
category_name: str
|
||||
description: str
|
||||
default_action: str
|
||||
keywords: Sequence[Mapping[str, str]]
|
||||
exceptions: Sequence[str]
|
||||
identifier_words: Sequence[str]
|
||||
always_block_keywords: Sequence[Mapping[str, str]]
|
||||
inherit_from: str
|
||||
additional_block_words: Sequence[str]
|
||||
phrase_patterns: Sequence[str]
|
||||
|
||||
|
||||
# Helper data structure for category-based detection
|
||||
class CategoryConfig:
|
||||
"""Configuration for a content category."""
|
||||
|
|
@ -92,13 +133,13 @@ class CategoryConfig:
|
|||
category_name: str,
|
||||
description: str,
|
||||
default_action: ContentFilterAction,
|
||||
keywords: list[dict[str, str]],
|
||||
exceptions: list[str],
|
||||
identifier_words: list[str] | None = None,
|
||||
always_block_keywords: list[dict[str, str]] | None = None,
|
||||
keywords: Sequence[Mapping[str, str]],
|
||||
exceptions: Sequence[str],
|
||||
identifier_words: Sequence[str] | None = None,
|
||||
always_block_keywords: Sequence[Mapping[str, str]] | None = None,
|
||||
inherit_from: str | None = None,
|
||||
additional_block_words: list[str] | None = None,
|
||||
phrase_patterns: list[str] | None = None,
|
||||
additional_block_words: Sequence[str] | None = None,
|
||||
phrase_patterns: Sequence[str] | None = None,
|
||||
):
|
||||
self.category_name = category_name
|
||||
self.description = description
|
||||
|
|
@ -151,7 +192,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
severity_threshold: str = "medium",
|
||||
llm_router: Router | None = None,
|
||||
image_model: str | None = None,
|
||||
competitor_intent_config: dict[str, Any] | None = None,
|
||||
competitor_intent_config: dict[str, object] | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
|
|
@ -194,9 +235,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
# Always-block keywords are checked after exceptions (exceptions take precedence)
|
||||
self.always_block_category_keywords: dict[str, tuple[str, str, ContentFilterAction]] = {}
|
||||
# Store conditional categories (identifier_words + block_words)
|
||||
self.conditional_categories: dict[
|
||||
str, dict[str, Any]
|
||||
] = {} # category_name -> {identifier_words, block_words, action, severity}
|
||||
self.conditional_categories: dict[str, ConditionalCategoryConfig] = {}
|
||||
|
||||
# Competitor intent checker (optional; airline uses major_airlines.json, generic requires competitors)
|
||||
self._competitor_intent_checker: BaseCompetitorIntentChecker | None = None
|
||||
|
|
@ -212,7 +251,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
normalized_blocked_words: Final = self._normalize_blocked_words(blocked_words)
|
||||
|
||||
# Compile regex patterns
|
||||
self.compiled_patterns: list[dict[str, Any]] = []
|
||||
self.compiled_patterns: list[CompiledPatternEntry] = []
|
||||
for pattern_config in normalized_patterns:
|
||||
self._add_pattern(pattern_config)
|
||||
|
||||
|
|
@ -250,7 +289,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
"Loaded %s categories with %s keywords", len(self.loaded_categories), len(self.category_keywords)
|
||||
)
|
||||
|
||||
def _init_competitor_intent_checker(self, competitor_intent_config: dict[str, Any]) -> None:
|
||||
def _init_competitor_intent_checker(self, competitor_intent_config: dict[str, object]) -> None:
|
||||
try:
|
||||
competitor_intent_type: Final = competitor_intent_config.get("competitor_intent_type", "airline")
|
||||
if competitor_intent_type == "generic":
|
||||
|
|
@ -293,6 +332,15 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
result.append(word)
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _category_config_view(cat_config: ContentFilterCategoryConfig) -> _CategoryConfigView:
|
||||
return {
|
||||
"category": cat_config.get("category"),
|
||||
"enabled": cat_config.get("enabled", True),
|
||||
"action": cat_config.get("action"),
|
||||
"category_file": cat_config.get("category_file"),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _assert_within_categories_dir(path: str, categories_dir: str) -> None:
|
||||
"""Raise ValueError if path escapes the categories directory."""
|
||||
|
|
@ -395,7 +443,8 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
categories_dir: Final = os.path.join(os.path.dirname(__file__), "categories")
|
||||
|
||||
for cat_config in categories:
|
||||
category_name = cat_config.get("category")
|
||||
view = self._category_config_view(cat_config)
|
||||
category_name = view["category"]
|
||||
if not category_name or not isinstance(category_name, str):
|
||||
verbose_proxy_logger.warning("Category name missing or invalid in config, skipping")
|
||||
continue
|
||||
|
|
@ -405,12 +454,12 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
verbose_proxy_logger.warning("Category name '%s' contains invalid characters, skipping", category_name)
|
||||
continue
|
||||
|
||||
enabled = cat_config.get("enabled", True)
|
||||
action = cat_config.get("action")
|
||||
enabled = view["enabled"]
|
||||
action = view["action"]
|
||||
severity_threshold = (
|
||||
cat_config.get("severity_threshold", self.severity_threshold) or self.severity_threshold
|
||||
)
|
||||
custom_file = cat_config.get("category_file")
|
||||
custom_file = view["category_file"]
|
||||
|
||||
if not enabled:
|
||||
verbose_proxy_logger.debug("Category %s is disabled, skipping", category_name)
|
||||
|
|
@ -514,7 +563,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
categories_dir: Directory containing category files
|
||||
"""
|
||||
try:
|
||||
block_words: Final = []
|
||||
block_words: Final[list[str]] = []
|
||||
inherit_from = category_config_obj.inherit_from
|
||||
|
||||
# Load inherited block words if specified
|
||||
|
|
@ -605,11 +654,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
"""
|
||||
if file_path.lower().endswith(".json"):
|
||||
return self._load_category_file_json(file_path)
|
||||
with open(file_path, "r") as f:
|
||||
data: Final = yaml.safe_load(f)
|
||||
|
||||
# Handle always_block_keywords if present
|
||||
always_block: Final = data.get("always_block_keywords", [])
|
||||
data: Final = self._read_category_yaml(file_path)
|
||||
|
||||
return CategoryConfig(
|
||||
category_name=data.get("category_name", "unknown"),
|
||||
|
|
@ -618,12 +663,17 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
keywords=data.get("keywords", []),
|
||||
exceptions=data.get("exceptions", []),
|
||||
identifier_words=data.get("identifier_words"),
|
||||
always_block_keywords=always_block,
|
||||
always_block_keywords=data.get("always_block_keywords", []),
|
||||
inherit_from=data.get("inherit_from"),
|
||||
additional_block_words=data.get("additional_block_words"),
|
||||
phrase_patterns=data.get("phrase_patterns"),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _read_category_yaml(file_path: str) -> CategoryFileData:
|
||||
with open(file_path, "r") as f:
|
||||
return yaml.safe_load(f)
|
||||
|
||||
def _load_category_file_json(self, file_path: str) -> CategoryConfig:
|
||||
"""
|
||||
Load a category from the harm_toxic_abuse-style JSON format.
|
||||
|
|
@ -682,13 +732,13 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
pattern_config: ContentFilterPattern configuration
|
||||
"""
|
||||
try:
|
||||
extra_config: dict[str, Any] = {}
|
||||
extra_config: _PatternExtraLookup = {"keyword_pattern": None, "allow_word_numbers": False}
|
||||
if pattern_config.pattern_type == "prebuilt":
|
||||
if not pattern_config.pattern_name:
|
||||
raise ValueError("pattern_name is required for prebuilt patterns")
|
||||
compiled = get_compiled_pattern(pattern_config.pattern_name)
|
||||
pattern_name = pattern_config.pattern_name
|
||||
extra_config = PATTERN_EXTRA_CONFIG.get(pattern_name, {}) or {}
|
||||
extra_config = self._lookup_pattern_extra(pattern_name)
|
||||
elif pattern_config.pattern_type == "regex":
|
||||
if not pattern_config.pattern:
|
||||
raise ValueError("pattern is required for regex patterns")
|
||||
|
|
@ -697,9 +747,8 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
else:
|
||||
raise ValueError(f"Unknown pattern_type: {pattern_config.pattern_type}")
|
||||
|
||||
keyword_regex: Pattern | None = None
|
||||
if extra_config.get("keyword_pattern"):
|
||||
keyword_regex = re.compile(extra_config["keyword_pattern"], re.IGNORECASE)
|
||||
keyword_pattern: Final = extra_config["keyword_pattern"]
|
||||
keyword_regex: Final = re.compile(keyword_pattern, re.IGNORECASE) if keyword_pattern else None
|
||||
|
||||
self.compiled_patterns.append(
|
||||
{
|
||||
|
|
@ -707,7 +756,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
"pattern_name": pattern_name,
|
||||
"action": pattern_config.action,
|
||||
"keyword_regex": keyword_regex,
|
||||
"allow_word_numbers": bool(extra_config.get("allow_word_numbers")),
|
||||
"allow_word_numbers": extra_config["allow_word_numbers"],
|
||||
}
|
||||
)
|
||||
verbose_proxy_logger.debug("Added pattern: %s with action %s", pattern_name, pattern_config.action)
|
||||
|
|
@ -715,6 +764,14 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
verbose_proxy_logger.error("Error adding pattern %s: %s", pattern_config, e)
|
||||
raise
|
||||
|
||||
@staticmethod
|
||||
def _lookup_pattern_extra(pattern_name: str) -> _PatternExtraLookup:
|
||||
extra: Final = PATTERN_EXTRA_CONFIG.get(pattern_name)
|
||||
return {
|
||||
"keyword_pattern": extra.get("keyword_pattern") if extra is not None else None,
|
||||
"allow_word_numbers": bool(extra.get("allow_word_numbers")) if extra is not None else False,
|
||||
}
|
||||
|
||||
def _load_blocked_words_file(self, file_path: str) -> None:
|
||||
"""
|
||||
Load blocked words from a YAML file.
|
||||
|
|
@ -754,18 +811,16 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
except Exception as e:
|
||||
raise Exception(f"Error loading blocked words file {file_path}: {e}")
|
||||
|
||||
def _find_pattern_spans(self, text: str, pattern_entry: dict[str, Any]) -> list[tuple[int, int]]:
|
||||
def _find_pattern_spans(self, text: str, pattern_entry: CompiledPatternEntry) -> list[tuple[int, int]]:
|
||||
"""Return all match spans for a pattern, applying contextual rules if required."""
|
||||
|
||||
regex: Final[Pattern] = pattern_entry["regex"]
|
||||
keyword_regex: Final[Pattern | None] = pattern_entry.get("keyword_regex")
|
||||
regex: Final[Pattern[str]] = pattern_entry["regex"]
|
||||
keyword_regex: Final[Pattern[str] | None] = pattern_entry.get("keyword_regex")
|
||||
allow_word_numbers: Final[bool] = pattern_entry.get("allow_word_numbers", False)
|
||||
|
||||
keyword_matches: list[re.Match] | None = None
|
||||
if keyword_regex is not None:
|
||||
keyword_matches = list(keyword_regex.finditer(text))
|
||||
if not keyword_matches:
|
||||
return []
|
||||
keyword_matches: Final = list(keyword_regex.finditer(text)) if keyword_regex is not None else None
|
||||
if keyword_matches is not None and not keyword_matches:
|
||||
return []
|
||||
|
||||
match_spans: Final[list[tuple[int, int]]] = []
|
||||
|
||||
|
|
@ -795,7 +850,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
self,
|
||||
value_start: int,
|
||||
value_end: int,
|
||||
keyword_matches: list[re.Match],
|
||||
keyword_matches: Sequence[re.Match[str]],
|
||||
text: str,
|
||||
) -> bool:
|
||||
"""Check if a value is separated from a keyword by an allowed gap."""
|
||||
|
|
@ -861,7 +916,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
def _convert_word_number_sequence(self, sequence: str) -> str | None:
|
||||
"""Convert a spelled-out digit sequence (e.g., 'One-Two') into digits."""
|
||||
|
||||
tokens: Final = WORD_NUMBER_TOKEN_FINDER.findall(sequence)
|
||||
tokens: Final[list[str]] = WORD_NUMBER_TOKEN_FINDER.findall(sequence)
|
||||
if not tokens:
|
||||
return None
|
||||
|
||||
|
|
@ -1328,7 +1383,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
HTTPException: If sensitive content is detected and action is BLOCK
|
||||
"""
|
||||
# Collect all exceptions from loaded categories
|
||||
all_exceptions: Final = []
|
||||
all_exceptions: Final[list[str]] = []
|
||||
for category in self.loaded_categories.values():
|
||||
all_exceptions.extend(category.exceptions)
|
||||
|
||||
|
|
@ -1404,7 +1459,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
if not (images and self.image_model and self.llm_router):
|
||||
return
|
||||
|
||||
tasks: Final = []
|
||||
tasks: Final[list[Coroutine[object, object, ModelResponse]]] = []
|
||||
for image in images:
|
||||
task = self.llm_router.acompletion(
|
||||
model=self.image_model,
|
||||
|
|
@ -1425,12 +1480,10 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
tasks.append(task)
|
||||
|
||||
responses: Final = await asyncio.gather(*tasks)
|
||||
descriptions: Final = []
|
||||
descriptions: Final[list[str]] = []
|
||||
for response in responses:
|
||||
choice = response.choices[0]
|
||||
message = getattr(choice, "message", None)
|
||||
if message and getattr(message, "content", None):
|
||||
image_description = message.content
|
||||
image_description = self._describe_image_response_content(response)
|
||||
if image_description:
|
||||
verbose_proxy_logger.debug("Image description: %s", image_description)
|
||||
descriptions.append(image_description)
|
||||
else:
|
||||
|
|
@ -1447,7 +1500,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
except HTTPException as e:
|
||||
# e.detail can be a string or dict
|
||||
if isinstance(e.detail, dict) and "error" in e.detail:
|
||||
detail_dict = cast(dict[str, Any], e.detail)
|
||||
detail_dict = cast(dict[str, str], e.detail)
|
||||
detail_dict["error"] = detail_dict["error"] + " (Image description): " + description
|
||||
elif isinstance(e.detail, str):
|
||||
e.detail = e.detail + " (Image description): " + description
|
||||
|
|
@ -1455,6 +1508,14 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
e.detail = "Content blocked: Image description detected" + description
|
||||
raise e
|
||||
|
||||
@staticmethod
|
||||
def _describe_image_response_content(response: ModelResponse) -> str | None:
|
||||
choice = response.choices[0]
|
||||
message = getattr(choice, "message", None)
|
||||
if message and getattr(message, "content", None):
|
||||
return message.content
|
||||
return None
|
||||
|
||||
def _count_masked_entities(
|
||||
self,
|
||||
detections: list[ContentFilterDetection],
|
||||
|
|
@ -1484,12 +1545,12 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
category = category_detection["category"]
|
||||
masked_entity_count[category] = masked_entity_count.get(category, 0) + 1
|
||||
|
||||
def _build_match_details(self, detections: list[ContentFilterDetection]) -> list[dict]:
|
||||
def _build_match_details(self, detections: list[ContentFilterDetection]) -> list[dict[str, object]]:
|
||||
"""Build match_details list from content filter detections."""
|
||||
match_details: Final[list[dict]] = []
|
||||
match_details: Final[list[dict[str, object]]] = []
|
||||
for detection in detections:
|
||||
action_taken = detection.get("action", detection.get("action_hint", ""))
|
||||
detail: dict = {"type": detection["type"], "action_taken": action_taken}
|
||||
detail: dict[str, object] = {"type": detection["type"], "action_taken": action_taken}
|
||||
if detection["type"] == "pattern":
|
||||
detail["detection_method"] = "regex"
|
||||
detail["snippet"] = cast(PatternDetection, detection).get("pattern_name", "")
|
||||
|
|
@ -1510,7 +1571,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
|
||||
def _get_detection_methods(self, detections: list[ContentFilterDetection]) -> str:
|
||||
"""Get comma-separated detection methods used."""
|
||||
methods: Final[set] = set()
|
||||
methods: Final[set[str]] = set()
|
||||
for detection in detections:
|
||||
if detection["type"] == "pattern":
|
||||
methods.add("regex")
|
||||
|
|
@ -1659,7 +1720,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
guardrail_json_response = exception_str if exception_str else [dict(detection) for detection in detections]
|
||||
|
||||
# Competitor intent: add confidence and classification to tracing if present
|
||||
tracing_kw: Final[dict[str, Any]] = {
|
||||
tracing_kw: Final[GuardrailTracingDetail] = {
|
||||
"guardrail_id": self.config_guardrail_id or self.guardrail_name,
|
||||
"policy_template": self.config_policy_template or self._get_policy_templates(),
|
||||
"detection_method": (self._get_detection_methods(detections) if detections else None),
|
||||
|
|
|
|||
|
|
@ -55,7 +55,7 @@ for pattern_data in _PATTERNS_DATA["patterns"]:
|
|||
PATTERN_EXTRA_CONFIG[pattern_data["name"]] = extra_config
|
||||
|
||||
|
||||
def get_compiled_pattern(pattern_name: str) -> Pattern:
|
||||
def get_compiled_pattern(pattern_name: str) -> Pattern[str]:
|
||||
"""
|
||||
Get a compiled regex pattern by name.
|
||||
|
||||
|
|
|
|||
|
|
@ -73,12 +73,14 @@ import hashlib
|
|||
import os
|
||||
import re
|
||||
import time
|
||||
from typing import Any, Final, Optional
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional
|
||||
|
||||
import jwt
|
||||
from cryptography.hazmat.primitives import serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey
|
||||
from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey, RSAPublicKey
|
||||
from typing_extensions import NotRequired, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching import DualCache
|
||||
|
|
@ -90,13 +92,28 @@ from litellm.proxy._types import UserAPIKeyAuth
|
|||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import CallTypesLiteral
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from jwt.types import Options
|
||||
|
||||
|
||||
class _OIDCDiscoveryDocument(TypedDict, total=False):
|
||||
jwks_uri: str
|
||||
|
||||
|
||||
class _JWTDecodeKwargs(TypedDict):
|
||||
algorithms: Sequence[str]
|
||||
options: "Options"
|
||||
audience: NotRequired[str]
|
||||
issuer: NotRequired[str]
|
||||
|
||||
|
||||
# Module-level singleton for the JWKS discovery endpoint to access.
|
||||
_mcp_jwt_signer_instance: Optional["MCPJWTSigner"] = None
|
||||
|
||||
_MCP_JWT_CALL_TYPES: Final = frozenset({"call_mcp_tool", "list_mcp_tools"})
|
||||
|
||||
# Simple in-memory JWKS cache: keyed by JWKS URI → (keys_list, fetched_at).
|
||||
_jwks_cache: Final[dict[str, tuple]] = {}
|
||||
_jwks_cache: Final[dict[str, tuple[Sequence[Mapping[str, object]], float]]] = {}
|
||||
_JWKS_CACHE_TTL: Final = 3600 # 1 hour
|
||||
|
||||
|
||||
|
|
@ -133,7 +150,7 @@ def _int_to_base64url(n: int) -> str:
|
|||
return base64.urlsafe_b64encode(n.to_bytes(byte_length, byteorder="big")).rstrip(b"=").decode("ascii")
|
||||
|
||||
|
||||
def _compute_kid(public_key: Any) -> str:
|
||||
def _compute_kid(public_key: RSAPublicKey) -> str:
|
||||
"""Derive a key ID from the public key's DER encoding (SHA-256, first 16 hex chars)."""
|
||||
der_bytes: Final = public_key.public_bytes(
|
||||
encoding=serialization.Encoding.DER,
|
||||
|
|
@ -142,7 +159,7 @@ def _compute_kid(public_key: Any) -> str:
|
|||
return hashlib.sha256(der_bytes).hexdigest()[:16]
|
||||
|
||||
|
||||
async def _fetch_jwks(jwks_uri: str) -> list[dict[str, Any]]:
|
||||
async def _fetch_jwks(jwks_uri: str) -> Sequence[Mapping[str, object]]:
|
||||
"""
|
||||
Fetch and cache a JWKS from the given URI.
|
||||
|
||||
|
|
@ -163,12 +180,13 @@ async def _fetch_jwks(jwks_uri: str) -> list[dict[str, Any]]:
|
|||
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
resp: Final = await client.get(jwks_uri, headers={"Accept": "application/json"})
|
||||
resp.raise_for_status()
|
||||
keys = resp.json().get("keys", [])
|
||||
_jwks_cache[jwks_uri] = (keys, now)
|
||||
return keys
|
||||
jwks_body: Final[Mapping[str, Sequence[Mapping[str, object]]]] = resp.json()
|
||||
fetched_keys: Final = jwks_body.get("keys", [])
|
||||
_jwks_cache[jwks_uri] = (fetched_keys, now)
|
||||
return fetched_keys
|
||||
|
||||
|
||||
async def _fetch_oidc_discovery(discovery_uri: str) -> dict[str, Any]:
|
||||
async def _fetch_oidc_discovery(discovery_uri: str) -> _OIDCDiscoveryDocument:
|
||||
"""Fetch an OIDC discovery document and return its parsed JSON."""
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
|
|
@ -178,7 +196,8 @@ async def _fetch_oidc_discovery(discovery_uri: str) -> dict[str, Any]:
|
|||
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
resp: Final = await client.get(discovery_uri, headers={"Accept": "application/json"})
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
document: Final[_OIDCDiscoveryDocument] = resp.json()
|
||||
return document
|
||||
|
||||
|
||||
class MCPJWTSigner(CustomGuardrail):
|
||||
|
|
@ -230,8 +249,8 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
# FR-12: End-user identity mapping
|
||||
end_user_claim_sources: list[str] | None = None,
|
||||
# FR-13: Claim operations
|
||||
add_claims: dict[str, Any] | None = None,
|
||||
set_claims: dict[str, Any] | None = None,
|
||||
add_claims: Mapping[str, object] | None = None,
|
||||
set_claims: Mapping[str, object] | None = None,
|
||||
remove_claims: list[str] | None = None,
|
||||
# FR-14: Two-token model
|
||||
channel_token_audience: str | None = None,
|
||||
|
|
@ -283,7 +302,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
self.verify_issuer: str | None = verify_issuer
|
||||
self.verify_audience: str | None = verify_audience
|
||||
# Cached OIDC discovery document (fetched lazily, TTL = 24 h)
|
||||
self._oidc_discovery_doc: dict[str, Any] | None = None
|
||||
self._oidc_discovery_doc: _OIDCDiscoveryDocument | None = None
|
||||
self._oidc_discovery_fetched_at: float = 0.0
|
||||
|
||||
# --- FR-12: End-user identity mapping ---
|
||||
|
|
@ -294,8 +313,8 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
]
|
||||
|
||||
# --- FR-13: Claim operations ---
|
||||
self.add_claims: dict[str, Any] = add_claims or {}
|
||||
self.set_claims: dict[str, Any] = set_claims or {}
|
||||
self.add_claims: Mapping[str, object] = add_claims or {}
|
||||
self.set_claims: Mapping[str, object] = set_claims or {}
|
||||
self.remove_claims: list[str] = remove_claims or []
|
||||
|
||||
# --- FR-14: Two-token model ---
|
||||
|
|
@ -347,7 +366,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
"""
|
||||
return 3600 if self._persistent_key else 300
|
||||
|
||||
def get_jwks(self) -> dict[str, Any]:
|
||||
def get_jwks(self) -> Mapping[str, Sequence[Mapping[str, str]]]:
|
||||
"""
|
||||
Return the JWKS for the RSA public key.
|
||||
Used by GET /.well-known/jwks.json so MCP servers can verify tokens.
|
||||
|
|
@ -374,7 +393,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
# the IdP, short enough to pick up jwks_uri changes after key rotation.
|
||||
_OIDC_DISCOVERY_TTL = 86400
|
||||
|
||||
async def _get_oidc_discovery(self) -> dict[str, Any]:
|
||||
async def _get_oidc_discovery(self) -> _OIDCDiscoveryDocument:
|
||||
"""Fetch and cache the OIDC discovery document with a 24-hour TTL.
|
||||
|
||||
Only caches when the doc contains a 'jwks_uri' so that a transient or
|
||||
|
|
@ -391,7 +410,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
return doc
|
||||
return self._oidc_discovery_doc or {}
|
||||
|
||||
async def _verify_incoming_jwt(self, raw_token: str) -> dict[str, Any]:
|
||||
async def _verify_incoming_jwt(self, raw_token: str) -> dict[str, object]:
|
||||
"""
|
||||
Verify an incoming Bearer JWT against the configured IdP's JWKS.
|
||||
|
||||
|
|
@ -438,8 +457,8 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
# it infers from the key type (RSAPublicKey → RS256).
|
||||
alg: Final = getattr(signing_jwk, "algorithm_name", None) or "RS256"
|
||||
|
||||
decode_options: Final[dict[str, Any]] = {"verify_exp": True}
|
||||
decode_kwargs: Final[dict[str, Any]] = {
|
||||
decode_options: Final[Options] = {"verify_exp": True}
|
||||
decode_kwargs: Final[_JWTDecodeKwargs] = {
|
||||
"algorithms": [alg],
|
||||
"options": decode_options,
|
||||
}
|
||||
|
|
@ -451,10 +470,10 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
if self.verify_issuer:
|
||||
decode_kwargs["issuer"] = self.verify_issuer
|
||||
|
||||
payload: Final[dict[str, Any]] = jwt.decode(raw_token, signing_jwk.key, **decode_kwargs)
|
||||
payload: Final[dict[str, object]] = jwt.decode(raw_token, signing_jwk.key, **decode_kwargs)
|
||||
return payload
|
||||
|
||||
async def _introspect_opaque_token(self, token: str) -> dict[str, Any]:
|
||||
async def _introspect_opaque_token(self, token: str) -> dict[str, object]:
|
||||
"""
|
||||
Perform RFC 7662 token introspection for opaque (non-JWT) tokens.
|
||||
|
||||
|
|
@ -479,7 +498,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
headers={"Accept": "application/json"},
|
||||
)
|
||||
resp.raise_for_status()
|
||||
result: Final[dict[str, Any]] = resp.json()
|
||||
result: Final[dict[str, object]] = resp.json()
|
||||
if not result.get("active", False):
|
||||
raise jwt.exceptions.ExpiredSignatureError(
|
||||
"MCPJWTSigner: incoming token is inactive (introspection returned active=false)"
|
||||
|
|
@ -492,7 +511,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
|
||||
def _validate_required_claims(
|
||||
self,
|
||||
jwt_claims: dict[str, Any] | None,
|
||||
jwt_claims: Mapping[str, object] | None,
|
||||
) -> None:
|
||||
"""
|
||||
Raise HTTP 403 if any required_claims are absent from the verified
|
||||
|
|
@ -522,7 +541,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
def _resolve_end_user_identity(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
jwt_claims: dict[str, Any] | None,
|
||||
jwt_claims: Mapping[str, object] | None,
|
||||
) -> str:
|
||||
"""
|
||||
Resolve the outbound JWT 'sub' using the ordered end_user_claim_sources list.
|
||||
|
|
@ -545,19 +564,19 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
value = str(raw) if raw else None
|
||||
|
||||
elif source == "litellm:user_id":
|
||||
uid = getattr(user_api_key_dict, "user_id", None)
|
||||
uid = user_api_key_dict.user_id
|
||||
value = str(uid) if uid else None
|
||||
|
||||
elif source == "litellm:email":
|
||||
email = getattr(user_api_key_dict, "user_email", None)
|
||||
email = user_api_key_dict.user_email
|
||||
value = str(email) if email else None
|
||||
|
||||
elif source == "litellm:end_user_id":
|
||||
eid = getattr(user_api_key_dict, "end_user_id", None)
|
||||
eid = user_api_key_dict.end_user_id
|
||||
value = str(eid) if eid else None
|
||||
|
||||
elif source == "litellm:team_id":
|
||||
tid = getattr(user_api_key_dict, "team_id", None)
|
||||
tid = user_api_key_dict.team_id
|
||||
value = str(tid) if tid else None
|
||||
|
||||
else:
|
||||
|
|
@ -568,7 +587,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
return value
|
||||
|
||||
# Final fallback for service accounts with no user identity
|
||||
token: Final = getattr(user_api_key_dict, "token", None) or getattr(user_api_key_dict, "api_key", None)
|
||||
token: Final = user_api_key_dict.token or user_api_key_dict.api_key
|
||||
if token:
|
||||
return "apikey:" + hashlib.sha256(str(token).encode()).hexdigest()[:16]
|
||||
return "litellm-proxy"
|
||||
|
|
@ -615,7 +634,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
# FR-13: Claim operations
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _apply_claim_operations(self, claims: dict[str, Any]) -> dict[str, Any]:
|
||||
def _apply_claim_operations(self, claims: dict[str, object]) -> dict[str, object]:
|
||||
"""Apply add_claims, set_claims, and remove_claims to the claim dict."""
|
||||
# add_claims: insert only when key is absent
|
||||
for k, v in self.add_claims.items():
|
||||
|
|
@ -637,9 +656,9 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
|
||||
def _passthrough_optional_claims(
|
||||
self,
|
||||
claims: dict[str, Any],
|
||||
jwt_claims: dict[str, Any] | None,
|
||||
) -> dict[str, Any]:
|
||||
claims: dict[str, object],
|
||||
jwt_claims: Mapping[str, object] | None,
|
||||
) -> dict[str, object]:
|
||||
"""Forward optional_claims from verified incoming token into the outbound JWT."""
|
||||
if not self.optional_claims or not jwt_claims:
|
||||
return claims
|
||||
|
|
@ -656,7 +675,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
data: dict,
|
||||
jwt_claims: dict[str, Any] | None = None,
|
||||
jwt_claims: Mapping[str, object] | None = None,
|
||||
call_type: CallTypesLiteral | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
|
|
@ -669,7 +688,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
jwt_claims if available. None for pure API-key requests.
|
||||
"""
|
||||
now: Final = int(time.time())
|
||||
claims: dict[str, Any] = {
|
||||
claims: dict[str, object] = {
|
||||
"iss": self.issuer,
|
||||
"aud": self.audience,
|
||||
"iat": now,
|
||||
|
|
@ -681,18 +700,18 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
claims["sub"] = self._resolve_end_user_identity(user_api_key_dict, jwt_claims)
|
||||
|
||||
# email passthrough when available from LiteLLM context
|
||||
user_email: Final = getattr(user_api_key_dict, "user_email", None)
|
||||
user_email: Final = user_api_key_dict.user_email
|
||||
if user_email:
|
||||
claims["email"] = user_email
|
||||
|
||||
# act — RFC 8693 delegation claim (team/org context)
|
||||
team_id: Final = getattr(user_api_key_dict, "team_id", None)
|
||||
org_id: Final = getattr(user_api_key_dict, "org_id", None)
|
||||
team_id: Final = user_api_key_dict.team_id
|
||||
org_id: Final = user_api_key_dict.org_id
|
||||
act_sub: Final = team_id or org_id or "litellm-proxy"
|
||||
claims["act"] = {"sub": act_sub}
|
||||
|
||||
# end_user_id when set separately from user_id
|
||||
end_user_id: Final = getattr(user_api_key_dict, "end_user_id", None)
|
||||
end_user_id: Final = user_api_key_dict.end_user_id
|
||||
if end_user_id:
|
||||
claims["end_user_id"] = end_user_id
|
||||
|
||||
|
|
@ -710,8 +729,8 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
|
||||
def _build_channel_token_claims(
|
||||
self,
|
||||
base_claims: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
base_claims: Mapping[str, object],
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Build claims for the channel token (FR-14 two-token model).
|
||||
|
||||
|
|
@ -776,7 +795,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
# ------------------------------------------------------------------
|
||||
# FR-5: Verify incoming token before re-signing
|
||||
# ------------------------------------------------------------------
|
||||
jwt_claims: dict[str, Any] | None = None
|
||||
jwt_claims: dict[str, object] | None = None
|
||||
raw_token: Final[str | None] = hook_data.get("incoming_bearer_token")
|
||||
|
||||
if self.access_token_discovery_uri and raw_token:
|
||||
|
|
@ -810,7 +829,7 @@ class MCPJWTSigner(CustomGuardrail):
|
|||
|
||||
# Fall back to LiteLLM-decoded JWT claims (available when proxy uses JWT auth).
|
||||
if jwt_claims is None:
|
||||
jwt_claims = getattr(user_api_key_dict, "jwt_claims", None)
|
||||
jwt_claims = user_api_key_dict.jwt_claims
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# FR-15: Validate required claims
|
||||
|
|
@ -896,7 +915,7 @@ async def inject_mcp_jwt_headers_for_upstream(
|
|||
if auth_hdr.lower().startswith("bearer "):
|
||||
incoming_bearer_token = auth_hdr[len("bearer ") :]
|
||||
|
||||
hook_data: Final[dict[str, Any]] = {
|
||||
hook_data: Final = {
|
||||
"mcp_tool_name": "" if for_list_tools else mcp_tool_name,
|
||||
"incoming_bearer_token": incoming_bearer_token,
|
||||
"extra_headers": merged,
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ Provides real-time threat detection, DLP, URL filtering, content masking, and po
|
|||
import json
|
||||
import os
|
||||
import re
|
||||
from collections.abc import AsyncIterable, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
|
||||
from urllib.parse import urlparse
|
||||
|
|
@ -166,7 +167,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
GuardrailEventHooks.during_mcp_call: GuardrailEventHooks.during_call,
|
||||
}
|
||||
|
||||
def should_run_guardrail(self, data: Any, event_type: GuardrailEventHooks) -> bool:
|
||||
def should_run_guardrail(self, data: Mapping[str, object], event_type: GuardrailEventHooks) -> bool:
|
||||
if super().should_run_guardrail(data, event_type):
|
||||
return True
|
||||
compat: Final = self._MCP_COMPAT_MAP.get(event_type)
|
||||
|
|
@ -175,7 +176,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
return True
|
||||
return False
|
||||
|
||||
def _extract_text_from_messages(self, messages: list[dict[str, Any]]) -> str:
|
||||
def _extract_text_from_messages(self, messages: Sequence[Mapping[str, object]]) -> str:
|
||||
"""Extract text content from messages array."""
|
||||
if not isinstance(messages, list) or not messages:
|
||||
return ""
|
||||
|
|
@ -242,10 +243,10 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
self,
|
||||
content: str = "",
|
||||
is_response: bool = False,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
call_id: str | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
call_id: object = None,
|
||||
tool_event: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""Call PANW Prisma AIRS API to scan content or a tool_event."""
|
||||
|
||||
if tool_event is None and not content.strip():
|
||||
|
|
@ -275,7 +276,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
else:
|
||||
app_name_value = self.app_name # Defaults to "LiteLLM"
|
||||
|
||||
panw_metadata: Final = {
|
||||
panw_metadata: Final[dict[str, object]] = {
|
||||
"app_user": (
|
||||
(metadata.get("app_user") or metadata.get("user") or "litellm_user") if metadata else "litellm_user"
|
||||
),
|
||||
|
|
@ -295,13 +296,13 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
panw_metadata["litellm_trace_id"] = metadata["litellm_trace_id"]
|
||||
|
||||
# Build contents: tool_event takes priority, else prompt/response text
|
||||
contents: list[dict[str, Any]]
|
||||
contents: Sequence[Mapping[str, object]]
|
||||
if tool_event is not None:
|
||||
contents = [{"tool_event": tool_event}]
|
||||
else:
|
||||
contents = [{"response" if is_response else "prompt": content}]
|
||||
|
||||
payload: Final = {
|
||||
payload: Final[dict[str, object]] = {
|
||||
"metadata": panw_metadata,
|
||||
"contents": contents,
|
||||
}
|
||||
|
|
@ -325,7 +326,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
# If neither profile_name nor profile_id is provided, PANW API will use the
|
||||
# profile linked to the API key (if configured in Strata Cloud Manager)
|
||||
if profile_name or profile_id:
|
||||
ai_profile: Final = {}
|
||||
ai_profile: Final[dict[str, object]] = {}
|
||||
if profile_id:
|
||||
ai_profile["profile_id"] = profile_id
|
||||
if profile_name:
|
||||
|
|
@ -333,7 +334,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
payload["ai_profile"] = ai_profile
|
||||
|
||||
if is_response and tool_event is None:
|
||||
payload["metadata"]["is_response"] = True
|
||||
panw_metadata["is_response"] = True
|
||||
|
||||
headers: Final = {
|
||||
"Content-Type": "application/json",
|
||||
|
|
@ -355,7 +356,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
)
|
||||
response.raise_for_status()
|
||||
|
||||
result: Final = response.json()
|
||||
result: Final[dict[str, object]] = response.json()
|
||||
|
||||
# Validate response format
|
||||
if "action" not in result:
|
||||
|
|
@ -489,7 +490,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
)
|
||||
return "unknown"
|
||||
|
||||
def _get_masked_text(self, scan_result: dict[str, Any], is_response: bool = False) -> str | None:
|
||||
def _get_masked_text(self, scan_result: Mapping[str, object], is_response: bool = False) -> str | None:
|
||||
"""Extract masked text from PANW scan result."""
|
||||
masked_key: Final = "response_masked_data" if is_response else "prompt_masked_data"
|
||||
masked_data: Final = scan_result.get(masked_key)
|
||||
|
|
@ -511,7 +512,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
@staticmethod
|
||||
def _apply_mcp_masking(
|
||||
request_data: dict,
|
||||
original_args: Any,
|
||||
original_args: object,
|
||||
masked_text: str,
|
||||
*,
|
||||
is_blocked: bool = True,
|
||||
|
|
@ -544,7 +545,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
# If the original args were structured, preserve the type.
|
||||
if isinstance(original_args, (dict, list)):
|
||||
try:
|
||||
parsed: Final = json.loads(masked_text)
|
||||
parsed: Final[object] = json.loads(masked_text)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
|
|
@ -556,7 +557,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
}
|
||||
},
|
||||
)
|
||||
masked_value: Any = parsed
|
||||
masked_value: object = parsed
|
||||
else:
|
||||
masked_value = masked_text
|
||||
|
||||
|
|
@ -572,7 +573,9 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
else:
|
||||
verbose_proxy_logger.info("PANW Prisma AIRS: MCP request allowed with PII masking applied")
|
||||
|
||||
def _apply_masking_to_messages(self, messages: list[dict[str, Any]], masked_text: str) -> list[dict[str, Any]]:
|
||||
def _apply_masking_to_messages(
|
||||
self, messages: list[dict[str, object]], masked_text: str
|
||||
) -> Sequence[Mapping[str, object]]:
|
||||
"""Apply masked text to the last user message."""
|
||||
if not messages:
|
||||
return messages
|
||||
|
|
@ -622,7 +625,9 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
if hasattr(choice.message.function_call, "arguments"):
|
||||
choice.message.function_call.arguments = masked_text
|
||||
|
||||
def _build_error_detail(self, scan_result: dict[str, Any], is_response: bool = False) -> dict[str, Any]:
|
||||
def _build_error_detail(
|
||||
self, scan_result: Mapping[str, object], is_response: bool = False
|
||||
) -> Mapping[str, Mapping[str, object]]:
|
||||
"""Build enhanced error detail with scan information."""
|
||||
action_type: Final = "Response" if is_response else "Prompt"
|
||||
code_suffix: Final = "_response_blocked" if is_response else "_blocked"
|
||||
|
|
@ -642,7 +647,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
},
|
||||
)
|
||||
|
||||
error_detail: Final = {
|
||||
error_detail: Final[dict[str, dict[str, object]]] = {
|
||||
"error": {
|
||||
"message": error_msg,
|
||||
"type": "guardrail_violation",
|
||||
|
|
@ -672,12 +677,12 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
|
||||
def _handle_api_error_with_logging(
|
||||
self,
|
||||
scan_result: dict[str, Any],
|
||||
data: dict[str, Any],
|
||||
scan_result: dict[str, object],
|
||||
data: dict[str, object],
|
||||
start_time: datetime,
|
||||
event_type: GuardrailEventHooks,
|
||||
is_response: bool = False,
|
||||
) -> dict[str, Any] | None:
|
||||
) -> None:
|
||||
"""Handle API errors with fail-open/fail-closed logic."""
|
||||
end_time: Final = datetime.now()
|
||||
duration: Final = (end_time - start_time).total_seconds()
|
||||
|
|
@ -722,7 +727,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=f"{self.guardrail_name}:unscanned"
|
||||
)
|
||||
return None
|
||||
return
|
||||
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
|
|
@ -783,7 +788,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
return metadata
|
||||
|
||||
@staticmethod
|
||||
def _extract_text_from_sse_bytes(chunks: list[bytes]) -> str:
|
||||
def _extract_text_from_sse_bytes(chunks: Sequence[bytes]) -> str:
|
||||
"""Extract text from Anthropic SSE byte chunks (content_block_delta → text_delta)."""
|
||||
texts: Final[list[str]] = []
|
||||
raw: Final = b"".join(chunks).decode("utf-8", errors="replace")
|
||||
|
|
@ -804,7 +809,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
return "".join(texts)
|
||||
|
||||
@staticmethod
|
||||
def _extract_text_from_streaming_events(chunks: list) -> str:
|
||||
def _extract_text_from_streaming_events(chunks: Sequence[object]) -> str:
|
||||
"""Extract text from /v1/responses streaming events (object or dict)."""
|
||||
|
||||
def _attr(c, key):
|
||||
|
|
@ -960,7 +965,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
cache: DualCache,
|
||||
data: dict[str, Any],
|
||||
call_type: CallTypesLiteral,
|
||||
) -> dict[str, Any] | None:
|
||||
) -> dict[str, object] | None:
|
||||
"""
|
||||
Pre-call hook to scan user prompts before sending to LLM.
|
||||
|
||||
|
|
@ -1075,10 +1080,10 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
@log_guardrail_information
|
||||
async def async_post_call_success_hook(
|
||||
self,
|
||||
data: dict[str, Any],
|
||||
data: dict[str, object],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: Any,
|
||||
) -> Any:
|
||||
response: object,
|
||||
) -> object:
|
||||
"""
|
||||
Post-call hook to scan LLM responses before returning to user.
|
||||
|
||||
|
|
@ -1193,7 +1198,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
assembled_model_response: ModelResponse,
|
||||
request_data: dict,
|
||||
start_time: datetime,
|
||||
) -> tuple[bool, ModelResponse, dict[str, Any]]:
|
||||
) -> tuple[bool, ModelResponse, dict[str, object]]:
|
||||
"""
|
||||
Scan assembled streaming response and apply masking if needed.
|
||||
Returns (content_was_modified, response, scan_result).
|
||||
|
|
@ -1255,8 +1260,8 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: Any,
|
||||
request_data: dict,
|
||||
response: AsyncIterable[object],
|
||||
request_data: dict[str, object],
|
||||
):
|
||||
"""
|
||||
Process streaming response chunks and scan the assembled response.
|
||||
|
|
@ -1367,7 +1372,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
# returns a proper JSON error response with the correct status code.
|
||||
# (Raising from a generator hits create_response's generic except → 500.)
|
||||
detail: Final = e.detail if isinstance(e.detail, dict) else {"message": str(e.detail)}
|
||||
error_obj: Final[dict[str, Any]] = dict(detail.get("error", detail))
|
||||
error_obj: Final[dict[str, object]] = dict(detail.get("error", detail))
|
||||
error_obj["code"] = e.status_code
|
||||
yield f"data: {json.dumps({'error': error_obj})}\n\n"
|
||||
except Exception as e:
|
||||
|
|
@ -1378,8 +1383,8 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
self,
|
||||
tool_calls: list,
|
||||
is_response: bool,
|
||||
metadata: dict[str, Any],
|
||||
call_id: str,
|
||||
metadata: Mapping[str, object],
|
||||
call_id: object,
|
||||
request_data: dict,
|
||||
start_time: datetime,
|
||||
) -> None:
|
||||
|
|
@ -1416,7 +1421,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
tool_name = func.get("name")
|
||||
|
||||
# --- build tool_event payload (canonical PANW schema) -----------
|
||||
tool_event: dict[str, Any] = {
|
||||
tool_event: dict[str, object] = {
|
||||
"metadata": {
|
||||
"ecosystem": "openai",
|
||||
"method": "tools/call",
|
||||
|
|
@ -1472,7 +1477,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
|
||||
@staticmethod
|
||||
def _is_anthropic_request(
|
||||
request_data: dict,
|
||||
request_data: Mapping[str, object],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> bool:
|
||||
"""Detect if the current request is an Anthropic /v1/messages call."""
|
||||
|
|
@ -1497,7 +1502,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
|
||||
def _use_latest_user_only(
|
||||
self,
|
||||
request_data: dict,
|
||||
request_data: Mapping[str, object],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> bool:
|
||||
"""Resolve whether to scan only the latest user message.
|
||||
|
|
@ -1515,8 +1520,8 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
|
||||
@staticmethod
|
||||
def _get_latest_user_text_indices(
|
||||
texts: list[str],
|
||||
messages: list,
|
||||
texts: Sequence[str],
|
||||
messages: Sequence[object],
|
||||
) -> set | None:
|
||||
"""Return text indices belonging to only the latest scannable human-authored (user or developer) message.
|
||||
|
||||
|
|
@ -1569,8 +1574,8 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
|
||||
@staticmethod
|
||||
def _get_scannable_text_indices(
|
||||
texts: list[str],
|
||||
structured_messages: list,
|
||||
texts: Sequence[str],
|
||||
structured_messages: Sequence[object],
|
||||
) -> set | None:
|
||||
"""Derive which ``texts`` indices originate from user/system messages.
|
||||
|
||||
|
|
@ -1627,7 +1632,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
request_data: dict[str, object],
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
|
|
@ -1798,7 +1803,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
# "mcp_tool_name"/"mcp_arguments". Check canonical first, then fallback.
|
||||
mcp_tool_name: Final = request_data.get("mcp_tool_name") or self._mcp_name_fallback(request_data)
|
||||
if mcp_tool_name and input_type == "request":
|
||||
mcp_tool_event: Final[dict[str, Any]] = {
|
||||
mcp_tool_event: Final[dict[str, object]] = {
|
||||
"metadata": {
|
||||
"ecosystem": "mcp",
|
||||
"method": "tools/call",
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ Pre-call hook that filters MCP tools semantically before LLM inference.
|
|||
Reduces context window size and improves tool selection accuracy.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -164,7 +165,7 @@ class SemanticToolFilterHook(CustomLogger):
|
|||
return [name for name in names if name]
|
||||
|
||||
@staticmethod
|
||||
def _narrow_mcp_references(tools: list[Any], selected_tool_names: list[str]) -> list[Any]:
|
||||
def _narrow_mcp_references(tools: Sequence[Mapping[str, object]], selected_tool_names: list[str]) -> list[object]:
|
||||
"""
|
||||
Restrict each litellm_proxy MCP reference to the semantically selected tools.
|
||||
|
||||
|
|
|
|||
|
|
@ -31,6 +31,8 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_utils import (
|
||||
ESTIMATED_OUTPUT_TOKENS_FIELD,
|
||||
get_estimated_output_tokens,
|
||||
get_key_tag_rpm_limit,
|
||||
get_model_rate_limit_from_metadata,
|
||||
)
|
||||
|
|
@ -562,6 +564,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
data: dict,
|
||||
model: str | None = None,
|
||||
min_configured_tpm_limit: int | None = None,
|
||||
configured_output_tokens: int | None = None,
|
||||
) -> int:
|
||||
"""
|
||||
Estimate total tokens this request will consume so we can reserve them
|
||||
|
|
@ -575,6 +578,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
provided, the no-``max_tokens`` output-budget floor is capped at a
|
||||
fraction of that limit so small TPM caps remain usable. Omit to
|
||||
preserve the unconstrained floor.
|
||||
|
||||
``configured_output_tokens`` is the operator-declared estimate resolved
|
||||
from key or team metadata. When provided it replaces the heuristic
|
||||
floor entirely, so the reservation reflects what this tenant's model
|
||||
actually emits rather than one constant shared by every tenant.
|
||||
"""
|
||||
messages = data.get("messages")
|
||||
prompt: Final = data.get("prompt")
|
||||
|
|
@ -604,7 +612,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
case (_, embeddings_input) if embeddings_input:
|
||||
# Embeddings have no output tokens
|
||||
max_tokens_estimate = 0
|
||||
case _ if total_chars == 0:
|
||||
case _ if total_chars == 0 and configured_output_tokens is None:
|
||||
# Fully contentless request (no messages, prompt, or input).
|
||||
# Don't apply the conservative output-budget floor here — it
|
||||
# would over-reserve and could push small TPM limits into a
|
||||
|
|
@ -619,7 +627,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
# so a small per-tenant TPM cap can't be tripped by the floor
|
||||
# alone.
|
||||
output_floor: Final = self._no_max_tokens_output_floor(min_configured_tpm_limit)
|
||||
max_tokens_estimate = max(estimated_input_tokens, output_floor)
|
||||
max_tokens_estimate = (
|
||||
configured_output_tokens
|
||||
if configured_output_tokens is not None
|
||||
else max(estimated_input_tokens, output_floor)
|
||||
)
|
||||
|
||||
total_estimated: Final = estimated_input_tokens + max_tokens_estimate
|
||||
|
||||
|
|
@ -2586,8 +2598,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
data.get("max_tokens") is not None or data.get("max_completion_tokens") is not None
|
||||
)
|
||||
is_embedding: Final = data.get("input") is not None
|
||||
configured_output_tokens: Final = get_estimated_output_tokens(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
model_name=requested_model,
|
||||
)
|
||||
if capped_floor < baseline_floor and not has_explicit_max_tokens and not is_embedding:
|
||||
data["max_tokens"] = capped_floor
|
||||
data["max_tokens"] = max(capped_floor, configured_output_tokens or 0)
|
||||
|
||||
# Floor at 1 token so contentless requests (/responses,
|
||||
# tool-call continuations, empty messages) still flow
|
||||
|
|
@ -2601,10 +2617,23 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
data=data,
|
||||
model=requested_model,
|
||||
min_configured_tpm_limit=min_configured_tpm_limit,
|
||||
configured_output_tokens=configured_output_tokens,
|
||||
),
|
||||
1,
|
||||
)
|
||||
|
||||
if configured_output_tokens is not None and estimated_tokens > min_configured_tpm_limit:
|
||||
verbose_proxy_logger.debug(
|
||||
"Reserving %s tokens for model %s (declared %s=%s plus the input estimate) exceeds the "
|
||||
"smallest TPM limit this request is charged against (%s), so it cannot be admitted even "
|
||||
"against an empty window. Lower the declared estimate or raise the TPM limit.",
|
||||
estimated_tokens,
|
||||
requested_model,
|
||||
ESTIMATED_OUTPUT_TOKENS_FIELD,
|
||||
configured_output_tokens,
|
||||
min_configured_tpm_limit,
|
||||
)
|
||||
|
||||
tpm_response: Final = await self.reserve_tpm_tokens(
|
||||
descriptors=descriptors,
|
||||
estimated_tokens=estimated_tokens,
|
||||
|
|
|
|||
|
|
@ -48,6 +48,15 @@ _UNATTRIBUTED_TRACKABLE_CALL_TYPES: Final[frozenset[str]] = frozenset(
|
|||
}
|
||||
)
|
||||
|
||||
# Both spellings, because call_type reaches the callback as str(...) of either the
|
||||
# enum member or its value.
|
||||
_CAPTURED_IDENTITY_CALL_TYPES: Final[frozenset[str]] = frozenset(
|
||||
(
|
||||
CallTypes.aretrieve_batch.value,
|
||||
str(CallTypes.aretrieve_batch),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class _ProxyDBLogger(CustomLogger):
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
|
|
@ -212,7 +221,10 @@ class _ProxyDBLogger(CustomLogger):
|
|||
# Only fetch key details when user_id wasn't already populated (e.g. direct MCP REST calls).
|
||||
# Avoids a cache/DB lookup on every normal LLM request.
|
||||
if metadata.get("user_api_key") and not metadata.get("user_api_key_user_id"):
|
||||
metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata=metadata)
|
||||
metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( # rebind-ok: enriched metadata replaces the original
|
||||
metadata=metadata,
|
||||
resolve_missing_key_identity=str(kwargs.get("call_type")) not in _CAPTURED_IDENTITY_CALL_TYPES,
|
||||
)
|
||||
_write_spend_metadata_to_kwargs(kwargs=kwargs, metadata=metadata)
|
||||
budget_reservation: Final = _get_budget_reservation_from_metadata(metadata=metadata)
|
||||
user_id: Final = cast(str | None, metadata.get("user_api_key_user_id", None))
|
||||
|
|
@ -337,7 +349,7 @@ class _ProxyDBLogger(CustomLogger):
|
|||
spend_log_error("Error in tracking cost callback - %s", str(e), exc=e)
|
||||
|
||||
@staticmethod
|
||||
async def _enrich_failure_metadata_with_key_info(metadata: dict) -> dict:
|
||||
async def _enrich_failure_metadata_with_key_info(metadata: dict, resolve_missing_key_identity: bool = True) -> dict:
|
||||
"""
|
||||
Enriches failure spend log metadata by looking up the key object (and team object)
|
||||
from cache/DB when key fields are missing.
|
||||
|
|
@ -349,6 +361,11 @@ class _ProxyDBLogger(CustomLogger):
|
|||
2. Post-auth failures (provider errors, rate limits): key fields are populated
|
||||
but team_alias is missing because LiteLLM_VerificationTokenView SQL view
|
||||
doesn't include it. We look up the team object to fill in team_alias.
|
||||
|
||||
Scenario 1 reads the key's identity as it stands right now, so it is only correct
|
||||
for a log emitted within the request it describes. Callers that log after a delay,
|
||||
against an identity captured earlier, pass resolve_missing_key_identity=False and
|
||||
keep their own user_id, team_id and org_id.
|
||||
"""
|
||||
api_key_hash: Final = metadata.get("user_api_key")
|
||||
if not api_key_hash:
|
||||
|
|
@ -361,7 +378,7 @@ class _ProxyDBLogger(CustomLogger):
|
|||
)
|
||||
|
||||
# Step 1: If key fields are missing, look up the full key object
|
||||
if metadata.get("user_api_key_alias") is None:
|
||||
if resolve_missing_key_identity and metadata.get("user_api_key_alias") is None:
|
||||
try:
|
||||
key_obj: Final = await get_key_object(
|
||||
hashed_token=api_key_hash,
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ import re
|
|||
import time
|
||||
from collections import OrderedDict
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
|
|
@ -66,6 +67,32 @@ _SESSION_ID_VALUE_RE: Final = re.compile(r"^[a-zA-Z0-9_\-]{8,}$")
|
|||
|
||||
_SHA256_HEX_RE: Final = re.compile(r"^[0-9a-f]{64}$")
|
||||
|
||||
# W3C Trace Context traceparent header: https://www.w3.org/TR/trace-context/
|
||||
# e.g. "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"
|
||||
_TRACEPARENT_RE: Final = re.compile(r"^[0-9a-f]{2}-([0-9a-f]{32})-[0-9a-f]{16}-[0-9a-f]{2}$", re.IGNORECASE)
|
||||
|
||||
|
||||
def _trace_id_from_traceparent(traceparent: str) -> str | None:
|
||||
"""Extract the trace-id from a W3C Trace Context traceparent header, e.g.
|
||||
"00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01" -> the 32-hex
|
||||
trace-id in the middle. An all-zero trace-id is invalid per spec and is
|
||||
rejected, matching how the OpenTelemetry SDK itself treats it."""
|
||||
match: Final = _TRACEPARENT_RE.match(traceparent.strip())
|
||||
if not match:
|
||||
return None
|
||||
trace_id: Final = match.group(1).lower()
|
||||
return trace_id if trace_id != "0" * 32 else None
|
||||
|
||||
|
||||
def _session_id_from_baggage(baggage: str) -> str | None:
|
||||
"""Extract a session.id entry from a W3C Baggage header
|
||||
(https://www.w3.org/TR/baggage/), e.g. "session.id=abc-123,user.id=42"."""
|
||||
for pair in baggage.split(","):
|
||||
key, _, value = pair.strip().partition("=")
|
||||
if key.strip() == "session.id" and value.strip():
|
||||
return value.strip()
|
||||
return None
|
||||
|
||||
|
||||
def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None:
|
||||
"""Only proxy-validated keys are stamped, proven by the unforgeable
|
||||
|
|
@ -210,6 +237,11 @@ _UNTRUSTED_ROOT_CONTROL_FIELDS: Final = (
|
|||
"_code_interpreter_interception_sandbox_key",
|
||||
"_code_interpreter_interception_session_scoped",
|
||||
"max_agentic_loops",
|
||||
# Recomputed below from the actual caller-controlled timeout sources (headers and
|
||||
# body fields); a client-forged value here would let a request either dodge cooldown
|
||||
# protection on a real deployment failure or force a false "not caller-controlled"
|
||||
# reading that lets its own bad timeout cool down deployments other tenants rely on.
|
||||
"client_side_timeout",
|
||||
)
|
||||
|
||||
_UNTRUSTED_METADATA_CONTROL_FIELDS: Final = (
|
||||
|
|
@ -850,6 +882,19 @@ class LiteLLMProxyRequestSetup:
|
|||
return float(stream_timeout_header)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _get_keepalive_seconds_from_request(headers: Mapping[str, str]) -> float | None:
|
||||
"""
|
||||
Get `keepalive_seconds` from the request headers, for clients (e.g. the
|
||||
Vercel AI SDK) that can set custom headers more easily than extra body
|
||||
fields. Subject to the same deployment-level allow_client_keepalive_override
|
||||
gate as the request body field: see _resolve_keepalive_seconds.
|
||||
"""
|
||||
keepalive_seconds_header: Final = headers.get("x-litellm-keepalive-seconds", None)
|
||||
if keepalive_seconds_header is not None:
|
||||
return float(keepalive_seconds_header)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _get_num_retries_from_request(headers: dict) -> int | None:
|
||||
"""
|
||||
|
|
@ -1035,6 +1080,7 @@ class LiteLLMProxyRequestSetup:
|
|||
def add_litellm_data_for_backend_llm_call(
|
||||
*,
|
||||
headers: dict,
|
||||
request_data: Mapping[str, Any],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
general_settings: dict[str, Any] | None = None,
|
||||
) -> LitellmDataForBackendLLMCall:
|
||||
|
|
@ -1053,18 +1099,38 @@ class LiteLLMProxyRequestSetup:
|
|||
if _organization is not None:
|
||||
data["organization"] = _organization
|
||||
|
||||
timeout: Final = LiteLLMProxyRequestSetup._get_timeout_from_request(headers)
|
||||
if timeout is not None:
|
||||
data["timeout"] = timeout
|
||||
header_timeout: Final = LiteLLMProxyRequestSetup._get_timeout_from_request(headers)
|
||||
if header_timeout is not None:
|
||||
data["timeout"] = header_timeout
|
||||
|
||||
stream_timeout: Final = LiteLLMProxyRequestSetup._get_stream_timeout_from_request(headers)
|
||||
if stream_timeout is not None:
|
||||
data["stream_timeout"] = stream_timeout
|
||||
header_stream_timeout: Final = LiteLLMProxyRequestSetup._get_stream_timeout_from_request(headers)
|
||||
if header_stream_timeout is not None:
|
||||
data["stream_timeout"] = header_stream_timeout
|
||||
|
||||
# Router._get_timeout resolves the effective per-attempt timeout from any of
|
||||
# kwargs["timeout"], kwargs["request_timeout"], or kwargs["stream_timeout"], and a
|
||||
# caller can supply any of those via the request body as well as the headers above.
|
||||
# A deliberately tiny value can force a 408 on every deployment in a fallback chain,
|
||||
# so this marker (never trusted verbatim from the client; stripped above) must cover
|
||||
# every source cooldown_handlers._trigger_cooldown_for_failed_deployment needs to
|
||||
# distinguish from a real deployment health signal.
|
||||
if (
|
||||
header_timeout is not None
|
||||
or header_stream_timeout is not None
|
||||
or request_data.get("timeout") is not None
|
||||
or request_data.get("request_timeout") is not None
|
||||
or request_data.get("stream_timeout") is not None
|
||||
):
|
||||
data["client_side_timeout"] = True
|
||||
|
||||
num_retries: Final = LiteLLMProxyRequestSetup._get_num_retries_from_request(headers)
|
||||
if num_retries is not None:
|
||||
data["num_retries"] = num_retries
|
||||
|
||||
keepalive_seconds: Final = LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request(headers)
|
||||
if keepalive_seconds is not None:
|
||||
data["keepalive_seconds"] = keepalive_seconds
|
||||
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -1113,6 +1179,33 @@ class LiteLLMProxyRequestSetup:
|
|||
body_metadata["user_id"] = session_id
|
||||
verbose_proxy_logger.debug("Extracted session_id from Anthropic metadata.user_id")
|
||||
|
||||
# Last-resort fallback: the W3C standards for trace/session propagation
|
||||
# (https://www.w3.org/TR/trace-context/, https://www.w3.org/TR/baggage/).
|
||||
# Lower priority than everything above - only fires when neither the
|
||||
# explicit litellm headers nor the Anthropic-metadata path found
|
||||
# anything - but lets a caller's existing traceparent/baggage headers
|
||||
# (from real OTel instrumentation) correlate with litellm's own logs
|
||||
# instead of generating an unrelated trace_id.
|
||||
normalized_headers: Final = MappingProxyType({k.lower(): v for k, v in headers.items() if isinstance(k, str)})
|
||||
if "litellm_trace_id" not in data:
|
||||
traceparent: Final = normalized_headers.get("traceparent")
|
||||
if isinstance(traceparent, str):
|
||||
trace_id_from_traceparent: Final = _trace_id_from_traceparent(traceparent)
|
||||
if trace_id_from_traceparent:
|
||||
metadata_from_headers["trace_id"] = trace_id_from_traceparent
|
||||
data["litellm_trace_id"] = trace_id_from_traceparent # rebind-ok: data is an out-param
|
||||
verbose_proxy_logger.debug(
|
||||
"Extracted trace_id from W3C traceparent header: %s", trace_id_from_traceparent
|
||||
)
|
||||
if "litellm_session_id" not in data:
|
||||
baggage: Final = normalized_headers.get("baggage")
|
||||
if isinstance(baggage, str):
|
||||
session_id_from_baggage: Final = _session_id_from_baggage(baggage)
|
||||
if session_id_from_baggage:
|
||||
metadata_from_headers["session_id"] = session_id_from_baggage
|
||||
data["litellm_session_id"] = session_id_from_baggage # rebind-ok: data is an out-param
|
||||
verbose_proxy_logger.debug("Extracted session_id from W3C baggage header")
|
||||
|
||||
if isinstance(data[_metadata_variable_name], dict):
|
||||
data[_metadata_variable_name].update(metadata_from_headers)
|
||||
return data
|
||||
|
|
@ -1545,6 +1638,7 @@ async def add_litellm_data_to_request(
|
|||
data.update(
|
||||
LiteLLMProxyRequestSetup.add_litellm_data_for_backend_llm_call(
|
||||
headers=_headers,
|
||||
request_data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
general_settings=general_settings,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -20,7 +20,10 @@ from fastapi import APIRouter, Depends, HTTPException
|
|||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
|
||||
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
|
||||
from litellm.proxy.management_endpoints.common_utils import (
|
||||
_user_has_admin_view,
|
||||
validate_budget_duration,
|
||||
)
|
||||
from litellm.proxy.utils import jsonify_object
|
||||
from litellm.repositories.budget_repository import BudgetRepository
|
||||
|
||||
|
|
@ -72,6 +75,8 @@ async def new_budget(
|
|||
detail={"error": f"soft_budget must be a non-negative finite number. Received: {budget_obj.soft_budget}"},
|
||||
)
|
||||
|
||||
validate_budget_duration(budget_obj.budget_duration)
|
||||
|
||||
# Validate model_max_budget if present
|
||||
if budget_obj.model_max_budget is not None and len(budget_obj.model_max_budget) > 0:
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
|
|
@ -153,6 +158,8 @@ async def update_budget(
|
|||
detail={"error": f"soft_budget must be a non-negative finite number. Received: {budget_obj.soft_budget}"},
|
||||
)
|
||||
|
||||
validate_budget_duration(budget_obj.budget_duration)
|
||||
|
||||
# Validate model_max_budget if present in update
|
||||
if budget_obj.model_max_budget is not None and len(budget_obj.model_max_budget) > 0:
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ import asyncio
|
|||
import json
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol
|
||||
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
|
@ -37,8 +37,26 @@ from litellm.types.management_endpoints import (
|
|||
CacheSettingsField,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
||||
class _CacheConfigRow(Protocol):
|
||||
cache_settings: str | Mapping[str, object] | None
|
||||
|
||||
|
||||
class _CacheConfigTable(Protocol):
|
||||
async def find_unique(self, where: Mapping[str, str]) -> _CacheConfigRow | None: ...
|
||||
|
||||
async def upsert(self, where: Mapping[str, str], data: Mapping[str, Mapping[str, str]]) -> _CacheConfigRow: ...
|
||||
|
||||
|
||||
def _cache_config_table(prisma_client: "PrismaClient") -> _CacheConfigTable:
|
||||
return CacheConfigRepository(prisma_client).table
|
||||
|
||||
|
||||
# Cache fields holding credentials. Masked on read so plaintext Redis /
|
||||
# Sentinel passwords never leave the server in a GET response. `url` is here
|
||||
# because a Redis/Valkey URL can embed a password inline
|
||||
|
|
@ -197,7 +215,7 @@ def _saved_secret_is_reusable(incoming: Mapping[str, object], saved: Mapping[str
|
|||
return True
|
||||
|
||||
|
||||
def _merge_over_saved(incoming: Mapping[str, object], saved: Mapping[str, object]) -> dict[str, Any]:
|
||||
def _merge_over_saved(incoming: Mapping[str, object], saved: Mapping[str, object]) -> Mapping[str, object]:
|
||||
"""Keep the stored secret behind any credential the caller echoed back redacted or omitted.
|
||||
|
||||
GET returns credentials as the marker and the form never re-prefills a
|
||||
|
|
@ -339,7 +357,7 @@ class CacheSettingsManager:
|
|||
return normalized1 == normalized2
|
||||
|
||||
@staticmethod
|
||||
async def init_cache_settings_in_db(prisma_client, proxy_config):
|
||||
async def init_cache_settings_in_db(prisma_client: "PrismaClient", proxy_config):
|
||||
"""
|
||||
Initialize cache settings from database into the router on startup.
|
||||
Only reinitializes if cache params have changed.
|
||||
|
|
@ -349,7 +367,7 @@ class CacheSettingsManager:
|
|||
try:
|
||||
cache_config: Final = await call_with_db_reconnect_retry(
|
||||
prisma_client,
|
||||
lambda: CacheConfigRepository(prisma_client).table.find_unique(where={"id": "cache_config"}),
|
||||
lambda: _cache_config_table(prisma_client).find_unique(where={"id": "cache_config"}),
|
||||
reason="init_cache_settings_in_db_lookup_failure",
|
||||
)
|
||||
if cache_config is not None and cache_config.cache_settings:
|
||||
|
|
@ -444,7 +462,7 @@ async def get_cache_settings(
|
|||
# Read the stored settings (decrypted); an env-only cache has none.
|
||||
stored: dict[str, object] = {}
|
||||
if prisma_client is not None:
|
||||
cache_config = await CacheConfigRepository(prisma_client).table.find_unique(where={"id": "cache_config"})
|
||||
cache_config = await _cache_config_table(prisma_client).find_unique(where={"id": "cache_config"})
|
||||
if cache_config is not None and cache_config.cache_settings:
|
||||
stored = proxy_config._decrypt_db_variables(
|
||||
variables_dict=_parse_stored_settings(cache_config.cache_settings)
|
||||
|
|
@ -511,9 +529,7 @@ async def test_cache_connection(
|
|||
saved_settings: dict[str, object] = {}
|
||||
if prisma_client is not None:
|
||||
try:
|
||||
existing_row: Final = await CacheConfigRepository(prisma_client).table.find_unique(
|
||||
where={"id": "cache_config"}
|
||||
)
|
||||
existing_row: Final = await _cache_config_table(prisma_client).find_unique(where={"id": "cache_config"})
|
||||
if existing_row is not None and existing_row.cache_settings:
|
||||
saved_settings = proxy_config._decrypt_db_variables(
|
||||
variables_dict=_parse_stored_settings(existing_row.cache_settings)
|
||||
|
|
@ -590,7 +606,7 @@ async def update_cache_settings(
|
|||
try:
|
||||
# Read the stored row first: its decrypted values back any credential the
|
||||
# caller echoed back redacted, and its key set drives the audit diff.
|
||||
existing_row: Final = await CacheConfigRepository(prisma_client).table.find_unique(where={"id": "cache_config"})
|
||||
existing_row: Final = await _cache_config_table(prisma_client).find_unique(where={"id": "cache_config"})
|
||||
before_settings: dict[str, object] | None = None
|
||||
saved_settings: dict[str, object] = {}
|
||||
if existing_row is not None and existing_row.cache_settings:
|
||||
|
|
@ -606,7 +622,7 @@ async def update_cache_settings(
|
|||
encrypted_settings: Final = proxy_config._encrypt_env_variables(environment_variables=cache_settings)
|
||||
|
||||
# Save to database
|
||||
await CacheConfigRepository(prisma_client).table.upsert(
|
||||
await _cache_config_table(prisma_client).upsert(
|
||||
where={"id": "cache_config"},
|
||||
data={
|
||||
"create": {
|
||||
|
|
|
|||
|
|
@ -8,7 +8,9 @@ from fastapi import HTTPException, status
|
|||
from typing_extensions import TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import PTU_SENTINEL_API_KEY
|
||||
from litellm.proxy._types import CommonProxyErrors
|
||||
from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.repositories.table_repositories import DeletedVerificationTokenRepository
|
||||
from litellm.repositories.verification_token_repository import (
|
||||
|
|
@ -140,6 +142,28 @@ class _GroupingSetsRow(SimpleNamespace):
|
|||
failed_requests: int | None
|
||||
|
||||
|
||||
def _reported_flat_cost(record: DailySpendRecord | _GroupingSetsRow) -> float:
|
||||
"""Flat cost a daily row reports, which is zero unless PTU cost attribution is enabled.
|
||||
|
||||
Both read paths funnel through here: the paginated path reads the ``ptu_flat_cost``
|
||||
column straight off the row, and the aggregated path reads the SUM() alias. Rows an
|
||||
operator accrued during an earlier opt-in stay in the table, so the gate lives on the
|
||||
read rather than on the query that produced the rows.
|
||||
|
||||
The row is checked before the flag because this runs once per metric accumulation, and
|
||||
a record fans out across roughly a dozen breakdowns. The flag reads through the secret
|
||||
manager, uncached, so consulting it for every accumulation put thousands of lookups on
|
||||
a shared endpoint that made none before. Only a row actually carrying flat cost, which
|
||||
is a sentinel row, reaches it now.
|
||||
"""
|
||||
raw: Final = getattr(record, "ptu_flat_cost", None) or 0.0
|
||||
if not raw:
|
||||
return 0.0
|
||||
if not is_ptu_cost_attribution_enabled():
|
||||
return 0.0
|
||||
return raw
|
||||
|
||||
|
||||
def update_metrics(existing_metrics: SpendMetrics, record: DailySpendRecord) -> SpendMetrics:
|
||||
"""Update metrics with new record data.
|
||||
|
||||
|
|
@ -150,6 +174,7 @@ def update_metrics(existing_metrics: SpendMetrics, record: DailySpendRecord) ->
|
|||
prompt_tokens: Final = record.prompt_tokens or 0
|
||||
completion_tokens: Final = record.completion_tokens or 0
|
||||
existing_metrics.spend += record.spend or 0.0
|
||||
existing_metrics.flat_cost += _reported_flat_cost(record)
|
||||
existing_metrics.prompt_tokens += prompt_tokens
|
||||
existing_metrics.completion_tokens += completion_tokens
|
||||
existing_metrics.total_tokens += prompt_tokens + completion_tokens
|
||||
|
|
@ -208,30 +233,43 @@ def update_breakdown_metrics(
|
|||
entity_id_field: str | None = None,
|
||||
entity_metadata_field: Mapping[str, dict[str, object]] | None = None,
|
||||
) -> BreakdownMetrics:
|
||||
"""Updates breakdown metrics for a single record using the existing update_metrics function"""
|
||||
"""Updates breakdown metrics for a single record using the existing update_metrics function.
|
||||
|
||||
PTU sentinel rows (api_key == PTU_SENTINEL_API_KEY) add their flat cost to every
|
||||
parent bucket but never appear as an api_key row, and are kept out of the
|
||||
per-request provider breakdown."""
|
||||
|
||||
is_ptu_sentinel: Final = record.api_key == PTU_SENTINEL_API_KEY
|
||||
|
||||
# A PTU sentinel row keys on the deployment id so a rename cannot move it, and carries
|
||||
# the operator-facing name in model_group. The breakdown key is rendered directly as a
|
||||
# label, so display the name; two deployments sharing one name merge here, which is
|
||||
# what the write path used to do by collapsing them into a single row.
|
||||
model_key: Final = (record.model_group or record.model) if is_ptu_sentinel else record.model
|
||||
|
||||
# Update model breakdown
|
||||
if record.model and record.model not in breakdown.models:
|
||||
breakdown.models[record.model] = MetricWithMetadata(
|
||||
if model_key and model_key not in breakdown.models:
|
||||
breakdown.models[model_key] = MetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=model_metadata.get(record.model, {}), # Add any model-specific metadata here
|
||||
metadata=model_metadata.get(model_key, {}), # Add any model-specific metadata here
|
||||
)
|
||||
if record.model:
|
||||
breakdown.models[record.model].metrics = update_metrics(breakdown.models[record.model].metrics, record)
|
||||
if model_key:
|
||||
breakdown.models[model_key].metrics = update_metrics(breakdown.models[model_key].metrics, record)
|
||||
|
||||
# Update API key breakdown for this model
|
||||
if record.api_key not in breakdown.models[record.model].api_key_breakdown:
|
||||
breakdown.models[record.model].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
),
|
||||
if not is_ptu_sentinel:
|
||||
# Update API key breakdown for this model
|
||||
if record.api_key not in breakdown.models[model_key].api_key_breakdown:
|
||||
breakdown.models[model_key].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
),
|
||||
)
|
||||
breakdown.models[model_key].api_key_breakdown[record.api_key].metrics = update_metrics(
|
||||
breakdown.models[model_key].api_key_breakdown[record.api_key].metrics,
|
||||
record,
|
||||
)
|
||||
breakdown.models[record.model].api_key_breakdown[record.api_key].metrics = update_metrics(
|
||||
breakdown.models[record.model].api_key_breakdown[record.api_key].metrics,
|
||||
record,
|
||||
)
|
||||
|
||||
# Update model group breakdown
|
||||
model_group_key: Final = record.model_group or record.model
|
||||
|
|
@ -245,19 +283,20 @@ def update_breakdown_metrics(
|
|||
breakdown.model_groups[model_group_key].metrics, record
|
||||
)
|
||||
|
||||
# Update API key breakdown for this model
|
||||
if record.api_key not in breakdown.model_groups[model_group_key].api_key_breakdown:
|
||||
breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
),
|
||||
if not is_ptu_sentinel:
|
||||
# Update API key breakdown for this model
|
||||
if record.api_key not in breakdown.model_groups[model_group_key].api_key_breakdown:
|
||||
breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
),
|
||||
)
|
||||
breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key].metrics = update_metrics(
|
||||
breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key].metrics,
|
||||
record,
|
||||
)
|
||||
breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key].metrics = update_metrics(
|
||||
breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key].metrics,
|
||||
record,
|
||||
)
|
||||
|
||||
if record.mcp_namespaced_tool_name:
|
||||
if record.mcp_namespaced_tool_name not in breakdown.mcp_servers:
|
||||
|
|
@ -288,28 +327,29 @@ def update_breakdown_metrics(
|
|||
record,
|
||||
)
|
||||
|
||||
# Update provider breakdown
|
||||
provider: Final = record.custom_llm_provider or "unknown"
|
||||
if provider not in breakdown.providers:
|
||||
breakdown.providers[provider] = MetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=provider_metadata.get(provider, {}), # Add any provider-specific metadata here
|
||||
)
|
||||
breakdown.providers[provider].metrics = update_metrics(breakdown.providers[provider].metrics, record)
|
||||
if not is_ptu_sentinel:
|
||||
# Update provider breakdown
|
||||
provider: Final = record.custom_llm_provider or "unknown"
|
||||
if provider not in breakdown.providers:
|
||||
breakdown.providers[provider] = MetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=provider_metadata.get(provider, {}), # Add any provider-specific metadata here
|
||||
)
|
||||
breakdown.providers[provider].metrics = update_metrics(breakdown.providers[provider].metrics, record)
|
||||
|
||||
# Update API key breakdown for this provider
|
||||
if record.api_key not in breakdown.providers[provider].api_key_breakdown:
|
||||
breakdown.providers[provider].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
),
|
||||
# Update API key breakdown for this provider
|
||||
if record.api_key not in breakdown.providers[provider].api_key_breakdown:
|
||||
breakdown.providers[provider].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
),
|
||||
)
|
||||
breakdown.providers[provider].api_key_breakdown[record.api_key].metrics = update_metrics(
|
||||
breakdown.providers[provider].api_key_breakdown[record.api_key].metrics,
|
||||
record,
|
||||
)
|
||||
breakdown.providers[provider].api_key_breakdown[record.api_key].metrics = update_metrics(
|
||||
breakdown.providers[provider].api_key_breakdown[record.api_key].metrics,
|
||||
record,
|
||||
)
|
||||
|
||||
# Update endpoint breakdown
|
||||
if record.endpoint:
|
||||
|
|
@ -336,16 +376,17 @@ def update_breakdown_metrics(
|
|||
record,
|
||||
)
|
||||
|
||||
# Update api key breakdown
|
||||
if record.api_key not in breakdown.api_keys:
|
||||
breakdown.api_keys[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
), # Add any api_key-specific metadata here
|
||||
)
|
||||
breakdown.api_keys[record.api_key].metrics = update_metrics(breakdown.api_keys[record.api_key].metrics, record)
|
||||
if not is_ptu_sentinel:
|
||||
# Update api key breakdown
|
||||
if record.api_key not in breakdown.api_keys:
|
||||
breakdown.api_keys[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
), # Add any api_key-specific metadata here
|
||||
)
|
||||
breakdown.api_keys[record.api_key].metrics = update_metrics(breakdown.api_keys[record.api_key].metrics, record)
|
||||
|
||||
# Update entity-specific metrics if entity_id_field is provided
|
||||
if entity_id_field:
|
||||
|
|
@ -358,19 +399,20 @@ def update_breakdown_metrics(
|
|||
)
|
||||
breakdown.entities[entity_value].metrics = update_metrics(breakdown.entities[entity_value].metrics, record)
|
||||
|
||||
# Update API key breakdown for this entity
|
||||
if record.api_key not in breakdown.entities[entity_value].api_key_breakdown:
|
||||
breakdown.entities[entity_value].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
),
|
||||
if not is_ptu_sentinel:
|
||||
# Update API key breakdown for this entity
|
||||
if record.api_key not in breakdown.entities[entity_value].api_key_breakdown:
|
||||
breakdown.entities[entity_value].api_key_breakdown[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=SpendMetrics(),
|
||||
metadata=KeyMetadata(
|
||||
key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None),
|
||||
team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None),
|
||||
),
|
||||
)
|
||||
breakdown.entities[entity_value].api_key_breakdown[record.api_key].metrics = update_metrics(
|
||||
breakdown.entities[entity_value].api_key_breakdown[record.api_key].metrics,
|
||||
record,
|
||||
)
|
||||
breakdown.entities[entity_value].api_key_breakdown[record.api_key].metrics = update_metrics(
|
||||
breakdown.entities[entity_value].api_key_breakdown[record.api_key].metrics,
|
||||
record,
|
||||
)
|
||||
|
||||
return breakdown
|
||||
|
||||
|
|
@ -599,6 +641,14 @@ def _build_aggregated_sql_query(
|
|||
# total_successful_requests metadata they feed) once the admin UI reads SGR
|
||||
# only from LiteLLM_DailyGatewayRequests. The remaining spend, token and
|
||||
# api_requests rollups are still served from here.
|
||||
#
|
||||
# Only LiteLLM_DailyTeamSpend carries ptu_flat_cost; other daily tables emit a
|
||||
# constant zero so the SpendMetrics.flat_cost response shape stays uniform.
|
||||
ptu_flat_cost_select: Final = (
|
||||
"SUM(ptu_flat_cost)::float AS ptu_flat_cost"
|
||||
if table_name == "litellm_dailyteamspend"
|
||||
else "0::float AS ptu_flat_cost"
|
||||
)
|
||||
sql_query: Final = f"""
|
||||
SELECT
|
||||
date,
|
||||
|
|
@ -612,6 +662,7 @@ def _build_aggregated_sql_query(
|
|||
custom_llm_provider, mcp_namespaced_tool_name,
|
||||
endpoint) AS group_level,
|
||||
SUM(spend)::float AS spend,
|
||||
{ptu_flat_cost_select},
|
||||
SUM(prompt_tokens)::bigint AS prompt_tokens,
|
||||
SUM(completion_tokens)::bigint AS completion_tokens,
|
||||
SUM(cache_read_input_tokens)::bigint AS cache_read_input_tokens,
|
||||
|
|
@ -707,7 +758,9 @@ async def _aggregate_spend_records(
|
|||
The per-row loop is offloaded to a worker thread via asyncio.to_thread so
|
||||
a large result set doesn't peg the event loop.
|
||||
"""
|
||||
api_keys: Final[set[str]] = {record.api_key for record in records if record.api_key}
|
||||
api_keys: Final[set[str]] = {
|
||||
record.api_key for record in records if record.api_key and record.api_key != PTU_SENTINEL_API_KEY
|
||||
}
|
||||
|
||||
api_key_metadata: dict[str, _KeyMetadataDict] = {}
|
||||
if api_keys:
|
||||
|
|
@ -754,6 +807,7 @@ def _record_to_spend_metrics(record: _GroupingSetsRow) -> SpendMetrics:
|
|||
completion_tokens: Final = record.completion_tokens or 0
|
||||
return SpendMetrics(
|
||||
spend=record.spend or 0.0,
|
||||
flat_cost=_reported_flat_cost(record),
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=prompt_tokens + completion_tokens,
|
||||
|
|
@ -820,6 +874,7 @@ def _aggregate_grouping_sets_records_sync(
|
|||
for record in records:
|
||||
level = record.group_level
|
||||
metrics = _record_to_spend_metrics(record)
|
||||
is_ptu_sentinel = record.api_key == PTU_SENTINEL_API_KEY
|
||||
|
||||
if level == _GROUP_GRAND_TOTAL:
|
||||
total_metrics = metrics
|
||||
|
|
@ -832,7 +887,7 @@ def _aggregate_grouping_sets_records_sync(
|
|||
breakdown = ensure_date(record.date)["breakdown"]
|
||||
|
||||
if level == _GROUP_DATE_API_KEY:
|
||||
if record.api_key:
|
||||
if record.api_key and not is_ptu_sentinel:
|
||||
breakdown.api_keys[record.api_key] = KeyMetricWithMetadata(
|
||||
metrics=metrics,
|
||||
metadata=_key_metadata(api_key_metadata, record.api_key),
|
||||
|
|
@ -841,13 +896,13 @@ def _aggregate_grouping_sets_records_sync(
|
|||
if record.model:
|
||||
assign_metric_with_metadata(breakdown.models, record.model, metrics)
|
||||
elif level == _GROUP_DATE_MODEL_API_KEY:
|
||||
if record.model and record.api_key:
|
||||
if record.model and record.api_key and not is_ptu_sentinel:
|
||||
assign_api_key_breakdown(breakdown.models, record.model, record.api_key, metrics)
|
||||
elif level == _GROUP_DATE_MODEL_GROUP:
|
||||
if record.model_group:
|
||||
assign_metric_with_metadata(breakdown.model_groups, record.model_group, metrics)
|
||||
elif level == _GROUP_DATE_MODEL_GROUP_API_KEY:
|
||||
if record.model_group and record.api_key:
|
||||
if record.model_group and record.api_key and not is_ptu_sentinel:
|
||||
assign_api_key_breakdown(
|
||||
breakdown.model_groups,
|
||||
record.model_group,
|
||||
|
|
@ -855,10 +910,17 @@ def _aggregate_grouping_sets_records_sync(
|
|||
metrics,
|
||||
)
|
||||
elif level == _GROUP_DATE_PROVIDER:
|
||||
# Only PTU sentinel rows carry ptu_flat_cost and they have no provider, so at
|
||||
# this level the sentinel's cost would land under "unknown". Withholding the
|
||||
# flat cost matches the per-row path, which skips sentinel rows outright. The
|
||||
# bucket itself is still assigned unconditionally: a legacy row predating the
|
||||
# api_requests column backfills to all zeroes, and skipping those would drop a
|
||||
# provider the base build reported.
|
||||
provider_metrics = metrics.model_copy(update={"flat_cost": 0.0}) # mutable-ok: pydantic update payload
|
||||
provider = record.custom_llm_provider or "unknown"
|
||||
assign_metric_with_metadata(breakdown.providers, provider, metrics)
|
||||
assign_metric_with_metadata(breakdown.providers, provider, provider_metrics)
|
||||
elif level == _GROUP_DATE_PROVIDER_API_KEY:
|
||||
if record.api_key:
|
||||
if record.api_key and not is_ptu_sentinel:
|
||||
provider = record.custom_llm_provider or "unknown"
|
||||
assign_api_key_breakdown(breakdown.providers, provider, record.api_key, metrics)
|
||||
elif level == _GROUP_DATE_MCP:
|
||||
|
|
@ -898,7 +960,7 @@ async def _aggregate_grouping_sets_records(
|
|||
records: Sequence[_GroupingSetsRow],
|
||||
) -> _AggregatedSpendData:
|
||||
"""Async wrapper: fetch api_key_metadata, then dispatch on a worker thread."""
|
||||
api_keys: Final[set[str]] = {r.api_key for r in records if r.api_key}
|
||||
api_keys: Final[set[str]] = {r.api_key for r in records if r.api_key and r.api_key != PTU_SENTINEL_API_KEY}
|
||||
|
||||
api_key_metadata: dict[str, _KeyMetadataDict] = {}
|
||||
if api_keys:
|
||||
|
|
@ -1008,6 +1070,7 @@ async def get_daily_activity(
|
|||
results=aggregated["results"],
|
||||
metadata=DailySpendMetadata(
|
||||
total_spend=metadata_metrics.spend,
|
||||
total_flat_cost=metadata_metrics.flat_cost,
|
||||
total_prompt_tokens=metadata_metrics.prompt_tokens,
|
||||
total_completion_tokens=metadata_metrics.completion_tokens,
|
||||
total_tokens=metadata_metrics.total_tokens,
|
||||
|
|
@ -1098,6 +1161,7 @@ async def get_daily_activity_aggregated(
|
|||
results=aggregated["results"],
|
||||
metadata=DailySpendMetadata(
|
||||
total_spend=aggregated["totals"].spend,
|
||||
total_flat_cost=aggregated["totals"].flat_cost,
|
||||
total_prompt_tokens=aggregated["totals"].prompt_tokens,
|
||||
total_completion_tokens=aggregated["totals"].completion_tokens,
|
||||
total_tokens=aggregated["totals"].total_tokens,
|
||||
|
|
|
|||
|
|
@ -22,6 +22,35 @@ def validate_finite_spend(spend: float | None) -> None:
|
|||
)
|
||||
|
||||
|
||||
def validate_budget_duration(budget_duration: str | None) -> None:
|
||||
"""Reject budget durations that can't be parsed, are non-positive, or
|
||||
overflow date math, so a bad value can't be persisted and later crash the
|
||||
budget reset job.
|
||||
|
||||
A non-positive duration also resolves to a reset time of "now", which leaves
|
||||
the row permanently due: the reset job re-reads it every tick and, once
|
||||
enough of them exist, they fill each batch and starve every other tenant's
|
||||
reset.
|
||||
"""
|
||||
if budget_duration is None:
|
||||
return
|
||||
|
||||
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
|
||||
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
|
||||
|
||||
try:
|
||||
if duration_in_seconds(budget_duration) <= 0:
|
||||
raise ValueError("budget_duration must be positive")
|
||||
get_budget_reset_time(budget_duration=budget_duration)
|
||||
except (ValueError, OverflowError):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"Invalid budget_duration '{budget_duration}'. Use a format like '1h', '24h', '7d', or '30d'."
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching import DualCache
|
||||
from litellm.proxy._types import (
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ from litellm.litellm_core_utils.duration_parser import duration_in_seconds
|
|||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity
|
||||
from litellm.proxy.management_endpoints.common_utils import validate_budget_duration
|
||||
from litellm.proxy.management_helpers.object_permission_utils import (
|
||||
_set_object_permission,
|
||||
handle_update_object_permission_common,
|
||||
|
|
@ -184,6 +185,7 @@ def new_budget_request(data: NewCustomerRequest) -> BudgetNewRequest | None:
|
|||
|
||||
if budget_kv_pairs:
|
||||
budget_request: Final = BudgetNewRequest(**budget_kv_pairs)
|
||||
validate_budget_duration(budget_request.budget_duration)
|
||||
if budget_request.budget_reset_at is None and budget_request.budget_duration is not None:
|
||||
budget_request.budget_reset_at = datetime.utcnow() + timedelta(
|
||||
seconds=duration_in_seconds(duration=budget_request.budget_duration)
|
||||
|
|
|
|||
|
|
@ -42,6 +42,7 @@ from litellm.proxy.management_endpoints.common_utils import (
|
|||
_is_user_team_admin,
|
||||
_user_has_admin_view,
|
||||
require_caller_user_id_for_non_admin,
|
||||
validate_budget_duration,
|
||||
validate_finite_spend,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
|
|
@ -506,6 +507,8 @@ async def new_user(
|
|||
status_code=500,
|
||||
detail=CommonProxyErrors.db_not_connected_error.value,
|
||||
)
|
||||
validate_budget_duration(data.budget_duration)
|
||||
|
||||
# Check for duplicate user_id or email
|
||||
await _check_duplicate_user_id(data.user_id, prisma_client)
|
||||
await _check_duplicate_user_email(data.user_email, prisma_client)
|
||||
|
|
@ -1185,6 +1188,7 @@ def _update_internal_user_params(data_json: dict, data: UpdateUserRequest | Upda
|
|||
if "budget_duration" in non_default_values:
|
||||
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
|
||||
|
||||
validate_budget_duration(non_default_values["budget_duration"])
|
||||
non_default_values["budget_reset_at"] = get_budget_reset_time(
|
||||
budget_duration=non_default_values["budget_duration"]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -55,7 +55,10 @@ from litellm.proxy.auth.auth_checks import (
|
|||
get_project_object,
|
||||
get_team_object,
|
||||
)
|
||||
from litellm.proxy.auth.auth_utils import abbreviate_api_key
|
||||
from litellm.proxy.auth.auth_utils import (
|
||||
abbreviate_api_key,
|
||||
enforce_output_token_estimates_are_admin_only,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
decrypt_callback_vars,
|
||||
|
|
@ -79,6 +82,7 @@ from litellm.proxy.management_endpoints.common_utils import (
|
|||
_set_object_metadata_field,
|
||||
_team_member_has_permission,
|
||||
_user_has_admin_view,
|
||||
validate_budget_duration,
|
||||
validate_finite_spend,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
|
|
@ -841,12 +845,21 @@ async def _common_key_generation_helper(
|
|||
premium_user=premium_user,
|
||||
)
|
||||
|
||||
validate_budget_duration(data.budget_duration)
|
||||
|
||||
if data.throttle_on_budget_exceeded is True and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={"error": "Only proxy admins can enable throttle_on_budget_exceeded on a key."},
|
||||
)
|
||||
|
||||
enforce_output_token_estimates_are_admin_only(
|
||||
data=data,
|
||||
existing_metadata=None,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
entity="key",
|
||||
)
|
||||
|
||||
if data.metadata is not None and data.metadata.get("service_account_id") is not None and data.team_id is None:
|
||||
await validate_team_id_used_in_service_account_request(
|
||||
team_id=data.team_id,
|
||||
|
|
@ -1014,7 +1027,7 @@ async def _common_key_generation_helper(
|
|||
|
||||
# Only set budget_duration on key when explicitly provided. Keys with budget_id
|
||||
# but no explicit budget_duration follow their linked budget tier's schedule;
|
||||
# reset_budget_for_keys_linked_to_budgets() resets them when the tier resets.
|
||||
# reset_budget_for_litellm_budget_table() resets them when the tier resets.
|
||||
# This avoids duplicating budget_duration on keys so tier updates apply automatically.
|
||||
if "budget_duration" in data_json:
|
||||
data_json["key_budget_duration"] = data_json.pop("budget_duration", None)
|
||||
|
|
@ -1584,6 +1597,8 @@ async def generate_key_fn(
|
|||
- budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}.
|
||||
- model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit.
|
||||
- model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit.
|
||||
- default_estimated_output_tokens: Optional[int] - Proxy admin only. Expected output tokens reserved for TPM limiting when a request omits max_tokens. Positive integer. Falls back to the team setting, then to the built-in estimate.
|
||||
- default_estimated_output_tokens_per_model: Optional[dict] - Proxy admin only. Per-model override of the above. Example - {"gpt-4": 4096, "gpt-3.5-turbo": 1024}. Takes precedence over the key-wide value.
|
||||
- mcp_rpm_limit: Optional[dict] - key-specific per-MCP-server rpm limit, keyed by MCP server name (alias if set, else the configured name). Example - {"github": 100, "slack": 200}. IF null or {} then no MCP-specific rpm limit.
|
||||
- tag_rpm_limit: Optional[dict] - key-specific per-request-tag rpm limit, keyed by request tag. Example - {"cell-1": 1000, "cell-2": 500}. Each tag gets an independent counter; requests whose tag is absent fall back to the key-level rpm limit.
|
||||
- tpm_limit_type: Optional[str] - Type of tpm limit. Options: "best_effort_throughput" (no error if we're overallocating tpm), "guaranteed_throughput" (raise an error if we're overallocating tpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput".
|
||||
|
|
@ -1793,6 +1808,8 @@ async def generate_service_account_key_fn(
|
|||
- budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}.
|
||||
- model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit.
|
||||
- model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit.
|
||||
- default_estimated_output_tokens: Optional[int] - Proxy admin only. Expected output tokens reserved for TPM limiting when a request omits max_tokens. Positive integer. Falls back to the team setting, then to the built-in estimate.
|
||||
- default_estimated_output_tokens_per_model: Optional[dict] - Proxy admin only. Per-model override of the above. Example - {"gpt-4": 4096, "gpt-3.5-turbo": 1024}. Takes precedence over the key-wide value.
|
||||
- mcp_rpm_limit: Optional[dict] - key-specific per-MCP-server rpm limit, keyed by MCP server name (alias if set, else the configured name). Example - {"github": 100, "slack": 200}. IF null or {} then no MCP-specific rpm limit.
|
||||
- tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
|
||||
- rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
|
||||
|
|
@ -2387,6 +2404,7 @@ async def _validate_update_key_data(
|
|||
"""Validate permissions and constraints for key update."""
|
||||
# Reject NaN/±inf spend before it can reach the DB / spend counter.
|
||||
validate_finite_spend(data.spend)
|
||||
validate_budget_duration(data.budget_duration)
|
||||
|
||||
_is_proxy_admin: Final = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
|
||||
|
||||
|
|
@ -2473,6 +2491,13 @@ async def _validate_update_key_data(
|
|||
detail={"error": "Only proxy admins can enable throttle_on_budget_exceeded on a key."},
|
||||
)
|
||||
|
||||
enforce_output_token_estimates_are_admin_only(
|
||||
data=data,
|
||||
existing_metadata=_existing_metadata if isinstance(_existing_metadata, dict) else None,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
entity="key",
|
||||
)
|
||||
|
||||
# Personal-key bypass: the caller both created the key AND still owns it
|
||||
# (user_id == caller). Checking only created_by would let a demoted admin
|
||||
# who originally created a key for another user continue editing it without
|
||||
|
|
@ -2655,6 +2680,8 @@ async def update_key_fn(
|
|||
- mcp_rpm_limit: Optional[dict] - Per-MCP-server RPM limits, keyed by MCP server name {"github": 100, "slack": 200}
|
||||
- tag_rpm_limit: Optional[dict] - Per-request-tag RPM limits, keyed by request tag {"cell-1": 1000, "cell-2": 500}. Each tag gets an independent counter; absent tags fall back to the key-level rpm limit.
|
||||
- model_tpm_limit: Optional[dict] - Model-specific TPM limits {"gpt-4": 100000, "claude-v1": 200000}
|
||||
- default_estimated_output_tokens: Optional[int] - Proxy admin only. Expected output tokens reserved for TPM limiting when a request omits max_tokens. Positive integer.
|
||||
- default_estimated_output_tokens_per_model: Optional[dict] - Proxy admin only. Per-model override of the above {"gpt-4": 4096, "gpt-3.5-turbo": 1024}
|
||||
- tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
|
||||
- rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
|
||||
- allowed_cache_controls: Optional[list] - List of allowed cache control values
|
||||
|
|
@ -4629,6 +4656,15 @@ async def _execute_virtual_key_regeneration(
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
if data is not None:
|
||||
_existing_key_metadata: Final = getattr(key_in_db, "metadata", None)
|
||||
enforce_output_token_estimates_are_admin_only(
|
||||
data=data,
|
||||
existing_metadata=_existing_key_metadata if isinstance(_existing_key_metadata, dict) else None,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
entity="key",
|
||||
)
|
||||
|
||||
new_token: Final = await get_new_token(data=data)
|
||||
new_token_hash: Final = hash_token(new_token)
|
||||
new_token_key_name: Final = abbreviate_api_key(api_key=new_token)
|
||||
|
|
|
|||
|
|
@ -4,7 +4,8 @@ from typing import Final
|
|||
from urllib.parse import urlencode
|
||||
|
||||
from fastapi import Request
|
||||
from fastapi.dependencies.utils import get_flat_dependant
|
||||
from fastapi.dependencies.utils import get_flat_params
|
||||
from fastapi.params import ParamTypes
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from litellm.types.proxy.management_endpoints.management_v1 import (
|
||||
|
|
@ -42,7 +43,13 @@ def _declared_query_params(request: Request) -> frozenset[str]:
|
|||
dependant: Final = getattr(route, "dependant", None)
|
||||
if dependant is None:
|
||||
return frozenset()
|
||||
return frozenset(field.alias for field in get_flat_dependant(dependant, skip_repeats=True).query_params)
|
||||
# fastapi>=0.140.7 removed get_flat_dependant(); get_flat_params() returns the
|
||||
# flattened (deduped) param list. Filter to query params to match the old behavior.
|
||||
return frozenset(
|
||||
field.alias
|
||||
for field in get_flat_params(dependant)
|
||||
if getattr(field.field_info, "in_", None) == ParamTypes.query
|
||||
)
|
||||
|
||||
|
||||
def escape_like(value: str) -> str:
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ import datetime
|
|||
import json
|
||||
from collections.abc import Awaitable, Mapping, Sequence
|
||||
from json import JSONDecodeError
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, Protocol, cast
|
||||
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
|
||||
|
|
@ -55,6 +56,10 @@ from litellm.proxy.management_endpoints.team_endpoints import (
|
|||
update_team as _legacy_update_team,
|
||||
)
|
||||
from litellm.proxy.management_helpers.audit_logs import create_object_audit_log
|
||||
from litellm.proxy.spend_tracking.ptu_feature_flag import (
|
||||
PTU_COST_ATTRIBUTION_ENV_VAR,
|
||||
is_ptu_cost_attribution_enabled,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
from litellm.repositories.model_repository import ModelRepository
|
||||
from litellm.repositories.table_repositories import ModelTableRepository
|
||||
|
|
@ -79,6 +84,7 @@ from litellm.types.router import (
|
|||
SPECIAL_MODEL_INFO_PARAMS,
|
||||
Deployment,
|
||||
GenericLiteLLMParams,
|
||||
ModelInfo,
|
||||
updateDeployment,
|
||||
)
|
||||
from litellm.utils import get_utc_datetime
|
||||
|
|
@ -233,7 +239,129 @@ def _raise_on_strategy_router_write_violation(
|
|||
)
|
||||
|
||||
|
||||
_PTU_MODEL_INFO_FIELDS: Final = ("ptu_count", "cost_per_ptu_per_hour", "ptu_effective_from", "ptu_effective_to")
|
||||
|
||||
|
||||
def _explicitly_cleared_ptu_fields(model_info: ModelInfo | None) -> frozenset[str]:
|
||||
"""The PTU fields a patch sends as an explicit null, which update_db_model drops.
|
||||
|
||||
Empty while the feature is off, so disabling pauses PTU rather than letting a client
|
||||
that round-trips a model_info blob erase a configuration set up during an earlier opt-in.
|
||||
"""
|
||||
if model_info is None or not is_ptu_cost_attribution_enabled():
|
||||
return frozenset()
|
||||
return frozenset(
|
||||
field
|
||||
for field in _PTU_MODEL_INFO_FIELDS
|
||||
if field in model_info.model_fields_set and getattr(model_info, field) is None
|
||||
)
|
||||
|
||||
|
||||
def _merged_ptu_model_info(*, db_model: Deployment, patch_data: updateDeployment) -> Mapping[str, object]:
|
||||
"""The model_info a patch would store, which is the stored blob updated by the patch.
|
||||
|
||||
A PTU invariant holds over the deployment as it will exist, not over whichever subset
|
||||
of fields a caller happened to send.
|
||||
"""
|
||||
empty: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
stored: Final = db_model.model_info.model_dump(exclude_none=True) if db_model.model_info else empty
|
||||
incoming: Final = patch_data.model_info.model_dump(exclude_none=True) if patch_data.model_info else empty
|
||||
cleared: Final = _explicitly_cleared_ptu_fields(patch_data.model_info)
|
||||
return MappingProxyType({k: v for k, v in {**stored, **incoming}.items() if k not in cleared})
|
||||
|
||||
|
||||
def _raise_if_ptu_cost_attribution_disabled(incoming_model_info: Mapping[str, object]) -> None:
|
||||
"""Reject PTU model_info fields unless the operator opted into PTU cost attribution.
|
||||
|
||||
Takes the incoming request's model_info rather than the merged deployment, so an
|
||||
unrelated patch of a model that still stores PTU config from an earlier opt-in is
|
||||
left alone. The fields are rejected rather than dropped so a caller never believes
|
||||
a flat cost was configured while the rollup that would price it is not running.
|
||||
|
||||
Only a value is rejected. An explicit null reaches the clear loop, which is gated on
|
||||
the same flag, so a disabled proxy neither writes PTU config nor erases what an
|
||||
earlier opt-in stored. Disabling pauses the feature rather than discarding its setup.
|
||||
"""
|
||||
if is_ptu_cost_attribution_enabled():
|
||||
return
|
||||
supplied: Final = tuple(field for field in _PTU_MODEL_INFO_FIELDS if incoming_model_info.get(field) is not None)
|
||||
if not supplied:
|
||||
return
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
f"PTU cost attribution is disabled, so {', '.join(supplied)} cannot be set. "
|
||||
f"Set {PTU_COST_ATTRIBUTION_ENV_VAR}=true to enable it."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _validate_ptu_model_info(model_info: Mapping[str, object]) -> None:
|
||||
"""Enforce the PTU cross-field invariant on the effective model_info.
|
||||
|
||||
ptu_count and cost_per_ptu_per_hour must be set together, and a team_id and a
|
||||
ptu_effective_from are required when they are. The start is mandatory rather than
|
||||
defaulted because flat cost accrues from it: inferring one would let a deployment
|
||||
configured today be billed for days it did not exist. Per-field bounds (positive
|
||||
count, non-negative rate) are enforced by ModelInfo itself.
|
||||
|
||||
Window ordering is checked before the count/rate gate. A patch that touches only one
|
||||
end of the window carries no count or rate, and ModelInfo sees one field at a time, so
|
||||
leaving it to either would let an inverted window reach the row; the next load then
|
||||
fails to parse it and drops the deployment out of the router, where no further patch
|
||||
can repair it because each one re-parses the stored value first.
|
||||
"""
|
||||
effective_from: Final = _coerce_ptu_datetime(model_info.get("ptu_effective_from"))
|
||||
effective_to: Final = _coerce_ptu_datetime(model_info.get("ptu_effective_to"))
|
||||
if effective_from is not None and effective_to is not None and effective_to <= effective_from:
|
||||
raise HTTPException(status_code=400, detail="ptu_effective_to must be after ptu_effective_from")
|
||||
|
||||
has_count: Final = model_info.get("ptu_count") is not None
|
||||
has_rate: Final = model_info.get("cost_per_ptu_per_hour") is not None
|
||||
if not has_count and not has_rate:
|
||||
return
|
||||
if has_count != has_rate:
|
||||
raise HTTPException(status_code=400, detail="ptu_count and cost_per_ptu_per_hour must be set together")
|
||||
if effective_from is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
"ptu_effective_from is required when PTU fields are set. Flat cost accrues from that "
|
||||
"instant, so without it the start would have to be inferred and a deployment configured "
|
||||
"today could be billed for days it did not exist"
|
||||
),
|
||||
)
|
||||
if not model_info.get("team_id"):
|
||||
raise HTTPException(
|
||||
status_code=400, detail="team_id is required when PTU fields are set (one model maps to one team)"
|
||||
)
|
||||
|
||||
|
||||
def _parse_ptu_datetime(value: object) -> datetime.datetime | None:
|
||||
"""``value`` as a datetime, parsing an ISO string, else None."""
|
||||
if isinstance(value, datetime.datetime):
|
||||
return value
|
||||
if not isinstance(value, str):
|
||||
return None
|
||||
try:
|
||||
return datetime.datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
|
||||
def _coerce_ptu_datetime(value: object) -> datetime.datetime | None:
|
||||
"""Coerce a model_info effective-window value (datetime or ISO string) to UTC, else None."""
|
||||
parsed: Final = _parse_ptu_datetime(value)
|
||||
if parsed is None:
|
||||
return None
|
||||
if parsed.tzinfo is None:
|
||||
return parsed.replace(tzinfo=datetime.timezone.utc)
|
||||
return parsed.astimezone(datetime.timezone.utc)
|
||||
|
||||
|
||||
def update_db_model(db_model: Deployment, updated_patch: updateDeployment) -> PrismaCompatibleUpdateDBModel:
|
||||
if updated_patch.model_info is not None:
|
||||
_raise_if_ptu_cost_attribution_disabled(updated_patch.model_info.model_dump(exclude_none=True))
|
||||
merged_model_name: Final = updated_patch.model_name or db_model.model_name
|
||||
merged_litellm_params: Final = db_model.litellm_params.model_dump(exclude_none=True)
|
||||
merged_model_info: Final = db_model.model_info.model_dump(exclude_none=True)
|
||||
|
|
@ -270,6 +398,10 @@ def update_db_model(db_model: Deployment, updated_patch: updateDeployment) -> Pr
|
|||
if field in SPECIAL_MODEL_INFO_PARAMS and getattr(updated_patch.model_info, field) is None:
|
||||
merged_model_info.pop(field, None)
|
||||
merged_litellm_params.pop(field, None)
|
||||
for field in _explicitly_cleared_ptu_fields(updated_patch.model_info):
|
||||
merged_model_info.pop(field, None)
|
||||
|
||||
_validate_ptu_model_info(merged_model_info)
|
||||
|
||||
# convert to prisma compatible format
|
||||
|
||||
|
|
@ -716,6 +848,18 @@ async def _update_team_model_in_db(
|
|||
premium_user=premium_user,
|
||||
)
|
||||
|
||||
# Validated before any write, beside the premium check the create path already runs
|
||||
# here. The team ACL is updated below and autocommits, so a validator that raises
|
||||
# further down would leave the team mutated and the deployment row never written.
|
||||
#
|
||||
# The merged view is what gets stored, so that is what has to satisfy the invariants.
|
||||
# Validating the patch alone rejected a partial edit of an already valid deployment:
|
||||
# raising the rate on a configured model carries no ptu_effective_from, which the
|
||||
# stored row supplies.
|
||||
if patch_data.model_info is not None:
|
||||
_raise_if_ptu_cost_attribution_disabled(patch_data.model_info.model_dump(exclude_none=True))
|
||||
_validate_ptu_model_info(_merged_ptu_model_info(db_model=db_model, patch_data=patch_data))
|
||||
|
||||
patch_team_id: Final = patch_data.model_info.team_id if patch_data.model_info else None
|
||||
|
||||
# No team_id in patch, proceed with standard update
|
||||
|
|
@ -1424,6 +1568,10 @@ async def add_new_model(
|
|||
|
||||
model_response: LiteLLM_ProxyModelTable | None = None
|
||||
# update DB
|
||||
incoming_model_info: Final = model_params.model_info.model_dump(exclude_none=True)
|
||||
_raise_if_ptu_cost_attribution_disabled(incoming_model_info)
|
||||
_validate_ptu_model_info(incoming_model_info)
|
||||
|
||||
if store_model_in_db is True:
|
||||
"""
|
||||
- store model_list in db
|
||||
|
|
|
|||
|
|
@ -82,6 +82,7 @@ from litellm.proxy.auth.auth_checks import (
|
|||
get_team_object,
|
||||
get_user_object,
|
||||
)
|
||||
from litellm.proxy.auth.auth_utils import enforce_output_token_estimates_are_admin_only
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars
|
||||
from litellm.proxy.common_utils.json_merge_patch import apply_json_merge_patch
|
||||
|
|
@ -95,6 +96,7 @@ from litellm.proxy.management_endpoints.common_utils import (
|
|||
_update_metadata_fields,
|
||||
_upsert_budget_and_membership,
|
||||
_user_has_admin_view,
|
||||
validate_budget_duration,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.organization_endpoints import (
|
||||
add_member_to_organization,
|
||||
|
|
@ -1153,6 +1155,8 @@ async def new_team(
|
|||
- metadata: Optional[dict] - Metadata for team, store information for team. Example metadata = {"extra_info": "some info"}
|
||||
- model_rpm_limit: Optional[Dict[str, int]] - The RPM (Requests Per Minute) limit for this team - applied across all keys for this team.
|
||||
- model_tpm_limit: Optional[Dict[str, int]] - The TPM (Tokens Per Minute) limit for this team - applied across all keys for this team.
|
||||
- default_estimated_output_tokens: Optional[int] - Expected output tokens reserved for TPM limiting when a request omits max_tokens, for keys on this team that do not set their own. Positive integer.
|
||||
- default_estimated_output_tokens_per_model: Optional[Dict[str, int]] - Per-model override of the above. Example: {"gpt-4": 4096, "gpt-3.5-turbo": 1024}
|
||||
- mcp_rpm_limit: Optional[Dict[str, int]] - Per-MCP-server RPM limit for this team, keyed by MCP server name (alias if set, else the configured name). Example: {"github": 100, "slack": 200}. Applied across all keys for this team.
|
||||
- tpm_limit: Optional[int] - The TPM (Tokens Per Minute) limit for this team - all keys with this team_id will have at max this TPM limit
|
||||
- rpm_limit: Optional[int] - The RPM (Requests Per Minute) limit for this team - all keys associated with this team_id will have at max this RPM limit
|
||||
|
|
@ -1255,6 +1259,9 @@ async def new_team(
|
|||
detail={"error": f"soft_budget must be a non-negative finite number. Received: {data.soft_budget}"},
|
||||
)
|
||||
|
||||
validate_budget_duration(data.budget_duration)
|
||||
validate_budget_duration(data.team_member_budget_duration)
|
||||
|
||||
if data.soft_budget is not None:
|
||||
if data.max_budget is not None:
|
||||
# If max_budget is set, soft_budget must be strictly lower than max_budget
|
||||
|
|
@ -1266,6 +1273,13 @@ async def new_team(
|
|||
},
|
||||
)
|
||||
|
||||
enforce_output_token_estimates_are_admin_only(
|
||||
data=data,
|
||||
existing_metadata=None,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
entity="team",
|
||||
)
|
||||
|
||||
# Check if license is over limit
|
||||
total_teams: Final = await _team_db(prisma_client).count()
|
||||
if total_teams and _license_check.is_team_count_over_limit(team_count=total_teams):
|
||||
|
|
@ -1863,6 +1877,8 @@ async def update_team(
|
|||
- allowed_passthrough_routes: Optional[List[str]] - List of allowed pass through routes for the team.
|
||||
- model_rpm_limit: Optional[Dict[str, int]] - The RPM (Requests Per Minute) limit per model for this team. Example: {"gpt-4": 100, "gpt-3.5-turbo": 200}
|
||||
- model_tpm_limit: Optional[Dict[str, int]] - The TPM (Tokens Per Minute) limit per model for this team. Example: {"gpt-4": 10000, "gpt-3.5-turbo": 20000}
|
||||
- default_estimated_output_tokens: Optional[int] - Expected output tokens reserved for TPM limiting when a request omits max_tokens, for keys on this team that do not set their own. Positive integer.
|
||||
- default_estimated_output_tokens_per_model: Optional[Dict[str, int]] - Per-model override of the above. Example: {"gpt-4": 4096, "gpt-3.5-turbo": 1024}
|
||||
- mcp_rpm_limit: Optional[Dict[str, int]] - Per-MCP-server RPM limit for this team, keyed by MCP server name (alias if set, else the configured name). Example: {"github": 100, "slack": 200}. Applied across all keys for this team.
|
||||
Example - update team TPM Limit
|
||||
- allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint.
|
||||
|
|
@ -1935,6 +1951,9 @@ async def update_team(
|
|||
detail={"error": f"soft_budget must be a non-negative finite number. Received: {data.soft_budget}"},
|
||||
)
|
||||
|
||||
validate_budget_duration(data.budget_duration)
|
||||
validate_budget_duration(data.team_member_budget_duration)
|
||||
|
||||
existing_team_row = await _team_db(prisma_client).find_unique(where={"team_id": data.team_id})
|
||||
|
||||
if existing_team_row is None:
|
||||
|
|
@ -1949,6 +1968,14 @@ async def update_team(
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
_existing_team_metadata: Final[object] = getattr(existing_team_row, "metadata", None)
|
||||
enforce_output_token_estimates_are_admin_only(
|
||||
data=data,
|
||||
existing_metadata=_existing_team_metadata if isinstance(_existing_team_metadata, dict) else None,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
entity="team",
|
||||
)
|
||||
|
||||
_check_passthrough_routes_caller_permission(data, user_api_key_dict, entity="team")
|
||||
|
||||
if data.soft_budget is not None:
|
||||
|
|
@ -2959,7 +2986,7 @@ async def team_member_add(
|
|||
except HTTPException as e:
|
||||
raise e
|
||||
|
||||
_validate_budget_duration(data.budget_duration)
|
||||
validate_budget_duration(data.budget_duration)
|
||||
|
||||
prisma_client = cast(PrismaClient, prisma_client)
|
||||
|
||||
|
|
@ -3262,29 +3289,6 @@ def _build_member_budget_patch(data: TeamMemberUpdateRequest) -> dict[str, objec
|
|||
}
|
||||
|
||||
|
||||
def _validate_budget_duration(budget_duration: str | None) -> None:
|
||||
"""Reject budget durations that can't be parsed, are non-positive, or
|
||||
overflow date math, so a bad value can't be persisted and later crash the
|
||||
budget reset job."""
|
||||
if budget_duration is None:
|
||||
return
|
||||
|
||||
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
|
||||
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
|
||||
|
||||
try:
|
||||
if duration_in_seconds(budget_duration) <= 0:
|
||||
raise ValueError("budget_duration must be positive")
|
||||
get_budget_reset_time(budget_duration=budget_duration)
|
||||
except (ValueError, OverflowError):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": f"Invalid budget_duration '{budget_duration}'. Use a format like '1h', '24h', '7d', or '30d'."
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/team/member_update",
|
||||
tags=["team management"],
|
||||
|
|
@ -3322,7 +3326,7 @@ async def team_member_update(
|
|||
detail={"error": "Either user_id or user_email needs to be passed in"},
|
||||
)
|
||||
|
||||
_validate_budget_duration(data.budget_duration)
|
||||
validate_budget_duration(data.budget_duration)
|
||||
|
||||
_existing_team_row: Final = await _team_db(prisma_client).find_unique(where={"team_id": data.team_id})
|
||||
|
||||
|
|
|
|||
|
|
@ -18,13 +18,16 @@ Scoping:
|
|||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, Final
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Final, Protocol
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import (
|
||||
CommonProxyErrors,
|
||||
LiteLLM_TeamTable,
|
||||
LitellmUserRoles,
|
||||
UserAPIKeyAuth,
|
||||
user_api_key_has_admin_view,
|
||||
|
|
@ -40,10 +43,56 @@ from litellm.types.memory_management import (
|
|||
MemoryUpdateRequest,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
||||
def _serialize_metadata_for_prisma(metadata: Any) -> str:
|
||||
class _MemoryRecord(Protocol):
|
||||
memory_id: str
|
||||
key: str
|
||||
value: str
|
||||
metadata: object
|
||||
user_id: str | None
|
||||
team_id: str | None
|
||||
created_at: datetime | None
|
||||
created_by: str | None
|
||||
updated_at: datetime | None
|
||||
updated_by: str | None
|
||||
|
||||
|
||||
class _MemoryTableActions(Protocol):
|
||||
async def create(self, data: Mapping[str, object]) -> _MemoryRecord: ...
|
||||
|
||||
async def find_many(
|
||||
self,
|
||||
where: Mapping[str, object] | None = ...,
|
||||
order: Mapping[str, str] | None = ...,
|
||||
skip: int = ...,
|
||||
take: int = ...,
|
||||
) -> Sequence[_MemoryRecord]: ...
|
||||
|
||||
async def count(self, where: Mapping[str, object] | None = ...) -> int: ...
|
||||
|
||||
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> _MemoryRecord: ...
|
||||
|
||||
async def delete(self, where: Mapping[str, object]) -> _MemoryRecord | None: ...
|
||||
|
||||
|
||||
def _memory_table(prisma_client: "PrismaClient") -> _MemoryTableActions:
|
||||
return MemoryRepository(prisma_client).table
|
||||
|
||||
|
||||
class _TeamTableActions(Protocol):
|
||||
async def find_unique(self, where: Mapping[str, str]) -> LiteLLM_TeamTable | None: ...
|
||||
|
||||
|
||||
def _team_table(prisma_client: "PrismaClient") -> _TeamTableActions:
|
||||
return TeamRepository(prisma_client).table
|
||||
|
||||
|
||||
def _serialize_metadata_for_prisma(metadata: object) -> str:
|
||||
"""
|
||||
Encode a `metadata` payload for the `Json?` column.
|
||||
|
||||
|
|
@ -62,25 +111,25 @@ def _is_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
|
|||
return user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
|
||||
|
||||
def _visibility_filter(user_api_key_dict: UserAPIKeyAuth) -> dict | None:
|
||||
def _visibility_filter(user_api_key_dict: UserAPIKeyAuth) -> Mapping[str, object] | None:
|
||||
"""
|
||||
Prisma `where` fragment restricting rows to those the caller can see.
|
||||
Returns None for admins (no restriction).
|
||||
"""
|
||||
if user_api_key_has_admin_view(user_api_key_dict):
|
||||
return None
|
||||
ors: Final[list[dict]] = []
|
||||
if user_api_key_dict.user_id:
|
||||
ors.append({"user_id": user_api_key_dict.user_id})
|
||||
if user_api_key_dict.team_id:
|
||||
ors.append({"team_id": user_api_key_dict.team_id})
|
||||
ors: Final = [
|
||||
{field: value}
|
||||
for field, value in (("user_id", user_api_key_dict.user_id), ("team_id", user_api_key_dict.team_id))
|
||||
if value
|
||||
]
|
||||
if not ors:
|
||||
# Caller has neither user_id nor team_id — match nothing.
|
||||
return {"memory_id": "__no_match__"}
|
||||
return {"OR": ors}
|
||||
|
||||
|
||||
def _row_to_model(row: Any) -> LiteLLM_MemoryRow:
|
||||
def _row_to_model(row: _MemoryRecord) -> LiteLLM_MemoryRow:
|
||||
return LiteLLM_MemoryRow(
|
||||
memory_id=row.memory_id,
|
||||
key=row.key,
|
||||
|
|
@ -95,7 +144,7 @@ def _row_to_model(row: Any) -> LiteLLM_MemoryRow:
|
|||
)
|
||||
|
||||
|
||||
def _require_prisma():
|
||||
def _require_prisma() -> "PrismaClient":
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
|
|
@ -113,7 +162,9 @@ def _internal_error(log_message: str, exc: Exception, default_detail: str) -> HT
|
|||
return HTTPException(status_code=500, detail=default_detail)
|
||||
|
||||
|
||||
async def _assert_write_access(prisma_client: Any, row: Any, user_api_key_dict: UserAPIKeyAuth) -> None:
|
||||
async def _assert_write_access(
|
||||
prisma_client: "PrismaClient", row: _MemoryRecord, user_api_key_dict: UserAPIKeyAuth
|
||||
) -> None:
|
||||
"""
|
||||
Enforce ownership for mutations (PUT/DELETE).
|
||||
|
||||
|
|
@ -153,7 +204,7 @@ async def _assert_write_access(prisma_client: Any, row: Any, user_api_key_dict:
|
|||
)
|
||||
|
||||
|
||||
async def _is_team_admin_for(prisma_client: Any, user_api_key_dict: UserAPIKeyAuth, team_id: str) -> bool:
|
||||
async def _is_team_admin_for(prisma_client: "PrismaClient", user_api_key_dict: UserAPIKeyAuth, team_id: str) -> bool:
|
||||
"""
|
||||
True if the caller is a team admin of `team_id`, or an org admin for the
|
||||
team's organization. Mirrors the auth pattern used by team-management
|
||||
|
|
@ -168,7 +219,7 @@ async def _is_team_admin_for(prisma_client: Any, user_api_key_dict: UserAPIKeyAu
|
|||
)
|
||||
|
||||
try:
|
||||
team_obj: Final = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
|
||||
team_obj: Final = await _team_table(prisma_client).find_unique(where={"team_id": team_id})
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error loading team for write-auth check (team_id=%s): %s", team_id, e)
|
||||
return False
|
||||
|
|
@ -269,7 +320,7 @@ async def create_memory(
|
|||
# `metadata` is a `Json?` column — prisma-client-python rejects raw
|
||||
# Python values, so JSON-encode any non-null payload and omit the field
|
||||
# entirely when None so the column defaults to SQL NULL.
|
||||
create_data: Final[dict] = {
|
||||
create_data: Final[dict[str, object]] = {
|
||||
"key": body.key,
|
||||
"value": body.value,
|
||||
"user_id": user_id,
|
||||
|
|
@ -281,7 +332,7 @@ async def create_memory(
|
|||
create_data["metadata"] = _serialize_metadata_for_prisma(body.metadata)
|
||||
|
||||
try:
|
||||
row: Final = await MemoryRepository(prisma_client).table.create(data=create_data)
|
||||
row: Final = await _memory_table(prisma_client).create(data=create_data)
|
||||
except Exception as e:
|
||||
# Key is globally unique. Any duplicate → 409.
|
||||
if _is_unique_violation(e):
|
||||
|
|
@ -325,14 +376,14 @@ async def list_memory(
|
|||
# top-level "AND" — safer than `dict.update` since future visibility
|
||||
# filters could grow an "OR" key that would clobber this one if merged
|
||||
# by key.
|
||||
key_filter: Final[dict] = {}
|
||||
key_filter: Final[dict[str, object]] = {}
|
||||
if key_prefix is not None:
|
||||
key_filter["key"] = {"startsWith": key_prefix}
|
||||
elif key is not None:
|
||||
key_filter["key"] = key
|
||||
|
||||
vis: Final = _visibility_filter(user_api_key_dict)
|
||||
where: dict
|
||||
where: Mapping[str, object]
|
||||
if vis is None:
|
||||
where = key_filter
|
||||
elif not key_filter:
|
||||
|
|
@ -341,8 +392,8 @@ async def list_memory(
|
|||
where = {"AND": [key_filter, vis]}
|
||||
|
||||
try:
|
||||
total: Final = await MemoryRepository(prisma_client).table.count(where=where)
|
||||
rows: Final = await MemoryRepository(prisma_client).table.find_many(
|
||||
total: Final = await _memory_table(prisma_client).count(where=where)
|
||||
rows: Final = await _memory_table(prisma_client).find_many(
|
||||
where=where,
|
||||
order={"updated_at": "desc"},
|
||||
skip=(page - 1) * page_size,
|
||||
|
|
@ -354,12 +405,14 @@ async def list_memory(
|
|||
return MemoryListResponse(memories=[_row_to_model(r) for r in rows], total=total)
|
||||
|
||||
|
||||
async def _find_memory_for_caller(prisma_client: Any, key: str, user_api_key_dict: UserAPIKeyAuth) -> Any:
|
||||
async def _find_memory_for_caller(
|
||||
prisma_client: "PrismaClient", key: str, user_api_key_dict: UserAPIKeyAuth
|
||||
) -> _MemoryRecord:
|
||||
"""Look up a memory row by key, scoped to the caller's visibility."""
|
||||
key_filter: Final[dict] = {"key": key}
|
||||
key_filter: Final[Mapping[str, object]] = {"key": key}
|
||||
vis: Final = _visibility_filter(user_api_key_dict)
|
||||
where: Final[dict] = key_filter if vis is None else {"AND": [key_filter, vis]}
|
||||
rows = await MemoryRepository(prisma_client).table.find_many(where=where, take=1, order={"updated_at": "desc"})
|
||||
where: Final[Mapping[str, object]] = key_filter if vis is None else {"AND": [key_filter, vis]}
|
||||
rows = await _memory_table(prisma_client).find_many(where=where, take=1, order={"updated_at": "desc"})
|
||||
if not rows:
|
||||
raise HTTPException(status_code=404, detail=f"Memory with key '{key}' not found")
|
||||
return rows[0]
|
||||
|
|
@ -415,7 +468,7 @@ async def upsert_memory(
|
|||
fields_sent: Final = body.model_fields_set
|
||||
metadata_in_payload: Final = "metadata" in fields_sent
|
||||
|
||||
data: Final[dict] = {}
|
||||
data: Final[dict[str, object]] = {}
|
||||
if body.value is not None:
|
||||
data["value"] = body.value
|
||||
if metadata_in_payload:
|
||||
|
|
@ -427,7 +480,7 @@ async def upsert_memory(
|
|||
)
|
||||
data["updated_by"] = user_api_key_dict.user_id
|
||||
|
||||
async def _find_existing() -> Any:
|
||||
async def _find_existing() -> _MemoryRecord | None:
|
||||
"""Return the caller-visible row for `key`, or None."""
|
||||
try:
|
||||
return await _find_memory_for_caller(prisma_client, key, user_api_key_dict)
|
||||
|
|
@ -444,7 +497,7 @@ async def upsert_memory(
|
|||
# their team) — otherwise a teammate could overwrite a personal
|
||||
# entry through the OR-based visibility filter.
|
||||
await _assert_write_access(prisma_client, existing, user_api_key_dict)
|
||||
row = await MemoryRepository(prisma_client).table.update(
|
||||
row = await _memory_table(prisma_client).update(
|
||||
where={"memory_id": existing.memory_id},
|
||||
data=data,
|
||||
)
|
||||
|
|
@ -459,7 +512,7 @@ async def upsert_memory(
|
|||
# Omit `metadata` when None so the column defaults to SQL NULL;
|
||||
# otherwise JSON-encode for Prisma — same pattern as
|
||||
# `create_memory` above.
|
||||
create_data: Final[dict] = {
|
||||
create_data: Final[dict[str, object]] = {
|
||||
"key": key,
|
||||
"value": body.value,
|
||||
"user_id": user_id,
|
||||
|
|
@ -470,7 +523,7 @@ async def upsert_memory(
|
|||
if body.metadata is not None:
|
||||
create_data["metadata"] = _serialize_metadata_for_prisma(body.metadata)
|
||||
try:
|
||||
row = await MemoryRepository(prisma_client).table.create(data=create_data)
|
||||
row = await _memory_table(prisma_client).create(data=create_data)
|
||||
except Exception as e:
|
||||
# Race: a concurrent PUT/POST created the row after our check.
|
||||
# Re-read and fall back to an update so the PUT stays idempotent
|
||||
|
|
@ -487,7 +540,7 @@ async def upsert_memory(
|
|||
)
|
||||
# Same write-authorization check as the non-race path.
|
||||
await _assert_write_access(prisma_client, existing_after_race, user_api_key_dict)
|
||||
row = await MemoryRepository(prisma_client).table.update(
|
||||
row = await _memory_table(prisma_client).update(
|
||||
where={"memory_id": existing_after_race.memory_id},
|
||||
data=data,
|
||||
)
|
||||
|
|
@ -515,7 +568,7 @@ async def delete_memory(
|
|||
# Visibility != write authority — see the upsert handler for the rationale.
|
||||
await _assert_write_access(prisma_client, row, user_api_key_dict)
|
||||
try:
|
||||
await MemoryRepository(prisma_client).table.delete(where={"memory_id": row.memory_id})
|
||||
await _memory_table(prisma_client).delete(where={"memory_id": row.memory_id})
|
||||
except Exception as e:
|
||||
raise _internal_error("Error deleting memory: %s", e, "Internal error deleting memory entry.")
|
||||
|
||||
|
|
|
|||
|
|
@ -68,6 +68,16 @@ class StorageBackendFileService:
|
|||
code=400,
|
||||
)
|
||||
|
||||
if target_model_names:
|
||||
managed_files_hook: Final = proxy_logging_obj.get_proxy_hook("managed_files")
|
||||
if not isinstance(managed_files_hook, BaseFileEndpoints):
|
||||
raise ProxyException(
|
||||
message="Uploading with target_model_names requires a database-connected proxy, and this proxy has no database configured",
|
||||
type="invalid_request_error",
|
||||
param="target_model_names",
|
||||
code=400,
|
||||
)
|
||||
|
||||
# Extract file information
|
||||
file_content: Final = file_data["content"]
|
||||
filename: Final = file_data.get("filename", "file")
|
||||
|
|
|
|||
|
|
@ -60,6 +60,7 @@ from .passthrough_endpoint_router import PassthroughEndpointRouter
|
|||
|
||||
vertex_llm_base: Final = VertexBase()
|
||||
router: Final = APIRouter()
|
||||
openai_passthrough_router: Final = APIRouter()
|
||||
default_vertex_config: Final = None
|
||||
|
||||
passthrough_endpoint_router: Final = PassthroughEndpointRouter()
|
||||
|
|
@ -1875,7 +1876,7 @@ async def vertex_proxy_route(
|
|||
)
|
||||
|
||||
|
||||
@router.api_route(
|
||||
@openai_passthrough_router.api_route(
|
||||
"/openai_passthrough/{endpoint:path}",
|
||||
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
|
||||
tags=["OpenAI Pass-through", "pass-through"],
|
||||
|
|
|
|||
|
|
@ -1,4 +1,6 @@
|
|||
import asyncio
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
from urllib.parse import urlparse
|
||||
|
|
@ -39,6 +41,32 @@ else:
|
|||
EndpointType = Any
|
||||
|
||||
|
||||
def _optional_str(value: object) -> str | None:
|
||||
return value if isinstance(value, str) else None
|
||||
|
||||
|
||||
def _optional_str_tuple(value: object) -> tuple[str, ...] | None:
|
||||
if not isinstance(value, list):
|
||||
return None
|
||||
items: Final = cast(list[object], value) # cast-ok: isinstance-narrowed; element type unknown
|
||||
return tuple(tag for tag in items if isinstance(tag, str))
|
||||
|
||||
|
||||
def _request_tags(request_metadata: Mapping[str, object]) -> tuple[str, ...] | None:
|
||||
"""Tags for the batch-cost spend row: the request's own tags when it sent any,
|
||||
otherwise the key's tags, which auth exposes as user_api_key_auth_metadata (a
|
||||
tagged key does not put its tags in the top-level metadata "tags" on the
|
||||
passthrough path)
|
||||
"""
|
||||
tags: Final = _optional_str_tuple(request_metadata.get("tags"))
|
||||
if tags:
|
||||
return tags
|
||||
key_auth_metadata: Final = request_metadata.get("user_api_key_auth_metadata")
|
||||
if isinstance(key_auth_metadata, dict):
|
||||
return _optional_str_tuple(key_auth_metadata.get("tags"))
|
||||
return None
|
||||
|
||||
|
||||
class VertexPassthroughLoggingHandler:
|
||||
@staticmethod
|
||||
def vertex_passthrough_handler(
|
||||
|
|
@ -657,11 +685,13 @@ class VertexPassthroughLoggingHandler:
|
|||
|
||||
# Store the managed object for cost tracking
|
||||
# This will be picked up by check_batch_cost polling mechanism
|
||||
is_batch_create: Final = url_route.split("?")[0].rstrip("/").endswith("batchPredictionJobs")
|
||||
VertexPassthroughLoggingHandler._store_batch_managed_object(
|
||||
unified_object_id=unified_object_id,
|
||||
batch_object=litellm_batch_response,
|
||||
model_object_id=batch_id,
|
||||
logging_obj=logging_obj,
|
||||
is_batch_create=is_batch_create,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
|
@ -779,17 +809,45 @@ class VertexPassthroughLoggingHandler:
|
|||
"kwargs": kwargs,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _log_batch_registration_result(
|
||||
finished: asyncio.Task, unified_object_id: str, model_object_id: str, is_batch_create: bool
|
||||
) -> None:
|
||||
error: Final = finished.exception() if not finished.cancelled() else None
|
||||
if finished.cancelled() or error is not None:
|
||||
consequence: Final = (
|
||||
"its cost will not be tracked" if is_batch_create else "its status and output file may be stale"
|
||||
)
|
||||
verbose_proxy_logger.error(
|
||||
"Failed to store batch managed object with unified_object_id=%s, batch_id=%s; %s: %s",
|
||||
unified_object_id,
|
||||
model_object_id,
|
||||
consequence,
|
||||
error,
|
||||
)
|
||||
return
|
||||
verbose_proxy_logger.info(
|
||||
"Stored batch managed object with unified_object_id=%s, batch_id=%s",
|
||||
unified_object_id,
|
||||
model_object_id,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _store_batch_managed_object(
|
||||
unified_object_id: str,
|
||||
batch_object: LiteLLMBatch,
|
||||
model_object_id: str,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
is_batch_create: bool,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
"""
|
||||
Store batch managed object for cost tracking.
|
||||
This will be picked up by the check_batch_cost polling mechanism.
|
||||
|
||||
A poll refreshes the batch status and file object but neither creates the row
|
||||
nor writes attribution, so the creating key and its tags are persisted from
|
||||
the create alone.
|
||||
"""
|
||||
try:
|
||||
# Get the managed files hook from the logging object
|
||||
|
|
@ -805,7 +863,7 @@ class VertexPassthroughLoggingHandler:
|
|||
|
||||
user_api_key_dict: Final = UserAPIKeyAuth(
|
||||
user_id=_request_metadata.get("user_api_key_user_id", "default-user"),
|
||||
api_key="",
|
||||
api_key=_optional_str(_request_metadata.get("user_api_key")),
|
||||
team_id=_request_metadata.get("user_api_key_team_id"),
|
||||
team_alias=None,
|
||||
user_role=LitellmUserRoles.CUSTOMER, # Use proper enum value
|
||||
|
|
@ -827,9 +885,7 @@ class VertexPassthroughLoggingHandler:
|
|||
)
|
||||
|
||||
# Store the unified object for batch cost tracking
|
||||
import asyncio
|
||||
|
||||
asyncio.create_task(
|
||||
task: Final = asyncio.create_task(
|
||||
managed_files_hook.store_unified_object_id(
|
||||
unified_object_id=unified_object_id,
|
||||
file_object=batch_object,
|
||||
|
|
@ -837,13 +893,15 @@ class VertexPassthroughLoggingHandler:
|
|||
model_object_id=model_object_id,
|
||||
file_purpose="batch",
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_tags=_request_tags(_request_metadata),
|
||||
persist_attribution=is_batch_create,
|
||||
create_if_missing=is_batch_create,
|
||||
)
|
||||
)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
"Stored batch managed object with unified_object_id=%s, batch_id=%s",
|
||||
unified_object_id,
|
||||
model_object_id,
|
||||
task.add_done_callback(
|
||||
lambda finished: VertexPassthroughLoggingHandler._log_batch_registration_result(
|
||||
finished, unified_object_id, model_object_id, is_batch_create
|
||||
)
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.warning(
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ import threading
|
|||
import time
|
||||
import traceback
|
||||
import warnings
|
||||
from collections.abc import AsyncGenerator, Callable, Mapping
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Mapping
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from types import MappingProxyType, UnionType
|
||||
from typing import (
|
||||
|
|
@ -23,6 +23,7 @@ from typing import (
|
|||
Any,
|
||||
Final,
|
||||
Literal,
|
||||
NamedTuple,
|
||||
Optional,
|
||||
TypedDict,
|
||||
Union,
|
||||
|
|
@ -536,6 +537,7 @@ from litellm.proxy.openai_files_endpoints.files_endpoints import (
|
|||
set_files_config,
|
||||
)
|
||||
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
||||
openai_passthrough_router,
|
||||
passthrough_endpoint_router,
|
||||
vertex_ai_live_websocket_passthrough,
|
||||
)
|
||||
|
|
@ -611,7 +613,7 @@ from litellm.secret_managers.main import (
|
|||
normalize_nonempty_secret_str,
|
||||
str_to_bool,
|
||||
)
|
||||
from litellm.types.integrations.slack_alerting import SlackAlertingArgs
|
||||
from litellm.types.integrations.slack_alerting import AlertType, SlackAlertingArgs
|
||||
from litellm.types.llms.anthropic import (
|
||||
AnthropicMessagesRequest,
|
||||
AnthropicResponse,
|
||||
|
|
@ -6859,9 +6861,19 @@ class ProxyConfig:
|
|||
guardrail_id = guardrail.get("guardrail_id")
|
||||
if guardrail_id:
|
||||
db_guardrail_ids.add(guardrail_id)
|
||||
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(
|
||||
guardrail=cast(Guardrail, guardrail),
|
||||
)
|
||||
try:
|
||||
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(
|
||||
guardrail=cast(Guardrail, guardrail),
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # one unloadable row must not stop the remaining guardrails
|
||||
verbose_proxy_logger.error(
|
||||
"litellm.proxy.proxy_server.py::ProxyConfig:_init_guardrails_in_db - "
|
||||
"skipping guardrail '%s' (ID: %s): %s: %s",
|
||||
guardrail.get("guardrail_name"),
|
||||
guardrail_id,
|
||||
type(e).__name__,
|
||||
e,
|
||||
)
|
||||
|
||||
# Drop in-memory DB-backed entries whose row was deleted on another
|
||||
# pod. Config-loaded entries are never touched.
|
||||
|
|
@ -7632,6 +7644,200 @@ def _pop_complete_sse_frame(buffer: str) -> tuple[str | None, str]:
|
|||
return buffer[:frame_end], buffer[frame_end:]
|
||||
|
||||
|
||||
_STREAM_KEEPALIVE: Final = object()
|
||||
|
||||
_KEEPALIVE_MIN_SECONDS: Final = 1.0
|
||||
_KEEPALIVE_MAX_SECONDS: Final = 300.0
|
||||
_EMPTY_MAPPING: Final[Mapping[str, Any]] = MappingProxyType({})
|
||||
|
||||
|
||||
async def _iter_with_keepalive(
|
||||
aiter: AsyncIterator[Any],
|
||||
resolve_keepalive_seconds: Callable[[object], float],
|
||||
keepalive_seconds: float,
|
||||
) -> AsyncGenerator[Any, None]:
|
||||
"""Wrap `aiter` with idle-gap heartbeats, re-resolving the interval after each
|
||||
real chunk via `resolve_keepalive_seconds`. A mid-stream router fallback can
|
||||
swap in a deployment with a different keepalive policy, including one that
|
||||
newly enables or newly disables heartbeats, partway through the same stream;
|
||||
re-resolving against each chunk's own identity (rather than trusting the
|
||||
interval picked before iteration started, or picked the last time it went
|
||||
inactive) keeps the heartbeat behavior in sync with whichever deployment
|
||||
actually produced it, in both directions. While the interval is <= 0, no
|
||||
task is created and no timeout is awaited: a chunk is forwarded the moment
|
||||
it arrives, at the same cost as a bare `async for`."""
|
||||
pending: asyncio.Task[Any] | None = None # rebind-ok: rebound each loop iteration
|
||||
current_keepalive_seconds = keepalive_seconds # rebind-ok: re-resolved after each chunk
|
||||
try:
|
||||
while True:
|
||||
if current_keepalive_seconds <= 0:
|
||||
try:
|
||||
item = await aiter.__anext__()
|
||||
except StopAsyncIteration:
|
||||
break
|
||||
yield item
|
||||
current_keepalive_seconds = resolve_keepalive_seconds(item)
|
||||
continue
|
||||
|
||||
if pending is None:
|
||||
pending = asyncio.create_task(aiter.__anext__())
|
||||
done, _ = await asyncio.wait((pending,), timeout=current_keepalive_seconds)
|
||||
if not done:
|
||||
yield _STREAM_KEEPALIVE
|
||||
continue
|
||||
try:
|
||||
item = pending.result()
|
||||
except StopAsyncIteration:
|
||||
break
|
||||
finally:
|
||||
pending = None
|
||||
yield item
|
||||
current_keepalive_seconds = resolve_keepalive_seconds(item)
|
||||
finally:
|
||||
if pending is not None and not pending.done():
|
||||
pending.cancel()
|
||||
try:
|
||||
await pending
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
|
||||
class _DeploymentKeepaliveConfig(NamedTuple):
|
||||
keepalive_seconds: Any
|
||||
allow_client_override: bool
|
||||
|
||||
|
||||
def _keepalive_from_deployment_config(
|
||||
request_data: Mapping[str, Any], response: object
|
||||
) -> _DeploymentKeepaliveConfig | None:
|
||||
if llm_router is None:
|
||||
return None
|
||||
|
||||
hidden: Final = get_hidden_params_dict(response)
|
||||
model_id: Final = hidden.get("model_id")
|
||||
if isinstance(model_id, str) and model_id:
|
||||
deployment: Final = llm_router.get_deployment(model_id=model_id)
|
||||
# A populated model_id names the specific deployment that served this
|
||||
# stream. If it no longer resolves (e.g. removed by a config reload
|
||||
# mid-stream), that's a stale identity, not an absent one: don't fall
|
||||
# through to guessing via model_name below, since a currently-live
|
||||
# sibling deployment's config was never what actually served this
|
||||
# stream.
|
||||
if deployment is None:
|
||||
return None
|
||||
return _DeploymentKeepaliveConfig(
|
||||
keepalive_seconds=getattr(deployment.litellm_params, "keepalive_seconds", None),
|
||||
allow_client_override=bool(getattr(deployment.litellm_params, "allow_client_keepalive_override", False)),
|
||||
)
|
||||
|
||||
# No model_id at all to pin down which deployment actually served this
|
||||
# stream: only trust the fallback when every deployment under this
|
||||
# model_name agrees on both keepalive_seconds and
|
||||
# allow_client_keepalive_override (including deployments that leave either
|
||||
# field unset), so a stream never inherits a sibling deployment's policy.
|
||||
configs: Final = frozenset(
|
||||
(
|
||||
(deployment_dict.get("litellm_params") or _EMPTY_MAPPING).get("keepalive_seconds"),
|
||||
bool(
|
||||
(deployment_dict.get("litellm_params") or _EMPTY_MAPPING).get("allow_client_keepalive_override", False)
|
||||
),
|
||||
)
|
||||
for deployment_dict in llm_router.get_model_list(model_name=request_data.get("model")) or ()
|
||||
)
|
||||
if len(configs) == 1:
|
||||
keepalive_seconds, allow_client_override = next(iter(configs))
|
||||
return _DeploymentKeepaliveConfig(
|
||||
keepalive_seconds=keepalive_seconds, allow_client_override=allow_client_override
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def _is_explicit_keepalive_disable(raw: object) -> bool:
|
||||
if not isinstance(raw, (int, float, str)):
|
||||
return False
|
||||
try:
|
||||
return float(raw) <= 0
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
|
||||
def _resolve_keepalive_seconds(request_data: Mapping[str, Any], response: object = None) -> float:
|
||||
deployment_config: Final = _keepalive_from_deployment_config(request_data, response)
|
||||
deployment_raw: Final = deployment_config.keepalive_seconds if deployment_config is not None else None
|
||||
allow_client_override: Final = deployment_config.allow_client_override if deployment_config is not None else False
|
||||
|
||||
# An operator setting keepalive_seconds: 0 on a deployment is an explicit hard
|
||||
# disable: an authenticated client must not be able to re-enable heartbeats
|
||||
# (and the idle-timeout evasion that comes with them) for a deployment the
|
||||
# operator opted out of, regardless of what the request body asks for.
|
||||
if _is_explicit_keepalive_disable(deployment_raw):
|
||||
return 0.0
|
||||
|
||||
# keepalive_seconds is operator-only unless the deployment explicitly opts in:
|
||||
# a client can't unilaterally enable heartbeats (and the LB-idle-timeout
|
||||
# evasion that comes with them) for a deployment that never configured this.
|
||||
client_supplied: Final = request_data.get("keepalive_seconds") if allow_client_override else None
|
||||
raw: Final = client_supplied if client_supplied is not None else deployment_raw
|
||||
try:
|
||||
value: Final = float(raw) if isinstance(raw, (int, float, str)) else 0.0
|
||||
except ValueError:
|
||||
return 0.0
|
||||
if value <= 0:
|
||||
return 0.0
|
||||
clamped: Final = max(_KEEPALIVE_MIN_SECONDS, min(value, _KEEPALIVE_MAX_SECONDS))
|
||||
if clamped != value:
|
||||
verbose_proxy_logger.info(
|
||||
"keepalive_seconds=%s clamped to %s [min=%s, max=%s]",
|
||||
value,
|
||||
clamped,
|
||||
_KEEPALIVE_MIN_SECONDS,
|
||||
_KEEPALIVE_MAX_SECONDS,
|
||||
)
|
||||
return clamped
|
||||
|
||||
|
||||
_KEEPALIVE_CACHE_TTL_SECONDS: Final = 5.0
|
||||
|
||||
|
||||
def _make_keepalive_resolver(request_data: Mapping[str, Any]) -> Callable[[object], float]:
|
||||
"""Wrap `_resolve_keepalive_seconds` with a memo keyed on the serving
|
||||
deployment's model_id. The steady-state case (no mid-stream fallback, the
|
||||
overwhelming majority of streams) sees the same model_id on every chunk, so
|
||||
this turns the per-chunk cost from a full `llm_router.get_deployment()`
|
||||
Pydantic rebuild into a cheap hidden-params read once per
|
||||
`_KEEPALIVE_CACHE_TTL_SECONDS` for that model_id. The cache expires on its
|
||||
own rather than living for the life of the stream, so an operator's live
|
||||
config change (disabling keepalive, revoking client override, or removing
|
||||
the deployment) is observed within a bounded window instead of being able
|
||||
to be evaded by an already-in-flight stream indefinitely. A missing/empty
|
||||
model_id can't be trusted as a cache key (see
|
||||
`_keepalive_from_deployment_config`'s model_name fallback, which reflects
|
||||
current router state rather than one deployment's fixed identity), so
|
||||
those chunks always resolve fresh, matching prior behavior exactly.
|
||||
"""
|
||||
last_model_id: str | None = None # rebind-ok: memoized identity of the last-resolved chunk
|
||||
last_value: float = 0.0 # rebind-ok: cached resolution for last_model_id
|
||||
last_resolved_at: float = float("-inf") # rebind-ok: monotonic timestamp of the last real resolution
|
||||
|
||||
def _resolve(item: object) -> float:
|
||||
nonlocal last_model_id, last_value, last_resolved_at
|
||||
model_id = get_hidden_params_dict(item).get("model_id")
|
||||
now: Final = time.monotonic()
|
||||
if (
|
||||
isinstance(model_id, str)
|
||||
and model_id
|
||||
and model_id == last_model_id
|
||||
and now - last_resolved_at < _KEEPALIVE_CACHE_TTL_SECONDS
|
||||
):
|
||||
return last_value
|
||||
value: Final = _resolve_keepalive_seconds(request_data, item)
|
||||
if isinstance(model_id, str) and model_id:
|
||||
last_model_id, last_value, last_resolved_at = model_id, value, now
|
||||
return value
|
||||
|
||||
return _resolve
|
||||
|
||||
|
||||
async def async_data_generator(
|
||||
response,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -7680,7 +7886,28 @@ async def async_data_generator(
|
|||
else:
|
||||
stream_iterator = response
|
||||
|
||||
async for chunk in stream_iterator:
|
||||
# A stream can start on a deployment with keepalive off and fall back
|
||||
# mid-stream to one that enables it: only skip wrapping altogether when
|
||||
# there's no router to ever fall back through in the first place (in
|
||||
# which case _resolve_keepalive_seconds can never return non-zero for
|
||||
# any chunk of this stream), not merely because the first chunk's
|
||||
# deployment happens to start with it off.
|
||||
resolve_keepalive_seconds: Final = _make_keepalive_resolver(request_data)
|
||||
stream_source: Final = (
|
||||
_iter_with_keepalive(
|
||||
stream_iterator.__aiter__(),
|
||||
resolve_keepalive_seconds,
|
||||
resolve_keepalive_seconds(response),
|
||||
)
|
||||
if llm_router is not None
|
||||
else stream_iterator
|
||||
)
|
||||
|
||||
async for item in stream_source:
|
||||
if item is _STREAM_KEEPALIVE:
|
||||
yield ": ping\n\n"
|
||||
continue
|
||||
chunk = cast(Any, item) # cast-ok: sentinel already handled above, item is a real chunk here
|
||||
if needs_per_chunk_hook:
|
||||
### CALL HOOKS ### - modify outgoing data
|
||||
chunk, _str_so_far = await _apply_streaming_chunk_hooks(
|
||||
|
|
@ -8469,6 +8696,47 @@ class ProxyStartupEvent:
|
|||
|
||||
await cls._initialize_spend_tracking_background_jobs(scheduler=scheduler)
|
||||
|
||||
### PTU DAILY ROLLUP ###
|
||||
from litellm.proxy.spend_tracking.ptu_feature_flag import (
|
||||
is_ptu_cost_attribution_enabled,
|
||||
)
|
||||
|
||||
if is_ptu_cost_attribution_enabled():
|
||||
from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import (
|
||||
PTU_ROLLUP_JOB_ID,
|
||||
run_scheduled_ptu_rollup,
|
||||
)
|
||||
|
||||
async def _alert_ptu_rollup_failure(message: str) -> None:
|
||||
await proxy_logging_obj.alerting_handler(
|
||||
message=message,
|
||||
level="High",
|
||||
alert_type=AlertType.failed_tracking_spend,
|
||||
)
|
||||
|
||||
async def _scheduled_ptu_rollup() -> None:
|
||||
# Reuse the PodLockManager from db_spend_update_writer so only one pod
|
||||
# reconciles a day; a multi-pod race could prune another pod's fresh rows
|
||||
await run_scheduled_ptu_rollup(
|
||||
prisma_client,
|
||||
pod_lock_manager=proxy_logging_obj.db_spend_update_writer.pod_lock_manager,
|
||||
alert=_alert_ptu_rollup_failure,
|
||||
)
|
||||
|
||||
scheduler.add_job(
|
||||
_scheduled_ptu_rollup,
|
||||
"cron",
|
||||
hour=0,
|
||||
minute=15,
|
||||
timezone="UTC",
|
||||
id=PTU_ROLLUP_JOB_ID,
|
||||
replace_existing=True,
|
||||
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
"PTU rollup job scheduled at 00:15 UTC daily (only models with PTU config accrue flat cost)"
|
||||
)
|
||||
|
||||
### SPEND LOG CLEANUP ###
|
||||
if (
|
||||
general_settings.get("maximum_spend_logs_retention_period") is not None
|
||||
|
|
@ -16607,6 +16875,7 @@ app.include_router(search_router)
|
|||
app.include_router(image_router)
|
||||
app.include_router(fine_tuning_router)
|
||||
app.include_router(credential_router)
|
||||
app.include_router(openai_passthrough_router)
|
||||
app.include_router(batches_router)
|
||||
app.include_router(openai_files_router)
|
||||
app.include_router(llm_passthrough_router)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -893,6 +893,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 +986,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?
|
||||
|
||||
|
|
|
|||
18
litellm/proxy/spend_tracking/ptu_feature_flag.py
Normal file
18
litellm/proxy/spend_tracking/ptu_feature_flag.py
Normal file
|
|
@ -0,0 +1,18 @@
|
|||
"""Opt-in flag for PTU (provisioned throughput unit) flat-cost attribution.
|
||||
|
||||
The whole feature is inert unless an operator sets
|
||||
``LITELLM_ENABLE_PTU_COST_ATTRIBUTION``: the daily rollup is not scheduled, the
|
||||
model endpoints reject PTU config, the daily activity read path reports zero flat
|
||||
cost, and the model form hides the PTU inputs.
|
||||
"""
|
||||
|
||||
from typing import Final
|
||||
|
||||
from litellm.secret_managers.main import get_secret_bool
|
||||
|
||||
PTU_COST_ATTRIBUTION_ENV_VAR: Final = "LITELLM_ENABLE_PTU_COST_ATTRIBUTION"
|
||||
|
||||
|
||||
def is_ptu_cost_attribution_enabled() -> bool:
|
||||
"""Report whether this deployment opted into PTU flat-cost attribution."""
|
||||
return get_secret_bool(PTU_COST_ATTRIBUTION_ENV_VAR, False) is True
|
||||
663
litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py
Normal file
663
litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py
Normal file
|
|
@ -0,0 +1,663 @@
|
|||
"""
|
||||
Daily rollup for per-model PTU (provisioned throughput) flat cost.
|
||||
|
||||
v1 reads PTU config straight off the model deployment
|
||||
(``LiteLLM_ProxyModelTable.model_info``): a deployment carrying ``ptu_count``
|
||||
and ``cost_per_ptu_per_hour`` accrues flat cost of
|
||||
``ptu_count * cost_per_ptu_per_hour * active_hours`` for a given UTC day, where
|
||||
``active_hours`` is the overlap between the day and the optional
|
||||
``[ptu_effective_from, ptu_effective_to)`` window (a window opening at 23:00
|
||||
charges one hour that day). The amount is written to ``LiteLLM_DailyTeamSpend``
|
||||
under a sentinel api_key so the rows are distinguishable from per-request rows
|
||||
and share the existing unique constraint.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from datetime import date, datetime, time, timedelta, timezone
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import (
|
||||
PTU_PRUNE_SKEW_GRACE_SECONDS,
|
||||
PTU_ROLLUP_JOB_ID,
|
||||
PTU_ROLLUP_LOCK_TTL_SECONDS,
|
||||
PTU_ROLLUP_MAX_BACKFILL_DAYS,
|
||||
PTU_SENTINEL_API_KEY,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled
|
||||
from litellm.types.router import ModelInfo
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
_HOURS_PER_DAY: Final = 24
|
||||
_UPSERT_ATTEMPTS: Final = 3
|
||||
_UPSERT_RETRY_BACKOFF_SECONDS: Final = 0.5
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RollupResult:
|
||||
day: date
|
||||
models_processed: int
|
||||
rows_written: int
|
||||
rows_failed: int = 0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BackfillResult:
|
||||
start: date
|
||||
end: date
|
||||
days_scanned: int
|
||||
rows_written: int
|
||||
rows_failed: int = 0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PTUModel:
|
||||
"""A model deployment carrying valid manual PTU config."""
|
||||
|
||||
model_id: str
|
||||
model_name: str
|
||||
team_id: str
|
||||
ptu_count: int
|
||||
cost_per_ptu_per_hour: float
|
||||
effective_from: datetime | None = None
|
||||
effective_to: datetime | None = None
|
||||
|
||||
|
||||
def _parse_utc_datetime(value: object) -> datetime | None:
|
||||
"""Parse a model_info datetime (ISO string or datetime) into a UTC-aware datetime, else None."""
|
||||
parsed: Final = _coerce_datetime(value)
|
||||
if parsed is None:
|
||||
return None
|
||||
if parsed.tzinfo is None:
|
||||
return parsed.replace(tzinfo=timezone.utc)
|
||||
return parsed.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def _coerce_datetime(value: object) -> datetime | None:
|
||||
"""``value`` as a datetime, parsing an ISO string, else None."""
|
||||
if isinstance(value, datetime):
|
||||
return value
|
||||
if not isinstance(value, str):
|
||||
return None
|
||||
try:
|
||||
return datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
|
||||
def _public_model_name(row: object, model_info: Mapping[str, object]) -> str:
|
||||
"""The name an operator recognises for this deployment.
|
||||
|
||||
Creating a team-scoped deployment rewrites model_name to a synthetic routing key
|
||||
(``model_name_<team_id>_<uuid4>``) and keeps the chosen name in
|
||||
``model_info.team_public_model_name``. PTU config is only accepted alongside a
|
||||
team_id, so every PTU deployment carries that synthetic name; keying the sentinel
|
||||
row on it would file each charge under a UUID that no usage view can resolve and
|
||||
that never lines up with the same model's request rows.
|
||||
"""
|
||||
public_name: Final = model_info.get("team_public_model_name")
|
||||
if isinstance(public_name, str) and public_name:
|
||||
return public_name
|
||||
return str(getattr(row, "model_name", "") or "")
|
||||
|
||||
|
||||
def _decode_model_info(raw: object) -> "Mapping[str, object] | None":
|
||||
"""A deployment's model_info as a dict, decoding a JSON string, else None."""
|
||||
if isinstance(raw, str):
|
||||
try:
|
||||
return json.loads(raw)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
if isinstance(raw, dict):
|
||||
return raw
|
||||
return None
|
||||
|
||||
|
||||
def _parse_ptu_model(row: object) -> PTUModel | None:
|
||||
"""Return a PTUModel when the deployment carries valid manual PTU config, else None.
|
||||
|
||||
Valid means model_info has a positive ptu_count, a non-negative
|
||||
cost_per_ptu_per_hour, and a team_id (1 model -> 1 team).
|
||||
"""
|
||||
raw_model_info: Final = getattr(row, "model_info", None)
|
||||
model_info: Final = _decode_model_info(raw_model_info)
|
||||
if model_info is None:
|
||||
return None
|
||||
ptu_count: Final = model_info.get("ptu_count")
|
||||
cost_per_hour: Final = model_info.get("cost_per_ptu_per_hour")
|
||||
team_id: Final = model_info.get("team_id")
|
||||
if ptu_count is None or cost_per_hour is None or not team_id:
|
||||
return None
|
||||
try:
|
||||
ptu_count_int: Final = int(ptu_count)
|
||||
cost_per_hour_float: Final = float(cost_per_hour)
|
||||
except (TypeError, ValueError, OverflowError):
|
||||
return None
|
||||
if not 0 < ptu_count_int <= ModelInfo.MAX_PTU_COUNT:
|
||||
return None
|
||||
if not 0 <= cost_per_hour_float <= ModelInfo.MAX_COST_PER_PTU_PER_HOUR:
|
||||
return None
|
||||
if model_info.get("ptu_effective_from") is None:
|
||||
# The endpoints require a start; a row without one predates that rule or was
|
||||
# written around them, and inferring one would bill days the deployment did not exist
|
||||
return None
|
||||
raw_from: Final = model_info.get("ptu_effective_from")
|
||||
raw_to: Final = model_info.get("ptu_effective_to")
|
||||
effective_from: Final = _parse_utc_datetime(raw_from)
|
||||
effective_to: Final = _parse_utc_datetime(raw_to)
|
||||
# A present-but-unparseable bound would read as "no bound" and silently widen the
|
||||
# window to the whole day, so the deployment is skipped until the config is fixed
|
||||
if (raw_from is not None and effective_from is None) or (raw_to is not None and effective_to is None):
|
||||
return None
|
||||
if effective_from is not None and effective_to is not None and effective_to <= effective_from:
|
||||
return None
|
||||
return PTUModel(
|
||||
model_id=str(getattr(row, "model_id", "") or ""),
|
||||
model_name=_public_model_name(row, model_info),
|
||||
team_id=str(team_id),
|
||||
ptu_count=ptu_count_int,
|
||||
cost_per_ptu_per_hour=cost_per_hour_float,
|
||||
effective_from=effective_from,
|
||||
effective_to=effective_to,
|
||||
)
|
||||
|
||||
|
||||
def _active_hours_on_day(model: PTUModel, day: date) -> float:
|
||||
"""Hours the model's PTU window overlaps ``day`` (UTC), clamped to [0, 24]."""
|
||||
day_start: Final = datetime.combine(day, time.min, tzinfo=timezone.utc)
|
||||
day_end: Final = day_start + timedelta(days=1)
|
||||
start: Final = max(day_start, model.effective_from) if model.effective_from else day_start
|
||||
end: Final = min(day_end, model.effective_to) if model.effective_to else day_end
|
||||
if end <= start:
|
||||
return 0.0
|
||||
return (end - start).total_seconds() / 3600.0
|
||||
|
||||
|
||||
def _compute_daily_flat_cost(model: PTUModel, day: date) -> float:
|
||||
"""Flat cost for ``day``: ptu_count * cost_per_ptu_per_hour * active_hours."""
|
||||
return float(model.ptu_count) * model.cost_per_ptu_per_hour * _active_hours_on_day(model, day)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _PTUCharge:
|
||||
"""One sentinel row's worth of flat cost for a deployment on a day.
|
||||
|
||||
``model_id`` is the row's identity and goes in the unique key; ``model_name`` is what
|
||||
an operator reads and rides alongside it. A deployment can be renamed, so keying on
|
||||
the name would let two runs holding different config views write the same day twice.
|
||||
"""
|
||||
|
||||
team_id: str
|
||||
model_id: str
|
||||
model_name: str
|
||||
flat_cost: float
|
||||
|
||||
|
||||
def _aggregate_charges(ptu_models: tuple[PTUModel, ...], day: date) -> tuple[_PTUCharge, ...]:
|
||||
"""One charge per deployment that accrues cost on ``day``. Zero-cost deployments are
|
||||
dropped, which keeps a day outside a window from writing a row.
|
||||
|
||||
Deployments sharing a public name inside a team no longer need collapsing: each keys
|
||||
its own row on its own id, and the read path merges them back under the shared name.
|
||||
"""
|
||||
return tuple(
|
||||
_PTUCharge(
|
||||
team_id=model.team_id,
|
||||
model_id=model.model_id,
|
||||
model_name=model.model_name,
|
||||
flat_cost=_compute_daily_flat_cost(model, day),
|
||||
)
|
||||
for model in sorted(ptu_models, key=lambda m: (m.team_id, m.model_id))
|
||||
if _compute_daily_flat_cost(model, day) > 0
|
||||
)
|
||||
|
||||
|
||||
async def _upsert_ptu_daily_row(
|
||||
prisma_client: "PrismaClient",
|
||||
*,
|
||||
team_id: str,
|
||||
model_id: str,
|
||||
model_name: str,
|
||||
date_str: str,
|
||||
flat_cost: float,
|
||||
) -> None:
|
||||
"""Idempotent upsert of a sentinel-api_key row on LiteLLM_DailyTeamSpend.
|
||||
|
||||
``model`` holds the deployment id because it is part of the table's unique key and a
|
||||
rename must not move the row. ``model_group`` carries the operator-facing name, which
|
||||
is outside the key and is what the usage views display.
|
||||
"""
|
||||
where: Final = { # mutable-ok: prisma upsert filter payload
|
||||
"team_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint": { # mutable-ok: prisma composite-key filter
|
||||
"team_id": team_id,
|
||||
"date": date_str,
|
||||
"api_key": PTU_SENTINEL_API_KEY,
|
||||
"model": model_id,
|
||||
"custom_llm_provider": "",
|
||||
"mcp_namespaced_tool_name": "",
|
||||
"endpoint": "",
|
||||
}
|
||||
}
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
await prisma_client.db.litellm_dailyteamspend.upsert(
|
||||
where=where,
|
||||
data={ # mutable-ok: prisma upsert data payload
|
||||
"create": { # mutable-ok: prisma create payload
|
||||
"team_id": team_id,
|
||||
"date": date_str,
|
||||
"api_key": PTU_SENTINEL_API_KEY,
|
||||
"model": model_id,
|
||||
"model_group": model_name,
|
||||
"custom_llm_provider": "",
|
||||
"mcp_namespaced_tool_name": "",
|
||||
"endpoint": "",
|
||||
"ptu_flat_cost": flat_cost,
|
||||
},
|
||||
"update": { # mutable-ok: prisma update payload
|
||||
"model_group": model_name,
|
||||
"ptu_flat_cost": flat_cost,
|
||||
"updated_at": now,
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def _upsert_charge_with_retry(
|
||||
prisma_client: "PrismaClient",
|
||||
*,
|
||||
charge: _PTUCharge,
|
||||
date_str: str,
|
||||
) -> bool:
|
||||
"""Write one charge, retrying transient failures. Returns False once attempts are spent.
|
||||
|
||||
The upsert is idempotent on the sentinel unique key, so a retry can only rewrite the
|
||||
same amount for the same day. Retrying in-run matters because the scheduled job moves
|
||||
on to the next date: a write lost here is a day of PTU cost that no later run replays.
|
||||
"""
|
||||
for attempt in range(1, _UPSERT_ATTEMPTS + 1):
|
||||
try:
|
||||
await _upsert_ptu_daily_row(
|
||||
prisma_client,
|
||||
team_id=charge.team_id,
|
||||
model_id=charge.model_id,
|
||||
model_name=charge.model_name,
|
||||
date_str=date_str,
|
||||
flat_cost=charge.flat_cost,
|
||||
)
|
||||
return True
|
||||
except Exception as exc: # noqa: BLE001 # one bad row must not stop the batch
|
||||
if attempt < _UPSERT_ATTEMPTS:
|
||||
verbose_proxy_logger.warning(
|
||||
"PTU rollup: upsert attempt %d/%d failed for team=%s model=%s day=%s: %s",
|
||||
attempt,
|
||||
_UPSERT_ATTEMPTS,
|
||||
charge.team_id,
|
||||
charge.model_name,
|
||||
date_str,
|
||||
exc,
|
||||
)
|
||||
await asyncio.sleep(_UPSERT_RETRY_BACKOFF_SECONDS * attempt)
|
||||
continue
|
||||
verbose_proxy_logger.error(
|
||||
"PTU rollup: upsert failed after %d attempts for team=%s model=%s day=%s "
|
||||
"(rerun the rollup for that date to recover): %s",
|
||||
_UPSERT_ATTEMPTS,
|
||||
charge.team_id,
|
||||
charge.model_name,
|
||||
date_str,
|
||||
exc,
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
async def _load_ptu_models(prisma_client: "PrismaClient") -> tuple[PTUModel, ...]:
|
||||
"""Every model deployment currently carrying valid manual PTU config."""
|
||||
rows: Final = await prisma_client.db.litellm_proxymodeltable.find_many()
|
||||
return tuple(parsed for parsed in (_parse_ptu_model(row) for row in rows) if parsed is not None)
|
||||
|
||||
|
||||
async def run_ptu_flat_cost_rollup(
|
||||
prisma_client: "PrismaClient",
|
||||
target_date: date | None = None,
|
||||
may_prune: bool = True,
|
||||
) -> RollupResult:
|
||||
"""Rollup one UTC day of flat PTU cost across all PTU-configured model deployments.
|
||||
|
||||
Defaults to yesterday UTC. Authoritative for the day: it upserts the current charges
|
||||
first, then deletes the day's sentinel rows this run did not refresh, so a
|
||||
since-removed, invalidated, or now-out-of-window deployment leaves no stale charge.
|
||||
|
||||
The prune predicate is ``updated_at < run_started`` rather than "not in the charge
|
||||
set I computed", which matters under concurrency: whether a row is garbage becomes a
|
||||
property of the row instead of one run's in-memory config snapshot, so a run can
|
||||
never delete a row a concurrent run just wrote. It is still skipped when any charge
|
||||
failed to write, since a row whose replacement never landed would look unrefreshed.
|
||||
"""
|
||||
day: Final = target_date or (datetime.now(timezone.utc).date() - timedelta(days=1))
|
||||
|
||||
if prisma_client is None:
|
||||
verbose_proxy_logger.warning("PTU rollup: prisma_client is None, skipping")
|
||||
return RollupResult(day=day, models_processed=0, rows_written=0)
|
||||
|
||||
date_str: Final = day.isoformat()
|
||||
run_started: Final = datetime.now(timezone.utc)
|
||||
|
||||
ptu_models: Final = await _load_ptu_models(prisma_client)
|
||||
charges: Final = _aggregate_charges(ptu_models, day)
|
||||
|
||||
landed: Final = tuple(
|
||||
[await _upsert_charge_with_retry(prisma_client, charge=charge, date_str=date_str) for charge in charges]
|
||||
)
|
||||
rows_written: Final = sum(landed)
|
||||
rows_failed: Final = len(charges) - rows_written
|
||||
|
||||
if not may_prune:
|
||||
verbose_proxy_logger.info(
|
||||
"PTU rollup for %s: ran without the cross-pod lock, skipping the prune so a "
|
||||
"concurrent pod's charges cannot be swept by this run's cutoff",
|
||||
date_str,
|
||||
)
|
||||
elif rows_failed:
|
||||
# A charge that never landed leaves its row looking unrefreshed, so the prune
|
||||
# would delete the very row the failed write was meant to replace
|
||||
verbose_proxy_logger.warning(
|
||||
"PTU rollup: %d charge(s) failed for %s, skipping the prune so a row whose "
|
||||
"replacement did not land is not deleted; rerun that date to reconcile",
|
||||
rows_failed,
|
||||
date_str,
|
||||
)
|
||||
else:
|
||||
await _prune_unrefreshed_sentinel_rows(prisma_client, date_str=date_str, run_started=run_started)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
"PTU rollup for %s: %d PTU models processed, %d rows written, %d rows failed",
|
||||
date_str,
|
||||
len(ptu_models),
|
||||
rows_written,
|
||||
rows_failed,
|
||||
)
|
||||
return RollupResult(
|
||||
day=day,
|
||||
models_processed=len(ptu_models),
|
||||
rows_written=rows_written,
|
||||
rows_failed=rows_failed,
|
||||
)
|
||||
|
||||
|
||||
def _backfill_window(ptu_models: tuple[PTUModel, ...], end: date) -> tuple[date, ...]:
|
||||
"""The UTC days the catch-up pass considers, oldest first, through ``end`` inclusive.
|
||||
|
||||
Starts at the earliest declared ``ptu_effective_from``, floored at
|
||||
``PTU_ROLLUP_MAX_BACKFILL_DAYS`` before ``end``. A start is required alongside the
|
||||
count and rate, so a deployment without one is not priced rather than being given the
|
||||
floor, which would bill it for the whole cap window. Empty when there is no PTU
|
||||
config, or when every declared window opens after ``end``.
|
||||
"""
|
||||
floor: Final = end - timedelta(days=PTU_ROLLUP_MAX_BACKFILL_DAYS)
|
||||
starts: Final = tuple(model.effective_from.date() for model in ptu_models if model.effective_from)
|
||||
if not starts:
|
||||
return ()
|
||||
start: Final = max(min(starts), floor)
|
||||
return tuple(start + timedelta(days=offset) for offset in range((end - start).days + 1))
|
||||
|
||||
|
||||
async def _existing_sentinel_keys(
|
||||
prisma_client: "PrismaClient",
|
||||
*,
|
||||
start: date,
|
||||
end: date,
|
||||
) -> frozenset[tuple[str, str, str]]:
|
||||
"""``(team_id, deployment id, date)`` of every PTU sentinel row within ``[start, end]``.
|
||||
|
||||
The row's ``model`` column holds the deployment id, so this is an exact identity and
|
||||
survives a rename. Nothing here reads the display name.
|
||||
"""
|
||||
date_range: Final = {"gte": start.isoformat(), "lte": end.isoformat()} # mutable-ok: prisma range filter
|
||||
rows: Final = await prisma_client.db.litellm_dailyteamspend.find_many(
|
||||
where={"api_key": PTU_SENTINEL_API_KEY, "date": date_range} # mutable-ok: prisma find filter
|
||||
)
|
||||
return frozenset(
|
||||
(
|
||||
str(getattr(row, "team_id", "") or ""),
|
||||
str(getattr(row, "model", "") or ""),
|
||||
str(getattr(row, "date", "") or ""),
|
||||
)
|
||||
for row in rows
|
||||
)
|
||||
|
||||
|
||||
async def run_ptu_flat_cost_backfill(
|
||||
prisma_client: "PrismaClient",
|
||||
today: date | None = None,
|
||||
) -> BackfillResult:
|
||||
"""Price the elapsed days of every PTU window that carry no sentinel row yet.
|
||||
|
||||
Writes only the charges that are missing and never rewrites or deletes an existing
|
||||
row, so a day already priced keeps the amount it was billed, whatever the config says
|
||||
now. A day counts as priced when a sentinel row exists for that deployment id, so
|
||||
renaming a deployment neither re-prices its history nor files a second charge beside
|
||||
the row already there. Zero-cost days write nothing, which leaves a day
|
||||
outside a window reconsidered on each run rather than recorded as done.
|
||||
|
||||
It deletes nothing. Removing a deployment stops it accruing new charges and leaves the
|
||||
days it was billed for standing, since those days were incurred.
|
||||
"""
|
||||
end: Final = (today or datetime.now(timezone.utc).date()) - timedelta(days=1)
|
||||
|
||||
if not prisma_client:
|
||||
verbose_proxy_logger.warning("PTU backfill: prisma_client is None, skipping")
|
||||
return BackfillResult(start=end, end=end, days_scanned=0, rows_written=0)
|
||||
|
||||
ptu_models: Final = await _load_ptu_models(prisma_client)
|
||||
days: Final = _backfill_window(ptu_models, end)
|
||||
|
||||
if not days:
|
||||
return BackfillResult(start=end, end=end, days_scanned=0, rows_written=0)
|
||||
|
||||
priced: Final = await _existing_sentinel_keys(prisma_client, start=days[0], end=days[-1])
|
||||
missing: Final = tuple(
|
||||
(day.isoformat(), charge)
|
||||
for day in days
|
||||
for charge in _aggregate_charges(ptu_models, day)
|
||||
if (charge.team_id, charge.model_id, day.isoformat()) not in priced
|
||||
)
|
||||
if not missing:
|
||||
return BackfillResult(start=days[0], end=days[-1], days_scanned=len(days), rows_written=0)
|
||||
|
||||
landed: Final = tuple(
|
||||
[
|
||||
await _upsert_charge_with_retry(prisma_client, charge=charge, date_str=date_str)
|
||||
for date_str, charge in missing
|
||||
]
|
||||
)
|
||||
rows_written: Final = sum(landed)
|
||||
verbose_proxy_logger.info(
|
||||
"PTU backfill for %s to %s: %d unpriced charge(s) found, %d written, %d failed",
|
||||
days[0].isoformat(),
|
||||
days[-1].isoformat(),
|
||||
len(missing),
|
||||
rows_written,
|
||||
len(missing) - rows_written,
|
||||
)
|
||||
return BackfillResult(
|
||||
start=days[0],
|
||||
end=days[-1],
|
||||
days_scanned=len(days),
|
||||
rows_written=rows_written,
|
||||
rows_failed=len(missing) - rows_written,
|
||||
)
|
||||
|
||||
|
||||
async def run_scheduled_ptu_rollup(
|
||||
prisma_client: "PrismaClient",
|
||||
pod_lock_manager: "PodLockManager | None" = None,
|
||||
target_date: date | None = None,
|
||||
alert: Callable[[str], Awaitable[None]] | None = None,
|
||||
) -> RollupResult | None:
|
||||
"""Run the daily rollup under a cross-pod lock so only one proxy reconciles a day.
|
||||
|
||||
Every proxy process schedules this cron, and the read-charge-prune sequence is not
|
||||
atomic: two pods reading different config snapshots can have the loser's prune delete
|
||||
a row the winner just wrote. Returns None when another pod holds the lock, since that
|
||||
pod is doing the work. A deployment without a Redis-backed lock manager runs
|
||||
unguarded, as ``SpendLogCleanup`` does, and so does a run that cannot reach Redis at
|
||||
all: the lock exists to avoid duplicate work, so no lock problem may cost a day.
|
||||
|
||||
The lease is a fixed TTL with no renewal, so a long scan can outlive it. That costs
|
||||
duplicate work rather than correctness: the upserts are idempotent on the sentinel
|
||||
key and the prune reads only the row's own timestamp, so a second pod arriving
|
||||
mid-run cannot corrupt the day.
|
||||
|
||||
Returns None without touching the database when PTU cost attribution is off. Proxy
|
||||
startup already skips scheduling the cron, so this guards the function itself rather
|
||||
than its one caller, and a deployment that never opted in accrues nothing whatever
|
||||
reaches it.
|
||||
"""
|
||||
if not is_ptu_cost_attribution_enabled():
|
||||
return None
|
||||
|
||||
if pod_lock_manager is None or pod_lock_manager.redis_cache is None:
|
||||
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False)
|
||||
|
||||
if not await pod_lock_manager.acquire_lock(cronjob_id=PTU_ROLLUP_JOB_ID, ttl=PTU_ROLLUP_LOCK_TTL_SECONDS):
|
||||
if await _lock_is_held(pod_lock_manager):
|
||||
verbose_proxy_logger.info("PTU rollup: another pod holds the rollup lock, skipping this run")
|
||||
return None
|
||||
# acquire_lock reports contention and a Redis outage the same way, so an
|
||||
# unreachable Redis would otherwise skip the day on every pod at once. The
|
||||
# reconcile is safe to run concurrently, so losing the lock costs duplicate
|
||||
# work; losing the day costs a team's charges
|
||||
verbose_proxy_logger.warning(
|
||||
"PTU rollup: could not take the rollup lock and no other pod holds it, "
|
||||
"running unguarded rather than skipping the day"
|
||||
)
|
||||
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False)
|
||||
|
||||
try:
|
||||
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=True)
|
||||
finally:
|
||||
await pod_lock_manager.release_lock(cronjob_id=PTU_ROLLUP_JOB_ID)
|
||||
|
||||
|
||||
async def _lock_is_held(pod_lock_manager: "PodLockManager") -> bool:
|
||||
"""True only when the rollup lock is readable and someone is holding it.
|
||||
|
||||
A Redis that cannot be read is reported as "not held" so the caller runs the day
|
||||
rather than skipping it; the cost of being wrong here is a duplicate reconcile.
|
||||
"""
|
||||
try:
|
||||
lock_key: Final = pod_lock_manager.get_redis_lock_key(PTU_ROLLUP_JOB_ID)
|
||||
return bool(await pod_lock_manager.redis_cache.async_get_cache(lock_key))
|
||||
except Exception as exc: # noqa: BLE001 # an unreadable lock must not skip the day
|
||||
verbose_proxy_logger.warning("PTU rollup: could not read the rollup lock: %s", exc)
|
||||
return False
|
||||
|
||||
|
||||
async def _run_and_alert(
|
||||
prisma_client: "PrismaClient",
|
||||
*,
|
||||
target_date: date | None,
|
||||
alert: "Callable[[str], Awaitable[None]] | None",
|
||||
may_prune: bool = True,
|
||||
) -> RollupResult:
|
||||
"""Reconcile the day, catch up any days left unpriced, and alert on charges that did not land.
|
||||
|
||||
A charge that exhausts its retries leaves that team showing no PTU cost for the date,
|
||||
and the scheduled job moves on to the next day rather than replaying it. That is a
|
||||
silent underbill unless someone is reading proxy logs, so it is escalated to whatever
|
||||
alerting the deployment has configured.
|
||||
|
||||
The catch-up pass runs only on the scheduled shape, where ``target_date`` is None. An
|
||||
explicit date means reconcile exactly that day, so it stays a single-day operation.
|
||||
Its failure is contained: the day's own result is returned either way.
|
||||
"""
|
||||
result: Final = await run_ptu_flat_cost_rollup(prisma_client, target_date=target_date, may_prune=may_prune)
|
||||
if result.rows_failed:
|
||||
await _deliver_alert(
|
||||
alert,
|
||||
f"PTU flat-cost rollup for {result.day.isoformat()}: {result.rows_failed} of "
|
||||
f"{result.rows_written + result.rows_failed} team charges failed to write. Those teams show no PTU "
|
||||
f"cost for that date until the rollup is rerun for it.",
|
||||
)
|
||||
if target_date is None:
|
||||
await _backfill_and_alert(prisma_client, alert=alert)
|
||||
return result
|
||||
|
||||
|
||||
async def _backfill_and_alert(
|
||||
prisma_client: "PrismaClient",
|
||||
*,
|
||||
alert: "Callable[[str], Awaitable[None]] | None",
|
||||
) -> None:
|
||||
"""Catch up unpriced PTU days, alerting on charges that did not land.
|
||||
|
||||
Never raises: the day's own rollup has already run and its result must reach the
|
||||
caller whatever the catch-up pass does.
|
||||
"""
|
||||
try:
|
||||
backfill: Final = await run_ptu_flat_cost_backfill(prisma_client)
|
||||
except Exception as exc: # noqa: BLE001 # the catch-up pass must not fail the day's rollup
|
||||
verbose_proxy_logger.error("PTU backfill: catch-up pass failed, the day's rollup still stands: %s", exc)
|
||||
return
|
||||
if backfill.rows_failed:
|
||||
await _deliver_alert(
|
||||
alert,
|
||||
f"PTU flat-cost backfill for {backfill.start.isoformat()} to {backfill.end.isoformat()}: "
|
||||
f"{backfill.rows_failed} of {backfill.rows_written + backfill.rows_failed} previously unpriced charges "
|
||||
f"failed to write. Those days stay unpriced until a later run picks them up.",
|
||||
)
|
||||
|
||||
|
||||
async def _deliver_alert(alert: "Callable[[str], Awaitable[None]] | None", message: str) -> None:
|
||||
"""Send an operator alert when one is configured, swallowing a broken channel."""
|
||||
if alert is None:
|
||||
return
|
||||
try:
|
||||
await alert(message)
|
||||
except Exception as exc: # noqa: BLE001 # a broken alert channel must not fail the rollup
|
||||
verbose_proxy_logger.error("PTU rollup: could not deliver the failed-charge alert: %s", exc)
|
||||
|
||||
|
||||
async def _prune_unrefreshed_sentinel_rows(
|
||||
prisma_client: "PrismaClient",
|
||||
*,
|
||||
date_str: str,
|
||||
run_started: datetime,
|
||||
) -> None:
|
||||
"""Delete the day's PTU sentinel rows this run did not refresh.
|
||||
|
||||
Every charge the run wrote bumps ``updated_at`` past ``run_started``, so anything
|
||||
left below that mark is a (team, model) the current config no longer prices. The mark
|
||||
is pulled back by ``PTU_PRUNE_SKEW_GRACE_SECONDS`` because the two timestamps come
|
||||
from different hosts: a stale row is hours old, a concurrently written one is seconds
|
||||
old, and the grace separates them without waiting on clocks agreeing. The
|
||||
predicate reads only the row, never the caller's config snapshot, which is what
|
||||
makes it safe to run twice, out of order, or beside another pod: a row written
|
||||
after this run began is out of reach of its delete. Mirrors the retention predicate
|
||||
``SpendLogCleanup`` deletes by."""
|
||||
cutoff: Final = run_started - timedelta(seconds=PTU_PRUNE_SKEW_GRACE_SECONDS)
|
||||
await prisma_client.db.litellm_dailyteamspend.delete_many(
|
||||
where={ # mutable-ok: prisma delete filter
|
||||
"date": date_str,
|
||||
"api_key": PTU_SENTINEL_API_KEY,
|
||||
"updated_at": {"lt": cutoff}, # mutable-ok: prisma comparison filter
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
__all__ = (
|
||||
"PTU_ROLLUP_JOB_ID",
|
||||
"PTU_SENTINEL_API_KEY",
|
||||
"BackfillResult",
|
||||
"PTUModel",
|
||||
"RollupResult",
|
||||
"run_ptu_flat_cost_backfill",
|
||||
"run_ptu_flat_cost_rollup",
|
||||
"run_scheduled_ptu_rollup",
|
||||
)
|
||||
|
|
@ -13,7 +13,7 @@ from typing import TYPE_CHECKING, Final, NamedTuple
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import _get_cost_per_unit, generic_cost_per_token
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router
|
||||
|
|
@ -26,29 +26,42 @@ class SavingsSpend(NamedTuple):
|
|||
autorouter: float = 0.0
|
||||
|
||||
|
||||
def _input_and_cache_read_cost(model: str | None, custom_llm_provider: str | None) -> tuple[float, float]:
|
||||
def _input_cache_read_and_write_cost(info: ModelInfo | None) -> tuple[float, float, float]:
|
||||
"""
|
||||
Return ``(input_cost_per_token, cache_read_cost_per_token)`` for a model.
|
||||
Return ``(input_cost, cache_read_cost, cache_write_cost)`` per token.
|
||||
|
||||
Falls open to ``(0.0, 0.0)`` when the model is unknown so savings degrade to
|
||||
zero rather than raising inside the spend writer. When a model has no
|
||||
separate cache-read price the cache-read cost mirrors the input cost, which
|
||||
yields zero caching savings.
|
||||
``info`` is whatever pricing the caller resolved -- deployment rates when the
|
||||
request came through a router deployment, public rates otherwise -- so a
|
||||
negotiated price is honoured here rather than silently replaced by the list rate.
|
||||
``None`` falls open to ``(0.0, 0.0, 0.0)`` so savings degrade to zero rather than
|
||||
raising inside the spend writer.
|
||||
|
||||
Prices are read through ``_get_cost_per_unit``, the same accessor the cost
|
||||
calculator uses, which coerces the string prices a ``config.yaml`` can produce
|
||||
(``"3e-7"``) and resolves service-tier suffixes.
|
||||
|
||||
An absent cache price mirrors the input cost, which yields a zero discount on the
|
||||
read leg and a zero premium on the write leg. Mirroring rather than taking
|
||||
``_get_cost_per_unit``'s 0.0 default is load-bearing on the write leg: a zero write
|
||||
price would make the premium ``0 - input_cost``, turning a model that simply has no
|
||||
write pricing into a spurious extra saving.
|
||||
|
||||
The two legs then differ on an explicit ``0.0``, and the asymmetry is deliberate. A
|
||||
free cache *write* does not exist -- entries carrying a literal zero (``deepseek-chat``
|
||||
does) mean "no separate price", so a falsy write price also mirrors input. A free
|
||||
cache *read* is real: 15 models charge for input and serve reads for nothing, which
|
||||
is the largest discount available, so the read leg keeps its literal zero.
|
||||
"""
|
||||
if not model:
|
||||
return 0.0, 0.0
|
||||
try:
|
||||
info: Final = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
|
||||
except Exception as e: # noqa: BLE001 # get_model_info raises bare Exception for unmapped models; degrade to zero savings
|
||||
verbose_proxy_logger.debug(
|
||||
"savings: no model info for provider=%s model=%s (%s)", custom_llm_provider, model, e
|
||||
)
|
||||
return 0.0, 0.0
|
||||
input_cost: Final = float(info.get("input_cost_per_token") or 0.0)
|
||||
cache_read_cost: Final = info.get("cache_read_input_token_cost")
|
||||
if cache_read_cost is None:
|
||||
return input_cost, input_cost
|
||||
return input_cost, float(cache_read_cost)
|
||||
if info is None:
|
||||
return 0.0, 0.0, 0.0
|
||||
input_cost: Final = _get_cost_per_unit(info, "input_cost_per_token") or 0.0
|
||||
cache_read_cost: Final = _get_cost_per_unit(info, "cache_read_input_token_cost", default_value=None)
|
||||
cache_write_cost: Final = _get_cost_per_unit(info, "cache_creation_input_token_cost", default_value=None)
|
||||
return (
|
||||
input_cost,
|
||||
input_cost if cache_read_cost is None else cache_read_cost,
|
||||
cache_write_cost if cache_write_cost else input_cost,
|
||||
)
|
||||
|
||||
|
||||
class _ModelIdentity(NamedTuple):
|
||||
|
|
@ -434,10 +447,28 @@ def compute_savings_spend(
|
|||
Dollar savings for one request, split by optimization driver.
|
||||
|
||||
Compression savings price the tokens compression removed at the model's
|
||||
input rate. Prompt-caching savings price the cache-read tokens at the
|
||||
difference between the input rate and the discounted cache-read rate; the
|
||||
read count is derived here from ``usage_object`` so no caller can hand in a
|
||||
count that disagrees with the usage record. Auto-router savings compare the
|
||||
input rate. Prompt-caching savings are NET: the cache-read discount minus the
|
||||
premium paid to write those entries, both derived here from ``usage_object`` so no
|
||||
caller can hand in a count that disagrees with the usage record.
|
||||
|
||||
The net form follows from what the request would have cost with caching off. The
|
||||
provider reports ``prompt_tokens`` as the inclusive total of three disjoint
|
||||
partitions (uncached text, cache reads, cache writes), so an uncached counterfactual
|
||||
bills every one of those tokens at the flat input rate::
|
||||
|
||||
would_have_cost = (text + reads + writes) * input
|
||||
actually_cost = text * input + reads * read_rate + writes * write_rate
|
||||
savings = reads * (input - read_rate) - writes * (write_rate - input)
|
||||
|
||||
So the write leg subtracts the write PREMIUM, not the whole write cost: those tokens
|
||||
had to be sent either way, and the counterfactual already pays the input rate for
|
||||
them. The premium stays signed, because a handful of models price writes below their
|
||||
input rate and there the write is a genuine extra saving.
|
||||
|
||||
A request that only writes cache and gets no hits therefore reports negative savings,
|
||||
which is accurate: it really did cost more than the uncached call would have. The
|
||||
daily rollup increments arithmetically, so those rows offset positive ones in the
|
||||
same bucket. Auto-router savings compare the
|
||||
served ``model`` against the counterfactual baseline the router recorded on
|
||||
its ``routing_decision``, and are zero unless the two differ. That record
|
||||
also says whether the conversation was already underway, which is what tells
|
||||
|
|
@ -454,10 +485,21 @@ def compute_savings_spend(
|
|||
the same way; that is pre-existing behaviour on two shipped drivers rather than
|
||||
something introduced here, and moving those numbers is its own change.
|
||||
"""
|
||||
input_cost, cache_read_cost = _input_and_cache_read_cost(model, custom_llm_provider)
|
||||
# Deployment rates when the request came through one, public rates otherwise --
|
||||
# `_effective_model_info` merges a deployment's configured prices over the built-in
|
||||
# map, so a negotiated price is not silently replaced by the list rate.
|
||||
router_instance: Router | None = llm_router() if llm_router else None
|
||||
identity: Final = _resolve_model(model, custom_llm_provider)
|
||||
pricing: Final = _effective_model_info(router_instance, model_id, model or "") or (
|
||||
_model_info(identity) if identity else None
|
||||
)
|
||||
input_cost, cache_read_cost, cache_write_cost = _input_cache_read_and_write_cost(pricing)
|
||||
compression: Final = max(compression_saved_tokens, 0) * input_cost
|
||||
cache_read_input_tokens: Final = extract_cache_read_tokens(usage_object)
|
||||
prompt_caching: Final = max(cache_read_input_tokens, 0) * max(input_cost - cache_read_cost, 0.0)
|
||||
cache_creation_input_tokens: Final = extract_cache_creation_tokens(usage_object)
|
||||
read_discount: Final = max(cache_read_input_tokens, 0) * max(input_cost - cache_read_cost, 0.0)
|
||||
write_premium: Final = max(cache_creation_input_tokens, 0) * (cache_write_cost - input_cost)
|
||||
prompt_caching: Final = read_discount - write_premium
|
||||
|
||||
usage: Final = _usage_from_spend_log(usage_object)
|
||||
if usage is None or not model:
|
||||
|
|
@ -480,9 +522,7 @@ def compute_savings_spend(
|
|||
# Absent means the router never recorded a shape, which is the conservative
|
||||
# reading: charge the cache write rather than claim a first turn's saving.
|
||||
conversation_continuing=decision.get("conversation_continuing") is not False,
|
||||
selected_info=_effective_model_info(
|
||||
(router_instance := llm_router() if llm_router else None), model_id, model or ""
|
||||
),
|
||||
selected_info=_effective_model_info(router_instance, model_id, model or ""),
|
||||
baseline_info=_effective_model_info(router_instance, baseline_id, baseline_model or ""),
|
||||
cost_breakdown=cost_breakdown,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ from litellm.proxy.config_resolvers.sso import (
|
|||
SSO_SECRET_FIELDS,
|
||||
resolve_sso_config,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled
|
||||
from litellm.proxy.utils import invalidate_config_param
|
||||
from litellm.repositories.config_repository import ConfigRepository
|
||||
from litellm.repositories.organization_repository import OrganizationRepository
|
||||
|
|
@ -307,6 +308,27 @@ ALLOWED_UI_SETTINGS_FIELDS: Final = {
|
|||
"enable_chat_ui",
|
||||
}
|
||||
|
||||
ENABLE_PTU_COST_ATTRIBUTION_UI_SETTING: Final = "enable_ptu_cost_attribution"
|
||||
|
||||
# UI settings derived from the deployment environment. Deliberately kept out of
|
||||
# ALLOWED_UI_SETTINGS_FIELDS: they are read-only, never persisted, and PATCH
|
||||
# rejects them so an admin cannot flip an env-gated feature at runtime.
|
||||
_DERIVED_UI_SETTINGS_FIELDS: Final[frozenset[str]] = frozenset({ENABLE_PTU_COST_ATTRIBUTION_UI_SETTING})
|
||||
|
||||
|
||||
def _derived_ui_setting_value(key: str) -> object:
|
||||
"""The environment-derived value GET reports for ``key``.
|
||||
|
||||
PATCH compares against this rather than rejecting the key outright, so the body GET
|
||||
hands back is still a valid PATCH body. Rejecting on presence broke read-modify-write:
|
||||
a client that edited one setting and sent the rest back unchanged got a 400 and lost
|
||||
the edit it actually wanted.
|
||||
"""
|
||||
if key == ENABLE_PTU_COST_ATTRIBUTION_UI_SETTING:
|
||||
return is_ptu_cost_attribution_enabled()
|
||||
return None
|
||||
|
||||
|
||||
# Flags that must be synced from the persisted UISettings into
|
||||
# general_settings at runtime (on both read and write).
|
||||
_RUNTIME_GENERAL_SETTINGS_FLAGS: Final = [
|
||||
|
|
@ -1345,21 +1367,15 @@ async def get_ui_settings():
|
|||
detail={"error": "Database not connected. Please connect a database."},
|
||||
)
|
||||
|
||||
ui_settings: Mapping[str, JsonValue] = {}
|
||||
|
||||
db_record: Final = await _ui_settings_db(UISettingsRepository(prisma_client)).find_unique(
|
||||
where={"id": "ui_settings"}
|
||||
)
|
||||
|
||||
if db_record and db_record.ui_settings:
|
||||
ui_settings_json: Final = db_record.ui_settings
|
||||
if isinstance(ui_settings_json, str):
|
||||
ui_settings = json.loads(ui_settings_json)
|
||||
else:
|
||||
ui_settings = dict(ui_settings_json)
|
||||
stored: Final = (db_record.ui_settings if db_record else None) or "{}"
|
||||
parsed: Final = json.loads(stored) if isinstance(stored, str) else stored
|
||||
|
||||
# Sanitize any unexpected keys from persisted config before returning
|
||||
ui_settings = {k: v for k, v in ui_settings.items() if k in ALLOWED_UI_SETTINGS_FIELDS}
|
||||
ui_settings: Final = {k: v for k, v in parsed.items() if k in ALLOWED_UI_SETTINGS_FIELDS}
|
||||
|
||||
# Sync runtime flags into general_settings so the proxy picks them up
|
||||
# at runtime (covers server restart scenarios).
|
||||
|
|
@ -1377,11 +1393,18 @@ async def get_ui_settings():
|
|||
# Build config-like object for schema helper
|
||||
config: Final[dict[str, object]] = {"litellm_settings": {"ui_settings": ui_settings}}
|
||||
|
||||
return await _get_settings_with_schema(
|
||||
settings: Final = await _get_settings_with_schema(
|
||||
settings_key="ui_settings",
|
||||
settings_class=_get_effective_ui_settings_class(),
|
||||
config=config,
|
||||
)
|
||||
return UISettingsResponse(
|
||||
values={
|
||||
**settings["values"],
|
||||
ENABLE_PTU_COST_ATTRIBUTION_UI_SETTING: is_ptu_cost_attribution_enabled(),
|
||||
},
|
||||
field_schema=settings["field_schema"],
|
||||
)
|
||||
|
||||
|
||||
@router.patch(
|
||||
|
|
@ -1418,6 +1441,20 @@ async def update_ui_settings(
|
|||
detail={"error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature."},
|
||||
)
|
||||
|
||||
conflicting_keys: Final = sorted(
|
||||
key
|
||||
for key, value in settings_body.items()
|
||||
if key in _DERIVED_UI_SETTINGS_FIELDS and value != _derived_ui_setting_value(key)
|
||||
)
|
||||
if conflicting_keys:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
f"Setting(s) {conflicting_keys} are derived from the deployment environment "
|
||||
"and cannot be changed from the UI."
|
||||
),
|
||||
)
|
||||
|
||||
# Validate against the same effective class GET advertises, so
|
||||
# enterprise-registered fields are typed consistently on both sides.
|
||||
effective_cls: Final = _get_effective_ui_settings_class()
|
||||
|
|
|
|||
|
|
@ -544,6 +544,11 @@ class ProxyLogging:
|
|||
for hook in PROXY_HOOKS:
|
||||
proxy_hook = get_proxy_hook(hook)
|
||||
expected_args = inspect.getfullargspec(proxy_hook).args
|
||||
if "prisma_client" in expected_args and prisma_client is None:
|
||||
verbose_proxy_logger.debug(
|
||||
"Skipping proxy hook %s: it requires a database and no prisma client is configured", hook
|
||||
)
|
||||
continue
|
||||
passed_in_args: dict[str, Any] = {}
|
||||
if "internal_usage_cache" in expected_args:
|
||||
passed_in_args["internal_usage_cache"] = self.internal_usage_cache
|
||||
|
|
@ -3481,13 +3486,15 @@ class PrismaClient:
|
|||
r.expires = r.expires.isoformat()
|
||||
elif query_type == "find_all" and expires is not None and reset_at is not None:
|
||||
response = await VerificationTokenRepository(self).table.find_many(
|
||||
take=limit,
|
||||
where={
|
||||
"OR": [
|
||||
{"expires": None},
|
||||
{"expires": {"gt": expires}},
|
||||
],
|
||||
"budget_reset_at": {"lt": reset_at},
|
||||
}
|
||||
"NOT": {"budget_duration": None},
|
||||
},
|
||||
)
|
||||
if response is not None and len(response) > 0:
|
||||
for r in response:
|
||||
|
|
@ -3537,6 +3544,7 @@ class PrismaClient:
|
|||
response = await UserRepository(self).table.find_many(where=key_val)
|
||||
elif query_type == "find_all" and reset_at is not None:
|
||||
response = await UserRepository(self).table.find_many(
|
||||
take=limit,
|
||||
where={
|
||||
# A user seeded from default_internal_user_params
|
||||
# (or created via /user/new without an explicit
|
||||
|
|
@ -3547,16 +3555,12 @@ class PrismaClient:
|
|||
# of the row, silently exceeding max_budget. Treat a
|
||||
# NULL budget_reset_at with a non-NULL budget_duration
|
||||
# as due, matching the budget-table query below.
|
||||
"NOT": {"budget_duration": None},
|
||||
"OR": [
|
||||
{
|
||||
"AND": [
|
||||
{"budget_reset_at": None},
|
||||
{"NOT": {"budget_duration": None}},
|
||||
]
|
||||
},
|
||||
{"budget_reset_at": None},
|
||||
{"budget_reset_at": {"lt": reset_at}},
|
||||
],
|
||||
}
|
||||
},
|
||||
)
|
||||
elif query_type == "find_all" and user_id_list is not None:
|
||||
response = await UserRepository(self).table.find_many(where={"user_id": {"in": user_id_list}})
|
||||
|
|
@ -3612,17 +3616,14 @@ class PrismaClient:
|
|||
elif table_name == "budget" and reset_at is not None:
|
||||
if query_type == "find_all":
|
||||
response = await BudgetRepository(self).table.find_many(
|
||||
take=limit,
|
||||
where={
|
||||
"NOT": {"budget_duration": None},
|
||||
"OR": [
|
||||
{
|
||||
"AND": [
|
||||
{"budget_reset_at": None},
|
||||
{"NOT": {"budget_duration": None}},
|
||||
]
|
||||
},
|
||||
{"budget_reset_at": None},
|
||||
{"budget_reset_at": {"lt": reset_at}},
|
||||
]
|
||||
}
|
||||
],
|
||||
},
|
||||
)
|
||||
return response
|
||||
|
||||
|
|
@ -3640,20 +3641,17 @@ class PrismaClient:
|
|||
)
|
||||
elif query_type == "find_all" and reset_at is not None:
|
||||
response = await TeamRepository(self).table.find_many(
|
||||
take=limit,
|
||||
where={
|
||||
# Same NULL budget_reset_at gap as the user query
|
||||
# above: a team with a budget_duration but no
|
||||
# initialized budget_reset_at would never be reset.
|
||||
"NOT": {"budget_duration": None},
|
||||
"OR": [
|
||||
{
|
||||
"AND": [
|
||||
{"budget_reset_at": None},
|
||||
{"NOT": {"budget_duration": None}},
|
||||
]
|
||||
},
|
||||
{"budget_reset_at": None},
|
||||
{"budget_reset_at": {"lt": reset_at}},
|
||||
],
|
||||
}
|
||||
},
|
||||
)
|
||||
elif query_type == "find_all" and user_id is not None:
|
||||
response = await TeamRepository(self).table.find_many(
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import Any, Final
|
||||
from typing import Annotated, Any, Final
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response
|
||||
|
||||
|
|
@ -18,7 +18,8 @@ from litellm.proxy.vector_store_endpoints.utils import (
|
|||
get_litellm_managed_vector_store,
|
||||
)
|
||||
from litellm.repositories.table_repositories import ManagedVectorStoreIndexRepository
|
||||
from litellm.types.vector_stores import IndexCreateRequest
|
||||
from litellm.types.vector_stores import IndexCreateRequest, IndexListResponse
|
||||
from litellm.vector_stores.vector_store_registry import VectorStoreIndexRegistry
|
||||
|
||||
router: Final = APIRouter()
|
||||
########################################################
|
||||
|
|
@ -549,14 +550,15 @@ async def index_create(
|
|||
Create an index. Just writes the index to the database.
|
||||
|
||||
```bash
|
||||
curl -L -X POST 'http://0.0.0.0:4000/indexes/create' \
|
||||
curl -L -X POST 'http://0.0.0.0:4000/v1/indexes' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'LiteLLM-Beta: indexes_beta=v1' \
|
||||
-d '{
|
||||
-d '{
|
||||
"index_name": "dall-e-3",
|
||||
"vector_store_index": "real-index-name",
|
||||
"vector_store_name": "azure-ai-search"
|
||||
"litellm_params": {
|
||||
"vector_store_index": "real-index-name",
|
||||
"vector_store_name": "azure-ai-search"
|
||||
}
|
||||
}'
|
||||
```
|
||||
"""
|
||||
|
|
@ -592,3 +594,36 @@ async def index_create(
|
|||
new_index = await ManagedVectorStoreIndexRepository(prisma_client).table.create(data=jsonify_object(index_data))
|
||||
|
||||
return new_index.model_dump()
|
||||
|
||||
|
||||
@router.get(
|
||||
"/v1/indexes",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=IndexListResponse,
|
||||
)
|
||||
async def index_list(
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
) -> IndexListResponse:
|
||||
"""
|
||||
List all vector store indexes. Proxy admin only.
|
||||
|
||||
```bash
|
||||
curl -L -X GET 'http://0.0.0.0:4000/v1/indexes' \
|
||||
-H 'Authorization: Bearer sk-1234'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
assert_proxy_admin_for_vector_store_index_management(
|
||||
user_api_key_dict,
|
||||
operation="list",
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail=CommonProxyErrors.db_not_connected_error.value,
|
||||
)
|
||||
|
||||
indexes: Final = await VectorStoreIndexRegistry._get_vector_store_indexes_from_db(prisma_client)
|
||||
return IndexListResponse(data=indexes)
|
||||
|
|
|
|||
|
|
@ -41,7 +41,7 @@ def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
|
|||
def assert_proxy_admin_for_vector_store_index_management(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
*,
|
||||
operation: Literal["create", "delete", "update"] = "create",
|
||||
operation: Literal["create", "delete", "update", "list"] = "create",
|
||||
) -> None:
|
||||
"""Raise 403 unless the caller is a proxy admin."""
|
||||
if _is_proxy_admin(user_api_key_dict):
|
||||
|
|
|
|||
|
|
@ -17,7 +17,8 @@ from __future__ import annotations
|
|||
|
||||
import hashlib
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -35,10 +36,32 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import httpx
|
||||
|
||||
from litellm import Router
|
||||
from litellm.types.rag import RAGIngestOptions
|
||||
|
||||
|
||||
class S3VectorDataPayload(TypedDict):
|
||||
float32: Sequence[float]
|
||||
|
||||
|
||||
class S3VectorEntry(TypedDict):
|
||||
key: str
|
||||
data: S3VectorDataPayload
|
||||
metadata: Mapping[str, str]
|
||||
|
||||
|
||||
class S3VectorsQueryMatch(TypedDict, total=False):
|
||||
key: str
|
||||
distance: float
|
||||
metadata: Mapping[str, str]
|
||||
|
||||
|
||||
class S3VectorsQueryResponse(TypedDict, total=False):
|
||||
vectors: Sequence[S3VectorsQueryMatch]
|
||||
|
||||
|
||||
class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
||||
"""
|
||||
S3 Vectors RAG ingestion using httpx + AWS SigV4 signing.
|
||||
|
|
@ -66,10 +89,10 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
BaseAWSLLM.__init__(self)
|
||||
|
||||
# Extract config
|
||||
self.vector_bucket_name = self.vector_store_config["vector_bucket_name"]
|
||||
self.index_name = self.vector_store_config.get("index_name")
|
||||
self.distance_metric = self.vector_store_config.get("distance_metric", S3_VECTORS_DEFAULT_DISTANCE_METRIC)
|
||||
self.non_filterable_metadata_keys = self.vector_store_config.get(
|
||||
self.vector_bucket_name: str = self.vector_store_config["vector_bucket_name"]
|
||||
self.index_name: str | None = self.vector_store_config.get("index_name")
|
||||
self.distance_metric: str = self.vector_store_config.get("distance_metric", S3_VECTORS_DEFAULT_DISTANCE_METRIC)
|
||||
self.non_filterable_metadata_keys: Sequence[str] = self.vector_store_config.get(
|
||||
"non_filterable_metadata_keys",
|
||||
S3_VECTORS_DEFAULT_NON_FILTERABLE_METADATA_KEYS,
|
||||
)
|
||||
|
|
@ -78,7 +101,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
self.dimension = self._get_dimension_from_config()
|
||||
|
||||
# Get AWS region using BaseAWSLLM method
|
||||
_aws_region: Final = self.vector_store_config.get("aws_region_name")
|
||||
_aws_region: Final[str | None] = self.vector_store_config.get("aws_region_name")
|
||||
self.aws_region_name = self.get_aws_region_name_for_non_llm_api_calls(
|
||||
aws_region_name=str(_aws_region) if _aws_region else None
|
||||
)
|
||||
|
|
@ -135,7 +158,8 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
Returns None if dimension should be auto-detected.
|
||||
"""
|
||||
if "dimension" in self.vector_store_config:
|
||||
return int(self.vector_store_config["dimension"])
|
||||
configured_dimension: Final[int] = self.vector_store_config["dimension"]
|
||||
return int(configured_dimension)
|
||||
return None
|
||||
|
||||
async def _ensure_config_initialized(self):
|
||||
|
|
@ -258,7 +282,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
get_body: Final = safe_dumps({"vectorBucketName": self.vector_bucket_name})
|
||||
|
||||
try:
|
||||
response = await self._sign_and_execute_request("POST", get_url, data=get_body)
|
||||
response: httpx.Response = await self._sign_and_execute_request("POST", get_url, data=get_body)
|
||||
if response.status_code == 200:
|
||||
verbose_logger.debug("Vector bucket %s exists", self.vector_bucket_name)
|
||||
return
|
||||
|
|
@ -294,7 +318,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
get_body: Final = safe_dumps({"vectorBucketName": self.vector_bucket_name, "indexName": self.index_name})
|
||||
|
||||
try:
|
||||
response = await self._sign_and_execute_request("POST", get_url, data=get_body)
|
||||
response: httpx.Response = await self._sign_and_execute_request("POST", get_url, data=get_body)
|
||||
if response.status_code == 200:
|
||||
verbose_logger.debug("Vector index %s exists", self.index_name)
|
||||
return
|
||||
|
|
@ -311,7 +335,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
)
|
||||
|
||||
# Prepare index configuration per AWS API docs
|
||||
index_config: Final = {
|
||||
index_config: Final[dict[str, object]] = {
|
||||
"vectorBucketName": self.vector_bucket_name,
|
||||
"indexName": self.index_name,
|
||||
"dataType": "float32",
|
||||
|
|
@ -336,7 +360,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
verbose_logger.exception("Error creating vector index: %s", e)
|
||||
raise
|
||||
|
||||
async def _put_vectors(self, vectors: list[dict[str, Any]]):
|
||||
async def _put_vectors(self, vectors: Sequence[S3VectorEntry]):
|
||||
"""
|
||||
Call PutVectors API to store vectors in S3 Vectors.
|
||||
|
||||
|
|
@ -355,7 +379,9 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
}
|
||||
|
||||
try:
|
||||
response: Final = await self._sign_and_execute_request("POST", url, data=safe_dumps(request_body))
|
||||
response: Final[httpx.Response] = await self._sign_and_execute_request(
|
||||
"POST", url, data=safe_dumps(request_body)
|
||||
)
|
||||
|
||||
if response.status_code in (200, 201):
|
||||
verbose_logger.info("Successfully stored %s vectors in index %s", len(vectors), self.index_name)
|
||||
|
|
@ -442,24 +468,18 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
raise ValueError(error_msg)
|
||||
|
||||
# Prepare vectors for PutVectors API
|
||||
vectors: Final = []
|
||||
for i, (chunk, embedding) in enumerate(zip(chunks, embeddings)):
|
||||
# Build metadata dict
|
||||
metadata: dict[str, str] = {
|
||||
"source_text": chunk, # Non-filterable (for reference)
|
||||
"chunk_index": str(i), # Filterable
|
||||
}
|
||||
|
||||
if filename:
|
||||
metadata["filename"] = filename # Filterable
|
||||
|
||||
vector_obj = {
|
||||
"key": f"{filename}_{i}" if filename else f"chunk_{i}",
|
||||
"data": {"float32": embedding},
|
||||
"metadata": metadata,
|
||||
}
|
||||
|
||||
vectors.append(vector_obj)
|
||||
vectors: Final = [
|
||||
S3VectorEntry(
|
||||
key=f"{filename}_{i}" if filename else f"chunk_{i}",
|
||||
data=S3VectorDataPayload(float32=embedding),
|
||||
metadata=(
|
||||
{"source_text": chunk, "chunk_index": str(i), "filename": filename}
|
||||
if filename
|
||||
else {"source_text": chunk, "chunk_index": str(i)}
|
||||
),
|
||||
)
|
||||
for i, (chunk, embedding) in enumerate(zip(chunks, embeddings))
|
||||
]
|
||||
|
||||
# Call PutVectors API
|
||||
await self._put_vectors(vectors)
|
||||
|
|
@ -468,7 +488,9 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
vector_store_id: Final = f"{self.vector_bucket_name}:{self.index_name}"
|
||||
return vector_store_id, filename
|
||||
|
||||
async def query_vector_store(self, vector_store_id: str, query: str, top_k: int = 5) -> dict[str, Any] | None:
|
||||
async def query_vector_store(
|
||||
self, vector_store_id: str, query: str, top_k: int = 5
|
||||
) -> S3VectorsQueryResponse | None:
|
||||
"""
|
||||
Query S3 Vectors using QueryVectors API.
|
||||
|
||||
|
|
@ -489,7 +511,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
embedding_model: Final = self.embedding_config.get("model", "text-embedding-3-small")
|
||||
|
||||
response = await litellm.aembedding(model=embedding_model, input=[query])
|
||||
query_embedding: Final = response.data[0]["embedding"]
|
||||
query_embedding: Final[Sequence[float]] = response.data[0]["embedding"]
|
||||
|
||||
# Call QueryVectors API
|
||||
url: Final = f"https://s3vectors.{self.aws_region_name}.api.aws/QueryVectors"
|
||||
|
|
@ -504,15 +526,18 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
}
|
||||
|
||||
try:
|
||||
response = await self._sign_and_execute_request("POST", url, data=safe_dumps(request_body))
|
||||
query_response: Final[httpx.Response] = await self._sign_and_execute_request(
|
||||
"POST", url, data=safe_dumps(request_body)
|
||||
)
|
||||
|
||||
if response.status_code == 200:
|
||||
results: Final = response.json()
|
||||
if query_response.status_code == 200:
|
||||
results: Final[S3VectorsQueryResponse] = query_response.json()
|
||||
matches: Final = results.get("vectors")
|
||||
verbose_logger.debug("Query returned %s results", len(results.get("vectors", [])))
|
||||
|
||||
# Check if query terms appear in results
|
||||
if results.get("vectors"):
|
||||
for result in results["vectors"]:
|
||||
if matches:
|
||||
for result in matches:
|
||||
metadata = result.get("metadata", {})
|
||||
source_text = metadata.get("source_text", "")
|
||||
if query.lower() in source_text.lower():
|
||||
|
|
@ -521,7 +546,9 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM):
|
|||
# Return results even if exact match not found
|
||||
return results
|
||||
else:
|
||||
verbose_logger.error("QueryVectors failed with status %s: %s", response.status_code, response.text)
|
||||
verbose_logger.error(
|
||||
"QueryVectors failed with status %s: %s", query_response.status_code, query_response.text
|
||||
)
|
||||
return None
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error querying vectors: %s", e)
|
||||
|
|
|
|||
|
|
@ -70,10 +70,14 @@ from litellm.repositories.table_repositories import (
|
|||
)
|
||||
from litellm.repositories.team_repository import TeamRepository
|
||||
from litellm.repositories.unit_of_work import (
|
||||
BudgetCascadeUnitOfWork,
|
||||
BudgetWindowWrites,
|
||||
KeySpendResetWrites,
|
||||
LinkedSpendResetWrites,
|
||||
SpendResetUnitOfWork,
|
||||
TeamSpendResetWrites,
|
||||
UserSpendResetWrites,
|
||||
budget_cascade_unit_of_work,
|
||||
spend_reset_unit_of_work,
|
||||
)
|
||||
from litellm.repositories.user_repository import UserRepository
|
||||
|
|
@ -88,7 +92,9 @@ __all__ = [
|
|||
"AgentsRepository",
|
||||
"AuditLogRepository",
|
||||
"BatchTable",
|
||||
"BudgetCascadeUnitOfWork",
|
||||
"BudgetRepository",
|
||||
"BudgetWindowWrites",
|
||||
"CacheConfigRepository",
|
||||
"ClaudeCodePluginRepository",
|
||||
"ConfigOverridesRepository",
|
||||
|
|
@ -107,6 +113,7 @@ __all__ = [
|
|||
"InvitationLinkRepository",
|
||||
"JWTKeyMappingRepository",
|
||||
"KeySpendResetWrites",
|
||||
"LinkedSpendResetWrites",
|
||||
"MCPServerRepository",
|
||||
"MCPToolsetRepository",
|
||||
"MCPUserCredentialsRepository",
|
||||
|
|
@ -149,5 +156,6 @@ __all__ = [
|
|||
"WorkflowEventRepository",
|
||||
"WorkflowMessageRepository",
|
||||
"WorkflowRunRepository",
|
||||
"budget_cascade_unit_of_work",
|
||||
"spend_reset_unit_of_work",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -29,6 +29,8 @@ class SpendLinkedTable(Protocol[RowT_co]):
|
|||
class BatchTable(Protocol):
|
||||
def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ...
|
||||
|
||||
def update_many(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ...
|
||||
|
||||
|
||||
class PrismaBatch(Protocol):
|
||||
@property
|
||||
|
|
@ -40,4 +42,19 @@ class PrismaBatch(Protocol):
|
|||
@property
|
||||
def litellm_teamtable(self) -> BatchTable: ...
|
||||
|
||||
@property
|
||||
def litellm_budgettable(self) -> BatchTable: ...
|
||||
|
||||
@property
|
||||
def litellm_teammembership(self) -> BatchTable: ...
|
||||
|
||||
@property
|
||||
def litellm_organizationtable(self) -> BatchTable: ...
|
||||
|
||||
@property
|
||||
def litellm_tagtable(self) -> BatchTable: ...
|
||||
|
||||
@property
|
||||
def litellm_endusertable(self) -> BatchTable: ...
|
||||
|
||||
async def commit(self) -> None: ...
|
||||
|
|
|
|||
|
|
@ -1,17 +1,21 @@
|
|||
"""
|
||||
Unit of work over a single Prisma batch.
|
||||
Units of work over a single Prisma batch.
|
||||
|
||||
``spend_reset_unit_of_work`` opens one ``db.batch_()`` and binds a typed write
|
||||
Each context manager here opens one ``db.batch_()`` and binds a typed write
|
||||
repository per table to it, so every update queued through the yielded object
|
||||
lands in the same transaction. The batch commits when the block exits cleanly
|
||||
and is abandoned, writing nothing, when the block raises.
|
||||
|
||||
Each write repository queues narrow ``{spend, budget_reset_at}`` updates
|
||||
``spend_reset_unit_of_work`` covers the per-row key/user/team resets;
|
||||
``budget_cascade_unit_of_work`` covers a budget tier's reset, where the
|
||||
dependent spend and the tier's next window have to move together.
|
||||
|
||||
Each write repository queues narrow ``{spend}`` / ``{budget_reset_at}`` updates
|
||||
instead of full-model writes, which trip ``prisma.errors.DataError`` on rows
|
||||
carrying fields the update input type rejects (see #27730).
|
||||
"""
|
||||
|
||||
from collections.abc import AsyncGenerator, Callable
|
||||
from collections.abc import AsyncGenerator, Callable, Mapping
|
||||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
|
|
@ -43,6 +47,24 @@ class TeamSpendResetWrites:
|
|||
self.table.update(where={"team_id": team_id}, data={"spend": 0, "budget_reset_at": budget_reset_at})
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class LinkedSpendResetWrites:
|
||||
table: BatchTable
|
||||
|
||||
def queue_spend_zero(self, where: Mapping[str, object]) -> None:
|
||||
self.table.update_many(where=where, data={"spend": 0})
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BudgetWindowWrites:
|
||||
table: BatchTable
|
||||
|
||||
def queue_window_advance(self, budget_id: str, budget_reset_at: datetime) -> None:
|
||||
"""``update_many`` so a tier deleted between the read and the commit is a
|
||||
no-op row count instead of a P2025 that aborts the whole chunk."""
|
||||
self.table.update_many(where={"budget_id": budget_id}, data={"budget_reset_at": budget_reset_at})
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SpendResetUnitOfWork:
|
||||
keys: KeySpendResetWrites
|
||||
|
|
@ -50,6 +72,23 @@ class SpendResetUnitOfWork:
|
|||
teams: TeamSpendResetWrites
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BudgetCascadeUnitOfWork:
|
||||
"""Every write a budget-tier reset performs, bound to one batch.
|
||||
|
||||
The dependent spend rows and the budget rows' ``budget_reset_at`` advance
|
||||
must land together: advancing the window without zeroing the spend it
|
||||
gates leaves the dependents pinned at their cap until the next window.
|
||||
"""
|
||||
|
||||
team_memberships: LinkedSpendResetWrites
|
||||
keys: LinkedSpendResetWrites
|
||||
organizations: LinkedSpendResetWrites
|
||||
tags: LinkedSpendResetWrites
|
||||
endusers: LinkedSpendResetWrites
|
||||
budgets: BudgetWindowWrites
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def spend_reset_unit_of_work(new_batch: Callable[[], PrismaBatch]) -> AsyncGenerator[SpendResetUnitOfWork, None]:
|
||||
batch = new_batch()
|
||||
|
|
@ -59,3 +98,19 @@ async def spend_reset_unit_of_work(new_batch: Callable[[], PrismaBatch]) -> Asyn
|
|||
teams=TeamSpendResetWrites(table=batch.litellm_teamtable),
|
||||
)
|
||||
await batch.commit()
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def budget_cascade_unit_of_work(
|
||||
new_batch: Callable[[], PrismaBatch],
|
||||
) -> AsyncGenerator[BudgetCascadeUnitOfWork, None]:
|
||||
batch = new_batch()
|
||||
yield BudgetCascadeUnitOfWork(
|
||||
team_memberships=LinkedSpendResetWrites(table=batch.litellm_teammembership),
|
||||
keys=LinkedSpendResetWrites(table=batch.litellm_verificationtoken),
|
||||
organizations=LinkedSpendResetWrites(table=batch.litellm_organizationtable),
|
||||
tags=LinkedSpendResetWrites(table=batch.litellm_tagtable),
|
||||
endusers=LinkedSpendResetWrites(table=batch.litellm_endusertable),
|
||||
budgets=BudgetWindowWrites(table=batch.litellm_budgettable),
|
||||
)
|
||||
await batch.commit()
|
||||
|
|
|
|||
|
|
@ -2,7 +2,10 @@ import re
|
|||
import traceback
|
||||
from collections.abc import Iterable, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, TypedDict, overload
|
||||
|
||||
from openai.types.chat import ChatCompletionToolParam
|
||||
from openai.types.responses.function_tool_param import FunctionToolParam
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
|
||||
|
|
@ -18,6 +21,7 @@ from litellm.types.llms.openai import (
|
|||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamingResponse,
|
||||
)
|
||||
from litellm.types.llms.openai import ToolParam as ResponsesToolParam
|
||||
from litellm.types.utils import (
|
||||
CallTypes,
|
||||
Choices,
|
||||
|
|
@ -36,10 +40,14 @@ else:
|
|||
MCPTool = Any
|
||||
|
||||
# NOTE: We intentionally keep ToolParam as a broad type here to avoid tight coupling
|
||||
# to optional OpenAI SDK typing symbols in environments that may not have them available.
|
||||
# `Any` is used to keep mypy compatible with the broader OpenAI tool union types
|
||||
# passed around in Responses API while still allowing dict-style access at runtime.
|
||||
ToolParam = Any
|
||||
ToolParam: TypeAlias = Mapping[str, object]
|
||||
|
||||
|
||||
class MCPToolResult(TypedDict):
|
||||
tool_call_id: str | None
|
||||
result: str
|
||||
name: str | None
|
||||
|
||||
|
||||
LITELLM_PROXY_MCP_SERVER_URL: Final = "litellm_proxy"
|
||||
LITELLM_PROXY_MCP_SERVER_URL_PREFIX: Final = f"{LITELLM_PROXY_MCP_SERVER_URL}/mcp/"
|
||||
|
|
@ -199,13 +207,12 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
_get_tools_from_mcp_servers,
|
||||
)
|
||||
|
||||
mcp_servers: Final[list[str]] = []
|
||||
if mcp_tools_with_litellm_proxy:
|
||||
for _tool in mcp_tools_with_litellm_proxy:
|
||||
# if user specifies servers as server_url: litellm_proxy/mcp/zapier,github then return zapier,github
|
||||
server_url = _tool.get("server_url", "") if isinstance(_tool, dict) else ""
|
||||
if isinstance(server_url, str) and server_url.startswith(LITELLM_PROXY_MCP_SERVER_URL_PREFIX):
|
||||
mcp_servers.append(server_url.split("/")[-1])
|
||||
mcp_servers: Final = [
|
||||
server_url.split("/")[-1]
|
||||
for _tool in (mcp_tools_with_litellm_proxy or ())
|
||||
for server_url in (_tool.get("server_url", "") if isinstance(_tool, dict) else "",)
|
||||
if isinstance(server_url, str) and server_url.startswith(LITELLM_PROXY_MCP_SERVER_URL_PREFIX)
|
||||
]
|
||||
|
||||
# Resolve toolset names: collect all toolset IDs first, then apply their
|
||||
# combined permissions in a single pass so multiple toolsets are unioned
|
||||
|
|
@ -279,15 +286,15 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
allowed_mcp_servers=allowed_mcp_servers,
|
||||
)
|
||||
|
||||
server_names: Final[list[str]] = []
|
||||
for server in allowed_mcp_servers:
|
||||
if server is None:
|
||||
continue
|
||||
server_name = (
|
||||
getattr(server, "server_name", None) or getattr(server, "alias", None) or getattr(server, "name", None)
|
||||
server_names: Final = [
|
||||
server_name
|
||||
for server in allowed_mcp_servers
|
||||
if server is not None
|
||||
for server_name in (
|
||||
getattr(server, "server_name", None) or getattr(server, "alias", None) or getattr(server, "name", None),
|
||||
)
|
||||
if isinstance(server_name, str):
|
||||
server_names.append(server_name)
|
||||
if isinstance(server_name, str)
|
||||
]
|
||||
|
||||
return tools, server_names
|
||||
|
||||
|
|
@ -305,8 +312,8 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
List of deduplicated MCP tools
|
||||
The returned dictionary maps each tool_name to the server_name
|
||||
"""
|
||||
seen_names: Final = set()
|
||||
deduplicated_tools: Final = []
|
||||
seen_names: Final[set[str]] = set()
|
||||
deduplicated_tools: Final[list[MCPTool]] = []
|
||||
tool_server_map: Final[dict[str, str]] = {}
|
||||
|
||||
for tool in mcp_tools:
|
||||
|
|
@ -331,7 +338,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
) -> list[MCPTool]:
|
||||
"""Filter MCP tools based on allowed_tools parameter from the original tool configs."""
|
||||
# Collect all allowed tool names from all MCP tool configs
|
||||
allowed_tool_names: Final = set()
|
||||
allowed_tool_names: Final[set[str]] = set()
|
||||
for tool_config in mcp_tools_with_litellm_proxy:
|
||||
if isinstance(tool_config, dict) and "allowed_tools" in tool_config:
|
||||
allowed_tools = tool_config.get("allowed_tools", [])
|
||||
|
|
@ -343,23 +350,13 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
return mcp_tools
|
||||
|
||||
# Filter tools based on allowed names
|
||||
filtered_tools: Final = []
|
||||
for mcp_tool in mcp_tools:
|
||||
if isinstance(mcp_tool, dict):
|
||||
tool_name = mcp_tool.get("name")
|
||||
else:
|
||||
tool_name = getattr(mcp_tool, "name", None)
|
||||
|
||||
if not tool_name:
|
||||
continue
|
||||
|
||||
if tool_name in allowed_tool_names:
|
||||
filtered_tools.append(mcp_tool)
|
||||
continue
|
||||
|
||||
unprefixed_name, _ = split_server_prefix_from_name(tool_name)
|
||||
if unprefixed_name in allowed_tool_names:
|
||||
filtered_tools.append(mcp_tool)
|
||||
filtered_tools: Final = [
|
||||
mcp_tool
|
||||
for mcp_tool in mcp_tools
|
||||
for tool_name in (mcp_tool.get("name") if isinstance(mcp_tool, dict) else getattr(mcp_tool, "name", None),)
|
||||
if tool_name
|
||||
and (tool_name in allowed_tool_names or split_server_prefix_from_name(tool_name)[0] in allowed_tool_names)
|
||||
]
|
||||
|
||||
return filtered_tools
|
||||
|
||||
|
|
@ -448,24 +445,37 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
|
||||
return deduplicated_mcp_tools, tool_server_map
|
||||
|
||||
@overload
|
||||
@staticmethod
|
||||
def _transform_mcp_tools_to_openai(
|
||||
mcp_tools: Sequence[MCPTool],
|
||||
target_format: Literal["responses"] = ...,
|
||||
) -> list[FunctionToolParam]: ...
|
||||
|
||||
@overload
|
||||
@staticmethod
|
||||
def _transform_mcp_tools_to_openai(
|
||||
mcp_tools: Sequence[MCPTool],
|
||||
target_format: Literal["chat"],
|
||||
) -> list[ChatCompletionToolParam]: ...
|
||||
|
||||
@staticmethod
|
||||
def _transform_mcp_tools_to_openai(
|
||||
mcp_tools: Sequence[MCPTool],
|
||||
target_format: Literal["responses", "chat"] = "responses",
|
||||
) -> list[Any]:
|
||||
) -> Sequence[FunctionToolParam | ChatCompletionToolParam]:
|
||||
"""Transform MCP tools to OpenAI-compatible format."""
|
||||
from litellm.experimental_mcp_client.tools import (
|
||||
transform_mcp_tool_to_openai_responses_api_tool,
|
||||
transform_mcp_tool_to_openai_tool,
|
||||
)
|
||||
|
||||
openai_tools: Final[list[Any]] = []
|
||||
for mcp_tool in mcp_tools:
|
||||
if target_format == "chat":
|
||||
openai_tool = transform_mcp_tool_to_openai_tool(mcp_tool)
|
||||
else:
|
||||
openai_tool = transform_mcp_tool_to_openai_responses_api_tool(mcp_tool)
|
||||
openai_tools.append(openai_tool)
|
||||
openai_tools: Final = [
|
||||
transform_mcp_tool_to_openai_tool(mcp_tool)
|
||||
if target_format == "chat"
|
||||
else transform_mcp_tool_to_openai_responses_api_tool(mcp_tool)
|
||||
for mcp_tool in mcp_tools
|
||||
]
|
||||
|
||||
return openai_tools
|
||||
|
||||
|
|
@ -496,9 +506,9 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
return True
|
||||
|
||||
@staticmethod
|
||||
def _extract_tool_calls_from_response(response: ResponsesAPIResponse) -> list[Any]:
|
||||
def _extract_tool_calls_from_response(response: ResponsesAPIResponse) -> list[object]:
|
||||
"""Extract tool calls from the response output."""
|
||||
tool_calls: Final[list[Any]] = []
|
||||
tool_calls: Final[list[object]] = []
|
||||
for output_item in response.output:
|
||||
# Check if this is a function call output item
|
||||
if isinstance(output_item, dict) and output_item.get("type") == "function_call":
|
||||
|
|
@ -533,7 +543,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
|
||||
@staticmethod
|
||||
def _extract_tool_call_details(
|
||||
tool_call,
|
||||
tool_call: object,
|
||||
) -> tuple[str | None, str | None, str | None]:
|
||||
"""Extract tool name, arguments, and call_id from a tool call."""
|
||||
if isinstance(tool_call, dict):
|
||||
|
|
@ -566,7 +576,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
return tool_name, tool_arguments, tool_call_id
|
||||
|
||||
@staticmethod
|
||||
def _parse_tool_arguments(tool_arguments: Any) -> dict[str, Any]:
|
||||
def _parse_tool_arguments(tool_arguments: str | None) -> dict[str, object]:
|
||||
"""Parse tool arguments, handling both string and dict formats."""
|
||||
import json
|
||||
|
||||
|
|
@ -591,23 +601,18 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
# Fallback to generic handling if MCP types not available
|
||||
return "Tool executed successfully"
|
||||
|
||||
text_parts: Final = []
|
||||
other_content_types: Final = []
|
||||
|
||||
for content_item in result.content:
|
||||
if isinstance(content_item, TextContent):
|
||||
# Text content - extract the text
|
||||
text_parts.append(str(content_item.text))
|
||||
elif isinstance(content_item, ImageContent):
|
||||
# Image content
|
||||
other_content_types.append("Image")
|
||||
elif isinstance(content_item, EmbeddedResource):
|
||||
# Embedded resource
|
||||
other_content_types.append("EmbeddedResource")
|
||||
else:
|
||||
# Other unknown content types
|
||||
content_type = type(content_item).__name__
|
||||
other_content_types.append(content_type)
|
||||
text_parts: Final = [
|
||||
str(content_item.text) for content_item in result.content if isinstance(content_item, TextContent)
|
||||
]
|
||||
other_content_types: Final = [
|
||||
"Image"
|
||||
if isinstance(content_item, ImageContent)
|
||||
else "EmbeddedResource"
|
||||
if isinstance(content_item, EmbeddedResource)
|
||||
else type(content_item).__name__
|
||||
for content_item in result.content
|
||||
if not isinstance(content_item, TextContent)
|
||||
]
|
||||
|
||||
# Combine text parts if any
|
||||
result_text = " ".join(text_parts) if text_parts else ""
|
||||
|
|
@ -631,7 +636,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
litellm_call_id: str | None = None,
|
||||
litellm_trace_id: str | None = None,
|
||||
request_tags: list[str] | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
) -> list[MCPToolResult]:
|
||||
"""Execute tool calls and return results."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
|
|
@ -645,11 +650,11 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
)
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
tool_results: Final = []
|
||||
tool_results: Final[list[MCPToolResult]] = []
|
||||
tool_call_id: str | None = None
|
||||
rules_obj: Final = Rules()
|
||||
for tool_call in tool_calls:
|
||||
logging_request_data: dict[str, Any] = {}
|
||||
logging_request_data: dict[str, object] = {}
|
||||
tool_name: str | None = None
|
||||
try:
|
||||
(
|
||||
|
|
@ -678,7 +683,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
sanitized_tool_name = strip_known_server_prefix(resolved_tool_name, mcp_server)
|
||||
|
||||
start_time = datetime.now()
|
||||
logging_input = [
|
||||
logging_input: Sequence[Mapping[str, object]] = [
|
||||
{
|
||||
"role": "tool",
|
||||
"content": {
|
||||
|
|
@ -688,13 +693,14 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
}
|
||||
]
|
||||
tool_logging_call_id = litellm_call_id or str(uuid.uuid4())
|
||||
logging_metadata: dict[str, object] = {
|
||||
"tool_call_id": tool_call_id,
|
||||
"tool_name": sanitized_tool_name,
|
||||
"server_name": server_name,
|
||||
}
|
||||
logging_request_data = {
|
||||
"model": f"MCP: {tool_name}",
|
||||
"metadata": {
|
||||
"tool_call_id": tool_call_id,
|
||||
"tool_name": sanitized_tool_name,
|
||||
"server_name": server_name,
|
||||
},
|
||||
"metadata": logging_metadata,
|
||||
"input": logging_input,
|
||||
"call_type": CallTypes.call_mcp_tool.value,
|
||||
"litellm_call_id": tool_logging_call_id,
|
||||
|
|
@ -712,7 +718,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
if litellm_trace_id:
|
||||
logging_request_data["litellm_trace_id"] = litellm_trace_id
|
||||
if request_tags:
|
||||
logging_request_data["metadata"]["tags"] = request_tags
|
||||
logging_metadata["tags"] = request_tags
|
||||
if user_api_key_auth is not None:
|
||||
from litellm.proxy.litellm_pre_call_utils import (
|
||||
LiteLLMProxyRequestSetup,
|
||||
|
|
@ -902,16 +908,16 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
|
||||
@staticmethod
|
||||
def _create_follow_up_messages_for_chat(
|
||||
original_messages: list[Any],
|
||||
original_messages: list[object],
|
||||
response: ModelResponse,
|
||||
tool_results: Sequence[Mapping[str, object]],
|
||||
) -> list[Any]:
|
||||
) -> Sequence[Mapping[str, object]]:
|
||||
"""Create follow-up chat messages that include tool execution results."""
|
||||
from copy import deepcopy
|
||||
|
||||
from litellm.utils import convert_list_message_to_dict
|
||||
|
||||
follow_up_messages: list[Any] = convert_list_message_to_dict(deepcopy(original_messages))
|
||||
follow_up_messages: list[dict[str, object]] = convert_list_message_to_dict(deepcopy(original_messages))
|
||||
|
||||
if not follow_up_messages:
|
||||
follow_up_messages = []
|
||||
|
|
@ -950,9 +956,9 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
response: ResponsesAPIResponse,
|
||||
tool_results: Sequence[Mapping[str, object]],
|
||||
original_input: str | ResponseInputParam | None = None,
|
||||
) -> list[Any]:
|
||||
) -> list[object]:
|
||||
"""Create follow-up input with tool results in proper format."""
|
||||
follow_up_input: Final[list[Any]] = []
|
||||
follow_up_input: Final[list[object]] = []
|
||||
|
||||
# Add original user input if available to maintain conversation context
|
||||
if original_input:
|
||||
|
|
@ -964,8 +970,8 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
follow_up_input.append(original_input)
|
||||
|
||||
# Add the assistant message with function calls
|
||||
assistant_message_content: Final[list[Any]] = []
|
||||
function_calls: Final[list[dict[str, Any]]] = []
|
||||
assistant_message_content: Final[list[object]] = []
|
||||
function_calls: Final[list[dict[str, object]]] = []
|
||||
|
||||
for output_item in response.output:
|
||||
if not isinstance(output_item, dict) and hasattr(output_item, "model_dump"):
|
||||
|
|
@ -1027,7 +1033,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
async def _make_follow_up_call(
|
||||
follow_up_input: list[Any],
|
||||
model: str,
|
||||
all_tools: list[Any] | None,
|
||||
all_tools: Sequence[ResponsesToolParam] | None,
|
||||
response_id: str,
|
||||
**call_params: Any,
|
||||
) -> ResponsesAPIResponse | BaseResponsesAPIStreamingIterator:
|
||||
|
|
@ -1044,7 +1050,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
async def _log_mcp_tool_failure(
|
||||
*,
|
||||
proxy_logging_obj: Optional["ProxyLogging"],
|
||||
user_api_key_auth: Any,
|
||||
user_api_key_auth: "UserAPIKeyAuth | None",
|
||||
request_data: dict[str, object],
|
||||
error: Exception,
|
||||
) -> None:
|
||||
|
|
@ -1072,7 +1078,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
all_tools: Sequence[object] | None,
|
||||
mcp_tools_with_litellm_proxy: list[Mapping[str, object]],
|
||||
mcp_discovery_events: list[ResponsesAPIStreamingResponse],
|
||||
call_params: dict[str, Any],
|
||||
call_params: Mapping[str, object],
|
||||
previous_response_id: str | None,
|
||||
tool_server_map: dict[str, str],
|
||||
**kwargs,
|
||||
|
|
@ -1115,10 +1121,10 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
input: str | ResponseInputParam,
|
||||
model: str,
|
||||
all_tools: Sequence[object] | None,
|
||||
call_params: dict[str, Any],
|
||||
call_params: Mapping[str, object],
|
||||
previous_response_id: str | None,
|
||||
**kwargs,
|
||||
) -> dict[str, Any]:
|
||||
**kwargs: object,
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Build a clean request parameters dictionary for MCP streaming.
|
||||
|
||||
|
|
@ -1126,7 +1132,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
in a clean, maintainable way.
|
||||
"""
|
||||
# Start with the core required parameters
|
||||
request_params: Final = {
|
||||
request_params: Final[dict[str, object]] = {
|
||||
"input": input,
|
||||
"model": model,
|
||||
"tools": all_tools,
|
||||
|
|
@ -1146,7 +1152,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
|
||||
@staticmethod
|
||||
def _create_tool_execution_events(
|
||||
tool_calls: Sequence[object], tool_results: list[dict[str, Any]]
|
||||
tool_calls: Sequence[object], tool_results: Sequence[MCPToolResult]
|
||||
) -> list[ResponsesAPIStreamingResponse]:
|
||||
"""
|
||||
Create MCP tool execution events for streaming.
|
||||
|
|
|
|||
|
|
@ -19,13 +19,13 @@ from litellm.types.llms.openai import (
|
|||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
ResponsesAPIStreamingResponse,
|
||||
ToolParam,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.responses.mcp.litellm_proxy_mcp_handler import MCPToolResult
|
||||
else:
|
||||
MCPTool = Any
|
||||
|
||||
|
|
@ -33,7 +33,7 @@ MAX_MCP_TOOL_CALL_ROUNDS: Final = 5
|
|||
|
||||
|
||||
async def create_mcp_list_tools_events(
|
||||
mcp_tools_with_litellm_proxy: list[ToolParam],
|
||||
mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]],
|
||||
user_api_key_auth: "UserAPIKeyAuth | None",
|
||||
base_item_id: str,
|
||||
pre_processed_mcp_tools: list[MCPTool],
|
||||
|
|
@ -44,13 +44,14 @@ async def create_mcp_list_tools_events(
|
|||
|
||||
try:
|
||||
# Extract MCP server names
|
||||
mcp_servers: Final = []
|
||||
for tool in mcp_tools_with_litellm_proxy:
|
||||
if isinstance(tool, dict) and "server_url" in tool:
|
||||
server_url = tool.get("server_url")
|
||||
if isinstance(server_url, str) and server_url.startswith("litellm_proxy/mcp/"):
|
||||
server_name = server_url.split("/")[-1]
|
||||
mcp_servers.append(server_name)
|
||||
_mcp_servers: Final = [
|
||||
server_url.split("/")[-1]
|
||||
for tool in mcp_tools_with_litellm_proxy
|
||||
if isinstance(tool, dict)
|
||||
and "server_url" in tool
|
||||
and isinstance(server_url := tool.get("server_url"), str)
|
||||
and server_url.startswith("litellm_proxy/mcp/")
|
||||
]
|
||||
|
||||
# Emit list tools in progress event
|
||||
in_progress_event: Final = MCPListToolsInProgressEvent(
|
||||
|
|
@ -65,15 +66,14 @@ async def create_mcp_list_tools_events(
|
|||
filtered_mcp_tools: Final = pre_processed_mcp_tools
|
||||
|
||||
# Convert tools to dict format for the event
|
||||
mcp_tools_dict: Final = []
|
||||
for tool in filtered_mcp_tools:
|
||||
if hasattr(tool, "model_dump") and callable(getattr(tool, "model_dump")):
|
||||
# Type cast to help mypy understand this is safe after hasattr check
|
||||
mcp_tools_dict.append(cast(Any, tool).model_dump())
|
||||
elif hasattr(tool, "__dict__"):
|
||||
mcp_tools_dict.append(tool.__dict__)
|
||||
else:
|
||||
mcp_tools_dict.append({"name": getattr(tool, "name", str(tool))})
|
||||
_mcp_tools_dict: Final = [
|
||||
tool.model_dump()
|
||||
if hasattr(tool, "model_dump") and callable(getattr(tool, "model_dump"))
|
||||
else tool.__dict__
|
||||
if hasattr(tool, "__dict__")
|
||||
else {"name": getattr(tool, "name", str(tool))}
|
||||
for tool in filtered_mcp_tools
|
||||
]
|
||||
|
||||
# Emit list tools completed event
|
||||
completed_event: Final = MCPListToolsCompletedEvent(
|
||||
|
|
@ -96,21 +96,18 @@ async def create_mcp_list_tools_events(
|
|||
server_label = str(server_label_value) if server_label_value is not None else ""
|
||||
|
||||
# Format tools for OpenAI output_item.done format
|
||||
formatted_tools: Final = []
|
||||
for tool in filtered_mcp_tools:
|
||||
tool_dict = {
|
||||
formatted_tools: Final = [
|
||||
{
|
||||
"name": getattr(tool, "name", "unknown"),
|
||||
"description": getattr(tool, "description", ""),
|
||||
"annotations": {"read_only": False},
|
||||
**dict.fromkeys(
|
||||
("input_schema",) if hasattr(tool, "inputSchema") or hasattr(tool, "input_schema") else (),
|
||||
getattr(tool, "inputSchema", getattr(tool, "input_schema", None)),
|
||||
),
|
||||
}
|
||||
|
||||
# Add input_schema if available
|
||||
if hasattr(tool, "inputSchema"):
|
||||
tool_dict["input_schema"] = getattr(tool, "inputSchema")
|
||||
elif hasattr(tool, "input_schema"):
|
||||
tool_dict["input_schema"] = getattr(tool, "input_schema")
|
||||
|
||||
formatted_tools.append(tool_dict)
|
||||
for tool in filtered_mcp_tools
|
||||
]
|
||||
|
||||
# Create the output_item.done event with MCP tools list
|
||||
output_item_done_event = OutputItemDoneEvent(
|
||||
|
|
@ -166,7 +163,7 @@ async def create_mcp_list_tools_events(
|
|||
|
||||
def create_mcp_call_events(
|
||||
tool_name: str,
|
||||
tool_call_id: str,
|
||||
tool_call_id: str | None,
|
||||
arguments: str,
|
||||
result: str | None = None,
|
||||
base_item_id: str | None = None,
|
||||
|
|
@ -256,9 +253,12 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
4. Emits tool execution events in the stream
|
||||
"""
|
||||
|
||||
model: str
|
||||
tool_results: "Sequence[MCPToolResult]"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_iterator: Any, # Can be None - will be created internally
|
||||
base_iterator: "BaseResponsesAPIStreamingIterator | ResponsesAPIResponse | None", # created internally when None
|
||||
mcp_events: list[ResponsesAPIStreamingResponse],
|
||||
tool_server_map: dict[str, str],
|
||||
mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]] | None = None,
|
||||
|
|
@ -285,7 +285,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
self.tool_server_map = tool_server_map
|
||||
|
||||
# Iterator references
|
||||
self.base_iterator: Any | ResponsesAPIResponse | None = base_iterator # Will be created when needed
|
||||
self.base_iterator: BaseResponsesAPIStreamingIterator | ResponsesAPIResponse | None = (
|
||||
base_iterator # Will be created when needed
|
||||
)
|
||||
|
||||
# Response collection for tool execution
|
||||
self.collected_response: ResponsesAPIResponse | None = None
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ import weakref
|
|||
from collections import defaultdict
|
||||
from collections.abc import AsyncGenerator, Callable, Generator, Mapping, Sequence
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeVar, Union, cast
|
||||
|
||||
import anyio
|
||||
|
|
@ -111,12 +112,14 @@ from litellm.router_utils.common_utils import (
|
|||
filter_web_search_deployments,
|
||||
resolve_model_group_alias,
|
||||
truncate_fallback_error_detail,
|
||||
warn_on_provider_credential_mismatch,
|
||||
)
|
||||
from litellm.router_utils.cooldown_cache import CooldownCache
|
||||
from litellm.router_utils.cooldown_handlers import (
|
||||
DEFAULT_COOLDOWN_TIME_SECONDS,
|
||||
_async_get_cooldown_deployments,
|
||||
_async_get_cooldown_deployments_with_debug_info,
|
||||
_first_present, # pyright: ignore[reportPrivateUsage] - shared internal helper across router_utils submodules, matching the other cooldown_handlers imports on this line
|
||||
_get_cooldown_deployments,
|
||||
_set_cooldown_deployments,
|
||||
is_advisor_orchestration_failure,
|
||||
|
|
@ -1898,7 +1901,7 @@ class Router:
|
|||
# Set per-deployment num_retries on exception for retry logic
|
||||
if deployment is not None:
|
||||
self._set_deployment_num_retries_on_exception(e, deployment)
|
||||
self._set_failed_deployment_id_on_exception(e, deployment)
|
||||
self._stamp_failed_deployment_id_with_effective_model_info(e, deployment, kwargs)
|
||||
raise e
|
||||
|
||||
def _get_silent_experiment_kwargs(self, **kwargs) -> dict:
|
||||
|
|
@ -2961,7 +2964,7 @@ class Router:
|
|||
# Set per-deployment num_retries on exception for retry logic
|
||||
if deployment is not None:
|
||||
self._set_deployment_num_retries_on_exception(e, deployment)
|
||||
self._set_failed_deployment_id_on_exception(e, deployment)
|
||||
self._stamp_failed_deployment_id_with_effective_model_info(e, deployment, kwargs)
|
||||
raise e
|
||||
except Exception as e:
|
||||
verbose_router_logger.info("litellm.acompletion(model=%s)\x1b[31m Exception %s\x1b[0m", model_name, e)
|
||||
|
|
@ -2970,7 +2973,7 @@ class Router:
|
|||
# Set per-deployment num_retries on exception for retry logic
|
||||
if deployment is not None:
|
||||
self._set_deployment_num_retries_on_exception(e, deployment)
|
||||
self._set_failed_deployment_id_on_exception(e, deployment)
|
||||
self._stamp_failed_deployment_id_with_effective_model_info(e, deployment, kwargs)
|
||||
raise e
|
||||
|
||||
def _update_kwargs_before_fallbacks(
|
||||
|
|
@ -3016,7 +3019,7 @@ class Router:
|
|||
except (ValueError, TypeError):
|
||||
pass # Skip if value can't be converted to int
|
||||
|
||||
def _set_failed_deployment_id_on_exception(self, exception: Exception, deployment: dict) -> None:
|
||||
def _set_failed_deployment_id_on_exception(self, exception: Exception, deployment: Mapping[str, Any]) -> None:
|
||||
"""
|
||||
Stamp the failed deployment's `model_info.id` on the exception so the
|
||||
fallback layer can exclude it from subsequent re-picks within the same
|
||||
|
|
@ -3035,6 +3038,16 @@ class Router:
|
|||
except Exception:
|
||||
pass
|
||||
|
||||
def _stamp_failed_deployment_id_with_effective_model_info(
|
||||
self, exception: Exception, deployment: Mapping[str, Any], kwargs: Mapping[str, Any]
|
||||
) -> None:
|
||||
# A client-side-credential call gets a dynamic deployment id generated inside
|
||||
# _update_kwargs_with_deployment and stamped into kwargs["model_info"]; stamping
|
||||
# the static shared deployment's id instead would let one tenant's bad credentials
|
||||
# cool down the deployment every other tenant sharing this config relies on.
|
||||
effective_model_info: Final = kwargs.get("model_info") or deployment.get("model_info") or MappingProxyType({})
|
||||
self._set_failed_deployment_id_on_exception(exception, MappingProxyType({"model_info": effective_model_info}))
|
||||
|
||||
def _update_kwargs_with_default_litellm_params(
|
||||
self, kwargs: dict, metadata_variable_name: str | None = "metadata"
|
||||
) -> None:
|
||||
|
|
@ -4521,10 +4534,11 @@ class Router:
|
|||
|
||||
passthrough_on_no_deployment: Final = kwargs.pop("passthrough_on_no_deployment", False)
|
||||
function_name: Final = "_ageneric_api_call_with_fallbacks"
|
||||
deployment = None # rebind-ok: pre-init so the except block can stamp a failure with no deployment picked
|
||||
try:
|
||||
parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs)
|
||||
try:
|
||||
deployment: Final = await self.async_get_available_deployment(
|
||||
deployment = await self.async_get_available_deployment( # rebind-ok: set on success, see pre-init above
|
||||
model=model,
|
||||
request_kwargs=kwargs,
|
||||
messages=kwargs.get("messages", None),
|
||||
|
|
@ -4601,6 +4615,8 @@ class Router:
|
|||
)
|
||||
if model is not None:
|
||||
self.fail_calls[model] += 1
|
||||
if deployment is not None:
|
||||
self._stamp_failed_deployment_id_with_effective_model_info(e, deployment, kwargs)
|
||||
raise e
|
||||
|
||||
async def _aresponses_with_streaming_fallbacks(
|
||||
|
|
@ -7078,7 +7094,9 @@ class Router:
|
|||
)
|
||||
|
||||
# Determine cooldown time with priority: deployment config > response header > router default
|
||||
deployment_cooldown: Final = litellm_params.get("cooldown_time", None)
|
||||
deployment_cooldown: Final = _first_present(
|
||||
_model_info if isinstance(_model_info, dict) else None, litellm_params, key="cooldown_time"
|
||||
)
|
||||
|
||||
header_cooldown = None
|
||||
if exception_headers is not None:
|
||||
|
|
@ -7523,6 +7541,7 @@ class Router:
|
|||
"""
|
||||
try:
|
||||
litellm_params: Final[LiteLLM_Params] = LiteLLM_Params(**_litellm_params)
|
||||
warn_on_provider_credential_mismatch(model_name=_model_name, litellm_params=_litellm_params)
|
||||
deployment = Deployment(
|
||||
**deployment_info,
|
||||
model_name=_model_name,
|
||||
|
|
@ -8215,6 +8234,11 @@ class Router:
|
|||
if _deployment_model_id and self.has_model_id(_deployment_model_id):
|
||||
return None
|
||||
|
||||
warn_on_provider_credential_mismatch(
|
||||
model_name=deployment.model_name,
|
||||
litellm_params=deployment.litellm_params.model_dump(exclude_none=True),
|
||||
)
|
||||
|
||||
# add to model list
|
||||
_deployment: Final = deployment.to_json(exclude_none=True)
|
||||
# initialize client
|
||||
|
|
@ -11800,6 +11824,23 @@ class Router:
|
|||
and allowed_fails_policy.BadRequestErrorAllowedFails is not None
|
||||
):
|
||||
return allowed_fails_policy.BadRequestErrorAllowedFails
|
||||
if (
|
||||
isinstance(exception, litellm.InternalServerError)
|
||||
and allowed_fails_policy.InternalServerErrorAllowedFails is not None
|
||||
):
|
||||
return allowed_fails_policy.InternalServerErrorAllowedFails
|
||||
if (
|
||||
isinstance(exception, litellm.ServiceUnavailableError)
|
||||
and allowed_fails_policy.ServiceUnavailableErrorAllowedFails is not None
|
||||
):
|
||||
return allowed_fails_policy.ServiceUnavailableErrorAllowedFails
|
||||
if (
|
||||
isinstance(exception, litellm.BadGatewayError)
|
||||
and allowed_fails_policy.BadGatewayErrorAllowedFails is not None
|
||||
):
|
||||
return allowed_fails_policy.BadGatewayErrorAllowedFails
|
||||
if isinstance(exception, litellm.NotFoundError) and allowed_fails_policy.NotFoundErrorAllowedFails is not None:
|
||||
return allowed_fails_policy.NotFoundErrorAllowedFails
|
||||
|
||||
def _initialize_alerting(self):
|
||||
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
|
||||
|
|
|
|||
|
|
@ -1,15 +1,18 @@
|
|||
import hashlib
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.llms.openai import OpenAIFileObject
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._logging import verbose_logger, verbose_router_logger
|
||||
from litellm.constants import ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS
|
||||
from litellm.exceptions import BadRequestError
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.types.router import CredentialLiteLLMParams
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
|
||||
def _is_proxy_admin_request(request_kwargs: Mapping[str, object] | None) -> bool:
|
||||
|
|
@ -210,3 +213,77 @@ def filter_web_search_deployments(
|
|||
if len(healthy_deployments) > 0 and len(final_deployments) == 0:
|
||||
verbose_logger.warning("No deployments support web search for request")
|
||||
return final_deployments
|
||||
|
||||
|
||||
# Credential params that only one provider family reads, paired with the providers
|
||||
# that read them. A deployment carrying them while resolving elsewhere is almost
|
||||
# always a missing route prefix: `model: claude-sonnet-5` with `aws_region_name`
|
||||
# set resolves to the first-party Anthropic API, silently ignores the AWS
|
||||
# credentials, and 401s at request time.
|
||||
_AWS_PROVIDERS: Final = frozenset(
|
||||
provider.value for provider in LlmProviders if provider.value.startswith(("bedrock", "sagemaker"))
|
||||
)
|
||||
_VERTEX_PROVIDERS: Final = frozenset(
|
||||
provider.value for provider in LlmProviders if provider.value.startswith("vertex_ai")
|
||||
)
|
||||
|
||||
PROVIDER_SCOPED_CREDENTIAL_PARAMS: Final[Mapping[str, frozenset[str]]] = MappingProxyType(
|
||||
{
|
||||
"aws_access_key_id": _AWS_PROVIDERS,
|
||||
"aws_profile_name": _AWS_PROVIDERS,
|
||||
"aws_region_name": _AWS_PROVIDERS,
|
||||
"aws_role_name": _AWS_PROVIDERS,
|
||||
"aws_secret_access_key": _AWS_PROVIDERS,
|
||||
"aws_session_name": _AWS_PROVIDERS,
|
||||
"aws_session_token": _AWS_PROVIDERS,
|
||||
"aws_web_identity_token": _AWS_PROVIDERS,
|
||||
"vertex_credentials": _VERTEX_PROVIDERS,
|
||||
"vertex_location": _VERTEX_PROVIDERS,
|
||||
"vertex_project": _VERTEX_PROVIDERS,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def warn_on_provider_credential_mismatch(model_name: str, litellm_params: Mapping[str, object]) -> str | None:
|
||||
"""
|
||||
Warn when a deployment carries one provider's credentials but resolves to another.
|
||||
|
||||
Returns the warning text (for tests), or None when the deployment is consistent
|
||||
or its provider cannot be resolved. Never raises: a deployment litellm cannot
|
||||
classify is left alone rather than blocking router startup.
|
||||
|
||||
Only inline credential params are examined. A deployment that sources them
|
||||
through ``litellm_credential_name`` resolves them after registration, so it
|
||||
carries none of these keys here and is left alone rather than warned about
|
||||
on incomplete information.
|
||||
"""
|
||||
model: Final = litellm_params.get("model")
|
||||
if not isinstance(model, str) or not model:
|
||||
return None
|
||||
scoped: Final = tuple(param for param in PROVIDER_SCOPED_CREDENTIAL_PARAMS if litellm_params.get(param) is not None)
|
||||
if not scoped:
|
||||
return None
|
||||
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
|
||||
try:
|
||||
_, resolved_provider, _, _ = get_llm_provider(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider if isinstance(custom_llm_provider, str) else None,
|
||||
)
|
||||
except BadRequestError:
|
||||
return None
|
||||
mismatched: Final = sorted(
|
||||
param for param in scoped if resolved_provider not in PROVIDER_SCOPED_CREDENTIAL_PARAMS[param]
|
||||
)
|
||||
if not mismatched:
|
||||
return None
|
||||
expected: Final = sorted(
|
||||
{provider for param in mismatched for provider in PROVIDER_SCOPED_CREDENTIAL_PARAMS[param]}
|
||||
)
|
||||
warning: Final = (
|
||||
f"Deployment '{model_name}' sets {mismatched} but 'model={model}' resolves to provider "
|
||||
f"'{resolved_provider}', which ignores them. Those params are read by {expected}, so this is "
|
||||
f"usually a missing route prefix (e.g. '{expected[0]}/{model}'); as written the request goes to "
|
||||
f"'{resolved_provider}' and will fail on that provider's credentials."
|
||||
)
|
||||
verbose_router_logger.warning(warning)
|
||||
return warning
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ Wrapper around router cache. Meant to handle model cooldown logic
|
|||
|
||||
import functools
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from typing_extensions import TypedDict
|
||||
|
|
@ -28,6 +29,12 @@ class CooldownCacheValue(TypedDict):
|
|||
cooldown_time: float
|
||||
|
||||
|
||||
# Cap on the corrected in-memory TTL set in `_corrected_active_cooldown`: re-checks the
|
||||
# real remaining cooldown against Redis at least this often, so an entry that later gets
|
||||
# deleted or extended in Redis before its original deadline is still noticed promptly.
|
||||
_MAX_CORRECTED_IN_MEMORY_TTL_SECONDS: Final = 60.0
|
||||
|
||||
|
||||
class CooldownCache:
|
||||
def __init__(self, cache: DualCache, default_cooldown_time: float):
|
||||
self.cache = cache
|
||||
|
|
@ -100,6 +107,30 @@ class CooldownCache:
|
|||
def get_cooldown_cache_key(model_id: str) -> str:
|
||||
return "deployment:" + model_id + ":cooldown"
|
||||
|
||||
def _corrected_active_cooldown(
|
||||
self,
|
||||
key: str,
|
||||
result: Mapping[str, Any],
|
||||
current_time: float,
|
||||
) -> CooldownCacheValue | None:
|
||||
"""
|
||||
Return a CooldownCacheValue if the cooldown is still active, or None if it has expired.
|
||||
|
||||
Also corrects the in-memory TTL when DualCache promotes a Redis entry using the
|
||||
default 600s TTL instead of the true remaining cooldown time.
|
||||
"""
|
||||
cooldown_cache_value: Final = CooldownCacheValue(**result) # pyright: ignore[reportUnknownArgumentType] - result comes from an untyped cache read, not from our own code
|
||||
remaining: Final = (cooldown_cache_value["timestamp"] + cooldown_cache_value["cooldown_time"]) - current_time
|
||||
if remaining <= 0:
|
||||
self.cache.in_memory_cache.delete_cache(key)
|
||||
return None
|
||||
current_expiry: Final = self.cache.in_memory_cache.ttl_dict.get(key)
|
||||
if current_expiry is not None and current_expiry > current_time + remaining + 5:
|
||||
corrected_ttl: Final = min(remaining, _MAX_CORRECTED_IN_MEMORY_TTL_SECONDS)
|
||||
self.cache.in_memory_cache.delete_cache(key)
|
||||
self.cache.in_memory_cache.set_cache(key, result, ttl=corrected_ttl)
|
||||
return cooldown_cache_value
|
||||
|
||||
async def async_get_active_cooldowns(
|
||||
self, model_ids: list[str], parent_otel_span: Span | None
|
||||
) -> list[tuple[str, CooldownCacheValue]]:
|
||||
|
|
@ -117,11 +148,13 @@ class CooldownCache:
|
|||
if results is None or all(v is None for v in results):
|
||||
return active_cooldowns
|
||||
|
||||
# Process the results
|
||||
current_time: Final = time.time()
|
||||
for model_id, result in zip(model_ids, results):
|
||||
if result and isinstance(result, dict):
|
||||
cooldown_cache_value = CooldownCacheValue(**result)
|
||||
active_cooldowns.append((model_id, cooldown_cache_value))
|
||||
key = CooldownCache.get_cooldown_cache_key(model_id)
|
||||
cooldown_cache_value = self._corrected_active_cooldown(key, result, current_time)
|
||||
if cooldown_cache_value is not None:
|
||||
active_cooldowns.append((model_id, cooldown_cache_value))
|
||||
|
||||
return active_cooldowns
|
||||
|
||||
|
|
@ -134,11 +167,13 @@ class CooldownCache:
|
|||
results: Final = self.cache.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
|
||||
|
||||
active_cooldowns: Final = []
|
||||
# Process the results
|
||||
current_time: Final = time.time()
|
||||
for model_id, result in zip(model_ids, results):
|
||||
if result and isinstance(result, dict):
|
||||
cooldown_cache_value = CooldownCacheValue(**result)
|
||||
active_cooldowns.append((model_id, cooldown_cache_value))
|
||||
key = CooldownCache.get_cooldown_cache_key(model_id)
|
||||
cooldown_cache_value = self._corrected_active_cooldown(key, result, current_time)
|
||||
if cooldown_cache_value is not None:
|
||||
active_cooldowns.append((model_id, cooldown_cache_value))
|
||||
|
||||
return active_cooldowns
|
||||
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@ Router cooldown handlers
|
|||
|
||||
import asyncio
|
||||
import math
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import litellm
|
||||
|
|
@ -58,6 +60,148 @@ def is_advisor_orchestration_failure(exception: BaseException | None) -> bool:
|
|||
return bool(getattr(exception, _ADVISOR_ORCHESTRATION_FAILURE_ATTR, False))
|
||||
|
||||
|
||||
_EXCEPTION_POLICY_FIELDS: Final[tuple[tuple[type, str], ...]] = (
|
||||
# ContentPolicyViolationError subclasses BadRequestError, so it must be checked first.
|
||||
(litellm.ContentPolicyViolationError, "ContentPolicyViolationErrorAllowedFails"),
|
||||
(litellm.BadRequestError, "BadRequestErrorAllowedFails"),
|
||||
(litellm.AuthenticationError, "AuthenticationErrorAllowedFails"),
|
||||
(litellm.Timeout, "TimeoutErrorAllowedFails"),
|
||||
(litellm.RateLimitError, "RateLimitErrorAllowedFails"),
|
||||
(litellm.InternalServerError, "InternalServerErrorAllowedFails"),
|
||||
(litellm.ServiceUnavailableError, "ServiceUnavailableErrorAllowedFails"),
|
||||
(litellm.BadGatewayError, "BadGatewayErrorAllowedFails"),
|
||||
(litellm.NotFoundError, "NotFoundErrorAllowedFails"),
|
||||
)
|
||||
|
||||
|
||||
def _first_present(*sources: Mapping[str, Any] | None, key: str) -> int | float | None:
|
||||
"""Return *key* from the first source mapping where it's set, so callers can
|
||||
support a setting living in more than one deployment config location. Sources
|
||||
are checked in order from most to least specific to that setting."""
|
||||
for source in sources:
|
||||
if source is None:
|
||||
continue
|
||||
value = source.get(key)
|
||||
if value is not None:
|
||||
return value
|
||||
return None
|
||||
|
||||
|
||||
def _get_deployment_cooldown_policy(
|
||||
litellm_router_instance: LitellmRouter,
|
||||
deployment: str,
|
||||
) -> tuple[Mapping[str, int] | None, int | None]:
|
||||
"""Return (allowed_fails_policy, allowed_fails) from deployment model_info, or (None, None).
|
||||
|
||||
`model_info` is the only supported location for these two fields (unlike
|
||||
`cooldown_time`, they have no pre-existing `litellm_params` precedent): `litellm_params`
|
||||
gets copied wholesale into the actual provider call kwargs (see e.g.
|
||||
Router._image_generation's `data = deployment["litellm_params"].copy()`), so a new
|
||||
field placed there would leak into the outgoing LLM request instead of staying
|
||||
router-internal.
|
||||
"""
|
||||
dep: Final = litellm_router_instance.get_model_info(id=deployment)
|
||||
if dep is None:
|
||||
return None, None
|
||||
mi: Final[Mapping[str, Any]] = dep.get("model_info") or MappingProxyType({})
|
||||
raw: Final = mi.get("allowed_fails_policy")
|
||||
policy: Final[Mapping[str, int] | None] = raw if isinstance(raw, dict) else None
|
||||
allowed: Final[int | None] = mi.get("allowed_fails")
|
||||
return policy, allowed
|
||||
|
||||
|
||||
def _resolve_allowed_fails_from_policy(
|
||||
policy: Mapping[str, int] | None,
|
||||
exception: Exception,
|
||||
) -> int | None:
|
||||
"""Match *exception* against *policy* and return the configured allowed-fail count, or None."""
|
||||
if policy is None:
|
||||
return None
|
||||
for exc_type, field in _EXCEPTION_POLICY_FIELDS:
|
||||
if isinstance(exception, exc_type):
|
||||
value = policy.get(field)
|
||||
if value is not None:
|
||||
return value
|
||||
return None
|
||||
|
||||
|
||||
def _should_cooldown_based_on_deployment_policy(
|
||||
litellm_router_instance: LitellmRouter,
|
||||
deployment: str,
|
||||
original_exception: Exception,
|
||||
dep_policy: Mapping[str, int] | None,
|
||||
dep_allowed_fails: int | None,
|
||||
is_single_deployment_model_group: bool,
|
||||
) -> bool:
|
||||
"""Resolve deployment-level allowed-fails and delegate to the shared counting logic.
|
||||
|
||||
When the deployment's policy doesn't cover *original_exception*'s type and no
|
||||
deployment-wide `allowed_fails` is set either, defer to router-level behavior
|
||||
instead of forcing an immediate cooldown.
|
||||
|
||||
A generic, deployment-wide `allowed_fails` predates this feature's per-exception-type
|
||||
policy and is a much less deliberate opt-in, so on a single-deployment model group it
|
||||
still defers to the "avoid cooldowns on single deployment model groups" safety net
|
||||
(see `_should_cooldown_deployment`'s BASE CASE) rather than silently disabling it. An
|
||||
explicit, named-exception-type `allowed_fails_policy` entry is unambiguous enough to
|
||||
override that safety net, matching `_has_explicit_allowed_fails_policy_for_exception`.
|
||||
"""
|
||||
allowed_fails_from_policy: Final = _resolve_allowed_fails_from_policy(dep_policy, original_exception)
|
||||
if allowed_fails_from_policy is None and dep_allowed_fails is not None and is_single_deployment_model_group:
|
||||
return False
|
||||
|
||||
allowed_fails_override: Final[int | None] = (
|
||||
allowed_fails_from_policy if allowed_fails_from_policy is not None else dep_allowed_fails
|
||||
)
|
||||
cache_key_suffix: Final[str | None] = (
|
||||
type(original_exception).__name__
|
||||
if allowed_fails_from_policy is not None
|
||||
else ("generic" if dep_allowed_fails is not None else None)
|
||||
)
|
||||
|
||||
dep: Final = litellm_router_instance.get_model_info(id=deployment)
|
||||
cooldown_time_override: Final = (
|
||||
_first_present(dep.get("model_info"), dep.get("litellm_params"), key="cooldown_time")
|
||||
if dep is not None
|
||||
else None
|
||||
)
|
||||
|
||||
return should_cooldown_based_on_allowed_fails_policy(
|
||||
litellm_router_instance=litellm_router_instance,
|
||||
deployment=deployment,
|
||||
original_exception=original_exception,
|
||||
allowed_fails_override=allowed_fails_override,
|
||||
cooldown_time_override=cooldown_time_override,
|
||||
cache_key_suffix=cache_key_suffix,
|
||||
)
|
||||
|
||||
|
||||
def _has_explicit_allowed_fails_policy_for_exception(
|
||||
litellm_router_instance: LitellmRouter,
|
||||
deployment: str | None,
|
||||
original_exception: Exception,
|
||||
) -> bool:
|
||||
"""True if this deployment has an explicit, deployment-level allowed_fails_policy
|
||||
entry matching *original_exception*'s type.
|
||||
|
||||
`_is_cooldown_required` skips cooldown evaluation for most 4XX errors (BadRequestError,
|
||||
ContentPolicyViolationError) by default, since a generic client error is usually not the
|
||||
deployment's fault. A deployment-level allowed_fails_policy entry naming that exact
|
||||
exception type is this PR's own per-deployment opt-in, so it overrides that default.
|
||||
|
||||
Deliberately scoped to the deployment level only, and to the named-exception-type
|
||||
policy dict rather than a plain `allowed_fails` integer: a pre-existing router-wide
|
||||
`allowed_fails_policy` (or a deployment's generic `allowed_fails`) predates this
|
||||
feature and must keep its existing behavior for 4XX types `_is_cooldown_required`
|
||||
already excludes, rather than silently start cooling down deployments whose configs
|
||||
never opted into this specific override.
|
||||
"""
|
||||
if deployment is None:
|
||||
return False
|
||||
dep_policy, _ = _get_deployment_cooldown_policy(litellm_router_instance, deployment)
|
||||
return _resolve_allowed_fails_from_policy(dep_policy, original_exception) is not None
|
||||
|
||||
|
||||
def _is_cooldown_required(
|
||||
litellm_router_instance: LitellmRouter,
|
||||
model_id: str,
|
||||
|
|
@ -155,6 +299,10 @@ def _should_run_cooldown_logic(
|
|||
model_id=deployment,
|
||||
exception_status=exception_status,
|
||||
exception_str=str(original_exception),
|
||||
) and not _has_explicit_allowed_fails_policy_for_exception(
|
||||
litellm_router_instance=litellm_router_instance,
|
||||
deployment=deployment,
|
||||
original_exception=original_exception,
|
||||
):
|
||||
verbose_router_logger.debug("Should Not Run Cooldown Logic: _is_cooldown_required returned False")
|
||||
return False
|
||||
|
|
@ -190,11 +338,24 @@ def _should_cooldown_deployment(
|
|||
|
||||
- v1 logic (Legacy): if allowed fails or allowed fail policy set, coolsdown if num fails in this minute > allowed fails
|
||||
"""
|
||||
## BASE CASE - single deployment
|
||||
model_group: Final = litellm_router_instance.get_model_group(id=deployment)
|
||||
is_single_deployment_model_group = False
|
||||
if model_group is not None and len(model_group) == 1:
|
||||
is_single_deployment_model_group = True
|
||||
|
||||
## CHECK DEPLOYMENT-LEVEL POLICY FIRST (overrides router-level)
|
||||
dep_policy, dep_allowed_fails = _get_deployment_cooldown_policy(litellm_router_instance, deployment)
|
||||
if dep_policy is not None or dep_allowed_fails is not None:
|
||||
return _should_cooldown_based_on_deployment_policy(
|
||||
litellm_router_instance,
|
||||
deployment,
|
||||
original_exception,
|
||||
dep_policy,
|
||||
dep_allowed_fails,
|
||||
is_single_deployment_model_group,
|
||||
)
|
||||
|
||||
## BASE CASE - single deployment
|
||||
if (
|
||||
litellm_router_instance.allowed_fails_policy is None
|
||||
and _is_allowed_fails_set_on_router(litellm_router_instance=litellm_router_instance) is False
|
||||
|
|
@ -382,29 +543,50 @@ def should_cooldown_based_on_allowed_fails_policy(
|
|||
litellm_router_instance: LitellmRouter,
|
||||
deployment: str,
|
||||
original_exception: Any,
|
||||
allowed_fails_override: int | None = None,
|
||||
cooldown_time_override: float | None = None,
|
||||
cache_key_suffix: str | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if fails are within the allowed limit and update the number of fails.
|
||||
|
||||
When *allowed_fails_override* / *cooldown_time_override* are supplied they
|
||||
take precedence over the router-level values (used by deployment-level overrides).
|
||||
|
||||
When *cache_key_suffix* is supplied the fail counter is keyed as
|
||||
``{deployment}:{cache_key_suffix}`` so that different exception types are
|
||||
tracked independently per deployment.
|
||||
|
||||
Returns:
|
||||
- True if fails exceed the allowed limit (should cooldown)
|
||||
- False if fails are within the allowed limit (should not cooldown)
|
||||
"""
|
||||
allowed_fails: Final = (
|
||||
litellm_router_instance.get_allowed_fails_from_policy(
|
||||
exception=original_exception,
|
||||
)
|
||||
or litellm_router_instance.allowed_fails
|
||||
allowed_fails_from_policy: Final = litellm_router_instance.get_allowed_fails_from_policy(
|
||||
exception=original_exception
|
||||
)
|
||||
allowed_fails: Final = (
|
||||
allowed_fails_override
|
||||
if allowed_fails_override is not None
|
||||
else (
|
||||
allowed_fails_from_policy
|
||||
if allowed_fails_from_policy is not None
|
||||
else litellm_router_instance.allowed_fails
|
||||
)
|
||||
)
|
||||
cooldown_time: Final = (
|
||||
cooldown_time_override
|
||||
if cooldown_time_override is not None
|
||||
else (litellm_router_instance.cooldown_time or DEFAULT_COOLDOWN_TIME_SECONDS)
|
||||
)
|
||||
cooldown_time: Final = litellm_router_instance.cooldown_time or DEFAULT_COOLDOWN_TIME_SECONDS
|
||||
|
||||
current_fails: Final = litellm_router_instance.failed_calls.get_cache(key=deployment) or 0
|
||||
cache_key: Final = f"{deployment}:{cache_key_suffix}" if cache_key_suffix else deployment
|
||||
current_fails: Final = litellm_router_instance.failed_calls.get_cache(key=cache_key) or 0
|
||||
updated_fails: Final = current_fails + 1
|
||||
|
||||
if updated_fails > allowed_fails:
|
||||
return True
|
||||
else:
|
||||
litellm_router_instance.failed_calls.set_cache(key=deployment, value=updated_fails, ttl=cooldown_time)
|
||||
litellm_router_instance.failed_calls.set_cache(key=cache_key, value=updated_fails, ttl=cooldown_time)
|
||||
|
||||
return False
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import hashlib
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
|
@ -12,6 +13,16 @@ from litellm.router_utils.add_retry_fallback_headers import (
|
|||
add_fallback_headers_to_response,
|
||||
get_fallback_error_info,
|
||||
)
|
||||
from litellm.router_utils.batch_utils import _get_router_metadata_variable_name
|
||||
from litellm.router_utils.cooldown_handlers import (
|
||||
_first_present, # pyright: ignore[reportPrivateUsage] - shared internal helper, used across router_utils
|
||||
_set_cooldown_deployments, # pyright: ignore[reportPrivateUsage] - shared helper, used across router_utils
|
||||
cast_exception_status_to_int,
|
||||
is_advisor_orchestration_failure,
|
||||
)
|
||||
from litellm.router_utils.router_callbacks.track_deployment_metrics import (
|
||||
increment_deployment_failures_for_current_minute,
|
||||
)
|
||||
from litellm.types.router import LiteLLMParamsTypedDict
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -21,6 +32,116 @@ if TYPE_CHECKING:
|
|||
else:
|
||||
LitellmRouter = Any
|
||||
|
||||
# Status codes a generic API call's caller-supplied resource id can trigger on its own
|
||||
# (e.g. a nonexistent file/batch/thread id), independent of the selected deployment's health.
|
||||
_REQUEST_SCOPED_STATUS_CODES: Final = frozenset((404,))
|
||||
|
||||
|
||||
def _trigger_cooldown_for_failed_deployment(
|
||||
litellm_router: LitellmRouter,
|
||||
kwargs: Mapping[str, Any],
|
||||
exception: Exception,
|
||||
) -> None:
|
||||
"""
|
||||
Trigger cooldown for a failed fallback deployment.
|
||||
|
||||
In the fallback path the normal failure-callback cooldown is skipped because the
|
||||
Logging object sets has_logged_async_failure=True after the first failure and
|
||||
blocks all subsequent failure callbacks. This helper ensures every failed
|
||||
fallback deployment is evaluated for cooldown regardless.
|
||||
"""
|
||||
try:
|
||||
if is_advisor_orchestration_failure(exception):
|
||||
verbose_router_logger.debug(
|
||||
"Not triggering cooldown for fallback deployment: failure originated "
|
||||
"from advisor orchestration, not the selected deployment."
|
||||
)
|
||||
return
|
||||
|
||||
exception_status: Final[str | int] = getattr(exception, "status_code", "")
|
||||
|
||||
# Generic API calls (files, batches, threads, rerank, ...) take a caller-supplied
|
||||
# resource id, so a 404 there usually means "that id doesn't exist" rather than
|
||||
# "this deployment is unhealthy". Left unguarded, one bad id would 404 every
|
||||
# deployment in the fallback chain and cool all of them down from a single request.
|
||||
if (
|
||||
kwargs.get("original_generic_function") is not None
|
||||
and cast_exception_status_to_int(exception_status) in _REQUEST_SCOPED_STATUS_CODES
|
||||
):
|
||||
verbose_router_logger.debug(
|
||||
"Not triggering cooldown for fallback deployment: status %s on a generic API "
|
||||
"call is caller-attributable, not a deployment health signal.",
|
||||
exception_status,
|
||||
)
|
||||
return
|
||||
|
||||
# The proxy's `x-litellm-timeout` header lets a caller set an arbitrarily short
|
||||
# timeout, which litellm.Timeout reports as status 408 regardless of the deployment's
|
||||
# actual health. Left unguarded, a caller could force a 408 on every deployment in
|
||||
# the fallback chain from a single request with a near-zero timeout.
|
||||
if kwargs.get("client_side_timeout") and cast_exception_status_to_int(exception_status) == 408:
|
||||
verbose_router_logger.debug(
|
||||
"Not triggering cooldown for fallback deployment: a caller-supplied "
|
||||
"x-litellm-timeout caused this 408, not deployment health."
|
||||
)
|
||||
return
|
||||
|
||||
# Only Router._set_failed_deployment_id_on_exception()'s server-stamped id is
|
||||
# trusted here: a metadata-bucket lookup (e.g. "metadata"/"litellm_metadata")
|
||||
# can't reliably tell a caller-supplied bucket from a router-authored one
|
||||
# without knowing this call's function_name, so a client with permission to
|
||||
# set metadata could otherwise get an arbitrary deployment cooled down.
|
||||
deployment_id: Final[str | None] = getattr(exception, "failed_deployment_id", None)
|
||||
|
||||
if deployment_id is None:
|
||||
verbose_router_logger.debug("Cannot trigger cooldown for fallback: no failed_deployment_id on exception")
|
||||
return
|
||||
|
||||
# Priority: deployment config > response header > router default, matching
|
||||
# Router.deployment_callback_on_failure's precedence for the primary path.
|
||||
deployment_dict: Final = litellm_router.get_model_info(id=deployment_id)
|
||||
deployment_cooldown: Final = (
|
||||
_first_present(
|
||||
deployment_dict.get("model_info"), deployment_dict.get("litellm_params"), key="cooldown_time"
|
||||
)
|
||||
if deployment_dict is not None
|
||||
else None
|
||||
)
|
||||
exception_headers: Final = litellm.litellm_core_utils.exception_mapping_utils._get_response_headers(
|
||||
original_exception=exception
|
||||
)
|
||||
_get_retry_after: Final = (
|
||||
litellm.utils._get_retry_after_from_exception_header # pyright: ignore[reportPrivateUsage] - as router.py
|
||||
)
|
||||
header_cooldown: Final = (
|
||||
_get_retry_after(response_headers=exception_headers) if exception_headers is not None else None
|
||||
)
|
||||
time_to_cooldown: Final = (
|
||||
deployment_cooldown
|
||||
if deployment_cooldown is not None and deployment_cooldown >= 0
|
||||
else (
|
||||
header_cooldown
|
||||
if header_cooldown is not None and header_cooldown >= 0
|
||||
else litellm_router.cooldown_time
|
||||
)
|
||||
)
|
||||
|
||||
increment_deployment_failures_for_current_minute(
|
||||
litellm_router_instance=litellm_router,
|
||||
deployment_id=deployment_id,
|
||||
)
|
||||
_set_cooldown_deployments(
|
||||
litellm_router_instance=litellm_router,
|
||||
exception_status=exception_status,
|
||||
original_exception=exception,
|
||||
deployment=deployment_id,
|
||||
time_to_cooldown=time_to_cooldown,
|
||||
)
|
||||
|
||||
verbose_router_logger.debug("Triggered cooldown for fallback deployment %s", deployment_id)
|
||||
except Exception as e: # noqa: BLE001 - best-effort cooldown trigger must never break the fallback response itself
|
||||
verbose_router_logger.debug("Error triggering cooldown for fallback deployment: %s", e)
|
||||
|
||||
|
||||
def fallback_attempt_key(fallback_target: object) -> str | None:
|
||||
"""
|
||||
|
|
@ -131,6 +252,28 @@ def get_fallback_model_group(fallbacks: list[Any], model_group: str) -> tuple[li
|
|||
return fallback_model_group, generic_fallback_idx
|
||||
|
||||
|
||||
PROVIDER_SCOPED_RESOURCE_KEYS: Final = ("input_file_id", "training_file")
|
||||
|
||||
|
||||
def _get_fallback_target_model_group(fallback_entry: str | Mapping[str, object]) -> str | None:
|
||||
if isinstance(fallback_entry, str):
|
||||
return fallback_entry
|
||||
target: Final = fallback_entry.get("model")
|
||||
return target if isinstance(target, str) else None
|
||||
|
||||
|
||||
def references_provider_scoped_resource(kwargs: Mapping[str, object]) -> bool:
|
||||
"""
|
||||
True when the request names a file that only exists under one provider's credentials.
|
||||
|
||||
Batch and fine-tuning jobs are created from a file the caller already uploaded, and
|
||||
that file lives in the account of the deployment that stored it. Handing the id to a
|
||||
different model group can only fail, and the second provider's error replaces the
|
||||
error the caller actually needs to see.
|
||||
"""
|
||||
return any(kwargs.get(key) for key in PROVIDER_SCOPED_RESOURCE_KEYS)
|
||||
|
||||
|
||||
async def run_async_fallback(
|
||||
*args: tuple[Any],
|
||||
litellm_router: LitellmRouter,
|
||||
|
|
@ -176,6 +319,10 @@ async def run_async_fallback(
|
|||
|
||||
error_from_fallbacks = original_exception
|
||||
fallback_errors = (get_fallback_error_info(original_exception),)
|
||||
metadata_variable_name: Final = _get_router_metadata_variable_name(
|
||||
function_name=getattr(kwargs.get("original_function"), "__name__", None)
|
||||
)
|
||||
same_model_group_only: Final = references_provider_scoped_resource(kwargs)
|
||||
# Read out of kwargs and narrowed here rather than declared as a parameter: every caller
|
||||
# reaches this function by spreading a loosely-typed kwargs dict, so a declared parameter
|
||||
# would carry an annotation that no call site can actually be checked against.
|
||||
|
|
@ -188,6 +335,13 @@ async def run_async_fallback(
|
|||
for mg in fallback_model_group:
|
||||
if mg == original_model_group:
|
||||
continue
|
||||
if same_model_group_only and _get_fallback_target_model_group(mg) != original_model_group:
|
||||
verbose_router_logger.info(
|
||||
"Skipping fallback to model_group = %s: request is pinned to model_group = %s by its uploaded file",
|
||||
mask_sensitive_structure(mg),
|
||||
original_model_group,
|
||||
)
|
||||
continue
|
||||
attempt_key = fallback_attempt_key(mg)
|
||||
if attempt_key is not None:
|
||||
if attempt_key in attempted:
|
||||
|
|
@ -205,9 +359,10 @@ async def run_async_fallback(
|
|||
kwargs["model"] = mg
|
||||
elif isinstance(mg, dict):
|
||||
kwargs.update(mg)
|
||||
kwargs.setdefault("metadata", {}).update(
|
||||
{"model_group": kwargs.get("model", None)}
|
||||
) # update model_group used, if fallbacks are done
|
||||
kwargs[metadata_variable_name] = {
|
||||
**(kwargs.get(metadata_variable_name) or {}),
|
||||
"model_group": kwargs.get("model", None),
|
||||
}
|
||||
fallback_depth = fallback_depth + 1
|
||||
kwargs["fallback_depth"] = fallback_depth
|
||||
kwargs["max_fallbacks"] = max_fallbacks
|
||||
|
|
@ -236,6 +391,13 @@ async def run_async_fallback(
|
|||
kwargs=kwargs,
|
||||
original_exception=original_exception,
|
||||
)
|
||||
logging_obj = kwargs.get("litellm_logging_obj")
|
||||
if logging_obj is not None and logging_obj.model_call_details.get("has_logged_async_failure", False):
|
||||
_trigger_cooldown_for_failed_deployment(
|
||||
litellm_router=litellm_router,
|
||||
kwargs=kwargs,
|
||||
exception=e,
|
||||
)
|
||||
raise error_from_fallbacks
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -3,7 +3,6 @@ Pydantic models for Memory management endpoints.
|
|||
"""
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
|
@ -12,7 +11,7 @@ class LiteLLM_MemoryRow(BaseModel):
|
|||
memory_id: str
|
||||
key: str
|
||||
value: str
|
||||
metadata: Any | None = None
|
||||
metadata: object | None = None
|
||||
user_id: str | None = None
|
||||
team_id: str | None = None
|
||||
created_at: datetime | None = None
|
||||
|
|
@ -24,7 +23,7 @@ class LiteLLM_MemoryRow(BaseModel):
|
|||
class MemoryCreateRequest(BaseModel):
|
||||
key: str = Field(..., description="Memory key (acts as the namespace in the URL).")
|
||||
value: str = Field(..., description="Memory content. Typically markdown/text for LLM context.")
|
||||
metadata: Any | None = Field(
|
||||
metadata: object | None = Field(
|
||||
default=None,
|
||||
description="Optional JSON metadata (tags, structured fields).",
|
||||
)
|
||||
|
|
@ -40,7 +39,7 @@ class MemoryCreateRequest(BaseModel):
|
|||
|
||||
class MemoryUpdateRequest(BaseModel):
|
||||
value: str | None = None
|
||||
metadata: Any | None = None
|
||||
metadata: object | None = None
|
||||
# Only honored on create (when the row doesn't yet exist) and only for
|
||||
# PROXY_ADMIN callers — mirrors MemoryCreateRequest so admins can bootstrap
|
||||
# rows scoped to another user/team via PUT, not just POST.
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ class GroupByDimension(str, Enum):
|
|||
|
||||
class SpendMetrics(BaseModel):
|
||||
spend: float = Field(default=0.0)
|
||||
flat_cost: float = Field(default=0.0)
|
||||
prompt_tokens: int = Field(default=0)
|
||||
completion_tokens: int = Field(default=0)
|
||||
cache_read_input_tokens: int = Field(default=0)
|
||||
|
|
@ -75,6 +76,7 @@ class DailySpendData(BaseModel):
|
|||
|
||||
class DailySpendMetadata(BaseModel):
|
||||
total_spend: float = Field(default=0.0)
|
||||
total_flat_cost: float = Field(default=0.0)
|
||||
total_prompt_tokens: int = Field(default=0)
|
||||
total_completion_tokens: int = Field(default=0)
|
||||
total_tokens: int = Field(default=0)
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ litellm.Router Types - includes RouterConfig, UpdateRouterConfig, ModelInfo etc
|
|||
import datetime
|
||||
import enum
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Final, Generic, Literal, TypeVar, get_type_hints
|
||||
from typing import Any, ClassVar, Final, Generic, Literal, TypeVar, get_type_hints
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
|
|
@ -127,6 +127,14 @@ class UpdateRouterConfig(BaseModel):
|
|||
model_config = ConfigDict(protected_namespaces=())
|
||||
|
||||
|
||||
def _as_utc(value: datetime.datetime | None) -> datetime.datetime | None:
|
||||
if value is None:
|
||||
return None
|
||||
if value.tzinfo is None:
|
||||
return value.replace(tzinfo=datetime.timezone.utc)
|
||||
return value.astimezone(datetime.timezone.utc)
|
||||
|
||||
|
||||
class ModelInfo(MirroredPricingParams):
|
||||
id: str | None # Allow id to be optional on input, but it will always be present as a str in the model instance
|
||||
db_model: bool = False # used for proxy - to separate models which are stored in the db vs. config.
|
||||
|
|
@ -151,6 +159,17 @@ class ModelInfo(MirroredPricingParams):
|
|||
# admin-toggled pause flag; mirrors LiteLLM_ProxyModelTable.blocked
|
||||
blocked: bool | None = None
|
||||
|
||||
# Bounds live on the model rather than litellm.constants: names there reach
|
||||
# litellm/__init__ through several modules' star re-exports, and a Final rebound that
|
||||
# way trips the basedpyright gate.
|
||||
MAX_PTU_COUNT: ClassVar[int] = 1_000_000
|
||||
MAX_COST_PER_PTU_PER_HOUR: ClassVar[float] = 1_000_000.0
|
||||
|
||||
ptu_count: int | None = None
|
||||
cost_per_ptu_per_hour: float | None = None
|
||||
ptu_effective_from: datetime.datetime | None = None
|
||||
ptu_effective_to: datetime.datetime | None = None
|
||||
|
||||
def __init__(self, id: str | int | None = None, **params) -> None:
|
||||
if id is None:
|
||||
id = str(uuid.uuid4()) # Generate a UUID if id is None or not provided
|
||||
|
|
@ -158,6 +177,23 @@ class ModelInfo(MirroredPricingParams):
|
|||
id = str(id)
|
||||
super().__init__(id=id, **params)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_ptu_bounds(self) -> "ModelInfo":
|
||||
if self.ptu_count is not None and not 0 < self.ptu_count <= self.MAX_PTU_COUNT:
|
||||
raise ValueError(f"ptu_count must be a positive integer no greater than {self.MAX_PTU_COUNT}")
|
||||
if (
|
||||
self.cost_per_ptu_per_hour is not None
|
||||
and not 0 <= self.cost_per_ptu_per_hour <= self.MAX_COST_PER_PTU_PER_HOUR
|
||||
):
|
||||
raise ValueError(
|
||||
f"cost_per_ptu_per_hour must be a finite number between 0 and {self.MAX_COST_PER_PTU_PER_HOUR}"
|
||||
)
|
||||
start: Final = _as_utc(self.ptu_effective_from)
|
||||
end: Final = _as_utc(self.ptu_effective_to)
|
||||
if start is not None and end is not None and end <= start:
|
||||
raise ValueError("ptu_effective_to must be after ptu_effective_from")
|
||||
return self
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
def __contains__(self, key) -> bool:
|
||||
|
|
@ -245,6 +281,12 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
|
|||
# Deployment budgets
|
||||
max_budget: float | None = None
|
||||
budget_duration: str | None = None
|
||||
keepalive_seconds: float | None = None
|
||||
# keepalive_seconds is operator-only by default: a client's request-level
|
||||
# value is ignored unless the deployment opts in here. Prevents a client
|
||||
# from unilaterally enabling heartbeats (and the LB-idle-timeout evasion
|
||||
# that comes with them) for a deployment that never configured them.
|
||||
allow_client_keepalive_override: bool | None = False
|
||||
use_in_pass_through: bool | None = False
|
||||
use_litellm_proxy: bool | None = False
|
||||
use_chat_completions_api: bool | None = None
|
||||
|
|
@ -421,6 +463,11 @@ class LiteLLMParamsTypedDict(TypedDict, total=False):
|
|||
# deployment budgets
|
||||
max_budget: float | None
|
||||
budget_duration: str | None
|
||||
keepalive_seconds: float | None
|
||||
allow_client_keepalive_override: bool | None
|
||||
|
||||
# per-deployment cooldown override
|
||||
cooldown_time: float | None
|
||||
|
||||
|
||||
class DeploymentTypedDict(TypedDict, total=False):
|
||||
|
|
@ -513,6 +560,9 @@ class AllowedFailsPolicy(BaseModel):
|
|||
RateLimitErrorAllowedFails: int | None = None
|
||||
ContentPolicyViolationErrorAllowedFails: int | None = None
|
||||
InternalServerErrorAllowedFails: int | None = None
|
||||
ServiceUnavailableErrorAllowedFails: int | None = None
|
||||
BadGatewayErrorAllowedFails: int | None = None
|
||||
NotFoundErrorAllowedFails: int | None = None
|
||||
|
||||
|
||||
class AlertingConfig(BaseModel):
|
||||
|
|
|
|||
|
|
@ -3129,6 +3129,7 @@ class StandardAuditLogPayload(TypedDict):
|
|||
class StandardLoggingPayload(TypedDict):
|
||||
id: str
|
||||
trace_id: str # Trace multiple LLM calls belonging to same overall request (e.g. fallbacks/retries)
|
||||
session_id: str # End-user/conversation session id (litellm_session_id), independent of trace_id
|
||||
litellm_call_id: str | None # UUID returned in x-litellm-call-id response header
|
||||
call_type: str
|
||||
stream: bool | None
|
||||
|
|
@ -3384,6 +3385,8 @@ agentic_loop_internal_litellm_params: Final = [
|
|||
"_code_interpreter_interception_sandbox_key",
|
||||
"_code_interpreter_interception_session_scoped",
|
||||
"_code_interpreter_interception_converted_stream",
|
||||
"_websearch_interception_emit_native_blocks",
|
||||
"_websearch_interception_converted_stream",
|
||||
]
|
||||
|
||||
# Proxy-owned callback credentials, stamped from admin-configured team/key callback
|
||||
|
|
@ -3398,6 +3401,8 @@ all_litellm_params = (
|
|||
+ [
|
||||
"metadata",
|
||||
"litellm_metadata",
|
||||
"keepalive_seconds",
|
||||
"allow_client_keepalive_override",
|
||||
"litellm_trace_id",
|
||||
"litellm_request_debug",
|
||||
"guardrails",
|
||||
|
|
|
|||
|
|
@ -277,6 +277,11 @@ class LiteLLM_ManagedVectorStoreIndex(BaseModel):
|
|||
updated_by: str | None = None
|
||||
|
||||
|
||||
class IndexListResponse(BaseModel):
|
||||
object: Literal["list"] = "list"
|
||||
data: tuple[LiteLLM_ManagedVectorStoreIndex, ...]
|
||||
|
||||
|
||||
class VectorStoreIndexType(str, Enum):
|
||||
"""Type of vector store index"""
|
||||
|
||||
|
|
|
|||
104
litellm/utils.py
104
litellm/utils.py
|
|
@ -711,14 +711,71 @@ def _remove_thought_signatures_from_messages(messages: list, thought_signature_s
|
|||
return processed_messages
|
||||
|
||||
|
||||
def _restore_correlation_context_if_supported(logging_obj: object) -> None:
|
||||
"""Call logging_obj._restore_correlation_context() if it's actually there.
|
||||
|
||||
Some call sites (tests, narrow unit paths) inject a minimal stand-in
|
||||
object as litellm_logging_obj instead of a real Logging instance - this
|
||||
method is new plumbing specific to request_correlation_in_logs, not part
|
||||
of any pre-existing stand-in's expected interface. `object` (not `Any`)
|
||||
is deliberate: the getattr() below is exactly how this stays type-safe
|
||||
while still tolerating a stand-in that lacks the method.
|
||||
"""
|
||||
restore: Final = getattr(logging_obj, "_restore_correlation_context", None)
|
||||
if restore is not None:
|
||||
restore()
|
||||
|
||||
|
||||
def _is_streaming_response_for_correlation(result: object) -> bool:
|
||||
"""True if `result` is a lazy stream wrapper rather than an already-complete response.
|
||||
|
||||
Only wrapper_async() consults this - it must NOT restore the originating
|
||||
Task's trace_id/session_id as soon as a streaming call returns this: the
|
||||
caller is about to iterate it over however many subsequent lines of their
|
||||
own code, and those log lines should still show this call's ids, not the
|
||||
pre-call ones. This is safe specifically because each async call already
|
||||
runs in its own asyncio Task with its own copy of the contextvars, so
|
||||
leaving it "open" can only affect that one Task, never a different,
|
||||
unrelated future request - Tasks, unlike a thread pool's worker threads,
|
||||
are never recycled across requests. The corresponding terminal handler
|
||||
(async_success_handler, dispatched once the full stream is actually
|
||||
assembled) is what restores it once streaming genuinely finishes.
|
||||
|
||||
wrapper() (the sync path) does NOT consult this at all: sync calls pass
|
||||
supports_correlation_logging=False into function_setup()/Logging(), so
|
||||
they never stamp trace_id/session_id in the first place - a plain OS
|
||||
thread has no per-call isolation the way an asyncio Task does, and a
|
||||
thread pool's worker threads *are* recycled across unrelated requests, so
|
||||
stamping ids there without a safe restore mechanism could permanently
|
||||
misattribute a later, unrelated request's logs. Full sync support is
|
||||
deferred to a follow-up PR with its own restore mechanism; see
|
||||
Logging.__init__'s supports_correlation_logging parameter.
|
||||
|
||||
Genuinely circular otherwise: utils.py -> streaming_handler.py ->
|
||||
redact_messages.py -> llms/vertex_ai/common_utils.py -> utils.py, which
|
||||
needs names (supports_response_schema, etc.) this module hasn't finished
|
||||
defining yet at that point in its own top-to-bottom execution.
|
||||
"""
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
return isinstance(result, CustomStreamWrapper)
|
||||
|
||||
|
||||
# Runs once per call to check if the user wants to send their data anywhere - PostHog/Sentry/Slack/etc.
|
||||
def function_setup(
|
||||
original_function: str, rules_obj, start_time, *args, **kwargs
|
||||
): # just run once to check if user wants to send their data anywhere - PostHog/Sentry/Slack/etc.
|
||||
original_function: str,
|
||||
rules_obj: Rules,
|
||||
start_time: datetime.datetime,
|
||||
*args: Any, # positional passthrough to the wrapped LLM call (ANN401 ignored, see ruff-strict.toml)
|
||||
is_async_call: bool = True,
|
||||
**kwargs: Any, # kwargs-ok: forwarded to Logging()/callbacks, varies per call_type
|
||||
) -> tuple[LiteLLMLoggingObject, dict[str, Any]]:
|
||||
### NOTICES ###
|
||||
if litellm.set_verbose is True:
|
||||
verbose_logger.warning(
|
||||
"`litellm.set_verbose` is deprecated. Please set `os.environ['LITELLM_LOG'] = 'DEBUG'` for debug logs."
|
||||
)
|
||||
logging_obj: LiteLLMLoggingObject | None = None # rebind-ok: set to the real object further down on success
|
||||
try:
|
||||
global callback_list, add_breadcrumb, user_logger_fn, Logging
|
||||
|
||||
|
|
@ -1001,7 +1058,8 @@ def function_setup(
|
|||
):
|
||||
stream = True
|
||||
get_litellm_logging_class: Final = getattr(sys.modules[__name__], "get_litellm_logging_class")
|
||||
logging_obj: Final = get_litellm_logging_class()( # Victim for object pool
|
||||
# Victim for object pool
|
||||
logging_obj = get_litellm_logging_class()( # rebind-ok: 2nd assignment to logging_obj (see initial None above)
|
||||
model=model,
|
||||
messages=messages,
|
||||
stream=stream,
|
||||
|
|
@ -1016,6 +1074,7 @@ def function_setup(
|
|||
dynamic_async_failure_callbacks=dynamic_async_failure_callbacks,
|
||||
kwargs=kwargs,
|
||||
applied_guardrails=applied_guardrails,
|
||||
supports_correlation_logging=is_async_call,
|
||||
)
|
||||
|
||||
## check if metadata is passed in
|
||||
|
|
@ -1040,6 +1099,15 @@ def function_setup(
|
|||
)
|
||||
return logging_obj, kwargs
|
||||
except Exception as e:
|
||||
# If Logging() was constructed above before this failed, its __init__ already
|
||||
# mutated trace_id_var/session_id_var - restore them *before* logging the
|
||||
# exception below, since we're about to raise without ever returning
|
||||
# logging_obj to the caller's wrapper()/wrapper_async() (which would
|
||||
# otherwise be the one doing this restore). Restoring first means this
|
||||
# diagnostic log line itself doesn't get stamped with a call's ids when
|
||||
# that call never actually produced a usable logging object.
|
||||
if logging_obj is not None:
|
||||
_restore_correlation_context_if_supported(logging_obj)
|
||||
verbose_logger.exception("litellm.utils.py::function_setup() - [Non-Blocking] Error in function_setup")
|
||||
raise e
|
||||
|
||||
|
|
@ -1296,7 +1364,9 @@ def client(original_function):
|
|||
|
||||
try:
|
||||
if logging_obj is None:
|
||||
logging_obj, kwargs = function_setup(original_function.__name__, rules_obj, start_time, *args, **kwargs)
|
||||
logging_obj, kwargs = function_setup(
|
||||
original_function.__name__, rules_obj, start_time, *args, is_async_call=False, **kwargs
|
||||
)
|
||||
|
||||
# Type assertion: logging_obj is guaranteed to be non-None after function_setup
|
||||
assert logging_obj is not None, "logging_obj should not be None after function_setup"
|
||||
|
|
@ -1807,9 +1877,11 @@ def client(original_function):
|
|||
kwargs["retry_strategy"] = "exponential_backoff_retry"
|
||||
elif isinstance(e, openai.APIError): # generic api error
|
||||
kwargs["retry_strategy"] = "constant_retry"
|
||||
return await litellm.acompletion_with_retries(*args, **kwargs)
|
||||
result = await litellm.acompletion_with_retries(*args, **kwargs)
|
||||
except Exception:
|
||||
pass
|
||||
else:
|
||||
return result
|
||||
elif (
|
||||
isinstance(e, litellm.exceptions.ContextWindowExceededError)
|
||||
and context_window_fallback_dict
|
||||
|
|
@ -1820,7 +1892,8 @@ def client(original_function):
|
|||
args[0] = context_window_fallback_dict[model]
|
||||
else:
|
||||
kwargs["model"] = context_window_fallback_dict[model]
|
||||
return await original_function(*args, **kwargs)
|
||||
result = await original_function(*args, **kwargs)
|
||||
return result
|
||||
elif call_type == CallTypes.aresponses.value:
|
||||
_is_litellm_router_call = "model_group" in (
|
||||
kwargs.get("metadata") or {}
|
||||
|
|
@ -1837,9 +1910,11 @@ def client(original_function):
|
|||
kwargs["retry_strategy"] = "exponential_backoff_retry"
|
||||
elif isinstance(e, openai.APIError): # generic api error
|
||||
kwargs["retry_strategy"] = "constant_retry"
|
||||
return await litellm.aresponses_with_retries(*args, **kwargs)
|
||||
result = await litellm.aresponses_with_retries(*args, **kwargs)
|
||||
except Exception:
|
||||
pass
|
||||
else:
|
||||
return result
|
||||
|
||||
deployment_num_retries: Final = kwargs.get("num_retries")
|
||||
if deployment_num_retries is not None:
|
||||
|
|
@ -1849,6 +1924,21 @@ def client(original_function):
|
|||
setattr(e, "timeout", timeout)
|
||||
raise e
|
||||
|
||||
finally:
|
||||
# Restore trace_id/session_id contextvars to their pre-call value once
|
||||
# this call (in this asyncio Task) is fully done - see
|
||||
# request_correlation_in_logs. Unlike wrapper()'s sync path, it's safe to
|
||||
# skip restoring when returning a stream: each async call already runs in
|
||||
# its own Task with its own copy of the contextvars (asyncio.create_task
|
||||
# copies context at creation), so leaving this Task's own view "open"
|
||||
# while the caller iterates the stream can only affect that one Task -
|
||||
# never a different, unrelated future request, since Tasks (unlike a
|
||||
# thread pool's worker threads) are never recycled across requests. The
|
||||
# corresponding terminal handler (async_success_handler) restores it once
|
||||
# streaming genuinely finishes; aclose()/__del__ cover early termination.
|
||||
if not _is_streaming_response_for_correlation(result):
|
||||
_restore_correlation_context_if_supported(logging_obj)
|
||||
|
||||
get_coroutine_checker: Final = getattr(sys.modules[__name__], "get_coroutine_checker")
|
||||
is_coroutine: Final = get_coroutine_checker().is_async_callable(original_function)
|
||||
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -1,13 +1,3 @@
|
|||
[[IgnoredVulns]]
|
||||
id = "GHSA-fwg2-594c-jp42"
|
||||
ignoreUntil = 2026-08-12
|
||||
reason = "pypdf 6.15.0 (the fix) published 2026-08-06 and is still inside the P3D exclude-newer window, so uv cannot lock it yet; bump and drop this entry from 2026-08-09"
|
||||
|
||||
[[IgnoredVulns]]
|
||||
id = "GHSA-fp3f-mc75-235c"
|
||||
ignoreUntil = 2026-08-12
|
||||
reason = "second pypdf advisory with the same 6.15.0 fix, published 2026-08-07 after the first; drop alongside GHSA-fwg2-594c-jp42 in the same bump"
|
||||
|
||||
[[IgnoredVulns]]
|
||||
id = "GHSA-w8v5-vhqr-4h9v"
|
||||
ignoreUntil = 2026-09-09
|
||||
|
|
|
|||
|
|
@ -1,18 +1,18 @@
|
|||
{
|
||||
"ANN001": {
|
||||
"limit": 3114
|
||||
"limit": 3106
|
||||
},
|
||||
"ANN002": {
|
||||
"limit": 71
|
||||
},
|
||||
"ANN003": {
|
||||
"limit": 834
|
||||
"limit": 832
|
||||
},
|
||||
"ANN201": {
|
||||
"limit": 2031
|
||||
"limit": 2023
|
||||
},
|
||||
"ANN202": {
|
||||
"limit": 865
|
||||
"limit": 860
|
||||
},
|
||||
"ANN204": {
|
||||
"limit": 713
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 133
|
||||
},
|
||||
"ANN401": {
|
||||
"limit": 1555
|
||||
"limit": 1495
|
||||
},
|
||||
"ASYNC230": {
|
||||
"limit": 11
|
||||
|
|
@ -39,7 +39,7 @@
|
|||
"limit": 505
|
||||
},
|
||||
"B009": {
|
||||
"limit": 81
|
||||
"limit": 79
|
||||
},
|
||||
"B010": {
|
||||
"limit": 190
|
||||
|
|
@ -78,7 +78,7 @@
|
|||
"limit": 1
|
||||
},
|
||||
"C901": {
|
||||
"limit": 314
|
||||
"limit": 313
|
||||
},
|
||||
"D419": {
|
||||
"limit": 6
|
||||
|
|
@ -234,16 +234,16 @@
|
|||
"limit": 5
|
||||
},
|
||||
"TID251": {
|
||||
"limit": 1238
|
||||
"limit": 1226
|
||||
},
|
||||
"TRY002": {
|
||||
"limit": 528
|
||||
"limit": 524
|
||||
},
|
||||
"TRY004": {
|
||||
"limit": 96
|
||||
},
|
||||
"TRY201": {
|
||||
"limit": 407
|
||||
"limit": 405
|
||||
},
|
||||
"TRY203": {
|
||||
"limit": 113
|
||||
|
|
|
|||
|
|
@ -16,6 +16,17 @@ external = [
|
|||
"PLC0415", "E402", "BLE001", "ARG002", "S102", "S324", "S606", "D401", "F403", "F405",
|
||||
]
|
||||
|
||||
[lint.per-file-ignores]
|
||||
# ANN401 (explicit `Any` disallowed) has no per-line/function-level ignore mechanism
|
||||
# in ruff, only file-level. These two files each have a handful of parameters that
|
||||
# are genuinely heterogeneous with no fitting concrete type: a response object that
|
||||
# varies across every LLM call type (completion/embedding/transcription/etc. each
|
||||
# return a different shape), and *args/**kwargs forwarded verbatim with no fixed
|
||||
# shape. Tried the closest existing union (CostResponseTypes) first; basedpyright
|
||||
# caught a real mismatch, confirming Any is correct here, not a shortcut.
|
||||
"litellm/litellm_core_utils/litellm_logging.py" = ["ANN401"]
|
||||
"litellm/utils.py" = ["ANN401"]
|
||||
|
||||
[lint.mccabe]
|
||||
max-complexity = 15
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -893,6 +893,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 +986,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?
|
||||
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue