mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge remote-tracking branch 'upstream/litellm_internal_staging' into deepkeep-as-internal
# Conflicts: # enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py
This commit is contained in:
commit
c71d7b4077
381 changed files with 30406 additions and 4088 deletions
|
|
@ -1610,14 +1610,14 @@ jobs:
|
|||
- run:
|
||||
name: Run helm lint
|
||||
command: |
|
||||
helm lint ./deploy/charts/litellm-helm
|
||||
helm lint ./helm/litellm-helm
|
||||
|
||||
# Run helm tests
|
||||
- run:
|
||||
name: Run helm tests
|
||||
command: |
|
||||
IMAGE_TAG=${CIRCLE_SHA1:-ci}
|
||||
helm install litellm ./deploy/charts/litellm-helm -f ./deploy/charts/litellm-helm/ci/test-values.yaml \
|
||||
helm install litellm ./helm/litellm-helm -f ./helm/litellm-helm/ci/test-values.yaml \
|
||||
--set image.repository=litellm-ci \
|
||||
--set image.tag=${IMAGE_TAG} \
|
||||
--set image.pullPolicy=Never
|
||||
|
|
|
|||
3
.github/pull_request_template.md
vendored
3
.github/pull_request_template.md
vendored
|
|
@ -13,7 +13,7 @@
|
|||
- [ ] I have added meaningful tests
|
||||
- [ ] My PR passes all CI/CD checks (e.g., lint, format, unit tests)
|
||||
- [ ] My PR's scope is as isolated as possible; it only solves 1 specific problem
|
||||
- [ ] I have requested a Greptile review by commenting `@greptileai` and received a **Confidence Score of at least 4/5** before requesting a maintainer review
|
||||
- [ ] I have received a Greptile **Confidence Score of at least 4/5** before requesting a maintainer review (Greptile reviews automatically once the PR is opened; only comment `@greptileai` to re-request a review after pushing changes)
|
||||
|
||||
## Delays in PR merge?
|
||||
|
||||
|
|
@ -24,6 +24,7 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac
|
|||
<!-- Include screenshots, screen recordings, or command (e.g., curl) + output demonstrating that your changes work as expected
|
||||
The proof must be completely e2e with no mocks, using, for example, actual LLM calls costing real $. `pytest` commands are not enough
|
||||
For bug fixes: show reproduction before the fix and passing behavior after
|
||||
Include the commit hash each proof was captured at, for both the before and the after runs
|
||||
For new features: show the feature working end-to-end
|
||||
For UI changes: include before/after screenshots -->
|
||||
|
||||
|
|
|
|||
4
.github/workflows/codspeed.yml
vendored
4
.github/workflows/codspeed.yml
vendored
|
|
@ -4,9 +4,11 @@ on:
|
|||
push:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
# Allow CodSpeed to trigger backtest performance analysis
|
||||
# in order to generate initial data
|
||||
workflow_dispatch:
|
||||
|
|
@ -22,7 +24,7 @@ concurrency:
|
|||
jobs:
|
||||
benchmarks:
|
||||
runs-on: ubuntu-24.04
|
||||
timeout-minutes: 15
|
||||
timeout-minutes: 60
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
|
|
|
|||
4
.github/workflows/helm_unit_test.yml
vendored
4
.github/workflows/helm_unit_test.yml
vendored
|
|
@ -38,4 +38,6 @@ jobs:
|
|||
echo "Helm unittest plugin integrity verified: $ACTUAL_SHA"
|
||||
|
||||
- name: Run unit tests
|
||||
run: helm unittest -f 'tests/*.yaml' deploy/charts/litellm-helm
|
||||
run: |
|
||||
helm unittest -f 'tests/*.yaml' helm/litellm-helm
|
||||
helm unittest -f 'tests/*.yaml' helm/litellm
|
||||
|
|
|
|||
113
.github/workflows/test-terraform-provider.yml
vendored
Normal file
113
.github/workflows/test-terraform-provider.yml
vendored
Normal file
|
|
@ -0,0 +1,113 @@
|
|||
name: Terraform Provider
|
||||
|
||||
on:
|
||||
push:
|
||||
paths:
|
||||
- "terraform/provider/**"
|
||||
- ".github/workflows/test-terraform-provider.yml"
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths:
|
||||
- "terraform/provider/**"
|
||||
- "litellm/proxy/**"
|
||||
- ".github/workflows/test-terraform-provider.yml"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
provider-checks:
|
||||
name: gofmt, vet, build, test
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
defaults:
|
||||
run:
|
||||
working-directory: terraform/provider
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- uses: actions/setup-go@7a3fe6cf4cb3a834922a1244abfce67bcef6a0c5 # v6.2.0
|
||||
with:
|
||||
go-version-file: terraform/provider/go.mod
|
||||
cache: true
|
||||
cache-dependency-path: terraform/provider/go.sum
|
||||
|
||||
- name: gofmt
|
||||
run: |
|
||||
UNFORMATTED=$(gofmt -l .)
|
||||
if [ -n "${UNFORMATTED}" ]; then
|
||||
echo "::error::gofmt required for: ${UNFORMATTED}"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
- name: go vet
|
||||
run: go vet ./...
|
||||
|
||||
- name: Build
|
||||
run: go build ./...
|
||||
|
||||
- name: Test
|
||||
run: go test -timeout 120s ./...
|
||||
|
||||
endpoint-drift:
|
||||
name: Provider endpoints vs proxy OpenAPI schema
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 20
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
- name: Cache uv dependencies
|
||||
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
|
||||
with:
|
||||
path: |
|
||||
~/.cache/uv
|
||||
.venv
|
||||
key: ${{ runner.os }}-uv-${{ hashFiles('uv.lock') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-uv-
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
|
||||
|
||||
- name: Generate Prisma client
|
||||
env:
|
||||
PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache
|
||||
run: |
|
||||
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
|
||||
|
||||
- name: Generate proxy OpenAPI schema
|
||||
run: |
|
||||
uv run --no-sync python terraform/provider/tools/dump_openapi.py "${RUNNER_TEMP}/openapi.json"
|
||||
|
||||
- uses: actions/setup-go@7a3fe6cf4cb3a834922a1244abfce67bcef6a0c5 # v6.2.0
|
||||
with:
|
||||
go-version-file: terraform/provider/go.mod
|
||||
cache: true
|
||||
cache-dependency-path: terraform/provider/go.sum
|
||||
|
||||
- name: Audit provider endpoints against the schema
|
||||
working-directory: terraform/provider
|
||||
run: go run ./tools/endpointaudit -provider-dir ./litellm -spec "${RUNNER_TEMP}/openapi.json"
|
||||
5
.gitignore
vendored
5
.gitignore
vendored
|
|
@ -52,9 +52,8 @@ ui/litellm-dashboard/node_modules
|
|||
ui/litellm-dashboard/next-env.d.ts
|
||||
ui/litellm-dashboard/package.json
|
||||
ui/litellm-dashboard/package-lock.json
|
||||
deploy/charts/litellm/*.tgz
|
||||
deploy/charts/litellm/charts/*
|
||||
deploy/charts/*.tgz
|
||||
helm/litellm-helm/*.tgz
|
||||
helm/*.tgz
|
||||
litellm/proxy/vertex_key.json
|
||||
**/.vim/
|
||||
**/node_modules
|
||||
|
|
|
|||
|
|
@ -21,11 +21,11 @@ End-to-end tests belong in `tests/e2e/` and must follow the harness conventions
|
|||
|
||||
When creating PRs, don't set base to `main`. `litellm_internal_staging` serves that purpose
|
||||
|
||||
When writing a PR body, treat the comments and imperative instructions inside @.github/pull_request_template.md as rules to follow, not just layout
|
||||
When writing a PR body, treat the comments and imperative instructions inside @.github/pull_request_template.md as rules to follow, not just layout. Agent harnesses may strip HTML comments from copies of that file injected into context, so read .github/pull_request_template.md from disk before writing a PR body to make sure you see every comment rule
|
||||
|
||||
If you're resolving a linear ticket, in the "## Linear ticket" section of the PR, say "Resolves LIT-1234", replacing "LIT-1234" with the actual ticket id that you're resolving. If you don't have the ticket id, don't make one up or search for it. Just leave the section blank
|
||||
|
||||
Never use `pytest` commands or the like as "Screenshots / Proof of Fix". We prefer curl'ing a live proxy instance running on localhost:4000 (I like to run it with `python litellm/proxy/proxy_cli.py --config litellm/proxy/dev_config.yaml --detailed_debug --reload --use_v2_migration_resolver 2>&1 | tee litellm.log`) and showing both the command run and the output. Also, it should hit real LLM provider APIs, not mocks, and cost real $$$ because that is the most realistic test. The proof of fix should be exactly what the end user / customer would see / do. The run logs in PR #27703 is a prime example of how to do it (not a huge fan of using a python test script that future me and the team will have no visibility into; I prefer just curl commands or a short list of bash commands (e.g., using `for`)). If it's a UI thing, just tell me which URLs to go to (e.g., http://localhost:4000/ui/?page=logs), where to click, what fields to fill out, etc. along with the other commands to run in an ordered list, and I'll do it myself and post the screenshots after you make the PR
|
||||
Never use `pytest` commands or the like as "Screenshots / Proof of Fix". We prefer curl'ing a live proxy instance running on localhost:4000 (I like to run it with `python litellm/proxy/proxy_cli.py --config litellm/proxy/dev_config.yaml --detailed_debug --reload --use_v2_migration_resolver 2>&1 | tee litellm.log`; the Admin UI dev server is `npm run dev` in `ui/litellm-dashboard`, served on port 3000) and showing both the command run and the output. Also, it should hit real LLM provider APIs, not mocks, and cost real $$$ because that is the most realistic test. The proof of fix should be exactly what the end user / customer would see / do. The run logs in PR #27703 is a prime example of how to do it (not a huge fan of using a python test script that future me and the team will have no visibility into; I prefer just curl commands or a short list of bash commands (e.g., using `for`)). If it's a UI thing, just tell me which URLs to go to (e.g., http://localhost:4000/ui/?page=logs), where to click, what fields to fill out, etc. along with the other commands to run in an ordered list, and I'll do it myself and post the screenshots after you make the PR
|
||||
|
||||
If you ever make public-facing PR descriptions, comments, issues, commit messages, etc., always follow these guidelines to sound less AI-y:
|
||||
- don't use emojis
|
||||
|
|
|
|||
|
|
@ -1,10 +1,10 @@
|
|||
# syntax=docker/dockerfile:1.7
|
||||
|
||||
# Base image for building
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f
|
||||
|
||||
# Runtime image
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f
|
||||
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
|
||||
# Pinned by digest like the other base images; bump explicitly on Node upgrades.
|
||||
ARG UI_BUILD_IMAGE=node:20.18-alpine3.20@sha256:3488b10bf958af7125a176419d2d8a9937d895bf124012aae811651988d2ffe6
|
||||
|
|
|
|||
2
Makefile
2
Makefile
|
|
@ -268,7 +268,7 @@ test-integration: install-test-deps
|
|||
$(UV_RUN) pytest tests/ -k "not test_litellm"
|
||||
|
||||
test-unit-helm: install-helm-unittest
|
||||
helm unittest -f 'tests/*.yaml' deploy/charts/litellm-helm
|
||||
helm unittest -f 'tests/*.yaml' helm/litellm-helm
|
||||
|
||||
# LLM Translation testing targets
|
||||
test-llm-translation: install-test-deps
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f
|
||||
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
|
||||
|
||||
FROM $UV_IMAGE AS uvbin
|
||||
|
|
|
|||
|
|
@ -99,7 +99,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 45895
|
||||
"limit": 45894
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 113
|
||||
|
|
|
|||
10
codecov.yaml
10
codecov.yaml
|
|
@ -15,6 +15,16 @@ ignore:
|
|||
flag_management:
|
||||
default_rules:
|
||||
carryforward: true
|
||||
# Dead flags no CI job uploads anymore: their carried-forward sessions were
|
||||
# measured against old revisions, and the stale line maps mark comment lines
|
||||
# of since-edited files as missed, sinking patch coverage on unrelated PRs.
|
||||
individual_flags:
|
||||
- name: proxy-mgmt-behavior
|
||||
carryforward: false
|
||||
- name: security
|
||||
carryforward: false
|
||||
- name: proxy-db-schema-migration
|
||||
carryforward: false
|
||||
|
||||
component_management:
|
||||
individual_components:
|
||||
|
|
|
|||
Binary file not shown.
|
|
@ -1,15 +0,0 @@
|
|||
{
|
||||
"$schema": "https://schema.management.azure.com/schemas/0.1.2-preview/CreateUIDefinition.MultiVm.json#",
|
||||
"handler": "Microsoft.Azure.CreateUIDef",
|
||||
"version": "0.1.2-preview",
|
||||
"parameters": {
|
||||
"config": {
|
||||
"isWizard": false,
|
||||
"basics": { }
|
||||
},
|
||||
"basics": [ ],
|
||||
"steps": [ ],
|
||||
"outputs": { },
|
||||
"resourceTypes": [ ]
|
||||
}
|
||||
}
|
||||
|
|
@ -1,63 +0,0 @@
|
|||
{
|
||||
"$schema": "https://schema.management.azure.com/schemas/2019-04-01/deploymentTemplate.json#",
|
||||
"contentVersion": "1.0.0.0",
|
||||
"parameters": {
|
||||
"imageName": {
|
||||
"type": "string",
|
||||
"defaultValue": "ghcr.io/berriai/litellm:main-latest"
|
||||
},
|
||||
"containerName": {
|
||||
"type": "string",
|
||||
"defaultValue": "litellm-container"
|
||||
},
|
||||
"dnsLabelName": {
|
||||
"type": "string",
|
||||
"defaultValue": "litellm"
|
||||
},
|
||||
"portNumber": {
|
||||
"type": "int",
|
||||
"defaultValue": 4000
|
||||
}
|
||||
},
|
||||
"resources": [
|
||||
{
|
||||
"type": "Microsoft.ContainerInstance/containerGroups",
|
||||
"apiVersion": "2021-03-01",
|
||||
"name": "[parameters('containerName')]",
|
||||
"location": "[resourceGroup().location]",
|
||||
"properties": {
|
||||
"containers": [
|
||||
{
|
||||
"name": "[parameters('containerName')]",
|
||||
"properties": {
|
||||
"image": "[parameters('imageName')]",
|
||||
"resources": {
|
||||
"requests": {
|
||||
"cpu": 1,
|
||||
"memoryInGB": 2
|
||||
}
|
||||
},
|
||||
"ports": [
|
||||
{
|
||||
"port": "[parameters('portNumber')]"
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
],
|
||||
"osType": "Linux",
|
||||
"restartPolicy": "Always",
|
||||
"ipAddress": {
|
||||
"type": "Public",
|
||||
"ports": [
|
||||
{
|
||||
"protocol": "tcp",
|
||||
"port": "[parameters('portNumber')]"
|
||||
}
|
||||
],
|
||||
"dnsNameLabel": "[parameters('dnsLabelName')]"
|
||||
}
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
|
|
@ -1,42 +0,0 @@
|
|||
param imageName string = 'ghcr.io/berriai/litellm:main-latest'
|
||||
param containerName string = 'litellm-container'
|
||||
param dnsLabelName string = 'litellm'
|
||||
param portNumber int = 4000
|
||||
|
||||
resource containerGroupName 'Microsoft.ContainerInstance/containerGroups@2021-03-01' = {
|
||||
name: containerName
|
||||
location: resourceGroup().location
|
||||
properties: {
|
||||
containers: [
|
||||
{
|
||||
name: containerName
|
||||
properties: {
|
||||
image: imageName
|
||||
resources: {
|
||||
requests: {
|
||||
cpu: 1
|
||||
memoryInGB: 2
|
||||
}
|
||||
}
|
||||
ports: [
|
||||
{
|
||||
port: portNumber
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
]
|
||||
osType: 'Linux'
|
||||
restartPolicy: 'Always'
|
||||
ipAddress: {
|
||||
type: 'Public'
|
||||
ports: [
|
||||
{
|
||||
protocol: 'tcp'
|
||||
port: portNumber
|
||||
}
|
||||
]
|
||||
dnsNameLabel: dnsLabelName
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -1,10 +1,10 @@
|
|||
# syntax=docker/dockerfile:1.7
|
||||
|
||||
# Base image for building
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f
|
||||
|
||||
# Runtime image
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f
|
||||
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
|
||||
# Pinned by digest like the other base images; bump explicitly on Node upgrades.
|
||||
ARG UI_BUILD_IMAGE=node:20.18-alpine3.20@sha256:3488b10bf958af7125a176419d2d8a9937d895bf124012aae811651988d2ffe6
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
# syntax=docker/dockerfile:1.7
|
||||
|
||||
# Base images
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f
|
||||
ARG PROXY_EXTRAS_SOURCE=published
|
||||
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
|
||||
# Pinned by digest like the other base images; bump explicitly on Node upgrades.
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ if TYPE_CHECKING:
|
|||
from litellm.proxy._types import LiteLLM_ManagedObjectTable
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
from litellm.router import Router
|
||||
from litellm.types.utils import LiteLLMBatch
|
||||
|
||||
|
||||
CHECK_BATCH_COST_USER_AGENT = "LiteLLM Proxy/CheckBatchCost"
|
||||
|
|
@ -57,9 +58,7 @@ class CheckBatchCost:
|
|||
"user_api_key_alias": getattr(user_row, "user_alias", None),
|
||||
}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
f"CheckBatchCost: could not look up user {user_id} for batch {batch_id}: {e}"
|
||||
)
|
||||
verbose_proxy_logger.error(f"CheckBatchCost: could not look up user {user_id} for batch {batch_id}: {e}")
|
||||
return {}
|
||||
|
||||
async def _cleanup_stale_managed_objects(self) -> None:
|
||||
|
|
@ -68,22 +67,11 @@ class CheckBatchCost:
|
|||
in non-terminal states as 'stale_expired'. These will never complete and
|
||||
should not be polled.
|
||||
"""
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(
|
||||
days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS
|
||||
)
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS)
|
||||
result = await self.prisma_client.db.litellm_managedobjecttable.update_many(
|
||||
where={
|
||||
"file_purpose": "batch",
|
||||
"status": {
|
||||
"not_in": [
|
||||
"completed",
|
||||
"complete",
|
||||
"failed",
|
||||
"expired",
|
||||
"cancelled",
|
||||
"stale_expired",
|
||||
]
|
||||
},
|
||||
"status": {"not_in": ["completed", "complete", "failed", "expired", "cancelled", "stale_expired"]},
|
||||
"created_at": {"lt": cutoff},
|
||||
},
|
||||
data={"status": "stale_expired"},
|
||||
|
|
@ -290,13 +278,20 @@ class CheckBatchCost:
|
|||
except Exception:
|
||||
return None
|
||||
|
||||
async def check_batch_cost(self):
|
||||
async def _track_completed_batch_cost(
|
||||
self,
|
||||
job: "LiteLLM_ManagedObjectTable",
|
||||
response: "LiteLLMBatch",
|
||||
model_id: str,
|
||||
batch_id: str,
|
||||
prom_logger: Optional["PrometheusLogger"],
|
||||
) -> Optional[Tuple[Optional[str], Optional[str]]]:
|
||||
"""
|
||||
Check if the batch JOB has been tracked.
|
||||
- get all status="validating" and file_purpose="batch" jobs
|
||||
- check if batch is now complete
|
||||
- if not, return False
|
||||
- if so, return True
|
||||
Fetch a completed batch's results, compute cost/usage, and emit the
|
||||
aretrieve_batch spend log. Returns (model_name, llm_provider) on
|
||||
success, None when the job can't be routed to a deployment. Raises on
|
||||
results-fetch or cost-computation failures so the caller can leave the
|
||||
job unprocessed and retry it on a later poll.
|
||||
"""
|
||||
from litellm.batches.batch_utils import (
|
||||
_get_file_content_as_dictionary,
|
||||
|
|
@ -309,14 +304,189 @@ class CheckBatchCost:
|
|||
_is_base64_encoded_unified_file_id,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"Batch ID: {batch_id} is complete, tracking cost and usage"
|
||||
)
|
||||
|
||||
# aretrieve_batch is called with the raw provider batch ID, so response.id
|
||||
# is the raw provider value (e.g. "batch_20260223-0518.234"). We need the
|
||||
# unified base64 ID in the S3 log so downstream consumers can correlate it
|
||||
# back to the batch they submitted via the proxy.
|
||||
#
|
||||
# CheckBatchCost builds its own LiteLLMLogging object (logging_obj below) and
|
||||
# calls async_success_handler(result=response) directly. That handler calls
|
||||
# _build_standard_logging_payload(response, ...) which reads response.id at
|
||||
# that point — so setting response.id here is sufficient.
|
||||
#
|
||||
# The HTTP endpoint does this substitution via the managed files hook
|
||||
# (async_post_call_success_hook). CheckBatchCost bypasses that hook entirely,
|
||||
# so we do it explicitly here.
|
||||
response.id = job.unified_object_id
|
||||
|
||||
# This background job runs as default_user_id, so going through the HTTP endpoint
|
||||
# would trigger check_managed_file_id_access and get 403. Instead, extract the raw
|
||||
# provider file ID and call afile_content directly with deployment credentials.
|
||||
raw_output_file_id = response.output_file_id
|
||||
decoded = _is_base64_encoded_unified_file_id(raw_output_file_id)
|
||||
if decoded:
|
||||
try:
|
||||
raw_output_file_id = decoded.split("llm_output_file_id,")[1].split(";")[0]
|
||||
except (IndexError, AttributeError):
|
||||
pass
|
||||
|
||||
credentials = self.llm_router.get_deployment_credentials_with_provider(model_id) or {}
|
||||
_file_content = await afile_content(
|
||||
file_id=raw_output_file_id,
|
||||
**credentials,
|
||||
)
|
||||
|
||||
# Access content - handle both direct attribute and method call
|
||||
if hasattr(_file_content, 'content'):
|
||||
content_bytes = _file_content.content # type: ignore[union-attr]
|
||||
elif hasattr(_file_content, 'read'):
|
||||
content_bytes = await _file_content.read() # type: ignore[misc]
|
||||
else:
|
||||
content_bytes = _file_content # type: ignore[assignment]
|
||||
|
||||
file_content_as_dict = _get_file_content_as_dictionary(
|
||||
content_bytes # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
# Record output file size
|
||||
if prom_logger and content_bytes:
|
||||
try:
|
||||
prom_logger.record_managed_file_size(
|
||||
size_bytes=len(content_bytes), # type: ignore
|
||||
purpose="batch",
|
||||
file_type="output",
|
||||
model=model_id,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
deployment_info = self.llm_router.get_deployment(model_id=model_id)
|
||||
if deployment_info is None:
|
||||
verbose_proxy_logger.info(
|
||||
f"Skipping job {job.unified_object_id} because it is not a valid deployment info"
|
||||
)
|
||||
self._record_error(prom_logger, "deployment_not_found")
|
||||
return None
|
||||
custom_llm_provider = deployment_info.litellm_params.custom_llm_provider
|
||||
litellm_model_name = deployment_info.litellm_params.model
|
||||
|
||||
model_name, llm_provider, _, _ = get_llm_provider(
|
||||
model=litellm_model_name,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
# CheckBatchCost bypasses async_post_call_success_hook, so convert raw
|
||||
# output/error file IDs to managed base64 IDs before the DB write here.
|
||||
managed_files_hook = self.proxy_logging_obj.get_proxy_hook("managed_files")
|
||||
if managed_files_hook is not None:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
_minimal_auth = UserAPIKeyAuth(
|
||||
user_id=job.created_by or "default-user-id",
|
||||
team_id=getattr(job, "team_id", None),
|
||||
)
|
||||
for _file_attr in ["output_file_id", "error_file_id"]:
|
||||
_raw_file_id = getattr(response, _file_attr, None)
|
||||
if _raw_file_id and not _is_base64_encoded_unified_file_id(_raw_file_id):
|
||||
try:
|
||||
_unified_file_id = managed_files_hook.get_unified_output_file_id(
|
||||
output_file_id=_raw_file_id,
|
||||
model_id=model_id,
|
||||
model_name=str(model_name) if model_name else deployment_info.model_name or None,
|
||||
)
|
||||
await managed_files_hook.store_unified_file_id(
|
||||
file_id=_unified_file_id,
|
||||
file_object=None,
|
||||
litellm_parent_otel_span=None,
|
||||
model_mappings={model_id: _raw_file_id},
|
||||
user_api_key_dict=_minimal_auth,
|
||||
)
|
||||
setattr(response, _file_attr, _unified_file_id)
|
||||
verbose_proxy_logger.info(
|
||||
f"CheckBatchCost: converted {_file_attr} "
|
||||
f"{_raw_file_id!r} -> managed ID for batch {batch_id}"
|
||||
)
|
||||
except Exception as _e:
|
||||
verbose_proxy_logger.warning(
|
||||
f"CheckBatchCost: failed to create managed file ID for "
|
||||
f"{_file_attr}={_raw_file_id!r}: {_e}"
|
||||
)
|
||||
|
||||
# Pass deployment model_info so custom batch pricing
|
||||
# (input_cost_per_token_batches etc.) is used for cost calc
|
||||
deployment_model_info = deployment_info.model_info.model_dump() if deployment_info.model_info else {}
|
||||
batch_cost, batch_usage, batch_models = (
|
||||
await calculate_batch_cost_and_usage(
|
||||
file_content_dictionary=file_content_as_dict,
|
||||
custom_llm_provider=llm_provider, # type: ignore
|
||||
model_name=model_name,
|
||||
model_info=deployment_model_info, # type: ignore[arg-type]
|
||||
)
|
||||
)
|
||||
logging_obj = LiteLLMLogging(
|
||||
model=batch_models[0],
|
||||
messages=[{"role": "user", "content": "<retrieve_batch>"}],
|
||||
stream=False,
|
||||
call_type="aretrieve_batch",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id=str(uuid.uuid4()),
|
||||
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
|
||||
"proxy_server_request": {
|
||||
"headers": {
|
||||
"user-agent": CHECK_BATCH_COST_USER_AGENT,
|
||||
}
|
||||
},
|
||||
"metadata": {
|
||||
"user_api_key_user_id": creator_user_id,
|
||||
**user_info,
|
||||
},
|
||||
},
|
||||
optional_params={},
|
||||
)
|
||||
|
||||
await logging_obj.async_success_handler(
|
||||
result=response,
|
||||
batch_cost=batch_cost,
|
||||
batch_usage=batch_usage,
|
||||
batch_models=batch_models,
|
||||
)
|
||||
|
||||
# Record batch duration (completed_at - created_at)
|
||||
if prom_logger and response.completed_at and response.created_at:
|
||||
duration_seconds = float(response.completed_at - response.created_at)
|
||||
if duration_seconds >= 0:
|
||||
prom_logger.record_managed_batch_duration(
|
||||
duration_seconds=duration_seconds,
|
||||
model=model_name,
|
||||
api_provider=str(llm_provider) if llm_provider else None,
|
||||
)
|
||||
|
||||
return model_name, str(llm_provider) if llm_provider else None
|
||||
|
||||
async def check_batch_cost(self):
|
||||
"""
|
||||
Check if the batch JOB has been tracked.
|
||||
- get all status="validating" and file_purpose="batch" jobs
|
||||
- check if batch is now complete
|
||||
- if not, return False
|
||||
- if so, return True
|
||||
"""
|
||||
try:
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
|
||||
prom_logger = PrometheusLogger.get_instance()
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
f"CheckBatchCost: could not get Prometheus logger: {e}"
|
||||
)
|
||||
verbose_proxy_logger.error(f"CheckBatchCost: could not get Prometheus logger: {e}")
|
||||
prom_logger = None
|
||||
|
||||
processed_models: List[Tuple[Optional[str], Optional[str]]] = []
|
||||
|
|
@ -355,11 +525,7 @@ class CheckBatchCost:
|
|||
order={"created_at": "asc"},
|
||||
)
|
||||
except Exception as query_err:
|
||||
if (
|
||||
"batch_processed" not in str(query_err).lower()
|
||||
and "unknown column" not in str(query_err).lower()
|
||||
and "does not exist" not in str(query_err).lower()
|
||||
):
|
||||
if "batch_processed" not in str(query_err).lower() and "unknown column" not in str(query_err).lower() and "does not exist" not in str(query_err).lower():
|
||||
raise
|
||||
# Permanent schema gap — cache the result so future cycles skip straight to fallback
|
||||
self._has_batch_processed_column = False
|
||||
|
|
@ -393,210 +559,34 @@ class CheckBatchCost:
|
|||
f"Skipping job {job.unified_object_id} because of error querying model ID: {model_id} for cost and usage of batch ID: {batch_id}: {e}"
|
||||
)
|
||||
if prom_logger:
|
||||
prom_logger.record_check_batch_cost_error(
|
||||
"provider_retrieval_error"
|
||||
)
|
||||
prom_logger.record_check_batch_cost_error("provider_retrieval_error")
|
||||
continue
|
||||
|
||||
## RETRIEVE THE BATCH JOB OUTPUT FILE
|
||||
if response.status == "completed" and response.output_file_id is not None:
|
||||
verbose_proxy_logger.info(
|
||||
f"Batch ID: {batch_id} is complete, tracking cost and usage"
|
||||
)
|
||||
|
||||
# aretrieve_batch is called with the raw provider batch ID, so response.id
|
||||
# is the raw provider value (e.g. "batch_20260223-0518.234"). We need the
|
||||
# unified base64 ID in the S3 log so downstream consumers can correlate it
|
||||
# back to the batch they submitted via the proxy.
|
||||
#
|
||||
# CheckBatchCost builds its own LiteLLMLogging object (logging_obj below) and
|
||||
# calls async_success_handler(result=response) directly. That handler calls
|
||||
# _build_standard_logging_payload(response, ...) which reads response.id at
|
||||
# that point — so setting response.id here is sufficient.
|
||||
#
|
||||
# The HTTP endpoint does this substitution via the managed files hook
|
||||
# (async_post_call_success_hook). CheckBatchCost bypasses that hook entirely,
|
||||
# so we do it explicitly here.
|
||||
response.id = job.unified_object_id
|
||||
|
||||
# This background job runs as default_user_id, so going through the HTTP endpoint
|
||||
# would trigger check_managed_file_id_access and get 403. Instead, extract the raw
|
||||
# provider file ID and call afile_content directly with deployment credentials.
|
||||
raw_output_file_id = response.output_file_id
|
||||
decoded = _is_base64_encoded_unified_file_id(raw_output_file_id)
|
||||
if decoded:
|
||||
try:
|
||||
raw_output_file_id = decoded.split("llm_output_file_id,")[
|
||||
1
|
||||
].split(";")[0]
|
||||
except (IndexError, AttributeError):
|
||||
pass
|
||||
|
||||
credentials = (
|
||||
self.llm_router.get_deployment_credentials_with_provider(model_id)
|
||||
or {}
|
||||
)
|
||||
_file_content = await afile_content(
|
||||
file_id=raw_output_file_id,
|
||||
**credentials,
|
||||
)
|
||||
|
||||
# Access content - handle both direct attribute and method call
|
||||
if hasattr(_file_content, "content"):
|
||||
content_bytes = _file_content.content # type: ignore[union-attr]
|
||||
elif hasattr(_file_content, "read"):
|
||||
content_bytes = await _file_content.read() # type: ignore[misc]
|
||||
else:
|
||||
content_bytes = _file_content # type: ignore[assignment]
|
||||
|
||||
file_content_as_dict = _get_file_content_as_dictionary(
|
||||
content_bytes # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
# Record output file size
|
||||
if prom_logger and content_bytes:
|
||||
try:
|
||||
prom_logger.record_managed_file_size(
|
||||
size_bytes=len(content_bytes), # type: ignore
|
||||
purpose="batch",
|
||||
file_type="output",
|
||||
model=model_id,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
deployment_info = self.llm_router.get_deployment(model_id=model_id)
|
||||
if deployment_info is None:
|
||||
verbose_proxy_logger.info(
|
||||
f"Skipping job {job.unified_object_id} because it is not a valid deployment info"
|
||||
if (
|
||||
response.status == "completed"
|
||||
and response.output_file_id is not None
|
||||
):
|
||||
try:
|
||||
tracked = await self._track_completed_batch_cost(
|
||||
job=job,
|
||||
response=response,
|
||||
model_id=model_id,
|
||||
batch_id=batch_id,
|
||||
prom_logger=prom_logger,
|
||||
)
|
||||
if prom_logger:
|
||||
prom_logger.record_check_batch_cost_error(
|
||||
"deployment_not_found"
|
||||
)
|
||||
except Exception as tracking_err:
|
||||
verbose_proxy_logger.error(
|
||||
f"CheckBatchCost: failed to track cost for batch {batch_id} "
|
||||
f"(job {job.id}); leaving it unprocessed so the next poll retries: {tracking_err}"
|
||||
)
|
||||
self._record_error(prom_logger, "cost_tracking_error")
|
||||
continue
|
||||
if tracked is None:
|
||||
continue
|
||||
custom_llm_provider = deployment_info.litellm_params.custom_llm_provider
|
||||
litellm_model_name = deployment_info.litellm_params.model
|
||||
|
||||
model_name, llm_provider, _, _ = get_llm_provider(
|
||||
model=litellm_model_name,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
# CheckBatchCost bypasses async_post_call_success_hook, so convert raw
|
||||
# output/error file IDs to managed base64 IDs before the DB write here.
|
||||
managed_files_hook = self.proxy_logging_obj.get_proxy_hook(
|
||||
"managed_files"
|
||||
)
|
||||
if managed_files_hook is not None:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
_minimal_auth = UserAPIKeyAuth(
|
||||
user_id=job.created_by or "default-user-id",
|
||||
team_id=getattr(job, "team_id", None),
|
||||
)
|
||||
for _file_attr in ["output_file_id", "error_file_id"]:
|
||||
_raw_file_id = getattr(response, _file_attr, None)
|
||||
if _raw_file_id and not _is_base64_encoded_unified_file_id(
|
||||
_raw_file_id
|
||||
):
|
||||
try:
|
||||
_unified_file_id = (
|
||||
managed_files_hook.get_unified_output_file_id(
|
||||
output_file_id=_raw_file_id,
|
||||
model_id=model_id,
|
||||
model_name=(
|
||||
str(model_name)
|
||||
if model_name
|
||||
else deployment_info.model_name or None
|
||||
),
|
||||
)
|
||||
)
|
||||
await managed_files_hook.store_unified_file_id(
|
||||
file_id=_unified_file_id,
|
||||
file_object=None,
|
||||
litellm_parent_otel_span=None,
|
||||
model_mappings={model_id: _raw_file_id},
|
||||
user_api_key_dict=_minimal_auth,
|
||||
)
|
||||
setattr(response, _file_attr, _unified_file_id)
|
||||
verbose_proxy_logger.info(
|
||||
f"CheckBatchCost: converted {_file_attr} "
|
||||
f"{_raw_file_id!r} -> managed ID for batch {batch_id}"
|
||||
)
|
||||
except Exception as _e:
|
||||
verbose_proxy_logger.warning(
|
||||
f"CheckBatchCost: failed to create managed file ID for "
|
||||
f"{_file_attr}={_raw_file_id!r}: {_e}"
|
||||
)
|
||||
|
||||
# Pass deployment model_info so custom batch pricing
|
||||
# (input_cost_per_token_batches etc.) is used for cost calc
|
||||
deployment_model_info = (
|
||||
deployment_info.model_info.model_dump()
|
||||
if deployment_info.model_info
|
||||
else {}
|
||||
)
|
||||
batch_cost, batch_usage, batch_models = (
|
||||
await calculate_batch_cost_and_usage(
|
||||
file_content_dictionary=file_content_as_dict,
|
||||
custom_llm_provider=llm_provider, # type: ignore
|
||||
model_name=model_name,
|
||||
model_info=deployment_model_info, # type: ignore[arg-type]
|
||||
)
|
||||
)
|
||||
logging_obj = LiteLLMLogging(
|
||||
model=batch_models[0],
|
||||
messages=[{"role": "user", "content": "<retrieve_batch>"}],
|
||||
stream=False,
|
||||
call_type="aretrieve_batch",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id=str(uuid.uuid4()),
|
||||
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
|
||||
"proxy_server_request": {
|
||||
"headers": {
|
||||
"user-agent": CHECK_BATCH_COST_USER_AGENT,
|
||||
}
|
||||
},
|
||||
"metadata": {
|
||||
"user_api_key_user_id": creator_user_id,
|
||||
**user_info,
|
||||
},
|
||||
},
|
||||
optional_params={},
|
||||
)
|
||||
|
||||
await logging_obj.async_success_handler(
|
||||
result=response,
|
||||
batch_cost=batch_cost,
|
||||
batch_usage=batch_usage,
|
||||
batch_models=batch_models,
|
||||
)
|
||||
|
||||
# Record batch duration (completed_at - created_at)
|
||||
if prom_logger and response.completed_at and response.created_at:
|
||||
duration_seconds = float(
|
||||
response.completed_at - response.created_at
|
||||
)
|
||||
if duration_seconds >= 0:
|
||||
prom_logger.record_managed_batch_duration(
|
||||
duration_seconds=duration_seconds,
|
||||
model=model_name,
|
||||
api_provider=str(llm_provider) if llm_provider else None,
|
||||
)
|
||||
|
||||
# Track this job for the final metrics summary
|
||||
processed_models.append(
|
||||
(model_name, str(llm_provider) if llm_provider else None)
|
||||
)
|
||||
processed_models.append(tracked)
|
||||
|
||||
# mark the job as complete
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.47"
|
||||
version = "0.1.48"
|
||||
description = "Package for LiteLLM Enterprise features"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.1.47"
|
||||
version = "0.1.48"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:c61ac6919b811ea53c4782d69f1fe05218ba3c25d53f01b6ab7892e621bd4370
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f
|
||||
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
|
||||
|
||||
FROM $UV_IMAGE AS uvbin
|
||||
|
|
|
|||
|
|
@ -25,17 +25,25 @@ DatabaseURLSettings.from_env().apply_to_env()
|
|||
|
||||
from litellm.proxy.proxy_server import app
|
||||
|
||||
from gateway.routes.allowlist import GATEWAY_EXACT_PATHS, GATEWAY_PATH_PREFIXES
|
||||
from gateway.routes.allowlist import (
|
||||
GATEWAY_EXACT_PATHS,
|
||||
GATEWAY_MOUNT_PATHS,
|
||||
GATEWAY_PATH_PREFIXES,
|
||||
)
|
||||
|
||||
|
||||
def _is_gateway_route(route) -> bool:
|
||||
"""Keep the route on the gateway if its path is in the LLM data-plane surface."""
|
||||
"""Keep the route on the gateway if its path is in the LLM data-plane surface.
|
||||
|
||||
Prometheus registers /metrics as a Mount (``app.mount("/metrics", make_asgi_app())``),
|
||||
so Mounts are matched against GATEWAY_MOUNT_PATHS instead of being dropped with
|
||||
the UI static mounts.
|
||||
"""
|
||||
path = getattr(route, "path", None)
|
||||
if path is None:
|
||||
return False
|
||||
if isinstance(route, Mount):
|
||||
# Gateway never serves the static UI or its asset bundles.
|
||||
return False
|
||||
return path in GATEWAY_MOUNT_PATHS
|
||||
if path in GATEWAY_EXACT_PATHS:
|
||||
return True
|
||||
return any(path.startswith(prefix) for prefix in GATEWAY_PATH_PREFIXES)
|
||||
|
|
|
|||
|
|
@ -106,7 +106,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
|
|||
# Health & ops
|
||||
"/health",
|
||||
"/metrics",
|
||||
"/watsonx"
|
||||
"/watsonx",
|
||||
)
|
||||
|
||||
GATEWAY_EXACT_PATHS: frozenset[str] = frozenset(
|
||||
|
|
@ -120,3 +120,9 @@ GATEWAY_EXACT_PATHS: frozenset[str] = frozenset(
|
|||
"/test",
|
||||
}
|
||||
)
|
||||
|
||||
GATEWAY_MOUNT_PATHS: frozenset[str] = frozenset(
|
||||
{
|
||||
"/metrics",
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -45,11 +45,16 @@ spec:
|
|||
value: /app/config/config.yaml
|
||||
{{- end }}
|
||||
{{- include "litellm.envFrom" .Values.backend | nindent 10 }}
|
||||
{{- if .Values.gateway.config.create }}
|
||||
{{- if or .Values.gateway.config.create .Values.backend.volumeMounts }}
|
||||
volumeMounts:
|
||||
{{- if .Values.gateway.config.create }}
|
||||
- name: gateway-config
|
||||
mountPath: /app/config/config.yaml
|
||||
subPath: config.yaml
|
||||
{{- end }}
|
||||
{{- with .Values.backend.volumeMounts }}
|
||||
{{- toYaml . | nindent 12 }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
{{- with .Values.backend.livenessProbe }}
|
||||
livenessProbe:
|
||||
|
|
@ -61,11 +66,16 @@ spec:
|
|||
{{- end }}
|
||||
resources:
|
||||
{{- toYaml .Values.backend.resources | nindent 12 }}
|
||||
{{- if .Values.gateway.config.create }}
|
||||
{{- if or .Values.gateway.config.create .Values.backend.volumes }}
|
||||
volumes:
|
||||
{{- if .Values.gateway.config.create }}
|
||||
- name: gateway-config
|
||||
configMap:
|
||||
name: {{ include "litellm.gateway.fullname" . }}-config
|
||||
{{- end }}
|
||||
{{- with .Values.backend.volumes }}
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
{{- with .Values.backend.nodeSelector }}
|
||||
nodeSelector:
|
||||
|
|
|
|||
|
|
@ -47,11 +47,16 @@ spec:
|
|||
value: {{ .Values.gateway.numWorkers | quote }}
|
||||
{{- end }}
|
||||
{{- include "litellm.envFrom" .Values.gateway | nindent 10 }}
|
||||
{{- if .Values.gateway.config.create }}
|
||||
{{- if or .Values.gateway.config.create .Values.gateway.volumeMounts }}
|
||||
volumeMounts:
|
||||
{{- if .Values.gateway.config.create }}
|
||||
- name: gateway-config
|
||||
mountPath: /app/config/config.yaml
|
||||
subPath: config.yaml
|
||||
{{- end }}
|
||||
{{- with .Values.gateway.volumeMounts }}
|
||||
{{- toYaml . | nindent 12 }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
{{- with .Values.gateway.livenessProbe }}
|
||||
livenessProbe:
|
||||
|
|
@ -63,11 +68,16 @@ spec:
|
|||
{{- end }}
|
||||
resources:
|
||||
{{- toYaml .Values.gateway.resources | nindent 12 }}
|
||||
{{- if .Values.gateway.config.create }}
|
||||
{{- if or .Values.gateway.config.create .Values.gateway.volumes }}
|
||||
volumes:
|
||||
{{- if .Values.gateway.config.create }}
|
||||
- name: gateway-config
|
||||
configMap:
|
||||
name: {{ include "litellm.gateway.fullname" . }}-config
|
||||
{{- end }}
|
||||
{{- with .Values.gateway.volumes }}
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- end }}
|
||||
{{- with .Values.gateway.nodeSelector }}
|
||||
nodeSelector:
|
||||
|
|
|
|||
|
|
@ -46,6 +46,10 @@ spec:
|
|||
{{- toYaml . | nindent 12 }}
|
||||
{{- end }}
|
||||
{{- include "litellm.envFrom" .Values.ui | nindent 10 }}
|
||||
{{- with .Values.ui.volumeMounts }}
|
||||
volumeMounts:
|
||||
{{- toYaml . | nindent 12 }}
|
||||
{{- end }}
|
||||
{{- with .Values.ui.livenessProbe }}
|
||||
livenessProbe:
|
||||
{{- toYaml . | nindent 12 }}
|
||||
|
|
@ -56,6 +60,10 @@ spec:
|
|||
{{- end }}
|
||||
resources:
|
||||
{{- toYaml .Values.ui.resources | nindent 12 }}
|
||||
{{- with .Values.ui.volumes }}
|
||||
volumes:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
{{- end }}
|
||||
{{- with .Values.ui.nodeSelector }}
|
||||
nodeSelector:
|
||||
{{- toYaml . | nindent 8 }}
|
||||
|
|
|
|||
172
helm/litellm/tests/deployment_volumes_tests.yaml
Normal file
172
helm/litellm/tests/deployment_volumes_tests.yaml
Normal file
|
|
@ -0,0 +1,172 @@
|
|||
suite: test deployment volumes and volumeMounts
|
||||
templates:
|
||||
- gateway/deployment.yaml
|
||||
- gateway/configmap.yaml
|
||||
- backend/deployment.yaml
|
||||
- ui/deployment.yaml
|
||||
values:
|
||||
- ./values/required.yaml
|
||||
tests:
|
||||
- it: gateway renders only the config volume by default
|
||||
template: gateway/deployment.yaml
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.template.spec.volumes
|
||||
value:
|
||||
- name: gateway-config
|
||||
configMap:
|
||||
name: RELEASE-NAME-litellm-gateway-config
|
||||
- equal:
|
||||
path: spec.template.spec.containers[0].volumeMounts
|
||||
value:
|
||||
- name: gateway-config
|
||||
mountPath: /app/config/config.yaml
|
||||
subPath: config.yaml
|
||||
|
||||
- it: gateway merges user volumes and volumeMounts with the config volume
|
||||
template: gateway/deployment.yaml
|
||||
set:
|
||||
gateway.volumes:
|
||||
- name: custom-callbacks
|
||||
configMap:
|
||||
name: custom-callbacks
|
||||
gateway.volumeMounts:
|
||||
- name: custom-callbacks
|
||||
mountPath: /app/custom_callbacks.py
|
||||
subPath: custom_callbacks.py
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.template.spec.volumes[0].name
|
||||
value: gateway-config
|
||||
- equal:
|
||||
path: spec.template.spec.volumes[1]
|
||||
value:
|
||||
name: custom-callbacks
|
||||
configMap:
|
||||
name: custom-callbacks
|
||||
- equal:
|
||||
path: spec.template.spec.containers[0].volumeMounts[0].name
|
||||
value: gateway-config
|
||||
- equal:
|
||||
path: spec.template.spec.containers[0].volumeMounts[1]
|
||||
value:
|
||||
name: custom-callbacks
|
||||
mountPath: /app/custom_callbacks.py
|
||||
subPath: custom_callbacks.py
|
||||
|
||||
- it: gateway renders user volumes even when config creation is disabled
|
||||
template: gateway/deployment.yaml
|
||||
set:
|
||||
gateway.config.create: false
|
||||
gateway.volumes:
|
||||
- name: certs
|
||||
secret:
|
||||
secretName: tls-certs
|
||||
gateway.volumeMounts:
|
||||
- name: certs
|
||||
mountPath: /etc/certs
|
||||
readOnly: true
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.template.spec.volumes
|
||||
value:
|
||||
- name: certs
|
||||
secret:
|
||||
secretName: tls-certs
|
||||
- equal:
|
||||
path: spec.template.spec.containers[0].volumeMounts
|
||||
value:
|
||||
- name: certs
|
||||
mountPath: /etc/certs
|
||||
readOnly: true
|
||||
|
||||
- it: gateway omits volumes when config creation is disabled and no user volumes are set
|
||||
template: gateway/deployment.yaml
|
||||
set:
|
||||
gateway.config.create: false
|
||||
asserts:
|
||||
- isNull:
|
||||
path: spec.template.spec.volumes
|
||||
- isNull:
|
||||
path: spec.template.spec.containers[0].volumeMounts
|
||||
|
||||
- it: backend merges user volumes and volumeMounts with the shared config volume
|
||||
template: backend/deployment.yaml
|
||||
set:
|
||||
backend.volumes:
|
||||
- name: sso-handler
|
||||
configMap:
|
||||
name: sso-handler
|
||||
backend.volumeMounts:
|
||||
- name: sso-handler
|
||||
mountPath: /app/custom_sso.py
|
||||
subPath: custom_sso.py
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.template.spec.volumes[0].name
|
||||
value: gateway-config
|
||||
- equal:
|
||||
path: spec.template.spec.volumes[1]
|
||||
value:
|
||||
name: sso-handler
|
||||
configMap:
|
||||
name: sso-handler
|
||||
- equal:
|
||||
path: spec.template.spec.containers[0].volumeMounts[1]
|
||||
value:
|
||||
name: sso-handler
|
||||
mountPath: /app/custom_sso.py
|
||||
subPath: custom_sso.py
|
||||
|
||||
- it: backend renders user volumes even when config creation is disabled
|
||||
template: backend/deployment.yaml
|
||||
set:
|
||||
gateway.config.create: false
|
||||
backend.volumes:
|
||||
- name: data
|
||||
emptyDir: {}
|
||||
backend.volumeMounts:
|
||||
- name: data
|
||||
mountPath: /data
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.template.spec.volumes
|
||||
value:
|
||||
- name: data
|
||||
emptyDir: {}
|
||||
- equal:
|
||||
path: spec.template.spec.containers[0].volumeMounts
|
||||
value:
|
||||
- name: data
|
||||
mountPath: /data
|
||||
|
||||
- it: ui renders no volumes by default
|
||||
template: ui/deployment.yaml
|
||||
asserts:
|
||||
- isNull:
|
||||
path: spec.template.spec.volumes
|
||||
- isNull:
|
||||
path: spec.template.spec.containers[0].volumeMounts
|
||||
|
||||
- it: ui renders user volumes and volumeMounts
|
||||
template: ui/deployment.yaml
|
||||
set:
|
||||
ui.volumes:
|
||||
- name: nginx-config
|
||||
configMap:
|
||||
name: custom-nginx
|
||||
ui.volumeMounts:
|
||||
- name: nginx-config
|
||||
mountPath: /etc/nginx/conf.d
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.template.spec.volumes
|
||||
value:
|
||||
- name: nginx-config
|
||||
configMap:
|
||||
name: custom-nginx
|
||||
- equal:
|
||||
path: spec.template.spec.containers[0].volumeMounts
|
||||
value:
|
||||
- name: nginx-config
|
||||
mountPath: /etc/nginx/conf.d
|
||||
4
helm/litellm/tests/values/required.yaml
Normal file
4
helm/litellm/tests/values/required.yaml
Normal file
|
|
@ -0,0 +1,4 @@
|
|||
database:
|
||||
writer:
|
||||
host: postgres.example.com
|
||||
dbname: litellm
|
||||
|
|
@ -124,6 +124,11 @@ gateway:
|
|||
extraEnv: [] # Add extra environment variables to the gateway
|
||||
envConfigMaps: [] # Add extra environment variables to the gateway from config maps
|
||||
envSecrets: [] # Add extra environment variables to the gateway from secrets
|
||||
# Additional volumes on the gateway Deployment (e.g. a ConfigMap holding
|
||||
# custom callback / SSO handler code, mounted next to the proxy config).
|
||||
volumes: []
|
||||
# Additional volumeMounts on the gateway container.
|
||||
volumeMounts: []
|
||||
config:
|
||||
create: true
|
||||
proxy_config: {}
|
||||
|
|
@ -167,6 +172,10 @@ backend:
|
|||
extraEnv: []
|
||||
envConfigMaps: []
|
||||
envSecrets: []
|
||||
# Additional volumes on the backend Deployment.
|
||||
volumes: []
|
||||
# Additional volumeMounts on the backend container.
|
||||
volumeMounts: []
|
||||
image:
|
||||
repository: ghcr.io/berriai/litellm-backend
|
||||
tag: ""
|
||||
|
|
@ -206,6 +215,10 @@ ui:
|
|||
extraEnv: []
|
||||
envConfigMaps: []
|
||||
envSecrets: []
|
||||
# Additional volumes on the ui Deployment.
|
||||
volumes: []
|
||||
# Additional volumeMounts on the ui container.
|
||||
volumeMounts: []
|
||||
image:
|
||||
repository: ghcr.io/berriai/litellm-ui
|
||||
tag: ""
|
||||
|
|
|
|||
|
|
@ -379,6 +379,7 @@ budget_duration: Optional[str] = (
|
|||
None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d").
|
||||
)
|
||||
default_soft_budget: float = DEFAULT_SOFT_BUDGET # by default all litellm proxy keys have a soft budget of 50.0
|
||||
budget_exceeded_throttle_percentage: Optional[float] = None
|
||||
forward_traceparent_to_llm_provider: bool = False
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -11,7 +11,6 @@ from litellm._logging import verbose_logger
|
|||
from litellm.a2a_protocol.cost_calculator import A2ACostCalculator
|
||||
from litellm.a2a_protocol.utils import A2ARequestUtils
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.thread_pool_executor import executor
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from a2a.types import SendStreamingMessageRequest, SendStreamingMessageResponse
|
||||
|
|
@ -128,22 +127,15 @@ class A2AStreamingIterator:
|
|||
|
||||
# Call success handlers - they will build standard_logging_object
|
||||
asyncio.create_task(
|
||||
self.logging_obj.async_success_handler(
|
||||
result=result,
|
||||
self.logging_obj.dispatch_success_handlers(
|
||||
result,
|
||||
start_time=self.start_time,
|
||||
end_time=end_time,
|
||||
cache_hit=None,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
|
||||
executor.submit(
|
||||
self.logging_obj.success_handler,
|
||||
result=result,
|
||||
cache_hit=None,
|
||||
start_time=self.start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
verbose_logger.info(
|
||||
f"A2A streaming completed: prompt_tokens={prompt_tokens}, "
|
||||
f"completion_tokens={completion_tokens}, total_tokens={total_tokens}, "
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ from typing import Any, Iterator, List, Literal, Optional, Tuple
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import _parse_prompt_tokens_details
|
||||
from litellm.types.llms.openai import Batch
|
||||
from litellm.types.utils import CallTypes, ModelInfo, Usage
|
||||
from litellm.utils import token_counter
|
||||
|
|
@ -34,7 +35,7 @@ async def calculate_batch_cost_and_usage(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
model_name=model_name,
|
||||
)
|
||||
batch_models = _get_batch_models_from_file_content(file_content_dictionary, model_name)
|
||||
batch_models = _get_batch_models_from_file_content(file_content_dictionary, model_name, custom_llm_provider)
|
||||
|
||||
return batch_cost, batch_usage, batch_models
|
||||
|
||||
|
|
@ -70,7 +71,7 @@ async def _handle_completed_batch(
|
|||
model_name=model_name,
|
||||
)
|
||||
|
||||
batch_models = _get_batch_models_from_file_content(file_content_dictionary, model_name)
|
||||
batch_models = _get_batch_models_from_file_content(file_content_dictionary, model_name, custom_llm_provider)
|
||||
|
||||
return batch_cost, batch_usage, batch_models
|
||||
|
||||
|
|
@ -78,6 +79,7 @@ async def _handle_completed_batch(
|
|||
def _get_batch_models_from_file_content(
|
||||
file_content_dictionary: List[dict],
|
||||
model_name: Optional[str] = None,
|
||||
custom_llm_provider: str = "openai",
|
||||
) -> List[str]:
|
||||
"""
|
||||
Get the models from the file content
|
||||
|
|
@ -86,8 +88,8 @@ def _get_batch_models_from_file_content(
|
|||
return [model_name]
|
||||
batch_models = []
|
||||
for _item in file_content_dictionary:
|
||||
if _batch_response_was_successful(_item):
|
||||
_response_body = _get_response_from_batch_job_output_file(_item)
|
||||
if _batch_response_was_successful(_item, custom_llm_provider):
|
||||
_response_body = _get_response_from_batch_job_output_file(_item, custom_llm_provider)
|
||||
_model = _response_body.get("model")
|
||||
if _model:
|
||||
batch_models.append(_model)
|
||||
|
|
@ -373,10 +375,10 @@ def _get_batch_job_cost_from_file_content(
|
|||
# parse the file content as json
|
||||
verbose_logger.debug("file_content_dictionary=%s", json.dumps(file_content_dictionary, indent=4))
|
||||
for _item in file_content_dictionary:
|
||||
if _batch_response_was_successful(_item):
|
||||
_response_body = _get_response_from_batch_job_output_file(_item)
|
||||
if model_info is not None:
|
||||
usage = _get_batch_job_usage_from_response_body(_response_body)
|
||||
if _batch_response_was_successful(_item, custom_llm_provider):
|
||||
_response_body = _get_response_from_batch_job_output_file(_item, custom_llm_provider)
|
||||
if model_info is not None or custom_llm_provider == "anthropic":
|
||||
usage = _get_batch_job_usage_from_response_body(_response_body, custom_llm_provider)
|
||||
model = _response_body.get("model", "")
|
||||
prompt_cost, completion_cost = batch_cost_calculator(
|
||||
usage=usage,
|
||||
|
|
@ -418,17 +420,31 @@ def _get_batch_job_total_usage_from_file_content(
|
|||
total_tokens: int = 0
|
||||
prompt_tokens: int = 0
|
||||
completion_tokens: int = 0
|
||||
cache_read_tokens: int = 0
|
||||
cache_creation_tokens: int = 0
|
||||
for _item in file_content_dictionary:
|
||||
if _batch_response_was_successful(_item):
|
||||
_response_body = _get_response_from_batch_job_output_file(_item)
|
||||
usage: Usage = _get_batch_job_usage_from_response_body(_response_body)
|
||||
if _batch_response_was_successful(_item, custom_llm_provider):
|
||||
_response_body = _get_response_from_batch_job_output_file(_item, custom_llm_provider)
|
||||
usage: Usage = _get_batch_job_usage_from_response_body(_response_body, custom_llm_provider)
|
||||
total_tokens += usage.total_tokens
|
||||
prompt_tokens += usage.prompt_tokens
|
||||
completion_tokens += usage.completion_tokens
|
||||
prompt_details = _parse_prompt_tokens_details(usage)
|
||||
cache_read_tokens += prompt_details["cache_hit_tokens"]
|
||||
cache_creation_tokens += prompt_details["cache_creation_tokens"]
|
||||
cache_token_params = {
|
||||
key: tokens
|
||||
for key, tokens in (
|
||||
("cache_read_input_tokens", cache_read_tokens),
|
||||
("cache_creation_input_tokens", cache_creation_tokens),
|
||||
)
|
||||
if tokens > 0
|
||||
}
|
||||
return Usage(
|
||||
total_tokens=total_tokens,
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
**cache_token_params,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -465,27 +481,51 @@ def _count_prompt_or_input_tokens(model: str, value: Any) -> int:
|
|||
return 0
|
||||
|
||||
|
||||
def _get_batch_job_usage_from_response_body(response_body: dict) -> Usage:
|
||||
def _get_batch_job_usage_from_response_body(response_body: dict, custom_llm_provider: str = "openai") -> Usage:
|
||||
"""
|
||||
Get the tokens of a batch job from the response body
|
||||
"""
|
||||
if custom_llm_provider == "anthropic":
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
|
||||
return AnthropicConfig().calculate_usage(
|
||||
usage_object=response_body.get("usage", None) or {},
|
||||
reasoning_content=None,
|
||||
)
|
||||
_usage_dict = response_body.get("usage", None) or {}
|
||||
usage: Usage = Usage(**_usage_dict)
|
||||
return usage
|
||||
|
||||
|
||||
def _get_response_from_batch_job_output_file(batch_job_output_file: dict) -> Any:
|
||||
def _get_anthropic_result_from_batch_results_line(batch_results_line: dict) -> dict:
|
||||
"""
|
||||
Get the ``result`` object from a line of an Anthropic message batch results JSONL file.
|
||||
|
||||
Anthropic batch results lines look like:
|
||||
``{"custom_id": ..., "result": {"type": "succeeded", "message": {..., "usage": {...}}}}``
|
||||
"""
|
||||
return batch_results_line.get("result", None) or {}
|
||||
|
||||
|
||||
def _get_response_from_batch_job_output_file(batch_job_output_file: dict, custom_llm_provider: str = "openai") -> Any:
|
||||
"""
|
||||
Get the response from the batch job output file
|
||||
"""
|
||||
if custom_llm_provider == "anthropic":
|
||||
return _get_anthropic_result_from_batch_results_line(batch_job_output_file).get("message", None) or {}
|
||||
_response: dict = batch_job_output_file.get("response", None) or {}
|
||||
_response_body = _response.get("body", None) or {}
|
||||
return _response_body
|
||||
|
||||
|
||||
def _batch_response_was_successful(batch_job_output_file: dict) -> bool:
|
||||
def _batch_response_was_successful(batch_job_output_file: dict, custom_llm_provider: str = "openai") -> bool:
|
||||
"""
|
||||
Check if the batch job response status == 200
|
||||
Check if the batch job response was successful
|
||||
|
||||
OpenAI-shaped output rows report ``response.status_code == 200``; Anthropic
|
||||
message batch results lines report ``result.type == "succeeded"``.
|
||||
"""
|
||||
if custom_llm_provider == "anthropic":
|
||||
return _get_anthropic_result_from_batch_results_line(batch_job_output_file).get("type") == "succeeded"
|
||||
_response: dict = batch_job_output_file.get("response", None) or {}
|
||||
return _response.get("status_code", None) == 200
|
||||
|
|
|
|||
|
|
@ -279,7 +279,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
print_verbose("No prompt provided for semantic caching")
|
||||
return
|
||||
|
||||
embedding = await self._get_async_embedding(prompt, **kwargs)
|
||||
embedding = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
await self._ensure_index_async(len(embedding))
|
||||
|
||||
doc_key = self._doc_key(key)
|
||||
|
|
@ -298,7 +298,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
embedding = await self._get_async_embedding(prompt, **kwargs)
|
||||
embedding = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
await self._ensure_index_async(len(embedding))
|
||||
|
||||
search_result = await self.async_client.ft(self.index_name).search(
|
||||
|
|
|
|||
|
|
@ -1504,6 +1504,7 @@ LITELLM_SETTINGS_SAFE_DB_OVERRIDES = [
|
|||
"public_model_groups_links",
|
||||
"cost_discount_config",
|
||||
"cost_margin_config",
|
||||
"budget_exceeded_throttle_percentage",
|
||||
]
|
||||
SPECIAL_LITELLM_AUTH_TOKEN = ["ui-token"]
|
||||
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60))
|
||||
|
|
|
|||
|
|
@ -2155,17 +2155,23 @@ def batch_cost_calculator(
|
|||
if input_cost_per_token_batches:
|
||||
total_prompt_cost = usage.prompt_tokens * input_cost_per_token_batches
|
||||
elif input_cost_per_token:
|
||||
details = _parse_prompt_tokens_details(usage)
|
||||
cache_read_tokens = details["cache_hit_tokens"]
|
||||
cache_creation_tokens = details["cache_creation_tokens"]
|
||||
|
||||
# Subtract cached tokens from prompt_tokens before calculating cost
|
||||
# Fixes issue where cached tokens are being charged again
|
||||
base_input_tokens = get_billable_input_tokens(usage) - cache_creation_tokens
|
||||
total_prompt_cost = (
|
||||
get_billable_input_tokens(usage) * (input_cost_per_token) / 2
|
||||
base_input_tokens * (input_cost_per_token) / 2
|
||||
) # batch cost is usually half of the regular token cost
|
||||
|
||||
# Add cache read cost if applicable
|
||||
details = _parse_prompt_tokens_details(usage)
|
||||
cache_read_tokens = details["cache_hit_tokens"]
|
||||
cache_read_cost_key = _get_service_tier_cost_key("cache_read_input_token_cost", None)
|
||||
total_prompt_cost += calculate_cost_component(model_info, cache_read_cost_key, cache_read_tokens) / 2
|
||||
|
||||
cache_creation_cost = model_info.get("cache_creation_input_token_cost") or input_cost_per_token
|
||||
total_prompt_cost += cache_creation_tokens * cache_creation_cost / 2
|
||||
if output_cost_per_token_batches:
|
||||
total_completion_cost = usage.completion_tokens * output_cost_per_token_batches
|
||||
elif output_cost_per_token:
|
||||
|
|
|
|||
|
|
@ -368,7 +368,11 @@ class OpenTelemetryV2(CustomLogger):
|
|||
# it (named provisionally) so it isn't leaked as an open span.
|
||||
carrier.span.end(end_time=to_ns(end_time))
|
||||
return None
|
||||
data = LLMCallSpanData.from_standard_logging_payload(payload, capture_content=self.config.capture_span_content)
|
||||
data = LLMCallSpanData.from_standard_logging_payload(
|
||||
payload,
|
||||
capture_content=self.config.capture_span_content,
|
||||
time_to_first_chunk_seconds=call.time_to_first_chunk_seconds,
|
||||
)
|
||||
end_time_ns = to_ns(end_time)
|
||||
if carrier.span is not None:
|
||||
# Born at the boundary: stamp attributes from the typed payload, set
|
||||
|
|
|
|||
|
|
@ -55,6 +55,7 @@ class GenAIMapper:
|
|||
GenAI.RESPONSE_MODEL: lambda d: d.response_model,
|
||||
GenAI.RESPONSE_ID: lambda d: d.response_id,
|
||||
GenAI.RESPONSE_FINISH_REASONS: lambda d: list(d.finish_reasons) if d.finish_reasons else None,
|
||||
GenAI.RESPONSE_TIME_TO_FIRST_CHUNK: lambda d: d.time_to_first_chunk_seconds,
|
||||
GenAI.USAGE_INPUT_TOKENS: lambda d: d.usage.input_tokens,
|
||||
GenAI.USAGE_OUTPUT_TOKENS: lambda d: d.usage.output_tokens,
|
||||
Error.TYPE: lambda d: d.error.error_type if d.error else None,
|
||||
|
|
|
|||
|
|
@ -41,7 +41,7 @@ from typing import TYPE_CHECKING, Any, Mapping, cast
|
|||
|
||||
from litellm.constants import LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL
|
||||
from litellm.integrations.otel.model.semconv import resolve_operation
|
||||
from litellm.integrations.otel.model.utils import as_str
|
||||
from litellm.integrations.otel.model.utils import as_str, to_seconds
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
|
|
@ -201,6 +201,7 @@ class LLMCallEvent:
|
|||
# span is renamed from the typed payload at close (``finish_span``); this only
|
||||
# needs to be reasonable for a span that never gets closed (a leak).
|
||||
provisional_span_name: str
|
||||
time_to_first_chunk_seconds: float | None
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, kwargs: Mapping[str, Any]) -> "LLMCallEvent":
|
||||
|
|
@ -214,9 +215,25 @@ class LLMCallEvent:
|
|||
dynamic_params=kwargs.get("standard_callback_dynamic_params"),
|
||||
is_no_upstream_call=bool(kwargs.get(LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL)),
|
||||
provisional_span_name=f"{operation.value} {model}".strip(),
|
||||
time_to_first_chunk_seconds=time_to_first_chunk_seconds(kwargs),
|
||||
)
|
||||
|
||||
|
||||
def time_to_first_chunk_seconds(kwargs: Mapping[str, Any]) -> float | None:
|
||||
"""Seconds from the upstream request being issued (``api_call_start_time``)
|
||||
to the first streamed chunk (``completion_start_time``); ``None`` for
|
||||
non-streaming calls, where ``completion_start_time`` is backfilled with the
|
||||
end time and would not measure first-chunk latency."""
|
||||
optional_params = cast(Mapping[str, Any], kwargs.get("optional_params") or {})
|
||||
if not optional_params.get("stream"):
|
||||
return None
|
||||
api_call_start = to_seconds(kwargs.get("api_call_start_time"))
|
||||
completion_start = to_seconds(kwargs.get("completion_start_time"))
|
||||
if api_call_start is None or completion_start is None:
|
||||
return None
|
||||
return completion_start - api_call_start
|
||||
|
||||
|
||||
def _call_id(payload: "StandardLoggingPayload | None", kwargs: Mapping[str, Any]) -> str | None:
|
||||
"""The call id from the payload (when closed) or the bare kwargs (at pre_call)."""
|
||||
if payload is not None:
|
||||
|
|
|
|||
|
|
@ -305,10 +305,14 @@ class LLMCallSpanData:
|
|||
messages_in: tuple[Mapping[str, object], ...] = ()
|
||||
choices_out: tuple[Mapping[str, object], ...] = ()
|
||||
system_fingerprint: str | None = None
|
||||
time_to_first_chunk_seconds: float | None = None
|
||||
|
||||
@classmethod
|
||||
def from_standard_logging_payload(
|
||||
cls, payload: "StandardLoggingPayload", capture_content: bool = False
|
||||
cls,
|
||||
payload: "StandardLoggingPayload",
|
||||
capture_content: bool = False,
|
||||
time_to_first_chunk_seconds: float | None = None,
|
||||
) -> "LLMCallSpanData":
|
||||
params = cast(Mapping[str, object], payload.get("model_parameters") or {})
|
||||
# The single parse of the request's metadata — the request-vs-provider
|
||||
|
|
@ -349,6 +353,7 @@ class LLMCallSpanData:
|
|||
messages_in=_dicts(payload.get("messages")) if capture_content else (),
|
||||
choices_out=choices_out if capture_content else (),
|
||||
system_fingerprint=as_str(response.get("system_fingerprint")),
|
||||
time_to_first_chunk_seconds=time_to_first_chunk_seconds,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -69,6 +69,7 @@ class GenAI:
|
|||
RESPONSE_ID: Final = "gen_ai.response.id"
|
||||
RESPONSE_MODEL: Final = "gen_ai.response.model"
|
||||
RESPONSE_FINISH_REASONS: Final = "gen_ai.response.finish_reasons"
|
||||
RESPONSE_TIME_TO_FIRST_CHUNK: Final = "gen_ai.response.time_to_first_chunk"
|
||||
# usage
|
||||
USAGE_INPUT_TOKENS: Final = "gen_ai.usage.input_tokens"
|
||||
USAGE_OUTPUT_TOKENS: Final = "gen_ai.usage.output_tokens"
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ from litellm.integrations.opentelemetry import (
|
|||
_build_metric_attribute_filter,
|
||||
_resolve_metric_attribute_filter,
|
||||
)
|
||||
from litellm.integrations.otel.model.metadata import time_to_first_chunk_seconds
|
||||
from litellm.integrations.otel.model.semconv import Metric, resolve_operation
|
||||
from litellm.integrations.otel.model.utils import to_seconds
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
|
@ -181,13 +182,10 @@ class GenAIMetricRecorder:
|
|||
self._metrics.token_usage.record(usage.get("completion_tokens", 0), attributes=out_attrs)
|
||||
|
||||
def _record_time_to_first_token(self, kwargs: Mapping[str, Any], common_attrs: dict) -> None:
|
||||
if not kwargs.get("optional_params", {}).get("stream", False):
|
||||
time_to_first_chunk = time_to_first_chunk_seconds(kwargs)
|
||||
if time_to_first_chunk is None:
|
||||
return
|
||||
api_call_start = to_seconds(kwargs.get("api_call_start_time"))
|
||||
completion_start = to_seconds(kwargs.get("completion_start_time"))
|
||||
if api_call_start is None or completion_start is None:
|
||||
return
|
||||
self._metrics.time_to_first_token.record(completion_start - api_call_start, attributes=common_attrs)
|
||||
self._metrics.time_to_first_token.record(time_to_first_chunk, attributes=common_attrs)
|
||||
|
||||
def _record_time_per_output_token(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -174,22 +174,15 @@ class InteractionsAPIStreamingIterator(BaseInteractionsAPIStreamingIterator):
|
|||
logging_response = copy.deepcopy(self.completed_response)
|
||||
|
||||
asyncio.create_task(
|
||||
self.logging_obj.async_success_handler(
|
||||
result=logging_response,
|
||||
self.logging_obj.dispatch_success_handlers(
|
||||
logging_response,
|
||||
start_time=self.start_time,
|
||||
end_time=datetime.now(),
|
||||
cache_hit=None,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
|
||||
executor.submit(
|
||||
self.logging_obj.success_handler,
|
||||
result=logging_response,
|
||||
cache_hit=None,
|
||||
start_time=self.start_time,
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
|
||||
class SyncInteractionsAPIStreamingIterator(BaseInteractionsAPIStreamingIterator):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -123,6 +123,34 @@ def process_audio_file(audio_file: FileTypes) -> ProcessedAudioFile:
|
|||
return ProcessedAudioFile(file_content=file_content, filename=filename, content_type=content_type)
|
||||
|
||||
|
||||
BARE_ISO_639_1_TO_BCP47 = {
|
||||
"en": "en-US",
|
||||
"es": "es-ES",
|
||||
"de": "de-DE",
|
||||
"fr": "fr-FR",
|
||||
"it": "it-IT",
|
||||
"pt": "pt-BR",
|
||||
"ja": "ja-JP",
|
||||
"ko": "ko-KR",
|
||||
"zh": "zh-CN",
|
||||
"ru": "ru-RU",
|
||||
"hi": "hi-IN",
|
||||
"ar": "ar-SA",
|
||||
}
|
||||
|
||||
|
||||
def normalize_transcription_language_to_bcp47(language: str) -> str:
|
||||
"""
|
||||
OpenAI's transcription `language` param accepts bare ISO-639-1 codes like
|
||||
``en``; speech APIs such as Google Speech-to-Text and NVIDIA Riva require
|
||||
BCP-47 like ``en-US``. Map the most common bare codes and pass through
|
||||
anything already region-qualified (or unknown, for a clear provider error).
|
||||
"""
|
||||
if "-" in language:
|
||||
return language
|
||||
return BARE_ISO_639_1_TO_BCP47.get(language.lower(), language)
|
||||
|
||||
|
||||
def get_audio_file_name(file_obj: FileTypes) -> str:
|
||||
"""
|
||||
Safely get the name of a file-like object or return its string representation.
|
||||
|
|
|
|||
|
|
@ -1944,7 +1944,7 @@ def _map_azure_exception(
|
|||
response=getattr(original_exception, "response", None),
|
||||
body=getattr(original_exception, "body", None),
|
||||
)
|
||||
elif "invalid_request_error" in error_str:
|
||||
elif "invalid_request_error" in error_str and getattr(original_exception, "status_code", None) in (None, 400):
|
||||
raise BadRequestError(
|
||||
message=f"AzureException BadRequestError - {message}",
|
||||
llm_provider="azure",
|
||||
|
|
@ -1986,6 +1986,14 @@ def _map_azure_exception(
|
|||
litellm_debug_info=extra_information,
|
||||
response=getattr(original_exception, "response", None),
|
||||
)
|
||||
elif original_exception.status_code == 404:
|
||||
raise NotFoundError(
|
||||
message=f"AzureException NotFoundError - {message}",
|
||||
llm_provider="azure",
|
||||
model=model,
|
||||
litellm_debug_info=extra_information,
|
||||
response=getattr(original_exception, "response", None),
|
||||
)
|
||||
elif original_exception.status_code == 408:
|
||||
raise Timeout(
|
||||
message=f"AzureException Timeout - {message}",
|
||||
|
|
@ -2173,7 +2181,7 @@ def exception_type( # type: ignore
|
|||
litellm_response_headers = _get_response_headers(original_exception=original_exception)
|
||||
try:
|
||||
error_str = redact_string(str(original_exception)) if _ENABLE_SECRET_REDACTION else str(original_exception)
|
||||
if model:
|
||||
if model or custom_llm_provider:
|
||||
if hasattr(original_exception, "message"):
|
||||
error_str = (
|
||||
redact_string(str(original_exception.message))
|
||||
|
|
|
|||
|
|
@ -36,6 +36,8 @@ OPTIONAL_KWARGS_KEYS = frozenset(
|
|||
"aws_bedrock_project_id",
|
||||
"tpm",
|
||||
"rpm",
|
||||
"itpm",
|
||||
"otpm",
|
||||
"use_xai_oauth",
|
||||
}
|
||||
)
|
||||
|
|
@ -74,6 +76,7 @@ def get_litellm_params(
|
|||
proxy_server_request=None,
|
||||
acompletion=None,
|
||||
aembedding=None,
|
||||
allm_passthrough_route=None,
|
||||
preset_cache_key=None,
|
||||
no_log=None,
|
||||
input_cost_per_second=None,
|
||||
|
|
@ -116,6 +119,7 @@ def get_litellm_params(
|
|||
# Build base dict with explicit parameters (always included)
|
||||
litellm_params = {
|
||||
"acompletion": acompletion,
|
||||
"allm_passthrough_route": allm_passthrough_route,
|
||||
"api_key": api_key,
|
||||
"force_timeout": force_timeout,
|
||||
"logger_fn": logger_fn,
|
||||
|
|
|
|||
|
|
@ -1530,6 +1530,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
and litellm_params.get(CallTypes.aembedding.value, False) is not True
|
||||
and litellm_params.get(CallTypes.aimage_generation.value, False) is not True
|
||||
and litellm_params.get(CallTypes.atranscription.value, False) is not True
|
||||
and litellm_params.get(CallTypes.allm_passthrough_route.value, False) is not True
|
||||
)
|
||||
|
||||
def _is_assembled_stream_success(self, result=None) -> bool:
|
||||
|
|
|
|||
|
|
@ -1,5 +1,4 @@
|
|||
import asyncio
|
||||
import concurrent.futures
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Protocol, Union, cast
|
||||
|
||||
|
|
@ -25,9 +24,6 @@ if TYPE_CHECKING:
|
|||
else:
|
||||
CLIENT_CONNECTION_CLASS = Any
|
||||
|
||||
# Create a thread pool with a maximum of 10 threads
|
||||
executor = concurrent.futures.ThreadPoolExecutor(max_workers=10)
|
||||
|
||||
|
||||
class RealtimeEventNormalizer(Protocol):
|
||||
def should_drop(self, event: object) -> bool: ...
|
||||
|
|
@ -315,13 +311,12 @@ class RealTimeStreaming:
|
|||
if self.session_tools or self.tool_calls:
|
||||
self.logging_obj.model_call_details["realtime_tools"] = self.session_tools
|
||||
self.logging_obj.model_call_details["realtime_tool_calls"] = self.tool_calls
|
||||
## ASYNC LOGGING
|
||||
# Route through the bounded logging worker (per-coroutine timeout +
|
||||
# concurrency cap) instead of a bare create_task, so a slow callback
|
||||
# can't leave suspended tasks pinning each call's response in memory.
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(self.logging_obj.async_success_handler(self.messages))
|
||||
## SYNC LOGGING
|
||||
executor.submit(self.logging_obj.success_handler(self.messages))
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
|
||||
self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True)
|
||||
)
|
||||
|
||||
async def _send_to_backend(self, message: str) -> bool:
|
||||
"""Send a message to the backend WebSocket.
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ from ..common_utils import AnthropicError, AnthropicModelInfo
|
|||
|
||||
ANTHROPIC_FILES_API_BASE = "https://api.anthropic.com"
|
||||
ANTHROPIC_FILES_BETA_HEADER = "files-api-2025-04-14"
|
||||
ANTHROPIC_MESSAGE_BATCH_ID_PREFIX = "msgbatch_"
|
||||
|
||||
|
||||
class AnthropicFilesConfig(BaseFilesConfig):
|
||||
|
|
@ -258,6 +259,8 @@ class AnthropicFilesConfig(BaseFilesConfig):
|
|||
file_id = file_content_request.get("file_id")
|
||||
api_base = AnthropicModelInfo.get_api_base(litellm_params.get("api_base")) or ANTHROPIC_FILES_API_BASE
|
||||
encoded_file_id = encode_url_path_segment(file_id, field_name="file_id")
|
||||
if file_id.startswith(ANTHROPIC_MESSAGE_BATCH_ID_PREFIX):
|
||||
return f"{api_base.rstrip('/')}/v1/messages/batches/{encoded_file_id}/results", {}
|
||||
return f"{api_base.rstrip('/')}/v1/files/{encoded_file_id}/content", {}
|
||||
|
||||
def transform_file_content_response(
|
||||
|
|
|
|||
|
|
@ -205,7 +205,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
#########################################################
|
||||
########## DELETE RESPONSE API TRANSFORMATION ##############
|
||||
#########################################################
|
||||
def _construct_url_for_response_id_in_path(self, api_base: str, response_id: str) -> str:
|
||||
def _construct_url_for_response_id_in_path(self, api_base: str, response_id: str, path_suffix: str = "") -> str:
|
||||
"""
|
||||
Constructs a URL for the API request with the response_id in the path.
|
||||
"""
|
||||
|
|
@ -218,14 +218,14 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
# Remove trailing slash if present to avoid double slashes
|
||||
path = parsed_url.path.rstrip("/")
|
||||
encoded_response_id = encode_url_path_segment(response_id, field_name="response_id")
|
||||
new_path = f"{path}/{encoded_response_id}"
|
||||
new_path = f"{path}/{encoded_response_id}{path_suffix}"
|
||||
|
||||
# Reconstruct the URL with all original components but with the modified path
|
||||
constructed_url = urlunparse(
|
||||
(
|
||||
parsed_url.scheme, # http, https
|
||||
parsed_url.netloc, # domain name, port
|
||||
new_path, # path with response_id added
|
||||
new_path,
|
||||
parsed_url.params, # parameters
|
||||
parsed_url.query, # query string
|
||||
parsed_url.fragment, # fragment
|
||||
|
|
@ -288,7 +288,9 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
limit: int = 20,
|
||||
order: Literal["asc", "desc"] = "desc",
|
||||
) -> Tuple[str, Dict]:
|
||||
url = self._construct_url_for_response_id_in_path(api_base=api_base, response_id=response_id) + "/input_items"
|
||||
url = self._construct_url_for_response_id_in_path(
|
||||
api_base=api_base, response_id=response_id, path_suffix="/input_items"
|
||||
)
|
||||
params: Dict[str, Any] = {}
|
||||
if after is not None:
|
||||
params["after"] = after
|
||||
|
|
@ -322,27 +324,8 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
This function handles URLs with query parameters by inserting the response_id
|
||||
at the correct location (before any query parameters).
|
||||
"""
|
||||
from urllib.parse import urlparse, urlunparse
|
||||
|
||||
# Parse the URL to separate its components
|
||||
parsed_url = urlparse(api_base)
|
||||
|
||||
# Insert the response_id and /cancel at the end of the path component
|
||||
# Remove trailing slash if present to avoid double slashes
|
||||
path = parsed_url.path.rstrip("/")
|
||||
encoded_response_id = encode_url_path_segment(response_id, field_name="response_id")
|
||||
new_path = f"{path}/{encoded_response_id}/cancel"
|
||||
|
||||
# Reconstruct the URL with all original components but with the modified path
|
||||
cancel_url = urlunparse(
|
||||
(
|
||||
parsed_url.scheme, # http, https
|
||||
parsed_url.netloc, # domain name, port
|
||||
new_path, # path with response_id and /cancel added
|
||||
parsed_url.params, # parameters
|
||||
parsed_url.query, # query string
|
||||
parsed_url.fragment, # fragment
|
||||
)
|
||||
cancel_url = self._construct_url_for_response_id_in_path(
|
||||
api_base=api_base, response_id=response_id, path_suffix="/cancel"
|
||||
)
|
||||
|
||||
data: Dict = {}
|
||||
|
|
|
|||
|
|
@ -7,13 +7,14 @@ The bedrock-mantle endpoint uses the Anthropic Messages API format but is served
|
|||
at a different endpoint (bedrock-mantle.{region}.api.aws) with AWS SigV4 auth.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
from typing import TYPE_CHECKING, Any, AsyncIterator, Iterator, List, Optional
|
||||
|
||||
from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import (
|
||||
AmazonAnthropicClaudeConfig,
|
||||
)
|
||||
from litellm.llms.bedrock.common_utils import build_mantle_messages_url
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
|
@ -91,10 +92,14 @@ class AmazonMantleConfig(AmazonAnthropicClaudeConfig):
|
|||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
# The parent strips "model" from the body (Invoke API puts it in URL).
|
||||
# The mantle endpoint (Messages API) requires "model" in the body.
|
||||
request["model"] = model_id
|
||||
return request
|
||||
# The parent strips "model" and "stream" from the body (Invoke API puts
|
||||
# the model in the URL and streams via a dedicated endpoint). The mantle
|
||||
# endpoint (Messages API) requires both in the body.
|
||||
return self._restore_mantle_body_fields(
|
||||
request=request,
|
||||
model_id=model_id,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
|
||||
async def async_transform_request(
|
||||
self,
|
||||
|
|
@ -114,5 +119,31 @@ class AmazonMantleConfig(AmazonAnthropicClaudeConfig):
|
|||
headers=headers,
|
||||
)
|
||||
await self._async_convert_document_url_sources_to_base64(request)
|
||||
request["model"] = model_id
|
||||
return request
|
||||
return self._restore_mantle_body_fields(
|
||||
request=request,
|
||||
model_id=model_id,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _restore_mantle_body_fields(request: dict, model_id: str, optional_params: dict) -> dict:
|
||||
stream_fields: dict = {"stream": True} if optional_params.get("stream") is True else {}
|
||||
return {**request, "model": model_id, **stream_fields}
|
||||
|
||||
@property
|
||||
def has_custom_stream_wrapper(self) -> bool:
|
||||
return False
|
||||
|
||||
def get_model_response_iterator(
|
||||
self,
|
||||
streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse,
|
||||
sync_stream: bool,
|
||||
json_mode: Optional[bool] = False,
|
||||
) -> Any:
|
||||
from litellm.llms.anthropic.chat.handler import ModelResponseIterator
|
||||
|
||||
return ModelResponseIterator(
|
||||
streaming_response=streaming_response,
|
||||
sync_stream=sync_stream,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -6,8 +6,13 @@ AmazonAnthropicClaudeMessagesConfig. Overrides only the URL and model-prefix
|
|||
stripping that are specific to the bedrock-mantle endpoint.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
|
||||
from typing import TYPE_CHECKING, Any, AsyncIterator, Dict, List, Optional, Tuple
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
|
||||
AnthropicMessagesConfig,
|
||||
)
|
||||
from litellm.llms.bedrock.common_utils import build_mantle_messages_url
|
||||
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
|
||||
AmazonAnthropicClaudeMessagesConfig,
|
||||
|
|
@ -89,8 +94,26 @@ class AmazonMantleMessagesConfig(AmazonAnthropicClaudeMessagesConfig):
|
|||
headers=headers,
|
||||
)
|
||||
|
||||
# Parent (AmazonAnthropicClaudeMessagesConfig) removes "model" from the
|
||||
# body (Bedrock Invoke puts model in the URL). The mantle endpoint
|
||||
# (Messages API) requires "model" in the request body.
|
||||
request["model"] = model_id
|
||||
return request
|
||||
# Parent (AmazonAnthropicClaudeMessagesConfig) removes "model" and
|
||||
# "stream" from the body (Bedrock Invoke puts the model in the URL and
|
||||
# streams via a dedicated endpoint). The mantle endpoint (Messages API)
|
||||
# requires both in the request body.
|
||||
stream_fields: dict[str, bool] = (
|
||||
{"stream": True} if anthropic_messages_optional_request_params.get("stream") is True else {}
|
||||
)
|
||||
return {**request, "model": model_id, **stream_fields}
|
||||
|
||||
def get_async_streaming_response_iterator(
|
||||
self,
|
||||
model: str,
|
||||
httpx_response: httpx.Response,
|
||||
request_body: dict,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
) -> AsyncIterator:
|
||||
return AnthropicMessagesConfig.get_async_streaming_response_iterator(
|
||||
self,
|
||||
model=model,
|
||||
httpx_response=httpx_response,
|
||||
request_body=request_body,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from litellm._logging import _redact_string, verbose_proxy_logger
|
|||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
|
||||
from ..base_aws_llm import BaseAWSLLM
|
||||
from ..common_utils import BedrockError
|
||||
from .transformation import BedrockRealtimeConfig
|
||||
|
||||
|
||||
|
|
@ -59,9 +60,7 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
InvokeModelWithBidirectionalStreamOperationInput,
|
||||
)
|
||||
from aws_sdk_bedrock_runtime.config import Config
|
||||
from smithy_aws_core.identity.environment import (
|
||||
EnvironmentCredentialsResolver,
|
||||
)
|
||||
from smithy_aws_core.identity import StaticCredentialsResolver
|
||||
except ImportError:
|
||||
raise ImportError("Missing aws_sdk_bedrock_runtime. Install with: pip install aws-sdk-bedrock-runtime")
|
||||
|
||||
|
|
@ -82,11 +81,36 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
|
||||
verbose_proxy_logger.debug(f"Bedrock Realtime: Connecting to {endpoint_uri} with model {model}")
|
||||
|
||||
credentials = self.get_credentials(
|
||||
aws_access_key_id=aws_access_key_id,
|
||||
aws_secret_access_key=aws_secret_access_key,
|
||||
aws_session_token=aws_session_token,
|
||||
aws_region_name=aws_region_name,
|
||||
aws_session_name=aws_session_name,
|
||||
aws_profile_name=aws_profile_name,
|
||||
aws_role_name=aws_role_name,
|
||||
aws_web_identity_token=aws_web_identity_token,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
aws_external_id=aws_external_id,
|
||||
)
|
||||
if credentials is None:
|
||||
raise BedrockError(
|
||||
status_code=401,
|
||||
message=(
|
||||
"No AWS credentials found for Bedrock realtime. Set aws_* params in litellm_params "
|
||||
"or configure credentials in the environment"
|
||||
),
|
||||
)
|
||||
frozen_credentials = credentials.get_frozen_credentials()
|
||||
|
||||
# Initialize Bedrock client with aws_sdk_bedrock_runtime
|
||||
config = Config(
|
||||
endpoint_uri=endpoint_uri,
|
||||
region=aws_region_name,
|
||||
aws_credentials_identity_resolver=EnvironmentCredentialsResolver(),
|
||||
aws_access_key_id=frozen_credentials.access_key,
|
||||
aws_secret_access_key=frozen_credentials.secret_key,
|
||||
aws_session_token=frozen_credentials.token,
|
||||
aws_credentials_identity_resolver=StaticCredentialsResolver(),
|
||||
)
|
||||
bedrock_client = BedrockRuntimeClient(config=config)
|
||||
|
||||
|
|
|
|||
|
|
@ -1300,16 +1300,15 @@ class BaseLLMHTTPHandler:
|
|||
if client is None or not isinstance(client, HTTPHandler):
|
||||
client = _get_httpx_client()
|
||||
|
||||
json_data = data if files is None and isinstance(data, dict) else None
|
||||
|
||||
try:
|
||||
# Make the POST request - clean and simple, always use data and files
|
||||
response = client.post(
|
||||
url=complete_url,
|
||||
headers=headers,
|
||||
data=data,
|
||||
data=data if json_data is None else None,
|
||||
files=files,
|
||||
json=(
|
||||
data if files is None and isinstance(data, dict) else None
|
||||
), # Use json param only when no files and data is dict
|
||||
json=json_data,
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as e:
|
||||
|
|
@ -1373,16 +1372,15 @@ class BaseLLMHTTPHandler:
|
|||
else:
|
||||
async_httpx_client = client
|
||||
|
||||
json_data = data if files is None and isinstance(data, dict) else None
|
||||
|
||||
try:
|
||||
# Make the async POST request - clean and simple, always use data and files
|
||||
response = await async_httpx_client.post(
|
||||
url=complete_url,
|
||||
headers=headers,
|
||||
data=data,
|
||||
data=data if json_data is None else None,
|
||||
files=files,
|
||||
json=(
|
||||
data if files is None and isinstance(data, dict) else None
|
||||
), # Use json param only when no files and data is dict
|
||||
json=json_data,
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as e:
|
||||
|
|
@ -2914,6 +2912,7 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
try:
|
||||
response = sync_httpx_client.get(url=url, headers=headers, params=data)
|
||||
response.raise_for_status()
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
|
|
@ -2985,9 +2984,9 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
try:
|
||||
response = await async_httpx_client.get(url=url, headers=headers, params=data)
|
||||
|
||||
response.raise_for_status()
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error retrieving response: {e}")
|
||||
verbose_logger.debug(f"Error retrieving response: {e}")
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
provider_config=responses_api_provider_config,
|
||||
|
|
@ -3078,6 +3077,7 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
try:
|
||||
response = sync_httpx_client.get(url=url, headers=headers, params=params)
|
||||
response.raise_for_status()
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=responses_api_provider_config)
|
||||
|
||||
|
|
@ -3151,6 +3151,7 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
try:
|
||||
response = await async_httpx_client.get(url=url, headers=headers, params=params)
|
||||
response.raise_for_status()
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=responses_api_provider_config)
|
||||
|
||||
|
|
@ -4737,6 +4738,13 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=provider_config)
|
||||
|
||||
if response.status_code >= 400:
|
||||
raise provider_config.get_error_class(
|
||||
error_message=response.text,
|
||||
status_code=response.status_code,
|
||||
headers=response.headers,
|
||||
)
|
||||
|
||||
return provider_config.transform_file_content_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -4793,6 +4801,13 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=provider_config)
|
||||
|
||||
if response.status_code >= 400:
|
||||
raise provider_config.get_error_class(
|
||||
error_message=response.text,
|
||||
status_code=response.status_code,
|
||||
headers=response.headers,
|
||||
)
|
||||
|
||||
return provider_config.transform_file_content_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@
|
|||
Support for gpt model family
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
|
|
@ -782,8 +783,30 @@ class OpenAIChatCompletionStreamingHandler(BaseModelResponseIterator):
|
|||
delta["reasoning_content"] = delta.pop("reasoning")
|
||||
return choices
|
||||
|
||||
@staticmethod
|
||||
def _extract_error_from_chunk(chunk: dict) -> Optional[tuple[str, int]]:
|
||||
"""OpenAI-compatible backends (vLLM, sglang) can return an HTTP 200
|
||||
stream whose body carries an error payload, e.g.
|
||||
``data: {"error": {"message": "...", "code": 400}}``."""
|
||||
error = chunk.get("error")
|
||||
if not error:
|
||||
return None
|
||||
if not isinstance(error, dict):
|
||||
return str(error), 500
|
||||
message = error.get("message")
|
||||
code = error.get("code")
|
||||
status_code = code if isinstance(code, int) and 400 <= code < 600 else 500
|
||||
return (message if isinstance(message, str) else json.dumps(error)), status_code
|
||||
|
||||
def chunk_parser(self, chunk: dict) -> ModelResponseStream:
|
||||
try:
|
||||
error_details = self._extract_error_from_chunk(chunk)
|
||||
if error_details is not None:
|
||||
error_message, error_status_code = error_details
|
||||
raise OpenAIError(
|
||||
status_code=error_status_code,
|
||||
message=error_message,
|
||||
)
|
||||
choices = chunk.get("choices", [])
|
||||
choices = self._map_reasoning_to_reasoning_content(choices)
|
||||
|
||||
|
|
|
|||
|
|
@ -197,12 +197,13 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
return data
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> List[str]:
|
||||
"""Extract tool names from Responses API request (tools[].name for function, tools[].server_label for mcp)."""
|
||||
"""Extract tool names from Responses API request (tools[].name for function
|
||||
and custom, tools[].server_label for mcp)."""
|
||||
names: List[str] = []
|
||||
for tool in data.get("tools") or []:
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
if tool.get("type") == "function" and tool.get("name"):
|
||||
if tool.get("type") in ("function", "custom") and tool.get("name"):
|
||||
names.append(str(tool["name"]))
|
||||
elif tool.get("type") == "mcp" and tool.get("server_label"):
|
||||
names.append(str(tool["server_label"]))
|
||||
|
|
|
|||
194
litellm/llms/vertex_ai/audio_transcription/transformation.py
Normal file
194
litellm/llms/vertex_ai/audio_transcription/transformation.py
Normal file
|
|
@ -0,0 +1,194 @@
|
|||
import base64
|
||||
|
||||
from httpx import Headers, Response
|
||||
|
||||
import litellm
|
||||
from litellm.exceptions import UnsupportedParamsError
|
||||
from litellm.litellm_core_utils.audio_utils.utils import (
|
||||
normalize_transcription_language_to_bcp47,
|
||||
process_audio_file,
|
||||
)
|
||||
from litellm.llms.base_llm.audio_transcription.transformation import (
|
||||
AudioTranscriptionRequestData,
|
||||
BaseAudioTranscriptionConfig,
|
||||
)
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.vertex_ai.common_utils import VertexAIError, validate_vertex_location
|
||||
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
OpenAIAudioTranscriptionOptionalParams,
|
||||
)
|
||||
from litellm.types.llms.vertex_ai_speech_to_text import (
|
||||
VertexSpeechToTextAutoDecodingConfig,
|
||||
VertexSpeechToTextRecognitionConfig,
|
||||
VertexSpeechToTextRecognitionFeatures,
|
||||
VertexSpeechToTextRecognizeRequest,
|
||||
VertexSpeechToTextRecognizeResponse,
|
||||
)
|
||||
from litellm.types.utils import FileTypes, TranscriptionResponse
|
||||
|
||||
DEFAULT_SPEECH_TO_TEXT_LOCATION = "us"
|
||||
AUTO_LANGUAGE_CODE = "auto"
|
||||
SUPPORTED_RESPONSE_FORMATS = ("json", "text")
|
||||
_URL_UNSAFE_PROJECT_CHARS = ("/", "?", "#", "\\", ":", " ", "\t", "\n", "\r")
|
||||
|
||||
|
||||
class VertexAIAudioTranscriptionConfig(BaseAudioTranscriptionConfig, VertexBase):
|
||||
def __init__(self) -> None:
|
||||
BaseAudioTranscriptionConfig.__init__(self)
|
||||
VertexBase.__init__(self)
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]:
|
||||
return ["language", "response_format"]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
supported_params = self.get_supported_openai_params(model)
|
||||
mapped = {
|
||||
**optional_params,
|
||||
**{k: v for k, v in non_default_params.items() if k in supported_params},
|
||||
}
|
||||
response_format = mapped.get("response_format")
|
||||
if response_format is None or response_format in SUPPORTED_RESPONSE_FORMATS:
|
||||
return mapped
|
||||
if drop_params or litellm.drop_params:
|
||||
return {k: v for k, v in mapped.items() if k != "response_format"}
|
||||
raise UnsupportedParamsError(
|
||||
status_code=400,
|
||||
message=(
|
||||
f"Google Speech-to-Text does not support response_format={response_format!r}. "
|
||||
f"Supported values: {', '.join(SUPPORTED_RESPONSE_FORMATS)}. "
|
||||
"To drop unsupported openai params from the call, set `litellm.drop_params = True`"
|
||||
),
|
||||
)
|
||||
|
||||
def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException:
|
||||
return VertexAIError(status_code=status_code, message=error_message, headers=headers)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
) -> dict:
|
||||
access_token, project_id = self._ensure_access_token(
|
||||
credentials=self.safe_get_vertex_ai_credentials(litellm_params),
|
||||
project_id=self.safe_get_vertex_ai_project(litellm_params),
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
return {
|
||||
**headers,
|
||||
"Authorization": f"Bearer {access_token}",
|
||||
"x-goog-user-project": project_id,
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: bool | None = None,
|
||||
) -> str:
|
||||
location = self._validate_location(self.safe_get_vertex_ai_location(litellm_params))
|
||||
project_id = self._validate_project_id(
|
||||
self.safe_get_vertex_ai_project(litellm_params) or self._resolve_project_id_from_credentials(litellm_params)
|
||||
)
|
||||
host = "speech.googleapis.com" if location == "global" else f"{location}-speech.googleapis.com"
|
||||
base_url = (api_base or f"https://{host}").rstrip("/")
|
||||
return f"{base_url}/v2/projects/{project_id}/locations/{location}/recognizers/_:recognize"
|
||||
|
||||
@staticmethod
|
||||
def _validate_location(location: str | None) -> str:
|
||||
try:
|
||||
return validate_vertex_location(location or DEFAULT_SPEECH_TO_TEXT_LOCATION)
|
||||
except ValueError as e:
|
||||
raise VertexAIError(status_code=400, message=str(e)) from e
|
||||
|
||||
@staticmethod
|
||||
def _validate_project_id(project_id: str) -> str:
|
||||
if not project_id or ".." in project_id or any(c in project_id for c in _URL_UNSAFE_PROJECT_CHARS):
|
||||
raise VertexAIError(status_code=400, message=f"Invalid vertex_project format: {project_id!r}")
|
||||
return project_id
|
||||
|
||||
def _resolve_project_id_from_credentials(self, litellm_params: dict) -> str:
|
||||
_, project_id = self._ensure_access_token(
|
||||
credentials=self.safe_get_vertex_ai_credentials(litellm_params),
|
||||
project_id=None,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
return project_id
|
||||
|
||||
def transform_audio_transcription_request(
|
||||
self,
|
||||
model: str,
|
||||
audio_file: FileTypes,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
) -> AudioTranscriptionRequestData:
|
||||
processed_audio = process_audio_file(audio_file)
|
||||
language = optional_params.get("language")
|
||||
language_codes = (
|
||||
[normalize_transcription_language_to_bcp47(language)]
|
||||
if isinstance(language, str) and language
|
||||
else [AUTO_LANGUAGE_CODE]
|
||||
)
|
||||
request_body = VertexSpeechToTextRecognizeRequest(
|
||||
config=VertexSpeechToTextRecognitionConfig(
|
||||
model=model.removeprefix("vertex_ai/"),
|
||||
languageCodes=language_codes,
|
||||
features=VertexSpeechToTextRecognitionFeatures(enableAutomaticPunctuation=True),
|
||||
autoDecodingConfig=VertexSpeechToTextAutoDecodingConfig(),
|
||||
),
|
||||
content=base64.b64encode(processed_audio.file_content).decode("utf-8"),
|
||||
)
|
||||
return AudioTranscriptionRequestData(data=dict(request_body))
|
||||
|
||||
def transform_audio_transcription_response(
|
||||
self,
|
||||
raw_response: Response,
|
||||
) -> TranscriptionResponse:
|
||||
try:
|
||||
response_json = raw_response.json()
|
||||
except ValueError:
|
||||
raise VertexAIError(
|
||||
status_code=raw_response.status_code,
|
||||
message=f"Received non-JSON response from Google Speech-to-Text: {raw_response.text}",
|
||||
)
|
||||
parsed = VertexSpeechToTextRecognizeResponse.model_validate(response_json)
|
||||
transcripts = tuple(
|
||||
result.alternatives[0].transcript
|
||||
for result in parsed.results
|
||||
if result.alternatives and result.alternatives[0].transcript
|
||||
)
|
||||
response = TranscriptionResponse(text=" ".join(transcripts))
|
||||
response["task"] = "transcribe"
|
||||
detected_language = next((result.languageCode for result in parsed.results if result.languageCode), None)
|
||||
if detected_language is not None:
|
||||
response["language"] = detected_language
|
||||
billed_duration = _parse_duration_seconds(parsed.metadata.totalBilledDuration if parsed.metadata else None)
|
||||
if billed_duration is not None:
|
||||
response["duration"] = billed_duration
|
||||
response._hidden_params = response_json
|
||||
return response
|
||||
|
||||
|
||||
def _parse_duration_seconds(duration: str | None) -> float | None:
|
||||
if duration is None or not duration.endswith("s"):
|
||||
return None
|
||||
try:
|
||||
return float(duration[:-1])
|
||||
except ValueError:
|
||||
return None
|
||||
|
|
@ -311,6 +311,28 @@ def get_vertex_base_model_name(model: str) -> str:
|
|||
return model
|
||||
|
||||
|
||||
def validate_vertex_location(vertex_location: Optional[str]) -> str:
|
||||
"""
|
||||
Validate a Vertex AI location before interpolating it into a request host or
|
||||
URL path.
|
||||
|
||||
``vertex_location`` is client-controllable on the proxy (it flows in from the
|
||||
request body), so it must never be trusted verbatim in a URL or an attacker
|
||||
could point the host at their own server and exfiltrate the admin's Google
|
||||
access token. Allow the special ``global`` control plane and otherwise require
|
||||
a lowercase alphanumeric-plus-hyphen token (e.g. ``us``, ``us-central1``,
|
||||
``eu``), which rejects host injection like ``attacker.example/`` or
|
||||
``evil.com#``.
|
||||
"""
|
||||
if vertex_location == "global":
|
||||
return vertex_location
|
||||
if vertex_location is None:
|
||||
raise ValueError("vertex_location is required")
|
||||
if not re.match(r"^[a-z][a-z0-9-]*$", vertex_location):
|
||||
raise ValueError("Invalid vertex_location format")
|
||||
return vertex_location
|
||||
|
||||
|
||||
def get_vertex_base_url(
|
||||
vertex_location: Optional[str],
|
||||
) -> str:
|
||||
|
|
@ -321,15 +343,12 @@ def get_vertex_base_url(
|
|||
- Multi-region geographies (e.g. ``us``, ``eu``) use ``aiplatform.{geo}.rep.googleapis.com``.
|
||||
- Regional locations (e.g. ``us-central1``) use ``{region}-aiplatform.googleapis.com``.
|
||||
"""
|
||||
if vertex_location == "global":
|
||||
validated_location = validate_vertex_location(vertex_location)
|
||||
if validated_location == "global":
|
||||
return "https://aiplatform.googleapis.com"
|
||||
if vertex_location is None:
|
||||
raise ValueError("vertex_location is required")
|
||||
if not re.match(r"^[a-z][a-z0-9-]*$", vertex_location):
|
||||
raise ValueError("Invalid vertex_location format")
|
||||
if "-" not in vertex_location:
|
||||
return f"https://aiplatform.{vertex_location}.rep.googleapis.com"
|
||||
return f"https://{vertex_location}-aiplatform.googleapis.com"
|
||||
if "-" not in validated_location:
|
||||
return f"https://aiplatform.{validated_location}.rep.googleapis.com"
|
||||
return f"https://{validated_location}-aiplatform.googleapis.com"
|
||||
|
||||
|
||||
def _get_embedding_url(
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ import json
|
|||
import os
|
||||
import threading
|
||||
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -315,6 +316,9 @@ class VertexBase:
|
|||
api_base=api_base,
|
||||
)
|
||||
|
||||
if partner == VertexPartnerProvider.llama:
|
||||
return default_api_base
|
||||
|
||||
if len(default_api_base.split(":")) > 1:
|
||||
endpoint = default_api_base.split(":")[-1]
|
||||
else:
|
||||
|
|
@ -615,7 +619,8 @@ class VertexBase:
|
|||
|
||||
Handles custom api_base for:
|
||||
1. Gemini (Google AI Studio) - constructs /models/{model}:{endpoint}
|
||||
2. Vertex AI with standard proxies - constructs {api_base}:{endpoint}
|
||||
2. Vertex AI with standard proxies - constructs {api_base}:{endpoint};
|
||||
if api_base has no path (bare host), grafts the default vertex URL path onto it
|
||||
3. Vertex AI with PSC endpoints - constructs full path structure
|
||||
{api_base}/v1/projects/{project}/locations/{location}/endpoints/{model}:{endpoint}
|
||||
(only when use_psc_endpoint_format=True)
|
||||
|
|
@ -660,8 +665,9 @@ class VertexBase:
|
|||
model_for_url,
|
||||
endpoint,
|
||||
)
|
||||
elif urlparse(api_base).path in ("", "/"):
|
||||
url = api_base.rstrip("/") + urlparse(url).path
|
||||
else:
|
||||
# Fallback to simple format if we don't have all parameters
|
||||
url = "{}:{}".format(api_base, endpoint)
|
||||
if stream is True:
|
||||
url = url + "?alt=sse"
|
||||
|
|
|
|||
|
|
@ -582,6 +582,7 @@ async def acompletion(
|
|||
"api_key": api_key,
|
||||
"model_list": model_list,
|
||||
"reasoning_effort": reasoning_effort,
|
||||
"verbosity": verbosity,
|
||||
"safety_identifier": safety_identifier,
|
||||
"service_tier": service_tier,
|
||||
"extra_headers": extra_headers,
|
||||
|
|
@ -1081,6 +1082,54 @@ def _build_custom_pricing_entry(
|
|||
return entry
|
||||
|
||||
|
||||
def _get_router_deployment_id(kwargs: dict) -> Optional[str]:
|
||||
for metadata_key in ("litellm_metadata", "metadata"):
|
||||
metadata = kwargs.get(metadata_key) or {}
|
||||
if not isinstance(metadata, dict):
|
||||
continue
|
||||
deployment_model_info = metadata.get("model_info") or {}
|
||||
if not isinstance(deployment_model_info, dict):
|
||||
continue
|
||||
deployment_id = deployment_model_info.get("id")
|
||||
if deployment_id is not None:
|
||||
return str(deployment_id)
|
||||
return None
|
||||
|
||||
|
||||
def _register_custom_pricing_for_request(
|
||||
model: str,
|
||||
custom_llm_provider: str,
|
||||
kwargs: dict,
|
||||
model_info: Optional[dict],
|
||||
) -> None:
|
||||
"""Register per-request custom pricing in litellm.model_cost.
|
||||
|
||||
Router-originated requests (identified by the deployment id the router puts
|
||||
in metadata) get their full pricing registered under that unique id only;
|
||||
the shared ``{provider}/{model}`` key receives the entry with pricing fields
|
||||
stripped, mirroring Router._create_deployment. This keeps one deployment's
|
||||
pricing overrides (e.g. a zero-cost wildcard) from clobbering built-in
|
||||
pricing used by sibling deployments of the same backend model. Direct SDK
|
||||
calls keep the legacy behavior of registering the shared key with pricing.
|
||||
"""
|
||||
entry = _build_custom_pricing_entry(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs=kwargs,
|
||||
model_info=model_info,
|
||||
)
|
||||
shared_key = f"{custom_llm_provider}/{model}"
|
||||
deployment_id = _get_router_deployment_id(kwargs)
|
||||
if deployment_id is None:
|
||||
litellm.register_model({shared_key: entry})
|
||||
return
|
||||
litellm.register_model(
|
||||
{
|
||||
deployment_id: entry,
|
||||
shared_key: CustomPricingLiteLLMParams.strip_custom_pricing_fields(entry),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _complete_azure(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
|
||||
_azure_detection_model = ctx._azure_detection_model
|
||||
acompletion = ctx.acompletion
|
||||
|
|
@ -5107,14 +5156,11 @@ def completion( # type: ignore
|
|||
if (
|
||||
input_cost_per_token is not None and output_cost_per_token is not None
|
||||
) or input_cost_per_second is not None:
|
||||
litellm.register_model(
|
||||
{
|
||||
f"{custom_llm_provider}/{model}": _build_custom_pricing_entry(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs=kwargs,
|
||||
model_info=model_info,
|
||||
)
|
||||
}
|
||||
_register_custom_pricing_for_request(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs=kwargs,
|
||||
model_info=model_info,
|
||||
)
|
||||
### BUILD CUSTOM PROMPT TEMPLATE -- IF GIVEN ###
|
||||
custom_prompt_dict = {} # type: ignore
|
||||
|
|
@ -5193,6 +5239,7 @@ def completion( # type: ignore
|
|||
"parallel_tool_calls": parallel_tool_calls,
|
||||
"messages": messages,
|
||||
"reasoning_effort": reasoning_effort,
|
||||
"verbosity": verbosity,
|
||||
"thinking": thinking,
|
||||
"web_search_options": web_search_options,
|
||||
"include_server_side_tool_invocations": (
|
||||
|
|
@ -5957,14 +6004,11 @@ def embedding(
|
|||
|
||||
### REGISTER CUSTOM MODEL PRICING -- IF GIVEN ###
|
||||
if (input_cost_per_token is not None and output_cost_per_token is not None) or input_cost_per_second is not None:
|
||||
litellm.register_model(
|
||||
{
|
||||
f"{custom_llm_provider}/{model}": _build_custom_pricing_entry(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs=kwargs,
|
||||
model_info=kwargs.get("model_info"),
|
||||
)
|
||||
}
|
||||
_register_custom_pricing_for_request(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs=kwargs,
|
||||
model_info=kwargs.get("model_info"),
|
||||
)
|
||||
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
|
|
|
|||
|
|
@ -5761,6 +5761,76 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"azure/us/gpt-5.4": {
|
||||
"cache_read_input_token_cost": 2.8e-07,
|
||||
"cache_read_input_token_cost_priority": 5.5e-07,
|
||||
"input_cost_per_token": 2.75e-06,
|
||||
"input_cost_per_token_priority": 5.5e-06,
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"output_cost_per_token_priority": 3.3e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"azure/eu/gpt-5.4": {
|
||||
"cache_read_input_token_cost": 2.8e-07,
|
||||
"cache_read_input_token_cost_priority": 5.5e-07,
|
||||
"input_cost_per_token": 2.75e-06,
|
||||
"input_cost_per_token_priority": 5.5e-06,
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"output_cost_per_token_priority": 3.3e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"azure/gpt-5.4-2026-03-05": {
|
||||
"cache_read_input_token_cost": 2.5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 5e-07,
|
||||
|
|
@ -5802,6 +5872,76 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"azure/us/gpt-5.4-2026-03-05": {
|
||||
"cache_read_input_token_cost": 2.8e-07,
|
||||
"cache_read_input_token_cost_priority": 5.5e-07,
|
||||
"input_cost_per_token": 2.75e-06,
|
||||
"input_cost_per_token_priority": 5.5e-06,
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"output_cost_per_token_priority": 3.3e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"azure/eu/gpt-5.4-2026-03-05": {
|
||||
"cache_read_input_token_cost": 2.8e-07,
|
||||
"cache_read_input_token_cost_priority": 5.5e-07,
|
||||
"input_cost_per_token": 2.75e-06,
|
||||
"input_cost_per_token_priority": 5.5e-06,
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"output_cost_per_token_priority": 3.3e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"azure/gpt-5.4-pro": {
|
||||
"cache_read_input_token_cost": 3e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 6e-06,
|
||||
|
|
@ -5917,6 +6057,90 @@
|
|||
"supports_xhigh_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": false
|
||||
},
|
||||
"azure/us/gpt-5.5": {
|
||||
"cache_read_input_token_cost": 5.5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 1.1e-06,
|
||||
"cache_read_input_token_cost_priority": 1.38e-06,
|
||||
"input_cost_per_token": 5.5e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 1.1e-05,
|
||||
"input_cost_per_token_priority": 1.375e-05,
|
||||
"output_cost_per_token": 3.3e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 4.95e-05,
|
||||
"output_cost_per_token_priority": 8.25e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_none_reasoning_effort": true,
|
||||
"supports_xhigh_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": false
|
||||
},
|
||||
"azure/eu/gpt-5.5": {
|
||||
"cache_read_input_token_cost": 5.5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 1.1e-06,
|
||||
"cache_read_input_token_cost_priority": 1.38e-06,
|
||||
"input_cost_per_token": 5.5e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 1.1e-05,
|
||||
"input_cost_per_token_priority": 1.375e-05,
|
||||
"output_cost_per_token": 3.3e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 4.95e-05,
|
||||
"output_cost_per_token_priority": 8.25e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_none_reasoning_effort": true,
|
||||
"supports_xhigh_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": false
|
||||
},
|
||||
"azure/gpt-5.5-2026-04-23": {
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 1e-06,
|
||||
|
|
@ -5959,6 +6183,84 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"azure/us/gpt-5.5-2026-04-23": {
|
||||
"cache_read_input_token_cost": 5.5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 1.1e-06,
|
||||
"cache_read_input_token_cost_priority": 1.38e-06,
|
||||
"input_cost_per_token": 5.5e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 1.1e-05,
|
||||
"input_cost_per_token_priority": 1.375e-05,
|
||||
"output_cost_per_token": 3.3e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 4.95e-05,
|
||||
"output_cost_per_token_priority": 8.25e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"azure/eu/gpt-5.5-2026-04-23": {
|
||||
"cache_read_input_token_cost": 5.5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 1.1e-06,
|
||||
"cache_read_input_token_cost_priority": 1.38e-06,
|
||||
"input_cost_per_token": 5.5e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 1.1e-05,
|
||||
"input_cost_per_token_priority": 1.375e-05,
|
||||
"output_cost_per_token": 3.3e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 4.95e-05,
|
||||
"output_cost_per_token_priority": 8.25e-05,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"azure/gpt-5.5-pro": {
|
||||
"cache_read_input_token_cost": 3e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 6e-06,
|
||||
|
|
@ -34655,6 +34957,19 @@
|
|||
"/v1/audio/speech"
|
||||
]
|
||||
},
|
||||
"vertex_ai/chirp_3": {
|
||||
"input_cost_per_second": 0.00026667,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"metadata": {
|
||||
"calculation": "$0.016/60 seconds = $0.00026667 per second",
|
||||
"original_pricing_per_minute": 0.016
|
||||
},
|
||||
"mode": "audio_transcription",
|
||||
"source": "https://cloud.google.com/speech-to-text/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/audio/transcriptions"
|
||||
]
|
||||
},
|
||||
"vertex_ai/claude-3-5-haiku": {
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "vertex_ai-anthropic_models",
|
||||
|
|
|
|||
|
|
@ -51,12 +51,12 @@ class LiteLLM_ManagedVectorStoreTable(LiteLLMPydanticObjectBase):
|
|||
class LiteLLM_ManagedVectorStoresTable(LiteLLMPydanticObjectBase):
|
||||
vector_store_id: str
|
||||
custom_llm_provider: str
|
||||
vector_store_name: Optional[str]
|
||||
vector_store_description: Optional[str]
|
||||
vector_store_metadata: Optional[Dict[str, Any]]
|
||||
created_at: Optional[datetime]
|
||||
updated_at: Optional[datetime]
|
||||
litellm_credential_name: Optional[str]
|
||||
litellm_params: Optional[Dict[str, Any]]
|
||||
team_id: Optional[str]
|
||||
user_id: Optional[str]
|
||||
vector_store_name: Optional[str] = None
|
||||
vector_store_description: Optional[str] = None
|
||||
vector_store_metadata: Optional[Dict[str, Any]] = None
|
||||
created_at: Optional[datetime] = None
|
||||
updated_at: Optional[datetime] = None
|
||||
litellm_credential_name: Optional[str] = None
|
||||
litellm_params: Optional[Dict[str, Any]] = None
|
||||
team_id: Optional[str] = None
|
||||
user_id: Optional[str] = None
|
||||
|
|
|
|||
|
|
@ -171,7 +171,6 @@ def llm_passthrough_route(
|
|||
api_key: Optional[str] = None,
|
||||
request_query_params: Optional[dict] = None,
|
||||
request_headers: Optional[dict] = None,
|
||||
allm_passthrough_route: bool = False,
|
||||
content: Optional[Any] = None,
|
||||
data: Optional[dict] = None,
|
||||
files: Optional[RequestFiles] = None,
|
||||
|
|
@ -198,7 +197,7 @@ def llm_passthrough_route(
|
|||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
_is_async = allm_passthrough_route
|
||||
_is_async = bool(kwargs.get("allm_passthrough_route", False))
|
||||
|
||||
litellm_logging_obj = cast("LiteLLMLoggingObj", kwargs.get("litellm_logging_obj"))
|
||||
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from starlette.types import Scope
|
|||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._types import (
|
||||
UI_TEAM_ID,
|
||||
LiteLLM_TeamTable,
|
||||
ProxyException,
|
||||
SpecialHeaders,
|
||||
|
|
@ -357,6 +358,7 @@ class MCPRequestHandler:
|
|||
# Inline imports avoid a circular dependency: mcp_server_manager imports
|
||||
# from this module.
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
MCPServerManager,
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
|
@ -382,7 +384,18 @@ class MCPRequestHandler:
|
|||
# fetches the upstream token automatically using stored credentials,
|
||||
# so allowing anonymous bypass would let any external caller invoke
|
||||
# tools authenticated as LiteLLM's service account.
|
||||
if server.has_client_credentials:
|
||||
#
|
||||
# Resolve the flow rather than reading has_client_credentials directly:
|
||||
# this is a security gate, and a legacy row whose oauth2_flow was never
|
||||
# stamped still carries the M2M credential shape (client_id/secret +
|
||||
# token_url, no authorization_url). Treating an unstamped-but-M2M-shaped
|
||||
# row as non-M2M here would reopen the anonymous bypass the explicit
|
||||
# column no longer closes on its own. Shares the one resolution helper
|
||||
# with the egress backstop and the anonymous-delegate allowlist; all fail
|
||||
# closed on the ambiguous shape and are removed together once no null rows
|
||||
# remain. A pure-PKCE delegate server (no stored credentials) resolves to a
|
||||
# non-M2M flow and keeps its bypass.
|
||||
if MCPServerManager.effective_oauth2_flow(server) == "client_credentials":
|
||||
return False
|
||||
return True
|
||||
|
||||
|
|
@ -726,6 +739,9 @@ class MCPRequestHandler:
|
|||
if not user_api_key_auth or not user_api_key_auth.team_id or not prisma_client:
|
||||
return None
|
||||
|
||||
if user_api_key_auth.team_id == UI_TEAM_ID:
|
||||
return None
|
||||
|
||||
# Get the team object (which has object_permission already loaded)
|
||||
team_obj: Optional[LiteLLM_TeamTable] = await get_team_object(
|
||||
team_id=user_api_key_auth.team_id,
|
||||
|
|
@ -1021,6 +1037,9 @@ class MCPRequestHandler:
|
|||
if user_api_key_auth is None or not user_api_key_auth.team_id or prisma_client is None:
|
||||
return []
|
||||
|
||||
if user_api_key_auth.team_id == UI_TEAM_ID:
|
||||
return []
|
||||
|
||||
team_obj: Optional[LiteLLM_TeamTable] = await get_team_object(
|
||||
team_id=user_api_key_auth.team_id,
|
||||
prisma_client=prisma_client,
|
||||
|
|
@ -1503,6 +1522,9 @@ class MCPRequestHandler:
|
|||
verbose_logger.debug("prisma_client is None")
|
||||
return []
|
||||
|
||||
if user_api_key_auth.team_id == UI_TEAM_ID:
|
||||
return []
|
||||
|
||||
try:
|
||||
team_obj: Optional[LiteLLM_TeamTable] = await get_team_object(
|
||||
team_id=user_api_key_auth.team_id,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import asyncio
|
||||
import html as _html
|
||||
import json
|
||||
import secrets
|
||||
import time
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
|
||||
|
|
@ -8,7 +9,7 @@ from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
|
|||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Form, HTTPException, Request
|
||||
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse
|
||||
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Response
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -137,6 +138,72 @@ def decode_state_hash(encrypted_state: str) -> dict:
|
|||
return state_data
|
||||
|
||||
|
||||
# LIT-4197: some upstream authorization servers reject an over-long ``state``
|
||||
# (the encrypted OAuth session blob routinely exceeds their limit). The upstream
|
||||
# only needs an opaque value it echoes back on ``/callback``, so we forward a
|
||||
# short random handle and keep the encrypted session in a per-flow HttpOnly
|
||||
# cookie bound to that handle. The browser carries the cookie across the
|
||||
# upstream round trip, so the flow stays correct with no server-side session
|
||||
# store (works across proxy replicas, unlike an in-process map).
|
||||
_OAUTH_STATE_COOKIE_PREFIX = "mcp_oauth_state_"
|
||||
_OAUTH_STATE_COOKIE_TTL_SECONDS = 600
|
||||
_OAUTH_STATE_HANDLE_BYTES = 32
|
||||
|
||||
|
||||
def _oauth_state_cookie_name(relay_state: str) -> str:
|
||||
return f"{_OAUTH_STATE_COOKIE_PREFIX}{relay_state}"
|
||||
|
||||
|
||||
def _oauth_state_cookie_path_and_secure(request: Request) -> tuple[str, bool]:
|
||||
parsed = urlparse(get_request_base_url(request))
|
||||
return parsed.path or "/", parsed.scheme == "https"
|
||||
|
||||
|
||||
def _set_oauth_state_cookie(
|
||||
response: Response,
|
||||
request: Request,
|
||||
relay_state: str,
|
||||
encoded_state: str,
|
||||
) -> None:
|
||||
path, secure = _oauth_state_cookie_path_and_secure(request)
|
||||
response.set_cookie(
|
||||
key=_oauth_state_cookie_name(relay_state),
|
||||
value=encoded_state,
|
||||
max_age=_OAUTH_STATE_COOKIE_TTL_SECONDS,
|
||||
path=path,
|
||||
secure=secure,
|
||||
httponly=True,
|
||||
samesite="lax",
|
||||
)
|
||||
|
||||
|
||||
def _resolve_encoded_oauth_state(request: Request, state: str) -> str:
|
||||
"""Return the encrypted OAuth session for a ``/callback`` request.
|
||||
|
||||
New flows carry it in a per-flow cookie keyed by the short handle we
|
||||
forwarded upstream (the IdP echoes that handle back as ``state``). Flows
|
||||
started before this change - or in flight across a deploy - carry the
|
||||
encrypted blob directly in ``state``, so fall back to it when the cookie
|
||||
is absent.
|
||||
"""
|
||||
cookie_value = request.cookies.get(_oauth_state_cookie_name(state))
|
||||
return cookie_value if cookie_value else state
|
||||
|
||||
|
||||
def _clear_oauth_state_cookie(response: Response, request: Request, state: str) -> None:
|
||||
cookie_name = _oauth_state_cookie_name(state)
|
||||
if cookie_name not in request.cookies:
|
||||
return
|
||||
path, secure = _oauth_state_cookie_path_and_secure(request)
|
||||
response.delete_cookie(
|
||||
key=cookie_name,
|
||||
path=path,
|
||||
secure=secure,
|
||||
httponly=True,
|
||||
samesite="lax",
|
||||
)
|
||||
|
||||
|
||||
def _get_validated_client_redirect_uri(request: Request, state_data: Dict[str, Any]) -> str:
|
||||
"""Return a trusted (same-origin, loopback, or ops-allowlisted)
|
||||
client redirect URI from OAuth state.
|
||||
|
|
@ -462,11 +529,12 @@ async def authorize_with_server(
|
|||
code_challenge_method=code_challenge_method,
|
||||
client_redirect_uri=redirect_uri,
|
||||
)
|
||||
relay_state = secrets.token_urlsafe(_OAUTH_STATE_HANDLE_BYTES)
|
||||
|
||||
params = {
|
||||
"client_id": mcp_server.client_id if mcp_server.client_id else client_id,
|
||||
"redirect_uri": f"{request_base_url}/callback",
|
||||
"state": encoded_state,
|
||||
"state": relay_state,
|
||||
"response_type": response_type or "code",
|
||||
}
|
||||
if scope:
|
||||
|
|
@ -483,7 +551,9 @@ async def authorize_with_server(
|
|||
existing_params = dict(parse_qsl(parsed_auth_url.query))
|
||||
existing_params.update(params)
|
||||
final_url = urlunparse(parsed_auth_url._replace(query=urlencode(existing_params)))
|
||||
return RedirectResponse(final_url)
|
||||
response = RedirectResponse(final_url)
|
||||
_set_oauth_state_cookie(response, request, relay_state, encoded_state)
|
||||
return response
|
||||
|
||||
|
||||
async def exchange_token_with_server(
|
||||
|
|
@ -772,6 +842,7 @@ async def _persist_dcr_client_registration(
|
|||
data=UpdateMCPServerRequest(
|
||||
server_id=mcp_server.server_id,
|
||||
credentials=credentials,
|
||||
oauth2_flow="authorization_code",
|
||||
**({"token_url": mcp_server.token_url} if mcp_server.token_url else {}),
|
||||
),
|
||||
touched_by="mcp_oauth_dcr",
|
||||
|
|
@ -1016,17 +1087,19 @@ async def callback(
|
|||
error_description,
|
||||
)
|
||||
if state:
|
||||
encoded_state = _resolve_encoded_oauth_state(request, state)
|
||||
try:
|
||||
state_data = decode_state_hash(state)
|
||||
state_data = decode_state_hash(encoded_state)
|
||||
original_state = state_data.get("original_state")
|
||||
redirect_uri = _get_validated_client_redirect_uri(request, state_data)
|
||||
except HTTPException:
|
||||
# Untrusted/invalid client redirect_uri — surface inline rather
|
||||
# than blindly forwarding the error to an attacker-controlled URL.
|
||||
return _render_oauth_error_html(error, error_description)
|
||||
except Exception:
|
||||
# State could not be decrypted (expired key, tampered, etc.).
|
||||
return _render_oauth_error_html(error, error_description)
|
||||
# Untrusted/invalid client redirect_uri (HTTPException), or an
|
||||
# undecryptable state (expired key, tampered): surface the IdP
|
||||
# error inline rather than forwarding it to an attacker-controlled
|
||||
# URL, and drop the one-time cookie we can no longer consume.
|
||||
response = _render_oauth_error_html(error, error_description)
|
||||
_clear_oauth_state_cookie(response, request, state)
|
||||
return response
|
||||
|
||||
params: Dict[str, str] = {"error": error}
|
||||
if error_description:
|
||||
|
|
@ -1036,7 +1109,9 @@ async def callback(
|
|||
if original_state is not None:
|
||||
params["state"] = original_state
|
||||
complete_returned_url = _append_query_params(redirect_uri, params)
|
||||
return RedirectResponse(url=complete_returned_url, status_code=302)
|
||||
response = RedirectResponse(url=complete_returned_url, status_code=302)
|
||||
_clear_oauth_state_cookie(response, request, state)
|
||||
return response
|
||||
|
||||
# No state — nothing to round-trip to. Show the user the error.
|
||||
return _render_oauth_error_html(error, error_description)
|
||||
|
|
@ -1052,7 +1127,8 @@ async def callback(
|
|||
|
||||
# 3. Successful authorization response.
|
||||
try:
|
||||
state_data = decode_state_hash(state)
|
||||
encoded_state = _resolve_encoded_oauth_state(request, state)
|
||||
state_data = decode_state_hash(encoded_state)
|
||||
original_state = state_data["original_state"]
|
||||
|
||||
# Re-validate the client redirect URI at the sink. /authorize
|
||||
|
|
@ -1065,14 +1141,18 @@ async def callback(
|
|||
|
||||
params = {"code": code, "state": original_state}
|
||||
complete_returned_url = _append_query_params(redirect_uri, params)
|
||||
return RedirectResponse(url=complete_returned_url, status_code=302)
|
||||
response = RedirectResponse(url=complete_returned_url, status_code=302)
|
||||
_clear_oauth_state_cookie(response, request, state)
|
||||
return response
|
||||
|
||||
except HTTPException:
|
||||
# Re-raise so a non-loopback base_url surfaces as 400 instead of
|
||||
# a generic "authentication incomplete" redirect.
|
||||
raise
|
||||
except Exception:
|
||||
return HTMLResponse("<html><body>Authentication incomplete. You can close this window.</body></html>")
|
||||
response = HTMLResponse("<html><body>Authentication incomplete. You can close this window.</body></html>")
|
||||
_clear_oauth_state_cookie(response, request, state)
|
||||
return response
|
||||
|
||||
|
||||
# ------------------------------
|
||||
|
|
|
|||
|
|
@ -253,6 +253,76 @@ def _without_authorization(
|
|||
return filtered or None
|
||||
|
||||
|
||||
def _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str:
|
||||
"""Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection."""
|
||||
if mcp_server.auth_type == MCPAuth.api_key:
|
||||
return f"ApiKey {mcp_auth_header}"
|
||||
if mcp_server.auth_type == MCPAuth.basic:
|
||||
return f"Basic {mcp_auth_header}"
|
||||
return f"Bearer {mcp_auth_header}"
|
||||
|
||||
|
||||
def _openapi_forwarded_extra_headers(
|
||||
mcp_server: MCPServer,
|
||||
raw_headers: Optional[dict[str, str]],
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
) -> Optional[dict[str, str]]:
|
||||
if not mcp_server.extra_headers or not raw_headers:
|
||||
return None
|
||||
normalized_raw = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)}
|
||||
skip_caller_authorization = _should_strip_caller_authorization(
|
||||
mcp_server=mcp_server,
|
||||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
forwarded: dict[str, str] = {}
|
||||
for header_name in mcp_server.extra_headers:
|
||||
if not isinstance(header_name, str):
|
||||
continue
|
||||
if skip_caller_authorization and header_name.lower() == "authorization":
|
||||
continue
|
||||
value = normalized_raw.get(header_name.lower())
|
||||
if value is not None:
|
||||
forwarded[header_name] = value
|
||||
return forwarded or None
|
||||
|
||||
|
||||
async def _resolve_byok_mcp_auth_header(
|
||||
mcp_server: MCPServer,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
mcp_auth_header: Optional[str],
|
||||
) -> Optional[str]:
|
||||
"""Resolve BYOK credential for tool calls that bypass ``execute_mcp_tool``."""
|
||||
if not mcp_server.is_byok:
|
||||
return mcp_auth_header
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_check_byok_credential,
|
||||
_get_byok_credential,
|
||||
)
|
||||
|
||||
if not mcp_auth_header:
|
||||
byok_cred = await _get_byok_credential(mcp_server, user_api_key_auth)
|
||||
if byok_cred is None:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail={
|
||||
"error": "byok_auth_required",
|
||||
"server_id": mcp_server.server_id,
|
||||
"server_name": mcp_server.server_name or mcp_server.name,
|
||||
"message": (
|
||||
"No stored credential found for this BYOK server. "
|
||||
"Complete the OAuth authorization flow to provide your API key."
|
||||
),
|
||||
},
|
||||
headers={"WWW-Authenticate": 'Bearer resource_metadata="/.well-known/oauth-protected-resource"'},
|
||||
)
|
||||
return byok_cred
|
||||
|
||||
await _check_byok_credential(mcp_server, user_api_key_auth)
|
||||
return mcp_auth_header
|
||||
|
||||
|
||||
def _extract_upstream_auth_failure(
|
||||
exc: BaseException,
|
||||
) -> Optional[tuple[int, Optional[str]]]:
|
||||
|
|
@ -511,6 +581,21 @@ def _create_elicitation_callback():
|
|||
class MCPServerManager:
|
||||
_STDIO_ENV_TEMPLATE_PATTERN = re.compile(r"^\$\{(X-[^}]+)\}$")
|
||||
|
||||
@staticmethod
|
||||
def _explicit_oauth2_flow(
|
||||
oauth2_flow: Optional[str],
|
||||
) -> Optional[Literal["client_credentials", "authorization_code"]]:
|
||||
"""DB rows persist their flow (write-time stamps plus the startup backfill) and
|
||||
config servers must declare it (validated at load), so both builds read the
|
||||
value verbatim: unknown or null resolves to None, which
|
||||
``needs_user_oauth_token`` already treats as interactive. Field-shape inference
|
||||
survives only in the request-time security helpers (``effective_oauth2_flow`` /
|
||||
``resolve_oauth2_flow_for_request``).
|
||||
"""
|
||||
if oauth2_flow in ("client_credentials", "authorization_code"):
|
||||
return cast(Literal["client_credentials", "authorization_code"], oauth2_flow)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _resolve_oauth2_flow(
|
||||
*,
|
||||
|
|
@ -521,11 +606,18 @@ class MCPServerManager:
|
|||
client_id: Optional[str],
|
||||
client_secret: Optional[str],
|
||||
) -> Optional[Literal["client_credentials", "authorization_code"]]:
|
||||
"""Infer oauth2_flow for legacy records that omit the field.
|
||||
"""Infer oauth2_flow from field shape when the value is omitted.
|
||||
|
||||
DB rows created before oauth2_flow support may have OAuth2 client
|
||||
credentials + token_url but a null oauth2_flow. Treat these as M2M,
|
||||
unless authorization_url is present (interactive OAuth).
|
||||
SECURITY-SENSITIVE: this is the shape-inference engine both request-time security
|
||||
helpers delegate to, so it is what decides M2M-vs-interactive for an unstamped row.
|
||||
Always access it through ``effective_oauth2_flow`` (boolean/enum decisions) or
|
||||
``resolve_oauth2_flow_for_request`` (the egress object backstop), which are the single
|
||||
choke points for request-time resolution; do not call it directly from security sites
|
||||
and do not weaken its M2M-shape branch without accounting for those callers. DB rows
|
||||
are stamped at write time and by the startup backfill, config servers must declare
|
||||
oauth2_flow (validated at load), and both builds read the value verbatim via
|
||||
``_explicit_oauth2_flow``. Delete this whole request-time layer only once the backstop
|
||||
warning stays silent in production.
|
||||
"""
|
||||
if oauth2_flow in ("client_credentials", "authorization_code"):
|
||||
return cast(Literal["client_credentials", "authorization_code"], oauth2_flow)
|
||||
|
|
@ -540,6 +632,51 @@ class MCPServerManager:
|
|||
return "client_credentials"
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def effective_oauth2_flow(server: "MCPServer") -> Optional[Literal["client_credentials", "authorization_code"]]:
|
||||
"""The oauth2_flow a security decision must use for ``server`` this request.
|
||||
|
||||
Column-first, shape-fallback: a stamped row returns its explicit value; an
|
||||
unstamped (null) row whose fields carry the M2M shape resolves to
|
||||
``client_credentials`` so it is treated as M2M and fails closed. Every
|
||||
security-sensitive reader (anonymous-delegate allowlist and gate, egress flow
|
||||
resolution) goes through this one helper rather than reading the bare
|
||||
``has_client_credentials`` column, which is unreliable for null rows.
|
||||
"""
|
||||
return MCPServerManager._resolve_oauth2_flow(
|
||||
auth_type=server.auth_type,
|
||||
oauth2_flow=server.oauth2_flow,
|
||||
token_url=server.token_url,
|
||||
authorization_url=server.authorization_url,
|
||||
client_id=server.client_id,
|
||||
client_secret=server.client_secret,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def resolve_oauth2_flow_for_request(server: "MCPServer") -> "MCPServer":
|
||||
"""Return ``server`` with its effective oauth2_flow applied, for egress paths.
|
||||
|
||||
A stamped row is returned unchanged (its effective flow equals the stored value).
|
||||
An unstamped M2M-shape row is returned as a per-request copy carrying
|
||||
``oauth2_flow=client_credentials`` so downstream ``has_client_credentials`` /
|
||||
``needs_user_oauth_token`` compute correctly and the stored client credentials are
|
||||
used instead of forwarding the caller's Authorization. Use this at every point that
|
||||
resolves an allowed server id into an ``MCPServer`` for a tool call or listing.
|
||||
"""
|
||||
effective = MCPServerManager.effective_oauth2_flow(server)
|
||||
if effective is None or effective == server.oauth2_flow:
|
||||
return server
|
||||
verbose_logger.warning(
|
||||
"MCP server %s has no persisted oauth2_flow but matches the %s shape; using the "
|
||||
"inferred flow for this request. The startup backfill leaves this ambiguous M2M "
|
||||
"shape unstamped on purpose, so it will NOT self-heal: set oauth2_flow explicitly "
|
||||
"in the dashboard or via PUT /v1/mcp/server (client_credentials for M2M, or "
|
||||
"authorization_code after an interactive sign-in).",
|
||||
server.server_id,
|
||||
effective,
|
||||
)
|
||||
return server.model_copy(update={"oauth2_flow": effective})
|
||||
|
||||
@staticmethod
|
||||
def _obo_needs_endpoint_discovery(
|
||||
auth_type: Optional[MCPAuthType],
|
||||
|
|
@ -772,6 +909,20 @@ class MCPServerManager:
|
|||
mcp_oauth_metadata.registration_url if mcp_oauth_metadata else None
|
||||
)
|
||||
|
||||
config_oauth2_flow = server_config.get("oauth2_flow", None)
|
||||
if auth_type == MCPAuth.oauth2 and config_oauth2_flow not in (
|
||||
"client_credentials",
|
||||
"authorization_code",
|
||||
):
|
||||
raise ValueError(
|
||||
f"Invalid config for MCP server '{server_name or server_id}': auth_type oauth2 "
|
||||
f"requires an explicit oauth2_flow (got {config_oauth2_flow!r}). Set "
|
||||
"oauth2_flow: client_credentials for machine-to-machine servers (the proxy mints "
|
||||
"a shared token at token_url using client_id/client_secret, no user interaction) "
|
||||
"or oauth2_flow: authorization_code for interactive servers (per-user tokens via "
|
||||
"browser sign-in, including delegate_auth_to_upstream)."
|
||||
)
|
||||
|
||||
new_server = MCPServer(
|
||||
server_id=server_id,
|
||||
name=name_for_prefix,
|
||||
|
|
@ -785,14 +936,7 @@ class MCPServerManager:
|
|||
# oauth specific fields
|
||||
client_id=server_config.get("client_id", None),
|
||||
client_secret=server_config.get("client_secret", None),
|
||||
oauth2_flow=self._resolve_oauth2_flow(
|
||||
auth_type=auth_type,
|
||||
oauth2_flow=server_config.get("oauth2_flow", None),
|
||||
token_url=resolved_token_url,
|
||||
authorization_url=resolved_authorization_url,
|
||||
client_id=server_config.get("client_id", None),
|
||||
client_secret=server_config.get("client_secret", None),
|
||||
),
|
||||
oauth2_flow=self._explicit_oauth2_flow(config_oauth2_flow),
|
||||
scopes=resolved_scopes,
|
||||
authorization_url=resolved_authorization_url,
|
||||
token_url=resolved_token_url,
|
||||
|
|
@ -1170,15 +1314,7 @@ class MCPServerManager:
|
|||
env_vars=env_vars_list,
|
||||
client_id=client_id_value or getattr(mcp_server, "client_id", None),
|
||||
client_secret=client_secret_value or getattr(mcp_server, "client_secret", None),
|
||||
oauth2_flow=self._resolve_oauth2_flow(
|
||||
auth_type=auth_type,
|
||||
oauth2_flow=getattr(mcp_server, "oauth2_flow", None),
|
||||
token_url=mcp_server.token_url or getattr(mcp_oauth_metadata, "token_url", None),
|
||||
authorization_url=mcp_server.authorization_url
|
||||
or getattr(mcp_oauth_metadata, "authorization_url", None),
|
||||
client_id=client_id_value or getattr(mcp_server, "client_id", None),
|
||||
client_secret=client_secret_value or getattr(mcp_server, "client_secret", None),
|
||||
),
|
||||
oauth2_flow=self._explicit_oauth2_flow(getattr(mcp_server, "oauth2_flow", None)),
|
||||
scopes=resolved_scopes,
|
||||
authorization_url=mcp_server.authorization_url or getattr(mcp_oauth_metadata, "authorization_url", None),
|
||||
token_url=mcp_server.token_url or getattr(mcp_oauth_metadata, "token_url", None),
|
||||
|
|
@ -1486,8 +1622,11 @@ class MCPServerManager:
|
|||
and getattr(server, "delegate_auth_to_upstream", False) is True
|
||||
# M2M servers must not be exposed anonymously: an
|
||||
# unauthenticated caller would get LiteLLM to proxy tool
|
||||
# calls using its stored client_credentials.
|
||||
and not server.has_client_credentials
|
||||
# calls using its stored client_credentials. Resolve the flow
|
||||
# rather than reading has_client_credentials so an unstamped
|
||||
# M2M-shape row (null column, verbatim-read as non-M2M) still
|
||||
# fails closed here, matching the anonymous-delegate auth gate.
|
||||
and MCPServerManager.effective_oauth2_flow(server) != "client_credentials"
|
||||
]
|
||||
combined_servers.update(delegate_server_ids)
|
||||
|
||||
|
|
@ -3861,6 +4000,15 @@ class MCPServerManager:
|
|||
start_time = datetime.datetime.now()
|
||||
mcp_server = self._resolve_mcp_server_for_tool_call(server_name, name)
|
||||
|
||||
# Resolved before any hook runs so a missing BYOK credential (401) never
|
||||
# leaves during-hook side effects (audit logging, rate-limit bookkeeping)
|
||||
# recorded against a call that ultimately fails.
|
||||
mcp_auth_header = await _resolve_byok_mcp_auth_header(
|
||||
mcp_server,
|
||||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# Pre MCP Tool Call Hook
|
||||
# Allow validation and modification of tool calls before execution
|
||||
|
|
@ -3907,9 +4055,25 @@ class MCPServerManager:
|
|||
server_name,
|
||||
)
|
||||
|
||||
auth_header_value = (
|
||||
_format_byok_openapi_auth_header(mcp_server, mcp_auth_header) if mcp_auth_header else None
|
||||
)
|
||||
forwarded_headers = _openapi_forwarded_extra_headers(mcp_server, raw_headers, user_api_key_auth)
|
||||
|
||||
async def _call_openapi_via_handler():
|
||||
async with self._limit_outbound_concurrency(mcp_server):
|
||||
return await self._call_openapi_tool_handler(mcp_server, name, arguments)
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
_request_auth_header,
|
||||
_request_extra_headers,
|
||||
)
|
||||
|
||||
auth_token = _request_auth_header.set(auth_header_value)
|
||||
extra_token = _request_extra_headers.set(forwarded_headers)
|
||||
try:
|
||||
async with self._limit_outbound_concurrency(mcp_server):
|
||||
return await self._call_openapi_tool_handler(mcp_server, name, arguments)
|
||||
finally:
|
||||
_request_auth_header.reset(auth_token)
|
||||
_request_extra_headers.reset(extra_token)
|
||||
|
||||
tasks.append(asyncio.create_task(_call_openapi_via_handler()))
|
||||
else:
|
||||
|
|
@ -4465,6 +4629,7 @@ class MCPServerManager:
|
|||
authorization_url=server.authorization_url,
|
||||
token_url=server.token_url,
|
||||
registration_url=server.registration_url,
|
||||
oauth2_flow=server.oauth2_flow,
|
||||
allow_all_keys=server.allow_all_keys,
|
||||
instructions=server.instructions,
|
||||
timeout=server.timeout,
|
||||
|
|
@ -4553,6 +4718,8 @@ class MCPServerManager:
|
|||
teams=[],
|
||||
mcp_access_groups=server.access_groups or [],
|
||||
allowed_tools=server.allowed_tools or [],
|
||||
tool_name_to_display_name=server.tool_name_to_display_name,
|
||||
tool_name_to_description=server.tool_name_to_description,
|
||||
extra_headers=server.extra_headers or [],
|
||||
mcp_info=server.mcp_info,
|
||||
static_headers=server.static_headers,
|
||||
|
|
@ -4566,6 +4733,7 @@ class MCPServerManager:
|
|||
authorization_url=server.authorization_url,
|
||||
token_url=server.token_url,
|
||||
registration_url=server.registration_url,
|
||||
oauth2_flow=server.oauth2_flow,
|
||||
allow_all_keys=server.allow_all_keys,
|
||||
available_on_public_internet=server.available_on_public_internet,
|
||||
delegate_auth_to_upstream=server.delegate_auth_to_upstream,
|
||||
|
|
|
|||
155
litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py
Normal file
155
litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py
Normal file
|
|
@ -0,0 +1,155 @@
|
|||
"""Startup backfill for oauth2 MCP server rows persisted before oauth2_flow was written.
|
||||
|
||||
Rows created before the write-side stamps (DCR persist, UI create, REST create) carry a
|
||||
null ``oauth2_flow`` and rely on read-time field-shape inference, which cannot tell a
|
||||
DCR-registered interactive server from an M2M server unless endpoint discovery succeeds
|
||||
first. This backfill classifies each null row once, at rest, using signals inference
|
||||
never had, and persists the result so the read path never has to infer again.
|
||||
|
||||
Signal order, strongest first:
|
||||
|
||||
1. Per-user OAuth token rows exist for the server: only the interactive flow mints
|
||||
per-user tokens, so this is definitive and immune to the discovery trap. BYOK API
|
||||
keys share the same table (``LiteLLM_MCPUserCredentials``), so only rows whose
|
||||
payload decodes as a ``type: oauth2`` token count as proof; bare keys and
|
||||
undecodable rows prove nothing about the flow.
|
||||
2. ``authorization_url`` persisted: interactive needs a user-facing authorization
|
||||
endpoint; M2M (RFC 6749 section 4.4) never has one.
|
||||
3. ``registration_url`` persisted: dynamic client registration (RFC 7591) exists to mint
|
||||
clients for the interactive flow; M2M servers are configured with static credentials.
|
||||
4. ``token_url`` plus decryptable ``client_id`` and ``client_secret``: ambiguous, left
|
||||
unstamped. The shape is shared by M2M servers and DCR-registered interactive servers
|
||||
whose authorization endpoint lives only in discovery (registered but never signed
|
||||
in), so stamping client_credentials here could permanently route per-user traffic
|
||||
through the proxy's stored client credential. The row keeps working through the
|
||||
request-time backstop and a warning names it with the one-line fix (set oauth2_flow
|
||||
via the dashboard or ``PUT /v1/mcp/server``); a completed interactive sign-in also
|
||||
heals it via rule 1 at the next boot.
|
||||
5. Anything else is interactive: matching how ``needs_user_oauth_token`` treats a null
|
||||
flow, so the stamp never changes runtime routing for rows no rule recognizes.
|
||||
|
||||
The backfill never stamps client_credentials: M2M is asserted by a human (config
|
||||
requires it, the API accepts it, the dashboard sets it), mirroring the config-level
|
||||
validation error. Runs before the first registry load on every boot and is idempotent:
|
||||
a healed fleet has no null rows and the backfill exits after one query.
|
||||
"""
|
||||
|
||||
import json
|
||||
from collections import Counter
|
||||
from typing import Any, Literal, Optional
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._experimental.mcp_server.db import _decode_oauth_payload, decrypt_credentials
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.types.mcp import MCPCredentials
|
||||
|
||||
OAuth2Flow = Literal["client_credentials", "authorization_code"]
|
||||
BackfillRule = Literal[
|
||||
"per_user_tokens",
|
||||
"authorization_url",
|
||||
"registration_url",
|
||||
"ambiguous_m2m_shape",
|
||||
"interactive_default",
|
||||
]
|
||||
|
||||
_BACKFILL_AUDIT_ACTOR = "oauth2_flow_backfill"
|
||||
|
||||
|
||||
def _decrypted_credentials(raw_credentials: Any) -> Optional[MCPCredentials]:
|
||||
if raw_credentials is None:
|
||||
return None
|
||||
if isinstance(raw_credentials, str):
|
||||
try:
|
||||
parsed = json.loads(raw_credentials)
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
else:
|
||||
parsed = raw_credentials
|
||||
if not isinstance(parsed, dict):
|
||||
return None
|
||||
return decrypt_credentials(credentials=dict(parsed))
|
||||
|
||||
|
||||
def classify_null_flow_row(
|
||||
*,
|
||||
has_per_user_tokens: bool,
|
||||
authorization_url: Optional[str],
|
||||
registration_url: Optional[str],
|
||||
token_url: Optional[str],
|
||||
credentials: Optional[MCPCredentials],
|
||||
) -> tuple[Optional[OAuth2Flow], BackfillRule]:
|
||||
if has_per_user_tokens:
|
||||
return "authorization_code", "per_user_tokens"
|
||||
if authorization_url:
|
||||
return "authorization_code", "authorization_url"
|
||||
if registration_url:
|
||||
return "authorization_code", "registration_url"
|
||||
if token_url and credentials and credentials.get("client_id") and credentials.get("client_secret"):
|
||||
return None, "ambiguous_m2m_shape"
|
||||
return "authorization_code", "interactive_default"
|
||||
|
||||
|
||||
async def backfill_null_oauth2_flows(prisma_client: PrismaClient) -> dict[BackfillRule, int]:
|
||||
"""Classify every ``auth_type=oauth2`` row whose ``oauth2_flow`` is null; stamp the provable
|
||||
ones, warn on the ambiguous ones, and return counts per rule."""
|
||||
null_rows: list[Any] = await prisma_client.db.litellm_mcpservertable.find_many(
|
||||
where={"auth_type": "oauth2", "oauth2_flow": None},
|
||||
)
|
||||
if not null_rows:
|
||||
return {}
|
||||
|
||||
server_ids = [row.server_id for row in null_rows]
|
||||
token_rows: list[Any] = await prisma_client.db.litellm_mcpusercredentials.find_many(
|
||||
where={"server_id": {"in": server_ids}},
|
||||
)
|
||||
server_ids_with_oauth_tokens: set[str] = {
|
||||
token_row.server_id for token_row in token_rows if _decode_oauth_payload(token_row.credential_b64) is not None
|
||||
}
|
||||
|
||||
classified = tuple(
|
||||
(
|
||||
row,
|
||||
classify_null_flow_row(
|
||||
has_per_user_tokens=row.server_id in server_ids_with_oauth_tokens,
|
||||
authorization_url=row.authorization_url,
|
||||
registration_url=row.registration_url,
|
||||
token_url=row.token_url,
|
||||
credentials=_decrypted_credentials(row.credentials),
|
||||
),
|
||||
)
|
||||
for row in null_rows
|
||||
)
|
||||
|
||||
for row, (flow, rule) in classified:
|
||||
if flow is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"oauth2_flow backfill: server_id=%s is ambiguous (client credentials + token_url, "
|
||||
"no interactive signal); left unstamped. Set oauth2_flow explicitly via the "
|
||||
"dashboard or PUT /v1/mcp/server: client_credentials if this server is M2M, or "
|
||||
"complete an interactive sign-in and it will be stamped authorization_code at the "
|
||||
"next boot.",
|
||||
row.server_id,
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.info(
|
||||
"oauth2_flow backfill: server_id=%s stamped %s (rule=%s)",
|
||||
row.server_id,
|
||||
flow,
|
||||
rule,
|
||||
)
|
||||
|
||||
stamped_flows = {flow for _, (flow, _) in classified if flow is not None}
|
||||
for stamped_flow in stamped_flows:
|
||||
server_ids_for_flow = [row.server_id for row, (row_flow, _) in classified if row_flow == stamped_flow]
|
||||
await prisma_client.db.litellm_mcpservertable.update_many(
|
||||
where={"server_id": {"in": server_ids_for_flow}, "oauth2_flow": None},
|
||||
data={"oauth2_flow": stamped_flow, "updated_by": _BACKFILL_AUDIT_ACTOR},
|
||||
)
|
||||
|
||||
counts: dict[BackfillRule, int] = dict(Counter(rule for _, (_, rule) in classified))
|
||||
verbose_proxy_logger.info(
|
||||
"oauth2_flow backfill: processed %d oauth2 server row(s): %s",
|
||||
len(null_rows),
|
||||
counts,
|
||||
)
|
||||
return counts
|
||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue